diff --git a/.gitattributes b/.gitattributes index 3f841bf7cad747229f3aa020197e91cd5d5c021d..b15fc291b72a092f96ec5c716a0da6b4edaedaa6 100644 --- a/.gitattributes +++ b/.gitattributes @@ -1,3 +1,6 @@ codex-rs/app-server-protocol/schema/** linguist-generated codex-rs/hooks/schema/generated/** linguist-generated third_party/voice/sources.json text eol=lf +.github/codex-cli-splash.png filter=lfs diff=lfs merge=lfs -text +codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst filter=lfs diff=lfs merge=lfs -text +codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst filter=lfs diff=lfs merge=lfs -text diff --git a/.github/codex-cli-splash.png b/.github/codex-cli-splash.png new file mode 100644 index 0000000000000000000000000000000000000000..bf71356354ebd9762c04113ac13f9e0ed5cbf58e --- /dev/null +++ b/.github/codex-cli-splash.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:15b86fa9c0790a779ecdd84f1e7dee029ab79bfb093bd3c4876998696925b013 +size 838131 diff --git a/codex-cli/bin/codex.js b/codex-cli/bin/codex.js new file mode 100644 index 0000000000000000000000000000000000000000..934600714b97f340cb5468e20e5fe560d751cbed --- /dev/null +++ b/codex-cli/bin/codex.js @@ -0,0 +1,295 @@ +#!/usr/bin/env node +// Unified entry point for the Codex CLI. + +import { spawn } from "node:child_process"; +import { existsSync, readFileSync, realpathSync } from "fs"; +import { createRequire } from "node:module"; +import path from "path"; +import { fileURLToPath } from "url"; + +// __dirname equivalent in ESM +const __filename = fileURLToPath(import.meta.url); +const __dirname = path.dirname(__filename); +const require = createRequire(import.meta.url); +const codexPackageRoot = realpathSync(path.join(__dirname, "..")); + +const PLATFORM_PACKAGE_BY_TARGET = { + "x86_64-unknown-linux-musl": "@openai/codex-linux-x64", + "aarch64-unknown-linux-musl": "@openai/codex-linux-arm64", + "x86_64-apple-darwin": "@openai/codex-darwin-x64", + "aarch64-apple-darwin": "@openai/codex-darwin-arm64", + "x86_64-pc-windows-msvc": "@openai/codex-win32-x64", + "aarch64-pc-windows-msvc": "@openai/codex-win32-arm64", +}; + +const { platform, arch } = process; + +let targetTriple = null; +switch (platform) { + case "linux": + case "android": + switch (arch) { + case "x64": + targetTriple = "x86_64-unknown-linux-musl"; + break; + case "arm64": + targetTriple = "aarch64-unknown-linux-musl"; + break; + default: + break; + } + break; + case "darwin": + switch (arch) { + case "x64": + targetTriple = "x86_64-apple-darwin"; + break; + case "arm64": + targetTriple = "aarch64-apple-darwin"; + break; + default: + break; + } + break; + case "win32": + switch (arch) { + case "x64": + targetTriple = "x86_64-pc-windows-msvc"; + break; + case "arm64": + targetTriple = "aarch64-pc-windows-msvc"; + break; + default: + break; + } + break; + default: + break; +} + +if (!targetTriple) { + throw new Error(`Unsupported platform: ${platform} (${arch})`); +} + +const platformPackage = PLATFORM_PACKAGE_BY_TARGET[targetTriple]; +if (!platformPackage) { + throw new Error(`Unsupported target triple: ${targetTriple}`); +} + +function findCodexExecutable() { + let vendorRoot; + try { + const packageJsonPath = require.resolve(`${platformPackage}/package.json`); + vendorRoot = path.join(path.dirname(packageJsonPath), "vendor"); + } catch { + vendorRoot = path.join(__dirname, "..", "vendor"); + } + + const codexExecutable = path.join( + vendorRoot, + targetTriple, + "bin", + process.platform === "win32" ? "codex.exe" : "codex", + ); + if (existsSync(codexExecutable)) { + return codexExecutable; + } + + const packageManager = detectPackageManager(); + const updateCommand = + packageManager === "bun" + ? "bun install -g @openai/codex@latest" + : packageManager === "pnpm" + ? "pnpm add -g @openai/codex@latest" + : packageManager === "vite-plus" + ? "vp install -g @openai/codex@latest" + : "npm install -g @openai/codex@latest"; + throw new Error( + `Missing optional dependency ${platformPackage}. Reinstall Codex: ${updateCommand}`, + ); +} + +const binaryPath = findCodexExecutable(); + +// Use an asynchronous spawn instead of spawnSync so that Node is able to +// respond to signals (e.g. Ctrl-C / SIGINT) while the native binary is +// executing. This allows us to forward those signals to the child process +// and guarantees that when either the child terminates or the parent +// receives a fatal signal, both processes exit in a predictable manner. + +function isPnpmOwnedCodexInstall(nodeModulesDir) { + if (!existsSync(path.join(nodeModulesDir, ".modules.yaml"))) { + return false; + } + + try { + return ( + realpathSync(path.join(nodeModulesDir, "@openai", "codex")) === + codexPackageRoot + ); + } catch { + return false; + } +} + +function isVitePlusOwnedCodexInstall(packagesDir) { + if (path.basename(packagesDir) !== "packages") { + return false; + } + + try { + const metadata = JSON.parse( + readFileSync(path.join(packagesDir, "@openai", "codex.json"), "utf8"), + ); + if (metadata.name !== "@openai/codex") { + return false; + } + + // Vite+ records the active global installation in packages/@openai/codex.json. + // Older installs have no ID or append a #-prefixed ID to the package name; + // newer installs put the ID in a subdirectory of the package prefix. + const installId = metadata.installId || ""; + const installDir = installId.startsWith("#") + ? path.join(packagesDir, `@openai/codex${installId}`) + : path.join(packagesDir, "@openai/codex", installId); + for (const nodeModulesDir of [ + path.join(installDir, "lib", "node_modules"), + path.join(installDir, "node_modules"), + ]) { + const packageRoot = path.join(nodeModulesDir, "@openai", "codex"); + if ( + existsSync(packageRoot) && + realpathSync(packageRoot) === codexPackageRoot + ) { + return true; + } + } + } catch { + // Missing or unreadable ownership metadata must not prevent Codex starting. + } + return false; +} + +/** + * Use heuristics to detect the package manager that was used to install Codex + * in order to give the user a hint about how to update it. + */ +function detectPackageManager() { + // Package-manager ownership metadata can be several parents above the package. + // Search ancestors of both the canonical package root and lexical entrypoint + // because the package manager may link either path. + const entrypointDir = path.dirname(path.resolve(process.argv[1])); + for (const startDir of new Set([codexPackageRoot, entrypointDir])) { + const filesystemRoot = path.parse(startDir).root; + for ( + let currentDir = startDir; + currentDir !== filesystemRoot; + currentDir = path.dirname(currentDir) + ) { + if (isVitePlusOwnedCodexInstall(currentDir)) { + return "vite-plus"; + } + if (isPnpmOwnedCodexInstall(path.join(currentDir, "node_modules"))) { + return "pnpm"; + } + } + + if (isPnpmOwnedCodexInstall(path.join(filesystemRoot, "node_modules"))) { + return "pnpm"; + } + } + + const userAgent = process.env.npm_config_user_agent || ""; + if (/\bbun\//.test(userAgent)) { + return "bun"; + } + + const execPath = process.env.npm_execpath || ""; + if (execPath.includes("bun")) { + return "bun"; + } + + if ( + __dirname.includes(".bun/install/global") || + __dirname.includes(".bun\\install\\global") + ) { + return "bun"; + } + + return userAgent ? "npm" : null; +} + +const packageManager = detectPackageManager(); +const packageManagerEnvVar = + packageManager === "bun" + ? "CODEX_MANAGED_BY_BUN" + : packageManager === "pnpm" + ? "CODEX_MANAGED_BY_PNPM" + : packageManager === "vite-plus" + ? "CODEX_MANAGED_BY_VITE_PLUS" + : "CODEX_MANAGED_BY_NPM"; +const env = { + ...process.env, + CODEX_MANAGED_PACKAGE_ROOT: codexPackageRoot, +}; +delete env.CODEX_MANAGED_BY_NPM; +delete env.CODEX_MANAGED_BY_BUN; +delete env.CODEX_MANAGED_BY_PNPM; +delete env.CODEX_MANAGED_BY_VITE_PLUS; +env[packageManagerEnvVar] = "1"; + +const child = spawn(binaryPath, process.argv.slice(2), { + stdio: "inherit", + env, +}); + +child.on("error", (err) => { + // Typically triggered when the binary is missing or not executable. + // Re-throwing here will terminate the parent with a non-zero exit code + // while still printing a helpful stack trace. + // eslint-disable-next-line no-console + console.error(err); + process.exit(1); +}); + +// Forward common termination signals to the child so that it shuts down +// gracefully. In the handler we temporarily disable the default behavior of +// exiting immediately; once the child has been signaled we simply wait for +// its exit event which will in turn terminate the parent (see below). +const forwardSignal = (signal) => { + if (child.killed) { + return; + } + try { + child.kill(signal); + } catch { + /* ignore */ + } +}; + +["SIGINT", "SIGTERM", "SIGHUP"].forEach((sig) => { + process.on(sig, () => forwardSignal(sig)); +}); + +// When the child exits, mirror its termination reason in the parent so that +// shell scripts and other tooling observe the correct exit status. +// Wrap the lifetime of the child process in a Promise so that we can await +// its termination in a structured way. The Promise resolves with an object +// describing how the child exited: either via exit code or due to a signal. +const childResult = await new Promise((resolve) => { + child.on("exit", (code, signal) => { + if (signal) { + resolve({ type: "signal", signal }); + } else { + resolve({ type: "code", exitCode: code ?? 1 }); + } + }); +}); + +if (childResult.type === "signal") { + // Re-emit the same signal so that the parent terminates with the expected + // semantics (this also sets the correct exit code of 128 + n). + process.kill(process.pid, childResult.signal); +} else { + process.exit(childResult.exitCode); +} diff --git a/codex-cli/scripts/README.md b/codex-cli/scripts/README.md new file mode 100644 index 0000000000000000000000000000000000000000..4877781c36b9d39b5fb6d469aa0ff72f611cef3e --- /dev/null +++ b/codex-cli/scripts/README.md @@ -0,0 +1,23 @@ +# npm releases + +Use the staging helper in the repo root to generate npm tarballs for a release. For +example, to stage the CLI, responses proxy, and SDK packages for version `0.6.0`: + +```bash +./scripts/stage_npm_packages.py \ + --release-version 0.6.0 \ + --package codex \ + --package codex-responses-api-proxy \ + --package codex-sdk +``` + +This downloads the required native package archive artifacts, hydrates `vendor/` for +each package, and writes tarballs to `dist/npm/`. + +When `--package codex` is provided, the staging helper builds the lightweight +`@openai/codex` meta package plus all platform-native `@openai/codex` variants +that are later published under platform-specific dist-tags. + +Direct `build_npm_package.py` invocations are still useful for package-specific +debugging, but native packages expect `--vendor-src` to point at a prehydrated +`vendor/` tree. Release packaging should use `scripts/stage_npm_packages.py` diff --git a/codex-cli/scripts/build_npm_package.py b/codex-cli/scripts/build_npm_package.py new file mode 100644 index 0000000000000000000000000000000000000000..12f4f9dc8993108f5315bb8632f99c6793539124 --- /dev/null +++ b/codex-cli/scripts/build_npm_package.py @@ -0,0 +1,461 @@ +#!/usr/bin/env python3 +"""Stage and optionally package the @openai/codex npm module.""" + +import argparse +import json +import os +import shutil +import subprocess +import sys +import tempfile +from pathlib import Path + +SCRIPT_DIR = Path(__file__).resolve().parent +CODEX_CLI_ROOT = SCRIPT_DIR.parent +REPO_ROOT = CODEX_CLI_ROOT.parent +RESPONSES_API_PROXY_NPM_ROOT = REPO_ROOT / "codex-rs" / "responses-api-proxy" / "npm" +CODEX_SDK_ROOT = REPO_ROOT / "sdk" / "typescript" +CODEX_NPM_NAME = "@openai/codex" +CODEX_PACKAGE_COMPONENT = "codex-package" + +# `npm_name` is the local optional-dependency alias consumed by `bin/codex.js`. +# The underlying package published to npm is always `@openai/codex`. +CODEX_PLATFORM_PACKAGES: dict[str, dict[str, str]] = { + "codex-linux-x64": { + "npm_name": "@openai/codex-linux-x64", + "npm_tag": "linux-x64", + "target_triple": "x86_64-unknown-linux-musl", + "os": "linux", + "cpu": "x64", + }, + "codex-linux-arm64": { + "npm_name": "@openai/codex-linux-arm64", + "npm_tag": "linux-arm64", + "target_triple": "aarch64-unknown-linux-musl", + "os": "linux", + "cpu": "arm64", + }, + "codex-darwin-x64": { + "npm_name": "@openai/codex-darwin-x64", + "npm_tag": "darwin-x64", + "target_triple": "x86_64-apple-darwin", + "os": "darwin", + "cpu": "x64", + }, + "codex-darwin-arm64": { + "npm_name": "@openai/codex-darwin-arm64", + "npm_tag": "darwin-arm64", + "target_triple": "aarch64-apple-darwin", + "os": "darwin", + "cpu": "arm64", + }, + "codex-win32-x64": { + "npm_name": "@openai/codex-win32-x64", + "npm_tag": "win32-x64", + "target_triple": "x86_64-pc-windows-msvc", + "os": "win32", + "cpu": "x64", + }, + "codex-win32-arm64": { + "npm_name": "@openai/codex-win32-arm64", + "npm_tag": "win32-arm64", + "target_triple": "aarch64-pc-windows-msvc", + "os": "win32", + "cpu": "arm64", + }, +} + +PACKAGE_EXPANSIONS: dict[str, list[str]] = { + "codex": ["codex", *CODEX_PLATFORM_PACKAGES], +} + +PACKAGE_NATIVE_COMPONENTS: dict[str, list[str]] = { + "codex": [], + "codex-linux-x64": [CODEX_PACKAGE_COMPONENT], + "codex-linux-arm64": [CODEX_PACKAGE_COMPONENT], + "codex-darwin-x64": [CODEX_PACKAGE_COMPONENT], + "codex-darwin-arm64": [CODEX_PACKAGE_COMPONENT], + "codex-win32-x64": [CODEX_PACKAGE_COMPONENT], + "codex-win32-arm64": [CODEX_PACKAGE_COMPONENT], + "codex-responses-api-proxy": ["codex-responses-api-proxy"], + "codex-sdk": [], +} + +PACKAGE_TARGET_FILTERS: dict[str, str] = { + package_name: package_config["target_triple"] + for package_name, package_config in CODEX_PLATFORM_PACKAGES.items() +} + +PACKAGE_CHOICES = tuple(PACKAGE_NATIVE_COMPONENTS) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Build or stage the Codex CLI npm package." + ) + parser.add_argument( + "--package", + choices=PACKAGE_CHOICES, + default="codex", + help="Which npm package to stage (default: codex).", + ) + parser.add_argument( + "--version", + help="Version number to write to package.json inside the staged package.", + ) + parser.add_argument( + "--release-version", + help=("Version to stage for npm release."), + ) + parser.add_argument( + "--staging-dir", + type=Path, + help=( + "Directory to stage the package contents. Defaults to a new temporary directory " + "if omitted. The directory must be empty when provided." + ), + ) + parser.add_argument( + "--tmp", + dest="staging_dir", + type=Path, + help=argparse.SUPPRESS, + ) + parser.add_argument( + "--pack-output", + type=Path, + help="Path where the generated npm tarball should be written.", + ) + parser.add_argument( + "--vendor-src", + type=Path, + help="Directory containing pre-installed native binaries to bundle (vendor root).", + ) + return parser.parse_args() + + +def main() -> int: + args = parse_args() + + package = args.package + version = args.version + release_version = args.release_version + if release_version: + if version and version != release_version: + raise RuntimeError( + "--version and --release-version must match when both are provided." + ) + version = release_version + + if not version: + raise RuntimeError("Must specify --version or --release-version.") + + staging_dir, created_temp = prepare_staging_dir(args.staging_dir) + + try: + stage_sources(staging_dir, version, package) + + vendor_src = args.vendor_src.resolve() if args.vendor_src else None + native_components = PACKAGE_NATIVE_COMPONENTS.get(package, []) + target_filter = PACKAGE_TARGET_FILTERS.get(package) + + if native_components: + if vendor_src is None: + components_str = ", ".join(native_components) + raise RuntimeError( + "Native components " + f"({components_str}) required for package '{package}'. Provide --vendor-src " + "pointing to a directory containing pre-installed binaries." + ) + + copy_native_binaries( + vendor_src, + staging_dir, + native_components, + target_filter={target_filter} if target_filter else None, + ) + + if release_version: + staging_dir_str = str(staging_dir) + if package == "codex": + print( + f"Staged version {version} for release in {staging_dir_str}\n\n" + "Verify the CLI:\n" + f" node {staging_dir_str}/bin/codex.js --version\n" + f" node {staging_dir_str}/bin/codex.js --help\n\n" + ) + elif package == "codex-responses-api-proxy": + print( + f"Staged version {version} for release in {staging_dir_str}\n\n" + "Verify the responses API proxy:\n" + f" node {staging_dir_str}/bin/codex-responses-api-proxy.js --help\n\n" + ) + elif package in CODEX_PLATFORM_PACKAGES: + print( + f"Staged version {version} for release in {staging_dir_str}\n\n" + "Verify native payload contents:\n" + f" ls {staging_dir_str}/vendor\n\n" + ) + else: + print( + f"Staged version {version} for release in {staging_dir_str}\n\n" + "Verify the SDK contents:\n" + f" ls {staging_dir_str}/dist\n" + " node -e \"import('./dist/index.js').then(() => console.log('ok'))\"\n\n" + ) + else: + print(f"Staged package in {staging_dir}") + + if args.pack_output is not None: + output_path = run_npm_pack(staging_dir, args.pack_output) + print(f"npm pack output written to {output_path}") + finally: + if created_temp: + # Preserve the staging directory for further inspection. + pass + + return 0 + + +def prepare_staging_dir(staging_dir: Path | None) -> tuple[Path, bool]: + if staging_dir is not None: + staging_dir = staging_dir.resolve() + staging_dir.mkdir(parents=True, exist_ok=True) + if any(staging_dir.iterdir()): + raise RuntimeError(f"Staging directory {staging_dir} is not empty.") + return staging_dir, False + + temp_dir = Path(tempfile.mkdtemp(prefix="codex-npm-stage-")) + return temp_dir, True + + +def stage_sources(staging_dir: Path, version: str, package: str) -> None: + package_json: dict + package_json_path: Path | None = None + + if package == "codex": + bin_dir = staging_dir / "bin" + bin_dir.mkdir(parents=True, exist_ok=True) + shutil.copy2(CODEX_CLI_ROOT / "bin" / "codex.js", bin_dir / "codex.js") + + readme_src = REPO_ROOT / "README.md" + if readme_src.exists(): + shutil.copy2(readme_src, staging_dir / "README.md") + + package_json_path = CODEX_CLI_ROOT / "package.json" + elif package in CODEX_PLATFORM_PACKAGES: + platform_package = CODEX_PLATFORM_PACKAGES[package] + platform_npm_tag = platform_package["npm_tag"] + platform_version = compute_platform_package_version(version, platform_npm_tag) + + readme_src = REPO_ROOT / "README.md" + if readme_src.exists(): + shutil.copy2(readme_src, staging_dir / "README.md") + + with open(CODEX_CLI_ROOT / "package.json", "r", encoding="utf-8") as fh: + codex_package_json = json.load(fh) + + package_json = { + "name": CODEX_NPM_NAME, + "version": platform_version, + "license": codex_package_json.get("license", "Apache-2.0"), + "os": [platform_package["os"]], + "cpu": [platform_package["cpu"]], + "files": ["vendor"], + "repository": codex_package_json.get("repository"), + } + + engines = codex_package_json.get("engines") + if isinstance(engines, dict): + package_json["engines"] = engines + + package_manager = codex_package_json.get("packageManager") + if isinstance(package_manager, str): + package_json["packageManager"] = package_manager + elif package == "codex-responses-api-proxy": + bin_dir = staging_dir / "bin" + bin_dir.mkdir(parents=True, exist_ok=True) + launcher_src = ( + RESPONSES_API_PROXY_NPM_ROOT / "bin" / "codex-responses-api-proxy.js" + ) + shutil.copy2(launcher_src, bin_dir / "codex-responses-api-proxy.js") + + readme_src = RESPONSES_API_PROXY_NPM_ROOT / "README.md" + if readme_src.exists(): + shutil.copy2(readme_src, staging_dir / "README.md") + + package_json_path = RESPONSES_API_PROXY_NPM_ROOT / "package.json" + elif package == "codex-sdk": + package_json_path = CODEX_SDK_ROOT / "package.json" + stage_codex_sdk_sources(staging_dir) + else: + raise RuntimeError(f"Unknown package '{package}'.") + + if package_json_path is not None: + with open(package_json_path, "r", encoding="utf-8") as fh: + package_json = json.load(fh) + package_json["version"] = version + + if package == "codex": + package_json["files"] = ["bin/codex.js"] + package_json["optionalDependencies"] = { + CODEX_PLATFORM_PACKAGES[platform_package]["npm_name"]: ( + f"npm:{CODEX_NPM_NAME}@" + f"{compute_platform_package_version(version, CODEX_PLATFORM_PACKAGES[platform_package]['npm_tag'])}" + ) + for platform_package in PACKAGE_EXPANSIONS["codex"] + if platform_package != "codex" + } + + elif package == "codex-sdk": + scripts = package_json.get("scripts") + if isinstance(scripts, dict): + scripts.pop("prepare", None) + + dependencies = package_json.get("dependencies") + if not isinstance(dependencies, dict): + dependencies = {} + dependencies[CODEX_NPM_NAME] = version + package_json["dependencies"] = dependencies + + with open(staging_dir / "package.json", "w", encoding="utf-8") as out: + json.dump(package_json, out, indent=2) + out.write("\n") + + +def compute_platform_package_version(version: str, platform_tag: str) -> str: + # npm forbids republishing the same package name/version, so each + # platform-specific tarball needs a unique version string. + return f"{version}-{platform_tag}" + + +def run_command(cmd: list[str], cwd: Path | None = None) -> None: + print("+", " ".join(cmd), flush=True) + subprocess.run(cmd, cwd=cwd, check=True) + + +def stage_codex_sdk_sources(staging_dir: Path) -> None: + package_root = CODEX_SDK_ROOT + + run_command(["pnpm", "install", "--frozen-lockfile"], cwd=package_root) + run_command(["pnpm", "run", "build"], cwd=package_root) + + dist_src = package_root / "dist" + if not dist_src.exists(): + raise RuntimeError("codex-sdk build did not produce a dist directory.") + + shutil.copytree(dist_src, staging_dir / "dist") + + readme_src = package_root / "README.md" + if readme_src.exists(): + shutil.copy2(readme_src, staging_dir / "README.md") + + license_src = REPO_ROOT / "LICENSE" + if license_src.exists(): + shutil.copy2(license_src, staging_dir / "LICENSE") + + +def copy_native_binaries( + vendor_src: Path, + staging_dir: Path, + components: list[str], + target_filter: set[str] | None = None, +) -> None: + vendor_src = vendor_src.resolve() + if not vendor_src.exists(): + raise RuntimeError(f"Vendor source directory not found: {vendor_src}") + + components_set = set(components) + if not components_set: + return + + vendor_dest = staging_dir / "vendor" + if vendor_dest.exists(): + shutil.rmtree(vendor_dest) + vendor_dest.mkdir(parents=True, exist_ok=True) + + copied_targets: set[str] = set() + + for target_dir in vendor_src.iterdir(): + if not target_dir.is_dir(): + continue + + if target_filter is not None and target_dir.name not in target_filter: + continue + + copied_targets.add(target_dir.name) + + dest_target_dir = vendor_dest / target_dir.name + + if CODEX_PACKAGE_COMPONENT in components_set: + if dest_target_dir.exists(): + shutil.rmtree(dest_target_dir) + shutil.copytree(target_dir, dest_target_dir) + else: + dest_target_dir.mkdir(parents=True, exist_ok=True) + + for component in sorted(components_set - {CODEX_PACKAGE_COMPONENT}): + src_component_dir = target_dir / component + if not src_component_dir.exists(): + raise RuntimeError( + f"Missing native component '{component}' in vendor source: {src_component_dir}" + ) + + dest_component_dir = dest_target_dir / component + if dest_component_dir.exists(): + shutil.rmtree(dest_component_dir) + shutil.copytree(src_component_dir, dest_component_dir) + + if target_filter is not None: + missing_targets = sorted(target_filter - copied_targets) + if missing_targets: + missing_list = ", ".join(missing_targets) + raise RuntimeError( + f"Missing target directories in vendor source: {missing_list}" + ) + + +def run_npm_pack(staging_dir: Path, output_path: Path) -> Path: + output_path = output_path.resolve() + output_path.parent.mkdir(parents=True, exist_ok=True) + + with tempfile.TemporaryDirectory(prefix="codex-npm-pack-") as pack_dir_str: + pack_dir = Path(pack_dir_str) + npm_cache_dir = pack_dir / "npm-cache" + npm_logs_dir = pack_dir / "npm-logs" + npm_cache_dir.mkdir() + npm_logs_dir.mkdir() + env = os.environ.copy() + env["NPM_CONFIG_CACHE"] = str(npm_cache_dir) + env["NPM_CONFIG_LOGS_DIR"] = str(npm_logs_dir) + stdout = subprocess.check_output( + ["npm", "pack", "--json", "--pack-destination", str(pack_dir)], + cwd=staging_dir, + env=env, + text=True, + ) + try: + pack_output = json.loads(stdout) + except json.JSONDecodeError as exc: + raise RuntimeError("Failed to parse npm pack output.") from exc + + if not pack_output: + raise RuntimeError("npm pack did not produce an output tarball.") + + tarball_name = pack_output[0].get("filename") or pack_output[0].get("name") + if not tarball_name: + raise RuntimeError("Unable to determine npm pack output filename.") + + tarball_path = pack_dir / tarball_name + if not tarball_path.exists(): + raise RuntimeError(f"Expected npm pack output not found: {tarball_path}") + + shutil.move(str(tarball_path), output_path) + + return output_path + + +if __name__ == "__main__": + import sys + + sys.exit(main()) diff --git a/codex-cli/scripts/init_firewall.sh b/codex-cli/scripts/init_firewall.sh new file mode 100644 index 0000000000000000000000000000000000000000..1251325f0147c527943282b9f5fc8439215c1c4b --- /dev/null +++ b/codex-cli/scripts/init_firewall.sh @@ -0,0 +1,115 @@ +#!/bin/bash +set -euo pipefail # Exit on error, undefined vars, and pipeline failures +IFS=$'\n\t' # Stricter word splitting + +# Read allowed domains from file +ALLOWED_DOMAINS_FILE="/etc/codex/allowed_domains.txt" +if [ -f "$ALLOWED_DOMAINS_FILE" ]; then + ALLOWED_DOMAINS=() + while IFS= read -r domain; do + ALLOWED_DOMAINS+=("$domain") + done < "$ALLOWED_DOMAINS_FILE" + echo "Using domains from file: ${ALLOWED_DOMAINS[*]}" +else + # Fallback to default domains + ALLOWED_DOMAINS=("api.openai.com") + echo "Domains file not found, using default: ${ALLOWED_DOMAINS[*]}" +fi + +# Ensure we have at least one domain +if [ ${#ALLOWED_DOMAINS[@]} -eq 0 ]; then + echo "ERROR: No allowed domains specified" + exit 1 +fi + +# Flush existing rules and delete existing ipsets +iptables -F +iptables -X +iptables -t nat -F +iptables -t nat -X +iptables -t mangle -F +iptables -t mangle -X +ipset destroy allowed-domains 2>/dev/null || true + +# First allow DNS and localhost before any restrictions +# Allow outbound DNS +iptables -A OUTPUT -p udp --dport 53 -j ACCEPT +# Allow inbound DNS responses +iptables -A INPUT -p udp --sport 53 -j ACCEPT +# Allow localhost +iptables -A INPUT -i lo -j ACCEPT +iptables -A OUTPUT -o lo -j ACCEPT + +# Create ipset with CIDR support +ipset create allowed-domains hash:net + +# Resolve and add other allowed domains +for domain in "${ALLOWED_DOMAINS[@]}"; do + echo "Resolving $domain..." + ips=$(dig +short A "$domain") + if [ -z "$ips" ]; then + echo "ERROR: Failed to resolve $domain" + exit 1 + fi + + while read -r ip; do + if [[ ! "$ip" =~ ^[0-9]{1,3}\.[0-9]{1,3}\.[0-9]{1,3}\.[0-9]{1,3}$ ]]; then + echo "ERROR: Invalid IP from DNS for $domain: $ip" + exit 1 + fi + echo "Adding $ip for $domain" + ipset add allowed-domains "$ip" + done < <(echo "$ips") +done + +# Get host IP from default route +HOST_IP=$(ip route | grep default | cut -d" " -f3) +if [ -z "$HOST_IP" ]; then + echo "ERROR: Failed to detect host IP" + exit 1 +fi + +HOST_NETWORK=$(echo "$HOST_IP" | sed "s/\.[0-9]*$/.0\/24/") +echo "Host network detected as: $HOST_NETWORK" + +# Set up remaining iptables rules +iptables -A INPUT -s "$HOST_NETWORK" -j ACCEPT +iptables -A OUTPUT -d "$HOST_NETWORK" -j ACCEPT + +# Set default policies to DROP first +iptables -P INPUT DROP +iptables -P FORWARD DROP +iptables -P OUTPUT DROP + +# First allow established connections for already approved traffic +iptables -A INPUT -m state --state ESTABLISHED,RELATED -j ACCEPT +iptables -A OUTPUT -m state --state ESTABLISHED,RELATED -j ACCEPT + +# Then allow only specific outbound traffic to allowed domains +iptables -A OUTPUT -m set --match-set allowed-domains dst -j ACCEPT + +# Append final REJECT rules for immediate error responses +# For TCP traffic, send a TCP reset; for UDP, send ICMP port unreachable. +iptables -A INPUT -p tcp -j REJECT --reject-with tcp-reset +iptables -A INPUT -p udp -j REJECT --reject-with icmp-port-unreachable +iptables -A OUTPUT -p tcp -j REJECT --reject-with tcp-reset +iptables -A OUTPUT -p udp -j REJECT --reject-with icmp-port-unreachable +iptables -A FORWARD -p tcp -j REJECT --reject-with tcp-reset +iptables -A FORWARD -p udp -j REJECT --reject-with icmp-port-unreachable + +echo "Firewall configuration complete" +echo "Verifying firewall rules..." +if curl --connect-timeout 5 https://example.com >/dev/null 2>&1; then + echo "ERROR: Firewall verification failed - was able to reach https://example.com" + exit 1 +else + echo "Firewall verification passed - unable to reach https://example.com as expected" +fi + +# Always verify OpenAI API access is working +if ! curl --connect-timeout 5 https://api.openai.com >/dev/null 2>&1; then + echo "ERROR: Firewall verification failed - unable to reach https://api.openai.com" + exit 1 +else + echo "Firewall verification passed - able to reach https://api.openai.com as expected" +fi diff --git a/codex-cli/scripts/run_in_container.sh b/codex-cli/scripts/run_in_container.sh new file mode 100644 index 0000000000000000000000000000000000000000..607ec297a6c8ad1ad0708dbd9a6e784c0c8bc6a0 --- /dev/null +++ b/codex-cli/scripts/run_in_container.sh @@ -0,0 +1,95 @@ +#!/bin/bash +set -e + +# Usage: +# ./run_in_container.sh [--work_dir directory] "COMMAND" +# +# Examples: +# ./run_in_container.sh --work_dir project/code "ls -la" +# ./run_in_container.sh "echo Hello, world!" + +# Default the work directory to WORKSPACE_ROOT_DIR if not provided. +WORK_DIR="${WORKSPACE_ROOT_DIR:-$(pwd)}" +# Default allowed domains - can be overridden with OPENAI_ALLOWED_DOMAINS env var +OPENAI_ALLOWED_DOMAINS="${OPENAI_ALLOWED_DOMAINS:-api.openai.com}" + +# Parse optional flag. +if [ "$1" = "--work_dir" ]; then + if [ -z "$2" ]; then + echo "Error: --work_dir flag provided but no directory specified." + exit 1 + fi + WORK_DIR="$2" + shift 2 +fi + +WORK_DIR=$(realpath "$WORK_DIR") + +# Generate a unique container name based on the normalized work directory +CONTAINER_NAME="codex_$(echo "$WORK_DIR" | sed 's/\//_/g' | sed 's/[^a-zA-Z0-9_-]//g')" + +# Define cleanup to remove the container on script exit, ensuring no leftover containers +cleanup() { + docker rm -f "$CONTAINER_NAME" >/dev/null 2>&1 || true +} +# Trap EXIT to invoke cleanup regardless of how the script terminates +trap cleanup EXIT + +# Ensure a command is provided. +if [ "$#" -eq 0 ]; then + echo "Usage: $0 [--work_dir directory] \"COMMAND\"" + exit 1 +fi + +# Check if WORK_DIR is set. +if [ -z "$WORK_DIR" ]; then + echo "Error: No work directory provided and WORKSPACE_ROOT_DIR is not set." + exit 1 +fi + +# Verify that OPENAI_ALLOWED_DOMAINS is not empty +if [ -z "$OPENAI_ALLOWED_DOMAINS" ]; then + echo "Error: OPENAI_ALLOWED_DOMAINS is empty." + exit 1 +fi + +# Kill any existing container for the working directory using cleanup(), centralizing removal logic. +cleanup + +# Run the container with the specified directory mounted at the same path inside the container. +docker run --name "$CONTAINER_NAME" -d \ + -e OPENAI_API_KEY \ + --cap-add=NET_ADMIN \ + --cap-add=NET_RAW \ + -v "$WORK_DIR:/app$WORK_DIR" \ + codex \ + sleep infinity + +# Write the allowed domains to a file in the container +docker exec --user root "$CONTAINER_NAME" bash -c "mkdir -p /etc/codex" +for domain in $OPENAI_ALLOWED_DOMAINS; do + # Validate domain format to prevent injection + if [[ ! "$domain" =~ ^[a-zA-Z0-9][a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$ ]]; then + echo "Error: Invalid domain format: $domain" + exit 1 + fi + echo "$domain" | docker exec --user root -i "$CONTAINER_NAME" bash -c "cat >> /etc/codex/allowed_domains.txt" +done + +# Set proper permissions on the domains file +docker exec --user root "$CONTAINER_NAME" bash -c "chmod 444 /etc/codex/allowed_domains.txt && chown root:root /etc/codex/allowed_domains.txt" + +# Initialize the firewall inside the container as root user +docker exec --user root "$CONTAINER_NAME" bash -c "/usr/local/bin/init_firewall.sh" + +# Remove the firewall script after running it +docker exec --user root "$CONTAINER_NAME" bash -c "rm -f /usr/local/bin/init_firewall.sh" + +# Execute the provided command in the container, ensuring it runs in the work directory. +# We use a parameterized bash command to safely handle the command and directory. + +quoted_args="" +for arg in "$@"; do + quoted_args+=" $(printf '%q' "$arg")" +done +docker exec -it "$CONTAINER_NAME" bash -c "cd \"/app$WORK_DIR\" && codex --sandbox workspace-write --ask-for-approval on-request ${quoted_args}" diff --git a/codex-rs/agent-identity/src/lib.rs b/codex-rs/agent-identity/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..14a9e35190118d3f62fe1197dd9735229f34e4cd --- /dev/null +++ b/codex-rs/agent-identity/src/lib.rs @@ -0,0 +1,1000 @@ +use std::collections::BTreeMap; +use std::error::Error as StdError; +use std::fmt; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use chrono::SecondsFormat; +use chrono::Utc; +use codex_http_client::HttpClient; +use codex_http_client::HttpError; +use codex_protocol::auth::PlanType as AuthPlanType; +use codex_protocol::protocol::SessionSource; +use crypto_box::SecretKey as Curve25519SecretKey; +use ed25519_dalek::Signer as _; +use ed25519_dalek::SigningKey; +use ed25519_dalek::VerifyingKey; +use ed25519_dalek::pkcs8::DecodePrivateKey; +use ed25519_dalek::pkcs8::EncodePrivateKey; +use http::StatusCode; +use jsonwebtoken::Algorithm; +use jsonwebtoken::DecodingKey; +use jsonwebtoken::Validation; +use jsonwebtoken::decode; +use jsonwebtoken::decode_header; +use jsonwebtoken::jwk::JwkSet; +use rand::TryRngCore; +use rand::rngs::OsRng; +use serde::Deserialize; +use serde::Serialize; +use serde::de::DeserializeOwned; +use sha2::Digest as _; +use sha2::Sha512; + +const AGENT_TASK_REGISTRATION_TIMEOUT: Duration = Duration::from_secs(30); +const AGENT_IDENTITY_JWKS_TIMEOUT: Duration = Duration::from_secs(10); +const AGENT_IDENTITY_JWT_AUDIENCE: &str = "codex-app-server"; +const AGENT_IDENTITY_JWT_ISSUER: &str = "https://chatgpt.com/codex-backend/agent-identity"; +const AGENT_REGISTRATION_TIMEOUT: Duration = Duration::from_secs(15); +const PROD_AGENT_IDENTITY_AUTHAPI_BASE_URL: &str = "https://auth.openai.com/api/accounts"; +const STAGING_AGENT_IDENTITY_AUTHAPI_BASE_URL: &str = "https://auth.api.openai.org/api/accounts"; +const AGENT_IDENTITY_KEY_SEED_BYTES: usize = 64; +const AGENT_IDENTITY_KEY_DERIVATION_CONTEXT: &[u8] = b"codex-agent-identity-ed25519-v1"; + +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum ChatGptEnvironment { + #[default] + Production, + Staging, +} + +impl ChatGptEnvironment { + pub fn from_chatgpt_base_url(chatgpt_base_url: &str) -> Result { + match chatgpt_base_url.trim_end_matches('/') { + "https://chatgpt.com" + | "https://chatgpt.com/backend-api" + | "https://chatgpt.com/codex" + | "https://chatgpt.com/backend-api/codex" + | "https://chat.openai.com" + | "https://chat.openai.com/backend-api" + | "https://chat.openai.com/codex" + | "https://chat.openai.com/backend-api/codex" => Ok(Self::Production), + "https://chatgpt-staging.com" + | "https://chatgpt-staging.com/backend-api" + | "https://chatgpt-staging.com/codex" + | "https://chatgpt-staging.com/backend-api/codex" => Ok(Self::Staging), + _ => anyhow::bail!( + "Agent Identity only supports production and staging ChatGPT environments" + ), + } + } + + pub fn chatgpt_base_url(self) -> &'static str { + match self { + Self::Production => "https://chatgpt.com/backend-api", + Self::Staging => "https://chatgpt-staging.com/backend-api", + } + } + + pub fn agent_identity_authapi_base_url(self) -> &'static str { + match self { + Self::Production => PROD_AGENT_IDENTITY_AUTHAPI_BASE_URL, + Self::Staging => STAGING_AGENT_IDENTITY_AUTHAPI_BASE_URL, + } + } +} + +/// Borrowed durable signing material for a registered agent identity. +/// +/// This intentionally does not include a task id. Task ids are scoped to a +/// single Codex run, while the agent runtime id and private key are the +/// reusable identity material used to register and sign that run task. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct AgentIdentityKey<'a> { + pub agent_runtime_id: &'a str, + pub private_key_pkcs8_base64: &'a str, +} + +#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] +pub struct AgentBillOfMaterials { + pub agent_version: String, + pub agent_harness_id: String, + pub running_location: String, +} + +pub struct GeneratedAgentKeyMaterial { + pub private_key_pkcs8_base64: String, + pub public_key_ssh: String, +} + +/// Claims carried by an Agent Identity JWT. +#[derive(Clone, Debug, Deserialize, PartialEq, Eq)] +pub struct AgentIdentityJwtClaims { + pub iss: String, + pub aud: String, + pub iat: usize, + pub exp: usize, + pub agent_runtime_id: String, + pub agent_private_key: String, + pub account_id: String, + pub chatgpt_user_id: String, + pub email: Option, + pub plan_type: AuthPlanType, + pub chatgpt_account_is_fedramp: bool, +} + +#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] +struct AgentAssertionEnvelope { + agent_runtime_id: String, + task_id: String, + timestamp: String, + signature: String, +} + +#[derive(Serialize)] +struct RegisterTaskRequest { + timestamp: String, + signature: String, +} + +#[derive(Deserialize)] +struct RegisterTaskResponse { + #[serde(default)] + task_id: Option, + #[serde(default, rename = "taskId")] + task_id_camel: Option, + #[serde(default)] + encrypted_task_id: Option, + #[serde(default, rename = "encryptedTaskId")] + encrypted_task_id_camel: Option, +} + +#[derive(Debug, Serialize)] +struct RegisterAgentRequest { + abom: AgentBillOfMaterials, + agent_public_key: String, + capabilities: Vec, + ttl: Option, +} + +#[derive(Debug, Deserialize)] +struct RegisterAgentResponse { + agent_runtime_id: String, +} + +/// HTTP status failure returned by Agent Identity registration endpoints. +#[derive(Debug)] +pub struct AgentIdentityRegistrationHttpError { + operation: &'static str, + status: StatusCode, + body: String, +} + +impl AgentIdentityRegistrationHttpError { + fn new(operation: &'static str, status: StatusCode, body: String) -> Self { + Self { + operation, + status, + body, + } + } + + /// HTTP status returned by the registration endpoint. + pub fn status(&self) -> StatusCode { + self.status + } +} + +impl fmt::Display for AgentIdentityRegistrationHttpError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + if self.body.is_empty() { + write!(f, "{} failed with status {}", self.operation, self.status) + } else { + write!( + f, + "{} failed with status {}: {}", + self.operation, self.status, self.body + ) + } + } +} + +impl StdError for AgentIdentityRegistrationHttpError {} + +/// Returns whether an Agent Identity registration error is safe to retry. +pub fn is_retryable_registration_error(error: &anyhow::Error) -> bool { + error.chain().any(is_retryable_registration_cause) +} + +fn is_retryable_registration_cause(cause: &(dyn StdError + 'static)) -> bool { + if let Some(error) = cause.downcast_ref::() { + return is_retryable_registration_status(error.status()); + } + + if let Some(error) = cause.downcast_ref::() { + if let Some(status) = error.status() { + return is_retryable_registration_status(status); + } + return error.is_timeout() || error.is_connect() || error.is_request(); + } + + false +} + +fn is_retryable_registration_status(status: StatusCode) -> bool { + status == StatusCode::TOO_MANY_REQUESTS || status.is_server_error() +} + +pub fn authorization_header_for_agent_task( + key: AgentIdentityKey<'_>, + task_id: &str, +) -> Result { + let timestamp = Utc::now().to_rfc3339_opts(SecondsFormat::Secs, true); + let envelope = AgentAssertionEnvelope { + agent_runtime_id: key.agent_runtime_id.to_string(), + task_id: task_id.to_string(), + timestamp: timestamp.clone(), + signature: sign_agent_assertion_payload(key, task_id, ×tamp)?, + }; + let serialized_assertion = serialize_agent_assertion(&envelope)?; + Ok(format!("AgentAssertion {serialized_assertion}")) +} + +pub async fn fetch_agent_identity_jwks( + client: &HttpClient, + agent_identity_jwt_base_url: &str, +) -> Result { + let response = client + .get(agent_identity_jwks_url(agent_identity_jwt_base_url)) + .timeout(AGENT_IDENTITY_JWKS_TIMEOUT) + .send() + .await + .context("failed to request agent identity JWKS")? + .error_for_status() + .context("agent identity JWKS endpoint returned an error")?; + + response + .json() + .await + .context("failed to decode agent identity JWKS") +} + +pub fn decode_agent_identity_jwt( + jwt: &str, + jwks: Option<&JwkSet>, +) -> Result { + let Some(jwks) = jwks else { + return decode_agent_identity_jwt_payload(jwt); + }; + + let header = decode_header(jwt).context("failed to decode agent identity JWT header")?; + let kid = header + .kid + .context("agent identity JWT header does not include a kid")?; + let jwk = jwks + .find(&kid) + .with_context(|| format!("agent identity JWT kid {kid} is not trusted"))?; + let decoding_key = DecodingKey::from_jwk(jwk).context("failed to build JWT decoding key")?; + let mut validation = Validation::new(Algorithm::RS256); + validation.set_audience(&[AGENT_IDENTITY_JWT_AUDIENCE]); + validation.set_issuer(&[AGENT_IDENTITY_JWT_ISSUER]); + validation.required_spec_claims.insert("iss".to_string()); + validation.required_spec_claims.insert("aud".to_string()); + decode::(jwt, &decoding_key, &validation) + .map(|data| data.claims) + .context("failed to verify agent identity JWT") +} + +fn decode_agent_identity_jwt_payload(jwt: &str) -> Result { + let mut parts = jwt.split('.'); + let (_header_b64, payload_b64, _sig_b64) = match (parts.next(), parts.next(), parts.next()) { + (Some(h), Some(p), Some(s)) if !h.is_empty() && !p.is_empty() && !s.is_empty() => (h, p, s), + _ => anyhow::bail!("invalid agent identity JWT format"), + }; + anyhow::ensure!(parts.next().is_none(), "invalid agent identity JWT format"); + + let payload_bytes = URL_SAFE_NO_PAD + .decode(payload_b64) + .context("agent identity JWT payload is not valid base64url")?; + serde_json::from_slice(&payload_bytes).context("agent identity JWT payload is not valid JSON") +} + +pub fn sign_task_registration_payload( + key: AgentIdentityKey<'_>, + timestamp: &str, +) -> Result { + let signing_key = signing_key_from_private_key_pkcs8_base64(key.private_key_pkcs8_base64)?; + let payload = format!("{}:{timestamp}", key.agent_runtime_id); + Ok(BASE64_STANDARD.encode(signing_key.sign(payload.as_bytes()).to_bytes())) +} + +pub async fn register_agent_task( + client: &HttpClient, + agent_identity_authapi_base_url: &str, + key: AgentIdentityKey<'_>, +) -> Result { + let timestamp = Utc::now().to_rfc3339_opts(SecondsFormat::Secs, true); + let request = RegisterTaskRequest { + signature: sign_task_registration_payload(key, ×tamp)?, + timestamp, + }; + let url = agent_task_registration_url(agent_identity_authapi_base_url, key.agent_runtime_id); + + let response = client + .post(url) + .timeout(AGENT_TASK_REGISTRATION_TIMEOUT) + .json(&request) + .send() + .await + .context("failed to register agent task")?; + if !response.status().is_success() { + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + let body = if body.len() > 512 { + format!("{}...", body.chars().take(512).collect::()) + } else { + body + }; + return Err(AgentIdentityRegistrationHttpError::new( + "agent task registration", + status, + body, + ) + .into()); + } + + let response = response + .json() + .await + .context("failed to decode agent task registration response")?; + + task_id_from_register_task_response(key, response) +} + +pub async fn register_agent_identity( + client: &HttpClient, + agent_identity_authapi_base_url: &str, + access_token: &str, + is_fedramp_account: bool, + key_material: &GeneratedAgentKeyMaterial, + abom: AgentBillOfMaterials, + capabilities: Vec, +) -> Result { + let url = agent_registration_url(agent_identity_authapi_base_url); + let request = RegisterAgentRequest { + abom, + agent_public_key: key_material.public_key_ssh.clone(), + capabilities, + ttl: None, + }; + + let mut request_builder = client + .post(&url) + .bearer_auth(access_token) + .json(&request) + .timeout(AGENT_REGISTRATION_TIMEOUT); + if is_fedramp_account { + request_builder = request_builder.header("X-OpenAI-Fedramp", "true"); + } + + let response = request_builder + .send() + .await + .with_context(|| format!("failed to send agent identity registration request to {url}"))? + .error_for_status() + .with_context(|| format!("agent identity registration failed for {url}"))? + .json::() + .await + .with_context(|| format!("failed to parse agent identity response from {url}"))?; + + Ok(response.agent_runtime_id) +} + +fn task_id_from_register_task_response( + key: AgentIdentityKey<'_>, + response: RegisterTaskResponse, +) -> Result { + if let Some(task_id) = response.task_id.or(response.task_id_camel) { + return Ok(task_id); + } + let encrypted_task_id = response + .encrypted_task_id + .or(response.encrypted_task_id_camel) + .context("agent task registration response omitted task id")?; + decrypt_task_id_response(key, &encrypted_task_id) +} + +pub fn decrypt_task_id_response( + key: AgentIdentityKey<'_>, + encrypted_task_id: &str, +) -> Result { + let signing_key = signing_key_from_private_key_pkcs8_base64(key.private_key_pkcs8_base64)?; + let ciphertext = BASE64_STANDARD + .decode(encrypted_task_id) + .context("encrypted task id is not valid base64")?; + let plaintext = curve25519_secret_key_from_signing_key(&signing_key) + .unseal(&ciphertext) + .map_err(|_| anyhow::anyhow!("failed to decrypt encrypted task id"))?; + String::from_utf8(plaintext).context("decrypted task id is not valid UTF-8") +} + +pub fn generate_agent_key_material() -> Result { + let mut seed_material = [0u8; AGENT_IDENTITY_KEY_SEED_BYTES]; + OsRng + .try_fill_bytes(&mut seed_material) + .context("failed to generate agent identity private key seed material")?; + // Ed25519 stores a 32-byte seed, so derive it from all sampled seed material. + let mut digest = Sha512::new(); + digest.update(AGENT_IDENTITY_KEY_DERIVATION_CONTEXT); + digest.update(seed_material); + let digest = digest.finalize(); + let mut secret_key_bytes = [0u8; 32]; + secret_key_bytes.copy_from_slice(&digest[..32]); + let signing_key = SigningKey::from_bytes(&secret_key_bytes); + let private_key_pkcs8 = signing_key + .to_pkcs8_der() + .context("failed to encode agent identity private key as PKCS#8")?; + + Ok(GeneratedAgentKeyMaterial { + private_key_pkcs8_base64: BASE64_STANDARD.encode(private_key_pkcs8.as_bytes()), + public_key_ssh: encode_ssh_ed25519_public_key(&signing_key.verifying_key()), + }) +} + +pub fn public_key_ssh_from_private_key_pkcs8_base64( + private_key_pkcs8_base64: &str, +) -> Result { + let signing_key = signing_key_from_private_key_pkcs8_base64(private_key_pkcs8_base64)?; + Ok(encode_ssh_ed25519_public_key(&signing_key.verifying_key())) +} + +pub fn verifying_key_from_private_key_pkcs8_base64( + private_key_pkcs8_base64: &str, +) -> Result { + let signing_key = signing_key_from_private_key_pkcs8_base64(private_key_pkcs8_base64)?; + Ok(signing_key.verifying_key()) +} + +pub fn curve25519_secret_key_from_private_key_pkcs8_base64( + private_key_pkcs8_base64: &str, +) -> Result { + let signing_key = signing_key_from_private_key_pkcs8_base64(private_key_pkcs8_base64)?; + Ok(curve25519_secret_key_from_signing_key(&signing_key)) +} + +pub fn agent_registration_url(agent_identity_authapi_base_url: &str) -> String { + agent_identity_authapi_url(agent_identity_authapi_base_url, "/v1/agent/register") +} + +pub fn agent_task_registration_url( + agent_identity_authapi_base_url: &str, + agent_runtime_id: &str, +) -> String { + agent_identity_authapi_url( + agent_identity_authapi_base_url, + &format!("/v1/agent/{agent_runtime_id}/task/register"), + ) +} + +pub fn agent_identity_jwks_url(agent_identity_jwt_base_url: &str) -> String { + let trimmed = agent_identity_jwt_base_url.trim_end_matches('/'); + if trimmed.contains("/backend-api") { + format!("{trimmed}/wham/agent-identities/jwks") + } else { + format!("{trimmed}/agent-identities/jwks") + } +} + +fn agent_identity_authapi_url(agent_identity_authapi_base_url: &str, api_path: &str) -> String { + let base_url = agent_identity_authapi_base_url.trim_end_matches('/'); + format!("{base_url}{api_path}") +} + +pub fn build_abom(session_source: SessionSource) -> AgentBillOfMaterials { + AgentBillOfMaterials { + agent_version: env!("CARGO_PKG_VERSION").to_string(), + agent_harness_id: match &session_source { + SessionSource::VSCode => "codex-app".to_string(), + SessionSource::Cli + | SessionSource::Exec + | SessionSource::Mcp + | SessionSource::Custom(_) + | SessionSource::Internal(_) + | SessionSource::SubAgent(_) + | SessionSource::Unknown => "codex-cli".to_string(), + }, + running_location: format!("{}-{}", session_source, std::env::consts::OS), + } +} + +pub fn encode_ssh_ed25519_public_key(verifying_key: &VerifyingKey) -> String { + let mut blob = Vec::with_capacity(4 + 11 + 4 + 32); + append_ssh_string(&mut blob, b"ssh-ed25519"); + append_ssh_string(&mut blob, verifying_key.as_bytes()); + format!("ssh-ed25519 {}", BASE64_STANDARD.encode(blob)) +} + +fn sign_agent_assertion_payload( + key: AgentIdentityKey<'_>, + task_id: &str, + timestamp: &str, +) -> Result { + let signing_key = signing_key_from_private_key_pkcs8_base64(key.private_key_pkcs8_base64)?; + let payload = format!("{}:{task_id}:{timestamp}", key.agent_runtime_id); + Ok(BASE64_STANDARD.encode(signing_key.sign(payload.as_bytes()).to_bytes())) +} + +fn serialize_agent_assertion(envelope: &AgentAssertionEnvelope) -> Result { + let payload = serde_json::to_vec(&BTreeMap::from([ + ("agent_runtime_id", envelope.agent_runtime_id.as_str()), + ("signature", envelope.signature.as_str()), + ("task_id", envelope.task_id.as_str()), + ("timestamp", envelope.timestamp.as_str()), + ])) + .context("failed to serialize agent assertion envelope")?; + Ok(URL_SAFE_NO_PAD.encode(payload)) +} + +fn curve25519_secret_key_from_signing_key(signing_key: &SigningKey) -> Curve25519SecretKey { + let digest = Sha512::digest(signing_key.to_bytes()); + let mut secret_key = [0u8; 32]; + secret_key.copy_from_slice(&digest[..32]); + secret_key[0] &= 248; + secret_key[31] &= 127; + secret_key[31] |= 64; + Curve25519SecretKey::from(secret_key) +} + +fn append_ssh_string(buf: &mut Vec, value: &[u8]) { + buf.extend_from_slice(&(value.len() as u32).to_be_bytes()); + buf.extend_from_slice(value); +} + +fn signing_key_from_private_key_pkcs8_base64(private_key_pkcs8_base64: &str) -> Result { + let private_key = BASE64_STANDARD + .decode(private_key_pkcs8_base64) + .context("stored agent identity private key is not valid base64")?; + SigningKey::from_pkcs8_der(&private_key) + .context("stored agent identity private key is not valid PKCS#8") +} + +#[cfg(test)] +mod tests { + use base64::Engine as _; + use ed25519_dalek::Signature; + use ed25519_dalek::Verifier as _; + use jsonwebtoken::EncodingKey; + use jsonwebtoken::Header; + use pretty_assertions::assert_eq; + + use codex_protocol::auth::KnownPlan; + + use super::*; + + #[test] + fn register_task_request_uses_single_run_task_shape() { + let request = RegisterTaskRequest { + timestamp: "2026-04-23T00:00:00Z".to_string(), + signature: "signature".to_string(), + }; + + let serialized = serde_json::to_value(request).expect("serialize request"); + + assert_eq!( + serialized, + serde_json::json!({ + "timestamp": "2026-04-23T00:00:00Z", + "signature": "signature", + }) + ); + } + + #[test] + fn authorization_header_for_agent_task_serializes_signed_agent_assertion() { + let signing_key = SigningKey::from_bytes(&[7u8; 32]); + let private_key = signing_key + .to_pkcs8_der() + .expect("encode test key material"); + let key = AgentIdentityKey { + agent_runtime_id: "agent-123", + private_key_pkcs8_base64: &BASE64_STANDARD.encode(private_key.as_bytes()), + }; + + let header = authorization_header_for_agent_task(key, "task-123") + .expect("build agent assertion header"); + let token = header + .strip_prefix("AgentAssertion ") + .expect("agent assertion scheme"); + let payload = URL_SAFE_NO_PAD + .decode(token) + .expect("valid base64url payload"); + let envelope: AgentAssertionEnvelope = + serde_json::from_slice(&payload).expect("valid assertion envelope"); + + assert_eq!( + envelope, + AgentAssertionEnvelope { + agent_runtime_id: "agent-123".to_string(), + task_id: "task-123".to_string(), + timestamp: envelope.timestamp.clone(), + signature: envelope.signature.clone(), + } + ); + let signature_bytes = BASE64_STANDARD + .decode(&envelope.signature) + .expect("valid base64 signature"); + let signature = Signature::from_slice(&signature_bytes).expect("valid signature bytes"); + signing_key + .verifying_key() + .verify( + format!( + "{}:{}:{}", + envelope.agent_runtime_id, envelope.task_id, envelope.timestamp + ) + .as_bytes(), + &signature, + ) + .expect("signature should verify"); + } + + #[test] + fn decode_agent_identity_jwt_reads_claims() { + let jwt = jwt_with_payload(serde_json::json!({ + "iss": AGENT_IDENTITY_JWT_ISSUER, + "aud": AGENT_IDENTITY_JWT_AUDIENCE, + "iat": 1_700_000_000usize, + "exp": 4_000_000_000usize, + "agent_runtime_id": "agent-runtime-id", + "agent_private_key": "private-key", + "account_id": "account-id", + "chatgpt_user_id": "user-id", + "email": "user@example.com", + "plan_type": "pro", + "chatgpt_account_is_fedramp": false, + })); + + let claims = decode_agent_identity_jwt(&jwt, /*jwks*/ None).expect("JWT should decode"); + + assert_eq!( + claims, + AgentIdentityJwtClaims { + iss: AGENT_IDENTITY_JWT_ISSUER.to_string(), + aud: AGENT_IDENTITY_JWT_AUDIENCE.to_string(), + iat: 1_700_000_000, + exp: 4_000_000_000, + agent_runtime_id: "agent-runtime-id".to_string(), + agent_private_key: "private-key".to_string(), + account_id: "account-id".to_string(), + chatgpt_user_id: "user-id".to_string(), + email: Some("user@example.com".to_string()), + plan_type: AuthPlanType::Known(KnownPlan::Pro), + chatgpt_account_is_fedramp: false, + } + ); + } + + #[test] + fn decode_agent_identity_jwt_accepts_missing_email() { + let jwt = jwt_with_payload(serde_json::json!({ + "iss": AGENT_IDENTITY_JWT_ISSUER, + "aud": AGENT_IDENTITY_JWT_AUDIENCE, + "iat": 1_700_000_000usize, + "exp": 4_000_000_000usize, + "agent_runtime_id": "agent-runtime-id", + "agent_private_key": "private-key", + "account_id": "account-id", + "chatgpt_user_id": "user-id", + "plan_type": "pro", + "chatgpt_account_is_fedramp": false, + })); + + let claims = decode_agent_identity_jwt(&jwt, /*jwks*/ None).expect("JWT should decode"); + + assert_eq!(claims.email, None); + } + + #[test] + fn decode_agent_identity_jwt_maps_raw_plan_aliases() { + let jwt = jwt_with_payload(serde_json::json!({ + "iss": AGENT_IDENTITY_JWT_ISSUER, + "aud": AGENT_IDENTITY_JWT_AUDIENCE, + "iat": 1_700_000_000usize, + "exp": 4_000_000_000usize, + "agent_runtime_id": "agent-runtime-id", + "agent_private_key": "private-key", + "account_id": "account-id", + "chatgpt_user_id": "user-id", + "email": "user@example.com", + "plan_type": "hc", + "chatgpt_account_is_fedramp": false, + })); + + let claims = decode_agent_identity_jwt(&jwt, /*jwks*/ None).expect("JWT should decode"); + + assert_eq!(claims.plan_type, AuthPlanType::Known(KnownPlan::Enterprise)); + } + + #[test] + fn decode_agent_identity_jwt_verifies_when_jwks_is_present() { + let jwks = test_jwks("test-key"); + let claims = AgentIdentityJwtClaims { + iss: AGENT_IDENTITY_JWT_ISSUER.to_string(), + aud: AGENT_IDENTITY_JWT_AUDIENCE.to_string(), + iat: 1_700_000_000, + exp: 4_000_000_000, + agent_runtime_id: "agent-runtime-id".to_string(), + agent_private_key: "private-key".to_string(), + account_id: "account-id".to_string(), + chatgpt_user_id: "user-id".to_string(), + email: Some("user@example.com".to_string()), + plan_type: AuthPlanType::Known(KnownPlan::Pro), + chatgpt_account_is_fedramp: false, + }; + let jwt = jsonwebtoken::encode( + &test_jwt_header("test-key"), + &serde_json::json!({ + "iss": claims.iss, + "aud": claims.aud, + "iat": claims.iat, + "exp": claims.exp, + "agent_runtime_id": claims.agent_runtime_id, + "agent_private_key": claims.agent_private_key, + "account_id": claims.account_id, + "chatgpt_user_id": claims.chatgpt_user_id, + "email": claims.email, + "plan_type": "pro", + "chatgpt_account_is_fedramp": claims.chatgpt_account_is_fedramp, + }), + &test_rsa_encoding_key(), + ) + .expect("JWT should encode"); + + let expected_claims = AgentIdentityJwtClaims { + iss: AGENT_IDENTITY_JWT_ISSUER.to_string(), + aud: AGENT_IDENTITY_JWT_AUDIENCE.to_string(), + iat: 1_700_000_000, + exp: 4_000_000_000, + agent_runtime_id: "agent-runtime-id".to_string(), + agent_private_key: "private-key".to_string(), + account_id: "account-id".to_string(), + chatgpt_user_id: "user-id".to_string(), + email: Some("user@example.com".to_string()), + plan_type: AuthPlanType::Known(KnownPlan::Pro), + chatgpt_account_is_fedramp: false, + }; + assert_eq!( + decode_agent_identity_jwt(&jwt, Some(&jwks)).expect("JWT should verify"), + expected_claims + ); + } + + #[test] + fn decode_agent_identity_jwt_rejects_untrusted_kid() { + let jwks = test_jwks("other-key"); + + let jwt = jsonwebtoken::encode( + &test_jwt_header("test-key"), + &serde_json::json!({ + "iss": AGENT_IDENTITY_JWT_ISSUER, + "aud": AGENT_IDENTITY_JWT_AUDIENCE, + "iat": 1_700_000_000, + "exp": 4_000_000_000usize, + "agent_runtime_id": "agent-runtime-id", + "agent_private_key": "private-key", + "account_id": "account-id", + "chatgpt_user_id": "user-id", + "email": "user@example.com", + "plan_type": "pro", + "chatgpt_account_is_fedramp": false, + }), + &test_rsa_encoding_key(), + ) + .expect("JWT should encode"); + + decode_agent_identity_jwt(&jwt, Some(&jwks)).expect_err("JWT should not verify"); + } + + #[test] + fn decode_agent_identity_jwt_requires_issuer_and_audience() { + let jwks = test_jwks("test-key"); + let jwt = jsonwebtoken::encode( + &test_jwt_header("test-key"), + &serde_json::json!({ + "iat": 1_700_000_000, + "exp": 4_000_000_000usize, + "agent_runtime_id": "agent-runtime-id", + "agent_private_key": "private-key", + "account_id": "account-id", + "chatgpt_user_id": "user-id", + "email": "user@example.com", + "plan_type": "pro", + "chatgpt_account_is_fedramp": false, + }), + &test_rsa_encoding_key(), + ) + .expect("JWT should encode"); + + decode_agent_identity_jwt(&jwt, Some(&jwks)).expect_err("JWT should not verify"); + } + + fn test_jwt_header(kid: &str) -> Header { + let mut header = Header::new(Algorithm::RS256); + header.kid = Some(kid.to_string()); + header + } + + fn test_rsa_encoding_key() -> EncodingKey { + EncodingKey::from_rsa_pem( + br#"-----BEGIN PRIVATE KEY----- +MIIEvgIBADANBgkqhkiG9w0BAQEFAASCBKgwggSkAgEAAoIBAQDWpAXYypOsYAwO +bvBduMk/mxaoYDze0AZSzaSzLuIlcsl2EKDgC3AabhIWXh/qTGEJLOU3VB1e5mO9 +FPbBlmIZSL3FQTbyt/hYutPFKfCou5PLmScw/TzILS3/RhT8UY9kxxZvXiEbTki9 +mvxRuZFpVqDFJHwfitIjKZGhXDCYVKurPTrxetYZJg0h8sQBLKjkZ0BqqaTUkAsg +0eBgZAlXEzG3By8PGhUqYLt6W1Q3KYw0FmGy/gTyzH1g0ukGgSJvOd8SkNT8MbOs +zl5kKxDNqpuEE6UZ3jbuJ+5382d31w+rOAJRzbf7QVdI9+luCSwJcDACYPQ4WNBa +uCpV0ovpAgMBAAECggEAVu84LwZdqYN9XpswX8VoPYrjMm9IODapWQBRpQFoNyK2 +1ksF3bjEPvA2Azk8U/l7k+vLKw22l6lY3EyRZPcz5GnB8xLm3ogE3mtNOp4yCyVu +RxhQ91aaN7mU17/a4BdorLi2LYVCg3zBmYociD1Q2AluNGsCmwPu+K7tfR2J0Sg8 +NjqiTbDG1XDpR/icwgC9t6vh8lZpCHDhF4tbQfLLVLeA/OdcuzXDyMCXbmdVIdBQ +rm4aIFmr2e1/2ctTbCg85S6AGFTH+pSLjrwTzyvf+F6NW5uNjLQAQLFj+EznBDxj +Xdx90cySrjsKK6PVWQF4RiTvkSW8eWL7R6B2FZbGwQKBgQDuVQRj72hWloR7mbEL +aUEEv3pIXTMXWEsoMBNczos/1L1RnAN1AI44TurznasPZAWvQj+kVbLDR+TAeZrL +iA8HIWswQUI18hFmgKzSkwIXGtubcKVrgsKeS4lMDKCM/Ef6WAYdeq6ronoY5lCN +YrJFmGp81W5zcV7lyiycgbSiGwKBgQDmjWYf6pZjrK7Z+OJ3X1AZfi2vss15SCvL +3fPgzIDbViztpGyQhc3DQZIsBNIu0xZp/veGce9TEeTds2ro9NfdJFeou8+fC7Pq +sOsM3amGFFi+ZW/9BWyjZEM88bgWWAjqLHbpfHDxjAf5CSxddqxgHlbP0Ytyb1Vg +gmPDn9YKSwKBgQDbTi3hC35WFuDHn0/zcSHcDZmnFuOZeqyFyV83yfMGhGrEuqvP +sPgtRikajJ3IZsB4WZyYSidZXEFY/0z6NjOl2xF38MTNQPbT/FmK1q1Yt2UWrlv5 +BvSwlk87RG9D7C0LZo4R+D7cPoDdgqjiwMvMEIkEX5zn641oI1ZTmWKuuwKBgQCD +KF+3unnRvHRAVoFnTZbA2fJdqMeRvogD04GhGlYX8V9f1hFY6nXTJaNlXVzA/J8c +r8ra9kgjJuPfZ+ljG58OFFW2DRohLcQtuHYPfK6rMzoFHqnl9EcIcMp7ijuionR3 +29HOJFgQYgxLFXfit9d6WugiE+BTupiEbckZif13HwKBgE/lAlkVHP6YahOO2Ljc +J1bwkqKZTB5dHolX9A58e/xXnfZ5P8f3Z83+Izap3FwqQulk7b1WO1MQcHuVg2NN +5da0D4h2rYOXnbYIg0BVu4spQbaM6ewsp66b8+MzLOBvj8SzWdt1Oyw0q/MRyQAR +8U4M2TSWCKUY/A6sT4W8+mT9 +-----END PRIVATE KEY-----"#, + ) + .expect("test RSA key should parse") + } + + fn test_jwks(kid: &str) -> jsonwebtoken::jwk::JwkSet { + serde_json::from_value(serde_json::json!({ + "keys": [{ + "kty": "RSA", + "kid": kid, + "use": "sig", + "alg": "RS256", + "n": "1qQF2MqTrGAMDm7wXbjJP5sWqGA83tAGUs2ksy7iJXLJdhCg4AtwGm4SFl4f6kxhCSzlN1QdXuZjvRT2wZZiGUi9xUE28rf4WLrTxSnwqLuTy5knMP08yC0t_0YU_FGPZMcWb14hG05IvZr8UbmRaVagxSR8H4rSIymRoVwwmFSrqz068XrWGSYNIfLEASyo5GdAaqmk1JALINHgYGQJVxMxtwcvDxoVKmC7eltUNymMNBZhsv4E8sx9YNLpBoEibznfEpDU_DGzrM5eZCsQzaqbhBOlGd427ifud_Nnd9cPqzgCUc23-0FXSPfpbgksCXAwAmD0OFjQWrgqVdKL6Q", + "e": "AQAB", + }] + })) + .expect("test JWKS should parse") + } + + #[test] + fn chatgpt_environment_maps_known_urls_to_authapi() -> anyhow::Result<()> { + assert_eq!( + ChatGptEnvironment::from_chatgpt_base_url("https://chatgpt.com/backend-api/codex")?, + ChatGptEnvironment::Production + ); + assert_eq!( + ChatGptEnvironment::Production.agent_identity_authapi_base_url(), + "https://auth.openai.com/api/accounts" + ); + assert_eq!( + ChatGptEnvironment::from_chatgpt_base_url("https://chatgpt-staging.com/backend-api")?, + ChatGptEnvironment::Staging + ); + assert_eq!( + ChatGptEnvironment::Staging.agent_identity_authapi_base_url(), + "https://auth.api.openai.org/api/accounts" + ); + Ok(()) + } + + #[test] + fn chatgpt_environment_rejects_custom_urls() { + assert!(ChatGptEnvironment::from_chatgpt_base_url("http://localhost:8080").is_err(),); + } + + #[test] + fn agent_registration_url_appends_to_authapi_base_url() { + assert_eq!( + agent_registration_url("https://auth.openai.com/api/accounts"), + "https://auth.openai.com/api/accounts/v1/agent/register" + ); + assert_eq!( + agent_registration_url("http://localhost:8080"), + "http://localhost:8080/v1/agent/register" + ); + assert_eq!( + agent_registration_url("http://localhost:8080/backend-api"), + "http://localhost:8080/backend-api/v1/agent/register" + ); + } + + #[test] + fn agent_task_registration_url_appends_to_authapi_base_url() { + assert_eq!( + agent_task_registration_url("https://auth.openai.com/api/accounts", "agent-runtime-id"), + "https://auth.openai.com/api/accounts/v1/agent/agent-runtime-id/task/register" + ); + assert_eq!( + agent_task_registration_url( + "https://auth.openai.com/api/accounts/", + "agent-runtime-id" + ), + "https://auth.openai.com/api/accounts/v1/agent/agent-runtime-id/task/register" + ); + assert_eq!( + agent_task_registration_url("http://localhost:8080", "agent-runtime-id"), + "http://localhost:8080/v1/agent/agent-runtime-id/task/register" + ); + } + + #[test] + fn retryable_registration_error_accepts_429_and_5xx() { + let too_many_requests = anyhow::Error::new(AgentIdentityRegistrationHttpError::new( + "agent registration", + StatusCode::TOO_MANY_REQUESTS, + "rate limited".to_string(), + )); + let unavailable = anyhow::Error::new(AgentIdentityRegistrationHttpError::new( + "agent registration", + StatusCode::SERVICE_UNAVAILABLE, + "try later".to_string(), + )); + + assert!(is_retryable_registration_error(&too_many_requests)); + assert!(is_retryable_registration_error(&unavailable)); + } + + #[test] + fn retryable_registration_error_rejects_hard_failures() { + let forbidden = anyhow::Error::new(AgentIdentityRegistrationHttpError::new( + "agent registration", + StatusCode::FORBIDDEN, + "not allowed".to_string(), + )); + let malformed = anyhow::anyhow!("failed to sign registration request"); + + assert!(!is_retryable_registration_error(&forbidden)); + assert!(!is_retryable_registration_error(&malformed)); + } + + #[test] + fn agent_identity_jwks_url_uses_agent_identity_jwt_route() { + assert_eq!( + agent_identity_jwks_url("https://chatgpt.com/backend-api"), + "https://chatgpt.com/backend-api/wham/agent-identities/jwks" + ); + assert_eq!( + agent_identity_jwks_url("https://chatgpt.com/backend-api/"), + "https://chatgpt.com/backend-api/wham/agent-identities/jwks" + ); + } + + #[test] + fn agent_identity_jwks_url_uses_jwt_issuer_base_url() { + assert_eq!( + agent_identity_jwks_url("http://localhost:8080/api/codex"), + "http://localhost:8080/api/codex/agent-identities/jwks" + ); + assert_eq!( + agent_identity_jwks_url("http://localhost:8080/api/codex/"), + "http://localhost:8080/api/codex/agent-identities/jwks" + ); + } + + fn jwt_with_payload(payload: serde_json::Value) -> String { + let encode = |bytes: &[u8]| URL_SAFE_NO_PAD.encode(bytes); + let header_b64 = encode(br#"{"alg":"none","typ":"JWT"}"#); + let payload_b64 = encode(&serde_json::to_vec(&payload).expect("payload should serialize")); + let signature_b64 = encode(b"sig"); + format!("{header_b64}.{payload_b64}.{signature_b64}") + } +} diff --git a/codex-rs/ansi-escape/src/lib.rs b/codex-rs/ansi-escape/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..b47cf14f8ea672785931bb7d5d3d5982fd227aaf --- /dev/null +++ b/codex-rs/ansi-escape/src/lib.rs @@ -0,0 +1,58 @@ +use ansi_to_tui::Error; +use ansi_to_tui::IntoText; +use ratatui::text::Line; +use ratatui::text::Text; + +// Expand tabs in a best-effort way for transcript rendering. +// Tabs can interact poorly with left-gutter prefixes in our TUI and CLI +// transcript views (e.g., `nl` separates line numbers from content with a tab). +// Replacing tabs with spaces avoids odd visual artifacts without changing +// semantics for our use cases. +fn expand_tabs(s: &str) -> std::borrow::Cow<'_, str> { + if s.contains('\t') { + // Keep it simple: replace each tab with 4 spaces. + // We do not try to align to tab stops since most usages (like `nl`) + // look acceptable with a fixed substitution and this avoids stateful math + // across spans. + std::borrow::Cow::Owned(s.replace('\t', " ")) + } else { + std::borrow::Cow::Borrowed(s) + } +} + +/// This function should be used when the contents of `s` are expected to match +/// a single line. If multiple lines are found, a warning is logged and only the +/// first line is returned. +pub fn ansi_escape_line(s: &str) -> Line<'static> { + // Normalize tabs to spaces to avoid odd gutter collisions in transcript mode. + let s = expand_tabs(s); + let text = ansi_escape(&s); + match text.lines.as_slice() { + [] => "".into(), + [only] => only.clone(), + [first, rest @ ..] => { + tracing::warn!("ansi_escape_line: expected a single line, got {first:?} and {rest:?}"); + first.clone() + } + } +} + +pub fn ansi_escape(s: &str) -> Text<'static> { + // to_text() claims to be faster, but introduces complex lifetime issues + // such that it's not worth it. + match s.into_text() { + Ok(text) => text, + Err(err) => match err { + Error::NomError(message) => { + tracing::error!( + "ansi_to_tui NomError docs claim should never happen when parsing `{s}`: {message}" + ); + panic!(); + } + Error::Utf8Error(utf8error) => { + tracing::error!("Utf8Error: {utf8error}"); + panic!(); + } + }, + } +} diff --git a/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst b/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst new file mode 100644 index 0000000000000000000000000000000000000000..b157722b0fe54b43fc0bea2860a5bab40fc87e29 --- /dev/null +++ b/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d6fb74c9af068055ac5e369ce116c27299e46ad2a0487541a7abbd52761f9832 +size 157496 diff --git a/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst b/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst new file mode 100644 index 0000000000000000000000000000000000000000..112166ea68828d7c7f59c4d9f43f5badfe53e17a --- /dev/null +++ b/codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be380114e24a7541e07df1af9b1d2d10f8749f1d97516496a7e9c59bd395c29d +size 151569 diff --git a/codex-rs/app-server-transport/src/connection_auth.rs b/codex-rs/app-server-transport/src/connection_auth.rs new file mode 100644 index 0000000000000000000000000000000000000000..afa7b4f3c7df7330128e489fb2c2b1a6bc6f41c5 --- /dev/null +++ b/codex-rs/app-server-transport/src/connection_auth.rs @@ -0,0 +1,51 @@ +//! Binds a transport connection to the authentication owner that established it. +//! Owner revisions invalidate queued work even before transport closure is delivered. + +use codex_login::AuthChangeState; +use codex_login::AuthManager; +use std::io; +use tokio::sync::watch; + +#[derive(Clone, Debug)] +pub struct ConnectionAuth { + changes: watch::Receiver, + owner_generation: u64, +} + +impl ConnectionAuth { + pub(crate) fn capture(auth_manager: &AuthManager) -> Self { + let changes = auth_manager.auth_change_state_receiver(); + let owner_generation = changes.borrow().owner_generation; + Self::new(changes, owner_generation) + } + + pub(crate) fn ensure_current(&self) -> io::Result<()> { + if self.is_current() { + Ok(()) + } else { + Err(io::Error::new( + io::ErrorKind::Interrupted, + "remote control authentication changed", + )) + } + } + + pub(crate) fn new(changes: watch::Receiver, owner_generation: u64) -> Self { + Self { + changes, + owner_generation, + } + } + + pub fn is_current(&self) -> bool { + self.changes.borrow().owner_generation == self.owner_generation + && self.changes.has_changed().is_ok() + } + + pub(crate) async fn invalidated(&self) { + let mut changes = self.changes.clone(); + let _ = changes + .wait_for(|state| state.owner_generation != self.owner_generation) + .await; + } +} diff --git a/codex-rs/app-server-transport/src/daemon_recovery.rs b/codex-rs/app-server-transport/src/daemon_recovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..c19c51d8a8df9bb663340662dd260cabc763cefd --- /dev/null +++ b/codex-rs/app-server-transport/src/daemon_recovery.rs @@ -0,0 +1,82 @@ +//! Shared on-disk candidate set for managed daemon restarts. + +use std::collections::BTreeMap; +use std::collections::BTreeSet; +use std::io; +use std::path::Path; + +use codex_core::path_utils::write_atomically; +use serde::Deserialize; +use serde::Serialize; + +#[derive(Debug, Default, Deserialize, Serialize, PartialEq, Eq)] +pub struct RecoverySnapshot { + #[serde(skip)] + pub loaded: BTreeSet, + pub interrupted: BTreeMap, +} + +#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)] +pub struct InterruptedTurn { + pub turn_id: String, + pub output_schema: Option, + pub service_tier: Option, + pub cyber_access_program: Option, + /// Only local execution with thread-owned configuration can continue automatically. + /// Older snapshots without this identity are reloaded without continuation. + pub local_environment: Option, +} + +// Old servers accept the array and skip this non-thread entry during best-effort +// restoration. Keeping metadata in the same atomic file avoids stale sidecars. +const INTERRUPTION_PREFIX: &str = "codex-interrupted-v1:"; + +pub fn read_snapshot(path: &Path) -> io::Result { + let mut loaded: BTreeSet = match std::fs::read(path) { + Ok(contents) => serde_json::from_slice(&contents).map_err(io::Error::other)?, + Err(err) if err.kind() == io::ErrorKind::NotFound => return Ok(RecoverySnapshot::default()), + Err(err) => return Err(err), + }; + let mut snapshot = RecoverySnapshot::default(); + loaded.retain(|entry| { + if let Some(metadata) = entry.strip_prefix(INTERRUPTION_PREFIX) { + if let Ok(saved) = serde_json::from_str::(metadata) { + snapshot = saved; + } + false + } else { + true + } + }); + snapshot.interrupted.retain(|id, _| loaded.contains(id)); + snapshot.loaded = loaded; + Ok(snapshot) +} + +pub fn read_candidates(path: &Path) -> io::Result> { + Ok(read_snapshot(path)?.loaded) +} + +pub fn write_candidates(path: &Path, candidates: &BTreeSet) -> io::Result<()> { + write_snapshot( + path, + &RecoverySnapshot { + loaded: candidates.clone(), + ..Default::default() + }, + ) +} + +pub fn write_snapshot(path: &Path, snapshot: &RecoverySnapshot) -> io::Result<()> { + let mut saved = snapshot.loaded.clone(); + if !snapshot.interrupted.is_empty() { + saved.insert(format!( + "{INTERRUPTION_PREFIX}{}", + serde_json::to_string(snapshot).map_err(io::Error::other)? + )); + } + write_atomically( + path, + &serde_json::to_string(&saved).map_err(io::Error::other)?, + ) +} diff --git a/codex-rs/app-server-transport/src/daemon_shutdown.rs b/codex-rs/app-server-transport/src/daemon_shutdown.rs new file mode 100644 index 0000000000000000000000000000000000000000..49e4551bb5b9ab12ba3954811f7097606db9aba0 --- /dev/null +++ b/codex-rs/app-server-transport/src/daemon_shutdown.rs @@ -0,0 +1,53 @@ +//! The detached Windows updater has no control socket, so its owner requests +//! termination through a file in its private state directory. Managed app-server +//! shutdown instead uses the protected local socket. + +use std::io; +use std::io::Read; +use std::path::Path; +#[cfg(windows)] +use std::path::PathBuf; +#[cfg(windows)] +use std::time::Duration; + +#[cfg(windows)] +pub const DAEMON_SHUTDOWN_FILE_ENV: &str = "CODEX_DAEMON_SHUTDOWN_FILE"; + +/// Waits for and consumes one updater shutdown request. +#[cfg(windows)] +pub async fn daemon_shutdown_signal() -> io::Result<()> { + let Some(path) = std::env::var_os(DAEMON_SHUTDOWN_FILE_ENV).map(PathBuf::from) else { + return std::future::pending().await; + }; + loop { + if take_shutdown_request(&path, std::process::id())? { + return Ok(()); + } + tokio::time::sleep(Duration::from_millis(50)).await; + } +} + +fn take_shutdown_request(path: &Path, pid: u32) -> io::Result { + let file = match std::fs::File::open(path) { + Ok(file) => file, + Err(err) if err.kind() == io::ErrorKind::NotFound => return Ok(false), + Err(err) => return Err(err), + }; + // Descendant app-servers may inherit the control path. Only the intended + // process can consume a request. Bound reads to a u32 PID plus one extra byte. + let mut contents = Vec::new(); + file.take(/*limit*/ 11).read_to_end(&mut contents)?; + if contents != pid.to_string().as_bytes() { + return Ok(false); + } + // Synchronous consumption cannot lose a request to select cancellation. + match std::fs::remove_file(path) { + Ok(()) => Ok(true), + Err(err) if err.kind() == io::ErrorKind::NotFound => Ok(false), + Err(err) => Err(err), + } +} + +#[cfg(test)] +#[path = "daemon_shutdown_tests.rs"] +mod tests; diff --git a/codex-rs/app-server-transport/src/daemon_shutdown_tests.rs b/codex-rs/app-server-transport/src/daemon_shutdown_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..7454dffcd1651c0ac91f0b0b16a4dd88c57eda5e --- /dev/null +++ b/codex-rs/app-server-transport/src/daemon_shutdown_tests.rs @@ -0,0 +1,22 @@ +use super::take_shutdown_request; + +#[test] +fn inherited_control_path_cannot_consume_another_process_request() { + let directory = tempfile::tempdir().expect("directory"); + let request = directory.path().join("shutdown"); + std::fs::write(&request, "1234").expect("request"); + assert!(!take_shutdown_request(&request, /*pid*/ 5678).expect("descendant probe")); + assert!(take_shutdown_request(&request, /*pid*/ 1234).expect("parent probe")); + assert!(!take_shutdown_request(&request, /*pid*/ 1234).expect("already consumed")); +} + +#[test] +fn incomplete_or_invalid_requests_are_not_consumed() { + let directory = tempfile::tempdir().expect("directory"); + let request = directory.path().join("shutdown"); + for contents in [b"".as_slice(), b"123", b"1234junk", &[0xff; 20]] { + std::fs::write(&request, contents).expect("request"); + assert!(!take_shutdown_request(&request, /*pid*/ 1234).expect("probe")); + assert!(request.exists()); + } +} diff --git a/codex-rs/app-server-transport/src/lib.rs b/codex-rs/app-server-transport/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..6eacd60a8b0c47a2b929a6b190daece03ae35a95 --- /dev/null +++ b/codex-rs/app-server-transport/src/lib.rs @@ -0,0 +1,45 @@ +pub mod daemon_recovery; +#[cfg(any(windows, test))] +mod daemon_shutdown; +#[cfg(windows)] +pub use daemon_shutdown::DAEMON_SHUTDOWN_FILE_ENV; +#[cfg(windows)] +pub use daemon_shutdown::daemon_shutdown_signal; +/// Only managed app-server launches accept the local socket shutdown request. +pub const DAEMON_SHUTDOWN_SOCKET_ENV: &str = "CODEX_DAEMON_SHUTDOWN_SOCKET"; +mod connection_auth; +mod outgoing_message; +mod transport; + +pub use connection_auth::ConnectionAuth; +pub use outgoing_message::ConnectionId; +pub use outgoing_message::OutgoingError; +pub use outgoing_message::OutgoingMessage; +pub use outgoing_message::OutgoingResponse; +pub use outgoing_message::QueuedOutgoingMessage; +pub use transport::AppServerStartupLock; +pub use transport::AppServerTransport; +pub use transport::AppServerTransportParseError; +pub use transport::CHANNEL_CAPACITY; +pub use transport::ConnectionOrigin; +pub use transport::DaemonShutdownAccess; +pub use transport::REMOTE_CONTROL_DISABLED_ENV_VAR; +pub use transport::RemoteControlDisabledByRequirements; +pub use transport::RemoteControlEnableError; +pub use transport::RemoteControlHandle; +pub use transport::RemoteControlPolicy; +pub use transport::RemoteControlStartConfig; +pub use transport::RemoteControlStartupMode; +pub use transport::RemoteControlUnavailable; +pub use transport::TransportEvent; +pub use transport::acquire_app_server_startup_lock; +pub use transport::app_server_control_socket_path; +pub use transport::app_server_startup_lock_path; +pub use transport::auth; +pub use transport::daemon_recovery_file_path; +pub use transport::prepare_control_socket_path; +pub use transport::start_control_socket_acceptor; +pub use transport::start_remote_control; +pub use transport::start_stdio_connection; +pub use transport::start_websocket_acceptor; +pub use transport::take_remote_control_disabled_env; diff --git a/codex-rs/app-server-transport/src/outgoing_message.rs b/codex-rs/app-server-transport/src/outgoing_message.rs new file mode 100644 index 0000000000000000000000000000000000000000..7f60ceb44bd264af8255e48aa757096ff891a0cc --- /dev/null +++ b/codex-rs/app-server-transport/src/outgoing_message.rs @@ -0,0 +1,59 @@ +use std::fmt; + +use codex_app_server_protocol::ClientResponsePayload; +use codex_app_server_protocol::JSONRPCErrorError; +use codex_app_server_protocol::RequestId; +use codex_app_server_protocol::ServerNotificationEnvelope; +use codex_app_server_protocol::ServerRequest; +use serde::Serialize; +use tokio::sync::oneshot; + +/// Stable identifier for a transport connection. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct ConnectionId(pub u64); + +impl fmt::Display for ConnectionId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + +/// Outgoing message from the server to the client. +#[derive(Debug, Clone, Serialize)] +#[serde(untagged)] +#[allow(clippy::large_enum_variant)] +pub enum OutgoingMessage { + Request(ServerRequest), + /// AppServerNotification is specific to the case where this is run as an + /// "app server" as opposed to an MCP server. + AppServerNotification(ServerNotificationEnvelope), + Response(OutgoingResponse), + Error(OutgoingError), +} + +#[derive(Debug, Clone, Serialize)] +pub struct OutgoingResponse { + pub id: RequestId, + pub result: Box, +} + +#[derive(Debug, Clone, PartialEq, Serialize)] +pub struct OutgoingError { + pub error: JSONRPCErrorError, + pub id: RequestId, +} + +#[derive(Debug)] +pub struct QueuedOutgoingMessage { + pub message: OutgoingMessage, + pub write_complete_tx: Option>, +} + +impl QueuedOutgoingMessage { + pub fn new(message: OutgoingMessage) -> Self { + Self { + message, + write_complete_tx: None, + } + } +} diff --git a/codex-rs/app-server-transport/src/transport/auth.rs b/codex-rs/app-server-transport/src/transport/auth.rs new file mode 100644 index 0000000000000000000000000000000000000000..eeccf21ed30a598835694d7b89d726d6db445ff0 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/auth.rs @@ -0,0 +1,751 @@ +use anyhow::Context; +use axum::http::HeaderMap; +use axum::http::StatusCode; +use axum::http::header::AUTHORIZATION; +use clap::Args; +use clap::ValueEnum; +use codex_utils_absolute_path::AbsolutePathBuf; +use constant_time_eq::constant_time_eq_32; +use jsonwebtoken::Algorithm; +use jsonwebtoken::DecodingKey; +use jsonwebtoken::Validation; +use jsonwebtoken::decode; +use serde::Deserialize; +use sha2::Digest; +use sha2::Sha256; +use std::io; +use std::io::ErrorKind; +use std::net::SocketAddr; +use std::path::Path; +use std::path::PathBuf; +use time::OffsetDateTime; + +const DEFAULT_MAX_CLOCK_SKEW_SECONDS: u64 = 30; +const MIN_SIGNED_BEARER_SECRET_BYTES: usize = 32; +const INVALID_AUTHORIZATION_HEADER_MESSAGE: &str = "invalid authorization header"; + +#[derive(Debug, Clone, Default, PartialEq, Eq, Args)] +pub struct AppServerWebsocketAuthArgs { + /// Websocket auth mode for non-loopback listeners. + #[arg(long = "ws-auth", value_name = "MODE", value_enum)] + pub ws_auth: Option, + + /// Absolute path to the capability-token file. + #[arg(long = "ws-token-file", value_name = "PATH")] + pub ws_token_file: Option, + + /// Hex-encoded SHA-256 digest of the capability token. + #[arg(long = "ws-token-sha256", value_name = "HEX")] + pub ws_token_sha256: Option, + + /// Absolute path to the shared secret file for signed JWT bearer tokens. + #[arg(long = "ws-shared-secret-file", value_name = "PATH")] + pub ws_shared_secret_file: Option, + + /// Expected issuer for signed JWT bearer tokens. + #[arg(long = "ws-issuer", value_name = "ISSUER")] + pub ws_issuer: Option, + + /// Expected audience for signed JWT bearer tokens. + #[arg(long = "ws-audience", value_name = "AUDIENCE")] + pub ws_audience: Option, + + /// Maximum clock skew when validating signed JWT bearer tokens. + #[arg(long = "ws-max-clock-skew-seconds", value_name = "SECONDS")] + pub ws_max_clock_skew_seconds: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, ValueEnum)] +pub enum WebsocketAuthCliMode { + CapabilityToken, + SignedBearerToken, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct AppServerWebsocketAuthSettings { + pub config: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum AppServerWebsocketAuthConfig { + CapabilityToken { + source: AppServerWebsocketCapabilityTokenSource, + }, + SignedBearerToken { + shared_secret_file: AbsolutePathBuf, + issuer: Option, + audience: Option, + max_clock_skew_seconds: u64, + }, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum AppServerWebsocketCapabilityTokenSource { + TokenFile { token_file: AbsolutePathBuf }, + TokenSha256 { token_sha256: [u8; 32] }, +} + +#[derive(Clone, Debug, Default)] +pub struct WebsocketAuthPolicy { + pub(crate) mode: Option, +} + +#[derive(Clone, Debug)] +pub(crate) enum WebsocketAuthMode { + CapabilityToken { + token_sha256: [u8; 32], + }, + SignedBearerToken { + shared_secret: Vec, + issuer: Option, + audience: Option, + max_clock_skew_seconds: i64, + }, +} + +#[derive(Debug)] +pub(crate) struct WebsocketAuthError { + status_code: StatusCode, + message: &'static str, +} + +#[derive(Deserialize)] +struct JwtClaims { + exp: i64, + nbf: Option, + iss: Option, + aud: Option, +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum JwtAudienceClaim { + Single(String), + Multiple(Vec), +} + +impl WebsocketAuthError { + pub(crate) fn status_code(&self) -> StatusCode { + self.status_code + } + + pub(crate) fn message(&self) -> &'static str { + self.message + } +} + +impl AppServerWebsocketAuthArgs { + pub fn try_into_settings(self) -> anyhow::Result { + let normalize = |value: Option| { + value.and_then(|value| { + let trimmed = value.trim(); + (!trimmed.is_empty()).then(|| trimmed.to_string()) + }) + }; + + let config = match self.ws_auth { + Some(WebsocketAuthCliMode::CapabilityToken) => { + if self.ws_shared_secret_file.is_some() + || self.ws_issuer.is_some() + || self.ws_audience.is_some() + || self.ws_max_clock_skew_seconds.is_some() + { + anyhow::bail!( + "`--ws-shared-secret-file`, `--ws-issuer`, `--ws-audience`, and `--ws-max-clock-skew-seconds` require `--ws-auth signed-bearer-token`" + ); + } + let source = match (self.ws_token_file, self.ws_token_sha256) { + (Some(_), Some(_)) => { + anyhow::bail!( + "`--ws-token-file` and `--ws-token-sha256` are mutually exclusive" + ); + } + (Some(token_file), None) => { + AppServerWebsocketCapabilityTokenSource::TokenFile { + token_file: absolute_path_arg("--ws-token-file", token_file)?, + } + } + (None, Some(token_sha256)) => { + AppServerWebsocketCapabilityTokenSource::TokenSha256 { + token_sha256: sha256_digest_arg("--ws-token-sha256", &token_sha256)?, + } + } + (None, None) => { + anyhow::bail!( + "`--ws-token-file` or `--ws-token-sha256` is required when `--ws-auth capability-token` is set" + ); + } + }; + Some(AppServerWebsocketAuthConfig::CapabilityToken { source }) + } + Some(WebsocketAuthCliMode::SignedBearerToken) => { + if self.ws_token_file.is_some() || self.ws_token_sha256.is_some() { + anyhow::bail!( + "`--ws-token-file` and `--ws-token-sha256` require `--ws-auth capability-token`, not `signed-bearer-token`" + ); + } + let shared_secret_file = self.ws_shared_secret_file.context( + "`--ws-shared-secret-file` is required when `--ws-auth signed-bearer-token` is set", + )?; + Some(AppServerWebsocketAuthConfig::SignedBearerToken { + shared_secret_file: absolute_path_arg( + "--ws-shared-secret-file", + shared_secret_file, + )?, + issuer: normalize(self.ws_issuer), + audience: normalize(self.ws_audience), + max_clock_skew_seconds: self + .ws_max_clock_skew_seconds + .unwrap_or(DEFAULT_MAX_CLOCK_SKEW_SECONDS), + }) + } + None => { + if self.ws_token_file.is_some() + || self.ws_token_sha256.is_some() + || self.ws_shared_secret_file.is_some() + || self.ws_issuer.is_some() + || self.ws_audience.is_some() + || self.ws_max_clock_skew_seconds.is_some() + { + anyhow::bail!( + "websocket auth flags require `--ws-auth capability-token` or `--ws-auth signed-bearer-token`" + ); + } + None + } + }; + + Ok(AppServerWebsocketAuthSettings { config }) + } +} + +pub fn policy_from_settings( + settings: &AppServerWebsocketAuthSettings, +) -> io::Result { + let mode = match settings.config.as_ref() { + Some(AppServerWebsocketAuthConfig::CapabilityToken { source }) => match source { + AppServerWebsocketCapabilityTokenSource::TokenFile { token_file } => { + let token = read_trimmed_secret(token_file.as_ref())?; + Some(WebsocketAuthMode::CapabilityToken { + token_sha256: sha256_digest(token.as_bytes()), + }) + } + AppServerWebsocketCapabilityTokenSource::TokenSha256 { token_sha256 } => { + Some(WebsocketAuthMode::CapabilityToken { + token_sha256: *token_sha256, + }) + } + }, + Some(AppServerWebsocketAuthConfig::SignedBearerToken { + shared_secret_file, + issuer, + audience, + max_clock_skew_seconds, + }) => { + let shared_secret = read_trimmed_secret(shared_secret_file.as_ref())?.into_bytes(); + validate_signed_bearer_secret(shared_secret_file.as_ref(), &shared_secret)?; + let max_clock_skew_seconds = i64::try_from(*max_clock_skew_seconds).map_err(|_| { + io::Error::new( + ErrorKind::InvalidInput, + "websocket auth clock skew must fit in a signed 64-bit integer", + ) + })?; + Some(WebsocketAuthMode::SignedBearerToken { + shared_secret, + issuer: issuer.clone(), + audience: audience.clone(), + max_clock_skew_seconds, + }) + } + None => None, + }; + + Ok(WebsocketAuthPolicy { mode }) +} + +pub(crate) fn is_unauthenticated_non_loopback_listener( + bind_address: SocketAddr, + policy: &WebsocketAuthPolicy, +) -> bool { + !bind_address.ip().is_loopback() && policy.mode.is_none() +} + +pub(crate) fn authorize_upgrade( + headers: &HeaderMap, + policy: &WebsocketAuthPolicy, +) -> Result<(), WebsocketAuthError> { + let Some(mode) = policy.mode.as_ref() else { + return Ok(()); + }; + + let token = bearer_token_from_headers(headers)?; + match mode { + WebsocketAuthMode::CapabilityToken { token_sha256 } => { + let actual_sha256 = sha256_digest(token.as_bytes()); + if constant_time_eq_32(token_sha256, &actual_sha256) { + Ok(()) + } else { + Err(unauthorized("invalid websocket bearer token")) + } + } + WebsocketAuthMode::SignedBearerToken { + shared_secret, + issuer, + audience, + max_clock_skew_seconds, + } => verify_signed_bearer_token( + token, + shared_secret, + issuer.as_deref(), + audience.as_deref(), + *max_clock_skew_seconds, + ), + } +} + +fn verify_signed_bearer_token( + token: &str, + shared_secret: &[u8], + issuer: Option<&str>, + audience: Option<&str>, + max_clock_skew_seconds: i64, +) -> Result<(), WebsocketAuthError> { + let claims = decode_jwt_claims(token, shared_secret)?; + validate_jwt_claims(&claims, issuer, audience, max_clock_skew_seconds) +} + +fn decode_jwt_claims(token: &str, shared_secret: &[u8]) -> Result { + let mut validation = Validation::new(Algorithm::HS256); + validation.required_spec_claims.clear(); + validation.validate_exp = false; + validation.validate_nbf = false; + validation.validate_aud = false; + + decode::(token, &DecodingKey::from_secret(shared_secret), &validation) + .map(|token_data| token_data.claims) + .map_err(|_| unauthorized("invalid websocket jwt")) +} + +fn validate_jwt_claims( + claims: &JwtClaims, + issuer: Option<&str>, + audience: Option<&str>, + max_clock_skew_seconds: i64, +) -> Result<(), WebsocketAuthError> { + let now = OffsetDateTime::now_utc().unix_timestamp(); + if now > claims.exp.saturating_add(max_clock_skew_seconds) { + return Err(unauthorized("expired websocket jwt")); + } + if let Some(nbf) = claims.nbf + && now < nbf.saturating_sub(max_clock_skew_seconds) + { + return Err(unauthorized("websocket jwt is not valid yet")); + } + if let Some(expected_issuer) = issuer + && claims.iss.as_deref() != Some(expected_issuer) + { + return Err(unauthorized("websocket jwt issuer mismatch")); + } + if let Some(expected_audience) = audience + && !audience_matches(claims.aud.as_ref(), expected_audience) + { + return Err(unauthorized("websocket jwt audience mismatch")); + } + + Ok(()) +} + +fn audience_matches(audience: Option<&JwtAudienceClaim>, expected_audience: &str) -> bool { + match audience { + Some(JwtAudienceClaim::Single(actual)) => actual == expected_audience, + Some(JwtAudienceClaim::Multiple(actual)) => { + actual.iter().any(|audience| audience == expected_audience) + } + None => false, + } +} + +fn bearer_token_from_headers(headers: &HeaderMap) -> Result<&str, WebsocketAuthError> { + let raw_header = headers + .get(AUTHORIZATION) + .ok_or_else(|| unauthorized("missing websocket bearer token"))?; + let header = raw_header + .to_str() + .map_err(|_| unauthorized(INVALID_AUTHORIZATION_HEADER_MESSAGE))?; + let Some((scheme, token)) = header.split_once(' ') else { + return Err(unauthorized(INVALID_AUTHORIZATION_HEADER_MESSAGE)); + }; + if !scheme.eq_ignore_ascii_case("Bearer") { + return Err(unauthorized(INVALID_AUTHORIZATION_HEADER_MESSAGE)); + } + let token = token.trim(); + if token.is_empty() { + return Err(unauthorized(INVALID_AUTHORIZATION_HEADER_MESSAGE)); + } + Ok(token) +} + +fn validate_signed_bearer_secret(path: &Path, shared_secret: &[u8]) -> io::Result<()> { + if shared_secret.len() < MIN_SIGNED_BEARER_SECRET_BYTES { + return Err(io::Error::new( + ErrorKind::InvalidInput, + format!( + "signed websocket bearer secret {} must be at least {MIN_SIGNED_BEARER_SECRET_BYTES} bytes", + path.display() + ), + )); + } + Ok(()) +} + +fn read_trimmed_secret(path: &std::path::Path) -> io::Result { + let raw = std::fs::read_to_string(path).map_err(|err| { + io::Error::new( + err.kind(), + format!( + "failed to read websocket auth secret {}: {err}", + path.display() + ), + ) + })?; + let trimmed = raw.trim(); + if trimmed.is_empty() { + return Err(io::Error::new( + ErrorKind::InvalidInput, + format!("websocket auth secret {} must not be empty", path.display()), + )); + } + Ok(trimmed.to_string()) +} + +fn absolute_path_arg(flag_name: &str, path: PathBuf) -> anyhow::Result { + AbsolutePathBuf::try_from(path).with_context(|| format!("{flag_name} must be an absolute path")) +} + +fn sha256_digest_arg(flag_name: &str, value: &str) -> anyhow::Result<[u8; 32]> { + let trimmed = value.trim(); + if trimmed.len() != 64 { + anyhow::bail!("{flag_name} must be a 64-character hex SHA-256 digest"); + } + + let mut digest = [0u8; 32]; + for (index, pair) in trimmed.as_bytes().chunks_exact(2).enumerate() { + let high = hex_nibble(flag_name, pair[0])?; + let low = hex_nibble(flag_name, pair[1])?; + digest[index] = (high << 4) | low; + } + Ok(digest) +} + +fn hex_nibble(flag_name: &str, byte: u8) -> anyhow::Result { + match byte { + b'0'..=b'9' => Ok(byte - b'0'), + b'a'..=b'f' => Ok(byte - b'a' + 10), + b'A'..=b'F' => Ok(byte - b'A' + 10), + _ => anyhow::bail!("{flag_name} must be a 64-character hex SHA-256 digest"), + } +} + +fn sha256_digest(input: &[u8]) -> [u8; 32] { + let mut digest = [0u8; 32]; + digest.copy_from_slice(&Sha256::digest(input)); + digest +} + +fn unauthorized(message: &'static str) -> WebsocketAuthError { + WebsocketAuthError { + status_code: StatusCode::UNAUTHORIZED, + message, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use axum::http::HeaderValue; + use base64::Engine; + use base64::engine::general_purpose::URL_SAFE_NO_PAD; + use hmac::Hmac; + use hmac::Mac; + use serde_json::json; + + type HmacSha256 = Hmac; + + fn signed_token(shared_secret: &[u8], claims: serde_json::Value) -> String { + let header = URL_SAFE_NO_PAD.encode(br#"{"alg":"HS256","typ":"JWT"}"#); + let claims_segment = URL_SAFE_NO_PAD.encode(serde_json::to_vec(&claims).unwrap()); + let payload = format!("{header}.{claims_segment}"); + let mut mac = HmacSha256::new_from_slice(shared_secret).unwrap(); + mac.update(payload.as_bytes()); + let signature = URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes()); + format!("{payload}.{signature}") + } + + #[test] + fn detects_unauthenticated_non_loopback_listener() { + let policy = WebsocketAuthPolicy::default(); + assert!(is_unauthenticated_non_loopback_listener( + "0.0.0.0:8765".parse().unwrap(), + &policy, + )); + assert!(!is_unauthenticated_non_loopback_listener( + "127.0.0.1:8765".parse().unwrap(), + &policy, + )); + assert!(!is_unauthenticated_non_loopback_listener( + "0.0.0.0:8765".parse().unwrap(), + &WebsocketAuthPolicy { + mode: Some(WebsocketAuthMode::CapabilityToken { + token_sha256: [0u8; 32], + }), + }, + )); + } + + #[test] + fn capability_token_args_require_token_file_or_hash() { + let err = AppServerWebsocketAuthArgs { + ws_auth: Some(WebsocketAuthCliMode::CapabilityToken), + ..Default::default() + } + .try_into_settings() + .expect_err("capability-token mode should require a token source"); + assert!( + err.to_string().contains("--ws-token-file") + && err.to_string().contains("--ws-token-sha256"), + "unexpected error: {err}" + ); + } + + #[test] + fn capability_token_args_accept_token_hash() { + let settings = AppServerWebsocketAuthArgs { + ws_auth: Some(WebsocketAuthCliMode::CapabilityToken), + ws_token_sha256: Some("ab".repeat(32)), + ..Default::default() + } + .try_into_settings() + .expect("capability-token hash args should parse"); + + assert_eq!( + settings, + AppServerWebsocketAuthSettings { + config: Some(AppServerWebsocketAuthConfig::CapabilityToken { + source: AppServerWebsocketCapabilityTokenSource::TokenSha256 { + token_sha256: [0xab; 32], + }, + }), + } + ); + } + + #[test] + fn capability_token_args_reject_multiple_token_sources() { + let err = AppServerWebsocketAuthArgs { + ws_auth: Some(WebsocketAuthCliMode::CapabilityToken), + ws_token_file: Some(PathBuf::from("/tmp/token")), + ws_token_sha256: Some("ab".repeat(32)), + ..Default::default() + } + .try_into_settings() + .expect_err("capability-token mode should reject multiple token sources"); + assert!( + err.to_string().contains("mutually exclusive"), + "unexpected error: {err}" + ); + } + + #[test] + fn capability_token_args_reject_malformed_token_hash() { + let err = AppServerWebsocketAuthArgs { + ws_auth: Some(WebsocketAuthCliMode::CapabilityToken), + ws_token_sha256: Some("not-a-sha256".to_string()), + ..Default::default() + } + .try_into_settings() + .expect_err("capability-token mode should reject malformed token hashes"); + assert!( + err.to_string().contains("64-character hex"), + "unexpected error: {err}" + ); + } + + #[test] + fn capability_token_hash_policy_authorizes_matching_bearer_token() { + let settings = AppServerWebsocketAuthSettings { + config: Some(AppServerWebsocketAuthConfig::CapabilityToken { + source: AppServerWebsocketCapabilityTokenSource::TokenSha256 { + token_sha256: sha256_digest(b"super-secret-token"), + }, + }), + }; + let policy = policy_from_settings(&settings).expect("hash policy should build"); + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_static("Bearer super-secret-token"), + ); + authorize_upgrade(&headers, &policy).expect("matching token should authorize"); + + headers.insert( + AUTHORIZATION, + HeaderValue::from_static("Bearer wrong-token"), + ); + let err = authorize_upgrade(&headers, &policy).expect_err("wrong token should fail"); + assert_eq!(err.status_code(), StatusCode::UNAUTHORIZED); + } + + #[test] + fn signed_bearer_args_require_mode_when_mode_specific_flags_are_set() { + let err = AppServerWebsocketAuthArgs { + ws_shared_secret_file: Some(PathBuf::from("/tmp/secret")), + ..Default::default() + } + .try_into_settings() + .expect_err("mode-specific flags should require --ws-auth"); + assert!( + err.to_string().contains("websocket auth flags require"), + "unexpected error: {err}" + ); + } + + #[test] + fn signed_bearer_args_default_clock_skew_and_trim_optional_claims() { + let settings = AppServerWebsocketAuthArgs { + ws_auth: Some(WebsocketAuthCliMode::SignedBearerToken), + ws_shared_secret_file: Some(PathBuf::from("/tmp/secret")), + ws_issuer: Some(" issuer ".to_string()), + ws_audience: Some(" ".to_string()), + ..Default::default() + } + .try_into_settings() + .expect("signed bearer args should parse"); + + assert_eq!( + settings, + AppServerWebsocketAuthSettings { + config: Some(AppServerWebsocketAuthConfig::SignedBearerToken { + shared_secret_file: AbsolutePathBuf::from_absolute_path("/tmp/secret") + .expect("absolute path"), + issuer: Some("issuer".to_string()), + audience: None, + max_clock_skew_seconds: DEFAULT_MAX_CLOCK_SKEW_SECONDS, + }), + } + ); + } + + #[test] + fn signed_bearer_token_verification_rejects_tampering() { + let shared_secret = b"0123456789abcdef0123456789abcdef"; + let token = signed_token( + shared_secret, + json!({ + "exp": OffsetDateTime::now_utc().unix_timestamp() + 60, + }), + ); + let tampered = token.replace(".eyJleHAi", ".eyJleHBi"); + let err = verify_signed_bearer_token( + &tampered, + shared_secret, + /*issuer*/ None, + /*audience*/ None, + /*max_clock_skew_seconds*/ 30, + ) + .expect_err("tampered jwt should fail"); + assert_eq!(err.status_code(), StatusCode::UNAUTHORIZED); + } + + #[test] + fn signed_bearer_token_verification_accepts_valid_token() { + let shared_secret = b"0123456789abcdef0123456789abcdef"; + let token = signed_token( + shared_secret, + json!({ + "exp": OffsetDateTime::now_utc().unix_timestamp() + 60, + "iss": "issuer", + "aud": "audience", + }), + ); + verify_signed_bearer_token( + &token, + shared_secret, + Some("issuer"), + Some("audience"), + /*max_clock_skew_seconds*/ 30, + ) + .expect("valid signed token should verify"); + } + + #[test] + fn signed_bearer_token_verification_accepts_multiple_audiences() { + let shared_secret = b"0123456789abcdef0123456789abcdef"; + let token = signed_token( + shared_secret, + json!({ + "exp": OffsetDateTime::now_utc().unix_timestamp() + 60, + "aud": ["other-audience", "audience"], + }), + ); + verify_signed_bearer_token( + &token, + shared_secret, + /*issuer*/ None, + Some("audience"), + /*max_clock_skew_seconds*/ 30, + ) + .expect("jwt audience arrays should verify"); + } + + #[test] + fn signed_bearer_token_verification_rejects_alg_none_tokens() { + let claims_segment = URL_SAFE_NO_PAD.encode( + serde_json::to_vec(&json!({ + "exp": OffsetDateTime::now_utc().unix_timestamp() + 60, + })) + .unwrap(), + ); + let header_segment = URL_SAFE_NO_PAD.encode(br#"{"alg":"none","typ":"JWT"}"#); + let token = format!("{header_segment}.{claims_segment}."); + let err = verify_signed_bearer_token( + &token, + b"0123456789abcdef0123456789abcdef", + /*issuer*/ None, + /*audience*/ None, + /*max_clock_skew_seconds*/ 30, + ) + .expect_err("alg=none jwt should be rejected"); + assert_eq!(err.status_code(), StatusCode::UNAUTHORIZED); + } + + #[test] + fn signed_bearer_token_verification_rejects_missing_exp() { + let shared_secret = b"0123456789abcdef0123456789abcdef"; + let token = signed_token( + shared_secret, + json!({ + "iss": "issuer", + }), + ); + let err = verify_signed_bearer_token( + &token, + shared_secret, + /*issuer*/ None, + /*audience*/ None, + /*max_clock_skew_seconds*/ 30, + ) + .expect_err("jwt without exp should be rejected"); + assert_eq!(err.status_code(), StatusCode::UNAUTHORIZED); + } + + #[test] + fn validate_signed_bearer_secret_rejects_short_secret() { + let err = validate_signed_bearer_secret(Path::new("/tmp/secret"), b"too-short") + .expect_err("short shared secret should be rejected"); + assert_eq!(err.kind(), ErrorKind::InvalidInput); + assert!( + err.to_string().contains("must be at least 32 bytes"), + "unexpected error: {err}" + ); + } +} diff --git a/codex-rs/app-server-transport/src/transport/mod.rs b/codex-rs/app-server-transport/src/transport/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..56faef69b0462c88efcac98c162f699943b4f160 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/mod.rs @@ -0,0 +1,602 @@ +pub mod auth; + +use crate::outgoing_message::ConnectionId; +use crate::outgoing_message::OutgoingError; +use crate::outgoing_message::OutgoingMessage; +use crate::outgoing_message::QueuedOutgoingMessage; +use codex_app_server_protocol::JSONRPCErrorError; +use codex_app_server_protocol::JSONRPCMessage; +use codex_app_server_protocol::RequestId; +use codex_core::config::find_codex_home; +use codex_utils_absolute_path::AbsolutePathBuf; +use std::net::SocketAddr; +use std::path::Path; +use std::path::PathBuf; +use std::str::FromStr; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; +use tracing::error; +use tracing::warn; + +/// Size of the bounded channels used to communicate between tasks. The value +/// is a balance between throughput and memory usage - 128 messages should be +/// plenty for an interactive CLI. +pub const CHANNEL_CAPACITY: usize = 128; + +mod remote_control; +mod stdio; +mod unix_socket; +#[cfg(test)] +mod unix_socket_tests; +mod websocket; + +pub use remote_control::REMOTE_CONTROL_DISABLED_ENV_VAR; +pub use remote_control::RemoteControlDisabledByRequirements; +pub use remote_control::RemoteControlEnableError; +pub use remote_control::RemoteControlHandle; +pub use remote_control::RemoteControlPolicy; +pub use remote_control::RemoteControlStartConfig; +pub use remote_control::RemoteControlStartupMode; +pub use remote_control::RemoteControlUnavailable; +pub use remote_control::start_remote_control; +pub use remote_control::take_remote_control_disabled_env; +pub use stdio::start_stdio_connection; +pub use unix_socket::AppServerStartupLock; +pub use unix_socket::DaemonShutdownAccess; +pub use unix_socket::acquire_app_server_startup_lock; +pub use unix_socket::prepare_control_socket_path; +pub use unix_socket::start_control_socket_acceptor; +pub use websocket::start_websocket_acceptor; + +const INTERNAL_ERROR_CODE: i64 = -32603; +const OVERLOADED_ERROR_CODE: i64 = -32001; + +const APP_SERVER_CONTROL_SOCKET_DIR_NAME: &str = "app-server-control"; +const APP_SERVER_CONTROL_SOCKET_FILE_NAME: &str = "app-server-control.sock"; +const APP_SERVER_STARTUP_LOCK_FILE_NAME: &str = "app-server-startup.lock"; +const DAEMON_RECOVERY_FILE_NAME: &str = "loaded-threads.json"; + +pub fn daemon_recovery_file_path(codex_home: &Path) -> PathBuf { + codex_home + .join("app-server-daemon") + .join(DAEMON_RECOVERY_FILE_NAME) +} + +pub fn app_server_control_socket_path(codex_home: &Path) -> std::io::Result { + AbsolutePathBuf::from_absolute_path( + codex_home + .join(APP_SERVER_CONTROL_SOCKET_DIR_NAME) + .join(APP_SERVER_CONTROL_SOCKET_FILE_NAME), + ) +} + +pub fn app_server_startup_lock_path(codex_home: &Path) -> std::io::Result { + AbsolutePathBuf::from_absolute_path( + codex_home + .join(APP_SERVER_CONTROL_SOCKET_DIR_NAME) + .join(APP_SERVER_STARTUP_LOCK_FILE_NAME), + ) +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum AppServerTransport { + Stdio, + UnixSocket { socket_path: AbsolutePathBuf }, + WebSocket { bind_address: SocketAddr }, + Off, +} + +#[derive(Debug, Clone, Eq, PartialEq)] +pub enum AppServerTransportParseError { + UnsupportedListenUrl(String), + InvalidUnixSocketPath { listen_url: String, message: String }, + InvalidWebSocketListenUrl(String), +} + +impl std::fmt::Display for AppServerTransportParseError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + AppServerTransportParseError::UnsupportedListenUrl(listen_url) => write!( + f, + "unsupported --listen URL `{listen_url}`; expected `stdio://`, `unix://`, `unix://PATH`, `ws://IP:PORT`, or `off`" + ), + AppServerTransportParseError::InvalidUnixSocketPath { + listen_url, + message, + } => write!( + f, + "invalid unix socket --listen URL `{listen_url}`; failed to resolve socket path: {message}" + ), + AppServerTransportParseError::InvalidWebSocketListenUrl(listen_url) => write!( + f, + "invalid websocket --listen URL `{listen_url}`; expected `ws://IP:PORT`" + ), + } + } +} + +impl std::error::Error for AppServerTransportParseError {} + +impl AppServerTransport { + pub const DEFAULT_LISTEN_URL: &'static str = "stdio://"; + + pub fn from_listen_url(listen_url: &str) -> Result { + if listen_url == Self::DEFAULT_LISTEN_URL { + return Ok(Self::Stdio); + } + + if let Some(raw_socket_path) = listen_url.strip_prefix("unix://") { + let socket_path = if raw_socket_path.is_empty() { + let codex_home = find_codex_home().map_err(|err| { + AppServerTransportParseError::InvalidUnixSocketPath { + listen_url: listen_url.to_string(), + message: format!("failed to resolve CODEX_HOME: {err}"), + } + })?; + app_server_control_socket_path(&codex_home).map_err(|err| { + AppServerTransportParseError::InvalidUnixSocketPath { + listen_url: listen_url.to_string(), + message: err.to_string(), + } + })? + } else { + AbsolutePathBuf::relative_to_current_dir(raw_socket_path).map_err(|err| { + AppServerTransportParseError::InvalidUnixSocketPath { + listen_url: listen_url.to_string(), + message: err.to_string(), + } + })? + }; + return Ok(Self::UnixSocket { socket_path }); + } + + if listen_url == "off" { + return Ok(Self::Off); + } + + if let Some(socket_addr) = listen_url.strip_prefix("ws://") { + let bind_address = socket_addr.parse::().map_err(|_| { + AppServerTransportParseError::InvalidWebSocketListenUrl(listen_url.to_string()) + })?; + return Ok(Self::WebSocket { bind_address }); + } + + Err(AppServerTransportParseError::UnsupportedListenUrl( + listen_url.to_string(), + )) + } +} + +impl FromStr for AppServerTransport { + type Err = AppServerTransportParseError; + + fn from_str(s: &str) -> Result { + Self::from_listen_url(s) + } +} + +#[derive(Debug)] +pub enum TransportEvent { + /// Accepted on the managed local control socket, outside JSON-RPC. + DaemonShutdown, + ConnectionOpened { + connection_id: ConnectionId, + origin: ConnectionOrigin, + auth: Option, + writer: mpsc::Sender, + disconnect_sender: Option, + }, + ConnectionClosed { + connection_id: ConnectionId, + }, + IncomingMessage { + connection_id: ConnectionId, + message: JSONRPCMessage, + }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ConnectionOrigin { + Stdio, + InProcess, + WebSocket, + RemoteControl, +} + +static CONNECTION_ID_COUNTER: AtomicU64 = AtomicU64::new(0); + +fn next_connection_id() -> ConnectionId { + ConnectionId(CONNECTION_ID_COUNTER.fetch_add(1, Ordering::Relaxed)) +} + +async fn forward_incoming_message( + transport_event_tx: &mpsc::Sender, + writer: &mpsc::Sender, + connection_id: ConnectionId, + payload: &str, +) -> bool { + match serde_json::from_str::(payload) { + Ok(message) => { + enqueue_incoming_message(transport_event_tx, writer, connection_id, message).await + } + Err(err) => { + error!("Failed to deserialize JSONRPCMessage: {err}"); + true + } + } +} + +async fn enqueue_incoming_message( + transport_event_tx: &mpsc::Sender, + writer: &mpsc::Sender, + connection_id: ConnectionId, + message: JSONRPCMessage, +) -> bool { + let event = TransportEvent::IncomingMessage { + connection_id, + message, + }; + match transport_event_tx.try_send(event) { + Ok(()) => true, + Err(mpsc::error::TrySendError::Closed(_)) => false, + Err(mpsc::error::TrySendError::Full(TransportEvent::IncomingMessage { + connection_id, + message: JSONRPCMessage::Request(request), + })) => { + let overload_error = OutgoingMessage::Error(OutgoingError { + id: request.id, + error: JSONRPCErrorError { + code: OVERLOADED_ERROR_CODE, + message: "Server overloaded; retry later.".to_string(), + data: None, + }, + }); + match writer.try_send(QueuedOutgoingMessage::new(overload_error)) { + Ok(()) => true, + Err(mpsc::error::TrySendError::Closed(_)) => false, + Err(mpsc::error::TrySendError::Full(_overload_error)) => { + warn!( + "dropping overload response for connection {:?}: outbound queue is full", + connection_id + ); + true + } + } + } + Err(mpsc::error::TrySendError::Full(event)) => transport_event_tx.send(event).await.is_ok(), + } +} + +fn serialize_outgoing_message(outgoing_message: OutgoingMessage) -> Option { + match serde_json::to_string(&outgoing_message) { + Ok(json) => Some(json), + Err(err) => { + error!("Failed to serialize JSONRPCMessage: {err}"); + let OutgoingMessage::Response(response) = outgoing_message else { + return None; + }; + serde_json::to_string(&response_serialization_error(response.id, err)) + .inspect_err(|err| error!("Failed to serialize JSONRPC error: {err}")) + .ok() + } + } +} + +fn response_serialization_error( + request_id: RequestId, + err: impl std::fmt::Display, +) -> OutgoingMessage { + OutgoingMessage::Error(OutgoingError { + id: request_id, + error: JSONRPCErrorError { + code: INTERNAL_ERROR_CODE, + message: format!("failed to serialize response: {err}"), + data: None, + }, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::outgoing_message::OutgoingResponse; + use codex_app_server_protocol::ClientResponsePayload; + use codex_app_server_protocol::ConfigWarningNotification; + use codex_app_server_protocol::JSONRPCNotification; + use codex_app_server_protocol::JSONRPCRequest; + use codex_app_server_protocol::JSONRPCResponse; + use codex_app_server_protocol::RequestId; + use codex_app_server_protocol::ServerNotification; + use codex_app_server_protocol::ServerNotificationEnvelope; + use codex_app_server_protocol::ThreadArchiveResponse; + use pretty_assertions::assert_eq; + use serde_json::json; + use tokio::time::Duration; + use tokio::time::timeout; + + #[test] + fn listen_off_parses_as_off_transport() { + assert_eq!( + AppServerTransport::from_listen_url("off"), + Ok(AppServerTransport::Off) + ); + } + + #[test] + fn serialize_outgoing_message_preserves_wire_shape() { + let message = OutgoingMessage::AppServerNotification(ServerNotificationEnvelope { + notification: ServerNotification::ConfigWarning(ConfigWarningNotification { + summary: "summary".to_string(), + details: None, + path: None, + range: None, + }), + emitted_at_ms: Some(1_234), + }); + + let json = serialize_outgoing_message(message).expect("message should serialize"); + assert_eq!( + serde_json::from_str::(&json).expect("message should be valid JSON"), + json!({ + "method": "configWarning", + "params": { + "summary": "summary", + "details": null, + }, + "emittedAtMs": 1_234, + }) + ); + } + + #[test] + fn serialize_typed_response_preserves_wire_shape() { + let message = OutgoingMessage::Response(OutgoingResponse { + id: RequestId::Integer(7), + result: Box::new(ClientResponsePayload::ThreadArchive( + ThreadArchiveResponse {}, + )), + }); + + let json = serialize_outgoing_message(message).expect("message should serialize"); + assert_eq!( + serde_json::from_str::(&json).expect("message should be valid JSON"), + json!({ "id": 7, "result": {} }) + ); + } + + #[cfg(unix)] + #[test] + fn serialize_invalid_typed_response_returns_jsonrpc_error() { + use std::ffi::OsString; + use std::os::unix::ffi::OsStringExt; + use std::path::PathBuf; + + let codex_home = + AbsolutePathBuf::from_absolute_path(PathBuf::from(OsString::from_vec(vec![ + b'/', b'b', b'a', b'd', 0xff, + ]))) + .expect("non-UTF-8 Unix paths are valid absolute paths"); + let message = OutgoingMessage::Response(OutgoingResponse { + id: RequestId::Integer(7), + result: Box::new(ClientResponsePayload::Initialize( + codex_app_server_protocol::InitializeResponse { + user_agent: "codex-test-agent".to_string(), + codex_home, + platform_family: "unix".to_string(), + platform_os: "linux".to_string(), + }, + )), + }); + + let json = serialize_outgoing_message(message) + .expect("invalid response should serialize as a JSON-RPC error"); + assert_eq!( + serde_json::from_str::(&json).expect("message should be valid JSON"), + json!({ + "id": 7, + "error": { + "code": -32603, + "message": "failed to serialize response: path contains invalid UTF-8 characters", + } + }) + ); + } + + #[tokio::test] + async fn enqueue_incoming_request_returns_overload_error_when_queue_is_full() { + let connection_id = ConnectionId(42); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(1); + let (writer_tx, mut writer_rx) = mpsc::channel(1); + + let first_message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + transport_event_tx + .send(TransportEvent::IncomingMessage { + connection_id, + message: first_message.clone(), + }) + .await + .expect("queue should accept first message"); + + let request = JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(7), + method: "config/read".to_string(), + params: Some(json!({ "includeLayers": false })), + trace: None, + }); + assert!( + enqueue_incoming_message(&transport_event_tx, &writer_tx, connection_id, request).await + ); + + let queued_event = transport_event_rx + .recv() + .await + .expect("first event should stay queued"); + match queued_event { + TransportEvent::IncomingMessage { + connection_id: queued_connection_id, + message, + } => { + assert_eq!(queued_connection_id, connection_id); + assert_eq!(message, first_message); + } + _ => panic!("expected queued incoming message"), + } + + let overload = writer_rx + .recv() + .await + .expect("request should receive overload error"); + let overload_json = + serde_json::to_value(overload.message).expect("serialize overload error"); + assert_eq!( + overload_json, + json!({ + "id": 7, + "error": { + "code": OVERLOADED_ERROR_CODE, + "message": "Server overloaded; retry later." + } + }) + ); + } + + #[tokio::test] + async fn enqueue_incoming_response_waits_instead_of_dropping_when_queue_is_full() { + let connection_id = ConnectionId(42); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(1); + let (writer_tx, _writer_rx) = mpsc::channel(1); + + let first_message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + transport_event_tx + .send(TransportEvent::IncomingMessage { + connection_id, + message: first_message.clone(), + }) + .await + .expect("queue should accept first message"); + + let response = JSONRPCMessage::Response(JSONRPCResponse { + id: RequestId::Integer(7), + result: json!({"ok": true}), + }); + let transport_event_tx_for_enqueue = transport_event_tx.clone(); + let writer_tx_for_enqueue = writer_tx.clone(); + let enqueue_handle = tokio::spawn(async move { + enqueue_incoming_message( + &transport_event_tx_for_enqueue, + &writer_tx_for_enqueue, + connection_id, + response, + ) + .await + }); + + let queued_event = transport_event_rx + .recv() + .await + .expect("first event should be dequeued"); + match queued_event { + TransportEvent::IncomingMessage { + connection_id: queued_connection_id, + message, + } => { + assert_eq!(queued_connection_id, connection_id); + assert_eq!(message, first_message); + } + _ => panic!("expected queued incoming message"), + } + + let enqueue_result = enqueue_handle.await.expect("enqueue task should not panic"); + assert!(enqueue_result); + + let forwarded_event = transport_event_rx + .recv() + .await + .expect("response should be forwarded instead of dropped"); + match forwarded_event { + TransportEvent::IncomingMessage { + connection_id: queued_connection_id, + message: JSONRPCMessage::Response(JSONRPCResponse { id, result }), + } => { + assert_eq!(queued_connection_id, connection_id); + assert_eq!(id, RequestId::Integer(7)); + assert_eq!(result, json!({"ok": true})); + } + _ => panic!("expected forwarded response message"), + } + } + + #[tokio::test] + async fn enqueue_incoming_request_does_not_block_when_writer_queue_is_full() { + let connection_id = ConnectionId(42); + let (transport_event_tx, _transport_event_rx) = mpsc::channel(1); + let (writer_tx, mut writer_rx) = mpsc::channel(1); + + transport_event_tx + .send(TransportEvent::IncomingMessage { + connection_id, + message: JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }), + }) + .await + .expect("transport queue should accept first message"); + + writer_tx + .send(QueuedOutgoingMessage::new( + OutgoingMessage::AppServerNotification(ServerNotificationEnvelope { + notification: ServerNotification::ConfigWarning(ConfigWarningNotification { + summary: "queued".to_string(), + details: None, + path: None, + range: None, + }), + emitted_at_ms: Some(1_234), + }), + )) + .await + .expect("writer queue should accept first message"); + + let request = JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(7), + method: "config/read".to_string(), + params: Some(json!({ "includeLayers": false })), + trace: None, + }); + + let enqueue_result = timeout( + Duration::from_millis(100), + enqueue_incoming_message(&transport_event_tx, &writer_tx, connection_id, request), + ) + .await + .expect("enqueue should not block while writer queue is full"); + assert!(enqueue_result); + + let queued_outgoing = writer_rx + .recv() + .await + .expect("writer queue should still contain original message"); + let queued_json = + serde_json::to_value(queued_outgoing.message).expect("serialize queued message"); + assert_eq!( + queued_json, + json!({ + "method": "configWarning", + "params": { + "summary": "queued", + "details": null, + }, + "emittedAtMs": 1_234, + }) + ); + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/auth.rs b/codex-rs/app-server-transport/src/transport/remote_control/auth.rs new file mode 100644 index 0000000000000000000000000000000000000000..93c728abec8ca71d202a4908cc55f0749b6bc323 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/auth.rs @@ -0,0 +1,282 @@ +//! Credentials and recovery bound to one remote-control login lifetime. +//! A request can refresh credentials, but cannot adopt a replacement authentication owner. + +use axum::http::HeaderMap; +use axum::http::HeaderValue; +use codex_api::SharedAuthProvider; +use codex_login::AuthManager; +use codex_login::UnauthorizedRecovery; +use std::io; +use std::io::ErrorKind; +use std::sync::Arc; +use tokio::sync::watch; +use tracing::info; +use tracing::warn; + +#[derive(Clone)] +pub(super) struct RemoteControlAuth { + manager: Arc, + pub(super) owner: crate::ConnectionAuth, +} + +pub(super) struct RemoteControlRecovery { + auth: RemoteControlAuth, + recovery: UnauthorizedRecovery, +} + +impl RemoteControlAuth { + pub(super) fn capture(manager: Arc) -> (Self, bool) { + loop { + let owner = crate::ConnectionAuth::capture(&manager); + let authenticated = manager + .auth_cached() + .is_some_and(|auth| auth.uses_codex_backend() && auth.get_account_id().is_some()); + if owner.is_current() { + return (Self { manager, owner }, authenticated); + } + } + } + + pub(super) fn ensure_current(&self) -> io::Result<()> { + self.owner.ensure_current() + } + + pub(super) fn unauthorized_recovery(&self) -> RemoteControlRecovery { + RemoteControlRecovery { + auth: self.clone(), + recovery: self.manager.unauthorized_recovery(), + } + } + + pub(super) fn auth_change_receiver(&self) -> watch::Receiver { + self.manager.auth_change_receiver() + } +} + +pub(super) const REMOTE_CONTROL_ACCOUNT_ID_HEADER: &str = "chatgpt-account-id"; + +pub(super) struct RemoteControlConnectionAuth { + pub(super) auth_provider: SharedAuthProvider, + pub(super) account_id: String, +} + +impl RemoteControlConnectionAuth { + pub(super) fn request_headers(&self) -> io::Result { + let mut headers = HeaderMap::new(); + self.auth_provider.add_auth_headers(&mut headers); + headers.insert( + REMOTE_CONTROL_ACCOUNT_ID_HEADER, + HeaderValue::from_str(&self.account_id).map_err(|err| { + io::Error::new( + ErrorKind::InvalidInput, + format!("invalid remote control account id header: {err}"), + ) + })?, + ); + Ok(headers) + } +} + +pub(super) async fn load_remote_control_auth( + auth: &RemoteControlAuth, +) -> io::Result { + auth.ensure_current()?; + let credentials = load_auth_manager(&auth.manager).await?; + auth.ensure_current()?; + Ok(credentials) +} + +async fn load_auth_manager( + auth_manager: &Arc, +) -> io::Result { + let mut reloaded = false; + let auth = loop { + let Some(auth) = auth_manager.auth().await else { + if reloaded { + return Err(io::Error::new( + ErrorKind::PermissionDenied, + "remote control requires ChatGPT authentication", + )); + } + auth_manager.reload().await; + reloaded = true; + continue; + }; + if !auth.uses_codex_backend() { + break auth; + } + if auth.get_account_id().is_none() && !reloaded { + auth_manager.reload().await; + reloaded = true; + continue; + } + break auth; + }; + + if !auth.uses_codex_backend() { + return Err(io::Error::new( + ErrorKind::PermissionDenied, + "remote control requires ChatGPT authentication; API key auth is not supported", + )); + } + + Ok(RemoteControlConnectionAuth { + auth_provider: codex_model_provider::auth_provider_from_auth(&auth), + account_id: auth.get_account_id().ok_or_else(|| { + io::Error::new( + ErrorKind::WouldBlock, + "remote control enrollment is waiting for a ChatGPT account id", + ) + })?, + }) +} + +pub(super) async fn recover_remote_control_auth( + recovery: &mut RemoteControlRecovery, + auth_change_rx: &mut watch::Receiver, +) -> bool { + if recovery.auth.ensure_current().is_err() { + return false; + } + let auth_recovery = &mut recovery.recovery; + if !auth_recovery.has_next() { + return false; + } + + let mode = auth_recovery.mode_name(); + let step = auth_recovery.step_name(); + let auth_change_revision_before_recovery = *auth_change_rx.borrow(); + match auth_recovery.next().await { + Ok(step_result) => { + if recovery.auth.ensure_current().is_err() { + return false; + } + if step_result.auth_state_changed() == Some(true) { + mark_recovery_auth_change_seen( + auth_change_rx, + auth_change_revision_before_recovery, + ); + } + info!( + "remote control auth recovery succeeded: mode={mode}, step={step}, auth_state_changed={:?}", + step_result.auth_state_changed() + ); + true + } + Err(err) => { + warn!("remote control auth recovery failed: mode={mode}, step={step}: {err}"); + false + } + } +} + +pub(super) fn mark_recovery_auth_change_seen( + auth_change_rx: &mut watch::Receiver, + auth_change_revision_before_recovery: u64, +) { + let auth_change_revision_after_recovery = *auth_change_rx.borrow(); + if auth_change_revision_after_recovery == auth_change_revision_before_recovery.wrapping_add(1) { + // Recovery updated the same watch that wakes the outer reconnect + // loop. Mark only that single revision seen; if more revisions + // arrived while recovery was in flight, leave them pending so the + // reconnect loop still reacts to the later external auth change. + auth_change_rx.borrow_and_update(); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use codex_api::AuthProvider; + use pretty_assertions::assert_eq; + + #[derive(Debug)] + struct TestAuthProvider { + account_ids: Vec<&'static str>, + } + + impl AuthProvider for TestAuthProvider { + fn add_auth_headers(&self, headers: &mut HeaderMap) { + headers.insert( + axum::http::header::AUTHORIZATION, + HeaderValue::from_static("Bearer test-token"), + ); + headers.insert("x-openai-fedramp", HeaderValue::from_static("true")); + for account_id in &self.account_ids { + headers.append("ChatGPT-Account-ID", HeaderValue::from_static(account_id)); + } + } + } + + fn remote_control_auth( + account_id: &str, + provider_account_ids: Vec<&'static str>, + ) -> RemoteControlConnectionAuth { + RemoteControlConnectionAuth { + auth_provider: Arc::new(TestAuthProvider { + account_ids: provider_account_ids, + }), + account_id: account_id.to_string(), + } + } + + #[test] + fn request_headers_adds_account_header_when_provider_omits_it() { + let headers = remote_control_auth("selected-account", Vec::new()) + .request_headers() + .expect("request headers should build"); + + assert_eq!( + headers + .get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER) + .iter() + .map(|value| value.to_str().expect("account header should be text")) + .collect::>(), + vec!["selected-account"] + ); + } + + #[test] + fn request_headers_replaces_provider_accounts_and_preserves_other_headers() { + let headers = remote_control_auth( + "selected-account", + vec!["provider-account-a", "provider-account-b"], + ) + .request_headers() + .expect("request headers should build"); + + assert_eq!( + headers + .get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER) + .iter() + .map(|value| value.to_str().expect("account header should be text")) + .collect::>(), + vec!["selected-account"] + ); + assert_eq!( + headers + .get(axum::http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some("Bearer test-token") + ); + assert_eq!( + headers + .get("x-openai-fedramp") + .and_then(|value| value.to_str().ok()), + Some("true") + ); + } + + #[test] + fn request_headers_rejects_invalid_account_header_value() { + let err = remote_control_auth("invalid\naccount", Vec::new()) + .request_headers() + .expect_err("invalid account header should fail"); + + assert_eq!(err.kind(), ErrorKind::InvalidInput); + assert!( + err.to_string() + .starts_with("invalid remote control account id header:") + ); + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/client_tracker.rs b/codex-rs/app-server-transport/src/transport/remote_control/client_tracker.rs new file mode 100644 index 0000000000000000000000000000000000000000..7c46632e3ddb8ed6e3d68f0d53fe716454d54781 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/client_tracker.rs @@ -0,0 +1,944 @@ +use super::CHANNEL_CAPACITY; +use super::TransportEvent; +use super::next_connection_id; +use super::protocol::ClientEnvelope; +pub use super::protocol::ClientEvent; +pub use super::protocol::ClientId; +use super::protocol::PongStatus; +use super::protocol::ServerEvent; +use super::protocol::StreamId; +use crate::outgoing_message::ConnectionId; +use crate::outgoing_message::QueuedOutgoingMessage; +use crate::transport::ConnectionOrigin; +use crate::transport::remote_control::QueuedServerEnvelope; +use codex_app_server_protocol::JSONRPCMessage; +use std::collections::HashMap; +use tokio::sync::mpsc; +use tokio::sync::watch; +use tokio::task::JoinHandle; +use tokio::task::JoinSet; +use tokio::time::Duration; +use tokio::time::Instant; +use tokio::time::timeout; +use tokio_util::sync::CancellationToken; +use tracing::info; +use tracing::warn; + +const REMOTE_CONTROL_CLIENT_IDLE_TIMEOUT: Duration = Duration::from_secs(10 * 60); +pub(crate) const REMOTE_CONTROL_IDLE_SWEEP_INTERVAL: Duration = Duration::from_secs(30); +#[cfg(not(test))] +const REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT: Duration = Duration::from_secs(5); +#[cfg(test)] +const REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT: Duration = Duration::from_millis(10); + +#[derive(Debug)] +pub(crate) struct Stopped; + +struct ClientState { + connection_id: ConnectionId, + disconnect_token: CancellationToken, + last_activity_at: Instant, + last_inbound_seq_id: Option, + status_tx: watch::Sender, +} + +pub(crate) struct ClientTracker { + pub(super) auth: Option, + clients: HashMap<(ClientId, StreamId), ClientState>, + legacy_stream_ids: HashMap, + join_set: JoinSet<(ClientId, StreamId)>, + server_event_tx: mpsc::Sender, + transport_event_tx: mpsc::Sender, + shutdown_token: CancellationToken, +} + +impl ClientTracker { + pub(crate) fn new( + server_event_tx: mpsc::Sender, + transport_event_tx: mpsc::Sender, + shutdown_token: &CancellationToken, + ) -> Self { + Self { + auth: None, + clients: HashMap::new(), + legacy_stream_ids: HashMap::new(), + join_set: JoinSet::new(), + server_event_tx, + transport_event_tx, + shutdown_token: shutdown_token.child_token(), + } + } + + pub(crate) async fn bookkeep_join_set(&mut self) -> Option<(ClientId, StreamId)> { + while let Some(join_result) = self.join_set.join_next().await { + let Ok(client_key) = join_result else { + continue; + }; + return Some(client_key); + } + futures::future::pending().await + } + + pub(crate) async fn shutdown(&mut self) { + self.shutdown_token.cancel(); + + while let Some(client_key) = self.clients.keys().next().cloned() { + let _ = self.close_client(&client_key).await; + } + + self.drain_join_set().await; + } + + async fn drain_join_set(&mut self) { + while self.join_set.join_next().await.is_some() {} + } + + pub(crate) async fn handle_message( + &mut self, + client_envelope: ClientEnvelope, + ) -> Result<(), Stopped> { + let ClientEnvelope { + client_id, + event, + stream_id, + seq_id, + cursor: _, + } = client_envelope; + let is_legacy_stream_id = stream_id.is_none(); + let is_initialize = matches!(&event, ClientEvent::ClientMessage { message } if remote_control_message_starts_connection(message)); + let stream_id = match stream_id { + Some(stream_id) => stream_id, + None if is_initialize => { + // TODO(ruslan): delete this fallback once all clients are updated to send stream_id. + self.legacy_stream_ids + .remove(&client_id) + .unwrap_or_else(StreamId::new_random) + } + None => self + .legacy_stream_ids + .get(&client_id) + .cloned() + .unwrap_or_else(|| { + if matches!(&event, ClientEvent::Ping) { + StreamId::new_random() + } else { + StreamId(String::new()) + } + }), + }; + if stream_id.0.is_empty() { + return Ok(()); + } + let client_key = (client_id.clone(), stream_id.clone()); + match event { + ClientEvent::ClientMessage { message } => { + if let Some(seq_id) = seq_id + && let Some(client) = self.clients.get(&client_key) + && client + .last_inbound_seq_id + .is_some_and(|last_seq_id| last_seq_id >= seq_id) + && !is_initialize + { + return Ok(()); + } + + if is_initialize && self.clients.contains_key(&client_key) { + self.close_client(&client_key).await?; + } + + if let Some(connection_id) = self.clients.get_mut(&client_key).map(|client| { + client.last_activity_at = Instant::now(); + client.connection_id + }) { + self.send_transport_event(TransportEvent::IncomingMessage { + connection_id, + message, + }) + .await?; + self.record_inbound_message_delivery(&client_key, seq_id); + return Ok(()); + } + + if !is_initialize { + return Ok(()); + } + + let connection_id = next_connection_id(); + let (writer_tx, writer_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let disconnect_token = self.shutdown_token.child_token(); + self.send_transport_event(TransportEvent::ConnectionOpened { + connection_id, + origin: ConnectionOrigin::RemoteControl, + auth: self.auth.clone(), + writer: writer_tx, + disconnect_sender: Some(disconnect_token.clone()), + }) + .await?; + + let (status_tx, status_rx) = watch::channel(PongStatus::Active); + self.join_set.spawn(Self::run_client_outbound( + client_id.clone(), + stream_id.clone(), + self.server_event_tx.clone(), + writer_rx, + status_rx, + disconnect_token.clone(), + )); + self.clients.insert( + client_key.clone(), + ClientState { + connection_id, + disconnect_token, + last_activity_at: Instant::now(), + last_inbound_seq_id: None, + status_tx, + }, + ); + if is_legacy_stream_id { + self.legacy_stream_ids.insert(client_id.clone(), stream_id); + } + if let Err(err) = self + .send_transport_event(TransportEvent::IncomingMessage { + connection_id, + message, + }) + .await + { + if let Some(client) = self.remove_client(&client_key) { + client.disconnect_token.cancel(); + // The initialize send already timed out on this queue; preserve close + // delivery without blocking reconnect on the same backpressure. + drop(self.spawn_connection_closed(client.connection_id)); + } + return Err(err); + } + if !is_legacy_stream_id { + self.record_inbound_message_delivery(&client_key, seq_id); + } + Ok(()) + } + ClientEvent::ClientMessageChunk { .. } | ClientEvent::Ack { .. } => Ok(()), + ClientEvent::Ping => { + if let Some(client) = self.clients.get_mut(&client_key) { + client.last_activity_at = Instant::now(); + let _ = client.status_tx.send(PongStatus::Active); + return Ok(()); + } + + let server_event_tx = self.server_event_tx.clone(); + tokio::spawn(async move { + let server_envelope = QueuedServerEnvelope { + event: ServerEvent::Pong { + status: PongStatus::Unknown, + }, + client_id, + stream_id, + write_complete_tx: None, + }; + let _ = server_event_tx.send(server_envelope).await; + }); + Ok(()) + } + ClientEvent::ClientClosed => self.close_client(&client_key).await, + } + } + + async fn run_client_outbound( + client_id: ClientId, + stream_id: StreamId, + server_event_tx: mpsc::Sender, + mut writer_rx: mpsc::Receiver, + mut status_rx: watch::Receiver, + disconnect_token: CancellationToken, + ) -> (ClientId, StreamId) { + loop { + let (event, write_complete_tx) = tokio::select! { + _ = disconnect_token.cancelled() => { + break; + } + queued_message = writer_rx.recv() => { + let Some(queued_message) = queued_message else { + break; + }; + let event = ServerEvent::ServerMessage { + message: Box::new(queued_message.message), + }; + (event, queued_message.write_complete_tx) + } + changed = status_rx.changed() => { + if changed.is_err() { + break; + } + let event = ServerEvent::Pong { status: status_rx.borrow().clone() }; + (event, None) + } + }; + let send_result = tokio::select! { + _ = disconnect_token.cancelled() => { + break; + } + send_result = server_event_tx.send(QueuedServerEnvelope { + event, + client_id: client_id.clone(), + stream_id: stream_id.clone(), + write_complete_tx, + }) => send_result, + }; + if send_result.is_err() { + break; + } + } + (client_id, stream_id) + } + + pub(crate) async fn close_expired_clients( + &mut self, + ) -> Result, Stopped> { + let now = Instant::now(); + let expired_client_ids: Vec<(ClientId, StreamId)> = self + .clients + .iter() + .filter_map(|(client_key, client)| { + (!remote_control_client_is_alive(client, now)).then_some(client_key.clone()) + }) + .collect(); + for client_key in &expired_client_ids { + self.close_client(client_key).await?; + } + Ok(expired_client_ids) + } + + pub(super) async fn close_client( + &mut self, + client_key: &(ClientId, StreamId), + ) -> Result<(), Stopped> { + let Some(client) = self.remove_client(client_key) else { + return Ok(()); + }; + client.disconnect_token.cancel(); + self.send_transport_event(TransportEvent::ConnectionClosed { + connection_id: client.connection_id, + }) + .await + } + + fn remove_client(&mut self, client_key: &(ClientId, StreamId)) -> Option { + let client = self.clients.remove(client_key)?; + if self + .legacy_stream_ids + .get(&client_key.0) + .is_some_and(|stream_id| stream_id == &client_key.1) + { + self.legacy_stream_ids.remove(&client_key.0); + } + Some(client) + } + + async fn send_transport_event(&self, event: TransportEvent) -> Result<(), Stopped> { + let event = match event { + TransportEvent::ConnectionClosed { connection_id } => { + return self.send_connection_closed(connection_id).await; + } + event => event, + }; + + let event_name = transport_event_name(&event); + match timeout( + REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT, + self.transport_event_tx.send(event), + ) + .await + { + Ok(Ok(())) => Ok(()), + Ok(Err(_)) => { + warn!( + transport_event = event_name, + "remote control transport event receiver dropped" + ); + Err(Stopped) + } + Err(_) => { + warn!( + transport_event = event_name, + timeout = ?REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT, + "timed out forwarding remote control transport event" + ); + Err(Stopped) + } + } + } + + fn record_inbound_message_delivery( + &mut self, + client_key: &(ClientId, StreamId), + seq_id: Option, + ) { + // Timed forwarding can fail, so only dedupe retries after app-server receives it. + if let Some(seq_id) = seq_id + && let Some(client) = self.clients.get_mut(client_key) + { + client.last_inbound_seq_id = Some(seq_id); + } + } + + async fn send_connection_closed(&self, connection_id: ConnectionId) -> Result<(), Stopped> { + // Worker shutdown can abort the caller; detach the cleanup event before awaiting it. + match self.spawn_connection_closed(connection_id).await { + Ok(result) => result, + Err(err) => { + warn!( + transport_event = "connection_closed", + ?err, + "remote control transport event forwarding task failed" + ); + Err(Stopped) + } + } + } + + fn spawn_connection_closed( + &self, + connection_id: ConnectionId, + ) -> JoinHandle> { + info!( + connection_id = ?connection_id, + "forwarding remote control connection closed transport event" + ); + let transport_event_tx = self.transport_event_tx.clone(); + tokio::spawn(async move { + transport_event_tx + .send(TransportEvent::ConnectionClosed { connection_id }) + .await + .map_err(|_| { + warn!( + transport_event = "connection_closed", + "remote control transport event receiver dropped" + ); + Stopped + }) + }) + } +} + +fn transport_event_name(event: &TransportEvent) -> &'static str { + match event { + TransportEvent::ConnectionOpened { .. } => "connection_opened", + TransportEvent::ConnectionClosed { .. } => "connection_closed", + TransportEvent::IncomingMessage { .. } => "incoming_message", + TransportEvent::DaemonShutdown => "daemon_shutdown", + } +} + +fn remote_control_message_starts_connection(message: &JSONRPCMessage) -> bool { + matches!( + message, + JSONRPCMessage::Request(codex_app_server_protocol::JSONRPCRequest { method, .. }) + if method == "initialize" + ) +} + +fn remote_control_client_is_alive(client: &ClientState, now: Instant) -> bool { + now.duration_since(client.last_activity_at) < REMOTE_CONTROL_CLIENT_IDLE_TIMEOUT +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::outgoing_message::OutgoingMessage; + use crate::transport::remote_control::protocol::ClientEnvelope; + use crate::transport::remote_control::protocol::ClientEvent; + use codex_app_server_protocol::ConfigWarningNotification; + use codex_app_server_protocol::JSONRPCRequest; + use codex_app_server_protocol::RequestId; + use codex_app_server_protocol::ServerNotification; + use codex_app_server_protocol::ServerNotificationEnvelope; + use pretty_assertions::assert_eq; + use serde_json::json; + use tokio::time::timeout; + + fn initialize_envelope(client_id: &str) -> ClientEnvelope { + initialize_envelope_with_stream_id(client_id, /*stream_id*/ None) + } + + fn initialize_envelope_with_stream_id( + client_id: &str, + stream_id: Option<&str>, + ) -> ClientEnvelope { + ClientEnvelope { + event: ClientEvent::ClientMessage { + message: JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(1), + method: "initialize".to_string(), + params: Some(json!({ + "clientInfo": { + "name": "remote-test-client", + "version": "0.1.0" + } + })), + trace: None, + }), + }, + client_id: ClientId(client_id.to_string()), + stream_id: stream_id.map(|stream_id| StreamId(stream_id.to_string())), + seq_id: Some(0), + cursor: None, + } + } + + fn initialized_notification() -> JSONRPCMessage { + JSONRPCMessage::Notification(codex_app_server_protocol::JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }) + } + + #[tokio::test] + async fn cancelled_outbound_task_emits_connection_closed() { + let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let mut client_tracker = + ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token); + + client_tracker + .handle_message(initialize_envelope("client-1")) + .await + .expect("initialize should open client"); + + let (connection_id, disconnect_sender) = match transport_event_rx + .recv() + .await + .expect("connection opened should be sent") + { + TransportEvent::ConnectionOpened { + connection_id, + disconnect_sender: Some(disconnect_sender), + .. + } => (connection_id, disconnect_sender), + other => panic!("expected connection opened, got {other:?}"), + }; + match transport_event_rx + .recv() + .await + .expect("initialize should be forwarded") + { + TransportEvent::IncomingMessage { + connection_id: incoming_connection_id, + .. + } => assert_eq!(incoming_connection_id, connection_id), + other => panic!("expected incoming initialize, got {other:?}"), + } + + disconnect_sender.cancel(); + let closed_client_id = timeout(Duration::from_secs(1), client_tracker.bookkeep_join_set()) + .await + .expect("bookkeeping should process the closed task") + .expect("closed task should return client id"); + assert_eq!(closed_client_id.0, ClientId("client-1".to_string())); + client_tracker + .close_client(&closed_client_id) + .await + .expect("closed client should emit connection closed"); + + match transport_event_rx + .recv() + .await + .expect("connection closed should be sent") + { + TransportEvent::ConnectionClosed { + connection_id: closed_connection_id, + } => assert_eq!(closed_connection_id, connection_id), + other => panic!("expected connection closed, got {other:?}"), + } + } + + #[tokio::test] + async fn shutdown_cancels_blocked_outbound_forwarding() { + let (server_event_tx, _server_event_rx) = mpsc::channel(1); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let mut client_tracker = + ClientTracker::new(server_event_tx.clone(), transport_event_tx, &shutdown_token); + + server_event_tx + .send(QueuedServerEnvelope { + event: ServerEvent::Pong { + status: PongStatus::Unknown, + }, + client_id: ClientId("queued-client".to_string()), + stream_id: StreamId("queued-stream".to_string()), + write_complete_tx: None, + }) + .await + .expect("server event queue should accept prefill"); + + client_tracker + .handle_message(initialize_envelope("client-1")) + .await + .expect("initialize should open client"); + + let writer = match transport_event_rx + .recv() + .await + .expect("connection opened should be sent") + { + TransportEvent::ConnectionOpened { writer, .. } => writer, + other => panic!("expected connection opened, got {other:?}"), + }; + let _ = transport_event_rx + .recv() + .await + .expect("initialize should be forwarded"); + + writer + .send(QueuedOutgoingMessage::new( + OutgoingMessage::AppServerNotification(ServerNotificationEnvelope { + notification: ServerNotification::ConfigWarning(ConfigWarningNotification { + summary: "test".to_string(), + details: None, + path: None, + range: None, + }), + emitted_at_ms: Some(1_234), + }), + )) + .await + .expect("writer should accept queued message"); + + timeout(Duration::from_secs(1), client_tracker.shutdown()) + .await + .expect("shutdown should not hang on blocked server forwarding"); + } + + #[tokio::test] + async fn non_close_transport_event_send_times_out_when_queue_stays_full() { + let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (transport_event_tx, _transport_event_rx) = mpsc::channel(1); + let shutdown_token = CancellationToken::new(); + let client_tracker = + ClientTracker::new(server_event_tx, transport_event_tx.clone(), &shutdown_token); + + transport_event_tx + .send(TransportEvent::ConnectionClosed { + connection_id: next_connection_id(), + }) + .await + .expect("transport event queue should accept prefill"); + + let send_result = client_tracker + .send_transport_event(TransportEvent::IncomingMessage { + connection_id: next_connection_id(), + message: initialized_notification(), + }) + .await; + + assert!(send_result.is_err()); + } + + #[tokio::test] + async fn incoming_message_timeout_does_not_advance_seq_id() { + let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(2); + let shutdown_token = CancellationToken::new(); + let mut client_tracker = + ClientTracker::new(server_event_tx, transport_event_tx.clone(), &shutdown_token); + + client_tracker + .handle_message(initialize_envelope_with_stream_id( + "client-1", + Some("stream-1"), + )) + .await + .expect("initialize should open client"); + let connection_id = match transport_event_rx.recv().await.expect("open event") { + TransportEvent::ConnectionOpened { connection_id, .. } => connection_id, + other => panic!("expected connection opened, got {other:?}"), + }; + let _ = transport_event_rx.recv().await.expect("initialize event"); + + for _ in 0..2 { + transport_event_tx + .send(TransportEvent::ConnectionClosed { + connection_id: next_connection_id(), + }) + .await + .expect("transport event queue should accept prefill"); + } + + let retry_envelope = ClientEnvelope { + event: ClientEvent::ClientMessage { + message: initialized_notification(), + }, + client_id: ClientId("client-1".to_string()), + stream_id: Some(StreamId("stream-1".to_string())), + seq_id: Some(1), + cursor: None, + }; + assert!( + client_tracker + .handle_message(retry_envelope.clone()) + .await + .is_err() + ); + for _ in 0..2 { + let _ = transport_event_rx.recv().await.expect("prefilled event"); + } + + client_tracker + .handle_message(retry_envelope) + .await + .expect("retry should forward after timeout"); + match transport_event_rx.recv().await.expect("retried event") { + TransportEvent::IncomingMessage { + connection_id: queued_connection_id, + .. + } => assert_eq!(queued_connection_id, connection_id), + other => panic!("expected incoming message, got {other:?}"), + } + } + + #[tokio::test(start_paused = true)] + async fn initialize_timeout_closes_open_connection() { + let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(1); + let shutdown_token = CancellationToken::new(); + let client_tracker = + ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token); + let handle_message = tokio::spawn(async move { + let mut client_tracker = client_tracker; + client_tracker + .handle_message(initialize_envelope_with_stream_id( + "client-1", + Some("stream-1"), + )) + .await + }); + + tokio::task::yield_now().await; + tokio::time::advance( + REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT + Duration::from_millis(1), + ) + .await; + + assert!(handle_message.await.expect("handle message task").is_err()); + let connection_id = match transport_event_rx.recv().await.expect("open event") { + TransportEvent::ConnectionOpened { connection_id, .. } => connection_id, + other => panic!("expected connection opened, got {other:?}"), + }; + + match transport_event_rx.recv().await.expect("close event") { + TransportEvent::ConnectionClosed { + connection_id: closed_connection_id, + } => assert_eq!(closed_connection_id, connection_id), + other => panic!("expected connection closed, got {other:?}"), + } + } + + #[tokio::test] + async fn close_client_waits_for_transport_event_queue_capacity() { + let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(2); + let shutdown_token = CancellationToken::new(); + let mut client_tracker = + ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token); + + client_tracker + .handle_message(initialize_envelope_with_stream_id( + "client-1", + Some("stream-1"), + )) + .await + .expect("initialize should open client"); + let connection_id = match transport_event_rx.recv().await.expect("open event") { + TransportEvent::ConnectionOpened { connection_id, .. } => connection_id, + other => panic!("expected connection opened, got {other:?}"), + }; + let _ = transport_event_rx.recv().await.expect("initialize event"); + + for _ in 0..2 { + client_tracker + .transport_event_tx + .send(TransportEvent::IncomingMessage { + connection_id, + message: initialized_notification(), + }) + .await + .expect("transport event queue should accept prefill"); + } + + let client_key = ( + ClientId("client-1".to_string()), + StreamId("stream-1".to_string()), + ); + let close_client = client_tracker.close_client(&client_key); + tokio::pin!(close_client); + assert!( + timeout(Duration::from_millis(20), &mut close_client) + .await + .is_err() + ); + + for _ in 0..2 { + match transport_event_rx.recv().await.expect("prefilled event") { + TransportEvent::IncomingMessage { + connection_id: queued_connection_id, + .. + } => assert_eq!(queued_connection_id, connection_id), + other => panic!("expected incoming message, got {other:?}"), + } + } + + close_client + .await + .expect("close should forward after queue drains"); + match transport_event_rx.recv().await.expect("close event") { + TransportEvent::ConnectionClosed { + connection_id: closed_connection_id, + } => assert_eq!(closed_connection_id, connection_id), + other => panic!("expected connection closed, got {other:?}"), + } + } + + #[tokio::test] + async fn close_client_keeps_forwarding_after_caller_is_aborted() { + let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(2); + let shutdown_token = CancellationToken::new(); + let mut client_tracker = + ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token); + + client_tracker + .handle_message(initialize_envelope_with_stream_id( + "client-1", + Some("stream-1"), + )) + .await + .expect("initialize should open client"); + let connection_id = match transport_event_rx.recv().await.expect("open event") { + TransportEvent::ConnectionOpened { connection_id, .. } => connection_id, + other => panic!("expected connection opened, got {other:?}"), + }; + let _ = transport_event_rx.recv().await.expect("initialize event"); + + for _ in 0..2 { + client_tracker + .transport_event_tx + .send(TransportEvent::IncomingMessage { + connection_id, + message: initialized_notification(), + }) + .await + .expect("transport event queue should accept prefill"); + } + + let client_key = ( + ClientId("client-1".to_string()), + StreamId("stream-1".to_string()), + ); + let mut close_client = + tokio::spawn(async move { client_tracker.close_client(&client_key).await }); + assert!( + timeout(Duration::from_millis(20), &mut close_client) + .await + .is_err() + ); + close_client.abort(); + let _ = close_client.await; + + for _ in 0..2 { + let _ = transport_event_rx.recv().await.expect("prefilled event"); + } + match timeout(Duration::from_secs(1), transport_event_rx.recv()) + .await + .expect("close should be delivered") + .expect("close event") + { + TransportEvent::ConnectionClosed { + connection_id: closed_connection_id, + } => assert_eq!(closed_connection_id, connection_id), + other => panic!("expected connection closed, got {other:?}"), + } + } + + #[tokio::test] + async fn initialize_with_new_stream_id_opens_new_connection_for_same_client() { + let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let mut client_tracker = + ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token); + + client_tracker + .handle_message(initialize_envelope_with_stream_id( + "client-1", + Some("stream-1"), + )) + .await + .expect("first initialize should open client"); + let first_connection_id = match transport_event_rx.recv().await.expect("open event") { + TransportEvent::ConnectionOpened { connection_id, .. } => connection_id, + other => panic!("expected connection opened, got {other:?}"), + }; + let _ = transport_event_rx.recv().await.expect("initialize event"); + + client_tracker + .handle_message(initialize_envelope_with_stream_id( + "client-1", + Some("stream-2"), + )) + .await + .expect("second initialize should open client"); + let second_connection_id = match transport_event_rx.recv().await.expect("open event") { + TransportEvent::ConnectionOpened { connection_id, .. } => connection_id, + other => panic!("expected connection opened, got {other:?}"), + }; + + assert_ne!(first_connection_id, second_connection_id); + } + + #[tokio::test] + async fn legacy_initialize_without_stream_id_resets_inbound_seq_id() { + let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (transport_event_tx, mut transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let mut client_tracker = + ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token); + + client_tracker + .handle_message(initialize_envelope("client-1")) + .await + .expect("initialize should open client"); + let connection_id = match transport_event_rx.recv().await.expect("open event") { + TransportEvent::ConnectionOpened { connection_id, .. } => connection_id, + other => panic!("expected connection opened, got {other:?}"), + }; + let _ = transport_event_rx.recv().await.expect("initialize event"); + + client_tracker + .handle_message(ClientEnvelope { + event: ClientEvent::ClientMessage { + message: JSONRPCMessage::Notification( + codex_app_server_protocol::JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }, + ), + }, + client_id: ClientId("client-1".to_string()), + stream_id: None, + seq_id: Some(0), + cursor: None, + }) + .await + .expect("legacy followup should be forwarded"); + + match transport_event_rx.recv().await.expect("followup event") { + TransportEvent::IncomingMessage { + connection_id: incoming_connection_id, + .. + } => assert_eq!(incoming_connection_id, connection_id), + other => panic!("expected incoming message, got {other:?}"), + } + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/clients.rs b/codex-rs/app-server-transport/src/transport/remote_control/clients.rs new file mode 100644 index 0000000000000000000000000000000000000000..37ec1c1a58272fe45e681488af0ed342142a60f3 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/clients.rs @@ -0,0 +1,303 @@ +use super::auth::RemoteControlAuth; +use super::auth::RemoteControlConnectionAuth; +use super::auth::load_remote_control_auth; +use super::auth::recover_remote_control_auth; +use super::enroll::format_headers; +use super::enroll::preview_remote_control_response_body; +use super::protocol::normalize_remote_control_base_url; +use axum::http::HeaderMap; +use codex_app_server_protocol::RemoteControlClient; +use codex_app_server_protocol::RemoteControlClientsListOrder; +use codex_app_server_protocol::RemoteControlClientsListParams; +use codex_app_server_protocol::RemoteControlClientsListResponse; +use codex_app_server_protocol::RemoteControlClientsRevokeParams; +use codex_app_server_protocol::RemoteControlClientsRevokeResponse; +use codex_login::default_client::create_client_without_request_logging; +use serde::Deserialize; +use std::io; +use std::io::ErrorKind; +use time::OffsetDateTime; +use time::format_description::well_known::Rfc3339; +use url::Url; + +const REMOTE_CONTROL_CLIENT_MANAGEMENT_TIMEOUT: std::time::Duration = + std::time::Duration::from_secs(30); + +#[derive(Debug, Deserialize)] +struct ListRemoteControlClientsResponse { + items: Vec, + #[serde(default)] + cursor: Option, +} + +#[derive(Debug, Deserialize)] +struct RemoteControlClientResponse { + client_id: String, + #[serde(default)] + display_name: Option, + #[serde(default)] + device_type: Option, + #[serde(default)] + platform: Option, + #[serde(default)] + os_version: Option, + #[serde(default)] + device_model: Option, + #[serde(default)] + app_version: Option, + #[serde(default)] + last_seen_at: Option, +} + +enum ClientManagementRequest<'a> { + List { + url: &'a Url, + params: &'a RemoteControlClientsListParams, + }, + Revoke { + url: &'a Url, + }, +} + +struct ClientManagementResponse { + status: axum::http::StatusCode, + headers: HeaderMap, + body: Vec, +} + +pub(super) async fn list_remote_control_clients( + remote_control_url: &str, + auth_manager: &RemoteControlAuth, + params: RemoteControlClientsListParams, +) -> io::Result { + if params.environment_id.is_empty() { + return Err(io::Error::new( + ErrorKind::InvalidInput, + "remote control client list requires environmentId", + )); + } + if params + .limit + .is_some_and(|limit| !(1..=100).contains(&limit)) + { + return Err(io::Error::new( + ErrorKind::InvalidInput, + "remote control client list limit must be between 1 and 100", + )); + } + let url = environment_clients_url(remote_control_url, ¶ms.environment_id)?; + let response = send_client_management_request( + auth_manager, + ClientManagementRequest::List { + url: &url, + params: ¶ms, + }, + "list remote control clients", + ) + .await?; + let ClientManagementResponse { + status, + headers, + body, + } = response; + let body_preview = preview_remote_control_response_body(&body); + ensure_success_response(status, &headers, &url, &body_preview, "client list")?; + let response = serde_json::from_slice::(&body).map_err( + |err| { + io::Error::other(format!( + "failed to parse remote control client list response from `{url}`: HTTP {status}, {}, body: {body_preview}, decode error: {err}", + format_headers(&headers) + )) + }, + )?; + Ok(RemoteControlClientsListResponse { + data: response + .items + .into_iter() + .map(RemoteControlClient::try_from) + .collect::>()?, + next_cursor: response.cursor, + }) +} + +pub(super) async fn revoke_remote_control_client( + remote_control_url: &str, + auth_manager: &RemoteControlAuth, + params: RemoteControlClientsRevokeParams, +) -> io::Result { + if params.environment_id.is_empty() { + return Err(io::Error::new( + ErrorKind::InvalidInput, + "remote control client revoke requires environmentId", + )); + } + if params.client_id.is_empty() { + return Err(io::Error::new( + ErrorKind::InvalidInput, + "remote control client revoke requires clientId", + )); + } + let mut url = environment_clients_url(remote_control_url, ¶ms.environment_id)?; + url.path_segments_mut() + .map_err(|()| { + io::Error::new( + ErrorKind::InvalidInput, + "remote control URL cannot be a base", + ) + })? + .push(¶ms.client_id); + let response = send_client_management_request( + auth_manager, + ClientManagementRequest::Revoke { url: &url }, + "revoke remote control client", + ) + .await?; + let ClientManagementResponse { + status, + headers, + body, + } = response; + let body_preview = preview_remote_control_response_body(&body); + ensure_success_response(status, &headers, &url, &body_preview, "client revoke")?; + Ok(RemoteControlClientsRevokeResponse {}) +} + +async fn send_client_management_request( + auth_manager: &RemoteControlAuth, + request: ClientManagementRequest<'_>, + action: &str, +) -> io::Result { + let mut auth_recovery = auth_manager.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + let auth = load_remote_control_auth(auth_manager).await?; + let response = send_client_management_request_once(&auth, &request, action).await?; + if response.status.as_u16() != 401 + || !recover_remote_control_auth(&mut auth_recovery, &mut auth_change_rx).await + { + return Ok(response); + } + let auth = load_remote_control_auth(auth_manager).await?; + send_client_management_request_once(&auth, &request, action).await +} + +async fn send_client_management_request_once( + auth: &RemoteControlConnectionAuth, + request: &ClientManagementRequest<'_>, + action: &str, +) -> io::Result { + let client = create_client_without_request_logging(); + let auth_headers = auth.request_headers()?; + let request = match request { + ClientManagementRequest::List { url, params } => { + let mut query = Vec::new(); + if let Some(cursor) = ¶ms.cursor { + query.push(("cursor", cursor.clone())); + } + if let Some(limit) = params.limit { + query.push(("limit", limit.to_string())); + } + if let Some(order) = params.order { + query.push(( + "order", + match order { + RemoteControlClientsListOrder::Asc => "asc", + RemoteControlClientsListOrder::Desc => "desc", + } + .to_string(), + )); + } + client.get((*url).clone()).query(&query) + } + ClientManagementRequest::Revoke { url } => client.delete((*url).clone()), + }; + let response = request + .timeout(REMOTE_CONTROL_CLIENT_MANAGEMENT_TIMEOUT) + .headers(auth_headers) + .send() + .await + .map_err(|err| io::Error::other(format!("failed to {action}: {err}")))?; + let headers = response.headers().clone(); + let status = response.status(); + let body = response + .bytes() + .await + .map_err(|err| io::Error::other(format!("failed to read {action} response: {err}")))? + .to_vec(); + Ok(ClientManagementResponse { + status, + headers, + body, + }) +} + +fn ensure_success_response( + status: axum::http::StatusCode, + headers: &HeaderMap, + url: &Url, + body_preview: &str, + response_kind: &str, +) -> io::Result<()> { + if status.is_success() { + return Ok(()); + } + let error_kind = match status.as_u16() { + 400 => ErrorKind::InvalidInput, + 401 | 403 => ErrorKind::PermissionDenied, + 404 => ErrorKind::NotFound, + _ => ErrorKind::Other, + }; + Err(io::Error::new( + error_kind, + format!( + "remote control {response_kind} failed at `{url}`: HTTP {status}, {}, body: {body_preview}", + format_headers(headers) + ), + )) +} + +fn environment_clients_url(remote_control_url: &str, environment_id: &str) -> io::Result { + let mut url = normalize_remote_control_base_url(remote_control_url)? + .join("wham/remote/control/environments") + .map_err(io::Error::other)?; + url.path_segments_mut() + .map_err(|()| { + io::Error::new( + ErrorKind::InvalidInput, + "remote control URL cannot be a base", + ) + })? + .push(environment_id) + .push("clients"); + Ok(url) +} + +impl TryFrom for RemoteControlClient { + type Error = io::Error; + + fn try_from(client: RemoteControlClientResponse) -> Result { + Ok(Self { + client_id: client.client_id, + display_name: client.display_name, + device_type: client.device_type, + platform: client.platform, + os_version: client.os_version, + device_model: client.device_model, + app_version: client.app_version, + last_seen_at: client + .last_seen_at + .map(|last_seen_at| { + OffsetDateTime::parse(&last_seen_at, &Rfc3339) + .map(OffsetDateTime::unix_timestamp) + .map_err(|err| { + io::Error::new( + ErrorKind::InvalidData, + format!( + "failed to parse remote control client last_seen_at `{last_seen_at}`: {err}" + ), + ) + }) + }) + .transpose()?, + }) + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/controller.rs b/codex-rs/app-server-transport/src/transport/remote_control/controller.rs new file mode 100644 index 0000000000000000000000000000000000000000..e7fb3c3091be5641643395403d4428e81ff1c715 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/controller.rs @@ -0,0 +1,368 @@ +//! Owns the current login's relay session and the process-facing handle. +//! Replacement happens before cleanup; retired sessions cannot publish into their replacements. + +use super::auth::RemoteControlAuth; +use super::*; +use futures::FutureExt; +use std::panic::AssertUnwindSafe; +use tokio_util::task::TaskTracker; + +#[derive(Clone)] +pub struct RemoteControlHandle { + pub(super) inner: Arc, +} + +struct CurrentSession { + authenticated: bool, + session: Arc, +} + +pub(super) struct RemoteControl { + config: RemoteControlStartConfig, + state_db: Option>, + auth_manager: Arc, + transport_event_tx: mpsc::Sender, + shutdown: CancellationToken, + tasks: TaskTracker, + current: StdMutex>, + startup: RemoteControlDesiredState, + persistence: RemoteControlPersistence, + client_name: RemoteControlPairingPersistenceKey, + requires_client_name: bool, + status: watch::Sender, + session_changed: watch::Sender<()>, +} + +impl RemoteControl { + pub(super) fn session(&self) -> Arc { + let mut current = self + .current + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(current) = current.as_ref() + && current.session.auth_manager.owner.is_current() + { + return current.session.clone(); + } + let (auth, authenticated) = RemoteControlAuth::capture(self.auth_manager.clone()); + let desired = match current.take() { + Some(previous) => { + previous.session.shutdown_token.cancel(); + if previous.authenticated { + RemoteControlDesiredState::Disabled + } else { + // Startup or an explicit enable while signed out may wait for first login. + *previous.session.desired_state_tx.borrow() + } + } + None => self.startup, + }; + let session = self.start_session(auth, desired); + self.status.send_replace(session.status()); + *current = Some(CurrentSession { + authenticated, + session: session.clone(), + }); + self.session_changed.send_replace(()); + session + } + + fn start_session( + &self, + auth_manager: RemoteControlAuth, + desired: RemoteControlDesiredState, + ) -> Arc { + let shutdown = self.shutdown.child_token(); + let (desired_state_tx, _) = watch::channel(desired); + let desired_state_tx = Arc::new(desired_state_tx); + let current_enrollment = + Arc::new(RemoteControlEnrollmentState::new(/*enrollment*/ None)); + let server_name = gethostname().to_string_lossy().trim().to_string(); + let (status_tx, _) = watch::channel(RemoteControlStatusChangedNotification { + status: if desired.is_enabled() { + RemoteControlConnectionStatus::Connecting + } else { + RemoteControlConnectionStatus::Disabled + }, + server_name: server_name.clone(), + installation_id: self.config.installation_id.clone(), + environment_id: None, + }); + let session = Arc::new(RemoteControlSession { + policy: self.config.policy, + shutdown_token: shutdown.clone(), + desired_state_tx: desired_state_tx.clone(), + desired_state_rpc_lock: Arc::new(Semaphore::new(1)), + persistence: self.persistence.clone(), + status_tx: Arc::new(status_tx.clone()), + state_db: self.state_db.clone(), + remote_control_url: self.config.remote_control_url.clone(), + current_enrollment: current_enrollment.clone(), + pairing_persistence_key: self.client_name.clone(), + pairing_persistence_key_required: self.requires_client_name, + auth_manager: auth_manager.clone(), + }); + let websocket = RemoteControlWebsocket::new( + websocket::RemoteControlWebsocketConfig { + remote_control_url: self.config.remote_control_url.clone(), + installation_id: self.config.installation_id.clone(), + remote_control_target: None, + server_name, + }, + self.state_db.clone(), + auth_manager, + RemoteControlChannels { + transport_event_tx: self.transport_event_tx.clone(), + status_publisher: RemoteControlStatusPublisher::new(status_tx), + current_enrollment, + pairing_persistence_key: self.client_name.clone(), + persistence: self.persistence.clone(), + }, + shutdown.clone(), + desired_state_tx, + ); + let client_name_rx = if self.requires_client_name { + let (tx, rx) = oneshot::channel(); + let mut names = self.client_name.subscribe(); + self.tasks.spawn(async move { + tokio::select! { + _ = shutdown.cancelled() => {} + name = names.wait_for(Option::is_some) => { + if let Ok(name) = name + && let Some(name) = name.as_ref() + { + let _ = tx.send(name.clone()); + } + } + } + }); + Some(rx) + } else { + None + }; + let process_shutdown = self.shutdown.clone(); + let failed_session = session.clone(); + self.tasks.spawn(async move { + if let Err(panic) = AssertUnwindSafe(websocket.run(client_name_rx)) + .catch_unwind() + .await + { + tracing::error!("remote control websocket task panicked"); + failed_session.publish_status(RemoteControlConnectionStatus::Disabled); + process_shutdown.cancel(); + std::panic::resume_unwind(panic); + } + }); + session + } +} + +impl RemoteControlSession { + async fn run( + &self, + operation: impl std::future::Future>, + ) -> io::Result { + self.auth_manager.ensure_current()?; + tokio::select! { + biased; + _ = self.auth_manager.owner.invalidated() => { + Err(io::Error::new(io::ErrorKind::Interrupted, "remote control authentication changed")) + } + result = operation => { + self.auth_manager.ensure_current()?; + result + } + } + } +} + +impl RemoteControlHandle { + pub fn ensure_remote_control_allowed(&self) -> Result<(), RemoteControlDisabledByRequirements> { + self.inner.session().ensure_remote_control_allowed() + } + + pub fn status(&self) -> RemoteControlStatusChangedNotification { + self.inner.session().status() + } + + pub fn status_receiver(&self) -> watch::Receiver { + self.inner.session(); + self.inner.status.subscribe() + } + + pub fn enable_ephemeral( + &self, + ) -> Result { + self.inner.session().enable_ephemeral() + } + + pub async fn disable_ephemeral(&self) -> RemoteControlStatusChangedNotification { + self.inner.session().disable_ephemeral().await + } + + pub async fn enable( + &self, + app_server_client_name: Option<&str>, + ) -> io::Result { + let session = self.inner.session(); + session.run(session.enable(app_server_client_name)).await + } + + pub async fn disable( + &self, + app_server_client_name: Option<&str>, + ) -> io::Result { + let session = self.inner.session(); + session.run(session.disable(app_server_client_name)).await + } + + pub async fn resolve_persisted_preference( + &self, + app_server_client_name: Option<&str>, + ) -> io::Result { + let session = self.inner.session(); + session + .run(session.resolve_persisted_preference(app_server_client_name)) + .await + } + + pub async fn start_pairing( + &self, + params: RemoteControlPairingStartParams, + app_server_client_name: Option<&str>, + ) -> io::Result { + let session = self.inner.session(); + session + .run(session.start_pairing(params, app_server_client_name)) + .await + } + + pub async fn pairing_status( + &self, + params: RemoteControlPairingStatusParams, + ) -> io::Result { + let session = self.inner.session(); + session.run(session.pairing_status(params)).await + } + + pub async fn list_clients( + &self, + params: RemoteControlClientsListParams, + ) -> io::Result { + let session = self.inner.session(); + session.run(session.list_clients(params)).await + } + + pub async fn revoke_client( + &self, + params: RemoteControlClientsRevokeParams, + ) -> io::Result { + let session = self.inner.session(); + session.run(session.revoke_client(params)).await + } +} + +pub async fn start_remote_control( + config: RemoteControlStartConfig, + state_db: Option>, + auth_manager: Arc, + transport_event_tx: mpsc::Sender, + shutdown_token: CancellationToken, + app_server_client_name_rx: Option>, + startup_mode: RemoteControlStartupMode, +) -> io::Result<(JoinHandle<()>, RemoteControlHandle)> { + let startup = + if config.policy == RemoteControlPolicy::DisabledByRequirements || state_db.is_none() { + RemoteControlDesiredState::Disabled + } else { + match startup_mode { + RemoteControlStartupMode::ResolvePersisted => RemoteControlDesiredState::Unknown, + RemoteControlStartupMode::DisabledEphemeral => RemoteControlDesiredState::Disabled, + RemoteControlStartupMode::EnabledEphemeral => { + normalize_remote_control_url(&config.remote_control_url)?; + RemoteControlDesiredState::Enabled { + persistence_preference: None, + } + } + } + }; + let (status, _) = watch::channel(RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Disabled, + server_name: gethostname().to_string_lossy().trim().to_string(), + installation_id: config.installation_id.clone(), + environment_id: None, + }); + let inner = Arc::new(RemoteControl { + config, + state_db, + auth_manager, + transport_event_tx, + shutdown: shutdown_token, + tasks: TaskTracker::new(), + current: StdMutex::new(None), + startup, + persistence: RemoteControlPersistence::default(), + client_name: watch::channel(None).0, + requires_client_name: app_server_client_name_rx.is_some(), + status, + session_changed: watch::channel(()).0, + }); + inner.session(); + if let Some(rx) = app_server_client_name_rx { + let names = inner.client_name.clone(); + let shutdown = inner.shutdown.clone(); + inner.tasks.spawn(async move { + tokio::select! { + _ = shutdown.cancelled() => {} + name = rx => match name { + Ok(name) => { names.send_replace(Some(name)); } + Err(_) => shutdown.cancel(), + } + } + }); + } + let handle = RemoteControlHandle { + inner: inner.clone(), + }; + let task = tokio::spawn(async move { + let mut session_changed = inner.session_changed.subscribe(); + loop { + // Reconcile on both API entry and notification delivery, so watcher latency is harmless. + let session = inner.session(); + let mut status = session.status_receiver(); + session_changed.borrow_and_update(); + { + let current = inner + .current + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if current + .as_ref() + .is_some_and(|current| Arc::ptr_eq(¤t.session, &session)) + && session.auth_manager.owner.is_current() + { + inner.status.send_if_modified(|current| { + let next = status.borrow_and_update().clone(); + if *current == next { + return false; + } + *current = next; + true + }); + } + } + tokio::select! { + biased; + _ = inner.shutdown.cancelled() => break, + _ = session.auth_manager.owner.invalidated() => {} + _ = session_changed.changed() => {} + _ = status.changed() => {} + } + } + inner.tasks.close(); + inner.tasks.wait().await; + inner.persistence.tasks.close(); + inner.persistence.tasks.wait().await; + }); + Ok((task, handle)) +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/desired_state.rs b/codex-rs/app-server-transport/src/transport/remote_control/desired_state.rs new file mode 100644 index 0000000000000000000000000000000000000000..86c41aa4d0b54d4198e02aa849d16184b66e25f2 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/desired_state.rs @@ -0,0 +1,155 @@ +use super::RemoteControlEnableError; +use super::RemoteControlSession; +use super::RemoteControlUnavailable; +use super::protocol::normalize_remote_control_url; +use super::publish_current_enrollment; +use super::websocket::RemoteControlStatusPublisher; +use codex_app_server_protocol::RemoteControlStatusChangedNotification; +use codex_state::RemoteControlEnrollmentRecord; +use std::io; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum RemoteControlDesiredState { + // `Unknown` exists only on plain startup before auth and enrollment scope resolve. Persisted + // `1` is `Enabled { persistence_preference: Some(true) }`; `0`, `NULL`, or no row are + // `Disabled`. Runtime-only enable is `Enabled { persistence_preference: None }`, so new rows + // keep `NULL`; durable RPC enable uses `Some(true)`, so new rows get `1`. Durable disable writes + // `0` before entering `Disabled`; runtime-only disable does not write. `Disabled` carries no + // preference because disabled sessions do not create enrollments. + Unknown, + Disabled, + Enabled { + persistence_preference: Option, + }, +} +impl RemoteControlDesiredState { + pub(super) fn is_enabled(self) -> bool { + matches!(self, Self::Enabled { .. }) + } +} + +pub(super) fn desired_state_from_persisted_enrollment( + enrollment: Option, +) -> RemoteControlDesiredState { + if enrollment.and_then(|enrollment| enrollment.remote_control_enabled) == Some(true) { + RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + } + } else { + RemoteControlDesiredState::Disabled + } +} + +impl RemoteControlSession { + pub async fn resolve_persisted_preference( + &self, + app_server_client_name: Option<&str>, + ) -> io::Result { + if self.ensure_remote_control_allowed().is_err() { + return Ok(false); + } + let _transition = self + .desired_state_rpc_lock + .acquire() + .await + .unwrap_or_else(|_| unreachable!()); + if !matches!( + *self.desired_state_tx.borrow(), + RemoteControlDesiredState::Unknown + ) { + return Ok(self.desired_state_tx.borrow().is_enabled()); + } + + let state_db = self + .state_db + .as_deref() + .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, RemoteControlUnavailable))?; + let auth = super::auth::load_remote_control_auth(&self.auth_manager).await?; + let remote_control_target = normalize_remote_control_url(&self.remote_control_url)?; + let app_server_client_name = self.pairing_persistence_key(app_server_client_name)?; + let _persistence = + super::persistence::read_lock(&self.auth_manager, &self.persistence).await?; + let enrollment = state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + &auth.account_id, + app_server_client_name.as_deref(), + ) + .await + .map_err(io::Error::other)?; + let desired_state = desired_state_from_persisted_enrollment(enrollment); + self.desired_state_tx.send_if_modified(|state| { + if !matches!(*state, RemoteControlDesiredState::Unknown) { + return false; + } + *state = desired_state; + true + }); + Ok(self.desired_state_tx.borrow().is_enabled()) + } + + pub async fn enable( + &self, + app_server_client_name: Option<&str>, + ) -> io::Result { + self.ensure_remote_control_allowed() + .map_err(|err| io::Error::new(io::ErrorKind::PermissionDenied, err))?; + let _transition = self + .desired_state_rpc_lock + .acquire() + .await + .unwrap_or_else(|_| unreachable!()); + let state_db = self + .state_db + .as_deref() + .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, RemoteControlUnavailable))?; + let mut auth = super::auth::load_remote_control_auth(&self.auth_manager).await?; + let remote_control_target = normalize_remote_control_url(&self.remote_control_url)?; + let app_server_client_name = self.pairing_persistence_key(app_server_client_name)?; + let app_server_client_name = app_server_client_name.as_deref(); + let status = self.status(); + let mut current_enrollment = self.current_enrollment.lock().await; + let (enrollment, _) = self + .load_or_enroll_server( + ¤t_enrollment, + &mut auth, + &status.installation_id, + &status.server_name, + app_server_client_name, + super::RemoteControlEnrollmentSelection::ReuseOrCreate, + ) + .await?; + + let current_auth = super::auth::load_remote_control_auth(&self.auth_manager).await?; + if current_auth.account_id != auth.account_id { + return Err(io::Error::new( + io::ErrorKind::Interrupted, + "remote control account changed during enrollment", + )); + } + + self.set_preference( + state_db, + &remote_control_target, + &auth.account_id, + app_server_client_name, + /*enabled*/ true, + Some(&enrollment), + ) + .await?; + publish_current_enrollment(&mut current_enrollment, &enrollment); + self.enable_with_preference(Some(true)).map_err(|err| { + let kind = match err { + RemoteControlEnableError::Unavailable(_) => io::ErrorKind::NotFound, + RemoteControlEnableError::AuthenticationChanged => io::ErrorKind::Interrupted, + RemoteControlEnableError::DisabledByRequirements(_) => { + io::ErrorKind::PermissionDenied + } + }; + io::Error::new(kind, err) + })?; + RemoteControlStatusPublisher::new(self.status_tx.as_ref().clone()) + .publish_environment_id(Some(enrollment.environment_id)); + Ok(self.status()) + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/enroll.rs b/codex-rs/app-server-transport/src/transport/remote_control/enroll.rs new file mode 100644 index 0000000000000000000000000000000000000000..638d797cf3a8c926925bda6bd67c4e2574bfd227 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/enroll.rs @@ -0,0 +1,770 @@ +use super::pairing_unavailable_error; +use super::protocol::RemoteControlPairingStatusRequest; +use super::protocol::RemoteControlPairingStatusResponse as BackendRemoteControlPairingStatusResponse; +use super::protocol::RemoteControlTarget; +use super::protocol::StartRemoteControlPairingRequest; +use super::protocol::StartRemoteControlPairingResponse; +use super::server_api::RemoteControlServerRequestError; +use super::server_api::retry_after_with_jitter; +use axum::http::HeaderMap; +use axum::http::StatusCode; +use codex_app_server_protocol::RemoteControlPairingStartResponse; +use codex_app_server_protocol::RemoteControlPairingStatusResponse; +use codex_login::default_client::create_client_without_request_logging; +use codex_state::RemoteControlEnrollmentRecord; +use codex_state::StateRuntime; +use std::io; +use std::io::ErrorKind; +use time::OffsetDateTime; +use time::format_description::well_known::Rfc3339; +use tracing::info; +use tracing::warn; + +const REMOTE_CONTROL_PAIRING_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30); +const REMOTE_CONTROL_RESPONSE_BODY_MAX_BYTES: usize = 4096; +const REMOTE_CONTROL_SERVER_TOKEN_REFRESH_SKEW_SECS: i64 = 5 * 60; + +const REQUEST_ID_HEADER: &str = "x-request-id"; +const OAI_REQUEST_ID_HEADER: &str = "x-oai-request-id"; +const CF_RAY_HEADER: &str = "cf-ray"; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(super) struct RemoteControlEnrollment { + pub(super) remote_control_target: RemoteControlTarget, + pub(super) account_id: String, + pub(super) environment_id: String, + pub(super) server_id: String, + pub(super) server_name: String, + pub(super) remote_control_token: Option, + pub(super) expires_at: Option, + pub(super) next_refresh_at: Option, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(super) enum RemoteControlServerTokenRefreshRequirement { + Required, + Proactive, + NotNeeded, +} + +impl RemoteControlEnrollment { + pub(super) async fn start_pairing( + &self, + request: StartRemoteControlPairingRequest, + ) -> io::Result { + if self.server_token_refresh_requirement() + == RemoteControlServerTokenRefreshRequirement::Required + { + return Err(pairing_unavailable_error()); + } + let remote_control_token = self + .remote_control_token + .as_deref() + .ok_or_else(pairing_unavailable_error)?; + + let response = create_client_without_request_logging() + .post(&self.remote_control_target.pair_url) + .timeout(REMOTE_CONTROL_PAIRING_TIMEOUT) + .bearer_auth(remote_control_token) + .json(&request) + .send() + .await + .map_err(|err| { + io::Error::other(format!( + "failed to start remote control pairing at `{}`: {err}", + self.remote_control_target.pair_url + )) + })?; + let headers = response.headers().clone(); + let status = response.status(); + let retry_at = retry_after_with_jitter(&headers, OffsetDateTime::now_utc()); + let body = response.bytes().await.map_err(|err| { + pairing_response_error( + format!( + "failed to read remote control pairing response from `{}`: {err}", + self.remote_control_target.pair_url + ), + status, + retry_at, + ErrorKind::Other, + ) + })?; + let body_preview = preview_remote_control_response_body(&body); + if !status.is_success() { + let error_kind = match status.as_u16() { + 401 | 403 => ErrorKind::PermissionDenied, + 404 => ErrorKind::NotFound, + _ => ErrorKind::Other, + }; + return Err(pairing_response_error( + format!( + "remote control pairing failed at `{}`: HTTP {status}, {}, body: {body_preview}", + self.remote_control_target.pair_url, + format_headers(&headers) + ), + status, + retry_at, + error_kind, + )); + } + + let pairing = serde_json::from_slice::(&body).map_err( + |err| { + io::Error::other(format!( + "failed to parse remote control pairing response from `{}`: HTTP {status}, {}, body: {body_preview}, decode error: {err}", + self.remote_control_target.pair_url, + format_headers(&headers) + )) + }, + )?; + let StartRemoteControlPairingResponse { + pairing_code, + manual_pairing_code, + server_id, + environment_id, + expires_at, + } = pairing; + if server_id != self.server_id || environment_id != self.environment_id { + return Err(io::Error::other(format!( + "remote control pairing returned mismatched enrollment: expected server_id={}, environment_id={}; got server_id={}, environment_id={}", + self.server_id, self.environment_id, server_id, environment_id + ))); + } + let expires_at = OffsetDateTime::parse(&expires_at, &Rfc3339) + .map_err(|err| { + io::Error::new( + ErrorKind::InvalidData, + format!( + "failed to parse remote control pairing response from `{}`: HTTP {status}, {}, body: {body_preview}, expires_at parse error: {err}", + self.remote_control_target.pair_url, + format_headers(&headers) + ), + ) + })? + .unix_timestamp(); + + Ok(RemoteControlPairingStartResponse { + pairing_code, + manual_pairing_code, + environment_id, + expires_at, + }) + } + + pub(super) async fn pairing_status( + &self, + request: RemoteControlPairingStatusRequest, + ) -> io::Result { + if self.server_token_refresh_requirement() + == RemoteControlServerTokenRefreshRequirement::Required + { + return Err(pairing_unavailable_error()); + } + let remote_control_token = self + .remote_control_token + .as_deref() + .ok_or_else(pairing_unavailable_error)?; + + let response = create_client_without_request_logging() + .post(&self.remote_control_target.pair_status_url) + .timeout(REMOTE_CONTROL_PAIRING_TIMEOUT) + .bearer_auth(remote_control_token) + .json(&request) + .send() + .await + .map_err(|err| { + io::Error::other(format!( + "failed to check remote control pairing status at `{}`: {err}", + self.remote_control_target.pair_status_url + )) + })?; + let headers = response.headers().clone(); + let status = response.status(); + let retry_at = retry_after_with_jitter(&headers, OffsetDateTime::now_utc()); + let body = response.bytes().await.map_err(|err| { + pairing_response_error( + format!( + "failed to read remote control pairing status response from `{}`: {err}", + self.remote_control_target.pair_status_url + ), + status, + retry_at, + ErrorKind::Other, + ) + })?; + let body_preview = preview_remote_control_response_body(&body); + if !status.is_success() { + let error_kind = match status.as_u16() { + 401 | 403 => ErrorKind::PermissionDenied, + 404 | 410 => ErrorKind::InvalidInput, + _ => ErrorKind::Other, + }; + return Err(pairing_response_error( + format!( + "remote control pairing status failed at `{}`: HTTP {status}, {}, body: {body_preview}", + self.remote_control_target.pair_status_url, + format_headers(&headers) + ), + status, + retry_at, + error_kind, + )); + } + + let response = serde_json::from_slice::(&body) + .map_err(|err| { + io::Error::other(format!( + "failed to parse remote control pairing status response from `{}`: HTTP {status}, {}, body: {body_preview}, decode error: {err}", + self.remote_control_target.pair_status_url, + format_headers(&headers) + )) + })?; + Ok(RemoteControlPairingStatusResponse { + claimed: response.claimed, + }) + } + + pub(super) fn server_token_refresh_requirement( + &self, + ) -> RemoteControlServerTokenRefreshRequirement { + self.server_token_refresh_requirement_at(OffsetDateTime::now_utc()) + } + + pub(super) fn should_refresh_server_token(&self) -> bool { + self.server_token_refresh_requirement() + != RemoteControlServerTokenRefreshRequirement::NotNeeded + } + + pub(super) fn server_token_refresh_requirement_at( + &self, + now: OffsetDateTime, + ) -> RemoteControlServerTokenRefreshRequirement { + let Some(expires_at) = self.remote_control_token.as_ref().and(self.expires_at) else { + return RemoteControlServerTokenRefreshRequirement::Required; + }; + if expires_at <= now { + return RemoteControlServerTokenRefreshRequirement::Required; + } + if expires_at > now + time::Duration::seconds(REMOTE_CONTROL_SERVER_TOKEN_REFRESH_SKEW_SECS) + || self + .next_refresh_at + .is_some_and(|next_refresh_at| next_refresh_at > now) + { + return RemoteControlServerTokenRefreshRequirement::NotNeeded; + } + RemoteControlServerTokenRefreshRequirement::Proactive + } + + pub(super) fn clear_server_token(&mut self) { + self.remote_control_token = None; + self.expires_at = None; + } +} + +fn pairing_response_error( + message: String, + status: StatusCode, + retry_at: Option, + fallback_kind: ErrorKind, +) -> io::Error { + if matches!( + status, + StatusCode::TOO_MANY_REQUESTS | StatusCode::SERVICE_UNAVAILABLE + ) { + RemoteControlServerRequestError::io_error( + message, + Some(status), + retry_at, + /*timed_out*/ false, + ) + } else { + io::Error::new(fallback_kind, message) + } +} + +pub(super) async fn load_persisted_remote_control_enrollment( + state_db: Option<&StateRuntime>, + remote_control_target: &RemoteControlTarget, + account_id: &str, + app_server_client_name: Option<&str>, +) -> io::Result> { + let Some(state_db) = state_db else { + return Err(io::Error::new( + ErrorKind::NotFound, + format!( + "remote control enrollment cache unavailable because sqlite state db is disabled: websocket_url={}, account_id={}, app_server_client_name={:?}", + remote_control_target.websocket_url, account_id, app_server_client_name + ), + )); + }; + let enrollment = match state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + account_id, + app_server_client_name, + ) + .await + { + Ok(enrollment) => enrollment, + Err(err) => { + warn!( + "failed to load persisted remote control enrollment: websocket_url={}, account_id={}, app_server_client_name={:?}, err={err}", + remote_control_target.websocket_url, account_id, app_server_client_name + ); + return Err(io::Error::other(err)); + } + }; + + match enrollment { + Some(enrollment) => { + info!( + "reusing persisted remote control enrollment: websocket_url={}, account_id={}, app_server_client_name={:?}, server_id={}, environment_id={}", + remote_control_target.websocket_url, + account_id, + app_server_client_name, + enrollment.server_id, + enrollment.environment_id + ); + Ok(Some(RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: enrollment.account_id, + environment_id: enrollment.environment_id, + server_id: enrollment.server_id, + server_name: enrollment.server_name, + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + })) + } + None => { + info!( + "no persisted remote control enrollment found: websocket_url={}, account_id={}, app_server_client_name={:?}", + remote_control_target.websocket_url, account_id, app_server_client_name + ); + Ok(None) + } + } +} + +pub(super) async fn update_persisted_remote_control_enrollment( + state_db: Option<&StateRuntime>, + remote_control_target: &RemoteControlTarget, + account_id: &str, + app_server_client_name: Option<&str>, + enrollment: Option<&RemoteControlEnrollment>, + remote_control_enabled: Option, +) -> io::Result<()> { + let Some(state_db) = state_db else { + return Err(io::Error::new( + ErrorKind::NotFound, + format!( + "remote control enrollment persistence unavailable because sqlite state db is disabled: websocket_url={}, account_id={}, app_server_client_name={:?}, has_enrollment={}", + remote_control_target.websocket_url, + account_id, + app_server_client_name, + enrollment.is_some() + ), + )); + }; + if let &Some(enrollment) = &enrollment + && enrollment.account_id != account_id + { + return Err(io::Error::other(format!( + "enrollment account_id does not match expected account_id `{account_id}`" + ))); + } + + if let Some(enrollment) = enrollment { + state_db + .upsert_remote_control_enrollment(&RemoteControlEnrollmentRecord { + websocket_url: remote_control_target.websocket_url.clone(), + account_id: account_id.to_string(), + app_server_client_name: app_server_client_name.map(str::to_string), + server_id: enrollment.server_id.clone(), + environment_id: enrollment.environment_id.clone(), + server_name: enrollment.server_name.clone(), + remote_control_enabled, + }) + .await + .map_err(io::Error::other)?; + info!( + "persisted remote control enrollment: websocket_url={}, account_id={}, app_server_client_name={:?}, server_id={}, environment_id={}", + remote_control_target.websocket_url, + account_id, + app_server_client_name, + enrollment.server_id, + enrollment.environment_id + ); + Ok(()) + } else { + let rows_affected = state_db + .delete_remote_control_enrollment( + &remote_control_target.websocket_url, + account_id, + app_server_client_name, + ) + .await + .map_err(io::Error::other)?; + info!( + "cleared persisted remote control enrollment: websocket_url={}, account_id={}, app_server_client_name={:?}, rows_affected={rows_affected}", + remote_control_target.websocket_url, account_id, app_server_client_name + ); + Ok(()) + } +} + +pub(crate) fn preview_remote_control_response_body(body: &[u8]) -> String { + let body = String::from_utf8_lossy(body); + let trimmed = body.trim(); + if trimmed.is_empty() { + return "".to_string(); + } + let redacted = redact_remote_control_response_body(trimmed); + if redacted.len() <= REMOTE_CONTROL_RESPONSE_BODY_MAX_BYTES { + return redacted; + } + + let mut cut = REMOTE_CONTROL_RESPONSE_BODY_MAX_BYTES; + while !redacted.is_char_boundary(cut) { + cut = cut.saturating_sub(1); + } + let mut truncated = redacted[..cut].to_string(); + truncated.push_str("..."); + truncated +} + +fn redact_remote_control_response_body(body: &str) -> String { + let Ok(mut body_json) = serde_json::from_str::(body) else { + return body.to_string(); + }; + let Some(body_object) = body_json.as_object_mut() else { + return body.to_string(); + }; + for sensitive_field in [ + "remote_control_token", + "pairing_code", + "manual_pairing_code", + ] { + if let Some(value) = body_object.get_mut(sensitive_field) { + *value = serde_json::Value::String("".to_string()); + } + } + body_json.to_string() +} + +pub(crate) fn format_headers(headers: &HeaderMap) -> String { + let request_id_str = headers + .get(REQUEST_ID_HEADER) + .or_else(|| headers.get(OAI_REQUEST_ID_HEADER)) + .map(|value| value.to_str().unwrap_or("").to_owned()) + .unwrap_or_else(|| "".to_owned()); + let cf_ray_str = headers + .get(CF_RAY_HEADER) + .map(|value| value.to_str().unwrap_or("").to_owned()) + .unwrap_or_else(|| "".to_owned()); + format!("request-id: {request_id_str}, cf-ray: {cf_ray_str}") +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::transport::remote_control::auth::RemoteControlConnectionAuth; + use crate::transport::remote_control::protocol::normalize_remote_control_url; + use crate::transport::remote_control::server_api::enroll_remote_control_server; + use codex_state::StateRuntime; + use codex_utils_absolute_path::test_support::PathExt; + use pretty_assertions::assert_eq; + use serde_json::json; + use std::sync::Arc; + use tempfile::TempDir; + use tokio::io::AsyncBufReadExt; + use tokio::io::AsyncWriteExt; + use tokio::io::BufReader; + use tokio::net::TcpListener; + use tokio::net::TcpStream; + use tokio::time::Duration; + use tokio::time::timeout; + + async fn remote_control_state_runtime(codex_home: &TempDir) -> Arc { + StateRuntime::init( + codex_state::SqliteConfig::new_for_testing(codex_home.path().abs()), + "test-provider".to_string(), + ) + .await + .expect("state runtime should initialize") + } + + #[test] + fn preview_remote_control_response_body_redacts_server_token() { + assert_eq!( + serde_json::from_str::(&preview_remote_control_response_body( + br#"{"server_id":"srv_e_test","remote_control_token":"secret","pairing_code":"pairing-code","manual_pairing_code":"ABCD-EFGH"}"# + )) + .expect("redacted response preview should stay valid json"), + json!({ + "server_id": "srv_e_test", + "remote_control_token": "", + "pairing_code": "", + "manual_pairing_code": "", + }) + ); + } + + #[tokio::test] + async fn persisted_remote_control_enrollment_round_trips_by_target_and_account() { + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let first_target = normalize_remote_control_url("https://chatgpt.com/remote/control") + .expect("first target should parse"); + let second_target = + normalize_remote_control_url("https://api.chatgpt-staging.com/other/control") + .expect("second target should parse"); + let first_enrollment = RemoteControlEnrollment { + remote_control_target: first_target.clone(), + account_id: "account-a".to_string(), + environment_id: "env_first".to_string(), + server_id: "srv_e_first".to_string(), + server_name: "first-server".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + let second_enrollment = RemoteControlEnrollment { + remote_control_target: second_target.clone(), + account_id: "account-a".to_string(), + environment_id: "env_second".to_string(), + server_id: "srv_e_second".to_string(), + server_name: "second-server".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &first_target, + "account-a", + Some("desktop-client"), + Some(&first_enrollment), + /*remote_control_enabled*/ None, + ) + .await + .expect("first enrollment should persist"); + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &second_target, + "account-a", + Some("desktop-client"), + Some(&second_enrollment), + /*remote_control_enabled*/ None, + ) + .await + .expect("second enrollment should persist"); + + assert_eq!( + load_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &first_target, + "account-a", + Some("desktop-client"), + ) + .await + .expect("first enrollment should load"), + Some(first_enrollment.clone()) + ); + assert_eq!( + load_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &first_target, + "account-b", + Some("desktop-client"), + ) + .await + .expect("missing account should load"), + None + ); + assert_eq!( + load_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &second_target, + "account-a", + Some("desktop-client"), + ) + .await + .expect("second enrollment should load"), + Some(second_enrollment) + ); + } + + #[tokio::test] + async fn clearing_persisted_remote_control_enrollment_removes_only_matching_entry() { + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let first_target = normalize_remote_control_url("https://chatgpt.com/remote/control") + .expect("first target should parse"); + let second_target = + normalize_remote_control_url("https://api.chatgpt-staging.com/other/control") + .expect("second target should parse"); + let first_enrollment = RemoteControlEnrollment { + remote_control_target: first_target.clone(), + account_id: "account-a".to_string(), + environment_id: "env_first".to_string(), + server_id: "srv_e_first".to_string(), + server_name: "first-server".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + let second_enrollment = RemoteControlEnrollment { + remote_control_target: second_target.clone(), + account_id: "account-a".to_string(), + environment_id: "env_second".to_string(), + server_id: "srv_e_second".to_string(), + server_name: "second-server".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &first_target, + "account-a", + /*app_server_client_name*/ None, + Some(&first_enrollment), + /*remote_control_enabled*/ None, + ) + .await + .expect("first enrollment should persist"); + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &second_target, + "account-a", + /*app_server_client_name*/ None, + Some(&second_enrollment), + /*remote_control_enabled*/ None, + ) + .await + .expect("second enrollment should persist"); + + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &first_target, + "account-a", + /*app_server_client_name*/ None, + /*enrollment*/ None, + /*remote_control_enabled*/ None, + ) + .await + .expect("matching enrollment should clear"); + + assert_eq!( + load_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &first_target, + "account-a", + /*app_server_client_name*/ None, + ) + .await + .expect("cleared enrollment should load"), + None + ); + assert_eq!( + load_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &second_target, + "account-a", + /*app_server_client_name*/ None, + ) + .await + .expect("remaining enrollment should load"), + Some(second_enrollment) + ); + } + + #[tokio::test] + async fn enroll_remote_control_server_parse_failure_includes_response_body() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = format!( + "http://127.0.0.1:{}/backend-api/", + listener + .local_addr() + .expect("listener should have a local addr") + .port() + ); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let enroll_url = remote_control_target.enroll_url.clone(); + let response_body = json!({ + "server_id": "srv_e_test", + "environment_id": "env_test", + }); + let expected_body = response_body.to_string(); + let server_task = tokio::spawn(async move { + let stream = accept_http_request(&listener).await; + respond_with_json(stream, response_body).await; + }); + + let err = enroll_remote_control_server( + &remote_control_target, + &RemoteControlConnectionAuth { + auth_provider: codex_model_provider::unauthenticated_auth_provider(), + account_id: "account_id".to_string(), + }, + "11111111-1111-4111-8111-111111111111", + "test-server", + ) + .await + .expect_err("invalid response should fail to parse"); + + server_task.await.expect("server task should succeed"); + assert_eq!( + err.to_string(), + format!( + "failed to parse remote control server enrollment response from `{enroll_url}`: HTTP 200 OK, request-id: , cf-ray: , body: {expected_body}, decode error: missing field `remote_control_token` at line 1 column {}", + expected_body.len() + ) + ); + } + + async fn accept_http_request(listener: &TcpListener) -> TcpStream { + let (stream, _) = timeout(Duration::from_secs(5), listener.accept()) + .await + .expect("HTTP request should arrive in time") + .expect("listener accept should succeed"); + let mut reader = BufReader::new(stream); + + let mut request_line = String::new(); + reader + .read_line(&mut request_line) + .await + .expect("request line should read"); + loop { + let mut line = String::new(); + reader + .read_line(&mut line) + .await + .expect("header line should read"); + if line == "\r\n" { + break; + } + } + + reader.into_inner() + } + + async fn respond_with_json(mut stream: TcpStream, body: serde_json::Value) { + let body = body.to_string(); + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}", + body.len() + ); + stream + .write_all(response.as_bytes()) + .await + .expect("response should write"); + stream.flush().await.expect("response should flush"); + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/host_device.rs b/codex-rs/app-server-transport/src/transport/remote_control/host_device.rs new file mode 100644 index 0000000000000000000000000000000000000000..02f9d9fd43d989317eb98bad1057f7c59d79baae --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/host_device.rs @@ -0,0 +1,74 @@ +#[cfg(any(target_os = "macos", test))] +use serde::Deserialize; + +pub(super) const REMOTE_CONTROL_HOST_DEVICE_KIND_HEADER: &str = "x-codex-host-device-kind"; +#[cfg(any(target_os = "macos", test))] +const MAC_MINI_HOST_DEVICE_KIND: &str = "mac_mini"; + +#[cfg(any(target_os = "macos", test))] +#[derive(Deserialize)] +struct MacHardwareProfile { + #[serde(rename = "SPHardwareDataType")] + hardware: Vec, +} + +#[cfg(any(target_os = "macos", test))] +#[derive(Deserialize)] +struct MacHardware { + machine_name: String, +} + +#[cfg(any(target_os = "macos", test))] +fn host_device_kind_from_profile(profile: &[u8]) -> serde_json::Result> { + let profile: MacHardwareProfile = serde_json::from_slice(profile)?; + Ok(profile + .hardware + .first() + .is_some_and(|hardware| hardware.machine_name == "Mac mini") + .then_some(MAC_MINI_HOST_DEVICE_KIND)) +} + +#[cfg(target_os = "macos")] +pub(super) async fn host_device_kind() -> Option<&'static str> { + use std::process::Stdio; + use std::time::Duration; + use tokio::process::Command; + use tokio::sync::OnceCell; + + static HOST_DEVICE_KIND: OnceCell> = OnceCell::const_new(); + + HOST_DEVICE_KIND + .get_or_try_init(|| async { + let output = tokio::time::timeout( + Duration::from_secs(2), + Command::new("/usr/sbin/system_profiler") + .args(["-detailLevel", "mini", "SPHardwareDataType", "-json"]) + .stdin(Stdio::null()) + .stderr(Stdio::null()) + .kill_on_drop(true) + .output(), + ) + .await + .map_err(|_| ())? + .map_err(|_| ())?; + + if !output.status.success() { + return Err(()); + } + + host_device_kind_from_profile(&output.stdout).map_err(|_| ()) + }) + .await + .ok() + .copied() + .flatten() +} + +#[cfg(not(target_os = "macos"))] +pub(super) async fn host_device_kind() -> Option<&'static str> { + None +} + +#[cfg(test)] +#[path = "host_device_tests.rs"] +mod tests; diff --git a/codex-rs/app-server-transport/src/transport/remote_control/host_device_tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/host_device_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4c02b2111cfa331e080b0f1352d1a1908fcefe97 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/host_device_tests.rs @@ -0,0 +1,38 @@ +use super::host_device_kind_from_profile; +use pretty_assertions::assert_eq; + +#[test] +fn recognizes_only_the_exact_mac_mini_hardware_name() { + let mac_mini_profile = br#"{"SPHardwareDataType":[{"machine_name":"Mac mini"}]}"#; + let macbook_profile = br#"{"SPHardwareDataType":[{"machine_name":"MacBook Pro"}]}"#; + let misleading_profile = br#"{"SPHardwareDataType":[{"machine_name":"Not a Mac mini"}]}"#; + + assert_eq!( + host_device_kind_from_profile(mac_mini_profile).expect("valid Mac mini hardware profile"), + Some("mac_mini") + ); + assert_eq!( + host_device_kind_from_profile(macbook_profile).expect("valid MacBook hardware profile"), + None + ); + assert_eq!( + host_device_kind_from_profile(misleading_profile) + .expect("valid non-Mac-mini hardware profile"), + None + ); +} + +#[test] +fn ignores_missing_hardware_profiles() { + assert_eq!( + host_device_kind_from_profile(br#"{"SPHardwareDataType":[]}"#) + .expect("valid empty hardware profile"), + None + ); +} + +#[test] +fn rejects_malformed_hardware_profiles_so_they_remain_retryable() { + assert!(host_device_kind_from_profile(br#"{"machine_name":"Mac mini"}"#).is_err()); + assert!(host_device_kind_from_profile(b"not json").is_err()); +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/mod.rs b/codex-rs/app-server-transport/src/transport/remote_control/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..aa672d36caa01edad5c3a2691565c674123d63ba --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/mod.rs @@ -0,0 +1,1002 @@ +mod auth; +mod client_tracker; +mod clients; +mod controller; +mod persistence; +pub use controller::RemoteControlHandle; +pub use controller::start_remote_control; +mod desired_state; +mod enroll; +mod host_device; +mod protocol; +mod segment; +mod server_api; +mod websocket; + +use self::auth::load_remote_control_auth; +use self::auth::recover_remote_control_auth; +use self::desired_state::RemoteControlDesiredState; +use self::enroll::RemoteControlEnrollment; +use self::enroll::load_persisted_remote_control_enrollment; +use self::persistence::RemoteControlPersistence; +use self::server_api::enroll_remote_control_server; +use self::server_api::refresh_remote_control_server; +use crate::transport::remote_control::websocket::RemoteControlChannels; +use crate::transport::remote_control::websocket::RemoteControlStatusPublisher; +use crate::transport::remote_control::websocket::RemoteControlWebsocket; + +pub use self::protocol::ClientId; +use self::protocol::RemoteControlPairingStatusCode; +use self::protocol::ServerEvent; +use self::protocol::StreamId; +use self::protocol::normalize_remote_control_url; +use super::CHANNEL_CAPACITY; +use super::TransportEvent; +use super::next_connection_id; +use codex_app_server_protocol::RemoteControlClientsListParams; +use codex_app_server_protocol::RemoteControlClientsListResponse; +use codex_app_server_protocol::RemoteControlClientsRevokeParams; +use codex_app_server_protocol::RemoteControlClientsRevokeResponse; +use codex_app_server_protocol::RemoteControlConnectionStatus; +use codex_app_server_protocol::RemoteControlPairingStartParams; +use codex_app_server_protocol::RemoteControlPairingStartResponse; +use codex_app_server_protocol::RemoteControlPairingStatusParams; +use codex_app_server_protocol::RemoteControlPairingStatusResponse; +use codex_app_server_protocol::RemoteControlStatusChangedNotification; +use codex_login::AuthManager; +use codex_state::StateRuntime; +use gethostname::gethostname; +use std::error::Error; +use std::fmt; +use std::io; +use std::ops::Deref; +use std::ops::DerefMut; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use tokio::sync::Semaphore; +use tokio::sync::SemaphorePermit; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::sync::watch; +use tokio::task::JoinHandle; +use tokio_util::sync::CancellationToken; +use tracing::info; +use tracing::warn; + +pub struct RemoteControlStartConfig { + pub remote_control_url: String, + pub installation_id: String, + pub policy: RemoteControlPolicy, +} + +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +pub enum RemoteControlPolicy { + #[default] + Allowed, + DisabledByRequirements, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RemoteControlStartupMode { + ResolvePersisted, + DisabledEphemeral, + EnabledEphemeral, +} + +/// Internal marker used by the daemon to disable remote control without requiring a new CLI flag. +pub const REMOTE_CONTROL_DISABLED_ENV_VAR: &str = + "CODEX_INTERNAL_APP_SERVER_REMOTE_CONTROL_DISABLED"; + +/// Reads and removes the daemon's internal disabled-start marker before worker threads start. +pub fn take_remote_control_disabled_env() -> bool { + let disabled = + std::env::var_os(REMOTE_CONTROL_DISABLED_ENV_VAR).is_some_and(|value| value == "1"); + // SAFETY: app-server calls this synchronously at process startup, before spawning threads. + unsafe { std::env::remove_var(REMOTE_CONTROL_DISABLED_ENV_VAR) }; + disabled +} + +pub(super) struct QueuedServerEnvelope { + pub(super) event: ServerEvent, + pub(super) client_id: ClientId, + pub(super) stream_id: StreamId, + pub(super) write_complete_tx: Option>, +} + +#[derive(Clone)] +struct RemoteControlSession { + policy: RemoteControlPolicy, + shutdown_token: CancellationToken, + desired_state_tx: Arc>, + desired_state_rpc_lock: Arc, + persistence: RemoteControlPersistence, + status_tx: Arc>, + state_db: Option>, + remote_control_url: String, + current_enrollment: CurrentRemoteControlEnrollment, + pairing_persistence_key: RemoteControlPairingPersistenceKey, + pairing_persistence_key_required: bool, + auth_manager: auth::RemoteControlAuth, +} + +// Pairing and websocket connect share one selected server so they cannot enroll or replace +// different persisted rows while either path is awaiting backend I/O. +type CurrentRemoteControlEnrollment = Arc; +type RemoteControlPairingPersistenceKey = watch::Sender>; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum RemoteControlEnrollmentSelection { + ReuseOrCreate, + ReplaceExisting, +} + +struct RemoteControlEnrollmentState { + enrollment: StdMutex>, + // Keep an observed server deadline across enrollment replacement and enable transitions. + retry_at: StdMutex>, + lock: Semaphore, +} + +impl RemoteControlEnrollmentState { + fn new(enrollment: Option) -> Self { + Self { + enrollment: StdMutex::new(enrollment), + retry_at: StdMutex::new(None), + lock: Semaphore::new(1), + } + } + + async fn lock_for_request(&self) -> io::Result> { + self.check_retry_after()?; + let lease = self.lock().await; + self.check_retry_after()?; + Ok(lease) + } + + // Check immediately before admitting network work. Requests admitted before + // an overload response publishes its deadline may finish concurrently. + fn check_retry_after(&self) -> io::Result<()> { + let retry_at = *self + .retry_at + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(retry_at) = retry_at + && retry_at > time::OffsetDateTime::now_utc() + { + return Err(server_api::RemoteControlServerRequestError::retry_deferred( + retry_at, + )); + } + Ok(()) + } + + fn record_retry_after(&self, result: io::Result) -> io::Result { + if let Err(err) = &result + && let Some(retry_at) = server_api::remote_control_retry_at(err) + { + let mut current_retry_at = self + .retry_at + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + *current_retry_at = Some(current_retry_at.map_or(retry_at, |at| at.max(retry_at))); + } + result + } + + async fn lock(&self) -> RemoteControlEnrollmentLease<'_> { + let permit = match self.lock.acquire().await { + Ok(permit) => permit, + Err(_) => unreachable!("remote control enrollment lock should stay open"), + }; + RemoteControlEnrollmentLease { + state: self, + enrollment: self.snapshot(), + _permit: permit, + } + } + + fn snapshot(&self) -> Option { + self.enrollment + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + } +} + +struct RemoteControlEnrollmentLease<'a> { + state: &'a RemoteControlEnrollmentState, + enrollment: Option, + _permit: SemaphorePermit<'a>, +} + +impl Deref for RemoteControlEnrollmentLease<'_> { + type Target = Option; + + fn deref(&self) -> &Self::Target { + &self.enrollment + } +} + +impl DerefMut for RemoteControlEnrollmentLease<'_> { + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.enrollment + } +} + +impl Drop for RemoteControlEnrollmentLease<'_> { + fn drop(&mut self) { + *self + .state + .enrollment + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = self.enrollment.take(); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteControlUnavailable; + +impl fmt::Display for RemoteControlUnavailable { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!( + f, + "remote control cannot be enabled because sqlite state db is unavailable" + ) + } +} + +impl Error for RemoteControlUnavailable {} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteControlDisabledByRequirements; + +impl fmt::Display for RemoteControlDisabledByRequirements { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "remote control is disabled by managed requirements") + } +} + +impl Error for RemoteControlDisabledByRequirements {} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RemoteControlEnableError { + Unavailable(RemoteControlUnavailable), + DisabledByRequirements(RemoteControlDisabledByRequirements), + AuthenticationChanged, +} + +impl fmt::Display for RemoteControlEnableError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Unavailable(err) => err.fmt(f), + Self::DisabledByRequirements(err) => err.fmt(f), + Self::AuthenticationChanged => f.write_str("remote control authentication changed"), + } + } +} + +impl Error for RemoteControlEnableError {} + +impl RemoteControlSession { + pub fn ensure_remote_control_allowed(&self) -> Result<(), RemoteControlDisabledByRequirements> { + match self.policy { + RemoteControlPolicy::Allowed => Ok(()), + RemoteControlPolicy::DisabledByRequirements => Err(RemoteControlDisabledByRequirements), + } + } + + fn ensure_remote_control_allowed_io(&self) -> io::Result<()> { + self.ensure_remote_control_allowed() + .map_err(|err| io::Error::new(io::ErrorKind::PermissionDenied, err)) + } + + pub fn enable_ephemeral( + &self, + ) -> Result { + self.enable_with_preference(/*persistence_preference*/ None) + } + + fn enable_with_preference( + &self, + persistence_preference: Option, + ) -> Result { + self.ensure_remote_control_allowed() + .map_err(RemoteControlEnableError::DisabledByRequirements)?; + if self.state_db.is_none() { + warn!("remote control cannot be enabled because sqlite state db is unavailable"); + return Err(RemoteControlEnableError::Unavailable( + RemoteControlUnavailable, + )); + } + + let mut effective_persistence_preference = persistence_preference; + self.auth_manager + .ensure_current() + .map_err(|_| RemoteControlEnableError::AuthenticationChanged)?; + let desired_state_changed = self.desired_state_tx.send_if_modified(|state| { + if effective_persistence_preference.is_none() + && matches!( + *state, + RemoteControlDesiredState::Enabled { + persistence_preference: Some(true) + } + ) + { + effective_persistence_preference = Some(true); + } + let next_state = RemoteControlDesiredState::Enabled { + persistence_preference: effective_persistence_preference, + }; + let changed = *state != next_state; + *state = next_state; + changed + }); + + let status = self.status(); + info!( + desired_state_changed, + ?effective_persistence_preference, + current_status = ?status.status, + environment_id = ?status.environment_id, + installation_id = %status.installation_id, + server_name = %status.server_name, + "remote control enable requested" + ); + if matches!( + status.status, + RemoteControlConnectionStatus::Connected | RemoteControlConnectionStatus::Connecting + ) { + return Ok(status); + } + + Ok(self.publish_status(RemoteControlConnectionStatus::Connecting)) + } + + pub async fn disable( + &self, + app_server_client_name: Option<&str>, + ) -> io::Result { + self.ensure_remote_control_allowed_io()?; + let _transition = self + .desired_state_rpc_lock + .acquire() + .await + .unwrap_or_else(|_| unreachable!()); + self.persist_preference( + app_server_client_name, + /*remote_control_enabled*/ false, + ) + .await?; + Ok(self.transition_disabled()) + } + + pub async fn disable_ephemeral(&self) -> RemoteControlStatusChangedNotification { + let _transition = self + .desired_state_rpc_lock + .acquire() + .await + .unwrap_or_else(|_| unreachable!()); + let _persistence = self.persistence.lock().await; + self.transition_disabled() + } + + fn transition_disabled(&self) -> RemoteControlStatusChangedNotification { + let desired_state_changed = self.desired_state_tx.send_if_modified(|state| { + let changed = *state != RemoteControlDesiredState::Disabled; + *state = RemoteControlDesiredState::Disabled; + changed + }); + let status = self.status(); + info!( + desired_state_changed, + current_status = ?status.status, + environment_id = ?status.environment_id, + installation_id = %status.installation_id, + server_name = %status.server_name, + "remote control disable requested" + ); + self.publish_status(RemoteControlConnectionStatus::Disabled) + } + + async fn persist_preference( + &self, + app_server_client_name: Option<&str>, + remote_control_enabled: bool, + ) -> io::Result<()> { + let state_db = self + .state_db + .as_deref() + .ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, RemoteControlUnavailable))?; + let auth = load_remote_control_auth(&self.auth_manager).await?; + let remote_control_target = normalize_remote_control_url(&self.remote_control_url)?; + let app_server_client_name = self.pairing_persistence_key(app_server_client_name)?; + self.set_preference( + state_db, + &remote_control_target, + &auth.account_id, + app_server_client_name.as_deref(), + remote_control_enabled, + /*fallback_enrollment*/ None, + ) + .await?; + Ok(()) + } + + pub fn status(&self) -> RemoteControlStatusChangedNotification { + self.status_tx.borrow().clone() + } + + pub fn status_receiver(&self) -> watch::Receiver { + self.status_tx.subscribe() + } + + pub async fn start_pairing( + &self, + params: RemoteControlPairingStartParams, + app_server_client_name: Option<&str>, + ) -> io::Result { + self.ensure_remote_control_allowed_io()?; + if !self.desired_state_tx.borrow().is_enabled() { + return Err(Self::pairing_disabled_error()); + } + let mut current_enrollment = self.current_enrollment.lock_for_request().await?; + let mut auth = load_remote_control_auth(&self.auth_manager) + .await + .map_err(|_| pairing_unavailable_error())?; + let status = self.status(); + let installation_id = status.installation_id; + let app_server_client_name = self.pairing_persistence_key(app_server_client_name)?; + let app_server_client_name = app_server_client_name.as_deref(); + let mut enrollment = self + .load_or_enroll_pairing_server( + &mut current_enrollment, + &mut auth, + &installation_id, + &status.server_name, + app_server_client_name, + RemoteControlEnrollmentSelection::ReuseOrCreate, + ) + .await?; + if enrollment.should_refresh_server_token() { + let refresh_result = refresh_pairing_enrollment( + &mut current_enrollment, + &self.auth_manager, + &mut auth, + &installation_id, + &mut enrollment, + ) + .await; + if refresh_result + .as_ref() + .is_err_and(|err| err.kind() == io::ErrorKind::NotFound) + { + enrollment = self + .load_or_enroll_pairing_server( + &mut current_enrollment, + &mut auth, + &installation_id, + &status.server_name, + app_server_client_name, + RemoteControlEnrollmentSelection::ReplaceExisting, + ) + .await?; + } else { + refresh_result?; + } + } + let pairing_request = || protocol::StartRemoteControlPairingRequest { + manual_code: params.manual_code, + }; + self.current_enrollment.check_retry_after()?; + let pairing_response = match enrollment.start_pairing(pairing_request()).await { + Err(err) if err.kind() == io::ErrorKind::PermissionDenied => { + clear_pairing_server_token(&mut current_enrollment, &mut enrollment)?; + refresh_pairing_enrollment( + &mut current_enrollment, + &self.auth_manager, + &mut auth, + &installation_id, + &mut enrollment, + ) + .await?; + self.current_enrollment.check_retry_after()?; + enrollment.start_pairing(pairing_request()).await + } + Err(err) if err.kind() == io::ErrorKind::NotFound => { + enrollment = self + .load_or_enroll_pairing_server( + &mut current_enrollment, + &mut auth, + &installation_id, + &status.server_name, + app_server_client_name, + RemoteControlEnrollmentSelection::ReplaceExisting, + ) + .await?; + self.current_enrollment.check_retry_after()?; + enrollment.start_pairing(pairing_request()).await + } + pairing_response => pairing_response, + }; + if let Err(err) = &pairing_response { + if server_api::remote_control_retry_at(err).is_some() { + return self.current_enrollment.record_retry_after(pairing_response); + } + match err.kind() { + io::ErrorKind::NotFound => { + self.load_or_enroll_pairing_server( + &mut current_enrollment, + &mut auth, + &installation_id, + &status.server_name, + app_server_client_name, + RemoteControlEnrollmentSelection::ReplaceExisting, + ) + .await?; + return Err(pairing_unavailable_error()); + } + io::ErrorKind::PermissionDenied => { + clear_pairing_server_token(&mut current_enrollment, &mut enrollment)?; + return Err(pairing_unavailable_error()); + } + _ => {} + } + } + let current_auth = load_remote_control_auth(&self.auth_manager) + .await + .map_err(|_| pairing_unavailable_error())?; + if current_auth.account_id != auth.account_id { + return Err(pairing_unavailable_error()); + } + if !self.desired_state_tx.borrow().is_enabled() { + return Err(Self::pairing_disabled_error()); + } + pairing_response + } + + async fn load_or_enroll_pairing_server( + &self, + current_enrollment: &mut Option, + auth: &mut auth::RemoteControlConnectionAuth, + installation_id: &str, + server_name: &str, + app_server_client_name: Option<&str>, + selection: RemoteControlEnrollmentSelection, + ) -> io::Result { + let (enrollment, created) = self + .load_or_enroll_server( + current_enrollment, + auth, + installation_id, + server_name, + app_server_client_name, + selection, + ) + .await?; + if !created { + publish_current_enrollment(current_enrollment, &enrollment); + return Ok(enrollment); + } + + let state_db = self + .state_db + .as_deref() + .ok_or_else(pairing_unavailable_error)?; + persistence::save_enrollment( + &self.auth_manager, + &self.persistence, + state_db, + &enrollment, + app_server_client_name, + &self.desired_state_tx, + ) + .await?; + publish_current_enrollment(current_enrollment, &enrollment); + Ok(enrollment) + } + + async fn load_or_enroll_server( + &self, + current_enrollment: &Option, + auth: &mut auth::RemoteControlConnectionAuth, + installation_id: &str, + server_name: &str, + app_server_client_name: Option<&str>, + selection: RemoteControlEnrollmentSelection, + ) -> io::Result<(RemoteControlEnrollment, bool)> { + let remote_control_target = normalize_remote_control_url(&self.remote_control_url)?; + match selection { + RemoteControlEnrollmentSelection::ReuseOrCreate => { + if let Some(enrollment) = current_enrollment + .as_ref() + .filter(|enrollment| enrollment.account_id == auth.account_id) + .cloned() + { + return Ok((enrollment, false)); + } + + let state_db = self + .state_db + .as_deref() + .ok_or_else(pairing_unavailable_error)?; + let _persistence = + persistence::read_lock(&self.auth_manager, &self.persistence).await?; + if let Some(mut enrollment) = load_persisted_remote_control_enrollment( + Some(state_db), + &remote_control_target, + &auth.account_id, + app_server_client_name, + ) + .await? + { + enrollment.server_name = server_name.to_string(); + return Ok((enrollment, false)); + } + } + RemoteControlEnrollmentSelection::ReplaceExisting => {} + } + + self.current_enrollment.check_retry_after()?; + // Reused enrollments must still reach durable persistence during shutdown. + let enrollment = tokio::select! { + biased; + _ = self.shutdown_token.cancelled() => { + return Err(io::Error::new( + io::ErrorKind::Interrupted, + "remote control is shutting down", + )); + } + result = enroll_pairing_server( + &self.current_enrollment, + &self.auth_manager, + auth, + &remote_control_target, + installation_id, + server_name, + ) => self.current_enrollment.record_retry_after(result)?, + }; + Ok((enrollment, true)) + } + + fn pairing_persistence_key( + &self, + app_server_client_name: Option<&str>, + ) -> io::Result> { + if self.pairing_persistence_key_required && self.pairing_persistence_key.borrow().is_none() + { + let app_server_client_name = + app_server_client_name.ok_or_else(pairing_unavailable_error)?; + self.pairing_persistence_key + .send_replace(Some(app_server_client_name.to_string())); + } + Ok(self.pairing_persistence_key.borrow().clone()) + } + + pub async fn pairing_status( + &self, + params: RemoteControlPairingStatusParams, + ) -> io::Result { + self.ensure_remote_control_allowed_io()?; + if !self.desired_state_tx.borrow().is_enabled() { + return Err(Self::pairing_disabled_error()); + } + let mut current_enrollment = self.current_enrollment.lock_for_request().await?; + let mut auth = load_remote_control_auth(&self.auth_manager) + .await + .map_err(|_| pairing_unavailable_error())?; + let app_server_client_name = self.pairing_persistence_key.borrow().clone(); + let app_server_client_name = app_server_client_name.as_deref(); + let mut enrollment = current_enrollment + .as_ref() + .filter(|enrollment| enrollment.account_id == auth.account_id) + .cloned() + .ok_or_else(pairing_unavailable_error)?; + let status = self.status(); + let installation_id = status.installation_id; + let server_name = status.server_name; + if enrollment.should_refresh_server_token() { + let refresh_result = refresh_pairing_enrollment( + &mut current_enrollment, + &self.auth_manager, + &mut auth, + &installation_id, + &mut enrollment, + ) + .await; + if refresh_result + .as_ref() + .is_err_and(|err| err.kind() == io::ErrorKind::NotFound) + { + self.load_or_enroll_pairing_server( + &mut current_enrollment, + &mut auth, + &installation_id, + &server_name, + app_server_client_name, + RemoteControlEnrollmentSelection::ReplaceExisting, + ) + .await?; + return Err(pairing_unavailable_error()); + } + refresh_result?; + } + let status_code = remote_control_pairing_status_code(¶ms)?; + let pairing_status_request = + || protocol::RemoteControlPairingStatusRequest::from(status_code.clone()); + self.current_enrollment.check_retry_after()?; + let pairing_status_response = + match enrollment.pairing_status(pairing_status_request()).await { + Err(err) if err.kind() == io::ErrorKind::PermissionDenied => { + clear_pairing_server_token(&mut current_enrollment, &mut enrollment)?; + refresh_pairing_enrollment( + &mut current_enrollment, + &self.auth_manager, + &mut auth, + &installation_id, + &mut enrollment, + ) + .await?; + self.current_enrollment.check_retry_after()?; + enrollment.pairing_status(pairing_status_request()).await + } + pairing_status_response => pairing_status_response, + }; + if let Err(err) = &pairing_status_response { + if server_api::remote_control_retry_at(err).is_some() { + return self + .current_enrollment + .record_retry_after(pairing_status_response); + } + match err.kind() { + io::ErrorKind::NotFound => { + self.load_or_enroll_pairing_server( + &mut current_enrollment, + &mut auth, + &installation_id, + &server_name, + app_server_client_name, + RemoteControlEnrollmentSelection::ReplaceExisting, + ) + .await?; + return Err(pairing_unavailable_error()); + } + io::ErrorKind::PermissionDenied => { + clear_pairing_server_token(&mut current_enrollment, &mut enrollment)?; + return Err(pairing_unavailable_error()); + } + _ => {} + } + } + if !self.desired_state_tx.borrow().is_enabled() { + return Err(Self::pairing_disabled_error()); + } + let current_auth = load_remote_control_auth(&self.auth_manager) + .await + .map_err(|_| pairing_unavailable_error())?; + if current_auth.account_id != auth.account_id { + return Err(pairing_unavailable_error()); + } + pairing_status_response + } + + pub async fn list_clients( + &self, + params: RemoteControlClientsListParams, + ) -> io::Result { + self.ensure_remote_control_allowed_io()?; + clients::list_remote_control_clients(&self.remote_control_url, &self.auth_manager, params) + .await + } + + pub async fn revoke_client( + &self, + params: RemoteControlClientsRevokeParams, + ) -> io::Result { + self.ensure_remote_control_allowed_io()?; + clients::revoke_remote_control_client(&self.remote_control_url, &self.auth_manager, params) + .await + } + + fn pairing_disabled_error() -> io::Error { + io::Error::new( + io::ErrorKind::InvalidInput, + "remote control pairing requires remote control to be enabled", + ) + } + + fn publish_status( + &self, + connection_status: RemoteControlConnectionStatus, + ) -> RemoteControlStatusChangedNotification { + let mut status_change = None; + self.status_tx.send_if_modified(|status| { + let next_status = + remote_control_status_with_connection_status(status, connection_status); + if *status == next_status { + return false; + } + + status_change = Some((status.clone(), next_status.clone())); + *status = next_status; + true + }); + if let Some((previous_status, next_status)) = status_change { + info!( + previous_status = ?previous_status.status, + next_status = ?next_status.status, + previous_environment_id = ?previous_status.environment_id, + next_environment_id = ?next_status.environment_id, + installation_id = %next_status.installation_id, + server_name = %next_status.server_name, + "remote control handle status changed" + ); + } + self.status() + } +} + +async fn enroll_pairing_server( + current_enrollment: &RemoteControlEnrollmentState, + auth_manager: &auth::RemoteControlAuth, + auth: &mut auth::RemoteControlConnectionAuth, + remote_control_target: &protocol::RemoteControlTarget, + installation_id: &str, + server_name: &str, +) -> io::Result { + match enroll_remote_control_server(remote_control_target, auth, installation_id, server_name) + .await + { + Ok(enrollment) => return Ok(enrollment), + Err(err) if err.kind() == io::ErrorKind::PermissionDenied => { + let mut auth_recovery = auth_manager.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + if !recover_remote_control_auth(&mut auth_recovery, &mut auth_change_rx).await { + return Err(err); + } + *auth = load_remote_control_auth(auth_manager) + .await + .map_err(|_| pairing_unavailable_error())?; + } + Err(err) => return Err(err), + } + current_enrollment.check_retry_after()?; + enroll_remote_control_server(remote_control_target, auth, installation_id, server_name).await +} + +fn remote_control_pairing_status_code( + params: &RemoteControlPairingStatusParams, +) -> io::Result { + match (¶ms.pairing_code, ¶ms.manual_pairing_code) { + (Some(pairing_code), None) => Ok(RemoteControlPairingStatusCode::PairingCode( + pairing_code.clone(), + )), + (None, Some(manual_pairing_code)) => Ok(RemoteControlPairingStatusCode::ManualPairingCode( + manual_pairing_code.clone(), + )), + (Some(_), Some(_)) => Err(io::Error::new( + io::ErrorKind::InvalidInput, + "remote control pairing status accepts either pairingCode or manualPairingCode, not both", + )), + (None, None) => Err(io::Error::new( + io::ErrorKind::InvalidInput, + "remote control pairing status requires pairingCode or manualPairingCode", + )), + } +} + +async fn refresh_pairing_enrollment( + current_enrollment: &mut RemoteControlEnrollmentLease<'_>, + auth_manager: &auth::RemoteControlAuth, + auth: &mut auth::RemoteControlConnectionAuth, + installation_id: &str, + enrollment: &mut RemoteControlEnrollment, +) -> io::Result<()> { + current_enrollment.state.check_retry_after()?; + let mut refresh_result = refresh_remote_control_server(auth, installation_id, enrollment).await; + if refresh_result + .as_ref() + .is_err_and(|err| err.kind() == io::ErrorKind::PermissionDenied) + { + let mut auth_recovery = auth_manager.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + if recover_remote_control_auth(&mut auth_recovery, &mut auth_change_rx).await { + match load_remote_control_auth(auth_manager).await { + Ok(recovered_auth) if recovered_auth.account_id == enrollment.account_id => { + *auth = recovered_auth; + current_enrollment.state.check_retry_after()?; + refresh_result = + refresh_remote_control_server(auth, installation_id, enrollment).await; + } + Ok(_) | Err(_) => { + enrollment.clear_server_token(); + refresh_result = Err(pairing_unavailable_error()); + } + } + } else { + enrollment.clear_server_token(); + } + } + if refresh_result + .as_ref() + .is_err_and(|err| err.kind() == io::ErrorKind::PermissionDenied) + { + enrollment.clear_server_token(); + } + if !replace_current_enrollment(current_enrollment, enrollment) { + Err(pairing_unavailable_error()) + } else { + current_enrollment.state.record_retry_after(refresh_result) + } +} + +fn clear_pairing_server_token( + current_enrollment: &mut Option, + enrollment: &mut RemoteControlEnrollment, +) -> io::Result<()> { + enrollment.clear_server_token(); + if replace_current_enrollment(current_enrollment, enrollment) { + Ok(()) + } else { + Err(pairing_unavailable_error()) + } +} + +fn pairing_unavailable_error() -> io::Error { + io::Error::new( + io::ErrorKind::InvalidInput, + "remote control pairing is unavailable until enrollment completes", + ) +} + +fn remote_control_status_with_connection_status( + status: &RemoteControlStatusChangedNotification, + connection_status: RemoteControlConnectionStatus, +) -> RemoteControlStatusChangedNotification { + RemoteControlStatusChangedNotification { + status: connection_status, + server_name: status.server_name.clone(), + installation_id: status.installation_id.clone(), + environment_id: if connection_status == RemoteControlConnectionStatus::Disabled { + None + } else { + status.environment_id.clone() + }, + } +} + +fn publish_current_enrollment( + current_enrollment: &mut Option, + enrollment: &RemoteControlEnrollment, +) { + *current_enrollment = Some(enrollment.clone()); +} + +fn replace_current_enrollment( + current_enrollment: &mut Option, + enrollment: &RemoteControlEnrollment, +) -> bool { + if !current_enrollment + .as_ref() + .is_some_and(|current| same_remote_control_enrollment(current, enrollment)) + { + return false; + } + *current_enrollment = Some(enrollment.clone()); + true +} + +fn same_remote_control_enrollment( + left: &RemoteControlEnrollment, + right: &RemoteControlEnrollment, +) -> bool { + // A refresh rotates only the bearer. Pairing remains current while the same persisted server + // record is still selected for the current account. + left.account_id == right.account_id + && left.server_id == right.server_id + && left.environment_id == right.environment_id +} + +#[cfg(test)] +mod segment_tests; +#[cfg(test)] +mod tests; diff --git a/codex-rs/app-server-transport/src/transport/remote_control/persistence.rs b/codex-rs/app-server-transport/src/transport/remote_control/persistence.rs new file mode 100644 index 0000000000000000000000000000000000000000..aa0038f4bcdf22d2bd7fb356b9dec7e385096c83 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/persistence.rs @@ -0,0 +1,165 @@ +//! Serializes enrollment storage across login sessions. +//! An admitted SQLite write keeps its permit until completion even if its caller is cancelled. +//! Process shutdown drains admitted writes after stopping the session workers. + +use super::RemoteControlSession; +use super::auth::RemoteControlAuth; +use super::desired_state::RemoteControlDesiredState; +use super::enroll::RemoteControlEnrollment; +use super::enroll::update_persisted_remote_control_enrollment; +use super::protocol::RemoteControlTarget; +use codex_state::StateRuntime; +use std::future::Future; +use std::io; +use std::sync::Arc; +use tokio::sync::Semaphore; +use tokio::sync::SemaphorePermit; +use tokio::sync::watch; +use tokio_util::task::TaskTracker; + +#[cfg(test)] +#[path = "persistence_tests.rs"] +mod tests; + +#[derive(Clone)] +pub(super) struct RemoteControlPersistence { + semaphore: Arc, + pub(super) tasks: TaskTracker, +} + +impl Default for RemoteControlPersistence { + fn default() -> Self { + Self { + semaphore: Arc::new(Semaphore::new(1)), + tasks: TaskTracker::new(), + } + } +} + +impl RemoteControlPersistence { + pub(super) async fn lock(&self) -> SemaphorePermit<'_> { + self.semaphore + .acquire() + .await + .unwrap_or_else(|_| unreachable!()) + } +} + +pub(super) async fn read_lock<'a>( + auth: &RemoteControlAuth, + lock: &'a RemoteControlPersistence, +) -> io::Result> { + let permit = lock.lock().await; + auth.ensure_current()?; + Ok(permit) +} + +async fn commit( + auth: &RemoteControlAuth, + lock: &RemoteControlPersistence, + operation: impl Future> + Send + 'static, +) -> io::Result { + // Register before checking admission. Shutdown either observes this token or rejects us. + let task = lock.tasks.token(); + if lock.tasks.is_closed() { + return Err(io::Error::new( + io::ErrorKind::Interrupted, + "remote control is shutting down", + )); + } + let auth = auth.clone(); + let semaphore = lock.semaphore.clone(); + tokio::spawn(async move { + let _task = task; + let _permit = semaphore.acquire_owned().await.map_err(io::Error::other)?; + auth.ensure_current()?; + operation.await + }) + .await + .map_err(io::Error::other)? +} + +pub(super) async fn save_enrollment( + auth: &RemoteControlAuth, + lock: &RemoteControlPersistence, + state_db: &StateRuntime, + enrollment: &RemoteControlEnrollment, + client_name: Option<&str>, + desired: &watch::Sender, +) -> io::Result<()> { + let state_db = state_db.clone(); + let enrollment = enrollment.clone(); + let client_name = client_name.map(str::to_owned); + let desired = desired.clone(); + commit(auth, lock, async move { + let preference = match *desired.borrow() { + RemoteControlDesiredState::Enabled { + persistence_preference, + } => persistence_preference, + RemoteControlDesiredState::Disabled | RemoteControlDesiredState::Unknown => { + return Err(io::Error::new( + io::ErrorKind::Interrupted, + "remote control disabled during enrollment", + )); + } + }; + update_persisted_remote_control_enrollment( + Some(&state_db), + &enrollment.remote_control_target, + &enrollment.account_id, + client_name.as_deref(), + Some(&enrollment), + preference, + ) + .await + }) + .await +} + +impl RemoteControlSession { + pub(super) async fn set_preference( + &self, + state_db: &StateRuntime, + target: &RemoteControlTarget, + account_id: &str, + client_name: Option<&str>, + enabled: bool, + fallback_enrollment: Option<&RemoteControlEnrollment>, + ) -> io::Result<()> { + let state_db = state_db.clone(); + let target = target.clone(); + let account_id = account_id.to_owned(); + let client_name = client_name.map(str::to_owned); + let enrollment = fallback_enrollment.cloned(); + let desired = self.desired_state_tx.as_ref().clone(); + commit(&self.auth_manager, &self.persistence, async move { + let updated = state_db + .set_remote_control_enabled( + &target.websocket_url, + &account_id, + client_name.as_deref(), + enabled, + ) + .await + .map_err(io::Error::other)?; + if updated == 0 + && let Some(enrollment) = enrollment + { + update_persisted_remote_control_enrollment( + Some(&state_db), + &target, + &account_id, + client_name.as_deref(), + Some(&enrollment), + Some(enabled), + ) + .await?; + } + if !enabled { + desired.send_replace(RemoteControlDesiredState::Disabled); + } + Ok(()) + }) + .await + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/persistence_tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/persistence_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..01386b19f553846b1ac8f7a9abc3e3de4fae6b2b --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/persistence_tests.rs @@ -0,0 +1,51 @@ +//! Checks the persistence boundary when an operation loses its caller. + +use super::*; +use codex_core::test_support::auth_manager_from_auth; +use codex_login::CodexAuth; +use futures::poll; +use pretty_assertions::assert_eq; +use tokio::sync::oneshot; + +#[tokio::test] +async fn cancelled_commit_keeps_its_permit_and_is_drained() -> io::Result<()> { + let auth = RemoteControlAuth::capture(auth_manager_from_auth( + CodexAuth::create_dummy_chatgpt_auth_for_testing(), + )) + .0; + let persistence = RemoteControlPersistence::default(); + let writes = Arc::new(tokio::sync::Mutex::new(Vec::new())); + let (started_tx, started_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel(); + let old_writes = writes.clone(); + let mut old = Box::pin(commit(&auth, &persistence, async move { + started_tx.send(()).expect("test receiver is open"); + release_rx.await.map_err(io::Error::other)?; + old_writes.lock().await.push("old"); + Ok(()) + })); + assert!(poll!(&mut old).is_pending()); + started_rx.await.map_err(io::Error::other)?; + drop(old); + + let new_writes = writes.clone(); + let mut new = Box::pin(commit(&auth, &persistence, async move { + new_writes.lock().await.push("new"); + Ok(()) + })); + assert!(poll!(&mut new).is_pending()); + persistence.tasks.close(); + let mut drained = Box::pin(persistence.tasks.wait()); + assert!(poll!(&mut drained).is_pending()); + release_tx.send(()).expect("commit still owns its receiver"); + tokio::time::timeout(std::time::Duration::from_secs(5), drained).await?; + new.await?; + assert_eq!(*writes.lock().await, vec!["old", "new"]); + let error = commit::<()>(&auth, &persistence, async { + panic!("closed storage must reject new work") + }) + .await + .expect_err("shutdown closes admission"); + assert_eq!(error.kind(), io::ErrorKind::Interrupted); + Ok(()) +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/protocol.rs b/codex-rs/app-server-transport/src/transport/remote_control/protocol.rs new file mode 100644 index 0000000000000000000000000000000000000000..e3c122c1f7c2e90638c083fe9f40638d5f718aca --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/protocol.rs @@ -0,0 +1,401 @@ +use crate::outgoing_message::OutgoingMessage; +use codex_app_server_protocol::JSONRPCMessage; +use serde::Deserialize; +use serde::Serialize; +use std::io; +use std::io::ErrorKind; +use url::Host; +use url::Url; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(super) struct RemoteControlTarget { + pub(super) websocket_url: String, + pub(super) enroll_url: String, + pub(super) refresh_url: String, + pub(super) pair_url: String, + pub(super) pair_status_url: String, +} + +#[derive(Debug, Serialize)] +pub(super) struct EnrollRemoteServerRequest { + pub(super) name: String, + pub(super) os: &'static str, + pub(super) arch: &'static str, + pub(super) app_server_version: &'static str, + pub(super) installation_id: String, +} + +#[derive(Debug, Deserialize)] +pub(super) struct EnrollRemoteServerResponse { + pub(super) server_id: String, + pub(super) environment_id: String, + pub(super) remote_control_token: String, + pub(super) expires_at: String, +} + +#[derive(Debug, Serialize)] +pub(super) struct RefreshRemoteServerRequest { + pub(super) server_id: String, + pub(super) installation_id: String, +} + +#[derive(Debug, Serialize)] +pub(super) struct StartRemoteControlPairingRequest { + pub(super) manual_code: bool, +} + +#[derive(Debug, Deserialize)] +pub(super) struct StartRemoteControlPairingResponse { + pub(super) pairing_code: String, + pub(super) manual_pairing_code: Option, + pub(super) server_id: String, + pub(super) environment_id: String, + pub(super) expires_at: String, +} + +#[derive(Debug, Serialize)] +pub(super) struct RemoteControlPairingStatusRequest { + #[serde(skip_serializing_if = "Option::is_none")] + pub(super) pairing_code: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub(super) manual_pairing_code: Option, +} + +#[derive(Clone)] +pub(super) enum RemoteControlPairingStatusCode { + PairingCode(String), + ManualPairingCode(String), +} + +impl From for RemoteControlPairingStatusRequest { + fn from(code: RemoteControlPairingStatusCode) -> Self { + match code { + RemoteControlPairingStatusCode::PairingCode(pairing_code) => Self { + pairing_code: Some(pairing_code), + manual_pairing_code: None, + }, + RemoteControlPairingStatusCode::ManualPairingCode(manual_pairing_code) => Self { + pairing_code: None, + manual_pairing_code: Some(manual_pairing_code), + }, + } + } +} + +#[derive(Debug, Deserialize)] +pub(super) struct RemoteControlPairingStatusResponse { + pub(super) claimed: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(transparent)] +pub struct ClientId(pub String); + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(transparent)] +pub struct StreamId(pub String); + +impl StreamId { + pub fn new_random() -> Self { + Self(uuid::Uuid::now_v7().to_string()) + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ClientEvent { + ClientMessage { + message: JSONRPCMessage, + }, + ClientMessageChunk { + segment_id: usize, + segment_count: usize, + message_size_bytes: usize, + message_chunk_base64: String, + }, + /// Backend-generated acknowledgement for all server envelopes addressed to + /// `client_id` and `stream_id` whose envelope `seq_id` is less than or equal + /// to this ack's `seq_id`. Chunk acknowledgements carry `segment_id` so the + /// sender can retain only the still-unacked wire chunks on reconnect. + Ack { + #[serde(skip_serializing_if = "Option::is_none")] + segment_id: Option, + }, + Ping, + ClientClosed, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) struct ClientEnvelope { + #[serde(flatten)] + pub(crate) event: ClientEvent, + #[serde(rename = "client_id")] + pub(crate) client_id: ClientId, + #[serde(rename = "stream_id", skip_serializing_if = "Option::is_none")] + pub(crate) stream_id: Option, + /// For `Ack`, this is the backend-generated per-stream cursor over + /// `ServerEnvelope.seq_id`. + #[serde(rename = "seq_id", skip_serializing_if = "Option::is_none")] + pub(crate) seq_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub(crate) cursor: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum PongStatus { + Active, + Unknown, +} + +#[derive(Debug, Clone, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ServerEvent { + ServerMessage { + message: Box, + }, + ServerMessageChunk { + segment_id: usize, + segment_count: usize, + message_size_bytes: usize, + message_chunk_base64: String, + }, + #[allow(dead_code)] + Ack, + Pong { + status: PongStatus, + }, +} + +impl ServerEvent { + pub(crate) fn segment_id(&self) -> Option { + match self { + Self::ServerMessageChunk { segment_id, .. } => Some(*segment_id), + Self::ServerMessage { .. } | Self::Ack | Self::Pong { .. } => None, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "snake_case")] +pub(crate) struct ServerEnvelope { + #[serde(flatten)] + pub(crate) event: ServerEvent, + #[serde(rename = "client_id")] + pub(crate) client_id: ClientId, + #[serde(rename = "stream_id")] + pub(crate) stream_id: StreamId, + #[serde(rename = "seq_id")] + pub(crate) seq_id: u64, +} + +fn is_allowed_remote_control_chatgpt_host(host: &Option>) -> bool { + let Some(Host::Domain(host)) = *host else { + return false; + }; + host == "chatgpt.com" + || host == "chatgpt-staging.com" + || host.ends_with(".chatgpt.com") + || host.ends_with(".chatgpt-staging.com") +} + +fn is_localhost(host: &Option>) -> bool { + match host { + Some(Host::Domain("localhost")) => true, + Some(Host::Ipv4(ip)) => ip.is_loopback(), + Some(Host::Ipv6(ip)) => ip.is_loopback(), + _ => false, + } +} + +pub(super) fn normalize_remote_control_url( + remote_control_url: &str, +) -> io::Result { + let remote_control_url = normalize_remote_control_base_url(remote_control_url)?; + let map_url_parse_error = |err: url::ParseError| -> io::Error { + io::Error::new( + ErrorKind::InvalidInput, + format!("invalid remote control URL `{remote_control_url}`: {err}"), + ) + }; + + let enroll_url = remote_control_url + .join("wham/remote/control/server/enroll") + .map_err(map_url_parse_error)?; + let refresh_url = remote_control_url + .join("wham/remote/control/server/refresh") + .map_err(map_url_parse_error)?; + let pair_url = remote_control_url + .join("wham/remote/control/server/pair") + .map_err(map_url_parse_error)?; + let pair_status_url = remote_control_url + .join("wham/remote/control/server/pair/status") + .map_err(map_url_parse_error)?; + let mut websocket_url = remote_control_url + .join("wham/remote/control/server") + .map_err(map_url_parse_error)?; + websocket_url + .set_scheme(if enroll_url.scheme() == "https" { + "wss" + } else { + "ws" + }) + .map_err(|()| { + io::Error::new( + ErrorKind::InvalidInput, + format!("invalid remote control URL `{remote_control_url}`"), + ) + })?; + + Ok(RemoteControlTarget { + websocket_url: websocket_url.to_string(), + enroll_url: enroll_url.to_string(), + refresh_url: refresh_url.to_string(), + pair_url: pair_url.to_string(), + pair_status_url: pair_status_url.to_string(), + }) +} + +pub(super) fn normalize_remote_control_base_url(remote_control_url: &str) -> io::Result { + let map_url_parse_error = |err: url::ParseError| -> io::Error { + io::Error::new( + ErrorKind::InvalidInput, + format!("invalid remote control URL `{remote_control_url}`: {err}"), + ) + }; + let map_scheme_error = |_: ()| -> io::Error { + io::Error::new( + ErrorKind::InvalidInput, + format!( + "invalid remote control URL `{remote_control_url}`; expected HTTPS URL for chatgpt.com or chatgpt-staging.com, or HTTP/HTTPS URL for localhost" + ), + ) + }; + + let mut remote_control_url = Url::parse(remote_control_url).map_err(map_url_parse_error)?; + if !remote_control_url.path().ends_with('/') { + let normalized_path = format!("{}/", remote_control_url.path()); + remote_control_url.set_path(&normalized_path); + } + + let host = remote_control_url.host(); + match remote_control_url.scheme() { + "https" if is_localhost(&host) || is_allowed_remote_control_chatgpt_host(&host) => {} + "http" if is_localhost(&host) => {} + _ => return Err(map_scheme_error(())), + } + + Ok(remote_control_url) +} + +#[cfg(test)] +mod tests { + use super::*; + use pretty_assertions::assert_eq; + + #[test] + fn normalize_remote_control_url_accepts_chatgpt_https_urls() { + assert_eq!( + normalize_remote_control_url("https://chatgpt.com/backend-api") + .expect("chatgpt.com URL should normalize"), + RemoteControlTarget { + websocket_url: "wss://chatgpt.com/backend-api/wham/remote/control/server" + .to_string(), + enroll_url: "https://chatgpt.com/backend-api/wham/remote/control/server/enroll" + .to_string(), + refresh_url: "https://chatgpt.com/backend-api/wham/remote/control/server/refresh" + .to_string(), + pair_url: "https://chatgpt.com/backend-api/wham/remote/control/server/pair" + .to_string(), + pair_status_url: + "https://chatgpt.com/backend-api/wham/remote/control/server/pair/status" + .to_string(), + } + ); + assert_eq!( + normalize_remote_control_url("https://api.chatgpt-staging.com/backend-api") + .expect("chatgpt-staging.com subdomain URL should normalize"), + RemoteControlTarget { + websocket_url: + "wss://api.chatgpt-staging.com/backend-api/wham/remote/control/server" + .to_string(), + enroll_url: + "https://api.chatgpt-staging.com/backend-api/wham/remote/control/server/enroll" + .to_string(), + refresh_url: + "https://api.chatgpt-staging.com/backend-api/wham/remote/control/server/refresh" + .to_string(), + pair_url: + "https://api.chatgpt-staging.com/backend-api/wham/remote/control/server/pair" + .to_string(), + pair_status_url: + "https://api.chatgpt-staging.com/backend-api/wham/remote/control/server/pair/status" + .to_string(), + } + ); + } + + #[test] + fn normalize_remote_control_url_accepts_localhost_urls() { + assert_eq!( + normalize_remote_control_url("http://localhost:8080/backend-api") + .expect("localhost http URL should normalize"), + RemoteControlTarget { + websocket_url: "ws://localhost:8080/backend-api/wham/remote/control/server" + .to_string(), + enroll_url: "http://localhost:8080/backend-api/wham/remote/control/server/enroll" + .to_string(), + refresh_url: "http://localhost:8080/backend-api/wham/remote/control/server/refresh" + .to_string(), + pair_url: "http://localhost:8080/backend-api/wham/remote/control/server/pair" + .to_string(), + pair_status_url: + "http://localhost:8080/backend-api/wham/remote/control/server/pair/status" + .to_string(), + } + ); + assert_eq!( + normalize_remote_control_url("https://localhost:8443/backend-api") + .expect("localhost https URL should normalize"), + RemoteControlTarget { + websocket_url: "wss://localhost:8443/backend-api/wham/remote/control/server" + .to_string(), + enroll_url: "https://localhost:8443/backend-api/wham/remote/control/server/enroll" + .to_string(), + refresh_url: + "https://localhost:8443/backend-api/wham/remote/control/server/refresh" + .to_string(), + pair_url: "https://localhost:8443/backend-api/wham/remote/control/server/pair" + .to_string(), + pair_status_url: + "https://localhost:8443/backend-api/wham/remote/control/server/pair/status" + .to_string(), + } + ); + } + + #[test] + fn normalize_remote_control_url_rejects_unsupported_urls() { + for remote_control_url in [ + "http://chatgpt.com/backend-api", + "http://example.com/backend-api", + "https://example.com/backend-api", + "https://chat.openai.com/backend-api", + "https://chatgpt.com.evil.com/backend-api", + "https://evilchatgpt.com/backend-api", + "https://foo.localhost/backend-api", + ] { + let err = normalize_remote_control_url(remote_control_url) + .expect_err("unsupported URL should be rejected"); + + assert_eq!(err.kind(), ErrorKind::InvalidInput); + assert_eq!( + err.to_string(), + format!( + "invalid remote control URL `{remote_control_url}`; expected HTTPS URL for chatgpt.com or chatgpt-staging.com, or HTTP/HTTPS URL for localhost" + ) + ); + } + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/segment.rs b/codex-rs/app-server-transport/src/transport/remote_control/segment.rs new file mode 100644 index 0000000000000000000000000000000000000000..f14d62e4f20ca8fd09d351941a61ca2dd783f33d --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/segment.rs @@ -0,0 +1,469 @@ +use super::protocol::ClientEnvelope; +use super::protocol::ClientEvent; +use super::protocol::ClientId; +use super::protocol::ServerEnvelope; +use super::protocol::ServerEvent; +use super::protocol::StreamId; +use crate::outgoing_message::OutgoingMessage; +use crate::transport::response_serialization_error; +use base64::DecodeSliceError; +use base64::Engine; +use codex_app_server_protocol::JSONRPCMessage; +use std::collections::HashMap; +use std::io; +use std::io::ErrorKind; +use std::io::Write; +use tokio::time::Instant; +use tracing::warn; + +pub(super) const REMOTE_CONTROL_SEGMENT_TARGET_BYTES: usize = 100 * 1024; +pub(super) const REMOTE_CONTROL_SEGMENT_MAX_BYTES: usize = 150 * 1024; +pub(super) const REMOTE_CONTROL_REASSEMBLED_MAX_BYTES: usize = 100 * 1024 * 1024; +pub(super) const REMOTE_CONTROL_SEGMENT_COUNT_MAX: usize = 1024; +const REMOTE_CONTROL_SEGMENT_ASSEMBLY_MAX_COUNT: usize = 128; + +#[derive(Debug)] +struct ClientSegmentAssembly { + stream_id: StreamId, + metadata: ClientSegmentMetadata, + raw: Vec, + next_segment_id: usize, + last_chunk_seen_at: Instant, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct ClientSegmentMetadata { + seq_id: u64, + segment_count: usize, + message_size_bytes: usize, +} + +#[derive(Default)] +pub(super) struct ClientSegmentReassembler { + assemblies: HashMap, +} + +pub(super) enum ClientSegmentObservation { + Forward(Box), + Pending, + Dropped, +} + +impl ClientSegmentReassembler { + pub(super) fn observe(&mut self, envelope: ClientEnvelope) -> ClientSegmentObservation { + let ClientEvent::ClientMessageChunk { + segment_id, + segment_count, + message_size_bytes, + message_chunk_base64, + } = &envelope.event + else { + return ClientSegmentObservation::Forward(Box::new(envelope)); + }; + let segment_id = *segment_id; + let segment_count = *segment_count; + let message_size_bytes = *message_size_bytes; + + let Some(metadata) = ClientSegmentMetadata::from_envelope(&envelope) else { + warn!( + client_id = envelope.client_id.0.as_str(), + "dropping segmented remote-control client envelope without seq_id" + ); + return ClientSegmentObservation::Dropped; + }; + let Some(stream_id) = envelope.stream_id.clone() else { + warn!( + client_id = envelope.client_id.0.as_str(), + "dropping segmented remote-control client envelope without stream_id" + ); + return ClientSegmentObservation::Dropped; + }; + if self.should_ignore_chunk(&envelope.client_id, &stream_id, metadata.seq_id, segment_id) { + return ClientSegmentObservation::Dropped; + } + if segment_count == 0 + || segment_count > REMOTE_CONTROL_SEGMENT_COUNT_MAX + || segment_id >= segment_count + || message_size_bytes == 0 + || message_size_bytes > REMOTE_CONTROL_REASSEMBLED_MAX_BYTES + || message_chunk_base64.is_empty() + { + warn!( + client_id = envelope.client_id.0.as_str(), + "dropping invalid segmented remote-control client envelope" + ); + self.remove_assembly(&envelope.client_id, &stream_id); + return ClientSegmentObservation::Dropped; + } + + let now = Instant::now(); + match self.assemblies.get(&envelope.client_id) { + Some(assembly) if assembly.stream_id != stream_id => { + warn!( + client_id = envelope.client_id.0.as_str(), + "resetting segmented remote-control client envelope after stream change" + ); + self.assemblies.insert( + envelope.client_id.clone(), + ClientSegmentAssembly { + stream_id: stream_id.clone(), + metadata: metadata.clone(), + raw: Vec::new(), + next_segment_id: 0, + last_chunk_seen_at: now, + }, + ); + } + Some(_) => {} + None => { + self.evict_assemblies_if_full(); + self.assemblies.insert( + envelope.client_id.clone(), + ClientSegmentAssembly { + stream_id: stream_id.clone(), + metadata: metadata.clone(), + raw: Vec::new(), + next_segment_id: 0, + last_chunk_seen_at: now, + }, + ); + } + } + let result = { + let Some(assembly) = self.assemblies.get_mut(&envelope.client_id) else { + warn!( + client_id = envelope.client_id.0.as_str(), + "dropping segmented remote-control client envelope without assembly" + ); + return ClientSegmentObservation::Dropped; + }; + if metadata.seq_id < assembly.metadata.seq_id { + AssemblyUpdate::Ignore + } else if assembly.metadata != metadata { + warn!( + client_id = envelope.client_id.0.as_str(), + "resetting segmented remote-control client envelope after metadata mismatch" + ); + AssemblyUpdate::Drop + } else if segment_id < assembly.next_segment_id { + AssemblyUpdate::Pending + } else if segment_id != assembly.next_segment_id { + warn!( + client_id = envelope.client_id.0.as_str(), + "dropping out-of-order segmented remote-control client envelope" + ); + AssemblyUpdate::Drop + } else { + assembly.last_chunk_seen_at = now; + let chunk_start = assembly.raw.len(); + let decoded_chunk_len = base64::decoded_len_estimate(message_chunk_base64.len()); + let chunk_end = usize::min( + message_size_bytes, + chunk_start.saturating_add(decoded_chunk_len), + ); + assembly.raw.resize(chunk_end, 0); + match base64::engine::general_purpose::STANDARD.decode_slice( + message_chunk_base64.as_bytes(), + &mut assembly.raw[chunk_start..], + ) { + Ok(decoded_chunk_len) => { + assembly.raw.truncate(chunk_start + decoded_chunk_len); + assembly.next_segment_id += 1; + if assembly.next_segment_id < segment_count { + AssemblyUpdate::Pending + } else if assembly.raw.len() != message_size_bytes { + warn!( + client_id = envelope.client_id.0.as_str(), + "dropping reassembled remote-control client envelope with mismatched size" + ); + AssemblyUpdate::Drop + } else { + match serde_json::from_slice::(&assembly.raw) { + Ok(message) => AssemblyUpdate::Complete(message), + Err(err) => { + warn!( + client_id = envelope.client_id.0.as_str(), + "dropping invalid reassembled remote-control client envelope: {err}" + ); + AssemblyUpdate::Drop + } + } + } + } + Err(DecodeSliceError::OutputSliceTooSmall) => { + warn!( + client_id = envelope.client_id.0.as_str(), + "dropping segmented remote-control client envelope after size overflow" + ); + AssemblyUpdate::Drop + } + Err(err) => { + warn!( + client_id = envelope.client_id.0.as_str(), + "dropping segmented remote-control client envelope with invalid base64: {err}" + ); + AssemblyUpdate::Drop + } + } + } + }; + + match result { + AssemblyUpdate::Pending => ClientSegmentObservation::Pending, + AssemblyUpdate::Ignore => ClientSegmentObservation::Dropped, + AssemblyUpdate::Drop => { + self.remove_assembly(&envelope.client_id, &stream_id); + ClientSegmentObservation::Dropped + } + AssemblyUpdate::Complete(message) => { + self.remove_assembly(&envelope.client_id, &stream_id); + ClientSegmentObservation::Forward(Box::new(ClientEnvelope { + event: ClientEvent::ClientMessage { message }, + ..envelope + })) + } + } + } + + pub(super) fn invalidate_stream(&mut self, client_id: &ClientId, stream_id: &StreamId) { + self.remove_assembly(client_id, stream_id); + } + + pub(super) fn invalidate_client(&mut self, client_id: &ClientId) { + self.assemblies.remove(client_id); + } + + pub(super) fn should_ignore_chunk( + &self, + client_id: &ClientId, + stream_id: &StreamId, + seq_id: u64, + segment_id: usize, + ) -> bool { + self.assemblies.get(client_id).is_some_and(|assembly| { + assembly.stream_id == *stream_id + && (seq_id < assembly.metadata.seq_id + || (seq_id == assembly.metadata.seq_id + && segment_id < assembly.next_segment_id)) + }) + } + + fn remove_assembly(&mut self, client_id: &ClientId, stream_id: &StreamId) { + if self + .assemblies + .get(client_id) + .is_some_and(|assembly| &assembly.stream_id == stream_id) + { + self.assemblies.remove(client_id); + } + } + + fn evict_assemblies_if_full(&mut self) { + while self.assemblies.len() >= REMOTE_CONTROL_SEGMENT_ASSEMBLY_MAX_COUNT { + let Some(client_id) = self + .assemblies + .iter() + .min_by_key(|(_, assembly)| assembly.last_chunk_seen_at) + .map(|(client_id, _)| client_id.clone()) + else { + return; + }; + self.assemblies.remove(&client_id); + } + } +} + +enum AssemblyUpdate { + Pending, + Ignore, + Drop, + Complete(JSONRPCMessage), +} + +impl ClientSegmentMetadata { + fn from_envelope(envelope: &ClientEnvelope) -> Option { + let ClientEvent::ClientMessageChunk { + segment_count, + message_size_bytes, + .. + } = &envelope.event + else { + return None; + }; + Some(Self { + seq_id: envelope.seq_id?, + segment_count: *segment_count, + message_size_bytes: *message_size_bytes, + }) + } +} + +pub(super) fn split_server_envelope_for_transport( + envelope: ServerEnvelope, +) -> io::Result> { + if !matches!(envelope.event, ServerEvent::ServerMessage { .. }) { + return Ok(vec![envelope]); + } + + let envelope_size_bytes = match serialized_len(&envelope) { + Ok(envelope_size_bytes) => envelope_size_bytes, + Err(err) => { + let ServerEvent::ServerMessage { message } = envelope.event else { + unreachable!("server message variant checked above"); + }; + let OutgoingMessage::Response(response) = *message else { + return Err(err); + }; + return Ok(vec![ServerEnvelope { + event: ServerEvent::ServerMessage { + message: Box::new(response_serialization_error(response.id, err)), + }, + client_id: envelope.client_id, + stream_id: envelope.stream_id, + seq_id: envelope.seq_id, + }]); + } + }; + if envelope_size_bytes <= REMOTE_CONTROL_SEGMENT_MAX_BYTES { + return Ok(vec![envelope]); + } + + let ServerEvent::ServerMessage { message } = envelope.event.clone() else { + unreachable!("server message variant checked above"); + }; + let raw = serde_json::to_vec(message.as_ref()).map_err(io::Error::other)?; + let message_size_bytes = raw.len(); + if message_size_bytes > REMOTE_CONTROL_REASSEMBLED_MAX_BYTES { + warn!("dropping remote-control server envelope that exceeds reassembled size limit"); + return Ok(Vec::new()); + } + + let minimal_segment_count = + usize::min(message_size_bytes.max(1), REMOTE_CONTROL_SEGMENT_COUNT_MAX); + let minimal_chunk = &raw[..usize::min(raw.len(), 1)]; + if serialized_chunk_len( + &envelope, + /*segment_id*/ 0, + minimal_segment_count, + message_size_bytes, + minimal_chunk, + )? > REMOTE_CONTROL_SEGMENT_MAX_BYTES + { + warn!("dropping remote-control server envelope that cannot fit within segment size limit"); + return Ok(Vec::new()); + } + + let mut segment_count = usize::max( + 2, + message_size_bytes.div_ceil(REMOTE_CONTROL_SEGMENT_TARGET_BYTES), + ); + loop { + let chunk_size = usize::max(1, message_size_bytes.div_ceil(segment_count)); + segment_count = message_size_bytes.div_ceil(chunk_size); + let segments_fit = raw + .chunks(chunk_size) + .enumerate() + .all(|(segment_id, chunk)| { + serialized_chunk_len( + &envelope, + segment_id, + segment_count, + message_size_bytes, + chunk, + ) + .is_ok_and(|size| size <= REMOTE_CONTROL_SEGMENT_MAX_BYTES) + }); + if segments_fit { + return raw + .chunks(chunk_size) + .enumerate() + .map(|(segment_id, chunk)| { + build_chunk_envelope( + &envelope, + segment_id, + segment_count, + message_size_bytes, + chunk, + ) + }) + .collect(); + } + if chunk_size == 1 { + warn!( + "dropping remote-control server envelope that cannot fit within segment size limit" + ); + return Ok(Vec::new()); + } + let next_segment_count = segment_count + 1; + let next_chunk_size = usize::max(1, message_size_bytes.div_ceil(next_segment_count)); + segment_count = if next_chunk_size == chunk_size { + message_size_bytes + } else { + next_segment_count + }; + } +} + +fn serialized_chunk_len( + envelope: &ServerEnvelope, + segment_id: usize, + segment_count: usize, + message_size_bytes: usize, + chunk: &[u8], +) -> io::Result { + serialized_len(&build_chunk_envelope( + envelope, + segment_id, + segment_count, + message_size_bytes, + chunk, + )?) +} + +#[derive(Default)] +struct CountingWriter { + len: usize, +} + +impl Write for CountingWriter { + fn write(&mut self, buf: &[u8]) -> io::Result { + self.len += buf.len(); + Ok(buf.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +fn serialized_len(value: &impl serde::Serialize) -> io::Result { + let mut writer = CountingWriter::default(); + serde_json::to_writer(&mut writer, value).map_err(io::Error::other)?; + Ok(writer.len) +} + +fn build_chunk_envelope( + envelope: &ServerEnvelope, + segment_id: usize, + segment_count: usize, + message_size_bytes: usize, + chunk: &[u8], +) -> io::Result { + if segment_count > REMOTE_CONTROL_SEGMENT_COUNT_MAX { + return Err(io::Error::new( + ErrorKind::InvalidData, + "remote-control segment count exceeds maximum", + )); + } + Ok(ServerEnvelope { + event: ServerEvent::ServerMessageChunk { + segment_id, + segment_count, + message_size_bytes, + message_chunk_base64: base64::engine::general_purpose::STANDARD.encode(chunk), + }, + client_id: envelope.client_id.clone(), + stream_id: envelope.stream_id.clone(), + seq_id: envelope.seq_id, + }) +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/segment_tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/segment_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..3b2406b6f12bfb63249d921bee06c8b140957b7c --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/segment_tests.rs @@ -0,0 +1,450 @@ +use super::protocol::ClientEnvelope; +use super::protocol::ClientEvent; +use super::protocol::ClientId; +use super::protocol::ServerEnvelope; +use super::protocol::ServerEvent; +use super::protocol::StreamId; +use super::segment::ClientSegmentObservation; +use super::segment::ClientSegmentReassembler; +use super::segment::REMOTE_CONTROL_SEGMENT_MAX_BYTES; +use super::segment::split_server_envelope_for_transport; +use crate::outgoing_message::OutgoingMessage; +#[cfg(unix)] +use crate::outgoing_message::OutgoingResponse; +use base64::Engine; +#[cfg(unix)] +use codex_app_server_protocol::ClientResponsePayload; +use codex_app_server_protocol::ConfigWarningNotification; +#[cfg(unix)] +use codex_app_server_protocol::InitializeResponse; +use codex_app_server_protocol::JSONRPCMessage; +use codex_app_server_protocol::JSONRPCNotification; +#[cfg(unix)] +use codex_app_server_protocol::RequestId; +use codex_app_server_protocol::ServerNotification; +use codex_app_server_protocol::ServerNotificationEnvelope; +#[cfg(unix)] +use codex_utils_absolute_path::AbsolutePathBuf; +use pretty_assertions::assert_eq; +#[cfg(unix)] +use serde_json::json; + +#[test] +fn reassembles_client_message_chunks() { + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let client_id = ClientId("client-1".to_string()); + let stream_id = Some(StreamId("stream-1".to_string())); + let mut reassembler = ClientSegmentReassembler::default(); + + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + stream_id.clone(), + /*seq_id*/ 7, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + )), + ClientSegmentObservation::Pending + )); + let reassembled = match reassembler.observe(chunk_envelope( + client_id.clone(), + stream_id, + /*seq_id*/ 7, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + )) { + ClientSegmentObservation::Forward(reassembled) => *reassembled, + ClientSegmentObservation::Pending | ClientSegmentObservation::Dropped => { + panic!("message should reassemble") + } + }; + assert_eq!(reassembled.client_id, client_id); + assert_eq!( + reassembled.stream_id, + Some(StreamId("stream-1".to_string())) + ); + assert_eq!(reassembled.seq_id, Some(7)); + assert_eq!(reassembled.cursor, None); + match reassembled.event { + ClientEvent::ClientMessage { + message: reassembled_message, + } => assert_eq!(reassembled_message, message), + other => panic!("expected client message, got {other:?}"), + } +} + +#[test] +fn splits_large_server_messages_into_wire_chunks() { + let envelope = ServerEnvelope { + event: ServerEvent::ServerMessage { + message: Box::new(OutgoingMessage::AppServerNotification( + ServerNotificationEnvelope { + notification: ServerNotification::ConfigWarning(ConfigWarningNotification { + summary: "x".repeat(REMOTE_CONTROL_SEGMENT_MAX_BYTES), + details: None, + path: None, + range: None, + }), + emitted_at_ms: Some(1_234), + }, + )), + }, + client_id: ClientId("client-1".to_string()), + stream_id: StreamId("stream-1".to_string()), + seq_id: 9, + }; + + let segments = split_server_envelope_for_transport(envelope).expect("split should succeed"); + + assert!(segments.len() > 1); + assert!( + segments + .iter() + .all(|segment| matches!(segment.event, ServerEvent::ServerMessageChunk { .. })) + ); + assert!(segments.iter().all(|segment| segment.seq_id == 9)); + assert!(segments.iter().all(|segment| { + serde_json::to_vec(segment) + .expect("segment should serialize") + .len() + <= REMOTE_CONTROL_SEGMENT_MAX_BYTES + })); +} + +#[cfg(unix)] +#[test] +fn invalid_response_becomes_remote_control_jsonrpc_error() { + use std::ffi::OsString; + use std::os::unix::ffi::OsStringExt; + use std::path::PathBuf; + + let codex_home = AbsolutePathBuf::from_absolute_path(PathBuf::from(OsString::from_vec(vec![ + b'/', b'b', b'a', b'd', 0xff, + ]))) + .expect("non-UTF-8 Unix paths are valid absolute paths"); + let envelope = ServerEnvelope { + event: ServerEvent::ServerMessage { + message: Box::new(OutgoingMessage::Response(OutgoingResponse { + id: RequestId::Integer(7), + result: Box::new(ClientResponsePayload::Initialize(InitializeResponse { + user_agent: "codex-test-agent".to_string(), + codex_home, + platform_family: "unix".to_string(), + platform_os: "linux".to_string(), + })), + })), + }, + client_id: ClientId("client-1".to_string()), + stream_id: StreamId("stream-1".to_string()), + seq_id: 9, + }; + + let envelopes = split_server_envelope_for_transport(envelope) + .expect("invalid response should become a remote-control JSON-RPC error"); + assert_eq!( + serde_json::to_value(envelopes).expect("error envelope should serialize"), + json!([{ + "type": "server_message", + "client_id": "client-1", + "stream_id": "stream-1", + "seq_id": 9, + "message": { + "id": 7, + "error": { + "code": -32603, + "message": "failed to serialize response: path contains invalid UTF-8 characters", + } + } + }]) + ); +} + +#[test] +fn invalidates_incomplete_stream_assemblies() { + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let client_id = ClientId("client-1".to_string()); + let stream_id = StreamId("stream-1".to_string()); + let mut reassembler = ClientSegmentReassembler::default(); + + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + Some(stream_id.clone()), + /*seq_id*/ 7, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + )), + ClientSegmentObservation::Pending + )); + reassembler.invalidate_stream(&client_id, &stream_id); + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id, + Some(stream_id), + /*seq_id*/ 7, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + )), + ClientSegmentObservation::Dropped + )); +} + +#[test] +fn resets_incomplete_client_assembly_when_stream_changes() { + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let client_id = ClientId("client-1".to_string()); + let first_stream_id = StreamId("stream-1".to_string()); + let second_stream_id = StreamId("stream-2".to_string()); + let mut reassembler = ClientSegmentReassembler::default(); + + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + Some(first_stream_id.clone()), + /*seq_id*/ 7, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + )), + ClientSegmentObservation::Pending + )); + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + Some(second_stream_id.clone()), + /*seq_id*/ 8, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + )), + ClientSegmentObservation::Pending + )); + let reassembled = match reassembler.observe(chunk_envelope( + client_id.clone(), + Some(second_stream_id), + /*seq_id*/ 8, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + )) { + ClientSegmentObservation::Forward(reassembled) => *reassembled, + ClientSegmentObservation::Pending | ClientSegmentObservation::Dropped => { + panic!("replacement stream should reassemble") + } + }; + assert_eq!( + reassembled.stream_id, + Some(StreamId("stream-2".to_string())) + ); + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id, + Some(first_stream_id), + /*seq_id*/ 7, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + )), + ClientSegmentObservation::Dropped + )); +} + +#[test] +fn ignores_stale_chunks_without_dropping_newer_assembly() { + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let client_id = ClientId("client-1".to_string()); + let stream_id = Some(StreamId("stream-1".to_string())); + let mut reassembler = ClientSegmentReassembler::default(); + + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + stream_id.clone(), + /*seq_id*/ 8, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + )), + ClientSegmentObservation::Pending + )); + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + stream_id.clone(), + /*seq_id*/ 7, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + )), + ClientSegmentObservation::Dropped + )); + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id, + stream_id, + /*seq_id*/ 8, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + )), + ClientSegmentObservation::Forward(_) + )); +} + +#[test] +fn ignores_invalid_stale_chunks_without_dropping_newer_assembly() { + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let client_id = ClientId("client-1".to_string()); + let stream_id = Some(StreamId("stream-1".to_string())); + let mut reassembler = ClientSegmentReassembler::default(); + + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + stream_id.clone(), + /*seq_id*/ 8, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + )), + ClientSegmentObservation::Pending + )); + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + stream_id.clone(), + /*seq_id*/ 7, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + b"", + )), + ClientSegmentObservation::Dropped + )); + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id, + stream_id, + /*seq_id*/ 8, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + )), + ClientSegmentObservation::Forward(_) + )); +} + +#[test] +fn ignores_invalid_duplicate_chunks_without_dropping_current_assembly() { + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let client_id = ClientId("client-1".to_string()); + let stream_id = Some(StreamId("stream-1".to_string())); + let mut reassembler = ClientSegmentReassembler::default(); + + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + stream_id.clone(), + /*seq_id*/ 8, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + )), + ClientSegmentObservation::Pending + )); + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id.clone(), + stream_id.clone(), + /*seq_id*/ 8, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + b"", + )), + ClientSegmentObservation::Dropped + )); + assert!(matches!( + reassembler.observe(chunk_envelope( + client_id, + stream_id, + /*seq_id*/ 8, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + )), + ClientSegmentObservation::Forward(_) + )); +} + +fn chunk_envelope( + client_id: ClientId, + stream_id: Option, + seq_id: u64, + segment_id: usize, + segment_count: usize, + message_size_bytes: usize, + chunk: &[u8], +) -> ClientEnvelope { + ClientEnvelope { + event: ClientEvent::ClientMessageChunk { + segment_id, + segment_count, + message_size_bytes, + message_chunk_base64: base64::engine::general_purpose::STANDARD.encode(chunk), + }, + client_id, + stream_id, + seq_id: Some(seq_id), + cursor: None, + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/server_api.rs b/codex-rs/app-server-transport/src/transport/remote_control/server_api.rs new file mode 100644 index 0000000000000000000000000000000000000000..ab7d7fd8a08c0f50aa82e022c9661b7dddc01d9d --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/server_api.rs @@ -0,0 +1,382 @@ +use super::auth::RemoteControlConnectionAuth; +use super::enroll::RemoteControlEnrollment; +use super::enroll::RemoteControlServerTokenRefreshRequirement; +use super::enroll::format_headers; +use super::enroll::preview_remote_control_response_body; +use super::protocol::EnrollRemoteServerRequest; +use super::protocol::EnrollRemoteServerResponse; +use super::protocol::RefreshRemoteServerRequest; +use super::protocol::RemoteControlTarget; +use axum::http::HeaderMap; +use axum::http::StatusCode; +use codex_login::default_client::create_client_without_request_logging; +use rand::Rng; +use serde::Serialize; +use serde::de::DeserializeOwned; +use std::fmt; +use std::io; +use std::io::ErrorKind; +use std::time::Duration; +use time::OffsetDateTime; +use time::format_description::well_known::Rfc3339; +use tracing::warn; + +const REMOTE_CONTROL_RETRY_AFTER_JITTER_MAX_MILLIS: u64 = 30_000; + +const REMOTE_CONTROL_ENROLL_TIMEOUT: Duration = Duration::from_secs(30); +const REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MIN_SECS: u64 = 24; +const REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MAX_SECS: u64 = 36; + +pub(super) const REMOTE_CONTROL_INSTALLATION_ID_HEADER: &str = "x-codex-installation-id"; + +#[derive(Debug)] +pub(super) struct RemoteControlServerRequestError { + message: String, + status: Option, + retry_at: Option, +} + +impl RemoteControlServerRequestError { + pub(super) fn retry_deferred(retry_at: OffsetDateTime) -> io::Error { + io::Error::new( + ErrorKind::WouldBlock, + Self { + message: format!("remote control retry deferred until {retry_at}"), + status: None, + retry_at: Some(retry_at), + }, + ) + } + + pub(super) fn io_error( + message: String, + status: Option, + retry_at: Option, + timed_out: bool, + ) -> io::Error { + let kind = match status { + Some(StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN) => ErrorKind::PermissionDenied, + Some(StatusCode::NOT_FOUND) => ErrorKind::NotFound, + Some(status) if timed_out && !status.is_client_error() => ErrorKind::TimedOut, + None if timed_out => ErrorKind::TimedOut, + Some(_) | None => ErrorKind::Other, + }; + io::Error::new( + kind, + Self { + message, + status, + retry_at, + }, + ) + } + + fn is_transient(&self, kind: ErrorKind) -> bool { + kind == ErrorKind::TimedOut + || self.status.is_none() + || self.status.is_some_and(|status| { + status == StatusCode::TOO_MANY_REQUESTS || status.is_server_error() + }) + } +} + +impl fmt::Display for RemoteControlServerRequestError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str(&self.message) + } +} + +impl std::error::Error for RemoteControlServerRequestError {} + +pub(super) async fn enroll_remote_control_server( + remote_control_target: &RemoteControlTarget, + auth: &RemoteControlConnectionAuth, + installation_id: &str, + server_name: &str, +) -> io::Result { + let enroll_url = &remote_control_target.enroll_url; + let request = EnrollRemoteServerRequest { + name: server_name.to_string(), + os: std::env::consts::OS, + arch: std::env::consts::ARCH, + app_server_version: env!("CARGO_PKG_VERSION"), + installation_id: installation_id.to_string(), + }; + let enrollment_response = send_remote_control_server_request::<_, EnrollRemoteServerResponse>( + enroll_url, + auth, + installation_id, + &request, + "enroll", + "server enrollment", + REMOTE_CONTROL_ENROLL_TIMEOUT, + ) + .await?; + let mut enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: auth.account_id.clone(), + environment_id: enrollment_response.environment_id, + server_id: enrollment_response.server_id, + server_name: server_name.to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + update_remote_control_server_token( + &mut enrollment, + enroll_url, + enrollment_response.remote_control_token, + enrollment_response.expires_at, + )?; + Ok(enrollment) +} + +pub(super) async fn refresh_remote_control_server( + auth: &RemoteControlConnectionAuth, + installation_id: &str, + enrollment: &mut RemoteControlEnrollment, +) -> io::Result<()> { + let now = OffsetDateTime::now_utc(); + let refresh_requirement = enrollment.server_token_refresh_requirement_at(now); + if refresh_requirement == RemoteControlServerTokenRefreshRequirement::NotNeeded { + return Ok(()); + } + if refresh_requirement == RemoteControlServerTokenRefreshRequirement::Required + && let Some(next_refresh_at) = enrollment.next_refresh_at + && next_refresh_at > now + { + return Err(RemoteControlServerRequestError::retry_deferred( + next_refresh_at, + )); + } + let refresh_url = enrollment.remote_control_target.refresh_url.clone(); + let request = RefreshRemoteServerRequest { + server_id: enrollment.server_id.clone(), + installation_id: installation_id.to_string(), + }; + let refreshed = match send_remote_control_server_request::<_, EnrollRemoteServerResponse>( + &refresh_url, + auth, + installation_id, + &request, + "refresh", + "server refresh", + REMOTE_CONTROL_ENROLL_TIMEOUT, + ) + .await + { + Ok(refreshed) => refreshed, + Err(err) => { + let Some(refresh_error) = remote_control_server_request_error(&err) else { + return Err(err); + }; + if !refresh_error.is_transient(err.kind()) { + return Err(err); + } + let now = OffsetDateTime::now_utc(); + let refresh_is_required = enrollment.server_token_refresh_requirement_at(now) + == RemoteControlServerTokenRefreshRequirement::Required; + let (refresh_delay, next_refresh_at) = refresh_deferral(refresh_error.retry_at, now); + enrollment.next_refresh_at = Some(next_refresh_at); + // An explicit server deadline takes precedence over valid-token fallback. + // Keep the token, but let callers defer new requests until this deadline. + if refresh_is_required + || refresh_error + .retry_at + .is_some_and(|retry_at| retry_at > now) + { + warn!( + refresh_url, + server_id = %enrollment.server_id, + environment_id = %enrollment.environment_id, + error = %err, + ?refresh_delay, + %next_refresh_at, + "remote control server token refresh failed; deferring new requests" + ); + return Err(err); + } + warn!( + refresh_url, + server_id = %enrollment.server_id, + environment_id = %enrollment.environment_id, + error = %err, + ?refresh_delay, + %next_refresh_at, + "proactive remote control server token refresh failed; continuing with valid token" + ); + return Ok(()); + } + }; + if refreshed.server_id != enrollment.server_id + || refreshed.environment_id != enrollment.environment_id + { + return Err(io::Error::other(format!( + "remote control server refresh returned mismatched enrollment: expected server_id={}, environment_id={}; got server_id={}, environment_id={}", + enrollment.server_id, + enrollment.environment_id, + refreshed.server_id, + refreshed.environment_id + ))); + } + + update_remote_control_server_token( + enrollment, + &refresh_url, + refreshed.remote_control_token, + refreshed.expires_at, + ) +} + +async fn send_remote_control_server_request( + url: &str, + auth: &RemoteControlConnectionAuth, + installation_id: &str, + request: &Request, + action: &str, + response_kind: &str, + timeout: Duration, +) -> io::Result +where + Request: Serialize, + Response: DeserializeOwned, +{ + let client = create_client_without_request_logging(); + let auth_headers = auth.request_headers()?; + let response = client + .post(url) + .timeout(timeout) + .headers(auth_headers) + .header(REMOTE_CONTROL_INSTALLATION_ID_HEADER, installation_id) + .json(request) + .send() + .await + .map_err(|err| { + let timed_out = err.is_timeout(); + RemoteControlServerRequestError::io_error( + format!("failed to {action} remote control server at `{url}`: {err}"), + /*status*/ None, + /*retry_at*/ None, + timed_out, + ) + })?; + let headers = response.headers().clone(); + let status = response.status(); + let retry_at = retry_after_with_jitter(&headers, OffsetDateTime::now_utc()); + let body = response.bytes().await.map_err(|err| { + let timed_out = err.is_timeout(); + RemoteControlServerRequestError::io_error( + format!("failed to read remote control {response_kind} response from `{url}`: {err}"), + Some(status), + retry_at, + timed_out, + ) + })?; + let body_preview = preview_remote_control_response_body(&body); + if !status.is_success() { + let headers_str = format_headers(&headers); + return Err(RemoteControlServerRequestError::io_error( + format!( + "remote control {response_kind} failed at `{url}`: HTTP {status}, {headers_str}, body: {body_preview}" + ), + Some(status), + retry_at, + /*timed_out*/ false, + )); + } + + serde_json::from_slice::(&body).map_err(|err| { + let headers_str = format_headers(&headers); + io::Error::other(format!( + "failed to parse remote control {response_kind} response from `{url}`: HTTP {status}, {headers_str}, body: {body_preview}, decode error: {err}" + )) + }) +} + +fn update_remote_control_server_token( + enrollment: &mut RemoteControlEnrollment, + url: &str, + token: String, + expires_at: String, +) -> io::Result<()> { + let expires_at = OffsetDateTime::parse(&expires_at, &Rfc3339).map_err(|err| { + io::Error::other(format!( + "failed to parse remote control server token expiry from `{url}`: {err}" + )) + })?; + enrollment.remote_control_token = Some(token); + enrollment.expires_at = Some(expires_at); + enrollment.next_refresh_at = None; + Ok(()) +} + +fn remote_control_server_request_error( + err: &io::Error, +) -> Option<&RemoteControlServerRequestError> { + err.get_ref()?.downcast_ref() +} + +pub(super) fn remote_control_retry_at(err: &io::Error) -> Option { + let request_error = remote_control_server_request_error(err)?; + if !request_error.is_transient(err.kind()) { + return None; + } + request_error.retry_at +} + +pub(super) fn remote_control_retry_delay(err: &io::Error) -> Option { + Duration::try_from(remote_control_retry_at(err)? - OffsetDateTime::now_utc()).ok() +} + +pub(super) fn retry_after_with_jitter( + headers: &HeaderMap, + received_at: OffsetDateTime, +) -> Option { + let retry_at = parse_retry_after(headers, received_at)?; + // Keep the server's deadline as a lower bound, then spread clients across + // the next 30 seconds. Sample once per response so retries share a deadline. + let jitter = time::Duration::milliseconds( + rand::rng().random_range(0..=REMOTE_CONTROL_RETRY_AFTER_JITTER_MAX_MILLIS) as i64, + ); + Some(retry_at.checked_add(jitter).unwrap_or(retry_at)) +} + +fn parse_retry_after(headers: &HeaderMap, received_at: OffsetDateTime) -> Option { + let retry_after = headers + .get(axum::http::header::RETRY_AFTER)? + .to_str() + .ok()?; + let retry_at = if let Ok(seconds) = retry_after.parse::() { + let seconds = i64::try_from(seconds).ok()?; + received_at.checked_add(time::Duration::seconds(seconds))? + } else { + OffsetDateTime::from(httpdate::parse_http_date(retry_after).ok()?) + }; + (retry_at >= received_at).then_some(retry_at) +} + +fn refresh_deferral( + retry_at: Option, + now: OffsetDateTime, +) -> (Duration, OffsetDateTime) { + if let Some(retry_at) = retry_at + && let Ok(delay) = Duration::try_from(retry_at - now) + && !delay.is_zero() + { + return (delay, retry_at); + } + let delay = remote_control_server_token_refresh_backoff(); + let next_refresh_at = now + time::Duration::seconds(delay.as_secs() as i64); + (delay, next_refresh_at) +} + +fn remote_control_server_token_refresh_backoff() -> Duration { + Duration::from_secs(rand::rng().random_range( + REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MIN_SECS + ..=REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MAX_SECS, + )) +} + +#[cfg(test)] +#[path = "server_api_tests.rs"] +mod tests; diff --git a/codex-rs/app-server-transport/src/transport/remote_control/server_api_tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/server_api_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..0c636d3c0e3a044a1589646b91bb43eafc660816 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/server_api_tests.rs @@ -0,0 +1,321 @@ +use super::*; +use crate::transport::remote_control::protocol::normalize_remote_control_url; +use pretty_assertions::assert_eq; +use serde_json::json; +use std::time::SystemTime; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::sync::oneshot; + +const TEST_REQUEST_TIMEOUT: Duration = Duration::from_millis(100); + +fn auth() -> RemoteControlConnectionAuth { + RemoteControlConnectionAuth { + auth_provider: codex_model_provider::unauthenticated_auth_provider(), + account_id: "account-a".to_string(), + } +} + +fn assert_transient_timeout(err: &io::Error, expected_status: Option) { + let request_error = remote_control_server_request_error(err) + .expect("request error should preserve refresh metadata"); + assert_eq!( + ( + err.kind(), + request_error.status, + request_error.is_transient(err.kind()), + ), + (ErrorKind::TimedOut, expected_status, true) + ); +} + +async fn timed_out_request(partial_response: Option<&'static [u8]>) -> io::Error { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let url = format!( + "http://{}/backend-api/wham/remote/control/server/refresh", + listener + .local_addr() + .expect("listener should have a local address") + ); + let (request_done_tx, request_done_rx) = oneshot::channel(); + let server_task = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("request should connect"); + if let Some(partial_response) = partial_response { + stream + .write_all(partial_response) + .await + .expect("partial response should write"); + } + request_done_rx + .await + .expect("test should report request completion"); + }); + + let err = send_remote_control_server_request::<_, serde_json::Value>( + &url, + &auth(), + "installation-id", + &json!({"server_id": "server-id"}), + "refresh", + "server refresh", + TEST_REQUEST_TIMEOUT, + ) + .await + .expect_err("incomplete response should time out"); + request_done_tx + .send(()) + .expect("server should wait for request completion"); + server_task.await.expect("server task should finish"); + err +} + +fn enrollment(now: OffsetDateTime) -> RemoteControlEnrollment { + RemoteControlEnrollment { + remote_control_target: normalize_remote_control_url("http://localhost/backend-api/") + .expect("target should normalize"), + account_id: "account-a".to_string(), + environment_id: "env_first".to_string(), + server_id: "srv_e_first".to_string(), + server_name: "first-server".to_string(), + remote_control_token: Some("token".to_string()), + expires_at: Some(now + time::Duration::seconds(300)), + next_refresh_at: None, + } +} + +#[test] +fn remote_control_enrollment_classifies_server_token_refresh_requirement() { + let now = + OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse"); + let enrollment = enrollment(now); + let cases = [ + ( + enrollment.clone(), + RemoteControlServerTokenRefreshRequirement::Proactive, + ), + ( + RemoteControlEnrollment { + expires_at: Some(now + time::Duration::seconds(301)), + ..enrollment.clone() + }, + RemoteControlServerTokenRefreshRequirement::NotNeeded, + ), + ( + RemoteControlEnrollment { + next_refresh_at: Some(now + time::Duration::seconds(30)), + ..enrollment.clone() + }, + RemoteControlServerTokenRefreshRequirement::NotNeeded, + ), + ( + RemoteControlEnrollment { + next_refresh_at: Some(now), + ..enrollment.clone() + }, + RemoteControlServerTokenRefreshRequirement::Proactive, + ), + ( + RemoteControlEnrollment { + remote_control_token: None, + ..enrollment.clone() + }, + RemoteControlServerTokenRefreshRequirement::Required, + ), + ( + RemoteControlEnrollment { + expires_at: None, + ..enrollment.clone() + }, + RemoteControlServerTokenRefreshRequirement::Required, + ), + ( + RemoteControlEnrollment { + expires_at: Some(now), + next_refresh_at: Some(now + time::Duration::hours(1)), + ..enrollment + }, + RemoteControlServerTokenRefreshRequirement::Required, + ), + ]; + + for (enrollment, expected) in cases { + assert_eq!( + enrollment.server_token_refresh_requirement_at(now), + expected + ); + } +} + +#[test] +fn remote_control_server_request_error_classifies_status_before_timeout() { + let cases = [ + (None, true, ErrorKind::TimedOut, true), + (Some(StatusCode::OK), true, ErrorKind::TimedOut, true), + ( + Some(StatusCode::TOO_MANY_REQUESTS), + false, + ErrorKind::Other, + true, + ), + (Some(StatusCode::BAD_GATEWAY), false, ErrorKind::Other, true), + ( + Some(StatusCode::UNAUTHORIZED), + true, + ErrorKind::PermissionDenied, + false, + ), + ( + Some(StatusCode::FORBIDDEN), + true, + ErrorKind::PermissionDenied, + false, + ), + ( + Some(StatusCode::NOT_FOUND), + true, + ErrorKind::NotFound, + false, + ), + (Some(StatusCode::BAD_REQUEST), true, ErrorKind::Other, false), + (None, false, ErrorKind::Other, true), + ]; + + for (status, timed_out, expected_kind, expected_transient) in cases { + let err = RemoteControlServerRequestError::io_error( + String::new(), + status, + /*retry_at*/ None, + timed_out, + ); + let request_error = remote_control_server_request_error(&err) + .expect("request error should preserve refresh metadata"); + assert_eq!( + (err.kind(), request_error.is_transient(err.kind())), + (expected_kind, expected_transient) + ); + } +} + +#[tokio::test] +async fn request_timeout_before_response_headers_is_transient() { + let err = timed_out_request(/*partial_response*/ None).await; + assert_transient_timeout(&err, /*expected_status*/ None); +} + +#[tokio::test] +async fn response_body_timeout_is_transient() { + let err = timed_out_request(Some(b"HTTP/1.1 200 OK\r\nContent-Length: 20\r\n\r\n{")).await; + assert_transient_timeout(&err, Some(StatusCode::OK)); +} + +#[test] +fn retry_after_supports_delta_seconds_and_http_dates() { + let now = + OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse"); + let mut headers = HeaderMap::new(); + headers.insert( + axum::http::header::RETRY_AFTER, + axum::http::HeaderValue::from_static("120"), + ); + assert_eq!( + parse_retry_after(&headers, now), + Some(now + time::Duration::seconds(120)) + ); + + let retry_at = now + time::Duration::seconds(90); + let retry_at_system = SystemTime::UNIX_EPOCH + Duration::from_secs(1_700_000_090); + headers.insert( + axum::http::header::RETRY_AFTER, + httpdate::fmt_http_date(retry_at_system) + .parse() + .expect("HTTP date should be a valid header value"), + ); + assert_eq!(parse_retry_after(&headers, now), Some(retry_at)); +} + +#[test] +fn invalid_or_expired_retry_after_uses_bounded_fallback() { + let now = + OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse"); + let mut headers = HeaderMap::new(); + headers.insert( + axum::http::header::RETRY_AFTER, + axum::http::HeaderValue::from_static("invalid"), + ); + assert_eq!(parse_retry_after(&headers, now), None); + + headers.insert( + axum::http::header::RETRY_AFTER, + httpdate::fmt_http_date(SystemTime::UNIX_EPOCH + Duration::from_secs(1_699_999_999)) + .parse() + .expect("HTTP date should be a valid header value"), + ); + assert_eq!(parse_retry_after(&headers, now), None); + + let expired_while_reading_body = Some(now + time::Duration::seconds(1)); + for retry_at in [None, expired_while_reading_body] { + let deferred_at = now + time::Duration::seconds(2); + let (delay, next_refresh_at) = refresh_deferral(retry_at, deferred_at); + assert!( + (Duration::from_secs(REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MIN_SECS) + ..=Duration::from_secs(REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MAX_SECS,)) + .contains(&delay) + ); + assert_eq!( + next_refresh_at, + deferred_at + time::Duration::seconds(delay.as_secs() as i64) + ); + } +} + +#[test] +fn http_date_retry_after_preserves_absolute_deadline() { + let received_at = + OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse"); + let retry_at = received_at + time::Duration::seconds(120); + let body_read_at = received_at + time::Duration::seconds(30); + + assert_eq!( + refresh_deferral(Some(retry_at), body_read_at), + (Duration::from_secs(90), retry_at) + ); +} + +#[test] +fn retry_after_jitter_never_shortens_the_server_deadline() { + let now = + OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse"); + for value in ["120", "Tue, 14 Nov 2023 22:15:20 GMT"] { + let mut headers = HeaderMap::new(); + headers.insert( + axum::http::header::RETRY_AFTER, + value.parse().expect("retry hint should be a valid header"), + ); + let deadline = retry_after_with_jitter(&headers, now) + .expect("a valid retry hint should produce a deadline"); + assert!( + (now + time::Duration::seconds(120)..=now + time::Duration::seconds(150)) + .contains(&deadline) + ); + // Reading the response body must not start a new relative wait or resample jitter. + let (delay, deferred_until) = + refresh_deferral(Some(deadline), now + time::Duration::seconds(30)); + assert_eq!(deferred_until, deadline); + assert!((Duration::from_secs(90)..=Duration::from_secs(120)).contains(&delay)); + } +} + +#[test] +fn zero_retry_after_preserves_bounded_jitter() { + let now = OffsetDateTime::UNIX_EPOCH; + let mut headers = HeaderMap::new(); + headers.insert( + axum::http::header::RETRY_AFTER, + axum::http::HeaderValue::from_static("0"), + ); + let deadline = retry_after_with_jitter(&headers, now) + .expect("a zero-second retry hint should produce a deadline"); + assert!((now..=now + time::Duration::seconds(30)).contains(&deadline)); +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..2337b2a658671f6bf2eebbba41d221a6bd748fca --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/tests.rs @@ -0,0 +1,3142 @@ +use super::auth::REMOTE_CONTROL_ACCOUNT_ID_HEADER; +use super::enroll::RemoteControlEnrollment; +use super::enroll::load_persisted_remote_control_enrollment; +use super::enroll::update_persisted_remote_control_enrollment; +use super::protocol::ClientEnvelope; +use super::protocol::ClientEvent; +use super::protocol::ClientId; +use super::protocol::StreamId; +use super::protocol::normalize_remote_control_url; +use super::server_api::REMOTE_CONTROL_INSTALLATION_ID_HEADER; +use super::websocket::REMOTE_CONTROL_PROTOCOL_VERSION; +use super::websocket::RemoteControlWebsocket; +use super::websocket::RemoteControlWebsocketConfig; +use super::*; +use crate::outgoing_message::OutgoingMessage; +use crate::outgoing_message::QueuedOutgoingMessage; +use crate::transport::CHANNEL_CAPACITY; +use crate::transport::ConnectionOrigin; +use crate::transport::TransportEvent; +use base64::Engine; +use codex_app_server_protocol::ConfigWarningNotification; +use codex_app_server_protocol::JSONRPCMessage; +use codex_app_server_protocol::RemoteControlConnectionStatus; +use codex_app_server_protocol::RemoteControlPairingStartParams; +use codex_app_server_protocol::RemoteControlPairingStatusParams; +use codex_app_server_protocol::RemoteControlStatusChangedNotification; +use codex_app_server_protocol::ServerNotification; +use codex_app_server_protocol::ServerNotificationEnvelope; +use codex_config::types::AuthCredentialsStoreMode; +use codex_core::test_support::auth_manager_from_auth; +use codex_core::test_support::auth_manager_from_auth_with_home; +use codex_login::AuthDotJson; +use codex_login::AuthKeyringBackendKind; +use codex_login::AuthManager; +use codex_login::CodexAuth; +use codex_login::save_auth; +use codex_login::token_data::TokenData; +use codex_login::token_data::parse_chatgpt_jwt_claims; +use codex_protocol::auth::AuthMode; +use codex_state::RemoteControlEnrollmentRecord; +use codex_state::StateRuntime; +use codex_utils_absolute_path::test_support::PathExt; +use futures::SinkExt; +use futures::StreamExt; +use gethostname::gethostname; +use pretty_assertions::assert_eq; +use serde_json::json; +use std::collections::BTreeMap; +use std::sync::Arc; +use tempfile::TempDir; +use time::OffsetDateTime; +use tokio::io::AsyncBufReadExt; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::io::BufReader; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::sync::watch; +use tokio::time::Duration; +use tokio::time::timeout; +use tokio_tungstenite::WebSocketStream; +use tokio_tungstenite::accept_async; +use tokio_tungstenite::accept_hdr_async; +use tokio_tungstenite::tungstenite; +use tokio_util::sync::CancellationToken; + +mod clients_tests; +mod pairing_tests; +#[path = "tests/retry_tests.rs"] +mod retry_tests; + +const TEST_INSTALLATION_ID: &str = "11111111-1111-4111-8111-111111111111"; +const TEST_REMOTE_CONTROL_URL: &str = "http://127.0.0.1:1/backend-api/wham/remote/control"; +const TEST_REMOTE_CONTROL_SERVER_TOKEN: &str = "Remote Control Token"; +const TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN: &str = "Refreshed Remote Control Token"; +const TEST_REMOTE_CONTROL_SERVER_TOKEN_EXPIRES_AT: &str = "2999-01-01T00:00:00Z"; + +fn remote_control_auth_manager() -> Arc { + auth_manager_from_auth(CodexAuth::create_dummy_chatgpt_auth_for_testing()) +} + +fn remote_control_auth_manager_with_home(codex_home: &TempDir) -> Arc { + auth_manager_from_auth_with_home( + CodexAuth::create_dummy_chatgpt_auth_for_testing(), + codex_home.path().to_path_buf(), + ) +} + +fn remote_control_auth_dot_json(account_id: Option<&str>) -> AuthDotJson { + #[derive(serde::Serialize)] + struct Header { + alg: &'static str, + typ: &'static str, + } + + let header = Header { + alg: "none", + typ: "JWT", + }; + let payload = serde_json::json!({ + "email": "user@example.com", + "https://api.openai.com/auth": { + "chatgpt_user_id": "user-12345", + "user_id": "user-12345", + "chatgpt_account_id": "account_id" + } + }); + let b64 = |bytes: &[u8]| base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes); + let header_b64 = b64(&serde_json::to_vec(&header).expect("header should serialize")); + let payload_b64 = b64(&serde_json::to_vec(&payload).expect("payload should serialize")); + let fake_jwt = format!("{header_b64}.{payload_b64}.sig"); + + AuthDotJson { + auth_mode: Some(AuthMode::Chatgpt), + openai_api_key: None, + tokens: Some(TokenData { + id_token: parse_chatgpt_jwt_claims(&fake_jwt).expect("fake jwt should parse"), + access_token: "Access Token".to_string(), + refresh_token: "refresh-token".to_string(), + account_id: account_id.map(str::to_string), + }), + last_refresh: Some(chrono::Utc::now()), + agent_identity: None, + personal_access_token: None, + bedrock_api_key: None, + bedrock_access_keys: None, + } +} + +async fn remote_control_state_runtime(codex_home: &TempDir) -> Arc { + StateRuntime::init( + codex_state::SqliteConfig::new_for_testing(codex_home.path().abs()), + "test-provider".to_string(), + ) + .await + .expect("state runtime should initialize") +} + +#[tokio::test] +async fn committed_disable_prevents_later_enrollment_from_restoring_preference() { + let home = TempDir::new().expect("temp dir"); + let state_db = remote_control_state_runtime(&home).await; + let session = remote_control_handle_with_current_enrollment( + TEST_REMOTE_CONTROL_URL, + remote_control_auth_manager(), + ); + let enrollment = session.current_enrollment.snapshot().expect("enrollment"); + session + .desired_state_tx + .send_replace(RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + }); + session + .set_preference( + &state_db, + &enrollment.remote_control_target, + &enrollment.account_id, + /*client_name*/ None, + /*enabled*/ false, + Some(&enrollment), + ) + .await + .expect("disable commits"); + // This is the window before the disable RPC resumes and publishes its status. + let error = persistence::save_enrollment( + &session.auth_manager, + &session.persistence, + &state_db, + &enrollment, + /*client_name*/ None, + &session.desired_state_tx, + ) + .await + .expect_err("enrollment cannot re-enable a committed disable"); + assert_eq!(error.kind(), std::io::ErrorKind::Interrupted); + let saved = state_db + .get_remote_control_enrollment( + &enrollment.remote_control_target.websocket_url, + &enrollment.account_id, + /*app_server_client_name*/ None, + ) + .await + .expect("read preference") + .expect("saved enrollment"); + assert_eq!(saved.remote_control_enabled, Some(false)); +} + +#[tokio::test] +async fn plain_start_resolves_persisted_remote_control_preference() { + let cases = [ + ("enabled", Some(Some(true))), + ("disabled", Some(Some(false))), + ("unset", Some(None)), + ("missing", None), + ]; + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let remote_control_target = normalize_remote_control_url(TEST_REMOTE_CONTROL_URL) + .expect("remote control target should normalize"); + for (name, stored_preference) in cases { + let Some(remote_control_enabled) = stored_preference else { + continue; + }; + state_db + .upsert_remote_control_enrollment(&RemoteControlEnrollmentRecord { + websocket_url: remote_control_target.websocket_url.clone(), + account_id: "account_id".to_string(), + app_server_client_name: Some(name.to_string()), + server_id: format!("server-{name}"), + environment_id: format!("environment-{name}"), + server_name: format!("server-name-{name}"), + remote_control_enabled, + }) + .await + .expect("enrollment should persist"); + } + let (transport_event_tx, _transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let (status_tx, _status_rx) = watch::channel(RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Disabled, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: None, + }); + let (desired_state_tx, _desired_state_rx) = watch::channel(RemoteControlDesiredState::Unknown); + let desired_state_tx = Arc::new(desired_state_tx); + let mut websocket = RemoteControlWebsocket::new( + RemoteControlWebsocketConfig { + remote_control_url: TEST_REMOTE_CONTROL_URL.to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + remote_control_target: None, + server_name: test_server_name(), + }, + Some(state_db), + auth::RemoteControlAuth::capture(remote_control_auth_manager()).0, + RemoteControlChannels { + transport_event_tx, + status_publisher: RemoteControlStatusPublisher::new(status_tx), + current_enrollment: Arc::new(RemoteControlEnrollmentState::new( + /*enrollment*/ None, + )), + pairing_persistence_key: watch::channel(None).0, + persistence: RemoteControlPersistence::default(), + }, + CancellationToken::new(), + desired_state_tx.clone(), + ); + + for (name, stored_preference) in cases { + desired_state_tx.send_replace(RemoteControlDesiredState::Unknown); + assert!(websocket.resolve_unknown_desired_state(Some(name)).await); + let expected = if stored_preference == Some(Some(true)) { + RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + } + } else { + RemoteControlDesiredState::Disabled + }; + assert_eq!(*desired_state_tx.borrow(), expected, "case {name}"); + } +} + +#[tokio::test] +async fn explicit_disabled_start_ignores_persisted_enable() { + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let remote_control_target = normalize_remote_control_url(TEST_REMOTE_CONTROL_URL) + .expect("remote control target should normalize"); + let enrollment = RemoteControlEnrollmentRecord { + websocket_url: remote_control_target.websocket_url, + account_id: "account_id".to_string(), + app_server_client_name: None, + server_id: "server-id".to_string(), + environment_id: "environment-id".to_string(), + server_name: "server-name".to_string(), + remote_control_enabled: Some(true), + }; + state_db + .upsert_remote_control_enrollment(&enrollment) + .await + .expect("enrollment should persist"); + let (transport_event_tx, _transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url: TEST_REMOTE_CONTROL_URL.to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::DisabledEphemeral, + ) + .await + .expect("remote control should start disabled"); + + assert_eq!( + remote_handle.status().status, + RemoteControlConnectionStatus::Disabled + ); + assert_eq!( + state_db + .get_remote_control_enrollment( + &enrollment.websocket_url, + &enrollment.account_id, + /*app_server_client_name*/ None, + ) + .await + .expect("enrollment should load"), + Some(enrollment) + ); + + shutdown_token.cancel(); + remote_task.await.expect("remote control task should join"); +} + +#[tokio::test] +async fn managed_disable_overrides_startup_and_persisted_enablement() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let remote_control_target = normalize_remote_control_url(&remote_control_url) + .expect("remote control target should normalize"); + let enrollment = RemoteControlEnrollmentRecord { + websocket_url: remote_control_target.websocket_url, + account_id: "account_id".to_string(), + app_server_client_name: None, + server_id: "server-id".to_string(), + environment_id: "environment-id".to_string(), + server_name: "server-name".to_string(), + remote_control_enabled: Some(true), + }; + state_db + .upsert_remote_control_enrollment(&enrollment) + .await + .expect("enrollment should persist"); + let (transport_event_tx, _transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::DisabledByRequirements, + }, + Some(state_db.clone()), + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start disabled"); + + assert_eq!( + remote_handle.status().status, + RemoteControlConnectionStatus::Disabled + ); + assert_eq!( + remote_handle.ensure_remote_control_allowed(), + Err(RemoteControlDisabledByRequirements) + ); + assert!( + !remote_handle + .resolve_persisted_preference(/*app_server_client_name*/ None) + .await + .expect("managed disable should resolve without loading persistence") + ); + assert_eq!( + remote_handle + .enable_ephemeral() + .expect_err("managed requirements should reject ephemeral enable"), + RemoteControlEnableError::DisabledByRequirements(RemoteControlDisabledByRequirements) + ); + let enable_error = remote_handle + .enable(/*app_server_client_name*/ None) + .await + .expect_err("managed requirements should reject durable enable"); + assert_eq!(enable_error.kind(), std::io::ErrorKind::PermissionDenied); + assert_eq!( + enable_error.to_string(), + "remote control is disabled by managed requirements" + ); + let disable_error = remote_handle + .disable(/*app_server_client_name*/ None) + .await + .expect_err("managed requirements should reject durable disable"); + assert_eq!(disable_error.kind(), std::io::ErrorKind::PermissionDenied); + assert_eq!( + disable_error.to_string(), + "remote control is disabled by managed requirements" + ); + assert_eq!( + state_db + .get_remote_control_enrollment( + &enrollment.websocket_url, + &enrollment.account_id, + /*app_server_client_name*/ None, + ) + .await + .expect("enrollment should load"), + Some(enrollment) + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("managed requirements should prevent backend contact"); + + shutdown_token.cancel(); + remote_task.await.expect("remote control task should join"); +} + +fn remote_control_url_for_listener(listener: &TcpListener) -> String { + let addr = listener + .local_addr() + .expect("listener should have a local addr"); + format!("http://{addr}/backend-api/") +} + +fn test_server_name() -> String { + gethostname().to_string_lossy().trim().to_string() +} + +pub(super) fn remote_control_handle_with_current_enrollment( + remote_control_url: &str, + auth_manager: Arc, +) -> RemoteControlSession { + let (desired_state_tx, _desired_state_rx) = + watch::channel(RemoteControlDesiredState::Enabled { + persistence_preference: None, + }); + let (status_tx, _status_rx) = watch::channel(RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connecting, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_test".to_string()), + }); + let remote_control_target = normalize_remote_control_url(remote_control_url) + .expect("remote control target should normalize"); + let current_enrollment = Arc::new(RemoteControlEnrollmentState::new(Some( + RemoteControlEnrollment { + remote_control_target, + account_id: "account_id".to_string(), + environment_id: "env_test".to_string(), + server_id: "srv_e_test".to_string(), + server_name: test_server_name(), + remote_control_token: Some(TEST_REMOTE_CONTROL_SERVER_TOKEN.to_string()), + expires_at: Some( + OffsetDateTime::from_unix_timestamp(33_336_362_096) + .expect("future timestamp should parse"), + ), + next_refresh_at: None, + }, + ))); + RemoteControlSession { + policy: RemoteControlPolicy::Allowed, + shutdown_token: CancellationToken::new(), + desired_state_tx: Arc::new(desired_state_tx), + desired_state_rpc_lock: Arc::new(Semaphore::new(1)), + persistence: RemoteControlPersistence::default(), + status_tx: Arc::new(status_tx), + state_db: None, + remote_control_url: remote_control_url.to_string(), + current_enrollment, + pairing_persistence_key: watch::channel(None).0, + pairing_persistence_key_required: false, + auth_manager: auth::RemoteControlAuth::capture(auth_manager).0, + } +} + +#[tokio::test] +async fn durable_enable_reuses_in_memory_enrollment_after_shutdown() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + remote_handle.state_db = Some(state_db.clone()); + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Disabled); + let enrollment = remote_handle + .current_enrollment + .snapshot() + .expect("in-memory enrollment should exist"); + let expected_record = RemoteControlEnrollmentRecord { + websocket_url: enrollment.remote_control_target.websocket_url.clone(), + account_id: enrollment.account_id.clone(), + app_server_client_name: None, + server_id: enrollment.server_id.clone(), + environment_id: enrollment.environment_id.clone(), + server_name: enrollment.server_name.clone(), + remote_control_enabled: Some(true), + }; + assert_eq!( + state_db + .get_remote_control_enrollment( + &expected_record.websocket_url, + &expected_record.account_id, + /*app_server_client_name*/ None, + ) + .await + .expect("enrollment should load"), + None + ); + remote_handle.shutdown_token.cancel(); + + let status = timeout( + Duration::from_secs(5), + remote_handle.enable(/*app_server_client_name*/ None), + ) + .await + .expect("cached enable should complete without network I/O") + .expect("shutdown should not cancel durable enable using in-memory enrollment"); + + assert_eq!( + state_db + .get_remote_control_enrollment( + &expected_record.websocket_url, + &expected_record.account_id, + /*app_server_client_name*/ None, + ) + .await + .expect("enabled enrollment should load"), + Some(expected_record) + ); + assert_eq!( + *remote_handle.desired_state_tx.borrow(), + RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + } + ); + assert_eq!( + status.environment_id.as_deref(), + Some(enrollment.environment_id.as_str()) + ); + assert_eq!( + remote_handle.current_enrollment.snapshot(), + Some(enrollment) + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("in-memory enrollment should prevent backend contact"); +} + +#[tokio::test] +async fn durable_enable_reuses_persisted_enrollment_after_shutdown() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = normalize_remote_control_url(&remote_control_url) + .expect("remote control target should normalize"); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let persisted_enrollment = RemoteControlEnrollmentRecord { + websocket_url: remote_control_target.websocket_url, + account_id: "account_id".to_string(), + app_server_client_name: None, + server_id: "persisted-server-id".to_string(), + environment_id: "persisted-environment-id".to_string(), + server_name: format!("{}-persisted", test_server_name()), + remote_control_enabled: Some(false), + }; + state_db + .upsert_remote_control_enrollment(&persisted_enrollment) + .await + .expect("disabled enrollment should persist"); + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + remote_handle.state_db = Some(state_db.clone()); + *remote_handle.current_enrollment.lock().await = None; + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Disabled); + remote_handle.shutdown_token.cancel(); + + let status = timeout( + Duration::from_secs(5), + remote_handle.enable(/*app_server_client_name*/ None), + ) + .await + .expect("cached enable should complete without network I/O") + .expect("shutdown should not cancel durable enable using persisted enrollment"); + + assert_eq!( + status.environment_id.as_deref(), + Some(persisted_enrollment.environment_id.as_str()) + ); + assert_eq!( + remote_handle + .current_enrollment + .snapshot() + .map(|enrollment| enrollment.server_id), + Some(persisted_enrollment.server_id.clone()) + ); + assert_eq!( + state_db + .get_remote_control_enrollment( + &persisted_enrollment.websocket_url, + &persisted_enrollment.account_id, + /*app_server_client_name*/ None, + ) + .await + .expect("enabled enrollment should load"), + Some(RemoteControlEnrollmentRecord { + remote_control_enabled: Some(true), + ..persisted_enrollment + }) + ); + assert_eq!( + *remote_handle.desired_state_tx.borrow(), + RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + } + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("persisted enrollment should prevent backend contact"); +} + +#[tokio::test] +async fn durable_enable_without_cached_enrollment_is_cancelled_after_shutdown() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = normalize_remote_control_url(&remote_control_url) + .expect("remote control target should normalize"); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + remote_handle.state_db = Some(state_db.clone()); + *remote_handle.current_enrollment.lock().await = None; + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Disabled); + remote_handle.shutdown_token.cancel(); + + let error = timeout( + Duration::from_secs(5), + remote_handle.enable(/*app_server_client_name*/ None), + ) + .await + .expect("shutdown should cancel network enrollment promptly") + .expect_err("enable without cached enrollment should be cancelled"); + + assert_eq!(error.kind(), std::io::ErrorKind::Interrupted); + assert_eq!(error.to_string(), "remote control is shutting down"); + assert_eq!(remote_handle.current_enrollment.snapshot(), None); + assert_eq!( + *remote_handle.desired_state_tx.borrow(), + RemoteControlDesiredState::Disabled + ); + assert_eq!( + state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + "account_id", + /*app_server_client_name*/ None, + ) + .await + .expect("enrollment should load"), + None + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("cancelled enrollment should prevent backend contact"); +} + +#[tokio::test] +async fn ephemeral_enable_preserves_durable_preference() { + let codex_home = TempDir::new().expect("temp dir should create"); + let mut remote_handle = remote_control_handle_with_current_enrollment( + TEST_REMOTE_CONTROL_URL, + remote_control_auth_manager(), + ); + remote_handle.state_db = Some(remote_control_state_runtime(&codex_home).await); + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + }); + + remote_handle + .enable_ephemeral() + .expect("ephemeral enable should succeed"); + assert_eq!( + *remote_handle.desired_state_tx.borrow(), + RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + } + ); + + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Disabled); + remote_handle + .enable_ephemeral() + .expect("ephemeral enable should succeed"); + assert_eq!( + *remote_handle.desired_state_tx.borrow(), + RemoteControlDesiredState::Enabled { + persistence_preference: None, + } + ); +} + +fn remote_control_server_token_response( + server_id: &str, + environment_id: &str, + remote_control_token: &str, +) -> serde_json::Value { + json!({ + "server_id": server_id, + "environment_id": environment_id, + "remote_control_token": remote_control_token, + "expires_at": TEST_REMOTE_CONTROL_SERVER_TOKEN_EXPIRES_AT, + }) +} + +async fn expect_remote_control_status( + status_rx: &mut watch::Receiver, + expected_status: Option, + expected_environment_id: Option<&str>, +) { + timeout(Duration::from_secs(5), status_rx.changed()) + .await + .expect("remote control status event should arrive in time") + .expect("remote control status watch should remain open"); + let status = status_rx.borrow(); + if let Some(expected_status) = expected_status { + assert_eq!(status.status, expected_status); + } + assert_eq!(status.server_name, test_server_name()); + assert_eq!(status.installation_id, TEST_INSTALLATION_ID); + assert_eq!(status.environment_id.as_deref(), expected_environment_id); +} + +async fn expect_remote_control_status_snapshot( + status_rx: &mut watch::Receiver, + expected_status: RemoteControlStatusChangedNotification, +) { + if *status_rx.borrow() == expected_status { + return; + } + + let expected_status_for_wait = expected_status.clone(); + let result = timeout(Duration::from_secs(5), async { + loop { + status_rx + .changed() + .await + .expect("remote control status watch should remain open"); + if *status_rx.borrow() == expected_status_for_wait { + return; + } + } + }) + .await; + assert!( + result.is_ok(), + "remote control status snapshot should arrive in time; expected {expected_status:?}, latest {:?}", + status_rx.borrow().clone() + ); +} + +#[tokio::test] +async fn remote_control_transport_manages_virtual_clients_and_routes_messages() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = normalize_remote_control_url(&remote_control_url) + .expect("remote control target should normalize"); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let (transport_event_tx, mut transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + let mut websocket = accept_remote_control_connection(&listener).await; + let enrollment = state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + "account_id", + /*app_server_client_name*/ None, + ) + .await + .expect("new enrollment should load") + .expect("new enrollment should exist"); + assert_eq!(enrollment.remote_control_enabled, None); + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_test"), + ) + .await; + + let client_id = ClientId("client-1".to_string()); + send_client_event( + &mut websocket, + ClientEnvelope { + event: ClientEvent::Ping, + client_id: client_id.clone(), + stream_id: None, + seq_id: None, + cursor: None, + }, + ) + .await; + assert_eq!( + read_server_event(&mut websocket).await, + json!({ + "type": "pong", + "client_id": "client-1", + "seq_id": 1, + "status": "unknown", + }) + ); + + send_client_event( + &mut websocket, + ClientEnvelope { + event: ClientEvent::ClientMessage { + message: JSONRPCMessage::Notification( + codex_app_server_protocol::JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }, + ), + }, + client_id: client_id.clone(), + stream_id: None, + seq_id: Some(0), + cursor: None, + }, + ) + .await; + assert!( + timeout(Duration::from_millis(100), transport_event_rx.recv()) + .await + .is_err(), + "non-initialize client messages should be ignored before connection creation" + ); + + let initialize_message = JSONRPCMessage::Request(codex_app_server_protocol::JSONRPCRequest { + id: codex_app_server_protocol::RequestId::Integer(1), + method: "initialize".to_string(), + params: Some(json!({ + "clientInfo": { + "name": "remote-test-client", + "version": "0.1.0" + } + })), + trace: None, + }); + send_client_event( + &mut websocket, + ClientEnvelope { + event: ClientEvent::ClientMessage { + message: initialize_message.clone(), + }, + client_id: client_id.clone(), + stream_id: None, + seq_id: Some(1), + cursor: None, + }, + ) + .await; + + let (connection_id, writer) = match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("connection open should arrive in time") + .expect("connection open should exist") + { + TransportEvent::ConnectionOpened { + connection_id, + origin, + writer, + .. + } => { + assert_eq!(origin, ConnectionOrigin::RemoteControl); + (connection_id, writer) + } + other => panic!("expected connection open event, got {other:?}"), + }; + + match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("initialize message should arrive in time") + .expect("initialize message should exist") + { + TransportEvent::IncomingMessage { + connection_id: incoming_connection_id, + message, + } => { + assert_eq!(incoming_connection_id, connection_id); + assert_eq!(message, initialize_message); + } + other => panic!("expected initialize incoming message, got {other:?}"), + } + + let followup_message = + JSONRPCMessage::Notification(codex_app_server_protocol::JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + send_client_event( + &mut websocket, + ClientEnvelope { + event: ClientEvent::ClientMessage { + message: followup_message.clone(), + }, + client_id: client_id.clone(), + stream_id: None, + seq_id: Some(2), + cursor: None, + }, + ) + .await; + match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("followup message should arrive in time") + .expect("followup message should exist") + { + TransportEvent::IncomingMessage { + connection_id: incoming_connection_id, + message, + } => { + assert_eq!(incoming_connection_id, connection_id); + assert_eq!(message, followup_message); + } + other => panic!("expected followup incoming message, got {other:?}"), + } + + send_client_event( + &mut websocket, + ClientEnvelope { + event: ClientEvent::Ping, + client_id: client_id.clone(), + stream_id: None, + seq_id: None, + cursor: None, + }, + ) + .await; + assert_eq!( + read_server_event(&mut websocket).await, + json!({ + "type": "pong", + "client_id": "client-1", + "seq_id": 1, + "status": "active", + }) + ); + + writer + .send(QueuedOutgoingMessage::new( + OutgoingMessage::AppServerNotification(ServerNotificationEnvelope { + notification: ServerNotification::ConfigWarning(ConfigWarningNotification { + summary: "test".to_string(), + details: None, + path: None, + range: None, + }), + emitted_at_ms: Some(1_234), + }), + )) + .await + .expect("remote writer should accept outgoing message"); + assert_eq!( + read_server_event(&mut websocket).await, + json!({ + "type": "server_message", + "client_id": "client-1", + "seq_id": 2, + "message": { + "method": "configWarning", + "params": { + "summary": "test", + "details": null, + }, + "emittedAtMs": 1_234, + } + }) + ); + + send_client_event( + &mut websocket, + ClientEnvelope { + event: ClientEvent::ClientClosed, + client_id: client_id.clone(), + stream_id: None, + seq_id: None, + cursor: None, + }, + ) + .await; + match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("connection close should arrive in time") + .expect("connection close should exist") + { + TransportEvent::ConnectionClosed { + connection_id: closed_connection_id, + } => { + assert_eq!(closed_connection_id, connection_id); + } + other => panic!("expected connection close event, got {other:?}"), + } + + send_client_event( + &mut websocket, + ClientEnvelope { + event: ClientEvent::Ping, + client_id, + stream_id: None, + seq_id: None, + cursor: None, + }, + ) + .await; + assert_eq!( + read_server_event(&mut websocket).await, + json!({ + "type": "pong", + "client_id": "client-1", + "seq_id": 1, + "status": "unknown", + }) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_transport_reconnects_after_disconnect() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let (transport_event_tx, mut transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(remote_control_state_runtime(&codex_home).await), + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + let (first_handshake_request, mut first_websocket) = + accept_remote_control_backend_connection(&listener).await; + assert_eq!( + first_handshake_request.headers.get("authorization"), + Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}")) + ); + first_websocket + .close(None) + .await + .expect("first websocket should close"); + drop(first_websocket); + + let (second_handshake_request, mut second_websocket) = + accept_remote_control_backend_connection(&listener).await; + assert_eq!( + second_handshake_request.headers.get("authorization"), + Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}")) + ); + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_test"), + ) + .await; + send_client_event( + &mut second_websocket, + ClientEnvelope { + event: ClientEvent::ClientMessage { + message: JSONRPCMessage::Request(codex_app_server_protocol::JSONRPCRequest { + id: codex_app_server_protocol::RequestId::Integer(2), + method: "initialize".to_string(), + params: Some(json!({ + "clientInfo": { + "name": "remote-test-client", + "version": "0.1.0" + } + })), + trace: None, + }), + }, + client_id: ClientId("client-2".to_string()), + stream_id: None, + seq_id: Some(0), + cursor: None, + }, + ) + .await; + + match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("reconnected initialize should arrive in time") + .expect("reconnected initialize should exist") + { + TransportEvent::ConnectionOpened { .. } => {} + other => panic!("expected connection open after reconnect, got {other:?}"), + } + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_transport_refreshes_server_token_after_websocket_unauthorized() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(remote_control_state_runtime(&codex_home).await), + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let websocket_request = accept_http_request(&listener).await; + assert_eq!( + websocket_request.request_line, + "GET /backend-api/wham/remote/control/server HTTP/1.1" + ); + assert_eq!( + websocket_request.headers.get("authorization"), + Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}")) + ); + respond_with_status(websocket_request.stream, "401 Unauthorized", "").await; + + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_json( + refresh_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let (handshake_request, _websocket) = accept_remote_control_backend_connection(&listener).await; + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_test"), + ) + .await; + assert_eq!( + handshake_request.headers.get("authorization"), + Some(&format!( + "Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}" + )) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_start_allows_remote_control_invalid_url_when_disabled() { + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, _remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url: "https://internal.example.com/backend-api/".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + /*state_db*/ None, + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::ResolvePersisted, + ) + .await + .expect("disabled remote control should not validate the URL at startup"); + + shutdown_token.cancel(); + timeout(Duration::from_secs(1), remote_task) + .await + .expect("remote control task should stop") + .expect("remote control task should join"); +} + +#[tokio::test] +async fn remote_control_start_allows_missing_auth_when_enabled() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, _remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(remote_control_state_runtime(&codex_home).await), + auth_manager, + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start before ChatGPT auth is available"); + + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("remote control should wait for auth before connecting"); + + shutdown_token.cancel(); + timeout(Duration::from_secs(1), remote_task) + .await + .expect("remote control task should stop") + .expect("remote control task should join"); +} + +#[tokio::test] +async fn remote_control_start_reports_missing_state_db_as_disabled_when_enabled() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + /*state_db*/ None, + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start disabled without sqlite state db"); + let mut status_rx = remote_handle.status_receiver(); + assert_eq!( + status_rx.borrow().clone(), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Disabled, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: None, + } + ); + + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("remote control should not connect without sqlite state db"); + + assert_eq!( + remote_handle + .enable_ephemeral() + .expect_err("enable should fail"), + RemoteControlEnableError::Unavailable(super::RemoteControlUnavailable) + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("remote control should remain disabled without sqlite state db"); + timeout(Duration::from_millis(20), status_rx.changed()) + .await + .expect_err("status should remain disabled without sqlite state db"); + + shutdown_token.cancel(); + timeout(Duration::from_secs(1), remote_task) + .await + .expect("remote control task should stop") + .expect("remote control task should join"); +} + +#[tokio::test] +async fn remote_control_handle_enable_disable_stops_and_restarts_connections() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(remote_control_state_runtime(&codex_home).await), + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + let mut first_websocket = accept_remote_control_connection(&listener).await; + expect_remote_control_status_snapshot( + &mut status_rx, + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connected, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_test".to_string()), + }, + ) + .await; + + assert_eq!( + remote_handle + .disable(Some("rpc-client")) + .await + .expect("disable should succeed"), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Disabled, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: None, + } + ); + expect_remote_control_status_snapshot( + &mut status_rx, + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Disabled, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: None, + }, + ) + .await; + timeout(Duration::from_secs(1), first_websocket.next()) + .await + .expect("disabling remote control should close the websocket"); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("disabled remote control should not reconnect"); + + assert_eq!( + remote_handle + .enable(Some("rpc-client")) + .await + .expect("enable should succeed"), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connecting, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_test".to_string()), + } + ); + expect_remote_control_status_snapshot( + &mut status_rx, + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connecting, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_test".to_string()), + }, + ) + .await; + let mut second_websocket = accept_remote_control_connection(&listener).await; + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_test"), + ) + .await; + second_websocket + .close(None) + .await + .expect("second websocket should close"); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_transport_clears_outgoing_buffer_when_backend_acks() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let (transport_event_tx, mut transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(remote_control_state_runtime(&codex_home).await), + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + + let enroll_request = accept_http_request(&listener).await; + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + let mut first_websocket = accept_remote_control_connection(&listener).await; + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_test"), + ) + .await; + + let client_id = ClientId("client-1".to_string()); + let initialize_message = JSONRPCMessage::Request(codex_app_server_protocol::JSONRPCRequest { + id: codex_app_server_protocol::RequestId::Integer(1), + method: "initialize".to_string(), + params: Some(json!({ + "clientInfo": { + "name": "remote-test-client", + "version": "0.1.0" + } + })), + trace: None, + }); + send_client_event( + &mut first_websocket, + ClientEnvelope { + event: ClientEvent::ClientMessage { + message: initialize_message, + }, + client_id: client_id.clone(), + stream_id: None, + seq_id: Some(0), + cursor: None, + }, + ) + .await; + + let writer = match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("connection open should arrive in time") + .expect("connection open should exist") + { + TransportEvent::ConnectionOpened { writer, .. } => writer, + other => panic!("expected connection open event, got {other:?}"), + }; + match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("initialize message should arrive in time") + .expect("initialize message should exist") + { + TransportEvent::IncomingMessage { .. } => {} + other => panic!("expected initialize incoming message, got {other:?}"), + } + + writer + .send(QueuedOutgoingMessage::new( + OutgoingMessage::AppServerNotification(ServerNotificationEnvelope { + notification: ServerNotification::ConfigWarning(ConfigWarningNotification { + summary: "stale".to_string(), + details: None, + path: None, + range: None, + }), + emitted_at_ms: Some(1_234), + }), + )) + .await + .expect("remote writer should accept outgoing message"); + let (server_event, stream_id) = read_server_event_with_stream_id(&mut first_websocket).await; + assert_eq!( + server_event, + json!({ + "type": "server_message", + "client_id": "client-1", + "seq_id": 1, + "message": { + "method": "configWarning", + "params": { + "summary": "stale", + "details": null, + }, + "emittedAtMs": 1_234, + } + }) + ); + + send_client_event( + &mut first_websocket, + ClientEnvelope { + event: ClientEvent::Ack { segment_id: None }, + client_id: client_id.clone(), + stream_id: Some(stream_id), + seq_id: Some(1), + cursor: None, + }, + ) + .await; + + send_client_event( + &mut first_websocket, + ClientEnvelope { + event: ClientEvent::ClientClosed, + client_id: client_id.clone(), + stream_id: None, + seq_id: None, + cursor: None, + }, + ) + .await; + match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("connection close should arrive in time") + .expect("connection close should exist") + { + TransportEvent::ConnectionClosed { .. } => {} + other => panic!("expected connection close event, got {other:?}"), + } + + first_websocket + .close(None) + .await + .expect("first websocket should close"); + drop(first_websocket); + + let mut second_websocket = accept_remote_control_connection(&listener).await; + send_client_event( + &mut second_websocket, + ClientEnvelope { + event: ClientEvent::Ping, + client_id, + stream_id: None, + seq_id: None, + cursor: None, + }, + ) + .await; + assert_eq!( + read_server_event(&mut second_websocket).await, + json!({ + "type": "pong", + "client_id": "client-1", + "seq_id": 1, + "status": "unknown", + }) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_http_mode_enrolls_before_connecting() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let (transport_event_tx, mut transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let expected_server_name = gethostname().to_string_lossy().trim().to_string(); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(remote_control_state_runtime(&codex_home).await), + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + assert_eq!( + enroll_request.headers.get("authorization"), + Some(&"Bearer Access Token".to_string()) + ); + assert_eq!( + enroll_request + .headers + .get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + vec!["account_id"] + ); + assert_eq!( + enroll_request + .headers + .get_all(REMOTE_CONTROL_INSTALLATION_ID_HEADER), + vec![TEST_INSTALLATION_ID] + ); + assert_eq!( + serde_json::from_str::(&enroll_request.body) + .expect("enroll body should deserialize"), + json!({ + "name": expected_server_name, + "os": std::env::consts::OS, + "arch": std::env::consts::ARCH, + "app_server_version": env!("CARGO_PKG_VERSION"), + "installation_id": TEST_INSTALLATION_ID, + }) + ); + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let (handshake_request, mut websocket) = + accept_remote_control_backend_connection(&listener).await; + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_test"), + ) + .await; + assert_eq!( + handshake_request.path, + "/backend-api/wham/remote/control/server" + ); + assert_eq!( + handshake_request.headers.get("authorization"), + Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}")) + ); + assert_eq!( + handshake_request + .headers + .get(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + None + ); + assert_eq!( + handshake_request + .headers + .get(REMOTE_CONTROL_INSTALLATION_ID_HEADER), + Some(&TEST_INSTALLATION_ID.to_string()) + ); + assert_eq!( + handshake_request.headers.get("x-codex-server-id"), + Some(&"srv_e_test".to_string()) + ); + assert_eq!( + handshake_request.headers.get("x-codex-name"), + Some(&base64::engine::general_purpose::STANDARD.encode(&expected_server_name)) + ); + assert_eq!( + handshake_request.headers.get("x-codex-protocol-version"), + Some(&REMOTE_CONTROL_PROTOCOL_VERSION.to_string()) + ); + + let backend_client_id = ClientId("backend-test-client".to_string()); + let writer = { + let initialize_message = + JSONRPCMessage::Request(codex_app_server_protocol::JSONRPCRequest { + id: codex_app_server_protocol::RequestId::Integer(11), + method: "initialize".to_string(), + params: Some(json!({ + "clientInfo": { + "name": "remote-backend-client", + "version": "0.1.0" + } + })), + trace: None, + }); + send_client_event( + &mut websocket, + ClientEnvelope { + event: ClientEvent::ClientMessage { + message: initialize_message.clone(), + }, + client_id: backend_client_id.clone(), + stream_id: None, + seq_id: Some(0), + cursor: None, + }, + ) + .await; + + let (connection_id, writer) = + match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("connection open should arrive in time") + .expect("connection open should exist") + { + TransportEvent::ConnectionOpened { + connection_id, + writer, + .. + } => (connection_id, writer), + other => panic!("expected connection open event, got {other:?}"), + }; + + match timeout(Duration::from_secs(5), transport_event_rx.recv()) + .await + .expect("initialize message should arrive in time") + .expect("initialize message should exist") + { + TransportEvent::IncomingMessage { + connection_id: incoming_connection_id, + message, + } => { + assert_eq!(incoming_connection_id, connection_id); + assert_eq!(message, initialize_message); + } + other => panic!("expected initialize incoming message, got {other:?}"), + } + writer + }; + + writer + .send(QueuedOutgoingMessage::new(OutgoingMessage::Response( + crate::outgoing_message::OutgoingResponse { + id: codex_app_server_protocol::RequestId::Integer(11), + result: Box::new( + codex_app_server_protocol::ClientResponsePayload::Initialize( + codex_app_server_protocol::InitializeResponse { + user_agent: "codex-test-agent".to_string(), + codex_home: codex_home.path().abs(), + platform_family: "test-family".to_string(), + platform_os: "test-os".to_string(), + }, + ), + ), + }, + ))) + .await + .expect("remote writer should accept initialize response"); + assert_eq!( + read_server_event(&mut websocket).await, + json!({ + "type": "server_message", + "client_id": backend_client_id.0.clone(), + "seq_id": 1, + "message": { + "id": 11, + "result": { + "userAgent": "codex-test-agent", + "codexHome": codex_home.path(), + "platformFamily": "test-family", + "platformOs": "test-os", + } + } + }) + ); + + writer + .send(QueuedOutgoingMessage::new( + OutgoingMessage::AppServerNotification(ServerNotificationEnvelope { + notification: ServerNotification::ConfigWarning(ConfigWarningNotification { + summary: "backend".to_string(), + details: None, + path: None, + range: None, + }), + emitted_at_ms: Some(1_234), + }), + )) + .await + .expect("remote writer should accept outgoing message"); + assert_eq!( + read_server_event(&mut websocket).await, + json!({ + "type": "server_message", + "client_id": backend_client_id.0.clone(), + "seq_id": 2, + "message": { + "method": "configWarning", + "params": { + "summary": "backend", + "details": null, + }, + "emittedAtMs": 1_234, + } + }) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_http_mode_refreshes_persisted_enrollment_before_connecting() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let persisted_enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_id".to_string(), + environment_id: "env_persisted".to_string(), + server_id: "srv_e_persisted".to_string(), + server_name: "persisted-server".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_id", + /*app_server_client_name*/ None, + Some(&persisted_enrollment), + /*remote_control_enabled*/ None, + ) + .await + .expect("persisted enrollment should save"); + + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, _remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + remote_control_auth_manager_with_home(&codex_home), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + assert_eq!( + refresh_request.headers.get("authorization"), + Some(&"Bearer Access Token".to_string()) + ); + assert_eq!( + refresh_request + .headers + .get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + vec!["account_id"] + ); + assert_eq!( + refresh_request + .headers + .get_all(REMOTE_CONTROL_INSTALLATION_ID_HEADER), + vec![TEST_INSTALLATION_ID] + ); + assert_eq!( + serde_json::from_str::(&refresh_request.body) + .expect("refresh body should deserialize"), + json!({ + "server_id": persisted_enrollment.server_id.clone(), + "installation_id": TEST_INSTALLATION_ID, + }) + ); + respond_with_json( + refresh_request.stream, + remote_control_server_token_response( + &persisted_enrollment.server_id, + &persisted_enrollment.environment_id, + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let (handshake_request, _websocket) = accept_remote_control_backend_connection(&listener).await; + assert_eq!( + handshake_request.path, + "/backend-api/wham/remote/control/server" + ); + assert_eq!( + handshake_request.headers.get("x-codex-server-id"), + Some(&persisted_enrollment.server_id) + ); + assert_eq!( + handshake_request.headers.get("authorization"), + Some(&format!( + "Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}" + )) + ); + assert_eq!( + load_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_id", + /*app_server_client_name*/ None, + ) + .await + .expect("persisted enrollment should load"), + Some(persisted_enrollment) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_stdio_mode_waits_for_client_name_before_connecting() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let app_server_client_name = "stdio-client"; + let persisted_enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_id".to_string(), + environment_id: "env_persisted".to_string(), + server_id: "srv_e_persisted".to_string(), + server_name: "persisted-server".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_id", + Some(app_server_client_name), + Some(&persisted_enrollment), + /*remote_control_enabled*/ None, + ) + .await + .expect("persisted enrollment should save"); + + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let (app_server_client_name_tx, app_server_client_name_rx) = oneshot::channel::(); + let shutdown_token = CancellationToken::new(); + let (remote_task, _remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + remote_control_auth_manager_with_home(&codex_home), + transport_event_tx, + shutdown_token.clone(), + Some(app_server_client_name_rx), + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("remote control should wait for the stdio client name"); + + let _ = app_server_client_name_tx.send(app_server_client_name.to_string()); + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_json( + refresh_request.stream, + remote_control_server_token_response( + &persisted_enrollment.server_id, + &persisted_enrollment.environment_id, + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + let (handshake_request, _websocket) = accept_remote_control_backend_connection(&listener).await; + assert_eq!( + handshake_request.headers.get("x-codex-server-id"), + Some(&persisted_enrollment.server_id) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_waits_for_account_id_before_enrolling() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json(/*account_id*/ None), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("auth without account id should save"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let expected_server_name = gethostname().to_string_lossy().trim().to_string(); + let expected_remote_control_target = normalize_remote_control_url(&remote_control_url) + .expect("remote control target should normalize"); + let expected_enrollment = RemoteControlEnrollment { + remote_control_target: expected_remote_control_target, + account_id: "account_id".to_string(), + environment_id: "env_ready".to_string(), + server_id: "srv_e_ready".to_string(), + server_name: expected_server_name, + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, _remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + auth_manager.clone(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start before account id is available"); + + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("remote control should wait for account id before enrolling"); + + save_auth( + codex_home.path(), + &remote_control_auth_dot_json(Some("account_id")), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("auth with account id should save"); + auth_manager.reload().await; + + let enroll_request = timeout(Duration::from_millis(100), accept_http_request(&listener)) + .await + .expect("auth change should wake remote control before the retry delay"); + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + &expected_enrollment.server_id, + &expected_enrollment.environment_id, + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let (handshake_request, _websocket) = accept_remote_control_backend_connection(&listener).await; + assert_eq!( + handshake_request.headers.get("x-codex-server-id"), + Some(&expected_enrollment.server_id) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn persisted_enable_does_not_follow_auth_to_an_account_without_a_preference() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json(Some("account_a")), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("account A auth should save"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_a".to_string(), + environment_id: "env_a".to_string(), + server_id: "srv_e_a".to_string(), + server_name: "server-a".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_a", + /*app_server_client_name*/ None, + Some(&enrollment), + /*remote_control_enabled*/ Some(true), + ) + .await + .expect("account A enrollment should save"); + + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + auth_manager.clone(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::ResolvePersisted, + ) + .await + .expect("remote control should start"); + + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_json( + refresh_request.stream, + remote_control_server_token_response( + &enrollment.server_id, + &enrollment.environment_id, + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + let (_handshake_request, mut websocket) = + accept_remote_control_backend_connection(&listener).await; + + save_auth( + codex_home.path(), + &remote_control_auth_dot_json(Some("account_b")), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("account B auth should save"); + auth_manager.reload().await; + let closed = timeout(Duration::from_secs(1), websocket.next()) + .await + .expect("account switch should close the backend websocket"); + assert!(matches!( + closed, + None | Some(Err(_)) | Some(Ok(tungstenite::Message::Close(_))) + )); + + let mut desired_state_rx = remote_handle.status_receiver(); + timeout( + Duration::from_secs(1), + desired_state_rx.wait_for(|state| state.status == RemoteControlConnectionStatus::Disabled), + ) + .await + .expect("account B missing preference should disable remote control") + .expect("desired state channel should stay open"); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("disabled account B should not enroll"); + assert_eq!( + state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + "account_b", + /*app_server_client_name*/ None, + ) + .await + .expect("account B enrollment should load"), + None + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_http_mode_reenrolls_when_refresh_reports_stale_enrollment() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let expected_server_name = gethostname().to_string_lossy().trim().to_string(); + let stale_enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_id".to_string(), + environment_id: "env_stale".to_string(), + server_id: "srv_e_stale".to_string(), + server_name: "stale-server".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + let refreshed_enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_id".to_string(), + environment_id: "env_refreshed".to_string(), + server_id: "srv_e_refreshed".to_string(), + server_name: expected_server_name, + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_id", + /*app_server_client_name*/ None, + Some(&stale_enrollment), + /*remote_control_enabled*/ Some(true), + ) + .await + .expect("stale enrollment should save"); + + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + remote_control_auth_manager_with_home(&codex_home), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::ResolvePersisted, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_stale"), + ) + .await; + respond_with_status(refresh_request.stream, "404 Not Found", "").await; + + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + &refreshed_enrollment.server_id, + &refreshed_enrollment.environment_id, + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let (handshake_request, _websocket) = accept_remote_control_backend_connection(&listener).await; + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_refreshed"), + ) + .await; + assert_eq!( + handshake_request.headers.get("x-codex-server-id"), + Some(&refreshed_enrollment.server_id) + ); + assert_eq!( + state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + "account_id", + /*app_server_client_name*/ None, + ) + .await + .expect("refreshed enrollment should load"), + Some(RemoteControlEnrollmentRecord { + websocket_url: remote_control_target.websocket_url.clone(), + account_id: "account_id".to_string(), + app_server_client_name: None, + server_id: refreshed_enrollment.server_id.clone(), + environment_id: refreshed_enrollment.environment_id.clone(), + server_name: refreshed_enrollment.server_name.clone(), + remote_control_enabled: Some(true), + }) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_http_mode_reenrolls_after_explicit_missing_server_404() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let expected_server_name = gethostname().to_string_lossy().trim().to_string(); + let stale_enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_id".to_string(), + environment_id: "env_stale".to_string(), + server_id: "srv_e_stale".to_string(), + server_name: "stale-server".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + let refreshed_enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_id".to_string(), + environment_id: "env_refreshed".to_string(), + server_id: "srv_e_refreshed".to_string(), + server_name: expected_server_name, + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_id", + /*app_server_client_name*/ None, + Some(&stale_enrollment), + /*remote_control_enabled*/ Some(true), + ) + .await + .expect("stale enrollment should save"); + + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + remote_control_auth_manager_with_home(&codex_home), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::ResolvePersisted, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_json( + refresh_request.stream, + remote_control_server_token_response( + &stale_enrollment.server_id, + &stale_enrollment.environment_id, + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let websocket_request = accept_http_request(&listener).await; + assert_eq!( + websocket_request.request_line, + "GET /backend-api/wham/remote/control/server HTTP/1.1" + ); + assert_eq!( + websocket_request.headers.get("x-codex-server-id"), + Some(&stale_enrollment.server_id) + ); + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_stale"), + ) + .await; + respond_with_status( + websocket_request.stream, + "404 Not Found", + &json!({"detail": "Remote app server not found"}).to_string(), + ) + .await; + + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + &refreshed_enrollment.server_id, + &refreshed_enrollment.environment_id, + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let (handshake_request, _websocket) = accept_remote_control_backend_connection(&listener).await; + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_refreshed"), + ) + .await; + assert_eq!( + handshake_request.headers.get("x-codex-server-id"), + Some(&refreshed_enrollment.server_id) + ); + assert_eq!( + state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + "account_id", + /*app_server_client_name*/ None, + ) + .await + .expect("refreshed enrollment should load"), + Some(RemoteControlEnrollmentRecord { + websocket_url: remote_control_target.websocket_url.clone(), + account_id: "account_id".to_string(), + app_server_client_name: None, + server_id: refreshed_enrollment.server_id.clone(), + environment_id: refreshed_enrollment.environment_id.clone(), + server_name: refreshed_enrollment.server_name.clone(), + remote_control_enabled: Some(true), + }) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_http_mode_preserves_stale_enrollment_when_reenrollment_fails() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let stale_enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_id".to_string(), + environment_id: "env_stale".to_string(), + server_id: "srv_e_stale".to_string(), + server_name: test_server_name(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_id", + /*app_server_client_name*/ None, + Some(&stale_enrollment), + /*remote_control_enabled*/ Some(true), + ) + .await + .expect("stale enrollment should save"); + + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + remote_control_auth_manager_with_home(&codex_home), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::ResolvePersisted, + ) + .await + .expect("remote control should start"); + + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_status(refresh_request.stream, "404 Not Found", "").await; + + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_status(enroll_request.stream, "500 Internal Server Error", "failed").await; + + let retry_refresh_request = accept_http_request(&listener).await; + assert_eq!( + retry_refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + let refresh_failed_at = OffsetDateTime::now_utc(); + respond_with_status( + retry_refresh_request.stream, + "500 Internal Server Error", + "failed", + ) + .await; + + let current_enrollment = remote_handle + .inner + .session() + .current_enrollment + .lock() + .await + .clone() + .expect("stale enrollment should remain available"); + let next_refresh_at = current_enrollment + .next_refresh_at + .expect("required refresh failure should set a retry deadline"); + assert!( + (refresh_failed_at + time::Duration::seconds(24) + ..=OffsetDateTime::now_utc() + time::Duration::seconds(36)) + .contains(&next_refresh_at) + ); + assert_eq!( + current_enrollment, + RemoteControlEnrollment { + next_refresh_at: Some(next_refresh_at), + ..stale_enrollment.clone() + } + ); + assert_eq!( + state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + "account_id", + /*app_server_client_name*/ None, + ) + .await + .expect("stale enrollment should load"), + Some(RemoteControlEnrollmentRecord { + websocket_url: remote_control_target.websocket_url, + account_id: "account_id".to_string(), + app_server_client_name: None, + server_id: stale_enrollment.server_id, + environment_id: stale_enrollment.environment_id, + server_name: stale_enrollment.server_name, + remote_control_enabled: Some(true), + }) + ); + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_http_mode_preserves_enrollment_after_generic_websocket_404() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let stale_enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_id".to_string(), + environment_id: "env_stale".to_string(), + server_id: "srv_e_stale".to_string(), + server_name: "stale-server".to_string(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_id", + /*app_server_client_name*/ None, + Some(&stale_enrollment), + /*remote_control_enabled*/ None, + ) + .await + .expect("stale enrollment should save"); + + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(state_db.clone()), + remote_control_auth_manager_with_home(&codex_home), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_json( + refresh_request.stream, + remote_control_server_token_response( + &stale_enrollment.server_id, + &stale_enrollment.environment_id, + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let websocket_request = accept_http_request(&listener).await; + assert_eq!( + websocket_request.request_line, + "GET /backend-api/wham/remote/control/server HTTP/1.1" + ); + assert_eq!( + websocket_request.headers.get("x-codex-server-id"), + Some(&stale_enrollment.server_id) + ); + assert_eq!( + websocket_request.headers.get("authorization"), + Some(&format!( + "Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}" + )) + ); + expect_remote_control_status( + &mut status_rx, + /*expected_status*/ None, + Some("env_stale"), + ) + .await; + respond_with_status_and_headers( + websocket_request.stream, + "404 Not Found", + &[("x-request-id", "request-404"), ("cf-ray", "ray-404")], + "Not Found", + ) + .await; + + assert_eq!( + load_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_id", + /*app_server_client_name*/ None, + ) + .await + .expect("stale enrollment should load"), + Some(stale_enrollment.clone()) + ); + + let (handshake_request, _websocket) = accept_remote_control_backend_connection(&listener).await; + assert_eq!( + handshake_request.headers.get("x-codex-server-id"), + Some(&stale_enrollment.server_id) + ); + assert_eq!( + handshake_request.headers.get("authorization"), + Some(&format!( + "Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}" + )) + ); + expect_remote_control_status_snapshot( + &mut status_rx, + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connected, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_stale".to_string()), + }, + ) + .await; + + shutdown_token.cancel(); + let _ = remote_task.await; +} + +#[derive(Debug)] +struct CapturedHttpRequest { + stream: TcpStream, + request_line: String, + headers: CapturedHttpHeaders, + body: String, +} + +#[derive(Debug, Default)] +struct CapturedHttpHeaders(Vec<(String, String)>); + +impl CapturedHttpHeaders { + fn append(&mut self, name: String, value: String) { + self.0.push((name, value)); + } + + fn get(&self, name: &str) -> Option<&String> { + self.0 + .iter() + .rev() + .find(|(candidate, _value)| candidate.eq_ignore_ascii_case(name)) + .map(|(_name, value)| value) + } + + fn get_all(&self, name: &str) -> Vec<&str> { + self.0 + .iter() + .filter(|(candidate, _value)| candidate.eq_ignore_ascii_case(name)) + .map(|(_name, value)| value.as_str()) + .collect() + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +struct CapturedWebSocketRequest { + path: String, + headers: BTreeMap, +} + +async fn accept_remote_control_connection(listener: &TcpListener) -> WebSocketStream { + let (stream, _) = timeout(Duration::from_secs(5), listener.accept()) + .await + .expect("remote control should connect in time") + .expect("listener accept should succeed"); + accept_async(stream) + .await + .expect("websocket handshake should succeed") +} + +async fn accept_http_request(listener: &TcpListener) -> CapturedHttpRequest { + let (stream, _) = timeout(Duration::from_secs(5), listener.accept()) + .await + .expect("HTTP request should arrive in time") + .expect("listener accept should succeed"); + let mut reader = BufReader::new(stream); + + let mut request_line = String::new(); + reader + .read_line(&mut request_line) + .await + .expect("request line should read"); + let request_line = request_line.trim_end_matches("\r\n").to_string(); + + let mut headers = CapturedHttpHeaders::default(); + loop { + let mut line = String::new(); + reader + .read_line(&mut line) + .await + .expect("header line should read"); + if line == "\r\n" { + break; + } + let line = line.trim_end_matches("\r\n"); + let (name, value) = line.split_once(':').expect("header should contain colon"); + headers.append(name.to_ascii_lowercase(), value.trim().to_string()); + } + + let content_length = headers + .get("content-length") + .and_then(|value| value.parse::().ok()) + .unwrap_or(0); + let mut body = vec![0; content_length]; + reader + .read_exact(&mut body) + .await + .expect("request body should read"); + + CapturedHttpRequest { + stream: reader.into_inner(), + request_line, + headers, + body: String::from_utf8(body).expect("body should be utf-8"), + } +} + +async fn respond_with_json(mut stream: TcpStream, body: serde_json::Value) { + let body = body.to_string(); + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}", + body.len() + ); + stream + .write_all(response.as_bytes()) + .await + .expect("response should write"); + stream.flush().await.expect("response should flush"); +} + +async fn respond_with_status(stream: TcpStream, status: &str, body: &str) { + respond_with_status_and_headers(stream, status, &[], body).await; +} + +async fn respond_with_status_and_headers( + mut stream: TcpStream, + status: &str, + headers: &[(&str, &str)], + body: &str, +) { + let extra_headers = headers + .iter() + .map(|(name, value)| format!("{name}: {value}\r\n")) + .collect::(); + let response = format!( + "HTTP/1.1 {status}\r\ncontent-type: text/plain\r\ncontent-length: {}\r\nconnection: close\r\n{extra_headers}\r\n{body}", + body.len(), + ); + stream + .write_all(response.as_bytes()) + .await + .expect("response should write"); + stream.flush().await.expect("response should flush"); +} + +async fn accept_remote_control_backend_connection( + listener: &TcpListener, +) -> (CapturedWebSocketRequest, WebSocketStream) { + let (stream, _) = timeout(Duration::from_secs(5), listener.accept()) + .await + .expect("websocket request should arrive in time") + .expect("listener accept should succeed"); + let captured_request = Arc::new(std::sync::Mutex::new(None::)); + let captured_request_for_callback = captured_request.clone(); + let websocket = accept_hdr_async( + stream, + move |request: &tungstenite::handshake::server::Request, + response: tungstenite::handshake::server::Response| { + let headers = request + .headers() + .iter() + .map(|(name, value)| { + ( + name.as_str().to_ascii_lowercase(), + value + .to_str() + .expect("header should be valid utf-8") + .to_string(), + ) + }) + .collect::>(); + *captured_request_for_callback + .lock() + .expect("capture lock should acquire") = Some(CapturedWebSocketRequest { + path: request.uri().path().to_string(), + headers, + }); + Ok(response) + }, + ) + .await + .expect("websocket handshake should succeed"); + let captured_request = captured_request + .lock() + .expect("capture lock should acquire") + .clone() + .expect("websocket request should be captured"); + (captured_request, websocket) +} + +async fn send_client_event( + websocket: &mut WebSocketStream, + client_envelope: ClientEnvelope, +) { + let payload = serde_json::to_string(&client_envelope).expect("client event should serialize"); + websocket + .send(tungstenite::Message::Text(payload.into())) + .await + .expect("client event should send"); +} + +async fn read_server_event(websocket: &mut WebSocketStream) -> serde_json::Value { + read_server_event_with_stream_id(websocket).await.0 +} + +async fn read_server_event_with_stream_id( + websocket: &mut WebSocketStream, +) -> (serde_json::Value, StreamId) { + loop { + let frame = timeout(Duration::from_secs(5), websocket.next()) + .await + .expect("server event should arrive in time") + .expect("websocket should stay open") + .expect("websocket frame should be readable"); + match frame { + tungstenite::Message::Text(text) => { + let mut event: serde_json::Value = + serde_json::from_str(text.as_ref()).expect("server event should deserialize"); + let stream_id = event + .as_object_mut() + .and_then(|event| event.remove("stream_id")) + .expect("stream_id should be present"); + let stream_id = stream_id + .as_str() + .expect("stream_id should be a string") + .to_string(); + return (event, StreamId(stream_id)); + } + tungstenite::Message::Ping(payload) => { + websocket + .send(tungstenite::Message::Pong(payload)) + .await + .expect("websocket pong should send"); + } + tungstenite::Message::Pong(_) => {} + tungstenite::Message::Close(frame) => { + panic!("unexpected websocket close frame: {frame:?}"); + } + tungstenite::Message::Binary(_) => { + panic!("unexpected binary websocket frame"); + } + tungstenite::Message::Frame(_) => {} + } + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..1aea56af42dd739819ea8920211fc36d8de9c6b6 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs @@ -0,0 +1,416 @@ +use super::super::clients::list_remote_control_clients; +use super::super::clients::revoke_remote_control_client; +use super::*; +use codex_app_server_protocol::RemoteControlClient; +use codex_app_server_protocol::RemoteControlClientsListOrder; +use codex_app_server_protocol::RemoteControlClientsListParams; +use codex_app_server_protocol::RemoteControlClientsListResponse; +use codex_app_server_protocol::RemoteControlClientsRevokeParams; +use codex_app_server_protocol::RemoteControlClientsRevokeResponse; +use codex_login::AuthKeyringBackendKind; +use pretty_assertions::assert_eq; + +fn client_management_handle( + remote_control_url: String, + auth_manager: Arc, +) -> RemoteControlSession { + let desired_state_tx = watch::channel(RemoteControlDesiredState::Disabled).0; + let (status_tx, _status_rx) = watch::channel(RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Disabled, + server_name: test_server_name(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: None, + }); + RemoteControlSession { + policy: RemoteControlPolicy::Allowed, + shutdown_token: CancellationToken::new(), + desired_state_tx: Arc::new(desired_state_tx), + desired_state_rpc_lock: Arc::new(Semaphore::new(1)), + persistence: RemoteControlPersistence::default(), + status_tx: Arc::new(status_tx), + state_db: None, + remote_control_url, + current_enrollment: Arc::new(RemoteControlEnrollmentState::new(/*enrollment*/ None)), + pairing_persistence_key: watch::channel(None).0, + pairing_persistence_key_required: false, + auth_manager: auth::RemoteControlAuth::capture(auth_manager).0, + } +} + +fn empty_client_list() -> serde_json::Value { + json!({ + "items": [], + "cursor": null, + }) +} + +#[tokio::test] +async fn remote_control_handle_lists_clients_while_disabled() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let request = accept_http_request(&listener).await; + assert_eq!( + request.request_line, + "GET /backend-api/wham/remote/control/environments/env%20%2F%3F/clients?cursor=cursor+%2F%3F&limit=10&order=asc HTTP/1.1" + ); + assert_eq!( + request.headers.get("authorization"), + Some(&"Bearer Access Token".to_string()) + ); + assert_eq!( + request.headers.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + vec!["account_id"] + ); + respond_with_json( + request.stream, + json!({ + "items": [{ + "client_id": "client-123", + "account_user_id": "user-123", + "enrollment_status": "enrolled_device_key", + "display_name": "Anton Phone", + "device_type": "phone", + "platform": "ios", + "os_version": "19.0", + "device_model": "iPhone", + "app_version": "1.2.3", + "last_seen_at": "2026-03-05T07:00:00Z", + "last_seen_city": "San Francisco", + }], + "cursor": "next-cursor", + }), + ) + .await; + }); + let handle = client_management_handle(remote_control_url, remote_control_auth_manager()); + + let response = handle + .list_clients(RemoteControlClientsListParams { + environment_id: "env /?".to_string(), + cursor: Some("cursor /?".to_string()), + limit: Some(10), + order: Some(RemoteControlClientsListOrder::Asc), + }) + .await + .expect("client list should succeed while remote control is disabled"); + server_task.await.expect("server task should finish"); + + assert_eq!( + response, + RemoteControlClientsListResponse { + data: vec![RemoteControlClient { + client_id: "client-123".to_string(), + display_name: Some("Anton Phone".to_string()), + device_type: Some("phone".to_string()), + platform: Some("ios".to_string()), + os_version: Some("19.0".to_string()), + device_model: Some("iPhone".to_string()), + app_version: Some("1.2.3".to_string()), + last_seen_at: Some(1_772_694_000), + }], + next_cursor: Some("next-cursor".to_string()), + } + ); +} + +#[tokio::test] +async fn remote_control_handle_revokes_client_while_disabled() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let request = accept_http_request(&listener).await; + assert_eq!( + request.request_line, + "DELETE /backend-api/wham/remote/control/environments/env%20%2F%3F/clients/client%20%2F%3F HTTP/1.1" + ); + assert_eq!( + request.headers.get("authorization"), + Some(&"Bearer Access Token".to_string()) + ); + assert_eq!( + request.headers.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + vec!["account_id"] + ); + respond_with_status(request.stream, "204 No Content", "").await; + }); + let handle = client_management_handle(remote_control_url, remote_control_auth_manager()); + + let response = handle + .revoke_client(RemoteControlClientsRevokeParams { + environment_id: "env /?".to_string(), + client_id: "client /?".to_string(), + }) + .await + .expect("client revoke should succeed while remote control is disabled"); + server_task.await.expect("server task should finish"); + + assert_eq!(response, RemoteControlClientsRevokeResponse {}); +} + +#[tokio::test] +async fn list_remote_control_clients_recovers_auth_after_unauthorized() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let stale_request = accept_http_request(&listener).await; + assert_eq!( + stale_request.headers.get("authorization"), + Some(&"Bearer stale-token".to_string()) + ); + assert_eq!( + stale_request + .headers + .get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + vec!["account_id"] + ); + respond_with_status(stale_request.stream, "401 Unauthorized", "").await; + + let recovered_request = accept_http_request(&listener).await; + assert_eq!( + recovered_request.headers.get("authorization"), + Some(&"Bearer fresh-token".to_string()) + ); + assert_eq!( + recovered_request + .headers + .get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + vec!["account_id"] + ); + respond_with_json(recovered_request.stream, empty_client_list()).await; + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let mut stale_auth = remote_control_auth_dot_json(Some("account_id")); + stale_auth + .tokens + .as_mut() + .expect("stale auth should include tokens") + .access_token = "stale-token".to_string(); + save_auth( + codex_home.path(), + &stale_auth, + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("stale auth should save"); + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let mut fresh_auth = remote_control_auth_dot_json(Some("account_id")); + fresh_auth + .tokens + .as_mut() + .expect("fresh auth should include tokens") + .access_token = "fresh-token".to_string(); + save_auth( + codex_home.path(), + &fresh_auth, + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("fresh auth should save"); + + let response = list_remote_control_clients( + &remote_control_url, + &auth::RemoteControlAuth::capture(auth_manager.clone()).0, + RemoteControlClientsListParams { + environment_id: "env-123".to_string(), + ..Default::default() + }, + ) + .await + .expect("client list should recover auth"); + server_task.await.expect("server task should finish"); + + assert_eq!( + response, + RemoteControlClientsListResponse { + data: Vec::new(), + next_cursor: None, + } + ); +} + +#[tokio::test] +async fn list_remote_control_clients_retries_unauthorized_only_once() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let stale_request = accept_http_request(&listener).await; + assert_eq!( + stale_request.headers.get("authorization"), + Some(&"Bearer stale-token".to_string()) + ); + assert_eq!( + stale_request + .headers + .get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + vec!["account_id"] + ); + respond_with_status(stale_request.stream, "401 Unauthorized", "").await; + + let recovered_request = accept_http_request(&listener).await; + assert_eq!( + recovered_request.headers.get("authorization"), + Some(&"Bearer fresh-token".to_string()) + ); + assert_eq!( + recovered_request + .headers + .get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + vec!["account_id"] + ); + respond_with_status(recovered_request.stream, "401 Unauthorized", "").await; + + assert!( + timeout(Duration::from_millis(100), accept_http_request(&listener)) + .await + .is_err() + ); + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let mut stale_auth = remote_control_auth_dot_json(Some("account_id")); + stale_auth + .tokens + .as_mut() + .expect("stale auth should include tokens") + .access_token = "stale-token".to_string(); + save_auth( + codex_home.path(), + &stale_auth, + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("stale auth should save"); + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let mut fresh_auth = remote_control_auth_dot_json(Some("account_id")); + fresh_auth + .tokens + .as_mut() + .expect("fresh auth should include tokens") + .access_token = "fresh-token".to_string(); + save_auth( + codex_home.path(), + &fresh_auth, + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("fresh auth should save"); + + let err = list_remote_control_clients( + &remote_control_url, + &auth::RemoteControlAuth::capture(auth_manager.clone()).0, + RemoteControlClientsListParams { + environment_id: "env-123".to_string(), + ..Default::default() + }, + ) + .await + .expect_err("second unauthorized response should fail"); + server_task.await.expect("server task should finish"); + + assert_eq!(err.kind(), std::io::ErrorKind::PermissionDenied); +} + +#[tokio::test] +async fn revoke_remote_control_client_does_not_retry_forbidden() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let request = accept_http_request(&listener).await; + assert_eq!( + request.headers.get("authorization"), + Some(&"Bearer Access Token".to_string()) + ); + assert_eq!( + request.headers.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER), + vec!["account_id"] + ); + respond_with_status_and_headers( + request.stream, + "403 Forbidden", + &[("x-request-id", "request-123"), ("cf-ray", "ray-123")], + "forbidden", + ) + .await; + }); + + let err = revoke_remote_control_client( + &remote_control_url, + &auth::RemoteControlAuth::capture(remote_control_auth_manager()).0, + RemoteControlClientsRevokeParams { + environment_id: "env-123".to_string(), + client_id: "client-123".to_string(), + }, + ) + .await + .expect_err("forbidden revoke should fail"); + server_task.await.expect("server task should finish"); + + assert_eq!(err.kind(), std::io::ErrorKind::PermissionDenied); + assert_eq!( + err.to_string(), + format!( + "remote control client revoke failed at `{remote_control_url}wham/remote/control/environments/env-123/clients/client-123`: HTTP 403 Forbidden, request-id: request-123, cf-ray: ray-123, body: forbidden" + ) + ); +} + +#[tokio::test] +async fn list_remote_control_clients_preserves_decode_error_context() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let request = accept_http_request(&listener).await; + respond_with_status(request.stream, "200 OK", "{").await; + }); + + let err = list_remote_control_clients( + &remote_control_url, + &auth::RemoteControlAuth::capture(remote_control_auth_manager()).0, + RemoteControlClientsListParams { + environment_id: "env-123".to_string(), + ..Default::default() + }, + ) + .await + .expect_err("malformed client list should fail"); + server_task.await.expect("server task should finish"); + + assert!( + err.to_string().contains( + "failed to parse remote control client list response from `http://127.0.0.1:" + ) + ); + assert!(err.to_string().contains("HTTP 200 OK")); + assert!(err.to_string().contains("body: {")); + assert!(err.to_string().contains("decode error:")); +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/tests/pairing_tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/tests/pairing_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..b343242c8a158c8331c29b35bb786b948eae23f4 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/tests/pairing_tests.rs @@ -0,0 +1,1137 @@ +use super::super::protocol::RemoteControlPairingStatusRequest; +use super::super::protocol::StartRemoteControlPairingRequest; +use super::*; +use codex_login::AuthKeyringBackendKind; +use pretty_assertions::assert_eq; +use std::io; + +fn remote_control_enrollment( + remote_control_url: &str, + remote_control_token: &str, +) -> RemoteControlEnrollment { + RemoteControlEnrollment { + remote_control_target: normalize_remote_control_url(remote_control_url) + .expect("target should normalize"), + account_id: "account-id".to_string(), + environment_id: "environment-id".to_string(), + server_id: "server-id".to_string(), + server_name: "server-name".to_string(), + remote_control_token: Some(remote_control_token.to_string()), + expires_at: Some( + OffsetDateTime::from_unix_timestamp(33_336_362_096) + .expect("future timestamp should parse"), + ), + next_refresh_at: None, + } +} + +async fn auth_manager_with_replacement( + codex_home: &TempDir, + replacement_account_id: &str, +) -> Arc { + let mut stale_auth = remote_control_auth_dot_json(Some("account_id")); + stale_auth + .tokens + .as_mut() + .expect("stale auth should include tokens") + .access_token = "stale-token".to_string(); + save_auth( + codex_home.path(), + &stale_auth, + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("stale auth should save"); + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let mut replacement_auth = remote_control_auth_dot_json(Some(replacement_account_id)); + replacement_auth + .tokens + .as_mut() + .expect("replacement auth should include tokens") + .access_token = "fresh-token".to_string(); + save_auth( + codex_home.path(), + &replacement_auth, + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("replacement auth should save"); + auth_manager +} + +fn pairing_response_json(server_id: &str, environment_id: &str) -> serde_json::Value { + json!({ + "pairing_code": "pairing-code", + "manual_pairing_code": "ABCD-EFGH", + "server_id": server_id, + "environment_id": environment_id, + "expires_at": "3026-05-22T12:34:56Z", + }) +} + +fn pairing_response(environment_id: &str) -> RemoteControlPairingStartResponse { + RemoteControlPairingStartResponse { + pairing_code: "pairing-code".to_string(), + manual_pairing_code: Some("ABCD-EFGH".to_string()), + environment_id: environment_id.to_string(), + expires_at: 33_336_362_096, + } +} + +async fn pairing_error(status: &'static str, body: &'static str) -> (String, String) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let expected_pair_url = normalize_remote_control_url(&remote_control_url) + .expect("target should normalize") + .pair_url; + let server_task = tokio::spawn(async move { + let pairing_request = accept_http_request(&listener).await; + respond_with_status_and_headers( + pairing_request.stream, + status, + &[("x-request-id", "request-123"), ("cf-ray", "ray-123")], + body, + ) + .await; + }); + + let err = remote_control_enrollment(&remote_control_url, "remote-control-token") + .start_pairing(StartRemoteControlPairingRequest { manual_code: false }) + .await + .expect_err("pairing should fail"); + server_task.await.expect("server task should finish"); + (err.to_string(), expected_pair_url) +} + +async fn pairing_response_error(body: serde_json::Value) -> String { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let pairing_request = accept_http_request(&listener).await; + respond_with_json(pairing_request.stream, body).await; + }); + + let err = remote_control_enrollment(&remote_control_url, "remote-control-token") + .start_pairing(StartRemoteControlPairingRequest { manual_code: false }) + .await + .expect_err("pairing should fail"); + server_task.await.expect("server task should finish"); + err.to_string() +} + +async fn pairing_status_error(status: &'static str, body: &'static str) -> (io::Error, String) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let expected_status_url = normalize_remote_control_url(&remote_control_url) + .expect("target should normalize") + .pair_status_url; + let server_task = tokio::spawn(async move { + let status_request = accept_http_request(&listener).await; + respond_with_status_and_headers( + status_request.stream, + status, + &[("x-request-id", "request-123"), ("cf-ray", "ray-123")], + body, + ) + .await; + }); + + let err = remote_control_enrollment(&remote_control_url, "remote-control-token") + .pairing_status(RemoteControlPairingStatusRequest { + pairing_code: Some("pairing-code".to_string()), + manual_pairing_code: None, + }) + .await + .expect_err("pairing status should fail"); + server_task.await.expect("server task should finish"); + (err, expected_status_url) +} + +#[tokio::test] +async fn remote_control_handle_starts_pairing_before_websocket_connects() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + assert_eq!( + serde_json::from_str::(&refresh_request.body) + .expect("refresh request body should deserialize"), + json!({ + "server_id": "srv_e_test", + "installation_id": TEST_INSTALLATION_ID, + }) + ); + respond_with_json( + refresh_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let pairing_request = accept_http_request(&listener).await; + assert_eq!( + pairing_request.request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + assert_eq!( + pairing_request.headers.get("authorization"), + Some(&format!( + "Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}" + )) + ); + assert_eq!( + serde_json::from_str::(&pairing_request.body) + .expect("pairing request body should deserialize"), + json!({ "manual_code": true }) + ); + respond_with_json( + pairing_request.stream, + pairing_response_json("srv_e_test", "env_test"), + ) + .await; + }); + let remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + remote_handle + .current_enrollment + .lock() + .await + .as_mut() + .expect("current enrollment should exist") + .expires_at = Some(OffsetDateTime::now_utc() + time::Duration::seconds(29)); + + let response = remote_handle + .start_pairing( + RemoteControlPairingStartParams { manual_code: true }, + /*app_server_client_name*/ None, + ) + .await + .expect("pairing should use the current server before websocket connect"); + server_task.await.expect("server task should finish"); + + assert_eq!(response, pairing_response("env_test")); +} + +#[tokio::test] +async fn proactive_refresh_rate_limit_uses_valid_token_for_pairing() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_status_and_headers( + refresh_request.stream, + "429 Too Many Requests", + &[], + "rate limited", + ) + .await; + + let pairing_request = accept_http_request(&listener).await; + assert_eq!( + pairing_request.request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + assert_eq!( + pairing_request.headers.get("authorization"), + Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}")) + ); + respond_with_json( + pairing_request.stream, + pairing_response_json("srv_e_test", "env_test"), + ) + .await; + }); + let remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + remote_handle + .current_enrollment + .lock() + .await + .as_mut() + .expect("current enrollment should exist") + .expires_at = Some(OffsetDateTime::now_utc() + time::Duration::minutes(4)); + + let response = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect("valid token should allow pairing after proactive refresh failure"); + server_task.await.expect("server task should finish"); + + assert_eq!(response, pairing_response("env_test")); + assert!( + remote_handle + .current_enrollment + .snapshot() + .and_then(|enrollment| enrollment.next_refresh_at) + .is_some() + ); +} + +#[tokio::test] +async fn required_refresh_deadline_blocks_pairing_without_request() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_status_and_headers( + refresh_request.stream, + "502 Bad Gateway", + &[("retry-after", "120")], + "upstream unavailable", + ) + .await; + listener + }); + let remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + remote_handle + .current_enrollment + .lock() + .await + .as_mut() + .expect("current enrollment should exist") + .expires_at = Some(OffsetDateTime::now_utc() - time::Duration::seconds(1)); + + let refresh_err = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("required refresh failure should block pairing"); + let listener = server_task.await.expect("server task should finish"); + let next_refresh_at = remote_handle + .current_enrollment + .snapshot() + .and_then(|enrollment| enrollment.next_refresh_at) + .expect("required pairing refresh should preserve the retry deadline"); + let deferred_err = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("required refresh deadline should block pairing"); + + assert!(refresh_err.to_string().contains("HTTP 502 Bad Gateway")); + assert_eq!(deferred_err.kind(), io::ErrorKind::WouldBlock); + assert!( + deferred_err + .to_string() + .contains(&next_refresh_at.to_string()) + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("pairing should not issue a request before the refresh deadline"); +} + +#[tokio::test] +async fn remote_control_pairing_status_returns_pending() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let status_request = accept_http_request(&listener).await; + assert_eq!( + status_request.request_line, + "POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1" + ); + assert_eq!( + status_request.headers.get("authorization"), + Some(&"Bearer remote-control-token".to_string()) + ); + assert_eq!( + serde_json::from_str::(&status_request.body) + .expect("status request body should deserialize"), + json!({ "pairing_code": "pairing-code" }) + ); + respond_with_json(status_request.stream, json!({ "claimed": false })).await; + }); + + let response = remote_control_enrollment(&remote_control_url, "remote-control-token") + .pairing_status(RemoteControlPairingStatusRequest { + pairing_code: Some("pairing-code".to_string()), + manual_pairing_code: None, + }) + .await + .expect("pairing status should succeed"); + server_task.await.expect("server task should finish"); + + assert!(!response.claimed); +} + +#[tokio::test] +async fn remote_control_pairing_status_accepts_manual_pairing_code() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let status_request = accept_http_request(&listener).await; + assert_eq!( + status_request.request_line, + "POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1" + ); + assert_eq!( + serde_json::from_str::(&status_request.body) + .expect("status request body should deserialize"), + json!({ "manual_pairing_code": "ABCD-EFGH" }) + ); + respond_with_json(status_request.stream, json!({ "claimed": false })).await; + }); + + let response = remote_control_enrollment(&remote_control_url, "remote-control-token") + .pairing_status(RemoteControlPairingStatusRequest { + pairing_code: None, + manual_pairing_code: Some("ABCD-EFGH".to_string()), + }) + .await + .expect("pairing status should succeed"); + server_task.await.expect("server task should finish"); + + assert!(!response.claimed); +} + +#[tokio::test] +async fn remote_control_pairing_status_returns_claimed() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let status_request = accept_http_request(&listener).await; + assert_eq!( + status_request.request_line, + "POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1" + ); + respond_with_json(status_request.stream, json!({ "claimed": true })).await; + }); + + let response = remote_control_enrollment(&remote_control_url, "remote-control-token") + .pairing_status(RemoteControlPairingStatusRequest { + pairing_code: Some("pairing-code".to_string()), + manual_pairing_code: None, + }) + .await + .expect("pairing status should succeed"); + server_task.await.expect("server task should finish"); + + assert!(response.claimed); +} + +#[tokio::test] +async fn remote_control_handle_refreshes_after_pairing_status_auth_failure() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let stale_status_request = accept_http_request(&listener).await; + assert_eq!( + stale_status_request.request_line, + "POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1" + ); + assert_eq!( + stale_status_request.headers.get("authorization"), + Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}")) + ); + respond_with_status(stale_status_request.stream, "401 Unauthorized", "").await; + + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_json( + refresh_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let refreshed_status_request = accept_http_request(&listener).await; + assert_eq!( + refreshed_status_request.request_line, + "POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1" + ); + assert_eq!( + refreshed_status_request.headers.get("authorization"), + Some(&format!( + "Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}" + )) + ); + respond_with_json(refreshed_status_request.stream, json!({ "claimed": true })).await; + }); + let remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + + let response = remote_handle + .pairing_status(RemoteControlPairingStatusParams { + pairing_code: Some("pairing-code".to_string()), + manual_pairing_code: None, + }) + .await + .expect("pairing status should refresh after server token auth failure"); + server_task.await.expect("server task should finish"); + + assert!(response.claimed); +} + +#[tokio::test] +async fn remote_control_pairing_status_maps_user_actionable_backend_errors() { + for (status, expected_kind) in [ + ("403 Forbidden", io::ErrorKind::PermissionDenied), + ("404 Not Found", io::ErrorKind::InvalidInput), + ("410 Gone", io::ErrorKind::InvalidInput), + ] { + let (err, _expected_status_url) = pairing_status_error(status, "not available").await; + assert_eq!(err.kind(), expected_kind); + } +} + +#[tokio::test] +async fn remote_control_pairing_status_preserves_decode_error_context() { + let (err, expected_status_url) = pairing_status_error("200 OK", "{").await; + let err = err.to_string(); + + assert!(err.contains(&format!( + "failed to parse remote control pairing status response from `{expected_status_url}`: HTTP 200 OK" + ))); + assert!(err.contains("request-id: request-123")); + assert!(err.contains("cf-ray: ray-123")); + assert!(err.contains("body: {")); + assert!(err.contains("decode error:")); +} + +#[tokio::test] +async fn remote_control_handle_refreshes_after_pairing_auth_failure() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let stale_pairing_request = accept_http_request(&listener).await; + assert_eq!( + stale_pairing_request.request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + assert_eq!( + stale_pairing_request.headers.get("authorization"), + Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}")) + ); + respond_with_status(stale_pairing_request.stream, "401 Unauthorized", "").await; + + let refresh_request = accept_http_request(&listener).await; + assert_eq!( + refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + assert_eq!( + refresh_request.headers.get("authorization"), + Some(&"Bearer Access Token".to_string()) + ); + respond_with_json( + refresh_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let refreshed_pairing_request = accept_http_request(&listener).await; + assert_eq!( + refreshed_pairing_request.request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + assert_eq!( + refreshed_pairing_request.headers.get("authorization"), + Some(&format!( + "Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}" + )) + ); + respond_with_json( + refreshed_pairing_request.stream, + pairing_response_json("srv_e_test", "env_test"), + ) + .await; + }); + let remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + + let response = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect("pairing should refresh after server token auth failure"); + server_task.await.expect("server task should finish"); + + assert_eq!(response, pairing_response("env_test")); +} + +#[tokio::test] +async fn pairing_auth_failure_preserves_refresh_deadline() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let pairing_request = accept_http_request(&listener).await; + assert_eq!( + pairing_request.request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + respond_with_status(pairing_request.stream, "401 Unauthorized", "").await; + }); + let remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + let next_refresh_at = OffsetDateTime::now_utc() + time::Duration::minutes(2); + remote_handle + .current_enrollment + .lock() + .await + .as_mut() + .expect("current enrollment should exist") + .next_refresh_at = Some(next_refresh_at); + let mut expected_enrollment = remote_handle + .current_enrollment + .snapshot() + .expect("current enrollment should exist"); + expected_enrollment.clear_server_token(); + + let err = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("refresh deadline should throttle recovery after token rejection"); + server_task.await.expect("server task should finish"); + + assert_eq!(err.kind(), io::ErrorKind::WouldBlock); + assert_eq!( + remote_handle.current_enrollment.snapshot(), + Some(expected_enrollment) + ); +} + +#[tokio::test] +async fn remote_control_handle_recovers_auth_before_refreshing_pairing() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let stale_refresh_request = accept_http_request(&listener).await; + assert_eq!( + stale_refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + assert_eq!( + stale_refresh_request.headers.get("authorization"), + Some(&"Bearer stale-token".to_string()) + ); + respond_with_status(stale_refresh_request.stream, "401 Unauthorized", "").await; + + let recovered_refresh_request = accept_http_request(&listener).await; + assert_eq!( + recovered_refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + assert_eq!( + recovered_refresh_request.headers.get("authorization"), + Some(&"Bearer fresh-token".to_string()) + ); + respond_with_json( + recovered_refresh_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let pairing_request = accept_http_request(&listener).await; + assert_eq!( + pairing_request.request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + assert_eq!( + pairing_request.headers.get("authorization"), + Some(&format!( + "Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}" + )) + ); + respond_with_json( + pairing_request.stream, + pairing_response_json("srv_e_test", "env_test"), + ) + .await; + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let auth_manager = auth_manager_with_replacement(&codex_home, "account_id").await; + let remote_handle = + remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager); + remote_handle + .current_enrollment + .lock() + .await + .as_mut() + .expect("current enrollment should exist") + .expires_at = Some(OffsetDateTime::now_utc() + time::Duration::seconds(29)); + + let response = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect("pairing should refresh after auth recovery"); + server_task.await.expect("server task should finish"); + + assert_eq!(response, pairing_response("env_test")); +} + +#[tokio::test] +async fn pairing_publishes_refresh_deferral_after_auth_recovery() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let stale_refresh_request = accept_http_request(&listener).await; + assert_eq!( + stale_refresh_request.headers.get("authorization"), + Some(&"Bearer stale-token".to_string()) + ); + respond_with_status(stale_refresh_request.stream, "401 Unauthorized", "").await; + + let recovered_refresh_request = accept_http_request(&listener).await; + assert_eq!( + recovered_refresh_request.headers.get("authorization"), + Some(&"Bearer fresh-token".to_string()) + ); + let response_started_at = OffsetDateTime::now_utc(); + respond_with_status_and_headers( + recovered_refresh_request.stream, + "502 Bad Gateway", + &[("retry-after", "120")], + "upstream unavailable", + ) + .await; + response_started_at + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let auth_manager = auth_manager_with_replacement(&codex_home, "account_id").await; + let remote_handle = + remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager); + remote_handle + .current_enrollment + .lock() + .await + .as_mut() + .expect("current enrollment should exist") + .expires_at = Some(OffsetDateTime::now_utc() - time::Duration::seconds(1)); + + let refresh_err = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("required refresh should remain strict after auth recovery"); + let refresh_completed_at = OffsetDateTime::now_utc(); + let deferred_err = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("published deadline should throttle the next pairing refresh"); + let response_started_at = server_task.await.expect("server task should finish"); + + assert!(refresh_err.to_string().contains("HTTP 502 Bad Gateway")); + assert_eq!(deferred_err.kind(), io::ErrorKind::WouldBlock); + let next_refresh_at = remote_handle + .current_enrollment + .snapshot() + .and_then(|enrollment| enrollment.next_refresh_at) + .expect("required refresh failure should publish its retry deadline"); + assert!( + (response_started_at + time::Duration::seconds(120) + ..=refresh_completed_at + time::Duration::seconds(150)) + .contains(&next_refresh_at) + ); +} + +#[tokio::test] +async fn pairing_auth_recovery_failure_publishes_cleared_server_token() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let server_task = tokio::spawn(async move { + let stale_refresh_request = accept_http_request(&listener).await; + assert_eq!( + stale_refresh_request.request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + assert_eq!( + stale_refresh_request.headers.get("authorization"), + Some(&"Bearer stale-token".to_string()) + ); + respond_with_status(stale_refresh_request.stream, "401 Unauthorized", "").await; + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let auth_manager = auth_manager_with_replacement(&codex_home, "different_account_id").await; + let remote_handle = + remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager); + remote_handle + .current_enrollment + .lock() + .await + .as_mut() + .expect("current enrollment should exist") + .expires_at = Some(OffsetDateTime::now_utc() + time::Duration::seconds(29)); + let mut expected_enrollment = remote_handle + .current_enrollment + .snapshot() + .expect("current enrollment should exist"); + expected_enrollment.clear_server_token(); + + let err = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("pairing should fail after auth changes account"); + server_task.await.expect("server task should finish"); + + assert_eq!(err.kind(), io::ErrorKind::PermissionDenied); + assert_eq!( + remote_handle.current_enrollment.snapshot(), + Some(expected_enrollment) + ); +} + +#[tokio::test] +async fn start_remote_control_pairing_preserves_backend_error_context() { + let (err, expected_pair_url) = + pairing_error("503 Service Unavailable", "pairing unavailable").await; + + assert_eq!( + err, + format!( + "remote control pairing failed at `{expected_pair_url}`: HTTP 503 Service Unavailable, request-id: request-123, cf-ray: ray-123, body: pairing unavailable" + ) + ); +} + +#[tokio::test] +async fn start_remote_control_pairing_preserves_decode_error_context() { + let (err, expected_pair_url) = pairing_error("200 OK", "{").await; + assert!(err.contains(&format!( + "failed to parse remote control pairing response from `{expected_pair_url}`: HTTP 200 OK" + ))); + assert!(err.contains("request-id: request-123")); + assert!(err.contains("cf-ray: ray-123")); + assert!(err.contains("body: {")); + assert!(err.contains("decode error:")); +} + +#[tokio::test] +async fn start_remote_control_pairing_rejects_mismatched_backend_enrollment() { + assert_eq!( + pairing_response_error(json!({ + "pairing_code": "pairing-code", + "manual_pairing_code": "ABCD-EFGH", + "server_id": "other-server-id", + "environment_id": "other-environment-id", + "expires_at": "3026-05-22T12:34:56Z", + })) + .await, + "remote control pairing returned mismatched enrollment: expected server_id=server-id, environment_id=environment-id; got server_id=other-server-id, environment_id=other-environment-id" + ); +} + +#[tokio::test] +async fn start_remote_control_pairing_preserves_expiry_parse_error_context() { + let err = pairing_response_error(json!({ + "pairing_code": "pairing-code", + "manual_pairing_code": "ABCD-EFGH", + "server_id": "server-id", + "environment_id": "environment-id", + "expires_at": "not-a-timestamp", + })) + .await; + + assert!(err.contains("failed to parse remote control pairing response")); + assert!(err.contains("HTTP 200 OK")); + assert!(err.contains("request-id: ")); + assert!(err.contains("cf-ray: ")); + assert!(err.contains("\"expires_at\":\"not-a-timestamp\"")); + assert!(err.contains("expires_at parse error:")); +} + +#[tokio::test] +async fn remote_control_handle_disable_keeps_current_enrollment() { + let remote_handle = remote_control_handle_with_current_enrollment( + TEST_REMOTE_CONTROL_URL, + remote_control_auth_manager(), + ); + + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Disabled); + assert!( + remote_handle.current_enrollment.lock().await.is_some(), + "disabled remote control should keep the selected pairing server" + ); +} + +#[tokio::test] +async fn remote_control_handle_reenrolls_after_stale_pairing_enrollment() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager_with_home(&codex_home), + ); + remote_handle.state_db = Some(state_db.clone()); + let stale_enrollment = remote_handle + .current_enrollment + .lock() + .await + .clone() + .expect("current enrollment should exist"); + let remote_control_target = stale_enrollment.remote_control_target.clone(); + let refreshed_enrollment = RemoteControlEnrollment { + remote_control_target: remote_control_target.clone(), + account_id: "account_id".to_string(), + environment_id: "env_refreshed".to_string(), + server_id: "srv_e_refreshed".to_string(), + server_name: test_server_name(), + remote_control_token: None, + expires_at: None, + next_refresh_at: None, + }; + update_persisted_remote_control_enrollment( + Some(state_db.as_ref()), + &remote_control_target, + "account_id", + /*app_server_client_name*/ None, + Some(&stale_enrollment), + /*remote_control_enabled*/ Some(true), + ) + .await + .expect("stale enrollment should save"); + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + }); + let server_refreshed_enrollment = refreshed_enrollment.clone(); + let server_task = tokio::spawn(async move { + let stale_pairing_request = accept_http_request(&listener).await; + assert_eq!( + stale_pairing_request.request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + assert_eq!( + stale_pairing_request.headers.get("authorization"), + Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}")) + ); + respond_with_status(stale_pairing_request.stream, "404 Not Found", "").await; + + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_json( + enroll_request.stream, + remote_control_server_token_response( + &server_refreshed_enrollment.server_id, + &server_refreshed_enrollment.environment_id, + TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + + let refreshed_pairing_request = accept_http_request(&listener).await; + assert_eq!( + refreshed_pairing_request.request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + assert_eq!( + refreshed_pairing_request.headers.get("authorization"), + Some(&format!( + "Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}" + )) + ); + respond_with_json( + refreshed_pairing_request.stream, + pairing_response_json( + &server_refreshed_enrollment.server_id, + &server_refreshed_enrollment.environment_id, + ), + ) + .await; + }); + let response = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect("pairing should re-enroll after stale enrollment"); + server_task.await.expect("server task should finish"); + + assert_eq!(response, pairing_response("env_refreshed")); + assert_eq!( + state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + "account_id", + /*app_server_client_name*/ None, + ) + .await + .expect("refreshed enrollment should load") + .expect("refreshed enrollment should exist") + .remote_control_enabled, + Some(true) + ); +} + +#[tokio::test] +async fn remote_control_handle_discards_pairing_response_after_auth_change() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json(Some("account_id")), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("initial auth should save"); + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let remote_handle = + remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager.clone()); + let pairing_task = tokio::spawn({ + let remote_handle = remote_handle.clone(); + async move { + remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + } + }); + + let pairing_request = accept_http_request(&listener).await; + save_auth( + codex_home.path(), + &remote_control_auth_dot_json(Some("next_account_id")), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("next auth should save"); + auth_manager.reload().await; + respond_with_json( + pairing_request.stream, + json!({ + "pairing_code": "stale-pairing-code", + "manual_pairing_code": "ABCD-EFGH", + "server_id": "srv_e_test", + "environment_id": "env_test", + "expires_at": "3026-05-22T12:34:56Z", + }), + ) + .await; + + assert_eq!( + pairing_task + .await + .expect("pairing task should join") + .expect_err("stale pairing response should be discarded") + .to_string(), + "remote control pairing is unavailable until enrollment completes" + ); +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/tests/retry_tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/tests/retry_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..8871f98c6bc0d5b816776fa3d6bbd0e0df6d5c39 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/tests/retry_tests.rs @@ -0,0 +1,203 @@ +//! Exercises overload retry deadlines over HTTP and WebSocket connections. +//! Auth reloads must respect the deadline; shutdown must remain prompt. + +use super::*; +use pretty_assertions::assert_eq; + +#[tokio::test] +async fn rate_limited_enrollment_respects_retry_after() { + assert_overload_retry_after("429 Too Many Requests", /*reject_enrollment*/ true).await; +} + +#[tokio::test] +async fn rate_limited_websocket_respects_retry_after() { + assert_overload_retry_after("429 Too Many Requests", /*reject_enrollment*/ false).await; +} + +#[tokio::test] +async fn enrollment_resumes_after_retry_after() { + assert_overload_retry_after("503 Service Unavailable", /*reject_enrollment*/ true).await; +} + +#[tokio::test] +async fn websocket_resumes_after_retry_after() { + assert_overload_retry_after("503 Service Unavailable", /*reject_enrollment*/ false).await; +} + +async fn assert_overload_retry_after(status: &str, reject_enrollment: bool) { + let verify_recovery = status == "503 Service Unavailable"; + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let codex_home = TempDir::new().expect("temp dir should create"); + let (transport_event_tx, _transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let auth_manager = remote_control_auth_manager_with_home(&codex_home); + let mut initial_auth = remote_control_auth_dot_json(Some("account_id")); + initial_auth + .tokens + .as_mut() + .expect("fixture should contain tokens") + .access_token = "Initial Access Token".to_string(); + save_auth( + codex_home.path(), + &initial_auth, + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("initial credentials should save"); + auth_manager.reload().await; + let (remote_task, remote_handle) = start_remote_control( + RemoteControlStartConfig { + remote_control_url: remote_control_url_for_listener(&listener), + installation_id: TEST_INSTALLATION_ID.to_string(), + policy: RemoteControlPolicy::Allowed, + }, + Some(remote_control_state_runtime(&codex_home).await), + auth_manager.clone(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + RemoteControlStartupMode::EnabledEphemeral, + ) + .await + .expect("remote control should start"); + let mut status_rx = remote_handle.status_receiver(); + let mut rejected_request = accept_http_request(&listener).await; + assert_eq!( + rejected_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + if !reject_enrollment { + respond_with_json( + rejected_request.stream, + remote_control_server_token_response( + "srv_e_test", + "env_test", + TEST_REMOTE_CONTROL_SERVER_TOKEN, + ), + ) + .await; + rejected_request = accept_http_request(&listener).await; + assert_eq!( + rejected_request.request_line, + "GET /backend-api/wham/remote/control/server HTTP/1.1" + ); + } + let response_started_at = std::time::Instant::now(); + respond_with_status_and_headers( + rejected_request.stream, + status, + &[("Retry-After", if verify_recovery { "3" } else { "120" })], + "overloaded", + ) + .await; + timeout( + Duration::from_secs(5), + status_rx.wait_for(|status| status.status == RemoteControlConnectionStatus::Errored), + ) + .await + .expect("the overload response should be processed") + .expect("the status channel should remain open"); + + let auth_changes = auth_manager.auth_change_receiver(); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json(Some("account_id")), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("updated credentials should save"); + auth_manager.reload().await; + assert!( + auth_changes + .has_changed() + .expect("auth watch should remain open") + ); + + if !verify_recovery { + let pairing_error = timeout( + Duration::from_secs(1), + remote_handle.start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ), + ) + .await + .expect("pairing should report the active retry delay promptly") + .expect_err("pairing must respect the shared server retry delay"); + let retry_delay = server_api::remote_control_retry_delay(&pairing_error) + .expect("pairing should preserve the overload retry deadline"); + assert!( + retry_delay >= Duration::from_secs(120).saturating_sub(response_started_at.elapsed()), + "pairing must preserve the full remaining Retry-After interval" + ); + assert!( + timeout(Duration::from_secs(2), listener.accept()) + .await + .is_err(), + "no enrollment, refresh or handshake should retry before Retry-After ({status})" + ); + remote_handle.disable_ephemeral().await; + assert!( + timeout(Duration::from_secs(2), listener.accept()) + .await + .is_err(), + "disabling remote control must stop connection attempts" + ); + remote_handle + .enable_ephemeral() + .expect("remote control should enable again"); + assert!( + timeout(Duration::from_secs(2), listener.accept()) + .await + .is_err(), + "re-enabling remote control must preserve the server retry delay" + ); + } + if !reject_enrollment { + assert_eq!( + remote_handle + .inner + .session() + .current_enrollment + .snapshot() + .and_then(|enrollment| enrollment.remote_control_token), + Some(TEST_REMOTE_CONTROL_SERVER_TOKEN.to_string()), + "overload must preserve the enrolled token" + ); + } + if verify_recovery { + remote_handle.disable_ephemeral().await; + remote_handle + .enable_ephemeral() + .expect("remote control should enable again"); + let (stream, _) = timeout(Duration::from_secs(45), listener.accept()) + .await + .expect("retry should resume after the server delay and bounded jitter") + .expect("retry connection should succeed"); + assert!( + response_started_at.elapsed() >= Duration::from_secs(3), + "the retry must wait for the full Retry-After interval" + ); + let mut reader = BufReader::new(stream); + let mut request_line = String::new(); + timeout(Duration::from_secs(5), reader.read_line(&mut request_line)) + .await + .expect("retry should send the request promptly") + .expect("retry request should read"); + assert_eq!( + request_line.trim_end(), + if reject_enrollment { + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + } else { + "GET /backend-api/wham/remote/control/server HTTP/1.1" + } + ); + } + shutdown_token.cancel(); + timeout(Duration::from_secs(1), remote_task) + .await + .expect("shutdown must interrupt the server retry delay") + .expect("remote task should finish"); +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/websocket.rs b/codex-rs/app-server-transport/src/transport/remote_control/websocket.rs new file mode 100644 index 0000000000000000000000000000000000000000..d50d31771d9a9936f0455752af3c51a0fed59383 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/websocket.rs @@ -0,0 +1,3535 @@ +use super::CurrentRemoteControlEnrollment; +use super::RemoteControlEnrollmentSelection; +use super::RemoteControlPairingPersistenceKey; +use super::auth::RemoteControlAuth; +use super::auth::RemoteControlRecovery; +use super::desired_state::RemoteControlDesiredState; +use super::desired_state::desired_state_from_persisted_enrollment; +use super::persistence::RemoteControlPersistence; +use super::protocol::ClientEnvelope; +use super::protocol::ClientEvent; +use super::protocol::ClientId; +use super::protocol::RemoteControlTarget; +use super::protocol::ServerEnvelope; +use super::protocol::StreamId; +use super::remote_control_status_with_connection_status; +use super::same_remote_control_enrollment; +use super::segment::ClientSegmentObservation; +use super::segment::ClientSegmentReassembler; +use super::segment::REMOTE_CONTROL_SEGMENT_MAX_BYTES; +use super::segment::split_server_envelope_for_transport; +use crate::transport::TransportEvent; +use crate::transport::remote_control::auth::RemoteControlConnectionAuth; +use crate::transport::remote_control::auth::load_remote_control_auth; +use crate::transport::remote_control::auth::recover_remote_control_auth; +use crate::transport::remote_control::client_tracker::ClientTracker; +use crate::transport::remote_control::client_tracker::REMOTE_CONTROL_IDLE_SWEEP_INTERVAL; +use crate::transport::remote_control::enroll::RemoteControlEnrollment; +use crate::transport::remote_control::enroll::format_headers; +use crate::transport::remote_control::enroll::load_persisted_remote_control_enrollment; +use crate::transport::remote_control::enroll::preview_remote_control_response_body; +use crate::transport::remote_control::host_device::REMOTE_CONTROL_HOST_DEVICE_KIND_HEADER; +use crate::transport::remote_control::host_device::host_device_kind; +use crate::transport::remote_control::server_api::RemoteControlServerRequestError; +use crate::transport::remote_control::server_api::enroll_remote_control_server; +use crate::transport::remote_control::server_api::refresh_remote_control_server; +use crate::transport::remote_control::server_api::remote_control_retry_delay; +use crate::transport::remote_control::server_api::retry_after_with_jitter; +use axum::http::HeaderValue; +use base64::Engine; +use codex_app_server_protocol::RemoteControlConnectionStatus; +use codex_app_server_protocol::RemoteControlStatusChangedNotification; +use codex_core::util::backoff; +use codex_state::StateRuntime; +use codex_utils_rustls_provider::ensure_rustls_crypto_provider; +use futures::SinkExt; +use futures::StreamExt; +use futures::stream::SplitSink; +use futures::stream::SplitStream; +use std::collections::HashMap; +use std::collections::VecDeque; +use std::io; +use std::io::ErrorKind; +use std::sync::Arc; +use tokio::net::TcpStream; +use tokio::sync::Mutex; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::sync::watch; +use tokio::time::MissedTickBehavior; +use tokio_tungstenite::MaybeTlsStream; +use tokio_tungstenite::WebSocketStream; +use tokio_tungstenite::connect_async; +use tokio_tungstenite::tungstenite; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_util::sync::CancellationToken; + +#[cfg(test)] +use super::RemoteControlEnrollmentState; +use tracing::error; +use tracing::info; +use tracing::warn; + +pub(super) const REMOTE_CONTROL_PROTOCOL_VERSION: &str = "3"; +pub(super) const REMOTE_CONTROL_INSTALLATION_ID_HEADER: &str = "x-codex-installation-id"; +const REMOTE_CONTROL_SUBSCRIBE_CURSOR_HEADER: &str = "x-codex-subscribe-cursor"; +const REMOTE_CONTROL_WEBSOCKET_PING_INTERVAL: std::time::Duration = + std::time::Duration::from_secs(10); +const REMOTE_CONTROL_WEBSOCKET_PONG_TIMEOUT: std::time::Duration = + std::time::Duration::from_secs(60); +const REMOTE_CONTROL_ACCOUNT_ID_RETRY_INTERVAL: std::time::Duration = + std::time::Duration::from_secs(1); +const REMOTE_CONTROL_RECONNECT_BACKOFF_CAP: std::time::Duration = + std::time::Duration::from_secs(30); +const REMOTE_CONTROL_WEBSOCKET_CONNECT_TIMEOUT: std::time::Duration = + std::time::Duration::from_secs(30); +const REMOTE_CONTROL_CONNECTION_SHUTDOWN_TIMEOUT: std::time::Duration = + std::time::Duration::from_secs(5); +const REMOTE_APP_SERVER_NOT_FOUND_DETAIL: &str = "Remote app server not found"; + +struct BoundedOutboundBuffer { + buffer_by_stream: HashMap<(ClientId, StreamId), VecDeque>, + used_tx: watch::Sender, +} + +impl BoundedOutboundBuffer { + fn new() -> (Self, watch::Receiver) { + let (used_tx, used_rx) = watch::channel(0); + let buffer = Self { + buffer_by_stream: HashMap::new(), + used_tx, + }; + (buffer, used_rx) + } + + fn insert(&mut self, server_envelope: &ServerEnvelope) { + self.buffer_by_stream + .entry(( + server_envelope.client_id.clone(), + server_envelope.stream_id.clone(), + )) + .or_default() + .push_back(server_envelope.clone()); + self.used_tx.send_modify(|used| *used += 1); + } + + fn ack( + &mut self, + client_id: &ClientId, + stream_id: &StreamId, + acked_seq_id: u64, + acked_segment_id: Option, + ) { + let key = (client_id.clone(), stream_id.clone()); + let Some(buffer) = self.buffer_by_stream.get_mut(&key) else { + return; + }; + let acked_cursor = (acked_seq_id, acked_segment_id.unwrap_or(usize::MAX)); + buffer.retain(|server_envelope| { + let envelope_cursor = ( + server_envelope.seq_id, + server_envelope.event.segment_id().unwrap_or_default(), + ); + let is_acked = envelope_cursor <= acked_cursor; + if is_acked { + self.used_tx.send_modify(|used| *used -= 1); + } + !is_acked + }); + if buffer.is_empty() { + self.buffer_by_stream.remove(&key); + } + } + + fn server_envelopes(&self) -> impl Iterator { + self.buffer_by_stream + .values() + .flat_map(|buffer| buffer.iter()) + } +} + +struct WebsocketState { + outbound_buffer: BoundedOutboundBuffer, + subscribe_cursor: Option, + next_seq_id_by_stream: HashMap<(ClientId, StreamId), u64>, + last_completed_client_chunk_seq_id_by_stream: HashMap<(ClientId, Option), u64>, + client_segment_reassembler: ClientSegmentReassembler, +} + +impl WebsocketState { + fn observe_client_message( + &mut self, + client_envelope: ClientEnvelope, + wire_size_bytes: usize, + ) -> ClientSegmentObservation { + let client_message_key = Self::client_message_key(&client_envelope); + if let Some((key, seq_id)) = client_message_key.as_ref() + && self + .last_completed_client_chunk_seq_id_by_stream + .get(key) + .is_some_and(|last_seq_id| last_seq_id >= seq_id) + { + return ClientSegmentObservation::Dropped; + } + if let ( + Some((_, seq_id)), + Some(stream_id), + ClientEvent::ClientMessageChunk { segment_id, .. }, + ) = ( + client_message_key.as_ref(), + client_envelope.stream_id.as_ref(), + &client_envelope.event, + ) && self.client_segment_reassembler.should_ignore_chunk( + &client_envelope.client_id, + stream_id, + *seq_id, + *segment_id, + ) { + return ClientSegmentObservation::Dropped; + } + if client_message_key.is_some() && wire_size_bytes > REMOTE_CONTROL_SEGMENT_MAX_BYTES { + warn!( + client_id = client_envelope.client_id.0.as_str(), + "dropping oversized segmented remote-control client envelope" + ); + if let Some(stream_id) = client_envelope.stream_id.as_ref() { + self.client_segment_reassembler + .invalidate_stream(&client_envelope.client_id, stream_id); + } + return ClientSegmentObservation::Dropped; + } + + self.client_segment_reassembler.observe(client_envelope) + } + + fn record_client_message_delivery( + &mut self, + client_envelope: &ClientEnvelope, + client_message_key: Option<((ClientId, Option), u64)>, + ) { + if let Some(cursor) = client_envelope.cursor.as_deref() { + self.subscribe_cursor = Some(cursor.to_string()); + } + if let Some((key, seq_id)) = client_message_key { + self.last_completed_client_chunk_seq_id_by_stream + .insert(key, seq_id); + } + if let ClientEvent::Ack { segment_id } = &client_envelope.event + && let Some(acked_seq_id) = client_envelope.seq_id + && let Some(stream_id) = client_envelope.stream_id.as_ref() + { + self.outbound_buffer.ack( + &client_envelope.client_id, + stream_id, + acked_seq_id, + *segment_id, + ); + } + } + + fn invalidate_client_message_stream(&mut self, client_id: &ClientId, stream_id: &StreamId) { + self.last_completed_client_chunk_seq_id_by_stream + .remove(&(client_id.clone(), Some(stream_id.clone()))); + } + + fn invalidate_client_message_client(&mut self, client_id: &ClientId) { + self.last_completed_client_chunk_seq_id_by_stream + .retain(|(cursor_client_id, _), _| cursor_client_id != client_id); + } + + fn client_message_key( + client_envelope: &ClientEnvelope, + ) -> Option<((ClientId, Option), u64)> { + let seq_id = match (&client_envelope.event, client_envelope.seq_id) { + (ClientEvent::ClientMessageChunk { .. }, Some(seq_id)) => seq_id, + _ => return None, + }; + Some(( + ( + client_envelope.client_id.clone(), + client_envelope.stream_id.clone(), + ), + seq_id, + )) + } +} + +pub(super) struct RemoteControlWebsocket { + remote_control_url: String, + installation_id: String, + server_name: String, + remote_control_target: Option, + state_db: Option>, + auth_manager: RemoteControlAuth, + status_publisher: RemoteControlStatusPublisher, + shutdown_token: CancellationToken, + reconnect_attempt: u64, + auth_recovery: RemoteControlRecovery, + auth_change_rx: watch::Receiver, + current_enrollment: CurrentRemoteControlEnrollment, + pairing_persistence_key: RemoteControlPairingPersistenceKey, + client_tracker: Arc>, + state: Arc>, + server_event_rx: Arc>>, + used_rx: watch::Receiver, + desired_state_tx: Arc>, + desired_state_rx: watch::Receiver, + persistence: RemoteControlPersistence, +} + +pub(super) struct RemoteControlWebsocketConfig { + pub(crate) remote_control_url: String, + pub(crate) installation_id: String, + pub(crate) remote_control_target: Option, + pub(crate) server_name: String, +} + +pub(super) struct RemoteControlAuthContext<'a> { + auth_manager: &'a RemoteControlAuth, + auth_recovery: &'a mut RemoteControlRecovery, + auth_change_rx: &'a mut watch::Receiver, +} + +struct RemoteControlEnrollmentAuthContext<'a, 'b> { + auth: &'a RemoteControlConnectionAuth, + recovery: &'a mut RemoteControlAuthContext<'b>, +} + +enum ConnectOutcome { + Connected(Box>>), + Disabled, + Shutdown, +} + +#[derive(Debug, Clone, Copy)] +enum ConnectionEndReason { + AuthOwnerChanged, + Shutdown, + Disabled, + EnabledWatchClosed, + ConnectionWorkerStopped, +} + +pub(super) struct RemoteControlChannels { + pub(super) transport_event_tx: mpsc::Sender, + pub(super) status_publisher: RemoteControlStatusPublisher, + pub(super) current_enrollment: CurrentRemoteControlEnrollment, + pub(super) pairing_persistence_key: RemoteControlPairingPersistenceKey, + pub(super) persistence: RemoteControlPersistence, +} + +#[derive(Clone)] +pub(super) struct RemoteControlStatusPublisher { + tx: watch::Sender, +} + +impl RemoteControlStatusPublisher { + pub(super) fn new(tx: watch::Sender) -> Self { + Self { tx } + } + + fn status(&self) -> RemoteControlStatusChangedNotification { + self.tx.borrow().clone() + } + + fn publish_status(&self, connection_status: RemoteControlConnectionStatus) { + let mut status_change = None; + self.tx.send_if_modified(|status| { + let next_status = + remote_control_status_with_connection_status(status, connection_status); + if *status == next_status { + return false; + } + + status_change = Some((status.clone(), next_status.clone())); + *status = next_status; + true + }); + if let Some((previous_status, next_status)) = status_change { + info!( + previous_status = ?previous_status.status, + next_status = ?next_status.status, + previous_environment_id = ?previous_status.environment_id, + next_environment_id = ?next_status.environment_id, + installation_id = %next_status.installation_id, + server_name = %next_status.server_name, + "remote control websocket status changed" + ); + } + } + + pub(super) fn publish_environment_id(&self, environment_id: Option) { + let mut status_change = None; + self.tx.send_if_modified(|status| { + if status.status == RemoteControlConnectionStatus::Disabled { + return false; + } + let next_status = RemoteControlStatusChangedNotification { + status: status.status, + server_name: status.server_name.clone(), + installation_id: status.installation_id.clone(), + environment_id, + }; + if *status == next_status { + return false; + } + + status_change = Some((status.clone(), next_status.clone())); + *status = next_status; + true + }); + if let Some((previous_status, next_status)) = status_change { + info!( + status = ?next_status.status, + previous_environment_id = ?previous_status.environment_id, + next_environment_id = ?next_status.environment_id, + installation_id = %next_status.installation_id, + server_name = %next_status.server_name, + "remote control websocket environment changed" + ); + } + } +} + +#[derive(Clone, Copy)] +pub(super) struct RemoteControlConnectOptions<'a> { + installation_id: &'a str, + server_name: &'a str, + subscribe_cursor: Option<&'a str>, + app_server_client_name: Option<&'a str>, + desired_state_tx: &'a watch::Sender, + persistence: &'a RemoteControlPersistence, +} + +impl RemoteControlWebsocket { + pub(super) fn new( + config: RemoteControlWebsocketConfig, + state_db: Option>, + auth_manager: RemoteControlAuth, + channels: RemoteControlChannels, + shutdown_token: CancellationToken, + desired_state_tx: Arc>, + ) -> Self { + let shutdown_token = shutdown_token.child_token(); + let (server_event_tx, server_event_rx) = mpsc::channel(super::CHANNEL_CAPACITY); + let client_tracker = ClientTracker::new( + server_event_tx, + channels.transport_event_tx, + &shutdown_token, + ); + let (outbound_buffer, used_rx) = BoundedOutboundBuffer::new(); + let auth_recovery = auth_manager.unauthorized_recovery(); + let auth_change_rx = auth_manager.auth_change_receiver(); + + let desired_state_rx = desired_state_tx.subscribe(); + Self { + remote_control_url: config.remote_control_url, + installation_id: config.installation_id, + server_name: config.server_name, + remote_control_target: config.remote_control_target, + state_db, + auth_manager, + status_publisher: channels.status_publisher, + shutdown_token, + reconnect_attempt: 0, + auth_recovery, + auth_change_rx, + current_enrollment: channels.current_enrollment, + pairing_persistence_key: channels.pairing_persistence_key, + client_tracker: Arc::new(Mutex::new(client_tracker)), + state: Arc::new(Mutex::new(WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + })), + server_event_rx: Arc::new(Mutex::new(server_event_rx)), + used_rx, + desired_state_tx, + desired_state_rx, + persistence: channels.persistence, + } + } + + #[expect( + clippy::await_holding_invalid_type, + reason = "remote-control client shutdown must serialize tracker state" + )] + pub(super) async fn run( + mut self, + app_server_client_name_rx: Option>, + ) { + let auth_owner = self.auth_manager.owner.clone(); + info!( + remote_control_url = %self.remote_control_url, + installation_id = %self.installation_id, + server_name = %self.server_name, + initial_desired_state = ?*self.desired_state_rx.borrow(), + "app-server remote control websocket loop started" + ); + let app_server_client_name = match self + .wait_for_app_server_client_name(app_server_client_name_rx) + .await + { + Ok(app_server_client_name) => app_server_client_name, + Err(_) => { + warn!( + remote_control_url = %self.remote_control_url, + installation_id = %self.installation_id, + server_name = %self.server_name, + shutdown_requested = self.shutdown_token.is_cancelled(), + "app-server remote control websocket loop stopped before client name was ready" + ); + self.client_tracker.lock().await.shutdown().await; + return; + } + }; + self.pairing_persistence_key + .send_replace(app_server_client_name.clone()); + loop { + let status = self.status_publisher.status(); + info!( + remote_control_url = %self.remote_control_url, + installation_id = %self.installation_id, + server_name = %self.server_name, + reconnect_attempt = self.reconnect_attempt.saturating_add(1), + current_status = ?status.status, + environment_id = ?status.environment_id, + "starting app-server remote control websocket connection cycle" + ); + let shutdown_token = self.shutdown_token.child_token(); + let connect_outcome = tokio::select! { + biased; + _ = shutdown_token.cancelled() => break, + _ = auth_owner.invalidated() => break, + outcome = async { + if matches!(*self.desired_state_rx.borrow(), RemoteControlDesiredState::Unknown) + && !self.resolve_unknown_desired_state(app_server_client_name.as_deref()).await + { + return ConnectOutcome::Shutdown; + } + if !self.wait_until_enabled().await { + return ConnectOutcome::Shutdown; + } + self.connect(&shutdown_token, app_server_client_name.as_deref()).await + } => outcome, + }; + let websocket_connection = match connect_outcome { + ConnectOutcome::Connected(websocket_connection) => *websocket_connection, + ConnectOutcome::Disabled => { + self.status_publisher + .publish_status(RemoteControlConnectionStatus::Disabled); + continue; + } + ConnectOutcome::Shutdown => break, + }; + + let connection_end_reason = self + .run_connection(websocket_connection, shutdown_token) + .await; + let status = self.status_publisher.status(); + info!( + remote_control_url = %self.remote_control_url, + installation_id = %self.installation_id, + server_name = %self.server_name, + connection_end_reason = ?connection_end_reason, + current_status = ?status.status, + environment_id = ?status.environment_id, + desired_state = ?*self.desired_state_rx.borrow(), + "app-server remote control websocket connection cycle ended" + ); + } + + self.client_tracker.lock().await.shutdown().await; + info!( + remote_control_url = %self.remote_control_url, + installation_id = %self.installation_id, + server_name = %self.server_name, + shutdown_requested = self.shutdown_token.is_cancelled(), + "app-server remote control websocket loop exited" + ); + } + + async fn wait_for_app_server_client_name( + &self, + app_server_client_name_rx: Option>, + ) -> Result, ()> { + match app_server_client_name_rx { + Some(app_server_client_name_rx) => { + tokio::select! { + _ = self.shutdown_token.cancelled() => Err(()), + app_server_client_name = app_server_client_name_rx => match app_server_client_name { + Ok(app_server_client_name) => Ok(Some(app_server_client_name)), + Err(_) => Err(()), + }, + } + } + None => Ok(None), + } + } + + pub(super) async fn resolve_unknown_desired_state( + &mut self, + app_server_client_name: Option<&str>, + ) -> bool { + let remote_control_target = match super::protocol::normalize_remote_control_url( + &self.remote_control_url, + ) { + Ok(remote_control_target) => remote_control_target, + Err(err) => { + warn!( + "remote control preference cannot be resolved because the URL is invalid: {err}" + ); + self.transition_unknown_to(RemoteControlDesiredState::Disabled); + return true; + } + }; + self.remote_control_target = Some(remote_control_target.clone()); + let Some(state_db) = self.state_db.clone() else { + self.transition_unknown_to(RemoteControlDesiredState::Disabled); + return true; + }; + + loop { + if !matches!( + *self.desired_state_rx.borrow(), + RemoteControlDesiredState::Unknown + ) { + return true; + } + let auth = match load_remote_control_auth(&self.auth_manager).await { + Ok(auth) => auth, + Err(err) => { + info!( + error = %err, + "waiting to resolve remote control preference until authentication is available" + ); + if !self.wait_for_preference_resolution_retry().await { + return false; + } + continue; + } + }; + let _persistence = + match super::persistence::read_lock(&self.auth_manager, &self.persistence).await { + Ok(permit) => permit, + Err(_) => return false, + }; + let enrollment = state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + &auth.account_id, + app_server_client_name, + ) + .await; + drop(_persistence); + let enrollment = match enrollment { + Ok(enrollment) => enrollment, + Err(err) => { + warn!( + error = %err, + "failed to resolve persisted remote control preference; retrying" + ); + if !self.wait_for_preference_resolution_retry().await { + return false; + } + continue; + } + }; + let desired_state = desired_state_from_persisted_enrollment(enrollment); + self.transition_unknown_to(desired_state); + return true; + } + } + + fn transition_unknown_to(&self, desired_state: RemoteControlDesiredState) { + self.desired_state_tx.send_if_modified(|state| { + if !matches!(*state, RemoteControlDesiredState::Unknown) { + return false; + } + *state = desired_state; + true + }); + } + + async fn wait_for_preference_resolution_retry(&mut self) -> bool { + tokio::select! { + _ = self.shutdown_token.cancelled() => false, + changed = self.desired_state_rx.changed() => changed.is_ok(), + _ = tokio::time::sleep(REMOTE_CONTROL_ACCOUNT_ID_RETRY_INTERVAL) => true, + } + } + + async fn wait_until_enabled(&mut self) -> bool { + tokio::select! { + _ = self.shutdown_token.cancelled() => false, + desired_state = self.desired_state_rx.wait_for(|state| state.is_enabled()) => desired_state.is_ok(), + } + } + + async fn connect( + &mut self, + shutdown_token: &CancellationToken, + app_server_client_name: Option<&str>, + ) -> ConnectOutcome { + self.status_publisher + .publish_status(RemoteControlConnectionStatus::Connecting); + let remote_control_target = match self.remote_control_target.as_ref() { + Some(remote_control_target) => remote_control_target.clone(), + None => match super::protocol::normalize_remote_control_url(&self.remote_control_url) { + Ok(remote_control_target) => { + self.remote_control_target = Some(remote_control_target.clone()); + remote_control_target + } + Err(err) => { + self.status_publisher + .publish_status(RemoteControlConnectionStatus::Errored); + warn!("remote control is enabled but the URL is invalid: {err}"); + tokio::select! { + _ = shutdown_token.cancelled() => return ConnectOutcome::Shutdown, + changed = self.desired_state_rx.wait_for(|state| !state.is_enabled()) => { + if changed.is_err() { + return ConnectOutcome::Shutdown; + } + return ConnectOutcome::Disabled; + } + } + } + }, + }; + + loop { + let subscribe_cursor = self.state.lock().await.subscribe_cursor.clone(); + let enrollment = self.current_enrollment.snapshot(); + info!( + websocket_url = %remote_control_target.websocket_url, + installation_id = %self.installation_id, + server_name = %self.server_name, + reconnect_attempt = self.reconnect_attempt.saturating_add(1), + has_enrollment = enrollment.is_some(), + server_id = ?enrollment.as_ref().map(|enrollment| enrollment.server_id.as_str()), + environment_id = ?enrollment.as_ref().map(|enrollment| enrollment.environment_id.as_str()), + subscribe_cursor_present = subscribe_cursor.is_some(), + app_server_client_name = ?app_server_client_name, + "connecting to app-server remote control websocket" + ); + let connect_options = RemoteControlConnectOptions { + installation_id: &self.installation_id, + server_name: &self.server_name, + subscribe_cursor: subscribe_cursor.as_deref(), + app_server_client_name, + desired_state_tx: &self.desired_state_tx, + persistence: &self.persistence, + }; + let auth_context = RemoteControlAuthContext { + auth_manager: &self.auth_manager, + auth_recovery: &mut self.auth_recovery, + auth_change_rx: &mut self.auth_change_rx, + }; + let mut disabled_rx = self.desired_state_rx.clone(); + let connect_result = tokio::select! { + _ = shutdown_token.cancelled() => return ConnectOutcome::Shutdown, + changed = disabled_rx.wait_for(|state| !state.is_enabled()) => { + if changed.is_err() { + return ConnectOutcome::Shutdown; + } + return ConnectOutcome::Disabled; + } + connect_result = async { + connect_remote_control_websocket( + &remote_control_target, + self.state_db.as_deref(), + auth_context, + &self.current_enrollment, + connect_options, + &self.status_publisher, + ) + .await + } => connect_result, + }; + + match connect_result { + Ok((websocket_connection, response)) => { + if !self.desired_state_rx.borrow().is_enabled() { + return ConnectOutcome::Disabled; + } + self.reconnect_attempt = 0; + self.auth_recovery = self.auth_manager.unauthorized_recovery(); + self.status_publisher + .publish_status(RemoteControlConnectionStatus::Connected); + let enrollment = self.current_enrollment.snapshot(); + info!( + websocket_url = %remote_control_target.websocket_url, + installation_id = %self.installation_id, + server_name = %self.server_name, + server_id = ?enrollment.as_ref().map(|enrollment| enrollment.server_id.as_str()), + environment_id = ?enrollment.as_ref().map(|enrollment| enrollment.environment_id.as_str()), + subscribe_cursor_present = subscribe_cursor.is_some(), + response_headers = %format_headers(response.headers()), + "connected to app-server remote control websocket" + ); + return ConnectOutcome::Connected(Box::new(websocket_connection)); + } + Err(err) => { + if !self.desired_state_rx.borrow().is_enabled() { + return ConnectOutcome::Disabled; + } + let server_retry_delay = remote_control_retry_delay(&err); + let reconnect_delay = if err.kind() == ErrorKind::WouldBlock { + server_retry_delay + .map_or(REMOTE_CONTROL_ACCOUNT_ID_RETRY_INTERVAL, |delay| { + REMOTE_CONTROL_ACCOUNT_ID_RETRY_INTERVAL.max(delay) + }) + } else { + self.status_publisher + .publish_status(RemoteControlConnectionStatus::Errored); + let reconnect_attempt = self.reconnect_attempt.saturating_add(1); + let (reconnect_delay, reconnect_backoff_reset) = + next_reconnect_delay(&mut self.reconnect_attempt); + let reconnect_delay = server_retry_delay + .map_or(reconnect_delay, |delay| reconnect_delay.max(delay)); + let enrollment = self.current_enrollment.snapshot(); + warn!( + websocket_url = %remote_control_target.websocket_url, + installation_id = %self.installation_id, + server_name = %self.server_name, + error = %err, + error_kind = ?err.kind(), + reconnect_attempt, + reconnect_delay = ?reconnect_delay, + reconnect_backoff_reset, + has_enrollment = enrollment.is_some(), + server_id = ?enrollment.as_ref().map(|enrollment| enrollment.server_id.as_str()), + environment_id = ?enrollment.as_ref().map(|enrollment| enrollment.environment_id.as_str()), + subscribe_cursor_present = subscribe_cursor.is_some(), + "failed to connect to app-server remote control websocket" + ); + if reconnect_backoff_reset { + info!( + reconnect_backoff_cap = ?REMOTE_CONTROL_RECONNECT_BACKOFF_CAP, + "reset app-server remote control websocket reconnect backoff after cap" + ); + } + reconnect_delay + }; + tokio::select! { + _ = shutdown_token.cancelled() => return ConnectOutcome::Shutdown, + changed = self.desired_state_rx.wait_for(|state| !state.is_enabled()) => { + if changed.is_err() { + return ConnectOutcome::Shutdown; + } + return ConnectOutcome::Disabled; + } + // Auth changes may shorten local backoff after the server deadline. + changed = wait_for_auth_change(&mut self.auth_change_rx, server_retry_delay) => { + if changed.is_err() { + return ConnectOutcome::Shutdown; + } + self.auth_recovery = self.auth_manager.unauthorized_recovery(); + self.reconnect_attempt = 0; + info!("retrying app-server remote control websocket after auth changed"); + } + _ = tokio::time::sleep(reconnect_delay) => {} + } + } + } + } + } + + async fn run_connection( + &self, + websocket_connection: WebSocketStream>, + shutdown_token: CancellationToken, + ) -> ConnectionEndReason { + if !self.auth_manager.owner.is_current() { + return ConnectionEndReason::AuthOwnerChanged; + } + self.client_tracker.lock().await.auth = Some(self.auth_manager.owner.clone()); + let (websocket_writer, websocket_reader) = websocket_connection.split(); + let mut join_set = tokio::task::JoinSet::new(); + + join_set.spawn(Self::run_server_writer( + self.state.clone(), + self.server_event_rx.clone(), + self.used_rx.clone(), + websocket_writer, + REMOTE_CONTROL_WEBSOCKET_PING_INTERVAL, + shutdown_token.clone(), + )); + join_set.spawn(Self::run_websocket_reader( + self.client_tracker.clone(), + self.state.clone(), + websocket_reader, + REMOTE_CONTROL_WEBSOCKET_PONG_TIMEOUT, + shutdown_token.clone(), + )); + + let mut desired_state_rx = self.desired_state_rx.clone(); + let connection_end_reason = tokio::select! { + biased; + _ = shutdown_token.cancelled() => ConnectionEndReason::Shutdown, + _ = self.auth_manager.owner.invalidated() => ConnectionEndReason::AuthOwnerChanged, + changed = desired_state_rx.wait_for(|state| !state.is_enabled()) => { + if changed.is_ok() { + self.status_publisher + .publish_status(RemoteControlConnectionStatus::Disabled); + ConnectionEndReason::Disabled + } else { + ConnectionEndReason::EnabledWatchClosed + } + } + _ = join_set.join_next() => ConnectionEndReason::ConnectionWorkerStopped, + }; + shutdown_token.cancel(); + + Self::join_connection_workers(&mut join_set, REMOTE_CONTROL_CONNECTION_SHUTDOWN_TIMEOUT) + .await; + connection_end_reason + } + + async fn join_connection_workers( + join_set: &mut tokio::task::JoinSet<()>, + shutdown_timeout: std::time::Duration, + ) { + if tokio::time::timeout(shutdown_timeout, Self::drain_join_set(join_set)) + .await + .is_ok() + { + return; + } + + warn!( + shutdown_timeout = ?shutdown_timeout, + remaining_workers = join_set.len(), + "timed out waiting for remote control connection workers to stop; aborting" + ); + join_set.abort_all(); + Self::drain_join_set(join_set).await; + } + + async fn drain_join_set(join_set: &mut tokio::task::JoinSet<()>) { + while join_set.join_next().await.is_some() {} + } + + async fn run_server_writer( + state: Arc>, + server_event_rx: Arc>>, + used_rx: watch::Receiver, + websocket_writer: SplitSink< + WebSocketStream>, + tungstenite::Message, + >, + ping_interval: std::time::Duration, + shutdown_token: CancellationToken, + ) { + let result = Self::run_server_writer_inner( + state, + server_event_rx, + used_rx, + websocket_writer, + ping_interval, + shutdown_token, + ) + .await; + if let Err(err) = result { + warn!("remote control websocket writer disconnected, err: {err}"); + } else { + warn!("remote control websocket writer was stopped"); + } + } + + #[expect( + clippy::await_holding_invalid_type, + reason = "remote-control server event receiver is shared across reconnects" + )] + async fn run_server_writer_inner( + state: Arc>, + server_event_rx: Arc>>, + mut used_rx: watch::Receiver, + mut websocket_writer: SplitSink< + WebSocketStream>, + tungstenite::Message, + >, + ping_interval: std::time::Duration, + shutdown_token: CancellationToken, + ) -> io::Result<()> { + let server_envelopes = state + .lock() + .await + .outbound_buffer + .server_envelopes() + .cloned() + .collect::>(); + for server_envelope in server_envelopes { + let payload = match serde_json::to_string(&server_envelope) { + Ok(payload) => payload, + Err(err) => { + error!("failed to serialize remote-control server event: {err}"); + continue; + } + }; + tokio::select! { + _ = shutdown_token.cancelled() => return Ok(()), + send_result = websocket_writer.send(tungstenite::Message::Text(payload.into())) => { + if let Err(err) = send_result { + return Err(io::Error::other(err)); + } + } + }; + } + + let mut ping_interval = + tokio::time::interval_at(tokio::time::Instant::now() + ping_interval, ping_interval); + ping_interval.set_missed_tick_behavior(MissedTickBehavior::Skip); + + let mut server_event_rx = server_event_rx.lock().await; + loop { + let outbound_has_capacity = *used_rx.borrow() < super::CHANNEL_CAPACITY; + let queued_server_envelope = tokio::select! { + _ = shutdown_token.cancelled() => return Ok(()), + _ = ping_interval.tick() => { + tokio::select! { + _ = shutdown_token.cancelled() => return Ok(()), + send_result = websocket_writer.send(tungstenite::Message::Ping(Vec::new().into())) => { + if let Err(err) = send_result { + return Err(io::Error::other(err)); + } + } + }; + continue; + } + wait_result = used_rx.changed(), if !outbound_has_capacity => + { + if wait_result.is_err() { + return Err(io::Error::new( + ErrorKind::UnexpectedEof, + "outbound buffer usage channel closed", + )); + } + continue; + } + recv_result = server_event_rx.recv(), if outbound_has_capacity => { + match recv_result { + Some(queued_server_envelope) => queued_server_envelope, + None => { + return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "server event channel closed")); + } + } + } + }; + let (payloads, write_complete_tx) = { + let mut state = state.lock().await; + let seq_key = ( + queued_server_envelope.client_id.clone(), + queued_server_envelope.stream_id.clone(), + ); + let seq_id = *state + .next_seq_id_by_stream + .entry(seq_key.clone()) + .or_insert(1); + + let server_envelope = ServerEnvelope { + event: queued_server_envelope.event, + client_id: queued_server_envelope.client_id, + seq_id, + stream_id: queued_server_envelope.stream_id, + }; + let server_envelopes = match split_server_envelope_for_transport(server_envelope) { + Ok(server_envelopes) => server_envelopes, + Err(err) => { + error!("failed to split remote-control server event: {err}"); + continue; + } + }; + let mut payloads = Vec::with_capacity(server_envelopes.len()); + for server_envelope in server_envelopes { + let payload = match serde_json::to_string(&server_envelope) { + Ok(payload) => payload, + Err(err) => { + error!("failed to serialize remote-control server event: {err}"); + continue; + } + }; + state.outbound_buffer.insert(&server_envelope); + payloads.push(payload); + } + state + .next_seq_id_by_stream + .insert(seq_key, seq_id.saturating_add(1)); + + (payloads, queued_server_envelope.write_complete_tx) + }; + + for payload in payloads { + tokio::select! { + _ = shutdown_token.cancelled() => return Ok(()), + send_result = websocket_writer.send(tungstenite::Message::Text(payload.into())) => { + if let Err(err) = send_result { + return Err(io::Error::other(err)); + } + } + } + } + if let Some(write_complete_tx) = write_complete_tx { + let _ = write_complete_tx.send(()); + } + } + } + + async fn run_websocket_reader( + client_tracker: Arc>, + state: Arc>, + websocket_reader: SplitStream>>, + pong_timeout: std::time::Duration, + shutdown_token: CancellationToken, + ) { + let result = Self::run_websocket_reader_inner( + client_tracker, + state, + websocket_reader, + pong_timeout, + shutdown_token, + ) + .await; + if let Err(err) = result { + warn!("remote control websocket reader disconnected, err: {err}"); + } else { + warn!("remote control websocket reader was stopped"); + } + } + + #[expect( + clippy::await_holding_invalid_type, + reason = "remote-control client tracking must stay serialized while processing inbound events" + )] + async fn run_websocket_reader_inner( + client_tracker: Arc>, + state: Arc>, + mut websocket_reader: SplitStream>>, + pong_timeout: std::time::Duration, + shutdown_token: CancellationToken, + ) -> io::Result<()> { + let mut client_tracker = client_tracker.lock().await; + let mut idle_sweep_interval = tokio::time::interval(REMOTE_CONTROL_IDLE_SWEEP_INTERVAL); + idle_sweep_interval.set_missed_tick_behavior(MissedTickBehavior::Skip); + let pong_deadline = tokio::time::sleep(pong_timeout); + tokio::pin!(pong_deadline); + + loop { + let incoming_message = tokio::select! { + _ = shutdown_token.cancelled() => return Ok(()), + _ = &mut pong_deadline => { + return Err(io::Error::new( + ErrorKind::TimedOut, + "remote control websocket pong timeout", + )); + } + client_key = client_tracker.bookkeep_join_set() => { + let Some(client_key) = client_key else { + continue; + }; + if client_tracker.close_client(&client_key).await.is_err() { + return Ok(()); + } + state + .lock() + .await + .client_segment_reassembler + .invalidate_stream(&client_key.0, &client_key.1); + state + .lock() + .await + .invalidate_client_message_stream(&client_key.0, &client_key.1); + continue; + } + _ = idle_sweep_interval.tick() => { + match client_tracker.close_expired_clients().await { + Ok(client_keys) => { + let mut websocket_state = state.lock().await; + for (client_id, stream_id) in client_keys { + websocket_state + .client_segment_reassembler + .invalidate_stream(&client_id, &stream_id); + websocket_state + .invalidate_client_message_stream(&client_id, &stream_id); + } + } + Err(_) => return Ok(()), + } + continue; + } + incoming_message = websocket_reader.next() => { + match incoming_message { + Some(incoming_message) => incoming_message, + None => return Err(io::Error::new(ErrorKind::UnexpectedEof, "websocket stream ended")), + } + } + }; + let (client_envelope, wire_size_bytes) = match incoming_message { + Ok(tungstenite::Message::Text(text)) => { + let wire_size_bytes = text.len(); + match serde_json::from_str::(&text) { + Ok(client_envelope) => (client_envelope, wire_size_bytes), + Err(err) => { + warn!("failed to deserialize remote-control client event: {err}"); + continue; + } + } + } + Ok(tungstenite::Message::Pong(_)) => { + pong_deadline + .as_mut() + .reset(tokio::time::Instant::now() + pong_timeout); + continue; + } + Ok(tungstenite::Message::Ping(_)) | Ok(tungstenite::Message::Frame(_)) => continue, + Ok(tungstenite::Message::Binary(_)) => { + warn!("dropping unsupported binary remote-control websocket message"); + continue; + } + Ok(tungstenite::Message::Close(_)) => { + return Err(io::Error::new( + ErrorKind::ConnectionAborted, + "websocket disconnected", + )); + } + Err(err) => { + return Err(io::Error::new( + ErrorKind::InvalidData, + format!("failed to read from websocket: {err}"), + )); + } + }; + + let client_message_key = WebsocketState::client_message_key(&client_envelope); + let observation = { + let mut websocket_state = state.lock().await; + websocket_state.observe_client_message(client_envelope, wire_size_bytes) + }; + let client_envelope = match observation { + ClientSegmentObservation::Forward(client_envelope) => *client_envelope, + ClientSegmentObservation::Pending | ClientSegmentObservation::Dropped => continue, + }; + + let closed_client = + matches!(&client_envelope.event, ClientEvent::ClientClosed).then(|| { + ( + client_envelope.client_id.clone(), + client_envelope.stream_id.clone(), + ) + }); + let delivered_client_envelope = client_envelope.clone(); + if client_tracker + .handle_message(client_envelope) + .await + .is_err() + { + return Ok(()); + } + state + .lock() + .await + .record_client_message_delivery(&delivered_client_envelope, client_message_key); + if let Some((client_id, stream_id)) = closed_client { + let mut websocket_state = state.lock().await; + if let Some(stream_id) = stream_id { + websocket_state + .client_segment_reassembler + .invalidate_stream(&client_id, &stream_id); + websocket_state.invalidate_client_message_stream(&client_id, &stream_id); + } else { + websocket_state + .client_segment_reassembler + .invalidate_client(&client_id); + websocket_state.invalidate_client_message_client(&client_id); + } + } + } + } +} + +fn set_remote_control_header( + headers: &mut tungstenite::http::HeaderMap, + name: &'static str, + value: &str, +) -> io::Result<()> { + let header_value = HeaderValue::from_str(value).map_err(|err| { + io::Error::new( + ErrorKind::InvalidInput, + format!("invalid remote control header `{name}`: {err}"), + ) + })?; + headers.insert(name, header_value); + Ok(()) +} + +async fn build_remote_control_websocket_request( + websocket_url: &str, + enrollment: &RemoteControlEnrollment, + installation_id: &str, + subscribe_cursor: Option<&str>, +) -> io::Result> { + let mut request = websocket_url.into_client_request().map_err(|err| { + io::Error::new( + ErrorKind::InvalidInput, + format!("invalid remote control websocket URL `{websocket_url}`: {err}"), + ) + })?; + let headers = request.headers_mut(); + set_remote_control_header(headers, "x-codex-server-id", &enrollment.server_id)?; + set_remote_control_header( + headers, + "x-codex-name", + &base64::engine::general_purpose::STANDARD.encode(&enrollment.server_name), + )?; + set_remote_control_header( + headers, + "x-codex-protocol-version", + REMOTE_CONTROL_PROTOCOL_VERSION, + )?; + set_remote_control_header( + headers, + "authorization", + &format!( + "Bearer {}", + enrollment + .remote_control_token + .as_deref() + .ok_or_else(|| io::Error::other("missing remote control server token"))? + ), + )?; + set_remote_control_header( + headers, + REMOTE_CONTROL_INSTALLATION_ID_HEADER, + installation_id, + )?; + if let Some(host_device_kind) = host_device_kind().await { + set_remote_control_header( + headers, + REMOTE_CONTROL_HOST_DEVICE_KIND_HEADER, + host_device_kind, + )?; + } + if let Some(subscribe_cursor) = subscribe_cursor { + set_remote_control_header( + headers, + REMOTE_CONTROL_SUBSCRIBE_CURSOR_HEADER, + subscribe_cursor, + )?; + } + Ok(request) +} + +async fn wait_for_auth_change( + auth_change_rx: &mut watch::Receiver, + server_retry_delay: Option, +) -> Result<(), watch::error::RecvError> { + if let Some(delay) = server_retry_delay { + tokio::time::sleep(delay).await; + } + auth_change_rx.changed().await +} + +fn next_reconnect_delay(reconnect_attempt: &mut u64) -> (std::time::Duration, bool) { + let reconnect_delay = backoff(*reconnect_attempt).min(REMOTE_CONTROL_RECONNECT_BACKOFF_CAP); + let reconnect_backoff_reset = reconnect_delay == REMOTE_CONTROL_RECONNECT_BACKOFF_CAP; + *reconnect_attempt = if reconnect_backoff_reset { + 0 + } else { + (*reconnect_attempt).saturating_add(1) + }; + (reconnect_delay, reconnect_backoff_reset) +} + +pub(super) async fn connect_remote_control_websocket( + remote_control_target: &RemoteControlTarget, + state_db: Option<&StateRuntime>, + mut auth_context: RemoteControlAuthContext<'_>, + current_enrollment: &CurrentRemoteControlEnrollment, + connect_options: RemoteControlConnectOptions<'_>, + status_publisher: &RemoteControlStatusPublisher, +) -> io::Result<( + WebSocketStream>, + tungstenite::http::Response<()>, +)> { + ensure_rustls_crypto_provider(); + + let (auth, enrollment) = { + let mut lease = current_enrollment.lock_for_request().await?; + let auth_result = prepare_remote_control_enrollment( + remote_control_target, + state_db, + &mut auth_context, + &mut lease, + connect_options, + status_publisher, + ) + .await; + let auth = current_enrollment.record_retry_after(auth_result)?; + let enrollment = lease.as_ref().cloned().ok_or_else(|| { + io::Error::other("missing remote control enrollment after enrollment step") + })?; + (auth, enrollment) + }; + let request = build_remote_control_websocket_request( + &remote_control_target.websocket_url, + &enrollment, + connect_options.installation_id, + connect_options.subscribe_cursor, + ) + .await?; + + current_enrollment.check_retry_after()?; + let websocket_connect_result = tokio::time::timeout( + REMOTE_CONTROL_WEBSOCKET_CONNECT_TIMEOUT, + connect_async(request), + ) + .await + .map_err(|_| { + io::Error::new( + ErrorKind::TimedOut, + format!( + "timed out connecting to remote control websocket at `{}` after {:?}", + remote_control_target.websocket_url, REMOTE_CONTROL_WEBSOCKET_CONNECT_TIMEOUT + ), + ) + })?; + + match websocket_connect_result { + Ok((websocket_stream, response)) => Ok((websocket_stream, response.map(|_| ()))), + Err(err) => { + match &err { + tungstenite::Error::Http(response) + if websocket_response_reports_missing_remote_app_server(response) => + { + info!( + "remote control websocket returned HTTP 404; replacing stale enrollment: websocket_url={}, account_id={}, server_id={}, environment_id={}", + remote_control_target.websocket_url, + auth.account_id, + enrollment.server_id, + enrollment.environment_id + ); + replace_remote_control_enrollment_if_matches( + state_db, + remote_control_target, + RemoteControlEnrollmentAuthContext { + auth: &auth, + recovery: &mut auth_context, + }, + current_enrollment, + &enrollment, + connect_options, + status_publisher, + ) + .await?; + } + tungstenite::Error::Http(response) if response.status().as_u16() == 404 => { + let response_body = response + .body() + .as_deref() + .map(preview_remote_control_response_body) + .unwrap_or_else(|| "".to_string()); + warn!( + websocket_url = %remote_control_target.websocket_url, + account_id = %auth.account_id, + server_id = %enrollment.server_id, + environment_id = %enrollment.environment_id, + response_status = %response.status(), + response_headers = %format_headers(response.headers()), + response_body = %response_body, + "remote control websocket returned unrecognized HTTP 404; preserving enrollment before retry" + ); + } + tungstenite::Error::Http(response) + if matches!(response.status().as_u16(), 429 | 503) => + { + return current_enrollment.record_retry_after(Err( + RemoteControlServerRequestError::io_error( + format_remote_control_websocket_connect_error( + &remote_control_target.websocket_url, + &err, + ), + Some(response.status()), + retry_after_with_jitter( + response.headers(), + time::OffsetDateTime::now_utc(), + ), + /*timed_out*/ false, + ), + )); + } + tungstenite::Error::Http(response) + if matches!(response.status().as_u16(), 401 | 403) => + { + clear_remote_control_server_token_if_matches(current_enrollment, &enrollment) + .await?; + return Err(io::Error::other(format!( + "remote control websocket auth failed with HTTP {}; refreshing server token before reconnect", + response.status() + ))); + } + _ => {} + } + Err(io::Error::other( + format_remote_control_websocket_connect_error( + &remote_control_target.websocket_url, + &err, + ), + )) + } + } +} + +async fn prepare_remote_control_enrollment( + remote_control_target: &RemoteControlTarget, + state_db: Option<&StateRuntime>, + auth_context: &mut RemoteControlAuthContext<'_>, + enrollment: &mut Option, + connect_options: RemoteControlConnectOptions<'_>, + status_publisher: &RemoteControlStatusPublisher, +) -> io::Result { + let Some(state_db) = state_db else { + *enrollment = None; + return Err(io::Error::new( + ErrorKind::NotFound, + "remote control requires sqlite state db", + )); + }; + + let auth = match load_remote_control_auth(auth_context.auth_manager).await { + Ok(auth) => auth, + Err(err) => { + if err.kind() == ErrorKind::PermissionDenied { + *enrollment = None; + status_publisher.publish_environment_id(/*environment_id*/ None); + } + return Err(err); + } + }; + if enrollment + .as_ref() + .is_some_and(|enrollment| enrollment.account_id != auth.account_id) + { + return Err(io::Error::new( + ErrorKind::Interrupted, + "enrollment belongs to another remote session", + )); + } + if let Some(enrollment) = enrollment.as_mut() { + enrollment.remote_control_target = remote_control_target.clone(); + } + + if let Some(enrollment) = enrollment.as_ref() { + status_publisher.publish_environment_id(Some(enrollment.environment_id.clone())); + } + + if enrollment.is_none() { + let _persistence = + super::persistence::read_lock(auth_context.auth_manager, connect_options.persistence) + .await?; + let loaded_enrollment = load_persisted_remote_control_enrollment( + Some(state_db), + remote_control_target, + &auth.account_id, + connect_options.app_server_client_name, + ) + .await?; + if let Some(loaded_enrollment) = loaded_enrollment.as_ref() { + status_publisher.publish_environment_id(Some(loaded_enrollment.environment_id.clone())); + } + *enrollment = loaded_enrollment.map(|mut enrollment| { + enrollment.server_name = connect_options.server_name.to_string(); + enrollment + }); + } + + enroll_and_persist_remote_control_server( + remote_control_target, + state_db, + RemoteControlEnrollmentAuthContext { + auth: &auth, + recovery: auth_context, + }, + enrollment, + connect_options, + status_publisher, + RemoteControlEnrollmentSelection::ReuseOrCreate, + ) + .await?; + + if enrollment + .as_ref() + .ok_or_else(|| io::Error::other("missing remote control enrollment after enrollment step"))? + .should_refresh_server_token() + { + let enrollment_ref = enrollment.as_ref().ok_or_else(|| { + io::Error::other("missing remote control enrollment after enrollment step") + })?; + let server_id = enrollment_ref.server_id.clone(); + let environment_id = enrollment_ref.environment_id.clone(); + + info!( + "refreshing remote control server token: websocket_url={}, refresh_url={}, account_id={}, server_id={}, environment_id={}", + remote_control_target.websocket_url, + remote_control_target.refresh_url, + auth.account_id, + server_id, + environment_id + ); + let enrollment_ref = enrollment.as_mut().ok_or_else(|| { + io::Error::other("missing remote control enrollment before server refresh") + })?; + match refresh_remote_control_server(&auth, connect_options.installation_id, enrollment_ref) + .await + { + Ok(()) => {} + Err(err) if err.kind() == ErrorKind::NotFound => { + info!( + "remote control server refresh returned HTTP 404; replacing stale enrollment: websocket_url={}, account_id={}, server_id={}, environment_id={}", + remote_control_target.websocket_url, auth.account_id, server_id, environment_id + ); + enroll_and_persist_remote_control_server( + remote_control_target, + state_db, + RemoteControlEnrollmentAuthContext { + auth: &auth, + recovery: auth_context, + }, + enrollment, + connect_options, + status_publisher, + RemoteControlEnrollmentSelection::ReplaceExisting, + ) + .await?; + } + Err(err) if err.kind() == ErrorKind::PermissionDenied => { + if recover_remote_control_auth( + auth_context.auth_recovery, + auth_context.auth_change_rx, + ) + .await + { + return Err(io::Error::other(format!( + "{err}; retrying after auth recovery" + ))); + } + enrollment_ref.clear_server_token(); + return Err(err); + } + Err(err) => return Err(err), + } + } + + Ok(auth) +} + +fn websocket_response_reports_missing_remote_app_server( + response: &tungstenite::http::Response>>, +) -> bool { + response.status().as_u16() == 404 + && response.body().as_deref().is_some_and(|body| { + serde_json::from_slice::(body).is_ok_and(|body| { + body.get("detail").and_then(serde_json::Value::as_str) + == Some(REMOTE_APP_SERVER_NOT_FOUND_DETAIL) + }) + }) +} + +async fn replace_remote_control_enrollment_if_matches( + state_db: Option<&StateRuntime>, + remote_control_target: &RemoteControlTarget, + auth_context: RemoteControlEnrollmentAuthContext<'_, '_>, + current_enrollment: &CurrentRemoteControlEnrollment, + enrollment: &RemoteControlEnrollment, + connect_options: RemoteControlConnectOptions<'_>, + status_publisher: &RemoteControlStatusPublisher, +) -> io::Result<()> { + let Some(state_db) = state_db else { + return Err(io::Error::new( + ErrorKind::NotFound, + "remote control requires sqlite state db", + )); + }; + let mut lease = current_enrollment.lock_for_request().await?; + if !lease + .as_ref() + .is_some_and(|current| same_remote_control_enrollment(current, enrollment)) + { + return Ok(()); + } + let result = enroll_and_persist_remote_control_server( + remote_control_target, + state_db, + auth_context, + &mut lease, + connect_options, + status_publisher, + RemoteControlEnrollmentSelection::ReplaceExisting, + ) + .await; + current_enrollment.record_retry_after(result) +} + +async fn clear_remote_control_server_token_if_matches( + current_enrollment: &CurrentRemoteControlEnrollment, + enrollment: &RemoteControlEnrollment, +) -> io::Result<()> { + let mut current_enrollment = current_enrollment.lock().await; + let current_enrollment = current_enrollment + .as_mut() + .filter(|current| same_remote_control_enrollment(current, enrollment)) + .ok_or_else(|| { + io::Error::other("missing remote control enrollment after websocket auth failure") + })?; + if current_enrollment.remote_control_token == enrollment.remote_control_token { + current_enrollment.clear_server_token(); + } + Ok(()) +} + +async fn enroll_and_persist_remote_control_server( + remote_control_target: &RemoteControlTarget, + state_db: &StateRuntime, + auth_context: RemoteControlEnrollmentAuthContext<'_, '_>, + enrollment: &mut Option, + connect_options: RemoteControlConnectOptions<'_>, + status_publisher: &RemoteControlStatusPublisher, + selection: RemoteControlEnrollmentSelection, +) -> io::Result<()> { + match selection { + RemoteControlEnrollmentSelection::ReuseOrCreate => { + if enrollment.is_some() { + return Ok(()); + } + } + RemoteControlEnrollmentSelection::ReplaceExisting => {} + } + if !connect_options.desired_state_tx.borrow().is_enabled() { + return Err(io::Error::new( + ErrorKind::Interrupted, + "remote control disabled before enrollment", + )); + } + + info!( + "creating new remote control enrollment: websocket_url={}, enroll_url={}, account_id={}", + remote_control_target.websocket_url, + remote_control_target.enroll_url, + auth_context.auth.account_id + ); + let new_enrollment = match enroll_remote_control_server( + remote_control_target, + auth_context.auth, + connect_options.installation_id, + connect_options.server_name, + ) + .await + { + Ok(new_enrollment) => new_enrollment, + Err(err) + if err.kind() == ErrorKind::PermissionDenied + && recover_remote_control_auth( + auth_context.recovery.auth_recovery, + auth_context.recovery.auth_change_rx, + ) + .await => + { + return Err(io::Error::other(format!( + "{err}; retrying after auth recovery" + ))); + } + Err(err) => return Err(err), + }; + super::persistence::save_enrollment( + auth_context.recovery.auth_manager, + connect_options.persistence, + state_db, + &new_enrollment, + connect_options.app_server_client_name, + connect_options.desired_state_tx, + ) + .await?; + info!( + "created new remote control enrollment: websocket_url={}, account_id={}, server_id={}, environment_id={}", + remote_control_target.websocket_url, + new_enrollment.account_id, + new_enrollment.server_id, + new_enrollment.environment_id + ); + status_publisher.publish_environment_id(Some(new_enrollment.environment_id.clone())); + *enrollment = Some(new_enrollment); + Ok(()) +} + +fn format_remote_control_websocket_connect_error( + websocket_url: &str, + err: &tungstenite::Error, +) -> String { + let mut message = + format!("failed to connect app-server remote control websocket `{websocket_url}`: {err}"); + let tungstenite::Error::Http(response) = err else { + return message; + }; + + message.push_str(&format!(", {}", format_headers(response.headers()))); + if let Some(body) = response.body().as_ref() + && !body.is_empty() + { + let body_preview = preview_remote_control_response_body(body); + message.push_str(&format!(", body: {body_preview}")); + } + + message +} + +#[cfg(test)] +#[path = "websocket_refresh_tests.rs"] +mod refresh_tests; + +#[cfg(test)] +mod tests { + use super::*; + use crate::outgoing_message::OutgoingMessage; + use crate::transport::remote_control::ServerEvent; + use crate::transport::remote_control::auth::mark_recovery_auth_change_seen; + use crate::transport::remote_control::protocol::StreamId; + use crate::transport::remote_control::protocol::normalize_remote_control_url; + use chrono::Utc; + use codex_app_server_protocol::ConfigWarningNotification; + use codex_app_server_protocol::JSONRPCMessage; + use codex_app_server_protocol::JSONRPCNotification; + use codex_app_server_protocol::ServerNotification; + use codex_app_server_protocol::ServerNotificationEnvelope; + use codex_config::types::AuthCredentialsStoreMode; + use codex_core::test_support::auth_manager_from_auth; + use codex_login::AuthDotJson; + use codex_login::AuthKeyringBackendKind; + use codex_login::AuthManager; + use codex_login::CodexAuth; + use codex_login::save_auth; + use codex_login::token_data::TokenData; + use codex_login::token_data::parse_chatgpt_jwt_claims; + use codex_protocol::auth::AuthMode; + use codex_state::StateRuntime; + use codex_utils_absolute_path::test_support::PathExt; + use futures::StreamExt; + use pretty_assertions::assert_eq; + use std::sync::Arc; + use tempfile::TempDir; + use tokio::io::AsyncBufReadExt; + use tokio::io::AsyncWriteExt; + use tokio::io::BufReader; + use tokio::net::TcpListener; + use tokio::net::TcpStream; + use tokio::sync::mpsc; + use tokio::time::Duration; + use tokio::time::timeout; + use tokio_tungstenite::accept_async; + + // Windows Bazel CI can take longer than a few seconds for the websocket + // client connection attempt to reach the local test listener. + #[cfg(windows)] + pub(super) const TEST_HTTP_ACCEPT_TIMEOUT: Duration = Duration::from_secs(30); + #[cfg(not(windows))] + pub(super) const TEST_HTTP_ACCEPT_TIMEOUT: Duration = Duration::from_secs(5); + pub(super) const TEST_INSTALLATION_ID: &str = "11111111-1111-4111-8111-111111111111"; + pub(super) const TEST_REMOTE_CONTROL_SERVER_TOKEN: &str = "Remote Control Token"; + + pub(super) fn remote_control_enrollment( + remote_control_token: Option<&str>, + ) -> RemoteControlEnrollment { + RemoteControlEnrollment { + remote_control_target: normalize_remote_control_url("http://localhost/backend-api/") + .expect("target should normalize"), + account_id: "account_id".to_string(), + environment_id: "env_test".to_string(), + server_id: "srv_e_test".to_string(), + server_name: "test-server".to_string(), + remote_control_token: remote_control_token.map(str::to_string), + expires_at: remote_control_token + .map(|_| time::OffsetDateTime::now_utc() + time::Duration::hours(1)), + next_refresh_at: None, + } + } + + pub(super) fn test_current_enrollment( + enrollment: Option, + ) -> CurrentRemoteControlEnrollment { + Arc::new(RemoteControlEnrollmentState::new(enrollment)) + } + + #[tokio::test(start_paused = true)] + async fn queued_auth_change_waits_for_server_delay() { + let (auth_change_tx, mut auth_change_rx) = watch::channel(0); + auth_change_tx.send_replace(1); + let started = tokio::time::Instant::now(); + let auth_change = wait_for_auth_change(&mut auth_change_rx, Some(Duration::from_secs(5))); + tokio::pin!(auth_change); + + assert!( + timeout(Duration::from_secs(4), auth_change.as_mut()) + .await + .is_err() + ); + timeout(Duration::from_secs(2), auth_change) + .await + .expect("queued credentials should be usable as soon as the server delay expires") + .expect("auth watch should remain open"); + assert_eq!(started.elapsed(), Duration::from_secs(5)); + } + + #[test] + fn next_reconnect_delay_resets_after_cap() { + let mut reconnect_attempt = 9; + + let (reconnect_delay, reconnect_backoff_reset) = + next_reconnect_delay(&mut reconnect_attempt); + + assert_eq!(reconnect_delay, REMOTE_CONTROL_RECONNECT_BACKOFF_CAP); + assert!(reconnect_backoff_reset); + assert_eq!(reconnect_attempt, 0); + + let (reconnect_delay, reconnect_backoff_reset) = + next_reconnect_delay(&mut reconnect_attempt); + + assert!(reconnect_delay >= Duration::from_millis(180)); + assert!(reconnect_delay <= Duration::from_millis(220)); + assert!(!reconnect_backoff_reset); + assert_eq!(reconnect_attempt, 1); + } + + #[test] + fn websocket_404_only_reports_explicit_missing_remote_app_server() { + let cases = [ + ( + Some(br#"{"detail":"Remote app server not found"}"#.to_vec()), + true, + ), + ( + Some(br#" { "detail": "Remote app server not found", "extra": true } "#.to_vec()), + true, + ), + (Some(br#"{"detail":"Not Found"}"#.to_vec()), false), + (Some(b"Not Found".to_vec()), false), + (Some(b"{".to_vec()), false), + (Some(Vec::new()), false), + (None, false), + ]; + + for (body, expected) in cases { + let response = tungstenite::http::Response::builder() + .status(/*status*/ 404) + .body(body) + .expect("response should build"); + assert_eq!( + websocket_response_reports_missing_remote_app_server(&response), + expected + ); + } + + let response = tungstenite::http::Response::builder() + .status(/*status*/ 503) + .body(Some( + br#"{"detail":"Remote app server not found"}"#.to_vec(), + )) + .expect("response should build"); + assert!(!websocket_response_reports_missing_remote_app_server( + &response + )); + } + + pub(super) fn remote_control_status_channel() -> ( + RemoteControlStatusPublisher, + watch::Receiver, + ) { + let (status_tx, status_rx) = watch::channel(RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connecting, + server_name: "test-server".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: None, + }); + (RemoteControlStatusPublisher::new(status_tx), status_rx) + } + + pub(super) fn enabled_desired_state_sender() -> watch::Sender { + watch::channel(RemoteControlDesiredState::Enabled { + persistence_preference: None, + }) + .0 + } + + #[test] + fn mark_recovery_auth_change_seen_marks_only_recovery_revision_seen() { + let (auth_change_tx, mut auth_change_rx) = watch::channel(0u64); + let auth_change_revision_before_recovery = *auth_change_rx.borrow(); + auth_change_tx.send_modify(|revision| *revision += 1); + + mark_recovery_auth_change_seen(&mut auth_change_rx, auth_change_revision_before_recovery); + + assert!( + !auth_change_rx + .has_changed() + .expect("auth change watch should remain open") + ); + } + + #[test] + fn mark_recovery_auth_change_seen_preserves_racing_auth_change() { + let (auth_change_tx, mut auth_change_rx) = watch::channel(0u64); + let auth_change_revision_before_recovery = *auth_change_rx.borrow(); + auth_change_tx.send_modify(|revision| *revision += 1); + auth_change_tx.send_modify(|revision| *revision += 1); + + mark_recovery_auth_change_seen(&mut auth_change_rx, auth_change_revision_before_recovery); + + assert!( + auth_change_rx + .has_changed() + .expect("auth change watch should remain open") + ); + } + + pub(super) async fn remote_control_state_runtime(codex_home: &TempDir) -> Arc { + StateRuntime::init( + codex_state::SqliteConfig::new_for_testing(codex_home.path().abs()), + "test-provider".to_string(), + ) + .await + .expect("state runtime should initialize") + } + + pub(super) fn remote_control_auth_manager() -> Arc { + auth_manager_from_auth(CodexAuth::create_dummy_chatgpt_auth_for_testing()) + } + + pub(super) fn remote_control_url_for_listener(listener: &TcpListener) -> String { + let addr = listener + .local_addr() + .expect("listener should have a local addr"); + format!("http://{addr}/backend-api/") + } + + pub(super) fn remote_control_auth_dot_json(access_token: &str) -> AuthDotJson { + #[derive(serde::Serialize)] + struct Header { + alg: &'static str, + typ: &'static str, + } + + let header = Header { + alg: "none", + typ: "JWT", + }; + let payload = serde_json::json!({ + "email": "user@example.com", + "https://api.openai.com/auth": { + "chatgpt_user_id": "user-12345", + "user_id": "user-12345", + "chatgpt_account_id": "account_id" + } + }); + let b64 = |bytes: &[u8]| base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes); + let header_b64 = b64(&serde_json::to_vec(&header).expect("header should serialize")); + let payload_b64 = b64(&serde_json::to_vec(&payload).expect("payload should serialize")); + let fake_jwt = format!("{header_b64}.{payload_b64}.sig"); + + AuthDotJson { + auth_mode: Some(AuthMode::Chatgpt), + openai_api_key: None, + tokens: Some(TokenData { + id_token: parse_chatgpt_jwt_claims(&fake_jwt).expect("fake jwt should parse"), + access_token: access_token.to_string(), + refresh_token: "refresh-token".to_string(), + account_id: Some("account_id".to_string()), + }), + last_refresh: Some(Utc::now()), + agent_identity: None, + personal_access_token: None, + bedrock_api_key: None, + bedrock_access_keys: None, + } + } + + #[tokio::test] + async fn connect_remote_control_websocket_includes_http_error_details() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let expected_error = format!( + "failed to connect app-server remote control websocket `{}`: HTTP error: 503 Service Unavailable, request-id: , cf-ray: , body: upstream unavailable", + remote_control_target.websocket_url + ); + let server_task = tokio::spawn(async move { + let (stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + "GET /backend-api/wham/remote/control/server HTTP/1.1" + ); + respond_with_status_and_headers( + stream, + "503 Service Unavailable", + &[("x-trace-id", "trace-503"), ("x-region", "us-east-1")], + "upstream unavailable", + ) + .await; + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = remote_control_auth_manager(); + let session_auth = RemoteControlAuth::capture(auth_manager.clone()).0; + let mut auth_recovery = session_auth.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + let current_enrollment = test_current_enrollment(Some(remote_control_enrollment(Some( + TEST_REMOTE_CONTROL_SERVER_TOKEN, + )))); + let (status_publisher, status_rx) = remote_control_status_channel(); + + let err = match connect_remote_control_websocket( + &remote_control_target, + Some(state_db.as_ref()), + RemoteControlAuthContext { + auth_manager: &session_auth, + auth_recovery: &mut auth_recovery, + auth_change_rx: &mut auth_change_rx, + }, + ¤t_enrollment, + RemoteControlConnectOptions { + installation_id: TEST_INSTALLATION_ID, + server_name: "test-server", + subscribe_cursor: None, + app_server_client_name: None, + desired_state_tx: &enabled_desired_state_sender(), + persistence: &RemoteControlPersistence::default(), + }, + &status_publisher, + ) + .await + { + Ok(_) => panic!("http error response should fail the websocket connect"), + Err(err) => err, + }; + + server_task.await.expect("server task should succeed"); + assert_eq!(err.to_string(), expected_error); + assert!(current_enrollment.lock().await.is_some()); + assert_eq!( + status_rx.borrow().clone(), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connecting, + server_name: "test-server".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_test".to_string()), + } + ); + } + + #[tokio::test] + async fn connect_remote_control_websocket_invalidates_unauthorized_server_token() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = remote_control_auth_manager(); + let session_auth = RemoteControlAuth::capture(auth_manager.clone()).0; + let mut auth_recovery = session_auth.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + let next_refresh_at = time::OffsetDateTime::now_utc() + time::Duration::minutes(2); + let mut enrollment = remote_control_enrollment(Some(TEST_REMOTE_CONTROL_SERVER_TOKEN)); + enrollment.next_refresh_at = Some(next_refresh_at); + let current_enrollment = test_current_enrollment(Some(enrollment)); + let (status_publisher, status_rx) = remote_control_status_channel(); + + let server_task = tokio::spawn(async move { + let (stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + "GET /backend-api/wham/remote/control/server HTTP/1.1" + ); + respond_with_status_and_headers(stream, "401 Unauthorized", &[], "unauthorized").await; + }); + + let err = connect_remote_control_websocket( + &remote_control_target, + Some(state_db.as_ref()), + RemoteControlAuthContext { + auth_manager: &session_auth, + auth_recovery: &mut auth_recovery, + auth_change_rx: &mut auth_change_rx, + }, + ¤t_enrollment, + RemoteControlConnectOptions { + installation_id: TEST_INSTALLATION_ID, + server_name: "test-server", + subscribe_cursor: None, + app_server_client_name: None, + desired_state_tx: &enabled_desired_state_sender(), + persistence: &RemoteControlPersistence::default(), + }, + &status_publisher, + ) + .await + .expect_err("unauthorized response should fail the websocket connect"); + + server_task.await.expect("server task should succeed"); + assert_eq!( + status_rx.borrow().clone(), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connecting, + server_name: "test-server".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_test".to_string()), + } + ); + assert_eq!( + err.to_string(), + "remote control websocket auth failed with HTTP 401 Unauthorized; refreshing server token before reconnect" + ); + let mut expected_enrollment = remote_control_enrollment(/*remote_control_token*/ None); + expected_enrollment.remote_control_target = remote_control_target; + expected_enrollment.next_refresh_at = Some(next_refresh_at); + assert_eq!(*current_enrollment.lock().await, Some(expected_enrollment)); + } + + #[tokio::test] + async fn connect_remote_control_websocket_recovers_after_unauthorized_enrollment() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let enroll_url = remote_control_target.enroll_url.clone(); + let server_task = tokio::spawn(async move { + let (stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_status_and_headers(stream, "401 Unauthorized", &[], "unauthorized").await; + }); + let codex_home = TempDir::new().expect("temp dir should create"); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json("stale-token"), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("stale auth should save"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let session_auth = RemoteControlAuth::capture(auth_manager.clone()).0; + let mut auth_recovery = session_auth.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + let current_enrollment = test_current_enrollment(/*enrollment*/ None); + let (status_publisher, status_rx) = remote_control_status_channel(); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json("fresh-token"), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("fresh auth should save"); + + let err = connect_remote_control_websocket( + &remote_control_target, + Some(state_db.as_ref()), + RemoteControlAuthContext { + auth_manager: &session_auth, + auth_recovery: &mut auth_recovery, + auth_change_rx: &mut auth_change_rx, + }, + ¤t_enrollment, + RemoteControlConnectOptions { + installation_id: TEST_INSTALLATION_ID, + server_name: "test-server", + subscribe_cursor: None, + app_server_client_name: None, + desired_state_tx: &enabled_desired_state_sender(), + persistence: &RemoteControlPersistence::default(), + }, + &status_publisher, + ) + .await + .expect_err("unauthorized enrollment should fail the websocket connect"); + + server_task.await.expect("server task should succeed"); + assert!( + !status_rx + .has_changed() + .expect("remote control status watch should remain open") + ); + assert_eq!( + err.to_string(), + format!( + "remote control server enrollment failed at `{enroll_url}`: HTTP 401 Unauthorized, request-id: , cf-ray: , body: unauthorized; retrying after auth recovery" + ) + ); + assert_eq!( + auth_manager + .auth() + .await + .expect("auth should remain available") + .get_token() + .expect("token should be readable"), + "fresh-token" + ); + assert!( + !auth_change_rx + .has_changed() + .expect("auth change watch should remain open"), + "recovery's own auth reload should not wake the reconnect loop" + ); + } + + #[tokio::test] + async fn connect_remote_control_websocket_recovers_after_unauthorized_refresh() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let refresh_url = remote_control_target.refresh_url.clone(); + let server_task = tokio::spawn(async move { + let (stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_status_and_headers(stream, "401 Unauthorized", &[], "unauthorized").await; + }); + let codex_home = TempDir::new().expect("temp dir should create"); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json("stale-token"), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("stale auth should save"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let session_auth = RemoteControlAuth::capture(auth_manager.clone()).0; + let mut auth_recovery = session_auth.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + let mut expected_enrollment = + remote_control_enrollment(Some(TEST_REMOTE_CONTROL_SERVER_TOKEN)); + expected_enrollment.remote_control_target = remote_control_target.clone(); + expected_enrollment.expires_at = + Some(time::OffsetDateTime::now_utc() + time::Duration::minutes(4)); + let current_enrollment = test_current_enrollment(Some(expected_enrollment.clone())); + let (status_publisher, status_rx) = remote_control_status_channel(); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json("fresh-token"), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("fresh auth should save"); + + let err = connect_remote_control_websocket( + &remote_control_target, + Some(state_db.as_ref()), + RemoteControlAuthContext { + auth_manager: &session_auth, + auth_recovery: &mut auth_recovery, + auth_change_rx: &mut auth_change_rx, + }, + ¤t_enrollment, + RemoteControlConnectOptions { + installation_id: TEST_INSTALLATION_ID, + server_name: "test-server", + subscribe_cursor: None, + app_server_client_name: None, + desired_state_tx: &enabled_desired_state_sender(), + persistence: &RemoteControlPersistence::default(), + }, + &status_publisher, + ) + .await + .expect_err("unauthorized refresh should fail the websocket connect"); + + server_task.await.expect("server task should succeed"); + assert_eq!( + status_rx.borrow().clone(), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connecting, + server_name: "test-server".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_test".to_string()), + } + ); + assert_eq!( + err.to_string(), + format!( + "remote control server refresh failed at `{refresh_url}`: HTTP 401 Unauthorized, request-id: , cf-ray: , body: unauthorized; retrying after auth recovery" + ) + ); + assert_eq!( + auth_manager + .auth() + .await + .expect("auth should remain available") + .get_token() + .expect("token should be readable"), + "fresh-token" + ); + assert_eq!(current_enrollment.snapshot(), Some(expected_enrollment)); + assert!( + !auth_change_rx + .has_changed() + .expect("auth change watch should remain open"), + "recovery's own auth reload should not wake the reconnect loop" + ); + } + + #[tokio::test] + async fn connect_remote_control_websocket_requires_sqlite_state_db() { + let remote_control_target = normalize_remote_control_url("http://127.0.0.1:9/backend-api/") + .expect("target should parse"); + let auth_manager = remote_control_auth_manager(); + let session_auth = RemoteControlAuth::capture(auth_manager.clone()).0; + let mut auth_recovery = session_auth.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + let current_enrollment = test_current_enrollment(Some(remote_control_enrollment(Some( + TEST_REMOTE_CONTROL_SERVER_TOKEN, + )))); + let (status_publisher, _status_rx) = remote_control_status_channel(); + + let err = connect_remote_control_websocket( + &remote_control_target, + /*state_db*/ None, + RemoteControlAuthContext { + auth_manager: &session_auth, + auth_recovery: &mut auth_recovery, + auth_change_rx: &mut auth_change_rx, + }, + ¤t_enrollment, + RemoteControlConnectOptions { + installation_id: TEST_INSTALLATION_ID, + server_name: "test-server", + subscribe_cursor: None, + app_server_client_name: None, + desired_state_tx: &enabled_desired_state_sender(), + persistence: &RemoteControlPersistence::default(), + }, + &status_publisher, + ) + .await + .expect_err("missing sqlite state db should fail remote control"); + + assert_eq!(err.kind(), ErrorKind::NotFound); + assert_eq!(err.to_string(), "remote control requires sqlite state db"); + assert_eq!(*current_enrollment.lock().await, None); + } + + #[tokio::test] + async fn connect_remote_control_websocket_requires_chatgpt_auth() { + let remote_control_target = normalize_remote_control_url("http://127.0.0.1:9/backend-api/") + .expect("target should parse"); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let session_auth = RemoteControlAuth::capture(auth_manager.clone()).0; + let mut auth_recovery = session_auth.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + let current_enrollment = test_current_enrollment(Some(remote_control_enrollment(Some( + TEST_REMOTE_CONTROL_SERVER_TOKEN, + )))); + let (status_publisher, mut status_rx) = remote_control_status_channel(); + status_publisher.publish_environment_id(Some("env_test".to_string())); + status_rx + .changed() + .await + .expect("remote control status watch should remain open"); + + let err = connect_remote_control_websocket( + &remote_control_target, + Some(state_db.as_ref()), + RemoteControlAuthContext { + auth_manager: &session_auth, + auth_recovery: &mut auth_recovery, + auth_change_rx: &mut auth_change_rx, + }, + ¤t_enrollment, + RemoteControlConnectOptions { + installation_id: TEST_INSTALLATION_ID, + server_name: "test-server", + subscribe_cursor: None, + app_server_client_name: None, + desired_state_tx: &enabled_desired_state_sender(), + persistence: &RemoteControlPersistence::default(), + }, + &status_publisher, + ) + .await + .expect_err("missing auth should fail remote control"); + + assert_eq!(err.kind(), ErrorKind::PermissionDenied); + assert_eq!( + err.to_string(), + "remote control requires ChatGPT authentication" + ); + assert_eq!(*current_enrollment.lock().await, None); + assert_eq!( + status_rx.borrow().clone(), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connecting, + server_name: "test-server".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: None, + } + ); + } + + #[tokio::test] + async fn run_remote_control_websocket_loop_shutdown_cancels_reconnect_backoff() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + drop(listener); + + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let (transport_event_tx, transport_event_rx) = mpsc::channel(1); + drop(transport_event_rx); + let (status_publisher, _status_rx) = remote_control_status_channel(); + let shutdown_token = CancellationToken::new(); + let (desired_state_tx, _desired_state_rx) = + watch::channel(RemoteControlDesiredState::Enabled { + persistence_preference: None, + }); + let websocket_task = tokio::spawn({ + let shutdown_token = shutdown_token.clone(); + async move { + RemoteControlWebsocket::new( + RemoteControlWebsocketConfig { + remote_control_url, + installation_id: TEST_INSTALLATION_ID.to_string(), + remote_control_target: Some(remote_control_target), + server_name: "test-server".to_string(), + }, + /*state_db*/ None, + RemoteControlAuth::capture(remote_control_auth_manager()).0, + RemoteControlChannels { + transport_event_tx, + status_publisher, + current_enrollment: test_current_enrollment(/*enrollment*/ None), + pairing_persistence_key: watch::channel(None).0, + persistence: RemoteControlPersistence::default(), + }, + shutdown_token, + Arc::new(desired_state_tx), + ) + .run(/*app_server_client_name_rx*/ None) + .await + } + }); + + tokio::time::sleep(Duration::from_millis(50)).await; + shutdown_token.cancel(); + + timeout(Duration::from_millis(100), websocket_task) + .await + .expect("shutdown should cancel reconnect backoff") + .expect("websocket task should join"); + } + + #[tokio::test] + async fn publish_status_if_changed_sends_only_status_changes() { + let (status_publisher, mut status_rx) = remote_control_status_channel(); + + status_publisher.publish_environment_id(/*environment_id*/ None); + assert!( + timeout(Duration::from_millis(20), status_rx.changed()) + .await + .is_err() + ); + + status_publisher.publish_environment_id(Some("env_first".to_string())); + status_rx + .changed() + .await + .expect("remote control status watch should remain open"); + assert_eq!( + status_rx.borrow().clone(), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connecting, + server_name: "test-server".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_first".to_string()), + } + ); + + status_publisher.publish_environment_id(Some("env_first".to_string())); + assert!( + timeout(Duration::from_millis(20), status_rx.changed()) + .await + .is_err() + ); + + status_publisher.publish_status(RemoteControlConnectionStatus::Connected); + status_rx + .changed() + .await + .expect("remote control status watch should remain open"); + assert_eq!( + status_rx.borrow().clone(), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connected, + server_name: "test-server".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: Some("env_first".to_string()), + } + ); + + status_publisher.publish_environment_id(/*environment_id*/ None); + status_rx + .changed() + .await + .expect("remote control status watch should remain open"); + assert_eq!( + status_rx.borrow().clone(), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Connected, + server_name: "test-server".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: None, + } + ); + + status_publisher.publish_environment_id(Some("env_disabled".to_string())); + status_publisher.publish_status(RemoteControlConnectionStatus::Disabled); + status_rx + .changed() + .await + .expect("remote control status watch should remain open"); + assert_eq!( + status_rx.borrow().clone(), + RemoteControlStatusChangedNotification { + status: RemoteControlConnectionStatus::Disabled, + server_name: "test-server".to_string(), + installation_id: TEST_INSTALLATION_ID.to_string(), + environment_id: None, + } + ); + + status_publisher.publish_environment_id(Some("env_disabled".to_string())); + assert!( + timeout(Duration::from_millis(20), status_rx.changed()) + .await + .is_err() + ); + } + + #[tokio::test] + async fn run_server_writer_inner_sends_periodic_ping_frames() { + let (client_stream, mut server_stream) = connected_websocket_pair().await; + let (websocket_writer, _websocket_reader) = client_stream.split(); + let (outbound_buffer, used_rx) = BoundedOutboundBuffer::new(); + let state = Arc::new(Mutex::new(WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + })); + let (_server_event_tx, server_event_rx) = mpsc::channel(super::super::CHANNEL_CAPACITY); + let server_event_rx = Arc::new(Mutex::new(server_event_rx)); + let shutdown_token = CancellationToken::new(); + let writer_task = tokio::spawn(RemoteControlWebsocket::run_server_writer_inner( + state, + server_event_rx, + used_rx, + websocket_writer, + Duration::from_millis(20), + shutdown_token.clone(), + )); + + let message = timeout(Duration::from_secs(5), server_stream.next()) + .await + .expect("ping frame should arrive in time") + .expect("server websocket should stay open") + .expect("ping frame should read"); + assert!(matches!(message, tungstenite::Message::Ping(_))); + + shutdown_token.cancel(); + writer_task + .await + .expect("writer task should join") + .expect("writer should stop cleanly"); + } + + #[tokio::test] + async fn join_connection_workers_aborts_stuck_worker_after_timeout() { + let mut join_set = tokio::task::JoinSet::new(); + join_set.spawn(futures::future::pending::<()>()); + + RemoteControlWebsocket::join_connection_workers(&mut join_set, Duration::from_millis(10)) + .await; + + assert!(join_set.is_empty()); + } + + #[tokio::test] + async fn run_server_writer_inner_assigns_contiguous_seq_ids_per_stream() { + let (client_stream, mut server_stream) = connected_websocket_pair().await; + let (websocket_writer, _websocket_reader) = client_stream.split(); + let (outbound_buffer, used_rx) = BoundedOutboundBuffer::new(); + let state = Arc::new(Mutex::new(WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + })); + let (server_event_tx, server_event_rx) = mpsc::channel(super::super::CHANNEL_CAPACITY); + let server_event_rx = Arc::new(Mutex::new(server_event_rx)); + let shutdown_token = CancellationToken::new(); + let writer_task = tokio::spawn(RemoteControlWebsocket::run_server_writer_inner( + state, + server_event_rx, + used_rx, + websocket_writer, + Duration::from_secs(60), + shutdown_token.clone(), + )); + + let client_id = ClientId("client-1".to_string()); + let first_stream = StreamId("stream-1".to_string()); + let second_stream = StreamId("stream-2".to_string()); + for stream_id in [&first_stream, &second_stream, &first_stream] { + server_event_tx + .send(super::super::QueuedServerEnvelope { + event: ServerEvent::Pong { + status: crate::transport::remote_control::protocol::PongStatus::Active, + }, + client_id: client_id.clone(), + stream_id: stream_id.clone(), + write_complete_tx: None, + }) + .await + .expect("server event should queue"); + } + + assert_eq!( + read_server_text_event(&mut server_stream).await, + serde_json::json!({ + "type": "pong", + "client_id": "client-1", + "stream_id": "stream-1", + "seq_id": 1, + "status": "active", + }) + ); + assert_eq!( + read_server_text_event(&mut server_stream).await, + serde_json::json!({ + "type": "pong", + "client_id": "client-1", + "stream_id": "stream-2", + "seq_id": 1, + "status": "active", + }) + ); + assert_eq!( + read_server_text_event(&mut server_stream).await, + serde_json::json!({ + "type": "pong", + "client_id": "client-1", + "stream_id": "stream-1", + "seq_id": 2, + "status": "active", + }) + ); + + shutdown_token.cancel(); + writer_task + .await + .expect("writer task should join") + .expect("writer should stop cleanly"); + } + + #[tokio::test] + async fn run_websocket_reader_inner_times_out_without_pong_frames() { + let (client_stream, _server_stream) = connected_websocket_pair().await; + let (_websocket_writer, websocket_reader) = client_stream.split(); + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let state = Arc::new(Mutex::new(WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + })); + let (server_event_tx, _server_event_rx) = mpsc::channel(super::super::CHANNEL_CAPACITY); + let (transport_event_tx, _transport_event_rx) = + mpsc::channel(super::super::CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let client_tracker = Arc::new(Mutex::new(ClientTracker::new( + server_event_tx, + transport_event_tx, + &shutdown_token, + ))); + + let err = timeout( + Duration::from_secs(5), + RemoteControlWebsocket::run_websocket_reader_inner( + client_tracker, + state, + websocket_reader, + Duration::from_millis(100), + shutdown_token, + ), + ) + .await + .expect("reader should time out waiting for pong") + .expect_err("missing pong should fail the websocket reader"); + + assert_eq!(err.kind(), ErrorKind::TimedOut); + assert_eq!(err.to_string(), "remote control websocket pong timeout"); + } + + #[test] + fn outbound_buffer_acks_by_stream_id() { + let (mut outbound_buffer, used_rx) = BoundedOutboundBuffer::new(); + let client_1 = ClientId("client-1".to_string()); + let client_2 = ClientId("client-2".to_string()); + let stream_1 = StreamId("stream-1".to_string()); + + outbound_buffer.insert(&server_envelope( + &client_1, + "stream-1", + /*seq_id*/ 1, + "first-client-old-stream", + )); + outbound_buffer.insert(&server_envelope( + &client_2, + "stream-1", + /*seq_id*/ 2, + "second-client", + )); + outbound_buffer.insert(&server_envelope( + &client_1, + "stream-2", + /*seq_id*/ 3, + "first-client-new-stream", + )); + + outbound_buffer.ack( + &client_1, &stream_1, /*acked_seq_id*/ 3, /*acked_segment_id*/ None, + ); + + let mut retained = outbound_buffer + .server_envelopes() + .map(|server_envelope| { + ( + server_envelope.client_id.0.as_str(), + server_envelope.stream_id.0.as_str(), + server_envelope.seq_id, + ) + }) + .collect::>(); + retained.sort_unstable(); + assert_eq!( + retained, + vec![("client-1", "stream-2", 3), ("client-2", "stream-1", 2)] + ); + assert_eq!(*used_rx.borrow(), 2); + } + + #[test] + fn outbound_buffer_retains_unacked_messages_until_ack_advances() { + let (mut outbound_buffer, used_rx) = BoundedOutboundBuffer::new(); + let client_1 = ClientId("client-1".to_string()); + let client_2 = ClientId("client-2".to_string()); + let stream_1 = StreamId("stream-1".to_string()); + + outbound_buffer.insert(&server_envelope( + &client_1, + "stream-1", + /*seq_id*/ 1, + "first-old", + )); + outbound_buffer.insert(&server_envelope( + &client_1, + "stream-2", + /*seq_id*/ 2, + "first-new", + )); + outbound_buffer.insert(&server_envelope( + &client_2, "stream-1", /*seq_id*/ 3, "second", + )); + + outbound_buffer.ack( + &client_1, &stream_1, /*acked_seq_id*/ 1, /*acked_segment_id*/ None, + ); + + let mut retained = outbound_buffer + .server_envelopes() + .map(|server_envelope| { + ( + server_envelope.client_id.0.as_str(), + server_envelope.stream_id.0.as_str(), + server_envelope.seq_id, + ) + }) + .collect::>(); + retained.sort_unstable(); + assert_eq!( + retained, + vec![("client-1", "stream-2", 2), ("client-2", "stream-1", 3)] + ); + assert_eq!(*used_rx.borrow(), 2); + } + + #[test] + fn outbound_buffer_advances_segmented_acks_by_wire_cursor() { + let (mut outbound_buffer, used_rx) = BoundedOutboundBuffer::new(); + let client_id = ClientId("client-1".to_string()); + let stream_id = StreamId("stream-1".to_string()); + + outbound_buffer.insert(&server_chunk_envelope( + &client_id, "stream-1", /*seq_id*/ 4, /*segment_id*/ 0, + )); + outbound_buffer.insert(&server_chunk_envelope( + &client_id, "stream-1", /*seq_id*/ 4, /*segment_id*/ 1, + )); + + outbound_buffer.ack( + &client_id, + &stream_id, + /*acked_seq_id*/ 4, + /*acked_segment_id*/ Some(1), + ); + + let retained = outbound_buffer + .server_envelopes() + .map(|server_envelope| server_envelope.event.segment_id()) + .collect::>(); + assert_eq!(retained, Vec::>::new()); + assert_eq!(*used_rx.borrow(), 0); + } + + #[test] + fn outbound_buffer_treats_segmentless_acks_as_seq_level_acks() { + let (mut outbound_buffer, used_rx) = BoundedOutboundBuffer::new(); + let client_id = ClientId("client-1".to_string()); + let stream_id = StreamId("stream-1".to_string()); + + outbound_buffer.insert(&server_chunk_envelope( + &client_id, "stream-1", /*seq_id*/ 4, /*segment_id*/ 0, + )); + outbound_buffer.insert(&server_chunk_envelope( + &client_id, "stream-1", /*seq_id*/ 4, /*segment_id*/ 1, + )); + + outbound_buffer.ack( + &client_id, &stream_id, /*acked_seq_id*/ 4, /*acked_segment_id*/ None, + ); + + let retained = outbound_buffer + .server_envelopes() + .map(|server_envelope| server_envelope.event.segment_id()) + .collect::>(); + assert_eq!(retained, Vec::>::new()); + assert_eq!(*used_rx.borrow(), 0); + } + + #[test] + fn websocket_state_drops_duplicate_client_chunks_while_pending() { + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let mut state = WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + }; + let first_chunk = client_chunk_envelope( + "client-1", "stream-1", /*seq_id*/ 4, /*segment_id*/ 0, + /*segment_count*/ 2, /*message_size_bytes*/ 2, b"x", + ); + let second_chunk = client_chunk_envelope( + "client-1", "stream-1", /*seq_id*/ 4, /*segment_id*/ 1, + /*segment_count*/ 2, /*message_size_bytes*/ 2, b"y", + ); + + assert!(matches!( + observe_client_message(&mut state, first_chunk.clone()), + ClientSegmentObservation::Pending + )); + assert!(matches!( + observe_client_message(&mut state, first_chunk.clone()), + ClientSegmentObservation::Dropped + )); + assert!(matches!( + observe_client_message(&mut state, second_chunk), + ClientSegmentObservation::Dropped + )); + assert!(matches!( + observe_client_message(&mut state, first_chunk), + ClientSegmentObservation::Pending + )); + } + + #[test] + fn websocket_state_drops_replayed_client_chunks_after_completion() { + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let mut state = WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + }; + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let first_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 4, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + ); + let second_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 4, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + ); + + assert!(matches!( + observe_client_message(&mut state, first_chunk.clone()), + ClientSegmentObservation::Pending + )); + let completed_envelope = match observe_client_message(&mut state, second_chunk) { + ClientSegmentObservation::Forward(client_envelope) => *client_envelope, + _ => panic!("expected completed client message"), + }; + state.record_client_message_delivery( + &completed_envelope, + Some(( + ( + ClientId("client-1".to_string()), + Some(StreamId("stream-1".to_string())), + ), + 4, + )), + ); + assert!(matches!( + observe_client_message(&mut state, first_chunk), + ClientSegmentObservation::Dropped + )); + } + + #[test] + fn websocket_state_allows_replay_before_completed_chunk_delivery() { + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let mut state = WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + }; + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let first_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 4, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + ); + let second_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 4, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + ); + + assert!(matches!( + observe_client_message(&mut state, first_chunk.clone()), + ClientSegmentObservation::Pending + )); + assert!(matches!( + observe_client_message(&mut state, second_chunk), + ClientSegmentObservation::Forward(_) + )); + assert!(matches!( + observe_client_message(&mut state, first_chunk), + ClientSegmentObservation::Pending + )); + } + + #[test] + fn websocket_state_allows_replay_after_rejected_out_of_order_chunk() { + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let mut state = WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + }; + let first_chunk = client_chunk_envelope( + "client-1", "stream-1", /*seq_id*/ 4, /*segment_id*/ 0, + /*segment_count*/ 2, /*message_size_bytes*/ 2, b"x", + ); + let second_chunk = client_chunk_envelope( + "client-1", "stream-1", /*seq_id*/ 4, /*segment_id*/ 1, + /*segment_count*/ 2, /*message_size_bytes*/ 2, b"y", + ); + + assert!(matches!( + observe_client_message(&mut state, second_chunk), + ClientSegmentObservation::Dropped + )); + assert!(matches!( + observe_client_message(&mut state, first_chunk), + ClientSegmentObservation::Pending + )); + } + + #[test] + fn websocket_state_allows_replay_after_later_chunk_drops() { + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let mut state = WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + }; + let first_chunk = client_chunk_envelope( + "client-1", "stream-1", /*seq_id*/ 4, /*segment_id*/ 0, + /*segment_count*/ 2, /*message_size_bytes*/ 2, b"x", + ); + let invalid_second_chunk = client_chunk_envelope( + "client-1", "stream-1", /*seq_id*/ 4, /*segment_id*/ 1, + /*segment_count*/ 2, /*message_size_bytes*/ 2, b"", + ); + + assert!(matches!( + observe_client_message(&mut state, first_chunk.clone()), + ClientSegmentObservation::Pending + )); + assert!(matches!( + observe_client_message(&mut state, invalid_second_chunk), + ClientSegmentObservation::Dropped + )); + assert!(matches!( + observe_client_message(&mut state, first_chunk), + ClientSegmentObservation::Pending + )); + } + + #[test] + fn websocket_state_drops_oversized_client_chunk_frames() { + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let mut state = WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + }; + let chunk = client_chunk_envelope( + "client-1", "stream-1", /*seq_id*/ 4, /*segment_id*/ 0, + /*segment_count*/ 1, /*message_size_bytes*/ 1, b"x", + ); + + assert!(matches!( + state.observe_client_message(chunk, REMOTE_CONTROL_SEGMENT_MAX_BYTES + 1), + ClientSegmentObservation::Dropped + )); + } + + #[test] + fn websocket_state_ignores_oversized_stale_chunks_without_dropping_newer_assembly() { + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let mut state = WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + }; + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let first_newer_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 8, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + ); + let oversized_stale_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 7, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + ); + let second_newer_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 8, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + ); + + assert!(matches!( + observe_client_message(&mut state, first_newer_chunk), + ClientSegmentObservation::Pending + )); + assert!(matches!( + state.observe_client_message( + oversized_stale_chunk, + REMOTE_CONTROL_SEGMENT_MAX_BYTES + 1, + ), + ClientSegmentObservation::Dropped + )); + assert!(matches!( + observe_client_message(&mut state, second_newer_chunk), + ClientSegmentObservation::Forward(_) + )); + } + + #[test] + fn websocket_state_ignores_oversized_duplicate_chunks_without_dropping_current_assembly() { + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let mut state = WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + }; + let message = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + let raw = serde_json::to_vec(&message).expect("message should serialize"); + let split = raw.len() / 2; + let first_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 8, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + ); + let oversized_duplicate_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 8, + /*segment_id*/ 0, + /*segment_count*/ 2, + raw.len(), + &raw[..split], + ); + let second_chunk = client_chunk_envelope( + "client-1", + "stream-1", + /*seq_id*/ 8, + /*segment_id*/ 1, + /*segment_count*/ 2, + raw.len(), + &raw[split..], + ); + + assert!(matches!( + observe_client_message(&mut state, first_chunk), + ClientSegmentObservation::Pending + )); + assert!(matches!( + state.observe_client_message( + oversized_duplicate_chunk, + REMOTE_CONTROL_SEGMENT_MAX_BYTES + 1, + ), + ClientSegmentObservation::Dropped + )); + assert!(matches!( + observe_client_message(&mut state, second_chunk), + ClientSegmentObservation::Forward(_) + )); + } + + #[test] + fn websocket_state_clears_chunk_cursor_when_stream_is_invalidated() { + let (outbound_buffer, _used_rx) = BoundedOutboundBuffer::new(); + let mut state = WebsocketState { + outbound_buffer, + subscribe_cursor: None, + next_seq_id_by_stream: HashMap::new(), + last_completed_client_chunk_seq_id_by_stream: HashMap::new(), + client_segment_reassembler: ClientSegmentReassembler::default(), + }; + let client_id = ClientId("client-1".to_string()); + let stream_id = StreamId("stream-1".to_string()); + + assert!(matches!( + observe_client_message( + &mut state, + client_chunk_envelope( + "client-1", "stream-1", /*seq_id*/ 4, /*segment_id*/ 0, + /*segment_count*/ 2, /*message_size_bytes*/ 2, b"x", + ) + ), + ClientSegmentObservation::Pending + )); + state.invalidate_client_message_stream(&client_id, &stream_id); + state + .client_segment_reassembler + .invalidate_stream(&client_id, &stream_id); + + assert!(matches!( + observe_client_message( + &mut state, + client_chunk_envelope( + "client-1", "stream-1", /*seq_id*/ 1, /*segment_id*/ 0, + /*segment_count*/ 2, /*message_size_bytes*/ 2, b"x", + ) + ), + ClientSegmentObservation::Pending + )); + } + + fn server_envelope( + client_id: &ClientId, + stream_id: &str, + seq_id: u64, + summary: &str, + ) -> ServerEnvelope { + ServerEnvelope { + event: ServerEvent::ServerMessage { + message: Box::new(OutgoingMessage::AppServerNotification( + ServerNotificationEnvelope { + notification: ServerNotification::ConfigWarning( + ConfigWarningNotification { + summary: summary.to_string(), + details: None, + path: None, + range: None, + }, + ), + emitted_at_ms: Some(1_234), + }, + )), + }, + client_id: client_id.clone(), + stream_id: StreamId(stream_id.to_string()), + seq_id, + } + } + + fn server_chunk_envelope( + client_id: &ClientId, + stream_id: &str, + seq_id: u64, + segment_id: usize, + ) -> ServerEnvelope { + ServerEnvelope { + event: ServerEvent::ServerMessageChunk { + segment_id, + segment_count: 2, + message_size_bytes: 2, + message_chunk_base64: String::new(), + }, + client_id: client_id.clone(), + stream_id: StreamId(stream_id.to_string()), + seq_id, + } + } + + fn client_chunk_envelope( + client_id: &str, + stream_id: &str, + seq_id: u64, + segment_id: usize, + segment_count: usize, + message_size_bytes: usize, + chunk: &[u8], + ) -> ClientEnvelope { + ClientEnvelope { + event: ClientEvent::ClientMessageChunk { + segment_id, + segment_count, + message_size_bytes, + message_chunk_base64: base64::engine::general_purpose::STANDARD.encode(chunk), + }, + client_id: ClientId(client_id.to_string()), + stream_id: Some(StreamId(stream_id.to_string())), + seq_id: Some(seq_id), + cursor: None, + } + } + + fn observe_client_message( + state: &mut WebsocketState, + envelope: ClientEnvelope, + ) -> ClientSegmentObservation { + let wire_size_bytes = serde_json::to_vec(&envelope) + .expect("client envelope should serialize") + .len(); + state.observe_client_message(envelope, wire_size_bytes) + } + + pub(super) async fn accept_http_request(listener: &TcpListener) -> (TcpStream, String) { + let (stream, _) = timeout(TEST_HTTP_ACCEPT_TIMEOUT, listener.accept()) + .await + .expect("HTTP request should arrive in time") + .expect("listener accept should succeed"); + let mut reader = BufReader::new(stream); + + let mut request_line = String::new(); + reader + .read_line(&mut request_line) + .await + .expect("request line should read"); + loop { + let mut line = String::new(); + reader + .read_line(&mut line) + .await + .expect("header line should read"); + if line == "\r\n" { + break; + } + } + + ( + reader.into_inner(), + request_line.trim_end_matches("\r\n").to_string(), + ) + } + + async fn connected_websocket_pair() -> ( + WebSocketStream>, + WebSocketStream, + ) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let connect_task = tokio::spawn(connect_async(format!( + "ws://{}", + listener + .local_addr() + .expect("listener should have a local addr") + ))); + let (server_stream, _) = listener + .accept() + .await + .expect("server should accept client"); + let server_stream = accept_async(server_stream) + .await + .expect("server websocket handshake should succeed"); + let (client_stream, _) = connect_task + .await + .expect("client connect task should join") + .expect("client websocket handshake should succeed"); + + (client_stream, server_stream) + } + + async fn read_server_text_event( + server_stream: &mut WebSocketStream, + ) -> serde_json::Value { + let message = timeout(Duration::from_secs(5), server_stream.next()) + .await + .expect("server event should arrive in time") + .expect("server websocket should stay open") + .expect("server event should read"); + let tungstenite::Message::Text(text) = message else { + panic!("expected text event, got {message:?}"); + }; + serde_json::from_str(text.as_ref()).expect("server event should deserialize") + } + + pub(super) async fn respond_with_status_and_headers( + mut stream: TcpStream, + status: &str, + headers: &[(&str, &str)], + body: &str, + ) { + let extra_headers = headers + .iter() + .map(|(name, value)| format!("{name}: {value}\r\n")) + .collect::(); + let response = format!( + "HTTP/1.1 {status}\r\ncontent-type: text/plain\r\ncontent-length: {}\r\nconnection: close\r\n{extra_headers}\r\n{body}", + body.len(), + ); + stream + .write_all(response.as_bytes()) + .await + .expect("response should write"); + stream.flush().await.expect("response should flush"); + } +} diff --git a/codex-rs/app-server-transport/src/transport/remote_control/websocket_refresh_tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/websocket_refresh_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..07abe80d17468799926a6b43efba497246d242b0 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/remote_control/websocket_refresh_tests.rs @@ -0,0 +1,768 @@ +use super::tests::TEST_HTTP_ACCEPT_TIMEOUT; +use super::tests::TEST_INSTALLATION_ID; +use super::tests::TEST_REMOTE_CONTROL_SERVER_TOKEN; +use super::tests::accept_http_request; +use super::tests::enabled_desired_state_sender; +use super::tests::remote_control_auth_dot_json; +use super::tests::remote_control_auth_manager; +use super::tests::remote_control_enrollment; +use super::tests::remote_control_state_runtime; +use super::tests::remote_control_status_channel; +use super::tests::remote_control_url_for_listener; +use super::tests::respond_with_status_and_headers; +use super::tests::test_current_enrollment; +use super::*; +use crate::transport::remote_control::protocol::normalize_remote_control_url; +use crate::transport::remote_control::server_api::remote_control_retry_at; +use crate::transport::remote_control::tests::remote_control_handle_with_current_enrollment; +use codex_app_server_protocol::RemoteControlPairingStartParams; +use codex_app_server_protocol::RemoteControlPairingStatusParams; +use codex_config::types::AuthCredentialsStoreMode; +use codex_login::AuthKeyringBackendKind; +use codex_login::AuthManager; +use codex_login::save_auth; +use pretty_assertions::assert_eq; +use tempfile::TempDir; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::time::Duration; +use tokio::time::timeout; +use tokio_tungstenite::WebSocketStream; +use tokio_tungstenite::accept_async; + +async fn connect_test_websocket( + remote_control_target: &RemoteControlTarget, + state_db: &StateRuntime, + auth_manager: &Arc, + current_enrollment: &CurrentRemoteControlEnrollment, +) -> io::Result<()> { + let session_auth = RemoteControlAuth::capture(auth_manager.clone()).0; + let mut auth_recovery = session_auth.unauthorized_recovery(); + let mut auth_change_rx = auth_manager.auth_change_receiver(); + let (status_publisher, _) = remote_control_status_channel(); + let desired_state_tx = enabled_desired_state_sender(); + let persistence = RemoteControlPersistence::default(); + connect_remote_control_websocket( + remote_control_target, + Some(state_db), + RemoteControlAuthContext { + auth_manager: &session_auth, + auth_recovery: &mut auth_recovery, + auth_change_rx: &mut auth_change_rx, + }, + current_enrollment, + RemoteControlConnectOptions { + installation_id: TEST_INSTALLATION_ID, + server_name: "test-server", + subscribe_cursor: None, + app_server_client_name: None, + desired_state_tx: &desired_state_tx, + persistence: &persistence, + }, + &status_publisher, + ) + .await + .map(|_| ()) +} + +#[tokio::test] +async fn proactive_refresh_failure_uses_valid_token_for_websocket_connect() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let server_task = tokio::spawn(async move { + let (stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_status_and_headers(stream, "502 Bad Gateway", &[], "upstream unavailable") + .await; + accept_test_websocket(&listener).await + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = remote_control_auth_manager(); + let mut enrollment = remote_control_enrollment(Some(TEST_REMOTE_CONTROL_SERVER_TOKEN)); + enrollment.expires_at = Some(time::OffsetDateTime::now_utc() + time::Duration::minutes(4)); + let current_enrollment = test_current_enrollment(Some(enrollment)); + + let refresh_started_at = time::OffsetDateTime::now_utc(); + connect_test_websocket( + &remote_control_target, + state_db.as_ref(), + &auth_manager, + ¤t_enrollment, + ) + .await + .expect("valid token should allow websocket connect after proactive refresh failure"); + let refresh_completed_at = time::OffsetDateTime::now_utc(); + let server_websocket = server_task.await.expect("server task should succeed"); + + let enrollment = current_enrollment + .lock() + .await + .clone() + .expect("enrollment should remain available"); + assert_eq!( + enrollment.remote_control_token.as_deref(), + Some(TEST_REMOTE_CONTROL_SERVER_TOKEN) + ); + let next_refresh_at = enrollment + .next_refresh_at + .expect("transient refresh should set a retry deadline"); + assert!( + (refresh_started_at + time::Duration::seconds(24) + ..=refresh_completed_at + time::Duration::seconds(36)) + .contains(&next_refresh_at) + ); + drop(server_websocket); +} + +#[tokio::test] +async fn proactive_refresh_connection_failure_uses_valid_token_for_websocket_connect() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let server_task = tokio::spawn(async move { + let (stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + drop(stream); + accept_test_websocket(&listener).await + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = remote_control_auth_manager(); + let mut enrollment = remote_control_enrollment(Some(TEST_REMOTE_CONTROL_SERVER_TOKEN)); + enrollment.expires_at = Some(time::OffsetDateTime::now_utc() + time::Duration::minutes(4)); + let current_enrollment = test_current_enrollment(Some(enrollment)); + + connect_test_websocket( + &remote_control_target, + state_db.as_ref(), + &auth_manager, + ¤t_enrollment, + ) + .await + .expect("valid token should allow websocket connect after refresh connection failure"); + let server_websocket = server_task.await.expect("server task should succeed"); + + assert!( + current_enrollment + .snapshot() + .and_then(|enrollment| enrollment.next_refresh_at) + .is_some(), + "connection failure should set a retry deadline" + ); + drop(server_websocket); +} + +#[tokio::test] +async fn websocket_retry_after_throttles_pairing_refresh() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let server_task = tokio::spawn(async move { + let (stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_status_and_headers( + stream, + "502 Bad Gateway", + &[("retry-after", "120")], + "upstream unavailable", + ) + .await; + listener + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = remote_control_auth_manager(); + let mut remote_handle = + remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager.clone()); + remote_handle.state_db = Some(state_db.clone()); + remote_handle + .current_enrollment + .lock() + .await + .as_mut() + .expect("current enrollment should exist") + .expires_at = Some(time::OffsetDateTime::now_utc() + time::Duration::minutes(4)); + let current_enrollment = remote_handle.current_enrollment.clone(); + let refresh_started_at = time::OffsetDateTime::now_utc(); + let refresh_error = connect_test_websocket( + &remote_control_target, + state_db.as_ref(), + &auth_manager, + ¤t_enrollment, + ) + .await + .expect_err("an explicit server deadline must defer the handshake even with a valid token"); + let refresh_completed_at = time::OffsetDateTime::now_utc(); + let next_refresh_at = current_enrollment + .snapshot() + .and_then(|enrollment| enrollment.next_refresh_at) + .expect("Retry-After should set a retry deadline"); + assert!( + (refresh_started_at + time::Duration::seconds(120) + ..=refresh_completed_at + time::Duration::seconds(150)) + .contains(&next_refresh_at) + ); + + let pairing_error = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("the refresh deadline must also defer pairing"); + let listener = server_task.await.expect("server task should succeed"); + assert_eq!( + remote_control_retry_at(&refresh_error), + Some(next_refresh_at) + ); + assert_eq!( + remote_control_retry_at(&pairing_error), + Some(next_refresh_at) + ); + assert_eq!( + current_enrollment + .snapshot() + .and_then(|enrollment| enrollment.remote_control_token), + Some(TEST_REMOTE_CONTROL_SERVER_TOKEN.to_string()) + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("no handshake or pairing should bypass the proactive refresh deadline"); +} + +#[tokio::test] +async fn pairing_http_date_retry_after_throttles_websocket_refresh() { + for status in [ + "429 Too Many Requests", + "503 Service Unavailable", + "502 Bad Gateway", + ] { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let retry_after = + httpdate::fmt_http_date(std::time::SystemTime::now() + Duration::from_secs(120)); + let expected_next_refresh_at = time::OffsetDateTime::from( + httpdate::parse_http_date(&retry_after).expect("Retry-After date should parse"), + ); + let server_task = tokio::spawn(async move { + let (refresh_stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + respond_with_status_and_headers( + refresh_stream, + status, + &[("retry-after", &retry_after)], + "upstream unavailable", + ) + .await; + listener + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = remote_control_auth_manager(); + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + auth_manager.clone(), + ); + remote_handle.state_db = Some(state_db.clone()); + remote_handle + .current_enrollment + .lock() + .await + .as_mut() + .expect("current enrollment should exist") + .expires_at = Some(time::OffsetDateTime::now_utc() + time::Duration::minutes(4)); + let current_enrollment = remote_handle.current_enrollment.clone(); + + let refresh_error = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("an explicit server deadline must defer pairing even with a valid token"); + let listener = server_task.await.expect("server task should succeed"); + let retry_at = remote_control_retry_at(&refresh_error) + .expect("the proactive refresh response should preserve its deadline"); + let enrollment = current_enrollment + .snapshot() + .expect("enrollment should remain"); + assert_eq!(enrollment.next_refresh_at, Some(retry_at)); + assert_eq!( + enrollment.remote_control_token.as_deref(), + Some(TEST_REMOTE_CONTROL_SERVER_TOKEN) + ); + assert!( + (expected_next_refresh_at..=expected_next_refresh_at + time::Duration::seconds(30)) + .contains(&retry_at) + ); + let connect_error = connect_test_websocket( + &remote_control_target, + state_db.as_ref(), + &auth_manager, + ¤t_enrollment, + ) + .await + .expect_err("the refresh deadline must also defer the handshake"); + assert_eq!(remote_control_retry_at(&connect_error), Some(retry_at)); + let pairing_error = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("another pairing request must retain the same deadline"); + assert_eq!(remote_control_retry_at(&pairing_error), Some(retry_at)); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("no pairing, refresh or handshake should bypass the deadline"); + } +} + +#[tokio::test] +async fn pairing_during_pending_handshake_respects_later_overload() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = remote_control_auth_manager(); + let mut remote_handle = + remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager.clone()); + remote_handle.state_db = Some(state_db.clone()); + let current_enrollment = remote_handle.current_enrollment.clone(); + let connect = connect_test_websocket( + &remote_control_target, + state_db.as_ref(), + &auth_manager, + ¤t_enrollment, + ); + tokio::pin!(connect); + let (handshake_stream, request_line) = tokio::select! { + result = &mut connect => panic!("handshake should wait for its response: {result:?}"), + request = accept_http_request(&listener) => request, + }; + assert_eq!( + request_line, + "GET /backend-api/wham/remote/control/server HTTP/1.1" + ); + + let pairing = remote_handle.start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ); + tokio::pin!(pairing); + let (pairing_stream, request_line) = tokio::select! { + result = &mut pairing => panic!("pairing should wait for its response: {result:?}"), + request = accept_http_request(&listener) => request, + }; + assert_eq!( + request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + respond_with_status_and_headers( + pairing_stream, + "200 OK", + &[], + r#"{"pairing_code":"pairing-code","manual_pairing_code":"ABCD-EFGH","server_id":"srv_e_test","environment_id":"env_test","expires_at":"3026-05-22T12:34:56Z"}"#, + ) + .await; + timeout(Duration::from_secs(1), pairing) + .await + .expect("pairing must stay responsive while the handshake is pending") + .expect("pairing should succeed before any overload is observed"); + respond_with_status_and_headers( + handshake_stream, + "503 Service Unavailable", + &[("Retry-After", "120")], + "overloaded", + ) + .await; + let connect_error = connect + .await + .expect_err("the handshake should report overload"); + let retry_at = remote_control_retry_at(&connect_error) + .expect("the handshake should preserve its retry deadline"); + let pairing_error = timeout( + Duration::from_secs(1), + remote_handle.start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ), + ) + .await + .expect("new pairing should report the deadline promptly") + .expect_err("new pairing must honor the handshake deadline"); + assert_eq!(remote_control_retry_at(&pairing_error), Some(retry_at)); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("new pairing must not bypass the newly recorded deadline"); +} + +#[tokio::test] +async fn pairing_auth_recovery_respects_concurrent_handshake_overload() { + for (pairing_status, auth_endpoint) in + [("401 Unauthorized", "refresh"), ("404 Not Found", "enroll")] + { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let codex_home = TempDir::new().expect("temp dir should create"); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json("stale-token"), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("stale auth should save"); + let auth_manager = AuthManager::shared( + codex_home.path().to_path_buf(), + /*enable_codex_api_key_env*/ false, + AuthCredentialsStoreMode::File, + /*forced_chatgpt_workspace_id*/ None, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + codex_login::test_support::transport_default_auth_route_config(), + ) + .await; + let state_db = remote_control_state_runtime(&codex_home).await; + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + auth_manager.clone(), + ); + remote_handle.state_db = Some(state_db.clone()); + let current_enrollment = remote_handle.current_enrollment.clone(); + let connect = connect_test_websocket( + &remote_control_target, + state_db.as_ref(), + &auth_manager, + ¤t_enrollment, + ); + tokio::pin!(connect); + let (handshake_stream, request_line) = tokio::select! { + result = &mut connect => panic!("handshake should wait for its response: {result:?}"), + request = accept_http_request(&listener) => request, + }; + assert_eq!( + request_line, + "GET /backend-api/wham/remote/control/server HTTP/1.1" + ); + + let pairing = remote_handle.start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ); + tokio::pin!(pairing); + let (pairing_stream, request_line) = tokio::select! { + result = &mut pairing => panic!("pairing should wait for its response: {result:?}"), + request = accept_http_request(&listener) => request, + }; + assert_eq!( + request_line, + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + ); + respond_with_status_and_headers(pairing_stream, pairing_status, &[], "retry auth").await; + let (auth_stream, request_line) = tokio::select! { + result = &mut pairing => panic!("auth should wait for its response: {result:?}"), + request = accept_http_request(&listener) => request, + }; + assert_eq!( + request_line, + format!("POST /backend-api/wham/remote/control/server/{auth_endpoint} HTTP/1.1") + ); + + respond_with_status_and_headers( + handshake_stream, + "503 Service Unavailable", + &[("Retry-After", "120")], + "overloaded", + ) + .await; + let connect_error = connect + .await + .expect_err("the handshake should report overload"); + let retry_at = remote_control_retry_at(&connect_error) + .expect("the handshake should preserve its retry deadline"); + save_auth( + codex_home.path(), + &remote_control_auth_dot_json("fresh-token"), + AuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .expect("replacement auth should save"); + respond_with_status_and_headers(auth_stream, "401 Unauthorized", &[], "stale auth").await; + let pairing_error = timeout(TEST_HTTP_ACCEPT_TIMEOUT, pairing) + .await + .expect("auth recovery should report the deadline promptly") + .expect_err("auth recovery must honor the handshake deadline"); + assert_eq!(remote_control_retry_at(&pairing_error), Some(retry_at)); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("auth recovery must not send another request during overload"); + } +} + +#[tokio::test] +async fn pairing_overload_defers_pairing_and_websocket_requests() { + for status in ["429 Too Many Requests", "503 Service Unavailable"] { + for check_status in [false, true] { + for incomplete_body in [false, true] { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let server_task = tokio::spawn(async move { + let (mut stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + if check_status { + "POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1" + } else { + "POST /backend-api/wham/remote/control/server/pair HTTP/1.1" + } + ); + if incomplete_body { + stream.write_all(format!( + "HTTP/1.1 {status}\r\nContent-Length: 100\r\nRetry-After: 120\r\nConnection: close\r\n\r\npartial" + ).as_bytes()).await.expect("partial response should send"); + stream.shutdown().await.expect("response should close"); + } else { + respond_with_status_and_headers( + stream, + status, + &[("Retry-After", "120")], + "overloaded", + ) + .await; + } + listener + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = remote_control_auth_manager(); + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + auth_manager.clone(), + ); + remote_handle.state_db = Some(state_db.clone()); + let current_enrollment = remote_handle.current_enrollment.clone(); + let status_params = || RemoteControlPairingStatusParams { + pairing_code: Some("pairing-code".to_string()), + manual_pairing_code: None, + }; + let response = if check_status { + remote_handle + .pairing_status(status_params()) + .await + .map(|_| ()) + } else { + remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .map(|_| ()) + }; + let error = response.expect_err("pairing should report overload"); + let listener = server_task.await.expect("server task should succeed"); + let retry_at = remote_control_retry_at(&error) + .expect("even a partial response must preserve the retry deadline"); + let pairing_error = remote_handle + .start_pairing( + RemoteControlPairingStartParams::default(), + /*app_server_client_name*/ None, + ) + .await + .expect_err("pairing must wait for the server deadline"); + let status_error = remote_handle + .pairing_status(status_params()) + .await + .expect_err("pairing status must wait for the server deadline"); + let connect_error = connect_test_websocket( + &remote_control_target, + state_db.as_ref(), + &auth_manager, + ¤t_enrollment, + ) + .await + .expect_err("the handshake must wait for the same deadline"); + for error in [pairing_error, status_error, connect_error] { + assert_eq!(remote_control_retry_at(&error), Some(retry_at)); + } + assert_eq!( + current_enrollment + .snapshot() + .and_then(|enrollment| enrollment.remote_control_token), + Some(TEST_REMOTE_CONTROL_SERVER_TOKEN.to_string()), + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("no remote-control request should bypass the pairing deadline"); + } + } + } +} + +async fn assert_refresh_failure_blocks_websocket( + expires_in: time::Duration, + response_delay: Duration, +) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = + normalize_remote_control_url(&remote_control_url).expect("target should parse"); + let (connects_done_tx, connects_done_rx) = oneshot::channel(); + let server_task = tokio::spawn(async move { + let (stream, request_line) = accept_http_request(&listener).await; + assert_eq!( + request_line, + "POST /backend-api/wham/remote/control/server/refresh HTTP/1.1" + ); + tokio::time::sleep(response_delay).await; + respond_with_status_and_headers( + stream, + "502 Bad Gateway", + &[("retry-after", "120")], + "upstream unavailable", + ) + .await; + assert_no_connection_until_connect_finishes(&listener, connects_done_rx).await; + }); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let auth_manager = remote_control_auth_manager(); + let mut enrollment = remote_control_enrollment(Some(TEST_REMOTE_CONTROL_SERVER_TOKEN)); + enrollment.expires_at = Some(time::OffsetDateTime::now_utc() + expires_in); + let current_enrollment = test_current_enrollment(Some(enrollment)); + + let refresh_started_at = time::OffsetDateTime::now_utc(); + let refresh_err = connect_test_websocket( + &remote_control_target, + state_db.as_ref(), + &auth_manager, + ¤t_enrollment, + ) + .await + .expect_err("required refresh failure should block websocket connect"); + let refresh_completed_at = time::OffsetDateTime::now_utc(); + let deferred_err = connect_test_websocket( + &remote_control_target, + state_db.as_ref(), + &auth_manager, + ¤t_enrollment, + ) + .await + .expect_err("required refresh deadline should block websocket reconnect"); + connects_done_tx + .send(()) + .expect("server should wait for connect attempts to finish"); + + server_task.await.expect("server task should succeed"); + assert!(refresh_err.to_string().contains("HTTP 502 Bad Gateway")); + assert_eq!(deferred_err.kind(), io::ErrorKind::WouldBlock); + let retry_at = remote_control_retry_at(&refresh_err) + .expect("the overload response should retain its retry deadline"); + assert_eq!(remote_control_retry_at(&deferred_err), Some(retry_at)); + let next_refresh_at = current_enrollment + .snapshot() + .and_then(|enrollment| enrollment.next_refresh_at) + .expect("required refresh failure should set a retry deadline"); + assert!( + (refresh_started_at + time::Duration::seconds(120) + ..=refresh_completed_at + time::Duration::seconds(150)) + .contains(&next_refresh_at) + ); +} + +#[tokio::test] +async fn expired_token_refresh_failure_throttles_reconnect_without_websocket() { + assert_refresh_failure_blocks_websocket(-time::Duration::seconds(1), Duration::ZERO).await; +} + +#[tokio::test] +async fn token_expiring_during_refresh_failure_throttles_reconnect_without_websocket() { + assert_refresh_failure_blocks_websocket( + time::Duration::seconds(1), + Duration::from_millis(1_200), + ) + .await; +} + +#[tokio::test] +async fn websocket_auth_failure_does_not_clear_rotated_server_token() { + let attempted_enrollment = remote_control_enrollment(Some("old-token")); + let mut rotated_enrollment = attempted_enrollment.clone(); + rotated_enrollment.remote_control_token = Some("new-token".to_string()); + rotated_enrollment.expires_at = + Some(time::OffsetDateTime::now_utc() + time::Duration::hours(1)); + let current_enrollment = test_current_enrollment(Some(rotated_enrollment.clone())); + + clear_remote_control_server_token_if_matches(¤t_enrollment, &attempted_enrollment) + .await + .expect("matching enrollment identity should remain available"); + + assert_eq!(current_enrollment.snapshot(), Some(rotated_enrollment)); +} + +async fn accept_test_websocket(listener: &TcpListener) -> WebSocketStream { + let (stream, _) = timeout(TEST_HTTP_ACCEPT_TIMEOUT, listener.accept()) + .await + .expect("websocket request should arrive in time") + .expect("listener accept should succeed"); + accept_async(stream) + .await + .expect("websocket handshake should succeed") +} + +async fn assert_no_connection_until_connect_finishes( + listener: &TcpListener, + mut connect_done_rx: oneshot::Receiver<()>, +) { + tokio::select! { + accepted = listener.accept() => { + accepted.expect("unexpected websocket connection should be accepted"); + panic!("required refresh failure must not proceed to websocket connect"); + } + connect_done = &mut connect_done_rx => { + connect_done.expect("connect completion should be reported"); + } + } +} diff --git a/codex-rs/app-server-transport/src/transport/stdio.rs b/codex-rs/app-server-transport/src/transport/stdio.rs new file mode 100644 index 0000000000000000000000000000000000000000..12552800cf2c79f6e6fe2b75a3e4e8e18c2aaace --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/stdio.rs @@ -0,0 +1,194 @@ +//! Stdio transport with EOF cleanup and Unix SIGTERM shutdown. Dedicated I/O +//! threads let SIGTERM abandon blocked pipes without holding the runtime open. +//! A dedicated signal thread arms the shared EOF/SIGTERM deadline even when +//! logging or the Tokio runtime is blocked. + +use super::CHANNEL_CAPACITY; +use super::ConnectionOrigin; +use super::TransportEvent; +use super::forward_incoming_message; +use super::next_connection_id; +use super::serialize_outgoing_message; +use crate::outgoing_message::QueuedOutgoingMessage; +use codex_app_server_protocol::InitializeParams; +use codex_app_server_protocol::JSONRPCMessage; +use codex_app_server_protocol::JSONRPCRequest; +use std::io::BufRead; +use std::io::ErrorKind; +use std::io::Result as IoResult; +use std::io::Write; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::task::JoinHandle; +use tokio_util::sync::CancellationToken; +use tracing::debug; +use tracing::error; +use tracing::info; + +pub async fn start_stdio_connection( + transport_event_tx: mpsc::Sender, + initialize_client_name_tx: oneshot::Sender, + install_shutdown_signal_handler: bool, +) -> IoResult> { + let shutdown_signal = CancellationToken::new(); + #[cfg(unix)] + if install_shutdown_signal_handler { + use signal_hook::consts::SIGTERM; + use signal_hook::iterator::Signals; + + // Register before accepting requests. Receiving SIGTERM and arming the + // watchdog must not depend on Tokio or synchronous transport logging. + let mut signals = Signals::new([SIGTERM])?; + let shutdown_signal = shutdown_signal.clone(); + std::thread::Builder::new() + .name("app-server-signal".to_string()) + .spawn(move || { + if signals.forever().next().is_some() { + start_shutdown_watchdog(); + shutdown_signal.cancel(); + } + })?; + } + let connection_id = next_connection_id(); + let (writer_tx, mut writer_rx) = mpsc::channel::(CHANNEL_CAPACITY); + let writer_tx_for_reader = writer_tx.clone(); + transport_event_tx + .send(TransportEvent::ConnectionOpened { + connection_id, + origin: ConnectionOrigin::Stdio, + auth: None, + writer: writer_tx, + disconnect_sender: None, + }) + .await + .map_err(|_| std::io::Error::new(ErrorKind::BrokenPipe, "processor unavailable"))?; + + // Tokio's stdin uses an uncancellable blocking task, which keeps its runtime + // alive while the client leaves stdin open. This process-owned thread may + // remain blocked until process exit, but is not joined by the runtime. + let (stdin_tx, mut stdin_rx) = mpsc::channel(/*buffer*/ 1); + std::thread::Builder::new() + .name("app-server-stdin".to_string()) + .spawn(move || { + for line in std::io::stdin().lock().lines() { + if stdin_tx.blocking_send(line).is_err() { + break; + } + } + })?; + + // Keep stdout's blocking writes off Tokio's pool too. The forwarding future + // owns writer_rx so cancelling it also releases producers stuck on a full queue. + let (stdout_tx, mut stdout_rx) = + mpsc::channel::<(String, oneshot::Sender<()>)>(/*buffer*/ 1); + std::thread::Builder::new() + .name("app-server-stdout".to_string()) + .spawn(move || { + let mut stdout = std::io::stdout().lock(); + while let Some((json, written_tx)) = stdout_rx.blocking_recv() { + if let Err(err) = stdout.write_all(json.as_bytes()) { + error!("Failed to write to stdout: {err}"); + break; + } + let _ = written_tx.send(()); + } + })?; + + let transport_event_tx_for_reader = transport_event_tx.clone(); + let read_messages = async move { + let mut initialize_client_name_tx = Some(initialize_client_name_tx); + while let Some(line) = stdin_rx.recv().await { + let line = match line { + Ok(line) => line, + Err(err) => { + error!("Failed reading stdin: {err}"); + break; + } + }; + if let Some(client_name) = stdio_initialize_client_name(&line) + && let Some(initialize_client_name_tx) = initialize_client_name_tx.take() + { + let _ = initialize_client_name_tx.send(client_name); + } + if !forward_incoming_message( + &transport_event_tx_for_reader, + &writer_tx_for_reader, + connection_id, + &line, + ) + .await + { + break; + } + } + + // EOF can finish the transport before RPC or runtime cleanup. Start + // the same process deadline even if no SIGTERM arrives. + if cfg!(unix) && install_shutdown_signal_handler { + start_shutdown_watchdog(); + } + let _ = transport_event_tx_for_reader + .send(TransportEvent::ConnectionClosed { connection_id }) + .await; + debug!("stdin reader finished (EOF)"); + }; + + let write_messages = async move { + while let Some(queued_message) = writer_rx.recv().await { + let Some(mut json) = serialize_outgoing_message(queued_message.message) else { + continue; + }; + json.push('\n'); + let (written_tx, written_rx) = oneshot::channel(); + if stdout_tx.send((json, written_tx)).await.is_err() || written_rx.await.is_err() { + break; + } + if let Some(write_complete_tx) = queued_message.write_complete_tx { + let _ = write_complete_tx.send(()); + } + } + info!("stdout writer exited (channel closed)"); + }; + + Ok(tokio::spawn(async move { + tokio::select! { + _ = shutdown_signal.cancelled() => { + // Cancelling both forwarding futures drops their queues before + // connection teardown, including when EOF already began draining. + info!("SIGTERM received; closing stdio connection (45s shutdown deadline)"); + let _ = transport_event_tx + .send(TransportEvent::ConnectionClosed { connection_id }) + .await; + } + _ = async move { tokio::join!(read_messages, write_messages); } => {} + } + })) +} + +fn start_shutdown_watchdog() { + // EOF and SIGTERM can both start cleanup; keep the first process deadline. + static STARTED: std::sync::Once = std::sync::Once::new(); + STARTED.call_once(|| { + std::thread::Builder::new() + .name("app-server-shutdown".to_string()) + .spawn(|| { + // Allow the processor's 30s RPC drain, then bound even Tokio's + // runtime teardown. Do not log here: stderr may also be blocked. + std::thread::sleep(std::time::Duration::from_secs(45)); + std::process::exit(/*code*/ 1); + }) + .unwrap_or_else(|_| std::process::exit(/*code*/ 1)); + }); +} + +fn stdio_initialize_client_name(line: &str) -> Option { + let message = serde_json::from_str::(line).ok()?; + let JSONRPCMessage::Request(JSONRPCRequest { method, params, .. }) = message else { + return None; + }; + if method != "initialize" { + return None; + } + let params = serde_json::from_value::(params?).ok()?; + Some(params.client_info.name) +} diff --git a/codex-rs/app-server-transport/src/transport/unix_socket.rs b/codex-rs/app-server-transport/src/transport/unix_socket.rs new file mode 100644 index 0000000000000000000000000000000000000000..0d55d58ea565465d19202eac3f12db6e9cf08397 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/unix_socket.rs @@ -0,0 +1,265 @@ +//! Control socket startup, guarded rendezvous paths, and WebSocket acceptance. + +use std::fs::OpenOptions; +use std::io::ErrorKind; +use std::io::Result as IoResult; +use std::path::Path; + +use super::TransportEvent; +use crate::transport::websocket::run_websocket_connection; +use codex_uds::UnixListener; +use codex_uds::UnixStream; +use codex_utils_absolute_path::AbsolutePathBuf; +use futures::SinkExt; +use futures::StreamExt; +use tokio::sync::mpsc; +use tokio::task::JoinHandle; +use tokio::time::Duration; +use tokio_tungstenite::accept_hdr_async; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::http::Response; +use tokio_tungstenite::tungstenite::http::StatusCode; +use tokio_util::sync::CancellationToken; +use tracing::error; +use tracing::info; +use tracing::warn; + +#[cfg(unix)] +const CONTROL_SOCKET_MODE: u32 = 0o600; + +#[derive(Clone, Copy)] +pub enum DaemonShutdownAccess { + Disabled, + Managed, +} + +pub async fn start_control_socket_acceptor( + socket_path: AbsolutePathBuf, + transport_event_tx: mpsc::Sender, + shutdown_token: CancellationToken, + daemon_shutdown_access: DaemonShutdownAccess, +) -> IoResult> { + #[cfg(windows)] + let (socket_path, directory_guard) = { + if let Some(parent) = socket_path.as_path().parent() { + codex_uds::prepare_private_socket_directory(parent).await?; + } + let (path, guard) = codex_uds::validate_private_socket_path(socket_path.as_path())?; + (AbsolutePathBuf::from_absolute_path_checked(path)?, guard) + }; + prepare_control_socket_path(socket_path.as_path()).await?; + let listener = UnixListener::bind(socket_path.as_path()).await?; + let socket_guard = ControlSocketFileGuard { + socket_path, + #[cfg(windows)] + _directory_guard: directory_guard, + }; + set_control_socket_permissions(socket_guard.socket_path.as_path()).await?; + info!( + socket_path = %socket_guard.socket_path.display(), + "app-server control socket listening" + ); + + Ok(tokio::spawn(run_control_socket_acceptor( + listener, + transport_event_tx, + shutdown_token, + socket_guard, + daemon_shutdown_access, + ))) +} + +async fn run_control_socket_acceptor( + mut listener: UnixListener, + transport_event_tx: mpsc::Sender, + shutdown_token: CancellationToken, + socket_guard: ControlSocketFileGuard, + daemon_shutdown_access: DaemonShutdownAccess, +) { + let _socket_guard = socket_guard; + loop { + let stream = tokio::select! { + _ = shutdown_token.cancelled() => { + break; + } + result = listener.accept() => { + match result { + Ok(stream) => stream, + Err(err) => { + if matches!( + err.kind(), + ErrorKind::ConnectionAborted | ErrorKind::ConnectionReset | ErrorKind::Interrupted + ) { + warn!("recoverable control socket accept error: {err}"); + continue; + } + error!("control socket accept error: {err}"); + tokio::time::sleep(Duration::from_secs(1)).await; + continue; + } + } + } + }; + + let transport_event_tx = transport_event_tx.clone(); + tokio::spawn(async move { + let mut shutdown_request = false; + let websocket_stream = match accept_hdr_async( + stream, + |request: &tokio_tungstenite::tungstenite::handshake::server::Request, response| { + if request.uri().path() == "/daemon/shutdown" { + if !matches!(daemon_shutdown_access, DaemonShutdownAccess::Managed) { + let mut rejection = Response::new(Some("unmanaged server".to_string())); + *rejection.status_mut() = StatusCode::FORBIDDEN; + return Err(rejection); + } + shutdown_request = true; + } + Ok(response) + }, + ) + .await + { + Ok(websocket_stream) => websocket_stream, + Err(err) => { + warn!("failed to upgrade control socket websocket connection: {err}"); + return; + } + }; + if shutdown_request { + run_daemon_shutdown(websocket_stream, transport_event_tx).await; + return; + } + let (websocket_writer, websocket_reader) = websocket_stream.split(); + run_websocket_connection(websocket_writer, websocket_reader, transport_event_tx).await; + }); + } + info!("control socket acceptor shutting down"); +} + +async fn run_daemon_shutdown( + mut websocket: tokio_tungstenite::WebSocketStream, + transport_event_tx: mpsc::Sender, +) { + let pid = std::process::id().to_string(); + if !matches!(websocket.next().await, Some(Ok(Message::Text(request))) if request == pid) { + return; + } + if websocket.send(Message::Text(pid.into())).await.is_err() { + return; + } + // Let the manager receive the acknowledgment before the main loop closes connections. + let _ = tokio::time::timeout(Duration::from_secs(2), websocket.next()).await; + let _ = transport_event_tx + .send(TransportEvent::DaemonShutdown) + .await; +} + +pub async fn prepare_control_socket_path(socket_path: &Path) -> IoResult<()> { + if let Some(parent) = socket_path.parent() { + codex_uds::prepare_private_socket_directory(parent).await?; + } + + #[cfg(windows)] + let (socket_path, _directory_guard) = codex_uds::validate_private_socket_path(socket_path)?; + #[cfg(windows)] + let socket_path = AbsolutePathBuf::from_absolute_path_checked(socket_path)?; + #[cfg(windows)] + let socket_path = socket_path.as_path(); + + match UnixStream::connect(socket_path).await { + Ok(_stream) => { + return Err(std::io::Error::new( + ErrorKind::AddrInUse, + format!( + "app-server control socket is already in use at {}", + socket_path.display() + ), + )); + } + Err(err) if err.kind() == ErrorKind::NotFound => return Ok(()), + Err(err) if err.kind() == ErrorKind::ConnectionRefused => {} + Err(err) => { + if !socket_path.exists() { + return Ok(()); + } + return Err(err); + } + } + + if !socket_path.try_exists()? { + return Ok(()); + } + + if !codex_uds::is_stale_socket_path(socket_path).await? { + return Err(std::io::Error::new( + ErrorKind::AlreadyExists, + format!( + "app-server control socket path exists and is not a socket: {}", + socket_path.display() + ), + )); + } + tokio::fs::remove_file(socket_path).await +} + +pub struct AppServerStartupLock { + _file: std::fs::File, +} + +pub async fn acquire_app_server_startup_lock( + startup_lock_path: AbsolutePathBuf, +) -> IoResult { + if let Some(parent) = startup_lock_path.as_path().parent() { + codex_uds::prepare_private_socket_directory(parent).await?; + } + tokio::task::spawn_blocking(move || { + let file = OpenOptions::new() + .create(true) + .truncate(false) + .read(true) + .write(true) + .open(startup_lock_path.as_path())?; + file.lock()?; + Ok(AppServerStartupLock { _file: file }) + }) + .await + .map_err(|err| std::io::Error::other(format!("startup lock task failed: {err}")))? +} + +#[cfg(unix)] +async fn set_control_socket_permissions(socket_path: &Path) -> IoResult<()> { + use std::os::unix::fs::PermissionsExt; + + tokio::fs::set_permissions( + socket_path, + std::fs::Permissions::from_mode(CONTROL_SOCKET_MODE), + ) + .await +} + +#[cfg(not(unix))] +async fn set_control_socket_permissions(_socket_path: &Path) -> IoResult<()> { + Ok(()) +} + +struct ControlSocketFileGuard { + socket_path: AbsolutePathBuf, + // Keep the directory pinned until after the socket file is removed in Drop. + #[cfg(windows)] + _directory_guard: std::os::windows::io::OwnedHandle, +} + +impl Drop for ControlSocketFileGuard { + fn drop(&mut self) { + if let Err(err) = std::fs::remove_file(self.socket_path.as_path()) + && err.kind() != ErrorKind::NotFound + { + warn!( + socket_path = %self.socket_path.display(), + %err, + "failed to remove app-server control socket file" + ); + } + } +} diff --git a/codex-rs/app-server-transport/src/transport/unix_socket_tests.rs b/codex-rs/app-server-transport/src/transport/unix_socket_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..55637c015ad31e91f972ba1c2170a50b375a1ddd --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/unix_socket_tests.rs @@ -0,0 +1,340 @@ +use super::AppServerTransport; +use super::CHANNEL_CAPACITY; +use super::DaemonShutdownAccess; +use super::TransportEvent; +use super::acquire_app_server_startup_lock; +use super::app_server_control_socket_path; +use super::start_control_socket_acceptor; +use codex_app_server_protocol::JSONRPCMessage; +use codex_app_server_protocol::JSONRPCNotification; +use codex_core::config::find_codex_home; +use codex_uds::UnixStream; +use codex_utils_absolute_path::AbsolutePathBuf; +use futures::SinkExt; +use futures::StreamExt; +use pretty_assertions::assert_eq; +use std::io::Result as IoResult; +use std::path::Path; +use tokio::sync::mpsc; +use tokio::time::Duration; +use tokio::time::timeout; +use tokio_tungstenite::client_async; +use tokio_tungstenite::tungstenite::Bytes; +use tokio_tungstenite::tungstenite::Message as WebSocketMessage; +use tokio_util::sync::CancellationToken; + +#[test] +fn listen_unix_socket_parses_as_unix_socket_transport() { + assert_eq!( + AppServerTransport::from_listen_url("unix://"), + Ok(AppServerTransport::UnixSocket { + socket_path: default_control_socket_path() + }) + ); +} + +#[test] +fn listen_unix_socket_accepts_absolute_custom_path() { + assert_eq!( + AppServerTransport::from_listen_url("unix:///tmp/codex.sock"), + Ok(AppServerTransport::UnixSocket { + socket_path: absolute_path("/tmp/codex.sock") + }) + ); +} + +#[test] +fn listen_unix_socket_accepts_relative_custom_path() { + assert_eq!( + AppServerTransport::from_listen_url("unix://codex.sock"), + Ok(AppServerTransport::UnixSocket { + socket_path: AbsolutePathBuf::relative_to_current_dir("codex.sock") + .expect("relative path should resolve") + }) + ); +} + +#[tokio::test] +async fn control_socket_acceptor_upgrades_and_forwards_websocket_text_messages_and_pings() { + let temp_dir = tempfile::TempDir::new().expect("temp dir"); + let socket_path = test_socket_path(temp_dir.path()); + let (transport_event_tx, mut transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let accept_handle = start_control_socket_acceptor( + socket_path.clone(), + transport_event_tx, + shutdown_token.clone(), + DaemonShutdownAccess::Disabled, + ) + .await + .expect("control socket acceptor should start"); + + let stream = connect_to_socket(socket_path.as_path()) + .await + .expect("client should connect"); + let (mut websocket, response) = client_async("ws://localhost/rpc", stream) + .await + .expect("websocket upgrade should complete"); + assert_eq!(response.status().as_u16(), 101); + + let opened = timeout(Duration::from_secs(1), transport_event_rx.recv()) + .await + .expect("connection opened event should arrive") + .expect("connection opened event"); + let connection_id = match opened { + TransportEvent::ConnectionOpened { connection_id, .. } => connection_id, + _ => panic!("expected connection opened event"), + }; + + let notification = JSONRPCMessage::Notification(JSONRPCNotification { + method: "initialized".to_string(), + params: None, + }); + websocket + .send(WebSocketMessage::Text( + serde_json::to_string(¬ification) + .expect("notification should serialize") + .into(), + )) + .await + .expect("notification should send"); + + let incoming = timeout(Duration::from_secs(1), transport_event_rx.recv()) + .await + .expect("incoming message event should arrive") + .expect("incoming message event"); + assert_eq!( + match incoming { + TransportEvent::IncomingMessage { + connection_id: incoming_connection_id, + message, + } => (incoming_connection_id, message), + _ => panic!("expected incoming message event"), + }, + (connection_id, notification) + ); + + websocket + .send(WebSocketMessage::Ping(Bytes::from_static(b"check"))) + .await + .expect("ping should send"); + let pong = timeout(Duration::from_secs(1), websocket.next()) + .await + .expect("pong should arrive") + .expect("pong frame") + .expect("pong should be valid"); + assert_eq!(pong, WebSocketMessage::Pong(Bytes::from_static(b"check"))); + + websocket.close(None).await.expect("close should send"); + let closed = timeout(Duration::from_secs(1), transport_event_rx.recv()) + .await + .expect("connection closed event should arrive") + .expect("connection closed event"); + assert!(matches!( + closed, + TransportEvent::ConnectionClosed { + connection_id: closed_connection_id, + } if closed_connection_id == connection_id + )); + + shutdown_token.cancel(); + accept_handle.await.expect("acceptor should join"); + assert_socket_path_removed(socket_path.as_path()); +} + +#[tokio::test] +async fn shutdown_is_only_accepted_on_managed_local_socket_for_its_own_pid() { + let temp_dir = tempfile::TempDir::new().expect("temp dir"); + let socket_path = test_socket_path(temp_dir.path()); + let (tx, mut rx) = mpsc::channel(CHANNEL_CAPACITY); + let shutdown = CancellationToken::new(); + let acceptor = start_control_socket_acceptor( + socket_path.clone(), + tx, + shutdown.clone(), + DaemonShutdownAccess::Disabled, + ) + .await + .expect("acceptor"); + + let stream = connect_to_socket(socket_path.as_path()) + .await + .expect("connect"); + assert!( + client_async("ws://localhost/daemon/shutdown", stream) + .await + .is_err() + ); + assert!(rx.try_recv().is_err()); + shutdown.cancel(); + acceptor.await.expect("acceptor shutdown"); + + let (tx, mut rx) = mpsc::channel(CHANNEL_CAPACITY); + let shutdown = CancellationToken::new(); + let acceptor = start_control_socket_acceptor( + socket_path.clone(), + tx, + shutdown.clone(), + DaemonShutdownAccess::Managed, + ) + .await + .expect("managed acceptor"); + let stream = connect_to_socket(socket_path.as_path()) + .await + .expect("connect"); + let (mut websocket, _) = client_async("ws://localhost/daemon/shutdown", stream) + .await + .expect("upgrade"); + websocket + .send(WebSocketMessage::Text("0".into())) + .await + .expect("wrong pid"); + assert!(!matches!( + websocket.next().await, + Some(Ok(WebSocketMessage::Text(_))) + )); + assert!(rx.try_recv().is_err()); + + let stream = connect_to_socket(socket_path.as_path()) + .await + .expect("connect"); + let (mut websocket, _) = client_async("ws://localhost/daemon/shutdown", stream) + .await + .expect("upgrade"); + let pid = std::process::id().to_string(); + websocket + .send(WebSocketMessage::Text(pid.clone().into())) + .await + .expect("request"); + assert_eq!( + websocket.next().await.expect("ack").expect("ack frame"), + WebSocketMessage::Text(pid.into()) + ); + assert!( + rx.try_recv().is_err(), + "server must wait until the ack is received" + ); + websocket.close(None).await.expect("confirm receipt"); + assert!(matches!( + timeout(Duration::from_secs(2), rx.recv()).await, + Ok(Some(TransportEvent::DaemonShutdown)) + )); + shutdown.cancel(); + acceptor.await.expect("acceptor shutdown"); +} + +#[tokio::test] +async fn app_server_startup_lock_serializes_waiters() { + let temp_dir = tempfile::TempDir::new().expect("temp dir"); + let lock_path = test_startup_lock_path(temp_dir.path()); + let first_lock = acquire_app_server_startup_lock(lock_path.clone()) + .await + .expect("first startup lock should succeed"); + let mut second_lock = tokio::spawn(acquire_app_server_startup_lock(lock_path)); + + assert!( + timeout(Duration::from_millis(100), &mut second_lock) + .await + .is_err() + ); + + drop(first_lock); + second_lock + .await + .expect("second startup lock task should join") + .expect("second startup lock should succeed"); +} + +#[cfg(unix)] +#[tokio::test] +async fn control_socket_file_is_private_after_bind() { + use std::os::unix::fs::PermissionsExt; + + let temp_dir = tempfile::TempDir::new().expect("temp dir"); + let socket_path = test_socket_path(temp_dir.path()); + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let accept_handle = start_control_socket_acceptor( + socket_path.clone(), + transport_event_tx, + shutdown_token.clone(), + DaemonShutdownAccess::Disabled, + ) + .await + .expect("control socket acceptor should start"); + + let metadata = tokio::fs::metadata(socket_path.as_path()) + .await + .expect("socket metadata should exist"); + assert_eq!(metadata.permissions().mode() & 0o777, 0o600); + + shutdown_token.cancel(); + accept_handle.await.expect("acceptor should join"); +} + +#[cfg(windows)] +#[tokio::test] +async fn control_socket_pins_directory_until_shutdown() { + let temp_dir = tempfile::TempDir::new().expect("temp dir"); + let socket_path = test_socket_path(temp_dir.path()); + let directory = socket_path.as_path().parent().unwrap(); + let moved = temp_dir.path().join("moved"); + let (tx, _rx) = mpsc::channel::(CHANNEL_CAPACITY); + let shutdown = CancellationToken::new(); + let acceptor = start_control_socket_acceptor( + socket_path.clone(), + tx, + shutdown.clone(), + DaemonShutdownAccess::Disabled, + ) + .await + .expect("acceptor"); + assert!(std::fs::rename(directory, &moved).is_err()); + shutdown.cancel(); + acceptor.await.expect("shutdown"); + std::fs::rename(directory, moved).expect("directory unpinned after cleanup"); +} + +fn absolute_path(path: &str) -> AbsolutePathBuf { + AbsolutePathBuf::from_absolute_path(path).expect("absolute path") +} + +fn default_control_socket_path() -> AbsolutePathBuf { + let codex_home = find_codex_home().expect("codex home"); + app_server_control_socket_path(&codex_home).expect("default control socket path") +} + +fn test_socket_path(temp_dir: &Path) -> AbsolutePathBuf { + AbsolutePathBuf::from_absolute_path( + temp_dir + .join("app-server-control") + .join("app-server-control.sock"), + ) + .expect("socket path should resolve") +} + +fn test_startup_lock_path(temp_dir: &Path) -> AbsolutePathBuf { + AbsolutePathBuf::from_absolute_path( + temp_dir + .join("app-server-control") + .join("app-server-startup.lock"), + ) + .expect("startup lock path should resolve") +} + +async fn connect_to_socket(socket_path: &Path) -> IoResult { + UnixStream::connect(socket_path).await +} + +#[cfg(unix)] +fn assert_socket_path_removed(socket_path: &Path) { + assert!(!socket_path.exists()); +} + +#[cfg(windows)] +fn assert_socket_path_removed(_socket_path: &Path) { + // uds_windows uses a regular filesystem path as its rendezvous point, + // but there is no Unix socket filesystem node to assert on. +} diff --git a/codex-rs/app-server-transport/src/transport/websocket.rs b/codex-rs/app-server-transport/src/transport/websocket.rs new file mode 100644 index 0000000000000000000000000000000000000000..eb28e0e4fd81628d3e8bc351d6741bb7d2e1e2e1 --- /dev/null +++ b/codex-rs/app-server-transport/src/transport/websocket.rs @@ -0,0 +1,389 @@ +use super::CHANNEL_CAPACITY; +use super::ConnectionOrigin; +use super::TransportEvent; +use super::auth::WebsocketAuthPolicy; +use super::auth::authorize_upgrade; +use super::auth::is_unauthenticated_non_loopback_listener; +use super::forward_incoming_message; +use super::next_connection_id; +use super::serialize_outgoing_message; +use crate::outgoing_message::ConnectionId; +use crate::outgoing_message::QueuedOutgoingMessage; +use axum::Router; +use axum::body::Body; +use axum::body::Bytes; +use axum::extract::ConnectInfo; +use axum::extract::State; +use axum::extract::ws::Message as AxumWebSocketMessage; +use axum::extract::ws::WebSocketUpgrade; +use axum::http::HeaderMap; +use axum::http::Request; +use axum::http::StatusCode; +use axum::http::header::ORIGIN; +use axum::middleware; +use axum::middleware::Next; +use axum::response::IntoResponse; +use axum::response::Response; +use axum::routing::any; +use axum::routing::get; +use futures::SinkExt; +use futures::StreamExt; +use owo_colors::OwoColorize; +use owo_colors::Stream; +use owo_colors::Style; +use std::io::Result as IoResult; +use std::net::SocketAddr; +use std::sync::Arc; +use tokio::net::TcpListener; +use tokio::sync::mpsc; +use tokio::task::JoinHandle; +use tokio_tungstenite::tungstenite::Message as TungsteniteWebSocketMessage; +use tokio_util::sync::CancellationToken; +use tracing::error; +use tracing::info; +use tracing::warn; + +/// WebSocket clients can briefly lag behind normal turn output bursts while the +/// writer task is healthy, so give them more headroom than internal channels. +const WEBSOCKET_OUTBOUND_CHANNEL_CAPACITY: usize = 32 * 1024; +const _: () = assert!(WEBSOCKET_OUTBOUND_CHANNEL_CAPACITY > CHANNEL_CAPACITY); + +fn colorize(text: &str, style: Style) -> String { + text.if_supports_color(Stream::Stderr, |value| value.style(style)) + .to_string() +} + +#[allow(clippy::print_stderr)] +fn print_websocket_startup_banner(addr: SocketAddr) { + let title = colorize("codex app-server (WebSockets)", Style::new().bold().cyan()); + let listening_label = colorize("listening on:", Style::new().dimmed()); + let listen_url = colorize(&format!("ws://{addr}"), Style::new().green()); + let ready_label = colorize("readyz:", Style::new().dimmed()); + let ready_url = colorize(&format!("http://{addr}/readyz"), Style::new().green()); + let health_label = colorize("healthz:", Style::new().dimmed()); + let health_url = colorize(&format!("http://{addr}/healthz"), Style::new().green()); + let note_label = colorize("note:", Style::new().dimmed()); + eprintln!("{title}"); + eprintln!(" {listening_label} {listen_url}"); + eprintln!(" {ready_label} {ready_url}"); + eprintln!(" {health_label} {health_url}"); + if addr.ip().is_loopback() { + eprintln!( + " {note_label} binds localhost only (use SSH port-forwarding for remote access)" + ); + } else { + eprintln!(" {note_label} websocket auth is required for non-localhost listeners"); + } +} + +#[derive(Clone)] +struct WebSocketListenerState { + transport_event_tx: mpsc::Sender, + auth_policy: Arc, +} + +async fn health_check_handler() -> StatusCode { + StatusCode::OK +} + +async fn reject_requests_with_origin_header( + request: Request, + next: Next, +) -> Result { + if request.headers().contains_key(ORIGIN) { + warn!( + method = %request.method(), + uri = %request.uri(), + "rejecting websocket listener request with Origin header" + ); + Err(StatusCode::FORBIDDEN) + } else { + Ok(next.run(request).await) + } +} + +async fn websocket_upgrade_handler( + websocket: WebSocketUpgrade, + ConnectInfo(peer_addr): ConnectInfo, + State(state): State, + headers: HeaderMap, +) -> impl IntoResponse { + if let Err(err) = authorize_upgrade(&headers, state.auth_policy.as_ref()) { + warn!( + %peer_addr, + message = err.message(), + "rejecting websocket client during upgrade" + ); + return (err.status_code(), err.message()).into_response(); + } + info!(%peer_addr, "websocket client connected"); + websocket + .on_upgrade(move |stream| async move { + let (websocket_writer, websocket_reader) = stream.split(); + run_websocket_connection(websocket_writer, websocket_reader, state.transport_event_tx) + .await; + }) + .into_response() +} + +pub async fn start_websocket_acceptor( + bind_address: SocketAddr, + transport_event_tx: mpsc::Sender, + shutdown_token: CancellationToken, + auth_policy: WebsocketAuthPolicy, +) -> IoResult> { + if is_unauthenticated_non_loopback_listener(bind_address, &auth_policy) { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!( + "refusing to start non-loopback websocket listener {bind_address} without auth; configure `--ws-auth capability-token` or `--ws-auth signed-bearer-token`" + ), + )); + } + let listener = TcpListener::bind(bind_address).await?; + let local_addr = listener.local_addr()?; + print_websocket_startup_banner(local_addr); + info!("app-server websocket listening on ws://{local_addr}"); + + let router = Router::new() + .route("/readyz", get(health_check_handler)) + .route("/healthz", get(health_check_handler)) + .fallback(any(websocket_upgrade_handler)) + .layer(middleware::from_fn(reject_requests_with_origin_header)) + .with_state(WebSocketListenerState { + transport_event_tx, + auth_policy: Arc::new(auth_policy), + }); + let server = axum::serve( + listener, + router.into_make_service_with_connect_info::(), + ) + .with_graceful_shutdown(async move { + shutdown_token.cancelled().await; + }); + Ok(tokio::spawn(async move { + if let Err(err) = server.await { + error!("websocket acceptor failed: {err}"); + } + info!("websocket acceptor shutting down"); + })) +} + +pub(crate) async fn run_websocket_connection( + websocket_writer: impl futures::sink::Sink + Send + 'static, + websocket_reader: impl futures::stream::Stream> + Send + 'static, + transport_event_tx: mpsc::Sender, +) where + M: AppServerWebSocketMessage + Send + 'static, + SinkError: Send + 'static, + StreamError: std::fmt::Display + Send + 'static, +{ + let connection_id = next_connection_id(); + let (writer_tx, writer_rx) = + mpsc::channel::(WEBSOCKET_OUTBOUND_CHANNEL_CAPACITY); + let writer_tx_for_reader = writer_tx.clone(); + let disconnect_token = CancellationToken::new(); + if transport_event_tx + .send(TransportEvent::ConnectionOpened { + connection_id, + origin: ConnectionOrigin::WebSocket, + auth: None, + writer: writer_tx, + disconnect_sender: Some(disconnect_token.clone()), + }) + .await + .is_err() + { + return; + } + + let (writer_control_tx, writer_control_rx) = mpsc::channel::(CHANNEL_CAPACITY); + let mut outbound_task = tokio::spawn(run_websocket_outbound_loop( + websocket_writer, + writer_rx, + writer_control_rx, + disconnect_token.clone(), + )); + let mut inbound_task = tokio::spawn(run_websocket_inbound_loop( + websocket_reader, + transport_event_tx.clone(), + writer_tx_for_reader, + writer_control_tx, + connection_id, + disconnect_token.clone(), + )); + + tokio::select! { + _ = &mut outbound_task => { + disconnect_token.cancel(); + inbound_task.abort(); + } + _ = &mut inbound_task => { + disconnect_token.cancel(); + outbound_task.abort(); + } + } + + let _ = transport_event_tx + .send(TransportEvent::ConnectionClosed { connection_id }) + .await; +} + +pub(crate) enum IncomingWebSocketMessage { + Text(String), + Binary, + Ping(Bytes), + Pong, + Close, +} + +/// Converts concrete WebSocket message types into the small message surface the +/// app-server transport needs, and constructs the only outbound frames it +/// sends directly. +pub(crate) trait AppServerWebSocketMessage: Sized { + fn text(text: String) -> Self; + fn pong(payload: Bytes) -> Self; + fn into_incoming(self) -> Option; +} + +impl AppServerWebSocketMessage for AxumWebSocketMessage { + fn text(text: String) -> Self { + Self::Text(text.into()) + } + + fn pong(payload: Bytes) -> Self { + Self::Pong(payload) + } + + fn into_incoming(self) -> Option { + Some(match self { + Self::Text(text) => IncomingWebSocketMessage::Text(text.to_string()), + Self::Binary(_) => IncomingWebSocketMessage::Binary, + Self::Ping(payload) => IncomingWebSocketMessage::Ping(payload), + Self::Pong(_) => IncomingWebSocketMessage::Pong, + Self::Close(_) => IncomingWebSocketMessage::Close, + }) + } +} + +impl AppServerWebSocketMessage for TungsteniteWebSocketMessage { + fn text(text: String) -> Self { + Self::Text(text.into()) + } + + fn pong(payload: Bytes) -> Self { + Self::Pong(payload) + } + + fn into_incoming(self) -> Option { + Some(match self { + Self::Text(text) => IncomingWebSocketMessage::Text(text.to_string()), + Self::Binary(_) => IncomingWebSocketMessage::Binary, + Self::Ping(payload) => IncomingWebSocketMessage::Ping(payload), + Self::Pong(_) => IncomingWebSocketMessage::Pong, + Self::Close(_) => IncomingWebSocketMessage::Close, + Self::Frame(_) => return None, + }) + } +} + +async fn run_websocket_outbound_loop( + websocket_writer: impl futures::sink::Sink + Send + 'static, + mut writer_rx: mpsc::Receiver, + mut writer_control_rx: mpsc::Receiver, + disconnect_token: CancellationToken, +) where + M: AppServerWebSocketMessage + Send + 'static, + SinkError: Send + 'static, +{ + tokio::pin!(websocket_writer); + loop { + tokio::select! { + _ = disconnect_token.cancelled() => { + break; + } + message = writer_control_rx.recv() => { + let Some(message) = message else { + break; + }; + if websocket_writer.send(message).await.is_err() { + break; + } + } + queued_message = writer_rx.recv() => { + let Some(queued_message) = queued_message else { + break; + }; + let Some(json) = serialize_outgoing_message(queued_message.message) else { + continue; + }; + if websocket_writer.send(M::text(json)).await.is_err() { + break; + } + if let Some(write_complete_tx) = queued_message.write_complete_tx { + let _ = write_complete_tx.send(()); + } + } + } + } +} + +async fn run_websocket_inbound_loop( + websocket_reader: impl futures::stream::Stream> + Send + 'static, + transport_event_tx: mpsc::Sender, + writer_tx_for_reader: mpsc::Sender, + writer_control_tx: mpsc::Sender, + connection_id: ConnectionId, + disconnect_token: CancellationToken, +) where + M: AppServerWebSocketMessage + Send + 'static, + StreamError: std::fmt::Display + Send + 'static, +{ + tokio::pin!(websocket_reader); + loop { + tokio::select! { + _ = disconnect_token.cancelled() => { + break; + } + incoming_message = websocket_reader.next() => { + match incoming_message { + Some(Ok(message)) => match message.into_incoming() { + Some(IncomingWebSocketMessage::Text(text)) + if !forward_incoming_message( + &transport_event_tx, + &writer_tx_for_reader, + connection_id, + &text, + ) + .await + => { + break; + } + Some(IncomingWebSocketMessage::Text(_)) => {} + Some(IncomingWebSocketMessage::Ping(payload)) => { + match writer_control_tx.try_send(M::pong(payload)) { + Ok(()) => {} + Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => break, + Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => { + warn!("websocket control queue full while replying to ping; closing connection"); + break; + } + } + } + Some(IncomingWebSocketMessage::Pong) => {} + Some(IncomingWebSocketMessage::Close) => break, + Some(IncomingWebSocketMessage::Binary) => { + warn!("dropping unsupported binary websocket message"); + } + None => {} + }, + None => break, + Some(Err(err)) => { + warn!("websocket receive error: {err}"); + break; + } + } + } + } + } +} diff --git a/codex-rs/bwrap/src/main.rs b/codex-rs/bwrap/src/main.rs new file mode 100644 index 0000000000000000000000000000000000000000..09c624aa9e584c76e9bc8a1e154067f83a2cb2ea --- /dev/null +++ b/codex-rs/bwrap/src/main.rs @@ -0,0 +1,45 @@ +#[cfg(all(target_os = "linux", bwrap_available))] +fn main() { + use std::ffi::CStr; + use std::ffi::CString; + use std::os::raw::c_char; + use std::os::unix::ffi::OsStrExt; + + unsafe extern "C" { + fn bwrap_main(argc: libc::c_int, argv: *const *const c_char) -> libc::c_int; + } + + let cstrings = std::env::args_os() + .map(|arg| { + CString::new(arg.as_os_str().as_bytes()) + .unwrap_or_else(|err| panic!("failed to convert argv to CString: {err}")) + }) + .collect::>(); + let mut argv_ptrs = cstrings + .iter() + .map(CString::as_c_str) + .map(CStr::as_ptr) + .collect::>(); + argv_ptrs.push(std::ptr::null()); + + // SAFETY: We provide a null-terminated argv vector whose pointers remain + // valid for the duration of the call. + let exit_code = unsafe { bwrap_main(cstrings.len() as libc::c_int, argv_ptrs.as_ptr()) }; + std::process::exit(exit_code); +} + +#[cfg(all(target_os = "linux", not(bwrap_available)))] +fn main() { + panic!( + r#"bubblewrap is not available in this build. +Notes: +- ensure the target OS is Linux +- libcap headers must be available via pkg-config +- bubblewrap sources expected at codex-rs/vendor/bubblewrap (default)"# + ); +} + +#[cfg(not(target_os = "linux"))] +fn main() { + panic!("bwrap is only supported on Linux"); +} diff --git a/codex-rs/codex-mcp/src/agent_plugin_config.rs b/codex-rs/codex-mcp/src/agent_plugin_config.rs new file mode 100644 index 0000000000000000000000000000000000000000..d2a445566d72639231f938b1bdcec7411b771526 --- /dev/null +++ b/codex-rs/codex-mcp/src/agent_plugin_config.rs @@ -0,0 +1,532 @@ +use super::PluginMcpConfigParseOutcome; +use super::PluginMcpServerParseError; +use codex_config::McpServerConfig; +use serde::Deserialize; +use serde_json::Map as JsonMap; +use serde_json::Value as JsonValue; +use std::collections::BTreeMap; +use std::ffi::OsString; +use std::path::Path; +use std::path::PathBuf; +use url::Host; + +// Published Agent Plugins v1 MCP schema: +// https://github.com/agentplugins/agent-plugins-spec/blob/main/schemas/1.0.0/mcp.schema.json +const AGENT_PLUGIN_MCP_SCHEMA_URI: &str = "https://agent-plugins.org/schemas/1.0.0/mcp.schema.json"; +const SUPPORTED_AGENT_PLUGIN_MCP_SCHEMA_URIS: &[&str] = &[AGENT_PLUGIN_MCP_SCHEMA_URI]; +const PLUGIN_ROOT_VARIABLE: &str = "PLUGIN_ROOT"; +const PLUGIN_DATA_VARIABLE: &str = "PLUGIN_DATA"; +const CLIENT_OWNED_HTTP_HEADERS: &[&str] = &[ + "accept", + "authorization", + "connection", + "content-encoding", + "content-length", + "content-type", + "host", + "last-event-id", + "mcp-protocol-version", + "mcp-session-id", + "proxy-authorization", + "te", + "trailer", + "transfer-encoding", + "upgrade", + "user-agent", +]; + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +struct AgentPluginMcpFile { + #[serde(rename = "$schema")] + schema: String, + mcp_servers: BTreeMap, +} + +#[derive(Debug, Deserialize)] +#[serde(tag = "type", deny_unknown_fields)] +enum AgentPluginMcpServer { + #[serde(rename = "stdio")] + Stdio { + command: String, + #[serde(default)] + args: Vec, + #[serde(default)] + env: BTreeMap, + cwd: Option, + }, + #[serde(rename = "streamable-http")] + StreamableHttp { + url: String, + headers: Option>, + }, + #[serde(rename = "sse")] + Sse { + #[serde(rename = "url")] + _url: String, + #[serde(rename = "headers")] + _headers: Option>, + }, +} + +/// Translates an Agent Plugins `mcp.json` into Codex MCP configuration. +pub fn parse_agent_plugin_mcp_config( + plugin_root: &Path, + plugin_data_root: &Path, + contents: &str, +) -> Result { + parse_agent_plugin_mcp_config_from(contents, plugin_root, plugin_data_root) +} + +fn parse_agent_plugin_mcp_config_from( + contents: &str, + plugin_root: &Path, + plugin_data_root: &Path, +) -> Result { + let AgentPluginMcpFile { + schema, + mcp_servers, + } = serde_json::from_str(contents)?; + if !SUPPORTED_AGENT_PLUGIN_MCP_SCHEMA_URIS.contains(&schema.as_str()) { + return Err(plugin_mcp_json_error(format!( + "unsupported Agent Plugins MCP schema `{schema}`; supported schemas: {}", + SUPPORTED_AGENT_PLUGIN_MCP_SCHEMA_URIS.join(", ") + ))); + } + + let mut outcome = PluginMcpConfigParseOutcome::default(); + for (name, value) in mcp_servers { + match normalize_agent_plugin_mcp_server(value, plugin_root, plugin_data_root) { + Ok(config) => { + outcome.servers.insert(name, config); + } + Err(message) => outcome + .errors + .push(PluginMcpServerParseError { name, message }), + } + } + Ok(outcome) +} + +fn normalize_agent_plugin_mcp_server( + value: JsonValue, + plugin_root: &Path, + plugin_data_root: &Path, +) -> Result { + let object = value + .as_object() + .ok_or_else(|| "Agent Plugins MCP server must be an object".to_string())?; + match object.get("type").and_then(JsonValue::as_str) { + Some("stdio") => reject_explicit_null(object, "cwd")?, + Some("streamable-http" | "sse") => reject_explicit_null(object, "headers")?, + _ => {} + } + let server = + serde_json::from_value::(value).map_err(|err| err.to_string())?; + let object = match server { + AgentPluginMcpServer::Stdio { + command, + args, + env, + cwd, + } => normalize_agent_plugin_stdio_server( + command, + args, + env, + cwd, + plugin_root, + plugin_data_root, + )?, + AgentPluginMcpServer::StreamableHttp { url, headers } => { + normalize_agent_plugin_http_server(url, headers)? + } + AgentPluginMcpServer::Sse { .. } => { + return Err("Agent Plugins legacy SSE transport is not supported by Codex".to_string()); + } + }; + serde_json::from_value(JsonValue::Object(object)).map_err(|err| err.to_string()) +} + +fn normalize_agent_plugin_stdio_server( + mut command: String, + mut args: Vec, + mut env: BTreeMap, + cwd: Option, + plugin_root: &Path, + plugin_data_root: &Path, +) -> Result, String> { + #[cfg(windows)] + let has_windows_path_prefix = matches!( + Path::new(&command).components().next(), + Some(std::path::Component::Prefix(_)) + ); + #[cfg(not(windows))] + let has_windows_path_prefix = false; + let is_bare_command = !command.is_empty() + && !command.contains('/') + && !command.contains('\\') + && !has_windows_path_prefix; + let is_plugin_relative_command = + command.starts_with("./") && is_portable_relative_path(&command); + if !is_bare_command && !is_plugin_relative_command { + return Err( + "Agent Plugins stdio command must be a bare executable name or a contained `./` path" + .to_string(), + ); + } + for reserved in [PLUGIN_ROOT_VARIABLE, PLUGIN_DATA_VARIABLE] { + if env + .keys() + .any(|name| environment_variable_names_match(name, reserved)) + { + return Err(format!( + "Agent Plugins stdio `env` cannot override reserved variable `{reserved}`" + )); + } + } + #[cfg(windows)] + { + let mut normalized_env = BTreeMap::new(); + for (name, value) in env { + let normalized_name = name.to_ascii_uppercase(); + if normalized_env.insert(normalized_name, value).is_some() { + return Err(format!( + "duplicate case-insensitive Agent Plugins environment variable `{name}`" + )); + } + } + env = normalized_env; + } + + let root_path = absolute_plugin_path(plugin_root)?; + let data_root_path = absolute_plugin_path(plugin_data_root)?; + let root = host_path_string(&root_path); + let data_root = host_path_string(&data_root_path); + if command.starts_with("./") { + command = host_path_string(&resolve_contained_host_path( + &command, &root_path, &root_path, + )?); + } + for arg in &mut args { + *arg = expand_agent_plugin_placeholders(arg, &root, &data_root); + } + for value in env.values_mut() { + *value = expand_agent_plugin_placeholders(value, &root, &data_root); + } + + let cwd = cwd.as_deref().unwrap_or("${PLUGIN_ROOT}"); + let Some(cwd_root) = parse_agent_plugin_cwd(cwd) else { + return Err( + "Agent Plugins stdio `cwd` must be a contained `./`, `${PLUGIN_ROOT}`, or `${PLUGIN_DATA}` path" + .to_string(), + ); + }; + let cwd = expand_agent_plugin_placeholders(cwd, &root, &data_root); + let cwd_root = match cwd_root { + AgentPluginCwdRoot::Package => &root_path, + AgentPluginCwdRoot::Data => &data_root_path, + }; + env.insert(PLUGIN_ROOT_VARIABLE.to_string(), root); + env.insert(PLUGIN_DATA_VARIABLE.to_string(), data_root); + + Ok(JsonMap::from_iter([ + ("command".to_string(), JsonValue::String(command)), + ( + "args".to_string(), + JsonValue::Array(args.into_iter().map(JsonValue::String).collect()), + ), + ("env".to_string(), string_map_value(env)), + ( + "cwd".to_string(), + JsonValue::String(host_path_string(&resolve_contained_host_path( + &cwd, cwd_root, cwd_root, + )?)), + ), + ])) +} + +fn reject_explicit_null(object: &JsonMap, field: &str) -> Result<(), String> { + if object.get(field).is_some_and(JsonValue::is_null) { + return Err(format!( + "Agent Plugins MCP `{field}` must use its declared type when present" + )); + } + Ok(()) +} + +fn environment_variable_names_match(left: &str, right: &str) -> bool { + if cfg!(windows) { + left.eq_ignore_ascii_case(right) + } else { + left == right + } +} + +fn normalize_agent_plugin_http_server( + url: String, + mut headers: Option>, +) -> Result, String> { + validate_agent_plugin_url(&url)?; + if let Some(configured_headers) = headers.as_mut() { + validate_agent_plugin_headers(configured_headers)?; + configured_headers.retain(|name, _| { + !CLIENT_OWNED_HTTP_HEADERS + .iter() + .any(|owned| name.eq_ignore_ascii_case(owned)) + }); + } + let mut object = JsonMap::from_iter([("url".to_string(), JsonValue::String(url))]); + if let Some(headers) = headers.filter(|headers| !headers.is_empty()) { + object.insert("http_headers".to_string(), string_map_value(headers)); + } + Ok(object) +} + +fn validate_agent_plugin_url(raw_url: &str) -> Result<(), String> { + if raw_url.is_empty() { + return Err("Agent Plugins HTTP server requires a non-empty `url`".to_string()); + } + let parsed = url::Url::parse(raw_url) + .map_err(|err| format!("invalid Agent Plugins MCP URL `{raw_url}`: {err}"))?; + if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() { + return Err("Agent Plugins MCP URL must be absolute HTTP or HTTPS".to_string()); + } + if !parsed.username().is_empty() || parsed.password().is_some() || parsed.fragment().is_some() { + return Err( + "Agent Plugins MCP URL must not contain user information or a fragment".to_string(), + ); + } + let is_loopback = match parsed.host() { + Some(Host::Domain(host)) => host == "localhost", + Some(Host::Ipv4(address)) => address.is_loopback(), + Some(Host::Ipv6(address)) => address.is_loopback(), + None => false, + }; + if parsed.scheme() == "http" && !is_loopback { + return Err("non-loopback Agent Plugins MCP endpoints must use HTTPS".to_string()); + } + Ok(()) +} + +fn validate_agent_plugin_headers(headers: &BTreeMap) -> Result<(), String> { + let mut seen = std::collections::HashSet::new(); + for (name, value) in headers { + if !seen.insert(name.to_ascii_lowercase()) { + return Err(format!( + "duplicate case-insensitive Agent Plugins HTTP header `{name}`" + )); + } + if !is_valid_http_header_name(name) { + return Err(format!("invalid Agent Plugins HTTP header name `{name}`")); + } + if value + .bytes() + .any(|byte| (byte < 32 && byte != b'\t') || byte == 127) + { + return Err(format!( + "invalid Agent Plugins HTTP header value for `{name}`" + )); + } + } + Ok(()) +} + +fn string_map_value(values: BTreeMap) -> JsonValue { + JsonValue::Object( + values + .into_iter() + .map(|(name, value)| (name, JsonValue::String(value))) + .collect(), + ) +} + +#[derive(Clone, Copy, Debug)] +enum AgentPluginCwdRoot { + Package, + Data, +} + +fn parse_agent_plugin_cwd(value: &str) -> Option { + if value == "./" { + return Some(AgentPluginCwdRoot::Package); + } + if let Some(relative) = value.strip_prefix("./") + && is_portable_path_suffix(relative) + { + return Some(AgentPluginCwdRoot::Package); + } + for (placeholder, root) in [ + ("${PLUGIN_ROOT}", AgentPluginCwdRoot::Package), + ("${PLUGIN_DATA}", AgentPluginCwdRoot::Data), + ] { + if value == placeholder { + return Some(root); + } + if let Some(relative) = value.strip_prefix(&format!("{placeholder}/")) + && (relative.is_empty() || is_portable_path_suffix(relative)) + { + return Some(root); + } + } + None +} + +fn expand_agent_plugin_placeholders(value: &str, plugin_root: &str, plugin_data: &str) -> String { + const ROOT: &str = "${PLUGIN_ROOT}"; + const DATA: &str = "${PLUGIN_DATA}"; + let mut output = String::with_capacity(value.len()); + let mut remaining = value; + loop { + let next = match (remaining.find(ROOT), remaining.find(DATA)) { + (Some(root), Some(data)) if root <= data => Some((root, ROOT, plugin_root)), + (Some(_), Some(data)) => Some((data, DATA, plugin_data)), + (Some(root), None) => Some((root, ROOT, plugin_root)), + (None, Some(data)) => Some((data, DATA, plugin_data)), + (None, None) => None, + }; + let Some((index, placeholder, replacement)) = next else { + output.push_str(remaining); + break; + }; + output.push_str(&remaining[..index]); + output.push_str(replacement); + remaining = &remaining[index + placeholder.len()..]; + } + output +} + +fn absolute_plugin_path(path: &Path) -> Result { + let absolute = if path.is_absolute() { + Ok(path.to_path_buf()) + } else { + std::env::current_dir() + .map(|cwd| cwd.join(path)) + .map_err(|err| format!("failed to resolve plugin path: {err}")) + }?; + resolve_existing_path_prefix(&absolute) +} + +fn resolve_contained_host_path( + value: &str, + root: &Path, + allowed_root: &Path, +) -> Result { + let value = Path::new(value); + let path = if value.is_absolute() { + value.to_path_buf() + } else { + root.join(value) + }; + let path = resolve_existing_path_prefix(&path)?; + if !path.starts_with(allowed_root) { + return Err(format!( + "expanded path `{}` must remain within `{}`", + value.display(), + allowed_root.display() + )); + } + Ok(path) +} + +fn resolve_existing_path_prefix(path: &Path) -> Result { + let mut existing = path.to_path_buf(); + let mut missing_components = Vec::::new(); + loop { + match std::fs::canonicalize(&existing) { + Ok(mut resolved) => { + for component in missing_components.iter().rev() { + resolved.push(component); + } + return Ok(lexical_normalize(&resolved)); + } + Err(err) if err.kind() == std::io::ErrorKind::NotFound => { + if std::fs::symlink_metadata(&existing) + .is_ok_and(|metadata| metadata.file_type().is_symlink()) + { + return Err(format!( + "failed to resolve symlinked path `{}`", + path.display() + )); + } + let Some(component) = existing.components().next_back() else { + return Err(format!( + "failed to resolve path `{}`: {err}", + path.display() + )); + }; + if matches!( + component, + std::path::Component::Prefix(_) | std::path::Component::RootDir + ) { + return Err(format!( + "failed to resolve path `{}`: {err}", + path.display() + )); + } + missing_components.push(component.as_os_str().to_os_string()); + if !existing.pop() { + return Err(format!( + "failed to resolve path `{}`: {err}", + path.display() + )); + } + } + Err(err) => { + return Err(format!( + "failed to resolve path `{}`: {err}", + path.display() + )); + } + } + } +} + +fn host_path_string(path: &Path) -> String { + let rendered = path.to_string_lossy(); + #[cfg(windows)] + if let Some(path) = rendered.strip_prefix(r"\\?\") { + return path + .strip_prefix(r"UNC\") + .map(|path| format!(r"\\{path}")) + .unwrap_or_else(|| path.to_string()); + } + rendered.into_owned() +} + +fn is_portable_relative_path(value: &str) -> bool { + value + .strip_prefix("./") + .is_some_and(is_portable_path_suffix) +} + +fn is_portable_path_suffix(value: &str) -> bool { + !value.is_empty() && !value.contains('\\') +} + +fn is_valid_http_header_name(name: &str) -> bool { + !name.is_empty() + && name + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || b"!#$%&'*+-.^_`|~".contains(&byte)) +} + +fn lexical_normalize(path: &Path) -> PathBuf { + let mut normalized = PathBuf::new(); + for component in path.components() { + match component { + std::path::Component::CurDir => {} + std::path::Component::ParentDir => { + normalized.pop(); + } + component => normalized.push(component.as_os_str()), + } + } + normalized +} + +fn plugin_mcp_json_error(message: impl Into) -> serde_json::Error { + serde_json::Error::io(std::io::Error::new( + std::io::ErrorKind::InvalidData, + message.into(), + )) +} diff --git a/codex-rs/codex-mcp/src/auth_changes.rs b/codex-rs/codex-mcp/src/auth_changes.rs new file mode 100644 index 0000000000000000000000000000000000000000..78498bd071dbeef2941df284e16afd982e3efeef --- /dev/null +++ b/codex-rs/codex-mcp/src/auth_changes.rs @@ -0,0 +1,62 @@ +//! Forwards auth invalidations without credentials. The managed client owns the watcher. + +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Result; +use codex_login::AuthChangeState; +use codex_rmcp_client::RmcpClient; +use rmcp::model::ServerCapabilities; +use serde_json::json; +use tokio::sync::watch; +use tokio_util::task::AbortOnDropHandle; + +pub(crate) const CAPABILITY: &str = "codex/auth-change"; +const NOTIFICATION: &str = "notifications/codex/authChanged"; +const SEND_TIMEOUT: Duration = Duration::from_secs(5); + +pub(crate) async fn start( + client: Arc, + capabilities: &ServerCapabilities, + changes: Option>, +) -> Result>>> { + let Some(mut changes) = changes.filter(|_| { + capabilities + .experimental + .as_ref() + .is_some_and(|capabilities| capabilities.contains_key(CAPABILITY)) + }) else { + return Ok(None); + }; + + notify(&client, &mut changes).await?; + let task = tokio::spawn(async move { + while changes.changed().await.is_ok() { + if notify(&client, &mut changes).await.is_err() { + tracing::warn!("MCP auth invalidation delivery failed; closing connection"); + client.shutdown().await; + break; + } + } + }); + Ok(Some(Arc::new(AbortOnDropHandle::new(task)))) +} + +async fn notify(client: &RmcpClient, changes: &mut watch::Receiver) -> Result<()> { + let state = *changes.borrow_and_update(); + tokio::time::timeout( + SEND_TIMEOUT, + client.send_custom_notification( + NOTIFICATION, + Some(json!({ + "generation": state.generation, + "ownerGeneration": state.owner_generation, + })), + ), + ) + .await? +} + +#[cfg(test)] +#[path = "auth_changes_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/auth_changes_tests.rs b/codex-rs/codex-mcp/src/auth_changes_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..36de21ba0c647d8ab0b05a674e426c7188f0d216 --- /dev/null +++ b/codex-rs/codex-mcp/src/auth_changes_tests.rs @@ -0,0 +1,104 @@ +use super::*; +use codex_rmcp_client::InProcessTransportFactory; +use futures::FutureExt; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; +use rmcp::ServiceExt; +use rmcp::model::ClientCapabilities; +use rmcp::model::CustomNotification; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use tokio::sync::mpsc; +use tokio::time::timeout; + +#[derive(Clone)] +struct NotificationServer(mpsc::Sender); + +impl rmcp::ServerHandler for NotificationServer { + async fn on_custom_notification( + &self, + notification: CustomNotification, + _context: rmcp::service::NotificationContext, + ) { + self.0.send(json!(notification)).await.unwrap(); + } +} + +impl InProcessTransportFactory for NotificationServer { + fn open(&self) -> BoxFuture<'static, std::io::Result> { + let server = self.clone(); + async move { + let (client, transport) = tokio::io::duplex(/*max_buf_size*/ 4096); + tokio::spawn(async move { + let service = server.serve(transport).await.unwrap(); + service.waiting().await.unwrap(); + }); + Ok(client) + } + .boxed() + } +} + +#[tokio::test] +async fn auth_notifications_require_opt_in_and_follow_client_lifetime() -> Result<()> { + let (notifications, mut received) = mpsc::channel(/*buffer*/ 8); + let client = Arc::new( + RmcpClient::new_in_process_client(Arc::new(NotificationServer(notifications))).await?, + ); + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("test", "1"), + ), + Some(SEND_TIMEOUT), + Box::new(|_, _| async { anyhow::bail!("unexpected elicitation") }.boxed()), + ) + .await?; + let (changes, receiver) = watch::channel(AuthChangeState::default()); + let mut capabilities = ServerCapabilities::default(); + assert!( + start(Arc::clone(&client), &capabilities, Some(receiver.clone())) + .await? + .is_none() + ); + capabilities.experimental = Some([(CAPABILITY.to_string(), Default::default())].into()); + assert!( + start(Arc::clone(&client), &capabilities, /*changes*/ None) + .await? + .is_none() + ); + assert_eq!(received.try_recv(), Err(mpsc::error::TryRecvError::Empty)); + let watcher = start(Arc::clone(&client), &capabilities, Some(receiver)) + .await? + .unwrap(); + assert_eq!( + timeout(SEND_TIMEOUT, received.recv()).await?, + Some( + json!({"method": NOTIFICATION, "params": {"_meta": {}, "generation": 0, "ownerGeneration": 0}}) + ), + ); + changes.send_modify(|state| state.generation += 1); + assert_eq!( + timeout(SEND_TIMEOUT, received.recv()).await?, + Some( + json!({"method": NOTIFICATION, "params": {"_meta": {}, "generation": 1, "ownerGeneration": 0}}) + ), + ); + for _ in 0..2 { + changes.send_modify(|state| { + state.generation += 1; + state.owner_generation += 1; + }); + } + assert_eq!( + timeout(SEND_TIMEOUT, received.recv()).await?, + Some( + json!({"method": NOTIFICATION, "params": {"_meta": {}, "generation": 3, "ownerGeneration": 2}}) + ), + ); + drop(watcher); + timeout(SEND_TIMEOUT, changes.closed()).await?; + client.shutdown().await; + Ok(()) +} diff --git a/codex-rs/codex-mcp/src/auth_elicitation.rs b/codex-rs/codex-mcp/src/auth_elicitation.rs new file mode 100644 index 0000000000000000000000000000000000000000..27ad8a3e8873751bf9cf83c3d3727c8ac8b1cf61 --- /dev/null +++ b/codex-rs/codex-mcp/src/auth_elicitation.rs @@ -0,0 +1,435 @@ +//! Auth elicitation helpers. +//! +//! This module owns protocol-neutral auth elicitation parsing and payload shaping. +//! Session orchestration stays in `codex-core`. + +use codex_protocol::mcp::CallToolResult; +use serde::Serialize; + +pub const MCP_TOOL_CODEX_APPS_META_KEY: &str = "_codex_apps"; +pub const CONNECTOR_AUTH_FAILURE_META_KEY: &str = "connector_auth_failure"; +pub const CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY: &str = "is_auth_failure"; +pub const CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY: &str = "auth_reason"; +pub const CONNECTOR_AUTH_FAILURE_CONNECTOR_ID_KEY: &str = "connector_id"; +pub const CONNECTOR_AUTH_FAILURE_LINK_ID_KEY: &str = "link_id"; +pub const CONNECTOR_AUTH_FAILURE_ERROR_CODE_KEY: &str = "error_code"; +pub const CONNECTOR_AUTH_FAILURE_ERROR_HTTP_STATUS_CODE_KEY: &str = "error_http_status_code"; +pub const CONNECTOR_AUTH_FAILURE_ERROR_ACTION_KEY: &str = "error_action"; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct CodexAppsConnectorAuthFailure { + pub connector_id: String, + pub connector_name: String, + pub install_url: String, + pub auth_reason: Option, + pub link_id: Option, + pub error_code: Option, + pub error_http_status_code: Option, + pub error_action: Option, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct CodexAppsAuthElicitation { + pub meta: serde_json::Value, + pub message: String, + pub url: String, + pub elicitation_id: String, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct CodexAppsAuthElicitationPlan { + pub auth_failure: CodexAppsConnectorAuthFailure, + pub elicitation: CodexAppsAuthElicitation, +} + +#[derive(Serialize)] +struct CodexAppsConnectorAuthFailureMeta<'a> { + is_auth_failure: bool, + connector_id: &'a str, + connector_name: &'a str, + install_url: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + auth_reason: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + link_id: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + error_code: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + error_http_status_code: Option, + #[serde(skip_serializing_if = "Option::is_none")] + error_action: Option<&'a str>, +} + +pub fn connector_auth_failure_from_tool_result( + result: &CallToolResult, + connector_id: Option<&str>, + connector_name: Option<&str>, + install_url: Option, +) -> Option { + let connector_id = connector_id + .map(str::trim) + .filter(|connector_id| !connector_id.is_empty())?; + let auth_failure = connector_auth_failure_metadata(result, connector_id)?; + let connector_name = connector_name + .map(str::trim) + .filter(|name| !name.is_empty()) + .unwrap_or(connector_id) + .to_string(); + + Some(CodexAppsConnectorAuthFailure { + connector_id: connector_id.to_string(), + connector_name, + install_url: install_url?, + auth_reason: string_auth_failure_field( + auth_failure, + CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY, + ), + link_id: string_auth_failure_field(auth_failure, CONNECTOR_AUTH_FAILURE_LINK_ID_KEY), + error_code: string_auth_failure_field(auth_failure, CONNECTOR_AUTH_FAILURE_ERROR_CODE_KEY), + error_http_status_code: auth_failure + .get(CONNECTOR_AUTH_FAILURE_ERROR_HTTP_STATUS_CODE_KEY) + .and_then(serde_json::Value::as_i64), + error_action: string_auth_failure_field( + auth_failure, + CONNECTOR_AUTH_FAILURE_ERROR_ACTION_KEY, + ), + }) +} + +pub fn is_connector_auth_failure_from_tool_result( + result: &CallToolResult, + connector_id: Option<&str>, +) -> bool { + connector_id + .map(str::trim) + .filter(|connector_id| !connector_id.is_empty()) + .is_some_and(|connector_id| connector_auth_failure_metadata(result, connector_id).is_some()) +} + +fn connector_auth_failure_metadata<'a>( + result: &'a CallToolResult, + connector_id: &str, +) -> Option<&'a serde_json::Map> { + if result.is_error != Some(true) { + return None; + } + + let auth_failure = result + .meta + .as_ref()? + .as_object()? + .get(MCP_TOOL_CODEX_APPS_META_KEY)? + .as_object()? + .get(CONNECTOR_AUTH_FAILURE_META_KEY)? + .as_object()?; + if auth_failure + .get(CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY) + .and_then(serde_json::Value::as_bool) + != Some(true) + { + return None; + } + if let Some(auth_failure_connector_id) = + string_auth_failure_field(auth_failure, CONNECTOR_AUTH_FAILURE_CONNECTOR_ID_KEY) + && auth_failure_connector_id != connector_id + { + return None; + } + + Some(auth_failure) +} + +pub fn build_auth_elicitation_plan( + call_id: &str, + result: &CallToolResult, + connector_id: Option<&str>, + connector_name: Option<&str>, + install_url: Option, +) -> Option { + let auth_failure = + connector_auth_failure_from_tool_result(result, connector_id, connector_name, install_url)?; + let elicitation = build_auth_elicitation(call_id, &auth_failure); + Some(CodexAppsAuthElicitationPlan { + auth_failure, + elicitation, + }) +} + +pub fn build_auth_elicitation( + call_id: &str, + auth_failure: &CodexAppsConnectorAuthFailure, +) -> CodexAppsAuthElicitation { + CodexAppsAuthElicitation { + meta: serde_json::json!({ + MCP_TOOL_CODEX_APPS_META_KEY: { + CONNECTOR_AUTH_FAILURE_META_KEY: CodexAppsConnectorAuthFailureMeta { + is_auth_failure: true, + connector_id: &auth_failure.connector_id, + connector_name: &auth_failure.connector_name, + install_url: &auth_failure.install_url, + auth_reason: auth_failure.auth_reason.as_deref(), + link_id: auth_failure.link_id.as_deref(), + error_code: auth_failure.error_code.as_deref(), + error_http_status_code: auth_failure.error_http_status_code, + error_action: auth_failure.error_action.as_deref(), + }, + }, + }), + message: auth_elicitation_message(auth_failure), + url: auth_failure.install_url.clone(), + elicitation_id: auth_elicitation_id(call_id), + } +} + +pub fn auth_elicitation_completed_result( + auth_failure: &CodexAppsConnectorAuthFailure, + meta: Option, +) -> CallToolResult { + CallToolResult { + content: vec![serde_json::json!({ + "type": "text", + "text": format!( + "Authentication for {} was requested and accepted. Retry this tool call now.", + auth_failure.connector_name + ), + })], + structured_content: None, + is_error: Some(true), + meta, + } +} + +pub fn auth_elicitation_id(call_id: &str) -> String { + format!("codex_apps_auth_{call_id}") +} + +fn string_auth_failure_field( + auth_failure: &serde_json::Map, + key: &str, +) -> Option { + auth_failure + .get(key) + .and_then(serde_json::Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToString::to_string) +} + +fn auth_elicitation_message(auth_failure: &CodexAppsConnectorAuthFailure) -> String { + match auth_failure.auth_reason.as_deref() { + Some("oauth_upgrade_required") => format!( + "Reconnect {} on ChatGPT to grant the permissions needed for this request.", + auth_failure.connector_name + ), + Some("reauthentication_required") => format!( + "Reconnect {} on ChatGPT to restore access for this request.", + auth_failure.connector_name + ), + Some("missing_link") => format!( + "Sign in to {} on ChatGPT to use it in Codex.", + auth_failure.connector_name + ), + _ => format!( + "Sign in to {} on ChatGPT to continue.", + auth_failure.connector_name + ), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use pretty_assertions::assert_eq; + + fn auth_failure_result() -> CallToolResult { + CallToolResult { + content: vec![serde_json::json!({ + "type": "text", + "text": "Connector reauthentication required", + })], + structured_content: None, + is_error: Some(true), + meta: Some(serde_json::json!({ + MCP_TOOL_CODEX_APPS_META_KEY: { + CONNECTOR_AUTH_FAILURE_META_KEY: { + CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY: true, + CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY: "reauthentication_required", + CONNECTOR_AUTH_FAILURE_CONNECTOR_ID_KEY: "connector_calendar", + "connector_name": "Untrusted Calendar", + CONNECTOR_AUTH_FAILURE_LINK_ID_KEY: "link_123", + CONNECTOR_AUTH_FAILURE_ERROR_CODE_KEY: "UNAUTHORIZED", + CONNECTOR_AUTH_FAILURE_ERROR_HTTP_STATUS_CODE_KEY: 401, + CONNECTOR_AUTH_FAILURE_ERROR_ACTION_KEY: "TRIGGER_REAUTHENTICATION", + }, + }, + })), + } + } + + #[test] + fn parses_auth_failure_from_trusted_connector_metadata() { + assert_eq!( + connector_auth_failure_from_tool_result( + &auth_failure_result(), + Some("connector_calendar"), + Some("Google Calendar"), + Some("https://chatgpt.com/apps/google-calendar/connector_calendar".to_string()), + ), + Some(CodexAppsConnectorAuthFailure { + connector_id: "connector_calendar".to_string(), + connector_name: "Google Calendar".to_string(), + install_url: "https://chatgpt.com/apps/google-calendar/connector_calendar" + .to_string(), + auth_reason: Some("reauthentication_required".to_string()), + link_id: Some("link_123".to_string()), + error_code: Some("UNAUTHORIZED".to_string()), + error_http_status_code: Some(401), + error_action: Some("TRIGGER_REAUTHENTICATION".to_string()), + }) + ); + } + + #[test] + fn rejects_missing_or_mismatched_connector_ids() { + assert_eq!( + connector_auth_failure_from_tool_result( + &auth_failure_result(), + /*connector_id*/ None, + Some("Google Calendar"), + Some("https://chatgpt.com/apps/google-calendar/connector_calendar".to_string()), + ), + None + ); + assert_eq!( + connector_auth_failure_from_tool_result( + &auth_failure_result(), + Some("connector_drive"), + Some("Google Drive"), + Some("https://chatgpt.com/apps/google-drive/connector_drive".to_string()), + ), + None + ); + } + + #[test] + fn detects_auth_failure_without_an_install_url() { + let result = auth_failure_result(); + + assert_eq!( + is_connector_auth_failure_from_tool_result(&result, Some("connector_calendar")), + true + ); + assert_eq!( + connector_auth_failure_from_tool_result( + &result, + Some("connector_calendar"), + Some("Google Calendar"), + /*install_url*/ None, + ), + None + ); + } + + #[test] + fn auth_failure_detection_requires_trusted_connector_identity_and_auth_flag() { + let result = auth_failure_result(); + assert_eq!( + is_connector_auth_failure_from_tool_result(&result, /*connector_id*/ None), + false + ); + assert_eq!( + is_connector_auth_failure_from_tool_result(&result, Some("connector_drive")), + false + ); + + let mut ordinary_error = result.clone(); + ordinary_error.meta.as_mut().expect("auth metadata")[MCP_TOOL_CODEX_APPS_META_KEY] + [CONNECTOR_AUTH_FAILURE_META_KEY][CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY] = + serde_json::Value::Bool(false); + assert_eq!( + is_connector_auth_failure_from_tool_result(&ordinary_error, Some("connector_calendar"),), + false + ); + + let mut successful_result = result; + successful_result.is_error = Some(false); + assert_eq!( + is_connector_auth_failure_from_tool_result( + &successful_result, + Some("connector_calendar"), + ), + false + ); + } + + #[test] + fn detects_each_supported_connector_auth_reason() { + for auth_reason in [ + "missing_link", + "oauth_upgrade_required", + "reauthentication_required", + ] { + let mut result = auth_failure_result(); + result.meta.as_mut().expect("auth metadata")[MCP_TOOL_CODEX_APPS_META_KEY] + [CONNECTOR_AUTH_FAILURE_META_KEY][CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY] = + serde_json::Value::String(auth_reason.to_string()); + + assert_eq!( + is_connector_auth_failure_from_tool_result(&result, Some("connector_calendar")), + true + ); + } + } + + #[test] + fn builds_url_elicitation_payload() { + let auth_failure = connector_auth_failure_from_tool_result( + &auth_failure_result(), + Some("connector_calendar"), + Some("Google Calendar"), + Some("https://chatgpt.com/apps/google-calendar/connector_calendar".to_string()), + ) + .expect("auth failure"); + + assert_eq!( + build_auth_elicitation("call_123", &auth_failure), + CodexAppsAuthElicitation { + meta: serde_json::json!({ + MCP_TOOL_CODEX_APPS_META_KEY: { + CONNECTOR_AUTH_FAILURE_META_KEY: { + CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY: true, + CONNECTOR_AUTH_FAILURE_CONNECTOR_ID_KEY: "connector_calendar", + "connector_name": "Google Calendar", + "install_url": + "https://chatgpt.com/apps/google-calendar/connector_calendar", + CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY: "reauthentication_required", + CONNECTOR_AUTH_FAILURE_LINK_ID_KEY: "link_123", + CONNECTOR_AUTH_FAILURE_ERROR_CODE_KEY: "UNAUTHORIZED", + CONNECTOR_AUTH_FAILURE_ERROR_HTTP_STATUS_CODE_KEY: 401, + CONNECTOR_AUTH_FAILURE_ERROR_ACTION_KEY: "TRIGGER_REAUTHENTICATION", + }, + }, + }), + message: "Reconnect Google Calendar on ChatGPT to restore access for this request." + .to_string(), + url: "https://chatgpt.com/apps/google-calendar/connector_calendar".to_string(), + elicitation_id: "codex_apps_auth_call_123".to_string(), + } + ); + } + + #[test] + fn builds_auth_elicitation_plan() { + let plan = build_auth_elicitation_plan( + "call_123", + &auth_failure_result(), + Some("connector_calendar"), + Some("Google Calendar"), + Some("https://chatgpt.com/apps/google-calendar/connector_calendar".to_string()), + ) + .expect("auth elicitation plan"); + + assert_eq!(plan.auth_failure.connector_name, "Google Calendar"); + assert_eq!(plan.elicitation.elicitation_id, "codex_apps_auth_call_123"); + } +} diff --git a/codex-rs/codex-mcp/src/binding.rs b/codex-rs/codex-mcp/src/binding.rs new file mode 100644 index 0000000000000000000000000000000000000000..ade3747729372a203a4ff60e9263de3316a994bf --- /dev/null +++ b/codex-rs/codex-mcp/src/binding.rs @@ -0,0 +1,389 @@ +//! Immutable MCP catalog and execution handles. + +use std::collections::HashMap; +use std::fmt; +use std::future::Future; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use codex_config::AppToolApproval; +use codex_protocol::mcp::CallToolResult; +use codex_protocol::models::PermissionProfile; +use rmcp::model::ListResourceTemplatesResult; +use rmcp::model::ListResourcesResult; +use rmcp::model::PaginatedRequestParams; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ReadResourceResult; +use rmcp::model::Resource; +use rmcp::model::ResourceTemplate; +use serde_json::Value as JsonValue; + +use crate::McpConfig; +use crate::binding_clients::McpBindingClients; +use crate::client_tool_catalog::ToolCatalogSnapshot; +use crate::connection_manager::McpConnectionSet; +use crate::rmcp_client::ManagedClient; +use crate::server::McpServerMetadata; +use crate::tools::ToolInfo; + +/// The exact tool catalog and execution handles shared by compatible sampling steps. +pub struct McpBinding { + connections: Arc, + clients: Arc, + config: Arc, + plugins_available: bool, + tools: Vec, + calls: HashMap<(String, String), PreparedMcpCall>, +} + +impl McpBinding { + /// Creates an empty binding for tests and callers without a materialized runtime. + pub fn empty(config: Arc) -> Self { + Self::new( + Arc::new(McpConnectionSet::empty(config.prefix_mcp_tool_names)), + Arc::new(McpBindingClients::new(HashMap::new())), + config, + /*plugins_available*/ false, + Vec::new(), + HashMap::new(), + ) + } + + pub(crate) fn new( + connections: Arc, + clients: Arc, + config: Arc, + plugins_available: bool, + tools: Vec, + calls: HashMap<(String, String), PreparedMcpCall>, + ) -> Self { + Self { + connections, + clients, + config, + plugins_available, + tools, + calls, + } + } + + pub fn config(&self) -> &Arc { + &self.config + } + + pub fn plugins_available(&self) -> bool { + self.plugins_available + } + + /// Returns the frozen model-visible catalog captured for this binding. + pub fn tools(&self) -> &[ToolInfo] { + &self.tools + } + + /// Returns permitted tool metadata, including app-only tools. + pub fn tool_info(&self, server: &str, tool: &str) -> Option<&ToolInfo> { + self.calls + .get(&(server.to_string(), tool.to_string())) + .map(PreparedMcpCall::tool_info) + } + + /// Binds a model-visible call to the exact client and metadata in this binding. + pub fn prepare_call(&self, server: &str, tool: &str) -> Option { + self.calls + .get(&(server.to_string(), tool.to_string())) + .filter(|call| crate::tool_is_model_visible(call.tool_info())) + .cloned() + } + + pub fn has_servers(&self) -> bool { + self.connections.has_servers() + } + + pub async fn list_resources( + &self, + server: &str, + params: Option, + ) -> Result { + if self.clients.client(server).is_some() { + self.clients.list_resources(server, params).await + } else { + self.connections.list_resources(server, params).await + } + } + + pub async fn list_all_resources( + &self, + include_server: impl Fn(&str) -> bool, + ) -> HashMap> { + self.clients.list_all_resources(include_server).await + } + + pub async fn list_resource_templates( + &self, + server: &str, + params: Option, + ) -> Result { + if self.clients.client(server).is_some() { + self.clients.list_resource_templates(server, params).await + } else { + self.connections + .list_resource_templates(server, params) + .await + } + } + + pub async fn list_all_resource_templates( + &self, + include_server: impl Fn(&str) -> bool, + ) -> HashMap> { + self.clients + .list_all_resource_templates(include_server) + .await + } + + pub async fn read_resource( + &self, + server: &str, + params: ReadResourceRequestParams, + ) -> Result { + if self.clients.client(server).is_some() { + self.clients.read_resource(server, params).await + } else { + self.connections.read_resource(server, params).await + } + } +} + +impl fmt::Debug for McpBinding { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("McpBinding") + .field("tools", &self.tools) + .field("prepared_call_count", &self.calls.len()) + .finish_non_exhaustive() + } +} + +/// A call bound to the exact client, tool, timeout, and server metadata seen by +/// one [`McpBinding`]. +#[derive(Clone)] +pub struct PreparedMcpCall { + connections: Arc, + client: Arc, + config: Arc, + catalog_snapshot: Arc, + tool_info: ToolInfo, + server_name: String, + server_metadata: McpServerMetadata, + plugin_id: Option, + selected_plugin_server: bool, +} + +impl PreparedMcpCall { + #[expect( + clippy::too_many_arguments, + reason = "the exact call authority stays together" + )] + pub(crate) fn new( + connections: Arc, + client: Arc, + config: Arc, + catalog_snapshot: Arc, + tool_info: ToolInfo, + server_metadata: McpServerMetadata, + plugin_id: Option, + selected_plugin_server: bool, + ) -> Option { + let server_name = tool_info.server_name.clone(); + config.permission_profile_for_server(&server_name)?; + Some(Self { + connections, + client, + config, + catalog_snapshot, + tool_info, + server_name, + server_metadata, + plugin_id, + selected_plugin_server, + }) + } + + pub fn tool_info(&self) -> &ToolInfo { + &self.tool_info + } + + /// Returns the configuration and approval authority captured with this client. + pub fn config(&self) -> &McpConfig { + &self.config + } + + /// Returns the owner permissions validated when this immutable call was prepared. + pub fn permission_profile(&self) -> &PermissionProfile { + let Some(permission_profile) = self.config.permission_profile_for_server(&self.server_name) + else { + unreachable!("prepared MCP calls retain their immutable permission authority"); + }; + permission_profile + } + + pub fn server_name(&self) -> &str { + &self.server_name + } + + /// Returns whether this call is bound to the host-owned Codex Apps server. + pub fn is_host_owned_apps(&self) -> bool { + self.config + .mcp_server_catalog + .server(&self.server_name) + .is_some_and(|registration| { + registration + .source() + .is_host_owned_apps(&self.server_name, registration.config()) + }) + } + + pub fn server_origin(&self) -> Option<&str> { + self.server_metadata + .origin + .as_ref() + .map(super::server::McpServerOrigin::as_str) + } + + pub fn server_environment_id(&self) -> &str { + &self.server_metadata.environment_id + } + + pub fn server_pollutes_memory(&self) -> bool { + self.server_metadata.pollutes_memory + } + + pub fn tool_approval_mode(&self) -> AppToolApproval { + self.server_metadata + .tool_approval_mode(&self.tool_info.tool.name) + } + + /// Returns the explicit output budget captured with this call's effective server config. + pub fn output_token_limit(&self) -> Option { + self.config + .mcp_server_catalog + .server(&self.server_name)? + .config() + .tools + .get(self.tool_info.tool.name.as_ref())? + .output_token_limit + .map(std::num::NonZeroUsize::get) + } + + pub fn plugin_id(&self) -> Option<&str> { + self.plugin_id.as_deref() + } + + pub fn is_selected_plugin_server(&self) -> bool { + self.selected_plugin_server + } + + pub async fn server_supports_sandbox_state_meta_capability(&self) -> Result { + Ok(self.client.server_supports_sandbox_state_meta_capability) + } + + pub async fn call( + &self, + arguments: Option, + meta: Option, + timeout: Option, + ) -> Result { + self.call_with_preparation(timeout, || async move { Ok((arguments, meta)) }) + .await + } + + /// Runs irreversible call preparation and execution under the authority of + /// this call's captured catalog and the extensions owned by the Codex session. + /// A caller-supplied timeout can further restrict the server's configured timeout. + pub async fn call_with_preparation( + &self, + requested_timeout: Option, + prepare: F, + ) -> Result + where + F: FnOnce() -> Fut, + Fut: Future, Option)>>, + { + let effective_timeout = match (self.client.tool_timeout, requested_timeout) { + (Some(server_timeout), Some(requested_timeout)) => { + Some(server_timeout.min(requested_timeout)) + } + (server_timeout, requested_timeout) => server_timeout.or(requested_timeout), + }; + let tool_name = self.tool_info.tool.name.to_string(); + self.client + .tool_catalog + .run_with_snapshot(&self.catalog_snapshot, || async { + let (arguments, meta) = prepare().await?; + let timeout_deadline = + effective_timeout.map(|timeout| tokio::time::Instant::now() + timeout); + let add_trusted_access_context = self.connections.add_trusted_access_context( + &self.tool_info, + &self.server_metadata, + arguments.as_ref(), + meta, + ); + let meta = match effective_timeout.zip(timeout_deadline) { + Some((timeout, deadline)) => { + tokio::time::timeout_at(deadline, add_trusted_access_context) + .await + .map_err(|_| { + anyhow::anyhow!("timed out awaiting tools/call after {timeout:.0?}") + })? + } + None => add_trusted_access_context.await, + }; + let remaining_timeout = match effective_timeout.zip(timeout_deadline) { + Some((timeout, deadline)) => { + let remaining = deadline.saturating_duration_since(tokio::time::Instant::now()); + if remaining.is_zero() { + return Err(anyhow::anyhow!( + "timed out awaiting tools/call after {timeout:.0?}" + )); + } + Some(remaining) + } + None => None, + }; + self.client + .client + .call_tool(tool_name.clone(), arguments, meta, remaining_timeout) + .await + .with_context(|| format!("tool call failed for `{}/{tool_name}`", self.server_name)) + }) + .await + .ok_or_else(|| anyhow::anyhow!( + "tool call rejected because the catalog changed after `{}/{tool_name}` was prepared", + self.server_name + ))? + .map(call_tool_result_from_rmcp) + } +} + +pub(crate) fn call_tool_result_from_rmcp(result: rmcp::model::CallToolResult) -> CallToolResult { + let content = result + .content + .into_iter() + .map(|content| { + serde_json::to_value(content) + .unwrap_or_else(|_| JsonValue::String("".to_string())) + }) + .collect(); + CallToolResult { + content, + structured_content: result.structured_content, + is_error: result.is_error, + meta: result.meta.and_then(|meta| serde_json::to_value(meta).ok()), + } +} + +#[cfg(test)] +#[path = "binding_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/binding_clients.rs b/codex-rs/codex-mcp/src/binding_clients.rs new file mode 100644 index 0000000000000000000000000000000000000000..b2b1c174b8d5518a6b46aca5ac79ecf74fd7e15a --- /dev/null +++ b/codex-rs/codex-mcp/src/binding_clients.rs @@ -0,0 +1,156 @@ +use std::collections::HashMap; +use std::sync::Arc; + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use rmcp::model::ListResourceTemplatesResult; +use rmcp::model::ListResourcesResult; +use rmcp::model::PaginatedRequestParams; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ReadResourceResult; +use rmcp::model::Resource; +use rmcp::model::ResourceTemplate; +use tokio::task::JoinSet; +use tracing::warn; + +use crate::pagination::collect_paginated; +use crate::rmcp_client::ManagedClient; + +/// The ready clients captured for one model step. +pub(crate) struct McpBindingClients { + clients: HashMap>, +} + +impl McpBindingClients { + pub(crate) fn new(clients: HashMap>) -> Self { + Self { clients } + } + + pub(crate) fn client(&self, server: &str) -> Option> { + self.clients.get(server).cloned() + } + + pub(crate) async fn list_resources( + &self, + server: &str, + params: Option, + ) -> Result { + let managed = self + .client(server) + .ok_or_else(|| anyhow!("MCP server '{server}' was not ready for this step"))?; + managed + .client + .list_resources(params, managed.tool_timeout) + .await + .with_context(|| format!("resources/list failed for `{server}`")) + } + + pub(crate) async fn list_resource_templates( + &self, + server: &str, + params: Option, + ) -> Result { + let managed = self + .client(server) + .ok_or_else(|| anyhow!("MCP server '{server}' was not ready for this step"))?; + managed + .client + .list_resource_templates(params, managed.tool_timeout) + .await + .with_context(|| format!("resources/templates/list failed for `{server}`")) + } + + pub(crate) async fn read_resource( + &self, + server: &str, + params: ReadResourceRequestParams, + ) -> Result { + let managed = self + .client(server) + .ok_or_else(|| anyhow!("MCP server '{server}' was not ready for this step"))?; + let uri = params.uri.clone(); + managed + .client + .read_resource(params, managed.tool_timeout) + .await + .with_context(|| format!("resources/read failed for `{server}` ({uri})")) + } + + pub(crate) async fn list_all_resources( + &self, + include_server: impl Fn(&str) -> bool, + ) -> HashMap> { + let mut join_set = JoinSet::new(); + for (server_name, managed) in self + .clients + .iter() + .filter(|(server_name, _)| include_server(server_name)) + { + let server_name = server_name.clone(); + let client = Arc::clone(&managed.client); + let timeout = managed.tool_timeout; + join_set.spawn(async move { + let resources = collect_paginated("resources/list", timeout, |params| { + let client = Arc::clone(&client); + async move { + let response = client.list_resources(params, timeout).await?; + Ok((response.resources, response.next_cursor)) + } + }) + .await; + (server_name, resources) + }); + } + collect_resource_results(&mut join_set, "resources").await + } + + pub(crate) async fn list_all_resource_templates( + &self, + include_server: impl Fn(&str) -> bool, + ) -> HashMap> { + let mut join_set = JoinSet::new(); + for (server_name, managed) in self + .clients + .iter() + .filter(|(server_name, _)| include_server(server_name)) + { + let server_name = server_name.clone(); + let client = Arc::clone(&managed.client); + let timeout = managed.tool_timeout; + join_set.spawn(async move { + let templates = collect_paginated("resources/templates/list", timeout, |params| { + let client = Arc::clone(&client); + async move { + let response = client.list_resource_templates(params, timeout).await?; + Ok((response.resource_templates, response.next_cursor)) + } + }) + .await; + (server_name, templates) + }); + } + collect_resource_results(&mut join_set, "resource templates").await + } +} + +async fn collect_resource_results( + join_set: &mut JoinSet<(String, Result>)>, + kind: &str, +) -> HashMap> { + let mut resources = HashMap::new(); + while let Some(result) = join_set.join_next().await { + match result { + Ok((server, Ok(server_resources))) => { + resources.insert(server, server_resources); + } + Ok((server, Err(error))) => { + warn!("Failed to list {kind} for MCP server '{server}': {error:#}"); + } + Err(error) => { + warn!("Task panic when listing {kind} for MCP server: {error:#}"); + } + } + } + resources +} diff --git a/codex-rs/codex-mcp/src/binding_tests.rs b/codex-rs/codex-mcp/src/binding_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..447a92364d5f3d78a45a4e1f7ca9a32c7e4f6846 --- /dev/null +++ b/codex-rs/codex-mcp/src/binding_tests.rs @@ -0,0 +1,401 @@ +use std::collections::HashMap; +use std::io; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use codex_config::AppToolApproval; +use codex_config::Constrained; +use codex_config::types::ApprovalsReviewer; +use codex_protocol::mcp::McpServerInfo; +use codex_protocol::models::PermissionProfile; +use codex_protocol::protocol::AskForApproval; +use codex_rmcp_client::InProcessTransportFactory; +use codex_rmcp_client::RmcpClient; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use rmcp::model::JsonObject; +use rmcp::model::Tool; +use tokio::io::DuplexStream; +use tokio::sync::Notify; + +use super::McpBinding; +use super::PreparedMcpCall; +use crate::binding_clients::McpBindingClients; +use crate::client_tool_catalog::ClientToolCatalog; +use crate::connection_manager::McpConnectionSet; +use crate::rmcp_client::ManagedClient; +use crate::server::McpServerMetadata; +use crate::server::McpServerOrigin; +use crate::tools::ToolInfo; + +const SERVER_NAME: &str = "docs"; +const TOOL_NAME: &str = "search"; + +struct TestInProcessTransportFactory; + +impl InProcessTransportFactory for TestInProcessTransportFactory { + fn open(&self) -> futures::future::BoxFuture<'static, io::Result> { + async { + let (client_stream, _server_stream) = tokio::io::duplex(1); + Ok(client_stream) + } + .boxed() + } +} + +struct TestStep { + step: Arc, + client: Arc, + tool_catalog: Arc, +} + +async fn test_step( + label: &str, + approval_mode: AppToolApproval, + supports_sandbox_state_meta: bool, +) -> TestStep { + let tool = ToolInfo { + server_name: SERVER_NAME.to_string(), + supports_parallel_tool_calls: false, + server_origin: None, + callable_name: TOOL_NAME.to_string(), + callable_namespace: SERVER_NAME.to_string(), + namespace_description: None, + tool: Tool::new( + TOOL_NAME.to_string(), + format!("{label} catalog"), + Arc::new(JsonObject::default()), + ), + openai_file_input_optional_fields: Default::default(), + connector_id: None, + connector_name: None, + plugin_display_names: Vec::new(), + }; + let client = Arc::new( + RmcpClient::new_in_process_client(Arc::new(TestInProcessTransportFactory)) + .await + .expect("create in-process MCP client"), + ); + let tool_catalog = Arc::new(ClientToolCatalog::new( + vec![tool.clone()], + /*updates*/ None, + )); + let managed_client = Arc::new(ManagedClient { + _auth_change_notifications: None, + client: Arc::clone(&client), + server_info: McpServerInfo { + name: label.to_string(), + title: Some(format!("{label} server")), + version: "1.0.0".to_string(), + description: None, + icons: None, + website_url: None, + }, + tool_catalog: Arc::clone(&tool_catalog), + tool_timeout: None, + server_instructions: None, + server_supports_sandbox_state_meta_capability: supports_sandbox_state_meta, + codex_apps_tools_cache_context: None, + }); + let clients = Arc::new(McpBindingClients::new(HashMap::from([( + SERVER_NAME.to_string(), + Arc::clone(&managed_client), + )]))); + let connections = Arc::new(McpConnectionSet::empty(/*prefix_mcp_tool_names*/ true)); + let mut config = crate::mcp::tests::test_mcp_config(std::env::temp_dir()); + if label == "old" { + config.approval_policy = Constrained::allow_any(AskForApproval::Never); + config.permission_profile = PermissionProfile::Disabled; + } else { + config.approvals_reviewer = ApprovalsReviewer::AutoReview; + } + config + .server_permission_profiles + .insert(SERVER_NAME.to_string(), config.permission_profile.clone()); + let config = Arc::new(config); + let prepared = PreparedMcpCall::new( + Arc::clone(&connections), + managed_client, + Arc::clone(&config), + tool_catalog.read(Arc::new).await, + tool.clone(), + McpServerMetadata { + environment_id: format!("{label}-environment"), + pollutes_memory: label == "old", + origin: Some(McpServerOrigin::StreamableHttp(format!( + "https://{label}.example" + ))), + supports_parallel_tool_calls: false, + default_tools_approval_mode: Some(approval_mode), + tool_approval_modes: HashMap::new(), + }, + Some(format!("{label}-plugin")), + label == "old", + ) + .expect("test call should retain its thread-owned permission profile"); + let calls = HashMap::from([((SERVER_NAME.to_string(), TOOL_NAME.to_string()), prepared)]); + + TestStep { + step: Arc::new(McpBinding::new( + connections, + clients, + config, + /*plugins_available*/ false, + vec![tool], + calls, + )), + client, + tool_catalog, + } +} + +#[tokio::test] +async fn prepared_call_keeps_captured_connection_and_authority_after_refresh() -> anyhow::Result<()> +{ + let old = test_step( + "old", + AppToolApproval::Prompt, + /*supports_sandbox_state_meta*/ true, + ) + .await; + let old_call = old + .step + .prepare_call(SERVER_NAME, TOOL_NAME) + .expect("old step should prepare the advertised tool"); + let old_connections = Arc::downgrade(&old.step.connections); + + let new = test_step( + "new", + AppToolApproval::Approve, + /*supports_sandbox_state_meta*/ false, + ) + .await; + let new_call = new + .step + .prepare_call(SERVER_NAME, TOOL_NAME) + .expect("new step should prepare the advertised tool"); + + assert_eq!( + ( + old.step.tools()[0].tool.description.as_deref(), + old_call.tool_info().tool.description.as_deref(), + old_call.server_origin(), + old_call.server_environment_id(), + old_call.server_pollutes_memory(), + old_call.tool_approval_mode(), + old_call.plugin_id(), + old_call.is_selected_plugin_server(), + old_call + .server_supports_sandbox_state_meta_capability() + .await?, + ), + ( + Some("old catalog"), + Some("old catalog"), + Some("https://old.example"), + "old-environment", + true, + AppToolApproval::Prompt, + Some("old-plugin"), + true, + true, + ) + ); + assert_eq!( + ( + new.step.tools()[0].tool.description.as_deref(), + new_call.tool_info().tool.description.as_deref(), + new_call.server_environment_id(), + new_call.tool_approval_mode(), + ), + ( + Some("new catalog"), + Some("new catalog"), + "new-environment", + AppToolApproval::Approve, + ) + ); + assert!(Arc::ptr_eq(&old_call.client.client, &old.client)); + assert!(!Arc::ptr_eq(&old.client, &new.client)); + assert_eq!( + ( + old_call.config().approval_policy.value(), + old_call.permission_profile(), + old_call.config().approvals_reviewer, + ), + ( + AskForApproval::Never, + &PermissionProfile::Disabled, + ApprovalsReviewer::User, + ) + ); + assert_eq!( + ( + new_call.config().approval_policy.value(), + new_call.config().approvals_reviewer, + ), + (AskForApproval::OnRequest, ApprovalsReviewer::AutoReview) + ); + + drop(old.step); + assert!( + old_connections.upgrade().is_some(), + "the prepared call should keep its captured connection set alive" + ); + drop(old_call); + assert!( + old_connections.upgrade().is_none(), + "the captured connection set should be released with the prepared call" + ); + Ok(()) +} + +#[tokio::test] +async fn prepared_call_does_not_reroute_after_captured_connection_closes() { + let old = test_step( + "old", + AppToolApproval::Prompt, + /*supports_sandbox_state_meta*/ true, + ) + .await; + let old_call = old + .step + .prepare_call(SERVER_NAME, TOOL_NAME) + .expect("old step should prepare the advertised tool"); + let new = test_step( + "new", + AppToolApproval::Approve, + /*supports_sandbox_state_meta*/ false, + ) + .await; + assert!(!Arc::ptr_eq(&old.client, &new.client)); + + old.client.shutdown().await; + + let error = old_call + .call( + Some(serde_json::json!({"query": "codex"})), + /*meta*/ None, + /*timeout*/ None, + ) + .await + .expect_err("a call bound to a closed connection must fail"); + assert!( + format!("{error:#}").contains("MCP client is shut down"), + "the prepared call should fail on its captured client: {error:#}" + ); +} + +#[tokio::test] +async fn prepared_call_is_rejected_after_catalog_refresh() { + let step = test_step( + "old", + AppToolApproval::Prompt, + /*supports_sandbox_state_meta*/ true, + ) + .await; + let prepared = step + .step + .prepare_call(SERVER_NAME, TOOL_NAME) + .expect("step should prepare the advertised tool"); + + step.tool_catalog + .refresh( + || async { Ok((step.step.tools().to_vec(), ())) }, + |_, ()| {}, + ) + .await + .expect("refresh tool catalog"); + + let error = prepared + .call( + Some(serde_json::json!({"query": "codex"})), + /*meta*/ None, + /*timeout*/ None, + ) + .await + .expect_err("a call from an older catalog must be rejected"); + assert!( + format!("{error:#}").contains("catalog changed"), + "unexpected error: {error:#}" + ); +} + +#[tokio::test] +async fn stale_prepared_call_does_not_run_preparation() { + let step = test_step( + "old", + AppToolApproval::Prompt, + /*supports_sandbox_state_meta*/ true, + ) + .await; + let prepared = step + .step + .prepare_call(SERVER_NAME, TOOL_NAME) + .expect("step should prepare the advertised tool"); + step.tool_catalog + .refresh( + || async { Ok((step.step.tools().to_vec(), ())) }, + |_, ()| {}, + ) + .await + .expect("refresh tool catalog"); + let prepared_side_effect_ran = Arc::new(AtomicBool::new(false)); + let marker = Arc::clone(&prepared_side_effect_ran); + + prepared + .call_with_preparation(/*requested_timeout*/ None, || async move { + marker.store(true, Ordering::SeqCst); + Ok((None, None)) + }) + .await + .expect_err("a call from an older catalog must be rejected"); + + assert!(!prepared_side_effect_ran.load(Ordering::SeqCst)); +} + +#[tokio::test] +async fn preparation_holds_catalog_authority_until_it_finishes() { + let step = test_step( + "old", + AppToolApproval::Prompt, + /*supports_sandbox_state_meta*/ true, + ) + .await; + let prepared = step + .step + .prepare_call(SERVER_NAME, TOOL_NAME) + .expect("step should prepare the advertised tool"); + let preparation_started = Arc::new(Notify::new()); + let finish_preparation = Arc::new(Notify::new()); + let started = Arc::clone(&preparation_started); + let finish = Arc::clone(&finish_preparation); + let call = tokio::spawn(async move { + prepared + .call_with_preparation(/*requested_timeout*/ None, || async move { + started.notify_one(); + finish.notified().await; + Err(anyhow::anyhow!("stop after preparation")) + }) + .await + }); + + preparation_started.notified().await; + let refresh = step.tool_catalog.refresh( + || async { Ok((step.step.tools().to_vec(), ())) }, + |_, ()| {}, + ); + tokio::pin!(refresh); + assert!( + futures::poll!(&mut refresh).is_pending(), + "catalog replacement must wait for irreversible call preparation" + ); + finish_preparation.notify_one(); + call.await + .expect("call task should finish") + .expect_err("the test preparation should stop the call"); + refresh + .await + .expect("catalog refresh should finish after preparation"); +} diff --git a/codex-rs/codex-mcp/src/catalog.rs b/codex-rs/codex-mcp/src/catalog.rs new file mode 100644 index 0000000000000000000000000000000000000000..486b7a9480be6d76057298957c27affbfff3678c --- /dev/null +++ b/codex-rs/codex-mcp/src/catalog.rs @@ -0,0 +1,652 @@ +use std::cmp::Reverse; +use std::collections::BTreeMap; +use std::collections::BTreeSet; +use std::collections::HashMap; + +use codex_config::McpServerConfig; +use codex_config::McpServerDisabledReason; +use codex_config::McpServerIdpOAuthConfig; +use codex_config::RequirementSource; +use codex_protocol::mcp_policy::EnvironmentMcpPolicy; +use codex_utils_path_uri::PathUri; + +use crate::CODEX_APPS_MCP_SERVER_NAME; +use crate::McpProtocolMode; + +/// Plugin identity retained with an MCP registration for tool attribution. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct McpPluginAttribution { + plugin_id: String, + display_name: String, + agent_plugin: bool, + host_root: Option, +} + +impl McpPluginAttribution { + pub fn new(plugin_id: String, display_name: String) -> Self { + Self { + plugin_id, + display_name, + agent_plugin: false, + host_root: None, + } + } + + pub fn agent_plugin(plugin_id: String, display_name: String) -> Self { + Self { + plugin_id, + display_name, + agent_plugin: true, + host_root: None, + } + } + + /// Records the exact host-discovered plugin root. + pub fn with_host_root(mut self, host_root: PathUri) -> Self { + self.host_root = Some(host_root); + self + } + + pub fn plugin_id(&self) -> &str { + &self.plugin_id + } + + pub fn display_name(&self) -> &str { + &self.display_name + } + + pub fn is_agent_plugin(&self) -> bool { + self.agent_plugin + } + + /// Returns the host-discovered root captured with this server registration. + pub fn host_root(&self) -> Option<&PathUri> { + self.host_root.as_ref() + } +} + +/// The component that declared an MCP server registration. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum McpServerSource { + /// A plugin discovered through the process-wide legacy plugin manager. + Plugin(McpPluginAttribution), + /// A plugin explicitly selected for this thread through a capability root. + SelectedPlugin(McpPluginAttribution), + Config, + Compatibility { + id: String, + }, + Extension { + id: String, + host_owned_apps: bool, + }, +} + +impl McpServerSource { + pub fn is_agent_plugin(&self) -> bool { + match self { + Self::Plugin(attribution) | Self::SelectedPlugin(attribution) => { + attribution.is_agent_plugin() + } + Self::Config | Self::Compatibility { .. } | Self::Extension { .. } => false, + } + } + + pub(crate) fn is_host_owned_apps(&self, name: &str, config: &McpServerConfig) -> bool { + name == CODEX_APPS_MCP_SERVER_NAME + && config.is_local_environment() + && matches!( + self, + Self::Compatibility { .. } + | Self::Extension { + host_owned_apps: true, + .. + } + ) + } + + fn disabled_registration_is_name_veto(&self) -> bool { + // A selected package's policy applies to its registration, not to a higher runtime source + // that happens to use the same logical server name. + !matches!(self, Self::SelectedPlugin(_)) + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)] +enum RegistrationPrecedence { + Plugin(Reverse), + SelectedPlugin(Reverse), + Config, + Compatibility, + Extension(usize), +} + +impl RegistrationPrecedence { + fn tier(self) -> u8 { + match self { + Self::Plugin(_) => 0, + Self::SelectedPlugin(_) => 1, + Self::Config => 2, + Self::Compatibility => 3, + Self::Extension(_) => 4, + } + } +} + +/// One named MCP server declaration before source resolution. +#[derive(Clone, Debug, PartialEq)] +pub struct McpServerRegistration { + name: String, + source: McpServerSource, + config: McpServerConfig, + protocol_mode: Option, + precedence: RegistrationPrecedence, +} + +impl McpServerRegistration { + pub fn from_config(name: String, config: McpServerConfig) -> Self { + Self::new( + name, + McpServerSource::Config, + config, + RegistrationPrecedence::Config, + ) + } + + pub fn from_plugin( + name: String, + attribution: McpPluginAttribution, + plugin_order: usize, + config: McpServerConfig, + ) -> Self { + Self::new( + name, + McpServerSource::Plugin(attribution), + config, + RegistrationPrecedence::Plugin(Reverse(plugin_order)), + ) + } + + /// Registers a thread-selected plugin above discovered plugins and below config. + pub fn from_selected_plugin( + name: String, + attribution: McpPluginAttribution, + selection_order: usize, + config: McpServerConfig, + ) -> Self { + Self::new( + name, + McpServerSource::SelectedPlugin(attribution), + config, + RegistrationPrecedence::SelectedPlugin(Reverse(selection_order)), + ) + } + + pub fn from_compatibility( + name: String, + id: impl Into, + config: McpServerConfig, + ) -> Self { + Self::new( + name, + McpServerSource::Compatibility { id: id.into() }, + config, + RegistrationPrecedence::Compatibility, + ) + } + + pub fn from_extension( + name: String, + id: impl Into, + contribution_order: usize, + config: McpServerConfig, + ) -> Self { + Self::new( + name, + McpServerSource::Extension { + id: id.into(), + host_owned_apps: false, + }, + config, + RegistrationPrecedence::Extension(contribution_order), + ) + } + + /// Overrides the protocol for this registration if it wins HTTP server resolution. + pub fn with_protocol_mode(mut self, protocol_mode: McpProtocolMode) -> Self { + self.protocol_mode = Some(protocol_mode); + self + } + + /// Registers the controller-owned Apps server contributed by a host extension. + pub fn from_hosted_apps( + id: impl Into, + contribution_order: usize, + config: McpServerConfig, + ) -> Self { + let host_owned_apps = config.is_local_environment(); + Self::new( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + McpServerSource::Extension { + id: id.into(), + host_owned_apps, + }, + config, + RegistrationPrecedence::Extension(contribution_order), + ) + } + + fn new( + name: String, + source: McpServerSource, + config: McpServerConfig, + precedence: RegistrationPrecedence, + ) -> Self { + Self { + name, + source, + config, + protocol_mode: None, + precedence, + } + } +} + +/// The authority available for MCP servers running in one environment. +#[derive(Clone, Copy, Debug)] +pub enum McpEnvironmentAuthority<'a> { + /// The selected environment adds no restrictions to the controller policy. + Unrestricted, + /// The owner supplied the final restrictions for this environment. + Restricted(&'a EnvironmentMcpPolicy), + /// An explicitly selected plugin can use its executor without attaching that executor. + SelectedPluginsOnly, + /// The attachment is pending or failed, so its owner policy is not available. + Unavailable, +} + +/// One side of an MCP server conflict, including whether it registers or +/// removes the server. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum McpServerConflictAction { + Register(McpServerSource), + Remove(McpServerSource), +} + +/// A same-tier name collision and the final outcome after all precedence is applied. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct McpServerConflict { + pub name: String, + pub outcome: McpServerConflictAction, + pub contenders: Vec, +} + +#[derive(Clone, Debug)] +enum CatalogAction { + Register(Box), + Remove { + name: String, + source: McpServerSource, + precedence: RegistrationPrecedence, + }, +} + +impl CatalogAction { + fn name(&self) -> &str { + match self { + Self::Register(registration) => ®istration.name, + Self::Remove { name, .. } => name, + } + } + + fn precedence(&self) -> RegistrationPrecedence { + match self { + Self::Register(registration) => registration.precedence, + Self::Remove { precedence, .. } => *precedence, + } + } + + fn conflict_action(&self) -> McpServerConflictAction { + match self { + Self::Register(registration) => { + McpServerConflictAction::Register(registration.source.clone()) + } + Self::Remove { source, .. } => McpServerConflictAction::Remove(source.clone()), + } + } +} + +/// Mutable inputs used to produce an immutable resolved catalog. +#[derive(Clone, Debug, Default)] +pub struct McpCatalogBuilder { + actions: Vec, + disabled_server_names: BTreeSet, + ema_idp: Option, +} + +impl McpCatalogBuilder { + /// Enables EMA with the IdP selected from trusted configuration. + /// Without this policy, finalization disables EMA registrations. + pub fn enable_ema(&mut self, idp: McpServerIdpOAuthConfig) { + self.ema_idp = Some(idp); + } + + pub fn register(&mut self, registration: McpServerRegistration) { + self.actions + .push(CatalogAction::Register(Box::new(registration))); + } + + /// Applies the legacy name-scoped disabled veto after source resolution. + pub fn disable(&mut self, name: String) { + self.disabled_server_names.insert(name); + } + + pub fn remove_compatibility(&mut self, name: String, id: impl Into) { + self.actions.push(CatalogAction::Remove { + name, + source: McpServerSource::Compatibility { id: id.into() }, + precedence: RegistrationPrecedence::Compatibility, + }); + } + + pub fn remove_extension( + &mut self, + name: String, + id: impl Into, + contribution_order: usize, + ) { + self.actions.push(CatalogAction::Remove { + name, + source: McpServerSource::Extension { + id: id.into(), + host_owned_apps: false, + }, + precedence: RegistrationPrecedence::Extension(contribution_order), + }); + } + + /// Applies environment authority before resolving immutable server registrations. + pub fn build_with_environment_authority<'a>( + mut self, + mut authority_for_environment: impl FnMut(&str) -> McpEnvironmentAuthority<'a>, + ) -> ResolvedMcpCatalog { + for action in &mut self.actions { + let CatalogAction::Register(registration) = action else { + continue; + }; + // Controller-owned Apps and existing managed denials are not attachment-owned. + if !registration.config.enabled + || registration + .source + .is_host_owned_apps(®istration.name, ®istration.config) + { + continue; + } + + let allowed = match authority_for_environment(®istration.config.environment_id) { + McpEnvironmentAuthority::Unrestricted => true, + McpEnvironmentAuthority::SelectedPluginsOnly => { + matches!(®istration.source, McpServerSource::SelectedPlugin(_)) + } + McpEnvironmentAuthority::Unavailable => false, + McpEnvironmentAuthority::Restricted(policy) => match ®istration.source { + McpServerSource::Config + | McpServerSource::Compatibility { .. } + | McpServerSource::Extension { .. } => { + policy.servers.as_ref().is_none_or(|requirements| { + requirements + .get(®istration.name) + .is_some_and(|requirement| { + registration.config.matches_requirement(requirement) + }) + }) + } + McpServerSource::Plugin(attribution) + | McpServerSource::SelectedPlugin(attribution) => { + // Empty server policy denies every plugin; otherwise use package policy. + !policy.servers.as_ref().is_some_and(BTreeMap::is_empty) + && policy + .plugins + .as_ref() + .filter(|plugins| { + plugins.values().any(|plugin| plugin.mcp_servers.is_some()) + }) + .is_none_or(|plugins| { + plugins + .get(attribution.plugin_id()) + .and_then(|plugin| plugin.mcp_servers.as_ref()) + .and_then(|requirements| { + requirements.get(®istration.name) + }) + .is_some_and(|requirement| { + registration.config.matches_requirement(requirement) + }) + }) + } + }, + }; + + if !allowed { + registration.config.enabled = false; + registration.config.disabled_reason = Some(McpServerDisabledReason::Requirements { + source: RequirementSource::Unknown, + }); + } + } + self.build() + } + + pub fn build(mut self) -> ResolvedMcpCatalog { + // Keep source actions unbound so later catalog revisions resolve afresh. + for action in &mut self.actions { + if let CatalogAction::Register(registration) = action + && let Some(oauth) = &mut registration.config.oauth + { + oauth.ema_registration = None; + } + } + // Stable sorting makes action order the tie-breaker when precedence is equal. + self.actions.sort_by_key(CatalogAction::precedence); + + let mut winners = BTreeMap::::new(); + let mut actions_by_name_and_tier = BTreeMap::<(String, u8), Vec<&CatalogAction>>::new(); + for action in &self.actions { + winners.insert(action.name().to_string(), action.clone()); + actions_by_name_and_tier + .entry((action.name().to_string(), action.precedence().tier())) + .or_default() + .push(action); + } + + let mut conflicts = Vec::new(); + for ((name, _), actions) in actions_by_name_and_tier { + if actions.len() < 2 { + continue; + } + let Some(outcome) = winners.get(&name).map(CatalogAction::conflict_action) else { + continue; + }; + conflicts.push(McpServerConflict { + name, + outcome, + contenders: actions + .into_iter() + .map(CatalogAction::conflict_action) + .collect(), + }); + } + + let mut disabled_server_names = self.disabled_server_names; + let ema_idp = self.ema_idp; + let servers = winners + .into_iter() + .filter_map(|(name, action)| match action { + CatalogAction::Register(registration) => { + let mut registration = *registration; + let persist_disabled_name = + registration.source.disabled_registration_is_name_veto() + && !matches!( + registration.config.disabled_reason, + Some(McpServerDisabledReason::EmaRegistration) + ); + if !registration.config.enabled || disabled_server_names.contains(&name) { + registration.config.enabled = false; + if persist_disabled_name { + // Preserve legacy disabled winners across later runtime overlays. + disabled_server_names.insert(name.clone()); + } + } + if matches!( + registration.config.auth, + codex_config::McpServerAuth::EmaAuth + ) { + let allowed = ema_idp.as_ref().is_some_and(|idp| { + registration.config.resolve_ema_registration(idp).is_ok() + }); + // EMA denial must not become a persistent name veto. + registration.config.enabled &= allowed; + } + Some(( + name, + ResolvedMcpServer { + source: registration.source, + config: registration.config, + protocol_mode: registration.protocol_mode, + }, + )) + } + CatalogAction::Remove { .. } => None, + }) + .collect(); + + ResolvedMcpCatalog { + actions: self.actions, + disabled_server_names, + ema_idp, + servers, + conflicts, + } + } +} + +/// A single winning MCP registration. +#[derive(Clone, Debug, PartialEq)] +pub struct ResolvedMcpServer { + source: McpServerSource, + config: McpServerConfig, + protocol_mode: Option, +} + +impl ResolvedMcpServer { + pub fn source(&self) -> &McpServerSource { + &self.source + } + + pub fn config(&self) -> &McpServerConfig { + &self.config + } + + pub fn protocol_mode(&self) -> Option { + self.protocol_mode + } +} + +/// Immutable result of MCP registration resolution. +#[derive(Clone, Debug, Default)] +pub struct ResolvedMcpCatalog { + actions: Vec, + disabled_server_names: BTreeSet, + ema_idp: Option, + servers: BTreeMap, + conflicts: Vec, +} + +impl ResolvedMcpCatalog { + pub fn builder() -> McpCatalogBuilder { + McpCatalogBuilder::default() + } + + pub fn to_builder(&self) -> McpCatalogBuilder { + McpCatalogBuilder { + actions: self.actions.clone(), + disabled_server_names: self.disabled_server_names.clone(), + ema_idp: self.ema_idp.clone(), + } + } + + pub fn server(&self, name: &str) -> Option<&ResolvedMcpServer> { + self.servers.get(name) + } + + pub fn configured_servers(&self) -> HashMap { + self.servers + .iter() + .map(|(name, server)| (name.clone(), server.config.clone())) + .collect() + } + + /// Returns whether both catalogs have the same winning servers, sources, and EMA policy. + pub fn has_same_servers(&self, other: &Self) -> bool { + self.servers == other.servers && self.ema_idp == other.ema_idp + } + + /// Replaces the resolved server set while preserving known sources and EMA policy. + /// + /// Names not present in the existing catalog are treated as config-owned. + pub fn with_materialized_servers(&self, servers: HashMap) -> Self { + let mut builder = McpCatalogBuilder { + ema_idp: self.ema_idp.clone(), + ..Default::default() + }; + for (name, config) in servers { + let previous = self.server(&name); + let source = previous + .map(|server| server.source.clone()) + .unwrap_or(McpServerSource::Config); + let precedence = match &source { + McpServerSource::Plugin(_) => RegistrationPrecedence::Plugin(Reverse(0)), + McpServerSource::SelectedPlugin(_) => { + RegistrationPrecedence::SelectedPlugin(Reverse(0)) + } + McpServerSource::Config => RegistrationPrecedence::Config, + McpServerSource::Compatibility { .. } => RegistrationPrecedence::Compatibility, + McpServerSource::Extension { .. } => RegistrationPrecedence::Extension(0), + }; + let mut registration = McpServerRegistration::new(name, source, config, precedence); + registration.protocol_mode = previous.and_then(ResolvedMcpServer::protocol_mode); + builder.register(registration); + } + builder.build() + } + + /// Returns package attribution for each winning plugin-owned server. + pub fn plugin_attributions_by_server_name(&self) -> HashMap { + self.servers + .iter() + .filter_map(|(name, server)| match server.source() { + McpServerSource::Plugin(attribution) + | McpServerSource::SelectedPlugin(attribution) => { + Some((name.clone(), attribution.clone())) + } + McpServerSource::Config + | McpServerSource::Compatibility { .. } + | McpServerSource::Extension { .. } => None, + }) + .collect() + } + + /// Returns the names of winning servers supplied by thread-selected plugins. + pub(crate) fn selected_plugin_server_names(&self) -> impl Iterator { + self.servers.iter().filter_map(|(name, server)| { + matches!(server.source(), McpServerSource::SelectedPlugin(_)).then_some(name.as_str()) + }) + } + + pub fn conflicts(&self) -> &[McpServerConflict] { + &self.conflicts + } +} + +#[cfg(test)] +#[path = "catalog_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/catalog_tests.rs b/codex-rs/codex-mcp/src/catalog_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4c5d1ed0cdd1245a6d195a45ee5335906f0521e9 --- /dev/null +++ b/codex-rs/codex-mcp/src/catalog_tests.rs @@ -0,0 +1,722 @@ +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::time::Duration; + +use codex_config::AppToolApproval; +use codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID; +use codex_config::McpServerAuth; +use codex_config::McpServerConfig; +use codex_config::McpServerDisabledReason; +use codex_config::McpServerIdpOAuthConfig; +use codex_config::McpServerToolConfig; +use codex_config::McpServerTransportConfig; +use codex_config::types::PluginMcpServerEmaAuthConfig; +use codex_protocol::mcp_policy::EnvironmentMcpPolicy; +use codex_protocol::mcp_policy::PluginMcpRequirements; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; + +use crate::CODEX_APPS_MCP_SERVER_NAME; +use crate::McpProtocolMode; + +use super::McpEnvironmentAuthority; +use super::McpPluginAttribution; +use super::McpServerConflict; +use super::McpServerConflictAction; +use super::McpServerRegistration; +use super::McpServerSource; +use super::ResolvedMcpCatalog; +use super::ResolvedMcpServer; + +fn server(url: &str) -> McpServerConfig { + McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: url.to_string(), + bearer_token_env_var: None, + http_headers: None, + env_http_headers: None, + http_headers_helper: None, + }, + environment_id: DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: true, + supports_parallel_tool_calls: true, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: Some(Duration::from_secs(7)), + tool_timeout_sec: Some(Duration::from_secs(11)), + default_tools_approval_mode: Some(AppToolApproval::Prompt), + enabled_tools: Some(vec!["read".to_string()]), + disabled_tools: Some(vec!["write".to_string()]), + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::from([( + "read".to_string(), + McpServerToolConfig { + approval_mode: Some(AppToolApproval::Approve), + ..Default::default() + }, + )]), + } +} + +fn plugin(plugin_id: &str) -> McpPluginAttribution { + McpPluginAttribution::new(plugin_id.to_string(), plugin_id.to_string()) +} + +fn plugin_source(plugin_id: &str) -> McpServerSource { + McpServerSource::Plugin(plugin(plugin_id)) +} + +fn selected_plugin_source(plugin_id: &str) -> McpServerSource { + McpServerSource::SelectedPlugin(plugin(plugin_id)) +} + +fn compatibility_source(id: &str) -> McpServerSource { + McpServerSource::Compatibility { id: id.to_string() } +} + +fn extension_source(id: &str) -> McpServerSource { + McpServerSource::Extension { + id: id.to_string(), + host_owned_apps: false, + } +} + +fn register(source: McpServerSource) -> McpServerConflictAction { + McpServerConflictAction::Register(source) +} + +fn remove(source: McpServerSource) -> McpServerConflictAction { + McpServerConflictAction::Remove(source) +} + +#[test] +fn plugin_host_root_is_retained_in_catalog_identity() { + let original_root = PathUri::parse("file:///plugins/original").expect("valid plugin root URI"); + let replacement_root = + PathUri::parse("file:///plugins/replacement").expect("valid plugin root URI"); + let catalog_for_root = |root| { + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_plugin( + "docs".to_string(), + plugin("plugin@test").with_host_root(root), + /*plugin_order*/ 0, + server("https://plugin.example/mcp"), + )); + builder.build() + }; + let original = catalog_for_root(original_root.clone()); + let replacement = catalog_for_root(replacement_root); + + let Some(McpServerSource::Plugin(attribution)) = + original.server("docs").map(ResolvedMcpServer::source) + else { + panic!("expected host-discovered plugin registration"); + }; + assert_eq!(attribution.host_root(), Some(&original_root)); + assert!(!original.has_same_servers(&replacement)); +} + +#[test] +fn ema_policy_survives_rebuilds_and_rebinds_materialized_servers() { + let idp = McpServerIdpOAuthConfig { + issuer: "https://idp.example".to_string(), + client_id: "enterprise-client".to_string(), + }; + let mut builder = ResolvedMcpCatalog::builder(); + builder.enable_ema(idp.clone()); + let enabled_empty = builder.build(); + assert!(!enabled_empty.has_same_servers(&ResolvedMcpCatalog::default())); + + let mut builder = enabled_empty.to_builder(); + let mut config = server("https://resource.example/mcp"); + config.auth = McpServerAuth::EmaAuth; + builder.register(McpServerRegistration::from_config( + "enterprise".to_string(), + config, + )); + let original = builder.build(); + assert!(original.has_same_servers(&original.to_builder().build())); + + let mut servers = original.configured_servers(); + let config = servers.get_mut("enterprise").expect("EMA server"); + let McpServerTransportConfig::StreamableHttp { url, .. } = &mut config.transport else { + panic!("expected HTTP server"); + }; + *url = "https://resource.example/revised".to_string(); + let materialized = original.with_materialized_servers(servers); + let config = materialized + .server("enterprise") + .expect("materialized EMA server") + .config(); + let registration = config + .ema_registration() + .expect("finalized EMA registration"); + assert_eq!( + ( + config.enabled, + registration.idp(), + registration.server_url() + ), + (true, &idp, "https://resource.example/revised") + ); + + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_config( + "enterprise".to_string(), + config.clone(), + )); + let denied = builder.build(); + let config = denied + .server("enterprise") + .expect("denied EMA server") + .config(); + assert_eq!((config.enabled, config.ema_registration()), (false, None)); +} + +#[test] +fn rejected_plugin_ema_registration_does_not_veto_hosted_apps() { + let idp = McpServerIdpOAuthConfig { + issuer: "https://idp.example".to_string(), + client_id: "enterprise-client".to_string(), + }; + for (url, resource) in [ + ("https://other.example/mcp", "https://resource.example"), + ("https://resource.example/mcp", " "), + ] { + for initially_enabled in [true, false] { + let mut rejected = server(url); + rejected.enabled = initially_enabled; + let policy = PluginMcpServerEmaAuthConfig { + url: "https://resource.example/mcp".to_string(), + client_id: "resource-client".to_string(), + authorization_server_issuer: "https://as.example".to_string(), + scopes: Vec::new(), + resource: resource.to_string(), + }; + policy.apply(&mut rejected); + assert!(rejected.resolve_ema_registration(&idp).is_err()); + assert_eq!( + ( + rejected.enabled, + rejected.auth.clone(), + rejected.ema_registration() + ), + (false, McpServerAuth::EmaAuth, None), + ); + assert_eq!( + rejected.disabled_reason, + initially_enabled.then_some(McpServerDisabledReason::EmaRegistration), + ); + + let mut builder = ResolvedMcpCatalog::builder(); + builder.enable_ema(idp.clone()); + builder.register(McpServerRegistration::from_plugin( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + plugin("plugin@test"), + /*plugin_order*/ 0, + rejected.clone(), + )); + let catalog = builder.build(); + assert_eq!( + catalog.server(CODEX_APPS_MCP_SERVER_NAME).unwrap().config(), + &rejected, + ); + + let materialized = catalog.with_materialized_servers(catalog.configured_servers()); + for catalog in [catalog, materialized] { + let mut builder = catalog.to_builder(); + let mut expected = server("https://chatgpt.com/mcp"); + builder.register(McpServerRegistration::from_hosted_apps( + "apps", + /*contribution_order*/ 0, + expected.clone(), + )); + expected.enabled = initially_enabled; + assert_eq!( + builder.build().server(CODEX_APPS_MCP_SERVER_NAME), + Some(&ResolvedMcpServer { + source: McpServerSource::Extension { + id: "apps".to_string(), + host_owned_apps: true, + }, + config: expected, + protocol_mode: None, + }), + ); + } + } + } +} + +#[test] +fn source_precedence_preserves_the_winning_registration() { + let extension = server("https://extension.example/mcp"); + let mut plugin_server = server("https://plugin.example/mcp"); + plugin_server.enabled = false; + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_extension( + "docs".to_string(), + "hosted", + /*contribution_order*/ 0, + extension.clone(), + )); + builder.register(McpServerRegistration::from_plugin( + "docs".to_string(), + plugin("plugin@test"), + /*plugin_order*/ 0, + plugin_server, + )); + builder.register(McpServerRegistration::from_plugin( + "docs".to_string(), + plugin("other-plugin@test"), + /*plugin_order*/ 1, + server("https://other-plugin.example/mcp"), + )); + builder.register(McpServerRegistration::from_compatibility( + "docs".to_string(), + "legacy", + server("https://compatibility.example/mcp"), + )); + builder.register(McpServerRegistration::from_config( + "docs".to_string(), + server("https://config.example/mcp"), + )); + + let catalog = builder.build(); + let resolved = catalog.server("docs").expect("resolved server"); + + assert_eq!( + resolved.source(), + &McpServerSource::Extension { + id: "hosted".to_string(), + host_owned_apps: false, + } + ); + assert_eq!(resolved.config(), &extension); + assert!(catalog.plugin_attributions_by_server_name().is_empty()); + assert_eq!( + catalog.conflicts(), + &[McpServerConflict { + name: "docs".to_string(), + outcome: register(extension_source("hosted")), + contenders: vec![ + register(plugin_source("other-plugin@test")), + register(plugin_source("plugin@test")), + ], + }] + ); +} + +#[test] +fn disabled_veto_only_disables_the_winning_registration() { + let extension = server("https://extension.example/mcp"); + let mut expected = extension.clone(); + expected.enabled = false; + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_extension( + "docs".to_string(), + "hosted", + /*contribution_order*/ 0, + extension, + )); + builder.disable("docs".to_string()); + + let actual = builder + .build() + .server("docs") + .expect("resolved server") + .config() + .clone(); + + assert_eq!(actual, expected); +} + +#[test] +fn disabled_winner_remains_a_veto_when_the_catalog_is_extended() { + let mut disabled = server("https://config.example/mcp"); + disabled.enabled = false; + let mut expected = server("https://extension.example/mcp"); + expected.enabled = false; + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_config( + "docs".to_string(), + disabled, + )); + let mut builder = builder.build().to_builder(); + builder.register(McpServerRegistration::from_extension( + "docs".to_string(), + "hosted", + /*contribution_order*/ 0, + server("https://extension.example/mcp"), + )); + + let resolved = builder.build(); + + assert_eq!( + resolved.server("docs"), + Some(&super::ResolvedMcpServer { + source: extension_source("hosted"), + config: expected, + protocol_mode: None, + }) + ); +} + +#[test] +fn disabled_discovered_plugin_remains_a_veto_for_runtime_overlays() { + let mut disabled = server("https://plugin.example/mcp"); + disabled.enabled = false; + let mut expected = server("https://extension.example/mcp"); + expected.enabled = false; + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_plugin( + "docs".to_string(), + plugin("plugin@test"), + /*plugin_order*/ 0, + disabled, + )); + let mut builder = builder.build().to_builder(); + builder.register(McpServerRegistration::from_extension( + "docs".to_string(), + "hosted", + /*contribution_order*/ 0, + server("https://extension.example/mcp"), + )); + + let resolved = builder.build(); + + assert_eq!( + resolved.server("docs"), + Some(&super::ResolvedMcpServer { + source: extension_source("hosted"), + config: expected, + protocol_mode: None, + }) + ); +} + +#[test] +fn earlier_plugin_wins_with_an_explicit_conflict() { + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_plugin( + "docs".to_string(), + plugin("alpha@test"), + /*plugin_order*/ 0, + server("https://alpha.example/mcp"), + )); + builder.register(McpServerRegistration::from_plugin( + "docs".to_string(), + plugin("beta@test"), + /*plugin_order*/ 1, + server("https://beta.example/mcp"), + )); + + let catalog = builder.build(); + + assert_eq!( + catalog.plugin_attributions_by_server_name(), + HashMap::from([("docs".to_string(), plugin("alpha@test"))]) + ); + assert_eq!( + catalog.conflicts(), + &[McpServerConflict { + name: "docs".to_string(), + outcome: register(plugin_source("alpha@test")), + contenders: vec![ + register(plugin_source("beta@test")), + register(plugin_source("alpha@test")), + ], + }] + ); +} + +#[test] +fn selected_plugins_override_discovered_plugins_but_not_config() { + let selected = server("https://selected-alpha.example/mcp"); + let mut discovered = server("https://local.example/mcp"); + discovered.enabled = false; + discovered.default_tools_approval_mode = Some(AppToolApproval::Auto); + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_plugin( + "docs".to_string(), + plugin("local@test"), + /*plugin_order*/ 0, + discovered, + )); + builder.register(McpServerRegistration::from_selected_plugin( + "docs".to_string(), + plugin("selected-beta"), + /*selection_order*/ 1, + server("https://selected-beta.example/mcp"), + )); + builder.register(McpServerRegistration::from_selected_plugin( + "docs".to_string(), + plugin("selected-alpha"), + /*selection_order*/ 0, + selected.clone(), + )); + + let catalog = builder.build(); + + assert_eq!( + catalog.server("docs"), + Some(&super::ResolvedMcpServer { + source: selected_plugin_source("selected-alpha"), + config: selected, + protocol_mode: None, + }) + ); + assert_eq!( + catalog.plugin_attributions_by_server_name(), + HashMap::from([("docs".to_string(), plugin("selected-alpha"))]) + ); + assert_eq!( + catalog.conflicts(), + &[McpServerConflict { + name: "docs".to_string(), + outcome: register(selected_plugin_source("selected-alpha")), + contenders: vec![ + register(selected_plugin_source("selected-beta")), + register(selected_plugin_source("selected-alpha")), + ], + }] + ); + + let refreshed = server("https://refreshed.example/mcp"); + let catalog = + catalog.with_materialized_servers(HashMap::from([("docs".to_string(), refreshed.clone())])); + assert_eq!( + catalog.server("docs"), + Some(&super::ResolvedMcpServer { + source: selected_plugin_source("selected-alpha"), + config: refreshed, + protocol_mode: None, + }) + ); + + let mut builder = catalog.to_builder(); + let configured = server("https://config.example/mcp"); + builder.register(McpServerRegistration::from_config( + "docs".to_string(), + configured.clone(), + )); + let catalog = builder.build(); + + assert_eq!( + catalog.server("docs"), + Some(&super::ResolvedMcpServer { + source: McpServerSource::Config, + config: configured, + protocol_mode: None, + }) + ); +} + +#[test] +fn disabled_selected_plugin_does_not_veto_runtime_overlays() { + let mut disabled = server("https://selected.example/mcp"); + disabled.enabled = false; + let extension = server("https://extension.example/mcp"); + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_selected_plugin( + "docs".to_string(), + plugin("selected"), + /*selection_order*/ 0, + disabled, + )); + let mut builder = builder.build().to_builder(); + builder.register(McpServerRegistration::from_extension( + "docs".to_string(), + "hosted", + /*contribution_order*/ 0, + extension.clone(), + )); + + let resolved = builder.build(); + + assert_eq!( + resolved.server("docs"), + Some(&super::ResolvedMcpServer { + source: extension_source("hosted"), + config: extension, + protocol_mode: None, + }) + ); +} + +#[test] +fn equal_precedence_uses_insertion_order_not_source_identity() { + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_compatibility( + "docs".to_string(), + "z-first", + server("https://first.example/mcp"), + )); + builder.register(McpServerRegistration::from_compatibility( + "docs".to_string(), + "a-second", + server("https://second.example/mcp"), + )); + + let catalog = builder.build(); + + assert_eq!( + catalog.server("docs"), + Some(&super::ResolvedMcpServer { + source: compatibility_source("a-second"), + config: server("https://second.example/mcp"), + protocol_mode: None, + }) + ); + let mut builder = catalog.to_builder(); + builder.remove_compatibility("docs".to_string(), "remove-last"); + + let catalog = builder.build(); + + assert_eq!(catalog.server("docs"), None); + assert_eq!( + catalog.conflicts(), + &[McpServerConflict { + name: "docs".to_string(), + outcome: remove(compatibility_source("remove-last")), + contenders: vec![ + register(compatibility_source("z-first")), + register(compatibility_source("a-second")), + remove(compatibility_source("remove-last")), + ], + }] + ); +} + +#[test] +fn extension_protocol_mode_follows_the_winner_through_materialization() { + let config = server("https://apps.example/mcp"); + let mut builder = ResolvedMcpCatalog::builder(); + builder.register( + McpServerRegistration::from_extension( + "apps".to_string(), + "loser", + /*contribution_order*/ 0, + config.clone(), + ) + .with_protocol_mode(McpProtocolMode::V20260728), + ); + builder.register(McpServerRegistration::from_extension( + "apps".to_string(), + "winner", + /*contribution_order*/ 1, + config.clone(), + )); + let without_override = builder.build(); + assert_eq!( + without_override + .server("apps") + .and_then(ResolvedMcpServer::protocol_mode), + None + ); + + let mut builder = ResolvedMcpCatalog::builder(); + builder.register( + McpServerRegistration::from_extension( + "apps".to_string(), + "winner", + /*contribution_order*/ 1, + config, + ) + .with_protocol_mode(McpProtocolMode::Legacy), + ); + let with_override = builder.build(); + assert!(!without_override.has_same_servers(&with_override)); + + let refreshed = server("https://refreshed.example/mcp"); + let materialized = with_override + .with_materialized_servers(HashMap::from([("apps".to_string(), refreshed.clone())])); + assert_eq!( + materialized.server("apps"), + Some(&ResolvedMcpServer { + source: extension_source("winner"), + config: refreshed, + protocol_mode: Some(McpProtocolMode::Legacy), + }) + ); +} + +#[test] +fn environment_policy_exempts_only_explicitly_host_owned_apps() { + let policy = EnvironmentMcpPolicy { + servers: Some(BTreeMap::new()), + plugins: None, + }; + for (registration, expected) in [ + ( + McpServerRegistration::from_extension( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + "apps", + /*contribution_order*/ 0, + server("https://apps.example/mcp"), + ), + false, + ), + ( + McpServerRegistration::from_hosted_apps( + "apps", + /*contribution_order*/ 0, + server("https://apps.example/mcp"), + ), + true, + ), + ] { + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(registration); + let catalog = builder + .build_with_environment_authority(|_| McpEnvironmentAuthority::Restricted(&policy)); + assert_eq!( + catalog + .server(CODEX_APPS_MCP_SERVER_NAME) + .expect("Apps registration") + .config() + .enabled, + expected + ); + } +} + +#[test] +fn environment_policy_preserves_selected_plugin_and_empty_server_allowlist_semantics() { + let mut builder = ResolvedMcpCatalog::builder(); + builder.register(McpServerRegistration::from_selected_plugin( + "selected".to_string(), + plugin("selected-plugin"), + /*selection_order*/ 0, + server("https://plugin.example/mcp"), + )); + let metadata_only_policy = EnvironmentMcpPolicy { + servers: None, + plugins: Some(BTreeMap::from([( + "metadata-only-plugin".to_string(), + PluginMcpRequirements { mcp_servers: None }, + )])), + }; + let deny_all_policy = EnvironmentMcpPolicy { + servers: Some(BTreeMap::new()), + plugins: None, + }; + + for (policy, expected) in [(&metadata_only_policy, true), (&deny_all_policy, false)] { + let resolved = builder + .clone() + .build_with_environment_authority(|_| McpEnvironmentAuthority::Restricted(policy)); + assert_eq!( + resolved + .server("selected") + .expect("selected plugin") + .config() + .enabled, + expected + ); + } +} diff --git a/codex-rs/codex-mcp/src/client_capabilities.rs b/codex-rs/codex-mcp/src/client_capabilities.rs new file mode 100644 index 0000000000000000000000000000000000000000..582891dbdaff0353057af80973397facee27de0d --- /dev/null +++ b/codex-rs/codex-mcp/src/client_capabilities.rs @@ -0,0 +1,73 @@ +use std::collections::HashMap; + +use codex_protocol::mcp::ClientMcpExtensions; +use codex_protocol::mcp::MCP_APP_UI_EXTENSION_ID; +use codex_protocol::mcp::OPENAI_ELICITATION_EXTENSION_ID; +use codex_protocol::mcp::OPENAI_FORM_EXTENSION_ID; +use codex_protocol::mcp::OPENAI_STANDARD_FORM_INPUT_EXTENSION_ID; +use serde_json::Map; +use serde_json::Value; + +/// Restricts verification to the host-owned plugin service, even when another server uses its name. +pub(crate) fn server_mcp_extensions( + extensions: &ClientMcpExtensions, + server_name: &str, + server: Option<&crate::catalog::ResolvedMcpServer>, +) -> ClientMcpExtensions { + let plugin_service = server.is_some_and(|server| { + server + .source() + .is_host_owned_apps(server_name, server.config()) + }); + ClientMcpExtensions::new(extensions.iter().map(|(id, settings)| { + let mut settings = settings.clone(); + if !plugin_service + && id == OPENAI_ELICITATION_EXTENSION_ID + && let Some(settings) = settings.as_object_mut() + { + settings.remove("userVerification"); + } + (id.to_string(), settings) + })) +} + +/// Selects the MCP extensions Codex supports from those declared by the app-server host. +/// +/// App-server clients may declare unrelated extensions. Codex retains only the +/// trusted extension namespaces it knows how to project downstream. The +/// legacy form capability is normalized into the same extension map. +pub fn client_mcp_extensions( + extensions: Option<&HashMap>, + legacy_openai_form_elicitation: bool, +) -> ClientMcpExtensions { + let mut selected = extensions + .into_iter() + .flat_map(HashMap::iter) + .filter(|(id, _)| { + matches!( + id.as_str(), + OPENAI_FORM_EXTENSION_ID + | OPENAI_ELICITATION_EXTENSION_ID + | OPENAI_STANDARD_FORM_INPUT_EXTENSION_ID + | MCP_APP_UI_EXTENSION_ID + ) + }) + .map(|(id, value)| (id.clone(), value.clone())) + .collect::>(); + if let Some(settings) = selected + .get_mut(OPENAI_ELICITATION_EXTENSION_ID) + .and_then(Value::as_object_mut) + { + settings.retain(|key, _| key == "form"); + } + if legacy_openai_form_elicitation { + selected + .entry(OPENAI_FORM_EXTENSION_ID.to_string()) + .or_insert_with(|| Value::Object(Map::new())); + } + ClientMcpExtensions::new(selected) +} + +#[cfg(test)] +#[path = "client_capabilities_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/client_capabilities_tests.rs b/codex-rs/codex-mcp/src/client_capabilities_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..be0733499217bf38c15182136d1b0a76d1f14f6f --- /dev/null +++ b/codex-rs/codex-mcp/src/client_capabilities_tests.rs @@ -0,0 +1,110 @@ +use std::collections::HashMap; + +use pretty_assertions::assert_eq; +use serde_json::json; + +use super::*; + +#[test] +fn selects_only_supported_mcp_extensions() { + let app_ui = json!({ + "mimeTypes": [ + "text/html;profile=mcp-app", + "text/x-dil;profile=mcp-app", + ], + "futureField": {"preserved": true}, + }); + let form = json!({"futureField": {"preserved": true}}); + let extensions = HashMap::from([ + ( + OPENAI_ELICITATION_EXTENSION_ID.to_string(), + json!({"form": form, "userVerification": {}, "unsupported": {}}), + ), + (MCP_APP_UI_EXTENSION_ID.to_string(), app_ui.clone()), + (OPENAI_FORM_EXTENSION_ID.to_string(), json!({})), + ( + OPENAI_STANDARD_FORM_INPUT_EXTENSION_ID.to_string(), + json!({}), + ), + ("example/other".to_string(), json!({"enabled": true})), + ]); + + assert_eq!( + client_mcp_extensions( + Some(&extensions), + /*legacy_openai_form_elicitation*/ false, + ), + ClientMcpExtensions::new(HashMap::from([ + ( + OPENAI_ELICITATION_EXTENSION_ID.to_string(), + json!({"form": form}), + ), + (MCP_APP_UI_EXTENSION_ID.to_string(), app_ui), + (OPENAI_FORM_EXTENSION_ID.to_string(), json!({})), + ( + OPENAI_STANDARD_FORM_INPUT_EXTENSION_ID.to_string(), + json!({}), + ), + ])) + ); +} + +#[test] +fn normalizes_legacy_form_capability_into_extensions() { + assert_eq!( + client_mcp_extensions( + /*extensions*/ None, /*legacy_openai_form_elicitation*/ true, + ), + ClientMcpExtensions::new(HashMap::from([( + OPENAI_FORM_EXTENSION_ID.to_string(), + json!({}), + )])) + ); +} + +#[test] +fn user_verification_is_projected_only_to_the_host_owned_plugin_service() { + use crate::catalog::McpServerRegistration; + use crate::catalog::ResolvedMcpCatalog; + use crate::mcp::CODEX_APPS_MCP_SERVER_NAME; + use crate::mcp::codex_apps_mcp_server_config; + + let config = codex_apps_mcp_server_config( + "https://example.com", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ); + let extensions = ClientMcpExtensions::new([( + OPENAI_ELICITATION_EXTENSION_ID.to_string(), + json!({"form": {}, "userVerification": {}}), + )]); + for (name, registration, settings) in [ + ( + CODEX_APPS_MCP_SERVER_NAME, + McpServerRegistration::from_hosted_apps( + "host", + /*contribution_order*/ 0, + config.clone(), + ), + json!({"form": {}, "userVerification": {}}), + ), + ( + CODEX_APPS_MCP_SERVER_NAME, + McpServerRegistration::from_config(CODEX_APPS_MCP_SERVER_NAME.into(), config.clone()), + json!({"form": {}}), + ), + ( + "attached", + McpServerRegistration::from_config("attached".into(), config), + json!({"form": {}}), + ), + ] { + let mut catalog = ResolvedMcpCatalog::builder(); + catalog.register(registration); + let catalog = catalog.build(); + assert_eq!( + server_mcp_extensions(&extensions, name, catalog.server(name)), + ClientMcpExtensions::new([(OPENAI_ELICITATION_EXTENSION_ID.to_string(), settings)]), + ); + } +} diff --git a/codex-rs/codex-mcp/src/client_tool_catalog.rs b/codex-rs/codex-mcp/src/client_tool_catalog.rs new file mode 100644 index 0000000000000000000000000000000000000000..0501f1c782d6906656a81ee6be3047c7b5313347 --- /dev/null +++ b/codex-rs/codex-mcp/src/client_tool_catalog.rs @@ -0,0 +1,241 @@ +//! Client-local catalog revisions and synchronization for MCP calls. +//! +//! Apps clients read the provider's current catalog without retaining old arrays. +//! Only active readers, calls, and frozen model bindings pin earlier snapshots. +//! Explicit refreshes publish after calls using this client's current revision finish. +//! Equivalent shared publications preserve calls, but explicit refreshes invalidate them. + +use std::collections::HashSet; +use std::future::Future; +use std::sync::Arc; + +use anyhow::Result; +use codex_connectors::ConnectorRuntimeSnapshot; +use tokio::sync::Mutex; +use tokio::sync::RwLock; +use tokio::sync::RwLockReadGuard; +use tokio::sync::watch; + +use crate::tools::ToolInfo; + +type ToolCatalogUpdates = watch::Receiver>>>; + +/// The exact Apps catalog returned by an awaited refresh of one published runtime. +pub struct CodexAppsToolSnapshot { + /// Raw installed tools, including tools hidden or disabled for the model. + pub tools: Vec, + /// Raw MCP tool names allowed by the same runtime's generic MCP policy. + /// App-specific policy is applied by the caller. + pub model_visible_tool_names: HashSet, +} + +pub(crate) struct ClientToolCatalog { + current: RwLock, + /// Serialize fetches without blocking calls against the current catalog. + refresh_lock: Mutex<()>, +} + +struct CatalogState { + revision: u64, + /// The catalog revision installed by this client's last explicit refresh. + last_refresh_revision: u64, + source: CatalogSource, +} + +/// Ordinary MCP clients own their tools; Apps clients observe the shared provider. +enum CatalogSource { + Local(Arc<[ToolInfo]>), + Live { + updates: ToolCatalogUpdates, + tools_version: u64, + }, +} + +pub(crate) struct ToolCatalogSnapshot { + pub(crate) revision: u64, + pub(crate) tools: Arc<[ToolInfo]>, +} + +impl ClientToolCatalog { + /// Live clients read the provider; callers publish startup tools before subscribing. + pub(crate) fn new( + tools: impl Into>, + updates: Option, + ) -> Self { + let source = match updates { + Some(mut updates) => { + let tools_version = updates + .borrow_and_update() + .as_ref() + .map_or(0, |snapshot| snapshot.tools_version()); + CatalogSource::Live { + updates, + tools_version, + } + } + None => CatalogSource::Local(tools.into()), + }; + Self { + current: RwLock::new(CatalogState { + revision: 0, + last_refresh_revision: 0, + source, + }), + refresh_lock: Mutex::new(()), + } + } + + pub(crate) async fn read(&self, read: impl FnOnce(ToolCatalogSnapshot) -> R) -> R { + let (_current, snapshot) = self.read_current().await; + read(snapshot) + } + + /// Captures tools and their client-local revision together, without retaining them in the client. + async fn read_current(&self) -> (RwLockReadGuard<'_, CatalogState>, ToolCatalogSnapshot) { + loop { + { + let current = self.current.read().await; + let tools = match ¤t.source { + CatalogSource::Local(tools) => Some(Arc::clone(tools)), + CatalogSource::Live { updates, .. } => { + // Hold the watch borrow while checking its version so the tools and + // revision come from the same publication. + let snapshot = updates.borrow(); + (!updates.has_changed().unwrap_or(false)).then(|| { + snapshot + .as_ref() + .map(|snapshot| snapshot.shared_tools()) + .unwrap_or_default() + }) + } + }; + if let Some(tools) = tools { + let snapshot = ToolCatalogSnapshot { + revision: current.revision, + tools, + }; + return (current, snapshot); + } + } + let mut current = self.current.write().await; + let changed = match &mut current.source { + CatalogSource::Live { + updates, + tools_version, + } => { + let version = updates + .borrow_and_update() + .as_ref() + .map_or(0, |snapshot| snapshot.tools_version()); + let changed = *tools_version != version; + *tools_version = version; + changed + } + CatalogSource::Local(_) => false, + }; + if changed { + current.revision += 1; + } + } + } + + /// Serialize fetching and publication, leaving the current catalog usable during the fetch. + /// The publication callback runs alongside the exact-client update under the write lock. + #[expect( + clippy::await_holding_invalid_type, + reason = "refreshes must remain serialized through fetching and catalog publication" + )] + pub(crate) async fn refresh(&self, fetch: F, publish: P) -> Result + where + F: FnOnce() -> Fut, + Fut: Future, C)>>, + P: FnOnce(&[ToolInfo], C) -> R, + { + let _refresh = self.refresh_lock.lock().await; + let (tools, context) = fetch().await?; + let mut current = self.current.write().await; + let result = publish(&tools, context); + match &mut current.source { + CatalogSource::Local(current_tools) => *current_tools = tools.into(), + CatalogSource::Live { + updates, + tools_version, + } => { + *tools_version = updates + .borrow_and_update() + .as_ref() + .map_or(0, |snapshot| snapshot.tools_version()); + } + } + current.revision += 1; + current.last_refresh_revision = current.revision; + Ok(result) + } + + /// Reject stale calls before preparation and hold catalog authority until execution finishes. + #[expect( + clippy::await_holding_invalid_type, + reason = "catalog publication must wait for call preparation and execution" + )] + pub(crate) async fn run_with_snapshot( + &self, + expected: &ToolCatalogSnapshot, + run: F, + ) -> Option + where + F: FnOnce() -> Fut, + Fut: Future, + { + let (current, snapshot) = self.read_current().await; + if current.last_refresh_revision > expected.revision + || (snapshot.revision != expected.revision + && !catalogs_match(&snapshot.tools, &expected.tools)) + { + return None; + } + let result = run().await; + drop(current); + Some(result) + } +} + +/// Compares complete definitions independently of tool-list order. Stable sorting keeps +/// conflicting duplicate identities in their original order, since deduplication can pick the +/// first definition. This slow path runs only after the client's revision has changed. +fn catalogs_match(left: &[ToolInfo], right: &[ToolInfo]) -> bool { + if left == right { + return true; + } + if left.len() != right.len() { + return false; + } + let mut left = left.iter().collect::>(); + let mut right = right.iter().collect::>(); + for tools in [&mut left, &mut right] { + tools.sort_by_key(|tool| { + ( + tool.server_name.as_str(), + tool.tool.name.as_ref(), + tool.connector_id.as_deref(), + tool.callable_namespace.as_str(), + tool.callable_name.as_str(), + ) + }); + } + left == right +} + +/// A binding cache key includes client identity, since new clients start at zero. +#[derive(Clone)] +pub(crate) struct ClientToolCatalogRevision { + pub(crate) catalog: Arc, + pub(crate) revision: u64, +} + +impl PartialEq for ClientToolCatalogRevision { + fn eq(&self, other: &Self) -> bool { + Arc::ptr_eq(&self.catalog, &other.catalog) && self.revision == other.revision + } +} + +impl Eq for ClientToolCatalogRevision {} diff --git a/codex-rs/codex-mcp/src/codex_apps.rs b/codex-rs/codex-mcp/src/codex_apps.rs new file mode 100644 index 0000000000000000000000000000000000000000..809b8ec93b8477b6173afa308ca0d7878caf53e4 --- /dev/null +++ b/codex-rs/codex-mcp/src/codex_apps.rs @@ -0,0 +1,70 @@ +//! Codex Apps support for the host-owned apps MCP server. +//! +//! This module owns the normalization that turns ChatGPT-hosted app +//! connector/tool metadata into model-visible MCP callable names. + +use codex_utils_plugins::mcp_connector::sanitize_name; + +mod file_params; + +pub use file_params::declared_openai_file_input_param_names; +pub(crate) use file_params::prepare_openai_file_params_for_model; + +pub(crate) fn normalize_codex_apps_tool_title(connector_name: Option<&str>, value: &str) -> String { + let Some(connector_name) = connector_name + .map(str::trim) + .filter(|name| !name.is_empty()) + else { + return value.to_string(); + }; + + let prefix = format!("{connector_name}_"); + if let Some(stripped) = value.strip_prefix(&prefix) + && !stripped.is_empty() + { + return stripped.to_string(); + } + + value.to_string() +} + +pub(crate) fn normalize_codex_apps_callable_name( + tool_name: &str, + connector_id: Option<&str>, + connector_name: Option<&str>, +) -> String { + let tool_name = sanitize_name(tool_name); + + if let Some(connector_name) = connector_name + .map(str::trim) + .map(sanitize_name) + .filter(|name| !name.is_empty()) + && let Some(stripped) = tool_name.strip_prefix(&connector_name) + && !stripped.is_empty() + { + return stripped.to_string(); + } + + if let Some(connector_id) = connector_id + .map(str::trim) + .map(sanitize_name) + .filter(|name| !name.is_empty()) + && let Some(stripped) = tool_name.strip_prefix(&connector_id) + && !stripped.is_empty() + { + return stripped.to_string(); + } + + tool_name +} + +pub(crate) fn normalize_codex_apps_callable_namespace( + server_name: &str, + connector_name: Option<&str>, +) -> String { + if let Some(connector_name) = connector_name { + format!("{}__{}", server_name, sanitize_name(connector_name)) + } else { + server_name.to_string() + } +} diff --git a/codex-rs/codex-mcp/src/codex_apps/file_params.rs b/codex-rs/codex-mcp/src/codex_apps/file_params.rs new file mode 100644 index 0000000000000000000000000000000000000000..98e2fa5d618a5625f0218a653d783378a3929276 --- /dev/null +++ b/codex-rs/codex-mcp/src/codex_apps/file_params.rs @@ -0,0 +1,219 @@ +//! Apps SDK `openai/fileParams` metadata and schema shaping. +//! +//! For each declared file argument, this module derives the provided-file fields +//! accepted by its input schema and records them on `ToolInfo` for execution-time +//! argument rewriting. It also presents file arguments to the model as local paths. +//! +//! See . + +use std::collections::HashMap; +use std::collections::HashSet; +use std::sync::Arc; + +use rmcp::model::Tool; +use serde_json::Map; +use serde_json::Value as JsonValue; + +use crate::tools::ToolInfo; + +const META_OPENAI_FILE_PARAMS: &str = "openai/fileParams"; + +#[derive(Default)] +struct OpenAiFileSchemaInfo { + accepts_mime_type: bool, + accepts_file_name: bool, +} + +pub fn declared_openai_file_input_param_names( + meta: Option<&Map>, +) -> Vec { + let Some(meta) = meta else { + return Vec::new(); + }; + + meta.get(META_OPENAI_FILE_PARAMS) + .and_then(JsonValue::as_array) + .into_iter() + .flatten() + .filter_map(JsonValue::as_str) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .collect() +} + +/// Derives execution-time file capabilities from the raw schema, then masks +/// declared file arguments as local paths for the model. +pub(crate) fn prepare_openai_file_params_for_model(tool_info: &mut ToolInfo) { + let file_params = declared_openai_file_input_param_names(tool_info.tool.meta.as_deref()); + tool_info.openai_file_input_optional_fields = + supported_openai_file_input_optional_fields(&tool_info.tool, &file_params); + + if file_params.is_empty() { + return; + } + + let mut tool = tool_info.tool.clone(); + let mut input_schema = JsonValue::Object(tool.input_schema.as_ref().clone()); + rewrite_input_schema_for_local_file_paths(&mut input_schema, &file_params); + if let JsonValue::Object(input_schema) = input_schema { + tool.input_schema = Arc::new(input_schema); + } + tool_info.tool = tool; +} + +fn supported_openai_file_input_optional_fields( + tool: &Tool, + file_params: &[String], +) -> HashMap> { + let properties = tool + .input_schema + .get("properties") + .and_then(JsonValue::as_object); + + file_params + .iter() + .map(|field_name| { + let optional_fields = properties + .and_then(|properties| properties.get(field_name)) + .map(|schema| { + let schema_info = openai_file_schema_info(schema, tool.input_schema.as_ref()); + let mut optional_fields = Vec::new(); + if schema_info.accepts_mime_type { + optional_fields.push("mime_type".to_string()); + } + if schema_info.accepts_file_name { + optional_fields.push("file_name".to_string()); + } + optional_fields + }) + .unwrap_or_default(); + (field_name.clone(), optional_fields) + }) + .collect() +} + +fn openai_file_schema_info( + schema: &JsonValue, + root_schema: &Map, +) -> OpenAiFileSchemaInfo { + let mut info = OpenAiFileSchemaInfo::default(); + let mut pending = vec![schema]; + let mut visited_refs = HashSet::new(); + + while let Some(schema) = pending.pop() { + let Some(schema) = schema.as_object() else { + continue; + }; + + if let Some(schema_ref) = schema.get("$ref").and_then(JsonValue::as_str) + && visited_refs.insert(schema_ref) + && let Some(referenced_schema) = resolve_local_schema_ref(root_schema, schema_ref) + { + pending.push(referenced_schema); + } + + for keyword in ["anyOf", "oneOf", "allOf"] { + if let Some(variants) = schema.get(keyword).and_then(JsonValue::as_array) { + pending.extend(variants); + } + } + + if schema.get("type").and_then(JsonValue::as_str) == Some("array") + || schema.contains_key("items") + { + if let Some(items) = schema.get("items") { + pending.push(items); + } + continue; + } + + let properties = schema.get("properties").and_then(JsonValue::as_object); + let is_object_schema = schema.get("type").and_then(JsonValue::as_str) == Some("object") + || properties.is_some() + || schema.contains_key("additionalProperties"); + if !is_object_schema { + continue; + } + let accepts_additional_properties = !matches!( + schema.get("additionalProperties"), + Some(JsonValue::Bool(false) | JsonValue::Object(_)) + ); + info.accepts_mime_type |= accepts_additional_properties + || properties.is_some_and(|properties| properties.contains_key("mime_type")); + info.accepts_file_name |= accepts_additional_properties + || properties.is_some_and(|properties| properties.contains_key("file_name")); + } + + info +} + +fn resolve_local_schema_ref<'a>( + root_schema: &'a Map, + schema_ref: &str, +) -> Option<&'a JsonValue> { + let pointer = schema_ref.strip_prefix("#/")?; + let mut segments = pointer.split('/'); + let first_segment = segments.next()?.replace("~1", "/").replace("~0", "~"); + let mut referenced_schema = root_schema.get(&first_segment)?; + + for segment in segments { + let segment = segment.replace("~1", "/").replace("~0", "~"); + referenced_schema = match referenced_schema { + JsonValue::Object(object) => object.get(&segment)?, + JsonValue::Array(array) => array.get(segment.parse::().ok()?)?, + _ => return None, + }; + } + + Some(referenced_schema) +} + +fn rewrite_input_schema_for_local_file_paths(input_schema: &mut JsonValue, file_params: &[String]) { + let Some(properties) = input_schema + .as_object_mut() + .and_then(|schema| schema.get_mut("properties")) + .and_then(JsonValue::as_object_mut) + else { + return; + }; + + for field_name in file_params { + let Some(property_schema) = properties.get_mut(field_name) else { + continue; + }; + rewrite_input_property_schema_as_local_file_path(property_schema); + } +} + +fn rewrite_input_property_schema_as_local_file_path(schema: &mut JsonValue) { + let Some(object) = schema.as_object_mut() else { + return; + }; + + let mut description = object + .get("description") + .and_then(JsonValue::as_str) + .map(str::to_string) + .unwrap_or_default(); + let guidance = "This parameter expects an absolute local file path. If you want to upload a file, provide the absolute path to that file here."; + if description.is_empty() { + description = guidance.to_string(); + } else if !description.contains(guidance) { + description = format!("{description} {guidance}"); + } + + let is_array = object.get("type").and_then(JsonValue::as_str) == Some("array") + || object.get("items").is_some(); + object.clear(); + object.insert("description".to_string(), JsonValue::String(description)); + if is_array { + object.insert("type".to_string(), JsonValue::String("array".to_string())); + object.insert("items".to_string(), serde_json::json!({ "type": "string" })); + } else { + object.insert("type".to_string(), JsonValue::String("string".to_string())); + } +} + +#[cfg(test)] +#[path = "file_params_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/codex_apps/file_params_tests.rs b/codex-rs/codex-mcp/src/codex_apps/file_params_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..3d7f326ef691e869386d4226c1e4e2ffb77f108f --- /dev/null +++ b/codex-rs/codex-mcp/src/codex_apps/file_params_tests.rs @@ -0,0 +1,284 @@ +use std::collections::HashMap; +use std::sync::Arc; + +use pretty_assertions::assert_eq; +use rmcp::model::JsonObject; +use rmcp::model::MetaObject; +use rmcp::model::Tool; + +use super::*; +use crate::tools::ToolInfo; + +fn tool_info(tool: Tool) -> ToolInfo { + ToolInfo { + server_name: "codex_apps".to_string(), + supports_parallel_tool_calls: false, + server_origin: None, + callable_name: tool.name.to_string(), + callable_namespace: "codex_apps".to_string(), + namespace_description: None, + tool, + openai_file_input_optional_fields: HashMap::new(), + connector_id: None, + connector_name: None, + plugin_display_names: Vec::new(), + } +} + +fn test_tool(name: &str) -> Tool { + Tool::new( + name.to_string(), + format!("Test tool: {name}"), + Arc::new(JsonObject::default()), + ) +} + +#[test] +fn declared_openai_file_fields_treat_names_literally() { + let meta = serde_json::json!({ + "openai/fileParams": ["file", "input_file", "attachments"] + }); + let meta = meta.as_object().expect("meta object"); + + assert_eq!( + declared_openai_file_input_param_names(Some(meta)), + vec![ + "file".to_string(), + "input_file".to_string(), + "attachments".to_string(), + ] + ); +} + +#[test] +fn prepare_openai_file_params_for_model_masks_file_params() { + let mut tool = test_tool("upload"); + tool.input_schema = Arc::new( + serde_json::json!({ + "type": "object", + "properties": { + "file": { + "type": "object", + "description": "Original file payload." + }, + "files": { + "type": "array", + "items": {"type": "object"} + } + } + }) + .as_object() + .expect("object") + .clone(), + ); + tool.meta = Some(MetaObject( + serde_json::json!({ + "openai/fileParams": ["file", "files"] + }) + .as_object() + .expect("object") + .clone(), + )); + let mut tool_info = tool_info(tool); + + prepare_openai_file_params_for_model(&mut tool_info); + + assert_eq!( + *tool_info.tool.input_schema, + serde_json::json!({ + "type": "object", + "properties": { + "file": { + "type": "string", + "description": "Original file payload. This parameter expects an absolute local file path. If you want to upload a file, provide the absolute path to that file here." + }, + "files": { + "type": "array", + "items": {"type": "string"}, + "description": "This parameter expects an absolute local file path. If you want to upload a file, provide the absolute path to that file here." + } + } + }) + .as_object() + .expect("object") + .clone() + ); +} + +#[test] +fn prepare_openai_file_params_for_model_derives_supported_optional_fields() { + let mut tool = Tool::new( + "upload".to_string(), + "Upload files".to_string(), + Arc::new( + serde_json::json!({ + "type": "object", + "$defs": { + "Rich/File": { + "type": "object", + "properties": { + "download_url": {"type": "string"}, + "file_id": {"type": "string"}, + "file_name": {"type": "string"} + }, + "additionalProperties": false + } + }, + "properties": { + "photoshop_image": { + "type": "object", + "properties": { + "download_url": {"type": "string"}, + "file_id": {"type": "string"} + }, + "additionalProperties": false + }, + "drive_import": { + "type": "object", + "properties": { + "download_url": {"type": "string"}, + "file_id": {"type": "string"}, + "mime_type": {"type": "string"}, + "file_name": {"type": "string"} + }, + "additionalProperties": false + }, + "attachments": { + "anyOf": [ + { + "type": "array", + "items": { + "oneOf": [ + { + "allOf": [ + { + "type": "object", + "properties": { + "download_url": {"type": "string"}, + "file_id": {"type": "string"} + } + }, + { + "type": "object", + "properties": { + "mime_type": {"type": "string"} + } + } + ] + }, + {"type": "null"} + ] + } + }, + {"type": "null"} + ] + }, + "referenced_file": { + "$ref": "#/$defs/Rich~1File" + }, + "custom_file": { + "type": "object", + "properties": { + "download_url": {"type": "string"}, + "file_id": {"type": "string"}, + "mime_type": {"type": "string"}, + "uri": {"type": "string"} + }, + "additionalProperties": false + }, + "open_file": { + "type": "object", + "properties": { + "download_url": {"type": "string"}, + "file_id": {"type": "string"} + } + }, + "explicitly_open_file": { + "type": "object", + "properties": { + "download_url": {"type": "string"}, + "file_id": {"type": "string"} + }, + "additionalProperties": true + }, + "items_only_files": { + "items": { + "type": "object", + "properties": { + "download_url": {"type": "string"}, + "file_id": {"type": "string"}, + "file_name": {"type": "string"} + }, + "additionalProperties": false + } + } + } + }) + .as_object() + .expect("object") + .clone(), + ), + ); + tool.meta = Some(MetaObject( + serde_json::json!({ + "openai/fileParams": [ + "photoshop_image", + "drive_import", + "attachments", + "referenced_file", + "custom_file", + "open_file", + "explicitly_open_file", + "items_only_files", + "missing_file" + ] + }) + .as_object() + .expect("object") + .clone(), + )); + let mut tool_info = tool_info(tool); + + prepare_openai_file_params_for_model(&mut tool_info); + + assert_eq!( + tool_info.openai_file_input_optional_fields, + HashMap::from([ + ("photoshop_image".to_string(), Vec::new()), + ( + "drive_import".to_string(), + vec!["mime_type".to_string(), "file_name".to_string()] + ), + ( + "attachments".to_string(), + vec!["mime_type".to_string(), "file_name".to_string()] + ), + ("referenced_file".to_string(), vec!["file_name".to_string()]), + ("custom_file".to_string(), vec!["mime_type".to_string()]), + ( + "open_file".to_string(), + vec!["mime_type".to_string(), "file_name".to_string()] + ), + ( + "explicitly_open_file".to_string(), + vec!["mime_type".to_string(), "file_name".to_string()] + ), + ( + "items_only_files".to_string(), + vec!["file_name".to_string()] + ), + ("missing_file".to_string(), Vec::new()), + ]) + ); +} + +#[test] +fn prepare_openai_file_params_for_model_leaves_tools_without_file_params_unchanged() { + let original_tool = test_tool("upload"); + let mut tool_info = tool_info(original_tool.clone()); + + prepare_openai_file_params_for_model(&mut tool_info); + + assert_eq!(tool_info.tool, original_tool); + assert!(tool_info.openai_file_input_optional_fields.is_empty()); +} diff --git a/codex-rs/codex-mcp/src/connection_manager.rs b/codex-rs/codex-mcp/src/connection_manager.rs new file mode 100644 index 0000000000000000000000000000000000000000..6022f1c1193a6490ea27ddf3541f6bfd9e88f828 --- /dev/null +++ b/codex-rs/codex-mcp/src/connection_manager.rs @@ -0,0 +1,1048 @@ +//! Aggregates MCP server connections for Codex. +//! +//! [`McpConnectionSet`] is the private connection set behind +//! [`crate::McpRuntime`] and [`crate::McpBinding`]. It coordinates startup status +//! events, keeps server metadata, and aggregates tools and resources across +//! running RMCP clients. + +#[path = "connection_manager/required.rs"] +mod required; +#[path = "connection_manager/resources.rs"] +mod resources; +#[path = "connection_manager/startup.rs"] +mod startup; +#[path = "connection_manager/status.rs"] +mod status; +#[path = "connection_manager/tool_catalog.rs"] +mod tool_catalog; + +use startup::chatgpt_auth_provider_for_server; +use startup::emit_update; +use startup::mcp_init_error_display; +use startup::mcp_startup_failure_reason; +use startup::should_share_codex_apps_tools_cache; +pub(crate) use tool_catalog::BindingCatalogRevision; +pub use tool_catalog::tool_is_model_visible; + +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::OnceLock; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use crate::binding::call_tool_result_from_rmcp; +use crate::catalog::McpServerSource; +use crate::elicitation::ElicitationRequestManager; +use crate::elicitation::ElicitationRequestRouter; +use crate::event_stream::EventStreamConnectionSettings; +use crate::mcp::CODEX_APPS_MCP_SERVER_NAME; +use crate::mcp::ToolPluginContext; +use crate::pagination::MAX_CODEX_APPS_TOOL_CATALOG_ITEMS; +use crate::pagination::MAX_MCP_CATALOG_ITEMS; +use crate::rmcp_client::AsyncManagedClient; +use crate::rmcp_client::DEFAULT_STARTUP_TIMEOUT; +use crate::rmcp_client::DEFAULT_TOOL_TIMEOUT; +use crate::rmcp_client::ManagedClient; +use crate::rmcp_client::StartupOutcomeError; +use crate::rmcp_client::prepare_codex_apps_tools_for_model; +use crate::rmcp_client::prepare_regular_mcp_tools_for_model; +use crate::runtime::McpPublicationGate; +use crate::runtime::McpRuntimeInput; +use crate::runtime::McpStartupPolicy; +use crate::server::McpServerConnectionIdentity; +use crate::server::McpServerMetadata; +use crate::tool_catalog_cache::McpToolCatalogCacheContext; +use crate::tools::ToolFilter; +use crate::tools::ToolInfo; +use crate::tools::filter_tools; +use crate::trusted_access::ENTITLEMENT_CONTEXT_KEY; +use crate::trusted_access::TrustedAccessContext; +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use codex_config::McpServerTransportConfig; +use codex_diagnostics::Gauge; +use codex_diagnostics::GaugeGuard; +use codex_protocol::mcp::CallToolResult; +use codex_protocol::mcp::McpServerInfo; +use codex_protocol::protocol::Event; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::McpStartupCompleteEvent; +use codex_protocol::protocol::McpStartupFailure; +use codex_protocol::protocol::McpStartupFailureReason; +use codex_protocol::protocol::McpStartupStatus; +use codex_protocol::protocol::McpStartupUpdateEvent; +use codex_rmcp_client::determine_streamable_http_auth_status_from_credentials; +use tokio::sync::watch; +use tokio::task::JoinSet; +use tracing::warn; + +static LIVE_CONNECTIONS: Gauge = Gauge::new("mcp.connections.live"); + +pub(crate) struct McpServerConnection { + identity: Option, + client: AsyncManagedClient, + // Startup-only budget; changing it must not replace a ready connection. + startup_timeout: Duration, + startup_trigger: Option>, + _diagnostics_guard: GaugeGuard, +} + +impl McpServerConnection { + async fn reusable_client( + &self, + desired: &McpServerConnectionIdentity, + ) -> Option { + let current = self.identity.as_ref()?; + if !current.has_same_connection_config(desired) { + return None; + } + if !self.client.startup_complete.load(Ordering::Acquire) { + return None; + } + let client = self.client.client().await.ok()?; + if client.client.is_closed().await { + return None; + } + if current == desired { + if matches!(desired.oauth_credentials(), Ok(None)) + && tokio::time::timeout(Duration::ZERO, client.client.managed_oauth_credentials()) + .await + .is_ok_and(|credentials| matches!(credentials, Some(Some(_)))) + { + return None; + } + return Some(client); + } + let Ok(desired_credentials) = desired.oauth_credentials() else { + return Some(client); + }; + let reusable = match client.client.managed_oauth_credentials().await { + Some(live_credentials) => live_credentials.as_ref() == desired_credentials, + None => current + .oauth_credentials() + .is_ok_and(|startup_credentials| startup_credentials == desired_credentials), + }; + if reusable { Some(client) } else { None } + } + + pub(crate) async fn client(&self) -> Result { + if let Some(startup_trigger) = &self.startup_trigger { + startup_trigger.send_replace(true); + } + self.client.client().await + } + + async fn shutdown(&self) { + self.client.shutdown().await; + } + + fn cancel_startup(&self) { + if !self.startup_is_dormant() && !self.client.startup_complete.load(Ordering::Acquire) { + self.client.cancel_token.cancel(); + } + } + + fn startup_is_dormant(&self) -> bool { + self.startup_trigger + .as_ref() + .is_some_and(|startup_trigger| !*startup_trigger.borrow()) + } +} + +impl Drop for McpServerConnection { + fn drop(&mut self) { + self.client.cancel_token.cancel(); + } +} + +#[derive(Clone)] +struct McpServerView { + connection: Arc, + protocol_mode: crate::McpProtocolMode, + metadata: McpServerMetadata, + tool_filter: ToolFilter, + tool_timeout: Option, + catalog_item_limit: usize, +} + +impl McpServerView { + async fn listed_tools( + &self, + tool_plugin_context: &ToolPluginContext, + ) -> Result, StartupOutcomeError> { + let tools = self.connection.client.listed_tools().await?; + let tools = filter_tools(tools, &self.tool_filter); + Ok(if self.connection.client.is_codex_apps_mcp_server { + prepare_codex_apps_tools_for_model(tools, tool_plugin_context) + } else { + prepare_regular_mcp_tools_for_model(tools, tool_plugin_context) + }) + } +} + +/// A published view over a set of running MCP server connections. +pub(crate) struct McpConnectionSet { + servers: HashMap, + pub(crate) event_stream_connection: Option>, + disabled_servers: Vec, + required_servers: Vec, + optional_startup_deadline: OnceLock, + tool_plugin_context: Arc, + prefix_mcp_tool_names: bool, + non_prefixed_mcp_tool_servers: Vec, + elicitation_requests: ElicitationRequestManager, + pub(crate) trusted_access: Option, +} + +impl McpConnectionSet { + /// Creates an MCP connection manager. Threadless callers can pass no `tx_event`; startup + /// notifications are then skipped and interactive elicitations are declined. + pub async fn new( + previous: Option<&Self>, + publication_gate: McpPublicationGate, + input: McpRuntimeInput, + elicitation_router: ElicitationRequestRouter, + ) -> Self { + let trusted_access = TrustedAccessContext::from_runtime(&input); + let McpRuntimeInput { + startup_policy, + config, + plugins_available: _, + ready_selected_capability_roots: _, + mcp_servers, + submit_id, + tx_event, + startup_cancellation_token, + runtime_context, + codex_apps_tools_cache, + tool_catalog_cache, + codex_apps_tools_cache_key, + client_mcp_extensions, + auth, + auth_manager, + elicitation_reviewer, + elicitation_lifecycle, + } = input; + let store_mode = config.mcp_oauth_credentials_store_mode; + let keyring_backend_kind = config.auth_keyring_backend_kind; + let oauth_refresh_mode = config.oauth_refresh_mode; + let codex_home = config.codex_home.clone(); + let prefix_mcp_tool_names = config.prefix_mcp_tool_names; + let non_prefixed_mcp_tool_servers = config.non_prefixed_mcp_tool_servers.clone(); + let default_protocol_mode = config.protocol_mode; + let host_owned_apps_protocol_mode = config.host_owned_apps_protocol_mode; + let client_elicitation_capability = config.client_elicitation_capability.clone(); + let tool_plugin_context = crate::mcp::tool_plugin_context(&config); + let auth = auth.as_ref(); + let mut servers = HashMap::new(); + let mut event_stream_connection = None; + let disabled_servers = mcp_servers + .iter() + .filter(|(_, server)| !server.enabled()) + .map(|(name, _)| name.clone()) + .collect(); + let mut required_servers = mcp_servers + .iter() + .filter(|(_, server)| server.enabled() && server.required()) + .map(|(server_name, _)| server_name.clone()) + .collect::>(); + required_servers.sort(); + let mut reused_ready = Vec::new(); + let mut join_set = JoinSet::new(); + // Explicit reconnects have no previous set and must replace their clients eagerly. + let allow_deferred_startup = + startup_policy == McpStartupPolicy::LazyWhenCached && previous.is_some(); + let reusable_previous = previous.filter(|previous| { + !previous.servers.is_empty() + && previous.elicitation_requests.update( + Arc::clone(&config), + elicitation_reviewer.clone(), + elicitation_lifecycle.clone(), + ) + }); + let elicitation_requests = if let Some(previous) = reusable_previous { + previous.elicitation_requests.clone() + } else { + ElicitationRequestManager::new( + Arc::clone(&config), + elicitation_reviewer, + elicitation_lifecycle, + elicitation_router, + ) + }; + let tool_plugin_context = Arc::new(tool_plugin_context); + let startup_submit_id = submit_id; + let static_chatgpt_auth_provider = auth + .filter(|auth| auth.uses_codex_backend()) + .map(codex_model_provider::auth_provider_from_auth); + let codex_apps_auth_provider = auth_manager.as_ref().and_then(|auth_manager| { + auth.filter(|auth| auth.uses_codex_backend()).map(|auth| { + codex_model_provider::auth_provider_from_auth_manager( + Arc::clone(auth_manager), + auth, + ) + }) + }); + for (server_name, server) in mcp_servers + .into_iter() + .filter(|(_, server)| server.enabled()) + { + let registration = config.mcp_server_catalog.server(&server_name); + let client_mcp_extensions = crate::client_capabilities::server_mcp_extensions( + &client_mcp_extensions, + &server_name, + registration, + ); + let is_host_owned_codex_apps = registration.is_some_and(|server| { + server + .source() + .is_host_owned_apps(&server_name, server.config()) + }); + let host_plugin_root = registration.and_then(|server| match server.source() { + McpServerSource::Plugin(plugin) => plugin.host_root(), + McpServerSource::SelectedPlugin(_) + | McpServerSource::Config + | McpServerSource::Compatibility { .. } + | McpServerSource::Extension { .. } => None, + }); + let catalog_item_limit = if is_host_owned_codex_apps { + MAX_CODEX_APPS_TOOL_CATALOG_ITEMS + } else { + MAX_MCP_CATALOG_ITEMS + }; + let metadata = McpServerMetadata::from(&server); + let configured_config = server.config().clone(); + let protocol_mode = if matches!( + &configured_config.transport, + McpServerTransportConfig::StreamableHttp { .. } + ) { + registration + .and_then(crate::ResolvedMcpServer::protocol_mode) + .unwrap_or(if is_host_owned_codex_apps { + host_owned_apps_protocol_mode + } else { + default_protocol_mode + }) + } else { + default_protocol_mode + }; + let configured_tool_filter = ToolFilter::from_config(&configured_config); + let startup_timeout = configured_config + .startup_timeout_sec + .unwrap_or(DEFAULT_STARTUP_TIMEOUT); + let configured_tool_timeout = Some( + configured_config + .tool_timeout_sec + .unwrap_or(DEFAULT_TOOL_TIMEOUT), + ); + let resolved_environment = + runtime_context.resolve_server_environment(&server_name, &configured_config); + // For built-in Codex Apps, `CODEX_CONNECTORS_TOKEN` is a debug + // override: it supplies runtime auth but bypasses the shared tools + // cache. + let uses_env_bearer_token = match &configured_config.transport { + McpServerTransportConfig::StreamableHttp { + bearer_token_env_var, + .. + } => bearer_token_env_var.is_some(), + McpServerTransportConfig::Stdio { .. } => false, + }; + let shares_codex_apps_tools_cache = is_host_owned_codex_apps + && should_share_codex_apps_tools_cache(&server_name, uses_env_bearer_token); + let codex_apps_tools_cache_context = shares_codex_apps_tools_cache.then(|| { + // Only equivalent discovery inputs may share executable Apps tools. + let mut transport = configured_config.transport.clone(); + if let McpServerTransportConfig::StreamableHttp { + http_headers: Some(headers), + .. + } = &mut transport + { + // mcp_server_config_for_url in codex-rs/codex-mcp/src/mcp/mod.rs + // adds thread attribution that threadless discovery does not carry. + headers.retain(|name, _| !name.eq_ignore_ascii_case("originator")); + } + let mut scope = serde_json::json!([ + transport, + &configured_config.auth, + protocol_mode.preferred_protocol_version().as_str(), + catalog_item_limit, + ]); + scope.sort_all_objects(); + codex_apps_tools_cache + .context(codex_home.clone(), codex_apps_tools_cache_key.clone()) + .with_live_scope(scope.to_string()) + }); + // The reserved Codex Apps registration follows the shared + // AuthManager across refreshes. In the hosted-plugin path, this + // is the ChatGPT /ps/mcp connection. User-configured MCP + // registrations keep their existing configured auth path. + let chatgpt_auth_provider = if server_name == CODEX_APPS_MCP_SERVER_NAME { + codex_apps_auth_provider + .clone() + .or_else(|| static_chatgpt_auth_provider.clone()) + } else { + static_chatgpt_auth_provider.clone() + }; + // If Codex Apps has an env bearer token, that is its auth path. Do + // not also attach the ambient CodexAuth provider. + let runtime_auth_provider = + if server_name == CODEX_APPS_MCP_SERVER_NAME && uses_env_bearer_token { + None + } else { + chatgpt_auth_provider_for_server(&server, chatgpt_auth_provider) + }; + if is_host_owned_codex_apps { + event_stream_connection = Some(Arc::new(EventStreamConnectionSettings { + server: server.clone(), + store_mode, + keyring_backend_kind, + oauth_refresh_mode, + runtime_context: runtime_context.clone(), + resolved_environment: resolved_environment.clone(), + auth_provider: runtime_auth_provider.clone(), + auth_manager: auth_manager.clone(), + auth: auth.cloned(), + protocol_mode, + client_mcp_extensions: client_mcp_extensions.clone(), + })); + } + let connection_identity = McpServerConnectionIdentity::new( + &server_name, + &server, + host_plugin_root, + store_mode, + keyring_backend_kind, + oauth_refresh_mode, + &resolved_environment, + &runtime_context, + runtime_auth_provider.as_ref(), + auth, + shares_codex_apps_tools_cache + .then(|| (codex_home.clone(), codex_apps_tools_cache_key.clone())), + client_elicitation_capability.clone(), + client_mcp_extensions.clone(), + previous + .and_then(|previous| previous.servers.get(&server_name)) + .and_then(|view| view.connection.identity.as_ref()), + ); + let expected_protocol_mode = match &configured_config.transport { + McpServerTransportConfig::StreamableHttp { .. } => Some(protocol_mode), + McpServerTransportConfig::Stdio { .. } + if protocol_mode == crate::McpProtocolMode::Legacy => + { + Some(crate::McpProtocolMode::Legacy) + } + McpServerTransportConfig::Stdio { env, .. } => match env + .as_ref() + .and_then(|variables| variables.get("CODEX_MCP_PROTOCOL_VERSION")) + { + None => Some(crate::McpProtocolMode::Legacy), + Some(version) + if version == rmcp::model::ProtocolVersion::V_2026_07_28.as_str() => + { + Some(protocol_mode) + } + Some(_) => None, + }, + }; + if let Some(previous_view) = + reusable_previous.and_then(|previous| previous.servers.get(&server_name)) + { + let connection = Arc::clone(&previous_view.connection); + let reusable_pending_startup = connection.identity.as_ref() + == Some(&connection_identity) + && !connection.client.startup_complete.load(Ordering::Acquire) + && connection.startup_timeout == startup_timeout + && !connection.startup_is_dormant() + && !connection.client.cancel_token.is_cancelled() + && previous_view.catalog_item_limit == catalog_item_limit + && expected_protocol_mode.is_some() + && previous_view.protocol_mode == protocol_mode; + let unchanged_auth_failure = if connection.identity.as_ref() + == Some(&connection_identity) + && connection_identity.oauth_store_was_contended + && previous_view.protocol_mode == protocol_mode + && connection.client.startup_complete.load(Ordering::Acquire) + { + connection + .client() + .await + .err() + .filter(StartupOutcomeError::is_authentication_required) + } else { + None + }; + if reusable_pending_startup + || unchanged_auth_failure.is_some() + || connection + .reusable_client(&connection_identity) + .await + .is_some_and(|client| { + previous_view.catalog_item_limit == catalog_item_limit + && expected_protocol_mode.is_some_and(|expected| { + client.client.protocol_mode() == expected + }) + }) + { + let pending_client = + reusable_pending_startup.then(|| connection.client.clone()); + servers.insert( + server_name.clone(), + McpServerView { + connection, + protocol_mode, + metadata, + tool_filter: configured_tool_filter, + tool_timeout: configured_tool_timeout, + catalog_item_limit, + }, + ); + if let Some(error) = unchanged_auth_failure { + let reason = connection_identity + .oauth_credentials() + .ok() + .flatten() + .map(|_| McpStartupFailureReason::ReauthenticationRequired); + let status = McpStartupStatus::Failed { + error: mcp_init_error_display( + &server_name, + Some(&configured_config), + &error, + reason, + ), + reason, + }; + let tx_event = tx_event.clone(); + let submit_id = startup_submit_id.clone(); + let publication_gate = publication_gate.clone(); + join_set.spawn(async move { + if !publication_gate.wait().await { + return (server_name, Err(StartupOutcomeError::Cancelled)); + } + if let Some(tx_event) = tx_event.as_ref() { + for status in [McpStartupStatus::Starting, status] { + let _ = emit_update( + submit_id.as_str(), + tx_event, + McpStartupUpdateEvent { + server: server_name.clone(), + status, + }, + ) + .await; + } + } + (server_name, Err(error)) + }); + } else if let Some(client) = pending_client { + let publication_gate = publication_gate.clone(); + join_set.spawn(async move { + if !publication_gate.wait().await { + return (server_name, Err(StartupOutcomeError::Cancelled)); + } + (server_name, client.client().await) + }); + } else { + reused_ready.push(server_name); + } + continue; + } + } + let cancel_token = startup_cancellation_token.child_token(); + let tool_catalog_cache_context = if server_name == CODEX_APPS_MCP_SERVER_NAME { + None + } else if let Ok(environment) = resolved_environment.as_ref() { + tool_catalog_cache.context( + &server_name, + &configured_config, + &runtime_context, + environment.as_ref(), + (&client_elicitation_capability, &client_mcp_extensions), + Some(( + &connection_identity, + protocol_mode, + server.is_agent_plugin(), + )), + ) + } else { + None + }; + let has_runtime_auth = runtime_auth_provider.is_some(); + let async_managed_client = AsyncManagedClient::new( + server_name.clone(), + startup_submit_id.clone(), + server, + store_mode, + keyring_backend_kind, + oauth_refresh_mode, + cancel_token.clone(), + tx_event.clone(), + elicitation_requests.clone(), + codex_apps_tools_cache_context, + tool_catalog_cache_context, + runtime_context.clone(), + resolved_environment, + runtime_auth_provider, + client_elicitation_capability.clone(), + client_mcp_extensions.clone(), + auth_manager + .as_ref() + .filter(|_| { + matches!( + &configured_config.transport, + McpServerTransportConfig::Stdio { .. } + ) + }) + .map(|manager| manager.auth_change_state_receiver()), + protocol_mode, + catalog_item_limit, + ); + let defer_startup = allow_deferred_startup + && !tool_plugin_context.is_selected_plugin_mcp_server(&server_name) + && async_managed_client + .tool_catalog_cache_context + .as_ref() + .and_then(McpToolCatalogCacheContext::current_tools) + .is_some_and(|tools| { + tools.into_iter().any(|tool| { + configured_tool_filter.allows(&tool.tool.name) + && tool_is_model_visible(&tool) + }) + }); + let (startup_trigger, startup_receiver) = if defer_startup { + let (trigger, receiver) = watch::channel(false); + (Some(trigger), Some(receiver)) + } else { + (None, None) + }; + servers.insert( + server_name.clone(), + McpServerView { + connection: Arc::new(McpServerConnection { + identity: Some(connection_identity), + client: async_managed_client.clone(), + startup_timeout, + startup_trigger, + _diagnostics_guard: LIVE_CONNECTIONS.track(), + }), + protocol_mode, + metadata, + tool_filter: configured_tool_filter, + tool_timeout: configured_tool_timeout, + catalog_item_limit, + }, + ); + let tx_event = tx_event.clone(); + let submit_id = startup_submit_id.clone(); + let publication_gate = publication_gate.clone(); + let startup = async move { + if let Some(mut startup_receiver) = startup_receiver + && tokio::select! { + started = startup_receiver.wait_for(|started| *started) => started.is_err(), + () = cancel_token.cancelled() => true, + } + { + return (server_name, Err(StartupOutcomeError::Cancelled)); + } + if !publication_gate.wait().await { + return (server_name, Err(StartupOutcomeError::Cancelled)); + } + if let Some(tx_event) = tx_event.as_ref() { + let _ = emit_update( + submit_id.as_str(), + tx_event, + McpStartupUpdateEvent { + server: server_name.clone(), + status: McpStartupStatus::Starting, + }, + ) + .await; + } + let mut outcome = async_managed_client.client().await; + if cancel_token.is_cancelled() { + outcome = Err(StartupOutcomeError::Cancelled); + } + if let Some(tx_event) = tx_event.as_ref() { + let auth_state = match &outcome { + Err(error) if error.is_authentication_required() && !has_runtime_auth => { + match &configured_config.transport { + McpServerTransportConfig::StreamableHttp { + url, + bearer_token_env_var, + http_headers, + env_http_headers, + .. + } => { + match determine_streamable_http_auth_status_from_credentials( + configured_config + .oauth_credential_name(&server_name) + .as_ref(), + url, + bearer_token_env_var.as_deref(), + http_headers.clone(), + env_http_headers.clone(), + store_mode, + keyring_backend_kind, + ) { + Ok(auth_state) => auth_state, + Err(error) => { + warn!( + "failed to read stored auth status for MCP server `{server_name}`: {error:?}" + ); + None + } + } + } + McpServerTransportConfig::Stdio { .. } => None, + } + } + Ok(_) | Err(_) => None, + }; + if cancel_token.is_cancelled() { + outcome = Err(StartupOutcomeError::Cancelled); + } + let status = match &outcome { + Ok(_) => McpStartupStatus::Ready, + Err(StartupOutcomeError::Cancelled) => McpStartupStatus::Cancelled, + Err(error) => { + let reason = mcp_startup_failure_reason(auth_state, error); + let error_str = mcp_init_error_display( + server_name.as_str(), + Some(&configured_config), + error, + reason, + ); + McpStartupStatus::Failed { + error: error_str, + reason, + } + } + }; + + let _ = emit_update( + submit_id.as_str(), + tx_event, + McpStartupUpdateEvent { + server: server_name.clone(), + status, + }, + ) + .await; + } + if cancel_token.is_cancelled() { + outcome = Err(StartupOutcomeError::Cancelled); + } + + if matches!(&outcome, Err(StartupOutcomeError::Failed { .. })) { + async_managed_client.reconnect_failed_startup().await; + } + + (server_name, outcome) + }; + if defer_startup { + // Dormant servers must not hold the initial startup summary open. + tokio::spawn(startup); + } else { + join_set.spawn(startup); + } + } + let manager = Self { + servers, + event_stream_connection, + disabled_servers, + required_servers, + optional_startup_deadline: OnceLock::new(), + tool_plugin_context, + prefix_mcp_tool_names, + non_prefixed_mcp_tool_servers, + elicitation_requests: elicitation_requests.clone(), + trusted_access, + }; + let summary_publication_gate = publication_gate; + tokio::spawn(async move { + let outcomes = join_set.join_all().await; + if let Some(tx_event) = tx_event { + if !summary_publication_gate.wait().await { + return; + } + let mut summary = McpStartupCompleteEvent { + ready: reused_ready, + ..Default::default() + }; + for server_name in &summary.ready { + let _ = emit_update( + startup_submit_id.as_str(), + &tx_event, + McpStartupUpdateEvent { + server: server_name.clone(), + status: McpStartupStatus::Ready, + }, + ) + .await; + } + for (server_name, outcome) in outcomes { + match outcome { + Ok(_) => summary.ready.push(server_name), + Err(StartupOutcomeError::Cancelled) => summary.cancelled.push(server_name), + Err(StartupOutcomeError::Failed { error, .. }) => { + summary.failed.push(McpStartupFailure { + server: server_name, + error, + }) + } + } + } + let _ = tx_event + .send(Event { + id: startup_submit_id, + msg: EventMsg::McpStartupComplete(summary), + }) + .await; + } + }); + manager + } + + pub fn empty(prefix_mcp_tool_names: bool) -> Self { + Self { + servers: HashMap::new(), + event_stream_connection: None, + disabled_servers: Vec::new(), + required_servers: Vec::new(), + optional_startup_deadline: OnceLock::new(), + tool_plugin_context: Arc::new(ToolPluginContext::default()), + prefix_mcp_tool_names, + non_prefixed_mcp_tool_servers: Vec::new(), + elicitation_requests: ElicitationRequestManager::default(), + trusted_access: None, + } + } + + pub fn has_servers(&self) -> bool { + !self.servers.is_empty() + } + + pub(crate) fn contains_server(&self, server_name: &str) -> bool { + self.servers.contains_key(server_name) + } + + pub(crate) async fn authentication_failed_servers(&self) -> Vec { + let mut failed_servers = Vec::new(); + for (server_name, view) in &self.servers { + if view + .connection + .client + .startup_complete + .load(Ordering::Acquire) + && let Err(error) = view.connection.client().await + && error.is_authentication_required() + { + failed_servers.push(server_name.clone()); + } + } + failed_servers + } + + pub(crate) async fn updated_oauth_credentials_after_auth_failure( + &self, + config: &crate::McpConfig, + ) -> Vec { + let mut candidates = Vec::new(); + for server_name in self.authentication_failed_servers().await { + if let Some(view) = self.servers.get(&server_name) + && let Some(identity) = view.connection.identity.as_ref() + && let Some(server) = config.mcp_server_catalog.server(&server_name) + { + candidates.push((server_name, identity.clone(), server.config().clone())); + } + } + if candidates.is_empty() { + return Vec::new(); + } + + match tokio::task::spawn_blocking(move || { + candidates + .into_iter() + .filter_map(|(server_name, identity, config)| { + identity + .oauth_credentials_changed(&server_name, &config) + .then_some(server_name) + }) + .collect() + }) + .await + { + Ok(recovered_servers) => recovered_servers, + Err(error) => { + warn!(%error, "failed to inspect stored MCP OAuth credentials"); + Vec::new() + } + } + } + + pub(crate) async fn wait_for_server_startup(&self, server_name: &str) -> bool { + let Some(view) = self.servers.get(server_name) else { + return false; + }; + view.connection.client.ready_transport().is_some() || view.connection.client().await.is_ok() + } + + /// Stop all MCP clients owned by this manager and terminate stdio server processes. + pub async fn shutdown(&self) { + let connections = self + .servers + .values() + .map(|view| Arc::clone(&view.connection)) + .collect::>(); + // Keep cleanup alive if an interrupt cancels the refresh that requested it. + let shutdown_task = tokio::spawn(async move { + for connection in connections { + connection.shutdown().await; + } + }); + if let Err(error) = shutdown_task.await { + warn!("MCP client shutdown task failed: {error}"); + } + } + + pub(crate) fn cancel_startup(&self) { + for view in self.servers.values() { + view.connection.cancel_startup(); + } + } + + pub fn plugin_id_for_mcp_server_name(&self, server_name: &str) -> Option<&str> { + self.tool_plugin_context + .plugin_id_for_mcp_server_name(server_name) + } + + pub fn is_selected_plugin_mcp_server(&self, server_name: &str) -> bool { + self.tool_plugin_context + .is_selected_plugin_mcp_server(server_name) + } + + pub async fn wait_for_server_ready(&self, server_name: &str, timeout: Duration) -> bool { + let Some(view) = self.servers.get(server_name) else { + return false; + }; + + match tokio::time::timeout(timeout, view.connection.client()).await { + Ok(Ok(_)) => true, + Ok(Err(_)) | Err(_) => false, + } + } + + /// Invoke the tool indicated by the (server, tool) pair. + #[allow(clippy::too_many_arguments)] + pub async fn call_tool( + &self, + server: &str, + tool: &str, + environment_id: Option<&str>, + arguments: Option, + mut meta: Option, + requested_timeout: Option, + wait_for_server: bool, + ) -> Result { + let view = self + .servers + .get(server) + .ok_or_else(|| anyhow!("unknown MCP server '{server}'"))?; + if let Some(environment_id) = environment_id + && view.metadata.environment_id != environment_id + { + bail!( + "MCP server `{server}` is running in environment `{}`, expected `{environment_id}`", + view.metadata.environment_id + ); + } + if !view.tool_filter.allows(tool) { + return Err(anyhow!( + "tool '{tool}' is disabled for MCP server '{server}'" + )); + } + let client = if wait_for_server { + view.connection + .client() + .await + .context("failed to get client")? + } else { + let client = view + .connection + .client + .ready_client() + .ok_or_else(|| anyhow!("MCP server '{server}' is not connected"))?; + if client.client.is_closed().await { + bail!("MCP server '{server}' is not connected"); + } + client + }; + + let effective_timeout = match (view.tool_timeout, requested_timeout) { + (Some(server_timeout), Some(requested_timeout)) => { + Some(server_timeout.min(requested_timeout)) + } + (server_timeout, requested_timeout) => server_timeout.or(requested_timeout), + }; + // Direct callers cannot supply host-owned entitlement metadata, even for unlisted tools. + if let Some(serde_json::Value::Object(meta)) = meta.as_mut() { + meta.remove(ENTITLEMENT_CONTEXT_KEY); + } + let result: rmcp::model::CallToolResult = client + .client + .call_tool(tool.to_string(), arguments, meta, effective_timeout) + .await + .with_context(|| format!("tool call failed for `{server}/{tool}`"))?; + + Ok(call_tool_result_from_rmcp(result)) + } + + /// Capabilities belong to the initialized connection, never a shared tool cache. + pub(crate) fn list_available_server_capabilities(&self) -> HashMap { + self.servers + .iter() + .filter_map(|(name, view)| { + view.connection + .client + .server_capabilities + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + .map(|capabilities| (name.clone(), capabilities)) + }) + .collect() + } + + /// Returns presentation metadata from the current connection. + /// Codex Apps metadata may come from its existing cache; regular MCP server information is + /// connection-specific, so pending regular clients are awaited. + pub(crate) async fn list_available_server_infos(&self) -> HashMap { + let mut server_infos = HashMap::new(); + for (server_name, view) in &self.servers { + let client = &view.connection.client; + if !client.startup_complete.load(Ordering::Acquire) + && let Some(server_info) = client.cached_server_info.clone() + { + server_infos.insert(server_name.clone(), server_info); + continue; + } + match view.connection.client().await { + Ok(managed_client) => { + server_infos.insert(server_name.clone(), managed_client.server_info); + } + Err(_) => { + if let Some(server_info) = client.cached_server_info.clone() { + server_infos.insert(server_name.clone(), server_info); + } + } + } + } + server_infos + } +} + +#[cfg(test)] +#[path = "connection_manager_tests.rs"] +pub(crate) mod tests; diff --git a/codex-rs/codex-mcp/src/connection_manager/required.rs b/codex-rs/codex-mcp/src/connection_manager/required.rs new file mode 100644 index 0000000000000000000000000000000000000000..2ba8cb371d5a2758693f77a7688e3d0e86221df9 --- /dev/null +++ b/codex-rs/codex-mcp/src/connection_manager/required.rs @@ -0,0 +1,67 @@ +use anyhow::Result; +use anyhow::anyhow; +use codex_protocol::protocol::McpStartupFailure; +use tracing::Instrument; +use tracing::info_span; + +use super::McpConnectionSet; +use crate::rmcp_client::StartupOutcomeError; + +impl McpConnectionSet { + /// Waits for every required server and reports their startup failures together. + /// + /// The manager must already be reachable through [`crate::McpRuntime`] so + /// startup-time elicitation can resolve while validation waits. + pub(crate) async fn validate_required_servers(&self) -> Result<()> { + let failures = async { + let mut failures = Vec::new(); + for server_name in &self.required_servers { + let Some(view) = self.servers.get(server_name) else { + failures.push(McpStartupFailure { + server: server_name.clone(), + error: format!("required MCP server `{server_name}` was not initialized"), + }); + continue; + }; + if view.connection.startup_is_dormant() && view.connection.client.has_cached_tools() + { + continue; + } + + match view.connection.client().await { + Ok(_) => {} + Err(error) => failures.push(McpStartupFailure { + server: server_name.clone(), + error: startup_outcome_error_message(error), + }), + } + } + failures + } + .instrument(info_span!( + "session_init.required_mcp_wait", + otel.name = "session_init.required_mcp_wait", + session_init.required_mcp_server_count = self.required_servers.len(), + )) + .await; + if failures.is_empty() { + return Ok(()); + } + + let details = failures + .iter() + .map(|failure| format!("{}: {}", failure.server, failure.error)) + .collect::>() + .join("; "); + Err(anyhow!( + "required MCP servers failed to initialize: {details}" + )) + } +} + +fn startup_outcome_error_message(error: StartupOutcomeError) -> String { + match error { + StartupOutcomeError::Cancelled => "MCP startup cancelled".to_string(), + StartupOutcomeError::Failed { error, .. } => error, + } +} diff --git a/codex-rs/codex-mcp/src/connection_manager/resources.rs b/codex-rs/codex-mcp/src/connection_manager/resources.rs new file mode 100644 index 0000000000000000000000000000000000000000..c84cc6b43f97383497cb9a03ea91cc06f56e40da --- /dev/null +++ b/codex-rs/codex-mcp/src/connection_manager/resources.rs @@ -0,0 +1,174 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use rmcp::model::ListResourceTemplatesResult; +use rmcp::model::ListResourcesResult; +use rmcp::model::PaginatedRequestParams; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ReadResourceResult; +use rmcp::model::Resource; +use rmcp::model::ResourceTemplate; +use tokio::task::JoinSet; +use tracing::warn; + +use super::McpConnectionSet; +use crate::pagination::collect_paginated; +use crate::rmcp_client::ManagedClient; + +impl McpConnectionSet { + /// Returns resources from servers selected by `include_server`. + pub async fn list_all_resources( + &self, + include_server: impl Fn(&str) -> bool, + ) -> HashMap> { + let mut join_set = JoinSet::new(); + for (server_name, view) in self + .servers + .iter() + .filter(|(server_name, _)| include_server(server_name)) + { + let server_name = server_name.clone(); + let Ok(managed_client) = view.connection.client().await else { + continue; + }; + let timeout = view.tool_timeout; + let client = managed_client.client; + join_set.spawn(async move { + let resources = collect_paginated("resources/list", timeout, |params| { + let client = Arc::clone(&client); + async move { + let response = client.list_resources(params, timeout).await?; + Ok((response.resources, response.next_cursor)) + } + }) + .await; + (server_name, resources) + }); + } + + let mut resources = HashMap::new(); + while let Some(result) = join_set.join_next().await { + match result { + Ok((server_name, Ok(server_resources))) => { + resources.insert(server_name, server_resources); + } + Ok((server_name, Err(error))) => { + warn!("Failed to list resources for MCP server '{server_name}': {error:#}"); + } + Err(error) => { + warn!("Task panic when listing resources for MCP server: {error:#}"); + } + } + } + resources + } + + /// Returns resource templates from servers selected by `include_server`. + pub async fn list_all_resource_templates( + &self, + include_server: impl Fn(&str) -> bool, + ) -> HashMap> { + let mut join_set = JoinSet::new(); + for (server_name, view) in self + .servers + .iter() + .filter(|(server_name, _)| include_server(server_name)) + { + let server_name = server_name.clone(); + let Ok(managed_client) = view.connection.client().await else { + continue; + }; + let timeout = view.tool_timeout; + let client = managed_client.client; + join_set.spawn(async move { + let templates = collect_paginated("resources/templates/list", timeout, |params| { + let client = Arc::clone(&client); + async move { + let response = client.list_resource_templates(params, timeout).await?; + Ok((response.resource_templates, response.next_cursor)) + } + }) + .await; + (server_name, templates) + }); + } + + let mut templates = HashMap::new(); + while let Some(result) = join_set.join_next().await { + match result { + Ok((server_name, Ok(server_templates))) => { + templates.insert(server_name, server_templates); + } + Ok((server_name, Err(error))) => { + warn!( + "Failed to list resource templates for MCP server '{server_name}': {error:#}" + ); + } + Err(error) => { + warn!("Task panic when listing resource templates for MCP server: {error:#}"); + } + } + } + templates + } + + pub async fn list_resources( + &self, + server: &str, + params: Option, + ) -> Result { + let (managed, timeout) = self.client_by_name(server).await?; + managed + .client + .list_resources(params, timeout) + .await + .with_context(|| format!("resources/list failed for `{server}`")) + } + + pub async fn list_resource_templates( + &self, + server: &str, + params: Option, + ) -> Result { + let (managed, timeout) = self.client_by_name(server).await?; + managed + .client + .list_resource_templates(params, timeout) + .await + .with_context(|| format!("resources/templates/list failed for `{server}`")) + } + + pub async fn read_resource( + &self, + server: &str, + params: ReadResourceRequestParams, + ) -> Result { + let (managed, timeout) = self.client_by_name(server).await?; + let uri = params.uri.clone(); + managed + .client + .read_resource(params, timeout) + .await + .with_context(|| format!("resources/read failed for `{server}` ({uri})")) + } + + pub(crate) async fn client_by_name( + &self, + name: &str, + ) -> Result<(ManagedClient, Option)> { + let view = self + .servers + .get(name) + .ok_or_else(|| anyhow!("unknown MCP server '{name}'"))?; + let client = view + .connection + .client() + .await + .context("failed to get client")?; + Ok((client, view.tool_timeout)) + } +} diff --git a/codex-rs/codex-mcp/src/connection_manager/startup.rs b/codex-rs/codex-mcp/src/connection_manager/startup.rs new file mode 100644 index 0000000000000000000000000000000000000000..5d4c1fdd539908ef3e75159b83ac5d169253dade --- /dev/null +++ b/codex-rs/codex-mcp/src/connection_manager/startup.rs @@ -0,0 +1,133 @@ +use std::collections::HashMap; + +use anyhow::Result; +use async_channel::Sender; +use codex_api::SharedAuthProvider; +use codex_config::McpServerAuth; +use codex_config::McpServerConfig; +use codex_config::McpServerTransportConfig; +use codex_protocol::protocol::Event; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::McpStartupFailureReason; +use codex_protocol::protocol::McpStartupUpdateEvent; +use codex_rmcp_client::McpAuthState; +use codex_rmcp_client::McpLoginRequirement; + +use crate::mcp::CODEX_APPS_MCP_SERVER_NAME; +use crate::rmcp_client::DEFAULT_STARTUP_TIMEOUT; +use crate::rmcp_client::StartupOutcomeError; +use crate::server::EffectiveMcpServer; + +/// Makes ChatGPT authentication available to servers that explicitly opt in. +pub(super) fn chatgpt_auth_provider_for_server( + server: &EffectiveMcpServer, + chatgpt_auth_provider: Option, +) -> Option { + if !matches!(&server.config().auth, McpServerAuth::ChatGpt) + || !server.config().is_local_environment() + { + return None; + } + chatgpt_auth_provider +} + +pub(super) fn should_share_codex_apps_tools_cache( + server_name: &str, + uses_env_bearer_token: bool, +) -> bool { + server_name == CODEX_APPS_MCP_SERVER_NAME && !uses_env_bearer_token +} + +pub(super) async fn emit_update( + submit_id: &str, + tx_event: &Sender, + update: McpStartupUpdateEvent, +) -> Result<(), async_channel::SendError> { + tx_event + .send(Event { + id: submit_id.to_string(), + msg: EventMsg::McpStartupUpdate(update), + }) + .await +} + +pub(super) fn mcp_startup_failure_reason( + auth_state: Option, + error: &StartupOutcomeError, +) -> Option { + if !error.is_authentication_required() { + return None; + } + match auth_state { + Some( + McpAuthState::LoggedOut(McpLoginRequirement::Reauthentication) | McpAuthState::OAuth, + ) => Some(McpStartupFailureReason::ReauthenticationRequired), + Some( + McpAuthState::Unsupported + | McpAuthState::Unknown + | McpAuthState::LoggedOut(McpLoginRequirement::Login) + | McpAuthState::BearerToken, + ) + | None => None, + } +} + +pub(super) fn mcp_init_error_display( + server_name: &str, + config: Option<&McpServerConfig>, + error: &StartupOutcomeError, + reason: Option, +) -> String { + let server_key = if !server_name.is_empty() + && server_name + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-') + { + server_name.to_string() + } else { + serde_json::Value::String(server_name.to_string()).to_string() + }; + if let Some(McpServerTransportConfig::StreamableHttp { + url, + bearer_token_env_var, + http_headers, + .. + }) = config.map(|config| &config.transport) + && url == "https://api.githubcopilot.com/mcp/" + && bearer_token_env_var.is_none() + && http_headers.as_ref().map(HashMap::is_empty).unwrap_or(true) + { + format!( + "GitHub MCP does not support OAuth. Log in by adding a personal access token (https://github.com/settings/personal-access-tokens) to your environment and config.toml:\n[mcp_servers.{server_key}]\nbearer_token_env_var = CODEX_GITHUB_PERSONAL_ACCESS_TOKEN" + ) + } else if error.is_authentication_required() { + let recovery_hint = if config.is_some_and(|config| !config.is_local_environment()) { + "Use your client's MCP OAuth sign-in flow.".to_string() + } else { + format!("Run `codex mcp login {server_name}`.") + }; + let auth_status = match reason { + Some(McpStartupFailureReason::ReauthenticationRequired) => { + "requires OAuth reauthentication" + } + None => "is not logged in", + }; + format!("The {server_name} MCP server {auth_status}. {recovery_hint}") + } else if matches!( + error, + StartupOutcomeError::Failed { error, .. } + if error.contains("request timed out") + || error.contains("timed out handshaking with MCP server") + || error.contains("MCP client startup timed out") + ) { + let startup_timeout_secs = config + .and_then(|config| config.startup_timeout_sec) + .unwrap_or(DEFAULT_STARTUP_TIMEOUT) + .as_secs(); + format!( + "MCP client for `{server_name}` timed out after {startup_timeout_secs} seconds. Add or adjust `startup_timeout_sec` in your config.toml:\n[mcp_servers.{server_key}]\nstartup_timeout_sec = XX" + ) + } else { + format!("MCP client for `{server_name}` failed to start: {error:#}") + } +} diff --git a/codex-rs/codex-mcp/src/connection_manager/status.rs b/codex-rs/codex-mcp/src/connection_manager/status.rs new file mode 100644 index 0000000000000000000000000000000000000000..6e8852b0ea32c63b89841f440d9cab49caee1b25 --- /dev/null +++ b/codex-rs/codex-mcp/src/connection_manager/status.rs @@ -0,0 +1,30 @@ +//! Read-only connection state; observing a server must not start it. + +use std::collections::HashMap; + +use codex_protocol::mcp::McpServerConnectionStatus; + +use super::McpConnectionSet; + +impl McpConnectionSet { + pub(crate) async fn connection_statuses(&self) -> HashMap { + use McpServerConnectionStatus as Status; + + let mut statuses = self + .disabled_servers + .iter() + .map(|name| (name.clone(), Status::Disabled)) + .collect::>(); + for (name, view) in &self.servers { + let connection = &view.connection; + let client = &connection.client; + let status = if connection.startup_is_dormant() && !client.cancel_token.is_cancelled() { + Status::NotStarted + } else { + client.connection_status().await + }; + statuses.insert(name.clone(), status); + } + statuses + } +} diff --git a/codex-rs/codex-mcp/src/connection_manager/tool_catalog.rs b/codex-rs/codex-mcp/src/connection_manager/tool_catalog.rs new file mode 100644 index 0000000000000000000000000000000000000000..58b810366cb7a64581e91475e4ba2aff8ce09faa --- /dev/null +++ b/codex-rs/codex-mcp/src/connection_manager/tool_catalog.rs @@ -0,0 +1,552 @@ +use std::collections::HashMap; +use std::collections::HashSet; +use std::sync::Arc; +use std::sync::atomic::Ordering; +use std::time::Instant; + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use codex_connectors::ConnectorRuntimeFetchSource; +use futures::future::join_all; +use tracing::Instrument; +use tracing::instrument; +use tracing::trace; +use tracing::trace_span; + +use super::McpConnectionSet; +use super::McpServerMetadata; +use crate::binding::McpBinding; +use crate::binding::PreparedMcpCall; +use crate::binding_clients::McpBindingClients; +use crate::client_tool_catalog::ClientToolCatalogRevision; +use crate::client_tool_catalog::CodexAppsToolSnapshot; +use crate::client_tool_catalog::ToolCatalogSnapshot; +use crate::mcp::CODEX_APPS_MCP_SERVER_NAME; +use crate::rmcp_client::CODEX_APPS_REFRESH_DURATION_METRIC; +use crate::rmcp_client::MCP_TOOLS_LIST_DURATION_METRIC; +use crate::rmcp_client::ManagedClient; +use crate::rmcp_client::list_tools_for_client_uncached; +use crate::rmcp_client::prepare_codex_apps_tools_for_model; +use crate::runtime::emit_duration; +use crate::tools::ToolInfo; +use crate::tools::filter_tools; +use crate::tools::normalize_tools_for_model_with_prefix; + +const MCP_UI_META_KEY: &str = "ui"; +const MCP_UI_VISIBILITY_META_KEY: &str = "visibility"; +const MCP_UI_MODEL_VISIBILITY: &str = "model"; +/// Returns whether a tool may be included in model-facing tool declarations. +/// +/// Tools without visibility metadata remain visible. Tools with visibility +/// metadata are hidden unless they explicitly include `model`. +/// +/// +pub fn tool_is_model_visible(tool: &ToolInfo) -> bool { + let Some(visibility) = tool + .tool + .meta + .as_deref() + .and_then(|meta| meta.get(MCP_UI_META_KEY)) + .and_then(serde_json::Value::as_object) + .and_then(|ui| ui.get(MCP_UI_VISIBILITY_META_KEY)) + .and_then(serde_json::Value::as_array) + else { + return true; + }; + visibility + .iter() + .any(|target| target.as_str() == Some(MCP_UI_MODEL_VISIBILITY)) +} + +/// Catalog identity within one published connection set, including dormant servers. +#[derive(PartialEq, Eq)] +pub(crate) enum BindingCatalogRevision { + Ready(ClientToolCatalogRevision), + Dormant(u64), +} + +impl McpConnectionSet { + pub(crate) async fn stable_catalog_revisions( + &self, + required_servers: &[String], + required_plugins: &HashSet, + ) -> Option> { + let mut revisions = HashMap::new(); + for (server_name, view) in &self.servers { + if !view + .connection + .client + .startup_complete + .load(Ordering::Acquire) + { + // A cached catalog can remain stable without starting its server. + // Explicit requirements must still start the server during binding capture. + if !view.connection.startup_is_dormant() + || view.connection.client.cancel_token.is_cancelled() + || required_servers + .iter() + .any(|required| required == server_name) + || (self.is_selected_plugin_mcp_server(server_name) + && self + .plugin_id_for_mcp_server_name(server_name) + .is_some_and(|plugin_id| required_plugins.contains(plugin_id))) + { + return None; + } + let revision = view + .connection + .client + .tool_catalog_cache_context + .as_ref()? + .current_revision()?; + revisions.insert( + server_name.clone(), + BindingCatalogRevision::Dormant(revision), + ); + continue; + } + let Some(client) = view.connection.client.ready_client() else { + if !view.connection.client.is_codex_apps_mcp_server + && self.required_servers.binary_search(server_name).is_err() + && matches!(view.connection.client.client.peek(), Some(Err(_))) + { + continue; + } + return None; + }; + if client.client.is_closed().await { + return None; + } + let revision = client.tool_catalog.read(|catalog| catalog.revision).await; + revisions.insert( + server_name.clone(), + BindingCatalogRevision::Ready(ClientToolCatalogRevision { + catalog: Arc::clone(&client.tool_catalog), + revision, + }), + ); + } + Some(revisions) + } + + /// Returns all tools with model-visible names normalized. + pub async fn list_all_tools(&self) -> Vec { + self.list_tools_with_errors().await.0 + } + + #[instrument(level = "trace", skip_all, fields(mcp_server_count = self.servers.len()))] + pub(crate) async fn list_tools_with_errors(&self) -> (Vec, HashMap) { + let mut tools = Vec::new(); + let mut errors = HashMap::new(); + let mut available_server_count = 0; + let mut unavailable_server_count = 0; + let server_results = join_all(self.servers.iter().map(|(server_name, view)| async move { + view.connection.client.reconnect_failed_startup().await; + let has_cached_tools = view.connection.client.has_cached_tools(); + let startup_complete = view + .connection + .client + .startup_complete + .load(Ordering::Acquire); + let server_tools = view + .listed_tools(&self.tool_plugin_context) + .instrument(trace_span!( + "list_tools_for_server", + server_name = %server_name, + has_cached_tools, + startup_complete + )) + .await; + let result = match server_tools { + Ok(server_tools) => Ok(server_tools + .into_iter() + .map(|tool| Self::with_server_metadata(tool, &view.metadata)) + .collect::>()), + Err(error) => { + trace!( + server_name = %server_name, + has_cached_tools, + startup_complete, + "MCP server tools unavailable while building tool list" + ); + Err(error) + } + }; + (server_name, result) + })) + .await; + for (server_name, server_tools) in server_results { + match server_tools { + Ok(server_tools) => { + available_server_count += 1; + tools.extend(server_tools); + } + Err(error) => { + unavailable_server_count += 1; + errors.insert(server_name.clone(), error.to_string()); + } + } + } + let tools = normalize_tools_for_model_with_prefix( + tools, + self.prefix_mcp_tool_names, + &self.non_prefixed_mcp_tool_servers, + ); + trace!( + available_server_count, + unavailable_server_count, + tool_count = tools.len(), + "built MCP tool list" + ); + (tools, errors) + } + + #[instrument(level = "trace", skip_all)] + pub(crate) async fn capture_binding_with_metadata( + self: &Arc, + config: Arc, + plugins_available: bool, + required_servers: &[String], + required_plugins: &HashSet, + ) -> McpBinding { + let mut listed_tools = Vec::new(); + let mut clients = HashMap::new(); + let optional_mcp_startup_grace = config.optional_mcp_startup_grace; + let server_snapshots = join_all(self.servers.iter().map(|(server_name, view)| async move { + if !view + .connection + .client + .startup_complete + .load(Ordering::Acquire) + { + let required = self.required_servers.binary_search(server_name).is_ok(); + // Keep the catalog that lets us skip startup even if it expires during the wait. + let cached_tools = view.connection.client.cached_tools().filter(|tools| { + view.connection.client.is_codex_apps_mcp_server || !tools.is_empty() + }); + let has_cached_tools = cached_tools.is_some(); + let must_wait_for_startup = (required + && (!view.connection.startup_is_dormant() || !has_cached_tools)) + || required_servers + .iter() + .any(|required| required == server_name) + || (self.is_selected_plugin_mcp_server(server_name) + && self + .plugin_id_for_mcp_server_name(server_name) + .is_some_and(|plugin_id| required_plugins.contains(plugin_id))) + || (server_name == CODEX_APPS_MCP_SERVER_NAME && !has_cached_tools); + if !must_wait_for_startup && has_cached_tools { + return (server_name, view, cached_tools); + } + if !must_wait_for_startup && optional_mcp_startup_grace.is_zero() { + if let Some(cache) = view.connection.client.tool_catalog_cache_context.as_ref() + { + cache.optional_startup_deadline( + tokio::time::Instant::now(), + optional_mcp_startup_grace, + ); + } + let _ = view.connection.client().await; + } else if !must_wait_for_startup { + let optional_startup_deadline = if view.connection.startup_is_dormant() { + tokio::time::Instant::now() + optional_mcp_startup_grace + } else { + *self.optional_startup_deadline.get_or_init(|| { + tokio::time::Instant::now() + optional_mcp_startup_grace + }) + }; + let startup_deadline = view + .connection + .client + .tool_catalog_cache_context + .as_ref() + .map(|cache| { + cache.optional_startup_deadline( + optional_startup_deadline, + optional_mcp_startup_grace, + ) + }) + .unwrap_or(optional_startup_deadline); + if tokio::time::timeout_at(startup_deadline, view.connection.client()) + .await + .is_err() + { + trace!(server_name = %server_name, "omitting pending optional MCP server"); + } + return (server_name, view, cached_tools); + } + let _ = view.connection.client().await; + return (server_name, view, cached_tools); + } + (server_name, view, None) + })) + .await; + let server_results = join_all(server_snapshots.into_iter().map(|(server_name, view, cached_tools)| async move { + let (client, server_tools) = if !view + .connection + .client + .startup_complete + .load(Ordering::Acquire) + { + (None, view.connection.client.cached_tools_or(cached_tools)?) + } else { + view.connection.client.reconnect_failed_startup().await; + let Ok(mut client) = view.connection.client().await else { + trace!(server_name = %server_name, "omitting MCP server without an exact ready client"); + return None; + }; + client.tool_timeout = view.tool_timeout; + let snapshot = client.tool_catalog.read(Arc::new).await; + let server_tools = snapshot.tools.to_vec(); + (Some((Arc::new(client), snapshot)), server_tools) + }; + let server_tools = filter_tools(server_tools, &view.tool_filter); + let server_tools = if server_name == CODEX_APPS_MCP_SERVER_NAME { + prepare_codex_apps_tools_for_model(server_tools, &self.tool_plugin_context) + } else { + crate::rmcp_client::prepare_regular_mcp_tools_for_model( + server_tools, + &self.tool_plugin_context, + ) + }; + let server_tools = server_tools + .into_iter() + .map(|mut tool| { + if client.is_none() + && let Some(annotations) = tool.tool.annotations.as_mut() + { + annotations.read_only_hint = None; + } + Self::with_server_metadata(tool, &view.metadata) + }) + .collect::>(); + Some((server_name.clone(), client, server_tools)) + })) + .await; + for (server_name, client, server_tools) in server_results.into_iter().flatten() { + if let Some((client, snapshot)) = client { + clients.insert(server_name, (client, snapshot)); + } + listed_tools.extend(server_tools); + } + let listed_tools = normalize_tools_for_model_with_prefix( + listed_tools, + self.prefix_mcp_tool_names, + &self.non_prefixed_mcp_tool_servers, + ); + let mut tools = Vec::with_capacity(listed_tools.len()); + let mut calls = std::collections::HashMap::with_capacity(listed_tools.len()); + for tool_info in listed_tools { + let model_visible = crate::tool_is_model_visible(&tool_info); + let Some((client, snapshot)) = clients.get(&tool_info.server_name) else { + if model_visible { + tools.push(tool_info); + } + continue; + }; + let Some(call) = self.prepare_call( + &tool_info, + Arc::clone(client), + Arc::clone(&config), + Arc::clone(snapshot), + ) else { + trace!( + server_name = %tool_info.server_name, + tool_name = %tool_info.tool.name, + "omitting MCP tool without an exact ready client" + ); + continue; + }; + calls.insert( + ( + tool_info.server_name.clone(), + tool_info.tool.name.to_string(), + ), + call, + ); + if model_visible { + tools.push(tool_info); + } + } + let clients = Arc::new(McpBindingClients::new( + clients + .into_iter() + .map(|(server_name, (client, _))| (server_name, client)) + .collect(), + )); + McpBinding::new( + Arc::clone(self), + clients, + config, + plugins_available, + tools, + calls, + ) + } + + fn prepare_call( + self: &Arc, + tool_info: &ToolInfo, + client: Arc, + config: Arc, + tool_catalog_snapshot: Arc, + ) -> Option { + let server_name = &tool_info.server_name; + let view = self.servers.get(server_name)?; + PreparedMcpCall::new( + Arc::clone(self), + client, + config, + tool_catalog_snapshot, + tool_info.clone(), + view.metadata.clone(), + self.plugin_id_for_mcp_server_name(server_name) + .map(str::to_string), + self.is_selected_plugin_mcp_server(server_name), + ) + } + + /// Refreshes one exact Apps catalog, preserving the raw inventory for app policy. + pub(crate) async fn refresh_codex_apps_client_catalog( + &self, + config: &crate::McpConfig, + ) -> Result { + let refresh_start = Instant::now(); + let view = self + .servers + .get(CODEX_APPS_MCP_SERVER_NAME) + .ok_or_else(|| anyhow!("unknown MCP server '{CODEX_APPS_MCP_SERVER_NAME}'"))?; + let (tools, _) = self.refresh_codex_apps_tool_catalog().await?; + let server_has_permission = config + .permission_profile_for_server(CODEX_APPS_MCP_SERVER_NAME) + .is_some(); + let model_visible_tool_names = tools + .iter() + .filter(|tool| { + server_has_permission + && view.tool_filter.allows(&tool.tool.name) + && self + .tool_plugin_context + .allows_connector_id(tool.connector_id.as_deref()) + && tool_is_model_visible(tool) + }) + .map(|tool| tool.tool.name.to_string()) + .collect(); + emit_duration( + CODEX_APPS_REFRESH_DURATION_METRIC, + refresh_start.elapsed(), + &[("path", "legacy"), ("trigger", "explicit")], + ); + Ok(CodexAppsToolSnapshot { + tools, + model_visible_tool_names, + }) + } + + /// Refreshes Apps tools and returns the prepared shared-cache winner for discovery. + pub async fn refresh_codex_apps_tools_for_discovery(&self) -> Result> { + let refresh_start = Instant::now(); + let view = self + .servers + .get(CODEX_APPS_MCP_SERVER_NAME) + .ok_or_else(|| anyhow!("unknown MCP server '{CODEX_APPS_MCP_SERVER_NAME}'"))?; + let (_, tools) = self.refresh_codex_apps_tool_catalog().await?; + let tools = prepare_codex_apps_tools_for_model( + filter_tools(tools, &view.tool_filter), + &self.tool_plugin_context, + ) + .into_iter() + .map(|tool| Self::with_server_metadata(tool, &view.metadata)); + let tools = normalize_tools_for_model_with_prefix( + tools, + self.prefix_mcp_tool_names, + &self.non_prefixed_mcp_tool_servers, + ); + emit_duration( + CODEX_APPS_REFRESH_DURATION_METRIC, + refresh_start.elapsed(), + &[("path", "legacy"), ("trigger", "explicit")], + ); + Ok(tools) + } + + /// Publishes the exact client catalog and returns both raw inventories. + async fn refresh_codex_apps_tool_catalog(&self) -> Result<(Vec, Vec)> { + let view = self + .servers + .get(CODEX_APPS_MCP_SERVER_NAME) + .ok_or_else(|| anyhow!("unknown MCP server '{CODEX_APPS_MCP_SERVER_NAME}'"))?; + let managed_client = view + .connection + .client() + .await + .context("failed to get client")?; + let (tools, list_start) = managed_client + .tool_catalog + .refresh( + || async { + let list_start = Instant::now(); + let fetch_ticket = managed_client.codex_apps_tools_cache_context.as_ref().map( + |cache_context| { + cache_context.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh) + }, + ); + let client_tools = list_tools_for_client_uncached( + CODEX_APPS_MCP_SERVER_NAME, + /*is_codex_apps_mcp_server*/ true, + /*codex_apps_refresh_trigger*/ "explicit", + &managed_client.client, + view.tool_timeout, + view.catalog_item_limit, + managed_client.server_instructions.as_deref(), + ) + .await + .with_context(|| { + format!( + "failed to refresh tools for MCP server '{CODEX_APPS_MCP_SERVER_NAME}'" + ) + })?; + Ok((client_tools, (fetch_ticket, list_start))) + }, + |client_tools, (fetch_ticket, list_start)| { + // Discovery can accept another scope's winner; executable catalogs + // receive only the latest successful fetch from their own scope. + let tools = match ( + managed_client.codex_apps_tools_cache_context.as_ref(), + fetch_ticket, + ) { + (Some(cache_context), Some(fetch_ticket)) => cache_context + .publish_if_newest_accepted( + fetch_ticket, + &managed_client.server_info, + client_tools.to_vec(), + ), + (None, None) => client_tools.to_vec(), + _ => unreachable!("Codex Apps fetch ticket requires cache context"), + }; + (tools, list_start) + }, + ) + .await?; + let client_tools = managed_client + .tool_catalog + .read(|catalog| catalog.tools.to_vec()) + .await; + emit_duration( + MCP_TOOLS_LIST_DURATION_METRIC, + list_start.elapsed(), + &[("cache", "miss")], + ); + Ok((client_tools, tools)) + } + + fn with_server_metadata(mut tool: ToolInfo, metadata: &McpServerMetadata) -> ToolInfo { + tool.supports_parallel_tool_calls = metadata.supports_parallel_tool_calls; + tool.server_origin = metadata + .origin + .as_ref() + .map(|origin| origin.as_str().to_string()); + tool + } +} diff --git a/codex-rs/codex-mcp/src/connection_manager_tests.rs b/codex-rs/codex-mcp/src/connection_manager_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..0d10a93e7077b48255f23f7c7c5d46dd9b5235f8 --- /dev/null +++ b/codex-rs/codex-mcp/src/connection_manager_tests.rs @@ -0,0 +1,6322 @@ +use super::*; +use crate::McpBinding; +use crate::client_tool_catalog::ClientToolCatalog; +use crate::elicitation::ElicitationLifecycle; +use crate::elicitation::ElicitationRequestManager; +use crate::elicitation::ElicitationRequestRouter; +use crate::elicitation::ElicitationReviewRequest; +use crate::elicitation::ElicitationReviewer; +use crate::elicitation::elicitation_is_rejected_by_policy; +use crate::mcp::tests::test_elicitation_config; +use crate::rmcp_client::AsyncManagedClient; +use crate::rmcp_client::CODEX_APPS_RECONNECT_INITIAL_BACKOFF; +use crate::rmcp_client::CodexAppsStartupReconnect; +use crate::rmcp_client::ManagedClient; +use crate::rmcp_client::ManagedClientFuture; +use crate::rmcp_client::StartupOutcomeError; +use crate::rmcp_client::list_tools_for_client_uncached; +use crate::runtime::McpRuntimeContext; +use crate::server::EffectiveMcpServer; +use crate::server::McpServerMetadata; +use crate::server::McpServerOrigin; +use crate::tool_catalog_cache::McpToolCatalogCache; +use crate::tools::ToolFilter; +use crate::tools::ToolInfo; +use crate::tools::filter_tools; +use crate::tools::normalize_tools_for_model_with_prefix; +use assert_matches::assert_matches; +use codex_config::AppToolApproval; +use codex_config::Constrained; +use codex_config::McpServerAuth; +use codex_config::McpServerConfig; +use codex_config::McpServerEnvVar; +use codex_config::McpServerToolConfig; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_connectors::ConnectorRuntimeContext; +use codex_connectors::ConnectorRuntimeContextKey; +use codex_connectors::ConnectorRuntimeFetchSource; +use codex_connectors::ConnectorRuntimeManager; +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +use codex_exec_server_test_support::environment_manager_without_environments; +use codex_login::AuthHeaders; +use codex_login::AuthManager; +use codex_login::CodexAuth; +use codex_protocol::ToolName; +use codex_protocol::approvals::ElicitationRequest; +use codex_protocol::mcp::ClientMcpExtensions; +use codex_protocol::mcp::McpServerInfo; +use codex_protocol::models::PermissionProfile; +use codex_protocol::protocol::AskForApproval; +use codex_protocol::protocol::GranularApprovalConfig; +use codex_protocol::protocol::McpStartupFailureReason; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::InProcessTransportFactory; +use codex_rmcp_client::McpAuthState; +use codex_rmcp_client::McpLoginRequirement; +use codex_rmcp_client::McpOAuthRefreshMode; +use codex_rmcp_client::RmcpClient; +use codex_utils_path_uri::PathUri; +use futures::FutureExt; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; +use rmcp::ErrorData as McpError; +use rmcp::RoleServer; +use rmcp::ServerHandler; +use rmcp::ServiceExt; +use rmcp::model::ClientCapabilities; +use rmcp::model::ElicitRequestParams; +use rmcp::model::ElicitationAction; +use rmcp::model::ElicitationCapability; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::JsonObject; +use rmcp::model::ListToolsResult; +use rmcp::model::NumberOrString; +use rmcp::model::PaginatedRequestParams; +use rmcp::model::ProtocolVersion; +use rmcp::model::ServerCapabilities; +use rmcp::model::ServerInfo; +use rmcp::model::Tool; +use rmcp::service::RequestContext; +use std::collections::HashMap; +use std::collections::HashSet; +use std::io; +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use tempfile::tempdir; +use tokio::io::DuplexStream; +use tokio::sync::Notify; +use tokio_util::sync::CancellationToken; + +impl McpConnectionSet { + fn new_uninitialized( + approval_policy: &Constrained, + permission_profile: &Constrained, + prefix_mcp_tool_names: bool, + ) -> Self { + Self { + servers: HashMap::new(), + event_stream_connection: None, + disabled_servers: Vec::new(), + required_servers: Vec::new(), + optional_startup_deadline: OnceLock::new(), + tool_plugin_context: Arc::new(ToolPluginContext::default()), + prefix_mcp_tool_names, + non_prefixed_mcp_tool_servers: Vec::new(), + elicitation_requests: ElicitationRequestManager::new( + test_elicitation_config( + "server", + approval_policy.value(), + permission_profile.get().clone(), + ), + /*reviewer*/ None, + /*lifecycle*/ None, + ElicitationRequestRouter::default(), + ), + trusted_access: None, + } + } + + pub(crate) fn insert_test_client( + &mut self, + name: impl Into, + client: AsyncManagedClient, + ) { + let name = name.into(); + self.servers.insert( + name, + McpServerView { + tool_filter: ToolFilter::default(), + protocol_mode: crate::McpProtocolMode::Legacy, + connection: Arc::new(McpServerConnection { + identity: None, + client, + startup_timeout: DEFAULT_STARTUP_TIMEOUT, + startup_trigger: None, + _diagnostics_guard: LIVE_CONNECTIONS.track(), + }), + metadata: McpServerMetadata { + environment_id: String::new(), + pollutes_memory: true, + origin: None, + supports_parallel_tool_calls: false, + default_tools_approval_mode: None, + tool_approval_modes: HashMap::new(), + }, + tool_timeout: None, + catalog_item_limit: crate::pagination::MAX_MCP_CATALOG_ITEMS, + }, + ); + } + + fn test_client(&self, name: &str) -> &AsyncManagedClient { + &self.servers[name].connection.client + } + + fn set_test_server_metadata(&mut self, name: &str, metadata: McpServerMetadata) { + self.servers + .get_mut(name) + .expect("test server exists") + .metadata = metadata; + } + + fn shares_test_connection_with(&self, other: &Self, name: &str) -> bool { + let Some(left) = self.servers.get(name) else { + return false; + }; + let Some(right) = other.servers.get(name) else { + return false; + }; + Arc::ptr_eq(&left.connection, &right.connection) + } +} + +fn create_test_tool(server_name: &str, tool_name: &str) -> ToolInfo { + ToolInfo { + server_name: server_name.to_string(), + supports_parallel_tool_calls: false, + server_origin: None, + callable_name: tool_name.to_string(), + callable_namespace: server_name.to_string(), + namespace_description: None, + tool: Tool::new( + tool_name.to_string(), + format!("Test tool: {tool_name}"), + Arc::new(JsonObject::default()), + ), + openai_file_input_optional_fields: Default::default(), + connector_id: None, + connector_name: None, + plugin_display_names: Vec::new(), + } +} + +fn create_codex_apps_tools_cache_context( + codex_home: PathBuf, + account_id: Option<&str>, + chatgpt_user_id: Option<&str>, +) -> ConnectorRuntimeContext { + ConnectorRuntimeManager::::default().context( + codex_home, + ConnectorRuntimeContextKey::personal( + account_id.map(ToOwned::to_owned), + chatgpt_user_id.map(ToOwned::to_owned), + ), + ) +} + +fn store_current_tools(cache_context: &ConnectorRuntimeContext, tools: Vec) { + let _ = cache_context.publish_if_newest_accepted( + cache_context.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh), + &create_test_server_info("Codex Apps"), + tools, + ); +} + +async fn capture_binding(manager: &Arc) -> McpBinding { + let mut config = crate::mcp::tests::test_mcp_config(std::env::temp_dir()); + config.server_permission_profiles = manager + .servers + .keys() + .map(|name| (name.clone(), PermissionProfile::default())) + .collect(); + manager + .capture_binding_with_metadata( + Arc::new(config), + /*plugins_available*/ false, + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await +} + +fn create_test_server_info(title: &str) -> McpServerInfo { + McpServerInfo { + name: "codex-apps".to_string(), + title: Some(title.to_string()), + version: "1.0.0".to_string(), + description: None, + icons: None, + website_url: None, + } +} + +struct TestInProcessTransportFactory; + +struct PendingHttpClient; + +impl HttpClient for PendingHttpClient { + fn http_request( + &self, + _params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + futures::future::pending().boxed() + } + + fn http_request_stream( + &self, + _params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + futures::future::pending().boxed() + } +} + +impl InProcessTransportFactory for TestInProcessTransportFactory { + fn open(&self) -> BoxFuture<'static, io::Result> { + async { + let (client_stream, _server_stream) = tokio::io::duplex(1); + Ok(client_stream) + } + .boxed() + } +} + +#[derive(Clone)] +struct RefreshTestTransportFactory { + tool: Tool, + list_started: Option>, + release_list: Option>, + next_cursor: Option, + list_requests: Arc, +} + +impl ServerHandler for RefreshTestTransportFactory { + fn get_info(&self) -> ServerInfo { + ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) + } + + async fn list_tools( + &self, + _request: Option, + _context: rmcp::service::RequestContext, + ) -> Result { + self.list_requests + .fetch_add(1, std::sync::atomic::Ordering::SeqCst); + if let Some(list_started) = &self.list_started { + list_started.notify_one(); + } + if let Some(release_list) = &self.release_list { + release_list.notified().await; + } + let mut result = ListToolsResult::with_all_items(vec![self.tool.clone()]); + result.next_cursor = self.next_cursor.clone(); + Ok(result) + } +} + +impl InProcessTransportFactory for RefreshTestTransportFactory { + fn open(&self) -> BoxFuture<'static, io::Result> { + let server = self.clone(); + async move { + let (client_stream, server_stream) = tokio::io::duplex(4096); + tokio::spawn(async move { + let server = server + .serve(server_stream) + .await + .expect("serve test MCP server"); + server.waiting().await.expect("wait for test MCP server"); + }); + Ok(client_stream) + } + .boxed() + } +} + +#[derive(Clone)] +struct MutableToolsServer { + tools: Arc>>, + block_tool_listing: Arc, +} + +impl ServerHandler for MutableToolsServer { + fn get_info(&self) -> ServerInfo { + ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) + } + + async fn list_tools( + &self, + _request: Option, + _context: RequestContext, + ) -> Result { + if self.block_tool_listing.load(Ordering::Acquire) { + std::future::pending::<()>().await; + } + Ok(ListToolsResult { + tools: self.tools.read().await.clone(), + ..Default::default() + }) + } +} + +struct MutableToolsTransportFactory { + server: MutableToolsServer, +} + +impl InProcessTransportFactory for MutableToolsTransportFactory { + fn open(&self) -> BoxFuture<'static, io::Result> { + let server = self.server.clone(); + async move { + let (client_stream, server_stream) = tokio::io::duplex(4096); + tokio::spawn(async move { + server + .serve(server_stream) + .await + .expect("serve mutable MCP tools") + .waiting() + .await + .expect("mutable MCP tools server completes"); + }); + Ok(client_stream) + } + .boxed() + } +} + +struct DisconnectingToolsTransportFactory { + server: MutableToolsServer, + disconnect: CancellationToken, +} + +impl InProcessTransportFactory for DisconnectingToolsTransportFactory { + fn open(&self) -> BoxFuture<'static, io::Result> { + let server = self.server.clone(); + let disconnect = self.disconnect.clone(); + async move { + let (client_stream, server_stream) = tokio::io::duplex(4096); + tokio::spawn(async move { + let server = server + .serve(server_stream) + .await + .expect("serve disconnecting MCP tools"); + let cancellation = server.cancellation_token(); + tokio::select! { + () = disconnect.cancelled() => cancellation.cancel(), + result = server.waiting() => { + result.expect("disconnecting MCP server should complete"); + } + } + }); + Ok(client_stream) + } + .boxed() + } +} + +#[tokio::test] +async fn legacy_tool_catalog_does_not_follow_pagination_cursor() -> anyhow::Result<()> { + let requests = Arc::new(AtomicUsize::new(0)); + let client = Arc::new( + RmcpClient::new_in_process_client(Arc::new(RefreshTestTransportFactory { + tool: create_test_tool("legacy", "first-page").tool, + list_started: None, + release_list: None, + next_cursor: Some("next-page".to_string()), + list_requests: Arc::clone(&requests), + })) + .await?, + ); + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("codex-test", "0.0.0-test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(5)), + Box::new(|_, _| async { Err(anyhow!("unexpected elicitation")) }.boxed()), + ) + .await?; + + let tools = list_tools_for_client_uncached( + "legacy", + /*is_codex_apps_mcp_server*/ false, + "test", + &client, + Some(Duration::from_secs(5)), + crate::pagination::MAX_MCP_CATALOG_ITEMS, + /*server_instructions*/ None, + ) + .await?; + + assert_eq!(tools.len(), 1); + assert_eq!(tools[0].tool.name.as_ref(), "first-page"); + assert_eq!(requests.load(std::sync::atomic::Ordering::SeqCst), 1); + client.shutdown().await; + Ok(()) +} + +async fn create_test_managed_client(tools: Vec) -> ManagedClient { + ManagedClient { + _auth_change_notifications: None, + client: Arc::new( + RmcpClient::new_in_process_client(Arc::new(TestInProcessTransportFactory)) + .await + .expect("create in-process RMCP client"), + ), + server_info: create_test_server_info("Ready"), + tool_catalog: Arc::new(ClientToolCatalog::new(tools, /*updates*/ None)), + tool_timeout: None, + server_instructions: None, + server_supports_sandbox_state_meta_capability: false, + codex_apps_tools_cache_context: None, + } +} + +#[tokio::test(start_paused = true)] +async fn prepared_call_timeout_includes_trusted_access_lookup() { + let mut tool = create_test_tool("docs", "access"); + tool.tool.annotations = Some(rmcp::model::ToolAnnotations::new().read_only(true)); + let mut tool_meta = rmcp::model::MetaObject::new(); + tool_meta.insert( + "openai/requestedEntitlements".to_string(), + serde_json::json!(["cyber_trusted_access"]), + ); + tool.tool.meta = Some(tool_meta); + + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + let mut config = crate::mcp::tests::test_mcp_config(std::env::temp_dir()); + let mut catalog = crate::ResolvedMcpCatalog::builder(); + catalog.register(crate::McpServerRegistration::from_plugin( + "docs".to_string(), + crate::McpPluginAttribution::new("docs@test".to_string(), "Docs".to_string()), + /*plugin_order*/ 0, + serde_json::from_value(serde_json::json!({ "command": "docs" })) + .expect("plugin MCP config"), + )); + config.mcp_server_catalog = catalog.build(); + config + .server_permission_profiles + .insert("docs".to_string(), PermissionProfile::default()); + manager.tool_plugin_context = Arc::new(crate::tool_plugin_context(&config)); + let auth = CodexAuth::create_dummy_chatgpt_auth_for_testing(); + manager.trusted_access = Some(TrustedAccessContext::new( + auth.clone(), + AuthManager::from_auth_for_testing(auth), + "https://chatgpt.com/backend-api".to_string(), + Arc::new(PendingHttpClient), + )); + let manager = Arc::new(manager); + let server_metadata = McpServerMetadata { + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + pollutes_memory: true, + origin: Some(McpServerOrigin::Stdio), + supports_parallel_tool_calls: false, + default_tools_approval_mode: None, + tool_approval_modes: HashMap::new(), + }; + let client = Arc::new(create_test_managed_client(vec![tool.clone()]).await); + let catalog_snapshot = client.tool_catalog.read(Arc::new).await; + let prepared = crate::PreparedMcpCall::new( + manager, + client, + Arc::new(config), + catalog_snapshot, + tool, + server_metadata, + Some("docs@test".to_string()), + /*selected_plugin_server*/ false, + ) + .expect("docs should retain its permission profile"); + + let started = tokio::time::Instant::now(); + let error = prepared + .call( + Some(serde_json::json!({})), + /*meta*/ None, + Some(Duration::from_secs(1)), + ) + .await + .expect_err("trusted access lookup should consume the call timeout"); + + assert_eq!(started.elapsed(), Duration::from_secs(1)); + assert!(format!("{error:#}").contains("timed out awaiting tools/call after 1s")); +} + +pub(crate) async fn create_ready_async_managed_client(tools: Vec) -> AsyncManagedClient { + AsyncManagedClient { + client: futures::future::ready::>(Ok( + create_test_managed_client(tools).await, + )) + .boxed() + .shared(), + is_codex_apps_mcp_server: false, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(true)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + } +} + +#[tokio::test] +async fn connection_statuses_observe_clients_without_starting_them() { + use codex_protocol::mcp::McpServerConnectionStatus as Status; + + let mut manager = McpConnectionSet::empty(/*prefix_mcp_tool_names*/ true); + manager.disabled_servers.push("disabled".to_string()); + let ready = create_ready_async_managed_client(Vec::new()).await; + ready.client().await.expect("ready client"); + manager.insert_test_client("connected", ready); + for (name, error) in [ + ( + "failed", + StartupOutcomeError::Failed { + error: "broken".to_string(), + is_authentication_required: false, + }, + ), + ( + "auth", + StartupOutcomeError::Failed { + error: "login".to_string(), + is_authentication_required: true, + }, + ), + ( + "flattened-auth", + StartupOutcomeError::from(anyhow!("Auth required for server")), + ), + ("cancelled", StartupOutcomeError::Cancelled), + ] { + let mut client = create_ready_async_managed_client(Vec::new()).await; + client.client = futures::future::ready(Err(error)).boxed().shared(); + assert!(client.client().await.is_err()); + manager.insert_test_client(name, client); + } + let mut pending = create_ready_async_managed_client(Vec::new()).await; + pending.client = futures::future::pending().boxed().shared(); + pending.cached_server_info = Some(create_test_server_info("Cached")); + manager.insert_test_client("starting", pending.clone()); + manager.insert_test_client("deferred", pending); + let (trigger, _receiver) = watch::channel(/*init*/ false); + Arc::get_mut(&mut manager.servers.get_mut("deferred").unwrap().connection) + .unwrap() + .startup_trigger = Some(trigger.clone()); + + let statuses = tokio::time::timeout( + Duration::from_millis(/*millis*/ 100), + manager.connection_statuses(), + ) + .await + .expect("status must not await startup"); + let mut expected = HashMap::from([ + ("connected".to_string(), Status::Connected), + ("failed".to_string(), Status::Failed), + ("auth".to_string(), Status::AuthenticationRequired), + ("flattened-auth".to_string(), Status::AuthenticationRequired), + ("cancelled".to_string(), Status::Cancelled), + ("starting".to_string(), Status::Starting), + ("deferred".to_string(), Status::NotStarted), + ("disabled".to_string(), Status::Disabled), + ]); + assert_eq!(statuses, expected); + assert!(!*trigger.borrow()); + manager.test_client("connected").cancel_token.cancel(); + expected.insert("connected".to_string(), Status::Cancelled); + assert_eq!(manager.connection_statuses().await, expected); +} + +#[tokio::test(start_paused = true)] +async fn connection_statuses_follow_latest_reconnect_outcome() { + use codex_protocol::mcp::McpServerConnectionStatus as Status; + + let recovered = create_test_managed_client(Vec::new()).await; + let attempts = Arc::new(AtomicUsize::new(0)); + let started = Arc::new(Notify::new()); + let release = Arc::new(Notify::new()); + let finished = Arc::new(Notify::new()); + let factory = { + let attempts = Arc::clone(&attempts); + let started = Arc::clone(&started); + let release = Arc::clone(&release); + let finished = Arc::clone(&finished); + Arc::new(move || { + let attempt = attempts.fetch_add(1, Ordering::SeqCst); + let recovered = recovered.clone(); + let started = Arc::clone(&started); + let release = Arc::clone(&release); + let finished = Arc::clone(&finished); + async move { + started.notify_one(); + release.notified().await; + finished.notify_one(); + match attempt { + 0 | 1 => Err(StartupOutcomeError::Failed { + error: "retry failed".to_string(), + is_authentication_required: attempt == 0, + }), + _ => Ok(recovered), + } + } + .boxed() + .shared() + }) + }; + let manager = create_test_manager_with_failed_apps_startup(Vec::new(), factory); + let client = manager.test_client(CODEX_APPS_MCP_SERVER_NAME); + assert!(client.client().await.is_err()); + let expected = |status| HashMap::from([(CODEX_APPS_MCP_SERVER_NAME.to_string(), status)]); + assert_eq!( + manager.connection_statuses().await, + expected(Status::Failed) + ); + + for status in [ + Status::AuthenticationRequired, + Status::Failed, + Status::Connected, + ] { + client.reconnect_failed_startup().await; + started.notified().await; + assert_eq!( + manager.connection_statuses().await, + expected(Status::Starting) + ); + release.notify_one(); + finished.notified().await; + assert_eq!(manager.connection_statuses().await, expected(status)); + tokio::time::advance(CODEX_APPS_RECONNECT_INITIAL_BACKOFF * 2).await; + } + assert_eq!(attempts.load(Ordering::SeqCst), 3); +} + +fn create_gated_async_managed_client( + client: ManagedClient, +) -> ( + AsyncManagedClient, + tokio::sync::oneshot::Receiver<()>, + tokio::sync::oneshot::Sender<()>, +) { + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = tokio::sync::oneshot::channel(); + let startup_complete = Arc::new(AtomicBool::new(false)); + let startup_complete_for_client = Arc::clone(&startup_complete); + let client = async move { + started_tx.send(()).expect("signal client startup"); + release_rx.await.expect("release client startup"); + startup_complete_for_client.store(true, std::sync::atomic::Ordering::Release); + Ok(client) + } + .boxed() + .shared(); + + ( + AsyncManagedClient { + client, + is_codex_apps_mcp_server: false, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: None, + startup_complete, + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + started_rx, + release_tx, + ) +} + +pub(crate) async fn create_test_manager_with_ready_apps_client( + cache_context: ConnectorRuntimeContext, + tool_name: &str, + list_started: Option>, + release_list: Option>, +) -> anyhow::Result> { + let tool = create_test_tool(CODEX_APPS_MCP_SERVER_NAME, tool_name); + let client = Arc::new( + RmcpClient::new_in_process_client(Arc::new(RefreshTestTransportFactory { + tool: tool.tool.clone(), + list_started, + release_list, + next_cursor: None, + list_requests: Arc::new(AtomicUsize::new(0)), + })) + .await?, + ); + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("codex-test", "0.0.0-test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(5)), + Box::new(|_, _| async { Err(anyhow!("unexpected elicitation")) }.boxed()), + ) + .await?; + + let managed_client = ManagedClient { + _auth_change_notifications: None, + client, + server_info: create_test_server_info("Codex Apps"), + tool_catalog: Arc::new(ClientToolCatalog::new(vec![tool], /*updates*/ None)), + tool_timeout: Some(Duration::from_secs(5)), + server_instructions: None, + server_supports_sandbox_state_meta_capability: false, + codex_apps_tools_cache_context: Some(cache_context.clone()), + }; + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: futures::future::ready::>(Ok( + managed_client, + )) + .boxed() + .shared(), + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: Some(create_test_server_info("Codex Apps")), + codex_apps_tools_cache_context: Some(cache_context), + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(true)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + manager.set_test_server_metadata( + CODEX_APPS_MCP_SERVER_NAME, + McpServerMetadata { + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + pollutes_memory: false, + origin: None, + supports_parallel_tool_calls: false, + default_tools_approval_mode: None, + tool_approval_modes: HashMap::new(), + }, + ); + Ok(Arc::new(manager)) +} + +fn create_test_manager_with_failed_apps_startup( + cached_tools: Vec, + reconnect_factory: Arc ManagedClientFuture + Send + Sync>, +) -> McpConnectionSet { + let client: ManagedClientFuture = futures::future::ready(Err(StartupOutcomeError::Failed { + error: "startup failed".to_string(), + is_authentication_required: false, + })) + .boxed() + .shared(); + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("reconnect-test-account"), + Some("reconnect-test-user"), + ); + store_current_tools(&cache_context, cached_tools); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client, + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: Some(cache_context), + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(true)), + startup_reconnect: Some(Arc::new(CodexAppsStartupReconnect::new(reconnect_factory))), + cancel_token: CancellationToken::new(), + }, + ); + manager +} + +fn model_tool_names(tools: &[ToolInfo]) -> HashSet { + tools + .iter() + .map(ToolInfo::canonical_tool_name) + .collect::>() +} + +fn model_tool_name_len(name: &ToolName) -> usize { + name.namespace + .as_deref() + .map_or(0, |namespace| namespace.len() + "__".len()) + + name.name.len() +} + +fn is_code_mode_compatible_tool_name(name: &ToolName) -> bool { + name.namespace + .as_deref() + .into_iter() + .chain(std::iter::once(name.name.as_str())) + .flat_map(str::chars) + .all(|c| c.is_ascii_alphanumeric() || c == '_') +} + +#[test] +fn elicitation_granular_policy_defaults_to_prompting() { + assert!(!elicitation_is_rejected_by_policy( + AskForApproval::OnRequest + )); + assert!(!elicitation_is_rejected_by_policy( + AskForApproval::UnlessTrusted + )); + assert!(elicitation_is_rejected_by_policy(AskForApproval::Granular( + GranularApprovalConfig { + sandbox_approval: true, + rules: true, + skill_approval: true, + request_permissions: true, + mcp_elicitations: false, + } + ))); +} + +#[test] +fn elicitation_granular_policy_respects_never_and_config() { + assert!(elicitation_is_rejected_by_policy(AskForApproval::Never)); + assert!(elicitation_is_rejected_by_policy(AskForApproval::Granular( + GranularApprovalConfig { + sandbox_approval: true, + rules: true, + skill_approval: true, + request_permissions: true, + mcp_elicitations: false, + } + ))); +} + +#[tokio::test] +async fn disabled_permissions_auto_accept_elicitation_with_empty_form_schema() { + let manager = ElicitationRequestManager::new( + test_elicitation_config("server", AskForApproval::Never, PermissionProfile::Disabled), + /*reviewer*/ None, + /*lifecycle*/ None, + ElicitationRequestRouter::default(), + ); + let (tx_event, _rx_event) = async_channel::bounded(1); + let sender = manager.make_sender( + "server".to_string(), + Some(tx_event), + &ClientMcpExtensions::default(), + ); + + let response = sender( + NumberOrString::Number(1), + codex_rmcp_client::Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta: None, + message: "Confirm?".to_string(), + requested_schema: rmcp::model::ElicitationSchema::builder() + .build() + .expect("schema should build"), + }), + ) + .await + .expect("elicitation should auto accept"); + + assert_eq!( + response, + ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(serde_json::json!({})), + meta: None, + } + ); +} + +#[tokio::test] +async fn disabled_permissions_do_not_auto_accept_elicitation_with_requested_fields() { + let manager = ElicitationRequestManager::new( + test_elicitation_config("server", AskForApproval::Never, PermissionProfile::Disabled), + /*reviewer*/ None, + /*lifecycle*/ None, + ElicitationRequestRouter::default(), + ); + let (tx_event, _rx_event) = async_channel::bounded(1); + let sender = manager.make_sender( + "server".to_string(), + Some(tx_event), + &ClientMcpExtensions::default(), + ); + + let response = sender( + NumberOrString::Number(1), + codex_rmcp_client::Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta: None, + message: "What should I say?".to_string(), + requested_schema: + rmcp::model::ElicitationSchema::builder() + .required_property( + "message", + rmcp::model::PrimitiveSchemaDefinition::String( + rmcp::model::StringSchema::new(), + ), + ) + .build() + .expect("schema should build"), + }), + ) + .await + .expect("elicitation should auto decline"); + + assert_eq!( + response, + ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + } + ); +} + +fn full_access_form_input_enabled_router() -> ElicitationRequestRouter { + let router = ElicitationRequestRouter::default(); + router.enable_full_access_form_input(); + router +} + +fn elicitation_meta(value: serde_json::Value) -> Option { + let serde_json::Value::Object(map) = value else { + panic!("elicitation metadata must be an object"); + }; + Some(rmcp::model::RequestMetaObject::from(map)) +} + +fn requested_user_input_schema() -> rmcp::model::ElicitationSchema { + rmcp::model::ElicitationSchema::builder() + .required_property( + "message", + rmcp::model::PrimitiveSchemaDefinition::String(rmcp::model::StringSchema::new()), + ) + .build() + .expect("schema should build") +} + +#[derive(Default)] +struct DecliningElicitationReviewer { + review_count: AtomicUsize, +} + +impl ElicitationReviewer for DecliningElicitationReviewer { + fn review( + &self, + _request: ElicitationReviewRequest, + ) -> BoxFuture<'static, anyhow::Result>> { + self.review_count.fetch_add(1, Ordering::SeqCst); + async { + Ok(Some(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + })) + } + .boxed() + } +} + +async fn assert_elicitation_declined_with_reviewer_calls( + approval_policy: AskForApproval, + server_name: &str, + elicitation: ElicitRequestParams, + expected_reviewer_calls: usize, +) { + let reviewer = Arc::new(DecliningElicitationReviewer::default()); + let manager = ElicitationRequestManager::new( + test_elicitation_config(server_name, approval_policy, PermissionProfile::Disabled), + Some(reviewer.clone()), + /*lifecycle*/ None, + full_access_form_input_enabled_router(), + ); + let (tx_event, rx_event) = async_channel::bounded(1); + let sender = manager.make_sender( + server_name.to_string(), + Some(tx_event), + &ClientMcpExtensions::default(), + ); + + let response = tokio::select! { + biased; + event = rx_event.recv() => { + panic!("elicitation unexpectedly reached the user: {event:?}"); + } + response = sender( + NumberOrString::Number(1), + codex_rmcp_client::Elicitation::Mcp(elicitation), + ) => response.expect("elicitation should be declined"), + }; + + assert_eq!( + response, + ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }, + ); + assert_eq!( + reviewer.review_count.load(Ordering::SeqCst), + expected_reviewer_calls + ); + assert!(rx_event.try_recv().is_err()); +} + +async fn assert_requested_user_input_is_declined( + approval_policy: AskForApproval, + permission_profile: PermissionProfile, + router: ElicitationRequestRouter, +) { + let manager = ElicitationRequestManager::new( + test_elicitation_config("server", approval_policy, permission_profile), + /*reviewer*/ None, + /*lifecycle*/ None, + router, + ); + let (tx_event, rx_event) = async_channel::bounded(1); + let sender = manager.make_sender( + "server".to_string(), + Some(tx_event), + &ClientMcpExtensions::default(), + ); + + let response = tokio::select! { + biased; + event = rx_event.recv() => { + panic!("user-input form unexpectedly reached the user: {event:?}"); + } + response = sender( + NumberOrString::Number(1), + codex_rmcp_client::Elicitation::Mcp( + ElicitRequestParams::FormElicitationParams { + meta: None, + message: "What should I say?".to_string(), + requested_schema: requested_user_input_schema(), + }, + ), + ) => response.expect("restricted user-input request should decline"), + }; + + assert_eq!( + response, + ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }, + ); + assert!(rx_event.try_recv().is_err()); +} + +#[tokio::test] +async fn disabled_permissions_do_not_surface_user_input_when_auto_denied() { + let router = full_access_form_input_enabled_router(); + router.set_auto_deny(/*auto_deny*/ true); + assert_requested_user_input_is_declined( + AskForApproval::Never, + PermissionProfile::Disabled, + router, + ) + .await; +} + +#[tokio::test] +async fn plugin_tool_suggestion_elicitations_are_declined_before_review() { + assert_elicitation_declined_with_reviewer_calls( + AskForApproval::OnRequest, + "server", + ElicitRequestParams::FormElicitationParams { + meta: elicitation_meta(serde_json::json!({ + "codex_approval_kind": "tool_suggestion", + })), + message: "Install this app?".to_string(), + requested_schema: rmcp::model::ElicitationSchema::builder() + .build() + .expect("schema should build"), + }, + /*expected_reviewer_calls*/ 0, + ) + .await; +} + +#[tokio::test] +async fn disabled_permissions_surface_requested_user_input_without_metadata() { + assert_disabled_permissions_surface_requested_user_input(/*meta*/ None).await; +} + +#[tokio::test] +async fn disabled_permissions_surface_requested_user_input_with_non_codex_approval_metadata() { + assert_disabled_permissions_surface_requested_user_input(elicitation_meta(serde_json::json!({ + "origin": "https://example.com", + "persist": "always", + }))) + .await; +} + +async fn assert_disabled_permissions_surface_requested_user_input( + meta: Option, +) { + let router = full_access_form_input_enabled_router(); + let reviewer = Arc::new(DecliningElicitationReviewer::default()); + let manager = ElicitationRequestManager::new( + test_elicitation_config("server", AskForApproval::Never, PermissionProfile::Disabled), + Some(reviewer.clone()), + /*lifecycle*/ None, + router.clone(), + ); + let (tx_event, rx_event) = async_channel::bounded(1); + let sender = manager.make_sender( + "server".to_string(), + Some(tx_event), + &ClientMcpExtensions::default(), + ); + let requested_schema = requested_user_input_schema(); + let mut pending = tokio::spawn(sender( + NumberOrString::Number(1), + codex_rmcp_client::Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta: meta.clone(), + message: "What should I say?".to_string(), + requested_schema: requested_schema.clone(), + }), + )); + let request = tokio::select! { + event = rx_event.recv() => { + let EventMsg::ElicitationRequest(request) = event.expect("user-input event").msg else { + panic!("expected MCP user-input elicitation"); + }; + request + } + response = &mut pending => { + panic!("user input resolved without reaching the user: {response:?}"); + } + }; + + assert_eq!( + request.request, + ElicitationRequest::Form { + meta: meta + .map(serde_json::to_value) + .transpose() + .expect("user-input metadata should serialize"), + message: "What should I say?".to_string(), + requested_schema: serde_json::to_value(requested_schema) + .expect("schema should serialize"), + }, + ); + assert_eq!(request.server_name, "server"); + assert_eq!(reviewer.review_count.load(Ordering::SeqCst), 0); + + let codex_protocol::mcp::RequestId::String(request_id) = request.id else { + panic!("expected Codex-owned string request ID"); + }; + let user_response = ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(serde_json::json!({ "message": "The actual user response." })), + meta: None, + }; + router + .resolve( + "server".to_string(), + NumberOrString::String(request_id.into()), + user_response.clone(), + ) + .await + .expect("actual user response should resolve the elicitation"); + assert_eq!( + pending + .await + .expect("user-input task should complete") + .expect("user input should resolve"), + user_response, + ); +} + +#[tokio::test] +async fn disabled_permissions_decline_requested_user_input_with_approval_metadata() { + assert_elicitation_declined_with_reviewer_calls( + AskForApproval::Never, + "node_repl", + ElicitRequestParams::FormElicitationParams { + meta: elicitation_meta(serde_json::json!({ + "codex_approval_kind": "mcp_tool_call", + "connector_id": "browser-use", + "tool_name": "access_browser_origin", + })), + message: "Allow Browser Use to access this website?".to_string(), + requested_schema: + rmcp::model::ElicitationSchema::builder() + .required_property( + "confirmation", + rmcp::model::PrimitiveSchemaDefinition::String( + rmcp::model::StringSchema::new(), + ), + ) + .build() + .expect("schema should build"), + }, + /*expected_reviewer_calls*/ 0, + ) + .await; +} + +#[tokio::test] +async fn restricted_never_policy_does_not_surface_requested_user_input() { + assert_requested_user_input_is_declined( + AskForApproval::Never, + PermissionProfile::default(), + full_access_form_input_enabled_router(), + ) + .await; +} + +#[tokio::test] +async fn granular_policy_does_not_surface_requested_user_input() { + assert_requested_user_input_is_declined( + AskForApproval::Granular(GranularApprovalConfig { + sandbox_approval: true, + rules: true, + skill_approval: true, + request_permissions: true, + mcp_elicitations: false, + }), + PermissionProfile::Disabled, + full_access_form_input_enabled_router(), + ) + .await; +} + +#[tokio::test] +async fn on_request_approval_forms_remain_with_the_reviewer() { + assert_elicitation_declined_with_reviewer_calls( + AskForApproval::OnRequest, + "server", + ElicitRequestParams::FormElicitationParams { + meta: elicitation_meta(serde_json::json!({ + "codex_request_type": "approval_request", + "codex_approval_kind": "mcp_tool_call", + "tool_name": "test_tool", + })), + message: "Approve this action?".to_string(), + requested_schema: rmcp::model::ElicitationSchema::builder() + .build() + .expect("schema should build"), + }, + /*expected_reviewer_calls*/ 1, + ) + .await; +} + +#[tokio::test] +async fn disabled_permissions_decline_user_input_without_an_event_channel() { + let manager = ElicitationRequestManager::new( + test_elicitation_config("server", AskForApproval::Never, PermissionProfile::Disabled), + /*reviewer*/ None, + /*lifecycle*/ None, + full_access_form_input_enabled_router(), + ); + let sender = manager.make_sender( + "server".to_string(), + /*tx_event*/ None, + &ClientMcpExtensions::default(), + ); + + let response = sender( + NumberOrString::Number(1), + codex_rmcp_client::Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta: None, + message: "What should I say?".to_string(), + requested_schema: requested_user_input_schema(), + }), + ) + .await + .expect("headless user-input request should decline"); + + assert_eq!( + response, + ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }, + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn concurrent_authority_updates_never_auto_approve_mixed_policy() { + let manager = ElicitationRequestManager::new( + test_elicitation_config( + "server", + AskForApproval::Never, + PermissionProfile::default(), + ), + /*reviewer*/ None, + /*lifecycle*/ None, + ElicitationRequestRouter::default(), + ); + let updating_manager = manager.clone(); + let updater = tokio::spawn(async move { + for _ in 0..1_000 { + assert!(updating_manager.update( + test_elicitation_config( + "server", + AskForApproval::OnRequest, + PermissionProfile::Disabled + ), + /*reviewer*/ None, + /*lifecycle*/ None, + )); + assert!(updating_manager.update( + test_elicitation_config( + "server", + AskForApproval::Never, + PermissionProfile::default() + ), + /*reviewer*/ None, + /*lifecycle*/ None, + )); + } + }); + let sender = manager.make_sender( + "server".to_string(), + /*tx_event*/ None, + &ClientMcpExtensions::default(), + ); + let elicitation = + codex_rmcp_client::Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta: None, + message: "Confirm?".to_string(), + requested_schema: rmcp::model::ElicitationSchema::builder() + .build() + .expect("schema should build"), + }); + + for _ in 0..1_000 { + let response = sender(NumberOrString::Number(1), elicitation.clone()) + .await + .expect("elicitation should resolve"); + assert_eq!( + response, + ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + } + ); + } + + updater.await.expect("authority updates should finish"); +} + +#[tokio::test] +async fn shared_elicitation_router_targets_the_exact_pending_request() { + struct Registration(Arc); + + impl Drop for Registration { + fn drop(&mut self) { + self.0.fetch_sub(1, std::sync::atomic::Ordering::SeqCst); + } + } + + let router = ElicitationRequestRouter::default(); + let outstanding = Arc::new(AtomicUsize::new(0)); + let lifecycle = ElicitationLifecycle::new({ + let outstanding = outstanding.clone(); + move || { + outstanding.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Registration(outstanding.clone()) + } + }); + let manager_a = ElicitationRequestManager::new( + test_elicitation_config( + "server", + AskForApproval::OnRequest, + PermissionProfile::default(), + ), + /*reviewer*/ None, + Some(lifecycle.clone()), + router.clone(), + ); + let manager_b = ElicitationRequestManager::new( + test_elicitation_config( + "server", + AskForApproval::OnRequest, + PermissionProfile::default(), + ), + /*reviewer*/ None, + Some(lifecycle), + router.clone(), + ); + let (tx_event, rx_event) = async_channel::bounded(2); + let sender_a = manager_a.make_sender( + "server".to_string(), + Some(tx_event.clone()), + &ClientMcpExtensions::default(), + ); + let sender_b = manager_b.make_sender( + "server".to_string(), + Some(tx_event), + &ClientMcpExtensions::default(), + ); + let elicitation = + codex_rmcp_client::Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta: None, + message: "Which runtime?".to_string(), + requested_schema: + rmcp::model::ElicitationSchema::builder() + .required_property( + "runtime", + rmcp::model::PrimitiveSchemaDefinition::String( + rmcp::model::StringSchema::new(), + ), + ) + .build() + .expect("schema should build"), + }); + + let pending_a = tokio::spawn(sender_a(NumberOrString::Number(1), elicitation.clone())); + let EventMsg::ElicitationRequest(request_a) = rx_event.recv().await.expect("request A").msg + else { + panic!("expected elicitation request"); + }; + let pending_b = tokio::spawn(sender_b(NumberOrString::Number(1), elicitation)); + let EventMsg::ElicitationRequest(request_b) = rx_event.recv().await.expect("request B").msg + else { + panic!("expected elicitation request"); + }; + assert_eq!(outstanding.load(std::sync::atomic::Ordering::SeqCst), 2); + let ( + codex_protocol::mcp::RequestId::String(request_a_id), + codex_protocol::mcp::RequestId::String(request_b_id), + ) = (request_a.id, request_b.id) + else { + panic!("expected Codex-owned string request IDs"); + }; + assert_ne!(request_a_id, request_b_id); + + let response_a = ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(serde_json::json!({"runtime": "a"})), + meta: None, + }; + router + .resolve( + "server".to_string(), + NumberOrString::String(request_a_id.into()), + response_a.clone(), + ) + .await + .expect("runtime B should route a response to runtime A"); + let response_b = ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(serde_json::json!({"runtime": "b"})), + meta: None, + }; + router + .resolve( + "server".to_string(), + NumberOrString::String(request_b_id.into()), + response_b.clone(), + ) + .await + .expect("runtime A should route a response to runtime B"); + + assert_eq!( + pending_a + .await + .expect("request A task") + .expect("request A response"), + response_a + ); + assert_eq!( + pending_b + .await + .expect("request B task") + .expect("request B response"), + response_b + ); + assert_eq!(outstanding.load(std::sync::atomic::Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn cancelled_elicitation_is_removed_without_affecting_other_pending_requests() { + let router = ElicitationRequestRouter::default(); + let manager = ElicitationRequestManager::new( + test_elicitation_config( + "server", + AskForApproval::OnRequest, + PermissionProfile::default(), + ), + /*reviewer*/ None, + /*lifecycle*/ None, + router.clone(), + ); + let (tx_event, rx_event) = async_channel::bounded(2); + let sender = manager.make_sender( + "server".to_string(), + Some(tx_event), + &ClientMcpExtensions::default(), + ); + let elicitation = + codex_rmcp_client::Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta: None, + message: "Confirm?".to_string(), + requested_schema: + rmcp::model::ElicitationSchema::builder() + .required_property( + "answer", + rmcp::model::PrimitiveSchemaDefinition::String( + rmcp::model::StringSchema::new(), + ), + ) + .build() + .expect("schema should build"), + }); + + let cancelled = tokio::spawn(sender(NumberOrString::Number(1), elicitation.clone())); + let EventMsg::ElicitationRequest(cancelled_request) = + rx_event.recv().await.expect("cancelled request event").msg + else { + panic!("expected elicitation request"); + }; + let pending = tokio::spawn(sender(NumberOrString::Number(2), elicitation)); + let EventMsg::ElicitationRequest(pending_request) = + rx_event.recv().await.expect("pending request event").msg + else { + panic!("expected elicitation request"); + }; + let ( + codex_protocol::mcp::RequestId::String(cancelled_id), + codex_protocol::mcp::RequestId::String(pending_id), + ) = (cancelled_request.id, pending_request.id) + else { + panic!("expected Codex-owned string request IDs"); + }; + + cancelled.abort(); + assert!( + cancelled + .await + .expect_err("cancelled request should be aborted") + .is_cancelled() + ); + + let response = ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(serde_json::json!({"answer": "yes"})), + meta: None, + }; + let error = router + .resolve( + "server".to_string(), + NumberOrString::String(cancelled_id.into()), + response.clone(), + ) + .await + .expect_err("cancelled request should be removed immediately"); + assert_eq!(error.to_string(), "elicitation request not found"); + + router + .resolve( + "server".to_string(), + NumberOrString::String(pending_id.into()), + response.clone(), + ) + .await + .expect("another pending request should remain routable"); + assert_eq!( + pending + .await + .expect("pending request task") + .expect("pending request response"), + response + ); +} + +#[test] +fn test_normalize_tools_short_non_duplicated_names() { + let tools = vec![ + create_test_tool("server1", "tool1"), + create_test_tool("server1", "tool2"), + ]; + + let model_tools = + normalize_tools_for_model_with_prefix(tools, /*prefix_mcp_tool_names*/ true, &[]); + + assert_eq!( + model_tool_names(&model_tools), + HashSet::from([ + ToolName::namespaced("mcp__server1", "tool1"), + ToolName::namespaced("mcp__server1", "tool2") + ]) + ); +} + +#[test] +fn test_normalize_tools_omits_prefix_only_for_selected_servers() { + let tools = vec![ + create_test_tool("history", "search"), + create_test_tool("notes", "read"), + create_test_tool("calendar", "list"), + ]; + + let model_tools = normalize_tools_for_model_with_prefix( + tools, + /*prefix_mcp_tool_names*/ true, + &["history".to_string(), "notes".to_string()], + ); + + assert_eq!( + model_tool_names(&model_tools), + HashSet::from([ + ToolName::namespaced("history", "search"), + ToolName::namespaced("notes", "read"), + ToolName::namespaced("mcp__calendar", "list"), + ]) + ); +} + +#[test] +fn test_normalize_tools_selects_raw_server_name() { + let mut tool = create_test_tool("codex_apps", "search"); + tool.callable_namespace = "codex_apps__calendar".to_string(); + + let model_tools = normalize_tools_for_model_with_prefix( + vec![tool], + /*prefix_mcp_tool_names*/ true, + &["codex_apps".to_string()], + ); + + assert_eq!( + model_tool_names(&model_tools), + HashSet::from([ToolName::namespaced("codex_apps__calendar", "search")]) + ); +} + +#[test] +fn test_normalize_tools_global_feature_omits_prefix_for_every_server() { + let tools = vec![ + create_test_tool("history", "search"), + create_test_tool("calendar", "list"), + ]; + + let model_tools = normalize_tools_for_model_with_prefix( + tools, + /*prefix_mcp_tool_names*/ false, + &["history".to_string()], + ); + + assert_eq!( + model_tool_names(&model_tools), + HashSet::from([ + ToolName::namespaced("history", "search"), + ToolName::namespaced("calendar", "list"), + ]) + ); +} + +#[test] +fn test_normalize_tools_duplicated_names_skipped() { + let tools = vec![ + create_test_tool("server1", "duplicate_tool"), + create_test_tool("server1", "duplicate_tool"), + ]; + + let model_tools = + normalize_tools_for_model_with_prefix(tools, /*prefix_mcp_tool_names*/ true, &[]); + + // Only the first tool should remain, the second is skipped + assert_eq!( + model_tool_names(&model_tools), + HashSet::from([ToolName::namespaced("mcp__server1", "duplicate_tool")]) + ); +} + +#[test] +fn test_normalize_tools_respects_responses_api_name_length_boundaries() { + let namespace = "mcp__codex_apps"; + let namespace_len = namespace.len() + "__".len(); + + for total_len in [128, 129] { + let tool_name = "a".repeat(total_len - namespace_len); + let model_tools = normalize_tools_for_model_with_prefix( + vec![create_test_tool("codex_apps", &tool_name)], + /*prefix_mcp_tool_names*/ true, + &[], + ); + let model_name = model_tools[0].canonical_tool_name(); + + assert_eq!(model_tool_name_len(&model_name), 128); + if total_len == 128 { + assert_eq!(model_name, ToolName::namespaced(namespace, tool_name)); + } else { + assert_ne!(model_name.name, tool_name); + } + } +} + +#[test] +fn test_normalize_tools_long_names_same_server() { + let server_name = "my_server"; + let first_name = "a".repeat(128); + let second_name = "b".repeat(128); + + let tools = vec![ + create_test_tool(server_name, &first_name), + create_test_tool(server_name, &second_name), + ]; + + let model_tools = + normalize_tools_for_model_with_prefix(tools, /*prefix_mcp_tool_names*/ true, &[]); + + assert_eq!(model_tools.len(), 2); + + let names = model_tool_names(&model_tools); + + assert!(names.iter().all(|name| model_tool_name_len(name) == 128)); + assert!( + names + .iter() + .all(|name| name.namespace.as_deref() == Some("mcp__my_server")) + ); + assert!( + names.iter().all(is_code_mode_compatible_tool_name), + "model-visible names must be code-mode compatible: {names:?}" + ); +} + +#[test] +fn test_normalize_tools_sanitizes_invalid_characters() { + let tools = vec![create_test_tool("server.one", "tool.two-three")]; + + let model_tools = + normalize_tools_for_model_with_prefix(tools, /*prefix_mcp_tool_names*/ true, &[]); + + assert_eq!(model_tools.len(), 1); + let tool = model_tools.into_iter().next().expect("one tool"); + let model_name = tool.canonical_tool_name(); + assert_eq!( + model_name, + ToolName::namespaced("mcp__server_one", "tool_two_three") + ); + assert_eq!( + ToolName::namespaced(tool.callable_namespace.clone(), tool.callable_name.clone()), + model_name + ); + // The callable parts are sanitized for model-visible tool calls, but the raw + // MCP name is preserved for the actual MCP call. + assert_eq!(tool.server_name, "server.one"); + assert_eq!(tool.callable_namespace, "mcp__server_one"); + assert_eq!(tool.callable_name, "tool_two_three"); + assert_eq!(tool.tool.name, "tool.two-three"); + + assert!( + is_code_mode_compatible_tool_name(&model_name), + "model-visible name must be code-mode compatible: {model_name:?}" + ); +} + +#[test] +fn test_normalize_tools_keeps_hyphenated_mcp_tools_callable() { + let tools = vec![create_test_tool("music-studio", "get-strudel-guide")]; + + let model_tools = + normalize_tools_for_model_with_prefix(tools, /*prefix_mcp_tool_names*/ true, &[]); + + assert_eq!(model_tools.len(), 1); + let tool = model_tools.into_iter().next().expect("one tool"); + assert_eq!( + tool.canonical_tool_name(), + ToolName::namespaced("mcp__music_studio", "get_strudel_guide") + ); + assert_eq!(tool.callable_namespace, "mcp__music_studio"); + assert_eq!(tool.callable_name, "get_strudel_guide"); + assert_eq!(tool.tool.name, "get-strudel-guide"); +} + +#[test] +fn test_normalize_tools_disambiguates_sanitized_namespace_collisions() { + let tools = vec![ + create_test_tool("basic-server", "lookup"), + create_test_tool("basic_server", "query"), + create_test_tool("npm:@scope/package.name", "lookup"), + create_test_tool("npm__scope_package_name", "lookup"), + ]; + + let model_tools = + normalize_tools_for_model_with_prefix(tools, /*prefix_mcp_tool_names*/ true, &[]); + + assert_eq!(model_tools.len(), 4); + let mut namespaces = model_tools + .iter() + .map(|tool| tool.callable_namespace.as_str()) + .collect::>(); + namespaces.sort(); + namespaces.dedup(); + assert_eq!(namespaces.len(), 4); + + let raw_servers = model_tools + .iter() + .map(|tool| tool.server_name.as_str()) + .collect::>(); + assert_eq!( + raw_servers, + HashSet::from([ + "basic-server", + "basic_server", + "npm:@scope/package.name", + "npm__scope_package_name", + ]) + ); + let model_names = model_tool_names(&model_tools); + assert!( + model_names.iter().all(is_code_mode_compatible_tool_name), + "model-visible names must be code-mode compatible: {model_names:?}" + ); +} + +#[test] +fn test_normalize_tools_disambiguates_sanitized_tool_name_collisions() { + let tools = vec![ + create_test_tool("server", "tool-name"), + create_test_tool("server", "tool_name"), + ]; + + let model_tools = + normalize_tools_for_model_with_prefix(tools, /*prefix_mcp_tool_names*/ true, &[]); + + assert_eq!(model_tools.len(), 2); + let raw_tool_names = model_tools + .iter() + .map(|tool| tool.tool.name.to_string()) + .collect::>(); + assert_eq!( + raw_tool_names, + HashSet::from(["tool-name".to_string(), "tool_name".to_string()]) + ); + let callable_tool_names = model_tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(); + assert_eq!(callable_tool_names.len(), 2); +} + +#[test] +fn tool_filter_allows_by_default() { + let filter = ToolFilter::default(); + + assert!(filter.allows("any")); +} + +#[test] +fn tool_filter_applies_enabled_list() { + let filter = ToolFilter { + enabled: Some(HashSet::from(["allowed".to_string()])), + disabled: HashSet::new(), + }; + + assert!(filter.allows("allowed")); + assert!(!filter.allows("denied")); +} + +#[test] +fn tool_filter_applies_disabled_list() { + let filter = ToolFilter { + enabled: None, + disabled: HashSet::from(["blocked".to_string()]), + }; + + assert!(!filter.allows("blocked")); + assert!(filter.allows("open")); +} + +#[test] +fn tool_filter_applies_enabled_then_disabled() { + let filter = ToolFilter { + enabled: Some(HashSet::from(["keep".to_string(), "remove".to_string()])), + disabled: HashSet::from(["remove".to_string()]), + }; + + assert!(filter.allows("keep")); + assert!(!filter.allows("remove")); + assert!(!filter.allows("unknown")); +} + +#[test] +fn filter_tools_applies_per_server_filters() { + let server1_tools = vec![ + create_test_tool("server1", "tool_a"), + create_test_tool("server1", "tool_b"), + ]; + let server2_tools = vec![create_test_tool("server2", "tool_a")]; + let server1_filter = ToolFilter { + enabled: Some(HashSet::from(["tool_a".to_string(), "tool_b".to_string()])), + disabled: HashSet::from(["tool_b".to_string()]), + }; + let server2_filter = ToolFilter { + enabled: None, + disabled: HashSet::from(["tool_a".to_string()]), + }; + + let filtered: Vec<_> = filter_tools(server1_tools, &server1_filter) + .into_iter() + .chain(filter_tools(server2_tools, &server2_filter)) + .collect(); + + assert_eq!(filtered.len(), 1); + assert_eq!(filtered[0].server_name, "server1"); + assert_eq!(filtered[0].callable_name, "tool_a"); +} + +#[test] +fn codex_apps_env_bearer_token_bypasses_shared_tools_cache() { + assert!(!should_share_codex_apps_tools_cache( + CODEX_APPS_MCP_SERVER_NAME, + /*uses_env_bearer_token*/ true, + )); +} + +#[tokio::test] +async fn hosted_apps_protocol_mode_is_independent_of_generic_mode() -> anyhow::Result<()> { + let codex_home = tempdir()?; + let server_config: McpServerConfig = + serde_json::from_value(serde_json::json!({ "url": "http://127.0.0.1:1/ps/mcp" }))?; + + for (generic_mode, hosted_mode) in [ + ( + crate::McpProtocolMode::Legacy, + crate::McpProtocolMode::V20260728, + ), + ( + crate::McpProtocolMode::V20260728, + crate::McpProtocolMode::Legacy, + ), + ] { + let mut config = crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()); + config.protocol_mode = generic_mode; + config.host_owned_apps_protocol_mode = hosted_mode; + let mut catalog = crate::ResolvedMcpCatalog::builder(); + catalog.register(crate::McpServerRegistration::from_hosted_apps( + "test-host", + /*contribution_order*/ 0, + server_config.clone(), + )); + catalog.register(crate::McpServerRegistration::from_config( + "third_party".to_string(), + server_config.clone(), + )); + config.mcp_server_catalog = catalog.build(); + + let startup_cancellation_token = CancellationToken::new(); + startup_cancellation_token.cancel(); + let manager = McpConnectionSet::new( + /*previous*/ None, + McpPublicationGate::already_published(), + McpRuntimeInput { + startup_policy: McpStartupPolicy::Eager, + config: Arc::new(config), + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + mcp_servers: HashMap::from([ + ( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + EffectiveMcpServer::configured(server_config.clone()), + ), + ( + "third_party".to_string(), + EffectiveMcpServer::configured(server_config.clone()), + ), + ]), + submit_id: "protocol-mode-scope".to_string(), + tx_event: None, + startup_cancellation_token, + runtime_context: McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + codex_home.path().to_path_buf(), + ), + codex_apps_tools_cache: ConnectorRuntimeManager::default(), + tool_catalog_cache: McpToolCatalogCache::default(), + codex_apps_tools_cache_key: ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ), + client_mcp_extensions: ClientMcpExtensions::default(), + auth: None, + auth_manager: None, + elicitation_reviewer: None, + elicitation_lifecycle: None, + }, + ElicitationRequestRouter::default(), + ) + .await; + + assert_eq!( + manager.servers[CODEX_APPS_MCP_SERVER_NAME].protocol_mode, + hosted_mode + ); + assert_eq!(manager.servers["third_party"].protocol_mode, generic_mode); + assert_eq!( + manager + .event_stream_connection + .as_ref() + .expect("hosted Apps event stream") + .protocol_mode, + hosted_mode + ); + } + + Ok(()) +} + +#[tokio::test] +async fn codex_apps_extension_does_not_share_host_owned_tools_cache() -> anyhow::Result<()> { + let codex_home = tempdir()?; + let cache_key = ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ); + let codex_apps_tools_cache = ConnectorRuntimeManager::::default(); + let cache_context = + codex_apps_tools_cache.context(codex_home.path().to_path_buf(), cache_key.clone()); + store_current_tools( + &cache_context, + vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "calendar_create_event", + )], + ); + + let server_config: McpServerConfig = + serde_json::from_value(serde_json::json!({ "url": "http://127.0.0.1:1" }))?; + for (hosted_mode, extension_mode, expected_mode) in [ + ( + crate::McpProtocolMode::V20260728, + None, + crate::McpProtocolMode::Legacy, + ), + ( + crate::McpProtocolMode::Legacy, + Some(crate::McpProtocolMode::V20260728), + crate::McpProtocolMode::V20260728, + ), + ] { + let mut config = crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()); + config.host_owned_apps_protocol_mode = hosted_mode; + let mut registration = crate::McpServerRegistration::from_extension( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + "test-extension", + /*contribution_order*/ 0, + server_config.clone(), + ); + if let Some(mode) = extension_mode { + registration = registration.with_protocol_mode(mode); + } + let mut catalog = crate::ResolvedMcpCatalog::builder(); + catalog.register(registration); + config.mcp_server_catalog = catalog.build(); + + let startup_cancellation_token = CancellationToken::new(); + startup_cancellation_token.cancel(); + let manager = McpConnectionSet::new( + /*previous*/ None, + McpPublicationGate::already_published(), + McpRuntimeInput { + startup_policy: McpStartupPolicy::Eager, + config: Arc::new(config), + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + mcp_servers: HashMap::from([( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + EffectiveMcpServer::configured(server_config.clone()), + )]), + submit_id: "cache-ownership-test".to_string(), + tx_event: None, + startup_cancellation_token, + runtime_context: McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + codex_home.path().to_path_buf(), + ), + codex_apps_tools_cache: codex_apps_tools_cache.clone(), + tool_catalog_cache: McpToolCatalogCache::default(), + codex_apps_tools_cache_key: cache_key.clone(), + client_mcp_extensions: ClientMcpExtensions::default(), + auth: None, + auth_manager: None, + elicitation_reviewer: None, + elicitation_lifecycle: None, + }, + ElicitationRequestRouter::default(), + ) + .await; + + let client = manager.test_client(CODEX_APPS_MCP_SERVER_NAME); + assert_eq!( + manager.servers[CODEX_APPS_MCP_SERVER_NAME].protocol_mode, expected_mode, + "an ordinary extension must use its own mode, not the hosted protocol default" + ); + assert!( + client.codex_apps_tools_cache_context.is_none(), + "an extension must not receive the host-owned Apps cache" + ); + assert!( + !client.has_cached_tools(), + "an extension must not expose cached host-owned Apps tools" + ); + } + + Ok(()) +} + +#[tokio::test] +async fn list_all_tools_uses_shared_codex_apps_cache_while_client_is_pending() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + store_current_tools( + &cache_context, + vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "calendar_create_event", + )], + ); + let pending_client = futures::future::pending::>() + .boxed() + .shared(); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: pending_client, + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: Some(cache_context), + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(false)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + + let tools = manager.list_all_tools().await; + let tool = tools + .iter() + .find(|tool| { + tool.canonical_tool_name() + == ToolName::namespaced("mcp__codex_apps", "calendar_create_event") + }) + .expect("tool from shared cache"); + assert_eq!(tool.server_name, CODEX_APPS_MCP_SERVER_NAME); + assert_eq!(tool.callable_name, "calendar_create_event"); +} + +#[tokio::test] +async fn capture_binding_uses_the_ready_clients_own_tools() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + store_current_tools( + &cache_context, + vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "shared_cached_tool", + )], + ); + let mut ready_client = create_test_managed_client(vec![ + create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "client_local_tool"), + create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "client_local_blocked"), + ]) + .await; + let tool_filter = ToolFilter { + enabled: None, + disabled: HashSet::from(["client_local_blocked".to_string()]), + }; + ready_client.codex_apps_tools_cache_context = Some(cache_context.clone()); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: futures::future::ready(Ok(ready_client)).boxed().shared(), + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: Some(cache_context), + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(true)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + manager + .servers + .get_mut(CODEX_APPS_MCP_SERVER_NAME) + .expect("test server exists") + .tool_filter = tool_filter; + manager.set_test_server_metadata( + CODEX_APPS_MCP_SERVER_NAME, + McpServerMetadata { + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + pollutes_memory: false, + origin: None, + supports_parallel_tool_calls: false, + default_tools_approval_mode: None, + tool_approval_modes: HashMap::new(), + }, + ); + let manager = Arc::new(manager); + + assert_eq!( + manager + .list_all_tools() + .await + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["shared_cached_tool"] + ); + let step = capture_binding(&manager).await; + assert_eq!( + step.tools() + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["client_local_tool"] + ); + assert!( + step.prepare_call(CODEX_APPS_MCP_SERVER_NAME, "client_local_tool") + .is_some() + ); + assert!( + step.prepare_call(CODEX_APPS_MCP_SERVER_NAME, "shared_cached_tool") + .is_none() + ); + assert!( + step.prepare_call(CODEX_APPS_MCP_SERVER_NAME, "client_local_blocked") + .is_none() + ); +} + +#[tokio::test] +async fn hard_refresh_keeps_client_catalog_local_when_shared_cache_loses_race() -> anyhow::Result<()> +{ + let codex_home = tempdir()?; + let shared_cache = ConnectorRuntimeManager::::default(); + let cache_key = ConnectorRuntimeContextKey::personal( + Some("shared-account".to_string()), + Some("shared-user".to_string()), + ); + let cache_context_a = shared_cache.context(codex_home.path().to_path_buf(), cache_key.clone()); + let cache_context_b = shared_cache.context(codex_home.path().to_path_buf(), cache_key); + let list_started = Arc::new(Notify::new()); + let release_list = Arc::new(Notify::new()); + let manager_a = create_test_manager_with_ready_apps_client( + cache_context_a.clone(), + "a_only", + Some(Arc::clone(&list_started)), + Some(Arc::clone(&release_list)), + ) + .await?; + let mut manager_b = create_test_manager_with_ready_apps_client( + cache_context_b, + "b_only", + /*list_started*/ None, + /*release_list*/ None, + ) + .await?; + + let manager_a_for_refresh = Arc::clone(&manager_a); + let refresh_a = tokio::spawn(async move { + manager_a_for_refresh + .refresh_codex_apps_tools_for_discovery() + .await + }); + list_started.notified().await; + let tools_b = manager_b.refresh_codex_apps_tools_for_discovery().await?; + release_list.notify_one(); + let tools_a = refresh_a.await??; + + assert_eq!( + tools_b + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["b_only"] + ); + assert_eq!( + tools_a + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["b_only"] + ); + assert_eq!( + cache_context_a + .current_tools() + .expect("shared cache tools") + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["b_only"] + ); + assert_eq!( + capture_binding(&manager_a) + .await + .tools() + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["a_only"] + ); + assert_eq!( + capture_binding(&manager_b) + .await + .tools() + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["b_only"] + ); + + let mut config = crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()); + config.server_permission_profiles.insert( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + PermissionProfile::default(), + ); + let manager_a_for_refresh = Arc::clone(&manager_a); + let config_for_refresh = config.clone(); + let refresh_a = tokio::spawn(async move { + manager_a_for_refresh + .refresh_codex_apps_client_catalog(&config_for_refresh) + .await + }); + list_started.notified().await; + manager_b.refresh_codex_apps_tools_for_discovery().await?; + release_list.notify_one(); + let snapshot_a = refresh_a.await??; + assert_eq!( + snapshot_a + .tools + .iter() + .map(|tool| tool.tool.name.as_ref()) + .collect::>(), + vec!["a_only"] + ); + assert_eq!( + snapshot_a.model_visible_tool_names, + HashSet::from(["a_only".to_string()]) + ); + assert_eq!( + cache_context_a + .current_tools() + .expect("shared cache tools") + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["b_only"] + ); + Arc::get_mut(&mut manager_b) + .expect("unshared manager") + .servers + .get_mut(CODEX_APPS_MCP_SERVER_NAME) + .expect("Apps server") + .tool_filter + .disabled + .insert("b_only".to_string()); + let snapshot_b = manager_b.refresh_codex_apps_client_catalog(&config).await?; + assert_eq!( + ( + snapshot_b + .tools + .iter() + .map(|tool| tool.tool.name.as_ref()) + .collect::>(), + snapshot_b.model_visible_tool_names, + ), + (vec!["b_only"], HashSet::new()) + ); + Ok(()) +} + +#[tokio::test(start_paused = true)] +async fn tool_catalog_cache_sanitizes_tools_and_tracks_environment_generation() { + let cache = McpToolCatalogCache::default(); + let environment_manager = Arc::new(environment_manager_without_environments()); + let replace_environment = |url: &str| { + environment_manager + .upsert_environment( + "remote".to_string(), + url.to_string(), + /*connect_timeout*/ None, + ) + .expect("replace environment"); + }; + replace_environment("ws://127.0.0.1:1"); + let runtime_context = + McpRuntimeContext::new(Arc::clone(&environment_manager), PathBuf::from("/tmp")); + let config: McpServerConfig = serde_json::from_value(serde_json::json!({ + "command": "docs-mcp", + "environment_id": "remote" + })) + .expect("MCP config"); + let resolve_environment = || { + runtime_context + .resolve_server_environment("docs", &config) + .expect("resolve environment") + .expect("remote environment") + }; + let cache_context = |environment: &Arc| { + cache + .context( + "docs", + &config, + &runtime_context, + Some(environment), + ( + &ElicitationCapability::default(), + &ClientMcpExtensions::default(), + ), + /*connection_identity*/ None, + ) + .expect("cache context") + }; + let first_environment = resolve_environment(); + let first_environment_weak = Arc::downgrade(&first_environment); + let first_context = cache_context(&first_environment); + first_context.publish_if_newest(first_context.begin_fetch(), &[]); + assert!(!first_context.has_tools()); + + let mut tool = create_test_tool("docs", "search"); + tool.tool.annotations = Some(rmcp::model::ToolAnnotations::new().read_only(true)); + first_context.publish_if_newest(first_context.begin_fetch(), &[tool]); + assert_eq!( + first_context.current_tools().expect("cached tools")[0] + .tool + .annotations, + None + ); + + drop(first_environment); + replace_environment("ws://127.0.0.1:2"); + assert!(first_environment_weak.upgrade().is_none()); + let replacement_environment = resolve_environment(); + assert!(!cache_context(&replacement_environment).has_tools()); + + let older = first_context.begin_fetch(); + let newer = first_context.begin_fetch(); + first_context.publish_if_newest(newer, &[create_test_tool("docs", "new")]); + first_context.publish_if_newest(older, &[create_test_tool("docs", "old")]); + assert_eq!( + first_context.current_tools().expect("cached tools")[0].callable_name, + "new" + ); + + tokio::time::advance(Duration::from_secs(30 * 60 + 1)).await; + assert!(!first_context.has_tools()); +} + +#[test] +fn tool_catalog_cache_bypasses_remote_sourced_environment_variables() { + let cache = McpToolCatalogCache::default(); + let runtime_context = McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + PathBuf::from("/tmp"), + ); + let config: McpServerConfig = serde_json::from_value(serde_json::json!({ + "command": "docs-mcp", + "env_vars": [McpServerEnvVar::Config { + name: "DOCS_TOKEN".to_string(), + source: Some("remote".to_string()), + }], + })) + .expect("MCP config"); + + assert!( + cache + .context( + "docs", + &config, + &runtime_context, + /*resolved_environment*/ None, + ( + &ElicitationCapability::default(), + &ClientMcpExtensions::default() + ), + /*connection_identity*/ None, + ) + .is_none() + ); +} + +#[test] +fn tool_catalog_cache_bypasses_http_headers_helpers() { + let cache = McpToolCatalogCache::default(); + let runtime_context = reusable_server_runtime_context(); + let mut config = reusable_server_config("https://example.com/mcp"); + let identity = reusable_server_identity("docs", &config, &runtime_context); + let context = |config: &McpServerConfig, identity: &McpServerConnectionIdentity| { + cache.context( + "docs", + config, + &runtime_context, + /*resolved_environment*/ None, + ( + &ElicitationCapability::default(), + &ClientMcpExtensions::default(), + ), + Some(( + identity, + crate::McpProtocolMode::Legacy, + /*agent_plugin*/ false, + )), + ) + }; + assert!(context(&config, &identity).is_some()); + + let McpServerTransportConfig::StreamableHttp { + http_headers_helper, + .. + } = &mut config.transport + else { + unreachable!("expected HTTP transport"); + }; + *http_headers_helper = Some("auth-cli headers".to_string()); + let identity = reusable_server_identity("docs", &config, &runtime_context); + assert!(context(&config, &identity).is_none()); +} + +#[tokio::test] +async fn list_available_server_infos_uses_cache_while_client_is_pending() { + let pending_client = futures::future::pending::>() + .boxed() + .shared(); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + let server_info = create_test_server_info("Codex Apps"); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: pending_client, + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: Some(server_info.clone()), + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(false)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + + let timeout_result = tokio::time::timeout( + Duration::from_millis(10), + manager.list_available_server_infos(), + ) + .await; + let server_infos = timeout_result.expect("server info lookup should not block on startup"); + assert_eq!( + server_infos.get(CODEX_APPS_MCP_SERVER_NAME), + Some(&server_info) + ); +} + +#[tokio::test] +async fn list_all_tools_accepts_canonical_namespaced_tool_names() { + let managed_client = + create_ready_async_managed_client(vec![create_test_tool("rmcp", "echo")]).await; + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ false, + ); + manager.insert_test_client("rmcp", managed_client); + + let tools = manager.list_all_tools().await; + let tool = tools + .iter() + .find(|tool| tool.canonical_tool_name() == ToolName::namespaced("rmcp", "echo")) + .expect("split MCP tool namespace and name should resolve"); + + let expected = ("rmcp", "rmcp", "echo", "echo"); + assert_eq!( + ( + tool.server_name.as_str(), + tool.callable_namespace.as_str(), + tool.callable_name.as_str(), + tool.tool.name.as_ref(), + ), + expected + ); +} + +#[tokio::test] +async fn capture_binding_exposes_cached_tools_before_startup() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let mut cached_tool = create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "shared_cached_tool"); + cached_tool.tool.annotations = Some( + rmcp::model::ToolAnnotations::new() + .read_only(true) + .destructive(false) + .open_world(false), + ); + store_current_tools(&cache_context, vec![cached_tool]); + let startup_complete = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let startup_complete_for_client = Arc::clone(&startup_complete); + let (startup_started, wait_for_startup) = tokio::sync::oneshot::channel(); + let (release_startup, startup_released) = tokio::sync::oneshot::channel(); + let pending_client = async move { + startup_started.send(()).expect("signal client startup"); + startup_released.await.expect("release client startup"); + startup_complete_for_client.store(true, std::sync::atomic::Ordering::Release); + Ok(create_test_managed_client(vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "client_local_tool", + )]) + .await) + } + .boxed() + .shared(); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: pending_client, + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: Some(cache_context), + tool_catalog_cache_context: None, + startup_complete, + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + manager.set_test_server_metadata( + CODEX_APPS_MCP_SERVER_NAME, + McpServerMetadata { + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + pollutes_memory: false, + origin: None, + supports_parallel_tool_calls: false, + default_tools_approval_mode: None, + tool_approval_modes: HashMap::new(), + }, + ); + let manager = Arc::new(manager); + let cached_binding = capture_binding(&manager).await; + assert_eq!( + cached_binding + .tools() + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["shared_cached_tool"] + ); + assert_eq!( + cached_binding.tools()[0].tool.annotations, + Some( + rmcp::model::ToolAnnotations::new() + .destructive(false) + .open_world(false) + ) + ); + assert!( + cached_binding + .prepare_call(CODEX_APPS_MCP_SERVER_NAME, "shared_cached_tool") + .is_none() + ); + + let manager_for_startup = Arc::clone(&manager); + let startup = tokio::spawn(async move { + manager_for_startup + .wait_for_server_startup(CODEX_APPS_MCP_SERVER_NAME) + .await + }); + + wait_for_startup.await.expect("client startup should begin"); + release_startup.send(()).expect("release client startup"); + assert!(startup.await.expect("startup task")); + + let step = capture_binding(&manager).await; + assert_eq!( + step.tools() + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["client_local_tool"] + ); +} + +#[tokio::test(start_paused = true)] +async fn capture_binding_skips_pending_optional_servers_after_configured_shared_startup_grace() { + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + let mut plugin_config = crate::mcp::tests::test_mcp_config(std::env::temp_dir()); + let mut catalog = crate::ResolvedMcpCatalog::builder(); + catalog.register(crate::McpServerRegistration::from_plugin( + "pending-one".to_string(), + crate::McpPluginAttribution::new("optional-plugin".to_string(), "Optional".to_string()), + /*plugin_order*/ 0, + serde_json::from_value(serde_json::json!({ "command": "optional-plugin" })) + .expect("optional plugin MCP config"), + )); + catalog.register(crate::McpServerRegistration::from_selected_plugin( + "pending-selected".to_string(), + crate::McpPluginAttribution::new("selected-plugin".to_string(), "Selected".to_string()), + /*selection_order*/ 0, + serde_json::from_value(serde_json::json!({ "command": "selected-plugin" })) + .expect("selected plugin MCP config"), + )); + plugin_config.mcp_server_catalog = catalog.build(); + plugin_config.optional_mcp_startup_grace = Duration::from_millis(250); + manager.tool_plugin_context = Arc::new(crate::tool_plugin_context(&plugin_config)); + for server_name in ["pending-one", "pending-two", "pending-selected"] { + manager.insert_test_client( + server_name.to_string(), + AsyncManagedClient { + client: futures::future::pending::>() + .boxed() + .shared(), + is_codex_apps_mcp_server: false, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: None, + startup_complete: Arc::new(AtomicBool::new(false)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + } + + let mut required_manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + required_manager.tool_plugin_context = Arc::clone(&manager.tool_plugin_context); + required_manager.insert_test_client( + "pending-selected", + manager.test_client("pending-selected").clone(), + ); + required_manager.required_servers = vec!["pending-selected".to_string()]; + + let manager = Arc::new(manager); + assert!( + manager + .stable_catalog_revisions( + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new() + ) + .await + .is_none() + ); + let started = tokio::time::Instant::now(); + let binding = tokio::time::timeout( + Duration::from_millis(500), + manager.capture_binding_with_metadata( + Arc::new(plugin_config), + /*plugins_available*/ false, + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ), + ) + .await + .expect("all optional servers should share the configured startup grace"); + assert!(binding.tools().is_empty()); + assert_eq!(started.elapsed(), Duration::from_millis(250)); + + let binding = tokio::time::timeout(Duration::from_millis(1), capture_binding(&manager)) + .await + .expect("later bindings must not restart the optional startup grace"); + assert!(binding.tools().is_empty()); + + assert!( + tokio::time::timeout( + Duration::from_millis(1), + binding.list_resources("pending-one", /*params*/ None), + ) + .await + .is_err(), + "resources must wait for an omitted server instead of failing immediately" + ); + assert!( + tokio::time::timeout( + Duration::from_millis(1), + binding.list_all_resources(|server| server == "pending-one"), + ) + .await + .is_ok(), + "resource discovery must not wait for an omitted optional server" + ); + + for server_name in ["pending-one", "pending-selected"] { + let required_servers = vec![server_name.to_string()]; + let binding = tokio::time::timeout( + Duration::from_millis(1), + manager.capture_binding_with_metadata( + Arc::new(crate::mcp::tests::test_mcp_config(std::env::temp_dir())), + /*plugins_available*/ false, + &required_servers, + /*required_plugins*/ &HashSet::new(), + ), + ) + .await; + assert!(binding.is_err(), "explicitly requested servers must wait"); + } + // A plugin mention must still require startup after the optional grace has elapsed. + for (plugin_id, must_wait) in [ + ("selected-plugin", true), + ("optional-plugin", false), + ("selected-plugin-other", false), + ] { + let required_plugins = HashSet::from([plugin_id.to_string()]); + let binding = tokio::time::timeout( + Duration::from_millis(1), + manager.capture_binding_with_metadata( + Arc::new(crate::mcp::tests::test_mcp_config(std::env::temp_dir())), + /*plugins_available*/ false, + /*required_servers*/ &[], + &required_plugins, + ), + ) + .await; + assert_eq!( + binding.is_err(), + must_wait, + "plugin requirement {plugin_id}" + ); + } + assert!( + tokio::time::timeout( + Duration::from_millis(1500), + capture_binding(&Arc::new(required_manager)), + ) + .await + .is_err(), + "configured-required selected plugin servers must wait beyond the optional grace" + ); +} + +#[tokio::test(start_paused = true)] +async fn capture_binding_waits_for_optional_startup_when_shared_grace_is_disabled() { + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + let (client, startup_started, release_startup) = create_gated_async_managed_client( + create_test_managed_client(vec![create_test_tool("optional", "echo")]).await, + ); + manager.insert_test_client("optional", client); + + let mut config = crate::mcp::tests::test_mcp_config(std::env::temp_dir()); + config.optional_mcp_startup_grace = Duration::ZERO; + config + .server_permission_profiles + .insert("optional".to_string(), PermissionProfile::default()); + let manager = Arc::new(manager); + let mut capture = tokio::spawn(async move { + manager + .capture_binding_with_metadata( + Arc::new(config), + /*plugins_available*/ false, + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await + }); + + startup_started.await.expect("client startup should begin"); + assert!( + tokio::time::timeout(Duration::from_millis(1), &mut capture) + .await + .is_err(), + "disabled shared grace should keep waiting for optional startup" + ); + release_startup.send(()).expect("release client startup"); + + let binding = capture.await.expect("capture binding task"); + assert!(binding.prepare_call("optional", "echo").is_some()); +} + +#[tokio::test] +async fn stable_catalog_revisions_ignore_terminal_optional_server_failures() { + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + let ready = create_ready_async_managed_client(vec![create_test_tool("ready", "echo")]).await; + assert!(ready.client().await.is_ok()); + let mut failed = ready.clone(); + manager.insert_test_client("ready", ready); + failed.client = futures::future::ready::>(Err( + StartupOutcomeError::Failed { + error: "optional startup failed".to_string(), + is_authentication_required: false, + }, + )) + .boxed() + .shared(); + assert!(failed.client().await.is_err()); + manager.insert_test_client("failed", failed); + + assert!( + manager + .stable_catalog_revisions( + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new() + ) + .await + .is_some() + ); + manager.required_servers.push("failed".to_string()); + assert!( + manager + .stable_catalog_revisions( + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new() + ) + .await + .is_none() + ); + manager.required_servers.clear(); + + let binding = capture_binding(&Arc::new(manager)).await; + assert!(binding.prepare_call("ready", "echo").is_some()); +} + +#[tokio::test(start_paused = true)] +async fn capture_binding_shares_optional_startup_grace_across_connection_sets() { + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let cache = McpToolCatalogCache::default(); + let runtime_context = McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + std::env::temp_dir(), + ); + let server_config: McpServerConfig = + serde_json::from_value(serde_json::json!({ "command": "pending-mcp" })) + .expect("pending MCP server configuration"); + let cache_context = cache + .context( + "pending", + &server_config, + &runtime_context, + /*resolved_environment*/ None, + ( + &ElicitationCapability::default(), + &ClientMcpExtensions::default(), + ), + /*connection_identity*/ None, + ) + .expect("shared pending MCP catalog"); + + let create_connection_set = || { + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + "pending", + AsyncManagedClient { + client: futures::future::pending::>() + .boxed() + .shared(), + is_codex_apps_mcp_server: false, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: Some(cache_context.clone()), + startup_complete: Arc::new(AtomicBool::new(false)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + Arc::new(manager) + }; + + let first_started = tokio::time::Instant::now(); + let first = tokio::time::timeout( + Duration::from_millis(1500), + capture_binding(&create_connection_set()), + ) + .await + .expect("the first thread should receive the optional startup grace"); + assert!(first.tools().is_empty()); + assert_eq!(first_started.elapsed(), Duration::from_secs(1)); + + let second = tokio::time::timeout( + Duration::from_millis(1), + capture_binding(&create_connection_set()), + ) + .await + .expect("the next thread must not restart the same server's startup grace"); + assert!(second.tools().is_empty()); + + let mut disabled_config = crate::mcp::tests::test_mcp_config(std::env::temp_dir()); + disabled_config.optional_mcp_startup_grace = Duration::ZERO; + let disabled_manager = create_connection_set(); + assert!( + tokio::time::timeout( + Duration::from_millis(1), + disabled_manager.capture_binding_with_metadata( + Arc::new(disabled_config), + /*plugins_available*/ false, + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ), + ) + .await + .is_err(), + "disabled grace should keep waiting for the pending optional server" + ); + + let restored_started = tokio::time::Instant::now(); + let restored = tokio::time::timeout( + Duration::from_millis(1500), + capture_binding(&create_connection_set()), + ) + .await + .expect("restoring the startup grace should create a fresh deadline"); + assert!(restored.tools().is_empty()); + assert_eq!(restored_started.elapsed(), Duration::from_secs(1)); + + let mut updated_config = crate::mcp::tests::test_mcp_config(std::env::temp_dir()); + updated_config.optional_mcp_startup_grace = Duration::from_millis(250); + let updated_manager = create_connection_set(); + let updated_started = tokio::time::Instant::now(); + let updated = tokio::time::timeout( + Duration::from_millis(500), + updated_manager.capture_binding_with_metadata( + Arc::new(updated_config), + /*plugins_available*/ false, + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ), + ) + .await + .expect("a changed startup grace should receive its newly configured deadline"); + assert!(updated.tools().is_empty()); + assert_eq!(updated_started.elapsed(), Duration::from_millis(250)); + + cache_context.publish_if_newest( + cache_context.begin_fetch(), + &[create_test_tool("pending", "cached_tool")], + ); + let deadline_after_publication = tokio::time::Instant::now() + Duration::from_secs(1); + assert_eq!( + cache_context.optional_startup_deadline(deadline_after_publication, Duration::from_secs(1)), + deadline_after_publication, + "publishing a catalog must not install a stale startup deadline" + ); + let cached_manager = create_connection_set(); + let cached = tokio::time::timeout(Duration::from_millis(1), capture_binding(&cached_manager)) + .await + .expect("cached tools should be immediately available to later threads"); + assert_eq!( + cached + .tools() + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["cached_tool"] + ); + + tokio::time::advance(Duration::from_secs(30 * 60 + 1)).await; + assert!( + tokio::time::timeout(Duration::from_millis(1), capture_binding(&cached_manager)) + .await + .is_err(), + "an expired catalog should receive a fresh startup grace" + ); + + cache_context.disable(); + for _ in 0..2 { + let started = tokio::time::Instant::now(); + let binding = tokio::time::timeout( + Duration::from_millis(1500), + capture_binding(&create_connection_set()), + ) + .await + .expect("non-cacheable servers should keep their per-thread startup grace"); + assert!(binding.tools().is_empty()); + assert_eq!(started.elapsed(), Duration::from_secs(1)); + } +} + +#[tokio::test(start_paused = true)] +async fn capture_binding_uses_cache_published_during_optional_startup() { + let runtime_context = McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + std::env::temp_dir(), + ); + let server_config: McpServerConfig = + serde_json::from_value(serde_json::json!({ "command": "pending-mcp" })) + .expect("server configuration"); + let cache_context = McpToolCatalogCache::default() + .context( + "pending", + &server_config, + &runtime_context, + /*resolved_environment*/ None, + ( + &ElicitationCapability::default(), + &ClientMcpExtensions::default(), + ), + /*connection_identity*/ None, + ) + .expect("shared catalog"); + let (mut client, started, _release) = create_gated_async_managed_client( + create_test_managed_client(vec![create_test_tool("pending", "live_tool")]).await, + ); + client.tool_catalog_cache_context = Some(cache_context.clone()); + let mut manager = McpConnectionSet::new_uninitialized( + &Constrained::allow_any(AskForApproval::OnRequest), + &Constrained::allow_any(PermissionProfile::default()), + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client("pending", client); + let manager = Arc::new(manager); + + let binding = capture_binding(&manager); + tokio::pin!(binding); + assert!(futures::poll!(&mut binding).is_pending()); + started + .await + .expect("optional startup began without a cache"); + cache_context.publish_if_newest( + cache_context.begin_fetch(), + &[create_test_tool("pending", "peer_tool")], + ); + tokio::time::advance(Duration::from_secs(/*secs*/ 2)).await; + + let binding = binding.await; + assert!( + !manager.servers["pending"] + .connection + .client + .startup_complete + .load(Ordering::Acquire) + ); + assert_eq!( + model_tool_names(binding.tools()), + HashSet::from([ToolName::namespaced("mcp__pending", "peer_tool")]), + ); +} + +#[tokio::test(start_paused = true)] +async fn capture_binding_retains_cached_tools_that_expire_while_waiting_for_another_server() { + let runtime_context = McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + std::env::temp_dir(), + ); + let server_config: McpServerConfig = + serde_json::from_value(serde_json::json!({ "command": "server-a-mcp" })) + .expect("MCP server configuration"); + let cache_context = McpToolCatalogCache::default() + .context( + "server_a", + &server_config, + &runtime_context, + /*resolved_environment*/ None, + ( + &ElicitationCapability::default(), + &ClientMcpExtensions::default(), + ), + /*connection_identity*/ None, + ) + .expect("server A cache context"); + let (mut client_a, _started_a, _release_a) = create_gated_async_managed_client( + create_test_managed_client(vec![create_test_tool("server_a", "live_tool")]).await, + ); + client_a.tool_catalog_cache_context = Some(cache_context.clone()); + let (client_b, started_b, release_b) = create_gated_async_managed_client( + create_test_managed_client(vec![create_test_tool("server_b", "tool_b")]).await, + ); + cache_context.publish_if_newest( + cache_context.begin_fetch(), + &[create_test_tool("server_a", "cached_tool")], + ); + tokio::time::advance(Duration::from_secs(30 * 60 - 1)).await; + + let mut manager = McpConnectionSet::new_uninitialized( + &Constrained::allow_any(AskForApproval::OnRequest), + &Constrained::allow_any(PermissionProfile::default()), + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client("server_a", client_a); + manager.insert_test_client("server_b", client_b); + Arc::get_mut(&mut manager.servers.get_mut("server_a").unwrap().connection) + .expect("unique server A connection") + .startup_trigger = Some(watch::channel(/*init*/ false).0); + manager.required_servers = vec!["server_a".to_string(), "server_b".to_string()]; + let manager = Arc::new(manager); + + let binding = capture_binding(&manager); + tokio::pin!(binding); + assert!(futures::poll!(&mut binding).is_pending()); + started_b.await.expect("server B startup should begin"); + assert!(manager.servers["server_a"].connection.startup_is_dormant()); + + // Expire A's catalog while capture is waiting for B's required startup. + tokio::time::advance(Duration::from_secs(/*secs*/ 2)).await; + assert!(cache_context.current_tools().is_none()); + release_b.send(()).expect("release server B startup"); + + let binding = tokio::time::timeout(Duration::from_secs(/*secs*/ 1), binding) + .await + .expect("cached server A should not need startup"); + assert_eq!( + model_tool_names(binding.tools()), + HashSet::from([ + ToolName::namespaced("mcp__server_a", "cached_tool"), + ToolName::namespaced("mcp__server_b", "tool_b"), + ]) + ); + assert!(manager.servers["server_a"].connection.startup_is_dormant()); +} + +#[tokio::test] +async fn capture_binding_omits_cache_disabled_while_waiting_for_another_server() { + let runtime_context = McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + std::env::temp_dir(), + ); + let server_config: McpServerConfig = + serde_json::from_value(serde_json::json!({ "command": "cached-mcp" })) + .expect("server configuration"); + let cache_context = McpToolCatalogCache::default() + .context( + "cached", + &server_config, + &runtime_context, + /*resolved_environment*/ None, + ( + &ElicitationCapability::default(), + &ClientMcpExtensions::default(), + ), + /*connection_identity*/ None, + ) + .expect("shared catalog"); + cache_context.publish_if_newest( + cache_context.begin_fetch(), + &[create_test_tool("cached", "cached_tool")], + ); + let (mut cached, _started, _release) = create_gated_async_managed_client( + create_test_managed_client(vec![create_test_tool("cached", "live_tool")]).await, + ); + cached.tool_catalog_cache_context = Some(cache_context.clone()); + let (waiting, started, release) = create_gated_async_managed_client( + create_test_managed_client(vec![create_test_tool("waiting", "ready_tool")]).await, + ); + let mut manager = McpConnectionSet::new_uninitialized( + &Constrained::allow_any(AskForApproval::OnRequest), + &Constrained::allow_any(PermissionProfile::default()), + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client("cached", cached); + manager.insert_test_client("waiting", waiting); + manager.required_servers = vec!["waiting".to_string()]; + let manager = Arc::new(manager); + + let binding = capture_binding(&manager); + tokio::pin!(binding); + assert!(futures::poll!(&mut binding).is_pending()); + started.await.expect("required server startup began"); + cache_context.disable(); + release.send(()).expect("release required server startup"); + + assert_eq!( + model_tool_names(binding.await.tools()), + HashSet::from([ToolName::namespaced("mcp__waiting", "ready_tool")]), + ); +} + +#[tokio::test] +async fn capture_binding_resolves_concurrently_and_rechecks_cached_clients() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + store_current_tools( + &cache_context, + vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "shared_cached_tool", + )], + ); + let ready_apps_client = create_test_managed_client(vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "client_local_tool", + )]) + .await; + let (mut apps_client, apps_started, release_apps) = + create_gated_async_managed_client(ready_apps_client); + apps_client.is_codex_apps_mcp_server = true; + apps_client.codex_apps_tools_cache_context = Some(cache_context); + let first_client = + create_test_managed_client(vec![create_test_tool("first", "first_tool")]).await; + let second_client = + create_test_managed_client(vec![create_test_tool("second", "second_tool")]).await; + let (first_client, first_started, release_first) = + create_gated_async_managed_client(first_client); + let (second_client, second_started, release_second) = + create_gated_async_managed_client(second_client); + + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client(CODEX_APPS_MCP_SERVER_NAME, apps_client); + manager.insert_test_client("first", first_client); + manager.insert_test_client("second", second_client); + let manager = Arc::new(manager); + + let manager_for_startup = Arc::clone(&manager); + let startup = tokio::spawn(async move { + manager_for_startup + .wait_for_server_startup(CODEX_APPS_MCP_SERVER_NAME) + .await + }); + tokio::time::timeout(Duration::from_secs(1), apps_started) + .await + .expect("Codex Apps startup should begin") + .expect("signal Codex Apps startup"); + + let manager_for_binding = Arc::clone(&manager); + let binding = tokio::spawn(async move { capture_binding(&manager_for_binding).await }); + tokio::time::timeout(Duration::from_secs(1), async { + first_started.await.expect("first server startup"); + second_started.await.expect("second server startup"); + }) + .await + .expect("both uncached servers should start before either is released"); + + release_apps.send(()).expect("release Codex Apps startup"); + assert!(startup.await.expect("Codex Apps startup task")); + release_first.send(()).expect("release first server"); + release_second.send(()).expect("release second server"); + + let binding = binding.await.expect("binding capture should complete"); + assert_eq!( + binding + .tools() + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + HashSet::from(["client_local_tool", "first_tool", "second_tool"]) + ); + assert!( + binding + .prepare_call(CODEX_APPS_MCP_SERVER_NAME, "client_local_tool") + .is_some() + ); + assert!( + binding + .prepare_call(CODEX_APPS_MCP_SERVER_NAME, "shared_cached_tool") + .is_none() + ); + assert!(binding.prepare_call("first", "first_tool").is_some()); + assert!(binding.prepare_call("second", "second_tool").is_some()); +} + +#[tokio::test] +async fn list_all_tools_applies_legacy_mcp_prefix_by_default() { + let managed_client = + create_ready_async_managed_client(vec![create_test_tool("rmcp", "echo")]).await; + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client("rmcp", managed_client); + + let tools = manager.list_all_tools().await; + let tool = tools + .iter() + .find(|tool| tool.canonical_tool_name() == ToolName::namespaced("mcp__rmcp", "echo")) + .expect("legacy-prefixed MCP tool name should resolve"); + + let expected = ("rmcp", "mcp__rmcp", "echo", "echo"); + assert_eq!( + ( + tool.server_name.as_str(), + tool.callable_namespace.as_str(), + tool.callable_name.as_str(), + tool.tool.name.as_ref(), + ), + expected + ); +} + +#[tokio::test] +async fn call_tool_requires_connection_without_waiting_for_startup() { + let client = create_test_managed_client(vec![create_test_tool("docs", "search")]).await; + let (client, startup_started, release_startup) = create_gated_async_managed_client(client); + let startup_client = client.clone(); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client("docs", client); + + let pending_call = tokio::time::timeout( + Duration::from_millis(50), + manager.call_tool( + "docs", + "search", + /*environment_id*/ None, + /*arguments*/ None, + /*meta*/ None, + Some(Duration::from_secs(5)), + /*wait_for_server*/ false, + ), + ) + .await + .expect("ready-only invocation must not wait for pending server startup") + .expect_err("pending server must not accept ready-only calls"); + assert!(pending_call.to_string().contains("not connected")); + + let startup = tokio::spawn(async move { startup_client.client().await }); + startup_started.await.expect("server startup should begin"); + release_startup.send(()).expect("release server startup"); + startup + .await + .expect("startup task should finish") + .expect("server startup should succeed"); + + let ready_error = manager + .call_tool( + "docs", + "search", + /*environment_id*/ None, + /*arguments*/ None, + /*meta*/ None, + Some(Duration::from_secs(5)), + /*wait_for_server*/ false, + ) + .await + .expect_err("ready server should reach the uninitialized test transport"); + assert!(format!("{ready_error:#}").contains("MCP client not initialized")); +} + +#[tokio::test] +async fn connected_call_respects_server_tool_filters() { + let client = create_ready_async_managed_client(vec![create_test_tool("docs", "search")]).await; + client.client().await.expect("server should be ready"); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client("docs", client); + manager + .servers + .get_mut("docs") + .expect("test server should exist") + .tool_filter + .disabled + .insert("search".to_string()); + + let filtered_error = manager + .call_tool( + "docs", + "search", + /*environment_id*/ None, + /*arguments*/ None, + /*meta*/ None, + Some(Duration::from_secs(5)), + /*wait_for_server*/ false, + ) + .await + .expect_err("disabled tools should not be callable"); + assert!(filtered_error.to_string().contains("disabled")); +} + +#[tokio::test] +async fn call_tool_validates_environment_without_waiting_for_ready_connections() { + let client = create_test_managed_client(vec![create_test_tool("docs", "search")]).await; + let (client, _, _) = create_gated_async_managed_client(client); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client("docs", client); + manager + .servers + .get_mut("docs") + .expect("test server should exist") + .metadata + .environment_id = "executor-a".to_string(); + + let mismatched_environment = manager + .call_tool( + "docs", + "search", + Some("executor-b"), + /*arguments*/ None, + /*meta*/ None, + Some(Duration::from_secs(5)), + /*wait_for_server*/ false, + ) + .await + .expect_err("calls must reject a server from a different environment"); + assert_eq!( + mismatched_environment.to_string(), + "MCP server `docs` is running in environment `executor-a`, expected `executor-b`" + ); + + let pending_call = tokio::time::timeout( + Duration::from_millis(50), + manager.call_tool( + "docs", + "search", + Some("executor-a"), + /*arguments*/ None, + /*meta*/ None, + Some(Duration::from_secs(5)), + /*wait_for_server*/ false, + ), + ) + .await + .expect("environment-scoped calls must not wait for pending server startup") + .expect_err("pending server must not accept environment-scoped calls"); + assert!(pending_call.to_string().contains("not connected")); +} + +#[tokio::test] +async fn list_all_tools_resolves_server_catalogs_concurrently() { + let first_client = create_test_managed_client(vec![create_test_tool("first", "search")]).await; + let second_client = + create_test_managed_client(vec![create_test_tool("second", "lookup")]).await; + let (first_client, first_started, release_first) = + create_gated_async_managed_client(first_client); + let (second_client, second_started, release_second) = + create_gated_async_managed_client(second_client); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client("first", first_client); + manager.insert_test_client("second", second_client); + let manager = Arc::new(manager); + let manager_for_listing = Arc::clone(&manager); + let listing = tokio::spawn(async move { manager_for_listing.list_all_tools().await }); + + tokio::time::timeout(Duration::from_secs(1), async { + first_started.await.expect("first server startup"); + second_started.await.expect("second server startup"); + }) + .await + .expect("both server catalogs should start before either is released"); + release_first.send(()).expect("release first server"); + release_second.send(()).expect("release second server"); + + let tools = listing.await.expect("tool listing should complete"); + assert_eq!( + model_tool_names(&tools), + HashSet::from([ + ToolName::namespaced("mcp__first", "search"), + ToolName::namespaced("mcp__second", "lookup"), + ]) + ); +} + +#[tokio::test] +async fn list_all_tools_blocks_while_client_is_pending_without_cached_tools() { + let pending_client = futures::future::pending::>() + .boxed() + .shared(); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: pending_client, + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(false)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + + let timeout_result = + tokio::time::timeout(Duration::from_millis(10), manager.list_all_tools()).await; + assert!(timeout_result.is_err()); +} + +#[tokio::test] +async fn cancelling_startup_does_not_disable_a_ready_client() { + let client = create_ready_async_managed_client(vec![create_test_tool("ready", "search")]).await; + + client.cancel_token.cancel(); + + let managed = client + .client() + .await + .expect("startup cancellation should not disable a ready client"); + assert_eq!( + managed + .tool_catalog + .read(|catalog| model_tool_names(&catalog.tools)) + .await, + HashSet::from([ToolName::namespaced("ready", "search")]) + ); +} + +#[tokio::test] +async fn shutdown_cancels_pending_tool_listing() { + let cancel_token = CancellationToken::new(); + let cancel_token_for_startup = cancel_token.clone(); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let pending_client = async move { + let _ = started_tx.send(()); + cancel_token_for_startup.cancelled().await; + Err(StartupOutcomeError::Cancelled) + } + .boxed() + .shared(); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: pending_client, + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(false)), + startup_reconnect: None, + cancel_token, + }, + ); + let manager = Arc::new(manager); + let manager_for_list = Arc::clone(&manager); + let list_task = tokio::spawn(async move { manager_for_list.list_all_tools().await }); + + started_rx.await.expect("tool listing should start"); + tokio::time::timeout(Duration::from_secs(1), manager.shutdown()) + .await + .expect("shutdown should cancel speculative tool listing"); + let tools = list_task.await.expect("tool listing task should not panic"); + assert!(tools.is_empty()); +} + +#[tokio::test] +async fn shutdown_continues_after_caller_is_aborted() { + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (completed_tx, completed_rx) = tokio::sync::oneshot::channel(); + let release = Arc::new(tokio::sync::Notify::new()); + let release_for_client = Arc::clone(&release); + let blocking_client = async move { + let _ = started_tx.send(()); + release_for_client.notified().await; + let _ = completed_tx.send(()); + Err(StartupOutcomeError::Cancelled) + } + .boxed() + .shared(); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: blocking_client, + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(false)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + let manager = Arc::new(manager); + let shutdown_task = tokio::spawn({ + let manager = Arc::clone(&manager); + async move { manager.shutdown().await } + }); + + started_rx.await.expect("client shutdown should start"); + shutdown_task.abort(); + let shutdown_error = shutdown_task + .await + .expect_err("caller shutdown task should be aborted"); + assert!(shutdown_error.is_cancelled()); + release.notify_one(); + + tokio::time::timeout(Duration::from_secs(1), completed_rx) + .await + .expect("client shutdown should survive caller cancellation") + .expect("client shutdown completion sender should stay alive"); +} + +#[tokio::test] +async fn list_all_tools_does_not_block_when_shared_codex_apps_cache_is_empty() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + store_current_tools(&cache_context, Vec::new()); + let pending_client = futures::future::pending::>() + .boxed() + .shared(); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: pending_client, + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: Some(cache_context), + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(false)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + + let timeout_result = + tokio::time::timeout(Duration::from_millis(10), manager.list_all_tools()).await; + let tools = timeout_result.expect("shared empty cache should not block"); + assert!(tools.is_empty()); +} + +#[tokio::test] +async fn list_all_tools_uses_shared_codex_apps_cache_when_client_startup_fails() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + store_current_tools( + &cache_context, + vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "calendar_create_event", + )], + ); + let server_info = create_test_server_info("Codex Apps"); + let failed_client = futures::future::ready::>(Err( + StartupOutcomeError::Failed { + error: "startup failed".to_string(), + is_authentication_required: false, + }, + )) + .boxed() + .shared(); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + let startup_complete = Arc::new(std::sync::atomic::AtomicBool::new(true)); + manager.insert_test_client( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + AsyncManagedClient { + client: failed_client, + is_codex_apps_mcp_server: true, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: Some(server_info.clone()), + codex_apps_tools_cache_context: Some(cache_context), + tool_catalog_cache_context: None, + startup_complete, + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + ); + + let tools = manager.list_all_tools().await; + let tool = tools + .iter() + .find(|tool| { + tool.canonical_tool_name() + == ToolName::namespaced("mcp__codex_apps", "calendar_create_event") + }) + .expect("tool from shared cache"); + assert_eq!(tool.server_name, CODEX_APPS_MCP_SERVER_NAME); + assert_eq!(tool.callable_name, "calendar_create_event"); + assert_eq!( + manager + .list_available_server_infos() + .await + .get(CODEX_APPS_MCP_SERVER_NAME), + Some(&server_info) + ); +} + +#[tokio::test] +async fn list_all_tools_reconnects_failed_codex_apps_startup_and_reuses_client() { + let recovered_client = create_test_managed_client(vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "drive_search", + )]) + .await; + let attempts = Arc::new(AtomicUsize::new(0)); + let attempts_for_reconnect = Arc::clone(&attempts); + let reconnect_finished = Arc::new(tokio::sync::Notify::new()); + let reconnect_finished_for_factory = Arc::clone(&reconnect_finished); + let reconnect_factory = Arc::new(move || { + attempts_for_reconnect.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + let reconnect_finished = Arc::clone(&reconnect_finished_for_factory); + let recovered_client = recovered_client.clone(); + async move { + reconnect_finished.notify_one(); + Ok(recovered_client) + } + .boxed() + .shared() + }); + let mut manager = create_test_manager_with_failed_apps_startup(Vec::new(), reconnect_factory); + manager + .servers + .get_mut(CODEX_APPS_MCP_SERVER_NAME) + .expect("test server exists") + .metadata = McpServerMetadata { + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + pollutes_memory: false, + origin: None, + supports_parallel_tool_calls: false, + default_tools_approval_mode: None, + tool_approval_modes: HashMap::new(), + }; + let manager = Arc::new(manager); + + assert!( + manager + .stable_catalog_revisions( + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new() + ) + .await + .is_none() + ); + let reconnect_finished_wait = reconnect_finished.notified(); + let tools = manager.list_all_tools().await; + assert!(tools.is_empty()); + reconnect_finished_wait.await; + + let tools = manager.list_all_tools().await; + assert_eq!( + tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["drive_search"] + ); + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 1); + assert!( + manager + .stable_catalog_revisions( + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new() + ) + .await + .is_some() + ); + + let step = capture_binding(&manager).await; + let prepared = step + .prepare_call(CODEX_APPS_MCP_SERVER_NAME, "drive_search") + .expect("recovered tool should have a prepared call"); + assert!( + !prepared + .server_supports_sandbox_state_meta_capability() + .await + .expect("prepared call should use the recovered client") + ); + + let tools = manager.list_all_tools().await; + assert_eq!( + tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["drive_search"] + ); + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 1); +} + +#[tokio::test(start_paused = true)] +async fn later_tool_list_retries_after_failed_reconnect_and_keeps_cached_tools() { + let recovered_client = create_test_managed_client(vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "drive_search", + )]) + .await; + let attempts = Arc::new(AtomicUsize::new(0)); + let attempts_for_reconnect = Arc::clone(&attempts); + let reconnect_finished = Arc::new(tokio::sync::Notify::new()); + let reconnect_finished_for_factory = Arc::clone(&reconnect_finished); + let reconnect_factory = Arc::new(move || { + let attempt = attempts_for_reconnect.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + let reconnect_finished = Arc::clone(&reconnect_finished_for_factory); + let recovered_client = recovered_client.clone(); + async move { + let result = if attempt < 2 { + Err(StartupOutcomeError::Failed { + error: "recreated startup failed".to_string(), + is_authentication_required: false, + }) + } else { + Ok(recovered_client) + }; + reconnect_finished.notify_one(); + result + } + .boxed() + .shared() + }); + let manager = create_test_manager_with_failed_apps_startup( + vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "cached_drive_search", + )], + reconnect_factory, + ); + + let first_reconnect_finished = reconnect_finished.notified(); + let tools = manager.list_all_tools().await; + assert_eq!( + tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["cached_drive_search"] + ); + first_reconnect_finished.await; + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 1); + + let tools = manager.list_all_tools().await; + assert_eq!( + tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["cached_drive_search"] + ); + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 1); + + tokio::time::advance(CODEX_APPS_RECONNECT_INITIAL_BACKOFF).await; + let second_reconnect_finished = reconnect_finished.notified(); + let tools = manager.list_all_tools().await; + assert_eq!( + tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["cached_drive_search"] + ); + second_reconnect_finished.await; + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 2); + + tokio::time::advance(CODEX_APPS_RECONNECT_INITIAL_BACKOFF).await; + let tools = manager.list_all_tools().await; + assert_eq!( + tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["cached_drive_search"] + ); + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 2); + + tokio::time::advance(CODEX_APPS_RECONNECT_INITIAL_BACKOFF).await; + let third_reconnect_finished = reconnect_finished.notified(); + let tools = manager.list_all_tools().await; + assert_eq!( + tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["cached_drive_search"] + ); + third_reconnect_finished.await; + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 3); + + let tools = manager.list_all_tools().await; + assert_eq!( + tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["drive_search"] + ); +} + +#[tokio::test] +async fn tool_lists_do_not_block_and_share_codex_apps_startup_reconnect() { + let recovered_client = create_test_managed_client(vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "drive_search", + )]) + .await; + let attempts = Arc::new(AtomicUsize::new(0)); + let attempts_for_reconnect = Arc::clone(&attempts); + let reconnect_started = Arc::new(tokio::sync::Notify::new()); + let reconnect_started_for_factory = Arc::clone(&reconnect_started); + let release_reconnect = Arc::new(tokio::sync::Notify::new()); + let release_reconnect_for_factory = Arc::clone(&release_reconnect); + let reconnect_factory = Arc::new(move || { + let recovered_client = recovered_client.clone(); + let attempts = Arc::clone(&attempts_for_reconnect); + let reconnect_started = Arc::clone(&reconnect_started_for_factory); + let release_reconnect = Arc::clone(&release_reconnect_for_factory); + async move { + attempts.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + reconnect_started.notify_one(); + release_reconnect.notified().await; + Ok(recovered_client) + } + .boxed() + .shared() + }); + let mut manager = create_test_manager_with_failed_apps_startup( + vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "cached_drive_search", + )], + reconnect_factory, + ); + manager.set_test_server_metadata( + CODEX_APPS_MCP_SERVER_NAME, + McpServerMetadata { + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + pollutes_memory: false, + origin: None, + supports_parallel_tool_calls: false, + default_tools_approval_mode: None, + tool_approval_modes: HashMap::new(), + }, + ); + let manager = Arc::new(manager); + let reconnect_started_wait = reconnect_started.notified(); + let first_tools = tokio::time::timeout(Duration::from_millis(10), manager.list_all_tools()) + .await + .expect("cached tools should not wait for reconnect"); + + reconnect_started_wait.await; + let second_tools = tokio::time::timeout(Duration::from_millis(10), manager.list_all_tools()) + .await + .expect("concurrent cached tools should not wait for reconnect"); + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 1); + assert_eq!( + first_tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["cached_drive_search"] + ); + assert_eq!( + second_tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["cached_drive_search"] + ); + let pending_step = tokio::time::timeout(Duration::from_millis(10), capture_binding(&manager)) + .await + .expect("step capture should not wait for reconnect"); + assert!( + pending_step.tools().is_empty(), + "a model step must not advertise cached tools without an exact ready client" + ); + + release_reconnect.notify_one(); + tokio::task::yield_now().await; + let tools = manager.list_all_tools().await; + assert_eq!( + tools + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["drive_search"] + ); + let recovered_step = capture_binding(&manager).await; + assert_eq!( + recovered_step + .tools() + .iter() + .map(|tool| tool.callable_name.as_str()) + .collect::>(), + vec!["drive_search"] + ); + assert!( + recovered_step + .prepare_call(CODEX_APPS_MCP_SERVER_NAME, "drive_search") + .is_some() + ); + assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 1); +} + +#[tokio::test] +async fn list_all_tools_adds_server_metadata_to_tools() { + let server_name = "docs"; + let managed_client = + create_ready_async_managed_client(vec![create_test_tool(server_name, "search")]).await; + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + manager.insert_test_client(server_name, managed_client); + manager.set_test_server_metadata( + server_name, + McpServerMetadata { + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + pollutes_memory: true, + origin: Some(McpServerOrigin::StreamableHttp( + "https://docs.example".to_string(), + )), + supports_parallel_tool_calls: true, + default_tools_approval_mode: None, + tool_approval_modes: HashMap::new(), + }, + ); + + let tools = manager.list_all_tools().await; + assert_eq!(tools.len(), 1); + let tool = &tools[0]; + assert_eq!(tool.server_name, server_name); + assert!(tool.supports_parallel_tool_calls); + assert_eq!(tool.server_origin.as_deref(), Some("https://docs.example")); +} + +#[test] +fn server_metadata_preserves_tool_approval_policy() { + let mut config = crate::codex_apps_mcp_server_config( + "https://docs.example", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ); + config.environment_id = "remote".to_string(); + config.default_tools_approval_mode = Some(AppToolApproval::Prompt); + config.tools.insert( + "search".to_string(), + McpServerToolConfig { + approval_mode: Some(AppToolApproval::Approve), + ..Default::default() + }, + ); + let metadata = McpServerMetadata::from(&EffectiveMcpServer::configured(config)); + + assert_eq!(metadata.environment_id, "remote"); + assert_eq!(metadata.tool_approval_mode("read"), AppToolApproval::Prompt); + assert_eq!( + metadata.tool_approval_mode("search"), + AppToolApproval::Approve + ); +} + +#[test] +fn hosted_actor_credentials_are_only_available_to_host_owned_mcp_servers() { + let bootstrap_auth = CodexAuth::create_dummy_chatgpt_auth_for_testing(); + let mut actor_headers = Default::default(); + codex_model_provider::auth_provider_from_auth(&bootstrap_auth) + .add_auth_headers(&mut actor_headers); + actor_headers.insert( + "x-openai-actor-authorization", + "hosted-actor-secret" + .parse() + .expect("valid actor authorization header"), + ); + let hosted_auth = CodexAuth::Headers(AuthHeaders::new(actor_headers)); + let provider = codex_model_provider::auth_provider_from_auth(&hosted_auth); + let mut local_config = crate::codex_apps_mcp_server_config( + "https://chatgpt.com", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ); + local_config.auth = McpServerAuth::ChatGpt; + + let local_server = EffectiveMcpServer::configured(local_config.clone()); + let local_provider = + chatgpt_auth_provider_for_server(&local_server, Some(Arc::clone(&provider))) + .expect("host-owned Codex Apps must retain hosted authentication"); + assert_eq!( + local_provider + .to_auth_headers() + .get("x-openai-actor-authorization") + .and_then(|value| value.to_str().ok()), + Some("hosted-actor-secret") + ); + + let mut remote_config = local_config; + remote_config.environment_id = "customer-executor".to_string(); + let remote_server = EffectiveMcpServer::configured(remote_config); + assert!( + chatgpt_auth_provider_for_server(&remote_server, Some(provider)).is_none(), + "customer-owned executors must never receive hosted actor credentials" + ); +} + +#[tokio::test] +async fn executor_owned_chatgpt_mcp_accepts_only_safe_explicit_authorization() -> anyhow::Result<()> +{ + let codex_home = tempdir()?; + let environment_manager = Arc::new(environment_manager_without_environments()); + environment_manager.upsert_environment( + "customer-executor".to_string(), + "ws://127.0.0.1:1".to_string(), + /*connect_timeout*/ None, + )?; + let runtime_context = + McpRuntimeContext::new(Arc::clone(&environment_manager), PathBuf::from("/tmp")); + let bootstrap_auth = CodexAuth::create_dummy_chatgpt_auth_for_testing(); + let mut actor_headers = Default::default(); + codex_model_provider::auth_provider_from_auth(&bootstrap_auth) + .add_auth_headers(&mut actor_headers); + actor_headers.insert( + "x-openai-actor-authorization", + "hosted-actor-secret" + .parse() + .expect("valid actor authorization header"), + ); + let hosted_auth = CodexAuth::Headers(AuthHeaders::new(actor_headers)); + let runtime_config = crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()); + let cases = [ + ("missing", None, None, None, false), + ("empty", Some(("Authorization", "")), None, None, false), + ( + "whitespace", + Some(("Authorization", " \t ")), + None, + None, + false, + ), + ( + "invalid newline", + Some(("Authorization", "Bearer executor\r\nsecret")), + None, + None, + false, + ), + ( + "invalid NUL", + Some(("Authorization", "Bearer executor\0secret")), + None, + None, + false, + ), + ( + "invalid DEL", + Some(("Authorization", "Bearer executor\u{007f}secret")), + None, + None, + false, + ), + ( + "environment header", + Some(("Authorization", "Bearer executor-secret")), + None, + Some(("aUtHoRiZaTiOn", "CODEX_TEST_HOSTED_SECRET")), + false, + ), + ( + "environment bearer", + Some(("Authorization", "Bearer executor-secret")), + Some("CODEX_TEST_HOSTED_SECRET"), + None, + false, + ), + ( + "mixed-case static header", + Some(("aUtHoRiZaTiOn", "Bearer executor-secret")), + None, + None, + true, + ), + ]; + + for (case, static_header, bearer_env_var, env_header, allows_executor_auth) in cases { + let mut server_json = serde_json::json!({ + "url": "https://chatgpt.com/backend-api/ps/mcp", + "auth": "chatgpt", + "environment_id": "customer-executor", + }); + if let Some((name, value)) = static_header { + server_json["http_headers"] = serde_json::json!({ name: value }); + } + if let Some(name) = bearer_env_var { + server_json["bearer_token_env_var"] = serde_json::json!(name); + } + if let Some((name, value)) = env_header { + server_json["env_http_headers"] = serde_json::json!({ name: value }); + } + let server_config = serde_json::from_value::(server_json)?; + let mcp_servers = crate::effective_mcp_servers_from_configured( + HashMap::from([("fake-first-party".to_string(), server_config)]), + &runtime_config, + Some(&hosted_auth), + ); + assert!(matches!( + mcp_servers["fake-first-party"].config().auth, + McpServerAuth::ChatGpt + )); + let remote_server = &mcp_servers["fake-first-party"]; + assert!( + chatgpt_auth_provider_for_server( + remote_server, + Some(codex_model_provider::auth_provider_from_auth(&hosted_auth)), + ) + .is_none(), + "{case}: executor-owned servers must never receive hosted actor credentials" + ); + let resolved_environment = + runtime_context.resolve_server_environment("fake-first-party", remote_server.config()); + let connection_identity = |keyring_backend_kind| { + McpServerConnectionIdentity::new( + "fake-first-party", + remote_server, + /*host_plugin_root*/ None, + OAuthCredentialsStoreMode::File, + keyring_backend_kind, + McpOAuthRefreshMode::Legacy, + &resolved_environment, + &runtime_context, + /*runtime_auth_provider*/ None, + Some(&hosted_auth), + /*codex_apps_cache_identity*/ None, + ElicitationCapability::default(), + ClientMcpExtensions::default(), + /*previous_identity*/ None, + ) + }; + let direct_keyring_identity = connection_identity(AuthKeyringBackendKind::Direct); + let secrets_keyring_identity = connection_identity(AuthKeyringBackendKind::Secrets); + assert!( + direct_keyring_identity.has_same_connection_config(&secrets_keyring_identity), + "{case}: executor-owned servers must not inspect orchestrator OAuth stores" + ); + assert!( + direct_keyring_identity + .oauth_credentials() + .expect("executor-owned ChatGPT authentication must skip OAuth lookup") + .is_none(), + "{case}: executor-owned servers must not retain hosted OAuth credentials" + ); + let auth_statuses = crate::compute_auth_statuses( + mcp_servers.iter(), + OAuthCredentialsStoreMode::default(), + AuthKeyringBackendKind::default(), + Some(&hosted_auth), + &runtime_context, + ) + .await; + let expected_auth_state = if allows_executor_auth { + McpAuthState::BearerToken + } else { + McpAuthState::Unsupported + }; + assert_eq!( + auth_statuses["fake-first-party"].auth_state, expected_auth_state, + "{case}: auth status must only accept safe executor-owned authorization" + ); + + let manager = McpConnectionSet::new( + /*previous*/ None, + McpPublicationGate::already_published(), + McpRuntimeInput { + startup_policy: McpStartupPolicy::Eager, + config: Arc::new(runtime_config.clone()), + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + mcp_servers, + submit_id: "security-test".to_string(), + tx_event: None, + startup_cancellation_token: CancellationToken::new(), + runtime_context: runtime_context.clone(), + codex_apps_tools_cache: ConnectorRuntimeManager::default(), + tool_catalog_cache: McpToolCatalogCache::default(), + codex_apps_tools_cache_key: ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ), + client_mcp_extensions: ClientMcpExtensions::default(), + auth: Some(hosted_auth.clone()), + auth_manager: None, + elicitation_reviewer: None, + elicitation_lifecycle: None, + }, + ElicitationRequestRouter::default(), + ) + .await; + let error = match manager.test_client("fake-first-party").client().await { + Ok(_) => panic!("{case}: the unreachable fake executor must not connect"), + Err(error) => error, + }; + let StartupOutcomeError::Failed { error, .. } = error else { + panic!("{case}: executor-owned authentication must fail rather than be cancelled"); + }; + if allows_executor_auth { + assert!( + error.contains("127.0.0.1:1"), + "{case}: safe explicit credentials should reach the executor: {error}" + ); + } else { + assert_eq!( + error, + "executor-owned MCP server `fake-first-party` cannot use hosted ChatGPT authentication; configure executor-owned credentials instead", + "{case}: unsafe credentials must fail before contacting the executor" + ); + } + } + + Ok(()) +} + +#[tokio::test] +async fn no_local_runtime_fails_local_stdio_but_keeps_local_http_server() { + let codex_home = tempdir().expect("tempdir"); + let mcp_servers = HashMap::from([ + ( + "stdio".to_string(), + EffectiveMcpServer::configured(McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::Stdio { + command: "echo".to_string(), + args: Vec::new(), + env: None, + env_vars: Vec::new(), + cwd: None, + }, + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + }), + ), + ( + "http".to_string(), + EffectiveMcpServer::configured(McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: "http://127.0.0.1:1".to_string(), + bearer_token_env_var: None, + http_headers: None, + env_http_headers: None, + http_headers_helper: None, + }, + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + }), + ), + ]); + + let cancel_token = CancellationToken::new(); + let manager = McpConnectionSet::new( + /*previous*/ None, + McpPublicationGate::already_published(), + McpRuntimeInput { + startup_policy: McpStartupPolicy::Eager, + config: Arc::new(crate::mcp::tests::test_mcp_config( + codex_home.path().to_path_buf(), + )), + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + mcp_servers, + submit_id: String::new(), + tx_event: None, + startup_cancellation_token: cancel_token.clone(), + runtime_context: McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + PathBuf::from("/tmp"), + ), + codex_apps_tools_cache: ConnectorRuntimeManager::::default(), + tool_catalog_cache: McpToolCatalogCache::default(), + codex_apps_tools_cache_key: ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ), + client_mcp_extensions: ClientMcpExtensions::default(), + auth: None, + auth_manager: None, + elicitation_reviewer: None, + elicitation_lifecycle: None, + }, + ElicitationRequestRouter::default(), + ) + .await; + + assert!(manager.contains_server("stdio")); + assert!(manager.contains_server("http")); + assert!( + !manager + .wait_for_server_ready("stdio", Duration::from_millis(10)) + .await + ); + let error = match manager.test_client("stdio").client().await { + Ok(_) => panic!("local stdio MCP startup should fail"), + Err(error) => error, + }; + let StartupOutcomeError::Failed { error, .. } = error else { + panic!("local stdio MCP startup should fail rather than be cancelled"); + }; + assert_eq!( + error, + "local stdio MCP server `stdio` requires a local environment" + ); + cancel_token.cancel(); +} + +#[test] +fn elicitation_capability_uses_2025_06_18_shape_for_form_only_support() { + let capability = Some(ElicitationCapability::default()); + assert_eq!( + serde_json::to_value(capability).expect("serialize elicitation capability"), + serde_json::json!({}) + ); +} + +#[test] +fn elicitation_capability_advertises_url_support_when_enabled() { + let capability = Some( + ElicitationCapability::new() + .with_form(rmcp::model::FormElicitationCapability::new()) + .with_url(rmcp::model::UrlElicitationCapability::new()), + ); + assert_eq!( + serde_json::to_value(capability).expect("serialize elicitation capability"), + serde_json::json!({ + "form": {}, + "url": {}, + }) + ); +} + +#[test] +fn mcp_init_error_display_prompts_for_github_pat() { + let server_name = "github"; + let config = McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: "https://api.githubcopilot.com/mcp/".to_string(), + bearer_token_env_var: None, + http_headers: None, + env_http_headers: None, + http_headers_helper: None, + }, + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + }; + let err: StartupOutcomeError = anyhow::anyhow!("OAuth is unsupported").into(); + + let display = mcp_init_error_display(server_name, Some(&config), &err, /*reason*/ None); + + let expected = format!( + "GitHub MCP does not support OAuth. Log in by adding a personal access token (https://github.com/settings/personal-access-tokens) to your environment and config.toml:\n[mcp_servers.{server_name}]\nbearer_token_env_var = CODEX_GITHUB_PERSONAL_ACCESS_TOKEN" + ); + + assert_eq!(expected, display); +} + +#[test] +fn mcp_init_error_display_prompts_for_login_when_auth_required() { + let server_name = "example"; + let expected = format!( + "The {server_name} MCP server is not logged in. Run `codex mcp login {server_name}`." + ); + let executor_config: McpServerConfig = serde_json::from_value(serde_json::json!({ + "url": "https://example.com/mcp", + "environment_id": "executor-1", + })) + .expect("executor MCP configuration should deserialize"); + + for error in [ + anyhow::anyhow!("Auth required for server").into(), + StartupOutcomeError::Failed { + error: "OAuth refresh token was rejected: invalid_grant".to_string(), + is_authentication_required: true, + }, + ] { + let display = mcp_init_error_display( + server_name, + /*config*/ None, + &error, + /*reason*/ None, + ); + assert_eq!(expected, display); + + let executor_display = mcp_init_error_display( + server_name, + Some(&executor_config), + &error, + /*reason*/ None, + ); + assert_eq!( + format!( + "The {server_name} MCP server is not logged in. Use your client's MCP OAuth sign-in flow." + ), + executor_display + ); + } +} + +#[test] +fn mcp_init_error_display_identifies_oauth_reauthentication() { + let server_name = "example"; + let error = StartupOutcomeError::Failed { + error: "authorization required: Bearer error=\"invalid_token\"".to_string(), + is_authentication_required: true, + }; + let executor_config: McpServerConfig = serde_json::from_value(serde_json::json!({ + "url": "https://example.com/mcp", + "environment_id": "executor-1", + })) + .expect("executor MCP configuration should deserialize"); + + for (config, recovery_hint) in [ + (None, "Run `codex mcp login example`."), + ( + Some(&executor_config), + "Use your client's MCP OAuth sign-in flow.", + ), + ] { + assert_eq!( + mcp_init_error_display( + server_name, + config, + &error, + Some(McpStartupFailureReason::ReauthenticationRequired), + ), + format!( + "The {server_name} MCP server requires OAuth reauthentication. {recovery_hint}" + ), + ); + } +} + +#[test] +fn mcp_startup_failure_reason_requires_existing_oauth_and_auth_failure() { + for (auth_state, is_authentication_required, expected) in [ + ( + Some(McpAuthState::LoggedOut( + McpLoginRequirement::Reauthentication, + )), + true, + Some(McpStartupFailureReason::ReauthenticationRequired), + ), + ( + Some(McpAuthState::LoggedOut( + McpLoginRequirement::Reauthentication, + )), + false, + None, + ), + ( + Some(McpAuthState::LoggedOut(McpLoginRequirement::Login)), + true, + None, + ), + (Some(McpAuthState::Unsupported), true, None), + (Some(McpAuthState::BearerToken), true, None), + ( + Some(McpAuthState::OAuth), + true, + Some(McpStartupFailureReason::ReauthenticationRequired), + ), + (Some(McpAuthState::OAuth), false, None), + (None, true, None), + ] { + let error = StartupOutcomeError::Failed { + error: "startup failed".to_string(), + is_authentication_required, + }; + + assert_eq!( + mcp_startup_failure_reason(auth_state, &error), + expected, + "auth_state={auth_state:?}, is_authentication_required={is_authentication_required}" + ); + } +} + +#[test] +fn mcp_init_error_display_reports_generic_errors() { + let server_name = "custom"; + let config = McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: "https://example.com".to_string(), + bearer_token_env_var: Some("TOKEN".to_string()), + http_headers: None, + env_http_headers: None, + http_headers_helper: None, + }, + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + }; + let err: StartupOutcomeError = anyhow::anyhow!("boom").into(); + + let display = mcp_init_error_display(server_name, Some(&config), &err, /*reason*/ None); + + let expected = format!("MCP client for `{server_name}` failed to start: {err:#}"); + + assert_eq!(expected, display); +} + +#[test] +fn mcp_init_error_display_quotes_server_names() { + let github_config: McpServerConfig = serde_json::from_value(serde_json::json!({ + "url": "https://api.githubcopilot.com/mcp/", + })) + .expect("GitHub MCP configuration should deserialize"); + let error: StartupOutcomeError = anyhow::anyhow!("request timed out").into(); + let mut displays = Vec::new(); + for server_name in ["npm:@scope/package.name", "server.name"] { + for config in [None, Some(&github_config)] { + displays.push(mcp_init_error_display( + server_name, + config, + &error, + /*reason*/ None, + )); + } + } + insta::assert_snapshot!(displays.join("\n\n")); +} + +#[test] +fn mcp_init_error_display_includes_startup_timeout_hint() { + let server_name = "slow"; + for error in [ + "request timed out", + "MCP client startup timed out after 30s", + ] { + let err: StartupOutcomeError = anyhow::anyhow!(error).into(); + + let display = mcp_init_error_display( + server_name, + /*config*/ None, + &err, + /*reason*/ None, + ); + + assert_eq!( + "MCP client for `slow` timed out after 30 seconds. Add or adjust `startup_timeout_sec` in your config.toml:\n[mcp_servers.slow]\nstartup_timeout_sec = XX", + display + ); + } +} + +fn reusable_server_config(url: &str) -> McpServerConfig { + McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: url.to_string(), + bearer_token_env_var: Some("CODEX_MCP_REUSE_TEST_TOKEN".to_string()), + http_headers: None, + env_http_headers: None, + http_headers_helper: None, + }, + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + } +} + +fn reusable_server_runtime_context() -> McpRuntimeContext { + McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + PathBuf::from("/tmp"), + ) +} + +fn reusable_server_identity( + server_name: &str, + config: &McpServerConfig, + runtime_context: &McpRuntimeContext, +) -> McpServerConnectionIdentity { + let server = EffectiveMcpServer::configured(config.clone()); + let resolved_environment = runtime_context.resolve_server_environment(server_name, config); + McpServerConnectionIdentity::new( + server_name, + &server, + /*host_plugin_root*/ None, + OAuthCredentialsStoreMode::default(), + AuthKeyringBackendKind::default(), + McpOAuthRefreshMode::Legacy, + &resolved_environment, + runtime_context, + /*runtime_auth_provider*/ None, + /*auth*/ None, + /*codex_apps_cache_identity*/ None, + ElicitationCapability::default(), + ClientMcpExtensions::default(), + /*previous_identity*/ None, + ) +} + +async fn manager_with_reusable_ready_server( + config: &McpServerConfig, + runtime_context: &McpRuntimeContext, + tools: Vec, +) -> McpConnectionSet { + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut manager = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + let server = EffectiveMcpServer::configured(config.clone()); + manager.servers.insert( + "docs".to_string(), + McpServerView { + protocol_mode: crate::McpProtocolMode::Legacy, + connection: Arc::new(McpServerConnection { + identity: Some(reusable_server_identity("docs", config, runtime_context)), + client: create_ready_async_managed_client(tools).await, + startup_timeout: config + .startup_timeout_sec + .unwrap_or(DEFAULT_STARTUP_TIMEOUT), + startup_trigger: None, + _diagnostics_guard: LIVE_CONNECTIONS.track(), + }), + metadata: McpServerMetadata::from(&server), + tool_filter: ToolFilter::from_config(config), + tool_timeout: Some(config.tool_timeout_sec.unwrap_or(DEFAULT_TOOL_TIMEOUT)), + catalog_item_limit: crate::pagination::MAX_MCP_CATALOG_ITEMS, + }, + ); + manager +} + +async fn reconcile_reusable_server( + previous: &McpConnectionSet, + config: McpServerConfig, + runtime_context: McpRuntimeContext, +) -> McpConnectionSet { + let codex_home = tempdir().expect("tempdir"); + reconcile_reusable_server_with_mcp_config( + previous, + "docs", + config, + runtime_context, + crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()), + ) + .await +} + +async fn reconcile_reusable_server_with_mcp_config( + previous: &McpConnectionSet, + server_name: &str, + config: McpServerConfig, + runtime_context: McpRuntimeContext, + mcp_config: crate::McpConfig, +) -> McpConnectionSet { + let (tx_event, _rx_event) = async_channel::unbounded(); + McpConnectionSet::new( + Some(previous), + McpPublicationGate::already_published(), + McpRuntimeInput { + startup_policy: McpStartupPolicy::Eager, + config: Arc::new(mcp_config), + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + mcp_servers: HashMap::from([( + server_name.to_string(), + EffectiveMcpServer::configured(config), + )]), + submit_id: "refresh".to_string(), + tx_event: Some(tx_event), + startup_cancellation_token: CancellationToken::new(), + runtime_context, + codex_apps_tools_cache: ConnectorRuntimeManager::default(), + tool_catalog_cache: McpToolCatalogCache::default(), + codex_apps_tools_cache_key: ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ), + client_mcp_extensions: ClientMcpExtensions::default(), + auth: None, + auth_manager: None, + elicitation_reviewer: None, + elicitation_lifecycle: None, + }, + ElicitationRequestRouter::default(), + ) + .await +} + +#[tokio::test] +async fn apps_catalog_broadcast_preserves_running_calls_and_rejects_stale_calls() +-> anyhow::Result<()> { + let codex_home = tempdir()?; + let context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + /*account_id*/ None, + /*chatgpt_user_id*/ None, + ) + .with_live_scope("apps".to_string()); + let original = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "original")]; + let updated = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "updated")]; + store_current_tools(&context, original.clone()); + let catalog = ClientToolCatalog::new(original.clone(), context.subscribe()); + let snapshot = catalog.read(Arc::new).await; + + store_current_tools(&context, original.clone()); + assert_eq!( + catalog.read(|catalog| catalog.revision).await, + 0, + "an unchanged startup broadcast must not invalidate prepared calls" + ); + + let (release, released) = tokio::sync::oneshot::channel::<()>(); + let running = catalog.run_with_snapshot(&snapshot, || async { released.await.unwrap() }); + tokio::pin!(running); + assert!(futures::poll!(&mut running).is_pending()); + store_current_tools(&context, updated.clone()); + let stale = catalog.run_with_snapshot(&snapshot, || async { panic!("stale preparation") }); + tokio::pin!(stale); + assert!( + futures::poll!(&mut stale).is_pending(), + "publication does not wait, but adoption must wait for the running call" + ); + release.send(()).unwrap(); + assert_eq!(running.await, Some(())); + assert_eq!(stale.await, None::<()>); + assert_eq!( + catalog + .read(|catalog| (catalog.revision, catalog.tools.to_vec())) + .await, + (1, updated.clone()) + ); + + // A client whose startup finishes late adopts the already-published result. + let late = ClientToolCatalog::new(original, context.subscribe()); + assert_eq!(late.read(|catalog| catalog.tools.to_vec()).await, updated); + Ok(()) +} + +#[tokio::test] +async fn apps_catalog_broadcast_restores_equivalent_calls() -> anyhow::Result<()> { + let codex_home = tempdir()?; + let context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + /*account_id*/ None, + /*chatgpt_user_id*/ None, + ) + .with_live_scope("apps".to_string()); + let original = vec![ + create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "first"), + create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "second"), + ]; + store_current_tools(&context, original.clone()); + let catalog = ClientToolCatalog::new(original.clone(), context.subscribe()); + let snapshot = catalog.read(Arc::new).await; + + let mut changed = original.clone(); + changed[0].tool.description = Some("Changed definition".into()); + store_current_tools(&context, changed); + assert_eq!( + catalog + .run_with_snapshot(&snapshot, || async { panic!("changed preparation") }) + .await, + None::<()> + ); + + let mut restored = original; + restored.reverse(); + store_current_tools(&context, restored); + assert_eq!( + catalog + .run_with_snapshot(&snapshot, || async { "prepared" }) + .await, + Some("prepared") + ); + catalog + .refresh( + || async { Ok((snapshot.tools.to_vec(), ())) }, + |tools, ()| store_current_tools(&context, tools.to_vec()), + ) + .await?; + assert_eq!( + catalog + .run_with_snapshot(&snapshot, || async { panic!("refreshed preparation") }) + .await, + None::<()> + ); + Ok(()) +} + +#[test] +fn idle_apps_clients_do_not_retain_replaced_tools() { + let context = ConnectorRuntimeManager::::new_without_cache() + .context( + PathBuf::from("unused"), + ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ), + ) + .with_live_scope("apps".into()); + let tool = create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "original"); + let schema = Arc::downgrade(&tool.tool.input_schema); + store_current_tools(&context, vec![tool]); + let clients = [ + ClientToolCatalog::new(Vec::new(), context.subscribe()), + ClientToolCatalog::new(Vec::new(), context.subscribe()), + ]; + store_current_tools( + &context, + vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "updated")], + ); + assert!( + schema.upgrade().is_none(), + "idle clients must not own the old schema" + ); + drop(clients); +} + +#[tokio::test] +async fn apps_catalog_broadcast_survives_an_older_local_refresh() -> anyhow::Result<()> { + let codex_home = tempdir()?; + let context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + /*account_id*/ None, + /*chatgpt_user_id*/ None, + ) + .with_live_scope("apps".to_string()); + store_current_tools(&context, Vec::new()); + let catalog = ClientToolCatalog::new(Vec::new(), context.subscribe()); + let older_ticket = context.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh); + let newer = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "newer")]; + store_current_tools(&context, newer.clone()); + assert_eq!(catalog.read(|catalog| catalog.tools.to_vec()).await, newer); + catalog + .refresh( + || async { Ok((Vec::new(), older_ticket)) }, + |tools, ticket| { + context.publish_if_newest_accepted( + ticket, + &create_test_server_info("Apps"), + tools.to_vec(), + ) + }, + ) + .await?; + assert_eq!(catalog.read(|catalog| catalog.tools.to_vec()).await, newer); + Ok(()) +} + +#[tokio::test] +async fn refreshed_catalog_follows_reused_client_without_mutating_old_bindings() +-> anyhow::Result<()> { + let codex_home = tempdir()?; + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + /*account_id*/ None, + /*chatgpt_user_id*/ None, + ); + let runtime_context = reusable_server_runtime_context(); + let config = reusable_server_config("http://127.0.0.1:1"); + let mut previous = create_test_manager_with_ready_apps_client( + cache_context, + "refreshed", + /*list_started*/ None, + /*release_list*/ None, + ) + .await?; + let manager = Arc::get_mut(&mut previous).expect("unshared manager"); + manager.insert_test_client( + "docs", + create_ready_async_managed_client(vec![create_test_tool("docs", "unrelated")]).await, + ); + let connection = Arc::get_mut( + &mut manager + .servers + .get_mut(CODEX_APPS_MCP_SERVER_NAME) + .expect("Apps server") + .connection, + ) + .expect("unshared Apps connection"); + connection.identity = Some(reusable_server_identity( + CODEX_APPS_MCP_SERVER_NAME, + &config, + &runtime_context, + )); + let client = connection.client().await?; + // Seed an earlier catalog before capturing the binding under test. + let startup_tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "startup")]; + client + .tool_catalog + .refresh(|| async { Ok((startup_tools, ())) }, |_, ()| {}) + .await?; + let old_binding = capture_binding(&previous).await; + let old_tools = serde_json::to_value(old_binding.tools())?; + let old_call = old_binding + .prepare_call(CODEX_APPS_MCP_SERVER_NAME, "startup") + .expect("startup call"); + let unrelated_call = old_binding + .prepare_call("docs", "unrelated") + .expect("unrelated call"); + + previous.refresh_codex_apps_tools_for_discovery().await?; + let republished = Arc::new( + reconcile_reusable_server_with_mcp_config( + &previous, + CODEX_APPS_MCP_SERVER_NAME, + config, + runtime_context, + crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()), + ) + .await, + ); + assert!(previous.shares_test_connection_with(&republished, CODEX_APPS_MCP_SERVER_NAME)); + assert_eq!( + model_tool_names(capture_binding(&republished).await.tools()), + HashSet::from([ToolName::namespaced("mcp__codex_apps", "refreshed")]), + ); + assert_eq!(serde_json::to_value(old_binding.tools())?, old_tools); + + let error = old_call + .call_with_preparation(/*requested_timeout*/ None, || async { + panic!("stale call preparation must not run"); + }) + .await + .expect_err("a call from the old catalog must be rejected"); + assert!(error.to_string().contains("catalog changed")); + let error = unrelated_call + .call_with_preparation(/*requested_timeout*/ None, || async { + Err(anyhow!("unrelated preparation reached")) + }) + .await + .expect_err("stop before the unrelated tool executes"); + assert!( + error.to_string().contains("unrelated preparation reached"), + "Apps refresh must not invalidate another client's calls" + ); + + Ok(()) +} + +#[tokio::test] +async fn reconciliation_reuses_connection_without_relisting_regular_tools() -> anyhow::Result<()> { + let tools = Arc::new(tokio::sync::RwLock::new(vec![Tool::new( + "old_search", + "old search", + Arc::new(JsonObject::default()), + )])); + let block_tool_listing = Arc::new(AtomicBool::new(false)); + let client = Arc::new( + RmcpClient::new_in_process_client(Arc::new(MutableToolsTransportFactory { + server: MutableToolsServer { + tools: Arc::clone(&tools), + block_tool_listing: Arc::clone(&block_tool_listing), + }, + })) + .await?, + ); + let initialize = client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("codex-test", "0.0.0-test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + /*timeout*/ None, + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }) + } + .boxed() + }), + ) + .await?; + let initial_tools = list_tools_for_client_uncached( + "docs", + /*is_codex_apps_mcp_server*/ false, + /*codex_apps_refresh_trigger*/ "test", + &client, + /*timeout*/ None, + crate::pagination::MAX_MCP_CATALOG_ITEMS, + initialize.instructions.as_deref(), + ) + .await?; + let managed_client = ManagedClient { + _auth_change_notifications: None, + client, + server_info: create_test_server_info("Mutable tools"), + tool_catalog: Arc::new(ClientToolCatalog::new(initial_tools, /*updates*/ None)), + tool_timeout: None, + server_instructions: initialize.instructions, + server_supports_sandbox_state_meta_capability: false, + codex_apps_tools_cache_context: None, + }; + let runtime_context = reusable_server_runtime_context(); + let config = reusable_server_config("http://127.0.0.1:1"); + let approval_policy = Constrained::allow_any(AskForApproval::OnRequest); + let permission_profile = Constrained::allow_any(PermissionProfile::default()); + let mut previous = McpConnectionSet::new_uninitialized( + &approval_policy, + &permission_profile, + /*prefix_mcp_tool_names*/ true, + ); + let server = EffectiveMcpServer::configured(config.clone()); + previous.servers.insert( + "docs".to_string(), + McpServerView { + protocol_mode: crate::McpProtocolMode::Legacy, + connection: Arc::new(McpServerConnection { + identity: Some(reusable_server_identity("docs", &config, &runtime_context)), + client: AsyncManagedClient { + client: futures::future::ready(Ok(managed_client)).boxed().shared(), + is_codex_apps_mcp_server: false, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(true)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + startup_timeout: config + .startup_timeout_sec + .unwrap_or(DEFAULT_STARTUP_TIMEOUT), + startup_trigger: None, + _diagnostics_guard: LIVE_CONNECTIONS.track(), + }), + metadata: McpServerMetadata::from(&server), + tool_filter: ToolFilter::from_config(&config), + tool_timeout: Some(config.tool_timeout_sec.unwrap_or(DEFAULT_TOOL_TIMEOUT)), + catalog_item_limit: crate::pagination::MAX_MCP_CATALOG_ITEMS, + }, + ); + let previous = Arc::new(previous); + let old_step = capture_binding(&previous).await; + *tools.write().await = vec![Tool::new( + "new_search", + "new search", + Arc::new(JsonObject::default()), + )]; + block_tool_listing.store(true, Ordering::Release); + + let reconciled = Arc::new( + tokio::time::timeout( + Duration::from_secs(1), + reconcile_reusable_server(&previous, config, runtime_context), + ) + .await + .expect("connection reuse must not wait for a tool-list request"), + ); + let new_step = capture_binding(&reconciled).await; + + assert!(previous.shares_test_connection_with(&reconciled, "docs")); + assert_eq!( + old_step + .tools() + .iter() + .map(|tool| tool.tool.name.to_string()) + .collect::>(), + vec!["old_search".to_string()] + ); + assert_eq!( + new_step + .tools() + .iter() + .map(|tool| tool.tool.name.to_string()) + .collect::>(), + vec!["old_search".to_string()] + ); + Ok(()) +} + +#[tokio::test] +async fn reconciliation_reuses_an_unchanged_ready_server() { + let runtime_context = reusable_server_runtime_context(); + let config = reusable_server_config("http://127.0.0.1:1"); + let previous = manager_with_reusable_ready_server( + &config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + + let reconciled = reconcile_reusable_server(&previous, config, runtime_context.clone()).await; + + assert!(previous.shares_test_connection_with(&reconciled, "docs")); + assert_eq!( + model_tool_names(&reconciled.list_all_tools().await), + HashSet::from([ToolName::namespaced("mcp__docs", "search")]) + ); +} + +#[tokio::test] +async fn reconciliation_reuses_an_unchanged_pending_server_without_waiting() -> anyhow::Result<()> { + let runtime_context = reusable_server_runtime_context(); + let mut config = reusable_server_config("http://127.0.0.1:1"); + let tools = vec![ + create_test_tool("docs", "search"), + create_test_tool("docs", "write"), + ]; + let mut previous = + manager_with_reusable_ready_server(&config, &runtime_context, tools.clone()).await; + let managed_client = create_test_managed_client(tools).await; + let (pending_client, startup_started, release_startup) = + create_gated_async_managed_client(managed_client); + let startup = tokio::spawn({ + let pending_client = pending_client.clone(); + async move { pending_client.client().await } + }); + startup_started.await?; + let connection = Arc::get_mut( + &mut previous + .servers + .get_mut("docs") + .expect("test server should exist") + .connection, + ) + .expect("test server should have one connection owner"); + connection.client = pending_client; + config.enabled_tools = Some(vec!["search".to_string()]); + config.startup_timeout_sec = Some(DEFAULT_STARTUP_TIMEOUT); + + let reconciled = tokio::time::timeout( + Duration::from_millis(100), + reconcile_reusable_server(&previous, config, runtime_context), + ) + .await + .expect("reconciliation must not wait for an unchanged pending MCP server"); + + assert!(previous.shares_test_connection_with(&reconciled, "docs")); + release_startup + .send(()) + .map_err(|()| anyhow!("pending startup should still be running"))?; + startup.await??; + assert_eq!( + model_tool_names(&reconciled.list_all_tools().await), + HashSet::from([ToolName::namespaced("mcp__docs", "search")]) + ); + Ok(()) +} + +#[tokio::test] +async fn reconciliation_cancels_a_reused_pending_server_when_disabled() -> anyhow::Result<()> { + let runtime_context = reusable_server_runtime_context(); + let mut config = reusable_server_config("http://127.0.0.1:1"); + let tools = vec![create_test_tool("docs", "search")]; + let mut previous = + manager_with_reusable_ready_server(&config, &runtime_context, tools.clone()).await; + let managed_client = create_test_managed_client(tools).await; + let (pending_client, startup_started, release_startup) = + create_gated_async_managed_client(managed_client); + let cancellation = pending_client.cancel_token.clone(); + let startup = tokio::spawn({ + let pending_client = pending_client.clone(); + async move { pending_client.client().await } + }); + startup_started.await?; + let connection = Arc::get_mut( + &mut previous + .servers + .get_mut("docs") + .expect("test server should exist") + .connection, + ) + .expect("test server should have one connection owner"); + connection.client = pending_client; + + let reused = + reconcile_reusable_server(&previous, config.clone(), runtime_context.clone()).await; + assert!(previous.shares_test_connection_with(&reused, "docs")); + + config.enabled = false; + let removed = reconcile_reusable_server(&reused, config, runtime_context).await; + assert!(!removed.servers.contains_key("docs")); + drop(previous); + drop(reused); + + assert!( + cancellation.is_cancelled(), + "disabling a reused pending MCP server should cancel its obsolete startup" + ); + release_startup + .send(()) + .map_err(|()| anyhow!("pending startup should remain available for test cleanup"))?; + startup.await??; + Ok(()) +} + +#[tokio::test] +async fn reconciliation_retries_non_oauth_authentication_failures() { + let runtime_context = reusable_server_runtime_context(); + let config = reusable_server_config("http://127.0.0.1:1"); + let mut previous = + manager_with_reusable_ready_server(&config, &runtime_context, Vec::new()).await; + let connection = + Arc::get_mut(&mut previous.servers.get_mut("docs").expect("server").connection) + .expect("test server has one connection owner"); + connection.client.client = futures::future::ready(Err(StartupOutcomeError::Failed { + error: "bearer token rejected".to_string(), + is_authentication_required: true, + })) + .boxed() + .shared(); + + let reconciled = reconcile_reusable_server(&previous, config, runtime_context).await; + + assert!(!previous.shares_test_connection_with(&reconciled, "docs")); +} + +#[test] +fn connection_identity_uses_effective_authorization_headers() { + let runtime_context = reusable_server_runtime_context(); + let missing_env_var = format!("CODEX_TEST_UNSET_MCP_AUTHORIZATION_{}", std::process::id()); + assert!(std::env::var_os(&missing_env_var).is_none()); + + for (static_header, environment_header, has_authorization) in [ + (Some("Bearer configured-token"), None, true), + (Some("invalid\nheader"), None, false), + (None, Some("PATH"), true), + (None, Some(missing_env_var.as_str()), false), + (None, None, false), + ] { + let mut config = reusable_server_config("http://127.0.0.1:1"); + config.transport = McpServerTransportConfig::StreamableHttp { + url: "http://127.0.0.1:1".to_string(), + bearer_token_env_var: None, + http_headers: static_header + .map(|value| HashMap::from([("aUtHoRiZaTiOn".to_string(), value.to_string())])), + env_http_headers: environment_header + .map(|value| HashMap::from([("aUtHoRiZaTiOn".to_string(), value.to_string())])), + http_headers_helper: None, + }; + let server = EffectiveMcpServer::configured(config); + let identity = |keyring_backend_kind, oauth_refresh_mode| { + McpServerConnectionIdentity::new( + "docs", + &server, + /*host_plugin_root*/ None, + OAuthCredentialsStoreMode::File, + keyring_backend_kind, + oauth_refresh_mode, + &Ok(None), + &runtime_context, + /*runtime_auth_provider*/ None, + /*auth*/ None, + /*codex_apps_cache_identity*/ None, + ElicitationCapability::default(), + ClientMcpExtensions::default(), + /*previous_identity*/ None, + ) + }; + + assert_eq!( + identity(AuthKeyringBackendKind::Direct, McpOAuthRefreshMode::Legacy) + .has_same_connection_config(&identity( + AuthKeyringBackendKind::Secrets, + McpOAuthRefreshMode::Legacy, + )), + has_authorization, + ); + assert_eq!( + identity(AuthKeyringBackendKind::Direct, McpOAuthRefreshMode::Legacy) + .has_same_connection_config(&identity( + AuthKeyringBackendKind::Direct, + McpOAuthRefreshMode::Coordinated, + )), + has_authorization, + ); + } +} + +#[tokio::test] +async fn reconciliation_reuses_legacy_stdio_server_with_existing_protocol_marker() { + let runtime_context = McpRuntimeContext::new( + Arc::new(codex_exec_server::EnvironmentManager::default_for_tests()), + PathBuf::from("/tmp"), + ); + let mut config = reusable_server_config("http://127.0.0.1:1"); + config.transport = McpServerTransportConfig::Stdio { + command: "legacy-server".to_string(), + args: Vec::new(), + env: Some(HashMap::from([( + "CODEX_MCP_PROTOCOL_VERSION".to_string(), + "1999-01-01".to_string(), + )])), + env_vars: Vec::new(), + cwd: None, + }; + let previous = manager_with_reusable_ready_server( + &config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + + let reconciled = reconcile_reusable_server(&previous, config, runtime_context).await; + + assert!(previous.shares_test_connection_with(&reconciled, "docs")); +} + +#[tokio::test] +async fn reconciliation_replaces_connection_when_auth_mode_changes() -> anyhow::Result<()> { + let environment_manager = Arc::new(environment_manager_without_environments()); + environment_manager.upsert_environment( + "customer-executor".to_string(), + "ws://127.0.0.1:1".to_string(), + /*connect_timeout*/ None, + )?; + let runtime_context = McpRuntimeContext::new(environment_manager, PathBuf::from("/tmp")); + let codex_home = tempdir()?; + let mcp_config = crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()); + let [config, refreshed_config] = [McpServerAuth::OAuth, McpServerAuth::ChatGpt].map(|auth| { + let mut config = reusable_server_config("https://chatgpt.com/backend-api/ps/mcp"); + config.environment_id = "customer-executor".to_string(); + config.auth = auth; + crate::effective_mcp_servers_from_configured( + HashMap::from([("docs".to_string(), config)]), + &mcp_config, + /*auth*/ None, + ) + .remove("docs") + .expect("configured server should survive auth projection") + .config() + .clone() + }); + let previous = manager_with_reusable_ready_server( + &config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + + let reconciled = reconcile_reusable_server(&previous, refreshed_config, runtime_context).await; + let outcome = reconciled + .servers + .get("docs") + .expect("refreshed server should exist") + .connection + .client() + .await; + assert_matches!( + outcome.err().expect("changed auth mode must be validated"), + StartupOutcomeError::Failed { + error, + is_authentication_required, + } => { + assert_eq!( + (error.as_str(), is_authentication_required), + ( + "executor-owned MCP server `docs` cannot use hosted ChatGPT authentication; configure executor-owned credentials instead", + false, + ) + ); + } + ); + assert_eq!( + model_tool_names(&reconciled.list_all_tools().await), + HashSet::new() + ); + Ok(()) +} + +#[tokio::test] +async fn reconciliation_replaces_connection_when_protocol_mode_changes() { + let runtime_context = reusable_server_runtime_context(); + let config = reusable_server_config("http://127.0.0.1:1"); + let previous = manager_with_reusable_ready_server( + &config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + let codex_home = tempdir().expect("tempdir"); + let mut mcp_config = crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()); + mcp_config.protocol_mode = codex_rmcp_client::McpProtocolMode::V20260728; + + let reconciled = McpConnectionSet::new( + Some(&previous), + McpPublicationGate::already_published(), + McpRuntimeInput { + startup_policy: McpStartupPolicy::Eager, + config: Arc::new(mcp_config), + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + mcp_servers: HashMap::from([( + "docs".to_string(), + EffectiveMcpServer::configured(config), + )]), + submit_id: "refresh".to_string(), + tx_event: None, + startup_cancellation_token: CancellationToken::new(), + runtime_context, + codex_apps_tools_cache: ConnectorRuntimeManager::default(), + tool_catalog_cache: McpToolCatalogCache::default(), + codex_apps_tools_cache_key: ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ), + client_mcp_extensions: ClientMcpExtensions::default(), + auth: None, + auth_manager: None, + elicitation_reviewer: None, + elicitation_lifecycle: None, + }, + ElicitationRequestRouter::default(), + ) + .await; + + assert!(!previous.shares_test_connection_with(&reconciled, "docs")); +} + +#[tokio::test] +async fn reconciliation_reuses_legacy_stdio_server_when_modern_protocol_is_enabled() { + let runtime_context = McpRuntimeContext::new( + Arc::new(codex_exec_server::EnvironmentManager::default_for_tests()), + PathBuf::from("/tmp"), + ); + let mut config = reusable_server_config("http://127.0.0.1:1"); + config.transport = McpServerTransportConfig::Stdio { + command: "legacy-server".to_string(), + args: Vec::new(), + env: None, + env_vars: Vec::new(), + cwd: None, + }; + let previous = manager_with_reusable_ready_server( + &config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + let codex_home = tempdir().expect("tempdir"); + let mut mcp_config = crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()); + mcp_config.protocol_mode = codex_rmcp_client::McpProtocolMode::V20260728; + + let reconciled = McpConnectionSet::new( + Some(&previous), + McpPublicationGate::already_published(), + McpRuntimeInput { + startup_policy: McpStartupPolicy::Eager, + config: Arc::new(mcp_config), + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + mcp_servers: HashMap::from([( + "docs".to_string(), + EffectiveMcpServer::configured(config), + )]), + submit_id: "refresh".to_string(), + tx_event: None, + startup_cancellation_token: CancellationToken::new(), + runtime_context, + codex_apps_tools_cache: ConnectorRuntimeManager::default(), + tool_catalog_cache: McpToolCatalogCache::default(), + codex_apps_tools_cache_key: ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ), + client_mcp_extensions: ClientMcpExtensions::default(), + auth: None, + auth_manager: None, + elicitation_reviewer: None, + elicitation_lifecycle: None, + }, + ElicitationRequestRouter::default(), + ) + .await; + + assert!(previous.shares_test_connection_with(&reconciled, "docs")); +} + +#[tokio::test] +async fn reconciliation_updates_elicitation_policy_without_restarting_ready_server() { + let runtime_context = reusable_server_runtime_context(); + let config = reusable_server_config("http://127.0.0.1:1"); + let previous = manager_with_reusable_ready_server( + &config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + { + let mut authority = previous + .elicitation_requests + .authority + .lock() + .expect("elicitation authority lock"); + let config = Arc::make_mut( + &mut authority + .as_mut() + .expect("test manager should have permission authority") + .config, + ); + config.approval_policy = Constrained::allow_any(AskForApproval::Never); + config.permission_profile = PermissionProfile::Disabled; + } + + let reconciled = reconcile_reusable_server(&previous, config, runtime_context).await; + + assert!(previous.shares_test_connection_with(&reconciled, "docs")); + let authority = reconciled + .elicitation_requests + .authority + .lock() + .expect("elicitation authority lock"); + let config = &authority + .as_ref() + .expect("reconciled manager should have permission authority") + .config; + assert_eq!(config.approval_policy.value(), AskForApproval::OnRequest); + assert_eq!(config.permission_profile, PermissionProfile::default()); +} + +#[tokio::test] +async fn reconciliation_reuses_ready_server_when_startup_timeout_changes() { + let runtime_context = reusable_server_runtime_context(); + let mut config = reusable_server_config("http://127.0.0.1:1"); + let previous = manager_with_reusable_ready_server( + &config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + config.startup_timeout_sec = Some(Duration::from_secs(60)); + + let reconciled = reconcile_reusable_server(&previous, config, runtime_context).await; + + assert_eq!( + model_tool_names(&reconciled.list_all_tools().await), + HashSet::from([ToolName::namespaced("mcp__docs", "search")]) + ); +} + +#[tokio::test] +async fn reconciliation_replaces_closed_connections() -> anyhow::Result<()> { + let runtime_context = reusable_server_runtime_context(); + let config = reusable_server_config("http://127.0.0.1:1"); + let mut previous = manager_with_reusable_ready_server( + &config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + let disconnect = CancellationToken::new(); + let client = Arc::new( + RmcpClient::new_in_process_client(Arc::new(DisconnectingToolsTransportFactory { + server: MutableToolsServer { + tools: Arc::new(tokio::sync::RwLock::new(vec![Tool::new( + "search", + "search", + Arc::new(JsonObject::default()), + )])), + block_tool_listing: Arc::new(AtomicBool::new(false)), + }, + disconnect: disconnect.clone(), + })) + .await?, + ); + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("codex-test", "0.0.0-test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + /*timeout*/ None, + Box::new(|_, _| async { Err(anyhow!("unexpected elicitation")) }.boxed()), + ) + .await?; + let view = previous + .servers + .get_mut("docs") + .expect("test server should exist"); + let mut connected_client = view.connection.client().await?; + connected_client.client = Arc::clone(&client); + view.connection = Arc::new(McpServerConnection { + identity: Some(reusable_server_identity("docs", &config, &runtime_context)), + client: AsyncManagedClient { + client: futures::future::ready(Ok(connected_client)) + .boxed() + .shared(), + is_codex_apps_mcp_server: false, + server_capabilities: Arc::new(std::sync::Mutex::new(None)), + cached_server_info: None, + codex_apps_tools_cache_context: None, + tool_catalog_cache_context: None, + startup_complete: Arc::new(std::sync::atomic::AtomicBool::new(true)), + startup_reconnect: None, + cancel_token: CancellationToken::new(), + }, + startup_timeout: config + .startup_timeout_sec + .unwrap_or(DEFAULT_STARTUP_TIMEOUT), + startup_trigger: None, + _diagnostics_guard: LIVE_CONNECTIONS.track(), + }); + + assert!(!client.is_closed().await); + disconnect.cancel(); + tokio::time::timeout(Duration::from_secs(2), async { + while !client.is_closed().await { + tokio::task::yield_now().await; + } + }) + .await + .expect("closed MCP transport should be detected"); + + let reconciled = reconcile_reusable_server(&previous, config, runtime_context).await; + + assert!(!previous.shares_test_connection_with(&reconciled, "docs")); + Ok(()) +} + +#[tokio::test] +async fn reconciliation_reconnects_when_connection_identity_changes() { + let runtime_context = reusable_server_runtime_context(); + let previous_config = reusable_server_config("http://127.0.0.1:1"); + let previous = manager_with_reusable_ready_server( + &previous_config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + + let reconciled = reconcile_reusable_server( + &previous, + reusable_server_config("http://127.0.0.1:2"), + runtime_context, + ) + .await; + + assert!(!previous.shares_test_connection_with(&reconciled, "docs")); +} + +#[tokio::test] +async fn reconciliation_reconnects_when_host_plugin_root_changes() { + let runtime_context = reusable_server_runtime_context(); + let server_config = reusable_server_config("http://127.0.0.1:1"); + let original_root = PathUri::parse("file:///plugins/original").expect("valid plugin root URI"); + let replacement_root = + PathUri::parse("file:///plugins/replacement").expect("valid plugin root URI"); + let mut previous = manager_with_reusable_ready_server( + &server_config, + &runtime_context, + vec![create_test_tool("docs", "search")], + ) + .await; + let server = EffectiveMcpServer::configured(server_config.clone()); + let resolved_environment = runtime_context.resolve_server_environment("docs", &server_config); + let original_identity = McpServerConnectionIdentity::new( + "docs", + &server, + Some(&original_root), + OAuthCredentialsStoreMode::default(), + AuthKeyringBackendKind::default(), + McpOAuthRefreshMode::Legacy, + &resolved_environment, + &runtime_context, + /*runtime_auth_provider*/ None, + /*auth*/ None, + /*codex_apps_cache_identity*/ None, + ElicitationCapability::default(), + ClientMcpExtensions::default(), + /*previous_identity*/ None, + ); + Arc::get_mut( + &mut previous + .servers + .get_mut("docs") + .expect("test server should exist") + .connection, + ) + .expect("test server should have one connection owner") + .identity = Some(original_identity); + + let codex_home = tempdir().expect("tempdir"); + let config_for_root = |root| { + let mut config = crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()); + let mut catalog = crate::ResolvedMcpCatalog::builder(); + catalog.register(crate::McpServerRegistration::from_plugin( + "docs".to_string(), + crate::McpPluginAttribution::new("docs@test".to_string(), "Docs".to_string()) + .with_host_root(root), + /*plugin_order*/ 0, + server_config.clone(), + )); + config.mcp_server_catalog = catalog.build(); + config + }; + + let unchanged = reconcile_reusable_server_with_mcp_config( + &previous, + "docs", + server_config.clone(), + runtime_context.clone(), + config_for_root(original_root), + ) + .await; + assert!(previous.shares_test_connection_with(&unchanged, "docs")); + + let replacement_config = config_for_root(replacement_root); + let replacement = reconcile_reusable_server_with_mcp_config( + &unchanged, + "docs", + server_config, + runtime_context, + replacement_config, + ) + .await; + assert!(!unchanged.shares_test_connection_with(&replacement, "docs")); +} + +#[tokio::test] +async fn connection_identity_distinguishes_accounts_with_the_same_token() -> anyhow::Result<()> { + let runtime_context = reusable_server_runtime_context(); + let config = reusable_server_config("http://127.0.0.1:1"); + let server = EffectiveMcpServer::configured(config); + let access_token = "header.e30.same"; + let previous_auth = CodexAuth::from_external_chatgpt_tokens( + access_token, + "account-a", + /*chatgpt_plan_type*/ None, + )?; + let changed_auth = CodexAuth::from_external_chatgpt_tokens( + access_token, + "account-b", + /*chatgpt_plan_type*/ None, + )?; + let connection_identity = |auth: &CodexAuth| { + let provider = codex_model_provider::auth_provider_from_auth(auth); + McpServerConnectionIdentity::new( + "docs", + &server, + /*host_plugin_root*/ None, + OAuthCredentialsStoreMode::default(), + AuthKeyringBackendKind::default(), + McpOAuthRefreshMode::Legacy, + &Ok(None), + &runtime_context, + Some(&provider), + Some(auth), + /*codex_apps_cache_identity*/ None, + ElicitationCapability::default(), + ClientMcpExtensions::default(), + /*previous_identity*/ None, + ) + }; + + assert_eq!(previous_auth, changed_auth); + assert_eq!(previous_auth.get_token()?, changed_auth.get_token()?); + assert!( + !connection_identity(&previous_auth) + .has_same_connection_config(&connection_identity(&changed_auth)) + ); + Ok(()) +} + +#[tokio::test] +async fn connection_identity_distinguishes_agent_account_runtime_and_task() -> anyhow::Result<()> { + let runtime_context = reusable_server_runtime_context(); + let config = reusable_server_config("http://127.0.0.1:1"); + let server = EffectiveMcpServer::configured(config); + let record = codex_login::auth::AgentIdentityAuthRecord { + agent_runtime_id: "agent-a".to_string(), + agent_private_key: "MC4CAQAwBQYDK2VwBCIEIJ7kFBaOujmoz1gvBNEC+BeM2IX87FFB0xmISOZ/XO0c" + .to_string(), + account_id: "account-a".to_string(), + chatgpt_user_id: "user-a".to_string(), + email: Some("agent@example.com".to_string()), + plan_type: codex_protocol::account::PlanType::Plus, + chatgpt_account_is_fedramp: false, + task_id: Some("task-a".to_string()), + }; + let auth_route_config = codex_login::test_support::transport_default_auth_route_config(); + let previous_auth = CodexAuth::AgentIdentity( + codex_login::auth::AgentIdentityAuth::from_record( + record.clone(), + "https://auth.openai.com/api/accounts", + &auth_route_config, + ) + .await?, + ); + let connection_identity = |auth: &CodexAuth| { + let provider = codex_model_provider::auth_provider_from_auth(auth); + McpServerConnectionIdentity::new( + CODEX_APPS_MCP_SERVER_NAME, + &server, + /*host_plugin_root*/ None, + OAuthCredentialsStoreMode::default(), + AuthKeyringBackendKind::default(), + McpOAuthRefreshMode::Legacy, + &Ok(None), + &runtime_context, + Some(&provider), + Some(auth), + /*codex_apps_cache_identity*/ None, + ElicitationCapability::default(), + ClientMcpExtensions::default(), + /*previous_identity*/ None, + ) + }; + let previous_identity = connection_identity(&previous_auth); + + for changed_record in [ + codex_login::auth::AgentIdentityAuthRecord { + account_id: "account-b".to_string(), + ..record.clone() + }, + codex_login::auth::AgentIdentityAuthRecord { + chatgpt_user_id: "user-b".to_string(), + ..record.clone() + }, + codex_login::auth::AgentIdentityAuthRecord { + chatgpt_account_is_fedramp: true, + ..record.clone() + }, + codex_login::auth::AgentIdentityAuthRecord { + agent_runtime_id: "agent-b".to_string(), + ..record.clone() + }, + codex_login::auth::AgentIdentityAuthRecord { + task_id: Some("task-b".to_string()), + ..record.clone() + }, + ] { + let changed_auth = CodexAuth::AgentIdentity( + codex_login::auth::AgentIdentityAuth::from_record( + changed_record, + "https://auth.openai.com/api/accounts", + &auth_route_config, + ) + .await?, + ); + assert_eq!(previous_auth, changed_auth); + assert!(!previous_identity.has_same_connection_config(&connection_identity(&changed_auth))); + } + + Ok(()) +} + +#[tokio::test] +async fn view_only_changes_reuse_connection_and_preserve_the_old_step() { + let runtime_context = reusable_server_runtime_context(); + let mut old_config = reusable_server_config("http://127.0.0.1:1"); + old_config.default_tools_approval_mode = Some(AppToolApproval::Prompt); + let previous = Arc::new( + manager_with_reusable_ready_server( + &old_config, + &runtime_context, + vec![ + create_test_tool("docs", "search"), + create_test_tool("docs", "write"), + ], + ) + .await, + ); + let old_step = capture_binding(&previous).await; + let old_call = old_step + .prepare_call("docs", "search") + .expect("old step should prepare search"); + + let mut new_config = old_config; + new_config.enabled_tools = Some(vec!["search".to_string()]); + new_config.default_tools_approval_mode = Some(AppToolApproval::Approve); + let reconciled = + Arc::new(reconcile_reusable_server(previous.as_ref(), new_config, runtime_context).await); + assert!(previous.shares_test_connection_with(&reconciled, "docs")); + + let new_step = capture_binding(&reconciled).await; + let new_call = new_step + .prepare_call("docs", "search") + .expect("new step should prepare search"); + drop(previous); + + assert_eq!( + old_step + .tools() + .iter() + .map(|tool| tool.tool.name.to_string()) + .collect::>(), + HashSet::from(["search".to_string(), "write".to_string()]) + ); + assert_eq!(old_call.tool_approval_mode(), AppToolApproval::Prompt); + assert_eq!( + new_step + .tools() + .iter() + .map(|tool| tool.tool.name.to_string()) + .collect::>(), + vec!["search".to_string()] + ); + assert_eq!(new_call.tool_approval_mode(), AppToolApproval::Approve); +} diff --git a/codex-rs/codex-mcp/src/elicitation.rs b/codex-rs/codex-mcp/src/elicitation.rs new file mode 100644 index 0000000000000000000000000000000000000000..ec3d3f0d0fbf66874857377b7f2ab546a4f2eca3 --- /dev/null +++ b/codex-rs/codex-mcp/src/elicitation.rs @@ -0,0 +1,548 @@ +//! MCP elicitation request tracking and policy handling. +//! +//! RMCP clients call into this module when a server asks Codex to elicit data +//! from the user. It decides whether the request can be automatically accepted, +//! must be declined by policy, or should be surfaced as a Codex protocol event +//! and later resolved through the stored responder. + +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; + +use crate::McpConfig; +use crate::mcp::McpPermissionPromptAutoApproveContext; +use crate::mcp::mcp_permission_prompt_is_auto_approved; +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use async_channel::Sender; +use codex_protocol::approvals::ElicitationRequest; +use codex_protocol::approvals::ElicitationRequestEvent; +use codex_protocol::mcp::ClientMcpExtensions; +use codex_protocol::mcp::OPENAI_ELICITATION_EXTENSION_ID; +use codex_protocol::mcp::RequestId as ProtocolRequestId; +use codex_protocol::mcp_approval_meta::APPROVAL_KIND_KEY; +use codex_protocol::mcp_approval_meta::APPROVAL_KIND_TOOL_SUGGESTION; +use codex_protocol::mcp_approval_meta::APPROVALS_REVIEWER_KEY; +use codex_protocol::mcp_approval_meta::STRICT_AUTO_REVIEW_KEY; +use codex_protocol::protocol::AskForApproval; +use codex_protocol::protocol::Event; +use codex_protocol::protocol::EventMsg; +use codex_rmcp_client::Elicitation; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::SendElicitation; +use futures::future::BoxFuture; +use futures::future::FutureExt; +use rmcp::model::ElicitationAction; +use rmcp::model::RequestId; +use serde_json::Value; +use tokio::sync::oneshot; + +static NEXT_ELICITATION_REQUEST_ID: AtomicU64 = AtomicU64::new(0); + +const STRICT_AUTO_REVIEW_DECLINE_MESSAGE: &str = "Automated review of this operation failed. Do not proceed without asking the user for explicit approval."; + +#[path = "user_verification_elicitation.rs"] +mod user_verification_elicitation; + +#[derive(Debug, Clone)] +pub struct ElicitationReviewRequest { + pub server_name: String, + pub request_id: RequestId, + pub elicitation: Elicitation, +} + +pub trait ElicitationReviewer: Send + Sync { + fn review( + &self, + request: ElicitationReviewRequest, + ) -> BoxFuture<'static, Result>>; +} + +pub type ElicitationReviewerHandle = Arc; + +/// Holds an owner-provided registration while an MCP elicitation is waiting for a response. +#[derive(Clone)] +pub struct ElicitationLifecycle { + register: Arc Box + Send + Sync>, +} + +impl ElicitationLifecycle { + pub fn new(register: impl Fn() -> T + Send + Sync + 'static) -> Self + where + T: Send + Sync + 'static, + { + Self { + register: Arc::new(move || Box::new(register())), + } + } + + fn start(&self) -> ActiveElicitation { + ActiveElicitation { + _registration: (self.register)(), + } + } +} + +struct ActiveElicitation { + _registration: Box, +} + +/// Routes model-visible elicitation response tokens to their exact pending responders. +/// +/// One router is shared by every MCP runtime created for a thread. The public response token is +/// generated by Codex rather than copied from the MCP connection, so separate runtimes may reuse +/// the same server request ID without colliding. +#[derive(Clone, Default)] +pub(crate) struct ElicitationRequestRouter { + requests: Arc>, + auto_deny: Arc, + full_access_form_input_enabled: Arc, +} + +struct PendingElicitationRequest { + router: ElicitationRequestRouter, + key: (String, RequestId), +} + +impl Drop for PendingElicitationRequest { + fn drop(&mut self) { + let responder = self + .router + .requests + .lock() + .ok() + .and_then(|mut requests| requests.remove(&self.key)); + drop(responder); + } +} + +impl ElicitationRequestRouter { + pub(crate) fn auto_deny(&self) -> bool { + self.auto_deny.load(Ordering::Relaxed) + } + + pub(crate) fn set_auto_deny(&self, auto_deny: bool) { + self.auto_deny.store(auto_deny, Ordering::Relaxed); + } + + pub(crate) fn full_access_form_input_enabled(&self) -> bool { + self.full_access_form_input_enabled.load(Ordering::Acquire) + } + + pub(crate) fn enable_full_access_form_input(&self) { + self.full_access_form_input_enabled + .store(true, Ordering::Release); + } + + pub(crate) async fn resolve( + &self, + server_name: String, + id: RequestId, + response: ElicitationResponse, + ) -> Result<()> { + let responder = self + .requests + .lock() + .map_err(|_| anyhow!("elicitation request router unavailable"))? + .remove(&(server_name, id)) + .ok_or_else(|| anyhow!("elicitation request not found"))?; + responder + .send(response) + .map_err(|_| anyhow!("elicitation response receiver closed")) + } +} + +#[derive(Clone)] +pub(crate) struct ElicitationAuthority { + pub(crate) config: Arc, + reviewer: Option, + lifecycle: Option, +} + +#[derive(Clone, Default)] +pub(crate) struct ElicitationRequestManager { + router: ElicitationRequestRouter, + pub(crate) authority: Arc>>, +} + +impl ElicitationRequestManager { + pub(crate) fn new( + config: Arc, + reviewer: Option, + lifecycle: Option, + router: ElicitationRequestRouter, + ) -> Self { + Self { + router, + authority: Arc::new(StdMutex::new(Some(ElicitationAuthority { + config, + reviewer, + lifecycle, + }))), + } + } + + pub(crate) fn update( + &self, + config: Arc, + reviewer: Option, + lifecycle: Option, + ) -> bool { + let Ok(mut authority) = self.authority.lock() else { + return false; + }; + *authority = Some(ElicitationAuthority { + config, + reviewer, + lifecycle, + }); + true + } + + pub(crate) fn make_sender( + &self, + server_name: String, + tx_event: Option>, + client_mcp_extensions: &ClientMcpExtensions, + ) -> SendElicitation { + // Event receivers such as codex mcp-server do not necessarily handle verification. + // Only wait for a response when trusted host activation enabled this exact route. + let user_verification_enabled = client_mcp_extensions + .get(OPENAI_ELICITATION_EXTENSION_ID) + .and_then(|settings| settings.get("userVerification")) + .is_some_and(Value::is_object); + let router = self.router.clone(); + let authority = self.authority.clone(); + Box::new(move |id, elicitation| { + let router = router.clone(); + let tx_event = tx_event.clone(); + let server_name = server_name.clone(); + let authority = authority.clone(); + async move { + if let Elicitation::UserVerification { + title, + description, + challenge, + } = elicitation + { + if !user_verification_enabled { + return Ok(ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None, + }); + } + return user_verification_elicitation::route( + router, + tx_event, + authority, + server_name, + ElicitationRequest::UserVerification { + title, + description, + challenge, + }, + ) + .await; + } + if router.auto_deny() { + return Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }); + } + + if elicitation + .meta() + .and_then(|meta| meta.get(APPROVAL_KIND_KEY)) + .and_then(serde_json::Value::as_str) + == Some(APPROVAL_KIND_TOOL_SUGGESTION) + { + return Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }); + } + + let Ok(Some(authority)) = authority.lock().map(|authority| authority.clone()) + else { + return Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }); + }; + let ElicitationAuthority { + config, + reviewer, + lifecycle, + } = authority; + let approval_policy = config.approval_policy.value(); + let Some(permission_profile) = config.permission_profile_for_server(&server_name) + else { + return Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }); + }; + + match elicitation + .meta() + .and_then(|meta| meta.get(STRICT_AUTO_REVIEW_KEY)) + { + Some(Value::Bool(true)) => { + if matches!( + approval_policy, + AskForApproval::Granular(config) if !config.allows_mcp_elicitations() + ) { + return Ok(strict_auto_review_decline()); + } + let Some(reviewer) = reviewer.as_ref() else { + return Ok(strict_auto_review_decline()); + }; + let _active_elicitation = + lifecycle.as_ref().map(ElicitationLifecycle::start); + return Ok( + match reviewer + .review(ElicitationReviewRequest { + server_name, + request_id: id, + elicitation, + }) + .await + { + Ok(Some(response)) + if (response.action == ElicitationAction::Accept + && response.content == Some(serde_json::json!({})) + || matches!( + response.action, + ElicitationAction::Decline | ElicitationAction::Cancel + ) && response.content.is_none()) + && response + .meta + .as_ref() + .and_then(Value::as_object) + .and_then(|meta| meta.get(APPROVALS_REVIEWER_KEY)) + .and_then(Value::as_str) + == Some("auto_review") => + { + response + } + Ok(Some(_)) | Ok(None) | Err(_) => strict_auto_review_decline(), + }, + ); + } + None | Some(Value::Bool(false)) => {} + Some(_) => return Ok(strict_auto_review_decline()), + } + + let permission_prompt_is_auto_approved = mcp_permission_prompt_is_auto_approved( + approval_policy, + permission_profile, + McpPermissionPromptAutoApproveContext::default(), + ); + if permission_prompt_is_auto_approved && can_auto_accept_elicitation(&elicitation) { + return Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(serde_json::json!({})), + meta: None, + }); + } + + let should_surface_form_in_full_access = router.full_access_form_input_enabled() + && permission_prompt_is_auto_approved + && !elicitation + .meta() + .is_some_and(|meta| meta.contains_key(APPROVAL_KIND_KEY)) + && match &elicitation { + Elicitation::Mcp( + rmcp::model::ElicitRequestParams::FormElicitationParams { + requested_schema, + .. + }, + ) => !requested_schema.properties.is_empty(), + Elicitation::OpenAiElicitationForm { + requested_schema, .. + } => requested_schema + .get("properties") + .and_then(Value::as_object) + .is_some_and(|properties| !properties.is_empty()), + Elicitation::Mcp(_) + | Elicitation::OpenAiForm { .. } + | Elicitation::UserVerification { .. } => false, + }; + + if !should_surface_form_in_full_access { + if elicitation_is_rejected_by_policy(approval_policy) { + return Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }); + } + + if let Some(reviewer) = reviewer { + let request = ElicitationReviewRequest { + server_name: server_name.clone(), + request_id: id.clone(), + elicitation: elicitation.clone(), + }; + if let Some(response) = reviewer.review(request).await? { + return Ok(response); + } + } + } + + let Some(tx_event) = tx_event else { + return Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }); + }; + + let public_request_id = format!( + "codex-mcp-elicitation-{}", + NEXT_ELICITATION_REQUEST_ID.fetch_add(1, Ordering::Relaxed) + ); + let routed_request_id = RequestId::String(public_request_id.clone().into()); + let request = match elicitation { + Elicitation::UserVerification { .. } => { + return Ok(ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None, + }); + } + Elicitation::Mcp(rmcp::model::ElicitRequestParams::FormElicitationParams { + meta, + message, + requested_schema, + }) => ElicitationRequest::Form { + meta: meta + .map(serde_json::to_value) + .transpose() + .context("failed to serialize MCP elicitation metadata")?, + message, + requested_schema: serde_json::to_value(requested_schema) + .context("failed to serialize MCP elicitation schema")?, + }, + Elicitation::Mcp(rmcp::model::ElicitRequestParams::UrlElicitationParams { + meta, + message, + url, + elicitation_id, + }) => ElicitationRequest::Url { + meta: meta + .map(serde_json::to_value) + .transpose() + .context("failed to serialize MCP elicitation metadata")?, + message, + url, + elicitation_id, + }, + Elicitation::Mcp(_) => { + return Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }); + } + Elicitation::OpenAiForm { + meta, + message, + requested_schema, + } => ElicitationRequest::OpenAiForm { + meta, + message, + requested_schema, + }, + Elicitation::OpenAiElicitationForm { + meta, + message, + requested_schema, + } => ElicitationRequest::OpenAiElicitationForm { + meta, + message, + requested_schema, + }, + }; + let (tx, rx) = oneshot::channel(); + let _active_elicitation = lifecycle.as_ref().map(ElicitationLifecycle::start); + let request_key = (server_name.clone(), routed_request_id); + router + .requests + .lock() + .map_err(|_| anyhow!("elicitation request router unavailable"))? + .insert(request_key.clone(), tx); + let _pending_request = PendingElicitationRequest { + router: router.clone(), + key: request_key, + }; + tx_event + .send(Event { + id: "mcp_elicitation_request".to_string(), + msg: EventMsg::ElicitationRequest(ElicitationRequestEvent { + turn_id: None, + server_name, + id: ProtocolRequestId::String(public_request_id), + request, + }), + }) + .await + .context("failed to deliver MCP elicitation request")?; + rx.await + .context("elicitation request channel closed unexpectedly") + } + .boxed() + }) + } +} + +fn strict_auto_review_decline() -> ElicitationResponse { + ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: Some(serde_json::json!({ + "message": STRICT_AUTO_REVIEW_DECLINE_MESSAGE, + })), + } +} + +pub(crate) fn elicitation_is_rejected_by_policy(approval_policy: AskForApproval) -> bool { + match approval_policy { + AskForApproval::Never => true, + AskForApproval::OnRequest => false, + AskForApproval::UnlessTrusted => false, + AskForApproval::Granular(granular_config) => !granular_config.allows_mcp_elicitations(), + } +} + +type ResponderMap = HashMap<(String, RequestId), oneshot::Sender>; + +fn can_auto_accept_elicitation(elicitation: &Elicitation) -> bool { + match elicitation { + Elicitation::Mcp(rmcp::model::ElicitRequestParams::FormElicitationParams { + requested_schema, + .. + }) => { + // Auto-accept confirm/approval elicitations without schema requirements. + requested_schema.properties.is_empty() + } + Elicitation::Mcp(_) + | Elicitation::OpenAiForm { .. } + | Elicitation::OpenAiElicitationForm { .. } + | Elicitation::UserVerification { .. } => false, + } +} + +#[cfg(test)] +#[path = "elicitation_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/elicitation_tests.rs b/codex-rs/codex-mcp/src/elicitation_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..3e9f0867f498d13e076f8117e8200194a1d99c06 --- /dev/null +++ b/codex-rs/codex-mcp/src/elicitation_tests.rs @@ -0,0 +1,638 @@ +use super::*; +use crate::mcp::tests::test_elicitation_config; +use async_channel::Receiver; +use codex_protocol::models::PermissionProfile; +use codex_protocol::protocol::GranularApprovalConfig; +use pretty_assertions::assert_eq; +use rmcp::model::ElicitRequestParams; +use rmcp::model::ElicitationSchema; +use rmcp::model::RequestMetaObject; +use serde_json::Map; +use serde_json::json; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering::Relaxed; + +type ReviewerResponse = std::result::Result, &'static str>; + +struct RecordingReviewer { + calls: AtomicUsize, + active_elicitations: Arc, + response: ReviewerResponse, +} + +impl RecordingReviewer { + fn new(response: ReviewerResponse) -> Arc { + Arc::new(Self { + calls: AtomicUsize::default(), + active_elicitations: Arc::default(), + response, + }) + } +} + +impl ElicitationReviewer for RecordingReviewer { + fn review( + &self, + request: ElicitationReviewRequest, + ) -> BoxFuture<'static, Result>> { + assert_eq!(request.server_name, "independent-mcp"); + self.calls.fetch_add(/*val*/ 1, Relaxed); + let active_elicitations = self.active_elicitations.clone(); + let response = self.response.clone(); + async move { + assert_eq!(active_elicitations.load(Relaxed), 1); + tokio::task::yield_now().await; + assert_eq!(active_elicitations.load(Relaxed), 1); + response.map_err(anyhow::Error::msg) + } + .boxed() + } +} + +struct LifecycleRegistration(Arc); + +impl Drop for LifecycleRegistration { + fn drop(&mut self) { + self.0.fetch_sub(/*val*/ 1, Relaxed); + } +} + +fn approved_response() -> ElicitationResponse { + ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: Some(json!({ "approvals_reviewer": "auto_review" })), + } +} + +fn elicitation_fixture( + approval_policy: AskForApproval, + permission_profile: PermissionProfile, + reviewer: Option>, +) -> (ElicitationRequestManager, Receiver, SendElicitation) { + let lifecycle = reviewer.as_ref().map(|reviewer| { + let active_elicitations = reviewer.active_elicitations.clone(); + ElicitationLifecycle::new(move || { + active_elicitations.fetch_add(/*val*/ 1, Relaxed); + LifecycleRegistration(active_elicitations.clone()) + }) + }); + let mut config = test_elicitation_config( + "independent-mcp", + approval_policy, + permission_profile.clone(), + ); + Arc::make_mut(&mut config) + .server_permission_profiles + .insert("another-independent-mcp".to_string(), permission_profile); + let manager = ElicitationRequestManager::new( + config, + reviewer.map(|reviewer| reviewer as Arc), + lifecycle, + ElicitationRequestRouter::default(), + ); + let (tx_event, events) = async_channel::bounded(1); + let sender = manager.make_sender( + "independent-mcp".to_string(), + Some(tx_event), + &ClientMcpExtensions::default(), + ); + (manager, events, sender) +} + +async fn send_elicitation(sender: &SendElicitation, marker: Option) -> ElicitationResponse { + let elicitation = Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta: marker.map(|value| { + RequestMetaObject::from(Map::from_iter([(STRICT_AUTO_REVIEW_KEY.into(), value)])) + }), + message: "Review this request".to_string(), + requested_schema: ElicitationSchema::builder().build().unwrap(), + }); + sender(RequestId::Number(7), elicitation) + .await + .expect("elicitation must receive a terminal response") +} + +async fn assert_declined(marker: Value, response: Option) { + let expected_calls = usize::from(marker == Value::Bool(true)); + let reviewer = response.map(RecordingReviewer::new); + let (_, events, sender) = elicitation_fixture( + AskForApproval::Never, + PermissionProfile::Disabled, + reviewer.clone(), + ); + assert_eq!( + send_elicitation(&sender, Some(marker)).await, + strict_auto_review_decline() + ); + if let Some(reviewer) = reviewer { + assert_eq!(reviewer.calls.load(Relaxed), expected_calls); + } + assert!(events.is_empty()); +} + +#[test] +fn closed_event_channel_immediately_cleans_up_pending_elicitation() { + let active_elicitations = Arc::new(AtomicUsize::new(0)); + let registrations = active_elicitations.clone(); + let lifecycle = ElicitationLifecycle::new(move || { + registrations.fetch_add(/*val*/ 1, Relaxed); + LifecycleRegistration(registrations.clone()) + }); + let (manager, events, sender) = elicitation_fixture( + AskForApproval::OnRequest, + PermissionProfile::Disabled, + /*reviewer*/ None, + ); + assert!(manager.update( + test_elicitation_config( + "independent-mcp", + AskForApproval::OnRequest, + PermissionProfile::Disabled + ), + /*reviewer*/ None, + Some(lifecycle), + )); + drop(events); + + let elicitation = Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta: None, + message: "Review this request".to_string(), + requested_schema: ElicitationSchema::builder().build().unwrap(), + }); + let error = sender(RequestId::Number(7), elicitation) + .now_or_never() + .expect("closed event channel must not leave an elicitation pending") + .expect_err("closed event channel must fail the elicitation"); + + assert_eq!( + error.to_string(), + "failed to deliver MCP elicitation request" + ); + assert!( + manager + .router + .requests + .lock() + .expect("pending request router should be available") + .is_empty() + ); + assert_eq!(active_elicitations.load(Relaxed), 0); +} + +#[tokio::test] +async fn strict_auto_review_respects_explicit_elicitation_denials() { + for policy in [ + AskForApproval::OnRequest, + AskForApproval::UnlessTrusted, + AskForApproval::Never, + AskForApproval::Granular(GranularApprovalConfig { + sandbox_approval: true, + rules: true, + skill_approval: true, + request_permissions: true, + mcp_elicitations: false, + }), + ] { + let explicitly_denied = matches!( + policy, + AskForApproval::Granular(config) if !config.allows_mcp_elicitations() + ); + let reviewer = RecordingReviewer::new(Ok(Some(approved_response()))); + let (manager, events, sender) = + elicitation_fixture(policy, PermissionProfile::Disabled, Some(reviewer.clone())); + assert_eq!( + send_elicitation(&sender, Some(json!(true))).await, + if explicitly_denied { + strict_auto_review_decline() + } else { + approved_response() + } + ); + if policy == AskForApproval::Never { + for (server_name, marker) in [ + ("independent-mcp", Some(json!(false))), + ("another-independent-mcp", None), + ] { + let sender = manager.make_sender( + server_name.into(), + /*tx_event*/ None, + &ClientMcpExtensions::default(), + ); + assert_eq!( + send_elicitation(&sender, marker).await, + ElicitationResponse { + meta: None, + ..approved_response() + }, + ); + } + } + manager.router.set_auto_deny(/*auto_deny*/ true); + assert_eq!( + send_elicitation(&sender, Some(json!(true))).await, + ElicitationResponse { + meta: None, + ..strict_auto_review_decline() + }, + ); + assert_eq!( + ( + reviewer.calls.load(Relaxed), + reviewer.active_elicitations.load(Relaxed) + ), + (usize::from(!explicitly_denied), 0), + ); + assert!(events.is_empty(), "strict review must not emit an event"); + } +} + +#[tokio::test] +async fn strict_auto_review_preserves_guardian_denials_and_cancellations() { + for response in [ + ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: Some(json!({ + "approvals_reviewer": "auto_review", + "message": "The user has not authorized sending this data. Ask the user for approval.", + })), + }, + ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: Some(json!({ "approvals_reviewer": "auto_review" })), + }, + ] { + let reviewer = RecordingReviewer::new(Ok(Some(response.clone()))); + let (_, events, sender) = elicitation_fixture( + AskForApproval::Never, + PermissionProfile::Disabled, + Some(reviewer.clone()), + ); + assert_eq!(send_elicitation(&sender, Some(json!(true))).await, response); + assert_eq!(reviewer.calls.load(Relaxed), 1); + assert!(events.is_empty(), "strict review must not emit an event"); + } +} + +#[tokio::test] +async fn strict_auto_review_fails_closed_without_a_canonical_decision() { + for marker in ["null", "\"true\"", "1", "{}", "[true]"] { + let marker = serde_json::from_str(marker).expect("valid malformed marker"); + assert_declined(marker, Some(Ok(Some(approved_response())))).await; + } + for response in [Ok(None), Err("reviewer failed")] { + assert_declined(json!(true), Some(response)).await; + } + let invalid_decisions: [fn(&mut ElicitationResponse); 6] = [ + |response| { + response.action = ElicitationAction::Decline; + response.meta = Some(json!({ "message": "Ask the user to approve this request." })); + }, + |response| response.action = ElicitationAction::Cancel, + |response| response.meta = None, + |response| response.meta = Some(json!({ "approvals_reviewer": "user" })), + |response| response.meta = Some(json!({ "approvals_reviewer": "guardian_subagent" })), + |response| response.content = Some(json!({ "approved_for_session": true })), + ]; + for make_invalid in invalid_decisions { + let mut response = approved_response(); + make_invalid(&mut response); + assert_declined(json!(true), Some(Ok(Some(response)))).await; + } + assert_declined(json!(true), /*response*/ None).await; +} + +#[tokio::test] +async fn reused_elicitation_senders_follow_each_servers_latest_permission_authority() { + let mut config = crate::mcp::tests::test_mcp_config(std::env::temp_dir()); + config.approval_policy = codex_config::Constrained::allow_any(AskForApproval::Never); + config.permission_profile = PermissionProfile::Disabled; + config.apps_enabled = true; + let auth = codex_login::CodexAuth::create_dummy_chatgpt_auth_for_testing(); + + let hosted_server = crate::codex_apps_mcp_server_config( + "https://example.com", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ); + let mut attached_server = hosted_server.clone(); + attached_server.environment_id = "attached".to_string(); + let mut catalog = crate::ResolvedMcpCatalog::builder(); + catalog.register(crate::McpServerRegistration::from_config( + "attached".to_string(), + attached_server, + )); + catalog.register(crate::McpServerRegistration::from_hosted_apps( + "host", + /*contribution_order*/ 0, + hosted_server, + )); + config.mcp_server_catalog = catalog.build(); + let servers = crate::effective_mcp_servers(&config, Some(&auth)); + config.set_server_permission_profiles( + &servers, + [("attached".to_string(), PermissionProfile::read_only())], + ); + + let manager = ElicitationRequestManager::new( + Arc::new(config.clone()), + /*reviewer*/ None, + /*lifecycle*/ None, + ElicitationRequestRouter::default(), + ); + let attached = manager.make_sender( + "attached".to_string(), + /*tx_event*/ None, + &ClientMcpExtensions::default(), + ); + let hosted = manager.make_sender( + crate::CODEX_APPS_MCP_SERVER_NAME.to_string(), + /*tx_event*/ None, + &ClientMcpExtensions::default(), + ); + + assert_eq!( + send_elicitation(&attached, /*marker*/ None).await.action, + ElicitationAction::Decline + ); + assert_eq!( + send_elicitation(&hosted, /*marker*/ None).await.action, + ElicitationAction::Accept + ); + + config.set_server_permission_profiles( + &servers, + [("attached".to_string(), PermissionProfile::Disabled)], + ); + assert!(manager.update( + Arc::new(config.clone()), + /*reviewer*/ None, + /*lifecycle*/ None, + )); + assert_eq!( + send_elicitation(&attached, /*marker*/ None).await.action, + ElicitationAction::Accept + ); + + let mut configured_servers = config.mcp_server_catalog.configured_servers(); + configured_servers + .get_mut("attached") + .expect("attached server should be registered") + .enabled = false; + config.mcp_server_catalog = config + .mcp_server_catalog + .with_materialized_servers(configured_servers); + let servers = crate::effective_mcp_servers(&config, Some(&auth)); + config.set_server_permission_profiles( + &servers, + [("attached".to_string(), PermissionProfile::Disabled)], + ); + assert!(manager.update( + Arc::new(config.clone()), + /*reviewer*/ None, + /*lifecycle*/ None, + )); + assert_eq!( + send_elicitation(&attached, /*marker*/ None).await.action, + ElicitationAction::Decline + ); + + let servers = crate::effective_mcp_servers(&config, /*auth*/ None); + config.set_server_permission_profiles(&servers, std::iter::empty()); + assert!(manager.update( + Arc::new(config.clone()), + /*reviewer*/ None, + /*lifecycle*/ None, + )); + assert_eq!( + send_elicitation(&hosted, /*marker*/ None).await.action, + ElicitationAction::Decline + ); +} + +fn verification_fixture( + approval_policy: AskForApproval, + reviewer: Option>, +) -> (ElicitationRequestManager, Receiver, SendElicitation) { + let mut config = test_elicitation_config( + crate::CODEX_APPS_MCP_SERVER_NAME, + approval_policy, + PermissionProfile::Disabled, + ); + let mut catalog = crate::catalog::ResolvedMcpCatalog::builder(); + catalog.register(crate::catalog::McpServerRegistration::from_hosted_apps( + "verification-test", + /*contribution_order*/ 0, + crate::mcp::codex_apps_mcp_server_config( + "https://example.com", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ), + )); + Arc::make_mut(&mut config).mcp_server_catalog = catalog.build(); + let manager = ElicitationRequestManager::new( + config, + reviewer.map(|reviewer| reviewer as Arc), + /*lifecycle*/ None, + ElicitationRequestRouter::default(), + ); + let (tx, events) = async_channel::bounded(1); + let sender = manager.make_sender( + crate::CODEX_APPS_MCP_SERVER_NAME.into(), + Some(tx), + &ClientMcpExtensions::new([( + OPENAI_ELICITATION_EXTENSION_ID.to_string(), + json!({"userVerification": {}}), + )]), + ); + (manager, events, sender) +} + +#[tokio::test] +async fn user_verification_requires_the_app_even_when_policy_would_approve_or_decline() { + for approval_policy in [AskForApproval::Never, AskForApproval::OnRequest] { + let reviewer = RecordingReviewer::new(Ok(Some(approved_response()))); + let (manager, events, sender) = + verification_fixture(approval_policy, Some(reviewer.clone())); + let pending = tokio::spawn(sender( + RequestId::Number(7), + Elicitation::UserVerification { + title: "Approve purchase".into(), + description: "Pay $200".into(), + challenge: "AQID".into(), + }, + )); + let event = events.recv().await.unwrap(); + let EventMsg::ElicitationRequest(request) = event.msg else { + panic!("expected typed elicitation"); + }; + assert_eq!( + request.request, + ElicitationRequest::UserVerification { + title: "Approve purchase".into(), + description: "Pay $200".into(), + challenge: "AQID".into(), + } + ); + assert_eq!(reviewer.calls.load(Relaxed), 0); + let ProtocolRequestId::String(id) = request.id else { + panic!("expected routed request id"); + }; + let proof = ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({"credentialId": "AQID", "signature": "BAUG"})), + meta: None, + }; + manager + .router + .resolve( + request.server_name.clone(), + RequestId::String(id.clone().into()), + proof.clone(), + ) + .await + .unwrap(); + assert_eq!(pending.await.unwrap().unwrap(), proof); + assert!( + manager + .router + .resolve(request.server_name, RequestId::String(id.into()), proof) + .await + .is_err() + ); + } +} + +#[tokio::test] +async fn user_verification_cancels_when_no_app_can_receive_the_request() { + let (manager, _, _) = verification_fixture(AskForApproval::OnRequest, /*reviewer*/ None); + let sender = manager.make_sender( + crate::CODEX_APPS_MCP_SERVER_NAME.into(), + /*tx_event*/ None, + &ClientMcpExtensions::new([( + OPENAI_ELICITATION_EXTENSION_ID.to_string(), + json!({"userVerification": {}}), + )]), + ); + assert_eq!( + sender( + RequestId::Number(7), + Elicitation::UserVerification { + title: "Approve".into(), + description: String::new(), + challenge: "AQID".into(), + } + ) + .await + .unwrap(), + ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None + }, + ); +} + +#[tokio::test] +async fn user_verification_cancels_for_an_event_receiver_without_host_activation() { + let (manager, _, _) = verification_fixture(AskForApproval::OnRequest, /*reviewer*/ None); + let (tx, events) = async_channel::bounded(1); + let sender = manager.make_sender( + crate::CODEX_APPS_MCP_SERVER_NAME.into(), + Some(tx), + &ClientMcpExtensions::default(), + ); + let response = tokio::time::timeout( + std::time::Duration::from_secs(1), + sender( + RequestId::Number(7), + Elicitation::UserVerification { + title: "Approve".into(), + description: String::new(), + challenge: "AQID".into(), + }, + ), + ) + .await + .expect("an inactive host must not wait for a response") + .unwrap(); + assert_eq!( + response, + ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None, + }, + ); + assert!(events.try_recv().is_err()); + assert!(manager.router.requests.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn user_verification_drops_pending_response_when_the_request_is_cancelled() { + let (manager, events, sender) = + verification_fixture(AskForApproval::OnRequest, /*reviewer*/ None); + let pending = tokio::spawn(sender( + RequestId::Number(7), + Elicitation::UserVerification { + title: "Approve".into(), + description: String::new(), + challenge: "AQID".into(), + }, + )); + events.recv().await.unwrap(); + pending.abort(); + assert!(pending.await.unwrap_err().is_cancelled()); + assert!(manager.router.requests.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn user_verification_rejects_attached_servers_even_if_they_use_the_plugin_service_name() { + for name in ["attached", crate::CODEX_APPS_MCP_SERVER_NAME] { + let (manager, _, _) = + verification_fixture(AskForApproval::OnRequest, /*reviewer*/ None); + { + let mut authority = manager.authority.lock().unwrap(); + let config = Arc::make_mut(&mut authority.as_mut().unwrap().config); + let mut catalog = crate::catalog::ResolvedMcpCatalog::builder(); + catalog.register(crate::catalog::McpServerRegistration::from_config( + name.into(), + crate::mcp::codex_apps_mcp_server_config( + "https://example.com", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ), + )); + config.mcp_server_catalog = catalog.build(); + } + let (tx, events) = async_channel::bounded(1); + let sender = manager.make_sender( + name.into(), + Some(tx), + &ClientMcpExtensions::new([( + OPENAI_ELICITATION_EXTENSION_ID.to_string(), + json!({"userVerification": {}}), + )]), + ); + assert_eq!( + sender( + RequestId::Number(7), + Elicitation::UserVerification { + title: "Approve".into(), + description: String::new(), + challenge: "AQID".into() + } + ) + .await + .unwrap(), + ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None + }, + ); + assert!(events.try_recv().is_err()); + } +} diff --git a/codex-rs/codex-mcp/src/event_stream.rs b/codex-rs/codex-mcp/src/event_stream.rs new file mode 100644 index 0000000000000000000000000000000000000000..ca52d20f6fd1c7f8f135854c1e82d6695005b869 --- /dev/null +++ b/codex-rs/codex-mcp/src/event_stream.rs @@ -0,0 +1,159 @@ +//! Opens MCP event streams with clients that remain connected after task unloading. + +use std::sync::Arc; + +use anyhow::Result; +use anyhow::bail; +use codex_api::SharedAuthProvider; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_login::AuthManager; +use codex_login::CodexAuth; +use codex_protocol::mcp::ClientMcpExtensions; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::McpOAuthRefreshMode; +use rmcp::model::ElicitationAction; +use rmcp::model::ElicitationCapability; +use serde_json::Map; +use serde_json::Value; +use tokio::sync::watch; + +use crate::CODEX_APPS_MCP_SERVER_NAME; +use crate::EffectiveMcpServer; +use crate::McpEventStream; +use crate::McpProtocolMode; +use crate::McpRuntimeContext; +use crate::rmcp_client::DEFAULT_STARTUP_TIMEOUT; +use crate::rmcp_client::make_rmcp_client; +use crate::rmcp_client::mcp_initialize_request_params; + +pub(crate) struct EventStreamConnectionSettings { + pub server: EffectiveMcpServer, + pub store_mode: OAuthCredentialsStoreMode, + pub keyring_backend_kind: AuthKeyringBackendKind, + pub oauth_refresh_mode: McpOAuthRefreshMode, + pub runtime_context: McpRuntimeContext, + pub resolved_environment: std::result::Result>, String>, + pub auth_provider: Option, + pub auth_manager: Option>, + pub auth: Option, + pub protocol_mode: McpProtocolMode, + pub client_mcp_extensions: ClientMcpExtensions, +} + +/// Opens each event stream with its own MCP client. +/// Owners watch `wait_for_access_change` to cancel streams when access changes. +#[derive(Clone)] +pub struct McpEventStreamOpener { + pub(crate) connection: Arc, + pub(crate) cancellation_receiver: watch::Receiver<()>, + pub(crate) cancel_event_streams_on_server_removal: watch::Sender<()>, +} + +impl McpEventStreamOpener { + /// Retains cancellation for this task's subscriptions across MCP runtime replacement. + pub fn event_stream_cancellation_sender(&self) -> watch::Sender<()> { + self.cancel_event_streams_on_server_removal.clone() + } + + /// Creates an MCP client and opens an event stream. + pub async fn open( + &self, + event_name: &str, + arguments: &Value, + request_meta: Option<&Map>, + ) -> Result { + tokio::select! { + biased; + () = self.wait_for_access_change() => bail!("event subscription access changed"), + result = async { + let connection = &self.connection; + if let Some(manager) = &connection.auth_manager { + let auth = manager.auth().await; + if !self.matches_auth(auth.as_ref()) { + bail!("event subscription account changed"); + } + } + + let startup_timeout = connection.server.config().startup_timeout_sec + .unwrap_or(DEFAULT_STARTUP_TIMEOUT); + let client = Arc::new(tokio::time::timeout(startup_timeout, make_rmcp_client( + CODEX_APPS_MCP_SERVER_NAME, + connection.server.clone(), + connection.store_mode, + connection.keyring_backend_kind, + connection.oauth_refresh_mode, + connection.runtime_context.clone(), + connection.resolved_environment.clone(), + connection.auth_provider.clone(), + connection.protocol_mode, + )).await??); + client.initialize( + mcp_initialize_request_params( + ElicitationCapability::default(), + connection.client_mcp_extensions.clone(), + ), + Some(startup_timeout), + Box::new(|_, _| Box::pin(async { + Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }) + })), + ).await?; + McpEventStream::open( + client, + self.cancellation_receiver.clone(), + event_name, + arguments, + request_meta, + ).await + } => result, + } + } + + /// Waits for an account change or removal of the event server from the task. + pub async fn wait_for_access_change(&self) { + let mut cancellation_receiver = self.cancellation_receiver.clone(); + let auth_change = async { + let Some(manager) = &self.connection.auth_manager else { + return std::future::pending().await; + }; + let mut changes = manager.auth_change_receiver(); + loop { + if !self.matches_auth(manager.auth_cached().as_ref()) { + return; + } + if changes.changed().await.is_err() { + return; + } + } + }; + tokio::select! { + Ok(()) = cancellation_receiver.changed() => {}, + () = auth_change => {}, + } + } + + fn matches_auth(&self, current: Option<&CodexAuth>) -> bool { + match (self.connection.auth.as_ref(), current) { + (Some(CodexAuth::AgentIdentity(expected)), Some(CodexAuth::AgentIdentity(current))) => { + expected.record() == current.record() + } + (Some(CodexAuth::AgentIdentity(_)), _) | (_, Some(CodexAuth::AgentIdentity(_))) => { + false + } + (Some(expected), Some(current)) => { + expected.get_account_id() == current.get_account_id() + && expected.get_chatgpt_user_id() == current.get_chatgpt_user_id() + && expected.is_workspace_account() == current.is_workspace_account() + && expected.is_fedramp_account() == current.is_fedramp_account() + && (expected.get_account_id().is_some() || expected == current) + } + (None, None) => true, + _ => false, + } + } +} diff --git a/codex-rs/codex-mcp/src/executor_environment_http_client.rs b/codex-rs/codex-mcp/src/executor_environment_http_client.rs new file mode 100644 index 0000000000000000000000000000000000000000..c2d50617c6e3865064ff88e06daee1631314a2ea --- /dev/null +++ b/codex-rs/codex-mcp/src/executor_environment_http_client.rs @@ -0,0 +1,45 @@ +use std::sync::Arc; + +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpHeader; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +use futures::future::BoxFuture; + +pub(crate) struct ExecutorEnvironmentHttpClient { + pub(crate) bearer_token_env_var: String, + pub(crate) http_client: Arc, +} + +impl ExecutorEnvironmentHttpClient { + fn attach_authorization(&self, params: &mut HttpRequestParams) { + params + .headers + .retain(|header| !header.name.eq_ignore_ascii_case("authorization")); + params.headers.push(HttpHeader { + name: "authorization".to_string(), + value: "Bearer ".to_string(), + value_env_var: Some(self.bearer_token_env_var.clone()), + }); + } +} + +impl HttpClient for ExecutorEnvironmentHttpClient { + fn http_request( + &self, + mut params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + self.attach_authorization(&mut params); + self.http_client.http_request(params) + } + + fn http_request_stream( + &self, + mut params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + self.attach_authorization(&mut params); + self.http_client.http_request_stream(params) + } +} diff --git a/codex-rs/codex-mcp/src/lib.rs b/codex-rs/codex-mcp/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..8a4f489c904b71df301f17125e55b4515f773e3f --- /dev/null +++ b/codex-rs/codex-mcp/src/lib.rs @@ -0,0 +1,124 @@ +pub use binding::McpBinding; +pub use binding::PreparedMcpCall; +pub use client_capabilities::client_mcp_extensions; +pub use client_tool_catalog::CodexAppsToolSnapshot; +pub use codex_rmcp_client::McpProtocolMode; +pub use connection_manager::tool_is_model_visible; +pub use elicitation::ElicitationLifecycle; +pub use elicitation::ElicitationReviewRequest; +pub use elicitation::ElicitationReviewer; +pub use elicitation::ElicitationReviewerHandle; +pub use event_stream::McpEventStreamOpener; +pub use resource_client::CodexAppsResourceListParams; +pub use resource_client::McpEventCatalogSnapshot; +pub use resource_client::McpEventDefinition; +pub use resource_client::McpEventNotification; +pub use resource_client::McpEventStream; +pub use resource_client::McpResourceClient; +pub use resource_client::McpResourceClientCacheKey; +pub use resource_client::McpResourcePage; +pub use resource_client::McpResourceReadResult; +pub use rmcp::model::ReadResourceRequestParams; +pub use rmcp_client::MCP_SANDBOX_STATE_META_CAPABILITY; +pub use runtime::McpRuntime; +pub use runtime::McpRuntimeContext; +pub use runtime::McpRuntimeInput; +pub use runtime::McpStartupPolicy; +pub use runtime::SandboxState; +pub use runtime::apply_http_headers_helper; +pub use tool_catalog_cache::McpToolCatalogCache; +pub use tools::ToolInfo; +pub use trusted_access::TrustedAccessContext; + +/// Backward-compatible name for the shared Codex Apps tools runtime. +pub type CodexAppsToolsCache = codex_connectors::ConnectorRuntimeManager; +/// Backward-compatible name for the Codex Apps runtime context key. +pub type CodexAppsToolsCacheKey = codex_connectors::ConnectorRuntimeContextKey; + +pub use catalog::McpCatalogBuilder; +pub use catalog::McpEnvironmentAuthority; +pub use catalog::McpPluginAttribution; +pub use catalog::McpServerConflict; +pub use catalog::McpServerConflictAction; +pub use catalog::McpServerRegistration; +pub use catalog::McpServerSource; +pub use catalog::ResolvedMcpCatalog; +pub use catalog::ResolvedMcpServer; + +pub use mcp::CODEX_APPS_MCP_SERVER_NAME; +pub use mcp::DEFAULT_OPTIONAL_MCP_STARTUP_GRACE; +pub use mcp::McpConfig; +pub use mcp::ToolPluginContext; +pub use server::EffectiveMcpServer; + +pub use auth_elicitation::CodexAppsAuthElicitation; +pub use auth_elicitation::CodexAppsAuthElicitationPlan; +pub use auth_elicitation::CodexAppsConnectorAuthFailure; +pub use auth_elicitation::MCP_TOOL_CODEX_APPS_META_KEY; +pub use auth_elicitation::auth_elicitation_completed_result; +pub use auth_elicitation::auth_elicitation_id; +pub use auth_elicitation::build_auth_elicitation; +pub use auth_elicitation::build_auth_elicitation_plan; +pub use auth_elicitation::connector_auth_failure_from_tool_result; +pub use auth_elicitation::is_connector_auth_failure_from_tool_result; +/// Backward-compatible name for the Codex Apps runtime context key builder. +pub use codex_connectors::connector_runtime_context_key as codex_apps_tools_cache_key; +pub use mcp::codex_apps_mcp_server_config; +pub use mcp::configured_mcp_servers; +pub use mcp::effective_mcp_servers; +pub use mcp::effective_mcp_servers_from_configured; +pub use mcp::host_owned_codex_apps_enabled; +pub use mcp::hosted_plugin_runtime_mcp_server_config; +pub use mcp::tool_plugin_context; +pub use plugin_config::PluginMcpConfigParseOutcome; +pub use plugin_config::PluginMcpServerParseError; +pub use plugin_config::parse_agent_plugin_mcp_config; +pub use plugin_config::parse_executor_plugin_mcp_config; +pub use plugin_config::parse_plugin_mcp_config; + +pub use mcp::McpServerStatusSnapshot; +pub use mcp::McpSnapshotDetail; +pub use mcp::collect_mcp_server_status_snapshot_with_detail; +pub use mcp::read_mcp_resource; + +pub use mcp::McpAuthStatusEntry; +pub use mcp::McpOAuthLoginConfig; +pub use mcp::McpOAuthLoginSupport; +pub use mcp::McpOAuthScopesSource; +pub use mcp::ResolvedMcpOAuthScopes; +pub use mcp::compute_auth_statuses; +pub use mcp::discover_supported_scopes; +pub use mcp::oauth_login_support; +pub use mcp::resolve_oauth_callback; +pub use mcp::resolve_oauth_scopes; +pub use mcp::should_retry_without_scopes; + +pub use codex_apps::declared_openai_file_input_param_names; +pub use mcp::McpPermissionPromptAutoApproveContext; +pub use mcp::mcp_permission_prompt_is_auto_approved; +pub use mcp::qualified_mcp_tool_name_prefix; + +mod auth_changes; +pub(crate) mod auth_elicitation; +mod binding; +pub(crate) mod binding_clients; +mod catalog; +mod client_capabilities; +mod client_tool_catalog; +pub(crate) mod codex_apps; +pub(crate) mod connection_manager; +pub(crate) mod elicitation; +mod event_stream; +mod executor_environment_http_client; +pub(crate) mod mcp; +mod openai_docs_source_attribution; +mod pagination; +mod plugin_config; +mod resource_client; +mod resource_origin; +pub(crate) mod rmcp_client; +pub(crate) mod runtime; +pub(crate) mod server; +mod tool_catalog_cache; +pub(crate) mod tools; +mod trusted_access; diff --git a/codex-rs/codex-mcp/src/mcp/auth.rs b/codex-rs/codex-mcp/src/mcp/auth.rs new file mode 100644 index 0000000000000000000000000000000000000000..8ecfe156d95651e1d709aa85af000dbb667ac5fd --- /dev/null +++ b/codex-rs/codex-mcp/src/mcp/auth.rs @@ -0,0 +1,497 @@ +use std::collections::HashMap; +use std::sync::Arc; + +use anyhow::Result; +use codex_config::McpServerAuth; +use codex_config::McpServerConfig; +use codex_config::McpServerTransportConfig; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::HttpClient; +use codex_login::CodexAuth; +use codex_rmcp_client::McpAuthState; +use codex_rmcp_client::McpOAuthCallbackMode; +use codex_rmcp_client::OAuthDiscoveryTimeout; +use codex_rmcp_client::OAuthProviderError; +use codex_rmcp_client::StreamableHttpRedirectMode; +use codex_rmcp_client::determine_streamable_http_auth_status; +use codex_rmcp_client::determine_streamable_http_auth_status_from_credentials; +use codex_rmcp_client::discover_streamable_http_oauth; +use codex_rmcp_client::resolve_mcp_oauth_callback_url; +use futures::FutureExt; +use futures::future::join_all; +use tracing::warn; + +use crate::runtime::McpRuntimeContext; +use crate::server::EffectiveMcpServer; +use crate::server::has_explicit_http_authorization; + +#[derive(Debug, Clone)] +pub struct McpOAuthLoginConfig { + pub url: String, + pub http_headers: Option>, + pub env_http_headers: Option>, + pub discovered_scopes: Option>, + pub callback_mode: McpOAuthCallbackMode, +} + +#[derive(Debug)] +pub enum McpOAuthLoginSupport { + Supported(McpOAuthLoginConfig), + Unsupported, + Unknown(anyhow::Error), +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum McpOAuthScopesSource { + Explicit, + Configured, + Discovered, + Empty, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ResolvedMcpOAuthScopes { + pub scopes: Vec, + pub source: McpOAuthScopesSource, +} + +/// Keeps registered callbacks tied to their client while preserving legacy redirects. +pub fn resolve_oauth_callback( + server: &McpServerConfig, + server_url: &str, + global_callback_url: Option<&str>, +) -> Result> { + if let Some(callback_url) = server + .oauth + .as_ref() + .and_then(|oauth| oauth.callback_url.as_deref()) + { + return Ok(Some(callback_url.to_string())); + } + + if server + .oauth_client_id() + .is_none_or(|client_id| client_id.trim().is_empty()) + { + return Ok(global_callback_url.map(ToOwned::to_owned)); + } + + resolve_mcp_oauth_callback_url( + server_url, + global_callback_url, + McpOAuthCallbackMode::CallbackSpecific, + ) + .map(Some) +} + +#[derive(Debug, Clone)] +pub struct McpAuthStatusEntry { + pub config: Option, + pub auth_state: McpAuthState, +} + +pub async fn oauth_login_support( + transport: &McpServerTransportConfig, + http_client: Arc, + discovery_timeout: OAuthDiscoveryTimeout, + redirect_mode: StreamableHttpRedirectMode, +) -> McpOAuthLoginSupport { + let Some(mut config) = oauth_login_candidate(transport) else { + return McpOAuthLoginSupport::Unsupported; + }; + match discover_streamable_http_oauth( + &config.url, + config.http_headers.clone(), + config.env_http_headers.clone(), + http_client, + discovery_timeout, + redirect_mode, + ) + .await + { + Ok(Some(discovery)) => { + config.discovered_scopes = discovery.scopes_supported; + config.callback_mode = discovery.callback_mode; + McpOAuthLoginSupport::Supported(config) + } + Ok(None) => McpOAuthLoginSupport::Unsupported, + Err(err) => McpOAuthLoginSupport::Unknown(err), + } +} + +fn oauth_login_candidate(transport: &McpServerTransportConfig) -> Option { + let McpServerTransportConfig::StreamableHttp { + url, + bearer_token_env_var, + http_headers, + env_http_headers, + .. + } = transport + else { + return None; + }; + if bearer_token_env_var.is_some() { + return None; + } + Some(McpOAuthLoginConfig { + url: url.clone(), + http_headers: http_headers.clone(), + env_http_headers: env_http_headers.clone(), + discovered_scopes: None, + callback_mode: McpOAuthCallbackMode::CallbackSpecific, + }) +} + +pub async fn discover_supported_scopes( + transport: &McpServerTransportConfig, + http_client: Arc, + discovery_timeout: OAuthDiscoveryTimeout, + redirect_mode: StreamableHttpRedirectMode, +) -> Option> { + match oauth_login_support(transport, http_client, discovery_timeout, redirect_mode).await { + McpOAuthLoginSupport::Supported(config) => config.discovered_scopes, + McpOAuthLoginSupport::Unsupported | McpOAuthLoginSupport::Unknown(_) => None, + } +} + +pub fn resolve_oauth_scopes( + explicit_scopes: Option>, + configured_scopes: Option>, + discovered_scopes: Option>, +) -> ResolvedMcpOAuthScopes { + if let Some(scopes) = explicit_scopes { + return ResolvedMcpOAuthScopes { + scopes, + source: McpOAuthScopesSource::Explicit, + }; + } + + if let Some(scopes) = configured_scopes { + return ResolvedMcpOAuthScopes { + scopes, + source: McpOAuthScopesSource::Configured, + }; + } + + if let Some(scopes) = discovered_scopes + && !scopes.is_empty() + { + return ResolvedMcpOAuthScopes { + scopes, + source: McpOAuthScopesSource::Discovered, + }; + } + + ResolvedMcpOAuthScopes { + scopes: Vec::new(), + source: McpOAuthScopesSource::Empty, + } +} + +pub fn should_retry_without_scopes(scopes: &ResolvedMcpOAuthScopes, error: &anyhow::Error) -> bool { + scopes.source == McpOAuthScopesSource::Discovered + && error.downcast_ref::().is_some() +} + +pub async fn compute_auth_statuses<'a, I>( + servers: I, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + auth: Option<&CodexAuth>, + runtime_context: &McpRuntimeContext, +) -> HashMap +where + I: IntoIterator, +{ + let futures = servers.into_iter().map(|(name, server)| { + let name = name.clone(); + let redirect_mode = if server.is_agent_plugin() { + StreamableHttpRedirectMode::AgentPluginV1 + } else { + StreamableHttpRedirectMode::Legacy + }; + let config = server.config().clone(); + let runtime_context = runtime_context.clone(); + async move { + let auth_state = match compute_auth_status( + &name, + &config, + store_mode, + keyring_backend_kind, + auth, + &runtime_context, + redirect_mode, + ) + .await + { + Ok(status) => status, + Err(error) => { + warn!("failed to determine auth status for MCP server `{name}`: {error:?}"); + McpAuthState::Unknown + } + }; + let entry = McpAuthStatusEntry { + config: Some(config), + auth_state, + }; + (name, entry) + } + }); + + join_all(futures).await.into_iter().collect() +} + +async fn compute_auth_status( + server_name: &str, + config: &McpServerConfig, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + auth: Option<&CodexAuth>, + runtime_context: &McpRuntimeContext, + redirect_mode: StreamableHttpRedirectMode, +) -> Result { + if !config.enabled { + return Ok(McpAuthState::Unsupported); + } + if matches!(config.auth, McpServerAuth::EmaAuth) { + // EMA connections are not enabled until the runtime stage of the stack. + return Ok(McpAuthState::Unsupported); + } + let has_runtime_auth = matches!(config.auth, McpServerAuth::ChatGpt) + && auth.is_some_and(CodexAuth::uses_codex_backend) + && matches!( + &config.transport, + McpServerTransportConfig::StreamableHttp { + bearer_token_env_var: None, + .. + } + ); + + if matches!(config.auth, McpServerAuth::ChatGpt) && !config.is_local_environment() { + return Ok(if has_explicit_http_authorization(config) { + McpAuthState::BearerToken + } else { + McpAuthState::Unsupported + }); + } + + if has_runtime_auth { + return Ok(McpAuthState::BearerToken); + } + + match &config.transport { + McpServerTransportConfig::Stdio { .. } => Ok(McpAuthState::Unsupported), + McpServerTransportConfig::StreamableHttp { + url, + bearer_token_env_var, + http_headers, + env_http_headers, + http_headers_helper, + } => { + if http_headers_helper.is_some() { + // Status inspection must not execute an arbitrary local helper. Existing + // credentials remain reportable; otherwise discovery waits for startup/login. + return Ok(determine_streamable_http_auth_status_from_credentials( + config.oauth_credential_name(server_name).as_ref(), + url, + bearer_token_env_var.as_deref(), + http_headers.clone(), + env_http_headers.clone(), + store_mode, + keyring_backend_kind, + )? + .unwrap_or(McpAuthState::Unknown)); + } + let http_client = runtime_context + .resolve_http_client(server_name, config) + .map_err(anyhow::Error::msg)?; + let discovery_timeout = if config.is_local_environment() { + OAuthDiscoveryTimeout::LOCAL + } else { + OAuthDiscoveryTimeout::Requested + }; + let oauth_credential_name = config.oauth_credential_name(server_name); + determine_streamable_http_auth_status( + oauth_credential_name.as_ref(), + url, + bearer_token_env_var.as_deref(), + http_headers.clone(), + env_http_headers.clone(), + store_mode, + keyring_backend_kind, + http_client, + discovery_timeout, + redirect_mode, + ) + .boxed() + .await + } + } +} + +#[cfg(test)] +mod tests { + use anyhow::anyhow; + use pretty_assertions::assert_eq; + + use super::McpOAuthScopesSource; + use super::OAuthProviderError; + use super::ResolvedMcpOAuthScopes; + use super::resolve_oauth_callback; + use super::resolve_oauth_scopes; + use super::should_retry_without_scopes; + + #[test] + fn callback_resolution_preserves_registered_and_legacy_clients() -> anyhow::Result<()> { + for (client_id, saved_callback, global_callback, expected_callback) in [ + ( + Some("registered-client"), + Some("http://127.0.0.1/callback"), + Some("https://override.example/callback"), + Some("http://127.0.0.1/callback"), + ), + ( + Some("legacy-client"), + None, + None, + Some("http://127.0.0.1/callback/epMNJ6P1xGQ9"), + ), + ( + None, + Some("https://plugin.example/callback"), + Some("https://override.example/callback"), + Some("https://plugin.example/callback"), + ), + ( + None, + None, + Some("https://override.example/callback"), + Some("https://override.example/callback"), + ), + ] { + let server = serde_json::from_value(serde_json::json!({ + "url": "https://mcp.example.com/mcp", + "oauth": { + "client_id": client_id, + "callback_url": saved_callback, + }, + }))?; + assert_eq!( + resolve_oauth_callback(&server, "https://mcp.example.com/mcp", global_callback)? + .as_deref(), + expected_callback + ); + } + + Ok(()) + } + + #[test] + fn resolve_oauth_scopes_prefers_explicit() { + let resolved = resolve_oauth_scopes( + Some(vec!["explicit".to_string()]), + Some(vec!["configured".to_string()]), + Some(vec!["discovered".to_string()]), + ); + + assert_eq!( + resolved, + ResolvedMcpOAuthScopes { + scopes: vec!["explicit".to_string()], + source: McpOAuthScopesSource::Explicit, + } + ); + } + + #[test] + fn resolve_oauth_scopes_prefers_configured_over_discovered() { + let resolved = resolve_oauth_scopes( + /*explicit_scopes*/ None, + Some(vec!["configured".to_string()]), + Some(vec!["discovered".to_string()]), + ); + + assert_eq!( + resolved, + ResolvedMcpOAuthScopes { + scopes: vec!["configured".to_string()], + source: McpOAuthScopesSource::Configured, + } + ); + } + + #[test] + fn resolve_oauth_scopes_uses_discovered_when_needed() { + let resolved = resolve_oauth_scopes( + /*explicit_scopes*/ None, + /*configured_scopes*/ None, + Some(vec!["discovered".to_string()]), + ); + + assert_eq!( + resolved, + ResolvedMcpOAuthScopes { + scopes: vec!["discovered".to_string()], + source: McpOAuthScopesSource::Discovered, + } + ); + } + + #[test] + fn resolve_oauth_scopes_preserves_explicitly_empty_configured_scopes() { + let resolved = resolve_oauth_scopes( + /*explicit_scopes*/ None, + Some(Vec::new()), + Some(vec!["ignored".into()]), + ); + + assert_eq!( + resolved, + ResolvedMcpOAuthScopes { + scopes: Vec::new(), + source: McpOAuthScopesSource::Configured, + } + ); + } + + #[test] + fn resolve_oauth_scopes_falls_back_to_empty() { + let resolved = resolve_oauth_scopes( + /*explicit_scopes*/ None, /*configured_scopes*/ None, + /*discovered_scopes*/ None, + ); + + assert_eq!( + resolved, + ResolvedMcpOAuthScopes { + scopes: Vec::new(), + source: McpOAuthScopesSource::Empty, + } + ); + } + + #[test] + fn should_retry_without_scopes_only_for_discovered_provider_errors() { + let discovered = ResolvedMcpOAuthScopes { + scopes: vec!["scope".to_string()], + source: McpOAuthScopesSource::Discovered, + }; + let provider_error = anyhow!(OAuthProviderError::new( + Some("invalid_scope".to_string()), + Some("scope rejected".to_string()), + )); + + assert!(should_retry_without_scopes(&discovered, &provider_error)); + + let configured = ResolvedMcpOAuthScopes { + scopes: vec!["scope".to_string()], + source: McpOAuthScopesSource::Configured, + }; + assert!(!should_retry_without_scopes(&configured, &provider_error)); + assert!(!should_retry_without_scopes( + &discovered, + &anyhow!("timed out waiting for OAuth callback"), + )); + } +} diff --git a/codex-rs/codex-mcp/src/mcp/mod.rs b/codex-rs/codex-mcp/src/mcp/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..43d9a43c0b06a30c51e2357da343aaf64be1aa92 --- /dev/null +++ b/codex-rs/codex-mcp/src/mcp/mod.rs @@ -0,0 +1,841 @@ +pub use auth::McpAuthStatusEntry; +pub use auth::McpOAuthLoginConfig; +pub use auth::McpOAuthLoginSupport; +pub use auth::McpOAuthScopesSource; +pub use auth::ResolvedMcpOAuthScopes; +pub use auth::compute_auth_statuses; +pub use auth::discover_supported_scopes; +pub use auth::oauth_login_support; +pub use auth::resolve_oauth_callback; +pub use auth::resolve_oauth_scopes; +pub use auth::should_retry_without_scopes; + +pub(crate) mod auth; + +use std::collections::HashMap; +use std::collections::HashSet; +use std::env; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::Duration; + +use codex_config::ConfigLayerStack; +use codex_config::Constrained; +use codex_config::McpEnterpriseManagedAuthConfig; +use codex_config::McpServerAuth; +use codex_config::McpServerConfig; +use codex_config::McpServerTransportConfig; +use codex_config::types::AppToolApproval; +use codex_config::types::ApprovalsReviewer; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_connectors::ConnectorRuntimeManager; +use codex_connectors::ConnectorSnapshot; +use codex_connectors::connector_runtime_context_key; +use codex_login::CodexAuth; +use codex_model_provider::CHATGPT_CODEX_BASE_URL; +use codex_protocol::mcp::ClientMcpExtensions; +use codex_protocol::mcp::McpServerInfo; +use codex_protocol::mcp::Resource; +use codex_protocol::mcp::ResourceTemplate; +use codex_protocol::mcp::Tool; +use codex_protocol::models::PermissionProfile; +use codex_protocol::protocol::AskForApproval; +use codex_protocol::protocol::McpAuthStatus; +use codex_rmcp_client::McpOAuthRefreshMode; +use codex_utils_path_uri::PathUri; +use rmcp::model::ElicitationCapability; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ReadResourceResult; +use serde_json::Value; +use tokio_util::sync::CancellationToken; + +use crate::McpProtocolMode; +use crate::McpServerSource; +use crate::ResolvedMcpCatalog; +use crate::connection_manager::McpConnectionSet; +use crate::runtime::McpPublicationGate; +use crate::runtime::McpRuntimeContext; +use crate::runtime::McpRuntimeInput; +use crate::runtime::McpStartupPolicy; +use crate::server::EffectiveMcpServer; +use crate::tools::ToolInfo; + +pub const CODEX_APPS_MCP_SERVER_NAME: &str = "codex_apps"; +const DEFAULT_CODEX_APPS_MCP_PRODUCT_SKU: &str = "codex"; +const MCP_TOOL_NAME_PREFIX: &str = "mcp"; +const MCP_TOOL_NAME_DELIMITER: &str = "__"; +const CODEX_CONNECTORS_TOKEN_ENV_VAR: &str = "CODEX_CONNECTORS_TOKEN"; + +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum McpSnapshotDetail { + #[default] + Full, + ToolsAndAuthOnly, +} + +impl McpSnapshotDetail { + fn include_resources(self) -> bool { + matches!(self, Self::Full) + } +} + +pub fn qualified_mcp_tool_name_prefix(server_name: &str) -> String { + sanitize_responses_api_tool_name(&format!( + "{MCP_TOOL_NAME_PREFIX}{MCP_TOOL_NAME_DELIMITER}{server_name}{MCP_TOOL_NAME_DELIMITER}" + )) +} + +/// Returns true when MCP permission prompts should resolve as approved instead +/// of being shown to the user. +pub fn mcp_permission_prompt_is_auto_approved( + approval_policy: AskForApproval, + permission_profile: &PermissionProfile, + context: McpPermissionPromptAutoApproveContext, +) -> bool { + if context.tool_approval_mode == Some(AppToolApproval::Approve) { + return true; + } + + if approval_policy != AskForApproval::Never { + return false; + } + + match permission_profile { + PermissionProfile::Disabled | PermissionProfile::External { .. } => true, + PermissionProfile::Managed { file_system, .. } => { + file_system.to_sandbox_policy().has_full_disk_write_access() + } + } +} + +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +pub struct McpPermissionPromptAutoApproveContext { + pub tool_approval_mode: Option, +} + +/// MCP runtime settings derived from `codex_core::config::Config`. +/// +/// Each published runtime and prepared call owns one immutable copy of these +/// settings, so its connection, approval policy, and sandbox authority cannot +/// change independently. Auth remains separate and is supplied explicitly to +/// runtime entry points such as [`effective_mcp_servers`]. +#[derive(Debug, Clone)] +pub struct McpConfig { + /// Base URL for ChatGPT-hosted app MCP servers, copied from the root config. + pub chatgpt_base_url: String, + /// Optional product SKU forwarded to the host-owned apps MCP server. + pub apps_mcp_product_sku: Option, + /// Codex home directory used for MCP OAuth state and app-tool cache files. + pub codex_home: PathBuf, + /// Trusted enterprise IdP inherited after normal catalog and policy resolution. + pub mcp_enterprise_managed_auth: Option, + pub xaa_enabled: bool, + /// Preferred credential store for MCP OAuth tokens. + pub mcp_oauth_credentials_store_mode: OAuthCredentialsStoreMode, + /// OAuth refresh ownership selected for new MCP connections. + pub oauth_refresh_mode: McpOAuthRefreshMode, + /// Backend used when MCP OAuth storage is configured for keyring-backed persistence. + pub auth_keyring_backend_kind: AuthKeyringBackendKind, + /// Optional fixed localhost callback port for MCP OAuth login. + pub mcp_oauth_callback_port: Option, + /// Optional OAuth redirect URI override for MCP login. + pub mcp_oauth_callback_url: Option, + /// How long a tool catalog capture waits for optional MCP servers to initialize. + /// + /// A zero duration disables the shared grace and waits for each server's + /// configured startup timeout instead. + pub optional_mcp_startup_grace: Duration, + /// Whether skill MCP dependency installation prompts are enabled. + pub skill_mcp_dependency_install_enabled: bool, + /// Approval policy used for MCP tool calls and MCP elicitation requests. + pub approval_policy: Constrained, + /// Permission profile captured with the connections and approval policy. + pub permission_profile: PermissionProfile, + /// Configuration layers used to evaluate Apps tool policy and reviewer selection. + pub config_layer_stack: ConfigLayerStack, + /// Default reviewer used when an Apps tool has no reviewer override. + pub approvals_reviewer: ApprovalsReviewer, + /// Working directories for the exact environment handles used by this runtime. + pub environment_cwds: HashMap, + /// Explicit server permissions; unresolved or unavailable servers have no entry. + pub server_permission_profiles: HashMap, + /// Optional path to `codex-linux-sandbox` for sandboxed MCP tool execution. + pub codex_linux_sandbox_exe: Option, + /// Whether to use legacy Landlock behavior in the MCP sandbox state. + // TODO(anp): Reconcile this runtime-wide copy with TurnEnvironment::sandbox_context + // for the environment that owns each MCP server. + pub use_legacy_landlock: bool, + /// Whether the app MCP integration is enabled by config. + /// + /// ChatGPT auth is checked separately before a materialized host-owned Apps + /// server can be used. + pub apps_enabled: bool, + /// Whether model-visible MCP tool namespaces should keep the legacy + /// `mcp__` prefix. + pub prefix_mcp_tool_names: bool, + /// MCP servers whose model-visible tool namespaces omit the `mcp__` prefix. + pub non_prefixed_mcp_tool_servers: Vec, + /// Protocol mode for servers other than the host-owned Codex Apps registration. + pub protocol_mode: McpProtocolMode, + /// Independent protocol mode for the trusted, HTTP Codex Apps registration. + pub host_owned_apps_protocol_mode: McpProtocolMode, + /// Client-side elicitation capabilities advertised during MCP initialization. + pub client_elicitation_capability: ElicitationCapability, + /// Resolved MCP registrations keyed by logical server name. + pub mcp_server_catalog: ResolvedMcpCatalog, + /// Plugin declarations used to attribute connector tools to plugin display names. + /// MCP registrations retain their own package attribution in the catalog. + pub connector_snapshot: ConnectorSnapshot, +} + +/// Default amount of time a tool catalog capture waits for optional MCP servers. +pub const DEFAULT_OPTIONAL_MCP_STARTUP_GRACE: Duration = Duration::from_secs(1); + +impl McpConfig { + /// Resolves enabled runtime servers against the exact attachment permissions being published. + pub fn set_server_permission_profiles( + &mut self, + servers: &HashMap, + environment_profiles: impl IntoIterator, + ) { + let environment_profiles = environment_profiles.into_iter().collect::>(); + self.server_permission_profiles = servers + .iter() + .filter(|(_, server)| server.enabled()) + .filter_map(|(server_name, _)| { + let server = self.mcp_server_catalog.server(server_name)?; + let permission_profile = if server + .source() + .is_host_owned_apps(server_name, server.config()) + { + &self.permission_profile + } else if let Some(permission_profile) = + environment_profiles.get(&server.config().environment_id) + { + permission_profile + } else if server.config().is_local_environment() + || matches!(server.source(), McpServerSource::SelectedPlugin(_)) + { + &self.permission_profile + } else { + return None; + }; + Some((server_name.clone(), permission_profile.clone())) + }) + .collect(); + } + + /// Returns this server's published permission authority. + pub fn permission_profile_for_server(&self, server_name: &str) -> Option<&PermissionProfile> { + self.server_permission_profiles.get(server_name) + } + + /// Standalone discovery and resource reads must not inherit thread execution authority. + pub fn for_threadless_operations(&self, servers: &HashMap) -> Self { + let mut config = self.clone(); + config.permission_profile = PermissionProfile::default(); + config.server_permission_profiles = servers + .iter() + .filter(|(_, server)| server.enabled()) + .map(|(name, _)| (name.clone(), PermissionProfile::default())) + .collect(); + config + } +} + +/// Plugin attribution and selection data derived from the current MCP configuration. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct ToolPluginContext { + plugin_display_names_by_connector_id: HashMap>, + disabled_connector_ids: HashSet, + plugin_display_names_by_mcp_server_name: HashMap>, + plugin_ids_by_mcp_server_name: HashMap, + selected_plugin_mcp_server_names: HashSet, +} + +impl ToolPluginContext { + pub(crate) fn allows_connector_id(&self, connector_id: Option<&str>) -> bool { + connector_id.is_none_or(|id| !self.disabled_connector_ids.contains(id)) + } + + pub fn plugin_display_names_for_connector_id(&self, connector_id: &str) -> &[String] { + self.plugin_display_names_by_connector_id + .get(connector_id) + .map(Vec::as_slice) + .unwrap_or(&[]) + } + + pub fn plugin_display_names_for_mcp_server_name(&self, server_name: &str) -> &[String] { + self.plugin_display_names_by_mcp_server_name + .get(server_name) + .map(Vec::as_slice) + .unwrap_or(&[]) + } + + pub fn plugin_id_for_mcp_server_name(&self, server_name: &str) -> Option<&str> { + self.plugin_ids_by_mcp_server_name + .get(server_name) + .map(String::as_str) + } + + pub(crate) fn is_selected_plugin_mcp_server(&self, server_name: &str) -> bool { + self.selected_plugin_mcp_server_names.contains(server_name) + } + + fn from_config(config: &McpConfig) -> Self { + let mut tool_plugin_context = Self { + disabled_connector_ids: config.connector_snapshot.disabled_connector_ids().clone(), + ..Self::default() + }; + for connector_id in config.connector_snapshot.connector_ids() { + tool_plugin_context + .plugin_display_names_by_connector_id + .insert( + connector_id.0.clone(), + config + .connector_snapshot + .plugin_display_names_for_connector_id(&connector_id.0) + .to_vec(), + ); + } + + for (server_name, attribution) in config + .mcp_server_catalog + .plugin_attributions_by_server_name() + { + tool_plugin_context + .plugin_display_names_by_mcp_server_name + .insert( + server_name.clone(), + vec![attribution.display_name().to_string()], + ); + tool_plugin_context + .plugin_ids_by_mcp_server_name + .insert(server_name, attribution.plugin_id().to_string()); + } + tool_plugin_context.selected_plugin_mcp_server_names.extend( + config + .mcp_server_catalog + .selected_plugin_server_names() + .map(str::to_string), + ); + + for plugin_names in tool_plugin_context + .plugin_display_names_by_connector_id + .values_mut() + .chain( + tool_plugin_context + .plugin_display_names_by_mcp_server_name + .values_mut(), + ) + { + plugin_names.sort_unstable(); + plugin_names.dedup(); + } + tool_plugin_context + } +} + +pub fn host_owned_codex_apps_enabled(config: &McpConfig, auth: Option<&CodexAuth>) -> bool { + config.apps_enabled && auth.is_some_and(CodexAuth::uses_codex_backend) +} + +pub fn configured_mcp_servers(config: &McpConfig) -> HashMap { + config.mcp_server_catalog.configured_servers() +} + +pub fn effective_mcp_servers( + config: &McpConfig, + auth: Option<&CodexAuth>, +) -> HashMap { + effective_mcp_servers_from_configured(configured_mcp_servers(config), config, auth) +} + +fn is_trusted_chatgpt_mcp_server( + transport: &McpServerTransportConfig, + chatgpt_base_url: &str, +) -> bool { + let McpServerTransportConfig::StreamableHttp { url, .. } = transport else { + return false; + }; + let Ok(server_url) = url::Url::parse(url) else { + return false; + }; + if !matches!(server_url.scheme(), "http" | "https") { + return false; + } + + if url::Url::parse(CHATGPT_CODEX_BASE_URL) + .ok() + .is_some_and(|chatgpt_url| server_url.origin() == chatgpt_url.origin()) + { + return true; + } + + url::Url::parse(chatgpt_base_url) + .ok() + .is_some_and(|staging_url| { + staging_url.scheme() == "https" + && staging_url.domain().is_some_and(|host| { + host == "chatgpt-staging.com" || host.ends_with(".chatgpt-staging.com") + }) + && server_url.origin() == staging_url.origin() + }) +} + +/// Converts a materialized server map to its auth-gated runtime view. +/// +/// Compatibility built-ins and extension overlays must already be reflected in +/// `configured_servers`; this function does not synthesize missing servers. +pub fn effective_mcp_servers_from_configured( + configured_servers: HashMap, + config: &McpConfig, + auth: Option<&CodexAuth>, +) -> HashMap { + let mut servers = configured_servers + .into_iter() + .map(|(name, mut server)| { + match server.auth.clone() { + McpServerAuth::ChatGpt => { + if !is_trusted_chatgpt_mcp_server(&server.transport, &config.chatgpt_base_url) { + server.auth = McpServerAuth::OAuth; + } + } + McpServerAuth::OAuth | McpServerAuth::EmaAuth => {} + } + let agent_plugin = config + .mcp_server_catalog + .server(&name) + .is_some_and(|server| server.source().is_agent_plugin()); + ( + name, + EffectiveMcpServer::configured(server).with_agent_plugin(agent_plugin), + ) + }) + .collect::>(); + if !host_owned_codex_apps_enabled(config, auth) { + servers.remove(CODEX_APPS_MCP_SERVER_NAME); + } + servers +} + +pub fn tool_plugin_context(config: &McpConfig) -> ToolPluginContext { + ToolPluginContext::from_config(config) +} + +pub async fn read_mcp_resource( + config: &McpConfig, + auth: Option<&CodexAuth>, + runtime_context: McpRuntimeContext, + codex_apps_tools_cache: ConnectorRuntimeManager, + tool_catalog_cache: crate::McpToolCatalogCache, + server: &str, + params: ReadResourceRequestParams, +) -> anyhow::Result { + let mut mcp_servers = effective_mcp_servers(config, auth); + mcp_servers.retain(|name, _| name == server); + let cancel_token = CancellationToken::new(); + let runtime_config = config.for_threadless_operations(&mcp_servers); + let manager = McpConnectionSet::new( + /*previous*/ None, + McpPublicationGate::already_published(), + McpRuntimeInput { + startup_policy: McpStartupPolicy::Eager, + config: Arc::new(runtime_config), + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + mcp_servers, + submit_id: String::new(), + tx_event: None, + startup_cancellation_token: cancel_token.clone(), + runtime_context, + codex_apps_tools_cache, + tool_catalog_cache, + codex_apps_tools_cache_key: connector_runtime_context_key(auth), + client_mcp_extensions: ClientMcpExtensions::default(), + auth: auth.cloned(), + auth_manager: None, + elicitation_reviewer: None, + elicitation_lifecycle: None, + }, + crate::elicitation::ElicitationRequestRouter::default(), + ) + .await; + + let result = manager.read_resource(server, params).await; + cancel_token.cancel(); + result +} + +#[derive(Debug, Clone)] +pub struct McpServerStatusSnapshot { + pub server_infos: HashMap, + pub server_capabilities: HashMap, + pub tools_by_server: HashMap>, + pub tools_errors: HashMap, + pub resources: HashMap>, + pub resource_templates: HashMap>, + pub auth_statuses: HashMap, + pub server_names: Vec, +} + +pub async fn collect_mcp_server_status_snapshot_with_detail( + config: &McpConfig, + auth: Option<&CodexAuth>, + submit_id: String, + runtime_context: McpRuntimeContext, + codex_apps_tools_cache: ConnectorRuntimeManager, + tool_catalog_cache: crate::McpToolCatalogCache, + detail: McpSnapshotDetail, +) -> McpServerStatusSnapshot { + let mcp_servers = effective_mcp_servers(config, auth); + if mcp_servers.is_empty() { + return McpServerStatusSnapshot { + server_infos: HashMap::new(), + server_capabilities: HashMap::new(), + tools_by_server: HashMap::new(), + tools_errors: HashMap::new(), + resources: HashMap::new(), + resource_templates: HashMap::new(), + auth_statuses: HashMap::new(), + server_names: Vec::new(), + }; + } + + let auth_status_entries = compute_auth_statuses( + mcp_servers.iter(), + config.mcp_oauth_credentials_store_mode, + config.auth_keyring_backend_kind, + auth, + &runtime_context, + ) + .await; + + let server_names = mcp_servers.keys().cloned().collect(); + + let cancel_token = CancellationToken::new(); + let runtime_config = config.for_threadless_operations(&mcp_servers); + let mcp_connection_manager = McpConnectionSet::new( + /*previous*/ None, + McpPublicationGate::already_published(), + McpRuntimeInput { + startup_policy: McpStartupPolicy::Eager, + config: Arc::new(runtime_config), + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + mcp_servers, + submit_id, + tx_event: None, + startup_cancellation_token: cancel_token.clone(), + runtime_context, + codex_apps_tools_cache, + tool_catalog_cache, + codex_apps_tools_cache_key: connector_runtime_context_key(auth), + client_mcp_extensions: ClientMcpExtensions::default(), + auth: auth.cloned(), + auth_manager: None, + elicitation_reviewer: None, + elicitation_lifecycle: None, + }, + crate::elicitation::ElicitationRequestRouter::default(), + ) + .await; + + let snapshot = collect_mcp_server_status_snapshot_from_manager( + &mcp_connection_manager, + auth_status_entries, + server_names, + detail, + ) + .await; + + cancel_token.cancel(); + + snapshot +} + +/// The Responses API requires tool names to match `^[a-zA-Z0-9_-]+$`. +/// MCP server/tool names are user-controlled, so sanitize the fully-qualified +/// name we expose to the model by replacing any disallowed character with `_`. +pub(crate) fn sanitize_responses_api_tool_name(name: &str) -> String { + let mut sanitized = String::with_capacity(name.len()); + for c in name.chars() { + if c.is_ascii_alphanumeric() || c == '_' { + sanitized.push(c); + } else { + sanitized.push('_'); + } + } + + if sanitized.is_empty() { + "_".to_string() + } else { + sanitized + } +} + +fn codex_apps_mcp_bearer_token_env_var() -> Option { + match env::var(CODEX_CONNECTORS_TOKEN_ENV_VAR) { + Ok(value) if !value.trim().is_empty() => Some(CODEX_CONNECTORS_TOKEN_ENV_VAR.to_string()), + Ok(_) => None, + Err(env::VarError::NotPresent) => None, + Err(env::VarError::NotUnicode(_)) => Some(CODEX_CONNECTORS_TOKEN_ENV_VAR.to_string()), + } +} + +fn normalize_codex_apps_base_url(base_url: &str) -> String { + let mut base_url = base_url.trim_end_matches('/').to_string(); + if (base_url.starts_with("https://chatgpt.com") + || base_url.starts_with("https://chat.openai.com")) + && !base_url.contains("/backend-api") + { + base_url = format!("{base_url}/backend-api"); + } + base_url +} + +fn codex_apps_mcp_url_for_base_url(base_url: &str) -> String { + let base_url = normalize_codex_apps_base_url(base_url); + let base_url = if base_url.contains("/backend-api") || base_url.contains("/api/codex") { + base_url + } else { + format!("{base_url}/api/codex") + }; + format!("{base_url}/ps/mcp") +} + +pub fn codex_apps_mcp_server_config( + chatgpt_base_url: &str, + apps_mcp_product_sku: Option<&str>, + originator: Option<&str>, +) -> McpServerConfig { + mcp_server_config_for_url( + codex_apps_mcp_url_for_base_url(chatgpt_base_url), + apps_mcp_product_sku, + originator, + McpServerAuth::ChatGpt, + ) +} + +/// Builds the ChatGPT-hosted plugin runtime served by plugin-service. +pub fn hosted_plugin_runtime_mcp_server_config( + chatgpt_base_url: &str, + apps_mcp_product_sku: Option<&str>, + originator: Option<&str>, +) -> McpServerConfig { + codex_apps_mcp_server_config(chatgpt_base_url, apps_mcp_product_sku, originator) +} + +fn mcp_server_config_for_url( + url: String, + apps_mcp_product_sku: Option<&str>, + originator: Option<&str>, + auth_mode: McpServerAuth, +) -> McpServerConfig { + let product_sku = apps_mcp_product_sku.unwrap_or(DEFAULT_CODEX_APPS_MCP_PRODUCT_SKU); + let mut http_headers = + HashMap::from([("X-OpenAI-Product-Sku".to_string(), product_sku.to_string())]); + if let Some(originator) = originator { + http_headers.insert("originator".to_string(), originator.to_string()); + } + let env_http_headers = None; + + McpServerConfig { + transport: McpServerTransportConfig::StreamableHttp { + url, + bearer_token_env_var: codex_apps_mcp_bearer_token_env_var(), + http_headers: Some(http_headers), + env_http_headers, + http_headers_helper: None, + }, + auth: auth_mode, + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: Some(Duration::from_secs(30)), + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + } +} + +fn protocol_tool_from_rmcp_tool(name: &str, tool: &rmcp::model::Tool) -> Option { + match serde_json::to_value(tool) { + Ok(value) => match Tool::from_mcp_value(value) { + Ok(tool) => Some(tool), + Err(err) => { + tracing::warn!("Failed to convert MCP tool '{name}': {err}"); + None + } + }, + Err(err) => { + tracing::warn!("Failed to serialize MCP tool '{name}': {err}"); + None + } + } +} + +fn auth_statuses_from_entries( + auth_status_entries: &HashMap, +) -> HashMap { + auth_status_entries + .iter() + .map(|(name, entry)| (name.clone(), McpAuthStatus::from(entry.auth_state))) + .collect::>() +} + +fn convert_mcp_resources( + resources: HashMap>, +) -> HashMap> { + resources + .into_iter() + .map(|(name, resources)| { + let resources = resources + .into_iter() + .filter_map(|resource| match serde_json::to_value(resource) { + Ok(value) => match Resource::from_mcp_value(value.clone()) { + Ok(resource) => Some(resource), + Err(err) => { + let (uri, resource_name) = match value { + Value::Object(obj) => ( + obj.get("uri") + .and_then(|v| v.as_str().map(ToString::to_string)), + obj.get("name") + .and_then(|v| v.as_str().map(ToString::to_string)), + ), + _ => (None, None), + }; + + tracing::warn!( + "Failed to convert MCP resource (uri={uri:?}, name={resource_name:?}): {err}" + ); + None + } + }, + Err(err) => { + tracing::warn!("Failed to serialize MCP resource: {err}"); + None + } + }) + .collect::>(); + (name, resources) + }) + .collect::>() +} + +fn convert_mcp_resource_templates( + resource_templates: HashMap>, +) -> HashMap> { + resource_templates + .into_iter() + .map(|(name, templates)| { + let templates = templates + .into_iter() + .filter_map(|template| match serde_json::to_value(template) { + Ok(value) => match ResourceTemplate::from_mcp_value(value.clone()) { + Ok(template) => Some(template), + Err(err) => { + let (uri_template, template_name) = match value { + Value::Object(obj) => ( + obj.get("uriTemplate") + .or_else(|| obj.get("uri_template")) + .and_then(|v| v.as_str().map(ToString::to_string)), + obj.get("name") + .and_then(|v| v.as_str().map(ToString::to_string)), + ), + _ => (None, None), + }; + + tracing::warn!( + "Failed to convert MCP resource template (uri_template={uri_template:?}, name={template_name:?}): {err}" + ); + None + } + }, + Err(err) => { + tracing::warn!("Failed to serialize MCP resource template: {err}"); + None + } + }) + .collect::>(); + (name, templates) + }) + .collect::>() +} + +async fn collect_mcp_server_status_snapshot_from_manager( + mcp_connection_manager: &McpConnectionSet, + auth_status_entries: HashMap, + server_names: Vec, + detail: McpSnapshotDetail, +) -> McpServerStatusSnapshot { + let ((server_infos, (tools, tools_errors)), resources, resource_templates) = tokio::join!( + async { + let server_infos = mcp_connection_manager.list_available_server_infos().await; + let tools = mcp_connection_manager.list_tools_with_errors().await; + (server_infos, tools) + }, + async { + if detail.include_resources() { + mcp_connection_manager.list_all_resources(|_| true).await + } else { + HashMap::new() + } + }, + async { + if detail.include_resources() { + mcp_connection_manager + .list_all_resource_templates(|_| true) + .await + } else { + HashMap::new() + } + }, + ); + + let mut tools_by_server = HashMap::>::new(); + for tool_info in tools { + let raw_tool_name = tool_info.tool.name.to_string(); + let Some(tool) = protocol_tool_from_rmcp_tool(&raw_tool_name, &tool_info.tool) else { + continue; + }; + let tool_name = tool.name.clone(); + tools_by_server + .entry(tool_info.server_name) + .or_default() + .insert(tool_name, tool); + } + + // Status-only discovery has no event channel. Report OAuth failures from the completed + // connection attempt instead of retaining the credential-presence status read beforehand. + let mut auth_statuses = auth_statuses_from_entries(&auth_status_entries); + for server_name in mcp_connection_manager.authentication_failed_servers().await { + if auth_statuses.get(&server_name) == Some(&McpAuthStatus::OAuth) { + auth_statuses.insert(server_name, McpAuthStatus::NotLoggedIn); + } + } + + McpServerStatusSnapshot { + server_capabilities: mcp_connection_manager.list_available_server_capabilities(), + server_infos, + tools_by_server, + tools_errors, + resources: convert_mcp_resources(resources), + resource_templates: convert_mcp_resource_templates(resource_templates), + auth_statuses, + server_names, + } +} + +#[cfg(test)] +#[path = "mod_tests.rs"] +pub(crate) mod tests; diff --git a/codex-rs/codex-mcp/src/mcp/mod_tests.rs b/codex-rs/codex-mcp/src/mcp/mod_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..e45d6bcacc872f0bd8f74bc60cf8871d4e0dc5e9 --- /dev/null +++ b/codex-rs/codex-mcp/src/mcp/mod_tests.rs @@ -0,0 +1,678 @@ +use super::*; +use crate::McpPluginAttribution; +use crate::McpServerRegistration; +use crate::connection_manager::tests::create_ready_async_managed_client; +use crate::mcp::auth::McpAuthStatusEntry; +use crate::rmcp_client::StartupOutcomeError; +use codex_config::Constrained; +use codex_config::types::AppToolApproval; +use codex_config::types::AuthKeyringBackendKind; +use codex_login::CodexAuth; +use codex_plugin::AppConnectorId; +use codex_plugin::PluginCapabilitySummary; +use codex_protocol::models::ManagedFileSystemPermissions; +use codex_protocol::models::PermissionProfile; +use codex_protocol::permissions::NetworkSandboxPolicy; +use codex_protocol::protocol::AskForApproval; +use codex_protocol::protocol::GranularApprovalConfig; +use codex_rmcp_client::McpAuthState; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use std::collections::HashMap; +use std::collections::HashSet; +use std::path::PathBuf; +use std::sync::Arc; + +#[tokio::test] +async fn status_snapshot_only_downgrades_oauth_authentication_failures() { + let auth_failure = StartupOutcomeError::Failed { + error: "OAuth refresh token was rejected".to_string(), + is_authentication_required: true, + }; + let mut manager = McpConnectionSet::empty(/*prefix_mcp_tool_names*/ true); + let mut auth_status_entries = HashMap::new(); + for (name, auth_state, startup_error) in [ + ( + "oauth-failed", + McpAuthState::OAuth, + Some(auth_failure.clone()), + ), + ("oauth-ready", McpAuthState::OAuth, None), + ( + "oauth-provider-error", + McpAuthState::OAuth, + Some(StartupOutcomeError::Failed { + error: "provider temporarily unavailable".to_string(), + is_authentication_required: false, + }), + ), + ("bearer", McpAuthState::BearerToken, Some(auth_failure)), + ] { + let mut client = create_ready_async_managed_client(Vec::new()).await; + if let Some(error) = startup_error { + client.client = futures::future::ready(Err(error)).boxed().shared(); + } + manager.insert_test_client(name, client); + auth_status_entries.insert( + name.to_string(), + McpAuthStatusEntry { + config: None, + auth_state, + }, + ); + } + + let server_names = auth_status_entries.keys().cloned().collect(); + let snapshot = collect_mcp_server_status_snapshot_from_manager( + &manager, + auth_status_entries, + server_names, + McpSnapshotDetail::ToolsAndAuthOnly, + ) + .await; + + assert_eq!( + snapshot.auth_statuses, + HashMap::from([ + ("oauth-failed".to_string(), McpAuthStatus::NotLoggedIn), + ("oauth-ready".to_string(), McpAuthStatus::OAuth), + ("oauth-provider-error".to_string(), McpAuthStatus::OAuth), + ("bearer".to_string(), McpAuthStatus::BearerToken), + ]), + ); +} + +pub(crate) fn test_mcp_config(codex_home: PathBuf) -> McpConfig { + McpConfig { + chatgpt_base_url: "https://chatgpt.com".to_string(), + apps_mcp_product_sku: None, + codex_home, + mcp_enterprise_managed_auth: None, + xaa_enabled: false, + mcp_oauth_credentials_store_mode: OAuthCredentialsStoreMode::default(), + oauth_refresh_mode: McpOAuthRefreshMode::Legacy, + auth_keyring_backend_kind: AuthKeyringBackendKind::default(), + mcp_oauth_callback_port: None, + mcp_oauth_callback_url: None, + optional_mcp_startup_grace: DEFAULT_OPTIONAL_MCP_STARTUP_GRACE, + skill_mcp_dependency_install_enabled: true, + approval_policy: Constrained::allow_any(AskForApproval::OnRequest), + permission_profile: PermissionProfile::default(), + config_layer_stack: codex_config::ConfigLayerStack::default(), + approvals_reviewer: codex_config::types::ApprovalsReviewer::default(), + environment_cwds: HashMap::new(), + server_permission_profiles: HashMap::new(), + codex_linux_sandbox_exe: None, + use_legacy_landlock: false, + apps_enabled: false, + prefix_mcp_tool_names: true, + non_prefixed_mcp_tool_servers: Vec::new(), + protocol_mode: McpProtocolMode::Legacy, + host_owned_apps_protocol_mode: McpProtocolMode::Legacy, + client_elicitation_capability: ElicitationCapability::default(), + mcp_server_catalog: ResolvedMcpCatalog::default(), + connector_snapshot: codex_connectors::ConnectorSnapshot::default(), + } +} + +pub(crate) fn test_elicitation_config( + server_name: &str, + approval_policy: AskForApproval, + permission_profile: PermissionProfile, +) -> Arc { + let mut config = test_mcp_config(PathBuf::new()); + config.approval_policy = Constrained::allow_any(approval_policy); + config.permission_profile = permission_profile.clone(); + config + .server_permission_profiles + .insert(server_name.to_string(), permission_profile); + Arc::new(config) +} + +#[test] +fn ema_catalog_supports_configured_installed_and_selected_plugins_without_widening_policy() { + let server: McpServerConfig = serde_json::from_value(serde_json::json!({ + "url": "https://resource.example/mcp", "auth": "ema_auth", + "oauth": { "client_id": "resource-client", "authorization_server_issuer": "https://as.example" } + })).unwrap(); + let plugin = McpPluginAttribution::agent_plugin("plugin@test".into(), "Plugin".into()); + let mut catalog = ResolvedMcpCatalog::builder(); + catalog.register(McpServerRegistration::from_config( + "configured".into(), + server.clone(), + )); + catalog.register(McpServerRegistration::from_plugin( + "installed".into(), + plugin.clone(), + /*plugin_order*/ 0, + server.clone(), + )); + catalog.register(McpServerRegistration::from_selected_plugin( + "selected".into(), + plugin, + /*selection_order*/ 0, + server, + )); + let mut config = test_mcp_config(PathBuf::new()); + let idp = codex_config::McpServerIdpOAuthConfig { + issuer: "https://idp.example".into(), + client_id: "enterprise-client".into(), + }; + let deny_all = codex_protocol::mcp_policy::EnvironmentMcpPolicy { + servers: Some(Default::default()), + plugins: None, + }; + for (xaa_enabled, denied) in [(true, false), (true, true), (false, false)] { + let mut catalog = catalog.clone(); + if xaa_enabled { + catalog.enable_ema(idp.clone()); + } + config.mcp_server_catalog = catalog.build_with_environment_authority(|_| { + if denied { + crate::McpEnvironmentAuthority::Restricted(&deny_all) + } else { + crate::McpEnvironmentAuthority::Unrestricted + } + }); + let servers = effective_mcp_servers(&config, /*auth*/ None); + for name in ["configured", "installed", "selected"] { + assert_eq!(servers[name].enabled(), xaa_enabled && !denied, "{name}"); + assert_eq!( + servers[name].config().oauth_idp(), + xaa_enabled.then_some(&idp) + ); + assert_eq!(servers[name].config().auth, McpServerAuth::EmaAuth); + } + } +} + +#[test] +fn qualified_mcp_tool_name_prefix_sanitizes_server_names_without_lowercasing() { + assert_eq!( + qualified_mcp_tool_name_prefix("Some-Server"), + "mcp__Some_Server__".to_string() + ); +} + +#[test] +fn mcp_server_permissions_handle_unattached_and_threadless_servers() { + let mut config = test_mcp_config(PathBuf::new()); + config.permission_profile = PermissionProfile::Disabled; + let mut missing_server = codex_apps_mcp_server_config( + "https://example.com", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ); + missing_server.environment_id = "missing".to_string(); + let mut selected_server = missing_server.clone(); + selected_server.environment_id = "unattached".to_string(); + let mut catalog = ResolvedMcpCatalog::builder(); + catalog.register(McpServerRegistration::from_config( + "missing".to_string(), + missing_server, + )); + catalog.register(McpServerRegistration::from_selected_plugin( + "selected".to_string(), + McpPluginAttribution::new("selected@test".to_string(), "Selected".to_string()), + /*selection_order*/ 0, + selected_server, + )); + config.mcp_server_catalog = catalog.build(); + let servers = effective_mcp_servers(&config, /*auth*/ None); + assert_eq!(config.permission_profile_for_server("selected"), None); + config.set_server_permission_profiles(&servers, std::iter::empty()); + + assert_eq!( + config.permission_profile_for_server("selected"), + Some(&PermissionProfile::Disabled) + ); + assert_eq!(config.permission_profile_for_server("missing"), None); + + let config = config.for_threadless_operations(&servers); + assert_eq!( + config.permission_profile_for_server("selected"), + Some(&PermissionProfile::default()) + ); +} + +#[test] +fn mcp_prompt_auto_approval_honors_unrestricted_managed_profiles() { + assert!(mcp_permission_prompt_is_auto_approved( + AskForApproval::Never, + &PermissionProfile::Managed { + file_system: ManagedFileSystemPermissions::Unrestricted, + network: NetworkSandboxPolicy::Enabled, + }, + McpPermissionPromptAutoApproveContext::default(), + )); + assert!(mcp_permission_prompt_is_auto_approved( + AskForApproval::Never, + &PermissionProfile::Managed { + file_system: ManagedFileSystemPermissions::Unrestricted, + network: NetworkSandboxPolicy::Restricted, + }, + McpPermissionPromptAutoApproveContext::default(), + )); + assert!(!mcp_permission_prompt_is_auto_approved( + AskForApproval::Never, + &PermissionProfile::read_only(), + McpPermissionPromptAutoApproveContext::default(), + )); + assert!(!mcp_permission_prompt_is_auto_approved( + AskForApproval::OnRequest, + &PermissionProfile::Managed { + file_system: ManagedFileSystemPermissions::Unrestricted, + network: NetworkSandboxPolicy::Enabled, + }, + McpPermissionPromptAutoApproveContext::default(), + )); +} + +#[test] +fn mcp_prompt_auto_approval_honors_approved_tools_in_all_permission_modes() { + for approval_policy in [ + AskForApproval::UnlessTrusted, + AskForApproval::OnRequest, + AskForApproval::Granular(GranularApprovalConfig { + sandbox_approval: true, + rules: true, + skill_approval: true, + request_permissions: true, + mcp_elicitations: true, + }), + AskForApproval::Never, + ] { + assert!(mcp_permission_prompt_is_auto_approved( + approval_policy, + &PermissionProfile::read_only(), + McpPermissionPromptAutoApproveContext { + tool_approval_mode: Some(AppToolApproval::Approve), + }, + )); + } + + assert!(!mcp_permission_prompt_is_auto_approved( + AskForApproval::OnRequest, + &PermissionProfile::read_only(), + McpPermissionPromptAutoApproveContext { + tool_approval_mode: Some(AppToolApproval::Auto), + }, + )); +} + +#[test] +fn mcp_prompt_auto_approval_rejects_auto_mode_in_default_permission_mode() { + assert!(!mcp_permission_prompt_is_auto_approved( + AskForApproval::OnRequest, + &PermissionProfile::read_only(), + McpPermissionPromptAutoApproveContext { + tool_approval_mode: Some(AppToolApproval::Auto), + }, + )); +} + +#[test] +fn tool_plugin_context_collects_app_and_mcp_sources() { + let mut config = test_mcp_config(PathBuf::new()); + let mut catalog = ResolvedMcpCatalog::builder(); + catalog.register(McpServerRegistration::from_plugin( + "alpha".to_string(), + McpPluginAttribution::new("alpha@test".to_string(), "alpha-plugin".to_string()), + /*plugin_order*/ 0, + codex_apps_mcp_server_config( + "https://alpha.example", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ), + )); + config.mcp_server_catalog = catalog.build(); + config.connector_snapshot = + codex_connectors::ConnectorSnapshot::from_plugin_capability_summaries(&[ + PluginCapabilitySummary { + config_name: "alpha@test".to_string(), + display_name: "alpha-plugin".to_string(), + plugin_namespace: None, + app_connector_ids: vec![AppConnectorId("connector_example".to_string())], + mcp_server_names: vec!["alpha".to_string()], + ..PluginCapabilitySummary::default() + }, + PluginCapabilitySummary { + config_name: "beta@test".to_string(), + display_name: "beta-plugin".to_string(), + plugin_namespace: None, + app_connector_ids: vec![ + AppConnectorId("connector_example".to_string()), + AppConnectorId("connector_gmail".to_string()), + ], + mcp_server_names: vec!["beta".to_string()], + ..PluginCapabilitySummary::default() + }, + ]); + let provenance = tool_plugin_context(&config); + + assert_eq!( + provenance, + ToolPluginContext { + disabled_connector_ids: HashSet::new(), + plugin_display_names_by_connector_id: HashMap::from([ + ( + "connector_example".to_string(), + vec!["alpha-plugin".to_string(), "beta-plugin".to_string()], + ), + ( + "connector_gmail".to_string(), + vec!["beta-plugin".to_string()], + ), + ]), + plugin_display_names_by_mcp_server_name: HashMap::from([( + "alpha".to_string(), + vec!["alpha-plugin".to_string()], + )]), + plugin_ids_by_mcp_server_name: HashMap::from([( + "alpha".to_string(), + "alpha@test".to_string(), + )]), + selected_plugin_mcp_server_names: HashSet::new(), + } + ); + assert_eq!( + provenance.plugin_id_for_mcp_server_name("alpha"), + Some("alpha@test") + ); + assert_eq!(provenance.plugin_id_for_mcp_server_name("beta"), None); +} + +#[test] +fn selected_mcp_attribution_does_not_join_an_unrelated_local_summary() { + let mut config = test_mcp_config(PathBuf::new()); + let mut catalog = ResolvedMcpCatalog::builder(); + catalog.register(McpServerRegistration::from_selected_plugin( + "github".to_string(), + McpPluginAttribution::new( + "shared-plugin-id".to_string(), + "Executor GitHub".to_string(), + ), + /*selection_order*/ 0, + codex_apps_mcp_server_config( + "https://github.example", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ), + )); + config.mcp_server_catalog = catalog.build(); + config.connector_snapshot = + codex_connectors::ConnectorSnapshot::from_plugin_capability_summaries(&[ + PluginCapabilitySummary { + config_name: "shared-plugin-id".to_string(), + display_name: "Local GitHub".to_string(), + plugin_namespace: None, + mcp_server_names: vec!["github".to_string()], + ..PluginCapabilitySummary::default() + }, + ]); + + let provenance = tool_plugin_context(&config); + + assert_eq!( + provenance, + ToolPluginContext { + disabled_connector_ids: HashSet::new(), + plugin_display_names_by_connector_id: HashMap::new(), + plugin_display_names_by_mcp_server_name: HashMap::from([( + "github".to_string(), + vec!["Executor GitHub".to_string()], + )]), + plugin_ids_by_mcp_server_name: HashMap::from([( + "github".to_string(), + "shared-plugin-id".to_string(), + )]), + selected_plugin_mcp_server_names: HashSet::from(["github".to_string()]), + } + ); + assert!(provenance.is_selected_plugin_mcp_server("github")); +} + +#[test] +fn codex_apps_mcp_url_for_base_url_uses_plugin_service_paths() { + assert_eq!( + codex_apps_mcp_url_for_base_url("https://chatgpt.com/backend-api"), + "https://chatgpt.com/backend-api/ps/mcp" + ); + assert_eq!( + codex_apps_mcp_url_for_base_url("https://chat.openai.com"), + "https://chat.openai.com/backend-api/ps/mcp" + ); + assert_eq!( + codex_apps_mcp_url_for_base_url("http://localhost:8080/api/codex"), + "http://localhost:8080/api/codex/ps/mcp" + ); + assert_eq!( + codex_apps_mcp_url_for_base_url("http://localhost:8080"), + "http://localhost:8080/api/codex/ps/mcp" + ); +} + +#[test] +fn codex_apps_server_config_uses_plugin_service_path() { + let config = codex_apps_mcp_server_config( + "https://chatgpt.com", + /*apps_mcp_product_sku*/ None, + /*originator*/ None, + ); + let url = match &config.transport { + McpServerTransportConfig::StreamableHttp { url, .. } => url, + _ => panic!("expected streamable http transport for codex apps"), + }; + + assert_eq!(url, "https://chatgpt.com/backend-api/ps/mcp"); +} + +#[test] +fn codex_apps_server_config_forwards_thread_originator_header() { + let config = codex_apps_mcp_server_config( + "https://chatgpt.com", + /*apps_mcp_product_sku*/ None, + Some("thread_originator"), + ); + + match &config.transport { + McpServerTransportConfig::StreamableHttp { + http_headers, + env_http_headers, + .. + } => { + assert_eq!( + http_headers, + &Some(HashMap::from([ + ("originator".to_string(), "thread_originator".to_string()), + ("X-OpenAI-Product-Sku".to_string(), "codex".to_string()), + ])) + ); + assert!(env_http_headers.is_none()); + } + other => panic!("expected streamable http transport, got {other:?}"), + } +} + +#[test] +fn codex_apps_server_config_sets_product_sku_header() { + for (configured_product_sku, expected_product_sku) in [(None, "codex"), (Some("tpp"), "tpp")] { + let config = codex_apps_mcp_server_config( + "https://chatgpt.com", + configured_product_sku, + /*originator*/ None, + ); + + match &config.transport { + McpServerTransportConfig::StreamableHttp { + http_headers, + env_http_headers, + .. + } => { + assert_eq!( + http_headers, + &Some(HashMap::from([( + "X-OpenAI-Product-Sku".to_string(), + expected_product_sku.to_string(), + )])) + ); + assert!(env_http_headers.is_none()); + } + other => panic!("expected streamable http transport, got {other:?}"), + } + } +} + +#[test] +fn codex_apps_server_config_forwards_originator_and_configured_product_sku_headers() { + let config = codex_apps_mcp_server_config( + "https://chatgpt.com", + Some("tpp"), + Some("thread_originator"), + ); + + match &config.transport { + McpServerTransportConfig::StreamableHttp { + http_headers, + env_http_headers, + .. + } => { + assert_eq!( + http_headers, + &Some(HashMap::from([ + ("originator".to_string(), "thread_originator".to_string()), + ("X-OpenAI-Product-Sku".to_string(), "tpp".to_string()), + ])) + ); + assert!(env_http_headers.is_none()); + } + other => panic!("expected streamable http transport, got {other:?}"), + } +} + +#[test] +fn effective_mcp_servers_preserve_chatgpt_auth_for_staging() { + for url in [ + "https://chatgpt-staging.com", + "https://preview.chatgpt-staging.com", + ] { + let mut config = test_mcp_config(PathBuf::new()); + config.chatgpt_base_url = url.to_string(); + let server = codex_apps_mcp_server_config( + url, /*apps_mcp_product_sku*/ None, /*originator*/ None, + ); + let configured = HashMap::from([("staging".to_string(), server)]); + let effective = + effective_mcp_servers_from_configured(configured, &config, /*auth*/ None); + + assert_eq!(effective["staging"].config().auth, McpServerAuth::ChatGpt); + } +} + +#[tokio::test] +async fn effective_mcp_servers_preserve_runtime_servers() { + let codex_home = tempfile::tempdir().expect("tempdir"); + let mut config = test_mcp_config(codex_home.path().to_path_buf()); + config.apps_enabled = true; + let auth = CodexAuth::create_dummy_chatgpt_auth_for_testing(); + + let mut catalog = ResolvedMcpCatalog::builder(); + catalog.register(McpServerRegistration::from_config( + "sample".to_string(), + McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: "https://user.example/mcp".to_string(), + bearer_token_env_var: None, + http_headers: None, + env_http_headers: None, + http_headers_helper: None, + }, + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + }, + )); + catalog.register(McpServerRegistration::from_config( + "docs".to_string(), + McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: "https://docs.example/mcp".to_string(), + bearer_token_env_var: None, + http_headers: None, + env_http_headers: None, + http_headers_helper: None, + }, + environment_id: codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + }, + )); + catalog.register(McpServerRegistration::from_config( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + codex_apps_mcp_server_config( + &config.chatgpt_base_url, + config.apps_mcp_product_sku.as_deref(), + /*originator*/ None, + ), + )); + config.mcp_server_catalog = catalog.build(); + + let effective = effective_mcp_servers(&config, Some(&auth)); + + let sample = effective.get("sample").expect("user server should exist"); + let docs = effective + .get("docs") + .expect("configured server should exist"); + let codex_apps = effective + .get(CODEX_APPS_MCP_SERVER_NAME) + .expect("codex apps server should exist"); + + let sample = sample.config(); + let docs = docs.config(); + let codex_apps = codex_apps.config(); + + match &sample.transport { + McpServerTransportConfig::StreamableHttp { url, .. } => { + assert_eq!(url, "https://user.example/mcp"); + } + other => panic!("expected streamable http transport, got {other:?}"), + } + match &docs.transport { + McpServerTransportConfig::StreamableHttp { url, .. } => { + assert_eq!(url, "https://docs.example/mcp"); + } + other => panic!("expected streamable http transport, got {other:?}"), + } + match &codex_apps.transport { + McpServerTransportConfig::StreamableHttp { url, .. } => { + assert_eq!(url, "https://chatgpt.com/backend-api/ps/mcp"); + } + other => panic!("expected streamable http transport, got {other:?}"), + } +} diff --git a/codex-rs/codex-mcp/src/openai_docs_source_attribution.rs b/codex-rs/codex-mcp/src/openai_docs_source_attribution.rs new file mode 100644 index 0000000000000000000000000000000000000000..a3baf5c87fc11780ee8755e52c338a3903c48ea9 --- /dev/null +++ b/codex-rs/codex-mcp/src/openai_docs_source_attribution.rs @@ -0,0 +1,56 @@ +use std::sync::Arc; + +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +use futures::future::BoxFuture; + +const OPENAI_DEVELOPER_DOCS_MCP_URL: &str = "https://developers.openai.com/mcp"; +const OPENAI_DEVELOPER_DOCS_MCP_CODEX_URL: &str = "https://developers.openai.com/mcp?source=codex"; + +pub(crate) fn maybe_with_openai_docs_source_attribution( + mcp_server_url: &str, + http_client: Arc, +) -> Arc { + if mcp_server_url == OPENAI_DEVELOPER_DOCS_MCP_URL { + Arc::new(OpenAiDocsHttpClient { http_client }) + } else { + http_client + } +} + +struct OpenAiDocsHttpClient { + http_client: Arc, +} + +impl OpenAiDocsHttpClient { + fn attribute_mcp_request(&self, params: &mut HttpRequestParams) { + if params.url == OPENAI_DEVELOPER_DOCS_MCP_URL { + params.url = OPENAI_DEVELOPER_DOCS_MCP_CODEX_URL.to_string(); + } + } +} + +impl HttpClient for OpenAiDocsHttpClient { + fn http_request( + &self, + mut params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + self.attribute_mcp_request(&mut params); + self.http_client.http_request(params) + } + + fn http_request_stream( + &self, + mut params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + self.attribute_mcp_request(&mut params); + self.http_client.http_request_stream(params) + } +} + +#[cfg(test)] +#[path = "openai_docs_source_attribution_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/openai_docs_source_attribution_tests.rs b/codex-rs/codex-mcp/src/openai_docs_source_attribution_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..7447debabf1b642ae0642ba01d73ed763e0225fc --- /dev/null +++ b/codex-rs/codex-mcp/src/openai_docs_source_attribution_tests.rs @@ -0,0 +1,92 @@ +use std::sync::Arc; +use std::sync::Mutex; + +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +use futures::FutureExt; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; + +use super::OPENAI_DEVELOPER_DOCS_MCP_CODEX_URL; +use super::OPENAI_DEVELOPER_DOCS_MCP_URL; +use super::maybe_with_openai_docs_source_attribution; + +#[derive(Default)] +struct RecordingHttpClient { + urls: Mutex>, +} + +impl HttpClient for RecordingHttpClient { + fn http_request( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + self.urls.lock().unwrap().push(params.url); + async { Err(ExecServerError::HttpRequest("test response".to_string())) }.boxed() + } + + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + self.urls.lock().unwrap().push(params.url); + async { Err(ExecServerError::HttpRequest("test response".to_string())) }.boxed() + } +} + +fn request(url: &str) -> HttpRequestParams { + HttpRequestParams { + method: "POST".to_string(), + url: url.to_string(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "test-request".to_string(), + stream_response: true, + } +} + +#[tokio::test] +async fn attributes_only_docs_mcp_requests() { + let recording_client = Arc::new(RecordingHttpClient::default()); + let http_client = maybe_with_openai_docs_source_attribution( + OPENAI_DEVELOPER_DOCS_MCP_URL, + recording_client.clone(), + ); + + let _ = http_client + .http_request_stream(request(OPENAI_DEVELOPER_DOCS_MCP_URL)) + .await; + let _ = http_client + .http_request(request( + "https://developers.openai.com/.well-known/oauth-protected-resource/mcp", + )) + .await; + + assert_eq!( + recording_client.urls.lock().unwrap().as_slice(), + [ + OPENAI_DEVELOPER_DOCS_MCP_CODEX_URL, + "https://developers.openai.com/.well-known/oauth-protected-resource/mcp", + ] + ); +} + +#[test] +fn leaves_other_mcp_clients_unwrapped() { + let recording_client = Arc::new(RecordingHttpClient::default()); + let http_client = maybe_with_openai_docs_source_attribution( + "https://example.com/mcp", + recording_client.clone(), + ); + + assert!(Arc::ptr_eq( + &http_client, + &(recording_client as Arc) + )); +} diff --git a/codex-rs/codex-mcp/src/pagination.rs b/codex-rs/codex-mcp/src/pagination.rs new file mode 100644 index 0000000000000000000000000000000000000000..9b22b3453e0d09494c1d6abd7c92cff16e61b3bc --- /dev/null +++ b/codex-rs/codex-mcp/src/pagination.rs @@ -0,0 +1,84 @@ +use std::collections::HashSet; +use std::future::Future; +use std::time::Duration; + +use anyhow::Result; +use anyhow::anyhow; +use rmcp::model::PaginatedRequestParams; + +const MAX_MCP_CATALOG_PAGES: usize = 100; +pub(crate) const MAX_MCP_CATALOG_ITEMS: usize = 2_048; +pub(crate) const MAX_CODEX_APPS_TOOL_CATALOG_ITEMS: usize = 8_192; +const MAX_MCP_PAGINATION_CURSOR_BYTES: usize = 64 * 1024; +const DEFAULT_MCP_PAGINATION_TIMEOUT: Duration = Duration::from_secs(30); + +pub(crate) async fn collect_paginated( + method: &str, + overall_timeout: Option, + fetch: F, +) -> Result> +where + F: FnMut(Option) -> Fut, + Fut: Future, Option)>>, +{ + collect_paginated_with_limit(method, overall_timeout, MAX_MCP_CATALOG_ITEMS, fetch).await +} + +pub(crate) async fn collect_paginated_with_limit( + method: &str, + overall_timeout: Option, + max_items: usize, + mut fetch: F, +) -> Result> +where + F: FnMut(Option) -> Fut, + Fut: Future, Option)>>, +{ + let collect = async { + let mut collected = Vec::new(); + let mut cursor = None; + let mut seen_cursors = HashSet::new(); + let mut page_count = 0; + + loop { + if page_count == MAX_MCP_CATALOG_PAGES { + return Err(anyhow!( + "{method} exceeded the pagination limit of {MAX_MCP_CATALOG_PAGES} pages" + )); + } + page_count += 1; + let params = cursor.as_ref().map(|next: &String| { + PaginatedRequestParams::default().with_cursor(Some(next.clone())) + }); + let (items, next_cursor) = fetch(params).await?; + if items.len() > max_items.saturating_sub(collected.len()) { + return Err(anyhow!( + "{method} exceeded the catalog limit of {max_items} items" + )); + } + collected.extend(items); + + let Some(next_cursor) = next_cursor else { + return Ok(collected); + }; + if next_cursor.len() > MAX_MCP_PAGINATION_CURSOR_BYTES { + return Err(anyhow!( + "{method} returned a pagination cursor exceeding {MAX_MCP_PAGINATION_CURSOR_BYTES} bytes" + )); + } + if !seen_cursors.insert(next_cursor.clone()) { + return Err(anyhow!("{method} returned a repeated pagination cursor")); + } + cursor = Some(next_cursor); + } + }; + + let timeout = overall_timeout.unwrap_or(DEFAULT_MCP_PAGINATION_TIMEOUT); + tokio::time::timeout(timeout, collect) + .await + .map_err(|_| anyhow!("{method} pagination timed out after {timeout:?}"))? +} + +#[cfg(test)] +#[path = "pagination_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/pagination_tests.rs b/codex-rs/codex-mcp/src/pagination_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..709acd299a06894c63769fefad1daefeb4822fed --- /dev/null +++ b/codex-rs/codex-mcp/src/pagination_tests.rs @@ -0,0 +1,226 @@ +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use anyhow::anyhow; +use pretty_assertions::assert_eq; + +use super::MAX_MCP_CATALOG_ITEMS; +use super::MAX_MCP_CATALOG_PAGES; +use super::MAX_MCP_PAGINATION_CURSOR_BYTES; +use super::collect_paginated; + +#[tokio::test] +async fn collects_all_pages_including_an_empty_cursor() { + let requests = Arc::new(Mutex::new(Vec::new())); + let observed = Arc::clone(&requests); + + let pages = collect_paginated("tools/list", /*overall_timeout*/ None, move |params| { + let observed = Arc::clone(&observed); + async move { + let cursor = params.and_then(|params| params.cursor); + observed.lock().expect("request lock").push(cursor.clone()); + match cursor.as_deref() { + None => Ok((vec!["first"], Some(String::new()))), + Some("") => Ok((vec!["second"], Some("last".to_string()))), + Some("last") => Ok((vec!["third"], None)), + Some(cursor) => Err(anyhow!("unexpected cursor: {cursor}")), + } + } + }) + .await + .expect("paginated request succeeds"); + + assert_eq!(pages, vec!["first", "second", "third"]); + assert_eq!( + *requests.lock().expect("request lock"), + vec![None, Some(String::new()), Some("last".to_string())] + ); +} + +#[tokio::test] +async fn rejects_nonconsecutive_repeated_cursors() { + let error = collect_paginated( + "resources/list", + /*overall_timeout*/ None, + |params| async move { + let cursor = params.and_then(|params| params.cursor); + let next = match cursor.as_deref() { + None => "first", + Some("first") => "second", + Some("second") => "first", + Some(cursor) => return Err(anyhow!("unexpected cursor: {cursor}")), + }; + Ok((Vec::<()>::new(), Some(next.to_string()))) + }, + ) + .await + .expect_err("a repeated cursor must fail"); + + assert_eq!( + error.to_string(), + "resources/list returned a repeated pagination cursor" + ); +} + +#[tokio::test] +async fn rejects_excessive_pagination_before_fetching_another_page() { + let requests = Arc::new(AtomicUsize::new(0)); + let observed = Arc::clone(&requests); + + let error = collect_paginated( + "tools/list", + /*overall_timeout*/ None, + move |_params| { + let observed = Arc::clone(&observed); + async move { + let page = observed.fetch_add(1, Ordering::Relaxed); + Ok((Vec::<()>::new(), Some(page.to_string()))) + } + }, + ) + .await + .expect_err("unbounded pagination must fail"); + + assert_eq!( + error.to_string(), + format!("tools/list exceeded the pagination limit of {MAX_MCP_CATALOG_PAGES} pages") + ); + assert_eq!(requests.load(Ordering::Relaxed), MAX_MCP_CATALOG_PAGES); +} + +#[tokio::test] +async fn rejects_a_page_exceeding_the_catalog_item_limit() { + let error = collect_paginated( + "tools/list", + /*overall_timeout*/ None, + |_params| async { Ok((vec![(); MAX_MCP_CATALOG_ITEMS + 1], None)) }, + ) + .await + .expect_err("an oversized catalog page must fail"); + + assert_eq!( + error.to_string(), + format!("tools/list exceeded the catalog limit of {MAX_MCP_CATALOG_ITEMS} items") + ); +} + +#[tokio::test] +async fn rejects_a_catalog_exceeding_the_item_limit_across_pages() { + let error = collect_paginated( + "tools/list", + /*overall_timeout*/ None, + |params| async move { + match params.and_then(|params| params.cursor) { + None => Ok((vec![(); MAX_MCP_CATALOG_ITEMS], Some("last".to_string()))), + Some(cursor) if cursor == "last" => Ok((vec![()], None)), + Some(cursor) => Err(anyhow!("unexpected cursor: {cursor}")), + } + }, + ) + .await + .expect_err("catalog items must be bounded across pages"); + + assert_eq!( + error.to_string(), + format!("tools/list exceeded the catalog limit of {MAX_MCP_CATALOG_ITEMS} items") + ); +} + +#[tokio::test] +async fn rejects_oversized_pagination_cursors_before_following_them() { + let requests = Arc::new(AtomicUsize::new(0)); + let observed = Arc::clone(&requests); + + let error = collect_paginated( + "resources/list", + /*overall_timeout*/ None, + move |_params| { + let observed = Arc::clone(&observed); + async move { + observed.fetch_add(1, Ordering::Relaxed); + Ok(( + Vec::<()>::new(), + Some("x".repeat(MAX_MCP_PAGINATION_CURSOR_BYTES + 1)), + )) + } + }, + ) + .await + .expect_err("an oversized pagination cursor must fail"); + + assert_eq!( + error.to_string(), + format!( + "resources/list returned a pagination cursor exceeding {MAX_MCP_PAGINATION_CURSOR_BYTES} bytes" + ) + ); + assert_eq!(requests.load(Ordering::Relaxed), 1); +} + +#[tokio::test] +async fn forwards_page_failures() { + let error = collect_paginated( + "resources/templates/list", + /*overall_timeout*/ None, + |_params| async { Err::<(Vec<()>, Option), _>(anyhow!("page failed")) }, + ) + .await + .expect_err("a page error must fail"); + + assert_eq!(error.to_string(), "page failed"); +} + +#[tokio::test(start_paused = true)] +async fn applies_a_default_timeout_when_no_timeout_is_configured() { + let error = collect_paginated( + "resources/list", + /*overall_timeout*/ None, + |_params| async { + tokio::time::sleep(Duration::from_secs(31)).await; + Ok((Vec::<()>::new(), None)) + }, + ) + .await + .expect_err("pagination without a configured timeout must still be bounded"); + + assert_eq!( + error.to_string(), + "resources/list pagination timed out after 30s" + ); +} + +#[tokio::test(start_paused = true)] +async fn applies_one_timeout_across_individually_timely_pages() { + let requests = Arc::new(Mutex::new(Vec::new())); + let observed = Arc::clone(&requests); + + let error = collect_paginated("tools/list", Some(Duration::from_secs(5)), move |params| { + let observed = Arc::clone(&observed); + async move { + let cursor = params.and_then(|params| params.cursor); + observed.lock().expect("request lock").push(cursor.clone()); + tokio::time::sleep(Duration::from_secs(2)).await; + + match cursor.as_deref() { + None => Ok((vec!["first"], Some("second".to_string()))), + Some("second") => Ok((vec!["second"], Some("third".to_string()))), + Some("third") => Ok((vec!["third"], None)), + Some(cursor) => Err(anyhow!("unexpected cursor: {cursor}")), + } + } + }) + .await + .expect_err("the combined page duration must exceed the shared timeout"); + + assert_eq!( + error.to_string(), + "tools/list pagination timed out after 5s" + ); + assert_eq!( + *requests.lock().expect("request lock"), + vec![None, Some("second".to_string()), Some("third".to_string())] + ); +} diff --git a/codex-rs/codex-mcp/src/plugin_config.rs b/codex-rs/codex-mcp/src/plugin_config.rs new file mode 100644 index 0000000000000000000000000000000000000000..f55ebf977df72f36ae90d33a3c97f2c0513efd1b --- /dev/null +++ b/codex-rs/codex-mcp/src/plugin_config.rs @@ -0,0 +1,303 @@ +use codex_config::McpServerAuth; +use codex_config::McpServerConfig; +use codex_config::McpServerEnvVar; +use codex_config::McpServerTransportConfig; +use codex_utils_path_uri::LegacyAppPathString; +use codex_utils_path_uri::PathUri; +use serde::Deserialize; +use serde_json::Map as JsonMap; +use serde_json::Value as JsonValue; +use std::collections::BTreeMap; +use std::path::Path; +use tracing::warn; + +#[path = "agent_plugin_config.rs"] +mod agent_plugin_config; + +pub use agent_plugin_config::parse_agent_plugin_mcp_config; + +#[derive(Clone, Copy, Debug)] +enum PluginMcpSource<'a> { + Host { + root: &'a Path, + }, + Environment { + root: &'a PathUri, + environment_id: &'a str, + }, +} + +/// One plugin MCP server that could not be normalized into runtime configuration. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct PluginMcpServerParseError { + pub name: String, + pub message: String, +} + +/// Valid servers and per-server errors parsed from one plugin MCP file. +#[derive(Debug, Default, PartialEq)] +pub struct PluginMcpConfigParseOutcome { + pub servers: BTreeMap, + pub errors: Vec, +} + +#[derive(Debug, Default, Deserialize)] +#[serde(rename_all = "camelCase")] +struct PluginMcpServersFile { + mcp_servers: BTreeMap, +} + +#[derive(Debug, Deserialize)] +#[serde(untagged)] +enum PluginMcpFile { + McpServersObject(PluginMcpServersFile), + ServerMap(BTreeMap), +} + +impl PluginMcpFile { + fn into_mcp_servers(self) -> BTreeMap { + match self { + Self::McpServersObject(file) => file.mcp_servers, + Self::ServerMap(mcp_servers) => mcp_servers, + } + } +} + +/// Parses the two supported plugin MCP file shapes and normalizes each server. +/// +/// Native plugin HTTP servers share the regular MCP transport configuration; +/// relative helper commands therefore use the session's local process cwd. +/// +/// Invalid individual servers are returned as errors without discarding valid +/// siblings. A malformed top-level document fails the whole parse. +pub fn parse_plugin_mcp_config( + plugin_root: &Path, + contents: &str, +) -> Result { + parse_plugin_mcp_config_from(contents, PluginMcpSource::Host { root: plugin_root }) +} + +/// Parses executor-owned plugin MCP config without interpreting the plugin root +/// as a path on the orchestrator host. +pub fn parse_executor_plugin_mcp_config( + plugin_root: &PathUri, + contents: &str, + environment_id: &str, +) -> Result { + parse_plugin_mcp_config_from( + contents, + PluginMcpSource::Environment { + root: plugin_root, + environment_id, + }, + ) +} + +impl PluginMcpSource<'_> { + fn display(self) -> String { + match self { + Self::Host { root } => root.display().to_string(), + Self::Environment { root, .. } => root.to_string(), + } + } +} + +fn parse_plugin_mcp_config_from( + contents: &str, + source: PluginMcpSource<'_>, +) -> Result { + let parsed = serde_json::from_str::(contents)?; + let mut outcome = PluginMcpConfigParseOutcome::default(); + + for (name, config_value) in parsed.into_mcp_servers() { + match normalize_plugin_mcp_server(config_value, source) { + Ok(config) => { + outcome.servers.insert(name, config); + } + Err(message) => outcome + .errors + .push(PluginMcpServerParseError { name, message }), + } + } + + Ok(outcome) +} + +fn normalize_plugin_mcp_server( + value: JsonValue, + source: PluginMcpSource<'_>, +) -> Result { + let mut object = normalize_plugin_mcp_server_value(value, source); + if let PluginMcpSource::Environment { + root, + environment_id, + } = source + { + object.insert( + "environment_id".to_string(), + JsonValue::String(environment_id.to_string()), + ); + if object.contains_key("command") { + match object.remove("cwd") { + Some(JsonValue::String(cwd)) => object.insert( + "cwd".to_string(), + JsonValue::String(environment_cwd(root, Some(&cwd))?.into_string()), + ), + Some(JsonValue::Null) | None => object.insert( + "cwd".to_string(), + JsonValue::String( + environment_cwd(root, /*configured_cwd*/ None)?.into_string(), + ), + ), + Some(value) => object.insert("cwd".to_string(), value), + }; + } + } + + let mut config = serde_json::from_value::(JsonValue::Object(object)) + .map_err(|err| err.to_string())?; + if matches!(config.auth, McpServerAuth::EmaAuth) { + return Err( + "plugin MCP declarations cannot select ema_auth; configure enterprise authentication in host policy" + .to_string(), + ); + } + if matches!(source, PluginMcpSource::Environment { .. }) { + bind_environment_env_vars(&mut config)?; + } + Ok(config) +} + +fn environment_cwd( + root: &PathUri, + configured_cwd: Option<&str>, +) -> Result { + let Some(configured_cwd) = configured_cwd else { + return Ok(root.clone().into()); + }; + let cwd = PathUri::parse(configured_cwd) + .or_else(|_| root.join(configured_cwd)) + .map_err(|err| format!("invalid cwd `{configured_cwd}`: {err}"))?; + if !cwd.starts_with(root) { + return Err(format!( + "cwd `{configured_cwd}` must remain within plugin root `{root}`" + )); + } + Ok(cwd.into()) +} + +fn bind_environment_env_vars(config: &mut McpServerConfig) -> Result<(), String> { + let is_local_environment = config.is_local_environment(); + let env_vars = match &mut config.transport { + McpServerTransportConfig::Stdio { env_vars, .. } => env_vars, + // Bearer credentials resolve on the executor; other header variables do not yet. + McpServerTransportConfig::StreamableHttp { + env_http_headers, .. + } => { + if is_local_environment { + return Ok(()); + } + if env_http_headers + .as_ref() + .is_some_and(|headers| !headers.is_empty()) + { + return Err( + "`env_http_headers` requires executor-side environment resolution for an executor-owned HTTP MCP" + .to_string(), + ); + } + return Ok(()); + } + }; + for env_var in env_vars { + match env_var { + McpServerEnvVar::Name(name) if !is_local_environment => { + *env_var = McpServerEnvVar::Config { + name: std::mem::take(name), + source: Some("remote".to_string()), + }; + } + McpServerEnvVar::Name(_) => {} + McpServerEnvVar::Config { name, source } => { + match (is_local_environment, source.as_deref()) { + (true, None | Some("local")) | (false, Some("remote")) => {} + (true, Some("remote")) => { + return Err(format!( + "env_vars entry `{name}` cannot use source `remote` in a local environment" + )); + } + (false, None) => *source = Some("remote".to_string()), + (false, Some("local")) => { + return Err(format!( + "env_vars entry `{name}` cannot use source `local` in an executor-owned plugin" + )); + } + (_, Some(source)) => unreachable!("validated env_vars source `{source}`"), + } + } + } + } + Ok(()) +} + +fn normalize_plugin_mcp_server_value( + value: JsonValue, + source: PluginMcpSource<'_>, +) -> JsonMap { + let mut object = match value { + JsonValue::Object(object) => object, + _ => return JsonMap::new(), + }; + + if let Some(JsonValue::String(transport_type)) = object.remove("type") { + match transport_type.as_str() { + "http" | "streamable_http" | "streamable-http" | "stdio" => {} + other => { + let plugin_display = source.display(); + warn!( + plugin = %plugin_display, + transport = other, + "plugin MCP server uses an unknown transport type" + ); + } + } + } + + if let Some(JsonValue::Object(mut oauth)) = object.remove("oauth") { + if let Some(callback_url) = oauth.remove("callbackUrl") { + oauth + .entry("callback_url".to_string()) + .or_insert(callback_url); + } + + if let Some(callback_port) = oauth.remove("callbackPort") { + oauth + .entry("callback_port".to_string()) + .or_insert(callback_port); + } + + if let Some(client_id) = oauth.remove("clientId") { + oauth.entry("client_id".to_string()).or_insert(client_id); + } + + if !oauth.is_empty() { + object.insert("oauth".to_string(), JsonValue::Object(oauth)); + } + } + + if let PluginMcpSource::Host { root } = source + && let Some(JsonValue::String(cwd)) = object.get("cwd") + && !Path::new(cwd).is_absolute() + { + object.insert( + "cwd".to_string(), + JsonValue::String(root.join(cwd).display().to_string()), + ); + } + + object +} + +#[cfg(test)] +#[path = "plugin_config_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/plugin_config_tests.rs b/codex-rs/codex-mcp/src/plugin_config_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..b6836e55c62d1ccf44831e6ce040fcb7f9c448e2 --- /dev/null +++ b/codex-rs/codex-mcp/src/plugin_config_tests.rs @@ -0,0 +1,967 @@ +use super::PluginMcpConfigParseOutcome; +use super::PluginMcpServerParseError; +use super::parse_agent_plugin_mcp_config; +use super::parse_executor_plugin_mcp_config; +use super::parse_plugin_mcp_config; +use codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID; +use codex_config::McpServerConfig; +use codex_config::McpServerEnvVar; +use codex_config::McpServerOAuthConfig; +use codex_config::McpServerTransportConfig; +use codex_utils_path_uri::LegacyAppPathString; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::path::Path; +use std::path::PathBuf; + +fn plugin_root() -> PathBuf { + std::env::current_dir() + .expect("current directory") + .join("plugin-root") +} + +fn plugin_root_uri(plugin_root: &Path) -> PathUri { + PathUri::from_host_native_path(plugin_root).expect("plugin root URI") +} + +#[test] +fn agent_plugin_placeholder_expansion_is_single_pass() { + let plugin_root = plugin_root().join("${PLUGIN_DATA}"); + let plugin_data_root = plugin_root + .parent() + .expect("plugin root parent") + .join("plugin-data"); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_data_root, + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers":{"demo":{ + "type":"stdio", + "command":"python", + "args":["${PLUGIN_ROOT}:${PLUGIN_DATA}"] + }} + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + let McpServerTransportConfig::Stdio { args, .. } = &outcome.servers["demo"].transport else { + panic!("expected stdio transport"); + }; + assert_eq!( + args, + &vec![format!( + "{}:{}", + plugin_root.display(), + plugin_data_root.display() + )] + ); +} + +#[test] +fn agent_plugin_mcp_expands_reserved_paths_and_maps_transports() { + let plugin_root = plugin_root(); + let plugin_data_root = plugin_root.parent().expect("parent").join("plugin-data"); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_data_root, + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers": { + "local": { + "type":"stdio", + "command":"python", + "args":["${PLUGIN_ROOT}/server.py", "${PLUGIN_DATA}/state.json"], + "env":{"CACHE":"${PLUGIN_DATA}/cache"}, + "cwd":"${PLUGIN_ROOT}/scripts" + }, + "remote": { + "type":"streamable-http", + "url":"https://example.com/mcp", + "headers":{"X-Plugin":"demo"} + } + } + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert!(outcome.errors.is_empty()); + let local = outcome.servers.get("local").expect("local server"); + let McpServerTransportConfig::Stdio { args, env, cwd, .. } = &local.transport else { + panic!("expected stdio transport"); + }; + assert_eq!( + args, + &vec![ + format!("{}/server.py", plugin_root.display()), + format!("{}/state.json", plugin_data_root.display()), + ] + ); + assert_eq!( + env.as_ref().expect("environment").get("PLUGIN_ROOT"), + Some(&plugin_root.display().to_string()) + ); + assert_eq!( + env.as_ref().expect("environment").get("PLUGIN_DATA"), + Some(&plugin_data_root.display().to_string()) + ); + assert_eq!( + cwd.as_ref(), + Some(&LegacyAppPathString::from_path( + &plugin_root.join("scripts") + )) + ); + + let remote = outcome.servers.get("remote").expect("remote server"); + let McpServerTransportConfig::StreamableHttp { http_headers, .. } = &remote.transport else { + panic!("expected HTTP transport"); + }; + assert_eq!( + http_headers + .as_ref() + .and_then(|headers| headers.get("X-Plugin")), + Some(&"demo".to_string()) + ); +} + +#[test] +fn agent_plugin_mcp_handles_portable_path_and_http_edge_cases() { + let plugin_root = plugin_root(); + let plugin_data_root = plugin_root.parent().expect("parent").join("plugin-data"); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_data_root, + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers": { + "contained":{"type":"stdio","command":"./bin/../server","cwd":"${PLUGIN_ROOT}/work/../data"}, + "redundant-separator":{"type":"stdio","command":".//bin/server"}, + "root-slash":{"type":"stdio","command":"python","cwd":"${PLUGIN_ROOT}/"}, + "data-slash":{"type":"stdio","command":"python","cwd":"${PLUGIN_DATA}/"}, + "headers":{"type":"streamable-http","url":"https://example.com/mcp","headers":{"aUtHoRiZaTiOn":"public-package-value","Content-Length":"0","HOST":"other.example.com","Proxy-Authorization":"public-package-value","Transfer-Encoding":"chunked","uSeR-aGeNt":"plugin-agent/1.0","X-Plugin":"demo","X-Plugin-Name":"café"}}, + "loopback":{"type":"streamable-http","url":"http://[::1]/mcp"} + } + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert!(outcome.errors.is_empty()); + assert_eq!( + outcome.servers.keys().collect::>(), + vec![ + "contained", + "data-slash", + "headers", + "loopback", + "redundant-separator", + "root-slash" + ] + ); + let McpServerTransportConfig::Stdio { command, cwd, .. } = + &outcome.servers["contained"].transport + else { + panic!("expected stdio transport"); + }; + assert_eq!(command, &plugin_root.join("server").display().to_string()); + assert_eq!( + cwd.as_ref(), + Some(&LegacyAppPathString::from_path(&plugin_root.join("data"))) + ); + let McpServerTransportConfig::Stdio { command, .. } = + &outcome.servers["redundant-separator"].transport + else { + panic!("expected stdio transport"); + }; + assert_eq!( + command, + &plugin_root.join("bin").join("server").display().to_string() + ); + for (server_name, expected_cwd) in [ + ("root-slash", plugin_root.as_path()), + ("data-slash", plugin_data_root.as_path()), + ] { + let McpServerTransportConfig::Stdio { cwd, .. } = &outcome.servers[server_name].transport + else { + panic!("expected stdio transport"); + }; + assert_eq!( + cwd.as_ref(), + Some(&LegacyAppPathString::from_path(expected_cwd)) + ); + } + let McpServerTransportConfig::StreamableHttp { http_headers, .. } = + &outcome.servers["headers"].transport + else { + panic!("expected HTTP transport"); + }; + assert_eq!( + http_headers, + &Some(HashMap::from([ + ("X-Plugin".to_string(), "demo".to_string()), + ("X-Plugin-Name".to_string(), "café".to_string()), + ])) + ); +} + +#[test] +fn agent_plugin_mcp_skips_invalid_server_without_disabling_siblings() { + let plugin_root = plugin_root(); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_root.join("data"), + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers": { + "valid":{"type":"stdio","command":"python"}, + "reserved":{"type":"stdio","command":"python","env":{"PLUGIN_ROOT":"bad"}} + } + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert_eq!(outcome.servers.keys().collect::>(), vec!["valid"]); + assert_eq!(outcome.errors.len(), 1); + assert_eq!(outcome.errors[0].name, "reserved"); + assert!( + outcome.errors[0] + .message + .contains("reserved variable `PLUGIN_ROOT`") + ); +} + +#[test] +fn agent_plugin_mcp_preserves_server_named_mcp_servers() { + let plugin_root = plugin_root(); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_root.join("data"), + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers": { + "mcpServers":{"type":"stdio","command":"first"}, + "sibling":{"type":"stdio","command":"second"} + } + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert!(outcome.errors.is_empty()); + assert_eq!( + outcome.servers.keys().collect::>(), + vec!["mcpServers", "sibling"] + ); +} + +#[test] +fn agent_plugin_mcp_preserves_arbitrary_server_names() { + let plugin_root = plugin_root(); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_root.join("data"), + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers":{"agent.smoke / local":{"type":"stdio","command":"python"}} + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert!(outcome.errors.is_empty()); + assert!(outcome.servers.contains_key("agent.smoke / local")); +} + +#[test] +fn agent_plugin_mcp_rejects_explicit_null_optional_fields() { + let plugin_root = plugin_root(); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_root.join("data"), + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers":{ + "cwd":{"type":"stdio","command":"python","cwd":null}, + "headers":{"type":"streamable-http","url":"https://example.com/mcp","headers":null} + } + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert!(outcome.servers.is_empty()); + assert_eq!(outcome.errors.len(), 2); + assert!( + outcome + .errors + .iter() + .any(|error| error.message.contains("`cwd`")) + ); + assert!( + outcome + .errors + .iter() + .any(|error| error.message.contains("`headers`")) + ); +} + +#[cfg(windows)] +#[test] +fn agent_plugin_mcp_rejects_reserved_environment_aliases_case_insensitively() { + let plugin_root = plugin_root(); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_root.join("data"), + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers":{"reserved":{"type":"stdio","command":"python","env":{"plugin_root":"bad"}}} + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert!(outcome.servers.is_empty()); + assert_eq!(outcome.errors.len(), 1); +} + +#[cfg(windows)] +#[test] +fn agent_plugin_mcp_overlays_windows_environment_case_insensitively() { + let plugin_root = plugin_root(); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_root.join("data"), + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers":{ + "configured":{"type":"stdio","command":"python","env":{"Path":"configured"}}, + "duplicate":{"type":"stdio","command":"python","env":{"PATH":"one","Path":"two"}} + } + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert_eq!( + outcome.servers.keys().collect::>(), + vec!["configured"] + ); + assert_eq!(outcome.errors.len(), 1); + let McpServerTransportConfig::Stdio { env, .. } = &outcome.servers["configured"].transport + else { + panic!("expected stdio transport"); + }; + assert_eq!( + env.as_ref().and_then(|env| env.get("PATH")), + Some(&"configured".to_string()) + ); +} + +#[cfg(unix)] +#[test] +fn agent_plugin_mcp_resolves_root_before_collapsing_parent_components() { + use std::os::unix::fs::symlink; + + let temp = tempfile::tempdir().expect("temporary directory"); + let base = temp.path().join("base"); + let outside = temp.path().join("outside"); + let outside_directory = outside.join("directory"); + let resolved_plugin_root = outside.join("plugin"); + let plugin_data_root = temp.path().join("plugin-data"); + std::fs::create_dir_all(&base).expect("create base directory"); + std::fs::create_dir_all(&outside_directory).expect("create symlink target"); + std::fs::create_dir_all(&resolved_plugin_root).expect("create resolved plugin root"); + std::fs::create_dir_all(&plugin_data_root).expect("create plugin data root"); + let canonical_plugin_root = resolved_plugin_root + .canonicalize() + .expect("canonical plugin root"); + symlink(&outside_directory, base.join("link")).expect("create root symlink"); + let plugin_root = base.join("link/../plugin"); + + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_data_root, + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers":{"demo":{"type":"stdio","command":"python"}} + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert!(outcome.errors.is_empty()); + let McpServerTransportConfig::Stdio { env, cwd, .. } = &outcome.servers["demo"].transport + else { + panic!("expected stdio transport"); + }; + assert_eq!( + env.as_ref().and_then(|env| env.get("PLUGIN_ROOT")), + Some(&canonical_plugin_root.display().to_string()) + ); + assert_eq!( + cwd.as_ref(), + Some(&LegacyAppPathString::from_path(&canonical_plugin_root)) + ); +} + +#[cfg(unix)] +#[test] +fn agent_plugin_mcp_rejects_missing_descendant_below_escaping_symlink() { + use std::os::unix::fs::symlink; + + let temp = tempfile::tempdir().expect("temporary directory"); + let plugin_root = temp.path().join("plugin"); + let plugin_data_root = temp.path().join("plugin-data"); + let outside = temp.path().join("outside"); + std::fs::create_dir_all(&plugin_root).expect("create plugin root"); + std::fs::create_dir_all(&plugin_data_root).expect("create plugin data root"); + std::fs::create_dir_all(&outside).expect("create outside directory"); + symlink(&outside, plugin_root.join("link")).expect("create escaping symlink"); + + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_data_root, + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers":{"escape":{"type":"stdio","command":"./link/missing"}} + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert!(outcome.servers.is_empty()); + assert_eq!(outcome.errors.len(), 1); + assert!(outcome.errors[0].message.contains("must remain within")); +} + +#[test] +fn agent_plugin_mcp_enforces_closed_transport_and_path_semantics() { + let plugin_root = plugin_root(); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_root.join("data"), + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers": { + "valid":{"type":"stdio","command":"python"}, + "command":{"type":"stdio","command":"../server"}, + "escape":{"type":"stdio","command":"./../server"}, + "cwd":{"type":"stdio","command":"python","cwd":"${PLUGIN_ROOT}/../outside"}, + "backslash":{"type":"stdio","command":"./scripts\\..\\outside"}, + "remote":{"type":"streamable-http","url":"http://example.com/mcp"}, + "header":{"type":"streamable-http","url":"https://example.com/mcp","headers":{"X-Demo":"one","x-demo":"two"}}, + "sse":{"type":"sse","url":"https://example.com/sse"}, + "unknown":{"type":"stdio","command":"python","future":true} + } + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert_eq!(outcome.servers.keys().collect::>(), vec!["valid"]); + assert_eq!(outcome.errors.len(), 8); +} + +#[cfg(windows)] +#[test] +fn agent_plugin_mcp_rejects_drive_relative_windows_command() { + let plugin_root = plugin_root(); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_root.join("data"), + r#"{"$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json","mcpServers":{"drive-relative":{"type":"stdio","command":"C:server.exe"}}}"#, + ) + .expect("parse Agent Plugins MCP config"); + assert!(outcome.servers.is_empty()); + assert_eq!(outcome.errors.len(), 1); +} + +#[test] +fn agent_plugin_mcp_treats_args_and_env_as_opaque_after_expansion() { + let plugin_root = plugin_root(); + let data_root = plugin_root.join("data"); + let outcome = parse_agent_plugin_mcp_config( + &plugin_root, + &data_root, + r#"{ + "$schema":"https://agent-plugins.org/schemas/1.0.0/mcp.schema.json", + "mcpServers":{"demo":{ + "type":"stdio", + "command":"python", + "args":["${PLUGIN_ROOT}/../opaque"], + "env":{"OPAQUE":"${PLUGIN_DATA}/../opaque"} + }} + }"#, + ) + .expect("parse Agent Plugins MCP config"); + + assert!(outcome.errors.is_empty()); + let McpServerTransportConfig::Stdio { args, env, .. } = &outcome.servers["demo"].transport + else { + panic!("expected stdio transport"); + }; + assert_eq!(args, &vec![format!("{}/../opaque", plugin_root.display())]); + assert_eq!( + env.as_ref().and_then(|env| env.get("OPAQUE")), + Some(&format!("{}/../opaque", data_root.display())) + ); +} + +#[test] +fn agent_plugin_mcp_rejects_unsupported_schema() { + let plugin_root = plugin_root(); + let error = parse_agent_plugin_mcp_config( + &plugin_root, + &plugin_root.join("data"), + r#"{"$schema":"https://agent-plugins.org/schemas/2.0.0/mcp.schema.json","mcpServers":{}}"#, + ) + .expect_err("unsupported schema"); + + assert!( + error + .to_string() + .contains("unsupported Agent Plugins MCP schema") + ); +} + +fn stdio_server( + command: &str, + environment_id: &str, + cwd: LegacyAppPathString, + env_vars: Vec, +) -> McpServerConfig { + McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::Stdio { + command: command.to_string(), + args: Vec::new(), + env: None, + env_vars, + cwd: Some(cwd), + }, + environment_id: environment_id.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + } +} + +#[test] +fn declared_placement_preserves_local_plugin_normalization() { + let plugin_root = plugin_root(); + let expected_stdio = stdio_server( + "demo-mcp", + DEFAULT_MCP_SERVER_ENVIRONMENT_ID, + LegacyAppPathString::from_path(&plugin_root.join("scripts")), + Vec::new(), + ); + let expected_http = McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: "https://example.com/mcp".to_string(), + bearer_token_env_var: None, + http_headers: None, + env_http_headers: None, + http_headers_helper: None, + }, + environment_id: DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: Some(McpServerOAuthConfig { + client_id: Some("client-id".to_string()), + callback_url: Some("http://127.0.0.1/callback/registered".to_string()), + callback_port: Some(9876), + ..Default::default() + }), + oauth_resource: None, + tools: HashMap::new(), + }; + let mut expected_helper = McpServerConfig { + oauth: None, + ..expected_http.clone() + }; + let McpServerTransportConfig::StreamableHttp { + http_headers_helper, + .. + } = &mut expected_helper.transport + else { + unreachable!("expected HTTP transport"); + }; + *http_headers_helper = Some("./auth.sh".to_string()); + + let outcome = parse_plugin_mcp_config( + &plugin_root, + r#"{ + "demo": { + "type": "stdio", + "command": "demo-mcp", + "cwd": "scripts" + }, + "hosted": { + "type": "http", + "url": "https://example.com/mcp", + "oauth": {"clientId": "client-id", "callbackUrl": "http://127.0.0.1/callback/registered", "callbackPort": 9876} + }, + "helper": {"type":"http","url":"https://example.com/mcp","http_headers_helper":"./auth.sh"} + }"#, + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::from([ + ("demo".to_string(), expected_stdio), + ("helper".to_string(), expected_helper), + ("hosted".to_string(), expected_http), + ]), + errors: Vec::new(), + } + ); +} + +#[test] +fn native_plugin_mcp_cannot_self_declare_ema_auth() { + let server = serde_json::json!({ + "type": "http", + "url": "https://resource.example/mcp", + "auth": "ema_auth" + }); + for contents in [ + serde_json::json!({"enterprise": server}), + serde_json::json!({"mcpServers": {"enterprise": server}}), + ] { + let outcome = parse_plugin_mcp_config(&plugin_root(), &contents.to_string()) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::new(), + errors: vec![PluginMcpServerParseError { + name: "enterprise".to_string(), + message: "plugin MCP declarations cannot select ema_auth; configure enterprise authentication in host policy".to_string(), + }], + } + ); + } +} + +#[test] +fn environment_placement_forces_authority_and_defaults_null_cwd() { + let plugin_root = plugin_root(); + let plugin_root_uri = plugin_root_uri(&plugin_root); + let outcome = parse_executor_plugin_mcp_config( + &plugin_root_uri, + r#"{ + "$schema":"https://example.com/plugin-mcp.schema.json", + "mcpServers":{"demo":{ + "command":"demo-mcp", + "environment_id":"local", + "cwd":null, + "env_vars":["EXECUTOR_TOKEN", {"name":"OTHER_TOKEN"}] + }} + }"#, + "executor-1", + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::from([( + "demo".to_string(), + stdio_server( + "demo-mcp", + "executor-1", + plugin_root_uri.into(), + vec![ + McpServerEnvVar::Config { + name: "EXECUTOR_TOKEN".to_string(), + source: Some("remote".to_string()), + }, + McpServerEnvVar::Config { + name: "OTHER_TOKEN".to_string(), + source: Some("remote".to_string()), + }, + ], + ), + )]), + errors: Vec::new(), + } + ); +} + +#[test] +fn environment_placement_resolves_relative_cwd_beneath_plugin_root() { + let plugin_root = plugin_root(); + let plugin_root_uri = plugin_root_uri(&plugin_root); + let outcome = parse_executor_plugin_mcp_config( + &plugin_root_uri, + r#"{"demo":{"command":"demo-mcp","cwd":"scripts"}}"#, + "executor-1", + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::from([( + "demo".to_string(), + stdio_server( + "demo-mcp", + "executor-1", + plugin_root_uri + .join("scripts") + .expect("plugin cwd URI") + .into(), + Vec::new(), + ), + )]), + errors: Vec::new(), + } + ); +} + +#[test] +fn executor_environment_placement_resolves_foreign_uri_cwd() { + let plugin_root = PathUri::parse("file:///C:/plugins/demo").expect("plugin root URI"); + let outcome = parse_executor_plugin_mcp_config( + &plugin_root, + r#"{"demo":{"command":"demo-mcp","cwd":"scripts"}}"#, + "executor-1", + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::from([( + "demo".to_string(), + stdio_server( + "demo-mcp", + "executor-1", + LegacyAppPathString::from( + plugin_root.join("scripts").expect("executor cwd URI"), + ), + Vec::new(), + ), + )]), + errors: Vec::new(), + } + ); +} + +#[test] +fn environment_placement_rejects_relative_cwd_that_escapes_package() { + let plugin_root = plugin_root(); + let plugin_root_uri = plugin_root_uri(&plugin_root); + let outcome = parse_executor_plugin_mcp_config( + &plugin_root_uri, + r#"{"demo":{"command":"demo-mcp","cwd":"../outside"}}"#, + "executor-1", + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::new(), + errors: vec![PluginMcpServerParseError { + name: "demo".to_string(), + message: format!( + "cwd `../outside` must remain within plugin root `{plugin_root_uri}`" + ), + }], + } + ); +} + +#[test] +fn environment_placement_rejects_orchestrator_env_vars() { + let plugin_root = plugin_root(); + let outcome = parse_executor_plugin_mcp_config( + &plugin_root_uri(&plugin_root), + r#"{"demo":{"command":"demo-mcp","env_vars":[{"name":"TOKEN","source":"local"}]}}"#, + "executor-1", + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::new(), + errors: vec![PluginMcpServerParseError { + name: "demo".to_string(), + message: + "env_vars entry `TOKEN` cannot use source `local` in an executor-owned plugin" + .to_string(), + }], + } + ); +} + +#[test] +fn remote_environment_placement_preserves_bearer_and_rejects_header_env_references() { + let plugin_root = plugin_root(); + let outcome = parse_executor_plugin_mcp_config( + &plugin_root_uri(&plugin_root), + r#"{ + "bearer": { + "url": "https://example.com/bearer", + "bearer_token_env_var": "TOKEN" + }, + "headers": { + "url": "https://example.com/headers", + "env_http_headers": {"Authorization": "TOKEN"} + } + }"#, + "executor-1", + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::from([( + "bearer".to_string(), + serde_json::from_value(serde_json::json!({ + "url": "https://example.com/bearer", + "bearer_token_env_var": "TOKEN", + "environment_id": "executor-1", + })) + .expect("executor-owned bearer configuration"), + )]), + errors: vec![PluginMcpServerParseError { + name: "headers".to_string(), + message: "`env_http_headers` requires executor-side environment resolution for an executor-owned HTTP MCP" + .to_string(), + }], + } + ); +} + +#[test] +fn local_environment_placement_preserves_http_env_references() { + let plugin_root = plugin_root(); + let outcome = parse_executor_plugin_mcp_config( + &plugin_root_uri(&plugin_root), + r#"{ + "demo": { + "url": "https://example.com/mcp", + "bearer_token_env_var": "TOKEN", + "env_http_headers": {"X-Account": "ACCOUNT_ID"} + } + }"#, + DEFAULT_MCP_SERVER_ENVIRONMENT_ID, + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::from([( + "demo".to_string(), + McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: "https://example.com/mcp".to_string(), + bearer_token_env_var: Some("TOKEN".to_string()), + http_headers: None, + env_http_headers: Some(HashMap::from([( + "X-Account".to_string(), + "ACCOUNT_ID".to_string(), + )])), + http_headers_helper: None, + }, + environment_id: DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + }, + )]), + errors: Vec::new(), + } + ); +} + +#[test] +fn local_environment_placement_preserves_local_env_vars() { + let plugin_root = plugin_root(); + let plugin_root_uri = plugin_root_uri(&plugin_root); + let outcome = parse_executor_plugin_mcp_config( + &plugin_root_uri, + r#"{"demo":{"command":"demo-mcp","env_vars":["TOKEN",{"name":"OTHER","source":"local"}]}}"#, + DEFAULT_MCP_SERVER_ENVIRONMENT_ID, + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::from([( + "demo".to_string(), + stdio_server( + "demo-mcp", + DEFAULT_MCP_SERVER_ENVIRONMENT_ID, + plugin_root_uri.into(), + vec![ + McpServerEnvVar::Name("TOKEN".to_string()), + McpServerEnvVar::Config { + name: "OTHER".to_string(), + source: Some("local".to_string()), + }, + ], + ), + )]), + errors: Vec::new(), + } + ); +} + +#[test] +fn local_environment_placement_rejects_remote_env_vars() { + let plugin_root = plugin_root(); + let outcome = parse_executor_plugin_mcp_config( + &plugin_root_uri(&plugin_root), + r#"{"demo":{"command":"demo-mcp","env_vars":[{"name":"TOKEN","source":"remote"}]}}"#, + DEFAULT_MCP_SERVER_ENVIRONMENT_ID, + ) + .expect("parse plugin MCP config"); + + assert_eq!( + outcome, + PluginMcpConfigParseOutcome { + servers: BTreeMap::new(), + errors: vec![PluginMcpServerParseError { + name: "demo".to_string(), + message: "env_vars entry `TOKEN` cannot use source `remote` in a local environment" + .to_string(), + }], + } + ); +} diff --git a/codex-rs/codex-mcp/src/resource_client.rs b/codex-rs/codex-mcp/src/resource_client.rs new file mode 100644 index 0000000000000000000000000000000000000000..3c13695c73eb24129199002ccb64f735770ab8e5 --- /dev/null +++ b/codex-rs/codex-mcp/src/resource_client.rs @@ -0,0 +1,373 @@ +use std::sync::Arc; +use std::sync::Weak; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use codex_protocol::mcp::Resource; +use codex_protocol::mcp::ResourceContent; +use codex_rmcp_client::CancellableEventStreamRequest; +use codex_rmcp_client::RmcpClient; +use rmcp::model::GetMeta; +use rmcp::model::PaginatedRequestParams; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ServerResult; +use rmcp::service::ServiceError; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Map; +use serde_json::Value; +use serde_json::json; +use tokio::runtime::Handle; +use tokio::sync::watch; + +use crate::McpEventStreamOpener; +use crate::McpRuntime; +use crate::connection_manager::McpConnectionSet; +use crate::mcp::CODEX_APPS_MCP_SERVER_NAME; + +/// One page of resources returned by an MCP server. +#[derive(Clone, Debug, PartialEq)] +pub struct McpResourcePage { + /// Resources advertised on this page. + pub resources: Vec, + /// Opaque cursor to supply when requesting the next page. + pub next_cursor: Option, +} + +/// Parameters for one Codex Apps resource page. +/// +/// Keep `mime_type` when requesting a continuation page: the server applies +/// the filter to each request separately. +#[derive(Clone, Debug, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct CodexAppsResourceListParams { + #[serde(skip_serializing_if = "Option::is_none")] + pub cursor: Option, + pub mime_type: String, +} + +/// Contents returned after reading one MCP resource. +#[derive(Clone, Debug, PartialEq)] +pub struct McpResourceReadResult { + /// Text or blob content returned for the requested resource. + pub contents: Vec, +} + +/// An event advertised by an MCP server. +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(rename_all = "camelCase")] +pub struct McpEventDefinition { + pub name: String, + pub description: String, + pub delivery: Vec, + pub input_schema: Value, + pub payload_schema: Value, +} + +/// Events returned from one stable MCP connection generation. +pub struct McpEventCatalogSnapshot { + pub cache_key: McpResourceClientCacheKey, + pub events: Vec, +} + +/// One unmodified lifecycle notification from an MCP event subscription. +#[derive(Clone, Debug, PartialEq)] +pub struct McpEventNotification { + pub method: String, + pub params: Option, +} + +/// Owns an MCP event subscription and cancels its request when dropped. +pub struct McpEventStream { + request: Option, + runtime_handle: Handle, + client: Option>, + cancel_event_streams_on_server_removal: watch::Receiver<()>, +} + +impl McpEventStream { + pub(crate) async fn open( + client: Arc, + cancel_event_streams_on_server_removal: watch::Receiver<()>, + event_name: &str, + arguments: &Value, + request_meta: Option<&Map>, + ) -> Result { + let mut params = json!({ "name": event_name, "arguments": arguments }); + if let Some(request_meta) = request_meta { + params["_meta"] = Value::Object(request_meta.clone()); + } + let request = client + .send_event_stream_request(Some(params)) + .await + .context("events/stream request failed")?; + Ok(Self { + request: Some(request), + runtime_handle: Handle::current(), + client: Some(client), + cancel_event_streams_on_server_removal, + }) + } + + /// Receives the next raw lifecycle notification for this subscription. + pub async fn recv(&mut self) -> Result> { + let Some(request) = self.request.as_mut() else { + return Ok(None); + }; + + tokio::select! { + biased; + + Ok(()) = self.cancel_event_streams_on_server_removal.changed() => { + self.cancel(); + Err(anyhow!("hosted MCP event server was removed")) + } + Some(notification) = request.notifications.recv() => { + let metadata = notification.get_meta().0.0.clone(); + let mut params = notification.params; + if !metadata.is_empty() { + params.get_or_insert_with(|| json!({}))["_meta"] = Value::Object(metadata); + } + Ok(Some(McpEventNotification { + method: notification.method, + params, + })) + } + response = &mut request.handle.rx => { + self.request = None; + self.client = None; + + match response { + Ok(Ok(_)) + | Ok(Err(ServiceError::Cancelled { .. })) + | Ok(Err(ServiceError::TransportClosed)) + | Err(_) => Ok(None), + Ok(Err(error)) => Err(error.into()), + } + } + } + } + + fn cancel(&mut self) { + if let Some(CancellableEventStreamRequest { + handle, + notifications, + }) = self.request.take() + { + drop(notifications); + let client = self.client.take(); + self.runtime_handle.spawn(async move { + let _ = tokio::time::timeout( + Duration::from_secs(30), + handle.cancel(Some("event subscription closed".to_string())), + ) + .await; + drop(client); + }); + } + } +} + +impl Drop for McpEventStream { + fn drop(&mut self) { + self.cancel(); + } +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase")] +struct McpEventListResult { + events: Vec, +} + +/// Access to MCP resources and event subscriptions through the latest runtime. +#[derive(Clone)] +pub struct McpResourceClient { + runtime: Arc, +} + +/// Opaque identity for the connection set currently used by an MCP resource client. +#[derive(Clone)] +pub struct McpResourceClientCacheKey(Weak); + +impl PartialEq for McpResourceClientCacheKey { + fn eq(&self, other: &Self) -> bool { + self.0.ptr_eq(&other.0) + } +} + +impl Eq for McpResourceClientCacheKey {} + +impl std::fmt::Debug for McpResourceClient { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("McpResourceClient") + .finish_non_exhaustive() + } +} + +impl McpResourceClient { + /// Creates a resource client that follows the thread's latest published runtime. + pub fn new(runtime: Arc) -> Self { + Self { runtime } + } + + /// Returns the identity of the connection set used by this client. + pub fn cache_key(&self) -> McpResourceClientCacheKey { + McpResourceClientCacheKey(Arc::downgrade(&self.runtime.latest_connections())) + } + + /// Returns whether this client can address the named server. + /// + /// This does not wait for server startup. + pub async fn has_server(&self, server: &str) -> bool { + self.runtime.latest_connections().contains_server(server) + } + + /// Lists one resource page from the named server. + pub async fn list_resources( + &self, + server: &str, + cursor: Option, + ) -> Result { + let params = + cursor.map(|cursor| PaginatedRequestParams::default().with_cursor(Some(cursor))); + let result = self + .runtime + .latest_connections() + .list_resources(server, params) + .await?; + let resources = result + .resources + .into_iter() + .map(resource_from_rmcp) + .collect::>>()?; + Ok(McpResourcePage { + resources, + next_cursor: result.next_cursor, + }) + } + + /// Lists one Codex Apps resource page using plugin-service's top-level `mimeType` parameter. + pub async fn list_codex_apps_resources( + &self, + params: CodexAppsResourceListParams, + ) -> Result { + let params = serde_json::to_value(params) + .context("failed to serialize Codex Apps resource params")?; + let connections = self.runtime.latest_host_owned_codex_apps_connections()?; + let (managed, timeout) = connections + .client_by_name(CODEX_APPS_MCP_SERVER_NAME) + .await?; + let result = managed + .client + .send_custom_request_with_timeout("resources/list", Some(params), timeout) + .await + .context("resources/list failed for `codex_apps`")?; + let result = match result { + ServerResult::ListResourcesResult(result) => result, + ServerResult::CustomResult(result) => result + .result_as::() + .context("resources/list returned invalid resources")?, + _ => return Err(anyhow!("resources/list returned an unexpected MCP result")), + }; + let resources = result + .resources + .into_iter() + .map(resource_from_rmcp) + .collect::>>()?; + Ok(McpResourcePage { + resources, + next_cursor: result.next_cursor, + }) + } + + /// Reads one resource from the named server. + pub async fn read_resource(&self, server: &str, uri: &str) -> Result { + let params = ReadResourceRequestParams::new(uri.to_string()); + let result = self + .runtime + .latest_connections() + .read_resource(server, params) + .await?; + let contents = result + .contents + .into_iter() + .map(resource_content_from_rmcp) + .collect::>>()?; + Ok(McpResourceReadResult { contents }) + } + + /// Lists the events advertised by the MCP event server. + pub async fn list_events(&self) -> Result { + let (connections, _) = self + .runtime + .latest_connections_for_event_server(CODEX_APPS_MCP_SERVER_NAME)?; + let cache_key = McpResourceClientCacheKey(Arc::downgrade(&connections)); + let (managed, request_timeout) = connections + .client_by_name(CODEX_APPS_MCP_SERVER_NAME) + .await?; + let result = managed + .client + .send_custom_request_with_timeout("events/list", /*params*/ None, request_timeout) + .await + .context("events/list request failed")?; + let ServerResult::CustomResult(result) = result else { + return Err(anyhow!("events/list returned an unexpected MCP result")); + }; + let result = result + .result_as::() + .context("events/list returned invalid event definitions")?; + + Ok(McpEventCatalogSnapshot { + cache_key, + events: result.events, + }) + } + + /// Opens an MCP event subscription with the supplied event arguments. + pub async fn open_event_stream( + &self, + event_name: &str, + arguments: &Value, + request_meta: Option<&Map>, + ) -> Result { + let (connections, cancel_event_streams_on_server_removal) = self + .runtime + .latest_connections_for_event_server(CODEX_APPS_MCP_SERVER_NAME)?; + let (managed, _) = connections + .client_by_name(CODEX_APPS_MCP_SERVER_NAME) + .await?; + McpEventStream::open( + managed.client, + cancel_event_streams_on_server_removal, + event_name, + arguments, + request_meta, + ) + .await + } + + /// Creates an event stream opener using the task's event server settings. + pub fn event_stream_opener(&self) -> Result { + self.runtime.event_stream_opener() + } + + /// Forwards event server removal to the owner of the task's subscriptions. + pub fn forward_event_server_removals_to(&self, cancellation: watch::Sender<()>) { + self.runtime.forward_event_server_removals_to(cancellation); + } +} + +fn resource_from_rmcp(resource: rmcp::model::Resource) -> Result { + let value = serde_json::to_value(resource).context("failed to serialize MCP resource")?; + Resource::from_mcp_value(value).context("failed to convert MCP resource") +} + +fn resource_content_from_rmcp(content: rmcp::model::ResourceContents) -> Result { + let value = + serde_json::to_value(content).context("failed to serialize MCP resource content")?; + serde_json::from_value(value).context("failed to convert MCP resource content") +} diff --git a/codex-rs/codex-mcp/src/resource_origin.rs b/codex-rs/codex-mcp/src/resource_origin.rs new file mode 100644 index 0000000000000000000000000000000000000000..0337bf761d39213b0b702e0706c0a7278c3a9ce1 --- /dev/null +++ b/codex-rs/codex-mcp/src/resource_origin.rs @@ -0,0 +1,316 @@ +//! Thread-owned, bounded provenance for app-hosted widget resources. + +use std::collections::VecDeque; + +use anyhow::Context; +use codex_connectors::AppToolPolicyEvaluator; +use codex_connectors::AppToolPolicyInput; +use codex_protocol::ThreadId; +use codex_protocol::items::McpToolCallStatus; +use codex_protocol::items::TurnItem; +use codex_protocol::mcp::McpResourceOrigin; +use codex_protocol::mcp::McpResourceOriginCheckpoint; +use codex_protocol::protocol::EventMsg; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ReadResourceResult; + +use crate::CODEX_APPS_MCP_SERVER_NAME; +use crate::McpBinding; + +const MAX_ORIGINS: usize = 64; +const MAX_ORIGIN_BYTES: usize = 1024; + +#[derive(Default)] +pub(crate) struct ResourceOrigins { + origins: VecDeque, + turns: VecDeque, + current_turn_id: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) struct ResourceOrigin { + call_id: String, + turn_id: Option, + tool: String, + connector_id: String, + link_id: Option, + uri: String, + ambiguous_account: bool, +} + +impl ResourceOrigins { + pub(crate) fn checkpoint(&self) -> Option { + (!self.origins.is_empty()).then(|| McpResourceOriginCheckpoint { + origins: self + .origins + .iter() + .map(|origin| McpResourceOrigin { + call_id: origin.call_id.clone(), + turn_id: origin.turn_id.clone(), + tool: origin.tool.clone(), + connector_id: origin.connector_id.clone(), + link_id: origin.link_id.clone(), + uri: origin.uri.clone(), + ambiguous_account: origin.ambiguous_account, + }) + .collect(), + turns: self.turns.iter().cloned().collect(), + current_turn_id: self.current_turn_id.clone(), + }) + } + + pub(crate) fn restore_checkpoint(&mut self, checkpoint: &McpResourceOriginCheckpoint) { + if checkpoint.origins.len() > MAX_ORIGINS + || checkpoint.turns.len() > MAX_ORIGINS + || checkpoint + .turns + .iter() + .any(|turn_id| turn_id.len() > MAX_ORIGIN_BYTES) + || checkpoint + .current_turn_id + .as_ref() + .is_some_and(|turn_id| turn_id.len() > MAX_ORIGIN_BYTES) + { + *self = Self::default(); + return; + } + + let origins = checkpoint + .origins + .iter() + .map(|origin| ResourceOrigin { + call_id: origin.call_id.clone(), + turn_id: origin.turn_id.clone(), + tool: origin.tool.clone(), + connector_id: origin.connector_id.clone(), + link_id: origin.link_id.clone(), + uri: origin.uri.clone(), + ambiguous_account: origin.ambiguous_account, + }) + .collect::>(); + if origins.iter().any(|origin| { + origin.byte_len() > MAX_ORIGIN_BYTES || origin.connector_id.trim().is_empty() + }) { + *self = Self::default(); + return; + } + + *self = Self { + origins, + turns: checkpoint.turns.iter().cloned().collect(), + current_turn_id: checkpoint.current_turn_id.clone(), + }; + } + + pub(crate) fn observe(&mut self, event: &EventMsg) { + match event { + EventMsg::TurnStarted(event) if event.turn_id.len() <= MAX_ORIGIN_BYTES => { + self.current_turn_id = Some(event.turn_id.clone()); + if self.turns.back() != Some(&event.turn_id) { + self.turns.push_back(event.turn_id.clone()); + if self.turns.len() > MAX_ORIGINS { + self.turns.pop_front(); + } + } + } + EventMsg::ItemCompleted(event) => { + let TurnItem::McpToolCall(item) = &event.item else { + return; + }; + if item.status == McpToolCallStatus::Completed { + self.remember( + &item.id, + Some(&event.turn_id), + &item.server, + &item.tool, + &item.arguments, + item.connector_id.as_deref(), + item.link_id.as_deref(), + item.mcp_app_resource_uri.as_deref(), + ); + } + } + EventMsg::McpToolCallEnd(event) + if event + .result + .as_ref() + .is_ok_and(|result| !result.is_error.unwrap_or(false)) => + { + let turn_id = self.current_turn_id.clone(); + self.remember( + &event.call_id, + turn_id.as_deref(), + &event.invocation.server, + &event.invocation.tool, + event + .invocation + .arguments + .as_ref() + .unwrap_or(&serde_json::Value::Null), + event.connector_id.as_deref(), + event.link_id.as_deref(), + event.mcp_app_resource_uri.as_deref(), + ); + } + EventMsg::ThreadRolledBack(event) => { + for _ in 0..event.num_turns { + let Some(turn_id) = self.turns.pop_back() else { + *self = Self::default(); + return; + }; + self.origins.retain(|origin| { + origin.turn_id.is_some() + && origin.turn_id.as_deref() != Some(turn_id.as_str()) + }); + } + self.current_turn_id = self.turns.back().cloned(); + } + _ => {} + } + } + + pub(crate) fn find(&self, call_id: &str) -> anyhow::Result { + self.origins + .iter() + .rev() + .find(|origin| origin.call_id == call_id) + .cloned() + .context("originating MCP tool call was not found or did not complete successfully") + } + + #[expect( + clippy::too_many_arguments, + reason = "only bounded provenance fields are retained" + )] + fn remember( + &mut self, + call_id: &str, + turn_id: Option<&str>, + server: &str, + tool: &str, + arguments: &serde_json::Value, + connector_id: Option<&str>, + link_id: Option<&str>, + uri: Option<&str>, + ) { + if server != CODEX_APPS_MCP_SERVER_NAME { + return; + } + let Some(connector_id) = connector_id.filter(|value| !value.trim().is_empty()) else { + return; + }; + let Some(uri) = uri else { + return; + }; + let link_id = link_id.filter(|value| !value.trim().is_empty()); + let ambiguous_account = arguments + .get("link_id") + .and_then(serde_json::Value::as_str) + .filter(|value| !value.trim().is_empty()) + .is_some_and(|argument_link_id| Some(argument_link_id) != link_id); + let origin = ResourceOrigin { + call_id: call_id.to_owned(), + turn_id: turn_id.map(str::to_owned), + tool: tool.to_owned(), + connector_id: connector_id.to_owned(), + link_id: link_id.map(str::to_owned), + uri: uri.to_owned(), + ambiguous_account, + }; + if origin.byte_len() > MAX_ORIGIN_BYTES { + return; + } + + if let Some(index) = self + .origins + .iter() + .position(|existing| existing.call_id == origin.call_id) + { + self.origins.remove(index); + } + if self.origins.len() >= MAX_ORIGINS { + self.origins.pop_front(); + } + self.origins.push_back(origin); + } +} + +impl ResourceOrigin { + fn byte_len(&self) -> usize { + self.call_id.len() + + self.turn_id.as_ref().map_or(0, String::len) + + self.tool.len() + + self.connector_id.len() + + self.link_id.as_ref().map_or(0, String::len) + + self.uri.len() + } + + pub(crate) async fn read( + &self, + binding: &McpBinding, + thread_id: ThreadId, + uri: &str, + ) -> anyhow::Result { + if self.uri != uri { + anyhow::bail!("originating MCP tool call does not match the requested resource"); + } + if self.ambiguous_account { + anyhow::bail!("originating MCP tool call has ambiguous account selection"); + } + + let tool_info = binding + .tool_info(CODEX_APPS_MCP_SERVER_NAME, &self.tool) + .context("originating MCP tool is unavailable")?; + if tool_info.connector_id.as_deref() != Some(self.connector_id.as_str()) { + anyhow::bail!("originating MCP tool connector does not match its app context"); + } + let tool_meta = tool_info.tool.meta.as_ref().map(|meta| &meta.0); + let current_link_id = tool_meta + .and_then(|meta| meta.get("link_id")) + .and_then(serde_json::Value::as_str) + .filter(|value| !value.trim().is_empty()); + if current_link_id != self.link_id.as_deref() { + anyhow::bail!("originating MCP tool link does not match its app context"); + } + if self.link_id.is_none() + && tool_meta + .and_then(|meta| meta.get("_codex_apps")) + .and_then(|meta| meta.get("requires_explicit_link_id")) + .and_then(serde_json::Value::as_bool) + == Some(true) + { + anyhow::bail!("originating MCP tool requires an explicit account link"); + } + + let annotations = tool_info.tool.annotations.as_ref(); + if !AppToolPolicyEvaluator::new(&binding.config().config_layer_stack) + .policy(AppToolPolicyInput { + connector_id: Some(&self.connector_id), + link_id: None, + tool_name: tool_info.tool.name.as_ref(), + tool_title: tool_info.tool.title.as_deref(), + destructive_hint: annotations.and_then(|value| value.destructive_hint), + open_world_hint: annotations.and_then(|value| value.open_world_hint), + }) + .enabled + { + anyhow::bail!("originating MCP tool is disabled by app configuration"); + } + + let meta = serde_json::from_value(serde_json::json!({ + "threadId": thread_id, + "x-codex-turn-metadata": { + "mcp_request_meta": { + "selected_connector_ids": [&self.connector_id], + "link_id": &self.link_id, + } + }, + }))?; + binding + .read_resource( + CODEX_APPS_MCP_SERVER_NAME, + ReadResourceRequestParams::new(uri).with_meta(meta), + ) + .await + } +} diff --git a/codex-rs/codex-mcp/src/rmcp_client.rs b/codex-rs/codex-mcp/src/rmcp_client.rs new file mode 100644 index 0000000000000000000000000000000000000000..06ea54e861939e7c4dd7fdb9ad3f6a3112a19eab --- /dev/null +++ b/codex-rs/codex-mcp/src/rmcp_client.rs @@ -0,0 +1,1475 @@ +//! RMCP client lifecycle for MCP server connections. +//! +//! This module owns startup of individual RMCP clients: building the transport, +//! initializing the server, listing raw tools, applying per-server tool filters, +//! and exposing cached Codex Apps tools while a client is still connecting. +//! Initialization capabilities survive tool-discovery failures and reset on a new attempt. +//! Higher-level aggregation and resource/tool APIs live in +//! [`crate::connection_manager`]. + +#[path = "rmcp_client/status.rs"] +mod status; + +use std::borrow::Cow; +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::env; +use std::ffi::OsString; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +use std::time::Duration; +use std::time::Instant; + +use crate::client_tool_catalog::ClientToolCatalog; +use crate::codex_apps::normalize_codex_apps_callable_name; +use crate::codex_apps::normalize_codex_apps_callable_namespace; +use crate::codex_apps::normalize_codex_apps_tool_title; +use crate::codex_apps::prepare_openai_file_params_for_model; +use crate::elicitation::ElicitationRequestManager; +use crate::executor_environment_http_client::ExecutorEnvironmentHttpClient; +use crate::mcp::CODEX_APPS_MCP_SERVER_NAME; +use crate::mcp::ToolPluginContext; +use crate::openai_docs_source_attribution::maybe_with_openai_docs_source_attribution; +use crate::pagination::collect_paginated_with_limit; +use crate::runtime::McpRuntimeContext; +use crate::runtime::emit_duration; +use crate::server::EffectiveMcpServer; +use crate::server::has_explicit_http_authorization; +use crate::tool_catalog_cache::McpToolCatalogCacheContext; +use crate::tool_catalog_cache::McpToolCatalogFetchTicket; +use crate::tools::ToolInfo; +use anyhow::Result; +use anyhow::anyhow; +use async_channel::Sender; +use codex_api::SharedAuthProvider; +use codex_async_utils::CancelErr; +use codex_async_utils::OrCancelExt; +use codex_config::McpServerAuth; +use codex_config::McpServerConfig; +use codex_config::McpServerTransportConfig; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_connectors::ConnectorRuntimeContext; +use codex_connectors::ConnectorRuntimeFetchSource; +use codex_exec_server::Environment; +use codex_login::AuthChangeState; +use codex_protocol::mcp::ClientMcpExtensions; +use codex_protocol::mcp::McpServerInfo; +use codex_protocol::protocol::Event; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::McpStartupStatus; +use codex_protocol::protocol::McpStartupUpdateEvent; +use codex_rmcp_client::ExecutorStdioServerLauncher; +use codex_rmcp_client::LocalStdioServerLauncher; +use codex_rmcp_client::McpOAuthRefreshMode; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::StdioServerLauncher; +use codex_rmcp_client::StreamableHttpBearerToken; +use codex_rmcp_client::StreamableHttpRedirectMode; +use codex_rmcp_client::ToolWithConnectorId; +use codex_rmcp_client::is_authentication_required_error; +use futures::future::BoxFuture; +use futures::future::FutureExt; +use futures::future::Shared; +use rmcp::model::ClientCapabilities; +use rmcp::model::ElicitationCapability; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use rmcp::model::ServerPeerInfo; +use rmcp::model::Tool as RmcpTool; +use tokio::sync::watch; +use tokio::time::Instant as TokioInstant; +use tokio_util::sync::CancellationToken; +use tokio_util::task::AbortOnDropHandle; +use tracing::Instrument; +use tracing::instrument; +use tracing::warn; + +/// MCP server capability indicating that Codex should include [`SandboxState`] +/// in tool-call request `_meta` under this key. +pub const MCP_SANDBOX_STATE_META_CAPABILITY: &str = "codex/sandbox-state-meta"; +/// Experimental MCP server capability for development and testing only; production servers should +/// not use it. Its `cacheable: false` property disables sharing tool definitions across connections. +const MCP_TOOL_CATALOG_CACHE_CAPABILITY: &str = "codex/tool-catalog-cache"; +const MCP_TOOL_CATALOG_CACHEABLE_PROPERTY: &str = "cacheable"; +pub(crate) const MCP_TOOLS_LIST_DURATION_METRIC: &str = "codex.mcp.tools.list.duration_ms"; +pub(crate) const MCP_TOOLS_FETCH_UNCACHED_DURATION_METRIC: &str = + "codex.mcp.tools.fetch_uncached.duration_ms"; +pub(crate) const CODEX_APPS_REFRESH_DURATION_METRIC: &str = "codex.apps.refresh.duration_ms"; +pub(crate) const DEFAULT_STARTUP_TIMEOUT: Duration = Duration::from_secs(30); +pub(crate) const DEFAULT_TOOL_TIMEOUT: Duration = Duration::from_secs(300); + +pub(crate) const CODEX_APPS_RECONNECT_INITIAL_BACKOFF: Duration = Duration::from_secs(1); +const CODEX_APPS_RECONNECT_MAX_BACKOFF: Duration = Duration::from_secs(30); + +const UNTRUSTED_CONNECTOR_META_KEYS: &[&str] = &[ + "connector_id", + "connector_name", + "connector_display_name", + "connector_description", + "connectorDescription", +]; + +#[derive(Clone)] +pub(crate) struct ManagedClient { + pub(crate) _auth_change_notifications: Option>>, + pub(crate) client: Arc, + pub(crate) server_info: McpServerInfo, + pub(crate) tool_catalog: Arc, + pub(crate) tool_timeout: Option, + pub(crate) server_instructions: Option, + pub(crate) server_supports_sandbox_state_meta_capability: bool, + pub(crate) codex_apps_tools_cache_context: Option>, +} + +impl ManagedClient { + pub(crate) async fn listed_tools(&self) -> Vec { + let total_start = Instant::now(); + self.tool_catalog + .read(|catalog| { + // Discovery may use the shared cache until this client is refreshed. + // Executable bindings always capture this client's own catalog. + if catalog.revision == 0 + && let Some(cache_context) = &self.codex_apps_tools_cache_context + { + let tools = cache_context.current_tools(); + emit_duration( + MCP_TOOLS_LIST_DURATION_METRIC, + total_start.elapsed(), + &[("cache", if tools.is_some() { "hit" } else { "miss" })], + ); + if let Some(tools) = tools { + return tools; + } + } + catalog.tools.to_vec() + }) + .await + } +} + +pub(crate) type ManagedClientFuture = + Shared>>; + +#[derive(Default)] +struct CodexAppsStartupReconnectState { + current_client: Option, + last_error: Option, + reconnect_in_flight: bool, + consecutive_failures: u32, + retry_not_before: Option, +} + +#[derive(Clone)] +struct CodexAppsStartupStatusContext { + submit_id: String, + server_name: String, + tx_event: Sender, +} + +pub(crate) struct CodexAppsStartupReconnect { + factory: Arc ManagedClientFuture + Send + Sync>, + state: StdMutex, + startup_status_context: Option, +} + +impl CodexAppsStartupReconnect { + pub(crate) fn new(factory: Arc ManagedClientFuture + Send + Sync>) -> Self { + Self { + factory, + state: StdMutex::new(CodexAppsStartupReconnectState::default()), + startup_status_context: None, + } + } + + fn with_startup_status_context( + mut self, + submit_id: String, + server_name: String, + tx_event: Option>, + ) -> Self { + self.startup_status_context = tx_event.map(|tx_event| CodexAppsStartupStatusContext { + submit_id, + server_name, + tx_event, + }); + self + } + + fn current_client(&self) -> Option { + self.state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .current_client + .clone() + } + + fn reconnect_in_background(self: &Arc) { + { + let mut state = self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if state.current_client.is_some() || state.reconnect_in_flight { + return; + } + if state + .retry_not_before + .is_some_and(|retry_not_before| TokioInstant::now() < retry_not_before) + { + return; + } + state.reconnect_in_flight = true; + } + + let reconnect = Arc::clone(self); + tokio::spawn(async move { + let result = (reconnect.factory)().await; + let startup_status_context = reconnect.startup_status_context.clone(); + let recovered = { + let mut state = reconnect + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + state.reconnect_in_flight = false; + match result { + Ok(client) => { + state.current_client = Some(client); + state.last_error = None; + state.consecutive_failures = 0; + state.retry_not_before = None; + true + } + Err(error) => { + state.last_error = Some(error.clone()); + state.consecutive_failures = state.consecutive_failures.saturating_add(1); + let retry_after = codex_apps_reconnect_backoff(state.consecutive_failures); + state.retry_not_before = Some(TokioInstant::now() + retry_after); + warn!( + error = %error, + retry_after_ms = retry_after.as_millis(), + "Apps MCP startup reconnect failed; continuing with cached tools" + ); + false + } + } + }; + + if recovered && let Some(context) = startup_status_context { + let _ = context + .tx_event + .send(Event { + id: context.submit_id, + msg: EventMsg::McpStartupUpdate(McpStartupUpdateEvent { + server: context.server_name, + status: McpStartupStatus::Ready, + }), + }) + .await; + } + }); + } +} + +fn codex_apps_reconnect_backoff(consecutive_failures: u32) -> Duration { + let exponent = consecutive_failures.saturating_sub(1).min(5); + CODEX_APPS_RECONNECT_INITIAL_BACKOFF + .saturating_mul(1 << exponent) + .min(CODEX_APPS_RECONNECT_MAX_BACKOFF) +} + +#[derive(Clone)] +struct ManagedClientStartup { + server_name: String, + server: EffectiveMcpServer, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + oauth_refresh_mode: McpOAuthRefreshMode, + tx_event: Option>, + elicitation_requests: ElicitationRequestManager, + codex_apps_tools_cache_context: Option>, + tool_catalog_cache_context: Option, + runtime_context: McpRuntimeContext, + resolved_environment: std::result::Result>, String>, + runtime_auth_provider: Option, + client_elicitation_capability: ElicitationCapability, + client_mcp_extensions: ClientMcpExtensions, + auth_changes: Option>, + protocol_mode: McpProtocolMode, + catalog_item_limit: usize, + cancel_token: CancellationToken, + startup_complete: Arc, + server_capabilities: Arc>>, +} + +impl ManagedClientStartup { + fn start(&self) -> ManagedClientFuture { + // A new attempt must not expose capabilities from an earlier connection. + *self + .server_capabilities + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = None; + let Self { + server_name, + server, + store_mode, + keyring_backend_kind, + oauth_refresh_mode, + tx_event, + elicitation_requests, + codex_apps_tools_cache_context, + tool_catalog_cache_context, + runtime_context, + resolved_environment, + runtime_auth_provider, + client_elicitation_capability, + client_mcp_extensions, + auth_changes, + protocol_mode, + catalog_item_limit, + cancel_token, + startup_complete, + server_capabilities, + } = self.clone(); + let is_codex_apps_mcp_server = server_name == CODEX_APPS_MCP_SERVER_NAME; + let startup_timeout = server + .config() + .startup_timeout_sec + .unwrap_or(DEFAULT_STARTUP_TIMEOUT); + let cancel_token_for_fut = cancel_token; + async move { + let tool_catalog_fetch_ticket = tool_catalog_cache_context + .as_ref() + .map(McpToolCatalogCacheContext::begin_fetch); + let refresh_start = is_codex_apps_mcp_server.then(Instant::now); + let outcome = match async { + if let Err(error) = validate_mcp_server_name(&server_name) { + return Err(error.into()); + } + + let client = match tokio::time::timeout( + startup_timeout, + make_rmcp_client( + &server_name, + server.clone(), + store_mode, + keyring_backend_kind, + oauth_refresh_mode, + runtime_context, + resolved_environment, + runtime_auth_provider, + protocol_mode, + ), + ) + .await + { + Ok(result) => Arc::new(result?), + Err(_) => { + return Err(StartupOutcomeError::from(anyhow!( + "MCP client startup timed out after {startup_timeout:?}" + ))); + } + }; + start_server_task( + server_name.clone(), + client, + StartServerTaskParams { + is_codex_apps_mcp_server, + startup_timeout: Some(startup_timeout), + tx_event, + elicitation_requests, + codex_apps_tools_cache_context, + tool_catalog_cache_context, + tool_catalog_fetch_ticket, + client_elicitation_capability, + client_mcp_extensions, + auth_changes, + catalog_item_limit, + server_capabilities, + }, + ) + .await + } + .or_cancel(&cancel_token_for_fut) + .await + { + Ok(result) => result, + Err(CancelErr::Cancelled) => Err(StartupOutcomeError::Cancelled), + }; + // Log once per startup attempt, including discovery without startup notifications. + if let Err(StartupOutcomeError::Failed { error, .. }) = &outcome { + warn!(server_name, %error, "MCP server startup failed"); + } + if outcome.is_ok() + && let Some(refresh_start) = refresh_start + { + emit_duration( + CODEX_APPS_REFRESH_DURATION_METRIC, + refresh_start.elapsed(), + &[("path", "legacy"), ("trigger", "initial")], + ); + } + + startup_complete.store(true, Ordering::Release); + outcome + } + .in_current_span() + .boxed() + .shared() + } +} + +#[derive(Clone)] +pub(crate) struct AsyncManagedClient { + pub(crate) client: ManagedClientFuture, + pub(crate) is_codex_apps_mcp_server: bool, + pub(crate) cached_server_info: Option, + /// Retained after initialization even if subsequent tool discovery fails. + pub(crate) server_capabilities: Arc>>, + pub(crate) codex_apps_tools_cache_context: Option>, + pub(crate) tool_catalog_cache_context: Option, + pub(crate) startup_complete: Arc, + pub(crate) startup_reconnect: Option>, + pub(crate) cancel_token: CancellationToken, +} + +impl AsyncManagedClient { + // Keep this constructor flat so the startup inputs remain readable at the + // single call site instead of introducing a one-off params wrapper. + #[instrument(level = "trace", skip_all, fields(server_name = %server_name))] + #[allow(clippy::too_many_arguments)] + pub(crate) fn new( + server_name: String, + startup_submit_id: String, + server: EffectiveMcpServer, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + oauth_refresh_mode: McpOAuthRefreshMode, + cancel_token: CancellationToken, + tx_event: Option>, + elicitation_requests: ElicitationRequestManager, + codex_apps_tools_cache_context: Option>, + tool_catalog_cache_context: Option, + runtime_context: McpRuntimeContext, + resolved_environment: std::result::Result>, String>, + runtime_auth_provider: Option, + client_elicitation_capability: ElicitationCapability, + client_mcp_extensions: ClientMcpExtensions, + auth_changes: Option>, + protocol_mode: McpProtocolMode, + catalog_item_limit: usize, + ) -> Self { + let is_codex_apps_mcp_server = server_name == CODEX_APPS_MCP_SERVER_NAME; + let reconnect_server_name = server_name.clone(); + let reconnect_tx_event = tx_event.clone(); + let cached_server_info = if is_codex_apps_mcp_server { + codex_apps_tools_cache_context + .as_ref() + .and_then(ConnectorRuntimeContext::cached_server_info) + } else { + None + }; + let startup_complete = Arc::new(AtomicBool::new(false)); + let server_capabilities = Arc::new(StdMutex::new(None)); + let startup = Arc::new(ManagedClientStartup { + server_name, + server, + store_mode, + keyring_backend_kind, + oauth_refresh_mode, + tx_event, + elicitation_requests, + codex_apps_tools_cache_context: codex_apps_tools_cache_context.clone(), + tool_catalog_cache_context: tool_catalog_cache_context.clone(), + runtime_context, + resolved_environment, + runtime_auth_provider, + client_elicitation_capability, + client_mcp_extensions, + auth_changes, + protocol_mode, + catalog_item_limit, + cancel_token: cancel_token.clone(), + startup_complete: Arc::clone(&startup_complete), + server_capabilities: Arc::clone(&server_capabilities), + }); + let client = startup.start(); + let startup_reconnect = is_codex_apps_mcp_server.then(|| { + let startup = Arc::clone(&startup); + Arc::new( + CodexAppsStartupReconnect::new(Arc::new(move || startup.start())) + .with_startup_status_context( + startup_submit_id, + reconnect_server_name, + reconnect_tx_event, + ), + ) + }); + Self { + client, + is_codex_apps_mcp_server, + cached_server_info, + server_capabilities, + codex_apps_tools_cache_context, + tool_catalog_cache_context, + startup_complete, + startup_reconnect, + cancel_token, + } + } + + pub(crate) async fn client(&self) -> Result { + if let Some(client) = self + .startup_reconnect + .as_ref() + .and_then(|reconnect| reconnect.current_client()) + { + return Ok(client); + } + self.client.clone().await + } + + /// Returns the current ready client, including its tool catalog and metadata, + /// without waiting for startup or initiating a reconnection. + pub(crate) fn ready_client(&self) -> Option { + self.startup_reconnect + .as_ref() + .and_then(|reconnect| reconnect.current_client()) + .or_else(|| { + self.client + .peek() + .and_then(|result| result.as_ref().ok()) + .cloned() + }) + } + + pub(crate) fn ready_transport(&self) -> Option> { + self.ready_client().map(|client| client.client) + } + + pub(crate) async fn reconnect_failed_startup(&self) { + let Some(startup_reconnect) = self.startup_reconnect.as_ref() else { + return; + }; + if !self.startup_complete.load(Ordering::Acquire) { + return; + } + if matches!(self.client().await, Err(StartupOutcomeError::Failed { .. })) { + startup_reconnect.reconnect_in_background(); + } + } + + pub(crate) async fn shutdown(&self) { + self.cancel_token.cancel(); + match self.client().await { + Ok(client) => client.client.shutdown().await, + Err(StartupOutcomeError::Cancelled) => {} + Err(error) => { + warn!("failed to initialize MCP client during shutdown: {error:#}"); + } + } + } + + pub(crate) fn has_cached_tools(&self) -> bool { + self.codex_apps_tools_cache_context + .as_ref() + .is_some_and(ConnectorRuntimeContext::has_current_tools) + || self + .tool_catalog_cache_context + .as_ref() + .is_some_and(McpToolCatalogCacheContext::has_tools) + } + + pub(crate) fn cached_tools(&self) -> Option> { + self.cached_tools_or(/*fallback*/ None) + } + + pub(crate) fn cached_tools_or(&self, fallback: Option>) -> Option> { + self.codex_apps_tools_cache_context + .as_ref() + .and_then(ConnectorRuntimeContext::current_tools) + .or_else(|| { + self.tool_catalog_cache_context + .as_ref() + .and_then(|cache| cache.current_tools_or(fallback)) + }) + } + + pub(crate) async fn listed_tools(&self) -> Result, StartupOutcomeError> { + // Plugin provenance is resolved per-session rather than stored in shared cache payloads. + if !self.startup_complete.load(Ordering::Acquire) + && let Some(startup_tools) = self.cached_tools() + { + Ok(startup_tools) + } else { + match self.client().await { + Ok(client) => Ok(client.listed_tools().await), + Err(error) if self.is_codex_apps_mcp_server => self.cached_tools().ok_or(error), + Err(error) => Err(error), + } + } + } +} + +#[derive(Debug, Clone, thiserror::Error)] +pub(crate) enum StartupOutcomeError { + #[error("MCP startup cancelled")] + Cancelled, + // We can't store the original error here because anyhow::Error doesn't implement + // `Clone`. + #[error("MCP startup failed: {error}")] + Failed { + error: String, + is_authentication_required: bool, + }, +} + +impl StartupOutcomeError { + pub(crate) fn is_authentication_required(&self) -> bool { + match self { + Self::Cancelled => false, + Self::Failed { + error, + is_authentication_required, + } => *is_authentication_required || error.contains("Auth required"), + } + } +} + +impl From for StartupOutcomeError { + fn from(error: anyhow::Error) -> Self { + let is_authentication_required = is_authentication_required_error(&error); + Self::Failed { + error: format!("{error:#}"), + is_authentication_required, + } + } +} + +#[instrument(level = "trace", skip_all, fields(server_name = %server_name))] +pub(crate) async fn list_tools_for_client_uncached( + server_name: &str, + is_codex_apps_mcp_server: bool, + codex_apps_refresh_trigger: &'static str, + client: &Arc, + timeout: Option, + catalog_item_limit: usize, + server_instructions: Option<&str>, +) -> Result> { + let fetch_start = Instant::now(); + let protocol_mode = client.protocol_mode(); + let tools = collect_paginated_with_limit("tools/list", timeout, catalog_item_limit, |params| { + let client = Arc::clone(client); + async move { + let response = client + .list_tools_with_connector_ids(params, timeout) + .await?; + let next_cursor = match protocol_mode { + McpProtocolMode::Legacy => None, + McpProtocolMode::V20260728 => response.next_cursor, + }; + Ok((response.tools, next_cursor)) + } + }) + .await? + .into_iter() + .map(|tool| { + tool_info_from_listed_tool( + server_name, + is_codex_apps_mcp_server, + server_instructions, + tool, + ) + }) + .collect(); + if is_codex_apps_mcp_server { + emit_duration( + MCP_TOOLS_FETCH_UNCACHED_DURATION_METRIC, + fetch_start.elapsed(), + &[("trigger", codex_apps_refresh_trigger)], + ); + } else { + emit_duration( + MCP_TOOLS_FETCH_UNCACHED_DURATION_METRIC, + fetch_start.elapsed(), + &[], + ); + } + Ok(tools) +} + +/// Filters disabled connectors, presents declared Codex Apps file parameters to the model as +/// local-path inputs, and adds plugin names to each tool. Plugin membership is resolved by +/// connector ID, falling back to the MCP server when absent. +pub(crate) fn prepare_codex_apps_tools_for_model( + mut tools: Vec, + tool_plugin_context: &ToolPluginContext, +) -> Vec { + tools.retain(|tool| tool_plugin_context.allows_connector_id(tool.connector_id.as_deref())); + for tool in &mut tools { + prepare_openai_file_params_for_model(tool); + let plugin_names = match tool.connector_id.as_deref() { + Some(connector_id) => { + tool_plugin_context.plugin_display_names_for_connector_id(connector_id) + } + None => tool_plugin_context + .plugin_display_names_for_mcp_server_name(tool.server_name.as_str()), + }; + add_plugin_provenance_to_tool(tool, plugin_names); + } + tools +} + +/// Stores plugin names on the tool and appends a model-visible plugin membership note. +fn add_plugin_provenance_to_tool(tool: &mut ToolInfo, plugin_names: &[String]) { + tool.plugin_display_names = plugin_names.to_vec(); + if plugin_names.is_empty() { + return; + } + + let plugin_source_note = if plugin_names.len() == 1 { + format!("This tool is part of plugin `{}`.", plugin_names[0]) + } else { + format!( + "This tool is part of plugins {}.", + plugin_names + .iter() + .map(|plugin_name| format!("`{plugin_name}`")) + .collect::>() + .join(", ") + ) + }; + let description = tool + .tool + .description + .as_deref() + .map(str::trim) + .unwrap_or(""); + let annotated_description = if description.is_empty() { + plugin_source_note + } else if matches!(description.chars().last(), Some('.' | '!' | '?')) { + format!("{description} {plugin_source_note}") + } else { + format!("{description}. {plugin_source_note}") + }; + tool.tool.description = Some(Cow::Owned(annotated_description)); +} + +/// Adds server-scoped plugin names to regular MCP tools without changing their input schemas. +pub(crate) fn prepare_regular_mcp_tools_for_model( + mut tools: Vec, + tool_plugin_context: &ToolPluginContext, +) -> Vec { + for tool in &mut tools { + let plugin_names = + tool_plugin_context.plugin_display_names_for_mcp_server_name(tool.server_name.as_str()); + add_plugin_provenance_to_tool(tool, plugin_names); + } + tools +} + +fn tool_info_from_listed_tool( + server_name: &str, + is_codex_apps_mcp_server: bool, + server_instructions: Option<&str>, + tool: ToolWithConnectorId, +) -> ToolInfo { + if is_codex_apps_mcp_server { + codex_apps_tool_info_from_listed_tool(server_name, server_instructions, tool) + } else { + regular_mcp_tool_info_from_listed_tool(server_name, server_instructions, tool) + } +} + +/// Converts a Codex Apps tool by preserving connector fields, removing connector prefixes from +/// model-visible names and titles, and using the connector description for its tool namespace. +fn codex_apps_tool_info_from_listed_tool( + server_name: &str, + server_instructions: Option<&str>, + tool: ToolWithConnectorId, +) -> ToolInfo { + let mut tool_def = tool.tool; + let connector_id = tool.connector_id; + let connector_name = tool.connector_name; + let connector_description = tool.connector_description; + let callable_name = normalize_codex_apps_callable_name( + &tool_def.name, + connector_id.as_deref(), + connector_name.as_deref(), + ); + let callable_namespace = + normalize_codex_apps_callable_namespace(server_name, connector_name.as_deref()); + if let Some(title) = tool_def.title.as_deref() { + let normalized_title = normalize_codex_apps_tool_title(connector_name.as_deref(), title); + if tool_def.title.as_deref() != Some(normalized_title.as_str()) { + tool_def.title = Some(normalized_title); + } + } + let has_connector_metadata = + connector_id.is_some() || connector_name.is_some() || connector_description.is_some(); + let namespace_description = if has_connector_metadata { + connector_description + } else { + server_instructions.map(str::to_string) + }; + ToolInfo { + server_name: server_name.to_owned(), + supports_parallel_tool_calls: false, + server_origin: None, + callable_name, + callable_namespace, + namespace_description, + tool: tool_def, + openai_file_input_optional_fields: HashMap::new(), + connector_id, + connector_name, + plugin_display_names: Vec::new(), + } +} + +/// Converts a regular MCP tool by removing reserved connector metadata, keeping its raw tool name, +/// and using the MCP server name and instructions for the model-visible namespace. +fn regular_mcp_tool_info_from_listed_tool( + server_name: &str, + server_instructions: Option<&str>, + tool: ToolWithConnectorId, +) -> ToolInfo { + let mut tool_def = tool.tool; + strip_untrusted_connector_meta(&mut tool_def); + ToolInfo { + server_name: server_name.to_owned(), + supports_parallel_tool_calls: false, + server_origin: None, + callable_name: tool_def.name.to_string(), + callable_namespace: server_name.to_string(), + namespace_description: server_instructions.map(str::to_string), + tool: tool_def, + openai_file_input_optional_fields: HashMap::new(), + connector_id: None, + connector_name: None, + plugin_display_names: Vec::new(), + } +} + +fn strip_untrusted_connector_meta(tool: &mut RmcpTool) { + if let Some(meta) = tool.meta.as_mut() { + meta.retain(|key, _| !is_untrusted_connector_meta_key(key)); + } +} + +fn is_untrusted_connector_meta_key(key: &str) -> bool { + UNTRUSTED_CONNECTOR_META_KEYS.contains(&key) +} + +fn resolve_bearer_token( + server_name: &str, + bearer_token_env_var: Option<&str>, +) -> Result> { + let Some(env_var) = bearer_token_env_var else { + return Ok(None); + }; + + match env::var(env_var) { + Ok(value) => { + if value.is_empty() { + Err(anyhow!( + "Environment variable {env_var} for MCP server '{server_name}' is empty" + )) + } else { + Ok(Some(value)) + } + } + Err(env::VarError::NotPresent) => Err(anyhow!( + "Environment variable {env_var} for MCP server '{server_name}' is not set" + )), + Err(env::VarError::NotUnicode(_)) => Err(anyhow!( + "Environment variable {env_var} for MCP server '{server_name}' contains invalid Unicode" + )), + } +} + +fn validate_mcp_server_name(server_name: &str) -> Result<()> { + let re = regex_lite::Regex::new(r"^[a-zA-Z0-9_:@/.-]+$")?; + if !re.is_match(server_name) { + return Err(anyhow!( + "Invalid MCP server name '{server_name}': must match pattern {pattern}", + pattern = re.as_str() + )); + } + Ok(()) +} + +#[instrument(level = "trace", skip_all, fields(server_name = %server_name))] +async fn start_server_task( + server_name: String, + client: Arc, + params: StartServerTaskParams, +) -> Result { + let StartServerTaskParams { + is_codex_apps_mcp_server, + startup_timeout, + tx_event, + elicitation_requests, + codex_apps_tools_cache_context, + tool_catalog_cache_context, + tool_catalog_fetch_ticket, + client_elicitation_capability, + client_mcp_extensions, + auth_changes, + catalog_item_limit, + server_capabilities, + } = params; + let send_elicitation = + elicitation_requests.make_sender(server_name.clone(), tx_event, &client_mcp_extensions); + let mut params = + mcp_initialize_request_params(client_elicitation_capability, client_mcp_extensions); + if auth_changes.is_some() { + params + .capabilities + .experimental + .get_or_insert_default() + .insert( + crate::auth_changes::CAPABILITY.to_string(), + Default::default(), + ); + } + + let requested_capabilities = params.capabilities.clone(); + let started_at = Instant::now(); + let initialize_result = client + .initialize(params, startup_timeout, send_elicitation) + .await; + record_protocol_discovery_metrics( + client.protocol_mode(), + is_codex_apps_mcp_server, + started_at, + &initialize_result, + ); + let initialize_result = initialize_result.map_err(StartupOutcomeError::from)?; + *server_capabilities + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = + Some(serde_json::json!(initialize_result.capabilities)); + + let auth_change_notifications = crate::auth_changes::start( + Arc::clone(&client), + &initialize_result.capabilities, + auth_changes, + ) + .await + .map_err(StartupOutcomeError::from)?; + + let server_disables_tool_catalog_cache = initialize_result + .capabilities + .experimental + .as_ref() + .and_then(|experimental| experimental.get(MCP_TOOL_CATALOG_CACHE_CAPABILITY)) + .and_then(|capability| capability.get(MCP_TOOL_CATALOG_CACHEABLE_PROPERTY)) + .and_then(serde_json::Value::as_bool) + == Some(false); + if server_disables_tool_catalog_cache + && let Some(cache_context) = tool_catalog_cache_context.as_ref() + { + cache_context.disable(); + } + let server_supports_sandbox_state_meta_capability = initialize_result + .capabilities + .experimental + .as_ref() + .and_then(|exp| exp.get(MCP_SANDBOX_STATE_META_CAPABILITY)) + .is_some(); + let codex_apps_tools_cache_context = codex_apps_tools_cache_context.map(|context| { + if server_disables_tool_catalog_cache { + context.without_live_scope() + } else { + // Converted tools can inherit these instructions. Server capabilities and the + // negotiated protocol can also differ between otherwise identical connections. + let mut scope = serde_json::json!([ + requested_capabilities, + initialize_result.protocol_version, + initialize_result.capabilities, + initialize_result.instructions, + initialize_result.server_info, + ]); + scope.sort_all_objects(); + context.with_live_scope(scope.to_string()) + } + }); + let list_start = Instant::now(); + let server_info = + mcp_server_info_from_implementation(&server_name, initialize_result.server_info); + let fetch_ticket = codex_apps_tools_cache_context + .as_ref() + .map(|context| context.begin_fetch(ConnectorRuntimeFetchSource::Startup)); + let tools = list_tools_for_client_uncached( + &server_name, + is_codex_apps_mcp_server, + /*codex_apps_refresh_trigger*/ "initial", + &client, + startup_timeout, + catalog_item_limit, + initialize_result.instructions.as_deref(), + ) + .await + .map_err(StartupOutcomeError::from)?; + let client_tools: Arc<[ToolInfo]> = + match (codex_apps_tools_cache_context.as_ref(), fetch_ticket) { + (Some(context), Some(ticket)) if server_disables_tool_catalog_cache => { + context.publish_runtime_if_newest_accepted(ticket, &server_info, tools.clone()); + tools.into() + } + (Some(context), Some(ticket)) => context + .publish_runtime_if_newest_accepted(ticket, &server_info, tools) + .shared_tools(), + (None, None) => tools.into(), + _ => unreachable!("Codex Apps fetch ticket requires cache context"), + }; + if let (Some(cache_context), Some(fetch_ticket)) = ( + tool_catalog_cache_context.as_ref(), + tool_catalog_fetch_ticket, + ) { + cache_context.publish_if_newest(fetch_ticket, &client_tools); + } + if is_codex_apps_mcp_server || tool_catalog_cache_context.is_some() { + emit_duration( + MCP_TOOLS_LIST_DURATION_METRIC, + list_start.elapsed(), + &[("cache", "miss")], + ); + } + let managed = ManagedClient { + _auth_change_notifications: auth_change_notifications, + client: Arc::clone(&client), + server_info, + tool_catalog: Arc::new(ClientToolCatalog::new( + client_tools, + codex_apps_tools_cache_context + .as_ref() + .and_then(ConnectorRuntimeContext::subscribe), + )), + tool_timeout: None, + server_instructions: initialize_result.instructions, + server_supports_sandbox_state_meta_capability, + codex_apps_tools_cache_context, + }; + + Ok(managed) +} + +fn record_protocol_discovery_metrics( + mode: McpProtocolMode, + is_codex_apps_mcp_server: bool, + started_at: Instant, + result: &Result, +) { + let Some(metrics) = codex_otel::global() else { + return; + }; + + let mode = match mode { + McpProtocolMode::Legacy => "legacy", + McpProtocolMode::V20260728 => "auto", + }; + let outcome = match result { + Ok(result) if result.protocol_version == ProtocolVersion::V_2026_07_28 => "modern", + Ok(_) => "legacy", + Err(_) => "failure", + }; + let mut tags = vec![("mode", mode), ("outcome", outcome)]; + if is_codex_apps_mcp_server { + tags.push(("server_kind", "openai_codex_apps")); + } + let _ = metrics.counter("codex.mcp.protocol_discovery", /*inc*/ 1, &tags); + let _ = metrics.record_duration( + "codex.mcp.protocol_discovery.duration_ms", + started_at.elapsed(), + &tags, + ); +} + +pub(crate) fn mcp_initialize_request_params( + client_elicitation_capability: ElicitationCapability, + client_mcp_extensions: ClientMcpExtensions, +) -> InitializeRequestParams { + let mut capabilities = ClientCapabilities::default(); + capabilities.elicitation = Some(client_elicitation_capability); + let extensions = client_mcp_extensions + .iter() + .filter_map(|(id, settings)| { + settings + .as_object() + .cloned() + .map(|settings| (id.to_string(), settings)) + }) + .collect::>(); + if !extensions.is_empty() { + capabilities.extensions = Some(extensions); + } + InitializeRequestParams::new( + capabilities, + Implementation::new("codex-mcp-client", env!("CARGO_PKG_VERSION")).with_title("Codex"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18) +} + +fn mcp_server_info_from_implementation( + server_name: &str, + server_info: Option, +) -> McpServerInfo { + let server_info = server_info.unwrap_or_else(|| Implementation::new(server_name, "")); + McpServerInfo { + name: server_info.name, + title: server_info.title, + version: server_info.version, + description: server_info.description, + icons: server_info.icons.map(|icons| { + icons + .into_iter() + .filter_map(|icon| serde_json::to_value(icon).ok()) + .collect() + }), + website_url: server_info.website_url, + } +} + +struct StartServerTaskParams { + server_capabilities: Arc>>, + is_codex_apps_mcp_server: bool, + startup_timeout: Option, // TODO: cancel_token should handle this. + tx_event: Option>, + elicitation_requests: ElicitationRequestManager, + codex_apps_tools_cache_context: Option>, + tool_catalog_cache_context: Option, + tool_catalog_fetch_ticket: Option, + client_elicitation_capability: ElicitationCapability, + client_mcp_extensions: ClientMcpExtensions, + auth_changes: Option>, + catalog_item_limit: usize, +} + +#[allow(clippy::too_many_arguments)] +#[instrument(level = "trace", skip_all, fields(server_name = %server_name))] +pub(crate) async fn make_rmcp_client( + server_name: &str, + server: EffectiveMcpServer, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + oauth_refresh_mode: McpOAuthRefreshMode, + runtime_context: McpRuntimeContext, + resolved_environment: std::result::Result>, String>, + runtime_auth_provider: Option, + protocol_mode: McpProtocolMode, +) -> Result { + let config = server.config().clone(); + if matches!(config.auth, McpServerAuth::EmaAuth) { + return Err(StartupOutcomeError::from(anyhow!( + "EMA MCP connections are not enabled in this version" + ))); + } + if matches!(config.auth, McpServerAuth::ChatGpt) + && !config.is_local_environment() + && !has_explicit_http_authorization(&config) + { + return Err(StartupOutcomeError::from(anyhow!( + "executor-owned MCP server `{server_name}` cannot use hosted ChatGPT authentication; configure executor-owned credentials instead" + ))); + } + let resolved_environment = + resolved_environment.map_err(|err| StartupOutcomeError::from(anyhow!(err)))?; + let is_local_environment = config.is_local_environment(); + let oauth_credential_name = config.oauth_credential_name(server_name); + let McpServerConfig { transport, .. } = config; + + match transport { + McpServerTransportConfig::Stdio { + command, + args, + env, + env_vars, + cwd, + } => { + let command_os: OsString = command.into(); + let args_os: Vec = args.into_iter().map(Into::into).collect(); + let env_os = env.map(|env| { + env.into_iter() + .map(|(key, value)| (key.into(), value.into())) + .collect::>() + }); + let launcher = if is_local_environment { + // TODO(starr): Unify local stdio MCP launch with + // `ExecutorStdioServerLauncher` once the executor-backed path + // preserves `LocalStdioServerLauncher` semantics. + Arc::new(LocalStdioServerLauncher::new( + runtime_context.local_process_cwd(), + )) as Arc + } else { + let Some(environment) = resolved_environment.as_ref() else { + unreachable!( + "non-local stdio MCP servers resolve an environment before launch" + ); + }; + Arc::new(ExecutorStdioServerLauncher::new( + environment.get_exec_backend(), + )) as Arc + }; + + let cwd = cwd.map(codex_utils_path_uri::LegacyAppPathString::into_string); + RmcpClient::new_stdio_client_with_protocol_mode( + command_os, + args_os, + env_os, + &env_vars, + cwd, + launcher, + protocol_mode, + ) + .await + .map_err(|err| StartupOutcomeError::from(anyhow!(err))) + } + McpServerTransportConfig::StreamableHttp { + url, + http_headers, + env_http_headers, + bearer_token_env_var, + http_headers_helper: _, + } => { + let http_client = runtime_context + .http_client_for_server(server.config(), resolved_environment.as_ref()) + .map_err(|error| StartupOutcomeError::from(anyhow!(error)))?; + let http_client = maybe_with_openai_docs_source_attribution(&url, http_client); + let executor_resolves_bearer_token = if !is_local_environment + && bearer_token_env_var.is_some() + { + let Some(environment) = resolved_environment.as_ref() else { + return Err(StartupOutcomeError::from(anyhow!( + "non-local HTTP MCP server `{server_name}` did not resolve an execution environment" + ))); + }; + environment + .info() + .await + .map_err(|error| StartupOutcomeError::from(anyhow!(error)))? + .capabilities + .http_header_env_vars + } else { + false + }; + let (http_client, resolved_bearer_token) = if executor_resolves_bearer_token + && let Some(env_var) = bearer_token_env_var.as_ref() + { + ( + Arc::new(ExecutorEnvironmentHttpClient { + bearer_token_env_var: env_var.clone(), + http_client, + }) as Arc, + Some(StreamableHttpBearerToken::ProvidedByHttpClient), + ) + } else { + let token = resolve_bearer_token(server_name, bearer_token_env_var.as_deref()) + .map_err(StartupOutcomeError::from)? + .map(StreamableHttpBearerToken::Resolved); + (http_client, token) + }; + let redirect_mode = if server.is_agent_plugin() { + StreamableHttpRedirectMode::AgentPluginV1 + } else { + StreamableHttpRedirectMode::Legacy + }; + RmcpClient::new_streamable_http_client_with_protocol_mode_and_redirect_mode( + oauth_credential_name.as_ref(), + &url, + resolved_bearer_token, + http_headers, + env_http_headers, + store_mode, + keyring_backend_kind, + http_client, + runtime_auth_provider, + protocol_mode, + redirect_mode, + oauth_refresh_mode, + ) + .await + .map_err(StartupOutcomeError::from) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use codex_protocol::mcp::MCP_APP_UI_EXTENSION_ID; + use codex_protocol::mcp::OPENAI_FORM_EXTENSION_ID; + use pretty_assertions::assert_eq; + use rmcp::model::JsonObject; + use rmcp::model::MetaObject; + use rmcp::transport::auth::AuthError; + + #[test] + fn startup_outcome_error_identifies_authentication_required() { + let error = anyhow::Error::new(AuthError::AuthorizationRequired) + .context("failed to initialize MCP server"); + + let error = StartupOutcomeError::from(error); + + assert!(error.is_authentication_required()); + } + + #[test] + fn missing_server_implementation_uses_configured_server_name() { + assert_eq!( + mcp_server_info_from_implementation("configured-server", /*server_info*/ None), + McpServerInfo { + name: "configured-server".to_string(), + title: None, + version: String::new(), + description: None, + icons: None, + website_url: None, + } + ); + } + + #[test] + fn advertised_server_implementation_takes_precedence_over_configured_name() { + assert_eq!( + mcp_server_info_from_implementation( + "configured-server", + Some( + Implementation::new("advertised-server", "1.2.3") + .with_title("Advertised server") + .with_description("Advertised description") + .with_website_url("https://example.com"), + ), + ), + McpServerInfo { + name: "advertised-server".to_string(), + title: Some("Advertised server".to_string()), + version: "1.2.3".to_string(), + description: Some("Advertised description".to_string()), + icons: None, + website_url: Some("https://example.com".to_string()), + } + ); + } + + #[test] + fn mcp_initialize_advertises_client_extensions() { + let unsupported = mcp_initialize_request_params( + ElicitationCapability::default(), + ClientMcpExtensions::default(), + ); + assert_eq!(unsupported.capabilities.extensions, None); + + let app_ui = serde_json::json!({ + "mimeTypes": ["text/html;profile=mcp-app"], + "futureField": {"preserved": true}, + }); + let supported = mcp_initialize_request_params( + ElicitationCapability::default(), + ClientMcpExtensions::new([ + (OPENAI_FORM_EXTENSION_ID.to_string(), serde_json::json!({})), + (MCP_APP_UI_EXTENSION_ID.to_string(), app_ui.clone()), + ]), + ); + assert_eq!( + supported.capabilities.extensions, + Some(BTreeMap::from([ + (OPENAI_FORM_EXTENSION_ID.to_string(), JsonObject::new()), + ( + MCP_APP_UI_EXTENSION_ID.to_string(), + app_ui.as_object().cloned().expect("app UI settings"), + ), + ])) + ); + } + + fn tool_with_connector_meta() -> RmcpTool { + RmcpTool::new( + "capture_file_upload", + "test tool", + Arc::new(JsonObject::default()), + ) + .with_meta(MetaObject( + serde_json::json!({ + "connector_id": "connector_gmail", + "connector_name": "Gmail", + "connector_display_name": "Gmail", + "connector_description": "Mail connector", + "connectorDescription": "Mail connector", + "connectorFutureField": "future connector metadata", + "CONNECTOR_UPPERCASE": "uppercase connector metadata", + "openai/fileParams": ["file"], + "custom": "kept" + }) + .as_object() + .expect("object") + .clone(), + )) + } + + #[test] + fn custom_mcp_connector_metadata_is_stripped() { + let mut tool = tool_with_connector_meta(); + + strip_untrusted_connector_meta(&mut tool); + + let meta = tool.meta.as_ref().expect("meta"); + for key in [ + "connector_id", + "connector_name", + "connector_display_name", + "connector_description", + "connectorDescription", + ] { + assert!(!meta.0.contains_key(key), "{key} should be stripped"); + } + assert!(meta.0.contains_key("connectorFutureField")); + assert!(meta.0.contains_key("CONNECTOR_UPPERCASE")); + assert!(meta.0.contains_key("openai/fileParams")); + assert_eq!( + meta.0.get("custom").and_then(|value| value.as_str()), + Some("kept") + ); + } + + #[test] + fn codex_apps_connector_metadata_is_preserved() { + let tool = tool_with_connector_meta(); + let expected_tool = tool.clone(); + + let tool_info = tool_info_from_listed_tool( + CODEX_APPS_MCP_SERVER_NAME, + /*is_codex_apps_mcp_server*/ true, + /*server_instructions*/ None, + ToolWithConnectorId { + tool, + connector_id: Some("connector_gmail".to_string()), + connector_name: Some("Gmail".to_string()), + connector_description: Some("Mail connector".to_string()), + }, + ); + + let expected = ToolInfo { + server_name: CODEX_APPS_MCP_SERVER_NAME.to_string(), + supports_parallel_tool_calls: false, + server_origin: None, + callable_name: "capture_file_upload".to_string(), + callable_namespace: "codex_apps__gmail".to_string(), + namespace_description: Some("Mail connector".to_string()), + tool: expected_tool, + openai_file_input_optional_fields: HashMap::new(), + connector_id: Some("connector_gmail".to_string()), + connector_name: Some("Gmail".to_string()), + plugin_display_names: Vec::new(), + }; + assert_eq!( + serde_json::to_value(tool_info).expect("serialize actual tool info"), + serde_json::to_value(expected).expect("serialize expected tool info") + ); + } +} diff --git a/codex-rs/codex-mcp/src/rmcp_client/status.rs b/codex-rs/codex-mcp/src/rmcp_client/status.rs new file mode 100644 index 0000000000000000000000000000000000000000..b166743c77e20a10b645941ce19101cbed7fe3e8 --- /dev/null +++ b/codex-rs/codex-mcp/src/rmcp_client/status.rs @@ -0,0 +1,46 @@ +//! Observe the latest startup attempt without polling it or initiating a retry. + +use codex_protocol::mcp::McpServerConnectionStatus as Status; + +use super::AsyncManagedClient; +use super::StartupOutcomeError; + +impl AsyncManagedClient { + pub(crate) async fn connection_status(&self) -> Status { + if self.cancel_token.is_cancelled() { + return Status::Cancelled; + } + let reconnect_outcome = if let Some(reconnect) = &self.startup_reconnect { + let state = reconnect + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if state.reconnect_in_flight { + return Status::Starting; + } + state + .current_client + .clone() + .map(Ok) + .or_else(|| state.last_error.clone().map(Err)) + } else { + None + }; + let outcome = reconnect_outcome.or_else(|| self.client.peek().cloned()); + match outcome { + Some(Ok(client)) => { + if client.client.is_closed().await { + Status::Failed + } else { + Status::Connected + } + } + Some(Err(error)) if error.is_authentication_required() => { + Status::AuthenticationRequired + } + Some(Err(StartupOutcomeError::Failed { .. })) => Status::Failed, + Some(Err(StartupOutcomeError::Cancelled)) => Status::Cancelled, + None => Status::Starting, + } + } +} diff --git a/codex-rs/codex-mcp/src/runtime.rs b/codex-rs/codex-mcp/src/runtime.rs new file mode 100644 index 0000000000000000000000000000000000000000..85cefccf0d3a493b5e4bdcd3a0821de89ebea0a6 --- /dev/null +++ b/codex-rs/codex-mcp/src/runtime.rs @@ -0,0 +1,1326 @@ +//! Runtime support for Model Context Protocol (MCP) servers. +//! +//! This module contains the thread-owned MCP runtime and data that describes the +//! environment in which MCP servers execute. Transport startup lives in +//! [`crate::rmcp_client`] and connection-set behavior lives in +//! [`crate::connection_manager`]. + +use std::collections::HashMap; +use std::collections::HashSet; +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::Weak; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use arc_swap::ArcSwap; +use async_channel::Sender; +use codex_config::types::McpServerDisabledReason; +use codex_connectors::ConnectorRuntimeContextKey; +use codex_connectors::ConnectorRuntimeManager; +use codex_exec_server::Environment; +use codex_exec_server::EnvironmentManager; +use codex_exec_server::HttpClient; +use codex_exec_server::RouteAwareHttpClient; +use codex_login::AuthManager; +use codex_login::CodexAuth; +use codex_protocol::ThreadId; +use codex_protocol::capabilities::SelectedCapabilityRoot; +use codex_protocol::mcp::CallToolResult; +use codex_protocol::mcp::ClientMcpExtensions; +use codex_protocol::mcp::McpResourceOriginCheckpoint; +use codex_protocol::models::PermissionProfile; +use codex_protocol::protocol::Event; +use codex_protocol::protocol::EventMsg; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::with_http_headers_helper; +use codex_utils_path_uri::PathUri; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ReadResourceResult; +use rmcp::model::RequestId; +use serde::Deserialize; +use serde::Serialize; +use tokio::sync::watch; +use tokio_util::sync::CancellationToken; + +use crate::McpConfig; +use crate::binding::McpBinding; +use crate::client_tool_catalog::CodexAppsToolSnapshot; +use crate::connection_manager::BindingCatalogRevision; +use crate::connection_manager::McpConnectionSet; +use crate::elicitation::ElicitationLifecycle; +use crate::elicitation::ElicitationRequestRouter; +use crate::elicitation::ElicitationReviewerHandle; +use crate::event_stream::McpEventStreamOpener; +use crate::mcp::CODEX_APPS_MCP_SERVER_NAME; +use crate::resource_origin::ResourceOrigins; +use crate::server::EffectiveMcpServer; +use crate::tool_catalog_cache::McpToolCatalogCache; +use crate::tools::ToolInfo; + +/// Controls when one task starts its eligible MCP servers. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum McpStartupPolicy { + /// Start configured servers when their task's MCP runtime is published. + Eager, + /// Start servers with cached tool definitions on first use. + LazyWhenCached, +} + +/// Everything needed to materialize one exact MCP configuration. +pub struct McpRuntimeInput { + pub startup_policy: McpStartupPolicy, + pub config: Arc, + pub plugins_available: bool, + pub ready_selected_capability_roots: Vec, + pub mcp_servers: HashMap, + pub submit_id: String, + pub tx_event: Option>, + pub startup_cancellation_token: CancellationToken, + pub runtime_context: McpRuntimeContext, + pub codex_apps_tools_cache: ConnectorRuntimeManager, + pub tool_catalog_cache: McpToolCatalogCache, + pub codex_apps_tools_cache_key: ConnectorRuntimeContextKey, + pub client_mcp_extensions: ClientMcpExtensions, + pub auth: Option, + pub auth_manager: Option>, + pub elicitation_reviewer: Option, + pub elicitation_lifecycle: Option, +} + +/// Owns all mutable MCP state for one Codex thread. +/// +/// Publication replaces the latest state atomically. Existing bindings retain +/// their exact connections and configuration for as long as they are needed. +pub struct McpRuntime { + current: ArcSwap, + event_stream_cancellation: Mutex, + reconnect_pending: AtomicBool, + elicitation_router: ElicitationRequestRouter, + resource_origins: Mutex, +} + +struct EventStreamCancellation { + event_server_available: bool, + cancel_event_streams_on_server_removal: watch::Sender<()>, + retained_subscription_cancellation: Option>, +} + +struct PublishedMcpRuntime { + connections: Arc, + config: Option>, + auth: Option, + auth_token: Option, + plugins_available: bool, + ready_selected_capability_roots: Vec, + selected_environments: HashMap>, + cached_binding: Mutex>, +} + +fn ensure_host_owned_apps_registration( + current: &PublishedMcpRuntime, + server: &str, +) -> anyhow::Result<()> { + if !current + .config + .as_ref() + .and_then(|config| config.mcp_server_catalog.server(server)) + .is_some_and(|registration| { + registration + .source() + .is_host_owned_apps(server, registration.config()) + }) + { + anyhow::bail!("MCP server '{server}' is not registered by the hosted runtime"); + } + Ok(()) +} + +struct CachedMcpBinding { + catalog_revisions: HashMap, + // Reuse a frozen binding while a model step or caller still needs it. + binding: Weak, +} + +struct McpReconnectGuard<'a> { + pending: &'a AtomicBool, + claimed: bool, +} + +impl Drop for McpReconnectGuard<'_> { + fn drop(&mut self) { + if self.claimed { + self.pending.store(true, Ordering::Release); + } + } +} + +#[derive(Clone)] +pub(crate) struct McpPublicationGate { + published: Option>, +} + +impl McpPublicationGate { + fn pending() -> (watch::Sender, Self) { + let (publish, published) = watch::channel(false); + ( + publish, + Self { + published: Some(published), + }, + ) + } + + pub(crate) fn already_published() -> Self { + Self { published: None } + } + + pub(crate) async fn wait(mut self) -> bool { + let Some(published) = self.published.as_mut() else { + return true; + }; + loop { + if *published.borrow() { + return true; + } + if published.changed().await.is_err() { + return false; + } + } + } +} + +impl McpRuntime { + /// Creates a runtime with no configured servers. + /// + /// This is useful while constructing a thread that must publish a stable + /// runtime handle before its full MCP inputs are available. + pub fn empty(prefix_mcp_tool_names: bool) -> Self { + Self { + current: ArcSwap::from_pointee(PublishedMcpRuntime { + connections: Arc::new(McpConnectionSet::empty(prefix_mcp_tool_names)), + config: None, + auth: None, + auth_token: None, + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + selected_environments: HashMap::new(), + cached_binding: Mutex::new(None), + }), + event_stream_cancellation: Mutex::new(EventStreamCancellation { + event_server_available: false, + cancel_event_streams_on_server_removal: watch::channel(()).0, + retained_subscription_cancellation: None, + }), + reconnect_pending: AtomicBool::new(false), + elicitation_router: ElicitationRequestRouter::default(), + resource_origins: Mutex::default(), + } + } + + /// Updates this thread's bounded resource provenance from a live or restored event. + pub fn observe_event(&self, event: &EventMsg) { + if !matches!( + event, + EventMsg::TurnStarted(_) + | EventMsg::ItemCompleted(_) + | EventMsg::McpToolCallEnd(_) + | EventMsg::ThreadRolledBack(_) + ) { + return; + } + self.resource_origins + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .observe(event); + } + + /// Captures bounded widget provenance for the next compaction checkpoint. + pub fn resource_origin_checkpoint(&self) -> Option { + self.resource_origins + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .checkpoint() + } + + /// Restores widget provenance retained by a compaction checkpoint. + pub fn restore_resource_origin_checkpoint(&self, checkpoint: &McpResourceOriginCheckpoint) { + self.resource_origins + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .restore_checkpoint(checkpoint); + } + + /// Reads a widget through the current binding of the app tool that produced it. + pub async fn read_resource_for_call( + &self, + thread_id: ThreadId, + call_id: &str, + uri: &str, + ) -> anyhow::Result { + let origin = self + .resource_origins + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .find(call_id)?; + let binding = self + .current_binding_for_call(crate::CODEX_APPS_MCP_SERVER_NAME) + .await + .ok_or_else(|| anyhow::anyhow!("codex_apps MCP server is unavailable"))?; + + origin.read(&binding, thread_id, uri).await + } + + pub async fn new(input: McpRuntimeInput) -> Self { + let runtime = Self::empty(input.config.prefix_mcp_tool_names); + runtime.replace(input).await; + runtime + } + + /// Reconciles configured servers and publishes their immutable runtime snapshot. + pub async fn replace(&self, input: McpRuntimeInput) { + let current = self.current.load_full(); + let mut reconnect = McpReconnectGuard { + pending: &self.reconnect_pending, + claimed: self.reconnect_pending.swap(false, Ordering::AcqRel), + }; + self.publish( + input, + (!reconnect.claimed).then_some(current.connections.as_ref()), + ) + .await; + reconnect.claimed = false; + } + + /// Starts fresh connections and returns their complete, refreshed Apps catalog. + pub async fn replace_fresh(&self, input: McpRuntimeInput) -> anyhow::Result> { + self.publish(input, /*previous*/ None).await; + self.latest_hard_refresh_codex_apps_tools_cache().await + } + + async fn publish(&self, input: McpRuntimeInput, previous: Option<&McpConnectionSet>) { + let (publish, publication_gate) = McpPublicationGate::pending(); + let config = Arc::clone(&input.config); + let auth = input.auth.clone(); + let auth_token = auth.as_ref().and_then(|auth| auth.get_token().ok()); + let plugins_available = input.plugins_available; + let ready_selected_capability_roots = input.ready_selected_capability_roots.clone(); + let selected_environments = input.runtime_context.selected_environments.clone(); + let connections = Arc::new( + McpConnectionSet::new( + previous, + publication_gate, + input, + self.elicitation_router.clone(), + ) + .await, + ); + let hosted_event_server_retained = connections.contains_server(CODEX_APPS_MCP_SERVER_NAME) + && config + .mcp_server_catalog + .server(CODEX_APPS_MCP_SERVER_NAME) + .is_some_and(|registration| { + registration + .source() + .is_host_owned_apps(CODEX_APPS_MCP_SERVER_NAME, registration.config()) + }); + let mut cancellation = self + .event_stream_cancellation + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + self.current.store(Arc::new(PublishedMcpRuntime { + connections, + config: Some(config), + auth, + auth_token, + plugins_available, + ready_selected_capability_roots, + selected_environments, + cached_binding: Mutex::new(None), + })); + let _ = publish.send(true); + cancellation.event_server_available = hosted_event_server_retained; + if !hosted_event_server_retained { + cancellation + .cancel_event_streams_on_server_removal + .send_replace(()); + if let Some(retained) = &cancellation.retained_subscription_cancellation { + retained.send_replace(()); + } + } + } + + /// Ensures the next refresh creates fresh connections for every configured server. + pub fn reconnect_on_next_refresh(&self) { + self.reconnect_pending.store(true, Ordering::Release); + } + + /// Captures the latest published configuration and live client handles. + pub async fn current_binding(&self) -> Option> { + self.current_binding_with_requirements(&[], &HashSet::new()) + .await + } + + /// Captures one runtime, waiting for explicitly required servers and selected plugins. + /// Plugin IDs are resolved by the captured connection set, even if a refresh publishes later. + pub async fn current_binding_with_requirements( + &self, + required_servers: &[String], + required_plugins: &HashSet, + ) -> Option> { + Self::binding_from_published_runtime( + self.current.load_full(), + required_servers, + required_plugins, + ) + .await + } + + async fn binding_from_published_runtime( + current: Arc, + required_servers: &[String], + required_plugins: &HashSet, + ) -> Option> { + let config = Arc::clone(current.config.as_ref()?); + let stable_catalog_revisions = current + .connections + .stable_catalog_revisions(required_servers, required_plugins) + .await; + if let Some(catalog_revisions) = &stable_catalog_revisions { + let cached = current + .cached_binding + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(cached) = cached.as_ref() + && &cached.catalog_revisions == catalog_revisions + && let Some(binding) = cached.binding.upgrade() + { + return Some(binding); + } + } + + let binding = Arc::new( + current + .connections + .capture_binding_with_metadata( + config, + current.plugins_available, + required_servers, + required_plugins, + ) + .await, + ); + if let Some(catalog_revisions) = stable_catalog_revisions + && current + .connections + .stable_catalog_revisions(required_servers, required_plugins) + .await + .as_ref() + == Some(&catalog_revisions) + { + let mut cached = current + .cached_binding + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(cached) = cached.as_ref() + && cached.catalog_revisions == catalog_revisions + && let Some(binding) = cached.binding.upgrade() + { + return Some(binding); + } + *cached = Some(CachedMcpBinding { + catalog_revisions, + binding: Arc::downgrade(&binding), + }); + } + Some(binding) + } + + /// Returns whether the published snapshot still belongs to the current credentials. + pub fn current_auth_matches(&self, auth: Option<&CodexAuth>) -> bool { + let current = self.current.load(); + match (current.auth.as_ref(), auth) { + (Some(previous), Some(latest)) => { + previous == latest + && previous.get_account_id() == latest.get_account_id() + && previous.get_chatgpt_user_id() == latest.get_chatgpt_user_id() + && previous.is_fedramp_account() == latest.is_fedramp_account() + && current.auth_token == latest.get_token().ok() + } + (None, None) => true, + (Some(_), None) | (None, Some(_)) => false, + } + } + + /// Detects newly saved credentials for servers whose startup failed authentication. + pub async fn updated_oauth_credentials_after_auth_failure(&self) -> Vec { + let current = self.current.load_full(); + let Some(config) = current.config.as_ref() else { + return Vec::new(); + }; + current + .connections + .updated_oauth_credentials_after_auth_failure(config) + .await + } + + /// Checks the current generation before retrying servers detected outside the refresh gate. + pub async fn has_authentication_failed_servers(&self, server_names: &[String]) -> bool { + self.current + .load_full() + .connections + .authentication_failed_servers() + .await + .into_iter() + .any(|server_name| server_names.contains(&server_name)) + } + + /// Waits for the selected server without capturing an execution binding. + pub async fn wait_for_server_startup(&self, server: &str) { + self.current + .load_full() + .connections + .wait_for_server_startup(server) + .await; + } + + /// Captures the current runtime after its selected server has finished startup. + pub async fn current_binding_for_call(&self, server: &str) -> Option> { + let current = self.current.load_full(); + current.config.as_ref()?; + if !current.connections.wait_for_server_startup(server).await { + return None; + } + Self::binding_from_published_runtime( + current, + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await + } + + /// Returns the latest published configuration without waiting for clients. + pub fn current_config(&self) -> Option> { + self.current.load().config.clone() + } + + pub fn current_ready_selected_capability_roots(&self) -> Vec { + self.current.load().ready_selected_capability_roots.clone() + } + + /// Whether this publication uses the currently ready environment handles. + pub fn current_environments_match( + &self, + environments: &HashMap>, + ) -> bool { + let current = self.current.load(); + current.config.is_some() + && current.selected_environments.len() == environments.len() + && environments.iter().all(|(id, environment)| { + current + .selected_environments + .get(id) + .is_some_and(|published| Arc::ptr_eq(published, environment)) + }) + } + + pub fn elicitations_auto_deny(&self) -> bool { + self.elicitation_router.auto_deny() + } + + pub fn set_elicitations_auto_deny(&self, auto_deny: bool) { + self.elicitation_router.set_auto_deny(auto_deny); + } + + pub fn enable_full_access_form_input(&self) { + self.elicitation_router.enable_full_access_form_input(); + } + + pub async fn resolve_elicitation( + &self, + server_name: String, + id: RequestId, + response: ElicitationResponse, + ) -> anyhow::Result<()> { + self.elicitation_router + .resolve(server_name, id, response) + .await + } + + pub async fn latest_hard_refresh_codex_apps_tools_cache( + &self, + ) -> anyhow::Result> { + self.latest_connections() + .refresh_codex_apps_tools_for_discovery() + .await + } + + /// Refreshes the published Apps client and returns its exact inventory and MCP eligibility. + pub async fn refresh_codex_apps_tools(&self) -> anyhow::Result { + let current = self.current.load_full(); + let config = current + .config + .as_deref() + .ok_or_else(|| anyhow::anyhow!("MCP runtime is not configured"))?; + current + .connections + .refresh_codex_apps_client_catalog(config) + .await + } + + /// Lists the latest known tools for non-model discovery surfaces. + /// + /// Unlike [`Self::current_binding`], this may return cached tools while their + /// client reconnects because callers only inspect tool metadata. + pub async fn latest_list_all_tools(&self) -> Vec { + self.latest_connections().list_all_tools().await + } + + #[allow(clippy::too_many_arguments)] + pub async fn latest_call_tool( + &self, + server: &str, + tool: &str, + environment_id: Option<&str>, + arguments: Option, + meta: Option, + requested_timeout: Option, + wait_for_server: bool, + ) -> anyhow::Result { + self.latest_connections() + .call_tool( + server, + tool, + environment_id, + arguments, + meta, + requested_timeout, + wait_for_server, + ) + .await + } + + pub async fn latest_read_resource( + &self, + server: &str, + params: ReadResourceRequestParams, + ) -> anyhow::Result { + self.latest_connections() + .read_resource(server, params) + .await + } + + pub async fn latest_wait_for_server_ready(&self, server: &str, timeout: Duration) -> bool { + self.latest_connections() + .wait_for_server_ready(server, timeout) + .await + } + + pub async fn validate_required_servers(&self) -> anyhow::Result<()> { + self.latest_connections().validate_required_servers().await + } + + pub fn cancel_startup(&self) { + self.current.load().connections.cancel_startup(); + } + + /// Observes matching published registrations without starting or reconnecting servers. + pub async fn connection_statuses( + &self, + config: &McpConfig, + ) -> std::collections::HashMap { + let current = self.current.load_full(); + let Some(published_config) = current.config.as_ref() else { + return HashMap::new(); + }; + let mut statuses = current.connections.connection_statuses().await; + statuses.retain(|name, _| { + published_config + .mcp_server_catalog + .server(name) + .is_some_and(|server| config.mcp_server_catalog.server(name) == Some(server)) + }); + statuses + } + + pub(crate) fn latest_connections(&self) -> Arc { + Arc::clone(&self.current.load().connections) + } + + pub(crate) fn latest_host_owned_codex_apps_connections( + &self, + ) -> anyhow::Result> { + let current = self.current.load(); + ensure_host_owned_apps_registration(¤t, CODEX_APPS_MCP_SERVER_NAME)?; + Ok(Arc::clone(¤t.connections)) + } + + pub(crate) fn latest_connections_for_event_server( + &self, + server: &str, + ) -> anyhow::Result<(Arc, watch::Receiver<()>)> { + let cancellation = self + .event_stream_cancellation + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let cancel_event_streams_on_server_removal = cancellation + .cancel_event_streams_on_server_removal + .subscribe(); + let current = self.current.load(); + if server == CODEX_APPS_MCP_SERVER_NAME { + ensure_host_owned_apps_registration(¤t, server)?; + } + Ok(( + Arc::clone(¤t.connections), + cancel_event_streams_on_server_removal, + )) + } + + pub(crate) fn event_stream_opener(&self) -> anyhow::Result { + let cancellation = self + .event_stream_cancellation + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let connection = self + .current + .load() + .connections + .event_stream_connection + .clone() + .ok_or_else(|| anyhow::anyhow!("Event subscriptions are unavailable for this task"))?; + let cancel_event_streams_on_server_removal = cancellation + .retained_subscription_cancellation + .as_ref() + .unwrap_or(&cancellation.cancel_event_streams_on_server_removal) + .clone(); + Ok(McpEventStreamOpener { + connection, + cancellation_receiver: cancel_event_streams_on_server_removal.subscribe(), + cancel_event_streams_on_server_removal, + }) + } + + pub(crate) fn forward_event_server_removals_to(&self, retained: watch::Sender<()>) { + let mut cancellation = self + .event_stream_cancellation + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if !cancellation.event_server_available { + retained.send_replace(()); + } + cancellation.retained_subscription_cancellation = Some(retained); + } + + pub async fn shutdown(&self) { + self.latest_connections().shutdown().await; + } +} + +#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct SandboxState { + pub permission_profile: PermissionProfile, + pub codex_linux_sandbox_exe: Option, + pub sandbox_cwd: PathUri, + #[serde(default)] + pub use_legacy_landlock: bool, +} + +/// Runtime context used when resolving per-server MCP environments. +/// +/// `McpConfig` describes what servers exist. This value carries the canonical +/// environment registry plus the host-local cwd used by local MCP processes. +#[derive(Clone)] +pub struct McpRuntimeContext { + environment_manager: Arc, + selected_environments: HashMap>, + local_process_cwd: PathBuf, + local_http_client: Arc, +} + +/// Applies the local HTTP headers helper configured for an MCP server. +/// +/// Callers retain ownership of selecting the underlying HTTP transport. This +/// function centralizes the helper-specific policy checks and decoration used +/// by both MCP runtime startup and standalone OAuth login. +pub fn apply_http_headers_helper( + client: Arc, + config: &codex_config::McpServerConfig, + local_process_cwd: PathBuf, +) -> Result, String> { + if matches!( + config.disabled_reason, + Some(McpServerDisabledReason::Requirements { .. }) + ) { + return Err("the MCP server is disabled by managed requirements".to_string()); + } + let codex_config::McpServerTransportConfig::StreamableHttp { + url, + http_headers_helper: Some(command), + .. + } = &config.transport + else { + return Ok(client); + }; + if !config.is_local_environment() { + return Err("HTTP headers helpers can only run in the local environment".to_string()); + } + with_http_headers_helper(client, url, command, local_process_cwd) + .map_err(|error| error.to_string()) +} + +impl McpRuntimeContext { + pub fn new(environment_manager: Arc, local_process_cwd: PathBuf) -> Self { + let local_http_client = Arc::new( + RouteAwareHttpClient::new(environment_manager.http_client_factory().clone()) + .with_tls_backend_fallback(), + ); + Self { + environment_manager, + selected_environments: HashMap::new(), + local_process_cwd, + local_http_client, + } + } + + /// Pins the concrete environment handles captured for this thread or model step. + pub fn with_selected_environments( + mut self, + selected_environments: HashMap>, + ) -> Self { + self.selected_environments = selected_environments; + self + } + + pub(crate) fn local_process_cwd(&self) -> PathBuf { + self.local_process_cwd.clone() + } + + pub(crate) fn local_http_client(&self) -> Arc { + Arc::clone(&self.local_http_client) + } + + pub(crate) fn resolve_server_environment( + &self, + server_name: &str, + config: &codex_config::McpServerConfig, + ) -> Result>, String> { + // Resolve `"local"` through the shared registry when available. Local + // HTTP is the one current exception: it can use the ambient HTTP client + // even when no local Environment is configured. + if let Some(environment) = self + .selected_environments + .get(&config.environment_id) + .cloned() + .or_else(|| { + self.environment_manager + .get_environment(&config.environment_id) + }) + { + return Ok(Some(environment)); + } + + if config.is_local_environment() { + return match config.transport { + codex_config::McpServerTransportConfig::Stdio { .. } => Err(format!( + "local stdio MCP server `{server_name}` requires a local environment" + )), + codex_config::McpServerTransportConfig::StreamableHttp { .. } => Ok(None), + }; + } + + Err(format!( + "MCP server `{server_name}` references unknown environment id `{}`", + config.environment_id + )) + } + + /// Resolves local MCP's specialized HTTP capability or the selected remote capability. + pub fn resolve_http_client( + &self, + server_name: &str, + config: &codex_config::McpServerConfig, + ) -> Result, String> { + let environment = self.resolve_server_environment(server_name, config)?; + self.http_client_for_server(config, environment.as_ref()) + } + + pub(crate) fn http_client_for_server( + &self, + config: &codex_config::McpServerConfig, + environment: Option<&Arc>, + ) -> Result, String> { + let client = match environment { + Some(environment) if environment.is_remote() => environment.get_http_client(), + Some(_) | None => self.local_http_client(), + }; + apply_http_headers_helper(client, config, self.local_process_cwd()) + } +} + +pub(crate) fn emit_duration(metric: &str, duration: Duration, tags: &[(&str, &str)]) { + if let Some(metrics) = codex_otel::global() { + let _ = metrics.record_duration(metric, duration, tags); + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID; + use codex_config::McpServerConfig; + use codex_config::McpServerTransportConfig; + use codex_exec_server::EnvironmentManager; + use codex_exec_server_test_support::environment_manager_without_environments; + use codex_utils_path_uri::LegacyAppPathString; + use pretty_assertions::assert_eq; + use serde_json::Value; + + use super::*; + + fn stdio_server(environment_id: &str) -> McpServerConfig { + McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::Stdio { + command: "echo".to_string(), + args: Vec::new(), + env: None, + env_vars: Vec::new(), + cwd: None, + }, + environment_id: environment_id.to_string(), + enabled: true, + required: false, + supports_parallel_tool_calls: false, + omit_tools_from: None, + disabled_reason: None, + startup_timeout_sec: None, + tool_timeout_sec: None, + default_tools_approval_mode: None, + enabled_tools: None, + disabled_tools: None, + scopes: None, + oauth: None, + oauth_resource: None, + tools: HashMap::new(), + } + } + + #[tokio::test] + async fn publication_gate_opens_only_for_the_winning_candidate() { + let (publish, gate) = McpPublicationGate::pending(); + let wait = tokio::spawn(gate.wait()); + tokio::task::yield_now().await; + assert!(!wait.is_finished()); + + publish.send(true).expect("publish candidate"); + assert!(wait.await.expect("gate task")); + + let (publish, gate) = McpPublicationGate::pending(); + drop(publish); + assert!(!gate.wait().await); + } + + #[tokio::test] + async fn cached_bindings_follow_the_clients_catalog_revision() -> anyhow::Result<()> { + let codex_home = tempfile::tempdir()?; + let cache_context = ConnectorRuntimeManager::::default().context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ), + ); + let connections = + crate::connection_manager::tests::create_test_manager_with_ready_apps_client( + cache_context, + "search", + /*list_started*/ None, + /*release_list*/ None, + ) + .await?; + // Complete the fixture's shared startup future before testing stable reuse. + connections.list_all_tools().await; + let mut config = crate::mcp::tests::test_mcp_config(codex_home.path().to_path_buf()); + config.server_permission_profiles.insert( + CODEX_APPS_MCP_SERVER_NAME.to_string(), + PermissionProfile::default(), + ); + let published = Arc::new(PublishedMcpRuntime { + connections: Arc::clone(&connections), + config: Some(Arc::new(config)), + auth: None, + auth_token: None, + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + selected_environments: HashMap::new(), + cached_binding: Mutex::new(None), + }); + let before = McpRuntime::binding_from_published_runtime( + Arc::clone(&published), + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await + .expect("initial binding"); + let repeated = McpRuntime::binding_from_published_runtime( + Arc::clone(&published), + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await + .expect("cached initial binding"); + assert!(Arc::ptr_eq(&before, &repeated)); + + connections.refresh_codex_apps_tools_for_discovery().await?; + + let refreshed = McpRuntime::binding_from_published_runtime( + Arc::clone(&published), + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await + .expect("refreshed binding"); + assert!(!Arc::ptr_eq(&before, &refreshed)); + let call = refreshed + .prepare_call(CODEX_APPS_MCP_SERVER_NAME, "search") + .expect("refreshed call"); + let error = call + .call_with_preparation(/*requested_timeout*/ None, || async { + Err(anyhow::anyhow!("reached refreshed call preparation")) + }) + .await + .expect_err("stop before dispatch"); + assert!( + error + .to_string() + .contains("reached refreshed call preparation") + ); + let repeated = McpRuntime::binding_from_published_runtime( + Arc::clone(&published), + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await + .expect("cached refreshed binding"); + assert!(Arc::ptr_eq(&refreshed, &repeated)); + let released = Arc::downgrade(&refreshed); + drop(refreshed); + drop(repeated); + assert!( + released.upgrade().is_none(), + "the runtime cache must not pin an unused binding" + ); + Ok(()) + } + + #[tokio::test] + async fn cached_bindings_are_scoped_to_the_published_runtime() { + let published = Arc::new(PublishedMcpRuntime { + connections: Arc::new(McpConnectionSet::empty(/*prefix_mcp_tool_names*/ true)), + config: Some(Arc::new(crate::mcp::tests::test_mcp_config( + std::env::temp_dir(), + ))), + auth: None, + auth_token: None, + plugins_available: false, + ready_selected_capability_roots: Vec::new(), + selected_environments: HashMap::new(), + cached_binding: Mutex::new(None), + }); + let first = McpRuntime::binding_from_published_runtime( + Arc::clone(&published), + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await + .expect("first binding"); + let repeated = McpRuntime::binding_from_published_runtime( + Arc::clone(&published), + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await + .expect("repeated binding"); + assert!(Arc::ptr_eq(&first, &repeated)); + + let previous = Arc::into_inner(published).expect("published runtime has no other owners"); + let republished = Arc::new(PublishedMcpRuntime { + cached_binding: Mutex::new(None), + ..previous + }); + let refreshed = McpRuntime::binding_from_published_runtime( + republished, + /*required_servers*/ &[], + /*required_plugins*/ &HashSet::new(), + ) + .await + .expect("republished binding"); + assert!(!Arc::ptr_eq(&first, &refreshed)); + } + + fn http_server(environment_id: &str) -> McpServerConfig { + McpServerConfig { + auth: Default::default(), + transport: McpServerTransportConfig::StreamableHttp { + url: "http://127.0.0.1:1".to_string(), + bearer_token_env_var: None, + http_headers: None, + env_http_headers: None, + http_headers_helper: None, + }, + environment_id: environment_id.to_string(), + ..stdio_server(environment_id) + } + } + + #[test] + fn sandbox_state_serializes_skip_missing_entries_as_missing_path_behavior() { + let sandbox_cwd = PathUri::from_host_native_path( + std::env::current_dir().expect("current directory should be available"), + ) + .expect("current directory should convert to a URI"); + let sandbox_state = SandboxState { + permission_profile: PermissionProfile::workspace_write(), + codex_linux_sandbox_exe: None, + sandbox_cwd, + use_legacy_landlock: false, + }; + + let serialized = serde_json::to_value(&sandbox_state).expect("serialize sandbox state"); + let serialized_text = serde_json::to_string(&serialized).expect("serialize JSON text"); + assert!( + !serialized_text.contains("generated_default_path"), + "MCP sandbox metadata must preserve FileSystemPath's stable wire variants" + ); + assert!( + !serialized_text.contains("generated_default_special"), + "MCP sandbox metadata must preserve FileSystemPath's stable wire variants" + ); + + let entries = serialized + .pointer("/permissionProfile/file_system/entries") + .and_then(Value::as_array) + .expect("workspace-write profile should contain filesystem entries"); + let skip_missing_entries = entries + .iter() + .filter(|entry| { + entry.get("missing_path_behavior").and_then(Value::as_str) == Some("skip") + }) + .collect::>(); + assert!( + !skip_missing_entries.is_empty(), + "skip-missing entries should be represented as optional missing_path_behavior" + ); + assert!( + skip_missing_entries.iter().all(|entry| { + matches!( + entry.pointer("/path/type").and_then(Value::as_str), + Some("path" | "special") + ) + }), + "skip-missing entries should use the stable path/special variants" + ); + + let deserialized: SandboxState = + serde_json::from_value(serialized).expect("deserialize sandbox state"); + assert_eq!( + deserialized.permission_profile, + sandbox_state.permission_profile + ); + } + + #[test] + fn local_stdio_requires_local_stdio_availability() { + let runtime_context = McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + PathBuf::from("/tmp"), + ); + + let error = match runtime_context + .resolve_server_environment("stdio", &stdio_server(DEFAULT_MCP_SERVER_ENVIRONMENT_ID)) + { + Ok(_) => panic!("local stdio MCP should require a local environment"), + Err(error) => error, + }; + assert_eq!( + error, + "local stdio MCP server `stdio` requires a local environment" + ); + } + + #[test] + fn local_http_does_not_require_local_stdio_availability() { + let runtime_context = McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + PathBuf::from("/tmp"), + ); + + let resolved_runtime = match runtime_context + .resolve_server_environment("http", &http_server(DEFAULT_MCP_SERVER_ENVIRONMENT_ID)) + { + Ok(resolved_runtime) => resolved_runtime, + Err(error) => panic!("local HTTP MCP should resolve: {error}"), + }; + assert!(resolved_runtime.is_none()); + } + + #[tokio::test] + async fn local_http_client_is_shared_across_resolution_and_context_clones() { + for environment_manager in [ + EnvironmentManager::default_for_tests(), + environment_manager_without_environments(), + ] { + let runtime_context = + McpRuntimeContext::new(Arc::new(environment_manager), PathBuf::from("/tmp")); + let config = http_server(DEFAULT_MCP_SERVER_ENVIRONMENT_ID); + let first_client = runtime_context + .resolve_http_client("http", &config) + .expect("first local HTTP capability should resolve"); + let repeated_client = runtime_context + .resolve_http_client("http", &config) + .expect("repeated local HTTP capability should resolve"); + let resolved_environment = runtime_context + .resolve_server_environment("http", &config) + .expect("local HTTP environment should resolve"); + let startup_client = runtime_context + .http_client_for_server(&config, resolved_environment.as_ref()) + .expect("startup local HTTP capability should resolve"); + let cloned_client = runtime_context + .clone() + .resolve_http_client("http", &config) + .expect("cloned local HTTP capability should resolve"); + + assert!(Arc::ptr_eq(&first_client, &repeated_client)); + assert!(Arc::ptr_eq(&first_client, &startup_client)); + assert!(Arc::ptr_eq(&first_client, &cloned_client)); + } + } + + #[test] + fn unknown_explicit_environment_is_rejected() { + let runtime_context = McpRuntimeContext::new( + Arc::new(environment_manager_without_environments()), + PathBuf::from("/tmp"), + ); + + let error = + match runtime_context.resolve_server_environment("stdio", &stdio_server("remote")) { + Ok(_) => panic!("unknown MCP environment should fail"), + Err(error) => error, + }; + assert_eq!( + error, + "MCP server `stdio` references unknown environment id `remote`" + ); + } + + #[tokio::test] + async fn explicit_remote_stdio_and_http_accept_named_environment() { + let runtime_context = McpRuntimeContext::new( + Arc::new( + EnvironmentManager::create_for_tests( + Some("ws://127.0.0.1:8765".to_string()), + /*local_runtime_paths*/ None, + ) + .await, + ), + PathBuf::from("/tmp"), + ); + + let mut remote_stdio = stdio_server("remote"); + let McpServerTransportConfig::Stdio { cwd, .. } = &mut remote_stdio.transport else { + unreachable!("stdio helper should build stdio transport"); + }; + *cwd = Some(LegacyAppPathString::from_path(&std::env::temp_dir())); + for resolved_runtime in [ + runtime_context.resolve_server_environment("stdio", &remote_stdio), + runtime_context.resolve_server_environment("http", &http_server("remote")), + ] { + let resolved_runtime = match resolved_runtime { + Ok(resolved_runtime) => resolved_runtime, + Err(error) => panic!("remote MCP should resolve: {error}"), + }; + assert!(resolved_runtime.is_some()); + } + + let mut remote_http_with_helper = http_server("remote"); + let McpServerTransportConfig::StreamableHttp { + http_headers_helper, + .. + } = &mut remote_http_with_helper.transport + else { + unreachable!("HTTP helper should build streamable HTTP transport"); + }; + *http_headers_helper = Some("helper-that-must-not-run".to_string()); + let error = match runtime_context.resolve_http_client("http", &remote_http_with_helper) { + Ok(_) => panic!("remote HTTP helper should be rejected"), + Err(error) => error, + }; + assert_eq!( + error, + "HTTP headers helpers can only run in the local environment" + ); + + let remote_http = http_server("remote"); + let remote_environment = runtime_context + .resolve_server_environment("http", &remote_http) + .expect("remote HTTP MCP should resolve") + .expect("remote HTTP MCP should have an environment"); + let remote_client = runtime_context + .resolve_http_client("http", &remote_http) + .expect("remote HTTP capability should resolve"); + assert!(Arc::ptr_eq( + &remote_client, + &remote_environment.get_http_client() + )); + } + + #[tokio::test] + async fn remote_stdio_accepts_foreign_absolute_cwd() { + let runtime_context = McpRuntimeContext::new( + Arc::new( + EnvironmentManager::create_for_tests( + Some("ws://127.0.0.1:8765".to_string()), + /*local_runtime_paths*/ None, + ) + .await, + ), + PathBuf::from("/tmp"), + ); + let mut remote_stdio = stdio_server("remote"); + let McpServerTransportConfig::Stdio { cwd, .. } = &mut remote_stdio.transport else { + unreachable!("stdio helper should build stdio transport"); + }; + *cwd = Some( + PathUri::parse("file:///C:/plugins/demo") + .expect("foreign cwd URI") + .into(), + ); + + let resolved_runtime = + match runtime_context.resolve_server_environment("stdio", &remote_stdio) { + Ok(resolved_runtime) => resolved_runtime, + Err(error) => panic!("foreign cwd should resolve: {error}"), + }; + assert!(resolved_runtime.is_some()); + } + + #[tokio::test] + async fn local_stdio_accepts_local_environment_when_available() { + let runtime_context = McpRuntimeContext::new( + Arc::new(EnvironmentManager::default_for_tests()), + PathBuf::from("/tmp"), + ); + + let resolved_runtime = match runtime_context + .resolve_server_environment("stdio", &stdio_server(DEFAULT_MCP_SERVER_ENVIRONMENT_ID)) + { + Ok(resolved_runtime) => resolved_runtime, + Err(error) => panic!("local stdio MCP should resolve: {error}"), + }; + assert!(resolved_runtime.is_some()); + } +} diff --git a/codex-rs/codex-mcp/src/server.rs b/codex-rs/codex-mcp/src/server.rs new file mode 100644 index 0000000000000000000000000000000000000000..1b6c1ecb5b8fd75914b57696090ac2f39d83e858 --- /dev/null +++ b/codex-rs/codex-mcp/src/server.rs @@ -0,0 +1,439 @@ +use std::collections::HashMap; +use std::ffi::OsString; +use std::path::PathBuf; +use std::sync::Arc; + +use crate::runtime::McpRuntimeContext; +use codex_api::SharedAuthProvider; +use codex_config::AppToolApproval; +use codex_config::McpServerAuth; +use codex_config::McpServerConfig; +use codex_config::McpServerTransportConfig; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_connectors::ConnectorRuntimeContextKey; +use codex_exec_server::Environment; +use codex_login::CodexAuth; +use codex_protocol::mcp::ClientMcpExtensions; +use codex_rmcp_client::McpOAuthRefreshMode; +use codex_rmcp_client::StoredOAuthCredentialSnapshot; +use codex_rmcp_client::StoredOAuthTokens; +use codex_utils_path_uri::PathUri; +use rmcp::model::ElicitationCapability; +use tracing::warn; + +/// MCP server after runtime additions have been applied. +#[derive(Debug, Clone)] +pub struct EffectiveMcpServer { + config: McpServerConfig, + agent_plugin: bool, +} + +impl EffectiveMcpServer { + pub fn configured(config: McpServerConfig) -> Self { + Self { + config, + agent_plugin: false, + } + } + + pub fn with_agent_plugin(mut self, agent_plugin: bool) -> Self { + self.agent_plugin = agent_plugin; + self + } + + pub fn config(&self) -> &McpServerConfig { + &self.config + } + + pub fn enabled(&self) -> bool { + self.config.enabled + } + + pub fn required(&self) -> bool { + self.config.required + } + + pub fn is_agent_plugin(&self) -> bool { + self.agent_plugin + } +} + +pub(crate) fn has_explicit_http_authorization(config: &McpServerConfig) -> bool { + let McpServerTransportConfig::StreamableHttp { + bearer_token_env_var, + http_headers, + env_http_headers, + .. + } = &config.transport + else { + return false; + }; + + if bearer_token_env_var.is_some() + || env_http_headers + .as_ref() + .is_some_and(|headers| !headers.is_empty()) + { + return false; + } + + http_headers.as_ref().is_some_and(|headers| { + headers.iter().any(|(name, value)| { + name.eq_ignore_ascii_case("authorization") + && !value.trim().is_empty() + && value + .bytes() + .all(|byte| byte == b'\t' || (byte >= b' ' && byte != 0x7f)) + }) + }) +} + +/// Inputs that determine the identity of a live MCP connection. +/// +/// Tool policy and presentation metadata intentionally do not appear here: +/// those belong to a publication and can change without reconnecting. +#[derive(Clone)] +pub(crate) struct McpServerConnectionIdentity { + auth: McpServerAuth, + transport: McpServerTransportConfig, + environment_id: String, + host_plugin_root: Option, + oauth_store: Option<(OAuthCredentialsStoreMode, AuthKeyringBackendKind)>, + oauth_refresh_mode: Option, + oauth_credentials: Result, String>, + pub(crate) oauth_store_was_contended: bool, + resolved_environment: Result>, String>, + local_stdio_fallback_cwd: Option, + referenced_environment_variables: Vec<(String, Option)>, + runtime_auth: Option, + runtime_auth_token: Option, + codex_apps_cache_identity: Option<(PathBuf, ConnectorRuntimeContextKey)>, + client_elicitation_capability: ElicitationCapability, + client_mcp_extensions: ClientMcpExtensions, + agent_plugin: bool, +} + +impl McpServerConnectionIdentity { + #[allow(clippy::too_many_arguments)] + pub(crate) fn new( + server_name: &str, + server: &EffectiveMcpServer, + host_plugin_root: Option<&PathUri>, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + oauth_refresh_mode: McpOAuthRefreshMode, + resolved_environment: &Result>, String>, + runtime_context: &McpRuntimeContext, + runtime_auth_provider: Option<&SharedAuthProvider>, + auth: Option<&CodexAuth>, + codex_apps_cache_identity: Option<(PathBuf, ConnectorRuntimeContextKey)>, + client_elicitation_capability: ElicitationCapability, + client_mcp_extensions: ClientMcpExtensions, + previous_identity: Option<&Self>, + ) -> Self { + let config = server.config(); + let valid_http_header_value = |value: &str| { + value + .bytes() + .all(|byte| byte == b'\t' || (byte >= b' ' && byte != 0x7f)) + }; + let stored_oauth_url = if runtime_auth_provider.is_none() + && !matches!(config.auth, McpServerAuth::EmaAuth) + && (!matches!(config.auth, McpServerAuth::ChatGpt) || config.is_local_environment()) + { + match &config.transport { + McpServerTransportConfig::StreamableHttp { + url, + bearer_token_env_var: None, + http_headers, + env_http_headers, + http_headers_helper: _, + } if !http_headers.as_ref().is_some_and(|headers| { + headers.iter().any(|(name, value)| { + name.eq_ignore_ascii_case("authorization") && valid_http_header_value(value) + }) + }) && !env_http_headers.as_ref().is_some_and(|headers| { + headers.iter().any(|(name, env_var)| { + name.eq_ignore_ascii_case("authorization") + && std::env::var(env_var).is_ok_and(|value| { + !value.trim().is_empty() && valid_http_header_value(&value) + }) + }) + }) => + { + Some(url) + } + McpServerTransportConfig::StreamableHttp { .. } + | McpServerTransportConfig::Stdio { .. } => None, + } + } else { + None + }; + let oauth_credentials = stored_oauth_url.map_or(Ok(None), |url| { + let credential_name = config.oauth_credential_name(server_name); + StoredOAuthCredentialSnapshot::for_runtime_refresh( + previous_identity.and_then(|previous_identity| { + previous_identity + .oauth_credentials + .as_ref() + .ok() + .and_then(Option::as_ref) + .filter(|_| { + previous_identity.oauth_store + == Some((store_mode, keyring_backend_kind)) + }) + }), + credential_name.as_ref(), + url, + store_mode, + keyring_backend_kind, + ) + .map_err(|error| { + warn!(server_name, %error, "failed to read stored MCP OAuth credentials"); + error.to_string() + }) + }); + let local_stdio_fallback_cwd = (config.is_local_environment() + && matches!( + config.transport, + McpServerTransportConfig::Stdio { cwd: None, .. } + | McpServerTransportConfig::StreamableHttp { + http_headers_helper: Some(_), + .. + } + )) + .then(|| runtime_context.local_process_cwd()); + let referenced_environment_variables = referenced_environment_variables(config); + let runtime_auth = runtime_auth_provider.and(auth).cloned(); + let runtime_auth_token = runtime_auth.as_ref().and_then(|auth| auth.get_token().ok()); + let oauth_store_was_contended = oauth_credentials + .as_ref() + .ok() + .and_then(Option::as_ref) + .is_some_and(StoredOAuthCredentialSnapshot::store_was_contended); + + Self { + auth: config.auth.clone(), + transport: config.transport.clone(), + environment_id: config.environment_id.clone(), + host_plugin_root: host_plugin_root.cloned(), + oauth_store: stored_oauth_url + .is_some() + .then_some((store_mode, keyring_backend_kind)), + oauth_refresh_mode: stored_oauth_url.is_some().then_some(oauth_refresh_mode), + oauth_credentials, + oauth_store_was_contended, + resolved_environment: resolved_environment.clone(), + local_stdio_fallback_cwd, + referenced_environment_variables, + runtime_auth, + runtime_auth_token, + codex_apps_cache_identity, + client_elicitation_capability, + client_mcp_extensions, + agent_plugin: server.is_agent_plugin(), + } + } + + pub(crate) fn has_same_connection_config(&self, other: &Self) -> bool { + let same_runtime_auth = match (&self.runtime_auth, &other.runtime_auth) { + (Some(CodexAuth::AgentIdentity(left)), Some(CodexAuth::AgentIdentity(right))) => { + left.record() == right.record() + } + (Some(left), Some(right)) => { + left == right + && left.get_account_id() == right.get_account_id() + && left.get_chatgpt_user_id() == right.get_chatgpt_user_id() + && left.is_fedramp_account() == right.is_fedramp_account() + } + (None, None) => true, + (Some(_), None) | (None, Some(_)) => false, + }; + self.auth == other.auth + && self.transport == other.transport + && self.environment_id == other.environment_id + && self.host_plugin_root == other.host_plugin_root + && self.oauth_store == other.oauth_store + && self.oauth_refresh_mode == other.oauth_refresh_mode + && same_resolved_environment(&self.resolved_environment, &other.resolved_environment) + && self.local_stdio_fallback_cwd == other.local_stdio_fallback_cwd + && self.referenced_environment_variables == other.referenced_environment_variables + && same_runtime_auth + && self.runtime_auth_token == other.runtime_auth_token + && self.codex_apps_cache_identity == other.codex_apps_cache_identity + && self.client_elicitation_capability == other.client_elicitation_capability + && self.client_mcp_extensions == other.client_mcp_extensions + && self.agent_plugin == other.agent_plugin + } + + pub(crate) fn oauth_credentials(&self) -> Result, &String> { + self.oauth_credentials.as_ref().map(|credentials| { + credentials + .as_ref() + .map(StoredOAuthCredentialSnapshot::credentials) + }) + } + + pub(crate) fn oauth_credentials_changed( + &self, + server_name: &str, + config: &McpServerConfig, + ) -> bool { + let Some((store_mode, keyring_backend_kind)) = self.oauth_store else { + return false; + }; + let McpServerTransportConfig::StreamableHttp { url, .. } = &self.transport else { + return false; + }; + + let credential_name = config.oauth_credential_name(server_name); + let current_credentials = match self.oauth_credentials.as_ref() { + Ok(Some(credentials)) => credentials.reload( + credential_name.as_ref(), + url, + store_mode, + keyring_backend_kind, + ), + Ok(None) | Err(_) => StoredOAuthCredentialSnapshot::for_runtime_refresh( + /*previous*/ None, + credential_name.as_ref(), + url, + store_mode, + keyring_backend_kind, + ) + .map(|snapshot| snapshot.map(|snapshot| snapshot.credentials().clone())), + }; + + match current_credentials { + Ok(Some(current_credentials)) => { + self.oauth_credentials() != Ok(Some(¤t_credentials)) + } + Ok(None) => false, + Err(error) => { + warn!(server_name, %error, "failed to read stored MCP OAuth credentials"); + false + } + } + } +} + +impl PartialEq for McpServerConnectionIdentity { + fn eq(&self, other: &Self) -> bool { + self.has_same_connection_config(other) && self.oauth_credentials == other.oauth_credentials + } +} + +fn same_resolved_environment( + left: &Result>, String>, + right: &Result>, String>, +) -> bool { + match (left, right) { + (Ok(Some(left)), Ok(Some(right))) => Arc::ptr_eq(left, right), + (Ok(None), Ok(None)) => true, + (Err(left), Err(right)) => left == right, + (Ok(_), Ok(_)) | (Ok(_), Err(_)) | (Err(_), Ok(_)) => false, + } +} + +fn referenced_environment_variables(config: &McpServerConfig) -> Vec<(String, Option)> { + let mut names = match &config.transport { + McpServerTransportConfig::Stdio { env_vars, .. } => env_vars + .iter() + .filter(|env_var| !env_var.is_remote_source()) + .map(|env_var| env_var.name().to_string()) + .collect::>(), + McpServerTransportConfig::StreamableHttp { + bearer_token_env_var, + env_http_headers, + .. + } => bearer_token_env_var + .iter() + .filter(|name| config.is_local_environment() || std::env::var_os(name).is_some()) + .chain(env_http_headers.iter().flat_map(|headers| headers.values())) + .cloned() + .collect(), + }; + names.sort(); + names.dedup(); + names + .into_iter() + .map(|name| { + let value = std::env::var_os(&name); + (name, value) + }) + .collect() +} + +/// Transport origin retained for metrics and diagnostics after server launch. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum McpServerOrigin { + Stdio, + StreamableHttp(String), +} + +impl McpServerOrigin { + pub fn as_str(&self) -> &str { + match self { + Self::Stdio => "stdio", + Self::StreamableHttp(origin) => origin, + } + } + + fn from_transport(transport: &McpServerTransportConfig) -> Option { + match transport { + McpServerTransportConfig::StreamableHttp { url, .. } => { + let parsed = url::Url::parse(url).ok()?; + Some(Self::StreamableHttp(parsed.origin().ascii_serialization())) + } + McpServerTransportConfig::Stdio { .. } => Some(Self::Stdio), + } + } +} + +/// Semantic metadata that must survive after the server is launched. +#[derive(Debug, Clone)] +pub(crate) struct McpServerMetadata { + pub environment_id: String, + pub pollutes_memory: bool, + pub origin: Option, + pub supports_parallel_tool_calls: bool, + pub default_tools_approval_mode: Option, + pub tool_approval_modes: HashMap, +} + +impl McpServerMetadata { + pub fn tool_approval_mode(&self, tool_name: &str) -> AppToolApproval { + self.tool_approval_modes + .get(tool_name) + .copied() + .or(self.default_tools_approval_mode) + .unwrap_or_default() + } +} + +impl From<&EffectiveMcpServer> for McpServerMetadata { + fn from(server: &EffectiveMcpServer) -> Self { + let config = server.config(); + Self { + environment_id: config.environment_id.clone(), + pollutes_memory: true, + origin: McpServerOrigin::from_transport(&config.transport), + supports_parallel_tool_calls: config.supports_parallel_tool_calls, + default_tools_approval_mode: config.default_tools_approval_mode, + tool_approval_modes: config + .tools + .iter() + .filter_map(|(name, config)| { + config + .approval_mode + .map(|approval_mode| (name.clone(), approval_mode)) + }) + .collect(), + } + } +} + +#[cfg(test)] +#[path = "server_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/server_tests.rs b/codex-rs/codex-mcp/src/server_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..7365f75573ec8cd1dc372eead321f43c931267a5 --- /dev/null +++ b/codex-rs/codex-mcp/src/server_tests.rs @@ -0,0 +1,43 @@ +use super::referenced_environment_variables; +use codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID; +use codex_config::McpServerConfig; +use pretty_assertions::assert_eq; + +#[test] +fn remote_http_connections_track_host_headers_but_not_executor_bearer_tokens() { + let mut config: McpServerConfig = serde_json::from_value(serde_json::json!({ + "url": "https://example.com/mcp", + "environment_id": "executor-1", + "bearer_token_env_var": "NODE_REPL_AUTH_TOKEN", + "env_http_headers": {"X-Api-Key": "PATH"}, + })) + .expect("remote MCP configuration should deserialize"); + + assert_eq!( + referenced_environment_variables(&config), + vec![("PATH".to_string(), std::env::var_os("PATH"))], + ); + + let remote_host_bearer: McpServerConfig = serde_json::from_value(serde_json::json!({ + "url": "https://example.com/mcp", + "environment_id": "executor-1", + "bearer_token_env_var": "PATH", + })) + .expect("host-resolved remote MCP configuration should deserialize"); + assert_eq!( + referenced_environment_variables(&remote_host_bearer), + vec![("PATH".to_string(), std::env::var_os("PATH"))], + ); + + config.environment_id = DEFAULT_MCP_SERVER_ENVIRONMENT_ID.to_string(); + assert_eq!( + referenced_environment_variables(&config), + vec![ + ( + "NODE_REPL_AUTH_TOKEN".to_string(), + std::env::var_os("NODE_REPL_AUTH_TOKEN"), + ), + ("PATH".to_string(), std::env::var_os("PATH")), + ], + ); +} diff --git a/codex-rs/codex-mcp/src/snapshots/codex_mcp__connection_manager__tests__mcp_init_error_display_quotes_server_names.snap b/codex-rs/codex-mcp/src/snapshots/codex_mcp__connection_manager__tests__mcp_init_error_display_quotes_server_names.snap new file mode 100644 index 0000000000000000000000000000000000000000..555a69efc47d92e9cfcafed9c3a7a70bfa0f5dee --- /dev/null +++ b/codex-rs/codex-mcp/src/snapshots/codex_mcp__connection_manager__tests__mcp_init_error_display_quotes_server_names.snap @@ -0,0 +1,19 @@ +--- +source: codex-mcp/src/connection_manager_tests.rs +expression: "displays.join(\"\\n\\n\")" +--- +MCP client for `npm:@scope/package.name` timed out after 30 seconds. Add or adjust `startup_timeout_sec` in your config.toml: +[mcp_servers."npm:@scope/package.name"] +startup_timeout_sec = XX + +GitHub MCP does not support OAuth. Log in by adding a personal access token (https://github.com/settings/personal-access-tokens) to your environment and config.toml: +[mcp_servers."npm:@scope/package.name"] +bearer_token_env_var = CODEX_GITHUB_PERSONAL_ACCESS_TOKEN + +MCP client for `server.name` timed out after 30 seconds. Add or adjust `startup_timeout_sec` in your config.toml: +[mcp_servers."server.name"] +startup_timeout_sec = XX + +GitHub MCP does not support OAuth. Log in by adding a personal access token (https://github.com/settings/personal-access-tokens) to your environment and config.toml: +[mcp_servers."server.name"] +bearer_token_env_var = CODEX_GITHUB_PERSONAL_ACCESS_TOKEN diff --git a/codex-rs/codex-mcp/src/tool_catalog_cache.rs b/codex-rs/codex-mcp/src/tool_catalog_cache.rs new file mode 100644 index 0000000000000000000000000000000000000000..72d723b9817a7afa85244f69838d1ad8855a464f --- /dev/null +++ b/codex-rs/codex-mcp/src/tool_catalog_cache.rs @@ -0,0 +1,405 @@ +use std::collections::BTreeMap; +use std::collections::hash_map::DefaultHasher; +use std::hash::Hash; +use std::hash::Hasher; +use std::num::NonZeroUsize; +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::MutexGuard; +use std::sync::Weak; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_config::McpServerAuth; +use codex_config::McpServerConfig; +use codex_config::McpServerTransportConfig; +use codex_exec_server::Environment; +use codex_protocol::mcp::ClientMcpExtensions; +use lru::LruCache; +use rmcp::model::ElicitationCapability; +use sha1::Digest; +use sha1::Sha1; +use tokio::time::Instant; + +use crate::McpProtocolMode; +use crate::McpRuntimeContext; +use crate::ToolInfo; +use crate::server::McpServerConnectionIdentity; +use crate::server::has_explicit_http_authorization; + +const TOOL_CATALOG_CACHE_CAPACITY: usize = 32; +const TOOL_CATALOG_CACHE_TTL: Duration = Duration::from_secs(30 * 60); + +/// Process-scoped cache of recent reusable tool definitions for MCP servers. +#[derive(Clone)] +pub struct McpToolCatalogCache { + entries: Arc>>>, +} + +impl Default for McpToolCatalogCache { + fn default() -> Self { + Self { + entries: Arc::new(Mutex::new(LruCache::new( + NonZeroUsize::new(TOOL_CATALOG_CACHE_CAPACITY).unwrap_or(NonZeroUsize::MIN), + ))), + } + } +} + +struct ToolCatalogCacheEntry { + state: Mutex, + next_fetch_generation: AtomicU64, +} + +#[derive(Default)] +struct ToolCatalogCacheState { + snapshot: Option, + optional_startup_deadline: Option, + last_accepted_generation: u64, + disabled_by_server: bool, +} + +struct OptionalStartupDeadline { + grace: Duration, + deadline: Instant, +} + +struct ToolCatalogSnapshot { + tools: Vec, + published_at: Instant, +} + +#[derive(Clone)] +pub(crate) struct McpToolCatalogCacheContext { + entry: Arc, +} + +pub(crate) struct McpToolCatalogFetchTicket { + generation: u64, +} + +impl McpToolCatalogCache { + pub(crate) fn context( + &self, + server_name: &str, + config: &McpServerConfig, + runtime_context: &McpRuntimeContext, + resolved_environment: Option<&Arc>, + client_context: (&ElicitationCapability, &ClientMcpExtensions), + connection_identity: Option<(&McpServerConnectionIdentity, McpProtocolMode, bool)>, + ) -> Option { + let identity = ToolCatalogIdentity::new( + server_name, + config, + runtime_context, + resolved_environment, + client_context, + connection_identity, + )?; + let entry = lock_unpoisoned(&self.entries) + .get_or_insert(identity, || Arc::new(ToolCatalogCacheEntry::default())) + .clone(); + Some(McpToolCatalogCacheContext { entry }) + } +} + +impl Default for ToolCatalogCacheEntry { + fn default() -> Self { + Self { + state: Mutex::new(ToolCatalogCacheState::default()), + next_fetch_generation: AtomicU64::new(0), + } + } +} + +impl McpToolCatalogCacheContext { + pub(crate) fn has_tools(&self) -> bool { + self.current_revision().is_some() + } + + /// Identifies the current usable catalog without cloning its tool definitions. + /// The entry is fixed for a connection; accepted publications advance its revision. + pub(crate) fn current_revision(&self) -> Option { + let state = lock_unpoisoned(&self.entry.state); + let snapshot = state.snapshot.as_ref()?; + (!state.disabled_by_server + && !snapshot.tools.is_empty() + && snapshot.published_at.elapsed() <= TOOL_CATALOG_CACHE_TTL) + .then_some(state.last_accepted_generation) + } + + pub(crate) fn optional_startup_deadline( + &self, + default_deadline: Instant, + startup_grace: Duration, + ) -> Instant { + let mut state = lock_unpoisoned(&self.entry.state); + if state.disabled_by_server + || state + .snapshot + .as_ref() + .is_some_and(|snapshot| snapshot.published_at.elapsed() <= TOOL_CATALOG_CACHE_TTL) + { + return default_deadline; + } + let cached_deadline = + state + .optional_startup_deadline + .get_or_insert(OptionalStartupDeadline { + grace: startup_grace, + deadline: default_deadline, + }); + if cached_deadline.grace != startup_grace { + *cached_deadline = OptionalStartupDeadline { + grace: startup_grace, + deadline: default_deadline, + }; + } + cached_deadline.deadline + } + + pub(crate) fn current_tools(&self) -> Option> { + self.current_tools_or(/*fallback*/ None) + } + + /// Prefers the current catalog, retaining a capture's fallback across expiry but not opt-out. + pub(crate) fn current_tools_or( + &self, + fallback: Option>, + ) -> Option> { + let state = lock_unpoisoned(&self.entry.state); + if state.disabled_by_server { + return None; + } + state + .snapshot + .as_ref() + .filter(|snapshot| snapshot.published_at.elapsed() <= TOOL_CATALOG_CACHE_TTL) + .map(|snapshot| snapshot.tools.clone()) + .or(fallback) + } + + pub(crate) fn begin_fetch(&self) -> McpToolCatalogFetchTicket { + McpToolCatalogFetchTicket { + generation: self + .entry + .next_fetch_generation + .fetch_add(1, Ordering::Relaxed) + + 1, + } + } + + pub(crate) fn disable(&self) { + let mut state = lock_unpoisoned(&self.entry.state); + state.disabled_by_server = true; + state.snapshot = None; + } + + pub(crate) fn publish_if_newest(&self, ticket: McpToolCatalogFetchTicket, tools: &[ToolInfo]) { + let mut state = lock_unpoisoned(&self.entry.state); + if state.disabled_by_server || ticket.generation <= state.last_accepted_generation { + return; + } + + let mut tools = tools.to_vec(); + for tool in &mut tools { + // Tool annotations affect approval and parallelism decisions, so only the live + // connection may supply them. + tool.tool.annotations = None; + } + state.last_accepted_generation = ticket.generation; + state.optional_startup_deadline = None; + state.snapshot = Some(ToolCatalogSnapshot { + tools, + published_at: Instant::now(), + }); + } +} + +fn lock_unpoisoned(mutex: &Mutex) -> MutexGuard<'_, T> { + mutex + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) +} + +struct ToolCatalogIdentity { + server_name: String, + transport: ToolCatalogTransportIdentity, + environment: Option>, + local_stdio_fallback_cwd: Option, +} + +impl PartialEq for ToolCatalogIdentity { + fn eq(&self, other: &Self) -> bool { + self.server_name == other.server_name + && self.transport == other.transport + && self.local_stdio_fallback_cwd == other.local_stdio_fallback_cwd + && match (&self.environment, &other.environment) { + (Some(environment), Some(other)) => Weak::ptr_eq(environment, other), + (None, None) => true, + _ => false, + } + } +} + +impl Eq for ToolCatalogIdentity {} + +impl Hash for ToolCatalogIdentity { + fn hash(&self, state: &mut H) { + self.server_name.hash(state); + self.transport.hash(state); + self.local_stdio_fallback_cwd.hash(state); + self.environment + .as_ref() + .map(|environment| Weak::as_ptr(environment) as usize) + .hash(state); + } +} + +impl ToolCatalogIdentity { + fn new( + server_name: &str, + config: &McpServerConfig, + runtime_context: &McpRuntimeContext, + environment: Option<&Arc>, + client_context: (&ElicitationCapability, &ClientMcpExtensions), + connection_identity: Option<(&McpServerConnectionIdentity, McpProtocolMode, bool)>, + ) -> Option { + let transport = + ToolCatalogTransportIdentity::new(config, client_context, connection_identity)?; + Some(Self { + server_name: server_name.to_string(), + transport, + environment: environment.map(Arc::downgrade), + local_stdio_fallback_cwd: matches!( + &config.transport, + McpServerTransportConfig::Stdio { cwd: None, .. } + ) + .then(|| runtime_context.local_process_cwd()), + }) + } +} + +#[derive(PartialEq, Eq, Hash)] +enum ToolCatalogTransportIdentity { + Stdio { fingerprint: [u8; 20] }, + StreamableHttp { fingerprint: [u8; 20] }, +} + +impl ToolCatalogTransportIdentity { + fn new( + config: &McpServerConfig, + client_context: (&ElicitationCapability, &ClientMcpExtensions), + connection_identity: Option<(&McpServerConnectionIdentity, McpProtocolMode, bool)>, + ) -> Option { + let (client_elicitation_capability, client_mcp_extensions) = client_context; + if let McpServerTransportConfig::StreamableHttp { + url, + bearer_token_env_var, + http_headers, + env_http_headers, + http_headers_helper, + } = &config.transport + { + // Helper output is a dynamic credential identity that cannot be represented by config. + if http_headers_helper.is_some() { + return None; + } + let (connection_identity, protocol_mode, agent_plugin) = connection_identity?; + if config.oauth.is_some() + || config.scopes.is_some() + || config.oauth_resource.is_some() + || (matches!(config.auth, McpServerAuth::ChatGpt) + && !has_explicit_http_authorization(config)) + || (!has_explicit_http_authorization(config) + && connection_identity.oauth_credentials().ok()?.is_some()) + { + return None; + } + + let mut hasher = Sha1::new(); + hasher.update( + serde_json::to_vec(&( + url, + bearer_token_env_var, + http_headers + .as_ref() + .map(|headers| headers.iter().collect::>()), + env_http_headers + .as_ref() + .map(|headers| headers.iter().collect::>()), + &config.auth, + &config.environment_id, + agent_plugin, + protocol_mode.preferred_protocol_version().as_str(), + client_elicitation_capability, + client_mcp_extensions.iter().collect::>(), + )) + .ok()?, + ); + let mut env_vars = bearer_token_env_var + .iter() + .chain(env_http_headers.iter().flat_map(|headers| headers.values())) + .collect::>(); + env_vars.sort_unstable(); + env_vars.dedup(); + for name in env_vars { + hasher.update(name.as_bytes()); + let mut value_hasher = DefaultHasher::new(); + std::env::var_os(name).hash(&mut value_hasher); + hasher.update(value_hasher.finish().to_le_bytes()); + } + return Some(Self::StreamableHttp { + fingerprint: hasher.finalize().into(), + }); + } + let McpServerTransportConfig::Stdio { + command, + args, + env, + env_vars, + cwd, + } = &config.transport + else { + return None; + }; + if env_vars + .iter() + .any(codex_config::McpServerEnvVar::is_remote_source) + { + return None; + } + + let mut hasher = Sha1::new(); + let env = env.as_ref().map(|env| { + env.iter() + .map(|(key, value)| (key.as_str(), value.as_str())) + .collect::>() + }); + hasher.update( + serde_json::to_vec(&( + command, + args, + env, + env_vars, + cwd, + &config.environment_id, + client_elicitation_capability, + client_mcp_extensions.iter().collect::>(), + )) + .ok()?, + ); + for env_var in env_vars { + hasher.update(env_var.name().as_bytes()); + let mut value_hasher = DefaultHasher::new(); + std::env::var_os(env_var.name()).hash(&mut value_hasher); + hasher.update(value_hasher.finish().to_le_bytes()); + } + + Some(Self::Stdio { + fingerprint: hasher.finalize().into(), + }) + } +} diff --git a/codex-rs/codex-mcp/src/tools.rs b/codex-rs/codex-mcp/src/tools.rs new file mode 100644 index 0000000000000000000000000000000000000000..4afc7561f88cc1a63317ecaaf1f37ee27383f680 --- /dev/null +++ b/codex-rs/codex-mcp/src/tools.rs @@ -0,0 +1,316 @@ +//! MCP tool metadata, filtering, and name normalization. +//! +//! Raw MCP tool identities must be preserved for protocol calls, while +//! model-visible tool names must be sanitized, deduplicated, and kept within API +//! limits. This module owns that translation as well as the shared [`ToolInfo`] +//! type. + +use std::collections::HashMap; +use std::collections::HashSet; + +use codex_config::McpServerConfig; +use codex_protocol::ToolName; +use rmcp::model::Tool; +use serde::Deserialize; +use serde::Serialize; +use sha1::Digest; +use sha1::Sha1; +use tracing::warn; + +use crate::mcp::sanitize_responses_api_tool_name; + +const LEGACY_MCP_TOOL_NAME_PREFIX: &str = "mcp__"; + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct ToolInfo { + /// Raw MCP server name used for routing the tool call. + pub server_name: String, + /// Whether calls routed to this server may run in parallel. + #[serde(default)] + pub supports_parallel_tool_calls: bool, + /// MCP server origin used for telemetry and diagnostics, when known. + #[serde(default)] + pub server_origin: Option, + /// Model-visible tool name used in Responses API tool declarations. + #[serde(rename = "tool_name", alias = "callable_name")] + pub callable_name: String, + /// Model-visible namespace used for deferred tool loading. + #[serde(rename = "tool_namespace", alias = "callable_namespace")] + pub callable_namespace: String, + /// Model-visible namespace description. + // Keep the old serialized field name readable for cached ToolInfo values. + #[serde(default, alias = "connector_description")] + pub namespace_description: Option, + /// Raw MCP tool definition; `tool.name` is sent back to the MCP server. + pub tool: Tool, + /// Optional provided-file fields accepted by each declared `openai/fileParams` + /// argument. This is derived from the raw MCP schema before file arguments are + /// masked as local paths for the model. + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub openai_file_input_optional_fields: HashMap>, + pub connector_id: Option, + pub connector_name: Option, + #[serde(default)] + pub plugin_display_names: Vec, +} + +impl ToolInfo { + pub fn canonical_tool_name(&self) -> ToolName { + ToolName::namespaced(self.callable_namespace.clone(), self.callable_name.clone()) + } +} + +/// A tool is allowed to be used if both are true: +/// 1. enabled is None (no allowlist is set) or the tool is explicitly enabled. +/// 2. The tool is not explicitly disabled. +#[derive(Default, Clone)] +pub(crate) struct ToolFilter { + pub(crate) enabled: Option>, + pub(crate) disabled: HashSet, +} + +impl ToolFilter { + pub(crate) fn from_config(cfg: &McpServerConfig) -> Self { + let enabled = cfg + .enabled_tools + .as_ref() + .map(|tools| tools.iter().cloned().collect::>()); + let disabled = cfg + .disabled_tools + .as_ref() + .map(|tools| tools.iter().cloned().collect::>()) + .unwrap_or_default(); + + Self { enabled, disabled } + } + + pub(crate) fn allows(&self, tool_name: &str) -> bool { + if let Some(enabled) = &self.enabled + && !enabled.contains(tool_name) + { + return false; + } + + !self.disabled.contains(tool_name) + } +} + +pub(crate) fn filter_tools(tools: Vec, filter: &ToolFilter) -> Vec { + tools + .into_iter() + .filter(|tool| filter.allows(&tool.tool.name)) + .collect() +} + +/// Returns MCP tools with model-visible names normalized. +/// +/// Raw MCP server/tool names are kept on each [`ToolInfo`] for protocol calls, while +/// `callable_namespace` / `callable_name` are sanitized and, when necessary, hashed so +/// every model-visible name is unique and <= 128 bytes. +/// +/// When `prefix_mcp_tool_names` is true, the historical `mcp__` namespace +/// prefix is added except for tools from `non_prefixed_mcp_tool_servers`. +pub(crate) fn normalize_tools_for_model_with_prefix( + tools: I, + prefix_mcp_tool_names: bool, + non_prefixed_mcp_tool_servers: &[String], +) -> Vec +where + I: IntoIterator, +{ + let mut seen_raw_names = HashSet::new(); + let mut candidates = Vec::new(); + for tool in tools { + let raw_namespace_identity = format!( + "{}\0{}\0{}", + tool.server_name, + tool.callable_namespace, + tool.connector_id.as_deref().unwrap_or_default() + ); + let raw_tool_identity = format!( + "{}\0{}\0{}", + raw_namespace_identity, tool.callable_name, tool.tool.name + ); + if !seen_raw_names.insert(raw_tool_identity.clone()) { + warn!("skipping duplicated tool {}", tool.tool.name); + continue; + } + + let callable_namespace = callable_namespace_with_prefix( + &sanitize_responses_api_tool_name(&tool.callable_namespace), + prefix_mcp_tool_names && !non_prefixed_mcp_tool_servers.contains(&tool.server_name), + ); + + candidates.push(CallableToolCandidate { + callable_namespace, + callable_name: sanitize_responses_api_tool_name(&tool.callable_name), + raw_namespace_identity, + raw_tool_identity, + tool, + }); + } + + let mut namespace_identities_by_base = HashMap::>::new(); + for candidate in &candidates { + namespace_identities_by_base + .entry(candidate.callable_namespace.clone()) + .or_default() + .insert(candidate.raw_namespace_identity.clone()); + } + let colliding_namespaces = namespace_identities_by_base + .into_iter() + .filter_map(|(namespace, identities)| (identities.len() > 1).then_some(namespace)) + .collect::>(); + for candidate in &mut candidates { + if colliding_namespaces.contains(&candidate.callable_namespace) { + candidate.callable_namespace = append_namespace_hash_suffix( + &candidate.callable_namespace, + &candidate.raw_namespace_identity, + ); + } + } + + let mut tool_identities_by_base = HashMap::<(String, String), HashSet>::new(); + for candidate in &candidates { + tool_identities_by_base + .entry(( + candidate.callable_namespace.clone(), + candidate.callable_name.clone(), + )) + .or_default() + .insert(candidate.raw_tool_identity.clone()); + } + let colliding_tools = tool_identities_by_base + .into_iter() + .filter_map(|(key, identities)| (identities.len() > 1).then_some(key)) + .collect::>(); + for candidate in &mut candidates { + if colliding_tools.contains(&( + candidate.callable_namespace.clone(), + candidate.callable_name.clone(), + )) { + candidate.callable_name = + append_hash_suffix(&candidate.callable_name, &candidate.raw_tool_identity); + } + } + + candidates.sort_by(|left, right| left.raw_tool_identity.cmp(&right.raw_tool_identity)); + + let mut used_names = HashSet::new(); + let mut model_tools = Vec::new(); + for mut candidate in candidates { + let (callable_namespace, callable_name) = unique_callable_parts( + &candidate.callable_namespace, + &candidate.callable_name, + &candidate.raw_tool_identity, + &mut used_names, + MCP_TOOL_NAME_DELIMITER.len(), + ); + candidate.tool.callable_namespace = callable_namespace; + candidate.tool.callable_name = callable_name; + model_tools.push(candidate.tool); + } + model_tools +} + +#[derive(Debug)] +struct CallableToolCandidate { + tool: ToolInfo, + raw_namespace_identity: String, + raw_tool_identity: String, + callable_namespace: String, + callable_name: String, +} + +const MCP_TOOL_NAME_DELIMITER: &str = "__"; +const MAX_TOOL_NAME_LENGTH: usize = 128; +const CALLABLE_NAME_HASH_LEN: usize = 12; +fn callable_namespace_with_prefix(namespace: &str, prefix_mcp_tool_names: bool) -> String { + if !prefix_mcp_tool_names || namespace.starts_with(LEGACY_MCP_TOOL_NAME_PREFIX) { + namespace.to_string() + } else { + format!("{LEGACY_MCP_TOOL_NAME_PREFIX}{namespace}") + } +} + +fn sha1_hex(s: &str) -> String { + let mut hasher = Sha1::new(); + hasher.update(s.as_bytes()); + let sha1 = hasher.finalize(); + format!("{sha1:x}") +} + +fn callable_name_hash_suffix(raw_identity: &str) -> String { + let hash = sha1_hex(raw_identity); + format!("_{}", &hash[..CALLABLE_NAME_HASH_LEN]) +} + +fn append_hash_suffix(value: &str, raw_identity: &str) -> String { + format!("{value}{}", callable_name_hash_suffix(raw_identity)) +} + +fn append_namespace_hash_suffix(namespace: &str, raw_identity: &str) -> String { + if let Some(namespace) = namespace.strip_suffix(MCP_TOOL_NAME_DELIMITER) { + format!( + "{}{}{}", + namespace, + callable_name_hash_suffix(raw_identity), + MCP_TOOL_NAME_DELIMITER + ) + } else { + append_hash_suffix(namespace, raw_identity) + } +} + +fn truncate_name(value: &str, max_len: usize) -> String { + value.chars().take(max_len).collect() +} + +fn fit_callable_parts_with_hash( + namespace: &str, + tool_name: &str, + raw_identity: &str, + reserved_len: usize, +) -> (String, String) { + let suffix = callable_name_hash_suffix(raw_identity); + let max_tool_len = MAX_TOOL_NAME_LENGTH.saturating_sub(namespace.len() + reserved_len); + if max_tool_len >= suffix.len() { + let prefix_len = max_tool_len - suffix.len(); + return ( + namespace.to_string(), + format!("{}{}", truncate_name(tool_name, prefix_len), suffix), + ); + } + + let max_namespace_len = MAX_TOOL_NAME_LENGTH.saturating_sub(suffix.len() + reserved_len); + (truncate_name(namespace, max_namespace_len), suffix) +} + +fn unique_callable_parts( + namespace: &str, + tool_name: &str, + raw_identity: &str, + used_names: &mut HashSet, + reserved_len: usize, +) -> (String, String) { + let model_name = format!("{namespace}{tool_name}"); + if model_name.len() + reserved_len <= MAX_TOOL_NAME_LENGTH && used_names.insert(model_name) { + return (namespace.to_string(), tool_name.to_string()); + } + + let mut attempt = 0_u32; + loop { + let hash_input = if attempt == 0 { + raw_identity.to_string() + } else { + format!("{raw_identity}\0{attempt}") + }; + let (namespace, tool_name) = + fit_callable_parts_with_hash(namespace, tool_name, &hash_input, reserved_len); + let model_name = format!("{namespace}{tool_name}"); + if used_names.insert(model_name) { + return (namespace, tool_name); + } + attempt = attempt.saturating_add(1); + } +} diff --git a/codex-rs/codex-mcp/src/trusted_access.rs b/codex-rs/codex-mcp/src/trusted_access.rs new file mode 100644 index 0000000000000000000000000000000000000000..baed4690cb8d2dc6fdc6e713766fbb89cbd18102 --- /dev/null +++ b/codex-rs/codex-mcp/src/trusted_access.rs @@ -0,0 +1,297 @@ +use std::sync::Arc; +use std::time::Duration; + +use crate::connection_manager::McpConnectionSet; +use crate::runtime::McpRuntimeInput; +use crate::server::McpServerMetadata; +use crate::server::McpServerOrigin; +use crate::tools::ToolInfo; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpHeader; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use codex_login::AuthManager; +use codex_login::CodexAuth; +use codex_protocol::auth::AuthMode; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Map; +use serde_json::Value; +use serde_json::json; + +pub(crate) const ENTITLEMENT_CONTEXT_KEY: &str = "openai/entitlementContext"; +const MAX_VERIFIED_ACCESS_RESPONSE_BYTES: usize = 1024 * 1024; +const REQUESTED_ENTITLEMENTS_KEY: &str = "openai/requestedEntitlements"; +const CYBER_TRUSTED_ACCESS_ENTITLEMENT: &str = "cyber_trusted_access"; +const TRUSTED_ACCESS_TIMEOUT: Duration = Duration::from_millis(2_500); + +impl McpConnectionSet { + /// Installed and task-selected plugins may request supported advisory entitlement metadata. + /// Model calls use the local, read-only, zero-argument boundary. + pub(crate) async fn add_trusted_access_context( + &self, + tool: &ToolInfo, + server: &McpServerMetadata, + arguments: Option<&Value>, + meta: Option, + ) -> Option { + if tool + .tool + .meta + .as_deref() + .and_then(|meta| meta.get(REQUESTED_ENTITLEMENTS_KEY)) + .and_then(Value::as_array) + .is_some_and(|entitlements| { + entitlements.iter().all(Value::is_string) + && entitlements.iter().any(|entitlement| { + entitlement.as_str() == Some(CYBER_TRUSTED_ACCESS_ENTITLEMENT) + }) + }) + && self + .plugin_id_for_mcp_server_name(&tool.server_name) + .is_some() + && arguments.is_none_or(|arguments| arguments.as_object().is_some_and(Map::is_empty)) + && server.environment_id == codex_config::DEFAULT_MCP_SERVER_ENVIRONMENT_ID + && matches!(server.origin, Some(McpServerOrigin::Stdio)) + && tool + .tool + .annotations + .as_ref() + .and_then(|annotations| annotations.read_only_hint) + == Some(true) + && let Some(context) = self.trusted_access.as_ref() + { + context.add_context(meta).await + } else { + meta + } + } +} + +#[derive(Deserialize)] +struct VerifiedAccessResponse { + programs: Vec, +} + +#[derive(Deserialize)] +struct VerifiedAccessProgram { + state: VerifiedAccessState, + grants: Vec, +} + +#[derive(Clone, Copy, Deserialize)] +#[serde(rename_all = "snake_case")] +enum VerifiedAccessState { + Active, + Inactive, + Unavailable, +} + +#[derive(Deserialize)] +struct VerifiedAccessGrant { + level: TrustedAccessLevel, + source: VerifiedAccessSource, +} + +#[derive(Deserialize, Serialize)] +#[serde(rename_all = "snake_case")] +enum TrustedAccessLevel { + Tac1, + Tac2, + Tac3, + Government, +} + +#[derive(Deserialize)] +#[serde(rename_all = "snake_case")] +enum VerifiedAccessSource { + Individual, + Organization, +} + +/// Fetches account-bound verified access for trusted, host-owned MCP metadata. +/// Callers must authorize the receiving plugin and tool before attaching it. +pub struct TrustedAccessContext { + auth: CodexAuth, + auth_manager: Arc, + chatgpt_base_url: String, + http_client: Arc, +} + +impl TrustedAccessContext { + pub(crate) fn from_runtime(input: &McpRuntimeInput) -> Option { + let auth = input.auth.as_ref()?; + if !matches!( + auth.api_auth_mode(), + AuthMode::Chatgpt | AuthMode::ChatgptAuthTokens + ) { + return None; + } + Some(Self::new( + auth.clone(), + input.auth_manager.clone()?, + input.config.chatgpt_base_url.clone(), + input.runtime_context.local_http_client(), + )) + } + + pub fn new( + auth: CodexAuth, + auth_manager: Arc, + chatgpt_base_url: String, + http_client: Arc, + ) -> Self { + Self { + auth, + auth_manager, + chatgpt_base_url, + http_client, + } + } + + /// Replaces caller-supplied entitlement metadata with a fresh verified result. + pub async fn add_context(&self, meta: Option) -> Option { + let mut meta = match meta { + Some(Value::Object(meta)) => meta, + None => Map::new(), + other => return other, + }; + meta.remove(ENTITLEMENT_CONTEXT_KEY); + + let status = tokio::time::timeout(TRUSTED_ACCESS_TIMEOUT, self.fetch_status()) + .await + .ok() + .flatten() + .unwrap_or_else(|| { + json!({ + "schemaVersion": 1, + "status": "unknown", + "grants": [], + "stale": false + }) + }); + meta.insert( + ENTITLEMENT_CONTEXT_KEY.to_string(), + json!({ + "schemaVersion": 1, + "entitlements": { "cyber_trusted_access": status } + }), + ); + Some(Value::Object(meta)) + } + + async fn fetch_status(&self) -> Option { + let auth = self.auth_manager.auth().await?; + let expected_account_id = self + .auth + .get_account_id() + .filter(|account_id| !account_id.trim().is_empty())?; + let account_id = auth + .get_account_id() + .filter(|account_id| !account_id.trim().is_empty())?; + if !matches!( + self.auth.api_auth_mode(), + AuthMode::Chatgpt | AuthMode::ChatgptAuthTokens + ) || !matches!( + auth.api_auth_mode(), + AuthMode::Chatgpt | AuthMode::ChatgptAuthTokens + ) || account_id != expected_account_id + || auth.get_chatgpt_user_id() != self.auth.get_chatgpt_user_id() + || auth.is_workspace_account() != self.auth.is_workspace_account() + || auth.is_fedramp_account() != self.auth.is_fedramp_account() + { + return None; + } + + let headers = codex_model_provider::auth_provider_from_auth(&auth) + .to_auth_headers() + .iter() + .map(|(name, value)| { + Some(HttpHeader { + name: name.as_str().to_string(), + value: value.to_str().ok()?.to_string(), + value_env_var: None, + }) + }) + .collect::>>()?; + let (response, mut response_body) = self + .http_client + .http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: format!( + "{}/accounts/verified_access", + self.chatgpt_base_url.trim_end_matches('/') + ), + headers, + body: None, + timeout_ms: Some(TRUSTED_ACCESS_TIMEOUT.as_millis() as u64), + redirect_policy: HttpRedirectPolicy::Stop, + request_id: "trusted-access-status".to_string(), + stream_response: true, + }) + .await + .ok()?; + if response.status != 200 { + return None; + } + let mut response_bytes = Vec::new(); + while let Some(chunk) = response_body.recv().await.ok()? { + if chunk.len() > MAX_VERIFIED_ACCESS_RESPONSE_BYTES - response_bytes.len() { + return None; + } + response_bytes.extend_from_slice(&chunk); + } + let response: VerifiedAccessResponse = serde_json::from_slice(&response_bytes).ok()?; + let mut programs = response + .programs + .into_iter() + .filter(|program| program.get("program").and_then(Value::as_str) == Some("cyber")); + let program = programs.next()?; + if programs.next().is_some() { + return None; + } + let program: VerifiedAccessProgram = serde_json::from_value(program).ok()?; + + if self.auth_manager.auth_cached().is_none_or(|current| { + !matches!( + current.api_auth_mode(), + AuthMode::Chatgpt | AuthMode::ChatgptAuthTokens + ) || current.get_account_id().as_deref() != Some(account_id.as_str()) + || current.get_chatgpt_user_id() != auth.get_chatgpt_user_id() + || current.is_workspace_account() != auth.is_workspace_account() + || current.is_fedramp_account() != auth.is_fedramp_account() + }) { + return None; + } + + let status = match program.state { + VerifiedAccessState::Active if !program.grants.is_empty() => "granted", + VerifiedAccessState::Inactive if program.grants.is_empty() => "not_granted", + VerifiedAccessState::Unavailable if program.grants.is_empty() => "unknown", + VerifiedAccessState::Active + | VerifiedAccessState::Inactive + | VerifiedAccessState::Unavailable => return None, + }; + let grants = program + .grants + .into_iter() + .map(|grant| { + let source = match grant.source { + VerifiedAccessSource::Individual => "user", + VerifiedAccessSource::Organization => "current_account", + }; + json!({ "level": grant.level, "source": source }) + }) + .collect::>(); + Some(json!({ + "schemaVersion": 1, + "status": status, + "grants": grants, + "stale": false + })) + } +} + +#[cfg(test)] +#[path = "trusted_access_tests.rs"] +mod tests; diff --git a/codex-rs/codex-mcp/src/trusted_access_tests.rs b/codex-rs/codex-mcp/src/trusted_access_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..63bc1d9d4a6b4b5bf74718b8395a8d2a0e041d2b --- /dev/null +++ b/codex-rs/codex-mcp/src/trusted_access_tests.rs @@ -0,0 +1,658 @@ +use std::collections::VecDeque; +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; + +use codex_exec_server::ByteChunk; +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpHeader; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +use codex_login::AuthCredentialsStoreMode; +use codex_login::AuthHeaders; +use codex_login::AuthKeyringBackendKind; +use codex_login::AuthManager; +use codex_login::CodexAuth; +use codex_login::ExternalAuth; +use codex_login::ExternalAuthFuture; +use codex_login::ExternalAuthRefreshContext; +use futures::FutureExt; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; +use serde_json::Value; +use serde_json::json; +use tokio::sync::Notify; + +use super::MAX_VERIFIED_ACCESS_RESPONSE_BYTES; +use super::TrustedAccessContext; + +struct RecordingHttpClient { + requests: Mutex>, + status: u16, + response: Vec, + response_chunks: Option>>, + response_gate: Option<(Notify, Notify)>, +} + +impl RecordingHttpClient { + fn new(status: u16, response: Value) -> Self { + Self { + requests: Mutex::new(Vec::new()), + status, + response: serde_json::to_vec(&response).expect("serialize response"), + response_chunks: None, + response_gate: None, + } + } +} + +impl HttpClient for RecordingHttpClient { + fn http_request( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + self.requests.lock().expect("record request").push(params); + let response = HttpRequestResponse { + status: self.status, + headers: Vec::new(), + body: ByteChunk(self.response.clone()), + }; + async move { + if let Some((requested, release)) = &self.response_gate { + requested.notify_one(); + release.notified().await; + } + Ok(response) + } + .boxed() + } + + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + self.requests.lock().expect("record request").push(params); + let response = HttpRequestResponse { + status: self.status, + headers: Vec::new(), + body: ByteChunk(Vec::new()), + }; + let chunks = self + .response_chunks + .clone() + .unwrap_or_else(|| vec![self.response.clone()]); + async move { + if let Some((requested, release)) = &self.response_gate { + requested.notify_one(); + release.notified().await; + } + Ok((response, HttpResponseBodyStream::from_chunks(chunks))) + } + .boxed() + } +} + +struct StaticExternalAuth(CodexAuth); + +impl ExternalAuth for StaticExternalAuth { + fn resolve(&self) -> ExternalAuthFuture<'_, CodexAuth> { + Box::pin(async { Ok(self.0.clone()) }) + } + + fn refresh(&self, _context: ExternalAuthRefreshContext) -> ExternalAuthFuture<'_, CodexAuth> { + self.resolve() + } +} + +struct SequencedExternalAuth(Mutex>); + +impl ExternalAuth for SequencedExternalAuth { + fn resolve(&self) -> ExternalAuthFuture<'_, CodexAuth> { + Box::pin(async { + let mut sequence = self.0.lock().expect("auth sequence"); + Ok(if sequence.len() > 1 { + sequence.pop_front().expect("next auth") + } else { + sequence.front().expect("last auth").clone() + }) + }) + } + + fn refresh(&self, _context: ExternalAuthRefreshContext) -> ExternalAuthFuture<'_, CodexAuth> { + self.resolve() + } +} + +fn header_auth() -> CodexAuth { + CodexAuth::Headers(AuthHeaders::new( + [ + ( + "authorization".parse().unwrap(), + "Bearer synthetic-pat".parse().unwrap(), + ), + ( + "chatgpt-account-id".parse().unwrap(), + "account-a".parse().unwrap(), + ), + ] + .into_iter() + .collect(), + )) +} + +fn chatgpt_auth(account_id: &str) -> CodexAuth { + CodexAuth::from_external_chatgpt_tokens( + "header.e30.same", + account_id, + /*chatgpt_plan_type*/ None, + ) + .expect("test auth") +} + +fn fedramp_chatgpt_auth(account_id: &str) -> CodexAuth { + CodexAuth::from_external_chatgpt_tokens( + "header.eyJodHRwczovL2FwaS5vcGVuYWkuY29tL2F1dGgiOnsiY2hhdGdwdF9hY2NvdW50X2lzX2ZlZHJhbXAiOnRydWV9fQ.same", + account_id, + /*chatgpt_plan_type*/ None, + ) + .expect("FedRAMP test auth") +} + +fn context(auth: CodexAuth, client: Arc) -> TrustedAccessContext { + TrustedAccessContext::new( + auth.clone(), + AuthManager::from_auth_for_testing(auth), + "https://chatgpt.com/backend-api".to_string(), + client, + ) +} + +fn cyber_response(state: &str, grants: Value) -> Value { + json!({ "programs": [{ "program": "cyber", "state": state, "grants": grants }] }) +} + +fn expected_metadata(status: &str, grants: Value) -> Value { + json!({ "openai/entitlementContext": { + "schemaVersion": 1, + "entitlements": { "cyber_trusted_access": { + "schemaVersion": 1, + "status": status, + "grants": grants, + "stale": false + } } + } }) +} + +#[tokio::test] +async fn delivers_account_bound_verified_access_to_the_calling_plugin() { + let client = Arc::new(RecordingHttpClient::new( + /*status*/ 200, + cyber_response( + "active", + json!([ + { "level": "tac2", "source": "individual" }, + { "level": "government", "source": "organization" } + ]), + ), + )); + let mut context = context( + CodexAuth::create_dummy_chatgpt_auth_for_testing(), + client.clone(), + ); + context.chatgpt_base_url.push('/'); + let metadata = context + .add_context(Some(json!({ "threadId": "thread-1" }))) + .await; + let mut expected = expected_metadata( + "granted", + json!([ + { "level": "tac2", "source": "user" }, + { "level": "government", "source": "current_account" } + ]), + ); + expected["threadId"] = json!("thread-1"); + assert_eq!(metadata, Some(expected)); + + let mut requests = client.requests.lock().expect("inspect recorded requests"); + for request in requests.iter_mut() { + request + .headers + .sort_by(|left, right| left.name.cmp(&right.name)); + } + assert_eq!( + *requests, + vec![HttpRequestParams { + method: "GET".to_string(), + url: "https://chatgpt.com/backend-api/accounts/verified_access".to_string(), + headers: vec![ + HttpHeader { + name: "authorization".to_string(), + value: "Bearer Access Token".to_string(), + value_env_var: None, + }, + HttpHeader { + name: "chatgpt-account-id".to_string(), + value: "account_id".to_string(), + value_env_var: None, + }, + ], + body: None, + timeout_ms: Some(2_500), + redirect_policy: HttpRedirectPolicy::Stop, + request_id: "trusted-access-status".to_string(), + stream_response: true, + }] + ); +} + +#[tokio::test] +async fn rejects_auth_without_a_nonempty_account_id() -> anyhow::Result<()> { + for account_id in [None, Some(""), Some(" ")] { + let home = tempfile::tempdir()?; + std::fs::write( + home.path().join("auth.json"), + serde_json::to_vec(&json!({ + "auth_mode": "chatgpt", + "tokens": { + "id_token": "header.eyJodHRwczovL2FwaS5vcGVuYWkuY29tL2F1dGgiOnsiY2hhdGdwdF91c2VyX2lkIjoidXNlci1hIn19.signature", + "access_token": "synthetic-access-token", + "refresh_token": "synthetic-refresh-token", + "account_id": account_id + }, + "last_refresh": "2099-01-01T00:00:00Z" + }))?, + )?; + let auth = CodexAuth::from_auth_storage( + home.path(), + AuthCredentialsStoreMode::File, + /*chatgpt_base_url*/ None, + AuthKeyringBackendKind::default(), + &codex_login::test_support::transport_default_auth_route_config(), + ) + .await? + .expect("managed ChatGPT auth"); + let client = Arc::new(RecordingHttpClient::new( + /*status*/ 200, + cyber_response( + "active", + json!([{ "level": "tac1", "source": "individual" }]), + ), + )); + assert_eq!( + context(auth, client.clone()) + .add_context(/*meta*/ None) + .await, + Some(expected_metadata("unknown", json!([]))), + "account_id={account_id:?}" + ); + assert!(client.requests.lock().expect("inspect requests").is_empty()); + } + Ok(()) +} + +#[tokio::test] +async fn rejects_api_key_authentication_without_sending_credentials() { + let client = Arc::new(RecordingHttpClient::new( + /*status*/ 200, + json!({ "programs": [] }), + )); + let context = context(CodexAuth::from_api_key("test-api-key"), client.clone()); + + let metadata = context + .add_context(Some(json!({ + "threadId": "thread-1", + "openai/entitlementContext": { "forged": true } + }))) + .await; + + let mut expected = expected_metadata("unknown", json!([])); + expected["threadId"] = json!("thread-1"); + assert_eq!(metadata, Some(expected)); + assert!(client.requests.lock().expect("inspect requests").is_empty()); +} + +#[tokio::test] +async fn rejects_initial_unsupported_auth_after_switch_to_chatgpt() { + let client = Arc::new(RecordingHttpClient::new( + /*status*/ 200, + cyber_response( + "active", + json!([{ "level": "tac1", "source": "individual" }]), + ), + )); + let mut context = context(header_auth(), client.clone()); + context.auth_manager = AuthManager::from_auth_for_testing(chatgpt_auth("account-a")); + + assert_eq!( + context.add_context(/*meta*/ None).await, + Some(expected_metadata("unknown", json!([]))) + ); + assert!(client.requests.lock().expect("inspect requests").is_empty()); +} + +#[tokio::test] +async fn maps_verified_access_states_and_rejects_invalid_provider_responses() { + let cases = [ + ( + 200, + json!({ "programs": [ + { "program": "future", "state": "pending", "grants": [{ "level": "premium", "source": "subscription" }] }, + { "program": "cyber", "state": "active", "grants": [{ "level": "tac1", "source": "individual" }] } + ] }), + expected_metadata("granted", json!([{ "level": "tac1", "source": "user" }])), + ), + ( + 200, + cyber_response("inactive", json!([])), + expected_metadata("not_granted", json!([])), + ), + ( + 200, + cyber_response("unavailable", json!([])), + expected_metadata("unknown", json!([])), + ), + ( + 200, + json!({ "programs": [{ "program": "other", "state": "active", "grants": [] }] }), + expected_metadata("unknown", json!([])), + ), + ( + 200, + cyber_response( + "active", + json!([{ "level": "admin", "source": "individual" }]), + ), + expected_metadata("unknown", json!([])), + ), + ( + 200, + cyber_response("active", json!([])), + expected_metadata("unknown", json!([])), + ), + ( + 403, + json!({ "error": "forbidden" }), + expected_metadata("unknown", json!([])), + ), + ]; + + for (response_status, response, expected) in cases { + let response_description = response.to_string(); + let context = context( + CodexAuth::create_dummy_chatgpt_auth_for_testing(), + Arc::new(RecordingHttpClient::new(response_status, response)), + ); + let metadata = context.add_context(/*meta*/ None).await.expect("metadata"); + assert_eq!( + metadata, expected, + "HTTP {response_status}: {response_description}" + ); + } +} + +#[tokio::test] +async fn replaces_untrusted_entitlements_on_malformed_response() { + let mut client = RecordingHttpClient::new(/*status*/ 200, Value::Null); + client.response = b"{\"programs\":[".to_vec(); + let context = context(chatgpt_auth("account-a"), Arc::new(client)); + let mut expected = expected_metadata("unknown", json!([])); + expected["threadId"] = json!("thread-1"); + + assert_eq!( + context + .add_context(Some(json!({ + "threadId": "thread-1", + "openai/entitlementContext": { "forged": true } + }))) + .await, + Some(expected) + ); +} + +#[tokio::test] +async fn rejects_verified_access_response_larger_than_one_mebibyte() { + let valid_response = serde_json::to_vec(&cyber_response( + "active", + json!([{ "level": "tac1", "source": "individual" }]), + )) + .expect("serialize response"); + let mut client = RecordingHttpClient::new(/*status*/ 200, Value::Null); + client.response_chunks = Some(vec![ + vec![b' '; MAX_VERIFIED_ACCESS_RESPONSE_BYTES], + valid_response, + ]); + client.response = client + .response_chunks + .as_ref() + .expect("response chunks") + .concat(); + + assert_eq!( + context(chatgpt_auth("account-a"), Arc::new(client)) + .add_context(/*meta*/ None) + .await, + Some(expected_metadata("unknown", json!([]))) + ); +} + +#[tokio::test(start_paused = true)] +async fn returns_unknown_at_the_lookup_deadline() { + let mut client = RecordingHttpClient::new(/*status*/ 200, Value::Null); + client.response_gate = Some((Notify::new(), Notify::new())); + let client = Arc::new(client); + let context = context(chatgpt_auth("account-a"), client.clone()); + let lookup = context.add_context(/*meta*/ None); + tokio::pin!(lookup); + + assert!(futures::poll!(lookup.as_mut()).is_pending()); + assert_eq!(client.requests.lock().expect("recorded requests").len(), 1); + tokio::time::advance(Duration::from_millis(2_499)).await; + assert!(futures::poll!(lookup.as_mut()).is_pending()); + tokio::time::advance(Duration::from_millis(1)).await; + assert_eq!( + lookup + .now_or_never() + .expect("lookup must finish at the deadline"), + Some(expected_metadata("unknown", json!([]))) + ); +} + +#[tokio::test] +async fn rejects_duplicate_cyber_programs() { + let active = json!({ + "program": "cyber", "state": "active", + "grants": [{ "level": "tac1", "source": "individual" }] + }); + for other in [ + active.clone(), + json!({ "program": "cyber", "state": "inactive", "grants": [] }), + json!({ "program": "cyber" }), + ] { + for programs in [json!([active, other]), json!([other, active])] { + let context = context( + CodexAuth::create_dummy_chatgpt_auth_for_testing(), + Arc::new(RecordingHttpClient::new( + /*status*/ 200, + json!({ "programs": programs }), + )), + ); + assert_eq!( + context.add_context(/*meta*/ None).await, + Some(expected_metadata("unknown", json!([]))), + "duplicate programs: {programs}" + ); + } + } +} + +#[tokio::test] +async fn rejects_identity_changes_before_sending_credentials() { + for (description, initial_auth, selected_auth) in [ + ( + "account switch", + chatgpt_auth("account-a"), + chatgpt_auth("account-b"), + ), + ( + "standard to FedRAMP", + chatgpt_auth("account-a"), + fedramp_chatgpt_auth("account-a"), + ), + ( + "FedRAMP to standard", + fedramp_chatgpt_auth("account-a"), + chatgpt_auth("account-a"), + ), + ] { + let client = Arc::new(RecordingHttpClient::new( + /*status*/ 200, + cyber_response( + "active", + json!([{ "level": "tac1", "source": "individual" }]), + ), + )); + let mut context = context(initial_auth, client.clone()); + context.auth_manager = AuthManager::from_auth_for_testing(selected_auth); + assert_eq!( + context.add_context(/*meta*/ None).await, + Some(expected_metadata("unknown", json!([]))), + "identity change before the request: {description}" + ); + assert!( + client.requests.lock().expect("inspect requests").is_empty(), + "identity change before the request: {description}" + ); + } +} + +#[tokio::test] +async fn uses_the_checked_auth_snapshot_for_request_headers() -> anyhow::Result<()> { + let client = Arc::new(RecordingHttpClient::new( + /*status*/ 200, + cyber_response("inactive", json!([])), + )); + let context = context(chatgpt_auth("account-a"), client.clone()); + let refreshed = CodexAuth::from_external_chatgpt_tokens( + "header.e30.refreshed", + "account-a", + /*chatgpt_plan_type*/ None, + )?; + context + .auth_manager + .set_external_auth(Arc::new(SequencedExternalAuth(Mutex::new(VecDeque::from( + [context.auth.clone(), refreshed, header_auth()], + ))))) + .await?; + + assert_eq!( + context.add_context(/*meta*/ None).await, + Some(expected_metadata("not_granted", json!([]))) + ); + let requests = client.requests.lock().expect("recorded requests"); + assert_eq!(requests.len(), 1); + let authorization = requests[0] + .headers + .iter() + .find(|header| header.name == "authorization"); + assert_eq!( + authorization, + Some(&HttpHeader { + name: "authorization".to_string(), + value: "Bearer header.e30.refreshed".to_string(), + value_env_var: None, + }) + ); + Ok(()) +} + +#[tokio::test] +async fn rejects_identity_changes_while_request_is_in_flight() -> anyhow::Result<()> { + let personal = chatgpt_auth("account-a"); + let workspace = + CodexAuth::from_external_chatgpt_tokens("header.e30.same", "account-a", Some("team"))?; + for (description, initial_auth, selected_auth) in [ + ( + "account switch", + personal.clone(), + Some(chatgpt_auth("account-b")), + ), + ("logout", personal.clone(), None), + ( + "personal to workspace", + personal.clone(), + Some(workspace.clone()), + ), + ("workspace to personal", workspace, Some(personal)), + ( + "auth mode change", + chatgpt_auth("account-a"), + Some(header_auth()), + ), + ( + "standard to FedRAMP", + chatgpt_auth("account-a"), + Some(fedramp_chatgpt_auth("account-a")), + ), + ( + "FedRAMP to standard", + fedramp_chatgpt_auth("account-a"), + Some(chatgpt_auth("account-a")), + ), + ] { + let mut client = RecordingHttpClient::new( + /*status*/ 200, + cyber_response( + "active", + json!([{ "level": "tac1", "source": "individual" }]), + ), + ); + client.response_gate = Some((Notify::new(), Notify::new())); + let client = Arc::new(client); + let request_is_fedramp = initial_auth.is_fedramp_account(); + let context = context(initial_auth, client.clone()); + let auth_manager = &context.auth_manager; + auth_manager + .set_external_auth(Arc::new(StaticExternalAuth(context.auth.clone()))) + .await?; + + let (metadata, auth_change) = tokio::join!(context.add_context(/*meta*/ None), async { + let (requested, release) = client.response_gate.as_ref().expect("gated response"); + tokio::time::timeout(Duration::from_secs(5), requested.notified()).await?; + { + let requests = client.requests.lock().expect("inspect request"); + assert_eq!(requests.len(), 1); + assert!(requests[0].headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("chatgpt-account-id") + && header.value == "account-a" + })); + assert_eq!( + requests[0] + .headers + .iter() + .any(|header| header.name.eq_ignore_ascii_case("x-openai-fedramp")), + request_is_fedramp + ); + } + if let Some(auth) = selected_auth { + auth_manager + .set_external_auth(Arc::new(StaticExternalAuth(auth))) + .await?; + } else { + auth_manager.clear_external_auth(); + } + release.notify_one(); + Ok::<(), anyhow::Error>(()) + }); + auth_change?; + + assert_eq!( + metadata, + Some(expected_metadata("unknown", json!([]))), + "auth change after the request: {description}" + ); + } + Ok(()) +} diff --git a/codex-rs/codex-mcp/src/user_verification_elicitation.rs b/codex-rs/codex-mcp/src/user_verification_elicitation.rs new file mode 100644 index 0000000000000000000000000000000000000000..5b2890400c25f5e011b41d29ed0e82fb4b28834f --- /dev/null +++ b/codex-rs/codex-mcp/src/user_verification_elicitation.rs @@ -0,0 +1,65 @@ +//! Routes device verification directly to the app, outside automated approval policy. + +use super::*; + +pub(super) async fn route( + router: ElicitationRequestRouter, + events: Option>, + authority: Arc>>, + server_name: String, + request: ElicitationRequest, +) -> Result { + let authority = authority + .lock() + .ok() + .and_then(|authority| authority.clone()); + let plugin_service = authority.as_ref().is_some_and(|authority| { + authority + .config + .mcp_server_catalog + .server(&server_name) + .is_some_and(|server| { + server + .source() + .is_host_owned_apps(&server_name, server.config()) + }) + }); + let Some(events) = events.filter(|_| plugin_service && !router.auto_deny()) else { + return Ok(ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None, + }); + }; + let id = format!( + "codex-mcp-elicitation-{}", + NEXT_ELICITATION_REQUEST_ID.fetch_add(1, Ordering::Relaxed) + ); + let key = (server_name.clone(), RequestId::String(id.clone().into())); + let (response, receiver) = oneshot::channel(); + router + .requests + .lock() + .map_err(|_| anyhow!("elicitation request router unavailable"))? + .insert(key.clone(), response); + let _pending = PendingElicitationRequest { router, key }; + let _active = authority + .as_ref() + .and_then(|authority| authority.lifecycle.as_ref()) + .map(ElicitationLifecycle::start); + events + .send(Event { + id: "mcp_elicitation_request".to_string(), + msg: EventMsg::ElicitationRequest(ElicitationRequestEvent { + turn_id: None, + server_name, + id: ProtocolRequestId::String(id), + request, + }), + }) + .await + .context("failed to deliver user-verification request")?; + receiver + .await + .context("user-verification response channel closed") +} diff --git a/codex-rs/connectors/src/accessible.rs b/codex-rs/connectors/src/accessible.rs new file mode 100644 index 0000000000000000000000000000000000000000..c42752f6d0cc9324281fd2b047144c1c5b2563eb --- /dev/null +++ b/codex-rs/connectors/src/accessible.rs @@ -0,0 +1,78 @@ +use std::collections::BTreeSet; +use std::collections::HashMap; + +use crate::AppInfo; +use crate::metadata::connector_install_url; +use crate::normalize_connector_value; + +pub struct AccessibleConnectorTool { + pub connector_id: String, + pub connector_name: Option, + pub connector_description: Option, + pub plugin_display_names: Vec, +} + +pub fn collect_accessible_connectors(tools: I) -> Vec +where + I: IntoIterator, +{ + let mut connectors: HashMap)> = HashMap::new(); + for tool in tools { + let connector_id = tool.connector_id; + let connector_name = normalize_connector_value(tool.connector_name.as_deref()) + .unwrap_or_else(|| connector_id.clone()); + let connector_description = + normalize_connector_value(tool.connector_description.as_deref()); + if let Some((existing, existing_plugin_display_names)) = connectors.get_mut(&connector_id) { + if existing.name == connector_id && connector_name != connector_id { + existing.name = connector_name; + } + if existing.description.is_none() && connector_description.is_some() { + existing.description = connector_description; + } + existing_plugin_display_names.extend(tool.plugin_display_names); + } else { + connectors.insert( + connector_id.clone(), + ( + AppInfo { + id: connector_id.clone(), + name: connector_name, + description: connector_description, + logo_url: None, + logo_url_dark: None, + icon_assets: None, + icon_dark_assets: None, + distribution_channel: None, + branding: None, + app_metadata: None, + labels: None, + install_url: None, + is_accessible: true, + is_enabled: true, + plugin_display_names: Vec::new(), + }, + tool.plugin_display_names + .into_iter() + .collect::>(), + ), + ); + } + } + let mut accessible: Vec = connectors + .into_values() + .map(|(mut connector, plugin_display_names)| { + connector.plugin_display_names = plugin_display_names.into_iter().collect(); + connector.install_url = Some(connector_install_url(&connector.name, &connector.id)); + connector + }) + .collect(); + accessible.sort_by(|left, right| { + right + .is_accessible + .cmp(&left.is_accessible) + .then_with(|| left.name.cmp(&right.name)) + .then_with(|| left.id.cmp(&right.id)) + }); + accessible +} diff --git a/codex-rs/connectors/src/app_info.rs b/codex-rs/connectors/src/app_info.rs new file mode 100644 index 0000000000000000000000000000000000000000..cffe20f4672655ec8328ba73a4a3c351f1b73de0 --- /dev/null +++ b/codex-rs/connectors/src/app_info.rs @@ -0,0 +1,111 @@ +//! Connector-domain app metadata used by directory discovery, caching, and tool selection. +//! +//! The Serde implementations decode connector-directory response metadata and persist normalized +//! app information in the connector-directory disk cache. They do not define the app-server wire +//! format; `codex-app-server-protocol` owns separate API types for that boundary. + +use serde::Deserialize; +use serde::Serialize; +use std::collections::HashMap; + +/// Branding supplied by the connector directory for an app. +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)] +#[serde(rename_all = "camelCase")] +pub struct AppBranding { + pub category: Option, + pub developer: Option, + pub website: Option, + pub privacy_policy: Option, + pub terms_of_service: Option, + pub is_discoverable_app: bool, +} + +/// Review state supplied by the connector directory for an app. +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)] +#[serde(rename_all = "camelCase")] +pub struct AppReview { + pub status: String, +} + +/// Screenshot metadata supplied by the connector directory for an app. +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)] +#[serde(rename_all = "camelCase")] +pub struct AppScreenshot { + pub url: Option, + #[serde(alias = "file_id")] + pub file_id: Option, + #[serde(alias = "user_prompt")] + pub user_prompt: String, +} + +/// Extended metadata supplied by the connector directory for an app. +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)] +#[serde(rename_all = "camelCase")] +pub struct AppMetadata { + pub review: Option, + pub categories: Option>, + pub sub_categories: Option>, + pub seo_description: Option, + pub screenshots: Option>, + pub developer: Option, + pub version: Option, + pub version_id: Option, + pub version_notes: Option, + pub first_party_requires_install: Option, + pub show_in_composer_when_unlinked: Option, +} + +/// Connector metadata used by connector discovery, caching, and tool selection. +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)] +#[serde(rename_all = "camelCase")] +pub struct AppInfo { + pub id: String, + pub name: String, + pub description: Option, + pub logo_url: Option, + pub logo_url_dark: Option, + pub icon_assets: Option>, + pub icon_dark_assets: Option>, + pub distribution_channel: Option, + pub branding: Option, + pub app_metadata: Option, + pub labels: Option>, + pub install_url: Option, + #[serde(default)] + pub is_accessible: bool, + #[serde(default = "default_enabled")] + pub is_enabled: bool, + #[serde(default)] + pub plugin_display_names: Vec, +} + +impl AppInfo { + pub fn category(&self) -> Option { + self.branding + .as_ref() + .and_then(|branding| non_empty_category(branding.category.as_deref())) + .or_else(|| { + self.app_metadata + .as_ref() + .and_then(|metadata| metadata.categories.as_ref()) + .and_then(|categories| { + categories + .iter() + .find_map(|category| non_empty_category(Some(category.as_str()))) + }) + }) + } +} + +const fn default_enabled() -> bool { + true +} + +fn non_empty_category(category: Option<&str>) -> Option { + let category = category?.trim(); + if category.is_empty() { + None + } else { + Some(category.to_string()) + } +} diff --git a/codex-rs/connectors/src/app_tool_policy.rs b/codex-rs/connectors/src/app_tool_policy.rs new file mode 100644 index 0000000000000000000000000000000000000000..a32e7903e18722f10ce00a6a40946e29fa791a89 --- /dev/null +++ b/codex-rs/connectors/src/app_tool_policy.rs @@ -0,0 +1,245 @@ +use codex_config::AppsRequirementsToml; +use codex_config::ConfigLayerStack; +use codex_config::types::AppToolApproval; +use codex_config::types::AppsConfigToml; +use serde::Deserialize; + +use crate::AppInfo; + +/// The effective enablement and approval policy for one app tool. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct AppToolPolicy { + pub enabled: bool, + pub approval: AppToolApproval, +} + +impl Default for AppToolPolicy { + fn default() -> Self { + Self { + enabled: true, + approval: AppToolApproval::Auto, + } + } +} + +/// App and account metadata used to evaluate one tool. +#[derive(Debug, Clone, Copy)] +pub struct AppToolPolicyInput<'a> { + pub connector_id: Option<&'a str>, + pub link_id: Option<&'a str>, + pub tool_name: &'a str, + pub tool_title: Option<&'a str>, + pub destructive_hint: Option, + pub open_world_hint: Option, +} + +/// Resolves app tool policy against one immutable config snapshot. +/// +/// Callers should construct one evaluator and reuse it for every tool in the +/// same exposure build so config layers are merged and decoded only once. +pub struct AppToolPolicyEvaluator<'a> { + apps_config: Option, + requirements_apps_config: Option<&'a AppsRequirementsToml>, +} + +impl<'a> AppToolPolicyEvaluator<'a> { + pub fn new(config_layer_stack: &'a ConfigLayerStack) -> Self { + let apps_config = apps_config_from_layer_stack(config_layer_stack); + let requirements_apps_config = config_layer_stack.requirements_toml().apps.as_ref(); + Self::from_parts(apps_config, requirements_apps_config) + } + + pub fn policy(&self, input: AppToolPolicyInput<'_>) -> AppToolPolicy { + let managed_approval = managed_app_tool_approval( + self.requirements_apps_config, + input.connector_id, + input.tool_name, + ); + app_tool_policy_from_apps_config(self.apps_config.as_ref(), input, managed_approval) + } + + /// Returns the effective local and managed enablement for one connector. + pub fn app_enabled(&self, connector_id: &str) -> bool { + self.apps_config + .as_ref() + .map(|apps_config| app_is_enabled(apps_config, Some(connector_id))) + .unwrap_or(true) + } + + /// Applies app policy without overriding source state for unconfigured apps. + pub fn apply_app_enabled_state(&self, mut apps: Vec) -> Vec { + let Some(apps_config) = self.apps_config.as_ref() else { + return apps; + }; + + for app in &mut apps { + if apps_config.default.is_some() || apps_config.apps.contains_key(app.id.as_str()) { + app.is_enabled = self.app_enabled(app.id.as_str()); + } + } + + apps + } + + fn from_parts( + apps_config: Option, + requirements_apps_config: Option<&'a AppsRequirementsToml>, + ) -> Self { + Self { + apps_config: effective_apps_config(apps_config, requirements_apps_config), + requirements_apps_config, + } + } +} + +/// Reads the merged, unmanaged Apps configuration from a config-layer stack. +pub fn apps_config_from_layer_stack( + config_layer_stack: &ConfigLayerStack, +) -> Option { + config_layer_stack + .effective_config() + .as_table() + .and_then(|table| table.get("apps")) + .cloned() + .and_then(|value| AppsConfigToml::deserialize(value).ok()) +} + +pub fn app_is_enabled(apps_config: &AppsConfigToml, connector_id: Option<&str>) -> bool { + let default_enabled = apps_config + .default + .as_ref() + .map(|defaults| defaults.enabled) + .unwrap_or(true); + + connector_id + .and_then(|connector_id| apps_config.apps.get(connector_id)) + .map(|app| app.enabled) + .unwrap_or(default_enabled) +} + +fn effective_apps_config( + apps_config: Option, + requirements_apps_config: Option<&AppsRequirementsToml>, +) -> Option { + let had_apps_config = apps_config.is_some(); + let mut apps_config = apps_config.unwrap_or_default(); + apply_requirements_apps_constraints(&mut apps_config, requirements_apps_config); + if had_apps_config || apps_config.default.is_some() || !apps_config.apps.is_empty() { + Some(apps_config) + } else { + None + } +} + +fn apply_requirements_apps_constraints( + apps_config: &mut AppsConfigToml, + requirements_apps_config: Option<&AppsRequirementsToml>, +) { + let Some(requirements_apps_config) = requirements_apps_config else { + return; + }; + + for (app_id, requirement) in &requirements_apps_config.apps { + if requirement.enabled == Some(false) { + let app = apps_config.apps.entry(app_id.clone()).or_default(); + app.enabled = false; + } + } +} + +fn managed_app_tool_approval( + requirements_apps_config: Option<&AppsRequirementsToml>, + connector_id: Option<&str>, + tool_name: &str, +) -> Option { + let connector_id = connector_id?; + requirements_apps_config? + .apps + .get(connector_id)? + .tools + .as_ref()? + .tools + .get(tool_name)? + .approval_mode +} + +fn app_tool_policy_from_apps_config( + apps_config: Option<&AppsConfigToml>, + input: AppToolPolicyInput<'_>, + managed_approval: Option, +) -> AppToolPolicy { + let Some(apps_config) = apps_config else { + return AppToolPolicy { + approval: managed_approval.unwrap_or(AppToolApproval::Auto), + ..Default::default() + }; + }; + + let app = input + .connector_id + .and_then(|connector_id| apps_config.apps.get(connector_id)); + let tools = app.and_then(|app| app.tools.as_ref()); + let tool_config = tools.and_then(|tools| { + tools + .tools + .get(input.tool_name) + .or_else(|| input.tool_title.and_then(|title| tools.tools.get(title))) + }); + let approval = managed_approval + .or_else(|| tool_config.and_then(|tool| tool.approval_mode)) + .or_else(|| { + input + .link_id + .and_then(|link_id| app?.links.as_ref()?.links.get(link_id)) + .and_then(|link| link.default_tools_approval_mode) + }) + .or_else(|| app.and_then(|app| app.default_tools_approval_mode)) + .or_else(|| { + input + .connector_id + .and(apps_config.default.as_ref()) + .and_then(|defaults| defaults.default_tools_approval_mode) + }) + .unwrap_or(AppToolApproval::Auto); + + if !app_is_enabled(apps_config, input.connector_id) { + return AppToolPolicy { + enabled: false, + approval, + }; + } + + if let Some(enabled) = tool_config.and_then(|tool| tool.enabled) { + return AppToolPolicy { enabled, approval }; + } + + if let Some(enabled) = app.and_then(|app| app.default_tools_enabled) { + return AppToolPolicy { enabled, approval }; + } + + let app_defaults = apps_config.default.as_ref(); + let destructive_enabled = app + .and_then(|app| app.destructive_enabled) + .unwrap_or_else(|| { + app_defaults + .map(|defaults| defaults.destructive_enabled) + .unwrap_or(true) + }); + let open_world_enabled = app + .and_then(|app| app.open_world_enabled) + .unwrap_or_else(|| { + app_defaults + .map(|defaults| defaults.open_world_enabled) + .unwrap_or(true) + }); + let destructive_hint = input.destructive_hint.unwrap_or(true); + let open_world_hint = input.open_world_hint.unwrap_or(true); + let enabled = + (destructive_enabled || !destructive_hint) && (open_world_enabled || !open_world_hint); + + AppToolPolicy { enabled, approval } +} + +#[cfg(test)] +#[path = "app_tool_policy_tests.rs"] +mod tests; diff --git a/codex-rs/connectors/src/app_tool_policy_tests.rs b/codex-rs/connectors/src/app_tool_policy_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..e255de49726c29cda6ed9a6707441ca5cd9054fb --- /dev/null +++ b/codex-rs/connectors/src/app_tool_policy_tests.rs @@ -0,0 +1,962 @@ +use std::collections::BTreeMap; +use std::collections::HashMap; + +use codex_config::AbsolutePathBuf; +use codex_config::AppRequirementToml; +use codex_config::AppToolRequirementToml; +use codex_config::AppToolsRequirementsToml; +use codex_config::AppsRequirementsToml; +use codex_config::CONFIG_TOML_FILE; +use codex_config::ConfigLayerStack; +use codex_config::ConfigRequirements; +use codex_config::ConfigRequirementsToml; +use codex_config::TomlValue; +use codex_config::types::AppConfig; +use codex_config::types::AppToolApproval; +use codex_config::types::AppToolConfig; +use codex_config::types::AppToolsConfig; +use codex_config::types::AppsConfigToml; +use codex_config::types::AppsDefaultConfig; +use pretty_assertions::assert_eq; + +use super::*; + +#[test] +fn evaluator_reuses_one_snapshot_across_tools() { + let apps_config = AppsConfigToml { + default: None, + apps: HashMap::from([( + "calendar".to_string(), + AppConfig { + enabled: true, + default_tools_enabled: Some(false), + tools: Some(AppToolsConfig { + tools: HashMap::from([( + "events/create".to_string(), + AppToolConfig { + enabled: Some(true), + approval_mode: Some(AppToolApproval::Prompt), + }, + )]), + }), + ..Default::default() + }, + )]), + }; + let requirements = AppsRequirementsToml { + apps: BTreeMap::from([( + "calendar".to_string(), + AppRequirementToml { + enabled: None, + tools: Some(AppToolsRequirementsToml { + tools: BTreeMap::from([( + "events/create".to_string(), + AppToolRequirementToml { + approval_mode: Some(AppToolApproval::Approve), + ..Default::default() + }, + )]), + }), + }, + )]), + }; + let evaluator = AppToolPolicyEvaluator::from_parts(Some(apps_config), Some(&requirements)); + + assert_eq!( + [ + evaluator.policy(input("events/create", /*tool_title*/ None)), + evaluator.policy(input("events/list", /*tool_title*/ None)), + evaluator.policy(input("calendar_events/create", Some("events/create"))), + ], + [ + AppToolPolicy { + enabled: true, + approval: AppToolApproval::Approve, + }, + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Auto, + }, + AppToolPolicy { + enabled: true, + approval: AppToolApproval::Prompt, + }, + ] + ); +} + +#[test] +fn evaluator_uses_global_defaults_for_destructive_hints() { + let apps_config = AppsConfigToml { + default: Some(defaults( + /*enabled*/ true, /*destructive_enabled*/ false, + /*open_world_enabled*/ true, + )), + apps: HashMap::new(), + }; + + assert_eq!( + policy_from_apps_config( + Some(&apps_config), + Some("calendar"), + "events/create", + /*tool_title*/ None, + Some(true), + /*open_world_hint*/ None, + /*managed_approval*/ None, + ), + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Auto, + } + ); +} + +#[test] +fn evaluator_defaults_missing_destructive_hint_to_true() { + let apps_config = AppsConfigToml { + default: Some(defaults( + /*enabled*/ true, /*destructive_enabled*/ false, + /*open_world_enabled*/ true, + )), + apps: HashMap::new(), + }; + + assert_eq!( + policy_from_apps_config( + Some(&apps_config), + Some("calendar"), + "events/create", + /*tool_title*/ None, + /*destructive_hint*/ None, + Some(false), + /*managed_approval*/ None, + ), + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Auto, + } + ); +} + +#[test] +fn evaluator_defaults_missing_open_world_hint_to_true() { + let apps_config = AppsConfigToml { + default: Some(defaults( + /*enabled*/ true, /*destructive_enabled*/ true, + /*open_world_enabled*/ false, + )), + apps: HashMap::new(), + }; + + assert_eq!( + policy_from_apps_config( + Some(&apps_config), + Some("calendar"), + "events/create", + /*tool_title*/ None, + Some(false), + /*open_world_hint*/ None, + /*managed_approval*/ None, + ), + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Auto, + } + ); +} + +#[test] +fn app_enablement_uses_defaults_and_per_app_overrides() { + let apps_config = AppsConfigToml { + default: Some(defaults( + /*enabled*/ false, /*destructive_enabled*/ true, + /*open_world_enabled*/ true, + )), + apps: HashMap::from([( + "calendar".to_string(), + AppConfig { + enabled: true, + ..Default::default() + }, + )]), + }; + + assert_eq!( + [ + app_is_enabled(&apps_config, Some("calendar")), + app_is_enabled(&apps_config, Some("drive")), + app_is_enabled(&apps_config, /*connector_id*/ None), + ], + [true, false, false] + ); + + let evaluator = AppToolPolicyEvaluator::from_parts( + Some(apps_config), + /*requirements_apps_config*/ None, + ); + assert_eq!( + evaluator.apply_app_enabled_state(vec![ + app("calendar", /*enabled*/ false), + app("drive", /*enabled*/ true), + ]), + vec![ + app("calendar", /*enabled*/ true), + app("drive", /*enabled*/ false), + ] + ); +} + +#[test] +fn app_enablement_preserves_source_state_and_honors_local_and_managed_overrides() { + let apps_config = AppsConfigToml { + default: None, + apps: HashMap::from([ + ( + "calendar".to_string(), + AppConfig { + enabled: true, + ..Default::default() + }, + ), + ( + "drive".to_string(), + AppConfig { + enabled: true, + ..Default::default() + }, + ), + ]), + }; + let requirements = app_enabled_requirement("drive", /*enabled*/ false); + let evaluator = AppToolPolicyEvaluator::from_parts(Some(apps_config), Some(&requirements)); + + assert_eq!( + evaluator.apply_app_enabled_state(vec![ + app("calendar", /*enabled*/ false), + app("drive", /*enabled*/ true), + app("slack", /*enabled*/ false), + app("gmail", /*enabled*/ true), + ]), + vec![ + app("calendar", /*enabled*/ true), + app("drive", /*enabled*/ false), + app("slack", /*enabled*/ false), + app("gmail", /*enabled*/ true), + ] + ); +} + +#[test] +fn managed_disable_overrides_enabled_app() { + let apps_config = AppsConfigToml { + default: None, + apps: HashMap::from([( + "connector_123123".to_string(), + AppConfig { + enabled: true, + ..Default::default() + }, + )]), + }; + let requirements = app_enabled_requirement("connector_123123", /*enabled*/ false); + + assert_eq!( + policy_from_config_parts( + Some(&apps_config), + Some(&requirements), + Some("connector_123123"), + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + ), + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Auto, + } + ); +} + +#[test] +fn managed_enable_does_not_override_disabled_app() { + let apps_config = AppsConfigToml { + default: None, + apps: HashMap::from([( + "connector_123123".to_string(), + AppConfig { + enabled: false, + ..Default::default() + }, + )]), + }; + let requirements = app_enabled_requirement("connector_123123", /*enabled*/ true); + + assert_eq!( + policy_from_config_parts( + Some(&apps_config), + Some(&requirements), + Some("connector_123123"), + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + ), + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Auto, + } + ); +} + +#[test] +fn managed_disable_applies_without_apps_config() { + let requirements = app_enabled_requirement("connector_123123", /*enabled*/ false); + + assert_eq!( + policy_from_config_parts( + /*apps_config*/ None, + Some(&requirements), + Some("connector_123123"), + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + ), + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Auto, + } + ); +} + +#[test] +fn evaluator_honors_default_app_enabled_false() { + let apps_config = AppsConfigToml { + default: Some(defaults( + /*enabled*/ false, /*destructive_enabled*/ true, + /*open_world_enabled*/ true, + )), + apps: HashMap::new(), + }; + + assert_eq!( + policy_from_apps_config( + Some(&apps_config), + Some("calendar"), + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + /*managed_approval*/ None, + ), + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Auto, + } + ); +} + +#[test] +fn evaluator_allows_per_app_enable_when_default_is_disabled() { + let apps_config = AppsConfigToml { + default: Some(defaults( + /*enabled*/ false, /*destructive_enabled*/ true, + /*open_world_enabled*/ true, + )), + apps: HashMap::from([( + "calendar".to_string(), + AppConfig { + enabled: true, + ..Default::default() + }, + )]), + }; + + assert_eq!( + policy_from_apps_config( + Some(&apps_config), + Some("calendar"), + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + /*managed_approval*/ None, + ), + AppToolPolicy::default() + ); +} + +#[test] +fn evaluator_uses_managed_approval_without_apps_config() { + assert_eq!( + policy_from_apps_config( + /*apps_config*/ None, + Some("calendar"), + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + Some(AppToolApproval::Approve), + ), + AppToolPolicy { + enabled: true, + approval: AppToolApproval::Approve, + } + ); +} + +#[test] +fn managed_approval_uses_raw_tool_name() { + let requirements = app_tool_requirements( + "connector_123123", + "calendar/list_events", + AppToolApproval::Approve, + ); + + assert_eq!( + [ + policy_from_config_parts( + /*apps_config*/ None, + Some(&requirements), + Some("connector_123123"), + "calendar/list_events", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + ), + policy_from_config_parts( + /*apps_config*/ None, + Some(&requirements), + Some("connector_123123"), + "calendar/create_event", + Some("calendar/list_events"), + /*destructive_hint*/ None, + /*open_world_hint*/ None, + ), + ], + [ + AppToolPolicy { + enabled: true, + approval: AppToolApproval::Approve, + }, + AppToolPolicy::default(), + ] + ); +} + +#[test] +fn managed_approval_overrides_user_tool_approval() { + let apps_config = AppsConfigToml { + default: None, + apps: HashMap::from([( + "connector_123123".to_string(), + AppConfig { + enabled: true, + tools: Some(AppToolsConfig { + tools: HashMap::from([( + "calendar/list_events".to_string(), + AppToolConfig { + enabled: None, + approval_mode: Some(AppToolApproval::Prompt), + }, + )]), + }), + ..Default::default() + }, + )]), + }; + let requirements = app_tool_requirements( + "connector_123123", + "calendar/list_events", + AppToolApproval::Approve, + ); + + assert_eq!( + policy_from_config_parts( + Some(&apps_config), + Some(&requirements), + Some("connector_123123"), + "calendar/list_events", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + ), + AppToolPolicy { + enabled: true, + approval: AppToolApproval::Approve, + } + ); +} + +#[test] +fn link_privacy_mode_overrides_app_default_and_preserves_tool_settings() { + let apps_config = serde_json::from_value(serde_json::json!({ + "calendar": { + "default_tools_approval_mode": "auto", + "tools": { "events/create": { "approval_mode": "writes" } }, + "links": { + "link_calendar": { "default_tools_approval_mode": "approve" }, + "link_other": { "default_tools_approval_mode": "prompt" }, + "link_without_privacy": {}, + }, + }, + "drive": { + "links": { "link_drive": { "default_tools_approval_mode": "prompt" } }, + }, + "without_links": { + "default_tools_approval_mode": "writes", + }, + "empty_links": { + "default_tools_approval_mode": "prompt", + "links": {}, + }, + })) + .expect("apps config"); + let evaluator = AppToolPolicyEvaluator::from_parts( + Some(apps_config), + /*requirements_apps_config*/ None, + ); + + for (link_id, approval) in [ + (Some("link_calendar"), AppToolApproval::Approve), + (Some("link_other"), AppToolApproval::Prompt), + (Some("link_without_privacy"), AppToolApproval::Auto), + (Some("link_drive"), AppToolApproval::Auto), + (None, AppToolApproval::Auto), + ] { + assert_eq!( + evaluator.policy(AppToolPolicyInput { + link_id, + ..input("events/list", /*tool_title*/ None) + }), + AppToolPolicy { + enabled: true, + approval, + } + ); + } + + assert_eq!( + evaluator.policy(AppToolPolicyInput { + link_id: Some("link_calendar"), + ..input("events/create", /*tool_title*/ None) + }), + AppToolPolicy { + enabled: true, + approval: AppToolApproval::Writes, + } + ); + + for (connector_id, approval) in [ + ("without_links", AppToolApproval::Writes), + ("empty_links", AppToolApproval::Prompt), + ] { + assert_eq!( + evaluator.policy(AppToolPolicyInput { + connector_id: Some(connector_id), + link_id: Some("link_calendar"), + ..input("events/list", /*tool_title*/ None) + }), + AppToolPolicy { + enabled: true, + approval, + } + ); + } +} + +#[test] +fn link_privacy_mode_preserves_managed_connector_requirements() { + let apps_config = serde_json::from_value(serde_json::json!({ + "calendar": { + "links": { "link_calendar": { "default_tools_approval_mode": "approve" } }, + }, + })) + .expect("apps config"); + let mut requirements = + app_tool_requirements("calendar", "events/create", AppToolApproval::Prompt); + requirements + .apps + .get_mut("calendar") + .expect("calendar requirement") + .enabled = Some(false); + let evaluator = AppToolPolicyEvaluator::from_parts(Some(apps_config), Some(&requirements)); + + assert_eq!( + evaluator.policy(AppToolPolicyInput { + link_id: Some("link_calendar"), + ..input("events/create", /*tool_title*/ None) + }), + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Prompt, + } + ); +} + +#[test] +fn per_tool_enable_overrides_app_level_hints() { + let apps_config = AppsConfigToml { + default: None, + apps: HashMap::from([( + "calendar".to_string(), + AppConfig { + enabled: true, + destructive_enabled: Some(false), + open_world_enabled: Some(false), + tools: Some(AppToolsConfig { + tools: HashMap::from([( + "events/create".to_string(), + AppToolConfig { + enabled: Some(true), + approval_mode: None, + }, + )]), + }), + ..Default::default() + }, + )]), + }; + + assert_eq!( + policy_from_apps_config( + Some(&apps_config), + Some("calendar"), + "events/create", + /*tool_title*/ None, + Some(true), + Some(true), + /*managed_approval*/ None, + ), + AppToolPolicy::default() + ); +} + +#[test] +fn default_tools_enable_overrides_app_level_hints() { + let mut app = AppConfig { + enabled: true, + destructive_enabled: Some(false), + open_world_enabled: Some(false), + default_tools_enabled: Some(true), + ..Default::default() + }; + let apps_config = |app: AppConfig| AppsConfigToml { + default: None, + apps: HashMap::from([("calendar".to_string(), app)]), + }; + + let enabled_policy = policy_from_apps_config( + Some(&apps_config(app.clone())), + Some("calendar"), + "events/create", + /*tool_title*/ None, + Some(true), + Some(true), + /*managed_approval*/ None, + ); + app.destructive_enabled = Some(true); + app.open_world_enabled = Some(true); + app.default_tools_enabled = Some(false); + app.default_tools_approval_mode = Some(AppToolApproval::Approve); + let disabled_policy = policy_from_apps_config( + Some(&apps_config(app)), + Some("calendar"), + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + /*managed_approval*/ None, + ); + + assert_eq!( + [enabled_policy, disabled_policy], + [ + AppToolPolicy::default(), + AppToolPolicy { + enabled: false, + approval: AppToolApproval::Approve, + }, + ] + ); +} + +#[test] +fn evaluator_uses_apps_default_tools_approval_mode_only_with_connector_id() { + let apps_config = AppsConfigToml { + default: Some(AppsDefaultConfig { + default_tools_approval_mode: Some(AppToolApproval::Prompt), + ..defaults( + /*enabled*/ true, /*destructive_enabled*/ true, + /*open_world_enabled*/ true, + ) + }), + apps: HashMap::new(), + }; + + assert_eq!( + [ + policy_from_apps_config( + Some(&apps_config), + Some("calendar"), + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + /*managed_approval*/ None, + ), + policy_from_apps_config( + Some(&apps_config), + /*connector_id*/ None, + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + /*managed_approval*/ None, + ), + ], + [ + AppToolPolicy { + enabled: true, + approval: AppToolApproval::Prompt, + }, + AppToolPolicy::default(), + ] + ); +} + +#[test] +fn evaluator_prefers_app_default_tools_approval_mode_over_apps_default() { + let apps_config = AppsConfigToml { + default: Some(AppsDefaultConfig { + default_tools_approval_mode: Some(AppToolApproval::Approve), + ..defaults( + /*enabled*/ true, /*destructive_enabled*/ true, + /*open_world_enabled*/ true, + ) + }), + apps: HashMap::from([( + "calendar".to_string(), + AppConfig { + enabled: true, + default_tools_approval_mode: Some(AppToolApproval::Prompt), + tools: Some(AppToolsConfig { + tools: HashMap::new(), + }), + ..Default::default() + }, + )]), + }; + + assert_eq!( + policy_from_apps_config( + Some(&apps_config), + Some("calendar"), + "events/list", + /*tool_title*/ None, + /*destructive_hint*/ None, + /*open_world_hint*/ None, + /*managed_approval*/ None, + ), + AppToolPolicy { + enabled: true, + approval: AppToolApproval::Prompt, + } + ); +} + +#[test] +fn evaluator_matches_tool_title_for_user_config() { + let apps_config = AppsConfigToml { + default: None, + apps: HashMap::from([( + "calendar".to_string(), + AppConfig { + enabled: true, + destructive_enabled: Some(false), + open_world_enabled: Some(false), + default_tools_approval_mode: Some(AppToolApproval::Auto), + default_tools_enabled: Some(false), + tools: Some(AppToolsConfig { + tools: HashMap::from([( + "events/create".to_string(), + AppToolConfig { + enabled: Some(true), + approval_mode: Some(AppToolApproval::Approve), + }, + )]), + }), + ..Default::default() + }, + )]), + }; + + assert_eq!( + policy_from_apps_config( + Some(&apps_config), + Some("calendar"), + "calendar_events/create", + Some("events/create"), + Some(true), + Some(true), + /*managed_approval*/ None, + ), + AppToolPolicy { + enabled: true, + approval: AppToolApproval::Approve, + } + ); +} + +fn input<'a>(tool_name: &'a str, tool_title: Option<&'a str>) -> AppToolPolicyInput<'a> { + AppToolPolicyInput { + connector_id: Some("calendar"), + link_id: None, + tool_name, + tool_title, + destructive_hint: Some(true), + open_world_hint: Some(true), + } +} + +fn app(id: &str, enabled: bool) -> AppInfo { + AppInfo { + id: id.to_string(), + name: id.to_string(), + description: None, + logo_url: None, + logo_url_dark: None, + icon_assets: None, + icon_dark_assets: None, + distribution_channel: None, + branding: None, + app_metadata: None, + labels: None, + install_url: None, + is_accessible: true, + is_enabled: enabled, + plugin_display_names: Vec::new(), + } +} + +fn policy_from_apps_config( + apps_config: Option<&AppsConfigToml>, + connector_id: Option<&str>, + tool_name: &str, + tool_title: Option<&str>, + destructive_hint: Option, + open_world_hint: Option, + managed_approval: Option, +) -> AppToolPolicy { + let requirements = managed_approval.map(|approval| { + app_tool_requirements( + connector_id.expect("managed approval requires a connector id"), + tool_name, + approval, + ) + }); + policy_from_config_parts( + apps_config, + requirements.as_ref(), + connector_id, + tool_name, + tool_title, + destructive_hint, + open_world_hint, + ) +} + +fn policy_from_config_parts( + apps_config: Option<&AppsConfigToml>, + requirements_apps_config: Option<&AppsRequirementsToml>, + connector_id: Option<&str>, + tool_name: &str, + tool_title: Option<&str>, + destructive_hint: Option, + open_world_hint: Option, +) -> AppToolPolicy { + let requirements = ConfigRequirementsToml { + apps: requirements_apps_config.cloned(), + ..Default::default() + }; + let config_layer_stack = + ConfigLayerStack::new(Vec::new(), ConfigRequirements::default(), requirements) + .expect("config layer stack"); + let config_layer_stack = if let Some(apps_config) = apps_config { + let mut user_config = TomlValue::Table(Default::default()); + user_config + .as_table_mut() + .expect("user config table") + .insert( + "apps".to_string(), + TomlValue::try_from(apps_config).expect("serialize apps config"), + ); + let config_toml_path = + AbsolutePathBuf::try_from(std::env::temp_dir().join(CONFIG_TOML_FILE)) + .expect("absolute config path"); + config_layer_stack + .with_user_config(&config_toml_path, user_config) + .expect("apps user config should be valid") + } else { + config_layer_stack + }; + AppToolPolicyEvaluator::new(&config_layer_stack).policy(AppToolPolicyInput { + connector_id, + link_id: None, + tool_name, + tool_title, + destructive_hint, + open_world_hint, + }) +} + +fn app_enabled_requirement(app_id: &str, enabled: bool) -> AppsRequirementsToml { + AppsRequirementsToml { + apps: BTreeMap::from([( + app_id.to_string(), + AppRequirementToml { + enabled: Some(enabled), + tools: None, + }, + )]), + } +} + +fn app_tool_requirements( + app_id: &str, + tool_name: &str, + approval_mode: AppToolApproval, +) -> AppsRequirementsToml { + AppsRequirementsToml { + apps: BTreeMap::from([( + app_id.to_string(), + AppRequirementToml { + enabled: None, + tools: Some(AppToolsRequirementsToml { + tools: BTreeMap::from([( + tool_name.to_string(), + AppToolRequirementToml { + approval_mode: Some(approval_mode), + ..Default::default() + }, + )]), + }), + }, + )]), + } +} + +fn defaults( + enabled: bool, + destructive_enabled: bool, + open_world_enabled: bool, +) -> AppsDefaultConfig { + AppsDefaultConfig { + enabled, + approvals_reviewer: None, + destructive_enabled, + open_world_enabled, + default_tools_approval_mode: None, + } +} diff --git a/codex-rs/connectors/src/connector_runtime/mod.rs b/codex-rs/connectors/src/connector_runtime/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..23858e82d3853fb4e06134590e403fac0bddc4dc --- /dev/null +++ b/codex-rs/connectors/src/connector_runtime/mod.rs @@ -0,0 +1,477 @@ +//! Shared runtime snapshot for connector-backed MCP tools. +//! +//! Runtime snapshots are process-local live state scoped by account and +//! workspace. Disk is best-effort cold-start persistence; a context reads it +//! once when created and never rereads it. Full connector metadata is +//! owned by the connector metadata store, not by this module. +//! Live catalog subscriptions publish only successful fetches from a matching +//! discovery scope, never disk snapshots or another scope's discovery winner. +//! Equivalent live contexts share one current catalog. Each fetch still contacts +//! the server; equal definitions retain their established storage and ordering. +//! Providers expire with their last context. + +use std::collections::HashMap; +use std::path::Path; +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::Weak; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; +use std::time::Duration; +use std::time::Instant; +use std::time::SystemTime; + +use arc_swap::ArcSwapOption; +use codex_login::CodexAuth; +use codex_protocol::mcp::McpServerInfo; +use serde::Deserialize; +use serde::Serialize; +use serde::de::DeserializeOwned; +use tokio::sync::watch; + +use self::persistence::load_cached_codex_apps_server_info; +use self::persistence::load_cached_connector_runtime_for_identity; +use self::persistence::persist_codex_apps_cache; +use self::persistence::server_info_cache_path; +use self::persistence::tools_cache_path; + +const MCP_TOOLS_CACHE_PUBLISH_DURATION_METRIC: &str = "codex.mcp.tools.cache_publish.duration_ms"; + +/// The current immutable tools for matching discovery inputs. +struct CatalogProvider { + scope: Vec, + updates: watch::Sender>>>, +} + +/// Values stored in the connector runtime's persisted tool snapshot. +/// +/// The runtime uses the connector-owned Codex Apps cache layout for every +/// serializable, cloneable payload. Equality determines whether fresh results can +/// retain the previous storage and tool version, so it must include all metadata +/// that affects readers or prepared calls. +pub trait ConnectorRuntimePayload: Clone + PartialEq + Serialize + DeserializeOwned {} + +impl ConnectorRuntimePayload for T where T: Clone + PartialEq + Serialize + DeserializeOwned {} + +/// The account and workspace identity of a connector runtime catalog. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +pub struct ConnectorRuntimeContextKey { + account_id: Option, + chatgpt_user_id: Option, + is_workspace_account: bool, +} + +impl ConnectorRuntimeContextKey { + pub fn personal(account_id: Option, chatgpt_user_id: Option) -> Self { + Self { + account_id, + chatgpt_user_id, + is_workspace_account: false, + } + } + + pub fn workspace(account_id: Option, chatgpt_user_id: Option) -> Self { + Self { + account_id, + chatgpt_user_id, + is_workspace_account: true, + } + } +} + +/// Builds the connector runtime context key for the active Codex auth. +pub fn connector_runtime_context_key(auth: Option<&CodexAuth>) -> ConnectorRuntimeContextKey { + let account_id = auth.and_then(CodexAuth::get_account_id); + let chatgpt_user_id = auth.and_then(CodexAuth::get_chatgpt_user_id); + if auth.is_some_and(CodexAuth::is_workspace_account) { + ConnectorRuntimeContextKey::workspace(account_id, chatgpt_user_id) + } else { + ConnectorRuntimeContextKey::personal(account_id, chatgpt_user_id) + } +} + +/// Returns the persisted connector runtime tools cache path for the active auth identity. +pub fn connector_runtime_cache_path(codex_home: &Path, auth: Option<&CodexAuth>) -> PathBuf { + let identity = ConnectorRuntimeIdentity { + codex_home: codex_home.to_path_buf(), + key: connector_runtime_context_key(auth), + }; + tools_cache_path(&identity) +} + +/// One atomically published connector runtime state. +/// +/// Tools remain raw and in response order. Local and managed configuration is +/// intentionally applied by readers rather than persisted in this snapshot. +#[derive(Debug, Clone)] +pub struct ConnectorRuntimeSnapshot { + tools: Arc<[T]>, + refreshed_at: SystemTime, + generation: u64, + tools_version: u64, +} + +impl ConnectorRuntimeSnapshot { + pub fn tools(&self) -> &[T] { + &self.tools + } + + pub fn shared_tools(&self) -> Arc<[T]> { + Arc::clone(&self.tools) + } + + /// Advances only when tool definitions change, so equal refreshes preserve prepared calls. + pub fn tools_version(&self) -> u64 { + self.tools_version + } + + pub fn refreshed_at(&self) -> SystemTime { + self.refreshed_at + } + + pub fn age(&self) -> Duration { + SystemTime::now() + .duration_since(self.refreshed_at) + .unwrap_or_default() + } +} + +/// Process-scoped registry of connector runtime state by account and workspace. +/// +/// Contexts with the same identity share one live entry. Different identities +/// remain independently available for clients that already hold their context. +pub struct ConnectorRuntimeManager { + entries: Arc>>>>, + disk_cache: ConnectorRuntimeDiskCache, +} + +impl Clone for ConnectorRuntimeManager { + fn clone(&self) -> Self { + Self { + entries: Arc::clone(&self.entries), + disk_cache: self.disk_cache, + } + } +} + +impl Default for ConnectorRuntimeManager { + fn default() -> Self { + Self { + entries: Arc::new(Mutex::new(HashMap::new())), + disk_cache: ConnectorRuntimeDiskCache::Enabled, + } + } +} + +impl ConnectorRuntimeManager { + /// Constructs a process-local connector runtime that never reads or writes the disk cache. + pub fn new_without_cache() -> Self { + Self { + entries: Arc::new(Mutex::new(HashMap::new())), + disk_cache: ConnectorRuntimeDiskCache::Disabled, + } + } + + pub fn current_snapshot( + &self, + codex_home: PathBuf, + key: ConnectorRuntimeContextKey, + ) -> Option>> { + self.context(codex_home, key).current_snapshot() + } + + pub fn context( + &self, + codex_home: PathBuf, + key: ConnectorRuntimeContextKey, + ) -> ConnectorRuntimeContext { + let identity = ConnectorRuntimeIdentity { codex_home, key }; + let mut entries = lock_unpoisoned(&self.entries); + let entry = entries + .entry(identity.clone()) + .or_insert_with(|| Arc::new(ConnectorRuntimeEntry::new(identity, self.disk_cache))) + .clone(); + ConnectorRuntimeContext { + entry, + live_catalog: None, + } + } +} + +/// Handle to one shared account/workspace connector runtime. +pub struct ConnectorRuntimeContext { + entry: Arc>, + live_catalog: Option>>, +} + +impl Clone for ConnectorRuntimeContext { + fn clone(&self) -> Self { + Self { + entry: Arc::clone(&self.entry), + live_catalog: self.live_catalog.clone(), + } + } +} + +impl ConnectorRuntimeContext { + /// Groups executable catalogs by equivalent discovery inputs within this account/home. + /// Include transport/auth, requested capabilities, and initialization results. + /// Calling this again refines the existing scope; it does not replace it. + /// Account-wide discovery reads are unaffected. + pub fn with_live_scope(mut self, scope: String) -> Self { + let mut scope_parts = self + .live_catalog + .as_ref() + .map(|catalog| catalog.scope.clone()) + .unwrap_or_default(); + scope_parts.push(scope); + let mut catalogs = lock_unpoisoned(&self.entry.live_catalogs); + catalogs.retain(|_, catalog| catalog.strong_count() > 0); + let catalog = catalogs.entry(scope_parts.clone()).or_default(); + self.live_catalog = Some(catalog.upgrade().unwrap_or_else(|| { + let provider = Arc::new(CatalogProvider { + scope: scope_parts, + updates: watch::channel(None).0, + }); + *catalog = Arc::downgrade(&provider); + provider + })); + drop(catalogs); + self + } + + /// Detaches a server that opts out of live catalog sharing, including publication. + pub fn without_live_scope(mut self) -> Self { + self.live_catalog = None; + self + } + + /// Subscribes to accepted live tools without refetching or waiting for other clients. + pub fn subscribe(&self) -> Option>>>> { + self.live_catalog + .as_ref() + .map(|catalog| catalog.updates.subscribe()) + } + + pub fn current_snapshot(&self) -> Option>> { + self.entry.current_snapshot.load_full() + } + + pub fn has_current_tools(&self) -> bool { + self.current_snapshot().is_some() + } + + pub fn begin_fetch(&self, source: ConnectorRuntimeFetchSource) -> ConnectorRuntimeFetchTicket { + ConnectorRuntimeFetchTicket { + generation: self + .entry + .next_fetch_generation + .fetch_add(1, Ordering::Relaxed) + + 1, + source, + } + } + + pub fn cached_server_info(&self) -> Option { + match self.entry.disk_cache { + ConnectorRuntimeDiskCache::Enabled => load_cached_codex_apps_server_info(self), + ConnectorRuntimeDiskCache::Disabled => None, + } + } + + fn tools_cache_path(&self) -> PathBuf { + tools_cache_path(&self.entry.identity) + } + + fn server_info_cache_path(&self) -> PathBuf { + server_info_cache_path(&self.entry.identity) + } + + pub fn current_tools(&self) -> Option> { + self.current_snapshot() + .map(|snapshot| snapshot.tools.to_vec()) + } + + pub fn publish_runtime_if_newest_accepted( + &self, + ticket: ConnectorRuntimeFetchTicket, + server_info: &McpServerInfo, + tools: Vec, + ) -> Arc> { + match self.entry.disk_cache { + ConnectorRuntimeDiskCache::Enabled => self.publish_runtime_if_newest_accepted_with( + ticket, + server_info, + tools, + persist_codex_apps_cache, + ), + ConnectorRuntimeDiskCache::Disabled => self.publish_runtime_if_newest_accepted_with( + ticket, + server_info, + tools, + |_, _, _| {}, + ), + } + } + + fn publish_runtime_if_newest_accepted_with( + &self, + ticket: ConnectorRuntimeFetchTicket, + server_info: &McpServerInfo, + tools: Vec, + persist: impl FnOnce(&ConnectorRuntimeContext, &McpServerInfo, &ConnectorRuntimeSnapshot), + ) -> Arc> { + let publish_start = Instant::now(); + let mut last_accepted_generation = lock_unpoisoned(&self.entry.last_accepted_generation); + let prior = self + .live_catalog + .as_ref() + .and_then(|catalog| catalog.updates.borrow().clone()); + let (tools, tools_version) = match prior { + Some(prior) + if prior.generation < ticket.generation && prior.tools() == tools.as_slice() => + { + (prior.shared_tools(), prior.tools_version) + } + _ => (tools.into(), ticket.generation), + }; + let snapshot = Arc::new(ConnectorRuntimeSnapshot { + tools, + tools_version, + refreshed_at: SystemTime::now(), + generation: ticket.generation, + }); + if let Some(live_catalog) = &self.live_catalog { + live_catalog.updates.send_if_modified(|current| { + if current + .as_ref() + .is_some_and(|current| current.generation >= ticket.generation) + { + return false; + } + *current = Some(Arc::clone(&snapshot)); + true + }); + } + if ticket.generation <= *last_accepted_generation + && let Some(snapshot) = self.current_snapshot() + { + drop(last_accepted_generation); + emit_duration( + MCP_TOOLS_CACHE_PUBLISH_DURATION_METRIC, + publish_start.elapsed(), + &[("source", ticket.source.as_str()), ("result", "stale")], + ); + return snapshot; + } + + *last_accepted_generation = ticket.generation; + self.entry + .current_snapshot + .store(Some(Arc::clone(&snapshot))); + // Keep the generation guard through persistence so accepted generations cannot reach disk + // out of order. + persist(self, server_info, snapshot.as_ref()); + drop(last_accepted_generation); + emit_duration( + MCP_TOOLS_CACHE_PUBLISH_DURATION_METRIC, + publish_start.elapsed(), + &[("source", ticket.source.as_str()), ("result", "published")], + ); + snapshot + } + + pub fn publish_if_newest_accepted( + &self, + ticket: ConnectorRuntimeFetchTicket, + server_info: &McpServerInfo, + tools: Vec, + ) -> Vec { + self.publish_runtime_if_newest_accepted(ticket, server_info, tools) + .tools + .to_vec() + } +} + +#[derive(Debug, Clone, Copy)] +pub enum ConnectorRuntimeFetchSource { + Startup, + HardRefresh, +} + +impl ConnectorRuntimeFetchSource { + fn as_str(self) -> &'static str { + match self { + Self::Startup => "startup", + Self::HardRefresh => "hard_refresh", + } + } +} + +pub struct ConnectorRuntimeFetchTicket { + generation: u64, + source: ConnectorRuntimeFetchSource, +} + +/// All live state owned by one connector identity. +struct ConnectorRuntimeEntry { + identity: ConnectorRuntimeIdentity, + disk_cache: ConnectorRuntimeDiskCache, + current_snapshot: ArcSwapOption>, + next_fetch_generation: AtomicU64, + last_accepted_generation: Mutex, + live_catalogs: Mutex, Weak>>>, +} + +impl ConnectorRuntimeEntry { + fn new(identity: ConnectorRuntimeIdentity, disk_cache: ConnectorRuntimeDiskCache) -> Self { + let current_snapshot = match disk_cache { + ConnectorRuntimeDiskCache::Enabled => { + load_cached_connector_runtime_for_identity(&identity).map(Arc::new) + } + ConnectorRuntimeDiskCache::Disabled => None, + }; + Self { + identity, + disk_cache, + current_snapshot: ArcSwapOption::from(current_snapshot), + next_fetch_generation: AtomicU64::new(0), + last_accepted_generation: Mutex::new(0), + live_catalogs: Mutex::new(HashMap::new()), + } + } +} + +#[derive(Clone, Copy)] +enum ConnectorRuntimeDiskCache { + Enabled, + Disabled, +} + +/// Everything that decides whether two connector runtime clients can share a snapshot. +/// +/// The auth key says whose runtime catalog we are reading. `codex_home` keeps +/// the persisted cache under the right home directory. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +struct ConnectorRuntimeIdentity { + codex_home: PathBuf, + key: ConnectorRuntimeContextKey, +} + +fn emit_duration(metric: &str, duration: Duration, tags: &[(&str, &str)]) { + if let Some(metrics) = codex_otel::global() { + let _ = metrics.record_duration(metric, duration, tags); + } +} + +fn lock_unpoisoned(mutex: &Mutex) -> std::sync::MutexGuard<'_, T> { + mutex + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) +} + +mod persistence; + +#[cfg(test)] +mod tests; diff --git a/codex-rs/connectors/src/connector_runtime/persistence.rs b/codex-rs/connectors/src/connector_runtime/persistence.rs new file mode 100644 index 0000000000000000000000000000000000000000..8abed04a7409b4749aeed59b8011a7bb6bef7609 --- /dev/null +++ b/codex-rs/connectors/src/connector_runtime/persistence.rs @@ -0,0 +1,273 @@ +//! Bounded, atomic persistence for connector runtime snapshots. + +use std::fs::File; +use std::io::Read; +use std::io::Write; +use std::path::Path; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::Instant; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +use anyhow::Context; +use anyhow::anyhow; +use codex_protocol::mcp::McpServerInfo; +use serde::Deserialize; +use serde::Serialize; +use sha1::Digest; +use sha1::Sha1; +use tempfile::NamedTempFile; +use tracing::instrument; + +use super::ConnectorRuntimeContext; +use super::ConnectorRuntimeIdentity; +use super::ConnectorRuntimePayload; +use super::ConnectorRuntimeSnapshot; +use super::emit_duration; + +const MCP_TOOLS_CACHE_WRITE_DURATION_METRIC: &str = "codex.mcp.tools.cache_write.duration_ms"; +const CODEX_APPS_TOOLS_CACHE_DIR: &str = "cache/codex_apps_tools"; +pub(crate) const CODEX_APPS_TOOLS_CACHE_SCHEMA_VERSION: u8 = 4; +const CODEX_APPS_SERVER_INFO_CACHE_DIR: &str = "cache/codex_apps_server_info"; +const CODEX_APPS_SERVER_INFO_CACHE_SCHEMA_VERSION: u8 = 1; +pub(crate) const CODEX_APPS_TOOLS_CACHE_MAX_BYTES: u64 = 32 * 1024 * 1024; + +pub(crate) fn tools_cache_path(identity: &ConnectorRuntimeIdentity) -> PathBuf { + cache_path_in(identity, CODEX_APPS_TOOLS_CACHE_DIR) +} + +pub(crate) fn server_info_cache_path(identity: &ConnectorRuntimeIdentity) -> PathBuf { + cache_path_in(identity, CODEX_APPS_SERVER_INFO_CACHE_DIR) +} + +fn cache_path_in(identity: &ConnectorRuntimeIdentity, cache_dir: &str) -> PathBuf { + // `codex_home` is already the parent directory. Keep it out of the + // filename hash so non-UTF-8 Unix paths cannot collapse distinct auth keys. + let identity_json = serde_json::to_string(&identity.key).unwrap_or_default(); + let identity_hash = sha1_hex(&identity_json); + identity + .codex_home + .join(cache_dir) + .join(format!("{identity_hash}.json")) +} + +#[instrument(level = "trace", skip_all)] +pub(crate) fn load_cached_connector_runtime_for_identity( + identity: &ConnectorRuntimeIdentity, +) -> Option> { + let cache_path = tools_cache_path(identity); + let (bytes, modified_at) = read_bounded_cache_file(&cache_path).ok()?; + let cache: CodexAppsToolsDiskCache = serde_json::from_slice(&bytes).ok()?; + (cache.schema_version == CODEX_APPS_TOOLS_CACHE_SCHEMA_VERSION).then_some( + ConnectorRuntimeSnapshot { + tools: cache.tools, + refreshed_at: modified_at, + generation: 0, + tools_version: 0, + }, + ) +} + +pub(crate) fn write_cached_connector_runtime( + cache_context: &ConnectorRuntimeContext, + snapshot: &ConnectorRuntimeSnapshot, +) -> anyhow::Result<()> +where + T: ConnectorRuntimePayload, +{ + let cache_path = cache_context.tools_cache_path(); + let bytes = serde_json::to_vec_pretty(&CodexAppsToolsDiskCache { + schema_version: CODEX_APPS_TOOLS_CACHE_SCHEMA_VERSION, + tools: snapshot.tools.clone(), + }) + .context("failed to serialize connector runtime cache")?; + write_codex_apps_cache_file(&cache_path, "runtime", bytes) +} + +#[instrument(level = "trace", skip_all)] +pub(crate) fn load_cached_codex_apps_server_info( + cache_context: &ConnectorRuntimeContext, +) -> Option { + let (bytes, _) = read_bounded_cache_file(&cache_context.server_info_cache_path()).ok()?; + let cache: CodexAppsServerInfoDiskCache = serde_json::from_slice(&bytes).ok()?; + (cache.schema_version == CODEX_APPS_SERVER_INFO_CACHE_SCHEMA_VERSION) + .then_some(cache.server_info) +} + +fn write_cached_codex_apps_server_info( + cache_context: &ConnectorRuntimeContext, + server_info: &McpServerInfo, +) -> anyhow::Result<()> { + let cache_path = cache_context.server_info_cache_path(); + let bytes = serde_json::to_vec_pretty(&CodexAppsServerInfoDiskCache { + schema_version: CODEX_APPS_SERVER_INFO_CACHE_SCHEMA_VERSION, + server_info: server_info.clone(), + }) + .context("failed to serialize Codex Apps server info cache")?; + write_codex_apps_cache_file(&cache_path, "server info", bytes) +} + +pub(crate) fn persist_codex_apps_cache( + cache_context: &ConnectorRuntimeContext, + server_info: &McpServerInfo, + snapshot: &ConnectorRuntimeSnapshot, +) where + T: ConnectorRuntimePayload, +{ + let cache_write_start = Instant::now(); + let tools_result = write_cached_connector_runtime(cache_context, snapshot); + if let Err(err) = &tools_result { + tracing::warn!("failed to write connector runtime cache: {err:#}"); + } + let server_info_result = write_cached_codex_apps_server_info(cache_context, server_info); + if let Err(err) = &server_info_result { + tracing::warn!("failed to write Codex Apps server info cache: {err:#}"); + } + let status = if tools_result.is_ok() && server_info_result.is_ok() { + "success" + } else { + "failure" + }; + emit_duration( + MCP_TOOLS_CACHE_WRITE_DURATION_METRIC, + cache_write_start.elapsed(), + &[("status", status)], + ); +} + +fn read_bounded_cache_file(cache_path: &Path) -> anyhow::Result<(Vec, SystemTime)> { + let mut file = File::open(cache_path) + .with_context(|| format!("failed to open cache `{}`", cache_path.display()))?; + let metadata = file + .metadata() + .with_context(|| format!("failed to stat cache `{}`", cache_path.display()))?; + if metadata.len() > CODEX_APPS_TOOLS_CACHE_MAX_BYTES { + return Err(anyhow!( + "cache `{}` is {} bytes, exceeding the {} byte limit", + cache_path.display(), + metadata.len(), + CODEX_APPS_TOOLS_CACHE_MAX_BYTES + )); + } + let mut bytes = Vec::with_capacity(metadata.len() as usize); + std::io::Read::by_ref(&mut file) + .take(CODEX_APPS_TOOLS_CACHE_MAX_BYTES + 1) + .read_to_end(&mut bytes) + .with_context(|| format!("failed to read cache `{}`", cache_path.display()))?; + if bytes.len() as u64 > CODEX_APPS_TOOLS_CACHE_MAX_BYTES { + return Err(anyhow!( + "cache `{}` grew beyond the {} byte limit while reading", + cache_path.display(), + CODEX_APPS_TOOLS_CACHE_MAX_BYTES + )); + } + Ok((bytes, metadata.modified().unwrap_or(UNIX_EPOCH))) +} + +fn write_codex_apps_cache_file( + cache_path: &Path, + cache_name: &str, + bytes: Vec, +) -> anyhow::Result<()> { + let parent = cache_path.parent().ok_or_else(|| { + anyhow!( + "Codex Apps {cache_name} cache path `{}` has no parent", + cache_path.display() + ) + })?; + std::fs::create_dir_all(parent).with_context(|| { + format!( + "failed to create Codex Apps {cache_name} cache directory `{}`", + parent.display() + ) + })?; + let mut temporary = NamedTempFile::new_in(parent).with_context(|| { + format!( + "failed to create temporary Codex Apps {cache_name} cache in `{}`", + parent.display() + ) + })?; + temporary.write_all(&bytes).with_context(|| { + format!( + "failed to write temporary Codex Apps {cache_name} cache for `{}`", + cache_path.display() + ) + })?; + temporary.persist(cache_path).map_err(|error| { + anyhow!( + "failed to atomically replace Codex Apps {cache_name} cache `{}`: {}", + cache_path.display(), + error.error + ) + })?; + Ok(()) +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +struct CodexAppsToolsDiskCache { + schema_version: u8, + tools: Arc<[T]>, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +struct CodexAppsServerInfoDiskCache { + schema_version: u8, + server_info: McpServerInfo, +} + +fn sha1_hex(s: &str) -> String { + let mut hasher = Sha1::new(); + hasher.update(s.as_bytes()); + let sha1 = hasher.finalize(); + format!("{sha1:x}") +} + +#[cfg(test)] +pub(crate) fn write_cached_codex_apps_tools_for_test( + cache_context: &ConnectorRuntimeContext, + server_info: &McpServerInfo, + tools: &[T], +) where + T: ConnectorRuntimePayload, +{ + let snapshot = ConnectorRuntimeSnapshot { + tools: tools.into(), + refreshed_at: SystemTime::now(), + generation: 0, + tools_version: 0, + }; + cache_context + .entry + .current_snapshot + .store(Some(Arc::new(snapshot.clone()))); + persist_codex_apps_cache(cache_context, server_info, &snapshot); +} + +#[cfg(test)] +pub(crate) fn read_cached_codex_apps_tools( + cache_context: &ConnectorRuntimeContext, +) -> Option> +where + T: ConnectorRuntimePayload, +{ + load_cached_connector_runtime_for_identity(&cache_context.entry.identity) + .map(|snapshot| snapshot.tools.to_vec()) +} + +#[cfg(test)] +pub(crate) fn write_cached_codex_apps_tools( + cache_context: &ConnectorRuntimeContext, + tools: &[T], +) -> anyhow::Result<()> +where + T: ConnectorRuntimePayload, +{ + let snapshot = ConnectorRuntimeSnapshot { + tools: tools.into(), + refreshed_at: SystemTime::now(), + generation: 0, + tools_version: 0, + }; + write_cached_connector_runtime(cache_context, &snapshot) +} diff --git a/codex-rs/connectors/src/connector_runtime/tests.rs b/codex-rs/connectors/src/connector_runtime/tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..e4a4996dc39559dcaeb860c1058e52fd32f50e61 --- /dev/null +++ b/codex-rs/connectors/src/connector_runtime/tests.rs @@ -0,0 +1,864 @@ +use super::persistence::CODEX_APPS_TOOLS_CACHE_MAX_BYTES; +use super::persistence::CODEX_APPS_TOOLS_CACHE_SCHEMA_VERSION; +use super::persistence::read_cached_codex_apps_tools; +use super::persistence::write_cached_codex_apps_tools; +use super::persistence::write_cached_codex_apps_tools_for_test; +use super::*; +use codex_protocol::mcp::McpServerInfo; +use pretty_assertions::assert_eq; +use serde::Deserialize; +use serde::Serialize; +#[cfg(unix)] +use std::os::unix::ffi::OsStringExt; +use std::path::PathBuf; +use std::sync::Arc; +use tempfile::tempdir; + +const CODEX_APPS_MCP_SERVER_NAME: &str = "codex_apps"; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +struct TestTool { + server_name: String, + callable_name: String, + connector_id: Option, + connector_name: Option, +} + +fn create_test_tool(server_name: &str, tool_name: &str) -> TestTool { + TestTool { + server_name: server_name.to_string(), + callable_name: tool_name.to_string(), + connector_id: None, + connector_name: None, + } +} + +fn create_test_tool_with_connector( + server_name: &str, + tool_name: &str, + connector_id: &str, + connector_name: Option<&str>, +) -> TestTool { + let mut tool = create_test_tool(server_name, tool_name); + tool.connector_id = Some(connector_id.to_string()); + tool.connector_name = connector_name.map(ToOwned::to_owned); + tool +} + +fn create_codex_apps_tools_cache_context( + codex_home: PathBuf, + account_id: Option<&str>, + chatgpt_user_id: Option<&str>, +) -> ConnectorRuntimeContext { + ConnectorRuntimeManager::::default().context( + codex_home, + ConnectorRuntimeContextKey { + account_id: account_id.map(ToOwned::to_owned), + chatgpt_user_id: chatgpt_user_id.map(ToOwned::to_owned), + is_workspace_account: false, + }, + ) +} + +fn create_test_server_info(title: &str) -> McpServerInfo { + McpServerInfo { + name: "codex-apps".to_string(), + title: Some(title.to_string()), + version: "1.0.0".to_string(), + description: None, + icons: None, + website_url: None, + } +} + +#[test] +fn codex_apps_tools_cache_is_overwritten_by_last_write() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let tools_gateway_1 = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "one")]; + let tools_gateway_2 = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "two")]; + + write_cached_codex_apps_tools(&cache_context, &tools_gateway_1).expect("write first cache"); + let cached_gateway_1 = + read_cached_codex_apps_tools(&cache_context).expect("cache entry exists for first write"); + assert_eq!(cached_gateway_1[0].callable_name, "one"); + + write_cached_codex_apps_tools(&cache_context, &tools_gateway_2).expect("write second cache"); + let cached_gateway_2 = + read_cached_codex_apps_tools(&cache_context).expect("cache entry exists for second write"); + assert_eq!(cached_gateway_2[0].callable_name, "two"); +} + +#[test] +fn codex_apps_tools_cache_is_scoped_per_user() { + let codex_home = tempdir().expect("tempdir"); + let cache_context_user_1 = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let cache_context_user_2 = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-two"), + Some("user-two"), + ); + let tools_user_1 = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "one")]; + let tools_user_2 = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "two")]; + + write_cached_codex_apps_tools(&cache_context_user_1, &tools_user_1) + .expect("write user one cache"); + write_cached_codex_apps_tools(&cache_context_user_2, &tools_user_2) + .expect("write user two cache"); + + let read_user_1 = + read_cached_codex_apps_tools(&cache_context_user_1).expect("cache entry for user one"); + let read_user_2 = + read_cached_codex_apps_tools(&cache_context_user_2).expect("cache entry for user two"); + + assert_eq!(read_user_1[0].callable_name, "one"); + assert_eq!(read_user_2[0].callable_name, "two"); + assert_ne!( + cache_context_user_1.tools_cache_path(), + cache_context_user_2.tools_cache_path(), + "each user should get an isolated cache file" + ); +} + +#[test] +fn codex_apps_tools_cache_preserves_formerly_disallowed_connectors() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let tools = vec![ + create_test_tool_with_connector( + CODEX_APPS_MCP_SERVER_NAME, + "formerly_blocked_tool", + "connector_2b0a9009c9c64bf9933a3dae3f2b1254", + Some("Formerly Blocked"), + ), + create_test_tool_with_connector( + CODEX_APPS_MCP_SERVER_NAME, + "calendar_tool", + "calendar", + Some("Calendar"), + ), + ]; + + write_cached_codex_apps_tools(&cache_context, &tools).expect("write cache"); + let cached = read_cached_codex_apps_tools(&cache_context).expect("cache entry exists for user"); + + assert_eq!( + cached + .iter() + .map(|tool| (tool.callable_name.as_str(), tool.connector_id.as_deref())) + .collect::>(), + vec![ + ( + "formerly_blocked_tool", + Some("connector_2b0a9009c9c64bf9933a3dae3f2b1254") + ), + ("calendar_tool", Some("calendar")), + ] + ); +} + +#[test] +fn codex_apps_tools_cache_is_ignored_when_schema_version_mismatches() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let cache_path = cache_context.tools_cache_path(); + if let Some(parent) = cache_path.parent() { + std::fs::create_dir_all(parent).expect("create parent"); + } + let bytes = serde_json::to_vec_pretty(&serde_json::json!({ + "schema_version": CODEX_APPS_TOOLS_CACHE_SCHEMA_VERSION + 1, + "tools": [create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "one")], + })) + .expect("serialize"); + std::fs::write(cache_path, bytes).expect("write"); + + assert!(read_cached_codex_apps_tools(&cache_context).is_none()); +} + +#[test] +fn codex_apps_tools_cache_is_ignored_when_json_is_invalid() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let cache_path = cache_context.tools_cache_path(); + if let Some(parent) = cache_path.parent() { + std::fs::create_dir_all(parent).expect("create parent"); + } + std::fs::write(cache_path, b"{not json").expect("write"); + + assert!(read_cached_codex_apps_tools(&cache_context).is_none()); +} + +#[test] +fn startup_cached_codex_apps_tools_loads_from_disk_cache() { + let codex_home = tempdir().expect("tempdir"); + let writer_cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let cached_tools = vec![create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "calendar_search", + )]; + let server_info = create_test_server_info("Codex Apps"); + write_cached_codex_apps_tools_for_test(&writer_cache_context, &server_info, &cached_tools); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + + let startup_tools = cache_context + .current_tools() + .expect("expected startup snapshot to load from cache"); + let cached_server_info = cache_context.cached_server_info(); + + assert_eq!(startup_tools.len(), 1); + assert_eq!(startup_tools[0].server_name, CODEX_APPS_MCP_SERVER_NAME); + assert_eq!(startup_tools[0].callable_name, "calendar_search"); + assert_eq!(cached_server_info, Some(server_info)); +} + +#[test] +fn startup_cached_codex_apps_tools_loads_without_server_info_cache() { + let codex_home = tempdir().expect("tempdir"); + let writer_cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let cache_path = writer_cache_context.tools_cache_path(); + if let Some(parent) = cache_path.parent() { + std::fs::create_dir_all(parent).expect("create parent"); + } + let bytes = serde_json::to_vec_pretty(&serde_json::json!({ + "schema_version": CODEX_APPS_TOOLS_CACHE_SCHEMA_VERSION, + "tools": [create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "calendar_search")], + })) + .expect("serialize"); + std::fs::write(cache_path, bytes).expect("write"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + + let startup_tools = cache_context + .current_tools() + .expect("legacy startup snapshot should remain available"); + let cached_server_info = cache_context.cached_server_info(); + + assert_eq!(startup_tools.len(), 1); + assert_eq!(startup_tools[0].callable_name, "calendar_search"); + assert_eq!(cached_server_info, None); +} + +#[test] +fn codex_apps_server_info_cache_survives_legacy_tools_cache_write() { + let codex_home = tempdir().expect("tempdir"); + let cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let server_info = create_test_server_info("Codex Apps"); + write_cached_codex_apps_tools_for_test( + &cache_context, + &server_info, + &[create_test_tool( + CODEX_APPS_MCP_SERVER_NAME, + "calendar_search", + )], + ); + + let cache_path = cache_context.tools_cache_path(); + if let Some(parent) = cache_path.parent() { + std::fs::create_dir_all(parent).expect("create parent"); + } + let bytes = serde_json::to_vec_pretty(&serde_json::json!({ + "schema_version": CODEX_APPS_TOOLS_CACHE_SCHEMA_VERSION - 1, + "tools": [create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "calendar_search")], + })) + .expect("serialize"); + std::fs::write(cache_path, bytes).expect("write legacy tools cache"); + let startup_cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + + assert_eq!( + startup_cache_context.cached_server_info(), + Some(server_info) + ); + assert!(startup_cache_context.current_tools().is_none()); +} + +#[test] +fn codex_apps_tools_cache_context_does_not_reread_disk_after_creation() { + let codex_home = tempdir().expect("tempdir"); + let writer_cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let cached_tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "cached")]; + write_cached_codex_apps_tools(&writer_cache_context, &cached_tools).expect("write cache"); + let reader_cache_context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let updated_tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "updated")]; + write_cached_codex_apps_tools(&writer_cache_context, &updated_tools).expect("rewrite cache"); + + assert_eq!( + reader_cache_context + .current_tools() + .expect("in-memory tools")[0] + .callable_name, + "cached" + ); + assert_eq!( + read_cached_codex_apps_tools(&writer_cache_context).expect("disk tools")[0].callable_name, + "updated" + ); +} + +#[test] +fn codex_apps_tools_cache_publishes_newest_shared_snapshot() { + let codex_home = tempdir().expect("tempdir"); + let cache = ConnectorRuntimeManager::::default(); + let cache_context_1 = cache.context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey { + account_id: Some("account-one".to_string()), + chatgpt_user_id: Some("user-one".to_string()), + is_workspace_account: false, + }, + ); + let cache_context_2 = cache.context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey { + account_id: Some("account-one".to_string()), + chatgpt_user_id: Some("user-one".to_string()), + is_workspace_account: false, + }, + ); + let older_ticket = cache_context_1.begin_fetch(ConnectorRuntimeFetchSource::Startup); + let newer_ticket = cache_context_2.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh); + let server_info = create_test_server_info("Codex Apps"); + let newer_tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "newer")]; + let older_tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "older")]; + + let published_tools = + cache_context_2.publish_if_newest_accepted(newer_ticket, &server_info, newer_tools); + assert_eq!(cache_context_1.current_tools(), Some(published_tools)); + let current_tools = + cache_context_1.publish_if_newest_accepted(older_ticket, &server_info, older_tools); + + assert_eq!(current_tools[0].callable_name, "newer"); + assert_eq!( + cache_context_2.current_tools().expect("shared snapshot")[0].callable_name, + "newer" + ); + assert_eq!( + read_cached_codex_apps_tools(&cache_context_1).expect("persisted snapshot")[0] + .callable_name, + "newer" + ); +} + +#[test] +fn codex_apps_tools_cache_keeps_live_publish_when_disk_persistence_fails() { + let codex_home = tempdir().expect("tempdir"); + let codex_home_file = codex_home.path().join("not-a-directory"); + std::fs::write(&codex_home_file, b"occupied").expect("create codex home file"); + let cache_context = ConnectorRuntimeManager::::default().context( + codex_home_file, + ConnectorRuntimeContextKey { + account_id: Some("account-one".to_string()), + chatgpt_user_id: Some("user-one".to_string()), + is_workspace_account: false, + }, + ); + let tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "live")]; + let published_tools = cache_context.publish_if_newest_accepted( + cache_context.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh), + &create_test_server_info("Codex Apps"), + tools.clone(), + ); + + assert_eq!(published_tools, tools); + assert_eq!(cache_context.current_tools(), Some(tools)); +} + +#[test] +fn connector_runtime_without_cache_ignores_disk_state() { + let codex_home = tempdir().expect("tempdir"); + let writer = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "cached")]; + let server_info = create_test_server_info("Codex Apps"); + write_cached_codex_apps_tools_for_test(&writer, &server_info, &tools); + let context = ConnectorRuntimeManager::::new_without_cache().context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey { + account_id: Some("account-one".to_string()), + chatgpt_user_id: Some("user-one".to_string()), + is_workspace_account: false, + }, + ); + + assert_eq!(context.current_tools(), None); + assert_eq!(context.cached_server_info(), None); +} + +#[test] +fn connector_runtime_without_cache_publishes_without_writing() { + let temp_dir = tempdir().expect("tempdir"); + let codex_home = temp_dir.path().join("codex-home"); + let context = ConnectorRuntimeManager::::new_without_cache().context( + codex_home.clone(), + ConnectorRuntimeContextKey { + account_id: Some("account-one".to_string()), + chatgpt_user_id: Some("user-one".to_string()), + is_workspace_account: false, + }, + ); + let tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "live")]; + let published_tools = context.publish_if_newest_accepted( + context.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh), + &create_test_server_info("Codex Apps"), + tools.clone(), + ); + + assert_eq!(published_tools, tools); + assert_eq!(context.current_tools(), Some(tools)); + assert!(!codex_home.exists()); +} + +#[cfg(unix)] +#[test] +fn codex_apps_tools_cache_scopes_non_utf8_home_disk_paths() { + let codex_home = PathBuf::from(std::ffi::OsString::from_vec( + b"/tmp/codex-home-\xff".to_vec(), + )); + let cache = ConnectorRuntimeManager::::default(); + let user_one_context = cache.context( + codex_home.clone(), + ConnectorRuntimeContextKey { + account_id: Some("account-one".to_string()), + chatgpt_user_id: Some("user-one".to_string()), + is_workspace_account: false, + }, + ); + let user_two_context = cache.context( + codex_home, + ConnectorRuntimeContextKey { + account_id: Some("account-two".to_string()), + chatgpt_user_id: Some("user-two".to_string()), + is_workspace_account: false, + }, + ); + let cache_paths = [ + user_one_context.tools_cache_path(), + user_two_context.tools_cache_path(), + ]; + + assert_ne!(cache_paths[0], cache_paths[1]); +} + +#[test] +fn live_catalogs_isolate_scopes_and_reject_older_refreshes() { + let codex_home = tempdir().expect("tempdir"); + let manager = ConnectorRuntimeManager::::new_without_cache(); + let context = manager.context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey::personal( + Some("account".to_string()), + /*chatgpt_user_id*/ None, + ), + ); + let scope_a = context.clone().with_live_scope("endpoint-a".to_string()); + let scope_b = context.with_live_scope("endpoint-b".to_string()); + let updates_a = scope_a.subscribe().expect("scope A subscription"); + let updates_b = scope_b.subscribe().expect("scope B subscription"); + let server_info = create_test_server_info("Codex Apps"); + let older_a = scope_a.begin_fetch(ConnectorRuntimeFetchSource::Startup); + let newer_a = scope_a.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh); + let newest_b = scope_b.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh); + let tools_a = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "newer-a")]; + let tools_b = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "newest-b")]; + + scope_b.publish_if_newest_accepted(newest_b, &server_info, tools_b.clone()); + assert!(updates_a.borrow().is_none()); + let discovery = scope_a.publish_if_newest_accepted(newer_a, &server_info, tools_a.clone()); + assert_eq!(discovery, tools_b); + scope_a.publish_if_newest_accepted( + older_a, + &server_info, + vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "stale-a")], + ); + assert_eq!( + updates_a.borrow().as_ref().expect("scope A tools").tools(), + &tools_a + ); + assert_eq!( + updates_b.borrow().as_ref().expect("scope B tools").tools(), + &tools_b + ); +} + +#[test] +fn contexts_for_different_identities_keep_isolated_snapshots() { + let codex_home = tempdir().expect("tempdir"); + let manager = ConnectorRuntimeManager::::default(); + let context_a = manager.context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey { + account_id: Some("account-a".to_string()), + chatgpt_user_id: Some("user-a".to_string()), + is_workspace_account: false, + }, + ); + let context_a = context_a.with_live_scope("same-endpoint".to_string()); + let updates_a = context_a.subscribe().expect("account A subscription"); + let tools_a = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "tool-a")]; + let snapshot_a = context_a.publish_runtime_if_newest_accepted( + context_a.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh), + &create_test_server_info("Codex Apps"), + tools_a.clone(), + ); + let older_ticket_a = context_a.begin_fetch(ConnectorRuntimeFetchSource::Startup); + let context_b = manager.context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey { + account_id: Some("account-b".to_string()), + chatgpt_user_id: Some("user-b".to_string()), + is_workspace_account: false, + }, + ); + let context_b = context_b.with_live_scope("same-endpoint".to_string()); + let updates_b = context_b.subscribe().expect("account B subscription"); + let same_context_a = manager.context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey { + account_id: Some("account-a".to_string()), + chatgpt_user_id: Some("user-a".to_string()), + is_workspace_account: false, + }, + ); + let same_context_a = same_context_a.with_live_scope("same-endpoint".to_string()); + + assert!(Arc::ptr_eq( + &snapshot_a, + &same_context_a + .current_snapshot() + .expect("context A snapshot") + )); + assert!(context_b.current_snapshot().is_none()); + + let tools_b = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "tool-b")]; + let snapshot_b = context_b.publish_runtime_if_newest_accepted( + context_b.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh), + &create_test_server_info("Codex Apps"), + tools_b.clone(), + ); + let newer_tools_a = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "newer-a")]; + let newer_snapshot_a = same_context_a.publish_runtime_if_newest_accepted( + same_context_a.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh), + &create_test_server_info("Codex Apps"), + newer_tools_a.clone(), + ); + let stale_snapshot_a = context_a.publish_runtime_if_newest_accepted( + older_ticket_a, + &create_test_server_info("Codex Apps"), + vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "stale-a")], + ); + + assert_eq!(snapshot_a.tools(), &tools_a); + assert_eq!(snapshot_b.tools(), &tools_b); + assert_eq!(newer_snapshot_a.tools(), &newer_tools_a); + assert!(Arc::ptr_eq(&newer_snapshot_a, &stale_snapshot_a)); + assert!(Arc::ptr_eq( + &newer_snapshot_a, + &context_a.current_snapshot().expect("context A snapshot") + )); + assert!(Arc::ptr_eq( + &snapshot_b, + &context_b.current_snapshot().expect("context B snapshot") + )); + assert_eq!( + updates_a + .borrow() + .as_ref() + .expect("account A tools") + .tools(), + &newer_tools_a + ); + assert_eq!( + updates_b + .borrow() + .as_ref() + .expect("account B tools") + .tools(), + &tools_b + ); +} + +#[test] +fn live_provider_reuses_equal_tools_only_within_its_scope() { + let manager = ConnectorRuntimeManager::::new_without_cache(); + let context = manager.context( + PathBuf::from("unused"), + ConnectorRuntimeContextKey::personal( + /*account_id*/ None, /*chatgpt_user_id*/ None, + ), + ); + let first = context + .clone() + .with_live_scope("endpoint".into()) + .with_live_scope("capabilities".into()); + let second = context + .clone() + .with_live_scope("endpoint".into()) + .with_live_scope("capabilities".into()); + let other = context + .clone() + .with_live_scope("other-endpoint".into()) + .with_live_scope("capabilities".into()); + let tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "search")]; + let publish = |context: &ConnectorRuntimeContext| { + context.publish_runtime_if_newest_accepted( + context.begin_fetch(ConnectorRuntimeFetchSource::Startup), + &create_test_server_info("Apps"), + tools.clone(), + ) + }; + let original = publish(&first); + let repeated = publish(&second); + let separate = publish(&other); + assert!(Arc::ptr_eq( + &original.shared_tools(), + &repeated.shared_tools() + )); + assert_eq!(original.tools_version(), repeated.tools_version()); + assert!(!Arc::ptr_eq( + &original.shared_tools(), + &separate.shared_tools() + )); + drop(first); + drop(second); + let restarted = context + .with_live_scope("endpoint".into()) + .with_live_scope("capabilities".into()); + assert!(restarted.current_snapshot().is_some()); + assert!(restarted.subscribe().unwrap().borrow().is_none()); +} + +#[test] +fn oversized_tools_cache_is_ignored_during_initial_load() { + let codex_home = tempdir().expect("tempdir"); + let context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let cache_path = context.tools_cache_path(); + std::fs::create_dir_all(cache_path.parent().expect("cache parent")) + .expect("create cache parent"); + let file = std::fs::File::create(cache_path).expect("create oversized cache"); + file.set_len(CODEX_APPS_TOOLS_CACHE_MAX_BYTES + 1) + .expect("size oversized cache"); + + let reloaded = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + + assert!(reloaded.current_snapshot().is_none()); +} + +#[test] +fn cold_loaded_snapshot_uses_cache_modification_time() { + let codex_home = tempdir().expect("tempdir"); + let writer = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "cached")]; + write_cached_codex_apps_tools(&writer, &tools).expect("write tools cache"); + let modified_at = std::fs::metadata(writer.tools_cache_path()) + .and_then(|metadata| metadata.modified()) + .expect("cache modification time"); + + let reloaded = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let snapshot = reloaded.current_snapshot().expect("cold-loaded snapshot"); + + assert_eq!(snapshot.tools(), &tools); + assert_eq!(snapshot.refreshed_at(), modified_at); +} +#[test] +fn accepted_generations_finish_persistence_in_order() { + let codex_home = tempdir().expect("tempdir"); + let context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let older_ticket = context.begin_fetch(ConnectorRuntimeFetchSource::Startup); + let newer_ticket = context.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh); + let (older_persisting_tx, older_persisting_rx) = std::sync::mpsc::channel(); + let (release_older_tx, release_older_rx) = std::sync::mpsc::channel(); + let older_context = context.clone(); + let older_publish = std::thread::spawn(move || { + older_context.publish_runtime_if_newest_accepted_with( + older_ticket, + &create_test_server_info("Codex Apps"), + vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "older")], + move |_, _, _| { + older_persisting_tx + .send(()) + .expect("signal older persistence"); + release_older_rx.recv().expect("release older persistence"); + }, + ) + }); + older_persisting_rx + .recv_timeout(Duration::from_secs(1)) + .expect("older generation should enter persistence"); + + let (newer_persisting_tx, newer_persisting_rx) = std::sync::mpsc::channel(); + let newer_context = context; + let newer_publish = std::thread::spawn(move || { + newer_context.publish_runtime_if_newest_accepted_with( + newer_ticket, + &create_test_server_info("Codex Apps"), + vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "newer")], + move |_, _, _| { + newer_persisting_tx + .send(()) + .expect("signal newer persistence"); + }, + ) + }); + + assert!( + newer_persisting_rx + .recv_timeout(Duration::from_millis(20)) + .is_err() + ); + release_older_tx + .send(()) + .expect("allow older persistence to finish"); + newer_persisting_rx + .recv_timeout(Duration::from_secs(1)) + .expect("newer generation should persist after older generation"); + + older_publish.join().expect("join older publish"); + let newer_snapshot = newer_publish.join().expect("join newer publish"); + assert_eq!(newer_snapshot.tools()[0].callable_name, "newer"); +} + +#[test] +fn personal_and_workspace_contexts_are_distinct_even_with_matching_ids() { + let codex_home = tempdir().expect("tempdir"); + let manager = ConnectorRuntimeManager::::default(); + let personal_context = manager.context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey { + account_id: Some("account".to_string()), + chatgpt_user_id: Some("user".to_string()), + is_workspace_account: false, + }, + ); + let personal_tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "personal")]; + let _ = personal_context.publish_runtime_if_newest_accepted( + personal_context.begin_fetch(ConnectorRuntimeFetchSource::Startup), + &create_test_server_info("Codex Apps"), + personal_tools.clone(), + ); + + let workspace_context = manager.context( + codex_home.path().to_path_buf(), + ConnectorRuntimeContextKey { + account_id: Some("account".to_string()), + chatgpt_user_id: Some("user".to_string()), + is_workspace_account: true, + }, + ); + + let workspace_tools = vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "workspace")]; + let _ = workspace_context.publish_runtime_if_newest_accepted( + workspace_context.begin_fetch(ConnectorRuntimeFetchSource::Startup), + &create_test_server_info("Codex Apps"), + workspace_tools.clone(), + ); + + assert_eq!(personal_context.current_tools(), Some(personal_tools)); + assert_eq!(workspace_context.current_tools(), Some(workspace_tools)); + assert_ne!( + personal_context.tools_cache_path(), + workspace_context.tools_cache_path() + ); +} + +#[test] +fn live_publish_sets_timestamp_and_stale_publish_preserves_it() { + let codex_home = tempdir().expect("tempdir"); + let context = create_codex_apps_tools_cache_context( + codex_home.path().to_path_buf(), + Some("account-one"), + Some("user-one"), + ); + let stale_ticket = context.begin_fetch(ConnectorRuntimeFetchSource::Startup); + let current_ticket = context.begin_fetch(ConnectorRuntimeFetchSource::HardRefresh); + let before = SystemTime::now(); + let current = context.publish_runtime_if_newest_accepted( + current_ticket, + &create_test_server_info("Codex Apps"), + vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "current")], + ); + let after = SystemTime::now(); + + assert!(current.refreshed_at() >= before); + assert!(current.refreshed_at() <= after); + + let stale = context.publish_runtime_if_newest_accepted( + stale_ticket, + &create_test_server_info("Codex Apps"), + vec![create_test_tool(CODEX_APPS_MCP_SERVER_NAME, "stale")], + ); + assert!(Arc::ptr_eq(¤t, &stale)); + assert_eq!(stale.refreshed_at(), current.refreshed_at()); +} diff --git a/codex-rs/connectors/src/directory_cache.rs b/codex-rs/connectors/src/directory_cache.rs new file mode 100644 index 0000000000000000000000000000000000000000..abaa8b049e0ae233eb015c1c3e0b39aa8b407560 --- /dev/null +++ b/codex-rs/connectors/src/directory_cache.rs @@ -0,0 +1,113 @@ +use std::path::PathBuf; + +use serde::Deserialize; +use serde::Serialize; +use sha1::Digest; +use sha1::Sha1; +use tracing::warn; + +use crate::AppInfo; +use crate::ConnectorDirectoryCacheKey; + +pub(crate) const CONNECTOR_DIRECTORY_DISK_CACHE_SCHEMA_VERSION: u8 = 1; +const CONNECTOR_DIRECTORY_DISK_CACHE_DIR: &str = "cache/codex_app_directory"; + +#[derive(Clone)] +pub struct ConnectorDirectoryCacheContext { + pub(crate) codex_home: PathBuf, + pub(crate) cache_key: ConnectorDirectoryCacheKey, +} + +impl ConnectorDirectoryCacheContext { + pub fn new(codex_home: PathBuf, cache_key: ConnectorDirectoryCacheKey) -> Self { + Self { + codex_home, + cache_key, + } + } + + /// Returns the persisted connector directory cache path for this identity. + pub fn cache_path(&self) -> PathBuf { + let cache_key_json = serde_json::to_string(&self.cache_key).unwrap_or_default(); + let cache_key_hash = sha1_hex(&cache_key_json); + self.codex_home + .join(CONNECTOR_DIRECTORY_DISK_CACHE_DIR) + .join(format!("{cache_key_hash}.json")) + } +} + +pub(crate) enum CachedConnectorDirectoryDiskLoad { + Hit { connectors: Vec }, + Missing, + Invalid, +} + +pub(crate) fn load_cached_directory_connectors_from_disk( + cache_context: &ConnectorDirectoryCacheContext, +) -> CachedConnectorDirectoryDiskLoad { + let cache_path = cache_context.cache_path(); + let bytes = match std::fs::read(&cache_path) { + Ok(bytes) => bytes, + Err(err) if err.kind() == std::io::ErrorKind::NotFound => { + return CachedConnectorDirectoryDiskLoad::Missing; + } + Err(err) => { + warn!( + cache_path = %cache_path.display(), + "failed to read connector directory disk cache: {err}" + ); + return CachedConnectorDirectoryDiskLoad::Invalid; + } + }; + let cache: ConnectorDirectoryDiskCache = match serde_json::from_slice(&bytes) { + Ok(cache) => cache, + Err(err) => { + warn!( + cache_path = %cache_path.display(), + "failed to parse connector directory disk cache: {err}" + ); + let _ = std::fs::remove_file(cache_path); + return CachedConnectorDirectoryDiskLoad::Invalid; + } + }; + if cache.schema_version != CONNECTOR_DIRECTORY_DISK_CACHE_SCHEMA_VERSION { + let _ = std::fs::remove_file(cache_path); + return CachedConnectorDirectoryDiskLoad::Invalid; + } + + CachedConnectorDirectoryDiskLoad::Hit { + connectors: cache.connectors, + } +} + +pub(crate) fn write_cached_directory_connectors_to_disk( + cache_context: &ConnectorDirectoryCacheContext, + connectors: &[AppInfo], +) { + let cache_path = cache_context.cache_path(); + if let Some(parent) = cache_path.parent() + && std::fs::create_dir_all(parent).is_err() + { + return; + } + let Ok(bytes) = serde_json::to_vec_pretty(&ConnectorDirectoryDiskCache { + schema_version: CONNECTOR_DIRECTORY_DISK_CACHE_SCHEMA_VERSION, + connectors: connectors.to_vec(), + }) else { + return; + }; + let _ = std::fs::write(cache_path, bytes); +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +struct ConnectorDirectoryDiskCache { + schema_version: u8, + connectors: Vec, +} + +fn sha1_hex(value: &str) -> String { + let mut hasher = Sha1::new(); + hasher.update(value.as_bytes()); + let sha1 = hasher.finalize(); + format!("{sha1:x}") +} diff --git a/codex-rs/connectors/src/filter.rs b/codex-rs/connectors/src/filter.rs new file mode 100644 index 0000000000000000000000000000000000000000..3fcabd6fb6fca35bef49d8a1c8abddd3d2c28894 --- /dev/null +++ b/codex-rs/connectors/src/filter.rs @@ -0,0 +1,129 @@ +use std::collections::HashSet; + +use crate::AppInfo; + +pub fn filter_tool_suggest_discoverable_connectors( + directory_connectors: Vec, + accessible_connectors: &[AppInfo], + discoverable_connector_ids: &HashSet, +) -> Vec { + let accessible_connector_ids: HashSet<&str> = accessible_connectors + .iter() + .filter(|connector| connector.is_accessible) + .map(|connector| connector.id.as_str()) + .collect(); + + let mut connectors = directory_connectors + .into_iter() + .filter(|connector| !accessible_connector_ids.contains(connector.id.as_str())) + .filter(|connector| discoverable_connector_ids.contains(connector.id.as_str())) + .collect::>(); + connectors.sort_by(|left, right| { + left.name + .cmp(&right.name) + .then_with(|| left.id.cmp(&right.id)) + }); + connectors +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::metadata::connector_install_url; + use pretty_assertions::assert_eq; + + fn app(id: &str) -> AppInfo { + AppInfo { + id: id.to_string(), + name: id.to_string(), + description: None, + logo_url: None, + logo_url_dark: None, + icon_assets: None, + icon_dark_assets: None, + distribution_channel: None, + install_url: None, + branding: None, + app_metadata: None, + labels: None, + is_accessible: false, + is_enabled: true, + plugin_display_names: Vec::new(), + } + } + + fn named_app(id: &str, name: &str) -> AppInfo { + AppInfo { + id: id.to_string(), + name: name.to_string(), + install_url: Some(connector_install_url(name, id)), + ..app(id) + } + } + + #[test] + fn filter_tool_suggest_discoverable_connectors_keeps_only_plugin_backed_uninstalled_apps() { + let filtered = filter_tool_suggest_discoverable_connectors( + vec![ + named_app( + "connector_2128aebfecb84f64a069897515042a44", + "Google Calendar", + ), + named_app("connector_68df038e0ba48191908c8434991bbac2", "Gmail"), + named_app("connector_other", "Other"), + ], + &[AppInfo { + is_accessible: true, + ..named_app( + "connector_2128aebfecb84f64a069897515042a44", + "Google Calendar", + ) + }], + &HashSet::from([ + "connector_2128aebfecb84f64a069897515042a44".to_string(), + "connector_68df038e0ba48191908c8434991bbac2".to_string(), + ]), + ); + + assert_eq!( + filtered, + vec![named_app( + "connector_68df038e0ba48191908c8434991bbac2", + "Gmail", + )] + ); + } + + #[test] + fn filter_tool_suggest_discoverable_connectors_excludes_accessible_apps_even_when_disabled() { + let filtered = filter_tool_suggest_discoverable_connectors( + vec![ + named_app( + "connector_2128aebfecb84f64a069897515042a44", + "Google Calendar", + ), + named_app("connector_68df038e0ba48191908c8434991bbac2", "Gmail"), + ], + &[ + AppInfo { + is_accessible: true, + ..named_app( + "connector_2128aebfecb84f64a069897515042a44", + "Google Calendar", + ) + }, + AppInfo { + is_accessible: true, + is_enabled: false, + ..named_app("connector_68df038e0ba48191908c8434991bbac2", "Gmail") + }, + ], + &HashSet::from([ + "connector_2128aebfecb84f64a069897515042a44".to_string(), + "connector_68df038e0ba48191908c8434991bbac2".to_string(), + ]), + ); + + assert_eq!(filtered, Vec::::new()); + } +} diff --git a/codex-rs/connectors/src/lib.rs b/codex-rs/connectors/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..ec8b19c9217b41998bf96bdf3420c87e2b797d04 --- /dev/null +++ b/codex-rs/connectors/src/lib.rs @@ -0,0 +1,992 @@ +use std::collections::HashMap; +use std::future::Future; +use std::sync::LazyLock; +use std::sync::Mutex as StdMutex; +use std::time::Duration; +use std::time::Instant; + +use serde::Deserialize; +use serde::Serialize; + +pub mod accessible; +mod app_info; +mod app_tool_policy; +mod connector_runtime; +mod directory_cache; +pub mod filter; +pub mod merge; +pub mod metadata; +mod metadata_store; +mod plugin_config; +mod runtime_projection; +mod snapshot; + +pub use app_info::AppBranding; +pub use app_info::AppInfo; +pub use app_info::AppMetadata; +pub use app_info::AppReview; +pub use app_info::AppScreenshot; +pub use app_tool_policy::AppToolPolicy; +pub use app_tool_policy::AppToolPolicyEvaluator; +pub use app_tool_policy::AppToolPolicyInput; +pub use app_tool_policy::app_is_enabled; +pub use app_tool_policy::apps_config_from_layer_stack; +pub use connector_runtime::ConnectorRuntimeContext; +pub use connector_runtime::ConnectorRuntimeContextKey; +pub use connector_runtime::ConnectorRuntimeFetchSource; +pub use connector_runtime::ConnectorRuntimeFetchTicket; +pub use connector_runtime::ConnectorRuntimeManager; +pub use connector_runtime::ConnectorRuntimePayload; +pub use connector_runtime::ConnectorRuntimeSnapshot; +pub use connector_runtime::connector_runtime_cache_path; +pub use connector_runtime::connector_runtime_context_key; +pub use directory_cache::ConnectorDirectoryCacheContext; +pub use metadata_store::ConnectorMetadata; +pub use metadata_store::ConnectorMetadataStore; +pub use metadata_store::ConnectorToolSummary; +pub use plugin_config::parse_plugin_app_config; +pub use plugin_config::parse_plugin_app_config_value; +pub use runtime_projection::ConnectorRuntimeTool; +pub use runtime_projection::InstalledConnectorRuntime; +pub use runtime_projection::connector_tool_is_synthetic; +pub use runtime_projection::installed_connector_runtime; +pub use snapshot::ConnectorSnapshot; +pub use snapshot::PluginConnectorSource; + +pub const CONNECTORS_CACHE_TTL: Duration = Duration::from_secs(3600); +/// TTL for app/read metadata; it starts aligned with the connector directory cache. +pub const CONNECTOR_METADATA_CACHE_TTL: Duration = CONNECTORS_CACHE_TTL; + +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct ConnectorDirectoryCacheKey { + chatgpt_base_url: String, + account_id: Option, + chatgpt_user_id: Option, + is_workspace_account: bool, +} + +impl ConnectorDirectoryCacheKey { + pub fn new( + chatgpt_base_url: String, + account_id: Option, + chatgpt_user_id: Option, + is_workspace_account: bool, + ) -> Self { + Self { + chatgpt_base_url, + account_id, + chatgpt_user_id, + is_workspace_account, + } + } +} + +#[derive(Clone)] +struct CachedConnectorDirectory { + key: ConnectorDirectoryCacheKey, + expires_at: Instant, + connectors: Vec, +} + +static CONNECTOR_DIRECTORY_CACHE: LazyLock>> = + LazyLock::new(|| StdMutex::new(None)); + +#[derive(Debug, Deserialize)] +pub struct DirectoryListResponse { + apps: Vec, + #[serde(alias = "nextToken")] + next_token: Option, +} + +#[derive(Debug, Deserialize, Clone)] +pub struct DirectoryApp { + id: String, + name: String, + description: Option, + #[serde(alias = "appMetadata")] + app_metadata: Option, + branding: Option, + labels: Option>, + #[serde(alias = "logoUrl")] + logo_url: Option, + #[serde(alias = "logoUrlDark")] + logo_url_dark: Option, + #[serde(alias = "iconAssets")] + icon_assets: Option>, + #[serde(alias = "iconDarkAssets")] + icon_dark_assets: Option>, + #[serde(alias = "distributionChannel")] + distribution_channel: Option, + visibility: Option, +} + +pub fn cached_directory_connectors( + cache_context: &ConnectorDirectoryCacheContext, +) -> Option> { + if let Some(cached_connectors) = cached_directory_connectors_in_memory(&cache_context.cache_key) + { + return Some(cached_connectors); + } + + let directory_cache::CachedConnectorDirectoryDiskLoad::Hit { connectors } = + directory_cache::load_cached_directory_connectors_from_disk(cache_context) + else { + return None; + }; + write_cached_directory_connectors_in_memory( + cache_context.cache_key.clone(), + &connectors, + Duration::ZERO, + ); + Some(connectors) +} + +fn cached_directory_connectors_in_memory( + cache_key: &ConnectorDirectoryCacheKey, +) -> Option> { + let cache_guard = CONNECTOR_DIRECTORY_CACHE + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + cache_guard + .as_ref() + .filter(|cached| cached.key == *cache_key) + .map(|cached| cached.connectors.clone()) +} + +fn unexpired_directory_connectors_in_memory( + cache_key: &ConnectorDirectoryCacheKey, +) -> Option> { + let cache_guard = CONNECTOR_DIRECTORY_CACHE + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let cached = cache_guard.as_ref()?; + if cached.key == *cache_key && Instant::now() < cached.expires_at { + return Some(cached.connectors.clone()); + } + None +} + +pub async fn list_all_connectors_with_options( + cache_context: ConnectorDirectoryCacheContext, + is_workspace_account: bool, + force_refetch: bool, + mut fetch_page: F, +) -> anyhow::Result> +where + F: FnMut(String) -> Fut, + Fut: Future>, +{ + if !force_refetch + && let Some(cached_connectors) = + unexpired_directory_connectors_in_memory(&cache_context.cache_key) + { + return Ok(cached_connectors); + } + + let apps = if is_workspace_account { + // The workspace directory is independent from the paginated public directory. + // Start both before awaiting either so workspace accounts do not pay for the + // two request chains back-to-back. + let workspace_connectors = + fetch_page("/connectors/directory/list_workspace?external_logos=true".to_string()); + let directory_connectors = list_directory_connectors(&mut fetch_page); + let (directory_connectors, workspace_connectors) = + tokio::join!(directory_connectors, workspace_connectors); + let mut apps = directory_connectors?; + if let Ok(response) = workspace_connectors { + apps.extend( + response + .apps + .into_iter() + .filter(|app| !is_hidden_directory_app(app)), + ); + } + apps + } else { + list_directory_connectors(&mut fetch_page).await? + }; + + let mut connectors = merge_directory_apps(apps) + .into_iter() + .map(directory_app_to_app_info) + .collect::>(); + for connector in &mut connectors { + let install_url = match connector.install_url.take() { + Some(install_url) => install_url, + None => connector_install_url(&connector.name, &connector.id), + }; + connector.name = normalize_connector_name(&connector.name, &connector.id); + connector.description = normalize_connector_value(connector.description.as_deref()); + connector.install_url = Some(install_url); + connector.is_accessible = false; + } + connectors.sort_by(|left, right| { + left.name + .cmp(&right.name) + .then_with(|| left.id.cmp(&right.id)) + }); + write_cached_directory_connectors(&cache_context, &connectors); + Ok(connectors) +} + +fn write_cached_directory_connectors( + cache_context: &ConnectorDirectoryCacheContext, + connectors: &[AppInfo], +) { + write_cached_directory_connectors_in_memory( + cache_context.cache_key.clone(), + connectors, + CONNECTORS_CACHE_TTL, + ); + directory_cache::write_cached_directory_connectors_to_disk(cache_context, connectors); +} + +fn write_cached_directory_connectors_in_memory( + cache_key: ConnectorDirectoryCacheKey, + connectors: &[AppInfo], + ttl: Duration, +) { + let mut cache_guard = CONNECTOR_DIRECTORY_CACHE + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + *cache_guard = Some(CachedConnectorDirectory { + key: cache_key, + expires_at: Instant::now() + ttl, + connectors: connectors.to_vec(), + }); +} + +async fn list_directory_connectors(fetch_page: &mut F) -> anyhow::Result> +where + F: FnMut(String) -> Fut, + Fut: Future>, +{ + let mut apps = Vec::new(); + let mut next_token: Option = None; + loop { + let path = match next_token.as_deref() { + Some(token) => { + let encoded_token = urlencoding::encode(token); + format!("/connectors/directory/list?token={encoded_token}&external_logos=true") + } + None => "/connectors/directory/list?external_logos=true".to_string(), + }; + let response = fetch_page(path).await?; + apps.extend( + response + .apps + .into_iter() + .filter(|app| !is_hidden_directory_app(app)), + ); + next_token = response + .next_token + .map(|token| token.trim().to_string()) + .filter(|token| !token.is_empty()); + if next_token.is_none() { + break; + } + } + Ok(apps) +} + +fn merge_directory_apps(apps: Vec) -> Vec { + let mut merged: HashMap = HashMap::new(); + for app in apps { + if let Some(existing) = merged.get_mut(&app.id) { + merge_directory_app(existing, app); + } else { + merged.insert(app.id.clone(), app); + } + } + merged.into_values().collect() +} + +fn merge_directory_app(existing: &mut DirectoryApp, incoming: DirectoryApp) { + let DirectoryApp { + id: _, + name, + description, + app_metadata, + branding, + labels, + logo_url, + logo_url_dark, + icon_assets, + icon_dark_assets, + distribution_channel, + visibility: _, + } = incoming; + + let incoming_name_is_empty = name.trim().is_empty(); + if existing.name.trim().is_empty() && !incoming_name_is_empty { + existing.name = name; + } + + let incoming_description_present = description + .as_deref() + .map(|value| !value.trim().is_empty()) + .unwrap_or(false); + if incoming_description_present { + existing.description = description; + } + + if existing.logo_url.is_none() && logo_url.is_some() { + existing.logo_url = logo_url; + } + if existing.logo_url_dark.is_none() && logo_url_dark.is_some() { + existing.logo_url_dark = logo_url_dark; + } + if existing.icon_assets.as_ref().is_none_or(HashMap::is_empty) + && icon_assets + .as_ref() + .is_some_and(|assets| !assets.is_empty()) + { + existing.icon_assets = icon_assets; + } + if existing + .icon_dark_assets + .as_ref() + .is_none_or(HashMap::is_empty) + && icon_dark_assets + .as_ref() + .is_some_and(|assets| !assets.is_empty()) + { + existing.icon_dark_assets = icon_dark_assets; + } + if existing.distribution_channel.is_none() && distribution_channel.is_some() { + existing.distribution_channel = distribution_channel; + } + + if let Some(incoming_branding) = branding { + if let Some(existing_branding) = existing.branding.as_mut() { + if existing_branding.category.is_none() && incoming_branding.category.is_some() { + existing_branding.category = incoming_branding.category; + } + if existing_branding.developer.is_none() && incoming_branding.developer.is_some() { + existing_branding.developer = incoming_branding.developer; + } + if existing_branding.website.is_none() && incoming_branding.website.is_some() { + existing_branding.website = incoming_branding.website; + } + if existing_branding.privacy_policy.is_none() + && incoming_branding.privacy_policy.is_some() + { + existing_branding.privacy_policy = incoming_branding.privacy_policy; + } + if existing_branding.terms_of_service.is_none() + && incoming_branding.terms_of_service.is_some() + { + existing_branding.terms_of_service = incoming_branding.terms_of_service; + } + if !existing_branding.is_discoverable_app && incoming_branding.is_discoverable_app { + existing_branding.is_discoverable_app = true; + } + } else { + existing.branding = Some(incoming_branding); + } + } + + if let Some(incoming_app_metadata) = app_metadata { + if let Some(existing_app_metadata) = existing.app_metadata.as_mut() { + if existing_app_metadata.review.is_none() && incoming_app_metadata.review.is_some() { + existing_app_metadata.review = incoming_app_metadata.review; + } + if existing_app_metadata.categories.is_none() + && incoming_app_metadata.categories.is_some() + { + existing_app_metadata.categories = incoming_app_metadata.categories; + } + if existing_app_metadata.sub_categories.is_none() + && incoming_app_metadata.sub_categories.is_some() + { + existing_app_metadata.sub_categories = incoming_app_metadata.sub_categories; + } + if existing_app_metadata.seo_description.is_none() + && incoming_app_metadata.seo_description.is_some() + { + existing_app_metadata.seo_description = incoming_app_metadata.seo_description; + } + if existing_app_metadata.screenshots.is_none() + && incoming_app_metadata.screenshots.is_some() + { + existing_app_metadata.screenshots = incoming_app_metadata.screenshots; + } + if existing_app_metadata.developer.is_none() + && incoming_app_metadata.developer.is_some() + { + existing_app_metadata.developer = incoming_app_metadata.developer; + } + if existing_app_metadata.version.is_none() && incoming_app_metadata.version.is_some() { + existing_app_metadata.version = incoming_app_metadata.version; + } + if existing_app_metadata.version_id.is_none() + && incoming_app_metadata.version_id.is_some() + { + existing_app_metadata.version_id = incoming_app_metadata.version_id; + } + if existing_app_metadata.version_notes.is_none() + && incoming_app_metadata.version_notes.is_some() + { + existing_app_metadata.version_notes = incoming_app_metadata.version_notes; + } + if existing_app_metadata.first_party_requires_install.is_none() + && incoming_app_metadata.first_party_requires_install.is_some() + { + existing_app_metadata.first_party_requires_install = + incoming_app_metadata.first_party_requires_install; + } + if existing_app_metadata + .show_in_composer_when_unlinked + .is_none() + && incoming_app_metadata + .show_in_composer_when_unlinked + .is_some() + { + existing_app_metadata.show_in_composer_when_unlinked = + incoming_app_metadata.show_in_composer_when_unlinked; + } + } else { + existing.app_metadata = Some(incoming_app_metadata); + } + } + + if existing.labels.is_none() && labels.is_some() { + existing.labels = labels; + } +} + +fn is_hidden_directory_app(app: &DirectoryApp) -> bool { + matches!(app.visibility.as_deref(), Some("HIDDEN")) +} + +fn directory_app_to_app_info(app: DirectoryApp) -> AppInfo { + AppInfo { + id: app.id, + name: app.name, + description: app.description, + logo_url: app.logo_url, + logo_url_dark: app.logo_url_dark, + icon_assets: app.icon_assets, + icon_dark_assets: app.icon_dark_assets, + distribution_channel: app.distribution_channel, + branding: app.branding, + app_metadata: app.app_metadata, + labels: app.labels, + install_url: None, + is_accessible: false, + is_enabled: true, + plugin_display_names: Vec::new(), + } +} + +fn connector_install_url(name: &str, connector_id: &str) -> String { + let chatgpt_base_url = std::env::var("CODEX_APP_SERVER_CHATGPT_BASE_URL") + .unwrap_or_else(|_| "https://chatgpt.com".to_string()); + let chatgpt_origin = chatgpt_base_url + .trim_end_matches('/') + .trim_end_matches("/backend-api"); + let slug = connector_name_slug(name); + format!("{chatgpt_origin}/apps/{slug}/{connector_id}") +} + +fn connector_name_slug(name: &str) -> String { + let mut normalized = String::with_capacity(name.len()); + for character in name.chars() { + if character.is_ascii_alphanumeric() { + normalized.push(character.to_ascii_lowercase()); + } else { + normalized.push('-'); + } + } + let normalized = normalized.trim_matches('-'); + if normalized.is_empty() { + "app".to_string() + } else { + normalized.to_string() + } +} + +fn normalize_connector_name(name: &str, connector_id: &str) -> String { + let trimmed = name.trim(); + if trimmed.is_empty() { + connector_id.to_string() + } else { + trimmed.to_string() + } +} + +fn normalize_connector_value(value: Option<&str>) -> Option { + value + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) +} + +#[cfg(test)] +mod tests { + use super::*; + use pretty_assertions::assert_eq; + use std::sync::Arc; + use std::sync::Mutex; + use std::sync::atomic::AtomicUsize; + use std::sync::atomic::Ordering; + use std::time::Duration; + use tempfile::TempDir; + use tokio::sync::Notify; + + static CONNECTOR_DIRECTORY_CACHE_TEST_LOCK: LazyLock> = + LazyLock::new(|| tokio::sync::Mutex::new(())); + + fn cache_key(id: &str) -> ConnectorDirectoryCacheKey { + ConnectorDirectoryCacheKey::new( + "https://chatgpt.example".to_string(), + Some(format!("account-{id}")), + Some(format!("user-{id}")), + /*is_workspace_account*/ true, + ) + } + + fn cache_context(codex_home: &TempDir, id: &str) -> ConnectorDirectoryCacheContext { + ConnectorDirectoryCacheContext::new(codex_home.path().to_path_buf(), cache_key(id)) + } + + fn clear_directory_memory_cache() { + let mut cache_guard = CONNECTOR_DIRECTORY_CACHE + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + *cache_guard = None; + } + + fn app(id: &str, name: &str) -> DirectoryApp { + DirectoryApp { + id: id.to_string(), + name: name.to_string(), + description: None, + app_metadata: None, + branding: None, + labels: None, + logo_url: None, + logo_url_dark: None, + icon_assets: None, + icon_dark_assets: None, + distribution_channel: None, + visibility: None, + } + } + + #[test] + fn connector_install_url_uses_configured_origin() { + let chatgpt_base_url = std::env::var("CODEX_APP_SERVER_CHATGPT_BASE_URL") + .unwrap_or_else(|_| "https://chatgpt.com".to_string()); + let chatgpt_origin = chatgpt_base_url + .trim_end_matches('/') + .trim_end_matches("/backend-api"); + assert_eq!( + connector_install_url("Google Calendar", "calendar"), + format!("{chatgpt_origin}/apps/google-calendar/calendar"), + ); + } + + #[test] + fn directory_app_icon_assets_reach_app_info() -> anyhow::Result<()> { + let response: DirectoryListResponse = serde_json::from_value(serde_json::json!({ + "apps": [{ + "id": "alpha", + "name": "Alpha", + "icon_assets": {}, + "icon_dark_assets": {} + }, { + "id": "alpha", + "name": "", + "icon_assets": { + "256_square": "https://example.com/alpha-square.png" + }, + "icon_dark_assets": { + "256_square": "https://example.com/alpha-square-dark.png" + } + }], + "next_token": null + }))?; + + let app_info = directory_app_to_app_info(merge_directory_apps(response.apps).remove(0)); + + assert_eq!( + serde_json::to_value(app_info)?, + serde_json::json!({ + "id": "alpha", + "name": "Alpha", + "description": null, + "logoUrl": null, + "logoUrlDark": null, + "iconAssets": { + "256_square": "https://example.com/alpha-square.png" + }, + "iconDarkAssets": { + "256_square": "https://example.com/alpha-square-dark.png" + }, + "distributionChannel": null, + "branding": null, + "appMetadata": null, + "labels": null, + "installUrl": null, + "isAccessible": false, + "isEnabled": true, + "pluginDisplayNames": [] + }) + ); + Ok(()) + } + + #[tokio::test] + #[expect( + clippy::await_holding_invalid_type, + reason = "test serializes access to the shared connector cache for its full duration" + )] + async fn list_all_connectors_uses_shared_directory_cache() -> anyhow::Result<()> { + let _cache_guard = CONNECTOR_DIRECTORY_CACHE_TEST_LOCK.lock().await; + + let calls = Arc::new(AtomicUsize::new(0)); + let call_counter = Arc::clone(&calls); + let codex_home = TempDir::new()?; + let cache_context = cache_context(&codex_home, "shared"); + + let first = list_all_connectors_with_options( + cache_context.clone(), + /*is_workspace_account*/ false, + /*force_refetch*/ false, + move |_path| { + let call_counter = Arc::clone(&call_counter); + async move { + call_counter.fetch_add(1, Ordering::SeqCst); + Ok(DirectoryListResponse { + apps: vec![app("alpha", "Alpha")], + next_token: None, + }) + } + }, + ) + .await?; + + let second = list_all_connectors_with_options( + cache_context, + /*is_workspace_account*/ false, + /*force_refetch*/ false, + move |_path| async move { + anyhow::bail!("cache should have been used"); + }, + ) + .await?; + + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert_eq!(first, second); + Ok(()) + } + + #[tokio::test] + #[expect( + clippy::await_holding_invalid_type, + reason = "test serializes access to the shared connector cache for its full duration" + )] + async fn list_all_connectors_merges_and_normalizes_directory_apps() -> anyhow::Result<()> { + let _cache_guard = CONNECTOR_DIRECTORY_CACHE_TEST_LOCK.lock().await; + + let codex_home = TempDir::new()?; + let cache_context = cache_context(&codex_home, "merged"); + let calls = Arc::new(AtomicUsize::new(0)); + let call_counter = Arc::clone(&calls); + + let connectors = list_all_connectors_with_options( + cache_context, + /*is_workspace_account*/ true, + /*force_refetch*/ true, + move |path| { + let call_counter = Arc::clone(&call_counter); + async move { + call_counter.fetch_add(1, Ordering::SeqCst); + if path.starts_with("/connectors/directory/list_workspace") { + Ok(DirectoryListResponse { + apps: vec![ + DirectoryApp { + description: Some("Merged description".to_string()), + branding: Some(AppBranding { + category: Some("calendar".to_string()), + developer: None, + website: None, + privacy_policy: None, + terms_of_service: None, + is_discoverable_app: true, + }), + ..app("alpha", "") + }, + DirectoryApp { + visibility: Some("HIDDEN".to_string()), + ..app("hidden", "Hidden") + }, + ], + next_token: None, + }) + } else { + Ok(DirectoryListResponse { + apps: vec![app("alpha", " Alpha "), app("beta", "Beta")], + next_token: None, + }) + } + } + }, + ) + .await?; + + assert_eq!(calls.load(Ordering::SeqCst), 2); + assert_eq!(connectors.len(), 2); + assert_eq!(connectors[0].id, "alpha"); + assert_eq!(connectors[0].name, "Alpha"); + assert_eq!( + connectors[0].description.as_deref(), + Some("Merged description") + ); + assert_eq!( + connectors[0].install_url.as_deref(), + Some("https://chatgpt.com/apps/alpha/alpha") + ); + assert_eq!( + connectors[0] + .branding + .as_ref() + .and_then(|branding| branding.category.as_deref()), + Some("calendar") + ); + assert_eq!(connectors[1].id, "beta"); + assert_eq!(connectors[1].name, "Beta"); + Ok(()) + } + + #[tokio::test] + #[expect( + clippy::await_holding_invalid_type, + reason = "test serializes access to the shared connector cache for its full duration" + )] + async fn list_all_connectors_overlaps_workspace_and_directory_requests() -> anyhow::Result<()> { + let _cache_guard = CONNECTOR_DIRECTORY_CACHE_TEST_LOCK.lock().await; + + let codex_home = TempDir::new()?; + let cache_context = cache_context(&codex_home, "overlap"); + let workspace_started = Arc::new(Notify::new()); + + // The public directory response waits until the workspace request is polled. + // Without overlap this future cannot complete; the timeout only bounds a + // regression instead of supplying the ordering. + let connectors = tokio::time::timeout( + Duration::from_secs(1), + list_all_connectors_with_options( + cache_context, + /*is_workspace_account*/ true, + /*force_refetch*/ true, + move |path| { + let workspace_started = Arc::clone(&workspace_started); + async move { + if path.starts_with("/connectors/directory/list_workspace") { + workspace_started.notify_one(); + Ok(DirectoryListResponse { + apps: vec![app("workspace", "Workspace")], + next_token: None, + }) + } else { + workspace_started.notified().await; + Ok(DirectoryListResponse { + apps: vec![app("directory", "Directory")], + next_token: None, + }) + } + } + }, + ), + ) + .await + .expect("workspace request should start while directory request is pending")?; + + assert_eq!( + connectors + .into_iter() + .map(|connector| connector.id) + .collect::>(), + vec!["directory".to_string(), "workspace".to_string()] + ); + Ok(()) + } + + #[tokio::test] + #[expect( + clippy::await_holding_invalid_type, + reason = "test serializes access to the shared connector cache for its full duration" + )] + async fn cached_directory_connectors_reads_directory_disk_cache() -> anyhow::Result<()> { + let _cache_guard = CONNECTOR_DIRECTORY_CACHE_TEST_LOCK.lock().await; + + let codex_home = TempDir::new()?; + let cache_context = cache_context(&codex_home, "disk"); + let calls = Arc::new(AtomicUsize::new(0)); + let call_counter = Arc::clone(&calls); + + let first = list_all_connectors_with_options( + cache_context.clone(), + /*is_workspace_account*/ false, + /*force_refetch*/ false, + move |_path| { + let call_counter = Arc::clone(&call_counter); + async move { + call_counter.fetch_add(1, Ordering::SeqCst); + Ok(DirectoryListResponse { + apps: vec![app("alpha", "Alpha")], + next_token: None, + }) + } + }, + ) + .await?; + + clear_directory_memory_cache(); + + let second = cached_directory_connectors(&cache_context).expect("disk cache should load"); + + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert_eq!(first, second); + Ok(()) + } + + #[tokio::test] + #[expect( + clippy::await_holding_invalid_type, + reason = "test serializes access to the shared connector cache for its full duration" + )] + async fn list_all_connectors_refreshes_when_only_directory_disk_cache_exists() + -> anyhow::Result<()> { + let _cache_guard = CONNECTOR_DIRECTORY_CACHE_TEST_LOCK.lock().await; + + let codex_home = TempDir::new()?; + let cache_context = cache_context(&codex_home, "disk-refresh"); + let calls = Arc::new(AtomicUsize::new(0)); + let call_counter = Arc::clone(&calls); + + list_all_connectors_with_options( + cache_context.clone(), + /*is_workspace_account*/ false, + /*force_refetch*/ false, + move |_path| { + let call_counter = Arc::clone(&call_counter); + async move { + call_counter.fetch_add(1, Ordering::SeqCst); + Ok(DirectoryListResponse { + apps: vec![app("alpha", "Alpha")], + next_token: None, + }) + } + }, + ) + .await?; + + clear_directory_memory_cache(); + let mut cached_expected = directory_app_to_app_info(app("alpha", "Alpha")); + cached_expected.install_url = Some(connector_install_url( + &cached_expected.name, + &cached_expected.id, + )); + assert_eq!( + cached_directory_connectors(&cache_context), + Some(vec![cached_expected]) + ); + let refreshed_calls = Arc::clone(&calls); + + let refreshed = list_all_connectors_with_options( + cache_context, + /*is_workspace_account*/ false, + /*force_refetch*/ false, + move |_path| { + let call_counter = Arc::clone(&refreshed_calls); + async move { + call_counter.fetch_add(1, Ordering::SeqCst); + Ok(DirectoryListResponse { + apps: vec![app("beta", "Beta")], + next_token: None, + }) + } + }, + ) + .await?; + + let mut expected = directory_app_to_app_info(app("beta", "Beta")); + expected.install_url = Some(connector_install_url(&expected.name, &expected.id)); + assert_eq!(calls.load(Ordering::SeqCst), 2); + assert_eq!(refreshed, vec![expected]); + Ok(()) + } + + #[tokio::test] + async fn cached_directory_connectors_drops_stale_disk_schema() -> anyhow::Result<()> { + let _cache_guard = CONNECTOR_DIRECTORY_CACHE_TEST_LOCK.lock().await; + + clear_directory_memory_cache(); + let codex_home = TempDir::new()?; + let cache_context = cache_context(&codex_home, "stale-schema"); + let cache_path = cache_context.cache_path(); + std::fs::create_dir_all(cache_path.parent().expect("cache parent"))?; + std::fs::write( + &cache_path, + serde_json::to_vec_pretty(&serde_json::json!({ + "schema_version": 0, + "connectors": [], + }))?, + )?; + + assert_eq!(cached_directory_connectors(&cache_context), None); + assert!(!cache_path.exists()); + Ok(()) + } + + #[tokio::test] + async fn list_directory_connectors_omits_tier_for_all_pages() -> anyhow::Result<()> { + let requested_paths: Arc>> = Arc::new(Mutex::new(Vec::new())); + let paths = Arc::clone(&requested_paths); + + let apps = list_directory_connectors(&mut move |path| { + let paths = Arc::clone(&paths); + async move { + paths + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .push(path.clone()); + if path == "/connectors/directory/list?external_logos=true" { + Ok(DirectoryListResponse { + apps: vec![app("alpha", "Alpha")], + next_token: Some("page 2".to_string()), + }) + } else { + assert_eq!( + path, + "/connectors/directory/list?token=page%202&external_logos=true" + ); + Ok(DirectoryListResponse { + apps: vec![app("beta", "Beta")], + next_token: None, + }) + } + } + }) + .await?; + + assert_eq!( + apps.iter().map(|app| app.id.as_str()).collect::>(), + vec!["alpha", "beta"] + ); + assert_eq!( + requested_paths + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .as_slice(), + &[ + "/connectors/directory/list?external_logos=true".to_string(), + "/connectors/directory/list?token=page%202&external_logos=true".to_string(), + ] + ); + Ok(()) + } +} diff --git a/codex-rs/connectors/src/merge.rs b/codex-rs/connectors/src/merge.rs new file mode 100644 index 0000000000000000000000000000000000000000..9f906afc50e42ff7de822a8202dcf725b32a29e9 --- /dev/null +++ b/codex-rs/connectors/src/merge.rs @@ -0,0 +1,220 @@ +use std::collections::HashMap; +use std::collections::HashSet; + +use crate::AppInfo; +use crate::metadata::connector_install_url; +use crate::metadata::sort_connectors_by_accessibility_and_name; + +pub fn merge_connectors( + connectors: Vec, + accessible_connectors: Vec, +) -> Vec { + let mut merged: HashMap = connectors + .into_iter() + .map(|mut connector| { + connector.is_accessible = false; + (connector.id.clone(), connector) + }) + .collect(); + + for mut connector in accessible_connectors { + connector.is_accessible = true; + let connector_id = connector.id.clone(); + if let Some(existing) = merged.get_mut(&connector_id) { + existing.is_accessible = true; + if existing.name == existing.id && connector.name != connector.id { + existing.name = connector.name; + } + if existing.description.is_none() && connector.description.is_some() { + existing.description = connector.description; + } + if existing.logo_url.is_none() && connector.logo_url.is_some() { + existing.logo_url = connector.logo_url; + } + if existing.logo_url_dark.is_none() && connector.logo_url_dark.is_some() { + existing.logo_url_dark = connector.logo_url_dark; + } + if existing.icon_assets.is_none() && connector.icon_assets.is_some() { + existing.icon_assets = connector.icon_assets; + } + if existing.icon_dark_assets.is_none() && connector.icon_dark_assets.is_some() { + existing.icon_dark_assets = connector.icon_dark_assets; + } + if existing.distribution_channel.is_none() && connector.distribution_channel.is_some() { + existing.distribution_channel = connector.distribution_channel; + } + existing + .plugin_display_names + .extend(connector.plugin_display_names); + } else { + merged.insert(connector_id, connector); + } + } + + let mut merged = merged.into_values().collect::>(); + for connector in &mut merged { + if connector.install_url.is_none() { + connector.install_url = Some(connector_install_url(&connector.name, &connector.id)); + } + connector.plugin_display_names.sort_unstable(); + connector.plugin_display_names.dedup(); + } + sort_connectors_by_accessibility_and_name(&mut merged); + merged +} + +pub fn merge_plugin_connectors(connectors: Vec, plugin_app_ids: I) -> Vec +where + I: IntoIterator, +{ + let mut merged = connectors; + let mut connector_ids = merged + .iter() + .map(|connector| connector.id.clone()) + .collect::>(); + + for connector_id in plugin_app_ids { + if connector_ids.insert(connector_id.clone()) { + merged.push(plugin_connector_to_app_info(connector_id)); + } + } + + sort_connectors_by_accessibility_and_name(&mut merged); + merged +} + +pub fn merge_plugin_connectors_with_accessible( + plugin_app_ids: I, + accessible_connectors: Vec, +) -> Vec +where + I: IntoIterator, +{ + let accessible_connector_ids: HashSet<&str> = accessible_connectors + .iter() + .map(|connector| connector.id.as_str()) + .collect(); + let plugin_connectors = plugin_app_ids + .into_iter() + .filter(|connector_id| accessible_connector_ids.contains(connector_id.as_str())) + .map(plugin_connector_to_app_info) + .collect::>(); + merge_connectors(plugin_connectors, accessible_connectors) +} + +pub fn plugin_connector_to_app_info(connector_id: String) -> AppInfo { + // Leave the placeholder name as the connector id so merge_connectors() can + // replace it with canonical app metadata from directory fetches or + // connector_name values from codex_apps tool discovery. + let name = connector_id.clone(); + AppInfo { + id: connector_id.clone(), + name: name.clone(), + description: None, + logo_url: None, + logo_url_dark: None, + icon_assets: None, + icon_dark_assets: None, + distribution_channel: None, + branding: None, + app_metadata: None, + labels: None, + install_url: Some(connector_install_url(&name, &connector_id)), + is_accessible: false, + is_enabled: true, + plugin_display_names: Vec::new(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::metadata::connector_install_url; + use crate::metadata::connector_mention_slug; + use pretty_assertions::assert_eq; + + fn plugin_names(names: &[&str]) -> Vec { + names.iter().map(ToString::to_string).collect() + } + + fn google_calendar_accessible_connector(plugin_display_names: &[&str]) -> AppInfo { + AppInfo { + id: "calendar".to_string(), + name: "Google Calendar".to_string(), + description: Some("Plan events".to_string()), + logo_url: Some("https://example.com/logo.png".to_string()), + logo_url_dark: Some("https://example.com/logo-dark.png".to_string()), + icon_assets: None, + icon_dark_assets: None, + distribution_channel: Some("workspace".to_string()), + branding: None, + app_metadata: None, + labels: None, + install_url: None, + is_accessible: true, + is_enabled: true, + plugin_display_names: plugin_names(plugin_display_names), + } + } + + #[test] + fn merge_connectors_replaces_plugin_placeholder_name_with_accessible_name() { + let plugin = plugin_connector_to_app_info("calendar".to_string()); + let accessible = google_calendar_accessible_connector(&[]); + + let merged = merge_connectors(vec![plugin], vec![accessible]); + + assert_eq!( + merged, + vec![AppInfo { + id: "calendar".to_string(), + name: "Google Calendar".to_string(), + description: Some("Plan events".to_string()), + logo_url: Some("https://example.com/logo.png".to_string()), + logo_url_dark: Some("https://example.com/logo-dark.png".to_string()), + icon_assets: None, + icon_dark_assets: None, + distribution_channel: Some("workspace".to_string()), + branding: None, + app_metadata: None, + labels: None, + install_url: Some(connector_install_url("calendar", "calendar")), + is_accessible: true, + is_enabled: true, + plugin_display_names: Vec::new(), + }] + ); + assert_eq!(connector_mention_slug(&merged[0]), "google-calendar"); + } + + #[test] + fn merge_connectors_unions_and_dedupes_plugin_display_names() { + let mut plugin = plugin_connector_to_app_info("calendar".to_string()); + plugin.plugin_display_names = plugin_names(&["sample", "alpha", "sample"]); + + let accessible = google_calendar_accessible_connector(&["beta", "alpha"]); + + let merged = merge_connectors(vec![plugin], vec![accessible]); + + assert_eq!( + merged, + vec![AppInfo { + id: "calendar".to_string(), + name: "Google Calendar".to_string(), + description: Some("Plan events".to_string()), + logo_url: Some("https://example.com/logo.png".to_string()), + logo_url_dark: Some("https://example.com/logo-dark.png".to_string()), + icon_assets: None, + icon_dark_assets: None, + distribution_channel: Some("workspace".to_string()), + branding: None, + app_metadata: None, + labels: None, + install_url: Some(connector_install_url("calendar", "calendar")), + is_accessible: true, + is_enabled: true, + plugin_display_names: plugin_names(&["alpha", "beta", "sample"]), + }] + ); + } +} diff --git a/codex-rs/connectors/src/metadata.rs b/codex-rs/connectors/src/metadata.rs new file mode 100644 index 0000000000000000000000000000000000000000..9deeabf04114b0ddf826a18158e60f5891f45545 --- /dev/null +++ b/codex-rs/connectors/src/metadata.rs @@ -0,0 +1,31 @@ +use crate::AppInfo; + +pub fn connector_display_label(connector: &AppInfo) -> String { + connector.name.clone() +} + +pub fn connector_mention_slug(connector: &AppInfo) -> String { + connector_mention_slug_from_name(&connector_display_label(connector)) +} + +pub fn connector_mention_slug_from_name(name: &str) -> String { + crate::connector_name_slug(name) +} + +pub fn connector_install_url(name: &str, connector_id: &str) -> String { + crate::connector_install_url(name, connector_id) +} + +pub fn sanitize_name(name: &str) -> String { + crate::connector_name_slug(name).replace("-", "_") +} + +pub(crate) fn sort_connectors_by_accessibility_and_name(connectors: &mut [AppInfo]) { + connectors.sort_by(|left, right| { + right + .is_accessible + .cmp(&left.is_accessible) + .then_with(|| left.name.cmp(&right.name)) + .then_with(|| left.id.cmp(&right.id)) + }); +} diff --git a/codex-rs/connectors/src/metadata_store.rs b/codex-rs/connectors/src/metadata_store.rs new file mode 100644 index 0000000000000000000000000000000000000000..c8a6449ac559e18606aba858e0c6740ba0f898a3 --- /dev/null +++ b/codex-rs/connectors/src/metadata_store.rs @@ -0,0 +1,144 @@ +use std::collections::HashMap; +use std::sync::LazyLock; +use std::sync::Mutex as StdMutex; +use std::time::Instant; + +use crate::CONNECTOR_METADATA_CACHE_TTL; + +/// Display-only summary of one app tool returned by the app batch-read API. +#[derive(Debug, Clone, PartialEq)] +pub struct ConnectorToolSummary { + pub name: String, + pub title: Option, + pub description: String, + pub is_enabled: bool, + pub disabled_reason: Option, + pub is_read_only: bool, +} + +/// Metadata returned by the app batch-read API. +/// +/// This intentionally excludes connector runtime state, full actions, and model descriptions. +/// Tool summaries contain display text and enabled/read-only state only, and icon URLs are already +/// projected as public URLs by the backend. +#[derive(Debug, Clone, PartialEq)] +pub struct ConnectorMetadata { + pub id: String, + pub name: String, + pub description: Option, + pub icon_url: Option, + pub icon_url_dark: Option, + pub distribution_channel: Option, + pub tool_summaries: Option>, +} + +/// A view of the process-wide metadata cache bound to one backend and auth identity. +/// +/// The active ChatGPT account id represents the selected personal account or workspace, while the +/// ChatGPT user id identifies the account principal. Keeping both plus workspace classification +/// matches the existing connector-directory cache partition. +pub struct ConnectorMetadataStore { + scope: ConnectorMetadataStoreScope, +} + +impl ConnectorMetadataStore { + pub fn new( + backend_base_url: String, + account_id: Option, + chatgpt_user_id: Option, + is_workspace_account: bool, + ) -> Self { + Self { + scope: ConnectorMetadataStoreScope { + backend_base_url, + account_id, + chatgpt_user_id, + is_workspace_account, + }, + } + } + + /// Returns only unexpired records for the requested ids, requiring tool summaries when asked. + /// + /// Expired entries are deliberately left in place so a failed refresh cannot mutate prior + /// cache state. + pub fn fresh_records( + &self, + ids: &[String], + include_tools: bool, + ) -> HashMap { + let cache = CONNECTOR_METADATA_CACHE + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let Some(records) = cache.get(&self.scope) else { + return HashMap::new(); + }; + let now = Instant::now(); + ids.iter() + .filter_map(|id| { + records + .get(id) + .filter(|record| { + now < record.expires_at + && (!include_tools || record.metadata.tool_summaries.is_some()) + }) + .map(|record| (id.clone(), record.metadata.clone())) + }) + .collect() + } + + /// Commits successfully fetched records without letting a late metadata-only response + /// replace fresh tool summaries. + pub fn commit(&self, records: &[ConnectorMetadata]) { + if records.is_empty() { + return; + } + + let now = Instant::now(); + let expires_at = now + CONNECTOR_METADATA_CACHE_TTL; + let mut cache = CONNECTOR_METADATA_CACHE + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + let scoped_records = cache.entry(self.scope.clone()).or_default(); + for metadata in records { + if metadata.tool_summaries.is_none() + && scoped_records.get(&metadata.id).is_some_and(|record| { + now < record.expires_at && record.metadata.tool_summaries.is_some() + }) + { + continue; + } + scoped_records.insert( + metadata.id.clone(), + CachedConnectorMetadata { + metadata: metadata.clone(), + expires_at, + }, + ); + } + } +} + +// `apps_mcp_product_sku` affects which tools the batch API returns, but is intentionally omitted +// from this key because we assume an app-server does not change its product SKU after launch. +// If that assumption changes, the SKU must be included in the cache scope. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +struct ConnectorMetadataStoreScope { + backend_base_url: String, + account_id: Option, + chatgpt_user_id: Option, + is_workspace_account: bool, +} + +struct CachedConnectorMetadata { + metadata: ConnectorMetadata, + expires_at: Instant, +} + +static CONNECTOR_METADATA_CACHE: LazyLock< + StdMutex>>, +> = LazyLock::new(|| StdMutex::new(HashMap::new())); + +#[cfg(test)] +#[path = "metadata_store_tests.rs"] +mod tests; diff --git a/codex-rs/connectors/src/metadata_store_tests.rs b/codex-rs/connectors/src/metadata_store_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..77dbdd208cb5d64129c2c5f6ff43294c668ae63b --- /dev/null +++ b/codex-rs/connectors/src/metadata_store_tests.rs @@ -0,0 +1,152 @@ +use pretty_assertions::assert_eq; + +use super::ConnectorMetadata; +use super::ConnectorMetadataStore; +use super::ConnectorToolSummary; + +fn metadata(id: &str) -> ConnectorMetadata { + ConnectorMetadata { + id: id.to_string(), + name: format!("{id} name"), + description: None, + icon_url: None, + icon_url_dark: None, + distribution_channel: None, + tool_summaries: None, + } +} + +#[test] +fn records_are_isolated_by_backend_account_user_and_workspace_scope() { + let requested_scope = ConnectorMetadataStore::new( + "https://backend-a.example".to_string(), + Some("account-a".to_string()), + Some("user-a".to_string()), + /*is_workspace_account*/ true, + ); + let other_backend = ConnectorMetadataStore::new( + "https://backend-b.example".to_string(), + Some("account-a".to_string()), + Some("user-a".to_string()), + /*is_workspace_account*/ true, + ); + let other_account = ConnectorMetadataStore::new( + "https://backend-a.example".to_string(), + Some("account-b".to_string()), + Some("user-a".to_string()), + /*is_workspace_account*/ true, + ); + let other_user = ConnectorMetadataStore::new( + "https://backend-a.example".to_string(), + Some("account-a".to_string()), + Some("user-b".to_string()), + /*is_workspace_account*/ true, + ); + let personal_account = ConnectorMetadataStore::new( + "https://backend-a.example".to_string(), + Some("account-a".to_string()), + Some("user-a".to_string()), + /*is_workspace_account*/ false, + ); + let ids = vec!["scoped-app".to_string()]; + + requested_scope.commit(&[metadata("scoped-app")]); + + assert_eq!( + requested_scope.fresh_records(&ids, /*include_tools*/ false), + std::collections::HashMap::from([("scoped-app".to_string(), metadata("scoped-app"))]) + ); + assert_eq!( + other_backend.fresh_records(&ids, /*include_tools*/ false), + Default::default() + ); + assert_eq!( + other_account.fresh_records(&ids, /*include_tools*/ false), + Default::default() + ); + assert_eq!( + other_user.fresh_records(&ids, /*include_tools*/ false), + Default::default() + ); + assert_eq!( + personal_account.fresh_records(&ids, /*include_tools*/ false), + Default::default() + ); +} + +#[test] +fn tool_inclusive_reads_require_cached_tool_summaries() { + let store = ConnectorMetadataStore::new( + "https://backend-tools.example".to_string(), + Some("account-tools".to_string()), + Some("user-tools".to_string()), + /*is_workspace_account*/ false, + ); + let metadata_only = metadata("metadata-only"); + let mut empty_tools = metadata("empty-tools"); + empty_tools.tool_summaries = Some(Vec::new()); + let mut with_tools = metadata("with-tools"); + with_tools.tool_summaries = Some(vec![ConnectorToolSummary { + name: "search".to_string(), + title: Some("Search".to_string()), + description: "Search the app".to_string(), + is_enabled: true, + disabled_reason: None, + is_read_only: true, + }]); + let ids = vec![ + "metadata-only".to_string(), + "empty-tools".to_string(), + "with-tools".to_string(), + ]; + + store.commit(&[ + metadata_only.clone(), + empty_tools.clone(), + with_tools.clone(), + ]); + + assert_eq!( + store.fresh_records(&ids, /*include_tools*/ false), + std::collections::HashMap::from([ + ("metadata-only".to_string(), metadata_only), + ("empty-tools".to_string(), empty_tools.clone()), + ("with-tools".to_string(), with_tools.clone()), + ]) + ); + assert_eq!( + store.fresh_records(&ids, /*include_tools*/ true), + std::collections::HashMap::from([ + ("empty-tools".to_string(), empty_tools), + ("with-tools".to_string(), with_tools), + ]) + ); +} + +#[test] +fn metadata_only_commit_does_not_replace_fresh_tool_summaries() { + let store = ConnectorMetadataStore::new( + "https://backend-tools-race.example".to_string(), + Some("account-tools-race".to_string()), + Some("user-tools-race".to_string()), + /*is_workspace_account*/ false, + ); + let mut with_tools = metadata("with-tools"); + with_tools.tool_summaries = Some(vec![ConnectorToolSummary { + name: "search".to_string(), + title: Some("Search".to_string()), + description: "Search the app".to_string(), + is_enabled: true, + disabled_reason: None, + is_read_only: true, + }]); + let ids = vec!["with-tools".to_string()]; + + store.commit(&[with_tools.clone()]); + store.commit(&[metadata("with-tools")]); + + assert_eq!( + store.fresh_records(&ids, /*include_tools*/ true), + std::collections::HashMap::from([("with-tools".to_string(), with_tools)]) + ); +} diff --git a/codex-rs/connectors/src/plugin_config.rs b/codex-rs/connectors/src/plugin_config.rs new file mode 100644 index 0000000000000000000000000000000000000000..6f3179be8359b904cfffe08a16750afe4453f48b --- /dev/null +++ b/codex-rs/connectors/src/plugin_config.rs @@ -0,0 +1,50 @@ +use codex_plugin::AppConnectorId; +use codex_plugin::AppDeclaration; +use indexmap::IndexMap; +use serde::Deserialize; +use serde_json::Value; + +#[derive(Debug, Default, Deserialize)] +#[serde(rename_all = "camelCase")] +struct PluginAppFile { + #[serde(default)] + apps: IndexMap, +} + +#[derive(Debug, Default, Deserialize)] +struct PluginAppConfig { + id: String, + category: Option, +} + +/// Parses connector declarations from a plugin app configuration file. +pub fn parse_plugin_app_config(contents: &str) -> serde_json::Result> { + serde_json::from_str(contents).map(app_declarations_from_file) +} + +/// Parses connector declarations from an already-decoded plugin app configuration. +pub fn parse_plugin_app_config_value(value: Value) -> serde_json::Result> { + serde_json::from_value(value).map(app_declarations_from_file) +} + +fn app_declarations_from_file(parsed: PluginAppFile) -> Vec { + parsed + .apps + .into_iter() + .map(|(name, app)| AppDeclaration { + name, + connector_id: AppConnectorId(app.id), + category: cleaned_category(app.category), + }) + .collect() +} + +fn cleaned_category(category: Option) -> Option { + category + .map(|category| category.trim().to_string()) + .filter(|category| !category.is_empty()) +} + +#[cfg(test)] +#[path = "plugin_config_tests.rs"] +mod tests; diff --git a/codex-rs/connectors/src/plugin_config_tests.rs b/codex-rs/connectors/src/plugin_config_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..60944ce76d2781ff475ff6790a2c3133d0f068dd --- /dev/null +++ b/codex-rs/connectors/src/plugin_config_tests.rs @@ -0,0 +1,53 @@ +use codex_plugin::AppConnectorId; +use codex_plugin::AppDeclaration; +use pretty_assertions::assert_eq; + +use super::parse_plugin_app_config; + +#[test] +fn parses_plugin_app_config_in_order_without_validating_connector_ids() { + let parsed = parse_plugin_app_config( + r#"{ + "apps": { + "calendar": { + "id": "connector_calendar", + "category": " productivity " + }, + "drive": { + "id": "connector_calendar", + "category": " " + }, + "blank": { + "id": " " + } + } + }"#, + ) + .expect("plugin app config should parse"); + + assert_eq!( + parsed, + vec![ + AppDeclaration { + name: "calendar".to_string(), + connector_id: AppConnectorId("connector_calendar".to_string()), + category: Some("productivity".to_string()), + }, + AppDeclaration { + name: "drive".to_string(), + connector_id: AppConnectorId("connector_calendar".to_string()), + category: None, + }, + AppDeclaration { + name: "blank".to_string(), + connector_id: AppConnectorId(" ".to_string()), + category: None, + }, + ] + ); +} + +#[test] +fn rejects_invalid_plugin_app_config() { + assert!(parse_plugin_app_config("not json").is_err()); +} diff --git a/codex-rs/connectors/src/runtime_projection.rs b/codex-rs/connectors/src/runtime_projection.rs new file mode 100644 index 0000000000000000000000000000000000000000..5e8f100dd0b564966e8d94eb278170b3bafbdd7b --- /dev/null +++ b/codex-rs/connectors/src/runtime_projection.rs @@ -0,0 +1,103 @@ +//! Connector-owned projection of raw runtime tools into installed app state. + +use std::collections::BTreeMap; + +use codex_config::ConfigLayerStack; + +use crate::AppToolPolicyEvaluator; +use crate::AppToolPolicyInput; + +/// Connector-relevant fields from one runtime tool. +/// +/// MCP owns the raw tool type and computes generic visibility/filter decisions. Connector +/// consumers adapt those fields into this view so connector policy stays out of MCP modules. +#[derive(Debug, Clone, Copy)] +pub struct ConnectorRuntimeTool<'a> { + pub connector_id: Option<&'a str>, + pub connector_name: Option<&'a str>, + pub tool_name: &'a str, + pub tool_title: Option<&'a str>, + pub destructive_hint: Option, + pub open_world_hint: Option, + pub synthetic: bool, + pub model_visible: bool, +} + +/// Installed state derived from one committed connector runtime snapshot. +/// +/// `enabled` and `callable` include local and managed app/tool configuration. Global feature and +/// workspace policy remain host concerns and are applied by the caller. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct InstalledConnectorRuntime { + pub id: String, + pub runtime_name: Option, + pub enabled: bool, + pub callable: bool, +} + +/// Projects raw runtime tools into one row per installed connector. +pub fn installed_connector_runtime<'a>( + config_layer_stack: &ConfigLayerStack, + tools: impl IntoIterator>, +) -> Vec { + let policy = AppToolPolicyEvaluator::new(config_layer_stack); + let mut apps = BTreeMap::, bool)>::new(); + + for tool in tools { + if tool.synthetic { + continue; + } + let Some(connector_id) = tool.connector_id.map(str::trim) else { + continue; + }; + if connector_id.is_empty() { + continue; + } + + let runtime_name = tool + .connector_name + .map(str::trim) + .filter(|name| !name.is_empty()) + .map(str::to_string); + let entry = apps + .entry(connector_id.to_string()) + .or_insert((None, false)); + if entry.0.is_none() { + entry.0 = runtime_name; + } + + let policy_allows_tool = policy + .policy(AppToolPolicyInput { + connector_id: Some(connector_id), + link_id: None, + tool_name: tool.tool_name, + tool_title: tool.tool_title, + destructive_hint: tool.destructive_hint, + open_world_hint: tool.open_world_hint, + }) + .enabled; + entry.1 |= tool.model_visible && policy_allows_tool; + } + + apps.into_iter() + .map(|(id, (runtime_name, callable))| InstalledConnectorRuntime { + enabled: policy.app_enabled(&id), + id, + runtime_name, + callable, + }) + .collect() +} + +/// Returns whether connector metadata marks a runtime tool as a synthetic link helper. +pub fn connector_tool_is_synthetic(connector_meta: Option<&serde_json::Value>) -> bool { + connector_meta + .and_then(serde_json::Value::as_object) + .and_then(|meta| meta.get("synthetic_link")) + .and_then(serde_json::Value::as_bool) + == Some(true) +} + +#[cfg(test)] +#[path = "runtime_projection_tests.rs"] +mod tests; diff --git a/codex-rs/connectors/src/runtime_projection_tests.rs b/codex-rs/connectors/src/runtime_projection_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..9438b4f6fcf634cea5c26e0a43bc11ddbad06bfe --- /dev/null +++ b/codex-rs/connectors/src/runtime_projection_tests.rs @@ -0,0 +1,113 @@ +use std::collections::BTreeMap; + +use codex_config::AppRequirementToml; +use codex_config::AppsRequirementsToml; +use codex_config::ConfigLayerStack; +use codex_config::ConfigRequirements; +use codex_config::ConfigRequirementsToml; +use pretty_assertions::assert_eq; + +use super::*; + +#[test] +fn projection_deduplicates_apps_and_ignores_non_runtime_tools() { + let config = ConfigLayerStack::new( + Vec::new(), + ConfigRequirements::default(), + ConfigRequirementsToml::default(), + ) + .expect("config layer stack"); + let apps = installed_connector_runtime( + &config, + [ + tool(Some(" drive "), /*connector_name*/ None, "files/list"), + tool(Some("drive"), Some(" Drive "), "files/get"), + ConnectorRuntimeTool { + synthetic: true, + ..tool(Some("synthetic"), Some("Synthetic"), "link") + }, + tool(Some(" "), Some("Empty"), "empty"), + tool(/*connector_id*/ None, Some("Missing"), "missing"), + ], + ); + + assert_eq!( + apps, + vec![InstalledConnectorRuntime { + id: "drive".to_string(), + runtime_name: Some("Drive".to_string()), + enabled: true, + callable: true, + }] + ); +} + +#[test] +fn projection_applies_managed_app_policy_and_model_visibility() { + let requirements = ConfigRequirementsToml { + apps: Some(AppsRequirementsToml { + apps: BTreeMap::from([( + "disabled".to_string(), + AppRequirementToml { + enabled: Some(false), + tools: None, + }, + )]), + }), + ..Default::default() + }; + let config = ConfigLayerStack::new(Vec::new(), ConfigRequirements::default(), requirements) + .expect("config layer stack"); + let apps = installed_connector_runtime( + &config, + [ + tool(Some("disabled"), Some("Disabled"), "disabled/tool"), + ConnectorRuntimeTool { + model_visible: false, + ..tool(Some("hidden"), Some("Hidden"), "hidden/tool") + }, + tool(Some("callable"), Some("Callable"), "callable/tool"), + ], + ); + + assert_eq!( + apps, + vec![ + InstalledConnectorRuntime { + id: "callable".to_string(), + runtime_name: Some("Callable".to_string()), + enabled: true, + callable: true, + }, + InstalledConnectorRuntime { + id: "disabled".to_string(), + runtime_name: Some("Disabled".to_string()), + enabled: false, + callable: false, + }, + InstalledConnectorRuntime { + id: "hidden".to_string(), + runtime_name: Some("Hidden".to_string()), + enabled: true, + callable: false, + }, + ] + ); +} + +fn tool<'a>( + connector_id: Option<&'a str>, + connector_name: Option<&'a str>, + tool_name: &'a str, +) -> ConnectorRuntimeTool<'a> { + ConnectorRuntimeTool { + connector_id, + connector_name, + tool_name, + tool_title: None, + destructive_hint: None, + open_world_hint: None, + synthetic: false, + model_visible: true, + } +} diff --git a/codex-rs/connectors/src/snapshot.rs b/codex-rs/connectors/src/snapshot.rs new file mode 100644 index 0000000000000000000000000000000000000000..4129f03988cc0fb88c4d6fdce85a974ffc661d8e --- /dev/null +++ b/codex-rs/connectors/src/snapshot.rs @@ -0,0 +1,155 @@ +use std::collections::HashMap; +use std::collections::HashSet; + +use codex_plugin::AppConnectorId; +use codex_plugin::AppDeclaration; +use codex_plugin::PluginCapabilitySummary; + +/// Connector declarations contributed by one plugin package. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct PluginConnectorSource { + plugin_id: String, + plugin_display_name: String, + connector_ids: Vec, +} + +impl PluginConnectorSource { + /// Creates one plugin source from parsed app declarations. + pub fn new( + plugin_id: impl Into, + plugin_display_name: impl Into, + declarations: impl IntoIterator, + ) -> Self { + Self::from_connector_ids( + plugin_id, + plugin_display_name, + declarations + .into_iter() + .map(|declaration| declaration.connector_id), + ) + } + + /// Creates one plugin source from connector IDs that were already parsed. + pub fn from_connector_ids( + plugin_id: impl Into, + plugin_display_name: impl Into, + connector_ids: impl IntoIterator, + ) -> Self { + let mut seen_connector_ids = HashSet::new(); + let connector_ids = connector_ids + .into_iter() + .filter(|connector_id| !connector_id.0.trim().is_empty()) + .filter(|connector_id| seen_connector_ids.insert(connector_id.clone())) + .collect(); + Self { + plugin_id: plugin_id.into(), + plugin_display_name: plugin_display_name.into(), + connector_ids, + } + } + + /// Returns the package name shown in connector provenance. + pub fn plugin_display_name(&self) -> &str { + &self.plugin_display_name + } + + /// Returns the connector IDs contributed by this package. + pub fn connector_ids(&self) -> &[AppConnectorId] { + &self.connector_ids + } +} + +impl From<&PluginCapabilitySummary> for PluginConnectorSource { + fn from(summary: &PluginCapabilitySummary) -> Self { + Self::from_connector_ids( + summary.config_name.clone(), + summary.display_name.clone(), + summary.app_connector_ids.clone(), + ) + } +} + +/// Immutable connector selection and provenance after applying disabled plugins. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct ConnectorSnapshot { + connector_ids: Vec, + plugin_display_names_by_connector_id: HashMap>, + disabled_connector_ids: HashSet, +} + +impl ConnectorSnapshot { + /// Builds the final selection from all plugin sources, preserving contribution order. + /// An enabled contributor preserves a shared connector unless its canonical owner is disabled. + pub fn from_plugin_sources( + sources: impl IntoIterator, + disabled_plugin_ids: &[String], + canonical_disabled_connector_ids: HashSet, + ) -> Self { + let mut connector_ids = Vec::new(); + let mut seen_connector_ids = HashSet::new(); + let mut plugin_display_names_by_connector_id: HashMap> = HashMap::new(); + let mut disabled_connector_ids = HashSet::new(); + + for source in sources { + if disabled_plugin_ids.contains(&source.plugin_id) { + disabled_connector_ids.extend(source.connector_ids.into_iter().map(|id| id.0)); + continue; + } + for connector_id in source.connector_ids() { + if canonical_disabled_connector_ids.contains(&connector_id.0) { + continue; + } + if seen_connector_ids.insert(connector_id.0.clone()) { + connector_ids.push(connector_id.clone()); + } + plugin_display_names_by_connector_id + .entry(connector_id.0.clone()) + .or_default() + .push(source.plugin_display_name().to_string()); + } + } + for plugin_names in plugin_display_names_by_connector_id.values_mut() { + plugin_names.sort_unstable(); + plugin_names.dedup(); + } + disabled_connector_ids.retain(|id| !seen_connector_ids.contains(id)); + disabled_connector_ids.extend(canonical_disabled_connector_ids); + + Self { + connector_ids, + plugin_display_names_by_connector_id, + disabled_connector_ids, + } + } + + /// Adapts the current host plugin summaries to the connector-owned snapshot. + pub fn from_plugin_capability_summaries(summaries: &[PluginCapabilitySummary]) -> Self { + Self::from_plugin_sources( + summaries.iter().map(PluginConnectorSource::from), + &[], + HashSet::new(), + ) + } + + /// Returns the connector IDs in source contribution order. + pub fn connector_ids(&self) -> &[AppConnectorId] { + &self.connector_ids + } + + /// Connector tools explicitly excluded or excluded by all contributing plugins. + pub fn disabled_connector_ids(&self) -> &HashSet { + &self.disabled_connector_ids + } + + /// Returns the package display names associated with one connector. + pub fn plugin_display_names_for_connector_id(&self, connector_id: &str) -> &[String] { + self.plugin_display_names_by_connector_id + .get(connector_id) + .map(Vec::as_slice) + .unwrap_or_default() + } +} + +#[cfg(test)] +#[path = "snapshot_tests.rs"] +mod tests; diff --git a/codex-rs/connectors/src/snapshot_tests.rs b/codex-rs/connectors/src/snapshot_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..e6eb2b6b071c2271ec0d3a07e5b99c42915e435d --- /dev/null +++ b/codex-rs/connectors/src/snapshot_tests.rs @@ -0,0 +1,97 @@ +use std::collections::HashSet; + +use codex_plugin::AppConnectorId; +use pretty_assertions::assert_eq; + +use super::ConnectorSnapshot; +use super::PluginConnectorSource; + +#[test] +fn snapshot_merges_sources_in_order_and_dedupes_provenance() { + let host = [ + source("skills", "Skills only", &[]), + source("host", "Zulu", &["calendar", "calendar"]), + ]; + let selected = [ + source("selected-a", "Alpha", &["drive", "calendar"]), + source("selected-b", "Alpha", &["calendar"]), + ]; + + let merged = ConnectorSnapshot::from_plugin_sources( + host.into_iter().chain(selected), + &[], + HashSet::new(), + ); + + assert_eq!( + merged.connector_ids(), + &[ + AppConnectorId("calendar".to_string()), + AppConnectorId("drive".to_string()), + ] + ); + assert_eq!( + merged.plugin_display_names_for_connector_id("calendar"), + &["Alpha".to_string(), "Zulu".to_string()] + ); + assert_eq!( + merged.plugin_display_names_for_connector_id("missing"), + &[] as &[String] + ); +} + +#[test] +fn disabled_plugins_preserve_shared_connectors() { + let sources = [ + source("alpha", "Alpha", &["exclusive", "shared"]), + source("beta", "Beta", &["other", "shared"]), + ]; + let filtered = ConnectorSnapshot::from_plugin_sources( + sources.clone(), + &["alpha".to_string()], + HashSet::new(), + ); + + let expected = ConnectorSnapshot { + disabled_connector_ids: HashSet::from(["exclusive".to_string()]), + ..ConnectorSnapshot::from_plugin_sources( + [source("beta", "Beta", &["other", "shared"])], + &[], + HashSet::new(), + ) + }; + assert_eq!(filtered, expected); + assert_eq!( + ConnectorSnapshot::from_plugin_sources( + sources.iter().rev().cloned(), + &["alpha".to_string()], + HashSet::new(), + ), + expected + ); + assert_eq!( + ConnectorSnapshot::from_plugin_sources( + sources, + &["alpha".to_string(), "beta".to_string()], + HashSet::new(), + ), + ConnectorSnapshot { + disabled_connector_ids: HashSet::from([ + "exclusive".to_string(), + "other".to_string(), + "shared".to_string(), + ]), + ..ConnectorSnapshot::default() + } + ); +} + +fn source(id: &str, display_name: &str, connector_ids: &[&str]) -> PluginConnectorSource { + PluginConnectorSource::from_connector_ids( + id, + display_name, + connector_ids + .iter() + .map(|id| AppConnectorId((*id).to_string())), + ) +} diff --git a/codex-rs/keyring-store/src/lib.rs b/codex-rs/keyring-store/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..ee91af11428476e84303b920cbf8dc186c2f7791 --- /dev/null +++ b/codex-rs/keyring-store/src/lib.rs @@ -0,0 +1,226 @@ +use keyring::Entry; +use keyring::Error as KeyringError; +use std::error::Error; +use std::fmt; +use std::fmt::Debug; +use tracing::trace; + +#[derive(Debug)] +pub enum CredentialStoreError { + Other(KeyringError), +} + +impl CredentialStoreError { + pub fn new(error: KeyringError) -> Self { + Self::Other(error) + } + + pub fn message(&self) -> String { + match self { + Self::Other(error) => error.to_string(), + } + } + + pub fn into_error(self) -> KeyringError { + match self { + Self::Other(error) => error, + } + } +} + +impl fmt::Display for CredentialStoreError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Other(error) => write!(f, "{error}"), + } + } +} + +impl Error for CredentialStoreError {} + +/// Shared credential store abstraction for keyring-backed implementations. +pub trait KeyringStore: Debug + Send + Sync { + fn load(&self, service: &str, account: &str) -> Result, CredentialStoreError>; + fn save(&self, service: &str, account: &str, value: &str) -> Result<(), CredentialStoreError>; + fn delete(&self, service: &str, account: &str) -> Result; +} + +#[derive(Debug, Clone, Copy)] +pub struct DefaultKeyringStore; + +impl KeyringStore for DefaultKeyringStore { + fn load(&self, service: &str, account: &str) -> Result, CredentialStoreError> { + trace!("keyring.load start, service={service}, account={account}"); + let entry = Entry::new(service, account).map_err(CredentialStoreError::new)?; + match entry.get_password() { + Ok(password) => { + trace!("keyring.load success, service={service}, account={account}"); + Ok(Some(password)) + } + Err(keyring::Error::NoEntry) => { + trace!("keyring.load no entry, service={service}, account={account}"); + Ok(None) + } + Err(error) => { + trace!("keyring.load error, service={service}, account={account}, error={error}"); + Err(CredentialStoreError::new(error)) + } + } + } + + fn save(&self, service: &str, account: &str, value: &str) -> Result<(), CredentialStoreError> { + trace!( + "keyring.save start, service={service}, account={account}, value_len={}", + value.len() + ); + let entry = Entry::new(service, account).map_err(CredentialStoreError::new)?; + match entry.set_password(value) { + Ok(()) => { + trace!("keyring.save success, service={service}, account={account}"); + Ok(()) + } + Err(error) => { + trace!("keyring.save error, service={service}, account={account}, error={error}"); + Err(CredentialStoreError::new(error)) + } + } + } + + fn delete(&self, service: &str, account: &str) -> Result { + trace!("keyring.delete start, service={service}, account={account}"); + let entry = Entry::new(service, account).map_err(CredentialStoreError::new)?; + match entry.delete_credential() { + Ok(()) => { + trace!("keyring.delete success, service={service}, account={account}"); + Ok(true) + } + Err(keyring::Error::NoEntry) => { + trace!("keyring.delete no entry, service={service}, account={account}"); + Ok(false) + } + Err(error) => { + trace!("keyring.delete error, service={service}, account={account}, error={error}"); + Err(CredentialStoreError::new(error)) + } + } + } +} + +pub mod tests { + use super::CredentialStoreError; + use super::KeyringStore; + use keyring::Error as KeyringError; + use keyring::credential::CredentialApi as _; + use keyring::mock::MockCredential; + use std::collections::HashMap; + use std::sync::Arc; + use std::sync::Mutex; + use std::sync::PoisonError; + + #[derive(Default, Clone, Debug)] + pub struct MockKeyringStore { + credentials: Arc>>>, + } + + impl MockKeyringStore { + pub fn credential(&self, account: &str) -> Arc { + let mut guard = self + .credentials + .lock() + .unwrap_or_else(PoisonError::into_inner); + guard + .entry(account.to_string()) + .or_insert_with(|| Arc::new(MockCredential::default())) + .clone() + } + + pub fn saved_value(&self, account: &str) -> Option { + let credential = { + let guard = self + .credentials + .lock() + .unwrap_or_else(PoisonError::into_inner); + guard.get(account).cloned() + }?; + credential.get_password().ok() + } + + pub fn set_error(&self, account: &str, error: KeyringError) { + let credential = self.credential(account); + credential.set_error(error); + } + + pub fn contains(&self, account: &str) -> bool { + let guard = self + .credentials + .lock() + .unwrap_or_else(PoisonError::into_inner); + guard.contains_key(account) + } + } + + impl KeyringStore for MockKeyringStore { + fn load( + &self, + _service: &str, + account: &str, + ) -> Result, CredentialStoreError> { + let credential = { + let guard = self + .credentials + .lock() + .unwrap_or_else(PoisonError::into_inner); + guard.get(account).cloned() + }; + + let Some(credential) = credential else { + return Ok(None); + }; + + match credential.get_password() { + Ok(password) => Ok(Some(password)), + Err(KeyringError::NoEntry) => Ok(None), + Err(error) => Err(CredentialStoreError::new(error)), + } + } + + fn save( + &self, + _service: &str, + account: &str, + value: &str, + ) -> Result<(), CredentialStoreError> { + let credential = self.credential(account); + credential + .set_password(value) + .map_err(CredentialStoreError::new) + } + + fn delete(&self, _service: &str, account: &str) -> Result { + let credential = { + let guard = self + .credentials + .lock() + .unwrap_or_else(PoisonError::into_inner); + guard.get(account).cloned() + }; + + let Some(credential) = credential else { + return Ok(false); + }; + + let removed = match credential.delete_credential() { + Ok(()) => Ok(true), + Err(KeyringError::NoEntry) => Ok(false), + Err(error) => Err(CredentialStoreError::new(error)), + }?; + + let mut guard = self + .credentials + .lock() + .unwrap_or_else(PoisonError::into_inner); + guard.remove(account); + Ok(removed) + } + } +} diff --git a/codex-rs/lmstudio/src/client.rs b/codex-rs/lmstudio/src/client.rs new file mode 100644 index 0000000000000000000000000000000000000000..8b0a89a93290ee9c3b65bdf16736efeb58bc7174 --- /dev/null +++ b/codex-rs/lmstudio/src/client.rs @@ -0,0 +1,424 @@ +use codex_core::config::Config; +use codex_http_client::ClientRouteClass; +use codex_http_client::RouteAwareClientPool; +use codex_model_provider_info::LMSTUDIO_OSS_PROVIDER_ID; +use std::io; +use std::path::Path; +use std::time::Duration; + +#[derive(Clone)] +pub struct LMStudioClient { + client: RouteAwareClientPool, + base_url: String, +} + +const LMSTUDIO_CONNECTION_ERROR: &str = "LM Studio is not responding. Install from https://lmstudio.ai/download and run 'lms server start'."; +const LMSTUDIO_CONNECTION_TIMEOUT: Duration = Duration::from_secs(5); + +impl LMStudioClient { + pub async fn try_from_provider(config: &Config) -> std::io::Result { + let provider = config + .model_providers + .get(LMSTUDIO_OSS_PROVIDER_ID) + .ok_or_else(|| { + io::Error::new( + io::ErrorKind::NotFound, + format!("Built-in provider {LMSTUDIO_OSS_PROVIDER_ID} not found",), + ) + })?; + let base_url = provider.base_url.as_ref().ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidData, + "oss provider must have a base_url", + ) + })?; + + let client = RouteAwareClientPool::with_connect_timeout( + config.http_client_factory(), + ClientRouteClass::Other, + LMSTUDIO_CONNECTION_TIMEOUT, + ); + + let client = LMStudioClient { + client, + base_url: base_url.to_string(), + }; + client.check_server().await?; + + Ok(client) + } + + async fn check_server(&self) -> io::Result<()> { + let url = format!("{}/models", self.base_url.trim_end_matches('/')); + let response = self.client.get(&url).send().await; + + if let Ok(resp) = response { + if resp.status().is_success() { + Ok(()) + } else { + Err(io::Error::other(format!( + "Server returned error: {} {LMSTUDIO_CONNECTION_ERROR}", + resp.status() + ))) + } + } else { + Err(io::Error::other(LMSTUDIO_CONNECTION_ERROR)) + } + } + + // Load a model by sending an empty request with max_tokens 1 + pub async fn load_model(&self, model: &str) -> io::Result<()> { + let url = format!("{}/responses", self.base_url.trim_end_matches('/')); + + let request_body = serde_json::json!({ + "model": model, + "input": "", + "max_output_tokens": 1 + }); + + let response = self + .client + .post(&url) + .header("Content-Type", "application/json") + .json(&request_body) + .send() + .await + .map_err(|e| io::Error::other(format!("Request failed: {e}")))?; + + if response.status().is_success() { + tracing::info!("Successfully loaded model '{model}'"); + Ok(()) + } else { + Err(io::Error::other(format!( + "Failed to load model: {}", + response.status() + ))) + } + } + + // Return the list of models available on the LM Studio server. + pub async fn fetch_models(&self) -> io::Result> { + let url = format!("{}/models", self.base_url.trim_end_matches('/')); + let response = self + .client + .get(&url) + .send() + .await + .map_err(|e| io::Error::other(format!("Request failed: {e}")))?; + + if response.status().is_success() { + let json: serde_json::Value = response.json().await.map_err(|e| { + io::Error::new(io::ErrorKind::InvalidData, format!("JSON parse error: {e}")) + })?; + let models = json["data"] + .as_array() + .ok_or_else(|| { + io::Error::new(io::ErrorKind::InvalidData, "No 'data' array in response") + })? + .iter() + .filter_map(|model| model["id"].as_str()) + .map(std::string::ToString::to_string) + .collect(); + Ok(models) + } else { + Err(io::Error::other(format!( + "Failed to fetch models: {}", + response.status() + ))) + } + } + + // Find lms, checking fallback paths if not in PATH + fn find_lms() -> std::io::Result { + Self::find_lms_with_home_dir(/*home_dir*/ None) + } + + fn find_lms_with_home_dir(home_dir: Option<&str>) -> std::io::Result { + // First try 'lms' in PATH + if which::which("lms").is_ok() { + return Ok("lms".to_string()); + } + + // Platform-specific fallback paths + let home = match home_dir { + Some(dir) => dir.to_string(), + None => { + #[cfg(unix)] + { + std::env::var("HOME").unwrap_or_default() + } + #[cfg(windows)] + { + std::env::var("USERPROFILE").unwrap_or_default() + } + } + }; + + #[cfg(unix)] + let fallback_path = format!("{home}/.lmstudio/bin/lms"); + + #[cfg(windows)] + let fallback_path = format!("{home}/.lmstudio/bin/lms.exe"); + + if Path::new(&fallback_path).exists() { + Ok(fallback_path) + } else { + Err(std::io::Error::new( + std::io::ErrorKind::NotFound, + "LM Studio not found. Please install LM Studio from https://lmstudio.ai/", + )) + } + } + + pub async fn download_model(&self, model: &str) -> std::io::Result<()> { + let lms = Self::find_lms()?; + eprintln!("Downloading model: {model}"); + + let status = std::process::Command::new(&lms) + .args(["get", "--yes", model]) + .stdout(std::process::Stdio::inherit()) + .stderr(std::process::Stdio::null()) + .status() + .map_err(|e| { + std::io::Error::other(format!("Failed to execute '{lms} get --yes {model}': {e}")) + })?; + + if !status.success() { + return Err(std::io::Error::other(format!( + "Model download failed with exit code: {}", + status.code().unwrap_or(-1) + ))); + } + + tracing::info!("Successfully downloaded model '{model}'"); + Ok(()) + } +} + +#[cfg(test)] +mod tests { + #![allow(clippy::expect_used, clippy::unwrap_used)] + use super::*; + + fn client_from_host_root( + host_root: impl Into, + connection_timeout: Duration, + ) -> LMStudioClient { + let client = RouteAwareClientPool::with_connect_timeout( + codex_http_client::HttpClientFactory::new( + codex_http_client::OutboundProxyPolicy::ReqwestDefault, + ), + ClientRouteClass::Other, + connection_timeout, + ); + LMStudioClient { + client, + base_url: host_root.into(), + } + } + + #[tokio::test] + async fn test_fetch_models_happy_path() { + if std::env::var(codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR).is_ok() { + tracing::info!( + "{} is set; skipping test_fetch_models_happy_path", + codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR + ); + return; + } + + let server = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/models")) + .respond_with( + wiremock::ResponseTemplate::new(200).set_body_raw( + serde_json::json!({ + "data": [ + {"id": "openai/gpt-oss-20b"}, + ] + }) + .to_string(), + "application/json", + ), + ) + .mount(&server) + .await; + + let client = client_from_host_root(server.uri(), LMSTUDIO_CONNECTION_TIMEOUT); + let models = client.fetch_models().await.expect("fetch models"); + assert!(models.contains(&"openai/gpt-oss-20b".to_string())); + } + + #[tokio::test] + async fn test_fetch_models_no_data_array() { + if std::env::var(codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR).is_ok() { + tracing::info!( + "{} is set; skipping test_fetch_models_no_data_array", + codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR + ); + return; + } + + let server = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/models")) + .respond_with( + wiremock::ResponseTemplate::new(200) + .set_body_raw(serde_json::json!({}).to_string(), "application/json"), + ) + .mount(&server) + .await; + + let client = client_from_host_root(server.uri(), LMSTUDIO_CONNECTION_TIMEOUT); + let result = client.fetch_models().await; + assert!(result.is_err()); + assert!( + result + .unwrap_err() + .to_string() + .contains("No 'data' array in response") + ); + } + + #[tokio::test] + async fn test_fetch_models_server_error() { + if std::env::var(codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR).is_ok() { + tracing::info!( + "{} is set; skipping test_fetch_models_server_error", + codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR + ); + return; + } + + let server = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/models")) + .respond_with(wiremock::ResponseTemplate::new(500)) + .mount(&server) + .await; + + let client = client_from_host_root(server.uri(), LMSTUDIO_CONNECTION_TIMEOUT); + let result = client.fetch_models().await; + assert!(result.is_err()); + assert!( + result + .unwrap_err() + .to_string() + .contains("Failed to fetch models: 500") + ); + } + + #[tokio::test] + async fn test_check_server_happy_path() { + if std::env::var(codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR).is_ok() { + tracing::info!( + "{} is set; skipping test_check_server_happy_path", + codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR + ); + return; + } + + let server = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/models")) + .respond_with(wiremock::ResponseTemplate::new(200)) + .mount(&server) + .await; + + let client = client_from_host_root(server.uri(), LMSTUDIO_CONNECTION_TIMEOUT); + client + .check_server() + .await + .expect("server check should pass"); + } + + #[tokio::test] + async fn test_check_server_allows_slow_response_after_connect() { + if std::env::var(codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR).is_ok() { + tracing::info!( + "{} is set; skipping test_check_server_allows_slow_response_after_connect", + codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR + ); + return; + } + + let server = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/models")) + .respond_with( + wiremock::ResponseTemplate::new(200).set_delay(Duration::from_millis(250)), + ) + .mount(&server) + .await; + + let client = client_from_host_root(server.uri(), Duration::from_millis(100)); + + client + .check_server() + .await + .expect("server check should allow a slow response after connecting"); + } + + #[tokio::test] + async fn test_check_server_error() { + if std::env::var(codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR).is_ok() { + tracing::info!( + "{} is set; skipping test_check_server_error", + codex_core::spawn::CODEX_SANDBOX_NETWORK_DISABLED_ENV_VAR + ); + return; + } + + let server = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/models")) + .respond_with(wiremock::ResponseTemplate::new(404)) + .mount(&server) + .await; + + let client = client_from_host_root(server.uri(), LMSTUDIO_CONNECTION_TIMEOUT); + let result = client.check_server().await; + assert!(result.is_err()); + assert!( + result + .unwrap_err() + .to_string() + .contains("Server returned error: 404") + ); + } + + #[test] + fn test_find_lms() { + let result = LMStudioClient::find_lms(); + + match result { + Ok(_) => { + // lms was found in PATH - that's fine + } + Err(e) => { + // Expected error when LM Studio not installed + assert!(e.to_string().contains("LM Studio not found")); + } + } + } + + #[test] + fn test_find_lms_with_mock_home() { + // Test fallback path construction without touching env vars + #[cfg(unix)] + { + let result = LMStudioClient::find_lms_with_home_dir(Some("/test/home")); + if let Err(e) = result { + assert!(e.to_string().contains("LM Studio not found")); + } + } + + #[cfg(windows)] + { + let result = LMStudioClient::find_lms_with_home_dir(Some("C:\\test\\home")); + if let Err(e) = result { + assert!(e.to_string().contains("LM Studio not found")); + } + } + } +} diff --git a/codex-rs/lmstudio/src/lib.rs b/codex-rs/lmstudio/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..fd4f82a728a3dcb6312e13f288c54bebace54cab --- /dev/null +++ b/codex-rs/lmstudio/src/lib.rs @@ -0,0 +1,46 @@ +mod client; + +pub use client::LMStudioClient; +use codex_core::config::Config; + +/// Default OSS model to use when `--oss` is passed without an explicit `-m`. +pub const DEFAULT_OSS_MODEL: &str = "openai/gpt-oss-20b"; + +/// Prepare the local OSS environment when `--oss` is selected. +/// +/// - Ensures a local LM Studio server is reachable. +/// - Checks if the model exists locally and downloads it if missing. +pub async fn ensure_oss_ready(config: &Config) -> std::io::Result<()> { + let model = match config.model.as_ref() { + Some(model) => model, + None => DEFAULT_OSS_MODEL, + }; + + // Verify local LM Studio is reachable. + let lmstudio_client = LMStudioClient::try_from_provider(config).await?; + + match lmstudio_client.fetch_models().await { + Ok(models) => { + if !models.iter().any(|m| m == model) { + lmstudio_client.download_model(model).await?; + } + } + Err(err) => { + // Not fatal; higher layers may still proceed and surface errors later. + tracing::warn!("Failed to query local models from LM Studio: {}.", err); + } + } + + // Load the model in the background + tokio::spawn({ + let client = lmstudio_client.clone(); + let model = model.to_string(); + async move { + if let Err(e) = client.load_model(&model).await { + tracing::warn!("Failed to load model {}: {}", model, e); + } + } + }); + + Ok(()) +} diff --git a/codex-rs/network-proxy/src/attribution.rs b/codex-rs/network-proxy/src/attribution.rs new file mode 100644 index 0000000000000000000000000000000000000000..85affff15b01798f0d94b52205b52e6de23533ca --- /dev/null +++ b/codex-rs/network-proxy/src/attribution.rs @@ -0,0 +1,145 @@ +use crate::state::NetworkProxyState; +use rama_core::Service; +use rama_core::error::BoxError; +use rama_core::extensions::ExtensionsMut; +use rama_tcp::TcpStream; +use std::io; +use std::io::Write; +use std::sync::Arc; +use std::time::Duration; +use tokio::io::AsyncReadExt; + +/// Internal handoff from the trusted Linux proxy bridge. +#[doc(hidden)] +pub const PROXY_ATTRIBUTION_TOKEN_ENV_KEY: &str = "CODEX_NETWORK_PROXY_ATTRIBUTION"; + +const ATTRIBUTION_FRAME_MAGIC: &[u8; 8] = b"\0CDXPXY1"; +const MAX_ATTRIBUTION_TOKEN_LEN: usize = 128; +const ATTRIBUTION_FRAME_TIMEOUT: Duration = Duration::from_secs(3); + +pub(crate) struct BindConnectionAttribution { + inner: S, + state: Arc, + environment_id: Option, +} + +impl BindConnectionAttribution { + pub(crate) fn new( + inner: S, + state: Arc, + environment_id: Option, + ) -> Self { + Self { + inner, + state, + environment_id, + } + } +} + +impl Service for BindConnectionAttribution +where + S: Service, + S::Error: Into, +{ + type Output = S::Output; + type Error = BoxError; + + async fn serve(&self, mut stream: TcpStream) -> Result { + let state = match read_attribution_token(&mut stream).await? { + Some(token) => self.state.for_execution_token(&token).ok_or_else(|| { + io::Error::new( + io::ErrorKind::PermissionDenied, + "unknown network proxy attribution token", + ) + })?, + None => self + .state + .for_environment_id(self.environment_id.as_deref()), + }; + if let Some(expected_environment_id) = self.environment_id.as_deref() + && state + .environment_id() + .is_some_and(|actual| actual != expected_environment_id) + { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "network proxy attribution environment mismatch", + ) + .into()); + } + stream.extensions_mut().insert(Arc::new(state)); + self.inner.serve(stream).await.map_err(Into::into) + } +} + +async fn read_attribution_token(stream: &mut TcpStream) -> Result, BoxError> { + let mut marker = [0_u8; 1]; + let read = stream.stream.peek(&mut marker).await?; + if read == 0 { + return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "empty proxy connection").into()); + } + if marker[0] != ATTRIBUTION_FRAME_MAGIC[0] { + return Ok(None); + } + + let token = tokio::time::timeout(ATTRIBUTION_FRAME_TIMEOUT, async { + let mut magic = [0_u8; ATTRIBUTION_FRAME_MAGIC.len()]; + stream.read_exact(&mut magic).await?; + if &magic != ATTRIBUTION_FRAME_MAGIC { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "invalid network proxy attribution frame", + )); + } + + let token_len = stream.read_u16().await? as usize; + if token_len == 0 || token_len > MAX_ATTRIBUTION_TOKEN_LEN { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "invalid network proxy attribution token length", + )); + } + let mut token = vec![0_u8; token_len]; + stream.read_exact(&mut token).await?; + String::from_utf8(token).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidData, + "network proxy attribution token is not UTF-8", + ) + }) + }) + .await + .map_err(|_| { + io::Error::new( + io::ErrorKind::TimedOut, + "network proxy attribution frame timed out", + ) + })??; + + Ok(Some(token)) +} + +/// Writes the trusted bridge preface consumed by the shared proxy ingress. +#[doc(hidden)] +pub fn write_attribution_frame(writer: &mut impl Write, token: &str) -> io::Result<()> { + if token.is_empty() || token.len() > MAX_ATTRIBUTION_TOKEN_LEN { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "invalid network proxy attribution token length", + )); + } + let token_len = u16::try_from(token.len()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "network proxy attribution token is too long", + ) + })?; + writer.write_all(ATTRIBUTION_FRAME_MAGIC)?; + writer.write_all(&token_len.to_be_bytes())?; + writer.write_all(token.as_bytes()) +} + +#[cfg(test)] +#[path = "attribution_tests.rs"] +mod tests; diff --git a/codex-rs/network-proxy/src/attribution_tests.rs b/codex-rs/network-proxy/src/attribution_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..fcb3c3771e157bbb53ae9ec5f9ec8611524f2e9a --- /dev/null +++ b/codex-rs/network-proxy/src/attribution_tests.rs @@ -0,0 +1,61 @@ +use super::BindConnectionAttribution; +use super::write_attribution_frame; +use crate::config::NetworkProxyConfig; +use crate::runtime::network_proxy_state_for_policy; +use crate::state::NetworkProxyState; +use pretty_assertions::assert_eq; +use rama_core::Service; +use rama_core::error::BoxError; +use rama_core::extensions::ExtensionsRef; +use rama_core::service::service_fn; +use rama_tcp::TcpStream as RamaTcpStream; +use std::io; +use std::sync::Arc; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::net::TcpStream; + +#[test] +fn attribution_frame_has_bounded_binary_prefix() -> io::Result<()> { + let mut frame = Vec::new(); + write_attribution_frame(&mut frame, "token-1")?; + + assert_eq!(&frame[..8], b"\0CDXPXY1"); + assert_eq!(u16::from_be_bytes([frame[8], frame[9]]), 7); + assert_eq!(&frame[10..], b"token-1"); + Ok(()) +} + +#[tokio::test] +async fn framed_connection_receives_registered_execution_state() -> Result<(), BoxError> { + let state = Arc::new(network_proxy_state_for_policy(NetworkProxyConfig::default())); + state.register_execution("token-1", "local", "execution-1"); + + let listener = TcpListener::bind("127.0.0.1:0").await?; + let addr = listener.local_addr()?; + let client = tokio::spawn(async move { + let mut stream = TcpStream::connect(addr).await?; + let mut frame = Vec::new(); + write_attribution_frame(&mut frame, "token-1")?; + stream.write_all(&frame).await + }); + + let (stream, _) = listener.accept().await?; + let service = BindConnectionAttribution::new( + service_fn(|stream: RamaTcpStream| async move { + let state = stream.extensions().get::>().cloned(); + Ok::<_, io::Error>(state) + }), + state, + Some("local".to_string()), + ); + let actual = service + .serve(RamaTcpStream::new(stream)) + .await? + .expect("connection state"); + client.await??; + + assert_eq!(actual.environment_id(), Some("local")); + assert_eq!(actual.execution_id().as_deref(), Some("execution-1")); + Ok(()) +} diff --git a/codex-rs/network-proxy/src/authorization_path.rs b/codex-rs/network-proxy/src/authorization_path.rs new file mode 100644 index 0000000000000000000000000000000000000000..7a692b8c2c452652cd61fa62917cb707d6726266 --- /dev/null +++ b/codex-rs/network-proxy/src/authorization_path.rs @@ -0,0 +1,64 @@ +/// Returns whether `path` has an unambiguous interpretation for authorization. +/// +/// MITM hooks authorize the request before the upstream server parses it. Reject +/// path forms that common upstreams may decode or normalize into a different +/// resource after a hook has matched. +pub(crate) fn is_safe_for_authorization(path: &str) -> bool { + path.split('/').all(is_safe_segment_for_authorization) +} + +fn is_safe_segment_for_authorization(segment: &str) -> bool { + let bytes = segment.as_bytes(); + let mut index = 0; + let mut decoded_dots = 0; + let mut has_non_dot = false; + while index < bytes.len() { + match bytes[index] { + b'.' => { + decoded_dots += 1; + index += 1; + } + b'\\' => return false, + b'%' => { + let Some(high) = bytes + .get(index + 1) + .and_then(|byte| decode_hex_digit(*byte)) + else { + return false; + }; + let Some(low) = bytes + .get(index + 2) + .and_then(|byte| decode_hex_digit(*byte)) + else { + return false; + }; + let decoded = high << 4 | low; + match decoded { + b'%' | b'/' | b'\\' => return false, + b'.' => decoded_dots += 1, + _ => has_non_dot = true, + } + index += 3; + } + _ => { + has_non_dot = true; + index += 1; + } + } + } + + has_non_dot || !matches!(decoded_dots, 1 | 2) +} + +fn decode_hex_digit(byte: u8) -> Option { + match byte { + b'0'..=b'9' => Some(byte - b'0'), + b'a'..=b'f' => Some(byte - b'a' + 10), + b'A'..=b'F' => Some(byte - b'A' + 10), + _ => None, + } +} + +#[cfg(test)] +#[path = "authorization_path_tests.rs"] +mod tests; diff --git a/codex-rs/network-proxy/src/brokered_tunnel_tests.rs b/codex-rs/network-proxy/src/brokered_tunnel_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..72affafb9a4cdce5195cd4d01b6d215b0fabde96 --- /dev/null +++ b/codex-rs/network-proxy/src/brokered_tunnel_tests.rs @@ -0,0 +1,544 @@ +//! End-to-end checks for credential scope and protocol compatibility on CONNECT and SOCKS5. + +use crate::CredentialProviderConfig; +use crate::NetworkMode; +use crate::NetworkProxyConfig; +use crate::NetworkProxyConstraints; +use crate::NetworkProxyState; +use crate::build_config_state; +use crate::connection_lifecycle::ProxyListeners; +use crate::runtime::ConfigReloader; +use crate::runtime::ConfigReloaderFuture; +use crate::runtime::ConfigState; +use pretty_assertions::assert_eq; +use rama_core::extensions::Extensions; +use rama_core::extensions::ExtensionsMut; +use rama_core::extensions::ExtensionsRef; +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::net::Ipv4Addr; +use std::net::SocketAddr; +use std::net::TcpListener as StdTcpListener; +use std::pin::Pin; +use std::sync::Arc; +use std::task::Context; +use std::task::Poll; +use tokio::io::AsyncRead; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWrite; +use tokio::io::AsyncWriteExt; +use tokio::io::DuplexStream; +use tokio::io::ReadBuf; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::time::Duration; +use tokio::time::timeout; + +struct TestStream { + inner: DuplexStream, + extensions: Extensions, +} + +impl TestStream { + fn new(inner: DuplexStream) -> Self { + Self { + inner, + extensions: Extensions::new(), + } + } +} + +impl AsyncRead for TestStream { + fn poll_read( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_read(cx, buf) + } +} + +impl AsyncWrite for TestStream { + fn poll_write( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_write(cx, buf) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_flush(cx) + } + + fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_shutdown(cx) + } +} + +impl ExtensionsRef for TestStream { + fn extensions(&self) -> &Extensions { + &self.extensions + } +} + +impl ExtensionsMut for TestStream { + fn extensions_mut(&mut self) -> &mut Extensions { + &mut self.extensions + } +} + +#[tokio::test] +async fn tls_prefix_detection_accumulates_fragmented_reads() { + let tls_prefix = [0x16, 0x03, 0x03, 0x00, 0x80]; + let (mut writer, reader) = tokio::io::duplex(16); + let writer_task = tokio::spawn(async move { + writer.write_all(&tls_prefix[..1]).await.unwrap(); + tokio::time::sleep(Duration::from_millis(300)).await; + writer.write_all(&tls_prefix[1..]).await.unwrap(); + }); + + let (protocol, mut stream) = crate::brokered_tunnel::peek_protocol( + TestStream::new(reader), + crate::brokered_tunnel::BrokeredProtocols { + tls: true, + http: true, + }, + ) + .await + .unwrap(); + let mut replayed = [0_u8; 5]; + stream.read_exact(&mut replayed).await.unwrap(); + + assert_eq!(protocol, crate::brokered_tunnel::TunnelProtocol::Tls); + assert_eq!(replayed, tls_prefix); + writer_task.await.unwrap(); +} + +const TOKEN: &str = "local_abcdefghijklmnopqrstuvwxyzabcdef"; + +struct StaticReloader(ConfigState); + +impl ConfigReloader for StaticReloader { + fn source_label(&self) -> String { + "HTTP tunnel test".to_string() + } + + fn maybe_reload(&self) -> ConfigReloaderFuture<'_, Option> { + Box::pin(async { Ok(None) }) + } + + fn reload_now(&self) -> ConfigReloaderFuture<'_, ConfigState> { + Box::pin(async { Ok(self.0.clone()) }) + } +} + +#[derive(Clone, Copy, Debug)] +enum Transport { + Connect, + Socks5, +} + +struct TestProxy { + addr: SocketAddr, + listeners: ProxyListeners, + dummy: String, + transport: Transport, +} + +impl Drop for TestProxy { + fn drop(&mut self) { + self.listeners.cancel(); + } +} + +impl TestProxy { + fn start(target: SocketAddr, transport: Transport, mode: NetworkMode) -> Self { + let mut config = NetworkProxyConfig { + enabled: true, + mode, + mitm: true, + credential_broker: true, + allow_local_binding: true, + credential_providers: BTreeMap::from([( + "local".to_string(), + CredentialProviderConfig { + env: vec!["LOCAL_TOKEN".to_string()], + patterns: vec!["^local_[a-z]{32}$".to_string()], + url_prefixes: vec![format!("http://{target}/v1")], + ..CredentialProviderConfig::default() + }, + )]), + ..NetworkProxyConfig::default() + }; + config.set_allowed_domains(vec!["127.0.0.1".to_string()]); + let config_state = build_config_state(config, NetworkProxyConstraints::default()).unwrap(); + let state = Arc::new(NetworkProxyState::with_reloader( + config_state.clone(), + Arc::new(StaticReloader(config_state)), + )); + let mut env = HashMap::from([("LOCAL_TOKEN".to_string(), TOKEN.to_string())]); + state.virtualize_snapshot_credentials(&mut env, /*environment_id*/ None); + let dummy = env.remove("LOCAL_TOKEN").unwrap(); + assert_ne!(dummy, TOKEN); + let listener = StdTcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let addr = listener.local_addr().unwrap(); + let mut listeners = ProxyListeners::new(); + match transport { + Transport::Connect => listeners.spawn(move |guard| { + crate::http_proxy::run_http_proxy_with_std_listener( + state, listener, /*policy_decider*/ None, /*environment_id*/ None, + guard, + ) + }), + Transport::Socks5 => listeners.spawn(move |guard| { + crate::socks5::run_socks5_with_std_listener( + state, listener, /*policy_decider*/ None, /*environment_id*/ None, + /*enable_socks5_udp*/ false, guard, + ) + }), + } + Self { + addr, + listeners, + dummy, + transport, + } + } + + async fn connect(&self, target: SocketAddr) -> TcpStream { + let mut stream = TcpStream::connect(self.addr).await.unwrap(); + match self.transport { + Transport::Connect => { + stream + .write_all( + format!("CONNECT {target} HTTP/1.1\r\nHost: {target}\r\n\r\n").as_bytes(), + ) + .await + .unwrap(); + assert!(read_headers(&mut stream).await.starts_with("HTTP/1.1 200")); + } + Transport::Socks5 => { + stream.write_all(&[5, 1, 0]).await.unwrap(); + let mut greeting = [0; 2]; + stream.read_exact(&mut greeting).await.unwrap(); + assert_eq!(greeting, [5, 0]); + let mut request = vec![5, 1, 0, 1, 127, 0, 0, 1]; + request.extend_from_slice(&target.port().to_be_bytes()); + stream.write_all(&request).await.unwrap(); + let mut response = [0; 10]; + stream.read_exact(&mut response).await.unwrap(); + assert_eq!(&response[..4], &[5, 0, 0, 1]); + } + } + stream + } +} + +async fn read_headers(stream: &mut TcpStream) -> String { + timeout(Duration::from_secs(3), async { + let mut bytes = Vec::new(); + while !bytes.ends_with(b"\r\n\r\n") { + assert!(bytes.len() < 16384); + bytes.push(stream.read_u8().await.unwrap()); + } + String::from_utf8(bytes).unwrap() + }) + .await + .unwrap() +} + +#[tokio::test] +async fn plaintext_tunnels_translate_only_the_configured_url_prefix() { + for transport in [Transport::Connect, Transport::Socks5] { + let target = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); + let addr = target.local_addr().unwrap(); + let proxy = TestProxy::start(addr, transport, NetworkMode::Full); + let mut client = proxy.connect(addr).await; + for (method, path, body, expected) in [ + ("GET", "/v1/models", "", TOKEN), + ("OPTIONS", "*", "", proxy.dummy.as_str()), + ("POST", "/v1/responses", "request body", TOKEN), + ("GET", "/v10/models", "", proxy.dummy.as_str()), + ("GET", "/v1/%2e%2e/private", "", proxy.dummy.as_str()), + ] { + client.write_all(format!("{method} {path} HTTP/1.1\r\nHost: {addr}\r\nAuthorization: Bearer {}\r\nContent-Length: {}\r\n\r\n{body}", proxy.dummy, body.len()).as_bytes()).await.unwrap(); + let (mut upstream, _) = timeout(Duration::from_secs(3), target.accept()) + .await + .unwrap() + .unwrap(); + let request = read_headers(&mut upstream).await; + assert!(request.starts_with(&format!("{method} {path} HTTP/1.1\r\n"))); + assert!( + request + .lines() + .any(|line| line + .eq_ignore_ascii_case(&format!("authorization: Bearer {expected}"))), + "{transport:?}, {path}: {request:?}" + ); + let mut received_body = vec![0; body.len()]; + timeout( + Duration::from_secs(3), + upstream.read_exact(&mut received_body), + ) + .await + .unwrap() + .unwrap(); + assert_eq!(received_body, body.as_bytes()); + upstream + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n") + .await + .unwrap(); + assert!(read_headers(&mut client).await.starts_with("HTTP/1.1 200")); + } + } +} + +#[tokio::test] +async fn plaintext_tunnels_accept_http2_clients() { + use rama_core::Service; + use rama_core::error::BoxError; + use rama_core::service::service_fn; + use rama_http::Body; + use rama_http::Request; + use rama_http::Version; + use rama_http_backend::client::HttpConnector; + use rama_net::client::EstablishedClientConnection; + for transport in [Transport::Connect, Transport::Socks5] { + let target = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); + let addr = target.local_addr().unwrap(); + let proxy = TestProxy::start(addr, transport, NetworkMode::Full); + let stream = Arc::new(tokio::sync::Mutex::new(Some(rama_tcp::TcpStream::from( + proxy.connect(addr).await, + )))); + let (captured_tx, mut captured_rx) = tokio::sync::mpsc::channel(1); + let target_task = tokio::spawn(async move { + let (stream, _) = target.accept().await.unwrap(); + rama_http_backend::server::HttpServer::auto(rama_core::rt::Executor::default()) + .service(service_fn(move |request: Request| { + let captured_tx = captured_tx.clone(); + async move { + captured_tx + .send(request.headers().get("authorization").cloned()) + .await + .unwrap(); + Ok::<_, std::convert::Infallible>(rama_http::Response::new(Body::empty())) + } + })) + .serve(rama_tcp::TcpStream::from(stream)) + .await + .unwrap(); + }); + let connector = HttpConnector::<_, Body>::new(service_fn(move |req: Request| { + let stream = stream.clone(); + async move { + Ok::<_, BoxError>(EstablishedClientConnection { + input: req, + conn: stream.lock().await.take().unwrap(), + }) + } + })); + let request = Request::builder() + .version(Version::HTTP_2) + .uri(format!("http://{addr}/v1/models")) + .header("authorization", format!("Bearer {}", proxy.dummy)) + .body(Body::empty()) + .unwrap(); + let EstablishedClientConnection { input, conn } = connector.serve(request).await.unwrap(); + let response = tokio::spawn(async move { conn.serve(input).await }); + assert_eq!( + timeout(Duration::from_secs(3), captured_rx.recv()) + .await + .unwrap() + .unwrap(), + Some(rama_http::HeaderValue::from_str(&format!("Bearer {TOKEN}")).unwrap()) + ); + assert_eq!( + timeout(Duration::from_secs(3), response) + .await + .unwrap() + .unwrap() + .unwrap() + .status(), + rama_http::StatusCode::OK + ); + target_task.abort(); + } +} + +#[tokio::test] +async fn plaintext_tunnels_reject_retargeting_and_nested_connect() { + for transport in [Transport::Connect, Transport::Socks5] { + let target = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); + let addr = target.local_addr().unwrap(); + let proxy = TestProxy::start(addr, transport, NetworkMode::Full); + for (method, uri, host, status) in [ + ("GET", "/v1".to_string(), "127.0.0.1:1".to_string(), 400), + ( + "GET", + "http://127.0.0.1:1/v1".to_string(), + addr.to_string(), + 400, + ), + ("GET", format!("https://{addr}/v1"), addr.to_string(), 400), + ("CONNECT", addr.to_string(), addr.to_string(), 405), + ] { + let mut client = proxy.connect(addr).await; + client.write_all(format!("{method} {uri} HTTP/1.1\r\nHost: {host}\r\nAuthorization: Bearer {}\r\nConnection: close\r\n\r\n", proxy.dummy).as_bytes()).await.unwrap(); + let response = read_headers(&mut client).await; + assert!( + response.starts_with(&format!("HTTP/1.1 {status}")), + "{transport:?}: {response:?}" + ); + } + assert!( + timeout(Duration::from_millis(50), target.accept()) + .await + .is_err() + ); + } +} + +#[tokio::test] +async fn plaintext_tunnels_preserve_server_first_and_client_first_opaque_bytes() { + for transport in [Transport::Connect, Transport::Socks5] { + for server_first in [true, false] { + let target = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); + let addr = target.local_addr().unwrap(); + let proxy = TestProxy::start(addr, transport, NetworkMode::Full); + let mut client = proxy.connect(addr).await; + if !server_first { + client.write_all(b"PING").await.unwrap(); + } + let (mut upstream, _) = timeout(Duration::from_secs(3), target.accept()) + .await + .unwrap() + .unwrap(); + upstream.write_all(b"SSH-2.0-server\r\n").await.unwrap(); + let mut banner = [0; 16]; + timeout(Duration::from_secs(3), client.read_exact(&mut banner)) + .await + .unwrap() + .unwrap(); + assert_eq!(&banner, b"SSH-2.0-server\r\n"); + if server_first { + client.write_all(b"PING").await.unwrap(); + } + let mut request = [0; 4]; + timeout(Duration::from_secs(3), upstream.read_exact(&mut request)) + .await + .unwrap() + .unwrap(); + assert_eq!(&request, b"PING"); + } + } +} + +#[tokio::test] +async fn plaintext_tunnels_preserve_http_upgrades() { + for (transport, protocol) in [Transport::Connect, Transport::Socks5] + .into_iter() + .flat_map(|transport| ["test-echo", "h2c"].map(move |protocol| (transport, protocol))) + { + let target = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); + let addr = target.local_addr().unwrap(); + let proxy = TestProxy::start(addr, transport, NetworkMode::Full); + let mut client = proxy.connect(addr).await; + client.write_all(format!("GET /v1 HTTP/1.1\r\nHost: {addr}\r\nAuthorization: Bearer {}\r\nConnection: upgrade, HTTP2-Settings, X-Remove\r\nUpgrade: {protocol}\r\nHTTP2-Settings: AAMAAABk\r\nX-Remove: hop-local\r\n\r\n", proxy.dummy).as_bytes()).await.unwrap(); + let (mut upstream, _) = timeout(Duration::from_secs(3), target.accept()) + .await + .unwrap() + .unwrap(); + let request = read_headers(&mut upstream).await; + assert!( + request + .lines() + .any(|line| line.eq_ignore_ascii_case(&format!("Authorization: Bearer {TOKEN}"))) + ); + assert!( + request + .lines() + .any(|line| line.eq_ignore_ascii_case(&format!("Upgrade: {protocol}"))) + ); + assert_eq!( + request + .lines() + .any(|line| line.eq_ignore_ascii_case("HTTP2-Settings: AAMAAABk")), + protocol == "h2c" + ); + let connection = if protocol == "h2c" { + "upgrade, http2-settings" + } else { + "upgrade" + }; + assert!( + request + .lines() + .any(|line| line.eq_ignore_ascii_case(&format!("Connection: {connection}"))) + ); + assert!(!request.to_ascii_lowercase().contains("x-remove")); + upstream.write_all(format!("HTTP/1.1 101 Switching Protocols\r\nConnection: upgrade\r\nUpgrade: {protocol}\r\n\r\nbanner").as_bytes()).await.unwrap(); + assert!(read_headers(&mut client).await.starts_with("HTTP/1.1 101")); + let mut banner = [0; 6]; + timeout(Duration::from_secs(3), client.read_exact(&mut banner)) + .await + .unwrap() + .unwrap(); + assert_eq!(&banner, b"banner"); + client.write_all(b"PING").await.unwrap(); + let mut payload = [0; 4]; + timeout(Duration::from_secs(3), upstream.read_exact(&mut payload)) + .await + .unwrap() + .unwrap(); + assert_eq!(&payload, b"PING"); + } +} + +#[tokio::test] +async fn limited_plaintext_tunnels_allow_reads_but_reject_writes_and_opaque_traffic() { + for transport in [Transport::Connect, Transport::Socks5] { + let target = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap(); + let addr = target.local_addr().unwrap(); + let proxy = TestProxy::start(addr, transport, NetworkMode::Limited); + for method in ["GET", "HEAD", "OPTIONS"] { + let mut client = proxy.connect(addr).await; + client.write_all(format!("{method} /v1 HTTP/1.1\r\nHost: {addr}\r\nAuthorization: Bearer {}\r\nConnection: close\r\n\r\n", proxy.dummy).as_bytes()).await.unwrap(); + let (mut upstream, _) = timeout(Duration::from_secs(3), target.accept()) + .await + .unwrap() + .unwrap(); + let request = read_headers(&mut upstream).await; + assert!( + request.lines().any( + |line| line.eq_ignore_ascii_case(&format!("Authorization: Bearer {TOKEN}")) + ) + ); + upstream + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n") + .await + .unwrap(); + assert!(read_headers(&mut client).await.starts_with("HTTP/1.1 200")); + } + for (method, extra_headers) in [ + ("POST", ""), + ("GET", "Upgrade: test-echo\r\nConnection: upgrade\r\n"), + ] { + let mut client = proxy.connect(addr).await; + client.write_all(format!("{method} /v1 HTTP/1.1\r\nHost: {addr}\r\nAuthorization: Bearer {}\r\nContent-Length: 0\r\n{extra_headers}\r\n", proxy.dummy).as_bytes()).await.unwrap(); + assert!(read_headers(&mut client).await.starts_with("HTTP/1.1 403")); + } + let mut client = proxy.connect(addr).await; + client.write_all(b"PING").await.unwrap(); + let mut byte = [0]; + let result = timeout(Duration::from_secs(3), client.read(&mut byte)) + .await + .unwrap(); + assert!(matches!(result, Ok(0) | Err(_))); + assert!( + timeout(Duration::from_millis(50), target.accept()) + .await + .is_err() + ); + } +} diff --git a/codex-rs/network-proxy/src/certs.rs b/codex-rs/network-proxy/src/certs.rs new file mode 100644 index 0000000000000000000000000000000000000000..9e6578e52c0f08f19cb1c9af0acb8a6f9b28ef7f --- /dev/null +++ b/codex-rs/network-proxy/src/certs.rs @@ -0,0 +1,1013 @@ +use anyhow::Context as _; +use anyhow::Result; +use anyhow::anyhow; +use base64::Engine as _; +use codex_utils_home_dir::find_codex_home; +use rama_net::tls::ApplicationProtocol; +use rama_tls_rustls::dep::pki_types::CertificateDer; +use rama_tls_rustls::dep::pki_types::PrivateKeyDer; +use rama_tls_rustls::dep::pki_types::pem::PemObject; +use rama_tls_rustls::dep::rcgen::BasicConstraints; +use rama_tls_rustls::dep::rcgen::CertificateParams; +use rama_tls_rustls::dep::rcgen::DistinguishedName; +use rama_tls_rustls::dep::rcgen::DnType; +use rama_tls_rustls::dep::rcgen::ExtendedKeyUsagePurpose; +use rama_tls_rustls::dep::rcgen::IsCa; +use rama_tls_rustls::dep::rcgen::Issuer; +use rama_tls_rustls::dep::rcgen::KeyPair; +use rama_tls_rustls::dep::rcgen::KeyUsagePurpose; +use rama_tls_rustls::dep::rcgen::PKCS_ECDSA_P256_SHA256; +use rama_tls_rustls::dep::rcgen::SanType; +use rama_tls_rustls::dep::rustls; +use rama_tls_rustls::server::TlsAcceptorData; +use sha2::Digest as _; +use sha2::Sha256; +use std::collections::HashMap; +use std::collections::HashSet; +use std::fs; +use std::fs::File; +use std::fs::OpenOptions; +use std::io::Write; +use std::net::IpAddr; +use std::path::Path; +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::LazyLock; +use std::sync::Mutex; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; +use tracing::info; +use tracing::warn; + +pub(super) struct ManagedMitmCa { + issuer: Issuer<'static, KeyPair>, + certificate_path: PathBuf, + _artifact_lease: File, +} + +static MANAGED_MITM_CAS: LazyLock>>> = + LazyLock::new(|| Mutex::new(HashMap::new())); + +impl ManagedMitmCa { + pub(super) fn load_or_create() -> Result> { + let proxy_dir = managed_ca_dir()?; + let mut managed_cas = MANAGED_MITM_CAS + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(ca) = managed_cas.get(&proxy_dir) { + return Ok(ca.clone()); + } + + let ca = Arc::new(Self::create(&proxy_dir)?); + managed_cas.insert(proxy_dir, ca.clone()); + Ok(ca) + } + + fn create(proxy_dir: &Path) -> Result { + fs::create_dir_all(proxy_dir) + .with_context(|| format!("failed to create {}", proxy_dir.display()))?; + + let (certificate_pem, private_key) = generate_ca()?; + let artifact_lock = match lock_managed_ca_artifacts(proxy_dir) { + Ok(lock) => Some(lock), + Err(err) => { + warn!("failed to lock managed MITM CA artifacts; skipping pruning: {err}"); + None + } + }; + let certificate_path = persist_managed_ca_certificate(proxy_dir, &certificate_pem)?; + let issuer = Issuer::from_ca_cert_pem(&certificate_pem, private_key) + .context("failed to parse managed MITM CA certificate")?; + let artifact_lease = lock_managed_ca_certificate(&certificate_path)?; + if artifact_lock.is_some() { + prune_managed_ca_artifacts(proxy_dir); + } + info!( + cert_path = %certificate_path.display(), + "generated process-local MITM CA" + ); + Ok(Self { + issuer, + certificate_path, + _artifact_lease: artifact_lease, + }) + } + + fn certificate_path(&self) -> &Path { + &self.certificate_path + } + + pub(super) fn tls_acceptor_data_for_host(&self, host: &str) -> Result { + let (cert_pem, key_pem) = issue_host_certificate_pem(host, &self.issuer)?; + let cert = CertificateDer::from_pem_slice(cert_pem.as_bytes()) + .context("failed to parse host cert PEM")?; + let key = PrivateKeyDer::from_pem_slice(key_pem.as_bytes()) + .context("failed to parse host key PEM")?; + let mut server_config = + rustls::ServerConfig::builder_with_protocol_versions(rustls::ALL_VERSIONS) + .with_no_client_auth() + .with_single_cert(vec![cert], key) + .context("failed to build rustls server config")?; + server_config.alpn_protocols = vec![ + ApplicationProtocol::HTTP_2.as_bytes().to_vec(), + ApplicationProtocol::HTTP_11.as_bytes().to_vec(), + ]; + + Ok(TlsAcceptorData::from(server_config)) + } +} + +fn issue_host_certificate_pem( + host: &str, + issuer: &Issuer<'_, KeyPair>, +) -> Result<(String, String)> { + let mut params = if let Ok(ip) = host.parse::() { + let mut params = CertificateParams::new(Vec::new()) + .map_err(|err| anyhow!("failed to create cert params: {err}"))?; + params.subject_alt_names.push(SanType::IpAddress(ip)); + params + } else { + CertificateParams::new(vec![host.to_string()]) + .map_err(|err| anyhow!("failed to create cert params: {err}"))? + }; + + params.extended_key_usages = vec![ExtendedKeyUsagePurpose::ServerAuth]; + params.key_usages = vec![ + KeyUsagePurpose::DigitalSignature, + KeyUsagePurpose::KeyEncipherment, + ]; + + let key_pair = KeyPair::generate_for(&PKCS_ECDSA_P256_SHA256) + .map_err(|err| anyhow!("failed to generate host key pair: {err}"))?; + let cert = params + .signed_by(&key_pair, issuer) + .map_err(|err| anyhow!("failed to sign host cert: {err}"))?; + + Ok((cert.pem(), key_pair.serialize_pem())) +} + +const MANAGED_MITM_CA_DIR: &str = "proxy"; +const MANAGED_MITM_CA_ARTIFACT_LOCK: &str = ".artifacts.lock"; +const MANAGED_MITM_CA_CERT_PREFIX: &str = "ca"; +const MANAGED_MITM_CA_TRUST_BUNDLE_PREFIX: &str = "ca-bundle"; +pub(crate) const SSL_CERT_DIR_ENV_KEY: &str = "SSL_CERT_DIR"; + +// Best-effort compatibility set for common child toolchains that accept a CA bundle path. +// This is intentionally curated rather than pretending to cover every TLS client. +pub const CUSTOM_CA_ENV_KEYS: [&str; 11] = [ + "CODEX_CA_CERTIFICATE", + "SSL_CERT_FILE", + "REQUESTS_CA_BUNDLE", + "CURL_CA_BUNDLE", + "NODE_EXTRA_CA_CERTS", + "GIT_SSL_CAINFO", + "CARGO_HTTP_CAINFO", + "PIP_CERT", + "BUNDLE_SSL_CA_CERT", + "npm_config_cafile", + "NPM_CONFIG_CAFILE", +]; + +pub(crate) fn ca_env_from_process() -> HashMap<&'static str, String> { + CUSTOM_CA_ENV_KEYS + .into_iter() + .chain([SSL_CERT_DIR_ENV_KEY]) + .filter_map(|key| std::env::var(key).ok().map(|value| (key, value))) + .collect() +} + +/// Immutable managed MITM CA bundle path plus startup TLS env values. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ManagedMitmCaTrustBundle { + pub(crate) path: PathBuf, + pub(crate) startup_env_values: HashMap<&'static str, String>, +} + +fn managed_ca_dir() -> Result { + let codex_home = + find_codex_home().context("failed to resolve CODEX_HOME for managed MITM CA")?; + Ok(codex_home.join(MANAGED_MITM_CA_DIR).to_path_buf()) +} + +pub(crate) fn managed_ca_trust_bundle( + env: &HashMap<&'static str, String>, +) -> Result { + let ca = ManagedMitmCa::load_or_create()?; + managed_ca_trust_bundle_for_cert_path(ca.certificate_path(), env) +} + +fn managed_ca_trust_bundle_for_cert_path( + cert_path: &Path, + env: &HashMap<&'static str, String>, +) -> Result { + let startup_env_values = startup_ca_file_env_values(env); + let startup_cert_dir = env + .get(SSL_CERT_DIR_ENV_KEY) + .filter(|value| !value.is_empty()) + .map(String::as_str); + let trust_bundle = + build_managed_ca_trust_bundle(cert_path, &startup_env_values, startup_cert_dir)?; + let path = persist_managed_ca_trust_bundle(cert_path, &trust_bundle)?; + + Ok(ManagedMitmCaTrustBundle { + path, + startup_env_values, + }) +} + +pub(crate) fn upstream_tls_root_store( + env: &HashMap<&'static str, String>, +) -> Result> { + let ca = ManagedMitmCa::load_or_create()?; + upstream_tls_root_store_for_cert_path(ca.certificate_path(), env) +} + +pub(crate) fn upstream_tls_root_store_for_cert_path( + managed_ca_cert_path: &Path, + env: &HashMap<&'static str, String>, +) -> Result> { + let startup_env_values = startup_ca_file_env_values(env); + let startup_cert_dir = env + .get(SSL_CERT_DIR_ENV_KEY) + .filter(|value| !value.is_empty()) + .map(String::as_str); + let certificates = load_platform_and_startup_root_certificates( + managed_ca_cert_path, + &startup_env_values, + startup_cert_dir, + )?; + let mut roots = rustls::RootCertStore::empty(); + let (_, ignored) = roots.add_parsable_certificates(certificates); + if ignored > 0 { + warn!( + ignored_root_count = ignored, + "ignored invalid platform or startup roots for MITM upstream TLS" + ); + } + Ok(Arc::new(roots)) +} + +fn startup_ca_file_env_values( + env: &HashMap<&'static str, String>, +) -> HashMap<&'static str, String> { + CUSTOM_CA_ENV_KEYS + .into_iter() + .filter_map(|key| { + env.get(key) + .filter(|value| !value.is_empty()) + .map(|value| (key, value.clone())) + }) + .collect() +} + +fn build_managed_ca_trust_bundle( + managed_ca_cert_path: &Path, + startup_env_values: &HashMap<&'static str, String>, + startup_cert_dir: Option<&str>, +) -> Result { + let mut trust_bundle = String::new(); + for cert in load_platform_and_startup_root_certificates( + managed_ca_cert_path, + startup_env_values, + startup_cert_dir, + )? { + push_certificate_pem(&mut trust_bundle, cert.as_ref()); + } + append_pem_file(&mut trust_bundle, managed_ca_cert_path)?; + Ok(trust_bundle) +} + +fn load_platform_and_startup_root_certificates( + managed_ca_cert_path: &Path, + startup_env_values: &HashMap<&'static str, String>, + startup_cert_dir: Option<&str>, +) -> Result>> { + let managed_ca_cert = fs::read(managed_ca_cert_path).with_context(|| { + format!( + "failed to read managed MITM CA certificate: {}", + managed_ca_cert_path.display() + ) + })?; + let managed_ca_cert = CertificateDer::from_pem_slice(&managed_ca_cert) + .context("failed to parse managed MITM CA certificate")?; + let rustls_native_certs::CertificateResult { certs, errors, .. } = + crate::native_certs::load_platform_native_certs(); + if !errors.is_empty() { + warn!( + native_root_error_count = errors.len(), + "encountered errors while loading native root certificates for MITM trust bundle" + ); + } + let mut certificates = certs; + let mut appended_startup_paths = HashSet::new(); + for path in CUSTOM_CA_ENV_KEYS + .into_iter() + .filter_map(|key| startup_env_values.get(key)) + .map(PathBuf::from) + { + if path != managed_ca_cert_path + && !is_current_generated_trust_bundle_path(&path, managed_ca_cert_path) + && appended_startup_paths.insert(path.clone()) + { + certificates.extend(read_ca_certificates(&path)?); + } + } + if let Some(startup_cert_dir) = startup_cert_dir { + for path in std::env::split_paths(startup_cert_dir) { + if appended_startup_paths.insert(path.clone()) { + certificates.extend(load_ca_directory_certificates(&path)); + } + } + } + let mut seen = HashSet::new(); + certificates.retain(|cert| cert != &managed_ca_cert && seen.insert(cert.as_ref().to_vec())); + Ok(certificates) +} + +fn read_ca_certificates(path: &Path) -> Result>> { + let pem = fs::read(path) + .with_context(|| format!("failed to read startup CA bundle: {}", path.display()))?; + let pem = String::from_utf8_lossy(&pem); + let contains_trusted_certificates = pem.contains("TRUSTED CERTIFICATE"); + let normalized_pem = pem + .replace("BEGIN TRUSTED CERTIFICATE", "BEGIN CERTIFICATE") + .replace("END TRUSTED CERTIFICATE", "END CERTIFICATE"); + let certs = CertificateDer::pem_slice_iter(normalized_pem.as_bytes()) + .collect::, _>>() + .with_context(|| format!("failed to parse startup CA bundle: {}", path.display()))?; + if certs.is_empty() { + return Err(anyhow!( + "startup CA bundle contained no certificates: {}", + path.display() + )); + } + certs + .into_iter() + .map(|cert| { + let cert = if contains_trusted_certificates { + first_der_item(cert.as_ref()).ok_or_else(|| { + anyhow!( + "startup CA bundle contained an invalid trusted certificate: {}", + path.display() + ) + })? + } else { + cert.as_ref() + }; + Ok(CertificateDer::from(cert.to_vec())) + }) + .collect() +} + +fn load_ca_directory_certificates(path: &Path) -> Vec> { + let rustls_native_certs::CertificateResult { certs, errors, .. } = + rustls_native_certs::load_certs_from_paths(None, Some(path)); + if !errors.is_empty() { + warn!( + ca_path = %path.display(), + ca_error_count = errors.len(), + "encountered errors while loading startup CA directory" + ); + } + certs +} + +fn first_der_item(der: &[u8]) -> Option<&[u8]> { + der_item_length(der).map(|length| &der[..length]) +} + +fn der_item_length(der: &[u8]) -> Option { + let &length_octet = der.get(1)?; + if length_octet & 0x80 == 0 { + return Some(2 + usize::from(length_octet)).filter(|length| *length <= der.len()); + } + + let length_octets = usize::from(length_octet & 0x7f); + if length_octets == 0 { + return None; + } + + let length_end = 2usize.checked_add(length_octets)?; + let mut content_length = 0usize; + for &byte in der.get(2..length_end)? { + content_length = content_length + .checked_mul(256)? + .checked_add(usize::from(byte))?; + } + length_end + .checked_add(content_length) + .filter(|length| *length <= der.len()) +} + +fn is_current_generated_trust_bundle_path(path: &Path, managed_ca_cert_path: &Path) -> bool { + let Some(proxy_dir) = managed_ca_cert_path.parent() else { + return false; + }; + if is_generated_trust_bundle_path(path, proxy_dir) { + return true; + } + let Some(file_name) = path.file_name().and_then(|file_name| file_name.to_str()) else { + return false; + }; + if path.parent() != Some(proxy_dir) + || !file_name.starts_with(MANAGED_MITM_CA_TRUST_BUNDLE_PREFIX) + || !file_name.ends_with(".pem") + { + return false; + } + let Ok(trust_bundle) = fs::read(path) else { + return false; + }; + let Ok(managed_ca_cert) = fs::read(managed_ca_cert_path) else { + return false; + }; + !managed_ca_cert.is_empty() + && trust_bundle + .windows(managed_ca_cert.len()) + .any(|window| window == managed_ca_cert) +} + +fn is_generated_trust_bundle_path(path: &Path, proxy_dir: &Path) -> bool { + is_generated_managed_ca_artifact_path(path, proxy_dir, MANAGED_MITM_CA_TRUST_BUNDLE_PREFIX) +} + +fn is_generated_managed_ca_artifact_path(path: &Path, proxy_dir: &Path, prefix: &str) -> bool { + let Some(file_name) = path.file_name().and_then(|file_name| file_name.to_str()) else { + return false; + }; + let Some(expected_hash) = file_name + .strip_prefix(prefix) + .and_then(|suffix| suffix.strip_prefix('-')) + .and_then(|suffix| suffix.strip_suffix(".pem")) + else { + return false; + }; + if path.parent() != Some(proxy_dir) + || expected_hash.len() != 64 + || !expected_hash.bytes().all(|byte| byte.is_ascii_hexdigit()) + { + return false; + } + let Ok(trust_bundle) = fs::read(path) else { + return false; + }; + format!("{:x}", Sha256::digest(trust_bundle)) == expected_hash +} + +/// Returns whether `path` points at a current Codex-generated MITM CA bundle. +pub fn is_managed_mitm_ca_trust_bundle_path(path: &str) -> bool { + let Ok(proxy_dir) = managed_ca_dir() else { + return false; + }; + is_generated_trust_bundle_path(Path::new(path), &proxy_dir) +} + +fn persist_managed_ca_trust_bundle( + managed_ca_cert_path: &Path, + trust_bundle: &str, +) -> Result { + let proxy_dir = managed_ca_cert_path + .parent() + .ok_or_else(|| anyhow!("managed MITM CA cert path is missing a parent"))?; + fs::create_dir_all(proxy_dir) + .with_context(|| format!("failed to create {}", proxy_dir.display()))?; + let hash = Sha256::digest(trust_bundle.as_bytes()); + let trust_bundle_path = proxy_dir.join(format!( + "{MANAGED_MITM_CA_TRUST_BUNDLE_PREFIX}-{hash:x}.pem" + )); + write_atomic_create_new_or_reuse( + &trust_bundle_path, + trust_bundle.as_bytes(), + /*mode*/ 0o644, + ) + .with_context(|| { + format!( + "failed to persist managed MITM CA trust bundle {}", + trust_bundle_path.display() + ) + })?; + Ok(trust_bundle_path) +} + +fn append_pem_file(bundle: &mut String, path: &Path) -> Result<()> { + if !bundle.ends_with('\n') { + bundle.push('\n'); + } + let pem = fs::read_to_string(path) + .with_context(|| format!("failed to read CA bundle {}", path.display()))?; + bundle.push_str(&pem); + if !bundle.ends_with('\n') { + bundle.push('\n'); + } + Ok(()) +} + +fn push_certificate_pem(bundle: &mut String, der: &[u8]) { + bundle.push_str("-----BEGIN CERTIFICATE-----\n"); + let encoded = base64::engine::general_purpose::STANDARD.encode(der); + for chunk in encoded.as_bytes().chunks(64) { + bundle.push_str(&String::from_utf8_lossy(chunk)); + bundle.push('\n'); + } + bundle.push_str("-----END CERTIFICATE-----\n"); +} + +fn persist_managed_ca_certificate(proxy_dir: &Path, cert_pem: &str) -> Result { + let hash = Sha256::digest(cert_pem.as_bytes()); + let cert_path = proxy_dir.join(format!("{MANAGED_MITM_CA_CERT_PREFIX}-{hash:x}.pem")); + write_atomic_create_new_or_reuse(&cert_path, cert_pem.as_bytes(), /*mode*/ 0o644) + .with_context(|| { + format!( + "failed to persist managed MITM CA certificate {}", + cert_path.display() + ) + })?; + Ok(cert_path) +} + +fn lock_managed_ca_certificate(certificate_path: &Path) -> Result { + let lock_path = managed_ca_certificate_lock_path(certificate_path) + .ok_or_else(|| anyhow!("managed MITM CA certificate path is missing a file name"))?; + let file = open_managed_ca_lock(&lock_path)?; + file.lock_shared() + .with_context(|| format!("failed to lock {}", lock_path.display()))?; + Ok(file) +} + +fn lock_managed_ca_artifacts(proxy_dir: &Path) -> Result { + let lock_path = proxy_dir.join(MANAGED_MITM_CA_ARTIFACT_LOCK); + let file = open_managed_ca_lock(&lock_path)?; + file.lock() + .with_context(|| format!("failed to lock {}", lock_path.display()))?; + Ok(file) +} + +fn managed_ca_certificate_lock_path(certificate_path: &Path) -> Option { + let file_name = certificate_path.file_name()?.to_string_lossy(); + Some(certificate_path.with_file_name(format!(".{file_name}.lock"))) +} + +fn open_managed_ca_lock(path: &Path) -> Result { + if fs::symlink_metadata(path) + .ok() + .is_some_and(|metadata| metadata.file_type().is_symlink()) + { + return Err(anyhow!( + "refusing to use symlink lock file {}", + path.display() + )); + } + + #[cfg(unix)] + use std::os::unix::fs::OpenOptionsExt; + + let mut options = OpenOptions::new(); + options.read(true).write(true).create(true).truncate(false); + #[cfg(unix)] + options.mode(0o600); + options + .open(path) + .with_context(|| format!("failed to open {}", path.display())) +} + +fn prune_managed_ca_artifacts(proxy_dir: &Path) { + for certificate_path in + generated_managed_ca_artifact_paths(proxy_dir, MANAGED_MITM_CA_CERT_PREFIX) + { + remove_inactive_managed_ca_certificate(&certificate_path); + } + + let remaining_certificates = + generated_managed_ca_artifact_paths(proxy_dir, MANAGED_MITM_CA_CERT_PREFIX) + .into_iter() + .filter_map(|path| fs::read(path).ok()) + .filter(|certificate| !certificate.is_empty()) + .collect::>(); + let bundle_paths = + generated_managed_ca_artifact_paths(proxy_dir, MANAGED_MITM_CA_TRUST_BUNDLE_PREFIX); + for bundle_path in bundle_paths { + let Ok(contents) = fs::read(&bundle_path) else { + continue; + }; + if remaining_certificates.iter().any(|certificate| { + contents + .windows(certificate.len()) + .any(|window| window == certificate) + }) { + continue; + } + if let Err(err) = fs::remove_file(&bundle_path) + && err.kind() != std::io::ErrorKind::NotFound + { + warn!( + path = %bundle_path.display(), + "failed to prune stale managed MITM CA trust bundle: {err}" + ); + } + } +} + +fn generated_managed_ca_artifact_paths(proxy_dir: &Path, prefix: &str) -> Vec { + let Ok(entries) = fs::read_dir(proxy_dir) else { + return Vec::new(); + }; + entries + .filter_map(std::result::Result::ok) + .filter_map(|entry| { + let path = entry.path(); + if !is_generated_managed_ca_artifact_path(&path, proxy_dir, prefix) { + return None; + } + Some(path) + }) + .collect() +} + +fn remove_inactive_managed_ca_certificate(certificate_path: &Path) { + let Some(lock_path) = managed_ca_certificate_lock_path(certificate_path) else { + return; + }; + let Ok(lock_file) = open_managed_ca_lock(&lock_path) else { + return; + }; + match lock_file.try_lock() { + Ok(()) => {} + Err(std::fs::TryLockError::WouldBlock) => return, + Err(err) => { + warn!( + path = %lock_path.display(), + "failed to inspect managed MITM CA artifact lease: {err}" + ); + return; + } + } + + let removed = match fs::remove_file(certificate_path) { + Ok(()) => true, + Err(err) if err.kind() == std::io::ErrorKind::NotFound => true, + Err(err) => { + warn!( + path = %certificate_path.display(), + "failed to prune stale managed MITM CA certificate: {err}" + ); + false + } + }; + drop(lock_file); + if removed + && let Err(err) = fs::remove_file(&lock_path) + && err.kind() != std::io::ErrorKind::NotFound + { + warn!( + path = %lock_path.display(), + "failed to prune stale managed MITM CA artifact lease: {err}" + ); + } +} + +fn generate_ca() -> Result<(String, KeyPair)> { + let mut params = CertificateParams::default(); + params.is_ca = IsCa::Ca(BasicConstraints::Unconstrained); + params.key_usages = vec![ + KeyUsagePurpose::KeyCertSign, + KeyUsagePurpose::DigitalSignature, + KeyUsagePurpose::KeyEncipherment, + ]; + let mut dn = DistinguishedName::new(); + dn.push(DnType::CommonName, "network_proxy MITM CA"); + params.distinguished_name = dn; + + let key_pair = KeyPair::generate_for(&PKCS_ECDSA_P256_SHA256) + .map_err(|err| anyhow!("failed to generate CA key pair: {err}"))?; + let cert = params + .self_signed(&key_pair) + .map_err(|err| anyhow!("failed to generate CA cert: {err}"))?; + Ok((cert.pem(), key_pair)) +} + +fn write_atomic_create_new(path: &Path, contents: &[u8], mode: u32) -> Result<()> { + let parent = path + .parent() + .ok_or_else(|| anyhow!("missing parent directory"))?; + + let nanos = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(); + let pid = std::process::id(); + let file_name = path.file_name().unwrap_or_default().to_string_lossy(); + let tmp_path = parent.join(format!(".{file_name}.tmp.{pid}.{nanos}")); + + let mut file = open_create_new_with_mode(&tmp_path, mode)?; + file.write_all(contents) + .with_context(|| format!("failed to write {}", tmp_path.display()))?; + file.sync_all() + .with_context(|| format!("failed to fsync {}", tmp_path.display()))?; + drop(file); + + // Create the final file using "create-new" semantics (no overwrite). `rename` on Unix can + // overwrite existing files, so prefer a hard-link, which fails if the destination exists. + match fs::hard_link(&tmp_path, path) { + Ok(()) => { + fs::remove_file(&tmp_path) + .with_context(|| format!("failed to remove {}", tmp_path.display()))?; + } + Err(err) if err.kind() == std::io::ErrorKind::AlreadyExists => { + let _ = fs::remove_file(&tmp_path); + return Err(anyhow!( + "refusing to overwrite existing file {}", + path.display() + )); + } + Err(_) => { + // Best-effort fallback for environments where hard links are not supported. + // This is still subject to a TOCTOU race, but the typical case is a private per-user + // config directory, where other users cannot create files anyway. + if path.exists() { + let _ = fs::remove_file(&tmp_path); + return Err(anyhow!( + "refusing to overwrite existing file {}", + path.display() + )); + } + fs::rename(&tmp_path, path).with_context(|| { + format!( + "failed to rename {} -> {}", + tmp_path.display(), + path.display() + ) + })?; + } + } + + sync_parent_dir(parent)?; + + Ok(()) +} + +#[cfg(not(windows))] +fn sync_parent_dir(parent: &Path) -> Result<()> { + // Best-effort durability: ensure the directory entry is persisted too. + let dir = File::open(parent).with_context(|| format!("failed to open {}", parent.display()))?; + dir.sync_all() + .with_context(|| format!("failed to fsync {}", parent.display())) +} + +#[cfg(windows)] +fn sync_parent_dir(_parent: &Path) -> Result<()> { + Ok(()) +} + +fn write_atomic_create_new_or_reuse(path: &Path, contents: &[u8], mode: u32) -> Result<()> { + if fs::symlink_metadata(path) + .ok() + .is_some_and(|metadata| metadata.file_type().is_symlink()) + { + return Err(anyhow!("refusing to reuse symlink {}", path.display())); + } + if fs::read(path).ok().as_deref() == Some(contents) { + return Ok(()); + } + if path.exists() { + return Err(anyhow!( + "refusing to reuse existing mismatched file {}", + path.display() + )); + } + match write_atomic_create_new(path, contents, mode) { + Ok(()) => Ok(()), + Err(_err) if fs::read(path).ok().as_deref() == Some(contents) => Ok(()), + Err(err) => Err(err), + } +} + +#[cfg(unix)] +fn open_create_new_with_mode(path: &Path, mode: u32) -> Result { + use std::os::unix::fs::OpenOptionsExt; + + OpenOptions::new() + .write(true) + .create_new(true) + .mode(mode) + .open(path) + .with_context(|| format!("failed to create {}", path.display())) +} + +#[cfg(not(unix))] +fn open_create_new_with_mode(path: &Path, _mode: u32) -> Result { + OpenOptions::new() + .write(true) + .create_new(true) + .open(path) + .with_context(|| format!("failed to create {}", path.display())) +} + +#[cfg(test)] +mod tests { + use super::*; + + use codex_utils_rustls_provider::ensure_rustls_crypto_provider; + use pretty_assertions::assert_eq; + use tempfile::tempdir; + + #[test] + fn managed_ca_private_key_is_not_persisted() { + ensure_rustls_crypto_provider(); + let dir = tempdir().unwrap(); + let ca = ManagedMitmCa::create(dir.path()).unwrap(); + ca.tls_acceptor_data_for_host("example.com").unwrap(); + let mut persisted_files = fs::read_dir(dir.path()) + .unwrap() + .map(|entry| entry.unwrap().path()) + .collect::>(); + persisted_files.sort(); + let mut expected_files = vec![ + ca.certificate_path().to_path_buf(), + managed_ca_certificate_lock_path(ca.certificate_path()).unwrap(), + dir.path().join(MANAGED_MITM_CA_ARTIFACT_LOCK), + ]; + expected_files.sort(); + + assert_eq!(persisted_files, expected_files); + assert_eq!( + fs::read(managed_ca_certificate_lock_path(ca.certificate_path()).unwrap()).unwrap(), + Vec::::new() + ); + } + + #[test] + fn managed_ca_artifact_pruning_preserves_only_active_certificates() { + let dir = tempdir().unwrap(); + let mut artifacts = Vec::new(); + let mut active_lease = None; + for index in 0..3 { + let certificate = format!("certificate {index}\n"); + let certificate_path = + persist_managed_ca_certificate(dir.path(), &certificate).unwrap(); + let lease = lock_managed_ca_certificate(&certificate_path).unwrap(); + if index == 0 { + active_lease = Some(lease); + } else { + drop(lease); + } + let bundle_path = persist_managed_ca_trust_bundle( + &certificate_path, + &format!("roots\n{certificate}"), + ) + .unwrap(); + artifacts.push((certificate_path, bundle_path)); + } + let unrelated_path = dir.path().join("ca-user.pem"); + fs::write(&unrelated_path, "user managed").unwrap(); + + prune_managed_ca_artifacts(dir.path()); + + let remaining_certificate_count = + generated_managed_ca_artifact_paths(dir.path(), MANAGED_MITM_CA_CERT_PREFIX).len(); + assert_eq!(remaining_certificate_count, 1); + assert!(artifacts[0].0.exists()); + assert!(artifacts[0].1.exists()); + assert!(!artifacts[1].0.exists()); + assert!(!artifacts[1].1.exists()); + assert!(!artifacts[2].0.exists()); + assert!(!artifacts[2].1.exists()); + assert!(unrelated_path.exists()); + + drop(active_lease.take()); + prune_managed_ca_artifacts(dir.path()); + + let remaining_certificates = + generated_managed_ca_artifact_paths(dir.path(), MANAGED_MITM_CA_CERT_PREFIX); + assert!(remaining_certificates.is_empty()); + assert!(!artifacts[0].0.exists()); + assert!(!artifacts[0].1.exists()); + } + + #[test] + fn current_generated_trust_bundle_path_rejects_stale_bundle() { + let dir = tempdir().unwrap(); + let managed_ca_cert_path = dir.path().join("ca.pem"); + let trust_bundle_path = dir.path().join("ca-bundle-123.pem"); + fs::write(&managed_ca_cert_path, "managed ca\n").unwrap(); + fs::write(&trust_bundle_path, "stale managed bundle\n").unwrap(); + assert!(!is_current_generated_trust_bundle_path( + &trust_bundle_path, + &managed_ca_cert_path, + )); + } + + #[test] + fn generated_trust_bundle_path_requires_matching_content_hash() { + let dir = tempdir().unwrap(); + let managed_ca_cert_path = dir.path().join("ca.pem"); + let trust_bundle_path = + persist_managed_ca_trust_bundle(&managed_ca_cert_path, "trusted roots").unwrap(); + + assert!(is_generated_trust_bundle_path( + &trust_bundle_path, + dir.path() + )); + fs::write(&trust_bundle_path, "tampered roots").unwrap(); + assert!(!is_generated_trust_bundle_path( + &trust_bundle_path, + dir.path() + )); + } + + #[test] + fn managed_ca_trust_bundle_appends_startup_file_and_directory_certificates() { + let dir = tempdir().unwrap(); + let managed_ca_cert_path = dir.path().join("ca.pem"); + let startup_ca_bundle_path = dir.path().join("startup-ca.pem"); + let startup_ca_dir = dir.path().join("startup-certs"); + let (managed_ca_cert, _) = generate_ca().unwrap(); + let (startup_ca_cert, startup_ca_key) = generate_ca().unwrap(); + let startup_ca_key = startup_ca_key.serialize_pem(); + let (directory_ca_cert, _) = generate_ca().unwrap(); + let mut trusted_ca_der = CertificateDer::from_pem_slice(startup_ca_cert.as_bytes()) + .unwrap() + .as_ref() + .to_vec(); + trusted_ca_der.extend_from_slice(&[0x30, 0x00]); + let mut trusted_ca_cert = String::new(); + push_certificate_pem(&mut trusted_ca_cert, &trusted_ca_der); + let trusted_ca_cert = trusted_ca_cert.replace("CERTIFICATE", "TRUSTED CERTIFICATE"); + fs::write(&managed_ca_cert_path, &managed_ca_cert).unwrap(); + fs::write( + &startup_ca_bundle_path, + format!("{trusted_ca_cert}{startup_ca_key}"), + ) + .unwrap(); + fs::create_dir(&startup_ca_dir).unwrap(); + fs::write(startup_ca_dir.join("directory-ca.pem"), &directory_ca_cert).unwrap(); + let startup_ca_bundle_path = startup_ca_bundle_path.display().to_string(); + let env = HashMap::from([ + ("SSL_CERT_FILE", startup_ca_bundle_path.clone()), + (SSL_CERT_DIR_ENV_KEY, startup_ca_dir.display().to_string()), + ]); + + let trust_bundle = + managed_ca_trust_bundle_for_cert_path(&managed_ca_cert_path, &env).unwrap(); + assert_eq!( + trust_bundle.startup_env_values, + HashMap::from([("SSL_CERT_FILE", startup_ca_bundle_path)]) + ); + let baseline_bundle = fs::read_to_string(&trust_bundle.path).unwrap(); + let baseline_certs = CertificateDer::pem_slice_iter(baseline_bundle.as_bytes()) + .collect::, _>>() + .unwrap(); + let expected_certs = [&startup_ca_cert, &directory_ca_cert, &managed_ca_cert] + .map(|cert| CertificateDer::from_pem_slice(cert.as_bytes()).unwrap()); + + assert!( + expected_certs + .iter() + .all(|cert| baseline_certs.contains(cert)) + ); + assert!(!baseline_bundle.contains(&startup_ca_key)); + assert!(!baseline_bundle.contains("TRUSTED CERTIFICATE")); + } + + #[test] + fn managed_ca_trust_bundle_skips_inherited_current_bundle() { + let dir = tempdir().unwrap(); + let managed_ca_cert_path = dir.path().join("ca.pem"); + let inherited_bundle_path = dir.path().join("ca-bundle-parent.pem"); + let (managed_ca_cert, _) = generate_ca().unwrap(); + fs::write(&managed_ca_cert_path, &managed_ca_cert).unwrap(); + fs::write( + &inherited_bundle_path, + format!("parent roots\n{managed_ca_cert}"), + ) + .unwrap(); + let env = HashMap::from([( + "REQUESTS_CA_BUNDLE", + inherited_bundle_path.display().to_string(), + )]); + + let trust_bundle = + managed_ca_trust_bundle_for_cert_path(&managed_ca_cert_path, &env).unwrap(); + let baseline_bundle = fs::read_to_string(&trust_bundle.path).unwrap(); + + assert_eq!(baseline_bundle.matches(&managed_ca_cert).count(), 1); + } + + #[cfg(unix)] + #[test] + fn write_atomic_create_new_or_reuse_rejects_matching_symlink_target() { + use std::os::unix::fs::symlink; + + let dir = tempdir().unwrap(); + let target = dir.path().join("real-bundle.pem"); + let link = dir.path().join("ca-bundle.pem"); + fs::write(&target, "bundle").unwrap(); + symlink(&target, &link).unwrap(); + + let err = write_atomic_create_new_or_reuse(&link, b"bundle", /*mode*/ 0o644).unwrap_err(); + + assert_eq!( + err.to_string(), + format!("refusing to reuse symlink {}", link.display()) + ); + } +} diff --git a/codex-rs/network-proxy/src/config.rs b/codex-rs/network-proxy/src/config.rs new file mode 100644 index 0000000000000000000000000000000000000000..0a880282edc7d0a19c9c09803dca7285215ef048 --- /dev/null +++ b/codex-rs/network-proxy/src/config.rs @@ -0,0 +1,1111 @@ +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use codex_utils_absolute_path::AbsolutePathBuf; +use serde::Deserialize; +use serde::Deserializer; +use serde::Serialize; +use serde::Serializer; +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::net::IpAddr; +use std::net::SocketAddr; +use std::path::Path; +use tracing::warn; +use url::Url; + +use crate::mitm_hook::MitmHookConfig; +use crate::policy::normalize_host; + +/// Variant order encodes effective precedence for duplicate patterns: +/// `None < Allow < Deny`, so deny wins over allow when entries conflict. +#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, PartialOrd, Ord)] +#[serde(rename_all = "lowercase")] +pub enum NetworkDomainPermission { + None, + Allow, + Deny, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct NetworkDomainPermissionEntry { + pub pattern: String, + pub permission: NetworkDomainPermission, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct NetworkDomainPermissions { + pub entries: Vec, +} + +impl Serialize for NetworkDomainPermissions { + fn serialize(&self, serializer: S) -> std::result::Result + where + S: Serializer, + { + self.effective_entries() + .into_iter() + .map(|entry| (entry.pattern, entry.permission)) + .collect::>() + .serialize(serializer) + } +} + +impl<'de> Deserialize<'de> for NetworkDomainPermissions { + fn deserialize(deserializer: D) -> std::result::Result + where + D: Deserializer<'de>, + { + let entries = BTreeMap::::deserialize(deserializer)? + .into_iter() + .map(|(pattern, permission)| NetworkDomainPermissionEntry { + pattern, + permission, + }) + .collect(); + Ok(Self { entries }) + } +} + +impl NetworkDomainPermissions { + fn effective_entries(&self) -> Vec { + let mut order = Vec::new(); + let mut effective_permissions = BTreeMap::new(); + + for entry in &self.entries { + if !effective_permissions.contains_key(&entry.pattern) { + order.push(entry.pattern.clone()); + } + + let permission = effective_permissions + .entry(entry.pattern.clone()) + .or_insert(entry.permission); + if entry.permission > *permission { + *permission = entry.permission; + } + } + + order + .into_iter() + .filter_map(|pattern| { + effective_permissions.remove(&pattern).map(|permission| { + NetworkDomainPermissionEntry { + pattern, + permission, + } + }) + }) + .collect() + } +} + +#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "lowercase")] +pub enum NetworkUnixSocketPermission { + Allow, + Deny, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)] +pub struct NetworkUnixSocketPermissions { + #[serde(flatten)] + pub entries: BTreeMap, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(default)] +pub struct NetworkProxyConfig { + #[serde(default)] + pub enabled: bool, + #[serde(default = "default_proxy_url")] + pub proxy_url: String, + pub enable_socks5: bool, + #[serde(default = "default_socks_url")] + pub socks_url: String, + pub enable_socks5_udp: bool, + pub allow_upstream_proxy: bool, + #[serde(default)] + pub dangerously_allow_non_loopback_proxy: bool, + #[serde(default)] + pub dangerously_allow_all_unix_sockets: bool, + #[serde(default)] + pub mode: NetworkMode, + #[serde(default)] + pub domains: Option, + #[serde(default)] + pub unix_sockets: Option, + pub allow_local_binding: bool, + #[serde(default)] + pub mitm: bool, + #[serde(default)] + pub credential_broker: bool, + /// Whether brokerage enabled MITM rather than inheriting an explicit setting. + #[serde(skip)] + pub credential_broker_enabled_mitm: bool, + #[serde(default, skip_serializing_if = "BTreeMap::is_empty")] + pub credential_providers: BTreeMap, + /// Trusted OpenAI endpoint derived from local configuration, never sent to remote executors. + #[serde(skip)] + pub credential_broker_openai_host: Option, + /// Trusted local destination context, never sent to remote executors or child environments. + #[serde(skip)] + pub credential_broker_context: crate::CredentialBrokerContext, + #[serde(default)] + pub dangerously_allow_plaintext_credential_injection: bool, + #[serde(default)] + pub mitm_hooks: Vec, +} + +impl Default for NetworkProxyConfig { + fn default() -> Self { + Self { + enabled: false, + proxy_url: default_proxy_url(), + enable_socks5: true, + socks_url: default_socks_url(), + enable_socks5_udp: true, + allow_upstream_proxy: true, + dangerously_allow_non_loopback_proxy: false, + dangerously_allow_all_unix_sockets: false, + mode: NetworkMode::default(), + domains: None, + unix_sockets: None, + allow_local_binding: false, + mitm: false, + credential_broker: false, + credential_broker_enabled_mitm: false, + credential_providers: BTreeMap::new(), + credential_broker_openai_host: None, + credential_broker_context: crate::CredentialBrokerContext::default(), + dangerously_allow_plaintext_credential_injection: false, + mitm_hooks: Vec::new(), + } + } +} + +impl NetworkProxyConfig { + pub fn set_credential_broker_enabled(&mut self, enabled: bool) { + self.credential_broker = enabled; + if enabled { + self.credential_broker_enabled_mitm |= !self.mitm; + self.mitm = true; + } else if self.credential_broker_enabled_mitm { + self.mitm = !self.mitm_hooks.is_empty(); + self.credential_broker_enabled_mitm = false; + } + } + + pub fn set_credential_broker_openai_base_url(&mut self, base_url: Option<&str>) { + self.credential_broker_openai_host = base_url.and_then(trusted_credential_broker_host); + } + + /// Retains trusted destination context without changing child environment policy. Conflicting + /// case-insensitive provider overrides disable brokerage on Windows. + pub fn configure_credential_broker_environment( + &mut self, + environment: &HashMap, + ) { + if cfg!(windows) + && self.credential_broker + && self.has_ambiguous_windows_credential_environment(environment) + { + warn!( + "credential brokerage disabled because shell environment overrides contain \ + conflicting case-insensitive provider keys" + ); + self.set_credential_broker_enabled(/*enabled*/ false); + } + self.credential_broker_context = if self.credential_broker { + crate::CredentialBrokerContext::capture(self, environment) + } else { + crate::CredentialBrokerContext::default() + }; + } + + fn has_ambiguous_windows_credential_environment( + &self, + environment: &HashMap, + ) -> bool { + environment.iter().any(|(key, value)| { + let is_provider_key = + crate::credential_broker::is_credential_broker_provider_env_key(key) + || self.credential_providers.values().any(|provider| { + provider + .env + .iter() + .chain(provider.url_prefix_from_env.iter()) + .any(|candidate| key.eq_ignore_ascii_case(candidate)) + }); + is_provider_key + && environment.iter().any(|(candidate, candidate_value)| { + key != candidate + && key.eq_ignore_ascii_case(candidate) + && value != candidate_value + }) + }) + } + + pub fn allowed_domains(&self) -> Option> { + self.domain_entries(NetworkDomainPermission::Allow) + } + + pub fn denied_domains(&self) -> Option> { + self.domain_entries(NetworkDomainPermission::Deny) + } + + fn domain_entries(&self, permission: NetworkDomainPermission) -> Option> { + self.domains + .as_ref() + .map(|domains| { + domains + .effective_entries() + .iter() + .filter(|entry| entry.permission == permission) + .map(|entry| entry.pattern.clone()) + .collect() + }) + .filter(|entries: &Vec| !entries.is_empty()) + } + + pub fn allow_unix_sockets(&self) -> Vec { + self.unix_sockets + .as_ref() + .map(|unix_sockets| { + unix_sockets + .entries + .iter() + .filter(|(_, permission)| { + matches!(permission, NetworkUnixSocketPermission::Allow) + }) + .map(|(path, _)| path.clone()) + .collect() + }) + .unwrap_or_default() + } + + pub fn set_allowed_domains(&mut self, allowed_domains: Vec) { + self.set_domain_entries(allowed_domains, NetworkDomainPermission::Allow); + } + + pub fn set_denied_domains(&mut self, denied_domains: Vec) { + self.set_domain_entries(denied_domains, NetworkDomainPermission::Deny); + } + + pub fn upsert_domain_permission( + &mut self, + host: String, + permission: NetworkDomainPermission, + normalize: impl Fn(&str) -> String, + ) { + let mut domains = self.domains.take().unwrap_or_default(); + let normalized_host = normalize(&host); + domains + .entries + .retain(|entry| normalize(&entry.pattern) != normalized_host); + domains.entries.push(NetworkDomainPermissionEntry { + pattern: host, + permission, + }); + self.domains = (!domains.entries.is_empty()).then_some(domains); + } + + pub fn set_allow_unix_sockets(&mut self, allow_unix_sockets: Vec) { + self.set_unix_socket_entries(allow_unix_sockets, NetworkUnixSocketPermission::Allow); + } + + fn set_domain_entries(&mut self, entries: Vec, permission: NetworkDomainPermission) { + let mut domains = self.domains.take().unwrap_or_default(); + domains + .entries + .retain(|entry| entry.permission != permission); + for entry in entries { + if !domains + .entries + .iter() + .any(|existing| existing.pattern == entry && existing.permission == permission) + { + domains.entries.push(NetworkDomainPermissionEntry { + pattern: entry, + permission, + }); + } + } + self.domains = (!domains.entries.is_empty()).then_some(domains); + } + + fn set_unix_socket_entries( + &mut self, + entries: Vec, + permission: NetworkUnixSocketPermission, + ) { + let mut unix_sockets = self.unix_sockets.take().unwrap_or_default(); + unix_sockets + .entries + .retain(|_, existing| *existing != permission); + for entry in entries { + unix_sockets.entries.insert(entry, permission); + } + self.unix_sockets = (!unix_sockets.entries.is_empty()).then_some(unix_sockets); + } +} + +pub(crate) fn trusted_credential_broker_host(base_url: &str) -> Option { + Url::parse(base_url) + .ok() + .filter(|url| { + url.scheme() == "https" && url.username().is_empty() && url.password().is_none() + }) + .and_then(|url| url.host_str().map(normalize_host)) +} + +#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Default)] +#[serde(rename_all = "lowercase")] +pub enum NetworkMode { + /// Limited (read-only) access: only GET/HEAD/OPTIONS are allowed for HTTP. HTTPS CONNECT is + /// blocked unless MITM is enabled so the proxy can enforce method policy on inner requests. + /// SOCKS5 UDP and non-HTTPS SOCKS5 TCP remain blocked in limited mode. + Limited, + /// Full network access: all HTTP methods are allowed. HTTPS CONNECTs are tunneled directly. + /// MITM hooks do not currently make full mode enter MITM. + #[default] + Full, +} + +impl NetworkMode { + pub fn allows_method(self, method: &str) -> bool { + match self { + Self::Full => true, + Self::Limited => matches!(method, "GET" | "HEAD" | "OPTIONS"), + } + } +} + +fn default_proxy_url() -> String { + "http://127.0.0.1:3128".to_string() +} + +fn default_socks_url() -> String { + "http://127.0.0.1:8081".to_string() +} + +/// Clamp non-loopback bind addresses to loopback unless explicitly allowed. +fn clamp_non_loopback( + addr: SocketAddr, + allow_non_loopback: bool, + name: &str, + override_setting_name: &str, +) -> SocketAddr { + if addr.ip().is_loopback() { + return addr; + } + + if allow_non_loopback { + warn!("DANGEROUS: {name} listening on non-loopback address {addr}"); + return addr; + } + + warn!( + "{name} requested non-loopback bind ({addr}); clamping to 127.0.0.1:{port} (set {override_setting_name} to override)", + port = addr.port() + ); + SocketAddr::from(([127, 0, 0, 1], addr.port())) +} + +pub(crate) fn clamp_bind_addrs( + http_addr: SocketAddr, + socks_addr: SocketAddr, + cfg: &NetworkProxyConfig, +) -> (SocketAddr, SocketAddr) { + let http_addr = clamp_non_loopback( + http_addr, + cfg.dangerously_allow_non_loopback_proxy, + "HTTP proxy", + "dangerously_allow_non_loopback_proxy", + ); + let socks_addr = clamp_non_loopback( + socks_addr, + cfg.dangerously_allow_non_loopback_proxy, + "SOCKS5 proxy", + "dangerously_allow_non_loopback_proxy", + ); + if cfg.allow_unix_sockets().is_empty() && !cfg.dangerously_allow_all_unix_sockets { + return (http_addr, socks_addr); + } + + // `x-unix-socket` is intentionally a local escape hatch. If the proxy is reachable from + // outside the machine, it can become a remote bridge into local daemons + // (e.g. docker.sock). To avoid footguns, enforce loopback binding whenever unix sockets + // are enabled. + if cfg.dangerously_allow_non_loopback_proxy && !http_addr.ip().is_loopback() { + warn!( + "unix socket proxying is enabled; ignoring dangerously_allow_non_loopback_proxy and clamping HTTP proxy to loopback" + ); + } + if cfg.dangerously_allow_non_loopback_proxy && !socks_addr.ip().is_loopback() { + warn!( + "unix socket proxying is enabled; ignoring dangerously_allow_non_loopback_proxy and clamping SOCKS5 proxy to loopback" + ); + } + ( + SocketAddr::from(([127, 0, 0, 1], http_addr.port())), + SocketAddr::from(([127, 0, 0, 1], socks_addr.port())), + ) +} + +pub struct RuntimeConfig { + pub http_addr: SocketAddr, + pub socks_addr: SocketAddr, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct UnixStyleAbsolutePath(String); + +impl UnixStyleAbsolutePath { + fn parse(value: &str) -> Option { + value.starts_with('/').then(|| Self(value.to_string())) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum ValidatedUnixSocketPath { + Native(AbsolutePathBuf), + UnixStyleAbsolute(UnixStyleAbsolutePath), +} + +impl ValidatedUnixSocketPath { + pub(crate) fn parse(socket_path: &str) -> Result { + let path = Path::new(socket_path); + if path.is_absolute() { + let path = AbsolutePathBuf::from_absolute_path(path) + .with_context(|| format!("failed to normalize unix socket path {socket_path:?}"))?; + return Ok(Self::Native(path)); + } + + if let Some(path) = UnixStyleAbsolutePath::parse(socket_path) { + return Ok(Self::UnixStyleAbsolute(path)); + } + + bail!("expected an absolute path, got {socket_path:?}"); + } +} + +pub(crate) fn validate_unix_socket_allowlist_paths(cfg: &NetworkProxyConfig) -> Result<()> { + for (index, socket_path) in cfg.allow_unix_sockets().iter().enumerate() { + ValidatedUnixSocketPath::parse(socket_path) + .with_context(|| format!("invalid network.allow_unix_sockets[{index}]"))?; + } + Ok(()) +} + +pub fn resolve_runtime(cfg: &NetworkProxyConfig) -> Result { + validate_unix_socket_allowlist_paths(cfg)?; + + let http_addr = resolve_addr(&cfg.proxy_url, /*default_port*/ 3128) + .with_context(|| format!("invalid network.proxy_url: {}", cfg.proxy_url))?; + let socks_addr = resolve_addr(&cfg.socks_url, /*default_port*/ 8081) + .with_context(|| format!("invalid network.socks_url: {}", cfg.socks_url))?; + let (http_addr, socks_addr) = clamp_bind_addrs(http_addr, socks_addr, cfg); + + Ok(RuntimeConfig { + http_addr, + socks_addr, + }) +} + +/// Returns the sorted loopback ports used by the configured managed proxy listeners. +pub fn managed_proxy_ports(cfg: &NetworkProxyConfig) -> Result> { + let runtime = resolve_runtime(cfg)?; + if runtime.http_addr.port() == 0 { + bail!("network.proxy_url must use a fixed non-zero port for managed proxy provisioning"); + } + let mut ports = vec![runtime.http_addr.port()]; + if cfg.enable_socks5 { + if runtime.socks_addr.port() == 0 { + bail!( + "network.socks_url must use a fixed non-zero port for managed proxy provisioning" + ); + } + ports.push(runtime.socks_addr.port()); + } + ports.sort_unstable(); + ports.dedup(); + Ok(ports) +} + +fn resolve_addr(url: &str, default_port: u16) -> Result { + let addr_parts = parse_host_port(url, default_port)?; + let host = if addr_parts.host.eq_ignore_ascii_case("localhost") { + "127.0.0.1".to_string() + } else { + addr_parts.host + }; + match host.parse::() { + Ok(ip) => Ok(SocketAddr::new(ip, addr_parts.port)), + Err(_) => Ok(SocketAddr::from(([127, 0, 0, 1], addr_parts.port))), + } +} + +pub fn host_and_port_from_network_addr(value: &str, default_port: u16) -> String { + let trimmed = value.trim(); + if trimmed.is_empty() { + return "".to_string(); + } + + let parts = match parse_host_port(trimmed, default_port) { + Ok(parts) => parts, + Err(_) => { + return format_host_and_port(trimmed, default_port); + } + }; + + format_host_and_port(&parts.host, parts.port) +} + +fn format_host_and_port(host: &str, port: u16) -> String { + if host.contains(':') { + format!("[{host}]:{port}") + } else { + format!("{host}:{port}") + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct SocketAddressParts { + host: String, + port: u16, +} + +fn parse_host_port(url: &str, default_port: u16) -> Result { + let trimmed = url.trim(); + if trimmed.is_empty() { + bail!("missing host in network proxy address: {url}"); + } + + // Avoid treating unbracketed IPv6 literals like "2001:db8::1" as scheme-prefixed URLs. + if matches!(trimmed.parse::(), Ok(IpAddr::V6(_))) && !trimmed.starts_with('[') { + return Ok(SocketAddressParts { + host: trimmed.to_string(), + port: default_port, + }); + } + + // Prefer the standard URL parser when the input is URL-like. Prefix a scheme when absent so + // we still accept loose host:port inputs. + let candidate = if trimmed.contains("://") { + trimmed.to_string() + } else { + format!("http://{trimmed}") + }; + if let Ok(parsed) = Url::parse(&candidate) + && let Some(host) = parsed.host_str() + { + let host = host.trim_matches(|c| c == '[' || c == ']'); + if host.is_empty() { + bail!("missing host in network proxy address: {url}"); + } + return Ok(SocketAddressParts { + host: host.to_string(), + port: parsed.port().unwrap_or(default_port), + }); + } + + parse_host_port_fallback(trimmed, default_port) +} + +fn parse_host_port_fallback(input: &str, default_port: u16) -> Result { + let without_scheme = input + .split_once("://") + .map(|(_, rest)| rest) + .unwrap_or(input); + let host_port = without_scheme.split('/').next().unwrap_or(without_scheme); + let host_port = host_port + .rsplit_once('@') + .map(|(_, rest)| rest) + .unwrap_or(host_port); + + if host_port.starts_with('[') + && let Some(end) = host_port.find(']') + { + let host = &host_port[1..end]; + let port = host_port[end + 1..] + .strip_prefix(':') + .and_then(|port| port.parse::().ok()) + .unwrap_or(default_port); + if host.is_empty() { + bail!("missing host in network proxy address: {input}"); + } + return Ok(SocketAddressParts { + host: host.to_string(), + port, + }); + } + + // Only treat `host:port` as such when there's a single `:`. This avoids + // accidentally interpreting unbracketed IPv6 addresses as `host:port`. + if host_port.bytes().filter(|b| *b == b':').count() == 1 + && let Some((host, port)) = host_port.rsplit_once(':') + { + if host.is_empty() { + bail!("missing host in network proxy address: {input}"); + } + return Ok(SocketAddressParts { + host: host.to_string(), + port: port.parse::().ok().unwrap_or(default_port), + }); + } + + if host_port.is_empty() { + bail!("missing host in network proxy address: {input}"); + } + Ok(SocketAddressParts { + host: host_port.to_string(), + port: default_port, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + use pretty_assertions::assert_eq; + + fn settings_with_unix_sockets(unix_sockets: &[&str]) -> NetworkProxyConfig { + let mut settings = NetworkProxyConfig::default(); + if !unix_sockets.is_empty() { + settings.set_allow_unix_sockets( + unix_sockets + .iter() + .map(|path| (*path).to_string()) + .collect(), + ); + } + settings + } + + #[test] + fn network_proxy_settings_default_matches_local_use_baseline() { + assert_eq!( + NetworkProxyConfig::default(), + NetworkProxyConfig { + enabled: false, + proxy_url: "http://127.0.0.1:3128".to_string(), + enable_socks5: true, + socks_url: "http://127.0.0.1:8081".to_string(), + enable_socks5_udp: true, + allow_upstream_proxy: true, + dangerously_allow_non_loopback_proxy: false, + dangerously_allow_all_unix_sockets: false, + mode: NetworkMode::Full, + domains: None, + unix_sockets: None, + allow_local_binding: false, + mitm: false, + credential_broker: false, + credential_broker_enabled_mitm: false, + credential_providers: BTreeMap::new(), + credential_broker_openai_host: None, + credential_broker_context: crate::CredentialBrokerContext::default(), + dangerously_allow_plaintext_credential_injection: false, + mitm_hooks: Vec::new(), + } + ); + } + + #[test] + fn disabling_credential_broker_restores_independent_mitm_setting() { + for (mitm, add_hook) in [(false, false), (true, false), (false, true)] { + let mut original = NetworkProxyConfig { + enabled: true, + mitm, + ..Default::default() + }; + let mut config = original.clone(); + for _ in 0..2 { + config.set_credential_broker_enabled(/*enabled*/ true); + } + if add_hook { + config.mitm_hooks.push(MitmHookConfig { + host: "api.example".to_string(), + ..Default::default() + }); + original.mitm_hooks.clone_from(&config.mitm_hooks); + original.mitm = true; + } + for _ in 0..2 { + config.set_credential_broker_enabled(/*enabled*/ false); + } + assert_eq!(config, original); + assert_eq!( + crate::RemoteNetworkProxyConfig::from_effective_config(&config).is_err(), + original.mitm + ); + } + } + + #[test] + #[cfg(windows)] + fn ambiguous_credential_environment_preserves_remote_proxy_support() { + let mut config = NetworkProxyConfig { + enabled: true, + ..Default::default() + }; + let expected = crate::RemoteNetworkProxyConfig::from_effective_config(&config).unwrap(); + config.set_credential_broker_enabled(/*enabled*/ true); + config.configure_credential_broker_environment(&HashMap::from([ + ("GH_HOST".to_string(), "first.example".to_string()), + ("gh_host".to_string(), "second.example".to_string()), + ])); + assert_eq!( + crate::RemoteNetworkProxyConfig::from_effective_config(&config).unwrap(), + expected + ); + } + + #[test] + #[cfg(unix)] + fn credential_broker_context_accepts_non_unicode_environment() { + use std::ffi::OsString; + use std::os::unix::ffi::OsStringExt; + use std::process::Command; + + const CHILD_ENV: &str = "CODEX_TEST_NON_UNICODE_BROKER_CONTEXT"; + if std::env::var_os(CHILD_ENV).is_none() { + let output = Command::new(std::env::current_exe().unwrap()) + .args([ + "--exact", + "config::tests::credential_broker_context_accepts_non_unicode_environment", + "--nocapture", + ]) + .env(CHILD_ENV, OsString::from_vec(vec![0xff])) + .env(OsString::from_vec(vec![0xfe]), "unrelated") + .env("GH_HOST", OsString::from_vec(vec![0xff])) + .output() + .unwrap(); + assert!(output.status.success(), "{output:?}"); + return; + } + + let mut config = NetworkProxyConfig::default(); + config.set_credential_broker_enabled(/*enabled*/ true); + config.configure_credential_broker_environment(&HashMap::new()); + assert!(config.credential_broker); + } + + #[test] + fn credential_broker_only_accepts_trusted_https_openai_endpoints() { + let mut config = NetworkProxyConfig::default(); + + for (base_url, expected_host) in [ + ( + Some("https://gateway.example.com/v1"), + Some("gateway.example.com"), + ), + ( + Some("https://gateway.example.com./v1"), + Some("gateway.example.com"), + ), + (Some("https://[2001:db8::1]/v1"), Some("2001:db8::1")), + (Some("http://gateway.example.com/v1"), None), + (Some("https://user@gateway.example.com/v1"), None), + (Some("not-a-url"), None), + (None, None), + ] { + config.set_credential_broker_openai_base_url(base_url); + assert_eq!( + config.credential_broker_openai_host.as_deref(), + expected_host + ); + } + } + + #[test] + fn managed_proxy_ports_reject_ephemeral_ports() { + let mut config = NetworkProxyConfig { + proxy_url: "http://127.0.0.1:0".to_string(), + ..Default::default() + }; + + assert_eq!( + managed_proxy_ports(&config).unwrap_err().to_string(), + "network.proxy_url must use a fixed non-zero port for managed proxy provisioning" + ); + + config.proxy_url = "http://127.0.0.1:3128".to_string(); + config.socks_url = "socks5h://127.0.0.1:48081".to_string(); + assert_eq!(managed_proxy_ports(&config).unwrap(), vec![3128, 48081]); + + config.socks_url = "socks5h://127.0.0.1:0".to_string(); + assert_eq!( + managed_proxy_ports(&config).unwrap_err().to_string(), + "network.socks_url must use a fixed non-zero port for managed proxy provisioning" + ); + + config.enable_socks5 = false; + assert_eq!(managed_proxy_ports(&config).unwrap(), vec![3128]); + } + + #[test] + fn network_proxy_config_uses_struct_defaults_for_missing_fields() { + let config: NetworkProxyConfig = serde_json::from_str(r#"{ "enabled": true }"#).unwrap(); + let expected = NetworkProxyConfig { + enabled: true, + ..NetworkProxyConfig::default() + }; + + assert_eq!(config, expected); + } + + #[test] + fn set_allowed_domains_preserves_existing_deny_for_same_pattern() { + let mut settings = NetworkProxyConfig::default(); + settings.set_denied_domains(vec!["example.com".to_string()]); + + settings.set_allowed_domains(vec!["example.com".to_string()]); + + assert_eq!(settings.allowed_domains(), None); + assert_eq!( + settings.denied_domains(), + Some(vec!["example.com".to_string()]) + ); + } + + #[test] + fn network_domain_permissions_serialize_to_effective_map_shape() { + let mut settings = NetworkProxyConfig::default(); + settings.set_denied_domains(vec!["example.com".to_string()]); + settings.set_allowed_domains(vec!["example.com".to_string()]); + let config = settings; + + let value = serde_json::to_value(&config).unwrap(); + + assert_eq!( + value, + serde_json::json!({ + "enabled": false, + "proxy_url": "http://127.0.0.1:3128", + "enable_socks5": true, + "socks_url": "http://127.0.0.1:8081", + "enable_socks5_udp": true, + "allow_upstream_proxy": true, + "dangerously_allow_non_loopback_proxy": false, + "dangerously_allow_all_unix_sockets": false, + "mode": "full", + "domains": { + "example.com": "deny", + }, + "unix_sockets": null, + "allow_local_binding": false, + "mitm": false, + "credential_broker": false, + "dangerously_allow_plaintext_credential_injection": false, + "mitm_hooks": [], + }) + ); + } + + #[test] + fn parse_host_port_defaults_for_empty_string() { + assert!(parse_host_port("", /*default_port*/ 1234).is_err()); + } + + #[test] + fn parse_host_port_defaults_for_whitespace() { + assert!(parse_host_port(" ", /*default_port*/ 5555).is_err()); + } + + #[test] + fn parse_host_port_parses_host_port_without_scheme() { + assert_eq!( + parse_host_port("127.0.0.1:8080", /*default_port*/ 3128).unwrap(), + SocketAddressParts { + host: "127.0.0.1".to_string(), + port: 8080, + } + ); + } + + #[test] + fn parse_host_port_parses_host_port_with_scheme_and_path() { + assert_eq!( + parse_host_port( + "http://example.com:8080/some/path", + /*default_port*/ 3128 + ) + .unwrap(), + SocketAddressParts { + host: "example.com".to_string(), + port: 8080, + } + ); + } + + #[test] + fn parse_host_port_strips_userinfo() { + assert_eq!( + parse_host_port( + "http://user:pass@host.example:5555", + /*default_port*/ 3128 + ) + .unwrap(), + SocketAddressParts { + host: "host.example".to_string(), + port: 5555, + } + ); + } + + #[test] + fn parse_host_port_parses_ipv6_with_brackets() { + assert_eq!( + parse_host_port("http://[::1]:9999", /*default_port*/ 3128).unwrap(), + SocketAddressParts { + host: "::1".to_string(), + port: 9999, + } + ); + } + + #[test] + fn parse_host_port_does_not_treat_unbracketed_ipv6_as_host_port() { + assert_eq!( + parse_host_port("2001:db8::1", /*default_port*/ 3128).unwrap(), + SocketAddressParts { + host: "2001:db8::1".to_string(), + port: 3128, + } + ); + } + + #[test] + fn parse_host_port_falls_back_to_default_port_when_port_is_invalid() { + assert_eq!( + parse_host_port("example.com:notaport", /*default_port*/ 3128).unwrap(), + SocketAddressParts { + host: "example.com".to_string(), + port: 3128, + } + ); + } + + #[test] + fn host_and_port_from_network_addr_defaults_for_empty_string() { + assert_eq!( + host_and_port_from_network_addr("", /*default_port*/ 1234), + "" + ); + } + + #[test] + fn host_and_port_from_network_addr_formats_ipv6() { + assert_eq!( + host_and_port_from_network_addr("http://[::1]:8080", /*default_port*/ 3128), + "[::1]:8080" + ); + } + + #[test] + fn resolve_addr_maps_localhost_to_loopback() { + assert_eq!( + resolve_addr("localhost", /*default_port*/ 3128).unwrap(), + "127.0.0.1:3128".parse::().unwrap() + ); + } + + #[test] + fn resolve_addr_parses_ip_literals() { + assert_eq!( + resolve_addr("1.2.3.4", /*default_port*/ 80).unwrap(), + "1.2.3.4:80".parse::().unwrap() + ); + } + + #[test] + fn resolve_addr_parses_ipv6_literals() { + assert_eq!( + resolve_addr("http://[::1]:8080", /*default_port*/ 3128).unwrap(), + "[::1]:8080".parse::().unwrap() + ); + } + + #[test] + fn resolve_addr_falls_back_to_loopback_for_hostnames() { + assert_eq!( + resolve_addr("http://example.com:5555", /*default_port*/ 3128).unwrap(), + "127.0.0.1:5555".parse::().unwrap() + ); + } + + #[test] + fn clamp_bind_addrs_allows_non_loopback_when_enabled() { + let cfg = NetworkProxyConfig { + dangerously_allow_non_loopback_proxy: true, + ..Default::default() + }; + let http_addr = "0.0.0.0:3128".parse::().unwrap(); + let socks_addr = "0.0.0.0:8081".parse::().unwrap(); + + let (http_addr, socks_addr) = clamp_bind_addrs(http_addr, socks_addr, &cfg); + + assert_eq!(http_addr, "0.0.0.0:3128".parse::().unwrap()); + assert_eq!(socks_addr, "0.0.0.0:8081".parse::().unwrap()); + } + + #[test] + fn clamp_bind_addrs_forces_loopback_when_unix_sockets_enabled() { + let cfg = { + let mut settings = settings_with_unix_sockets(&["/tmp/docker.sock"]); + settings.dangerously_allow_non_loopback_proxy = true; + settings + }; + let http_addr = "0.0.0.0:3128".parse::().unwrap(); + let socks_addr = "0.0.0.0:8081".parse::().unwrap(); + + let (http_addr, socks_addr) = clamp_bind_addrs(http_addr, socks_addr, &cfg); + + assert_eq!(http_addr, "127.0.0.1:3128".parse::().unwrap()); + assert_eq!(socks_addr, "127.0.0.1:8081".parse::().unwrap()); + } + + #[test] + fn clamp_bind_addrs_forces_loopback_when_all_unix_sockets_enabled() { + let cfg = NetworkProxyConfig { + dangerously_allow_non_loopback_proxy: true, + dangerously_allow_all_unix_sockets: true, + ..Default::default() + }; + let http_addr = "0.0.0.0:3128".parse::().unwrap(); + let socks_addr = "0.0.0.0:8081".parse::().unwrap(); + + let (http_addr, socks_addr) = clamp_bind_addrs(http_addr, socks_addr, &cfg); + + assert_eq!(http_addr, "127.0.0.1:3128".parse::().unwrap()); + assert_eq!(socks_addr, "127.0.0.1:8081".parse::().unwrap()); + } + + #[test] + fn resolve_runtime_rejects_relative_allow_unix_sockets_entries() { + let cfg = settings_with_unix_sockets(&["relative.sock"]); + + let err = match resolve_runtime(&cfg) { + Ok(runtime) => panic!( + "relative allow_unix_sockets should fail, but resolve_runtime succeeded: {:?}", + runtime.http_addr + ), + Err(err) => err, + }; + assert!( + err.to_string().contains("network.allow_unix_sockets[0]"), + "error should point at the invalid allow_unix_sockets entry: {err:#}" + ); + } + + #[test] + fn resolve_runtime_accepts_unix_style_absolute_allow_unix_sockets_entries() { + let cfg = settings_with_unix_sockets(&["/private/tmp/example.sock"]); + + assert!( + resolve_runtime(&cfg).is_ok(), + "unix-style absolute allow_unix_sockets entry should be accepted" + ); + } +} diff --git a/codex-rs/network-proxy/src/connect_policy.rs b/codex-rs/network-proxy/src/connect_policy.rs new file mode 100644 index 0000000000000000000000000000000000000000..53b267595e8ba9873e5566d8af5c59a664226427 --- /dev/null +++ b/codex-rs/network-proxy/src/connect_policy.rs @@ -0,0 +1,238 @@ +use crate::policy::is_non_public_ip; +use crate::runtime::HostBlockDecision; +use crate::state::NetworkProxyState; +use rama_core::Service; +use rama_core::error::BoxError; +use rama_core::error::ErrorExt as _; +use rama_core::error::OpaqueError; +use rama_core::extensions::ExtensionsMut; +use rama_net::address::Host; +use rama_net::address::HostWithPort; +use rama_net::address::ProxyAddress; +use rama_net::client::EstablishedClientConnection; +use rama_net::transport::TryRefIntoTransportContext; +use rama_tcp::TcpStream; +use rama_tcp::client::TcpStreamConnector; +use rama_tcp::client::service::TcpConnector; +use std::io; +use std::net::SocketAddr; +use std::sync::Arc; + +#[derive(Clone)] +pub(crate) struct TargetCheckedTcpConnector { + state: Arc, +} + +impl TargetCheckedTcpConnector { + pub(crate) fn new(state: Arc) -> Self { + Self { state } + } +} + +impl Service for TargetCheckedTcpConnector +where + Input: TryRefIntoTransportContext + Send + ExtensionsMut + 'static, + Input::Error: Into + Send + Sync + 'static, +{ + type Output = EstablishedClientConnection; + type Error = BoxError; + + async fn serve(&self, input: Input) -> Result { + if input.extensions().get::().is_some() { + return TcpConnector::new().serve(input).await; + } + + let target = input + .try_ref_into_transport_ctx() + .map_err(|err| OpaqueError::from_boxed(err.into()).context("read network target"))? + .host_with_port() + .ok_or_else(|| OpaqueError::from_display("network target is missing a port"))?; + + TcpConnector::new() + .with_connector(TargetCheckedStreamConnector { + state: self.state.clone(), + target, + }) + .serve(input) + .await + } +} + +#[derive(Clone)] +struct TargetCheckedStreamConnector { + state: Arc, + target: HostWithPort, +} + +impl TcpStreamConnector for TargetCheckedStreamConnector { + type Error = BoxError; + + async fn connect(&self, addr: SocketAddr) -> Result { + if is_non_public_ip(addr.ip()) && !self.allows_non_public_target(addr).await? { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "network target rejected by policy", + ) + .into()); + } + + tokio::net::TcpStream::connect(addr) + .await + .map(TcpStream::from) + .map_err(Into::into) + } +} + +impl TargetCheckedStreamConnector { + async fn allows_non_public_target(&self, addr: SocketAddr) -> Result { + if self.state.allow_local_binding().await.map_err(|err| { + let err: BoxError = err.into(); + OpaqueError::from_boxed(err) + .context("read network proxy config") + .into_boxed() + })? { + return Ok(true); + } + + if !target_matches_non_public_addr(&self.target.host, addr.ip()) { + return Ok(false); + } + + self.state + .host_blocked(&self.target.host.to_string(), self.target.port) + .await + .map(|decision| decision == HostBlockDecision::Allowed) + .map_err(|err| { + let err: BoxError = err.into(); + OpaqueError::from_boxed(err) + .context("evaluate network proxy target") + .into_boxed() + }) + } +} + +pub(crate) fn is_non_public_target(host: &Host) -> bool { + match host { + Host::Address(ip) => is_non_public_ip(*ip), + Host::Name(name) => name + .as_str() + .trim_end_matches('.') + .eq_ignore_ascii_case("localhost"), + } +} + +fn target_matches_non_public_addr(host: &Host, addr: std::net::IpAddr) -> bool { + match host { + Host::Address(ip) => *ip == addr, + Host::Name(name) => { + name.as_str() + .trim_end_matches('.') + .eq_ignore_ascii_case("localhost") + && addr.is_loopback() + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::NetworkProxyConfig; + use crate::state::network_proxy_state_for_policy; + use rama_net::address::HostWithPort; + use std::net::Ipv4Addr; + use tokio::net::TcpListener; + + #[tokio::test(flavor = "current_thread")] + async fn direct_connector_rejects_non_public_target_when_local_binding_disabled() { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)) + .await + .expect("bind local listener"); + let target = listener.local_addr().expect("local addr"); + let connector = TargetCheckedTcpConnector::new(Arc::new(network_proxy_state_for_policy( + NetworkProxyConfig::default(), + ))); + + let request: rama_tcp::client::Request = + rama_tcp::client::Request::new(HostWithPort::from(target)); + let err = Service::serve(&connector, request) + .await + .expect_err("local target should be rejected"); + + assert!( + format!("{err:?}").contains("network target rejected by policy"), + "unexpected error: {err:?}" + ); + } + + #[tokio::test(flavor = "current_thread")] + async fn direct_connector_allows_non_public_target_when_local_binding_enabled() { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)) + .await + .expect("bind local listener"); + let target = listener.local_addr().expect("local addr"); + let connector = TargetCheckedTcpConnector::new(Arc::new(network_proxy_state_for_policy( + NetworkProxyConfig { + allow_local_binding: true, + ..NetworkProxyConfig::default() + }, + ))); + + let request: rama_tcp::client::Request = + rama_tcp::client::Request::new(HostWithPort::from(target)); + let result = Service::serve(&connector, request).await; + + assert!(result.is_ok(), "local target should be allowed: {result:?}"); + } + + #[tokio::test(flavor = "current_thread")] + async fn direct_connector_allows_explicitly_allowlisted_non_public_target() { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)) + .await + .expect("bind local listener"); + let target = listener.local_addr().expect("local addr"); + let mut config = NetworkProxyConfig::default(); + config.set_allowed_domains(vec![target.ip().to_string()]); + let connector = + TargetCheckedTcpConnector::new(Arc::new(network_proxy_state_for_policy(config))); + + let request: rama_tcp::client::Request = + rama_tcp::client::Request::new(HostWithPort::from(target)); + let result = Service::serve(&connector, request).await; + + assert!( + result.is_ok(), + "explicitly allowlisted local target should be allowed: {result:?}" + ); + } + + #[tokio::test(flavor = "current_thread")] + async fn direct_connector_allows_explicitly_allowlisted_localhost_target() { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)) + .await + .expect("bind local listener"); + let target = listener.local_addr().expect("local addr"); + let mut config = NetworkProxyConfig::default(); + config.set_allowed_domains(vec!["localhost".to_string()]); + let connector = + TargetCheckedTcpConnector::new(Arc::new(network_proxy_state_for_policy(config))); + + let request: rama_tcp::client::Request = + rama_tcp::client::Request::new(HostWithPort::new(Host::LOCALHOST_NAME, target.port())); + let result = Service::serve(&connector, request).await; + + assert!( + result.is_ok(), + "explicitly allowlisted localhost target should be allowed: {result:?}" + ); + } + + #[test] + fn resolved_private_address_does_not_match_allowlisted_hostname() { + let host = Host::Name("example.com".parse().expect("valid domain")); + + assert!(!target_matches_non_public_addr( + &host, + Ipv4Addr::LOCALHOST.into() + )); + } +} diff --git a/codex-rs/network-proxy/src/connection_lifecycle/lifecycle_tests.rs b/codex-rs/network-proxy/src/connection_lifecycle/lifecycle_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..60be13c30ffa71c709dc85e165934e232f0abd9a --- /dev/null +++ b/codex-rs/network-proxy/src/connection_lifecycle/lifecycle_tests.rs @@ -0,0 +1,266 @@ +//! Exercises proxy teardown through live TCP connections, including half-closed tunnels. + +use crate::NetworkProxy; +use crate::NetworkProxyConfig; +use crate::state::network_proxy_state_for_policy; +use anyhow::Context; +use anyhow::Result; +use pretty_assertions::assert_eq; +use std::collections::HashMap; +use std::future::Future; +use std::future::poll_fn; +use std::io::ErrorKind; +use std::net::Ipv4Addr; +use std::net::SocketAddr; +use std::sync::Arc; +use std::task::Poll; +use std::time::Duration; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::time::timeout; + +#[derive(Clone, Copy, Debug)] +enum ConnectionKind { + HttpConnect(TunnelState), + Socks5(TunnelState), + HttpKeepAlive, +} + +#[derive(Clone, Copy, Debug)] +enum TunnelState { + Open, + ClientWriteClosed, + UpstreamWriteClosed, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum PeerSide { + Client, + Upstream, +} + +struct Connection { + client: TcpStream, + upstream: TcpStream, + eof_received_by: Option, +} + +#[derive(Clone, Copy, Debug)] +enum StopProxy { + Shutdown, + Drop, + CancelWait, +} + +#[tokio::test] +async fn shutdown_closes_proxy_connections() -> Result<()> { + assert_stop_closes_connections(StopProxy::Shutdown).await +} + +#[tokio::test] +async fn dropping_handle_closes_proxy_connections() -> Result<()> { + assert_stop_closes_connections(StopProxy::Drop).await +} + +#[tokio::test] +async fn canceling_wait_closes_proxy_connections() -> Result<()> { + assert_stop_closes_connections(StopProxy::CancelWait).await +} + +async fn assert_stop_closes_connections(stop: StopProxy) -> Result<()> { + let mut config = NetworkProxyConfig { + allow_local_binding: true, + allow_upstream_proxy: false, + enable_socks5_udp: false, + ..NetworkProxyConfig::default() + }; + config.set_allowed_domains(vec![Ipv4Addr::LOCALHOST.to_string()]); + let proxy = NetworkProxy::builder() + .state(Arc::new(network_proxy_state_for_policy(config))) + .managed_by_codex(!cfg!(target_os = "windows")) + .http_addr(SocketAddr::from((Ipv4Addr::LOCALHOST, 0))) + .socks_addr(SocketAddr::from((Ipv4Addr::LOCALHOST, 0))) + .build() + .await?; + let handle = proxy.run().await?; + // Environment listeners reserve ephemeral ports on every platform. Managed main + // listeners do so off Windows, whose shared ingress has separate route tests. + let prepared = proxy.prepare_for_remote_environment(HashMap::new(), "lifecycle-test")?; + let scopes = [ + ("environment", prepared.env), + #[cfg(not(target_os = "windows"))] + ( + "main", + proxy + .prepare_for_optional_environment(HashMap::new(), /*environment_id*/ None)? + .env, + ), + ]; + let mut connections = Vec::new(); + for (scope, env) in scopes { + for kind in [ + ConnectionKind::HttpConnect(TunnelState::Open), + ConnectionKind::HttpConnect(TunnelState::ClientWriteClosed), + ConnectionKind::HttpConnect(TunnelState::UpstreamWriteClosed), + ConnectionKind::Socks5(TunnelState::Open), + ConnectionKind::Socks5(TunnelState::ClientWriteClosed), + ConnectionKind::Socks5(TunnelState::UpstreamWriteClosed), + ConnectionKind::HttpKeepAlive, + ] { + let (key, prefix) = match kind { + ConnectionKind::HttpConnect(_) | ConnectionKind::HttpKeepAlive => { + ("HTTP_PROXY", "http://") + } + ConnectionKind::Socks5(_) => ("ALL_PROXY", "socks5h://"), + }; + let addr = env[key] + .strip_prefix(prefix) + .context("proxy URL scheme")? + .parse()?; + let connection = timeout(Duration::from_secs(5), open_connection(kind, addr)) + .await + .with_context(|| format!("{scope} {kind:?} connection did not become ready"))??; + connections.push((scope, kind, connection)); + } + } + + match stop { + StopProxy::Shutdown => timeout(Duration::from_secs(2), handle.shutdown()) + .await + .context("proxy shutdown did not finish with active connections")??, + StopProxy::Drop => drop(handle), + StopProxy::CancelWait => { + let mut wait = Box::pin(handle.wait()); + let state = poll_fn(|cx| Poll::Ready(wait.as_mut().poll(cx))).await; + assert!( + state.is_pending(), + "running proxy unexpectedly stopped: {state:?}" + ); + drop(wait); + } + } + + // Retain every endpoint until all checks finish: dropping a test peer could + // otherwise cause the EOF that a later assertion attributes to proxy teardown. + for (scope, kind, connection) in &mut connections { + for (side, stream) in [ + (PeerSide::Client, &mut connection.client), + (PeerSide::Upstream, &mut connection.upstream), + ] { + if connection.eof_received_by == Some(side) { + continue; + } + let mut byte = [0_u8; 1]; + let result = timeout(Duration::from_secs(2), stream.read(&mut byte)) + .await + .with_context(|| { + format!("{scope} {kind:?} {side:?} connection remained open after {stop:?}") + })?; + match result { + Ok(0) => {} + Err(error) + if matches!( + error.kind(), + ErrorKind::ConnectionReset | ErrorKind::ConnectionAborted + ) => {} + result => anyhow::bail!( + "{scope} {kind:?} {side:?} did not close after {stop:?}: {result:?}" + ), + } + } + } + Ok(()) +} + +async fn open_connection(kind: ConnectionKind, proxy_addr: SocketAddr) -> Result { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await?; + let target = listener.local_addr()?; + let mut client = TcpStream::connect(proxy_addr).await?; + let request = match kind { + ConnectionKind::HttpConnect(_) => { + format!("CONNECT {target} HTTP/1.1\r\nHost: {target}\r\n\r\n").into_bytes() + } + ConnectionKind::HttpKeepAlive => { + format!("GET http://{target}/ HTTP/1.1\r\nHost: {target}\r\n\r\n").into_bytes() + } + ConnectionKind::Socks5(_) => { + client.write_all(&[5, 1, 0]).await?; + let mut greeting = [0_u8; 2]; + client.read_exact(&mut greeting).await?; + assert_eq!(greeting, [5, 0]); + let [port_high, port_low] = target.port().to_be_bytes(); + vec![5, 1, 0, 1, 127, 0, 0, 1, port_high, port_low] + } + }; + client.write_all(&request).await?; + let (upstream, _) = listener.accept().await?; + let mut connection = Connection { + client, + upstream, + eof_received_by: None, + }; + let tunnel_state = match kind { + ConnectionKind::HttpConnect(state) => { + let response = read_http_headers(&mut connection.client).await?; + assert!(response.starts_with("HTTP/1.1 200 "), "{response:?}"); + state + } + ConnectionKind::Socks5(state) => { + let mut response = [0_u8; 10]; + connection.client.read_exact(&mut response).await?; + assert_eq!(&response[..4], &[5, 0, 0, 1]); + state + } + ConnectionKind::HttpKeepAlive => { + read_http_headers(&mut connection.upstream).await?; + connection + .upstream + .write_all(b"HTTP/1.1 204 No Content\r\n\r\n") + .await?; + let response = read_http_headers(&mut connection.client).await?; + assert!(response.starts_with("HTTP/1.1 204 "), "{response:?}"); + return Ok(connection); + } + }; + + let (sender, receiver, receiver_side) = match tunnel_state { + TunnelState::Open | TunnelState::ClientWriteClosed => ( + &mut connection.client, + &mut connection.upstream, + PeerSide::Upstream, + ), + TunnelState::UpstreamWriteClosed => ( + &mut connection.upstream, + &mut connection.client, + PeerSide::Client, + ), + }; + sender.write_all(b"request").await?; + let mut request = [0_u8; 7]; + receiver.read_exact(&mut request).await?; + assert_eq!(&request, b"request"); + match tunnel_state { + TunnelState::Open => {} + TunnelState::ClientWriteClosed | TunnelState::UpstreamWriteClosed => { + sender.shutdown().await?; + assert_eq!(receiver.read(&mut [0_u8; 1]).await?, 0); + connection.eof_received_by = Some(receiver_side); + } + } + // Traffic in the opposite direction remains valid after a write half closes. + receiver.write_all(b"response").await?; + let mut response = [0_u8; 8]; + sender.read_exact(&mut response).await?; + assert_eq!(&response, b"response"); + Ok(connection) +} + +async fn read_http_headers(stream: &mut TcpStream) -> Result { + let mut headers = Vec::new(); + while !headers.ends_with(b"\r\n\r\n") { + headers.push(stream.read_u8().await?); + } + Ok(String::from_utf8(headers)?) +} diff --git a/codex-rs/network-proxy/src/connection_lifecycle/listeners.rs b/codex-rs/network-proxy/src/connection_lifecycle/listeners.rs new file mode 100644 index 0000000000000000000000000000000000000000..5c414bcd58d84ac97cc6b2bdaaa78e69e0b4443b --- /dev/null +++ b/codex-rs/network-proxy/src/connection_lifecycle/listeners.rs @@ -0,0 +1,47 @@ +//! Keeps listener tasks and their connection scope alive under the same runtime owner. + +use super::scope::ConnectionLifecycle; +use anyhow::Result; +use rama_core::graceful::ShutdownGuard; +use std::future::Future; +use tokio::task::JoinSet; + +pub(crate) struct ProxyListeners { + connections: ConnectionLifecycle, + listeners: JoinSet>, +} + +impl ProxyListeners { + pub(crate) fn new() -> Self { + Self { + connections: ConnectionLifecycle::new(), + listeners: JoinSet::new(), + } + } + + pub(crate) fn spawn(&mut self, listener: F) + where + F: FnOnce(ShutdownGuard) -> Fut, + Fut: Future> + Send + 'static, + { + self.listeners.spawn(listener(self.connections.guard())); + } + + pub(crate) fn cancel(&mut self) { + self.connections.cancel(); + self.listeners.abort_all(); + } + + pub(crate) async fn wait(&mut self) -> Result<()> { + while let Some(result) = self.listeners.join_next().await { + result??; + } + Ok(()) + } + + pub(crate) async fn shutdown(mut self) { + self.cancel(); + self.listeners.shutdown().await; + self.connections.shutdown().await; + } +} diff --git a/codex-rs/network-proxy/src/connection_lifecycle/mod.rs b/codex-rs/network-proxy/src/connection_lifecycle/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..b0219bbcd4983f54665b4ba04b98c5922214e2e3 --- /dev/null +++ b/codex-rs/network-proxy/src/connection_lifecycle/mod.rs @@ -0,0 +1,12 @@ +mod listeners; +mod scope; +mod service; + +pub(crate) use listeners::ProxyListeners; +#[cfg(test)] +pub(crate) use scope::ConnectionLifecycle; +pub(crate) use service::CancelOnShutdown; + +#[cfg(test)] +#[path = "lifecycle_tests.rs"] +mod tests; diff --git a/codex-rs/network-proxy/src/connection_lifecycle/scope.rs b/codex-rs/network-proxy/src/connection_lifecycle/scope.rs new file mode 100644 index 0000000000000000000000000000000000000000..2af4b500c871498d5b1ddc25f1e9e6feb9273efa --- /dev/null +++ b/codex-rs/network-proxy/src/connection_lifecycle/scope.rs @@ -0,0 +1,35 @@ +//! Owns cancellation and completion of a proxy's accepted connections and descendant tasks. + +use rama_core::graceful::Shutdown; +use rama_core::graceful::ShutdownGuard; +use tokio::sync::oneshot; + +pub(crate) struct ConnectionLifecycle { + // Only the runtime owner holds the sender. Dropping it also signals shutdown when an + // enclosing wait/shutdown future is cancelled, without spawning another cleanup task. + cancel: Option>, + shutdown: Shutdown, +} + +impl ConnectionLifecycle { + pub(crate) fn new() -> Self { + let (cancel, signal) = oneshot::channel(); + Self { + cancel: Some(cancel), + shutdown: Shutdown::new(signal), + } + } + + pub(crate) fn guard(&self) -> ShutdownGuard { + self.shutdown.guard() + } + + pub(crate) fn cancel(&mut self) { + self.cancel.take(); + } + + pub(crate) async fn shutdown(mut self) { + self.cancel(); + self.shutdown.shutdown().await; + } +} diff --git a/codex-rs/network-proxy/src/connection_lifecycle/service.rs b/codex-rs/network-proxy/src/connection_lifecycle/service.rs new file mode 100644 index 0000000000000000000000000000000000000000..2444b9cb671ace0eb84bd8df7cf561e401f9b915 --- /dev/null +++ b/codex-rs/network-proxy/src/connection_lifecycle/service.rs @@ -0,0 +1,43 @@ +//! Cancels connection work when its executor's proxy scope ends, including HTTP upgrades. + +use rama_core::Service; +use rama_core::extensions::ExtensionsRef; +use rama_core::rt::Executor; + +#[derive(Clone)] +pub(crate) struct CancelOnShutdown { + inner: S, +} + +impl CancelOnShutdown { + pub(crate) fn new(inner: S) -> Self { + Self { inner } + } +} + +impl Service for CancelOnShutdown +where + S: Service, + Request: ExtensionsRef + Send + 'static, +{ + type Output = (); + type Error = S::Error; + + async fn serve(&self, request: Request) -> Result<(), Self::Error> { + let guard = request + .extensions() + .get::() + .and_then(Executor::guard) + .cloned(); + match guard { + Some(guard) => { + tokio::select! { + biased; + _ = guard.cancelled() => Ok(()), + result = self.inner.serve(request) => result, + } + } + None => self.inner.serve(request).await, + } + } +} diff --git a/codex-rs/network-proxy/src/credential_broker.rs b/codex-rs/network-proxy/src/credential_broker.rs new file mode 100644 index 0000000000000000000000000000000000000000..39b545f33c4be79c5138f48a4c1e8aa5eec67d49 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker.rs @@ -0,0 +1,1639 @@ +mod configured; +mod destination; +mod environment; +mod matching; +mod provider_config; +mod providers; +mod registry; +mod replacement; + +use crate::config::NetworkProxyConfig; +use crate::policy::normalize_host; +use environment::update_brokered_credentials_marker; +use rama_http::HeaderMap; +use registry::ActiveCredentialSource; +use registry::BrokeredCredentialProvider; +use registry::active_credential_sources; +use registry::is_builtin_shaped_credential; +use registry::prioritized_credentials; +use registry::registered_credential_source; +use registry::select_credentials; +use replacement::Replacements; +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::collections::HashSet; +use std::sync::Arc; +use std::sync::RwLock; +use url::Url; + +pub use environment::CredentialBrokerContext; +pub use environment::CredentialBrokerEnvironment; +pub use environment::brokered_credential_binding_env_keys; +pub use environment::brokered_credential_dummy_env_keys; +pub use environment::brokered_credential_env_keys; +pub use environment::brokered_credential_marker_env_keys; +pub use environment::brokered_credential_value_env_keys; +pub use environment::credential_broker_provider_context_env_keys; +pub use environment::credential_broker_provider_sources_allowed; +pub use environment::is_credential_broker_provider_env_key; +pub(crate) use environment::marked_credential_dummy_env_keys; +pub use provider_config::CredentialAuthMethod; +pub use provider_config::CredentialProviderConfig; + +pub const CREDENTIAL_BROKER_ACTIVE_ENV_KEY: &str = "CODEX_NETWORK_PROXY_CREDENTIAL_BROKER_ACTIVE"; +pub(crate) const BROKERED_CREDENTIALS_ENV_KEY: &str = "CODEX_NETWORK_PROXY_BROKERED_CREDENTIALS"; +const MIN_EMBEDDED_CREDENTIAL_LENGTH: usize = 16; +const BROKERED_CREDENTIAL_ALIAS_MARKER_PREFIX: &str = "@alias:"; + +#[derive(Clone)] +pub(crate) struct CredentialBroker { + state: Arc>, +} + +#[derive(Default)] +struct CredentialBrokerState { + config_revision: u64, + enabled: bool, + allow_local_binding: bool, + openai_api_host: Option, + context: CredentialBrokerContext, + configured_provider_configs: BTreeMap, + configured_providers: Vec>, + credentials: Vec, + credential_owners: Vec, + credential_aliases: Vec, +} + +struct CredentialOwner { + env_var: String, + real_value: String, +} + +struct CredentialRecord { + env_var: String, + provider: BrokeredCredentialProvider, + host_binding: providers::CredentialHostBinding, + additional_host_bindings: Vec, + fallback_host_bindings: Vec, + environment_id: Option, + real_value: String, + dummy_value: String, + // Canonical identities accompanying discovery, distinct from an alias's own credential. + source_values: Vec, + generated_aliases: Arc>>, +} + +impl CredentialRecord { + fn has_carried_identity( + &self, + env: &HashMap, + credentials: &[CredentialRecord], + sources: &[ActiveCredentialSource], + ) -> bool { + let mut source_present = false; + let source_matches = sources + .iter() + .filter(|source| source_accepts_credential(source, self)) + .flat_map(|source| &source.env_vars) + .filter_map(|key| env_value(env, key)) + .any(|value| { + source_present = true; + let value = match &self.provider { + BrokeredCredentialProvider::Builtin(_) => value.trim(), + BrokeredCredentialProvider::Configured(_) => value, + }; + let value = credentials + .iter() + .find(|source| source.dummy_value == value) + .map_or(value, |source| source.real_value.as_str()); + value == self.real_value + || value == self.dummy_value + || self.source_values.iter().any(|source| source == value) + }); + let source_changed = source_present && !source_matches; + env_contains_credential_value(env, &self.dummy_value) + || (source_changed || self.invalidates_host_binding(env)) + && env.iter().any(|(key, value)| { + !env_key_matches(key, CREDENTIAL_BROKER_ACTIVE_ENV_KEY) + && !env_key_matches(key, BROKERED_CREDENTIALS_ENV_KEY) + && !crate::is_managed_proxy_env_var(key, value) + && (value == &self.real_value + || self.contains_embedded_value(value, &self.real_value)) + }) + } + + fn invalidates_host_binding(&self, env: &HashMap) -> bool { + match &self.provider { + BrokeredCredentialProvider::Builtin(provider) => { + provider.sources().iter().any(|source| { + source + .env_vars + .iter() + .any(|key| env_key_matches(key, &self.env_var)) + && (source.invalidates_host_binding)(env) + }) + } + BrokeredCredentialProvider::Configured(provider) => { + provider + .config + .url_prefix_from_env + .as_deref() + .is_some_and(|key| env_value(env, key).is_some()) + && provider.dynamic_destination(env).is_none() + } + } + } + + fn observe_host_binding( + &mut self, + host_binding: providers::CredentialHostBinding, + env: &HashMap, + ) { + if self.invalidates_host_binding(env) { + self.additional_host_bindings.clear(); + self.fallback_host_bindings.clear(); + self.host_binding = host_binding; + return; + } + if self.host_binding == host_binding { + return; + } + // Track the latest source separately, but keep prior destinations usable by + // concurrent commands sharing this dummy in the same environment. + self.additional_host_bindings + .retain(|binding| binding != &host_binding); + self.additional_host_bindings + .push(std::mem::replace(&mut self.host_binding, host_binding)); + } + + fn host_bindings(&self) -> impl Iterator { + std::iter::once(&self.host_binding).chain(&self.additional_host_bindings) + } + + fn belongs_to_environment(&self, environment_id: Option<&str>) -> bool { + self.environment_id.as_deref() == environment_id + } +} + +fn source_accepts_credential( + source: &ActiveCredentialSource, + credential: &CredentialRecord, +) -> bool { + credential.provider.same_provider(&source.provider) + && source + .env_vars + .iter() + .any(|env_var| env_key_matches(env_var, &credential.env_var)) +} + +fn source_tracks_credential( + source: &ActiveCredentialSource, + credential: &CredentialRecord, + env: &HashMap, +) -> bool { + if !source_accepts_credential(source, credential) { + return false; + } + // Aliases may hold a different credential from the canonical source. + if source.host_binding == credential.host_binding { + return true; + } + let mut source_present = false; + let contains_credential = source + .env_vars + .iter() + .filter_map(|env_var| env_value(env, env_var)) + .any(|value| { + source_present = true; + let value = match &source.provider { + BrokeredCredentialProvider::Builtin(_) => value.trim(), + BrokeredCredentialProvider::Configured(_) => value, + }; + value == credential.dummy_value + || value == credential.real_value + || credential + .source_values + .iter() + .any(|source| source == value) + }); + !source_present || contains_credential +} + +struct CredentialAlias { + env_var: String, + dummy_value: String, +} + +enum CredentialEnvironment { + Child, + Snapshot, +} + +fn env_key_matches(candidate: &str, expected: &str) -> bool { + if cfg!(windows) { + candidate.eq_ignore_ascii_case(expected) + } else { + candidate == expected + } +} + +fn env_entry<'a>(env: &'a HashMap, key: &str) -> Option<(&'a str, &'a str)> { + env.iter() + .find(|(candidate, _)| env_key_matches(candidate, key)) + .map(|(key, value)| (key.as_str(), value.as_str())) +} + +pub(super) fn env_value<'a>(env: &'a HashMap, key: &str) -> Option<&'a str> { + env_entry(env, key).map(|(_, value)| value) +} + +fn set_env_value(env: &mut HashMap, key: &str, value: String) { + if cfg!(windows) { + env.retain(|candidate, _| !env_key_matches(candidate, key)); + } + env.insert(key.to_string(), value); +} + +fn remove_env_value(env: &mut HashMap, key: &str) { + if cfg!(windows) { + env.retain(|candidate, _| !env_key_matches(candidate, key)); + } else { + env.remove(key); + } +} + +fn env_contains_credential_value(env: &HashMap, credential: &str) -> bool { + env.iter().any(|(key, value)| { + !env_key_matches(key, CREDENTIAL_BROKER_ACTIVE_ENV_KEY) + && !env_key_matches(key, BROKERED_CREDENTIALS_ENV_KEY) + && value.contains(credential) + && !crate::is_managed_proxy_env_var(key, value) + }) +} + +impl CredentialBroker { + pub(crate) fn new(enabled: bool) -> Self { + Self { + state: Arc::new(RwLock::new(CredentialBrokerState { + enabled, + ..CredentialBrokerState::default() + })), + } + } + + pub(crate) fn configure(&self, config: &NetworkProxyConfig) { + let mut state = self.write_state(); + let previous_context_sources = (state.context != config.credential_broker_context) + .then(|| active_credential_sources(&state, &HashMap::new())); + if state.enabled != config.credential_broker + || state.openai_api_host != config.credential_broker_openai_host + || state.allow_local_binding != config.allow_local_binding + || state.configured_provider_configs != config.credential_providers + || previous_context_sources.is_some() + { + state.config_revision += 1; + } + state.context.clone_from(&config.credential_broker_context); + state.allow_local_binding = config.allow_local_binding; + if state.enabled != config.credential_broker { + state.enabled = config.credential_broker; + state.credentials.clear(); + state.credential_owners.clear(); + state.credential_aliases.clear(); + } + if state.openai_api_host != config.credential_broker_openai_host { + state + .openai_api_host + .clone_from(&config.credential_broker_openai_host); + state.credentials.retain(|credential| { + !matches!( + &credential.provider, + BrokeredCredentialProvider::Builtin(provider) + if provider.reset_on_configuration_change + ) + }); + } + if state.configured_provider_configs != config.credential_providers { + let mut configured_providers: Vec> = + Vec::new(); + let mut provider_configs = config.credential_providers.iter().collect::>(); + provider_configs.sort_by_key(|(id, provider_config)| { + std::cmp::Reverse( + state.configured_provider_configs.get(id.as_str()) == Some(*provider_config) + && state + .configured_providers + .iter() + .any(|provider| provider.id == id.as_str()), + ) + }); + for (id, provider_config) in provider_configs { + if configured_providers.iter().any(|existing| { + existing.config.env.iter().any(|existing_key| { + provider_config + .env + .iter() + .any(|key| env_key_matches(existing_key, key)) + }) + }) { + tracing::warn!( + provider = %id, + "ignoring credential provider with an overlapping environment source" + ); + continue; + } + + let existing = (state.configured_provider_configs.get(id) == Some(provider_config)) + .then(|| { + state + .configured_providers + .iter() + .find(|provider| provider.id == *id) + .cloned() + }) + .flatten(); + match existing { + Some(provider) => configured_providers.push(provider), + None => { + match configured::ConfiguredCredentialProvider::compile(id, provider_config) + { + Ok(provider) => configured_providers.push(Arc::new(provider)), + Err(error) => { + tracing::warn!(provider = %id, %error, "ignoring invalid credential provider"); + } + } + } + } + } + state + .credentials + .retain(|credential| match &credential.provider { + BrokeredCredentialProvider::Builtin(_) => true, + BrokeredCredentialProvider::Configured(provider) => configured_providers + .iter() + .any(|current| Arc::ptr_eq(current, provider)), + }); + state.configured_providers = configured_providers; + state + .configured_provider_configs + .clone_from(&config.credential_providers); + } + if let Some(previous_sources) = previous_context_sources { + // Update registrations that used the old fallback, leaving distinct captured bindings intact. + let context = state.context.with_fallbacks(&HashMap::new()).into_owned(); + let sources = active_credential_sources(&state, &context); + state.credentials.retain_mut(|credential| { + let Some(previous_source) = previous_sources.iter().find(|source| { + source_accepts_credential(source, credential) + && credential + .host_bindings() + .any(|binding| binding == &source.host_binding) + }) else { + return true; + }; + let source = sources + .iter() + .find(|source| source_accepts_credential(source, credential)); + if source.is_some_and(|source| source.host_binding == previous_source.host_binding) + { + return true; + } + if !credential + .fallback_host_bindings + .contains(&previous_source.host_binding) + { + credential + .fallback_host_bindings + .push(previous_source.host_binding.clone()); + } + if source.is_none() || credential.invalidates_host_binding(&context) { + // Clearing a fallback revokes its history, not independently captured hosts. + credential + .additional_host_bindings + .retain(|binding| !credential.fallback_host_bindings.contains(binding)); + if credential + .fallback_host_bindings + .contains(&credential.host_binding) + { + let replacement = source + .map(|source| source.host_binding.clone()) + .or_else(|| credential.additional_host_bindings.pop()); + let Some(replacement) = replacement else { + return false; + }; + credential.host_binding = replacement; + } + credential.fallback_host_bindings.clear(); + } + if credential.host_binding == previous_source.host_binding { + let Some(source) = source else { + return false; + }; + credential.observe_host_binding(source.host_binding.clone(), &context); + } else { + // A captured primary destination must stay primary when its older fallback changes. + if let Some(source) = source + && !credential + .host_bindings() + .any(|binding| binding == &source.host_binding) + { + credential + .additional_host_bindings + .push(source.host_binding.clone()); + } + } + if let Some(source) = source + && !credential + .fallback_host_bindings + .contains(&source.host_binding) + { + credential + .fallback_host_bindings + .push(source.host_binding.clone()); + } + true + }); + } + } + + pub(crate) fn config_revision(&self) -> u64 { + self.read_state().config_revision + } + + #[cfg(test)] + pub(crate) fn discover_parent_credentials( + &self, + parent_env: &HashMap, + child_env: &HashMap, + ) { + self.discover_parent_credentials_for_environment( + parent_env, child_env, /*environment_id*/ None, + ); + } + + pub(crate) fn discover_parent_credentials_for_environment( + &self, + parent_env: &HashMap, + child_env: &HashMap, + environment_id: Option<&str>, + ) { + let mut state = self.write_state(); + if !state.enabled { + return; + } + state.observe_credential_owners(parent_env); + + let active_sources = active_credential_sources(&state, child_env); + for source in &active_sources { + for env_var in &source.env_vars { + let Some(real_value) = + brokerable_credential_value(parent_env, &state, env_var, &source.provider) + .map(str::to_string) + else { + continue; + }; + // Reconcile carried identities in virtualize_env, preserving their registered + // destinations instead of binding the parent's value to a rotated child source. + if state.credentials.iter().any(|credential| { + env_key_matches(&credential.env_var, env_var) + && credential.provider.same_provider(&source.provider) + && credential.real_value == real_value + && credential.has_carried_identity( + child_env, + &state.credentials, + &active_sources, + ) + }) { + continue; + } + if env_value(child_env, env_var) != Some(real_value.as_str()) + && child_env.values().any(|value| { + value == &real_value + || is_builtin_shaped_credential(&real_value) + && value.contains(&real_value) + || source.provider.contains_embedded_value(value, &real_value) + }) + { + let _ = state.register( + env_var, + source.provider.clone(), + source.host_binding.clone(), + environment_id, + &real_value, + parent_env, + ); + } + } + } + } + + #[cfg(test)] + pub(crate) fn virtualize_child_env(&self, env: &mut HashMap) { + self.virtualize_child_env_for_environment(env, /*environment_id*/ None); + } + + pub(crate) fn virtualize_child_env_for_environment( + &self, + env: &mut HashMap, + environment_id: Option<&str>, + ) { + self.virtualize_env(env, environment_id, CredentialEnvironment::Child); + } + + pub(crate) fn virtualize_snapshot_env( + &self, + env: &mut HashMap, + environment_id: Option<&str>, + ) { + self.virtualize_env(env, environment_id, CredentialEnvironment::Snapshot); + } + + fn virtualize_env( + &self, + env: &mut HashMap, + environment_id: Option<&str>, + destination: CredentialEnvironment, + ) { + let mut state = self.write_state(); + if !state.enabled { + remove_env_value(env, CREDENTIAL_BROKER_ACTIVE_ENV_KEY); + remove_env_value(env, BROKERED_CREDENTIALS_ENV_KEY); + return; + } + set_env_value(env, CREDENTIAL_BROKER_ACTIVE_ENV_KEY, "1".to_string()); + state.observe_credential_owners(env); + + let active_sources = active_credential_sources(&state, env); + let carried_dummies = state + .credentials + .iter() + .filter(|credential| { + credential.has_carried_identity(env, &state.credentials, &active_sources) + }) + .map(|credential| credential.dummy_value.clone()) + .collect::>(); + for credential in state.credentials.iter_mut().filter(|credential| { + credential.belongs_to_environment(environment_id) + && carried_dummies.contains(&credential.dummy_value) + }) { + if let Some(source) = registered_credential_source(credential, &active_sources, env) + .filter(|source| { + credential.invalidates_host_binding(env) + || source_tracks_credential(source, credential, env) + }) + { + credential.observe_host_binding(source.host_binding, env); + } + } + let stale_credentials = state + .credentials + .iter() + .filter(|credential| { + credential.belongs_to_environment(environment_id) + && carried_dummies.contains(&credential.dummy_value) + && !registered_credential_source(credential, &active_sources, env) + .is_some_and(|source| source_tracks_credential(&source, credential, env)) + }) + .map(|credential| { + ( + credential.dummy_value.clone(), + credential.real_value.clone(), + ) + }) + .collect::>(); + state.replace_child_env_dummies(env, &stale_credentials); + state.credentials.retain(|credential| { + !credential.belongs_to_environment(environment_id) + || !stale_credentials + .iter() + .any(|(dummy_value, _)| credential.dummy_value == *dummy_value) + }); + + let mut owned_dummies = state + .credentials + .iter() + .filter(|credential| { + credential.belongs_to_environment(environment_id) + && carried_dummies.contains(&credential.dummy_value) + }) + .map(|credential| credential.dummy_value.clone()) + .collect::>(); + let unbound_inherited_credentials = state + .credentials + .iter() + .filter(|credential| { + !credential.belongs_to_environment(environment_id) + && !owned_dummies.contains(&credential.dummy_value) + && env_contains_credential_value(env, &credential.dummy_value) + && registered_credential_source(credential, &active_sources, env).is_none() + }) + .map(|credential| { + ( + credential.dummy_value.clone(), + credential.real_value.clone(), + ) + }) + .collect::>(); + state.replace_child_env_dummies(env, &unbound_inherited_credentials); + + let inherited_credentials = state + .credentials + .iter() + .filter(|credential| { + !credential.belongs_to_environment(environment_id) + && !owned_dummies.contains(&credential.dummy_value) + && carried_dummies.contains(&credential.dummy_value) + }) + .filter_map(|credential| { + registered_credential_source(credential, &active_sources, env).map(|source| { + ( + credential.env_var.clone(), + source.provider, + source.host_binding, + credential.real_value.clone(), + credential.dummy_value.clone(), + credential.source_values.clone(), + Arc::clone(&credential.generated_aliases), + ) + }) + }) + .collect::>(); + let mut rebound_dummies = Vec::new(); + for ( + env_var, + provider, + host_binding, + real_value, + inherited_dummy, + source_values, + generated_aliases, + ) in inherited_credentials + { + let current_dummy = if let Some(credential) = + state.credentials.iter_mut().find(|credential| { + credential.belongs_to_environment(environment_id) + && credential.provider.same_provider(&provider) + && env_key_matches(&credential.env_var, &env_var) + && credential.real_value == real_value + }) { + credential.observe_host_binding(host_binding, env); + credential.dummy_value.clone() + } else { + state.credentials.push(CredentialRecord { + env_var, + provider, + host_binding, + additional_host_bindings: Vec::new(), + fallback_host_bindings: Vec::new(), + environment_id: environment_id.map(str::to_string), + real_value, + dummy_value: inherited_dummy.clone(), + source_values, + generated_aliases, + }); + inherited_dummy.clone() + }; + if current_dummy != inherited_dummy { + rebound_dummies.push((inherited_dummy, current_dummy.clone())); + } + owned_dummies.insert(current_dummy); + } + state.replace_child_env_dummies(env, &rebound_dummies); + for credential in state + .credentials + .iter_mut() + .filter(|credential| credential.belongs_to_environment(environment_id)) + { + for source in &mut credential.source_values { + if let Some((_, rebound)) = + rebound_dummies.iter().find(|(dummy, _)| dummy == source) + { + source.clone_from(rebound); + } + } + } + + for source in &active_sources { + for env_var in &source.env_vars { + virtualize_env_var( + env, + &mut state, + env_var, + source.provider.clone(), + source.host_binding.clone(), + environment_id, + ); + } + } + let provider_context_keys = state.environment(env).provider_context_keys; + let mut known_registrations = Vec::new(); + let binding_env = state.context.with_fallbacks(env); + let discoverable_values = env + .iter() + .filter(|(key, value)| { + !key.eq_ignore_ascii_case("PATH") + && !key.to_ascii_uppercase().ends_with("_PATH") + && !crate::is_managed_proxy_env_var(key, value) + && !provider_context_keys + .iter() + .any(|context_key| env_key_matches(key, context_key)) + }) + .map(|(_, value)| { + let known = matching::known_credential_matches(&state, value, env); + for matched in &known { + if matched.value == matched.real_value + && !state.credentials.iter().any(|credential| { + owned_dummies.contains(&credential.dummy_value) + && credential.real_value == matched.real_value + && credential.provider.same_provider(&matched.provider) + && env_key_matches(&credential.env_var, matched.env_var) + }) + && let Some(source) = active_sources.iter().find(|source| { + source.provider.same_provider(&matched.provider) + && source + .env_vars + .iter() + .any(|key| env_key_matches(key, matched.env_var)) + }) + { + known_registrations.push(( + matched.env_var.to_string(), + matched.provider.clone(), + source.host_binding.clone(), + matched.real_value.to_string(), + )); + } + } + matching::mask_known_credentials(value, &known) + }) + .collect::>(); + // Keep normal destination reconciliation for exact identities before discovering new ones. + for (env_var, provider, host_binding, real_value) in known_registrations { + let _ = state.register( + &env_var, + provider, + host_binding, + environment_id, + &real_value, + env, + ); + } + let configured_discoveries = state + .configured_providers + .iter() + .flat_map(|provider| { + discoverable_values.iter().flat_map(move |value| { + provider + .find_discoverable_credentials(value) + .filter(|matched| !value[matched.range.clone()].contains('\0')) + .map(move |matched| (Arc::clone(provider), value, matched)) + }) + }) + .collect::>(); + let mut builtin_spans = Vec::new(); + for provider in providers::credential_providers() { + let brokered_provider = BrokeredCredentialProvider::Builtin(provider); + let Some((source, host_binding)) = provider.sources().iter().rev().find_map(|source| { + (source.host_binding)(&binding_env, state.openai_api_host.as_deref()) + .map(|binding| (source, binding)) + }) else { + continue; + }; + for value in &discoverable_values { + for prefix in provider.credential_prefixes { + for (start, _) in value.match_indices(prefix) { + let credential = + matching::builtin_credential_candidate(provider, value, start); + let configured_start = configured_discoveries + .iter() + .filter(|(_, candidate_value, matched)| { + *candidate_value == value + && matched.has_distinctive_prefix + && matched.range.start > start + && matched.range.start < start + credential.len() + }) + .map(|(_, _, matched)| matched.range.start - start) + .filter(|end| { + credential[..*end].trim_end_matches(['_', '-']).len() + >= provider.minimum_credential_len + }) + .min(); + let credential = configured_start.map_or(credential, |end| { + credential[..end].trim_end_matches(['_', '-']) + }); + if credential.len() < provider.minimum_credential_len + || matching::is_operational_path_match( + value, + start, + start + credential.len(), + ) + || provider.ignored_credential_prefixes.iter().any(|ignored| { + credential.starts_with(ignored) + && provider + .credential_watermark + .is_none_or(|watermark| !credential.contains(watermark)) + }) + || provider.request_header_value(credential).is_none() + { + continue; + } + // Retain ownership spans even when registration fails or is ambiguous. + let span = start..start + credential.len(); + builtin_spans.push((value, span.clone())); + if configured_discoveries + .iter() + .any(|(_, candidate_value, matched)| { + *candidate_value == value + && matched.range.start == start + && matched.range.end >= span.end + }) + || state.is_dummy_value(credential) + || state.credentials.iter().any(|existing| { + existing.belongs_to_environment(environment_id) + && existing.provider.same_provider(&brokered_provider) + && existing.host_binding == host_binding + && existing.real_value == credential + }) + || state.credential_owners.iter().any(|existing| { + existing.real_value == credential + && (source.binding_env_vars.is_empty() + && !source + .env_vars + .iter() + .any(|key| env_key_matches(key, &existing.env_var)) + || !provider.sources().iter().any(|candidate| { + candidate + .env_vars + .iter() + .any(|key| env_key_matches(key, &existing.env_var)) + })) + }) + || provider.sources().iter().any(|candidate| { + candidate + .env_vars + .iter() + .any(|key| env_value(env, key) == Some(credential)) + && (candidate.host_binding)( + &binding_env, + state.openai_api_host.as_deref(), + ) + .is_none() + }) + { + continue; + } + let _ = state.register( + source.env_vars[0], + brokered_provider.clone(), + host_binding.clone(), + environment_id, + credential, + env, + ); + } + } + } + } + for (provider, value, matched) in &configured_discoveries { + let Some(host_binding) = provider.host_binding(&binding_env) else { + continue; + }; + let Some(env_var) = provider.config.env.first() else { + continue; + }; + let brokered_provider = BrokeredCredentialProvider::Configured(provider.clone()); + let range = &matched.range; + let credential = &value[range.clone()]; + if matching::is_operational_path_match(value, range.start, range.end) + || builtin_spans.iter().any(|(candidate_value, span)| { + candidate_value == value && span.start < range.end && range.start < span.end + }) + || configured_discoveries + .iter() + .any(|(other, candidate_value, matched)| { + !Arc::ptr_eq(provider, other) + && candidate_value == value + && matched.range.start < range.end + && range.start < matched.range.end + }) + || state.is_dummy_value(credential) + || state.credentials.iter().any(|existing| { + existing.belongs_to_environment(environment_id) + && existing.provider.same_provider(&brokered_provider) + && existing.host_binding == host_binding + && existing.real_value == credential + }) + || credential_provider_definitions(&state).any(|other| { + !other.same_provider(&brokered_provider) + && other.recognizes_credential_in(credential) + }) + || brokered_provider.request_header_value(credential).is_none() + { + continue; + } + let _ = state.register( + env_var, + brokered_provider.clone(), + host_binding.clone(), + environment_id, + credential, + env, + ); + } + let credentials = prioritized_credentials(&state, env) + .into_iter() + .filter(|credential| { + credential.belongs_to_environment(environment_id) + && (owned_dummies.contains(&credential.dummy_value) + || active_sources.iter().any(|source| { + credential.provider.same_provider(&source.provider) + && credential.host_binding == source.host_binding + })) + }) + .collect::>(); + let mut credential_aliases = Vec::new(); + for (key, value) in env.iter_mut() { + if crate::is_managed_proxy_env_var(key, value) { + continue; + } + let mut replacements = Replacements::default(); + for credential in &credentials { + for original in [&credential.real_value, &credential.dummy_value] { + replacements.add( + credential.value_match_ranges(value, original), + &credential.dummy_value, + Some(credential), + ); + } + } + if replacements.render(value) { + credential_aliases.push(CredentialAlias { + env_var: key.clone(), + dummy_value: value.clone(), + }); + } + } + for alias in credential_aliases { + if !state.credential_aliases.iter().any(|existing| { + env_key_matches(&existing.env_var, &alias.env_var) + && existing.dummy_value == alias.dummy_value + }) { + state.credential_aliases.push(alias); + } + } + if state.allow_local_binding && matches!(destination, CredentialEnvironment::Child) { + // Clients bypass the proxy for local destinations, so their credentials must stay real. + state.restore_child_env(env, |credential| { + credential.belongs_to_environment(environment_id) + && credential + .host_bindings() + .any(providers::CredentialHostBinding::bypasses_proxy_with_local_binding) + }); + } + update_brokered_credentials_marker(&state, env); + } + + pub(crate) fn restore_child_env( + &self, + env: &mut HashMap, + _command: &mut [String], + ) { + let state = self.read_state(); + if !state.enabled || env_value(env, CREDENTIAL_BROKER_ACTIVE_ENV_KEY) != Some("1") { + return; + } + state.restore_child_env(env, |_| true); + } + + pub(crate) fn restore_and_disable_child_env( + &self, + env: &mut HashMap, + command: &mut [String], + ) { + self.restore_child_env(env, command); + remove_env_value(env, CREDENTIAL_BROKER_ACTIVE_ENV_KEY); + remove_env_value(env, BROKERED_CREDENTIALS_ENV_KEY); + } + + pub(crate) fn child_alias_matches( + &self, + key: &str, + value: &str, + snapshot_value: &str, + environment_id: Option<&str>, + ) -> bool { + let state = self.read_state(); + if !state.enabled { + return false; + } + let mut expected = HashMap::from([(key.to_string(), snapshot_value.to_string())]); + state.restore_child_env(&mut expected, |_| true); + let referenced_credentials = |text: &str| { + state + .credentials + .iter() + .filter(|credential| { + text == credential.dummy_value + || credential.contains_embedded_value(text, &credential.dummy_value) + }) + .collect::>() + }; + let expected_credentials = referenced_credentials(snapshot_value); + if expected_credentials.is_empty() { + return false; + } + state.credential_aliases.iter().any(|alias| { + if !env_key_matches(key, &alias.env_var) { + return false; + } + let alias_credentials = referenced_credentials(&alias.dummy_value); + let same_owner = |left: &&CredentialRecord, right: &&CredentialRecord| { + left.provider.same_provider(&right.provider) + && env_key_matches(&left.env_var, &right.env_var) + && left.real_value == right.real_value + }; + if !expected_credentials.iter().all(|expected| { + alias_credentials + .iter() + .any(|actual| same_owner(expected, actual)) + }) || !alias_credentials.iter().all(|actual| { + expected_credentials + .iter() + .any(|expected| same_owner(expected, actual)) + }) { + return false; + } + let mut candidate = HashMap::from([(key.to_string(), alias.dummy_value.clone())]); + let mut identity = candidate.clone(); + if state.allow_local_binding { + state.restore_child_env(&mut candidate, |credential| { + credential.belongs_to_environment(environment_id) + && credential.host_bindings().any( + providers::CredentialHostBinding::bypasses_proxy_with_local_binding, + ) + }); + } + if env_value(&candidate, key) != Some(value) { + return false; + } + state.restore_child_env(&mut identity, |_| true); + identity == expected + }) + } + + #[cfg(test)] + pub(crate) fn host_requires_mitm(&self, host: &str, port: u16) -> bool { + self.host_protocols_for_environment(host, port, /*environment_id*/ None) + .tls + } + + pub(crate) fn host_protocols_for_environment( + &self, + host: &str, + port: u16, + environment_id: Option<&str>, + ) -> crate::brokered_tunnel::BrokeredProtocols { + let normalized_host = normalize_host(host); + let state = self.read_state(); + let mut protocols = crate::brokered_tunnel::BrokeredProtocols::default(); + if state.enabled { + for credential in state + .credentials + .iter() + .filter(|credential| credential.belongs_to_environment(environment_id)) + { + for binding in credential.host_bindings() { + protocols.tls |= binding.requires_mitm(&normalized_host, port); + if let providers::CredentialHostBinding::ConfiguredHosts(destinations) = binding + { + protocols.http |= destinations.iter().any(|destination| { + destination.requires_http_interception(&normalized_host, port) + }); + } + } + } + } + protocols + } + + pub(crate) fn virtualize_text(&self, text: &mut String, env: &HashMap) -> bool { + let state = self.read_state(); + matching::virtualize_text(&state, text, env) + } + + pub(crate) fn restore_text(&self, text: &mut String) -> bool { + let state = self.read_state(); + if !state.enabled { + return false; + } + + let mut credentials = state.credentials.iter().collect::>(); + credentials + .sort_unstable_by_key(|credential| std::cmp::Reverse(credential.dummy_value.len())); + let mut replacements = Replacements::default(); + for credential in credentials { + replacements.add( + credential.value_match_ranges(text, &credential.dummy_value), + &credential.real_value, + /*dummy*/ None, + ); + } + replacements.render(text) + } + + pub(crate) fn environment(&self, env: &HashMap) -> CredentialBrokerEnvironment { + self.read_state().environment(env) + } + + pub(crate) fn environment_for_text( + &self, + text: &str, + env: &HashMap, + ) -> CredentialBrokerEnvironment { + let state = self.read_state(); + let mut scoped_env = env.clone(); + for credential in &state.credentials { + remove_env_value(&mut scoped_env, &credential.env_var); + } + for credential in &state.credentials { + if text == credential.dummy_value + || credential.contains_embedded_value(text, &credential.dummy_value) + { + set_env_value( + &mut scoped_env, + &credential.env_var, + credential.dummy_value.clone(), + ); + } + } + update_brokered_credentials_marker(&state, &mut scoped_env); + state.environment(&scoped_env) + } + + pub(crate) fn source_matches_text(&self, source: &str, source_value: &str, text: &str) -> bool { + let state = self.read_state(); + state.enabled + && state.credentials.iter().any(|credential| { + env_key_matches(&credential.env_var, source) + && credential.real_value == source_value + && (text == credential.real_value + || credential.contains_embedded_value(text, &credential.real_value)) + }) + } + + pub(crate) fn provider_sources_allowed( + &self, + value: &str, + virtualized: &str, + source_env: &HashMap, + is_allowed: impl Fn(&str) -> bool, + ) -> bool { + let state = self.read_state(); + let known = matching::known_credential_matches(&state, value, source_env); + if known.iter().any(|matched| { + !known.iter().any(|equivalent| { + equivalent.range == matched.range + && equivalent.provider.same_provider(&matched.provider) + && equivalent.real_value == matched.real_value + && is_allowed(equivalent.env_var) + }) + }) { + return false; + } + let uncovered = matching::mask_known_credentials(value, &known); + let value = uncovered.as_str(); + let builtin_recognized = providers::credential_providers().any(|provider| { + provider.credential_prefixes.iter().any(|prefix| { + value.match_indices(*prefix).any(|(start, _)| { + matching::recognized_credential_match(provider, value, virtualized, start) + .is_some() + }) + }) + }); + if builtin_recognized + && !credential_broker_provider_sources_allowed( + value, + virtualized, + source_env, + &is_allowed, + ) + { + return false; + } + + let mut configured_recognized = false; + for provider in &state.configured_providers { + for source in &provider.config.env { + if let Some(credential) = env_value(source_env, source) + && provider.matches_value(credential) + && !provider + .credential_value_match_ranges(value, credential) + .is_empty() + && provider + .credential_value_match_ranges(virtualized, credential) + .is_empty() + { + configured_recognized = true; + if !provider.config.env.iter().any(|equivalent| { + env_value(source_env, equivalent) == Some(credential) + && is_allowed(equivalent) + }) { + return false; + } + } + } + for credential in provider.find_credentials(value).filter(|credential| { + !credential.contains('\0') && !virtualized.contains(credential) + }) { + configured_recognized = true; + let sources = provider + .config + .env + .iter() + .filter(|source| env_value(source_env, source) == Some(credential)) + .collect::>(); + let allowed = if sources.is_empty() { + provider.config.env.iter().all(|source| is_allowed(source)) + } else { + sources.iter().any(|source| is_allowed(source)) + }; + if !allowed { + return false; + } + } + } + + !known.is_empty() || builtin_recognized || configured_recognized + } + + #[cfg(test)] + pub(crate) fn inject_request_headers(&self, destination: &str, headers: &mut HeaderMap) { + self.inject_request_headers_for_environment( + destination, + headers, + /*environment_id*/ None, + ); + } + + pub(crate) fn inject_request_headers_for_environment( + &self, + destination: &str, + headers: &mut HeaderMap, + environment_id: Option<&str>, + ) { + let request = if destination.contains("://") { + let Ok(request) = Url::parse(destination) else { + return; + }; + Some(request) + } else { + None + }; + let normalized_host = normalize_host( + request + .as_ref() + .and_then(Url::host_str) + .unwrap_or(destination), + ); + let state = self.read_state(); + if !state.enabled { + return; + } + + let credentials = select_credentials( + headers, + &normalized_host, + request.as_ref(), + &state.credentials, + environment_id, + ); + if credentials.iter().any(|(credential, _, _)| { + matches!( + &credential.provider, + BrokeredCredentialProvider::Configured(_) + ) + }) && let Some(request) = request.as_ref() + { + let Ok(raw_request) = destination.parse::() else { + return; + }; + let raw_path = raw_request.path(); + if raw_path != request.path() + || !crate::authorization_path::is_safe_for_authorization(raw_path) + { + return; + } + } + for (credential, header_name, header_value) in credentials { + credential + .provider + .insert_request_header(headers, header_name, header_value); + } + } + + fn read_state(&self) -> std::sync::RwLockReadGuard<'_, CredentialBrokerState> { + self.state + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } + + fn write_state(&self) -> std::sync::RwLockWriteGuard<'_, CredentialBrokerState> { + self.state + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } +} + +fn virtualize_env_var( + env: &mut HashMap, + state: &mut CredentialBrokerState, + env_var: &str, + provider: BrokeredCredentialProvider, + host_binding: providers::CredentialHostBinding, + environment_id: Option<&str>, +) { + let previous_dummy = env_value(env, env_var) + .filter(|value| state.is_dummy_value(value)) + .map(str::to_string); + let Some(real_value) = brokerable_credential_value(env, state, env_var, &provider) + .map(str::to_string) + .or_else(|| { + let dummy = previous_dummy.as_deref()?; + state + .credentials + .iter() + .find(|credential| { + credential.dummy_value == dummy + && credential.provider.same_provider(&provider) + && env_key_matches(&credential.env_var, env_var) + }) + .map(|credential| credential.real_value.clone()) + }) + else { + return; + }; + + if let Some(dummy_value) = state.register( + env_var, + provider, + host_binding, + environment_id, + &real_value, + env, + ) { + if let Some(previous_dummy) = previous_dummy + && previous_dummy != dummy_value + { + state.replace_child_env_dummies(env, &[(previous_dummy, dummy_value.clone())]); + } + set_env_value(env, env_var, dummy_value); + } +} + +fn brokerable_credential_value<'a>( + env: &'a HashMap, + state: &CredentialBrokerState, + env_var: &str, + provider: &BrokeredCredentialProvider, +) -> Option<&'a str> { + let real_value = env_value(env, env_var)?; + let real_value = match provider { + BrokeredCredentialProvider::Builtin(_) => real_value.trim(), + BrokeredCredentialProvider::Configured(_) => real_value, + }; + (!real_value.is_empty() + && !state.is_dummy_value(real_value) + && provider.request_header_value(real_value).is_some()) + .then_some(real_value) +} + +impl CredentialBrokerState { + fn restore_child_env( + &self, + env: &mut HashMap, + should_restore: impl Fn(&CredentialRecord) -> bool, + ) { + let credentials = self + .credentials + .iter() + .filter(|credential| should_restore(credential)) + .filter(|credential| { + env.iter().any(|(key, value)| { + !env_key_matches(key, CREDENTIAL_BROKER_ACTIVE_ENV_KEY) + && !env_key_matches(key, BROKERED_CREDENTIALS_ENV_KEY) + && (value == &credential.dummy_value + || credential.contains_embedded_value(value, &credential.dummy_value)) + }) + }) + .collect::>(); + for (key, value) in env.iter_mut() { + if env_key_matches(key, CREDENTIAL_BROKER_ACTIVE_ENV_KEY) + || env_key_matches(key, BROKERED_CREDENTIALS_ENV_KEY) + { + continue; + } + let canonical_credential = self + .credentials + .iter() + .any(|credential| env_key_matches(key, &credential.env_var)); + if !canonical_credential + && !self.credential_aliases.iter().any(|alias| { + env_key_matches(key, &alias.env_var) && value == &alias.dummy_value + }) + { + continue; + } + let mut replacements = Replacements::default(); + for credential in &credentials { + if canonical_credential + && !self.credentials.iter().any(|candidate| { + env_key_matches(key, &candidate.env_var) + && candidate.provider.same_provider(&credential.provider) + && candidate.host_binding == credential.host_binding + && candidate.real_value == credential.real_value + }) + { + continue; + } + replacements.add( + credential.value_match_ranges(value, &credential.dummy_value), + &credential.real_value, + /*dummy*/ None, + ); + } + // Carry surviving dummy spans through partial local-destination restoration. + for credential in &self.credentials { + replacements.add( + credential.value_match_ranges(value, &credential.dummy_value), + &credential.dummy_value, + Some(credential), + ); + } + replacements.render(value); + } + } + + fn observe_credential_owners(&mut self, env: &HashMap) { + // Ownership is known even when a destination is missing or invalid. + let sources = providers::credential_providers() + .flat_map(|provider| { + provider.sources().iter().flat_map(move |source| { + source + .env_vars + .iter() + .map(move |key| (*key, BrokeredCredentialProvider::Builtin(provider))) + }) + }) + .chain(self.configured_providers.iter().flat_map(|provider| { + provider.config.env.iter().map(move |key| { + ( + key.as_str(), + BrokeredCredentialProvider::Configured(Arc::clone(provider)), + ) + }) + })); + let owners = sources + .filter_map(|(key, provider)| { + brokerable_credential_value(env, self, key, &provider) + .map(|value| (key.to_string(), value.to_string())) + }) + .collect::>(); + for (key, value) in owners { + self.remember_credential_owner(&key, &value); + } + } + + fn remember_credential_owner(&mut self, env_var: &str, real_value: &str) { + if !self + .credential_owners + .iter() + .any(|owner| env_key_matches(&owner.env_var, env_var) && owner.real_value == real_value) + { + self.credential_owners.push(CredentialOwner { + env_var: env_var.to_string(), + real_value: real_value.to_string(), + }); + } + } + + fn register( + &mut self, + env_var: &str, + provider: BrokeredCredentialProvider, + host_binding: providers::CredentialHostBinding, + environment_id: Option<&str>, + real_value: &str, + existing_env: &HashMap, + ) -> Option { + self.remember_credential_owner(env_var, real_value); + if let Some(existing) = self.credentials.iter_mut().find(|credential| { + env_key_matches(&credential.env_var, env_var) + && credential.provider.same_provider(&provider) + && credential.belongs_to_environment(environment_id) + && credential.real_value == real_value + }) { + // Existing dummies were reconciled before fresh-value discovery. + if !env_contains_credential_value(existing_env, &existing.dummy_value) { + existing.observe_host_binding(host_binding, existing_env); + } + return Some(existing.dummy_value.clone()); + } + if self.credentials.iter().any(|credential| { + is_builtin_shaped_credential(real_value) && credential.dummy_value.contains(real_value) + || provider.contains_embedded_value(&credential.dummy_value, real_value) + }) { + tracing::warn!( + env_var, + "credential brokerage skipped: credential overlaps an existing dummy" + ); + return None; + } + + let Some(dummy_value) = (0..64).find_map(|_| { + let candidate = provider.dummy_value(real_value)?; + (candidate != real_value + // These records restore embedded dummies without configured regex boundaries. + && (!is_builtin_shaped_credential(real_value) + || candidate.len() >= MIN_EMBEDDED_CREDENTIAL_LENGTH) + && !existing_env + .values() + .any(|value| value.contains(&candidate)) + && !self.credentials.iter().any(|credential| { + credential.dummy_value == candidate + || credential.real_value == candidate + || credential.contains_embedded_value(&candidate, &credential.dummy_value) + || credential.contains_embedded_value(&candidate, &credential.real_value) + }) + && !credential_provider_definitions(self).any(|other| { + !other.same_provider(&provider) + && other.recognizes_strictly_embedded_dummy_collision_in(&candidate) + })) + .then_some(candidate) + }) else { + tracing::warn!( + env_var, + "credential brokerage skipped: unable to generate a unique dummy credential" + ); + return None; + }; + let source_values = env_value(existing_env, env_var) + .map(|value| { + let value = match &provider { + BrokeredCredentialProvider::Builtin(_) => value.trim(), + BrokeredCredentialProvider::Configured(_) => value, + }; + if value == real_value { + vec![real_value.to_string(), dummy_value.clone()] + } else if let Some(source) = self + .credentials + .iter() + .find(|source| source.dummy_value == value) + { + vec![source.real_value.clone(), source.dummy_value.clone()] + } else { + vec![value.to_string()] + } + }) + .unwrap_or_default(); + self.credentials.push(CredentialRecord { + env_var: env_var.to_string(), + provider, + host_binding, + additional_host_bindings: Vec::new(), + fallback_host_bindings: Vec::new(), + environment_id: environment_id.map(str::to_string), + real_value: real_value.to_string(), + dummy_value: dummy_value.clone(), + source_values, + generated_aliases: Arc::default(), + }); + Some(dummy_value) + } + + fn is_dummy_value(&self, value: &str) -> bool { + self.credentials + .iter() + .any(|credential| credential.dummy_value == value) + } +} + +fn credential_provider_definitions( + state: &CredentialBrokerState, +) -> impl Iterator + '_ { + providers::credential_providers() + .map(BrokeredCredentialProvider::Builtin) + .chain( + state + .configured_providers + .iter() + .cloned() + .map(BrokeredCredentialProvider::Configured), + ) +} + +#[cfg(test)] +#[path = "credential_broker_tests.rs"] +mod tests; + +#[cfg(test)] +#[path = "credential_broker/configured_tests.rs"] +mod configured_tests; diff --git a/codex-rs/network-proxy/src/credential_broker/configured.rs b/codex-rs/network-proxy/src/credential_broker/configured.rs new file mode 100644 index 0000000000000000000000000000000000000000..b4f0630d1e435ff11800dadc42b5c8a9c8414019 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/configured.rs @@ -0,0 +1,578 @@ +use super::CredentialAuthMethod; +use super::CredentialProviderConfig; +use super::destination::CredentialDestination; +use super::env_value; +use super::providers; +use super::providers::CredentialHostBinding; +use anyhow::Context; +use anyhow::Result; +use anyhow::ensure; +use base64::Engine as _; +use rama_http::HeaderMap; +use rama_http::HeaderName; +use rama_http::HeaderValue; +use rama_http::header::AUTHORIZATION; +use rand::Rng as _; +use regex::Regex; +use regex::RegexBuilder; +use regex_automata::Anchored; +use regex_automata::Input; +use regex_automata::MatchKind; +use regex_automata::nfa::thompson; +use regex_automata::nfa::thompson::pikevm::PikeVM; +use regex_syntax::hir::Class; +use regex_syntax::hir::ClassBytes; +use regex_syntax::hir::ClassBytesRange; +use regex_syntax::hir::ClassUnicode; +use regex_syntax::hir::ClassUnicodeRange; +use regex_syntax::hir::Hir; +use regex_syntax::hir::HirKind; +use regex_syntax::hir::Look; +use std::collections::HashMap; + +const MAX_REGEX_BYTES: usize = 2048; +const MAX_REGEX_REPEAT: u32 = 64; +const MAX_DUMMY_ATTEMPTS: usize = 64; +const MIN_DISTINCTIVE_CREDENTIAL_PREFIX_LENGTH: usize = 4; +const MAX_DUMMY_VALUE_BYTES: usize = 2048; + +pub(super) struct ConfiguredCredentialProvider { + pub(super) id: String, + pub(super) config: CredentialProviderConfig, + patterns: Vec, + header: Option, +} + +struct ConfiguredCredentialPattern { + matcher: Regex, + longest_matcher: PikeVM, + known_value_matcher: Regex, + full_matcher: Regex, + ascii_generator: Option, + generator: rand_regex::Regex, + has_distinctive_prefix: bool, +} + +pub(super) struct DiscoveredCredential { + pub(super) range: std::ops::Range, + pub(super) has_distinctive_prefix: bool, +} + +impl ConfiguredCredentialPattern { + fn candidates(&self) -> impl Iterator + '_ { + self.ascii_generator + .iter() + .chain(std::iter::once(&self.generator)) + .flat_map(|generator| { + (0..MAX_DUMMY_ATTEMPTS).map(move |_| rand::rng().sample::(generator)) + }) + .filter(move |candidate| { + candidate.len() <= MAX_DUMMY_VALUE_BYTES + && !candidate.contains('\0') + && self.full_matcher.is_match(candidate) + }) + } +} + +#[derive(Clone, Copy)] +enum PatternPurpose { + Discovery, + AsciiDummy, + UnicodeDummy, +} + +impl ConfiguredCredentialProvider { + pub(super) fn compile(id: &str, config: &CredentialProviderConfig) -> Result { + ensure!(!id.is_empty(), "credential provider name must not be empty"); + ensure!( + !config.env.is_empty(), + "credential provider `{id}` has no environment keys" + ); + ensure!( + !config.patterns.is_empty(), + "credential provider `{id}` has no credential patterns" + ); + ensure!( + !config.url_prefixes.is_empty() || config.url_prefix_from_env.is_some(), + "credential provider `{id}` has no destination URL prefixes" + ); + ensure!( + config.env.iter().all(|key| valid_environment_key(key)), + "credential provider `{id}` has an invalid environment key" + ); + ensure!( + !config + .env + .iter() + .any(|key| super::is_credential_broker_provider_env_key(key)), + "credential provider `{id}` overlaps a built-in credential source" + ); + if let Some(key) = config.url_prefix_from_env.as_deref() { + ensure!( + valid_environment_key(key) + && !config.env.iter().any(|source| { + source.as_str() == key || cfg!(windows) && source.eq_ignore_ascii_case(key) + }), + "credential provider `{id}` has an invalid host environment key" + ); + } + for destination in &config.url_prefixes { + CredentialDestination::parse(destination) + .with_context(|| format!("invalid destination for credential provider `{id}`"))?; + } + let methods = auth_methods(config); + let header = config + .header + .as_deref() + .map(|header| HeaderName::from_bytes(header.as_bytes())) + .transpose() + .with_context(|| format!("invalid header for credential provider `{id}`"))?; + ensure!( + methods.contains(&CredentialAuthMethod::Header) == header.is_some(), + "credential provider `{id}` requires a header name exactly when using header authentication" + ); + ensure!( + config.prefix.is_none() || header.is_some(), + "credential provider `{id}` requires header authentication to use a header prefix" + ); + let patterns = config + .patterns + .iter() + .map(|pattern| { + ensure!( + pattern.len() <= MAX_REGEX_BYTES, + "credential pattern exceeds {MAX_REGEX_BYTES} bytes" + ); + let parsed = regex_syntax::parse(pattern)?; + let searchable = + prepare_credential_pattern(parsed.clone(), PatternPurpose::Discovery); + let has_distinctive_prefix = regex_syntax::hir::literal::Extractor::new() + .extract(&searchable) + .literals() + .is_some_and(|prefixes| { + !prefixes.is_empty() + && prefixes.iter().all(|prefix| { + prefix.len() >= MIN_DISTINCTIVE_CREDENTIAL_PREFIX_LENGTH + }) + }); + let searchable_pattern = searchable.to_string(); + let matcher = RegexBuilder::new(&searchable_pattern) + .size_limit(1 << 20) + .dfa_size_limit(1 << 20) + .build()?; + let longest_matcher = PikeVM::builder() + .configure(PikeVM::config().match_kind(MatchKind::All)) + .thompson(thompson::Config::new().nfa_size_limit(Some(1 << 20))) + .build(&searchable_pattern)?; + let known_value_matcher = + RegexBuilder::new(&format!(r"\A(?s:.)(?:{searchable_pattern})(?s:.)\z")) + .size_limit(1 << 20) + .dfa_size_limit(1 << 20) + .build()?; + let ascii_generator = rand_regex::Regex::with_hir( + prepare_credential_pattern(parsed.clone(), PatternPurpose::AsciiDummy), + MAX_REGEX_REPEAT, + ) + .ok(); + let generator = rand_regex::Regex::with_hir( + prepare_credential_pattern(parsed.clone(), PatternPurpose::UnicodeDummy), + MAX_REGEX_REPEAT, + )?; + let full_pattern = + Hir::concat(vec![Hir::look(Look::Start), parsed, Hir::look(Look::End)]) + .to_string(); + let full_matcher = RegexBuilder::new(&full_pattern) + .size_limit(1 << 20) + .dfa_size_limit(1 << 20) + .build()?; + ensure!( + searchable + .properties() + .minimum_len() + .is_some_and(|length| length > 0), + "credential pattern must not match an empty value" + ); + ensure!( + generator.is_utf8() && generator.capacity() <= MAX_REGEX_BYTES, + "credential pattern must generate bounded UTF-8 values" + ); + let compiled = ConfiguredCredentialPattern { + matcher, + longest_matcher, + known_value_matcher, + full_matcher, + ascii_generator, + generator, + has_distinctive_prefix, + }; + ensure!( + compiled.candidates().next().is_some(), + "credential pattern could not independently generate a matching dummy" + ); + Ok(compiled) + }) + .collect::>>() + .with_context(|| format!("invalid credential pattern for provider `{id}`"))?; + + Ok(Self { + id: id.to_string(), + config: config.clone(), + patterns, + header, + }) + } + + pub(super) fn matches_value(&self, value: &str) -> bool { + self.patterns + .iter() + .any(|pattern| pattern.full_matcher.is_match(value)) + } + + pub(super) fn find_credential_matches<'a>( + &'a self, + value: &'a str, + ) -> impl Iterator> + 'a { + self.credential_matches(value).map(|matched| matched.range) + } + + fn credential_matches<'a>( + &'a self, + value: &'a str, + ) -> impl Iterator + 'a { + self.patterns.iter().flat_map(move |pattern| { + let mut offset = 0; + let mut cache = pattern.longest_matcher.create_cache(); + std::iter::from_fn(move || { + let start = pattern.matcher.find_at(value, offset)?.start(); + // Anchor at the first match, then explore every alternative to its longest end. + // Keep the full input so word-boundary assertions retain their context. + let input = Input::new(value).range(start..).anchored(Anchored::Yes); + let matched = pattern.longest_matcher.find(&mut cache, input)?; + offset = matched.end(); + Some(matched.range()) + }) + .filter(move |matched| { + matched.len() >= super::MIN_EMBEDDED_CREDENTIAL_LENGTH + || pattern.has_distinctive_prefix + }) + .map(|range| DiscoveredCredential { + range, + has_distinctive_prefix: pattern.has_distinctive_prefix, + }) + }) + } + + pub(super) fn find_discoverable_credentials<'a>( + &'a self, + value: &'a str, + ) -> impl Iterator + 'a { + // Registration needs a complete token; redaction must still catch embedded matches. + self.credential_matches(value).filter(move |matched| { + value + .as_bytes() + .get(matched.range.end) + .is_none_or(|byte| !byte.is_ascii_alphanumeric() && !matches!(byte, b'_' | b'-')) + }) + } + + pub(super) fn credential_value_match_ranges( + &self, + text: &str, + credential: &str, + ) -> Vec> { + let Some(first) = credential.chars().next() else { + return Vec::new(); + }; + let mut offset = 0; + let mut ranges = Vec::new(); + while let Some(relative) = text[offset..].find(credential) { + let start = offset + relative; + let end = start + credential.len(); + // Force the full known span to match while retaining adjacent word boundaries. + let before = text[..start].chars().next_back().unwrap_or('\0'); + let after = text[end..].chars().next().unwrap_or('\0'); + let context = format!("{before}{credential}{after}"); + offset = if self + .patterns + .iter() + .any(|pattern| pattern.known_value_matcher.is_match(&context)) + { + ranges.push(start..end); + end + } else { + start + first.len_utf8() + }; + } + ranges + } + + pub(super) fn contains_strictly_embedded_pattern_match(&self, value: &str) -> bool { + self.patterns.iter().any(|pattern| { + pattern + .matcher + .find_iter(value) + .any(|matched| matched.start() > 0 || matched.end() < value.len()) + }) + } + + pub(super) fn find_credentials<'a>(&'a self, value: &'a str) -> impl Iterator { + self.find_credential_matches(value) + .map(|matched| &value[matched]) + } + + pub(super) fn host_binding( + &self, + env: &HashMap, + ) -> Option { + let mut destinations = self + .config + .url_prefixes + .iter() + .filter_map(|destination| CredentialDestination::parse(destination).ok()) + .collect::>(); + if let Some(destination) = self.dynamic_destination(env) + && !destinations.contains(&destination) + { + destinations.push(destination); + } + (!destinations.is_empty()).then_some(CredentialHostBinding::ConfiguredHosts(destinations)) + } + + pub(super) fn dynamic_destination( + &self, + env: &HashMap, + ) -> Option { + let key = self.config.url_prefix_from_env.as_deref()?; + CredentialDestination::parse(env_value(env, key)?) + .ok() + .filter(|destination| !destination.is_wildcard()) + } + + pub(super) fn dummy_value(&self, real_value: &str) -> Option { + let pattern = self + .patterns + .iter() + .find(|pattern| pattern.full_matcher.is_match(real_value))?; + pattern.candidates().find(|candidate| { + candidate != real_value && self.preserves_usable_auth_methods(real_value, candidate) + }) + } + + pub(super) fn request_header_value(&self, value: &str) -> Option { + auth_methods(&self.config) + .iter() + .find_map(|method| self.request_header_value_for_method(*method, value)) + } + + fn preserves_usable_auth_methods(&self, real_value: &str, dummy_value: &str) -> bool { + auth_methods(&self.config).iter().all(|method| { + // Preserve whether clients serialize a whole Basic pair or one component. + if *method == CredentialAuthMethod::Basic + && real_value.contains(':') != dummy_value.contains(':') + { + return false; + } + let Some(real_header) = self.request_header_value_for_method(*method, real_value) + else { + return true; + }; + let Some(dummy_header) = self.request_header_value_for_method(*method, dummy_value) + else { + return false; + }; + if dummy_header.as_bytes().trim_ascii() != dummy_header.as_bytes() { + return false; + } + let header = match method { + CredentialAuthMethod::Bearer + | CredentialAuthMethod::Token + | CredentialAuthMethod::Basic => AUTHORIZATION, + CredentialAuthMethod::Header => { + let Some(header) = self.header.as_ref() else { + return false; + }; + header.clone() + } + }; + self.translate_request_headers( + &HeaderMap::from_iter([(header.clone(), dummy_header)]), + dummy_value, + real_value, + ) + .contains(&(header, real_header)) + }) + } + + fn request_header_value_for_method( + &self, + method: CredentialAuthMethod, + value: &str, + ) -> Option { + let value = match method { + CredentialAuthMethod::Bearer => format!("Bearer {value}"), + CredentialAuthMethod::Token => format!("token {value}"), + CredentialAuthMethod::Basic => format!( + "Basic {}", + base64::engine::general_purpose::STANDARD.encode(value) + ), + CredentialAuthMethod::Header => { + format!( + "{}{value}", + self.config.prefix.as_deref().unwrap_or_default() + ) + } + }; + HeaderValue::from_str(&value) + .ok() + .filter(|header| header.to_str().is_ok()) + } + + pub(super) fn translate_request_headers( + &self, + headers: &HeaderMap, + expected_value: &str, + replacement_value: &str, + ) -> Vec<(HeaderName, HeaderValue)> { + let mut translated = Vec::new(); + for method in auth_methods(&self.config) { + match method { + CredentialAuthMethod::Bearer + | CredentialAuthMethod::Token + | CredentialAuthMethod::Basic => { + let Some(header) = headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + else { + continue; + }; + let Some((scheme, _)) = header.split_once(' ') else { + continue; + }; + let expected_scheme = match method { + CredentialAuthMethod::Bearer => "bearer", + CredentialAuthMethod::Token => "token", + CredentialAuthMethod::Basic => "basic", + CredentialAuthMethod::Header => continue, + }; + if !scheme.eq_ignore_ascii_case(expected_scheme) { + continue; + } + if let Some(value) = providers::translate_standard_request_header( + headers, + expected_value, + replacement_value, + ) && !translated.iter().any(|(header, _)| header == AUTHORIZATION) + { + translated.push((AUTHORIZATION, value)); + } + } + CredentialAuthMethod::Header => { + let Some(header) = self.header.as_ref() else { + continue; + }; + let Some(value) = headers.get(header).and_then(|value| value.to_str().ok()) + else { + continue; + }; + let prefix = self.config.prefix.as_deref().unwrap_or_default(); + if value == format!("{prefix}{expected_value}") + && !translated.iter().any(|(name, _)| name == header) + && let Ok(value) = + HeaderValue::from_str(&format!("{prefix}{replacement_value}")) + { + translated.push((header.clone(), value)); + } + } + } + } + translated + } +} + +fn prepare_credential_pattern(hir: Hir, purpose: PatternPurpose) -> Hir { + match hir.into_kind() { + HirKind::Empty => Hir::empty(), + HirKind::Literal(literal) => { + if matches!(purpose, PatternPurpose::AsciiDummy) && !literal.0.is_ascii() { + Hir::fail() + } else { + Hir::literal(literal.0) + } + } + HirKind::Class(mut class) => { + if matches!(purpose, PatternPurpose::AsciiDummy) { + match &mut class { + Class::Unicode(class) => { + class.intersect(&ClassUnicode::new([ClassUnicodeRange::new(' ', '~')])) + } + Class::Bytes(class) => { + class.intersect(&ClassBytes::new([ClassBytesRange::new(b' ', b'~')])) + } + } + } + Hir::class(class) + } + HirKind::Look(look) => { + if matches!(purpose, PatternPurpose::Discovery) + && !matches!( + look, + Look::Start + | Look::End + | Look::StartLF + | Look::EndLF + | Look::StartCRLF + | Look::EndCRLF + ) + { + Hir::look(look) + } else { + Hir::empty() + } + } + HirKind::Repetition(mut repetition) => { + if matches!(purpose, PatternPurpose::Discovery) { + // Discovery needs the whole credential, not a lazy prefix of it. + repetition.greedy = true; + } + repetition.sub = Box::new(prepare_credential_pattern(*repetition.sub, purpose)); + Hir::repetition(repetition) + } + HirKind::Capture(mut capture) => { + capture.sub = Box::new(prepare_credential_pattern(*capture.sub, purpose)); + Hir::capture(capture) + } + HirKind::Concat(expressions) => Hir::concat( + expressions + .into_iter() + .map(|expression| prepare_credential_pattern(expression, purpose)) + .collect(), + ), + HirKind::Alternation(expressions) => { + let mut expressions = expressions + .into_iter() + .map(|expression| prepare_credential_pattern(expression, purpose)) + .filter(|expression| expression.properties().minimum_len().is_some()) + .collect::>(); + expressions.sort_by_key(|expression| { + std::cmp::Reverse(expression.properties().maximum_len().unwrap_or(usize::MAX)) + }); + Hir::alternation(expressions) + } + } +} + +fn auth_methods(config: &CredentialProviderConfig) -> &[CredentialAuthMethod] { + if config.auth.is_empty() { + &[CredentialAuthMethod::Bearer] + } else { + &config.auth + } +} + +pub(super) fn valid_environment_key(key: &str) -> bool { + let mut chars = key.chars(); + chars + .next() + .is_some_and(|character| character == '_' || character.is_ascii_alphabetic()) + && chars.all(|character| character == '_' || character.is_ascii_alphanumeric()) +} diff --git a/codex-rs/network-proxy/src/credential_broker/configured_tests.rs b/codex-rs/network-proxy/src/credential_broker/configured_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4578e55d2711ecfaba4485c03455fec2e1387e82 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/configured_tests.rs @@ -0,0 +1,3092 @@ +use super::BrokeredCredentialProvider; +use super::CredentialAuthMethod; +use super::CredentialBroker; +use super::CredentialProviderConfig; +use super::brokered_credential_marker_env_keys; +use super::brokered_credential_value_env_keys; +use crate::NetworkProxyConfig; +use base64::Engine as _; +use pretty_assertions::assert_eq; +use rama_http::HeaderMap; +use rama_http::HeaderValue; +use rama_http::header::AUTHORIZATION; +use std::collections::BTreeMap; +use std::collections::HashMap; + +fn broker_for(provider: CredentialProviderConfig) -> CredentialBroker { + let broker = CredentialBroker::new(/*enabled*/ true); + broker.configure(&NetworkProxyConfig { + credential_broker: true, + credential_providers: BTreeMap::from([("custom".to_string(), provider)]), + ..NetworkProxyConfig::default() + }); + broker +} + +#[test] +fn adjacent_credentials_keep_original_boundaries() { + let broker = broker_for(CredentialProviderConfig { + env: vec!["FIRST_TOKEN".to_string(), "SECOND_TOKEN".to_string()], + patterns: vec![ + r"^(?:first_a{30}-|first_[b-z]{30}x)$".to_string(), + r"\bsecond_[a-z]{20}\b".to_string(), + ], + url_prefixes: vec!["https://provider.example/v1".to_string()], + ..CredentialProviderConfig::default() + }); + let first = format!("first_{}-", "a".repeat(30)); + let second = format!("second_{}", "a".repeat(20)); + let original = format!("{first}{second}"); + let mut env = HashMap::from([ + ("FIRST_TOKEN".to_string(), first.clone()), + ("SECOND_TOKEN".to_string(), second.clone()), + ("BUNDLE".to_string(), original.clone()), + ]); + + broker.virtualize_child_env(&mut env); + assert_ne!(env["FIRST_TOKEN"], first); + assert_ne!(env["SECOND_TOKEN"], second); + assert!(env["FIRST_TOKEN"].ends_with('x')); + + let expected = format!("{}{}", env["FIRST_TOKEN"], env["SECOND_TOKEN"]); + let mut text = original.clone(); + assert!(broker.virtualize_text(&mut text, &env)); + assert_eq!((&env["BUNDLE"], &text), (&expected, &expected)); + + let mut snapshot = format!("export BUNDLE='{original}'\n"); + assert!(broker.virtualize_text(&mut snapshot, &env)); + assert_eq!(snapshot, format!("export BUNDLE='{expected}'\n")); + assert!(broker.restore_text(&mut snapshot)); + assert_eq!(snapshot, format!("export BUNDLE='{original}'\n")); + assert!(broker.provider_sources_allowed(&expected, "", &env, |_| true)); + for denied in ["FIRST_TOKEN", "SECOND_TOKEN"] { + assert!(!broker.provider_sources_allowed(&expected, "", &env, |key| key != denied)); + } + + let mut boundary_control = format!("x{second}"); + assert!(broker.virtualize_text(&mut boundary_control, &env)); + assert_eq!(boundary_control, format!("x{second}")); + let mut dummy_control = format!("x{}", env["SECOND_TOKEN"]); + assert!(!broker.restore_text(&mut dummy_control)); + assert_eq!(dummy_control, format!("x{}", env["SECOND_TOKEN"])); + + env.remove("FIRST_TOKEN"); + env.remove("SECOND_TOKEN"); + broker.virtualize_child_env_for_environment(&mut env, Some("child")); + assert_eq!(env["BUNDLE"], expected); + assert_eq!( + brokered_credential_marker_env_keys(&env), + vec!["BUNDLE", "FIRST_TOKEN", "SECOND_TOKEN"] + ); + assert_eq!(brokered_credential_value_env_keys(&env), vec!["BUNDLE"]); + assert!(broker.child_alias_matches("BUNDLE", &expected, &expected, /*environment_id*/ None)); + let mut restored = env.clone(); + restored.insert("UNOBSERVED".to_string(), expected.clone()); + broker.restore_and_disable_child_env(&mut restored, &mut []); + assert_eq!( + restored, + HashMap::from([ + ("BUNDLE".to_string(), original), + ("UNOBSERVED".to_string(), expected), + ]) + ); +} + +#[test] +fn adjacent_aliases_survive_rebinding_to_existing_environment() { + for retain_sources in [false, true] { + for target in [None, Some("child")] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["FIRST_TOKEN".to_string(), "SECOND_TOKEN".to_string()], + patterns: vec![ + r"^(?:first_a{30}-|first_[b-z]{30}x)$".to_string(), + r"\bsecond_[a-z]{20}\b".to_string(), + ], + url_prefixes: vec!["https://provider.example/v1".to_string()], + ..CredentialProviderConfig::default() + }); + let first = format!("first_{}-", "a".repeat(30)); + let second = format!("second_{}", "a".repeat(20)); + let original = format!("{first}{second}"); + let mut parent = HashMap::from([ + ("FIRST_TOKEN".to_string(), first), + ("SECOND_TOKEN".to_string(), second), + ("BUNDLE".to_string(), original.clone()), + ]); + let mut expected = parent.clone(); + let mut child = parent.clone(); + child.remove("BUNDLE"); + broker.virtualize_child_env_for_environment(&mut parent, Some("parent")); + broker.virtualize_child_env_for_environment(&mut child, target); + assert_ne!(parent["FIRST_TOKEN"], child["FIRST_TOKEN"]); + if !retain_sources { + for key in ["FIRST_TOKEN", "SECOND_TOKEN"] { + parent.remove(key); + expected.remove(key); + } + } + broker.virtualize_child_env_for_environment(&mut parent, target); + assert_eq!( + parent["BUNDLE"], + format!("{}{}", child["FIRST_TOKEN"], child["SECOND_TOKEN"]) + ); + assert_eq!( + brokered_credential_marker_env_keys(&parent), + vec!["BUNDLE", "FIRST_TOKEN", "SECOND_TOKEN"] + ); + let mut text = parent["BUNDLE"].clone(); + assert!(broker.restore_text(&mut text)); + assert_eq!(text, original); + broker.restore_and_disable_child_env(&mut parent, &mut []); + assert_eq!(parent, expected); + } + } +} + +#[test] +fn generated_alias_spans_recognize_overlapping_dummy_candidates() { + let broker = broker_for(CredentialProviderConfig { + // Register the short dummy before the longer dummy containing its bytes. + env: vec!["SECOND_TOKEN".to_string(), "FIRST_TOKEN".to_string()], + patterns: vec![ + r"^(?:first_a{30}-|first_b{31})$".to_string(), + r"\b(?:a{16}|b{16})\b".to_string(), + ], + url_prefixes: vec!["https://provider.example/v1".to_string()], + ..CredentialProviderConfig::default() + }); + let first = format!("first_{}-", "a".repeat(30)); + let second = "a".repeat(16); + let original = format!("{first}{second}"); + let mut env = HashMap::from([ + ("FIRST_TOKEN".to_string(), first), + ("SECOND_TOKEN".to_string(), second), + ("BUNDLE".to_string(), original.clone()), + ]); + broker.virtualize_child_env(&mut env); + assert_eq!(env["BUNDLE"], format!("first_{}", "b".repeat(47))); + let state = broker.read_state(); + let second = state + .credentials + .iter() + .find(|credential| credential.env_var == "SECOND_TOKEN") + .unwrap(); + assert_eq!(second.generated_dummy_ranges(&env["BUNDLE"]), vec![37..53]); + drop(state); + assert!(!broker.provider_sources_allowed(&env["BUNDLE"], "", &env, |key| key == "FIRST_TOKEN")); + let mut repeated = format!("\u{03bb}{};{}", env["BUNDLE"], env["BUNDLE"]); + assert!(broker.restore_text(&mut repeated)); + assert_eq!(repeated, format!("\u{03bb}{original};{original}")); + env.remove("FIRST_TOKEN"); + env.remove("SECOND_TOKEN"); + broker.virtualize_child_env(&mut env); + assert_eq!( + brokered_credential_marker_env_keys(&env), + vec!["BUNDLE", "FIRST_TOKEN", "SECOND_TOKEN"] + ); + broker.restore_and_disable_child_env(&mut env, &mut []); + assert_eq!(env, HashMap::from([("BUNDLE".to_string(), original)])); +} + +#[test] +fn adjacent_dummy_restoration_keeps_original_boundaries() { + let broker = broker_for(CredentialProviderConfig { + env: vec!["FIRST_TOKEN".to_string(), "SECOND_TOKEN".to_string()], + patterns: vec![ + r"^(?:first_a{30}x|first_[b-z]{30}-)$".to_string(), + r"\bsecond_[a-z]{20}\b".to_string(), + ], + url_prefixes: vec!["https://provider.example/v1".to_string()], + ..CredentialProviderConfig::default() + }); + let first = format!("first_{}x", "a".repeat(30)); + let second = format!("second_{}", "a".repeat(20)); + let mut env = HashMap::from([ + ("FIRST_TOKEN".to_string(), first.clone()), + ("SECOND_TOKEN".to_string(), second.clone()), + ]); + broker.virtualize_child_env(&mut env); + assert!(env["FIRST_TOKEN"].ends_with('-')); + let mut text = format!("{}{}", env["FIRST_TOKEN"], env["SECOND_TOKEN"]); + env.insert("BUNDLE".to_string(), text.clone()); + broker.virtualize_child_env(&mut env); + let mut partial = env.clone(); + broker + .read_state() + .restore_child_env(&mut partial, |credential| { + credential.env_var == "FIRST_TOKEN" + }); + let mut partial_text = partial["BUNDLE"].clone(); + assert_eq!(partial_text, format!("{first}{}", env["SECOND_TOKEN"])); + assert!( + !broker.provider_sources_allowed(&partial_text, "", &partial, |key| key == "FIRST_TOKEN") + ); + assert!(broker.restore_text(&mut partial_text)); + assert_eq!(partial_text, format!("{first}{second}")); + assert!(broker.restore_text(&mut text)); + assert_eq!(text, format!("{first}{second}")); + broker.restore_and_disable_child_env(&mut env, &mut []); + assert_eq!( + env, + HashMap::from([ + ("FIRST_TOKEN".to_string(), first.clone()), + ("SECOND_TOKEN".to_string(), second.clone()), + ("BUNDLE".to_string(), format!("{first}{second}")) + ]) + ); +} + +#[test] +fn local_proxy_bypass_preserves_credentials_and_aliases_across_reload() { + let token = "vendor_abcdefghijklmnopqrstuvwx"; + let github = "ghp_abcdefghijklmnopqrstuvwxyz0123456789"; + for (destination, bypassed) in [ + ("http://localhost:1234/v1", true), + ("http://127.0.0.1:1234/v1", true), + ("http://[::1]:1234/v1", true), + ("https://10.1.2.3/v1", true), + ("https://172.16.2.3/v1", true), + ("https://192.168.2.3/v1", true), + ("https://api.localhost/v1", true), + ("https://127.0.0.2/v1", false), + ("https://172.32.2.3/v1", false), + ("https://api.vendor.example/v1", false), + ] { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut config = NetworkProxyConfig { + credential_broker: true, + credential_providers: BTreeMap::from([( + "vendor".to_string(), + CredentialProviderConfig { + env: vec!["VENDOR_TOKEN".to_string()], + patterns: vec!["^vendor_[a-z]{24}$".to_string()], + url_prefixes: vec!["https://public.vendor.example/v1".to_string()], + url_prefix_from_env: Some("VENDOR_URL".to_string()), + ..CredentialProviderConfig::default() + }, + )]), + ..NetworkProxyConfig::default() + }; + broker.configure(&config); + let mut env = HashMap::from([ + ("VENDOR_TOKEN".to_string(), token.to_string()), + ("VENDOR_URL".to_string(), destination.to_string()), + ("AUTH_HEADER".to_string(), format!("Bearer {token}")), + ("GH_TOKEN".to_string(), github.to_string()), + ]); + broker.virtualize_child_env(&mut env); + let dummy = env["VENDOR_TOKEN"].clone(); + let github_dummy = env["GH_TOKEN"].clone(); + assert_ne!(dummy, token); + assert_ne!(github_dummy, github); + + config.allow_local_binding = true; + let revision = broker.config_revision(); + broker.configure(&config); + assert_eq!(broker.config_revision(), revision + 1); + for _ in 0..2 { + broker.virtualize_child_env(&mut env); + let expected = if bypassed { token } else { &dummy }; + assert_eq!( + (&env["VENDOR_TOKEN"], &env["AUTH_HEADER"], &env["GH_TOKEN"]), + ( + &expected.to_string(), + &format!("Bearer {expected}"), + &github_dummy + ), + "{destination}" + ); + assert_eq!( + brokered_credential_value_env_keys(&env), + if bypassed { + vec!["GH_TOKEN"] + } else { + vec!["AUTH_HEADER", "GH_TOKEN", "VENDOR_TOKEN"] + } + ); + } + config.allow_local_binding = false; + let mut snapshot_env = env.clone(); + broker.virtualize_snapshot_env(&mut snapshot_env, /*environment_id*/ None); + assert_eq!(snapshot_env["VENDOR_TOKEN"], dummy); + assert_eq!(snapshot_env["AUTH_HEADER"], format!("Bearer {dummy}")); + broker.configure(&config); + broker.virtualize_child_env(&mut env); + assert_eq!(env["VENDOR_TOKEN"], dummy); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {dummy}")); + } +} + +#[test] +fn child_alias_identity_survives_scoped_dummies_and_partial_direct_restoration() { + let local = "local_abcdefghijklmnopqrstuvwx"; + let remote = "ghp_abcdefghijklmnopqrstuvwxyz0123456789"; + for allow_local_binding in [false, true] { + let broker = CredentialBroker::new(/*enabled*/ true); + broker.configure(&NetworkProxyConfig { + credential_broker: true, + allow_local_binding, + credential_providers: BTreeMap::from([( + "local".to_string(), + CredentialProviderConfig { + env: vec!["LOCAL_TOKEN".to_string()], + patterns: vec!["^local_[a-z]{24}$".to_string()], + url_prefixes: vec!["http://127.0.0.1:1234".to_string()], + ..CredentialProviderConfig::default() + }, + )]), + ..NetworkProxyConfig::default() + }); + let real_alias = format!("Local {local}; Remote {remote}"); + let mut env = HashMap::from([ + ("LOCAL_TOKEN".to_string(), local.to_string()), + ("GH_TOKEN".to_string(), remote.to_string()), + ("AUTH_HEADER".to_string(), real_alias.clone()), + ]); + let mut snapshot = env.clone(); + broker.virtualize_snapshot_env(&mut snapshot, Some("snapshot")); + let snapshot_alias = snapshot["AUTH_HEADER"].clone(); + broker.virtualize_child_env_for_environment(&mut env, Some("child")); + env.remove("LOCAL_TOKEN"); + env.remove("GH_TOKEN"); + assert_eq!(env["AUTH_HEADER"].contains(local), allow_local_binding); + assert!(!env["AUTH_HEADER"].contains(remote)); + assert!(broker.child_alias_matches( + "AUTH_HEADER", + &env["AUTH_HEADER"], + &snapshot_alias, + Some("child") + )); + assert!(!broker.child_alias_matches( + "AUTH_HEADER", + &real_alias, + &snapshot_alias, + Some("child") + )); + assert!(!broker.child_alias_matches( + "AUTH_HEADER", + &format!("{} altered", env["AUTH_HEADER"]), + &snapshot_alias, + Some("child") + )); + assert!(!broker.child_alias_matches( + "OTHER_HEADER", + &env["AUTH_HEADER"], + &snapshot_alias, + Some("child") + )); + } +} + +#[test] +fn local_proxy_bypass_is_scoped_for_inherited_credentials_and_aliases() { + let token = "vendor_abcdefghijklmnopqrstuvwx"; + for parent_is_local in [false, true] { + for retain_source in [false, true] { + let broker = CredentialBroker::new(/*enabled*/ true); + broker.configure(&NetworkProxyConfig { + credential_broker: true, + allow_local_binding: true, + credential_providers: BTreeMap::from([( + "vendor".to_string(), + CredentialProviderConfig { + env: vec!["VENDOR_TOKEN".to_string()], + patterns: vec!["^vendor_[a-z]{24}$".to_string()], + url_prefix_from_env: Some("VENDOR_URL".to_string()), + ..CredentialProviderConfig::default() + }, + )]), + ..NetworkProxyConfig::default() + }); + let local = "http://127.0.0.1:1234/v1"; + let public = "https://api.vendor.example/v1"; + let mut snapshot = HashMap::from([ + ("VENDOR_TOKEN".to_string(), token.to_string()), + ( + "VENDOR_URL".to_string(), + if parent_is_local { local } else { public }.to_string(), + ), + ("AUTH_HEADER".to_string(), format!("Bearer {token}")), + ]); + broker.virtualize_snapshot_env(&mut snapshot, Some("parent")); + let dummy = snapshot["VENDOR_TOKEN"].clone(); + let snapshot_alias = snapshot["AUTH_HEADER"].clone(); + let mut child = snapshot.clone(); + child.insert( + "VENDOR_URL".to_string(), + if parent_is_local { public } else { local }.to_string(), + ); + if !retain_source { + child.remove("VENDOR_TOKEN"); + } + for (id, env, bypasses) in [ + ("child", &mut child, !parent_is_local), + ("parent", &mut snapshot, parent_is_local), + ] { + broker.virtualize_child_env_for_environment(env, Some(id)); + let expected = if bypasses { token } else { &dummy }; + assert_eq!( + env.get("VENDOR_TOKEN").map(String::as_str), + (retain_source || id == "parent").then_some(expected) + ); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {expected}")); + assert!(broker.child_alias_matches( + "AUTH_HEADER", + &env["AUTH_HEADER"], + &snapshot_alias, + Some(id) + )); + if !bypasses { + assert!(!broker.child_alias_matches( + "AUTH_HEADER", + &format!("Bearer {token}"), + &snapshot_alias, + Some(id) + )); + } + } + } + } +} + +#[test] +fn configured_provider_virtualizes_credentials_aliases_and_snapshots() { + let token = "stripe_live_abcdefghijklmnopqrstuvwx"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["STRIPE_API_KEY".to_string()], + patterns: vec!["^stripe_live_[a-z]{24}$".to_string()], + url_prefixes: vec!["api.stripe.com".to_string(), "*.stripe.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([ + ("STRIPE_API_KEY".to_string(), token.to_string()), + ("AUTH_HEADER".to_string(), format!("Bearer {token}")), + ]); + + broker.virtualize_child_env(&mut env); + + let dummy = &env["STRIPE_API_KEY"]; + assert_ne!(dummy, token); + assert!( + regex::Regex::new("^stripe_live_[a-z]{24}$") + .expect("valid credential pattern") + .is_match(dummy) + ); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {dummy}")); + assert_eq!( + brokered_credential_value_env_keys(&env), + vec!["AUTH_HEADER", "STRIPE_API_KEY"] + ); + assert_eq!( + broker.environment(&env).credential_keys, + vec!["STRIPE_API_KEY".to_string()] + ); + let mut snapshot = format!("token={token}"); + assert!(broker.virtualize_text(&mut snapshot, &env)); + assert_eq!(snapshot, format!("token={dummy}")); + let mut unknown = "stripe_live_yyyyyyyyyyyyyyyyyyyyyyyy".to_string(); + assert!(!broker.virtualize_text(&mut unknown, &env)); + assert!(unknown.is_empty()); + assert!(broker.host_requires_mitm("api.stripe.com", /*port*/ 443)); + assert!(broker.host_requires_mitm("billing.stripe.example", /*port*/ 443)); + assert!(!broker.host_requires_mitm("stripe.example", /*port*/ 443)); + assert!(!broker.host_requires_mitm("attacker.example", /*port*/ 443)); + let dummy_header = format!("Bearer {dummy}"); + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&dummy_header).expect("valid dummy authentication"), + ); + broker.inject_request_headers("https://attacker.example/", &mut headers); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(dummy_header.as_str()) + ); + + let mut unmanaged_env = env.clone(); + crate::strip_managed_proxy_env(&mut unmanaged_env); + assert!(!unmanaged_env.contains_key("STRIPE_API_KEY")); + env.remove("STRIPE_API_KEY"); + broker.virtualize_child_env(&mut env); + assert_eq!( + brokered_credential_marker_env_keys(&env), + vec!["AUTH_HEADER", "STRIPE_API_KEY"] + ); + assert_eq!( + brokered_credential_value_env_keys(&env), + vec!["AUTH_HEADER"] + ); +} + +#[test] +fn configured_provider_discovers_credentials_without_canonical_variables() { + let token = "stripe_live_abcdefghijklmnopqrstuvwx"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["STRIPE_API_KEY".to_string()], + patterns: vec!["^stripe_live_[a-z]{24}$".to_string()], + url_prefixes: vec!["api.stripe.com".to_string()], + ..CredentialProviderConfig::default() + }); + let authorization_header = format!("Bearer {token}"); + let mut env = HashMap::from([("AUTH_HEADER".to_string(), authorization_header.clone())]); + + broker.virtualize_child_env(&mut env); + + assert!(!env.contains_key("STRIPE_API_KEY")); + let dummy_header = &env["AUTH_HEADER"]; + assert_ne!(dummy_header, &authorization_header); + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(dummy_header).expect("valid dummy authentication"), + ); + broker.inject_request_headers("https://api.stripe.com/", &mut headers); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(authorization_header.as_str()) + ); + let mut snapshot = format!("export AUTH_HEADER='{authorization_header}'"); + assert!(broker.virtualize_text(&mut snapshot, &env)); + assert_eq!(snapshot, format!("export AUTH_HEADER='{dummy_header}'")); + + broker.restore_child_env(&mut env, &mut []); + assert_eq!(env["AUTH_HEADER"], authorization_header); + assert!(!env.contains_key("STRIPE_API_KEY")); +} + +#[test] +fn configured_provider_uses_private_destination_context() { + let token = "pin_abcdefgh"; + let destination = "https://api.vendor.example/v2"; + let mut config = NetworkProxyConfig { + credential_broker: true, + credential_providers: BTreeMap::from([( + "vendor".to_string(), + CredentialProviderConfig { + env: vec!["VENDOR_PASSWORD".to_string()], + patterns: vec!["pin_[a-z]{8}".to_string()], + url_prefix_from_env: Some("VENDOR_HOST".to_string()), + ..CredentialProviderConfig::default() + }, + )]), + ..NetworkProxyConfig::default() + }; + config.configure_credential_broker_environment(&HashMap::from([ + ("VENDOR_HOST".to_string(), destination.to_string()), + ("UNRELATED_SECRET".to_string(), token.to_string()), + ])); + assert!(!format!("{config:?}").contains(destination)); + assert!( + !serde_json::to_string(&config) + .unwrap() + .contains(destination) + ); + let broker = CredentialBroker::new(/*enabled*/ true); + broker.configure(&config); + let mut env = HashMap::from([ + ("VENDOR_PASSWORD".to_string(), token.to_string()), + ("AUTH_HEADER".to_string(), format!("Bearer {token}")), + ]); + + for environment_id in [None, Some("child")] { + broker.virtualize_child_env_for_environment(&mut env, environment_id); + assert!(!env.contains_key("VENDOR_HOST")); + assert!(!env.contains_key("UNRELATED_SECRET")); + assert_ne!(env["AUTH_HEADER"], format!("Bearer {token}")); + let dummy_header = env["AUTH_HEADER"].clone(); + env.remove("VENDOR_PASSWORD"); + broker.virtualize_child_env_for_environment(&mut env, environment_id); + assert_eq!(env["AUTH_HEADER"], dummy_header); + let mut snapshot = format!("export AUTH_HEADER='Bearer {token}'"); + assert!(broker.virtualize_text(&mut snapshot, &env)); + assert_eq!(snapshot, format!("export AUTH_HEADER='{dummy_header}'")); + for host in [destination, "https://other.example/v2"] { + let mut headers = HeaderMap::new(); + headers.insert(AUTHORIZATION, HeaderValue::from_str(&dummy_header).unwrap()); + broker.inject_request_headers_for_environment(host, &mut headers, environment_id); + assert_eq!( + headers[AUTHORIZATION], + if host == destination { + format!("Bearer {token}") + } else { + dummy_header.clone() + } + ); + } + } + + for host in ["https://override.example/v3", "", "not a valid destination"] { + let mut env = HashMap::from([ + ("VENDOR_PASSWORD".to_string(), token.to_string()), + ("VENDOR_HOST".to_string(), host.to_string()), + ]); + broker.virtualize_child_env(&mut env); + assert_eq!(env["VENDOR_HOST"], host); + assert_eq!( + env["VENDOR_PASSWORD"] == token, + !host.starts_with("https://") + ); + let dummy_header = format!("Bearer {}", env["VENDOR_PASSWORD"]); + let mut headers = HeaderMap::new(); + headers.insert(AUTHORIZATION, HeaderValue::from_str(&dummy_header).unwrap()); + broker.inject_request_headers(destination, &mut headers); + assert_eq!(headers[AUTHORIZATION], format!("Bearer {token}")); + if host.starts_with("https://") { + headers.insert(AUTHORIZATION, HeaderValue::from_str(&dummy_header).unwrap()); + broker.inject_request_headers(host, &mut headers); + assert_eq!(headers[AUTHORIZATION], format!("Bearer {token}")); + } + } +} + +#[test] +fn configured_provider_preserves_credentials_without_a_resolved_destination() { + let token = "stripe_live_abcdefghijklmnopqrstuvwx"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["STRIPE_API_KEY".to_string()], + patterns: vec!["^stripe_live_[a-z]{24}$".to_string()], + url_prefix_from_env: Some("STRIPE_HOST".to_string()), + ..CredentialProviderConfig::default() + }); + let mut bound_env = HashMap::from([ + ("STRIPE_API_KEY".to_string(), token.to_string()), + ("STRIPE_HOST".to_string(), "api.stripe.com".to_string()), + ]); + broker.virtualize_child_env(&mut bound_env); + assert_ne!(bound_env["STRIPE_API_KEY"], token); + + let mut env = HashMap::from([("STRIPE_API_KEY".to_string(), token.to_string())]); + + broker.virtualize_child_env(&mut env); + + assert_eq!(env["STRIPE_API_KEY"], token); + assert!(broker.environment(&env).credential_keys.is_empty()); +} + +#[test] +fn configured_provider_does_not_guess_ambiguous_alias_destinations() { + for (token, pattern, overlaps_builtin) in [ + ( + "shared_abcdefghijklmnopqrstuvwxyz", + "^shared_[a-z]{26}$", + false, + ), + ( + "sk-proj-abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789", + "^sk-proj-[A-Za-z0-9]+$", + true, + ), + ] { + let first = CredentialProviderConfig { + env: vec!["FIRST_TOKEN".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: vec!["first.example".to_string()], + ..CredentialProviderConfig::default() + }; + let mut providers = BTreeMap::from([("first".to_string(), first.clone())]); + if !overlaps_builtin { + providers.insert( + "second".to_string(), + CredentialProviderConfig { + env: vec!["SECOND_TOKEN".to_string()], + url_prefixes: vec!["second.example".to_string()], + ..first + }, + ); + } + let broker = CredentialBroker::new(/*enabled*/ true); + broker.configure(&NetworkProxyConfig { + credential_broker: true, + credential_providers: providers, + ..NetworkProxyConfig::default() + }); + let authorization_header = format!("Bearer {token}"); + let mut env = HashMap::from([("AUTH_HEADER".to_string(), authorization_header.clone())]); + + broker.virtualize_child_env(&mut env); + + assert_eq!(env["AUTH_HEADER"], authorization_header); + assert!(!broker.host_requires_mitm("first.example", /*port*/ 443)); + assert!(!broker.host_requires_mitm("second.example", /*port*/ 443)); + if overlaps_builtin { + assert!(!broker.host_requires_mitm("api.openai.com", /*port*/ 443)); + } + + env.insert("FIRST_TOKEN".to_string(), token.to_string()); + broker.virtualize_child_env(&mut env); + assert_ne!(env["AUTH_HEADER"], authorization_header); + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&env["AUTH_HEADER"]).expect("valid dummy authentication"), + ); + broker.inject_request_headers("https://first.example/", &mut headers); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(authorization_header.as_str()) + ); + } +} + +#[test] +fn configured_owner_keeps_prefixed_builtin_shaped_aliases_scoped() { + let aliases = [ + ("COPIED", "x", ""), + ("SUFFIXED", "", "_prod"), + ("WRAPPED", "backup_", "-staging"), + ( + "LONG_SUFFIX", + "", + "_production_shared_service_account_westus2", + ), + ]; + for (pattern, token_length, dummy_prefix) in [ + (r"\bsk-proj-[a-z]{64}\b", 64, Some("sk-proj-")), + // The fixed real branch forces an independently generated vendor-shaped dummy. + ( + r"\b(?:sk-proj-a{64}|vendor_[a-z]{64})\b", + 64, + Some("vendor_"), + ), + // With no distinct long alternative, registration must remain fail-open. + (r"\b(?:sk-proj-a{64}|pin)\b", 64, None), + (r"\b(?:sk-proj-[a-z]{64}|pin)\b", 64, Some("sk-proj-")), + (r"^sk-proj-[a-z]{7}$", 7, Some("sk-proj-")), + ] { + let token = format!("sk-proj-{}", "a".repeat(token_length)); + for (environment_id, retain_source) in [ + (None, true), + (Some("local"), true), + (None, false), + (Some("local"), false), + ] { + for (static_destination, destination) in [ + (None, Some("https://configured.example/v1")), + (Some("https://configured.example/v1"), None), + (None, None), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["VENDOR_KEY".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: static_destination.into_iter().map(str::to_string).collect(), + url_prefix_from_env: Some("VENDOR_URL".to_string()), + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("VENDOR_KEY".to_string(), token.clone())]); + env.extend(aliases.map(|(key, prefix, suffix)| { + (key.to_string(), format!("{prefix}{token}{suffix}")) + })); + if let Some(destination) = destination { + env.insert("VENDOR_URL".to_string(), destination.to_string()); + } + if !retain_source { + let parent_env = env.clone(); + env.remove("VENDOR_KEY"); + broker.discover_parent_credentials_for_environment( + &parent_env, + &env, + environment_id, + ); + } + let mut original_env = env.clone(); + + broker.virtualize_child_env_for_environment(&mut env, environment_id); + + // Arrive after registration so environment collision checks cannot mask a short dummy. + let mut ordinary = "spinning".to_string(); + assert!(!broker.restore_text(&mut ordinary)); + assert_eq!(ordinary, "spinning"); + env.insert("UNRELATED".to_string(), ordinary.clone()); + original_env.insert("UNRELATED".to_string(), ordinary); + let brokered = (static_destination.is_some() || destination.is_some()) + && dummy_prefix.is_some(); + let value = env["COPIED"].strip_prefix('x').unwrap(); + assert_eq!( + env.get("VENDOR_KEY").map(String::as_str), + retain_source.then_some(value) + ); + assert_eq!(value != token, brokered); + assert!( + !broker + .host_protocols_for_environment( + "api.openai.com", + /*port*/ 443, + environment_id, + ) + .tls + ); + let headers_for = |value: &str| { + HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {value}")).unwrap(), + )]) + }; + for target in [ + "https://api.openai.com/v1/models", + "https://api.github.com/v1/models", + "https://configured.example/v2/models", + "https://configured.example/v1/models", + ] { + let mut headers = headers_for(value); + broker.inject_request_headers_for_environment( + target, + &mut headers, + environment_id, + ); + let expected = if brokered && target == "https://configured.example/v1/models" { + &token + } else { + value + }; + assert_eq!(headers, headers_for(expected)); + } + let mut alias_only = env.clone(); + alias_only.remove("VENDOR_KEY"); + broker.virtualize_child_env_for_environment(&mut alias_only, environment_id); + for (key, prefix, suffix) in aliases { + assert_eq!(env[key], format!("{prefix}{value}{suffix}")); + assert_eq!(alias_only[key], env[key]); + assert!(broker.provider_sources_allowed( + &original_env[key], + &env[key], + &original_env, + |source| source == "VENDOR_KEY" + )); + assert!(!broker.provider_sources_allowed( + &original_env[key], + &env[key], + &original_env, + |source| source == "OPENAI_API_KEY" + )); + let mut snapshot_alias = original_env[key].clone(); + if brokered || token.len() >= super::MIN_EMBEDDED_CREDENTIAL_LENGTH { + assert_eq!( + broker.virtualize_text(&mut snapshot_alias, &alias_only), + brokered, + ); + assert!(!snapshot_alias.contains(&token)); + } + if brokered { + assert!(value.starts_with(dummy_prefix.unwrap())); + assert!( + value.len() >= token.len().min(super::MIN_EMBEDDED_CREDENTIAL_LENGTH) + ); + assert!( + brokered_credential_marker_env_keys(&alias_only) + .contains(&key.to_string()) + ); + assert_eq!(snapshot_alias, env[key]); + assert!( + broker + .environment_for_text(&snapshot_alias, &env) + .credential_keys + .contains(&"VENDOR_KEY".to_string()) + ); + assert!(broker.source_matches_text( + "VENDOR_KEY", + &token, + &original_env[key] + )); + assert!(broker.child_alias_matches( + key, + &env[key], + &snapshot_alias, + environment_id, + )); + assert!(broker.restore_text(&mut snapshot_alias)); + assert_eq!(snapshot_alias, original_env[key]); + } + } + broker.restore_and_disable_child_env(&mut env, &mut []); + assert_eq!(env, original_env); + } + } + } +} + +#[test] +fn known_credential_spans_preserve_overlaps_and_uncovered_adjacent_credentials() { + let short = format!("sk-proj-{}", "a".repeat(64)); + let long = format!("{short}bbbbbbbb"); + let known = format!(":known_{}", "b".repeat(24)); + let github = "ghp_spinningabcdefghijklmnopqrstuvwxyz0123456789"; + for (separator, extra, extra_pattern) in [ + ("", "extra_abcdefghijklmnopqrstuvwx", "^extra_[a-z]{24}$"), + ("_", "extra_abcdefghijklmnopqrstuvwx", "^extra_[a-z]{24}$"), + ("-", "extra_abcdefghijklmnopqrstuvwx", "^extra_[a-z]{24}$"), + ("", "extra_abcdefghijklmnopqrstuvwx_", "^extra_[a-z]{24}_$"), + ("", "extra_abcdefghijklmnopqrstuvwx-", "^extra_[a-z]{24}-$"), + ] { + let broker = CredentialBroker::new(/*enabled*/ true); + broker.configure(&NetworkProxyConfig { + credential_broker: true, + credential_providers: BTreeMap::from([ + ( + "vendor".to_string(), + CredentialProviderConfig { + env: vec![ + "SHORT_KEY".to_string(), + "LONG_KEY".to_string(), + "KNOWN_KEY".to_string(), + "PIN_KEY".to_string(), + ], + patterns: vec![ + r"\bsk-proj-[ab]{64,72}\b".to_string(), + "^:known_[a-z]{24}$".to_string(), + r"\bpin\b".to_string(), + ], + url_prefixes: vec!["https://vendor.example/v1".to_string()], + ..CredentialProviderConfig::default() + }, + ), + ( + "extra".to_string(), + CredentialProviderConfig { + env: vec!["EXTRA_KEY".to_string()], + // The broad Basic-capable alternative must not discover masked known values. + patterns: vec![extra_pattern.to_string(), "(?s)^.{80,}$".to_string()], + url_prefixes: vec!["https://extra.example/v1".to_string()], + auth: vec![CredentialAuthMethod::Bearer, CredentialAuthMethod::Basic], + ..CredentialProviderConfig::default() + }, + ), + ]), + ..NetworkProxyConfig::default() + }); + let original = HashMap::from([ + ("SHORT_KEY".to_string(), short.clone()), + ("LONG_KEY".to_string(), long.clone()), + ("KNOWN_KEY".to_string(), known.clone()), + ("PIN_KEY".to_string(), "pin".to_string()), + // This observed short generic owner must not hide the distinct full GitHub token. + ("GH_TOKEN".to_string(), "ghp_".to_string()), + ( + "BUNDLE".to_string(), + format!("{github}{separator}{long}{separator}{extra}{known}"), + ), + ("OVERLAP".to_string(), format!("{long}_{short}")), + ]); + assert!(!broker.provider_sources_allowed( + &format!("{extra}{known}"), + "", + &original, + |source| source == "KNOWN_KEY" + )); + let mut env = original.clone(); + broker.virtualize_child_env(&mut env); + let state = broker.read_state(); + assert_eq!(state.credentials.len(), 6); + assert!( + state + .credentials + .iter() + .all(|credential| !credential.real_value.contains('\0')) + ); + let dummy = |real| { + state + .credentials + .iter() + .find(|credential| credential.real_value == real) + .unwrap_or_else(|| { + panic!("missing registration for {real}, separator {separator:?}") + }) + .dummy_value + .clone() + }; + let github_dummy = dummy(github); + let extra_dummy = dummy(extra); + assert_eq!( + env["BUNDLE"], + format!( + "{github_dummy}{separator}{}{separator}{extra_dummy}{}", + env["LONG_KEY"], env["KNOWN_KEY"] + ) + ); + assert_eq!( + env["OVERLAP"], + format!("{}_{}", env["LONG_KEY"], env["SHORT_KEY"]) + ); + drop(state); + for (key, allowed) in [("LONG_KEY", true), ("SHORT_KEY", false)] { + assert_eq!( + broker.provider_sources_allowed( + &long, + &env["LONG_KEY"], + &original, + |source| source == key + ), + allowed + ); + } + assert!(!broker.provider_sources_allowed( + &original["BUNDLE"], + &env["BUNDLE"], + &original, + |source| source == "LONG_KEY" + )); + assert!(broker.provider_sources_allowed( + &original["BUNDLE"], + &env["BUNDLE"], + &original, + |_| true + )); + for (real, dummy, destination) in [ + (github, &github_dummy, "https://api.github.com/v1"), + (extra, &extra_dummy, "https://extra.example/v1"), + ] { + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")).unwrap(), + )]); + broker.inject_request_headers(destination, &mut headers); + assert_eq!(headers[AUTHORIZATION], format!("Bearer {real}")); + } + broker.restore_and_disable_child_env(&mut env, &mut []); + assert_eq!(env, original); + } +} + +#[test] +fn configured_provider_does_not_claim_embedded_builtin_credentials() { + let github_token = "ghp_abcdefghijklmnopqrstuvwxyz1234567890"; + let openai_token = "sk-proj-abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["COMBINED_TOKEN".to_string()], + patterns: vec!["^combo_[A-Za-z0-9_-]{80,180}$".to_string()], + url_prefixes: vec!["combo.example".to_string()], + ..CredentialProviderConfig::default() + }); + let combined = format!("combo_{github_token}_{openai_token}"); + let mut env = HashMap::from([("AUTH_HEADER".to_string(), combined.clone())]); + + broker.virtualize_child_env(&mut env); + + assert_ne!(env["AUTH_HEADER"], combined); + assert!(!env["AUTH_HEADER"].contains(github_token)); + assert!(!env["AUTH_HEADER"].contains(openai_token)); + assert!(!broker.host_requires_mitm("combo.example", /*port*/ 443)); + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&env["AUTH_HEADER"]).expect("valid dummy authentication"), + ); + broker.inject_request_headers("https://combo.example/", &mut headers); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(env["AUTH_HEADER"].as_str()) + ); +} + +#[test] +fn configured_discovery_does_not_claim_overlapping_credential_spans() { + let suffix = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUV"; + for token in [format!("sk-{suffix}"), format!("sk-proj-{suffix}{suffix}")] { + for patterns in [ + vec!["^[A-Za-z0-9]{48}$".to_string()], + vec![ + "^vendor_[a-z]{24}$".to_string(), + "^[A-Za-z0-9]{48}$".to_string(), + ], + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["VENDOR_TOKEN".to_string()], + patterns, + url_prefixes: vec!["vendor.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("AUTH_HEADER".to_string(), format!("Bearer {token}"))]); + broker.virtualize_child_env(&mut env); + assert!(!broker.host_requires_mitm("vendor.example", /*port*/ 443)); + assert!(broker.read_state().credentials.iter().all(|credential| { + !matches!( + credential.provider, + BrokeredCredentialProvider::Configured(_) + ) + })); + + // An explicit canonical source still identifies the provider unambiguously. + env.insert("VENDOR_TOKEN".to_string(), suffix.to_string()); + broker.virtualize_child_env(&mut env); + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {}", env["VENDOR_TOKEN"])).unwrap(), + )]); + broker.inject_request_headers("https://vendor.example/", &mut headers); + assert_eq!(headers[AUTHORIZATION], format!("Bearer {suffix}")); + } + } + + for (value, outer, inner) in [ + ( + format!("outer_{suffix}"), + "^outer_[A-Za-z0-9]{48}$", + "^[A-Za-z0-9]{48}$", + ), + ( + "first_aaaa.second_bbbb/cccc".to_string(), + r"^first_[a-z]{4}\.second_[a-z]{4}$", + "^second_[a-z]{4}/[a-z]{4}$", + ), + ] { + let broker = CredentialBroker::new(/*enabled*/ true); + broker.configure(&NetworkProxyConfig { + credential_broker: true, + credential_providers: [("outer", outer), ("inner", inner)] + .into_iter() + .map(|(id, pattern)| { + ( + id.to_string(), + CredentialProviderConfig { + env: vec![format!("{}_TOKEN", id.to_uppercase())], + patterns: vec![pattern.to_string()], + url_prefixes: vec![format!("{id}.example")], + ..CredentialProviderConfig::default() + }, + ) + }) + .collect(), + ..NetworkProxyConfig::default() + }); + let mut env = HashMap::from([("AUTH_HEADER".to_string(), format!("Bearer {value}"))]); + broker.virtualize_child_env(&mut env); + assert!(broker.read_state().credentials.is_empty()); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {value}")); + } +} + +#[test] +fn configured_dummy_does_not_embed_another_provider_credential() { + let broker = CredentialBroker::new(/*enabled*/ true); + broker.configure(&NetworkProxyConfig { + credential_broker: true, + credential_providers: BTreeMap::from([ + ( + "backup".to_string(), + CredentialProviderConfig { + env: vec!["BACKUP_TOKEN".to_string()], + patterns: vec!["^backup_(?:stage|prod)$".to_string()], + url_prefixes: vec!["backup.example".to_string()], + ..CredentialProviderConfig::default() + }, + ), + ( + "tenant".to_string(), + CredentialProviderConfig { + env: vec!["TENANT_TOKEN".to_string()], + patterns: vec!["^[a-z]{4}$".to_string()], + url_prefixes: vec!["tenant.example".to_string()], + ..CredentialProviderConfig::default() + }, + ), + ]), + ..NetworkProxyConfig::default() + }); + let real_value = "backup_stage"; + let mut env = HashMap::from([("BACKUP_TOKEN".to_string(), real_value.to_string())]); + + broker.virtualize_child_env(&mut env); + + assert_eq!(env["BACKUP_TOKEN"], real_value); +} + +#[test] +fn discovered_credential_aliases_rebind_when_the_destination_changes() { + for (token, host_key) in [ + ("stripe_live_abcdefghijklmnopqrstuvwx", "STRIPE_HOST"), + ("ghp_abcdefghijklmnopqrstuvwxyz1234567890", "GH_HOST"), + ( + "sk-proj-abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789", + "OPENAI_BASE_URL", + ), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["STRIPE_API_KEY".to_string()], + patterns: vec!["^stripe_live_[a-z]{24}$".to_string()], + url_prefix_from_env: Some("STRIPE_HOST".to_string()), + ..CredentialProviderConfig::default() + }); + let authorization_header = format!("Bearer {token}"); + for host in ["first.example", "second.example"] { + let host_value = if host_key == "OPENAI_BASE_URL" { + format!("https://{host}/v1") + } else { + host.to_string() + }; + let mut env = HashMap::from([ + (host_key.to_string(), host_value), + ("AUTH_HEADER".to_string(), authorization_header.clone()), + ]); + + broker.virtualize_child_env(&mut env); + + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&env["AUTH_HEADER"]).expect("valid dummy authentication"), + ); + broker.inject_request_headers(&format!("https://{host}/v1"), &mut headers); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(authorization_header.as_str()), + "provider host: {host}" + ); + } + + let mut unbound = + HashMap::from([("AUTH_HEADER".to_string(), authorization_header.clone())]); + broker.virtualize_child_env(&mut unbound); + if host_key == "OPENAI_BASE_URL" { + assert_ne!(unbound["AUTH_HEADER"], authorization_header); + } else { + assert_eq!(unbound["AUTH_HEADER"], authorization_header); + assert!(!broker.host_requires_mitm("api.github.com", /*port*/ 443)); + } + } +} + +#[test] +fn credential_rotation_preserves_alias_destination_ownership() { + for (key, rotated_key, host_key, prefix, static_host) in [ + ( + "PROVIDER_TOKEN", + "PROVIDER_TOKEN", + "PROVIDER_URL", + "provider_", + None, + ), + ( + "PROVIDER_TOKEN", + "PROVIDER_TOKEN_FALLBACK", + "PROVIDER_URL", + "provider_", + None, + ), + ( + "PROVIDER_TOKEN", + "PROVIDER_TOKEN", + "PROVIDER_URL", + "provider_", + Some("static.example"), + ), + ( + "GH_ENTERPRISE_TOKEN", + "GH_ENTERPRISE_TOKEN", + "GH_HOST", + "ghp_", + None, + ), + ( + "GH_ENTERPRISE_TOKEN", + "GITHUB_ENTERPRISE_TOKEN", + "GH_HOST", + "ghp_", + None, + ), + ( + "OPENAI_API_KEY", + "OPENAI_API_KEY", + "OPENAI_BASE_URL", + "sk-proj-", + Some("api.openai.com"), + ), + ] { + for (environment, real_aliases, discover_parent) in [ + ("parent", false, false), + ("child", false, false), + ("parent", true, false), + ("child", true, false), + ("parent", true, true), + ("child", true, true), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec![ + "PROVIDER_TOKEN".to_string(), + "PROVIDER_TOKEN_FALLBACK".to_string(), + ], + patterns: vec!["^provider_[a-z]{36}$".to_string()], + url_prefixes: static_host + .map(|host| format!("https://{host}")) + .into_iter() + .collect(), + url_prefix_from_env: Some("PROVIDER_URL".to_string()), + ..CredentialProviderConfig::default() + }); + let suffix_len = if key == "OPENAI_API_KEY" { 64 } else { 36 }; + let first = format!("{prefix}{}", "a".repeat(suffix_len)); + let second = format!("{prefix}{}", "b".repeat(suffix_len)); + let auxiliary = format!("{prefix}{}", "c".repeat(suffix_len)); + let endpoint = |host: &str| { + if host_key == "GH_HOST" { + host.to_string() + } else { + format!("https://{host}/v1") + } + }; + let mut env = HashMap::from([ + (key.to_string(), first.clone()), + (host_key.to_string(), endpoint("first.example")), + ("AUTH_HEADER".to_string(), format!("Bearer {first}")), + ( + "SECOND_AUTH_HEADER".to_string(), + format!("Bearer {auxiliary}"), + ), + ]); + let parent_env = env.clone(); + broker.virtualize_child_env_for_environment(&mut env, Some("parent")); + let old_dummy = env[key].clone(); + let alias_dummy = env["SECOND_AUTH_HEADER"] + .strip_prefix("Bearer ") + .unwrap() + .to_string(); + assert_ne!(old_dummy, first); + assert_ne!(alias_dummy, auxiliary); + env.remove(key); + env.insert(rotated_key.to_string(), second.clone()); + env.insert(host_key.to_string(), endpoint("second.example")); + if real_aliases { + env.insert("AUTH_HEADER".to_string(), format!("Bearer {first}")); + env.insert( + "SECOND_AUTH_HEADER".to_string(), + format!("Bearer {auxiliary}"), + ); + } + if discover_parent { + broker.discover_parent_credentials_for_environment( + &parent_env, + &env, + Some(environment), + ); + } + broker.virtualize_child_env_for_environment(&mut env, Some(environment)); + let new_dummy = env[rotated_key].clone(); + assert_eq!( + env["AUTH_HEADER"], + format!("Bearer {old_dummy}"), + "{key}, {environment}, real_aliases={real_aliases}, discover_parent={discover_parent}" + ); + assert_eq!(env["SECOND_AUTH_HEADER"], format!("Bearer {alias_dummy}")); + assert!(!env.values().any(|value| { + [&first, &second, &auxiliary] + .iter() + .any(|real| value.contains(real.as_str())) + })); + for (dummy, real, host, allowed) in [ + (&old_dummy, &first, "second.example", false), + (&old_dummy, &first, "first.example", true), + (&alias_dummy, &auxiliary, "second.example", false), + (&alias_dummy, &auxiliary, "first.example", true), + (&new_dummy, &second, "first.example", false), + (&new_dummy, &second, "second.example", true), + ] { + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")).unwrap(), + )]); + broker.inject_request_headers_for_environment( + &format!("https://{host}/v1"), + &mut headers, + Some(environment), + ); + assert_eq!( + headers[AUTHORIZATION], + format!("Bearer {}", if allowed { real } else { dummy }), + "{key}, {environment}, {host}, real_aliases={real_aliases}, discover_parent={discover_parent}" + ); + } + + env.insert(host_key.to_string(), String::new()); + if real_aliases { + env.insert("AUTH_HEADER".to_string(), format!("Bearer {first}")); + env.insert( + "SECOND_AUTH_HEADER".to_string(), + format!("Bearer {auxiliary}"), + ); + } + if discover_parent { + broker.discover_parent_credentials_for_environment( + &parent_env, + &env, + Some(environment), + ); + } + broker.virtualize_child_env_for_environment(&mut env, Some(environment)); + for host in ["first.example", "second.example"] + .into_iter() + .chain(static_host) + { + for (dummy, real) in [ + (&old_dummy, &first), + (&alias_dummy, &auxiliary), + (&new_dummy, &second), + ] { + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")).unwrap(), + )]); + broker.inject_request_headers_for_environment( + &format!("https://{host}/v1"), + &mut headers, + Some(environment), + ); + assert_eq!( + headers[AUTHORIZATION], + format!( + "Bearer {}", + if Some(host) == static_host { + real + } else { + dummy + } + ) + ); + } + } + let mut parent_headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {old_dummy}")).unwrap(), + )]); + broker.inject_request_headers_for_environment( + "https://first.example/v1", + &mut parent_headers, + Some("parent"), + ); + assert_eq!( + parent_headers[AUTHORIZATION], + format!( + "Bearer {}", + if environment == "child" { + &first + } else { + &old_dummy + } + ) + ); + } + } +} + +#[test] +fn configured_provider_accepts_anchored_alternatives_and_ascii_character_classes() { + for (pattern, token) in [ + ("(?i)^token_[a-z]{24}$", "TOKEN_abcdefghijklmnopqrstuvwx"), + ("(?i:^token_[a-z]{24}$)", "TOKEN_abcdefghijklmnopqrstuvwx"), + ( + "(?x) ^ token_[a-z]{24} $", + "token_abcdefghijklmnopqrstuvwx", + ), + ( + "(?x)^token_[a-z]{24}$ # provider token", + "token_abcdefghijklmnopqrstuvwx", + ), + ( + "(?x)^token_[a-z]{24} # [legacy provider\n$", + "token_abcdefghijklmnopqrstuvwx", + ), + (r"\btoken_[a-z]{24}\b", "token_abcdefghijklmnopqrstuvwx"), + (r"\Atoken_[a-z]{24}\z", "token_abcdefghijklmnopqrstuvwx"), + ( + "^(token_[a-z]{8}|token_[a-z]{24})$", + "token_abcdefghijklmnopqrstuvwx", + ), + (r"^token_\d{24}$", "token_012345678901234567890123"), + ("^token_.{24}$", "token_abcdefghijklmnopqrstuvwx"), + ("^token_[^:]{24}$", "token_abcdefghijklmnopqrstuvwx"), + ("^token_[]$a-z]{8}$", "token_a$bcd]ef"), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("PROVIDER_TOKEN".to_string(), token.to_string())]); + + broker.virtualize_child_env(&mut env); + + let dummy = &env["PROVIDER_TOKEN"]; + assert_ne!(dummy, token, "credential pattern: {pattern}"); + assert!(dummy.is_ascii(), "credential pattern: {pattern}"); + assert!( + regex::Regex::new(pattern) + .expect("valid credential pattern") + .is_match(dummy), + "credential pattern: {pattern}" + ); + + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")).expect("valid dummy authentication"), + ); + broker.inject_request_headers("https://api.provider.example/", &mut headers); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(format!("Bearer {token}").as_str()), + "credential pattern: {pattern}" + ); + + let replacement = if token.ends_with(|character: char| character.is_ascii_digit()) { + '9' + } else { + 'z' + }; + let unknown = format!("{}{replacement}", &token[..token.len() - 1]); + let mut copied = format!("AUTH_HEADER=Bearer {unknown}"); + assert!( + !broker.virtualize_text(&mut copied, &env), + "credential pattern: {pattern}" + ); + assert!(!copied.contains(&unknown), "credential pattern: {pattern}"); + } +} + +#[test] +fn configured_provider_preserves_word_boundaries_during_unbound_discovery() { + let token = "token_abcdefghijklmnopqrstuvwx"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec![r"\btoken_[a-z]{24}\b".to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let prefixed = format!("my{token}"); + let mut env = HashMap::from([ + ("APPLICATION_ID".to_string(), prefixed.clone()), + ("AUTH_HEADER".to_string(), format!("Bearer {token}")), + ]); + + broker.virtualize_child_env(&mut env); + + assert_eq!(env["APPLICATION_ID"], prefixed); + assert_ne!(env["AUTH_HEADER"], format!("Bearer {token}")); +} + +#[test] +fn configured_provider_replaces_exact_known_values_with_lazy_patterns() { + for (pattern, token) in [ + (r"^token_[a-z]+?$", "token_abcdefghijklmnopqrstuvwx"), + ( + r"^token_(?:[a-z]|[a-z]{24})$", + "token_abcdefghijklmnopqrstuvwx", + ), + (r"\btoken_[a-z]+?\b", "token_abcdefghijklmnopqrstuvwx"), + (r"^token_\B[a-z]+?$", "token_abcdefghijklmnopqrstuvwx"), + (r"^[a-z]{24}\b$", "abcabcabcabcabcabcabcabc"), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([ + ("PROVIDER_TOKEN".to_string(), token.to_string()), + ("AUTH_HEADER".to_string(), format!("Bearer {token}")), + ("REPEATED".to_string(), format!("{token}\n{token}")), + ("ADJACENT".to_string(), format!("é{token}é")), + ("OVERLAPPED".to_string(), format!("Bearer abc{token}")), + ]); + + broker.virtualize_child_env(&mut env); + + let dummy = &env["PROVIDER_TOKEN"]; + assert_ne!(dummy, token); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {dummy}")); + assert_eq!(env["REPEATED"], format!("{dummy}\n{dummy}")); + assert_eq!( + env["ADJACENT"], + format!( + "é{}é", + if pattern.contains(r"\b") { + token + } else { + dummy + } + ) + ); + assert_eq!( + env["OVERLAPPED"], + format!( + "Bearer abc{}", + if pattern.starts_with(r"\b") { + token + } else { + dummy + } + ) + ); + let mut snapshot = format!("{token}\nBearer {token}\n{token}"); + assert!(broker.virtualize_text(&mut snapshot, &env)); + assert_eq!(snapshot, format!("{dummy}\nBearer {dummy}\n{dummy}")); + } +} + +#[test] +fn configured_provider_discovers_complete_unregistered_values() { + for (pattern, token) in [ + (r"^token_[a-z]+?$", "token_abcdefghijklmnopqrstuvwx"), + (r"\btoken_[a-z]+?\b", "token_abcdefghijklmnopqrstuvwx"), + (r"^token_\B[a-z]+?$", "token_abcdefghijklmnopqrstuvwx"), + ( + r"^token_(?:[a-z]{8,32}|[a-z]{24}[0-9]{8})$", + "token_abcdefghijklmnopqrstuvwx12345678", + ), + ( + r"^token_(?:[a-z]+|[a-z]+[0-9]+)$", + "token_abcdefghijklmnopqrstuvwx12345678", + ), + ( + r"^token_(?:[a-z]{8,32}|[a-z]{24}_[0-9]{7})$", + "token_abcdefghijklmnopqrstuvwx_1234567", + ), + ( + r"^token_(?:[a-z]{8,32}|[a-z]{24}-[0-9]{7})$", + "token_abcdefghijklmnopqrstuvwx-1234567", + ), + ( + r"^token_(?:[a-z]+|[a-z]+/[a-z]+)$", + "token_abcdefghijklmnopqrstuvwx/zyxwvutsrqponmlkjihgfedcb", + ), + ( + r"^token_(?:[a-z]{8,128}|[a-z]{24}\.[a-z]{24})$", + "token_abcdefghijklmnopqrstuvwx.zyxwvutsrqponmlkjihgfedc", + ), + ( + r"\btoken_(?:[a-z]+|[a-z]+\+[a-z]+)\b", + "token_abcdefghijklmnopqrstuvwx+zyxwvutsrqponmlkjihgfedcb", + ), + ( + r"^token_(?:[a-z]+|[a-z]+:[a-z]+)$", + "token_abcdefghijklmnopqrstuvwx:zyxwvutsrqponmlkjihgfedcb", + ), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([ + ("AUTH_HEADER".to_string(), format!("Bearer {token}")), + ("REPEATED".to_string(), format!("{token},{token}")), + ]); + + let mut unknown_snapshot = format!("export AUTH_HEADER='Bearer {token}'"); + assert!( + !broker.virtualize_text(&mut unknown_snapshot, &HashMap::new()), + "{pattern}" + ); + assert_eq!( + unknown_snapshot, "export AUTH_HEADER='Bearer '", + "{pattern}" + ); + broker.virtualize_child_env(&mut env); + + let dummy = env["AUTH_HEADER"].strip_prefix("Bearer ").unwrap(); + assert_ne!(dummy, token, "{pattern}"); + assert_eq!(env["REPEATED"], format!("{dummy},{dummy}"), "{pattern}"); + let mut snapshot = format!("Bearer {token}"); + assert!(broker.virtualize_text(&mut snapshot, &env), "{pattern}"); + assert_eq!(snapshot, env["AUTH_HEADER"], "{pattern}"); + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&env["AUTH_HEADER"]).unwrap(), + )]); + broker.inject_request_headers("https://api.provider.example/", &mut headers); + assert_eq!( + headers, + HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {token}")).unwrap(), + )]), + "{pattern}" + ); + } +} + +#[test] +fn mixed_provider_aliases_keep_credentials_on_separate_destinations() { + let github = "ghp_abcdefghijklmnopqrstuvwxyz1234567890"; + for (pattern, vendor, separator) in [ + ("^vendor_[a-z]{24}$", "vendor_abcdefghijklmnopqrstuvwx"), + ( + r"^vendor_[a-z]{8}(?:\.[a-z]{8})?$", + "vendor_abcdefgh.ijklmnop", + ), + ] + .into_iter() + .flat_map(|(pattern, vendor)| ["_", "-", ""].map(|separator| (pattern, vendor, separator))) + { + let broker = broker_for(CredentialProviderConfig { + env: vec!["VENDOR_TOKEN".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: vec!["api.vendor.example".to_string()], + ..CredentialProviderConfig::default() + }); + let composite = format!("{github}{separator}{vendor}"); + let mut env = HashMap::from([("AUTH_BUNDLE".to_string(), composite.clone())]); + broker.virtualize_child_env(&mut env); + let records = broker.read_state(); + let github_dummy = records + .credentials + .iter() + .find(|record| record.real_value == github) + .expect("separate GitHub credential") + .dummy_value + .clone(); + let vendor_dummy = records + .credentials + .iter() + .find(|record| record.real_value == vendor) + .expect("separate vendor credential") + .dummy_value + .clone(); + drop(records); + assert_eq!( + env["AUTH_BUNDLE"], + format!("{github_dummy}{separator}{vendor_dummy}") + ); + let mut snapshot = composite; + // Without canonical sources, the snapshot cannot authorize a multi-token alias. + assert!(!broker.virtualize_text(&mut snapshot, &env)); + assert!(!snapshot.contains(github)); + assert!(!snapshot.contains(vendor)); + for (host, dummy, expected) in [ + ("api.github.com", &github_dummy, github), + ("api.vendor.example", &vendor_dummy, vendor), + ("api.github.com", &vendor_dummy, vendor_dummy.as_str()), + ] { + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")).unwrap(), + )]); + broker.inject_request_headers(&format!("https://{host}/"), &mut headers); + assert_eq!( + headers, + HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {expected}")).unwrap(), + )]) + ); + } + let composite = format!("{github}{separator}{vendor}"); + let mut canonical = HashMap::from([("GH_TOKEN".to_string(), composite.clone())]); + broker.virtualize_child_env(&mut canonical); + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {}", canonical["GH_TOKEN"])).unwrap(), + )]); + broker.inject_request_headers("api.github.com", &mut headers); + assert_eq!( + headers, + HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {composite}")).unwrap(), + )]) + ); + } +} + +#[test] +fn configured_provider_virtualizes_short_credentials_in_aliases_and_snapshots() { + for (pattern, token, unregistered) in [ + ("^pin_[a-z]{8}$", "pin_abcdefgh", "pin_hgfedcba"), + (r"\bpin_[a-z]{8}\b", "pin_abcdefgh", "pin_hgfedcba"), + ( + "^(pin_[a-z]{4}|pin_[a-z]{8})$", + "pin_abcdefgh", + "pin_hgfedcba", + ), + ("(?i)^pin_[a-z]{8}$", "PIN_abcdefgh", "PIN_hgfedcba"), + ( + "^(pin_[a-z]{8}|key_[a-z]{8})$", + "key_abcdefgh", + "key_hgfedcba", + ), + ("^pin/[a-z]{8}$", "pin/abcdefgh", "pin/hgfedcba"), + (r"^pin\\[a-z]{8}$", r"pin\abcdefgh", r"pin\hgfedcba"), + ] { + let preserves_word_boundaries = pattern.contains(r"\b"); + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_PASSWORD".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([ + ("PROVIDER_PASSWORD".to_string(), token.to_string()), + ("AUTH_HEADER".to_string(), format!("Bearer {token}")), + ("BACKUP".to_string(), format!("backup_{token}_prod")), + ( + "PATH".to_string(), + format!("/opt/venvs/build_{token}_prod/bin:/usr/bin"), + ), + ( + "TOOL_PATH".to_string(), + format!("tools/{token}/bin:/usr/bin"), + ), + ( + "WINDOWS_TOOL_PATH".to_string(), + format!(r"tools\{token}\bin"), + ), + ("UNRELATED".to_string(), "pin_setting".to_string()), + ]); + + broker.virtualize_child_env(&mut env); + + let dummy = env["PROVIDER_PASSWORD"].clone(); + assert_ne!(dummy, token, "credential pattern: {pattern}"); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {dummy}")); + let mut alias_only = env.clone(); + alias_only.remove("PROVIDER_PASSWORD"); + broker.virtualize_child_env(&mut alias_only); + assert!( + brokered_credential_marker_env_keys(&alias_only) + .contains(&"PROVIDER_PASSWORD".to_string()) + ); + assert!( + brokered_credential_marker_env_keys(&alias_only).contains(&"AUTH_HEADER".to_string()) + ); + assert!( + brokered_credential_value_env_keys(&alias_only).contains(&"AUTH_HEADER".to_string()) + ); + assert_eq!( + env["BACKUP"], + if preserves_word_boundaries { + format!("backup_{token}_prod") + } else { + format!("backup_{dummy}_prod") + } + ); + assert_eq!( + env["PATH"], + format!("/opt/venvs/build_{token}_prod/bin:/usr/bin") + ); + assert_eq!(env["TOOL_PATH"], format!("tools/{token}/bin:/usr/bin")); + assert_eq!(env["WINDOWS_TOOL_PATH"], format!(r"tools\{token}\bin")); + assert_eq!(env["UNRELATED"], "pin_setting"); + let mut snapshot = format!("export AUTH_HEADER='Bearer {token}'"); + assert!(broker.virtualize_text(&mut snapshot, &env)); + assert_eq!(snapshot, format!("export AUTH_HEADER='Bearer {dummy}'")); + for original_path in [ + "/opt/venvs/pin_hgfedcba/bin:/usr/bin", + "tools/pin_hgfedcba/bin:/usr/bin", + r"tools\pin_hgfedcba\bin", + ] { + let mut path = original_path.to_string(); + assert!(broker.virtualize_text(&mut path, &env)); + assert_eq!(path, original_path); + } + for (mut copied, expected_allowed) in [ + ( + format!("BACKUP=backup_{unregistered}_prod"), + preserves_word_boundaries, + ), + ( + format!("WEBHOOK_URL=https://api.vendor.example/{unregistered}"), + false, + ), + ] { + assert_eq!( + broker.virtualize_text(&mut copied, &env), + expected_allowed, + "credential pattern: {pattern}", + ); + assert_eq!( + copied.contains(unregistered), + expected_allowed, + "credential pattern: {pattern}", + ); + } + + broker.restore_child_env(&mut env, &mut []); + assert_eq!(env["PROVIDER_PASSWORD"], token); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {token}")); + assert_eq!(env["BACKUP"], format!("backup_{token}_prod")); + assert_eq!( + env["PATH"], + format!("/opt/venvs/build_{token}_prod/bin:/usr/bin") + ); + assert_eq!(env["TOOL_PATH"], format!("tools/{token}/bin:/usr/bin")); + assert_eq!(env["WINDOWS_TOOL_PATH"], format!(r"tools\{token}\bin")); + } +} + +#[test] +fn configured_provider_generates_bounded_independent_dummies() { + let long_token = format!("token_{}", "a".repeat(4096)); + let long_broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["^token_[a-z]+$".to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut long_env = HashMap::from([("PROVIDER_TOKEN".to_string(), long_token.clone())]); + + long_broker.virtualize_child_env(&mut long_env); + + let long_dummy = &long_env["PROVIDER_TOKEN"]; + assert_ne!(long_dummy, &long_token); + assert!(long_dummy.len() <= 2048); + + let narrow_token = format!("token_{}", "a".repeat(64)); + for _ in 0..3 { + let narrow_broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["^token_[ab]{64}$".to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut narrow_env = HashMap::from([("PROVIDER_TOKEN".to_string(), narrow_token.clone())]); + + narrow_broker.virtualize_child_env(&mut narrow_env); + + let changed = narrow_token + .bytes() + .zip(narrow_env["PROVIDER_TOKEN"].bytes()) + .filter(|(real, dummy)| real != dummy) + .count(); + assert!(changed >= 12, "dummy changed only {changed} bytes"); + } +} + +#[test] +fn configured_provider_rejects_aliases_containing_disallowed_credentials() { + let first = "token_abcdefghijklmnopqrstuvwx"; + let second = "token_zyxwvutsrqponmlkjihgfedc"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["FIRST_TOKEN".to_string(), "SECOND_TOKEN".to_string()], + patterns: vec!["token_[a-z]{24}".to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let source_env = HashMap::from([("FIRST_TOKEN".to_string(), first.to_string())]); + + assert!( + broker + .provider_sources_allowed(first, "", &source_env, |source| { source == "FIRST_TOKEN" }) + ); + assert!(!broker.provider_sources_allowed( + &format!("{first}|{second}"), + "", + &source_env, + |source| source == "FIRST_TOKEN", + )); + + let overlapping = broker_for(CredentialProviderConfig { + env: vec!["FIRST_TOKEN".to_string(), "SECOND_TOKEN".to_string()], + patterns: vec!["^(token_[a-z]{4}|token_[a-z]{8})$".to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let first = "token_abcd"; + let second = "token_abcdefgh"; + for source_env in [ + HashMap::from([ + ("FIRST_TOKEN".to_string(), first.to_string()), + ("SECOND_TOKEN".to_string(), second.to_string()), + ]), + HashMap::from([("FIRST_TOKEN".to_string(), first.to_string())]), + ] { + assert!( + !overlapping.provider_sources_allowed(second, "", &source_env, |source| source + == "FIRST_TOKEN",) + ); + } +} + +#[test] +fn configured_provider_preserves_operational_values_while_redacting_url_credentials() { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_PIN".to_string()], + patterns: vec![r"^\d{3}$".to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([ + ("PROVIDER_PIN".to_string(), "127".to_string()), + ("AUTH_HEADER".to_string(), "Bearer 127".to_string()), + ( + "HTTP_PROXY".to_string(), + "http://127.0.0.1:3128".to_string(), + ), + ("PATH".to_string(), "/opt/sdk127/bin:/usr/bin".to_string()), + ("HTTP_STATUS".to_string(), "443".to_string()), + ("RETRIES".to_string(), "200".to_string()), + ( + "WEBHOOK_URL".to_string(), + "https://127.0.0.1/auth/127".to_string(), + ), + ]); + + broker.virtualize_child_env(&mut env); + + let dummy = env["PROVIDER_PIN"].clone(); + assert_ne!(dummy, "127"); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {dummy}")); + assert_eq!(env["HTTP_PROXY"], "http://127.0.0.1:3128"); + assert_eq!(env["PATH"], "/opt/sdk127/bin:/usr/bin"); + assert_eq!(env["HTTP_STATUS"], "443"); + assert_eq!(env["RETRIES"], "200"); + assert_eq!( + env["WEBHOOK_URL"], + format!("https://127.0.0.1/auth/{dummy}") + ); + let mut ordinary_settings = "export HTTP_STATUS=443\nexport RETRIES=200\n".to_string(); + assert!(broker.virtualize_text(&mut ordinary_settings, &env)); + assert_eq!( + ordinary_settings, + "export HTTP_STATUS=443\nexport RETRIES=200\n" + ); + let mut credential_alias = "AUTH_HEADER=Bearer 127".to_string(); + assert!(broker.virtualize_text(&mut credential_alias, &env)); + assert_eq!(credential_alias, format!("AUTH_HEADER=Bearer {dummy}")); + + broker.restore_child_env(&mut env, &mut []); + assert_eq!(env["PROVIDER_PIN"], "127"); + assert_eq!(env["AUTH_HEADER"], "Bearer 127"); + assert_eq!(env["HTTP_PROXY"], "http://127.0.0.1:3128"); + assert_eq!(env["PATH"], "/opt/sdk127/bin:/usr/bin"); + assert_eq!(env["HTTP_STATUS"], "443"); + assert_eq!(env["RETRIES"], "200"); + assert_eq!(env["WEBHOOK_URL"], "https://127.0.0.1/auth/127"); +} + +#[test] +fn configured_provider_does_not_rescan_registered_credential_prefixes() { + let token = "token_abcdefghijkl"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["token_[a-z]{4}".to_string(), "token_[a-z]{12}".to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("PROVIDER_TOKEN".to_string(), token.to_string())]); + broker.virtualize_child_env(&mut env); + let dummy = &env["PROVIDER_TOKEN"]; + + let mut registered = format!("AUTH_HEADER=Bearer {dummy}"); + assert!(broker.virtualize_text(&mut registered, &env)); + assert_eq!(registered, format!("AUTH_HEADER=Bearer {dummy}")); + + let mut unknown = "AUTH_HEADER=Bearer token_zzzzzzzzzzzz".to_string(); + assert!(!broker.virtualize_text(&mut unknown, &env)); + assert_eq!(unknown, "AUTH_HEADER=Bearer "); +} + +#[test] +fn configured_provider_does_not_register_a_credential_inside_an_existing_dummy() { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PRIMARY_TOKEN".to_string(), "SECONDARY_TOKEN".to_string()], + patterns: vec!["token_[a-z]{4}".to_string(), "token_[a-z]{12}".to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([( + "PRIMARY_TOKEN".to_string(), + "token_abcdefghijkl".to_string(), + )]); + broker.virtualize_child_env(&mut env); + let primary_dummy = env["PRIMARY_TOKEN"].clone(); + let overlapping_real = primary_dummy[..10].to_string(); + env.insert("SECONDARY_TOKEN".to_string(), overlapping_real.clone()); + + broker.virtualize_child_env(&mut env); + + assert_eq!(env["PRIMARY_TOKEN"], primary_dummy); + assert_eq!(env["SECONDARY_TOKEN"], overlapping_real); +} + +#[test] +fn configured_provider_reload_preserves_credentials_and_rejects_overlapping_sources() { + let token = "provider_abcdefghijklmnopqrstuvwx"; + let provider = CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["provider_[a-z]{24}".to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + ..CredentialProviderConfig::default() + }; + let broker = broker_for(provider.clone()); + let mut env = HashMap::from([("PROVIDER_TOKEN".to_string(), token.to_string())]); + broker.virtualize_child_env(&mut env); + let dummy = env["PROVIDER_TOKEN"].clone(); + + let overlapping_builtin = CredentialProviderConfig { + env: vec!["GH_TOKEN".to_string()], + url_prefixes: vec!["builtin-overlap.example".to_string()], + ..provider.clone() + }; + let overlapping_configured = CredentialProviderConfig { + url_prefixes: vec!["configured-overlap.example".to_string()], + ..provider.clone() + }; + let unrelated = CredentialProviderConfig { + env: vec!["ANOTHER_TOKEN".to_string()], + url_prefixes: vec!["another.example".to_string()], + ..provider.clone() + }; + broker.configure(&NetworkProxyConfig { + credential_broker: true, + credential_providers: BTreeMap::from([ + ("aaa-overlap".to_string(), overlapping_configured), + ("builtin".to_string(), overlapping_builtin), + ("custom".to_string(), provider), + ("second".to_string(), unrelated), + ]), + ..NetworkProxyConfig::default() + }); + + broker.virtualize_child_env(&mut env); + assert_eq!(env["PROVIDER_TOKEN"], dummy); + assert!(!broker.host_requires_mitm("builtin-overlap.example", /*port*/ 443)); + assert!(!broker.host_requires_mitm("configured-overlap.example", /*port*/ 443)); + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")).expect("valid authentication"), + ); + broker.inject_request_headers("https://api.provider.example/", &mut headers); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(format!("Bearer {token}").as_str()) + ); +} + +#[test] +fn configured_provider_preserves_bearer_token_basic_and_custom_header_auth() { + let token = "provider_abcdefghijklmnopqrstuvwx"; + let github_token = "ghp_abcdefghijklmnopqrstuvwxyz1234567890"; + for (method, header_name, header_value) in [ + ( + CredentialAuthMethod::Bearer, + "authorization", + format!("Bearer {token}"), + ), + ( + CredentialAuthMethod::Token, + "authorization", + format!("token {token}"), + ), + ( + CredentialAuthMethod::Basic, + "authorization", + format!( + "Basic {}", + base64::engine::general_purpose::STANDARD.encode(format!("user:{token}")) + ), + ), + ( + CredentialAuthMethod::Header, + "x-api-key", + format!("Key {token}"), + ), + ] { + let host = if method == CredentialAuthMethod::Header { + "api.github.com" + } else { + "api.provider.example" + }; + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["provider_[a-z]{24}".to_string()], + url_prefixes: vec![host.to_string()], + auth: if method == CredentialAuthMethod::Header { + vec![CredentialAuthMethod::Bearer, method] + } else { + vec![method] + }, + header: (method == CredentialAuthMethod::Header).then(|| "x-api-key".to_string()), + prefix: (method == CredentialAuthMethod::Header).then(|| "Key ".to_string()), + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("PROVIDER_TOKEN".to_string(), token.to_string())]); + if method == CredentialAuthMethod::Header { + env.insert("GH_TOKEN".to_string(), github_token.to_string()); + } + broker.virtualize_child_env(&mut env); + let dummy = &env["PROVIDER_TOKEN"]; + let dummy_header = if method == CredentialAuthMethod::Basic { + format!( + "Basic {}", + base64::engine::general_purpose::STANDARD.encode(format!("user:{dummy}")) + ) + } else { + header_value.replace(token, dummy) + }; + let mut headers = HeaderMap::new(); + headers.insert( + rama_http::HeaderName::from_bytes(header_name.as_bytes()).expect("valid header"), + HeaderValue::from_str(&dummy_header).expect("valid dummy authentication"), + ); + if method == CredentialAuthMethod::Header { + headers.insert( + AUTHORIZATION, + HeaderValue::from_static("Bearer unrelated-session-token"), + ); + } + + broker.inject_request_headers(&format!("https://{host}/"), &mut headers); + + assert_eq!( + headers + .get(header_name) + .and_then(|value| value.to_str().ok()), + Some(header_value.as_str()), + "authentication method: {method:?}" + ); + if method == CredentialAuthMethod::Header { + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")) + .expect("valid dummy authentication"), + ); + headers.insert( + "x-api-key", + HeaderValue::from_str(&dummy_header).expect("valid dummy authentication"), + ); + broker.inject_request_headers(&format!("https://{host}/"), &mut headers); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(format!("Bearer {token}").as_str()) + ); + assert_eq!( + headers + .get("x-api-key") + .and_then(|value| value.to_str().ok()), + Some(header_value.as_str()) + ); + + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {}", env["GH_TOKEN"])) + .expect("valid dummy authentication"), + ); + headers.insert( + "x-api-key", + HeaderValue::from_str(&dummy_header).expect("valid dummy authentication"), + ); + broker.inject_request_headers(&format!("https://{host}/"), &mut headers); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(format!("Bearer {github_token}").as_str()) + ); + assert_eq!( + headers + .get("x-api-key") + .and_then(|value| value.to_str().ok()), + Some(header_value.as_str()) + ); + } + } +} + +#[test] +fn configured_provider_translates_full_basic_pairs_only_on_exact_match() { + let token = "user:provider_abcdefghijklmnopqrstuvwx:abcd"; + for (pattern, brokered) in [ + ("^user:provider_[a-z]{24}:[a-z]{4}$", true), + ("^[a-z]{4}(?::provider_[a-z]{24}:[a-z]{4})?$", true), + ( + "^(?:user:provider_abcdefghijklmnopqrstuvwx:abcd|[a-z]{32})$", + false, + ), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_CREDENTIALS".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + auth: vec![CredentialAuthMethod::Basic], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("PROVIDER_CREDENTIALS".to_string(), token.to_string())]); + broker.virtualize_child_env(&mut env); + let dummy = &env["PROVIDER_CREDENTIALS"]; + assert_eq!(dummy != token, brokered); + + for (destination, suffix, injected) in [ + ("https://api.provider.example/", "", true), + ("https://api.provider.example/", "extra", false), + ("https://other.example/", "", false), + ] { + // curl --user adds an empty password when the argument has no colon. + let pair = if dummy.contains(':') { + dummy.clone() + } else { + format!("{dummy}:") + }; + let input = format!("{pair}{suffix}"); + let header = |value: &str| { + HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!( + "bAsIc {}", + base64::engine::general_purpose::STANDARD.encode(value) + )) + .expect("valid Basic authentication"), + )]) + }; + let mut headers = header(&input); + + broker.inject_request_headers(destination, &mut headers); + + assert_eq!(headers, header(if injected { token } else { &input })); + } + } +} + +#[test] +fn configured_provider_generates_unicode_basic_dummies() { + for (pattern, secret, token) in [ + ( + "^pass_wörd[a-z]{8} $", + "pass_wördabcdefgh ".to_string(), + "pass_wördabcdefgh ".to_string(), + ), + ( + r"^pass_\B[éöü]{24}[a-z]$", + "éöü".repeat(8), + format!("pass_{}a", "éöü".repeat(8)), + ), + (r"^(?:α{32}|β{32})$", "α".repeat(32), "α".repeat(32)), + ( + r"^(?:α{32}|β{32}|\x00{32})$", + "α".repeat(32), + "α".repeat(32), + ), + ( + r"^pass_\B(?:[éöü]{24}|[-+]{24})[a-z]$", + "éöü".repeat(8), + format!("pass_{}a", "éöü".repeat(8)), + ), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_PASSWORD".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: vec!["api.provider.example".to_string()], + auth: vec![CredentialAuthMethod::Basic], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("PROVIDER_PASSWORD".to_string(), token.clone())]); + + broker.virtualize_child_env(&mut env); + + let dummy = &env["PROVIDER_PASSWORD"]; + assert_ne!(dummy, &token); + assert!(!dummy.contains('\0')); + assert!(!dummy.contains(&secret)); + assert!(regex::Regex::new(pattern).unwrap().is_match(dummy)); + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!( + "Basic {}", + base64::engine::general_purpose::STANDARD.encode(format!("user:{dummy}")) + )) + .expect("valid dummy authentication"), + )]); + broker.inject_request_headers("https://api.provider.example/", &mut headers); + let expected = format!( + "Basic {}", + base64::engine::general_purpose::STANDARD.encode(format!("user:{token}")) + ); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(expected.as_str()) + ); + } +} + +#[test] +fn configured_provider_does_not_break_usable_auth_methods_when_generating_dummies() { + let token = "provider_abcdefghijklmnopqrstuvwx"; + for (alternative, method) in [ + "β{32}", + "provider_[a-z]{23} ", + " provider_[a-z]{23}", + "provider_[a-z]{23}\\t", + "provider_[a-z]{12}:[a-z]{12}", + ] + .into_iter() + .flat_map(|alternative| { + [CredentialAuthMethod::Bearer, CredentialAuthMethod::Header] + .map(|method| (alternative, method)) + }) + .chain(std::iter::once((r"\x00{24}", CredentialAuthMethod::Basic))) + { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_PASSWORD".to_string()], + patterns: vec![format!("^(?:{token}|{alternative})$")], + url_prefixes: vec!["api.provider.example".to_string()], + auth: vec![method, CredentialAuthMethod::Basic], + header: (method == CredentialAuthMethod::Header).then(|| "x-api-key".to_string()), + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("PROVIDER_PASSWORD".to_string(), token.to_string())]); + + broker.virtualize_child_env(&mut env); + + assert_eq!(env["PROVIDER_PASSWORD"], token); + } +} + +#[test] +fn configured_basic_dummies_preserve_username_password_and_whole_value_auth() { + let token = "provider_abcdefghijklmnopqrstuvwx"; + let second = "provider_zyxwvutsrqponmlkjihgfedcb"; + let github = "ghp_abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGH"; + let broker = broker_for(CredentialProviderConfig { + env: vec![ + "PROVIDER_TOKEN".to_string(), + "PROVIDER_SECRET".to_string(), + "PROVIDER_COPY".to_string(), + ], + patterns: vec!["^provider_(?:[a-z]{24}|[a-z]{12}:[a-z]{12})$".to_string()], + url_prefixes: vec!["https://api.github.com/v1".to_string()], + auth: vec![CredentialAuthMethod::Basic], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([ + ("PROVIDER_TOKEN".to_string(), token.to_string()), + ("PROVIDER_SECRET".to_string(), second.to_string()), + ("PROVIDER_COPY".to_string(), token.to_string()), + ("GH_TOKEN".to_string(), github.to_string()), + ]); + broker.virtualize_child_env(&mut env); + let dummy = &env["PROVIDER_TOKEN"]; + assert_ne!(dummy, token); + assert!(!dummy.contains(':')); + for (input, expected) in [ + ( + format!("{dummy}:x-oauth-basic"), + format!("{token}:x-oauth-basic"), + ), + (format!("user:{dummy}"), format!("user:{token}")), + (format!("{dummy}:{dummy}"), format!("{token}:{token}")), + ( + format!("{dummy}:{}", env["PROVIDER_COPY"]), + format!("{token}:{token}"), + ), + ( + format!("{}:{dummy}", env["PROVIDER_COPY"]), + format!("{token}:{token}"), + ), + ( + format!("{dummy}:{}", env["PROVIDER_SECRET"]), + format!("{token}:{second}"), + ), + ( + format!("{}:{dummy}", env["PROVIDER_SECRET"]), + format!("{second}:{token}"), + ), + ( + format!("{}:{dummy}", env["GH_TOKEN"]), + format!("{github}:{token}"), + ), + ( + format!("{dummy}:{}", env["GH_TOKEN"]), + format!("{token}:{github}"), + ), + (dummy.clone(), token.to_string()), + ] { + let headers_for = |value: &str| { + HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!( + "Basic {}", + base64::engine::general_purpose::STANDARD.encode(value) + )) + .unwrap(), + )]) + }; + let mut headers = headers_for(&input); + broker.inject_request_headers("https://other.example/", &mut headers); + assert_eq!(headers, headers_for(&input)); + broker.inject_request_headers( + "https://api.github.com/public/%2e%2e/v1/models", + &mut headers, + ); + assert_eq!(headers, headers_for(&input)); + broker.inject_request_headers("https://api.github.com/v1/models", &mut headers); + assert_eq!(headers, headers_for(&expected)); + } +} + +#[test] +fn configured_provider_destination_history_is_scoped_to_the_environment() { + let token = "stripe_live_abcdefghijklmnopqrstuvwx"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["STRIPE_API_KEY".to_string()], + patterns: vec!["^stripe_live_[a-z]{24}$".to_string()], + url_prefixes: vec!["static.example".to_string()], + url_prefix_from_env: Some("STRIPE_HOST".to_string()), + ..CredentialProviderConfig::default() + }); + let mut first_env = HashMap::from([ + ("STRIPE_API_KEY".to_string(), token.to_string()), + ("STRIPE_HOST".to_string(), "first.example".to_string()), + ]); + broker.virtualize_child_env_for_environment(&mut first_env, Some("first-environment")); + let first_dummy = first_env["STRIPE_API_KEY"].clone(); + let mut second_env = HashMap::from([ + ("STRIPE_HOST".to_string(), "second.example".to_string()), + ("AUTH_HEADER".to_string(), format!("Bearer {first_dummy}")), + ]); + + broker.virtualize_child_env_for_environment(&mut second_env, Some("second-environment")); + + let second_dummy = second_env["AUTH_HEADER"] + .strip_prefix("Bearer ") + .expect("dummy bearer credential") + .to_string(); + assert_eq!(second_dummy, first_dummy); + assert!(!second_env.contains_key("STRIPE_API_KEY")); + assert_eq!(second_env["AUTH_HEADER"], format!("Bearer {second_dummy}")); + let translated = |environment_id: &str, destination: &str, dummy: &str| { + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")).expect("valid dummy authentication"), + )]); + broker.inject_request_headers_for_environment( + &format!("https://{destination}/"), + &mut headers, + Some(environment_id), + ); + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .expect("authorization header") + .to_string() + }; + assert_eq!( + translated("first-environment", "first.example", &first_dummy), + format!("Bearer {token}") + ); + assert_eq!( + translated("second-environment", "second.example", &second_dummy), + format!("Bearer {token}") + ); + assert_eq!( + translated("first-environment", "static.example", &first_dummy), + format!("Bearer {token}") + ); + + second_env.insert("STRIPE_HOST".to_string(), "third.example".to_string()); + broker.virtualize_child_env_for_environment(&mut second_env, Some("second-environment")); + + assert_eq!(second_env["AUTH_HEADER"], format!("Bearer {second_dummy}")); + assert_eq!( + translated("second-environment", "second.example", &second_dummy), + format!("Bearer {token}") + ); + assert_eq!( + translated("second-environment", "third.example", &second_dummy), + format!("Bearer {token}") + ); + assert_eq!( + translated("first-environment", "first.example", &first_dummy), + format!("Bearer {token}") + ); + assert_eq!( + translated("second-environment", "first.example", &second_dummy), + format!("Bearer {second_dummy}") + ); + for host in ["second.example", "third.example"] { + assert_eq!( + translated("first-environment", host, &first_dummy), + format!("Bearer {first_dummy}") + ); + } + + let revision = broker.config_revision(); + broker.configure(&NetworkProxyConfig { + credential_broker: true, + ..NetworkProxyConfig::default() + }); + assert_eq!(broker.config_revision(), revision + 1); + for host in ["second.example", "third.example", "static.example"] { + assert_eq!( + translated("second-environment", host, &second_dummy), + format!("Bearer {second_dummy}") + ); + assert!( + !broker + .host_protocols_for_environment(host, /*port*/ 443, Some("second-environment"),) + .tls + ); + } +} + +#[test] +fn configured_provider_preserves_filtered_destinations_but_honors_explicit_overrides() { + let token = "stripe_live_abcdefghijklmnopqrstuvwx"; + let second_token = "stripe_live_zyxwvutsrqponmlkjihgfedc"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["STRIPE_API_KEY".to_string()], + patterns: vec!["^stripe_live_[a-z]{24}$".to_string()], + url_prefix_from_env: Some("STRIPE_HOST".to_string()), + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([ + ("STRIPE_API_KEY".to_string(), token.to_string()), + ("STRIPE_HOST".to_string(), "first.example".to_string()), + ("AUTH_HEADER".to_string(), format!("Bearer {token}")), + ( + "SECOND_AUTH_HEADER".to_string(), + format!("Bearer {second_token}"), + ), + ]); + broker.virtualize_child_env_for_environment(&mut env, Some("environment")); + let dummy = env["STRIPE_API_KEY"].clone(); + + let mut filtered_child = env.clone(); + filtered_child.remove("STRIPE_HOST"); + let expected = filtered_child.clone(); + broker.virtualize_child_env_for_environment(&mut filtered_child, Some("filtered-child")); + assert_eq!(filtered_child, expected); + for (alias, real) in [("AUTH_HEADER", token), ("SECOND_AUTH_HEADER", second_token)] { + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&filtered_child[alias]).unwrap(), + )]); + broker.inject_request_headers_for_environment( + "https://first.example/", + &mut headers, + Some("filtered-child"), + ); + assert_eq!(headers[AUTHORIZATION], format!("Bearer {real}")); + } + + let mut child_env = env.clone(); + child_env.insert("STRIPE_HOST".to_string(), "child.example".to_string()); + broker.virtualize_child_env_for_environment(&mut child_env, Some("child")); + child_env.remove("STRIPE_HOST"); + env.remove("STRIPE_HOST"); + for (environment, destination, current_env) in [ + ("environment", "first.example", &mut env), + ("child", "child.example", &mut child_env), + ] { + let expected = current_env.clone(); + broker.virtualize_child_env_for_environment(current_env, Some(environment)); + assert_eq!(*current_env, expected); + for (alias, real) in [("AUTH_HEADER", token), ("SECOND_AUTH_HEADER", second_token)] { + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(¤t_env[alias]).unwrap(), + )]); + broker.inject_request_headers_for_environment( + &format!("https://{destination}/"), + &mut headers, + Some(environment), + ); + assert_eq!(headers[AUTHORIZATION], format!("Bearer {real}")); + } + } + + env.insert("STRIPE_HOST".to_string(), String::new()); + broker.virtualize_child_env_for_environment(&mut env, Some("environment")); + + assert_eq!(env["STRIPE_API_KEY"], token); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {token}")); + assert_eq!(env["SECOND_AUTH_HEADER"], format!("Bearer {second_token}")); + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")).expect("valid dummy authentication"), + )]); + broker.inject_request_headers_for_environment( + "https://first.example/", + &mut headers, + Some("environment"), + ); + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(format!("Bearer {dummy}").as_str()) + ); + assert!( + !broker + .host_protocols_for_environment("first.example", /*port*/ 443, Some("environment")) + .tls + ); +} + +#[test] +fn explicit_invalid_destinations_clear_previous_dynamic_bindings() { + for (key, host_key, token, static_host) in [ + ( + "GH_ENTERPRISE_TOKEN", + "GH_HOST", + "ghp_abcdefghijklmnopqrstuvwxyz0123456789", + None, + ), + ( + "GITHUB_ENTERPRISE_TOKEN", + "GH_HOST", + "ghp_abcdefghijklmnopqrstuvwxyz0123456789", + None, + ), + ( + "OPENAI_API_KEY", + "OPENAI_BASE_URL", + "sk-proj-abcdefghijklmnopqrstuvwxyz0123456789", + Some("api.openai.com"), + ), + ( + "PROVIDER_TOKEN", + "PROVIDER_ENDPOINT", + "provider_abcdefghijklmnopqrstuvwx", + Some("static.example"), + ), + ( + "PROVIDER_TOKEN", + "PROVIDER_ENDPOINT", + "provider_abcdefghijklmnopqrstuvwx", + None, + ), + ] { + for use_real in [false, true] { + for use_alias in [false, true] { + for invalid in [ + "", + " ", + "not a valid destination", + "http://untrusted.example", + ] + .into_iter() + .chain((key == "PROVIDER_TOKEN").then_some("https://*.example")) + .chain((host_key == "GH_HOST").then_some("github.com")) + .chain( + (host_key == "GH_HOST") + .then_some([ + "second.example:bad-port", + "127.0.0.1:bad-port", + "[::1]bad", + "[::1]:bad-port", + "second.example:99999", + "[example.com]", + ]) + .into_iter() + .flatten(), + ) { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["^provider_[a-z]{24}$".to_string()], + url_prefixes: static_host + .map(|host| format!("https://{host}")) + .into_iter() + .collect(), + url_prefix_from_env: Some("PROVIDER_ENDPOINT".to_string()), + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([ + (key.to_string(), token.to_string()), + ( + host_key.to_string(), + if host_key == "GH_HOST" { + "first.example".to_string() + } else { + "https://first.example/v1".to_string() + }, + ), + ]); + broker.virtualize_child_env(&mut env); + let dummy = env[key].clone(); + let second_host = static_host.unwrap_or("second.example"); + env.insert( + host_key.to_string(), + if host_key == "GH_HOST" { + second_host.to_string() + } else { + format!("https://{second_host}/v1") + }, + ); + broker.virtualize_child_env(&mut env); + assert!(broker.host_requires_mitm("first.example", /*port*/ 443)); + env.clear(); + let value = if use_real { token } else { &dummy }; + let (input_key, input_value) = if use_alias { + ("AUTH_HEADER", format!("Bearer {value}")) + } else { + (key, value.to_string()) + }; + env.insert(input_key.to_string(), input_value); + env.insert(host_key.to_string(), invalid.to_string()); + broker.virtualize_child_env(&mut env); + + for host in ["first.example", second_host] { + let expected = if static_host == Some(host) { + token + } else { + &dummy + }; + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {dummy}")).unwrap(), + )]); + broker.inject_request_headers(&format!("https://{host}/v1"), &mut headers); + assert_eq!( + headers[AUTHORIZATION], + format!("Bearer {expected}"), + "{key}, {invalid:?}, {host}, real={use_real}, alias={use_alias}" + ); + } + assert!(!broker.host_requires_mitm("first.example", /*port*/ 443)); + if host_key == "GH_HOST" { + for _ in 0..2 { + broker.virtualize_child_env(&mut env); + assert_eq!( + env[input_key], + if use_alias { + format!("Bearer {token}") + } else { + token.to_string() + }, + "{key}, {invalid:?}, real={use_real}, alias={use_alias}" + ); + for cloud in ["api.github.com", "github.com", "tenant.ghe.com"] { + assert!(!broker.host_requires_mitm(cloud, /*port*/ 443)); + } + } + + env.insert(host_key.to_string(), "restored.example".to_string()); + broker.virtualize_child_env(&mut env); + let restored_dummy = if use_alias { + env[input_key].strip_prefix("Bearer ").unwrap() + } else { + &env[input_key] + }; + assert_ne!(restored_dummy, token); + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {restored_dummy}")).unwrap(), + )]); + broker.inject_request_headers("https://restored.example/", &mut headers); + assert_eq!(headers[AUTHORIZATION], format!("Bearer {token}")); + assert!(!broker.host_requires_mitm("first.example", /*port*/ 443)); + assert!(!broker.host_requires_mitm("second.example", /*port*/ 443)); + assert!(!broker.host_requires_mitm("api.github.com", /*port*/ 443)); + env.clear(); + env.insert("GH_TOKEN".to_string(), token.to_string()); + broker.virtualize_child_env(&mut env); + let mut headers = HeaderMap::from_iter([( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {}", env["GH_TOKEN"])).unwrap(), + )]); + broker.inject_request_headers("https://api.github.com/", &mut headers); + assert_eq!(headers[AUTHORIZATION], format!("Bearer {token}")); + } + } + } + } + } +} + +#[test] +fn configured_provider_limits_injection_to_url_prefixes() { + let token = "provider_abcdefghijklmnopqrstuvwx"; + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["provider_[a-z]{24}".to_string()], + url_prefixes: vec![ + "https://root.provider.example".to_string(), + "https://api.provider.example/v1".to_string(), + "enterprise.example/v2/".to_string(), + "https://*.provider.example:8443/private".to_string(), + "http://localhost:443/v1".to_string(), + "127.0.0.1:443/v1".to_string(), + "http://[::1]:443/v1".to_string(), + "http://localhost:8080/v1".to_string(), + "localhost/v2".to_string(), + "https://localhost:443/v3".to_string(), + "https://localhost:80/v4".to_string(), + ], + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("PROVIDER_TOKEN".to_string(), token.to_string())]); + broker.virtualize_child_env(&mut env); + let dummy = &env["PROVIDER_TOKEN"]; + assert!(broker.host_requires_mitm("team.provider.example", /*port*/ 8443)); + assert!(!broker.host_requires_mitm("team.provider.example", /*port*/ 443)); + + for (destination, injected) in [ + ("root.provider.example", false), + ("https://root.provider.example/models", true), + ("https://api.provider.example/v1", true), + ("https://api.provider.example/v1/models?limit=1", true), + ("https://api.provider.example/v10/models", false), + ("https://api.provider.example/private", false), + ( + "https://api.provider.example/public/%2e%2e/v1/models", + false, + ), + ("https://api.provider.example/v1%2f../private", false), + ("http://api.provider.example/v1", false), + ("https://enterprise.example/v2/models", true), + ("https://enterprise.example/v2", false), + ("https://team.provider.example:8443/private/models", true), + ("https://team.provider.example/private/models", false), + ("https://provider.example:8443/private/models", false), + ("http://localhost:443/v1/models", true), + ("http://localhost/v1/models", false), + ("https://localhost/v1/models", false), + ("http://127.0.0.1:443/v1/models", true), + ("http://127.0.0.1/v1/models", false), + ("http://[::1]:443/v1/models", true), + ("http://[::1]/v1/models", false), + ("http://localhost:8080/v1/models", true), + ("http://localhost/v2/models", true), + ("http://localhost:443/v2/models", false), + ("https://localhost/v3/models", true), + ("http://localhost:443/v3/models", false), + ("https://localhost:80/v4/models", true), + ("https://localhost/v4/models", false), + ] { + let dummy_header = format!("Bearer {dummy}"); + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&dummy_header).expect("valid dummy authentication"), + ); + + broker.inject_request_headers(destination, &mut headers); + + let expected = if injected { + format!("Bearer {token}") + } else { + dummy_header + }; + assert_eq!( + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()), + Some(expected.as_str()), + "destination: {destination}" + ); + } +} + +#[test] +fn configured_provider_accepts_hostname_or_https_url_from_one_environment_key() { + let token = "provider_abcdefghijklmnopqrstuvwx"; + for (host_value, expected_host) in [ + ("enterprise.example", Some("enterprise.example")), + ("https://gateway.example/v1", Some("gateway.example")), + ("http://plaintext.example/v1", None), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["provider_[a-z]{24}".to_string()], + url_prefix_from_env: Some("PROVIDER_HOST".to_string()), + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([ + ("PROVIDER_TOKEN".to_string(), token.to_string()), + ("PROVIDER_HOST".to_string(), host_value.to_string()), + ]); + + broker.virtualize_child_env(&mut env); + + match expected_host { + Some(host) => { + assert_ne!(env["PROVIDER_TOKEN"], token); + assert!(broker.host_requires_mitm(host, /*port*/ 443)); + assert_eq!( + broker.environment(&env).binding_keys, + vec!["PROVIDER_HOST".to_string()] + ); + assert_eq!( + broker.environment(&env).configured_provider_context_keys, + vec!["PROVIDER_HOST".to_string()] + ); + } + None => { + assert_eq!(env["PROVIDER_TOKEN"], token); + let mut snapshot = token.to_string(); + assert!(broker.virtualize_text(&mut snapshot, &env)); + assert_eq!(snapshot, token); + } + } + } +} + +#[test] +fn configured_provider_rejects_unsafe_destinations_and_unusable_dummy_patterns() { + let impossible_assertion = CredentialProviderConfig { + env: vec!["PROVIDER_PASSWORD".to_string()], + patterns: vec![r"^pass_\b[a-z]{24}$".to_string()], + url_prefixes: vec!["api.example".to_string()], + ..CredentialProviderConfig::default() + }; + assert!( + super::configured::ConfiguredCredentialProvider::compile("custom", &impossible_assertion) + .is_err() + ); + + for (pattern, url_prefixes, first, second) in [ + ( + r"\b|token_[a-z]{24}", + vec!["api.example"], + "token_abcdefghijklmnopqrstuvwx", + None, + ), + ( + "provider_[a-z]{24}", + vec!["*"], + "provider_abcdefghijklmnopqrstuvwx", + None, + ), + ( + "token_[01]", + vec!["api.example"], + "token_0", + Some("token_1"), + ), + ( + "only_one_token", + vec!["api.example"], + "only_one_token", + None, + ), + ] { + let broker = broker_for(CredentialProviderConfig { + env: vec!["FIRST_TOKEN".to_string(), "SECOND_TOKEN".to_string()], + patterns: vec![pattern.to_string()], + url_prefixes: url_prefixes.into_iter().map(str::to_string).collect(), + ..CredentialProviderConfig::default() + }); + let mut env = HashMap::from([("FIRST_TOKEN".to_string(), first.to_string())]); + if let Some(second) = second { + env.insert("SECOND_TOKEN".to_string(), second.to_string()); + } + + broker.virtualize_child_env(&mut env); + + assert_eq!(env["FIRST_TOKEN"], first); + if let Some(second) = second { + assert_eq!(env["SECOND_TOKEN"], second); + } + assert!(broker.environment(&env).credential_keys.is_empty()); + let mut snapshot = first.to_string(); + assert!(broker.virtualize_text(&mut snapshot, &env)); + assert_eq!(snapshot, first); + } +} diff --git a/codex-rs/network-proxy/src/credential_broker/destination.rs b/codex-rs/network-proxy/src/credential_broker/destination.rs new file mode 100644 index 0000000000000000000000000000000000000000..116ed8d083e178ae3446cf28a6557304700b3eb9 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/destination.rs @@ -0,0 +1,162 @@ +use crate::authorization_path::is_safe_for_authorization; +use crate::policy::Host; +use crate::policy::is_loopback_host; +use crate::policy::normalize_host; +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use anyhow::ensure; +use url::Url; + +#[derive(Clone, PartialEq, Eq)] +pub(super) struct CredentialDestination { + scheme: CredentialDestinationScheme, + host: String, + wildcard: bool, + port: u16, + path_prefix: String, +} + +#[derive(Clone, Copy, PartialEq, Eq)] +enum CredentialDestinationScheme { + Http, + Https, +} + +impl CredentialDestinationScheme { + fn as_str(self) -> &'static str { + match self { + Self::Http => "http", + Self::Https => "https", + } + } + + fn default_port(self) -> u16 { + match self { + Self::Http => 80, + Self::Https => 443, + } + } +} + +impl CredentialDestination { + pub(super) fn parse(value: &str) -> Result { + let value = value.trim(); + let (explicit_scheme, authority) = value + .split_once("://") + .map_or((None, value), |(scheme, authority)| { + (Some(scheme), authority) + }); + let (wildcard, authority) = authority + .strip_prefix("*.") + .map_or((false, authority), |suffix| (true, suffix)); + let authority = if wildcard { + format!("wildcard.{authority}") + } else { + authority.to_string() + }; + let url = Url::parse(&format!("https://{authority}"))?; + ensure!( + url.username().is_empty() + && url.password().is_none() + && url.query().is_none() + && url.fragment().is_none(), + "credential destination cannot include user information, a query, or a fragment" + ); + let host = url + .host_str() + .context("credential destination has no hostname")?; + let host = if wildcard { + host.strip_prefix("wildcard.") + .context("credential destination has an invalid wildcard")? + } else { + host + }; + let host = normalize_host(host); + ensure!( + !host.is_empty() && !host.contains('*') && !host.contains('/'), + "credential destination must have an exact hostname or scoped wildcard" + ); + let loopback = !wildcard && is_loopback_host(&Host::parse(&host)?); + let scheme = match explicit_scheme { + None if loopback => CredentialDestinationScheme::Http, + None => CredentialDestinationScheme::Https, + Some(scheme) if scheme.eq_ignore_ascii_case("https") => { + CredentialDestinationScheme::Https + } + Some(scheme) if scheme.eq_ignore_ascii_case("http") && loopback => { + CredentialDestinationScheme::Http + } + Some(_) => { + bail!("credential destination must use HTTPS unless it targets loopback over HTTP") + } + }; + // HTTPS parsing removes an explicit :443, which HTTP must retain. + let url = match scheme { + CredentialDestinationScheme::Http => Url::parse(&format!("http://{authority}"))?, + CredentialDestinationScheme::Https => url, + }; + + Ok(Self { + scheme, + host, + wildcard, + port: url.port().unwrap_or_else(|| scheme.default_port()), + path_prefix: url.path().to_string(), + }) + } + + pub(super) fn is_wildcard(&self) -> bool { + self.wildcard + } + + pub(super) fn bypasses_proxy_with_local_binding(&self) -> bool { + (!self.wildcard || self.host.parse::().is_err()) + && crate::policy::is_default_proxy_bypass_host(&self.host) + } + + pub(super) fn matches_host(&self, host: &str, port: u16) -> bool { + if self.port != port { + return false; + } + if self.wildcard { + host.strip_suffix(&self.host) + .is_some_and(|prefix| prefix.ends_with('.')) + } else { + host == self.host + } + } + + pub(super) fn requires_mitm(&self, host: &str, port: u16) -> bool { + self.scheme == CredentialDestinationScheme::Https && self.matches_host(host, port) + } + + pub(super) fn requires_http_interception(&self, host: &str, port: u16) -> bool { + self.scheme == CredentialDestinationScheme::Http && self.matches_host(host, port) + } + + pub(super) fn matches_request(&self, host: &str, request: Option<&Url>) -> bool { + let Some(request) = request else { + return false; + }; + if request.scheme() != self.scheme.as_str() + || !self.matches_host( + host, + request + .port_or_known_default() + .unwrap_or_else(|| self.scheme.default_port()), + ) + { + return false; + } + let path = request.path(); + if !is_safe_for_authorization(path) { + return false; + } + self.path_prefix == "/" + || path == self.path_prefix + || path + .strip_prefix(&self.path_prefix) + .is_some_and(|suffix| self.path_prefix.ends_with('/') || suffix.starts_with('/')) + } +} diff --git a/codex-rs/network-proxy/src/credential_broker/environment.rs b/codex-rs/network-proxy/src/credential_broker/environment.rs new file mode 100644 index 0000000000000000000000000000000000000000..4455e217b00a00fee4cc76c2561d7b81c55e2b2f --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/environment.rs @@ -0,0 +1,470 @@ +use super::BROKERED_CREDENTIAL_ALIAS_MARKER_PREFIX; +use super::BROKERED_CREDENTIALS_ENV_KEY; +use super::CREDENTIAL_BROKER_ACTIVE_ENV_KEY; +use super::CredentialBrokerState; +use super::MIN_EMBEDDED_CREDENTIAL_LENGTH; +use super::configured::valid_environment_key; +use super::env_entry; +use super::env_key_matches; +use super::env_value; +use super::matching; +use super::providers; +use super::registry::BrokeredCredentialProvider; +use super::remove_env_value; +use super::set_env_value; +use crate::NetworkProxy; +use crate::NetworkProxyConfig; +use crate::PreparedManagedNetwork; +use std::borrow::Cow; +use std::collections::HashMap; + +/// Local destination hints retained by the broker, not added to child environments. +#[derive(Clone, Default, Eq, PartialEq)] +pub struct CredentialBrokerContext(HashMap); + +impl From> for CredentialBrokerContext { + fn from(env: HashMap) -> Self { + Self(env) + } +} + +impl std::fmt::Debug for CredentialBrokerContext { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("CredentialBrokerContext()") + } +} + +impl CredentialBrokerContext { + /// Uses private destination hints during brokerage and returns an environment ready to spawn. + /// Context values override routing inputs, but never change the child's visible values. + pub fn prepare_child_environment( + self, + proxy: &NetworkProxy, + mut env: HashMap, + environment_id: Option<&str>, + ) -> anyhow::Result { + let previous = self + .0 + .into_iter() + .map(|(key, value)| { + let previous = + env_entry(&env, &key).map(|(key, value)| (key.to_string(), value.to_string())); + set_env_value(&mut env, &key, value); + (key, previous) + }) + .collect::>(); + let mut prepared = proxy.prepare_for_optional_environment(env, environment_id)?; + for (key, previous) in previous { + remove_env_value(&mut prepared.env, &key); + if let Some((key, value)) = previous { + prepared.env.insert(key, value); + } + } + Ok(prepared) + } + + pub(crate) fn capture( + config: &NetworkProxyConfig, + overrides: &HashMap, + ) -> Self { + Self( + credential_broker_provider_context_env_keys() + .map(str::to_string) + .chain( + config + .credential_providers + .values() + .filter_map(|provider| provider.url_prefix_from_env.clone()), + ) + .filter_map(|key| { + env_value(overrides, &key) + .map(str::to_string) + .or_else(|| std::env::var(&key).ok()) + .map(|value| (key, value)) + }) + .collect(), + ) + } + + /// Resolves destination hints without applying or changing child environment policy. + pub fn with_fallbacks<'a>( + &self, + env: &'a HashMap, + ) -> Cow<'a, HashMap> { + let mut context = Cow::Borrowed(env); + for (key, value) in &self.0 { + // A present value, including an empty one, overrides the trusted fallback. + if env_value(&context, key).is_none() { + set_env_value(context.to_mut(), key, value.clone()); + } + } + context + } +} + +/// Trusted provider metadata and active credential bindings for one child environment. +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct CredentialBrokerEnvironment { + pub credential_keys: Vec, + pub binding_keys: Vec, + pub context_keys: Vec, + pub provider_keys: Vec, + pub provider_context_keys: Vec, + pub configured_provider_context_keys: Vec, +} + +impl CredentialBrokerState { + pub(super) fn environment(&self, env: &HashMap) -> CredentialBrokerEnvironment { + let mut metadata = CredentialBrokerEnvironment::default(); + for key in + providers::credential_env_keys().chain(credential_broker_provider_context_env_keys()) + { + push_unique_key(&mut metadata.provider_keys, key); + } + for key in credential_broker_provider_context_env_keys() { + push_unique_key(&mut metadata.provider_context_keys, key); + } + for provider in &self.configured_providers { + for key in &provider.config.env { + push_unique_key(&mut metadata.provider_keys, key); + } + if let Some(key) = provider.config.url_prefix_from_env.as_deref() { + push_unique_key(&mut metadata.provider_keys, key); + push_unique_key(&mut metadata.provider_context_keys, key); + push_unique_key(&mut metadata.configured_provider_context_keys, key); + } + } + + if !self.enabled { + return metadata; + } + + metadata.credential_keys = brokered_credential_dummy_env_keys(env); + let marked_credentials = env_value(env, BROKERED_CREDENTIALS_ENV_KEY) + .and_then(|marker| serde_json::from_str::>(marker).ok()) + .unwrap_or_default(); + for credential in &self.credentials { + if matches!( + credential.provider, + BrokeredCredentialProvider::Configured(_) + ) && env_value(env, &credential.env_var) == Some(credential.dummy_value.as_str()) + && marked_credentials.iter().any(|(key, value)| { + env_key_matches(key, &credential.env_var) && value == &credential.dummy_value + }) + { + push_unique_key(&mut metadata.credential_keys, &credential.env_var); + } + } + + for key in providers::credential_context_env_keys(&metadata.credential_keys) { + if env_value(env, key).is_some() { + push_unique_key(&mut metadata.context_keys, key); + } + } + for key in providers::credential_binding_env_keys(&metadata.credential_keys) { + if env_value(env, key).is_some() { + push_unique_key(&mut metadata.binding_keys, key); + } + } + for provider in &self.configured_providers { + if provider.config.env.iter().any(|provider_key| { + metadata + .credential_keys + .iter() + .any(|credential_key| env_key_matches(provider_key, credential_key)) + }) && let Some(key) = provider.config.url_prefix_from_env.as_deref() + && env_value(env, key).is_some() + { + push_unique_key(&mut metadata.context_keys, key); + push_unique_key(&mut metadata.binding_keys, key); + } + } + + metadata + } +} + +fn push_unique_key(keys: &mut Vec, key: &str) { + if !keys.iter().any(|candidate| env_key_matches(candidate, key)) { + keys.push(key.to_string()); + } +} + +pub(super) fn update_brokered_credentials_marker( + state: &CredentialBrokerState, + env: &mut HashMap, +) { + let credentials = state + .credentials + .iter() + .filter(|credential| { + env.iter().any(|(key, value)| { + !env_key_matches(key, CREDENTIAL_BROKER_ACTIVE_ENV_KEY) + && !env_key_matches(key, BROKERED_CREDENTIALS_ENV_KEY) + && (value == &credential.dummy_value + || credential.contains_embedded_value(value, &credential.dummy_value)) + }) + }) + .collect::>(); + let mut brokered = credentials + .iter() + .map(|credential| (credential.env_var.clone(), credential.dummy_value.clone())) + .collect::>(); + brokered.extend(state.credential_aliases.iter().flat_map(|alias| { + credentials + .iter() + .filter(|credential| { + alias.dummy_value == credential.dummy_value + || credential + .contains_embedded_value(&alias.dummy_value, &credential.dummy_value) + }) + .map(|credential| { + ( + format!("{BROKERED_CREDENTIAL_ALIAS_MARKER_PREFIX}{}", alias.env_var), + credential.dummy_value.clone(), + ) + }) + })); + brokered.sort_unstable(); + brokered.dedup(); + match serde_json::to_string(&brokered) { + Ok(marker) => { + set_env_value(env, BROKERED_CREDENTIALS_ENV_KEY, marker); + } + Err(_) => { + remove_env_value(env, BROKERED_CREDENTIALS_ENV_KEY); + } + } +} + +/// Returns supported environment keys whose current values still match the child-scoped dummy +/// values recorded by the credential broker. +/// +/// The broker marker is treated as untrusted: malformed metadata, unsupported keys, and values +/// replaced by the user are ignored. The environment is not mutated; callers own the decision to +/// remove the returned keys. +pub fn brokered_credential_dummy_env_keys(env: &HashMap) -> Vec { + let mut keys = marked_credential_dummy_env_keys(env) + .into_iter() + .filter(|key| { + providers::credential_env_keys().any(|candidate| env_key_matches(key, candidate)) + }) + .collect::>(); + let context_bound = |key: &str| { + providers::credential_providers().any(|provider| { + provider.sources().iter().any(|source| { + source + .env_vars + .iter() + .any(|candidate| env_key_matches(key, candidate)) + && source + .binding_env_vars + .iter() + .any(|context| env_value(env, context).is_some()) + }) + }) + }; + keys.sort_unstable_by(|left, right| { + context_bound(right) + .cmp(&context_bound(left)) + .then_with(|| left.cmp(right)) + }); + keys +} + +pub(crate) fn marked_credential_dummy_env_keys(env: &HashMap) -> Vec { + env_value(env, BROKERED_CREDENTIALS_ENV_KEY) + .and_then(|marker| serde_json::from_str::>(marker).ok()) + .unwrap_or_default() + .into_iter() + .filter_map(|(key, dummy_value)| { + let (actual_key, actual_value) = env_entry(env, &key)?; + (actual_value == dummy_value.as_str()).then(|| actual_key.to_string()) + }) + .collect() +} + +/// Returns canonical credential and known alias keys recorded for an active brokered child, +/// including configured provider keys and keys that are currently absent. +pub fn brokered_credential_marker_env_keys(env: &HashMap) -> Vec { + if env_value(env, CREDENTIAL_BROKER_ACTIVE_ENV_KEY) != Some("1") { + return Vec::new(); + } + let entries = env_value(env, BROKERED_CREDENTIALS_ENV_KEY) + .and_then(|marker| serde_json::from_str::>(marker).ok()) + .unwrap_or_default(); + let dummy_values = entries + .iter() + .filter(|(key, value)| valid_environment_key(key) && !value.is_empty()) + .map(|(_, dummy_value)| dummy_value.clone()) + .collect::>(); + let mut keys = entries + .into_iter() + .filter_map(|(key, value)| { + if valid_environment_key(&key) { + return (!value.is_empty()).then_some(key); + } + let alias_key = key.strip_prefix(BROKERED_CREDENTIAL_ALIAS_MARKER_PREFIX)?; + let known_alias = dummy_values + .iter() + .any(|dummy_value| value.contains(dummy_value)); + (known_alias && valid_environment_key(alias_key)).then(|| alias_key.to_string()) + }) + .collect::>(); + keys.sort_unstable(); + keys.dedup(); + keys +} + +/// Returns environment keys whose current values are brokered dummies or aliases containing them, +/// including aliases retained after their canonical source variable was removed. +pub fn brokered_credential_value_env_keys(env: &HashMap) -> Vec { + if env_value(env, CREDENTIAL_BROKER_ACTIVE_ENV_KEY) != Some("1") { + return Vec::new(); + } + let marked_values = env_value(env, BROKERED_CREDENTIALS_ENV_KEY) + .and_then(|marker| serde_json::from_str::>(marker).ok()) + .unwrap_or_default() + .into_iter() + .filter(|(key, dummy_value)| { + !dummy_value.is_empty() + && (valid_environment_key(key) + || key + .strip_prefix(BROKERED_CREDENTIAL_ALIAS_MARKER_PREFIX) + .is_some_and(valid_environment_key)) + }) + .collect::>(); + let mut keys = env + .iter() + .filter(|(key, _)| { + !env_key_matches(key, CREDENTIAL_BROKER_ACTIVE_ENV_KEY) + && !env_key_matches(key, BROKERED_CREDENTIALS_ENV_KEY) + }) + .filter(|(key, value)| { + marked_values.iter().any(|(marked_key, dummy)| { + value.as_str() == dummy.as_str() + || (dummy.len() >= MIN_EMBEDDED_CREDENTIAL_LENGTH + || marked_key + .strip_prefix(BROKERED_CREDENTIAL_ALIAS_MARKER_PREFIX) + .is_some_and(|alias_key| env_key_matches(key, alias_key))) + && value.contains(dummy.as_str()) + }) + }) + .map(|(key, _)| key.clone()) + .collect::>(); + keys.sort_unstable(); + keys +} + +/// Returns environment keys used to bind currently brokered credentials to hosts. +pub fn brokered_credential_binding_env_keys( + env: &HashMap, +) -> impl Iterator { + let brokered_keys = brokered_credential_dummy_env_keys(env); + providers::credential_binding_env_keys(&brokered_keys) + .filter(|key| env_value(env, key).is_some()) + .collect::>() + .into_iter() +} + +/// Returns environment keys used to bind registered credential providers to their destinations. +pub fn credential_broker_provider_context_env_keys() -> impl Iterator { + providers::credential_providers().flat_map(|provider| provider.context_env_vars.iter().copied()) +} + +/// Checks whether each credential in a value has an allowed source environment key. +pub fn credential_broker_provider_sources_allowed( + value: &str, + virtualized: &str, + source_env: &HashMap, + is_allowed: impl Fn(&str) -> bool, +) -> bool { + let mut recognized = false; + let allowed = providers::credential_providers() + .filter(move |provider| { + provider.credential_prefixes.iter().any(|prefix| { + value.match_indices(*prefix).any(|(start, _)| { + matching::recognized_credential_match(provider, value, virtualized, start) + .is_some() + }) + }) + }) + .all(|provider| { + recognized = true; + let actual_sources = provider + .sources() + .iter() + .flat_map(|source| source.env_vars.iter().copied()) + .filter(|source| { + env_value(source_env, source).is_some_and(|source_value| { + source_value.len() >= provider.minimum_credential_len + && value.contains(source_value) + }) + }) + .collect::>(); + let unattributed = provider.credential_prefixes.iter().any(|prefix| { + value.match_indices(*prefix).any(|(start, _)| { + matching::recognized_credential_match(provider, value, virtualized, start) + .is_some_and(|credential| { + !actual_sources.iter().any(|source| { + env_value(source_env, source).is_some_and(|source_value| { + credential == source_value + || credential + .strip_prefix(source_value) + .is_some_and(|suffix| suffix.starts_with(['_', '-'])) + }) + }) + }) + }) + }); + if actual_sources.is_empty() || unattributed { + provider + .sources() + .iter() + .flat_map(|source| source.env_vars.iter().copied()) + .all(&is_allowed) + } else { + actual_sources.iter().all(|source| { + actual_sources.iter().any(|equivalent| { + env_value(source_env, source) == env_value(source_env, equivalent) + && is_allowed(equivalent) + }) + }) + } + }); + allowed && recognized +} + +/// Returns whether an environment key belongs to a supported credential provider. +pub fn is_credential_broker_provider_env_key(key: &str) -> bool { + providers::credential_providers().any(|provider| { + provider + .sources() + .iter() + .flat_map(|source| source.env_vars.iter().copied()) + .chain(provider.context_env_vars.iter().copied()) + .any(|candidate| env_key_matches(key, candidate)) + }) +} + +/// Returns credential keys plus provider context keys already present in an environment with an +/// active broker. +pub fn brokered_credential_env_keys( + env: &HashMap, +) -> impl Iterator { + let active = env_value(env, CREDENTIAL_BROKER_ACTIVE_ENV_KEY).is_some_and(|value| value == "1"); + let mut keys = Vec::new(); + if active { + let brokered_keys = brokered_credential_dummy_env_keys(env); + keys.extend(providers::credential_env_keys().filter(|key| { + brokered_keys + .iter() + .any(|brokered_key| env_key_matches(brokered_key, key)) + })); + keys.extend( + providers::credential_context_env_keys(&brokered_keys) + .filter(|key| env_value(env, key).is_some()), + ); + } + keys.into_iter() +} diff --git a/codex-rs/network-proxy/src/credential_broker/matching.rs b/codex-rs/network-proxy/src/credential_broker/matching.rs new file mode 100644 index 0000000000000000000000000000000000000000..7e205010244be39a1c963f4c5e1daa0063b0c89e --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/matching.rs @@ -0,0 +1,541 @@ +use super::CredentialBrokerState; +use super::MIN_EMBEDDED_CREDENTIAL_LENGTH; +use super::env_key_matches; +use super::env_value; +use super::prioritized_credentials; +use super::providers; +use super::registry::BrokeredCredentialProvider; +use super::registry::is_builtin_shaped_credential; +use super::replacement::Replacements; +use std::collections::HashMap; +use std::path::Path; +use url::Position; +use url::Url; + +pub(super) struct KnownCredentialMatch<'a> { + pub(super) range: std::ops::Range, + pub(super) env_var: &'a str, + pub(super) provider: BrokeredCredentialProvider, + pub(super) real_value: &'a str, + pub(super) value: &'a str, +} + +pub(super) fn known_credential_matches<'a>( + state: &'a CredentialBrokerState, + text: &str, + source_env: &'a HashMap, +) -> Vec> { + let mut matches = Vec::new(); + let mut add = |env_var: &'a str, + provider: &BrokeredCredentialProvider, + real_value: &'a str, + value: &'a str| { + if value.is_empty() { + return; + } + let ranges = if text == value { + std::iter::once(0..value.len()).collect() + } else if value.len() >= MIN_EMBEDDED_CREDENTIAL_LENGTH + && (is_builtin_shaped_credential(real_value) + || matches!(provider, BrokeredCredentialProvider::Builtin(_))) + { + text.match_indices(value) + .map(|(start, _)| start..start + value.len()) + .collect() + } else if let BrokeredCredentialProvider::Configured(provider) = provider { + // Configured short values retain their pattern boundaries and path exclusions. + provider + .credential_value_match_ranges(text, value) + .into_iter() + .filter(|range| { + value.len() >= MIN_EMBEDDED_CREDENTIAL_LENGTH + || !is_operational_path_match(text, range.start, range.end) + // A complete longer configured token still needs its own source permission. + && !provider.find_discoverable_credentials(text).any(|candidate| { + candidate.range.start <= range.start && candidate.range.end > range.end + }) + }) + .collect() + } else { + Vec::new() + }; + matches.extend(ranges.into_iter().map(|range| KnownCredentialMatch { + range, + env_var, + provider: provider.clone(), + real_value, + value, + })); + }; + for credential in &state.credentials { + for value in [&credential.real_value, &credential.dummy_value] { + add( + &credential.env_var, + &credential.provider, + &credential.real_value, + value, + ); + } + } + for provider in super::credential_provider_definitions(state) { + let owns_key = |key: &str| match &provider { + BrokeredCredentialProvider::Builtin(provider) => { + provider.sources().iter().any(|source| { + source + .env_vars + .iter() + .any(|source| env_key_matches(source, key)) + }) + } + BrokeredCredentialProvider::Configured(provider) => provider + .config + .env + .iter() + .any(|source| env_key_matches(source, key)), + }; + for owner in &state.credential_owners { + if owns_key(&owner.env_var) { + add( + &owner.env_var, + &provider, + &owner.real_value, + &owner.real_value, + ); + } + } + for key in source_env.keys() { + if owns_key(key) + && let Some(real) = + super::brokerable_credential_value(source_env, state, key, &provider) + { + add(key, &provider, real, real); + } + } + } + for credential in &state.credentials { + matches.extend( + credential + .generated_dummy_ranges(text) + .into_iter() + .map(|range| KnownCredentialMatch { + range, + env_var: &credential.env_var, + provider: credential.provider.clone(), + real_value: &credential.real_value, + value: &credential.dummy_value, + }), + ); + } + // Longest exact identities win over overlapping shorter identities, as in replacement. + matches.sort_unstable_by_key(|matched| { + (std::cmp::Reverse(matched.range.len()), matched.range.start) + }); + let mut selected = Vec::>::new(); + for matched in matches { + if selected.iter().any(|known| { + known.range.start < matched.range.end + && matched.range.start < known.range.end + && (known.range != matched.range + || known.provider.same_provider(&matched.provider) + && env_key_matches(known.env_var, matched.env_var) + && known.real_value == matched.real_value) + }) { + continue; + } + selected.push(matched); + } + selected +} + +pub(super) fn mask_known_credentials(text: &str, matches: &[KnownCredentialMatch<'_>]) -> String { + let mut uncovered = text.to_string(); + for matched in matches { + uncovered.replace_range(matched.range.clone(), &"\0".repeat(matched.range.len())); + } + uncovered +} + +pub(super) fn virtualize_text( + state: &CredentialBrokerState, + output: &mut String, + env: &HashMap, +) -> bool { + if !state.enabled { + return true; + } + + let allowed_keys = state.environment(env).credential_keys; + let credentials = prioritized_credentials(state, env); + let text = output.as_str(); + let source_env = HashMap::new(); + let known = known_credential_matches(state, text, &source_env); + let mut replacements = Replacements::default(); + let binding_env = state.context.with_fallbacks(env); + let mut allowed = true; + for credential in &credentials { + let ranges = known + .iter() + .filter(|matched| { + matched.provider.same_provider(&credential.provider) + && env_key_matches(matched.env_var, &credential.env_var) + && matched.real_value == credential.real_value + && (matched.value == credential.real_value + || matched.value == credential.dummy_value) + }) + .map(|matched| matched.range.clone()) + .collect::>(); + if ranges.is_empty() { + continue; + } + + let replacement = credentials.iter().copied().find(|candidate| { + candidate.provider.same_provider(&credential.provider) + && candidate.real_value == credential.real_value + && (allowed_keys.iter().any(|key| { + env_key_matches(key, &candidate.env_var) + && env_value(env, key) == Some(candidate.dummy_value.as_str()) + }) || state.credential_aliases.iter().any(|alias| { + env_value(env, &alias.env_var) == Some(alias.dummy_value.as_str()) + && (alias.dummy_value == candidate.dummy_value + || candidate.contains_embedded_value( + &alias.dummy_value, + &candidate.dummy_value, + )) + && !state.credentials.iter().any(|other| { + other.dummy_value != candidate.dummy_value + && (alias.dummy_value == other.dummy_value + || other.contains_embedded_value( + &alias.dummy_value, + &other.dummy_value, + )) + }) + })) + }); + if replacement.is_none() { + allowed = false; + } + replacements.add( + ranges, + replacement.map_or("", |candidate| candidate.dummy_value.as_str()), + replacement, + ); + } + + let original = text; + let mut uncovered = replacements.masked(original); + let text = &mut uncovered; + // Startup can copy a supported credential before unsetting its source variable. + for provider in providers::credential_providers() { + for prefix in provider.credential_prefixes { + let mut offset = 0; + while let Some(position) = text[offset..].find(prefix) { + let start = offset + position; + if let Some(length) = state + .credentials + .iter() + .flat_map(|credential| { + [&credential.real_value, &credential.dummy_value] + .map(|value| (credential, value)) + }) + .filter(|(_, value)| text[start..].starts_with(value.as_str())) + .filter(|(credential, value)| { + value.len() >= MIN_EMBEDDED_CREDENTIAL_LENGTH + || matches!(&credential.provider, BrokeredCredentialProvider::Configured(provider) + if provider.credential_value_match_ranges(text, value) + .iter().any(|range| range.start == start + && !is_operational_path_match(text, range.start, range.end))) + }) + .map(|(_, value)| value.len()) + .max() + { + offset = start + length; + continue; + } + let candidate = builtin_credential_candidate(provider, text, start); + let length = candidate.len(); + if provider + .ignored_credential_prefixes + .iter() + .any(|prefix| candidate.starts_with(prefix)) + && provider + .credential_watermark + .is_none_or(|watermark| !candidate.contains(watermark)) + { + let known_length = env + .values() + .filter(|known| { + known.len() >= provider.minimum_credential_len + && candidate + .strip_prefix(known.as_str()) + .is_some_and(|suffix| { + provider.credential_prefixes.iter().any(|prefix| { + suffix.match_indices(prefix).any(|(offset, _)| { + suffix.len() - offset + >= provider.minimum_credential_len + }) + }) + }) + }) + .map(String::len) + .min(); + let embedded_supported = provider + .credential_prefixes + .iter() + .flat_map(|prefix| candidate.match_indices(prefix)) + .filter_map(|(offset, _)| { + (offset > 0 + && candidate.len() - offset >= provider.minimum_credential_len + && !provider + .ignored_credential_prefixes + .iter() + .any(|ignored| candidate[offset..].starts_with(ignored))) + .then_some(offset) + }) + .min(); + offset = start + + known_length + .into_iter() + .chain(embedded_supported) + .min() + .unwrap_or(length); + continue; + } + let end = start + length; + let credential = &text[start..end]; + let ignored_credential_match = + ignored_credential_match(provider, text, start, credential); + if length >= provider.minimum_credential_len && !ignored_credential_match { + replacements.add(std::iter::once(start..end), "", /*dummy*/ None); + text.replace_range(start..end, &"\0".repeat(length)); + offset = end; + allowed = false; + } else { + offset = start + + if ignored_credential_match { + prefix.len() + } else { + length + }; + } + } + } + } + + for provider in &state.configured_providers { + if provider.host_binding(&binding_env).is_none() { + continue; + } + let mut matches = provider + .find_credential_matches(text) + .map(|matched| (matched.start, matched.end)) + .collect::>(); + matches.sort_unstable_by_key(|(start, end)| (std::cmp::Reverse(end - start), *start)); + matches.dedup(); + let mut removed_ranges = Vec::<(usize, usize)>::new(); + for (start, end) in matches { + let credential = &text[start..end]; + if credential.contains('\0') + || !provider + .credential_value_match_ranges(original, credential) + .contains(&(start..end)) + || is_operational_path_match(original, start, end) + || state + .credentials + .iter() + .flat_map(|known| [&known.real_value, &known.dummy_value]) + .flat_map(|known| { + text.match_indices(known) + .map(move |(index, _)| (index, known)) + }) + .any(|(index, known)| index <= start && end <= index + known.len()) + || provider + .config + .env + .iter() + .any(|key| env_value(env, key).is_some_and(|value| value == credential)) + || removed_ranges.iter().any(|(removed_start, removed_end)| { + start < *removed_end && *removed_start < end + }) + { + continue; + } + removed_ranges.push((start, end)); + allowed = false; + } + removed_ranges.sort_unstable_by_key(|(start, _)| std::cmp::Reverse(*start)); + for (start, end) in removed_ranges { + replacements.add(std::iter::once(start..end), "", /*dummy*/ None); + text.replace_range(start..end, &"\0".repeat(end - start)); + } + } + + replacements.render(output); + allowed +} + +pub(super) fn is_operational_path_match(text: &str, start: usize, end: usize) -> bool { + let is_value_boundary = |character: char| { + character.is_ascii_whitespace() || matches!(character, '"' | '\'' | '=' | '`') + }; + let value_start = text[..start] + .rfind(is_value_boundary) + .map_or(0, |index| index + 1); + let value_end = text[end..] + .find(is_value_boundary) + .map_or(text.len(), |index| end + index); + let value = &text[value_start..value_end]; + if let Ok(url) = Url::parse(value) + && url.has_host() + { + let relative_start = start - value_start; + let relative_end = end - value_start; + return relative_start >= url[..Position::BeforeHost].len() + && relative_end <= url[..Position::AfterPort].len(); + } + + let relative_start = start - value_start; + let relative_end = end - value_start; + if !value[..relative_start].contains(['/', '\\']) + && !value[relative_end..].contains(['/', '\\']) + { + return false; + } + + let path = Path::new(value); + path.has_root() || path.components().count() > 1 || value.contains('\\') +} + +fn ignored_credential_match( + provider: &providers::CredentialProvider, + text: &str, + start: usize, + credential: &str, +) -> bool { + provider + .ignored_credential_prefixes + .iter() + .any(|prefix| credential.starts_with(prefix)) + && provider + .credential_watermark + .is_none_or(|watermark| !credential.contains(watermark)) + || credential.starts_with("sk-") + && provider + .credential_watermark + .is_none_or(|watermark| !credential.contains(watermark)) + && credential[3..].split(['-', '_']).any(|segment| { + segment.len() == 64 + && segment.bytes().all(|byte| byte.is_ascii_hexdigit()) + && credential + .find(segment) + .is_some_and(|offset| offset < provider.minimum_credential_len) + }) + && text[..start] + .rsplit(|character: char| { + character.is_ascii() && !character.is_ascii_alphanumeric() + }) + .next() + .is_some_and(|word| { + ((1..=3).contains(&word.len()) + || word + .get(word.len().saturating_sub(2)..) + .is_some_and(|suffix| { + suffix.eq_ignore_ascii_case("di") + || suffix.eq_ignore_ascii_case("ta") + && !word.eq_ignore_ascii_case("data") + }) + || credential.strip_prefix("sk-").is_some_and(|hash| { + hash.len() == 64 && hash.bytes().all(|byte| byte.is_ascii_hexdigit()) + })) + && word.bytes().all(|byte| byte.is_ascii_alphabetic()) + && text[..start] + .rsplit_once(['/', '\\']) + .is_some_and(|(_, component)| { + component.chars().all(|character| { + !character.is_ascii() + || character.is_ascii_alphanumeric() + || matches!(character, '_' | '-') + }) + }) + }) +} + +pub(super) fn recognized_credential_match<'a>( + provider: &providers::CredentialProvider, + value: &'a str, + virtualized: &str, + start: usize, +) -> Option<&'a str> { + let credential = builtin_credential_candidate(provider, value, start); + let length = credential.len(); + let enclosing_start = value.as_bytes()[..start] + .iter() + .rposition(|byte| !byte.is_ascii_alphanumeric() && !matches!(byte, b'_' | b'-')) + .map_or(0, |offset| offset + 1); + let enclosing = &value[enclosing_start..start + length]; + (length >= provider.minimum_credential_len || !virtualized.contains(credential)) + .then_some(credential) + .filter(|credential| !ignored_credential_match(provider, value, start, credential)) + .filter(|_| { + !provider.ignored_credential_prefixes.iter().any(|prefix| { + enclosing.match_indices(prefix).any(|(offset, _)| { + let ignored = &enclosing[offset..]; + offset <= start - enclosing_start + && ignored_credential_match( + provider, + value, + enclosing_start + offset, + ignored, + ) + && virtualized.contains(ignored) + }) + }) + }) +} + +pub(super) fn builtin_credential_candidate<'a>( + provider: &providers::CredentialProvider, + value: &'a str, + start: usize, +) -> &'a str { + let mut length = value.as_bytes()[start..] + .iter() + .take_while(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-')) + .count(); + let candidate = &value[start..start + length]; + let current_prefix_len = provider + .credential_prefixes + .iter() + .filter(|prefix| candidate.starts_with(**prefix)) + .map(|prefix| prefix.len()) + .max() + .unwrap_or(0); + if let Some(separator) = providers::credential_providers() + .flat_map(|candidate_provider| { + candidate_provider + .credential_prefixes + .iter() + .map(move |prefix| (prefix, candidate_provider.minimum_credential_len)) + }) + .filter_map(|(candidate_prefix, minimum_length)| { + candidate[current_prefix_len..] + .match_indices(*candidate_prefix) + .find_map(|(offset, _)| { + let offset = current_prefix_len + offset; + (matches!(candidate.as_bytes()[offset - 1], b'_' | b'-') + && offset > provider.minimum_credential_len + && candidate.len() - offset >= minimum_length) + .then_some(offset - 1) + }) + }) + .min() + { + length = separator; + } + let candidate = &value[start..start + length]; + if value.as_bytes().get(start + length) == Some(&b'\0') { + // A masked known span may follow an alias separator, not part of this builtin token. + candidate.trim_end_matches(['_', '-']) + } else { + candidate + } +} diff --git a/codex-rs/network-proxy/src/credential_broker/provider_config.rs b/codex-rs/network-proxy/src/credential_broker/provider_config.rs new file mode 100644 index 0000000000000000000000000000000000000000..566fa227807bbf158f396540598d360b9fb7614b --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/provider_config.rs @@ -0,0 +1,40 @@ +use schemars::JsonSchema; +use serde::Deserialize; +use serde::Serialize; + +/// Declarative description of an environment-backed credential family. +#[derive(Clone, Debug, Default, Deserialize, Eq, JsonSchema, PartialEq, Serialize)] +#[serde(default, deny_unknown_fields)] +pub struct CredentialProviderConfig { + pub env: Vec, + pub patterns: Vec, + /// URL prefixes authorized for injection. Bare loopback hosts imply HTTP; other bare hosts + /// imply HTTPS. + pub url_prefixes: Vec, + /// Environment variable containing an additional URL prefix or hostname. + #[serde(skip_serializing_if = "Option::is_none")] + pub url_prefix_from_env: Option, + pub auth: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub header: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub prefix: Option, +} + +impl CredentialProviderConfig { + /// Validate a complete provider definition using the broker's compilation rules. + pub fn validate(&self, id: &str) -> anyhow::Result<()> { + super::configured::ConfiguredCredentialProvider::compile(id, self)?; + Ok(()) + } +} + +/// Authentication formats supported by declarative credential providers. +#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum CredentialAuthMethod { + Bearer, + Token, + Basic, + Header, +} diff --git a/codex-rs/network-proxy/src/credential_broker/providers.rs b/codex-rs/network-proxy/src/credential_broker/providers.rs new file mode 100644 index 0000000000000000000000000000000000000000..aff11dc1bcea6432297c28321ef922b2b3a888d6 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/providers.rs @@ -0,0 +1,205 @@ +mod github; +mod openai; + +use super::destination::CredentialDestination; +use rama_http::HeaderMap; +use rama_http::HeaderValue; +use rand::Rng as _; +use std::collections::HashMap; +use url::Url; + +const DUMMY_ALPHANUMERIC: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"; + +type TranslateRequestHeader = fn(&HeaderMap, &str, &str) -> Option; + +/// Describes how one credential family is recognized and injected. +/// +/// Providers must be declared as `static` values because the broker uses their addresses as stable +/// identities when deduplicating credential records. +pub(super) struct CredentialProvider { + pub(super) context_env_vars: &'static [&'static str], + pub(super) credential_prefixes: &'static [&'static str], + pub(super) ignored_credential_prefixes: &'static [&'static str], + pub(super) credential_watermark: Option<&'static str>, + pub(super) minimum_credential_len: usize, + sources: &'static [CredentialSource], + pub(super) reset_on_configuration_change: bool, + dummy_value: fn(&str) -> String, + translate_request_header: TranslateRequestHeader, + request_header_value: fn(&str) -> Option, + insert_request_header: fn(&mut HeaderMap, HeaderValue), +} + +#[derive(Clone, PartialEq, Eq)] +pub(super) enum CredentialHostBinding { + ExactHost(String), + ExactHosts(Vec), + HostPattern { + exact_hosts: &'static [&'static str], + suffixes: &'static [&'static str], + }, + ConfiguredHosts(Vec), +} + +pub(super) struct CredentialSource { + pub(super) env_vars: &'static [&'static str], + pub(super) binding_env_vars: &'static [&'static str], + // An explicit invalid hint can clear dynamic destinations while leaving a static fallback. + pub(super) invalidates_host_binding: fn(&HashMap) -> bool, + pub(super) host_binding: + fn(&HashMap, Option<&str>) -> Option, +} + +const CREDENTIAL_PROVIDERS: &[&CredentialProvider] = &[&github::PROVIDER, &openai::PROVIDER]; + +impl CredentialProvider { + pub(super) fn sources(&self) -> &[CredentialSource] { + self.sources + } + + pub(super) fn dummy_value(&self, real_value: &str) -> String { + (self.dummy_value)(real_value) + } + + pub(super) fn translate_request_header( + &self, + headers: &HeaderMap, + expected_value: &str, + replacement_value: &str, + ) -> Option { + (self.translate_request_header)(headers, expected_value, replacement_value) + } + + pub(super) fn request_header_value(&self, value: &str) -> Option { + (self.request_header_value)(value) + } + + pub(super) fn insert_request_header(&self, headers: &mut HeaderMap, value: HeaderValue) { + (self.insert_request_header)(headers, value); + } +} + +impl CredentialHostBinding { + pub(super) fn bypasses_proxy_with_local_binding(&self) -> bool { + use crate::policy::is_default_proxy_bypass_host; + + match self { + Self::ExactHost(host) => is_default_proxy_bypass_host(host), + Self::ExactHosts(hosts) => hosts.iter().any(|host| is_default_proxy_bypass_host(host)), + Self::HostPattern { exact_hosts, .. } => exact_hosts + .iter() + .any(|host| is_default_proxy_bypass_host(host)), + Self::ConfiguredHosts(destinations) => destinations + .iter() + .any(CredentialDestination::bypasses_proxy_with_local_binding), + } + } + + pub(super) fn matches_host(&self, host: &str, port: u16) -> bool { + match self { + Self::ExactHost(expected_host) => host == expected_host, + Self::ExactHosts(expected_hosts) => { + expected_hosts.iter().any(|expected| host == expected) + } + Self::HostPattern { + exact_hosts, + suffixes, + } => { + exact_hosts.contains(&host) || suffixes.iter().any(|suffix| host.ends_with(suffix)) + } + Self::ConfiguredHosts(destinations) => destinations + .iter() + .any(|destination| destination.matches_host(host, port)), + } + } + + pub(super) fn requires_mitm(&self, host: &str, port: u16) -> bool { + match self { + Self::ConfiguredHosts(destinations) => destinations + .iter() + .any(|destination| destination.requires_mitm(host, port)), + Self::ExactHost(_) | Self::ExactHosts(_) | Self::HostPattern { .. } => { + self.matches_host(host, port) + } + } + } + + pub(super) fn matches_request(&self, host: &str, request: Option<&Url>) -> bool { + match self { + Self::ConfiguredHosts(destinations) => destinations + .iter() + .any(|destination| destination.matches_request(host, request)), + Self::ExactHost(_) | Self::ExactHosts(_) | Self::HostPattern { .. } => { + request.is_none_or(|request| request.scheme() == "https") + && self.matches_host( + host, + request.and_then(Url::port_or_known_default).unwrap_or(443), + ) + } + } + } +} + +pub(super) fn credential_context_env_keys( + brokered_keys: &[String], +) -> impl Iterator + '_ { + credential_providers() + .filter(move |provider| { + provider.sources().iter().any(|source| { + source.env_vars.iter().any(|key| { + brokered_keys + .iter() + .any(|brokered_key| super::env_key_matches(brokered_key, key)) + }) + }) + }) + .flat_map(|provider| provider.context_env_vars.iter().copied()) +} + +pub(super) fn credential_binding_env_keys( + brokered_keys: &[String], +) -> impl Iterator + '_ { + credential_providers() + .flat_map(CredentialProvider::sources) + .filter(move |source| { + source.env_vars.iter().any(|key| { + brokered_keys + .iter() + .any(|brokered_key| super::env_key_matches(brokered_key, key)) + }) + }) + .flat_map(|source| source.binding_env_vars.iter().copied()) +} + +pub(super) fn credential_env_keys() -> impl Iterator { + credential_providers() + .flat_map(CredentialProvider::sources) + .flat_map(|source| source.env_vars.iter().copied()) +} + +pub(super) fn credential_providers() -> impl Iterator { + CREDENTIAL_PROVIDERS.iter().copied() +} + +pub(super) fn translate_standard_request_header( + headers: &HeaderMap, + expected_value: &str, + replacement_value: &str, +) -> Option { + github::translate_request_header(headers, expected_value, replacement_value) +} + +fn shaped_dummy_value(real_value: &str, prefix: &str, minimum_len: usize) -> String { + let target_len = real_value.len().max(minimum_len).max(prefix.len() + 16); + let mut rng = rand::rng(); + let mut dummy = String::with_capacity(target_len); + dummy.push_str(prefix); + for index in prefix.len()..target_len { + let character = match real_value.as_bytes().get(index).copied() { + Some(template) if !template.is_ascii_alphanumeric() => template, + _ => DUMMY_ALPHANUMERIC[rng.random_range(0..DUMMY_ALPHANUMERIC.len())], + }; + dummy.push(char::from(character)); + } + dummy +} diff --git a/codex-rs/network-proxy/src/credential_broker/providers/github.rs b/codex-rs/network-proxy/src/credential_broker/providers/github.rs new file mode 100644 index 0000000000000000000000000000000000000000..e79dc56d746cdbb101e7c6cd75046b2aebac4643 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/providers/github.rs @@ -0,0 +1,171 @@ +use super::super::env_value; +use super::CredentialHostBinding; +use super::CredentialProvider; +use super::CredentialSource; +use super::shaped_dummy_value; +use crate::policy::normalize_host; +use base64::Engine as _; +use rama_http::HeaderMap; +use rama_http::HeaderValue; +use rama_http::header::AUTHORIZATION; +use std::collections::HashMap; + +const GH_HOST_ENV_VAR: &str = "GH_HOST"; +const GITHUB_TOKEN_PREFIXES: &[&str] = &["github_pat_", "ghp_", "gho_", "ghu_", "ghs_", "ghr_"]; +const GITHUB_TOKEN_MIN_LEN: usize = 40; +const GITHUB_CLOUD_TOKEN_ENV_VARS: &[&str] = &["GH_TOKEN", "GITHUB_TOKEN"]; +const GITHUB_ENTERPRISE_TOKEN_ENV_VARS: &[&str] = + &["GH_ENTERPRISE_TOKEN", "GITHUB_ENTERPRISE_TOKEN"]; +const GITHUB_CLOUD_HOSTS: &[&str] = &["api.github.com", "github.com", "uploads.github.com"]; +const GITHUB_CLOUD_HOST_SUFFIXES: &[&str] = &[".ghe.com"]; + +pub(super) static PROVIDER: CredentialProvider = CredentialProvider { + context_env_vars: &[GH_HOST_ENV_VAR], + credential_prefixes: GITHUB_TOKEN_PREFIXES, + ignored_credential_prefixes: &[], + credential_watermark: None, + minimum_credential_len: GITHUB_TOKEN_MIN_LEN, + sources: &[ + CredentialSource { + env_vars: GITHUB_CLOUD_TOKEN_ENV_VARS, + binding_env_vars: &[], + invalidates_host_binding: |_| false, + host_binding: github_cloud_binding, + }, + CredentialSource { + env_vars: GITHUB_ENTERPRISE_TOKEN_ENV_VARS, + binding_env_vars: &[GH_HOST_ENV_VAR], + invalidates_host_binding: |env| { + env_value(env, GH_HOST_ENV_VAR).is_some() + && github_host_hint(env).is_none_or(|host| github_cloud_host(&host)) + }, + host_binding: github_enterprise_binding, + }, + ], + reset_on_configuration_change: false, + dummy_value, + translate_request_header, + request_header_value, + insert_request_header, +}; + +fn dummy_value(real_value: &str) -> String { + shaped_dummy_value( + real_value, + github_token_prefix(real_value), + GITHUB_TOKEN_MIN_LEN, + ) +} + +fn request_header_value(value: &str) -> Option { + HeaderValue::from_str(&format!("Bearer {value}")).ok() +} + +fn insert_request_header(headers: &mut HeaderMap, value: HeaderValue) { + headers.insert(AUTHORIZATION, value); +} + +pub(super) fn translate_request_header( + headers: &HeaderMap, + expected_value: &str, + replacement_value: &str, +) -> Option { + let header = headers.get(AUTHORIZATION)?.to_str().ok()?; + let (scheme, value) = header.split_once(' ')?; + let value = value.trim(); + if scheme.eq_ignore_ascii_case("basic") { + let decoded = base64::engine::general_purpose::STANDARD + .decode(value) + .ok()?; + let mut translated = Vec::with_capacity(decoded.len() + replacement_value.len()); + if decoded == expected_value.as_bytes() { + translated.extend_from_slice(replacement_value.as_bytes()); + } else { + let separator = decoded.iter().position(|byte| *byte == b':')?; + let username = &decoded[..separator]; + let password = &decoded[separator + 1..]; + if username != expected_value.as_bytes() && password != expected_value.as_bytes() { + return None; + } + translated.extend_from_slice(if username == expected_value.as_bytes() { + replacement_value.as_bytes() + } else { + username + }); + translated.push(b':'); + translated.extend_from_slice(if password == expected_value.as_bytes() { + replacement_value.as_bytes() + } else { + password + }); + } + let encoded = base64::engine::general_purpose::STANDARD.encode(translated); + HeaderValue::from_str(&format!("{scheme} {encoded}")).ok() + } else if (scheme.eq_ignore_ascii_case("bearer") || scheme.eq_ignore_ascii_case("token")) + && value == expected_value + { + HeaderValue::from_str(&format!("{scheme} {replacement_value}")).ok() + } else { + None + } +} + +fn github_cloud_binding( + _: &HashMap, + _: Option<&str>, +) -> Option { + Some(CredentialHostBinding::HostPattern { + exact_hosts: GITHUB_CLOUD_HOSTS, + suffixes: GITHUB_CLOUD_HOST_SUFFIXES, + }) +} + +fn github_enterprise_binding( + env: &HashMap, + _: Option<&str>, +) -> Option { + github_host_hint(env) + .filter(|host| !github_cloud_host(host)) + .map(CredentialHostBinding::ExactHost) +} + +fn github_cloud_host(host: &str) -> bool { + GITHUB_CLOUD_HOSTS.contains(&host) + || GITHUB_CLOUD_HOST_SUFFIXES + .iter() + .any(|suffix| host.ends_with(suffix)) +} + +fn github_token_prefix(value: &str) -> &str { + GITHUB_TOKEN_PREFIXES + .iter() + .copied() + .find(|prefix| value.starts_with(prefix)) + .unwrap_or("ghp_") +} + +fn github_host_hint(env: &HashMap) -> Option { + let value = env_value(env, GH_HOST_ENV_VAR)?.trim(); + if value.is_empty() || value.contains(['/', '\\', '@', '?', '#', '*']) { + return None; + } + let authority = if value.parse::().is_ok() + || crate::policy::unscoped_ip_literal(value).is_some() + { + format!("[{value}]") + } else { + value.to_string() + }; + let parsed = authority.parse::().ok()?; + let suffix = authority.strip_prefix(parsed.host())?; + if !suffix.is_empty() { + suffix.strip_prefix(':')?.parse::().ok()?; + } + let host = normalize_host(parsed.host()); + let ip_literal = host.parse::().is_ok() + || crate::policy::unscoped_ip_literal(&host).is_some(); + if !ip_literal && (parsed.host().starts_with('[') || url::Host::parse(&host).is_err()) { + return None; + } + Some(host) +} diff --git a/codex-rs/network-proxy/src/credential_broker/providers/openai.rs b/codex-rs/network-proxy/src/credential_broker/providers/openai.rs new file mode 100644 index 0000000000000000000000000000000000000000..3b6808aaf7d95c34f422801af5ec77ac18021e7a --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/providers/openai.rs @@ -0,0 +1,95 @@ +use super::super::env_value; +use super::CredentialHostBinding; +use super::CredentialProvider; +use super::CredentialSource; +use super::shaped_dummy_value; +use crate::config::trusted_credential_broker_host; +use rama_http::HeaderMap; +use rama_http::HeaderValue; +use rama_http::header::AUTHORIZATION; +use std::collections::HashMap; + +const OPENAI_API_KEY_ENV_VARS: &[&str] = &["OPENAI_API_KEY"]; +const OPENAI_API_KEY_PREFIXES: &[&str] = &["sk-proj-", "sk-svcacct-", "sk-admin-", "sk-"]; +const OPENAI_BASE_URL_ENV_VAR: &str = "OPENAI_BASE_URL"; +const OPENAI_API_KEY_MIN_LEN: usize = 51; +const OPENAI_API_HOST: &str = "api.openai.com"; + +pub(super) static PROVIDER: CredentialProvider = CredentialProvider { + context_env_vars: &[OPENAI_BASE_URL_ENV_VAR], + credential_prefixes: OPENAI_API_KEY_PREFIXES, + ignored_credential_prefixes: &["sk-ant-", "sk-or-"], + credential_watermark: Some("T3BlbkFJ"), + minimum_credential_len: OPENAI_API_KEY_MIN_LEN, + sources: &[CredentialSource { + env_vars: OPENAI_API_KEY_ENV_VARS, + binding_env_vars: &[OPENAI_BASE_URL_ENV_VAR], + invalidates_host_binding: |env| { + env_value(env, OPENAI_BASE_URL_ENV_VAR) + .is_some_and(|value| trusted_credential_broker_host(value).is_none()) + }, + host_binding, + }], + reset_on_configuration_change: true, + dummy_value, + translate_request_header, + request_header_value, + insert_request_header, +}; + +fn dummy_value(real_value: &str) -> String { + shaped_dummy_value( + real_value, + openai_api_key_prefix(real_value), + OPENAI_API_KEY_MIN_LEN, + ) +} + +fn translate_request_header( + headers: &HeaderMap, + expected_value: &str, + replacement_value: &str, +) -> Option { + headers + .get(AUTHORIZATION) + .and_then(|header| header.to_str().ok()) + .filter(|header| header.contains(expected_value))?; + request_header_value(replacement_value) +} + +fn request_header_value(value: &str) -> Option { + HeaderValue::from_str(&format!("Bearer {value}")).ok() +} + +fn insert_request_header(headers: &mut HeaderMap, value: HeaderValue) { + headers.insert(AUTHORIZATION, value); +} + +fn host_binding( + env: &HashMap, + configured_host: Option<&str>, +) -> Option { + let mut hosts = vec![OPENAI_API_HOST.to_string()]; + for host in configured_host + .map(str::to_string) + .into_iter() + .chain(env_value(env, OPENAI_BASE_URL_ENV_VAR).and_then(trusted_credential_broker_host)) + { + if !hosts.contains(&host) { + hosts.push(host); + } + } + + Some(match hosts.as_slice() { + [host] => CredentialHostBinding::ExactHost(host.clone()), + _ => CredentialHostBinding::ExactHosts(hosts), + }) +} + +fn openai_api_key_prefix(value: &str) -> &str { + OPENAI_API_KEY_PREFIXES + .iter() + .copied() + .find(|prefix| value.starts_with(prefix)) + .unwrap_or("sk-") +} diff --git a/codex-rs/network-proxy/src/credential_broker/registry.rs b/codex-rs/network-proxy/src/credential_broker/registry.rs new file mode 100644 index 0000000000000000000000000000000000000000..d9e1555f034dcf2d607de99abeca27d930f29435 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/registry.rs @@ -0,0 +1,442 @@ +use super::CredentialBrokerState; +use super::CredentialRecord; +use super::MIN_EMBEDDED_CREDENTIAL_LENGTH; +use super::configured::ConfiguredCredentialProvider; +use super::env_entry; +use super::env_key_matches; +use super::env_value; +use super::matching; +use super::matching::is_operational_path_match; +use super::providers; +use super::source_accepts_credential; +use super::source_tracks_credential; +use base64::Engine; +use rama_http::HeaderMap; +use rama_http::HeaderName; +use rama_http::HeaderValue; +use rama_http::header::AUTHORIZATION; +use std::collections::HashMap; +use std::sync::Arc; +use url::Url; + +#[derive(Clone)] +pub(super) enum BrokeredCredentialProvider { + Builtin(&'static providers::CredentialProvider), + Configured(Arc), +} + +#[derive(Clone)] +pub(super) struct ActiveCredentialSource { + pub(super) provider: BrokeredCredentialProvider, + pub(super) host_binding: providers::CredentialHostBinding, + pub(super) env_vars: Vec, +} + +impl BrokeredCredentialProvider { + pub(super) fn same_provider(&self, other: &Self) -> bool { + match (self, other) { + (Self::Builtin(left), Self::Builtin(right)) => std::ptr::eq(*left, *right), + (Self::Configured(left), Self::Configured(right)) => Arc::ptr_eq(left, right), + (Self::Builtin(_), Self::Configured(_)) | (Self::Configured(_), Self::Builtin(_)) => { + false + } + } + } + + pub(super) fn dummy_value(&self, real_value: &str) -> Option { + match self { + Self::Builtin(provider) => Some(provider.dummy_value(real_value)), + Self::Configured(provider) => provider.dummy_value(real_value), + } + } + + pub(super) fn recognizes_credential_in(&self, value: &str) -> bool { + match self { + Self::Builtin(provider) => provider.credential_prefixes.iter().any(|prefix| { + value.match_indices(*prefix).any(|(start, _)| { + matching::recognized_credential_match(provider, value, "", start).is_some() + }) + }), + Self::Configured(provider) => provider.find_credential_matches(value).next().is_some(), + } + } + + pub(super) fn recognizes_strictly_embedded_dummy_collision_in(&self, value: &str) -> bool { + match self { + Self::Builtin(provider) => provider.credential_prefixes.iter().any(|prefix| { + value.match_indices(*prefix).any(|(start, _)| { + matching::recognized_credential_match(provider, value, "", start) + .is_some_and(|credential| start > 0 || credential.len() < value.len()) + }) + }), + Self::Configured(provider) => provider.contains_strictly_embedded_pattern_match(value), + } + } + + pub(super) fn contains_embedded_value(&self, text: &str, value: &str) -> bool { + match self { + Self::Builtin(_) => { + value.len() >= MIN_EMBEDDED_CREDENTIAL_LENGTH && text.contains(value) + } + Self::Configured(provider) => provider + .credential_value_match_ranges(text, value) + .into_iter() + .any(|range| { + value.len() >= MIN_EMBEDDED_CREDENTIAL_LENGTH + || !is_operational_path_match(text, range.start, range.end) + }), + } + } + + pub(super) fn request_header_value(&self, value: &str) -> Option { + match self { + Self::Builtin(provider) => provider.request_header_value(value), + Self::Configured(provider) => provider + .matches_value(value) + .then(|| provider.request_header_value(value)) + .flatten(), + } + } + + pub(super) fn translate_request_headers( + &self, + headers: &HeaderMap, + expected_value: &str, + replacement_value: &str, + ) -> Vec<(HeaderName, HeaderValue)> { + match self { + Self::Builtin(provider) => provider + .translate_request_header(headers, expected_value, replacement_value) + .map(|value| (AUTHORIZATION, value)), + Self::Configured(provider) => { + return provider.translate_request_headers( + headers, + expected_value, + replacement_value, + ); + } + } + .into_iter() + .collect() + } + + pub(super) fn insert_request_header( + &self, + headers: &mut HeaderMap, + name: HeaderName, + value: HeaderValue, + ) { + match self { + Self::Builtin(provider) => provider.insert_request_header(headers, value), + Self::Configured(_) => { + headers.insert(name, value); + } + } + } +} + +impl CredentialRecord { + pub(super) fn contains_embedded_value(&self, text: &str, value: &str) -> bool { + !self.value_match_ranges(text, value).is_empty() + } + + pub(super) fn value_match_ranges( + &self, + text: &str, + value: &str, + ) -> Vec> { + let mut ranges = if text == value { + std::iter::once(0..text.len()).collect() + } else if is_builtin_shaped_credential(&self.real_value) + || matches!(self.provider, BrokeredCredentialProvider::Builtin(_)) + && value.len() >= MIN_EMBEDDED_CREDENTIAL_LENGTH + { + text.match_indices(value) + .map(|(start, _)| start..start + value.len()) + .collect() + } else if let BrokeredCredentialProvider::Configured(provider) = &self.provider { + provider + .credential_value_match_ranges(text, value) + .into_iter() + .filter(|range| { + value.len() >= MIN_EMBEDDED_CREDENTIAL_LENGTH + || !is_operational_path_match(text, range.start, range.end) + }) + .collect() + } else { + Vec::new() + }; + if value == self.dummy_value { + ranges.extend(self.generated_dummy_ranges(text)); + ranges.sort_by_key(|range| range.start); + ranges.dedup(); + } + ranges + } +} + +pub(super) fn is_builtin_shaped_credential(value: &str) -> bool { + providers::credential_providers().any(|provider| { + value.len() >= provider.minimum_credential_len + && provider + .credential_prefixes + .iter() + .any(|prefix| value.starts_with(prefix)) + && matching::recognized_credential_match(provider, value, "", /*start*/ 0) + .is_some_and(|credential| credential == value) + }) +} + +pub(super) fn select_credentials<'a>( + headers: &HeaderMap, + host: &str, + request: Option<&Url>, + credentials: &'a [CredentialRecord], + environment_id: Option<&str>, +) -> Vec<(&'a CredentialRecord, HeaderName, HeaderValue)> { + let mut translated_matches = Vec::<(&CredentialRecord, HeaderName, HeaderValue)>::new(); + for credential in credentials.iter().filter(|credential| { + credential.belongs_to_environment(environment_id) + && credential + .host_bindings() + .any(|binding| binding.matches_request(host, request)) + }) { + for candidate in credentials.iter().filter(|candidate| { + candidate.belongs_to_environment(environment_id) + && candidate.provider.same_provider(&credential.provider) + && candidate.real_value == credential.real_value + && candidate + .host_bindings() + .any(|binding| binding.matches_request(host, request)) + }) { + for (name, value) in credential.provider.translate_request_headers( + headers, + &candidate.dummy_value, + &credential.real_value, + ) { + if !translated_matches + .iter() + .any(|(existing, header, translated)| { + existing.provider.same_provider(&credential.provider) + && existing.real_value == credential.real_value + && *header == name + && *translated == value + }) + { + translated_matches.push((credential, name, value)); + } + } + } + } + let mut selected = Vec::<(&CredentialRecord, HeaderName, HeaderValue)>::new(); + let mut ambiguous_headers = Vec::::new(); + for (credential, header_name, header_value) in translated_matches { + if ambiguous_headers.contains(&header_name) { + continue; + } + let Some(existing) = selected + .iter() + .position(|(_, selected_header, _)| selected_header == header_name) + else { + selected.push((credential, header_name, header_value)); + continue; + }; + if !credential + .provider + .same_provider(&selected[existing].0.provider) + || credential.real_value != selected[existing].0.real_value + || header_value != selected[existing].2 + { + if header_name == AUTHORIZATION + && let Some(merged) = headers.get(AUTHORIZATION).and_then(|original| { + merge_basic_auth_fields(original, &selected[existing].2, &header_value) + }) + { + // Keep the caller's configured-provider raw-path validation for either field. + if matches!( + credential.provider, + BrokeredCredentialProvider::Configured(_) + ) { + selected[existing].0 = credential; + } + selected[existing].2 = merged; + continue; + } + selected.swap_remove(existing); + ambiguous_headers.push(header_name); + } + } + selected +} + +fn merge_basic_auth_fields( + original: &HeaderValue, + left: &HeaderValue, + right: &HeaderValue, +) -> Option { + let decode = |header: &HeaderValue| { + let (scheme, value) = header.to_str().ok()?.split_once(' ')?; + scheme.eq_ignore_ascii_case("basic").then_some(())?; + let decoded = base64::engine::general_purpose::STANDARD + .decode(value.trim()) + .ok()?; + let separator = decoded.iter().position(|byte| *byte == b':')?; + Some(( + decoded[..separator].to_vec(), + decoded[separator + 1..].to_vec(), + )) + }; + let (original_user, original_password) = decode(original)?; + let (left_user, left_password) = decode(left)?; + let (right_user, right_password) = decode(right)?; + // Only combine independent fields. Overlapping translations remain ambiguous. + let (user, password) = if left_user != original_user + && left_password == original_password + && right_user == original_user + && right_password != original_password + { + (left_user, right_password) + } else if right_user != original_user + && right_password == original_password + && left_user == original_user + && left_password != original_password + { + (right_user, left_password) + } else { + return None; + }; + let encoded = + base64::engine::general_purpose::STANDARD.encode([user, vec![b':'], password].concat()); + HeaderValue::from_str(&format!("Basic {encoded}")).ok() +} + +pub(super) fn active_credential_sources( + state: &CredentialBrokerState, + env: &HashMap, +) -> Vec { + let binding_env = state.context.with_fallbacks(env); + let mut sources = Vec::new(); + for provider in providers::credential_providers() { + for source in provider.sources() { + if let Some(host_binding) = + (source.host_binding)(&binding_env, state.openai_api_host.as_deref()) + { + sources.push(ActiveCredentialSource { + provider: BrokeredCredentialProvider::Builtin(provider), + host_binding, + env_vars: source + .env_vars + .iter() + .map(|key| (*key).to_string()) + .collect(), + }); + } + } + } + for provider in &state.configured_providers { + if let Some(host_binding) = provider.host_binding(&binding_env) { + sources.push(ActiveCredentialSource { + provider: BrokeredCredentialProvider::Configured(provider.clone()), + host_binding, + env_vars: provider.config.env.clone(), + }); + } + } + sources +} + +pub(super) fn registered_credential_source( + credential: &CredentialRecord, + sources: &[ActiveCredentialSource], + env: &HashMap, +) -> Option { + let retained_env_vars = match &credential.provider { + BrokeredCredentialProvider::Builtin(provider) => provider + .sources() + .iter() + .find(|source| { + source + .env_vars + .iter() + .any(|key| env_key_matches(key, &credential.env_var)) + && !source.binding_env_vars.is_empty() + && source + .binding_env_vars + .iter() + .all(|key| env_value(env, key).is_none()) + }) + .map(|source| { + source + .env_vars + .iter() + .map(|key| (*key).to_string()) + .collect() + }), + BrokeredCredentialProvider::Configured(provider) => provider + .config + .url_prefix_from_env + .as_deref() + .filter(|key| env_value(env, key).is_none()) + .map(|_| provider.config.env.clone()), + }; + // Filtering a destination hint out of the child must not undo an existing + // registration or rebind it to the broker's ambient fallback. + if let Some(env_vars) = retained_env_vars { + Some(ActiveCredentialSource { + provider: credential.provider.clone(), + host_binding: credential.host_binding.clone(), + env_vars, + }) + } else { + sources + .iter() + .find(|source| source_accepts_credential(source, credential)) + .cloned() + .map(|mut source| { + // A new canonical token does not authorize an older alias at its destination. + if !credential.invalidates_host_binding(env) + && !source_tracks_credential(&source, credential, env) + { + source.host_binding = credential.host_binding.clone(); + } + source + }) + } +} + +pub(super) fn prioritized_credentials<'a>( + state: &'a CredentialBrokerState, + env: &HashMap, +) -> Vec<&'a CredentialRecord> { + let binding_env = state.context.with_fallbacks(env); + let mut credentials = state.credentials.iter().collect::>(); + credentials.sort_unstable_by_key(|credential| { + let active = env_entry(env, &credential.env_var) + .is_some_and(|(_, value)| value == credential.dummy_value); + let has_matching_host_binding = match &credential.provider { + BrokeredCredentialProvider::Builtin(provider) => { + provider.sources().iter().any(|source| { + !source.binding_env_vars.is_empty() + && (source.host_binding)(&binding_env, state.openai_api_host.as_deref()) + .is_some_and(|binding| binding == credential.host_binding) + }) + } + BrokeredCredentialProvider::Configured(provider) => provider + .config + .url_prefix_from_env + .as_deref() + .is_some_and(|key| { + env_value(&binding_env, key).is_some() + && provider + .host_binding(&binding_env) + .is_some_and(|binding| binding == credential.host_binding) + }), + }; + ( + std::cmp::Reverse(credential.real_value.len()), + std::cmp::Reverse(active), + std::cmp::Reverse(has_matching_host_binding), + ) + }); + credentials +} diff --git a/codex-rs/network-proxy/src/credential_broker/replacement.rs b/codex-rs/network-proxy/src/credential_broker/replacement.rs new file mode 100644 index 0000000000000000000000000000000000000000..fde7c42ab93cb438c313ac6be32615bf386b2b6c --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker/replacement.rs @@ -0,0 +1,155 @@ +use super::CredentialBrokerState; +use super::CredentialRecord; +use sha2::Digest; +use sha2::Sha256; +use std::collections::HashMap; +use std::ops::Range; + +impl CredentialBrokerState { + pub(super) fn replace_child_env_dummies( + &self, + env: &mut HashMap, + replacements: &[(String, String)], + ) { + if replacements.is_empty() { + return; + } + let credentials = self + .credentials + .iter() + .map(|credential| { + let replacement = replacements + .iter() + .find(|(dummy, _)| *dummy == credential.dummy_value) + .map_or(credential.dummy_value.as_str(), |(_, replacement)| { + replacement + }); + let target = self + .credentials + .iter() + .find(|candidate| candidate.dummy_value == replacement); + (credential, replacement, target) + }) + .collect::>(); + for value in env.values_mut() { + let mut spans = Replacements::default(); + for (credential, replacement, target) in &credentials { + spans.add( + credential.value_match_ranges(value, &credential.dummy_value), + replacement, + *target, + ); + } + spans.render(value); + } + } +} + +// Retain emitted context without keeping another copy of any real credentials in it. +#[derive(Clone, PartialEq, Eq)] +pub(super) struct GeneratedAlias { + digest: [u8; 32], + len: usize, + range: Range, +} + +impl CredentialRecord { + pub(super) fn generated_dummy_ranges(&self, text: &str) -> Vec> { + let aliases = self + .generated_aliases + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner); + text.char_indices() + .filter(|(start, _)| text[*start..].starts_with(&self.dummy_value)) + .filter_map(|(start, _)| { + aliases + .iter() + .any(|alias| { + start.checked_sub(alias.range.start).is_some_and(|offset| { + text.get(offset..offset + alias.len).is_some_and(|context| { + <[u8; 32]>::from(Sha256::digest(context)) == alias.digest + }) + }) + }) + .then_some(start..start + self.dummy_value.len()) + }) + .collect() + } +} + +#[derive(Default)] +pub(super) struct Replacements<'a> { + spans: Vec<(Range, &'a str, Option<&'a CredentialRecord>)>, +} + +impl<'a> Replacements<'a> { + pub(super) fn add( + &mut self, + ranges: impl IntoIterator>, + value: &'a str, + dummy: Option<&'a CredentialRecord>, + ) { + self.spans + .extend(ranges.into_iter().map(|range| (range, value, dummy))); + } + + pub(super) fn masked(&self, text: &str) -> String { + let mut masked = text.to_string(); + for (range, _, _) in &self.spans { + masked.replace_range(range.clone(), &"\0".repeat(range.len())); + } + masked + } + + pub(super) fn render(mut self, text: &mut String) -> bool { + if self.spans.is_empty() { + return false; + } + self.spans + .sort_by_key(|(range, _, _)| std::cmp::Reverse(range.len())); + let mut selected = Vec::<(Range, &str, Option<&CredentialRecord>)>::new(); + for span in self.spans { + if !selected + .iter() + .any(|(range, _, _)| range.start < span.0.end && span.0.start < range.end) + { + selected.push(span); + } + } + selected.sort_by_key(|(range, _, _)| range.start); + let mut output = String::with_capacity(text.len()); + let mut previous = 0; + let mut dummies = Vec::new(); + for (range, value, dummy) in selected { + output.push_str(&text[previous..range.start]); + let start = output.len(); + output.push_str(value); + if let Some(dummy) = dummy { + dummies.push((dummy, start..output.len())); + } + previous = range.end; + } + output.push_str(&text[previous..]); + for (dummy, range) in dummies { + if !dummy + .value_match_ranges(&output, &dummy.dummy_value) + .contains(&range) + { + let alias = GeneratedAlias { + digest: Sha256::digest(&output).into(), + len: output.len(), + range, + }; + let mut aliases = dummy + .generated_aliases + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if !aliases.contains(&alias) { + aliases.push(alias); + } + } + } + *text = output; + true + } +} diff --git a/codex-rs/network-proxy/src/credential_broker_tests.rs b/codex-rs/network-proxy/src/credential_broker_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..a368d2258e4a8f497469e0bf2bb8442937c01cb2 --- /dev/null +++ b/codex-rs/network-proxy/src/credential_broker_tests.rs @@ -0,0 +1,1601 @@ +use super::*; + +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD; +use pretty_assertions::assert_eq; +use rama_http::HeaderValue; +use rama_http::header::AUTHORIZATION; + +fn env_map(entries: [(&str, &str); N]) -> HashMap { + entries + .into_iter() + .map(|(key, value)| (key.to_string(), value.to_string())) + .collect() +} + +fn headers_with_bearer(value: &str) -> HeaderMap { + headers_with_authorization(&format!("Bearer {value}")) +} + +fn headers_with_authorization(value: &str) -> HeaderMap { + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(value).expect("valid authorization header"), + ); + headers +} + +fn authorization(headers: &HeaderMap) -> Option<&str> { + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()) +} + +fn assert_credential_shape(real_value: &str, dummy_value: &str, prefix: &str) { + assert_ne!(dummy_value, real_value); + assert_eq!(dummy_value.len(), real_value.len()); + assert_eq!(&dummy_value[..prefix.len()], prefix); + let same_shape = real_value + .bytes() + .zip(dummy_value.bytes()) + .skip(prefix.len()) + .all(|(real, dummy)| { + real.is_ascii_alphanumeric() && dummy.is_ascii_alphanumeric() || real == dummy + }); + assert!(same_shape); +} + +#[test] +fn virtualize_child_env_replaces_supported_credentials() { + let broker = CredentialBroker::new(/*enabled*/ true); + let github_token = "github_pat_11AA0bbCC_abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGH"; + let openai_api_key = "sk-proj-abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789-_"; + let authorization = format!("Bearer {github_token}"); + let mut env = env_map([ + ("GH_TOKEN", github_token), + ("HOMEBREW_GITHUB_API_TOKEN", github_token), + ("AUTH_HEADER", authorization.as_str()), + ("OPENAI_API_KEY", openai_api_key), + ("GH_ENTERPRISE_TOKEN", github_token), + ]); + + broker.virtualize_child_env(&mut env); + + let github_dummy = env.get("GH_TOKEN").expect("dummy GitHub token"); + let openai_dummy = env.get("OPENAI_API_KEY").expect("dummy OpenAI API key"); + assert_credential_shape(github_token, github_dummy, "github_pat_"); + assert_credential_shape(openai_api_key, openai_dummy, "sk-proj-"); + assert_eq!(env.get("HOMEBREW_GITHUB_API_TOKEN"), Some(github_dummy)); + assert_eq!(env.get("GH_ENTERPRISE_TOKEN"), Some(github_dummy)); + assert_eq!( + env.get("AUTH_HEADER"), + Some(&format!("Bearer {github_dummy}")) + ); + let mut persisted_credentials = format!("{github_token}\n{openai_api_key}"); + assert!(broker.virtualize_text(&mut persisted_credentials, &env)); + assert_eq!( + persisted_credentials, + format!("{github_dummy}\n{openai_dummy}") + ); + let mut filtered_env = env.clone(); + filtered_env.remove("OPENAI_API_KEY"); + let mut excluded_credentials = format!("{github_token}\n{openai_api_key}"); + assert!(!broker.virtualize_text(&mut excluded_credentials, &filtered_env)); + assert_eq!(excluded_credentials, format!("{github_dummy}\n")); + let mut excluded_dummies = format!("{github_dummy}\n{openai_dummy}"); + assert!(!broker.virtualize_text(&mut excluded_dummies, &filtered_env)); + assert_eq!(excluded_dummies, format!("{github_dummy}\n")); + let unknown_github_token = "ghp_0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefgh"; + let unknown_openai_key = "sk-proj-0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefgh"; + let unknown_legacy_openai_key = format!("sk-{}", "a".repeat(48)); + let mut unregistered = format!( + "{unknown_github_token}\n{unknown_openai_key}\n{unknown_legacy_openai_key}\nghp_x sk-proj-x" + ); + assert!(!broker.virtualize_text(&mut unregistered, &env)); + assert_eq!(unregistered, "\n\n\nghp_x sk-proj-x"); + for key in [ + "sk-ant-api03-0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmno", + "sk-ant-oat01-0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmno", + "sk-or-v1-0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmno", + ] { + let mut unrelated_provider = key.to_string(); + assert!(broker.virtualize_text(&mut unrelated_provider, &env)); + assert_eq!(unrelated_provider, key); + assert!(!credential_broker_provider_sources_allowed( + key, + key, + &HashMap::new(), + |_| true, + )); + } + let mut embedded_tokens = format!("v1{unknown_github_token}\n/opt/ta{unknown_openai_key}/bin"); + assert!(!broker.virtualize_text(&mut embedded_tokens, &env)); + assert_eq!(embedded_tokens, "v1\n/opt/ta/bin"); + let qualified_legacy = format!( + "sk-{}-{}T3BlbkFJ{}", + "a".repeat(20), + "b".repeat(19), + "c".repeat(20) + ); + let legacy_broker = CredentialBroker::new(/*enabled*/ true); + let mut legacy_env = env_map([("OPENAI_API_KEY", qualified_legacy.as_str())]); + legacy_broker.virtualize_child_env(&mut legacy_env); + assert_credential_shape(&qualified_legacy, &legacy_env["OPENAI_API_KEY"], "sk-"); + let unmarked_legacy = format!("sk-{}-{}", "a".repeat(15), "b".repeat(35)); + let mut unmarked_legacy_alias = format!("Bearer {unmarked_legacy}"); + assert!(!broker.virtualize_text(&mut unmarked_legacy_alias, &env)); + assert_eq!(unmarked_legacy_alias, "Bearer "); + let collision = format!( + "sk-{}-sk-{}T3BlbkFJ{}", + "a".repeat(20), + "b".repeat(16), + "c".repeat(20) + ); + let mut collision_alias = format!("Bearer {collision}"); + assert!(!broker.virtualize_text(&mut collision_alias, &env)); + assert_eq!(collision_alias, "Bearer "); + let unrelated_collision = format!("sk-ant-api03-{}-sk-{}", "a".repeat(40), "b".repeat(48)); + let mut redacted_collision = unrelated_collision.clone(); + assert!(!broker.virtualize_text(&mut redacted_collision, &env)); + assert_eq!( + redacted_collision, + format!("sk-ant-api03-{}-", "a".repeat(40)) + ); + assert!(credential_broker_provider_sources_allowed( + &unrelated_collision, + &redacted_collision, + &HashMap::new(), + |_| true, + )); + let unrelated = "sk-ant-api03-0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmno"; + let mut known_provider_env = env.clone(); + known_provider_env.insert("ANTHROPIC_API_KEY".to_string(), unrelated.to_string()); + for separator in ["_", "__", "--", "-_", "_-", "", "_openai_", "_Bearer_"] { + let mut known_adjacent = format!("{unrelated}{separator}{unknown_legacy_openai_key}"); + assert!(!broker.virtualize_text(&mut known_adjacent, &known_provider_env)); + assert_eq!(known_adjacent, format!("{unrelated}{separator}")); + } + let first_bundle = format!("{unrelated}_{unknown_legacy_openai_key}"); + known_provider_env.insert("FIRST_BUNDLE".to_string(), first_bundle.clone()); + let openrouter = "sk-or-v1-0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmno"; + for ignored in [unrelated, openrouter] { + for separator in ["_", "-", "__", "--", "", "_openai_", "_Bearer_"] { + let mut unknown_adjacent = format!("{ignored}{separator}{unknown_legacy_openai_key}"); + assert!(!broker.virtualize_text(&mut unknown_adjacent, &env)); + assert_eq!(unknown_adjacent, format!("{ignored}{separator}")); + } + } + let mut nested_bundle = format!("{first_bundle}_{openrouter}"); + assert!(!broker.virtualize_text(&mut nested_bundle, &known_provider_env)); + assert_eq!(nested_bundle, format!("{unrelated}__{openrouter}")); + let mixed_providers = format!("{unknown_github_token}_{unrelated_collision}"); + let mut virtualized_providers = mixed_providers.clone(); + assert!(!broker.virtualize_text(&mut virtualized_providers, &env)); + assert!(!credential_broker_provider_sources_allowed( + &mixed_providers, + &virtualized_providers, + &HashMap::new(), + |source| source != "OPENAI_API_KEY", + )); + let mixed_credentials = + format!("{unknown_github_token}_{unknown_legacy_openai_key}_{unrelated}"); + let mut virtualized_credentials = mixed_credentials.clone(); + assert!(!broker.virtualize_text(&mut virtualized_credentials, &env)); + assert!(!credential_broker_provider_sources_allowed( + &mixed_credentials, + &virtualized_credentials, + &HashMap::new(), + |source| source != "OPENAI_API_KEY", + )); + let equivalent_sources = env_map([ + ("GH_TOKEN", unknown_github_token), + ("GITHUB_TOKEN", unknown_github_token), + ]); + assert!(credential_broker_provider_sources_allowed( + unknown_github_token, + "", + &equivalent_sources, + |source| source == "GH_TOKEN", + )); + let distinct_github_token = "ghp_abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGH"; + let distinct_sources = env_map([ + ("GH_TOKEN", unknown_github_token), + ("GITHUB_TOKEN", distinct_github_token), + ]); + assert!(!credential_broker_provider_sources_allowed( + &format!("{unknown_github_token}\n{distinct_github_token}"), + "", + &distinct_sources, + |source| source == "GH_TOKEN", + )); + for unrelated in [ + unrelated, + "sk-or-v1-0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmno", + ] { + for separator in ['_', '-'] { + for (mut adjacent, expected) in [ + ( + format!("{unrelated}{separator}{qualified_legacy}"), + format!("{unrelated}{separator}"), + ), + ( + format!("{unknown_legacy_openai_key}{separator}{unrelated}"), + format!("{separator}{unrelated}"), + ), + ] { + assert!(!broker.virtualize_text(&mut adjacent, &env)); + assert_eq!(adjacent, expected); + } + } + } + let hash = "a".repeat(64); + let copied_openai_key = format!("sk-proj-{hash}"); + let mut copied_credential = format!("ta{copied_openai_key}"); + assert!(!broker.virtualize_text(&mut copied_credential, &env)); + assert_eq!(copied_credential, "ta"); + let mut path_embedded_credentials = + format!("/prefix/{unknown_legacy_openai_key}:/next\n/prefix/sk-proj-{hash}"); + assert!(!broker.virtualize_text(&mut path_embedded_credentials, &env)); + assert_eq!(path_embedded_credentials, "/prefix/:/next\n/prefix/"); + for separator in ['-', '_'] { + let mut embedded_path = format!("/workspace/token{separator}sk-proj-{hash}/bin"); + assert!(!broker.virtualize_text(&mut embedded_path, &env)); + assert_eq!(embedded_path, format!("/workspace/token{separator}/bin")); + } + let mut adjacent_path = format!("/workspace/tokensk-proj-{hash}/bin"); + assert!(!broker.virtualize_text(&mut adjacent_path, &env)); + assert_eq!(adjacent_path, "/workspace/token/bin"); + let mut word_adjacent_path = format!("/workspace/datask-proj-{hash}/bin"); + assert!(!broker.virtualize_text(&mut word_adjacent_path, &env)); + assert_eq!(word_adjacent_path, "/workspace/data/bin"); + for prefix in ["a", "di", "ma", "ri", "bri"] { + let mut embedded_credential = format!("Bearer {prefix}sk-{hash}"); + assert!(!broker.virtualize_text(&mut embedded_credential, &env)); + assert_eq!(embedded_credential, format!("Bearer {prefix}")); + } + for component in ["\u{e9}task", "e\u{301}task"] { + let mut unicode_adjacent_path = format!("/workspace/{component}-proj-{hash}/bin"); + assert!(!broker.virtualize_text(&mut unicode_adjacent_path, &env)); + assert_eq!( + unicode_adjacent_path, + format!("/workspace/{}/bin", component.strip_suffix("sk").unwrap()) + ); + } + let mut hashed_credential_path = + format!("/workspace/task-proj-0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefgh-{hash}/bin"); + assert!(!broker.virtualize_text(&mut hashed_credential_path, &env)); + assert_eq!(hashed_credential_path, "/workspace/ta/bin"); + let mut watermarked_path = format!("/workspace/tokensk-proj-{hash}-T3BlbkFJsuffix/bin"); + assert!(!broker.virtualize_text(&mut watermarked_path, &env)); + assert_eq!(watermarked_path, "/workspace/token/bin"); + for suffix in ["build-", "", "proj-", "admin-", "svcacct-"] { + let path_value = + format!("/task-{suffix}{hash}:/task-{suffix}{hash}/x:/task-{suffix}{hash}"); + for path in [ + path_value, + format!("declare -x PATH=\"/task-{suffix}{hash}-build:/task-{suffix}{hash}-release\""), + format!("export -UT PATH path=(/task-{suffix}{hash} /usr/bin)"), + format!("alias activate='source /task-{suffix}{hash}/bin/activate'"), + format!(r"C:\task-{suffix}{hash}\Scripts"), + ] { + let mut virtualized_path = path.clone(); + assert!(broker.virtualize_text(&mut virtualized_path, &env)); + assert_eq!(virtualized_path, path); + } + } + for component in [ + "my_task", + "flask", + "disk", + "mask", + "risk", + "brisk", + "subtask", + "mytask", + "devtask", + "multitask", + "buildtask", + "mydisk", + "harddisk", + "MY_Task", + "Flask", + "\u{e9}_task", + "e\u{301}_task", + ] { + for suffix in ["", "proj-", "admin-", "svcacct-"] { + for path in [ + format!("/workspace/{component}-{suffix}{hash}/bin"), + format!("VIRTUAL_ENV=/workspace/{component}-{suffix}{hash}"), + format!( + "alias activate='source /workspace/{component}-{suffix}{hash}/bin/activate'" + ), + format!(r"C:\\workspace\\{component}-{suffix}{hash}\\Scripts"), + ] { + let mut virtualized_path = path.clone(); + assert!(broker.virtualize_text(&mut virtualized_path, &env)); + assert_eq!(virtualized_path, path); + } + } + } + let registered_hex_credential = format!("sk-{hash}"); + let registered_broker = CredentialBroker::new(/*enabled*/ true); + let mut registered_env = env_map([("OPENAI_API_KEY", registered_hex_credential.as_str())]); + registered_broker.virtualize_child_env(&mut registered_env); + let registered_dummy = ®istered_env["OPENAI_API_KEY"]; + let mut credential_path = format!("/workspace/multita{registered_hex_credential}/bin"); + assert!(registered_broker.virtualize_text(&mut credential_path, ®istered_env)); + assert_eq!( + credential_path, + format!("/workspace/multita{registered_dummy}/bin") + ); + for (key, placeholder, credential) in [ + ("GH_TOKEN", "ghp_", unknown_github_token), + ("OPENAI_API_KEY", "sk-", unknown_legacy_openai_key.as_str()), + ] { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([(key, placeholder)]); + broker.virtualize_child_env(&mut env); + let mut alias = format!("Bearer {credential}"); + assert!(!broker.virtualize_text(&mut alias, &env)); + assert!(!alias.contains(credential)); + } + let mut command = vec![ + format!("Authorization: Bearer {github_dummy}"), + format!("Authorization: Bearer {openai_dummy}"), + ]; + let github_dummy = github_dummy.clone(); + let openai_dummy = openai_dummy.clone(); + env.insert("OPENAI_API_KEY".to_string(), "sk-user-override".to_string()); + env.insert( + "GIT_CONFIG_VALUE_0".to_string(), + format!("Authorization: Bearer {github_dummy}"), + ); + assert_eq!( + brokered_credential_dummy_env_keys(&env), + vec!["GH_TOKEN".to_string()] + ); + + broker.restore_child_env(&mut env, &mut command); + assert_eq!(env.get("GH_TOKEN").map(String::as_str), Some(github_token)); + assert_eq!( + env.get("HOMEBREW_GITHUB_API_TOKEN").map(String::as_str), + Some(github_token) + ); + assert_eq!( + env.get("GH_ENTERPRISE_TOKEN").map(String::as_str), + Some(github_token) + ); + assert_eq!(env.get("AUTH_HEADER"), Some(&authorization)); + assert_eq!( + env.get("OPENAI_API_KEY").map(String::as_str), + Some("sk-user-override") + ); + assert_eq!( + env.get("GIT_CONFIG_VALUE_0"), + Some(&format!("Authorization: Bearer {github_dummy}")) + ); + assert_eq!( + command, + vec![ + format!("Authorization: Bearer {github_dummy}"), + format!("Authorization: Bearer {openai_dummy}"), + ] + ); + + env.insert("GH_TOKEN".to_string(), openai_dummy.clone()); + env.insert("OPENAI_API_KEY".to_string(), github_dummy.clone()); + broker.restore_child_env(&mut env, &mut []); + assert_eq!(env.get("GH_TOKEN"), Some(&openai_dummy)); + assert_eq!(env.get("OPENAI_API_KEY"), Some(&github_dummy)); +} + +#[test] +fn unsupported_children_restore_credentials_and_disable_brokerage() { + let broker = CredentialBroker::new(/*enabled*/ true); + let github_token = "ghp_abcdefghijklmnopqrstuvwxyz1234567890"; + let mut env = env_map([("GH_TOKEN", github_token)]); + broker.virtualize_child_env(&mut env); + + assert_ne!(env.get("GH_TOKEN").map(String::as_str), Some(github_token)); + assert_eq!( + env.get(CREDENTIAL_BROKER_ACTIVE_ENV_KEY) + .map(String::as_str), + Some("1") + ); + assert!(env.contains_key(BROKERED_CREDENTIALS_ENV_KEY)); + + broker.restore_and_disable_child_env(&mut env, &mut []); + + assert_eq!(env.get("GH_TOKEN").map(String::as_str), Some(github_token)); + assert!(!env.contains_key(CREDENTIAL_BROKER_ACTIVE_ENV_KEY)); + assert!(!env.contains_key(BROKERED_CREDENTIALS_ENV_KEY)); +} + +#[cfg(windows)] +#[test] +fn brokered_credentials_match_environment_keys_case_insensitively_on_windows() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([ + ("gh_host", "github.example.com"), + ("gh_enterprise_token", "ghp-enterprise-real"), + ]); + + broker.virtualize_child_env(&mut env); + let dummy = env + .get("GH_ENTERPRISE_TOKEN") + .expect("dummy GitHub enterprise token"); + let mut headers = headers_with_bearer(dummy); + broker.inject_request_headers("github.example.com", &mut headers); + + assert_eq!( + brokered_credential_dummy_env_keys(&env), + vec!["GH_ENTERPRISE_TOKEN".to_string()] + ); + assert_eq!(authorization(&headers), Some("Bearer ghp-enterprise-real")); +} + +#[test] +fn virtualize_child_env_preserves_live_dummy_mappings() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut first_env = env_map([("GH_TOKEN", "ghp-real-one")]); + let mut second_env = env_map([("GH_TOKEN", "ghp-real-two")]); + + broker.virtualize_child_env(&mut first_env); + broker.virtualize_child_env(&mut second_env); + let first_dummy = first_env.get("GH_TOKEN").expect("first dummy token"); + let second_dummy = second_env.get("GH_TOKEN").expect("second dummy token"); + let mut first_headers = headers_with_bearer(first_dummy); + let mut second_headers = headers_with_bearer(second_dummy); + + broker.inject_request_headers("api.github.com", &mut first_headers); + broker.inject_request_headers("api.github.com", &mut second_headers); + + assert_eq!(authorization(&first_headers), Some("Bearer ghp-real-one")); + assert_eq!(authorization(&second_headers), Some("Bearer ghp-real-two")); + + let mut alias_only = env_map([("HOMEBREW_GITHUB_API_TOKEN", "ghp-real-one")]); + broker.virtualize_child_env(&mut alias_only); + assert_eq!( + alias_only.get("HOMEBREW_GITHUB_API_TOKEN"), + Some(first_dummy) + ); + broker.restore_child_env(&mut alias_only, &mut []); + assert_eq!(alias_only["HOMEBREW_GITHUB_API_TOKEN"], "ghp-real-one"); + + let mut overridden = env_map([ + ("GH_TOKEN", "ghp-real-two"), + ("HOMEBREW_GITHUB_API_TOKEN", "ghp-real-one"), + ]); + broker.virtualize_child_env(&mut overridden); + assert_eq!(overridden.get("GH_TOKEN"), Some(second_dummy)); + assert_eq!( + overridden.get("HOMEBREW_GITHUB_API_TOKEN"), + Some(first_dummy) + ); + + let mut cloud_alias = env_map([ + ("GH_TOKEN", "ghp-real-one"), + ("GITHUB_TOKEN", "ghp-real-one"), + ]); + broker.virtualize_child_env(&mut cloud_alias); + cloud_alias.insert("GITHUB_TOKEN".to_string(), first_dummy.clone()); + broker.restore_child_env(&mut cloud_alias, &mut []); + assert_eq!(cloud_alias["GITHUB_TOKEN"], "ghp-real-one"); + + let mut distinct_credentials = env_map([ + ("GH_TOKEN", "ghp-primary-secret"), + ("GITHUB_TOKEN", "ghp-secondary-secret"), + ]); + broker.virtualize_child_env(&mut distinct_credentials); + let secondary_dummy = distinct_credentials["GITHUB_TOKEN"].clone(); + distinct_credentials.remove("GITHUB_TOKEN"); + distinct_credentials.insert("GH_TOKEN".to_string(), secondary_dummy.clone()); + broker.virtualize_child_env(&mut distinct_credentials); + broker.restore_child_env(&mut distinct_credentials, &mut []); + assert_eq!(distinct_credentials["GH_TOKEN"], secondary_dummy); +} + +#[test] +fn unbound_enterprise_aliases_retain_source_ownership() { + let token = "ghp_abcdefghijklmnopqrstuvwxyz0123456789"; + let real_header = format!("Bearer {token}"); + for source in ["GH_ENTERPRISE_TOKEN", "GITHUB_ENTERPRISE_TOKEN"] { + for host in [None, Some("")] { + for parent_discovery in [false, true] { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut parent = env_map([(source, token)]); + if let Some(host) = host { + parent.insert("GH_HOST".to_string(), host.to_string()); + } + let mut child = env_map([("AUTH_HEADER", real_header.as_str())]); + if parent_discovery { + broker.discover_parent_credentials(&parent, &child); + } else { + broker.virtualize_child_env(&mut parent); + } + broker.virtualize_child_env(&mut child); + assert_eq!(child["AUTH_HEADER"], real_header, "{source}, {host:?}"); + assert!(broker.read_state().credentials.is_empty()); + + // A later explicit source and destination can still register normally. + child.insert(source.to_string(), token.to_string()); + child.insert("GH_HOST".to_string(), "enterprise.example".to_string()); + broker.virtualize_child_env(&mut child); + let dummy = &child[source]; + assert_ne!(dummy, token); + let mut headers = headers_with_bearer(dummy); + let original = headers.clone(); + broker.inject_request_headers("api.github.com", &mut headers); + assert_eq!(headers, original); + broker.inject_request_headers("enterprise.example", &mut headers); + assert_eq!(authorization(&headers), Some(real_header.as_str())); + } + } + } +} + +#[test] +fn virtualize_child_env_replaces_aliases_of_filtered_parent_credentials() { + let broker = CredentialBroker::new(/*enabled*/ true); + let github_token = "ghp_abcdefghijklmnopqrstuvwxyz1234567890"; + let authorization_header = format!("Bearer {github_token}"); + let parent_env = env_map([ + ("GH_TOKEN", github_token), + ("HOMEBREW_GITHUB_API_TOKEN", github_token), + ]); + let mut child_env = env_map([ + ("HOMEBREW_GITHUB_API_TOKEN", github_token), + ("AUTH_HEADER", authorization_header.as_str()), + ]); + + broker.discover_parent_credentials(&parent_env, &child_env); + broker.virtualize_child_env(&mut child_env); + + let dummy = child_env["HOMEBREW_GITHUB_API_TOKEN"].clone(); + assert_ne!(dummy, github_token); + assert_eq!(child_env["AUTH_HEADER"], format!("Bearer {dummy}")); + assert!(!child_env.contains_key("GH_TOKEN")); + + let mut virtualized_alias = child_env["AUTH_HEADER"].clone(); + assert!(broker.virtualize_text(&mut virtualized_alias, &child_env)); + assert_eq!(virtualized_alias, child_env["AUTH_HEADER"]); + + let mut headers = headers_with_bearer(&dummy); + broker.inject_request_headers("api.github.com", &mut headers); + assert_eq!(authorization(&headers), Some(authorization_header.as_str())); + + broker.restore_child_env(&mut child_env, &mut []); + assert_eq!(child_env["HOMEBREW_GITHUB_API_TOKEN"], github_token); + assert_eq!(child_env["AUTH_HEADER"], authorization_header); + assert!(!child_env.contains_key("GH_TOKEN")); + + let openai_token = "sk-proj-abcdefghijklmnopqrstuvwxyz1234567890"; + let mixed_bundle = format!("GitHub {github_token}\nOpenAI {openai_token}"); + let mut mixed_env = env_map([ + ("GH_TOKEN", github_token), + ("OPENAI_API_KEY", openai_token), + ("AUTH_BUNDLE", mixed_bundle.as_str()), + ]); + broker.virtualize_child_env(&mut mixed_env); + let mut mixed_alias = mixed_env["AUTH_BUNDLE"].clone(); + let excluded_dummy = mixed_env.remove("OPENAI_API_KEY").expect("OpenAI dummy"); + assert!(!broker.virtualize_text(&mut mixed_alias, &mixed_env)); + assert!(!mixed_alias.contains(&excluded_dummy)); +} + +#[test] +fn virtualize_child_env_preserves_paths_unless_the_credential_is_known() { + let broker = CredentialBroker::new(/*enabled*/ true); + let real = format!("sk-proj-{}", "a".repeat(64)); + let path = format!("/workspace/my_ta{real}/bin"); + let mut env = env_map([("VIRTUAL_ENV", &path)]); + + broker.virtualize_child_env(&mut env); + + assert_eq!( + env, + env_map([ + ("VIRTUAL_ENV", &path), + (CREDENTIAL_BROKER_ACTIVE_ENV_KEY, "1"), + (BROKERED_CREDENTIALS_ENV_KEY, "[]"), + ]) + ); + + let copied = format!("ta{real}"); + let mut copied_env = env_map([("COPIED", &copied)]); + broker.virtualize_child_env(&mut copied_env); + let dummy = copied_env["COPIED"].strip_prefix("ta").unwrap(); + assert_credential_shape(&real, dummy, "sk-proj-"); + let mut headers = headers_with_bearer(dummy); + broker.inject_request_headers("api.openai.com", &mut headers); + assert_eq!( + authorization(&headers), + Some(format!("Bearer {real}").as_str()) + ); + + broker.virtualize_child_env(&mut env); + assert_eq!(env["VIRTUAL_ENV"], path.replace(&real, dummy)); +} + +#[test] +fn virtualize_child_env_discovers_credentials_without_canonical_variables() { + for (token, canonical_key, host) in [ + ( + "ghp_abcdefghijklmnopqrstuvwxyz1234567890", + "GH_TOKEN", + "api.github.com", + ), + ( + "sk-proj-abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789", + "OPENAI_API_KEY", + "api.openai.com", + ), + ] { + let broker = CredentialBroker::new(/*enabled*/ true); + let authorization_header = format!("Bearer {token}"); + let mut env = env_map([("AUTH_HEADER", authorization_header.as_str())]); + + broker.virtualize_child_env(&mut env); + + let dummy_header = &env["AUTH_HEADER"]; + assert_ne!(dummy_header, &authorization_header); + assert!(!env.contains_key(canonical_key)); + let mut headers = headers_with_authorization(dummy_header); + broker.inject_request_headers(host, &mut headers); + assert_eq!(authorization(&headers), Some(authorization_header.as_str())); + + let mut snapshot = format!("export AUTH_HEADER='{authorization_header}'"); + assert!(broker.virtualize_text(&mut snapshot, &env)); + assert_eq!(snapshot, format!("export AUTH_HEADER='{dummy_header}'")); + + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([ + (canonical_key, "another-canonical-credential"), + ("AUTH_HEADER", authorization_header.as_str()), + ]); + broker.virtualize_child_env(&mut env); + assert_ne!(env["AUTH_HEADER"], authorization_header); + let mut headers = headers_with_authorization(&env["AUTH_HEADER"]); + broker.inject_request_headers(host, &mut headers); + assert_eq!(authorization(&headers), Some(authorization_header.as_str())); + } +} + +#[test] +fn virtualize_child_env_preserves_operational_paths_during_credential_discovery() { + let broker = CredentialBroker::new(/*enabled*/ true); + let hash = "a".repeat(64); + let virtual_env = format!("/workspace/multitask-proj-{hash}/bin"); + let github_host = format!("multitask-proj-{hash}.enterprise.example"); + let token = "sk-proj-abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"; + let authorization_header = format!("Bearer {token}"); + let mut env = env_map([ + ("VIRTUAL_ENV", virtual_env.as_str()), + ("GH_HOST", github_host.as_str()), + ("NO_PROXY", github_host.as_str()), + ("AUTH_HEADER", authorization_header.as_str()), + ]); + + broker.virtualize_child_env(&mut env); + + assert_eq!(env["VIRTUAL_ENV"], virtual_env); + assert_eq!(env["GH_HOST"], github_host); + assert_eq!(env["NO_PROXY"], github_host); + assert_ne!(env["AUTH_HEADER"], authorization_header); +} + +#[test] +fn virtualize_child_env_keeps_adjacent_provider_credentials_separate() { + let broker = CredentialBroker::new(/*enabled*/ true); + let github_token = "ghp_abcdefghijklmnopqrstuvwxyz1234567890ABCD"; + let openai_token = "sk-proj-abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"; + let real_bundle = format!("{github_token}_{openai_token}"); + let mut env = env_map([("AUTH_BUNDLE", real_bundle.as_str())]); + + broker.virtualize_child_env(&mut env); + + let dummy_bundle = env["AUTH_BUNDLE"].clone(); + assert!(!dummy_bundle.contains(github_token)); + assert!(!dummy_bundle.contains(openai_token)); + let github_dummy = &dummy_bundle[..github_token.len()]; + let openai_dummy = &dummy_bundle[github_token.len() + 1..]; + + let mut github_headers = headers_with_bearer(github_dummy); + broker.inject_request_headers("api.github.com", &mut github_headers); + assert_eq!( + authorization(&github_headers), + Some(format!("Bearer {github_token}").as_str()) + ); + + let mut openai_headers = headers_with_bearer(openai_dummy); + broker.inject_request_headers("api.openai.com", &mut openai_headers); + assert_eq!( + authorization(&openai_headers), + Some(format!("Bearer {openai_token}").as_str()) + ); + + let mut bundled_headers = headers_with_bearer(&dummy_bundle); + broker.inject_request_headers("api.github.com", &mut bundled_headers); + assert_eq!( + authorization(&bundled_headers), + Some(format!("Bearer {dummy_bundle}").as_str()) + ); + + let mut restored_bundle = dummy_bundle; + assert!(broker.restore_text(&mut restored_bundle)); + assert_eq!(restored_bundle, real_bundle); +} + +#[test] +fn virtualize_child_env_binds_filtered_enterprise_credentials_to_child_host() { + let github_token = "ghp_abcdefghijklmnopqrstuvwxyz1234567890"; + let authorization_header = format!("Bearer {github_token}"); + + for (parent_host, include_cloud_token) in [ + (None, false), + (Some("github.previous.example"), false), + (None, true), + ] { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut parent_env = env_map([("GH_ENTERPRISE_TOKEN", github_token)]); + if let Some(parent_host) = parent_host { + parent_env.insert("GH_HOST".to_string(), parent_host.to_string()); + } + if include_cloud_token { + parent_env.insert("GH_TOKEN".to_string(), github_token.to_string()); + } + let mut child_env = env_map([ + ("GH_HOST", "github.current.example"), + ("AUTH_HEADER", authorization_header.as_str()), + ]); + + broker.discover_parent_credentials(&parent_env, &child_env); + broker.virtualize_child_env(&mut child_env); + + assert_ne!(child_env["AUTH_HEADER"], authorization_header); + let mut headers = headers_with_authorization(&child_env["AUTH_HEADER"]); + broker.inject_request_headers("github.current.example", &mut headers); + assert_eq!(authorization(&headers), Some(authorization_header.as_str())); + + let mut previous_headers = headers_with_authorization(&child_env["AUTH_HEADER"]); + broker.inject_request_headers("github.previous.example", &mut previous_headers); + assert_eq!( + authorization(&previous_headers), + Some(child_env["AUTH_HEADER"].as_str()) + ); + } +} + +#[test] +fn brokered_credential_env_keys_only_include_registered_credentials() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([ + ("OPENAI_API_KEY", "sk-real"), + ("GH_TOKEN", ""), + ("GH_HOST", "github.example.com"), + ]); + + broker.virtualize_child_env(&mut env); + env.insert( + "GH_TOKEN".to_string(), + "ghp_added_after_brokerage".to_string(), + ); + + assert_eq!( + brokered_credential_env_keys(&env).collect::>(), + vec!["OPENAI_API_KEY"] + ); +} + +#[test] +fn brokered_credential_value_env_keys_include_dummy_aliases() { + let broker = CredentialBroker::new(/*enabled*/ true); + let real = "sk-proj-abcdefghijklmnopqrstuvwxyz"; + let alias = format!("Bearer {real}"); + let mut env = env_map([("OPENAI_API_KEY", real), ("AUTH_HEADER", alias.as_str())]); + + broker.virtualize_child_env(&mut env); + + let marker: Vec<(String, String)> = + serde_json::from_str(&env[BROKERED_CREDENTIALS_ENV_KEY]).unwrap(); + assert_eq!( + marker, + vec![ + ( + "@alias:AUTH_HEADER".to_string(), + env["OPENAI_API_KEY"].clone() + ), + ( + "@alias:OPENAI_API_KEY".to_string(), + env["OPENAI_API_KEY"].clone() + ), + ("OPENAI_API_KEY".to_string(), env["OPENAI_API_KEY"].clone()), + ] + ); + + assert_eq!( + brokered_credential_value_env_keys(&env), + vec!["AUTH_HEADER".to_string(), "OPENAI_API_KEY".to_string()] + ); + + let mut absent_alias_env = env.clone(); + absent_alias_env.remove("AUTH_HEADER"); + broker.virtualize_child_env(&mut absent_alias_env); + assert_eq!( + brokered_credential_marker_env_keys(&absent_alias_env), + vec!["AUTH_HEADER".to_string(), "OPENAI_API_KEY".to_string()] + ); + + env.remove("OPENAI_API_KEY"); + broker.virtualize_child_env(&mut env); + assert!(brokered_credential_dummy_env_keys(&env).is_empty()); + assert_eq!( + brokered_credential_marker_env_keys(&env), + vec!["AUTH_HEADER".to_string(), "OPENAI_API_KEY".to_string()] + ); + assert_eq!( + brokered_credential_value_env_keys(&env), + vec!["AUTH_HEADER".to_string()] + ); +} + +#[test] +fn virtualize_child_env_uses_fresh_dummy_capabilities() { + let mut first_env = env_map([("OPENAI_API_KEY", "sk-proj-abcdefghijklmnopqrstuvwxyz")]); + let mut second_env = first_env.clone(); + + CredentialBroker::new(/*enabled*/ true).virtualize_child_env(&mut first_env); + CredentialBroker::new(/*enabled*/ true).virtualize_child_env(&mut second_env); + + assert_ne!(first_env["OPENAI_API_KEY"], second_env["OPENAI_API_KEY"]); +} + +#[test] +fn child_without_dummy_cannot_use_previous_child_credential() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut first_env = env_map([("OPENAI_API_KEY", "sk-real")]); + let mut second_env = HashMap::new(); + + broker.virtualize_child_env(&mut first_env); + broker.virtualize_child_env(&mut second_env); + let mut headers = HeaderMap::new(); + + broker.inject_request_headers("api.openai.com", &mut headers); + + assert_eq!(authorization(&headers), None); +} + +#[test] +fn virtualize_child_env_keeps_unbound_enterprise_token_out_of_persisted_text() { + let broker = CredentialBroker::new(/*enabled*/ true); + let token = "ghp_abcdefghijklmnopqrstuvwxyz1234567890"; + let authorization_header = format!("Bearer {token}"); + let mut env = env_map([ + ("GH_ENTERPRISE_TOKEN", token), + ("AUTH_HEADER", authorization_header.as_str()), + ]); + + broker.virtualize_child_env(&mut env); + assert_eq!(env["GH_ENTERPRISE_TOKEN"], token); + assert_eq!(env["AUTH_HEADER"], authorization_header); + for alias in [ + format!("export GH_ENTERPRISE_TOKEN={token}"), + format!("export AUTH_HEADER='Bearer {token}_suffix'"), + format!("export AUTH_HEADER='Bearer {token}-suffix'"), + ] { + let mut persisted = alias; + assert!(!broker.virtualize_text(&mut persisted, &env)); + assert!(!persisted.contains(token)); + } + let distinct_token = "ghp_ABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789abcdefghijkl"; + let mut adjacent = format!("{token}_{distinct_token}"); + assert!(!broker.virtualize_text(&mut adjacent, &env)); + assert_eq!(adjacent, "_"); + + let mut truncated_env = env_map([("GH_ENTERPRISE_TOKEN", "ghp_abcdefghijkl")]); + broker.virtualize_child_env(&mut truncated_env); + let mut hidden = token.to_string(); + assert!(!broker.virtualize_text(&mut hidden, &truncated_env)); + assert!(hidden.is_empty()); + assert!(!credential_broker_provider_sources_allowed( + token, + "", + &truncated_env, + |source| source != "GH_TOKEN", + )); + assert!(!credential_broker_provider_sources_allowed( + token, + "", + &HashMap::new(), + |source| source != "GH_TOKEN", + )); + + let fine_grained = "github_pat_11AA0bbCC_abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGH"; + let mut truncated_env = env_map([( + "GH_ENTERPRISE_TOKEN", + "github_pat_11AA0bbCC_abcdefghijklmnopqrs", + )]); + broker.virtualize_child_env(&mut truncated_env); + let mut hidden = fine_grained.to_string(); + assert!(!broker.virtualize_text(&mut hidden, &truncated_env)); + assert!(hidden.is_empty()); + assert!(!credential_broker_provider_sources_allowed( + fine_grained, + "", + &env_map([("GH_TOKEN", "github_pat_11AA0bbCC")]), + |source| source == "GH_TOKEN", + )); + let mut headers = headers_with_bearer(token); + broker.inject_request_headers("attacker.example", &mut headers); + + assert_eq!(env["GH_ENTERPRISE_TOKEN"], token); + assert_eq!(headers, headers_with_bearer(token)); + assert!(!broker.host_requires_mitm("attacker.example", /*port*/ 443)); + + env.insert("GH_HOST".to_string(), "github.example.com".to_string()); + broker.virtualize_child_env(&mut env); + let mut headers = headers_with_bearer(&env["GH_ENTERPRISE_TOKEN"]); + broker.inject_request_headers("github.example.com", &mut headers); + assert_eq!( + authorization(&headers), + Some(format!("Bearer {token}").as_str()) + ); +} + +#[test] +fn inject_request_headers_requires_dummy_to_select_ambiguous_github_credential() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([ + ("GH_TOKEN", "ghp-real-one"), + ("GITHUB_TOKEN", "ghp-real-two"), + ]); + broker.virtualize_child_env(&mut env); + let github_token = env.get("GITHUB_TOKEN").expect("dummy github token"); + let mut headers = HeaderMap::new(); + + broker.inject_request_headers("api.github.com", &mut headers); + assert_eq!(authorization(&headers), None); + + headers = headers_with_bearer(github_token); + + broker.inject_request_headers("api.github.com", &mut headers); + + assert_eq!(authorization(&headers), Some("Bearer ghp-real-two")); +} + +#[test] +fn request_translation_preserves_provider_scheme_and_host_binding() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([("GH_TOKEN", "ghp-real")]); + broker.virtualize_child_env(&mut env); + let gh = &env["GH_TOKEN"]; + let basic_dummy = STANDARD.encode(format!("x-access-token:{gh}")); + let basic_real = STANDARD.encode("x-access-token:ghp-real"); + let basic_username_dummy = STANDARD.encode(format!("{gh}:x-oauth-basic")); + let basic_username_real = STANDARD.encode("ghp-real:x-oauth-basic"); + let basic_dummy = basic_dummy.as_str(); + let basic_real = basic_real.as_str(); + let basic_username_dummy = basic_username_dummy.as_str(); + let basic_username_real = basic_username_real.as_str(); + + for (host, scheme, input, expected) in [ + ("github.com", "Basic", basic_dummy, basic_real), + ("example.com", "Basic", basic_dummy, basic_dummy), + ( + "github.com", + "Basic", + basic_username_dummy, + basic_username_real, + ), + ( + "example.com", + "Basic", + basic_username_dummy, + basic_username_dummy, + ), + ("api.github.com", "Bearer", gh.as_str(), "ghp-real"), + ("uploads.github.com", "Bearer", gh.as_str(), "ghp-real"), + ("api.github.com", "token", gh.as_str(), "ghp-real"), + ] { + let mut headers = headers_with_authorization(&format!("{scheme} {input}")); + broker.inject_request_headers(host, &mut headers); + let expected = format!("{scheme} {expected}"); + assert_eq!(authorization(&headers), Some(expected.as_str()), "{host}"); + } +} + +#[test] +fn inject_request_headers_requires_dummy_and_preserves_explicit_authorization() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([("OPENAI_API_KEY", "sk-real")]); + broker.virtualize_child_env(&mut env); + let openai_api_key = env.get("OPENAI_API_KEY").expect("dummy OpenAI API key"); + let mut headers = HeaderMap::new(); + + broker.inject_request_headers("api.openai.com", &mut headers); + assert_eq!(authorization(&headers), None); + + headers = headers_with_bearer(openai_api_key); + broker.inject_request_headers("api.openai.com", &mut headers); + assert_eq!(authorization(&headers), Some("Bearer sk-real")); + + let mut explicit_headers = headers_with_bearer("sk-explicit"); + broker.inject_request_headers("api.openai.com", &mut explicit_headers); + + assert_eq!(authorization(&explicit_headers), Some("Bearer sk-explicit")); +} + +#[test] +fn concurrent_commands_preserve_discovered_credential_destinations() { + for (key, host_key, token) in [ + ( + "GH_ENTERPRISE_TOKEN", + "GH_HOST", + "ghp_abcdefghijklmnopqrstuvwxyz1234567890", + ), + ( + "OPENAI_API_KEY", + "OPENAI_BASE_URL", + "sk-proj-abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789", + ), + ( + "PROVIDER_TOKEN", + "PROVIDER_ENDPOINT", + "provider_abcdefghijklmnopqrstuvwx", + ), + ] { + let broker = CredentialBroker::new(/*enabled*/ true); + broker.configure(&NetworkProxyConfig { + credential_broker: true, + credential_providers: BTreeMap::from([( + "custom".to_string(), + CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["^provider_[a-z]{24}$".to_string()], + url_prefix_from_env: Some("PROVIDER_ENDPOINT".to_string()), + ..CredentialProviderConfig::default() + }, + )]), + ..NetworkProxyConfig::default() + }); + let mut commands = ["first.example", "second.example"].map(|host| { + let destination = if host_key == "GH_HOST" { + host.to_string() + } else { + format!("https://{host}/v1") + }; + env_map([ + (key, token), + (host_key, &destination), + ("AUTH_HEADER", &format!("Bearer {token}")), + ]) + }); + for env in &mut commands { + broker.virtualize_child_env_for_environment(env, Some("shared-environment")); + } + let dummy = commands[0][key].clone(); + assert_ne!(dummy, token); + assert_eq!(commands[1][key], dummy); + let assert_destinations = || { + for (environment_id, host, injected) in [ + ("shared-environment", "first.example", true), + ("shared-environment", "second.example", true), + ("shared-environment", "unrelated.example", false), + ("other-environment", "first.example", false), + ("other-environment", "second.example", false), + ] { + let mut headers = headers_with_bearer(&dummy); + broker.inject_request_headers_for_environment( + &format!("https://{host}/v1/models"), + &mut headers, + Some(environment_id), + ); + assert_eq!( + headers, + headers_with_bearer(if injected { token } else { &dummy }), + "{key}, {environment_id}, {host}" + ); + assert_eq!( + broker + .host_protocols_for_environment( + host, + /*port*/ 443, + Some(environment_id), + ) + .tls, + injected, + "{key}, {environment_id}, {host}" + ); + } + }; + assert_destinations(); + for env in &mut commands { + env.remove(key); + broker.virtualize_child_env_for_environment(env, Some("shared-environment")); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {dummy}")); + assert_destinations(); + } + + commands[0].insert(key.to_string(), token.to_string()); + broker.virtualize_child_env_for_environment(&mut commands[0], Some("shared-environment")); + assert_eq!(commands[0][key], dummy); + assert_destinations(); + + let mut inherited = env_map([(key, token), (host_key, &commands[0][host_key])]); + broker.virtualize_child_env_for_environment(&mut inherited, Some("parent-environment")); + let parent_dummy = inherited[key].clone(); + assert_ne!(parent_dummy, dummy); + broker.virtualize_child_env_for_environment(&mut inherited, Some("shared-environment")); + assert_eq!(inherited[key], dummy); + assert_destinations(); + let mut headers = headers_with_bearer(&parent_dummy); + broker.inject_request_headers_for_environment( + "https://second.example/v1/models", + &mut headers, + Some("parent-environment"), + ); + assert_eq!(headers, headers_with_bearer(&parent_dummy)); + } +} + +#[test] +fn builtin_credentials_use_private_destination_context() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut config = NetworkProxyConfig::default(); + config.set_credential_broker_enabled(/*enabled*/ true); + config.configure_credential_broker_environment(&env_map([ + ("GH_HOST", "github.enterprise.example"), + ("OPENAI_BASE_URL", "https://gateway.example/v1"), + ])); + broker.configure(&config); + for (key, context_key, token, host) in [ + ( + "GH_ENTERPRISE_TOKEN", + "GH_HOST", + "ghp-real", + "github.enterprise.example", + ), + ( + "OPENAI_API_KEY", + "OPENAI_BASE_URL", + "sk-real", + "gateway.example", + ), + ] { + let mut env = env_map([(key, token)]); + broker.virtualize_child_env(&mut env); + assert_ne!(env[key], token); + assert!(!env.contains_key("GH_HOST")); + assert!(!env.contains_key("OPENAI_BASE_URL")); + let mut headers = headers_with_bearer(&env[key]); + broker.inject_request_headers(host, &mut headers); + assert_eq!( + authorization(&headers), + Some(format!("Bearer {token}").as_str()) + ); + + let snapshot_destination = if context_key == "GH_HOST" { + "snapshot.example" + } else { + "https://snapshot.example/v1" + }; + env.insert(context_key.to_string(), snapshot_destination.to_string()); + broker.virtualize_child_env(&mut env); + env.remove(context_key); + env.insert( + "NEW_AUTH_HEADER".to_string(), + format!("Bearer {}", env[key]), + ); + env.insert("REAL_COPY".to_string(), token.to_string()); + let dummy = env[key].clone(); + broker.virtualize_child_env(&mut env); + assert_eq!((&env[key], &env["REAL_COPY"]), (&dummy, &dummy)); + assert!(!env.contains_key(context_key)); + for destination in ["snapshot.example", host] { + let mut headers = headers_with_bearer(&dummy); + broker.inject_request_headers(destination, &mut headers); + assert_eq!( + authorization(&headers), + Some(format!("Bearer {token}").as_str()) + ); + } + let mut inherited_env = env.clone(); + broker.virtualize_child_env_for_environment(&mut inherited_env, Some("child")); + assert_eq!(inherited_env[key], dummy); + for (destination, injected) in [("snapshot.example", true), (host, false)] { + let mut headers = headers_with_bearer(&dummy); + broker.inject_request_headers_for_environment(destination, &mut headers, Some("child")); + assert_eq!( + headers, + headers_with_bearer(if injected { token } else { &dummy }), + ); + } + broker.restore_child_env(&mut env, &mut []); + assert_eq!(env["NEW_AUTH_HEADER"], format!("Bearer {token}")); + } +} + +#[test] +fn private_destination_updates_reconcile_registered_fallbacks() { + for ((key, context_key, token), previously_used_fallback) in [ + ("GH_ENTERPRISE_TOKEN", "GH_HOST", "ghp-real"), + ("OPENAI_API_KEY", "OPENAI_BASE_URL", "sk-real"), + ( + "VENDOR_TOKEN", + "VENDOR_HOST", + "vendor_abcdefghijklmnopqrstuvwx", + ), + ] + .into_iter() + .flat_map(|source| [(source, false), (source, true)]) + { + let destination = |host: &str| { + if context_key == "GH_HOST" { + host.to_string() + } else { + format!("https://{host}") + } + }; + let broker = CredentialBroker::new(/*enabled*/ true); + let mut config = NetworkProxyConfig { + credential_broker: true, + credential_providers: BTreeMap::from([( + "vendor".to_string(), + CredentialProviderConfig { + env: vec!["VENDOR_TOKEN".to_string()], + patterns: vec!["^vendor_[a-z]{24}$".to_string()], + url_prefix_from_env: Some("VENDOR_HOST".to_string()), + ..CredentialProviderConfig::default() + }, + )]), + ..NetworkProxyConfig::default() + }; + config.configure_credential_broker_environment(&env_map([( + context_key, + &destination("first.example"), + )])); + broker.configure(&config); + let mut env = env_map([(key, token)]); + broker.virtualize_child_env(&mut env); + let dummy = env[key].clone(); + assert_ne!(dummy, token); + + let mut explicit_env = env_map([(key, token)]); + if previously_used_fallback { + broker.virtualize_child_env_for_environment(&mut explicit_env, Some("explicit")); + } + explicit_env.insert( + context_key.to_string(), + destination("other-explicit.example"), + ); + broker.virtualize_child_env_for_environment(&mut explicit_env, Some("explicit")); + explicit_env.insert(context_key.to_string(), destination("explicit.example")); + broker.virtualize_child_env_for_environment(&mut explicit_env, Some("explicit")); + explicit_env.remove(context_key); + let explicit_dummy = explicit_env[key].clone(); + let assert_inherited_destination = |child: &str| { + let mut inherited_env = explicit_env.clone(); + inherited_env.insert("CREDENTIAL_COPY".to_string(), token.to_string()); + broker.virtualize_child_env_for_environment(&mut inherited_env, Some(child)); + assert_eq!( + inherited_env["CREDENTIAL_COPY"], explicit_dummy, + "{key}: {child}" + ); + for host in ["explicit.example", "first.example", "second.example"] { + let mut headers = headers_with_bearer(&explicit_dummy); + broker.inject_request_headers_for_environment( + &format!("https://{host}/v1"), + &mut headers, + Some(child), + ); + assert_eq!( + headers, + headers_with_bearer(if host == "explicit.example" { + token + } else { + &explicit_dummy + }), + "{key}: {child}, {host}" + ); + } + }; + let unrelated_key = if context_key == "GH_HOST" { + "OPENAI_BASE_URL" + } else { + "GH_HOST" + }; + config.configure_credential_broker_environment(&env_map([ + (context_key, &destination("first.example")), + (unrelated_key, "https://unrelated.example"), + ])); + broker.configure(&config); + assert_inherited_destination("unchanged-fallback-child"); + + config.configure_credential_broker_environment(&env_map([( + context_key, + &destination("second.example"), + )])); + let revision = broker.config_revision(); + broker.configure(&config); + assert_eq!(broker.config_revision(), revision + 1); + assert_inherited_destination("updated-fallback-child"); + broker.virtualize_child_env(&mut env); + broker.virtualize_child_env_for_environment(&mut explicit_env, Some("explicit")); + assert_eq!(env[key], dummy); + for (destination, environment_id, value, injected) in [ + ("first.example", None, &dummy, true), + ("second.example", None, &dummy, true), + ("explicit.example", Some("explicit"), &explicit_dummy, true), + ( + "second.example", + Some("explicit"), + &explicit_dummy, + previously_used_fallback, + ), + ] { + let mut headers = headers_with_bearer(value); + broker.inject_request_headers_for_environment( + &format!("https://{destination}/v1"), + &mut headers, + environment_id, + ); + assert_eq!( + headers, + headers_with_bearer(if injected { token } else { value }), + "{key}: {destination}, {environment_id:?}" + ); + } + + config.configure_credential_broker_environment(&env_map([(context_key, "")])); + broker.configure(&config); + for destination in ["first.example", "second.example"] { + for (environment, value) in [(None, &dummy), (Some("explicit"), &explicit_dummy)] { + let mut headers = headers_with_bearer(value); + broker.inject_request_headers_for_environment( + &format!("https://{destination}/v1"), + &mut headers, + environment, + ); + assert_eq!(headers, headers_with_bearer(value)); + } + } + let mut headers = headers_with_bearer(&explicit_dummy); + broker.inject_request_headers_for_environment( + "https://explicit.example/v1", + &mut headers, + Some("explicit"), + ); + assert_eq!(headers, headers_with_bearer(token)); + let mut headers = headers_with_bearer(&explicit_dummy); + broker.inject_request_headers_for_environment( + "https://other-explicit.example/v1", + &mut headers, + Some("explicit"), + ); + assert_eq!(headers, headers_with_bearer(token)); + + // A fallback can become primary again without erasing captured destinations. + config.configure_credential_broker_environment(&env_map([( + context_key, + &destination("third.example"), + )])); + broker.configure(&config); + explicit_env.insert(context_key.to_string(), destination("third.example")); + broker.virtualize_child_env_for_environment(&mut explicit_env, Some("explicit")); + config.configure_credential_broker_environment(&env_map([(context_key, "")])); + broker.configure(&config); + for host in [ + "explicit.example", + "other-explicit.example", + "third.example", + ] { + let mut headers = headers_with_bearer(&explicit_dummy); + broker.inject_request_headers_for_environment( + &format!("https://{host}/v1"), + &mut headers, + Some("explicit"), + ); + assert_eq!( + headers, + headers_with_bearer(if host == "third.example" { + &explicit_dummy + } else { + token + }) + ); + } + } +} + +#[test] +fn openai_credentials_bind_only_to_default_and_configured_trusted_hosts() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut config = NetworkProxyConfig::default(); + config.set_credential_broker_enabled(/*enabled*/ true); + config.set_credential_broker_openai_base_url( + /*base_url*/ Some("https://gateway.example.com./v1"), + ); + broker.configure(&config); + + let mut env = env_map([ + ("OPENAI_API_KEY", "sk-real"), + ("OPENAI_BASE_URL", "https://sdk.example.com./v1"), + ("GH_TOKEN", "ghp-real"), + ]); + broker.virtualize_child_env(&mut env); + assert!(brokered_credential_env_keys(&env).any(|key| key == "OPENAI_BASE_URL")); + assert!(brokered_credential_binding_env_keys(&env).any(|key| key == "OPENAI_BASE_URL")); + let dummy = &env["OPENAI_API_KEY"]; + + for (host, expected_credential) in [ + ("api.openai.com", "sk-real"), + ("gateway.example.com", "sk-real"), + ("sdk.example.com", "sk-real"), + ("attacker.example", dummy.as_str()), + ] { + let mut headers = headers_with_bearer(dummy); + broker.inject_request_headers(host, &mut headers); + let expected = format!("Bearer {expected_credential}"); + assert_eq!(authorization(&headers), Some(expected.as_str()), "{host}"); + } + + config.set_credential_broker_openai_base_url( + /*base_url*/ Some("https://replacement.example/v1"), + ); + broker.configure(&config); + + let mut github_headers = headers_with_bearer(&env["GH_TOKEN"]); + broker.inject_request_headers("api.github.com", &mut github_headers); + assert_eq!(authorization(&github_headers), Some("Bearer ghp-real")); + + let mut openai_headers = headers_with_bearer(dummy); + broker.inject_request_headers("gateway.example.com", &mut openai_headers); + assert_eq!( + authorization(&openai_headers), + Some(format!("Bearer {dummy}").as_str()) + ); +} + +#[test] +fn github_cloud_credentials_match_ghe_com_host_hint() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([("GH_HOST", "astemu.ghe.com"), ("GH_TOKEN", "ghp-real")]); + broker.virtualize_child_env(&mut env); + assert!(!brokered_credential_binding_env_keys(&env).any(|key| key == "GH_HOST")); + let github_token = env.get("GH_TOKEN").expect("dummy GitHub token"); + let mut headers = headers_with_bearer(github_token); + + broker.inject_request_headers("api.astemu.ghe.com", &mut headers); + + assert_eq!(authorization(&headers), Some("Bearer ghp-real")); +} + +#[test] +fn github_cloud_credentials_do_not_bind_to_ghes_host_hint() { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([("GH_HOST", "github.example.com"), ("GH_TOKEN", "ghp-real")]); + broker.virtualize_child_env(&mut env); + let github_token = env.get("GH_TOKEN").expect("dummy github token"); + let expected_authorization = format!("Bearer {github_token}"); + let mut headers = headers_with_bearer(github_token); + + broker.inject_request_headers("github.example.com", &mut headers); + + assert_eq!( + authorization(&headers), + Some(expected_authorization.as_str()) + ); + assert!(!broker.host_requires_mitm("github.example.com", /*port*/ 443)); + assert!(broker.host_requires_mitm("api.github.com", /*port*/ 443)); +} + +#[test] +fn github_enterprise_credentials_bind_to_gh_host() { + for (hint, host) in [ + (" GitHub.Example.Com.:8443 ", "github.example.com"), + ("127.0.0.1:8443", "127.0.0.1"), + ("[::1]:8443", "::1"), + ("::1", "::1"), + ("[fe80::1%en0]:8443", "fe80::1%en0"), + ("fe80::1%en0", "fe80::1%en0"), + ] { + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([ + ("GH_HOST", hint), + ("GH_ENTERPRISE_TOKEN", "ghp-enterprise-real"), + ]); + broker.virtualize_child_env(&mut env); + let dummy = env["GH_ENTERPRISE_TOKEN"].clone(); + assert_ne!(dummy, "ghp-enterprise-real", "{hint}"); + env.remove("GH_HOST"); + env.insert( + "GH_ENTERPRISE_TOKEN".to_string(), + "ghp-enterprise-real".to_string(), + ); + broker.virtualize_child_env(&mut env); + let mut headers = headers_with_bearer(&dummy); + broker.inject_request_headers(host, &mut headers); + assert_eq!( + authorization(&headers), + Some("Bearer ghp-enterprise-real"), + "{hint}" + ); + } + let broker = CredentialBroker::new(/*enabled*/ true); + let mut env = env_map([ + ("GH_HOST", "github.example.com"), + ("GH_TOKEN", "ghp-enterprise-real"), + ("GH_ENTERPRISE_TOKEN", "ghp-enterprise-real"), + ("AUTH_HEADER", "Bearer ghp-enterprise-real"), + ]); + broker.virtualize_child_env(&mut env); + let github_dummy = env["GH_TOKEN"].clone(); + assert!(brokered_credential_env_keys(&env).any(|key| key == "GH_HOST")); + assert!(brokered_credential_binding_env_keys(&env).any(|key| key == "GH_HOST")); + let github_token = env + .get("GH_ENTERPRISE_TOKEN") + .expect("dummy GitHub enterprise token"); + assert_ne!(github_token, &github_dummy); + let mut headers = headers_with_bearer(github_token); + + broker.inject_request_headers("github.example.com", &mut headers); + + assert_eq!(authorization(&headers), Some("Bearer ghp-enterprise-real")); + assert_eq!(env["AUTH_HEADER"], format!("Bearer {github_token}")); + let mut alias_headers = headers_with_authorization(&env["AUTH_HEADER"]); + broker.inject_request_headers("github.example.com", &mut alias_headers); + assert_eq!( + authorization(&alias_headers), + Some("Bearer ghp-enterprise-real") + ); + let mut persisted_alias = "Bearer ghp-enterprise-real".to_string(); + assert!(broker.virtualize_text(&mut persisted_alias, &env)); + assert_eq!(persisted_alias, format!("Bearer {github_token}")); + assert_eq!( + brokered_credential_dummy_env_keys(&env).first(), + Some(&"GH_ENTERPRISE_TOKEN".to_string()) + ); + let mut cloud_headers = headers_with_bearer(&github_dummy); + broker.inject_request_headers("github.example.com", &mut cloud_headers); + assert_eq!(cloud_headers, headers_with_bearer(&github_dummy)); + let mut enterprise_headers = headers_with_bearer(github_token); + broker.inject_request_headers("api.github.com", &mut enterprise_headers); + assert_eq!(enterprise_headers, headers_with_bearer(github_token)); + let mut cloud_only = env_map([ + ("GH_HOST", "github.example.com"), + ("GH_TOKEN", "ghp-enterprise-real"), + ("AUTH_HEADER", "Bearer ghp-enterprise-real"), + ]); + broker.virtualize_child_env(&mut cloud_only); + assert_eq!(cloud_only["AUTH_HEADER"], format!("Bearer {github_dummy}")); + assert!(broker.host_requires_mitm("github.example.com", /*port*/ 443)); + assert!(broker.host_requires_mitm("api.github.com", /*port*/ 443)); + + env.insert("GH_HOST".to_string(), "attacker.example".to_string()); + env.insert("GH_ENTERPRISE_TOKEN".to_string(), github_dummy.clone()); + broker.virtualize_child_env(&mut env); + let mut attacker_headers = headers_with_bearer(&github_dummy); + broker.inject_request_headers("attacker.example", &mut attacker_headers); + assert_eq!(attacker_headers, headers_with_bearer(&github_dummy)); + assert!(!broker.host_requires_mitm("attacker.example", /*port*/ 443)); + + let mut alternate_enterprise_key = env_map([ + ("GH_HOST", "github.alternate.example"), + ("GH_TOKEN", "ghp-alternate-real"), + ("GITHUB_ENTERPRISE_TOKEN", "ghp-alternate-real"), + ("AUTH_HEADER", "Bearer ghp-alternate-real"), + ]); + broker.virtualize_child_env(&mut alternate_enterprise_key); + assert_eq!( + alternate_enterprise_key["AUTH_HEADER"], + format!( + "Bearer {}", + alternate_enterprise_key["GITHUB_ENTERPRISE_TOKEN"] + ) + ); + assert_eq!( + brokered_credential_dummy_env_keys(&alternate_enterprise_key).first(), + Some(&"GITHUB_ENTERPRISE_TOKEN".to_string()) + ); +} diff --git a/codex-rs/network-proxy/src/environment_policy.rs b/codex-rs/network-proxy/src/environment_policy.rs new file mode 100644 index 0000000000000000000000000000000000000000..365f7b2690b31395c44ad62b6c480759c3d21882 --- /dev/null +++ b/codex-rs/network-proxy/src/environment_policy.rs @@ -0,0 +1,84 @@ +use crate::NetworkDomainPermissions; +use crate::NetworkProxyConfig; +use crate::NetworkUnixSocketPermission; +use crate::NetworkUnixSocketPermissions; +use serde::Deserialize; +use serde::Serialize; + +/// Traffic restrictions supplied by the owner of one execution environment. +/// +/// Proxy enablement, listeners, network mode, MITM, and credentials remain outside +/// attachment-owned traffic policy. +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct EnvironmentNetworkPolicy { + pub domains: Option, + pub unix_sockets: Option, + pub allow_upstream_proxy: bool, + pub dangerously_allow_all_unix_sockets: bool, + pub allow_local_binding: bool, + pub managed_allowed_domains_only: bool, +} + +impl EnvironmentNetworkPolicy { + /// Captures portable traffic restrictions without exposing controller runtime settings. + pub fn from_config(config: &NetworkProxyConfig, managed_allowed_domains_only: bool) -> Self { + Self { + domains: config.domains.clone(), + unix_sockets: config.unix_sockets.clone(), + allow_upstream_proxy: config.allow_upstream_proxy, + dangerously_allow_all_unix_sockets: config.dangerously_allow_all_unix_sockets, + allow_local_binding: config.allow_local_binding, + managed_allowed_domains_only, + } + } + + /// Applies attachment-owned traffic settings while preserving inherited denials and proxy setup. + pub fn apply_to(&self, config: &mut NetworkProxyConfig) { + // Use the owner's domain rules without dropping controller denials. + let inherited_denials = config.denied_domains().unwrap_or_default(); + config.domains.clone_from(&self.domains); + for domain in inherited_denials { + config.upsert_domain_permission( + domain, + crate::NetworkDomainPermission::Deny, + crate::normalize_host, + ); + } + let inherited_sockets = config.unix_sockets.take().unwrap_or_default(); + let mut effective_sockets = self.unix_sockets.clone().unwrap_or_default(); + + // "Allow all" cannot override a socket denied by either policy. + let inherited_permits_all = config.dangerously_allow_all_unix_sockets + && !inherited_sockets + .entries + .values() + .any(|permission| matches!(permission, NetworkUnixSocketPermission::Deny)); + let owner_permits_all = self.dangerously_allow_all_unix_sockets + && !effective_sockets + .entries + .values() + .any(|permission| matches!(permission, NetworkUnixSocketPermission::Deny)); + + // Keep shared socket grants; controller denials always take priority. + effective_sockets.entries.retain(|path, permission| { + matches!(permission, NetworkUnixSocketPermission::Deny) + || inherited_permits_all + || matches!( + inherited_sockets.entries.get(path), + Some(NetworkUnixSocketPermission::Allow) + ) + }); + for (path, permission) in inherited_sockets.entries { + if owner_permits_all || matches!(permission, NetworkUnixSocketPermission::Deny) { + effective_sockets.entries.insert(path, permission); + } + } + + // Enable permissions only when both controller and owner allow them. + config.unix_sockets = (!effective_sockets.entries.is_empty()).then_some(effective_sockets); + config.dangerously_allow_all_unix_sockets = inherited_permits_all && owner_permits_all; + config.allow_upstream_proxy &= self.allow_upstream_proxy; + config.allow_local_binding &= self.allow_local_binding; + } +} diff --git a/codex-rs/network-proxy/src/http_proxy.rs b/codex-rs/network-proxy/src/http_proxy.rs new file mode 100644 index 0000000000000000000000000000000000000000..52263eafcf2e3f30c8471be6201a96f44d7fad36 --- /dev/null +++ b/codex-rs/network-proxy/src/http_proxy.rs @@ -0,0 +1,1917 @@ +use crate::attribution::BindConnectionAttribution; +use crate::config::NetworkMode; +use crate::connect_policy::TargetCheckedTcpConnector; +use crate::connection_lifecycle::CancelOnShutdown; +use crate::mitm; +use crate::network_policy::BlockDecisionAuditEventArgs; +use crate::network_policy::NetworkDecision; +use crate::network_policy::NetworkDecisionSource; +use crate::network_policy::NetworkPolicyDecider; +use crate::network_policy::NetworkPolicyDecision; +use crate::network_policy::NetworkPolicyRequest; +use crate::network_policy::NetworkPolicyRequestArgs; +use crate::network_policy::NetworkProtocol; +use crate::network_policy::emit_allow_decision_audit_event; +use crate::network_policy::emit_block_decision_audit_event; +use crate::network_policy::evaluate_host_policy; +use crate::policy::normalize_host; +use crate::reasons::REASON_METHOD_NOT_ALLOWED; +use crate::reasons::REASON_MITM_REQUIRED; +use crate::reasons::REASON_NOT_ALLOWED; +use crate::reasons::REASON_PROXY_DISABLED; +use crate::reasons::REASON_UNIX_SOCKET_UNSUPPORTED; +use crate::request_disconnect::NetworkRequestDisconnect; +use crate::responses::PolicyDecisionDetails; +use crate::responses::blocked_header_value; +use crate::responses::blocked_message_with_policy; +use crate::responses::blocked_text_response_with_policy; +use crate::responses::json_response; +use crate::runtime::HostMitmRequirement; +use crate::runtime::unix_socket_permissions_supported; +use crate::state::BlockedRequest; +use crate::state::BlockedRequestArgs; +use crate::state::NetworkProxyState; +use crate::upstream::UpstreamClient; +use crate::upstream::proxy_for_connect; +use anyhow::Context as _; +use anyhow::Result; +use codex_utils_rustls_provider::ensure_rustls_crypto_provider; +use rama_core::Layer; +use rama_core::Service; +use rama_core::error::ErrorExt as _; +use rama_core::error::OpaqueError; +use rama_core::extensions::ExtensionsMut; +use rama_core::extensions::ExtensionsRef; +use rama_core::graceful::ShutdownGuard; +use rama_core::service::BoxService; +use rama_core::service::service_fn; +use rama_core::stream::Stream; +use rama_http::Body; +use rama_http::HeaderMap; +use rama_http::HeaderName; +use rama_http::HeaderValue; +use rama_http::Request; +use rama_http::Response; +use rama_http::StatusCode; +use rama_http::header; +use rama_http::headers::HeaderMapExt; +use rama_http::headers::Host; +use rama_http::layer::remove_header::RemoveResponseHeaderLayer; +use rama_http::matcher::MethodMatcher; +use rama_http_backend::client::proxy::layer::HttpProxyConnector; +use rama_http_backend::server::HttpServer; +use rama_http_backend::server::layer::upgrade::UpgradeLayer; +use rama_http_backend::server::layer::upgrade::Upgraded; +use rama_net::Protocol; +use rama_net::client::ConnectorService; +use rama_net::client::EstablishedClientConnection; +use rama_net::http::RequestContext; +use rama_net::proxy::ProxyRequest; +use rama_net::proxy::ProxyTarget; +use rama_net::proxy::StreamForwardService; +use rama_net::stream::SocketInfo; +use rama_tcp::TcpStream; +use rama_tcp::client::Request as TcpRequest; +use rama_tcp::server::TcpListener; +use rama_tls_rustls::client::TlsConnectorDataBuilder; +use rama_tls_rustls::client::TlsConnectorLayer; +use serde::Serialize; +use std::convert::Infallible; +use std::net::SocketAddr; +use std::net::TcpListener as StdTcpListener; +use std::sync::Arc; +use std::time::Instant; +use tracing::error; +use tracing::info; +use tracing::warn; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum ConnectMitmMode { + Disabled, + Enabled, + DetectProtocol(crate::brokered_tunnel::BrokeredProtocols), +} + +pub async fn run_http_proxy( + state: Arc, + addr: SocketAddr, + policy_decider: Option>, + environment_id: Option, + guard: ShutdownGuard, +) -> Result<()> { + let listener = TcpListener::build() + .bind(addr) + .await + // Rama's `BoxError` is a `Box` without an explicit `'static` + // lifetime bound, which means it doesn't satisfy `anyhow::Context`'s `StdError` constraint. + // Wrap it in Rama's `OpaqueError` so we can preserve the original error as a source and + // still use `anyhow` for chaining. + .map_err(rama_core::error::OpaqueError::from) + .map_err(anyhow::Error::from) + .with_context(|| format!("bind HTTP proxy: {addr}"))?; + + run_http_proxy_with_listener(state, listener, policy_decider, environment_id, guard).await +} + +pub async fn run_http_proxy_with_std_listener( + state: Arc, + listener: StdTcpListener, + policy_decider: Option>, + environment_id: Option, + guard: ShutdownGuard, +) -> Result<()> { + let listener = + TcpListener::try_from(listener).context("convert std listener to HTTP proxy listener")?; + run_http_proxy_with_listener(state, listener, policy_decider, environment_id, guard).await +} + +async fn run_http_proxy_with_listener( + state: Arc, + listener: TcpListener, + policy_decider: Option>, + environment_id: Option, + guard: ShutdownGuard, +) -> Result<()> { + let addr = listener + .local_addr() + .context("read HTTP proxy listener local addr")?; + + info!("HTTP proxy listening on {addr}"); + + listener + .serve_graceful( + guard, + CancelOnShutdown::new(http_proxy_service(state, policy_decider, environment_id)), + ) + .await; + Ok(()) +} + +pub(crate) fn http_proxy_service( + state: Arc, + policy_decider: Option>, + environment_id: Option, +) -> BoxService { + ensure_rustls_crypto_provider(); + + // This proxy listener only needs HTTP/1 proxy semantics. Using Rama's auto builder + // forces every accepted socket through the HTTP version sniffing pre-read path before proxy + // request parsing, which can stall some local clients on macOS before CONNECT/absolute-form + // handling runs at all. + let http_service = HttpServer::http1().service( + ( + UpgradeLayer::new( + MethodMatcher::CONNECT, + service_fn({ + let policy_decider = policy_decider.clone(); + let environment_id = environment_id.clone(); + move |req| { + http_connect_accept(policy_decider.clone(), environment_id.clone(), req) + } + }), + CancelOnShutdown::new(service_fn(http_connect_proxy)), + ), + RemoveResponseHeaderLayer::hop_by_hop(), + ) + .into_layer(service_fn({ + let policy_decider = policy_decider.clone(); + let environment_id = environment_id.clone(); + move |req| http_plain_proxy(policy_decider.clone(), environment_id.clone(), req) + })), + ); + + BindConnectionAttribution::new(http_service, state, environment_id).boxed() +} + +async fn http_connect_accept( + policy_decider: Option>, + environment_id: Option, + mut req: Request, +) -> Result<(Response, Request), Response> { + let started_at = Instant::now(); + let app_state = req + .extensions() + .get::>() + .cloned() + .ok_or_else(|| text_response(StatusCode::INTERNAL_SERVER_ERROR, "missing state"))?; + + let authority = match RequestContext::try_from(&req).map(|ctx| ctx.host_with_port()) { + Ok(authority) => authority, + Err(err) => { + warn!("CONNECT missing authority: {err}"); + return Err(text_response(StatusCode::BAD_REQUEST, "missing authority")); + } + }; + + let host = normalize_host(&authority.host.to_string()); + if host.is_empty() { + return Err(text_response(StatusCode::BAD_REQUEST, "invalid host")); + } + + let client = client_addr(&req); + let enabled = app_state + .enabled() + .await + .map_err(|err| internal_error("failed to read enabled state", err))?; + if !enabled { + let client = client.as_deref().unwrap_or_default(); + warn!("CONNECT blocked; proxy disabled (client={client}, host={host})"); + return Err(proxy_disabled_response( + &app_state, + host, + authority.port, + client_addr(&req), + Some("CONNECT".to_string()), + NetworkProtocol::HttpsConnect, + /*audit_endpoint_override*/ None, + ) + .await); + } + + let disconnect = NetworkRequestDisconnect::default(); + let mut request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::HttpsConnect, + host: host.clone(), + port: authority.port, + environment_id, + client_addr: client.clone(), + method: Some("CONNECT".to_string()), + command: None, + exec_policy_hint: None, + }); + + request.disconnect = Some(disconnect.clone()); + match disconnect + .track_http_request( + started_at, + evaluate_host_policy(&app_state, policy_decider.as_ref(), &request), + ) + .await + { + Ok(NetworkDecision::Deny { + reason, + source, + decision, + }) => { + let details = PolicyDecisionDetails { + decision, + reason: &reason, + source, + protocol: NetworkProtocol::HttpsConnect, + host: &host, + port: authority.port, + }; + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: reason.clone(), + client: client.clone(), + method: Some("CONNECT".to_string()), + mode: None, + protocol: "http-connect".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(authority.port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!("CONNECT blocked (client={client}, host={host}, reason={reason})"); + return Err(blocked_text_with_details(&reason, &details)); + } + Ok(NetworkDecision::Allow) => { + let client = client.as_deref().unwrap_or_default(); + info!("CONNECT allowed (client={client}, host={host})"); + } + Err(err) => { + error!("failed to evaluate host for CONNECT {host}: {err}"); + return Err(text_response(StatusCode::INTERNAL_SERVER_ERROR, "error")); + } + } + + let mode = app_state + .network_mode() + .await + .map_err(|err| internal_error("failed to read network mode", err))?; + + let mitm_state = match app_state.mitm_state().await { + Ok(state) => state, + Err(err) => { + error!("failed to load MITM state: {err}"); + return Err(text_response(StatusCode::INTERNAL_SERVER_ERROR, "error")); + } + }; + let host_mitm_requirement = match app_state.host_mitm_requirement(&host, authority.port).await { + Ok(requirement) => requirement, + Err(err) => { + error!("failed to inspect MITM requirements for {host}: {err}"); + return Err(text_response(StatusCode::INTERNAL_SERVER_ERROR, "error")); + } + }; + let brokered_http = matches!(host_mitm_requirement, HostMitmRequirement::Credential(protocols) if protocols.http); + let connect_mitm_mode = if mode == NetworkMode::Limited && !brokered_http { + ConnectMitmMode::Enabled + } else { + match host_mitm_requirement { + HostMitmRequirement::None => ConnectMitmMode::Disabled, + HostMitmRequirement::Credential(protocols) => { + ConnectMitmMode::DetectProtocol(protocols) + } + HostMitmRequirement::Always => ConnectMitmMode::Enabled, + } + }; + + if connect_mitm_mode == ConnectMitmMode::Enabled && mitm_state.is_none() { + // Limited-mode enforcement and host-specific hooks require interception. Credential-only + // interception is deferred until the upgraded stream presents a supported protocol. + emit_http_block_decision_audit_event( + &app_state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ModeGuard, + reason: REASON_MITM_REQUIRED, + protocol: NetworkProtocol::HttpsConnect, + server_address: host.as_str(), + server_port: authority.port, + method: Some("CONNECT"), + client_addr: client.as_deref(), + }, + ); + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Deny, + reason: REASON_MITM_REQUIRED, + source: NetworkDecisionSource::ModeGuard, + protocol: NetworkProtocol::HttpsConnect, + host: &host, + port: authority.port, + }; + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: REASON_MITM_REQUIRED.to_string(), + client: client.clone(), + method: Some("CONNECT".to_string()), + mode: Some(mode), + protocol: "http-connect".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(authority.port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!( + "CONNECT blocked; MITM required to enforce HTTPS policy (client={client}, host={host}, mode={mode:?}, host_mitm_requirement={host_mitm_requirement:?})" + ); + return Err(blocked_text_with_details(REASON_MITM_REQUIRED, &details)); + } + + req.extensions_mut().insert(ProxyTarget(authority)); + req.extensions_mut().insert(connect_mitm_mode); + req.extensions_mut().insert(mode); + if connect_mitm_mode != ConnectMitmMode::Disabled + && let Some(mitm_state) = mitm_state + { + req.extensions_mut().insert(mitm_state); + } + + Ok(( + Response::builder() + .status(StatusCode::OK) + .body(Body::empty()) + .unwrap_or_else(|_| Response::new(Body::empty())), + req, + )) +} + +async fn http_connect_proxy(upgraded: Upgraded) -> Result<(), Infallible> { + let mode = upgraded + .extensions() + .get::() + .copied() + .unwrap_or(NetworkMode::Full); + let connect_mitm_mode = upgraded + .extensions() + .get::() + .copied() + .unwrap_or(ConnectMitmMode::Disabled); + let result: Result<(), OpaqueError> = match connect_mitm_mode { + ConnectMitmMode::Disabled => forward_connect_tunnel(upgraded).await, + ConnectMitmMode::Enabled => mitm_connect_tunnel(upgraded).await, + ConnectMitmMode::DetectProtocol(protocols) => { + match crate::brokered_tunnel::peek_protocol(upgraded, protocols).await { + Ok((crate::brokered_tunnel::TunnelProtocol::Tls, stream)) => { + mitm_connect_tunnel(stream).await + } + Ok((crate::brokered_tunnel::TunnelProtocol::Http, stream)) => { + mitm::mitm_stream(stream, rama_http::uri::Scheme::HTTP) + .await + .map_err(|err| { + OpaqueError::from_display(format!("HTTP tunnel error: {err}")) + }) + } + Ok((crate::brokered_tunnel::TunnelProtocol::Opaque, stream)) => { + if mode == NetworkMode::Limited { + Err(OpaqueError::from_display( + "opaque tunnels are not allowed in limited mode", + )) + } else { + forward_connect_tunnel(stream).await + } + } + Err(err) => Err(OpaqueError::from_display(format!( + "detect tunnel protocol: {err:#}" + ))), + } + } + }; + if let Err(err) = result { + warn!("CONNECT tunnel error: {err}"); + } + Ok(()) +} + +async fn mitm_connect_tunnel(stream: S) -> Result<(), OpaqueError> +where + S: Stream + Unpin + ExtensionsMut, +{ + let target = stream + .extensions() + .get::() + .map(|target| target.0.clone()) + .ok_or_else(|| OpaqueError::from_display("missing MITM authority"))?; + let host = normalize_host(&target.host.to_string()); + let port = target.port; + let mode = stream + .extensions() + .get::() + .copied() + .unwrap_or(NetworkMode::Full); + if stream.extensions().get::>().is_none() { + return Err(OpaqueError::from_display(format!( + "cannot enable MITM without state (host={host}, port={port})" + ))); + } + + info!("CONNECT MITM enabled (host={host}, port={port}, mode={mode:?})"); + mitm::mitm_stream(stream, rama_http::uri::Scheme::HTTPS) + .await + .map_err(|err| OpaqueError::from_display(format!("MITM tunnel error: {err}"))) +} + +async fn forward_connect_tunnel(upgraded: S) -> Result<(), OpaqueError> +where + S: Stream + Unpin + ExtensionsMut, +{ + let authority = upgraded + .extensions() + .get::() + .map(|target| target.0.clone()) + .ok_or_else(|| OpaqueError::from_display("missing forward authority"))?; + let app_state = upgraded + .extensions() + .get::>() + .cloned() + .ok_or_else(|| OpaqueError::from_display("missing app state"))?; + let allow_upstream_proxy = match app_state.allow_upstream_proxy().await { + Ok(allowed) => allowed, + Err(err) => { + error!("failed to read upstream proxy setting: {err}"); + false + } + }; + let proxy = if allow_upstream_proxy { + proxy_for_connect(&authority) + } else { + None + }; + match proxy.as_ref() { + Some(proxy) => info!( + "CONNECT route selected (host={}, port={}, route=upstream_proxy, proxy={})", + authority.host, authority.port, proxy.address + ), + None => info!( + "CONNECT route selected (host={}, port={}, route=direct)", + authority.host, authority.port + ), + } + + let mut extensions = upgraded.extensions().clone(); + if let Some(proxy) = proxy { + extensions.insert(proxy); + } + + let req = TcpRequest::new_with_extensions(authority.clone(), extensions) + .with_protocol(Protocol::HTTPS); + let proxy_connector = HttpProxyConnector::optional(TargetCheckedTcpConnector::new(app_state)); + let tls_config = TlsConnectorDataBuilder::new() + .with_alpn_protocols_http_auto() + .build(); + let connector = TlsConnectorLayer::tunnel(None) + .with_connector_data(tls_config) + .into_layer(proxy_connector); + info!("CONNECT upstream dial started (target={authority})"); + let connect_started_at = Instant::now(); + let EstablishedClientConnection { conn: target, .. } = match connector.connect(req).await { + Ok(connection) => { + info!( + "CONNECT upstream dial established (target={authority}, elapsed_ms={})", + connect_started_at.elapsed().as_millis() + ); + connection + } + Err(err) => { + warn!( + "CONNECT upstream dial failed (target={authority}, elapsed_ms={})", + connect_started_at.elapsed().as_millis() + ); + return Err(OpaqueError::from_boxed(err) + .with_context(|| format!("establish CONNECT tunnel to {authority}"))); + } + }; + + let proxy_req = ProxyRequest { + source: upgraded, + target, + }; + info!("CONNECT tunnel forwarding started (target={authority})"); + let forward_started_at = Instant::now(); + StreamForwardService::default() + .serve(proxy_req) + .await + .map(|_| { + info!( + "CONNECT tunnel forwarding completed (target={authority}, elapsed_ms={})", + forward_started_at.elapsed().as_millis() + ); + }) + .map_err(|err| { + warn!( + "CONNECT tunnel forwarding failed (target={authority}, elapsed_ms={})", + forward_started_at.elapsed().as_millis() + ); + OpaqueError::from_boxed(err.into()) + .with_context(|| format!("forward CONNECT tunnel to {authority}")) + }) +} + +async fn http_plain_proxy( + policy_decider: Option>, + environment_id: Option, + mut req: Request, +) -> Result { + let started_at = Instant::now(); + let app_state = match req.extensions().get::>().cloned() { + Some(state) => state, + None => { + error!("missing app state"); + return Ok(text_response(StatusCode::INTERNAL_SERVER_ERROR, "error")); + } + }; + let client = client_addr(&req); + let method_allowed = match app_state + .method_allowed(req.method().as_str()) + .await + .map_err(|err| internal_error("failed to evaluate method policy", err)) + { + Ok(allowed) => allowed, + Err(resp) => return Ok(resp), + }; + + // `x-unix-socket` is an escape hatch for talking to local daemons. We keep it tightly scoped: + // macOS-only + explicit allowlist by default, to avoid turning the proxy into a general local + // capability escalation mechanism. + if let Some(unix_socket_header) = req.headers().get("x-unix-socket") { + let socket_path = match unix_socket_header.to_str() { + Ok(value) => value.to_string(), + Err(_) => { + warn!("invalid x-unix-socket header value (non-UTF8)"); + return Ok(text_response( + StatusCode::BAD_REQUEST, + "invalid x-unix-socket header", + )); + } + }; + let enabled = match app_state + .enabled() + .await + .map_err(|err| internal_error("failed to read enabled state", err)) + { + Ok(enabled) => enabled, + Err(resp) => return Ok(resp), + }; + if !enabled { + let client = client.as_deref().unwrap_or_default(); + warn!("unix socket blocked; proxy disabled (client={client}, path={socket_path})"); + return Ok(proxy_disabled_response( + &app_state, + socket_path, + /*port*/ 0, + client_addr(&req), + Some(req.method().as_str().to_string()), + NetworkProtocol::Http, + Some(("unix-socket", 0)), + ) + .await); + } + if !method_allowed { + emit_http_block_decision_audit_event( + &app_state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ModeGuard, + reason: REASON_METHOD_NOT_ALLOWED, + protocol: NetworkProtocol::Http, + server_address: "unix-socket", + server_port: 0, + method: Some(req.method().as_str()), + client_addr: client.as_deref(), + }, + ); + let client = client.as_deref().unwrap_or_default(); + let method = req.method(); + warn!( + "unix socket blocked by method policy (client={client}, method={method}, mode=limited, allowed_methods=GET, HEAD, OPTIONS)" + ); + return Ok(json_blocked( + "unix-socket", + REASON_METHOD_NOT_ALLOWED, + /*details*/ None, + )); + } + + if !unix_socket_permissions_supported() { + emit_http_block_decision_audit_event( + &app_state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ProxyState, + reason: REASON_UNIX_SOCKET_UNSUPPORTED, + protocol: NetworkProtocol::Http, + server_address: "unix-socket", + server_port: 0, + method: Some(req.method().as_str()), + client_addr: client.as_deref(), + }, + ); + warn!("unix socket proxy unsupported on this platform (path={socket_path})"); + return Ok(text_response( + StatusCode::NOT_IMPLEMENTED, + "unix sockets unsupported", + )); + } + + return match app_state.is_unix_socket_allowed(&socket_path).await { + Ok(true) => { + emit_http_allow_decision_audit_event( + &app_state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ProxyState, + reason: "allow", + protocol: NetworkProtocol::Http, + server_address: "unix-socket", + server_port: 0, + method: Some(req.method().as_str()), + client_addr: client.as_deref(), + }, + ); + let client = client.as_deref().unwrap_or_default(); + info!("unix socket allowed (client={client}, path={socket_path})"); + match proxy_via_unix_socket(req, &socket_path).await { + Ok(resp) => Ok(resp), + Err(err) => { + warn!("unix socket proxy failed: {err}"); + Ok(text_response( + StatusCode::BAD_GATEWAY, + "unix socket proxy failed", + )) + } + } + } + Ok(false) => { + emit_http_block_decision_audit_event( + &app_state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ProxyState, + reason: REASON_NOT_ALLOWED, + protocol: NetworkProtocol::Http, + server_address: "unix-socket", + server_port: 0, + method: Some(req.method().as_str()), + client_addr: client.as_deref(), + }, + ); + let client = client.as_deref().unwrap_or_default(); + warn!("unix socket blocked (client={client}, path={socket_path})"); + Ok(json_blocked( + "unix-socket", + REASON_NOT_ALLOWED, + /*details*/ None, + )) + } + Err(err) => { + warn!("unix socket check failed: {err}"); + Ok(text_response(StatusCode::INTERNAL_SERVER_ERROR, "error")) + } + }; + } + + let request_ctx = match RequestContext::try_from(&req) { + Ok(request_ctx) => request_ctx, + Err(err) => { + warn!("missing host: {err}"); + return Ok(text_response(StatusCode::BAD_REQUEST, "missing host")); + } + }; + let authority = request_ctx.host_with_port(); + let host = normalize_host(&authority.host.to_string()); + let port = authority.port; + if let Err(reason) = validate_absolute_form_host_header(&req, &request_ctx) { + let client = client.as_deref().unwrap_or_default(); + let host_header = req + .headers() + .get(header::HOST) + .and_then(|value| value.to_str().ok()) + .unwrap_or(""); + warn!( + "request rejected due to mismatched Host header (client={client}, target={host}:{port}, host_header={host_header}, reason={reason})" + ); + return Ok(text_response(StatusCode::BAD_REQUEST, reason)); + } + let enabled = match app_state + .enabled() + .await + .map_err(|err| internal_error("failed to read enabled state", err)) + { + Ok(enabled) => enabled, + Err(resp) => return Ok(resp), + }; + if !enabled { + let client = client.as_deref().unwrap_or_default(); + let method = req.method(); + warn!("request blocked; proxy disabled (client={client}, host={host}, method={method})"); + return Ok(proxy_disabled_response( + &app_state, + host, + port, + client_addr(&req), + Some(req.method().as_str().to_string()), + NetworkProtocol::Http, + /*audit_endpoint_override*/ None, + ) + .await); + } + + let disconnect = NetworkRequestDisconnect::default(); + let mut request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Http, + host: host.clone(), + port, + environment_id, + client_addr: client.clone(), + method: Some(req.method().as_str().to_string()), + command: None, + exec_policy_hint: None, + }); + + request.disconnect = Some(disconnect.clone()); + match disconnect + .track_http_request( + started_at, + evaluate_host_policy(&app_state, policy_decider.as_ref(), &request), + ) + .await + { + Ok(NetworkDecision::Deny { + reason, + source, + decision, + }) => { + let details = PolicyDecisionDetails { + decision, + reason: &reason, + source, + protocol: NetworkProtocol::Http, + host: &host, + port, + }; + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: reason.clone(), + client: client.clone(), + method: Some(req.method().as_str().to_string()), + mode: None, + protocol: "http".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!("request blocked (client={client}, host={host}, reason={reason})"); + return Ok(json_blocked(&host, &reason, Some(&details))); + } + Ok(NetworkDecision::Allow) => {} + Err(err) => { + error!("failed to evaluate host for {host}: {err}"); + return Ok(text_response(StatusCode::INTERNAL_SERVER_ERROR, "error")); + } + } + + let host_mitm_requirement = match app_state.host_mitm_requirement(&host, port).await { + Ok(requirement) => requirement, + Err(err) => { + return Ok(internal_error("failed to inspect MITM requirements", err)); + } + }; + if host_mitm_requirement == HostMitmRequirement::Always { + emit_http_block_decision_audit_event( + &app_state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ModeGuard, + reason: REASON_MITM_REQUIRED, + protocol: NetworkProtocol::Http, + server_address: host.as_str(), + server_port: port, + method: Some(req.method().as_str()), + client_addr: client.as_deref(), + }, + ); + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Deny, + reason: REASON_MITM_REQUIRED, + source: NetworkDecisionSource::ModeGuard, + protocol: NetworkProtocol::Http, + host: &host, + port, + }; + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: REASON_MITM_REQUIRED.to_string(), + client: client.clone(), + method: Some(req.method().as_str().to_string()), + mode: None, + protocol: "http".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!( + "request blocked; MITM required to enforce host policy (client={client}, host={host}, method={})", + req.method() + ); + return Ok(json_blocked(&host, REASON_MITM_REQUIRED, Some(&details))); + } + + if !method_allowed { + emit_http_block_decision_audit_event( + &app_state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ModeGuard, + reason: REASON_METHOD_NOT_ALLOWED, + protocol: NetworkProtocol::Http, + server_address: host.as_str(), + server_port: port, + method: Some(req.method().as_str()), + client_addr: client.as_deref(), + }, + ); + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Deny, + reason: REASON_METHOD_NOT_ALLOWED, + source: NetworkDecisionSource::ModeGuard, + protocol: NetworkProtocol::Http, + host: &host, + port, + }; + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: REASON_METHOD_NOT_ALLOWED.to_string(), + client: client.clone(), + method: Some(req.method().as_str().to_string()), + mode: Some(NetworkMode::Limited), + protocol: "http".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + let method = req.method(); + warn!( + "request blocked by method policy (client={client}, host={host}, method={method}, mode=limited, allowed_methods=GET, HEAD, OPTIONS)" + ); + return Ok(json_blocked( + &host, + REASON_METHOD_NOT_ALLOWED, + Some(&details), + )); + } + + if let Err(err) = + inject_forward_request_credentials(app_state.as_ref(), &request_ctx, &mut req).await + { + return Ok(internal_error( + "failed to read plaintext credential injection config", + err, + )); + } + + let client = client.as_deref().unwrap_or_default(); + let method = req.method(); + info!("request allowed (client={client}, host={host}, method={method})"); + + let allow_upstream_proxy = match app_state + .allow_upstream_proxy() + .await + .map_err(|err| internal_error("failed to read upstream proxy config", err)) + { + Ok(allow) => allow, + Err(resp) => return Ok(resp), + }; + let client = if allow_upstream_proxy { + UpstreamClient::from_env_proxy(app_state.clone()) + } else { + UpstreamClient::direct(app_state.clone()) + }; + + // Strip hop-by-hop headers only after extracting metadata used for policy correlation. + remove_hop_by_hop_request_headers(req.headers_mut()); + match client.serve(req).await { + Ok(resp) => Ok(resp), + Err(err) => { + warn!("upstream request failed: {err}"); + Ok(text_response(StatusCode::BAD_GATEWAY, "upstream failure")) + } + } +} + +async fn inject_forward_request_credentials( + app_state: &NetworkProxyState, + context: &RequestContext, + req: &mut Request, +) -> Result<()> { + let authority = context.host_with_port(); + let scheme = req.uri().scheme_str().unwrap_or("http"); + // Server-wide OPTIONS may use root-scoped credentials without changing its wire target. + let request_path = req + .uri() + .path_and_query() + .map(rama_http::uri::PathAndQuery::as_str) + .filter(|path| *path != "*") + .unwrap_or("/"); + let destination = format!("{scheme}://{authority}{request_path}"); + let unrestricted_plaintext = app_state.plaintext_credential_injection_enabled().await?; + app_state.inject_request_credentials(&destination, req.headers_mut()); + if unrestricted_plaintext { + app_state.inject_request_credentials( + &normalize_host(&authority.host.to_string()), + req.headers_mut(), + ); + } + Ok(()) +} + +async fn proxy_via_unix_socket(req: Request, socket_path: &str) -> Result { + #[cfg(target_os = "macos")] + { + let client = UpstreamClient::unix_socket(socket_path); + + let (mut parts, body) = req.into_parts(); + let path = parts + .uri + .path_and_query() + .map(rama_http::uri::PathAndQuery::as_str) + .unwrap_or("/"); + parts.uri = path + .parse() + .with_context(|| format!("invalid unix socket request path: {path}"))?; + parts.headers.remove("x-unix-socket"); + remove_hop_by_hop_request_headers(&mut parts.headers); + + let req = Request::from_parts(parts, body); + client.serve(req).await.map_err(anyhow::Error::from) + } + #[cfg(not(target_os = "macos"))] + { + let _ = req; + let _ = socket_path; + Err(anyhow::anyhow!("unix sockets not supported")) + } +} + +fn client_addr(input: &T) -> Option { + input + .extensions() + .get::() + .map(|info| info.peer_addr().to_string()) +} + +fn validate_absolute_form_host_header( + req: &Request, + request_ctx: &RequestContext, +) -> Result<(), &'static str> { + if req.uri().scheme_str().is_none() { + return Ok(()); + } + + let Some(host_header) = req + .headers() + .typed_try_get::() + .map_err(|_| "invalid Host header")? + else { + return Ok(()); + }; + + if host_header.0.host != request_ctx.authority.host { + return Err("Host header does not match request target"); + } + + if let Some(host_port) = host_header.0.port { + if Some(host_port) != request_ctx.authority.port { + return Err("Host header does not match request target"); + } + return Ok(()); + } + + if !request_ctx.authority_has_default_port() { + return Err("Host header does not match request target"); + } + + Ok(()) +} +pub(crate) fn remove_hop_by_hop_request_headers(headers: &mut HeaderMap) { + while let Some(raw_connection) = headers.get(header::CONNECTION).cloned() { + headers.remove(header::CONNECTION); + if let Ok(raw_connection) = raw_connection.to_str() { + let connection_headers: Vec = raw_connection + .split(',') + .map(str::trim) + .filter(|token| !token.is_empty()) + .map(ToOwned::to_owned) + .collect(); + for token in connection_headers { + if let Ok(name) = HeaderName::from_bytes(token.as_bytes()) { + headers.remove(name); + } + } + } + } + for name in [ + &header::KEEP_ALIVE, + &header::PROXY_CONNECTION, + &header::PROXY_AUTHORIZATION, + &header::TRAILER, + &header::TRANSFER_ENCODING, + &header::UPGRADE, + ] { + headers.remove(name); + } + + // codespell:ignore te,TE + // 0x74,0x65 is ASCII "te" (the HTTP TE hop-by-hop header). + if let Ok(short_hop_header_name) = HeaderName::from_bytes(&[0x74, 0x65]) { + headers.remove(short_hop_header_name); + } +} + +fn json_blocked(host: &str, reason: &str, details: Option<&PolicyDecisionDetails<'_>>) -> Response { + let (message, decision, source, protocol, port) = details + .map(|details| { + ( + Some(blocked_message_with_policy(reason, details)), + Some(details.decision.as_str()), + Some(details.source.as_str()), + Some(details.protocol.as_policy_protocol()), + Some(details.port), + ) + }) + .unwrap_or((None, None, None, None, None)); + let response = BlockedResponse { + status: "blocked", + host, + reason, + decision, + source, + protocol, + port, + message, + }; + let mut resp = json_response(&response); + *resp.status_mut() = StatusCode::FORBIDDEN; + resp.headers_mut().insert( + "x-proxy-error", + HeaderValue::from_static(blocked_header_value(reason)), + ); + resp +} + +fn blocked_text_with_details(reason: &str, details: &PolicyDecisionDetails<'_>) -> Response { + blocked_text_response_with_policy(reason, details) +} + +async fn proxy_disabled_response( + app_state: &NetworkProxyState, + host: String, + port: u16, + client: Option, + method: Option, + protocol: NetworkProtocol, + audit_endpoint_override: Option<(&str, u16)>, +) -> Response { + let (audit_server_address, audit_server_port) = + audit_endpoint_override.unwrap_or((host.as_str(), port)); + emit_http_block_decision_audit_event( + app_state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ProxyState, + reason: REASON_PROXY_DISABLED, + protocol, + server_address: audit_server_address, + server_port: audit_server_port, + method: method.as_deref(), + client_addr: client.as_deref(), + }, + ); + + let blocked_host = host.clone(); + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: blocked_host, + reason: REASON_PROXY_DISABLED.to_string(), + client, + method, + mode: None, + protocol: protocol.as_policy_protocol().to_string(), + decision: Some("deny".to_string()), + source: Some("proxy_state".to_string()), + port: Some(port), + })) + .await; + + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Deny, + reason: REASON_PROXY_DISABLED, + source: NetworkDecisionSource::ProxyState, + protocol, + host: &host, + port, + }; + text_response( + StatusCode::SERVICE_UNAVAILABLE, + &blocked_message_with_policy(REASON_PROXY_DISABLED, &details), + ) +} + +fn internal_error(context: &str, err: impl std::fmt::Display) -> Response { + error!("{context}: {err}"); + text_response(StatusCode::INTERNAL_SERVER_ERROR, "error") +} + +fn text_response(status: StatusCode, body: &str) -> Response { + Response::builder() + .status(status) + .header("content-type", "text/plain") + .body(Body::from(body.to_string())) + .unwrap_or_else(|_| Response::new(Body::from(body.to_string()))) +} + +fn emit_http_block_decision_audit_event( + app_state: &NetworkProxyState, + args: BlockDecisionAuditEventArgs<'_>, +) { + emit_block_decision_audit_event(app_state, args); +} + +fn emit_http_allow_decision_audit_event( + app_state: &NetworkProxyState, + args: BlockDecisionAuditEventArgs<'_>, +) { + emit_allow_decision_audit_event(app_state, args); +} + +#[derive(Serialize)] +struct BlockedResponse<'a> { + status: &'static str, + host: &'a str, + reason: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + decision: Option<&'static str>, + #[serde(skip_serializing_if = "Option::is_none")] + source: Option<&'static str>, + #[serde(skip_serializing_if = "Option::is_none")] + protocol: Option<&'static str>, + #[serde(skip_serializing_if = "Option::is_none")] + port: Option, + #[serde(skip_serializing_if = "Option::is_none")] + message: Option, +} + +#[cfg(test)] +mod tests { + use super::*; + + use crate::CredentialProviderConfig; + use crate::config::NetworkMode; + use crate::config::NetworkProxyConfig; + use crate::runtime::network_proxy_state_for_policy; + use pretty_assertions::assert_eq; + use rama_http::Method; + use rama_http::Request; + use std::collections::BTreeMap; + use std::collections::HashMap; + use std::net::Ipv4Addr; + use std::net::TcpListener as StdTcpListener; + use std::sync::Arc; + use std::sync::Mutex; + use tokio::io::AsyncReadExt; + use tokio::io::AsyncWriteExt; + use tokio::net::TcpListener as TokioTcpListener; + use tokio::time::Duration; + use tokio::time::timeout; + + #[tokio::test] + async fn http_connect_accept_blocks_in_limited_mode() { + let policy = { + let mut policy = NetworkProxyConfig::default(); + policy.set_allowed_domains(vec!["example.com".to_string()]); + policy + }; + let state = Arc::new(network_proxy_state_for_policy(policy)); + state.set_network_mode(NetworkMode::Limited).await.unwrap(); + + let mut req = Request::builder() + .method(Method::CONNECT) + .uri("https://example.com:443") + .header("host", "example.com:443") + .body(Body::empty()) + .unwrap(); + req.extensions_mut().insert(state); + + let response = http_connect_accept( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await + .unwrap_err(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response.headers().get("x-proxy-error").unwrap(), + "blocked-by-mitm-required" + ); + } + + #[tokio::test] + async fn http_connect_accept_allows_allowlisted_host_in_full_mode() { + let policy = { + let mut policy = NetworkProxyConfig { + allow_local_binding: true, + ..NetworkProxyConfig::default() + }; + policy.set_allowed_domains(vec!["example.com".to_string()]); + policy + }; + let state = Arc::new(network_proxy_state_for_policy(policy)); + + let mut req = Request::builder() + .method(Method::CONNECT) + .uri("https://example.com:443") + .header("host", "example.com:443") + .body(Body::empty()) + .unwrap(); + req.extensions_mut().insert(state); + + let (response, _request) = http_connect_accept( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn http_connect_accept_passes_environment_id_to_decider() { + let state = Arc::new(network_proxy_state_for_policy(NetworkProxyConfig::default())); + let seen_environment_id = Arc::new(Mutex::new(None)); + let decider: Arc = Arc::new({ + let seen_environment_id = seen_environment_id.clone(); + move |request: NetworkPolicyRequest| { + *seen_environment_id + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = request.environment_id; + async { NetworkDecision::Allow } + } + }); + + let mut req = Request::builder() + .method(Method::CONNECT) + .uri("https://example.com:443") + .header("host", "example.com:443") + .body(Body::empty()) + .unwrap(); + req.extensions_mut().insert(state); + + let (response, _request) = + http_connect_accept(Some(decider), Some("remote".to_string()), req) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + seen_environment_id + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .as_deref(), + Some("remote") + ); + } + + #[tokio::test] + async fn http_connect_accept_defers_brokered_host_mitm_until_protocol_detection() { + let mut policy = NetworkProxyConfig { + credential_broker: true, + mitm: true, + ..NetworkProxyConfig::default() + }; + policy.set_allowed_domains(vec!["github.com".to_string()]); + let state = Arc::new(network_proxy_state_for_policy(policy)); + let mut env = HashMap::from([("GH_TOKEN".to_string(), "ghp-real".to_string())]); + state.virtualize_child_credentials(&mut env); + + let mut req = Request::builder() + .method(Method::CONNECT) + .uri("https://github.com:22") + .header("host", "github.com:22") + .body(Body::empty()) + .unwrap(); + req.extensions_mut().insert(state); + + let (response, request) = http_connect_accept( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await + .expect("brokered credentials should defer MITM until protocol detection"); + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + request.extensions().get::().copied(), + Some(ConnectMitmMode::DetectProtocol( + crate::brokered_tunnel::BrokeredProtocols { + tls: true, + http: false + } + )) + ); + } + + #[tokio::test] + async fn plaintext_credential_injection_requires_explicit_opt_in() { + let real_token = "ghp-real"; + for enabled in [false, true] { + let state = network_proxy_state_for_policy(NetworkProxyConfig { + credential_broker: true, + dangerously_allow_plaintext_credential_injection: enabled, + mitm: true, + ..NetworkProxyConfig::default() + }); + let mut env = HashMap::from([("GH_TOKEN".to_string(), real_token.to_string())]); + state.virtualize_child_credentials(&mut env); + let dummy = env.get("GH_TOKEN").expect("dummy GitHub token"); + let mut req = Request::builder() + .uri("http://api.github.com/") + .header(header::AUTHORIZATION, format!("Bearer {dummy}")) + .body(Body::empty()) + .unwrap(); + let context = RequestContext::try_from(&req).unwrap(); + + inject_forward_request_credentials(&state, &context, &mut req) + .await + .unwrap(); + + let expected = if enabled { real_token } else { dummy }; + assert_eq!( + req.headers()[header::AUTHORIZATION], + format!("Bearer {expected}") + ); + } + } + + #[tokio::test] + async fn configured_forward_requests_match_destination_scheme_and_port() { + let real_token = "provider_abcdefghijklmnopqrstuvwx"; + for (host, port, prefix, scope) in [ + ("localhost", 8080, ""), + ("127.0.0.1", 8081, ""), + ("api.provider.example", 443, "https://"), + ("api.provider.example", 8443, "https://"), + ] + .into_iter() + .flat_map(|(host, port, prefix)| ["/v1", "/"].map(move |scope| (host, port, prefix, scope))) + { + let state = network_proxy_state_for_policy(NetworkProxyConfig { + credential_broker: true, + dangerously_allow_plaintext_credential_injection: true, + credential_providers: BTreeMap::from([( + "custom".to_string(), + CredentialProviderConfig { + env: vec!["PROVIDER_TOKEN".to_string()], + patterns: vec!["^provider_[a-z]{24}$".to_string()], + url_prefix_from_env: Some("PROVIDER_URL".to_string()), + ..CredentialProviderConfig::default() + }, + )]), + ..NetworkProxyConfig::default() + }); + let mut env = HashMap::from([ + ("PROVIDER_TOKEN".to_string(), real_token.to_string()), + ( + "PROVIDER_URL".to_string(), + format!("{prefix}{host}:{port}{scope}"), + ), + ]); + state.virtualize_child_credentials(&mut env); + let dummy = env.get("PROVIDER_TOKEN").expect("dummy provider token"); + assert_ne!(dummy, real_token); + assert_eq!( + state.host_mitm_requirement(host, port).await.unwrap(), + HostMitmRequirement::Credential(crate::brokered_tunnel::BrokeredProtocols { + tls: !prefix.is_empty(), + http: prefix.is_empty() + }) + ); + assert_eq!( + state.host_mitm_requirement(host, port + 1).await.unwrap(), + HostMitmRequirement::None + ); + assert_eq!( + state + .for_environment_id(Some("other")) + .host_mitm_requirement(host, port) + .await + .unwrap(), + HostMitmRequirement::None + ); + for (uri, inject) in [ + (format!("http://{host}:{port}/v1/models"), prefix.is_empty()), + (format!("http://{host}:{}/v1/models", port + 1), false), + (format!("https://{host}:{}/v1/models", port + 1), false), + ( + format!("https://{host}:{port}/v1/models"), + !prefix.is_empty(), + ), + ("/v1/models".to_string(), prefix.is_empty()), + ( + format!("https://{host}:{port}/other"), + !prefix.is_empty() && scope == "/", + ), + ("*".to_string(), prefix.is_empty() && scope == "/"), + ] { + let request_uri: rama_http::Uri = uri.parse().unwrap(); + let host_header = request_uri + .authority() + .map(|authority| authority.as_str().to_string()) + .unwrap_or_else(|| format!("{host}:{port}")); + let mut req = Request::builder() + .method(if uri == "*" { + Method::OPTIONS + } else { + Method::GET + }) + .uri(request_uri) + .header(header::HOST, host_header) + .header(header::AUTHORIZATION, format!("Bearer {dummy}")) + .body(Body::empty()) + .unwrap(); + let context = RequestContext::try_from(&req).unwrap(); + + inject_forward_request_credentials(&state, &context, &mut req) + .await + .unwrap(); + + let expected = if inject { real_token } else { dummy }; + assert_eq!( + req.headers()[header::AUTHORIZATION], + format!("Bearer {expected}"), + "request: {uri}, configured scope: {scope}" + ); + assert_eq!(req.uri().to_string(), uri); + } + } + } + + #[tokio::test] + async fn http_connect_accept_blocks_hooked_host_in_full_mode_without_mitm_state() { + let mut policy = NetworkProxyConfig { + mitm: true, + mitm_hooks: vec![crate::mitm_hook::MitmHookConfig { + host: "api.github.com".to_string(), + matcher: crate::mitm_hook::MitmHookMatchConfig { + methods: vec!["POST".to_string()], + path_prefixes: vec!["/repos/openai/".to_string()], + ..crate::mitm_hook::MitmHookMatchConfig::default() + }, + actions: crate::mitm_hook::MitmHookActionsConfig::default(), + }], + ..Default::default() + }; + policy.set_allowed_domains(vec!["api.github.com".to_string()]); + let state = Arc::new(network_proxy_state_for_policy(policy)); + + let mut req = Request::builder() + .method(Method::CONNECT) + .uri("https://api.github.com:8443") + .header("host", "api.github.com:8443") + .body(Body::empty()) + .unwrap(); + req.extensions_mut().insert(state); + + let response = http_connect_accept( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await + .unwrap_err(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response.headers().get("x-proxy-error").unwrap(), + "blocked-by-mitm-required" + ); + } + + #[tokio::test] + async fn brokered_connect_forwards_server_first_opaque_protocol_without_mitm() { + let server_banner = b"SSH-2.0-server\r\n"; + let target_listener = TokioTcpListener::bind((Ipv4Addr::LOCALHOST, 0)) + .await + .expect("target listener should bind"); + let target_addr = target_listener + .local_addr() + .expect("target listener should expose local addr"); + let target_task = tokio::spawn(async move { + let (mut stream, _) = target_listener + .accept() + .await + .expect("target listener should accept"); + stream + .write_all(server_banner) + .await + .expect("target should write opaque server bytes"); + }); + + let state = Arc::new(network_proxy_state_for_policy({ + let mut network = NetworkProxyConfig { + credential_broker: true, + mitm: true, + ..NetworkProxyConfig::default() + }; + network.set_allowed_domains(vec!["127.0.0.1".to_string()]); + network.allow_local_binding = true; + network + })); + let mut env = HashMap::from([ + ("GH_HOST".to_string(), "127.0.0.1".to_string()), + ( + "GH_ENTERPRISE_TOKEN".to_string(), + "ghp-enterprise-real".to_string(), + ), + ]); + state.virtualize_child_credentials(&mut env); + let listener = + StdTcpListener::bind((Ipv4Addr::LOCALHOST, 0)).expect("proxy listener should bind"); + let proxy_addr = listener + .local_addr() + .expect("proxy listener should expose local addr"); + let lifecycle = crate::connection_lifecycle::ConnectionLifecycle::new(); + let proxy_task = tokio::spawn(run_http_proxy_with_std_listener( + state, + listener, + /*policy_decider*/ None, + /*environment_id*/ None, + lifecycle.guard(), + )); + + let mut stream = tokio::net::TcpStream::connect(proxy_addr) + .await + .expect("client should connect to proxy"); + let request = format!( + "CONNECT 127.0.0.1:{port} HTTP/1.1\r\nHost: 127.0.0.1:{port}\r\n\r\n", + port = target_addr.port() + ); + stream + .write_all(request.as_bytes()) + .await + .expect("client should write CONNECT request"); + + let mut buf = [0_u8; 256]; + let bytes_read = timeout(Duration::from_secs(2), stream.read(&mut buf)) + .await + .expect("proxy should respond before timeout") + .expect("client should read proxy response"); + let response = String::from_utf8_lossy(&buf[..bytes_read]); + assert!( + response.starts_with("HTTP/1.1 200 OK\r\n"), + "unexpected proxy response: {response:?}" + ); + + let mut buf = vec![0_u8; server_banner.len()]; + timeout(Duration::from_secs(2), stream.read_exact(&mut buf)) + .await + .expect("opaque server bytes should arrive before timeout") + .expect("client should read opaque server bytes"); + assert_eq!(buf, server_banner); + + drop(stream); + proxy_task.abort(); + let _ = proxy_task.await; + target_task.await.expect("target task should finish"); + } + + #[tokio::test] + async fn http_proxy_blocks_absolute_form_https_for_hooked_host() { + let target_listener = TokioTcpListener::bind((Ipv4Addr::LOCALHOST, 0)) + .await + .expect("target listener should bind"); + let target_addr = target_listener + .local_addr() + .expect("target listener should expose local addr"); + let target_task = tokio::spawn(async move { + timeout(Duration::from_secs(1), target_listener.accept()) + .await + .is_ok() + }); + + let state = Arc::new(network_proxy_state_for_policy({ + let mut network = NetworkProxyConfig { + allow_local_binding: true, + mitm: true, + mitm_hooks: vec![crate::mitm_hook::MitmHookConfig { + host: "127.0.0.1".to_string(), + matcher: crate::mitm_hook::MitmHookMatchConfig { + methods: vec!["GET".to_string()], + path_prefixes: vec!["/repos/openai/ALLOWED".to_string()], + ..crate::mitm_hook::MitmHookMatchConfig::default() + }, + actions: crate::mitm_hook::MitmHookActionsConfig::default(), + }], + ..NetworkProxyConfig::default() + }; + network.set_allowed_domains(vec!["127.0.0.1".to_string()]); + network + })); + let listener = + StdTcpListener::bind((Ipv4Addr::LOCALHOST, 0)).expect("proxy listener should bind"); + let proxy_addr = listener + .local_addr() + .expect("proxy listener should expose local addr"); + let lifecycle = crate::connection_lifecycle::ConnectionLifecycle::new(); + let proxy_task = tokio::spawn(run_http_proxy_with_std_listener( + state.clone(), + listener, + /*policy_decider*/ None, + /*environment_id*/ None, + lifecycle.guard(), + )); + + let mut stream = tokio::net::TcpStream::connect(proxy_addr) + .await + .expect("client should connect to proxy"); + let request = format!( + "GET https://127.0.0.1:{port}/repos/openai/UNAUTHORIZED HTTP/1.1\r\nHost: 127.0.0.1:{port}\r\nConnection: close\r\n\r\n", + port = target_addr.port() + ); + stream + .write_all(request.as_bytes()) + .await + .expect("client should write absolute-form HTTPS request"); + + let mut buf = [0_u8; 512]; + let bytes_read = timeout(Duration::from_secs(2), stream.read(&mut buf)) + .await + .expect("proxy should respond before timeout") + .expect("client should read proxy response"); + let response = String::from_utf8_lossy(&buf[..bytes_read]); + assert!( + response.starts_with("HTTP/1.1 403 Forbidden\r\n"), + "unexpected proxy response: {response:?}" + ); + assert!(response.contains("x-proxy-error: blocked-by-mitm-required\r\n")); + assert!( + !target_task.await.expect("target task should finish"), + "blocked request must not reach upstream" + ); + + let blocked = state.drain_blocked().await.unwrap(); + assert_eq!(blocked.len(), 1); + assert_eq!(blocked[0].reason, REASON_MITM_REQUIRED); + + drop(stream); + proxy_task.abort(); + let _ = proxy_task.await; + } + + #[tokio::test(flavor = "current_thread")] + async fn http_plain_proxy_blocks_unix_socket_when_method_not_allowed() { + let state = Arc::new(network_proxy_state_for_policy(NetworkProxyConfig::default())); + state + .set_network_mode(NetworkMode::Limited) + .await + .expect("network mode should update"); + + let mut req = Request::builder() + .method(Method::POST) + .uri("http://example.com") + .header("x-unix-socket", "/tmp/test.sock") + .body(Body::empty()) + .expect("request should build"); + req.extensions_mut().insert(state); + + let response = http_plain_proxy( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response.headers().get("x-proxy-error").unwrap(), + "blocked-by-method-policy" + ); + } + + #[tokio::test(flavor = "current_thread")] + async fn http_plain_proxy_rejects_unix_socket_when_not_allowlisted() { + let state = Arc::new(network_proxy_state_for_policy(NetworkProxyConfig::default())); + + let mut req = Request::builder() + .method(Method::GET) + .uri("http://example.com") + .header("x-unix-socket", "/tmp/test.sock") + .body(Body::empty()) + .expect("request should build"); + req.extensions_mut().insert(state); + + let response = http_plain_proxy( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await + .unwrap(); + + if cfg!(target_os = "macos") { + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response.headers().get("x-proxy-error").unwrap(), + "blocked-by-allowlist" + ); + } else { + assert_eq!(response.status(), StatusCode::NOT_IMPLEMENTED); + } + } + + #[cfg(target_os = "macos")] + #[tokio::test(flavor = "current_thread")] + async fn http_plain_proxy_attempts_allowed_unix_socket_proxy() { + let state = Arc::new(network_proxy_state_for_policy({ + let mut network = NetworkProxyConfig::default(); + network.set_allow_unix_sockets(vec!["/tmp/test.sock".to_string()]); + network + })); + + let mut req = Request::builder() + .method(Method::GET) + .uri("http://example.com") + .header("x-unix-socket", "/tmp/test.sock") + .body(Body::empty()) + .expect("request should build"); + req.extensions_mut().insert(state); + + let response = http_plain_proxy( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::BAD_GATEWAY); + } + + #[tokio::test] + async fn http_connect_accept_denies_denylisted_host() { + let policy = { + let mut policy = NetworkProxyConfig::default(); + policy.set_allowed_domains(vec!["**.openai.com".to_string()]); + policy.set_denied_domains(vec!["api.openai.com".to_string()]); + policy + }; + let state = Arc::new(network_proxy_state_for_policy(policy)); + + let mut req = Request::builder() + .method(Method::CONNECT) + .uri("https://api.openai.com:443") + .header("host", "api.openai.com:443") + .body(Body::empty()) + .unwrap(); + req.extensions_mut().insert(state); + + let response = http_connect_accept( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await + .unwrap_err(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response.headers().get("x-proxy-error").unwrap(), + "blocked-by-denylist" + ); + } + + #[tokio::test] + async fn http_plain_proxy_rejects_absolute_uri_host_header_mismatch() { + let state = Arc::new(network_proxy_state_for_policy(NetworkProxyConfig::default())); + let mut req = Request::builder() + .method(Method::GET) + .uri("http://raw.githubusercontent.com/openai/codex/main/README.md") + .header(header::HOST, "api.github.com") + .body(Body::empty()) + .unwrap(); + req.extensions_mut().insert(state); + + let response = http_plain_proxy( + /*policy_decider*/ None, /*environment_id*/ None, req, + ) + .await; + assert_eq!(response.unwrap().status(), StatusCode::BAD_REQUEST); + } + + #[test] + fn validate_absolute_form_host_header_allows_matching_default_port() { + let req = Request::builder() + .method(Method::GET) + .uri("http://example.com/") + .header("host", "example.com") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + validate_absolute_form_host_header(&req, &RequestContext::try_from(&req).unwrap(),), + Ok(()) + ); + } + + #[test] + fn validate_absolute_form_host_header_rejects_mismatched_host() { + let req = Request::builder() + .method(Method::GET) + .uri("http://raw.githubusercontent.com/") + .header("host", "api.github.com") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + validate_absolute_form_host_header(&req, &RequestContext::try_from(&req).unwrap(),), + Err("Host header does not match request target") + ); + } + + #[test] + fn validate_absolute_form_host_header_rejects_missing_non_default_port() { + let req = Request::builder() + .method(Method::GET) + .uri("http://example.com:8080/") + .header("host", "example.com") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + validate_absolute_form_host_header(&req, &RequestContext::try_from(&req).unwrap(),), + Err("Host header does not match request target") + ); + } + + #[test] + fn remove_hop_by_hop_request_headers_keeps_forwarding_headers() { + let mut headers = HeaderMap::new(); + headers.insert( + header::CONNECTION, + HeaderValue::from_static("x-hop, keep-alive"), + ); + headers.insert("x-hop", HeaderValue::from_static("1")); + headers.insert( + header::PROXY_AUTHORIZATION, + HeaderValue::from_static("Basic abc"), + ); + headers.insert( + &header::X_FORWARDED_FOR, + HeaderValue::from_static("127.0.0.1"), + ); + headers.insert(header::HOST, HeaderValue::from_static("example.com")); + + remove_hop_by_hop_request_headers(&mut headers); + + assert_eq!(headers.get(header::CONNECTION), None); + assert_eq!(headers.get("x-hop"), None); + assert_eq!(headers.get(header::PROXY_AUTHORIZATION), None); + assert_eq!( + headers.get(&header::X_FORWARDED_FOR), + Some(&HeaderValue::from_static("127.0.0.1")) + ); + assert_eq!( + headers.get(header::HOST), + Some(&HeaderValue::from_static("example.com")) + ); + } +} diff --git a/codex-rs/network-proxy/src/lib.rs b/codex-rs/network-proxy/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..dd8df15aca051c04b0af568c401798d5c0b8799d --- /dev/null +++ b/codex-rs/network-proxy/src/lib.rs @@ -0,0 +1,119 @@ +#![deny(clippy::print_stdout, clippy::print_stderr)] + +mod attribution; +mod authorization_path; +mod brokered_tunnel; +mod certs; +mod config; +mod connect_policy; +mod connection_lifecycle; +mod credential_broker; +mod environment_policy; +mod http_proxy; +mod mitm; +mod mitm_hook; +mod native_certs; +mod network_policy; +mod policy; +mod process_log_metadata; +mod proxy; +mod reasons; +mod remote_config; +mod request_cancellation; +mod request_disconnect; +mod responses; +mod runtime; +mod socks5; +mod state; +mod upstream; +#[cfg(target_os = "windows")] +mod windows_proxy_ingress; +#[cfg(target_os = "windows")] +mod windows_tcp_attribution; + +pub use attribution::PROXY_ATTRIBUTION_TOKEN_ENV_KEY; +pub use attribution::write_attribution_frame; +pub use certs::CUSTOM_CA_ENV_KEYS; +pub use certs::is_managed_mitm_ca_trust_bundle_path; +pub use config::NetworkDomainPermission; +pub use config::NetworkDomainPermissionEntry; +pub use config::NetworkDomainPermissions; +pub use config::NetworkMode; +pub use config::NetworkProxyConfig; +pub use config::NetworkUnixSocketPermission; +pub use config::NetworkUnixSocketPermissions; +pub use config::host_and_port_from_network_addr; +pub use config::managed_proxy_ports; +pub use credential_broker::CREDENTIAL_BROKER_ACTIVE_ENV_KEY; +pub use credential_broker::CredentialAuthMethod; +pub use credential_broker::CredentialBrokerContext; +pub use credential_broker::CredentialBrokerEnvironment; +pub use credential_broker::CredentialProviderConfig; +pub use credential_broker::brokered_credential_binding_env_keys; +pub use credential_broker::brokered_credential_dummy_env_keys; +pub use credential_broker::brokered_credential_env_keys; +pub use credential_broker::brokered_credential_marker_env_keys; +pub use credential_broker::brokered_credential_value_env_keys; +pub use credential_broker::credential_broker_provider_context_env_keys; +pub use credential_broker::credential_broker_provider_sources_allowed; +pub use credential_broker::is_credential_broker_provider_env_key; +pub use environment_policy::EnvironmentNetworkPolicy; +pub use mitm_hook::InjectedHeaderConfig; +pub use mitm_hook::MitmHookActionsConfig; +pub use mitm_hook::MitmHookBodyConfig; +pub use mitm_hook::MitmHookConfig; +pub use mitm_hook::MitmHookMatchConfig; +pub use network_policy::NetworkDecision; +pub use network_policy::NetworkDecisionSource; +pub use network_policy::NetworkPolicyAuditEvent; +pub use network_policy::NetworkPolicyAuditObserver; +pub use network_policy::NetworkPolicyDecider; +pub use network_policy::NetworkPolicyDeciderFuture; +pub use network_policy::NetworkPolicyDecision; +pub use network_policy::NetworkPolicyRequest; +pub use network_policy::NetworkPolicyRequestArgs; +pub use network_policy::NetworkProtocol; +pub use policy::normalize_host; +pub use process_log_metadata::ExecutorLogIdentity; +pub use process_log_metadata::NetworkProxyProcessLogMetadata; +pub use proxy::ALL_PROXY_ENV_KEYS; +pub use proxy::ALLOW_LOCAL_BINDING_ENV_KEY; +pub use proxy::Args; +#[cfg(target_os = "macos")] +pub use proxy::CODEX_PROXY_GIT_SSH_COMMAND_MARKER; +pub use proxy::DEFAULT_NO_PROXY_VALUE; +pub use proxy::ManagedNetworkSandboxContext; +pub use proxy::ManagedProxyRouting; +pub use proxy::NO_PROXY_ENV_KEYS; +pub use proxy::NetworkProxy; +pub use proxy::NetworkProxyBuilder; +pub use proxy::NetworkProxyHandle; +pub use proxy::PROXY_ACTIVE_ENV_KEY; +pub use proxy::PROXY_ENV_KEYS; +#[cfg(target_os = "macos")] +pub use proxy::PROXY_GIT_SSH_COMMAND_ENV_KEY; +pub use proxy::PROXY_URL_ENV_KEYS; +pub use proxy::PreparedManagedNetwork; +pub use proxy::has_proxy_url_env_vars; +pub use proxy::is_managed_proxy_env_var; +pub use proxy::proxy_url_env_value; +pub use proxy::strip_managed_proxy_env; +pub use remote_config::RemoteNetworkProxyConfig; +pub use remote_config::RemoteNetworkProxyLaunchConfig; +pub use request_cancellation::NetworkRequestCancellation; +pub use request_cancellation::NetworkRequestCancellationReason; +pub use request_disconnect::NetworkRequestDisconnect; +pub use runtime::BlockedRequest; +pub use runtime::BlockedRequestArgs; +pub use runtime::BlockedRequestObserver; +pub use runtime::BlockedRequestObserverFuture; +pub use runtime::ConfigReloader; +pub use runtime::ConfigReloaderFuture; +pub use runtime::ConfigState; +pub use runtime::NetworkProxyState; +pub use state::NetworkProxyAuditMetadata; +pub use state::NetworkProxyConstraintError; +pub use state::NetworkProxyConstraints; +pub use state::PartialNetworkProxyConfig; +pub use state::build_config_state; +pub use state::validate_policy_against_constraints; diff --git a/codex-rs/network-proxy/src/mitm_hook.rs b/codex-rs/network-proxy/src/mitm_hook.rs new file mode 100644 index 0000000000000000000000000000000000000000..6e2ee5cf8a1d2fd195c1daafe1f42de7773822f7 --- /dev/null +++ b/codex-rs/network-proxy/src/mitm_hook.rs @@ -0,0 +1,1086 @@ +#![cfg_attr(not(test), allow(dead_code))] + +use crate::authorization_path::is_safe_for_authorization; +use crate::config::NetworkProxyConfig; +use crate::policy::normalize_host; +use anyhow::Context as _; +use anyhow::Result; +use anyhow::anyhow; +use codex_utils_absolute_path::AbsolutePathBuf; +use globset::GlobBuilder; +use globset::GlobMatcher; +use rama_http::HeaderValue; +use rama_http::Request; +use rama_http::header::HeaderName; +use serde::Deserialize; +use serde::Serialize; +use std::collections::BTreeMap; +use std::env; +use std::fs; +use std::path::Path; +use url::form_urlencoded; + +const PATTERN_PREFIX: &str = "pattern:"; +const LITERAL_PREFIX: &str = "literal:"; + +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)] +#[serde(default)] +pub struct MitmHookConfig { + pub host: String, + #[serde(rename = "match", default)] + pub matcher: MitmHookMatchConfig, + #[serde(default)] + pub actions: MitmHookActionsConfig, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)] +#[serde(default)] +pub struct MitmHookMatchConfig { + pub methods: Vec, + pub path_prefixes: Vec, + pub query: BTreeMap>, + pub headers: BTreeMap>, + pub body: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)] +#[serde(default)] +pub struct MitmHookActionsConfig { + pub strip_request_headers: Vec, + pub inject_request_headers: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)] +#[serde(default)] +pub struct InjectedHeaderConfig { + pub name: String, + pub secret_env_var: Option, + pub secret_file: Option, + pub prefix: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(transparent)] +pub struct MitmHookBodyConfig(pub serde_json::Value); + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MitmHook { + pub host: String, + pub matcher: MitmHookMatcher, + pub actions: MitmHookActions, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MitmHookMatcher { + pub methods: Vec, + pub path_prefixes: Vec, + pub query: Vec, + pub headers: Vec, + pub body: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct QueryConstraint { + pub name: String, + pub allowed_values: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct HeaderConstraint { + pub name: HeaderName, + pub allowed_values: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MitmHookActions { + pub strip_request_headers: Vec, + pub inject_request_headers: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ResolvedInjectedHeader { + pub name: HeaderName, + pub value: HeaderValue, + pub source: SecretSource, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum SecretSource { + EnvVar(String), + File(AbsolutePathBuf), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MitmHookBodyMatcher { + pub raw: serde_json::Value, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum PathMatcher { + Prefix(String), + Glob(CompiledGlobMatcher), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ValueMatcher { + Exact(String), + Glob(CompiledGlobMatcher), +} + +enum MatcherPattern<'a> { + Literal(&'a str), + Glob(&'a str), +} + +#[derive(Clone)] +pub struct CompiledGlobMatcher { + pattern: String, + matcher: GlobMatcher, +} + +impl std::fmt::Debug for CompiledGlobMatcher { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("CompiledGlobMatcher") + .field("pattern", &self.pattern) + .finish() + } +} + +impl PartialEq for CompiledGlobMatcher { + fn eq(&self, other: &Self) -> bool { + self.pattern == other.pattern + } +} + +impl Eq for CompiledGlobMatcher {} + +impl CompiledGlobMatcher { + fn is_match(&self, candidate: &str) -> bool { + self.matcher.is_match(candidate) + } +} + +pub type MitmHooksByHost = BTreeMap>; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum HookEvaluation { + NoHooksForHost, + Matched { actions: MitmHookActions }, + HookedHostNoMatch, +} + +pub(crate) fn validate_mitm_hook_config(config: &NetworkProxyConfig) -> Result<()> { + let hooks = &config.mitm_hooks; + if hooks.is_empty() { + return Ok(()); + } + + if !config.mitm { + return Err(anyhow!("network.mitm_hooks requires network.mitm = true")); + } + + for (hook_index, hook) in hooks.iter().enumerate() { + let host = normalize_hook_host(&hook.host) + .with_context(|| format!("invalid network.mitm_hooks[{hook_index}].host"))?; + + let methods = normalize_methods(&hook.matcher.methods) + .with_context(|| format!("invalid network.mitm_hooks[{hook_index}].match.methods"))?; + if methods.is_empty() { + return Err(anyhow!( + "network.mitm_hooks[{hook_index}].match.methods must not be empty" + )); + } + + let path_prefixes = + compile_path_matchers(&hook.matcher.path_prefixes).with_context(|| { + format!("invalid network.mitm_hooks[{hook_index}].match.path_prefixes") + })?; + if path_prefixes.is_empty() { + return Err(anyhow!( + "network.mitm_hooks[{hook_index}].match.path_prefixes must not be empty" + )); + } + + if let Some(body) = hook.matcher.body.as_ref() { + let _ = body; + return Err(anyhow!( + "network.mitm_hooks[{hook_index}].match.body is reserved for a future release and is not yet supported" + )); + } + + validate_query_constraints(&hook.matcher.query) + .with_context(|| format!("invalid network.mitm_hooks[{hook_index}].match.query"))?; + validate_header_constraints(&hook.matcher.headers) + .with_context(|| format!("invalid network.mitm_hooks[{hook_index}].match.headers"))?; + validate_strip_request_headers(&hook.actions.strip_request_headers).with_context(|| { + format!("invalid network.mitm_hooks[{hook_index}].actions.strip_request_headers") + })?; + validate_injected_headers(&hook.actions.inject_request_headers).with_context(|| { + format!("invalid network.mitm_hooks[{hook_index}].actions.inject_request_headers") + })?; + + if host.is_empty() { + return Err(anyhow!( + "network.mitm_hooks[{hook_index}].host must not be empty" + )); + } + } + + Ok(()) +} + +pub(crate) fn compile_mitm_hooks(config: &NetworkProxyConfig) -> Result { + compile_mitm_hooks_with_resolvers( + config, + |name| env::var(name).ok(), + |path| { + let value = fs::read_to_string(path.as_path()).with_context(|| { + format!("failed to read secret file {}", path.as_path().display()) + })?; + Ok(value.trim().to_string()) + }, + ) +} + +pub(crate) fn evaluate_mitm_hooks( + hooks_by_host: &MitmHooksByHost, + host: &str, + req: &Request, +) -> HookEvaluation { + let normalized_host = normalize_host(host); + let Some(hooks) = hooks_by_host.get(&normalized_host) else { + return HookEvaluation::NoHooksForHost; + }; + + for hook in hooks { + if hook_matches(hook, req) { + return HookEvaluation::Matched { + actions: hook.actions.clone(), + }; + } + } + + HookEvaluation::HookedHostNoMatch +} + +fn compile_mitm_hooks_with_resolvers( + config: &NetworkProxyConfig, + resolve_env_var: EnvFn, + read_secret_file: FileFn, +) -> Result +where + EnvFn: Fn(&str) -> Option, + FileFn: Fn(&AbsolutePathBuf) -> Result, +{ + validate_mitm_hook_config(config)?; + + let mut hooks_by_host = MitmHooksByHost::new(); + for hook in &config.mitm_hooks { + let host = normalize_hook_host(&hook.host)?; + let methods = normalize_methods(&hook.matcher.methods)?; + let path_prefixes = compile_path_matchers(&hook.matcher.path_prefixes)?; + let query = hook + .matcher + .query + .iter() + .map(|(name, values)| { + Ok(QueryConstraint { + name: normalize_query_name(name)?, + allowed_values: compile_value_matchers(values)?, + }) + }) + .collect::>>()?; + let headers = hook + .matcher + .headers + .iter() + .map(|(name, values)| { + Ok(HeaderConstraint { + name: parse_header_name(name)?, + allowed_values: compile_value_matchers(values)?, + }) + }) + .collect::>>()?; + let strip_request_headers = hook + .actions + .strip_request_headers + .iter() + .map(|name| parse_header_name(name)) + .collect::>>()?; + let inject_request_headers = hook + .actions + .inject_request_headers + .iter() + .map(|header| { + compile_injected_header(header, &resolve_env_var, &read_secret_file) + .with_context(|| format!("failed to compile injected header {}", header.name)) + }) + .collect::>>()?; + + hooks_by_host + .entry(host.clone()) + .or_default() + .push(MitmHook { + host, + matcher: MitmHookMatcher { + methods, + path_prefixes, + query, + headers, + body: None, + }, + actions: MitmHookActions { + strip_request_headers, + inject_request_headers, + }, + }); + } + + Ok(hooks_by_host) +} + +fn compile_injected_header( + header: &InjectedHeaderConfig, + resolve_env_var: &EnvFn, + read_secret_file: &FileFn, +) -> Result +where + EnvFn: Fn(&str) -> Option, + FileFn: Fn(&AbsolutePathBuf) -> Result, +{ + let name = parse_header_name(&header.name)?; + let (secret, source) = match ( + header.secret_env_var.as_deref(), + header.secret_file.as_deref(), + ) { + (Some(env_var), None) => { + let value = resolve_env_var(env_var) + .ok_or_else(|| anyhow!("missing required environment variable {env_var}"))?; + (value, SecretSource::EnvVar(env_var.to_string())) + } + (None, Some(secret_file)) => { + let path = parse_secret_file(secret_file)?; + let value = read_secret_file(&path)?; + (value, SecretSource::File(path)) + } + _ => { + return Err(anyhow!( + "expected exactly one of secret_env_var or secret_file" + )); + } + }; + + let prefix = header.prefix.clone().unwrap_or_default(); + let value = HeaderValue::from_str(&format!("{prefix}{secret}")) + .with_context(|| format!("invalid value for injected header {}", header.name))?; + + Ok(ResolvedInjectedHeader { + name, + value, + source, + }) +} + +fn hook_matches(hook: &MitmHook, req: &Request) -> bool { + let method = req.method().as_str().to_ascii_uppercase(); + if !hook + .matcher + .methods + .iter() + .any(|allowed| allowed == &method) + { + return false; + } + + let path = req.uri().path(); + if !is_safe_for_authorization(path) || !path_matches(&hook.matcher.path_prefixes, path) { + return false; + } + + if !query_matches(&hook.matcher.query, req) { + return false; + } + + headers_match(&hook.matcher.headers, req) +} + +fn query_matches(query_constraints: &[QueryConstraint], req: &Request) -> bool { + if query_constraints.is_empty() { + return true; + } + + let actual_query = req.uri().query().unwrap_or_default(); + let mut actual_values: BTreeMap> = BTreeMap::new(); + for (name, value) in form_urlencoded::parse(actual_query.as_bytes()) { + actual_values + .entry(name.into_owned()) + .or_default() + .push(value.into_owned()); + } + + query_constraints.iter().all(|constraint| { + actual_values.get(&constraint.name).is_some_and(|actual| { + actual.iter().any(|candidate| { + constraint + .allowed_values + .iter() + .any(|allowed| allowed.matches(candidate)) + }) + }) + }) +} + +fn headers_match(header_constraints: &[HeaderConstraint], req: &Request) -> bool { + header_constraints.iter().all(|constraint| { + let actual = req.headers().get_all(&constraint.name); + if actual.iter().next().is_none() { + return false; + } + if constraint.allowed_values.is_empty() { + return true; + } + + actual.iter().any(|value| { + value.to_str().ok().is_some_and(|candidate| { + constraint + .allowed_values + .iter() + .any(|allowed| allowed.matches(candidate)) + }) + }) + }) +} + +fn path_matches(path_prefixes: &[PathMatcher], path: &str) -> bool { + path_prefixes.iter().any(|matcher| matcher.matches(path)) +} + +impl PathMatcher { + fn matches(&self, candidate: &str) -> bool { + match self { + Self::Prefix(prefix) => candidate.starts_with(prefix), + Self::Glob(glob) => glob.is_match(candidate), + } + } +} + +impl ValueMatcher { + fn matches(&self, candidate: &str) -> bool { + match self { + Self::Exact(value) => value == candidate, + Self::Glob(glob) => glob.is_match(candidate), + } + } +} + +fn compile_path_matchers(path_prefixes: &[String]) -> Result> { + path_prefixes + .iter() + .map(|prefix| { + match parse_matcher_pattern(prefix)? { + MatcherPattern::Literal(prefix) => { + if prefix.is_empty() { + return Err(anyhow!("path_prefixes must not contain empty entries")); + } + Ok(PathMatcher::Prefix(prefix.to_string())) + } + MatcherPattern::Glob(glob_pattern) => Ok(PathMatcher::Glob(compile_glob_matcher( + glob_pattern, + /*literal_separator*/ true, + )?)), + } + }) + .collect() +} + +fn compile_value_matchers(values: &[String]) -> Result> { + values + .iter() + .map(|value| match parse_matcher_pattern(value)? { + MatcherPattern::Literal(value) => Ok(ValueMatcher::Exact(value.to_string())), + MatcherPattern::Glob(glob_pattern) => Ok(ValueMatcher::Glob(compile_glob_matcher( + glob_pattern, + /*literal_separator*/ false, + )?)), + }) + .collect() +} + +fn parse_matcher_pattern(pattern: &str) -> Result> { + if let Some(literal) = pattern.strip_prefix(LITERAL_PREFIX) { + return Ok(MatcherPattern::Literal(literal)); + } + let Some(glob_pattern) = pattern.strip_prefix(PATTERN_PREFIX) else { + return Ok(MatcherPattern::Literal(pattern)); + }; + if glob_pattern.is_empty() { + return Err(anyhow!("glob pattern must not be empty")); + } + Ok(MatcherPattern::Glob(glob_pattern)) +} + +fn compile_glob_matcher(pattern: &str, literal_separator: bool) -> Result { + let mut builder = GlobBuilder::new(pattern); + builder + .backslash_escape(true) + .literal_separator(literal_separator); + builder + .build() + .map(|glob| CompiledGlobMatcher { + pattern: pattern.to_string(), + matcher: glob.compile_matcher(), + }) + .map_err(|err| anyhow!("invalid glob pattern {pattern:?}: {err}")) +} + +fn normalize_hook_host(host: &str) -> Result { + let normalized = normalize_host(host); + if normalized.is_empty() { + return Err(anyhow!("host must not be empty")); + } + if normalized.contains('*') { + return Err(anyhow!( + "MITM hook hosts must be exact hosts and cannot contain wildcards" + )); + } + Ok(normalized) +} + +fn normalize_methods(methods: &[String]) -> Result> { + methods + .iter() + .map(|method| { + let normalized = method.trim().to_ascii_uppercase(); + if normalized.is_empty() { + return Err(anyhow!("methods must not contain empty entries")); + } + Ok(normalized) + }) + .collect() +} + +fn validate_query_constraints(query: &BTreeMap>) -> Result<()> { + for (name, values) in query { + let normalized = normalize_query_name(name)?; + if normalized.is_empty() { + return Err(anyhow!("query keys must not be empty")); + } + if values.is_empty() { + return Err(anyhow!( + "query key {name:?} must list at least one allowed value" + )); + } + let _ = compile_value_matchers(values) + .with_context(|| format!("invalid matcher for query key {name:?}"))?; + } + Ok(()) +} + +fn normalize_query_name(name: &str) -> Result { + if name.is_empty() { + return Err(anyhow!("query keys must not be empty")); + } + Ok(name.to_string()) +} + +fn validate_header_constraints(headers: &BTreeMap>) -> Result<()> { + for (name, values) in headers { + let _ = parse_header_name(name)?; + let _ = compile_value_matchers(values) + .with_context(|| format!("invalid matcher for header {name:?}"))?; + } + Ok(()) +} + +fn validate_strip_request_headers(header_names: &[String]) -> Result<()> { + for name in header_names { + let _ = parse_header_name(name)?; + } + Ok(()) +} + +fn validate_injected_headers(headers: &[InjectedHeaderConfig]) -> Result<()> { + for header in headers { + let _ = parse_header_name(&header.name)?; + match ( + header.secret_env_var.as_deref(), + header.secret_file.as_deref(), + ) { + (Some(secret_env_var), None) => { + if secret_env_var.trim().is_empty() { + return Err(anyhow!("secret_env_var must not be empty")); + } + } + (None, Some(secret_file)) => { + let _ = parse_secret_file(secret_file)?; + } + _ => { + return Err(anyhow!( + "expected exactly one of secret_env_var or secret_file" + )); + } + } + } + Ok(()) +} + +fn parse_header_name(name: &str) -> Result { + HeaderName::from_bytes(name.as_bytes()) + .map_err(|err| anyhow!("invalid header name {name:?}: {err}")) +} + +fn parse_secret_file(path: &str) -> Result { + if path.trim().is_empty() { + return Err(anyhow!("secret_file must not be empty")); + } + let path = Path::new(path); + if !path.is_absolute() { + return Err(anyhow!("secret_file must be an absolute path: {path:?}")); + } + AbsolutePathBuf::from_absolute_path(path) + .with_context(|| format!("secret_file must be an absolute path: {path:?}")) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::NetworkMode; + use crate::config::NetworkProxyConfig; + use pretty_assertions::assert_eq; + use rama_http::Body; + use rama_http::Method; + use tempfile::NamedTempFile; + + fn base_config() -> NetworkProxyConfig { + NetworkProxyConfig { + mitm: true, + mode: NetworkMode::Limited, + ..NetworkProxyConfig::default() + } + } + + fn github_hook() -> MitmHookConfig { + MitmHookConfig { + host: "api.github.com".to_string(), + matcher: MitmHookMatchConfig { + methods: vec!["POST".to_string(), "PUT".to_string()], + path_prefixes: vec!["/repos/openai/".to_string()], + ..MitmHookMatchConfig::default() + }, + actions: MitmHookActionsConfig { + strip_request_headers: vec!["authorization".to_string()], + inject_request_headers: vec![InjectedHeaderConfig { + name: "authorization".to_string(), + secret_env_var: Some("CODEX_GITHUB_TOKEN".to_string()), + secret_file: None, + prefix: Some("Bearer ".to_string()), + }], + }, + } + } + + #[test] + fn validate_requires_mitm_for_hooks() { + let mut config = base_config(); + config.mitm = false; + config.mitm_hooks = vec![github_hook()]; + + let err = validate_mitm_hook_config(&config).expect_err("hooks require mitm"); + assert!( + err.to_string() + .contains("network.mitm_hooks requires network.mitm = true") + ); + } + + #[test] + fn validate_allows_hooks_in_full_mode() { + let mut config = base_config(); + config.mode = NetworkMode::Full; + config.mitm_hooks = vec![github_hook()]; + + validate_mitm_hook_config(&config).expect("hooks should be allowed in full mode"); + } + + #[test] + fn validate_rejects_body_matchers_for_now() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.matcher.body = Some(MitmHookBodyConfig(serde_json::json!({ + "repository": "openai/codex" + }))); + config.mitm_hooks = vec![hook]; + + let err = validate_mitm_hook_config(&config).expect_err("body matchers are reserved"); + assert!(err.to_string().contains("match.body is reserved")); + } + + #[test] + fn validate_rejects_relative_secret_file() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.actions.inject_request_headers[0].secret_env_var = None; + hook.actions.inject_request_headers[0].secret_file = Some("token.txt".to_string()); + config.mitm_hooks = vec![hook]; + + let err = validate_mitm_hook_config(&config).expect_err("secret file must be absolute"); + assert!(format!("{err:#}").contains("secret_file must be an absolute path")); + } + + #[test] + fn validate_rejects_dual_secret_sources() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.actions.inject_request_headers[0].secret_file = Some("/tmp/github-token".to_string()); + config.mitm_hooks = vec![hook]; + + let err = validate_mitm_hook_config(&config).expect_err("dual secret sources invalid"); + assert!(format!("{err:#}").contains("exactly one of secret_env_var or secret_file")); + } + + #[test] + fn compile_resolves_env_backed_injected_headers() { + let mut config = base_config(); + config.mitm_hooks = vec![github_hook()]; + + let hooks = compile_mitm_hooks_with_resolvers( + &config, + |name| (name == "CODEX_GITHUB_TOKEN").then(|| "ghp-secret".to_string()), + |_| Err(anyhow!("unexpected file lookup")), + ) + .unwrap(); + + let compiled = hooks.get("api.github.com").unwrap(); + assert_eq!(compiled.len(), 1); + assert_eq!( + compiled[0].actions.inject_request_headers[0].source, + SecretSource::EnvVar("CODEX_GITHUB_TOKEN".to_string()) + ); + assert_eq!( + compiled[0].actions.inject_request_headers[0].value, + HeaderValue::from_static("Bearer ghp-secret") + ); + } + + #[test] + fn compile_resolves_file_backed_injected_headers() { + let secret_file = NamedTempFile::new().unwrap(); + std::fs::write(secret_file.path(), "ghp-file-secret\n").unwrap(); + + let mut config = base_config(); + let mut hook = github_hook(); + hook.actions.inject_request_headers[0].secret_env_var = None; + hook.actions.inject_request_headers[0].secret_file = + Some(secret_file.path().display().to_string()); + config.mitm_hooks = vec![hook]; + + let hooks = compile_mitm_hooks(&config).unwrap(); + let compiled = hooks.get("api.github.com").unwrap(); + assert_eq!( + compiled[0].actions.inject_request_headers[0].value, + HeaderValue::from_static("Bearer ghp-file-secret") + ); + } + + #[test] + fn evaluate_returns_first_matching_hook() { + let mut config = base_config(); + let mut first = github_hook(); + first.matcher.path_prefixes = vec!["/repos/openai/".to_string()]; + let mut second = github_hook(); + second.actions.inject_request_headers[0].prefix = Some("Token ".to_string()); + config.mitm_hooks = vec![first, second]; + + let hooks = compile_mitm_hooks_with_resolvers( + &config, + |_| Some("abc".to_string()), + |_| Err(anyhow!("unexpected file lookup")), + ) + .unwrap(); + let req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/codex/issues") + .header("x-trace", "1") + .body(Body::empty()) + .unwrap(); + + let evaluation = evaluate_mitm_hooks(&hooks, "api.github.com", &req); + let HookEvaluation::Matched { actions } = evaluation else { + panic!("expected a matching hook"); + }; + + assert_eq!( + actions.inject_request_headers[0].value, + HeaderValue::from_static("Bearer abc") + ); + } + + #[test] + fn evaluate_matches_query_and_header_constraints() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.matcher.query = BTreeMap::from([( + "state".to_string(), + vec!["open".to_string(), "triage".to_string()], + )]); + hook.matcher.headers = BTreeMap::from([( + "x-github-api-version".to_string(), + vec!["2022-11-28".to_string()], + )]); + config.mitm_hooks = vec![hook]; + + let hooks = compile_mitm_hooks_with_resolvers( + &config, + |_| Some("abc".to_string()), + |_| Err(anyhow!("unexpected file lookup")), + ) + .unwrap(); + let req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/codex/issues?state=open&per_page=10") + .header("x-github-api-version", "2022-11-28") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + evaluate_mitm_hooks(&hooks, "api.github.com", &req), + HookEvaluation::Matched { + actions: hooks.get("api.github.com").unwrap()[0].actions.clone(), + } + ); + } + + #[test] + fn evaluate_matches_wildcard_path_query_and_header_constraints() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.matcher.path_prefixes = vec!["pattern:/repos/*/codex/issues*".to_string()]; + hook.matcher.query = + BTreeMap::from([("state".to_string(), vec!["pattern:op*".to_string()])]); + hook.matcher.headers = BTreeMap::from([( + "x-github-api-version".to_string(), + vec!["pattern:2022*preview".to_string()], + )]); + config.mitm_hooks = vec![hook]; + + let hooks = compile_mitm_hooks_with_resolvers( + &config, + |_| Some("abc".to_string()), + |_| Err(anyhow!("unexpected file lookup")), + ) + .unwrap(); + let req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/codex/issues?state=open") + .header("x-github-api-version", "2022-11-28-preview") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + evaluate_mitm_hooks(&hooks, "api.github.com", &req), + HookEvaluation::Matched { + actions: hooks.get("api.github.com").unwrap()[0].actions.clone(), + } + ); + } + + #[test] + fn validate_rejects_invalid_wildcard_path_pattern() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.matcher.path_prefixes = vec!["pattern:/repos/[".to_string()]; + config.mitm_hooks = vec![hook]; + + let err = validate_mitm_hook_config(&config).expect_err("invalid glob should fail"); + assert!(format!("{err:#}").contains("invalid glob pattern")); + } + + #[test] + fn evaluate_path_wildcard_does_not_cross_segment_boundaries() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.matcher.path_prefixes = vec!["pattern:/repos/*/codex/issues*".to_string()]; + config.mitm_hooks = vec![hook]; + + let hooks = compile_mitm_hooks_with_resolvers( + &config, + |_| Some("abc".to_string()), + |_| Err(anyhow!("unexpected file lookup")), + ) + .unwrap(); + let nested_req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/private/codex/issues") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + evaluate_mitm_hooks(&hooks, "api.github.com", &nested_req), + HookEvaluation::HookedHostNoMatch + ); + } + + #[test] + fn evaluate_rejects_paths_that_upstream_may_normalize() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.matcher.methods = vec!["GET".to_string()]; + hook.matcher.path_prefixes = vec!["pattern:/openai/openai/**".to_string()]; + config.mitm_hooks = vec![hook]; + + let hooks = compile_mitm_hooks_with_resolvers( + &config, + |_| Some("abc".to_string()), + |_| Err(anyhow!("unexpected file lookup")), + ) + .unwrap(); + let paths = [ + "/openai/openai/../codex", + "/openai/openai/%2e%2e/codex", + "/openai/openai/%2E%2E/codex", + "/openai/openai/.%2e/codex", + "/openai/openai/%2e./codex", + "/openai/openai/%252e%252e/codex", + "/openai/openai/%2f..%2fcodex", + "/openai/openai/%5c..%5ccodex", + "/openai/openai/%2e%2e/%2e%2e/microsoft/vscode", + ]; + let actual = paths + .iter() + .map(|path| { + let req = Request::builder() + .method(Method::GET) + .uri(*path) + .body(Body::empty()) + .unwrap(); + evaluate_mitm_hooks(&hooks, "api.github.com", &req) + }) + .collect::>(); + + assert_eq!(actual, vec![HookEvaluation::HookedHostNoMatch; paths.len()]); + } + + #[test] + fn evaluate_treats_glob_metacharacters_as_literal_without_glob_prefix() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.matcher.path_prefixes = vec!["/repos/[draft]/".to_string()]; + hook.matcher.query = BTreeMap::from([("state".to_string(), vec!["op*".to_string()])]); + hook.matcher.headers = BTreeMap::from([( + "x-github-api-version".to_string(), + vec!["2022-11-28[preview]".to_string()], + )]); + config.mitm_hooks = vec![hook]; + + let hooks = compile_mitm_hooks_with_resolvers( + &config, + |_| Some("abc".to_string()), + |_| Err(anyhow!("unexpected file lookup")), + ) + .unwrap(); + let exact_req = Request::builder() + .method(Method::POST) + .uri("/repos/[draft]/codex/issues?state=op*") + .header("x-github-api-version", "2022-11-28[preview]") + .body(Body::empty()) + .unwrap(); + let non_literal_req = Request::builder() + .method(Method::POST) + .uri("/repos/draft/codex/issues?state=open") + .header("x-github-api-version", "2022-11-28-preview") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + evaluate_mitm_hooks(&hooks, "api.github.com", &exact_req), + HookEvaluation::Matched { + actions: hooks.get("api.github.com").unwrap()[0].actions.clone(), + } + ); + assert_eq!( + evaluate_mitm_hooks(&hooks, "api.github.com", &non_literal_req), + HookEvaluation::HookedHostNoMatch + ); + } + + #[test] + fn evaluate_allows_literal_values_with_reserved_prefixes() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.matcher.query = + BTreeMap::from([("state".to_string(), vec!["literal:pattern:*".to_string()])]); + hook.matcher.headers = BTreeMap::from([( + "x-github-api-version".to_string(), + vec!["literal:pattern:*".to_string()], + )]); + config.mitm_hooks = vec![hook]; + + let hooks = compile_mitm_hooks_with_resolvers( + &config, + |_| Some("abc".to_string()), + |_| Err(anyhow!("unexpected file lookup")), + ) + .unwrap(); + let exact_req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/codex/issues?state=pattern%3A%2A") + .header("x-github-api-version", "pattern:*") + .body(Body::empty()) + .unwrap(); + let non_literal_req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/codex/issues?state=pattern%3Aopen") + .header("x-github-api-version", "pattern:preview") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + evaluate_mitm_hooks(&hooks, "api.github.com", &exact_req), + HookEvaluation::Matched { + actions: hooks.get("api.github.com").unwrap()[0].actions.clone(), + } + ); + assert_eq!( + evaluate_mitm_hooks(&hooks, "api.github.com", &non_literal_req), + HookEvaluation::HookedHostNoMatch + ); + } + + #[test] + fn evaluate_returns_hooked_host_no_match_when_query_constraint_fails() { + let mut config = base_config(); + let mut hook = github_hook(); + hook.matcher.query = BTreeMap::from([("state".to_string(), vec!["open".to_string()])]); + config.mitm_hooks = vec![hook]; + + let hooks = compile_mitm_hooks_with_resolvers( + &config, + |_| Some("abc".to_string()), + |_| Err(anyhow!("unexpected file lookup")), + ) + .unwrap(); + let req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/codex/issues?state=closed") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + evaluate_mitm_hooks(&hooks, "api.github.com", &req), + HookEvaluation::HookedHostNoMatch + ); + } + + #[test] + fn evaluate_returns_no_hooks_for_unconfigured_host() { + let req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/codex/issues") + .body(Body::empty()) + .unwrap(); + + assert_eq!( + evaluate_mitm_hooks(&MitmHooksByHost::new(), "api.github.com", &req), + HookEvaluation::NoHooksForHost + ); + } +} diff --git a/codex-rs/network-proxy/src/mitm_tests.rs b/codex-rs/network-proxy/src/mitm_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..ce02f0e0131c84b6430e038514562733bee592a7 --- /dev/null +++ b/codex-rs/network-proxy/src/mitm_tests.rs @@ -0,0 +1,392 @@ +use super::*; + +use crate::config::NetworkProxyConfig; +use crate::reasons::REASON_METHOD_NOT_ALLOWED; +use crate::reasons::REASON_MITM_HOOK_DENIED; +use crate::reasons::REASON_NOT_ALLOWED_LOCAL; +use crate::runtime::network_proxy_state_for_policy; +use codex_utils_absolute_path::AbsolutePathBuf; +use pretty_assertions::assert_eq; +use rama_http::Body; +use rama_http::HeaderMap; +use rama_http::HeaderValue; +use rama_http::Method; +use rama_http::Request; +use rama_http::StatusCode; +use rama_http::header::HeaderName; +use tempfile::NamedTempFile; + +fn github_write_hook() -> crate::mitm_hook::MitmHookConfig { + crate::mitm_hook::MitmHookConfig { + host: "api.github.com".to_string(), + matcher: crate::mitm_hook::MitmHookMatchConfig { + methods: vec!["POST".to_string(), "PUT".to_string()], + path_prefixes: vec!["/repos/openai/".to_string()], + ..crate::mitm_hook::MitmHookMatchConfig::default() + }, + actions: crate::mitm_hook::MitmHookActionsConfig { + strip_request_headers: vec!["authorization".to_string()], + inject_request_headers: vec![crate::mitm_hook::InjectedHeaderConfig { + name: "authorization".to_string(), + secret_env_var: Some("CODEX_GITHUB_TOKEN".to_string()), + secret_file: None, + prefix: Some("Bearer ".to_string()), + }], + }, + } +} + +fn policy_ctx( + app_state: Arc, + mode: NetworkMode, + target_host: &str, + target_port: u16, +) -> MitmPolicyContext { + MitmPolicyContext { + target_host: target_host.to_string(), + target_port, + scheme: Scheme::HTTPS, + mode, + app_state, + } +} + +#[tokio::test] +async fn mitm_policy_blocks_disallowed_method_and_records_telemetry() { + let app_state = Arc::new(network_proxy_state_for_policy({ + let mut network = NetworkProxyConfig::default(); + network.set_allowed_domains(vec!["example.com".to_string()]); + network + })); + let ctx = policy_ctx( + app_state.clone(), + NetworkMode::Limited, + "example.com", + /*target_port*/ 443, + ); + let req = Request::builder() + .method(Method::POST) + .uri("/v1/responses?api_key=secret") + .header(HOST, "example.com") + .body(Body::empty()) + .unwrap(); + + let response = mitm_blocking_response(&req, &ctx) + .await + .unwrap() + .expect("POST should be blocked in limited mode"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response.headers().get("x-proxy-error").unwrap(), + "blocked-by-method-policy" + ); + + let blocked = app_state.drain_blocked().await.unwrap(); + assert_eq!(blocked.len(), 1); + assert_eq!(blocked[0].reason, REASON_METHOD_NOT_ALLOWED); + assert_eq!(blocked[0].method.as_deref(), Some("POST")); + assert_eq!(blocked[0].host, "example.com"); + assert_eq!(blocked[0].port, Some(443)); +} + +#[tokio::test] +async fn mitm_policy_rejects_host_mismatch() { + let app_state = Arc::new(network_proxy_state_for_policy({ + let mut network = NetworkProxyConfig::default(); + network.set_allowed_domains(vec!["example.com".to_string()]); + network + })); + let ctx = policy_ctx( + app_state.clone(), + NetworkMode::Full, + "example.com", + /*target_port*/ 443, + ); + let req = Request::builder() + .method(Method::GET) + .uri("/") + .header(HOST, "evil.example") + .body(Body::empty()) + .unwrap(); + + let response = mitm_blocking_response(&req, &ctx) + .await + .unwrap() + .expect("mismatched host should be rejected"); + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!(app_state.blocked_snapshot().await.unwrap().len(), 0); +} + +#[tokio::test] +async fn mitm_policy_rechecks_local_private_target_after_connect() { + let app_state = Arc::new(network_proxy_state_for_policy({ + let mut network = NetworkProxyConfig::default(); + network.set_allowed_domains(vec!["example.com".to_string()]); + network.allow_local_binding = false; + network + })); + let ctx = policy_ctx( + app_state.clone(), + NetworkMode::Full, + "10.0.0.1", + /*target_port*/ 443, + ); + let req = Request::builder() + .method(Method::GET) + .uri("/health?token=secret") + .header(HOST, "10.0.0.1") + .body(Body::empty()) + .unwrap(); + + let response = mitm_blocking_response(&req, &ctx) + .await + .unwrap() + .expect("local/private target should be blocked on inner request"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + + let blocked = app_state.drain_blocked().await.unwrap(); + assert_eq!(blocked.len(), 1); + assert_eq!(blocked[0].reason, REASON_NOT_ALLOWED_LOCAL); + assert_eq!(blocked[0].host, "10.0.0.1"); + assert_eq!(blocked[0].port, Some(443)); +} + +#[tokio::test] +async fn mitm_policy_allows_matching_hooked_write_in_full_mode() { + let secret_file = NamedTempFile::new().unwrap(); + std::fs::write(secret_file.path(), "ghp-secret\n").unwrap(); + let mut hook = github_write_hook(); + hook.actions.inject_request_headers[0].secret_env_var = None; + hook.actions.inject_request_headers[0].secret_file = + Some(secret_file.path().display().to_string()); + let mut network = NetworkProxyConfig { + mitm: true, + mitm_hooks: vec![hook], + mode: NetworkMode::Full, + ..NetworkProxyConfig::default() + }; + network.set_allowed_domains(vec!["api.github.com".to_string()]); + let app_state = Arc::new(network_proxy_state_for_policy(network)); + let ctx = policy_ctx( + app_state.clone(), + NetworkMode::Full, + "api.github.com", + /*target_port*/ 443, + ); + let req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/codex/issues") + .header(HOST, "api.github.com") + .body(Body::empty()) + .unwrap(); + + let response = mitm_blocking_response(&req, &ctx).await.unwrap(); + + assert!( + response.is_none(), + "matching hook should bypass method clamp" + ); + assert_eq!(app_state.blocked_snapshot().await.unwrap().len(), 0); +} + +#[tokio::test] +async fn mitm_policy_blocks_encoded_path_traversal_for_repository_allowlist() { + let mut hook = github_write_hook(); + hook.host = "github.com".to_string(); + hook.matcher.methods = vec!["GET".to_string()]; + hook.matcher.path_prefixes = vec!["pattern:/openai/openai/**".to_string()]; + hook.actions.inject_request_headers.clear(); + let mut network = NetworkProxyConfig { + mitm: true, + mitm_hooks: vec![hook], + mode: NetworkMode::Full, + ..NetworkProxyConfig::default() + }; + network.set_allowed_domains(vec!["github.com".to_string()]); + let app_state = Arc::new(network_proxy_state_for_policy(network)); + let ctx = policy_ctx( + app_state.clone(), + NetworkMode::Full, + "github.com", + /*target_port*/ 443, + ); + let paths = [ + "/openai/openai/issues", + "/openai/codex", + "/openai/openai/%2e%2e/codex", + "/openai/openai/%2e%2e/%2e%2e/microsoft/vscode", + ]; + let mut actual = Vec::with_capacity(paths.len()); + for path in paths { + let req = Request::builder() + .method(Method::GET) + .uri(path) + .header(HOST, "github.com") + .body(Body::empty()) + .unwrap(); + let response = mitm_blocking_response(&req, &ctx).await.unwrap(); + actual.push(response.map(|response| { + ( + response.status(), + response.headers().get("x-proxy-error").cloned(), + ) + })); + } + + assert_eq!( + actual, + vec![ + None, + Some(( + StatusCode::FORBIDDEN, + Some(HeaderValue::from_static("blocked-by-mitm-hook")), + )), + Some(( + StatusCode::FORBIDDEN, + Some(HeaderValue::from_static("blocked-by-mitm-hook")), + )), + Some(( + StatusCode::FORBIDDEN, + Some(HeaderValue::from_static("blocked-by-mitm-hook")), + )), + ] + ); + let blocked = app_state.drain_blocked().await.unwrap(); + assert_eq!(blocked.len(), 3); + assert!( + blocked + .iter() + .all(|request| request.reason == REASON_MITM_HOOK_DENIED) + ); +} + +#[tokio::test] +async fn mitm_policy_blocks_matching_hooked_write_in_limited_mode() { + let mut hook = github_write_hook(); + hook.actions.inject_request_headers.clear(); + let mut network = NetworkProxyConfig { + mitm: true, + mitm_hooks: vec![hook], + mode: NetworkMode::Limited, + ..NetworkProxyConfig::default() + }; + network.set_allowed_domains(vec!["api.github.com".to_string()]); + let app_state = Arc::new(network_proxy_state_for_policy(network)); + let ctx = policy_ctx( + app_state.clone(), + NetworkMode::Limited, + "api.github.com", + /*target_port*/ 443, + ); + let req = Request::builder() + .method(Method::POST) + .uri("/repos/openai/codex/issues") + .header(HOST, "api.github.com") + .body(Body::empty()) + .unwrap(); + + let response = mitm_blocking_response(&req, &ctx) + .await + .unwrap() + .expect("matching POST hook should still be blocked in limited mode"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response.headers().get("x-proxy-error").unwrap(), + "blocked-by-method-policy" + ); + + let blocked = app_state.drain_blocked().await.unwrap(); + assert_eq!(blocked.len(), 1); + assert_eq!(blocked[0].reason, REASON_METHOD_NOT_ALLOWED); + assert_eq!(blocked[0].method.as_deref(), Some("POST")); + assert_eq!(blocked[0].host, "api.github.com"); + assert_eq!(blocked[0].port, Some(443)); +} + +#[tokio::test] +async fn mitm_policy_blocks_hook_miss_for_hooked_host_and_records_telemetry_in_full_mode() { + let secret_file = NamedTempFile::new().unwrap(); + std::fs::write(secret_file.path(), "ghp-secret\n").unwrap(); + let mut hook = github_write_hook(); + hook.actions.inject_request_headers[0].secret_env_var = None; + hook.actions.inject_request_headers[0].secret_file = + Some(secret_file.path().display().to_string()); + let mut network = NetworkProxyConfig { + mitm: true, + mitm_hooks: vec![hook], + mode: NetworkMode::Full, + ..NetworkProxyConfig::default() + }; + network.set_allowed_domains(vec!["api.github.com".to_string()]); + let app_state = Arc::new(network_proxy_state_for_policy(network)); + let ctx = policy_ctx( + app_state.clone(), + NetworkMode::Full, + "api.github.com", + /*target_port*/ 443, + ); + let req = Request::builder() + .method(Method::GET) + .uri("/repos/openai/codex/issues?token=secret") + .header(HOST, "api.github.com") + .header("authorization", "Bearer user-supplied") + .body(Body::empty()) + .unwrap(); + + let response = mitm_blocking_response(&req, &ctx) + .await + .unwrap() + .expect("hook miss should be blocked"); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!( + response.headers().get("x-proxy-error").unwrap(), + "blocked-by-mitm-hook" + ); + + let blocked = app_state.drain_blocked().await.unwrap(); + assert_eq!(blocked.len(), 1); + assert_eq!(blocked[0].reason, REASON_MITM_HOOK_DENIED); + assert_eq!(blocked[0].method.as_deref(), Some("GET")); + assert_eq!(blocked[0].host, "api.github.com"); + assert_eq!(blocked[0].port, Some(443)); +} + +#[test] +fn apply_mitm_hook_actions_replaces_authorization_header() { + let mut headers = HeaderMap::new(); + headers.append( + HeaderName::from_static("authorization"), + HeaderValue::from_static("Bearer user-supplied"), + ); + headers.append( + HeaderName::from_static("x-request-id"), + HeaderValue::from_static("req_123"), + ); + + let actions = crate::mitm_hook::MitmHookActions { + strip_request_headers: vec![HeaderName::from_static("authorization")], + inject_request_headers: vec![crate::mitm_hook::ResolvedInjectedHeader { + name: HeaderName::from_static("authorization"), + value: HeaderValue::from_static("Bearer secret-token"), + source: crate::mitm_hook::SecretSource::File( + AbsolutePathBuf::try_from("/tmp/github-token").unwrap(), + ), + }], + }; + + apply_mitm_hook_actions(&mut headers, Some(&actions)); + + assert_eq!( + headers.get("authorization"), + Some(&HeaderValue::from_static("Bearer secret-token")) + ); + assert_eq!( + headers.get("x-request-id"), + Some(&HeaderValue::from_static("req_123")) + ); +} diff --git a/codex-rs/network-proxy/src/network_policy.rs b/codex-rs/network-proxy/src/network_policy.rs new file mode 100644 index 0000000000000000000000000000000000000000..a47e1249b95c496562b64959f0834a8f0dce4045 --- /dev/null +++ b/codex-rs/network-proxy/src/network_policy.rs @@ -0,0 +1,1114 @@ +use crate::reasons::REASON_POLICY_DENIED; +use crate::request_cancellation::NetworkRequestCancellation; +use crate::request_disconnect::NetworkRequestDisconnect; +use crate::runtime::HostBlockDecision; +use crate::runtime::HostBlockReason; +use crate::state::NetworkProxyState; +use anyhow::Result; +use chrono::SecondsFormat; +use chrono::Utc; +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; + +const AUDIT_TARGET: &str = "codex_otel.network_proxy"; +const POLICY_DECISION_EVENT_NAME: &str = "codex.network_proxy.policy_decision"; +const POLICY_SCOPE_DOMAIN: &str = "domain"; +const POLICY_SCOPE_NON_DOMAIN: &str = "non_domain"; +const POLICY_DECISION_ALLOW: &str = "allow"; +const POLICY_DECISION_DENY: &str = "deny"; +const POLICY_REASON_ALLOW: &str = "allow"; +const DEFAULT_METHOD: &str = "none"; +const DEFAULT_CLIENT_ADDRESS: &str = "unknown"; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum NetworkProtocol { + Http, + HttpsConnect, + Socks5Tcp, + Socks5Udp, +} + +/// A completed network-policy audit decision without tenant or session identity. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct NetworkPolicyAuditEvent { + pub timestamp: String, + pub scope: String, + pub decision: String, + pub source: String, + pub reason: String, + pub protocol: NetworkProtocol, + pub host: String, + pub port: u16, + pub method: Option, + pub client: Option, + pub policy_override: bool, +} + +/// Observes final network-policy decisions without delaying or altering enforcement. +/// +/// Implementations must return immediately and treat notification delivery as best effort. +pub type NetworkPolicyAuditObserver = Arc; + +impl NetworkProtocol { + pub const fn as_policy_protocol(self) -> &'static str { + match self { + Self::Http => "http", + Self::HttpsConnect => "https_connect", + Self::Socks5Tcp => "socks5_tcp", + Self::Socks5Udp => "socks5_udp", + } + } +} + +#[derive(Clone, Copy, Debug, serde::Deserialize, serde::Serialize, PartialEq, Eq)] +#[serde(rename_all = "lowercase")] +pub enum NetworkPolicyDecision { + Deny, + Ask, +} + +impl NetworkPolicyDecision { + pub const fn as_str(self) -> &'static str { + match self { + Self::Deny => "deny", + Self::Ask => "ask", + } + } +} + +#[derive(Clone, Copy, Debug, serde::Deserialize, serde::Serialize, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum NetworkDecisionSource { + BaselinePolicy, + ModeGuard, + ProxyState, + Decider, +} + +impl NetworkDecisionSource { + pub const fn as_str(self) -> &'static str { + match self { + Self::BaselinePolicy => "baseline_policy", + Self::ModeGuard => "mode_guard", + Self::ProxyState => "proxy_state", + Self::Decider => "decider", + } + } +} + +#[derive(Clone, Debug)] +pub struct NetworkPolicyRequest { + pub protocol: NetworkProtocol, + pub host: String, + pub port: u16, + pub environment_id: Option, + pub client_addr: Option, + pub method: Option, + pub command: Option, + pub exec_policy_hint: Option, + pub execution_id: Option, + /// Present only when the local HTTP transport can identify an abandoned request. + pub disconnect: Option, + /// Controller-owned cause, published before an abandoned decision future is dropped. + pub cancellation: Option, +} + +pub struct NetworkPolicyRequestArgs { + pub protocol: NetworkProtocol, + pub host: String, + pub port: u16, + pub environment_id: Option, + pub client_addr: Option, + pub method: Option, + pub command: Option, + pub exec_policy_hint: Option, +} + +impl NetworkPolicyRequest { + pub fn new(args: NetworkPolicyRequestArgs) -> Self { + let NetworkPolicyRequestArgs { + protocol, + host, + port, + environment_id, + client_addr, + method, + command, + exec_policy_hint, + } = args; + Self { + protocol, + host, + port, + environment_id, + client_addr, + method, + command, + exec_policy_hint, + execution_id: None, + disconnect: None, + cancellation: None, + } + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum NetworkDecision { + Allow, + Deny { + reason: String, + source: NetworkDecisionSource, + decision: NetworkPolicyDecision, + }, +} + +impl NetworkDecision { + pub fn deny(reason: impl Into) -> Self { + Self::deny_with_source(reason, NetworkDecisionSource::Decider) + } + + pub fn ask(reason: impl Into) -> Self { + Self::ask_with_source(reason, NetworkDecisionSource::Decider) + } + + pub fn deny_with_source(reason: impl Into, source: NetworkDecisionSource) -> Self { + let reason = reason.into(); + let reason = if reason.is_empty() { + REASON_POLICY_DENIED.to_string() + } else { + reason + }; + Self::Deny { + reason, + source, + decision: NetworkPolicyDecision::Deny, + } + } + + pub fn ask_with_source(reason: impl Into, source: NetworkDecisionSource) -> Self { + let reason = reason.into(); + let reason = if reason.is_empty() { + REASON_POLICY_DENIED.to_string() + } else { + reason + }; + Self::Deny { + reason, + source, + decision: NetworkPolicyDecision::Ask, + } + } +} + +pub(crate) struct BlockDecisionAuditEventArgs<'a> { + pub source: NetworkDecisionSource, + pub reason: &'a str, + pub protocol: NetworkProtocol, + pub server_address: &'a str, + pub server_port: u16, + pub method: Option<&'a str>, + pub client_addr: Option<&'a str>, +} + +pub(crate) fn emit_block_decision_audit_event( + state: &NetworkProxyState, + args: BlockDecisionAuditEventArgs<'_>, +) { + emit_non_domain_policy_decision_audit_event(state, args, POLICY_DECISION_DENY); +} + +pub(crate) fn emit_allow_decision_audit_event( + state: &NetworkProxyState, + args: BlockDecisionAuditEventArgs<'_>, +) { + emit_non_domain_policy_decision_audit_event(state, args, POLICY_DECISION_ALLOW); +} + +fn emit_non_domain_policy_decision_audit_event( + state: &NetworkProxyState, + args: BlockDecisionAuditEventArgs<'_>, + decision: &'static str, +) { + let execution_id = state.execution_id(); + emit_policy_audit_event( + state, + PolicyAuditEventArgs { + scope: POLICY_SCOPE_NON_DOMAIN, + decision, + source: args.source.as_str(), + reason: args.reason, + protocol: args.protocol, + server_address: args.server_address, + server_port: args.server_port, + method: args.method, + client_addr: args.client_addr, + execution_id: execution_id.as_deref(), + policy_override: false, + }, + ); +} + +struct PolicyAuditEventArgs<'a> { + scope: &'static str, + decision: &'a str, + source: &'a str, + reason: &'a str, + protocol: NetworkProtocol, + server_address: &'a str, + server_port: u16, + method: Option<&'a str>, + client_addr: Option<&'a str>, + execution_id: Option<&'a str>, + policy_override: bool, +} + +fn emit_policy_audit_event(state: &NetworkProxyState, args: PolicyAuditEventArgs<'_>) { + let audit_metadata = state.audit_metadata(); + let process_log_metadata = &state.process_log_metadata; + let conversation_id = process_log_metadata + .thread_id + .as_deref() + .or(audit_metadata.conversation_id.as_deref()); + let timestamp = audit_timestamp(); + let launch_trace_id = state + .launch_span_context + .as_ref() + .map(|context| context.trace_id().to_string()); + let launch_span_id = state + .launch_span_context + .as_ref() + .map(|context| context.span_id().to_string()); + let executor_identity = process_log_metadata.executor_identity.as_ref(); + tracing::event!( + target: AUDIT_TARGET, + tracing::Level::INFO, + event.name = POLICY_DECISION_EVENT_NAME, + event.timestamp = %timestamp, + launch.trace_id = launch_trace_id.as_deref(), + launch.span_id = launch_span_id.as_deref(), + conversation.id = conversation_id, + tool.call_id = process_log_metadata.tool_call_id.as_deref(), + executor.environment_id = executor_identity.map(|identity| identity.environment_id.as_str()), + executor.registration_id = executor_identity.map(|identity| identity.registration_id.as_str()), + app.version = audit_metadata.app_version.as_deref(), + auth_mode = audit_metadata.auth_mode.as_deref(), + originator = audit_metadata.originator.as_deref(), + user.account_id = audit_metadata.user_account_id.as_deref(), + user.email = audit_metadata.user_email.as_deref(), + terminal.type = audit_metadata.terminal_type.as_deref(), + model = audit_metadata.model.as_deref(), + slug = audit_metadata.slug.as_deref(), + network.policy.scope = args.scope, + network.policy.decision = args.decision, + network.policy.source = args.source, + network.policy.reason = args.reason, + network.transport.protocol = args.protocol.as_policy_protocol(), + server.address = args.server_address, + server.port = args.server_port, + http.request.method = args.method.unwrap_or(DEFAULT_METHOD), + client.address = args.client_addr.unwrap_or(DEFAULT_CLIENT_ADDRESS), + execution.id = args.execution_id, + network.policy.override = args.policy_override, + ); + if let Some(observer) = &state.policy_audit_observer { + observer(NetworkPolicyAuditEvent { + timestamp, + scope: args.scope.to_string(), + decision: args.decision.to_string(), + source: args.source.to_string(), + reason: args.reason.to_string(), + protocol: args.protocol, + host: args.server_address.to_string(), + port: args.server_port, + method: args.method.map(str::to_string), + client: args.client_addr.map(str::to_string), + policy_override: args.policy_override, + }); + } +} + +fn audit_timestamp() -> String { + Utc::now().to_rfc3339_opts(SecondsFormat::Millis, true) +} + +/// Decide whether a network request should be allowed. +/// +/// If `command` or `exec_policy_hint` is provided, callers can map exec-policy +/// approvals to network access (e.g., allow all requests for commands matching +/// approved prefixes like `curl *`). +pub trait NetworkPolicyDecider: Send + Sync + 'static { + fn decide(&self, req: NetworkPolicyRequest) -> NetworkPolicyDeciderFuture<'_>; +} + +pub type NetworkPolicyDeciderFuture<'a> = + Pin + Send + 'a>>; + +impl NetworkPolicyDecider for Arc { + fn decide(&self, req: NetworkPolicyRequest) -> NetworkPolicyDeciderFuture<'_> { + Box::pin(async move { (**self).decide(req).await }) + } +} + +impl NetworkPolicyDecider for F +where + F: Fn(NetworkPolicyRequest) -> Fut + Send + Sync + 'static, + Fut: Future + Send + 'static, +{ + fn decide(&self, req: NetworkPolicyRequest) -> NetworkPolicyDeciderFuture<'_> { + Box::pin((self)(req)) + } +} + +pub(crate) async fn evaluate_host_policy( + state: &NetworkProxyState, + decider: Option<&Arc>, + request: &NetworkPolicyRequest, +) -> Result { + let execution_id = state.execution_id(); + let host_decision = state.host_blocked(&request.host, request.port).await?; + let (decision, policy_override) = match host_decision { + HostBlockDecision::Allowed => (NetworkDecision::Allow, false), + HostBlockDecision::Blocked(HostBlockReason::NotAllowed) => { + if let Some(decider) = decider { + let mut request = request.clone(); + if request.environment_id.is_none() + && let Some(environment_id) = state.environment_id() + { + request.environment_id = Some(environment_id.to_string()); + } + request.execution_id = execution_id.clone(); + let decider_decision = map_decider_decision(decider.decide(request).await); + let policy_override = matches!(decider_decision, NetworkDecision::Allow); + (decider_decision, policy_override) + } else { + ( + NetworkDecision::deny_with_source( + HostBlockReason::NotAllowed.as_str(), + NetworkDecisionSource::BaselinePolicy, + ), + false, + ) + } + } + HostBlockDecision::Blocked(reason) => ( + NetworkDecision::deny_with_source( + reason.as_str(), + NetworkDecisionSource::BaselinePolicy, + ), + false, + ), + }; + + let (policy_decision, source, reason) = match &decision { + NetworkDecision::Allow => ( + POLICY_DECISION_ALLOW, + if policy_override { + NetworkDecisionSource::Decider + } else { + NetworkDecisionSource::BaselinePolicy + }, + if policy_override { + HostBlockReason::NotAllowed.as_str() + } else { + POLICY_REASON_ALLOW + }, + ), + NetworkDecision::Deny { + reason, + source, + decision, + } => (decision.as_str(), *source, reason.as_str()), + }; + + emit_policy_audit_event( + state, + PolicyAuditEventArgs { + scope: POLICY_SCOPE_DOMAIN, + decision: policy_decision, + source: source.as_str(), + reason, + protocol: request.protocol, + server_address: request.host.as_str(), + server_port: request.port, + method: request.method.as_deref(), + client_addr: request.client_addr.as_deref(), + execution_id: execution_id.as_deref(), + policy_override, + }, + ); + + Ok(decision) +} + +fn map_decider_decision(decision: NetworkDecision) -> NetworkDecision { + match decision { + NetworkDecision::Allow => NetworkDecision::Allow, + NetworkDecision::Deny { + reason, decision, .. + } => NetworkDecision::Deny { + reason, + source: NetworkDecisionSource::Decider, + decision, + }, + } +} + +#[cfg(test)] +pub(crate) mod test_support { + pub(crate) const POLICY_DECISION_EVENT_NAME: &str = super::POLICY_DECISION_EVENT_NAME; + + use std::collections::BTreeMap; + use std::fmt; + use std::future::Future; + use std::sync::Arc; + use std::sync::Mutex; + use std::sync::atomic::AtomicU64; + use std::sync::atomic::Ordering; + use tracing::Event; + use tracing::Id; + use tracing::Metadata; + use tracing::Subscriber; + use tracing::field::Field; + use tracing::field::Visit; + use tracing::instrument::WithSubscriber; + use tracing::span::Attributes; + use tracing::span::Record; + use tracing::subscriber::Interest; + + #[derive(Clone, Debug, PartialEq, Eq)] + pub(crate) struct CapturedEvent { + pub target: String, + pub fields: BTreeMap, + } + + impl CapturedEvent { + pub fn field(&self, name: &str) -> Option<&str> { + self.fields.get(name).map(String::as_str) + } + } + + #[derive(Clone, Default)] + struct EventCollector { + events: Arc>>, + next_span_id: Arc, + } + + impl EventCollector { + fn events(&self) -> Vec { + self.events + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + } + } + + impl Subscriber for EventCollector { + fn enabled(&self, _metadata: &Metadata<'_>) -> bool { + true + } + + fn register_callsite(&self, _metadata: &'static Metadata<'static>) -> Interest { + Interest::always() + } + + fn max_level_hint(&self) -> Option { + Some(tracing::level_filters::LevelFilter::TRACE) + } + + fn new_span(&self, _span: &Attributes<'_>) -> Id { + Id::from_u64(self.next_span_id.fetch_add(1, Ordering::Relaxed) + 1) + } + + fn record(&self, _span: &Id, _values: &Record<'_>) {} + + fn record_follows_from(&self, _span: &Id, _follows: &Id) {} + + fn event(&self, event: &Event<'_>) { + let mut visitor = FieldVisitor::default(); + event.record(&mut visitor); + self.events + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .push(CapturedEvent { + target: event.metadata().target().to_string(), + fields: visitor.fields, + }); + } + + fn enter(&self, _span: &Id) {} + + fn exit(&self, _span: &Id) {} + } + + #[derive(Default)] + struct FieldVisitor { + fields: BTreeMap, + } + + impl FieldVisitor { + fn insert(&mut self, field: &Field, value: impl Into) { + self.fields.insert(field.name().to_string(), value.into()); + } + } + + impl Visit for FieldVisitor { + fn record_str(&mut self, field: &Field, value: &str) { + self.insert(field, value); + } + + fn record_bool(&mut self, field: &Field, value: bool) { + self.insert(field, value.to_string()); + } + + fn record_i64(&mut self, field: &Field, value: i64) { + self.insert(field, value.to_string()); + } + + fn record_u64(&mut self, field: &Field, value: u64) { + self.insert(field, value.to_string()); + } + + fn record_i128(&mut self, field: &Field, value: i128) { + self.insert(field, value.to_string()); + } + + fn record_u128(&mut self, field: &Field, value: u128) { + self.insert(field, value.to_string()); + } + + fn record_f64(&mut self, field: &Field, value: f64) { + self.insert(field, value.to_string()); + } + + fn record_error(&mut self, field: &Field, value: &(dyn std::error::Error + 'static)) { + self.insert(field, value.to_string()); + } + + fn record_debug(&mut self, field: &Field, value: &dyn fmt::Debug) { + self.insert(field, format!("{value:?}")); + } + } + + pub(crate) async fn capture_events(f: F) -> (T, Vec) + where + F: FnOnce() -> Fut, + Fut: Future, + { + let collector = EventCollector::default(); + // Keep tracing out of its single-subscriber fast path: concurrent tests + // without a subscriber can otherwise cache this callsite as disabled. + let _interest_dispatch = tracing::Dispatch::new(collector.clone()); + let output = async { + tracing::callsite::rebuild_interest_cache(); + f().await + } + .with_subscriber(collector.clone()) + .await; + let events = collector.events(); + (output, events) + } + + pub(crate) fn find_event_by_name<'a>( + events: &'a [CapturedEvent], + event_name: &str, + ) -> Option<&'a CapturedEvent> { + events + .iter() + .find(|event| event.field("event.name") == Some(event_name)) + } +} + +#[cfg(test)] +mod tests { + use super::test_support::capture_events; + use super::test_support::find_event_by_name; + use super::*; + use crate::ExecutorLogIdentity; + use crate::NetworkProxyProcessLogMetadata; + use crate::config::NetworkMode; + use crate::config::NetworkProxyConfig; + use crate::reasons::REASON_DENIED; + use crate::reasons::REASON_METHOD_NOT_ALLOWED; + use crate::reasons::REASON_NOT_ALLOWED; + use crate::reasons::REASON_NOT_ALLOWED_LOCAL; + use crate::runtime::ConfigReloader; + use crate::runtime::ConfigReloaderFuture; + use crate::runtime::ConfigState; + use crate::runtime::NetworkProxyAuditMetadata; + use crate::state::NetworkProxyConstraints; + use crate::state::build_config_state; + use crate::state::network_proxy_state_for_policy; + use pretty_assertions::assert_eq; + use std::sync::Arc; + use std::sync::atomic::AtomicUsize; + use std::sync::atomic::Ordering; + + const LEGACY_DOMAIN_POLICY_DECISION_EVENT_NAME: &str = + "codex.network_proxy.domain_policy_decision"; + const LEGACY_BLOCK_DECISION_EVENT_NAME: &str = "codex.network_proxy.block_decision"; + + #[derive(Clone)] + struct StaticReloader { + state: ConfigState, + } + + impl ConfigReloader for StaticReloader { + fn maybe_reload(&self) -> ConfigReloaderFuture<'_, Option> { + Box::pin(async { Ok(None) }) + } + + fn reload_now(&self) -> ConfigReloaderFuture<'_, ConfigState> { + Box::pin(async { Ok(self.state.clone()) }) + } + + fn source_label(&self) -> String { + "static test reloader".to_string() + } + } + + fn state_with_metadata(metadata: NetworkProxyAuditMetadata) -> NetworkProxyState { + let network = NetworkProxyConfig { + enabled: true, + mode: NetworkMode::Full, + ..NetworkProxyConfig::default() + }; + let config = network; + let state = build_config_state(config, NetworkProxyConstraints::default()).unwrap(); + let reloader = Arc::new(StaticReloader { + state: state.clone(), + }); + NetworkProxyState::with_reloader_and_audit_metadata(state, reloader, metadata) + } + + fn is_rfc3339_utc_millis(timestamp: &str) -> bool { + let bytes = timestamp.as_bytes(); + if bytes.len() != 24 { + return false; + } + bytes[4] == b'-' + && bytes[7] == b'-' + && bytes[10] == b'T' + && bytes[13] == b':' + && bytes[16] == b':' + && bytes[19] == b'.' + && bytes[23] == b'Z' + && bytes.iter().enumerate().all(|(idx, value)| match idx { + 4 | 7 | 10 | 13 | 16 | 19 | 23 => true, + _ => value.is_ascii_digit(), + }) + } + + #[tokio::test(flavor = "current_thread")] + async fn policy_audit_observer_receives_domain_and_non_domain_decisions() { + let mut state = network_proxy_state_for_policy(NetworkProxyConfig::default()); + let (captured_tx, captured_rx) = std::sync::mpsc::channel(); + state.set_policy_audit_observer(Arc::new(move |event| { + captured_tx + .send(event) + .expect("observer should capture the policy decision"); + })); + let decider: Arc = + Arc::new(|_request| async { NetworkDecision::Allow }); + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Http, + host: "example.com".to_string(), + port: 80, + environment_id: None, + client_addr: None, + method: None, + command: None, + exec_policy_hint: None, + }); + evaluate_host_policy(&state, Some(&decider), &request) + .await + .expect("evaluate domain policy"); + emit_block_decision_audit_event( + &state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ModeGuard, + reason: REASON_METHOD_NOT_ALLOWED, + protocol: NetworkProtocol::Http, + server_address: "unix-socket", + server_port: 0, + method: Some("POST"), + client_addr: None, + }, + ); + + let events: Vec<_> = captured_rx.try_iter().collect(); + assert_eq!( + events + .iter() + .map(|event| (event.scope.as_str(), event.decision.as_str())) + .collect::>(), + vec![("domain", "allow"), ("non_domain", "deny")] + ); + } + + #[tokio::test(flavor = "current_thread")] + async fn evaluate_host_policy_emits_domain_event_for_decider_allow_override() { + let state = network_proxy_state_for_policy(NetworkProxyConfig::default()); + let calls = Arc::new(AtomicUsize::new(0)); + let decider: Arc = Arc::new({ + let calls = calls.clone(); + move |_req| { + calls.fetch_add(1, Ordering::SeqCst); + // The default policy denies all; the decider is consulted for not_allowed + // requests and can override that decision. + async { NetworkDecision::Allow } + } + }); + + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Http, + host: "example.com".to_string(), + port: 80, + environment_id: None, + client_addr: None, + method: None, + command: None, + exec_policy_hint: None, + }); + + let (decision, events) = capture_events(|| async { + evaluate_host_policy(&state, Some(&decider), &request) + .await + .unwrap() + }) + .await; + assert_eq!(decision, NetworkDecision::Allow); + assert_eq!(calls.load(Ordering::SeqCst), 1); + + let event = find_event_by_name(&events, POLICY_DECISION_EVENT_NAME) + .expect("expected policy decision audit event"); + assert_eq!(event.target, AUDIT_TARGET); + assert!(event.target.starts_with("codex_otel.")); + assert_eq!( + event.field("network.policy.scope"), + Some(POLICY_SCOPE_DOMAIN) + ); + assert_eq!(event.field("network.policy.decision"), Some("allow")); + assert_eq!(event.field("network.policy.source"), Some("decider")); + assert_eq!( + event.field("network.policy.reason"), + Some(REASON_NOT_ALLOWED) + ); + assert_eq!(event.field("network.transport.protocol"), Some("http")); + assert_eq!(event.field("server.address"), Some("example.com")); + assert_eq!(event.field("server.port"), Some("80")); + assert_eq!(event.field("http.request.method"), Some(DEFAULT_METHOD)); + assert_eq!(event.field("client.address"), Some(DEFAULT_CLIENT_ADDRESS)); + assert_eq!(event.field("network.policy.override"), Some("true")); + let timestamp = event + .field("event.timestamp") + .expect("event timestamp should be present"); + assert!(is_rfc3339_utc_millis(timestamp)); + assert_eq!( + find_event_by_name(&events, LEGACY_DOMAIN_POLICY_DECISION_EVENT_NAME), + None + ); + assert_eq!( + find_event_by_name(&events, LEGACY_BLOCK_DECISION_EVENT_NAME), + None + ); + } + + #[tokio::test(flavor = "current_thread")] + async fn evaluate_host_policy_emits_execution_id_for_baseline_allow() { + let state = network_proxy_state_for_policy({ + let mut network = NetworkProxyConfig::default(); + network.set_allowed_domains(vec!["example.com".to_string()]); + network + }); + state.register_execution("token-baseline-allow", "local", "execution-baseline-allow"); + let state = state + .for_execution_token("token-baseline-allow") + .expect("expected registered execution"); + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Http, + host: "example.com".to_string(), + port: 80, + environment_id: None, + client_addr: None, + method: None, + command: None, + exec_policy_hint: None, + }); + + let (decision, events) = capture_events(|| async { + evaluate_host_policy(&state, /*decider*/ None, &request) + .await + .unwrap() + }) + .await; + assert_eq!(decision, NetworkDecision::Allow); + + let event = find_event_by_name(&events, POLICY_DECISION_EVENT_NAME) + .expect("expected policy decision audit event"); + assert_eq!(event.field("network.policy.decision"), Some("allow")); + assert_eq!( + event.field("execution.id"), + Some("execution-baseline-allow") + ); + assert_ne!(event.field("execution.id"), Some("token-baseline-allow")); + } + + #[tokio::test(flavor = "current_thread")] + async fn evaluate_host_policy_emits_domain_event_for_baseline_deny() { + let state = network_proxy_state_for_policy({ + let mut network = NetworkProxyConfig::default(); + network.set_allowed_domains(vec!["example.com".to_string()]); + network.set_denied_domains(vec!["blocked.com".to_string()]); + network + }); + state.register_execution("token-baseline-deny", "local", "execution-baseline-deny"); + let state = state + .for_execution_token("token-baseline-deny") + .expect("expected registered execution"); + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Http, + host: "blocked.com".to_string(), + port: 80, + environment_id: None, + client_addr: Some("127.0.0.1:1234".to_string()), + method: Some("GET".to_string()), + command: None, + exec_policy_hint: None, + }); + + let (decision, events) = capture_events(|| async { + evaluate_host_policy(&state, /*decider*/ None, &request) + .await + .unwrap() + }) + .await; + assert_eq!( + decision, + NetworkDecision::Deny { + reason: REASON_DENIED.to_string(), + source: NetworkDecisionSource::BaselinePolicy, + decision: NetworkPolicyDecision::Deny, + } + ); + + let event = find_event_by_name(&events, POLICY_DECISION_EVENT_NAME) + .expect("expected policy decision audit event"); + assert_eq!(event.field("network.policy.decision"), Some("deny")); + assert_eq!( + event.field("network.policy.source"), + Some("baseline_policy") + ); + assert_eq!(event.field("network.policy.reason"), Some(REASON_DENIED)); + assert_eq!(event.field("network.policy.override"), Some("false")); + assert_eq!(event.field("http.request.method"), Some("GET")); + assert_eq!(event.field("client.address"), Some("127.0.0.1:1234")); + assert_eq!(event.field("execution.id"), Some("execution-baseline-deny")); + assert_ne!(event.field("execution.id"), Some("token-baseline-deny")); + } + + #[tokio::test(flavor = "current_thread")] + async fn evaluate_host_policy_emits_domain_event_for_decider_ask() { + let state = network_proxy_state_for_policy(NetworkProxyConfig::default()); + let decider: Arc = + Arc::new(|_req| async { NetworkDecision::ask(REASON_NOT_ALLOWED) }); + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Http, + host: "example.com".to_string(), + port: 80, + environment_id: None, + client_addr: None, + method: Some("GET".to_string()), + command: None, + exec_policy_hint: None, + }); + + let (decision, events) = capture_events(|| async { + evaluate_host_policy(&state, Some(&decider), &request) + .await + .unwrap() + }) + .await; + assert_eq!( + decision, + NetworkDecision::Deny { + reason: REASON_NOT_ALLOWED.to_string(), + source: NetworkDecisionSource::Decider, + decision: NetworkPolicyDecision::Ask, + } + ); + + let event = find_event_by_name(&events, POLICY_DECISION_EVENT_NAME) + .expect("expected policy decision audit event"); + assert_eq!(event.field("network.policy.decision"), Some("ask")); + assert_eq!(event.field("network.policy.source"), Some("decider")); + assert_eq!( + event.field("network.policy.reason"), + Some(REASON_NOT_ALLOWED) + ); + assert_eq!(event.field("network.policy.override"), Some("false")); + } + + #[tokio::test(flavor = "current_thread")] + async fn evaluate_host_policy_emits_metadata_fields() { + let metadata = NetworkProxyAuditMetadata { + conversation_id: Some("conversation-1".to_string()), + app_version: Some("1.2.3".to_string()), + user_account_id: Some("acct-1".to_string()), + auth_mode: Some("Chatgpt".to_string()), + originator: Some("codex_cli_rs".to_string()), + user_email: Some("test@example.com".to_string()), + terminal_type: Some("iTerm.app/3.6.5".to_string()), + model: Some("gpt-5.3-codex".to_string()), + slug: Some("gpt-5.3-codex".to_string()), + }; + let mut state = state_with_metadata(metadata); + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Http, + host: "example.com".to_string(), + port: 80, + environment_id: None, + client_addr: None, + method: Some("GET".to_string()), + command: None, + exec_policy_hint: None, + }); + + for (thread_id, expected_conversation_id) in [ + (None, "conversation-1"), + (Some("process-thread-1"), "process-thread-1"), + (None, "conversation-1"), + ] { + state.set_process_log_metadata(NetworkProxyProcessLogMetadata { + thread_id: thread_id.map(str::to_string), + tool_call_id: Some("call-1".to_string()), + executor_identity: Some(ExecutorLogIdentity { + environment_id: "environment-1".to_string(), + registration_id: "registration-1".to_string(), + }), + }); + let (_decision, events) = capture_events(|| async { + evaluate_host_policy(&state, /*decider*/ None, &request) + .await + .unwrap() + }) + .await; + + let event = find_event_by_name(&events, POLICY_DECISION_EVENT_NAME) + .expect("expected policy decision audit event"); + assert_eq!( + event.field("conversation.id"), + Some(expected_conversation_id) + ); + assert_eq!(event.field("tool.call_id"), Some("call-1")); + assert_eq!( + event.field("executor.environment_id"), + Some("environment-1") + ); + assert_eq!( + event.field("executor.registration_id"), + Some("registration-1") + ); + assert_eq!(event.field("app.version"), Some("1.2.3")); + assert_eq!(event.field("auth_mode"), Some("Chatgpt")); + assert_eq!(event.field("originator"), Some("codex_cli_rs")); + assert_eq!(event.field("user.account_id"), Some("acct-1")); + assert_eq!(event.field("user.email"), Some("test@example.com")); + assert_eq!(event.field("terminal.type"), Some("iTerm.app/3.6.5")); + assert_eq!(event.field("model"), Some("gpt-5.3-codex")); + assert_eq!(event.field("slug"), Some("gpt-5.3-codex")); + } + } + + #[tokio::test(flavor = "current_thread")] + async fn emit_block_decision_audit_event_emits_non_domain_event() { + let state = network_proxy_state_for_policy(NetworkProxyConfig::default()); + + let (_, events) = capture_events(|| async { + emit_block_decision_audit_event( + &state, + BlockDecisionAuditEventArgs { + source: NetworkDecisionSource::ModeGuard, + reason: REASON_METHOD_NOT_ALLOWED, + protocol: NetworkProtocol::Http, + server_address: "unix-socket", + server_port: 0, + method: Some("POST"), + client_addr: None, + }, + ); + }) + .await; + + let event = find_event_by_name(&events, POLICY_DECISION_EVENT_NAME) + .expect("expected policy decision audit event"); + assert_eq!(event.target, AUDIT_TARGET); + assert_eq!( + event.field("network.policy.scope"), + Some(POLICY_SCOPE_NON_DOMAIN) + ); + assert_eq!( + event.field("network.policy.decision"), + Some(POLICY_DECISION_DENY) + ); + assert_eq!(event.field("network.policy.source"), Some("mode_guard")); + assert_eq!( + event.field("network.policy.reason"), + Some(REASON_METHOD_NOT_ALLOWED) + ); + assert_eq!(event.field("network.transport.protocol"), Some("http")); + assert_eq!(event.field("server.address"), Some("unix-socket")); + assert_eq!(event.field("server.port"), Some("0")); + assert_eq!(event.field("http.request.method"), Some("POST")); + assert_eq!(event.field("client.address"), Some(DEFAULT_CLIENT_ADDRESS)); + assert_eq!(event.field("network.policy.override"), Some("false")); + assert_eq!( + find_event_by_name(&events, LEGACY_BLOCK_DECISION_EVENT_NAME), + None + ); + } + + #[tokio::test(flavor = "current_thread")] + async fn evaluate_host_policy_still_denies_not_allowed_local_without_decider_override() { + let state = network_proxy_state_for_policy({ + let mut network = NetworkProxyConfig::default(); + network.set_allowed_domains(vec!["example.com".to_string()]); + network.allow_local_binding = false; + network + }); + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Http, + host: "127.0.0.1".to_string(), + port: 80, + environment_id: None, + client_addr: None, + method: Some("GET".to_string()), + command: None, + exec_policy_hint: None, + }); + + let decision = evaluate_host_policy(&state, /*decider*/ None, &request) + .await + .unwrap(); + assert_eq!( + decision, + NetworkDecision::Deny { + reason: REASON_NOT_ALLOWED_LOCAL.to_string(), + source: NetworkDecisionSource::BaselinePolicy, + decision: NetworkPolicyDecision::Deny, + } + ); + } + + #[test] + fn ask_uses_decider_source_and_ask_decision() { + assert_eq!( + NetworkDecision::ask(REASON_NOT_ALLOWED), + NetworkDecision::Deny { + reason: REASON_NOT_ALLOWED.to_string(), + source: NetworkDecisionSource::Decider, + decision: NetworkPolicyDecision::Ask, + } + ); + } +} diff --git a/codex-rs/network-proxy/src/proxy/execution_scope.rs b/codex-rs/network-proxy/src/proxy/execution_scope.rs new file mode 100644 index 0000000000000000000000000000000000000000..e01e115e7863a12c307c69bfbca6b284e600fabd --- /dev/null +++ b/codex-rs/network-proxy/src/proxy/execution_scope.rs @@ -0,0 +1,53 @@ +use super::*; + +pub(super) struct ExecutionScope { + pub(super) environment_id: String, + pub(super) execution_id: String, + pub(super) attribution_token: String, + pub(super) environment_policy: Option, + // Dropping the execution scope closes this channel and cancels remote reviews. + pub(super) lifetime_tx: tokio::sync::watch::Sender<()>, + state: Arc, +} + +impl Drop for ExecutionScope { + fn drop(&mut self) { + self.state.unregister_execution(&self.attribution_token); + } +} + +impl NetworkProxy { + /// Returns a proxy with execution-specific attribution and attachment policy. + pub fn for_execution( + &self, + environment_id: &str, + execution_id: &str, + attribution_token: String, + environment_policy: Option, + fallback_policy_decider: Option>, + ) -> Result { + anyhow::ensure!( + self.execution_scope.is_none(), + "cannot scope an execution-scoped network proxy" + ); + self.state + .register_execution(&attribution_token, environment_id, execution_id); + + let (lifetime_tx, _) = tokio::sync::watch::channel(()); + let mut proxy = self.clone(); + proxy.policy_decider = proxy.policy_decider.or(fallback_policy_decider); + // Strict attachment allowlists cannot be expanded through approval callbacks. + if matches!(&environment_policy, Some(policy) if policy.managed_allowed_domains_only) { + proxy.policy_decider = None; + } + proxy.execution_scope = Some(Arc::new(ExecutionScope { + environment_id: environment_id.to_string(), + execution_id: execution_id.to_string(), + attribution_token, + environment_policy, + lifetime_tx, + state: Arc::clone(&self.state), + })); + Ok(proxy) + } +} diff --git a/codex-rs/network-proxy/src/proxy/managed_routing_tests.rs b/codex-rs/network-proxy/src/proxy/managed_routing_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..55ca34c3905a91c448b2233cd964f13a35ff7ff4 --- /dev/null +++ b/codex-rs/network-proxy/src/proxy/managed_routing_tests.rs @@ -0,0 +1,123 @@ +//! Dedicated managed proxy routing regression coverage. +//! Each proxy must use distinct loopback listeners while preserving HTTP and SOCKS policy. + +use super::*; +use crate::config::NetworkProxyConfig; +use crate::state::network_proxy_state_for_policy; +use pretty_assertions::assert_eq; +use std::net::Ipv4Addr; +use std::time::Duration; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpStream; + +#[tokio::test] +async fn dedicated_listeners_preserve_http_and_socks_policy_without_restricted_tokens() -> Result<()> +{ + let origin = tokio::net::TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await?; + let origin_port = origin.local_addr()?.port(); + let origin_task = tokio::spawn(async move { + while let Ok((mut stream, _)) = origin.accept().await { + tokio::spawn(async move { + let mut request = [0; 1024]; + let _ = stream.read(&mut request).await; + let _ = stream + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n") + .await; + }); + } + }); + let mut handles = Vec::new(); + let mut addresses = Vec::new(); + for allowed in [true, false] { + let mut config = NetworkProxyConfig { + enabled: true, + mode: crate::NetworkMode::Full, + ..NetworkProxyConfig::default() + }; + config.set_allowed_domains(vec![if allowed { + "127.0.0.1".to_string() + } else { + "unreachable.invalid".to_string() + }]); + let proxy = NetworkProxy::builder() + .state(Arc::new(network_proxy_state_for_policy(config))) + .managed_proxy_routing(ManagedProxyRouting::DedicatedListeners) + .build() + .await?; + handles.push(proxy.run().await?); + let proxy = proxy.for_execution( + "environment", + "execution", + "token".to_string(), + /*environment_policy*/ None, + /*fallback_policy_decider*/ None, + )?; + let prepared = + proxy.prepare_for_optional_environment(HashMap::new(), /*environment_id*/ None)?; + let http = prepared.env["HTTP_PROXY"] + .trim_start_matches("http://") + .parse::()?; + let socks = prepared.env["ALL_PROXY"] + .trim_start_matches("socks5h://") + .parse::()?; + let mut ports = vec![http.port(), socks.port()]; + ports.sort_unstable(); + assert_eq!( + prepared.sandbox_context, + ManagedNetworkSandboxContext { + loopback_ports: ports, + allow_local_binding: false, + ..ManagedNetworkSandboxContext::default() + } + ); + assert_eq!( + (http.ip(), socks.ip()), + (Ipv4Addr::LOCALHOST.into(), Ipv4Addr::LOCALHOST.into()) + ); + for addr in [http, socks] { + assert!(!addresses.contains(&addr)); + addresses.push(addr); + } + #[cfg(target_os = "windows")] + { + assert_eq!( + proxy.network_proxy_restricting_sid(/*environment_id*/ None), + None + ); + assert!( + !prepared + .env + .contains_key(WINDOWS_SANDBOX_PROXY_PORTS_ENV_KEY) + ); + } + tokio::time::timeout(Duration::from_secs(5), async { + let mut stream = TcpStream::connect(http).await?; + stream.write_all(format!( + "GET http://127.0.0.1:{origin_port}/ HTTP/1.1\r\nHost: 127.0.0.1:{origin_port}\r\nConnection: close\r\n\r\n" + ).as_bytes()).await?; + let mut response = Vec::new(); + stream.read_to_end(&mut response).await?; + let status = if allowed { "200" } else { "403" }; + assert_eq!(String::from_utf8(response)?.split_whitespace().nth(1), Some(status)); + + let mut stream = TcpStream::connect(socks).await?; + stream.write_all(&[5, 1, 0]).await?; + let mut greeting = [0; 2]; + stream.read_exact(&mut greeting).await?; + assert_eq!(greeting, [5, 0]); + let mut connect = vec![5, 1, 0, 1, 127, 0, 0, 1]; + connect.extend_from_slice(&origin_port.to_be_bytes()); + stream.write_all(&connect).await?; + let mut reply = [0; 4]; + stream.read_exact(&mut reply).await?; + assert_eq!(reply[1] == 0, allowed); + Ok::<(), anyhow::Error>(()) + }).await??; + } + for handle in handles { + handle.shutdown().await?; + } + origin_task.abort(); + Ok(()) +} diff --git a/codex-rs/network-proxy/src/remote_config.rs b/codex-rs/network-proxy/src/remote_config.rs new file mode 100644 index 0000000000000000000000000000000000000000..3bb5101a616c9cea58d1a7cb3816f63af0b5baf5 --- /dev/null +++ b/codex-rs/network-proxy/src/remote_config.rs @@ -0,0 +1,116 @@ +use anyhow::Result; +use anyhow::ensure; +use serde::Deserialize; +use serde::Serialize; + +use crate::NetworkDomainPermissions; +use crate::NetworkMode; +use crate::NetworkProxyAuditMetadata; +use crate::NetworkProxyConfig; +use crate::NetworkUnixSocketPermissions; + +/// Executor-local proxy launch inputs transported with one process start. +/// +/// Unlike [`crate::ManagedNetworkSandboxContext`], this describes how the executor should create +/// proxy listeners. The sandbox context is materialized only after those listeners are running. +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(rename_all = "camelCase")] +#[non_exhaustive] +pub struct RemoteNetworkProxyLaunchConfig { + pub proxy: RemoteNetworkProxyConfig, + #[serde(default)] + pub audit_metadata: NetworkProxyAuditMetadata, + #[serde(default)] + pub environment_id: Option, + #[serde(default)] + pub execution_id: Option, + /// Controller-side policy decision budget. The executor adds transport overhead. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub policy_decision_timeout_ms: Option, +} + +impl RemoteNetworkProxyLaunchConfig { + pub fn new(proxy: RemoteNetworkProxyConfig) -> Self { + Self { + proxy, + audit_metadata: NetworkProxyAuditMetadata::default(), + environment_id: None, + execution_id: None, + policy_decision_timeout_ms: None, + } + } + + pub fn with_audit_metadata(mut self, audit_metadata: NetworkProxyAuditMetadata) -> Self { + self.audit_metadata = audit_metadata; + self + } + + pub fn for_execution(mut self, environment_id: String, execution_id: String) -> Self { + self.environment_id = Some(environment_id); + self.execution_id = Some(execution_id); + self + } +} + +/// Effective network proxy settings that are safe to send to a remote executor. +/// +/// Listener addresses are deliberately omitted because the executor chooses its own loopback +/// ports. MITM, credential injection, and hooks are not represented so their configuration cannot +/// cross the exec-server boundary accidentally. +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(rename_all = "camelCase")] +#[non_exhaustive] +pub struct RemoteNetworkProxyConfig { + pub enabled: bool, + pub enable_socks5: bool, + pub enable_socks5_udp: bool, + pub allow_upstream_proxy: bool, + pub dangerously_allow_all_unix_sockets: bool, + pub mode: NetworkMode, + pub domains: Option, + pub unix_sockets: Option, + pub allow_local_binding: bool, +} + +impl RemoteNetworkProxyConfig { + pub fn from_effective_config(config: &NetworkProxyConfig) -> Result { + ensure!( + !config.enabled + || (!config.mitm + && !config.credential_broker + && !config.dangerously_allow_plaintext_credential_injection + && config.mitm_hooks.is_empty()), + "remote exec-server network proxy does not support MITM, credential injection, or MITM hooks" + ); + Ok(Self { + enabled: config.enabled, + enable_socks5: config.enable_socks5, + enable_socks5_udp: config.enable_socks5_udp, + allow_upstream_proxy: config.allow_upstream_proxy, + dangerously_allow_all_unix_sockets: config.dangerously_allow_all_unix_sockets, + mode: config.mode, + domains: config.domains.clone(), + unix_sockets: config.unix_sockets.clone(), + allow_local_binding: config.allow_local_binding, + }) + } + + pub(crate) fn into_network_proxy_config(self) -> NetworkProxyConfig { + NetworkProxyConfig { + enabled: self.enabled, + enable_socks5: self.enable_socks5, + enable_socks5_udp: self.enable_socks5_udp, + allow_upstream_proxy: self.allow_upstream_proxy, + dangerously_allow_all_unix_sockets: self.dangerously_allow_all_unix_sockets, + mode: self.mode, + domains: self.domains, + unix_sockets: self.unix_sockets, + allow_local_binding: self.allow_local_binding, + ..NetworkProxyConfig::default() + } + } +} + +#[cfg(test)] +#[path = "remote_config_tests.rs"] +mod tests; diff --git a/codex-rs/network-proxy/src/remote_config_tests.rs b/codex-rs/network-proxy/src/remote_config_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..5094acfaade237cc27e8f8c0eea1fa710d5f5c1d --- /dev/null +++ b/codex-rs/network-proxy/src/remote_config_tests.rs @@ -0,0 +1,138 @@ +use pretty_assertions::assert_eq; + +use super::RemoteNetworkProxyConfig; +use super::RemoteNetworkProxyLaunchConfig; +use crate::MitmHookConfig; +use crate::NetworkMode; +use crate::NetworkProxyAuditMetadata; +use crate::NetworkProxyConfig; +use crate::NetworkProxyState; + +#[test] +fn round_trip_preserves_supported_effective_settings() { + let mut config = NetworkProxyConfig { + enabled: true, + enable_socks5: false, + enable_socks5_udp: false, + allow_upstream_proxy: false, + dangerously_allow_all_unix_sockets: true, + mode: NetworkMode::Limited, + allow_local_binding: true, + ..NetworkProxyConfig::default() + }; + config.set_allowed_domains(vec!["example.com".into()]); + config.set_denied_domains(vec!["blocked.example.com".into()]); + config.set_allow_unix_sockets(vec!["/var/run/example.sock".into()]); + + let remote = + RemoteNetworkProxyConfig::from_effective_config(&config).expect("supported remote config"); + let round_trip = remote.into_network_proxy_config(); + + assert_eq!(round_trip, config); +} + +#[test] +fn rejects_unsupported_configuration() { + let cases = [ + ( + "MITM", + NetworkProxyConfig { + enabled: true, + mitm: true, + ..NetworkProxyConfig::default() + }, + ), + ( + "credential broker", + NetworkProxyConfig { + enabled: true, + credential_broker: true, + ..NetworkProxyConfig::default() + }, + ), + ( + "plaintext credential injection", + NetworkProxyConfig { + enabled: true, + dangerously_allow_plaintext_credential_injection: true, + ..NetworkProxyConfig::default() + }, + ), + ( + "MITM hooks", + NetworkProxyConfig { + enabled: true, + mitm_hooks: vec![MitmHookConfig::default()], + ..NetworkProxyConfig::default() + }, + ), + ]; + + for (feature, config) in cases { + assert!( + RemoteNetworkProxyConfig::from_effective_config(&config).is_err(), + "{feature} must not cross the remote executor boundary" + ); + } +} + +#[test] +fn accepts_unsupported_configuration_when_proxy_is_disabled() { + let config = NetworkProxyConfig { + mitm: true, + credential_broker: true, + dangerously_allow_plaintext_credential_injection: true, + mitm_hooks: vec![MitmHookConfig::default()], + ..NetworkProxyConfig::default() + }; + + let remote = RemoteNetworkProxyConfig::from_effective_config(&config) + .expect("disabled proxy configuration does not cross the executor boundary"); + + assert!(!remote.enabled); +} + +#[test] +fn launch_config_materializes_audit_and_execution_attribution() { + let proxy = RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig { + enabled: true, + ..NetworkProxyConfig::default() + }) + .expect("supported remote config"); + let audit_metadata = NetworkProxyAuditMetadata { + conversation_id: Some("conversation-1".to_string()), + user_account_id: Some("account-1".to_string()), + originator: Some("codex_cli_rs".to_string()), + model: Some("model-1".to_string()), + ..NetworkProxyAuditMetadata::default() + }; + let state = NetworkProxyState::from_remote_launch_config(RemoteNetworkProxyLaunchConfig { + proxy, + audit_metadata: audit_metadata.clone(), + environment_id: Some("remote".to_string()), + execution_id: Some("execution-1".to_string()), + policy_decision_timeout_ms: None, + }) + .expect("remote launch state"); + + assert_eq!(state.audit_metadata(), &audit_metadata); + assert_eq!(state.environment_id(), Some("remote")); + assert_eq!(state.execution_id().as_deref(), Some("execution-1")); +} + +#[test] +fn policy_decision_callback_timeout_round_trips() { + let config = RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig::default()) + .expect("supported remote config"); + let mut launch = RemoteNetworkProxyLaunchConfig::new(config); + let without_timeout = serde_json::to_value(&launch).expect("serialize launch config"); + assert_eq!(without_timeout.get("policyDecisionTimeoutMs"), None); + launch.policy_decision_timeout_ms = Some(900_000); + let with_timeout = serde_json::to_value(&launch).expect("serialize launch timeout"); + assert_eq!(with_timeout["policyDecisionTimeoutMs"], 900_000); + assert_eq!( + serde_json::from_value::(with_timeout) + .expect("deserialize launch timeout"), + launch + ); +} diff --git a/codex-rs/network-proxy/src/request_cancellation.rs b/codex-rs/network-proxy/src/request_cancellation.rs new file mode 100644 index 0000000000000000000000000000000000000000..90d481acf1a310e2b99a7f614912407d23f27274 --- /dev/null +++ b/codex-rs/network-proxy/src/request_cancellation.rs @@ -0,0 +1,32 @@ +//! Records why a policy request was withdrawn before its decision future is dropped. + +use std::sync::Arc; +use std::sync::OnceLock; + +/// The controller's first known reason for withdrawing a policy request. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum NetworkRequestCancellationReason { + /// The process exited and its output streams closed normally. + ProcessFinished, + /// The process was explicitly terminated or its handle was abandoned. + ProcessCancelled, + /// The connection to the executor was lost. + ConnectionClosed, + /// The policy decision deadline expired. + TimedOut, +} + +/// Shared, in-process metadata; recording a reason does not authorize network access. +#[derive(Clone, Debug, Default)] +pub struct NetworkRequestCancellation(Arc>); + +impl NetworkRequestCancellation { + pub fn reason(&self) -> Option { + self.0.get().copied() + } + + /// Publish before dropping the decision future. Cleanup cannot replace an earlier cause. + pub fn record(&self, reason: NetworkRequestCancellationReason) { + let _ = self.0.set(reason); + } +} diff --git a/codex-rs/network-proxy/src/request_disconnect_tests.rs b/codex-rs/network-proxy/src/request_disconnect_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..f9ebc3a15e6353465018995fae72e94039fdda68 --- /dev/null +++ b/codex-rs/network-proxy/src/request_disconnect_tests.rs @@ -0,0 +1,41 @@ +use super::NetworkRequestDisconnect; +use pretty_assertions::assert_eq; +use std::future::pending; +use std::time::Duration; +use std::time::Instant; + +#[tokio::test] +async fn disconnect_is_published_before_policy_cleanup() { + struct ObserveOnDrop(NetworkRequestDisconnect); + impl Drop for ObserveOnDrop { + fn drop(&mut self) { + assert!(self.0.elapsed().is_some()); + } + } + + let disconnect = NetworkRequestDisconnect::default(); + let observer = ObserveOnDrop(disconnect.clone()); + let started_at = Instant::now(); + let decision = disconnect.track_http_request(started_at, async move { + let _observer = observer; + pending::<()>().await; + }); + assert!( + tokio::time::timeout(Duration::from_millis(1), decision) + .await + .is_err() + ); + assert!(disconnect.elapsed().expect("disconnect time") <= started_at.elapsed()); +} + +#[tokio::test] +async fn completed_policy_request_is_not_a_disconnect() { + let disconnect = NetworkRequestDisconnect::default(); + assert_eq!( + disconnect + .track_http_request(Instant::now(), async { 42 }) + .await, + 42 + ); + assert_eq!(disconnect.elapsed(), None); +} diff --git a/codex-rs/network-proxy/src/responses.rs b/codex-rs/network-proxy/src/responses.rs new file mode 100644 index 0000000000000000000000000000000000000000..d2aeb990e8f9deff4ad56c8b014aba4b7315fe9b --- /dev/null +++ b/codex-rs/network-proxy/src/responses.rs @@ -0,0 +1,121 @@ +use crate::network_policy::NetworkDecisionSource; +use crate::network_policy::NetworkPolicyDecision; +use crate::network_policy::NetworkProtocol; +use crate::reasons::REASON_DENIED; +use crate::reasons::REASON_METHOD_NOT_ALLOWED; +use crate::reasons::REASON_MITM_HOOK_DENIED; +use crate::reasons::REASON_MITM_REQUIRED; +use crate::reasons::REASON_NOT_ALLOWED; +use crate::reasons::REASON_NOT_ALLOWED_LOCAL; +use crate::reasons::REASON_PROXY_DISABLED; +use rama_http::Body; +use rama_http::Response; +use rama_http::StatusCode; +use serde::Serialize; +use tracing::error; + +pub struct PolicyDecisionDetails<'a> { + pub decision: NetworkPolicyDecision, + pub reason: &'a str, + pub source: NetworkDecisionSource, + pub protocol: NetworkProtocol, + pub host: &'a str, + pub port: u16, +} + +pub fn text_response(status: StatusCode, body: &str) -> Response { + Response::builder() + .status(status) + .header("content-type", "text/plain") + .body(Body::from(body.to_string())) + .unwrap_or_else(|_| Response::new(Body::from(body.to_string()))) +} + +pub fn json_response(value: &T) -> Response { + let body = match serde_json::to_string(value) { + Ok(body) => body, + Err(err) => { + error!("failed to serialize JSON response: {err}"); + "{}".to_string() + } + }; + Response::builder() + .status(StatusCode::OK) + .header("content-type", "application/json") + .body(Body::from(body)) + .unwrap_or_else(|err| { + error!("failed to build JSON response: {err}"); + Response::new(Body::from("{}")) + }) +} + +pub fn blocked_header_value(reason: &str) -> &'static str { + match reason { + REASON_NOT_ALLOWED | REASON_NOT_ALLOWED_LOCAL => "blocked-by-allowlist", + REASON_DENIED => "blocked-by-denylist", + REASON_METHOD_NOT_ALLOWED => "blocked-by-method-policy", + REASON_MITM_HOOK_DENIED => "blocked-by-mitm-hook", + REASON_MITM_REQUIRED => "blocked-by-mitm-required", + _ => "blocked-by-policy", + } +} + +pub fn blocked_message(reason: &str) -> &'static str { + match reason { + REASON_NOT_ALLOWED => "Domain not in allowlist.", + REASON_NOT_ALLOWED_LOCAL => "Sandbox policy blocks local/private network addresses.", + REASON_DENIED => "Domain denied by the sandbox policy.", + REASON_METHOD_NOT_ALLOWED => "Method not allowed in limited mode.", + REASON_MITM_HOOK_DENIED => "HTTPS request denied by MITM hook policy.", + REASON_MITM_REQUIRED => "MITM required for limited HTTPS.", + REASON_PROXY_DISABLED => "network proxy is disabled", + _ => "Request blocked by network policy.", + } +} + +pub fn blocked_text_response(reason: &str) -> Response { + Response::builder() + .status(StatusCode::FORBIDDEN) + .header("content-type", "text/plain") + .header("x-proxy-error", blocked_header_value(reason)) + .body(Body::from(blocked_message(reason))) + .unwrap_or_else(|_| Response::new(Body::from("blocked"))) +} +pub fn blocked_message_with_policy(reason: &str, details: &PolicyDecisionDetails<'_>) -> String { + let _ = (details.reason, details.host); + blocked_message(reason).to_string() +} + +pub fn blocked_text_response_with_policy( + reason: &str, + details: &PolicyDecisionDetails<'_>, +) -> Response { + Response::builder() + .status(StatusCode::FORBIDDEN) + .header("content-type", "text/plain") + .header("x-proxy-error", blocked_header_value(reason)) + .body(Body::from(blocked_message_with_policy(reason, details))) + .unwrap_or_else(|_| Response::new(Body::from("blocked"))) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::reasons::REASON_NOT_ALLOWED; + use pretty_assertions::assert_eq; + + #[test] + fn blocked_message_with_policy_returns_human_message() { + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Ask, + reason: REASON_NOT_ALLOWED, + source: NetworkDecisionSource::Decider, + protocol: NetworkProtocol::HttpsConnect, + host: "api.example.com", + port: 443, + }; + + let message = blocked_message_with_policy(REASON_NOT_ALLOWED, &details); + assert_eq!(message, "Domain not in allowlist."); + } +} diff --git a/codex-rs/network-proxy/src/socks5.rs b/codex-rs/network-proxy/src/socks5.rs new file mode 100644 index 0000000000000000000000000000000000000000..82a7093dabc331445694b187f45a1a2ec4ac01fe --- /dev/null +++ b/codex-rs/network-proxy/src/socks5.rs @@ -0,0 +1,1202 @@ +use crate::attribution::BindConnectionAttribution; +use crate::config::NetworkMode; +use crate::connect_policy::TargetCheckedTcpConnector; +use crate::connection_lifecycle::CancelOnShutdown; +use crate::mitm; +use crate::network_policy::BlockDecisionAuditEventArgs; +use crate::network_policy::NetworkDecision; +use crate::network_policy::NetworkDecisionSource; +use crate::network_policy::NetworkPolicyDecider; +use crate::network_policy::NetworkPolicyDecision; +use crate::network_policy::NetworkPolicyRequest; +use crate::network_policy::NetworkPolicyRequestArgs; +use crate::network_policy::NetworkProtocol; +use crate::network_policy::emit_block_decision_audit_event; +use crate::network_policy::evaluate_host_policy; +use crate::policy::normalize_host; +use crate::reasons::REASON_METHOD_NOT_ALLOWED; +use crate::reasons::REASON_MITM_REQUIRED; +use crate::reasons::REASON_PROXY_DISABLED; +use crate::responses::PolicyDecisionDetails; +use crate::responses::blocked_message_with_policy; +use crate::runtime::HostMitmRequirement; +use crate::state::BlockedRequest; +use crate::state::BlockedRequestArgs; +use crate::state::NetworkProxyState; +use anyhow::Context as _; +use anyhow::Result; +use rama_core::Service; +use rama_core::error::BoxError; +use rama_core::extensions::Extensions; +use rama_core::extensions::ExtensionsMut; +use rama_core::extensions::ExtensionsRef; +use rama_core::graceful::ShutdownGuard; +use rama_core::service::BoxService; +use rama_core::service::service_fn; +use rama_net::address::HostWithPort; +use rama_net::client::EstablishedClientConnection; +use rama_net::proxy::ProxyRequest; +use rama_net::proxy::ProxyTarget; +use rama_net::proxy::StreamForwardService; +use rama_net::stream::Socket; +use rama_net::stream::SocketInfo; +use rama_socks5::Socks5Acceptor; +use rama_socks5::server::DefaultConnector; +use rama_socks5::server::DefaultUdpRelay; +use rama_socks5::server::udp::RelayRequest; +use rama_socks5::server::udp::RelayResponse; +use rama_tcp::TcpStream; +use rama_tcp::client::Request as TcpRequest; +use rama_tcp::server::TcpListener; +use std::io; +use std::net::SocketAddr; +use std::net::TcpListener as StdTcpListener; +use std::pin::Pin; +use std::sync::Arc; +use std::task::Context as TaskContext; +use std::task::Poll; +use std::time::Instant; +use tokio::io::AsyncRead; +use tokio::io::AsyncWrite; +use tokio::io::ReadBuf; +use tracing::error; +use tracing::info; +use tracing::warn; + +pub async fn run_socks5( + state: Arc, + addr: SocketAddr, + policy_decider: Option>, + environment_id: Option, + enable_socks5_udp: bool, + guard: ShutdownGuard, +) -> Result<()> { + let listener = TcpListener::build() + .bind(addr) + .await + // See `http_proxy.rs` for details on why we wrap `BoxError` before converting to anyhow. + .map_err(rama_core::error::OpaqueError::from) + .map_err(anyhow::Error::from) + .with_context(|| format!("bind SOCKS5 proxy: {addr}"))?; + + run_socks5_with_listener( + state, + listener, + policy_decider, + environment_id, + enable_socks5_udp, + guard, + ) + .await +} + +pub async fn run_socks5_with_std_listener( + state: Arc, + listener: StdTcpListener, + policy_decider: Option>, + environment_id: Option, + enable_socks5_udp: bool, + guard: ShutdownGuard, +) -> Result<()> { + let listener = + TcpListener::try_from(listener).context("convert std listener to SOCKS5 proxy listener")?; + run_socks5_with_listener( + state, + listener, + policy_decider, + environment_id, + enable_socks5_udp, + guard, + ) + .await +} + +async fn run_socks5_with_listener( + state: Arc, + listener: TcpListener, + policy_decider: Option>, + environment_id: Option, + enable_socks5_udp: bool, + guard: ShutdownGuard, +) -> Result<()> { + let addr = listener + .local_addr() + .context("read SOCKS5 listener local addr")?; + + info!("SOCKS5 proxy listening on {addr}"); + + match state.network_mode().await { + Ok(NetworkMode::Limited) => { + info!( + "SOCKS5 UDP and non-HTTPS SOCKS5 TCP are blocked in limited mode; HTTPS SOCKS5 TCP requires MITM inspection" + ); + } + Ok(NetworkMode::Full) => {} + Err(err) => { + warn!("failed to read network mode: {err}"); + } + } + + listener + .serve_graceful( + guard, + CancelOnShutdown::new(socks5_proxy_service( + state, + policy_decider, + environment_id, + enable_socks5_udp, + )), + ) + .await; + Ok(()) +} + +pub(crate) fn socks5_proxy_service( + state: Arc, + policy_decider: Option>, + environment_id: Option, + enable_socks5_udp: bool, +) -> BoxService { + let tcp_connector = TargetCheckedTcpConnector::new(state.clone()); + let policy_tcp_connector = service_fn({ + let policy_decider = policy_decider.clone(); + let environment_id = environment_id.clone(); + move |req: TcpRequest| { + let tcp_connector = tcp_connector.clone(); + let policy_decider = policy_decider.clone(); + let environment_id = environment_id.clone(); + async move { handle_socks5_tcp(req, tcp_connector, policy_decider, environment_id).await } + } + }); + + let socks_proxy = service_fn(|request| async move { proxy_socks5_tcp(request).await }); + let socks_connector = DefaultConnector::default() + .with_connector(policy_tcp_connector) + .with_service(socks_proxy); + let base = Socks5Acceptor::new().with_connector(socks_connector); + + if enable_socks5_udp { + let udp_state = state.clone(); + let udp_decider = policy_decider.clone(); + let udp_relay = + DefaultUdpRelay::default().with_async_inspector(service_fn({ + let environment_id = environment_id.clone(); + move |request: RelayRequest| { + let udp_state = udp_state.clone(); + let udp_decider = udp_decider.clone(); + let environment_id = environment_id.clone(); + async move { + inspect_socks5_udp(request, udp_state, udp_decider, environment_id).await + } + } + })); + let socks_acceptor = base.with_udp_associator(udp_relay); + BindConnectionAttribution::new(socks_acceptor, state, environment_id).boxed() + } else { + BindConnectionAttribution::new(base, state, environment_id).boxed() + } +} + +async fn handle_socks5_tcp( + req: TcpRequest, + tcp_connector: TargetCheckedTcpConnector, + policy_decider: Option>, + environment_id: Option, +) -> Result, BoxError> { + let app_state = req + .extensions() + .get::>() + .cloned() + .ok_or_else(|| io::Error::other("missing state"))?; + + let host = normalize_host(&req.authority.host.to_string()); + let port = req.authority.port; + let target = req.authority.clone(); + if host.is_empty() { + return Err(io::Error::new(io::ErrorKind::InvalidInput, "invalid host").into()); + } + + let client = req + .extensions() + .get::() + .map(|info| info.peer_addr().to_string()); + + match app_state.enabled().await { + Ok(true) => {} + Ok(false) => { + emit_socks_block_decision_audit_event( + &app_state, + NetworkDecisionSource::ProxyState, + REASON_PROXY_DISABLED, + NetworkProtocol::Socks5Tcp, + host.as_str(), + port, + client.as_deref(), + ); + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Deny, + reason: REASON_PROXY_DISABLED, + source: NetworkDecisionSource::ProxyState, + protocol: NetworkProtocol::Socks5Tcp, + host: &host, + port, + }; + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: REASON_PROXY_DISABLED.to_string(), + client: client.clone(), + method: None, + mode: None, + protocol: "socks5".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!("SOCKS blocked; proxy disabled (client={client}, host={host})"); + return Err(policy_denied_error(REASON_PROXY_DISABLED, &details).into()); + } + Err(err) => { + error!("failed to read enabled state: {err}"); + return Err(io::Error::other("proxy error").into()); + } + } + + let mode = match app_state.network_mode().await { + Ok(mode) => mode, + Err(err) => { + error!("failed to evaluate method policy: {err}"); + return Err(io::Error::other("proxy error").into()); + } + }; + let host_mitm_requirement = match app_state.host_mitm_requirement(&host, port).await { + Ok(requirement) => requirement, + Err(err) => { + error!("failed to inspect MITM requirements for {host}: {err}"); + return Err(io::Error::other("proxy error").into()); + } + }; + // Otherwise retain the existing limited-mode restriction to the default HTTPS port. + let brokered_http = matches!(host_mitm_requirement, HostMitmRequirement::Credential(protocols) if protocols.http); + let socks5_tcp_target_is_https = port == 443; + if mode == NetworkMode::Limited && !socks5_tcp_target_is_https && !brokered_http { + emit_socks_block_decision_audit_event( + &app_state, + NetworkDecisionSource::ModeGuard, + REASON_METHOD_NOT_ALLOWED, + NetworkProtocol::Socks5Tcp, + host.as_str(), + port, + client.as_deref(), + ); + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Deny, + reason: REASON_METHOD_NOT_ALLOWED, + source: NetworkDecisionSource::ModeGuard, + protocol: NetworkProtocol::Socks5Tcp, + host: &host, + port, + }; + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: REASON_METHOD_NOT_ALLOWED.to_string(), + client: client.clone(), + method: None, + mode: Some(NetworkMode::Limited), + protocol: "socks5".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!( + "SOCKS blocked; limited mode only supports HTTPS MITM (client={client}, host={host}, port={port})" + ); + return Err(policy_denied_error(REASON_METHOD_NOT_ALLOWED, &details).into()); + } + + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Socks5Tcp, + host: host.clone(), + port, + environment_id, + client_addr: client.clone(), + method: None, + command: None, + exec_policy_hint: None, + }); + + match evaluate_host_policy(&app_state, policy_decider.as_ref(), &request).await { + Ok(NetworkDecision::Deny { + reason, + source, + decision, + }) => { + let details = PolicyDecisionDetails { + decision, + reason: &reason, + source, + protocol: NetworkProtocol::Socks5Tcp, + host: &host, + port, + }; + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: reason.clone(), + client: client.clone(), + method: None, + mode: None, + protocol: "socks5".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!("SOCKS blocked (client={client}, host={host}, reason={reason})"); + return Err(policy_denied_error(&reason, &details).into()); + } + Ok(NetworkDecision::Allow) => { + let client = client.as_deref().unwrap_or_default(); + info!("SOCKS allowed (client={client}, host={host}, port={port})"); + } + Err(err) => { + error!("failed to evaluate host: {err}"); + return Err(io::Error::other("proxy error").into()); + } + } + + let mitm_state = match app_state.mitm_state().await { + Ok(state) => state, + Err(err) => { + error!("failed to load MITM state: {err}"); + return Err(io::Error::other("proxy error").into()); + } + }; + let socks_mitm_mode = if mode == NetworkMode::Limited && !brokered_http { + SocksMitmMode::Enabled + } else { + match host_mitm_requirement { + HostMitmRequirement::None => SocksMitmMode::Disabled, + HostMitmRequirement::Credential(protocols) => SocksMitmMode::DetectProtocol(protocols), + HostMitmRequirement::Always => SocksMitmMode::Enabled, + } + }; + let unsupported_hook_protocol = + host_mitm_requirement == HostMitmRequirement::Always && !socks5_tcp_target_is_https; + if unsupported_hook_protocol + || (socks_mitm_mode != SocksMitmMode::Disabled && mitm_state.is_none()) + { + emit_socks_block_decision_audit_event( + &app_state, + NetworkDecisionSource::ModeGuard, + REASON_MITM_REQUIRED, + NetworkProtocol::Socks5Tcp, + host.as_str(), + port, + client.as_deref(), + ); + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Deny, + reason: REASON_MITM_REQUIRED, + source: NetworkDecisionSource::ModeGuard, + protocol: NetworkProtocol::Socks5Tcp, + host: &host, + port, + }; + let _ = app_state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: REASON_MITM_REQUIRED.to_string(), + client: client.clone(), + method: None, + mode: Some(mode), + protocol: "socks5".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!( + "SOCKS blocked; MITM required to enforce HTTPS policy (client={client}, host={host}, mode={mode:?}, host_mitm_requirement={host_mitm_requirement:?}, https_target={socks5_tcp_target_is_https})" + ); + return Err(policy_denied_error(REASON_MITM_REQUIRED, &details).into()); + } + + if let Some(mitm_state) = mitm_state { + let client = client.as_deref().unwrap_or_default(); + let conn = match socks_mitm_mode { + SocksMitmMode::Disabled => None, + SocksMitmMode::Enabled => Some(Socks5TcpConnection::Mitm { + target, + mode, + mitm: mitm_state, + extensions: Extensions::new(), + }), + SocksMitmMode::DetectProtocol(protocols) => Some(Socks5TcpConnection::DetectProtocol { + protocols, + target, + mode, + mitm: mitm_state, + state: app_state, + extensions: Extensions::new(), + }), + }; + if let Some(conn) = conn { + info!( + "SOCKS MITM selected (client={client}, host={host}, port={port}, mode={mode:?}, mitm_mode={socks_mitm_mode:?})" + ); + return Ok(EstablishedClientConnection { input: req, conn }); + } + } + + info!("SOCKS upstream dial started (host={host}, port={port})"); + let connect_started_at = Instant::now(); + let result = tcp_connector.serve(req).await.map(|connection| { + let EstablishedClientConnection { input, conn } = connection; + EstablishedClientConnection { + input, + conn: Socks5TcpConnection::Direct(conn), + } + }); + match &result { + Ok(_) => info!( + "SOCKS upstream dial established (host={host}, port={port}, elapsed_ms={})", + connect_started_at.elapsed().as_millis() + ), + Err(_) => warn!( + "SOCKS upstream dial failed (host={host}, port={port}, elapsed_ms={})", + connect_started_at.elapsed().as_millis() + ), + } + result +} + +/// Internal connector output for SOCKS5 TCP. MITM requests do not dial upstream before the +/// inner HTTPS request is inspected, so they carry the target metadata instead of a socket. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum SocksMitmMode { + Disabled, + Enabled, + DetectProtocol(crate::brokered_tunnel::BrokeredProtocols), +} + +#[derive(Debug)] +enum Socks5TcpConnection { + Direct(TcpStream), + Mitm { + target: HostWithPort, + mode: NetworkMode, + mitm: Arc, + extensions: Extensions, + }, + DetectProtocol { + protocols: crate::brokered_tunnel::BrokeredProtocols, + target: HostWithPort, + mode: NetworkMode, + mitm: Arc, + state: Arc, + extensions: Extensions, + }, +} + +impl AsyncRead for Socks5TcpConnection { + fn poll_read( + self: Pin<&mut Self>, + cx: &mut TaskContext<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + match self.get_mut() { + Self::Direct(stream) => Pin::new(stream).poll_read(cx, buf), + Self::Mitm { .. } | Self::DetectProtocol { .. } => Poll::Ready(Ok(())), + } + } +} + +impl AsyncWrite for Socks5TcpConnection { + fn poll_write( + self: Pin<&mut Self>, + cx: &mut TaskContext<'_>, + buf: &[u8], + ) -> Poll> { + match self.get_mut() { + Self::Direct(stream) => Pin::new(stream).poll_write(cx, buf), + Self::Mitm { .. } | Self::DetectProtocol { .. } => Poll::Ready(Ok(buf.len())), + } + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { + match self.get_mut() { + Self::Direct(stream) => Pin::new(stream).poll_flush(cx), + Self::Mitm { .. } | Self::DetectProtocol { .. } => Poll::Ready(Ok(())), + } + } + + fn poll_shutdown(self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll> { + match self.get_mut() { + Self::Direct(stream) => Pin::new(stream).poll_shutdown(cx), + Self::Mitm { .. } | Self::DetectProtocol { .. } => Poll::Ready(Ok(())), + } + } +} + +impl Socket for Socks5TcpConnection { + fn local_addr(&self) -> io::Result { + match self { + Self::Direct(stream) => stream.local_addr(), + Self::Mitm { .. } | Self::DetectProtocol { .. } => { + Ok(SocketAddr::from(([0, 0, 0, 0], 0))) + } + } + } + + fn peer_addr(&self) -> io::Result { + match self { + Self::Direct(stream) => stream.peer_addr(), + Self::Mitm { .. } | Self::DetectProtocol { .. } => { + Ok(SocketAddr::from(([0, 0, 0, 0], 0))) + } + } + } +} + +impl ExtensionsRef for Socks5TcpConnection { + fn extensions(&self) -> &Extensions { + match self { + Self::Direct(stream) => stream.extensions(), + Self::Mitm { extensions, .. } | Self::DetectProtocol { extensions, .. } => extensions, + } + } +} + +impl ExtensionsMut for Socks5TcpConnection { + fn extensions_mut(&mut self) -> &mut Extensions { + match self { + Self::Direct(stream) => stream.extensions_mut(), + Self::Mitm { extensions, .. } | Self::DetectProtocol { extensions, .. } => extensions, + } + } +} + +async fn proxy_socks5_tcp( + request: ProxyRequest, +) -> Result<(), BoxError> { + let ProxyRequest { mut source, target } = request; + match target { + Socks5TcpConnection::Direct(target) => StreamForwardService::default() + .serve(ProxyRequest { source, target }) + .await + .map_err(Into::into), + Socks5TcpConnection::Mitm { + target, mode, mitm, .. + } => { + source.extensions_mut().insert(ProxyTarget(target)); + source.extensions_mut().insert(mode); + source.extensions_mut().insert(mitm); + mitm::mitm_stream(source, rama_http::uri::Scheme::HTTPS) + .await + .map_err(Into::into) + } + Socks5TcpConnection::DetectProtocol { + protocols, + target, + mode, + mitm, + state, + .. + } => { + source.extensions_mut().insert(ProxyTarget(target.clone())); + source.extensions_mut().insert(mode); + source.extensions_mut().insert(mitm); + let (protocol, source) = crate::brokered_tunnel::peek_protocol(source, protocols) + .await + .map_err(|err| -> BoxError { err.into() })?; + match protocol { + crate::brokered_tunnel::TunnelProtocol::Tls => { + mitm::mitm_stream(source, rama_http::uri::Scheme::HTTPS) + .await + .map_err(Into::into) + } + crate::brokered_tunnel::TunnelProtocol::Http => { + mitm::mitm_stream(source, rama_http::uri::Scheme::HTTP) + .await + .map_err(Into::into) + } + crate::brokered_tunnel::TunnelProtocol::Opaque => { + if mode == NetworkMode::Limited { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "opaque tunnels are not allowed in limited mode", + ) + .into()); + } + info!("SOCKS opaque upstream dial started (target={target})"); + let connect_started_at = Instant::now(); + let EstablishedClientConnection { conn: upstream, .. } = + TargetCheckedTcpConnector::new(state) + .serve(TcpRequest::new(target.clone())) + .await?; + info!( + "SOCKS opaque upstream dial established (target={target}, elapsed_ms={})", + connect_started_at.elapsed().as_millis() + ); + StreamForwardService::default() + .serve(ProxyRequest { + source, + target: upstream, + }) + .await + .map_err(Into::into) + } + } + } + } +} + +async fn inspect_socks5_udp( + request: RelayRequest, + state: Arc, + policy_decider: Option>, + environment_id: Option, +) -> io::Result { + let RelayRequest { + server_address, + payload, + extensions, + .. + } = request; + + let host = normalize_host(&server_address.ip_addr.to_string()); + let port = server_address.port; + if host.is_empty() { + return Err(io::Error::new(io::ErrorKind::InvalidInput, "invalid host")); + } + + let client = extensions + .get::() + .map(|info| info.peer_addr().to_string()); + + match state.enabled().await { + Ok(true) => {} + Ok(false) => { + emit_socks_block_decision_audit_event( + &state, + NetworkDecisionSource::ProxyState, + REASON_PROXY_DISABLED, + NetworkProtocol::Socks5Udp, + host.as_str(), + port, + client.as_deref(), + ); + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Deny, + reason: REASON_PROXY_DISABLED, + source: NetworkDecisionSource::ProxyState, + protocol: NetworkProtocol::Socks5Udp, + host: &host, + port, + }; + let _ = state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: REASON_PROXY_DISABLED.to_string(), + client: client.clone(), + method: None, + mode: None, + protocol: "socks5-udp".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!("SOCKS UDP blocked; proxy disabled (client={client}, host={host})"); + return Err(policy_denied_error(REASON_PROXY_DISABLED, &details)); + } + Err(err) => { + error!("failed to read enabled state: {err}"); + return Err(io::Error::other("proxy error")); + } + } + + match state.network_mode().await { + Ok(NetworkMode::Limited) => { + emit_socks_block_decision_audit_event( + &state, + NetworkDecisionSource::ModeGuard, + REASON_METHOD_NOT_ALLOWED, + NetworkProtocol::Socks5Udp, + host.as_str(), + port, + client.as_deref(), + ); + let details = PolicyDecisionDetails { + decision: NetworkPolicyDecision::Deny, + reason: REASON_METHOD_NOT_ALLOWED, + source: NetworkDecisionSource::ModeGuard, + protocol: NetworkProtocol::Socks5Udp, + host: &host, + port, + }; + let _ = state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: REASON_METHOD_NOT_ALLOWED.to_string(), + client: client.clone(), + method: None, + mode: Some(NetworkMode::Limited), + protocol: "socks5-udp".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + return Err(policy_denied_error(REASON_METHOD_NOT_ALLOWED, &details)); + } + Ok(NetworkMode::Full) => {} + Err(err) => { + error!("failed to evaluate method policy: {err}"); + return Err(io::Error::other("proxy error")); + } + } + + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::Socks5Udp, + host: host.clone(), + port, + environment_id, + client_addr: client.clone(), + method: None, + command: None, + exec_policy_hint: None, + }); + + match evaluate_host_policy(&state, policy_decider.as_ref(), &request).await { + Ok(NetworkDecision::Deny { + reason, + source, + decision, + }) => { + let details = PolicyDecisionDetails { + decision, + reason: &reason, + source, + protocol: NetworkProtocol::Socks5Udp, + host: &host, + port, + }; + let _ = state + .record_blocked(BlockedRequest::new(BlockedRequestArgs { + host: host.clone(), + reason: reason.clone(), + client: client.clone(), + method: None, + mode: None, + protocol: "socks5-udp".to_string(), + decision: Some(details.decision.as_str().to_string()), + source: Some(details.source.as_str().to_string()), + port: Some(port), + })) + .await; + let client = client.as_deref().unwrap_or_default(); + warn!("SOCKS UDP blocked (client={client}, host={host}, reason={reason})"); + Err(policy_denied_error(&reason, &details)) + } + Ok(NetworkDecision::Allow) => Ok(RelayResponse { + maybe_payload: Some(payload), + extensions, + }), + Err(err) => { + error!("failed to evaluate UDP host: {err}"); + Err(io::Error::other("proxy error")) + } + } +} + +fn emit_socks_block_decision_audit_event( + state: &NetworkProxyState, + source: NetworkDecisionSource, + reason: &str, + protocol: NetworkProtocol, + host: &str, + port: u16, + client_addr: Option<&str>, +) { + emit_block_decision_audit_event( + state, + BlockDecisionAuditEventArgs { + source, + reason, + protocol, + server_address: host, + server_port: port, + method: None, + client_addr, + }, + ); +} + +fn policy_denied_error(reason: &str, details: &PolicyDecisionDetails<'_>) -> io::Error { + io::Error::new( + io::ErrorKind::PermissionDenied, + blocked_message_with_policy(reason, details), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::NetworkMode; + use crate::config::NetworkProxyConfig; + use crate::mitm_hook::MitmHookConfig; + use crate::mitm_hook::MitmHookMatchConfig; + use crate::network_policy::test_support::POLICY_DECISION_EVENT_NAME; + use crate::network_policy::test_support::capture_events; + use crate::network_policy::test_support::find_event_by_name; + use crate::runtime::ConfigReloader; + use crate::runtime::ConfigReloaderFuture; + use crate::runtime::ConfigState; + use crate::state::NetworkProxyConstraints; + use crate::state::build_config_state; + use pretty_assertions::assert_eq; + use rama_core::extensions::Extensions; + use rama_core::extensions::ExtensionsMut; + use rama_net::address::HostWithPort; + use rama_net::address::SocketAddress; + use rama_socks5::server::udp::RelayDirection; + use std::collections::HashMap; + use std::net::IpAddr; + use std::net::Ipv4Addr; + use std::sync::Arc; + use std::sync::Mutex; + + // Managed MITM CA files live under the shared test CODEX_HOME, so MITM-enabled config state + // must be materialized one test at a time. + static MITM_CONFIG_STATE_LOCK: Mutex<()> = Mutex::new(()); + + #[derive(Clone)] + struct StaticReloader { + state: ConfigState, + } + + impl ConfigReloader for StaticReloader { + fn maybe_reload(&self) -> ConfigReloaderFuture<'_, Option> { + Box::pin(async { Ok(None) }) + } + + fn reload_now(&self) -> ConfigReloaderFuture<'_, ConfigState> { + Box::pin(async { Ok(self.state.clone()) }) + } + + fn source_label(&self) -> String { + "static test reloader".to_string() + } + } + + fn state_for_settings(network: NetworkProxyConfig) -> Arc { + let config = network; + let _mitm_config_state_guard = config.mitm.then(|| MITM_CONFIG_STATE_LOCK.lock().unwrap()); + let state = build_config_state(config, NetworkProxyConstraints::default()).unwrap(); + let reloader = Arc::new(StaticReloader { + state: state.clone(), + }); + Arc::new(NetworkProxyState::with_reloader(state, reloader)) + } + + #[tokio::test(flavor = "current_thread")] + async fn handle_socks5_tcp_emits_block_decision_for_proxy_disabled() { + let state = state_for_settings(NetworkProxyConfig { + enabled: false, + mode: NetworkMode::Full, + ..NetworkProxyConfig::default() + }); + let mut request = + TcpRequest::new(HostWithPort::try_from("example.com:443").expect("valid authority")); + request.extensions_mut().insert(state.clone()); + + let (result, events) = capture_events(|| async { + handle_socks5_tcp( + request, + TargetCheckedTcpConnector::new(state.clone()), + /*policy_decider*/ None, + /*environment_id*/ None, + ) + .await + }) + .await; + assert!(result.is_err(), "proxy-disabled request should be denied"); + + let event = find_event_by_name(&events, POLICY_DECISION_EVENT_NAME) + .expect("expected policy decision event"); + assert_eq!(event.field("network.policy.scope"), Some("non_domain")); + assert_eq!(event.field("network.policy.decision"), Some("deny")); + assert_eq!(event.field("network.policy.source"), Some("proxy_state")); + assert_eq!( + event.field("network.policy.reason"), + Some(REASON_PROXY_DISABLED) + ); + assert_eq!( + event.field("network.transport.protocol"), + Some("socks5_tcp") + ); + assert_eq!(event.field("server.address"), Some("example.com")); + assert_eq!(event.field("server.port"), Some("443")); + assert_eq!(event.field("http.request.method"), Some("none")); + assert_eq!(event.field("client.address"), Some("unknown")); + } + + #[tokio::test(flavor = "current_thread")] + async fn handle_socks5_tcp_uses_mitm_in_limited_mode() { + let mut settings = NetworkProxyConfig { + enabled: true, + mode: NetworkMode::Limited, + mitm: true, + ..NetworkProxyConfig::default() + }; + settings.set_allowed_domains(vec!["example.com".to_string()]); + let state = state_for_settings(settings); + let mut request = + TcpRequest::new(HostWithPort::try_from("example.com:443").expect("valid authority")); + request.extensions_mut().insert(state.clone()); + + let result = handle_socks5_tcp( + request, + TargetCheckedTcpConnector::new(state), + /*policy_decider*/ None, + /*environment_id*/ None, + ) + .await + .expect("limited-mode HTTPS should use MITM"); + + assert!(matches!(result.conn, Socks5TcpConnection::Mitm { .. })); + } + + #[tokio::test(flavor = "current_thread")] + async fn handle_socks5_tcp_blocks_non_https_in_limited_mode() { + let mut settings = NetworkProxyConfig { + enabled: true, + mode: NetworkMode::Limited, + ..NetworkProxyConfig::default() + }; + settings.set_allowed_domains(vec!["example.com".to_string()]); + let state = state_for_settings(settings); + let mut request = + TcpRequest::new(HostWithPort::try_from("example.com:80").expect("valid authority")); + request.extensions_mut().insert(state.clone()); + + let (result, events) = capture_events(|| async { + handle_socks5_tcp( + request, + TargetCheckedTcpConnector::new(state), + /*policy_decider*/ None, + /*environment_id*/ None, + ) + .await + }) + .await; + assert!( + result.is_err(), + "limited-mode non-HTTPS SOCKS should be denied" + ); + + let event = find_event_by_name(&events, POLICY_DECISION_EVENT_NAME) + .expect("expected policy decision event"); + assert_eq!(event.field("network.policy.scope"), Some("non_domain")); + assert_eq!(event.field("network.policy.decision"), Some("deny")); + assert_eq!(event.field("network.policy.source"), Some("mode_guard")); + assert_eq!( + event.field("network.policy.reason"), + Some(REASON_METHOD_NOT_ALLOWED) + ); + assert_eq!( + event.field("network.transport.protocol"), + Some("socks5_tcp") + ); + assert_eq!(event.field("server.address"), Some("example.com")); + assert_eq!(event.field("server.port"), Some("80")); + assert_eq!(event.field("http.request.method"), Some("none")); + assert_eq!(event.field("client.address"), Some("unknown")); + } + + #[tokio::test(flavor = "current_thread")] + async fn handle_socks5_tcp_detects_tls_for_brokered_nonstandard_port_in_full_mode() { + let mut settings = NetworkProxyConfig { + enabled: true, + mode: NetworkMode::Full, + mitm: true, + credential_broker: true, + ..NetworkProxyConfig::default() + }; + settings.set_allowed_domains(vec!["api.openai.com".to_string()]); + let state = state_for_settings(settings); + let mut env = HashMap::from([("OPENAI_API_KEY".to_string(), "sk-real".to_string())]); + state.virtualize_child_credentials(&mut env); + let mut request = TcpRequest::new( + HostWithPort::try_from("api.openai.com:8443").expect("valid authority"), + ); + request.extensions_mut().insert(state.clone()); + + let result = handle_socks5_tcp( + request, + TargetCheckedTcpConnector::new(state), + /*policy_decider*/ None, + /*environment_id*/ None, + ) + .await + .expect("brokered TLS should defer MITM until protocol detection"); + + assert!(matches!( + result.conn, + Socks5TcpConnection::DetectProtocol { .. } + )); + } + + #[tokio::test(flavor = "current_thread")] + async fn handle_socks5_tcp_blocks_limited_mode_without_mitm_state() { + let mut settings = NetworkProxyConfig { + enabled: true, + mode: NetworkMode::Limited, + ..NetworkProxyConfig::default() + }; + settings.set_allowed_domains(vec!["example.com".to_string()]); + let state = state_for_settings(settings); + let mut request = + TcpRequest::new(HostWithPort::try_from("example.com:443").expect("valid authority")); + request.extensions_mut().insert(state.clone()); + + let err = handle_socks5_tcp( + request, + TargetCheckedTcpConnector::new(state), + /*policy_decider*/ None, + /*environment_id*/ None, + ) + .await + .expect_err("limited-mode HTTPS requires MITM"); + + assert!( + format!("{err:?}").contains("MITM required"), + "unexpected error: {err:?}" + ); + } + + #[tokio::test(flavor = "current_thread")] + async fn handle_socks5_tcp_uses_mitm_for_hooked_host_in_full_mode() { + let mut settings = NetworkProxyConfig { + enabled: true, + mode: NetworkMode::Full, + mitm: true, + mitm_hooks: vec![MitmHookConfig { + host: "api.github.com".to_string(), + matcher: MitmHookMatchConfig { + methods: vec!["POST".to_string()], + path_prefixes: vec!["/repos/openai/".to_string()], + ..MitmHookMatchConfig::default() + }, + ..MitmHookConfig::default() + }], + ..NetworkProxyConfig::default() + }; + settings.set_allowed_domains(vec!["api.github.com".to_string()]); + let state = state_for_settings(settings); + let mut request = + TcpRequest::new(HostWithPort::try_from("api.github.com:443").expect("valid authority")); + request.extensions_mut().insert(state.clone()); + + let result = handle_socks5_tcp( + request, + TargetCheckedTcpConnector::new(state), + /*policy_decider*/ None, + /*environment_id*/ None, + ) + .await + .expect("hooked HTTPS should use MITM"); + + assert!(matches!(result.conn, Socks5TcpConnection::Mitm { .. })); + } + + #[tokio::test(flavor = "current_thread")] + async fn handle_socks5_tcp_blocks_hooked_non_https_host_in_full_mode() { + let mut settings = NetworkProxyConfig { + enabled: true, + mode: NetworkMode::Full, + mitm: true, + mitm_hooks: vec![MitmHookConfig { + host: "api.github.com".to_string(), + matcher: MitmHookMatchConfig { + methods: vec!["POST".to_string()], + path_prefixes: vec!["/repos/openai/".to_string()], + ..MitmHookMatchConfig::default() + }, + ..MitmHookConfig::default() + }], + ..NetworkProxyConfig::default() + }; + settings.set_allowed_domains(vec!["api.github.com".to_string()]); + let state = state_for_settings(settings); + let mut request = + TcpRequest::new(HostWithPort::try_from("api.github.com:80").expect("valid authority")); + request.extensions_mut().insert(state.clone()); + + let err = handle_socks5_tcp( + request, + TargetCheckedTcpConnector::new(state), + /*policy_decider*/ None, + /*environment_id*/ None, + ) + .await + .expect_err("hooked non-HTTPS SOCKS should require MITM"); + + assert!( + format!("{err:?}").contains("MITM required"), + "unexpected error: {err:?}" + ); + } + + #[tokio::test(flavor = "current_thread")] + async fn inspect_socks5_udp_emits_block_decision_for_mode_guard_deny() { + let state = state_for_settings(NetworkProxyConfig { + enabled: true, + mode: NetworkMode::Limited, + ..NetworkProxyConfig::default() + }); + let request = RelayRequest { + direction: RelayDirection::South, + server_address: SocketAddress::new(IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34)), 53), + payload: Default::default(), + extensions: Extensions::new(), + }; + + let (result, events) = capture_events(|| async { + inspect_socks5_udp( + request, state, /*policy_decider*/ None, /*environment_id*/ None, + ) + .await + }) + .await; + assert!(result.is_err(), "limited-mode UDP request should be denied"); + + let event = find_event_by_name(&events, POLICY_DECISION_EVENT_NAME) + .expect("expected policy decision event"); + assert_eq!(event.field("network.policy.scope"), Some("non_domain")); + assert_eq!(event.field("network.policy.decision"), Some("deny")); + assert_eq!(event.field("network.policy.source"), Some("mode_guard")); + assert_eq!( + event.field("network.policy.reason"), + Some(REASON_METHOD_NOT_ALLOWED) + ); + assert_eq!( + event.field("network.transport.protocol"), + Some("socks5_udp") + ); + assert_eq!(event.field("server.address"), Some("93.184.216.34")); + assert_eq!(event.field("server.port"), Some("53")); + assert_eq!(event.field("http.request.method"), Some("none")); + assert_eq!(event.field("client.address"), Some("unknown")); + } +} diff --git a/codex-rs/network-proxy/src/state.rs b/codex-rs/network-proxy/src/state.rs new file mode 100644 index 0000000000000000000000000000000000000000..90731a8b4e8103fb2f8e1e60413539949085aac8 --- /dev/null +++ b/codex-rs/network-proxy/src/state.rs @@ -0,0 +1,451 @@ +use crate::config::NetworkDomainPermissions; +use crate::config::NetworkMode; +use crate::config::NetworkProxyConfig; +use crate::config::NetworkUnixSocketPermissions; +use crate::mitm::MitmState; +use crate::mitm::MitmUpstreamConfig; +use crate::mitm_hook::MitmHookConfig; +use crate::mitm_hook::compile_mitm_hooks; +use crate::mitm_hook::validate_mitm_hook_config; +use crate::policy::DomainPattern; +use crate::policy::compile_allowlist_globset; +use crate::policy::compile_denylist_globset; +use crate::policy::is_global_wildcard_domain_pattern; +use crate::runtime::ConfigState; +use serde::Deserialize; +use std::collections::HashSet; +use std::sync::Arc; + +pub use crate::runtime::BlockedRequest; +pub use crate::runtime::BlockedRequestArgs; +pub use crate::runtime::NetworkProxyAuditMetadata; +pub use crate::runtime::NetworkProxyState; +#[cfg(test)] +pub(crate) use crate::runtime::network_proxy_state_for_policy; + +#[derive(Debug, Default, Clone, PartialEq, Eq)] +pub struct NetworkProxyConstraints { + pub enabled: Option, + pub mode: Option, + pub allow_upstream_proxy: Option, + pub dangerously_allow_non_loopback_proxy: Option, + pub dangerously_allow_all_unix_sockets: Option, + pub allowed_domains: Option>, + pub allowlist_expansion_enabled: Option, + pub denied_domains: Option>, + pub denylist_expansion_enabled: Option, + pub allow_unix_sockets: Option>, + pub allow_local_binding: Option, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct PartialNetworkProxyConfig { + pub enabled: Option, + pub mode: Option, + pub allow_upstream_proxy: Option, + pub dangerously_allow_non_loopback_proxy: Option, + pub dangerously_allow_all_unix_sockets: Option, + #[serde(default)] + pub domains: Option, + #[serde(default)] + pub unix_sockets: Option, + pub allow_local_binding: Option, + pub mitm: Option, + pub credential_broker: Option, + pub dangerously_allow_plaintext_credential_injection: Option, + #[serde(default)] + pub mitm_hooks: Option>, +} + +pub fn build_config_state( + mut config: NetworkProxyConfig, + constraints: NetworkProxyConstraints, +) -> anyhow::Result { + if constraints.enabled == Some(false) { + config.credential_broker = false; + } + let brokerage_created_proxy = config.credential_broker && !config.enabled; + let brokerage_created_default_allowlist = brokerage_created_proxy + && config.allowed_domains().is_none() + && constraints.allowed_domains.is_none(); + if brokerage_created_proxy { + config.enabled = true; + } + if brokerage_created_default_allowlist { + config.set_allowed_domains(vec!["*".to_string()]); + } + crate::config::validate_unix_socket_allowlist_paths(&config)?; + anyhow::ensure!( + !config.credential_broker || config.mitm, + "network.credential_broker requires network.mitm = true" + ); + let allowed_domains = config.allowed_domains().unwrap_or_default(); + let denied_domains = config.denied_domains().unwrap_or_default(); + validate_non_global_wildcard_domain_patterns("network.denied_domains", &denied_domains) + .map_err(NetworkProxyConstraintError::into_anyhow)?; + let deny_set = compile_denylist_globset(&denied_domains)?; + let allow_set = compile_allowlist_globset(&allowed_domains)?; + let mitm_hooks = compile_mitm_hooks(&config)?; + let mitm = if config.mitm { + Some(Arc::new(MitmState::new(MitmUpstreamConfig { + allow_upstream_proxy: config.allow_upstream_proxy, + })?)) + } else { + None + }; + Ok(ConfigState { + config, + brokerage_created_default_allowlist, + allow_set, + deny_set, + mitm, + mitm_hooks, + constraints, + blocked: std::collections::VecDeque::new(), + blocked_total: 0, + }) +} + +pub fn validate_policy_against_constraints( + config: &NetworkProxyConfig, + constraints: &NetworkProxyConstraints, +) -> Result<(), NetworkProxyConstraintError> { + fn invalid_value( + field_name: &'static str, + candidate: impl Into, + allowed: impl Into, + ) -> NetworkProxyConstraintError { + NetworkProxyConstraintError::InvalidValue { + field_name, + candidate: candidate.into(), + allowed: allowed.into(), + } + } + + fn validate( + candidate: T, + validator: impl FnOnce(&T) -> Result<(), NetworkProxyConstraintError>, + ) -> Result<(), NetworkProxyConstraintError> { + validator(&candidate) + } + + let enabled = config.enabled; + let config_allowed_domains = config.allowed_domains().unwrap_or_default(); + let config_denied_domains = config.denied_domains().unwrap_or_default(); + let denied_domain_overrides: HashSet = config_denied_domains + .iter() + .map(|entry| entry.to_ascii_lowercase()) + .collect(); + let config_allow_unix_sockets = config.allow_unix_sockets(); + validate_mitm_hook_config(config).map_err(invalid_mitm_hook_configuration)?; + validate_non_global_wildcard_domain_patterns("network.denied_domains", &config_denied_domains)?; + if let Some(max_enabled) = constraints.enabled { + validate(enabled, move |candidate| { + if *candidate && !max_enabled { + Err(invalid_value( + "network.enabled", + "true", + "false (disabled by managed config)", + )) + } else { + Ok(()) + } + })?; + } + + if let Some(max_mode) = constraints.mode { + validate(config.mode, move |candidate| { + if network_mode_rank(*candidate) > network_mode_rank(max_mode) { + Err(invalid_value( + "network.mode", + format!("{candidate:?}"), + format!("{max_mode:?} or more restrictive"), + )) + } else { + Ok(()) + } + })?; + } + + let allow_upstream_proxy = constraints.allow_upstream_proxy; + validate( + config.allow_upstream_proxy, + move |candidate| match allow_upstream_proxy { + Some(true) | None => Ok(()), + Some(false) => { + if *candidate { + Err(invalid_value( + "network.allow_upstream_proxy", + "true", + "false (disabled by managed config)", + )) + } else { + Ok(()) + } + } + }, + )?; + + let allow_non_loopback_proxy = constraints.dangerously_allow_non_loopback_proxy; + validate( + config.dangerously_allow_non_loopback_proxy, + move |candidate| match allow_non_loopback_proxy { + Some(true) | None => Ok(()), + Some(false) => { + if *candidate { + Err(invalid_value( + "network.dangerously_allow_non_loopback_proxy", + "true", + "false (disabled by managed config)", + )) + } else { + Ok(()) + } + } + }, + )?; + + let allow_all_unix_sockets = constraints + .dangerously_allow_all_unix_sockets + .unwrap_or(constraints.allow_unix_sockets.is_none()); + validate( + config.dangerously_allow_all_unix_sockets, + move |candidate| { + if *candidate && !allow_all_unix_sockets { + Err(invalid_value( + "network.dangerously_allow_all_unix_sockets", + "true", + "false (disabled by managed config)", + )) + } else { + Ok(()) + } + }, + )?; + + if let Some(allow_local_binding) = constraints.allow_local_binding { + validate(config.allow_local_binding, move |candidate| { + if *candidate && !allow_local_binding { + Err(invalid_value( + "network.allow_local_binding", + "true", + "false (disabled by managed config)", + )) + } else { + Ok(()) + } + })?; + } + + if let Some(allowed_domains) = &constraints.allowed_domains { + validate_non_global_wildcard_domain_patterns("network.allowed_domains", allowed_domains)?; + match constraints.allowlist_expansion_enabled { + Some(true) => { + let required_set: HashSet = allowed_domains + .iter() + .map(|entry| entry.to_ascii_lowercase()) + .collect(); + validate(config_allowed_domains, |candidate| { + let candidate_set: HashSet = candidate + .iter() + .map(|entry| entry.to_ascii_lowercase()) + .collect(); + let missing: Vec = required_set + .iter() + .filter(|entry| { + !candidate_set.contains(*entry) + && !denied_domain_overrides.contains(*entry) + }) + .cloned() + .collect(); + if missing.is_empty() { + Ok(()) + } else { + Err(invalid_value( + "network.allowed_domains", + "missing managed allowed_domains entries", + format!("{missing:?}"), + )) + } + })?; + } + Some(false) => { + let required_set: HashSet = allowed_domains + .iter() + .map(|entry| entry.to_ascii_lowercase()) + .collect(); + validate(config_allowed_domains, |candidate| { + let candidate_set: HashSet = candidate + .iter() + .map(|entry| entry.to_ascii_lowercase()) + .collect(); + let expected_set: HashSet = required_set + .difference(&denied_domain_overrides) + .cloned() + .collect(); + if candidate_set == expected_set { + Ok(()) + } else { + Err(invalid_value( + "network.allowed_domains", + format!("{candidate:?}"), + "must match managed allowed_domains", + )) + } + })?; + } + None => { + let managed_patterns: Vec = allowed_domains + .iter() + .map(|entry| DomainPattern::parse_for_constraints(entry)) + .collect(); + validate(config_allowed_domains, move |candidate| { + let mut invalid = Vec::new(); + for entry in candidate { + let candidate_pattern = DomainPattern::parse_for_constraints(entry); + if !managed_patterns + .iter() + .any(|managed| managed.allows(&candidate_pattern)) + { + invalid.push(entry.clone()); + } + } + if invalid.is_empty() { + Ok(()) + } else { + Err(invalid_value( + "network.allowed_domains", + format!("{invalid:?}"), + "subset of managed allowed_domains", + )) + } + })?; + } + } + } + + if let Some(denied_domains) = &constraints.denied_domains { + validate_non_global_wildcard_domain_patterns("network.denied_domains", denied_domains)?; + let required_set: HashSet = denied_domains + .iter() + .map(|s| s.to_ascii_lowercase()) + .collect(); + match constraints.denylist_expansion_enabled { + Some(false) => { + validate(config_denied_domains, move |candidate| { + let candidate_set: HashSet = candidate + .iter() + .map(|entry| entry.to_ascii_lowercase()) + .collect(); + if candidate_set == required_set { + Ok(()) + } else { + Err(invalid_value( + "network.denied_domains", + format!("{candidate:?}"), + "must match managed denied_domains", + )) + } + })?; + } + Some(true) | None => { + validate(config_denied_domains, move |candidate| { + let candidate_set: HashSet = + candidate.iter().map(|s| s.to_ascii_lowercase()).collect(); + let missing: Vec = required_set + .iter() + .filter(|entry| !candidate_set.contains(*entry)) + .cloned() + .collect(); + if missing.is_empty() { + Ok(()) + } else { + Err(invalid_value( + "network.denied_domains", + "missing managed denied_domains entries", + format!("{missing:?}"), + )) + } + })?; + } + } + } + + if let Some(allow_unix_sockets) = &constraints.allow_unix_sockets { + let allowed_set: HashSet = allow_unix_sockets + .iter() + .map(|s| s.to_ascii_lowercase()) + .collect(); + validate(config_allow_unix_sockets, move |candidate| { + let mut invalid = Vec::new(); + for entry in candidate { + if !allowed_set.contains(&entry.to_ascii_lowercase()) { + invalid.push(entry.clone()); + } + } + if invalid.is_empty() { + Ok(()) + } else { + Err(invalid_value( + "network.allow_unix_sockets", + format!("{invalid:?}"), + "subset of managed allow_unix_sockets", + )) + } + })?; + } + + Ok(()) +} + +fn invalid_mitm_hook_configuration(err: anyhow::Error) -> NetworkProxyConstraintError { + NetworkProxyConstraintError::InvalidValue { + field_name: "network.mitm_hooks", + candidate: err.to_string(), + allowed: "valid MITM hook configuration".to_string(), + } +} + +fn validate_non_global_wildcard_domain_patterns( + field_name: &'static str, + patterns: &[String], +) -> Result<(), NetworkProxyConstraintError> { + if let Some(pattern) = patterns + .iter() + .find(|pattern| is_global_wildcard_domain_pattern(pattern)) + { + return Err(NetworkProxyConstraintError::InvalidValue { + field_name, + candidate: pattern.trim().to_string(), + allowed: "exact hosts or scoped wildcards like *.example.com or **.example.com" + .to_string(), + }); + } + Ok(()) +} + +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +pub enum NetworkProxyConstraintError { + #[error("invalid value for {field_name}: {candidate} (allowed {allowed})")] + InvalidValue { + field_name: &'static str, + candidate: String, + allowed: String, + }, +} + +impl NetworkProxyConstraintError { + pub fn into_anyhow(self) -> anyhow::Error { + anyhow::anyhow!(self) + } +} + +fn network_mode_rank(mode: NetworkMode) -> u8 { + match mode { + NetworkMode::Limited => 0, + NetworkMode::Full => 1, + } +} + +#[cfg(test)] +mod tests {} diff --git a/codex-rs/network-proxy/src/upstream.rs b/codex-rs/network-proxy/src/upstream.rs new file mode 100644 index 0000000000000000000000000000000000000000..f523327427cce06df99c1ac933709f89bc322f71 --- /dev/null +++ b/codex-rs/network-proxy/src/upstream.rs @@ -0,0 +1,287 @@ +use crate::connect_policy::TargetCheckedTcpConnector; +use crate::connect_policy::is_non_public_target; +use crate::state::NetworkProxyState; +use codex_utils_rustls_provider::ensure_rustls_crypto_provider; +use rama_core::Layer; +use rama_core::Service; +use rama_core::error::BoxError; +use rama_core::error::ErrorExt as _; +use rama_core::error::OpaqueError; +use rama_core::extensions::ExtensionsMut; +use rama_core::extensions::ExtensionsRef; +use rama_core::service::BoxService; +use rama_http::Body; +use rama_http::Request; +use rama_http::Response; +use rama_http::layer::version_adapter::RequestVersionAdapter; +use rama_http_backend::client::HttpClientService; +use rama_http_backend::client::HttpConnector; +use rama_http_backend::client::proxy::layer::HttpProxyConnectorLayer; +use rama_net::address::HostWithPort; +use rama_net::address::ProxyAddress; +use rama_net::client::EstablishedClientConnection; +use rama_net::http::RequestContext; +use rama_tls_rustls::client::TlsConnectorDataBuilder; +use rama_tls_rustls::client::TlsConnectorLayer; +use rama_tls_rustls::client::client_root_certs; +use rama_tls_rustls::dep::rustls; +use std::sync::Arc; +use std::time::Instant; +use tracing::info; +use tracing::warn; + +#[cfg(target_os = "macos")] +use rama_unix::client::UnixConnector; + +#[derive(Clone, Default)] +struct ProxyConfig { + http: Option, + https: Option, + all: Option, +} + +impl ProxyConfig { + fn from_env() -> Self { + let http = read_proxy_env(&["HTTP_PROXY", "http_proxy"]); + let https = read_proxy_env(&["HTTPS_PROXY", "https_proxy"]); + let all = read_proxy_env(&["ALL_PROXY", "all_proxy"]); + Self { http, https, all } + } + + fn proxy_for_protocol(&self, is_secure: bool) -> Option { + if is_secure { + self.https + .clone() + .or_else(|| self.http.clone()) + .or_else(|| self.all.clone()) + } else { + self.http.clone().or_else(|| self.all.clone()) + } + } + + fn proxy_for_target(&self, target: &HostWithPort, is_secure: bool) -> Option { + if is_non_public_target(&target.host) { + return None; + } + self.proxy_for_protocol(is_secure) + } +} + +fn read_proxy_env(keys: &[&str]) -> Option { + for key in keys { + let Ok(value) = std::env::var(key) else { + continue; + }; + let value = value.trim(); + if value.is_empty() { + continue; + } + match ProxyAddress::try_from(value) { + Ok(proxy) => { + if proxy + .protocol + .as_ref() + .map(rama_net::Protocol::is_http) + .unwrap_or(true) + { + return Some(proxy); + } + warn!("ignoring {key}: non-http proxy protocol"); + } + Err(err) => { + warn!("ignoring {key}: invalid proxy address ({err})"); + } + } + } + None +} + +pub(crate) fn proxy_for_connect(target: &HostWithPort) -> Option { + ProxyConfig::from_env().proxy_for_target(target, /*is_secure*/ true) +} + +#[derive(Clone)] +pub(crate) struct UpstreamClient { + connector: BoxService< + Request, + EstablishedClientConnection, Request>, + BoxError, + >, + proxy_config: ProxyConfig, +} + +impl UpstreamClient { + pub(crate) fn direct(state: Arc) -> Self { + Self::new( + ProxyConfig::default(), + TargetCheckedTcpConnector::new(state), + client_root_certs(), + ) + } + + pub(crate) fn from_env_proxy(state: Arc) -> Self { + Self::new( + ProxyConfig::from_env(), + TargetCheckedTcpConnector::new(state), + client_root_certs(), + ) + } + + pub(crate) fn direct_with_tls_root_store( + state: Arc, + tls_root_store: Arc, + ) -> Self { + Self::new( + ProxyConfig::default(), + TargetCheckedTcpConnector::new(state), + tls_root_store, + ) + } + + pub(crate) fn from_env_proxy_with_tls_root_store( + state: Arc, + tls_root_store: Arc, + ) -> Self { + Self::new( + ProxyConfig::from_env(), + TargetCheckedTcpConnector::new(state), + tls_root_store, + ) + } + + #[cfg(target_os = "macos")] + pub(crate) fn unix_socket(path: &str) -> Self { + let connector = build_unix_connector(path); + Self { + connector, + proxy_config: ProxyConfig::default(), + } + } + + fn new( + proxy_config: ProxyConfig, + transport: TargetCheckedTcpConnector, + tls_root_store: Arc, + ) -> Self { + let connector = build_http_connector(transport, tls_root_store); + Self { + connector, + proxy_config, + } + } +} + +impl Service> for UpstreamClient { + type Output = Response; + type Error = OpaqueError; + + async fn serve(&self, mut req: Request) -> Result { + let request_context = RequestContext::try_from(&req).ok(); + let authority = request_context + .as_ref() + .map(|ctx| ctx.host_with_port().to_string()) + .unwrap_or_else(|| "".to_string()); + let proxy = request_context.as_ref().map_or_else( + || self.proxy_config.proxy_for_protocol(/*is_secure*/ false), + |ctx| { + self.proxy_config + .proxy_for_target(&ctx.host_with_port(), ctx.protocol.is_secure()) + }, + ); + match proxy.as_ref() { + Some(proxy) => info!( + "HTTP upstream route selected (target={authority}, route=upstream_proxy, proxy={})", + proxy.address + ), + None => info!("HTTP upstream route selected (target={authority}, route=direct)"), + } + if let Some(proxy) = proxy { + req.extensions_mut().insert(proxy); + } + + let uri = req.uri().clone(); + let connect_started_at = Instant::now(); + let EstablishedClientConnection { + input: mut req, + conn: http_connection, + } = match self.connector.serve(req).await { + Ok(connection) => { + info!( + "HTTP upstream connection established (target={authority}, elapsed_ms={})", + connect_started_at.elapsed().as_millis() + ); + connection + } + Err(err) => { + warn!( + "HTTP upstream connection failed (target={authority}, elapsed_ms={})", + connect_started_at.elapsed().as_millis() + ); + return Err(OpaqueError::from_boxed(err)); + } + }; + + req.extensions_mut() + .extend(http_connection.extensions().clone()); + + let request_started_at = Instant::now(); + match http_connection.serve(req).await { + Ok(resp) => { + info!( + "HTTP upstream response headers received (target={authority}, elapsed_ms={})", + request_started_at.elapsed().as_millis() + ); + Ok(resp) + } + Err(err) => { + warn!( + "HTTP upstream response headers failed (target={authority}, elapsed_ms={})", + request_started_at.elapsed().as_millis() + ); + Err(OpaqueError::from_boxed(err) + .context(format!("http request failure for uri: {uri}"))) + } + } + } +} + +fn build_http_connector( + transport: TargetCheckedTcpConnector, + tls_root_store: Arc, +) -> BoxService< + Request, + EstablishedClientConnection, Request>, + BoxError, +> { + ensure_rustls_crypto_provider(); + let proxy = HttpProxyConnectorLayer::optional().into_layer(transport); + let client_config = rustls::ClientConfig::builder_with_protocol_versions(rustls::ALL_VERSIONS) + .with_root_certificates(tls_root_store) + .with_no_client_auth(); + let tls_config = TlsConnectorDataBuilder::from(client_config) + .with_alpn_protocols_http_auto() + .build(); + let tls = TlsConnectorLayer::auto() + .with_connector_data(tls_config) + .into_layer(proxy); + let tls = RequestVersionAdapter::new(tls); + let connector = HttpConnector::new(tls); + connector.boxed() +} + +#[cfg(test)] +#[path = "upstream_tests.rs"] +mod tests; + +#[cfg(target_os = "macos")] +fn build_unix_connector( + path: &str, +) -> BoxService< + Request, + EstablishedClientConnection, Request>, + BoxError, +> { + let transport = UnixConnector::fixed(path); + let connector = HttpConnector::new(transport); + connector.boxed() +} diff --git a/codex-rs/network-proxy/src/windows_tcp_attribution.rs b/codex-rs/network-proxy/src/windows_tcp_attribution.rs new file mode 100644 index 0000000000000000000000000000000000000000..5146d5e6cec61bccfc3ffa0405941f304073ea1c --- /dev/null +++ b/codex-rs/network-proxy/src/windows_tcp_attribution.rs @@ -0,0 +1,315 @@ +use std::ffi::c_void; +use std::io; +use std::mem::offset_of; +use std::mem::size_of; +use std::net::Ipv4Addr; +use std::net::SocketAddr; +use std::net::SocketAddrV4; +use std::os::windows::io::AsRawHandle; +use std::os::windows::io::FromRawHandle; +use std::os::windows::io::OwnedHandle; +use std::os::windows::io::RawHandle; + +use windows_sys::Win32::Foundation::ERROR_INSUFFICIENT_BUFFER; +use windows_sys::Win32::Foundation::GetLastError; +use windows_sys::Win32::Foundation::HANDLE; +use windows_sys::Win32::Foundation::HLOCAL; +use windows_sys::Win32::Foundation::LocalFree; +use windows_sys::Win32::Foundation::NO_ERROR; +use windows_sys::Win32::Foundation::PSID; +use windows_sys::Win32::NetworkManagement::IpHelper::GetExtendedTcpTable; +use windows_sys::Win32::NetworkManagement::IpHelper::MIB_TCPROW_OWNER_PID; +use windows_sys::Win32::NetworkManagement::IpHelper::MIB_TCPTABLE_OWNER_PID; +use windows_sys::Win32::NetworkManagement::IpHelper::TCP_TABLE_OWNER_PID_CONNECTIONS; +use windows_sys::Win32::Networking::WinSock::AF_INET; +use windows_sys::Win32::Security::Authorization::ConvertSidToStringSidW; +use windows_sys::Win32::Security::GetTokenInformation; +use windows_sys::Win32::Security::SID_AND_ATTRIBUTES; +use windows_sys::Win32::Security::TOKEN_GROUPS; +use windows_sys::Win32::Security::TOKEN_QUERY; +use windows_sys::Win32::Security::TokenRestrictedSids; +use windows_sys::Win32::System::Threading::OpenProcess; +use windows_sys::Win32::System::Threading::OpenProcessToken; +use windows_sys::Win32::System::Threading::PROCESS_QUERY_LIMITED_INFORMATION; + +/// Returns the restricting SIDs on the process that opened an accepted loopback connection. +/// +/// `accepted_local_addr` and `accepted_peer_addr` must come from the accepted server socket. The +/// owning-PID table describes the client side in the opposite direction, so the lookup matches the +/// exact reversed four-tuple. +pub(crate) fn restricting_sids_for_tcp_connection( + accepted_local_addr: SocketAddr, + accepted_peer_addr: SocketAddr, +) -> io::Result> { + let (SocketAddr::V4(accepted_local_addr), SocketAddr::V4(accepted_peer_addr)) = + (accepted_local_addr, accepted_peer_addr) + else { + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "Windows proxy connection attribution currently supports IPv4 only", + )); + }; + + let process_id = owning_process_id(accepted_local_addr, accepted_peer_addr)?; + restricting_sids_for_process(process_id) +} + +fn owning_process_id( + accepted_local_addr: SocketAddrV4, + accepted_peer_addr: SocketAddrV4, +) -> io::Result { + let mut byte_len = 0_u32; + let result = unsafe { + GetExtendedTcpTable( + std::ptr::null_mut(), + &mut byte_len, + 0, + AF_INET as u32, + TCP_TABLE_OWNER_PID_CONNECTIONS, + 0, + ) + }; + if result != ERROR_INSUFFICIENT_BUFFER { + return Err(win32_error("query IPv4 TCP owner table size", result)); + } + + let buffer = loop { + let mut buffer = aligned_buffer(byte_len as usize)?; + let result = unsafe { + GetExtendedTcpTable( + buffer.as_mut_ptr().cast::(), + &mut byte_len, + 0, + AF_INET as u32, + TCP_TABLE_OWNER_PID_CONNECTIONS, + 0, + ) + }; + match result { + NO_ERROR => break buffer, + ERROR_INSUFFICIENT_BUFFER => continue, + _ => return Err(win32_error("read IPv4 TCP owner table", result)), + } + }; + + let rows = parse_tcp_owner_rows(&buffer, byte_len as usize)?; + unique_client_process_id(rows, accepted_local_addr, accepted_peer_addr) +} + +fn parse_tcp_owner_rows(buffer: &[usize], byte_len: usize) -> io::Result<&[MIB_TCPROW_OWNER_PID]> { + let rows_offset = offset_of!(MIB_TCPTABLE_OWNER_PID, table); + if byte_len > size_of_val(buffer) || byte_len < rows_offset { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "invalid IPv4 TCP owner table length", + )); + } + + let row_count = unsafe { std::ptr::read_unaligned(buffer.as_ptr().cast::()) } as usize; + let rows_byte_len = row_count + .checked_mul(size_of::()) + .and_then(|len| rows_offset.checked_add(len)) + .ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidData, + "IPv4 TCP owner table length overflow", + ) + })?; + if rows_byte_len > byte_len { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "truncated IPv4 TCP owner table", + )); + } + + let rows = unsafe { + let rows_ptr = buffer + .as_ptr() + .cast::() + .add(rows_offset) + .cast::(); + std::slice::from_raw_parts(rows_ptr, row_count) + }; + Ok(rows) +} + +fn unique_client_process_id( + rows: &[MIB_TCPROW_OWNER_PID], + accepted_local_addr: SocketAddrV4, + accepted_peer_addr: SocketAddrV4, +) -> io::Result { + let mut matching_process_ids = rows + .iter() + .filter(|row| client_row_matches(row, accepted_local_addr, accepted_peer_addr)) + .map(|row| row.dwOwningPid); + let process_id = matching_process_ids.next().ok_or_else(|| { + io::Error::new( + io::ErrorKind::NotFound, + "accepted connection is absent from the IPv4 TCP owner table", + ) + })?; + if matching_process_ids.next().is_some() { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "accepted connection has multiple IPv4 TCP owner rows", + )); + } + Ok(process_id) +} + +fn client_row_matches( + row: &MIB_TCPROW_OWNER_PID, + accepted_local_addr: SocketAddrV4, + accepted_peer_addr: SocketAddrV4, +) -> bool { + ipv4_addr_matches(row.dwLocalAddr, *accepted_peer_addr.ip()) + && tcp_port(row.dwLocalPort) == accepted_peer_addr.port() + && ipv4_addr_matches(row.dwRemoteAddr, *accepted_local_addr.ip()) + && tcp_port(row.dwRemotePort) == accepted_local_addr.port() +} + +fn ipv4_addr_matches(table_addr: u32, socket_addr: Ipv4Addr) -> bool { + table_addr.to_ne_bytes() == socket_addr.octets() +} + +fn tcp_port(table_port: u32) -> u16 { + u16::from_be(table_port as u16) +} + +fn restricting_sids_for_process(process_id: u32) -> io::Result> { + let process_handle = unsafe { OpenProcess(PROCESS_QUERY_LIMITED_INFORMATION, 0, process_id) }; + let process = owned_handle(process_handle, "open proxy client process")?; + + let mut token_handle: HANDLE = 0; + let opened = unsafe { + OpenProcessToken( + process.as_raw_handle() as HANDLE, + TOKEN_QUERY, + &mut token_handle, + ) + }; + if opened == 0 { + return Err(last_error("open proxy client process token")); + } + let token = owned_handle(token_handle, "open proxy client process token")?; + + let mut byte_len = 0_u32; + let queried = unsafe { + GetTokenInformation( + token.as_raw_handle() as HANDLE, + TokenRestrictedSids, + std::ptr::null_mut(), + 0, + &mut byte_len, + ) + }; + if queried != 0 || unsafe { GetLastError() } != ERROR_INSUFFICIENT_BUFFER { + return Err(last_error("query proxy client restricting SID buffer size")); + } + + let mut buffer = aligned_buffer(byte_len as usize)?; + let queried = unsafe { + GetTokenInformation( + token.as_raw_handle() as HANDLE, + TokenRestrictedSids, + buffer.as_mut_ptr().cast::(), + byte_len, + &mut byte_len, + ) + }; + if queried == 0 { + return Err(last_error("read proxy client restricting SIDs")); + } + + parse_token_groups(&buffer, byte_len as usize)? + .iter() + .map(|entry| sid_to_string(entry.Sid)) + .collect() +} + +fn parse_token_groups(buffer: &[usize], byte_len: usize) -> io::Result<&[SID_AND_ATTRIBUTES]> { + let groups_offset = offset_of!(TOKEN_GROUPS, Groups); + if byte_len > size_of_val(buffer) || byte_len < groups_offset { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "invalid restricting SID buffer length", + )); + } + + let group_count = unsafe { std::ptr::read_unaligned(buffer.as_ptr().cast::()) } as usize; + let groups_byte_len = group_count + .checked_mul(size_of::()) + .and_then(|len| groups_offset.checked_add(len)) + .ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidData, + "restricting SID buffer length overflow", + ) + })?; + if groups_byte_len > byte_len { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "truncated restricting SID buffer", + )); + } + + let groups = unsafe { + let groups_ptr = buffer + .as_ptr() + .cast::() + .add(groups_offset) + .cast::(); + std::slice::from_raw_parts(groups_ptr, group_count) + }; + Ok(groups) +} + +fn sid_to_string(sid: PSID) -> io::Result { + let mut string_sid = std::ptr::null_mut(); + if unsafe { ConvertSidToStringSidW(sid, &mut string_sid) } == 0 { + return Err(last_error("convert proxy client restricting SID to string")); + } + + let value = unsafe { + let mut len = 0; + while *string_sid.add(len) != 0 { + len += 1; + } + String::from_utf16_lossy(std::slice::from_raw_parts(string_sid, len)) + }; + unsafe { + LocalFree(string_sid as HLOCAL); + } + Ok(value) +} + +fn aligned_buffer(byte_len: usize) -> io::Result> { + if byte_len == 0 { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "Windows API returned an empty buffer length", + )); + } + Ok(vec![0; byte_len.div_ceil(size_of::())]) +} + +fn owned_handle(handle: HANDLE, operation: &str) -> io::Result { + if handle == 0 { + return Err(last_error(operation)); + } + Ok(unsafe { OwnedHandle::from_raw_handle(handle as RawHandle) }) +} + +fn win32_error(operation: &str, error_code: u32) -> io::Error { + let error = io::Error::from_raw_os_error(error_code as i32); + io::Error::new(error.kind(), format!("{operation}: {error}")) +} + +fn last_error(operation: &str) -> io::Error { + let error = io::Error::last_os_error(); + io::Error::new(error.kind(), format!("{operation}: {error}")) +} + +#[cfg(test)] +#[path = "windows_tcp_attribution_tests.rs"] +mod tests; diff --git a/codex-rs/network-proxy/src/windows_tcp_attribution_tests.rs b/codex-rs/network-proxy/src/windows_tcp_attribution_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..ebbc4b8109701c92899ba468b46d25853696bee2 --- /dev/null +++ b/codex-rs/network-proxy/src/windows_tcp_attribution_tests.rs @@ -0,0 +1,112 @@ +use super::*; +use pretty_assertions::assert_eq; +use std::net::TcpListener; +use std::net::TcpStream; + +#[test] +fn parses_owner_table_and_matches_reversed_client_tuple() -> io::Result<()> { + let proxy_addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 3128); + let client_addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152); + let rows = [ + tcp_row(proxy_addr, client_addr, 100), + tcp_row(client_addr, proxy_addr, 200), + ]; + let (buffer, byte_len) = owner_table_buffer(&rows); + + let parsed = parse_tcp_owner_rows(&buffer, byte_len)?; + + assert_eq!( + unique_client_process_id(parsed, proxy_addr, client_addr)?, + 200 + ); + Ok(()) +} + +#[test] +fn rejects_multiple_matching_owner_rows() { + let proxy_addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 3128); + let client_addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152); + let rows = [ + tcp_row(client_addr, proxy_addr, 200), + tcp_row(client_addr, proxy_addr, 201), + ]; + + let error = unique_client_process_id(&rows, proxy_addr, client_addr) + .expect_err("duplicate connection rows should fail closed"); + + assert_eq!(error.kind(), io::ErrorKind::InvalidData); +} + +#[test] +fn rejects_truncated_owner_table() { + let byte_len = offset_of!(MIB_TCPTABLE_OWNER_PID, table); + let mut buffer = aligned_buffer(byte_len).expect("aligned table buffer"); + unsafe { + std::ptr::write_unaligned(buffer.as_mut_ptr().cast::(), 1); + } + + let Err(error) = parse_tcp_owner_rows(&buffer, byte_len) else { + panic!("truncated connection row should fail closed"); + }; + + assert_eq!(error.kind(), io::ErrorKind::InvalidData); +} + +#[test] +fn resolves_loopback_connection_to_current_process() -> io::Result<()> { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0))?; + let client = TcpStream::connect(listener.local_addr()?)?; + let (accepted, _) = listener.accept()?; + let local_addr = accepted.local_addr()?; + let peer_addr = accepted.peer_addr()?; + + let process_id = owning_process_id(socket_addr_v4(local_addr)?, socket_addr_v4(peer_addr)?)?; + let restricting_sids = restricting_sids_for_tcp_connection(local_addr, peer_addr)?; + + assert_eq!(process_id, std::process::id()); + assert!(restricting_sids.iter().all(|sid| sid.starts_with("S-"))); + drop(client); + Ok(()) +} + +fn socket_addr_v4(addr: SocketAddr) -> io::Result { + match addr { + SocketAddr::V4(addr) => Ok(addr), + SocketAddr::V6(_) => Err(io::Error::new( + io::ErrorKind::InvalidData, + "test listener unexpectedly used IPv6", + )), + } +} + +fn tcp_row( + local_addr: SocketAddrV4, + remote_addr: SocketAddrV4, + process_id: u32, +) -> MIB_TCPROW_OWNER_PID { + MIB_TCPROW_OWNER_PID { + dwState: 0, + dwLocalAddr: u32::from_ne_bytes(local_addr.ip().octets()), + dwLocalPort: local_addr.port().to_be() as u32, + dwRemoteAddr: u32::from_ne_bytes(remote_addr.ip().octets()), + dwRemotePort: remote_addr.port().to_be() as u32, + dwOwningPid: process_id, + } +} + +fn owner_table_buffer(rows: &[MIB_TCPROW_OWNER_PID]) -> (Vec, usize) { + let rows_offset = offset_of!(MIB_TCPTABLE_OWNER_PID, table); + let rows_byte_len = size_of_val(rows); + let byte_len = rows_offset + rows_byte_len; + let mut buffer = aligned_buffer(byte_len).expect("aligned table buffer"); + unsafe { + let buffer_ptr = buffer.as_mut_ptr().cast::(); + std::ptr::write_unaligned(buffer_ptr.cast::(), rows.len() as u32); + std::ptr::copy_nonoverlapping( + rows.as_ptr().cast::(), + buffer_ptr.add(rows_offset), + rows_byte_len, + ); + } + (buffer, byte_len) +} diff --git a/codex-rs/rmcp-client/src/auth_status.rs b/codex-rs/rmcp-client/src/auth_status.rs new file mode 100644 index 0000000000000000000000000000000000000000..479ae184f4e8cde560dad75d33c0837d63b1cb55 --- /dev/null +++ b/codex-rs/rmcp-client/src/auth_status.rs @@ -0,0 +1,935 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Result; +use codex_exec_server::HttpClient; +use codex_protocol::protocol::McpAuthStatus; +use futures::FutureExt; +use http::HeaderMap; +use http::header::AUTHORIZATION; +use rmcp::transport::AuthorizationManager; +use rmcp::transport::auth::AuthError; +use tracing::debug; + +use crate::http_client_adapter::StreamableHttpRedirectMode; +use crate::oauth::StoredOAuthTokenStatus; +use crate::oauth::oauth_token_status; +use crate::oauth_callback::McpOAuthCallbackMode; +use crate::oauth_callback::callback_mode; +use crate::oauth_http_client::OAuthHttpClientAdapter; +use crate::utils::build_default_headers; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; + +const DISCOVERY_TIMEOUT: Duration = Duration::from_secs(5); + +/// Timeout policy for OAuth metadata discovery through a supplied HTTP client. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum OAuthDiscoveryTimeout { + /// Preserve the timeout requested by the OAuth implementation. + Requested, + /// Cap OAuth discovery requests at the supplied duration. + Capped(Duration), +} + +impl OAuthDiscoveryTimeout { + /// Preserves the existing timeout for local OAuth discovery. + pub const LOCAL: Self = Self::Capped(DISCOVERY_TIMEOUT); +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct StreamableHttpOAuthDiscovery { + pub scopes_supported: Option>, + pub callback_mode: McpOAuthCallbackMode, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum McpLoginRequirement { + Login, + Reauthentication, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum McpAuthState { + Unsupported, + Unknown, + LoggedOut(McpLoginRequirement), + BearerToken, + OAuth, +} + +impl From for McpAuthStatus { + fn from(value: McpAuthState) -> Self { + match value { + McpAuthState::Unsupported => Self::Unsupported, + McpAuthState::Unknown => Self::Unknown, + McpAuthState::LoggedOut(_) => Self::NotLoggedIn, + McpAuthState::BearerToken => Self::BearerToken, + McpAuthState::OAuth => Self::OAuth, + } + } +} + +enum AuthStatusCheck { + Complete(McpAuthState), + Discover(HeaderMap), +} + +/// Determine authentication status while routing OAuth discovery through the +/// provided HTTP client. +#[allow(clippy::too_many_arguments)] +pub async fn determine_streamable_http_auth_status( + server_name: &str, + url: &str, + bearer_token_env_var: Option<&str>, + http_headers: Option>, + env_http_headers: Option>, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_client: Arc, + discovery_timeout: OAuthDiscoveryTimeout, + redirect_mode: StreamableHttpRedirectMode, +) -> Result { + let has_configured_headers = has_configured_headers(&http_headers, &env_http_headers); + let default_headers = match auth_status_before_discovery( + server_name, + url, + bearer_token_env_var, + http_headers, + env_http_headers, + store_mode, + keyring_backend_kind, + )? { + AuthStatusCheck::Complete(status) => return Ok(status), + AuthStatusCheck::Discover(default_headers) => default_headers, + }; + determine_auth_status_from_discovery( + server_name, + url, + discover_streamable_http_oauth_with_headers_and_http_client( + url, + default_headers, + http_client, + discovery_timeout, + has_configured_headers, + redirect_mode, + ) + .await, + ) +} + +/// Determine authentication status using only configured and stored credentials. +/// +/// Returns `None` when determining the status would require OAuth metadata discovery. +pub fn determine_streamable_http_auth_status_from_credentials( + server_name: &str, + url: &str, + bearer_token_env_var: Option<&str>, + http_headers: Option>, + env_http_headers: Option>, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result> { + match auth_status_before_discovery( + server_name, + url, + bearer_token_env_var, + http_headers, + env_http_headers, + store_mode, + keyring_backend_kind, + )? { + AuthStatusCheck::Complete(status) => Ok(Some(status)), + AuthStatusCheck::Discover(_) => Ok(None), + } +} + +fn auth_status_before_discovery( + server_name: &str, + url: &str, + bearer_token_env_var: Option<&str>, + http_headers: Option>, + env_http_headers: Option>, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result { + if bearer_token_env_var.is_some() { + return Ok(AuthStatusCheck::Complete(McpAuthState::BearerToken)); + } + + let default_headers = build_default_headers(http_headers, env_http_headers)?; + if default_headers.contains_key(AUTHORIZATION) { + return Ok(AuthStatusCheck::Complete(McpAuthState::BearerToken)); + } + + match oauth_token_status(server_name, url, store_mode, keyring_backend_kind)? { + StoredOAuthTokenStatus::Usable => { + return Ok(AuthStatusCheck::Complete(McpAuthState::OAuth)); + } + StoredOAuthTokenStatus::AuthorizationRequired => { + return Ok(AuthStatusCheck::Complete(McpAuthState::LoggedOut( + McpLoginRequirement::Reauthentication, + ))); + } + StoredOAuthTokenStatus::Missing => {} + } + + Ok(AuthStatusCheck::Discover(default_headers)) +} + +fn determine_auth_status_from_discovery( + server_name: &str, + url: &str, + discovery: Result>, +) -> Result { + match discovery { + Ok(Some(_)) => Ok(McpAuthState::LoggedOut(McpLoginRequirement::Login)), + Ok(None) => Ok(McpAuthState::Unsupported), + Err(error) => { + debug!( + "failed to detect OAuth support for MCP server `{server_name}` at {url}: {error:?}" + ); + Err(error) + } + } +} + +pub async fn discover_streamable_http_oauth( + url: &str, + http_headers: Option>, + env_http_headers: Option>, + http_client: Arc, + discovery_timeout: OAuthDiscoveryTimeout, + redirect_mode: StreamableHttpRedirectMode, +) -> Result> { + let has_configured_headers = has_configured_headers(&http_headers, &env_http_headers); + let default_headers = build_default_headers(http_headers, env_http_headers)?; + discover_streamable_http_oauth_with_headers_and_http_client( + url, + default_headers, + http_client, + discovery_timeout, + has_configured_headers, + redirect_mode, + ) + .await +} + +async fn discover_streamable_http_oauth_with_headers_and_http_client( + url: &str, + default_headers: HeaderMap, + http_client: Arc, + discovery_timeout: OAuthDiscoveryTimeout, + has_configured_headers: bool, + redirect_mode: StreamableHttpRedirectMode, +) -> Result> { + let oauth_http_client = match discovery_timeout { + OAuthDiscoveryTimeout::Requested => OAuthHttpClientAdapter::new_with_redirect_mode( + http_client, + default_headers, + url, + has_configured_headers, + redirect_mode, + )?, + OAuthDiscoveryTimeout::Capped(max_timeout) => { + OAuthHttpClientAdapter::new_with_max_timeout_and_redirect_mode( + http_client, + default_headers, + url, + max_timeout, + has_configured_headers, + redirect_mode, + )? + } + }; + let mut authorization_manager = + AuthorizationManager::new_with_oauth_http_client(url, Arc::new(oauth_http_client)).await?; + authorization_manager.set_allow_missing_issuer(true); + discover_streamable_http_oauth_with_manager(&authorization_manager).await +} + +fn has_configured_headers( + http_headers: &Option>, + env_http_headers: &Option>, +) -> bool { + http_headers + .as_ref() + .is_some_and(|headers| !headers.is_empty()) + || env_http_headers + .as_ref() + .is_some_and(|headers| !headers.is_empty()) +} + +async fn discover_streamable_http_oauth_with_manager( + authorization_manager: &AuthorizationManager, +) -> Result> { + match authorization_manager.resolve_metadata().boxed().await { + Ok(resolution) if !resolution.source.is_discovered() => Ok(None), + Ok(resolution) => { + let metadata = resolution.metadata; + Ok(Some(StreamableHttpOAuthDiscovery { + callback_mode: callback_mode(&metadata) + .unwrap_or(McpOAuthCallbackMode::CallbackSpecific), + scopes_supported: normalize_scopes(metadata.scopes_supported), + })) + } + Err(AuthError::NoAuthorizationSupport) => Ok(None), + Err(err) => Err(err.into()), + } +} + +fn normalize_scopes(scopes_supported: Option>) -> Option> { + let scopes_supported = scopes_supported?; + + let mut normalized = Vec::new(); + for scope in scopes_supported { + let scope = scope.trim(); + if scope.is_empty() { + continue; + } + let scope = scope.to_string(); + if !normalized.contains(&scope) { + normalized.push(scope); + } + } + + if normalized.is_empty() { + None + } else { + Some(normalized) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use axum::Json; + use axum::Router; + use axum::http::StatusCode; + use axum::http::header::WWW_AUTHENTICATE; + use axum::routing::get; + use codex_exec_server::ExecServerError; + use codex_exec_server::HttpRedirectPolicy; + use codex_exec_server::HttpRequestParams; + use codex_exec_server::HttpRequestResponse; + use codex_exec_server::HttpResponseBodyStream; + use codex_exec_server::RouteAwareHttpClient; + use codex_http_client::HttpClientFactory; + use codex_http_client::OutboundProxyPolicy; + use futures::future::BoxFuture; + use pretty_assertions::assert_eq; + use serial_test::serial; + use std::collections::HashMap; + use std::ffi::OsString; + use std::sync::Mutex; + use tokio::task::JoinHandle; + use wiremock::Mock; + use wiremock::MockServer; + use wiremock::ResponseTemplate; + use wiremock::matchers::header; + use wiremock::matchers::method; + use wiremock::matchers::path; + + struct TestServer { + url: String, + handle: JoinHandle<()>, + } + + fn test_http_client() -> Arc { + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))) + } + + impl Drop for TestServer { + fn drop(&mut self) { + self.handle.abort(); + } + } + + #[derive(Default)] + struct RecordingHttpClient { + headers: Mutex>>, + redirect_policy: Mutex>, + timeout_ms: Mutex>>, + } + + impl HttpClient for RecordingHttpClient { + fn http_request( + &self, + _params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + Box::pin(async { + Err(ExecServerError::HttpRequest( + "unexpected buffered request".to_string(), + )) + }) + } + + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> + { + *self + .headers + .lock() + .expect("header recorder lock should not be poisoned") = Some( + params + .headers + .iter() + .map(|header| (header.name.clone(), header.value.clone())) + .collect(), + ); + *self + .timeout_ms + .lock() + .expect("timeout recorder lock should not be poisoned") = Some(params.timeout_ms); + *self + .redirect_policy + .lock() + .expect("redirect policy recorder lock should not be poisoned") = + Some(params.redirect_policy); + Box::pin(async { + Err(ExecServerError::HttpRequest( + "expected discovery request failure".to_string(), + )) + }) + } + } + + fn assert_recorded_discovery_failure(discovery: Result>) { + let error = discovery.expect_err("the recording HTTP client rejects OAuth discovery"); + assert!( + matches!( + error.downcast_ref::(), + Some(AuthError::MetadataError(reason)) + if reason.contains("expected discovery request failure") + ), + "OAuth discovery must preserve the executor transport failure: {error:#}" + ); + } + + async fn spawn_oauth_discovery_server(metadata: serde_json::Value) -> TestServer { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let address = listener.local_addr().expect("listener should have address"); + let mut metadata = metadata; + if let Some(metadata) = metadata.as_object_mut() { + metadata + .entry("issuer") + .or_insert_with(|| format!("http://{address}/mcp").into()); + } + let app = Router::new().route( + "/.well-known/oauth-authorization-server/mcp", + get({ + let metadata = metadata.clone(); + move || { + let metadata = metadata.clone(); + async move { Json(metadata) } + } + }), + ); + let handle = tokio::spawn(async move { + axum::serve(listener, app).await.expect("server should run"); + }); + + TestServer { + url: format!("http://{address}/mcp"), + handle, + } + } + + struct EnvVarGuard { + key: String, + original: Option, + } + + impl EnvVarGuard { + fn set(key: &str, value: &str) -> Self { + let original = std::env::var_os(key); + unsafe { + std::env::set_var(key, value); + } + Self { + key: key.to_string(), + original, + } + } + } + + impl Drop for EnvVarGuard { + fn drop(&mut self) { + if let Some(value) = &self.original { + unsafe { + std::env::set_var(&self.key, value); + } + } else { + unsafe { + std::env::remove_var(&self.key); + } + } + } + } + + #[tokio::test] + async fn determine_auth_status_uses_bearer_token_when_authorization_header_present() { + let status = determine_streamable_http_auth_status( + "server", + "not-a-url", + /*bearer_token_env_var*/ None, + Some(HashMap::from([( + "Authorization".to_string(), + "Bearer token".to_string(), + )])), + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::Keyring, + AuthKeyringBackendKind::default(), + test_http_client(), + OAuthDiscoveryTimeout::Requested, + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect("status should compute"); + + assert_eq!(status, McpAuthState::BearerToken); + } + + #[tokio::test] + #[serial(auth_status_env)] + async fn determine_auth_status_uses_bearer_token_when_env_authorization_header_present() { + let _guard = EnvVarGuard::set("CODEX_RMCP_CLIENT_AUTH_STATUS_TEST_TOKEN", "Bearer token"); + let status = determine_streamable_http_auth_status( + "server", + "not-a-url", + /*bearer_token_env_var*/ None, + /*http_headers*/ None, + Some(HashMap::from([( + "Authorization".to_string(), + "CODEX_RMCP_CLIENT_AUTH_STATUS_TEST_TOKEN".to_string(), + )])), + OAuthCredentialsStoreMode::Keyring, + AuthKeyringBackendKind::default(), + test_http_client(), + OAuthDiscoveryTimeout::Requested, + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect("status should compute"); + + assert_eq!(status, McpAuthState::BearerToken); + } + + #[tokio::test] + async fn oauth_metadata_preserves_login_without_probing_anonymous_tools() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let address = listener.local_addr().expect("listener should have address"); + let metadata = serde_json::json!({ + "issuer": format!("http://{address}/mcp"), + "authorization_endpoint": format!("http://{address}/authorize"), + "token_endpoint": format!("http://{address}/token"), + }); + let app = Router::new() + .route( + "/mcp", + get(|| async { StatusCode::METHOD_NOT_ALLOWED }).post( + |Json(request): Json| async move { + let result = match request["method"].as_str() { + Some("initialize") => serde_json::json!({ + "protocolVersion": "2024-11-05", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "oauth", "version": "1"}, + }), + Some("tools/list") => serde_json::json!({"tools": []}), + _ => serde_json::json!({}), + }; + Json(serde_json::json!({ + "jsonrpc": "2.0", + "id": request["id"], + "result": result, + })) + }, + ), + ) + .route( + "/.well-known/oauth-authorization-server/mcp", + get(move || async move { Json(metadata) }), + ); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.expect("server should run"); + }); + let url = format!("http://{address}/mcp"); + let discovery = discover_streamable_http_oauth( + &url, + /*http_headers*/ None, + /*env_http_headers*/ None, + test_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await; + assert_eq!( + determine_auth_status_from_discovery("server", &url, discovery) + .expect("auth status should compute"), + McpAuthState::LoggedOut(McpLoginRequirement::Login) + ); + server.abort(); + } + + #[tokio::test] + async fn oauth_discovery_does_not_follow_cross_origin_redirects() { + let redirect_target = MockServer::start().await; + let redirect_url = format!("{}/redirect-target", redirect_target.uri()); + Mock::given(method("GET")) + .and(path("/redirect-target")) + .and(header("x-api-key", "sensitive-key")) + .respond_with(ResponseTemplate::new(200)) + .expect(0) + .mount(&redirect_target) + .await; + + let resource_server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/mcp")) + .and(header("x-api-key", "sensitive-key")) + .respond_with( + ResponseTemplate::new(302).insert_header("location", redirect_url.clone()), + ) + .expect(1) + .mount(&resource_server) + .await; + + let error = discover_streamable_http_oauth( + &format!("{}/mcp", resource_server.uri()), + Some(HashMap::from([( + "x-api-key".to_string(), + "sensitive-key".to_string(), + )])), + /*env_http_headers*/ None, + test_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect_err("cross-origin OAuth discovery redirects must be rejected"); + + assert!( + matches!( + error.downcast_ref::(), + Some(AuthError::MetadataError(reason)) + if reason.contains("OAuth discovery redirect to non-same-origin URL rejected") + && reason.contains(&redirect_url) + ), + "OAuth discovery must preserve the cross-origin redirect rejection: {error:#}" + ); + redirect_target.verify().await; + resource_server.verify().await; + } + + #[tokio::test] + async fn determine_auth_status_preserves_transient_http_errors() { + for status in [ + StatusCode::REQUEST_TIMEOUT, + StatusCode::TOO_EARLY, + StatusCode::TOO_MANY_REQUESTS, + ] { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/mcp")) + .respond_with(ResponseTemplate::new(status.as_u16())) + .expect(1) + .mount(&server) + .await; + + let error = determine_streamable_http_auth_status( + "transient-http-error", + &format!("{}/mcp", server.uri()), + /*bearer_token_env_var*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + test_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect_err("transient OAuth discovery failures must not become unsupported access"); + + assert!( + matches!( + error.downcast_ref::(), + Some(AuthError::MetadataError(reason)) if reason.contains(status.as_str()) + ), + "auth-status discovery must preserve HTTP {status}: {error:#}" + ); + server.verify().await; + } + } + + #[tokio::test] + async fn discover_streamable_http_oauth_returns_normalized_scopes() { + let server = spawn_oauth_discovery_server(serde_json::json!({ + "authorization_endpoint": "https://example.com/authorize", + "token_endpoint": "https://example.com/token", + "authorization_response_iss_parameter_supported": true, + "scopes_supported": ["profile", " email ", "profile", "", " "], + })) + .await; + + let discovery = discover_streamable_http_oauth( + &server.url, + /*http_headers*/ None, + /*env_http_headers*/ None, + test_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect("discovery should succeed") + .expect("oauth support should be detected"); + + assert_eq!( + discovery, + StreamableHttpOAuthDiscovery { + scopes_supported: Some(vec!["profile".to_string(), "email".to_string()]), + callback_mode: McpOAuthCallbackMode::IssuerBound, + } + ); + } + + #[tokio::test] + async fn issuer_support_without_a_metadata_issuer_falls_back_to_distinct_callbacks() { + let server = spawn_oauth_discovery_server(serde_json::json!({ + "issuer": null, + "authorization_endpoint": "https://example.com/authorize", + "token_endpoint": "https://example.com/token", + "authorization_response_iss_parameter_supported": true, + })) + .await; + + let discovery = discover_streamable_http_oauth( + &server.url, + /*http_headers*/ None, + /*env_http_headers*/ None, + test_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect("discovery should succeed") + .expect("oauth support should be detected"); + + assert_eq!( + discovery, + StreamableHttpOAuthDiscovery { + scopes_supported: None, + callback_mode: McpOAuthCallbackMode::CallbackSpecific, + } + ); + } + + #[tokio::test] + async fn routed_oauth_discovery_caps_local_discovery_timeout() { + let http_client = Arc::new(RecordingHttpClient::default()); + + let discovery = discover_streamable_http_oauth( + "http://example.com/mcp", + /*http_headers*/ None, + /*env_http_headers*/ None, + http_client.clone(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await; + + assert_recorded_discovery_failure(discovery); + assert_eq!( + *http_client + .timeout_ms + .lock() + .expect("timeout recorder lock should not be poisoned"), + Some(Some( + u64::try_from(DISCOVERY_TIMEOUT.as_millis()) + .expect("discovery timeout should fit in u64") + )) + ); + } + + #[tokio::test] + async fn routed_oauth_discovery_preserves_requested_timeout() { + let http_client = Arc::new(RecordingHttpClient::default()); + + let discovery = discover_streamable_http_oauth( + "http://example.com/mcp", + /*http_headers*/ None, + /*env_http_headers*/ None, + http_client.clone(), + OAuthDiscoveryTimeout::Requested, + StreamableHttpRedirectMode::Legacy, + ) + .await; + + assert_recorded_discovery_failure(discovery); + assert_eq!( + *http_client + .timeout_ms + .lock() + .expect("timeout recorder lock should not be poisoned"), + Some(Some(30_000)) + ); + } + + #[tokio::test] + async fn routed_agent_plugin_oauth_discovery_stops_with_configured_headers() { + let http_client = Arc::new(RecordingHttpClient::default()); + + let discovery = discover_streamable_http_oauth( + "http://example.com/mcp", + Some(HashMap::from([( + "X-Mcp-Discovery".to_string(), + "configured-value".to_string(), + )])), + /*env_http_headers*/ None, + http_client.clone(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::AgentPluginV1, + ) + .await; + + assert_recorded_discovery_failure(discovery); + let headers = http_client + .headers + .lock() + .expect("header recorder lock should not be poisoned") + .clone() + .expect("discovery should issue an HTTP request"); + assert_eq!( + headers + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("x-mcp-discovery")) + .map(|(_, value)| value.as_str()), + Some("configured-value") + ); + assert_eq!( + *http_client + .redirect_policy + .lock() + .expect("redirect policy recorder lock should not be poisoned"), + Some(HttpRedirectPolicy::Stop) + ); + } + + #[tokio::test] + async fn discover_streamable_http_oauth_follows_protected_resource_metadata() { + let authorization_server = spawn_oauth_discovery_server(serde_json::json!({ + "authorization_endpoint": "https://example.com/authorize", + "token_endpoint": "https://example.com/token", + "scopes_supported": ["read", " write ", "read"], + })) + .await; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let address = listener.local_addr().expect("listener should have address"); + let resource_metadata_url = format!("http://{address}/oauth-resource"); + let challenge = format!("Bearer resource_metadata=\"{resource_metadata_url}\""); + let authorization_server_url = authorization_server.url.clone(); + let app = Router::new() + .route( + "/mcp", + get(move || { + let challenge = challenge.clone(); + async move { (StatusCode::UNAUTHORIZED, [(WWW_AUTHENTICATE, challenge)]) } + }), + ) + .route( + "/oauth-resource", + get(move || { + let authorization_server_url = authorization_server_url.clone(); + async move { + Json(serde_json::json!({ + "resource": format!("http://{address}/mcp"), + "authorization_servers": [authorization_server_url], + })) + } + }), + ); + let handle = tokio::spawn(async move { + axum::serve(listener, app).await.expect("server should run"); + }); + let resource_server = TestServer { + url: format!("http://{address}/mcp"), + handle, + }; + + let discovery = discover_streamable_http_oauth( + &resource_server.url, + /*http_headers*/ None, + /*env_http_headers*/ None, + test_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect("discovery should succeed") + .expect("oauth support should be detected"); + + assert_eq!( + discovery.scopes_supported, + Some(vec!["read".to_string(), "write".to_string()]) + ); + } + + #[tokio::test] + async fn discover_streamable_http_oauth_ignores_empty_scopes() { + let server = spawn_oauth_discovery_server(serde_json::json!({ + "authorization_endpoint": "https://example.com/authorize", + "token_endpoint": "https://example.com/token", + "scopes_supported": ["", " "], + })) + .await; + + let discovery = discover_streamable_http_oauth( + &server.url, + /*http_headers*/ None, + /*env_http_headers*/ None, + test_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect("discovery should succeed") + .expect("oauth support should be detected"); + + assert_eq!(discovery.scopes_supported, None); + } + + #[tokio::test] + async fn supports_oauth_login_does_not_require_scopes_supported() { + let server = spawn_oauth_discovery_server(serde_json::json!({ + "authorization_endpoint": "https://example.com/authorize", + "token_endpoint": "https://example.com/token", + })) + .await; + + let supported = discover_streamable_http_oauth( + &server.url, + /*http_headers*/ None, + /*env_http_headers*/ None, + test_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect("support check should succeed") + .is_some(); + + assert!(supported); + } +} diff --git a/codex-rs/rmcp-client/src/bin/rmcp_test_server.rs b/codex-rs/rmcp-client/src/bin/rmcp_test_server.rs new file mode 100644 index 0000000000000000000000000000000000000000..a2ab794f5ab76dc5745c6c74a0822091ff82ccf0 --- /dev/null +++ b/codex-rs/rmcp-client/src/bin/rmcp_test_server.rs @@ -0,0 +1,151 @@ +use std::borrow::Cow; +use std::collections::HashMap; +use std::sync::Arc; + +use rmcp::ErrorData as McpError; +use rmcp::ServiceExt; +use rmcp::handler::server::ServerHandler; +use rmcp::model::CallToolRequestParams; +use rmcp::model::CallToolResult; +use rmcp::model::JsonObject; +use rmcp::model::ListToolsResult; +use rmcp::model::PaginatedRequestParams; +use rmcp::model::ServerCapabilities; +use rmcp::model::ServerInfo; +use rmcp::model::Tool; +use serde::Deserialize; +use serde_json::json; +use tokio::task; + +#[derive(Clone)] +struct TestToolServer { + tools: Arc>, +} +pub fn stdio() -> (tokio::io::Stdin, tokio::io::Stdout) { + (tokio::io::stdin(), tokio::io::stdout()) +} +impl TestToolServer { + fn new() -> Self { + let tools = vec![Self::echo_tool()]; + Self { + tools: Arc::new(tools), + } + } + + fn echo_tool() -> Tool { + #[expect(clippy::expect_used)] + let schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { + "message": { "type": "string" }, + "env_var": { "type": "string" } + }, + "required": ["message"], + "additionalProperties": false + })) + .expect("echo tool schema should deserialize"); + + let mut tool = Tool::new( + Cow::Borrowed("echo"), + Cow::Borrowed("Echo back the provided message and include environment data."), + Arc::new(schema), + ); + #[expect(clippy::expect_used)] + let output_schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { + "echo": { "type": "string" }, + "env": { + "anyOf": [ + { "type": "string" }, + { "type": "null" } + ] + } + }, + "required": ["echo", "env"], + "additionalProperties": false + })) + .expect("echo tool output schema should deserialize"); + tool.output_schema = Some(Arc::new(output_schema)); + tool + } +} + +#[derive(Deserialize)] +struct EchoArgs { + message: String, + env_var: Option, +} + +impl ServerHandler for TestToolServer { + fn get_info(&self) -> ServerInfo { + ServerInfo::new( + ServerCapabilities::builder() + .enable_tools() + .enable_tool_list_changed() + .build(), + ) + } + + fn list_tools( + &self, + _request: Option, + _context: rmcp::service::RequestContext, + ) -> impl std::future::Future> + Send + '_ { + let tools = self.tools.clone(); + async move { Ok(ListToolsResult::with_all_items((*tools).clone())) } + } + + async fn call_tool( + &self, + request: CallToolRequestParams, + _context: rmcp::service::RequestContext, + ) -> Result { + match request.name.as_ref() { + "echo" => { + let args: EchoArgs = match request.arguments { + Some(arguments) => serde_json::from_value(serde_json::Value::Object( + arguments.into_iter().collect(), + )) + .map_err(|err| McpError::invalid_params(err.to_string(), None))?, + None => { + return Err(McpError::invalid_params( + "missing arguments for echo tool", + None, + )); + } + }; + + let env_snapshot: HashMap = std::env::vars().collect(); + let env_name = args.env_var.as_deref().unwrap_or("MCP_TEST_VALUE"); + let structured_content = json!({ + "echo": args.message, + "env": env_snapshot.get(env_name), + }); + + let mut result = CallToolResult::success(Vec::new()); + result.structured_content = Some(structured_content); + Ok(result.into()) + } + other => Err(McpError::invalid_params( + format!("unknown tool: {other}"), + None, + )), + } + } +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + eprintln!("starting rmcp test server"); + // Run the server with STDIO transport. If the client disconnects we simply + // bubble up the error so the process exits. + let service = TestToolServer::new(); + let running = service.serve(stdio()).await?; + + // Wait for the client to finish interacting with the server. + running.waiting().await?; + // Drain background tasks to ensure clean shutdown. + task::yield_now().await; + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/bin/test_mcp_2026_discovery_stdio_server.rs b/codex-rs/rmcp-client/src/bin/test_mcp_2026_discovery_stdio_server.rs new file mode 100644 index 0000000000000000000000000000000000000000..fc3af65944258d18731493994fcf0a18e5f8ce1a --- /dev/null +++ b/codex-rs/rmcp-client/src/bin/test_mcp_2026_discovery_stdio_server.rs @@ -0,0 +1,66 @@ +use std::io; +use std::io::BufRead; +use std::io::Write; + +use anyhow::Result; +use serde_json::Value; +use serde_json::json; + +fn main() -> Result<()> { + assert!(std::env::var_os("CODEX_MCP_PROTOCOL_VERSION").is_none()); + let mut stdout = io::stdout().lock(); + + for line in io::stdin().lock().lines() { + let request: Value = serde_json::from_str(&line?)?; + let result = match request["method"].as_str() { + Some("server/discover") => { + assert_eq!( + request["params"]["_meta"]["io.modelcontextprotocol/protocolVersion"], + "2026-07-28" + ); + json!({ + "resultType": "complete", + "supportedVersions": ["2026-07-28"], + "capabilities": {"tools": {}, "resources": {}}, + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "strict-stdio-discovery", + "version": "1.0.0", + }, + }, + "ttlMs": 0, + "cacheScope": "private", + }) + } + Some("tools/list") => json!({ + "resultType": "complete", + "tools": [{ + "name": "stdio_echo", + "inputSchema": {"type": "object"}, + }], + }), + Some("resources/list") => json!({ + "resultType": "complete", + "resources": [{ + "uri": "test://stdio/resource", + "name": "stdio resource", + }], + }), + Some("resources/templates/list") => json!({ + "resultType": "complete", + "resourceTemplates": [{ + "uriTemplate": "test://stdio/{name}", + "name": "stdio resource template", + }], + }), + method => anyhow::bail!("unexpected modern stdio discovery method: {method:?}"), + }; + serde_json::to_writer( + &mut stdout, + &json!({"jsonrpc": "2.0", "id": request["id"], "result": result}), + )?; + stdout.write_all(b"\n")?; + stdout.flush()?; + } + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/bin/test_mcp_2026_stdio_server.rs b/codex-rs/rmcp-client/src/bin/test_mcp_2026_stdio_server.rs new file mode 100644 index 0000000000000000000000000000000000000000..3f6afd99f3691e27cb3e24b570698218a5009221 --- /dev/null +++ b/codex-rs/rmcp-client/src/bin/test_mcp_2026_stdio_server.rs @@ -0,0 +1,250 @@ +use std::io; +use std::io::BufRead; +use std::io::Write; + +use anyhow::Result; +use serde_json::Value; +use serde_json::json; + +const MODERN_VERSION: &str = "2026-07-28"; +const LEGACY_VERSION: &str = "2025-06-18"; +const MAX_MESSAGE_BYTES: usize = 8 * 1024 * 1024; + +fn write_message(stdout: &mut io::StdoutLock<'_>, message: Value) -> Result<()> { + serde_json::to_writer(&mut *stdout, &message)?; + stdout.write_all(b"\n")?; + stdout.flush()?; + Ok(()) +} + +fn assert_modern_metadata(message: &Value) { + let metadata = &message["params"]["_meta"]; + assert_eq!( + metadata["io.modelcontextprotocol/protocolVersion"], + MODERN_VERSION + ); + assert!(metadata["io.modelcontextprotocol/clientInfo"].is_object()); + assert!(metadata["io.modelcontextprotocol/clientCapabilities"].is_object()); +} + +fn legacy_schema_defaults() -> Value { + json!({ + "name": "John Doe", + "age": 30, + "score": 95.5, + "status": "active", + "verified": true, + }) +} + +fn main() -> Result<()> { + let mode = std::env::args().nth(1).unwrap_or_default(); + assert!(matches!( + mode.as_str(), + "modern" | "legacy" | "legacy-fallback" | "oversized-stdout" | "oversized-stdout-legacy" + )); + assert!(std::env::var_os("CODEX_MCP_PROTOCOL_VERSION").is_none()); + let is_legacy = mode.starts_with("legacy") || mode.ends_with("-legacy"); + let oversized = mode.starts_with("oversized-stdout"); + let mut stdout = io::stdout().lock(); + let mut initialized = false; + let mut pending_legacy_call = None; + + for line in io::stdin().lock().lines() { + let message: Value = serde_json::from_str(&line?)?; + let Some(method) = message.get("method").and_then(Value::as_str) else { + let pending_call = pending_legacy_call + .take() + .ok_or_else(|| anyhow::anyhow!("unexpected server-request response"))?; + assert_eq!(message["id"], "legacy-approval"); + assert_eq!(message["result"]["action"], "accept"); + assert_eq!(message["result"]["content"], legacy_schema_defaults()); + write_message( + &mut stdout, + json!({ + "jsonrpc": "2.0", + "id": pending_call, + "result": {"content": [{"type": "text", "text": "legacy approved"}]}, + }), + )?; + continue; + }; + + match method { + "server/discover" if is_legacy => { + assert_eq!(mode, "legacy-fallback"); + write_message( + &mut stdout, + json!({ + "jsonrpc": "2.0", + "id": message["id"], + "error": {"code": -32601, "message": "method not found"}, + }), + )?; + } + "server/discover" => { + assert!(!initialized); + assert_modern_metadata(&message); + initialized = true; + write_message( + &mut stdout, + json!({ + "jsonrpc": "2.0", + "id": message["id"], + "result": { + "resultType": "complete", + "supportedVersions": [MODERN_VERSION], + "capabilities": {"tools": {}, "resources": {}}, + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "strict-stdio-test", + "version": "1.0.0", + }, + }, + "ttlMs": 0, + "cacheScope": "private", + }, + }), + )?; + } + "initialize" => { + assert!(is_legacy); + assert_eq!(message["params"]["protocolVersion"], LEGACY_VERSION); + initialized = true; + write_message( + &mut stdout, + json!({ + "jsonrpc": "2.0", + "id": message["id"], + "result": { + "protocolVersion": LEGACY_VERSION, + "capabilities": {"tools": {}}, + "serverInfo": {"name": "legacy-stdio-test", "version": "1.0.0"}, + }, + }), + )?; + } + "notifications/initialized" => assert!(is_legacy), + "tools/list" => { + assert!(initialized); + if !is_legacy { + assert_modern_metadata(&message); + } + let description = if oversized { + "x".repeat(MAX_MESSAGE_BYTES + 1) + } else { + "echo a value".to_owned() + }; + let mut result = json!({ + "tools": [{ + "name": "echo", + "description": description, + "inputSchema": {"type": "object"}, + }], + }); + if !is_legacy { + result["resultType"] = json!("complete"); + result["ttlMs"] = json!(0); + result["cacheScope"] = json!("private"); + } + write_message( + &mut stdout, + json!({"jsonrpc": "2.0", "id": message["id"], "result": result}), + )?; + } + "tools/call" if is_legacy => { + assert!(initialized); + pending_legacy_call = Some(message["id"].clone()); + write_message( + &mut stdout, + json!({ + "jsonrpc": "2.0", + "id": "legacy-approval", + "method": "elicitation/create", + "params": { + "message": "Accept the schema defaults", + "requestedSchema": { + "type": "object", + "properties": { + "name": {"type": "string", "default": "John Doe"}, + "age": {"type": "integer", "default": 30}, + "score": {"type": "number", "default": 95.5}, + "status": { + "type": "string", + "enum": ["active", "inactive"], + "default": "active", + }, + "verified": {"type": "boolean", "default": true}, + }, + "required": [], + }, + }, + }), + )?; + } + "tools/call" => { + assert!(initialized); + assert_modern_metadata(&message); + if message.pointer("/params/inputResponses").is_none() { + write_message( + &mut stdout, + json!({ + "jsonrpc": "2.0", + "id": message["id"], + "result": { + "resultType": "input_required", + "inputRequests": { + "approval": { + "method": "elicitation/create", + "params": { + "mode": "form", + "message": "Approve stdio call?", + "requestedSchema": { + "type": "object", + "properties": { + "approved": {"type": "boolean"}, + }, + }, + }, + }, + }, + "requestState": "stdio-state", + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "strict-stdio-test", + "version": "1.0.0", + }, + }, + }, + }), + )?; + continue; + } + + assert_eq!(message["params"]["requestState"], "stdio-state"); + assert_eq!( + message["params"]["inputResponses"]["approval"]["action"], + "accept" + ); + assert_eq!( + message["params"]["inputResponses"]["approval"]["content"]["approved"], + true + ); + write_message( + &mut stdout, + json!({ + "jsonrpc": "2.0", + "id": message["id"], + "result": { + "resultType": "complete", + "content": [{"type": "text", "text": "modern approved"}], + }, + }), + )?; + } + other => anyhow::bail!("unexpected MCP method {other}"), + } + } + + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/bin/test_stdio_server.rs b/codex-rs/rmcp-client/src/bin/test_stdio_server.rs new file mode 100644 index 0000000000000000000000000000000000000000..79a2e1c066c849aa3fd1cd91090e6c1ef11730a5 --- /dev/null +++ b/codex-rs/rmcp-client/src/bin/test_stdio_server.rs @@ -0,0 +1,1062 @@ +use std::borrow::Cow; +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::collections::hash_map::Entry; +use std::sync::Arc; +use std::sync::OnceLock; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use rmcp::ErrorData as McpError; +use rmcp::ServiceExt; +use rmcp::handler::server::ServerHandler; +use rmcp::model::CallToolRequestParams; +use rmcp::model::CallToolResult; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::InitializeResult; +use rmcp::model::JsonObject; +use rmcp::model::ListResourceTemplatesResult; +use rmcp::model::ListResourcesResult; +use rmcp::model::ListToolsResult; +use rmcp::model::MetaObject; +use rmcp::model::PaginatedRequestParams; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ReadResourceResult; +use rmcp::model::Resource; +use rmcp::model::ResourceContents; +use rmcp::model::ResourceTemplate; +use rmcp::model::ServerCapabilities; +use rmcp::model::ServerInfo; +use rmcp::model::Tool; +use rmcp::model::ToolAnnotations; +use serde::Deserialize; +use serde_json::json; +use tokio::sync::Barrier; +use tokio::task; +use tokio::time::sleep; + +#[derive(Clone)] +struct TestToolServer { + tools: Arc>, + resources: Arc>, + resource_templates: Arc>, + supports_openai_form_elicitation: Arc, +} + +const MEMO_URI: &str = "memo://codex/example-note"; +const MEMO_CONTENT: &str = "This is a sample MCP resource served by the rmcp test server."; +const SANDBOX_STATE_META_CAPABILITY: &str = "codex/sandbox-state-meta"; +const SMALL_PNG_BASE64: &str = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR4nGP4z8DwHwAFAAH/iZk9HQAAAABJRU5ErkJggg=="; +const APP_ONLY_CWD_MARKER_FILE_ENV: &str = "MCP_TEST_APP_ONLY_CWD_MARKER_FILE"; +const DYNAMIC_SERVER_METADATA_ENV: &str = "MCP_TEST_DYNAMIC_SERVER_METADATA"; +const INITIALIZE_BARRIER_FILE_ENV: &str = "MCP_TEST_INITIALIZE_BARRIER_FILE"; +const SERVER_INSTRUCTIONS_ENV: &str = "MCP_TEST_SERVER_INSTRUCTIONS"; + +fn dynamic_server_process_label() -> Option { + std::env::var_os(DYNAMIC_SERVER_METADATA_ENV) + .is_some() + .then(|| format!("rmcp-test-process-{}", std::process::id())) +} + +pub fn stdio() -> (tokio::io::Stdin, tokio::io::Stdout) { + (tokio::io::stdin(), tokio::io::stdout()) +} + +impl TestToolServer { + fn new() -> Self { + #[expect(clippy::expect_used)] + let sandbox_meta_schema: JsonObject = serde_json::from_value(serde_json::json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })) + .expect("sandbox_meta tool schema should deserialize"); + let mut sandbox_meta_tool = Tool::new( + Cow::Borrowed("sandbox_meta"), + Cow::Borrowed("Return the MCP request metadata received by this test server."), + Arc::new(sandbox_meta_schema), + ); + sandbox_meta_tool.annotations = Some(ToolAnnotations::new().read_only(true)); + let entitlement_tools = std::env::var("MCP_TEST_DAYBREAK_READ_ONLY") + .ok() + .into_iter() + .flat_map(|read_only| { + ["get_codex_security_daybreak_access", "get_daybreak_access"].map(|name| { + let mut tool = sandbox_meta_tool.clone(); + tool.name = Cow::Borrowed(name); + tool.description = + Some(Cow::Borrowed("Return requested account access metadata.")); + tool.annotations = Some(ToolAnnotations::new().read_only(read_only == "true")); + if name == "get_codex_security_daybreak_access" { + let mut meta = MetaObject::new(); + meta.insert( + "openai/requestedEntitlements".to_string(), + json!(["cyber_trusted_access"]), + ); + tool.meta = Some(meta); + } + tool + }) + }) + .collect::>(); + + #[expect(clippy::expect_used)] + let thread_hint_schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })) + .expect("thread_hint tool schema should deserialize"); + let mut thread_hint_tool = Tool::new( + Cow::Borrowed("thread_hint"), + Cow::Borrowed("Return an unstructured history hint for a thread."), + Arc::new(thread_hint_schema), + ); + thread_hint_tool.annotations = Some(ToolAnnotations::new().read_only(true)); + let mut thread_hint_meta = MetaObject::new(); + thread_hint_meta.insert("ui".to_string(), json!({ "visibility": [] })); + thread_hint_tool.meta = Some(thread_hint_meta); + + #[expect(clippy::expect_used)] + let encrypted_output_schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })) + .expect("encrypted_output tool schema should deserialize"); + let mut encrypted_output_tool = Tool::new( + Cow::Borrowed("encrypted_output"), + Cow::Borrowed("Return mixed plaintext and encrypted content for integration tests."), + Arc::new(encrypted_output_schema), + ); + encrypted_output_tool.annotations = Some(ToolAnnotations::new().read_only(true)); + + let mut tools = vec![ + Self::echo_tool(), + Self::echo_dash_tool(), + encrypted_output_tool, + thread_hint_tool, + Self::client_capabilities_tool(), + Self::cwd_tool(), + Self::sync_tool(), + Self::sync_readonly_tool(), + Self::image_tool(), + Self::image_scenario_tool(), + sandbox_meta_tool, + ]; + tools.extend(entitlement_tools); + if std::env::var_os("MCP_TEST_ENABLE_NODE_REPL_JS").is_some() { + #[expect(clippy::expect_used)] + let schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { "code": { "type": "string" } }, + "required": ["code"], + "additionalProperties": false + })) + .expect("js tool schema should deserialize"); + let mut tool = Tool::new( + Cow::Borrowed("js"), + Cow::Borrowed("Run JavaScript in the test Node REPL."), + Arc::new(schema), + ); + tool.annotations = Some(ToolAnnotations::new().read_only(true)); + tools.push(tool); + } + if let Some(process_label) = dynamic_server_process_label() + && let Some(echo) = tools.iter_mut().find(|tool| tool.name == "echo") + { + echo.description = Some(Cow::Owned(format!("Echo from {process_label}."))); + } + if std::env::var_os("MCP_TEST_OVERSIZED_TOOL_DESCRIPTION").is_some() + && let Some(echo) = tools.iter_mut().find(|tool| tool.name == "echo") + { + echo.description = Some(Cow::Owned("x".repeat(8 * 1024 * 1024 + 1))); + } + let resources = vec![Self::memo_resource()]; + let resource_templates = vec![Self::memo_template()]; + Self { + tools: Arc::new(tools), + resources: Arc::new(resources), + resource_templates: Arc::new(resource_templates), + supports_openai_form_elicitation: Arc::new(AtomicBool::new(false)), + } + } + + fn echo_tool() -> Tool { + Self::build_echo_tool( + "echo", + "Echo back the provided message and include environment data.", + ) + } + + fn echo_dash_tool() -> Tool { + Self::build_echo_tool( + "echo-tool", + "Echo back the provided message via a tool name that is not a legal JS identifier.", + ) + } + + fn build_echo_tool(name: &'static str, description: &'static str) -> Tool { + #[expect(clippy::expect_used)] + let schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { + "message": { "type": "string" }, + "env_var": { "type": "string" } + }, + "required": ["message"], + "additionalProperties": false + })) + .expect("echo tool schema should deserialize"); + + let mut tool = Tool::new( + Cow::Borrowed(name), + Cow::Borrowed(description), + Arc::new(schema), + ); + #[expect(clippy::expect_used)] + let output_schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { + "echo": { "type": "string" }, + "env": { + "anyOf": [ + { "type": "string" }, + { "type": "null" } + ] + }, + }, + "required": ["echo", "env"], + "additionalProperties": false + })) + .expect("echo tool output schema should deserialize"); + tool.output_schema = Some(Arc::new(output_schema)); + tool.annotations = Some(ToolAnnotations::new().read_only(true)); + tool + } + + fn cwd_tool() -> Tool { + #[expect(clippy::expect_used)] + let schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })) + .expect("cwd tool schema should deserialize"); + + let mut tool = Tool::new( + Cow::Borrowed("cwd"), + Cow::Borrowed("Return the current working directory of this test server process."), + Arc::new(schema), + ); + #[expect(clippy::expect_used)] + let output_schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { + "cwd": { "type": "string" } + }, + "required": ["cwd"], + "additionalProperties": false + })) + .expect("cwd tool output schema should deserialize"); + tool.output_schema = Some(Arc::new(output_schema)); + tool.annotations = Some(ToolAnnotations::new().read_only(true)); + tool + } + + fn client_capabilities_tool() -> Tool { + #[expect(clippy::expect_used)] + let schema: JsonObject = serde_json::from_value(serde_json::json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })) + .expect("client capabilities tool schema should deserialize"); + + let mut tool = Tool::new( + Cow::Borrowed("client_capabilities"), + Cow::Borrowed("Return capabilities advertised by the MCP client."), + Arc::new(schema), + ); + tool.annotations = Some(ToolAnnotations::new().read_only(true)); + tool + } + + fn sync_tool() -> Tool { + #[expect(clippy::expect_used)] + let schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { + "sleep_before_ms": { "type": "number" }, + "sleep_after_ms": { "type": "number" }, + "barrier": { + "type": "object", + "properties": { + "id": { "type": "string" }, + "participants": { "type": "number" }, + "timeout_ms": { "type": "number" } + }, + "required": ["id", "participants"], + "additionalProperties": false + } + }, + "additionalProperties": false + })) + .expect("sync tool schema should deserialize"); + + let mut tool = Tool::new( + Cow::Borrowed("sync"), + Cow::Borrowed( + "Synchronize concurrent test calls and optionally delay before or after the barrier.", + ), + Arc::new(schema), + ); + #[expect(clippy::expect_used)] + let output_schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { + "result": { "type": "string" } + }, + "required": ["result"], + "additionalProperties": false + })) + .expect("sync tool output schema should deserialize"); + tool.output_schema = Some(Arc::new(output_schema)); + tool + } + + fn sync_readonly_tool() -> Tool { + let mut tool = Self::sync_tool(); + tool.name = Cow::Borrowed("sync_readonly"); + tool.annotations = Some(ToolAnnotations::new().read_only(true)); + tool + } + + fn image_tool() -> Tool { + #[expect(clippy::expect_used)] + let schema: JsonObject = serde_json::from_value(serde_json::json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + })) + .expect("image tool schema should deserialize"); + + let mut tool = Tool::new( + Cow::Borrowed("image"), + Cow::Borrowed("Return a single image content block."), + Arc::new(schema), + ); + tool.annotations = Some(ToolAnnotations::new().read_only(true)); + tool + } + + /// Tool intended for manual testing of Codex TUI rendering for MCP image tool results. + /// + /// This exists to exercise edge cases where a `CallToolResult.content` includes image blocks + /// that aren't the first item (or includes invalid image blocks before a valid image). + /// + /// Manual testing approach (Codex TUI): + /// - Build this binary: `cargo build -p codex-rmcp-client --bin test_stdio_server` + /// - Register it: + /// - `codex mcp add mcpimg -- /abs/path/to/test_stdio_server` + /// - Then in Codex TUI, ask it to call: + /// - `mcpimg.image_scenario({"scenario":"image_only"})` + /// - `mcpimg.image_scenario({"scenario":"image_only_original_detail"})` + /// - `mcpimg.image_scenario({"scenario":"text_then_image","caption":"Here is the image:"})` + /// - `mcpimg.image_scenario({"scenario":"invalid_base64_then_image"})` + /// - `mcpimg.image_scenario({"scenario":"invalid_image_bytes_then_image"})` + /// - `mcpimg.image_scenario({"scenario":"multiple_valid_images"})` + /// - `mcpimg.image_scenario({"scenario":"image_then_text","caption":"Here is the image:"})` + /// - `mcpimg.image_scenario({"scenario":"text_only","caption":"Here is the image:"})` + /// - You should see an extra history cell: `tool result (image output)`. + fn image_scenario_tool() -> Tool { + #[expect(clippy::expect_used)] + let schema: JsonObject = serde_json::from_value(serde_json::json!({ + "type": "object", + "properties": { + "scenario": { + "type": "string", + "enum": [ + "image_only", + "image_only_original_detail", + "text_then_image", + "invalid_base64_then_image", + "invalid_image_bytes_then_image", + "multiple_valid_images", + "image_then_text", + "text_only" + ] + }, + "caption": { "type": "string" }, + "data_url": { + "type": "string", + "description": "Optional data URL like data:image/png;base64,AAAA...; if omitted, uses a built-in tiny PNG." + } + }, + "required": ["scenario"], + "additionalProperties": false + })) + .expect("image_scenario tool schema should deserialize"); + + let mut tool = Tool::new( + Cow::Borrowed("image_scenario"), + Cow::Borrowed( + "Return content blocks for manual testing of MCP image rendering scenarios.", + ), + Arc::new(schema), + ); + tool.annotations = Some(ToolAnnotations::new().read_only(true)); + tool + } + + fn memo_resource() -> Resource { + Resource::new(MEMO_URI, "example-note") + .with_title("Example Note") + .with_description("A sample MCP resource exposed for integration tests.") + .with_mime_type("text/plain") + } + + fn memo_template() -> ResourceTemplate { + ResourceTemplate::new("memo://codex/{slug}", "codex-memo") + .with_title("Codex Memo") + .with_description("Template for memo://codex/{slug} resources used in tests.") + .with_mime_type("text/plain") + } + + fn memo_text() -> &'static str { + MEMO_CONTENT + } +} + +#[derive(Deserialize)] +struct EchoArgs { + message: String, + env_var: Option, +} + +#[derive(Deserialize)] +struct JsArgs { + code: String, +} + +const DEFAULT_SYNC_TIMEOUT_MS: u64 = 1_000; + +static SYNC_BARRIERS: OnceLock>> = + OnceLock::new(); + +struct SyncBarrierState { + barrier: Arc, + participants: usize, +} + +#[derive(Debug, Deserialize)] +struct SyncBarrierArgs { + id: String, + participants: usize, + #[serde(default = "default_sync_timeout_ms")] + timeout_ms: u64, +} + +#[derive(Debug, Deserialize)] +struct SyncArgs { + #[serde(default)] + sleep_before_ms: Option, + #[serde(default)] + sleep_after_ms: Option, + #[serde(default)] + barrier: Option, +} + +fn default_sync_timeout_ms() -> u64 { + DEFAULT_SYNC_TIMEOUT_MS +} + +fn sync_barrier_map() -> &'static tokio::sync::Mutex> { + SYNC_BARRIERS.get_or_init(|| tokio::sync::Mutex::new(HashMap::new())) +} + +#[derive(Deserialize, Debug)] +#[serde(rename_all = "snake_case")] +/// Scenarios for `image_scenario`, intended to exercise Codex TUI handling of MCP image outputs. +/// +/// The key behavior under test is that the TUI should render an image output cell if *any* +/// decodable image block exists in the tool result content, even if the first block is text or an +/// invalid image. +enum ImageScenario { + ImageOnly, + ImageOnlyOriginalDetail, + TextThenImage, + InvalidBase64ThenImage, + InvalidImageBytesThenImage, + MultipleValidImages, + ImageThenText, + TextOnly, +} + +#[derive(Deserialize, Debug)] +struct ImageScenarioArgs { + scenario: ImageScenario, + #[serde(default)] + caption: Option, + #[serde(default)] + data_url: Option, +} + +impl ServerHandler for TestToolServer { + async fn initialize( + &self, + request: InitializeRequestParams, + context: rmcp::service::RequestContext, + ) -> Result { + if let Ok(barrier_file) = std::env::var(INITIALIZE_BARRIER_FILE_ENV) { + while !std::path::Path::new(&barrier_file).is_file() { + sleep(Duration::from_millis(10)).await; + } + } + self.supports_openai_form_elicitation.store( + request + .capabilities + .extensions + .as_ref() + .is_some_and(|extensions| extensions.contains_key("openai/form")), + Ordering::Relaxed, + ); + context.peer.set_peer_info(request); + Ok(self.get_info()) + } + + fn get_info(&self) -> ServerInfo { + let mut capabilities = ServerCapabilities::builder() + .enable_tools() + .enable_tool_list_changed() + .enable_resources() + .build(); + capabilities.experimental = Some(BTreeMap::from([( + SANDBOX_STATE_META_CAPABILITY.to_string(), + JsonObject::new(), + )])); + + let server_info = ServerInfo::new(capabilities); + let server_info = match dynamic_server_process_label() { + Some(process_label) => server_info + .with_server_info( + Implementation::new("codex-rmcp-test-server", env!("CARGO_PKG_VERSION")) + .with_title(process_label.clone()), + ) + .with_instructions(format!("Use the tools from {process_label}.")), + None => { + server_info.with_instructions("Use these tools to exercise the rmcp test server.") + } + }; + match std::env::var(SERVER_INSTRUCTIONS_ENV) { + Ok(instructions) => server_info.with_instructions(instructions), + Err(_) => server_info, + } + } + + fn list_tools( + &self, + request: Option, + _context: rmcp::service::RequestContext, + ) -> impl std::future::Future> + Send + '_ { + let tools = self.tools.clone(); + async move { + let mut tools = (*tools).clone(); + if let Some(marker_file) = std::env::var_os(APP_ONLY_CWD_MARKER_FILE_ENV) + && std::path::Path::new(&marker_file).is_file() + && let Some(cwd) = tools.iter_mut().find(|tool| tool.name == "cwd") + { + cwd.meta + .get_or_insert_with(MetaObject::new) + .insert("ui".to_string(), json!({ "visibility": ["app"] })); + } + let mut result = ListToolsResult::with_all_items(tools); + match ( + std::env::var("MCP_TEST_TOOL_PAGINATION").as_deref(), + request.and_then(|request| request.cursor).as_deref(), + ) { + (Ok("two-pages"), None) => { + result.tools.retain(|tool| tool.name == "echo"); + result.next_cursor = Some("second".to_string()); + } + (Ok("two-pages"), Some("second")) => { + result.tools.retain(|tool| tool.name == "sync"); + } + (Ok("oversized-cursor"), None) => { + result.tools.retain(|tool| tool.name == "echo"); + result.next_cursor = Some("x".repeat(65_537)); + } + _ => {} + } + Ok(result) + } + } + + fn list_resources( + &self, + _request: Option, + _context: rmcp::service::RequestContext, + ) -> impl std::future::Future> + Send + '_ { + let resources = self.resources.clone(); + async move { Ok(ListResourcesResult::with_all_items((*resources).clone())) } + } + + async fn list_resource_templates( + &self, + _request: Option, + _context: rmcp::service::RequestContext, + ) -> Result { + Ok(ListResourceTemplatesResult::with_all_items( + (*self.resource_templates).clone(), + )) + } + + async fn read_resource( + &self, + ReadResourceRequestParams { uri, .. }: ReadResourceRequestParams, + _context: rmcp::service::RequestContext, + ) -> Result { + if uri == MEMO_URI { + Ok( + ReadResourceResult::new(vec![ResourceContents::TextResourceContents { + uri, + mime_type: Some("text/plain".to_string()), + text: Self::memo_text().to_string(), + meta: None, + }]) + .into(), + ) + } else { + Err(McpError::resource_not_found( + "resource_not_found", + Some(json!({ "uri": uri })), + )) + } + } + + async fn call_tool( + &self, + request: CallToolRequestParams, + context: rmcp::service::RequestContext, + ) -> Result { + match request.name.as_ref() { + "js" => { + let args = Self::parse_call_args::(&request, "js")?; + if args.code == "nodeRepl.fail()" { + Ok(CallToolResult::error(vec![ + rmcp::model::ContentBlock::text("guardian-hidden-failed-result"), + ])) + } else if args.code == "nodeRepl.empty()" { + Ok(CallToolResult::success(vec![ + rmcp::model::ContentBlock::text(" "), + ])) + } else if args.code == "await nodeRepl.emitImage(await tab.screenshot())" { + let mut meta = MetaObject::new(); + meta.insert("codex/imageDetail".to_string(), json!("low")); + Ok(CallToolResult::success(vec![ + rmcp::model::ContentBlock::text("guardian-visible-before-image"), + rmcp::model::ContentBlock::Image( + rmcp::model::ImageContent::new(SMALL_PNG_BASE64, "IMAGE/PNG") + .with_meta(meta), + ), + rmcp::model::ContentBlock::text("guardian-visible-after-image"), + ])) + } else if let Some(text) = args.code.strip_prefix("nodeRepl.write(") + && let Some(text) = text.strip_suffix(')') + { + let text = serde_json::from_str::(text) + .map_err(|error| McpError::invalid_params(error.to_string(), None))?; + let mut result = + CallToolResult::success(vec![rmcp::model::ContentBlock::text(text)]); + result.structured_content = + Some(json!({ "text": "guardian-hidden-structured-override" })); + let mut meta = MetaObject::new(); + meta.insert("ui".to_string(), json!("guardian-hidden-ui-preview")); + result.meta = Some(meta); + Ok(result) + } else { + Err(McpError::invalid_params("unsupported test js source", None)) + } + } + "client_capabilities" => Ok(Self::structured_result(json!({ + "supportsOpenaiFormElicitation": self + .supports_openai_form_elicitation + .load(Ordering::Relaxed), + }))), + "sandbox_meta" | "get_codex_security_daybreak_access" | "get_daybreak_access" => Ok( + Self::structured_result(serde_json::Value::Object(context.meta.0.0)), + ), + "cwd" => { + let cwd = std::env::current_dir() + .map(|path| path.to_string_lossy().into_owned()) + .map_err(|err| McpError::internal_error(err.to_string(), None))?; + Ok(Self::structured_result(json!({ "cwd": cwd }))) + } + "thread_hint" => { + let thread_id = context + .meta + .0 + .get("threadId") + .and_then(serde_json::Value::as_str) + .ok_or_else(|| { + McpError::invalid_params("missing threadId metadata".to_string(), None) + })?; + Ok(CallToolResult::success(vec![ + rmcp::model::ContentBlock::text(format!( + "manual history hint for thread {thread_id}" + )), + rmcp::model::ContentBlock::text( + "unstructured notes/thread_hint fixture result", + ), + ])) + } + "echo" | "echo-tool" => { + let args: EchoArgs = match request.arguments { + Some(arguments) => serde_json::from_value(serde_json::Value::Object( + arguments.into_iter().collect(), + )) + .map_err(|err| McpError::invalid_params(err.to_string(), None))?, + None => { + return Err(McpError::invalid_params( + format!("missing arguments for {} tool", request.name), + None, + )); + } + }; + + let env_snapshot: HashMap = std::env::vars().collect(); + let env_name = args.env_var.as_deref().unwrap_or("MCP_TEST_VALUE"); + let echo = dynamic_server_process_label() + .unwrap_or_else(|| format!("ECHOING: {}", args.message)); + let structured_content = json!({ + "echo": echo, + "env": env_snapshot.get(env_name), + }); + + Ok(Self::structured_result(structured_content)) + } + "encrypted_output" => { + let mut meta = MetaObject::new(); + meta.insert("codex/encryptedContent".to_string(), json!(true)); + let mut result = CallToolResult::success(vec![ + rmcp::model::ContentBlock::text("Lookup completed"), + rmcp::model::ContentBlock::Text( + rmcp::model::TextContent::new("gAAAA-test").with_meta(meta), + ), + ]); + result.structured_content = Some(json!({"encrypted_output": "ignored"})); + Ok(result) + } + "image" => { + // Read a data URL (e.g. data:image/png;base64,AAA...) from env and convert to + // an MCP image content block. Tests set MCP_TEST_IMAGE_DATA_URL. + let data_url = std::env::var("MCP_TEST_IMAGE_DATA_URL").map_err(|_| { + McpError::invalid_params( + "missing MCP_TEST_IMAGE_DATA_URL env var for image tool", + None, + ) + })?; + + let (mime_type, data_b64) = parse_data_url(&data_url).ok_or_else(|| { + McpError::invalid_params( + format!("invalid data URL for image tool: {data_url}"), + None, + ) + })?; + + Ok(CallToolResult::success(vec![ + rmcp::model::ContentBlock::image(data_b64, mime_type), + ])) + } + "image_scenario" => { + let args = Self::parse_call_args::(&request, "image_scenario")?; + Self::image_scenario_result(args) + } + "sync" => { + let args = Self::parse_call_args::(&request, "sync")?; + Self::sync_result(args).await + } + "sync_readonly" => { + let args = Self::parse_call_args::(&request, "sync_readonly")?; + Self::sync_result(args).await + } + other => Err(McpError::invalid_params( + format!("unknown tool: {other}"), + None, + )), + } + .map(Into::into) + } +} + +impl TestToolServer { + fn parse_call_args Deserialize<'de>>( + request: &CallToolRequestParams, + tool_name: &'static str, + ) -> Result { + match request.arguments.as_ref() { + Some(arguments) => serde_json::from_value(serde_json::Value::Object( + arguments.clone().into_iter().collect(), + )) + .map_err(|err| McpError::invalid_params(err.to_string(), None)), + None => Err(McpError::invalid_params( + format!("missing arguments for {tool_name} tool"), + None, + )), + } + } + + fn image_scenario_result(args: ImageScenarioArgs) -> Result { + let (mime_type, valid_data_b64) = if let Some(data_url) = &args.data_url { + parse_data_url(data_url).ok_or_else(|| { + McpError::invalid_params( + format!("invalid data_url for image_scenario tool: {data_url}"), + None, + ) + })? + } else { + ("image/png".to_string(), SMALL_PNG_BASE64.to_string()) + }; + + let caption = args + .caption + .unwrap_or_else(|| "Here is the image:".to_string()); + + let mut content = Vec::new(); + match args.scenario { + ImageScenario::ImageOnly => { + content.push(rmcp::model::ContentBlock::image(valid_data_b64, mime_type)); + } + ImageScenario::ImageOnlyOriginalDetail => { + let mut meta = MetaObject::new(); + meta.insert( + "codex/imageDetail".to_string(), + serde_json::json!("original"), + ); + content.push(rmcp::model::ContentBlock::Image( + rmcp::model::ImageContent::new(valid_data_b64, mime_type).with_meta(meta), + )); + } + ImageScenario::TextThenImage => { + content.push(rmcp::model::ContentBlock::text(caption)); + content.push(rmcp::model::ContentBlock::image(valid_data_b64, mime_type)); + } + ImageScenario::InvalidBase64ThenImage => { + content.push(rmcp::model::ContentBlock::image( + "not-base64".to_string(), + "image/png".to_string(), + )); + content.push(rmcp::model::ContentBlock::image(valid_data_b64, mime_type)); + } + ImageScenario::InvalidImageBytesThenImage => { + let oversized = std::env::var("MCP_TEST_OVERSIZED_INVALID_IMAGE") == Ok("1".into()); + content.push(rmcp::model::ContentBlock::image( + if oversized { + "A".repeat(8 * 1024 * 1024 - 24) + } else { + "bm90IGFuIGltYWdl".to_string() + }, + "image/png".to_string(), + )); + let (mime_type, valid_data_b64) = std::env::var("MCP_TEST_IMAGE_DATA_URL") + .ok() + .and_then(|data_url| parse_data_url(&data_url)) + .unwrap_or((mime_type, valid_data_b64)); + content.push(rmcp::model::ContentBlock::image(valid_data_b64, mime_type)); + } + ImageScenario::MultipleValidImages => { + content.push(rmcp::model::ContentBlock::image( + valid_data_b64.clone(), + mime_type.clone(), + )); + content.push(rmcp::model::ContentBlock::image(valid_data_b64, mime_type)); + } + ImageScenario::ImageThenText => { + content.push(rmcp::model::ContentBlock::image(valid_data_b64, mime_type)); + content.push(rmcp::model::ContentBlock::text(caption)); + } + ImageScenario::TextOnly => { + content.push(rmcp::model::ContentBlock::text(caption)); + } + } + + Ok(CallToolResult::success(content)) + } + + async fn sync_result(args: SyncArgs) -> Result { + if let Some(delay) = args.sleep_before_ms + && delay > 0 + { + sleep(Duration::from_millis(delay)).await; + } + + if let Some(barrier) = args.barrier { + wait_on_sync_barrier(barrier).await?; + } + + if let Some(delay) = args.sleep_after_ms + && delay > 0 + { + sleep(Duration::from_millis(delay)).await; + } + + Ok(Self::structured_result(json!({ "result": "ok" }))) + } + + fn structured_result(value: serde_json::Value) -> CallToolResult { + let mut result = CallToolResult::success(Vec::new()); + result.structured_content = Some(value); + result + } +} + +async fn wait_on_sync_barrier(args: SyncBarrierArgs) -> Result<(), McpError> { + if args.participants == 0 { + return Err(McpError::invalid_params( + "barrier participants must be greater than zero", + None, + )); + } + + if args.timeout_ms == 0 { + return Err(McpError::invalid_params( + "barrier timeout must be greater than zero", + None, + )); + } + + let barrier_id = args.id.clone(); + let barrier = { + let mut map = sync_barrier_map().lock().await; + match map.entry(barrier_id.clone()) { + Entry::Occupied(entry) => { + let state = entry.get(); + if state.participants != args.participants { + let existing = state.participants; + return Err(McpError::invalid_params( + format!( + "barrier {barrier_id} already registered with {existing} participants" + ), + None, + )); + } + state.barrier.clone() + } + Entry::Vacant(entry) => { + let barrier = Arc::new(Barrier::new(args.participants)); + entry.insert(SyncBarrierState { + barrier: barrier.clone(), + participants: args.participants, + }); + barrier + } + } + }; + + let wait_result = + match tokio::time::timeout(Duration::from_millis(args.timeout_ms), barrier.wait()).await { + Ok(wait_result) => wait_result, + Err(_) => { + remove_sync_barrier_if_current(&barrier_id, &barrier).await; + return Err(McpError::invalid_params( + "sync barrier wait timed out", + None, + )); + } + }; + + if wait_result.is_leader() { + remove_sync_barrier_if_current(&barrier_id, &barrier).await; + } + + Ok(()) +} + +async fn remove_sync_barrier_if_current(barrier_id: &str, barrier: &Arc) { + let mut map = sync_barrier_map().lock().await; + if let Some(state) = map.get(barrier_id) + && Arc::ptr_eq(&state.barrier, barrier) + { + map.remove(barrier_id); + } +} + +fn parse_data_url(url: &str) -> Option<(String, String)> { + let rest = url.strip_prefix("data:")?; + let (mime_and_opts, data) = rest.split_once(',')?; + let (mime, _opts) = mime_and_opts.split_once(';').unwrap_or((mime_and_opts, "")); + Some((mime.to_string(), data.to_string())) +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + #[cfg(windows)] + if std::env::var_os("MCP_TEST_DESCENDANT_ROLE").is_some() { + tokio::time::sleep(Duration::from_secs(30)).await; + return Ok(()); + } + + eprintln!("starting rmcp test server"); + if let Ok(pid_file) = std::env::var("MCP_TEST_PID_FILE") { + std::fs::write(pid_file, std::process::id().to_string())?; + } + #[cfg(windows)] + if let Ok(marker_file) = std::env::var("MCP_TEST_BREAKAWAY_DENIED_FILE") { + use std::os::windows::process::CommandExt; + + const CREATE_BREAKAWAY_FROM_JOB: u32 = 0x0100_0000; + const ERROR_ACCESS_DENIED: i32 = 5; + + let escaped = std::process::Command::new(std::env::current_exe()?) + .creation_flags(CREATE_BREAKAWAY_FROM_JOB) + .env("MCP_TEST_DESCENDANT_ROLE", "1") + .stdin(std::process::Stdio::null()) + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) + .spawn(); + match escaped { + Err(error) if error.raw_os_error() == Some(ERROR_ACCESS_DENIED) => { + std::fs::write(marker_file, "denied")?; + } + Err(error) => return Err(error.into()), + Ok(mut child) => { + let _ = child.kill(); + let _ = child.wait(); + return Err("MCP descendant unexpectedly escaped its Windows job".into()); + } + } + } + #[cfg(windows)] + if let Ok(pid_file) = std::env::var("MCP_TEST_DESCENDANT_PID_FILE") { + let child = std::process::Command::new(std::env::current_exe()?) + .env("MCP_TEST_DESCENDANT_ROLE", "1") + .stdin(std::process::Stdio::null()) + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) + .spawn()?; + std::fs::write(pid_file, child.id().to_string())?; + } + // Run the server with STDIO transport. If the client disconnects we simply + // bubble up the error so the process exits. + let service = TestToolServer::new(); + let running = service.serve(stdio()).await?; + + // A test can close an initialized transport without killing an arbitrary PID. + let exit_file = std::env::var_os("MCP_TEST_EXIT_FILE"); + tokio::select! { + result = running.waiting() => { result?; } + _ = async { + let Some(exit_file) = exit_file else { + return std::future::pending::<()>().await; + }; + while !std::path::Path::new(&exit_file).exists() { + sleep(Duration::from_millis(/*millis*/ 20)).await; + } + } => std::process::exit(0), + } + // Drain background tasks to ensure clean shutdown. + task::yield_now().await; + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs b/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs new file mode 100644 index 0000000000000000000000000000000000000000..9702794bd67bee4cb912e63862736829e7057c62 --- /dev/null +++ b/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs @@ -0,0 +1,590 @@ +use std::borrow::Cow; +use std::collections::HashMap; +use std::fs; +use std::io::ErrorKind; +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; + +use axum::Router; +use axum::body::Body; +use axum::body::to_bytes; +use axum::extract::Json; +use axum::extract::State; +use axum::http::HeaderMap; +use axum::http::HeaderValue; +use axum::http::Method; +use axum::http::Request; +use axum::http::StatusCode; +use axum::http::header::AUTHORIZATION; +use axum::http::header::CONTENT_TYPE; +use axum::http::header::HOST; +use axum::http::header::WWW_AUTHENTICATE; +use axum::middleware; +use axum::middleware::Next; +use axum::response::Response; +use axum::routing::get; +use axum::routing::post; +use rmcp::ErrorData as McpError; +use rmcp::handler::server::ServerHandler; +use rmcp::model::CallToolRequestParams; +use rmcp::model::CallToolResult; +use rmcp::model::JsonObject; +use rmcp::model::ListResourceTemplatesResult; +use rmcp::model::ListResourcesResult; +use rmcp::model::ListToolsResult; +use rmcp::model::PaginatedRequestParams; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ReadResourceResult; +use rmcp::model::Resource; +use rmcp::model::ResourceContents; +use rmcp::model::ResourceTemplate; +use rmcp::model::ServerCapabilities; +use rmcp::model::ServerInfo; +use rmcp::model::Tool; +use rmcp::model::ToolAnnotations; +use rmcp::transport::StreamableHttpServerConfig; +use rmcp::transport::StreamableHttpService; +use rmcp::transport::streamable_http_server::session::local::LocalSessionManager; +use serde::Deserialize; +use serde_json::Value; +use serde_json::json; +use tokio::sync::Mutex; +use tokio::task; +use tokio::time::sleep; + +#[derive(Clone)] +struct TestToolServer { + tools: Arc>, + resources: Arc>, + resource_templates: Arc>, +} + +const MEMO_URI: &str = "memo://codex/example-note"; +const MEMO_CONTENT: &str = "This is a sample MCP resource served by the rmcp test server."; +const MCP_SESSION_ID_HEADER: &str = "mcp-session-id"; +const SESSION_POST_FAILURE_CONTROL_PATH: &str = "/test/control/session-post-failure"; +const INITIALIZE_POST_FAILURE_CONTROL_PATH: &str = "/test/control/initialize-post-failure"; +const INITIALIZED_NOTIFICATION_POST_FAILURE_CONTROL_PATH: &str = + "/test/control/initialized-notification-post-failure"; +const MAX_MCP_POST_BODY_BYTES: usize = 1024 * 1024; + +#[derive(Clone, Default)] +struct PostFailureState { + armed_failure: Arc>>, +} + +#[derive(Clone, Copy, Debug)] +enum ArmedFailureTarget { + Initialize, + InitializedNotification, + Session, +} + +#[derive(Clone, Debug)] +struct ArmedFailure { + target: ArmedFailureTarget, + status: StatusCode, + remaining: usize, + /// Raw `WWW-Authenticate` challenge header field values returned with the failure. + www_authenticate_headers: Vec, + content_type: Option, + body: Option, +} + +#[derive(Debug, Deserialize)] +struct ArmSessionPostFailureRequest { + status: u16, + remaining: usize, + /// Raw `WWW-Authenticate` challenge header field values to add to the failure. + #[serde(default)] + www_authenticate_headers: Vec, + content_type: Option, + body: Option, +} + +#[derive(Deserialize)] +struct EchoArgs { + message: String, + #[allow(dead_code)] + env_var: Option, +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + let mut args = std::env::args_os().skip(1); + match args.next().as_deref() { + Some(value) if value == std::ffi::OsStr::new("--http-headers-helper") => { + if std::env::var_os("MCP_TEST_AMBIENT_SECRET").is_some() { + return Err("helper inherited ambient secret".into()); + } + if let Some(invocations) = args.next() { + let count = fs::read_to_string(&invocations).unwrap_or_default().len() + 1; + fs::write(invocations, "x".repeat(count))?; + let header = + if args.next().as_deref() == Some(std::ffi::OsStr::new("--authorization")) { + "Authorization" + } else { + "Proxy-Authorization" + }; + println!( + r#"{{"{header}":"Bearer gateway-token","X-Helper-Generation":"{count}"}}"# + ); + } else { + println!(r#"{{"Proxy-Authorization":"Bearer gateway-token"}}"#); + } + return Ok(()); + } + _ => {} + } + let bind_addr = parse_bind_addr()?; + let post_failure_state = PostFailureState::default(); + const MAX_BIND_RETRIES: u32 = 20; + const BIND_RETRY_DELAY: Duration = Duration::from_millis(50); + + let mut bind_retries = 0; + let listener = loop { + match tokio::net::TcpListener::bind(&bind_addr).await { + Ok(listener) => break listener, + Err(err) if err.kind() == ErrorKind::PermissionDenied => { + eprintln!( + "failed to bind to {bind_addr}: {err}. make sure the process has network access" + ); + return Ok(()); + } + Err(err) if err.kind() == ErrorKind::AddrInUse && bind_retries < MAX_BIND_RETRIES => { + bind_retries += 1; + sleep(BIND_RETRY_DELAY).await; + } + Err(err) => return Err(err.into()), + } + }; + let actual_bind_addr = listener.local_addr()?; + if let Ok(bound_addr_file) = std::env::var("MCP_STREAMABLE_HTTP_BOUND_ADDR_FILE") { + fs::write(bound_addr_file, actual_bind_addr.to_string())?; + } + eprintln!("starting rmcp streamable http test server on http://{actual_bind_addr}/mcp"); + + let router = Router::new() + .route( + SESSION_POST_FAILURE_CONTROL_PATH, + post(arm_session_post_failure), + ) + .route( + INITIALIZE_POST_FAILURE_CONTROL_PATH, + post(arm_initialize_post_failure), + ) + .route( + INITIALIZED_NOTIFICATION_POST_FAILURE_CONTROL_PATH, + post(arm_initialized_notification_post_failure), + ) + .route( + "/.well-known/oauth-authorization-server/mcp", + get({ + move |headers: HeaderMap| async move { + let metadata_base = headers + .get(HOST) + .and_then(|value| value.to_str().ok()) + .map(|host| format!("http://{host}")) + .unwrap_or_else(|| format!("http://{actual_bind_addr}")); + #[expect(clippy::expect_used)] + Response::builder() + .status(StatusCode::OK) + .header(CONTENT_TYPE, "application/json") + .body(Body::from( + serde_json::to_vec(&json!({ + "issuer": format!("{metadata_base}/mcp"), + "authorization_endpoint": format!("{metadata_base}/oauth/authorize"), + "token_endpoint": format!("{metadata_base}/oauth/token"), + "scopes_supported": [""], + })).expect("failed to serialize metadata"), + )) + .expect("valid metadata response") + } + }), + ) + .route( + "/oauth/token", + post(|| async { + ( + StatusCode::BAD_REQUEST, + Json(json!({ + "error": "invalid_grant", + "error_description": "refresh token expired or revoked", + })), + ) + }), + ) + .nest_service( + "/mcp", + StreamableHttpService::new( + || Ok(TestToolServer::new()), + Arc::new(LocalSessionManager::default()), + StreamableHttpServerConfig::default(), + ), + ) + .layer(middleware::from_fn_with_state( + post_failure_state.clone(), + fail_mcp_post_when_armed, + )) + .with_state(post_failure_state); + + let router = if let Ok(token) = std::env::var("MCP_EXPECT_BEARER") { + let expected = Arc::new(format!("Bearer {token}")); + router.layer(middleware::from_fn_with_state(expected, require_bearer)) + } else { + router + }; + let router = if let Ok(token) = std::env::var("MCP_EXPECT_GATEWAY_BEARER") { + let expected = Arc::new(format!("Bearer {token}")); + router.layer(middleware::from_fn_with_state( + expected, + require_gateway_bearer, + )) + } else { + router + }; + + axum::serve(listener, router).await?; + task::yield_now().await; + Ok(()) +} + +impl ServerHandler for TestToolServer { + fn get_info(&self) -> ServerInfo { + ServerInfo::new( + ServerCapabilities::builder() + .enable_tools() + .enable_tool_list_changed() + .enable_resources() + .build(), + ) + } + + fn list_tools( + &self, + _request: Option, + _context: rmcp::service::RequestContext, + ) -> impl std::future::Future> + Send + '_ { + let tools = self.tools.clone(); + async move { Ok(ListToolsResult::with_all_items((*tools).clone())) } + } + + fn list_resources( + &self, + _request: Option, + _context: rmcp::service::RequestContext, + ) -> impl std::future::Future> + Send + '_ { + let resources = self.resources.clone(); + async move { Ok(ListResourcesResult::with_all_items((*resources).clone())) } + } + + async fn list_resource_templates( + &self, + _request: Option, + _context: rmcp::service::RequestContext, + ) -> Result { + Ok(ListResourceTemplatesResult::with_all_items( + (*self.resource_templates).clone(), + )) + } + + async fn read_resource( + &self, + ReadResourceRequestParams { uri, .. }: ReadResourceRequestParams, + _context: rmcp::service::RequestContext, + ) -> Result { + if uri == MEMO_URI { + Ok( + ReadResourceResult::new(vec![ResourceContents::TextResourceContents { + uri, + mime_type: Some("text/plain".to_string()), + text: Self::memo_text().to_string(), + meta: None, + }]) + .into(), + ) + } else { + Err(McpError::resource_not_found( + "resource_not_found", + Some(json!({ "uri": uri })), + )) + } + } + + async fn call_tool( + &self, + request: CallToolRequestParams, + _context: rmcp::service::RequestContext, + ) -> Result { + match request.name.as_ref() { + "echo" => { + let args: EchoArgs = match request.arguments { + Some(arguments) => serde_json::from_value(serde_json::Value::Object( + arguments.into_iter().collect(), + )) + .map_err(|err| McpError::invalid_params(err.to_string(), None))?, + None => { + return Err(McpError::invalid_params( + "missing arguments for echo tool", + None, + )); + } + }; + + let env_snapshot: HashMap = std::env::vars().collect(); + let structured_content = json!({ + "echo": format!("ECHOING: {}", args.message), + "env": env_snapshot.get("MCP_TEST_VALUE"), + }); + + let mut result = CallToolResult::success(Vec::new()); + result.structured_content = Some(structured_content); + Ok(result.into()) + } + other => Err(McpError::invalid_params( + format!("unknown tool: {other}"), + None, + )), + } + } +} + +impl TestToolServer { + fn new() -> Self { + let tools = vec![Self::echo_tool()]; + let resources = vec![Self::memo_resource()]; + let resource_templates = vec![Self::memo_template()]; + Self { + tools: Arc::new(tools), + resources: Arc::new(resources), + resource_templates: Arc::new(resource_templates), + } + } + + fn echo_tool() -> Tool { + #[expect(clippy::expect_used)] + let schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { + "message": { "type": "string" }, + "env_var": { "type": "string" } + }, + "required": ["message"], + "additionalProperties": false + })) + .expect("echo tool schema should deserialize"); + + let mut tool = Tool::new( + Cow::Borrowed("echo"), + Cow::Borrowed("Echo back the provided message and include environment data."), + Arc::new(schema), + ); + #[expect(clippy::expect_used)] + let output_schema: JsonObject = serde_json::from_value(json!({ + "type": "object", + "properties": { + "echo": { "type": "string" }, + "env": { + "anyOf": [ + { "type": "string" }, + { "type": "null" } + ] + } + }, + "required": ["echo", "env"], + "additionalProperties": false + })) + .expect("echo tool output schema should deserialize"); + tool.output_schema = Some(Arc::new(output_schema)); + tool.annotations = Some(ToolAnnotations::new().read_only(true)); + tool + } + + fn memo_resource() -> Resource { + Resource::new(MEMO_URI, "example-note") + .with_title("Example Note") + .with_description("A sample MCP resource exposed for integration tests.") + .with_mime_type("text/plain") + } + + fn memo_template() -> ResourceTemplate { + ResourceTemplate::new("memo://codex/{slug}", "codex-memo") + .with_title("Codex Memo") + .with_description("Template for memo://codex/{slug} resources used in tests.") + .with_mime_type("text/plain") + } + + fn memo_text() -> &'static str { + MEMO_CONTENT + } +} + +fn parse_bind_addr() -> Result> { + let default_addr = "127.0.0.1:3920"; + let bind_addr = std::env::var("MCP_STREAMABLE_HTTP_BIND_ADDR") + .or_else(|_| std::env::var("BIND_ADDR")) + .unwrap_or_else(|_| default_addr.to_string()); + Ok(bind_addr.parse()?) +} + +async fn require_bearer( + State(expected): State>, + request: Request, + next: Next, +) -> Result { + if request.uri().path().contains("/.well-known/") || request.uri().path() == "/oauth/token" { + return Ok(next.run(request).await); + } + if request + .headers() + .get(AUTHORIZATION) + .is_some_and(|value| value.as_bytes() == expected.as_bytes()) + { + Ok(next.run(request).await) + } else { + Err(StatusCode::UNAUTHORIZED) + } +} + +async fn require_gateway_bearer( + State(expected): State>, + request: Request, + next: Next, +) -> Result { + if !request.uri().path().starts_with("/mcp") { + return Ok(next.run(request).await); + } + if request + .headers() + .get("proxy-authorization") + .is_some_and(|value| value.as_bytes() == expected.as_bytes()) + { + Ok(next.run(request).await) + } else { + Err(StatusCode::UNAUTHORIZED) + } +} + +async fn arm_session_post_failure( + State(state): State, + Json(request): Json, +) -> Result { + arm_post_failure(state, request, ArmedFailureTarget::Session).await +} + +async fn arm_initialize_post_failure( + State(state): State, + Json(request): Json, +) -> Result { + arm_post_failure(state, request, ArmedFailureTarget::Initialize).await +} + +async fn arm_initialized_notification_post_failure( + State(state): State, + Json(request): Json, +) -> Result { + arm_post_failure(state, request, ArmedFailureTarget::InitializedNotification).await +} + +async fn arm_post_failure( + state: PostFailureState, + request: ArmSessionPostFailureRequest, + target: ArmedFailureTarget, +) -> Result { + let status = StatusCode::from_u16(request.status).map_err(|_| StatusCode::BAD_REQUEST)?; + let www_authenticate_headers = request + .www_authenticate_headers + .into_iter() + .map(|value| HeaderValue::from_str(&value).map_err(|_| StatusCode::BAD_REQUEST)) + .collect::, _>>()?; + let content_type = request + .content_type + .map(|value| HeaderValue::from_str(&value).map_err(|_| StatusCode::BAD_REQUEST)) + .transpose()?; + let armed_failure = if request.remaining == 0 { + None + } else { + Some(ArmedFailure { + target, + status, + remaining: request.remaining, + www_authenticate_headers, + content_type, + body: request.body, + }) + }; + *state.armed_failure.lock().await = armed_failure; + Ok(StatusCode::NO_CONTENT) +} + +async fn fail_mcp_post_when_armed( + State(state): State, + request: Request, + next: Next, +) -> Response { + if request.uri().path() != "/mcp" || request.method() != Method::POST { + return next.run(request).await; + } + let (parts, body) = request.into_parts(); + let body_bytes = match to_bytes(body, MAX_MCP_POST_BODY_BYTES).await { + Ok(body_bytes) => body_bytes, + Err(_) => { + let mut response = Response::new(Body::from("failed to read request body")); + *response.status_mut() = StatusCode::BAD_REQUEST; + return response; + } + }; + let has_session_id = parts.headers.contains_key(MCP_SESSION_ID_HEADER); + let mcp_method = request_mcp_method(&body_bytes); + + { + let mut armed_failure = state.armed_failure.lock().await; + if let Some(failure) = armed_failure.as_mut() + && failure.remaining > 0 + && match failure.target { + ArmedFailureTarget::Initialize => !has_session_id, + ArmedFailureTarget::InitializedNotification => { + has_session_id && mcp_method.as_deref() == Some("notifications/initialized") + } + ArmedFailureTarget::Session => { + has_session_id && mcp_method.as_deref() != Some("notifications/initialized") + } + } + { + failure.remaining -= 1; + let status = failure.status; + let www_authenticate_headers = failure.www_authenticate_headers.clone(); + let content_type = failure.content_type.clone(); + let body = failure + .body + .clone() + .unwrap_or_else(|| format!("forced session failure with status {status}")); + if failure.remaining == 0 { + *armed_failure = None; + } + let mut response = Response::new(Body::from(body)); + *response.status_mut() = status; + if let Some(content_type) = content_type { + response.headers_mut().insert(CONTENT_TYPE, content_type); + } + for www_authenticate_header in www_authenticate_headers { + response + .headers_mut() + .append(WWW_AUTHENTICATE, www_authenticate_header); + } + return response; + } + } + + next.run(Request::from_parts(parts, Body::from(body_bytes))) + .await +} + +fn request_mcp_method(body: &[u8]) -> Option { + serde_json::from_slice::(body) + .ok()? + .get("method")? + .as_str() + .map(ToString::to_string) +} diff --git a/codex-rs/rmcp-client/src/bounded_stdio_transport.rs b/codex-rs/rmcp-client/src/bounded_stdio_transport.rs new file mode 100644 index 0000000000000000000000000000000000000000..fdde9f698ef29e4a8c0183e4b32ab2d835df8e74 --- /dev/null +++ b/codex-rs/rmcp-client/src/bounded_stdio_transport.rs @@ -0,0 +1,132 @@ +//! Bounded, compatibility-aware transport for modern locally spawned MCP servers. + +use std::future::Future; +use std::io; +use std::sync::Arc; + +use memchr::memchr; +use rmcp::service::RoleClient; +use rmcp::service::RxJsonRpcMessage; +use rmcp::service::TxJsonRpcMessage; +use rmcp::transport::Transport; +use tokio::io::AsyncBufReadExt; +use tokio::io::AsyncWriteExt; +use tokio::io::BufReader; +use tokio::process::ChildStdin; +use tokio::process::ChildStdout; +use tokio::sync::Mutex; +use tracing::debug; +use tracing::warn; + +/// Match the existing executor stdio transport and the production MCP limit. +pub(crate) const MAX_MCP_STDIO_LINE_BYTES: usize = 8 * 1024 * 1024; + +pub(super) struct BoundedStdioTransport { + stdin: Arc>>, + stdout: BufReader, + pending_line: Vec, + program_name: String, +} + +impl BoundedStdioTransport { + pub(super) fn new(stdin: ChildStdin, stdout: ChildStdout, program_name: String) -> Self { + Self { + stdin: Arc::new(Mutex::new(Some(stdin))), + stdout: BufReader::new(stdout), + pending_line: Vec::new(), + program_name, + } + } + + async fn receive_message(&mut self) -> Option> { + loop { + let bytes = match self.stdout.fill_buf().await { + Ok(bytes) => bytes, + Err(error) => { + warn!( + "Failed to read MCP server stdout ({}): {error}", + self.program_name + ); + return None; + } + }; + + if bytes.is_empty() { + if self.pending_line.is_empty() { + return None; + } + return self.decode_pending_message(); + } + + let newline = memchr(b'\n', bytes); + let content_len = newline.unwrap_or(bytes.len()); + if content_len > MAX_MCP_STDIO_LINE_BYTES.saturating_sub(self.pending_line.len()) { + warn!( + "MCP stdio line exceeds {MAX_MCP_STDIO_LINE_BYTES} bytes ({}); closing transport", + self.program_name + ); + self.pending_line.clear(); + return None; + } + + self.pending_line.extend_from_slice(&bytes[..content_len]); + let consumed = content_len + usize::from(newline.is_some()); + self.stdout.consume(consumed); + + if newline.is_some() + && let Some(message) = self.decode_pending_message() + { + return Some(message); + } + } + } + + fn decode_pending_message(&mut self) -> Option> { + let line = std::mem::take(&mut self.pending_line); + let line = line.strip_suffix(b"\r").unwrap_or(&line); + match serde_json::from_slice(line) { + Ok(message) => Some(message), + Err(error) => { + debug!( + "Failed to parse local MCP server message ({}): {error}", + self.program_name + ); + None + } + } + } +} + +impl Transport for BoundedStdioTransport { + type Error = io::Error; + + #[expect( + clippy::await_holding_invalid_type, + reason = "complete JSON-RPC frames must hold the stdin lock across writes" + )] + fn send( + &mut self, + item: TxJsonRpcMessage, + ) -> impl Future> + Send + 'static { + let stdin = Arc::clone(&self.stdin); + async move { + let mut message = serde_json::to_vec(&item).map_err(io::Error::other)?; + message.push(b'\n'); + let mut guard = stdin.lock().await; + let stdin = guard + .as_mut() + .ok_or_else(|| io::Error::new(io::ErrorKind::BrokenPipe, "MCP stdin closed"))?; + stdin.write_all(&message).await?; + stdin.flush().await + } + } + + fn receive(&mut self) -> impl Future>> + Send { + self.receive_message() + } + + async fn close(&mut self) -> Result<(), Self::Error> { + self.stdin.lock().await.take(); + Ok(()) + } +} diff --git a/codex-rs/rmcp-client/src/elicitation_client_service.rs b/codex-rs/rmcp-client/src/elicitation_client_service.rs new file mode 100644 index 0000000000000000000000000000000000000000..497ca0652d0c82d800ac75bc29d0acfca686499a --- /dev/null +++ b/codex-rs/rmcp-client/src/elicitation_client_service.rs @@ -0,0 +1,699 @@ +use std::collections::HashMap; +use std::collections::HashSet; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; + +use codex_protocol::mcp::OPENAI_ELICITATION_EXTENSION_ID; + +use rmcp::RoleClient; +use rmcp::model::ClientInfo; +use rmcp::model::ClientResult; +use rmcp::model::CustomRequest; +use rmcp::model::CustomResult; +use rmcp::model::ElicitResult; +use rmcp::model::ElicitationAction; +use rmcp::model::MetaObject; +use rmcp::model::ProtocolVersion; +use rmcp::model::RequestId; +use rmcp::model::RequestMetaObject; +use rmcp::model::RequestParamsMeta; +use rmcp::model::ServerNotification; +use rmcp::model::ServerRequest; +use rmcp::service::NotificationContext; +use rmcp::service::RequestContext; +use rmcp::service::Service; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Map; +use serde_json::Value; +use tokio::sync::oneshot; + +use crate::logging_client_handler::LoggingClientHandler; +use crate::rmcp_client::Elicitation; +use crate::rmcp_client::ElicitationPauseState; +use crate::rmcp_client::ElicitationResponse; +use crate::rmcp_client::SendElicitation; + +const MCP_PROGRESS_TOKEN_META_KEY: &str = "progressToken"; +const MCP_ELICITATION_CREATE_METHOD: &str = "elicitation/create"; +const OPENAI_FORM_METHOD: &str = "openai/form"; +const OPENAI_ELICITATION_METHOD: &str = "openai/elicitation/create"; + +#[derive(Deserialize)] +#[serde(tag = "mode")] +enum OpenAiElicitationRequestParams { + #[serde(rename = "form")] + Form(OpenAiFormRequestParams), +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase")] +struct OpenAiFormRequestParams { + #[serde(rename = "_meta")] + meta: Option, + message: String, + requested_schema: Value, +} + +#[derive(Clone)] +pub(crate) struct ElicitationClientService { + handler: LoggingClientHandler, + supports_openai_form: bool, + supports_openai_elicitation_form: bool, + supports_user_verification: bool, + send_elicitation: Arc, + pause_state: ElicitationPauseState, + pending_verifications: Arc>, +} + +// A notification handler can run before its request handler. Never evict an early +// cancellation: after saturation, cancel new elicitations for this connection. +const MAX_EARLY_CANCELLATIONS: usize = 1024; + +#[derive(Default)] +struct VerificationCancellations { + pending: HashMap>, + early: HashSet, + saturated: bool, +} + +struct PendingVerification { + request_id: RequestId, + cancellations: Arc>, +} + +impl Drop for PendingVerification { + fn drop(&mut self) { + self.cancellations + .lock() + .unwrap_or_else(PoisonError::into_inner) + .pending + .remove(&self.request_id); + } +} + +impl ElicitationClientService { + pub(crate) fn new( + client_info: ClientInfo, + send_elicitation: SendElicitation, + pause_state: ElicitationPauseState, + ) -> Self { + let supports_openai_form = client_info + .capabilities + .extensions + .as_ref() + .is_some_and(|extensions| extensions.contains_key(OPENAI_FORM_METHOD)); + let supports_openai_elicitation_form = client_info + .capabilities + .extensions + .as_ref() + .and_then(|extensions| extensions.get(OPENAI_ELICITATION_EXTENSION_ID)) + .and_then(|settings| settings.get("form")) + .is_some_and(Value::is_object); + let send_elicitation = Arc::new(send_elicitation); + let supports_user_verification = client_info + .capabilities + .extensions + .as_ref() + .and_then(|extensions| extensions.get(OPENAI_ELICITATION_EXTENSION_ID)) + .and_then(|settings| settings.get("userVerification")) + .is_some_and(Value::is_object); + Self { + handler: LoggingClientHandler::new( + client_info, + clone_send_elicitation(Arc::clone(&send_elicitation)), + ), + supports_openai_form, + supports_openai_elicitation_form, + supports_user_verification, + send_elicitation, + pause_state, + pending_verifications: Arc::default(), + } + } + + async fn create_elicitation( + &self, + request: Elicitation, + context: RequestContext, + ) -> Result { + let RequestContext { id, meta, ct, .. } = context; + let request = restore_context_meta(request, meta); + let user_verification = matches!(&request, Elicitation::UserVerification { .. }); + let (cancel_tx, cancel_rx) = oneshot::channel(); + let _pending = { + let mut cancellations = self + .pending_verifications + .lock() + .unwrap_or_else(PoisonError::into_inner); + if cancellations.saturated || cancellations.early.remove(&id) { + return Ok(ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None, + }); + } + cancellations.pending.insert(id.clone(), cancel_tx); + PendingVerification { + request_id: id.clone(), + cancellations: Arc::clone(&self.pending_verifications), + } + }; + let _pause = self.pause_state.enter(); + let response = tokio::select! { + biased; + _ = ct.cancelled() => { + return Ok(ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None, + }); + } + _ = cancel_rx => { + return Ok(ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None, + }); + } + response = (self.send_elicitation)(id, request) => response, + } + .map_err(|err| rmcp::ErrorData::internal_error(err.to_string(), None))?; + Ok(if user_verification { + crate::user_verification::validate_response(response) + } else { + response + }) + } +} + +fn clone_send_elicitation(send_elicitation: Arc) -> SendElicitation { + Box::new(move |request_id, request| send_elicitation(request_id, request)) +} + +impl Service for ElicitationClientService { + async fn handle_request( + &self, + request: ServerRequest, + context: RequestContext, + ) -> Result { + match request { + ServerRequest::ElicitRequest(request) => { + let modern_session = context + .peer + .peer_info() + .is_some_and(|info| info.protocol_version >= ProtocolVersion::V_2026_07_28); + let response = self + .create_elicitation(Elicitation::Mcp(request.params), context) + .await?; + if modern_session { + Ok(ClientResult::ElicitResult(typed_elicitation_result( + response, + )?)) + } else { + Ok(ClientResult::CustomResult(elicitation_response_result( + response, + )?)) + } + } + ServerRequest::CustomRequest(request) + if request.method == MCP_ELICITATION_CREATE_METHOD => + { + let modern_session = context + .peer + .peer_info() + .is_some_and(|info| info.protocol_version >= ProtocolVersion::V_2026_07_28); + let response = self + .create_elicitation(custom_mcp_elicitation(request)?, context) + .await?; + if modern_session { + Ok(ClientResult::ElicitResult(typed_elicitation_result( + response, + )?)) + } else { + Ok(ClientResult::CustomResult(elicitation_response_result( + response, + )?)) + } + } + ServerRequest::CustomRequest(request) + if request.method == OPENAI_ELICITATION_METHOD + && self.supports_user_verification + && request + .params + .as_ref() + .and_then(|params| params.get("mode")) + .and_then(Value::as_str) + == Some(crate::user_verification::MODE) => + { + let response = self + .create_elicitation(crate::user_verification::parse_request(request)?, context) + .await?; + Ok(ClientResult::CustomResult(elicitation_response_result( + response, + )?)) + } + ServerRequest::CustomRequest(request) + if request.method == OPENAI_FORM_METHOD && self.supports_openai_form => + { + let response = self + .create_elicitation(openai_form_elicitation(request)?, context) + .await?; + Ok(ClientResult::CustomResult(elicitation_response_result( + response, + )?)) + } + ServerRequest::CustomRequest(request) + if request.method == OPENAI_ELICITATION_METHOD + && self.supports_openai_elicitation_form => + { + let response = self + .create_elicitation(openai_elicitation_form(request)?, context) + .await?; + Ok(ClientResult::CustomResult(elicitation_response_result( + response, + )?)) + } + ServerRequest::CustomRequest(request) + if request.method == OPENAI_ELICITATION_METHOD + && self.supports_user_verification => + { + Err(rmcp::ErrorData::invalid_params( + "invalid elicitation mode", + /*data*/ None, + )) + } + request => { + >::handle_request( + &self.handler, + request, + context, + ) + .await + } + } + } + + async fn handle_notification( + &self, + notification: ServerNotification, + context: NotificationContext, + ) -> Result<(), rmcp::ErrorData> { + if let ServerNotification::CancelledNotification(cancelled) = ¬ification + && let Some(request_id) = cancelled.params.request_id.as_ref() + { + let mut cancellations = self + .pending_verifications + .lock() + .unwrap_or_else(PoisonError::into_inner); + if let Some(cancel) = cancellations.pending.remove(request_id) { + let _ = cancel.send(()); + } else if !cancellations.saturated && !cancellations.early.contains(request_id) { + // Bound both the number and size of retained request IDs. + if cancellations.early.len() == MAX_EARLY_CANCELLATIONS + || matches!(request_id, RequestId::String(id) if id.len() > 1024) + { + cancellations.saturated = true; + cancellations.early.clear(); + } else { + cancellations.early.insert(request_id.clone()); + } + } + } + >::handle_notification( + &self.handler, + notification, + context, + ) + .await + } + + fn get_info(&self) -> ClientInfo { + >::get_info(&self.handler) + } +} + +fn custom_mcp_elicitation(request: CustomRequest) -> Result { + let raw_params = request + .params + .ok_or_else(|| rmcp::ErrorData::invalid_params("missing params", None))?; + let params: rmcp::model::ElicitRequestParams = serde_json::from_value(raw_params) + .map_err(|err| rmcp::ErrorData::invalid_params(err.to_string(), None))?; + Ok(Elicitation::Mcp(params)) +} + +fn openai_form_elicitation(request: CustomRequest) -> Result { + let params = request + .params_as::() + .map_err(|err| rmcp::ErrorData::invalid_params(err.to_string(), None))? + .ok_or_else(|| rmcp::ErrorData::invalid_params("missing params", None))?; + Ok(Elicitation::OpenAiForm { + meta: params.meta, + message: params.message, + requested_schema: params.requested_schema, + }) +} + +pub(crate) fn openai_elicitation_form( + request: CustomRequest, +) -> Result { + let params = request + .params_as::() + .map_err(|err| rmcp::ErrorData::invalid_params(err.to_string(), /*data*/ None))? + .ok_or_else(|| rmcp::ErrorData::invalid_params("missing params", /*data*/ None))?; + let OpenAiElicitationRequestParams::Form(params) = params; + Ok(Elicitation::OpenAiElicitationForm { + meta: params.meta, + message: params.message, + requested_schema: params.requested_schema, + }) +} + +fn restore_context_meta( + mut request: Elicitation, + mut context_meta: RequestMetaObject, +) -> Elicitation { + // RMCP lifts JSON-RPC `_meta` into RequestContext before invoking services. + context_meta.remove(MCP_PROGRESS_TOKEN_META_KEY); + if context_meta.is_empty() { + return request; + } + + match &mut request { + Elicitation::UserVerification { .. } => {} + Elicitation::Mcp(request) => request + .meta_mut() + .get_or_insert_with(RequestMetaObject::new) + .extend(context_meta), + Elicitation::OpenAiForm { meta, .. } | Elicitation::OpenAiElicitationForm { meta, .. } => { + let meta = meta + .get_or_insert_with(|| Value::Object(Map::new())) + .as_object_mut(); + if let Some(meta) = meta { + meta.extend(context_meta.0.0); + } + } + } + request +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +struct CreateElicitationResultWithMeta { + action: ElicitationAction, + #[serde(skip_serializing_if = "Option::is_none")] + content: Option, + #[serde(rename = "_meta", skip_serializing_if = "Option::is_none")] + meta: Option, +} + +fn elicitation_response_result( + response: ElicitationResponse, +) -> Result { + let ElicitationResponse { + action, + content, + meta, + } = response; + let result = CreateElicitationResultWithMeta { + action, + content, + meta, + }; + + serde_json::to_value(result) + .map(CustomResult) + .map_err(|err| rmcp::ErrorData::internal_error(err.to_string(), None)) +} + +fn typed_elicitation_result( + response: ElicitationResponse, +) -> Result { + let ElicitationResponse { + action, + content, + meta, + } = response; + let mut result = ElicitResult::new(action); + result.content = content; + result.meta = match meta { + None => None, + Some(Value::Object(meta)) => Some(MetaObject::from(meta)), + Some(meta) => { + return Err(rmcp::ErrorData::invalid_params( + format!("MCP elicitation response _meta must be an object, got {meta}"), + None, + )); + } + }; + Ok(result) +} + +#[cfg(test)] +mod tests { + use pretty_assertions::assert_eq; + use rmcp::model::BooleanSchema; + use rmcp::model::ElicitRequestParams; + use rmcp::model::ElicitationSchema; + use rmcp::model::PrimitiveSchemaDefinition; + use serde_json::Value; + use serde_json::json; + + use super::*; + + #[test] + fn restore_context_meta_adds_elicitation_meta_and_removes_progress_token() { + let request = restore_context_meta( + Elicitation::Mcp(form_request(/*meta*/ None)), + meta(json!({ + "progressToken": "progress-token", + "persist": ["session", "always"], + })), + ); + + assert_eq!( + request, + Elicitation::Mcp(form_request(Some(meta(json!({ + "persist": ["session", "always"], + }))))) + ); + } + + #[test] + fn legacy_sep1034_elicitation_without_mode_preserves_schema_defaults() { + let request = json!({ + "method": "elicitation/create", + "params": { + "message": "Confirm the default values", + "requestedSchema": { + "type": "object", + "properties": { + "name": {"type": "string", "default": "John Doe"}, + "age": {"type": "integer", "default": 30}, + "score": {"type": "number", "default": 95.5}, + "status": { + "type": "string", + "enum": ["active", "inactive"], + "default": "active", + }, + "verified": {"type": "boolean", "default": true}, + }, + "required": [], + }, + }, + }); + + let request = serde_json::from_value::(request) + .expect("legacy form elicitations must deserialize without a mode"); + let ServerRequest::ElicitRequest(request) = request else { + panic!("legacy elicitation/create must dispatch to the typed handler"); + }; + let ElicitRequestParams::FormElicitationParams { + requested_schema, .. + } = request.params + else { + panic!("an omitted legacy elicitation mode must default to form"); + }; + + assert_eq!( + serde_json::to_value(requested_schema) + .expect("legacy schema defaults must remain serializable"), + json!({ + "type": "object", + "properties": { + "name": {"type": "string", "default": "John Doe"}, + "age": {"type": "integer", "default": 30}, + "score": {"type": "number", "default": 95.5}, + "status": { + "type": "string", + "enum": ["active", "inactive"], + "default": "active", + }, + "verified": {"type": "boolean", "default": true}, + }, + "required": [], + }) + ); + } + + #[test] + fn parses_legacy_custom_elicitation_without_mode() { + let request = CustomRequest::new( + MCP_ELICITATION_CREATE_METHOD, + Some(json!({ + "message": "Confirm?", + "requestedSchema": { + "type": "object", + "properties": { + "confirmed": {"type": "boolean"}, + "age": {"type": "integer", "minimum": 1, "maximum": 99, "default": 30}, + }, + "required": ["confirmed"], + }, + })), + ); + let Elicitation::Mcp(rmcp::model::ElicitRequestParams::FormElicitationParams { + requested_schema, + .. + }) = custom_mcp_elicitation(request) + .expect("legacy custom elicitation parameters must deserialize") + else { + panic!("omitted legacy elicitation mode must default to form"); + }; + + assert_eq!( + serde_json::to_value(requested_schema).expect("schema must serialize"), + json!({ + "type": "object", + "properties": { + "confirmed": {"type": "boolean"}, + "age": {"type": "integer", "minimum": 1, "maximum": 99, "default": 30}, + }, + "required": ["confirmed"], + }) + ); + } + + #[test] + fn parses_openai_form_custom_requests() { + let elicitation = openai_form_elicitation(CustomRequest::new( + OPENAI_FORM_METHOD, + Some(json!({ + "message": "Select a template", + "requestedSchema": { + "type": "object", + "properties": { + "template": { + "type": "openai/imagePicker", + "items": [{ + "id": "monthly-review", + "title": "Monthly review", + "image": "data:image/svg+xml;base64,PHN2ZyB4bWxucz0iaHR0cDovL3d3dy53My5vcmcvMjAwMC9zdmciLz4=" + }] + } + } + } + })), + )) + .expect("valid openai/form request"); + + assert_eq!( + elicitation, + Elicitation::OpenAiForm { + meta: None, + message: "Select a template".to_string(), + requested_schema: json!({ + "type": "object", + "properties": { + "template": { + "type": "openai/imagePicker", + "items": [{ + "id": "monthly-review", + "title": "Monthly review", + "image": "data:image/svg+xml;base64,PHN2ZyB4bWxucz0iaHR0cDovL3d3dy53My5vcmcvMjAwMC9zdmciLz4=" + }] + } + } + }), + } + ); + } + + #[test] + fn elicitation_response_result_serializes_response_meta() { + let result = rmcp::model::ClientResult::CustomResult( + elicitation_response_result(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({ "confirmed": true })), + meta: Some(json!({ "persist": "always" })), + }) + .expect("elicitation response should serialize"), + ); + + assert_eq!( + serde_json::to_value(result).expect("client result should serialize"), + json!({ + "action": "accept", + "content": { "confirmed": true }, + "_meta": { "persist": "always" }, + }) + ); + } + + #[test] + fn typed_elicitation_result_preserves_response_meta() { + let result = typed_elicitation_result(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({ "confirmed": true })), + meta: Some(json!({ "persist": "always" })), + }) + .expect("modern elicitation response should serialize"); + + assert_eq!( + serde_json::to_value(result).expect("typed elicitation result should serialize"), + json!({ + "action": "accept", + "content": { "confirmed": true }, + "_meta": { "persist": "always" }, + }) + ); + } + + #[test] + fn typed_elicitation_result_rejects_non_object_meta() { + let error = typed_elicitation_result(ElicitationResponse { + action: ElicitationAction::Accept, + content: None, + meta: Some(json!(["invalid"])), + }) + .expect_err("modern elicitation metadata must be an object"); + + assert_eq!(error.code, rmcp::model::ErrorCode::INVALID_PARAMS); + } + + fn form_request(meta: Option) -> ElicitRequestParams { + ElicitRequestParams::FormElicitationParams { + meta, + message: "Confirm?".to_string(), + requested_schema: ElicitationSchema::builder() + .required_property( + "confirmed", + PrimitiveSchemaDefinition::Boolean(BooleanSchema::new()), + ) + .build() + .expect("schema should build"), + } + } + + fn meta(value: Value) -> RequestMetaObject { + let Value::Object(map) = value else { + panic!("meta must be an object"); + }; + RequestMetaObject::from(map) + } +} + +#[cfg(test)] +#[path = "user_verification_dispatch_tests.rs"] +mod user_verification_dispatch_tests; diff --git a/codex-rs/rmcp-client/src/ema_auth_policy.rs b/codex-rs/rmcp-client/src/ema_auth_policy.rs new file mode 100644 index 0000000000000000000000000000000000000000..42057b1f78733dac3718355246ba44d2316b5bfe --- /dev/null +++ b/codex-rs/rmcp-client/src/ema_auth_policy.rs @@ -0,0 +1,134 @@ +//! Credential-destination policy for enterprise MCP OAuth. + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use serde_json::Value; +use url::Host; +use url::Url; + +/// A sanitized enterprise-auth failure that callers may handle without parsing text. +#[derive(Debug, PartialEq, Eq, thiserror::Error)] +pub enum EmaAuthFailure { + #[error("invalid_grant")] + InvalidGrant { grant_source: EmaInvalidGrantSource }, + #[error("insufficient_user_authentication")] + InsufficientUserAuthentication, + #[error("enterprise identity requires authentication")] + ReauthenticationRequired, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum EmaInvalidGrantSource { + EnterpriseIdentity, + ResourceAuthorization, +} + +pub(crate) fn ema_reauthentication_required(message: &'static str) -> anyhow::Error { + anyhow::Error::new(EmaAuthFailure::ReauthenticationRequired).context(message) +} + +pub(crate) fn safe_oauth_error_code(code: Option<&str>) -> &str { + code.filter(|code| { + matches!( + *code, + "invalid_request" + | "invalid_client" + | "invalid_grant" + | "invalid_scope" + | "invalid_target" + | "unauthorized_client" + | "unsupported_grant_type" + | "access_denied" + | "temporarily_unavailable" + | "server_error" + | "insufficient_user_authentication" + ) + }) + .unwrap_or("OAuth token request rejected") +} + +pub(crate) fn validate_ema_public_client_auth( + advertised_methods: Option<&Value>, + issuer_description: &str, +) -> Result<()> { + let advertised_methods = advertised_methods.ok_or_else(|| { + anyhow!( + "{issuer_description} does not explicitly advertise public-client token endpoint authentication" + ) + })?; + let methods = advertised_methods.as_array().ok_or_else(|| { + anyhow!("{issuer_description} advertised malformed token endpoint authentication methods") + })?; + if !methods.iter().any(|method| method.as_str() == Some("none")) { + bail!("{issuer_description} does not support public-client token endpoint authentication"); + } + Ok(()) +} + +pub(crate) fn validate_ema_oauth_endpoint(endpoint: &str, description: &str) -> Result<()> { + let url = Url::parse(endpoint).with_context(|| format!("{description} is not a valid URL"))?; + validate_credential_destination(&url, description) +} + +fn validate_credential_destination(url: &Url, description: &str) -> Result<()> { + let loopback = match url.host() { + Some(Host::Domain(host)) => host.eq_ignore_ascii_case("localhost"), + Some(Host::Ipv4(address)) => address.is_loopback(), + Some(Host::Ipv6(address)) => address.is_loopback(), + None => false, + }; + if url.scheme() != "https" && !(url.scheme() == "http" && loopback) { + bail!("{description} must use HTTPS or an HTTP loopback address"); + } + if !url.username().is_empty() || url.password().is_some() || url.fragment().is_some() { + bail!("{description} contains disallowed credentials or a URL fragment"); + } + Ok(()) +} + +pub(crate) fn advertised_capability( + value: Option<&Value>, + expected: &str, + description: &str, +) -> Result> { + value + .map(|value| { + let values = value + .as_array() + .ok_or_else(|| anyhow!("{description} is malformed"))?; + Ok(values.iter().any(|value| value.as_str() == Some(expected))) + }) + .transpose() +} + +/// Resource indicators must describe the configured MCP origin, query and path. +pub fn validate_ema_auth_resource(server_url: &str, resource: Option<&str>) -> Result<()> { + let server = Url::parse(server_url).context("enterprise MCP server URL is invalid")?; + validate_credential_destination(&server, "enterprise MCP server URL")?; + let Some(resource) = resource.filter(|resource| !resource.trim().is_empty()) else { + return Ok(()); + }; + let resource = Url::parse(resource).context("enterprise MCP resource indicator is invalid")?; + validate_credential_destination(&resource, "enterprise MCP resource indicator")?; + if resource.origin() != server.origin() || resource.query() != server.query() { + bail!( + "enterprise MCP resource indicator must match the configured MCP server origin and query" + ); + } + let resource_path = resource.path().trim_end_matches('/'); + let server_path = server.path().trim_end_matches('/'); + if server_path != resource_path + && !server_path + .strip_prefix(resource_path) + .is_some_and(|suffix| suffix.starts_with('/')) + { + bail!("enterprise MCP resource indicator path must contain the configured MCP server path"); + } + Ok(()) +} + +#[cfg(test)] +#[path = "ema_auth_policy_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/ema_auth_policy_tests.rs b/codex-rs/rmcp-client/src/ema_auth_policy_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..3bdaed367dab65f12445ff2d13abf85719f1d758 --- /dev/null +++ b/codex-rs/rmcp-client/src/ema_auth_policy_tests.rs @@ -0,0 +1,51 @@ +use pretty_assertions::assert_eq; +use serde_json::json; + +use super::*; + +#[test] +fn resource_origin_query_and_path_are_bound() { + let server = "https://mcp.example/enterprise/tools?tenant=one"; + for (resource, valid) in [ + ("https://mcp.example/enterprise?tenant=one", true), + ("https://other.example/enterprise?tenant=one", false), + ("https://mcp.example/enterprise-admin?tenant=one", false), + ("https://mcp.example/enterprise?tenant=two", false), + ("https://mcp.example/enterprise", false), + ("http://localhost:4000/enterprise?tenant=one", false), + ] { + assert_eq!( + validate_ema_auth_resource(server, Some(resource)).is_ok(), + valid, + "{resource}" + ); + } + for (server, valid) in [ + ("https://mcp.example/mcp", true), + ("http://localhost:4000/mcp", true), + ("http://127.0.0.1:4000/mcp", true), + ("http://mcp.example/mcp", false), + ] { + assert_eq!( + validate_ema_auth_resource(server, /*resource*/ None).is_ok(), + valid, + "{server}" + ); + } +} + +#[test] +fn public_clients_require_an_explicitly_advertised_auth_method() { + for (advertised, valid) in [ + (None, false), + (Some(json!(["none"])), true), + (Some(json!(["client_secret_basic", "none"])), true), + (Some(json!(["private_key_jwt"])), false), + (Some(json!("none")), false), + ] { + assert_eq!( + validate_ema_public_client_auth(advertised.as_ref(), "IdP").is_ok(), + valid + ); + } +} diff --git a/codex-rs/rmcp-client/src/ema_claims.rs b/codex-rs/rmcp-client/src/ema_claims.rs new file mode 100644 index 0000000000000000000000000000000000000000..886f148a32a45d80c6cd61b049721d8129d6501b --- /dev/null +++ b/codex-rs/rmcp-client/src/ema_claims.rs @@ -0,0 +1,250 @@ +//! Validate token routing and signed authorization before forwarding an ID-JAG. + +use std::collections::HashSet; +use std::time::Duration; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use base64::Engine; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use serde::Deserialize; +use serde::de::DeserializeOwned; + +use crate::ema_exchange::EmaAccessToken; + +pub(crate) const ID_JAG_TOKEN_TYPE: &str = "urn:ietf:params:oauth:token-type:id-jag"; + +#[derive(Deserialize)] +struct JwtHeader { + alg: String, + typ: Option, +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum OAuthResource { + Single(String), + Multiple(Vec), +} + +impl OAuthResource { + fn is_exact(&self, expected: &str) -> bool { + match self { + Self::Single(value) => value == expected, + Self::Multiple(values) => values.as_slice() == [expected], + } + } +} + +fn signed_jwt(token: &str) -> Result<(JwtHeader, T)> { + let mut parts = token.split('.'); + let (Some(header), Some(payload), Some(signature), None) = + (parts.next(), parts.next(), parts.next(), parts.next()) + else { + bail!("identity assertion is not a compact signed JWT"); + }; + if header.is_empty() || payload.is_empty() || signature.is_empty() { + bail!("identity assertion contains an empty JWT segment"); + } + let header: JwtHeader = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(header)?) + .map_err(|_| anyhow!("invalid identity assertion JWT header"))?; + if header.alg.trim().is_empty() || header.alg.eq_ignore_ascii_case("none") { + bail!("identity assertion is unsigned"); + } + let claims = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(payload)?) + .map_err(|_| anyhow!("invalid identity assertion JWT claims"))?; + Ok((header, claims)) +} + +#[derive(Deserialize)] +pub(crate) struct OidcClaims { + iss: String, + sub: String, + aud: OAuthResource, + azp: Option, + exp: u64, +} + +pub(crate) fn oidc_identity( + assertion: &str, + expected_issuer: &str, + expected_audience: &str, +) -> Result { + let (_, claims): (_, OidcClaims) = signed_jwt(assertion)?; + if claims.iss != expected_issuer || claims.sub.trim().is_empty() { + bail!("OIDC identity assertion issuer or subject does not match the enterprise IdP"); + } + let (audience_matches, multiple_audiences) = match &claims.aud { + OAuthResource::Single(value) => (value == expected_audience, false), + OAuthResource::Multiple(values) => ( + values.iter().any(|value| value == expected_audience), + values.len() > 1, + ), + }; + if !audience_matches + || claims + .azp + .as_deref() + .is_some_and(|party| party != expected_audience) + || multiple_audiences && claims.azp.as_deref() != Some(expected_audience) + { + bail!("OIDC identity assertion audience or authorized party does not match the IdP client"); + } + Ok(claims) +} + +pub fn validate_oidc_identity_assertion( + assertion: &str, + expected_issuer: &str, + expected_audience: &str, +) -> Result<()> { + let claims = oidc_identity(assertion, expected_issuer, expected_audience)?; + let now = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs(); + if claims.exp <= now { + bail!("OIDC identity assertion is expired"); + } + Ok(()) +} + +#[derive(Deserialize)] +struct IdJagClaims { + iss: String, + sub: String, + aud: OAuthResource, + client_id: String, + jti: String, + exp: u64, + iat: u64, + resource: OAuthResource, + scope: Option, +} + +pub(crate) struct IdJagBinding<'a> { + pub issuer: &'a str, + pub audience: &'a str, + pub client_id: &'a str, + pub resource: &'a str, + /// Empty means the scope parameter was omitted, not an empty authorization ceiling. + pub requested_scopes: &'a HashSet<&'a str>, +} + +#[derive(Deserialize)] +pub(crate) struct IdJagResponse { + pub access_token: String, + issued_token_type: String, + token_type: String, + resource: Option, + scope: Option, + refresh_token: Option, +} + +impl IdJagResponse { + pub(crate) fn validate(&self, binding: IdJagBinding<'_>) -> Result> { + if self.issued_token_type != ID_JAG_TOKEN_TYPE + || self.token_type != "N_A" + || self.refresh_token.is_some() + { + bail!("enterprise IdP returned an unsupported ID-JAG token type or refresh token"); + } + let (header, claims): (_, IdJagClaims) = signed_jwt(&self.access_token)?; + if header.typ.as_deref() != Some("oauth-id-jag+jwt") + || claims.iss != binding.issuer + || !claims.aud.is_exact(binding.audience) + || claims.client_id != binding.client_id + || claims.sub.trim().is_empty() + || claims.jti.trim().is_empty() + { + bail!("ID-JAG type, issuer, audience, client, subject, or JWT ID is invalid"); + } + let now = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs(); + if claims.exp <= now || claims.iat > now.saturating_add(60) { + bail!("enterprise IdP returned an expired or future-issued ID-JAG"); + } + if !claims.resource.is_exact(binding.resource) + || self + .resource + .as_ref() + .is_some_and(|value| !value.is_exact(binding.resource)) + { + bail!("ID-JAG must authorize exactly the configured MCP resource"); + } + let granted = match claims.scope.as_deref() { + Some(scope) => parse_scope(scope)?, + None if binding.requested_scopes.is_empty() => HashSet::new(), + None => bail!("ID-JAG is missing the requested scope authorization"), + }; + if !binding.requested_scopes.is_empty() && !granted.is_subset(binding.requested_scopes) { + bail!("ID-JAG contains a scope outside the enterprise authorization request"); + } + match self.scope.as_deref() { + Some(scope) if parse_scope(scope)? != granted => { + bail!("enterprise IdP token response scope does not match the signed ID-JAG") + } + None if !binding.requested_scopes.is_empty() + && granted != *binding.requested_scopes => + { + bail!("enterprise IdP token response omitted its narrowed scope") + } + _ => {} + } + Ok(granted.into_iter().map(str::to_string).collect()) + } +} + +fn parse_scope(scope: &str) -> Result> { + let scopes = scope.split_ascii_whitespace().collect::>(); + if scopes.is_empty() || scopes.len() != scope.split_ascii_whitespace().count() { + bail!("enterprise authorization contains malformed or duplicate scopes"); + } + Ok(scopes) +} + +#[derive(Deserialize)] +pub(crate) struct McpAccessTokenResponse { + access_token: String, + token_type: String, + expires_in: Option, + resource: Option, + scope: Option, + refresh_token: Option, +} + +impl McpAccessTokenResponse { + pub(crate) fn validate( + self, + resource: &str, + id_jag_scopes: &HashSet, + ) -> Result { + if !self.token_type.eq_ignore_ascii_case("bearer") || self.access_token.trim().is_empty() { + bail!("MCP authorization server returned an invalid bearer token"); + } + if self.refresh_token.is_some() || self.expires_in == Some(0) { + bail!("MCP authorization server returned a refresh token or zero token lifetime"); + } + // The stable EMA response does not require the resource to be echoed; + // when present, it must agree with the resource bound in the ID-JAG. + if self + .resource + .as_ref() + .is_some_and(|returned| !returned.is_exact(resource)) + { + bail!("MCP access token must authorize exactly the configured MCP resource"); + } + // RFC 6749 defines an omitted scope as unchanged from the request. Here + // that authority is the scope carried by the validated ID-JAG. + if let Some(scope) = self.scope.as_deref() + && !parse_scope(scope)? + .iter() + .all(|scope| id_jag_scopes.contains(*scope)) + { + bail!("MCP authorization server granted a scope outside the ID-JAG authorization"); + } + Ok(EmaAccessToken { + access_token: self.access_token, + expires_in: self.expires_in.map(Duration::from_secs), + }) + } +} diff --git a/codex-rs/rmcp-client/src/ema_exchange.rs b/codex-rs/rmcp-client/src/ema_exchange.rs new file mode 100644 index 0000000000000000000000000000000000000000..7b6131653a9727c4a667c3a086557199bd33f595 --- /dev/null +++ b/codex-rs/rmcp-client/src/ema_exchange.rs @@ -0,0 +1,219 @@ +//! Non-interactive ID-JAG exchange against explicitly trusted OAuth endpoints. + +use std::collections::HashSet; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use codex_exec_server::HttpClient; +use rmcp::transport::auth::OAuthHttpRedirectPolicy; +use serde::Deserialize; +use serde::de::DeserializeOwned; + +use crate::ema_auth_policy::EmaAuthFailure; +use crate::ema_auth_policy::EmaInvalidGrantSource; +use crate::ema_auth_policy::safe_oauth_error_code; +use crate::ema_auth_policy::validate_ema_oauth_endpoint; +use crate::ema_claims::ID_JAG_TOKEN_TYPE; +use crate::ema_claims::IdJagBinding; +use crate::ema_claims::IdJagResponse; +use crate::ema_claims::McpAccessTokenResponse; +use crate::http_client_adapter::StreamableHttpRedirectMode; +use crate::oauth_http_client::OAuthHttpClientAdapter; +use crate::utils::build_default_headers; + +pub(crate) const TOKEN_EXCHANGE_GRANT_TYPE: &str = + "urn:ietf:params:oauth:grant-type:token-exchange"; +pub(crate) const JWT_BEARER_GRANT_TYPE: &str = "urn:ietf:params:oauth:grant-type:jwt-bearer"; + +/// A resource-bound bearer and its server-reported lifetime, with redacted diagnostics. +#[derive(Clone, PartialEq, Eq)] +pub struct EmaAccessToken { + pub access_token: String, + pub expires_in: Option, +} + +impl std::fmt::Debug for EmaAccessToken { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("EmaAccessToken") + .field("access_token", &"[REDACTED]") + .field("expires_in", &self.expires_in) + .finish() + } +} + +/// The caller supplies trusted authorization-server metadata and an IdP credential. +/// This primitive does not perform resource discovery or interactive login. +pub struct EmaIdJagExchangeRequest<'a> { + pub resource: &'a str, + pub scopes: &'a [String], + pub mcp_client_id: &'a str, + pub authorization_server_issuer: &'a str, + pub authorization_server_token_endpoint: &'a str, + pub idp_token_endpoint: &'a str, + pub idp_issuer: &'a str, + pub idp_client_id: &'a str, + pub refresh_token: String, + pub idp_http_client: Arc, + pub resource_http_client: Arc, +} + +/// Exchanges an enterprise IdP credential for a resource-bound MCP bearer token. +pub async fn exchange_id_jag(request: EmaIdJagExchangeRequest<'_>) -> Result { + for (endpoint, description) in [ + (request.resource, "enterprise MCP resource"), + ( + request.authorization_server_issuer, + "MCP authorization server issuer", + ), + ( + request.authorization_server_token_endpoint, + "MCP token endpoint", + ), + (request.idp_issuer, "enterprise IdP issuer"), + (request.idp_token_endpoint, "enterprise IdP token endpoint"), + ] { + validate_ema_oauth_endpoint(endpoint, description)?; + } + if request.authorization_server_issuer == request.idp_issuer { + bail!("enterprise IdP and MCP authorization server issuers must be different for ID-JAG"); + } + if request.mcp_client_id.trim().is_empty() || request.idp_client_id.trim().is_empty() { + bail!("enterprise authorization requires the registered IdP and MCP client IDs"); + } + if request.refresh_token.trim().is_empty() { + bail!("enterprise IdP refresh token must not be empty"); + } + let requested_scopes = request + .scopes + .iter() + .map(String::as_str) + .collect::>(); + if requested_scopes.len() != request.scopes.len() + || request + .scopes + .iter() + .any(|scope| scope.is_empty() || scope.chars().any(char::is_whitespace)) + { + bail!("enterprise MCP authorization scopes must be distinct, non-empty scope tokens"); + } + let scope = (!request.scopes.is_empty()).then(|| request.scopes.join(" ")); + let mut params = vec![ + ("grant_type", TOKEN_EXCHANGE_GRANT_TYPE), + ("requested_token_type", ID_JAG_TOKEN_TYPE), + ("audience", request.authorization_server_issuer), + ("resource", request.resource), + ("subject_token", request.refresh_token.as_str()), + ( + "subject_token_type", + "urn:ietf:params:oauth:token-type:refresh_token", + ), + ]; + if let Some(scope) = scope.as_deref() { + params.push(("scope", scope)); + } + let id_jag: IdJagResponse = post_form( + &request.idp_http_client, + request.idp_token_endpoint, + ¶ms, + request.idp_client_id, + EmaInvalidGrantSource::EnterpriseIdentity, + "enterprise IdP ID-JAG exchange", + ) + .await?; + let granted_scopes = id_jag.validate(IdJagBinding { + issuer: request.idp_issuer, + audience: request.authorization_server_issuer, + client_id: request.mcp_client_id, + resource: request.resource, + requested_scopes: &requested_scopes, + })?; + // Only the signed assertion carries authority to the Resource AS. Repeating + // the requested resource or scopes could undo enterprise policy narrowing. + let access_token: McpAccessTokenResponse = post_form( + &request.resource_http_client, + request.authorization_server_token_endpoint, + &[ + ("grant_type", JWT_BEARER_GRANT_TYPE), + ("assertion", id_jag.access_token.as_str()), + ], + request.mcp_client_id, + EmaInvalidGrantSource::ResourceAuthorization, + "MCP JWT bearer exchange", + ) + .await?; + access_token.validate(request.resource, &granted_scopes) +} + +#[derive(Deserialize)] +struct OAuthErrorResponse { + error: Option, +} + +pub(crate) async fn post_form( + http_client: &Arc, + url: &str, + params: &[(&str, &str)], + client_id: &str, + invalid_grant_source: EmaInvalidGrantSource, + operation: &str, +) -> Result { + let client = OAuthHttpClientAdapter::new_with_redirect_mode( + Arc::clone(http_client), + build_default_headers(/*http_headers*/ None, /*env_http_headers*/ None)?, + url, + /*has_configured_headers*/ false, + StreamableHttpRedirectMode::Legacy, + )?; + let body = { + let mut form = url::form_urlencoded::Serializer::new(String::new()); + form.extend_pairs(params.iter().copied()); + form.append_pair("client_id", client_id); + form.finish().into_bytes() + }; + let builder = oauth2::http::Request::builder() + .method("POST") + .uri(url) + .header("content-type", "application/x-www-form-urlencoded") + .header("accept", "application/json"); + let response = client + .execute_request( + builder.body(body)?, + OAuthHttpRedirectPolicy::Stop, + Some(Duration::from_secs(30)), + ) + .await + .map_err(|error| anyhow!("{operation} request failed: {error}"))?; + if !response.status().is_success() { + let error = serde_json::from_slice::(response.body()).ok(); + // Provider-controlled text may reflect the submitted assertion or secret. + // Only known OAuth codes may reach callers. + let code = safe_oauth_error_code(error.as_ref().and_then(|error| error.error.as_deref())); + if code == "invalid_grant" { + return Err(anyhow::Error::new(EmaAuthFailure::InvalidGrant { + grant_source: invalid_grant_source, + }) + .context(format!( + "{operation} returned HTTP {}: invalid_grant", + response.status() + ))); + } + if code == "insufficient_user_authentication" { + return Err( + anyhow::Error::new(EmaAuthFailure::InsufficientUserAuthentication).context( + format!("{operation} returned HTTP {}: {code}", response.status()), + ), + ); + } + bail!("{operation} returned HTTP {}: {code}", response.status()); + } + serde_json::from_slice(response.body()) + .map_err(|_| anyhow!("failed to parse {operation} response")) +} + +#[cfg(test)] +#[path = "ema_exchange_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/ema_exchange_tests.rs b/codex-rs/rmcp-client/src/ema_exchange_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..46717907108c60e540e660028afa457341f0ebf9 --- /dev/null +++ b/codex-rs/rmcp-client/src/ema_exchange_tests.rs @@ -0,0 +1,455 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +use base64::Engine; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use codex_exec_server::RouteAwareHttpClient; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use serde_json::Value; +use serde_json::json; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::*; +use crate::ema_claims::validate_oidc_identity_assertion; + +fn http_client() -> Arc { + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))) +} + +fn unique_form_fields(body: &[u8]) -> HashMap { + let pairs = url::form_urlencoded::parse(body) + .into_owned() + .collect::>(); + let fields = pairs.iter().cloned().collect::>(); + assert_eq!( + pairs.len(), + fields.len(), + "OAuth form must not contain duplicate fields" + ); + fields +} + +fn jwt(claims: &Value) -> String { + format!( + "{}.{}.signature", + URL_SAFE_NO_PAD.encode(br#"{"alg":"ES256","typ":"oauth-id-jag+jwt"}"#), + URL_SAFE_NO_PAD.encode(serde_json::to_vec(claims).expect("serialize claims")) + ) +} + +fn claims(issuer: &str, audience: &str, resource: &str) -> Value { + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("current time") + .as_secs(); + json!({"iss":issuer,"aud":audience,"sub":"user","client_id":"mcp-client", + "jti":"unique-jag","iat":now,"exp":now + 3600,"resource":resource,"scope":"files.read"}) +} + +fn jag_response(claims: &Value) -> Value { + let mut response = json!({"access_token":jwt(claims),"issued_token_type":ID_JAG_TOKEN_TYPE, + "token_type":"N_A","resource":claims["resource"]}); + if let Some(scope) = claims.get("scope") { + response["scope"] = scope.clone(); + } + response +} + +fn token_response() -> Value { + json!({"access_token":"resource-token","token_type":"Bearer","expires_in":300}) +} + +#[tokio::test] +async fn public_client_round_trip_preserves_signed_narrowing() -> Result<()> { + let requested_scopes = ["files.read".to_string(), "files.write".to_string()]; + let client = http_client(); + for (scopes, echo_scope, refresh_token) in [ + (requested_scopes.as_slice(), true, "opaque-refresh-token"), + (&[], true, "opaque-refresh-token"), + (&[], false, "opaque-refresh-token"), + (requested_scopes.as_slice(), true, ""), + (requested_scopes.as_slice(), true, " \t"), + ] { + let server = MockServer::start().await; + let issuer = format!("{}/idp", server.uri()); + let audience = format!("{}/as", server.uri()); + let resource = format!("{}/mcp", server.uri()); + let mut jag = jag_response(&claims(&issuer, &audience, &resource)); + if !echo_scope { + jag.as_object_mut() + .expect("ID-JAG response") + .remove("scope"); + } + let valid = !refresh_token.trim().is_empty(); + for (endpoint, response) in [("/idp/token", jag.clone()), ("/as/token", token_response())] { + Mock::given(method("POST")) + .and(path(endpoint)) + .respond_with(ResponseTemplate::new(200).set_body_json(response)) + .expect(u64::from(valid)) + .mount(&server) + .await; + } + let expected_subject = refresh_token.to_string(); + let result = exchange_id_jag(EmaIdJagExchangeRequest { + resource: &resource, + scopes, + mcp_client_id: "mcp-client", + authorization_server_issuer: &audience, + authorization_server_token_endpoint: &format!("{audience}/token"), + idp_token_endpoint: &format!("{issuer}/token"), + idp_issuer: &issuer, + idp_client_id: "idp-client", + refresh_token: refresh_token.to_string(), + idp_http_client: Arc::clone(&client), + resource_http_client: Arc::clone(&client), + }) + .boxed() + .await; + if !valid { + assert!(result.is_err(), "invalid subject must fail before HTTP"); + assert!( + server + .received_requests() + .await + .expect("requests") + .is_empty() + ); + continue; + } + assert_eq!( + result?, + EmaAccessToken { + access_token: "resource-token".to_string(), + expires_in: Some(Duration::from_secs(300)), + } + ); + let requests = server.received_requests().await.expect("requests"); + assert_eq!(requests.len(), 2); + assert!( + requests + .iter() + .all(|request| request.headers.get("authorization").is_none()) + ); + let mut forms = requests + .iter() + .map(|request| unique_form_fields(&request.body)) + .collect::>(); + assert_eq!( + forms[0].remove("scope"), + (!scopes.is_empty()).then(|| scopes.join(" ")) + ); + assert_eq!( + forms[0], + HashMap::from([ + ( + "grant_type".to_string(), + TOKEN_EXCHANGE_GRANT_TYPE.to_string() + ), + ( + "requested_token_type".to_string(), + ID_JAG_TOKEN_TYPE.to_string(), + ), + ("subject_token".to_string(), expected_subject), + ( + "subject_token_type".to_string(), + "urn:ietf:params:oauth:token-type:refresh_token".to_string(), + ), + ("audience".to_string(), audience.clone()), + ("resource".to_string(), resource.clone()), + ("client_id".to_string(), "idp-client".to_string()), + ]) + ); + assert_eq!( + forms[1], + HashMap::from([ + ("grant_type".to_string(), JWT_BEARER_GRANT_TYPE.to_string()), + ( + "assertion".to_string(), + jag["access_token"].as_str().expect("JAG").to_string() + ), + ("client_id".to_string(), "mcp-client".to_string()), + ]) + ); + } + Ok(()) +} + +#[test] +fn signed_claims_and_resource_tokens_cannot_widen_authority() -> Result<()> { + let requested = HashSet::from(["files.read", "files.write"]); + let original = claims( + "https://idp.example", + "https://as.example", + "https://mcp.example", + ); + let binding = || IdJagBinding { + issuer: "https://idp.example", + audience: "https://as.example", + client_id: "mcp-client", + resource: "https://mcp.example", + requested_scopes: &requested, + }; + let valid: IdJagResponse = serde_json::from_value(jag_response(&original))?; + let granted = valid.validate(binding())?; + assert_eq!(granted, HashSet::from(["files.read".to_string()])); + for (requested_scopes, signed_scope, response_scope, valid) in [ + ("", Some("files.read"), Some("files.read"), true), + ("", Some("files.read"), None, true), + ("", None, None, true), + ("files.read", Some("files.read"), None, true), + ("files.read files.write", Some("files.read"), None, false), + ( + "files.read", + Some("files.read files.write"), + Some("files.read files.write"), + false, + ), + ("files.read", None, None, false), + ("", Some("files.read"), Some("files.write"), false), + ("", Some(" \t"), None, false), + ("", Some("files.read files.read"), Some("files.read"), false), + ("", Some("files.read"), Some(" \t"), false), + ("", Some("files.read"), Some("files.read files.read"), false), + ] { + let requested_scopes = requested_scopes.split_ascii_whitespace().collect(); + let mut scoped_claims = original.clone(); + scoped_claims + .as_object_mut() + .expect("ID-JAG claims") + .remove("scope"); + if let Some(scope) = signed_scope { + scoped_claims["scope"] = json!(scope); + } + let mut response = jag_response(&scoped_claims); + response + .as_object_mut() + .expect("ID-JAG response") + .remove("scope"); + if let Some(scope) = response_scope { + response["scope"] = json!(scope); + } + let response: IdJagResponse = serde_json::from_value(response)?; + let result = response.validate(IdJagBinding { + requested_scopes: &requested_scopes, + ..binding() + }); + assert_eq!( + result.is_ok(), + valid, + "requested {requested_scopes:?}, signed {signed_scope:?}, response {response_scope:?}" + ); + if let Ok(granted) = result { + assert_eq!( + granted, + signed_scope + .into_iter() + .map(str::to_string) + .collect::>() + ); + for (scope, valid_token) in [ + (None, true), + (Some("files.read"), signed_scope.is_some()), + (Some("files.read files.write"), false), + ] { + let mut response = token_response(); + if let Some(scope) = scope { + response["scope"] = json!(scope); + } + let response: McpAccessTokenResponse = serde_json::from_value(response)?; + assert_eq!( + response.validate("https://mcp.example", &granted).is_ok(), + valid_token, + "ID-JAG {signed_scope:?}, bearer {scope:?}" + ); + } + } + } + let mut rotated = jag_response(&original); + rotated["refresh_token"] = json!("unsupported-jag-refresh-token"); + let response: IdJagResponse = serde_json::from_value(rotated)?; + assert!(response.validate(binding()).is_err()); + for header in [ + json!({"alg": "ES256", "typ": "JWT"}), + json!({"alg": "ES256"}), + json!({"alg": "none", "typ": "oauth-id-jag+jwt"}), + ] { + let mut changed = jag_response(&original); + changed["access_token"] = json!(format!( + "{}.{}.signature", + URL_SAFE_NO_PAD.encode(serde_json::to_vec(&header)?), + URL_SAFE_NO_PAD.encode(serde_json::to_vec(&original)?), + )); + let response: IdJagResponse = serde_json::from_value(changed)?; + assert!( + response.validate(binding()).is_err(), + "accepted invalid ID-JAG header {header}" + ); + } + for (field, value) in [ + ("iss", json!("https://attacker.example")), + ("aud", json!("https://attacker.example")), + ("client_id", json!("other-client")), + ("sub", json!("")), + ("jti", json!("")), + ("exp", json!(0)), + ("iat", json!(u64::MAX)), + ("scope", json!("files.admin")), + ( + "resource", + json!(["https://mcp.example", "https://other.example"]), + ), + ] { + let mut changed = original.clone(); + changed[field] = value; + let response: IdJagResponse = serde_json::from_value(jag_response(&changed))?; + assert!( + response.validate(binding()).is_err(), + "accepted changed {field}" + ); + } + for (field, value) in [ + ("scope", json!("files.read files.write")), + ("scope", json!("files.read files.read")), + ("resource", json!("https://other.example")), + ("expires_in", json!(0)), + ("refresh_token", json!("refresh")), + ("token_type", json!("N_A")), + ("access_token", json!("")), + ] { + let mut changed = token_response(); + changed[field] = value; + let response: McpAccessTokenResponse = serde_json::from_value(changed)?; + assert!( + response.validate("https://mcp.example", &granted).is_err(), + "accepted changed {field}" + ); + } + let mut explicit_binding = token_response(); + explicit_binding["resource"] = json!("https://mcp.example"); + explicit_binding["scope"] = json!("files.read"); + let response: McpAccessTokenResponse = serde_json::from_value(explicit_binding)?; + assert_eq!( + response.validate("https://mcp.example", &granted)?, + EmaAccessToken { + access_token: "resource-token".to_string(), + expires_in: Some(Duration::from_secs(300)), + } + ); + Ok(()) +} + +#[test] +fn identity_and_credential_destinations_are_bound() { + for endpoint in [ + "http://idp.example/token", + "https://user:pass@idp.example/token", + "https://idp.example/token#fragment", + ] { + assert!( + validate_ema_oauth_endpoint(endpoint, "IdP").is_err(), + "accepted {endpoint}" + ); + } + let original = claims("https://idp.example", "idp-client", "https://mcp.example"); + assert!( + validate_oidc_identity_assertion(&jwt(&original), "https://idp.example", "idp-client") + .is_ok() + ); + for (field, value) in [ + ("iss", json!("https://other.example")), + ("aud", json!(["idp-client", "other"])), + ("azp", json!("other")), + ("exp", json!(0)), + ("sub", json!("")), + ] { + let mut changed = original.clone(); + changed[field] = value; + assert!( + validate_oidc_identity_assertion(&jwt(&changed), "https://idp.example", "idp-client") + .is_err(), + "accepted changed {field}" + ); + } +} + +#[tokio::test] +async fn provider_errors_cannot_reflect_credentials() { + const SENTINEL: &str = "secret-assertion-sentinel"; + for (code, expected) in [ + (SENTINEL, "OAuth token request rejected"), + ("invalid_grant", "invalid_grant"), + ( + "insufficient_user_authentication", + "insufficient_user_authentication", + ), + ] { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({ + "error":code,"error_description":SENTINEL, + }))) + .mount(&server) + .await; + let error = post_form::( + &http_client(), + &format!("{}/token", server.uri()), + &[("subject_token", SENTINEL)], + "test-client", + EmaInvalidGrantSource::EnterpriseIdentity, + "test token exchange", + ) + .await + .expect_err("provider error should fail"); + match code { + "invalid_grant" => assert_eq!( + error.downcast_ref::(), + Some(&EmaAuthFailure::InvalidGrant { + grant_source: EmaInvalidGrantSource::EnterpriseIdentity, + }) + ), + "insufficient_user_authentication" => assert_eq!( + error.downcast_ref::(), + Some(&EmaAuthFailure::InsufficientUserAuthentication) + ), + _ => assert_eq!(error.downcast_ref::(), None), + } + let error = error.to_string(); + assert!(!error.contains(SENTINEL), "provider reflected a credential"); + assert!(error.ends_with(expected), "{error}"); + } + let server = MockServer::start().await; + let mut malformed = token_response(); + malformed["expires_in"] = json!(SENTINEL); + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(malformed)) + .mount(&server) + .await; + let error = post_form::( + &http_client(), + &format!("{}/token", server.uri()), + &[], + "test-client", + EmaInvalidGrantSource::EnterpriseIdentity, + "test token exchange", + ) + .await + .err() + .expect("malformed response should fail"); + assert!( + !format!("{error:#}").contains(SENTINEL), + "parser reflected a credential" + ); +} diff --git a/codex-rs/rmcp-client/src/ema_identity.rs b/codex-rs/rmcp-client/src/ema_identity.rs new file mode 100644 index 0000000000000000000000000000000000000000..f906e553f7a0e4e104eaa8e2201d8b83418cb6a1 --- /dev/null +++ b/codex-rs/rmcp-client/src/ema_identity.rs @@ -0,0 +1,167 @@ +//! Shared enterprise IdP discovery and stored refresh-token resolution. +//! Login and token exchange require published metadata bound to the configured issuer. + +use std::sync::Arc; + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use codex_exec_server::HttpClient; +use codex_keyring_store::DefaultKeyringStore; +use codex_keyring_store::KeyringStore; +use oauth2::TokenResponse; +use rmcp::transport::AuthorizationManager; +use rmcp::transport::auth::AuthorizationMetadata; +use rmcp::transport::auth::OAuthHttpClient; + +use crate::ema_auth_policy::advertised_capability; +use crate::ema_auth_policy::ema_reauthentication_required; +use crate::ema_auth_policy::validate_ema_oauth_endpoint; +use crate::ema_auth_policy::validate_ema_public_client_auth; +use crate::ema_claims::ID_JAG_TOKEN_TYPE; +use crate::ema_exchange::TOKEN_EXCHANGE_GRANT_TYPE; +use crate::http_client_adapter::StreamableHttpRedirectMode; +use crate::oauth::RefreshCredentialLock; +use crate::oauth::StoredOAuthCredentialSnapshot; +use crate::oauth::StoredOAuthTokens; +use crate::oauth::stored_oidc_identity; +use crate::oauth_http_client::OAuthHttpClientAdapter; +use crate::utils::build_default_headers; + +pub struct EmaIdpIdentityRequest<'a> { + pub issuer: &'a str, + pub client_id: &'a str, + pub credentials: &'a StoredOAuthCredentialSnapshot, + pub http_client: Arc, + pub redirect_mode: StreamableHttpRedirectMode, +} + +/// An opaque IdP refresh token whose credential lock is held through token exchange. +#[allow(dead_code)] +pub struct EmaIdpIdentity { + pub(crate) token_endpoint: String, + pub(crate) refresh_token: String, + pub(crate) credential_lock: RefreshCredentialLock, +} + +/// Returns whether stored credentials contain a refresh token bound to the configured login. +pub fn stored_ema_identity_is_usable( + tokens: &StoredOAuthTokens, + issuer: &str, + client_id: &str, +) -> bool { + tokens.url == issuer + && tokens.bound_issuer() == Some(issuer) + && tokens.client_id == client_id + && tokens.has_refresh_token() + && stored_oidc_identity(tokens).is_ok() +} + +pub(crate) async fn resolve_ema_idp_authorization_manager( + issuer: &str, + http_client: Arc, +) -> Result<(AuthorizationManager, AuthorizationMetadata)> { + validate_ema_oauth_endpoint(issuer, "enterprise IdP issuer")?; + let mut manager = AuthorizationManager::new_with_oauth_http_client(issuer, http_client) + .await + .context("failed to create enterprise IdP metadata discovery client")?; + manager.set_allow_missing_issuer(false); + let resolution = manager + .resolve_metadata() + .await + .context("failed to discover enterprise IdP authorization metadata")?; + if !resolution.source.is_discovered() { + bail!("enterprise IdP must publish authorization metadata"); + } + let metadata = resolution.metadata; + if metadata.issuer.as_deref() != Some(issuer) { + bail!("enterprise IdP authorization metadata issuer does not match configuration"); + } + validate_ema_oauth_endpoint(&metadata.token_endpoint, "enterprise IdP token endpoint")?; + validate_ema_public_client_auth( + metadata + .additional_fields + .get("token_endpoint_auth_methods_supported"), + "enterprise IdP", + )?; + Ok((manager, metadata)) +} + +/// Resolve a stored refresh-token subject against the configured enterprise IdP metadata. +pub async fn resolve_ema_idp_identity( + request: EmaIdpIdentityRequest<'_>, +) -> Result { + resolve_ema_idp_identity_in(request, &DefaultKeyringStore).await +} + +async fn resolve_ema_idp_identity_in( + request: EmaIdpIdentityRequest<'_>, + keyring_store: &K, +) -> Result { + if request.issuer.trim().is_empty() || request.client_id.trim().is_empty() { + bail!("ema_auth requires a non-empty enterprise IdP issuer and client ID"); + } + let credentials = request.credentials.credentials(); + if credentials.url != request.issuer + || credentials.bound_issuer() != Some(request.issuer) + || credentials.client_id != request.client_id + { + bail!("stored enterprise IdP credentials do not match the configured issuer and client"); + } + stored_oidc_identity(credentials)?; + let client = OAuthHttpClientAdapter::new_with_redirect_mode( + request.http_client, + build_default_headers(/*http_headers*/ None, /*env_http_headers*/ None)?, + request.issuer, + /*has_configured_headers*/ false, + request.redirect_mode, + )?; + let (_, metadata) = + resolve_ema_idp_authorization_manager(request.issuer, Arc::new(client)).await?; + for (name, expected) in [ + ( + "identity_chaining_requested_token_types_supported", + ID_JAG_TOKEN_TYPE, + ), + ("grant_types_supported", TOKEN_EXCHANGE_GRANT_TYPE), + ] { + if advertised_capability(metadata.additional_fields.get(name), expected, name)? + == Some(false) + { + bail!( + "enterprise IdP does not advertise the required ID-JAG token exchange capability" + ); + } + } + let credential_lock = + RefreshCredentialLock::acquire_for_server(&credentials.server_name, &credentials.url) + .await?; + let snapshot = request.credentials.clone(); + let keyring_store = keyring_store.clone(); + // The worker retains the guard even if its caller stops waiting for the reread. + tokio::task::spawn_blocking(move || { + let latest = snapshot.load_ema_credentials(&keyring_store)?; + let refresh_token = latest + .token_response + .0 + .refresh_token() + .filter(|token| !token.secret().trim().is_empty()) + .ok_or_else(|| { + ema_reauthentication_required( + "enterprise IdP session has no refresh token; sign in again", + ) + })?; + Ok(EmaIdpIdentity { + token_endpoint: metadata.token_endpoint, + refresh_token: refresh_token.secret().to_string(), + credential_lock, + }) + }) + .await + .map_err(|_| anyhow!("enterprise IdP credential reread task failed"))? +} + +#[cfg(test)] +#[path = "ema_identity_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/ema_identity_tests.rs b/codex-rs/rmcp-client/src/ema_identity_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..53c7ef497ad4ac210849bd00af1e1058308db4ac --- /dev/null +++ b/codex-rs/rmcp-client/src/ema_identity_tests.rs @@ -0,0 +1,435 @@ +use std::io; +use std::sync::Mutex; +use std::sync::PoisonError; +use std::sync::mpsc; +use std::thread; +use std::time::Duration; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +use base64::Engine; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use codex_config::types::AuthKeyringBackendKind; +use codex_exec_server::RouteAwareHttpClient; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_keyring_store::CredentialStoreError; +use codex_keyring_store::tests::MockKeyringStore; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use pretty_assertions::assert_ne; +use serde_json::json; +use tokio::sync::oneshot; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::*; +use crate::EmaAuthFailure; +use crate::WrappedOAuthTokenResponse; +use crate::oauth::ResolvedOAuthCredentialStore; +use crate::oauth::test_support::TempCodexHome; + +fn credentials(issuer: &str, subject: &str, expires_at: u64) -> StoredOAuthTokens { + let assertion = format!( + "{}.{}.signature", + URL_SAFE_NO_PAD.encode(br#"{"alg":"ES256"}"#), + URL_SAFE_NO_PAD.encode( + serde_json::to_vec(&json!({ + "iss":issuer,"aud":"idp-client","sub":subject,"exp":expires_at, + })) + .expect("claims") + ) + ); + StoredOAuthTokens { + server_name: "ema-idp:enterprise-test".to_string(), + url: issuer.to_string(), + issuer: Some(issuer.to_string()), + client_id: "idp-client".to_string(), + token_response: WrappedOAuthTokenResponse( + serde_json::from_value(json!({ + "access_token":"unused","token_type":"Bearer","id_token":assertion, + "refresh_token":"stored-refresh","scope":"openid offline_access", + })) + .expect("credentials"), + ), + expires_at: None, + } +} + +fn request<'a>( + issuer: &'a str, + credentials: &'a StoredOAuthCredentialSnapshot, +) -> EmaIdpIdentityRequest<'a> { + EmaIdpIdentityRequest { + issuer, + client_id: "idp-client", + credentials, + http_client: Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + redirect_mode: StreamableHttpRedirectMode::Legacy, + } +} + +async fn discovery() -> (MockServer, String) { + let server = MockServer::start().await; + let issuer = format!("{}/idp", server.uri()); + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/idp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "issuer":issuer,"authorization_endpoint":format!("{issuer}/authorize"), + "token_endpoint":format!("{issuer}/token"), + "identity_chaining_requested_token_types_supported":[ID_JAG_TOKEN_TYPE], + "grant_types_supported":[TOKEN_EXCHANGE_GRANT_TYPE], + "token_endpoint_auth_methods_supported":["none"], + }))) + .mount(&server) + .await; + (server, issuer) +} + +#[tokio::test] +async fn enterprise_discovery_shares_idp_checks_but_keeps_login_endpoint_policy() -> Result<()> { + for case in [ + "valid", + "missing-discovery", + "issuer", + "missing-issuer", + "token-endpoint", + "public-client", + "authorization-endpoint", + ] { + let server = MockServer::start().await; + let issuer = format!("{}/idp", server.uri()); + // Login need not advertise ID-JAG exchange capabilities. + let mut metadata = json!({ + "issuer": issuer, "authorization_endpoint": format!("{issuer}/authorize"), + "token_endpoint": format!("{issuer}/token"), + "token_endpoint_auth_methods_supported": ["none"], + "grant_types_supported": ["authorization_code"], + }); + match case { + "issuer" => metadata["issuer"] = json!("https://other.example"), + "missing-issuer" => { + metadata.as_object_mut().expect("metadata").remove("issuer"); + } + "token-endpoint" => metadata["token_endpoint"] = json!("http://unsafe.example/token"), + "public-client" => { + metadata["token_endpoint_auth_methods_supported"] = json!(["client_secret_basic"]) + } + "authorization-endpoint" => { + metadata["authorization_endpoint"] = json!("http://unsafe.example/authorize") + } + "valid" | "missing-discovery" => {} + _ => unreachable!(), + } + if case != "missing-discovery" { + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/idp")) + .respond_with(ResponseTemplate::new(200).set_body_json(metadata)) + .mount(&server) + .await; + } + let client = Arc::new(OAuthHttpClientAdapter::new_with_redirect_mode( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + build_default_headers(/*http_headers*/ None, /*env_http_headers*/ None)?, + &issuer, + /*has_configured_headers*/ false, + StreamableHttpRedirectMode::Legacy, + )?); + let shared = resolve_ema_idp_authorization_manager(&issuer, client.clone()).await; + let login = crate::enterprise_oauth_login::resolve_enterprise_authorization_manager( + &issuer, client, + ) + .await; + assert_eq!( + (shared.is_ok(), login.is_ok()), + ( + matches!(case, "valid" | "authorization-endpoint"), + case == "valid" + ), + "{case}", + ); + if case == "missing-discovery" { + assert_eq!( + shared + .err() + .expect("reject synthesized metadata") + .to_string(), + "enterprise IdP must publish authorization metadata", + ); + } + } + Ok(()) +} + +const REREAD_TEST_TIMEOUT: Duration = Duration::from_secs(/*secs*/ 5); +const REREAD_PANIC_SENTINEL: &str = "credential-panic-payload-sentinel"; + +#[derive(Clone, Copy, Debug)] +enum ReadOutcome { + Stored, + BackendError, + Panic, +} + +#[derive(Clone, Debug)] +struct GatedKeyringStore { + inner: MockKeyringStore, + executor_thread: thread::ThreadId, + entered: Arc>>>, + release: Arc>>, +} + +impl KeyringStore for GatedKeyringStore { + fn load(&self, service: &str, account: &str) -> Result, CredentialStoreError> { + // Fail before blocking if a regression puts this read on the current-thread executor. + assert_ne!(thread::current().id(), self.executor_thread); + if let Some(entered) = self + .entered + .lock() + .unwrap_or_else(PoisonError::into_inner) + .take() + { + let _ = entered.send(()); + } + let outcome = match self + .release + .lock() + .unwrap_or_else(PoisonError::into_inner) + .recv_timeout(REREAD_TEST_TIMEOUT) + { + Ok(outcome) => outcome, + Err(mpsc::RecvTimeoutError::Disconnected) => ReadOutcome::Stored, + Err(mpsc::RecvTimeoutError::Timeout) => panic!("credential reread gate timed out"), + }; + match outcome { + ReadOutcome::Stored => self.inner.load(service, account), + ReadOutcome::BackendError => Err(CredentialStoreError::new( + keyring::Error::PlatformFailure(Box::new(io::Error::new( + io::ErrorKind::PermissionDenied, + "credential backend unavailable", + ))), + )), + ReadOutcome::Panic => panic!("{REREAD_PANIC_SENTINEL}"), + } + } + + fn save(&self, service: &str, account: &str, value: &str) -> Result<(), CredentialStoreError> { + self.inner.save(service, account, value) + } + + fn delete(&self, service: &str, account: &str) -> Result { + self.inner.delete(service, account) + } +} + +#[tokio::test(flavor = "current_thread")] +async fn refresh_subject_reread_is_cancellable_and_releases_guard_on_failure() -> Result<()> { + let _home = TempCodexHome::new(); + let (_server, issuer) = discovery().await; + let store = ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct); + let stored = credentials(&issuer, "user", /*expires_at*/ 0); + let snapshot = StoredOAuthCredentialSnapshot::new(stored.clone(), store); + for outcome in [ + ReadOutcome::Stored, + ReadOutcome::BackendError, + ReadOutcome::Panic, + ] { + let inner = MockKeyringStore::default(); + store.save(&inner, &stored.server_name, &stored)?; + let (entered_tx, entered_rx) = oneshot::channel(); + // Dropping the sole sender releases the worker on any early return or panic. + let (release_tx, release_rx) = mpsc::channel(); + let keyring = GatedKeyringStore { + inner, + executor_thread: thread::current().id(), + entered: Arc::new(Mutex::new(Some(entered_tx))), + release: Arc::new(Mutex::new(release_rx)), + }; + let request = request(&issuer, &snapshot); + crate::oauth::test_support::warm_http_client(request.http_client.as_ref()).await?; + let mut identity = Box::pin(resolve_ema_idp_identity_in(request, &keyring)); + tokio::select! { + result = &mut identity => { + result?; + bail!("credential reread completed before its gate was released"); + } + entered = tokio::time::timeout(REREAD_TEST_TIMEOUT, entered_rx) => { entered??; } + } + + if matches!(outcome, ReadOutcome::Stored) { + // The caller times out while the already-started blocking read remains gated. + assert!( + tokio::time::timeout(Duration::ZERO, identity) + .await + .is_err() + ); + assert!( + RefreshCredentialLock::acquire_for_server(&stored.server_name, &issuer) + .now_or_never() + .is_none() + ); + release_tx.send(outcome)?; + } else { + release_tx.send(outcome)?; + let error = tokio::time::timeout(REREAD_TEST_TIMEOUT, identity) + .await? + .err() + .context("credential reread should fail")?; + if matches!(outcome, ReadOutcome::Panic) { + assert_eq!( + format!("{error:#}"), + "enterprise IdP credential reread task failed" + ); + } else { + assert!(error.to_string().contains("refusing file fallback")); + // The store's transparent wrapper exposes a platform error's source. + assert!(error.chain().any(|cause| { + matches!( + cause.downcast_ref::(), + Some(error) if error.kind() == io::ErrorKind::PermissionDenied + ) + })); + } + } + // Drain the detached read before TempCodexHome changes the process environment. + let _released = tokio::time::timeout( + REREAD_TEST_TIMEOUT, + RefreshCredentialLock::acquire_for_server(&stored.server_name, &issuer), + ) + .await??; + } + Ok(()) +} + +#[tokio::test] +async fn refresh_subject_rereads_pinned_credentials_after_id_token_expiry() -> Result<()> { + let _home = TempCodexHome::new(); + let (_server, issuer) = discovery().await; + let store = ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct); + let stored = credentials(&issuer, "user", /*expires_at*/ 0); + let snapshot = StoredOAuthCredentialSnapshot::new(stored.clone(), store); + let keyring = MockKeyringStore::default(); + store.save(&keyring, &stored.server_name, &stored)?; + + let identity = resolve_ema_idp_identity_in(request(&issuer, &snapshot), &keyring).await?; + assert_eq!( + (&identity.token_endpoint, identity.refresh_token.as_str()), + (&format!("{issuer}/token"), "stored-refresh") + ); + assert!( + tokio::time::timeout( + Duration::from_millis(/*millis*/ 50), + RefreshCredentialLock::acquire_for_server(&stored.server_name, &issuer), + ) + .await + .is_err() + ); + drop(identity); + let _released = RefreshCredentialLock::acquire_for_server(&stored.server_name, &issuer).await?; + Ok(()) +} + +#[tokio::test] +async fn refresh_subject_rejects_removed_replaced_or_missing_credentials() -> Result<()> { + let _home = TempCodexHome::new(); + let (server, issuer) = discovery().await; + let now = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs(); + let store = ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct); + let original = credentials(&issuer, "user", /*expires_at*/ 0); + for change in [ + "deleted", + "subject", + "issuer", + "client", + "refresh", + "login", + "missing-refresh", + "file", + ] { + let keyring = MockKeyringStore::default(); + let snapshot = StoredOAuthCredentialSnapshot::new( + original.clone(), + if change == "file" { + ResolvedOAuthCredentialStore::File + } else { + store + }, + ); + let mut latest = original.clone(); + match change { + "subject" => latest = credentials(&issuer, "other-user", /*expires_at*/ 0), + "issuer" => latest.issuer = Some("https://other.example".to_string()), + "client" => latest.client_id = "other-client".to_string(), + "refresh" => { + latest + .token_response + .0 + .set_refresh_token(Some(oauth2::RefreshToken::new( + "other-users-refresh".to_string(), + ))) + } + "login" => latest = credentials(&issuer, "user", now + 3600), + "missing-refresh" => latest.token_response.0.set_refresh_token(None), + "deleted" | "file" => {} + _ => panic!("unexpected change"), + } + if change != "deleted" { + store.save(&keyring, &original.server_name, &latest)?; + } + let error = resolve_ema_idp_identity_in(request(&issuer, &snapshot), &keyring) + .await + .err() + .expect("must not reuse the stale snapshot or fall back to its ID token"); + if change == "file" { + assert!(error.to_string().contains("require keyring storage")); + } else { + assert_eq!( + error.downcast_ref::(), + Some(&EmaAuthFailure::ReauthenticationRequired), + "{change}" + ); + } + } + assert!( + server + .received_requests() + .await + .expect("requests") + .iter() + .all(|request| request.method.as_str() == "GET") + ); + Ok(()) +} + +#[test] +fn stored_identity_usability_requires_a_bound_refresh_token_not_a_current_id_token() { + let issuer = "https://idp.example"; + let expired = credentials(issuer, "user", /*expires_at*/ 0); + assert!(stored_ema_identity_is_usable( + &expired, + issuer, + "idp-client" + )); + for change in ["url", "issuer", "client", "refresh"] { + let mut changed = expired.clone(); + match change { + "url" => changed.url = "https://other.example".to_string(), + "issuer" => changed.issuer = Some("https://other.example".to_string()), + "client" => changed.client_id = "other-client".to_string(), + "refresh" => changed.token_response.0.set_refresh_token(None), + _ => panic!("unexpected change"), + } + assert!(!stored_ema_identity_is_usable( + &changed, + issuer, + "idp-client" + )); + } +} diff --git a/codex-rs/rmcp-client/src/enterprise_oauth_login.rs b/codex-rs/rmcp-client/src/enterprise_oauth_login.rs new file mode 100644 index 0000000000000000000000000000000000000000..6c805fd460648491d9a46e082048f4864d7fcf6f --- /dev/null +++ b/codex-rs/rmcp-client/src/enterprise_oauth_login.rs @@ -0,0 +1,371 @@ +//! Enterprise OIDC policy and staged, keyring-only credential commits. +//! Browser completion never persists a grant; its owner revalidates authority under the commit lock. + +use std::future::Future; +use std::net::IpAddr; +use std::net::Ipv4Addr; +use std::sync::Arc; + +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::HttpClient; +use http::Method; +use http::header::CONTENT_LENGTH; +use http::header::CONTENT_TYPE; +use rmcp::transport::AuthorizationManager; +use rmcp::transport::auth::AuthorizationMetadata; +use rmcp::transport::auth::OAuthHttpClient; +use rmcp::transport::auth::OAuthHttpClientFuture; +use rmcp::transport::auth::OAuthHttpRequest; +use tracing::instrument::WithSubscriber; +use url::Host; +use url::Url; + +use crate::StoredOAuthTokens; +use crate::ema_auth_policy::validate_ema_oauth_endpoint; +use crate::ema_claims::validate_oidc_identity_assertion; +use crate::ema_identity::resolve_ema_idp_authorization_manager; +use crate::http_client_adapter::StreamableHttpRedirectMode; +use crate::oauth::EnterpriseOAuthGeneration; +use crate::oauth::EnterpriseOAuthGenerationFile; +use crate::oauth::RefreshCredentialLock; +use crate::oauth::delete_oauth_tokens_with_lock_held; +use crate::oauth::save_oauth_tokens_with_lock_held; +use crate::oauth::validate_authorization_server_endpoints; +use crate::oauth_client_registration::McpOAuthClientRegistration; +use crate::perform_oauth_login::OAuthHttpContext; +use crate::perform_oauth_login::OAuthLoginPurpose; +use crate::perform_oauth_login::OauthLoginFlow; + +/// An exclusive credential mutation guard. Hold it through primary account logout +/// so a competing process cannot commit between deleting the grant and signing out. +/// Coordination, like ordinary OAuth refresh, is scoped to the same CODEX_HOME. +pub struct EnterpriseOAuthCredentialGuard { + credential_name: String, + issuer: String, + keyring_backend: AuthKeyringBackendKind, + generation_file: EnterpriseOAuthGenerationFile, + _lock: RefreshCredentialLock, +} + +impl EnterpriseOAuthCredentialGuard { + pub async fn acquire( + credential_name: &str, + issuer: &str, + keyring_backend: AuthKeyringBackendKind, + ) -> Result { + let lock = RefreshCredentialLock::acquire_for_server(credential_name, issuer) + .with_subscriber(tracing::subscriber::NoSubscriber::default()) + .await + .map_err(|_| anyhow!("failed to lock enterprise credentials"))?; + let generation_file = EnterpriseOAuthGenerationFile::open(credential_name, issuer, &lock) + .map_err(|_| anyhow!("failed to open enterprise login generation"))?; + Ok(Self { + credential_name: credential_name.to_owned(), + issuer: issuer.to_owned(), + keyring_backend, + generation_file, + _lock: lock, + }) + } + + /// Invalidate pending logins and delete the grant, if present. A false return + /// means no grant was stored; earlier login attempts are still invalidated. + pub fn delete_tokens(&self) -> Result { + // Invalidate attempts that have not stored a grant yet, including other processes. + // Persist before deletion so no successful logout can admit an earlier login. + self.generation_file + .replace() + .map_err(|_| anyhow!("failed to invalidate pending enterprise sign-ins"))?; + // Underlying keyring diagnostics include the account-scoped key. Suppress + // them and discard their error chain, not merely the outer error message. + tracing::subscriber::with_default(tracing::subscriber::NoSubscriber::default(), || { + delete_oauth_tokens_with_lock_held( + &self._lock, + &self.credential_name, + &self.issuer, + OAuthCredentialsStoreMode::Keyring, + self.keyring_backend, + ) + }) + .map_err(|_| anyhow!("failed to delete enterprise credentials")) + } +} + +/// Invalidate earlier enterprise logins across processes, then delete any stored grant. +pub async fn delete_enterprise_oauth_tokens( + credential_name: &str, + issuer: &str, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result { + EnterpriseOAuthCredentialGuard::acquire(credential_name, issuer, keyring_backend_kind) + .await? + .delete_tokens() +} + +/// A browser login with no detached persistence worker. Dropping this handle or +/// its wait future unblocks the callback listener and cannot write credentials. +pub struct EnterpriseOAuthLoginHandle { + flow: OauthLoginFlow, + keyring_backend: AuthKeyringBackendKind, + generation: EnterpriseOAuthGeneration, +} + +impl EnterpriseOAuthLoginHandle { + pub fn authorization_url(&self) -> String { + self.flow.authorization_url() + } + + pub async fn wait(self) -> Result { + let stored = self + .flow + .complete(/*emit_browser_url*/ false) + .with_subscriber(tracing::subscriber::NoSubscriber::default()) + .await + .map_err(|_| anyhow!("enterprise IdP authorization failed"))?; + validate_enterprise_credentials(&stored)?; + Ok(EnterpriseOAuthCredentials { + stored, + keyring_backend: self.keyring_backend, + generation: self.generation, + }) + } +} + +/// A validated grant that is not yet stored. It intentionally does not expose tokens. +pub struct EnterpriseOAuthCredentials { + stored: StoredOAuthTokens, + keyring_backend: AuthKeyringBackendKind, + generation: EnterpriseOAuthGeneration, +} + +impl EnterpriseOAuthCredentials { + /// Revalidate the current attempt, account and configuration under the same + /// exclusive lock used by logout, then persist without an intervening await. + /// Return the authority proof so the caller can retire its attempt while still + /// holding the commit gate, before notifying clients or refreshing runtimes. + pub async fn commit_if(self, is_current: F) -> Result + where + F: FnOnce() -> Fut, + Fut: Future>, + { + let guard = EnterpriseOAuthCredentialGuard::acquire( + &self.stored.server_name, + &self.stored.url, + self.keyring_backend, + ) + .await?; + let generation = guard + .generation_file + .current() + .map_err(|_| anyhow!("failed to read enterprise login generation"))?; + if generation.as_ref() != Some(&self.generation) { + bail!("enterprise sign-in was invalidated by logout"); + } + let Some(authority) = is_current().await else { + bail!("enterprise sign-in no longer matches the active account or configuration"); + }; + tracing::subscriber::with_default(tracing::subscriber::NoSubscriber::default(), || { + save_oauth_tokens_with_lock_held( + &guard._lock, + &self.stored.server_name, + &self.stored, + OAuthCredentialsStoreMode::Keyring, + self.keyring_backend, + ) + }) + .map_err(|_| anyhow!("failed to store enterprise credentials"))?; + Ok(authority) + } +} + +/// Start a host-registered enterprise login without launching the browser or +/// persisting a credential. The caller owns both the attempt and its later commit. +pub async fn perform_enterprise_oauth_login_return_url( + request: EnterpriseOAuthLoginRequest<'_>, +) -> Result { + // Capture before discovery or browser setup, then release the lock while the user signs in. + let generation = { + let guard = EnterpriseOAuthCredentialGuard::acquire( + request.credential_name, + request.issuer, + request.keyring_backend_kind, + ) + .await?; + match guard.generation_file.current() { + Ok(Some(generation)) => generation, + Ok(None) => guard + .generation_file + .replace() + .map_err(|_| anyhow!("failed to initialize enterprise login generation"))?, + Err(_) => bail!("failed to read enterprise login generation"), + } + }; + let flow = OauthLoginFlow::new( + request.credential_name, + request.issuer, + OAuthCredentialsStoreMode::Keyring, + request.keyring_backend_kind, + OAuthHttpContext { + http_headers: None, + env_http_headers: None, + http_client: request.http_client, + redirect_mode: request.redirect_mode, + }, + &["openid".to_string(), "offline_access".to_string()], + Some(request.client_id), + OAuthLoginPurpose::EnterpriseIdp, + McpOAuthClientRegistration::Auto, + /*oauth_resource*/ None, + /*launch_browser*/ false, + request.callback_port, + request.callback_url, + /*global_callback_url*/ None, + request.timeout_secs, + ) + .with_subscriber(tracing::subscriber::NoSubscriber::default()) + .await + .map_err(|_| anyhow!("failed to start enterprise IdP authorization"))?; + Ok(EnterpriseOAuthLoginHandle { + flow, + keyring_backend: request.keyring_backend_kind, + generation, + }) +} + +pub(crate) fn enterprise_callback_settings( + issuer: &str, + client_id: Option<&str>, + callback_url: Option<&str>, + callback_port: Option, +) -> Result<(IpAddr, Option)> { + validate_ema_oauth_endpoint(issuer, "enterprise IdP issuer")?; + if client_id.is_none_or(|client_id| client_id.trim().is_empty()) { + bail!("enterprise IdP login requires its registered client ID"); + } + let (ip, registered_port) = if let Some(callback_url) = callback_url { + validate_ema_oauth_endpoint(callback_url, "enterprise IdP callback URL")?; + let callback = Url::parse(callback_url)?; + let ip = match (callback.scheme(), callback.host()) { + ("http", Some(Host::Domain("localhost"))) => Ipv4Addr::LOCALHOST.into(), + ("http", Some(Host::Ipv4(ip))) if ip.is_loopback() => ip.into(), + ("http", Some(Host::Ipv6(ip))) if ip.is_loopback() => ip.into(), + _ => bail!("enterprise IdP callback URL must use an HTTP loopback address"), + }; + (ip, callback.port()) + } else { + (Ipv4Addr::LOCALHOST.into(), None) + }; + if callback_port + .zip(registered_port) + .is_some_and(|(configured, registered)| configured != registered) + { + bail!("enterprise IdP callback URL and listener specify different ports"); + } + Ok((ip, callback_port.or(registered_port))) +} + +fn validate_enterprise_credentials(stored: &StoredOAuthTokens) -> Result<()> { + if !stored.has_refresh_token() { + bail!("enterprise IdP login did not return a refresh token"); + } + let assertion = stored + .token_response + .0 + .extra_fields() + .0 + .get("id_token") + .and_then(serde_json::Value::as_str) + .ok_or_else(|| anyhow!("enterprise IdP login did not return an OIDC identity assertion"))?; + validate_oidc_identity_assertion( + assertion, + stored.issuer.as_deref().unwrap_or(&stored.url), + &stored.client_id, + ) + .map_err(|_| anyhow!("enterprise IdP returned an invalid OIDC identity assertion")) +} + +pub(crate) async fn resolve_enterprise_authorization_manager( + issuer: &str, + http_client: Arc, +) -> Result<(AuthorizationManager, AuthorizationMetadata)> { + let (manager, metadata) = resolve_ema_idp_authorization_manager( + issuer, + Arc::new(EnterpriseOAuthHttpClient(http_client)), + ) + .await?; + validate_authorization_server_endpoints(&metadata)?; + validate_ema_oauth_endpoint( + &metadata.authorization_endpoint, + "enterprise IdP authorization endpoint", + )?; + Ok((manager, metadata)) +} + +pub(crate) fn enterprise_authorization_url(auth_url: &str) -> Result { + let mut url = Url::parse(auth_url)?; + let query = url::form_urlencoded::Serializer::new(String::new()) + .extend_pairs( + url.query_pairs() + .filter(|(key, _)| key != "resource" && key != "prompt"), + ) + .append_pair("prompt", "consent") + .finish(); + url.set_query(Some(&query)); + Ok(url.to_string()) +} + +/// rmcp supplies a resource indicator for MCP OAuth, but the independent OIDC +/// login must not request the IdP issuer as a protected-resource audience. +struct EnterpriseOAuthHttpClient(Arc); + +impl OAuthHttpClient for EnterpriseOAuthHttpClient { + fn execute(&self, mut request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + if request.request.method() == Method::POST + && request + .request + .headers() + .get(CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| { + value.split(';').next().is_some_and(|mime| { + mime.trim() + .eq_ignore_ascii_case("application/x-www-form-urlencoded") + }) + }) + && url::form_urlencoded::parse(request.request.body()) + .any(|(key, value)| key == "grant_type" && value == "authorization_code") + { + let body = url::form_urlencoded::Serializer::new(String::new()) + .extend_pairs( + url::form_urlencoded::parse(request.request.body()) + .filter(|(key, _)| key != "resource"), + ) + .finish(); + *request.request.body_mut() = body.into_bytes(); + request.request.headers_mut().remove(CONTENT_LENGTH); + } + self.0.execute(request) + } +} + +/// Enterprise registration and host-owned login settings. The flow always uses +/// OpenID Connect, strict issuer validation, and keyring-only credential storage. +pub struct EnterpriseOAuthLoginRequest<'a> { + pub credential_name: &'a str, + pub issuer: &'a str, + pub client_id: &'a str, + pub keyring_backend_kind: AuthKeyringBackendKind, + pub callback_port: Option, + pub callback_url: Option<&'a str>, + pub timeout_secs: Option, + pub http_client: Arc, + pub redirect_mode: StreamableHttpRedirectMode, +} + +#[cfg(test)] +#[path = "enterprise_oauth_login_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/enterprise_oauth_login_tests.rs b/codex-rs/rmcp-client/src/enterprise_oauth_login_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..b651d76e11329893b7a2bd08a1ceed5fe0a6f254 --- /dev/null +++ b/codex-rs/rmcp-client/src/enterprise_oauth_login_tests.rs @@ -0,0 +1,480 @@ +//! Public enterprise login exercises the real callback and keyring adapter in an isolated process. + +use std::any::Any; +use std::collections::HashMap; +use std::sync::Mutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use base64::Engine; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use codex_exec_server::RouteAwareHttpClient; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use keyring::credential::Credential; +use keyring::credential::CredentialApi; +use keyring::credential::CredentialBuilderApi; +use keyring::credential::CredentialPersistence; +use keyring::mock::MockCredential; +use oauth2::TokenResponse; +use pretty_assertions::assert_eq; +use serde_json::json; +use sha2::Digest; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tracing_test::traced_test; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::*; + +const SECRET: &str = "enterprise-secret-sentinel"; +const CREDENTIAL_NAME: &str = "ema-idp:enterprise-secret-sentinel"; + +#[path = "enterprise_oauth_logout_tests.rs"] +mod logout; + +async fn isolated_process(test_name: &str) -> Result { + const CHILD: &str = "CODEX_ENTERPRISE_LOGIN_TEST_CHILD"; + if std::env::var_os(CHILD).is_some() { + return Ok(false); + } + let home = tempfile::tempdir()?; + let output = tokio::process::Command::new(std::env::current_exe()?) + .args(["--exact", test_name, "--nocapture"]) + .env(CHILD, "1") + .env("CODEX_HOME", home.path()) + .current_dir(home.path()) + .output() + .await?; + assert!( + output.status.success(), + "{}\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + Ok(true) +} + +async fn login(issuer: &str, callback_url: Option<&str>) -> Result { + perform_enterprise_oauth_login_return_url(EnterpriseOAuthLoginRequest { + credential_name: CREDENTIAL_NAME, + issuer, + client_id: "enterprise-client", + keyring_backend_kind: AuthKeyringBackendKind::Direct, + callback_port: None, + callback_url, + timeout_secs: Some(5), + http_client: Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + redirect_mode: StreamableHttpRedirectMode::Legacy, + }) + .await +} + +async fn complete_login(issuer: &str) -> Result { + let handle = login(issuer, /*callback_url*/ None).await?; + callback( + &handle.authorization_url(), + issuer, + /*provider_error*/ false, + ) + .await?; + handle.wait().await +} + +async fn metadata(server: &MockServer, issuer: &str) { + Mock::given(method("GET")).and(path("/.well-known/oauth-authorization-server/idp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "issuer": issuer, "authorization_endpoint": format!("{issuer}/authorize?prompt=none"), + "token_endpoint": format!("{issuer}/token"), "token_endpoint_auth_methods_supported": ["none"], + }))).mount(server).await; +} + +async fn callback(authorization_url: &str, issuer: &str, provider_error: bool) -> Result<()> { + let query = Url::parse(authorization_url)? + .query_pairs() + .into_owned() + .collect::>(); + assert_eq!(query.get("prompt").map(String::as_str), Some("consent")); + assert!(!query.contains_key("resource")); + let mut callback = Url::parse(&query["redirect_uri"])?; + let mut pairs = callback.query_pairs_mut(); + if provider_error { + pairs + .append_pair("error", SECRET) + .append_pair("error_description", SECRET); + } else { + pairs + .append_pair("code", SECRET) + .append_pair("state", &query["state"]) + .append_pair("iss", issuer); + } + drop(pairs); + let port = callback.port().expect("actual listener port"); + let mut stream = tokio::net::TcpStream::connect(("127.0.0.1", port)).await?; + stream + .write_all( + format!( + "GET {}?{} HTTP/1.1\r\nHost: 127.0.0.1:{port}\r\nConnection: close\r\n\r\n", + callback.path(), + callback.query().expect("query") + ) + .as_bytes(), + ) + .await?; + stream.read_to_end(&mut Vec::new()).await?; + Ok(()) +} + +#[tokio::test] +#[traced_test] +async fn enterprise_callback_errors_and_sdk_logs_exclude_credentials() -> Result<()> { + if isolated_process( + "enterprise_oauth_login::tests::enterprise_callback_errors_and_sdk_logs_exclude_credentials", + ) + .await? + { + return Ok(()); + } + for (callback_error, token_success) in [(true, false), (false, false), (false, true)] { + let server = MockServer::start().await; + let issuer = format!("{}/idp", server.uri()); + metadata(&server, &issuer).await; + let response = if token_success { + ResponseTemplate::new(200).set_body_json(json!({ + "access_token": SECRET, "refresh_token": SECRET, "id_token": SECRET, "token_type": "Bearer", + })) + } else { + ResponseTemplate::new(400) + .set_body_json(json!({"error": "invalid_grant", "error_description": SECRET})) + }; + Mock::given(method("POST")) + .and(path("/idp/token")) + .respond_with(response) + .expect(u64::from(!callback_error)) + .mount(&server) + .await; + let login = login(&issuer, /*callback_url*/ None).await?; + callback(&login.authorization_url(), &issuer, callback_error).await?; + tracing::trace!(target: "rmcp::transport::auth", "SDK trace capture enabled"); + let error = login.wait().await.err().expect("reject provider response"); + assert!(!format!("{error} {error:?} {error:#}").contains(SECRET)); + assert!(logs_contain("SDK trace capture enabled")); + assert!(!logs_contain(SECRET)); + } + Ok(()) +} + +#[test] +fn enterprise_callback_requires_loopback() -> Result<()> { + let issuer = "https://idp.example"; + let client_id = Some("enterprise-client"); + for (callback, port, expected_ip, expected_port) in [ + (None, None, "127.0.0.1", None), + (None, Some(8080), "127.0.0.1", Some(8080)), + (Some("http://localhost/callback"), None, "127.0.0.1", None), + (Some("http://127.0.0.2/callback"), None, "127.0.0.2", None), + (Some("http://[::1]/callback"), None, "::1", None), + ( + Some("http://localhost:8080/callback"), + None, + "127.0.0.1", + Some(8080), + ), + ( + Some("http://localhost:8080/callback"), + Some(8080), + "127.0.0.1", + Some(8080), + ), + ] { + assert_eq!( + enterprise_callback_settings(issuer, client_id, callback, port)?, + (expected_ip.parse::()?, expected_port), + ); + } + for callback in [ + "http://0.0.0.0/callback", + "http://[::]/callback", + "https://127.0.0.1/callback", + "http://remote.example/callback", + ] { + assert!( + enterprise_callback_settings( + issuer, + client_id, + Some(callback), + /*callback_port*/ None + ) + .is_err() + ); + } + assert!( + enterprise_callback_settings( + issuer, + client_id, + Some("http://localhost:8080/callback"), + Some(9090), + ) + .is_err() + ); + Ok(()) +} + +#[tokio::test] +#[traced_test] +async fn enterprise_public_api_storage_and_privacy() -> Result<()> { + if isolated_process("enterprise_oauth_login::tests::enterprise_public_api_storage_and_privacy") + .await? + { + return Ok(()); + } + let keyring = TestKeyring::default(); + keyring::set_default_credential_builder(Box::new(keyring.clone())); + let server = MockServer::start().await; + let issuer = format!("{}/idp", server.uri()); + metadata(&server, &issuer).await; + assert!( + login(&issuer, Some(&format!("{}/callback", server.uri()))) + .await + .is_err(), + "an occupied registered callback port must not silently move" + ); + let assertion = format!( + "{}.{}.signature", + URL_SAFE_NO_PAD.encode(r#"{"alg":"ES256"}"#), + URL_SAFE_NO_PAD.encode( + json!({"iss":issuer,"sub":"user","aud":"enterprise-client","exp":4102444800_u64}) + .to_string() + ) + ); + Mock::given(method("POST")) + .and(path("/idp/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token":SECRET,"refresh_token":SECRET,"id_token":assertion,"token_type":"Bearer", + }))) + .mount(&server) + .await; + + // Callback, exchange and DefaultKeyringStore all run through public entrypoints. + for (login_index, callback_url) in ["http://localhost/callback", "http://127.0.0.1/callback"] + .into_iter() + .enumerate() + { + let login = login(&issuer, Some(callback_url)).await?; + let authorization_url = login.authorization_url(); + let authorization_query = Url::parse(&authorization_url)? + .query_pairs() + .into_owned() + .collect::>(); + callback(&authorization_url, &issuer, /*provider_error*/ false).await?; + let credentials = login.wait().await?; + let token_requests = server + .received_requests() + .await + .expect("recorded OAuth requests") + .into_iter() + .filter(|request| { + request.method == http::Method::POST && request.url.path() == "/idp/token" + }) + .collect::>(); + assert_eq!(token_requests.len(), login_index + 1); + let token_form = url::form_urlencoded::parse(&token_requests[login_index].body) + .into_owned() + .collect::>(); + assert_eq!( + token_form.get("grant_type").map(String::as_str), + Some("authorization_code") + ); + assert!(!token_form.contains_key("resource")); + assert_eq!( + token_form.get("redirect_uri"), + Some( + authorization_query + .get("redirect_uri") + .expect("authorization redirect URI") + ) + ); + assert_eq!( + authorization_query + .get("code_challenge_method") + .map(String::as_str), + Some("S256") + ); + let verifier = token_form.get("code_verifier").expect("PKCE verifier"); + let challenge = URL_SAFE_NO_PAD.encode(sha2::Sha256::digest(verifier.as_bytes())); + assert_eq!(authorization_query.get("code_challenge"), Some(&challenge)); + assert!( + keyring + .values + .lock() + .unwrap() + .values() + .all(|value| value.get_secret().is_err()) + ); + let attempt = tokio::sync::Mutex::new(Some("active")); + { + let mut authority = credentials + .commit_if(|| async { Some(attempt.lock().await) }) + .await?; + assert!( + attempt.try_lock().is_err(), + "commit retains the attempt gate until the caller retires it" + ); + *authority = None; + } + assert_eq!(*attempt.lock().await, None); + let stored = tracing::subscriber::with_default( + tracing::subscriber::NoSubscriber::default(), + || { + crate::stored_oauth_credentials( + CREDENTIAL_NAME, + &issuer, + OAuthCredentialsStoreMode::Keyring, + AuthKeyringBackendKind::Direct, + ) + }, + )? + .expect("stored grant"); + assert_eq!( + stored + .token_response + .0 + .refresh_token() + .map(|token| token.secret().as_str()), + Some(SECRET) + ); + assert!( + delete_enterprise_oauth_tokens( + CREDENTIAL_NAME, + &issuer, + AuthKeyringBackendKind::Direct + ) + .await? + ); + } + + // Cancellation while blocked on the actual credential lock cannot leave a + // detached persistence worker that writes after the other process releases it. + let canceled = complete_login(&issuer).await?; + let guard = EnterpriseOAuthCredentialGuard::acquire( + CREDENTIAL_NAME, + &issuer, + AuthKeyringBackendKind::Direct, + ) + .await?; + let mut commit = Box::pin(canceled.commit_if(|| async { Some(()) })); + assert!(futures::poll!(&mut commit).is_pending()); + drop(commit); + drop(guard); + assert!( + keyring + .values + .lock() + .unwrap() + .values() + .all(|value| value.get_secret().is_err()) + ); + + // Rejected old attempts neither write nor delete a newer grant. + let old = complete_login(&issuer).await?; + complete_login(&issuer) + .await? + .commit_if(|| async { Some(()) }) + .await?; + assert!(old.commit_if(|| async { None::<()> }).await.is_err()); + assert!( + keyring + .values + .lock() + .unwrap() + .values() + .any(|value| value.get_secret().is_ok()) + ); + + // Inject raw account identifiers into the actual keyring adapter's error chain. + keyring.fail.store(true, Ordering::SeqCst); + let save_error = complete_login(&issuer) + .await? + .commit_if(|| async { Some(()) }) + .await + .unwrap_err(); + let delete_error = + delete_enterprise_oauth_tokens(CREDENTIAL_NAME, &issuer, AuthKeyringBackendKind::Direct) + .await + .unwrap_err(); + for error in [save_error, delete_error] { + assert!(!format!("{error} {error:?} {error:#}").contains(SECRET)); + } + keyring.fail.store(false, Ordering::SeqCst); + let home = std::path::PathBuf::from(std::env::var("CODEX_HOME")?); + assert!( + !home.join(".credentials.json").exists(), + "enterprise storage never falls back to plaintext" + ); + let locks = home.join("mcp-oauth-locks"); + std::fs::rename(&locks, home.join("held-locks"))?; + std::fs::write(&locks, b"not a directory")?; + let lock_error = + delete_enterprise_oauth_tokens(CREDENTIAL_NAME, &issuer, AuthKeyringBackendKind::Direct) + .await + .unwrap_err(); + assert!(!format!("{lock_error} {lock_error:?} {lock_error:#}").contains(SECRET)); + assert!(!logs_contain(SECRET)); + Ok(()) +} + +#[derive(Clone, Default)] +struct TestKeyring { + values: Arc>>>, + fail: Arc, +} + +struct TestCredential(Arc); + +impl CredentialApi for TestCredential { + fn set_secret(&self, secret: &[u8]) -> keyring::Result<()> { + self.0.set_secret(secret) + } + fn get_secret(&self) -> keyring::Result> { + self.0.get_secret() + } + fn delete_credential(&self) -> keyring::Result<()> { + self.0.delete_credential() + } + fn as_any(&self) -> &dyn Any { + self + } +} + +impl CredentialBuilderApi for TestKeyring { + fn build( + &self, + _target: Option<&str>, + _service: &str, + user: &str, + ) -> keyring::Result> { + if self.fail.load(Ordering::SeqCst) { + return Err(keyring::Error::Invalid("account".into(), user.into())); + } + Ok(Box::new(TestCredential( + self.values + .lock() + .unwrap() + .entry(user.into()) + .or_default() + .clone(), + ))) + } + fn as_any(&self) -> &dyn Any { + self + } + fn persistence(&self) -> CredentialPersistence { + CredentialPersistence::ProcessOnly + } +} diff --git a/codex-rs/rmcp-client/src/enterprise_oauth_logout_tests.rs b/codex-rs/rmcp-client/src/enterprise_oauth_logout_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..974ba726956544b1d7ebec8a8bdd2f165160d178 --- /dev/null +++ b/codex-rs/rmcp-client/src/enterprise_oauth_logout_tests.rs @@ -0,0 +1,184 @@ +//! Cross-process logout exercises staged login and the production keyring adapter. + +use std::fs; +use std::path::PathBuf; + +use pretty_assertions::assert_eq; +use sha2::Digest; +use sha2::Sha256; + +use super::*; + +const TEST: &str = + "enterprise_oauth_login::tests::logout::logout_invalidates_pending_login_across_processes"; +const LOGOUT_ISSUER: &str = "CODEX_ENTERPRISE_LOGOUT_TEST_ISSUER"; + +#[tokio::test] +async fn logout_invalidates_pending_login_across_processes() -> Result<()> { + if isolated_process(TEST).await? { + return Ok(()); + } + let home = PathBuf::from(std::env::var("CODEX_HOME")?); + let keyring = FileKeyring(home.join("test-keyring")); + fs::create_dir_all(&keyring.0)?; + keyring::set_default_credential_builder(Box::new(keyring)); + if let Ok(issuer) = std::env::var(LOGOUT_ISSUER) { + delete_enterprise_oauth_tokens(CREDENTIAL_NAME, &issuer, AuthKeyringBackendKind::Direct) + .await?; + return Ok(()); + } + let server = MockServer::start().await; + let issuer = format!("{}/idp", server.uri()); + metadata(&server, &issuer).await; + let assertion = format!( + "{}.{}.signature", + URL_SAFE_NO_PAD.encode(r#"{"alg":"ES256"}"#), + URL_SAFE_NO_PAD.encode( + json!({"iss":issuer,"sub":"user","aud":"enterprise-client","exp":4102444800_u64}) + .to_string() + ) + ); + Mock::given(method("POST")) + .and(path("/idp/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token":SECRET,"refresh_token":SECRET,"id_token":assertion,"token_type":"Bearer", + }))) + .mount(&server) + .await; + + // The first logout has no grant to delete. The second deletes the fresh grant + // from the previous iteration. Both must invalidate pending callbacks and staged grants. + for _ in 0..2 { + let pending = login(&issuer, /*callback_url*/ None).await?; + let staged = complete_login(&issuer).await?; + let output = tokio::process::Command::new(std::env::current_exe()?) + .args(["--exact", TEST, "--nocapture"]) + .env(LOGOUT_ISSUER, &issuer) + .output() + .await?; + assert!( + output.status.success(), + "{}\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + assert!(stored(&issuer)?.is_none()); + // The logout process has exited; its invalidation must survive that exit. + assert!(staged.commit_if(|| async { Some(()) }).await.is_err()); + assert!(stored(&issuer)?.is_none()); + + // A fresh login after logout remains usable, including for the same account. + complete_login(&issuer) + .await? + .commit_if(|| async { Some(()) }) + .await?; + let fresh = stored(&issuer)?.expect("fresh grant"); + callback( + &pending.authorization_url(), + &issuer, + /*provider_error*/ false, + ) + .await?; + assert!( + pending + .wait() + .await? + .commit_if(|| async { Some(()) }) + .await + .is_err() + ); + assert_eq!(stored(&issuer)?, Some(fresh)); + } + + let stale = complete_login(&issuer).await?; + let generation_path = fs::read_dir(home.join("mcp-oauth-locks"))? + .filter_map(std::result::Result::ok) + .map(|entry| entry.path()) + .find(|path| { + path.extension() + .is_some_and(|ext| ext == "enterprise-generation") + }) + .expect("persistent generation"); + fs::write(&generation_path, b"incomplete write")?; + assert!(stale.commit_if(|| async { Some(()) }).await.is_err()); + assert!(login(&issuer, /*callback_url*/ None).await.is_err()); + + // A lost marker never revives a pre-logout attempt or a pre-reset generation. + fs::remove_file(&generation_path)?; + let stale = complete_login(&issuer).await?; + fs::remove_file(&generation_path)?; + complete_login(&issuer) + .await? + .commit_if(|| async { Some(()) }) + .await?; + assert!(stale.commit_if(|| async { Some(()) }).await.is_err()); + + // Metadata failures are logout errors and must leave the stored grant intact. + let before = stored(&issuer)?; + fs::remove_file(&generation_path)?; + fs::create_dir(&generation_path)?; + assert!( + delete_enterprise_oauth_tokens(CREDENTIAL_NAME, &issuer, AuthKeyringBackendKind::Direct) + .await + .is_err() + ); + assert_eq!(stored(&issuer)?, before); + assert!(!home.join(".credentials.json").exists()); + Ok(()) +} + +fn stored(issuer: &str) -> Result> { + crate::stored_oauth_credentials( + CREDENTIAL_NAME, + issuer, + OAuthCredentialsStoreMode::Keyring, + AuthKeyringBackendKind::Direct, + ) +} + +// Only synthetic fixture credentials are persisted here. Every process still uses +// DefaultKeyringStore and the production credential lock, serializer and logout API. +struct FileKeyring(PathBuf); +struct FileCredential(PathBuf); + +impl CredentialBuilderApi for FileKeyring { + fn build( + &self, + _target: Option<&str>, + _service: &str, + user: &str, + ) -> keyring::Result> { + let name = format!("{:x}", Sha256::digest(user.as_bytes())); + Ok(Box::new(FileCredential(self.0.join(name)))) + } + + fn as_any(&self) -> &dyn Any { + self + } +} + +impl CredentialApi for FileCredential { + fn set_secret(&self, secret: &[u8]) -> keyring::Result<()> { + fs::write(&self.0, secret).map_err(keyring_error) + } + + fn get_secret(&self) -> keyring::Result> { + fs::read(&self.0).map_err(keyring_error) + } + + fn delete_credential(&self) -> keyring::Result<()> { + fs::remove_file(&self.0).map_err(keyring_error) + } + + fn as_any(&self) -> &dyn Any { + self + } +} + +fn keyring_error(error: std::io::Error) -> keyring::Error { + if error.kind() == std::io::ErrorKind::NotFound { + keyring::Error::NoEntry + } else { + keyring::Error::PlatformFailure(Box::new(error)) + } +} diff --git a/codex-rs/rmcp-client/src/event_notification_transport.rs b/codex-rs/rmcp-client/src/event_notification_transport.rs new file mode 100644 index 0000000000000000000000000000000000000000..f997376b4cf4be989a400679c24e2d2728c2d6fa --- /dev/null +++ b/codex-rs/rmcp-client/src/event_notification_transport.rs @@ -0,0 +1,264 @@ +use std::collections::HashMap; +use std::future::Future; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; + +use rmcp::model::ClientRequest; +use rmcp::model::CustomNotification; +use rmcp::model::JsonRpcMessage; +use rmcp::model::RequestId; +use rmcp::model::ServerNotification; +use rmcp::service::RoleClient; +use rmcp::service::RxJsonRpcMessage; +use rmcp::service::TxJsonRpcMessage; +use rmcp::transport::IntoTransport; +use rmcp::transport::Transport; +use tokio::sync::OwnedSemaphorePermit; +use tokio::sync::Semaphore; +use tokio::sync::mpsc; +use tracing::warn; + +pub(crate) const MAX_EVENT_NOTIFICATION_BYTES: usize = 1024 * 1024; +const MAX_QUEUED_EVENT_NOTIFICATION_BYTES: usize = 2 * 1024 * 1024; + +#[derive(Clone)] +pub(crate) struct EventNotificationSender { + notifications: mpsc::UnboundedSender, + available_bytes: Arc, +} + +pub struct EventNotificationReceiver { + notifications: mpsc::UnboundedReceiver, + available_bytes: Arc, +} + +struct QueuedEventNotification { + notification: CustomNotification, + _bytes: OwnedSemaphorePermit, +} + +pub(crate) fn event_notification_channel() -> (EventNotificationSender, EventNotificationReceiver) { + let (notifications_tx, notifications_rx) = mpsc::unbounded_channel(); + let available_bytes = Arc::new(Semaphore::new(MAX_QUEUED_EVENT_NOTIFICATION_BYTES)); + + ( + EventNotificationSender { + notifications: notifications_tx, + available_bytes: Arc::clone(&available_bytes), + }, + EventNotificationReceiver { + notifications: notifications_rx, + available_bytes, + }, + ) +} + +impl EventNotificationSender { + fn send(&self, notification: CustomNotification, notification_bytes: usize) -> Result<(), ()> { + let bytes = u32::try_from(notification_bytes).map_err(|_| ())?; + let permit = Arc::clone(&self.available_bytes) + .try_acquire_many_owned(bytes) + .map_err(|_| ())?; + self.notifications + .send(QueuedEventNotification { + notification, + _bytes: permit, + }) + .map_err(|_| ()) + } + + fn close(&self) { + self.available_bytes.close(); + } +} + +impl EventNotificationReceiver { + pub async fn recv(&mut self) -> Option { + self.notifications + .recv() + .await + .map(|queued| queued.notification) + } +} + +impl Drop for EventNotificationReceiver { + fn drop(&mut self) { + self.available_bytes.close(); + } +} + +/// Consumes Plugin Runtime event notifications before rmcp can log their payloads. +pub(crate) fn capture_event_notifications( + transport: T, +) -> impl Transport + 'static +where + T: IntoTransport, + E: std::error::Error + Send + Sync + 'static, +{ + EventNotificationTransport { + inner: transport.into_transport(), + routes: Arc::default(), + } +} + +struct EventNotificationTransport { + inner: T, + routes: Arc>>, +} + +impl Transport for EventNotificationTransport +where + T: Transport + 'static, +{ + type Error = T::Error; + + fn send( + &mut self, + mut message: TxJsonRpcMessage, + ) -> impl Future> + Send + 'static { + let request_route = match &mut message { + JsonRpcMessage::Request(envelope) => match &mut envelope.request { + ClientRequest::CustomRequest(request) if request.method == "events/stream" => { + request + .extensions + .remove::() + .map(|sender| (envelope.id.clone(), sender)) + } + _ => None, + }, + JsonRpcMessage::Notification(envelope) => { + if let rmcp::model::ClientNotification::CancelledNotification(cancelled) = + &envelope.notification + && let Some(request_id) = cancelled.params.request_id.as_ref() + && let Some(route) = self + .routes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove(request_id) + { + route.close(); + } + None + } + JsonRpcMessage::Response(_) | JsonRpcMessage::Error(_) => None, + }; + + let route_id = request_route + .as_ref() + .map(|(request_id, _)| request_id.clone()); + if let Some((request_id, sender)) = request_route + && let Some(previous) = self + .routes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .insert(request_id, sender) + { + previous.close(); + } + + let routes = Arc::clone(&self.routes); + let send = self.inner.send(message); + async move { + let result = send.await; + if result.is_err() + && let Some(route_id) = route_id + && let Some(route) = routes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove(&route_id) + { + route.close(); + } + result + } + } + + async fn receive(&mut self) -> Option> { + loop { + let message = self.inner.receive().await?; + if let Some(request_id) = response_id(&message) { + if let Some(route) = self + .routes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove(request_id) + { + route.close(); + } + return Some(message); + } + + let JsonRpcMessage::Notification(envelope) = &message else { + return Some(message); + }; + let ServerNotification::CustomNotification(notification) = &envelope.notification + else { + return Some(message); + }; + if !notification.method.starts_with("notifications/events/") { + return Some(message); + } + let Some(subscription_id) = + rmcp::model::GetMeta::get_meta(notification).subscription_id() + else { + continue; + }; + + let route = self + .routes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .get(&subscription_id) + .cloned(); + let Some(route) = route else { + continue; + }; + let notification_bytes = serde_json::to_vec(&message) + .map(|message| message.len()) + .unwrap_or(usize::MAX); + if notification_bytes > MAX_EVENT_NOTIFICATION_BYTES { + warn!( + notification_bytes, + "discarding oversized MCP event notification" + ); + continue; + } + if route + .send(notification.clone(), notification_bytes) + .is_err() + && let Some(route) = self + .routes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove(&subscription_id) + { + let _ = route.send( + CustomNotification::new( + "notifications/events/terminated", + /*params*/ None, + ), + /*notification_bytes*/ 0, + ); + route.close(); + } + } + } + + async fn close(&mut self) -> Result<(), Self::Error> { + let routes = + std::mem::take(&mut *self.routes.lock().unwrap_or_else(PoisonError::into_inner)); + for route in routes.into_values() { + route.close(); + } + self.inner.close().await + } +} + +fn response_id(message: &RxJsonRpcMessage) -> Option<&RequestId> { + match message { + JsonRpcMessage::Response(response) => Some(&response.id), + JsonRpcMessage::Error(error) => error.id.as_ref(), + JsonRpcMessage::Request(_) | JsonRpcMessage::Notification(_) => None, + } +} diff --git a/codex-rs/rmcp-client/src/executor_process_transport.rs b/codex-rs/rmcp-client/src/executor_process_transport.rs new file mode 100644 index 0000000000000000000000000000000000000000..b1fddbdf7518862ff3b2b853fd7ddb16a970d181 --- /dev/null +++ b/codex-rs/rmcp-client/src/executor_process_transport.rs @@ -0,0 +1,527 @@ +//! rmcp transport adapter for an executor-managed MCP stdio process. +//! +//! This module owns the lower-level byte translation after +//! `stdio_server_launcher` has already started a process through +//! `ExecBackend::start`. It does not choose where the MCP server runs and it +//! does not implement MCP lifecycle behavior. MCP protocol ownership stays in +//! `RmcpClient` and rmcp: +//! +//! 1. rmcp serializes a JSON-RPC message and calls [`Transport::send`]. +//! 2. This transport appends the stdio newline delimiter and writes those bytes +//! to executor `process/write`. +//! 3. The executor writes the bytes to the child process stdin. +//! 4. The child writes newline-delimited JSON-RPC messages to stdout. +//! 5. The executor reports stdout bytes through pushed process events. +//! 6. This transport buffers stdout until it has one full line, deserializes +//! that line, and returns the rmcp message from [`Transport::receive`]. +//! +//! Stderr is deliberately not part of the MCP byte stream. It is logged for +//! diagnostics only, matching the local stdio implementation. + +use std::future::Future; +use std::io; +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; + +use bytes::BytesMut; +use codex_exec_server::ExecOutputStream; +use codex_exec_server::ExecProcess; +use codex_exec_server::ExecProcessEvent; +use codex_exec_server::ExecProcessEventReceiver; +use codex_exec_server::ProcessId; +use codex_exec_server::ProcessOutputChunk; +use codex_exec_server::WriteStatus; +use memchr::memchr; +use rmcp::service::RoleClient; +use rmcp::service::RxJsonRpcMessage; +use rmcp::service::TxJsonRpcMessage; +use rmcp::transport::Transport; +use serde_json::to_vec; +use tokio::runtime::Handle; +use tokio::sync::Semaphore; +use tokio::sync::broadcast; +use tracing::debug; +use tracing::info; +use tracing::warn; + +static PROCESS_COUNTER: AtomicUsize = AtomicUsize::new(1); +// Tool results can make valid MCP responses large, so keep the protocol +// ceiling well above ordinary messages while still bounding hostile input. +const MAX_MCP_STDOUT_LINE_BYTES: usize = 8 * 1024 * 1024; +// Stderr is diagnostic only and does not need the protocol stream's allowance. +const MAX_MCP_STDERR_LINE_BYTES: usize = 1024 * 1024; + +#[cfg_attr(test, derive(Debug, PartialEq, Eq))] +struct LineBuffer { + bytes: BytesMut, + /// Prefix already scanned and known not to contain a newline. + scanned_len: usize, + /// Bytes after the last buffered newline. + pending_line_bytes: usize, + max_line_bytes: usize, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct LineTooLong { + max_line_bytes: usize, +} + +impl Default for LineBuffer { + fn default() -> Self { + Self::new(MAX_MCP_STDOUT_LINE_BYTES) + } +} + +impl LineBuffer { + fn new(max_line_bytes: usize) -> Self { + Self { + bytes: BytesMut::new(), + scanned_len: 0, + pending_line_bytes: 0, + max_line_bytes, + } + } + + fn extend_from_slice(&mut self, bytes: &[u8]) -> Result<(), LineTooLong> { + let mut remaining = bytes; + while let Some(newline_index) = memchr(b'\n', remaining) { + if newline_index > self.max_line_bytes.saturating_sub(self.pending_line_bytes) { + self.discard_pending_line(); + return Err(LineTooLong { + max_line_bytes: self.max_line_bytes, + }); + } + + let segment_len = newline_index + 1; + self.bytes.extend_from_slice(&remaining[..segment_len]); + self.pending_line_bytes = 0; + remaining = &remaining[segment_len..]; + } + if remaining.len() > self.max_line_bytes.saturating_sub(self.pending_line_bytes) { + self.discard_pending_line(); + return Err(LineTooLong { + max_line_bytes: self.max_line_bytes, + }); + } + + self.bytes.extend_from_slice(remaining); + self.pending_line_bytes += remaining.len(); + Ok(()) + } + + fn discard_pending_line(&mut self) { + let complete_line_bytes = self.bytes.len().saturating_sub(self.pending_line_bytes); + self.bytes.truncate(complete_line_bytes); + self.scanned_len = self.scanned_len.min(complete_line_bytes); + self.pending_line_bytes = 0; + } + + fn take_line(&mut self) -> Option { + let Some(relative_index) = memchr(b'\n', &self.bytes[self.scanned_len..]) else { + self.scanned_len = self.bytes.len(); + return None; + }; + + let newline_index = self.scanned_len + relative_index; + let mut line = self.bytes.split_to(newline_index + 1); + line.truncate(newline_index); + self.scanned_len = 0; + Some(line) + } + + fn take_remaining(&mut self) -> Option { + if self.bytes.is_empty() { + return None; + } + + self.scanned_len = 0; + self.pending_line_bytes = 0; + Some(self.bytes.split()) + } + + fn clear(&mut self) { + self.bytes = BytesMut::new(); + self.scanned_len = 0; + self.pending_line_bytes = 0; + } +} + +// Remote public implementation. + +/// A client-side rmcp transport backed by an executor-managed process. +/// +/// The orchestrator owns this value and calls rmcp on it. The process it wraps +/// may be local or remote depending on the `ExecBackend` used to create it, but +/// for remote MCP stdio the process lives on the executor and all interaction +/// crosses the executor process RPC boundary. +pub(super) struct ExecutorProcessTransport { + /// Logical process handle returned by the executor process API. + /// + /// `write` forwards stdin bytes. `terminate` stops the child when rmcp + /// closes the transport. + process: Arc, + + /// Prevents concurrent rmcp send futures from issuing overlapping stdin writes. + /// The single-slot semaphore gives mutex semantics while its permit can safely cross `.await`. + stdin_write_semaphore: Arc, + + /// Pushed output/lifecycle stream for the process. + /// + /// The executor process API still supports retained-output reads, but MCP + /// stdio is naturally streaming. This receiver lets rmcp wait for stdout + /// chunks without issuing `process/read` after each output notification. + events: ExecProcessEventReceiver, + + /// Human-readable program name used only in diagnostics. + program_name: String, + + /// Buffered child stdout bytes that have not yet formed a complete + /// newline-delimited JSON-RPC message. + stdout: LineBuffer, + + /// Buffered stderr bytes for diagnostic logging. + stderr: LineBuffer, + + /// Whether the executor has reported process closure or a terminal + /// subscription failure. Once closed, any remaining partial stdout line is + /// flushed once and then rmcp receives EOF. + closed: bool, + + /// Whether this transport already asked the executor to terminate the MCP + /// server process. + terminated: bool, + + /// Highest executor process event sequence observed by this transport. + /// + /// When the pushed event stream lags, use this as the retained-output read + /// cursor to recover missed stdout/stderr chunks from the executor. + last_seq: u64, +} + +impl ExecutorProcessTransport { + pub(super) fn new(process: Arc, program_name: String) -> Self { + // Subscribe before returning the transport to rmcp. Some test servers + // can emit output or exit quickly after `process/start`, and the + // process event log will replay anything that landed before this + // subscriber was attached. + let events = process.subscribe_events(); + Self { + process, + stdin_write_semaphore: Arc::new(Semaphore::new(1)), + events, + program_name, + stdout: LineBuffer::default(), + stderr: LineBuffer::new(MAX_MCP_STDERR_LINE_BYTES), + closed: false, + terminated: false, + last_seq: 0, + } + } + + pub(super) fn next_process_id() -> ProcessId { + // Process IDs are logical handles scoped to the executor connection, + // not OS pids. A monotonic client-side id is enough to avoid + // collisions between MCP servers started in the same session. + let index = PROCESS_COUNTER.fetch_add(1, Ordering::Relaxed); + ProcessId::from(format!("mcp-stdio-{index}")) + } +} + +impl Transport for ExecutorProcessTransport { + type Error = io::Error; + + fn send( + &mut self, + item: TxJsonRpcMessage, + ) -> impl Future> + Send + 'static { + let process = Arc::clone(&self.process); + let stdin_write_semaphore = Arc::clone(&self.stdin_write_semaphore); + async move { + let _stdin_write_permit = stdin_write_semaphore + .acquire() + .await + .map_err(io::Error::other)?; + // rmcp hands us a structured JSON-RPC message. Stdio transport on + // the wire is JSON plus one newline delimiter. + let mut bytes = to_vec(&item).map_err(io::Error::other)?; + bytes.push(b'\n'); + let response = process.write(bytes).await.map_err(io::Error::other)?; + match response.status { + WriteStatus::Accepted => Ok(()), + WriteStatus::UnknownProcess => { + Err(io::Error::new(io::ErrorKind::BrokenPipe, "unknown process")) + } + WriteStatus::StdinClosed => { + Err(io::Error::new(io::ErrorKind::BrokenPipe, "stdin closed")) + } + WriteStatus::Starting => Err(io::Error::new( + io::ErrorKind::WouldBlock, + "process is starting", + )), + } + } + } + + fn receive(&mut self) -> impl Future>> + Send { + self.receive_message() + } + + async fn close(&mut self) -> std::result::Result<(), Self::Error> { + self.process.terminate().await.map_err(io::Error::other)?; + self.terminated = true; + Ok(()) + } +} + +impl ExecutorProcessTransport { + async fn receive_message(&mut self) -> Option> { + loop { + // rmcp stdio framing is line-oriented JSON. We first drain any + // complete line already buffered from an earlier process event. + if let Some(message) = self.take_stdout_message(/*allow_partial*/ self.closed) { + return Some(message); + } + if self.closed { + self.flush_stderr(); + return None; + } + + match self.events.recv().await { + Ok(ExecProcessEvent::Output(chunk)) => { + // The executor pushes raw process bytes. This is the only + // place where those bytes are split back into the stdout + // protocol stream and stderr diagnostics. + self.push_process_output_if_new(chunk); + } + Ok(ExecProcessEvent::Exited { seq, .. }) => { + self.note_seq(seq); + // Wait for `Closed` before ending the rmcp stream so any + // output flushed during process shutdown can still be + // decoded into JSON-RPC messages. + } + Ok(ExecProcessEvent::Closed { seq }) => { + self.note_seq(seq); + self.closed = true; + } + Ok(ExecProcessEvent::Failed(message)) => { + warn!( + "Remote MCP server process failed ({}): {message}", + self.program_name + ); + self.closed = true; + } + Err(broadcast::error::RecvError::Lagged(skipped)) => { + warn!( + "Remote MCP server output stream lagged ({}): skipped {skipped} events", + self.program_name + ); + if let Err(error) = self.recover_lagged_events().await { + warn!( + "Failed to recover remote MCP server output stream ({}): {error}", + self.program_name + ); + self.closed = true; + } + } + Err(broadcast::error::RecvError::Closed) => { + self.closed = true; + } + } + } + } + + fn note_seq(&mut self, seq: u64) { + self.last_seq = self.last_seq.max(seq); + } + + fn should_accept_seq(&mut self, seq: u64) -> bool { + if seq <= self.last_seq { + return false; + } + self.last_seq = seq; + true + } + + async fn recover_lagged_events(&mut self) -> io::Result<()> { + let response = self + .process + .read( + Some(self.last_seq), + /*max_bytes*/ None, + /*wait_ms*/ Some(0), + ) + .await + .map_err(io::Error::other)?; + for chunk in response.chunks { + let expected_seq = self.last_seq.saturating_add(1); + if chunk.seq > expected_seq { + return Err(self.close_for_lost_output(expected_seq, chunk.seq)); + } + self.push_process_output_if_new(chunk); + if self.closed { + return Ok(()); + } + } + // Process reads include output chunks but not the sequenced `Exited` + // and `Closed` events. Account for those terminal events without + // allowing an evicted output chunk to be silently spliced into MCP. + let terminal_event_count = u64::from(response.exited) + u64::from(response.closed); + let next_output_seq = response.next_seq.saturating_sub(terminal_event_count); + let expected_seq = self.last_seq.saturating_add(1); + if next_output_seq > expected_seq { + return Err(self.close_for_lost_output(expected_seq, next_output_seq)); + } + self.last_seq = self.last_seq.max(response.next_seq.saturating_sub(1)); + if let Some(message) = response.failure { + warn!( + "Remote MCP server process failed ({}): {message}", + self.program_name + ); + self.closed = true; + } else if response.closed { + self.closed = true; + } + Ok(()) + } + + fn close_for_lost_output(&mut self, expected_seq: u64, received_seq: u64) -> io::Error { + self.stdout.clear(); + self.stderr.clear(); + self.closed = true; + io::Error::new( + io::ErrorKind::InvalidData, + format!( + "remote MCP server output stream lost process events: expected sequence {expected_seq}, received {received_seq}" + ), + ) + } + + fn push_process_output_if_new(&mut self, chunk: ProcessOutputChunk) { + if !self.should_accept_seq(chunk.seq) { + return; + } + self.push_process_output(chunk); + } + + fn push_process_output(&mut self, chunk: ProcessOutputChunk) { + let bytes = chunk.chunk.into_inner(); + match chunk.stream { + // MCP stdio uses stdout as the protocol stream. PTY output is + // accepted defensively because the executor process API has a + // unified stream enum, but remote MCP starts with `tty=false`. + ExecOutputStream::Stdout | ExecOutputStream::Pty => { + if let Err(error) = self.stdout.extend_from_slice(&bytes) { + self.close_for_oversized_line("stdout", error); + } + } + // Stderr is intentionally out-of-band. It should help debug server + // startup failures without entering rmcp framing. + ExecOutputStream::Stderr => { + if let Err(error) = self.push_stderr(&bytes) { + self.stdout.clear(); + self.close_for_oversized_line("stderr", error); + } + } + } + } + + fn close_for_oversized_line(&mut self, stream_name: &str, error: LineTooLong) { + let max_line_bytes = error.max_line_bytes; + warn!( + "Remote MCP server {stream_name} line exceeds {max_line_bytes} bytes ({}); closing transport", + self.program_name + ); + self.stderr.clear(); + // Returning EOF makes rmcp drop the transport, whose Drop implementation + // terminates the executor-managed process. + self.closed = true; + } + + fn take_stdout_message(&mut self, allow_partial: bool) -> Option> { + // A normal MCP stdio server emits one JSON-RPC message per newline. + // If the process has already closed, accept a final unterminated line + // so EOF after a complete JSON object behaves like local rmcp's + // `decode_eof` handling. + loop { + let line = match self.stdout.take_line() { + Some(line) => line, + None if allow_partial => self.stdout.take_remaining()?, + None => return None, + }; + let line = Self::trim_trailing_carriage_return(line); + match serde_json::from_slice(&line) { + Ok(message) => return Some(message), + Err(error) => { + debug!( + "Failed to parse remote MCP server message ({}): {error}", + self.program_name + ); + } + } + } + } + + fn push_stderr(&mut self, bytes: &[u8]) -> Result<(), LineTooLong> { + // Keep stderr line-oriented in logs so a chatty MCP server does not + // produce one log record per byte chunk. + self.stderr.extend_from_slice(bytes)?; + while let Some(line) = self.stderr.take_line() { + let line = Self::trim_trailing_carriage_return(line); + info!( + "MCP server stderr ({}): {}", + self.program_name, + String::from_utf8_lossy(&line) + ); + } + Ok(()) + } + + fn flush_stderr(&mut self) { + let Some(line) = self.stderr.take_remaining() else { + return; + }; + info!( + "MCP server stderr ({}): {}", + self.program_name, + String::from_utf8_lossy(&line) + ); + } + + fn trim_trailing_carriage_return(mut line: BytesMut) -> BytesMut { + if line.last() == Some(&b'\r') { + line.truncate(line.len() - 1); + } + line + } +} + +#[cfg(test)] +#[path = "executor_process_transport_tests.rs"] +mod tests; + +impl Drop for ExecutorProcessTransport { + fn drop(&mut self) { + if self.terminated { + return; + } + + let process = Arc::clone(&self.process); + let program_name = self.program_name.clone(); + let Ok(handle) = Handle::try_current() else { + warn!( + "Could not schedule remote MCP server process termination on drop ({}): no Tokio runtime is available", + self.program_name + ); + return; + }; + + std::mem::drop(handle.spawn(async move { + if let Err(error) = process.terminate().await { + warn!( + "Failed to terminate remote MCP server process on drop ({program_name}): {error}" + ); + } + })); + } +} diff --git a/codex-rs/rmcp-client/src/executor_process_transport_tests.rs b/codex-rs/rmcp-client/src/executor_process_transport_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..82fdaddcc93616a65c13ea644dea54f7d54c86d4 --- /dev/null +++ b/codex-rs/rmcp-client/src/executor_process_transport_tests.rs @@ -0,0 +1,501 @@ +use bytes::BytesMut; +use codex_exec_server::ExecOutputStream; +use codex_exec_server::ExecProcess; +use codex_exec_server::ExecProcessEventReceiver; +use codex_exec_server::ExecProcessFuture; +use codex_exec_server::ProcessId; +use codex_exec_server::ProcessOutputChunk; +use codex_exec_server::ProcessSignal; +use codex_exec_server::ReadResponse; +use codex_exec_server::WriteResponse; +use codex_exec_server::WriteStatus; +use pretty_assertions::assert_eq; +use rmcp::service::RoleClient; +use rmcp::service::TxJsonRpcMessage; +use rmcp::transport::Transport; +use serde_json::json; +use std::io; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +use std::task::Poll; +use tokio::sync::watch; + +use super::ExecutorProcessTransport; +use super::LineBuffer; +use super::LineTooLong; +use super::MAX_MCP_STDERR_LINE_BYTES; +use super::MAX_MCP_STDOUT_LINE_BYTES; + +struct BlockingFirstWriteProcess { + process_id: ProcessId, + writes: StdMutex>>, + release_first_write: AtomicBool, +} + +impl BlockingFirstWriteProcess { + fn writes(&self) -> Vec> { + self.writes.lock().expect("writes lock").clone() + } +} + +impl ExecProcess for BlockingFirstWriteProcess { + fn process_id(&self) -> &ProcessId { + &self.process_id + } + + fn subscribe_wake(&self) -> watch::Receiver { + watch::channel(0).1 + } + + fn subscribe_events(&self) -> ExecProcessEventReceiver { + ExecProcessEventReceiver::empty() + } + + fn read( + &self, + _after_seq: Option, + _max_bytes: Option, + _wait_ms: Option, + ) -> ExecProcessFuture<'_, ReadResponse> { + Box::pin(async { unreachable!("send test should not read process output") }) + } + + fn write(&self, chunk: Vec) -> ExecProcessFuture<'_, WriteResponse> { + let first_write = { + let mut writes = self.writes.lock().expect("writes lock"); + writes.push(chunk); + writes.len() == 1 + }; + Box::pin(std::future::poll_fn(move |_| { + if first_write && !self.release_first_write.load(Ordering::Acquire) { + return Poll::Pending; + } + Poll::Ready(Ok(WriteResponse { + status: WriteStatus::Accepted, + })) + })) + } + + fn signal(&self, _signal: ProcessSignal) -> ExecProcessFuture<'_, ()> { + Box::pin(async { Ok(()) }) + } + + fn terminate(&self) -> ExecProcessFuture<'_, ()> { + Box::pin(async { Ok(()) }) + } +} + +struct RetainedReadProcess { + process_id: ProcessId, + response: ReadResponse, +} + +impl ExecProcess for RetainedReadProcess { + fn process_id(&self) -> &ProcessId { + &self.process_id + } + + fn subscribe_wake(&self) -> watch::Receiver { + watch::channel(0).1 + } + + fn subscribe_events(&self) -> ExecProcessEventReceiver { + ExecProcessEventReceiver::empty() + } + + fn read( + &self, + _after_seq: Option, + _max_bytes: Option, + _wait_ms: Option, + ) -> ExecProcessFuture<'_, ReadResponse> { + Box::pin(async { Ok(self.response.clone()) }) + } + + fn write(&self, _chunk: Vec) -> ExecProcessFuture<'_, WriteResponse> { + Box::pin(async { unreachable!("recovery tests should not write process input") }) + } + + fn signal(&self, _signal: ProcessSignal) -> ExecProcessFuture<'_, ()> { + Box::pin(async { Ok(()) }) + } + + fn terminate(&self) -> ExecProcessFuture<'_, ()> { + Box::pin(async { Ok(()) }) + } +} + +fn recovered_output(seq: u64, stream: ExecOutputStream, bytes: &[u8]) -> ProcessOutputChunk { + ProcessOutputChunk { + seq, + stream, + chunk: bytes.to_vec().into(), + } +} + +fn retained_read_response(chunks: Vec, next_seq: u64) -> ReadResponse { + ReadResponse { + chunks, + next_seq, + exited: false, + exit_code: None, + closed: false, + failure: None, + sandbox_denied: false, + } +} + +#[tokio::test] +async fn serializes_concurrent_stdin_writes() { + let process = Arc::new(BlockingFirstWriteProcess { + process_id: ProcessId::from("mcp-stdio-test"), + writes: StdMutex::new(Vec::new()), + release_first_write: AtomicBool::new(false), + }); + let mut transport = + ExecutorProcessTransport::new(process.clone(), "mcp-stdio-test".to_string()); + let first_message: TxJsonRpcMessage = + serde_json::from_value(json!({ "jsonrpc": "2.0", "id": 1, "method": "ping" })) + .expect("first MCP message should deserialize"); + let second_message: TxJsonRpcMessage = + serde_json::from_value(json!({ "jsonrpc": "2.0", "id": 2, "method": "ping" })) + .expect("second MCP message should deserialize"); + + // Drive both sends explicitly so task scheduling cannot hide an overlapping write. + let first_send = transport.send(first_message); + tokio::pin!(first_send); + assert!(matches!(futures::poll!(first_send.as_mut()), Poll::Pending)); + assert_eq!( + process.writes(), + vec![b"{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}\n".to_vec()] + ); + + let second_send = transport.send(second_message); + tokio::pin!(second_send); + assert!(matches!( + futures::poll!(second_send.as_mut()), + Poll::Pending + )); + assert_eq!( + process.writes(), + vec![b"{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}\n".to_vec()] + ); + + process.release_first_write.store(true, Ordering::Release); + assert!(matches!( + futures::poll!(first_send.as_mut()), + Poll::Ready(Ok(())) + )); + assert!(matches!( + futures::poll!(second_send.as_mut()), + Poll::Ready(Ok(())) + )); + assert_eq!( + process.writes(), + vec![ + b"{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}\n".to_vec(), + b"{\"jsonrpc\":\"2.0\",\"id\":2,\"method\":\"ping\"}\n".to_vec(), + ] + ); +} + +#[tokio::test] +async fn rejects_lagged_recovery_when_retained_output_skips_a_sequence() { + let process = Arc::new(RetainedReadProcess { + process_id: ProcessId::from("mcp-stdio-lost-output"), + response: retained_read_response( + vec![recovered_output( + /*seq*/ 6, + ExecOutputStream::Stdout, + br#"{"jsonrpc":"2.0","id":1,"result":{}}"#, + )], + /*next_seq*/ 7, + ), + }); + let mut transport = ExecutorProcessTransport::new(process, "mcp-stdio-lost-output".to_string()); + transport.last_seq = 4; + transport + .stdout + .extend_from_slice(br#"{"jsonrpc":"2.0","id":1,"result":""#) + .expect("partial stdout should fit"); + transport + .stderr + .extend_from_slice(b"partial diagnostic") + .expect("partial stderr should fit"); + + let error = transport + .recover_lagged_events() + .await + .expect_err("evicted output must close the transport"); + + assert_eq!(error.kind(), io::ErrorKind::InvalidData); + assert_eq!( + error.to_string(), + "remote MCP server output stream lost process events: expected sequence 5, received 6" + ); + assert_eq!(transport.stdout, LineBuffer::default()); + assert_eq!(transport.stderr, LineBuffer::new(MAX_MCP_STDERR_LINE_BYTES)); + assert!(transport.closed); +} + +#[tokio::test] +async fn rejects_lagged_recovery_when_all_missing_output_was_evicted() { + let process = Arc::new(RetainedReadProcess { + process_id: ProcessId::from("mcp-stdio-evicted-output"), + response: retained_read_response(Vec::new(), /*next_seq*/ 7), + }); + let mut transport = + ExecutorProcessTransport::new(process, "mcp-stdio-evicted-output".to_string()); + transport.last_seq = 4; + + let error = transport + .recover_lagged_events() + .await + .expect_err("missing retained output must close the transport"); + + assert_eq!(error.kind(), io::ErrorKind::InvalidData); + assert_eq!( + error.to_string(), + "remote MCP server output stream lost process events: expected sequence 5, received 7" + ); + assert!(transport.closed); +} + +#[tokio::test] +async fn recovers_contiguous_output_after_a_duplicate_replayed_event() { + let process = Arc::new(RetainedReadProcess { + process_id: ProcessId::from("mcp-stdio-duplicate-output"), + response: retained_read_response( + vec![ + recovered_output( + /*seq*/ 4, + ExecOutputStream::Stdout, + b"duplicate output\n", + ), + recovered_output( + /*seq*/ 5, + ExecOutputStream::Stdout, + b"recovered stdout\n", + ), + recovered_output( + /*seq*/ 6, + ExecOutputStream::Stderr, + b"recovered stderr\n", + ), + ], + /*next_seq*/ 7, + ), + }); + let mut transport = + ExecutorProcessTransport::new(process, "mcp-stdio-duplicate-output".to_string()); + transport.last_seq = 4; + + transport + .recover_lagged_events() + .await + .expect("duplicate replay must not prevent recovery of contiguous output"); + + assert_eq!( + transport.stdout.take_line(), + Some(BytesMut::from(&b"recovered stdout"[..])) + ); + assert_eq!(transport.stdout.take_line(), None); + assert_eq!(transport.stderr, LineBuffer::new(MAX_MCP_STDERR_LINE_BYTES)); + assert_eq!(transport.last_seq, 6); + assert!(!transport.closed); +} + +#[tokio::test] +async fn recovers_contiguous_output_and_accounts_for_terminal_events() { + let mut response = retained_read_response( + vec![ + recovered_output( + /*seq*/ 5, + ExecOutputStream::Stdout, + b"recovered stdout\n", + ), + recovered_output( + /*seq*/ 6, + ExecOutputStream::Stderr, + b"recovered stderr\n", + ), + ], + /*next_seq*/ 9, + ); + response.exited = true; + response.exit_code = Some(0); + response.closed = true; + let process = Arc::new(RetainedReadProcess { + process_id: ProcessId::from("mcp-stdio-recovered-output"), + response, + }); + let mut transport = + ExecutorProcessTransport::new(process, "mcp-stdio-recovered-output".to_string()); + transport.last_seq = 4; + + transport + .recover_lagged_events() + .await + .expect("contiguous retained output should remain recoverable"); + + assert_eq!( + transport.stdout.take_line(), + Some(BytesMut::from(&b"recovered stdout"[..])) + ); + assert_eq!(transport.stderr, LineBuffer::new(MAX_MCP_STDERR_LINE_BYTES)); + assert_eq!(transport.last_seq, 8); + assert!(transport.closed); +} + +#[tokio::test] +async fn closes_on_an_unsequenced_process_failure_without_inventing_missing_output() { + let mut response = retained_read_response(Vec::new(), /*next_seq*/ 5); + response.exited = true; + response.closed = true; + response.failure = Some("executor disconnected".to_string()); + let process = Arc::new(RetainedReadProcess { + process_id: ProcessId::from("mcp-stdio-failed-process"), + response, + }); + let mut transport = + ExecutorProcessTransport::new(process, "mcp-stdio-failed-process".to_string()); + transport.last_seq = 4; + + transport + .recover_lagged_events() + .await + .expect("an unsequenced executor failure is not an output sequence gap"); + + assert_eq!(transport.last_seq, 4); + assert!(transport.closed); +} + +#[test] +fn searches_only_new_bytes_after_partial_line() { + let mut buffer = LineBuffer::default(); + + buffer + .extend_from_slice(b"partial") + .expect("partial line should fit"); + assert_eq!(buffer.take_line(), None); + assert_eq!( + buffer, + LineBuffer { + bytes: BytesMut::from(&b"partial"[..]), + scanned_len: 7, + pending_line_bytes: 7, + max_line_bytes: MAX_MCP_STDOUT_LINE_BYTES, + } + ); + + buffer + .extend_from_slice(b" line") + .expect("partial line should fit"); + assert_eq!(buffer.take_line(), None); + assert_eq!( + buffer, + LineBuffer { + bytes: BytesMut::from(&b"partial line"[..]), + scanned_len: 12, + pending_line_bytes: 12, + max_line_bytes: MAX_MCP_STDOUT_LINE_BYTES, + } + ); + + buffer + .extend_from_slice(b"\nnext") + .expect("completed line should fit"); + assert_eq!( + buffer.take_line(), + Some(BytesMut::from(&b"partial line"[..])) + ); + assert_eq!( + buffer, + LineBuffer { + bytes: BytesMut::from(&b"next"[..]), + scanned_len: 0, + pending_line_bytes: 4, + max_line_bytes: MAX_MCP_STDOUT_LINE_BYTES, + } + ); +} + +#[test] +fn splits_multiple_lines_and_retains_partial_tail() { + let mut buffer = LineBuffer::default(); + buffer + .extend_from_slice(b"first\nsecond\npartial") + .expect("lines should fit"); + + assert_eq!(buffer.take_line(), Some(BytesMut::from(&b"first"[..]))); + assert_eq!(buffer.take_line(), Some(BytesMut::from(&b"second"[..]))); + assert_eq!(buffer.take_line(), None); + assert_eq!( + buffer, + LineBuffer { + bytes: BytesMut::from(&b"partial"[..]), + scanned_len: 7, + pending_line_bytes: 7, + max_line_bytes: MAX_MCP_STDOUT_LINE_BYTES, + } + ); +} + +#[test] +fn takes_unterminated_remaining_bytes_at_eof() { + let mut buffer = LineBuffer::default(); + buffer + .extend_from_slice(b"remaining") + .expect("remaining line should fit"); + assert_eq!(buffer.take_line(), None); + + assert_eq!( + buffer.take_remaining(), + Some(BytesMut::from(&b"remaining"[..])) + ); + assert_eq!(buffer, LineBuffer::default()); +} + +#[test] +fn rejects_oversized_line_without_retaining_its_prefix() { + let mut buffer = LineBuffer::new(/*max_line_bytes*/ 5); + buffer + .extend_from_slice(b"12345") + .expect("line at the limit should fit"); + assert_eq!(buffer.take_line(), None); + + assert_eq!( + buffer.extend_from_slice(b"6"), + Err(LineTooLong { max_line_bytes: 5 }) + ); + assert_eq!(buffer, LineBuffer::new(/*max_line_bytes*/ 5)); +} + +#[test] +fn retains_complete_lines_before_an_oversized_line() { + let mut buffer = LineBuffer::new(/*max_line_bytes*/ 5); + + assert_eq!( + buffer.extend_from_slice(b"first\n123456"), + Err(LineTooLong { max_line_bytes: 5 }) + ); + + assert_eq!(buffer.take_line(), Some(BytesMut::from(&b"first"[..]))); + assert_eq!(buffer.take_remaining(), None); +} + +#[test] +fn accepts_input_larger_than_limit_when_each_line_is_bounded() { + let mut buffer = LineBuffer::new(/*max_line_bytes*/ 5); + + buffer + .extend_from_slice(b"12345\nabcde\ntail") + .expect("each individual line should fit"); + + assert_eq!(buffer.take_line(), Some(BytesMut::from(&b"12345"[..]))); + assert_eq!(buffer.take_line(), Some(BytesMut::from(&b"abcde"[..]))); + assert_eq!(buffer.take_line(), None); + assert_eq!(buffer.take_remaining(), Some(BytesMut::from(&b"tail"[..]))); +} diff --git a/codex-rs/rmcp-client/src/http_client_adapter.rs b/codex-rs/rmcp-client/src/http_client_adapter.rs new file mode 100644 index 0000000000000000000000000000000000000000..fee686286b6421765564d18e7d1006a54768355a --- /dev/null +++ b/codex-rs/rmcp-client/src/http_client_adapter.rs @@ -0,0 +1,1069 @@ +//! RMCP Streamable HTTP adapter built on top of the shared `HttpClient` +//! capability. +//! +//! This module runs in the orchestrator process. It turns high-level RMCP +//! operations like `post_message` and `get_stream` into calls on +//! `Arc`, which may be: +//! - a local HTTP client that issues requests from the orchestrator, or +//! - a remote HTTP client that forwards requests to the remote runtime + +use std::collections::HashMap; +use std::io; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; +use std::time::Duration; +use std::time::Instant; + +use bytes::Bytes; +use codex_api::SharedAuthProvider; +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpHeader; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpResponseBodyStream; +use futures::StreamExt; +use futures::stream; +use futures::stream::BoxStream; +use http::HeaderMap; +use http::HeaderName; +use http::HeaderValue; +use http::StatusCode; +use http::header::ACCEPT; +use http::header::AUTHORIZATION; +use http::header::CONTENT_TYPE; +use http::header::WWW_AUTHENTICATE; +use rmcp::model::ClientJsonRpcMessage; +use rmcp::model::ClientNotification; +use rmcp::model::ConstString; +use rmcp::model::DiscoverRequestMethod; +use rmcp::model::ErrorCode; +use rmcp::model::ErrorData; +use rmcp::model::JsonRpcMessage; +use rmcp::model::ProtocolVersion; +use rmcp::model::RequestId; +use rmcp::model::ServerJsonRpcMessage; +use rmcp::model::ServerResult; +use rmcp::transport::common::http_header::HEADER_MCP_PROTOCOL_VERSION; +use rmcp::transport::streamable_http_client::AuthRequiredError; +use rmcp::transport::streamable_http_client::InsufficientScopeError; +use rmcp::transport::streamable_http_client::StreamableHttpClient; +use rmcp::transport::streamable_http_client::StreamableHttpError; +use rmcp::transport::streamable_http_client::StreamableHttpPostResponse; +use sse_stream::Sse; +use sse_stream::SseStream; +use tokio::sync::oneshot; + +use crate::bounded_stdio_transport::MAX_MCP_STDIO_LINE_BYTES; +use crate::event_notification_transport::MAX_EVENT_NOTIFICATION_BYTES; +use crate::http_client_redirect::SameOriginRedirectHttpClient; + +use crate::www_authenticate::insufficient_scope_challenge; + +const EVENT_STREAM_MIME_TYPE: &str = "text/event-stream"; +const JSON_MIME_TYPE: &str = "application/json"; +const HEADER_SESSION_ID: &str = "Mcp-Session-Id"; +const NON_JSON_RESPONSE_BODY_PREVIEW_BYTES: usize = 8_192; +const LEGACY_HTTP_PREVALIDATION_ERROR_CODE: ErrorCode = ErrorCode(-32000); +const EVENT_STREAM_RESPONSE_TIMEOUT: Duration = Duration::from_secs(30); + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum StreamableHttpRedirectMode { + Legacy, + AgentPluginV1, +} + +#[derive(Clone)] +pub(crate) struct StreamableHttpClientAdapter { + http_client: Arc, + default_headers: HeaderMap, + auth_provider: Option, + event_stream_cancellations: Arc>>>, + has_configured_headers: bool, + redirect_mode: StreamableHttpRedirectMode, + initialize_deadline: Arc>>, +} + +struct EventStreamCancellation { + request_id: RequestId, + cancellations: Arc>>>, +} + +impl Drop for EventStreamCancellation { + fn drop(&mut self) { + self.cancellations + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove(&self.request_id); + } +} + +#[derive(Debug, thiserror::Error)] +pub(crate) enum StreamableHttpClientAdapterError { + #[error("streamable HTTP session expired with 404 Not Found")] + SessionExpired404, + #[error(transparent)] + HttpRequest(#[from] ExecServerError), + #[error("invalid HTTP header: {0}")] + Header(String), + #[error("MCP response body exceeds {maximum_bytes} bytes")] + ResponseTooLarge { maximum_bytes: usize }, +} + +impl StreamableHttpClientAdapter { + pub(crate) fn new( + http_client: Arc, + default_headers: HeaderMap, + auth_provider: Option, + has_configured_headers: bool, + redirect_mode: StreamableHttpRedirectMode, + initialize_deadline: Arc>>, + ) -> Self { + Self { + http_client: Arc::new(SameOriginRedirectHttpClient::new(http_client)), + default_headers, + auth_provider, + event_stream_cancellations: Arc::default(), + has_configured_headers, + redirect_mode, + initialize_deadline, + } + } + + fn redirect_policy(&self, headers: &HeaderMap) -> HttpRedirectPolicy { + mcp_redirect_policy(self.redirect_mode, headers, self.has_configured_headers) + } +} + +impl StreamableHttpClient for StreamableHttpClientAdapter { + type Error = StreamableHttpClientAdapterError; + + async fn post_message( + &self, + uri: Arc, + message: ClientJsonRpcMessage, + session_id: Option>, + auth_token: Option, + custom_headers: HashMap, + ) -> std::result::Result> { + let (mcp_method, mcp_request_id) = client_jsonrpc_message_fields(&message); + let has_session_id = session_id.is_some(); + let mut headers = self.default_headers.clone(); + headers.extend(custom_headers); + self.add_auth_headers(&mut headers); + insert_header( + &mut headers, + ACCEPT, + [EVENT_STREAM_MIME_TYPE, JSON_MIME_TYPE].join(", "), + StreamableHttpClientAdapterError::Header, + )?; + insert_header( + &mut headers, + CONTENT_TYPE, + JSON_MIME_TYPE.to_string(), + StreamableHttpClientAdapterError::Header, + )?; + if let Some(auth_token) = auth_token { + insert_header( + &mut headers, + AUTHORIZATION, + format!("Bearer {auth_token}"), + StreamableHttpClientAdapterError::Header, + )?; + } + if let Some(session_id_value) = session_id.as_ref() { + insert_header( + &mut headers, + HeaderName::from_static("mcp-session-id"), + session_id_value.to_string(), + StreamableHttpClientAdapterError::Header, + )?; + } + + let is_discovery_request = mcp_method.as_deref() == Some(DiscoverRequestMethod::VALUE); + let is_event_stream_request = mcp_method.as_deref() == Some("events/stream"); + let uses_modern_protocol = headers + .get(HEADER_MCP_PROTOCOL_VERSION) + .and_then(|value| value.to_str().ok()) + == Some(ProtocolVersion::V_2026_07_28.as_str()); + let maximum_response_bytes = if is_event_stream_request { + Some(MAX_EVENT_NOTIFICATION_BYTES) + } else { + (is_discovery_request || uses_modern_protocol).then_some(MAX_MCP_STDIO_LINE_BYTES) + }; + let redirect_policy = if is_discovery_request { + HttpRedirectPolicy::Stop + } else { + self.redirect_policy(&headers) + }; + let timeout_ms = if matches!( + mcp_method.as_deref(), + Some("initialize" | "notifications/initialized") + ) || mcp_method.as_deref() == Some(DiscoverRequestMethod::VALUE) + { + self.initialize_deadline + .lock() + .unwrap_or_else(PoisonError::into_inner) + .map(|deadline| { + u64::try_from( + deadline + .saturating_duration_since(Instant::now()) + .as_millis(), + ) + .unwrap_or(u64::MAX) + .max(1) + }) + } else { + None + }; + + let body = serde_json::to_vec(&message).map_err(StreamableHttpError::Deserialize)?; + let has_authorization_header = headers.contains_key(AUTHORIZATION); + if let JsonRpcMessage::Notification(notification) = &message + && let ClientNotification::CancelledNotification(cancelled) = ¬ification.notification + && let Some(request_id) = cancelled.params.request_id.as_ref() + && let Some(cancellation) = self + .event_stream_cancellations + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove(request_id) + { + let _ = cancellation.send(()); + return Ok(StreamableHttpPostResponse::Accepted); + } + + let request = self.http_client.http_request_stream(HttpRequestParams { + method: "POST".to_string(), + url: uri.to_string(), + headers: protocol_headers(&headers), + body: Some(body.into()), + timeout_ms, + redirect_policy, + request_id: "buffered-request".to_string(), + stream_response: true, + }); + let response = if is_event_stream_request { + tokio::time::timeout(EVENT_STREAM_RESPONSE_TIMEOUT, request) + .await + .map_err(|_| { + StreamableHttpError::UnexpectedServerResponse( + "timed out waiting for MCP event stream response headers".into(), + ) + })? + } else { + request.await + }; + let (response, mut body_stream) = match response { + Ok(response) => response, + Err(error) => { + log_post_message_http_error( + &uri, + mcp_method.as_deref(), + mcp_request_id.as_deref(), + has_session_id, + has_authorization_header, + ); + return Err(StreamableHttpError::Client( + StreamableHttpClientAdapterError::from(error), + )); + } + }; + + if response.status == StatusCode::NOT_FOUND.as_u16() && session_id.is_some() { + return Err(StreamableHttpError::Client( + StreamableHttpClientAdapterError::SessionExpired404, + )); + } + if response.status == StatusCode::UNAUTHORIZED.as_u16() { + let challenges = response + .headers + .iter() + .filter(|header| header.name.eq_ignore_ascii_case(WWW_AUTHENTICATE.as_str())) + .map(|header| header.value.as_str()) + .collect::>(); + if !challenges.is_empty() { + // RFC 9110 allows combining these list-based fields; keep challenges after the first. + return Err(StreamableHttpError::AuthRequired(AuthRequiredError::new( + challenges.join(", "), + ))); + } + } + if response.status == StatusCode::FORBIDDEN.as_u16() + && let Some(challenge) = insufficient_scope_challenge(&response.headers) + { + return Err(StreamableHttpError::InsufficientScope( + InsufficientScopeError::new( + challenge.www_authenticate_header, + challenge.required_scope, + ), + )); + } + if matches!( + StatusCode::from_u16(response.status).ok(), + Some(StatusCode::ACCEPTED | StatusCode::NO_CONTENT) + ) { + return Ok(StreamableHttpPostResponse::Accepted); + } + + let content_type = response_header(&response.headers, CONTENT_TYPE); + let session_id = response_header(&response.headers, HEADER_SESSION_ID); + if !status_is_success(response.status) { + let body = collect_body(&mut body_stream, maximum_response_bytes).await?; + if !retryable_post_response_status(mcp_method.as_deref(), response.status) + && (content_type + .as_deref() + .is_some_and(|content_type| content_type.starts_with(JSON_MIME_TYPE)) + || (mcp_method.as_deref() == Some(DiscoverRequestMethod::VALUE) + && response.status == StatusCode::BAD_REQUEST.as_u16() + && !has_session_id)) + && let Some(response_message) = parse_json_rpc_error(&body) + { + return Ok(StreamableHttpPostResponse::Json( + legacy_discovery_fallback_response( + &message, + response_message, + response.status == StatusCode::BAD_REQUEST.as_u16() && !has_session_id, + ), + session_id, + )); + } + if mcp_method.as_deref() == Some(DiscoverRequestMethod::VALUE) + && !has_session_id + && matches!( + StatusCode::from_u16(response.status).ok(), + Some(StatusCode::NOT_FOUND | StatusCode::METHOD_NOT_ALLOWED) + ) + && content_type + .as_deref() + .is_none_or(|content_type| !content_type.starts_with(JSON_MIME_TYPE)) + && let JsonRpcMessage::Request(request) = &message + { + let legacy_error = ServerJsonRpcMessage::error( + ErrorData::new( + ErrorCode::METHOD_NOT_FOUND, + "legacy MCP endpoint does not support server/discover", + None, + ), + Some(request.id.clone()), + ); + return Ok(StreamableHttpPostResponse::Json(legacy_error, session_id)); + } + return Err(StreamableHttpError::UnexpectedServerResponse( + format!( + "HTTP {}: {}", + response.status, + body_preview(String::from_utf8_lossy(&body).to_string()) + ) + .into(), + )); + } + match content_type.as_deref() { + Some(content_type) if content_type.starts_with(EVENT_STREAM_MIME_TYPE) => { + let mut event_stream = sse_stream_from_body(body_stream, maximum_response_bytes); + if mcp_method.as_deref() == Some(DiscoverRequestMethod::VALUE) { + while let Some(event) = event_stream.next().await { + let event = event.map_err(StreamableHttpError::Sse)?; + if !matches!(event.event.as_deref(), None | Some("") | Some("message")) { + continue; + } + let Some(data) = event.data.as_deref() else { + continue; + }; + if data.trim().is_empty() { + continue; + } + + let response = serde_json::from_slice(data.as_bytes()) + .map_err(StreamableHttpError::Deserialize)?; + let response = legacy_discovery_fallback_response( + &message, response, /*allow_uncorrelated_http_rejection*/ false, + ); + if matches!( + &response, + JsonRpcMessage::Response(_) | JsonRpcMessage::Error(_) + ) { + return Ok(StreamableHttpPostResponse::Json(response, session_id)); + } + } + + return Err(StreamableHttpError::UnexpectedServerResponse( + "empty sse stream".into(), + )); + } + if is_event_stream_request && let JsonRpcMessage::Request(request) = &message { + let (cancel, cancelled) = oneshot::channel(); + let cancellation = EventStreamCancellation { + request_id: request.id.clone(), + cancellations: Arc::clone(&self.event_stream_cancellations), + }; + cancellation + .cancellations + .lock() + .unwrap_or_else(PoisonError::into_inner) + .insert(cancellation.request_id.clone(), cancel); + + event_stream = stream::unfold( + Some((event_stream, cancelled, cancellation)), + |state| async move { + let (mut event_stream, mut cancelled, cancellation) = state?; + + tokio::select! { + biased; + + _ = &mut cancelled => None, + event = event_stream.next() => event.map(|event| { + (event, Some((event_stream, cancelled, cancellation))) + }), + } + }, + ) + .boxed(); + } + Ok(StreamableHttpPostResponse::Sse(event_stream, session_id)) + } + Some(content_type) if content_type.starts_with(JSON_MIME_TYPE) => { + let body = collect_body(&mut body_stream, maximum_response_bytes).await?; + let response_message = + serde_json::from_slice(&body).map_err(StreamableHttpError::Deserialize)?; + Ok(StreamableHttpPostResponse::Json( + legacy_discovery_fallback_response( + &message, + response_message, + /*allow_uncorrelated_http_rejection*/ false, + ), + session_id, + )) + } + _ => { + let body = collect_body(&mut body_stream, maximum_response_bytes).await?; + let content_type = content_type.unwrap_or_else(|| "missing-content-type".into()); + Err(StreamableHttpError::UnexpectedContentType(Some(format!( + "{content_type}; body: {}", + body_preview(String::from_utf8_lossy(&body).to_string()) + )))) + } + } + } + + async fn delete_session( + &self, + uri: Arc, + session: Arc, + auth_token: Option, + custom_headers: HashMap, + ) -> std::result::Result<(), StreamableHttpError> { + let mut headers = self.default_headers.clone(); + headers.extend(custom_headers); + self.add_auth_headers(&mut headers); + if let Some(auth_token) = auth_token { + insert_header( + &mut headers, + AUTHORIZATION, + format!("Bearer {auth_token}"), + StreamableHttpClientAdapterError::Header, + )?; + } + insert_header( + &mut headers, + HeaderName::from_static("mcp-session-id"), + session.to_string(), + StreamableHttpClientAdapterError::Header, + )?; + let redirect_policy = self.redirect_policy(&headers); + + let response = self + .http_client + .http_request(HttpRequestParams { + method: "DELETE".to_string(), + url: uri.to_string(), + headers: protocol_headers(&headers), + body: None, + timeout_ms: None, + redirect_policy, + request_id: "buffered-request".to_string(), + stream_response: false, + }) + .await + .map_err(StreamableHttpClientAdapterError::from) + .map_err(StreamableHttpError::Client)?; + + if response.status == StatusCode::METHOD_NOT_ALLOWED.as_u16() { + return Ok(()); + } + if !status_is_success(response.status) { + return Err(StreamableHttpError::UnexpectedServerResponse( + format!("DELETE returned HTTP {}", response.status).into(), + )); + } + Ok(()) + } + + async fn get_stream( + &self, + uri: Arc, + session_id: Option>, + last_event_id: Option, + auth_token: Option, + custom_headers: HashMap, + ) -> std::result::Result< + BoxStream<'static, std::result::Result>, + StreamableHttpError, + > { + let mut headers = self.default_headers.clone(); + headers.extend(custom_headers); + self.add_auth_headers(&mut headers); + insert_header( + &mut headers, + ACCEPT, + [EVENT_STREAM_MIME_TYPE, JSON_MIME_TYPE].join(", "), + StreamableHttpClientAdapterError::Header, + )?; + if let Some(session_id) = session_id { + insert_header( + &mut headers, + HeaderName::from_static("mcp-session-id"), + session_id.to_string(), + StreamableHttpClientAdapterError::Header, + )?; + } + if let Some(last_event_id) = last_event_id { + insert_header( + &mut headers, + HeaderName::from_static("last-event-id"), + last_event_id, + StreamableHttpClientAdapterError::Header, + )?; + } + if let Some(auth_token) = auth_token { + insert_header( + &mut headers, + AUTHORIZATION, + format!("Bearer {auth_token}"), + StreamableHttpClientAdapterError::Header, + )?; + } + let redirect_policy = self.redirect_policy(&headers); + + let (response, body_stream) = self + .http_client + .http_request_stream(HttpRequestParams { + method: "GET".to_string(), + url: uri.to_string(), + headers: protocol_headers(&headers), + body: None, + timeout_ms: None, + redirect_policy, + request_id: "buffered-request".to_string(), + stream_response: true, + }) + .await + .map_err(StreamableHttpClientAdapterError::from) + .map_err(StreamableHttpError::Client)?; + + if response.status == StatusCode::METHOD_NOT_ALLOWED.as_u16() { + return Err(StreamableHttpError::ServerDoesNotSupportSse); + } + if response.status == StatusCode::NOT_FOUND.as_u16() { + return Err(StreamableHttpError::Client( + StreamableHttpClientAdapterError::SessionExpired404, + )); + } + if !status_is_success(response.status) { + return Err(StreamableHttpError::UnexpectedServerResponse( + format!("GET returned HTTP {}", response.status).into(), + )); + } + + match response_header(&response.headers, CONTENT_TYPE).as_deref() { + Some(content_type) if is_streamable_http_content_type(content_type) => {} + Some(content_type) => { + return Err(StreamableHttpError::UnexpectedContentType(Some( + content_type.to_string(), + ))); + } + None => { + return Err(StreamableHttpError::UnexpectedContentType(None)); + } + } + + let uses_modern_protocol = headers + .get(HEADER_MCP_PROTOCOL_VERSION) + .and_then(|value| value.to_str().ok()) + .is_some_and(|version| version == ProtocolVersion::V_2026_07_28.as_str()); + let maximum_response_bytes = uses_modern_protocol.then_some(MAX_MCP_STDIO_LINE_BYTES); + Ok(sse_stream_from_body(body_stream, maximum_response_bytes)) + } +} + +impl StreamableHttpClientAdapter { + fn add_auth_headers(&self, headers: &mut HeaderMap) { + if let Some(auth_provider) = &self.auth_provider { + headers.extend(auth_provider.to_auth_headers()); + } + } +} + +fn body_preview(body: impl Into) -> String { + let mut body_preview = body.into(); + let body_len = body_preview.len(); + if body_len > NON_JSON_RESPONSE_BODY_PREVIEW_BYTES { + let mut boundary = NON_JSON_RESPONSE_BODY_PREVIEW_BYTES; + while !body_preview.is_char_boundary(boundary) { + boundary = boundary.saturating_sub(1); + } + body_preview.truncate(boundary); + body_preview.push_str(&format!( + "... (truncated {} bytes)", + body_len.saturating_sub(boundary) + )); + } + body_preview +} + +fn client_jsonrpc_message_fields( + message: &ClientJsonRpcMessage, +) -> (Option, Option) { + match message { + JsonRpcMessage::Request(request) => ( + Some(request.request.method().to_string()), + Some(request.id.to_string()), + ), + JsonRpcMessage::Response(response) => (None, Some(response.id.to_string())), + JsonRpcMessage::Notification(notification) => { + let method = match ¬ification.notification { + ClientNotification::CancelledNotification(notification) => { + notification.method.as_str() + } + ClientNotification::ProgressNotification(notification) => { + notification.method.as_str() + } + ClientNotification::InitializedNotification(notification) => { + notification.method.as_str() + } + ClientNotification::RootsListChangedNotification(notification) => { + notification.method.as_str() + } + ClientNotification::CustomNotification(notification) => { + notification.method.as_str() + } + _ => return (None, None), + }; + (Some(method.to_string()), None) + } + JsonRpcMessage::Error(error) => (None, error.id.as_ref().map(ToString::to_string)), + } +} + +fn log_post_message_http_error( + uri: &str, + mcp_method: Option<&str>, + mcp_request_id: Option<&str>, + has_session_id: bool, + has_authorization_header: bool, +) { + let parsed_url = url::Url::parse(uri).ok(); + tracing::warn!( + endpoint_scheme = parsed_url + .as_ref() + .map(url::Url::scheme) + .unwrap_or(""), + endpoint_host = parsed_url + .as_ref() + .and_then(url::Url::host_str) + .unwrap_or(""), + endpoint_path = parsed_url + .as_ref() + .map(url::Url::path) + .unwrap_or(""), + endpoint_has_query = parsed_url.as_ref().is_some_and(|url| url.query().is_some()), + mcp_method = mcp_method.unwrap_or(""), + mcp_request_id = mcp_request_id.unwrap_or(""), + has_session_id = has_session_id, + has_authorization_header = has_authorization_header, + "streamable HTTP post_message failed" + ); +} + +fn insert_header( + headers: &mut HeaderMap, + name: HeaderName, + value: String, + map_error: impl FnOnce(String) -> Error, +) -> std::result::Result<(), StreamableHttpError> +where + Error: std::error::Error + Send + Sync + 'static, +{ + let value = HeaderValue::from_str(&value) + .map_err(|error| StreamableHttpError::Client(map_error(error.to_string())))?; + headers.insert(name, value); + Ok(()) +} + +fn is_streamable_http_content_type(content_type: &str) -> bool { + content_type + .as_bytes() + .starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) + || content_type + .as_bytes() + .starts_with(JSON_MIME_TYPE.as_bytes()) +} + +fn protocol_headers(headers: &HeaderMap) -> Vec { + headers + .iter() + .filter_map(|(name, value)| { + Some(HttpHeader { + name: name.as_str().to_string(), + value: std::str::from_utf8(value.as_bytes()).ok()?.to_string(), + value_env_var: None, + }) + }) + .collect() +} + +fn response_header(headers: &[HttpHeader], name: impl AsRef) -> Option { + let name = name.as_ref(); + headers + .iter() + .find(|header| header.name.eq_ignore_ascii_case(name)) + .map(|header| header.value.clone()) +} + +fn status_is_success(status: u16) -> bool { + StatusCode::from_u16(status).is_ok_and(|status| status.is_success()) +} + +fn retryable_post_response_status(mcp_method: Option<&str>, status: u16) -> bool { + let Ok(status) = StatusCode::from_u16(status) else { + return false; + }; + is_retryable_http_status(status) + && matches!( + mcp_method, + Some( + DiscoverRequestMethod::VALUE + | "initialize" + | "notifications/initialized" + | "tools/list" + ) + ) +} + +fn is_retryable_http_status(status: StatusCode) -> bool { + matches!( + status, + StatusCode::REQUEST_TIMEOUT + | StatusCode::TOO_MANY_REQUESTS + | StatusCode::INTERNAL_SERVER_ERROR + | StatusCode::BAD_GATEWAY + | StatusCode::SERVICE_UNAVAILABLE + | StatusCode::GATEWAY_TIMEOUT + ) +} + +fn parse_json_rpc_error(body: &[u8]) -> Option { + match serde_json::from_slice::(body) { + Ok(message @ JsonRpcMessage::Error(_)) => Some(message), + _ => None, + } +} + +fn mcp_redirect_policy( + mode: StreamableHttpRedirectMode, + headers: &HeaderMap, + has_configured_headers: bool, +) -> HttpRedirectPolicy { + if headers + .get(HEADER_MCP_PROTOCOL_VERSION) + .and_then(|value| value.to_str().ok()) + == Some(ProtocolVersion::V_2026_07_28.as_str()) + || (mode == StreamableHttpRedirectMode::AgentPluginV1 + && (has_configured_headers || headers.contains_key(AUTHORIZATION))) + { + HttpRedirectPolicy::Stop + } else { + HttpRedirectPolicy::Follow + } +} + +// rmcp's automatic lifecycle does not yet recognize deployed legacy discovery +// rejection shapes. Remove this compatibility shim once the SDK does: +// https://github.com/modelcontextprotocol/rust-sdk/issues/1040 +fn legacy_discovery_fallback_response( + request: &ClientJsonRpcMessage, + response: ServerJsonRpcMessage, + allow_uncorrelated_http_rejection: bool, +) -> ServerJsonRpcMessage { + let JsonRpcMessage::Request(request) = request else { + return response; + }; + if request.request.method() != DiscoverRequestMethod::VALUE { + return response; + } + + if let JsonRpcMessage::Error(error) = &response + && error.error.code == ErrorCode::METHOD_NOT_FOUND + && error.id.as_ref() != Some(&request.id) + { + return ServerJsonRpcMessage::error( + ErrorData::new( + ErrorCode::HEADER_MISMATCH, + "server/discover method-not-found response did not match its request ID", + None, + ), + Some(request.id.clone()), + ); + } + + let requires_legacy_initialization = match &response { + JsonRpcMessage::Response(response) if response.id == request.id => match &response.result { + ServerResult::DiscoverResult(result) => { + only_known_legacy_protocol_versions(&result.supported_versions) + } + _ => false, + }, + JsonRpcMessage::Error(error) if error.id.as_ref() == Some(&request.id) => { + (error.error.code == ErrorCode::UNSUPPORTED_PROTOCOL_VERSION + && error + .error + .data + .as_ref() + .and_then(|data| data.get("supported")) + .and_then(|supported| { + serde_json::from_value::>(supported.clone()).ok() + }) + .is_some_and(|supported| only_known_legacy_protocol_versions(&supported))) + || matches!( + error.error.code, + ErrorCode::UNSUPPORTED_PROTOCOL_VERSION + | ErrorCode::INVALID_REQUEST + | ErrorCode::INVALID_PARAMS + ) && explicitly_rejects_modern_protocol_version(&error.error.message) + } + JsonRpcMessage::Error(error) + if allow_uncorrelated_http_rejection + && error.id.is_none() + && error.error.code == LEGACY_HTTP_PREVALIDATION_ERROR_CODE => + { + has_legacy_fallback_evidence(&error.error.message) + } + _ => false, + }; + + if requires_legacy_initialization { + ServerJsonRpcMessage::error( + ErrorData::new( + ErrorCode::METHOD_NOT_FOUND, + "MCP discovery requires legacy initialization", + None, + ), + Some(request.id.clone()), + ) + } else { + let mut response = response; + if let JsonRpcMessage::Error(error) = &mut response + && !matches!( + error.error.code, + ErrorCode::METHOD_NOT_FOUND + | ErrorCode::UNSUPPORTED_PROTOCOL_VERSION + | ErrorCode::HEADER_MISMATCH + | ErrorCode::MISSING_REQUIRED_CLIENT_CAPABILITY + ) + { + // rmcp 3.1.3 falls back on other discovery errors, so mark unproven + // rejections as modern failures while preserving their diagnostics. + error.error.code = ErrorCode::HEADER_MISMATCH; + } + response + } +} + +fn only_known_legacy_protocol_versions(versions: &[ProtocolVersion]) -> bool { + !versions.is_empty() + && versions.iter().all(|version| { + ProtocolVersion::KNOWN_VERSIONS.contains(version) + && version < &ProtocolVersion::V_2026_07_28 + }) +} + +fn explicitly_rejects_modern_protocol_version(message: &str) -> bool { + message + .trim() + .eq_ignore_ascii_case("unsupported protocol version: 2026-07-28") +} + +// Some legacy servers reject `server/discover` before assigning a JSON-RPC ID. +// A null-ID HTTP 400/-32000 does not, by itself, justify a downgrade. +// Retry `initialize` only for the exact missing-session error or a list of +// exclusively legacy versions that includes a version rmcp supports. +// These are compatibility hints, not proof of server identity; `initialize` +// negotiates the actual version, and 2025-06-18 is only our initial proposal. +fn has_legacy_fallback_evidence(message: &str) -> bool { + if message == "Bad Request: No valid session ID provided" { + return true; + } + + let Some(supported) = message + .strip_prefix("Bad Request: Unsupported protocol version: 2026-07-28 (supported versions: ") + .or_else(|| { + message.strip_prefix("Bad Request: Unsupported protocol version (supported versions: ") + }) + .and_then(|supported| supported.strip_suffix(')')) + else { + return false; + }; + + let versions = supported.split(',').map(str::trim).collect::>(); + !versions.is_empty() + && ProtocolVersion::KNOWN_VERSIONS + .iter() + .any(|known| versions.contains(&known.as_str())) + && versions.iter().all(|version| { + let bytes = version.as_bytes(); + bytes.len() == 10 + && bytes[4] == b'-' + && bytes[7] == b'-' + && bytes + .iter() + .enumerate() + .all(|(index, byte)| matches!(index, 4 | 7) || byte.is_ascii_digit()) + && *version < "2026-07-28" + }) +} + +async fn collect_body( + body_stream: &mut HttpResponseBodyStream, + maximum_bytes: Option, +) -> std::result::Result, StreamableHttpError> { + let mut body = Vec::new(); + while let Some(chunk) = body_stream + .recv() + .await + .map_err(StreamableHttpClientAdapterError::from) + .map_err(StreamableHttpError::Client)? + { + if let Some(maximum_bytes) = maximum_bytes + && chunk.len() > maximum_bytes.saturating_sub(body.len()) + { + return Err(StreamableHttpError::Client( + StreamableHttpClientAdapterError::ResponseTooLarge { maximum_bytes }, + )); + } + body.extend_from_slice(&chunk); + } + Ok(body) +} + +fn sse_stream_from_body( + body_stream: HttpResponseBodyStream, + maximum_event_bytes: Option, +) -> BoxStream<'static, std::result::Result> { + SseStream::from_bytes_stream(stream::unfold( + (body_stream, SseEventSizeLimit::new(maximum_event_bytes)), + |(mut body_stream, mut size_limit)| async move { + match body_stream.recv().await { + Ok(Some(bytes)) => { + if let Err(error) = size_limit.observe(&bytes) { + Some((Err(error), (body_stream, size_limit))) + } else { + Some((Ok(Bytes::from(bytes)), (body_stream, size_limit))) + } + } + Ok(None) => None, + Err(error) => Some((Err(io::Error::other(error)), (body_stream, size_limit))), + } + }, + )) + .boxed() +} + +struct SseEventSizeLimit { + maximum_bytes: Option, + retained_bytes: usize, + line_bytes: usize, + line_is_comment: bool, + previous_was_carriage_return: bool, + failed: bool, +} + +impl SseEventSizeLimit { + fn new(maximum_bytes: Option) -> Self { + Self { + maximum_bytes, + retained_bytes: 0, + line_bytes: 0, + line_is_comment: false, + previous_was_carriage_return: false, + failed: false, + } + } + + fn observe(&mut self, bytes: &[u8]) -> io::Result<()> { + if self.failed { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "oversized MCP SSE event was already rejected", + )); + } + let Some(maximum_bytes) = self.maximum_bytes else { + return Ok(()); + }; + + for &byte in bytes { + if self.previous_was_carriage_return { + self.previous_was_carriage_return = false; + if byte == b'\n' { + continue; + } + } + + match byte { + b'\r' => { + self.finish_line(maximum_bytes)?; + self.previous_was_carriage_return = true; + } + b'\n' => self.finish_line(maximum_bytes)?, + _ => { + if self.line_bytes == 0 { + self.line_is_comment = byte == b':'; + } + self.line_bytes = self.line_bytes.saturating_add(1); + self.check_limit(maximum_bytes)?; + } + } + } + Ok(()) + } + + fn finish_line(&mut self, maximum_bytes: usize) -> io::Result<()> { + if self.line_bytes == 0 { + self.retained_bytes = 0; + } else if !self.line_is_comment { + // The SSE parser inserts a newline when joining multiple data fields. + self.retained_bytes = self + .retained_bytes + .saturating_add(self.line_bytes) + .saturating_add(1); + } + + self.line_bytes = 0; + self.line_is_comment = false; + self.check_limit(maximum_bytes) + } + + fn check_limit(&mut self, maximum_bytes: usize) -> io::Result<()> { + if self.retained_bytes.saturating_add(self.line_bytes) > maximum_bytes { + self.failed = true; + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("MCP response body exceeds {maximum_bytes} bytes"), + )); + } + Ok(()) + } +} + +#[cfg(test)] +#[path = "http_client_adapter_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/http_client_adapter_tests.rs b/codex-rs/rmcp-client/src/http_client_adapter_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..c3effd99831abccd3fbf8b6665a6af38555f0f10 --- /dev/null +++ b/codex-rs/rmcp-client/src/http_client_adapter_tests.rs @@ -0,0 +1,285 @@ +use std::io::ErrorKind; + +use codex_exec_server::HttpRedirectPolicy; +use http::HeaderMap; +use http::HeaderValue; +use http::header::AUTHORIZATION; +use pretty_assertions::assert_eq; + +use super::HttpHeader; +use super::SseEventSizeLimit; +use super::StreamableHttpRedirectMode; +use super::mcp_redirect_policy; +use super::protocol_headers; + +#[test] +fn protocol_headers_preserve_utf8_values() { + let mut headers = HeaderMap::new(); + headers.insert( + "x-plugin-name", + HeaderValue::from_str("café").expect("valid HTTP field value"), + ); + + assert_eq!( + protocol_headers(&headers), + vec![HttpHeader { + name: "x-plugin-name".to_string(), + value: "café".to_string(), + value_env_var: None, + }] + ); +} + +#[test] +fn legacy_configured_headers_follow_redirects() { + assert_eq!( + mcp_redirect_policy( + StreamableHttpRedirectMode::Legacy, + &HeaderMap::new(), + /*has_configured_headers*/ true, + ), + HttpRedirectPolicy::Follow + ); +} + +#[test] +fn agent_plugin_configured_headers_stop_redirects() { + assert_eq!( + mcp_redirect_policy( + StreamableHttpRedirectMode::AgentPluginV1, + &HeaderMap::new(), + /*has_configured_headers*/ true, + ), + HttpRedirectPolicy::Stop + ); +} + +#[test] +fn requests_without_sensitive_headers_follow_redirects() { + assert_eq!( + mcp_redirect_policy( + StreamableHttpRedirectMode::AgentPluginV1, + &HeaderMap::new(), + /*has_configured_headers*/ false, + ), + HttpRedirectPolicy::Follow + ); +} + +#[test] +fn authorization_redirects_depend_on_mode() { + let mut headers = HeaderMap::new(); + headers.insert(AUTHORIZATION, HeaderValue::from_static("Bearer secret")); + assert_eq!( + mcp_redirect_policy( + StreamableHttpRedirectMode::Legacy, + &headers, + /*has_configured_headers*/ false, + ), + HttpRedirectPolicy::Follow + ); + assert_eq!( + mcp_redirect_policy( + StreamableHttpRedirectMode::AgentPluginV1, + &headers, + /*has_configured_headers*/ false, + ), + HttpRedirectPolicy::Stop + ); +} + +#[test] +fn lf_terminators_reset_the_event_limit() { + let mut limit = SseEventSizeLimit::new(Some(8)); + + limit + .observe(b"data: a\n\ndata: b\n\n") + .expect("LF-terminated events must have independent size limits"); + + assert_eq!((limit.retained_bytes, limit.line_bytes), (0, 0)); +} + +#[test] +fn carriage_return_terminators_reset_the_event_limit() { + let mut limit = SseEventSizeLimit::new(Some(8)); + + limit + .observe(b"data: a\r\rdata: b\r\r") + .expect("CR-terminated events must have independent size limits"); + + assert_eq!((limit.retained_bytes, limit.line_bytes), (0, 0)); +} + +#[test] +fn crlf_terminators_split_across_chunks_reset_the_event_limit() { + let mut limit = SseEventSizeLimit::new(Some(8)); + + limit + .observe(b"data: a\r") + .expect("a CR must finish the first event field"); + limit + .observe(b"\n\r") + .expect("a split CRLF must not finish another field"); + limit + .observe(b"\ndata: b\r") + .expect("the blank CRLF must reset the first event"); + limit + .observe(b"\n\r\n") + .expect("the second split CRLF event must remain within the limit"); + + assert_eq!((limit.retained_bytes, limit.line_bytes), (0, 0)); +} + +#[test] +fn event_at_the_exact_size_limit_is_accepted() { + let mut limit = SseEventSizeLimit::new(Some(9)); + + limit + .observe(b"data: ab\n\n") + .expect("an event at the exact limit must be accepted"); + + assert_eq!((limit.retained_bytes, limit.line_bytes), (0, 0)); +} + +#[test] +fn completed_keepalive_comments_do_not_accumulate() { + let mut limit = SseEventSizeLimit::new(Some(6)); + + limit + .observe(&b": ping\n".repeat(/*n*/ 64)) + .expect("completed keepalive comments must not accumulate"); + + assert_eq!((limit.retained_bytes, limit.line_bytes), (0, 0)); +} + +#[test] +fn split_keepalive_comments_are_discarded_only_when_complete() { + let mut limit = SseEventSizeLimit::new(Some(6)); + + limit + .observe(b": pi") + .expect("an incomplete comment within the limit must be retained"); + assert_eq!((limit.retained_bytes, limit.line_bytes), (0, 4)); + + limit + .observe(b"ng\n: ping\n") + .expect("only completed comments may be excluded from the limit"); + + assert_eq!((limit.retained_bytes, limit.line_bytes), (0, 0)); +} + +#[test] +fn comments_do_not_reset_accumulated_event_data() { + let mut limit = SseEventSizeLimit::new(Some(14)); + + limit + .observe(b"data: a\n") + .expect("the first data field must fit"); + limit + .observe(b": ping\n") + .expect("a completed comment must not count as event data"); + + let error = limit + .observe(b"data: b\n") + .expect_err("a comment must not reset previously retained data"); + + assert_eq!( + (error.kind(), error.to_string()), + ( + ErrorKind::InvalidData, + "MCP response body exceeds 14 bytes".to_string(), + ) + ); +} + +#[test] +fn an_unterminated_comment_remains_size_limited() { + let mut limit = SseEventSizeLimit::new(Some(6)); + + limit + .observe(b": ping") + .expect("a comment at the exact limit must be accepted"); + + let error = limit + .observe(b"!") + .expect_err("an unterminated comment must not bypass the limit"); + + assert_eq!( + (error.kind(), error.to_string()), + ( + ErrorKind::InvalidData, + "MCP response body exceeds 6 bytes".to_string(), + ) + ); +} + +#[test] +fn multiline_data_counts_parser_inserted_newlines() { + let mut limit = SseEventSizeLimit::new(Some(18)); + + let error = limit + .observe(b"data: aaa\ndata: bbb\n\n") + .expect_err("joined multiline event data must remain size limited"); + + assert_eq!( + (error.kind(), error.to_string()), + ( + ErrorKind::InvalidData, + "MCP response body exceeds 18 bytes".to_string(), + ) + ); +} + +#[test] +fn legacy_sse_streams_remain_unlimited() { + let mut limit = SseEventSizeLimit::new(/*maximum_bytes*/ None); + + limit + .observe(&[b'x'; 128]) + .expect("legacy SSE streams must not gain a message limit"); + + assert_eq!( + ( + limit.retained_bytes, + limit.line_bytes, + limit.line_is_comment, + limit.previous_was_carriage_return, + limit.failed, + ), + (0, 0, false, false, false) + ); +} + +#[test] +fn rejected_events_remain_rejected() { + let mut limit = SseEventSizeLimit::new(Some(8)); + + limit + .observe(b"data: abc") + .expect_err("an oversized event must be rejected"); + + let error = limit + .observe(b"") + .expect_err("a rejected event must not be resumed"); + + assert_eq!( + (error.kind(), error.to_string()), + ( + ErrorKind::InvalidData, + "oversized MCP SSE event was already rejected".to_string(), + ) + ); +} + +#[test] +fn size_accounting_saturates_without_overflow() { + let mut limit = SseEventSizeLimit::new(Some(usize::MAX - 1)); + limit.retained_bytes = usize::MAX - 2; + limit.line_bytes = 1; + + let error = limit + .observe(b"x") + .expect_err("saturating counts must still reject an oversized event"); + + assert_eq!((error.kind(), limit.failed), (ErrorKind::InvalidData, true)); +} diff --git a/codex-rs/rmcp-client/src/http_client_redirect.rs b/codex-rs/rmcp-client/src/http_client_redirect.rs new file mode 100644 index 0000000000000000000000000000000000000000..5d8cfee05f889421aa003bc37bae074d4c1a1475 --- /dev/null +++ b/codex-rs/rmcp-client/src/http_client_redirect.rs @@ -0,0 +1,222 @@ +//! Restricts MCP HTTP redirects to the configured server's origin. +//! +//! MCP requests can carry sensitive headers and tool-call bodies. Following a +//! cross-origin redirect would send them to another server or an internal +//! service, so this transport validates each redirect before following it. + +use std::sync::Arc; +use std::time::Duration; +use std::time::Instant; + +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpHeader; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +use futures::FutureExt; +use futures::future::BoxFuture; +use http::HeaderValue; +use http::StatusCode; +use url::Url; + +const MAX_REDIRECTS: usize = 10; + +pub(crate) struct SameOriginRedirectHttpClient { + inner: Arc, +} + +enum RedirectResponse { + Buffered(HttpRequestResponse), + Streaming(HttpRequestResponse, HttpResponseBodyStream), +} + +impl SameOriginRedirectHttpClient { + pub(crate) fn new(inner: Arc) -> Self { + Self { inner } + } + + async fn execute( + &self, + mut params: HttpRequestParams, + ) -> Result { + // Inspect each Location ourselves before the underlying client can send + // sensitive request data to the redirect destination. + params.redirect_policy = HttpRedirectPolicy::Stop; + let mut current_url = Url::parse(¶ms.url) + .map_err(|error| ExecServerError::HttpRequest(error.to_string()))?; + let original_origin = current_url.origin(); + // Redirect hops share the original timeout instead of restarting it. + let deadline = params + .timeout_ms + .map(|timeout_ms| Instant::now() + Duration::from_millis(timeout_ms)); + let mut redirects = 0; + + loop { + if let Some(deadline) = deadline { + let Some(remaining) = deadline.checked_duration_since(Instant::now()) else { + return Err(ExecServerError::HttpRequest( + "MCP HTTP request timed out".to_string(), + )); + }; + params.timeout_ms = Some( + u64::try_from(remaining.as_millis()) + .unwrap_or(u64::MAX) + .max(1), + ); + } + + let result = if params.stream_response { + let (response, stream) = self.inner.http_request_stream(params.clone()).await?; + RedirectResponse::Streaming(response, stream) + } else { + RedirectResponse::Buffered(self.inner.http_request(params.clone()).await?) + }; + let response = match &result { + RedirectResponse::Buffered(response) | RedirectResponse::Streaming(response, _) => { + response + } + }; + let status = StatusCode::from_u16(response.status).ok(); + if !matches!( + status, + Some( + StatusCode::MOVED_PERMANENTLY + | StatusCode::FOUND + | StatusCode::SEE_OTHER + | StatusCode::TEMPORARY_REDIRECT + | StatusCode::PERMANENT_REDIRECT + ) + ) { + return Ok(result); + } + + let Some(location) = response + .headers + .iter() + .find(|header| header.name.eq_ignore_ascii_case("location")) + .map(|header| header.value.as_str()) + else { + return Ok(result); + }; + let next_url = current_url + .join(location) + .map_err(|error| ExecServerError::HttpRequest(error.to_string()))?; + // Compare every hop with the configured origin so redirect chains + // cannot escape it through a later cross-origin Location. + if next_url.origin() != original_origin { + return Err(ExecServerError::HttpRequest( + "MCP HTTP redirect to a different origin is not allowed".to_string(), + )); + } + // A hostname can resolve to a different address on the next HTTP + // connection. HTTPS authenticates the hostname; localhost and IP + // literals do not depend on attacker-controlled DNS. + if next_url.scheme() == "http" + && matches!(next_url.host(), Some(url::Host::Domain(host)) if host != "localhost") + { + return Err(ExecServerError::HttpRequest( + "MCP HTTP redirects for non-loopback hostnames require HTTPS".to_string(), + )); + } + if redirects >= MAX_REDIRECTS { + return Err(ExecServerError::HttpRequest( + "MCP HTTP request exceeded the redirect limit".to_string(), + )); + } + + // Preserve normal HTTP behavior: 307/308 retain the method and body; + // 303 and legacy 301/302 POST redirects switch to a bodyless GET. + let drop_body = matches!( + status, + Some(StatusCode::MOVED_PERMANENTLY | StatusCode::FOUND) + ) && params.method.eq_ignore_ascii_case("POST") + || status == Some(StatusCode::SEE_OTHER); + if drop_body { + if !params.method.eq_ignore_ascii_case("HEAD") { + params.method = "GET".to_string(); + } + params.body = None; + params.headers.retain(|header| { + ![ + "content-type", + "content-length", + "content-encoding", + "transfer-encoding", + ] + .iter() + .any(|name| header.name.eq_ignore_ascii_case(name)) + }); + } + + // HTTPS keeps application proxy credentials inside the tunnel to + // the already-validated origin. Plaintext routes can change across + // requests, so their proxy credentials must never be replayed. + params.headers.retain(|header| { + (current_url.scheme() == "https" + || !header.name.eq_ignore_ascii_case("proxy-authorization")) + && !header.name.eq_ignore_ascii_case("referer") + }); + let mut referer = current_url.clone(); + let _ = referer.set_username(""); + let _ = referer.set_password(None); + referer.set_fragment(None); + if HeaderValue::from_str(referer.as_str()).is_ok() { + params.headers.push(HttpHeader { + name: "referer".to_string(), + value: referer.to_string(), + value_env_var: None, + }); + } + + params.url = next_url.to_string(); + current_url = next_url; + redirects += 1; + } + } +} + +impl HttpClient for SameOriginRedirectHttpClient { + fn http_request( + &self, + mut params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + async move { + params.stream_response = false; + if params.redirect_policy == HttpRedirectPolicy::Stop { + return self.inner.http_request(params).await; + } + match self.execute(params).await? { + RedirectResponse::Buffered(response) => Ok(response), + RedirectResponse::Streaming(_, _) => Err(ExecServerError::Protocol( + "MCP buffered HTTP request returned a streamed response".to_string(), + )), + } + } + .boxed() + } + + fn http_request_stream( + &self, + mut params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + async move { + params.stream_response = true; + if params.redirect_policy == HttpRedirectPolicy::Stop { + return self.inner.http_request_stream(params).await; + } + match self.execute(params).await? { + RedirectResponse::Streaming(response, stream) => Ok((response, stream)), + RedirectResponse::Buffered(_) => Err(ExecServerError::Protocol( + "MCP streamed HTTP request returned a buffered response".to_string(), + )), + } + } + .boxed() + } +} + +#[cfg(test)] +#[path = "http_client_redirect_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/http_client_redirect_tests.rs b/codex-rs/rmcp-client/src/http_client_redirect_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..966fcfcfcbc87cda084bdcd27b05cc3dced44ccc --- /dev/null +++ b/codex-rs/rmcp-client/src/http_client_redirect_tests.rs @@ -0,0 +1,391 @@ +use std::sync::Arc; +use std::sync::Mutex; + +use codex_exec_server::Environment; +use codex_exec_server::HttpHeader; +use pretty_assertions::assert_eq; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::path; + +use crate::http_headers::with_http_headers_helper; + +use super::*; + +const PROXY_HEADERS_HELPER: &str = if cfg!(windows) { + r#"echo {"Proxy-Authorization":"Bearer proxy-token"}"# +} else { + r#"printf '{"Proxy-Authorization":"Bearer proxy-token"}'"# +}; + +fn request(url: impl Into) -> HttpRequestParams { + HttpRequestParams { + method: "POST".to_string(), + url: url.into(), + headers: Vec::new(), + body: Some(b"sensitive-body".to_vec().into()), + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "redirect-test".to_string(), + stream_response: false, + } +} + +fn headers(headers: [(&str, &str); N]) -> Vec { + headers + .into_iter() + .map(|(name, value)| HttpHeader { + name: name.to_string(), + value: value.to_string(), + value_env_var: None, + }) + .collect() +} + +#[derive(Default)] +struct RecordingRedirectHttpClient { + requests: Mutex>, + delay: Duration, + loop_redirects: bool, +} + +impl HttpClient for RecordingRedirectHttpClient { + fn http_request( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + let mut requests = self.requests.lock().expect("request recorder lock"); + let redirect = requests.is_empty() || self.loop_redirects; + let delay = self.delay; + requests.push(params); + async move { + if !delay.is_zero() { + tokio::time::sleep(delay).await; + } + Ok(HttpRequestResponse { + status: if redirect { 307 } else { 200 }, + headers: if redirect { + headers([("location", "/final")]) + } else { + Vec::new() + }, + body: Vec::new().into(), + }) + } + .boxed() + } + + fn http_request_stream( + &self, + _params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + async { + Err(ExecServerError::HttpRequest( + "unexpected streaming request".to_string(), + )) + } + .boxed() + } +} + +#[tokio::test] +async fn plaintext_hostname_redirects_are_rejected_before_dns_can_rebind() -> anyhow::Result<()> { + for (url, expected_error) in [ + ( + "http://mcp.example/start", + Some("non-loopback hostnames require HTTPS"), + ), + ("http://localhost/start", None), + ("http://127.0.0.1/start", None), + ("http://[::1]/start", None), + ("https://mcp.example/start", None), + ] { + let recorder = Arc::new(RecordingRedirectHttpClient::default()); + let client = SameOriginRedirectHttpClient::new(recorder.clone()); + let response = client + .http_request(HttpRequestParams { + headers: headers([("authorization", "Bearer sensitive-token")]), + ..request(url) + }) + .await; + + if let Some(expected_error) = expected_error { + let error = response.expect_err("plaintext hostname redirects must fail"); + assert!(error.to_string().contains(expected_error), "{error}"); + } else { + assert_eq!(response?.status, 200); + } + let requests = recorder.requests.lock().expect("request recorder lock"); + assert_eq!(requests.len(), if expected_error.is_some() { 1 } else { 2 }); + } + + Ok(()) +} + +#[tokio::test] +async fn same_origin_redirects_enforce_shared_timeout_and_hop_limit() { + for (delay, timeout_ms, expected_error, expected_requests) in [ + (Duration::from_millis(100), 175, "timed out", 2), + (Duration::ZERO, 5_000, "redirect limit", MAX_REDIRECTS + 1), + ] { + let recorder = Arc::new(RecordingRedirectHttpClient { + delay, + loop_redirects: true, + ..Default::default() + }); + let client = SameOriginRedirectHttpClient::new(recorder.clone()); + let error = client + .http_request(HttpRequestParams { + method: "GET".to_string(), + body: None, + timeout_ms: Some(timeout_ms), + ..request("https://mcp.example/loop") + }) + .await + .expect_err("redirect loops must respect both limits"); + + assert!(error.to_string().contains(expected_error), "{error}"); + let requests = recorder.requests.lock().expect("request recorder lock"); + assert_eq!(requests.len(), expected_requests); + } +} + +#[tokio::test] +async fn same_origin_redirects_preserve_method_body_and_headers() -> anyhow::Result<()> { + for (method, redirect_status, redirected_method) in [ + ("POST", 301, "GET"), + ("POST", 302, "GET"), + ("POST", 303, "GET"), + ("HEAD", 303, "HEAD"), + ("POST", 307, "POST"), + ("GET", 307, "GET"), + ("DELETE", 307, "DELETE"), + ("POST", 308, "POST"), + ] { + let server = MockServer::start().await; + Mock::given(path("/start")) + .respond_with( + ResponseTemplate::new(redirect_status).insert_header("location", "/final"), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(path("/final")) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&server) + .await; + + let client = + SameOriginRedirectHttpClient::new(Environment::default_for_tests().get_http_client()); + let expected_content_length = b"sensitive-body".len().to_string(); + let mut params = HttpRequestParams { + method: method.to_string(), + headers: headers([ + ("x-api-key", "sensitive-key"), + ("proxy-authorization", "sensitive-proxy-credentials"), + ("referer", "https://stale.example/"), + ]), + body: (method == "POST").then(|| b"sensitive-body".to_vec().into()), + stream_response: method != "DELETE", + ..request(format!("{}/start", server.uri())) + }; + if method == "POST" { + params.headers.extend(headers([ + ("content-type", "application/json"), + ("content-encoding", "identity"), + ("content-length", &expected_content_length), + ])); + } + let status = if method == "DELETE" { + client.http_request(params).await?.status + } else { + client.http_request_stream(params).await?.0.status + }; + assert_eq!(status, 200); + + let requests = server.received_requests().await.expect("recorded requests"); + let redirected = &requests[1]; + let expected_referer = format!("{}/start", server.uri()); + let body_preserved = redirected_method == "POST"; + for (name, expected) in [ + ("x-api-key", Some("sensitive-key")), + ("proxy-authorization", None), + ("referer", Some(expected_referer.as_str())), + ("content-type", body_preserved.then_some("application/json")), + ("content-encoding", body_preserved.then_some("identity")), + ( + "content-length", + body_preserved.then_some(expected_content_length.as_str()), + ), + ] { + assert_eq!( + redirected + .headers + .get(name) + .and_then(|value| value.to_str().ok()), + expected, + "unexpected redirected {name} header" + ); + } + assert_eq!( + (redirected.method.as_str(), redirected.body.as_slice()), + ( + redirected_method, + if body_preserved { + b"sensitive-body".as_slice() + } else { + b"".as_slice() + } + ) + ); + } + + Ok(()) +} + +#[tokio::test] +async fn https_redirects_preserve_configured_and_helper_proxy_authorization() -> anyhow::Result<()> +{ + for helper_enabled in [false, true] { + let recorder = Arc::new(RecordingRedirectHttpClient::default()); + let inner: Arc = recorder.clone(); + let directory = tempfile::tempdir()?; + let url = "https://mcp.example/start"; + let inner = if helper_enabled { + with_http_headers_helper( + inner, + url, + PROXY_HEADERS_HELPER, + directory.path().to_path_buf(), + )? + } else { + inner + }; + let client = SameOriginRedirectHttpClient::new(inner); + let response = client + .http_request(HttpRequestParams { + headers: if helper_enabled { + Vec::new() + } else { + headers([("Proxy-Authorization", "Bearer proxy-token")]) + }, + ..request(url) + }) + .await?; + + let requests = recorder.requests.lock().expect("request recorder lock"); + assert_eq!( + ( + response.status, + requests.len(), + requests[1].url.as_str(), + requests[1] + .headers + .iter() + .find(|header| header.name.eq_ignore_ascii_case("proxy-authorization")) + .map(|header| header.value.as_str()), + ), + ( + 200, + 2, + "https://mcp.example/final", + Some("Bearer proxy-token"), + ), + ); + } + + Ok(()) +} + +#[tokio::test] +async fn plaintext_helper_redirects_block_mcp_but_preserve_oauth_stop() -> anyhow::Result<()> { + for (request_id, redirect_policy) in [ + ("mcp-request-1", HttpRedirectPolicy::Follow), + ("oauth-request-1", HttpRedirectPolicy::Stop), + ] { + let server = MockServer::start().await; + Mock::given(path("/start")) + .respond_with(ResponseTemplate::new(307).insert_header("location", "/final")) + .expect(1) + .mount(&server) + .await; + + let directory = tempfile::tempdir()?; + let url = format!("{}/start", server.uri()); + let helper = with_http_headers_helper( + Environment::default_for_tests().get_http_client(), + &url, + PROXY_HEADERS_HELPER, + directory.path().to_path_buf(), + )?; + let client = SameOriginRedirectHttpClient::new(helper); + let response = client + .http_request_stream(HttpRequestParams { + request_id: request_id.to_string(), + redirect_policy, + stream_response: true, + ..request(url) + }) + .await; + + match (redirect_policy, response) { + (HttpRedirectPolicy::Stop, Ok((response, _))) => assert_eq!(response.status, 307), + (HttpRedirectPolicy::Follow, Err(error)) => { + assert!(error.to_string().contains("Proxy-Authorization")); + } + _ => panic!("unexpected plaintext proxy-credential redirect behavior"), + } + assert_eq!(server.received_requests().await.unwrap().len(), 1); + } + Ok(()) +} + +#[tokio::test] +async fn cross_origin_redirects_never_reach_their_destination() -> anyhow::Result<()> { + for status in [307, 308] { + for method in ["POST", "GET", "DELETE"] { + let destination = MockServer::start().await; + let server = MockServer::start().await; + Mock::given(path("/start")) + .respond_with( + ResponseTemplate::new(status) + .insert_header("location", format!("{}/private", destination.uri())), + ) + .expect(1) + .mount(&server) + .await; + + let client = SameOriginRedirectHttpClient::new( + Environment::default_for_tests().get_http_client(), + ); + let params = HttpRequestParams { + method: method.to_string(), + headers: headers([("x-api-key", "sensitive-key")]), + body: (method == "POST").then(|| b"sensitive-body".to_vec().into()), + stream_response: method != "DELETE", + ..request(format!("{}/start", server.uri())) + }; + let error = if method == "DELETE" { + client + .http_request(params) + .await + .expect_err("cross-origin DELETE redirect must fail") + } else { + match client.http_request_stream(params).await { + Ok(_) => panic!("cross-origin {method} redirect must fail"), + Err(error) => error, + } + }; + + assert!( + error.to_string().contains("different origin"), + "cross-origin {method} redirect must explain its rejection: {error}" + ); + assert!(destination.received_requests().await.unwrap().is_empty()); + } + } + + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/http_headers.rs b/codex-rs/rmcp-client/src/http_headers.rs new file mode 100644 index 0000000000000000000000000000000000000000..d0b2d28c989be47b2f992f54e0423c544c882cdf --- /dev/null +++ b/codex-rs/rmcp-client/src/http_headers.rs @@ -0,0 +1,572 @@ +use std::collections::HashSet; +use std::fmt; +#[cfg(windows)] +use std::os::windows::process::CommandExt; +use std::path::Path; +use std::path::PathBuf; +use std::process::Stdio; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; +use std::time::Duration; + +use anyhow::Result; +use anyhow::anyhow; +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpHeader; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +#[cfg(all(unix, not(target_os = "macos")))] +use codex_utils_pty::process_group::kill_process_group; +#[cfg(target_os = "macos")] +use codex_utils_pty::process_group::kill_process_group_with_member_fallback as kill_process_group; +use futures::FutureExt; +use futures::future::BoxFuture; +use futures::future::Shared; +use http::HeaderMap; +use http::HeaderName; +use http::HeaderValue; +use serde::Deserialize; +use serde::de::MapAccess; +use serde::de::Visitor; +use tokio::io::AsyncReadExt; +use tokio::process::Child; +use tokio::process::Command; +use tokio::time::Instant; +use url::Origin; +use url::Url; + +use crate::utils::create_env_for_mcp_server; +use crate::www_authenticate::insufficient_scope_challenge; + +const HELPER_TIMEOUT: Duration = Duration::from_secs(10); +const MAX_HELPER_OUTPUT_BYTES: usize = 64 * 1024; +type CachedHeaders = Shared, Arc>>>; + +struct HttpHeadersProvider { + server_origin: Origin, + command: String, + cwd: PathBuf, + cache: Mutex, +} + +struct HeadersCache { + // Identifies a cohort of rejected requests, not a credential version. Advance even after + // failed or unchanged refreshes so concurrent rejections share one helper invocation. + refresh_epoch: u64, + current: CachedHeaders, + refresh: Option, +} + +struct RequestHeaders { + refresh_epoch: u64, + values: Arc, +} + +struct HttpHeadersClient { + inner: Arc, + provider: HttpHeadersProvider, +} + +struct HelperProcess { + child: Child, + #[cfg(unix)] + process_group_id: u32, + #[cfg(windows)] + job: codex_utils_pty::JobObject, +} + +struct RawHeaderEntries { + // Keep raw entries because ordinary map deserialization collapses exact duplicate keys. + entries: Vec<(String, String)>, + has_exact_duplicate: bool, +} + +struct RawHeaderEntriesVisitor; + +impl<'de> Visitor<'de> for RawHeaderEntriesVisitor { + type Value = RawHeaderEntries; + + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("a JSON object of string header names and values") + } + + fn visit_map(self, mut map: A) -> Result + where + A: MapAccess<'de>, + { + let mut entries = Vec::with_capacity(map.size_hint().unwrap_or_default()); + let mut names = HashSet::new(); + let mut has_exact_duplicate = false; + while let Some((name, value)) = map.next_entry::()? { + has_exact_duplicate |= !names.insert(name.clone()); + entries.push((name, value)); + } + Ok(RawHeaderEntries { + entries, + has_exact_duplicate, + }) + } +} + +impl<'de> Deserialize<'de> for RawHeaderEntries { + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + deserializer.deserialize_map(RawHeaderEntriesVisitor) + } +} + +impl Drop for HelperProcess { + fn drop(&mut self) { + #[cfg(unix)] + let _ = kill_process_group(self.process_group_id); + + #[cfg(windows)] + let _ = self.job.terminate(); + + let _ = self.child.start_kill(); + } +} + +impl HttpHeadersProvider { + fn new(server_url: &str, command: &str, cwd: PathBuf) -> Result { + let command = command.to_string(); + let cached = Self::helper_attempt(command.clone(), cwd.clone()); + Ok(Self { + server_origin: Url::parse(server_url)?.origin(), + command, + cwd, + cache: Mutex::new(HeadersCache { + refresh_epoch: 0, + current: cached, + refresh: None, + }), + }) + } + + fn helper_attempt(command: String, cwd: PathBuf) -> CachedHeaders { + async move { + // Keep helper cleanup running after caller cancellation; dropping the set aborts it. + let mut tasks = tokio::task::JoinSet::new(); + tasks.spawn(async move { + run_helper(&command, &cwd) + .await + .map(Arc::new) + .map_err(|error| Arc::::from(error.to_string())) + }); + tasks + .join_next() + .await + .ok_or_else(|| Arc::::from("MCP HTTP headers helper task was unavailable"))? + .map_err(|_| Arc::::from("MCP HTTP headers helper task failed"))? + } + .boxed() + .shared() + } + + async fn headers(&self) -> Result { + let (refresh_epoch, current) = { + let cache = self.cache.lock().unwrap_or_else(PoisonError::into_inner); + (cache.refresh_epoch, cache.current.clone()) + }; + current + .await + .map(|values| RequestHeaders { + refresh_epoch, + values, + }) + .map_err(|error| ExecServerError::HttpRequest(error.to_string())) + } + + async fn refresh(&self, rejected_epoch: u64) -> Result, ExecServerError> { + let attempt = { + let mut cache = self.cache.lock().unwrap_or_else(PoisonError::into_inner); + if cache.refresh_epoch != rejected_epoch { + cache.current.clone() + } else { + cache + .refresh + .get_or_insert_with(|| { + Self::helper_attempt(self.command.clone(), self.cwd.clone()) + }) + .clone() + } + }; + let result = attempt.clone().await; + let mut cache = self.cache.lock().unwrap_or_else(PoisonError::into_inner); + if cache.refresh_epoch == rejected_epoch + && cache + .refresh + .as_ref() + .is_some_and(|refresh| refresh.ptr_eq(&attempt)) + { + cache.refresh_epoch = rejected_epoch.saturating_add(/*rhs*/ 1); + if result.is_ok() { + cache.current = attempt; + } + cache.refresh = None; + } + result.map_err(|error| ExecServerError::HttpRequest(error.to_string())) + } +} + +/// Refreshes helper headers once after a same-origin POST returns 401/403, retrying only +/// when helper-provided values change. OAuth challenges survive failed or unchanged refreshes. +/// Helper Authorization is used only when the request has no explicit bearer/OAuth credential. +/// The helper is a short-lived process; independently managed daemons remain caller-owned. +pub fn with_http_headers_helper( + inner: Arc, + server_url: &str, + command: &str, + cwd: PathBuf, +) -> Result> { + let provider = HttpHeadersProvider::new(server_url, command, cwd)?; + Ok(Arc::new(HttpHeadersClient { inner, provider })) +} + +impl HttpHeadersClient { + async fn prepare_request( + &self, + params: HttpRequestParams, + ) -> Result<(HttpRequestParams, Option, Option), ExecServerError> { + let Ok(url) = Url::parse(¶ms.url) else { + return Ok((params, None, None)); + }; + if self.provider.server_origin != url.origin() { + return Ok((params, None, None)); + } + + let deadline = params + .timeout_ms + .map(|timeout_ms| Instant::now() + Duration::from_millis(timeout_ms)); + let headers = match deadline { + Some(deadline) => tokio::time::timeout_at(deadline, self.provider.headers()) + .await + .map_err(|_| { + ExecServerError::HttpRequest("HTTP request timed out".to_string()) + })??, + None => self.provider.headers().await?, + }; + let params = self.apply_headers(params, &headers.values, deadline)?; + Ok((params, Some(headers), deadline)) + } + + fn apply_headers( + &self, + mut params: HttpRequestParams, + headers: &HeaderMap, + deadline: Option, + ) -> Result { + // TODO: Follow same-origin redirects once later hops cannot leak helper headers. + params.redirect_policy = HttpRedirectPolicy::Stop; + for (name, value) in headers { + // Explicit bearer tokens, MCP OAuth and token-endpoint client authentication win. + if name == http::header::AUTHORIZATION + && params + .headers + .iter() + .any(|header| header.name.eq_ignore_ascii_case("authorization")) + { + continue; + } + params + .headers + .retain(|header| !header.name.eq_ignore_ascii_case(name.as_str())); + params.headers.push(HttpHeader { + name: name.to_string(), + value: std::str::from_utf8(value.as_bytes()) + .map_err(|error| ExecServerError::HttpRequest(error.to_string()))? + .to_string(), + value_env_var: None, + }); + } + if let Some(deadline) = deadline { + let remaining = deadline.saturating_duration_since(Instant::now()); + params.timeout_ms = Some( + u64::try_from(remaining.as_millis()) + .unwrap_or(u64::MAX) + .max(1), + ); + } + Ok(params) + } + + fn reject_proxy_authorization_redirect( + response: &HttpRequestResponse, + ) -> Result<(), ExecServerError> { + if (300..400).contains(&response.status) + && response + .headers + .iter() + .any(|header| header.name.eq_ignore_ascii_case("location")) + { + return Err(ExecServerError::HttpRequest( + "MCP HTTP redirect cannot safely replay Proxy-Authorization credentials" + .to_string(), + )); + } + Ok(()) + } + + async fn retry_request( + &self, + params: HttpRequestParams, + rejected: RequestHeaders, + deadline: Option, + response: &HttpRequestResponse, + ) -> Option { + if !matches!(response.status, 401 | 403) + || response.status == 403 && insufficient_scope_challenge(&response.headers).is_some() + { + return None; + } + // A rejection may be an OAuth challenge, so preserve it if the helper cannot improve it. + let refreshed = match deadline { + Some(deadline) => { + tokio::time::timeout_at(deadline, self.provider.refresh(rejected.refresh_epoch)) + .await + .ok()? + .ok()? + } + None => self.provider.refresh(rejected.refresh_epoch).await.ok()?, + }; + // Header names are user-defined. Compare effective helper values, excluding an + // Authorization that would be overridden, without treating JSON key order as a change. + let explicit_authorization = params + .headers + .iter() + .any(|header| header.name.eq_ignore_ascii_case("authorization")); + let differs = |left: &HeaderMap, right: &HeaderMap| { + left.iter().any(|(name, value)| { + (!explicit_authorization || name != http::header::AUTHORIZATION) + && right.get(name) != Some(value) + }) + }; + if !differs(&rejected.values, &refreshed) && !differs(&refreshed, &rejected.values) { + return None; + } + self.apply_headers(params, &refreshed, deadline).ok() + } + + async fn request<'a, T>( + &'a self, + params: HttpRequestParams, + send: impl Fn(HttpRequestParams) -> BoxFuture<'a, Result>, + response_headers: impl Fn(&T) -> &HttpRequestResponse, + ) -> Result { + // OAuth validates redirect responses itself; only intercept MCP replay. + let mcp_redirect_was_stopped = params.redirect_policy == HttpRedirectPolicy::Stop + && !params.request_id.starts_with("oauth-request-"); + let original_headers = params + .method + .eq_ignore_ascii_case("POST") + .then(|| params.headers.clone()); + let (params, headers, deadline) = self.prepare_request(params).await?; + let original = original_headers + .filter(|_| headers.is_some()) + .map(|headers| { + let mut original = params.clone(); + original.headers = headers; + original + }); + let needs_redirect_check = |params: &HttpRequestParams| { + mcp_redirect_was_stopped + && Url::parse(¶ms.url).is_ok_and(|url| url.scheme() == "http") + && params + .headers + .iter() + .any(|header| header.name.eq_ignore_ascii_case("proxy-authorization")) + }; + let prevent_proxy_authorization_redirect = needs_redirect_check(¶ms); + let response = send(params).await?; + if prevent_proxy_authorization_redirect { + Self::reject_proxy_authorization_redirect(response_headers(&response))?; + } + if let (Some(original), Some(headers)) = (original, headers) + && let Some(retry) = self + .retry_request(original, headers, deadline, response_headers(&response)) + .await + { + drop(response); + let prevent_proxy_authorization_redirect = needs_redirect_check(&retry); + let response = send(retry).await?; + if prevent_proxy_authorization_redirect { + Self::reject_proxy_authorization_redirect(response_headers(&response))?; + } + return Ok(response); + } + Ok(response) + } +} + +impl HttpClient for HttpHeadersClient { + fn http_request( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + self.request( + params, + |params| self.inner.http_request(params), + |response| response, + ) + .boxed() + } + + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + self.request( + params, + |params| self.inner.http_request_stream(params), + |response| &response.0, + ) + .boxed() + } +} + +async fn run_helper(command: &str, cwd: &Path) -> Result { + #[cfg(windows)] + let shell = std::env::var_os("COMSPEC").unwrap_or_else(|| "cmd.exe".into()); + #[cfg(not(windows))] + let shell = "sh"; + + // Match the repository's existing shell-command convention. The command is ordinary + // configuration and may be visible in local process metadata; credentials belong in the + // JSON output rather than in the command text. + let mut process = Command::new(shell); + #[cfg(windows)] + { + process.args(["/Q", "/D", "/C"]); + process.as_std_mut().raw_arg(format!(r#""{command}""#)); + } + #[cfg(not(windows))] + process.args(["-c", command]); + #[cfg(unix)] + process.process_group(0); + process + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::null()) + .current_dir(cwd) + // Match local MCP subprocess policy; arbitrary ambient variables are not inherited. + .env_clear() + .envs(create_env_for_mcp_server(/*extra_env*/ None, &[])?) + .kill_on_drop(true); + + #[cfg(windows)] + let (child, job) = { + let job = codex_utils_pty::JobObject::create_without_breakaway() + .map_err(|error| anyhow!("MCP HTTP headers helper containment failed: {error}"))?; + let child = job + .spawn_contained(&mut process) + .map_err(|error| anyhow!("MCP HTTP headers helper failed to start: {error}"))?; + (child, job) + }; + #[cfg(not(windows))] + let child = process + .spawn() + .map_err(|error| anyhow!("MCP HTTP headers helper failed to start: {error}"))?; + let mut process = HelperProcess { + #[cfg(unix)] + process_group_id: child + .id() + .ok_or_else(|| anyhow!("MCP HTTP headers helper process id was unavailable"))?, + child, + #[cfg(windows)] + job, + }; + let output = tokio::time::timeout(HELPER_TIMEOUT, async { + let stdout = process + .child + .stdout + .take() + .ok_or_else(|| anyhow!("MCP HTTP headers helper stdout was unavailable"))?; + let mut output = Vec::new(); + stdout + .take((MAX_HELPER_OUTPUT_BYTES + 1) as u64) + .read_to_end(&mut output) + .await?; + if output.len() > MAX_HELPER_OUTPUT_BYTES { + return Err(anyhow!("MCP HTTP headers helper output exceeds 64 KiB")); + } + let status = process.child.wait().await?; + if !status.success() { + return Err(anyhow!( + "MCP HTTP headers helper exited with status {status}" + )); + } + Ok(output) + }) + .await + .map_err(|_| anyhow!("MCP HTTP headers helper timed out after 10 seconds"))??; + + parse_helper_output(output) +} + +fn parse_helper_output(stdout: Vec) -> Result { + let stdout = String::from_utf8(stdout) + .map_err(|_| anyhow!("MCP HTTP headers helper wrote non-UTF-8 data"))?; + let mut deserializer = serde_json::Deserializer::from_str(stdout.trim()); + let headers = RawHeaderEntries::deserialize(&mut deserializer) + .and_then(|headers| { + deserializer.end()?; + Ok(headers) + }) + .map_err(|_| anyhow!("MCP HTTP headers helper must output a JSON object of strings"))?; + if headers.has_exact_duplicate { + return Err(anyhow!( + "MCP HTTP headers helper returned duplicate header names" + )); + } + let mut parsed = HeaderMap::with_capacity(headers.entries.len()); + for (name, value) in headers.entries { + let name = HeaderName::from_bytes(name.as_bytes()) + .map_err(|_| anyhow!("MCP HTTP headers helper returned an invalid header name"))?; + // Helper values replace same-name configured headers, except explicit Authorization. + // Google IAP uses Proxy-Authorization alongside application Authorization. For HTTPS MCP + // URLs it is sent through the forward-proxy tunnel to IAP, not used as CONNECT auth. + if matches!( + name.as_str(), + "accept" + | "connection" + | "content-encoding" + | "content-length" + | "content-type" + | "host" + | "keep-alive" + | "last-event-id" + | "mcp-protocol-version" + | "mcp-session-id" + | "origin" + | "proxy-connection" + | "referer" + | "te" + | "trailer" + | "transfer-encoding" + | "upgrade" + ) { + return Err(anyhow!( + "MCP HTTP headers helper returned a reserved header" + )); + } + if parsed.contains_key(&name) { + return Err(anyhow!( + "MCP HTTP headers helper returned duplicate header names" + )); + } + let value = HeaderValue::from_str(&value) + .map_err(|_| anyhow!("MCP HTTP headers helper returned an invalid header value"))?; + parsed.insert(name, value); + } + Ok(parsed) +} + +#[cfg(test)] +#[path = "http_headers_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/http_headers_tests.rs b/codex-rs/rmcp-client/src/http_headers_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..c27c5d10022de340f23bdf3e642ff7767943e1b2 --- /dev/null +++ b/codex-rs/rmcp-client/src/http_headers_tests.rs @@ -0,0 +1,533 @@ +use super::*; +use pretty_assertions::assert_eq; + +#[test] +fn helper_output_errors_do_not_echo_secrets() { + for output in [ + br#"{"Host":"secret"}"#.as_slice(), + br#"{"secret":"secret","secret":"secret"}"#.as_slice(), + br#"{"secret":"secret","Secret":"secret"}"#.as_slice(), + ] { + let error = parse_helper_output(output.to_vec()).expect_err("invalid helper output"); + assert!(!error.to_string().contains("secret")); + } +} + +#[cfg(unix)] +#[tokio::test] +async fn helper_attempt_is_shared_after_cancellation() { + use tempfile::tempdir; + + let temp = tempdir().expect("temporary helper directory"); + let cwd = temp + .path() + .canonicalize() + .expect("canonical helper directory"); + let cancelled_invocations = cwd.join("cancelled-invocations"); + let cancelled_finished = cwd.join("cancelled-helper-finished"); + let cancelled = HttpHeadersProvider::new( + "https://example.com", + &format!( + "test \"$(pwd)\" = '{0}'; test -n \"$HOME\"; test -n \"$PATH\"; \ + printf x >> '{1}'; sleep 0.2; printf x > '{2}'; \ + printf '{{\"X-Gateway\":\"token\"}}'", + cwd.display(), + cancelled_invocations.display(), + cancelled_finished.display(), + ), + cwd.clone(), + ) + .expect("cancelled provider"); + assert!( + tokio::time::timeout(Duration::from_millis(20), cancelled.headers()) + .await + .is_err() + ); + tokio::time::sleep(Duration::from_millis(500)).await; + assert!(cancelled_finished.exists()); + assert!(cancelled.headers().await.is_ok()); + assert_eq!( + std::fs::read_to_string(&cancelled_invocations).expect("cancelled invocation count"), + "x" + ); + std::fs::remove_file(&cancelled_finished).expect("reset helper completion marker"); + let headers = cancelled.headers().await.expect("cached headers"); + assert!( + tokio::time::timeout( + Duration::from_millis(/*millis*/ 20), + cancelled.refresh(headers.refresh_epoch) + ) + .await + .is_err() + ); + tokio::time::sleep(Duration::from_millis(/*millis*/ 500)).await; + assert!(cancelled_finished.exists()); + assert_eq!( + cancelled.refresh(headers.refresh_epoch).await.unwrap(), + headers.values + ); + assert_eq!( + std::fs::read_to_string(cancelled_invocations).expect("refresh invocation count"), + "xx" + ); + + let dropped_started = cwd.join("dropped-helper-started"); + let dropped_finished = cwd.join("dropped-helper-finished"); + let dropped = HttpHeadersProvider::new( + "https://example.com", + &format!( + "printf x > '{}'; sleep 1; printf x > '{}'", + dropped_started.display(), + dropped_finished.display(), + ), + cwd, + ) + .expect("dropped provider"); + assert!( + tokio::time::timeout(Duration::from_millis(500), dropped.headers()) + .await + .is_err() + ); + assert!(dropped_started.exists()); + drop(dropped); + tokio::time::sleep(Duration::from_millis(1_100)).await; + assert!(!dropped_finished.exists()); +} + +#[tokio::test] +async fn nonzero_helper_exit_is_cached() { + let temp = tempfile::tempdir().expect("temporary helper directory"); + let failed_invocations = temp.path().join("failed-invocations"); + let command = if cfg!(windows) { + format!( + r#"echo x>>"{0}" & echo {{"X-Gateway":"valid"}} & exit /b 23"#, + failed_invocations.display() + ) + } else { + format!( + "echo x >> '{0}'; printf '{{\"X-Gateway\":\"valid\"}}'; exit 23", + failed_invocations.display() + ) + }; + let failed = HttpHeadersProvider::new( + "https://example.com/mcp", + &command, + temp.path().to_path_buf(), + ) + .expect("failed provider"); + let first = failed.headers().await.err().expect("failed helper"); + let second = failed.headers().await.err().expect("cached failure"); + assert_eq!(first.to_string(), second.to_string()); + let invocations = std::fs::read_to_string(failed_invocations).expect("failed invocation count"); + assert_eq!(invocations.lines().count(), 1); +} + +#[cfg(unix)] +#[tokio::test] +async fn failed_refresh_is_shared_without_replacing_credentials() { + let temp = tempfile::tempdir().expect("temporary helper directory"); + let invocations = temp.path().join("invocations"); + let command = format!( + "printf x >> '{0}'; count=$(wc -c < '{0}'); \ + if [ \"$count\" -gt 1 ]; then exit 23; fi; \ + printf '{{\"X-Gateway\":\"token\"}}'", + invocations.display(), + ); + let provider = HttpHeadersProvider::new( + "https://example.com/mcp", + &command, + temp.path().to_path_buf(), + ) + .expect("headers provider"); + let original = provider.headers().await.expect("initial helper headers"); + assert!(provider.refresh(original.refresh_epoch).await.is_err()); + let stale = provider.refresh(original.refresh_epoch).await.unwrap(); + assert_eq!(stale, original.values); + let next = provider.headers().await.expect("preserved credentials"); + assert_eq!(next.values, original.values); + assert!(provider.refresh(next.refresh_epoch).await.is_err()); + assert_eq!(std::fs::read_to_string(invocations).unwrap(), "xxx"); +} + +#[cfg(unix)] +#[tokio::test] +async fn connection_headers_are_cached_and_origin_bound() { + use axum::Router; + use axum::http::StatusCode; + use axum::response::Redirect; + use axum::routing::get; + use axum::routing::post; + use codex_exec_server::RouteAwareHttpClient; + use codex_http_client::HttpClientFactory; + use codex_http_client::OutboundProxyPolicy; + use std::sync::Arc; + use tempfile::tempdir; + use tokio::net::TcpListener; + + async fn handle(headers: axum::http::HeaderMap) -> StatusCode { + assert_eq!( + headers.get("proxy-authorization"), + Some(&HeaderValue::from_static("Bearer token")) + ); + assert_eq!( + headers.get("x-label"), + Some(&HeaderValue::from_bytes("café".as_bytes()).unwrap()) + ); + StatusCode::NO_CONTENT + } + let temp = tempdir().expect("temporary helper directory"); + let invocation_file = temp.path().join("invocations"); + let cross_listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind cross-origin server"); + let cross_url = format!("http://{}/start", cross_listener.local_addr().unwrap()); + tokio::spawn(async move { + axum::serve( + cross_listener, + Router::new() + .route("/start", get(|| async { Redirect::temporary("/final") })) + .route("/final", get(|| async { StatusCode::NO_CONTENT })), + ) + .await + .expect("serve cross-origin requests"); + }); + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind test server"); + let url = format!("http://{}/mcp", listener.local_addr().unwrap()); + let redirect_url = cross_url.clone(); + let app = Router::new().route("/mcp", post(handle)).route( + "/redirect", + get(move || std::future::ready(Redirect::temporary(&redirect_url))), + ); + tokio::spawn(async move { + axum::serve(listener, app) + .await + .expect("serve test requests"); + }); + let command = format!( + "printf x >> '{}'; printf '{{\"Proxy-Authorization\":\"Bearer token\",\"X-Label\":\"café\"}}'", + invocation_file.display(), + ); + let inner: Arc = Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))); + let client = with_http_headers_helper(inner, &url, &command, temp.path().to_path_buf()) + .expect("headers helper client"); + let request = |session: &str| HttpRequestParams { + method: "POST".to_string(), + url: url.clone(), + headers: Vec::new(), + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: session.to_string(), + stream_response: true, + }; + let mut cross_request = request("cross-origin"); + cross_request.method = "GET".to_string(); + cross_request.url = cross_url; + assert_eq!( + client.http_request(cross_request).await.unwrap().status, + 204 + ); + assert!(!invocation_file.exists()); + let mut redirect_request = request("redirect"); + redirect_request.method = "GET".to_string(); + redirect_request.url = url.replace("/mcp", "/redirect"); + assert_eq!( + client.http_request(redirect_request).await.unwrap().status, + 307 + ); + let (left, right) = tokio::join!( + client.http_request_stream(request("session-a")), + client.http_request_stream(request("session-b")) + ); + assert_eq!(left.expect("left request").0.status, 204); + assert_eq!(right.expect("right request").0.status, 204); + assert_eq!( + std::fs::read_to_string(&invocation_file).expect("helper invocation count"), + "x" + ); +} + +#[cfg(unix)] +#[tokio::test] +async fn concurrent_rejected_posts_share_one_headers_refresh() { + use axum::Router; + use axum::extract::State; + use axum::http::StatusCode; + use axum::routing::post; + use codex_exec_server::RouteAwareHttpClient; + use codex_http_client::HttpClientFactory; + use codex_http_client::OutboundProxyPolicy; + use tempfile::tempdir; + use tokio::net::TcpListener; + use tokio::sync::Barrier; + + for (rejection, credential_header) in [ + (StatusCode::UNAUTHORIZED, "Proxy-Authorization"), + (StatusCode::FORBIDDEN, "x-litellm-api-key"), + (StatusCode::UNAUTHORIZED, "Authorization"), + ] { + let temp = tempdir().expect("temporary helper directory"); + let invocations = temp.path().join("invocations"); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("http://{}/mcp", listener.local_addr().unwrap()); + let barrier = Arc::new(Barrier::new(/*n*/ 2)); + let app = + Router::new() + .route( + "/mcp", + post( + move |State(barrier): State>, + headers: HeaderMap, + body: String| async move { + assert_eq!(body, "original request body"); + assert_eq!( + headers.get("mcp-session-id"), + Some(&HeaderValue::from_static("existing-session")) + ); + match headers[credential_header].to_str().unwrap() { + "Bearer token-1" => { + barrier.wait().await; + rejection + } + "Bearer token-2" => { + assert_eq!(headers["x-helper-only"], "configured"); + StatusCode::NO_CONTENT + } + unexpected => panic!("unexpected test credential: {unexpected}"), + } + }, + ), + ) + .with_state(barrier); + tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + let command = format!( + "printf x >> '{0}'; count=$(wc -c < '{0}'); count=$((count)); \ + if [ \"$count\" -eq 1 ]; then \ + printf '{{\"{credential_header}\":\"Bearer token-1\",\"X-Helper-Only\":\"stale\"}}'; \ + else printf '{{\"{credential_header}\":\"Bearer token-%s\"}}' \"$count\"; fi", + invocations.display(), + ); + let inner: Arc = Arc::new(RouteAwareHttpClient::new( + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )); + let client = with_http_headers_helper(inner, &url, &command, temp.path().to_path_buf()) + .expect("headers helper client"); + let request = |request_id: &str| HttpRequestParams { + method: "POST".to_string(), + url: url.clone(), + headers: vec![ + HttpHeader { + name: "X-Helper-Only".to_string(), + value: "configured".to_string(), + value_env_var: None, + }, + HttpHeader { + name: "mcp-session-id".to_string(), + value: "existing-session".to_string(), + value_env_var: None, + }, + ], + body: Some(b"original request body".to_vec().into()), + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: request_id.to_string(), + stream_response: true, + }; + + let (left, right) = tokio::join!( + client.http_request(request("left")), + client.http_request_stream(request("right")), + ); + assert_eq!(left.expect("left request").status, 204); + assert_eq!(right.expect("right request").0.status, 204); + assert_eq!( + std::fs::read_to_string(invocations).expect("helper invocation count"), + "xx" + ); + } +} + +#[cfg(unix)] +#[tokio::test] +async fn helper_refresh_preserves_oauth_challenges_and_retries_at_most_once() { + use codex_exec_server::RouteAwareHttpClient; + use codex_http_client::HttpClientFactory; + use codex_http_client::OutboundProxyPolicy; + use tempfile::tempdir; + use wiremock::Mock; + use wiremock::MockServer; + use wiremock::ResponseTemplate; + use wiremock::matchers::header; + use wiremock::matchers::method; + + let unchanged = "printf '{\"Proxy-Authorization\":\"Bearer token\"}'"; + let changed = "printf '{\"Proxy-Authorization\":\"Bearer token-%s\"}' \"$count\""; + let failed = "if [ \"$count\" -gt 1 ]; then exit 23; fi; printf '{\"Proxy-Authorization\":\"Bearer token\"}'"; + let ignored_authorization = "if [ \"$count\" -eq 1 ]; then \ + printf '{\"Authorization\":\"old\",\"X-A\":\"a\",\"X-B\":\"b\"}'; \ + else printf '{\"X-B\":\"b\",\"X-A\":\"a\",\"Authorization\":\"new\"}'; fi"; + for streamed in [false, true] { + for (helper_output, status, expected_mcp_requests) in [ + (unchanged, 401, 1), + (failed, 401, 1), + (changed, 401, 2), + (changed, 403, 1), + (ignored_authorization, 401, 1), + ] { + let (challenge, expected_helper_invocations) = if status == 403 { + ( + r#"Bearer error="insufficient_scope", scope="tools:write""#, + "x", + ) + } else { + (r#"Bearer realm="oauth""#, "xx") + }; + let temp = tempdir().expect("temporary helper directory"); + let invocations = temp.path().join("invocations"); + let server = MockServer::start().await; + let url = format!("{}/mcp", server.uri()); + Mock::given(method("POST")) + .and(header("authorization", "Bearer oauth-token")) + .respond_with( + ResponseTemplate::new(status) + .insert_header("www-authenticate", challenge) + .set_body_string("original OAuth challenge"), + ) + .expect(expected_mcp_requests) + .mount(&server) + .await; + let command = format!( + "printf x >> '{0}'; count=$(wc -c < '{0}'); {helper_output}", + invocations.display(), + ); + let inner: Arc = Arc::new(RouteAwareHttpClient::new( + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )); + let client = with_http_headers_helper(inner, &url, &command, temp.path().to_path_buf()) + .expect("headers helper client"); + let params = HttpRequestParams { + method: "POST".to_string(), + url: url.clone(), + headers: vec![HttpHeader { + name: "aUtHoRiZaTiOn".to_string(), + value: "Bearer oauth-token".to_string(), + value_env_var: None, + }], + body: None, + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Follow, + request_id: "oauth-challenge".to_string(), + stream_response: streamed, + }; + let (response, body) = if streamed { + let (response, mut body) = client + .http_request_stream(params) + .await + .expect("original OAuth response"); + let mut bytes = Vec::new(); + while let Some(chunk) = body.recv().await.expect("response body chunk") { + bytes.extend(chunk); + } + (response, bytes) + } else { + let response = client + .http_request(params) + .await + .expect("original OAuth response"); + let bytes = response.body.0.clone(); + (response, bytes) + }; + + assert_eq!(response.status, status); + assert!(response.headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("www-authenticate") && header.value == challenge + })); + assert_eq!(body, b"original OAuth challenge"); + assert_eq!( + std::fs::read_to_string(invocations).expect("helper invocation count"), + expected_helper_invocations + ); + } + } +} + +#[cfg(unix)] +#[tokio::test] +async fn refresh_retry_rechecks_deadline_and_redirects() { + use codex_exec_server::RouteAwareHttpClient; + use codex_http_client::HttpClientFactory; + use codex_http_client::OutboundProxyPolicy; + + let temp = tempfile::tempdir().expect("helper directory"); + let command = "if [ -e invoked ]; then printf '{\"Proxy-Authorization\":\"Bearer fresh\"}'; else touch invoked; printf '{}'; fi"; + let client = HttpHeadersClient { + inner: Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + provider: HttpHeadersProvider::new( + "http://example.com/mcp", + command, + temp.path().to_path_buf(), + ) + .expect("provider"), + }; + let requests = Mutex::new(Vec::new()); + let error = client + .request( + HttpRequestParams { + method: "POST".to_string(), + url: "http://example.com/mcp".to_string(), + headers: Vec::new(), + body: Some(b"original request".to_vec().into()), + timeout_ms: Some(5_000), + redirect_policy: HttpRedirectPolicy::Stop, + request_id: "mcp-request".to_string(), + stream_response: false, + }, + |params| { + let mut requests = requests.lock().expect("requests lock"); + let status = if requests.is_empty() { 401 } else { 307 }; + requests.push(params); + async move { + tokio::time::sleep(Duration::from_millis(/*millis*/ 20)).await; + Ok(HttpRequestResponse { + status, + headers: vec![HttpHeader { + name: "Location".to_string(), + value: "http://example.com/redirected".to_string(), + value_env_var: None, + }], + body: Vec::new().into(), + }) + } + .boxed() + }, + |response| response, + ) + .await + .expect_err("refreshed proxy redirect must fail"); + assert!( + error + .to_string() + .contains("cannot safely replay Proxy-Authorization") + ); + let requests = requests.into_inner().expect("recorded requests"); + let [original, retry] = requests.as_slice() else { + panic!("expected exactly two requests"); + }; + assert!(retry.timeout_ms < original.timeout_ms); + let mut expected_retry = original.clone(); + expected_retry.timeout_ms = retry.timeout_ms; + expected_retry.headers = vec![HttpHeader { + name: "proxy-authorization".to_string(), + value: "Bearer fresh".to_string(), + value_env_var: None, + }]; + assert_eq!(retry, &expected_retry); +} diff --git a/codex-rs/rmcp-client/src/in_process_transport.rs b/codex-rs/rmcp-client/src/in_process_transport.rs new file mode 100644 index 0000000000000000000000000000000000000000..f78d4ce0b528c12c16a68e3ac5cf335fa1ab72f7 --- /dev/null +++ b/codex-rs/rmcp-client/src/in_process_transport.rs @@ -0,0 +1,14 @@ +use std::io; + +use futures::future::BoxFuture; +use tokio::io::DuplexStream; + +/// Recreates a fresh in-process MCP byte stream whenever the client needs one. +/// +/// Implementations are expected to start the paired server side before +/// returning the client stream. The factory is retained by [`crate::RmcpClient`] +/// so reconnects can rebuild the transport without knowing which built-in +/// server produced it. +pub trait InProcessTransportFactory: Send + Sync { + fn open(&self) -> BoxFuture<'static, io::Result>; +} diff --git a/codex-rs/rmcp-client/src/lib.rs b/codex-rs/rmcp-client/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..51d94292c8f255e69466c9a752c965e766971956 --- /dev/null +++ b/codex-rs/rmcp-client/src/lib.rs @@ -0,0 +1,95 @@ +mod auth_status; +mod bounded_stdio_transport; +mod elicitation_client_service; +mod ema_auth_policy; +mod ema_claims; +mod ema_exchange; +mod ema_identity; +mod enterprise_oauth_login; +mod event_notification_transport; +mod executor_process_transport; +mod http_client_adapter; +mod http_client_redirect; +mod http_headers; +mod in_process_transport; +mod local_child; +mod local_stdio_transport; +mod logging_client_handler; +mod oauth; +mod oauth_callback; +mod oauth_client_registration; +mod oauth_http_client; +mod oauth_refresh_mode; +mod perform_oauth_login; +mod program_resolver; +mod protocol_mode; +mod rmcp_client; +mod service_error; +mod startup_error; +mod stdio_server_launcher; +mod tool_input; +mod user_verification; +mod utils; +mod www_authenticate; + +pub use auth_status::McpAuthState; +pub use auth_status::McpLoginRequirement; +pub use auth_status::OAuthDiscoveryTimeout; +pub use auth_status::StreamableHttpOAuthDiscovery; +pub use auth_status::determine_streamable_http_auth_status; +pub use auth_status::determine_streamable_http_auth_status_from_credentials; +pub use auth_status::discover_streamable_http_oauth; +pub use codex_protocol::protocol::McpAuthStatus; +pub use ema_auth_policy::EmaAuthFailure; +pub use ema_auth_policy::EmaInvalidGrantSource; +pub use ema_auth_policy::validate_ema_auth_resource; +pub use ema_claims::validate_oidc_identity_assertion; +pub use ema_exchange::EmaAccessToken; +pub use ema_exchange::EmaIdJagExchangeRequest; +pub use ema_exchange::exchange_id_jag; +pub use ema_identity::EmaIdpIdentity; +pub use ema_identity::EmaIdpIdentityRequest; +pub use ema_identity::resolve_ema_idp_identity; +pub use ema_identity::stored_ema_identity_is_usable; +pub use enterprise_oauth_login::EnterpriseOAuthCredentialGuard; +pub use enterprise_oauth_login::EnterpriseOAuthCredentials; +pub use enterprise_oauth_login::EnterpriseOAuthLoginHandle; +pub use enterprise_oauth_login::EnterpriseOAuthLoginRequest; +pub use enterprise_oauth_login::delete_enterprise_oauth_tokens; +pub use enterprise_oauth_login::perform_enterprise_oauth_login_return_url; +pub use event_notification_transport::EventNotificationReceiver; +pub use http_client_adapter::StreamableHttpRedirectMode; +pub use http_headers::with_http_headers_helper; +pub use in_process_transport::InProcessTransportFactory; +pub use oauth::StoredOAuthCredentialSnapshot; +pub use oauth::StoredOAuthTokens; +pub use oauth::WrappedOAuthTokenResponse; +pub use oauth::delete_oauth_tokens; +pub use oauth::save_oauth_tokens; +pub use oauth::stored_oauth_credential_snapshot; +pub use oauth::stored_oauth_credentials; +pub use oauth_callback::McpOAuthCallbackMode; +pub use oauth_callback::resolve_mcp_oauth_callback_url; +pub use oauth_client_registration::McpOAuthClientRegistration; +pub use oauth_refresh_mode::McpOAuthRefreshMode; +pub use perform_oauth_login::OAuthProviderError; +pub use perform_oauth_login::OauthLoginHandle; +pub use perform_oauth_login::perform_oauth_login; +pub use perform_oauth_login::perform_oauth_login_return_url; +pub use perform_oauth_login::perform_oauth_login_silent; +pub use perform_oauth_login::perform_oauth_login_with_callback_input; +pub use protocol_mode::McpProtocolMode; +pub use rmcp::model::ElicitationAction; +pub use rmcp_client::CancellableEventStreamRequest; +pub use rmcp_client::Elicitation; +pub use rmcp_client::ElicitationResponse; +pub use rmcp_client::ListToolsWithConnectorIdResult; +pub use rmcp_client::RmcpClient; +pub use rmcp_client::SendElicitation; +pub use rmcp_client::StreamableHttpBearerToken; +pub use rmcp_client::ToolWithConnectorId; +pub use service_error::mcp_error; +pub use startup_error::is_authentication_required_error; +pub use stdio_server_launcher::ExecutorStdioServerLauncher; +pub use stdio_server_launcher::LocalStdioServerLauncher; +pub use stdio_server_launcher::StdioServerLauncher; diff --git a/codex-rs/rmcp-client/src/local_child.rs b/codex-rs/rmcp-client/src/local_child.rs new file mode 100644 index 0000000000000000000000000000000000000000..fdb98303e5e89d7266dbd4f681ca250c780ea184 --- /dev/null +++ b/codex-rs/rmcp-client/src/local_child.rs @@ -0,0 +1,34 @@ +//! Uniform child-process API for local MCP servers, with platform-specific spawning. +//! +//! Commands must use the launcher's cleared environment, process group, and default +//! argv[0]. Both implementations expose Tokio stdio handles and kill on drop. + +use std::io; +use std::process::Stdio; + +use tokio::process::Command; + +#[cfg(target_os = "macos")] +#[path = "macos_stdio.rs"] +mod macos; + +#[cfg(target_os = "macos")] +pub(super) use macos::LocalChild; +#[cfg(not(target_os = "macos"))] +pub(super) use tokio::process::Child as LocalChild; + +pub(super) fn spawn(mut command: Command) -> io::Result { + command + .kill_on_drop(true) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()); + #[cfg(target_os = "macos")] + { + LocalChild::spawn(command) + } + #[cfg(not(target_os = "macos"))] + { + command.spawn() + } +} diff --git a/codex-rs/rmcp-client/src/local_stdio_transport.rs b/codex-rs/rmcp-client/src/local_stdio_transport.rs new file mode 100644 index 0000000000000000000000000000000000000000..b94c42584be88b6d288edf8b22170b3757eac44b --- /dev/null +++ b/codex-rs/rmcp-client/src/local_stdio_transport.rs @@ -0,0 +1,97 @@ +//! Protocol framing and child lifetime for locally spawned MCP servers. +//! +//! Process creation is platform-specific; protocol selection and shutdown are shared. + +use std::future::Future; +use std::io; +use std::time::Duration; + +use futures::FutureExt; +use rmcp::service::RoleClient; +use rmcp::service::RxJsonRpcMessage; +use rmcp::service::TxJsonRpcMessage; +use rmcp::transport::Transport; +use rmcp::transport::async_rw::AsyncRwTransport; +use tokio::process::ChildStderr; +use tokio::process::ChildStdin; +use tokio::process::ChildStdout; +use tokio::process::Command; + +use crate::bounded_stdio_transport::BoundedStdioTransport; +use crate::local_child; +use crate::local_child::LocalChild; +use crate::protocol_mode::McpProtocolMode; + +pub(super) struct LocalStdioTransport { + child: LocalChild, + transport: StdioTransport, +} + +enum StdioTransport { + /// Preserve rmcp's existing framing for servers using the initialize handshake. + Legacy(AsyncRwTransport), + /// Bound frames and skip messages unknown to the client during 2026-07-28 discovery. + V20260728(BoundedStdioTransport), +} + +impl LocalStdioTransport { + pub(super) fn spawn( + command: Command, + program_name: String, + protocol_mode: McpProtocolMode, + ) -> io::Result<(Self, Option)> { + let mut child = local_child::spawn(command)?; + let stdin = child + .stdin + .take() + .ok_or_else(|| io::Error::other("MCP server stdin was not piped"))?; + let stdout = child + .stdout + .take() + .ok_or_else(|| io::Error::other("MCP server stdout was not piped"))?; + let stderr = child.stderr.take(); + let transport = match protocol_mode { + McpProtocolMode::Legacy => StdioTransport::Legacy(AsyncRwTransport::new(stdout, stdin)), + McpProtocolMode::V20260728 => { + StdioTransport::V20260728(BoundedStdioTransport::new(stdin, stdout, program_name)) + } + }; + Ok((Self { child, transport }, stderr)) + } + + pub(super) fn id(&self) -> Option { + self.child.id() + } +} + +impl Transport for LocalStdioTransport { + type Error = io::Error; + + fn send( + &mut self, + item: TxJsonRpcMessage, + ) -> impl Future> + Send + 'static { + match &mut self.transport { + StdioTransport::Legacy(transport) => transport.send(item).boxed(), + StdioTransport::V20260728(transport) => transport.send(item).boxed(), + } + } + + fn receive(&mut self) -> impl Future>> + Send { + match &mut self.transport { + StdioTransport::Legacy(transport) => transport.receive().boxed(), + StdioTransport::V20260728(transport) => transport.receive().boxed(), + } + } + + async fn close(&mut self) -> io::Result<()> { + match &mut self.transport { + StdioTransport::Legacy(transport) => transport.close().await?, + StdioTransport::V20260728(transport) => transport.close().await?, + } + match tokio::time::timeout(Duration::from_secs(3), self.child.wait()).await { + Ok(status) => status.map(|_| ()), + Err(_) => self.child.kill().await, + } + } +} diff --git a/codex-rs/rmcp-client/src/logging_client_handler.rs b/codex-rs/rmcp-client/src/logging_client_handler.rs new file mode 100644 index 0000000000000000000000000000000000000000..98887d749011085803c6539e611429bc0e578c0f --- /dev/null +++ b/codex-rs/rmcp-client/src/logging_client_handler.rs @@ -0,0 +1,141 @@ +use std::sync::Arc; + +use rmcp::ClientHandler; +use rmcp::RoleClient; +use rmcp::model::CancelledNotificationParam; +use rmcp::model::ClientInfo; +use rmcp::model::ElicitRequestParams; +use rmcp::model::ElicitResult; +#[allow(deprecated)] +use rmcp::model::LoggingLevel; +#[allow(deprecated)] +use rmcp::model::LoggingMessageNotificationParam; +use rmcp::model::ProgressNotificationParam; +use rmcp::model::ResourceUpdatedNotificationParam; +use rmcp::service::NotificationContext; +use rmcp::service::RequestContext; +use tracing::debug; +use tracing::error; +use tracing::info; +use tracing::warn; + +use crate::rmcp_client::Elicitation; +use crate::rmcp_client::SendElicitation; + +#[derive(Clone)] +pub(crate) struct LoggingClientHandler { + client_info: ClientInfo, + send_elicitation: Arc, +} + +impl LoggingClientHandler { + pub(crate) fn new(client_info: ClientInfo, send_elicitation: SendElicitation) -> Self { + Self { + client_info, + send_elicitation: Arc::new(send_elicitation), + } + } +} + +impl ClientHandler for LoggingClientHandler { + async fn create_elicitation( + &self, + request: ElicitRequestParams, + context: RequestContext, + ) -> Result { + (self.send_elicitation)(context.id, Elicitation::Mcp(request)) + .await + .map(Into::into) + .map_err(|err| rmcp::ErrorData::internal_error(err.to_string(), None)) + } + + async fn on_cancelled( + &self, + params: CancelledNotificationParam, + _context: NotificationContext, + ) { + info!( + "MCP server cancelled request (request_id: {:?}, reason: {:?})", + params.request_id, params.reason + ); + } + + async fn on_progress( + &self, + params: ProgressNotificationParam, + _context: NotificationContext, + ) { + info!( + "MCP server progress notification (token: {:?}, progress: {}, total: {:?}, message: {:?})", + params.progress_token, params.progress, params.total, params.message + ); + } + + async fn on_resource_updated( + &self, + params: ResourceUpdatedNotificationParam, + _context: NotificationContext, + ) { + info!("MCP server resource updated (uri: {})", params.uri); + } + + async fn on_resource_list_changed(&self, _context: NotificationContext) { + info!("MCP server resource list changed"); + } + + async fn on_tool_list_changed(&self, _context: NotificationContext) { + info!("MCP server tool list changed"); + } + + async fn on_prompt_list_changed(&self, _context: NotificationContext) { + info!("MCP server prompt list changed"); + } + + fn get_info(&self) -> ClientInfo { + self.client_info.clone() + } + + #[allow(deprecated)] + async fn on_logging_message( + &self, + params: LoggingMessageNotificationParam, + _context: NotificationContext, + ) { + let LoggingMessageNotificationParam { + level, + logger, + data, + .. + } = params; + let logger = logger.as_deref(); + match level { + LoggingLevel::Emergency + | LoggingLevel::Alert + | LoggingLevel::Critical + | LoggingLevel::Error => { + error!( + "MCP server log message (level: {:?}, logger: {:?}, data: {})", + level, logger, data + ); + } + LoggingLevel::Warning => { + warn!( + "MCP server log message (level: {:?}, logger: {:?}, data: {})", + level, logger, data + ); + } + LoggingLevel::Notice | LoggingLevel::Info => { + info!( + "MCP server log message (level: {:?}, logger: {:?}, data: {})", + level, logger, data + ); + } + LoggingLevel::Debug => { + debug!( + "MCP server log message (level: {:?}, logger: {:?}, data: {})", + level, logger, data + ); + } + } + } +} diff --git a/codex-rs/rmcp-client/src/macos_stdio.rs b/codex-rs/rmcp-client/src/macos_stdio.rs new file mode 100644 index 0000000000000000000000000000000000000000..af471320f36a561bb247f11fd835cba06be6dce4 --- /dev/null +++ b/codex-rs/rmcp-client/src/macos_stdio.rs @@ -0,0 +1,415 @@ +//! Native spawning for macOS MCP executables, without rewriting script paths. +//! +//! Rust falls back to fork for a historical relative-path/cwd bug in Apple's +//! `posix_spawnp`. Calling `posix_spawn` directly avoids that wrapper. This module only +//! accepts the launcher's cleared-environment command shape, with piped stdio, +//! a new process group, and default `argv[0]`. Bare commands search the child's +//! PATH. Unsuccessful searches and executable files without shebangs retain the +//! existing launcher. Each native child owns its PID until it has been reaped. + +use std::ffi::CString; +use std::ffi::OsStr; +use std::io; +use std::os::fd::AsRawFd; +use std::os::fd::FromRawFd; +use std::os::fd::OwnedFd; +use std::os::unix::ffi::OsStrExt; +use std::os::unix::process::ExitStatusExt; +use std::path::Path; +use std::process::ExitStatus; +use std::ptr; + +use tokio::process::Child; +use tokio::process::ChildStderr; +use tokio::process::ChildStdin; +use tokio::process::ChildStdout; +use tokio::process::Command; +use tokio::signal::unix::Signal; +use tokio::signal::unix::SignalKind; +use tokio::signal::unix::signal; + +/// Matches Tokio's child API while keeping native spawning private to macOS. +pub(crate) struct LocalChild { + inner: ChildKind, + pub(crate) stdin: Option, + pub(crate) stdout: Option, + pub(crate) stderr: Option, +} + +enum ChildKind { + Tokio(Child), + Native(NativeChild), +} + +impl LocalChild { + /// Uses native spawning for relative paths and bare names, retaining Tokio's + /// fallback for unsuccessful PATH searches and executable text without a shebang. + pub(super) fn spawn(mut command: Command) -> io::Result { + let program = command.as_std().get_program(); + if Path::new(program).is_relative() + && !program.is_empty() + && let Some((child, stdin, stdout, stderr)) = NativeChild::spawn(command.as_std())? + { + return Ok(Self { + inner: ChildKind::Native(child), + stdin: Some(stdin), + stdout: Some(stdout), + stderr: Some(stderr), + }); + } + let mut child = command.spawn()?; + Ok(Self { + stdin: child.stdin.take(), + stdout: child.stdout.take(), + stderr: child.stderr.take(), + inner: ChildKind::Tokio(child), + }) + } + + pub(crate) fn id(&self) -> Option { + match &self.inner { + ChildKind::Tokio(child) => child.id(), + ChildKind::Native(child) => child.id(), + } + } + + pub(crate) async fn wait(&mut self) -> io::Result { + self.stdin.take(); + match &mut self.inner { + ChildKind::Tokio(child) => child.wait().await, + ChildKind::Native(child) => child.wait().await, + } + } + + pub(crate) async fn kill(&mut self) -> io::Result<()> { + self.stdin.take(); + match &mut self.inner { + ChildKind::Tokio(child) => child.kill().await, + ChildKind::Native(child) => child.kill().await, + } + } +} + +// libc does not expose this Apple extension. It is available since macOS 10.15, +// before Codex's minimum supported macOS version (12). +unsafe extern "C" { + fn posix_spawn_file_actions_addchdir_np( + actions: *mut libc::posix_spawn_file_actions_t, + path: *const libc::c_char, + ) -> libc::c_int; +} + +/// Owns a child PID until reaping, so cancellation cannot lose or reuse it. +/// Dropping a live child kills it and reaps it independently of the Tokio runtime. +struct NativeChild { + pid: Option, + status: Option, + sigchld: Signal, +} + +impl NativeChild { + /// Spawns the MCP command without changing its executable path or `argv[0]`. + /// The caller must clear inherited environment variables before setting the + /// child's environment, because only explicit command entries are copied. + fn spawn( + command: &std::process::Command, + ) -> io::Result> { + let program = c_string(command.get_program())?; + let search_path = !program.as_bytes().contains(&b'/'); + let args = std::iter::once(command.get_program()) + .chain(command.get_args()) + .map(c_string) + .collect::>>()?; + let argv = args + .iter() + .map(|arg| arg.as_ptr().cast_mut()) + .chain(std::iter::once(ptr::null_mut())) + .collect::>(); + let env = command + .get_envs() + .filter_map(|(key, value)| value.map(|value| (key, value))) + .map(|(key, value)| { + let mut entry = key.to_os_string(); + entry.push("="); + entry.push(value); + c_string(&entry) + }) + .collect::>>()?; + let envp = env + .iter() + .map(|entry| entry.as_ptr().cast_mut()) + .chain(std::iter::once(ptr::null_mut())) + .collect::>(); + let cwd = command + .get_current_dir() + .map(|cwd| c_string(cwd.as_os_str())) + .transpose()?; + + // Subscribe before spawning so a child that exits immediately cannot be missed. + let sigchld = signal(SignalKind::child())?; + let (stdin_read, stdin_write) = io::pipe()?; + let (stdout_read, stdout_write) = io::pipe()?; + let (stderr_read, stderr_write) = io::pipe()?; + let child_fds = [ + child_fd(stdin_read.into())?, + child_fd(stdout_write.into())?, + child_fd(stderr_write.into())?, + ]; + let stdin = ChildStdin::from_std(OwnedFd::from(stdin_write).into())?; + let stdout = ChildStdout::from_std(OwnedFd::from(stdout_read).into())?; + let stderr = ChildStderr::from_std(OwnedFd::from(stderr_read).into())?; + + let mut actions = FileActions(ptr::null_mut()); + let mut attrs = Attributes(ptr::null_mut()); + let mut pid = 0; + // SAFETY: All C strings and pipe descriptors outlive this synchronous + // spawn. The initialized action/attribute objects are destroyed by RAII. + let result = unsafe { + cvt(libc::posix_spawn_file_actions_init(&mut actions.0))?; + cvt(libc::posix_spawnattr_init(&mut attrs.0))?; + if let Some(cwd) = &cwd { + cvt(posix_spawn_file_actions_addchdir_np( + &mut actions.0, + cwd.as_ptr(), + ))?; + } + for (target, source) in child_fds.iter().enumerate() { + cvt(libc::posix_spawn_file_actions_adddup2( + &mut actions.0, + source.as_raw_fd(), + target as i32, + ))?; + } + cvt(libc::posix_spawnattr_setpgroup( + &mut attrs.0, + /*pgroup*/ 0, + ))?; + let mut defaults = 0; + cvt_errno(libc::sigemptyset(&mut defaults))?; + cvt_errno(libc::sigaddset(&mut defaults, libc::SIGPIPE))?; + cvt(libc::posix_spawnattr_setsigdefault(&mut attrs.0, &defaults))?; + // Match Command's descriptor inheritance: honor FD_CLOEXEC rather + // than introducing a different policy with CLOEXEC_DEFAULT. + cvt(libc::posix_spawnattr_setflags( + &mut attrs.0, + (libc::POSIX_SPAWN_SETPGROUP | libc::POSIX_SPAWN_SETSIGDEF) as _, + ))?; + let mut spawn = |executable: &CString| { + libc::posix_spawn( + &mut pid, + executable.as_ptr(), + &actions.0, + &attrs.0, + argv.as_ptr(), + envp.as_ptr(), + ) + }; + if !search_path { + spawn(&program) + } else { + // posix_spawnp searches the parent's PATH, not envp. Search the + // child's PATH ourselves, using Apple's default when it is unset. + let path = command + .get_envs() + .find(|(key, _)| *key == "PATH") + .and_then(|(_, value)| value) + .unwrap_or(OsStr::new("/usr/bin:/bin")); + let mut result = libc::ENOENT; + for directory in std::env::split_paths(path) { + let mut executable = directory.into_os_string(); + if executable.is_empty() { + executable.push("."); + } + // Preserve the spelling execvp would pass to a shebang + // interpreter, including empty entries and trailing slashes. + executable.push("/"); + executable.push(command.get_program()); + if executable.as_bytes().len() >= libc::PATH_MAX as usize { + return Ok(None); + } + result = spawn(&c_string(&executable)?); + if !matches!( + result, + libc::ENOENT + | libc::ENOTDIR + | libc::EACCES + | libc::ELOOP + | libc::ENAMETOOLONG + ) { + break; + } + } + result + } + }; + // Retain Command's shell fallback and exact PATH search errors. + if result == libc::ENOEXEC || (search_path && result != 0) { + return Ok(None); + } + cvt(result)?; + let child = Self { + pid: Some(pid), + status: None, + sigchld, + }; + Ok(Some((child, stdin, stdout, stderr))) + } + + fn id(&self) -> Option { + self.pid.map(|pid| pid as u32) + } + + /// Polls and caches the exit status, relinquishing the PID once it is reaped. + /// `ECHILD` also relinquishes it to prevent later signaling of a reused PID. + fn try_wait(&mut self) -> io::Result> { + if self.status.is_some() { + return Ok(self.status); + } + let pid = self + .pid + .ok_or_else(|| io::Error::from_raw_os_error(libc::ECHILD))?; + let mut status = 0; + // SAFETY: We own this child PID and provide writable status storage. + match unsafe { libc::waitpid(pid, &mut status, libc::WNOHANG) } { + 0 => Ok(None), + -1 => { + let error = io::Error::last_os_error(); + if error.raw_os_error() == Some(libc::ECHILD) { + self.pid = None; + } + Err(error) + } + _ => { + self.pid = None; + self.status = Some(ExitStatus::from_raw(status)); + Ok(self.status) + } + } + } + + /// Waits without transferring child ownership into the future, so callers + /// may cancel a wait and then wait again or kill the same child. + async fn wait(&mut self) -> io::Result { + loop { + match self.try_wait() { + Ok(Some(status)) => return Ok(status), + Ok(None) => {} + Err(error) if error.kind() == io::ErrorKind::Interrupted => continue, + Err(error) => return Err(error), + } + self.sigchld + .recv() + .await + .ok_or_else(|| io::Error::other("SIGCHLD stream closed"))?; + } + } + + /// Sends SIGKILL if still owned, then waits for the child to be reaped. + async fn kill(&mut self) -> io::Result<()> { + if let Some(pid) = self.pid { + // SAFETY: An unreaped child retains its PID, even after it exits. + let result = unsafe { libc::kill(pid, libc::SIGKILL) }; + if result == -1 && io::Error::last_os_error().raw_os_error() != Some(libc::ESRCH) { + return Err(io::Error::last_os_error()); + } + } + self.wait().await.map(|_| ()) + } +} + +impl Drop for NativeChild { + fn drop(&mut self) { + let _ = self.try_wait(); + let Some(pid) = self.pid.take() else { return }; + // SAFETY: This child has not been reaped, so its PID cannot be reused. + unsafe { + libc::kill(pid, libc::SIGKILL); + } + // Drop may run during runtime shutdown. Reap independently of Tokio, + // without blocking its worker threads while the killed process exits. + if std::thread::Builder::new() + .name("mcp-child-reaper".into()) + .spawn(move || reap(pid)) + .is_err() + { + // Resource exhaustion must not turn a dropped child into a zombie. + reap(pid); + } + } +} + +/// Reaps the child transferred by `Drop`, retrying interrupted waits without +/// requiring a live Tokio runtime. The caller relinquishes ownership of `pid`. +fn reap(pid: libc::pid_t) { + loop { + // SAFETY: The caller transferred exclusive ownership of this unreaped PID. + let result = unsafe { + libc::waitpid(pid, ptr::null_mut(), /*options*/ 0) + }; + if result != -1 || io::Error::last_os_error().kind() != io::ErrorKind::Interrupted { + break; + } + } +} + +fn c_string(value: &OsStr) -> io::Result { + CString::new(value.as_bytes()) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "nul byte in MCP command")) +} + +/// Keeps a child pipe source above stdio so `dup2` cannot clobber another source +/// when the parent has closed standard descriptors. Close-on-exec disposes of +/// this extra descriptor after the spawn actions duplicate it onto stdio. +fn child_fd(fd: OwnedFd) -> io::Result { + // SAFETY: fcntl duplicates this live descriptor; the returned fd is newly owned. + let duplicate = unsafe { libc::fcntl(fd.as_raw_fd(), libc::F_DUPFD_CLOEXEC, 3) }; + cvt_errno(duplicate)?; + Ok(unsafe { OwnedFd::from_raw_fd(duplicate) }) +} + +/// Converts a spawn API's returned error number, which does not use `errno`. +fn cvt(result: libc::c_int) -> io::Result<()> { + if result == 0 { + Ok(()) + } else { + Err(io::Error::from_raw_os_error(result)) + } +} + +/// Converts a syscall's `-1` sentinel using the thread's current `errno`. +fn cvt_errno(result: libc::c_int) -> io::Result<()> { + if result == -1 { + Err(io::Error::last_os_error()) + } else { + Ok(()) + } +} + +struct FileActions(libc::posix_spawn_file_actions_t); +struct Attributes(libc::posix_spawnattr_t); + +impl Drop for FileActions { + fn drop(&mut self) { + if !self.0.is_null() { + // SAFETY: This object was initialized by posix_spawn_file_actions_init. + unsafe { + libc::posix_spawn_file_actions_destroy(&mut self.0); + } + } + } +} + +impl Drop for Attributes { + fn drop(&mut self) { + if !self.0.is_null() { + // SAFETY: This object was initialized by posix_spawnattr_init. + unsafe { + libc::posix_spawnattr_destroy(&mut self.0); + } + } + } +} + +#[cfg(test)] +#[path = "macos_stdio_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/macos_stdio_tests.rs b/codex-rs/rmcp-client/src/macos_stdio_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4bf9ed063450f7d8233c3eacb4846fc5155b732a --- /dev/null +++ b/codex-rs/rmcp-client/src/macos_stdio_tests.rs @@ -0,0 +1,288 @@ +//! Regression coverage for native spawn compatibility, descriptor inheritance, and reaping. + +use super::*; +use pretty_assertions::assert_eq; +use std::fs; +use std::os::unix::ffi::OsStringExt; +use std::os::unix::fs::PermissionsExt; +use std::os::unix::fs::symlink; +use std::time::Duration; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; + +async fn native_output(command: Command) -> anyhow::Result { + let mut child = crate::local_child::spawn(command)?; + assert!(matches!(child.inner, ChildKind::Native(_))); + drop(child.stdin.take()); + let mut stdout = child.stdout.take().expect("piped stdout"); + let mut stderr = child.stderr.take().expect("piped stderr"); + let mut output = Vec::new(); + let mut diagnostic = Vec::new(); + let (_, _, status) = tokio::try_join!( + stdout.read_to_end(&mut output), + stderr.read_to_end(&mut diagnostic), + child.wait() + )?; + Ok(std::process::Output { + status, + stdout: output, + stderr: diagnostic, + }) +} + +#[tokio::test] +async fn bare_script_search_matches_child_path_and_preserves_script_spelling() -> anyhow::Result<()> +{ + let root = tempfile::tempdir()?; + let program = std::ffi::OsString::from("server=é"); + let bin = root.path().join("bin"); + let blocked = root.path().join("blocked"); + fs::create_dir(&bin)?; + fs::create_dir(&blocked)?; + fs::write(blocked.join(&program), "not executable")?; + let script = bin.join(&program); + fs::write( + &script, + "#!/bin/sh\nprintf '%s\\n' \"$0\" \"$1\" \"$MCP_TEST\"; /bin/pwd; printf diagnostic >&2; exit 23\n", + )?; + fs::set_permissions(&script, fs::Permissions::from_mode(/*mode*/ 0o755))?; + symlink(&script, root.path().join(&program))?; + for path in [ + bin.into_os_string(), + "bin/".into(), + "missing:blocked:bin".into(), + "missing:".into(), + "".into(), + ] { + let mut command = Command::new(&program); + command + .current_dir(root.path()) + .env_clear() + .env("PATH", path) + .env("MCP_TEST", "kept") + .arg("spaces ; literal $arg"); + let expected = command.output().await?; + assert_eq!(native_output(command).await?, expected); + } + Ok(()) +} + +#[tokio::test] +async fn bare_executable_uses_default_path_and_preserves_argv0() -> anyhow::Result<()> { + let mut command = Command::new("sh"); + command.env_clear().args(["-c", "printf '%s' \"$0\""]); + let expected = command.output().await?; + assert_eq!(native_output(command).await?, expected); + Ok(()) +} + +#[tokio::test] +async fn relative_script_preserves_paths_stdio_environment_and_process_group() -> anyhow::Result<()> +{ + let root = tempfile::tempdir()?; + fs::create_dir_all(root.path().join("actual/bin"))?; + symlink("actual/bin", root.path().join("link"))?; + let script = root.path().join("actual/server"); + fs::write( + &script, + "#!/bin/sh\nread -r input\nprintf '%s\\n' \"$0\" \"$1\" \"$2\" \"$MCP_TEST\" \"$input\"\nprintf diagnostic >&2\nexit 23\n", + )?; + fs::set_permissions(script, fs::Permissions::from_mode(0o755))?; + let mut command = std::process::Command::new("./link/../server"); + command + .current_dir(root.path()) + .env_clear() + .env("MCP_TEST", "kept") + .arg("spaces ; literal $arg") + .arg(std::ffi::OsString::from_vec(b"raw-\xff".to_vec())); + let (mut child, mut stdin, mut stdout, mut stderr) = + NativeChild::spawn(&command)?.expect("native child"); + let pid = child.id().expect("live PID") as libc::pid_t; + // SAFETY: getpgid only inspects the live child, which waits for input below. + assert_eq!(unsafe { libc::getpgid(pid) }, pid); + stdin.write_all(b"hello\n").await?; + drop(stdin); + let mut output = Vec::new(); + let mut diagnostic = String::new(); + let (_, _, status) = tokio::try_join!( + stdout.read_to_end(&mut output), + stderr.read_to_string(&mut diagnostic), + child.wait() + )?; + assert_eq!( + (output.as_slice(), diagnostic.as_str(), status.code()), + ( + b"./link/../server\nspaces ; literal $arg\nraw-\xff\nkept\nhello\n".as_slice(), + "diagnostic", + Some(23) + ) + ); + assert_eq!(child.wait().await?, status); + assert_eq!(child.id(), None); + Ok(()) +} + +#[tokio::test] +async fn native_executable_preserves_argv0() -> anyhow::Result<()> { + let root = tempfile::tempdir()?; + symlink("/bin/sh", root.path().join("shell"))?; + let program = Path::new("./shell"); + let mut command = std::process::Command::new(program); + command + .current_dir(root.path()) + .env_clear() + .args(["-c", "printf '%s' \"$0\""]); + let (mut child, stdin, mut stdout, _stderr) = + NativeChild::spawn(&command)?.expect("native child"); + drop(stdin); + let mut output = Vec::new(); + stdout.read_to_end(&mut output).await?; + assert!(child.wait().await?.success()); + assert_eq!(output, program.as_os_str().as_bytes()); + Ok(()) +} + +#[tokio::test] +async fn cancelled_wait_can_still_kill_and_reap_child() -> anyhow::Result<()> { + let root = tempfile::tempdir()?; + symlink("/bin/cat", root.path().join("server"))?; + let mut command = std::process::Command::new("./server"); + command.current_dir(root.path()).env_clear(); + let (mut child, _stdin, _stdout, _stderr) = + NativeChild::spawn(&command)?.expect("native child"); + assert!( + tokio::time::timeout(Duration::from_millis(20), child.wait()) + .await + .is_err() + ); + child.kill().await?; + assert_eq!(child.wait().await?.signal(), Some(libc::SIGKILL)); + assert_eq!(child.id(), None); + Ok(()) +} + +#[tokio::test] +async fn descriptor_inheritance_matches_command() -> anyhow::Result<()> { + let root = tempfile::tempdir()?; + symlink("/bin/sh", root.path().join("shell"))?; + let file = fs::File::open("/dev/null")?; + for (operation, expected) in [ + (libc::F_DUPFD, "inherited"), + (libc::F_DUPFD_CLOEXEC, "closed"), + ] { + // SAFETY: Duplicate a harmless descriptor with the requested inheritance flag. + let fd = unsafe { libc::fcntl(file.as_raw_fd(), operation, 200) }; + cvt_errno(fd)?; + // SAFETY: fcntl returned a new owned descriptor. + let _sentinel = unsafe { OwnedFd::from_raw_fd(fd) }; + let mut command = std::process::Command::new("./shell"); + command + .current_dir(root.path()) + .env_clear() + .env("SENTINEL", fd.to_string()) + .args([ + "-c", + "if [ -e /dev/fd/\"$SENTINEL\" ]; then printf inherited; else printf closed; fi", + ]); + let (mut child, stdin, mut stdout, _stderr) = + NativeChild::spawn(&command)?.expect("native child"); + drop(stdin); + let mut output = String::new(); + stdout.read_to_string(&mut output).await?; + assert!(child.wait().await?.success()); + assert_eq!(output, expected); + assert_eq!(command.output()?.stdout, output.as_bytes()); + } + Ok(()) +} + +#[tokio::test] +async fn launch_failures_preserve_os_errors() -> anyhow::Result<()> { + let root = tempfile::tempdir()?; + fs::write(root.path().join("not-executable"), "#!/bin/sh\nexit 0\n")?; + for (program, cwd, errno) in [ + ("./missing", root.path().to_path_buf(), libc::ENOENT), + ("./not-executable", root.path().to_path_buf(), libc::EACCES), + ("missing", root.path().to_path_buf(), libc::ENOENT), + ("not-executable", root.path().to_path_buf(), libc::EACCES), + ("", root.path().to_path_buf(), libc::ENOENT), + ("/bin/sh", root.path().join("missing"), libc::ENOENT), + ] { + let mut command = Command::new(program); + command + .env_clear() + .env("PATH", root.path()) + .current_dir(cwd); + let error = crate::local_child::spawn(command) + .err() + .expect("spawn should fail"); + assert_eq!(error.raw_os_error(), Some(errno)); + } + Ok(()) +} + +#[tokio::test] +async fn executable_text_without_shebang_retains_command_fallback() -> anyhow::Result<()> { + let root = tempfile::tempdir()?; + let script = root.path().join("server"); + fs::write(&script, "printf '%s' \"$0\"\n")?; + fs::set_permissions(script, fs::Permissions::from_mode(0o755))?; + for program in ["./server", "server"] { + let mut command = Command::new(program); + command + .current_dir(root.path()) + .env_clear() + .env("PATH", "."); + let mut child = crate::local_child::spawn(command)?; + let mut output = Vec::new(); + child + .stdout + .take() + .expect("piped stdout") + .read_to_end(&mut output) + .await?; + assert!(child.wait().await?.success()); + assert_eq!(output, b"./server"); + } + Ok(()) +} + +#[test] +fn dropping_after_runtime_shutdown_kills_and_reaps_child() -> anyhow::Result<()> { + let root = tempfile::tempdir()?; + symlink("/bin/cat", root.path().join("server"))?; + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build()?; + let (child, stdin, _stdout, _stderr) = runtime + .block_on(async { + let mut command = std::process::Command::new("./server"); + command.current_dir(root.path()).env_clear(); + NativeChild::spawn(&command) + })? + .expect("native child"); + let pid = child.id().expect("live PID") as libc::pid_t; + drop(runtime); + drop(child); + drop(stdin); + let deadline = std::time::Instant::now() + Duration::from_secs(5); + loop { + // SAFETY: Signal zero probes existence without changing the process. + if unsafe { libc::kill(pid, 0) } == -1 { + assert_eq!(io::Error::last_os_error().raw_os_error(), Some(libc::ESRCH)); + break; + } + anyhow::ensure!(std::time::Instant::now() < deadline, "child was not reaped"); + std::thread::sleep(Duration::from_millis(10)); + } + // SAFETY: WNOHANG verifies our drop reaper already collected this child. + assert_eq!( + unsafe { libc::waitpid(pid, ptr::null_mut(), libc::WNOHANG) }, + -1 + ); + assert_eq!( + io::Error::last_os_error().raw_os_error(), + Some(libc::ECHILD) + ); + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/oauth.rs b/codex-rs/rmcp-client/src/oauth.rs new file mode 100644 index 0000000000000000000000000000000000000000..86822efa1964ccbe88b64c758f34b7e77611c6fd --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth.rs @@ -0,0 +1,1834 @@ +//! This file handles all logic related to managing MCP OAuth credentials. +//! All credentials are stored using the keyring crate which uses os-specific keyring services. +//! https://crates.io/crates/keyring +//! macOS: macOS keychain. +//! Windows: Windows Credential Manager +//! Linux: DBus-based Secret Service, the kernel keyutils, and a combo of the two +//! FreeBSD, OpenBSD: DBus-based Secret Service +//! +//! For Linux, we use linux-native-async-persistent which uses both keyutils and async-secret-service (see below) for storage. +//! See the docs for the keyutils_persistent module for a full explanation of why both are used. Because this store uses the +//! async-secret-service, you must specify the additional features required by that store +//! +//! async-secret-service provides access to the DBus-based Secret Service storage on Linux, FreeBSD, and OpenBSD. This is an asynchronous +//! keystore that always encrypts secrets when they are transferred across the bus. If DBus isn't installed the keystore will fall back to the json +//! file because we don't use the "vendored" feature. +//! +//! If the keyring is not available or fails, we fall back to CODEX_HOME/.credentials.json which is consistent with other coding CLI agents. + +mod credential_store; +mod ema_identity; +mod enterprise_generation; +mod issuer_binding; +mod refresh_lock; +mod refresh_transaction; +mod resolved_store; +mod runtime; +mod store_lock; + +#[cfg(test)] +#[path = "oauth/test_support.rs"] +pub(crate) mod test_support; + +use anyhow::Context; +use anyhow::Error; +use anyhow::Result; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_secrets::LocalSecretsNamespace; +use codex_secrets::SecretName; +use codex_secrets::SecretScope; +use codex_secrets::SecretsBackendKind; +use codex_secrets::SecretsManager; +use oauth2::AccessToken; +use oauth2::RefreshToken; +use oauth2::Scope; +use oauth2::TokenResponse; +use oauth2::basic::BasicTokenType; +use rmcp::transport::auth::OAuthTokenResponse; +use rmcp::transport::auth::VendorExtraTokenFields; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Value; +use serde_json::map::Map as JsonMap; +use sha2::Digest; +use sha2::Sha256; +use std::collections::BTreeMap; +use std::fs; +use std::io::ErrorKind; +use std::io::Write; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::Duration; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; +use tracing::warn; + +use self::store_lock::OAuthStore; +use self::store_lock::OAuthStoreLock; +use self::store_lock::OAuthStoreLockFailure; + +use codex_keyring_store::DefaultKeyringStore; +use codex_keyring_store::KeyringStore; +use rmcp::transport::auth::AuthorizationManager; +use tokio::sync::Mutex; + +use codex_utils_home_dir::find_codex_home; + +pub(crate) use self::credential_store::OAuthCredentialStore; +pub(crate) use self::ema_identity::stored_oidc_identity; +pub(crate) use self::enterprise_generation::EnterpriseOAuthGeneration; +pub(crate) use self::enterprise_generation::EnterpriseOAuthGenerationFile; +pub(crate) use self::issuer_binding::validate_authorization_server_endpoints; +pub(crate) use self::issuer_binding::validate_refresh_token_issuer; +pub(crate) use self::refresh_lock::RefreshCredentialLock; +pub(crate) use self::refresh_transaction::install_tokens_in_manager; +pub(crate) use self::resolved_store::ResolvedOAuthCredentialStore; +pub(crate) use self::resolved_store::ResolvedOAuthTokens; +pub(crate) use self::resolved_store::resolve_oauth_tokens_from_store_policy; +use self::resolved_store::try_resolve_oauth_tokens_from_store_policy; +pub(crate) use self::runtime::OAuthRuntime; + +const KEYRING_SERVICE: &str = "Codex MCP Credentials"; +const MCP_OAUTH_SECRET_PREFIX: &str = "MCP_OAUTH"; +const REFRESH_SKEW_MILLIS: u64 = 30_000; + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct StoredOAuthTokens { + pub server_name: String, + pub url: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub issuer: Option, + pub client_id: String, + pub token_response: WrappedOAuthTokenResponse, + #[serde(default)] + pub expires_at: Option, +} + +impl StoredOAuthTokens { + pub(crate) fn has_refresh_token(&self) -> bool { + self.token_response + .0 + .refresh_token() + .is_some_and(|refresh_token| !refresh_token.secret().trim().is_empty()) + } + + pub(crate) fn bound_issuer(&self) -> Option<&str> { + self.issuer + .as_deref() + .filter(|issuer| !issuer.trim().is_empty()) + } + + pub(crate) fn access_token_is_usable_without_refresh(&self) -> bool { + !token_needs_refresh(self.expires_at) + && !self + .token_response + .0 + .access_token() + .secret() + .trim() + .is_empty() + } +} + +/// OAuth credentials paired with the concrete store selected for their client lifecycle. +#[derive(Debug, Clone)] +pub struct StoredOAuthCredentialSnapshot { + credentials: StoredOAuthTokens, + store: ResolvedOAuthCredentialStore, + store_was_contended: bool, +} + +impl PartialEq for StoredOAuthCredentialSnapshot { + fn eq(&self, other: &Self) -> bool { + self.credentials == other.credentials && self.store == other.store + } +} + +impl StoredOAuthCredentialSnapshot { + pub(crate) fn new( + mut credentials: StoredOAuthTokens, + store: ResolvedOAuthCredentialStore, + ) -> Self { + credentials.token_response.0.set_expires_in(None); + Self { + credentials, + store, + store_was_contended: false, + } + } + + /// Returns the normalized credentials originally read from the selected store. + pub fn credentials(&self) -> &StoredOAuthTokens { + &self.credentials + } + + /// Returns whether this snapshot was retained because its store could not be read. + pub fn store_was_contended(&self) -> bool { + self.store_was_contended + } + + /// Refreshes a runtime snapshot without waiting or discarding its last known authority. + pub fn for_runtime_refresh( + previous: Option<&Self>, + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + ) -> Result> { + match try_resolve_oauth_tokens_from_store_policy( + &DefaultKeyringStore, + server_name, + url, + store_mode, + keyring_backend_kind, + ) { + Ok(Some(mut resolved)) => { + resolved.tokens.token_response.0.set_expires_in(None); + Ok(Some(Self { + credentials: resolved.tokens, + store: resolved.store, + store_was_contended: false, + })) + } + Ok(None) => Ok(None), + Err(error) if oauth_store_is_contended(&error) => Ok(previous + .filter(|previous| { + previous.credentials.server_name == server_name + && previous.credentials.url == url + }) + .map(|previous| Self { + store_was_contended: true, + ..previous.clone() + })), + Err(error) => Err(error), + } + } + + /// Rereads the selected authority without waiting for a contended credential-store lock. + pub fn reload( + &self, + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + ) -> Result> { + if self.store == ResolvedOAuthCredentialStore::File + && store_mode == OAuthCredentialsStoreMode::Auto + { + return Self::for_runtime_refresh( + /*previous*/ None, + server_name, + url, + store_mode, + keyring_backend_kind, + ) + .map(|snapshot| snapshot.map(|snapshot| snapshot.credentials)); + } + + let credentials = match self.store.try_load(&DefaultKeyringStore, server_name, url) { + Ok(credentials) => credentials, + Err(error) if oauth_store_is_contended(&error) => return Ok(None), + Err(error) => return Err(error), + }; + Ok(normalized_oauth_credentials(credentials.as_ref())) + } +} + +/// Wrap OAuthTokenResponse to allow for partial equality comparison. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct WrappedOAuthTokenResponse(pub OAuthTokenResponse); + +impl PartialEq for WrappedOAuthTokenResponse { + fn eq(&self, other: &Self) -> bool { + match (serde_json::to_value(self), serde_json::to_value(other)) { + (Ok(s1), Ok(s2)) => s1 == s2, + _ => false, + } + } +} + +#[derive(Debug, PartialEq, Eq)] +pub(crate) enum StoredOAuthTokenStatus { + Missing, + Usable, + AuthorizationRequired, +} + +pub(crate) fn oauth_token_status( + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result { + let resolved = resolve_oauth_tokens_from_store_policy( + &DefaultKeyringStore, + server_name, + url, + store_mode, + keyring_backend_kind, + )?; + Ok(match resolved.as_ref().map(|resolved| &resolved.tokens) { + None => StoredOAuthTokenStatus::Missing, + Some(tokens) if oauth_tokens_are_usable(tokens) => StoredOAuthTokenStatus::Usable, + Some(_) => StoredOAuthTokenStatus::AuthorizationRequired, + }) +} + +/// Returns stored OAuth credentials without their derived expiration interval. +pub fn stored_oauth_credentials( + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result> { + Ok( + stored_oauth_credential_snapshot(server_name, url, store_mode, keyring_backend_kind)? + .map(|snapshot| snapshot.credentials), + ) +} + +/// Loads OAuth credentials together with the concrete authority selected by store policy. +pub fn stored_oauth_credential_snapshot( + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result> { + let Some(resolved) = resolve_oauth_tokens_from_store_policy( + &DefaultKeyringStore, + server_name, + url, + store_mode, + keyring_backend_kind, + )? + else { + return Ok(None); + }; + Ok(Some(StoredOAuthCredentialSnapshot::new( + resolved.tokens, + resolved.store, + ))) +} + +fn oauth_store_is_contended(error: &Error) -> bool { + matches!( + error.downcast_ref::(), + Some(OAuthStoreLockFailure::Timeout { acquire_timeout, .. }) + if acquire_timeout.is_zero() + ) +} + +fn normalized_oauth_credentials(tokens: Option<&StoredOAuthTokens>) -> Option { + tokens.map(|tokens| { + let mut tokens = tokens.clone(); + tokens.token_response.0.set_expires_in(None); + tokens + }) +} + +fn oauth_tokens_are_usable(tokens: &StoredOAuthTokens) -> bool { + if tokens.client_id.trim().is_empty() { + return false; + } + + if token_needs_refresh(tokens.expires_at) { + return tokens.bound_issuer().is_some() && tokens.has_refresh_token(); + } + + tokens.access_token_is_usable_without_refresh() +} + +fn refresh_expires_in_from_timestamp(tokens: &mut StoredOAuthTokens) { + let Some(expires_at) = tokens.expires_at else { + return; + }; + + match expires_in_from_timestamp(expires_at) { + Some(seconds) => { + let duration = Duration::from_secs(seconds); + tokens.token_response.0.set_expires_in(Some(&duration)); + } + None => { + // RMCP treats a missing expiry as unknown and uses the access token + // as-is. Treat a known-expired timestamp as an explicit zero so + // startup refreshes the token before the first request. + tokens + .token_response + .0 + .set_expires_in(Some(&Duration::ZERO)); + } + } +} + +fn load_oauth_tokens_from_keyring( + keyring_store: &K, + keyring_backend_kind: AuthKeyringBackendKind, + server_name: &str, + url: &str, +) -> std::result::Result, OAuthKeyringLoadError> { + match keyring_backend_kind { + AuthKeyringBackendKind::Direct => { + load_oauth_tokens_from_direct_keyring(keyring_store, server_name, url) + .map_err(OAuthKeyringLoadError::Backend) + } + AuthKeyringBackendKind::Secrets => { + load_oauth_tokens_from_secrets_keyring(keyring_store, server_name, url) + } + } +} + +fn load_oauth_tokens_from_direct_keyring( + keyring_store: &K, + server_name: &str, + url: &str, +) -> Result> { + let key = compute_store_key(server_name, url)?; + match keyring_store.load(KEYRING_SERVICE, &key) { + Ok(Some(serialized)) => { + let mut tokens: StoredOAuthTokens = serde_json::from_str(&serialized) + .context("failed to deserialize OAuth tokens from keyring")?; + refresh_expires_in_from_timestamp(&mut tokens); + Ok(Some(tokens)) + } + Ok(None) => Ok(None), + Err(error) => Err(Error::new(error.into_error())), + } +} + +fn load_oauth_tokens_from_secrets_keyring( + keyring_store: &K, + server_name: &str, + url: &str, +) -> std::result::Result, OAuthKeyringLoadError> { + let _store_lock = OAuthStoreLock::acquire_for_read(OAuthStore::Secrets)?; + load_oauth_tokens_from_secrets_keyring_with_lock_held(keyring_store, server_name, url) +} + +fn load_oauth_tokens_from_secrets_keyring_with_lock_held( + keyring_store: &K, + server_name: &str, + url: &str, +) -> std::result::Result, OAuthKeyringLoadError> { + let codex_home = find_codex_home().map_err(anyhow::Error::from)?; + let manager = SecretsManager::new_with_keyring_store_and_namespace( + codex_home.to_path_buf(), + SecretsBackendKind::Local, + Arc::new(keyring_store.clone()), + LocalSecretsNamespace::McpOAuth, + ); + let secret_name = compute_secret_name(server_name, url)?; + match manager + .get(&SecretScope::Global, &secret_name) + .context("failed to load MCP OAuth tokens from encrypted storage")? + { + Some(serialized) => { + let mut tokens: StoredOAuthTokens = serde_json::from_str(&serialized) + .context("failed to deserialize OAuth tokens from encrypted storage")?; + refresh_expires_in_from_timestamp(&mut tokens); + Ok(Some(tokens)) + } + None => Ok(None), + } +} + +/// Classifies keyring load failures that affect Auto fallback policy. +#[derive(Debug, thiserror::Error)] +enum OAuthKeyringLoadError { + /// Store coordination failed, so consulting another authority would be unsafe. + #[error(transparent)] + StoreLock(#[from] OAuthStoreLockFailure), + /// The selected keyring backend itself was unavailable or its data was invalid. + #[error(transparent)] + Backend(#[from] anyhow::Error), +} + +pub async fn save_oauth_tokens( + server_name: &str, + tokens: &StoredOAuthTokens, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result<()> { + let lock = RefreshCredentialLock::acquire_for_server(server_name, &tokens.url).await?; + save_oauth_tokens_with_lock_held(&lock, server_name, tokens, store_mode, keyring_backend_kind) +} + +/// Save while retaining the matching credential lock acquired by the caller. +pub(crate) fn save_oauth_tokens_with_lock_held( + _lock: &RefreshCredentialLock, + server_name: &str, + tokens: &StoredOAuthTokens, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result<()> { + let keyring_store = DefaultKeyringStore; + match store_mode { + OAuthCredentialsStoreMode::Auto => save_oauth_tokens_with_keyring_with_fallback_to_file( + &keyring_store, + keyring_backend_kind, + server_name, + tokens, + ), + OAuthCredentialsStoreMode::File => save_oauth_tokens_to_file(tokens), + OAuthCredentialsStoreMode::Keyring => save_oauth_tokens_with_keyring_and_cleanup_file( + &keyring_store, + keyring_backend_kind, + server_name, + tokens, + ), + } +} + +fn save_oauth_tokens_with_keyring( + keyring_store: &K, + keyring_backend_kind: AuthKeyringBackendKind, + server_name: &str, + tokens: &StoredOAuthTokens, +) -> Result<()> { + // This exact-store writer is used after a client resolves its authority. Only login-time + // policy resolution may clean up or update the non-selected store. + match keyring_backend_kind { + AuthKeyringBackendKind::Direct => { + save_oauth_tokens_to_direct_keyring(keyring_store, server_name, tokens) + } + AuthKeyringBackendKind::Secrets => { + save_oauth_tokens_to_secrets_keyring(keyring_store, server_name, tokens) + } + } +} + +fn save_oauth_tokens_to_direct_keyring( + keyring_store: &K, + server_name: &str, + tokens: &StoredOAuthTokens, +) -> Result<()> { + let serialized = serde_json::to_string(tokens).context("failed to serialize OAuth tokens")?; + + let key = compute_store_key(server_name, &tokens.url)?; + match keyring_store.save(KEYRING_SERVICE, &key, &serialized) { + Ok(()) => Ok(()), + Err(error) => { + let message = format!( + "failed to write OAuth tokens to keyring: {}", + error.message() + ); + warn!("{message}"); + Err(Error::new(error.into_error()).context(message)) + } + } +} + +/// Saves one credential while holding the Secrets aggregate-store lock across the mutation. +fn save_oauth_tokens_to_secrets_keyring( + keyring_store: &K, + server_name: &str, + tokens: &StoredOAuthTokens, +) -> Result<()> { + let serialized = serde_json::to_string(tokens).context("failed to serialize OAuth tokens")?; + let _store_lock = OAuthStoreLock::acquire_for_write(OAuthStore::Secrets)?; + save_oauth_tokens_to_secrets_keyring_with_lock_held( + keyring_store, + server_name, + tokens, + &serialized, + ) +} + +/// Writes one credential to Secrets. The caller must hold the Secrets aggregate-store lock. +fn save_oauth_tokens_to_secrets_keyring_with_lock_held( + keyring_store: &K, + server_name: &str, + tokens: &StoredOAuthTokens, + serialized: &str, +) -> Result<()> { + let codex_home = find_codex_home()?; + let manager = SecretsManager::new_with_keyring_store_and_namespace( + codex_home.to_path_buf(), + SecretsBackendKind::Local, + Arc::new(keyring_store.clone()), + LocalSecretsNamespace::McpOAuth, + ); + let secret_name = compute_secret_name(server_name, &tokens.url)?; + manager + .set(&SecretScope::Global, &secret_name, serialized) + .context("failed to write OAuth tokens to encrypted storage") +} + +/// Saves to the selected keyring backend, then best-effort removes the fallback File entry. +fn save_oauth_tokens_with_keyring_and_cleanup_file( + keyring_store: &K, + keyring_backend_kind: AuthKeyringBackendKind, + server_name: &str, + tokens: &StoredOAuthTokens, +) -> Result<()> { + save_oauth_tokens_with_keyring(keyring_store, keyring_backend_kind, server_name, tokens)?; + let key = compute_store_key(server_name, &tokens.url)?; + if let Err(error) = delete_oauth_tokens_from_file(&key) { + warn!( + server_name, + keyring_backend = ?keyring_backend_kind, + error = %error, + "failed to remove OAuth tokens from fallback storage" + ); + } + Ok(()) +} + +fn save_oauth_tokens_with_keyring_with_fallback_to_file( + keyring_store: &K, + keyring_backend_kind: AuthKeyringBackendKind, + server_name: &str, + tokens: &StoredOAuthTokens, +) -> Result<()> { + match save_oauth_tokens_with_keyring_and_cleanup_file( + keyring_store, + keyring_backend_kind, + server_name, + tokens, + ) { + Ok(()) => Ok(()), + // As on load, a store lock failure is a coordination failure rather than evidence that + // the keyring backend is unavailable. Falling back could leave a newer File token hidden + // behind a stale Secrets entry. + Err(error) if error.downcast_ref::().is_some() => Err(error), + Err(error) => { + let message = error.to_string(); + warn!("falling back to file storage for OAuth tokens: {message}"); + save_oauth_tokens_to_file(tokens) + .with_context(|| format!("failed to write OAuth tokens to keyring: {message}")) + } + } +} + +pub async fn delete_oauth_tokens( + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result { + let lock = RefreshCredentialLock::acquire_for_server(server_name, url).await?; + delete_oauth_tokens_with_lock_held(&lock, server_name, url, store_mode, keyring_backend_kind) +} + +/// Delete while retaining the matching credential lock acquired by the caller. +pub(crate) fn delete_oauth_tokens_with_lock_held( + _lock: &RefreshCredentialLock, + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result { + let keyring_store = DefaultKeyringStore; + delete_oauth_tokens_from_keyring_and_file( + &keyring_store, + store_mode, + keyring_backend_kind, + server_name, + url, + ) +} + +fn delete_oauth_tokens_from_keyring_and_file( + keyring_store: &K, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + server_name: &str, + url: &str, +) -> Result { + let key = compute_store_key(server_name, url)?; + let keyring_result = + delete_oauth_tokens_from_keyring(keyring_store, keyring_backend_kind, server_name, url); + let keyring_removed = match keyring_result { + Ok(removed) => removed, + Err(error) => { + let message = error.to_string(); + warn!("failed to delete OAuth tokens from keyring: {message}"); + match store_mode { + OAuthCredentialsStoreMode::Auto | OAuthCredentialsStoreMode::Keyring => { + return Err(error).context("failed to delete OAuth tokens from keyring"); + } + OAuthCredentialsStoreMode::File => false, + } + } + }; + + let file_removed = delete_oauth_tokens_from_file(&key)?; + Ok(keyring_removed || file_removed) +} + +fn delete_oauth_tokens_from_keyring( + keyring_store: &K, + keyring_backend_kind: AuthKeyringBackendKind, + server_name: &str, + url: &str, +) -> Result { + match keyring_backend_kind { + AuthKeyringBackendKind::Direct => { + delete_oauth_tokens_from_direct_keyring(keyring_store, server_name, url) + } + AuthKeyringBackendKind::Secrets => { + let direct_removed = + delete_oauth_tokens_from_direct_keyring(keyring_store, server_name, url)?; + let secrets_removed = + delete_oauth_tokens_from_secrets_keyring(keyring_store, server_name, url)?; + Ok(direct_removed || secrets_removed) + } + } +} + +fn delete_oauth_tokens_from_direct_keyring( + keyring_store: &K, + server_name: &str, + url: &str, +) -> Result { + let key = compute_store_key(server_name, url)?; + keyring_store + .delete(KEYRING_SERVICE, &key) + .map_err(|error| Error::new(error.into_error())) +} + +fn delete_oauth_tokens_from_secrets_keyring( + keyring_store: &K, + server_name: &str, + url: &str, +) -> Result { + let _store_lock = OAuthStoreLock::acquire_for_write(OAuthStore::Secrets)?; + let codex_home = find_codex_home()?; + let manager = SecretsManager::new_with_keyring_store_and_namespace( + codex_home.to_path_buf(), + SecretsBackendKind::Local, + Arc::new(keyring_store.clone()), + LocalSecretsNamespace::McpOAuth, + ); + let secret_name = compute_secret_name(server_name, url)?; + let secrets_removed = manager + .delete(&SecretScope::Global, &secret_name) + .context("failed to delete OAuth tokens from encrypted storage")?; + Ok(secrets_removed) +} + +#[derive(Clone)] +pub(crate) struct OAuthPersistor { + inner: Arc, +} + +struct OAuthPersistorInner { + server_name: String, + url: String, + authorization_manager: Arc>, + credential_store: ResolvedOAuthCredentialStore, + last_credentials: Mutex>, +} + +impl OAuthPersistor { + pub(crate) fn new( + server_name: String, + url: String, + authorization_manager: Arc>, + credential_store: ResolvedOAuthCredentialStore, + initial_credentials: Option, + ) -> Self { + Self { + inner: Arc::new(OAuthPersistorInner { + server_name, + url, + authorization_manager, + credential_store, + last_credentials: Mutex::new(initial_credentials), + }), + } + } + + pub(crate) async fn stored_credentials(&self) -> Option { + let credentials = self.inner.last_credentials.lock().await; + normalized_oauth_credentials(credentials.as_ref()) + } + + /// Persists RMCP-managed credential changes back to this client's resolved authority. + #[expect( + clippy::await_holding_invalid_type, + reason = "AuthorizationManager async access must be serialized through its mutex" + )] + pub(crate) async fn persist_if_needed(&self) -> Result<()> { + let (client_id, maybe_credentials) = { + let manager = self.inner.authorization_manager.clone(); + let guard = manager.lock().await; + guard.get_credentials().await + }?; + + match maybe_credentials { + Some(credentials) => { + let mut last_credentials = self.inner.last_credentials.lock().await; + let new_token_response = WrappedOAuthTokenResponse(credentials.clone()); + let same_token = last_credentials + .as_ref() + .map(|previous| previous.token_response == new_token_response) + .unwrap_or(false); + let expires_at = if same_token { + last_credentials + .as_ref() + .and_then(|previous| previous.expires_at) + } else { + compute_expires_at_millis(&credentials) + }; + let stored = StoredOAuthTokens { + server_name: self.inner.server_name.clone(), + url: self.inner.url.clone(), + issuer: last_credentials + .as_ref() + .and_then(|previous| previous.issuer.clone()), + client_id, + token_response: new_token_response, + expires_at, + }; + if last_credentials.as_ref() != Some(&stored) { + self.inner.credential_store.save( + &DefaultKeyringStore, + &self.inner.server_name, + &stored, + )?; + *last_credentials = Some(stored); + } + } + None => { + let mut last_credentials = self.inner.last_credentials.lock().await; + if last_credentials.take().is_some() + && let Err(error) = self.inner.credential_store.delete( + &DefaultKeyringStore, + &self.inner.server_name, + &self.inner.url, + ) + { + warn!( + server_name = %self.inner.server_name, + error = %error, + "failed to remove MCP OAuth credentials from the resolved store" + ); + } + } + } + + Ok(()) + } +} + +const FALLBACK_FILENAME: &str = ".credentials.json"; +const MCP_SERVER_TYPE: &str = "http"; + +type FallbackFile = BTreeMap; + +#[derive(Debug, Clone, Serialize, Deserialize)] +struct FallbackTokenEntry { + server_name: String, + server_url: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + issuer: Option, + client_id: String, + access_token: String, + #[serde(default)] + expires_at: Option, + #[serde(default)] + refresh_token: Option, + #[serde(default)] + scopes: Vec, + // Legacy host entries omit this marker, so executor lookups fail closed. + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + executor_owned: bool, +} + +fn load_oauth_tokens_from_file(server_name: &str, url: &str) -> Result> { + let _store_lock = OAuthStoreLock::acquire_for_read(OAuthStore::File)?; + load_oauth_tokens_from_file_with_lock_held(server_name, url) +} + +fn load_oauth_tokens_from_file_with_lock_held( + server_name: &str, + url: &str, +) -> Result> { + let Some(store) = read_fallback_file_unlocked()? else { + return Ok(None); + }; + + let key = compute_store_key(server_name, url)?; + let local_server_name = server_name.strip_prefix("local:").unwrap_or(server_name); + + for (stored_key, entry) in &store { + let matches_credential = if server_name.starts_with("executor:") { + stored_key == &key + && entry.executor_owned + && entry.server_name == server_name + && entry.server_url == url + } else if entry.executor_owned { + false + } else { + entry.server_url == url + // Escaped names may also match another server's stored, escaped name. + // Only accept a legacy unescaped entry under this identity's own key. + && (!server_name.starts_with("local:") || stored_key == &key) + && (entry.server_name == local_server_name + || (stored_key == &key && entry.server_name == server_name)) + }; + if !matches_credential { + continue; + } + + let mut token_response = OAuthTokenResponse::new( + AccessToken::new(entry.access_token.clone()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + + if let Some(refresh) = entry.refresh_token.clone() { + token_response.set_refresh_token(Some(RefreshToken::new(refresh))); + } + + let scopes = entry.scopes.clone(); + if !scopes.is_empty() { + token_response.set_scopes(Some(scopes.into_iter().map(Scope::new).collect())); + } + + let mut stored = StoredOAuthTokens { + server_name: entry.server_name.clone(), + url: entry.server_url.clone(), + issuer: entry.issuer.clone(), + client_id: entry.client_id.clone(), + token_response: WrappedOAuthTokenResponse(token_response), + expires_at: entry.expires_at, + }; + refresh_expires_in_from_timestamp(&mut stored); + + return Ok(Some(stored)); + } + + Ok(None) +} + +/// Saves one credential while holding the File aggregate-store lock across the full +/// read-modify-write operation. +fn save_oauth_tokens_to_file(tokens: &StoredOAuthTokens) -> Result<()> { + let _store_lock = OAuthStoreLock::acquire_for_write(OAuthStore::File)?; + save_oauth_tokens_to_file_with_lock_held(tokens) +} + +/// Updates the fallback File. The caller must hold the File aggregate-store lock. +fn save_oauth_tokens_to_file_with_lock_held(tokens: &StoredOAuthTokens) -> Result<()> { + let key = compute_store_key(&tokens.server_name, &tokens.url)?; + let mut store = read_fallback_file_unlocked()?.unwrap_or_default(); + let executor_owned = tokens.server_name.starts_with("executor:"); + if executor_owned && store.get(&key).is_some_and(|entry| !entry.executor_owned) { + anyhow::bail!("executor OAuth credential key conflicts with a host-owned credential"); + } + + let token_response = &tokens.token_response.0; + let expires_at = tokens + .expires_at + .or_else(|| compute_expires_at_millis(token_response)); + let refresh_token = token_response + .refresh_token() + .map(|token| token.secret().to_string()); + let scopes = token_response + .scopes() + .map(|s| s.iter().map(|s| s.to_string()).collect()) + .unwrap_or_default(); + let entry = FallbackTokenEntry { + server_name: tokens.server_name.clone(), + server_url: tokens.url.clone(), + issuer: tokens.issuer.clone(), + client_id: tokens.client_id.clone(), + access_token: token_response.access_token().secret().to_string(), + expires_at, + refresh_token, + scopes, + executor_owned, + }; + + store.insert(key, entry); + write_fallback_file(&store) +} + +fn delete_oauth_tokens_from_file(key: &str) -> Result { + let _store_lock = OAuthStoreLock::acquire_for_write(OAuthStore::File)?; + let mut store = match read_fallback_file_unlocked()? { + Some(store) => store, + None => return Ok(false), + }; + + if key.starts_with("executor:") + && !key.contains('|') + && store.get(key).is_some_and(|entry| !entry.executor_owned) + { + anyhow::bail!("executor OAuth credential key conflicts with a host-owned credential"); + } + + let removed = store.remove(key).is_some(); + + if removed { + write_fallback_file(&store)?; + } + + Ok(removed) +} + +pub(crate) fn compute_expires_at_millis(response: &OAuthTokenResponse) -> Option { + let expires_in = response.expires_in()?; + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_else(|_| Duration::from_secs(0)); + let expiry = now.checked_add(expires_in)?; + let millis = expiry.as_millis(); + if millis > u128::from(u64::MAX) { + Some(u64::MAX) + } else { + Some(millis as u64) + } +} + +fn expires_in_from_timestamp(expires_at: u64) -> Option { + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_else(|_| Duration::from_secs(0)); + let now_ms = now.as_millis() as u64; + + if expires_at <= now_ms { + None + } else { + Some((expires_at - now_ms) / 1000) + } +} + +fn token_needs_refresh(expires_at: Option) -> bool { + let Some(expires_at) = expires_at else { + return false; + }; + + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_else(|_| Duration::from_secs(0)) + .as_millis() as u64; + + now.saturating_add(REFRESH_SKEW_MILLIS) >= expires_at +} + +fn compute_store_key(server_name: &str, server_url: &str) -> Result { + let executor_owned = server_name.starts_with("executor:"); + let enterprise_owned = server_name.starts_with("ema-idp:"); + let server_name = server_name.strip_prefix("local:").unwrap_or(server_name); + let mut payload = JsonMap::new(); + payload.insert( + "type".to_string(), + Value::String(MCP_SERVER_TYPE.to_string()), + ); + payload.insert("url".to_string(), Value::String(server_url.to_string())); + payload.insert("headers".to_string(), Value::Object(JsonMap::new())); + let payload = if enterprise_owned { + // The OS keyring is shared across homes. Keep enterprise sessions + // isolated by Codex profile as well as authenticated user and workspace. + let codex_home = find_codex_home()?; + fs::create_dir_all(&codex_home)?; + payload.insert( + "codex_home".to_string(), + serde_json::to_value(codex_home.as_path().canonicalize()?)?, + ); + // Different binaries can enable different serde_json ordering features. + serde_json::to_value(payload.into_iter().collect::>())? + } else { + Value::Object(payload) + }; + let truncated = sha_256_prefix(&payload)?; + let separator = if executor_owned { ':' } else { '|' }; + Ok(format!("{server_name}{separator}{truncated}")) +} + +/// Derive a valid secret-store name from the MCP OAuth store key. +/// +/// `compute_store_key` intentionally includes readable identity components and +/// punctuation, but `SecretName` only allows `A-Z`, `0-9`, and `_`. +/// Re-hashing keeps the secret key deterministic while satisfying that +/// restricted alphabet. +fn compute_secret_name(server_name: &str, server_url: &str) -> Result { + let key = compute_store_key(server_name, server_url)?; + let mut hasher = Sha256::new(); + hasher.update(key.as_bytes()); + let digest = hasher.finalize(); + let hex = format!("{digest:X}"); + SecretName::new(&format!("{MCP_OAUTH_SECRET_PREFIX}_{}", &hex[..32])) +} + +fn fallback_file_path() -> Result { + Ok(find_codex_home()?.join(FALLBACK_FILENAME).to_path_buf()) +} + +fn read_fallback_file_unlocked() -> Result> { + let path = fallback_file_path()?; + let contents = match fs::read_to_string(&path) { + Ok(contents) => contents, + Err(err) if err.kind() == ErrorKind::NotFound => return Ok(None), + Err(err) => { + return Err(err).context(format!( + "failed to read credentials file at {}", + path.display() + )); + } + }; + + match serde_json::from_str::(&contents) { + Ok(store) => Ok(Some(store)), + Err(e) => Err(e).context(format!( + "failed to parse credentials file at {}", + path.display() + )), + } +} + +fn open_fallback_file_for_write(path: &std::path::Path) -> Result { + let mut options = fs::OpenOptions::new(); + options.write(true).create(true).truncate(false); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + options + .mode(0o600) + .custom_flags(libc::O_NOFOLLOW | libc::O_NONBLOCK); + } + #[cfg(windows)] + { + use std::os::windows::fs::OpenOptionsExt; + use windows_sys::Win32::Storage::FileSystem::FILE_FLAG_OPEN_REPARSE_POINT; + options.custom_flags(FILE_FLAG_OPEN_REPARSE_POINT); + } + let file = options.open(path)?; + anyhow::ensure!( + file.metadata()?.is_file(), + "credentials path is not a regular file" + ); + Ok(file) +} + +fn write_fallback_file(store: &FallbackFile) -> Result<()> { + let path = fallback_file_path()?; + + if store.is_empty() { + if path.exists() { + fs::remove_file(path)?; + } + return Ok(()); + } + + let parent = path + .parent() + .context("credentials file path has no parent directory")?; + fs::create_dir_all(parent)?; + + let serialized = serde_json::to_string(store)?; + let mut file = open_fallback_file_for_write(&path)?; + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + file.set_permissions(fs::Permissions::from_mode(0o600))?; + } + file.set_len(/*size*/ 0)?; + file.write_all(serialized.as_bytes())?; + + Ok(()) +} + +fn sha_256_prefix(value: &Value) -> Result { + let serialized = + serde_json::to_string(&value).context("failed to serialize MCP OAuth key payload")?; + let mut hasher = Sha256::new(); + hasher.update(serialized.as_bytes()); + let digest = hasher.finalize(); + let hex = format!("{digest:x}"); + let truncated = &hex[..16]; + Ok(truncated.to_string()) +} + +#[cfg(test)] +mod tests { + use super::*; + use anyhow::Result; + use codex_keyring_store::tests::MockKeyringStore; + use codex_secrets::compute_keyring_account; + use keyring::Error as KeyringError; + use pretty_assertions::assert_eq; + use std::sync::Arc; + #[path = "credential_store_tests.rs"] + mod credential_store_tests; + #[path = "persistor_tests.rs"] + mod persistor_tests; + + use super::test_support::TempCodexHome; + + #[test] + fn stored_oauth_credentials_ignore_derived_expiration_and_track_token_changes() -> Result<()> { + let _env = TempCodexHome::new(); + let mut tokens = sample_tokens(); + let credentials = super::normalized_oauth_credentials(Some(&tokens)); + tokens + .token_response + .0 + .set_expires_in(Some(&Duration::from_secs(1))); + assert_eq!( + credentials, + super::normalized_oauth_credentials(Some(&tokens)) + ); + super::save_oauth_tokens_to_file(&tokens)?; + assert_eq!( + credentials, + super::stored_oauth_credentials( + &tokens.server_name, + &tokens.url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + )? + ); + + tokens + .token_response + .0 + .set_access_token(AccessToken::new("new-access-token".to_string())); + super::save_oauth_tokens_to_file(&tokens)?; + assert_ne!( + credentials, + super::stored_oauth_credentials( + &tokens.server_name, + &tokens.url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + )? + ); + Ok(()) + } + + #[test] + fn resolve_oauth_tokens_from_store_policy_uses_keyring_when_available() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let expected = tokens.clone(); + let serialized = serde_json::to_string(&tokens)?; + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + store.save(KEYRING_SERVICE, &key, &serialized)?; + + let resolved = super::resolve_oauth_tokens_from_store_policy( + &store, + &tokens.server_name, + &tokens.url, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + )? + .expect("tokens should load from keyring"); + assert_eq!( + resolved.store, + ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct) + ); + assert_tokens_match_without_expiry(&resolved.tokens, &expected); + Ok(()) + } + + #[test] + fn load_oauth_tokens_falls_back_when_missing_in_keyring() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let expected = tokens.clone(); + + super::save_oauth_tokens_to_file(&tokens)?; + + let resolved = super::resolve_oauth_tokens_from_store_policy( + &store, + &tokens.server_name, + &tokens.url, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + )? + .expect("tokens should load from fallback"); + assert_eq!(resolved.store, ResolvedOAuthCredentialStore::File); + assert_tokens_match_without_expiry(&resolved.tokens, &expected); + Ok(()) + } + + #[test] + fn load_oauth_tokens_falls_back_when_keyring_errors() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let expected = tokens.clone(); + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + store.set_error(&key, KeyringError::Invalid("error".into(), "load".into())); + + super::save_oauth_tokens_to_file(&tokens)?; + + let resolved = super::resolve_oauth_tokens_from_store_policy( + &store, + &tokens.server_name, + &tokens.url, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + )? + .expect("tokens should load from fallback"); + assert_eq!(resolved.store, ResolvedOAuthCredentialStore::File); + assert_tokens_match_without_expiry(&resolved.tokens, &expected); + Ok(()) + } + + #[test] + fn exact_store_operations_do_not_adopt_or_mutate_the_other_store() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let file_tokens = sample_tokens(); + let mut keyring_tokens = file_tokens.clone(); + keyring_tokens + .token_response + .0 + .set_access_token(AccessToken::new("keyring-access-token".to_string())); + + super::save_oauth_tokens_to_file(&file_tokens)?; + let fallback_path = super::fallback_file_path()?; + let fallback_before = fs::read(&fallback_path)?; + super::save_oauth_tokens_with_keyring( + &store, + AuthKeyringBackendKind::Direct, + &keyring_tokens.server_name, + &keyring_tokens, + )?; + + assert_eq!(fs::read(fallback_path)?, fallback_before); + let loaded = ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct) + .load(&store, &keyring_tokens.server_name, &keyring_tokens.url)? + .expect("tokens should load from the selected keyring store"); + assert_tokens_match_without_expiry(&loaded, &keyring_tokens); + Ok(()) + } + + #[test] + fn save_oauth_tokens_prefers_keyring_when_available() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + + super::save_oauth_tokens_to_file(&tokens)?; + + super::save_oauth_tokens_with_keyring_with_fallback_to_file( + &store, + AuthKeyringBackendKind::Direct, + &tokens.server_name, + &tokens, + )?; + + let fallback_path = super::fallback_file_path()?; + assert!(!fallback_path.exists(), "fallback file should be removed"); + let stored = store.saved_value(&key).expect("value saved to keyring"); + assert_eq!(serde_json::from_str::(&stored)?, tokens); + Ok(()) + } + + #[test] + fn save_oauth_tokens_writes_fallback_when_keyring_fails() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + store.set_error(&key, KeyringError::Invalid("error".into(), "save".into())); + + super::save_oauth_tokens_with_keyring_with_fallback_to_file( + &store, + AuthKeyringBackendKind::Direct, + &tokens.server_name, + &tokens, + )?; + + let fallback_path = super::fallback_file_path()?; + assert!(fallback_path.exists(), "fallback file should be created"); + let saved = super::read_fallback_file_unlocked()?.expect("fallback file should load"); + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + let entry = saved.get(&key).expect("entry for key"); + assert_eq!(entry.server_name, tokens.server_name); + assert_eq!(entry.server_url, tokens.url); + assert_eq!(entry.client_id, tokens.client_id); + assert_eq!( + entry.access_token, + tokens.token_response.0.access_token().secret().as_str() + ); + assert!(store.saved_value(&key).is_none()); + Ok(()) + } + + #[cfg(unix)] + #[test] + fn fallback_file_is_private_at_creation() -> Result<()> { + use std::os::unix::fs::PermissionsExt; + const CHILD: &str = "CODEX_TEST_OAUTH_PERMISSIVE_UMASK"; + + if std::env::var_os(CHILD).is_none() { + // Change umask only in the child running this one test. + let status = std::process::Command::new("/bin/sh") + .args(["-c", "umask 000; exec \"$@\"", "sh"]) + .arg(std::env::current_exe()?) + .args([ + "--exact", + "oauth::tests::fallback_file_is_private_at_creation", + ]) + .env(CHILD, "1") + .status()?; + anyhow::ensure!(status.success(), "creation-permissions test failed"); + return Ok(()); + } + + let _env = TempCodexHome::new(); + let path = fallback_file_path()?; + let file = open_fallback_file_for_write(&path)?; + assert_eq!(file.metadata()?.permissions().mode() & 0o777, 0o600); + Ok(()) + } + + #[test] + fn fallback_file_updates_the_existing_file() -> Result<()> { + #[cfg(unix)] + use std::os::unix::fs::PermissionsExt; + + let env = TempCodexHome::new(); + save_oauth_tokens_to_file(&sample_tokens())?; + let path = fallback_file_path()?; + let original = env.path().join("original-file"); + fs::hard_link(&path, &original)?; + #[cfg(unix)] + fs::set_permissions(&path, fs::Permissions::from_mode(0o644))?; + + let mut store = read_fallback_file_unlocked()?.expect("saved credentials"); + store.values_mut().next().unwrap().access_token = "new".to_string(); + write_fallback_file(&store)?; + + let expected = serde_json::to_vec(&store)?; + assert_eq!( + [fs::read(original)?, fs::read(&path)?], + [expected.clone(), expected] + ); + #[cfg(unix)] + assert_eq!(fs::metadata(&path)?.permissions().mode() & 0o777, 0o600); + Ok(()) + } + + #[cfg(any(unix, windows))] + #[test] + fn fallback_file_write_does_not_follow_symlinks() -> Result<()> { + #[cfg(unix)] + use std::os::unix::fs::symlink; + #[cfg(windows)] + use std::os::windows::fs::symlink_file as symlink; + + let env = TempCodexHome::new(); + let path = fallback_file_path()?; + let target = env.path().join("symlink-target"); + fs::write(&target, "synthetic credentials")?; + let linked = symlink(&target, &path); + #[cfg(windows)] + if linked + .as_ref() + .is_err_and(|error| error.raw_os_error() == Some(1314)) + { + eprintln!("Skipping symlink test: Windows symlink privilege unavailable"); + return Ok(()); + } + linked?; + + assert!(open_fallback_file_for_write(&path).is_err()); + + assert_eq!(fs::read_to_string(target)?, "synthetic credentials"); + assert!(fs::symlink_metadata(path)?.file_type().is_symlink()); + Ok(()) + } + + #[test] + fn save_oauth_tokens_with_secrets_backend_writes_encrypted_storage() -> Result<()> { + let env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + let serialized = serde_json::to_string(&tokens)?; + store.save(KEYRING_SERVICE, &key, &serialized)?; + super::save_oauth_tokens_to_file(&tokens)?; + + super::save_oauth_tokens_with_keyring_with_fallback_to_file( + &store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens, + )?; + + let manager = SecretsManager::new_with_keyring_store_and_namespace( + env.path().to_path_buf(), + SecretsBackendKind::Local, + Arc::new(store.clone()), + LocalSecretsNamespace::McpOAuth, + ); + let secret_name = super::compute_secret_name(&tokens.server_name, &tokens.url)?; + let stored = manager + .get(&SecretScope::Global, &secret_name)? + .expect("tokens should be saved to encrypted storage"); + assert_eq!(serde_json::from_str::(&stored)?, tokens); + assert_eq!(store.saved_value(&key), Some(serialized)); + assert!(env.path().join("secrets").join("mcp_oauth.age").exists()); + assert!(!env.path().join("secrets").join("local.age").exists()); + assert!(!super::fallback_file_path()?.exists()); + Ok(()) + } + + #[test] + fn load_oauth_tokens_with_secrets_backend_reads_encrypted_storage() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let expected = tokens.clone(); + + super::save_oauth_tokens_with_keyring( + &store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens, + )?; + + let loaded = super::load_oauth_tokens_from_keyring( + &store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens.url, + )? + .expect("tokens should load from encrypted storage"); + assert_tokens_match_without_expiry(&loaded, &expected); + Ok(()) + } + + #[test] + fn load_oauth_tokens_with_secrets_backend_ignores_direct_entry() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + let serialized = serde_json::to_string(&tokens)?; + store.save(KEYRING_SERVICE, &key, &serialized)?; + + let loaded = super::load_oauth_tokens_from_keyring( + &store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens.url, + )?; + + assert!(loaded.is_none()); + Ok(()) + } + + #[test] + fn save_oauth_tokens_with_secrets_backend_falls_back_to_file_when_keyring_fails() -> Result<()> + { + let env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + store.set_error( + &compute_keyring_account(env.path()), + KeyringError::Invalid("error".into(), "save".into()), + ); + let tokens = sample_tokens(); + + super::save_oauth_tokens_with_keyring_with_fallback_to_file( + &store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens, + )?; + + let saved = super::read_fallback_file_unlocked()?.expect("fallback file should load"); + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + assert!(saved.contains_key(&key)); + Ok(()) + } + + #[test] + fn delete_oauth_tokens_with_secrets_backend_removes_secrets_and_file() -> Result<()> { + let env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let serialized = serde_json::to_string(&tokens)?; + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + store.save(KEYRING_SERVICE, &key, &serialized)?; + super::save_oauth_tokens_with_keyring( + &store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens, + )?; + store.save(KEYRING_SERVICE, &key, &serialized)?; + super::save_oauth_tokens_to_file(&tokens)?; + + let removed = super::delete_oauth_tokens_from_keyring_and_file( + &store, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens.url, + )?; + + let manager = SecretsManager::new_with_keyring_store_and_namespace( + env.path().to_path_buf(), + SecretsBackendKind::Local, + Arc::new(store.clone()), + LocalSecretsNamespace::McpOAuth, + ); + let secret_name = super::compute_secret_name(&tokens.server_name, &tokens.url)?; + assert!(removed); + assert!(manager.get(&SecretScope::Global, &secret_name)?.is_none()); + assert!(store.saved_value(&key).is_none()); + assert!(!super::fallback_file_path()?.exists()); + Ok(()) + } + + #[test] + fn delete_oauth_tokens_removes_all_storage() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let serialized = serde_json::to_string(&tokens)?; + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + store.save(KEYRING_SERVICE, &key, &serialized)?; + super::save_oauth_tokens_to_file(&tokens)?; + + let removed = super::delete_oauth_tokens_from_keyring_and_file( + &store, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + &tokens.server_name, + &tokens.url, + )?; + assert!(removed); + assert!(!store.contains(&key)); + assert!(!super::fallback_file_path()?.exists()); + Ok(()) + } + + #[test] + fn delete_oauth_tokens_file_mode_removes_keyring_only_entry() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let serialized = serde_json::to_string(&tokens)?; + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + store.save(KEYRING_SERVICE, &key, &serialized)?; + assert!(store.contains(&key)); + + let removed = super::delete_oauth_tokens_from_keyring_and_file( + &store, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + &tokens.server_name, + &tokens.url, + )?; + assert!(removed); + assert!(!store.contains(&key)); + assert!(!super::fallback_file_path()?.exists()); + Ok(()) + } + + #[test] + fn delete_oauth_tokens_propagates_keyring_errors() -> Result<()> { + let _env = TempCodexHome::new(); + let store = MockKeyringStore::default(); + let tokens = sample_tokens(); + let key = super::compute_store_key(&tokens.server_name, &tokens.url)?; + store.set_error(&key, KeyringError::Invalid("error".into(), "delete".into())); + super::save_oauth_tokens_to_file(&tokens).unwrap(); + + let result = super::delete_oauth_tokens_from_keyring_and_file( + &store, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + &tokens.server_name, + &tokens.url, + ); + assert!(result.is_err()); + assert!(super::fallback_file_path().unwrap().exists()); + Ok(()) + } + + #[test] + fn refresh_expires_in_from_timestamp_restores_future_durations() { + let mut tokens = sample_tokens(); + let expires_at = tokens.expires_at.expect("expires_at should be set"); + + tokens.token_response.0.set_expires_in(None); + super::refresh_expires_in_from_timestamp(&mut tokens); + + let actual = tokens + .token_response + .0 + .expires_in() + .expect("expires_in should be restored") + .as_secs(); + let expected = super::expires_in_from_timestamp(expires_at) + .expect("expires_at should still be in the future"); + let diff = actual.abs_diff(expected); + assert!(diff <= 1, "expires_in drift too large: diff={diff}"); + } + + #[test] + fn refresh_expires_in_from_timestamp_marks_expired_tokens() { + let mut tokens = sample_tokens(); + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_else(|_| Duration::from_secs(0)); + let expired_at = now.as_millis() as u64; + tokens.expires_at = Some(expired_at.saturating_sub(1000)); + + let duration = Duration::from_secs(600); + tokens.token_response.0.set_expires_in(Some(&duration)); + + super::refresh_expires_in_from_timestamp(&mut tokens); + + assert_eq!(tokens.token_response.0.expires_in(), Some(Duration::ZERO)); + } + + #[test] + fn oauth_tokens_are_usable_when_expiry_is_unknown() { + let mut tokens = sample_tokens(); + tokens.expires_at = None; + tokens.token_response.0.set_refresh_token(None); + + assert!(super::oauth_tokens_are_usable(&tokens)); + } + + #[test] + fn oauth_tokens_are_usable_when_unexpired_without_refresh_token() { + let mut tokens = sample_tokens(); + tokens.token_response.0.set_refresh_token(None); + + assert!(super::oauth_tokens_are_usable(&tokens)); + } + + #[test] + fn oauth_tokens_are_usable_when_expired_but_refreshable() { + let mut tokens = sample_tokens(); + tokens.expires_at = Some(0); + + assert!(super::oauth_tokens_are_usable(&tokens)); + } + + #[test] + fn oauth_tokens_are_not_usable_when_expired_and_unrefreshable() { + let mut tokens = sample_tokens(); + tokens.expires_at = Some(0); + tokens.token_response.0.set_refresh_token(None); + + assert!(!super::oauth_tokens_are_usable(&tokens)); + } + + #[test] + fn oauth_tokens_are_not_usable_when_near_expiry_and_unrefreshable() { + let mut tokens = sample_tokens(); + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_else(|_| Duration::from_secs(0)) + .as_millis() as u64; + tokens.expires_at = Some(now.saturating_add(REFRESH_SKEW_MILLIS - 1)); + tokens.token_response.0.set_refresh_token(None); + + assert!(!super::oauth_tokens_are_usable(&tokens)); + } + + #[test] + fn oauth_tokens_are_not_usable_when_client_id_is_blank() { + let mut tokens = sample_tokens(); + tokens.client_id = " ".to_string(); + + assert!(!super::oauth_tokens_are_usable(&tokens)); + } + + #[test] + fn oauth_tokens_are_not_usable_when_access_token_is_blank() { + let mut tokens = sample_tokens(); + tokens + .token_response + .0 + .set_access_token(AccessToken::new(" ".to_string())); + + assert!(!super::oauth_tokens_are_usable(&tokens)); + } + + #[test] + fn oauth_tokens_are_not_usable_when_required_refresh_token_is_blank() { + let mut tokens = sample_tokens(); + tokens.expires_at = Some(0); + tokens + .token_response + .0 + .set_refresh_token(Some(RefreshToken::new(" ".to_string()))); + + assert!(!super::oauth_tokens_are_usable(&tokens)); + } + + fn assert_tokens_match_without_expiry( + actual: &StoredOAuthTokens, + expected: &StoredOAuthTokens, + ) { + assert_eq!(actual.server_name, expected.server_name); + assert_eq!(actual.url, expected.url); + assert_eq!(actual.issuer, expected.issuer); + assert_eq!(actual.client_id, expected.client_id); + assert_eq!(actual.expires_at, expected.expires_at); + assert_token_response_match_without_expiry( + &actual.token_response, + &expected.token_response, + ); + } + + fn assert_token_response_match_without_expiry( + actual: &WrappedOAuthTokenResponse, + expected: &WrappedOAuthTokenResponse, + ) { + let actual_response = &actual.0; + let expected_response = &expected.0; + + assert_eq!( + actual_response.access_token().secret(), + expected_response.access_token().secret() + ); + assert_eq!(actual_response.token_type(), expected_response.token_type()); + assert_eq!( + actual_response.refresh_token().map(RefreshToken::secret), + expected_response.refresh_token().map(RefreshToken::secret), + ); + assert_eq!(actual_response.scopes(), expected_response.scopes()); + assert_eq!( + actual_response.extra_fields().0, + expected_response.extra_fields().0 + ); + assert_eq!( + actual_response.expires_in().is_some(), + expected_response.expires_in().is_some() + ); + } + + fn sample_tokens() -> StoredOAuthTokens { + let mut response = OAuthTokenResponse::new( + AccessToken::new("access-token".to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new("refresh-token".to_string()))); + response.set_scopes(Some(vec![ + Scope::new("scope-a".to_string()), + Scope::new("scope-b".to_string()), + ])); + let expires_in = Duration::from_secs(3600); + response.set_expires_in(Some(&expires_in)); + let expires_at = super::compute_expires_at_millis(&response); + + StoredOAuthTokens { + server_name: "test-server".to_string(), + url: "https://example.test".to_string(), + issuer: Some("https://issuer.example.test".to_string()), + client_id: "client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at, + } + } +} diff --git a/codex-rs/rmcp-client/src/oauth/credential_store.rs b/codex-rs/rmcp-client/src/oauth/credential_store.rs new file mode 100644 index 0000000000000000000000000000000000000000..0429efde4241cc34d7deb9361dcb1bb853d36dbe --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/credential_store.rs @@ -0,0 +1,324 @@ +//! Adapts Codex's pinned credential storage to RMCP-owned OAuth refreshes. +//! +//! Attach this store only after `AuthorizationManager::initialize_from_store` completes +//! against an in-memory store. Initialization may save tokenless client credentials, which +//! this refresh adapter does not support. +//! +//! Ordinary token reads use the cached credentials. Refresh-guard acquisition rereads +//! the pinned store before RMCP exchanges the token and saves the result. Codex +//! preparation rechecks freshness under that guard before asking RMCP to refresh. +//! Saves and clears require an active guard; no store operation reacquires the lock. +//! The runtime snapshot advances only for credentials compatible with this connection, +//! so replacement logins still cause a rebuild. +//! Synchronous store operations run off the Tokio workers so caller deadlines remain pollable. +//! Blocking mutations retain the caller's transaction guard even if their await is cancelled. + +use std::sync::Arc; +use std::sync::Weak; + +use anyhow::Context; +use anyhow::Result; +use codex_keyring_store::DefaultKeyringStore; +use codex_keyring_store::KeyringStore; +use futures::future::BoxFuture; +use oauth2::Scope; +use oauth2::TokenResponse; +use rmcp::transport::auth::AuthError; +use rmcp::transport::auth::AuthorizationManager; +use rmcp::transport::auth::CredentialRefreshGuard; +use rmcp::transport::auth::CredentialStore; +use rmcp::transport::auth::StoredCredentials; +use tokio::sync::Mutex; +use tracing::warn; + +use crate::oauth_http_client::PROACTIVE_REFRESH_TIMEOUT; + +use super::RefreshCredentialLock; +use super::ResolvedOAuthCredentialStore; +use super::StoredOAuthTokens; +use super::WrappedOAuthTokenResponse; +use super::normalized_oauth_credentials; +use super::refresh_expires_in_from_timestamp; +use super::refresh_transaction::REFRESH_REQUEST_TIMEOUT; +use super::token_needs_refresh; +use super::validate_refresh_token_issuer; + +#[derive(Clone)] +pub(crate) struct OAuthCredentialStore { + inner: Arc>, + held_refresh_guard: Option>, +} + +struct OAuthCredentialStoreInner { + server_name: String, + url: String, + client_id: String, + issuer: Option, + store: ResolvedOAuthCredentialStore, + keyring: K, + last_credentials: Mutex>, + refresh_guard: Mutex>, +} + +impl OAuthCredentialStore { + pub(crate) fn new( + tokens: StoredOAuthTokens, + store: ResolvedOAuthCredentialStore, + keyring: K, + ) -> Self { + Self { + inner: Arc::new(OAuthCredentialStoreInner { + server_name: tokens.server_name.clone(), + url: tokens.url.clone(), + client_id: tokens.client_id.clone(), + issuer: tokens.bound_issuer().map(str::to_owned), + store, + keyring, + last_credentials: Mutex::new(Some(tokens)), + refresh_guard: Mutex::new(Weak::new()), + }), + held_refresh_guard: None, + } + } + + pub(crate) async fn refresh_if_needed(&self, manager: &mut AuthorizationManager) -> Result<()> { + let guard = self.acquire_transaction_guard().await?; + let tokens = self + .inner + .last_credentials + .lock() + .await + .clone() + .ok_or(AuthError::AuthorizationRequired)?; + if !token_needs_refresh(tokens.expires_at) { + return Ok(()); + } + let metadata = manager.resolve_metadata().await?.metadata; + validate_refresh_token_issuer(&metadata, &tokens)?; + manager.set_metadata(metadata); + manager.configure_client_id(&tokens.client_id)?; + // Reuse the guard for RMCP's exchange and save after the locked freshness check. + manager.set_credential_store(Self { + inner: Arc::clone(&self.inner), + held_refresh_guard: Some(guard), + }); + let result = PROACTIVE_REFRESH_TIMEOUT + .scope(REFRESH_REQUEST_TIMEOUT, manager.refresh_token()) + .await; + manager.set_credential_store(self.clone()); + match result { + Ok(_) => Ok(()), + Err(AuthError::TokenRefreshRejected(_)) => Err(AuthError::AuthorizationRequired.into()), + Err(error) => Err(error.into()), + } + } + + pub(crate) async fn stored_credentials(&self) -> Option { + normalized_oauth_credentials(self.inner.last_credentials.lock().await.as_ref()) + } + + pub(crate) async fn acquire_transaction_guard( + &self, + ) -> Result, AuthError> { + let guard = Arc::new( + RefreshCredentialLock::acquire_for_server(&self.inner.server_name, &self.inner.url) + .await + .map_err(credential_store_error)?, + ); + let inner = Arc::clone(&self.inner); + let tokens = tokio::task::spawn_blocking(move || { + inner + .store + .load(&inner.keyring, &inner.server_name, &inner.url) + }) + .await + .context("OAuth credential load task failed") + .map_err(credential_store_error)? + .map_err(credential_store_error)? + .ok_or(AuthError::AuthorizationRequired)?; + self.validate_connection(&tokens)?; + *self.inner.last_credentials.lock().await = Some(tokens); + *self.inner.refresh_guard.lock().await = Arc::downgrade(&guard); + Ok(guard) + } + + fn validate_connection(&self, tokens: &StoredOAuthTokens) -> Result<(), AuthError> { + if tokens.client_id != self.inner.client_id + || tokens.bound_issuer() != self.inner.issuer.as_deref() + || tokens.has_refresh_token() && tokens.bound_issuer().is_none() + { + warn!( + server_name = %self.inner.server_name, + "stored OAuth credentials no longer match this connection's client and issuer; authorization must be rebuilt" + ); + return Err(AuthError::AuthorizationRequired); + } + if !tokens.has_refresh_token() && !tokens.access_token_is_usable_without_refresh() { + return Err(AuthError::TokenExpired); + } + Ok(()) + } +} + +// Spell out RMCP's object-safe boxed-future interface without introducing async-trait here. +impl CredentialStore for OAuthCredentialStore { + fn load<'life0, 'async_trait>( + &'life0 self, + ) -> BoxFuture<'async_trait, Result, AuthError>> + where + 'life0: 'async_trait, + Self: 'async_trait, + { + Box::pin(async move { + Ok(self + .inner + .last_credentials + .lock() + .await + .as_ref() + .map(rmcp_credentials)) + }) + } + + fn save<'life0, 'async_trait>( + &'life0 self, + credentials: StoredCredentials, + ) -> BoxFuture<'async_trait, Result<(), AuthError>> + where + 'life0: 'async_trait, + Self: 'async_trait, + { + Box::pin(async move { + let mut token_response = credentials + .token_response + .ok_or(AuthError::AuthorizationRequired)?; + // Codex stores granted scopes in the token response rather than a separate field. + if token_response.scopes().is_none() && !credentials.granted_scopes.is_empty() { + token_response.set_scopes(Some( + credentials + .granted_scopes + .into_iter() + .map(Scope::new) + .collect(), + )); + } + // The SDK's receipt time, rather than the time this save completes, owns expiry. + if credentials.token_received_at.is_none() { + token_response.set_expires_in(None); + } + let expires_at = credentials.token_received_at.and_then(|received_at| { + token_response.expires_in().map(|expires_in| { + received_at + .saturating_mul(1000) + .saturating_add(u64::try_from(expires_in.as_millis()).unwrap_or(u64::MAX)) + }) + }); + let tokens = StoredOAuthTokens { + server_name: self.inner.server_name.clone(), + url: self.inner.url.clone(), + issuer: credentials.issuer, + client_id: credentials.client_id, + token_response: WrappedOAuthTokenResponse(token_response), + expires_at, + }; + self.validate_connection(&tokens)?; + let inner = Arc::clone(&self.inner); + let guard = self + .inner + .refresh_guard + .lock() + .await + .upgrade() + .context("OAuth credential mutation requires an active refresh guard") + .map_err(credential_store_error)?; + let tokens = tokio::task::spawn_blocking(move || -> Result { + let _guard = guard; + inner + .store + .save(&inner.keyring, &inner.server_name, &tokens)?; + Ok(tokens) + }) + .await + .context("OAuth credential save task failed") + .map_err(credential_store_error)? + .map_err(credential_store_error)?; + *self.inner.last_credentials.lock().await = Some(tokens); + Ok(()) + }) + } + + fn clear<'life0, 'async_trait>(&'life0 self) -> BoxFuture<'async_trait, Result<(), AuthError>> + where + 'life0: 'async_trait, + Self: 'async_trait, + { + Box::pin(async move { + let inner = Arc::clone(&self.inner); + let guard = self + .inner + .refresh_guard + .lock() + .await + .upgrade() + .context("OAuth credential mutation requires an active refresh guard") + .map_err(credential_store_error)?; + tokio::task::spawn_blocking(move || -> Result<()> { + let _guard = guard; + inner + .store + .delete(&inner.keyring, &inner.server_name, &inner.url)?; + Ok(()) + }) + .await + .context("OAuth credential removal task failed") + .map_err(credential_store_error)? + .map_err(credential_store_error) + }) + } + + fn acquire_refresh_guard<'life0, 'async_trait>( + &'life0 self, + ) -> BoxFuture<'async_trait, Result, AuthError>> + where + 'life0: 'async_trait, + Self: 'async_trait, + { + Box::pin(async move { + let guard = match &self.held_refresh_guard { + Some(guard) => Arc::clone(guard), + None => self.acquire_transaction_guard().await?, + }; + Ok(Some(CredentialRefreshGuard::new(guard))) + }) + } +} + +fn credential_store_error(error: anyhow::Error) -> AuthError { + AuthError::CredentialStoreError(format!("{error:#}")) +} + +fn rmcp_credentials(tokens: &StoredOAuthTokens) -> StoredCredentials { + let mut tokens = tokens.clone(); + refresh_expires_in_from_timestamp(&mut tokens); + let token_received_at = tokens.expires_at.map(|expires_at| { + let remaining = tokens.token_response.0.expires_in().unwrap_or_default(); + // Reconstruct the receipt time from the authority's deadline; a second clock read + // could otherwise extend the grant if this task paused between the two reads. + expires_at.saturating_sub(u64::try_from(remaining.as_millis()).unwrap_or(u64::MAX)) / 1000 + }); + if token_received_at.is_none() { + tokens.token_response.0.set_expires_in(None); + } + let token_response = tokens.token_response.0; + let granted_scopes = token_response + .scopes() + .map(|scopes| scopes.iter().map(|scope| scope.to_string()).collect()) + .unwrap_or_default(); + StoredCredentials::new( + tokens.client_id, + Some(token_response), + granted_scopes, + token_received_at, + ) + .with_issuer(tokens.issuer) +} diff --git a/codex-rs/rmcp-client/src/oauth/ema_identity.rs b/codex-rs/rmcp-client/src/oauth/ema_identity.rs new file mode 100644 index 0000000000000000000000000000000000000000..6f804a245f63d8862b7953a1b38c4dad1df0779b --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/ema_identity.rs @@ -0,0 +1,78 @@ +//! Reread enterprise credentials without changing the connection's pinned identity. + +use anyhow::Result; +use anyhow::bail; +use codex_keyring_store::DefaultKeyringStore; +use codex_keyring_store::KeyringStore; + +use super::ResolvedOAuthCredentialStore; +use super::StoredOAuthCredentialSnapshot; +use super::StoredOAuthTokens; +use crate::ema_auth_policy::ema_reauthentication_required; +use crate::ema_claims::OidcClaims; +use crate::ema_claims::oidc_identity; + +pub(crate) fn stored_oidc_identity(tokens: &StoredOAuthTokens) -> Result { + let assertion = tokens + .token_response + .0 + .extra_fields() + .0 + .get("id_token") + .and_then(serde_json::Value::as_str) + .ok_or_else(|| { + ema_reauthentication_required( + "enterprise IdP session has no OIDC ID token; sign in again", + ) + })?; + // The ID token binds the login identity. Its expiry does not determine + // whether the IdP will accept the independently valid refresh token. + oidc_identity(assertion, &tokens.url, &tokens.client_id).map_err(|error| { + ema_reauthentication_required("stored enterprise IdP identity is invalid; sign in again") + .context(error.to_string()) + }) +} + +impl StoredOAuthCredentialSnapshot { + /// Reject deletion or replacement before using this session's cached resource bearer. + /// Call from a blocking task: the pinned keyring backend may perform blocking I/O. + pub fn validate_current_ema_credentials(&self) -> Result<()> { + self.load_ema_credentials(&DefaultKeyringStore).map(|_| ()) + } + + /// Reread only the pinned keyring authority; exchange callers hold its credential lock. + pub(crate) fn load_ema_credentials( + &self, + keyring_store: &K, + ) -> Result { + if !matches!(self.store, ResolvedOAuthCredentialStore::Keyring(_)) { + bail!("enterprise IdP credentials require keyring storage"); + } + let previous = &self.credentials; + let mut latest = self + .store + .load(keyring_store, &previous.server_name, &previous.url)? + .ok_or_else(|| { + ema_reauthentication_required( + "enterprise IdP credentials were removed; sign in again", + ) + })?; + // There is no refresh-token rotation writer in the supported EMA profile. + // Pin the whole atomic login record, not just the claims of its ID token. + latest.token_response.0.set_expires_in(None); + if latest != *previous + || previous.bound_issuer() != Some(previous.url.as_str()) + || !latest.has_refresh_token() + { + return Err(ema_reauthentication_required( + "enterprise IdP identity changed; sign in again and reconnect", + )); + } + stored_oidc_identity(&latest)?; + Ok(latest) + } +} + +#[cfg(test)] +#[path = "ema_identity_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/oauth/ema_identity_tests.rs b/codex-rs/rmcp-client/src/oauth/ema_identity_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..c66721f3959b966ef1054daabaac572c0b925302 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/ema_identity_tests.rs @@ -0,0 +1,85 @@ +use codex_config::types::AuthKeyringBackendKind; +use codex_keyring_store::tests::MockKeyringStore; +use serde_json::json; + +use super::*; +use crate::oauth::RefreshCredentialLock; +use crate::oauth::compute_store_key; +use crate::oauth::test_support::TempCodexHome; + +#[tokio::test] +async fn keyring_failure_does_not_reuse_the_pinned_refresh_token() -> Result<()> { + let _home = TempCodexHome::new(); + let tokens: StoredOAuthTokens = serde_json::from_value(json!({ + "server_name": "ema-idp:keyring-failure", + "url": "https://idp.example", + "issuer": "https://idp.example", + "client_id": "client", + "token_response": { + "access_token": "unused", + "token_type": "Bearer", + "refresh_token": "stale-refresh", + }, + }))?; + let key = compute_store_key(&tokens.server_name, &tokens.url)?; + let _lock = RefreshCredentialLock::acquire_for_server(&tokens.server_name, &tokens.url).await?; + let snapshot = StoredOAuthCredentialSnapshot::new( + tokens, + ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct), + ); + let keyring = MockKeyringStore::default(); + keyring.set_error( + &key, + keyring::Error::Invalid("backend".into(), "unavailable".into()), + ); + let error = snapshot + .load_ema_credentials(&keyring) + .expect_err("keyring failure must be terminal"); + assert!(error.to_string().contains("refusing file fallback")); + Ok(()) +} + +#[test] +fn ordinary_oauth_names_cannot_alias_enterprise_credential_keys() -> Result<()> { + let _home = TempCodexHome::new(); + let issuer = "https://idp.example"; + let enterprise_name = "ema-idp:synthetic-identity"; + let ordinary: codex_config::McpServerConfig = serde_json::from_value(json!({ + "url": issuer, + "oauth": {"client_id": "idp-client"}, + }))?; + let ordinary_name = ordinary.oauth_credential_name(enterprise_name); + + pretty_assertions::assert_ne!( + compute_store_key(&ordinary_name, issuer)?, + compute_store_key(enterprise_name, issuer)?, + "an ordinary server name must not select the enterprise credential namespace" + ); + let legacy_key = compute_store_key("ordinary-server", issuer)?; + let legacy_hash = legacy_key.split_once('|').expect("legacy key separator").1; + pretty_assertions::assert_eq!( + compute_store_key(&ordinary_name, issuer)?, + format!("{enterprise_name}|{legacy_hash}"), + "escaping the reserved prefix preserves the pre-EMA ordinary credential key" + ); + + let keyring = MockKeyringStore::default(); + let store = ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct); + let enterprise_tokens: StoredOAuthTokens = serde_json::from_value(json!({ + "server_name": enterprise_name, "url": issuer, "issuer": issuer, + "client_id": "idp-client", "token_response": { + "access_token": "unused", "token_type": "Bearer", "refresh_token": "enterprise-refresh" + } + }))?; + store.save(&keyring, enterprise_name, &enterprise_tokens)?; + let mut ordinary_tokens = enterprise_tokens.clone(); + ordinary_tokens.server_name = ordinary_name.to_string(); + store.save(&keyring, &ordinary_name, &ordinary_tokens)?; + assert!(store.delete(&keyring, &ordinary_name, issuer)?); + pretty_assertions::assert_eq!( + store.load(&keyring, enterprise_name, issuer)?, + Some(enterprise_tokens), + "ordinary OAuth save/logout must not overwrite or remove the enterprise entry" + ); + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/oauth/enterprise_generation.rs b/codex-rs/rmcp-client/src/oauth/enterprise_generation.rs new file mode 100644 index 0000000000000000000000000000000000000000..cd227cfbebd98847527e4c396ab7f3e263093a9c --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/enterprise_generation.rs @@ -0,0 +1,85 @@ +//! Persistent, non-secret logout generations for staged enterprise logins. +//! All reads and writes require the credential lock; missing state never admits an old attempt. + +use std::fs::File; +use std::fs::OpenOptions; +use std::io::Read; +use std::io::Seek; +use std::io::Write; + +use anyhow::Result; +use anyhow::ensure; +use codex_utils_home_dir::find_codex_home; +use oauth2::CsrfToken; +use sha2::Digest; +use sha2::Sha256; + +use super::RefreshCredentialLock; + +#[derive(PartialEq, Eq)] +pub(crate) struct EnterpriseOAuthGeneration([u8; 32]); + +pub(crate) struct EnterpriseOAuthGenerationFile { + file: File, +} + +impl EnterpriseOAuthGenerationFile { + pub(crate) fn open( + credential_name: &str, + issuer: &str, + _lock: &RefreshCredentialLock, + ) -> Result { + let key = super::compute_store_key(credential_name, issuer)?; + let name = format!("{:x}.enterprise-generation", Sha256::digest(key.as_bytes())); + let path = find_codex_home()?.join("mcp-oauth-locks").join(name); + let mut options = OpenOptions::new(); + options.read(true).write(true).create(true).truncate(false); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + options + .mode(0o600) + .custom_flags(libc::O_NOFOLLOW | libc::O_NONBLOCK); + } + #[cfg(windows)] + { + use std::os::windows::fs::OpenOptionsExt; + use windows_sys::Win32::Storage::FileSystem::FILE_FLAG_OPEN_REPARSE_POINT; + options.custom_flags(FILE_FLAG_OPEN_REPARSE_POINT); + } + let file = options.open(path)?; + ensure!( + file.metadata()?.is_file(), + "invalid enterprise generation file" + ); + Ok(Self { file }) + } + + pub(crate) fn current(&self) -> Result> { + let mut file = &self.file; + file.rewind()?; + match file.metadata()?.len() { + 0 => Ok(None), + 32 => { + let mut generation = [0; 32]; + file.read_exact(&mut generation)?; + Ok(Some(EnterpriseOAuthGeneration(generation))) + } + _ => anyhow::bail!("invalid enterprise generation file"), + } + } + + pub(crate) fn replace(&self) -> Result { + // Reuse OAuth's randomness without storing any credential. Random initialization + // also prevents an old attempt becoming valid if metadata is removed or truncated. + let random = CsrfToken::new_random_len(/*num_bytes*/ 32); + let generation = + EnterpriseOAuthGeneration(Sha256::digest(random.secret().as_bytes()).into()); + let mut file = &self.file; + file.set_len(/*size*/ 0)?; + file.rewind()?; + file.write_all(&generation.0)?; + file.sync_all()?; + Ok(generation) + } +} diff --git a/codex-rs/rmcp-client/src/oauth/issuer_binding.rs b/codex-rs/rmcp-client/src/oauth/issuer_binding.rs new file mode 100644 index 0000000000000000000000000000000000000000..5a12bf4024e1ee0484e7b033f7589143ec74a1b9 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/issuer_binding.rs @@ -0,0 +1,121 @@ +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use rmcp::transport::auth::AuthError; +use rmcp::transport::auth::AuthorizationMetadata; +use url::Url; + +use super::StoredOAuthTokens; + +/// Reject authorization endpoints that cannot be bound to their actual issuer. +pub(crate) fn validate_authorization_server_endpoints( + metadata: &AuthorizationMetadata, +) -> Result<()> { + let authorization_endpoint = Url::parse(&metadata.authorization_endpoint) + .context("OAuth authorization endpoint must be a valid URL")?; + let issuer = metadata + .issuer + .as_deref() + .filter(|issuer| !issuer.trim().is_empty()) + .map(Url::parse) + .transpose() + .context("OAuth authorization server issuer must be a valid URL")?; + let issuer_bound_callbacks = metadata + .additional_fields + .get("authorization_response_iss_parameter_supported") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false); + + if issuer_bound_callbacks { + if issuer.is_none() { + bail!("OAuth issuer-bound callbacks require an authorization server issuer"); + } + return Ok(()); + } + + let token_endpoint = + Url::parse(&metadata.token_endpoint).context("OAuth token endpoint must be a valid URL")?; + + if let Some(issuer) = issuer { + if authorization_endpoint.origin() == issuer.origin() + || authorization_endpoint.origin() == token_endpoint.origin() + // Remove these narrow compatibility exceptions once both providers support RFC 9207. + || matches!( + ( + issuer.as_str(), + authorization_endpoint.origin().ascii_serialization().as_str(), + token_endpoint.origin().ascii_serialization().as_str(), + ), + ( + "https://api.figma.com/", + "https://www.figma.com", + "https://api.figma.com", + ) | ( + "https://agent.robinhood.com/mcp/trading", + "https://robinhood.com", + "https://api.robinhood.com", + ) + ) + { + return Ok(()); + } + bail!( + "OAuth authorization endpoint origin does not match the authorization server origin without issuer-bound callbacks" + ); + } + + if token_endpoint.origin() != authorization_endpoint.origin() { + bail!( + "OAuth token endpoint origin does not match the authorization server origin without issuer-bound callbacks" + ); + } + + Ok(()) +} + +/// Verifies that a stored refresh token remains bound to its original issuer. +/// +/// Call this with the same metadata snapshot that RMCP will use for the credentials. Missing or +/// changed issuers require a new login rather than risking sending a refresh token to a different +/// authorization server. +pub(crate) fn validate_refresh_token_issuer( + metadata: &AuthorizationMetadata, + tokens: &StoredOAuthTokens, +) -> Result<()> { + if !tokens.has_refresh_token() { + return Ok(()); + } + + let Some(stored_issuer) = tokens.bound_issuer() else { + return Err(AuthError::AuthorizationRequired).with_context(|| { + format!( + "OAuth refresh credentials for server {} are missing an authorization server issuer; authorization required", + tokens.server_name + ) + }); + }; + + let Some(current_issuer) = metadata + .issuer + .as_deref() + .filter(|issuer| !issuer.trim().is_empty()) + else { + return Err(AuthError::AuthorizationRequired).with_context(|| { + format!( + "OAuth metadata for server {} did not include an authorization server issuer; authorization required", + tokens.server_name + ) + }); + }; + + if current_issuer != stored_issuer { + return Err(AuthError::AuthorizationRequired).with_context(|| { + format!( + "OAuth authorization server issuer changed for server {}; authorization required", + tokens.server_name + ) + }); + } + + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/oauth/refresh_lock.rs b/codex-rs/rmcp-client/src/oauth/refresh_lock.rs new file mode 100644 index 0000000000000000000000000000000000000000..e038dc2a1e34a9d8acfdf5b076664eaf18fdae0c --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/refresh_lock.rs @@ -0,0 +1,104 @@ +//! Cross-process serialization for one MCP OAuth credential's refresh transaction. +//! +//! The guard is intentionally acquired before the authoritative credential reread and retained +//! through provider refresh and persistence. This prevents two processes from replaying the same +//! rotating refresh token or observing a partially persisted transaction. + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use codex_utils_home_dir::find_codex_home; +use sha2::Digest; +use sha2::Sha256; +use std::fs; +use std::fs::File; +use std::fs::OpenOptions; +use std::path::Path; +use std::time::Duration; +use tokio::time::sleep; +use tokio::time::timeout; + +const REFRESH_LOCK_DIR: &str = "mcp-oauth-locks"; +const REFRESH_LOCK_ACQUIRE_TIMEOUT: Duration = Duration::from_secs(/*secs*/ 60); +const REFRESH_LOCK_RETRY_SLEEP: Duration = Duration::from_millis(/*millis*/ 50); +// Keep this internal target stable so diagnostics and cross-process tests can distinguish actual +// WouldBlock contention from a contender that merely started late and observed persisted tokens. +const LOCK_CONTENTION_EVENT_TARGET: &str = "codex_rmcp_client::oauth::refresh_lock::contention"; + +pub(crate) struct RefreshCredentialLock { + _file: File, +} + +impl RefreshCredentialLock { + pub(crate) async fn acquire_for_server(server_name: &str, url: &str) -> Result { + let store_key = super::compute_store_key(server_name, url)?; + let codex_home = find_codex_home()?; + Self::acquire_in(&codex_home, &store_key, REFRESH_LOCK_ACQUIRE_TIMEOUT) + .await + .with_context(|| format!("failed to acquire OAuth credential lock for {server_name}")) + } + + async fn acquire_in( + codex_home: &Path, + store_key: &str, + acquire_timeout: Duration, + ) -> Result { + // Scope coordination to CODEX_HOME alongside File and Secrets state. Direct keyring + // coordination across homes needs a separate cross-platform rendezvous. + // TODO(stevenlee): define that rendezvous before expanding this lock's scope. + let mut hasher = Sha256::new(); + hasher.update(store_key.as_bytes()); + let path = codex_home + .join(REFRESH_LOCK_DIR) + .join(format!("{:x}.lock", hasher.finalize())); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent)?; + } + + let file = OpenOptions::new() + .read(true) + .write(true) + .create(true) + .truncate(false) + .open(&path) + .with_context(|| format!("failed to open OAuth refresh lock {}", path.display()))?; + + // Bound every contender, but keep the acquired lock for the full provider request and + // persistence transaction. Releasing it while awaiting the provider would allow concurrent + // use of a rotating refresh token. + let mut reported_contention = false; + timeout(acquire_timeout, async { + loop { + match file.try_lock() { + Ok(()) => return Ok(()), + Err(std::fs::TryLockError::WouldBlock) => { + if !reported_contention { + tracing::debug!( + target: LOCK_CONTENTION_EVENT_TARGET, + lock_path = %path.display(), + "waiting for another process to finish refreshing MCP OAuth credentials" + ); + reported_contention = true; + } + sleep(REFRESH_LOCK_RETRY_SLEEP).await; + } + Err(error) => return Err(std::io::Error::from(error)), + } + } + }) + .await + .map_err(|_| { + anyhow!( + "timed out after {acquire_timeout:?} waiting for OAuth refresh lock {}", + path.display() + ) + })? + .with_context(|| format!("failed to lock OAuth refresh lock {}", path.display()))?; + + Ok(Self { _file: file }) + } +} + +#[cfg(test)] +#[path = "refresh_lock_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/oauth/refresh_lock_tests.rs b/codex-rs/rmcp-client/src/oauth/refresh_lock_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..6344bd036cbdab80ca1c20e7833889f58664c8e1 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/refresh_lock_tests.rs @@ -0,0 +1,40 @@ +use super::RefreshCredentialLock; +use anyhow::Result; +use std::time::Duration; +use tempfile::tempdir; + +#[tokio::test] +async fn acquisition_times_out_without_stealing() -> Result<()> { + let codex_home = tempdir()?; + let store_key = "test-store-key"; + let held_lock = RefreshCredentialLock::acquire_in( + codex_home.path(), + store_key, + Duration::from_millis(/*millis*/ 100), + ) + .await?; + + let error = RefreshCredentialLock::acquire_in( + codex_home.path(), + store_key, + Duration::from_millis(/*millis*/ 50), + ) + .await + .err() + .expect("contending lock acquisition should time out"); + assert!( + error + .to_string() + .contains("timed out after 50ms waiting for OAuth refresh lock"), + "unexpected error: {error:#}" + ); + + drop(held_lock); + let _reacquired = RefreshCredentialLock::acquire_in( + codex_home.path(), + store_key, + Duration::from_millis(/*millis*/ 100), + ) + .await?; + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/oauth/refresh_transaction.rs b/codex-rs/rmcp-client/src/oauth/refresh_transaction.rs new file mode 100644 index 0000000000000000000000000000000000000000..e4804dd57cd5398878d0fe26df52d77fe8691ebd --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/refresh_transaction.rs @@ -0,0 +1,391 @@ +//! Serialized read-refresh-write transactions for MCP OAuth credentials. + +use std::time::Duration; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +use anyhow::Context; +use anyhow::Error; +use anyhow::Result; +use codex_keyring_store::DefaultKeyringStore; +use codex_keyring_store::KeyringStore; +use oauth2::TokenResponse; +use rmcp::transport::auth::AuthError; +use rmcp::transport::auth::AuthorizationManager; +use rmcp::transport::auth::CredentialStore as _; +use rmcp::transport::auth::InMemoryCredentialStore; +use rmcp::transport::auth::OAuthTokenResponse; +use rmcp::transport::auth::StoredCredentials; +use tokio::time::timeout; +use tracing::debug; +use tracing::warn; + +use super::OAuthPersistor; +use super::OAuthPersistorInner; +use super::StoredOAuthTokens; +use super::WrappedOAuthTokenResponse; +use super::compute_expires_at_millis; +use super::expires_in_from_timestamp; +use super::refresh_lock::RefreshCredentialLock; +use super::token_needs_refresh; +use super::validate_refresh_token_issuer; + +pub(super) const REFRESH_REQUEST_TIMEOUT: Duration = Duration::from_secs(45); + +impl OAuthPersistor { + pub(crate) async fn refresh_if_needed(&self) -> Result<()> { + self.refresh_if_needed_in(&DefaultKeyringStore, REFRESH_REQUEST_TIMEOUT) + .await + } + + /// Injects the credential backend and provider timeout for deterministic failure-path tests. + pub(super) async fn refresh_if_needed_in( + &self, + keyring_store: &K, + refresh_request_timeout: Duration, + ) -> Result<()> { + let expires_at = { + let guard = self.inner.last_credentials.lock().await; + guard.as_ref().and_then(|tokens| tokens.expires_at) + }; + + if !token_needs_refresh(expires_at) { + return Ok(()); + } + + let persistor = self.clone(); + let keyring_store = keyring_store.clone(); + // Once the provider can consume a rotating token, caller cancellation must not cancel + // persistence. The owned task continues with independently bounded lock and request waits. + // A provider timeout leaves the outcome unknown and permits a later serialized retry: + // provider grace may recover, otherwise reauthorization is unavoidable. This residual + // risk is preferred to holding the credential lock indefinitely. + let transaction_task = tokio::spawn(async move { + let result = persistor + .refresh_transaction(&keyring_store, refresh_request_timeout) + .await; + + // Keep this summary inside the owned task so caller cancellation cannot suppress it. + if let Err(error) = &result { + warn!( + server_name = %persistor.inner.server_name, + refresh_reason = "expiry", + error = %error, + "MCP OAuth refresh transaction failed" + ); + } + + result + }); + transaction_task.await.with_context(|| { + format!( + "OAuth refresh task failed for server {}", + self.inner.server_name + ) + })? + } + + #[expect( + clippy::await_holding_invalid_type, + reason = "AuthorizationManager async access must be serialized through its Tokio mutex" + )] + #[tracing::instrument( + level = "debug", + skip_all, + fields( + server_name = %self.inner.server_name, + refresh_reason = "expiry", + ), + err + )] + async fn refresh_transaction( + &self, + keyring_store: &K, + refresh_request_timeout: Duration, + ) -> Result<()> { + debug!("waiting for the MCP OAuth credential transaction lock"); + let _lock = + RefreshCredentialLock::acquire_for_server(&self.inner.server_name, &self.inner.url) + .await?; + debug!("acquired the MCP OAuth credential transaction lock"); + + // Stay on the lifecycle-pinned store. A failure is surfaced rather than falling back and + // possibly replaying an older rotating refresh token from the other store. + debug!("rereading authoritative MCP OAuth credentials"); + let latest = self.inner.credential_store.load( + keyring_store, + &self.inner.server_name, + &self.inner.url, + )?; + + // The pre-lock snapshot is only a hint. This locked reread is authoritative, so adopt a + // winner from another process rather than refreshing its predecessor. + let Some(latest) = latest else { + let manager = self.inner.authorization_manager.clone(); + manager + .lock() + .await + .set_credential_store(InMemoryCredentialStore::new()); + *self.inner.last_credentials.lock().await = None; + return Err(AuthError::AuthorizationRequired).with_context(|| { + format!( + "OAuth tokens for server {} were removed before refresh; authorization required", + self.inner.server_name + ) + }); + }; + + if !token_needs_refresh(latest.expires_at) { + debug!("adopting newer MCP OAuth credentials without contacting the provider"); + let manager = self.inner.authorization_manager.clone(); + let mut guard = manager.lock().await; + if latest.has_refresh_token() { + let previous = self.inner.last_credentials.lock().await; + let expected_issuer = previous.as_ref().and_then(StoredOAuthTokens::bound_issuer); + let latest_issuer = latest.bound_issuer(); + if latest_issuer.is_none() || latest_issuer != expected_issuer { + return Err(AuthError::AuthorizationRequired).with_context(|| { + format!( + "OAuth refresh credentials for server {} could not be bound to the previously validated issuer; authorization required", + self.inner.server_name + ) + }); + } + } + install_tokens_in_manager(&mut guard, &latest).await?; + *self.inner.last_credentials.lock().await = Some(latest); + return Ok(()); + } + + // Without a refresh token, authorization is required before contacting the provider. + if !latest.has_refresh_token() { + return Err(AuthError::AuthorizationRequired).with_context(|| { + format!( + "OAuth tokens for server {} cannot be refreshed; authorization required", + self.inner.server_name + ) + }); + } + + let manager = self.inner.authorization_manager.clone(); + // The provider uses a separate HTTP client and cannot re-enter `AuthClient`. Retain this + // async guard so requests cannot observe credentials while they are staged and committed. + let mut guard = manager.lock().await; + let metadata = guard + .resolve_metadata() + .await + .context("failed to resolve OAuth metadata before using stored refresh credentials")? + .metadata; + validate_refresh_token_issuer(&metadata, &latest)?; + guard.set_metadata(metadata); + install_tokens_in_manager(&mut guard, &latest) + .await + .context("failed to stage OAuth credentials for refresh")?; + // The owned task prevents caller deadlines from canceling after possible token rotation; + // this timeout independently bounds the provider request. + debug!( + timeout_ms = refresh_request_timeout.as_millis(), + "requesting refreshed MCP OAuth credentials from the provider" + ); + let refreshed = match timeout(refresh_request_timeout, guard.refresh_token()).await { + Ok(Ok(token_response)) => { + debug!("received refreshed MCP OAuth credentials from the provider"); + refreshed_tokens(token_response, &latest, &self.inner) + } + Ok(Err(error @ AuthError::TokenRefreshRejected(_))) => { + // Definitive rejection requires authorization even if the access token has not + // expired yet. Other refresh failures below check actual access-token expiry. + warn!( + error = %error, + "MCP OAuth refresh token was rejected; reauthorization required" + ); + return Err(AuthError::AuthorizationRequired).with_context(|| { + format!( + "failed to refresh OAuth tokens for server {}: {error}", + self.inner.server_name + ) + }); + } + Ok(Err(error)) => { + warn!( + error = %error, + "MCP OAuth provider refresh failed" + ); + let error = Error::new(error).context(format!( + "failed to refresh OAuth tokens for server {}", + self.inner.server_name + )); + return self + .recover_after_failed_refresh(keyring_store, &mut guard, &latest, error) + .await; + } + Err(_) => { + warn!( + timeout_ms = refresh_request_timeout.as_millis(), + "MCP OAuth provider refresh timed out; the outcome is unknown and a later serialized retry is permitted" + ); + let error = anyhow::anyhow!( + "timed out after {refresh_request_timeout:?} refreshing OAuth tokens for server {}", + self.inner.server_name + ); + return self + .recover_after_failed_refresh(keyring_store, &mut guard, &latest, error) + .await; + } + }; + + // Persist to the pinned source before exposing the refreshed token. On failure, restore + // the prior in-process credential and return the error; serving an unpersisted token would + // hide the root cause until a later process restart. If the provider already consumed the + // prior token, the next refresh may require reauthorization. That is the deliberate + // fail-closed policy. + // TODO: Add a bounded persistence retry only if telemetry shows this is common; never + // silently switch stores or continue with an unpersisted credential. + debug!("persisting refreshed MCP OAuth credentials to the resolved store"); + if let Err(error) = + self.inner + .credential_store + .save(keyring_store, &self.inner.server_name, &refreshed) + { + warn!( + error = %error, + "failed to persist refreshed MCP OAuth credentials; returning the error and restoring the previous in-process credentials" + ); + install_tokens_in_manager(&mut guard, &latest) + .await + .context( + "failed to restore previous OAuth credentials after refresh persistence failed", + )?; + return Err(error); + } + + // This layer retains RMCP's legacy persistence hook. Install the same merged response + // (including carried-forward refresh token/scopes) so that hook cannot overwrite durable + // credentials with the provider's partial response. + install_tokens_in_manager(&mut guard, &refreshed) + .await + .context( + "refreshed OAuth tokens were persisted but could not be installed in the authorization manager", + )?; + *self.inner.last_credentials.lock().await = Some(refreshed); + drop(guard); + debug!("persisted refreshed MCP OAuth credentials and completed the transaction"); + Ok(()) + } + + async fn recover_after_failed_refresh( + &self, + keyring_store: &K, + manager: &mut AuthorizationManager, + previous: &StoredOAuthTokens, + error: Error, + ) -> Result<()> { + // A failed proactive refresh does not require a new login while the access token is + // still valid. Use actual expiry, not the 30-second refresh buffer. + if !token_has_expired(previous.expires_at) { + return Err(error); + } + + // Browser login can finish while the provider request is pending. Reread the pinned + // authority before prompting, and never delete or overwrite a replacement credential. + let replacement = self.inner.credential_store.load( + keyring_store, + &self.inner.server_name, + &self.inner.url, + )?; + if let Some(replacement) = replacement + && !token_has_expired(replacement.expires_at) + && !replacement.client_id.trim().is_empty() + && !replacement + .token_response + .0 + .access_token() + .secret() + .trim() + .is_empty() + { + // Match the existing pre-refresh adoption rule: refresh credentials must remain + // bound to the issuer already validated for this authorization manager. + if replacement.has_refresh_token() + && (replacement.bound_issuer().is_none() + || replacement.bound_issuer() != previous.bound_issuer()) + { + return Err(AuthError::AuthorizationRequired).context( + "replacement MCP OAuth refresh credentials do not match the validated issuer", + ); + } + debug!("adopting new MCP OAuth credentials after a failed refresh"); + install_tokens_in_manager(manager, &replacement).await?; + *self.inner.last_credentials.lock().await = Some(replacement); + return Ok(()); + } + + warn!("MCP OAuth access token is expired and refresh failed; reauthorization required"); + // Keep AuthorizationRequired as the source so both startup classification and runtime + // tool-call recovery recognize it; retain the original failure as diagnostic context. + Err(Error::new(AuthError::AuthorizationRequired).context(error)) + } +} + +fn token_has_expired(expires_at: Option) -> bool { + expires_at.is_some_and(|expires_at| expires_in_from_timestamp(expires_at).is_none()) +} + +/// Installs tokens without resolving metadata again, so callers can pin the validated snapshot. +pub(crate) async fn install_tokens_in_manager( + authorization_manager: &mut AuthorizationManager, + tokens: &StoredOAuthTokens, +) -> Result<()> { + let store = InMemoryCredentialStore::new(); + let token_response = tokens.token_response.0.clone(); + let granted_scopes = token_response + .scopes() + .map(|scopes| scopes.iter().map(|scope| scope.to_string()).collect()) + .unwrap_or_default(); + let token_received_at = SystemTime::now() + .duration_since(UNIX_EPOCH) + .ok() + .map(|duration| duration.as_secs()); + store + .save( + StoredCredentials::new( + tokens.client_id.clone(), + Some(token_response), + granted_scopes, + token_received_at, + ) + .with_issuer(tokens.issuer.clone()), + ) + .await + .context("failed to stage OAuth tokens for authorization manager")?; + + authorization_manager.set_credential_store(store); + // TODO(stevenlee): Add an RMCP adoption API that atomically updates credentials, client ID, + // and private `current_scopes`; this path cannot synchronize RMCP's scope-upgrade state. + authorization_manager + .initialize_from_store() + .await + .context("failed to adopt refreshed OAuth tokens")?; + Ok(()) +} + +fn refreshed_tokens( + mut token_response: OAuthTokenResponse, + previous: &StoredOAuthTokens, + inner: &OAuthPersistorInner, +) -> StoredOAuthTokens { + if token_response.refresh_token().is_none() { + token_response.set_refresh_token(previous.token_response.0.refresh_token().cloned()); + } + if token_response.scopes().is_none() { + token_response.set_scopes(previous.token_response.0.scopes().cloned()); + } + StoredOAuthTokens { + server_name: inner.server_name.clone(), + url: inner.url.clone(), + issuer: previous.issuer.clone(), + client_id: previous.client_id.clone(), + expires_at: compute_expires_at_millis(&token_response), + token_response: WrappedOAuthTokenResponse(token_response), + } +} diff --git a/codex-rs/rmcp-client/src/oauth/resolved_store.rs b/codex-rs/rmcp-client/src/oauth/resolved_store.rs new file mode 100644 index 0000000000000000000000000000000000000000..904f5d967056998706ee880e3aa07f9c5b2c98a3 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/resolved_store.rs @@ -0,0 +1,238 @@ +//! Resolves the configured MCP OAuth store and pins that concrete source for one client lifecycle. + +use anyhow::Context; +use anyhow::Result; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_keyring_store::KeyringStore; +use tracing::warn; + +use super::OAuthKeyringLoadError; +use super::OAuthStore; +use super::OAuthStoreLock; +use super::OAuthStoreLockFailure; +use super::StoredOAuthTokens; +use super::compute_store_key; +use super::delete_oauth_tokens_from_direct_keyring; +use super::delete_oauth_tokens_from_file; +use super::delete_oauth_tokens_from_secrets_keyring; +use super::load_oauth_tokens_from_file; +use super::load_oauth_tokens_from_file_with_lock_held; +use super::load_oauth_tokens_from_keyring; +use super::load_oauth_tokens_from_secrets_keyring_with_lock_held; +use super::save_oauth_tokens_to_file; +use super::save_oauth_tokens_with_keyring; + +/// Concrete credential store resolved for one MCP OAuth client lifecycle. +/// +/// This is intentionally not durable. `Auto` may resolve differently in a later process, but a +/// client that loaded credentials from one store must reread, refresh, persist, and remove only +/// through that store. A mid-lifecycle backend failure is unexpected and must return an error +/// rather than falling back to another possibly stale refresh token. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ResolvedOAuthCredentialStore { + File, + Keyring(AuthKeyringBackendKind), +} + +impl ResolvedOAuthCredentialStore { + /// Loads credentials only from this already-resolved authority. + /// + /// Unlike `resolve_oauth_tokens_from_store_policy`, this never evaluates configured + /// `Auto` fallback policy. + pub(crate) fn load( + self, + keyring_store: &K, + server_name: &str, + url: &str, + ) -> Result> { + match self { + Self::File => load_oauth_tokens_from_file(server_name, url) + .context("failed to reread OAuth tokens from resolved file storage"), + Self::Keyring(keyring_backend_kind) => load_oauth_tokens_from_keyring( + keyring_store, + keyring_backend_kind, + server_name, + url, + ) + .map_err(anyhow::Error::from) + .context( + "failed to reread OAuth tokens from resolved keyring storage; refusing file fallback", + ), + } + } + + /// Reads the selected authority without waiting for its aggregate-store lock. + pub(crate) fn try_load( + self, + keyring_store: &K, + server_name: &str, + url: &str, + ) -> Result> { + match self { + Self::File => { + let _store_lock = OAuthStoreLock::try_acquire_for_read(OAuthStore::File)?; + load_oauth_tokens_from_file_with_lock_held(server_name, url) + .context("failed to probe OAuth tokens from resolved file storage") + } + Self::Keyring(AuthKeyringBackendKind::Direct) => { + self.load(keyring_store, server_name, url) + } + Self::Keyring(AuthKeyringBackendKind::Secrets) => { + let _store_lock = OAuthStoreLock::try_acquire_for_read(OAuthStore::Secrets)?; + load_oauth_tokens_from_secrets_keyring_with_lock_held( + keyring_store, + server_name, + url, + ) + .map_err(anyhow::Error::from) + } + } + } + + /// Saves credentials only to this already-resolved authority. + pub(crate) fn save( + self, + keyring_store: &K, + server_name: &str, + tokens: &StoredOAuthTokens, + ) -> Result<()> { + match self { + Self::File => save_oauth_tokens_to_file(tokens), + Self::Keyring(keyring_backend_kind) => save_oauth_tokens_with_keyring( + keyring_store, + keyring_backend_kind, + server_name, + tokens, + ), + } + } + + /// Deletes credentials only from this already-resolved authority. + pub(crate) fn delete( + self, + keyring_store: &K, + server_name: &str, + url: &str, + ) -> Result { + match self { + Self::File => { + let key = compute_store_key(server_name, url)?; + delete_oauth_tokens_from_file(&key) + } + Self::Keyring(AuthKeyringBackendKind::Direct) => { + delete_oauth_tokens_from_direct_keyring(keyring_store, server_name, url) + } + Self::Keyring(AuthKeyringBackendKind::Secrets) => { + delete_oauth_tokens_from_secrets_keyring(keyring_store, server_name, url) + } + } + } +} + +#[derive(Debug)] +pub(crate) struct ResolvedOAuthTokens { + pub(crate) tokens: StoredOAuthTokens, + pub(crate) store: ResolvedOAuthCredentialStore, +} + +pub(crate) fn resolve_oauth_tokens_from_store_policy( + keyring_store: &K, + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result> { + match store_mode { + OAuthCredentialsStoreMode::Auto => { + // Auto remains keyring-first at lifecycle startup. The returned source is then pinned + // by the client transport recipe and OAuth persistor so retries, recovery, and + // refresh work cannot hot-switch stores. + // TODO(stevenlee): Different processes can still resolve Auto to different stores + // when keyring availability differs. Solving that safely requires durable backend + // selection or reconciliation of legacy entries and is intentionally outside this + // stack. + match load_oauth_tokens_from_keyring( + keyring_store, + keyring_backend_kind, + server_name, + url, + ) { + Ok(Some(tokens)) => Ok(Some(ResolvedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::Keyring(keyring_backend_kind), + })), + Ok(None) => Ok( + load_oauth_tokens_from_file(server_name, url)?.map(|tokens| { + ResolvedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::File, + } + }), + ), + // Auto may fall back when the keyring backend is unavailable, but a Secrets + // aggregate-lock failure means authority may be changing. Consulting File in + // that state could replay credentials hidden behind a newer Secrets entry. + Err(OAuthKeyringLoadError::StoreLock(error)) => Err(error.into()), + Err(error) => { + warn!("failed to read OAuth tokens from keyring: {error}"); + Ok(load_oauth_tokens_from_file(server_name, url) + .with_context(|| { + format!("failed to read OAuth tokens from keyring: {error}") + })? + .map(|tokens| ResolvedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::File, + })) + } + } + } + OAuthCredentialsStoreMode::File => Ok(load_oauth_tokens_from_file(server_name, url)?.map( + |tokens| ResolvedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::File, + }, + )), + OAuthCredentialsStoreMode::Keyring => Ok(load_oauth_tokens_from_keyring( + keyring_store, + keyring_backend_kind, + server_name, + url, + ) + .map_err(anyhow::Error::from) + .context("failed to read OAuth tokens from keyring")? + .map(|tokens| ResolvedOAuthTokens { + tokens, + store: ResolvedOAuthCredentialStore::Keyring(keyring_backend_kind), + })), + } +} + +pub(crate) fn try_resolve_oauth_tokens_from_store_policy( + keyring_store: &K, + server_name: &str, + url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, +) -> Result> { + let load = |store: ResolvedOAuthCredentialStore| { + store + .try_load(keyring_store, server_name, url) + .map(|tokens| tokens.map(|tokens| ResolvedOAuthTokens { tokens, store })) + }; + let keyring = ResolvedOAuthCredentialStore::Keyring(keyring_backend_kind); + match store_mode { + OAuthCredentialsStoreMode::File => load(ResolvedOAuthCredentialStore::File), + OAuthCredentialsStoreMode::Keyring => load(keyring), + OAuthCredentialsStoreMode::Auto => match load(keyring) { + Ok(Some(tokens)) => Ok(Some(tokens)), + Ok(None) => load(ResolvedOAuthCredentialStore::File), + Err(error) if error.downcast_ref::().is_some() => Err(error), + Err(error) => { + warn!("failed to read OAuth tokens from keyring: {error}"); + load(ResolvedOAuthCredentialStore::File) + .with_context(|| format!("failed to read OAuth tokens from keyring: {error}")) + } + }, + } +} diff --git a/codex-rs/rmcp-client/src/oauth/runtime.rs b/codex-rs/rmcp-client/src/oauth/runtime.rs new file mode 100644 index 0000000000000000000000000000000000000000..fbe43c09fc11f48a8fe4da1cd229738be0372ee3 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/runtime.rs @@ -0,0 +1,64 @@ +//! Keeps one OAuth refresh and persistence owner for the lifetime of a connection. +//! +//! Both modes use cached expiry to prepare credentials before the MCP request budget. +//! Coordinated preparation owns a task so caller cancellation cannot interrupt persistence. +//! It locks the manager before the credential store, matching RMCP's lock order. +//! RMCP commits through the pinned store while Legacy retains Codex's persistor. + +use std::sync::Arc; + +use anyhow::Context; +use anyhow::Result; +use rmcp::transport::auth::AuthorizationManager; +use tokio::sync::Mutex; +use tracing::warn; + +use super::OAuthPersistor; +use super::credential_store::OAuthCredentialStore; +use super::token_needs_refresh; + +#[derive(Clone)] +pub(crate) enum OAuthRuntime { + Legacy(OAuthPersistor), + Coordinated { + auth_manager: Arc>, + store: OAuthCredentialStore, + }, +} + +impl OAuthRuntime { + #[expect( + clippy::await_holding_invalid_type, + reason = "AuthorizationManager access follows RMCP's manager-before-credential lock order" + )] + pub(crate) async fn refresh_if_needed(&self) -> Result<()> { + match self { + Self::Legacy(persistor) => persistor.refresh_if_needed().await, + Self::Coordinated { + auth_manager, + store, + } => { + let expires_at = store + .stored_credentials() + .await + .and_then(|tokens| tokens.expires_at); + if !token_needs_refresh(expires_at) { + return Ok(()); + } + let auth_manager = Arc::clone(auth_manager); + let store = store.clone(); + tokio::spawn(async move { + let mut manager = auth_manager.lock().await; + let result = store.refresh_if_needed(&mut manager).await; + if result.is_err() { + // Keep the summary in the owned task without logging credential/provider data. + warn!("MCP OAuth preparation failed"); + } + result + }) + .await + .context("OAuth refresh task failed")? + } + } + } +} diff --git a/codex-rs/rmcp-client/src/oauth/store_lock.rs b/codex-rs/rmcp-client/src/oauth/store_lock.rs new file mode 100644 index 0000000000000000000000000000000000000000..f06b9b9adc53be373b12d5d97ad1d03d3a7a6f2a --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/store_lock.rs @@ -0,0 +1,212 @@ +//! Cross-process serialization for MCP OAuth stores shared by multiple credentials. +//! +//! File and Secrets each keep credentials for multiple MCP servers in one aggregate document. +//! Writers hold an exclusive lock across the complete read-modify-write operation. Readers share +//! the same lock so concurrent MCP startup and status checks do not serialize behind one another. +//! Direct keyring entries are already stored independently per credential and do not use this lock. + +use std::fs; +use std::fs::File; +use std::fs::OpenOptions; +use std::io; +use std::path::Path; +use std::path::PathBuf; +use std::time::Duration; +use std::time::Instant; + +use codex_utils_home_dir::find_codex_home; + +const OAUTH_LOCK_DIR: &str = "mcp-oauth-locks"; +const STORE_LOCK_ACQUIRE_TIMEOUT: Duration = Duration::from_secs(60); +const STORE_LOCK_RETRY_SLEEP: Duration = Duration::from_millis(50); +// Tests listen for this event so they prove a contender reached the real WouldBlock branch. +const LOCK_CONTENTION_EVENT_TARGET: &str = "codex_rmcp_client::oauth::store_lock::contention"; + +#[derive(Clone, Copy, Debug)] +pub(super) enum OAuthStore { + File, + Secrets, +} + +impl OAuthStore { + fn lock_filename(self) -> &'static str { + match self { + Self::File => "file-store.lock", + Self::Secrets => "secrets-store.lock", + } + } +} + +impl std::fmt::Display for OAuthStore { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::File => f.write_str("fallback file"), + Self::Secrets => f.write_str("encrypted secrets"), + } + } +} + +/// Serializes one complete operation on an aggregate OAuth credential store. +pub(super) struct OAuthStoreLock { + _file: File, +} + +#[derive(Clone, Copy, Debug)] +enum OAuthStoreLockMode { + Shared, + Exclusive, +} + +impl OAuthStoreLock { + pub(super) fn acquire_for_write(store: OAuthStore) -> Result { + Self::acquire_with_timeout( + store, + STORE_LOCK_ACQUIRE_TIMEOUT, + OAuthStoreLockMode::Exclusive, + ) + } + + pub(super) fn acquire_for_read(store: OAuthStore) -> Result { + Self::acquire_with_timeout( + store, + STORE_LOCK_ACQUIRE_TIMEOUT, + OAuthStoreLockMode::Shared, + ) + } + + pub(super) fn try_acquire_for_read(store: OAuthStore) -> Result { + Self::acquire_with_timeout(store, Duration::ZERO, OAuthStoreLockMode::Shared) + } + + fn acquire_with_timeout( + store: OAuthStore, + acquire_timeout: Duration, + mode: OAuthStoreLockMode, + ) -> Result { + // This lock intentionally follows the existing local File/Secrets credential-store + // authority. Those stores are CODEX_HOME-backed today: if CODEX_HOME is unset they use + // the default home (`~/.codex`), and if an embedder has no local home/filesystem authority + // those stores already cannot operate. A future provider-backed credential store should + // provide its own matching lock authority instead of using this local path. + let codex_home = find_codex_home() + .map_err(|source| OAuthStoreLockFailure::CodexHome { store, source })?; + Self::acquire_in_with_mode(&codex_home, store, acquire_timeout, mode) + } + + fn acquire_in_with_mode( + codex_home: &Path, + store: OAuthStore, + acquire_timeout: Duration, + mode: OAuthStoreLockMode, + ) -> Result { + let path = oauth_store_lock_path(codex_home, store); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).map_err(|source| OAuthStoreLockFailure::CreateDir { + store, + path: parent.to_path_buf(), + source, + })?; + } + + let file = OpenOptions::new() + .read(true) + .write(true) + .create(true) + .truncate(false) + .open(&path) + .map_err(|source| OAuthStoreLockFailure::Open { + store, + path: path.clone(), + source, + })?; + let started = Instant::now(); + let mut reported_contention = false; + + loop { + let result = match mode { + OAuthStoreLockMode::Shared => file.try_lock_shared(), + OAuthStoreLockMode::Exclusive => file.try_lock(), + }; + match result { + Ok(()) => return Ok(Self { _file: file }), + Err(std::fs::TryLockError::WouldBlock) if started.elapsed() >= acquire_timeout => { + return Err(OAuthStoreLockFailure::Timeout { + store, + path, + acquire_timeout, + }); + } + Err(std::fs::TryLockError::WouldBlock) => { + if !reported_contention { + tracing::debug!( + target: LOCK_CONTENTION_EVENT_TARGET, + store = %store, + lock_path = %path.display(), + "waiting for another process to finish updating MCP OAuth store state" + ); + reported_contention = true; + } + std::thread::sleep(STORE_LOCK_RETRY_SLEEP.min(acquire_timeout)); + } + Err(error) => { + return Err(OAuthStoreLockFailure::Lock { + store, + path, + source: io::Error::from(error), + }); + } + } + } + } +} + +/// Auto may fall back when the configured keyring backend is unavailable, but it must surface a +/// lock failure. Falling back while another process owns the aggregate-store lock could leave the +/// newer credential in File while a stale Secrets entry remains preferred. +#[derive(Debug, thiserror::Error)] +pub(super) enum OAuthStoreLockFailure { + #[error("failed to resolve CODEX_HOME for MCP OAuth {store} aggregate-store lock")] + CodexHome { + store: OAuthStore, + #[source] + source: io::Error, + }, + #[error("failed to create MCP OAuth {store} aggregate-store lock directory {}", path.display())] + CreateDir { + store: OAuthStore, + path: PathBuf, + #[source] + source: io::Error, + }, + #[error("failed to open MCP OAuth {store} aggregate-store lock {}", path.display())] + Open { + store: OAuthStore, + path: PathBuf, + #[source] + source: io::Error, + }, + #[error( + "timed out after {acquire_timeout:?} waiting for MCP OAuth {store} aggregate-store lock {}", + path.display() + )] + Timeout { + store: OAuthStore, + path: PathBuf, + acquire_timeout: Duration, + }, + #[error("failed to lock MCP OAuth {store} aggregate-store lock {}", path.display())] + Lock { + store: OAuthStore, + path: PathBuf, + #[source] + source: io::Error, + }, +} + +fn oauth_store_lock_path(codex_home: &Path, store: OAuthStore) -> PathBuf { + codex_home.join(OAUTH_LOCK_DIR).join(store.lock_filename()) +} + +#[cfg(test)] +#[path = "tests/store_lock_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/oauth/test_support.rs b/codex-rs/rmcp-client/src/oauth/test_support.rs new file mode 100644 index 0000000000000000000000000000000000000000..208d11b458c621dc5f912e77b248a9f35b0999c0 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/test_support.rs @@ -0,0 +1,69 @@ +use std::path::Path; +use std::sync::Mutex; +use std::sync::MutexGuard; +use std::sync::OnceLock; +use std::sync::PoisonError; + +use tempfile::tempdir; + +/// Build the actual client's Stop-policy route pool before measuring protocol deadlines. +/// Client construction loads platform/custom CA state and can be slow on developer hosts. +/// Use a separate local endpoint so warmup cannot affect the test server's request assertions; +/// the same production HTTP capability still applies all proxy and certificate policy. +pub(crate) async fn warm_http_client( + client: &dyn codex_exec_server::HttpClient, +) -> anyhow::Result<()> { + let server = wiremock::MockServer::start().await; + client + .http_request(codex_exec_server::HttpRequestParams { + method: "GET".to_string(), + url: server.uri(), + headers: Vec::new(), + body: None, + timeout_ms: None, + redirect_policy: codex_exec_server::HttpRedirectPolicy::Stop, + request_id: "test-client-warmup".to_string(), + stream_response: false, + }) + .await?; + Ok(()) +} + +/// Serializes tests that mutate process-wide CODEX_HOME. +/// +/// Keep OAuth tests on this one guard instead of defining per-module helpers; otherwise +/// concurrently running test modules can point File/Secrets storage at different homes. +pub(crate) struct TempCodexHome { + _guard: MutexGuard<'static, ()>, + _dir: tempfile::TempDir, +} + +impl TempCodexHome { + pub(crate) fn new() -> Self { + static LOCK: OnceLock> = OnceLock::new(); + let guard = LOCK + .get_or_init(Mutex::default) + .lock() + .unwrap_or_else(PoisonError::into_inner); + let dir = tempdir().expect("create CODEX_HOME temp dir"); + unsafe { + std::env::set_var("CODEX_HOME", dir.path()); + } + Self { + _guard: guard, + _dir: dir, + } + } + + pub(super) fn path(&self) -> &Path { + self._dir.path() + } +} + +impl Drop for TempCodexHome { + fn drop(&mut self) { + unsafe { + std::env::remove_var("CODEX_HOME"); + } + } +} diff --git a/codex-rs/rmcp-client/src/oauth/tests/credential_store_tests.rs b/codex-rs/rmcp-client/src/oauth/tests/credential_store_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..d8b4e3bdb1b98526b5daf8000e2b68215ea1d076 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/tests/credential_store_tests.rs @@ -0,0 +1,282 @@ +//! Verifies guarded mutations, pinned storage, token mapping, and saved runtime snapshots. + +use std::time::Duration; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +use anyhow::Result; +use codex_config::types::AuthKeyringBackendKind; +use keyring::Error as KeyringError; +use oauth2::AccessToken; +use oauth2::RefreshToken; +use oauth2::TokenResponse; +use pretty_assertions::assert_eq; +use rmcp::transport::auth::AuthError; +use rmcp::transport::auth::CredentialStore; + +use super::MockKeyringStore; +use super::TempCodexHome; +use super::assert_tokens_match_without_expiry; +use super::sample_tokens; +use crate::oauth::OAuthCredentialStore; +use crate::oauth::RefreshCredentialLock; +use crate::oauth::ResolvedOAuthCredentialStore; +use crate::oauth::compute_store_key; +use crate::oauth::load_oauth_tokens_from_file; +use crate::oauth::normalized_oauth_credentials; +use crate::oauth::save_oauth_tokens_to_file; +use crate::oauth::store_lock::OAuthStore; +use crate::oauth::store_lock::OAuthStoreLock; + +#[tokio::test(flavor = "current_thread")] +async fn mutations_require_and_retain_the_transaction_guard() -> Result<()> { + for clear in [false, true] { + let _env = TempCodexHome::new(); + let initial = sample_tokens(); + save_oauth_tokens_to_file(&initial)?; + let store = OAuthCredentialStore::new( + initial.clone(), + ResolvedOAuthCredentialStore::File, + MockKeyringStore::default(), + ); + let guard = store.acquire_transaction_guard().await?; + let credentials = store.load().await?.unwrap(); + if !clear { + store.clear().await?; + } + drop(guard); + let result = if clear { + store.clear().await + } else { + store.save(credentials.clone()).await + }; + assert!( + matches!(result, Err(AuthError::CredentialStoreError(error)) if error.contains("active refresh guard")) + ); + assert_eq!( + load_oauth_tokens_from_file(&initial.server_name, &initial.url)?.is_none(), + !clear, + ); + // Guard acquisition now rereads storage. Restore the deleted credential first, + // then delete it again under the guard so the queued save must recreate it. + save_oauth_tokens_to_file(&initial)?; + let guard = CredentialStore::acquire_refresh_guard(&store) + .await? + .expect("coordinated refresh guard"); + if !clear { + store.clear().await?; + } + let aggregate_lock = OAuthStoreLock::acquire_for_write(OAuthStore::File)?; + let mut mutation = if clear { + store.clear() + } else { + store.save(credentials) + }; + assert!(futures::poll!(&mut mutation).is_pending()); + drop(mutation); + drop(guard); + + // The aggregate lock gates I/O, but only the retained refresh guard blocks this probe. + let mut contender = Box::pin(RefreshCredentialLock::acquire_for_server( + &initial.server_name, + &initial.url, + )); + assert!(futures::poll!(&mut contender).is_pending()); + drop(aggregate_lock); + let _guard = tokio::time::timeout(Duration::from_secs(/*secs*/ 5), contender).await??; + let durable = load_oauth_tokens_from_file(&initial.server_name, &initial.url)?; + assert_eq!(durable.is_none(), clear); + } + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn save_publishes_only_persisted_credentials() -> Result<()> { + for fail_save in [false, true] { + let _env = TempCodexHome::new(); + let initial = sample_tokens(); + let keyring = MockKeyringStore::default(); + let authority = ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct); + authority.save(&keyring, &initial.server_name, &initial)?; + let mut fallback = initial.clone(); + fallback + .token_response + .0 + .set_access_token(AccessToken::new("fallback-token".into())); + save_oauth_tokens_to_file(&fallback)?; + let store = OAuthCredentialStore::new(initial.clone(), authority, keyring.clone()); + let _guard = store.acquire_transaction_guard().await?; + let mut credentials = store.load().await?.expect("stored credentials"); + let token_response = credentials + .token_response + .as_mut() + .expect("stored token response"); + token_response.set_access_token(AccessToken::new("rotated-access".into())); + token_response.set_refresh_token(Some(RefreshToken::new("rotated-refresh".into()))); + // The adapter must encode the separate authoritative grant in the stored response. + token_response.set_scopes(/*scopes*/ None); + if fail_save { + let key = compute_store_key(&initial.server_name, &initial.url)?; + keyring.set_error(&key, KeyringError::Invalid("test".into(), "save".into())); + } + let result = store.save(credentials).await; + let durable = authority + .load(&keyring, &initial.server_name, &initial.url)? + .expect("durable keyring credentials"); + let mut expected = initial.clone(); + if fail_save { + assert!(matches!(result, Err(AuthError::CredentialStoreError(_)))); + } else { + result?; + expected + .token_response + .0 + .set_access_token(AccessToken::new("rotated-access".into())); + expected + .token_response + .0 + .set_refresh_token(Some(RefreshToken::new("rotated-refresh".into()))); + expected.expires_at = durable.expires_at; + } + assert_tokens_match_without_expiry(&durable, &expected); + assert_eq!( + store.stored_credentials().await, + normalized_oauth_credentials(Some(&expected)) + ); + assert_tokens_match_without_expiry( + &load_oauth_tokens_from_file(&initial.server_name, &initial.url)?.unwrap(), + &fallback, + ); + } + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn pinned_read_failure_does_not_adopt_fallback_credentials() -> Result<()> { + let _env = TempCodexHome::new(); + let initial = sample_tokens(); + save_oauth_tokens_to_file(&initial)?; + let keyring = MockKeyringStore::default(); + let key = compute_store_key(&initial.server_name, &initial.url)?; + keyring.set_error(&key, KeyringError::Invalid("test".into(), "load".into())); + let store = OAuthCredentialStore::new( + initial.clone(), + ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct), + keyring, + ); + let cached = store.load().await?.expect("cached credentials"); + assert_eq!( + cached.token_response.unwrap().access_token().secret(), + initial.token_response.0.access_token().secret() + ); + assert!(matches!( + store.acquire_transaction_guard().await, + Err(AuthError::CredentialStoreError(_)) + )); + assert_eq!( + store.stored_credentials().await, + normalized_oauth_credentials(Some(&initial)) + ); + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn replacement_or_removal_does_not_acknowledge_a_new_runtime_snapshot() -> Result<()> { + let _env = TempCodexHome::new(); + let initial = sample_tokens(); + save_oauth_tokens_to_file(&initial)?; + let store = OAuthCredentialStore::new( + initial.clone(), + ResolvedOAuthCredentialStore::File, + MockKeyringStore::default(), + ); + let original_snapshot = store.stored_credentials().await; + for (client_id, issuer) in [ + ("replacement-client", initial.issuer.clone()), + ( + initial.client_id.as_str(), + Some("https://replacement.example.test".into()), + ), + (initial.client_id.as_str(), None), + ] { + let mut replacement = initial.clone(); + replacement.client_id = client_id.into(); + replacement.issuer = issuer; + save_oauth_tokens_to_file(&replacement)?; + assert!(matches!( + store.acquire_transaction_guard().await, + Err(AuthError::AuthorizationRequired) + )); + assert_eq!(store.stored_credentials().await, original_snapshot); + } + save_oauth_tokens_to_file(&initial)?; + let guard = store.acquire_transaction_guard().await?; + store.clear().await?; + assert!(load_oauth_tokens_from_file(&initial.server_name, &initial.url)?.is_none()); + drop(guard); + assert!(matches!( + store.acquire_transaction_guard().await, + Err(AuthError::AuthorizationRequired) + )); + assert!(store.load().await?.is_some()); + assert_eq!(store.stored_credentials().await, original_snapshot); + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn storage_roundtrip_preserves_absolute_and_unknown_expiry() -> Result<()> { + let _env = TempCodexHome::new(); + let initial = sample_tokens(); + save_oauth_tokens_to_file(&initial)?; + let store = OAuthCredentialStore::new( + initial.clone(), + ResolvedOAuthCredentialStore::File, + MockKeyringStore::default(), + ); + let future_received_at = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs(); + let future_deadline = (future_received_at + 120) * 1000; + for (received_at, expected_deadline) in [ + (Some(future_received_at), Some(future_deadline)), + (None, None), + ] { + let guard = store.acquire_transaction_guard().await?; + let mut credentials = store.load().await?.unwrap(); + credentials.token_received_at = received_at; + credentials + .token_response + .as_mut() + .unwrap() + .set_expires_in(Some(&Duration::from_secs(/*secs*/ 120))); + store.save(credentials).await?; + let durable = load_oauth_tokens_from_file(&initial.server_name, &initial.url)?.unwrap(); + let mut expected = initial.clone(); + expected.expires_at = expected_deadline; + expected + .token_response + .0 + .set_expires_in(expected_deadline.map(|_| Duration::ZERO).as_ref()); + assert_tokens_match_without_expiry(&durable, &expected); + drop(guard); + let _guard = store.acquire_transaction_guard().await?; + let reloaded = store.load().await?.unwrap(); + let expires_in = reloaded.token_response.unwrap().expires_in(); + if let Some(expected_deadline) = expected_deadline { + let expires_in = expires_in.expect("future token expiry"); + assert!(!expires_in.is_zero()); + let reconstructed_deadline = + Duration::from_secs(reloaded.token_received_at.expect("token receipt time")) + + expires_in; + let expected_deadline = Duration::from_millis(expected_deadline); + // Receipt times use whole seconds, so reconstruction may lose less than a second. + assert!(reconstructed_deadline <= expected_deadline); + assert!(expected_deadline - reconstructed_deadline < Duration::from_secs(/*secs*/ 1)); + } else { + assert_eq!(expires_in, None); + } + assert_eq!( + store.stored_credentials().await, + normalized_oauth_credentials(Some(&expected)) + ); + } + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/oauth/tests/persistor_tests.rs b/codex-rs/rmcp-client/src/oauth/tests/persistor_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..057fb7c6fdbd1805f630c9c96b7c1d08f4e9a3cb --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/tests/persistor_tests.rs @@ -0,0 +1,977 @@ +use std::sync::Arc; +use std::sync::mpsc; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::RouteAwareHttpClient; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_keyring_store::DefaultKeyringStore; +use http::HeaderMap; +use keyring::Error as KeyringError; +use oauth2::AccessToken; +use oauth2::RefreshToken; +use oauth2::TokenResponse; +use pretty_assertions::assert_eq; +use rmcp::transport::auth::AuthClient; +use rmcp::transport::auth::AuthError; +use rmcp::transport::auth::AuthorizationManager; +use rmcp::transport::auth::OAuthState; +use rmcp::transport::streamable_http_client::StreamableHttpClient; +use tokio::sync::Mutex as TokioMutex; +use tracing::Event; +use tracing::Id; +use tracing::Metadata; +use tracing::Subscriber; +use tracing::span::Attributes; +use tracing::span::Record; +use tracing::subscriber::Interest; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::body_string_contains; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::MockKeyringStore; +use super::TempCodexHome; +use super::assert_tokens_match_without_expiry; +use super::sample_tokens; +use crate::http_client_adapter::StreamableHttpClientAdapter; +use crate::http_client_adapter::StreamableHttpRedirectMode; +use crate::oauth::OAuthCredentialStore; +use crate::oauth::OAuthPersistor; +use crate::oauth::OAuthRuntime; +use crate::oauth::ResolvedOAuthCredentialStore; +use crate::oauth::StoredOAuthTokens; +use crate::oauth::WrappedOAuthTokenResponse; +use crate::oauth::compute_expires_at_millis; +use crate::oauth::compute_store_key; +use crate::oauth::delete_oauth_tokens; +use crate::oauth::load_oauth_tokens_from_file; +use crate::oauth::refresh_lock::RefreshCredentialLock; +use crate::oauth::save_oauth_tokens; +use crate::oauth::save_oauth_tokens_to_file; +use crate::oauth::stored_oauth_credentials; +use crate::oauth_http_client::OAuthHttpClientAdapter; +use crate::oauth_http_client::PROACTIVE_REFRESH_TIMEOUT; +use crate::startup_error::is_authentication_required_error; + +const REFRESH_LOCK_CONTENTION_EVENT_TARGET: &str = + "codex_rmcp_client::oauth::refresh_lock::contention"; + +#[tokio::test(flavor = "current_thread")] +async fn login_and_logout_follow_the_completed_refresh() -> Result<()> { + let (_env, _server, initial) = test_context().await?; + let mut replacement = initial.clone(); + replacement + .token_response + .0 + .set_access_token(AccessToken::new("replacement-access-token".into())); + replacement + .token_response + .0 + .set_refresh_token(Some(RefreshToken::new("replacement-refresh-token".into()))); + let mut refreshed = initial.clone(); + refreshed + .token_response + .0 + .set_refresh_token(Some(RefreshToken::new("rotated-refresh-token".into()))); + + for expected in [Some(replacement), None] { + save_oauth_tokens_to_file(&initial)?; + let held_lock = + RefreshCredentialLock::acquire_for_server(&initial.server_name, &initial.url).await?; + let (contended_tx, contended_rx) = mpsc::channel(); + let _subscriber_guard = + tracing::subscriber::set_default(LockContentionSubscriber { contended_tx }); + let mutation = tokio::spawn({ + let initial = initial.clone(); + let expected = expected.clone(); + async move { + match expected { + Some(tokens) => { + save_oauth_tokens( + &tokens.server_name, + &tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + ) + .await + } + None => { + assert!( + delete_oauth_tokens( + &initial.server_name, + &initial.url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + ) + .await? + ); + Ok(()) + } + } + } + }); + wait_for_lock_contention(contended_rx, /*expected_count*/ 1).await?; + assert!(!mutation.is_finished()); + assert_eq!( + load_oauth_tokens_from_file(&initial.server_name, &initial.url)?, + Some(initial.clone()) + ); + // Complete the in-flight refresh before allowing the login or logout to write. + save_oauth_tokens_to_file(&refreshed)?; + drop(held_lock); + mutation.await??; + assert_eq!( + load_oauth_tokens_from_file(&initial.server_name, &initial.url)?, + expected + ); + } + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn fresh_preparation_skips_storage_and_refresh_lock() -> Result<()> { + let (env, server, mut initial) = test_context().await?; + let fresh = sample_tokens(); + initial.token_response = fresh.token_response; + initial.expires_at = fresh.expires_at; + std::fs::create_dir(env.path().join(".credentials.json"))?; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .respond_with(ResponseTemplate::new(500)) + .expect(0) + .mount(&server) + .await; + + for mode in [ + crate::McpOAuthRefreshMode::Legacy, + crate::McpOAuthRefreshMode::Coordinated, + ] { + let runtime = runtime_for(&initial, mode).await?; + let held_lock = + RefreshCredentialLock::acquire_for_server(&initial.server_name, &initial.url).await?; + tokio::time::timeout( + Duration::from_millis(/*millis*/ 100), + runtime.refresh_if_needed(), + ) + .await + .context("fresh preparation must not wait for the refresh lock")??; + drop(held_lock); + } + server.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn coordinated_401_refresh_rereads_and_persists_before_retry() -> Result<()> { + let (_env, server, mut initial) = test_context().await?; + let fresh = sample_tokens(); + initial.token_response = fresh.token_response; + initial.expires_at = fresh.expires_at; + save_oauth_tokens_to_file(&initial)?; + let mut manager = authorization_manager_for(&initial).await?; + manager.set_credential_store(OAuthCredentialStore::new( + initial.clone(), + ResolvedOAuthCredentialStore::File, + DefaultKeyringStore, + )); + let client = AuthClient::new( + StreamableHttpClientAdapter::new( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + HeaderMap::new(), + /*auth_provider*/ None, + /*has_configured_headers*/ false, + StreamableHttpRedirectMode::Legacy, + Arc::default(), + ), + manager, + ); + let mut durable = initial.clone(); + durable + .token_response + .0 + .set_access_token(AccessToken::new("durable-access-token".into())); + durable + .token_response + .0 + .set_refresh_token(Some(RefreshToken::new("durable-refresh-token".into()))); + save_oauth_tokens_to_file(&durable)?; + + Mock::given(method("POST")) + .and(path("/mcp")) + .and(header("authorization", "Bearer access-token")) + .respond_with(ResponseTemplate::new(401).insert_header("www-authenticate", "Bearer")) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("refresh_token=durable-refresh-token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "refreshed-access-token", + "refresh_token": "rotated-refresh-token", + "token_type": "Bearer", + "expires_in": 3600, + }))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/mcp")) + .and(header("authorization", "Bearer refreshed-access-token")) + .respond_with(move |_: &wiremock::Request| { + let saved = load_oauth_tokens_from_file(&durable.server_name, &durable.url) + .expect("read persisted refresh before retry") + .expect("refresh must be persisted before retry"); + let mut expected = durable.token_response.0.clone(); + expected.set_access_token(AccessToken::new("refreshed-access-token".into())); + expected.set_refresh_token(Some(RefreshToken::new("rotated-refresh-token".into()))); + expected.set_expires_in(saved.token_response.0.expires_in().as_ref()); + assert_eq!(saved.token_response, WrappedOAuthTokenResponse(expected)); + ResponseTemplate::new(202) + }) + .expect(1) + .mount(&server) + .await; + client + .post_message( + initial.url.into(), + serde_json::from_value( + serde_json::json!({"jsonrpc": "2.0", "id": 1, "method": "ping"}), + )?, + /*session_id*/ None, + /*auth_token*/ None, + Default::default(), + ) + .await?; + server.verify().await; + Ok(()) +} + +struct LockContentionSubscriber { + contended_tx: mpsc::Sender<()>, +} + +impl Subscriber for LockContentionSubscriber { + fn enabled(&self, metadata: &Metadata<'_>) -> bool { + metadata.target() == REFRESH_LOCK_CONTENTION_EVENT_TARGET + } + + fn register_callsite(&self, metadata: &'static Metadata<'static>) -> Interest { + if self.enabled(metadata) { + Interest::always() + } else { + Interest::never() + } + } + + fn max_level_hint(&self) -> Option { + Some(tracing::level_filters::LevelFilter::DEBUG) + } + + fn new_span(&self, _span: &Attributes<'_>) -> Id { + Id::from_u64(/*u*/ 1) + } + + fn record(&self, _span: &Id, _values: &Record<'_>) {} + + fn record_follows_from(&self, _span: &Id, _follows_from: &Id) {} + + fn event(&self, event: &Event<'_>) { + if self.enabled(event.metadata()) { + self.contended_tx + .send(()) + .expect("signal actual OAuth credential-lock contention"); + } + } + + fn enter(&self, _span: &Id) {} + + fn exit(&self, _span: &Id) {} +} + +#[tokio::test(flavor = "current_thread")] +async fn concurrent_refreshes_call_provider_once_and_carry_omitted_fields() -> Result<()> { + assert_concurrent_refreshes(crate::McpOAuthRefreshMode::Legacy).await?; + assert_concurrent_refreshes(crate::McpOAuthRefreshMode::Coordinated).await +} + +async fn assert_concurrent_refreshes(mode: crate::McpOAuthRefreshMode) -> Result<()> { + let (_env, server, initial) = test_context().await?; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains("refresh_token=refresh-token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "refreshed-access-token", + "token_type": "Bearer", + "expires_in": 3600, + }))) + .expect(1) + .mount(&server) + .await; + save_oauth_tokens_to_file(&initial)?; + + // Hold the real credential lock until both refresh transactions report WouldBlock. This makes + // the lock assertion independent of task scheduling and ensures removing transaction locking + // makes the test fail before either request can reach the provider. + let held_lock = + RefreshCredentialLock::acquire_for_server(&initial.server_name, &initial.url).await?; + let (contended_tx, contended_rx) = mpsc::channel(); + let _subscriber_guard = + tracing::subscriber::set_default(LockContentionSubscriber { contended_tx }); + + let first = runtime_for(&initial, mode).await?; + let second = runtime_for(&initial, mode).await?; + let first_task = tokio::spawn({ + let first = first.clone(); + async move { first.refresh_if_needed().await } + }); + let second_task = tokio::spawn({ + let second = second.clone(); + async move { second.refresh_if_needed().await } + }); + + wait_for_lock_contention(contended_rx, /*expected_count*/ 2).await?; + drop(held_lock); + first_task.await??; + second_task.await??; + server.verify().await; + + // Layer 2 still invokes the legacy RMCP persistence hook after operations. Exercise that hook + // so a raw provider response that omitted refresh token/scopes cannot overwrite the merged + // authoritative credential. + if let OAuthRuntime::Legacy(persistor) = &first { + persistor.persist_if_needed().await?; + } + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)? + .expect("refreshed credentials should be stored"); + assert_eq!(stored.issuer, initial.issuer); + let live_credentials = match &first { + OAuthRuntime::Legacy(persistor) => persistor.stored_credentials().await, + OAuthRuntime::Coordinated { store, .. } => store.stored_credentials().await, + }; + let disk_credentials = stored_oauth_credentials( + &initial.server_name, + &initial.url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + )?; + assert_eq!(live_credentials, disk_credentials); + let mut expected_response = initial.token_response.0.clone(); + expected_response.set_access_token(AccessToken::new("refreshed-access-token".to_string())); + // File loads derive `expires_in` from stable `expires_at`, so it may tick down before this + // assertion. Normalize only that derived field and compare the complete token response so + // omitted refresh-token and scope carry-forward remain covered. + expected_response.set_expires_in(stored.token_response.0.expires_in().as_ref()); + assert_eq!( + stored.token_response, + WrappedOAuthTokenResponse(expected_response) + ); + Ok(()) +} + +#[expect( + clippy::await_holding_invalid_type, + reason = "AuthorizationManager async access must be serialized through its Tokio mutex" +)] +#[tokio::test(flavor = "current_thread")] +async fn resolved_keyring_read_error_preserves_in_memory_credentials() -> Result<()> { + let (_env, _server, initial) = test_context().await?; + let keyring_store = MockKeyringStore::default(); + let key = compute_store_key(&initial.server_name, &initial.url)?; + keyring_store.set_error(&key, KeyringError::Invalid("error".into(), "load".into())); + let manager = Arc::new(TokioMutex::new(authorization_manager_for(&initial).await?)); + let persistor = OAuthPersistor::new( + initial.server_name.clone(), + initial.url.clone(), + Arc::clone(&manager), + ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct), + Some(initial.clone()), + ); + + let error = persistor + .refresh_if_needed_in(&keyring_store, Duration::from_secs(/*secs*/ 45)) + .await + .expect_err("the resolved keyring read error should abort refresh"); + assert!( + error + .to_string() + .contains("failed to reread OAuth tokens from resolved keyring storage"), + "unexpected error: {error:#}" + ); + let guard = manager.lock().await; + let (_client_id, token_response) = guard.get_credentials().await?; + assert_eq!( + WrappedOAuthTokenResponse(token_response.expect("manager should retain credentials")), + initial.token_response + ); + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn missing_authoritative_credentials_require_reauthorization() -> Result<()> { + let (_env, _server, initial) = test_context().await?; + let persistor = persistor_for(&initial).await?; + + let error = persistor + .refresh_if_needed() + .await + .expect_err("a removed authoritative credential should abort refresh"); + assert!(error.chain().any(|source| matches!( + source.downcast_ref::(), + Some(AuthError::AuthorizationRequired) + ))); + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn rejected_refresh_token_requires_reauthorization() -> Result<()> { + let (_env, server, initial) = test_context().await?; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains("refresh_token=refresh-token")) + .respond_with(ResponseTemplate::new(400).set_body_json(serde_json::json!({ + "error": "invalid_grant", + "error_description": "refresh token expired or revoked", + }))) + .expect(1) + .mount(&server) + .await; + save_oauth_tokens_to_file(&initial)?; + let persistor = persistor_for(&initial).await?; + + let error = persistor + .refresh_if_needed() + .await + .expect_err("a provider-rejected refresh token should require reauthorization"); + assert!(is_authentication_required_error(&error)); + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)? + .expect("rejected refresh must preserve the durable credentials"); + assert_tokens_match_without_expiry(&stored, &initial); + server.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn changed_issuer_requires_reauthorization_before_refresh() -> Result<()> { + let (_env, server, mut initial) = test_context().await?; + initial.issuer = Some("https://original-issuer.example.test".to_string()); + Mock::given(method("POST")) + .and(path("/oauth/token")) + .respond_with(ResponseTemplate::new(500)) + .expect(0) + .mount(&server) + .await; + save_oauth_tokens_to_file(&initial)?; + let persistor = persistor_for(&initial).await?; + + let error = persistor + .refresh_if_needed() + .await + .expect_err("a changed issuer must abort refresh"); + assert!(is_authentication_required_error(&error)); + server.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn missing_issuer_requires_reauthorization_before_refresh() -> Result<()> { + let (_env, server, mut initial) = test_context().await?; + initial.issuer = None; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .respond_with(ResponseTemplate::new(500)) + .expect(0) + .mount(&server) + .await; + save_oauth_tokens_to_file(&initial)?; + let persistor = persistor_for(&initial).await?; + + let error = persistor + .refresh_if_needed() + .await + .expect_err("a missing issuer must abort refresh"); + assert!(is_authentication_required_error(&error)); + server.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn issuerless_newer_credentials_are_not_adopted_before_refresh() -> Result<()> { + let (_env, server, initial) = test_context().await?; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .respond_with(ResponseTemplate::new(500)) + .expect(0) + .mount(&server) + .await; + save_oauth_tokens_to_file(&initial)?; + let persistor = persistor_for(&initial).await?; + + let mut latest = initial.clone(); + latest.issuer = None; + latest.expires_at = Some(u64::MAX); + save_oauth_tokens_to_file(&latest)?; + + let error = persistor + .refresh_if_needed() + .await + .expect_err("an issuer-less refresh token must not be adopted"); + assert!(is_authentication_required_error(&error)); + server.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn proactive_refresh_failure_with_unexpired_token_does_not_require_reauthorization() +-> Result<()> { + let (_env, server, mut initial) = test_context().await?; + initial + .token_response + .0 + .set_expires_in(Some(&Duration::from_secs(/*secs*/ 30))); + initial.expires_at = compute_expires_at_millis(&initial.token_response.0); + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains("refresh_token=refresh-token")) + .respond_with(ResponseTemplate::new(503).set_body_json(serde_json::json!({ + "error": "temporarily_unavailable", + "error_description": "provider is temporarily unavailable", + }))) + .expect(1) + .mount(&server) + .await; + save_oauth_tokens_to_file(&initial)?; + let persistor = persistor_for(&initial).await?; + + let error = persistor + .refresh_if_needed() + .await + .expect_err("a transient provider failure should not erase valid credentials"); + assert!(!is_authentication_required_error(&error)); + assert!(error.chain().any(|source| matches!( + source.downcast_ref::(), + Some(AuthError::TokenRefreshFailed(_)) + ))); + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)? + .expect("a transient refresh failure must preserve durable credentials"); + assert_tokens_match_without_expiry(&stored, &initial); + server.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn caller_cancellation_does_not_cancel_refresh_persistence() -> Result<()> { + assert_caller_cancellation(crate::McpOAuthRefreshMode::Legacy).await?; + assert_caller_cancellation(crate::McpOAuthRefreshMode::Coordinated).await +} + +async fn assert_caller_cancellation(mode: crate::McpOAuthRefreshMode) -> Result<()> { + let (_env, server, initial) = test_context().await?; + let (request_received_tx, request_received_rx) = mpsc::channel(); + let (release_response_tx, release_response_rx) = mpsc::channel(); + let release_response_rx = Arc::new(std::sync::Mutex::new(release_response_rx)); + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains("refresh_token=refresh-token")) + .respond_with(move |_request: &wiremock::Request| { + request_received_tx + .send(()) + .expect("signal OAuth refresh request"); + release_response_rx + .lock() + .expect("lock OAuth refresh response gate") + .recv() + .expect("wait to release OAuth refresh response"); + ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "cancel-safe-access-token", + "token_type": "Bearer", + "expires_in": 3600, + })) + }) + .expect(1) + .mount(&server) + .await; + save_oauth_tokens_to_file(&initial)?; + let runtime = runtime_for(&initial, mode).await?; + let caller = tokio::spawn(async move { runtime.refresh_if_needed().await }); + + tokio::task::spawn_blocking(move || { + request_received_rx + .recv_timeout(Duration::from_secs(/*secs*/ 5)) + .context("timed out waiting for OAuth refresh request") + }) + .await??; + caller.abort(); + assert!( + caller + .await + .expect_err("caller should be cancelled") + .is_cancelled() + ); + + release_response_tx + .send(()) + .context("release OAuth refresh response")?; + + // Reacquiring the same credential lock waits for the detached refresh task to persist and + // release it, avoiding a scheduler-sensitive sleep after cancellation. + let _lock = tokio::time::timeout( + Duration::from_secs(/*secs*/ 2), + RefreshCredentialLock::acquire_for_server(&initial.server_name, &initial.url), + ) + .await + .context("detached refresh did not release its credential lock")??; + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)? + .expect("detached refresh should persist credentials"); + assert_eq!( + stored.token_response.0.access_token().secret(), + "cancel-safe-access-token" + ); + server.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn failed_refresh_adopts_login_completed_during_request() -> Result<()> { + for expires_in in [Some(Duration::from_secs(/*secs*/ 3600)), None] { + let (_env, server, initial) = test_context().await?; + let mut replacement = initial.clone(); + replacement.client_id = "new-login-client".to_string(); + replacement + .token_response + .0 + .set_access_token(AccessToken::new("new-login-access-token".to_string())); + replacement + .token_response + .0 + .set_expires_in(expires_in.as_ref()); + replacement.expires_at = compute_expires_at_millis(&replacement.token_response.0); + save_oauth_tokens_to_file(&initial)?; + let persistor = persistor_for(&initial).await?; + let provider_replacement = replacement.clone(); + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("grant_type=refresh_token")) + .respond_with(move |_request: &wiremock::Request| { + // A browser login writes independently of the refresh transaction lock. + save_oauth_tokens_to_file(&provider_replacement).expect("complete new login"); + ResponseTemplate::new(200).set_body_string("not an OAuth token response") + }) + .expect(1) + .mount(&server) + .await; + + persistor.refresh_if_needed().await?; + let live = persistor.stored_credentials().await.expect("adopted login"); + let mut expected_live = replacement.clone(); + expected_live.token_response.0.set_expires_in(None); + assert_eq!(live, expected_live); + persistor.persist_if_needed().await?; + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)? + .expect("new login must remain stored"); + assert_tokens_match_without_expiry(&stored, &replacement); + server.verify().await; + } + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn failed_refresh_does_not_adopt_unbound_replacement_credentials() -> Result<()> { + for issuer in [None, Some("https://different-issuer.example.test")] { + let (_env, server, initial) = test_context().await?; + let mut replacement = initial.clone(); + replacement.issuer = issuer.map(str::to_string); + replacement.expires_at = Some(u64::MAX); + replacement + .token_response + .0 + .set_access_token(AccessToken::new("replacement-access-token".to_string())); + save_oauth_tokens_to_file(&initial)?; + let persistor = persistor_for(&initial).await?; + let provider_replacement = replacement.clone(); + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("grant_type=refresh_token")) + .respond_with(move |_request: &wiremock::Request| { + save_oauth_tokens_to_file(&provider_replacement).expect("complete new login"); + ResponseTemplate::new(200).set_body_string("invalid JSON") + }) + .expect(1) + .mount(&server) + .await; + + let error = persistor + .refresh_if_needed() + .await + .expect_err("an unbound replacement must not enter the authorization manager"); + assert!(is_authentication_required_error(&error)); + let live = persistor + .stored_credentials() + .await + .expect("original credentials"); + let mut expected_live = initial.clone(); + expected_live.token_response.0.set_expires_in(None); + assert_eq!(live, expected_live); + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)? + .expect("the replacement must not be deleted or overwritten"); + assert_tokens_match_without_expiry(&stored, &replacement); + server.verify().await; + } + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn expired_refresh_failure_preserves_credentials_for_a_later_retry() -> Result<()> { + let (_env, server, initial) = test_context().await?; + save_oauth_tokens_to_file(&initial)?; + let persistor = persistor_for(&initial).await?; + let failure = Mock::given(method("POST")) + .and(path("/oauth/token")) + .respond_with(ResponseTemplate::new(200).set_body_string("invalid JSON")) + .expect(1) + .mount_as_scoped(&server) + .await; + + let error = persistor + .refresh_if_needed() + .await + .expect_err("refresh failed"); + assert!(is_authentication_required_error(&error)); + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)? + .expect("failed refresh must preserve stored credentials"); + assert_tokens_match_without_expiry(&stored, &initial); + drop(failure); + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("refresh_token=refresh-token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "retry-access-token", + "token_type": "Bearer", + "expires_in": 3600, + }))) + .expect(1) + .mount(&server) + .await; + + persistor.refresh_if_needed().await?; + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)? + .expect("successful retry must persist new credentials"); + assert_eq!( + stored.token_response.0.access_token().secret(), + "retry-access-token" + ); + server.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn provider_timeout_releases_lock_and_preserves_durable_credentials() -> Result<()> { + let (_env, server, initial) = test_context().await?; + mount_delayed_refresh(&server, "late-access-token").await; + save_oauth_tokens_to_file(&initial)?; + let persistor = persistor_for(&initial).await?; + + let error = persistor + .refresh_if_needed_in( + &MockKeyringStore::default(), + Duration::from_millis(/*millis*/ 50), + ) + .await + .expect_err("provider request should reach its explicit timeout"); + assert!(error.to_string().contains("timed out after 50ms")); + assert!(is_authentication_required_error(&error)); + + let _lock = tokio::time::timeout( + Duration::from_millis(/*millis*/ 100), + RefreshCredentialLock::acquire_for_server(&initial.server_name, &initial.url), + ) + .await + .context("provider timeout did not release the credential lock")??; + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)? + .expect("timed-out refresh must leave durable credentials present"); + assert_tokens_match_without_expiry(&stored, &initial); + server.verify().await; + Ok(()) +} + +#[expect( + clippy::await_holding_invalid_type, + reason = "AuthorizationManager async access must be serialized through its Tokio mutex" +)] +#[tokio::test(flavor = "current_thread")] +async fn coordinated_provider_timeout_excludes_lock_wait() -> Result<()> { + let (_env, server, initial) = test_context().await?; + mount_delayed_refresh(&server, "late-access-token").await; + save_oauth_tokens_to_file(&initial)?; + let (auth_manager, _) = coordinated_manager_for(&initial).await?; + let held_lock = + RefreshCredentialLock::acquire_for_server(&initial.server_name, &initial.url).await?; + let (contended_tx, contended_rx) = mpsc::channel(); + let _subscriber_guard = + tracing::subscriber::set_default(LockContentionSubscriber { contended_tx }); + let mut refresh = tokio::spawn( + PROACTIVE_REFRESH_TIMEOUT.scope(Duration::from_millis(/*millis*/ 50), async move { + auth_manager.lock().await.get_access_token().await + }), + ); + wait_for_lock_contention(contended_rx, /*expected_count*/ 1).await?; + assert!( + tokio::time::timeout(Duration::from_millis(/*millis*/ 100), &mut refresh) + .await + .is_err(), + "lock wait must not consume the provider timeout" + ); + drop(held_lock); + let error = refresh + .await? + .expect_err("the scoped provider request must time out"); + assert!(matches!(error, AuthError::TokenRefreshFailed(_))); + let _lock = tokio::time::timeout( + Duration::from_millis(/*millis*/ 100), + RefreshCredentialLock::acquire_for_server(&initial.server_name, &initial.url), + ) + .await??; + let stored = load_oauth_tokens_from_file(&initial.server_name, &initial.url)?.unwrap(); + assert_tokens_match_without_expiry(&stored, &initial); + server.verify().await; + Ok(()) +} + +async fn runtime_for( + tokens: &StoredOAuthTokens, + mode: crate::McpOAuthRefreshMode, +) -> Result { + Ok(match mode { + crate::McpOAuthRefreshMode::Legacy => OAuthRuntime::Legacy(persistor_for(tokens).await?), + crate::McpOAuthRefreshMode::Coordinated => { + let (auth_manager, store) = coordinated_manager_for(tokens).await?; + OAuthRuntime::Coordinated { + auth_manager, + store, + } + } + }) +} + +async fn coordinated_manager_for( + tokens: &StoredOAuthTokens, +) -> Result<(Arc>, OAuthCredentialStore)> { + let mut manager = authorization_manager_for(tokens).await?; + let store = OAuthCredentialStore::new( + tokens.clone(), + ResolvedOAuthCredentialStore::File, + DefaultKeyringStore, + ); + manager.set_credential_store(store.clone()); + Ok((Arc::new(TokioMutex::new(manager)), store)) +} + +async fn persistor_for(tokens: &StoredOAuthTokens) -> Result { + Ok(OAuthPersistor::new( + tokens.server_name.clone(), + tokens.url.clone(), + Arc::new(TokioMutex::new(authorization_manager_for(tokens).await?)), + ResolvedOAuthCredentialStore::File, + Some(tokens.clone()), + )) +} + +async fn test_context() -> Result<(TempCodexHome, MockServer, StoredOAuthTokens)> { + let env = TempCodexHome::new(); + let server = MockServer::start().await; + mount_oauth_metadata(&server).await; + let tokens = expired_tokens(&format!("{}/mcp", server.uri())); + Ok((env, server, tokens)) +} + +async fn authorization_manager_for(tokens: &StoredOAuthTokens) -> Result { + let oauth_http_client = Arc::new(OAuthHttpClientAdapter::new( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + HeaderMap::new(), + &tokens.url, + )); + let mut state = + OAuthState::new_with_oauth_http_client(tokens.url.clone(), oauth_http_client).await?; + state + .set_credentials(&tokens.client_id, tokens.token_response.0.clone()) + .await?; + let manager = match state { + OAuthState::Authorized(manager) | OAuthState::Unauthorized(manager) => manager, + OAuthState::Session(_) | OAuthState::AuthorizedHttpClient(_) => { + anyhow::bail!("unexpected OAuth state") + } + _ => anyhow::bail!("unexpected OAuth state"), + }; + Ok(manager) +} + +async fn mount_oauth_metadata(server: &MockServer) { + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "issuer": format!("{}/mcp", server.uri()), + "authorization_endpoint": format!("{}/oauth/authorize", server.uri()), + "token_endpoint": format!("{}/oauth/token", server.uri()), + "scopes_supported": ["scope-a", "scope-b"], + }))) + .mount(server) + .await; +} + +async fn mount_delayed_refresh(server: &MockServer, response_access_token: &str) { + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains("refresh_token=refresh-token")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(Duration::from_millis(/*millis*/ 200)) + .set_body_json(serde_json::json!({ + "access_token": response_access_token, + "token_type": "Bearer", + "expires_in": 3600, + })), + ) + .expect(1) + .mount(server) + .await; +} + +async fn wait_for_lock_contention(rx: mpsc::Receiver<()>, expected_count: usize) -> Result<()> { + tokio::task::spawn_blocking(move || { + for _ in 0..expected_count { + rx.recv_timeout(Duration::from_secs(/*secs*/ 5)) + .context("timed out waiting for lock contention")?; + } + Ok(()) + }) + .await? +} + +fn expired_tokens(url: &str) -> StoredOAuthTokens { + let mut tokens = sample_tokens(); + tokens.url = url.to_string(); + tokens.issuer = Some(url.to_string()); + tokens.expires_at = Some(0); + tokens + .token_response + .0 + .set_expires_in(Some(&Duration::ZERO)); + tokens +} diff --git a/codex-rs/rmcp-client/src/oauth/tests/store_lock_tests.rs b/codex-rs/rmcp-client/src/oauth/tests/store_lock_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..b8428f5e5e972f961b93004a1c43f299468c999b --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth/tests/store_lock_tests.rs @@ -0,0 +1,707 @@ +use std::process::Command; +use std::sync::mpsc; +use std::time::Duration; +use std::time::Instant; + +use anyhow::Context; +use anyhow::Result; +use codex_config::types::AuthKeyringBackendKind; +use codex_keyring_store::KeyringStore; +use codex_keyring_store::tests::MockKeyringStore; +use oauth2::AccessToken; +use oauth2::RefreshToken; +use oauth2::Scope; +use oauth2::TokenResponse; +use oauth2::basic::BasicTokenType; +use pretty_assertions::assert_eq; +use rmcp::transport::auth::OAuthTokenResponse; +use rmcp::transport::auth::VendorExtraTokenFields; +use tracing::Event; +use tracing::Id; +use tracing::Metadata; +use tracing::Subscriber; +use tracing::span::Attributes; +use tracing::span::Record; +use tracing::subscriber::Interest; + +use super::OAuthStore; +use super::OAuthStoreLock; +use super::OAuthStoreLockFailure; +use super::OAuthStoreLockMode; +use crate::oauth::StoredOAuthCredentialSnapshot; +use crate::oauth::StoredOAuthTokens; +use crate::oauth::WrappedOAuthTokenResponse; +use crate::oauth::fallback_file_path; +use crate::oauth::load_oauth_tokens_from_file; +use crate::oauth::load_oauth_tokens_from_keyring; +use crate::oauth::resolve_oauth_tokens_from_store_policy; +use crate::oauth::save_oauth_tokens_to_file; +use crate::oauth::save_oauth_tokens_to_file_with_lock_held; +use crate::oauth::save_oauth_tokens_to_secrets_keyring_with_lock_held; +use crate::oauth::save_oauth_tokens_with_keyring; +use crate::oauth::save_oauth_tokens_with_keyring_with_fallback_to_file; +use crate::oauth::stored_oauth_credential_snapshot; +use crate::oauth::test_support::TempCodexHome; +use codex_config::types::OAuthCredentialsStoreMode; + +const STORE_LOCK_CONTENTION_EVENT_TARGET: &str = "codex_rmcp_client::oauth::store_lock::contention"; +// Contention is proven by the tracing event emitted after a real WouldBlock. Keep the timeout +// generous because it only bounds a failed test; it must not turn worker scheduling latency into +// a false failure on loaded CI hosts. +const STORE_LOCK_CONTENTION_EVENT_TIMEOUT: Duration = Duration::from_secs(/*secs*/ 10); + +fn assert_tokens_match_without_expiry(actual: &StoredOAuthTokens, expected: &StoredOAuthTokens) { + assert_eq!(actual.server_name, expected.server_name); + assert_eq!(actual.url, expected.url); + assert_eq!(actual.client_id, expected.client_id); + assert_eq!(actual.expires_at, expected.expires_at); + assert_token_response_match_without_expiry(&actual.token_response, &expected.token_response); +} + +fn assert_token_response_match_without_expiry( + actual: &WrappedOAuthTokenResponse, + expected: &WrappedOAuthTokenResponse, +) { + let actual_response = &actual.0; + let expected_response = &expected.0; + + assert_eq!( + actual_response.access_token().secret(), + expected_response.access_token().secret() + ); + assert_eq!(actual_response.token_type(), expected_response.token_type()); + assert_eq!( + actual_response.refresh_token().map(RefreshToken::secret), + expected_response.refresh_token().map(RefreshToken::secret), + ); + assert_eq!(actual_response.scopes(), expected_response.scopes()); + assert_eq!( + actual_response.extra_fields().0, + expected_response.extra_fields().0 + ); + assert_eq!( + actual_response.expires_in().is_some(), + expected_response.expires_in().is_some() + ); +} + +fn sample_tokens() -> StoredOAuthTokens { + let mut response = OAuthTokenResponse::new( + AccessToken::new("access-token".to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new("refresh-token".to_string()))); + response.set_scopes(Some(vec![ + Scope::new("scope-a".to_string()), + Scope::new("scope-b".to_string()), + ])); + let expires_in = Duration::from_secs(3600); + response.set_expires_in(Some(&expires_in)); + let expires_at = crate::oauth::compute_expires_at_millis(&response); + + StoredOAuthTokens { + server_name: "test-server".to_string(), + url: "https://example.test".to_string(), + issuer: Some("https://issuer.example.test".to_string()), + client_id: "client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at, + } +} + +#[test] +fn file_credentials_keep_repeated_local_prefixes_isolated() -> Result<()> { + let _env = TempCodexHome::new(); + let config: codex_config::McpServerConfig = serde_json::from_value(serde_json::json!({ + "url": "https://example.test", + }))?; + let mut first = sample_tokens(); + first.server_name = config.oauth_credential_name("local:foo").into_owned(); + first.expires_at = None; + first.token_response.0.set_expires_in(None); + save_oauth_tokens_to_file(&first)?; + + let second_name = config.oauth_credential_name("local:local:foo"); + assert_eq!(load_oauth_tokens_from_file(&second_name, &first.url)?, None); + + let mut second = first.clone(); + second.server_name = second_name.into_owned(); + second + .token_response + .0 + .set_access_token(AccessToken::new("second-token".to_string())); + save_oauth_tokens_to_file(&second)?; + + for expected in [first, second] { + assert_eq!( + load_oauth_tokens_from_file(&expected.server_name, &expected.url)?, + Some(expected) + ); + } + Ok(()) +} + +#[test] +fn legacy_rmcp_oauth_keyring_credentials_remain_readable() -> Result<()> { + let _env = TempCodexHome::new(); + let keyring_store = MockKeyringStore::default(); + let mut expected = sample_tokens(); + expected.expires_at = None; + expected.token_response.0.set_expires_in(None); + + let serialized = serde_json::json!({ + "server_name": "test-server", + "url": "https://example.test", + "client_id": "client-id", + "token_response": { + "access_token": "access-token", + "token_type": "Bearer", + "refresh_token": "refresh-token", + "scope": "scope-a scope-b", + }, + }) + .to_string(); + let key = crate::oauth::compute_store_key(&expected.server_name, &expected.url)?; + keyring_store.save(crate::oauth::KEYRING_SERVICE, &key, &serialized)?; + + let resolved = resolve_oauth_tokens_from_store_policy( + &keyring_store, + &expected.server_name, + &expected.url, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + )? + .expect("OAuth credentials written before the rmcp upgrade should remain readable"); + + assert_eq!( + resolved.store, + crate::oauth::ResolvedOAuthCredentialStore::Keyring(AuthKeyringBackendKind::Direct) + ); + assert_tokens_match_without_expiry(&resolved.tokens, &expected); + assert!(crate::oauth::oauth_tokens_are_usable(&resolved.tokens)); + Ok(()) +} + +const LOCK_HOLDER_CHILD_TEST: &str = + "oauth::store_lock::tests::store_lock_is_released_when_holder_process_exits_child"; +const LOCK_HOLDER_READY_PATH_ENV: &str = "CODEX_OAUTH_STORE_LOCK_CHILD_READY_PATH"; + +#[test] +fn store_lock_is_released_when_holder_process_exits() -> Result<()> { + let env = TempCodexHome::new(); + let ready_file = env.path().join("lock-holder-ready"); + let mut child = Command::new(std::env::current_exe()?) + .arg("--exact") + .arg(LOCK_HOLDER_CHILD_TEST) + .arg("--ignored") + .env("CODEX_HOME", env.path()) + .env(LOCK_HOLDER_READY_PATH_ENV, &ready_file) + .spawn() + .context("spawn OAuth store lock holder test process")?; + + let test_result = (|| -> Result<()> { + let started = Instant::now(); + while !ready_file.exists() { + if started.elapsed() > Duration::from_secs(/*secs*/ 5) { + anyhow::bail!("timed out waiting for child process to acquire OAuth store lock"); + } + std::thread::sleep(Duration::from_millis(/*millis*/ 20)); + } + + let error = match OAuthStoreLock::acquire_in_with_mode( + env.path(), + OAuthStore::File, + Duration::from_millis(/*millis*/ 100), + OAuthStoreLockMode::Exclusive, + ) { + Ok(_) => { + anyhow::bail!("live holder process should keep the OAuth store lock unavailable") + } + Err(error) => error, + }; + assert!(matches!(error, OAuthStoreLockFailure::Timeout { .. })); + + child + .kill() + .context("kill OAuth store lock holder process")?; + let status = child + .wait() + .context("wait for killed OAuth store lock holder process")?; + assert!(!status.success()); + let _lock = OAuthStoreLock::acquire_in_with_mode( + env.path(), + OAuthStore::File, + Duration::from_secs(/*secs*/ 1), + OAuthStoreLockMode::Exclusive, + )?; + Ok(()) + })(); + + if let Ok(None) = child.try_wait() { + let _ = child.kill(); + let _ = child.wait(); + } + + test_result +} + +#[test] +#[ignore = "child process for store_lock_is_released_when_holder_process_exits"] +fn store_lock_is_released_when_holder_process_exits_child() -> Result<()> { + let ready_file = match std::env::var_os(LOCK_HOLDER_READY_PATH_ENV) { + Some(path) => std::path::PathBuf::from(path), + None => return Ok(()), + }; + let _lock = OAuthStoreLock::acquire_for_write(OAuthStore::File)?; + std::fs::write(ready_file, b"ready")?; + loop { + std::thread::sleep(Duration::from_secs(/*secs*/ 60)); + } +} + +#[test] +fn auto_save_secrets_lock_failure_does_not_fall_back_to_file() -> Result<()> { + let env = TempCodexHome::new(); + let lock_dir = env.path().join("mcp-oauth-locks"); + std::fs::create_dir_all(&lock_dir)?; + // Break only the Secrets lock path. The distinct File lock remains usable, so Auto would + // successfully write fallback credentials if it mistook coordination failure for backend + // unavailability. + std::fs::create_dir(lock_dir.join("secrets-store.lock"))?; + let keyring_store = MockKeyringStore::default(); + let tokens = sample_tokens(); + + let error = save_oauth_tokens_with_keyring_with_fallback_to_file( + &keyring_store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens, + ) + .expect_err("aggregate-store lock failure must abort Auto persistence"); + + assert!(error.downcast_ref::().is_some()); + assert!(!fallback_file_path()?.exists()); + save_oauth_tokens_to_file(&tokens)?; + let loaded = load_oauth_tokens_from_file(&tokens.server_name, &tokens.url)? + .expect("fallback File should remain independently writable"); + assert_tokens_match_without_expiry(&loaded, &tokens); + Ok(()) +} + +#[test] +fn auto_load_secrets_lock_failure_does_not_fall_back_to_file() -> Result<()> { + let env = TempCodexHome::new(); + let keyring_store = MockKeyringStore::default(); + let tokens = sample_tokens(); + save_oauth_tokens_to_file(&tokens)?; + + let lock_dir = env.path().join("mcp-oauth-locks"); + std::fs::create_dir(lock_dir.join("secrets-store.lock"))?; + let error = resolve_oauth_tokens_from_store_policy( + &keyring_store, + &tokens.server_name, + &tokens.url, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Secrets, + ) + .expect_err("aggregate-store lock failure must abort Auto resolution"); + + assert!(error.downcast_ref::().is_some()); + let loaded = load_oauth_tokens_from_file(&tokens.server_name, &tokens.url)? + .expect("fallback File should remain independently readable"); + assert_tokens_match_without_expiry(&loaded, &tokens); + Ok(()) +} + +#[test] +fn oauth_credential_probes_skip_contended_file_and_secrets_stores() -> Result<()> { + let _env = TempCodexHome::new(); + let tokens = sample_tokens(); + let file = OAuthCredentialsStoreMode::File; + let auto = OAuthCredentialsStoreMode::Auto; + let direct = AuthKeyringBackendKind::Direct; + let secrets = AuthKeyringBackendKind::Secrets; + save_oauth_tokens_to_file(&tokens)?; + let snapshot = + stored_oauth_credential_snapshot(&tokens.server_name, &tokens.url, file, direct)? + .expect("fallback credentials should exist before contention"); + let reload = |store_mode, keyring_backend_kind| { + snapshot.reload( + &tokens.server_name, + &tokens.url, + store_mode, + keyring_backend_kind, + ) + }; + let probe = |store_mode, keyring_backend_kind| { + StoredOAuthCredentialSnapshot::for_runtime_refresh( + /*previous*/ None, + &tokens.server_name, + &tokens.url, + store_mode, + keyring_backend_kind, + ) + }; + let refresh = |url| { + StoredOAuthCredentialSnapshot::for_runtime_refresh( + Some(&snapshot), + &tokens.server_name, + url, + file, + direct, + ) + }; + + { + let _lock = OAuthStoreLock::acquire_for_write(OAuthStore::File)?; + assert_eq!(reload(file, direct)?, None); + assert_eq!(reload(auto, secrets)?, None); + assert_eq!(probe(file, direct)?, None); + let contended = refresh(&tokens.url)?.expect("the previous snapshot should be retained"); + assert!(contended.store_was_contended()); + assert_eq!(contended, snapshot); + assert_eq!( + refresh("https://another.example.test")?, + None, + "store contention must not replay credentials for another endpoint", + ); + } + + assert_eq!(reload(file, direct)?, Some(snapshot.credentials().clone())); + let refreshed = refresh(&tokens.url)?.expect("the unlocked store should still contain tokens"); + assert!(!refreshed.store_was_contended()); + + let _lock = OAuthStoreLock::acquire_for_write(OAuthStore::Secrets)?; + assert_eq!( + probe(auto, secrets)?, + None, + "a locked Secrets authority must not fall back to the stale File entry", + ); + Ok(()) +} + +struct LockContentionSubscriber { + contended_tx: mpsc::Sender<()>, +} + +impl Subscriber for LockContentionSubscriber { + fn enabled(&self, metadata: &Metadata<'_>) -> bool { + metadata.target() == STORE_LOCK_CONTENTION_EVENT_TARGET + } + + fn register_callsite(&self, metadata: &'static Metadata<'static>) -> Interest { + if self.enabled(metadata) { + Interest::always() + } else { + Interest::never() + } + } + + fn max_level_hint(&self) -> Option { + Some(tracing::level_filters::LevelFilter::DEBUG) + } + + fn new_span(&self, _span: &Attributes<'_>) -> Id { + Id::from_u64(/*u*/ 1) + } + + fn record(&self, _span: &Id, _values: &Record<'_>) {} + + fn record_follows_from(&self, _span: &Id, _follows_from: &Id) {} + + fn event(&self, event: &Event<'_>) { + if self.enabled(event.metadata()) { + self.contended_tx + .send(()) + .expect("signal actual OAuth store lock contention"); + } + } + + fn enter(&self, _span: &Id) {} + + fn exit(&self, _span: &Id) {} +} + +fn complete_after_store_lock_contention( + codex_home: &std::path::Path, + store: OAuthStore, + while_locked: impl FnOnce() -> Result<()>, + operation: impl FnOnce() -> Result + Send + 'static, +) -> Result +where + T: Send + 'static, +{ + std::thread::scope(|scope| { + let held_lock = OAuthStoreLock::acquire_in_with_mode( + codex_home, + store, + Duration::from_millis(/*millis*/ 100), + OAuthStoreLockMode::Exclusive, + )?; + let (contended_tx, contended_rx) = mpsc::channel(); + let worker = scope.spawn(move || { + tracing::subscriber::with_default(LockContentionSubscriber { contended_tx }, operation) + }); + + // This event is emitted only after `try_lock()` returns WouldBlock, so the test fails if the + // operation stops acquiring the aggregate-store lock. + contended_rx + .recv_timeout(STORE_LOCK_CONTENTION_EVENT_TIMEOUT) + .context("timed out waiting for actual OAuth store lock contention")?; + while_locked()?; + drop(held_lock); + worker + .join() + .expect("contending OAuth store worker should finish") + }) +} + +#[test] +fn aggregate_store_readers_share_access_while_writers_remain_exclusive() -> Result<()> { + let env = TempCodexHome::new(); + + for store in [OAuthStore::File, OAuthStore::Secrets] { + let readers = (0..2) + .map(|_| { + OAuthStoreLock::acquire_in_with_mode( + env.path(), + store, + Duration::from_millis(/*millis*/ 100), + OAuthStoreLockMode::Shared, + ) + }) + .collect::, _>>()?; + + let writer_error = match OAuthStoreLock::acquire_in_with_mode( + env.path(), + store, + Duration::from_millis(/*millis*/ 100), + OAuthStoreLockMode::Exclusive, + ) { + Ok(_) => anyhow::bail!("an active OAuth store reader must exclude writers"), + Err(error) => error, + }; + assert!(matches!( + writer_error, + OAuthStoreLockFailure::Timeout { .. } + )); + drop(readers); + let _writer = OAuthStoreLock::acquire_in_with_mode( + env.path(), + store, + Duration::from_millis(/*millis*/ 100), + OAuthStoreLockMode::Exclusive, + )?; + } + + Ok(()) +} + +#[test] +fn aggregate_store_credential_loads_can_share_an_existing_reader() -> Result<()> { + let env = TempCodexHome::new(); + let keyring_store = MockKeyringStore::default(); + let tokens = sample_tokens(); + save_oauth_tokens_to_file(&tokens)?; + save_oauth_tokens_with_keyring( + &keyring_store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens, + )?; + + let file_reader = OAuthStoreLock::acquire_in_with_mode( + env.path(), + OAuthStore::File, + Duration::from_millis(/*millis*/ 100), + OAuthStoreLockMode::Shared, + )?; + let file_tokens = load_oauth_tokens_from_file(&tokens.server_name, &tokens.url)? + .expect("a File credential read should coexist with another reader"); + assert_tokens_match_without_expiry(&file_tokens, &tokens); + drop(file_reader); + + let secrets_reader = OAuthStoreLock::acquire_in_with_mode( + env.path(), + OAuthStore::Secrets, + Duration::from_millis(/*millis*/ 100), + OAuthStoreLockMode::Shared, + )?; + let secrets_tokens = load_oauth_tokens_from_keyring( + &keyring_store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens.url, + )? + .expect("a Secrets credential read should coexist with another reader"); + assert_tokens_match_without_expiry(&secrets_tokens, &tokens); + drop(secrets_reader); + + Ok(()) +} + +#[test] +fn file_store_lock_preserves_updates_for_different_servers() -> Result<()> { + let env = TempCodexHome::new(); + let first = sample_tokens(); + let mut second = sample_tokens(); + second.server_name = "second-server".to_string(); + second.url = "https://second.example.test".to_string(); + + let second_for_writer = second.clone(); + complete_after_store_lock_contention( + env.path(), + OAuthStore::File, + || save_oauth_tokens_to_file_with_lock_held(&first), + move || save_oauth_tokens_to_file(&second_for_writer), + )?; + + let loaded_first = load_oauth_tokens_from_file(&first.server_name, &first.url)? + .expect("first server tokens should remain stored"); + let loaded_second = load_oauth_tokens_from_file(&second.server_name, &second.url)? + .expect("second server tokens should be stored"); + assert_tokens_match_without_expiry(&loaded_first, &first); + assert_tokens_match_without_expiry(&loaded_second, &second); + Ok(()) +} + +#[test] +fn file_store_load_and_delete_observe_aggregate_lock() -> Result<()> { + let env = TempCodexHome::new(); + let tokens = sample_tokens(); + save_oauth_tokens_to_file(&tokens)?; + + let server_name = tokens.server_name.clone(); + let url = tokens.url.clone(); + let loaded = complete_after_store_lock_contention( + env.path(), + OAuthStore::File, + || Ok(()), + move || load_oauth_tokens_from_file(&server_name, &url), + )? + .expect("file credentials should remain readable after contention"); + assert_tokens_match_without_expiry(&loaded, &tokens); + + let key = crate::oauth::compute_store_key(&tokens.server_name, &tokens.url)?; + let removed = complete_after_store_lock_contention( + env.path(), + OAuthStore::File, + || Ok(()), + move || crate::oauth::delete_oauth_tokens_from_file(&key), + )?; + assert!(removed); + assert!(load_oauth_tokens_from_file(&tokens.server_name, &tokens.url)?.is_none()); + Ok(()) +} + +#[test] +fn secrets_store_lock_preserves_updates_for_different_servers() -> Result<()> { + let env = TempCodexHome::new(); + let keyring_store = MockKeyringStore::default(); + let first = sample_tokens(); + let mut second = sample_tokens(); + second.server_name = "second-server".to_string(); + second.url = "https://second.example.test".to_string(); + + let store_for_writer = keyring_store.clone(); + let second_for_writer = second.clone(); + complete_after_store_lock_contention( + env.path(), + OAuthStore::Secrets, + || { + let first_serialized = serde_json::to_string(&first)?; + save_oauth_tokens_to_secrets_keyring_with_lock_held( + &keyring_store, + &first.server_name, + &first, + &first_serialized, + ) + }, + move || { + save_oauth_tokens_with_keyring( + &store_for_writer, + AuthKeyringBackendKind::Secrets, + &second_for_writer.server_name, + &second_for_writer, + ) + }, + )?; + + let loaded_first = load_oauth_tokens_from_keyring( + &keyring_store, + AuthKeyringBackendKind::Secrets, + &first.server_name, + &first.url, + )? + .expect("first server tokens should remain stored"); + let loaded_second = load_oauth_tokens_from_keyring( + &keyring_store, + AuthKeyringBackendKind::Secrets, + &second.server_name, + &second.url, + )? + .expect("second server tokens should be stored"); + assert_tokens_match_without_expiry(&loaded_first, &first); + assert_tokens_match_without_expiry(&loaded_second, &second); + Ok(()) +} + +#[test] +fn secrets_store_load_and_delete_observe_aggregate_lock() -> Result<()> { + let env = TempCodexHome::new(); + let keyring_store = MockKeyringStore::default(); + let tokens = sample_tokens(); + save_oauth_tokens_with_keyring( + &keyring_store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens, + )?; + + let store_for_load = keyring_store.clone(); + let server_name = tokens.server_name.clone(); + let url = tokens.url.clone(); + let loaded = complete_after_store_lock_contention( + env.path(), + OAuthStore::Secrets, + || Ok(()), + move || { + Ok(load_oauth_tokens_from_keyring( + &store_for_load, + AuthKeyringBackendKind::Secrets, + &server_name, + &url, + )?) + }, + )? + .expect("encrypted credentials should remain readable after contention"); + assert_tokens_match_without_expiry(&loaded, &tokens); + + let store_for_delete = keyring_store.clone(); + let server_name = tokens.server_name.clone(); + let url = tokens.url.clone(); + let removed = complete_after_store_lock_contention( + env.path(), + OAuthStore::Secrets, + || Ok(()), + move || { + crate::oauth::delete_oauth_tokens_from_secrets_keyring( + &store_for_delete, + &server_name, + &url, + ) + }, + )?; + assert!(removed); + assert!( + load_oauth_tokens_from_keyring( + &keyring_store, + AuthKeyringBackendKind::Secrets, + &tokens.server_name, + &tokens.url, + )? + .is_none() + ); + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/oauth_callback.rs b/codex-rs/rmcp-client/src/oauth_callback.rs new file mode 100644 index 0000000000000000000000000000000000000000..eb3d7f31765b1d4e6e441e0f53b4ae761ca989c8 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_callback.rs @@ -0,0 +1,148 @@ +//! OAuth callback identity and authorization-server mix-up protection. +//! +//! Codex can authorize against many independent MCP servers. If those servers +//! share a callback URL and a response does not identify its authorization +//! server, Codex could associate an authorization code with the wrong server +//! and send that code to an attacker-controlled token endpoint. RFC 9700 calls +//! this an authorization-server mix-up attack: +//! https://www.rfc-editor.org/rfc/rfc9700#section-4.4 +//! +//! MCP prefers issuer identification: authorization servers SHOULD include +//! `iss` in authorization responses; servers that include `iss` MUST advertise +//! that support, and clients MUST validate any returned issuer before +//! exchanging the code. This also lets a CIMD document advertise one stable +//! redirect instead of separate redirects for every authorization server: +//! https://modelcontextprotocol.io/specification/2026-07-28/basic/authorization#authorization-response-validation +//! +//! This module prefers issuer binding when available and otherwise retains a +//! callback-specific compatibility fallback: +//! +//! - `IssuerBound`: reuse a stable callback only when authorization metadata +//! advertises `authorization_response_iss_parameter_supported` and contains +//! its issuer. RMCP validates the response's `iss` against that metadata +//! issuer before exchanging the authorization code, rejecting missing or +//! mismatched issuers: +//! https://www.rfc-editor.org/rfc/rfc9700#section-4.4.2.1 +//! https://www.rfc-editor.org/rfc/rfc9207#section-2.4 +//! - `CallbackSpecific`: append an ID derived from the complete MCP server URL. +//! Distinct callback paths bind each response to its intended server. This is +//! the required fallback when issuer identification is not supported: +//! https://www.rfc-editor.org/rfc/rfc9700#section-4.4.2.2 + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use base64::Engine; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use rmcp::transport::auth::AuthorizationMetadata; +use sha2::Digest; +use sha2::Sha256; +use url::Url; + +/// The OAuth mix-up defense associated with a newly registered callback. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum McpOAuthCallbackMode { + /// A server-specific redirect path identifies the authorization server. + CallbackSpecific, + /// A validated authorization-response issuer permits a shared redirect. + IssuerBound, +} + +pub(crate) fn callback_mode(metadata: &AuthorizationMetadata) -> Result { + let issuer_response_supported = metadata + .additional_fields + .get("authorization_response_iss_parameter_supported") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false); + + if !issuer_response_supported { + return Ok(McpOAuthCallbackMode::CallbackSpecific); + } + if metadata.issuer.is_none() { + bail!("OAuth authorization server advertises issuer support without a metadata issuer"); + } + + Ok(McpOAuthCallbackMode::IssuerBound) +} + +/// Resolves the registered callback independently of runtime listener ports. +pub fn resolve_mcp_oauth_callback_url( + server_url: &str, + callback_url: Option<&str>, + callback_mode: McpOAuthCallbackMode, +) -> Result { + let callback_url = callback_url.unwrap_or("http://127.0.0.1/callback"); + + match callback_mode { + McpOAuthCallbackMode::IssuerBound => { + Url::parse(callback_url) + .with_context(|| format!("invalid redirect URI `{callback_url}`"))?; + Ok(callback_url.to_string()) + } + McpOAuthCallbackMode::CallbackSpecific => { + let callback_id = callback_id_from_server_url(server_url)?; + append_callback_id_to_redirect_uri(callback_url, &callback_id) + } + } +} + +pub(crate) fn callback_id_from_server_url(server_url: &str) -> Result { + // Native Codex callback IDs intentionally hash the complete MCP URL (minus its fragment) + // with SHA-256. Python connector callback IDs use SHAKE-256 over the origin and are distinct. + let mut parsed = + Url::parse(server_url).with_context(|| format!("invalid MCP server URL `{server_url}`"))?; + parsed + .host_str() + .ok_or_else(|| anyhow!("MCP server URL `{server_url}` must include a host"))?; + parsed.set_fragment(None); + + let digest = Sha256::digest(parsed.as_str().as_bytes()); + Ok(URL_SAFE_NO_PAD.encode(&digest[..9])) +} + +pub(crate) fn append_callback_id_to_redirect_uri( + redirect_uri: &str, + callback_id: &str, +) -> Result { + let mut parsed = Url::parse(redirect_uri) + .with_context(|| format!("invalid redirect URI `{redirect_uri}`"))?; + if parsed + .path_segments() + .and_then(|mut segments| segments.next_back()) + == Some(callback_id) + { + return Ok(parsed.to_string()); + } + let path = parsed.path(); + let new_path = if path.ends_with('/') { + format!("{path}{callback_id}") + } else { + format!("{path}/{callback_id}") + }; + parsed.set_path(&new_path); + Ok(parsed.to_string()) +} + +pub(crate) fn validate_callback_redirect( + redirect_uri: &str, + callback_id: &str, + callback_mode: McpOAuthCallbackMode, +) -> Result<()> { + let has_expected_callback_id = Url::parse(redirect_uri)? + .path_segments() + .and_then(|mut segments| segments.next_back()) + == Some(callback_id); + + if !has_expected_callback_id && callback_mode != McpOAuthCallbackMode::IssuerBound { + bail!( + "OAuth callback requires its expected callback ID or authorization response issuer support" + ); + } + + Ok(()) +} + +#[cfg(test)] +#[path = "oauth_callback_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/oauth_callback_input.rs b/codex-rs/rmcp-client/src/oauth_callback_input.rs new file mode 100644 index 0000000000000000000000000000000000000000..535876413781a69ca3a14b3da4bbefae70d71f6c --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_callback_input.rs @@ -0,0 +1,178 @@ +//! Accepts a pasted OAuth redirect without navigating to it. The prepared callback address +//! is checked before the existing OAuth flow validates state/issuer and exchanges the code. + +use super::CallbackResult; +use super::OAuthHttpContext; +use super::OAuthLoginPurpose; +use super::OAuthProviderError; +use super::OauthCallbackResult; +use super::OauthLoginFlow; +use crate::McpOAuthClientRegistration; +use crate::StreamableHttpRedirectMode; +use crate::save_oauth_tokens; +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::HttpClient; +use std::collections::HashMap; +use std::future::Future; +use std::sync::Arc; +use tokio::time::timeout; +use url::Url; + +/// Runs MCP OAuth without opening a browser, accepting either the HTTP callback or a +/// full redirect URL returned by `read_callback`. The reader receives the authorization +/// URL and must release terminal state when its future is dropped (callback or timeout). +#[allow(clippy::too_many_arguments)] +pub async fn perform_oauth_login_with_callback_input( + server_name: &str, + server_url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_headers: Option>, + env_http_headers: Option>, + scopes: &[String], + oauth_client_id: Option<&str>, + client_registration: McpOAuthClientRegistration, + oauth_resource: Option<&str>, + callback_port: Option, + callback_url: Option<&str>, + global_callback_url: Option<&str>, + http_client: Arc, + read_callback: impl FnOnce(String) -> F, +) -> Result<()> +where + F: Future>, +{ + let mut flow = OauthLoginFlow::new( + server_name, + server_url, + store_mode, + keyring_backend_kind, + OAuthHttpContext { + http_headers, + env_http_headers, + http_client, + redirect_mode: StreamableHttpRedirectMode::Legacy, + }, + scopes, + oauth_client_id, + OAuthLoginPurpose::Mcp, + client_registration, + oauth_resource, + /*launch_browser*/ false, + callback_port, + callback_url, + global_callback_url, + /*timeout_secs*/ None, + ) + .await?; + let authorization_url = flow.authorization_url(); + let callback = timeout(flow.timeout, async { + tokio::select! { + callback = &mut flow.rx => callback.context("OAuth callback was cancelled"), + input = read_callback(authorization_url.clone()) => { + parse_callback_url(&input?, &flow.redirect_uri, &authorization_url) + } + } + }) + .await + .context("timed out waiting for OAuth callback")??; + // RMCP's issuer-mismatch error includes the received value. Reject it here + // without echoing any part of a pasted callback into terminal diagnostics. + if let CallbackResult::Success(callback) = &callback + && callback + .issuer + .as_deref() + .is_some_and(|issuer| Some(issuer) != flow.authorization_server_issuer.as_deref()) + { + bail!("OAuth callback issuer does not match this login"); + } + let stored = flow.complete_callback(callback).await?; + save_oauth_tokens(server_name, &stored, store_mode, keyring_backend_kind).await +} + +fn parse_callback_url( + input: &str, + redirect_uri: &str, + authorization_url: &str, +) -> Result { + if input.len() > 64 * 1024 { + bail!("OAuth callback URL exceeds 64 KiB"); + } + let mut callback = Url::parse(input.trim()).context("Invalid OAuth callback URL")?; + let mut expected = Url::parse(redirect_uri).context("Invalid OAuth redirect URI")?; + if callback.fragment().is_some() + || !callback.username().is_empty() + || callback.password().is_some() + { + bail!("OAuth callback URL must not contain credentials or a fragment"); + } + let mut response_params: Vec<_> = callback.query_pairs().into_owned().collect(); + // Configured redirect query parameters must survive unchanged. OAuth response + // parameters are appended by the authorization server, not part of that address. + let expected_params: Vec<_> = expected.query_pairs().into_owned().collect(); + for expected_param in &expected_params { + let position = response_params + .iter() + .position(|param| param == expected_param) + .ok_or_else(|| { + anyhow!("OAuth callback URL does not match this login's redirect URI") + })?; + response_params.remove(position); + } + if response_params + .iter() + .any(|(name, _)| expected_params.iter().any(|(key, _)| key == name)) + { + bail!("OAuth callback URL changes this login's redirect query parameters"); + } + callback.set_query(/*query*/ None); + expected.set_query(/*query*/ None); + if callback != expected { + bail!("OAuth callback URL does not match this login's redirect URI"); + } + + let mut params = HashMap::new(); + for (name, value) in response_params { + if params.insert(name, value).is_some() { + bail!("OAuth callback URL contains duplicate parameters"); + } + } + let state = params + .remove("state") + .filter(|state| !state.is_empty()) + .ok_or_else(|| anyhow!("OAuth callback URL is missing state"))?; + if params.contains_key("error") { + if params.contains_key("code") { + bail!("OAuth callback URL contains both a code and an error"); + } + let authorization = Url::parse(authorization_url)?; + if !authorization + .query_pairs() + .any(|(key, value)| key == "state" && value == state) + { + bail!("OAuth callback state does not match this login"); + } + // Do not print arbitrary pasted error descriptions or other callback values. + return Ok(CallbackResult::Error(OAuthProviderError::new( + /*error*/ None, /*error_description*/ None, + ))); + } + let code = params + .remove("code") + .filter(|code| !code.is_empty()) + .ok_or_else(|| anyhow!("OAuth callback URL is missing an authorization code"))?; + Ok(CallbackResult::Success(OauthCallbackResult { + code, + state, + issuer: params.remove("iss"), + })) +} + +#[cfg(test)] +#[path = "oauth_callback_input_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/oauth_callback_input_tests.rs b/codex-rs/rmcp-client/src/oauth_callback_input_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..21f2b3486e01c65be9ae4ddf1ac940a58510c7c0 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_callback_input_tests.rs @@ -0,0 +1,70 @@ +//! Validates pasted callback addresses, parameter ambiguity, and redacted errors. + +use super::*; +use pretty_assertions::assert_eq; + +#[test] +fn pasted_callback_preserves_configured_query_and_decodes_response() { + let CallbackResult::Success(callback) = parse_callback_url( + " https://callback.example:9443/registered?tenant=a&code=a%2Bb&state=csrf&iss=https%3A%2F%2Fissuer.example ", + "https://callback.example:9443/registered?tenant=a", + "https://issuer.example/authorize?state=csrf", + ) + .unwrap() else { + panic!("expected successful callback"); + }; + assert_eq!( + callback, + OauthCallbackResult { + code: "a+b".to_string(), + state: "csrf".to_string(), + issuer: Some("https://issuer.example".to_string()), + } + ); +} + +#[test] +fn pasted_callback_rejects_ambiguous_or_unbound_urls_without_echoing_values() { + let redirect = "http://127.0.0.1:1234/callback/server?tenant=a"; + for input in [ + "http://127.0.0.1:1234/other?tenant=a&code=secret&state=csrf", + "http://127.0.0.1:1235/callback/server?tenant=a&code=secret&state=csrf", + "http://127.0.0.1:1234/callback/server?tenant=b&code=secret&state=csrf", + "http://127.0.0.1:1234/callback/server?tenant=a&tenant=b&code=secret&state=csrf", + "http://127.0.0.1:1234/callback/server?tenant=a&code=secret&%63ode=second&state=csrf", + "http://127.0.0.1:1234/callback/server?tenant=a&code=secret&state=csrf&state=second", + "http://127.0.0.1:1234/callback/server?tenant=a&code=secret&state=csrf&iss=a&iss=b", + "http://secret@127.0.0.1:1234/callback/server?tenant=a&code=secret&state=csrf", + "http://127.0.0.1:1234/callback/server?tenant=a&code=secret&state=csrf#fragment", + "http://127.0.0.1:1234/callback/server?tenant=a&code=secret&state=", + "http://127.0.0.1:1234/callback/server?tenant=a&code=secret&state=csrf&error=denied", + ] { + let error = + parse_callback_url(input, redirect, "https://issuer.example?state=csrf").unwrap_err(); + assert!(!format!("{error:#}").contains("secret")); + } + assert!(parse_callback_url(&"x".repeat(/*n*/ 65_537), redirect, "").is_err()); +} + +#[test] +fn pasted_provider_error_requires_matching_state_and_redacts_description() { + let redirect = "http://127.0.0.1:1234/callback/server"; + let input = format!("{redirect}?error=access_denied&error_description=secret&state=csrf"); + let CallbackResult::Error(error) = parse_callback_url( + &input, + redirect, + "https://issuer.example/authorize?state=csrf", + ) + .unwrap() else { + panic!("expected provider error"); + }; + assert_eq!(error.to_string(), "OAuth provider returned an error"); + assert!( + parse_callback_url( + &input, + redirect, + "https://issuer.example/authorize?state=other", + ) + .is_err() + ); +} diff --git a/codex-rs/rmcp-client/src/oauth_callback_tests.rs b/codex-rs/rmcp-client/src/oauth_callback_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..3d70c780a08ac29d783534d371701551c4863348 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_callback_tests.rs @@ -0,0 +1,63 @@ +use pretty_assertions::assert_eq; + +use super::McpOAuthCallbackMode; +use super::callback_id_from_server_url; +use super::resolve_mcp_oauth_callback_url; +use super::validate_callback_redirect; + +#[test] +fn resolved_callbacks_follow_the_selected_mix_up_defense() { + let server_url = "https://mcp.example.com/mcp?tenant=one"; + let callback_id = callback_id_from_server_url(server_url).expect("resolve callback ID"); + let distinct_callback = format!("http://127.0.0.1/callback/{callback_id}"); + + for (callback, mode, expected) in [ + ( + None, + McpOAuthCallbackMode::CallbackSpecific, + distinct_callback.as_str(), + ), + ( + None, + McpOAuthCallbackMode::IssuerBound, + "http://127.0.0.1/callback", + ), + ( + Some("http://127.0.0.1:8080/oauth/callback"), + McpOAuthCallbackMode::IssuerBound, + "http://127.0.0.1:8080/oauth/callback", + ), + ] { + assert_eq!( + resolve_mcp_oauth_callback_url(server_url, callback, mode) + .expect("resolve registered callback"), + expected + ); + } +} + +#[test] +fn callback_redirect_requires_a_server_specific_id_or_issuer_support() { + for (redirect_uri, mode, expected_valid) in [ + ( + "http://127.0.0.1/callback/expected-id", + McpOAuthCallbackMode::CallbackSpecific, + true, + ), + ( + "http://127.0.0.1/callback", + McpOAuthCallbackMode::IssuerBound, + true, + ), + ( + "http://127.0.0.1/callback/wrong-id", + McpOAuthCallbackMode::CallbackSpecific, + false, + ), + ] { + assert_eq!( + validate_callback_redirect(redirect_uri, "expected-id", mode).is_ok(), + expected_valid + ); + } +} diff --git a/codex-rs/rmcp-client/src/oauth_client_registration.rs b/codex-rs/rmcp-client/src/oauth_client_registration.rs new file mode 100644 index 0000000000000000000000000000000000000000..d9de5e21f7129f06b09ee66079f4bda8fdb7a4aa --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_client_registration.rs @@ -0,0 +1,139 @@ +use std::sync::Arc; + +use anyhow::Result; +use anyhow::bail; +use rmcp::transport::AuthorizationManager; +use rmcp::transport::AuthorizationRequest; +use rmcp::transport::AuthorizationSession; +use rmcp::transport::auth::OAuthHttpClient; +use rmcp::transport::auth::OAuthState; +use url::Url; + +use crate::oauth::validate_authorization_server_endpoints; +use crate::oauth_callback::McpOAuthCallbackMode; +use crate::oauth_callback::append_callback_id_to_redirect_uri; +use crate::oauth_callback::callback_mode; +use crate::oauth_callback::validate_callback_redirect; + +/// OAuth client-registration strategy for one interactive HTTP MCP login. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub enum McpOAuthClientRegistration { + /// Prefer a supported native CIMD and otherwise use advertised DCR. + #[default] + Auto, + /// Require a ChatGPT-hosted Codex public native Client ID Metadata Document. + Cimd, + /// Require the authorization server's Dynamic Client Registration endpoint. + Dcr, +} + +/// OAuth state prepared from one authorization-server metadata resolution. +pub(crate) struct PreparedOAuthLogin { + pub(crate) oauth_state: OAuthState, + pub(crate) authorization_server_issuer: Option, + pub(crate) redirect_uri: String, +} + +pub(crate) async fn start_authorization( + server_url: &str, + http_client: Arc, + scopes: &[&str], + redirect_uri: &str, + callback_id: &str, + client_registration: McpOAuthClientRegistration, +) -> Result { + let mut auth_manager = + AuthorizationManager::new_with_oauth_http_client(server_url, http_client).await?; + auth_manager.set_allow_missing_issuer(true); + let metadata = auth_manager.resolve_metadata().await?.metadata; + validate_authorization_server_endpoints(&metadata)?; + let authorization_server_issuer = metadata.issuer.clone(); + let callback_mode = callback_mode(&metadata)?; + + let cimd_advertised = metadata + .additional_fields + .get("client_id_metadata_document_supported") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false); + let public_client_auth_supported = metadata + .additional_fields + .get("token_endpoint_auth_methods_supported") + .and_then(serde_json::Value::as_array) + .is_some_and(|methods| methods.iter().any(|method| method.as_str() == Some("none"))); + + let uses_shared_callback = callback_mode == McpOAuthCallbackMode::IssuerBound; + let redirect_uri = if uses_shared_callback { + redirect_uri.to_string() + } else { + append_callback_id_to_redirect_uri(redirect_uri, callback_id)? + }; + let parsed_redirect_uri = Url::parse(&redirect_uri)?; + let expected_callback_path = if uses_shared_callback { + "/callback".to_string() + } else { + format!("/callback/{callback_id}") + }; + let native_redirect_supported = parsed_redirect_uri.scheme() == "http" + && matches!( + parsed_redirect_uri.host_str(), + Some("127.0.0.1" | "localhost") + ) + && parsed_redirect_uri.port().is_some_and(|port| port > 0) + && parsed_redirect_uri.path() == expected_callback_path + && parsed_redirect_uri.query().is_none() + && parsed_redirect_uri.fragment().is_none() + && parsed_redirect_uri.username().is_empty() + && parsed_redirect_uri.password().is_none(); + validate_callback_redirect(&redirect_uri, callback_id, callback_mode)?; + // MCP 2026-07-28 priority: pre-registered clients never reach this path; offer + // advertised CIMD here and otherwise let rmcp fall back to DCR. + // https://modelcontextprotocol.io/specification/2026-07-28/basic/authorization/client-registration + let offer_cimd = match client_registration { + McpOAuthClientRegistration::Auto => { + cimd_advertised && native_redirect_supported && public_client_auth_supported + } + McpOAuthClientRegistration::Cimd => { + if !cimd_advertised || !public_client_auth_supported { + bail!( + "MCP authorization server does not advertise CIMD with token endpoint auth method `none`" + ); + } + if !native_redirect_supported { + bail!( + "MCP OAuth CIMD requires an ephemeral loopback callback at `{expected_callback_path}`" + ); + } + true + } + McpOAuthClientRegistration::Dcr => false, + }; + + auth_manager.set_metadata(metadata); + let mut request = AuthorizationRequest::new(redirect_uri.clone()) + .with_scopes(scopes.iter().copied()) + .with_client_name("Codex"); + if offer_cimd { + // CIMD is an active IETF Internet-Draft: this HTTPS client identifier resolves + // to its self-referential JSON metadata document. + // https://datatracker.ietf.org/doc/draft-ietf-oauth-client-id-metadata-document/ + let client_metadata_url = if uses_shared_callback { + "https://chatgpt.com/oauth/codex/client.json".to_string() + } else { + format!("https://chatgpt.com/oauth/codex/{callback_id}/client.json") + }; + request = request.with_client_metadata_url(client_metadata_url); + } + let session = AuthorizationSession::new(auth_manager, request) + .await + .map_err(|(_auth_manager, error)| error)?; + + Ok(PreparedOAuthLogin { + oauth_state: OAuthState::Session(session), + authorization_server_issuer, + redirect_uri, + }) +} + +#[cfg(test)] +#[path = "oauth_client_registration_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/oauth_client_registration_tests.rs b/codex-rs/rmcp-client/src/oauth_client_registration_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..8835449f5de108495c354f0cfa673735d09dc478 --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_client_registration_tests.rs @@ -0,0 +1,571 @@ +use std::collections::HashMap; +use std::sync::Arc; + +use anyhow::Result; +use base64::Engine; +use base64::engine::general_purpose::STANDARD; +use codex_exec_server::RouteAwareHttpClient; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use http::HeaderMap; +use pretty_assertions::assert_eq; +use rmcp::transport::auth::AuthorizationMetadata; +use rmcp::transport::auth::OAuthState; +use serde_json::Value; +use serde_json::json; +use url::Url; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::Request; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::McpOAuthClientRegistration; +use super::start_authorization; +use crate::oauth::validate_authorization_server_endpoints; +use crate::oauth_http_client::OAuthHttpClientAdapter; +use crate::utils::MCP_USER_AGENT; +use crate::utils::build_default_headers; + +const CALLBACK_ID: &str = "abc123ABC_-x"; + +async fn oauth_server(overrides: Value) -> MockServer { + let server = MockServer::start().await; + let base_url = server.uri(); + let mut metadata = json!({ + "authorization_endpoint": format!("{base_url}/authorize"), + "token_endpoint": format!("{base_url}/token"), + "registration_endpoint": format!("{base_url}/register"), + "client_id_metadata_document_supported": true, + "token_endpoint_auth_methods_supported": ["none"], + "code_challenge_methods_supported": ["S256"], + "scopes_supported": ["read", "offline_access"], + }); + metadata + .as_object_mut() + .expect("metadata should be an object") + .extend( + overrides + .as_object() + .expect("overrides should be an object") + .clone(), + ); + if metadata["authorization_response_iss_parameter_supported"] == json!(true) + && metadata.get("issuer").is_none() + { + metadata["issuer"] = json!(format!("{base_url}/mcp")); + } + + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(metadata)) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/register")) + .respond_with(|request: &Request| { + let registration: Value = serde_json::from_slice(&request.body) + .expect("dynamic registration should contain JSON"); + ResponseTemplate::new(200).set_body_json(json!({ + "client_id": "dcr-client", + "redirect_uris": registration["redirect_uris"], + })) + }) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "test-access-token", + "token_type": "Bearer", + "refresh_token": "test-refresh-token", + }))) + .mount(&server) + .await; + + server +} + +async fn requests_to(server: &MockServer, request_path: &str) -> Vec { + server + .received_requests() + .await + .expect("mock server should record requests") + .into_iter() + .filter(|request| request.url.path() == request_path) + .collect() +} + +async fn authorization( + server: &MockServer, + redirect_uri: &str, + registration: McpOAuthClientRegistration, +) -> Result<(OAuthState, HashMap)> { + let prepared = start_authorization( + &format!("{}/mcp", server.uri()), + Arc::new(OAuthHttpClientAdapter::new( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + HeaderMap::new(), + &format!("{}/mcp", server.uri()), + )), + &["read"], + redirect_uri, + CALLBACK_ID, + registration, + ) + .await?; + let state = prepared.oauth_state; + let query = Url::parse(&state.get_authorization_url().await?)? + .query_pairs() + .into_owned() + .collect(); + + Ok((state, query)) +} + +#[tokio::test] +async fn automatic_cimd_uses_stable_or_callback_specific_identity() -> Result<()> { + for (host, supports_issuer, expected_client_id, expected_redirect) in [ + ( + "127.0.0.1", + false, + "https://chatgpt.com/oauth/codex/abc123ABC_-x/client.json", + "http://127.0.0.1:43123/callback/abc123ABC_-x", + ), + ( + "localhost", + false, + "https://chatgpt.com/oauth/codex/abc123ABC_-x/client.json", + "http://localhost:43123/callback/abc123ABC_-x", + ), + ( + "127.0.0.1", + true, + "https://chatgpt.com/oauth/codex/client.json", + "http://127.0.0.1:43123/callback", + ), + ] { + let server = oauth_server(json!({ + "authorization_response_iss_parameter_supported": supports_issuer, + })) + .await; + let redirect = format!("http://{host}:43123/callback"); + let (mut state, query) = + authorization(&server, &redirect, McpOAuthClientRegistration::Auto).await?; + assert_eq!(query["client_id"], expected_client_id); + assert_eq!(query["redirect_uri"], expected_redirect); + assert_eq!(query["code_challenge_method"], "S256"); + assert_eq!(query["scope"], "read offline_access"); + + state + .handle_callback_with_issuer( + "valid-authorization-code", + &query["state"], + supports_issuer + .then(|| format!("{}/mcp", server.uri())) + .as_deref(), + ) + .await?; + let token_requests = requests_to(&server, "/token").await; + assert_eq!(token_requests.len(), 1); + let request = &token_requests[0]; + let body: HashMap<_, _> = url::form_urlencoded::parse(&request.body) + .into_owned() + .collect(); + assert_eq!(body["client_id"], expected_client_id); + assert_eq!(body["redirect_uri"], expected_redirect); + assert_eq!(body["grant_type"], "authorization_code"); + assert!(body.contains_key("code_verifier")); + assert!(!body.contains_key("client_secret")); + assert!(!request.headers.contains_key("authorization")); + assert!(requests_to(&server, "/register").await.is_empty()); + assert_eq!( + requests_to(&server, "/.well-known/oauth-authorization-server/mcp") + .await + .len(), + 1 + ); + } + + Ok(()) +} + +#[tokio::test] +async fn registration_selection_preserves_dcr_capabilities_and_exact_redirects() -> Result<()> { + let native = "http://localhost:43123/callback/abc123ABC_-x"; + let shared_native = "http://localhost:43123/callback"; + let custom = "https://callbacks.example.com/oauth/callback/abc123ABC_-x"; + for (metadata, redirect, registration, expected_redirect) in [ + ( + json!({"client_id_metadata_document_supported": false}), + native, + McpOAuthClientRegistration::Auto, + native, + ), + ( + json!({"token_endpoint_auth_methods_supported": null}), + native, + McpOAuthClientRegistration::Auto, + native, + ), + ( + json!({"token_endpoint_auth_methods_supported": ["private_key_jwt"]}), + native, + McpOAuthClientRegistration::Auto, + native, + ), + (json!({}), custom, McpOAuthClientRegistration::Auto, custom), + ( + json!({"authorization_response_iss_parameter_supported": true}), + shared_native, + McpOAuthClientRegistration::Dcr, + shared_native, + ), + ] { + let server = oauth_server(metadata).await; + let (_, query) = authorization(&server, redirect, registration).await?; + assert_eq!(query["client_id"], "dcr-client"); + assert_eq!(query["redirect_uri"], expected_redirect); + let registrations = requests_to(&server, "/register").await; + assert_eq!(registrations.len(), 1); + let registration: Value = serde_json::from_slice(®istrations[0].body)?; + assert_eq!(registration["redirect_uris"], json!([expected_redirect])); + } + + Ok(()) +} + +#[test] +fn legacy_provider_exceptions_require_exact_issuer_and_endpoint_origins() -> Result<()> { + for (issuer, authorization_endpoint, token_endpoint, accepted) in [ + ( + "https://api.figma.com", + "https://www.figma.com/oauth/mcp", + "https://api.figma.com/v1/oauth/token", + true, + ), + ( + "https://agent.robinhood.com/mcp/trading", + "https://robinhood.com/oauth", + "https://api.robinhood.com/oauth2/token/", + true, + ), + ( + "https://api.figma.com.attacker.example", + "https://www.figma.com/oauth/mcp", + "https://api.figma.com.attacker.example/token", + false, + ), + ( + "https://api.figma.com", + "https://www.figma.com/oauth/mcp", + "https://attacker.example/token", + false, + ), + ( + "http://api.figma.com", + "https://www.figma.com/oauth/mcp", + "http://api.figma.com/v1/oauth/token", + false, + ), + ( + "https://agent.robinhood.com/mcp/attacker", + "https://robinhood.com/oauth", + "https://api.robinhood.com/oauth2/token/", + false, + ), + ( + "https://agent.robinhood.com/mcp/trading", + "https://robinhood.com.attacker.example/oauth", + "https://api.robinhood.com/oauth2/token/", + false, + ), + ] { + let metadata: AuthorizationMetadata = serde_json::from_value(json!({ + "issuer": issuer, + "authorization_endpoint": authorization_endpoint, + "token_endpoint": token_endpoint, + }))?; + assert_eq!( + validate_authorization_server_endpoints(&metadata).is_ok(), + accepted, + "unexpected validation result for issuer {issuer}", + ); + } + + Ok(()) +} + +#[tokio::test] +async fn verified_issuer_can_delegate_authorization_and_token_to_one_origin() -> Result<()> { + let issuer_server = MockServer::start().await; + let endpoint_server = MockServer::start().await; + let issuer = format!("{}/mcp", issuer_server.uri()); + let redirect_uri = format!("http://127.0.0.1:43123/callback/{CALLBACK_ID}"); + + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "issuer": issuer, + "authorization_endpoint": format!("{}/authorize", endpoint_server.uri()), + "token_endpoint": format!("{}/token", endpoint_server.uri()), + "registration_endpoint": format!("{}/register", endpoint_server.uri()), + "token_endpoint_auth_methods_supported": ["none"], + "code_challenge_methods_supported": ["S256"], + }))) + .expect(1) + .mount(&issuer_server) + .await; + Mock::given(method("POST")) + .and(path("/register")) + .respond_with(|request: &Request| { + let registration: Value = serde_json::from_slice(&request.body) + .expect("dynamic registration should contain JSON"); + ResponseTemplate::new(200).set_body_json(json!({ + "client_id": "delegated-provider-client", + "redirect_uris": registration["redirect_uris"], + })) + }) + .expect(1) + .mount(&endpoint_server) + .await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "delegated-provider-token", + "token_type": "Bearer", + }))) + .expect(1) + .mount(&endpoint_server) + .await; + + let (mut state, query) = authorization( + &issuer_server, + &redirect_uri, + McpOAuthClientRegistration::Dcr, + ) + .await?; + assert_eq!(query["client_id"], "delegated-provider-client"); + assert_eq!( + Url::parse(&state.get_authorization_url().await?)?.origin(), + Url::parse(&endpoint_server.uri())?.origin(), + ); + + state + .handle_callback_with_issuer("valid-authorization-code", &query["state"], None) + .await?; + let token_requests = requests_to(&endpoint_server, "/token").await; + assert_eq!(token_requests.len(), 1); + assert!( + url::form_urlencoded::parse(&token_requests[0].body) + .any(|(name, _)| name == "code_verifier") + ); + issuer_server.verify().await; + endpoint_server.verify().await; + Ok(()) +} + +#[tokio::test] +async fn resource_headers_follow_same_origin_registration_redirect_and_sdk_auth_wins() -> Result<()> +{ + const RESOURCE_AUTHORIZATION: &str = "Bearer resource-only-secret"; + const RESOURCE_API_KEY: &str = "resource-api-key-secret"; + // These dummy OAuth credentials exist only in this test's local mock server. + const DUMMY_CLIENT_ID: &str = "dummy-test-client-id"; + const DUMMY_CLIENT_SECRET: &str = "dummy-test-client-secret"; + let sdk_authorization = format!( + "Basic {}", + STANDARD.encode(format!("{DUMMY_CLIENT_ID}:{DUMMY_CLIENT_SECRET}")) + ); + + let server = MockServer::start().await; + let authorization_server = MockServer::start().await; + let resource_url = format!("{}/mcp", server.uri()); + let resource_metadata_url = format!("{}/resource-metadata", server.uri()); + let redirect_uri = format!("http://127.0.0.1:43123/callback/{CALLBACK_ID}"); + + Mock::given(method("GET")) + .and(path("/mcp")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "resource": resource_url, + "authorization_servers": [authorization_server.uri()], + }))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server")) + .and(header("user-agent", MCP_USER_AGENT)) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "issuer": authorization_server.uri(), + "authorization_endpoint": format!("{}/authorize", authorization_server.uri()), + "token_endpoint": format!("{}/token", server.uri()), + "registration_endpoint": format!("{}/register", server.uri()), + "token_endpoint_auth_methods_supported": ["client_secret_basic"], + "code_challenge_methods_supported": ["S256"], + }))) + .expect(1) + .mount(&authorization_server) + .await; + Mock::given(method("POST")) + .and(path("/register")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("content-type", "application/json")) + .respond_with(ResponseTemplate::new(307).insert_header("location", "/register/")) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/register/")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .respond_with(ResponseTemplate::new(201).set_body_json(json!({ + "client_id": DUMMY_CLIENT_ID, + "client_secret": DUMMY_CLIENT_SECRET, + "redirect_uris": [redirect_uri], + }))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/token")) + .and(header("authorization", sdk_authorization)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "test-access-token", + "token_type": "Bearer", + }))) + .expect(1) + .mount(&server) + .await; + + let default_headers = build_default_headers( + Some(HashMap::from([ + ( + "Authorization".to_string(), + RESOURCE_AUTHORIZATION.to_string(), + ), + ("X-Api-Key".to_string(), RESOURCE_API_KEY.to_string()), + ])), + /*env_http_headers*/ None, + )?; + let prepared = start_authorization( + &resource_url, + Arc::new(OAuthHttpClientAdapter::new( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + default_headers, + &resource_url, + )), + &[], + &redirect_uri, + CALLBACK_ID, + McpOAuthClientRegistration::Dcr, + ) + .await?; + assert_eq!( + prepared.authorization_server_issuer.as_deref(), + Some(authorization_server.uri().as_str()) + ); + let mut state = prepared.oauth_state; + let csrf_state = Url::parse(&state.get_authorization_url().await?)? + .query_pairs() + .find(|(name, _)| name == "state") + .expect("authorization request should contain a CSRF state") + .1 + .into_owned(); + state + .handle_callback_with_issuer("valid-authorization-code", &csrf_state, None) + .await?; + + let authorization_requests = authorization_server + .received_requests() + .await + .expect("authorization server should record requests"); + assert_eq!(authorization_requests.len(), 1); + assert_eq!(authorization_requests[0].headers.get("authorization"), None); + assert_eq!(authorization_requests[0].headers.get("x-api-key"), None); + server.verify().await; + authorization_server.verify().await; + Ok(()) +} + +#[tokio::test] +async fn invalid_cimd_metadata_and_redirects_fail_without_dynamic_registration() { + let valid = "http://127.0.0.1:43123/callback/abc123ABC_-x"; + for (metadata, redirect, expected_error) in [ + ( + json!({ + "authorization_response_iss_parameter_supported": true, + "issuer": null, + }), + valid, + "issuer-bound callbacks require an authorization server issuer", + ), + ( + json!({"token_endpoint_auth_methods_supported": ["private_key_jwt"]}), + valid, + "token endpoint auth method `none`", + ), + ( + json!({}), + "http://127.0.0.1.evil.example:43123/callback/abc123ABC_-x", + "ephemeral loopback callback", + ), + ( + json!({}), + "http://127.0.0.1/callback/abc123ABC_-x", + "ephemeral loopback callback", + ), + ( + json!({}), + "http://127.0.0.1:43123/callback/wrong-id", + "ephemeral loopback callback", + ), + ( + json!({}), + "http://127.0.0.1:43123/callback/abc123ABC_-x?unexpected=true", + "ephemeral loopback callback", + ), + ( + json!({}), + "http://[::1]:43123/callback/abc123ABC_-x", + "ephemeral loopback callback", + ), + ] { + let server = oauth_server(metadata).await; + let error = authorization(&server, redirect, McpOAuthClientRegistration::Cimd) + .await + .err() + .expect("invalid CIMD metadata or callback should fail"); + assert!(error.to_string().contains(expected_error)); + assert!(requests_to(&server, "/register").await.is_empty()); + assert!(requests_to(&server, "/token").await.is_empty()); + } + + let server = oauth_server(json!({"registration_endpoint": null})).await; + let error = authorization(&server, valid, McpOAuthClientRegistration::Dcr) + .await + .err() + .expect("explicit DCR should require an advertised registration endpoint"); + assert!(error.to_string().contains("registration not supported")); +} diff --git a/codex-rs/rmcp-client/src/oauth_http_client.rs b/codex-rs/rmcp-client/src/oauth_http_client.rs new file mode 100644 index 0000000000000000000000000000000000000000..f827e1720f647e7fc50b989beaed59131434170c --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_http_client.rs @@ -0,0 +1,496 @@ +use std::sync::Arc; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; +use std::time::Duration; +use std::time::Instant; + +use codex_exec_server::HttpClient; +use codex_exec_server::HttpHeader; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use http::HeaderMap; +use http::HeaderValue; +use http::Method; +use http::StatusCode; +use http::header::AUTHORIZATION; +use http::header::CONTENT_ENCODING; +use http::header::CONTENT_LENGTH; +use http::header::CONTENT_TYPE; +use http::header::LOCATION; +use http::header::TRANSFER_ENCODING; +use http::header::USER_AGENT; +use oauth2::HttpRequest; +use oauth2::HttpResponse; +use rmcp::transport::auth::AuthorizationMetadata; +use rmcp::transport::auth::OAuthHttpClient; +use rmcp::transport::auth::OAuthHttpClientError; +use rmcp::transport::auth::OAuthHttpClientFuture; +use rmcp::transport::auth::OAuthHttpRedirectPolicy; +use rmcp::transport::auth::OAuthHttpRequest; +use url::Origin; +use url::Url; + +use crate::auth_status::OAuthDiscoveryTimeout; +use crate::http_client_adapter::StreamableHttpRedirectMode; +use crate::utils::MCP_USER_AGENT; + +const MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES: usize = 1024 * 1024; +const MAX_OAUTH_HTTP_REDIRECTS: usize = 10; +static NEXT_OAUTH_REQUEST_ID: AtomicU64 = AtomicU64::new(0); + +tokio::task_local! { + /// Bounds provider HTTP work during preparation, excluding credential lock waits and saves. + pub(crate) static PROACTIVE_REFRESH_TIMEOUT: Duration; +} + +#[derive(Debug, thiserror::Error)] +enum OAuthHttpClientAdapterError { + #[error("unsupported OAuth HTTP redirect policy")] + UnsupportedRedirectPolicy, + #[error("OAuth HTTP request timed out")] + TimedOut, + #[error("OAuth HTTP request exceeded {MAX_OAUTH_HTTP_REDIRECTS} redirects")] + TooManyRedirects, + #[error("OAuth HTTP response body exceeds {maximum_bytes} bytes")] + ResponseBodyTooLarge { maximum_bytes: usize }, + #[error("OAuth authorization server issuer does not match authorization metadata origin")] + AuthorizationMetadataIssuerOriginMismatch, +} + +fn oauth_http_client_error( + error: impl std::error::Error + Send + Sync + 'static, +) -> OAuthHttpClientError { + Box::new(error) +} + +#[derive(Clone)] +pub(crate) struct OAuthHttpClientAdapter { + http_client: Arc, + default_headers: HeaderMap, + resource_origin: Origin, + timeout: OAuthDiscoveryTimeout, + has_configured_headers: bool, + redirect_mode: StreamableHttpRedirectMode, +} + +impl OAuthHttpClientAdapter { + /// Recover only a candidate-local 503 without turning failed discovery into + /// missing metadata (which would enable RMCP's legacy endpoint fallback). + async fn execute_with_metadata_fallback( + &self, + request: HttpRequest, + redirect_policy: OAuthHttpRedirectPolicy, + timeout: Option, + ) -> Result { + let issuer_path = request + .uri() + .path() + .strip_prefix("/.well-known/oauth-authorization-server") + .filter(|suffix| suffix.is_empty() || suffix.starts_with('/')); + let Some(issuer_path) = issuer_path.filter(|_| { + request.method() == Method::GET + && matches!(redirect_policy, OAuthHttpRedirectPolicy::Stop) + }) else { + return self + .execute_request(request, redirect_policy, timeout) + .await; + }; + let mut candidates = vec![format!("/.well-known/openid-configuration{issuer_path}")]; + if !issuer_path.is_empty() { + candidates.push(format!("{issuer_path}/.well-known/openid-configuration")); + } + let mut candidate_url = + Url::parse(&request.uri().to_string()).map_err(oauth_http_client_error)?; + candidate_url.set_query(None); + candidate_url.set_fragment(None); + let operation = async { + let original = self + .execute_request(request.clone(), redirect_policy, timeout) + .await?; + if original.status() != StatusCode::SERVICE_UNAVAILABLE { + return Ok(original); + } + // These are the same issuer's OIDC candidates, in RMCP's discovery + // order. Reuse the adapter's header, origin and body-size checks. + // Do not follow redirects here: RMCP owns discovery redirect policy. + for path in candidates { + candidate_url.set_path(&path); + let mut candidate = request.clone(); + *candidate.uri_mut() = candidate_url + .as_str() + .parse() + .map_err(oauth_http_client_error)?; + let response = self + .execute_request(candidate, OAuthHttpRedirectPolicy::Stop, timeout) + .await?; + match response.status() { + StatusCode::OK => { + if serde_json::from_slice::(response.body()).is_ok() + { + // RMCP still validates the exact expected issuer + // before using these endpoints or saved credentials. + return Ok(response); + } + } + StatusCode::NOT_FOUND + | StatusCode::METHOD_NOT_ALLOWED + | StatusCode::SERVICE_UNAVAILABLE => {} + StatusCode::REQUEST_TIMEOUT + | StatusCode::TOO_EARLY + | StatusCode::TOO_MANY_REQUESTS => return Ok(response), + status if status.is_server_error() => return Ok(response), + _ => return Ok(original), + } + } + Ok(original) + }; + // Additional candidates share the original request's time budget. + let timeout = match self.timeout { + OAuthDiscoveryTimeout::Requested => timeout, + OAuthDiscoveryTimeout::Capped(cap) => { + Some(timeout.map_or(cap, |timeout| timeout.min(cap))) + } + }; + match timeout { + Some(timeout) => tokio::time::timeout(timeout, operation) + .await + .map_err(|_| oauth_http_client_error(OAuthHttpClientAdapterError::TimedOut))?, + None => operation.await, + } + } + + #[cfg(test)] + pub(crate) fn new( + http_client: Arc, + default_headers: HeaderMap, + resource_url: &str, + ) -> Self { + Self::new_with_redirect_mode( + http_client, + default_headers, + resource_url, + /*has_configured_headers*/ false, + StreamableHttpRedirectMode::Legacy, + ) + .expect("OAuth resource URL should be valid") + } + + pub(crate) fn new_with_redirect_mode( + http_client: Arc, + default_headers: HeaderMap, + resource_url: &str, + has_configured_headers: bool, + redirect_mode: StreamableHttpRedirectMode, + ) -> Result { + Ok(Self { + http_client, + default_headers, + resource_origin: Url::parse(resource_url)?.origin(), + timeout: OAuthDiscoveryTimeout::Requested, + has_configured_headers, + redirect_mode, + }) + } + + pub(crate) fn new_with_max_timeout_and_redirect_mode( + http_client: Arc, + default_headers: HeaderMap, + resource_url: &str, + max_timeout: Duration, + has_configured_headers: bool, + redirect_mode: StreamableHttpRedirectMode, + ) -> Result { + Ok(Self { + http_client, + default_headers, + resource_origin: Url::parse(resource_url)?.origin(), + timeout: OAuthDiscoveryTimeout::Capped(max_timeout), + has_configured_headers, + redirect_mode, + }) + } + + pub(crate) async fn execute_request( + &self, + request: HttpRequest, + redirect_policy: OAuthHttpRedirectPolicy, + timeout: Option, + ) -> Result { + let redirect_policy = match redirect_policy { + OAuthHttpRedirectPolicy::Follow => HttpRedirectPolicy::Follow, + OAuthHttpRedirectPolicy::Stop => HttpRedirectPolicy::Stop, + _ => { + return Err(oauth_http_client_error( + OAuthHttpClientAdapterError::UnsupportedRedirectPolicy, + )); + } + }; + let (parts, body) = request.into_parts(); + let mut request_url = + Url::parse(&parts.uri.to_string()).map_err(oauth_http_client_error)?; + let is_resource_origin = request_url.origin() == self.resource_origin; + let mut headers = if is_resource_origin { + self.default_headers.clone() + } else { + HeaderMap::new() + }; + for name in parts.headers.keys() { + headers.remove(name); + } + let has_resource_only_headers = is_resource_origin + && headers.iter().any(|(name, value)| { + name != USER_AGENT || value != HeaderValue::from_static(MCP_USER_AGENT) + }); + headers.extend(parts.headers); + if !is_resource_origin { + headers.insert(USER_AGENT, HeaderValue::from_static(MCP_USER_AGENT)); + } + let redirect_policy = oauth_redirect_policy( + self.redirect_mode, + &headers, + self.has_configured_headers, + redirect_policy, + ); + // The executor can only follow every redirect or none, so replay credentialed + // requests ourselves after checking that each destination stays on the resource origin. + let follow_same_origin_redirects = + has_resource_only_headers && redirect_policy == HttpRedirectPolicy::Follow; + + let headers = headers + .iter() + .map(|(name, value)| { + Ok(HttpHeader { + name: name.as_str().to_string(), + value: value.to_str().map_err(oauth_http_client_error)?.to_string(), + value_env_var: None, + }) + }) + .collect::, OAuthHttpClientError>>()?; + let timeout = match self.timeout { + OAuthDiscoveryTimeout::Requested => timeout, + OAuthDiscoveryTimeout::Capped(max_timeout) => { + Some(timeout.map_or(max_timeout, |timeout| timeout.min(max_timeout))) + } + }; + let timeout_ms = timeout.map(|timeout| { + u64::try_from(timeout.as_millis()) + .unwrap_or(u64::MAX) + .max(1) + }); + let deadline = timeout.map(|timeout| Instant::now() + timeout); + let mut params = HttpRequestParams { + method: parts.method.to_string(), + url: parts.uri.to_string(), + headers, + body: (!body.is_empty()).then_some(body.into()), + timeout_ms, + redirect_policy: if follow_same_origin_redirects { + HttpRedirectPolicy::Stop + } else { + redirect_policy + }, + request_id: String::new(), + stream_response: true, + }; + let mut redirects = 0; + let (response, body) = loop { + let request_id = NEXT_OAUTH_REQUEST_ID.fetch_add(1, Ordering::Relaxed); + params.request_id = format!("oauth-request-{request_id}"); + let (response, mut body_stream) = self + .http_client + .http_request_stream(params.clone()) + .await + .map_err(oauth_http_client_error)?; + let mut body = Vec::new(); + while let Some(chunk) = body_stream.recv().await.map_err(oauth_http_client_error)? { + if chunk.len() > MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES - body.len() { + return Err(oauth_http_client_error( + OAuthHttpClientAdapterError::ResponseBodyTooLarge { + maximum_bytes: MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES, + }, + )); + } + body.extend_from_slice(&chunk); + } + let Ok(status) = StatusCode::from_u16(response.status) else { + break (response, body); + }; + if !follow_same_origin_redirects + || !matches!( + status, + StatusCode::MOVED_PERMANENTLY + | StatusCode::FOUND + | StatusCode::SEE_OTHER + | StatusCode::TEMPORARY_REDIRECT + | StatusCode::PERMANENT_REDIRECT + ) + { + break (response, body); + } + let Some(next_url) = response + .headers + .iter() + .find(|header| header.name.eq_ignore_ascii_case(LOCATION.as_str())) + .and_then(|header| request_url.join(&header.value).ok()) + .filter(|url| url.origin() == self.resource_origin) + else { + break (response, body); + }; + if redirects >= MAX_OAUTH_HTTP_REDIRECTS { + return Err(oauth_http_client_error( + OAuthHttpClientAdapterError::TooManyRedirects, + )); + } + if status == StatusCode::SEE_OTHER + || matches!(status, StatusCode::MOVED_PERMANENTLY | StatusCode::FOUND) + && params.method == Method::POST.as_str() + { + if params.method != Method::HEAD.as_str() { + params.method = Method::GET.to_string(); + } + params.body = None; + params.headers.retain(|header| { + ![ + CONTENT_TYPE, + CONTENT_LENGTH, + CONTENT_ENCODING, + TRANSFER_ENCODING, + ] + .iter() + .any(|name| header.name.eq_ignore_ascii_case(name.as_str())) + }); + } + params.url = next_url.to_string(); + if let Some(deadline) = deadline { + let remaining = + deadline + .checked_duration_since(Instant::now()) + .ok_or_else(|| { + oauth_http_client_error(OAuthHttpClientAdapterError::TimedOut) + })?; + params.timeout_ms = Some( + u64::try_from(remaining.as_millis()) + .unwrap_or(u64::MAX) + .max(1), + ); + } + request_url = next_url; + redirects += 1; + }; + if response.status == StatusCode::OK.as_u16() + && let Ok(metadata) = serde_json::from_slice::(&body) + && let Some(issuer) = metadata.issuer.as_deref() + && Url::parse(issuer) + .map_err(oauth_http_client_error)? + .origin() + != request_url.origin() + { + return Err(oauth_http_client_error( + OAuthHttpClientAdapterError::AuthorizationMetadataIssuerOriginMismatch, + )); + } + let mut builder = oauth2::http::Response::builder().status(response.status); + for header in response.headers { + builder = builder.header(header.name, header.value); + } + builder.body(body).map_err(oauth_http_client_error) + } +} + +fn oauth_redirect_policy( + mode: StreamableHttpRedirectMode, + headers: &HeaderMap, + has_configured_headers: bool, + requested_policy: HttpRedirectPolicy, +) -> HttpRedirectPolicy { + if mode == StreamableHttpRedirectMode::AgentPluginV1 + && (has_configured_headers || headers.contains_key(AUTHORIZATION)) + { + HttpRedirectPolicy::Stop + } else { + requested_policy + } +} + +impl OAuthHttpClient for OAuthHttpClientAdapter { + fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + Box::pin(async move { + let operation = self.execute_with_metadata_fallback( + request.request, + request.redirect_policy, + request.timeout, + ); + match PROACTIVE_REFRESH_TIMEOUT.try_with(|duration| *duration) { + Ok(duration) => tokio::time::timeout(duration, operation) + .await + .map_err(|_| oauth_http_client_error(OAuthHttpClientAdapterError::TimedOut))?, + Err(_) => operation.await, + } + }) + } +} + +#[cfg(test)] +#[path = "oauth_http_client_security_tests.rs"] +mod security_tests; + +#[cfg(test)] +mod tests { + use http::HeaderValue; + use pretty_assertions::assert_eq; + + use super::*; + + #[test] + fn agent_plugin_oauth_stops_only_for_sensitive_headers() { + assert_eq!( + policy( + StreamableHttpRedirectMode::AgentPluginV1, + /*has_configured_headers*/ true, + /*has_authorization*/ false, + ), + HttpRedirectPolicy::Stop + ); + assert_eq!( + policy( + StreamableHttpRedirectMode::AgentPluginV1, + /*has_configured_headers*/ false, + /*has_authorization*/ true, + ), + HttpRedirectPolicy::Stop + ); + assert_eq!( + policy( + StreamableHttpRedirectMode::AgentPluginV1, + /*has_configured_headers*/ false, + /*has_authorization*/ false, + ), + HttpRedirectPolicy::Follow + ); + assert_eq!( + policy( + StreamableHttpRedirectMode::Legacy, + /*has_configured_headers*/ true, + /*has_authorization*/ true, + ), + HttpRedirectPolicy::Follow + ); + } + + fn policy( + mode: StreamableHttpRedirectMode, + has_configured_headers: bool, + has_authorization: bool, + ) -> HttpRedirectPolicy { + let mut headers = HeaderMap::new(); + if has_authorization { + headers.insert(AUTHORIZATION, HeaderValue::from_static("Bearer secret")); + } + oauth_redirect_policy( + mode, + &headers, + has_configured_headers, + HttpRedirectPolicy::Follow, + ) + } +} diff --git a/codex-rs/rmcp-client/src/oauth_http_client_security_tests.rs b/codex-rs/rmcp-client/src/oauth_http_client_security_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..d4b8a64444bbeccbd1050c09683136a33046a76b --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_http_client_security_tests.rs @@ -0,0 +1,444 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Result; +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +use codex_exec_server::RouteAwareHttpClient; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; +use rmcp::transport::auth::AuthorizationManager; +use rmcp::transport::auth::AuthorizationMetadata; +use rmcp::transport::auth::OAuthHttpRedirectPolicy; +use serde_json::json; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES; +use super::OAuthHttpClientAdapter; +use crate::http_client_adapter::StreamableHttpRedirectMode; +use crate::utils::MCP_USER_AGENT; +use crate::utils::build_default_headers; + +async fn metadata_manager( + resource: &MockServer, + authorization: &MockServer, + issuer_path: &str, +) -> Result { + let resource_url = format!("{}/mcp", resource.uri()); + Mock::given(method("GET")) + .and(path("/mcp")) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!( + "Bearer resource_metadata=\"{}/resource-metadata\"", + resource.uri() + ), + )) + .mount(resource) + .await; + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "resource": resource_url, + "authorization_servers": [format!("{}{issuer_path}", authorization.uri())] + }))) + .mount(resource) + .await; + let adapter = OAuthHttpClientAdapter::new( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + build_default_headers( + Some(HashMap::from([( + "X-Api-Key".to_string(), + "resource-secret".to_string(), + )])), + /*env_http_headers*/ None, + )?, + &resource_url, + ); + Ok(AuthorizationManager::new_with_oauth_http_client(resource_url, Arc::new(adapter)).await?) +} + +#[tokio::test] +async fn oauth_metadata_503_falls_back_to_oidc() -> Result<()> { + for (issuer_path, first_oidc_response) in [ + ("", None), + ("/adfs", Some(ResponseTemplate::new(404))), + ("/adfs", Some(ResponseTemplate::new(405))), + ("/adfs", Some(ResponseTemplate::new(503))), + ( + "/adfs", + Some(ResponseTemplate::new(200).set_body_string("not metadata")), + ), + ] { + let resource = MockServer::start().await; + let authorization = MockServer::start().await; + let manager = metadata_manager(&resource, &authorization, issuer_path).await?; + let oauth_path = format!("/.well-known/oauth-authorization-server{issuer_path}"); + let oidc_path = format!("{issuer_path}/.well-known/openid-configuration"); + let metadata = json!({ + "issuer": format!("{}{issuer_path}", authorization.uri()), + "authorization_endpoint": format!("{}/authorize", authorization.uri()), + "token_endpoint": format!("{}/token", authorization.uri()) + }); + Mock::given(method("GET")) + .and(path(oauth_path.clone())) + .respond_with(ResponseTemplate::new(503)) + .expect(1) + .mount(&authorization) + .await; + if let Some(response) = first_oidc_response { + Mock::given(method("GET")) + .and(path(format!( + "/.well-known/openid-configuration{issuer_path}" + ))) + .respond_with(response) + .expect(1) + .mount(&authorization) + .await; + } + Mock::given(method("GET")) + .and(path(oidc_path.clone())) + .respond_with(ResponseTemplate::new(200).set_body_json(&metadata)) + .expect(1) + .mount(&authorization) + .await; + + let resolved = manager.resolve_metadata().await?; + let expected: AuthorizationMetadata = serde_json::from_value(metadata)?; + assert_eq!( + serde_json::to_value(resolved.metadata)?, + serde_json::to_value(expected)? + ); + let requests = authorization.received_requests().await.unwrap(); + let mut expected_paths = vec![oauth_path]; + if !issuer_path.is_empty() { + expected_paths.push(format!("/.well-known/openid-configuration{issuer_path}")); + } + expected_paths.push(oidc_path); + assert_eq!( + requests + .iter() + .map(|request| request.url.path().to_string()) + .collect::>(), + expected_paths + ); + assert!( + requests + .iter() + .all(|request| !request.headers.contains_key("x-api-key")) + ); + authorization.verify().await; + } + Ok(()) +} + +#[tokio::test] +async fn oauth_metadata_fallback_preserves_terminal_failures() -> Result<()> { + let resource = MockServer::start().await; + let authorization = MockServer::start().await; + let manager = metadata_manager(&resource, &authorization, "/adfs").await?; + let redirect_target = MockServer::start().await; + let wrong_issuer = json!({ + "issuer": format!("{}/other-tenant", authorization.uri()), + "authorization_endpoint": format!("{}/authorize", authorization.uri()), + "token_endpoint": format!("{}/token", authorization.uri()) + }); + for (fallback, expected_error, request_count) in [ + (ResponseTemplate::new(404), "503", 3), + ( + ResponseTemplate::new(200).set_body_json(wrong_issuer), + "issuer mismatch", + 2, + ), + ( + ResponseTemplate::new(307).insert_header("location", redirect_target.uri()), + "503", + 2, + ), + (ResponseTemplate::new(408), "408", 2), + (ResponseTemplate::new(425), "425", 2), + (ResponseTemplate::new(429), "429", 2), + (ResponseTemplate::new(500), "500", 2), + ] { + authorization.reset().await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/adfs")) + .respond_with(ResponseTemplate::new(503)) + .expect(1) + .mount(&authorization) + .await; + Mock::given(method("GET")) + .and(path("/.well-known/openid-configuration/adfs")) + .respond_with(fallback) + .expect(1) + .mount(&authorization) + .await; + let error = manager + .resolve_metadata() + .await + .expect_err("discovery must fail closed"); + assert!(error.to_string().contains(expected_error), "{error}"); + assert_eq!( + authorization.received_requests().await.unwrap().len(), + request_count + ); + authorization.verify().await; + } + assert_eq!(redirect_target.received_requests().await.unwrap().len(), 0); + Ok(()) +} + +#[tokio::test] +async fn oauth_metadata_fallback_does_not_retry_other_requests_or_statuses() -> Result<()> { + for (verb, request_path, status) in [ + ("POST", "/.well-known/oauth-authorization-server/adfs", 503), + ("GET", "/.well-known/oauth-protected-resource/mcp", 503), + ("GET", "/.well-known/oauth-authorization-server-other", 503), + ("GET", "/.well-known/oauth-authorization-server/adfs", 500), + ] { + let server = MockServer::start().await; + Mock::given(method(verb)) + .and(path(request_path)) + .respond_with(ResponseTemplate::new(status)) + .mount(&server) + .await; + let adapter = OAuthHttpClientAdapter::new( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + Default::default(), + &format!("{}/mcp", server.uri()), + ); + let response = adapter + .execute_with_metadata_fallback( + oauth2::http::Request::builder() + .method(verb) + .uri(format!("{}{request_path}", server.uri())) + .body(Vec::new())?, + OAuthHttpRedirectPolicy::Stop, + Some(Duration::from_secs(/*secs*/ 5)), + ) + .await + .map_err(|error| anyhow::anyhow!(error))?; + assert_eq!(response.status().as_u16(), status); + assert_eq!(server.received_requests().await.unwrap().len(), 1); + } + Ok(()) +} + +struct DelayedMetadataHttpClient; + +impl HttpClient for DelayedMetadataHttpClient { + fn http_request( + &self, + _params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + Box::pin(async { panic!("OAuth requests must stream responses") }) + } + + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + Box::pin(async move { + // Each response fits the per-request timeout. Only the shared + // deadline can prevent the fallback from finishing. + tokio::time::sleep(Duration::from_millis(/*millis*/ 100)).await; + Ok(( + HttpRequestResponse { + status: if params.url.contains("oauth-authorization-server") { + 503 + } else { + 404 + }, + headers: Vec::new(), + body: Vec::new().into(), + }, + HttpResponseBodyStream::from_chunks(Vec::new()), + )) + }) + } +} + +#[tokio::test] +async fn oauth_metadata_fallback_shares_the_original_timeout() -> Result<()> { + let adapter = OAuthHttpClientAdapter::new( + Arc::new(DelayedMetadataHttpClient), + Default::default(), + "https://resource.example/mcp", + ); + let error = adapter + .execute_with_metadata_fallback( + oauth2::http::Request::builder() + .method("GET") + .uri("https://issuer.example/.well-known/oauth-authorization-server/adfs") + .body(Vec::new())?, + OAuthHttpRedirectPolicy::Stop, + Some(Duration::from_millis(/*millis*/ 150)), + ) + .await + .expect_err("fallback requests must share the original deadline"); + assert!(error.to_string().contains("timed out"), "{error}"); + Ok(()) +} + +#[tokio::test] +async fn oauth_registration_redirects_never_forward_resource_only_headers() -> Result<()> { + const RESOURCE_API_KEY: &str = "resource-api-key-secret"; + const RESOURCE_USER_AGENT: &str = "resource-only-user-agent"; + + for (redirect_mode, has_resource_only_headers) in [ + (StreamableHttpRedirectMode::Legacy, true), + (StreamableHttpRedirectMode::AgentPluginV1, true), + (StreamableHttpRedirectMode::Legacy, false), + ] { + let resource_server = MockServer::start().await; + let redirect_target = MockServer::start().await; + let resource_url = format!("{}/mcp", resource_server.uri()); + + Mock::given(method("POST")) + .and(path("/register")) + .and(header("content-type", "application/json")) + .and(header( + "user-agent", + if has_resource_only_headers { + RESOURCE_USER_AGENT + } else { + MCP_USER_AGENT + }, + )) + .respond_with(ResponseTemplate::new(307).insert_header( + "location", + format!("{}/redirected-register", redirect_target.uri()), + )) + .expect(1) + .mount(&resource_server) + .await; + Mock::given(method("POST")) + .and(path("/redirected-register")) + .and(header("content-type", "application/json")) + .and(header("user-agent", MCP_USER_AGENT)) + .respond_with(ResponseTemplate::new(201)) + .expect(u64::from(!has_resource_only_headers)) + .mount(&redirect_target) + .await; + + let configured_headers = if has_resource_only_headers { + HashMap::from([ + ("X-Api-Key".to_string(), RESOURCE_API_KEY.to_string()), + ("User-Agent".to_string(), RESOURCE_USER_AGENT.to_string()), + ]) + } else { + HashMap::from([( + "Content-Type".to_string(), + "resource-only-content-type".to_string(), + )]) + }; + let adapter = OAuthHttpClientAdapter::new_with_redirect_mode( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + build_default_headers(Some(configured_headers), /*env_http_headers*/ None)?, + &resource_url, + /*has_configured_headers*/ true, + redirect_mode, + )?; + let response = adapter + .execute_request( + oauth2::http::Request::builder() + .method("POST") + .uri(format!("{}/register", resource_server.uri())) + .header("content-type", "application/json") + .body(br#"{"client_name":"Codex"}"#.to_vec())?, + OAuthHttpRedirectPolicy::Follow, + /*timeout*/ None, + ) + .await + .map_err(|error| anyhow::anyhow!(error))?; + + assert_eq!( + response.status(), + if has_resource_only_headers { + oauth2::http::StatusCode::TEMPORARY_REDIRECT + } else { + oauth2::http::StatusCode::CREATED + } + ); + resource_server.verify().await; + redirect_target.verify().await; + } + + Ok(()) +} + +#[tokio::test] +async fn same_origin_redirects_preserve_timeout_and_response_body_limits() -> Result<()> { + for oversized_redirect_body in [false, true] { + let server = MockServer::start().await; + let resource_url = format!("{}/mcp", server.uri()); + let redirect = ResponseTemplate::new(307).insert_header("location", "/register/"); + let redirect = if oversized_redirect_body { + redirect.set_body_bytes(vec![0; MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES + 1]) + } else { + redirect.set_delay(Duration::from_millis(/*millis*/ 400)) + }; + Mock::given(method("POST")) + .and(path("/register")) + .respond_with(redirect) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/register/")) + .respond_with( + ResponseTemplate::new(201).set_delay(Duration::from_millis(/*millis*/ 400)), + ) + .expect(u64::from(!oversized_redirect_body)) + .mount(&server) + .await; + + let adapter = OAuthHttpClientAdapter::new( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + build_default_headers( + Some(HashMap::from([( + "X-Api-Key".to_string(), + "resource-api-key-secret".to_string(), + )])), + /*env_http_headers*/ None, + )?, + &resource_url, + ); + let error = adapter + .execute_request( + oauth2::http::Request::builder() + .method("POST") + .uri(format!("{}/register", server.uri())) + .body(Vec::new())?, + OAuthHttpRedirectPolicy::Follow, + (!oversized_redirect_body).then_some(Duration::from_millis(/*millis*/ 700)), + ) + .await + .expect_err("redirects must preserve request timeout and response body limits"); + if oversized_redirect_body { + assert!(error.to_string().contains("exceeds")); + } + server.verify().await; + } + + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/oauth_refresh_mode.rs b/codex-rs/rmcp-client/src/oauth_refresh_mode.rs new file mode 100644 index 0000000000000000000000000000000000000000..84b2a98370a43e194dc2f69229ed29302ef3d63c --- /dev/null +++ b/codex-rs/rmcp-client/src/oauth_refresh_mode.rs @@ -0,0 +1,11 @@ +//! Selects the owner of MCP OAuth refresh and credential persistence. + +/// MCP OAuth policy pinned for the lifetime of a connection. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum McpOAuthRefreshMode { + /// Keep Codex's existing refresh and persistence path. + #[default] + Legacy, + /// Let RMCP coordinate refresh through Codex's credential store. + Coordinated, +} diff --git a/codex-rs/rmcp-client/src/perform_oauth_login.rs b/codex-rs/rmcp-client/src/perform_oauth_login.rs new file mode 100644 index 0000000000000000000000000000000000000000..a1b094a041c38fe92553abcc5dfa127b11661e3f --- /dev/null +++ b/codex-rs/rmcp-client/src/perform_oauth_login.rs @@ -0,0 +1,1536 @@ +use std::collections::HashMap; +use std::net::SocketAddr; +use std::string::String; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use codex_exec_server::HttpClient; +use rmcp::transport::AuthorizationManager; +use rmcp::transport::AuthorizationSession; +use rmcp::transport::auth::AuthorizationMetadata; +use rmcp::transport::auth::OAuthClientConfig; +use rmcp::transport::auth::OAuthHttpClient; +use rmcp::transport::auth::OAuthState; +use tiny_http::Response; +use tiny_http::Server; +use tokio::sync::oneshot; +use tokio::time::timeout; +use url::Url; +use urlencoding::decode; + +use crate::StoredOAuthTokens; +use crate::WrappedOAuthTokenResponse; +use crate::enterprise_oauth_login::enterprise_authorization_url; +use crate::enterprise_oauth_login::enterprise_callback_settings; +use crate::enterprise_oauth_login::resolve_enterprise_authorization_manager; +use crate::http_client_adapter::StreamableHttpRedirectMode; +use crate::oauth::compute_expires_at_millis; +use crate::oauth::validate_authorization_server_endpoints; +use crate::oauth_callback::McpOAuthCallbackMode; +use crate::oauth_callback::append_callback_id_to_redirect_uri; +use crate::oauth_callback::callback_id_from_server_url; +use crate::oauth_callback::callback_mode; +use crate::oauth_callback::resolve_mcp_oauth_callback_url; +use crate::oauth_callback::validate_callback_redirect; +use crate::oauth_client_registration::McpOAuthClientRegistration; +use crate::oauth_client_registration::PreparedOAuthLogin; +use crate::oauth_client_registration::start_authorization as start_client_registration; +use crate::oauth_http_client::OAuthHttpClientAdapter; +use crate::save_oauth_tokens; +use crate::utils::build_default_headers; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; + +#[path = "oauth_callback_input.rs"] +mod callback_input; +pub use callback_input::perform_oauth_login_with_callback_input; + +#[derive(Clone, Copy)] +pub(crate) enum OAuthLoginPurpose { + Mcp, + EnterpriseIdp, +} + +pub(crate) struct OAuthHttpContext { + pub(crate) http_headers: Option>, + pub(crate) env_http_headers: Option>, + pub(crate) http_client: Arc, + pub(crate) redirect_mode: StreamableHttpRedirectMode, +} + +struct CallbackServerGuard { + server: Arc, +} + +impl Drop for CallbackServerGuard { + fn drop(&mut self) { + self.server.unblock(); + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct OAuthProviderError { + error: Option, + error_description: Option, +} + +impl OAuthProviderError { + pub fn new(error: Option, error_description: Option) -> Self { + Self { + error, + error_description, + } + } +} + +impl std::fmt::Display for OAuthProviderError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match (self.error.as_deref(), self.error_description.as_deref()) { + (Some(error), Some(error_description)) => { + write!(f, "OAuth provider returned `{error}`: {error_description}") + } + (Some(error), None) => write!(f, "OAuth provider returned `{error}`"), + (None, Some(error_description)) => write!(f, "OAuth error: {error_description}"), + (None, None) => write!(f, "OAuth provider returned an error"), + } + } +} + +impl std::error::Error for OAuthProviderError {} + +#[allow(clippy::too_many_arguments)] +pub async fn perform_oauth_login( + server_name: &str, + server_url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_headers: Option>, + env_http_headers: Option>, + scopes: &[String], + oauth_client_id: Option<&str>, + client_registration: McpOAuthClientRegistration, + oauth_resource: Option<&str>, + callback_port: Option, + callback_url: Option<&str>, + global_callback_url: Option<&str>, + http_client: Arc, +) -> Result<()> { + perform_oauth_login_with_browser_output( + server_name, + server_url, + store_mode, + keyring_backend_kind, + http_headers, + env_http_headers, + scopes, + oauth_client_id, + client_registration, + oauth_resource, + callback_port, + callback_url, + global_callback_url, + http_client, + /*emit_browser_url*/ true, + StreamableHttpRedirectMode::Legacy, + ) + .await +} + +#[allow(clippy::too_many_arguments)] +pub async fn perform_oauth_login_silent( + server_name: &str, + server_url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_headers: Option>, + env_http_headers: Option>, + scopes: &[String], + oauth_client_id: Option<&str>, + client_registration: McpOAuthClientRegistration, + oauth_resource: Option<&str>, + callback_port: Option, + callback_url: Option<&str>, + global_callback_url: Option<&str>, + http_client: Arc, + redirect_mode: StreamableHttpRedirectMode, +) -> Result<()> { + perform_oauth_login_with_browser_output( + server_name, + server_url, + store_mode, + keyring_backend_kind, + http_headers, + env_http_headers, + scopes, + oauth_client_id, + client_registration, + oauth_resource, + callback_port, + callback_url, + global_callback_url, + http_client, + /*emit_browser_url*/ false, + redirect_mode, + ) + .await +} + +#[allow(clippy::too_many_arguments)] +async fn perform_oauth_login_with_browser_output( + server_name: &str, + server_url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_headers: Option>, + env_http_headers: Option>, + scopes: &[String], + oauth_client_id: Option<&str>, + client_registration: McpOAuthClientRegistration, + oauth_resource: Option<&str>, + callback_port: Option, + callback_url: Option<&str>, + global_callback_url: Option<&str>, + http_client: Arc, + emit_browser_url: bool, + redirect_mode: StreamableHttpRedirectMode, +) -> Result<()> { + let http_context = OAuthHttpContext { + http_headers, + env_http_headers, + http_client, + redirect_mode, + }; + OauthLoginFlow::new( + server_name, + server_url, + store_mode, + keyring_backend_kind, + http_context, + scopes, + oauth_client_id, + OAuthLoginPurpose::Mcp, + client_registration, + oauth_resource, + /*launch_browser*/ true, + callback_port, + callback_url, + global_callback_url, + /*timeout_secs*/ None, + ) + .await? + .finish(emit_browser_url) + .await +} + +#[allow(clippy::too_many_arguments)] +pub async fn perform_oauth_login_return_url( + server_name: &str, + server_url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_headers: Option>, + env_http_headers: Option>, + scopes: &[String], + oauth_client_id: Option<&str>, + client_registration: McpOAuthClientRegistration, + oauth_resource: Option<&str>, + timeout_secs: Option, + callback_port: Option, + callback_url: Option<&str>, + global_callback_url: Option<&str>, + http_client: Arc, + redirect_mode: StreamableHttpRedirectMode, +) -> Result { + let http_context = OAuthHttpContext { + http_headers, + env_http_headers, + http_client, + redirect_mode, + }; + let flow = OauthLoginFlow::new( + server_name, + server_url, + store_mode, + keyring_backend_kind, + http_context, + scopes, + oauth_client_id, + OAuthLoginPurpose::Mcp, + client_registration, + oauth_resource, + /*launch_browser*/ false, + callback_port, + callback_url, + global_callback_url, + timeout_secs, + ) + .await?; + + let authorization_url = flow.authorization_url(); + let completion = flow.spawn(); + + Ok(OauthLoginHandle::new(authorization_url, completion)) +} + +fn spawn_callback_server( + server: Arc, + tx: oneshot::Sender, + expected_callback_path: String, +) { + tokio::task::spawn_blocking(move || { + while let Ok(request) = server.recv() { + let path = request.url().to_string(); + match parse_oauth_callback(&path, &expected_callback_path) { + CallbackOutcome::Success(OauthCallbackResult { + code, + state, + issuer, + }) => { + let response = Response::from_string( + "Authentication complete. You may close this window.", + ); + if let Err(err) = request.respond(response) { + eprintln!("Failed to respond to OAuth callback: {err}"); + } + if let Err(message) = send_oauth_callback( + tx, + CallbackResult::Success(OauthCallbackResult { + code, + state, + issuer, + }), + ) { + eprintln!("{message}"); + } + break; + } + CallbackOutcome::Error(error) => { + let response = Response::from_string(error.to_string()).with_status_code(400); + if let Err(err) = request.respond(response) { + eprintln!("Failed to respond to OAuth callback: {err}"); + } + if let Err(message) = send_oauth_callback(tx, CallbackResult::Error(error)) { + eprintln!("{message}"); + } + break; + } + CallbackOutcome::Invalid => { + let response = + Response::from_string("Invalid OAuth callback").with_status_code(400); + if let Err(err) = request.respond(response) { + eprintln!("Failed to respond to OAuth callback: {err}"); + } + } + } + } + }); +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct OauthCallbackResult { + code: String, + state: String, + issuer: Option, +} + +#[derive(Debug)] +enum CallbackResult { + Success(OauthCallbackResult), + Error(OAuthProviderError), +} + +fn send_oauth_callback( + tx: oneshot::Sender, + result: CallbackResult, +) -> std::result::Result<(), &'static str> { + tx.send(result) + .map_err(|_| "OAuth callback receiver closed") +} + +#[derive(Debug, PartialEq, Eq)] +enum CallbackOutcome { + Success(OauthCallbackResult), + Error(OAuthProviderError), + Invalid, +} + +fn parse_oauth_callback(path: &str, expected_callback_path: &str) -> CallbackOutcome { + let Some((route, query)) = path.split_once('?') else { + return CallbackOutcome::Invalid; + }; + if route != expected_callback_path { + return CallbackOutcome::Invalid; + } + + let mut code = None; + let mut state = None; + let mut error = None; + let mut error_description = None; + let mut issuer = None; + + for pair in query.split('&') { + let Some((key, value)) = pair.split_once('=') else { + continue; + }; + let Ok(decoded) = decode(value) else { + continue; + }; + let decoded = decoded.into_owned(); + match key { + "code" => code = Some(decoded), + "state" => state = Some(decoded), + "error" => error = Some(decoded), + "error_description" => error_description = Some(decoded), + "iss" => issuer = Some(decoded), + _ => {} + } + } + + if let (Some(code), Some(state)) = (code, state) { + return CallbackOutcome::Success(OauthCallbackResult { + code, + state, + issuer, + }); + } + + if error.is_some() || error_description.is_some() { + return CallbackOutcome::Error(OAuthProviderError::new(error, error_description)); + } + + CallbackOutcome::Invalid +} + +pub struct OauthLoginHandle { + authorization_url: String, + completion: oneshot::Receiver>, +} + +impl OauthLoginHandle { + fn new(authorization_url: String, completion: oneshot::Receiver>) -> Self { + Self { + authorization_url, + completion, + } + } + + pub fn authorization_url(&self) -> &str { + &self.authorization_url + } + + pub fn into_parts(self) -> (String, oneshot::Receiver>) { + (self.authorization_url, self.completion) + } + + pub async fn wait(self) -> Result<()> { + self.completion + .await + .map_err(|err| anyhow!("OAuth login task was cancelled: {err}"))? + } +} + +pub(crate) struct OauthLoginFlow { + auth_url: String, + redirect_uri: String, + oauth_state: OAuthState, + authorization_server_issuer: Option, + rx: oneshot::Receiver, + guard: CallbackServerGuard, + server_name: String, + server_url: String, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + launch_browser: bool, + timeout: Duration, +} + +fn resolve_callback_port(callback_port: Option) -> Result> { + if let Some(config_port) = callback_port { + if config_port == 0 { + bail!( + "invalid MCP OAuth callback port `{config_port}`: port must be between 1 and 65535" + ); + } + return Ok(Some(config_port)); + } + + Ok(None) +} + +fn local_redirect_uri(server: &Server) -> Result { + match server.server_addr() { + tiny_http::ListenAddr::IP(std::net::SocketAddr::V4(addr)) => { + let ip = addr.ip(); + let port = addr.port(); + Ok(format!("http://{ip}:{port}/callback")) + } + tiny_http::ListenAddr::IP(std::net::SocketAddr::V6(addr)) => { + let ip = addr.ip(); + let port = addr.port(); + Ok(format!("http://[{ip}]:{port}/callback")) + } + #[cfg(not(target_os = "windows"))] + _ => Err(anyhow!("unable to determine callback address")), + } +} + +fn resolve_redirect_uri(server: &Server, callback_url: Option<&str>) -> Result { + let Some(callback_url) = callback_url else { + return local_redirect_uri(server); + }; + let mut parsed = Url::parse(callback_url) + .with_context(|| format!("invalid MCP OAuth callback URL `{callback_url}`"))?; + + // Registered loopback callbacks omit the temporary listener port because + // the OS can assign a different port on every login. Add the active port + // only to this authorization request; RFC 8252 requires authorization + // servers to accept any request-time port for loopback IP redirects. + // https://www.rfc-editor.org/rfc/rfc8252#section-7.3 + if parsed.scheme() == "http" + && parsed.host_str() == Some("127.0.0.1") + && parsed.port().is_none() + { + let listener_port = server + .server_addr() + .to_ip() + .ok_or_else(|| anyhow!("unable to determine OAuth callback listener port"))? + .port(); + parsed + .set_port(Some(listener_port)) + .map_err(|()| anyhow!("unable to set OAuth callback listener port"))?; + return Ok(parsed.to_string()); + } + + Ok(callback_url.to_string()) +} + +fn callback_path_from_redirect_uri(redirect_uri: &str) -> Result { + let parsed = Url::parse(redirect_uri) + .with_context(|| format!("invalid redirect URI `{redirect_uri}`"))?; + Ok(parsed.path().to_string()) +} + +fn callback_bind_host(callback_url: Option<&str>) -> &'static str { + let Some(callback_url) = callback_url else { + return "127.0.0.1"; + }; + + let Ok(parsed) = Url::parse(callback_url) else { + return "127.0.0.1"; + }; + + match parsed.host_str() { + Some("localhost" | "127.0.0.1" | "::1") | None => "127.0.0.1", + Some(_) => "0.0.0.0", + } +} + +impl OauthLoginFlow { + #[allow(clippy::too_many_arguments)] + pub(crate) async fn new( + server_name: &str, + server_url: &str, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_context: OAuthHttpContext, + scopes: &[String], + oauth_client_id: Option<&str>, + purpose: OAuthLoginPurpose, + client_registration: McpOAuthClientRegistration, + oauth_resource: Option<&str>, + launch_browser: bool, + callback_port: Option, + callback_url: Option<&str>, + global_callback_url: Option<&str>, + timeout_secs: Option, + ) -> Result { + const DEFAULT_OAUTH_TIMEOUT_SECS: i64 = 300; + + let callback_port = resolve_callback_port(callback_port)?; + let is_enterprise_idp = matches!(purpose, OAuthLoginPurpose::EnterpriseIdp); + let (enterprise_bind_ip, callback_port) = if is_enterprise_idp { + let (ip, port) = enterprise_callback_settings( + server_url, + oauth_client_id, + callback_url, + callback_port, + )?; + (Some(ip), port) + } else { + (None, callback_port) + }; + let callback_id = callback_id_from_server_url(server_url)?; + let oauth_client_id = oauth_client_id.filter(|client_id| !client_id.trim().is_empty()); + let configured_callback = if oauth_client_id.is_some() { + callback_url + .map(|callback_url| { + Url::parse(callback_url) + .with_context(|| format!("invalid MCP OAuth callback URL `{callback_url}`")) + }) + .transpose()? + } else { + None + }; + + let OAuthHttpContext { + http_headers, + env_http_headers, + http_client, + redirect_mode, + } = http_context; + let has_configured_headers = http_headers + .as_ref() + .is_some_and(|headers| !headers.is_empty()) + || env_http_headers + .as_ref() + .is_some_and(|headers| !headers.is_empty()); + let default_headers = build_default_headers(http_headers, env_http_headers)?; + let oauth_http_client: Arc = + Arc::new(OAuthHttpClientAdapter::new_with_redirect_mode( + http_client, + default_headers, + server_url, + has_configured_headers, + redirect_mode, + )?); + let registered_authorization = if oauth_client_id.is_some() { + Some( + resolve_authorization_manager(server_url, Arc::clone(&oauth_http_client), purpose) + .await?, + ) + } else { + None + }; + let registered_callback_mode = registered_authorization + .as_ref() + .map(|(_, metadata)| callback_mode(metadata)) + .transpose()?; + let use_legacy_fallback = !is_enterprise_idp + && registered_callback_mode == Some(McpOAuthCallbackMode::CallbackSpecific) + && configured_callback.as_ref().is_some_and(|callback_url| { + callback_url + .path_segments() + .and_then(|mut segments| segments.next_back()) + != Some(callback_id.as_str()) + }); + let callback_url = if use_legacy_fallback { + // Any preregistered client's callback can lack its required ID when + // the authorization server does not support issuer binding. This + // especially affects plugins, whose callbacks are configured before + // metadata discovery. Preserve compatibility and avoid making every + // login fail by using the global/default callback instead; its + // required server-specific callback ID is appended below. + global_callback_url + } else { + callback_url + }; + + let bind_ip = match enterprise_bind_ip { + Some(ip) => ip, + None => callback_bind_host(callback_url).parse()?, + }; + // Port zero asks the OS for a free ephemeral port; the resolved + // redirect receives that port after the listener has been bound. + let bind_addr = SocketAddr::new(bind_ip, callback_port.unwrap_or(0)); + let server = Arc::new(Server::http(bind_addr).map_err(|err| anyhow!(err))?); + let guard = CallbackServerGuard { + server: Arc::clone(&server), + }; + let redirect_uri = resolve_redirect_uri(&server, callback_url)?; + let redirect_uri = if is_enterprise_idp { + let listener_port = server + .server_addr() + .to_ip() + .ok_or_else(|| anyhow!("unable to determine enterprise callback listener port"))? + .port(); + let mut redirect = Url::parse(&redirect_uri)?; + redirect + .set_port(Some(listener_port)) + .map_err(|()| anyhow!("invalid enterprise callback port"))?; + redirect.to_string() + } else { + redirect_uri + }; + + let scope_refs: Vec<&str> = scopes.iter().map(String::as_str).collect(); + let PreparedOAuthLogin { + oauth_state, + authorization_server_issuer, + redirect_uri, + } = if let Some((oauth_client_id, (auth_manager, metadata))) = + oauth_client_id.zip(registered_authorization) + { + let redirect_uri = if is_enterprise_idp { + resolve_mcp_oauth_callback_url( + server_url, + Some(&redirect_uri), + callback_mode(&metadata)?, + )? + } else if callback_url.is_some() && !use_legacy_fallback { + redirect_uri + } else { + append_callback_id_to_redirect_uri(&redirect_uri, &callback_id)? + }; + start_authorization( + auth_manager, + metadata, + &scope_refs, + &redirect_uri, + &callback_id, + oauth_client_id, + purpose, + ) + .await? + } else { + start_client_registration( + server_url, + oauth_http_client, + &scope_refs, + &redirect_uri, + &callback_id, + client_registration, + ) + .await? + }; + let callback_path = callback_path_from_redirect_uri(&redirect_uri)?; + let (tx, rx) = oneshot::channel(); + spawn_callback_server(server, tx, callback_path); + let auth_url = append_query_param( + &oauth_state.get_authorization_url().await?, + "resource", + oauth_resource, + ); + let timeout_secs = timeout_secs.unwrap_or(DEFAULT_OAUTH_TIMEOUT_SECS).max(1); + let timeout = Duration::from_secs(timeout_secs as u64); + + Ok(Self { + auth_url, + redirect_uri, + oauth_state, + authorization_server_issuer, + rx, + guard, + server_name: server_name.to_string(), + server_url: server_url.to_string(), + store_mode, + keyring_backend_kind, + launch_browser, + timeout, + }) + } + + pub(crate) fn authorization_url(&self) -> String { + self.auth_url.clone() + } + + async fn finish(self, emit_browser_url: bool) -> Result<()> { + let store_mode = self.store_mode; + let keyring_backend_kind = self.keyring_backend_kind; + let stored = self.complete(emit_browser_url).await?; + save_oauth_tokens( + &stored.server_name, + &stored, + store_mode, + keyring_backend_kind, + ) + .await + } + + pub(crate) async fn complete(mut self, emit_browser_url: bool) -> Result { + if self.launch_browser { + let server_name = &self.server_name; + let auth_url = &self.auth_url; + if emit_browser_url { + println!( + "Authorize `{server_name}` by opening this URL in your browser:\n{auth_url}\n" + ); + } + + if webbrowser::open(auth_url).is_err() { + if !emit_browser_url { + eprintln!( + "Authorize `{server_name}` by opening this URL in your browser:\n{auth_url}\n" + ); + } + eprintln!("(Browser launch failed; please copy the URL above manually.)"); + } + } + + let callback = timeout(self.timeout, &mut self.rx) + .await + .context("timed out waiting for OAuth callback")? + .context("OAuth callback was cancelled")?; + self.complete_callback(callback).await + } + + async fn complete_callback(mut self, callback: CallbackResult) -> Result { + let result = async { + let OauthCallbackResult { + code, + state: csrf_state, + issuer, + } = match callback { + CallbackResult::Success(callback) => callback, + CallbackResult::Error(error) => return Err(anyhow!(error)), + }; + + self.oauth_state + .handle_callback_with_issuer(&code, &csrf_state, issuer.as_deref()) + .await + .context("failed to handle OAuth callback")?; + + let (client_id, credentials_opt) = self + .oauth_state + .get_credentials() + .await + .context("failed to retrieve OAuth credentials")?; + let credentials = credentials_opt + .ok_or_else(|| anyhow!("OAuth provider did not return credentials"))?; + let expires_at = compute_expires_at_millis(&credentials); + Ok(StoredOAuthTokens { + server_name: self.server_name.clone(), + url: self.server_url.clone(), + issuer: self.authorization_server_issuer.clone(), + client_id, + token_response: WrappedOAuthTokenResponse(credentials), + expires_at, + }) + } + .await; + + drop(self.guard); + result + } + + fn spawn(self) -> oneshot::Receiver> { + let server_name = self.server_name.clone(); + let (tx, rx) = oneshot::channel(); + + tokio::spawn(async move { + let result = self.finish(/*emit_browser_url*/ false).await; + if let Err(err) = &result { + eprintln!("Failed to complete OAuth login for '{server_name}': {err:#}"); + } + + let _ = tx.send(result); + }); + + rx + } +} + +async fn resolve_authorization_manager( + server_url: &str, + http_client: Arc, + purpose: OAuthLoginPurpose, +) -> Result<(AuthorizationManager, AuthorizationMetadata)> { + if matches!(purpose, OAuthLoginPurpose::EnterpriseIdp) { + return resolve_enterprise_authorization_manager(server_url, http_client).await; + } + let mut auth_manager = + AuthorizationManager::new_with_oauth_http_client(server_url, http_client).await?; + auth_manager.set_allow_missing_issuer(true); + let metadata = auth_manager.resolve_metadata().await?.metadata; + validate_authorization_server_endpoints(&metadata)?; + Ok((auth_manager, metadata)) +} + +async fn start_authorization( + mut auth_manager: AuthorizationManager, + metadata: AuthorizationMetadata, + scopes: &[&str], + redirect_uri: &str, + callback_id: &str, + oauth_client_id: &str, + purpose: OAuthLoginPurpose, +) -> Result { + let strict_enterprise_idp = matches!(purpose, OAuthLoginPurpose::EnterpriseIdp); + let authorization_server_issuer = metadata.issuer.clone(); + validate_callback_redirect(redirect_uri, callback_id, callback_mode(&metadata)?)?; + auth_manager.set_metadata(metadata); + let client_config = OAuthClientConfig::new(oauth_client_id, redirect_uri) + .with_scopes(scopes.iter().map(|scope| (*scope).to_string()).collect()); + auth_manager.configure_client(client_config)?; + let auth_url = auth_manager.get_authorization_url(scopes).await?; + let auth_url = if strict_enterprise_idp { + enterprise_authorization_url(&auth_url)? + } else { + auth_url + }; + + Ok(PreparedOAuthLogin { + oauth_state: OAuthState::Session(AuthorizationSession::for_scope_upgrade( + auth_manager, + auth_url, + redirect_uri, + )), + authorization_server_issuer, + redirect_uri: redirect_uri.to_string(), + }) +} + +fn append_query_param(url: &str, key: &str, value: Option<&str>) -> String { + let Some(value) = value else { + return url.to_string(); + }; + let value = value.trim(); + if value.is_empty() { + return url.to_string(); + } + if let Ok(mut parsed) = Url::parse(url) { + parsed.query_pairs_mut().append_pair(key, value); + return parsed.to_string(); + } + let encoded = urlencoding::encode(value); + let separator = if url.contains('?') { "&" } else { "?" }; + format!("{url}{separator}{key}={encoded}") +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + use std::io::Read; + use std::io::Write; + use std::net::TcpStream; + use std::sync::Arc; + use std::sync::atomic::AtomicUsize; + use std::sync::atomic::Ordering; + + use axum::Json; + use axum::Router; + use axum::routing::get; + use axum::routing::post; + use codex_config::types::AuthKeyringBackendKind; + use codex_config::types::OAuthCredentialsStoreMode; + use codex_exec_server::ExecServerError; + use codex_exec_server::HttpClient; + use codex_exec_server::HttpRequestParams; + use codex_exec_server::HttpRequestResponse; + use codex_exec_server::HttpResponseBodyStream; + use codex_exec_server::RouteAwareHttpClient; + use codex_http_client::HttpClientFactory; + use codex_http_client::OutboundProxyPolicy; + use futures::future::BoxFuture; + use http::HeaderMap; + use oauth2::TokenResponse; + use pretty_assertions::assert_eq; + use serde_json::json; + use tokio::net::TcpListener; + use url::Url; + + use super::CallbackOutcome; + use super::McpOAuthClientRegistration; + use super::OAuthHttpClientAdapter; + use super::OAuthHttpContext; + use super::OAuthLoginPurpose; + use super::OAuthProviderError; + use super::OauthLoginFlow; + use super::StreamableHttpRedirectMode; + use super::append_callback_id_to_redirect_uri; + use super::append_query_param; + use super::callback_id_from_server_url; + use super::callback_path_from_redirect_uri; + use super::parse_oauth_callback; + use super::perform_oauth_login; + use super::perform_oauth_login_silent; + use super::resolve_authorization_manager; + use super::start_authorization; + use crate::oauth::stored_oauth_credentials; + use crate::oauth::test_support::TempCodexHome; + + #[derive(Default)] + struct RecordingHttpClient { + requests: AtomicUsize, + } + + impl HttpClient for RecordingHttpClient { + fn http_request( + &self, + _params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + Box::pin(async { + Err(ExecServerError::HttpRequest( + "unexpected buffered OAuth request".to_string(), + )) + }) + } + + fn http_request_stream( + &self, + _params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> + { + self.requests.fetch_add(1, Ordering::SeqCst); + Box::pin(async { + Err(ExecServerError::HttpRequest( + "configured OAuth client was used".to_string(), + )) + }) + } + } + + async fn spawn_oauth_metadata_server() -> (String, Arc) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind metadata listener"); + let addr = listener.local_addr().expect("read metadata listener addr"); + let base_url = format!("http://{addr}"); + let metadata = json!({ + "issuer": format!("{base_url}/mcp"), + "authorization_endpoint": format!("{base_url}/oauth/authorize"), + "token_endpoint": format!("{base_url}/oauth/token"), + "registration_endpoint": format!("{base_url}/oauth/register"), + "scopes_supported": ["read", "offline_access"], + }); + let registration_requests = Arc::new(AtomicUsize::new(0)); + let captured_registration_requests = Arc::clone(®istration_requests); + let path_scoped_metadata = metadata.clone(); + let app = Router::new() + .route( + "/.well-known/oauth-authorization-server/mcp", + get(move || { + let metadata = path_scoped_metadata.clone(); + async move { Json(metadata) } + }), + ) + .route( + "/.well-known/oauth-authorization-server", + get(move || { + let metadata = metadata.clone(); + async move { Json(metadata) } + }), + ) + .route( + "/oauth/register", + post(move || { + let registration_requests = Arc::clone(&captured_registration_requests); + async move { + registration_requests.fetch_add(1, Ordering::SeqCst); + Json(json!({"client_id": "unexpected-dynamic-client"})) + } + }), + ) + .route( + "/oauth/token", + post(|| async { + Json(json!({ + "access_token": "test-access-token", + "token_type": "Bearer", + })) + }), + ); + + tokio::spawn(async move { + axum::serve(listener, app) + .await + .expect("serve oauth metadata"); + }); + + (base_url, registration_requests) + } + + async fn send_oauth_callback(callback_url: Url) -> anyhow::Result<()> { + tokio::task::spawn_blocking(move || -> anyhow::Result<()> { + let host = callback_url + .host_str() + .ok_or_else(|| anyhow::anyhow!("callback URL should include a host"))?; + let port = callback_url + .port() + .ok_or_else(|| anyhow::anyhow!("callback URL should include a port"))?; + let mut stream = TcpStream::connect((host, port))?; + let mut path = callback_url.path().to_string(); + if let Some(query) = callback_url.query() { + path.push('?'); + path.push_str(query); + } + write!( + stream, + "GET {path} HTTP/1.1\r\nHost: {host}:{port}\r\nConnection: close\r\n\r\n" + )?; + let mut response = String::new(); + stream.read_to_string(&mut response)?; + anyhow::ensure!( + response.starts_with("HTTP/1.1 200"), + "OAuth callback failed: {response}" + ); + Ok(()) + }) + .await? + } + + #[tokio::test] + async fn ordinary_oauth_login_persists_issuer_without_a_refresh_token() -> anyhow::Result<()> { + let _env = TempCodexHome::new(); + let (base_url, _registration_requests) = spawn_oauth_metadata_server().await; + let server_url = format!("{base_url}/mcp"); + let flow = OauthLoginFlow::new( + "issuer-persistence-test", + &server_url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + OAuthHttpContext { + http_headers: None, + env_http_headers: None, + http_client: Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + redirect_mode: StreamableHttpRedirectMode::Legacy, + }, + &[], + Some("test-client"), + OAuthLoginPurpose::Mcp, + McpOAuthClientRegistration::Auto, + /*oauth_resource*/ None, + /*launch_browser*/ false, + /*callback_port*/ None, + /*callback_url*/ None, + /*global_callback_url*/ None, + Some(/*timeout_secs*/ 5), + ) + .await?; + let authorization_url = Url::parse(&flow.authorization_url())?; + let query = authorization_url.query_pairs().collect::>(); + let redirect_uri = query + .get("redirect_uri") + .ok_or_else(|| anyhow::anyhow!("authorization URL should include redirect_uri"))?; + let state = query + .get("state") + .ok_or_else(|| anyhow::anyhow!("authorization URL should include state"))?; + let mut callback_url = Url::parse(redirect_uri)?; + callback_url + .query_pairs_mut() + .append_pair("code", "test-code") + .append_pair("state", state); + send_oauth_callback(callback_url).await?; + flow.finish(/*emit_browser_url*/ false).await?; + + let stored = stored_oauth_credentials( + "issuer-persistence-test", + &server_url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + )? + .expect("OAuth login should persist credentials"); + assert_eq!(stored.issuer.as_deref(), Some(server_url.as_str())); + assert!(stored.token_response.0.refresh_token().is_none()); + Ok(()) + } + + #[tokio::test] + async fn configured_client_preserves_exact_scopes_and_redirect_without_registration() { + for (scopes, expected_scope) in [(&[][..], None), (&["read"][..], Some("read"))] { + let (base_url, registration_requests) = spawn_oauth_metadata_server().await; + let redirect_uri = "http://127.0.0.1:43123/callback/configured-client"; + let (auth_manager, metadata) = resolve_authorization_manager( + &format!("{base_url}/mcp"), + Arc::new(OAuthHttpClientAdapter::new( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + HeaderMap::new(), + &format!("{base_url}/mcp"), + )), + OAuthLoginPurpose::Mcp, + ) + .await + .expect("resolve pre-registered OAuth metadata"); + let prepared = start_authorization( + auth_manager, + metadata, + scopes, + redirect_uri, + "configured-client", + "eci-prd-pub-codex-123", + OAuthLoginPurpose::Mcp, + ) + .await + .expect("start pre-registered OAuth authorization"); + let oauth_state = prepared.oauth_state; + + let authorization_url = oauth_state + .get_authorization_url() + .await + .expect("read authorization URL"); + let query = Url::parse(&authorization_url) + .expect("authorization URL should parse") + .query_pairs() + .into_owned() + .collect::>(); + + assert_eq!( + query.get("client_id").map(String::as_str), + Some("eci-prd-pub-codex-123") + ); + assert_eq!( + query.get("redirect_uri").map(String::as_str), + Some(redirect_uri) + ); + assert_eq!(query.get("scope").map(String::as_str), expected_scope); + assert_eq!(registration_requests.load(Ordering::SeqCst), 0); + } + } + #[tokio::test] + async fn oauth_callback_validates_rfc_9207_issuer_before_token_exchange() { + for (supports_issuer, callback_issuer, expected_token_requests) in [ + (true, Some("matching"), 1), + (true, Some("mismatched"), 0), + (true, None, 0), + (false, Some("mismatched"), 0), + (false, None, 1), + ] { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind authorization metadata listener"); + let issuer = format!( + "http://{}", + listener.local_addr().expect("read listener address") + ); + let token_requests = Arc::new(AtomicUsize::new(0)); + let captured_token_requests = Arc::clone(&token_requests); + let authorization_issuer = format!("{issuer}/mcp"); + let metadata = json!({ + "issuer": authorization_issuer, + "authorization_endpoint": format!("{issuer}/authorize"), + "token_endpoint": format!("{issuer}/token"), + "authorization_response_iss_parameter_supported": supports_issuer, + }); + let scoped_metadata = metadata.clone(); + let app = Router::new() + .route( + "/.well-known/oauth-authorization-server/mcp", + get(move || { + let metadata = scoped_metadata.clone(); + async move { Json(metadata) } + }), + ) + .route( + "/.well-known/oauth-authorization-server", + get(move || { + let metadata = metadata.clone(); + async move { Json(metadata) } + }), + ) + .route( + "/token", + post(move || { + let token_requests = Arc::clone(&captured_token_requests); + async move { + token_requests.fetch_add(1, Ordering::SeqCst); + Json(json!({"access_token":"test-token","token_type":"Bearer"})) + } + }), + ); + let server = tokio::spawn(async move { + axum::serve(listener, app) + .await + .expect("serve authorization metadata fixture"); + }); + let redirect_uri = if supports_issuer { + "http://127.0.0.1/callback" + } else { + "http://127.0.0.1/callback/test-callback" + }; + let (auth_manager, metadata) = resolve_authorization_manager( + &format!("{issuer}/mcp"), + Arc::new(OAuthHttpClientAdapter::new( + Arc::new(RouteAwareHttpClient::new(HttpClientFactory::new( + OutboundProxyPolicy::ReqwestDefault, + ))), + HeaderMap::new(), + &format!("{issuer}/mcp"), + )), + OAuthLoginPurpose::Mcp, + ) + .await + .expect("resolve issuer-aware authorization metadata"); + let prepared = start_authorization( + auth_manager, + metadata, + &[], + redirect_uri, + "test-callback", + "test-client", + OAuthLoginPurpose::Mcp, + ) + .await + .expect("start issuer-aware authorization"); + let mut state = prepared.oauth_state; + let csrf_state = Url::parse( + &state + .get_authorization_url() + .await + .expect("retrieve authorization URL"), + ) + .expect("parse authorization URL") + .query_pairs() + .find(|(key, _)| key == "state") + .map(|(_, value)| value.into_owned()) + .expect("authorization URL should contain state"); + let callback_issuer = match callback_issuer { + Some("matching") => Some(authorization_issuer.as_str()), + Some(_) => Some("https://unexpected.example"), + None => None, + }; + let result = state + .handle_callback_with_issuer("test-code", &csrf_state, callback_issuer) + .await; + + assert_eq!( + token_requests.load(Ordering::SeqCst), + expected_token_requests + ); + assert_eq!(result.is_ok(), expected_token_requests == 1); + + if expected_token_requests == 0 { + state + .handle_callback_with_issuer( + "legitimate-code", + &csrf_state, + Some(authorization_issuer.as_str()), + ) + .await + .expect("issuer validation failures must preserve OAuth authorization state"); + assert_eq!(token_requests.load(Ordering::SeqCst), 1); + } + server.abort(); + } + } + + #[tokio::test] + async fn interactive_oauth_login_uses_supplied_http_client() { + let http_client = Arc::new(RecordingHttpClient::default()); + perform_oauth_login( + "configured-client", + "http://127.0.0.1:1/mcp", + OAuthCredentialsStoreMode::default(), + AuthKeyringBackendKind::default(), + /*http_headers*/ None, + /*env_http_headers*/ None, + &[], + /*oauth_client_id*/ None, + McpOAuthClientRegistration::Auto, + /*oauth_resource*/ None, + /*callback_port*/ None, + /*callback_url*/ None, + /*global_callback_url*/ None, + http_client.clone(), + ) + .await + .expect_err("OAuth metadata discovery should fail through the supplied client"); + + assert!(http_client.requests.load(Ordering::SeqCst) > 0); + } + + #[tokio::test] + async fn silent_oauth_login_uses_supplied_http_client() { + let http_client = Arc::new(RecordingHttpClient::default()); + perform_oauth_login_silent( + "configured-client", + "http://127.0.0.1:1/mcp", + OAuthCredentialsStoreMode::default(), + AuthKeyringBackendKind::default(), + /*http_headers*/ None, + /*env_http_headers*/ None, + &[], + /*oauth_client_id*/ None, + McpOAuthClientRegistration::Auto, + /*oauth_resource*/ None, + /*callback_port*/ None, + /*callback_url*/ None, + /*global_callback_url*/ None, + http_client.clone(), + StreamableHttpRedirectMode::Legacy, + ) + .await + .expect_err("OAuth metadata discovery should fail through the supplied client"); + + assert!(http_client.requests.load(Ordering::SeqCst) > 0); + } + + #[test] + fn parse_oauth_callback_accepts_default_path() { + let parsed = parse_oauth_callback("/callback?code=abc&state=xyz", "/callback"); + assert!(matches!(parsed, CallbackOutcome::Success(_))); + } + + #[test] + fn parse_oauth_callback_preserves_rfc_9207_issuer() { + let parsed = parse_oauth_callback( + "/callback?code=abc&state=xyz&iss=https%3A%2F%2Fissuer.example", + "/callback", + ); + assert_eq!( + parsed, + CallbackOutcome::Success(super::OauthCallbackResult { + code: "abc".to_string(), + state: "xyz".to_string(), + issuer: Some("https://issuer.example".to_string()), + }) + ); + } + + #[test] + fn parse_oauth_callback_accepts_custom_path() { + let parsed = parse_oauth_callback("/oauth/callback?code=abc&state=xyz", "/oauth/callback"); + assert!(matches!(parsed, CallbackOutcome::Success(_))); + } + + #[test] + fn parse_oauth_callback_accepts_callback_id_path() { + let parsed = + parse_oauth_callback("/callback/abc123?code=abc&state=xyz", "/callback/abc123"); + assert!(matches!(parsed, CallbackOutcome::Success(_))); + } + + #[test] + fn parse_oauth_callback_rejects_missing_callback_id_path() { + let parsed = parse_oauth_callback("/callback?code=abc&state=xyz", "/callback/abc123"); + assert!(matches!(parsed, CallbackOutcome::Invalid)); + } + + #[test] + fn parse_oauth_callback_rejects_wrong_path() { + let parsed = parse_oauth_callback("/callback?code=abc&state=xyz", "/oauth/callback"); + assert!(matches!(parsed, CallbackOutcome::Invalid)); + } + + #[test] + fn parse_oauth_callback_returns_provider_error() { + let parsed = parse_oauth_callback( + "/callback?error=invalid_scope&error_description=scope%20rejected", + "/callback", + ); + + assert_eq!( + parsed, + CallbackOutcome::Error(OAuthProviderError::new( + Some("invalid_scope".to_string()), + Some("scope rejected".to_string()), + )) + ); + } + + #[test] + fn callback_path_comes_from_redirect_uri() { + let path = callback_path_from_redirect_uri("https://example.com/oauth/callback") + .expect("redirect URI should parse"); + assert_eq!(path, "/oauth/callback"); + } + + #[test] + fn callback_id_is_bound_to_server_url() { + let callback_id = callback_id_from_server_url("https://mcp.example.com/mcp?tenant=one") + .expect("server URL should parse"); + let same_without_fragment = + callback_id_from_server_url("https://mcp.example.com/mcp?tenant=one#unused") + .expect("server URL should parse"); + let different_path = callback_id_from_server_url("https://mcp.example.com/sse?tenant=one") + .expect("server URL should parse"); + let different_query = callback_id_from_server_url("https://mcp.example.com/mcp?tenant=two") + .expect("server URL should parse"); + let different_origin = callback_id_from_server_url("https://mcp.example.com:8443/mcp") + .expect("server URL should parse"); + + assert_eq!(callback_id, same_without_fragment); + assert_ne!(callback_id, different_path); + assert_ne!(callback_id, different_query); + assert_ne!(callback_id, different_origin); + assert_eq!(callback_id, "XuuuHAzzHOni"); + } + + #[test] + fn callback_id_is_appended_to_redirect_uri_path() { + let redirect_uri = + append_callback_id_to_redirect_uri("http://127.0.0.1:1234/callback", "abc123") + .expect("redirect URI should parse"); + + assert_eq!(redirect_uri, "http://127.0.0.1:1234/callback/abc123"); + assert_eq!( + append_callback_id_to_redirect_uri(&redirect_uri, "abc123") + .expect("resolved redirect URI should parse"), + redirect_uri + ); + } + + #[test] + fn callback_id_is_appended_before_redirect_uri_query() { + let redirect_uri = append_callback_id_to_redirect_uri( + "https://callbacks.example.com/oauth/callback?provider=github", + "abc123", + ) + .expect("redirect URI should parse"); + + assert_eq!( + redirect_uri, + "https://callbacks.example.com/oauth/callback/abc123?provider=github" + ); + } + + #[test] + fn portless_loopback_callbacks_use_the_active_listener_port() { + let server = tiny_http::Server::http("127.0.0.1:0").expect("start callback listener"); + let listener_port = server + .server_addr() + .to_ip() + .expect("resolve callback listener address") + .port(); + + for path in ["/callback", "/callback/callback-id", "/custom/callback"] { + let callback = format!("http://127.0.0.1{path}"); + assert_eq!( + super::resolve_redirect_uri(&server, Some(&callback)) + .expect("insert active listener port"), + format!("http://127.0.0.1:{listener_port}{path}") + ); + } + + for callback in [ + "http://localhost/callback", + "http://127.0.0.1:3080/callback", + "https://127.0.0.1/callback", + "https://devbox.example.com/callback", + ] { + assert_eq!( + super::resolve_redirect_uri(&server, Some(callback)) + .expect("preserve configured callback origin"), + callback + ); + } + } + + #[test] + fn append_query_param_adds_resource_to_absolute_url() { + let url = append_query_param( + "https://example.com/authorize?scope=read", + "resource", + Some("https://api.example.com"), + ); + + assert_eq!( + url, + "https://example.com/authorize?scope=read&resource=https%3A%2F%2Fapi.example.com" + ); + } + + #[test] + fn append_query_param_ignores_empty_values() { + let url = append_query_param( + "https://example.com/authorize?scope=read", + "resource", + Some(" "), + ); + + assert_eq!(url, "https://example.com/authorize?scope=read"); + } + + #[test] + fn append_query_param_handles_unparseable_url() { + let url = append_query_param("not a url", "resource", Some("api/resource")); + + assert_eq!(url, "not a url?resource=api%2Fresource"); + } +} diff --git a/codex-rs/rmcp-client/src/program_resolver.rs b/codex-rs/rmcp-client/src/program_resolver.rs new file mode 100644 index 0000000000000000000000000000000000000000..53db522166c22ebf02c7adea51448557f1f113eb --- /dev/null +++ b/codex-rs/rmcp-client/src/program_resolver.rs @@ -0,0 +1,261 @@ +//! Platform-specific program resolution for MCP server execution. +//! +//! This module provides a unified interface for resolving executable paths +//! across different operating systems. The key challenge it addresses is that +//! Windows cannot execute script files (e.g., `.cmd`, `.bat`) directly through +//! `Command::new()` without their file extensions, while Unix systems handle +//! scripts natively through shebangs. +//! +//! The `resolve` function abstracts these platform differences: +//! - On Unix: Returns the program unchanged (OS handles script execution) +//! - On Windows: Uses the `which` crate to resolve full paths including extensions + +use std::collections::HashMap; +use std::ffi::OsString; +use std::path::Path; + +/// Resolves a program to its executable path on Unix systems. +/// +/// Unix systems handle PATH resolution and script execution natively through +/// the kernel's shebang (`#!`) mechanism, so this function simply returns +/// the program name unchanged. +#[cfg(unix)] +pub fn resolve( + program: OsString, + _env: &HashMap, + _cwd: &Path, +) -> std::io::Result { + Ok(program) +} + +/// Resolves a program to its executable path on Windows systems. +/// +/// Windows requires explicit file extensions for script execution. This function +/// uses the `which` crate to search the `PATH` environment variable and find +/// the full path to the executable, including necessary script extensions +/// (`.cmd`, `.bat`, etc.) defined in `PATHEXT`. +/// +/// This enables tools like `npx`, `pnpm`, and `yarn` to work correctly on Windows +/// without requiring users to specify full paths or extensions in their configuration. +#[cfg(windows)] +pub fn resolve( + program: OsString, + env: &HashMap, + cwd: &Path, +) -> std::io::Result { + // Extract PATH from environment for search locations + let search_path = env.iter().find_map(|(name, value)| { + name.to_string_lossy() + .eq_ignore_ascii_case("PATH") + .then_some(value) + }); + + // Attempt resolution via which crate + match which::which_in(&program, search_path, cwd) { + Ok(resolved) => { + tracing::debug!("Resolved {program:?} to {resolved:?}"); + Ok(resolved.into_os_string()) + } + Err(e) => { + tracing::debug!("Failed to resolve {program:?}: {e}. Using original path"); + // Fallback to original program - let Command::new() handle the error + Ok(program) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::utils::create_env_for_mcp_server; + use anyhow::Result; + use std::fs; + use std::path::Path; + use tempfile::TempDir; + use tokio::process::Command; + + /// Unix: Verifies the OS handles script execution without file extensions. + #[cfg(unix)] + #[tokio::test] + async fn test_unix_executes_script_without_extension() -> Result<()> { + let env = TestExecutableEnv::new()?; + // Linux can transiently report ETXTBSY while the freshly written test + // script is becoming executable on the backing filesystem. + let mut retries = 0; + let output = loop { + let mut cmd = Command::new(&env.program_name); + cmd.envs(&env.mcp_env); + + let output = cmd.output().await; + if !output + .as_ref() + .is_err_and(|err| err.kind() == std::io::ErrorKind::ExecutableFileBusy) + || retries == 2 + { + break output; + } + retries += 1; + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + }; + + assert!( + output.is_ok(), + "Unix should execute PATH-resolved scripts directly: {output:?}" + ); + Ok(()) + } + + /// Windows: Verifies scripts fail to execute without the proper extension. + #[cfg(windows)] + #[tokio::test] + async fn test_windows_fails_without_extension() -> Result<()> { + let env = TestExecutableEnv::new()?; + let mut cmd = Command::new(&env.program_name); + cmd.envs(&env.mcp_env); + + let output = cmd.output().await; + assert!( + output.is_err(), + "Windows requires .cmd/.bat extension for direct execution" + ); + Ok(()) + } + + /// Windows: Verifies scripts with an explicit extension execute correctly. + #[cfg(windows)] + #[tokio::test] + async fn test_windows_succeeds_with_extension() -> Result<()> { + let env = TestExecutableEnv::new()?; + // Append the `.cmd` extension to the program name + let program_with_ext = format!("{}.cmd", env.program_name); + let mut cmd = Command::new(&program_with_ext); + cmd.envs(&env.mcp_env); + + let output = cmd.output().await; + assert!( + output.is_ok(), + "Windows should execute scripts when the extension is provided" + ); + Ok(()) + } + + /// Verifies program resolution enables successful execution on all platforms. + #[tokio::test] + async fn test_resolved_program_executes_successfully() -> Result<()> { + let env = TestExecutableEnv::new()?; + #[cfg(windows)] + let env = { + let mut env = env; + let path = env + .mcp_env + .remove(std::ffi::OsStr::new("PATH")) + .expect("test environment should include PATH"); + env.mcp_env.insert(OsString::from("Path"), path); + env + }; + let program = OsString::from(&env.program_name); + + // Apply platform-specific resolution + let resolved = resolve(program, &env.mcp_env, std::env::current_dir()?.as_path())?; + + // Verify resolved path executes successfully + let mut cmd = Command::new(resolved); + cmd.envs(&env.mcp_env); + let output = cmd.output().await; + + assert!( + output.is_ok(), + "Resolved program should execute successfully" + ); + Ok(()) + } + + // Test fixture for creating temporary executables in a controlled environment. + struct TestExecutableEnv { + // Held to prevent the temporary directory from being deleted. + _temp_dir: TempDir, + program_name: String, + mcp_env: HashMap, + } + + impl TestExecutableEnv { + const TEST_PROGRAM: &'static str = "test_mcp_server"; + + fn new() -> Result { + let temp_dir = TempDir::new()?; + let dir_path = temp_dir.path(); + + Self::create_executable(dir_path)?; + + // Build a clean environment with the temp dir in the PATH. + let mut extra_env = HashMap::new(); + extra_env.insert(OsString::from("PATH"), Self::build_path_env_var(dir_path)); + + #[cfg(windows)] + extra_env.insert(OsString::from("PATHEXT"), Self::ensure_cmd_extension()); + + let mcp_env = create_env_for_mcp_server(Some(extra_env), &[])?; + + Ok(Self { + _temp_dir: temp_dir, + program_name: Self::TEST_PROGRAM.to_string(), + mcp_env, + }) + } + + /// Creates a simple, platform-specific executable script. + fn create_executable(dir: &Path) -> Result<()> { + #[cfg(windows)] + { + let file = dir.join(format!("{}.cmd", Self::TEST_PROGRAM)); + fs::write(&file, "@echo off\nexit 0")?; + } + + #[cfg(unix)] + { + let file = dir.join(Self::TEST_PROGRAM); + fs::write(&file, "#!/bin/sh\nexit 0")?; + Self::set_executable(&file)?; + } + + Ok(()) + } + + #[cfg(unix)] + fn set_executable(path: &Path) -> Result<()> { + use std::os::unix::fs::PermissionsExt; + let mut perms = fs::metadata(path)?.permissions(); + perms.set_mode(0o755); + fs::set_permissions(path, perms)?; + Ok(()) + } + + /// Prepends the given directory to the system's PATH variable. + fn build_path_env_var(dir: &Path) -> OsString { + let mut path = OsString::from(dir.as_os_str()); + if let Some(current) = std::env::var_os("PATH") { + let sep = if cfg!(windows) { ";" } else { ":" }; + path.push(sep); + path.push(current); + } + path + } + + /// Ensures `.CMD` is in the `PATHEXT` variable on Windows for script discovery. + #[cfg(windows)] + fn ensure_cmd_extension() -> OsString { + let current = std::env::var_os("PATHEXT").unwrap_or_default(); + if current + .to_string_lossy() + .to_ascii_uppercase() + .contains(".CMD") + { + current + } else { + let mut path_ext = OsString::from(".CMD;"); + path_ext.push(current); + path_ext + } + } + } +} diff --git a/codex-rs/rmcp-client/src/protocol_mode.rs b/codex-rs/rmcp-client/src/protocol_mode.rs new file mode 100644 index 0000000000000000000000000000000000000000..656ba8f6e97eb2fb6ccd433827e17642f9daf028 --- /dev/null +++ b/codex-rs/rmcp-client/src/protocol_mode.rs @@ -0,0 +1,128 @@ +use std::ffi::OsStr; +use std::io; + +use rmcp::model::ProtocolVersion; +use rmcp::service::ClientLifecycleMode; + +/// MCP compatibility policy selected once when a Codex session is created. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum McpProtocolMode { + /// Preserve the existing MCP initialization and OAuth behavior. + #[default] + Legacy, + /// Allow the MCP 2026-07-28 discovery and request lifecycle. + V20260728, +} + +impl McpProtocolMode { + /// Returns the newest protocol version this compatibility policy can use. + pub fn preferred_protocol_version(self) -> ProtocolVersion { + match self { + Self::Legacy => ProtocolVersion::V_2025_06_18, + Self::V20260728 => ProtocolVersion::V_2026_07_28, + } + } + + pub(crate) fn client_lifecycle(self) -> ClientLifecycleMode { + match self { + Self::Legacy => ClientLifecycleMode::Initialize, + Self::V20260728 => ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_06_18), + }, + } + } + + pub(crate) fn stdio_mode(self, requested_version: Option<&OsStr>) -> io::Result { + match (self, requested_version) { + (Self::Legacy, _) => Ok(Self::Legacy), + (_, None) => Ok(Self::Legacy), + (Self::V20260728, Some(version)) if version == OsStr::new("2026-07-28") => { + Ok(Self::V20260728) + } + (_, Some(version)) => Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!( + "unsupported CODEX_MCP_PROTOCOL_VERSION `{}` for stdio MCP server; expected `2026-07-28`", + version.to_string_lossy() + ), + )), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn protocol_modes_select_compatible_sdk_lifecycles() { + assert_eq!( + McpProtocolMode::Legacy.preferred_protocol_version(), + ProtocolVersion::V_2025_06_18 + ); + assert_eq!( + McpProtocolMode::Legacy.client_lifecycle(), + ClientLifecycleMode::Initialize + ); + assert_eq!( + McpProtocolMode::V20260728.preferred_protocol_version(), + ProtocolVersion::V_2026_07_28 + ); + assert_eq!( + McpProtocolMode::V20260728.client_lifecycle(), + ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_06_18), + } + ); + } + + #[test] + fn stdio_requires_both_the_modern_feature_and_a_server_opt_in() { + let modern_version = Some(OsStr::new("2026-07-28")); + + assert_eq!( + McpProtocolMode::Legacy + .stdio_mode(/*requested_version*/ None) + .unwrap(), + McpProtocolMode::Legacy + ); + assert_eq!( + McpProtocolMode::Legacy.stdio_mode(modern_version).unwrap(), + McpProtocolMode::Legacy + ); + assert_eq!( + McpProtocolMode::V20260728 + .stdio_mode(/*requested_version*/ None) + .unwrap(), + McpProtocolMode::Legacy + ); + assert_eq!( + McpProtocolMode::V20260728 + .stdio_mode(modern_version) + .unwrap(), + McpProtocolMode::V20260728 + ); + } + + #[test] + fn stdio_rejects_unknown_protocol_markers() { + let error = McpProtocolMode::V20260728 + .stdio_mode(Some(OsStr::new("1999-01-01"))) + .expect_err("an unknown protocol marker must fail"); + + assert_eq!(error.kind(), io::ErrorKind::InvalidInput); + assert!(error.to_string().contains("1999-01-01")); + } + + #[test] + fn legacy_stdio_does_not_interpret_existing_protocol_markers() { + assert_eq!( + McpProtocolMode::Legacy + .stdio_mode(Some(OsStr::new("1999-01-01"))) + .unwrap(), + McpProtocolMode::Legacy + ); + } +} diff --git a/codex-rs/rmcp-client/src/rmcp_client.rs b/codex-rs/rmcp-client/src/rmcp_client.rs new file mode 100644 index 0000000000000000000000000000000000000000..004f55d0cf94e5c102283ae830d3743259cc2648 --- /dev/null +++ b/codex-rs/rmcp-client/src/rmcp_client.rs @@ -0,0 +1,1710 @@ +use std::collections::HashMap; +use std::ffi::OsStr; +use std::ffi::OsString; +use std::future::Future; +use std::io; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::OnceLock; +use std::sync::PoisonError; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; +use std::time::Instant; + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use codex_api::SharedAuthProvider; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::McpServerEnvVar; +use codex_exec_server::HttpClient; +use codex_keyring_store::DefaultKeyringStore; +use futures::FutureExt; +use futures::future::BoxFuture; +use http::HeaderMap; +use http::header::AUTHORIZATION; +use oauth2::TokenResponse; +use rmcp::model::CallToolRequestParams; +use rmcp::model::CallToolResult; +use rmcp::model::ClientNotification; +use rmcp::model::ClientRequest; +use rmcp::model::ContentBlock; +use rmcp::model::CustomNotification; +use rmcp::model::CustomRequest; +use rmcp::model::ElicitRequestParams; +use rmcp::model::ElicitResult; +use rmcp::model::ElicitationAction; +use rmcp::model::Extensions; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ListResourceTemplatesResult; +use rmcp::model::ListResourcesResult; +use rmcp::model::ListToolsResult; +use rmcp::model::MetaObject; +use rmcp::model::PaginatedRequestParams; +use rmcp::model::ProtocolVersion; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ReadResourceResult; +use rmcp::model::RequestId; +use rmcp::model::RequestMetaObject; +use rmcp::model::RequestParamsMeta; +use rmcp::model::ServerPeerInfo; +use rmcp::model::ServerResult; +use rmcp::model::Tool; +use rmcp::service::ClientCacheConfig; +use rmcp::service::ClientServiceExt; +use rmcp::service::RequestHandle; +use rmcp::service::RoleClient; +use rmcp::service::RunningService; +use rmcp::service::ServiceError; +use rmcp::transport::AuthorizationManager; +use rmcp::transport::StreamableHttpClientTransport; +use rmcp::transport::auth::AuthClient; +use rmcp::transport::auth::AuthError; +use rmcp::transport::streamable_http_client::StreamableHttpClientTransportConfig; +use rmcp::transport::streamable_http_client::StreamableHttpError; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Value; +use tokio::sync::Mutex; +use tokio::sync::Semaphore; +use tokio::sync::watch; +use tokio::time; +use tracing::instrument; +use tracing::warn; + +use crate::elicitation_client_service::ElicitationClientService; +use crate::event_notification_transport::capture_event_notifications; +use crate::event_notification_transport::event_notification_channel; +use crate::http_client_adapter::StreamableHttpClientAdapter; +use crate::http_client_adapter::StreamableHttpClientAdapterError; +use crate::http_client_adapter::StreamableHttpRedirectMode; +use crate::in_process_transport::InProcessTransportFactory; +use crate::oauth::OAuthCredentialStore; +use crate::oauth::OAuthPersistor; +use crate::oauth::OAuthRuntime; +use crate::oauth::ResolvedOAuthCredentialStore; +use crate::oauth::ResolvedOAuthTokens; +use crate::oauth::StoredOAuthTokens; +use crate::oauth::install_tokens_in_manager; +use crate::oauth::resolve_oauth_tokens_from_store_policy; +use crate::oauth::validate_refresh_token_issuer; +use crate::oauth_http_client::OAuthHttpClientAdapter; +use crate::oauth_refresh_mode::McpOAuthRefreshMode; +use crate::protocol_mode::McpProtocolMode; +use crate::startup_error::is_authentication_required_error; +use crate::stdio_server_launcher::StdioServerCommand; +use crate::stdio_server_launcher::StdioServerLauncher; +use crate::stdio_server_launcher::StdioServerProcessHandle; +use crate::stdio_server_launcher::StdioServerTransport; +use crate::utils::build_default_headers; +use codex_config::types::OAuthCredentialsStoreMode; + +#[path = "streamable_http_retry.rs"] +mod streamable_http_retry; + +use self::streamable_http_retry::HandshakeError; +use self::streamable_http_retry::STREAMABLE_HTTP_RETRY_DELAYS_MS; +use self::streamable_http_retry::sleep_with_retry_deadline; + +enum PendingTransport { + InProcess { + transport: tokio::io::DuplexStream, + }, + Stdio { + transport: Box, + }, + StreamableHttp { + transport: StreamableHttpClientTransport, + }, + StreamableHttpWithOAuth { + transport: StreamableHttpClientTransport>, + oauth_runtime: OAuthRuntime, + }, + StreamableHttpWithAccessTokenOnly { + transport: StreamableHttpClientTransport>, + }, +} + +enum ClientState { + Connecting { + transport: Option, + }, + Ready { + service: Arc>, + oauth: Option, + }, + Closed, +} + +/// Bearer authentication applied directly or by the selected HTTP transport. +#[derive(Clone)] +pub enum StreamableHttpBearerToken { + /// A token already resolved in the current process. + Resolved(String), + /// The HTTP client attaches credentials when it sends each request. + ProvidedByHttpClient, +} + +#[derive(Clone)] +enum TransportRecipe { + InProcess { + factory: Arc, + }, + Stdio { + command: StdioServerCommand, + launcher: Arc, + }, + StreamableHttp { + server_name: String, + url: String, + bearer_token: Option, + http_headers: Option>, + env_http_headers: Option>, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + pinned_credential_store: Arc>, + http_client: Arc, + auth_provider: Option, + redirect_mode: StreamableHttpRedirectMode, + oauth_refresh_mode: McpOAuthRefreshMode, + initialize_deadline: Arc>>, + }, +} + +struct InitializeDeadlineGuard { + deadline: Arc>>, +} + +impl Drop for InitializeDeadlineGuard { + fn drop(&mut self) { + *self.deadline.lock().unwrap_or_else(PoisonError::into_inner) = None; + } +} + +#[derive(Clone)] +struct InitializeContext { + timeout: Option, + client_info: InitializeRequestParams, + send_elicitation: Arc, +} + +#[derive(Clone)] +pub(crate) struct ElicitationPauseState { + active_count: Arc, + paused: watch::Sender, +} + +impl ElicitationPauseState { + pub(crate) fn new() -> Self { + let (paused, _rx) = watch::channel(false); + Self { + active_count: Arc::new(AtomicUsize::new(0)), + paused, + } + } + + pub(crate) fn enter(&self) -> ElicitationPauseGuard { + if self.active_count.fetch_add(1, Ordering::AcqRel) == 0 { + self.paused.send_replace(true); + } + ElicitationPauseGuard { + pause_state: self.clone(), + } + } + + fn subscribe(&self) -> watch::Receiver { + self.paused.subscribe() + } +} + +pub(crate) struct ElicitationPauseGuard { + pause_state: ElicitationPauseState, +} + +impl Drop for ElicitationPauseGuard { + fn drop(&mut self) { + if self.pause_state.active_count.fetch_sub(1, Ordering::AcqRel) == 1 { + self.pause_state.paused.send_replace(false); + } + } +} + +async fn active_time_timeout( + duration: Duration, + mut pause_state: watch::Receiver, + operation: Fut, +) -> std::result::Result +where + Fut: Future, +{ + let mut remaining = duration; + tokio::pin!(operation); + + loop { + if *pause_state.borrow_and_update() { + tokio::select! { + result = &mut operation => return Ok(result), + changed = pause_state.changed() => { + if changed.is_err() { + return time::timeout(remaining, operation).await.map_err(|_| ()); + } + let _paused = *pause_state.borrow_and_update(); + } + } + continue; + } + + let active_start = Instant::now(); + tokio::select! { + result = &mut operation => return Ok(result), + _ = time::sleep(remaining) => { + return Err(()); + } + changed = pause_state.changed() => { + if changed.is_err() { + return time::timeout(remaining, operation).await.map_err(|_| ()); + } + if *pause_state.borrow_and_update() { + remaining = remaining.saturating_sub(active_start.elapsed()); + if remaining.is_zero() { + return Err(()); + } + } + } + } + } +} + +#[derive(Debug, thiserror::Error)] +pub(crate) enum ClientOperationError { + #[error(transparent)] + Service(#[from] rmcp::service::ServiceError), + #[error("timed out awaiting {label} after {duration:.0?}")] + Timeout { label: String, duration: Duration }, +} + +fn remaining_operation_timeout( + label: &str, + timeout: Option, + deadline: Option, +) -> std::result::Result, ClientOperationError> { + let Some(deadline) = deadline else { + return Ok(None); + }; + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + Err(ClientOperationError::Timeout { + label: label.to_string(), + duration: timeout.unwrap_or(remaining), + }) + } else { + Ok(Some(remaining)) + } +} + +#[derive(Debug, Clone, PartialEq)] +pub enum Elicitation { + Mcp(ElicitRequestParams), + OpenAiForm { + meta: Option, + message: String, + requested_schema: serde_json::Value, + }, + OpenAiElicitationForm { + meta: Option, + message: String, + requested_schema: serde_json::Value, + }, + UserVerification { + title: String, + description: String, + challenge: String, + }, +} + +impl Elicitation { + pub fn meta(&self) -> Option<&serde_json::Map> { + match self { + Self::Mcp(request) => request.meta().map(|meta| &meta.0.0), + Self::OpenAiForm { meta, .. } | Self::OpenAiElicitationForm { meta, .. } => { + meta.as_ref().and_then(serde_json::Value::as_object) + } + Self::UserVerification { .. } => None, + } + } +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ElicitationResponse { + pub action: ElicitationAction, + pub content: Option, + #[serde(rename = "_meta")] + pub meta: Option, +} + +impl From for ElicitationResponse { + fn from(value: ElicitResult) -> Self { + Self { + action: value.action, + content: value.content, + meta: value.meta.map(|meta| Value::Object(meta.0)), + } + } +} + +impl From for ElicitResult { + fn from(value: ElicitationResponse) -> Self { + let mut result = Self::new(value.action); + result.content = value.content; + result.meta = value.meta.and_then(|meta| match meta { + Value::Object(meta) => Some(rmcp::model::MetaObject::from(meta)), + _ => None, + }); + result + } +} + +/// Interface for sending elicitation requests to the UI and awaiting a response. +pub type SendElicitation = Box< + dyn Fn(RequestId, Elicitation) -> BoxFuture<'static, Result> + Send + Sync, +>; + +pub struct ToolWithConnectorId { + pub tool: Tool, + pub connector_id: Option, + pub connector_name: Option, + pub connector_description: Option, +} + +pub struct ListToolsWithConnectorIdResult { + pub next_cursor: Option, + pub tools: Vec, +} + +/// An active Plugin Runtime event request and its request-scoped notifications. +pub struct CancellableEventStreamRequest { + pub handle: RequestHandle, + pub notifications: crate::EventNotificationReceiver, +} + +/// MCP client implemented on top of the official `rmcp` SDK. +/// https://github.com/modelcontextprotocol/rust-sdk +pub struct RmcpClient { + state: Mutex, + stdio_process: Option, + transport_recipe: TransportRecipe, + protocol_mode: McpProtocolMode, + initialize_context: Mutex>, + session_recovery_lock: Semaphore, + elicitation_pause_state: ElicitationPauseState, +} + +impl RmcpClient { + /// Returns the protocol compatibility policy captured when this client was created. + pub fn protocol_mode(&self) -> McpProtocolMode { + self.protocol_mode + } + + pub async fn new_in_process_client( + factory: Arc, + ) -> io::Result { + let transport_recipe = TransportRecipe::InProcess { factory }; + let transport = Self::create_pending_transport(&transport_recipe) + .await + .map_err(io::Error::other)?; + + Ok(Self { + state: Mutex::new(ClientState::Connecting { + transport: Some(transport), + }), + stdio_process: None, + transport_recipe, + protocol_mode: McpProtocolMode::Legacy, + initialize_context: Mutex::new(None), + session_recovery_lock: Semaphore::new(/*permits*/ 1), + elicitation_pause_state: ElicitationPauseState::new(), + }) + } + + pub async fn new_stdio_client( + program: OsString, + args: Vec, + env: Option>, + env_vars: &[McpServerEnvVar], + cwd: Option, + launcher: Arc, + ) -> io::Result { + Self::new_stdio_client_with_protocol_mode( + program, + args, + env, + env_vars, + cwd, + launcher, + McpProtocolMode::Legacy, + ) + .await + } + + /// Constructs a stdio client with an explicitly selected compatibility policy. + #[allow(clippy::too_many_arguments)] + pub async fn new_stdio_client_with_protocol_mode( + program: OsString, + args: Vec, + mut env: Option>, + env_vars: &[McpServerEnvVar], + cwd: Option, + launcher: Arc, + protocol_mode: McpProtocolMode, + ) -> io::Result { + let requested_stdio_version = match protocol_mode { + McpProtocolMode::Legacy => None, + McpProtocolMode::V20260728 => env + .as_mut() + .and_then(|env| env.remove(OsStr::new("CODEX_MCP_PROTOCOL_VERSION"))), + }; + let protocol_mode = protocol_mode.stdio_mode(requested_stdio_version.as_deref())?; + let transport_recipe = TransportRecipe::Stdio { + command: StdioServerCommand::new( + program, + args, + env, + env_vars.to_vec(), + cwd, + protocol_mode, + ), + launcher, + }; + let transport = Self::create_pending_transport(&transport_recipe) + .await + .map_err(io::Error::other)?; + let stdio_process = match &transport { + PendingTransport::Stdio { transport } => Some(transport.process_handle()), + PendingTransport::InProcess { .. } + | PendingTransport::StreamableHttp { .. } + | PendingTransport::StreamableHttpWithOAuth { .. } + | PendingTransport::StreamableHttpWithAccessTokenOnly { .. } => None, + }; + + Ok(Self { + state: Mutex::new(ClientState::Connecting { + transport: Some(transport), + }), + stdio_process, + transport_recipe, + protocol_mode, + initialize_context: Mutex::new(None), + session_recovery_lock: Semaphore::new(/*permits*/ 1), + elicitation_pause_state: ElicitationPauseState::new(), + }) + } + + #[allow(clippy::too_many_arguments)] + pub async fn new_streamable_http_client( + server_name: &str, + url: &str, + bearer_token: Option, + http_headers: Option>, + env_http_headers: Option>, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_client: Arc, + auth_provider: Option, + ) -> Result { + Self::new_streamable_http_client_with_protocol_mode( + server_name, + url, + bearer_token, + http_headers, + env_http_headers, + store_mode, + keyring_backend_kind, + http_client, + auth_provider, + McpProtocolMode::Legacy, + ) + .await + } + + /// Constructs a streamable HTTP client with an explicitly selected compatibility policy. + #[allow(clippy::too_many_arguments)] + pub async fn new_streamable_http_client_with_protocol_mode( + server_name: &str, + url: &str, + bearer_token: Option, + http_headers: Option>, + env_http_headers: Option>, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_client: Arc, + auth_provider: Option, + protocol_mode: McpProtocolMode, + ) -> Result { + Self::new_streamable_http_client_with_protocol_mode_and_redirect_mode( + server_name, + url, + bearer_token.map(StreamableHttpBearerToken::Resolved), + http_headers, + env_http_headers, + store_mode, + keyring_backend_kind, + http_client, + auth_provider, + protocol_mode, + StreamableHttpRedirectMode::Legacy, + McpOAuthRefreshMode::Legacy, + ) + .await + } + + #[allow(clippy::too_many_arguments)] + pub async fn new_streamable_http_client_with_protocol_mode_and_redirect_mode( + server_name: &str, + url: &str, + bearer_token: Option, + http_headers: Option>, + env_http_headers: Option>, + store_mode: OAuthCredentialsStoreMode, + keyring_backend_kind: AuthKeyringBackendKind, + http_client: Arc, + auth_provider: Option, + protocol_mode: McpProtocolMode, + redirect_mode: StreamableHttpRedirectMode, + oauth_refresh_mode: McpOAuthRefreshMode, + ) -> Result { + let transport_recipe = TransportRecipe::StreamableHttp { + server_name: server_name.to_string(), + url: url.to_string(), + bearer_token, + http_headers, + env_http_headers, + store_mode, + keyring_backend_kind, + pinned_credential_store: Arc::new(OnceLock::new()), + http_client, + auth_provider, + redirect_mode, + oauth_refresh_mode, + initialize_deadline: Arc::new(StdMutex::new(None)), + }; + let transport = Self::create_pending_transport(&transport_recipe).await?; + Ok(Self { + state: Mutex::new(ClientState::Connecting { + transport: Some(transport), + }), + stdio_process: None, + transport_recipe, + protocol_mode, + initialize_context: Mutex::new(None), + session_recovery_lock: Semaphore::new(/*permits*/ 1), + elicitation_pause_state: ElicitationPauseState::new(), + }) + } + + /// Perform the initialization handshake with the MCP server. + /// https://modelcontextprotocol.io/specification/2025-06-18/basic/lifecycle#initialization + #[instrument(level = "trace", skip_all)] + pub async fn initialize( + &self, + params: InitializeRequestParams, + timeout: Option, + send_elicitation: SendElicitation, + ) -> Result { + let context = InitializeContext { + timeout, + client_info: params, + send_elicitation: Arc::new(send_elicitation), + }; + let pending_transport = { + let mut guard = self.state.lock().await; + match &mut *guard { + ClientState::Connecting { transport } => match transport.take() { + Some(transport) => transport, + None => return Err(anyhow!("client already initializing")), + }, + ClientState::Ready { .. } => return Err(anyhow!("client already initialized")), + ClientState::Closed => return Err(anyhow!("MCP client is shut down")), + } + }; + + let (service, oauth_runtime) = self + .connect_pending_transport_with_initialize_retries(pending_transport, &context) + .await?; + + let initialize_result_rmcp = service + .peer() + .peer_info() + .ok_or_else(|| anyhow!("handshake succeeded but server info was missing"))?; + let initialize_result = initialize_result_rmcp.as_ref().clone(); + + { + let mut initialize_context = self.initialize_context.lock().await; + *initialize_context = Some(context); + } + + { + let mut guard = self.state.lock().await; + if matches!(*guard, ClientState::Closed) { + return Err(anyhow!("MCP client is shut down")); + } + *guard = ClientState::Ready { + service, + oauth: oauth_runtime.clone(), + }; + } + + if let Some(OAuthRuntime::Legacy(runtime)) = oauth_runtime + && let Err(error) = runtime.persist_if_needed().await + { + warn!("failed to persist OAuth tokens after initialize: {error}"); + } + + Ok(initialize_result) + } + + pub async fn list_tools( + &self, + params: Option, + timeout: Option, + ) -> Result { + self.refresh_oauth_if_needed().await?; + let result = self + .run_service_operation("tools/list", timeout, move |service| { + let params = params.clone(); + async move { service.list_tools(params).await }.boxed() + }) + .await?; + self.persist_oauth_tokens().await; + Ok(result) + } + + #[instrument(level = "trace", skip_all)] + pub async fn list_tools_with_connector_ids( + &self, + params: Option, + timeout: Option, + ) -> Result { + self.refresh_oauth_if_needed().await?; + let result = self + .run_service_operation("tools/list", timeout, move |service| { + let params = params.clone(); + async move { service.list_tools(params).await }.boxed() + }) + .await?; + let tools = result + .tools + .into_iter() + .map(|tool| { + let meta = tool.meta.as_ref(); + let connector_id = Self::meta_string(meta, "connector_id"); + let connector_name = Self::meta_string(meta, "connector_name") + .or_else(|| Self::meta_string(meta, "connector_display_name")); + let connector_description = Self::meta_string(meta, "connector_description") + .or_else(|| Self::meta_string(meta, "connectorDescription")); + Ok(ToolWithConnectorId { + tool, + connector_id, + connector_name, + connector_description, + }) + }) + .collect::>>()?; + self.persist_oauth_tokens().await; + Ok(ListToolsWithConnectorIdResult { + next_cursor: result.next_cursor, + tools, + }) + } + + fn meta_string(meta: Option<&rmcp::model::MetaObject>, key: &str) -> Option { + meta.and_then(|meta| meta.get(key)) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + } + + pub async fn list_resources( + &self, + params: Option, + timeout: Option, + ) -> Result { + self.refresh_oauth_if_needed().await?; + let result = self + .run_service_operation("resources/list", timeout, move |service| { + let params = params.clone(); + async move { service.list_resources(params).await }.boxed() + }) + .await?; + self.persist_oauth_tokens().await; + Ok(result) + } + + pub async fn list_resource_templates( + &self, + params: Option, + timeout: Option, + ) -> Result { + self.refresh_oauth_if_needed().await?; + let result = self + .run_service_operation("resources/templates/list", timeout, move |service| { + let params = params.clone(); + async move { service.list_resource_templates(params).await }.boxed() + }) + .await?; + self.persist_oauth_tokens().await; + Ok(result) + } + + pub async fn read_resource( + &self, + params: ReadResourceRequestParams, + timeout: Option, + ) -> Result { + self.refresh_oauth_if_needed().await?; + let requested_modern = self.protocol_mode == McpProtocolMode::V20260728; + let result = self + .run_service_operation("resources/read", timeout, move |service| { + let params = params.clone(); + async move { + let modern_session = requested_modern + && service.peer().peer_info().is_some_and(|info| { + info.protocol_version == ProtocolVersion::V_2026_07_28 + }); + if modern_session { + service.read_resource(params).await + } else { + service.peer().read_resource(params).await + } + } + .boxed() + }) + .await?; + self.persist_oauth_tokens().await; + Ok(result) + } + + pub async fn call_tool( + &self, + name: String, + arguments: Option, + meta: Option, + timeout: Option, + ) -> Result { + let authentication_required_result = |error| { + if !is_authentication_required_error(&error) { + return Err(error); + } + // Local expiry and server rejection use the same reconnect signal without + // exposing token-endpoint or transport details in the tool result. + let mut result = CallToolResult::error(vec![ContentBlock::text( + "MCP authentication required. Reconnect to continue using this server.", + )]); + result.meta = Some( + serde_json::Map::from_iter([( + "mcp/www_authenticate".to_string(), + Value::String("Bearer error=\"invalid_token\"".to_string()), + )]) + .into(), + ); + Ok(result) + }; + if let Err(error) = self.refresh_oauth_if_needed().await { + return authentication_required_result(error); + } + let arguments = match arguments { + Some(Value::Object(map)) => Some(map), + Some(other) => { + return Err(anyhow!( + "MCP tool arguments must be a JSON object, got {other}" + )); + } + None => None, + }; + let meta = match meta { + Some(Value::Object(map)) => Some(RequestMetaObject::from(map)), + Some(other) => { + return Err(anyhow!( + "MCP tool request _meta must be a JSON object, got {other}" + )); + } + None => None, + }; + let mut rmcp_params = CallToolRequestParams::new(name); + rmcp_params.arguments = arguments; + let requested_modern = self.protocol_mode == McpProtocolMode::V20260728; + match self + .run_service_operation("tools/call", timeout, move |service| { + let mut rmcp_params = rmcp_params.clone(); + let meta = meta.clone(); + async move { + let modern_session = requested_modern + && service.peer().peer_info().is_some_and(|info| { + info.protocol_version == ProtocolVersion::V_2026_07_28 + }); + if modern_session { + rmcp_params.meta = meta; + return crate::tool_input::call_tool(&service, rmcp_params).await; + } + let mut options = rmcp::service::PeerRequestOptions::no_options(); + options.meta = meta; + let result = service + .peer() + .send_request_with_option( + ClientRequest::CallToolRequest(rmcp::model::CallToolRequest::new( + rmcp_params, + )), + options, + ) + .await? + .await_response() + .await?; + match result { + ServerResult::CallToolResult(result) => Ok(result), + _ => Err(rmcp::service::ServiceError::UnexpectedResponse), + } + } + .boxed() + }) + .await + { + Ok(result) => { + self.persist_oauth_tokens().await; + Ok(result) + } + Err(error) => { + let Some(ClientOperationError::Service(ServiceError::TransportSend(transport))) = + error.downcast_ref() + else { + return authentication_required_result(error); + }; + let Some(StreamableHttpError::AuthRequired(challenge)) = + transport + .error + .downcast_ref::>() + else { + return authentication_required_result(error); + }; + // The transport has already handled automatic refresh. Preserve the challenge + // for interactive login without replaying the rejected tool call. + let mut result = + CallToolResult::error(vec![ContentBlock::text("Authentication required")]); + result.meta = Some(MetaObject::from(serde_json::Map::from_iter([( + "mcp/www_authenticate".to_string(), + serde_json::json!([challenge.www_authenticate_header]), + )]))); + Ok(result) + } + } + } + + pub async fn send_custom_notification( + &self, + method: &str, + params: Option, + ) -> Result<()> { + self.refresh_oauth_if_needed().await?; + self.run_service_operation( + "notifications/custom", + /*timeout*/ None, + move |service| { + let params = params.clone(); + async move { + service + .send_notification(ClientNotification::CustomNotification( + CustomNotification { + method: method.to_string(), + params, + extensions: Extensions::new(), + }, + )) + .await + } + .boxed() + }, + ) + .await?; + self.persist_oauth_tokens().await; + Ok(()) + } + + pub async fn send_custom_request( + &self, + method: &str, + params: Option, + ) -> Result { + self.send_custom_request_with_timeout(method, params, /*timeout*/ None) + .await + } + + pub async fn send_custom_request_with_timeout( + &self, + method: &str, + params: Option, + timeout: Option, + ) -> Result { + self.refresh_oauth_if_needed().await?; + let response = self + .run_service_operation("requests/custom", timeout, move |service| { + let params = params.clone(); + async move { + service + .send_request(ClientRequest::CustomRequest(CustomRequest::new( + method, params, + ))) + .await + } + .boxed() + }) + .await?; + self.persist_oauth_tokens().await; + Ok(response) + } + + /// Starts a Plugin Runtime event stream without waiting for its final response. + pub async fn send_event_stream_request( + &self, + params: Option, + ) -> Result { + let service = self.service().await?; + let (sender, notifications) = event_notification_channel(); + let mut request = CustomRequest::new("events/stream", params); + request.extensions.insert(sender); + let handle = service + .peer() + .send_cancellable_request( + ClientRequest::CustomRequest(request), + rmcp::service::PeerRequestOptions::no_options(), + ) + .await?; + + Ok(CancellableEventStreamRequest { + handle, + notifications, + }) + } + + async fn service(&self) -> Result>> { + let guard = self.state.lock().await; + match &*guard { + ClientState::Ready { service, .. } => Ok(Arc::clone(service)), + ClientState::Connecting { .. } => Err(anyhow!("MCP client not initialized")), + ClientState::Closed => Err(anyhow!("MCP client is shut down")), + } + } + + async fn oauth_runtime(&self) -> Option { + let guard = self.state.lock().await; + match &*guard { + ClientState::Ready { + oauth: Some(runtime), + .. + } => Some(runtime.clone()), + _ => None, + } + } + + /// Returns `None` when this client does not manage stored OAuth credentials. + pub async fn managed_oauth_credentials(&self) -> Option> { + let runtime = self.oauth_runtime().await?; + Some(match runtime { + OAuthRuntime::Legacy(persistor) => persistor.stored_credentials().await, + OAuthRuntime::Coordinated { store, .. } => store.stored_credentials().await, + }) + } + + /// Returns whether an initialized transport or its underlying service has stopped. + pub async fn is_closed(&self) -> bool { + let state = self.state.lock().await; + match &*state { + ClientState::Ready { service, .. } => { + service.is_closed() || service.peer().is_transport_closed() + } + ClientState::Connecting { .. } => false, + ClientState::Closed => true, + } + } + + /// Stop the MCP transport and any stdio server process owned by this client. + pub async fn shutdown(&self) { + let previous_state = { + let mut guard = self.state.lock().await; + std::mem::replace(&mut *guard, ClientState::Closed) + }; + + if let Some(process) = &self.stdio_process + && let Err(error) = process.terminate().await + { + warn!("failed to terminate MCP stdio server process: {error}"); + } + + drop(previous_state); + } + + /// This should be called after every tool call so that if a given tool call triggered + /// a refresh of the OAuth tokens, they are persisted. + async fn persist_oauth_tokens(&self) { + if let Some(OAuthRuntime::Legacy(runtime)) = self.oauth_runtime().await + && let Err(error) = runtime.persist_if_needed().await + { + warn!("failed to persist OAuth tokens: {error}"); + } + } + + /// Prepares managed OAuth before the MCP operation timeout starts. + async fn refresh_oauth_if_needed(&self) -> Result<()> { + if let Some(runtime) = self.oauth_runtime().await { + runtime.refresh_if_needed().await?; + } + Ok(()) + } + + async fn create_pending_transport( + transport_recipe: &TransportRecipe, + ) -> Result { + match transport_recipe { + TransportRecipe::InProcess { factory } => { + let transport = factory.open().await?; + Ok(PendingTransport::InProcess { transport }) + } + TransportRecipe::Stdio { command, launcher } => { + let transport = launcher.launch(command.clone()).await?; + Ok(PendingTransport::Stdio { + transport: Box::new(transport), + }) + } + TransportRecipe::StreamableHttp { + server_name, + url, + bearer_token, + http_headers, + env_http_headers, + store_mode, + keyring_backend_kind, + pinned_credential_store, + http_client, + auth_provider, + redirect_mode, + oauth_refresh_mode, + initialize_deadline, + } => { + let has_configured_headers = matches!( + bearer_token, + Some(StreamableHttpBearerToken::ProvidedByHttpClient) + ) || http_headers + .as_ref() + .is_some_and(|headers| !headers.is_empty()) + || env_http_headers + .as_ref() + .is_some_and(|headers| !headers.is_empty()); + let default_headers = + build_default_headers(http_headers.clone(), env_http_headers.clone())?; + let auth_provider = + if bearer_token.is_some() || default_headers.contains_key(AUTHORIZATION) { + None + } else { + auth_provider.clone() + }; + + let resolved_oauth_tokens = if bearer_token.is_none() + && auth_provider.is_none() + && !default_headers.contains_key(AUTHORIZATION) + { + let oauth_server_name = server_name.clone(); + let oauth_url = url.clone(); + let oauth_store_mode = *store_mode; + let oauth_keyring_backend_kind = *keyring_backend_kind; + let pinned_credential_store = Arc::clone(pinned_credential_store); + + tokio::task::spawn_blocking(move || -> Result> { + if let Some(store) = pinned_credential_store.get().copied() { + // Rebuilds reread the source selected during first construction. Only + // initial construction below evaluates configured store policy. + return store + .load(&DefaultKeyringStore, &oauth_server_name, &oauth_url) + .map(|tokens| tokens.map(|tokens| ResolvedOAuthTokens { tokens, store })); + } + + match resolve_oauth_tokens_from_store_policy( + &DefaultKeyringStore, + &oauth_server_name, + &oauth_url, + oauth_store_mode, + oauth_keyring_backend_kind, + ) { + Ok(tokens) => { + if let Some(resolved) = tokens.as_ref() { + // Retries and session recovery rebuild this transport. Pin the + // first concrete source so Auto is not reevaluated mid-client. + pinned_credential_store.set(resolved.store).map_err(|_| { + anyhow!( + "OAuth credential store pinned concurrently for MCP server `{oauth_server_name}`" + ) + })?; + } + Ok(tokens) + } + Err(err) => { + warn!( + "failed to read tokens for server `{oauth_server_name}`: {err}" + ); + Ok(None) + } + } + }) + .await + .map_err(|error| anyhow!("OAuth credential loading task failed: {error}"))?? + } else { + None + }; + + if let Some(ResolvedOAuthTokens { + tokens: initial_tokens, + store: credential_store, + }) = resolved_oauth_tokens + { + match create_oauth_transport_and_runtime( + server_name, + url, + initial_tokens.clone(), + credential_store, + default_headers.clone(), + Arc::clone(http_client), + has_configured_headers, + *redirect_mode, + *oauth_refresh_mode, + Arc::clone(initialize_deadline), + ) + .await + { + Ok(pending_transport) => Ok(pending_transport), + Err(err) + if err.downcast_ref::().is_some_and(|auth_err| { + matches!(auth_err, AuthError::NoAuthorizationSupport) + }) => + { + let access_token = initial_tokens + .token_response + .0 + .access_token() + .secret() + .to_string(); + warn!( + "OAuth metadata discovery is unavailable for MCP server `{server_name}`; falling back to stored bearer token authentication" + ); + let http_config = + StreamableHttpClientTransportConfig::with_uri(url.clone()) + .auth_header(access_token); + let transport = StreamableHttpClientTransport::with_client( + StreamableHttpClientAdapter::new( + Arc::clone(http_client), + default_headers, + /*auth_provider*/ None, + has_configured_headers, + *redirect_mode, + Arc::clone(initialize_deadline), + ), + http_config, + ); + Ok(PendingTransport::StreamableHttp { transport }) + } + Err(err) => Err(err), + } + } else { + let mut http_config = + StreamableHttpClientTransportConfig::with_uri(url.clone()); + if let Some(StreamableHttpBearerToken::Resolved(bearer_token)) = bearer_token { + http_config = http_config.auth_header(bearer_token.clone()); + } + + let transport = StreamableHttpClientTransport::with_client( + StreamableHttpClientAdapter::new( + Arc::clone(http_client), + default_headers, + auth_provider, + has_configured_headers, + *redirect_mode, + Arc::clone(initialize_deadline), + ), + http_config, + ); + Ok(PendingTransport::StreamableHttp { transport }) + } + } + } + } + + async fn connect_pending_transport( + &self, + pending_transport: PendingTransport, + initialize_context: &InitializeContext, + timeout: Option, + ) -> Result<( + Arc>, + Option, + )> { + // Request IDs and remembered cancellations belong to this connection, including + // when a failed initialization or expired HTTP session creates a new transport. + let send_elicitation = Arc::clone(&initialize_context.send_elicitation); + let client_service = ElicitationClientService::new( + initialize_context.client_info.clone(), + Box::new(move |id, request| send_elicitation(id, request)), + self.elicitation_pause_state.clone(), + ); + let _initialize_deadline = match &self.transport_recipe { + TransportRecipe::StreamableHttp { + initialize_deadline, + .. + } => { + *initialize_deadline + .lock() + .unwrap_or_else(PoisonError::into_inner) = + timeout.and_then(|duration| Instant::now().checked_add(duration)); + Some(InitializeDeadlineGuard { + deadline: Arc::clone(initialize_deadline), + }) + } + TransportRecipe::InProcess { .. } | TransportRecipe::Stdio { .. } => None, + }; + let lifecycle = self.protocol_mode.client_lifecycle(); + let (transport, oauth_runtime) = match pending_transport { + PendingTransport::InProcess { transport } => ( + client_service + .serve_with_lifecycle(transport, lifecycle) + .boxed(), + None, + ), + PendingTransport::Stdio { transport } => ( + client_service + .serve_with_lifecycle(*transport, lifecycle) + .boxed(), + None, + ), + PendingTransport::StreamableHttp { transport } => ( + client_service + .serve_with_lifecycle(capture_event_notifications(transport), lifecycle) + .boxed(), + None, + ), + PendingTransport::StreamableHttpWithOAuth { + transport, + oauth_runtime, + } => ( + client_service + .serve_with_lifecycle(transport, lifecycle) + .boxed(), + Some(oauth_runtime), + ), + PendingTransport::StreamableHttpWithAccessTokenOnly { transport } => ( + client_service + .serve_with_lifecycle(transport, lifecycle) + .boxed(), + None, + ), + }; + + let service_result = match timeout { + Some(duration) => match time::timeout(duration, transport).await { + Ok(result) => { + result.map_err(|source| anyhow::Error::from(HandshakeError { source })) + } + Err(_elapsed) => Err(anyhow!( + "timed out handshaking with MCP server after {duration:?}" + )), + }, + None => transport + .await + .map_err(|source| anyhow::Error::from(HandshakeError { source })), + }; + let service = match service_result { + Ok(service) => service, + Err(error) => { + if let Some(OAuthRuntime::Legacy(runtime)) = oauth_runtime.as_ref() + && let Err(persist_error) = runtime.persist_if_needed().await + { + warn!( + "failed to persist OAuth tokens after failed initialize: {persist_error}" + ); + } + return Err(error); + } + }; + + // Preserve Codex's existing snapshot and request-freshness behavior. rmcp 3 + // enables response caching and stale-on-error fallback by default. + service + .peer() + .set_response_cache_config(ClientCacheConfig::disabled()) + .await; + + Ok((Arc::new(service), oauth_runtime)) + } + + async fn run_service_operation( + &self, + label: &str, + timeout: Option, + operation: F, + ) -> Result + where + F: Fn(Arc>) -> Fut, + Fut: std::future::Future>, + { + let service = self.service().await?; + match Self::run_service_operation_with_transient_retries( + Arc::clone(&service), + label, + timeout, + self.elicitation_pause_state.clone(), + &operation, + ) + .await + { + Ok(result) => Ok(result), + Err(error) if Self::is_session_expired_404(&error) => { + self.reinitialize_after_session_expiry(&service).await?; + let recovered_service = self.service().await?; + Self::run_service_operation_with_transient_retries( + recovered_service, + label, + timeout, + self.elicitation_pause_state.clone(), + &operation, + ) + .await + .map_err(Into::into) + } + Err(error) => Err(error.into()), + } + } + + async fn run_service_operation_with_transient_retries( + service: Arc>, + label: &str, + timeout: Option, + pause_state: ElicitationPauseState, + operation: &F, + ) -> std::result::Result + where + F: Fn(Arc>) -> Fut, + Fut: std::future::Future>, + { + let retry_deadline = timeout.map(|duration| Instant::now() + duration); + for (attempt, retry_delay_ms) in STREAMABLE_HTTP_RETRY_DELAYS_MS + .iter() + .copied() + .map(Some) + .chain(std::iter::once(None)) + .enumerate() + { + let attempt_timeout = remaining_operation_timeout(label, timeout, retry_deadline)?; + match Self::run_service_operation_once( + Arc::clone(&service), + label, + attempt_timeout, + pause_state.clone(), + operation, + ) + .await + { + Ok(result) => return Ok(result), + Err(error) if Self::is_retryable_tools_list_error(label, &error) => { + let Some(retry_delay_ms) = retry_delay_ms else { + return Err(error); + }; + let delay = Duration::from_millis(retry_delay_ms); + warn!( + attempt = attempt + 1, + max_attempts = STREAMABLE_HTTP_RETRY_DELAYS_MS.len() + 1, + delay_ms = delay.as_millis(), + error = %error, + "streamable HTTP MCP tools/list failed with a retryable error; retrying" + ); + if !sleep_with_retry_deadline(delay, retry_deadline).await { + return Err(ClientOperationError::Timeout { + label: label.to_string(), + duration: timeout.unwrap_or(delay), + }); + } + } + Err(error) => return Err(error), + } + } + + unreachable!("service operation retry loop should return on success or final error") + } + + async fn run_service_operation_once( + service: Arc>, + label: &str, + timeout: Option, + pause_state: ElicitationPauseState, + operation: &F, + ) -> std::result::Result + where + F: Fn(Arc>) -> Fut, + Fut: std::future::Future>, + { + match timeout { + Some(duration) => { + active_time_timeout(duration, pause_state.subscribe(), operation(service)) + .await + .map_err(|_| ClientOperationError::Timeout { + label: label.to_string(), + duration, + })? + .map_err(ClientOperationError::from) + } + None => operation(service).await.map_err(ClientOperationError::from), + } + } + + fn is_retryable_tools_list_error(label: &str, error: &ClientOperationError) -> bool { + if label != "tools/list" { + return false; + } + let ClientOperationError::Service(rmcp::service::ServiceError::TransportSend(error)) = + error + else { + return false; + }; + + error + .error + .downcast_ref::>() + .is_some_and(Self::is_retryable_streamable_http_error) + } + + fn is_session_expired_404(error: &ClientOperationError) -> bool { + let ClientOperationError::Service(rmcp::service::ServiceError::TransportSend(error)) = + error + else { + return false; + }; + + error + .error + .downcast_ref::>() + .is_some_and(|error| { + matches!( + error, + StreamableHttpError::Client( + StreamableHttpClientAdapterError::SessionExpired404 + ) + ) + }) + } + + async fn reinitialize_after_session_expiry( + &self, + failed_service: &Arc>, + ) -> Result<()> { + let _recovery_guard = self + .session_recovery_lock + .acquire() + .await + .map_err(|_| anyhow!("MCP client recovery semaphore closed"))?; + + { + let guard = self.state.lock().await; + match &*guard { + ClientState::Ready { service, .. } if !Arc::ptr_eq(service, failed_service) => { + return Ok(()); + } + ClientState::Ready { .. } => {} + ClientState::Connecting { .. } => { + return Err(anyhow!("MCP client not initialized")); + } + ClientState::Closed => { + return Err(anyhow!("MCP client is shut down")); + } + } + } + + let initialize_context = self + .initialize_context + .lock() + .await + .clone() + .ok_or_else(|| anyhow!("MCP client cannot recover before initialize succeeds"))?; + let pending_transport = Self::create_pending_transport(&self.transport_recipe).await?; + let (service, oauth_runtime) = self + .connect_pending_transport_with_initialize_retries( + pending_transport, + &initialize_context, + ) + .await?; + service + .peer() + .peer_info() + .ok_or_else(|| anyhow!("recovered handshake succeeded but server info was missing"))?; + + { + let mut guard = self.state.lock().await; + if matches!(*guard, ClientState::Closed) { + return Err(anyhow!("MCP client is shut down")); + } + *guard = ClientState::Ready { + service, + oauth: oauth_runtime.clone(), + }; + } + + if let Some(OAuthRuntime::Legacy(runtime)) = oauth_runtime + && let Err(error) = runtime.persist_if_needed().await + { + warn!("failed to persist OAuth tokens after session recovery: {error}"); + } + + Ok(()) + } +} + +#[allow(clippy::too_many_arguments)] +async fn create_oauth_transport_and_runtime( + server_name: &str, + url: &str, + initial_tokens: StoredOAuthTokens, + credential_store: ResolvedOAuthCredentialStore, + default_headers: HeaderMap, + http_client: Arc, + has_configured_headers: bool, + redirect_mode: StreamableHttpRedirectMode, + oauth_refresh_mode: McpOAuthRefreshMode, + initialize_deadline: Arc>>, +) -> Result { + let oauth_http_client = Arc::new(OAuthHttpClientAdapter::new_with_redirect_mode( + http_client.clone(), + default_headers.clone(), + url, + has_configured_headers, + redirect_mode, + )?); + let mut manager = + AuthorizationManager::new_with_oauth_http_client(url.to_string(), oauth_http_client) + .await?; + manager.set_allow_missing_issuer(true); + let metadata = manager + .resolve_metadata() + .await + .context("failed to resolve OAuth metadata before using stored credentials")? + .metadata; + let use_stored_access_token_only = + match validate_refresh_token_issuer(&metadata, &initial_tokens) { + Ok(()) => false, + Err(_error) if initial_tokens.access_token_is_usable_without_refresh() => true, + Err(error) => return Err(error), + }; + manager.set_metadata(metadata); + let mut runtime_tokens = initial_tokens.clone(); + if use_stored_access_token_only { + runtime_tokens.token_response.0.set_refresh_token(None); + runtime_tokens.issuer = None; + } + install_tokens_in_manager(&mut manager, &runtime_tokens).await?; + let coordinated_store = match oauth_refresh_mode { + McpOAuthRefreshMode::Coordinated if !use_stored_access_token_only => { + let store = OAuthCredentialStore::new( + initial_tokens.clone(), + credential_store, + DefaultKeyringStore, + ); + manager.set_credential_store(store.clone()); + Some(store) + } + McpOAuthRefreshMode::Legacy | McpOAuthRefreshMode::Coordinated => None, + }; + + let auth_client = AuthClient::new( + StreamableHttpClientAdapter::new( + http_client, + default_headers, + /*auth_provider*/ None, + has_configured_headers, + redirect_mode, + initialize_deadline, + ), + manager, + ); + let auth_manager = auth_client.auth_manager.clone(); + + let transport = StreamableHttpClientTransport::with_client( + auth_client, + StreamableHttpClientTransportConfig::with_uri(url.to_string()), + ); + + if use_stored_access_token_only { + warn!( + "stored OAuth refresh credentials could not be bound to their issuer for MCP server `{server_name}`; using the stored access token without refresh" + ); + return Ok(PendingTransport::StreamableHttpWithAccessTokenOnly { transport }); + } + + let runtime = match coordinated_store { + Some(store) => OAuthRuntime::Coordinated { + auth_manager, + store, + }, + None => OAuthRuntime::Legacy(OAuthPersistor::new( + server_name.to_string(), + url.to_string(), + auth_manager, + credential_store, + Some(initial_tokens), + )), + }; + + Ok(PendingTransport::StreamableHttpWithOAuth { + transport, + oauth_runtime: runtime, + }) +} + +#[cfg(test)] +#[path = "tool_input_tests.rs"] +mod tool_input_tests; + +#[cfg(test)] +#[path = "user_verification_cancellation_tests.rs"] +mod user_verification_cancellation_tests; + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use pretty_assertions::assert_eq; + use tokio::time; + + use super::*; + + #[test] + fn client_operation_timeout_rounds_duration() { + let error = ClientOperationError::Timeout { + label: "tools/list".to_string(), + duration: Duration::from_nanos(29_999_999_875), + }; + + assert_eq!(error.to_string(), "timed out awaiting tools/list after 30s"); + } + + #[tokio::test] + async fn active_time_timeout_pauses_while_elicitation_is_pending() { + let pause_state = ElicitationPauseState::new(); + let pause = pause_state.enter(); + tokio::spawn(async move { + time::sleep(Duration::from_millis(75)).await; + drop(pause); + }); + + let result = + active_time_timeout(Duration::from_millis(50), pause_state.subscribe(), async { + time::sleep(Duration::from_millis(90)).await; + "done" + }) + .await; + + assert_eq!(Ok("done"), result); + } +} diff --git a/codex-rs/rmcp-client/src/service_error.rs b/codex-rs/rmcp-client/src/service_error.rs new file mode 100644 index 0000000000000000000000000000000000000000..65da799f1478cb5157af1a935fa5791e080b8649 --- /dev/null +++ b/codex-rs/rmcp-client/src/service_error.rs @@ -0,0 +1,15 @@ +//! Access structured MCP errors without dropping their original protocol data. + +use anyhow::Error; +use rmcp::ErrorData; +use rmcp::service::ServiceError; + +use crate::rmcp_client::ClientOperationError; + +/// Returns the server's structured protocol error, including its original data. +pub fn mcp_error(error: &Error) -> Option<&ErrorData> { + let ClientOperationError::Service(ServiceError::McpError(error)) = error.downcast_ref()? else { + return None; + }; + Some(error) +} diff --git a/codex-rs/rmcp-client/src/startup_error.rs b/codex-rs/rmcp-client/src/startup_error.rs new file mode 100644 index 0000000000000000000000000000000000000000..744e361e9b7dead5cc6c6aeac6878f6eb3aaf1b1 --- /dev/null +++ b/codex-rs/rmcp-client/src/startup_error.rs @@ -0,0 +1,61 @@ +use anyhow::Error; +use rmcp::service::ClientInitializeError; +use rmcp::service::ServiceError; +use rmcp::transport::DynamicTransportError; +use rmcp::transport::auth::AuthError; +use rmcp::transport::streamable_http_client::StreamableHttpError; + +use crate::http_client_adapter::StreamableHttpClientAdapterError; +use crate::rmcp_client::ClientOperationError; + +/// Returns whether an RMCP client error indicates that authentication is required. +/// +/// This does not distinguish first-time login from reauthentication. +/// Streamable HTTP initialization errors are stored inside RMCP's dynamic +/// transport error, which is not part of the standard error source chain. +pub fn is_authentication_required_error(error: &Error) -> bool { + error.chain().any(|source| { + source + .downcast_ref::() + .is_some_and(auth_error_requires_authentication) + || source + .downcast_ref::() + .is_some_and(|mut error| { + while let ClientInitializeError::LegacyFallbackFailed { fallback, .. } = error { + error = fallback; + } + matches!( + error, + ClientInitializeError::TransportError { error, .. } + if transport_error_requires_authentication(error) + ) + }) + || source + .downcast_ref::() + .is_some_and(|error| { + matches!( + error, + ClientOperationError::Service(ServiceError::TransportSend(error)) + if transport_error_requires_authentication(error) + ) + }) + }) +} + +fn transport_error_requires_authentication(error: &DynamicTransportError) -> bool { + error + .error + .downcast_ref::>() + .is_some_and(|error| match error { + StreamableHttpError::AuthRequired(_) => true, + StreamableHttpError::Auth(auth_error) => auth_error_requires_authentication(auth_error), + _ => false, + }) +} + +fn auth_error_requires_authentication(error: &AuthError) -> bool { + matches!( + error, + AuthError::AuthorizationRequired | AuthError::TokenExpired + ) +} diff --git a/codex-rs/rmcp-client/src/stdio_server_launcher.rs b/codex-rs/rmcp-client/src/stdio_server_launcher.rs new file mode 100644 index 0000000000000000000000000000000000000000..84cc008e0f6070305fd1f37e51e99c932063f382 --- /dev/null +++ b/codex-rs/rmcp-client/src/stdio_server_launcher.rs @@ -0,0 +1,782 @@ +//! Launch MCP stdio servers and return the transport rmcp should use. +//! +//! This module owns the "where does the server process run?" decision: +//! +//! - [`LocalStdioServerLauncher`] starts the configured command as a child of +//! the orchestrator process. +//! - [`ExecutorStdioServerLauncher`] starts the configured command through the +//! executor process API. +//! +//! Both paths return [`StdioServerTransport`], so `RmcpClient` can hand the +//! resulting byte stream to rmcp without knowing where the process lives. The +//! executor-specific byte adaptation lives in `executor_process_transport`. + +use std::collections::HashMap; +use std::ffi::OsString; +use std::future::Future; +use std::io; +#[cfg(windows)] +use std::os::windows::io::OwnedHandle; +use std::path::Path; +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +#[cfg(unix)] +use std::thread::sleep; +#[cfg(unix)] +use std::thread::spawn; +use std::time::Duration; + +use anyhow::Result; +use anyhow::anyhow; +use codex_config::types::McpServerEnvVar; +use codex_exec_server::ExecBackend; +use codex_exec_server::ExecEnvPolicy; +use codex_exec_server::ExecParams; +use codex_exec_server::ExecProcess; +use codex_protocol::config_types::ShellEnvironmentPolicyInherit; +use codex_utils_path_uri::LegacyAppPathString; +use codex_utils_path_uri::PathUri; +#[cfg(all(unix, not(target_os = "macos")))] +use codex_utils_pty::process_group::kill_process_group; +#[cfg(target_os = "macos")] +use codex_utils_pty::process_group::kill_process_group_with_member_fallback as kill_process_group; +#[cfg(all(unix, not(target_os = "macos")))] +use codex_utils_pty::process_group::terminate_process_group; +#[cfg(target_os = "macos")] +use codex_utils_pty::process_group::terminate_process_group_with_member_fallback as terminate_process_group; +use futures::FutureExt; +use futures::future::BoxFuture; +use rmcp::service::RoleClient; +use rmcp::service::RxJsonRpcMessage; +use rmcp::service::TxJsonRpcMessage; +use rmcp::transport::Transport; +use tokio::io::AsyncBufReadExt; +use tokio::io::BufReader; +use tokio::process::Command; +use tokio::sync::watch; +use tokio::time::Instant; +use tracing::info; +use tracing::warn; + +use crate::executor_process_transport::ExecutorProcessTransport; +use crate::local_stdio_transport::LocalStdioTransport; +use crate::program_resolver; +use crate::protocol_mode::McpProtocolMode; +use crate::utils::create_env_for_mcp_server; +use crate::utils::create_env_overlay_for_remote_mcp_server; +use crate::utils::remote_mcp_env_var_names; + +// General purpose public code. + +/// Launches an MCP stdio server and returns the transport for rmcp. +/// +/// This trait is the boundary between MCP lifecycle code and process placement. +/// `RmcpClient` owns MCP operations such as `initialize` and `tools/list`; the +/// launcher owns starting the configured command and producing an rmcp +/// [`Transport`] over the server's stdin/stdout bytes. +pub trait StdioServerLauncher: private::Sealed + Send + Sync { + /// Start the configured stdio server and return its rmcp-facing transport. + fn launch( + &self, + command: StdioServerCommand, + ) -> BoxFuture<'static, io::Result>; +} + +/// Command-line process shape shared by stdio server launchers. +#[derive(Clone)] +pub struct StdioServerCommand { + program: OsString, + args: Vec, + env: Option>, + env_vars: Vec, + cwd: Option, + protocol_mode: McpProtocolMode, +} + +/// Client-side rmcp transport for a launched MCP stdio server. +/// +/// The concrete process placement stays private to this module. `RmcpClient` +/// only sees the standard rmcp transport abstraction and can pass this value +/// directly to `rmcp::service::serve_client`. +pub struct StdioServerTransport { + inner: StdioServerTransportInner, + process: StdioServerProcessHandle, +} + +enum StdioServerTransportInner { + Local(LocalStdioTransport), + Executor(ExecutorProcessTransport), +} + +impl Transport for StdioServerTransport { + type Error = io::Error; + + fn send( + &mut self, + item: TxJsonRpcMessage, + ) -> impl Future> + Send + 'static { + // Both variants already implement rmcp's transport contract. This + // wrapper keeps process placement private while leaving rmcp's send + // semantics unchanged. + match &mut self.inner { + StdioServerTransportInner::Local(transport) => transport.send(item).boxed(), + StdioServerTransportInner::Executor(transport) => transport.send(item).boxed(), + } + } + + fn receive(&mut self) -> impl Future>> + Send { + // rmcp reads from the same transport shape for both placements. The + // executor variant turns pushed process-output events back into the + // line-delimited JSON stream expected by rmcp. + match &mut self.inner { + StdioServerTransportInner::Local(transport) => transport.receive().boxed(), + StdioServerTransportInner::Executor(transport) => transport.receive().boxed(), + } + } + + async fn close(&mut self) -> std::result::Result<(), Self::Error> { + self.process.terminate().await?; + match &mut self.inner { + StdioServerTransportInner::Local(transport) => transport.close().await, + StdioServerTransportInner::Executor(transport) => transport.close().await, + } + } +} + +impl StdioServerTransport { + pub(crate) fn process_handle(&self) -> StdioServerProcessHandle { + self.process.clone() + } +} + +impl StdioServerCommand { + /// Build the stdio process parameters before choosing where the process + /// runs. + pub(super) fn new( + program: OsString, + args: Vec, + env: Option>, + env_vars: Vec, + cwd: Option, + protocol_mode: McpProtocolMode, + ) -> Self { + Self { + program, + args, + env, + env_vars, + cwd, + protocol_mode, + } + } +} + +// Local public implementation. + +/// Starts MCP stdio servers as local child processes. +/// +/// This is the existing behavior for local MCP servers: the orchestrator +/// process spawns the configured command and rmcp talks to the child's local +/// stdin/stdout pipes directly. +#[derive(Clone)] +pub struct LocalStdioServerLauncher { + fallback_cwd: PathBuf, +} + +impl LocalStdioServerLauncher { + /// Creates a local stdio launcher. + /// + /// `fallback_cwd` is used when the MCP server config omits `cwd`, so + /// relative commands resolve from the caller's runtime working directory. + pub fn new(fallback_cwd: PathBuf) -> Self { + Self { fallback_cwd } + } +} + +impl StdioServerLauncher for LocalStdioServerLauncher { + fn launch( + &self, + command: StdioServerCommand, + ) -> BoxFuture<'static, io::Result> { + let fallback_cwd = self.fallback_cwd.clone(); + async move { + // Keep synchronous program resolution and process creation from blocking the + // caller's startup deadline. + tokio::task::spawn_blocking(move || Self::launch_server(command, fallback_cwd)) + .await + .map_err(io::Error::other)? + } + .boxed() + } +} + +// Local private implementation. + +#[cfg(unix)] +const PROCESS_GROUP_TERM_GRACE_PERIOD: Duration = Duration::from_secs(2); + +// Keep queued stderr diagnostics before closing the reader, even when an +// escaped descendant prevents the pipe from reaching EOF. +const STDERR_READER_DRAIN_GRACE_PERIOD: Duration = Duration::from_millis(250); + +#[cfg(unix)] +struct LocalProcessTerminator { + process_group_id: u32, +} + +#[cfg(windows)] +enum LocalProcessTerminator { + Job(codex_utils_pty::JobObject), + Process(OwnedHandle), +} + +#[cfg(not(any(unix, windows)))] +struct LocalProcessTerminator; + +#[derive(Clone)] +pub(crate) struct StdioServerProcessHandle { + inner: Arc, +} + +struct StdioServerProcessHandleInner { + program_name: String, + kind: StdioServerProcessKind, + terminated: AtomicBool, + // An escaped descendant can keep stderr open after the MCP server exits. + stderr_reader: Option>, +} + +enum StdioServerProcessKind { + Local(Option), + Executor(Arc), +} + +mod private { + pub trait Sealed {} +} + +impl private::Sealed for LocalStdioServerLauncher {} + +impl LocalStdioServerLauncher { + fn launch_server( + command: StdioServerCommand, + fallback_cwd: PathBuf, + ) -> io::Result { + let StdioServerCommand { + program, + args, + env, + env_vars, + cwd, + protocol_mode, + } = command; + let program_name = program.to_string_lossy().into_owned(); + let envs = create_env_for_mcp_server(env, &env_vars).map_err(io::Error::other)?; + let cwd = cwd.map(PathBuf::from).unwrap_or(fallback_cwd); + let resolved_program = + program_resolver::resolve(program, &envs, &cwd).map_err(io::Error::other)?; + + let build_command = || { + let mut command = Command::new(&resolved_program); + command + .current_dir(&cwd) + .env_clear() + .envs(&envs) + .args(&args); + #[cfg(unix)] + command.process_group(0); + command + }; + #[cfg(windows)] + let mut command = build_command(); + #[cfg(not(windows))] + let command = build_command(); + #[cfg(windows)] + let job = match codex_utils_pty::JobObject::create_without_breakaway() { + Ok(job) => { + job.prepare_suspended_spawn(&mut command); + Some(job) + } + Err(error) => { + warn!("Windows MCP process job containment unavailable: {error}"); + None + } + }; + + let spawn_transport = |command: Command| -> io::Result<( + StdioServerTransportInner, + Option, + Option, + )> { + let (transport, stderr) = + LocalStdioTransport::spawn(command, program_name.clone(), protocol_mode)?; + let process_id = transport.id(); + Ok(( + StdioServerTransportInner::Local(transport), + stderr, + process_id, + )) + }; + let (transport, stderr, process_id) = spawn_transport(command)?; + #[cfg(windows)] + let (transport, stderr, process_id, job) = match job { + Some(job) => match process_id + .ok_or_else(|| io::Error::other("missing suspended MCP server process id")) + .and_then(|process_id| job.assign_and_resume_process(process_id)) + { + Ok(true) => (transport, stderr, process_id, Some(job)), + Ok(false) => (transport, stderr, process_id, None), + Err(error) => { + warn!( + "Windows MCP process job containment failed; retrying without it: {error}" + ); + drop(stderr); + drop(transport); + drop(job); + let (transport, stderr, process_id) = spawn_transport(build_command())?; + (transport, stderr, process_id, None) + } + }, + None => (transport, stderr, process_id, None), + }; + #[cfg(windows)] + let terminator = match job { + Some(job) => Some(LocalProcessTerminator::Job(job)), + None => process_id.and_then(|process_id| { + match codex_utils_pty::JobObject::open_process_handle(process_id) { + Ok(handle) => Some(LocalProcessTerminator::Process(handle)), + Err(error) => { + warn!("Windows MCP process handle unavailable: {error}"); + None + } + } + }), + }; + #[cfg(not(windows))] + let terminator = process_id.map(LocalProcessTerminator::new); + let stderr_reader = stderr.map(|stderr| { + let program_name = program_name.clone(); + let (stop_tx, mut stop_rx) = watch::channel(()); + std::mem::drop(tokio::spawn(async move { + let mut reader = BufReader::new(stderr).lines(); + // Give queued diagnostics time to reach the logs without waiting + // indefinitely for a descendant that still has stderr open. + let drain_deadline = tokio::time::sleep(STDERR_READER_DRAIN_GRACE_PERIOD); + tokio::pin!(drain_deadline); + let mut draining = false; + loop { + tokio::select! { + biased; + _ = &mut drain_deadline, if draining => break, + _ = stop_rx.changed(), if !draining => { + draining = true; + drain_deadline.as_mut().reset( + Instant::now() + STDERR_READER_DRAIN_GRACE_PERIOD + ); + } + line = reader.next_line() => { + match line { + Ok(Some(line)) => { + info!("MCP server stderr ({program_name}): {line}"); + } + Ok(None) => break, + Err(error) => { + warn!("Failed to read MCP server stderr ({program_name}): {error}"); + break; + } + } + }, + } + } + })); + stop_tx + }); + let process = StdioServerProcessHandle::local(program_name, terminator, stderr_reader); + + Ok(StdioServerTransport { + inner: transport, + process, + }) + } +} + +impl LocalProcessTerminator { + #[cfg(not(windows))] + fn new(process_group_id: u32) -> Self { + #[cfg(unix)] + { + Self { process_group_id } + } + #[cfg(not(any(unix, windows)))] + { + let _ = process_group_id; + Self + } + } + + #[cfg(unix)] + fn terminate(&self) { + let process_group_id = self.process_group_id; + let should_escalate = match terminate_process_group(process_group_id) { + Ok(exists) => exists, + Err(error) => { + warn!("Failed to terminate MCP process group {process_group_id}: {error}"); + false + } + }; + if should_escalate { + spawn(move || { + sleep(PROCESS_GROUP_TERM_GRACE_PERIOD); + if let Err(error) = kill_process_group(process_group_id) { + warn!("Failed to kill MCP process group {process_group_id}: {error}"); + } + }); + } + } + + #[cfg(windows)] + fn terminate(&self) { + let result = match self { + Self::Job(job) => job.terminate(), + Self::Process(process_handle) => { + codex_utils_pty::JobObject::terminate_process_handle(process_handle) + } + }; + if let Err(error) = result { + warn!("Failed to terminate Windows MCP process: {error}"); + } + } + + #[cfg(not(any(unix, windows)))] + fn terminate(&self) {} +} + +impl StdioServerProcessHandle { + fn local( + program_name: String, + terminator: Option, + stderr_reader: Option>, + ) -> Self { + Self { + inner: Arc::new(StdioServerProcessHandleInner { + program_name, + kind: StdioServerProcessKind::Local(terminator), + terminated: AtomicBool::new(false), + stderr_reader, + }), + } + } + + pub(crate) fn executor(program_name: String, process: Arc) -> Self { + Self { + inner: Arc::new(StdioServerProcessHandleInner { + program_name, + kind: StdioServerProcessKind::Executor(process), + terminated: AtomicBool::new(false), + stderr_reader: None, + }), + } + } + + pub(crate) async fn terminate(&self) -> io::Result<()> { + if self.inner.terminated.swap(true, Ordering::AcqRel) { + return Ok(()); + } + + let result = match &self.inner.kind { + StdioServerProcessKind::Local(Some(terminator)) => { + terminator.terminate(); + Ok(()) + } + StdioServerProcessKind::Local(None) => Ok(()), + StdioServerProcessKind::Executor(process) => match process.terminate().await { + Ok(()) => Ok(()), + Err(error) => { + self.inner.terminated.store(false, Ordering::Release); + Err(io::Error::other(error)) + } + }, + }; + if let Some(stderr_reader) = &self.inner.stderr_reader { + stderr_reader.send_replace(()); + } + result + } +} + +impl Drop for StdioServerProcessHandleInner { + fn drop(&mut self) { + if self.terminated.swap(true, Ordering::AcqRel) { + return; + } + + match &self.kind { + StdioServerProcessKind::Local(Some(terminator)) => { + terminator.terminate(); + } + StdioServerProcessKind::Local(None) => {} + StdioServerProcessKind::Executor(process) => { + let process = Arc::clone(process); + let program_name = self.program_name.clone(); + let Ok(handle) = tokio::runtime::Handle::try_current() else { + warn!( + "Could not schedule remote MCP server process termination on drop ({}): no Tokio runtime is available", + self.program_name + ); + return; + }; + + std::mem::drop(handle.spawn(async move { + if let Err(error) = process.terminate().await { + warn!( + "Failed to terminate remote MCP server process on drop ({program_name}): {error}" + ); + } + })); + } + } + if let Some(stderr_reader) = &self.stderr_reader { + stderr_reader.send_replace(()); + } + } +} + +// Remote public implementation. + +/// Starts MCP stdio servers through the executor process API. +/// +/// MCP framing still runs in the orchestrator. The executor only owns the +/// child process and transports raw stdin/stdout/stderr bytes, so it does not +/// need to know about MCP methods such as `initialize` or `tools/list`. +/// +/// Windows executor-backed servers retain the executor's normal descendant +/// lifetime. MCP-specific containment requires negotiated process ownership: +/// caller-controlled process IDs cannot safely select a destructive policy, +/// and a wrapper may exit while its descendants continue serving requests. +#[derive(Clone)] +pub struct ExecutorStdioServerLauncher { + exec_backend: Arc, +} + +impl ExecutorStdioServerLauncher { + /// Creates a stdio server launcher backed by the executor process API. + pub fn new(exec_backend: Arc) -> Self { + Self { exec_backend } + } +} + +impl StdioServerLauncher for ExecutorStdioServerLauncher { + fn launch( + &self, + command: StdioServerCommand, + ) -> BoxFuture<'static, io::Result> { + let exec_backend = Arc::clone(&self.exec_backend); + async move { Self::launch_server(command, exec_backend).await }.boxed() + } +} + +// Remote private implementation. + +impl private::Sealed for ExecutorStdioServerLauncher {} + +impl ExecutorStdioServerLauncher { + async fn launch_server( + command: StdioServerCommand, + exec_backend: Arc, + ) -> io::Result { + let StdioServerCommand { + program, + args, + env, + env_vars, + cwd, + protocol_mode: _, + } = command; + let Some(cwd) = cwd else { + return Err(io::Error::other( + "executor stdio server requires an explicit cwd", + )); + }; + let cwd: PathUri = LegacyAppPathString::from_path(Path::new(&cwd)) + .try_into() + .map_err(|err| io::Error::new(io::ErrorKind::InvalidInput, err))?; + let program_name = program.to_string_lossy().into_owned(); + let envs = create_env_overlay_for_remote_mcp_server(env, &env_vars); + let remote_env_vars = remote_mcp_env_var_names(&env_vars); + // The executor protocol carries argv/env as UTF-8 strings. Local stdio can + // accept arbitrary OsString values because it calls the OS directly; remote + // stdio must reject non-Unicode command, argument, or environment data + // before sending an executor request. + let argv = Self::process_api_argv(&program, &args).map_err(io::Error::other)?; + let env = Self::process_api_env(envs).map_err(io::Error::other)?; + let process_id = ExecutorProcessTransport::next_process_id(); + // Start the MCP server process on the executor with raw pipes. `tty=false` + // keeps stdout as a clean protocol stream, while `pipe_stdin=true` lets + // rmcp write JSON-RPC requests after the process starts. + let started = exec_backend + .start(ExecParams { + metadata: Default::default(), + process_id, + argv, + cwd, + shell_snapshot: None, + env_policy: Some(Self::remote_env_policy(&remote_env_vars)), + env, + tty: false, + pipe_stdin: true, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await + .map_err(io::Error::other)?; + + let process = + StdioServerProcessHandle::executor(program_name.clone(), Arc::clone(&started.process)); + Ok(StdioServerTransport { + inner: StdioServerTransportInner::Executor(ExecutorProcessTransport::new( + started.process, + program_name, + )), + process, + }) + } + + fn process_api_argv(program: &OsString, args: &[OsString]) -> Result> { + let mut argv = Vec::with_capacity(args.len() + 1); + argv.push(Self::os_string_to_process_api_string( + program.clone(), + "command", + )?); + for arg in args { + argv.push(Self::os_string_to_process_api_string( + arg.clone(), + "argument", + )?); + } + Ok(argv) + } + + fn process_api_env(env: HashMap) -> Result> { + env.into_iter() + .map(|(key, value)| { + Ok(( + Self::os_string_to_process_api_string(key, "environment variable name")?, + Self::os_string_to_process_api_string(value, "environment variable value")?, + )) + }) + .collect() + } + + fn os_string_to_process_api_string(value: OsString, label: &str) -> Result { + value + .into_string() + .map_err(|_| anyhow!("{label} must be valid Unicode for remote MCP stdio")) + } + + fn remote_env_policy(remote_env_vars: &[String]) -> ExecEnvPolicy { + let include_only = if remote_env_vars.is_empty() { + Vec::new() + } else { + // `source = "remote"` means the value is read from the executor's + // environment, not copied from Codex. Start from `All` only so the + // named remote variable is available to the filter below; the + // effective child env is still limited by `include_only`. + crate::utils::DEFAULT_ENV_VARS + .iter() + .map(|name| (*name).to_string()) + .chain(remote_env_vars.iter().cloned()) + .collect() + }; + ExecEnvPolicy { + inherit: if remote_env_vars.is_empty() { + ShellEnvironmentPolicyInherit::Core + } else { + ShellEnvironmentPolicyInherit::All + }, + ignore_default_excludes: true, + exclude: Vec::new(), + r#set: HashMap::new(), + include_only, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use codex_protocol::config_types::EnvironmentVariablePattern; + use codex_protocol::config_types::ShellEnvironmentPolicy; + use codex_protocol::shell_environment; + + #[test] + fn remote_env_policy_uses_core_env_without_remote_source_vars() { + let policy = ExecutorStdioServerLauncher::remote_env_policy(&[]); + + assert_eq!(policy.inherit, ShellEnvironmentPolicyInherit::Core); + assert!(policy.include_only.is_empty()); + } + + #[test] + fn remote_env_policy_includes_remote_source_vars_without_full_env() { + let policy = ExecutorStdioServerLauncher::remote_env_policy(&["REMOTE_TOKEN".to_string()]); + + assert_eq!(policy.inherit, ShellEnvironmentPolicyInherit::All); + assert!( + policy.include_only.contains(&"REMOTE_TOKEN".to_string()), + "remote source var should be included in executor env policy" + ); + assert!( + policy + .include_only + .contains(&crate::utils::DEFAULT_ENV_VARS[0].to_string()), + "remote default env vars should remain available" + ); + } + + #[test] + fn remote_env_policy_effectively_filters_unrequested_vars() { + let exec_policy = + ExecutorStdioServerLauncher::remote_env_policy(&["REMOTE_TOKEN".to_string()]); + let policy = ShellEnvironmentPolicy { + inherit: exec_policy.inherit, + ignore_default_excludes: exec_policy.ignore_default_excludes, + exclude: exec_policy + .exclude + .iter() + .map(|pattern| EnvironmentVariablePattern::new_case_insensitive(pattern)) + .collect(), + r#set: exec_policy.r#set, + include_only: exec_policy + .include_only + .iter() + .map(|pattern| EnvironmentVariablePattern::new_case_insensitive(pattern)) + .collect(), + use_profile: false, + }; + + let env = shell_environment::create_env_from_vars( + [ + ("PATH".to_string(), "/remote/bin".to_string()), + ("REMOTE_TOKEN".to_string(), "remote-secret".to_string()), + ( + "UNREQUESTED_SECRET".to_string(), + "must-not-pass".to_string(), + ), + ], + &policy, + /*thread_id*/ None, + ); + + assert_eq!(env.get("PATH").map(String::as_str), Some("/remote/bin")); + assert_eq!( + env.get("REMOTE_TOKEN").map(String::as_str), + Some("remote-secret") + ); + assert!(!env.contains_key("UNREQUESTED_SECRET")); + } +} diff --git a/codex-rs/rmcp-client/src/streamable_http_retry.rs b/codex-rs/rmcp-client/src/streamable_http_retry.rs new file mode 100644 index 0000000000000000000000000000000000000000..66dda2cb59c362f3cf8c78f9bcc595ac6652bdce --- /dev/null +++ b/codex-rs/rmcp-client/src/streamable_http_retry.rs @@ -0,0 +1,253 @@ +use std::sync::Arc; +use std::time::Duration; +use std::time::Instant; + +use anyhow::Result; +use anyhow::anyhow; +use codex_exec_server::ExecServerError; +use http::StatusCode; +use rmcp::service::RoleClient; +use rmcp::service::RunningService; +use rmcp::transport::streamable_http_client::StreamableHttpError; +use tokio::time; +use tracing::warn; + +use crate::elicitation_client_service::ElicitationClientService; +use crate::http_client_adapter::StreamableHttpClientAdapterError; +use crate::oauth::OAuthRuntime; + +use super::InitializeContext; +use super::PendingTransport; +use super::RmcpClient; + +const JSON_RPC_INTERNAL_ERROR_CODE: i64 = -32603; +pub(super) const STREAMABLE_HTTP_RETRY_DELAYS_MS: [u64; 2] = [250, 1_000]; + +impl RmcpClient { + pub(super) async fn connect_pending_transport_with_initialize_retries( + &self, + initial_transport: PendingTransport, + initialize_context: &InitializeContext, + ) -> Result<( + Arc>, + Option, + )> { + let timeout = initialize_context.timeout; + let should_retry = match &initial_transport { + PendingTransport::InProcess { .. } | PendingTransport::Stdio { .. } => false, + PendingTransport::StreamableHttp { .. } + | PendingTransport::StreamableHttpWithOAuth { .. } + | PendingTransport::StreamableHttpWithAccessTokenOnly { .. } => true, + }; + let mut retry_deadline = timeout.map(|duration| Instant::now() + duration); + let mut pending_transport = Some(initial_transport); + + for (attempt, retry_delay_ms) in STREAMABLE_HTTP_RETRY_DELAYS_MS + .iter() + .copied() + .map(Some) + .chain(std::iter::once(None)) + .enumerate() + { + let transport = match pending_transport.take() { + Some(transport) => transport, + None => { + let remaining = remaining_initialize_timeout(timeout, retry_deadline)?; + match remaining { + Some(remaining) => time::timeout( + remaining, + Self::create_pending_transport(&self.transport_recipe), + ) + .await + .map_err(|_| initialize_timeout_error(timeout, remaining))??, + None => Self::create_pending_transport(&self.transport_recipe).await?, + } + } + }; + if let PendingTransport::StreamableHttpWithOAuth { oauth_runtime, .. } = &transport { + // OAuth refresh has its own lock and provider request bounds. Exclude it from the + // MCP handshake budget, and finish persistence before attempting initialize. + let refresh_started_at = Instant::now(); + oauth_runtime.refresh_if_needed().await?; + if let Some(deadline) = retry_deadline.as_mut() { + *deadline += refresh_started_at.elapsed(); + } + } + let attempt_timeout = remaining_initialize_timeout(timeout, retry_deadline)?; + + match self + .connect_pending_transport(transport, initialize_context, attempt_timeout) + .await + { + Ok(result) => return Ok(result), + Err(error) if should_retry && Self::is_retryable_initialize_error(&error) => { + let Some(retry_delay_ms) = retry_delay_ms else { + return Err(error); + }; + let delay = Duration::from_millis(retry_delay_ms); + warn!( + attempt = attempt + 1, + max_attempts = STREAMABLE_HTTP_RETRY_DELAYS_MS.len() + 1, + delay_ms = delay.as_millis(), + error = %error, + "streamable HTTP MCP initialize failed with a retryable error; retrying" + ); + if !sleep_with_retry_deadline(delay, retry_deadline).await { + let duration = timeout.unwrap_or(delay); + return Err(anyhow!( + "timed out handshaking with MCP server after {duration:?}" + )); + } + } + Err(error) => return Err(error), + } + } + + unreachable!("initialize retry loop should return on success or final error") + } + + fn is_retryable_initialize_error(error: &anyhow::Error) -> bool { + error.chain().any(|source| { + source + .downcast_ref::() + .is_some_and(|error| Self::is_retryable_client_initialize_error(&error.source)) + || source + .downcast_ref::() + .is_some_and(Self::is_retryable_client_initialize_error) + }) + } + + fn is_retryable_client_initialize_error(error: &rmcp::service::ClientInitializeError) -> bool { + match error { + rmcp::service::ClientInitializeError::LegacyFallbackFailed { fallback, .. } => { + Self::is_retryable_client_initialize_error(fallback) + } + rmcp::service::ClientInitializeError::TransportError { error, context } + if matches!( + context.as_ref(), + "send initialize request" | "send discover request" + ) => + { + error + .error + .downcast_ref::>() + .is_some_and(Self::is_retryable_streamable_http_error) + } + rmcp::service::ClientInitializeError::TransportError { error, context } + if context.as_ref() == "send initialized notification" => + { + error + .error + .downcast_ref::>() + .is_some_and(|error| { + matches!(error, StreamableHttpError::TransportChannelClosed) + || Self::is_retryable_streamable_http_error(error) + }) + } + _ => false, + } + } + + pub(super) fn is_retryable_streamable_http_error( + error: &StreamableHttpError, + ) -> bool { + match error { + StreamableHttpError::Client(StreamableHttpClientAdapterError::HttpRequest( + ExecServerError::HttpRequest(_), + )) => true, + StreamableHttpError::Client(StreamableHttpClientAdapterError::HttpRequest( + ExecServerError::Server { code, message }, + )) => { + *code == JSON_RPC_INTERNAL_ERROR_CODE && message.starts_with("http/request failed:") + } + StreamableHttpError::Client(StreamableHttpClientAdapterError::HttpRequest( + ExecServerError::Protocol(message), + )) => message.starts_with("http response stream `") && message.contains("` failed:"), + StreamableHttpError::UnexpectedServerResponse(message) => { + is_retryable_unexpected_server_response(message.as_ref()) + } + StreamableHttpError::AuthRequired(_) + | StreamableHttpError::InsufficientScope(_) + | StreamableHttpError::SessionExpired + | StreamableHttpError::UnexpectedContentType(_) + | StreamableHttpError::ServerDoesNotSupportSse + | StreamableHttpError::Deserialize(_) + | StreamableHttpError::Client(StreamableHttpClientAdapterError::SessionExpired404) + | StreamableHttpError::Client(StreamableHttpClientAdapterError::Header(_)) => false, + _ => false, + } + } +} + +fn is_retryable_unexpected_server_response(message: &str) -> bool { + let Some(message) = message.strip_prefix("HTTP ") else { + return false; + }; + let status_code = message + .chars() + .take_while(char::is_ascii_digit) + .collect::(); + let Ok(status) = status_code.parse::() else { + return false; + }; + let Ok(status) = StatusCode::from_u16(status) else { + return false; + }; + is_retryable_http_status(status) +} + +fn is_retryable_http_status(status: StatusCode) -> bool { + matches!( + status, + StatusCode::REQUEST_TIMEOUT + | StatusCode::TOO_MANY_REQUESTS + | StatusCode::INTERNAL_SERVER_ERROR + | StatusCode::BAD_GATEWAY + | StatusCode::SERVICE_UNAVAILABLE + | StatusCode::GATEWAY_TIMEOUT + ) +} + +fn remaining_initialize_timeout( + timeout: Option, + deadline: Option, +) -> Result> { + let Some(deadline) = deadline else { + return Ok(None); + }; + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + Err(initialize_timeout_error(timeout, remaining)) + } else { + Ok(Some(remaining)) + } +} + +fn initialize_timeout_error(timeout: Option, fallback: Duration) -> anyhow::Error { + let duration = timeout.unwrap_or(fallback); + anyhow!("timed out handshaking with MCP server after {duration:?}") +} + +pub(super) async fn sleep_with_retry_deadline(delay: Duration, deadline: Option) -> bool { + if let Some(deadline) = deadline { + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + return false; + } + time::timeout(remaining, time::sleep(delay)).await.is_ok() + } else { + time::sleep(delay).await; + true + } +} + +#[derive(Debug, thiserror::Error)] +#[error("handshaking with MCP server failed: {source}")] +pub(super) struct HandshakeError { + #[source] + pub(super) source: rmcp::service::ClientInitializeError, +} + +#[cfg(test)] +#[path = "streamable_http_retry_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/streamable_http_retry_tests.rs b/codex-rs/rmcp-client/src/streamable_http_retry_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..efda30dad60416e9e96c08acf80877d1ca3ea749 --- /dev/null +++ b/codex-rs/rmcp-client/src/streamable_http_retry_tests.rs @@ -0,0 +1,107 @@ +use std::any::TypeId; + +use codex_exec_server::ExecServerError; +use pretty_assertions::assert_eq; +use rmcp::transport::DynamicTransportError; +use rmcp::transport::streamable_http_client::AuthRequiredError; +use rmcp::transport::streamable_http_client::StreamableHttpError; + +use crate::http_client_adapter::StreamableHttpClientAdapterError; +use crate::rmcp_client::ClientOperationError; + +use super::*; + +#[test] +fn retryable_initialize_error_includes_discovery_and_initialized_notification_context() { + let contexts = [ + "send discover request", + "send initialize request", + "send initialized notification", + "receive initialize response", + ]; + + assert_eq!( + contexts.map(|context| { + RmcpClient::is_retryable_client_initialize_error(&retryable_initialize_error(context)) + }), + [true, true, true, false], + ); +} + +#[test] +fn retryable_streamable_http_error_includes_remote_body_stream_failure() { + let errors = [ + StreamableHttpError::Client(StreamableHttpClientAdapterError::HttpRequest( + ExecServerError::HttpRequest("error sending request for url".to_string()), + )), + StreamableHttpError::Client(StreamableHttpClientAdapterError::HttpRequest( + ExecServerError::Server { + code: JSON_RPC_INTERNAL_ERROR_CODE, + message: "http/request failed: error sending request for url".to_string(), + }, + )), + StreamableHttpError::Client(StreamableHttpClientAdapterError::HttpRequest( + ExecServerError::Protocol( + "http response stream `http-1` failed: exec-server transport disconnected" + .to_string(), + ), + )), + StreamableHttpError::Client(StreamableHttpClientAdapterError::HttpRequest( + ExecServerError::Protocol( + "http response stream `http-1` received seq 2, expected 1".to_string(), + ), + )), + StreamableHttpError::UnexpectedServerResponse("HTTP 502: upstream failure".into()), + StreamableHttpError::UnexpectedServerResponse("HTTP 400: bad request".into()), + ]; + + assert_eq!( + errors.map(|error| RmcpClient::is_retryable_streamable_http_error(&error)), + [true, true, true, false, true, false], + ); +} + +#[test] +fn startup_http_authentication_challenges_require_reauthorization() { + let transport_error = || { + DynamicTransportError::from_parts( + "streamable_http", + TypeId::of::<()>(), + Box::new( + StreamableHttpError::::AuthRequired( + AuthRequiredError::new("Bearer error=\"invalid_token\"".to_string()), + ), + ), + ) + }; + let errors = [ + anyhow::Error::new(rmcp::service::ClientInitializeError::TransportError { + error: transport_error(), + context: "send initialize request".into(), + }), + anyhow::Error::new(ClientOperationError::from( + rmcp::service::ServiceError::TransportSend(transport_error()), + )), + ]; + + for error in errors { + assert!(crate::startup_error::is_authentication_required_error( + &error + )); + } +} + +fn retryable_initialize_error(context: &'static str) -> rmcp::service::ClientInitializeError { + rmcp::service::ClientInitializeError::TransportError { + error: DynamicTransportError::from_parts( + "streamable_http", + TypeId::of::<()>(), + Box::new(StreamableHttpError::Client( + StreamableHttpClientAdapterError::HttpRequest(ExecServerError::HttpRequest( + "error sending request for url".to_string(), + )), + )), + ), + context: context.into(), + } +} diff --git a/codex-rs/rmcp-client/src/tool_input.rs b/codex-rs/rmcp-client/src/tool_input.rs new file mode 100644 index 0000000000000000000000000000000000000000..ffda73988dfc6a0ec4c9a87ab81375ecbf22544c --- /dev/null +++ b/codex-rs/rmcp-client/src/tool_input.rs @@ -0,0 +1,208 @@ +//! Drives modern tool continuations, including the native verification extension. +//! Inputs use the existing client service; proofs never bypass its validation. +//! Session recovery cannot replay a submitted proof; connection closure drops pending inputs. + +use std::collections::BTreeMap; +use std::time::Duration; + +use rmcp::RoleClient; +use rmcp::model::CallToolRequest; +use rmcp::model::CallToolRequestParams; +use rmcp::model::CallToolResult; +use rmcp::model::ClientRequest; +use rmcp::model::DEFAULT_MRTR_MAX_ROUNDS; +use rmcp::model::GetExtensions; +use rmcp::model::GetMeta; +use rmcp::model::InputRequest; +use rmcp::model::RequestId; +use rmcp::model::ServerRequest; +use rmcp::model::ServerResult; +use rmcp::service::PeerRequestOptions; +use rmcp::service::RequestContext; +use rmcp::service::RunningService; +use rmcp::service::Service; +use rmcp::service::ServiceError; +use rmcp::transport::streamable_http_client::StreamableHttpError; +use serde::Deserialize; +use serde_json::Value; + +use crate::elicitation_client_service::ElicitationClientService; +use crate::http_client_adapter::StreamableHttpClientAdapterError; + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase")] +struct ToolInput { + input_requests: Option>, + request_state: Option, +} + +pub(crate) async fn call_tool( + service: &RunningService, + mut params: CallToolRequestParams, +) -> Result { + let mut state_only_rounds = 0; + let mut submitted_verification_proof = false; + for round in 0..DEFAULT_MRTR_MAX_ROUNDS { + let response = async { + let handle = service + .peer() + .send_request_with_option( + ClientRequest::CallToolRequest(CallToolRequest::new(params.clone())), + PeerRequestOptions::no_options(), + ) + .await?; + let id = handle.id.clone(); + Ok::<_, ServiceError>((id, handle.await_response().await?)) + } + .await; + let (id, result) = response.map_err(|error| { + if submitted_verification_proof + && let ServiceError::TransportSend(transport) = &error + && !matches!( + transport + .error + .downcast_ref::>(), + Some(StreamableHttpError::AuthRequired(_)) + ) + { + // Session recovery must not replay the original operation after a + // continuation (which may already have consumed a signed proof). + // Auth challenges already return without replaying the operation. + invalid("MCP tool continuation failed; the tool was not restarted") + } else { + error + } + })?; + let input = match result { + ServerResult::CallToolResult(result) => return Ok(result), + ServerResult::InputRequiredResult(result) => serde_json::to_value(result) + .map_err(|_| invalid("invalid MCP tool input request"))?, + // RMCP's InputRequest union intentionally excludes custom methods. + ServerResult::CustomResult(result) + if result.0.get("resultType").and_then(Value::as_str) == Some("input_required") => + { + result.0 + } + _ => return Err(ServiceError::UnexpectedResponse), + }; + if round + 1 == DEFAULT_MRTR_MAX_ROUNDS { + break; + } + let input: ToolInput = + serde_json::from_value(input).map_err(|_| invalid("invalid MCP tool input request"))?; + let requests = input.input_requests.unwrap_or_default(); + if requests.is_empty() && input.request_state.is_none() { + return Err(ServiceError::UnexpectedResponse); + } + // Parse the whole round before presenting any prompts. Only standard MCP + // inputs and supported OpenAI elicitation modes are accepted. + let requests = requests + .into_iter() + .map(|(key, value)| parse_input(value).map(|request| (key, request))) + .collect::, _>>()?; + if requests.is_empty() { + let millis = (50_u64 << state_only_rounds.min(/*other*/ 3)).min(/*other*/ 250); + tokio::time::sleep(Duration::from_millis(millis)).await; + state_only_rounds += 1; + } else { + state_only_rounds = 0; + } + let responses = futures::future::try_join_all(requests.into_iter().enumerate().map( + |(index, (key, mut request))| { + // Native prompts need distinct UI/cancellation ownership when + // concurrent tool calls use the same server-assigned input key. + let native_verification = matches!(&request, ServerRequest::CustomRequest(request) + if request.params.as_ref().and_then(|params| params.get("mode")) + .and_then(Value::as_str) == Some(crate::user_verification::MODE)); + let request_id = if native_verification { + RequestId::String(format!("tool-input/{id}/{index}").into()) + } else { + RequestId::String(key.clone().into()) + }; + let mut context = RequestContext::new(request_id, service.peer().clone()); + context.meta = std::mem::take(request.get_meta_mut()); + context.extensions = std::mem::take(request.extensions_mut()); + async move { + let result = service + .service() + .handle_request(request, context) + .await + .map_err(ServiceError::McpError)?; + let result = serde_json::to_value(result) + .map_err(|_| invalid("invalid MCP tool input response"))?; + let contains_proof = native_verification + && result.get("action").and_then(Value::as_str) == Some("accept"); + Ok::<_, ServiceError>((key, result, contains_proof)) + } + }, + )); + let responses = tokio::select! { + biased; + _ = async { + // RMCP exposes closure status but no awaitable closure notification. + // Include explicit cancellation and transport EOF; dropping the joined + // handlers releases both their UI callbacks and timeout pause guards. + while !service.is_closed() && !service.peer().is_transport_closed() { + tokio::time::sleep(Duration::from_millis(/*millis*/ 50)).await; + } + } => return Err(ServiceError::TransportClosed), + responses = responses => responses?, + }; + params.input_responses = (!responses.is_empty()).then(|| { + responses + .into_iter() + .map(|(key, result, contains_proof)| { + submitted_verification_proof |= contains_proof; + (key, result) + }) + .collect() + }); + params.request_state = input.request_state; + } + Err(ServiceError::InputRequiredRoundsExceeded { + max_rounds: DEFAULT_MRTR_MAX_ROUNDS, + }) +} + +fn parse_input(value: Value) -> Result { + if let Ok(request) = serde_json::from_value::(value.clone()) { + return match request { + InputRequest::CreateMessage(request) => { + Ok(ServerRequest::CreateMessageRequest(request)) + } + InputRequest::Elicitation(request) => Ok(ServerRequest::ElicitRequest(request)), + InputRequest::ListRoots(request) => Ok(ServerRequest::ListRootsRequest(request)), + _ => Err(ServiceError::UnexpectedResponse), + }; + } + let request: ServerRequest = + serde_json::from_value(value).map_err(|_| invalid("invalid MCP tool input request"))?; + if let ServerRequest::CustomRequest(request) = &request + && request.method == "openai/elicitation/create" + { + match request + .params + .as_ref() + .and_then(|params| params.get("mode")) + .and_then(Value::as_str) + { + Some("form") => { + crate::elicitation_client_service::openai_elicitation_form(request.clone()) + .map_err(ServiceError::McpError)?; + } + Some(crate::user_verification::MODE) => { + crate::user_verification::parse_request(request.clone()) + .map_err(ServiceError::McpError)?; + } + _ => return Err(invalid("unsupported OpenAI elicitation mode")), + } + return Ok(ServerRequest::CustomRequest(request.clone())); + } + Err(invalid("unsupported MCP tool input request")) +} + +fn invalid(message: &'static str) -> ServiceError { + ServiceError::McpError(rmcp::ErrorData::invalid_request( + message, /*data*/ None, + )) +} diff --git a/codex-rs/rmcp-client/src/tool_input_tests.rs b/codex-rs/rmcp-client/src/tool_input_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..ebbdc962f69a4bdf3e861f5241882a3d0e67ac27 --- /dev/null +++ b/codex-rs/rmcp-client/src/tool_input_tests.rs @@ -0,0 +1,101 @@ +//! Exercises pending MRTR input cleanup over a real duplex transport. + +use super::ElicitationClientService; +use crate::rmcp_client::ElicitationPauseState; +use crate::tool_input::call_tool; +use codex_protocol::mcp::OPENAI_ELICITATION_EXTENSION_ID; +use rmcp::RoleServer; +use rmcp::model::CallToolRequestParams; +use rmcp::model::ClientInfo; +use rmcp::model::ClientJsonRpcMessage; +use rmcp::model::CustomResult; +use rmcp::model::ServerJsonRpcMessage; +use rmcp::model::ServerResult; +use rmcp::service::ServiceError; +use rmcp::service::serve_directly; +use rmcp::transport::IntoTransport; +use rmcp::transport::Transport; +use serde_json::json; +use std::time::Duration; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::time::timeout; + +#[tokio::test] +async fn connection_closure_releases_pending_inputs_and_timeout_pause() -> anyhow::Result<()> { + for mode in ["form", "openai/userVerification"] { + for close_transport in [false, true] { + let pause_state = ElicitationPauseState::new(); + let mut paused = pause_state.subscribe(); + let (route_tx, mut route_rx) = mpsc::unbounded_channel(); + let mut info = ClientInfo::default(); + info.capabilities.extensions = Some( + [( + OPENAI_ELICITATION_EXTENSION_ID.into(), + serde_json::Map::from_iter([ + ("userVerification".into(), json!({})), + ("form".into(), json!({})), + ]), + )] + .into_iter() + .collect(), + ); + let service = ElicitationClientService::new( + info, + Box::new(move |_, _| { + let (tx, rx) = oneshot::channel(); + route_tx.send(tx).unwrap(); + Box::pin(async move { Ok(rx.await?) }) + }), + pause_state, + ); + let (client_transport, server_transport) = + tokio::io::duplex(/*max_buf_size*/ 4096); + let client = serve_directly(service, client_transport, /*peer_info*/ None); + let mut server = IntoTransport::::into_transport(server_transport); + let call = call_tool(&client, CallToolRequestParams::new("test")); + tokio::pin!(call); + let request = tokio::select! { + result = &mut call => panic!("tool completed before input: {result:?}"), + request = timeout(Duration::from_secs(/*secs*/ 5), server.receive()) => request?, + }; + let Some(ClientJsonRpcMessage::Request(request)) = request else { + anyhow::bail!("expected tool call"); + }; + let params = match mode { + "form" => json!({ + "mode": mode, "message": "Confirm", + "requestedSchema": {"type": "object", "properties": {}}, + }), + _ => json!({ + "mode": mode, "title": "Verify", "description": "", "challenge": "AQID", + }), + }; + server.send(ServerJsonRpcMessage::response( + ServerResult::CustomResult(CustomResult(json!({ + "resultType": "input_required", + "inputRequests": {"input": {"method": "openai/elicitation/create", "params": params}}, + }))), + request.id, + )).await?; + let reply = tokio::select! { + result = &mut call => panic!("tool completed before prompting: {result:?}"), + reply = timeout(Duration::from_secs(/*secs*/ 5), route_rx.recv()) => reply?.unwrap(), + }; + assert!(*paused.borrow()); + if close_transport { + drop(server); + } else { + client.cancellation_token().cancel(); + } + let result = timeout(Duration::from_secs(/*secs*/ 5), call).await?; + assert!( + matches!(result, Err(ServiceError::TransportClosed)), + "{result:?}" + ); + assert!(reply.is_closed()); + assert!(!*paused.borrow_and_update()); + } + } + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/user_verification.rs b/codex-rs/rmcp-client/src/user_verification.rs new file mode 100644 index 0000000000000000000000000000000000000000..fb93b42ff64da7fae7565f15aec634dcd59dfd92 --- /dev/null +++ b/codex-rs/rmcp-client/src/user_verification.rs @@ -0,0 +1,91 @@ +//! Validation for the user-verification elicitation extension. + +use base64::Engine as _; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use rmcp::model::CustomRequest; +use rmcp::model::ElicitationAction; +use serde::Deserialize; + +use crate::rmcp_client::Elicitation; +use crate::rmcp_client::ElicitationResponse; + +pub(crate) const MODE: &str = "openai/userVerification"; + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +struct RequestParams { + mode: String, + title: String, + description: String, + challenge: String, +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +struct Proof<'a> { + credential_id: &'a str, + signature: &'a str, +} + +pub(crate) fn parse_request(request: CustomRequest) -> Result { + let params = request + .params_as::() + .ok() + .flatten() + .ok_or_else(invalid_request)?; + if params.mode != MODE + || params.title.is_empty() + || params.title.len() > 256 + || params.description.len() > 4096 + || !valid_bytes(¶ms.challenge, /*max_decoded_bytes*/ 4096) + { + return Err(invalid_request()); + } + Ok(Elicitation::UserVerification { + title: params.title, + description: params.description, + challenge: params.challenge, + }) +} + +/// Accept only a bounded proof, and never return proof material for cancellation or rejection. +pub(crate) fn validate_response(mut response: ElicitationResponse) -> ElicitationResponse { + response.meta = None; + match response.action { + ElicitationAction::Accept => { + let proof = response + .content + .as_ref() + .and_then(|content| Proof::deserialize(content).ok()); + if proof.is_some_and(|proof| { + !proof.credential_id.is_empty() + && proof.credential_id.len() <= 1024 + && valid_bytes(proof.signature, /*max_decoded_bytes*/ 128) + }) { + return response; + } + tracing::warn!("user-verification acceptance omitted a valid proof; cancelling"); + response.action = ElicitationAction::Cancel; + } + ElicitationAction::Decline | ElicitationAction::Cancel => {} + _ => response.action = ElicitationAction::Cancel, + } + response.content = None; + response +} + +fn valid_bytes(encoded: &str, max_decoded_bytes: usize) -> bool { + !encoded.is_empty() + && encoded.len() <= max_decoded_bytes.div_ceil(3) * 4 + && URL_SAFE_NO_PAD + .decode(encoded) + .is_ok_and(|bytes| !bytes.is_empty() && bytes.len() <= max_decoded_bytes) +} + +fn invalid_request() -> rmcp::ErrorData { + rmcp::ErrorData::invalid_params("invalid user-verification request", /*data*/ None) +} + +#[cfg(test)] +#[path = "user_verification_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/user_verification_cancellation_tests.rs b/codex-rs/rmcp-client/src/user_verification_cancellation_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..2fc8a2f1464d0c898b07197d6bd82093beacccbc --- /dev/null +++ b/codex-rs/rmcp-client/src/user_verification_cancellation_tests.rs @@ -0,0 +1,407 @@ +//! Ensures MCP cancellation releases pending elicitation routes and tool timeout pauses. + +use super::ElicitationClientService; +use super::ElicitationPauseState; +use super::ElicitationResponse; +use super::RmcpClient; +use crate::InProcessTransportFactory; +use codex_protocol::mcp::OPENAI_ELICITATION_EXTENSION_ID; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; +use rmcp::RoleServer; +use rmcp::ServerHandler; +use rmcp::ServiceExt; +use rmcp::model::CancelledNotification; +use rmcp::model::CancelledNotificationParam; +use rmcp::model::ClientInfo; +use rmcp::model::ClientJsonRpcMessage; +use rmcp::model::CustomRequest; +use rmcp::model::ElicitationAction; +use rmcp::model::ProtocolVersion; +use rmcp::model::RequestId; +use rmcp::model::ServerJsonRpcMessage; +use rmcp::model::ServerNotification; +use rmcp::model::ServerRequest; +use rmcp::service::PeerRequestOptions; +use rmcp::service::RunningService; +use rmcp::service::ServerInitializeError; +use rmcp::service::serve_directly; +use rmcp::transport::IntoTransport; +use rmcp::transport::Transport; +use serde_json::json; +use std::sync::Arc; +use std::time::Duration; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::time::timeout; +use tracing_test::traced_test; + +struct ElicitationTestServer; + +impl ServerHandler for ElicitationTestServer {} + +struct ElicitationTestTransport { + servers: mpsc::UnboundedSender< + Result, ServerInitializeError>, + >, +} + +impl InProcessTransportFactory for ElicitationTestTransport { + fn open(&self) -> BoxFuture<'static, std::io::Result> { + let servers = self.servers.clone(); + Box::pin(async move { + let (client, server) = tokio::io::duplex(/*max_buf_size*/ 4096); + tokio::spawn(async move { + let _ = servers.send(ElicitationTestServer.serve(server).await); + }); + Ok(client) + }) + } +} + +#[tokio::test] +#[traced_test] +async fn recovered_connections_accept_elicitations_with_previously_cancelled_ids() +-> anyhow::Result<()> { + for params in [ + json!({ + "mode": "form", + "message": "Confirm", + "requestedSchema": {"type": "object", "properties": {}}, + }), + json!({ + "mode": "url", + "message": "Authorize", + "url": "https://example.com/authorize", + "elicitationId": "authorization", + }), + ] { + let (servers, mut server_rx) = mpsc::unbounded_channel(); + let client = + RmcpClient::new_in_process_client(Arc::new(ElicitationTestTransport { servers })) + .await?; + client + .initialize( + ClientInfo::default().with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(/*secs*/ 5)), + Box::new(|_, _| { + Box::pin(async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: None, + meta: None, + }) + }) + }), + ) + .await?; + let server = timeout(Duration::from_secs(/*secs*/ 5), server_rx.recv()) + .await? + .unwrap()?; + let reason = format!("late cancellation for {}", params["mode"]); + let request = + ServerRequest::CustomRequest(CustomRequest::new("elicitation/create", Some(params))); + let original = server + .send_request_with_option(request.clone(), PeerRequestOptions::no_options()) + .await?; + let id = original.id.clone(); + assert_eq!( + serde_json::to_value( + timeout(Duration::from_secs(/*secs*/ 5), original.await_response()).await?? + )?, + json!({"action": "accept"}), + ); + server + .notify_cancelled(CancelledNotificationParam::new( + Some(id.clone()), + Some(reason.clone()), + )) + .await?; + // Notifications run in independent tasks. The handler logs only after recording + // the cancellation, so wait for that acknowledgement before starting recovery. + let handled = format!( + "MCP server cancelled request (request_id: Some({id:?}), reason: Some({reason:?}))" + ); + timeout(Duration::from_secs(/*secs*/ 5), async { + while !logs_contain(&handled) { + tokio::task::yield_now().await; + } + }) + .await?; + + let previous = client.service().await?; + client.reinitialize_after_session_expiry(&previous).await?; + let recovered_server = timeout(Duration::from_secs(/*secs*/ 5), server_rx.recv()) + .await? + .unwrap()?; + let recovered = recovered_server + .send_request_with_option(request, PeerRequestOptions::no_options()) + .await?; + assert_eq!(recovered.id, id); + assert_eq!( + serde_json::to_value( + timeout(Duration::from_secs(/*secs*/ 5), recovered.await_response()).await?? + )?, + json!({"action": "accept"}), + ); + client.shutdown().await; + drop(previous); + server.cancel().await?; + recovered_server.cancel().await?; + } + Ok(()) +} + +#[tokio::test] +async fn ordinary_elicitations_release_pending_responses_on_cancellation() -> anyhow::Result<()> { + for params in [ + json!({ + "mode": "form", + "message": "Confirm", + "requestedSchema": {"type": "object", "properties": {}}, + }), + json!({ + "mode": "url", + "message": "Authorize", + "url": "https://example.com/authorize", + "elicitationId": "authorization", + }), + ] { + let pause_state = ElicitationPauseState::new(); + let mut paused = pause_state.subscribe(); + let (route_tx, mut route_rx) = mpsc::unbounded_channel(); + let service = ElicitationClientService::new( + ClientInfo::default(), + Box::new(move |_, _| { + let (response_tx, response_rx) = oneshot::channel(); + route_tx.send(response_tx).expect("observe elicitation"); + Box::pin(async move { Ok(response_rx.await?) }) + }), + pause_state, + ); + let (client_transport, server_transport) = tokio::io::duplex(/*max_buf_size*/ 4096); + let client = serve_directly(service, client_transport, /*peer_info*/ None); + let mut server = IntoTransport::::into_transport(server_transport); + let request = + ServerRequest::CustomRequest(CustomRequest::new("elicitation/create", Some(params))); + server + .send(ServerJsonRpcMessage::request( + request.clone(), + RequestId::Number(1), + )) + .await?; + let mut response_tx = timeout(Duration::from_secs(/*secs*/ 5), route_rx.recv()) + .await? + .expect("elicitation reached UI without user-verification capability"); + assert!(*paused.borrow()); + + server + .send(ServerJsonRpcMessage::notification( + ServerNotification::CancelledNotification(CancelledNotification::new( + CancelledNotificationParam::new( + Some(RequestId::Number(1)), + /*reason*/ None, + ), + )), + )) + .await?; + timeout(Duration::from_secs(/*secs*/ 5), response_tx.closed()).await?; + timeout( + Duration::from_secs(/*secs*/ 5), + paused.wait_for(|paused| !*paused), + ) + .await??; + let response = timeout(Duration::from_secs(/*secs*/ 5), server.receive()) + .await? + .expect("cancelled elicitation returned a response"); + assert_eq!( + serde_json::to_value(response)?, + json!({"jsonrpc": "2.0", "id": 1, "result": {"action": "cancel"}}), + ); + + server + .send(ServerJsonRpcMessage::request(request, RequestId::Number(2))) + .await?; + let mut response_tx = timeout(Duration::from_secs(/*secs*/ 5), route_rx.recv()) + .await? + .expect("connection still accepts elicitations after cancellation"); + assert!(*paused.borrow()); + + client.cancel().await?; + + timeout(Duration::from_secs(/*secs*/ 5), response_tx.closed()).await?; + timeout( + Duration::from_secs(/*secs*/ 5), + paused.wait_for(|paused| !*paused), + ) + .await??; + } + Ok(()) +} + +#[tokio::test] +async fn user_verification_service_cancellation_drops_pending_response() -> anyhow::Result<()> { + let pause_state = ElicitationPauseState::new(); + let mut paused = pause_state.subscribe(); + let (route_tx, mut route_rx) = mpsc::unbounded_channel(); + let mut client_info = ClientInfo::default(); + client_info.capabilities.extensions = Some( + [( + OPENAI_ELICITATION_EXTENSION_ID.to_string(), + serde_json::Map::from_iter([("userVerification".to_string(), json!({}))]), + )] + .into_iter() + .collect(), + ); + let service = ElicitationClientService::new( + client_info, + Box::new(move |_, _| { + let (response_tx, response_rx) = oneshot::channel(); + route_tx + .send(response_tx) + .expect("observe pending verification"); + Box::pin(async move { Ok(response_rx.await?) }) + }), + pause_state, + ); + let (client_transport, server_transport) = tokio::io::duplex(/*max_buf_size*/ 4096); + let client = serve_directly(service, client_transport, /*peer_info*/ None); + let mut server = IntoTransport::::into_transport(server_transport); + server + .send(ServerJsonRpcMessage::request( + ServerRequest::CustomRequest(CustomRequest::new( + "openai/elicitation/create", + Some(json!({ + "mode": "openai/userVerification", + "title": "Approve", + "description": "", + "challenge": "AQID", + })), + )), + RequestId::Number(1), + )) + .await?; + let mut response_tx = timeout(Duration::from_secs(/*secs*/ 5), route_rx.recv()) + .await? + .expect("verification was routed to the UI"); + assert!(*paused.borrow()); + + client.cancel().await?; + + timeout(Duration::from_secs(/*secs*/ 5), response_tx.closed()).await?; + timeout( + Duration::from_secs(/*secs*/ 5), + paused.wait_for(|paused| !*paused), + ) + .await??; + Ok(()) +} + +#[tokio::test] +async fn cancelling_one_verification_leaves_the_mcp_connection_and_other_requests_alive() +-> anyhow::Result<()> { + let pause_state = ElicitationPauseState::new(); + let mut paused = pause_state.subscribe(); + let (route_tx, mut route_rx) = mpsc::unbounded_channel(); + let mut client_info = ClientInfo::default(); + client_info.capabilities.extensions = Some( + [( + OPENAI_ELICITATION_EXTENSION_ID.to_string(), + serde_json::Map::from_iter([("userVerification".to_string(), json!({}))]), + )] + .into_iter() + .collect(), + ); + let service = ElicitationClientService::new( + client_info, + Box::new(move |id, _| { + let (response_tx, response_rx) = oneshot::channel(); + route_tx + .send((id, response_tx)) + .expect("observe verification"); + Box::pin(async move { Ok(response_rx.await?) }) + }), + pause_state, + ); + let (client_transport, server_transport) = tokio::io::duplex(/*max_buf_size*/ 4096); + let client = serve_directly(service, client_transport, /*peer_info*/ None); + let mut server = IntoTransport::::into_transport(server_transport); + + for id in [1, 2] { + server + .send(ServerJsonRpcMessage::request( + ServerRequest::CustomRequest(CustomRequest::new( + "openai/elicitation/create", + Some(json!({ + "mode": "openai/userVerification", + "title": "Approve", + "description": "", + "challenge": "AQID", + })), + )), + RequestId::Number(id), + )) + .await?; + } + let (first_id, mut first) = timeout(Duration::from_secs(/*secs*/ 5), route_rx.recv()) + .await? + .expect("first verification reached UI"); + let (second_id, second) = timeout(Duration::from_secs(/*secs*/ 5), route_rx.recv()) + .await? + .expect("second verification reached UI"); + assert_ne!(first_id, second_id); + assert!(*paused.borrow()); + + server + .send(ServerJsonRpcMessage::notification( + ServerNotification::CancelledNotification(CancelledNotification::new( + CancelledNotificationParam::new(Some(first_id.clone()), None), + )), + )) + .await?; + timeout(Duration::from_secs(/*secs*/ 5), first.closed()).await?; + assert!(!second.is_closed()); + assert!(*paused.borrow()); + let cancel = timeout(Duration::from_secs(/*secs*/ 5), server.receive()) + .await? + .expect("cancelled verification still sends a response"); + let ClientJsonRpcMessage::Response(cancel) = cancel else { + anyhow::bail!("expected a cancellation response"); + }; + assert_eq!(cancel.id, first_id); + assert_eq!(serde_json::to_value(cancel.result)?["action"], "cancel"); + + for request_id in [first_id, RequestId::Number(999)] { + server + .send(ServerJsonRpcMessage::notification( + ServerNotification::CancelledNotification(CancelledNotification::new( + CancelledNotificationParam::new(Some(request_id), None), + )), + )) + .await?; + } + assert!(!second.is_closed()); + + second + .send(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({"credentialId": "AQID", "signature": "BAUG"})), + meta: None, + }) + .expect("second verification remains pending"); + let accepted = timeout(Duration::from_secs(/*secs*/ 5), server.receive()) + .await? + .expect("second request survived cancellation"); + let ClientJsonRpcMessage::Response(accepted) = accepted else { + anyhow::bail!("expected the second verification response"); + }; + assert_eq!(accepted.id, second_id); + assert_eq!(serde_json::to_value(accepted.result)?["action"], "accept"); + timeout( + Duration::from_secs(/*secs*/ 5), + paused.wait_for(|paused| !*paused), + ) + .await??; + client.cancel().await?; + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/user_verification_dispatch_tests.rs b/codex-rs/rmcp-client/src/user_verification_dispatch_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..ea1ac316106fa832111ecbdddee6a66bdb936d5b --- /dev/null +++ b/codex-rs/rmcp-client/src/user_verification_dispatch_tests.rs @@ -0,0 +1,164 @@ +use super::*; +use pretty_assertions::assert_eq; +use rmcp::RoleServer; +use rmcp::model::ClientJsonRpcMessage; +use rmcp::model::ServerJsonRpcMessage; +use rmcp::service::serve_directly; +use rmcp::transport::IntoTransport; +use rmcp::transport::Transport; +use serde_json::json; +use std::time::Duration; +use tokio::time::timeout; + +fn service() -> ElicitationClientService { + let mut info = ClientInfo::default(); + info.capabilities.extensions = Some( + [( + OPENAI_ELICITATION_EXTENSION_ID.into(), + Map::from_iter([("userVerification".into(), json!({}))]), + )] + .into_iter() + .collect(), + ); + ElicitationClientService::new( + info, + Box::new(|_, _| panic!("cancelled or malformed verification must not reach the UI")), + ElicitationPauseState::new(), + ) +} + +#[tokio::test] +async fn user_verification_remembers_cancellation_before_request_handler_runs() -> anyhow::Result<()> +{ + let service = service(); + let (client_transport, server_transport) = tokio::io::duplex(/*max_buf_size*/ 4096); + let client = serve_directly(service.clone(), client_transport, /*peer_info*/ None); + let mut server = IntoTransport::::into_transport(server_transport); + server + .send(serde_json::from_value::(json!({ + "jsonrpc": "2.0", "method": "notifications/cancelled", "params": {"requestId": 1} + }))?) + .await?; + // Force the scheduling order that RMCP's independently spawned handlers allow. + timeout(Duration::from_secs(/*secs*/ 5), async { + loop { + if service + .pending_verifications + .lock() + .unwrap() + .early + .contains(&RequestId::Number(1)) + { + break; + } + tokio::task::yield_now().await; + } + }) + .await?; + server.send(serde_json::from_value::(json!({ + "jsonrpc": "2.0", "id": 1, "method": OPENAI_ELICITATION_METHOD, + "params": {"mode": crate::user_verification::MODE, "title": "Approve", "description": "", "challenge": "AQID"} + }))?).await?; + let response = timeout(Duration::from_secs(/*secs*/ 5), server.receive()) + .await? + .unwrap(); + assert_eq!( + serde_json::to_value(response)?, + json!({ + "jsonrpc": "2.0", "id": 1, "result": {"action": "cancel"} + }) + ); + assert!( + service + .pending_verifications + .lock() + .unwrap() + .early + .is_empty() + ); + client.cancel().await?; + Ok(()) +} + +#[tokio::test] +async fn user_verification_early_cancellation_storage_fails_closed_at_capacity() +-> anyhow::Result<()> { + let service = service(); + let (client_transport, server_transport) = tokio::io::duplex(/*max_buf_size*/ 4096); + let client = serve_directly(service.clone(), client_transport, /*peer_info*/ None); + let mut server = IntoTransport::::into_transport(server_transport); + for id in 0..=MAX_EARLY_CANCELLATIONS { + server + .send(serde_json::from_value::(json!({ + "jsonrpc": "2.0", "method": "notifications/cancelled", "params": {"requestId": id} + }))?) + .await?; + } + timeout(Duration::from_secs(/*secs*/ 5), async { + while !service.pending_verifications.lock().unwrap().saturated { + tokio::task::yield_now().await; + } + }) + .await?; + assert!( + service + .pending_verifications + .lock() + .unwrap() + .early + .is_empty() + ); + let response = service.handle_request( + ServerRequest::CustomRequest(CustomRequest::new(OPENAI_ELICITATION_METHOD, Some(json!({ + "mode": crate::user_verification::MODE, "title": "Approve", "description": "", "challenge": "AQID" + })))), + RequestContext::new(RequestId::Number(5000), client.peer().clone()), + ).await?; + assert_eq!(serde_json::to_value(response)?, json!({"action": "cancel"})); + // Ordinary requests still work after verification fails closed. + server + .send(serde_json::from_value::(json!({ + "jsonrpc": "2.0", "id": 5001, "method": "ping" + }))?) + .await?; + let response = timeout(Duration::from_secs(/*secs*/ 5), server.receive()) + .await? + .unwrap(); + assert_eq!( + serde_json::to_value(response)?, + json!({"jsonrpc": "2.0", "id": 5001, "result": {}}) + ); + client.cancel().await?; + Ok(()) +} + +#[tokio::test] +async fn user_verification_dispatch_rejects_malformed_modes_as_invalid_params() -> anyhow::Result<()> +{ + let (client_transport, server_transport) = tokio::io::duplex(/*max_buf_size*/ 4096); + let client = serve_directly(service(), client_transport, /*peer_info*/ None); + let mut server = IntoTransport::::into_transport(server_transport); + for params in [json!({}), json!({"mode": 7}), json!({"mode": "unknown"})] { + server + .send(ServerJsonRpcMessage::request( + ServerRequest::CustomRequest(CustomRequest::new( + OPENAI_ELICITATION_METHOD, + Some(params), + )), + RequestId::Number(1), + )) + .await?; + let response = timeout(Duration::from_secs(/*secs*/ 5), server.receive()) + .await? + .unwrap(); + let ClientJsonRpcMessage::Error(response) = response else { + anyhow::bail!("expected invalid params"); + }; + assert_eq!( + response.error, + rmcp::ErrorData::invalid_params("invalid elicitation mode", /*data*/ None) + ); + } + client.cancel().await?; + Ok(()) +} diff --git a/codex-rs/rmcp-client/src/user_verification_tests.rs b/codex-rs/rmcp-client/src/user_verification_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..1465f041dc5ad7ddfeb0db4bb3dc7916bb1b3b4a --- /dev/null +++ b/codex-rs/rmcp-client/src/user_verification_tests.rs @@ -0,0 +1,112 @@ +use super::*; +use pretty_assertions::assert_eq; +use serde_json::json; + +#[test] +fn user_verification_request_preserves_signed_bytes_and_display_text() { + assert_eq!( + parse_request(CustomRequest::new( + "openai/elicitation/create", + Some(json!({ + "mode": MODE, + "title": "Approve purchase", + "description": "Pay $200 to Example Store", + "challenge": "AQID", + })), + )) + .unwrap(), + Elicitation::UserVerification { + title: "Approve purchase".to_string(), + description: "Pay $200 to Example Store".to_string(), + challenge: "AQID".to_string(), + } + ); +} + +#[test] +fn user_verification_rejects_invalid_or_unbounded_requests_without_echoing_values() { + for (field, value) in [ + ("mode", json!("unknown")), + ("title", json!("")), + ("title", json!("x".repeat(257))), + ("description", json!("x".repeat(4097))), + ("challenge", json!("")), + ("challenge", json!("not base64url!")), + ("challenge", json!(URL_SAFE_NO_PAD.encode(vec![0; 4097]))), + ("unexpected", json!("secret")), + ] { + let mut params = json!({ + "mode": MODE, "title": "Approve", "description": "", "challenge": "AQID" + }); + params[field] = value; + assert_eq!( + parse_request(CustomRequest::new( + "openai/elicitation/create", + Some(params) + )), + Err(invalid_request()), + ); + } +} + +#[test] +fn user_verification_acceptance_returns_proof_only_in_content() { + let content = Some(json!({"credentialId": "credential:用户/1=", "signature": "BAUG"})); + assert_eq!( + validate_response(ElicitationResponse { + action: ElicitationAction::Accept, + content: content.clone(), + meta: Some(json!({"untrusted": "ignored"})), + }), + ElicitationResponse { + action: ElicitationAction::Accept, + content, + meta: None, + } + ); +} + +#[test] +fn user_verification_cancels_acceptance_without_a_valid_bounded_proof() { + for content in [ + None, + Some(json!({})), + Some(json!({"credentialId": "AQID", "signature": ""})), + Some(json!({"credentialId": "AQID", "signature": "BAUG="})), + Some(json!({"credentialId": "é".repeat(513), "signature": "BAUG"})), + Some(json!({"credentialId": "", "signature": "BAUG"})), + Some(json!({"credentialId": "AQID", "signature": URL_SAFE_NO_PAD.encode(vec![0; 129])})), + Some(json!({"credentialId": "AQID", "signature": "BAUG", "extra": true})), + ] { + assert_eq!( + validate_response(ElicitationResponse { + action: ElicitationAction::Accept, + content, + meta: None, + }), + ElicitationResponse { + action: ElicitationAction::Cancel, + content: None, + meta: None, + } + ); + } +} + +#[test] +fn user_verification_decline_and_cancel_discard_proof_material() { + for action in [ElicitationAction::Decline, ElicitationAction::Cancel] { + assert_eq!( + validate_response(ElicitationResponse { + action: action.clone(), + content: Some(json!({"credentialId": "AQID", "signature": "BAUG"})), + meta: Some(json!({"secret": "discarded"})), + }), + ElicitationResponse { + action, + content: None, + meta: None + } + ); + } +} diff --git a/codex-rs/rmcp-client/src/utils.rs b/codex-rs/rmcp-client/src/utils.rs new file mode 100644 index 0000000000000000000000000000000000000000..dac42ca6118f9b1f620030645447407532582383 --- /dev/null +++ b/codex-rs/rmcp-client/src/utils.rs @@ -0,0 +1,357 @@ +use anyhow::Result; +use anyhow::anyhow; +use codex_config::types::McpServerEnvVar; +use codex_network_proxy::CUSTOM_CA_ENV_KEYS; +use codex_protocol::shell_environment::is_non_inheritable_env_var; +use http::HeaderMap; +use http::HeaderName; +use http::HeaderValue; +use http::header::USER_AGENT; +use std::collections::HashMap; +use std::env; +use std::ffi::OsString; + +pub(crate) const MCP_USER_AGENT: &str = concat!("codex-mcp-client/", env!("CARGO_PKG_VERSION")); + +pub(crate) fn create_env_for_mcp_server( + extra_env: Option>, + env_vars: &[McpServerEnvVar], +) -> Result> { + let additional_env_vars = local_stdio_env_var_names(env_vars)?; + let mut env: HashMap = DEFAULT_ENV_VARS + .iter() + .copied() + .chain(additional_env_vars) + .filter_map(|var| env::var_os(var).map(|value| (OsString::from(var), value))) + .collect(); + for name in CUSTOM_CA_ENV_KEYS { + let Some(value) = env::var_os(name) else { + continue; + }; + if value.is_empty() { + continue; + } + let value = std::path::absolute(value)?.into_os_string(); + #[cfg(windows)] + env.retain(|key, _| !key.to_string_lossy().eq_ignore_ascii_case(name)); + env.insert(OsString::from(name), value); + } + for (name, value) in extra_env.unwrap_or_default() { + if cfg!(windows) + || name.to_str().is_some_and(|name| { + CUSTOM_CA_ENV_KEYS + .iter() + .any(|ca_name| ca_name.eq_ignore_ascii_case(name)) + }) + { + env.retain(|key, _| { + !key.to_string_lossy() + .eq_ignore_ascii_case(&name.to_string_lossy()) + }); + } + env.insert(name, value); + } + env.retain(|name, _| { + name.to_str() + .is_none_or(|name| !is_non_inheritable_env_var(name)) + }); + Ok(env) +} + +pub(crate) fn create_env_overlay_for_remote_mcp_server( + extra_env: Option>, + env_vars: &[McpServerEnvVar], +) -> HashMap { + // Remote stdio should inherit PATH/HOME/etc. from the executor side, not + // from the orchestrator process. Only forward variables explicitly named + // by the MCP config plus literal env overrides from that config. + let mut env: HashMap = env_vars + .iter() + .filter(|var| !var.is_remote_source()) + .filter_map(|var| env::var_os(var.name()).map(|value| (OsString::from(var.name()), value))) + .chain(extra_env.unwrap_or_default()) + .collect(); + env.retain(|name, _| { + name.to_str() + .is_none_or(|name| !is_non_inheritable_env_var(name)) + }); + env +} + +pub(crate) fn remote_mcp_env_var_names(env_vars: &[McpServerEnvVar]) -> Vec { + env_vars + .iter() + .filter(|var| var.is_remote_source()) + .filter(|var| !is_non_inheritable_env_var(var.name())) + .map(|var| var.name().to_string()) + .collect() +} + +fn local_stdio_env_var_names(env_vars: &[McpServerEnvVar]) -> Result> { + if let Some(remote_var) = env_vars.iter().find(|var| var.is_remote_source()) { + return Err(anyhow!( + "env_vars entry `{}` uses source `remote`, which requires remote MCP stdio", + remote_var.name() + )); + } + Ok(env_vars + .iter() + .map(McpServerEnvVar::name) + .filter(|name| !is_non_inheritable_env_var(name))) +} + +pub(crate) fn build_default_headers( + http_headers: Option>, + env_http_headers: Option>, +) -> Result { + let mut headers = HeaderMap::new(); + headers.insert(USER_AGENT, HeaderValue::from_static(MCP_USER_AGENT)); + + if let Some(static_headers) = http_headers { + for (name, value) in static_headers { + let header_name = match HeaderName::from_bytes(name.as_bytes()) { + Ok(name) => name, + Err(err) => { + tracing::warn!("invalid HTTP header name `{name}`: {err}"); + continue; + } + }; + let header_value = match HeaderValue::from_str(value.as_str()) { + Ok(value) => value, + Err(err) => { + tracing::warn!("invalid HTTP header value for `{name}`: {err}"); + continue; + } + }; + headers.insert(header_name, header_value); + } + } + + if let Some(env_headers) = env_http_headers { + for (name, env_var) in env_headers { + if let Ok(value) = env::var(&env_var) { + if value.trim().is_empty() { + continue; + } + + let header_name = match HeaderName::from_bytes(name.as_bytes()) { + Ok(name) => name, + Err(err) => { + tracing::warn!("invalid HTTP header name `{name}`: {err}"); + continue; + } + }; + + let header_value = match HeaderValue::from_str(value.as_str()) { + Ok(value) => value, + Err(err) => { + tracing::warn!( + "invalid HTTP header value read from {env_var} for `{name}`: {err}" + ); + continue; + } + }; + headers.insert(header_name, header_value); + } + } + } + + Ok(headers) +} + +#[cfg(unix)] +pub(crate) const DEFAULT_ENV_VARS: &[&str] = &[ + "HOME", + "LOGNAME", + "PATH", + "SHELL", + "USER", + "__CF_USER_TEXT_ENCODING", + "LANG", + "LC_ALL", + "TERM", + "TMPDIR", + "TZ", +]; + +#[cfg(windows)] +pub(crate) const DEFAULT_ENV_VARS: &[&str] = + codex_protocol::shell_environment::WINDOWS_CORE_ENV_VARS; + +#[cfg(test)] +mod tests { + use super::*; + use pretty_assertions::assert_eq; + + use serial_test::serial; + use std::ffi::OsStr; + + struct EnvVarGuard { + key: String, + original: Option, + } + + impl EnvVarGuard { + fn set(key: &str, value: impl AsRef) -> Self { + let original = std::env::var_os(key); + unsafe { + std::env::set_var(key, value.as_ref()); + } + Self { + key: key.to_string(), + original, + } + } + } + + impl Drop for EnvVarGuard { + fn drop(&mut self) { + if let Some(value) = &self.original { + unsafe { + std::env::set_var(&self.key, value); + } + } else { + unsafe { + std::env::remove_var(&self.key); + } + } + } + } + + #[tokio::test] + async fn create_env_honors_overrides() { + let value = "custom".to_string(); + let expected = OsString::from(&value); + let env = create_env_for_mcp_server( + Some(HashMap::from([ + (OsString::from("TZ"), expected.clone()), + ( + OsString::from("openai_identity_token_file"), + OsString::from("/run/identity-token"), + ), + ])), + &[], + ) + .expect("local MCP env should build"); + assert_eq!(env.get(OsStr::new("TZ")), Some(&expected)); + assert!(!env.contains_key(OsStr::new("openai_identity_token_file"))); + } + + #[test] + #[serial(extra_rmcp_env)] + fn create_env_includes_additional_whitelisted_variables() { + let custom_var = "EXTRA_RMCP_ENV"; + let value = "from-env"; + let expected = OsString::from(value); + let _guard = EnvVarGuard::set(custom_var, value); + let env = create_env_for_mcp_server(/*extra_env*/ None, &[custom_var.into()]) + .expect("local MCP env should build"); + assert_eq!(env.get(OsStr::new(custom_var)), Some(&expected)); + } + + #[test] + #[serial(extra_rmcp_env)] + fn create_remote_env_overlay_only_forwards_explicit_variables() { + let default_var = DEFAULT_ENV_VARS[0]; + let custom_var = "EXTRA_REMOTE_RMCP_ENV"; + let custom_value = OsString::from("from-env"); + let _default_guard = EnvVarGuard::set(default_var, "from-default"); + let _custom_guard = EnvVarGuard::set(custom_var, &custom_value); + + let env = create_env_overlay_for_remote_mcp_server( + Some(HashMap::from([( + OsString::from("OpenAI_Federation_Rule_Id"), + OsString::from("rule"), + )])), + &[custom_var.into()], + ); + + assert_eq!( + env, + HashMap::from([(OsString::from(custom_var), custom_value)]) + ); + } + + #[test] + #[serial(extra_rmcp_env)] + fn create_remote_env_overlay_does_not_copy_remote_source_variables() { + let remote_var = "REMOTE_ONLY_RMCP_ENV"; + let local_var = "LOCAL_RMCP_ENV"; + let local_value = OsString::from("from-local-env"); + let _remote_guard = EnvVarGuard::set(remote_var, "should-not-be-copied"); + let _local_guard = EnvVarGuard::set(local_var, &local_value); + + let env = create_env_overlay_for_remote_mcp_server( + /*extra_env*/ None, + &[ + McpServerEnvVar::Config { + name: remote_var.to_string(), + source: Some("remote".to_string()), + }, + McpServerEnvVar::Config { + name: local_var.to_string(), + source: Some("local".to_string()), + }, + ], + ); + + assert_eq!( + env, + HashMap::from([(OsString::from(local_var), local_value)]) + ); + } + + #[test] + fn remote_mcp_env_var_names_returns_remote_source_names() { + let names = remote_mcp_env_var_names(&[ + "LEGACY".into(), + McpServerEnvVar::Config { + name: "LOCAL".to_string(), + source: Some("local".to_string()), + }, + McpServerEnvVar::Config { + name: "REMOTE".to_string(), + source: Some("remote".to_string()), + }, + McpServerEnvVar::Config { + name: "openai_identity_token_file".to_string(), + source: Some("remote".to_string()), + }, + ]); + + assert_eq!(names, vec!["REMOTE".to_string()]); + } + + #[test] + fn create_local_env_rejects_remote_source_variables() { + let err = create_env_for_mcp_server( + /*extra_env*/ None, + &[McpServerEnvVar::Config { + name: "REMOTE".to_string(), + source: Some("remote".to_string()), + }], + ) + .expect_err("remote source should require remote stdio"); + + assert!( + err.to_string().contains("requires remote MCP stdio"), + "unexpected error: {err}" + ); + } + + #[cfg(unix)] + #[test] + #[serial(extra_rmcp_env)] + fn create_env_preserves_path_when_it_is_not_utf8() { + use std::os::unix::ffi::OsStrExt; + + let raw_path = std::ffi::OsStr::from_bytes(b"/tmp/codex-\xFF/bin"); + let expected = raw_path.to_os_string(); + let _guard = EnvVarGuard::set("PATH", raw_path); + + let env = + create_env_for_mcp_server(/*extra_env*/ None, &[]).expect("local MCP env should build"); + + assert_eq!(env.get(OsStr::new("PATH")), Some(&expected)); + } +} diff --git a/codex-rs/rmcp-client/src/www_authenticate.rs b/codex-rs/rmcp-client/src/www_authenticate.rs new file mode 100644 index 0000000000000000000000000000000000000000..c4382c8b1b29dc094f9c5c3399265a8e2b894529 --- /dev/null +++ b/codex-rs/rmcp-client/src/www_authenticate.rs @@ -0,0 +1,233 @@ +use codex_exec_server::HttpHeader; +use http::header::WWW_AUTHENTICATE; + +#[derive(Debug, PartialEq, Eq)] +pub(crate) struct InsufficientScopeChallenge { + pub(crate) www_authenticate_header: String, + pub(crate) required_scope: Option, +} + +#[derive(Debug, PartialEq, Eq)] +struct BearerInsufficientScope { + required_scope: Option, +} + +type AuthParameter<'a> = (&'a str, Option); +type ChallengeStart<'a> = (&'a str, Option>); + +#[derive(Default)] +enum Parameter { + #[default] + Missing, + Value(String), + Invalid, +} + +#[derive(Default)] +struct BearerChallenge { + error: Parameter, + scope: Parameter, +} + +impl BearerChallenge { + fn add_parameter(&mut self, name: &str, value: Option) { + let parameter = if name.eq_ignore_ascii_case("error") { + &mut self.error + } else if name.eq_ignore_ascii_case("scope") { + &mut self.scope + } else { + return; + }; + + *parameter = match (&*parameter, value) { + (Parameter::Missing, Some(value)) => Parameter::Value(value), + (Parameter::Missing, None) | (Parameter::Value(_), _) | (Parameter::Invalid, _) => { + Parameter::Invalid + } + }; + } + + fn into_insufficient_scope(self) -> Option { + match self.error { + Parameter::Value(error) if error == "insufficient_scope" => { + Some(BearerInsufficientScope { + required_scope: match self.scope { + Parameter::Value(scope) if valid_scope(&scope) => Some(scope), + Parameter::Missing | Parameter::Value(_) | Parameter::Invalid => None, + }, + }) + } + Parameter::Missing | Parameter::Value(_) | Parameter::Invalid => None, + } + } +} + +/// Finds a Bearer insufficient-scope challenge among all `WWW-Authenticate` +/// response header field values. +pub(crate) fn insufficient_scope_challenge( + headers: &[HttpHeader], +) -> Option { + headers + .iter() + .filter(|header| header.name.eq_ignore_ascii_case(WWW_AUTHENTICATE.as_str())) + .find_map(|header| { + parse_bearer_insufficient_scope(&header.value).map(|challenge| { + InsufficientScopeChallenge { + www_authenticate_header: header.value.clone(), + required_scope: challenge.required_scope, + } + }) + }) +} + +/// Parses a Bearer `WWW-Authenticate` challenge with an `insufficient_scope` +/// error and extracts its optional required scope. +/// +/// RFC 9110 section 11.2 defines challenge parameters as `auth-param` values +/// whose values are either `token` or `quoted-string`. Quoted strings use HTTP +/// syntax rather than JSON: section 5.6.4 requires recipients to replace each +/// `quoted-pair` with its escaped octet. +/// +/// RFC 6750 section 3 permits `scope` in the Bearer challenge at most once. +/// After HTTP quoted-string processing, each scope token can contain `%x21`, +/// `%x23-5B`, or `%x5D-7E`, with `%x20` separating multiple tokens. Therefore +/// returned scopes cannot contain `"` or `\`, even when those characters occur +/// in the header encoding. +/// +/// RMCP has related parsing logic, but it is private to that crate. +fn parse_bearer_insufficient_scope(header: &str) -> Option { + let segments = split_unquoted_segments(header)?; + let mut bearer_challenge: Option = None; + + for segment in segments { + if let Some((name, value)) = parse_auth_param(segment) { + if let Some(challenge) = bearer_challenge.as_mut() { + challenge.add_parameter(name, value); + } + continue; + } + + if let Some(challenge) = bearer_challenge + .take() + .and_then(BearerChallenge::into_insufficient_scope) + { + return Some(challenge); + } + + let (scheme, parameter) = parse_challenge_start(segment)?; + if scheme.eq_ignore_ascii_case("Bearer") { + let mut challenge = BearerChallenge::default(); + if let Some((name, value)) = parameter { + challenge.add_parameter(name, value); + } + bearer_challenge = Some(challenge); + } + } + + bearer_challenge.and_then(BearerChallenge::into_insufficient_scope) +} + +fn parse_challenge_start(segment: &str) -> Option> { + let segment = segment.trim(); + let parameter_start = segment.find(char::is_whitespace); + let (scheme, parameter) = match parameter_start { + Some(parameter_start) => ( + &segment[..parameter_start], + parse_auth_param(&segment[parameter_start..]), + ), + None => (segment, None), + }; + + is_http_token(scheme).then_some((scheme, parameter)) +} + +fn parse_auth_param(segment: &str) -> Option> { + let (name, value) = segment.trim().split_once('=')?; + let name = name.trim(); + is_http_token(name).then_some((name, parse_auth_param_value(value.trim()))) +} + +fn parse_auth_param_value(value: &str) -> Option { + if let Some(quoted_value) = value.strip_prefix('"') { + let quoted_value = quoted_value.strip_suffix('"')?; + let mut decoded = String::with_capacity(quoted_value.len()); + let mut characters = quoted_value.chars(); + while let Some(character) = characters.next() { + if character == '\\' { + decoded.push(characters.next()?); + } else { + decoded.push(character); + } + } + Some(decoded) + } else { + is_http_token(value).then(|| value.to_string()) + } +} + +fn split_unquoted_segments(header: &str) -> Option> { + let mut segments = Vec::new(); + let mut segment_start = 0; + let mut in_quotes = false; + let mut escaped = false; + + for (position, character) in header.char_indices() { + if escaped { + escaped = false; + continue; + } + match character { + '\\' if in_quotes => escaped = true, + '"' => in_quotes = !in_quotes, + ',' | ';' if !in_quotes => { + segments.push(&header[segment_start..position]); + segment_start = position + character.len_utf8(); + } + _ => {} + } + } + + if in_quotes || escaped { + None + } else { + segments.push(&header[segment_start..]); + Some(segments) + } +} + +fn valid_scope(scope: &str) -> bool { + scope.split(' ').all(|token| { + !token.is_empty() + && token + .bytes() + .all(|byte| matches!(byte, b'!' | b'#'..=b'[' | b']'..=b'~')) + }) +} + +fn is_http_token(value: &str) -> bool { + !value.is_empty() + && value.bytes().all(|byte| { + byte.is_ascii_alphanumeric() + || matches!( + byte, + b'!' | b'#' + | b'$' + | b'%' + | b'&' + | b'\'' + | b'*' + | b'+' + | b'-' + | b'.' + | b'^' + | b'_' + | b'`' + | b'|' + | b'~' + ) + }) +} + +#[cfg(test)] +#[path = "www_authenticate_tests.rs"] +mod tests; diff --git a/codex-rs/rmcp-client/src/www_authenticate_tests.rs b/codex-rs/rmcp-client/src/www_authenticate_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..7ecfae9560d5221ccadada742a4f2412320a6369 --- /dev/null +++ b/codex-rs/rmcp-client/src/www_authenticate_tests.rs @@ -0,0 +1,126 @@ +use codex_exec_server::HttpHeader; +use pretty_assertions::assert_eq; + +use super::BearerInsufficientScope; +use super::InsufficientScopeChallenge; +use super::insufficient_scope_challenge; +use super::parse_bearer_insufficient_scope; + +#[test] +fn extracts_scope_from_bearer_insufficient_scope_challenges() { + let cases = [ + ( + r#"Bearer error="insufficient_scope", scope="files:read files:write""#, + "files:read files:write", + ), + ( + r#"Bearer error="insufficient_scope", ScOpE = "files:read""#, + "files:read", + ), + ( + r#"Bearer scope="read:data", error="insufficient_scope""#, + "read:data", + ), + (r#"Bearer error="insufficient_scope", scope=read"#, "read"), + ( + r#"Bearer error="insufficient_scope", scope="files:read\ files:write""#, + "files:read files:write", + ), + ( + r#"Bearer error="insufficient_scope", error_description="request scope=admin, not \"root\"", scope="files:read""#, + "files:read", + ), + ( + r#"Basic realm="example", Bearer error="insufficient_scope", scope="files:read""#, + "files:read", + ), + ( + r#"Newauth scope="wrong", Bearer error="insufficient_scope", scope="files:read""#, + "files:read", + ), + ]; + + for (header, expected_scope) in cases { + assert_eq!( + parse_bearer_insufficient_scope(header), + Some(BearerInsufficientScope { + required_scope: Some(expected_scope.to_string()), + }), + "header: {header}" + ); + } +} + +#[test] +fn does_not_treat_other_bearer_errors_as_insufficient_scope() { + assert_eq!( + parse_bearer_insufficient_scope(r#"Bearer error="invalid_token", scope="files:read""#), + None + ); +} + +#[test] +fn rejects_invalid_or_ambiguous_scope_parameters() { + let cases = [ + r#"Bearer error="insufficient_scope", scope="#, + r#"Bearer error="insufficient_scope", scope="read\"write""#, + r#"Bearer error="insufficient_scope", scope="read\\write""#, + r#"Bearer error="insufficient_scope", scope="read write""#, + r#"Bearer error="insufficient_scope", scope=read:data"#, + r#"Bearer error="insufficient_scope", scope=files:read files:write"#, + r#"Bearer error="insufficient_scope", scope=read=value"#, + r#"Bearer error="insufficient_scope", scope="read", scope="write""#, + ]; + + for header in cases { + assert_eq!( + parse_bearer_insufficient_scope(header), + Some(BearerInsufficientScope { + required_scope: None, + }), + "header: {header}" + ); + } +} + +#[test] +fn ignores_scope_text_outside_a_scope_parameter() { + let cases = [ + r#"Bearer error_description="request scope=admin""#, + r#"Bearer resource_scope="admin""#, + r#"Bearer "scope=admin""#, + r#"Bearer error_description="unterminated scope=admin"#, + ]; + + for header in cases { + assert_eq!( + parse_bearer_insufficient_scope(header), + None, + "header: {header}" + ); + } +} + +#[test] +fn selects_bearer_challenge_from_a_later_www_authenticate_field_value() { + let headers = vec![ + HttpHeader { + name: "www-authenticate".to_string(), + value: r#"Basic realm="example""#.to_string(), + value_env_var: None, + }, + HttpHeader { + name: "WWW-Authenticate".to_string(), + value: r#"Bearer error="insufficient_scope", scope="files:read""#.to_string(), + value_env_var: None, + }, + ]; + + assert_eq!( + insufficient_scope_challenge(&headers), + Some(InsufficientScopeChallenge { + www_authenticate_header: headers[1].value.clone(), + required_scope: Some("files:read".to_string()), + }) + ); +} diff --git a/codex-rs/rmcp-client/tests/foreign_stdio_cwd.rs b/codex-rs/rmcp-client/tests/foreign_stdio_cwd.rs new file mode 100644 index 0000000000000000000000000000000000000000..84aca7b92a5e52f55075d9ea879ad9cbe67cb33a --- /dev/null +++ b/codex-rs/rmcp-client/tests/foreign_stdio_cwd.rs @@ -0,0 +1,67 @@ +use std::ffi::OsString; +use std::sync::Arc; +use std::sync::Mutex; + +use codex_exec_server::ExecBackend; +use codex_exec_server::ExecBackendFuture; +use codex_exec_server::ExecParams; +use codex_exec_server::ExecServerError; +use codex_rmcp_client::ExecutorStdioServerLauncher; +use codex_rmcp_client::RmcpClient; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; + +#[derive(Default)] +struct RecordingExecBackend { + params: Mutex>, +} + +impl ExecBackend for RecordingExecBackend { + fn start(&self, params: ExecParams) -> ExecBackendFuture<'_> { + let mut recorded_params = match self.params.lock() { + Ok(recorded_params) => recorded_params, + Err(poisoned) => poisoned.into_inner(), + }; + *recorded_params = Some(params); + Box::pin(async { + Err(ExecServerError::Protocol( + "stop after recording executor request".to_string(), + )) + }) + } +} + +#[tokio::test] +async fn executor_stdio_forwards_foreign_absolute_cwd_as_path_uri() { + #[cfg(not(windows))] + let cwd = r"C:\Users\openai\share"; + #[cfg(windows)] + let cwd = "/home/openai/share"; + #[cfg(not(windows))] + let expected_cwd: PathUri = "file:///C:/Users/openai/share" + .parse() + .expect("expected cwd should be a path URI"); + #[cfg(windows)] + let expected_cwd: PathUri = "file:///home/openai/share" + .parse() + .expect("expected cwd should be a path URI"); + let backend = Arc::new(RecordingExecBackend::default()); + let launcher = Arc::new(ExecutorStdioServerLauncher::new(backend.clone())); + + let _ = RmcpClient::new_stdio_client( + OsString::from("echo"), + Vec::new(), + /*env*/ None, + &[], + Some(cwd.to_string()), + launcher, + ) + .await; + let params = backend + .params + .lock() + .expect("recorded params lock should not be poisoned") + .take() + .expect("executor start request should be recorded"); + assert_eq!(params.cwd, expected_cwd); +} diff --git a/codex-rs/rmcp-client/tests/mcp_2026_discovery.rs b/codex-rs/rmcp-client/tests/mcp_2026_discovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..1495438d75efb9989cba066a04c014ddbae20632 --- /dev/null +++ b/codex-rs/rmcp-client/tests/mcp_2026_discovery.rs @@ -0,0 +1,1479 @@ +use std::collections::HashMap; +use std::ffi::OsString; +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; + +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::LocalStdioServerLauncher; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use rmcp::model::ClientCapabilities; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use rmcp::model::ServerCapabilities; +use rmcp::model::ServerPeerInfo; +use serde_json::Value; +use serde_json::json; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::Request; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +const MODERN_VERSION: &str = "2026-07-28"; +const LEGACY_VERSION: &str = "2025-06-18"; +const MAX_MCP_MESSAGE_BYTES: usize = 8 * 1024 * 1024; + +fn initialize_params() -> InitializeRequestParams { + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("codex-discovery-test", "0.0.0"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18) +} + +async fn create_client(server: &MockServer, mode: McpProtocolMode) -> anyhow::Result { + RmcpClient::new_streamable_http_client_with_protocol_mode( + "discovery-test", + &format!("{}/mcp", server.uri()), + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + mode, + ) + .await +} + +async fn initialize_client(client: &RmcpClient) -> anyhow::Result { + client + .initialize( + initialize_params(), + Some(Duration::from_secs(5)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: None, + }) + } + .boxed() + }), + ) + .await +} + +fn modern_discover_response(request: &Value) -> ResponseTemplate { + ResponseTemplate::new(200).set_body_json(modern_discover_result(request)) +} + +fn modern_discover_result(request: &Value) -> Value { + json!({ + "jsonrpc": "2.0", + "id": request["id"], + "result": { + "resultType": "complete", + "supportedVersions": [MODERN_VERSION], + "capabilities": {"tools": {}}, + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "modern-test", + "version": "1.0.0", + }, + }, + "ttlMs": 0, + "cacheScope": "private", + }, + }) +} + +fn legacy_initialize_response(request: &Value) -> ResponseTemplate { + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": request["id"], + "result": { + "protocolVersion": LEGACY_VERSION, + "capabilities": {"tools": {}}, + "serverInfo": {"name": "legacy-test", "version": "1.0.0"}, + }, + })) +} + +#[tokio::test] +async fn modern_mode_uses_sdk_discovery_and_self_contained_request_metadata() -> anyhow::Result<()> +{ + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + recorded.lock().expect("requests lock").push(body.clone()); + match body["method"].as_str() { + Some("server/discover") => { + assert_eq!( + body.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion"), + Some(&json!(MODERN_VERSION)) + ); + assert_eq!( + body.pointer("/params/_meta/io.modelcontextprotocol~1clientInfo/name"), + Some(&json!("codex-discovery-test")) + ); + modern_discover_response(&body) + } + Some("tools/list") => { + assert_eq!( + body.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion"), + Some(&json!(MODERN_VERSION)) + ); + assert_eq!( + body.pointer("/params/_meta/io.modelcontextprotocol~1clientInfo/name"), + Some(&json!("codex-discovery-test")) + ); + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": {"resultType": "complete", "tools": []}, + })) + } + other => panic!("unexpected modern lifecycle request: {other:?}"), + } + }) + .expect(2) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + let tools = client + .list_tools(/*params*/ None, Some(Duration::from_secs(5))) + .await?; + assert!(tools.tools.is_empty()); + + let methods = observed + .lock() + .expect("requests lock") + .iter() + .map(|request| request["method"].as_str().unwrap_or_default().to_owned()) + .collect::>(); + assert_eq!(methods, vec!["server/discover", "tools/list"]); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_does_not_follow_redirects_with_sensitive_headers() -> anyhow::Result<()> { + let redirect_target = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/forwarded")) + .and(header("x-api-key", "sensitive-key")) + .respond_with(ResponseTemplate::new(200)) + .expect(0) + .mount(&redirect_target) + .await; + + let resource_server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .and(header("x-api-key", "sensitive-key")) + .respond_with( + ResponseTemplate::new(307) + .insert_header("location", format!("{}/forwarded", redirect_target.uri())), + ) + .expect(1) + .mount(&resource_server) + .await; + + let client = RmcpClient::new_streamable_http_client_with_protocol_mode( + "discovery-redirect-test", + &format!("{}/mcp", resource_server.uri()), + /*bearer_token*/ None, + Some(HashMap::from([( + "x-api-key".to_string(), + "sensitive-key".to_string(), + )])), + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + McpProtocolMode::V20260728, + ) + .await?; + let error = initialize_client(&client) + .await + .expect_err("modern MCP discovery must not follow redirects"); + assert!( + error.to_string().contains("307"), + "redirect rejection should report its HTTP status: {error:#}" + ); + redirect_target.verify().await; + resource_server.verify().await; + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn legacy_mcp_requests_follow_same_origin_redirects_with_configured_headers() +-> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/forwarded")) + .and(header("x-api-key", "sensitive-key")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("initialize") => legacy_initialize_response(&body), + Some("notifications/initialized") => ResponseTemplate::new(202), + other => panic!("unexpected redirected legacy method: {other:?}"), + } + }) + .expect(2) + .mount(&server) + .await; + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(ResponseTemplate::new(307).insert_header("location", "/forwarded")) + .expect(2) + .mount(&server) + .await; + + let client = RmcpClient::new_streamable_http_client_with_protocol_mode( + "legacy-redirect-test", + &format!("{}/mcp", server.uri()), + /*bearer_token*/ None, + Some(HashMap::from([( + "x-api-key".to_string(), + "sensitive-key".to_string(), + )])), + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + McpProtocolMode::Legacy, + ) + .await?; + initialize_client(&client).await?; + server.verify().await; + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn legacy_fallback_does_not_follow_cross_origin_tool_redirects() -> anyhow::Result<()> { + for mode in [McpProtocolMode::Legacy, McpProtocolMode::V20260728] { + let destination = MockServer::start().await; + let server = MockServer::start().await; + let destination_url = format!("{}/private", destination.uri()); + Mock::given(method("POST")) + .and(path("/redirected")) + .respond_with(ResponseTemplate::new(307).insert_header("location", destination_url)) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("server/discover") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32601, "message": "method not found"}, + })), + Some("initialize") => legacy_initialize_response(&body), + Some("notifications/initialized") => ResponseTemplate::new(202), + Some("tools/call") => { + ResponseTemplate::new(307).insert_header("location", "/redirected") + } + other => panic!("unexpected legacy MCP method: {other:?}"), + } + }) + .mount(&server) + .await; + + let client = create_client(&server, mode).await?; + initialize_client(&client).await?; + let error = client + .call_tool( + "query_logs_sql".to_string(), + Some(json!({"query": "sensitive-query"})), + /*meta*/ None, + Some(Duration::from_secs(5)), + ) + .await + .expect_err("cross-origin tools/call redirect must fail"); + assert!( + format!("{error:#}").contains("different origin"), + "cross-origin tools/call redirect must explain its rejection: {error:#}" + ); + assert!(destination.received_requests().await.unwrap().is_empty()); + client.shutdown().await; + } + + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_accepts_metadata_namespaced_server_identity() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + modern_discover_response(&body) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + let server_info = initialize_client(&client).await?; + assert_eq!( + server_info.server_info, + Some(Implementation::new("modern-test", "1.0.0")) + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_accepts_missing_server_identity() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + let mut response = modern_discover_result(&body); + response["result"] + .as_object_mut() + .expect("discovery object") + .remove("_meta"); + ResponseTemplate::new(200).set_body_json(response) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + let server_info = initialize_client(&client).await?; + assert_eq!( + server_info, + ServerPeerInfo::new( + ProtocolVersion::V_2026_07_28, + ServerCapabilities::builder().enable_tools().build(), + ) + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_accepts_native_sse_responses() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + let response = modern_discover_result(&body); + ResponseTemplate::new(200).set_body_raw( + format!(": keepalive\n\nevent: message\ndata: {response}\n\n"), + "text/event-stream; charset=utf-8", + ) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_accepts_metadata_namespaced_server_identity_over_sse() +-> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + let response = modern_discover_result(&body); + ResponseTemplate::new(200).set_body_raw( + format!("event: message\ndata: {response}\n\n"), + "text/event-stream; charset=utf-8", + ) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_http_rejects_oversized_discovery_and_tool_json_responses() -> anyhow::Result<()> { + for oversized_method in ["server/discover", "tools/list", "tools/list:error"] { + let server = MockServer::start().await; + let padding = Arc::new("x".repeat(MAX_MCP_MESSAGE_BYTES + 1)); + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("server/discover") => { + let mut response = modern_discover_result(&body); + if oversized_method == "server/discover" { + response["result"]["instructions"] = json!(padding.as_str()); + } + ResponseTemplate::new(200).set_body_json(response) + } + Some("tools/list") if oversized_method == "tools/list:error" => { + ResponseTemplate::new(500).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32603, "message": padding.as_str()}, + })) + } + Some("tools/list") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "resultType": "complete", + "tools": [{ + "name": "oversized_tool", + "description": padding.as_str(), + "inputSchema": {"type": "object"}, + }], + }, + })), + other => panic!("unexpected MCP method {other:?}"), + } + }) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + let error = if oversized_method == "server/discover" { + initialize_client(&client) + .await + .expect_err("oversized discovery response must be rejected") + } else { + initialize_client(&client).await?; + client + .list_tools(/*params*/ None, Some(Duration::from_secs(5))) + .await + .expect_err("oversized tools response must be rejected") + }; + assert!( + format!("{error:#}").contains("8388608 bytes"), + "expected bounded {oversized_method} response, got {error:#}" + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn legacy_http_keeps_existing_large_json_response_behavior() -> anyhow::Result<()> { + let server = MockServer::start().await; + let padding = Arc::new("x".repeat(MAX_MCP_MESSAGE_BYTES + 1)); + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("initialize") => legacy_initialize_response(&body), + Some("notifications/initialized") => ResponseTemplate::new(202), + Some("tools/list") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "tools": [{ + "name": "legacy_large_tool", + "description": padding.as_str(), + "inputSchema": {"type": "object"}, + }], + }, + })), + other => panic!("unexpected legacy MCP method {other:?}"), + } + }) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::Legacy).await?; + initialize_client(&client).await?; + let tools = client + .list_tools(/*params*/ None, Some(Duration::from_secs(15))) + .await?; + assert_eq!( + tools.tools[0].description.as_deref().map(str::len), + Some(MAX_MCP_MESSAGE_BYTES + 1) + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_retries_transient_service_unavailable() -> anyhow::Result<()> { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + let attempt = { + let mut requests = recorded.lock().expect("requests lock"); + requests.push(body.clone()); + requests.len() + }; + if attempt == 1 { + ResponseTemplate::new(503).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": { + "code": -32000, + "message": "service unavailable", + }, + })) + } else { + modern_discover_response(&body) + } + }) + .expect(2) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + { + let requests = observed.lock().expect("requests lock"); + assert_eq!( + requests + .iter() + .map(|request| request["method"].as_str()) + .collect::>(), + vec![Some("server/discover"), Some("server/discover")], + "transient discovery failures must retry without falling back to legacy initialization" + ); + } + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_retries_json_unsupported_protocol_errors() -> anyhow::Result<()> { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + let attempt = { + let mut requests = recorded.lock().expect("requests lock"); + requests.push(body.clone()); + requests.len() + }; + if attempt == 1 { + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": { + "code": -32022, + "message": "unsupported protocol version", + "data": {"supported": [LEGACY_VERSION, MODERN_VERSION]}, + }, + })) + } else { + modern_discover_response(&body) + } + }) + .expect(2) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + { + let requests = observed.lock().expect("requests lock"); + assert_eq!(requests.len(), 2); + assert_ne!( + requests[0]["id"], requests[1]["id"], + "each SDK discovery attempt must use a distinct JSON-RPC request ID" + ); + } + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_sdk_emits_standard_headers_and_rejects_invalid_header_schemas() -> anyhow::Result<()> +{ + let server = MockServer::start().await; + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + assert_eq!( + request + .headers + .get("mcp-protocol-version") + .and_then(|value| value.to_str().ok()), + Some(MODERN_VERSION) + ); + assert_eq!( + request + .headers + .get("mcp-method") + .and_then(|value| value.to_str().ok()), + Some(method) + ); + + match method { + "server/discover" => modern_discover_response(&body), + "tools/list" => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "resultType": "complete", + "tools": [ + { + "name": "echo", + "inputSchema": { + "type": "object", + "properties": { + "tenant": {"type": "string", "x-mcp-header": "Tenant"}, + }, + }, + }, + { + "name": "invalid-header-schema", + "inputSchema": { + "type": "object", + "properties": { + "tenant": {"type": "string", "x-mcp-header": true}, + }, + }, + }, + ], + }, + })), + "tools/call" => { + assert_eq!( + request + .headers + .get("mcp-name") + .and_then(|value| value.to_str().ok()), + Some("echo") + ); + assert_eq!( + request + .headers + .get("mcp-param-tenant") + .and_then(|value| value.to_str().ok()), + Some("acme") + ); + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "resultType": "complete", + "content": [{"type": "text", "text": "headers work"}], + }, + })) + } + other => panic!("unexpected modern request: {other}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + let tools = client + .list_tools(/*params*/ None, Some(Duration::from_secs(5))) + .await?; + assert_eq!(tools.tools.len(), 1); + assert_eq!(tools.tools[0].name.as_ref(), "echo"); + + let result = client + .call_tool( + "echo".to_string(), + Some(json!({"tenant": "acme"})), + /*meta*/ None, + Some(Duration::from_secs(5)), + ) + .await?; + assert_eq!( + result.content[0].as_text().map(|text| text.text.as_str()), + Some("headers work") + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn legacy_mode_never_probes_server_discovery() -> anyhow::Result<()> { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + match method { + "initialize" => { + assert_eq!(body["params"]["protocolVersion"], LEGACY_VERSION); + legacy_initialize_response(&body) + } + "notifications/initialized" => ResponseTemplate::new(202), + other => panic!("legacy mode must not use modern lifecycle: {other}"), + } + }) + .expect(2) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::Legacy).await?; + initialize_client(&client).await?; + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["initialize", "notifications/initialized"] + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_mode_falls_back_only_for_a_json_rpc_method_not_found() -> anyhow::Result<()> { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + match method { + "server/discover" => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32601, "message": "method not found"}, + })), + "initialize" => { + assert_eq!(body["params"]["protocolVersion"], LEGACY_VERSION); + legacy_initialize_response(&body) + } + "notifications/initialized" => ResponseTemplate::new(202), + other => panic!("unexpected legacy fallback request: {other}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["server/discover", "initialize", "notifications/initialized"] + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_legacy_fallback_preserves_authentication_required_error() -> anyhow::Result<()> { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + match method { + "server/discover" => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32601, "message": "method not found"}, + })), + "initialize" => ResponseTemplate::new(401) + .insert_header("www-authenticate", "Bearer error=\"invalid_token\""), + other => panic!("unexpected authentication fallback request: {other}"), + } + }) + .expect(2) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + let error = initialize_client(&client) + .await + .expect_err("legacy fallback should surface the authentication challenge"); + assert!( + codex_rmcp_client::is_authentication_required_error(&error), + "legacy fallback must preserve the authentication-required classification: {error:#}" + ); + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["server/discover", "initialize"] + ); + Ok(()) +} + +#[tokio::test] +async fn modern_legacy_fallback_retries_transient_initialize_failures() -> anyhow::Result<()> { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + let initialize_attempt = { + let mut requests = recorded.lock().expect("requests lock"); + requests.push(method.into()); + requests + .iter() + .filter(|request| request.as_str() == "initialize") + .count() + }; + match method { + "server/discover" => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32601, "message": "method not found"}, + })), + "initialize" if initialize_attempt == 1 => { + ResponseTemplate::new(503).set_body_string("service unavailable") + } + "initialize" => legacy_initialize_response(&body), + "notifications/initialized" => ResponseTemplate::new(202), + other => panic!("unexpected retry fallback request: {other}"), + } + }) + .expect(5) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + assert_eq!( + *observed.lock().expect("requests lock"), + vec![ + "server/discover", + "initialize", + "server/discover", + "initialize", + "notifications/initialized", + ] + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_mode_falls_back_when_discovery_only_advertises_legacy_versions() +-> anyhow::Result<()> { + for legacy_response in ["discover-result", "unsupported-protocol"] { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + match method { + "server/discover" if legacy_response == "discover-result" => { + let mut response = modern_discover_result(&body); + response["result"]["supportedVersions"] = json!(["2025-11-25"]); + ResponseTemplate::new(200).set_body_json(response) + } + "server/discover" => ResponseTemplate::new(400).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": { + "code": -32022, + "message": "only legacy MCP protocols are supported", + "data": {"supported": [LEGACY_VERSION, "2025-11-25"]}, + }, + })), + "initialize" => { + assert_eq!(body["params"]["protocolVersion"], LEGACY_VERSION); + legacy_initialize_response(&body) + } + "notifications/initialized" => ResponseTemplate::new(202), + other => { + panic!("unexpected {legacy_response} fallback request: {other}") + } + } + }) + .expect(3) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["server/discover", "initialize", "notifications/initialized"], + "{legacy_response} should permit legacy initialization fallback" + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_never_downgrades_without_valid_legacy_only_version_evidence() +-> anyhow::Result<()> { + for (case, supported) in [ + ("missing-versions", None), + ("empty-versions", Some(json!([]))), + ( + "mixed-unknown-version", + Some(json!([LEGACY_VERSION, "1999-01-01"])), + ), + ("mixed-invalid-version", Some(json!([LEGACY_VERSION, 42]))), + ] { + let server = MockServer::start().await; + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + let mut error = json!({ + "code": -32022, + "message": "unsupported protocol version", + "data": {}, + }); + if let Some(supported) = &supported { + error["data"]["supported"] = supported.clone(); + } + ResponseTemplate::new(400).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": error, + })) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + assert!( + initialize_client(&client).await.is_err(), + "{case} must not silently downgrade to legacy" + ); + } + Ok(()) +} + +#[tokio::test] +async fn modern_mode_falls_back_when_legacy_server_explicitly_rejects_modern_version() +-> anyhow::Result<()> { + for error_code in [-32022, -32600, -32602] { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + match method { + "server/discover" => ResponseTemplate::new(400).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": { + "code": error_code, + "message": "Unsupported protocol version: 2026-07-28", + }, + })), + "initialize" => legacy_initialize_response(&body), + "notifications/initialized" => ResponseTemplate::new(202), + other => panic!("unexpected legacy fallback method: {other}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["server/discover", "initialize", "notifications/initialized"] + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_rejects_uncorrelated_legacy_version_rejection() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + ResponseTemplate::new(400).set_body_json(json!({ + "jsonrpc": "2.0", + "id": "unrelated-request", + "error": { + "code": -32602, + "message": "Unsupported protocol version: 2026-07-28", + }, + })) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + assert!( + initialize_client(&client).await.is_err(), + "an unrelated response ID must never authorize protocol downgrade" + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_rejects_uncorrelated_method_not_found_errors() -> anyhow::Result<()> { + for response_id in [Value::Null, json!("unrelated-request")] { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": response_id, + "error": {"code": -32601, "message": "method not found"}, + })) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + let error = initialize_client(&client) + .await + .expect_err("an unrelated error must not authorize protocol downgrade"); + assert!( + error.to_string().contains("did not match its request ID"), + "uncorrelated method-not-found error should explain its rejection: {error:#}" + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn modern_mode_falls_back_for_legacy_http_prevalidation_errors() -> anyhow::Result<()> { + for (message, content_type) in [ + ( + "Bad Request: Unsupported protocol version: 2026-07-28 (supported versions: 2025-11-25, 2025-06-18, 2025-03-26, 2024-11-05, 2024-10-07)", + "application/json", + ), + // tinymcp.dev omits the rejected version from its legacy error shape. + ( + "Bad Request: Unsupported protocol version (supported versions: 2025-06-18, 2025-03-26, 2024-11-05, 2024-10-07)", + "application/json", + ), + ( + "Bad Request: Unsupported protocol version (supported versions: 2025-06-18, 2025-03-26, 2024-11-05, 2024-10-07)", + "text/plain", + ), + ( + "Bad Request: Unsupported protocol version (supported versions: 2025-06-18, 2025-03-26, 2024-11-05, 2024-10-07)", + "application/octet-stream", + ), + ( + "Bad Request: No valid session ID provided", + "application/json", + ), + ] { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + match method { + "server/discover" => ResponseTemplate::new(400).set_body_raw( + serde_json::to_vec(&json!({ + "jsonrpc": "2.0", + "id": null, + "error": {"code": -32000, "message": message}, + })) + .expect("legacy discovery error serializes"), + content_type, + ), + "initialize" => legacy_initialize_response(&body), + "notifications/initialized" => ResponseTemplate::new(202), + other => panic!("unexpected legacy fallback method: {other}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["server/discover", "initialize", "notifications/initialized"] + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn modern_mode_negotiates_a_server_selected_legacy_protocol_version() -> anyhow::Result<()> { + for (message, content_type, negotiated_version) in [ + ( + "Bad Request: Unsupported protocol version: 2026-07-28 (supported versions: 2025-11-25)", + "application/json", + "2025-11-25", + ), + ( + "Bad Request: Unsupported protocol version (supported versions: 2025-11-25)", + "application/json", + "2025-11-25", + ), + ( + "Bad Request: Unsupported protocol version (supported versions: 2025-11-25)", + "text/plain", + "2025-11-25", + ), + ( + "Bad Request: Unsupported protocol version (supported versions: 2025-03-26)", + "application/json", + "2025-03-26", + ), + ( + "Bad Request: Unsupported protocol version (supported versions: 2024-11-05)", + "application/json", + "2024-11-05", + ), + ] { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + match method { + "server/discover" => ResponseTemplate::new(400).set_body_raw( + serde_json::to_vec(&json!({ + "jsonrpc": "2.0", + "id": null, + "error": {"code": -32000, "message": message}, + })) + .expect("legacy discovery error serializes"), + content_type, + ), + "initialize" => { + assert_eq!(body["params"]["protocolVersion"], LEGACY_VERSION); + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "protocolVersion": negotiated_version, + "capabilities": {"tools": {}}, + "serverInfo": {"name": "legacy-test", "version": "1.0.0"}, + }, + })) + } + "notifications/initialized" => { + assert_eq!( + request + .headers + .get("mcp-protocol-version") + .and_then(|value| value.to_str().ok()), + Some(negotiated_version) + ); + ResponseTemplate::new(202) + } + other => panic!("unexpected legacy fallback method: {other}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + let initialized = initialize_client(&client).await?; + assert_eq!(initialized.protocol_version.as_str(), negotiated_version); + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["server/discover", "initialize", "notifications/initialized"] + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_rejects_unrecognized_or_modern_null_id_errors() -> anyhow::Result<()> { + for (status, message, content_type) in [ + (400, "Bad Request: malformed request", "application/json"), + (400, "Bad Request: malformed request", "text/plain"), + ( + 400, + "Bad Request: Unsupported protocol version: 2026-07-28 (supported versions: 2025-06-18, 2099-01-01)", + "application/json", + ), + ( + 400, + "Bad Request: Unsupported protocol version (supported versions: 2025-06-18, 2099-01-01)", + "application/json", + ), + ( + 400, + "Bad Request: Unsupported protocol version (supported versions: 2025-06-18, 2099-01-01)", + "text/plain", + ), + ( + 400, + "Bad Request: Unsupported protocol version (supported versions: 2024-10-07)", + "application/json", + ), + ( + 200, + "Bad Request: No valid session ID provided", + "application/json", + ), + ] { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + ResponseTemplate::new(status).set_body_raw( + serde_json::to_vec(&json!({ + "jsonrpc": "2.0", + "id": null, + "error": {"code": -32000, "message": message}, + })) + .expect("legacy discovery error serializes"), + content_type, + ) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + assert!( + initialize_client(&client).await.is_err(), + "unproven legacy downgrade must fail (HTTP {status}, {message})" + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_rejects_non_legacy_http_prevalidation_errors() -> anyhow::Result<()> { + for (status, error_code, response_id, content_type) in [ + (400, -32022, None, "application/json"), + (400, -32021, None, "application/json"), + (400, -32020, None, "application/json"), + (400, -32601, None, "application/json"), + (400, -32602, None, "application/json"), + (400, -32020, None, "text/plain"), + (400, -32000, Some("unrelated-request"), "application/json"), + (401, -32000, None, "application/json"), + (403, -32000, None, "application/json"), + (200, -32000, None, "application/json"), + ] { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + ResponseTemplate::new(status).set_body_raw( + serde_json::to_vec(&json!({ + "jsonrpc": "2.0", + "id": response_id, + "error": {"code": error_code, "message": "discovery rejected"}, + })) + .expect("discovery error serializes"), + content_type, + ) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + assert!( + initialize_client(&client).await.is_err(), + "non-legacy discovery error must not downgrade (HTTP {status}, error {error_code}, response ID {response_id:?})" + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn modern_mode_falls_back_for_plain_404_and_405_discovery_responses() -> anyhow::Result<()> { + for status in [404, 405] { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + match method { + "server/discover" => { + ResponseTemplate::new(status).set_body_string("legacy MCP endpoint") + } + "initialize" => legacy_initialize_response(&body), + "notifications/initialized" => ResponseTemplate::new(202), + other => panic!("unexpected HTTP {status} fallback request: {other}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + initialize_client(&client).await?; + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["server/discover", "initialize", "notifications/initialized"], + "HTTP {status} should permit legacy initialization fallback" + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn modern_discovery_rejection_does_not_downgrade_to_legacy() -> anyhow::Result<()> { + for (status, error_code) in [(400, -32602), (403, -32021), (200, -32020)] { + let server = MockServer::start().await; + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + ResponseTemplate::new(status).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": error_code, "message": "modern discovery rejected"}, + })) + }) + .expect(1) + .mount(&server) + .await; + + let client = create_client(&server, McpProtocolMode::V20260728).await?; + let error = initialize_client(&client) + .await + .expect_err("modern discovery errors must not silently downgrade"); + assert!( + error.to_string().contains("modern discovery rejected"), + "unexpected discovery rejection for HTTP {status}: {error:#}" + ); + } + Ok(()) +} + +#[tokio::test] +async fn stdio_protocol_marker_rejects_unknown_versions_before_launch() -> anyhow::Result<()> { + let env = HashMap::from([( + OsString::from("CODEX_MCP_PROTOCOL_VERSION"), + OsString::from("1999-01-01"), + )]); + let error = RmcpClient::new_stdio_client_with_protocol_mode( + OsString::from("never-launch"), + Vec::new(), + Some(env), + &[], + /*cwd*/ None, + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)), + McpProtocolMode::V20260728, + ) + .await + .err() + .expect("an unsupported stdio protocol must fail before launch"); + + assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput); + assert!(error.to_string().contains("1999-01-01")); + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/mcp_2026_message_limits.rs b/codex-rs/rmcp-client/tests/mcp_2026_message_limits.rs new file mode 100644 index 0000000000000000000000000000000000000000..785aa4e1956c95a1b3319315ce1b135be6e0f713 --- /dev/null +++ b/codex-rs/rmcp-client/tests/mcp_2026_message_limits.rs @@ -0,0 +1,292 @@ +use std::collections::HashMap; +use std::ffi::OsString; +use std::sync::Arc; +use std::time::Duration; + +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::ExecutorStdioServerLauncher; +use codex_rmcp_client::LocalStdioServerLauncher; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::StdioServerLauncher; +use futures::FutureExt; +use rmcp::model::ClientCapabilities; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use serde_json::Value; +use serde_json::json; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::Request; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +const MAX_MCP_MESSAGE_BYTES: usize = 8 * 1024 * 1024; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum OversizedHttpResponse { + DiscoveryJson, + ToolJson, + ToolError, + ToolSseEvent, +} + +fn initialize_params() -> InitializeRequestParams { + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("codex-message-limit-test", "0.0.0-test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18) +} + +async fn initialize(client: &RmcpClient) -> anyhow::Result<()> { + client + .initialize( + initialize_params(), + Some(Duration::from_secs(15)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: None, + }) + } + .boxed() + }), + ) + .await?; + Ok(()) +} + +fn discovery_response(request: &Value, padding: Option<&str>) -> Value { + json!({ + "jsonrpc": "2.0", + "id": request["id"], + "result": { + "resultType": "complete", + "supportedVersions": ["2026-07-28"], + "capabilities": {"tools": {}}, + "instructions": padding.unwrap_or_default(), + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "message-limit-test", + "version": "1.0.0", + }, + }, + "ttlMs": 0, + "cacheScope": "private", + }, + }) +} + +fn tools_response(request: &Value, padding: &str) -> Value { + json!({ + "jsonrpc": "2.0", + "id": request["id"], + "result": { + "resultType": "complete", + "tools": [{ + "name": "oversized_tool", + "description": padding, + "inputSchema": {"type": "object"}, + }], + }, + }) +} + +async fn http_client(server: &MockServer, mode: McpProtocolMode) -> anyhow::Result { + RmcpClient::new_streamable_http_client_with_protocol_mode( + "message-limits", + &format!("{}/mcp", server.uri()), + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + mode, + ) + .await +} + +#[tokio::test] +async fn modern_http_rejects_oversized_json_error_and_sse_bodies() -> anyhow::Result<()> { + for oversized in [ + OversizedHttpResponse::DiscoveryJson, + OversizedHttpResponse::ToolJson, + OversizedHttpResponse::ToolError, + OversizedHttpResponse::ToolSseEvent, + ] { + let server = MockServer::start().await; + let padding = Arc::new("x".repeat(MAX_MCP_MESSAGE_BYTES + 1)); + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("JSON-RPC request"); + match body["method"].as_str() { + Some("server/discover") => { + ResponseTemplate::new(200).set_body_json(discovery_response( + &body, + (oversized == OversizedHttpResponse::DiscoveryJson) + .then_some(padding.as_str()), + )) + } + Some("tools/list") if oversized == OversizedHttpResponse::ToolError => { + ResponseTemplate::new(500).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32603, "message": padding.as_str()}, + })) + } + Some("tools/list") if oversized == OversizedHttpResponse::ToolSseEvent => { + ResponseTemplate::new(200).set_body_raw( + format!( + "event: message\ndata: {}\n\n", + tools_response(&body, &padding), + ), + "text/event-stream; charset=utf-8", + ) + } + Some("tools/list") => { + ResponseTemplate::new(200).set_body_json(tools_response(&body, &padding)) + } + other => panic!("unexpected MCP method {other:?}"), + } + }) + .mount(&server) + .await; + + let client = http_client(&server, McpProtocolMode::V20260728).await?; + if oversized == OversizedHttpResponse::DiscoveryJson { + let error = initialize(&client) + .await + .expect_err("oversized discovery must be rejected"); + assert!(format!("{error:#}").contains("8388608 bytes")); + } else { + initialize(&client).await?; + let error = client + .list_tools(/*params*/ None, Some(Duration::from_secs(5))) + .await + .expect_err("oversized tools response must be rejected"); + let error = format!("{error:#}"); + assert!( + error.contains("8388608 bytes") + || (oversized == OversizedHttpResponse::ToolSseEvent + && (error.contains("timed out") || error.contains("Transport closed"))), + "unexpected {oversized:?} error: {error}" + ); + } + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn legacy_http_keeps_existing_large_json_response_behavior() -> anyhow::Result<()> { + let server = MockServer::start().await; + let padding = Arc::new("x".repeat(MAX_MCP_MESSAGE_BYTES + 1)); + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("JSON-RPC request"); + match body["method"].as_str() { + Some("initialize") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "protocolVersion": "2025-06-18", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "legacy-large-response", "version": "1.0.0"}, + }, + })), + Some("notifications/initialized") => ResponseTemplate::new(202), + Some("tools/list") => { + ResponseTemplate::new(200).set_body_json(tools_response(&body, &padding)) + } + other => panic!("unexpected legacy MCP method {other:?}"), + } + }) + .mount(&server) + .await; + + let client = http_client(&server, McpProtocolMode::Legacy).await?; + initialize(&client).await?; + let tools = client + .list_tools(/*params*/ None, Some(Duration::from_secs(15))) + .await?; + assert_eq!( + tools.tools[0].description.as_deref().map(str::len), + Some(MAX_MCP_MESSAGE_BYTES + 1) + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn modern_local_and_executor_stdio_reject_oversized_lines() -> anyhow::Result<()> { + let server = codex_utils_cargo_bin::cargo_bin("test_mcp_2026_stdio_server")?; + + for (modern, executor) in [(false, false), (false, true), (true, false), (true, true)] { + let mode = if modern { + McpProtocolMode::V20260728 + } else { + McpProtocolMode::Legacy + }; + let env = modern.then(|| { + HashMap::from([( + OsString::from("CODEX_MCP_PROTOCOL_VERSION"), + OsString::from("2026-07-28"), + )]) + }); + let launcher: Arc = if executor { + Arc::new(ExecutorStdioServerLauncher::new( + Environment::default_for_tests().get_exec_backend(), + )) + } else { + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)) + }; + let fixture_mode = if modern { + "oversized-stdout" + } else { + "oversized-stdout-legacy" + }; + + let client = RmcpClient::new_stdio_client_with_protocol_mode( + server.clone().into(), + vec![OsString::from(fixture_mode)], + env, + &[], + Some(std::env::current_dir()?.to_string_lossy().into_owned()), + launcher, + mode, + ) + .await?; + initialize(&client).await?; + let result = client + .list_tools(/*params*/ None, Some(Duration::from_secs(10))) + .await; + if !modern && !executor { + let tools = result?; + assert_eq!( + tools.tools[0].description.as_deref().map(str::len), + Some(MAX_MCP_MESSAGE_BYTES + 1), + "legacy local stdio must preserve the existing unbounded native codec" + ); + } else { + assert!( + result.is_err(), + "oversized stdio line must be rejected (modern={modern}, executor={executor})" + ); + } + client.shutdown().await; + } + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/mcp_2026_mrtr.rs b/codex-rs/rmcp-client/tests/mcp_2026_mrtr.rs new file mode 100644 index 0000000000000000000000000000000000000000..5d86e92219cb96db749fe43d143921a991496329 --- /dev/null +++ b/codex-rs/rmcp-client/tests/mcp_2026_mrtr.rs @@ -0,0 +1,736 @@ +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; + +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_rmcp_client::Elicitation; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use rmcp::model::ClientCapabilities; +use rmcp::model::ElicitRequestParams; +use rmcp::model::ElicitationCapability; +use rmcp::model::FormElicitationCapability; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ServerResult; +use rmcp::model::UrlElicitationCapability; +use serde_json::Value; +use serde_json::json; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::Request; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +const MODERN_VERSION: &str = "2026-07-28"; +const OPAQUE_STATE: &str = " opaque/\u{2603}/=?base64?literal?=\n"; + +#[path = "mcp_2026_mrtr/native_verification_tests.rs"] +mod native_verification; + +fn discover_response(body: &Value) -> ResponseTemplate { + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "resultType": "complete", + "supportedVersions": [MODERN_VERSION], + "capabilities": {"tools": {}, "resources": {}}, + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "mrtr-test", + "version": "1.0.0", + }, + }, + "ttlMs": 0, + "cacheScope": "private", + }, + })) +} + +fn result_response(body: &Value, result: Value) -> ResponseTemplate { + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": result, + })) +} + +fn sse_result_response(body: &Value, result: Value) -> ResponseTemplate { + let message = json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": result, + }); + ResponseTemplate::new(200).set_body_raw( + format!("event: message\ndata: {message}\n\n"), + "text/event-stream; charset=utf-8", + ) +} + +fn elicitation_request(mode: &str) -> Value { + let params = match mode { + "form" => json!({ + "mode": "form", + "message": "Confirm the MCP request.", + "requestedSchema": { + "type": "object", + "properties": {"confirmed": {"type": "boolean"}}, + "required": ["confirmed"], + }, + "_meta": {"inputContext": "server-context"}, + }), + "url" => json!({ + "mode": "url", + "message": "Approve the MCP request.", + "url": "https://example.test/approve", + "elicitationId": "approval-123", + "_meta": {"inputContext": "server-context"}, + }), + other => panic!("unexpected elicitation mode: {other}"), + }; + json!({"method": "elicitation/create", "params": params}) +} + +async fn create_client( + server: &MockServer, + elicitation_modes: Arc>>, +) -> anyhow::Result { + let client = RmcpClient::new_streamable_http_client_with_protocol_mode( + "mrtr-test", + &format!("{}/mcp", server.uri()), + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + McpProtocolMode::V20260728, + ) + .await?; + + let mut capabilities = ClientCapabilities::default(); + capabilities.elicitation = Some( + ElicitationCapability::new() + .with_form(FormElicitationCapability::new()) + .with_url(UrlElicitationCapability::new()), + ); + client + .initialize( + InitializeRequestParams::new( + capabilities, + Implementation::new("codex-mrtr-test", "0.0.0"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(5)), + Box::new(move |request_id, request| { + let elicitation_modes = Arc::clone(&elicitation_modes); + async move { + assert_eq!(request_id.to_string(), "confirmation"); + let Elicitation::Mcp(request) = request else { + anyhow::bail!("MRTR input must be a standard MCP elicitation"); + }; + let mode = match request { + ElicitRequestParams::FormElicitationParams { meta, .. } => { + assert_eq!( + meta.and_then(|meta| meta.get("inputContext").cloned()), + Some(json!("server-context")) + ); + "form" + } + ElicitRequestParams::UrlElicitationParams { meta, .. } => { + assert_eq!( + meta.and_then(|meta| meta.get("inputContext").cloned()), + Some(json!("server-context")) + ); + "url" + } + _ => anyhow::bail!("unsupported elicitation mode"), + }; + elicitation_modes + .lock() + .map_err(|_| anyhow::anyhow!("elicitation lock was poisoned"))? + .push(mode.to_owned()); + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({"confirmed": true})), + meta: Some(json!({"clientContext": "preserved"})), + }) + } + .boxed() + }), + ) + .await?; + Ok(client) +} + +#[tokio::test] +async fn modern_sse_input_required_preserves_response_metadata() -> anyhow::Result<()> { + let server = MockServer::start().await; + let expected_result = json!({ + "resultType": "input_required", + "requestState": OPAQUE_STATE, + "_meta": {"responseContext": "preserved"}, + }); + let result = expected_result.clone(); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") => sse_result_response(&body, result.clone()), + other => panic!("unexpected MRTR request: {other:?}"), + } + }) + .expect(2) + .mount(&server) + .await; + + let client = create_client(&server, Arc::new(Mutex::new(Vec::new()))).await?; + let result = client + .send_custom_request_with_timeout( + "tools/call", + Some(json!({"name": "confirm", "arguments": {}})), + Some(Duration::from_secs(5)), + ) + .await?; + let ServerResult::InputRequiredResult(result) = result else { + anyhow::bail!("SSE response must remain an input_required result"); + }; + + assert_eq!(serde_json::to_value(result)?, expected_result); + + client.shutdown().await; + server.verify().await; + + Ok(()) +} + +#[tokio::test] +async fn modern_tool_mrtr_drives_form_and_url_elicitation_and_preserves_metadata() +-> anyhow::Result<()> { + for mode in ["form", "url"] { + let server = MockServer::start().await; + let calls = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&calls); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") => { + recorded.lock().expect("requests lock").push(body.clone()); + assert_eq!( + body.pointer("/params/_meta/requestContext"), + Some(&json!("caller-context")) + ); + assert_eq!( + body.pointer("/params/_meta/io.modelcontextprotocol~1protocolVersion"), + Some(&json!(MODERN_VERSION)) + ); + if body.pointer("/params/inputResponses").is_none() { + result_response( + &body, + json!({ + "resultType": "input_required", + "inputRequests": { + "confirmation": elicitation_request(mode), + }, + "requestState": OPAQUE_STATE, + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "mrtr-test", + "version": "1.0.0", + }, + "responseContext": "preserved", + }, + }), + ) + } else { + assert_eq!( + body.pointer("/params/requestState"), + Some(&json!(OPAQUE_STATE)) + ); + assert_eq!( + body.pointer("/params/inputResponses/confirmation"), + Some(&json!({ + "action": "accept", + "content": {"confirmed": true}, + "_meta": {"clientContext": "preserved"}, + })) + ); + result_response( + &body, + json!({ + "resultType": "complete", + "content": [{"type": "text", "text": "MRTR completed"}], + }), + ) + } + } + other => panic!("unexpected MRTR request: {other:?}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let elicitation_modes = Arc::new(Mutex::new(Vec::new())); + let client = create_client(&server, Arc::clone(&elicitation_modes)).await?; + let result = client + .call_tool( + "confirm".into(), + Some(json!({"mode": mode})), + Some(json!({"requestContext": "caller-context"})), + Some(Duration::from_secs(5)), + ) + .await?; + assert_eq!( + result.content[0].as_text().map(|text| text.text.as_str()), + Some("MRTR completed") + ); + assert_eq!( + *elicitation_modes.lock().expect("elicitation lock"), + vec![mode] + ); + + { + let calls = calls.lock().expect("requests lock"); + assert_eq!(calls.len(), 2); + assert_ne!(calls[0]["id"], calls[1]["id"]); + } + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn modern_tool_mrtr_uses_recovered_protocol_after_legacy_session_expiry() -> anyhow::Result<()> +{ + let server = MockServer::start().await; + let calls = Arc::new(Mutex::new(Vec::::new())); + let sessions = Arc::new(Mutex::new(Vec::::new())); + let recorded_calls = Arc::clone(&calls); + let recorded_sessions = Arc::clone(&sessions); + + Mock::given(method("GET")) + .and(path("/mcp")) + .respond_with(ResponseTemplate::new(405)) + .mount(&server) + .await; + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("server/discover") => { + let mut sessions = recorded_sessions.lock().expect("sessions lock"); + if sessions.is_empty() { + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32601, "message": "method not found"}, + })) + } else { + sessions.push("modern-session".to_owned()); + discover_response(&body).insert_header("mcp-session-id", "modern-session") + } + } + Some("initialize") => { + assert_eq!( + body.pointer("/params/protocolVersion"), + Some(&json!(ProtocolVersion::V_2025_06_18.as_str())) + ); + recorded_sessions + .lock() + .expect("sessions lock") + .push("legacy-session".to_owned()); + ResponseTemplate::new(200) + .set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "protocolVersion": ProtocolVersion::V_2025_06_18.as_str(), + "capabilities": {"tools": {}}, + "serverInfo": {"name": "mrtr-test", "version": "1.0.0"}, + }, + })) + .insert_header("mcp-session-id", "legacy-session") + } + Some("notifications/initialized") => ResponseTemplate::new(202), + Some("tools/call") => { + assert_eq!( + body.pointer("/params/_meta/requestContext"), + Some(&json!("caller-context")) + ); + let attempt = { + let mut calls = recorded_calls.lock().expect("requests lock"); + calls.push(body.clone()); + calls.len() + }; + + match attempt { + 1 => { + assert_eq!( + body.pointer( + "/params/_meta/io.modelcontextprotocol~1protocolVersion" + ), + None + ); + ResponseTemplate::new(404) + } + 2 => { + assert_eq!( + body.pointer( + "/params/_meta/io.modelcontextprotocol~1protocolVersion" + ), + Some(&json!(MODERN_VERSION)) + ); + result_response( + &body, + json!({ + "resultType": "input_required", + "inputRequests": { + "confirmation": elicitation_request("form"), + }, + "requestState": OPAQUE_STATE, + }), + ) + } + 3 => { + assert_eq!( + body.pointer( + "/params/_meta/io.modelcontextprotocol~1protocolVersion" + ), + Some(&json!(MODERN_VERSION)) + ); + assert_eq!( + body.pointer("/params/requestState"), + Some(&json!(OPAQUE_STATE)) + ); + assert_eq!( + body.pointer("/params/inputResponses/confirmation"), + Some(&json!({ + "action": "accept", + "content": {"confirmed": true}, + "_meta": {"clientContext": "preserved"}, + })) + ); + result_response( + &body, + json!({ + "resultType": "complete", + "content": [{"type": "text", "text": "recovered MRTR completed"}], + }), + ) + } + other => panic!("unexpected recovered MRTR attempt: {other}"), + } + } + other => panic!("unexpected recovered MRTR request: {other:?}"), + } + }) + .mount(&server) + .await; + + let elicitation_modes = Arc::new(Mutex::new(Vec::new())); + let client = create_client(&server, Arc::clone(&elicitation_modes)).await?; + let result = client + .call_tool( + "confirm".into(), + Some(json!({})), + Some(json!({"requestContext": "caller-context"})), + Some(Duration::from_secs(5)), + ) + .await?; + + assert_eq!( + result.content[0].as_text().map(|text| text.text.as_str()), + Some("recovered MRTR completed") + ); + assert_eq!( + *sessions.lock().expect("sessions lock"), + vec!["legacy-session", "modern-session"] + ); + assert_eq!(calls.lock().expect("requests lock").len(), 3); + assert_eq!( + *elicitation_modes.lock().expect("elicitation lock"), + vec!["form"] + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_resource_mrtr_retries_with_opaque_request_state() -> anyhow::Result<()> { + let server = MockServer::start().await; + let calls = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&calls); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("resources/read") => { + recorded.lock().expect("requests lock").push(body.clone()); + if body.pointer("/params/inputResponses").is_none() { + result_response( + &body, + json!({ + "resultType": "input_required", + "inputRequests": { + "confirmation": elicitation_request("form"), + }, + "requestState": OPAQUE_STATE, + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "mrtr-test", + "version": "1.0.0", + }, + "responseContext": "preserved", + }, + }), + ) + } else { + assert_eq!( + body.pointer("/params/requestState"), + Some(&json!(OPAQUE_STATE)) + ); + assert_eq!( + body.pointer("/params/inputResponses/confirmation/_meta/clientContext"), + Some(&json!("preserved")) + ); + result_response( + &body, + json!({ + "resultType": "complete", + "contents": [{ + "uri": "memo://requires-confirmation", + "mimeType": "text/plain", + "text": "approved resource", + }], + }), + ) + } + } + other => panic!("unexpected resource MRTR request: {other:?}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let elicitation_modes = Arc::new(Mutex::new(Vec::new())); + let client = create_client(&server, Arc::clone(&elicitation_modes)).await?; + let result = client + .read_resource( + ReadResourceRequestParams::new("memo://requires-confirmation"), + Some(Duration::from_secs(5)), + ) + .await?; + assert_eq!(result.contents.len(), 1); + assert_eq!( + *elicitation_modes.lock().expect("elicitation lock"), + vec!["form"] + ); + + { + let calls = calls.lock().expect("requests lock"); + assert_eq!(calls.len(), 2); + assert_ne!(calls[0]["id"], calls[1]["id"]); + } + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_tool_mrtr_accepts_sse_discovery_and_metadata_bearing_rounds() -> anyhow::Result<()> +{ + let server = MockServer::start().await; + let calls = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&calls); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("server/discover") => sse_result_response( + &body, + json!({ + "resultType": "complete", + "supportedVersions": [MODERN_VERSION], + "capabilities": {"tools": {}}, + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "sse-mrtr-test", + "version": "1.0.0", + }, + }, + "ttlMs": 0, + "cacheScope": "private", + }), + ), + Some("tools/call") => { + recorded.lock().expect("requests lock").push(body.clone()); + if body.pointer("/params/inputResponses").is_none() { + sse_result_response( + &body, + json!({ + "resultType": "input_required", + "inputRequests": { + "confirmation": elicitation_request("form"), + }, + "requestState": OPAQUE_STATE, + "_meta": {"responseContext": "sse-round"}, + }), + ) + } else { + assert_eq!( + body.pointer("/params/requestState"), + Some(&json!(OPAQUE_STATE)) + ); + assert_eq!( + body.pointer("/params/inputResponses/confirmation/action"), + Some(&json!("accept")) + ); + sse_result_response( + &body, + json!({ + "resultType": "complete", + "content": [{"type": "text", "text": "SSE MRTR completed"}], + }), + ) + } + } + other => panic!("unexpected SSE MRTR request: {other:?}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let elicitation_modes = Arc::new(Mutex::new(Vec::new())); + let client = create_client(&server, Arc::clone(&elicitation_modes)).await?; + let result = client + .call_tool( + "confirm".to_owned(), + Some(json!({})), + /*meta*/ None, + Some(Duration::from_secs(5)), + ) + .await?; + + assert_eq!( + result.content[0].as_text().map(|text| text.text.as_str()), + Some("SSE MRTR completed") + ); + assert_eq!( + *elicitation_modes.lock().expect("elicitation lock"), + vec!["form"] + ); + assert_eq!(calls.lock().expect("requests lock").len(), 2); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_tool_mrtr_preserves_state_only_rounds_without_elicitation() -> anyhow::Result<()> { + let server = MockServer::start().await; + let states = Arc::new(Mutex::new(Vec::>::new())); + let recorded = Arc::clone(&states); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") => { + let state = body + .pointer("/params/requestState") + .and_then(Value::as_str) + .map(str::to_owned); + recorded.lock().expect("states lock").push(state.clone()); + match state.as_deref() { + None => result_response( + &body, + json!({ + "resultType": "input_required", + "requestState": "state-1", + "_meta": {"responseContext": "first-round"}, + }), + ), + Some("state-1") => result_response( + &body, + json!({ + "resultType": "input_required", + "requestState": "state-2", + "_meta": {"responseContext": "second-round"}, + }), + ), + Some("state-2") => result_response( + &body, + json!({ + "resultType": "complete", + "content": [{"type": "text", "text": "state completed"}], + }), + ), + Some(other) => panic!("unexpected opaque state: {other}"), + } + } + other => panic!("unexpected state-only MRTR request: {other:?}"), + } + }) + .expect(4) + .mount(&server) + .await; + + let elicitation_modes = Arc::new(Mutex::new(Vec::new())); + let client = create_client(&server, Arc::clone(&elicitation_modes)).await?; + let result = client + .call_tool( + "state-only".into(), + Some(json!({})), + /*meta*/ None, + Some(Duration::from_secs(5)), + ) + .await?; + assert_eq!( + result.content[0].as_text().map(|text| text.text.as_str()), + Some("state completed") + ); + assert!( + elicitation_modes + .lock() + .expect("elicitation lock") + .is_empty() + ); + assert_eq!( + *states.lock().expect("states lock"), + vec![ + None, + Some("state-1".to_string()), + Some("state-2".to_string()) + ] + ); + client.shutdown().await; + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/mcp_2026_mrtr/native_verification_tests.rs b/codex-rs/rmcp-client/tests/mcp_2026_mrtr/native_verification_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..8f26c649860a614dbcb99ed80f132639066a7fed --- /dev/null +++ b/codex-rs/rmcp-client/tests/mcp_2026_mrtr/native_verification_tests.rs @@ -0,0 +1,578 @@ +//! Exercises native verification through the HTTP MRTR tool-input envelope. + +use super::*; +use codex_protocol::mcp::OPENAI_ELICITATION_EXTENSION_ID; +use codex_rmcp_client::SendElicitation; +use pretty_assertions::assert_eq; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::time::timeout; + +fn native_request() -> Value { + json!({ + "method": "openai/elicitation/create", + "params": { + "mode": "openai/userVerification", "title": "Verify action", + "description": "Verify the requested operation", "challenge": "AQID" + } + }) +} + +fn rich_form_request() -> Value { + let mut request = elicitation_request("form"); + request["method"] = json!("openai/elicitation/create"); + request +} + +fn input_required(request: Value) -> Value { + json!({ + "resultType": "input_required", "content": [], "isError": false, + "inputRequests": {"verification": request}, "requestState": OPAQUE_STATE + }) +} + +fn capabilities() -> ClientCapabilities { + let mut capabilities = ClientCapabilities::default(); + capabilities.extensions = Some( + [( + OPENAI_ELICITATION_EXTENSION_ID.into(), + serde_json::Map::from_iter([("userVerification".into(), json!({}))]), + )] + .into_iter() + .collect(), + ); + capabilities +} + +async fn client( + server: &MockServer, + capabilities: ClientCapabilities, + handler: SendElicitation, +) -> anyhow::Result> { + let client = RmcpClient::new_streamable_http_client_with_protocol_mode( + "native-mrtr-test", + &format!("{}/mcp", server.uri()), + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + McpProtocolMode::V20260728, + ) + .await?; + client + .initialize( + InitializeRequestParams::new(capabilities, Implementation::new("test", "0.0.0")), + Some(Duration::from_secs(/*secs*/ 5)), + handler, + ) + .await?; + Ok(Arc::new(client)) +} + +async fn call(client: &RmcpClient) -> anyhow::Result { + Ok(serde_json::to_value( + client + .call_tool( + "verified_operation".into(), + Some(json!({"resource": "example"})), + Some(json!({"caller": "preserved"})), + Some(Duration::from_secs(/*secs*/ 5)), + ) + .await?, + )?) +} + +fn response(action: ElicitationAction, content: Value) -> ElicitationResponse { + ElicitationResponse { + action, + content: Some(content), + meta: Some(json!({"discard": true})), + } +} + +#[tokio::test] +async fn native_mrtr_returns_validated_proof_or_cancellation_in_content() -> anyhow::Result<()> { + let proof = json!({"credentialId": "test-credential", "signature": "BAUG"}); + for (action, content, expected) in [ + ( + ElicitationAction::Accept, + proof.clone(), + json!({"action": "accept", "content": proof}), + ), + ( + ElicitationAction::Cancel, + proof.clone(), + json!({"action": "cancel"}), + ), + ( + ElicitationAction::Decline, + proof, + json!({"action": "decline"}), + ), + ( + ElicitationAction::Accept, + json!({"signature": "invalid"}), + json!({"action": "cancel"}), + ), + ] { + let server = MockServer::start().await; + let final_result = + json!({"resultType": "complete", "content": [{"type": "text", "text": "done"}]}); + let final_response = final_result.clone(); + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().unwrap(); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") if body.pointer("/params/inputResponses").is_none() => { + let mut result = input_required(native_request()); + result["inputRequests"]["confirmation"] = elicitation_request("form"); + result["inputRequests"]["richConfirmation"] = rich_form_request(); + sse_result_response(&body, result) + } + Some("tools/call") => { + assert_eq!( + body["params"]["inputResponses"], + json!({ + "verification": expected, + "confirmation": {"action": "accept", "content": {"confirmed": true}}, + "richConfirmation": {"action": "accept", "content": {"confirmed": true}} + }) + ); + assert_eq!(body["params"]["requestState"], OPAQUE_STATE); + assert_eq!(body["params"]["name"], "verified_operation"); + assert_eq!(body["params"]["arguments"], json!({"resource": "example"})); + assert_eq!(body["params"]["_meta"]["caller"], "preserved"); + result_response(&body, final_response.clone()) + } + other => panic!("unexpected request: {other:?}"), + } + }) + .expect(/*r*/ 3) + .mount(&server) + .await; + let mut capabilities = capabilities(); + capabilities + .extensions + .as_mut() + .unwrap() + .get_mut(OPENAI_ELICITATION_EXTENSION_ID) + .unwrap() + .insert("form".into(), json!({})); + capabilities.elicitation = + Some(ElicitationCapability::new().with_form(FormElicitationCapability::new())); + let client = client( + &server, + capabilities, + Box::new(move |_, request| { + if let Elicitation::OpenAiElicitationForm { meta, message, requested_schema } = request { + let expected = rich_form_request(); + assert_eq!( + json!({"_meta": meta, "message": message, "requestedSchema": requested_schema}), + json!({ + "_meta": expected["params"]["_meta"], + "message": expected["params"]["message"], + "requestedSchema": expected["params"]["requestedSchema"] + }) + ); + return Box::pin(async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({"confirmed": true})), + meta: None, + }) + }); + } + if let Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + meta, .. + }) = request + { + assert_eq!( + meta.and_then(|meta| meta.get("inputContext").cloned()), + Some(json!("server-context")) + ); + return Box::pin(async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({"confirmed": true})), + meta: None, + }) + }); + } + assert_eq!( + request, + Elicitation::UserVerification { + title: "Verify action".into(), + description: "Verify the requested operation".into(), + challenge: "AQID".into(), + } + ); + let response = response(action.clone(), content.clone()); + Box::pin(async move { Ok(response) }) + }), + ) + .await?; + assert_eq!(call(&client).await?, final_result); + let calls: Vec<_> = server + .received_requests() + .await + .unwrap() + .into_iter() + .map(|request| request.body_json::().unwrap()) + .filter(|body| body["method"] == "tools/call") + .collect(); + assert_ne!(calls[0]["id"], calls[1]["id"]); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn native_mrtr_bounds_repeated_input_requests() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().unwrap(); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") => result_response(&body, input_required(native_request())), + other => panic!("unexpected request: {other:?}"), + } + }) + .expect(1 + rmcp::model::DEFAULT_MRTR_MAX_ROUNDS as u64) + .mount(&server) + .await; + let client = client( + &server, + capabilities(), + Box::new(|_, _| { + Box::pin(async { + Ok(response( + ElicitationAction::Accept, + json!({"credentialId": "test", "signature": "BAUG"}), + )) + }) + }), + ) + .await?; + let error = call(&client).await.unwrap_err(); + assert!(error.to_string().contains("MRTR"), "{error}"); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn native_mrtr_rejects_unadvertised_and_malformed_requests_without_prompting() +-> anyhow::Result<()> { + let mut malformed = native_request(); + malformed["params"]["challenge"] = json!("not base64!"); + for (capabilities, request) in [ + (ClientCapabilities::default(), native_request()), + (capabilities(), rich_form_request()), + (capabilities(), malformed), + ( + capabilities(), + json!({"method": "arbitrary/custom", "params": {}}), + ), + ] { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |http: &Request| { + let body: Value = http.body_json().unwrap(); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") => result_response(&body, input_required(request.clone())), + other => panic!("unexpected request: {other:?}"), + } + }) + .expect(/*r*/ 2) + .mount(&server) + .await; + let client = client( + &server, + capabilities, + Box::new(|_, _| panic!("must not prompt")), + ) + .await?; + assert!(call(&client).await.is_err()); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn mrtr_rejects_unsupported_elicitation_modes_without_prompting() -> anyhow::Result<()> { + for mode in [None, Some(json!(null)), Some(json!(0)), Some(json!("url"))] { + let server = MockServer::start().await; + let mut input = native_request(); + if let Some(mode) = mode { + input["params"]["mode"] = mode; + } else { + input["params"].as_object_mut().unwrap().remove("mode"); + } + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().unwrap(); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") => result_response(&body, input_required(input.clone())), + other => panic!("unexpected request: {other:?}"), + } + }) + .expect(/*r*/ 2) + .mount(&server) + .await; + let client = client( + &server, + capabilities(), + Box::new(|_, _| panic!("must not prompt")), + ) + .await?; + let error = call(&client).await.unwrap_err(); + assert_eq!( + codex_rmcp_client::mcp_error(&error), + Some(&rmcp::ErrorData::invalid_request( + "unsupported OpenAI elicitation mode", + /*data*/ None + )) + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test] +async fn mrtr_preserves_auth_challenges_after_input_responses() -> anyhow::Result<()> { + for (request, action, content) in [ + ( + elicitation_request("form"), + ElicitationAction::Accept, + json!({"confirmed": true}), + ), + ( + rich_form_request(), + ElicitationAction::Accept, + json!({"confirmed": true}), + ), + ( + native_request(), + ElicitationAction::Accept, + json!({"credentialId": "test", "signature": "BAUG"}), + ), + (native_request(), ElicitationAction::Cancel, Value::Null), + (native_request(), ElicitationAction::Decline, Value::Null), + ] { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |http: &Request| { + let body: Value = http.body_json().unwrap(); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") if body.pointer("/params/inputResponses").is_none() => { + result_response(&body, input_required(request.clone())) + } + Some("tools/call") => ResponseTemplate::new(/*s*/ 401) + .insert_header("www-authenticate", "Bearer error=\"invalid_token\""), + other => panic!("unexpected request: {other:?}"), + } + }) + .expect(/*r*/ 3) + .mount(&server) + .await; + let mut capabilities = capabilities(); + capabilities + .extensions + .as_mut() + .unwrap() + .get_mut(OPENAI_ELICITATION_EXTENSION_ID) + .unwrap() + .insert("form".into(), json!({})); + capabilities.elicitation = + Some(ElicitationCapability::new().with_form(FormElicitationCapability::new())); + let client = client( + &server, + capabilities, + Box::new(move |_, _| { + let response = response(action.clone(), content.clone()); + Box::pin(async move { Ok(response) }) + }), + ) + .await?; + assert_eq!( + call(&client).await?, + serde_json::to_value( + rmcp::model::CallToolResult::error(vec![rmcp::model::ContentBlock::text( + "Authentication required", + )]) + .with_meta(Some(rmcp::model::MetaObject::from( + serde_json::Map::from_iter([( + "mcp/www_authenticate".into(), + json!(["Bearer error=\"invalid_token\""]), + )]) + ))) + )? + ); + client.shutdown().await; + server.verify().await; + } + Ok(()) +} + +#[tokio::test] +async fn native_mrtr_does_not_restart_after_continuation_session_expiry() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().unwrap(); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") if body.pointer("/params/inputResponses").is_none() => { + result_response(&body, input_required(native_request())) + } + Some("tools/call") => ResponseTemplate::new(/*s*/ 404), + other => panic!("unexpected request: {other:?}"), + } + }) + .expect(/*r*/ 3) + .mount(&server) + .await; + let client = client( + &server, + capabilities(), + Box::new(|_, _| { + Box::pin(async { + Ok(response( + ElicitationAction::Accept, + json!({"credentialId": "test", "signature": "BAUG"}), + )) + }) + }), + ) + .await?; + let error = call(&client).await.unwrap_err(); + assert!( + error.to_string().contains("the tool was not restarted"), + "{error}" + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn native_mrtr_does_not_restart_after_proof_then_state_only_round() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().unwrap(); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") if body.pointer("/params/requestState").is_none() => { + result_response(&body, input_required(native_request())) + } + Some("tools/call") if body.pointer("/params/inputResponses").is_some() => { + result_response( + &body, + json!({"resultType": "input_required", "requestState": "after-proof"}), + ) + } + Some("tools/call") => ResponseTemplate::new(/*s*/ 404), + other => panic!("unexpected request: {other:?}"), + } + }) + .expect(/*r*/ 4) + .mount(&server) + .await; + let client = client( + &server, + capabilities(), + Box::new(|_, _| { + Box::pin(async { + Ok(response( + ElicitationAction::Accept, + json!({"credentialId": "test", "signature": "BAUG"}), + )) + }) + }), + ) + .await?; + let error = call(&client).await.unwrap_err(); + assert!( + error.to_string().contains("the tool was not restarted"), + "{error}" + ); + client.shutdown().await; + server.verify().await; + Ok(()) +} + +#[tokio::test] +async fn native_mrtr_concurrent_prompts_have_independent_cancellation() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().unwrap(); + match body["method"].as_str() { + Some("server/discover") => discover_response(&body), + Some("tools/call") if body.pointer("/params/inputResponses").is_none() => { + result_response(&body, input_required(native_request())) + } + Some("tools/call") => { + assert_eq!( + body["params"]["inputResponses"], + json!({"verification": {"action": "cancel"}}) + ); + result_response(&body, json!({"resultType": "complete", "content": []})) + } + other => panic!("unexpected request: {other:?}"), + } + }) + .expect(/*r*/ 4) + .mount(&server) + .await; + let (tx, mut rx) = mpsc::unbounded_channel(); + let client = client( + &server, + capabilities(), + Box::new(move |id, _| { + let (reply_tx, reply_rx) = oneshot::channel(); + tx.send((id, reply_tx)).unwrap(); + Box::pin(async move { Ok(reply_rx.await?) }) + }), + ) + .await?; + let first_client = Arc::clone(&client); + let first_call = tokio::spawn(async move { call(&first_client).await }); + let (first_id, mut first_reply) = timeout(Duration::from_secs(/*secs*/ 5), rx.recv()) + .await? + .unwrap(); + let second_client = Arc::clone(&client); + let second_call = tokio::spawn(async move { call(&second_client).await }); + let (second_id, second_reply) = timeout(Duration::from_secs(/*secs*/ 5), rx.recv()) + .await? + .unwrap(); + assert_ne!(first_id, second_id); + first_call.abort(); + assert!(first_call.await.unwrap_err().is_cancelled()); + timeout(Duration::from_secs(/*secs*/ 5), first_reply.closed()).await?; + assert!(!second_reply.is_closed()); + second_reply + .send(response(ElicitationAction::Cancel, Value::Null)) + .unwrap(); + assert_eq!( + timeout(Duration::from_secs(/*secs*/ 5), second_call).await???, + json!({"resultType": "complete", "content": []}) + ); + client.shutdown().await; + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/mcp_2026_oauth_discovery.rs b/codex-rs/rmcp-client/tests/mcp_2026_oauth_discovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..a26dc7fbee6fc6563ad6f244c389efce8032f49f --- /dev/null +++ b/codex-rs/rmcp-client/tests/mcp_2026_oauth_discovery.rs @@ -0,0 +1,687 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Context as _; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_exec_server::HttpClient; +use codex_rmcp_client::McpOAuthCallbackMode; +use codex_rmcp_client::McpOAuthClientRegistration; +use codex_rmcp_client::OAuthDiscoveryTimeout; +use codex_rmcp_client::StreamableHttpOAuthDiscovery; +use codex_rmcp_client::StreamableHttpRedirectMode; +use codex_rmcp_client::discover_streamable_http_oauth; +use codex_rmcp_client::perform_oauth_login_return_url; +use pretty_assertions::assert_eq; +use rmcp::transport::auth::AuthError; +use serde_json::json; +use tokio::time::timeout; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +const RESOURCE_AUTHORIZATION: &str = "Bearer resource-only-secret"; +const RESOURCE_API_KEY: &str = "resource-api-key-secret"; +const RESOURCE_USER_AGENT: &str = "resource-only-user-agent"; +const MCP_USER_AGENT: &str = concat!("codex-mcp-client/", env!("CARGO_PKG_VERSION")); +// This is a test safety ceiling, not rmcp's private redirect limit. +const MAX_METADATA_REDIRECT_REQUESTS: u64 = 100; +const REDIRECT_DISCOVERY_TEST_TIMEOUT: Duration = Duration::from_secs(5); + +type DiscoveryResult = anyhow::Result>; + +#[derive(Clone, Copy)] +enum AuthorizationMetadataIssuer { + Matching, + Missing, + Mismatched, +} + +#[derive(Clone, Copy)] +enum MetadataDelivery { + Direct, + SameOriginRedirects, +} + +fn resource_headers() -> Option> { + Some(HashMap::from([ + ( + "Authorization".to_string(), + RESOURCE_AUTHORIZATION.to_string(), + ), + ("X-Api-Key".to_string(), RESOURCE_API_KEY.to_string()), + ("User-Agent".to_string(), RESOURCE_USER_AGENT.to_string()), + ])) +} + +fn local_http_client() -> Arc { + Environment::default_for_tests().get_http_client() +} + +async fn discover_with_local_http_client( + resource_url: &str, + redirect_mode: StreamableHttpRedirectMode, +) -> DiscoveryResult { + discover_streamable_http_oauth( + resource_url, + resource_headers(), + /*env_http_headers*/ None, + local_http_client(), + OAuthDiscoveryTimeout::LOCAL, + redirect_mode, + ) + .await +} + +async fn assert_authorization_requests_exclude_resource_headers( + authorization_server: &MockServer, +) -> anyhow::Result<()> { + let requests = authorization_server + .received_requests() + .await + .context("authorization-server request recording should be enabled")?; + assert!( + !requests.is_empty(), + "OAuth discovery must contact the authorization server" + ); + for request in requests { + assert_eq!(request.headers.get("authorization"), None); + assert_eq!(request.headers.get("x-api-key"), None); + assert_eq!( + request + .headers + .get("user-agent") + .map(wiremock::http::HeaderValue::as_bytes), + Some(MCP_USER_AGENT.as_bytes()) + ); + } + Ok(()) +} + +fn assert_cross_origin_redirect_rejected( + discovery: DiscoveryResult, + redirect_target: &str, +) -> anyhow::Result<()> { + let error = discovery + .err() + .context("cross-origin OAuth metadata redirects must be rejected")?; + assert!( + matches!( + error.downcast_ref::(), + Some(AuthError::MetadataError(reason)) + if reason.contains("OAuth discovery redirect to non-same-origin URL rejected") + && reason.contains(redirect_target) + ), + "expected the cross-origin redirect rejection for `{redirect_target}`: {error:#}", + ); + Ok(()) +} + +async fn assert_legacy_oauth_without_starting_an_mcp_session( + metadata_issuer: AuthorizationMetadataIssuer, + metadata_delivery: MetadataDelivery, +) -> anyhow::Result<()> { + let resource_server = MockServer::start().await; + let authorization_server = MockServer::start().await; + let resource_url = format!("{}/mcp", resource_server.uri()); + let resource_metadata_url = format!("{}/resource-metadata", resource_server.uri()); + + Mock::given(method("GET")) + .and(path("/mcp")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(2) + .mount(&resource_server) + .await; + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(ResponseTemplate::new(500)) + .expect(0) + .mount(&resource_server) + .await; + + let (resource_metadata_path, authorization_metadata_path) = match metadata_delivery { + MetadataDelivery::Direct => ( + "/resource-metadata", + "/.well-known/oauth-authorization-server", + ), + MetadataDelivery::SameOriginRedirects => { + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .respond_with( + ResponseTemplate::new(302) + .insert_header("location", "/redirected-resource-metadata"), + ) + .expect(2) + .mount(&resource_server) + .await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server")) + .respond_with( + ResponseTemplate::new(302) + .insert_header("location", "/redirected-authorization-metadata"), + ) + .expect(2) + .mount(&authorization_server) + .await; + ( + "/redirected-resource-metadata", + "/redirected-authorization-metadata", + ) + } + }; + + Mock::given(method("GET")) + .and(path(resource_metadata_path)) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "resource": resource_url, + "authorization_servers": [authorization_server.uri()], + }))) + .expect(2) + .mount(&resource_server) + .await; + + let mut metadata = json!({ + "authorization_endpoint": format!("{}/authorize", authorization_server.uri()), + "token_endpoint": format!("{}/token", authorization_server.uri()), + "scopes_supported": ["mcp:read"], + "code_challenge_methods_supported": ["S256"], + }); + match metadata_issuer { + AuthorizationMetadataIssuer::Matching => { + metadata["issuer"] = json!(authorization_server.uri()); + } + AuthorizationMetadataIssuer::Missing => {} + AuthorizationMetadataIssuer::Mismatched => { + metadata["issuer"] = json!("https://unexpected-issuer.example"); + } + } + Mock::given(method("GET")) + .and(path(authorization_metadata_path)) + .respond_with(ResponseTemplate::new(200).set_body_json(metadata)) + .expect(2) + .mount(&authorization_server) + .await; + + for redirect_mode in [ + StreamableHttpRedirectMode::Legacy, + StreamableHttpRedirectMode::AgentPluginV1, + ] { + let discovery = discover_with_local_http_client(&resource_url, redirect_mode).await; + match metadata_issuer { + AuthorizationMetadataIssuer::Matching | AuthorizationMetadataIssuer::Missing => { + assert_eq!( + discovery?, + Some(StreamableHttpOAuthDiscovery { + scopes_supported: Some(vec!["mcp:read".to_string()]), + callback_mode: McpOAuthCallbackMode::CallbackSpecific, + }), + ); + } + AuthorizationMetadataIssuer::Mismatched => { + let error = discovery + .err() + .context("a mismatched issuer must not be accepted")?; + assert!( + matches!( + error.downcast_ref::(), + Some(AuthError::MetadataError(reason)) + if reason.contains("issuer does not match authorization metadata origin") + ), + "expected the original authorization-server issuer to remain bound: {error:#}", + ); + } + } + } + + resource_server.verify().await; + authorization_server.verify().await; + assert_authorization_requests_exclude_resource_headers(&authorization_server).await?; + Ok(()) +} + +#[tokio::test] +async fn oauth_discovery_uses_get_first_without_starting_a_legacy_mcp_session() -> anyhow::Result<()> +{ + assert_legacy_oauth_without_starting_an_mcp_session( + AuthorizationMetadataIssuer::Matching, + MetadataDelivery::Direct, + ) + .await +} + +#[tokio::test] +async fn legacy_oauth_discovery_follows_same_origin_metadata_redirects() -> anyhow::Result<()> { + assert_legacy_oauth_without_starting_an_mcp_session( + AuthorizationMetadataIssuer::Matching, + MetadataDelivery::SameOriginRedirects, + ) + .await +} + +#[tokio::test] +async fn legacy_oauth_discovery_rejects_cross_origin_authorization_metadata_redirects() +-> anyhow::Result<()> { + let resource_server = MockServer::start().await; + let authorization_server = MockServer::start().await; + let redirect_target = MockServer::start().await; + let resource_url = format!("{}/mcp", resource_server.uri()); + let resource_metadata_url = format!("{}/resource-metadata", resource_server.uri()); + + Mock::given(method("GET")) + .and(path("/mcp")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(1) + .mount(&resource_server) + .await; + + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "resource": resource_url, + "authorization_servers": [authorization_server.uri()], + }))) + .expect(1) + .mount(&resource_server) + .await; + + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server")) + .respond_with(ResponseTemplate::new(302).insert_header( + "location", + format!( + "{}/redirected-authorization-metadata", + redirect_target.uri() + ), + )) + .expect(1) + .mount(&authorization_server) + .await; + + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(200)) + .expect(0) + .mount(&redirect_target) + .await; + + let discovery = + discover_with_local_http_client(&resource_url, StreamableHttpRedirectMode::Legacy).await; + + assert_cross_origin_redirect_rejected(discovery, &redirect_target.uri())?; + assert!( + redirect_target + .received_requests() + .await + .context("cross-origin request recording should be enabled")? + .is_empty(), + "OAuth authorization-server metadata discovery must not contact a cross-origin redirect target", + ); + redirect_target.verify().await; + authorization_server.verify().await; + resource_server.verify().await; + assert_authorization_requests_exclude_resource_headers(&authorization_server).await?; + Ok(()) +} + +#[tokio::test] +async fn legacy_oauth_discovery_rejects_authorization_metadata_redirect_cycles() +-> anyhow::Result<()> { + let resource_server = MockServer::start().await; + let authorization_server = MockServer::start().await; + let resource_url = format!("{}/mcp", resource_server.uri()); + let resource_metadata_url = format!("{}/resource-metadata", resource_server.uri()); + + Mock::given(method("GET")) + .and(path("/mcp")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(1) + .mount(&resource_server) + .await; + + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "resource": resource_url, + "authorization_servers": [authorization_server.uri()], + }))) + .expect(1) + .mount(&resource_server) + .await; + + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server")) + .respond_with( + ResponseTemplate::new(302) + .insert_header("location", "/.well-known/oauth-authorization-server"), + ) + .expect(2..=MAX_METADATA_REDIRECT_REQUESTS) + .mount(&authorization_server) + .await; + + let discovery = timeout( + REDIRECT_DISCOVERY_TEST_TIMEOUT, + discover_with_local_http_client(&resource_url, StreamableHttpRedirectMode::Legacy), + ) + .await + .context("OAuth metadata redirect cycles must fail within the bounded discovery timeout")?; + let error = discovery + .err() + .context("OAuth metadata redirect cycles must be rejected")?; + assert!( + matches!( + error.downcast_ref::(), + Some(AuthError::MetadataError(reason)) + if reason.contains("OAuth discovery exceeded ") && reason.contains(" redirects") + ), + "expected the SDK to report its bounded OAuth discovery redirect limit: {error:#}", + ); + authorization_server.verify().await; + resource_server.verify().await; + assert_authorization_requests_exclude_resource_headers(&authorization_server).await?; + Ok(()) +} + +#[tokio::test] +async fn legacy_oauth_discovery_rejects_cross_origin_resource_metadata_redirects() +-> anyhow::Result<()> { + let resource_server = MockServer::start().await; + let redirect_target = MockServer::start().await; + let resource_url = format!("{}/mcp", resource_server.uri()); + let resource_metadata_url = format!("{}/resource-metadata", resource_server.uri()); + + Mock::given(method("GET")) + .and(path("/mcp")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(1) + .mount(&resource_server) + .await; + + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .and(header("authorization", RESOURCE_AUTHORIZATION)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .respond_with(ResponseTemplate::new(302).insert_header( + "location", + format!("{}/redirect-target", redirect_target.uri()), + )) + .expect(1) + .mount(&resource_server) + .await; + + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(200)) + .expect(0) + .mount(&redirect_target) + .await; + + let discovery = + discover_with_local_http_client(&resource_url, StreamableHttpRedirectMode::Legacy).await; + + assert_cross_origin_redirect_rejected(discovery, &redirect_target.uri())?; + assert!( + redirect_target + .received_requests() + .await + .context("cross-origin request recording should be enabled")? + .is_empty(), + "OAuth discovery must not contact a cross-origin redirect target", + ); + redirect_target.verify().await; + resource_server.verify().await; + Ok(()) +} + +#[tokio::test] +async fn legacy_oauth_discovery_accepts_authorization_metadata_without_an_issuer() +-> anyhow::Result<()> { + for metadata_delivery in [ + MetadataDelivery::Direct, + MetadataDelivery::SameOriginRedirects, + ] { + assert_legacy_oauth_without_starting_an_mcp_session( + AuthorizationMetadataIssuer::Missing, + metadata_delivery, + ) + .await?; + } + Ok(()) +} + +#[tokio::test] +async fn legacy_oauth_discovery_rejects_an_explicit_mismatched_issuer() -> anyhow::Result<()> { + for metadata_delivery in [ + MetadataDelivery::Direct, + MetadataDelivery::SameOriginRedirects, + ] { + assert_legacy_oauth_without_starting_an_mcp_session( + AuthorizationMetadataIssuer::Mismatched, + metadata_delivery, + ) + .await?; + } + Ok(()) +} + +#[tokio::test] +async fn oauth_discovery_does_not_invent_support_for_an_unauthenticated_legacy_server() +-> anyhow::Result<()> { + let resource_server = MockServer::start().await; + + let server_url = format!("{}/mcp", resource_server.uri()); + let local_discovery = discover_streamable_http_oauth( + &server_url, + /*http_headers*/ None, + /*env_http_headers*/ None, + local_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await?; + + assert_eq!(local_discovery, None); + Ok(()) +} + +#[tokio::test] +async fn interactive_oauth_rejects_untrusted_authorization_metadata() -> anyhow::Result<()> { + for (metadata_issuer, authorization_metadata_path, issuer_bound_callbacks) in [ + ( + AuthorizationMetadataIssuer::Missing, + "/.well-known/untrusted-provider", + false, + ), + ( + AuthorizationMetadataIssuer::Mismatched, + "/.well-known/untrusted-provider", + true, + ), + ( + AuthorizationMetadataIssuer::Mismatched, + "/metadata.json", + false, + ), + ( + AuthorizationMetadataIssuer::Mismatched, + "/metadata.json", + true, + ), + ( + AuthorizationMetadataIssuer::Mismatched, + "/.well-known/oauth-authorization-server", + false, + ), + ( + AuthorizationMetadataIssuer::Mismatched, + "/.well-known/oauth-authorization-server", + true, + ), + ( + AuthorizationMetadataIssuer::Mismatched, + "/.well-known/openid-configuration", + true, + ), + ( + AuthorizationMetadataIssuer::Matching, + "/.well-known/untrusted-provider", + false, + ), + ] { + let resource_server = MockServer::start().await; + let authorization_server = MockServer::start().await; + let attacker_token_server = MockServer::start().await; + let resource_url = format!("{}/mcp", resource_server.uri()); + let resource_metadata_url = format!("{}/resource-metadata", resource_server.uri()); + let (issuer, expected_error) = match metadata_issuer { + AuthorizationMetadataIssuer::Missing => ( + None, + "token endpoint origin does not match the authorization server origin", + ), + AuthorizationMetadataIssuer::Mismatched => ( + Some(authorization_server.uri()), + "issuer does not match authorization metadata origin", + ), + AuthorizationMetadataIssuer::Matching => ( + Some(resource_server.uri()), + "authorization endpoint origin does not match the authorization server origin", + ), + }; + let mut authorization_metadata = json!({ + "authorization_endpoint": format!("{}/authorize", authorization_server.uri()), + "registration_endpoint": format!("{}/register", authorization_server.uri()), + "token_endpoint": format!("{}/token", attacker_token_server.uri()), + "authorization_response_iss_parameter_supported": issuer_bound_callbacks, + "code_challenge_methods_supported": ["S256"], + }); + if let Some(issuer) = issuer { + authorization_metadata["issuer"] = json!(issuer); + } + + Mock::given(method("GET")) + .and(path("/mcp")) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(2) + .mount(&resource_server) + .await; + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "resource": resource_url, + "authorization_servers": [format!( + "{}/.well-known/untrusted-provider", + resource_server.uri() + )], + }))) + .expect(2) + .mount(&resource_server) + .await; + if authorization_metadata_path != "/.well-known/untrusted-provider" { + Mock::given(method("GET")) + .and(path("/.well-known/untrusted-provider")) + .respond_with( + ResponseTemplate::new(302) + .insert_header("location", authorization_metadata_path), + ) + .expect(2) + .mount(&resource_server) + .await; + } + Mock::given(method("GET")) + .and(path(authorization_metadata_path)) + .respond_with(ResponseTemplate::new(200).set_body_json(authorization_metadata)) + .expect(2) + .mount(&resource_server) + .await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200)) + .expect(0) + .mount(&attacker_token_server) + .await; + + for oauth_client_id in [None, Some("preregistered-client")] { + let error = perform_oauth_login_return_url( + "untrusted-oauth-metadata", + &resource_url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + /*http_headers*/ None, + /*env_http_headers*/ None, + /*scopes*/ &[], + oauth_client_id, + McpOAuthClientRegistration::Dcr, + /*oauth_resource*/ None, + Some(/*timeout_secs*/ 5), + /*callback_port*/ None, + /*callback_url*/ None, + /*global_callback_url*/ None, + local_http_client(), + StreamableHttpRedirectMode::Legacy, + ) + .await + .err() + .context("untrusted OAuth authorization metadata must fail")?; + + assert!( + format!("{error:#}").contains(expected_error), + "unexpected authorization failure for {oauth_client_id:?}: {error:#}", + ); + } + + assert!( + attacker_token_server + .received_requests() + .await + .context("attacker token server request recording should be enabled")? + .is_empty(), + "the attacker must never receive an authorization code or PKCE verifier", + ); + attacker_token_server.verify().await; + authorization_server.verify().await; + resource_server.verify().await; + } + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/mcp_2026_sse_discovery.rs b/codex-rs/rmcp-client/tests/mcp_2026_sse_discovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..92910a34a2a7fe0d3f6a2f433f165d4df5755727 --- /dev/null +++ b/codex-rs/rmcp-client/tests/mcp_2026_sse_discovery.rs @@ -0,0 +1,203 @@ +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; + +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use rmcp::model::ClientCapabilities; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use serde_json::Value; +use serde_json::json; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::Request; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +const MODERN_VERSION: &str = "2026-07-28"; +const LEGACY_VERSION: &str = "2025-06-18"; + +async fn initialize_modern_client(server: &MockServer) -> anyhow::Result { + let client = RmcpClient::new_streamable_http_client_with_protocol_mode( + "sse-discovery-test", + &format!("{}/mcp", server.uri()), + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + McpProtocolMode::V20260728, + ) + .await?; + + let params = InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("codex-sse-discovery-test", "0.0.0"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18); + + client + .initialize( + params, + Some(Duration::from_secs(5)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + Ok(client) +} + +fn sse_response(message: Value) -> ResponseTemplate { + ResponseTemplate::new(200).set_body_raw( + format!("event: message\ndata: {message}\n\n"), + "text/event-stream", + ) +} + +#[tokio::test] +async fn modern_sse_discovery_accepts_metadata_namespaced_server_identity() -> anyhow::Result<()> { + let server = MockServer::start().await; + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + assert_eq!(body["method"], "server/discover"); + + sse_response(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "resultType": "complete", + "supportedVersions": [MODERN_VERSION], + "capabilities": {"tools": {}}, + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "modern-sse-test", + "version": "1.0.0", + }, + }, + "ttlMs": 0, + "cacheScope": "private", + }, + })) + }) + .expect(1) + .mount(&server) + .await; + + let client = initialize_modern_client(&server).await?; + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_sse_discovery_falls_back_for_correlated_method_not_found() -> anyhow::Result<()> { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + + match method { + "server/discover" => sse_response(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "error": {"code": -32601, "message": "method not found"}, + })), + "initialize" => { + assert_eq!(body["params"]["protocolVersion"], LEGACY_VERSION); + ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "protocolVersion": LEGACY_VERSION, + "capabilities": {"tools": {}}, + "serverInfo": { + "name": "legacy-sse-test", + "version": "1.0.0", + }, + }, + })) + } + "notifications/initialized" => ResponseTemplate::new(202), + other => panic!("unexpected legacy SSE fallback request: {other}"), + } + }) + .expect(3) + .mount(&server) + .await; + + let client = initialize_modern_client(&server).await?; + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["server/discover", "initialize", "notifications/initialized"] + ); + client.shutdown().await; + Ok(()) +} + +#[tokio::test] +async fn modern_sse_discovery_rejects_uncorrelated_method_not_found() -> anyhow::Result<()> { + for (case, rejected_id) in [("null", json!(null)), ("mismatched", json!("unrelated"))] { + let server = MockServer::start().await; + let observed = Arc::new(Mutex::new(Vec::::new())); + let recorded = Arc::clone(&observed); + + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + let method = body["method"].as_str().expect("JSON-RPC method"); + recorded.lock().expect("requests lock").push(method.into()); + assert_eq!(method, "server/discover"); + + sse_response(json!({ + "jsonrpc": "2.0", + "id": rejected_id.clone(), + "error": {"code": -32601, "message": "method not found"}, + })) + }) + .expect(1) + .mount(&server) + .await; + + assert!( + initialize_modern_client(&server).await.is_err(), + "{case} SSE discovery error must not downgrade to legacy" + ); + assert_eq!( + *observed.lock().expect("requests lock"), + vec!["server/discover"], + "{case} SSE discovery error must not initialize a legacy session" + ); + } + + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/mcp_2026_stdio.rs b/codex-rs/rmcp-client/tests/mcp_2026_stdio.rs new file mode 100644 index 0000000000000000000000000000000000000000..53ba7fbb9588e36785156480b733639448b86bc0 --- /dev/null +++ b/codex-rs/rmcp-client/tests/mcp_2026_stdio.rs @@ -0,0 +1,256 @@ +use std::collections::HashMap; +use std::ffi::OsString; +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_exec_server::Environment; +use codex_rmcp_client::Elicitation; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::ExecutorStdioServerLauncher; +use codex_rmcp_client::LocalStdioServerLauncher; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::StdioServerLauncher; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use rmcp::model::ClientCapabilities; +use rmcp::model::ElicitRequestParams; +use rmcp::model::ElicitationCapability; +use rmcp::model::FormElicitationCapability; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use serde_json::json; + +async fn exercise_stdio_server( + server_mode: &str, + protocol_mode: McpProtocolMode, + opt_in: bool, + use_executor: bool, +) -> anyhow::Result<()> { + let server = codex_utils_cargo_bin::cargo_bin("test_mcp_2026_stdio_server")?; + let launcher: Arc = if use_executor { + Arc::new(ExecutorStdioServerLauncher::new( + Environment::default_for_tests().get_exec_backend(), + )) + } else { + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)) + }; + let mut env = HashMap::new(); + if opt_in { + env.insert( + OsString::from("CODEX_MCP_PROTOCOL_VERSION"), + OsString::from("2026-07-28"), + ); + } + let cwd = std::env::current_dir()?; + #[cfg(unix)] + let root = tempfile::tempdir()?; + #[cfg(unix)] + let (server, cwd) = if use_executor { + (server, cwd) + } else { + use std::os::unix::fs::PermissionsExt; + let wrapper = root.path().join("server"); + std::fs::write( + &wrapper, + "#!/bin/sh\nprintf '%s' \"$0\" > argv0\nexec \"$MCP_SERVER\" \"$@\"\n", + )?; + std::fs::set_permissions(wrapper, std::fs::Permissions::from_mode(0o755))?; + env.insert(OsString::from("MCP_SERVER"), server.into_os_string()); + ( + std::path::PathBuf::from("./server"), + root.path().to_path_buf(), + ) + }; + let client = RmcpClient::new_stdio_client_with_protocol_mode( + server.into(), + vec![OsString::from(server_mode)], + Some(env), + &[], + Some(cwd.to_string_lossy().into_owned()), + launcher, + protocol_mode, + ) + .await?; + + let mut capabilities = ClientCapabilities::default(); + capabilities.elicitation = + Some(ElicitationCapability::new().with_form(FormElicitationCapability::new())); + let elicitation_count = Arc::new(AtomicUsize::new(0)); + let observed_elicitations = Arc::clone(&elicitation_count); + let legacy_session = server_mode.starts_with("legacy"); + let initialized = client + .initialize( + InitializeRequestParams::new( + capabilities, + Implementation::new("codex-mcp-client", "0.0.0-test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(5)), + Box::new(move |_request_id, request| { + let observed_elicitations = Arc::clone(&observed_elicitations); + async move { + observed_elicitations.fetch_add(1, Ordering::Relaxed); + let Elicitation::Mcp(ElicitRequestParams::FormElicitationParams { + requested_schema, + .. + }) = request + else { + anyhow::bail!("expected a standard MCP form elicitation"); + }; + let content = if legacy_session { + assert_eq!( + serde_json::to_value(requested_schema)?, + json!({ + "type": "object", + "properties": { + "name": {"type": "string", "default": "John Doe"}, + "age": {"type": "integer", "default": 30}, + "score": {"type": "number", "default": 95.5}, + "status": { + "type": "string", + "enum": ["active", "inactive"], + "default": "active", + }, + "verified": {"type": "boolean", "default": true}, + }, + "required": [], + }), + ); + json!({ + "name": "John Doe", + "age": 30, + "score": 95.5, + "status": "active", + "verified": true, + }) + } else { + json!({"approved": true}) + }; + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(content), + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + let expected_version = if legacy_session { + ProtocolVersion::V_2025_06_18 + } else { + ProtocolVersion::V_2026_07_28 + }; + assert_eq!(initialized.protocol_version, expected_version); + assert_eq!( + initialized + .server_info + .as_ref() + .map(|server_info| server_info.name.as_str()), + Some(if legacy_session { + "legacy-stdio-test" + } else { + "strict-stdio-test" + }) + ); + let tools = client + .list_tools(/*params*/ None, Some(Duration::from_secs(5))) + .await?; + assert_eq!( + tools + .tools + .iter() + .map(|tool| tool.name.as_ref()) + .collect::>(), + vec!["echo"] + ); + let result = client + .call_tool( + "echo".to_owned(), + Some(json!({"message": "hello stdio"})), + /*meta*/ None, + Some(Duration::from_secs(5)), + ) + .await?; + assert_eq!( + result.content[0].as_text().map(|text| text.text.as_str()), + Some(if legacy_session { + "legacy approved" + } else { + "modern approved" + }) + ); + assert_eq!(elicitation_count.load(Ordering::Relaxed), 1); + #[cfg(unix)] + if !use_executor { + assert_eq!( + std::fs::read_to_string(root.path().join("argv0"))?, + "./server" + ); + } + client.shutdown().await; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn modern_local_stdio_discovers_metadata_only_identity_and_drives_mrtr() -> anyhow::Result<()> +{ + exercise_stdio_server( + "modern", + McpProtocolMode::V20260728, + /*opt_in*/ true, + /*use_executor*/ false, + ) + .await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn modern_executor_stdio_discovers_metadata_only_identity_and_drives_mrtr() +-> anyhow::Result<()> { + exercise_stdio_server( + "modern", + McpProtocolMode::V20260728, + /*opt_in*/ true, + /*use_executor*/ true, + ) + .await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn legacy_stdio_supports_sep1034_defaults_without_the_rollout_flag() -> anyhow::Result<()> { + exercise_stdio_server( + "legacy", + McpProtocolMode::Legacy, + /*opt_in*/ false, + /*use_executor*/ false, + ) + .await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn rollout_flag_alone_preserves_legacy_stdio_and_sep1034_defaults() -> anyhow::Result<()> { + exercise_stdio_server( + "legacy", + McpProtocolMode::V20260728, + /*opt_in*/ false, + /*use_executor*/ false, + ) + .await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn modern_stdio_safely_falls_back_to_legacy_elicitation() -> anyhow::Result<()> { + exercise_stdio_server( + "legacy-fallback", + McpProtocolMode::V20260728, + /*opt_in*/ true, + /*use_executor*/ false, + ) + .await +} diff --git a/codex-rs/rmcp-client/tests/mcp_2026_stdio_discovery.rs b/codex-rs/rmcp-client/tests/mcp_2026_stdio_discovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..8cf62a2697cf505946e283ae3a6c975532ab9d38 --- /dev/null +++ b/codex-rs/rmcp-client/tests/mcp_2026_stdio_discovery.rs @@ -0,0 +1,285 @@ +use std::collections::HashMap; +use std::ffi::OsString; +use std::process::Command; +use std::sync::Arc; +use std::time::Duration; + +use codex_exec_server::Environment; +use codex_network_proxy::CUSTOM_CA_ENV_KEYS; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::ExecutorStdioServerLauncher; +use codex_rmcp_client::LocalStdioServerLauncher; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::StdioServerLauncher; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use rmcp::model::ClientCapabilities; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use serde_json::json; + +#[test] +fn local_stdio_inherits_ca_certificate_variables() -> anyhow::Result<()> { + let directories = tempfile::tempdir()?; + let source_dir = directories.path().join("source"); + let server_dir = directories.path().join("server"); + std::fs::create_dir_all(source_dir.join("certs"))?; + std::fs::create_dir_all(&server_dir)?; + + #[cfg(windows)] + let requests_ca_bundle = match source_dir.components().next() { + Some(std::path::Component::Prefix(prefix)) => match prefix.kind() { + std::path::Prefix::Disk(drive) | std::path::Prefix::VerbatimDisk(drive) => { + format!("{}:certs\\custom-ca.pem", char::from(drive)) + } + _ => "certs/custom-ca.pem".to_string(), + }, + _ => "certs/custom-ca.pem".to_string(), + }; + #[cfg(not(windows))] + let requests_ca_bundle = "certs/custom-ca.pem"; + + let output = Command::new(std::env::current_exe()?) + .arg("--exact") + .arg("local_stdio_inherits_ca_certificate_variables_child") + .arg("--ignored") + .current_dir(&source_dir) + .envs( + CUSTOM_CA_ENV_KEYS + .into_iter() + .map(|name| (name, "certs/custom-ca.pem")), + ) + .env("CODEX_CA_CERTIFICATE", "") + .env("REQUESTS_CA_BUNDLE", requests_ca_bundle) + .env("CODEX_MCP_TEST_SERVER_CWD", &server_dir) + .output()?; + + assert!( + output.status.success(), + "MCP subprocess inheritance test failed:\n{}\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "child process for local_stdio_inherits_ca_certificate_variables"] +async fn local_stdio_inherits_ca_certificate_variables_child() -> anyhow::Result<()> { + let server = codex_utils_cargo_bin::cargo_bin("test_stdio_server")?; + let source_dir = std::env::current_dir()?; + let server_dir = std::env::var("CODEX_MCP_TEST_SERVER_CWD")?; + let expected = source_dir.join("certs").join("custom-ca.pem"); + let npm_override = HashMap::from([( + OsString::from("NPM_CONFIG_CAFILE"), + OsString::from("explicit/custom-ca.pem"), + )]); + + for (overrides, env_name, expected) in [ + ( + None, + "SSL_CERT_FILE", + Some(expected.to_string_lossy().into_owned()), + ), + ( + None, + "REQUESTS_CA_BUNDLE", + Some(expected.to_string_lossy().into_owned()), + ), + (None, "CODEX_CA_CERTIFICATE", None), + ( + Some(npm_override.clone()), + "NPM_CONFIG_CAFILE", + Some("explicit/custom-ca.pem".to_string()), + ), + (Some(npm_override), "npm_config_cafile", None), + ] { + let client = RmcpClient::new_stdio_client( + server.clone().into(), + Vec::new(), + overrides, + &[], + Some(server_dir.clone()), + Arc::new(LocalStdioServerLauncher::new(source_dir.clone())), + ) + .await?; + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("stdio-ca-inheritance-test", "1.0.0"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(10)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }) + } + .boxed() + }), + ) + .await?; + let result = client + .call_tool( + "echo".to_string(), + Some(json!({ "message": "ca inheritance", "env_var": env_name })), + /*meta*/ None, + Some(Duration::from_secs(10)), + ) + .await?; + assert_eq!( + result.structured_content, + Some(json!({ "echo": "ECHOING: ca inheritance", "env": expected })) + ); + client.shutdown().await; + } + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn modern_local_and_executor_stdio_discover_metadata_identity_and_catalogs() +-> anyhow::Result<()> { + let server = codex_utils_cargo_bin::cargo_bin("test_mcp_2026_discovery_stdio_server")?; + + for executor in [false, true] { + let launcher: Arc = if executor { + Arc::new(ExecutorStdioServerLauncher::new( + Environment::default_for_tests().get_exec_backend(), + )) + } else { + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)) + }; + let client = RmcpClient::new_stdio_client_with_protocol_mode( + server.clone().into(), + Vec::new(), + Some(HashMap::from([( + OsString::from("CODEX_MCP_PROTOCOL_VERSION"), + OsString::from("2026-07-28"), + )])), + &[], + Some(std::env::current_dir()?.to_string_lossy().into_owned()), + launcher, + McpProtocolMode::V20260728, + ) + .await?; + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("stdio-discovery-test", "1.0.0"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(10)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + let tools = client + .list_tools(/*params*/ None, Some(Duration::from_secs(10))) + .await?; + assert_eq!(tools.tools[0].name.as_ref(), "stdio_echo"); + let resources = client + .list_resources(/*params*/ None, Some(Duration::from_secs(10))) + .await?; + assert_eq!(resources.resources[0].uri, "test://stdio/resource"); + let templates = client + .list_resource_templates(/*params*/ None, Some(Duration::from_secs(10))) + .await?; + assert_eq!( + templates.resource_templates[0].name, + "stdio resource template" + ); + client.shutdown().await; + } + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn legacy_stdio_preserves_existing_protocol_marker_environment() -> anyhow::Result<()> { + let server = codex_utils_cargo_bin::cargo_bin("test_stdio_server")?; + + for executor in [false, true] { + for version in ["2026-07-28", "1999-01-01"] { + let launcher: Arc = if executor { + Arc::new(ExecutorStdioServerLauncher::new( + Environment::default_for_tests().get_exec_backend(), + )) + } else { + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)) + }; + let client = RmcpClient::new_stdio_client_with_protocol_mode( + server.clone().into(), + Vec::new(), + Some(HashMap::from([( + OsString::from("CODEX_MCP_PROTOCOL_VERSION"), + OsString::from(version), + )])), + &[], + Some(std::env::current_dir()?.to_string_lossy().into_owned()), + launcher, + McpProtocolMode::Legacy, + ) + .await?; + + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("stdio-legacy-environment-test", "1.0.0"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(10)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + let result = client + .call_tool( + "echo".to_string(), + Some(json!({ + "message": "legacy environment", + "env_var": "CODEX_MCP_PROTOCOL_VERSION", + })), + /*meta*/ None, + Some(Duration::from_secs(10)), + ) + .await?; + assert_eq!( + result.structured_content, + Some(json!({ + "echo": "ECHOING: legacy environment", + "env": version, + })) + ); + client.shutdown().await; + } + } + + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/mcp_events.rs b/codex-rs/rmcp-client/tests/mcp_events.rs new file mode 100644 index 0000000000000000000000000000000000000000..3b015a01d274c90109f490b6314755240664fc35 --- /dev/null +++ b/codex-rs/rmcp-client/tests/mcp_events.rs @@ -0,0 +1,242 @@ +mod streamable_http_test_support; + +use std::convert::Infallible; +use std::time::Duration; + +use anyhow::Context as _; +use axum::Json; +use axum::Router; +use axum::http::StatusCode; +use axum::response::IntoResponse; +use axum::response::sse::Event; +use axum::response::sse::Sse; +use axum::routing::post; +use futures::StreamExt as _; +use futures::stream; +use pretty_assertions::assert_eq; +use serde_json::Value; +use serde_json::json; +use tokio::net::TcpListener; +use tokio::sync::Notify; +use tokio::sync::mpsc; +use tokio::time::timeout; + +use streamable_http_test_support::create_client; + +struct StreamClosed { + event_name: String, + closed: mpsc::UnboundedSender, +} + +impl Drop for StreamClosed { + fn drop(&mut self) { + let _ = self.closed.send(self.event_name.clone()); + } +} + +fn initialize_response(message: &Value) -> axum::response::Response { + Json(json!({ + "jsonrpc": "2.0", + "id": message["id"], + "result": { + "protocolVersion": "2025-06-18", + "capabilities": {}, + "serverInfo": {"name": "plugin-runtime", "version": "1.0.0"}, + }, + })) + .into_response() +} + +fn plugin_runtime_notifications(event_name: &str, request_id: Value) -> [Value; 2] { + let metadata = json!({"io.modelcontextprotocol/subscriptionId": request_id}); + [ + json!({ + "jsonrpc": "2.0", + "method": "notifications/events/active", + "params": {"cursor": null, "truncated": false, "_meta": metadata}, + }), + json!({ + "jsonrpc": "2.0", + "method": "notifications/events/event", + "params": { + "eventId": format!("event-{event_name}"), + "name": event_name, + "timestamp": "2026-08-07T12:00:00Z", + "data": {"issue": 42}, + "cursor": null, + "_meta": metadata, + }, + }), + ] +} + +#[tokio::test] +async fn plugin_runtime_event_streams_are_isolated_and_cancel_locally() -> anyhow::Result<()> { + let (stream_closed_tx, mut stream_closed_rx) = mpsc::unbounded_channel::(); + let router = Router::new().route( + "/mcp", + post(move |Json(message): Json| { + let stream_closed_tx = stream_closed_tx.clone(); + async move { + match message["method"].as_str() { + Some("initialize") => initialize_response(&message), + Some("notifications/initialized") => StatusCode::ACCEPTED.into_response(), + Some("events/stream") => { + let event_name = message["params"]["name"] + .as_str() + .expect("event stream request must contain a name") + .to_string(); + assert_eq!( + message.pointer("/params/arguments"), + Some(&json!({"project": "codex"})) + ); + assert!(message["params"].get("cursor").is_none()); + + let closed = StreamClosed { + event_name: event_name.clone(), + closed: stream_closed_tx, + }; + let events = stream::iter( + plugin_runtime_notifications(&event_name, message["id"].clone()) + .into_iter() + .map(|notification| { + Ok::<_, Infallible>( + Event::default() + .event("message") + .data(notification.to_string()), + ) + }), + ) + .chain(stream::pending()) + .map(move |event| { + let _ = &closed; + event + }); + + Sse::new(events).into_response() + } + Some("notifications/cancelled") => { + panic!("event cancellation should close the local stream without a POST") + } + method => panic!("unexpected Plugin Runtime request: {method:?}"), + } + } + }), + ); + + let listener = TcpListener::bind("127.0.0.1:0").await?; + let base_url = format!("http://{}", listener.local_addr()?); + let server = tokio::spawn(async move { axum::serve(listener, router).await }); + let client = create_client(&base_url).await?; + + let mut first = client + .send_event_stream_request(Some(json!({ + "name": "github.pull_request.opened", + "arguments": {"project": "codex"}, + }))) + .await?; + let mut second = client + .send_event_stream_request(Some(json!({ + "name": "gmail.message.received", + "arguments": {"project": "codex"}, + }))) + .await?; + + for (event_name, request) in [ + ("github.pull_request.opened", &mut first), + ("gmail.message.received", &mut second), + ] { + let active = timeout(Duration::from_secs(5), request.notifications.recv()) + .await? + .context("event stream closed before activation")?; + assert_eq!(active.method, "notifications/events/active"); + assert_eq!( + active.params, + Some(json!({"cursor": null, "truncated": false})) + ); + + let event = timeout(Duration::from_secs(5), request.notifications.recv()) + .await? + .context("event stream closed before delivery")?; + assert_eq!(event.method, "notifications/events/event"); + assert_eq!( + event.params.as_ref().and_then(|params| params.get("name")), + Some(&json!(event_name)) + ); + } + + first + .handle + .cancel(Some("event subscription closed".to_string())) + .await?; + let closed = timeout(Duration::from_secs(5), stream_closed_rx.recv()) + .await? + .context("cancelled stream did not close")?; + assert_eq!(closed, "github.pull_request.opened"); + + second + .handle + .cancel(Some("event subscription closed".to_string())) + .await?; + let closed = timeout(Duration::from_secs(5), stream_closed_rx.recv()) + .await? + .context("cancelled stream did not close")?; + assert_eq!(closed, "gmail.message.received"); + + client.shutdown().await; + server.abort(); + let _ = server.await; + Ok(()) +} + +#[tokio::test] +async fn plugin_runtime_event_stream_times_out_before_stalled_headers() -> anyhow::Result<()> { + let request_started = std::sync::Arc::new(Notify::new()); + let server_request_started = std::sync::Arc::clone(&request_started); + let router = Router::new().route( + "/mcp", + post(move |Json(message): Json| { + let request_started = std::sync::Arc::clone(&server_request_started); + async move { + match message["method"].as_str() { + Some("initialize") => initialize_response(&message), + Some("notifications/initialized") => StatusCode::ACCEPTED.into_response(), + Some("events/stream") => { + request_started.notify_one(); + std::future::pending().await + } + method => panic!("unexpected Plugin Runtime request: {method:?}"), + } + } + }), + ); + + let listener = TcpListener::bind("127.0.0.1:0").await?; + let base_url = format!("http://{}", listener.local_addr()?); + let server = tokio::spawn(async move { axum::serve(listener, router).await }); + let client = create_client(&base_url).await?; + let request = client + .send_event_stream_request(Some(json!({ + "name": "github.pull_request.opened", + "arguments": {}, + }))) + .await?; + + timeout(Duration::from_secs(5), request_started.notified()) + .await + .context("event stream request did not reach the server")?; + tokio::time::pause(); + tokio::time::advance(Duration::from_secs(31)).await; + let error = request + .handle + .rx + .await? + .expect_err("stalled response headers must time out"); + tokio::time::resume(); + assert!(error.to_string().contains("timed out")); + + client.shutdown().await; + server.abort(); + let _ = server.await; + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/process_group_cleanup.rs b/codex-rs/rmcp-client/tests/process_group_cleanup.rs new file mode 100644 index 0000000000000000000000000000000000000000..1dcd5586ef5f73c28a6d028a59172b8c7b003d47 --- /dev/null +++ b/codex-rs/rmcp-client/tests/process_group_cleanup.rs @@ -0,0 +1,182 @@ +#![cfg(unix)] + +use std::collections::HashMap; +use std::ffi::OsString; +use std::fs; +use std::os::unix::fs::PermissionsExt; +use std::path::Path; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::LocalStdioServerLauncher; +use codex_rmcp_client::RmcpClient; +use futures::FutureExt as _; +use rmcp::model::ClientCapabilities; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use serde_json::json; + +fn stdio_server_bin() -> Result { + codex_utils_cargo_bin::cargo_bin("test_stdio_server").map_err(Into::into) +} + +fn init_params() -> InitializeRequestParams { + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("codex-test", "0.0.0-test").with_title("Codex rmcp shutdown test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18) +} + +fn process_exists(pid: u32) -> bool { + std::process::Command::new("kill") + .arg("-0") + .arg(pid.to_string()) + .stderr(std::process::Stdio::null()) + .status() + .map(|status| status.success()) + .unwrap_or(false) +} + +async fn wait_for_pid_file(path: &Path) -> Result { + for _ in 0..50 { + match fs::read_to_string(path) { + Ok(content) => { + let trimmed = content.trim(); + if trimmed.is_empty() { + tokio::time::sleep(Duration::from_millis(100)).await; + continue; + } + + let pid = trimmed + .parse::() + .with_context(|| format!("failed to parse pid from {}", path.display()))?; + return Ok(pid); + } + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + tokio::time::sleep(Duration::from_millis(100)).await; + } + Err(error) => { + return Err(error).with_context(|| format!("failed to read {}", path.display())); + } + } + } + + anyhow::bail!("timed out waiting for child pid file at {}", path.display()); +} + +async fn wait_for_process_exit(pid: u32) -> Result<()> { + for _ in 0..50 { + if !process_exists(pid) { + return Ok(()); + } + tokio::time::sleep(Duration::from_millis(100)).await; + } + + anyhow::bail!("process {pid} still running after timeout"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn drop_kills_wrapper_process_group() -> Result<()> { + let temp_dir = tempfile::tempdir()?; + let child_pid_file = temp_dir.path().join("child.pid"); + let child_pid_file_str = child_pid_file.to_string_lossy().into_owned(); + let wrapper = temp_dir.path().join("wrapper"); + fs::write( + &wrapper, + "#!/bin/sh\nsleep 300 & child_pid=$!; echo \"$child_pid\" > \"$CHILD_PID_FILE\"; cat >/dev/null\n", + )?; + fs::set_permissions(wrapper, fs::Permissions::from_mode(0o755))?; + + let client = RmcpClient::new_stdio_client( + OsString::from("./wrapper"), + vec![], + Some(HashMap::from([( + OsString::from("CHILD_PID_FILE"), + OsString::from(child_pid_file_str), + )])), + &[], + Some(temp_dir.path().to_string_lossy().into_owned()), + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)), + ) + .await?; + + let grandchild_pid = wait_for_pid_file(&child_pid_file).await?; + assert!( + process_exists(grandchild_pid), + "expected grandchild process {grandchild_pid} to be running before dropping client" + ); + + drop(client); + + wait_for_process_exit(grandchild_pid).await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn shutdown_kills_initialized_stdio_server_with_in_flight_operation() -> Result<()> { + let temp_dir = tempfile::tempdir()?; + let server_pid_file = temp_dir.path().join("server.pid"); + let server_pid_file_str = server_pid_file.to_string_lossy().into_owned(); + + let client = Arc::new( + RmcpClient::new_stdio_client( + stdio_server_bin()?.into(), + Vec::::new(), + Some(HashMap::from([( + OsString::from("MCP_TEST_PID_FILE"), + OsString::from(server_pid_file_str), + )])), + &[], + /*cwd*/ None, + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)), + ) + .await?, + ); + + client + .initialize( + init_params(), + Some(Duration::from_secs(5)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + let server_pid = wait_for_pid_file(&server_pid_file).await?; + assert!( + process_exists(server_pid), + "expected MCP server process {server_pid} to be running before shutdown" + ); + + let call_client = Arc::clone(&client); + let call_task = tokio::spawn(async move { + call_client + .call_tool( + "sync".to_string(), + Some(json!({ "sleep_after_ms": 300_000 })), + /*meta*/ None, + Some(Duration::from_secs(300)), + ) + .await + }); + tokio::time::sleep(Duration::from_millis(200)).await; + + client.shutdown().await; + + wait_for_process_exit(server_pid).await?; + let _ = tokio::time::timeout(Duration::from_secs(5), call_task).await?; + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/resources.rs b/codex-rs/rmcp-client/tests/resources.rs new file mode 100644 index 0000000000000000000000000000000000000000..da280ee29145613f9699a3371150d8b31cb440e0 --- /dev/null +++ b/codex-rs/rmcp-client/tests/resources.rs @@ -0,0 +1,144 @@ +use std::ffi::OsString; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::Duration; + +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::LocalStdioServerLauncher; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::mcp_error; +use codex_utils_cargo_bin::CargoBinError; +use futures::FutureExt as _; +use pretty_assertions::assert_eq; +use rmcp::model::ClientCapabilities; +use rmcp::model::ElicitationCapability; +use rmcp::model::FormElicitationCapability; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ListResourceTemplatesResult; +use rmcp::model::ProtocolVersion; +use rmcp::model::ReadResourceRequestParams; +use rmcp::model::ResourceContents; +use serde_json::json; + +const RESOURCE_URI: &str = "memo://codex/example-note"; + +fn stdio_server_bin() -> Result { + codex_utils_cargo_bin::cargo_bin("test_stdio_server") +} + +fn init_params() -> InitializeRequestParams { + let mut capabilities = ClientCapabilities::default(); + capabilities.elicitation = + Some(ElicitationCapability::new().with_form(FormElicitationCapability::new())); + InitializeRequestParams::new( + capabilities, + Implementation::new("codex-test", "0.0.0-test").with_title("Codex rmcp resource test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18) +} + +async fn resource_client() -> anyhow::Result { + let client = RmcpClient::new_stdio_client( + stdio_server_bin()?.into(), + Vec::::new(), + /*env*/ None, + &[], + /*cwd*/ None, + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)), + ) + .await?; + + client + .initialize( + init_params(), + Some(Duration::from_secs(5)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + Ok(client) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn rmcp_client_can_list_and_read_resources() -> anyhow::Result<()> { + let client = resource_client().await?; + let list = client + .list_resources(/*params*/ None, Some(Duration::from_secs(5))) + .await?; + let memo = list + .resources + .iter() + .find(|resource| resource.uri == RESOURCE_URI) + .expect("memo resource present"); + assert_eq!( + memo, + &rmcp::model::Resource::new(RESOURCE_URI, "example-note") + .with_title("Example Note") + .with_description("A sample MCP resource exposed for integration tests.") + .with_mime_type("text/plain") + ); + let templates = client + .list_resource_templates(/*params*/ None, Some(Duration::from_secs(5))) + .await?; + let mut expected_templates = ListResourceTemplatesResult::with_all_items(vec![ + rmcp::model::ResourceTemplate::new("memo://codex/{slug}", "codex-memo") + .with_title("Codex Memo") + .with_description("Template for memo://codex/{slug} resources used in tests.") + .with_mime_type("text/plain"), + ]); + expected_templates.result_type = None; + assert_eq!(templates, expected_templates); + + let read = client + .read_resource( + ReadResourceRequestParams::new(RESOURCE_URI), + Some(Duration::from_secs(5)), + ) + .await?; + let text = read.contents.first().expect("resource contents present"); + assert_eq!( + text, + &ResourceContents::TextResourceContents { + uri: RESOURCE_URI.to_string(), + mime_type: Some("text/plain".to_string()), + text: "This is a sample MCP resource served by the rmcp test server.".to_string(), + meta: None, + } + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn rmcp_client_preserves_each_resource_error() -> anyhow::Result<()> { + let client = resource_client().await?; + for uri in ["memo://codex/missing-first", "memo://codex/missing-second"] { + let error = client + .read_resource( + ReadResourceRequestParams::new(uri), + Some(Duration::from_secs(5)), + ) + .await + .expect_err("missing resource must return a protocol error") + .context("resources/read failed"); + assert_eq!( + mcp_error(&error), + Some(&rmcp::ErrorData::resource_not_found( + "resource_not_found", + Some(json!({ "uri": uri })), + )) + ); + } + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/stdio_message_limits.rs b/codex-rs/rmcp-client/tests/stdio_message_limits.rs new file mode 100644 index 0000000000000000000000000000000000000000..7c8ec6102bbb4a3238fb8e61d697cb09f234b77f --- /dev/null +++ b/codex-rs/rmcp-client/tests/stdio_message_limits.rs @@ -0,0 +1,226 @@ +use std::collections::HashMap; +use std::ffi::OsString; +#[cfg(windows)] +use std::os::windows::io::AsRawHandle; +#[cfg(windows)] +use std::os::windows::io::FromRawHandle; +#[cfg(windows)] +use std::os::windows::io::OwnedHandle; +use std::sync::Arc; +use std::time::Duration; + +use codex_exec_server::Environment; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::ExecutorStdioServerLauncher; +use codex_rmcp_client::LocalStdioServerLauncher; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::StdioServerLauncher; +use futures::FutureExt; +use rmcp::model::ClientCapabilities; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; + +#[cfg(windows)] +#[link(name = "kernel32")] +unsafe extern "system" { + fn OpenProcess( + desired_access: u32, + inherit_handle: i32, + process_id: u32, + ) -> *mut std::ffi::c_void; + fn TerminateProcess(handle: *mut std::ffi::c_void, exit_code: u32) -> i32; + fn WaitForSingleObject(handle: *mut std::ffi::c_void, milliseconds: u32) -> u32; +} + +#[cfg(windows)] +fn open_process_for_wait(process_id: u32) -> std::io::Result { + let handle = unsafe { + OpenProcess( + /*desired_access*/ 0x0010_0001, + /*inherit_handle*/ 0, + process_id, + ) + }; + if handle.is_null() { + return Err(std::io::Error::last_os_error()); + } + Ok(unsafe { OwnedHandle::from_raw_handle(handle.cast()) }) +} + +#[cfg(windows)] +fn wait_for_process_exit(process: &OwnedHandle) -> std::io::Result<()> { + match unsafe { + WaitForSingleObject(process.as_raw_handle().cast(), /*milliseconds*/ 5_000) + } { + 0 => Ok(()), + u32::MAX => Err(std::io::Error::last_os_error()), + _ => Err(std::io::Error::new( + std::io::ErrorKind::TimedOut, + "process did not exit", + )), + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn stdio_message_limits_preserve_legacy_local_compatibility() -> anyhow::Result<()> { + let server = codex_utils_cargo_bin::cargo_bin("test_stdio_server")?; + + for (executor, protocol_mode, accepts_oversized) in [ + (false, McpProtocolMode::Legacy, true), + (false, McpProtocolMode::V20260728, false), + (true, McpProtocolMode::Legacy, false), + ] { + let launcher: Arc = if executor { + Arc::new(ExecutorStdioServerLauncher::new( + Environment::default_for_tests().get_exec_backend(), + )) + } else { + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)) + }; + let mut env = HashMap::from([( + OsString::from("MCP_TEST_OVERSIZED_TOOL_DESCRIPTION"), + OsString::from("1"), + )]); + if protocol_mode == McpProtocolMode::V20260728 { + env.insert( + OsString::from("CODEX_MCP_PROTOCOL_VERSION"), + OsString::from("2026-07-28"), + ); + } + let client = RmcpClient::new_stdio_client_with_protocol_mode( + server.clone().into(), + Vec::new(), + Some(env), + &[], + Some(std::env::current_dir()?.to_string_lossy().into_owned()), + launcher, + protocol_mode, + ) + .await?; + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("stdio-limit-test", "1.0.0"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(10)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + let result = client + .list_tools(/*params*/ None, Some(Duration::from_secs(10))) + .await; + assert_eq!( + result.is_ok(), + accepts_oversized, + "unexpected stdio size handling (executor={executor}, mode={protocol_mode:?})" + ); + client.shutdown().await; + } + Ok(()) +} + +#[cfg(windows)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn local_stdio_shutdown_terminates_descendants_after_server_exit() -> anyhow::Result<()> { + let server = codex_utils_cargo_bin::cargo_bin("test_stdio_server")?; + + for protocol_mode in [McpProtocolMode::Legacy, McpProtocolMode::V20260728] { + let temp_dir = tempfile::tempdir()?; + let server_pid_file = temp_dir.path().join("server.pid"); + let descendant_pid_file = temp_dir.path().join("descendant.pid"); + let breakaway_denied_file = temp_dir.path().join("breakaway.denied"); + let mut env = HashMap::from([ + ( + OsString::from("MCP_TEST_PID_FILE"), + server_pid_file.clone().into(), + ), + ( + OsString::from("MCP_TEST_DESCENDANT_PID_FILE"), + descendant_pid_file.clone().into(), + ), + ( + OsString::from("MCP_TEST_BREAKAWAY_DENIED_FILE"), + breakaway_denied_file.clone().into(), + ), + ]); + if protocol_mode == McpProtocolMode::V20260728 { + env.insert( + OsString::from("CODEX_MCP_PROTOCOL_VERSION"), + OsString::from("2026-07-28"), + ); + } + + let client = RmcpClient::new_stdio_client_with_protocol_mode( + server.clone().into(), + Vec::new(), + Some(env), + &[], + Some(std::env::current_dir()?.to_string_lossy().into_owned()), + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)), + protocol_mode, + ) + .await?; + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("stdio-cleanup-test", "1.0.0"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(10)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + assert_eq!( + std::fs::read_to_string(breakaway_denied_file)?, + "denied", + "MCP descendant escaped its Windows job (mode={protocol_mode:?})" + ); + let server_pid = std::fs::read_to_string(server_pid_file)? + .trim() + .parse::()?; + let descendant_pid = std::fs::read_to_string(descendant_pid_file)? + .trim() + .parse::()?; + let server_process = open_process_for_wait(server_pid)?; + let descendant_process = open_process_for_wait(descendant_pid)?; + let terminated = unsafe { + TerminateProcess(server_process.as_raw_handle().cast(), /*exit_code*/ 0) + }; + assert_ne!(terminated, 0, "failed to terminate test MCP server"); + wait_for_process_exit(&server_process)?; + + client.shutdown().await; + + assert!( + wait_for_process_exit(&descendant_process).is_ok(), + "MCP descendant survived its exited server (mode={protocol_mode:?})" + ); + } + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/stdio_stderr_cleanup.rs b/codex-rs/rmcp-client/tests/stdio_stderr_cleanup.rs new file mode 100644 index 0000000000000000000000000000000000000000..fa8b63c48cc010011b5507c950b838da1a32107f --- /dev/null +++ b/codex-rs/rmcp-client/tests/stdio_stderr_cleanup.rs @@ -0,0 +1,222 @@ +//! The local stderr reader must close when its MCP transport ends, even when +//! a descendant outside the server process group still has stderr open. + +#![cfg(unix)] + +use std::fs; +use std::io::Write; +use std::path::Path; +use std::sync::Arc; +use std::sync::Mutex; +use std::time::Duration; + +use anyhow::Result; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::LocalStdioServerLauncher; +use codex_rmcp_client::RmcpClient; +use futures::FutureExt as _; +use pretty_assertions::assert_eq; +use rmcp::model::ClientCapabilities; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; + +struct EscapedChild(i32); + +enum Cleanup { + Shutdown, + Drop, +} + +#[derive(Clone)] +struct TestLogWriter(Arc>>); + +impl Write for TestLogWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.0 + .lock() + .map_err(|_| std::io::Error::other("log buffer lock"))? + .extend_from_slice(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +impl Drop for EscapedChild { + fn drop(&mut self) { + // Only this test's PID, written after setsid(), is eligible for cleanup. + // The fixture also exits independently after 30 seconds. + unsafe { libc::kill(self.0, libc::SIGKILL) }; + } +} + +fn fd_count() -> Result { + let directory = if cfg!(target_os = "linux") { + "/proc/self/fd" + } else { + "/dev/fd" + }; + Ok(fs::read_dir(directory)?.count()) +} + +async fn read_pid(path: &Path) -> Result { + for _ in 0..100 { + if let Ok(text) = fs::read_to_string(path) + && let Ok(pid) = text.trim().parse::() + && pid > 1 + { + return Ok(pid); + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + anyhow::bail!("fixture PID unavailable: {}", path.display()) +} + +#[tokio::test(flavor = "current_thread")] +async fn client_teardown_closes_stderr_fd_and_logs_queued_diagnostics() -> Result<()> { + let logs = Arc::new(Mutex::new(Vec::new())); + let writer_logs = Arc::clone(&logs); + // The reader task runs on this thread, so its logs use this subscriber. + let _subscriber_guard = tracing::subscriber::set_default( + tracing_subscriber::fmt() + .with_ansi(false) + .with_max_level(tracing::Level::INFO) + .with_writer(move || TestLogWriter(Arc::clone(&writer_logs))) + .finish(), + ); + // FD counts are process-wide, so run the two teardown paths sequentially. + assert_stderr_reader_closed(Cleanup::Shutdown, &logs).await?; + assert_stderr_reader_closed(Cleanup::Drop, &logs).await +} + +async fn assert_stderr_reader_closed(cleanup: Cleanup, logs: &Arc>>) -> Result<()> { + let temporary = tempfile::tempdir()?; + let server_pid_file = temporary.path().join("server.pid"); + let child_pid_file = temporary.path().join("escaped.pid"); + let trigger_file = temporary.path().join("emit-stderr"); + let written_file = temporary.path().join("stderr-written"); + let diagnostic = match &cleanup { + Cleanup::Shutdown => "queued diagnostic before shutdown", + Cleanup::Drop => "queued diagnostic before drop", + }; + let script = temporary.path().join("server.py"); + fs::write( + &script, + r#"import json, os, pathlib, signal, sys, threading, time +# Keep either process bounded if test cleanup fails. +MAX_FIXTURE_LIFETIME_SECONDS = 30 +server_pid_path = pathlib.Path(sys.argv[1]) +escaped_pid_path = pathlib.Path(sys.argv[2]) +trigger_path = pathlib.Path(sys.argv[3]) +written_path = pathlib.Path(sys.argv[4]) +signal.alarm(MAX_FIXTURE_LIFETIME_SECONDS) +server_pid_path.write_text(str(os.getpid())) +if os.fork() == 0: # fork() returns zero in the child. + os.setsid() + signal.alarm(MAX_FIXTURE_LIFETIME_SECONDS) + os.close(sys.stdin.fileno()) + os.close(sys.stdout.fileno()) + escaped_pid_path.write_text(str(os.getpid())) + # Exit even if SIGALRM was ignored; the test normally kills this child sooner. + time.sleep(MAX_FIXTURE_LIFETIME_SECONDS) + os._exit(os.EX_OK) +def emit_diagnostic(): + while not trigger_path.exists(): + time.sleep(0.01) + print(sys.argv[5], file=sys.stderr, flush=True) + written_path.write_text("written") +threading.Thread(target=emit_diagnostic, daemon=True).start() +for line in sys.stdin: + message = json.loads(line) + if message.get("method") == "initialize": + result = {"protocolVersion": "2025-06-18", "capabilities": {}, "serverInfo": {"name": "stderr-reader-fixture", "version": "1"}} + print(json.dumps({"jsonrpc": "2.0", "id": message["id"], "result": result}), flush=True) +"#, + )?; + let baseline = fd_count()?; + let client = RmcpClient::new_stdio_client( + which::which("python3")?.into(), + vec![ + script.into_os_string(), + server_pid_file.clone().into_os_string(), + child_pid_file.clone().into_os_string(), + trigger_file.clone().into_os_string(), + written_file.clone().into_os_string(), + diagnostic.into(), + ], + /*env*/ None, + &[], + /*cwd*/ None, + Arc::new(LocalStdioServerLauncher::new(std::env::current_dir()?)), + ) + .await?; + let server_pid = read_pid(&server_pid_file).await?; + let escaped = EscapedChild(read_pid(&child_pid_file).await?); + assert_eq!(unsafe { libc::getpgid(server_pid) }, server_pid); + assert_eq!(unsafe { libc::getpgid(escaped.0) }, escaped.0); + client + .initialize( + InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("stderr-reader-test", "1"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18), + Some(Duration::from_secs(5)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Decline, + content: None, + meta: None, + }) + } + .boxed() + }), + ) + .await?; + // Keep the current-thread Tokio runtime blocked until the server has + // written stderr. The reader cannot run before teardown starts. + fs::write(&trigger_file, "")?; + for _ in 0..250 { + if written_file.exists() { + break; + } + std::thread::sleep(Duration::from_millis(20)); + } + assert!(written_file.exists(), "fixture did not write stderr"); + // For Shutdown, keep the client alive until after the FD assertion so + // final Drop cannot hide missing cleanup in the explicit shutdown path. + if matches!(cleanup, Cleanup::Shutdown) { + client.shutdown().await; + } else { + drop(client); + } + let mut after_shutdown = fd_count()?; + for _ in 0..250 { + if after_shutdown == baseline { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + after_shutdown = fd_count()?; + } + assert_eq!( + after_shutdown, baseline, + "stderr reader still owns a file descriptor while the escaped child is alive" + ); + assert_eq!(unsafe { libc::kill(escaped.0, 0) }, 0); + let captured_logs = String::from_utf8( + logs.lock() + .map_err(|_| anyhow::anyhow!("log buffer lock"))? + .clone(), + )?; + assert!( + captured_logs.contains(diagnostic), + "stderr diagnostic written before teardown was lost: {captured_logs}" + ); + drop(escaped); + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs b/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs new file mode 100644 index 0000000000000000000000000000000000000000..374e3232c22e257a014c444688fa454e6611eb97 --- /dev/null +++ b/codex-rs/rmcp-client/tests/streamable_http_oauth_startup.rs @@ -0,0 +1,1131 @@ +mod streamable_http_test_support; + +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_rmcp_client::McpAuthState; +use codex_rmcp_client::McpLoginRequirement; +use codex_rmcp_client::McpOAuthRefreshMode; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::OAuthDiscoveryTimeout; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::StoredOAuthTokens; +use codex_rmcp_client::StreamableHttpRedirectMode; +use codex_rmcp_client::WrappedOAuthTokenResponse; +use codex_rmcp_client::determine_streamable_http_auth_status; +use codex_rmcp_client::is_authentication_required_error; +use codex_rmcp_client::save_oauth_tokens; +use codex_rmcp_client::with_http_headers_helper; +use codex_utils_cargo_bin::cargo_bin; +use oauth2::AccessToken; +use oauth2::RefreshToken; +use oauth2::basic::BasicTokenType; +use pretty_assertions::assert_eq; +use rmcp::transport::auth::OAuthTokenResponse; +use rmcp::transport::auth::VendorExtraTokenFields; +use serde_json::Value; +use serde_json::json; +use tempfile::TempDir; +use tokio::process::Command; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::Request; +use wiremock::ResponseTemplate; +use wiremock::matchers::body_string_contains; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use streamable_http_test_support::initialize_client; +use streamable_http_test_support::initialize_client_with_timeout; + +const SERVER_NAME: &str = "test-streamable-http-oauth-startup"; +const EXPIRED_ACCESS_TOKEN: &str = "expired-access-token"; +const REFRESH_TOKEN: &str = "valid-refresh-token"; +const REFRESHED_ACCESS_TOKEN: &str = "refreshed-access-token"; +const RESOURCE_API_KEY: &str = "resource-api-key-secret"; +const RESOURCE_USER_AGENT: &str = "resource-only-user-agent"; +const MCP_USER_AGENT: &str = concat!("codex-mcp-client/", env!("CARGO_PKG_VERSION")); +const CHILD_SERVER_URL_ENV: &str = "MCP_TEST_OAUTH_STARTUP_SERVER_URL"; +const CHILD_REFRESH_MODE_ENV: &str = "MCP_TEST_OAUTH_REFRESH_MODE"; +const CHILD_HELPER_COMMAND_ENV: &str = "MCP_TEST_OAUTH_STARTUP_HELPER_COMMAND"; +const CHILD_RESOURCE_API_KEY_ENV: &str = "MCP_TEST_OAUTH_STARTUP_RESOURCE_API_KEY"; +const CHILD_STORED_ISSUER_ENV: &str = "MCP_TEST_OAUTH_STARTUP_STORED_ISSUER"; +const CHILD_ACCESS_TOKEN_EXPIRY_ENV: &str = "MCP_TEST_OAUTH_STARTUP_ACCESS_TOKEN_EXPIRY"; +const CHILD_REFRESH_SUCCEEDS_ENV: &str = "MCP_TEST_REFRESH_SUCCEEDS"; +const LEGACY_REFRESHABLE_SERVER_URL: &str = "https://legacy-refreshable.example/mcp"; +const UNEXPIRED_SERVER_URL: &str = "https://unexpired.example/mcp"; +const REFRESHABLE_SERVER_URL: &str = "https://refreshable.example/mcp"; + +#[derive(Clone, Copy)] +enum OAuthStartupScenario { + DirectAuthorizationMetadata, + OidcMetadataAfter503, + GatewayHeadersHelper, + SameOriginGatewayHeadersHelper, + ProtectedResourceMetadata, +} + +#[derive(Clone, Copy)] +enum IssuerMismatchAccessToken { + Expired, + Unexpired, +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn tool_call_preserves_challenge_only_after_silent_refresh_fails() -> anyhow::Result<()> { + for refresh_succeeds in [true, false] { + let server = MockServer::start().await; + let server_url = format!("{}/mcp", server.uri()); + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "issuer": server_url, + "authorization_endpoint": format!("{}/authorize", server.uri()), + "token_endpoint": format!("{}/token", server.uri()), + }))) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/token")) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains(format!( + "refresh_token={REFRESH_TOKEN}" + ))) + .respond_with(if refresh_succeeds { + ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": REFRESHED_ACCESS_TOKEN, + "token_type": "Bearer", + "expires_in": 7200, + "refresh_token": REFRESH_TOKEN, + })) + } else { + ResponseTemplate::new(/*s*/ 400).set_body_json(json!({"error": "invalid_grant"})) + }) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(|request: &Request| { + let body: Value = request.body_json().unwrap(); + let result = match body["method"].as_str() { + Some("initialize") => json!({ + "protocolVersion": body["params"]["protocolVersion"], + "capabilities": {}, + "serverInfo": {"name": "oauth-tool-call-test", "version": "1"}, + }), + Some("notifications/initialized") => { + return ResponseTemplate::new(/*s*/ 202); + } + Some("tools/call") => { + if request.headers.get("authorization").unwrap() + != format!("Bearer {REFRESHED_ACCESS_TOKEN}").as_str() + { + return ResponseTemplate::new(/*s*/ 401).insert_header( + "www-authenticate", + r#"Bearer error="invalid_token""#, + ); + } + json!({"content": [{"type": "text", "text": "refreshed"}]}) + } + _ => return ResponseTemplate::new(/*s*/ 400), + }; + ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "jsonrpc": "2.0", "id": body["id"], "result": result, + })) + }) + .mount(&server) + .await; + + // The child owns its credential store; the parallel test runner's environment is unchanged. + let codex_home = TempDir::new()?; + let status = Command::new(std::env::current_exe()?) + .args([ + "oauth_tool_call_child", + "--exact", + "--ignored", + "--nocapture", + ]) + .env("CODEX_HOME", codex_home.path()) + .env(CHILD_SERVER_URL_ENV, server_url) + .env(CHILD_REFRESH_SUCCEEDS_ENV, refresh_succeeds.to_string()) + .status() + .await?; + assert!(status.success(), "OAuth tool call child failed: {status}"); + server.verify().await; + let tool_calls = server + .received_requests() + .await + .unwrap() + .into_iter() + .filter(|request| { + request + .body_json::() + .ok() + .is_some_and(|body| body["method"] == "tools/call") + }) + .count(); + assert_eq!(tool_calls, if refresh_succeeds { 2 } else { 1 }); + } + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[ignore = "spawned by tool_call_preserves_challenge_only_after_silent_refresh_fails"] +async fn oauth_tool_call_child() -> anyhow::Result<()> { + let server_url = std::env::var(CHILD_SERVER_URL_ENV)?; + let refresh_succeeds = std::env::var(CHILD_REFRESH_SUCCEEDS_ENV)? == "true"; + let mut response = OAuthTokenResponse::new( + AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new(REFRESH_TOKEN.to_string()))); + response.set_expires_in(Some(&Duration::from_secs(/*secs*/ 7200))); + save_oauth_tokens( + SERVER_NAME, + &StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: server_url.clone(), + issuer: Some(server_url.clone()), + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some( + (SystemTime::now().duration_since(UNIX_EPOCH)?.as_millis() + 7_200_000) as u64, + ), + }, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .await?; + let client = RmcpClient::new_streamable_http_client( + SERVER_NAME, + &server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await?; + initialize_client(&client).await?; + let result = client + .call_tool( + "probe".to_string(), + /*arguments*/ None, + /*meta*/ None, + Some(Duration::from_secs(/*secs*/ 5)), + ) + .await?; + if refresh_succeeds { + assert_eq!( + serde_json::to_value(result)?, + json!({ + "content": [{"type": "text", "text": "refreshed"}], + }) + ); + } else { + assert_eq!(result.is_error, Some(true)); + assert_eq!( + serde_json::to_value(result.meta)?, + json!({ + "mcp/www_authenticate": [r#"Bearer error="invalid_token""#], + }) + ); + } + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn refreshes_expired_persisted_token_before_initialize() -> anyhow::Result<()> { + assert_expired_token_refresh( + OAuthStartupScenario::DirectAuthorizationMetadata, + McpOAuthRefreshMode::Coordinated, + ) + .await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn refreshes_expired_persisted_token_after_oidc_fallback() -> anyhow::Result<()> { + for refresh_mode in [ + McpOAuthRefreshMode::Legacy, + McpOAuthRefreshMode::Coordinated, + ] { + assert_expired_token_refresh(OAuthStartupScenario::OidcMetadataAfter503, refresh_mode) + .await?; + } + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn refreshes_oauth_with_gateway_headers_helper() -> anyhow::Result<()> { + assert_expired_token_refresh( + OAuthStartupScenario::GatewayHeadersHelper, + McpOAuthRefreshMode::Legacy, + ) + .await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn refreshes_oauth_after_gateway_rejects_token_request() -> anyhow::Result<()> { + assert_expired_token_refresh( + OAuthStartupScenario::SameOriginGatewayHeadersHelper, + McpOAuthRefreshMode::Legacy, + ) + .await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn refresh_uses_discovered_protected_resource_audience() -> anyhow::Result<()> { + assert_expired_token_refresh( + OAuthStartupScenario::ProtectedResourceMetadata, + McpOAuthRefreshMode::Legacy, + ) + .await +} + +async fn assert_expired_token_refresh( + scenario: OAuthStartupScenario, + refresh_mode: McpOAuthRefreshMode, +) -> anyhow::Result<()> { + let server = MockServer::start().await; + let authorization_server = MockServer::start().await; + let same_origin_gateway = matches!( + scenario, + OAuthStartupScenario::SameOriginGatewayHeadersHelper + ); + let token_server = if same_origin_gateway { + &server + } else { + &authorization_server + }; + let resource_url = format!("{}/mcp", server.uri()); + let (server_url, mcp_path, authorization_metadata_path) = match scenario { + OAuthStartupScenario::DirectAuthorizationMetadata + | OAuthStartupScenario::GatewayHeadersHelper + | OAuthStartupScenario::SameOriginGatewayHeadersHelper => ( + resource_url.clone(), + "/mcp", + "/.well-known/oauth-authorization-server/mcp", + ), + OAuthStartupScenario::ProtectedResourceMetadata => ( + format!("{resource_url}/?oauth=initialize"), + "/mcp/", + "/.well-known/oauth-authorization-server", + ), + OAuthStartupScenario::OidcMetadataAfter503 => ( + resource_url.clone(), + "/mcp", + "/mcp/.well-known/openid-configuration", + ), + }; + + if matches!(scenario, OAuthStartupScenario::OidcMetadataAfter503) { + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(503)) + .expect(2) + .mount(&server) + .await; + } + + if matches!(scenario, OAuthStartupScenario::ProtectedResourceMetadata) { + let resource_metadata_url = format!("{}/resource-metadata", server.uri()); + Mock::given(method("GET")) + .and(path(mcp_path)) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(2) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "resource": resource_url, + "authorization_servers": [server.uri()], + }))) + .expect(2) + .mount(&server) + .await; + } + + let mut authorization_metadata = json!({ + "issuer": server_url.clone(), + "authorization_endpoint": format!("{}/oauth/authorize", token_server.uri()), + "token_endpoint": format!("{}/oauth/token", token_server.uri()), + "scopes_supported": [""], + }); + if matches!(scenario, OAuthStartupScenario::ProtectedResourceMetadata) { + authorization_metadata["issuer"] = json!(server.uri()); + } + let authorization_server_issuer = authorization_metadata["issuer"] + .as_str() + .ok_or_else(|| anyhow::anyhow!("authorization metadata should include issuer"))? + .to_string(); + + let helper_directory = TempDir::new()?; + let helper_invocations = helper_directory.path().join("helper-invocations"); + Mock::given(method("GET")) + .and(path(authorization_metadata_path)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .respond_with(ResponseTemplate::new(200).set_body_json(authorization_metadata)) + .expect(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(header( + "user-agent", + if same_origin_gateway { + RESOURCE_USER_AGENT + } else { + MCP_USER_AGENT + }, + )) + .and(body_string_contains("grant_type=refresh_token")) + .and(body_string_contains(format!( + "refresh_token={REFRESH_TOKEN}" + ))) + .and({ + let expected_resource = resource_url.clone(); + move |request: &Request| { + url::form_urlencoded::parse(&request.body) + .any(|(name, value)| name == "resource" && value == expected_resource) + } + }) + .respond_with(move |request: &Request| { + if same_origin_gateway + && request + .headers + .get("x-helper-generation") + .is_some_and(|generation| generation == "1") + { + ResponseTemplate::new(/*s*/ 401) + } else { + let response = ResponseTemplate::new(/*s*/ 200).set_body_json(json!({ + "access_token": REFRESHED_ACCESS_TOKEN, + "token_type": "Bearer", + "expires_in": 7200, + "refresh_token": REFRESH_TOKEN, + })); + if refresh_mode == McpOAuthRefreshMode::Coordinated { + // Longer than the child's handshake timeout: refresh must finish first. + response.set_delay(Duration::from_secs(/*secs*/ 2)) + } else { + response + } + } + }) + .expect(if same_origin_gateway { 2 } else { 1 }) + .mount(token_server) + .await; + Mock::given(method("POST")) + .and(path(mcp_path)) + .and(header("user-agent", RESOURCE_USER_AGENT)) + .and(header("x-api-key", RESOURCE_API_KEY)) + .and(header( + "authorization", + format!("Bearer {REFRESHED_ACCESS_TOKEN}"), + )) + .respond_with(|request: &Request| { + let body: Value = match request.body_json() { + Ok(body) => body, + Err(_) => { + return ResponseTemplate::new(400).set_body_string("invalid JSON-RPC request"); + } + }; + match body.get("method").and_then(Value::as_str) { + Some("initialize") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body.get("id").cloned().unwrap_or(Value::Null), + "result": { + "protocolVersion": body + .pointer("/params/protocolVersion") + .cloned() + .unwrap_or_else(|| json!("2025-06-18")), + "capabilities": {}, + "serverInfo": { + "name": "oauth-startup-test", + "version": "0.0.0-test", + }, + }, + })), + Some("notifications/initialized") => ResponseTemplate::new(202), + method => ResponseTemplate::new(400) + .set_body_string(format!("unexpected JSON-RPC method: {method:?}")), + } + }) + .expect(2) + .mount(&server) + .await; + + let codex_home = TempDir::new()?; + let with_headers_helper = matches!( + scenario, + OAuthStartupScenario::GatewayHeadersHelper + | OAuthStartupScenario::SameOriginGatewayHeadersHelper + ); + + // Credential storage resolves CODEX_HOME from the process environment. + // Run the client half of the test in an ignored helper test so it can use + // an isolated home without mutating the parent test runner's environment. + let mut command = Command::new(std::env::current_exe()?); + command + .args(["oauth_startup_child", "--exact", "--ignored", "--nocapture"]) + .env("CODEX_HOME", codex_home.path()) + .env(CHILD_SERVER_URL_ENV, server_url) + .env(CHILD_STORED_ISSUER_ENV, authorization_server_issuer) + .env(CHILD_RESOURCE_API_KEY_ENV, RESOURCE_API_KEY) + .env( + CHILD_REFRESH_MODE_ENV, + if refresh_mode == McpOAuthRefreshMode::Coordinated { + "coordinated" + } else { + "legacy" + }, + ) + .env("MCP_TEST_AMBIENT_SECRET", "must-not-reach-helper"); + if with_headers_helper { + command.env( + CHILD_HELPER_COMMAND_ENV, + format!( + "\"{}\" --http-headers-helper \"{}\"", + cargo_bin("test_streamable_http_server")?.display(), + helper_invocations.display(), + ), + ); + } + let status = command.status().await?; + assert!(status.success(), "OAuth startup child failed: {status}"); + if with_headers_helper { + assert_eq!( + std::fs::read_to_string(helper_invocations)?, + if same_origin_gateway { "xx" } else { "x" } + ); + let requests = server.received_requests().await.unwrap_or_default(); + assert!(requests.iter().all(|request| { + request + .headers + .get("proxy-authorization") + .is_some_and(|value| value == "Bearer gateway-token") + })); + } + let authorization_requests = authorization_server + .received_requests() + .await + .ok_or_else(|| anyhow::anyhow!("authorization server should record requests"))?; + server.verify().await; + authorization_server.verify().await; + if same_origin_gateway { + assert!(authorization_requests.is_empty()); + return Ok(()); + } + assert_eq!(authorization_requests.len(), 1); + assert_eq!(authorization_requests[0].headers.get("x-api-key"), None); + assert_eq!( + authorization_requests[0].headers.get("proxy-authorization"), + None + ); + assert_eq!( + authorization_requests[0] + .headers + .get("user-agent") + .map(http::HeaderValue::as_bytes), + Some(MCP_USER_AGENT.as_bytes()) + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn rejects_refresh_when_authorization_server_issuer_changes_before_startup() +-> anyhow::Result<()> { + assert_issuer_mismatch_startup(IssuerMismatchAccessToken::Expired).await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn does_not_refresh_unexpired_token_after_issuer_change_401() -> anyhow::Result<()> { + assert_issuer_mismatch_startup(IssuerMismatchAccessToken::Unexpired).await +} + +async fn assert_issuer_mismatch_startup( + access_token: IssuerMismatchAccessToken, +) -> anyhow::Result<()> { + let issuer_a = MockServer::start().await; + let issuer_b = MockServer::start().await; + let mcp_server = MockServer::start().await; + let server_url = format!("{}/mcp", mcp_server.uri()); + let resource_metadata_url = format!("{}/resource-metadata", mcp_server.uri()); + + Mock::given(method("GET")) + .and(path("/mcp")) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(1) + .mount(&mcp_server) + .await; + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "resource": server_url, + "authorization_servers": [issuer_b.uri()], + }))) + .expect(1) + .mount(&mcp_server) + .await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "issuer": issuer_b.uri(), + "authorization_endpoint": format!("{}/oauth/authorize", issuer_b.uri()), + "token_endpoint": format!("{}/oauth/token", issuer_b.uri()), + "scopes_supported": [""], + }))) + .expect(1) + .mount(&issuer_b) + .await; + Mock::given(method("POST")) + .and(path("/mcp")) + .and(header( + "authorization", + format!("Bearer {EXPIRED_ACCESS_TOKEN}"), + )) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(match access_token { + IssuerMismatchAccessToken::Expired => 0, + IssuerMismatchAccessToken::Unexpired => 1, + }) + .mount(&mcp_server) + .await; + + let stored_issuer = issuer_a.uri(); + let status = run_issuer_startup_child( + &server_url, + &stored_issuer, + match access_token { + IssuerMismatchAccessToken::Expired => "expired", + IssuerMismatchAccessToken::Unexpired => "unexpired", + }, + ) + .await?; + + assert!( + status.success(), + "issuer mismatch startup child failed: {status}" + ); + let issuer_b_requests = issuer_b.received_requests().await.unwrap_or_default(); + assert!( + issuer_b_requests + .iter() + .all(|request| request.url.path() != "/oauth/token"), + "stored refresh token must not be posted to the replacement issuer" + ); + mcp_server.verify().await; + issuer_b.verify().await; + Ok(()) +} + +async fn run_issuer_startup_child( + server_url: &str, + stored_issuer: &str, + access_token_expiry: &str, +) -> anyhow::Result { + let codex_home = TempDir::new()?; + Ok(Command::new(std::env::current_exe()?) + .args([ + "issuer_startup_child", + "--exact", + "--ignored", + "--nocapture", + ]) + .env("CODEX_HOME", codex_home.path()) + .env(CHILD_SERVER_URL_ENV, server_url) + .env(CHILD_STORED_ISSUER_ENV, stored_issuer) + .env(CHILD_ACCESS_TOKEN_EXPIRY_ENV, access_token_expiry) + .status() + .await?) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn does_not_refresh_against_metadata_from_second_discovery() -> anyhow::Result<()> { + let issuer_a = MockServer::start().await; + let issuer_b = MockServer::start().await; + let mcp_server = MockServer::start().await; + let server_url = format!("{}/mcp", mcp_server.uri()); + let resource_metadata_url = format!("{}/resource-metadata", mcp_server.uri()); + let discovery_count = Arc::new(AtomicUsize::new(0)); + let issuer_a_url = issuer_a.uri(); + let issuer_b_url = issuer_b.uri(); + + Mock::given(method("GET")) + .and(path("/mcp")) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(1) + .mount(&mcp_server) + .await; + Mock::given(method("GET")) + .and(path("/resource-metadata")) + .respond_with({ + let discovery_count = Arc::clone(&discovery_count); + let issuer_a_url = issuer_a_url.clone(); + let issuer_b_url = issuer_b_url.clone(); + let metadata_resource_url = server_url.clone(); + move |_request: &Request| { + let issuer = if discovery_count.fetch_add(1, Ordering::SeqCst) == 0 { + &issuer_a_url + } else { + &issuer_b_url + }; + ResponseTemplate::new(200).set_body_json(json!({ + "resource": metadata_resource_url, + "authorization_servers": [issuer], + })) + } + }) + .expect(1) + .mount(&mcp_server) + .await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "issuer": issuer_a.uri(), + "authorization_endpoint": format!("{}/oauth/authorize", issuer_a.uri()), + "token_endpoint": format!("{}/oauth/token", issuer_a.uri()), + "scopes_supported": [""], + }))) + .expect(1) + .mount(&issuer_a) + .await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "issuer": issuer_b.uri(), + "authorization_endpoint": format!("{}/oauth/authorize", issuer_b.uri()), + "token_endpoint": format!("{}/oauth/token", issuer_b.uri()), + "scopes_supported": [""], + }))) + .mount(&issuer_b) + .await; + Mock::given(method("POST")) + .and(path("/mcp")) + .and(header( + "authorization", + format!("Bearer {EXPIRED_ACCESS_TOKEN}"), + )) + .respond_with(ResponseTemplate::new(401).insert_header( + "www-authenticate", + format!("Bearer resource_metadata=\"{resource_metadata_url}\""), + )) + .expect(1) + .mount(&mcp_server) + .await; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(body_string_contains(format!( + "refresh_token={REFRESH_TOKEN}" + ))) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({ + "error": "invalid_grant", + }))) + .expect(1) + .mount(&issuer_a) + .await; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .respond_with(ResponseTemplate::new(500)) + .expect(0) + .mount(&issuer_b) + .await; + + let status = run_issuer_startup_child(&server_url, &issuer_a_url, "unexpired").await?; + + assert!( + status.success(), + "metadata swap startup child failed: {status}" + ); + mcp_server.verify().await; + issuer_a.verify().await; + issuer_b.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn reports_auth_status_for_persisted_credentials() -> anyhow::Result<()> { + let codex_home = TempDir::new()?; + + let status = Command::new(std::env::current_exe()?) + .args([ + "persisted_credentials_auth_status_child", + "--exact", + "--ignored", + "--nocapture", + ]) + .env("CODEX_HOME", codex_home.path()) + .status() + .await?; + + assert!( + status.success(), + "persisted credentials auth status child failed: {status}" + ); + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn identifies_expired_unrefreshable_token_startup_error() -> anyhow::Result<()> { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "authorization_endpoint": format!("{}/oauth/authorize", server.uri()), + "token_endpoint": format!("{}/oauth/token", server.uri()), + }))) + .expect(1) + .mount(&server) + .await; + + let codex_home = TempDir::new()?; + let status = Command::new(std::env::current_exe()?) + .args([ + "expired_unrefreshable_startup_child", + "--exact", + "--ignored", + "--nocapture", + ]) + .env("CODEX_HOME", codex_home.path()) + .env(CHILD_SERVER_URL_ENV, format!("{}/mcp", server.uri())) + .status() + .await?; + + assert!( + status.success(), + "expired OAuth startup child failed: {status}" + ); + server.verify().await; + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[ignore = "spawned by reports_auth_status_for_persisted_credentials"] +async fn persisted_credentials_auth_status_child() -> anyhow::Result<()> { + let first_login_server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "authorization_endpoint": format!("{}/oauth/authorize", first_login_server.uri()), + "token_endpoint": format!("{}/oauth/token", first_login_server.uri()), + }))) + .expect(1) + .mount(&first_login_server) + .await; + + let status = auth_status(&format!("{}/mcp", first_login_server.uri())).await?; + assert_eq!(status, McpAuthState::LoggedOut(McpLoginRequirement::Login)); + first_login_server.verify().await; + + let mut response = OAuthTokenResponse::new( + AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new(REFRESH_TOKEN.to_string()))); + let tokens = StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: LEGACY_REFRESHABLE_SERVER_URL.to_string(), + issuer: None, + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some(0), + }; + save_oauth_tokens( + SERVER_NAME, + &tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .await?; + + let status = auth_status(LEGACY_REFRESHABLE_SERVER_URL).await?; + assert_eq!( + status, + McpAuthState::LoggedOut(McpLoginRequirement::Reauthentication) + ); + + let response = OAuthTokenResponse::new( + AccessToken::new("unexpired-access-token".to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_else(|_| Duration::from_secs(0)) + .as_millis() as u64; + let tokens = StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: UNEXPIRED_SERVER_URL.to_string(), + issuer: None, + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some(now.saturating_add(/*rhs*/ 60_000)), + }; + save_oauth_tokens( + SERVER_NAME, + &tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .await?; + + let status = auth_status(UNEXPIRED_SERVER_URL).await?; + assert_eq!(status, McpAuthState::OAuth); + + let mut response = OAuthTokenResponse::new( + AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new(REFRESH_TOKEN.to_string()))); + let tokens = StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: REFRESHABLE_SERVER_URL.to_string(), + issuer: Some("https://issuer.example.test".to_string()), + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some(0), + }; + save_oauth_tokens( + SERVER_NAME, + &tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .await?; + + let status = auth_status(REFRESHABLE_SERVER_URL).await?; + assert_eq!(status, McpAuthState::OAuth); + Ok(()) +} + +async fn auth_status(server_url: &str) -> anyhow::Result { + determine_streamable_http_auth_status( + SERVER_NAME, + server_url, + /*bearer_token_env_var*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + OAuthDiscoveryTimeout::LOCAL, + StreamableHttpRedirectMode::Legacy, + ) + .await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[ignore = "spawned by refreshes_expired_persisted_token_before_initialize"] +async fn oauth_startup_child() -> anyhow::Result<()> { + let server_url = std::env::var(CHILD_SERVER_URL_ENV)?; + let refresh_mode = if std::env::var(CHILD_REFRESH_MODE_ENV).as_deref() == Ok("coordinated") { + McpOAuthRefreshMode::Coordinated + } else { + McpOAuthRefreshMode::Legacy + }; + + // Save an expired access token with a valid refresh token so startup must + // refresh before sending the initialize request. + let mut response = OAuthTokenResponse::new( + AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new(REFRESH_TOKEN.to_string()))); + response.set_expires_in(Some(&Duration::from_secs(7200))); + let tokens = StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: server_url.clone(), + issuer: Some(std::env::var(CHILD_STORED_ISSUER_ENV)?), + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some(0), + }; + save_oauth_tokens( + SERVER_NAME, + &tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .await?; + + // This mirrors create_client's transport and initialization setup, except + // it omits the direct bearer token. Supplying that token would bypass the + // persisted OAuth credentials and the startup refresh under test. + let mut http_client = Environment::default_for_tests().get_http_client(); + if let Ok(helper_command) = std::env::var(CHILD_HELPER_COMMAND_ENV) { + http_client = with_http_headers_helper( + http_client, + &server_url, + &helper_command, + std::env::current_dir()?, + )?; + } + let client = RmcpClient::new_streamable_http_client_with_protocol_mode_and_redirect_mode( + SERVER_NAME, + &server_url, + /*bearer_token*/ None, + Some(HashMap::from([( + "User-Agent".to_string(), + RESOURCE_USER_AGENT.to_string(), + )])), + Some(HashMap::from([( + "X-Api-Key".to_string(), + CHILD_RESOURCE_API_KEY_ENV.to_string(), + )])), + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + http_client, + /*auth_provider*/ None, + McpProtocolMode::Legacy, + StreamableHttpRedirectMode::Legacy, + refresh_mode, + ) + .await?; + + if refresh_mode == McpOAuthRefreshMode::Coordinated { + initialize_client_with_timeout(&client, Duration::from_secs(/*secs*/ 1)).await?; + } else { + initialize_client(&client).await?; + } + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[ignore = "spawned by identifies_expired_unrefreshable_token_startup_error"] +async fn expired_unrefreshable_startup_child() -> anyhow::Result<()> { + let server_url = std::env::var(CHILD_SERVER_URL_ENV)?; + let response = OAuthTokenResponse::new( + AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + let tokens = StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: server_url.clone(), + issuer: None, + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some(0), + }; + save_oauth_tokens( + SERVER_NAME, + &tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .await?; + + let client = RmcpClient::new_streamable_http_client( + SERVER_NAME, + &server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await?; + + let error = initialize_client(&client) + .await + .expect_err("expired token without a refresh token should fail startup"); + assert!(is_authentication_required_error(&error)); + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[ignore = "spawned by issuer startup tests"] +async fn issuer_startup_child() -> anyhow::Result<()> { + let server_url = std::env::var(CHILD_SERVER_URL_ENV)?; + let mut response = OAuthTokenResponse::new( + AccessToken::new(EXPIRED_ACCESS_TOKEN.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + response.set_refresh_token(Some(RefreshToken::new(REFRESH_TOKEN.to_string()))); + response.set_expires_in(Some(&Duration::from_secs(7200))); + let expires_at = match std::env::var(CHILD_ACCESS_TOKEN_EXPIRY_ENV).as_deref() { + Ok("expired") | Err(_) => Some(0), + Ok("unexpired") => { + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_else(|_| Duration::from_secs(0)) + .as_millis() as u64; + Some(now.saturating_add(/*rhs*/ 60_000)) + } + Ok(value) => anyhow::bail!("unexpected access token expiry fixture: {value}"), + }; + let tokens = StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: server_url.clone(), + issuer: Some(std::env::var(CHILD_STORED_ISSUER_ENV)?), + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at, + }; + save_oauth_tokens( + SERVER_NAME, + &tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .await?; + + let client = RmcpClient::new_streamable_http_client( + SERVER_NAME, + &server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await; + + let error = match client { + Ok(client) => initialize_client(&client) + .await + .expect_err("stored refresh token must fail when startup cannot refresh safely"), + Err(error) => error, + }; + assert!( + is_authentication_required_error(&error), + "unexpected issuer startup error: {error:#}" + ); + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/streamable_http_oauth_store_pinning.rs b/codex-rs/rmcp-client/tests/streamable_http_oauth_store_pinning.rs new file mode 100644 index 0000000000000000000000000000000000000000..1656774650042ca28ffbeac1295d4bf367bd8a4c --- /dev/null +++ b/codex-rs/rmcp-client/tests/streamable_http_oauth_store_pinning.rs @@ -0,0 +1,375 @@ +mod streamable_http_test_support; + +use std::any::Any; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpRedirectPolicy; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::StoredOAuthTokens; +use codex_rmcp_client::StreamableHttpBearerToken; +use codex_rmcp_client::StreamableHttpRedirectMode; +use codex_rmcp_client::WrappedOAuthTokenResponse; +use codex_rmcp_client::save_oauth_tokens; +use codex_rmcp_client::stored_oauth_credential_snapshot; +use futures::future::BoxFuture; +use keyring::credential::Credential; +use keyring::credential::CredentialApi; +use keyring::credential::CredentialBuilderApi; +use keyring::credential::CredentialPersistence; +use oauth2::AccessToken; +use oauth2::basic::BasicTokenType; +use pretty_assertions::assert_eq; +use rmcp::transport::auth::OAuthTokenResponse; +use rmcp::transport::auth::VendorExtraTokenFields; +use tempfile::TempDir; +use tokio::process::Command; + +use streamable_http_test_support::arm_session_post_failure; +use streamable_http_test_support::call_echo_tool; +use streamable_http_test_support::expected_echo_result; +use streamable_http_test_support::initialize_client; +use streamable_http_test_support::spawn_streamable_http_server; + +const SERVER_NAME: &str = "test-streamable-http-oauth-store-pinning"; +const CHILD_SERVER_URL_ENV: &str = "MCP_TEST_OAUTH_PINNED_STORE_SERVER_URL"; +const KEYRING_ACCESS_TOKEN: &str = "keyring-access-token"; +const FILE_ACCESS_TOKEN: &str = "stale-file-access-token"; + +#[derive(Clone)] +struct RecordingHttpClient { + inner: Arc, + bearer_tokens: Arc>>, + redirect_policies: Arc>>, +} + +impl RecordingHttpClient { + fn new(inner: Arc) -> Self { + Self { + inner, + bearer_tokens: Arc::new(Mutex::new(Vec::new())), + redirect_policies: Arc::new(Mutex::new(Vec::new())), + } + } + + fn record_request(&self, params: &HttpRequestParams) { + self.redirect_policies + .lock() + .unwrap_or_else(PoisonError::into_inner) + .push(params.redirect_policy); + let Some(header) = params + .headers + .iter() + .find(|header| header.name.eq_ignore_ascii_case("authorization")) + else { + return; + }; + self.bearer_tokens + .lock() + .unwrap_or_else(PoisonError::into_inner) + .push(header.value.clone()); + } + + fn bearer_tokens(&self) -> Vec { + self.bearer_tokens + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone() + } +} + +impl HttpClient for RecordingHttpClient { + fn http_request( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + self.record_request(¶ms); + self.inner.http_request(params) + } + + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + self.record_request(¶ms); + self.inner.http_request_stream(params) + } +} + +#[derive(Debug, Default)] +struct TestKeyringState { + secret: Mutex>>, + fail_reads: AtomicBool, +} + +#[derive(Clone, Debug)] +struct TestCredential { + state: Arc, +} + +impl CredentialApi for TestCredential { + fn set_secret(&self, secret: &[u8]) -> keyring::Result<()> { + *self + .state + .secret + .lock() + .unwrap_or_else(PoisonError::into_inner) = Some(secret.to_vec()); + Ok(()) + } + + fn get_secret(&self) -> keyring::Result> { + if self.state.fail_reads.load(Ordering::SeqCst) { + return Err(keyring::Error::Invalid( + "simulated keyring read failure".to_string(), + "load".to_string(), + )); + } + + self.state + .secret + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone() + .ok_or(keyring::Error::NoEntry) + } + + fn delete_credential(&self) -> keyring::Result<()> { + self.state + .secret + .lock() + .unwrap_or_else(PoisonError::into_inner) + .take() + .map(|_| ()) + .ok_or(keyring::Error::NoEntry) + } + + fn as_any(&self) -> &dyn Any { + self + } +} + +#[derive(Debug)] +struct TestCredentialBuilder { + state: Arc, +} + +impl CredentialBuilderApi for TestCredentialBuilder { + fn build( + &self, + _target: Option<&str>, + _service: &str, + _user: &str, + ) -> keyring::Result> { + Ok(Box::new(TestCredential { + state: Arc::clone(&self.state), + })) + } + + fn as_any(&self) -> &dyn Any { + self + } + + fn persistence(&self) -> CredentialPersistence { + CredentialPersistence::ProcessOnly + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn transport_provided_bearer_token_avoids_placeholder_headers_and_redirects() +-> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let http_client = RecordingHttpClient::new(Environment::default_for_tests().get_http_client()); + let client = RmcpClient::new_streamable_http_client_with_protocol_mode_and_redirect_mode( + "transport-provided-bearer-token", + &format!("{base_url}/mcp"), + Some(StreamableHttpBearerToken::ProvidedByHttpClient), + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + Arc::new(http_client.clone()), + /*auth_provider*/ None, + McpProtocolMode::Legacy, + StreamableHttpRedirectMode::AgentPluginV1, + codex_rmcp_client::McpOAuthRefreshMode::Legacy, + ) + .await?; + + initialize_client(&client).await?; + assert_eq!( + call_echo_tool(&client, "transport-provided").await?, + expected_echo_result("transport-provided") + ); + assert_eq!(http_client.bearer_tokens(), Vec::::new()); + + let redirect_policies = http_client + .redirect_policies + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone(); + assert!(!redirect_policies.is_empty()); + assert_eq!( + redirect_policies, + vec![HttpRedirectPolicy::Stop; redirect_policies.len()] + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn auto_store_remains_pinned_across_session_recovery() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let codex_home = TempDir::new()?; + + let status = Command::new(std::env::current_exe()?) + .args([ + "auto_store_remains_pinned_across_session_recovery_child", + "--exact", + "--ignored", + "--nocapture", + ]) + .env("CODEX_HOME", codex_home.path()) + .env(CHILD_SERVER_URL_ENV, &base_url) + .status() + .await?; + + assert!( + status.success(), + "OAuth store-pinning child failed: {status}" + ); + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[ignore = "spawned by auto_store_remains_pinned_across_session_recovery"] +async fn auto_store_remains_pinned_across_session_recovery_child() -> anyhow::Result<()> { + let state = Arc::new(TestKeyringState::default()); + keyring::set_default_credential_builder(Box::new(TestCredentialBuilder { + state: Arc::clone(&state), + })); + + let base_url = std::env::var(CHILD_SERVER_URL_ENV)?; + let server_url = format!("{base_url}/mcp"); + let file_tokens = stored_tokens(&server_url, FILE_ACCESS_TOKEN); + save_oauth_tokens( + SERVER_NAME, + &file_tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + ) + .await?; + let file_snapshot = stored_oauth_credential_snapshot( + SERVER_NAME, + &server_url, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + )? + .expect("Auto should initially resolve the fallback File"); + let keyring_tokens = stored_tokens(&server_url, KEYRING_ACCESS_TOKEN); + save_oauth_tokens( + SERVER_NAME, + &keyring_tokens, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + ) + .await?; + save_oauth_tokens( + SERVER_NAME, + &file_tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::Direct, + ) + .await?; + assert_eq!( + file_snapshot.reload( + SERVER_NAME, + &server_url, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + )?, + Some(keyring_tokens.clone()), + ); + let http_client = RecordingHttpClient::new(Environment::default_for_tests().get_http_client()); + + let client = RmcpClient::new_streamable_http_client( + SERVER_NAME, + &server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::Auto, + AuthKeyringBackendKind::Direct, + Arc::new(http_client.clone()), + /*auth_provider*/ None, + ) + .await?; + initialize_client(&client).await?; + assert_eq!( + call_echo_tool(&client, "warmup").await?, + expected_echo_result("warmup") + ); + + arm_session_post_failure( + &base_url, + /*status*/ 404, + /*remaining*/ 1, + /*www_authenticate_headers*/ &[], + ) + .await?; + // The selected keyring becomes unavailable only after initial construction. If recovery + // reevaluates Auto, it adopts the stale File token and this operation incorrectly succeeds. + state.fail_reads.store(true, Ordering::SeqCst); + + match call_echo_tool(&client, "recovery-must-not-fallback").await { + Ok(result) => assert_eq!(result, expected_echo_result("recovery-must-not-fallback")), + Err(error) => { + let error_chain = format!("{error:#}"); + assert!( + error_chain.contains("failed to reread OAuth tokens from resolved keyring storage"), + "unexpected recovery error: {error_chain}" + ); + } + } + + let bearer_tokens = http_client.bearer_tokens(); + assert!( + bearer_tokens + .iter() + .any(|token| token == &format!("Bearer {KEYRING_ACCESS_TOKEN}")), + "expected requests authenticated by the keyring token: {bearer_tokens:?}" + ); + assert!( + bearer_tokens + .iter() + .all(|token| token != &format!("Bearer {FILE_ACCESS_TOKEN}")), + "stale File token must never be sent during recovery: {bearer_tokens:?}" + ); + Ok(()) +} + +fn stored_tokens(server_url: &str, access_token: &str) -> StoredOAuthTokens { + StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: server_url.to_string(), + issuer: None, + client_id: "test-client-id".to_string(), + token_response: WrappedOAuthTokenResponse(OAuthTokenResponse::new( + AccessToken::new(access_token.to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + )), + expires_at: None, + } +} diff --git a/codex-rs/rmcp-client/tests/streamable_http_oauth_tool_call.rs b/codex-rs/rmcp-client/tests/streamable_http_oauth_tool_call.rs new file mode 100644 index 0000000000000000000000000000000000000000..c620410d1440b3469daa579e084c6fc6f57a582a --- /dev/null +++ b/codex-rs/rmcp-client/tests/streamable_http_oauth_tool_call.rs @@ -0,0 +1,387 @@ +//! Exercises OAuth recovery through a real MCP client, with isolated file credentials. + +mod streamable_http_test_support; + +use std::time::Duration; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_rmcp_client::RmcpClient; +use codex_rmcp_client::StoredOAuthTokens; +use codex_rmcp_client::WrappedOAuthTokenResponse; +use codex_rmcp_client::is_authentication_required_error; +use codex_rmcp_client::save_oauth_tokens; +use codex_rmcp_client::stored_oauth_credentials; +use oauth2::AccessToken; +use oauth2::RefreshToken; +use oauth2::basic::BasicTokenType; +use pretty_assertions::assert_eq; +use rmcp::model::CallToolResult; +use rmcp::model::ContentBlock; +use rmcp::transport::auth::OAuthTokenResponse; +use rmcp::transport::auth::VendorExtraTokenFields; +use serde_json::Value; +use serde_json::json; +use tempfile::TempDir; +use tokio::process::Command; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::Request; +use wiremock::ResponseTemplate; +use wiremock::matchers::body_string_contains; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use streamable_http_test_support::initialize_client; + +const SERVER_NAME: &str = "runtime-oauth"; +const SCENARIO_ENV: &str = "MCP_TEST_RUNTIME_OAUTH_SCENARIO"; + +#[tokio::test] +async fn runtime_oauth_recovery() -> anyhow::Result<()> { + run_scenarios(&[ + "no-refresh", + "rejected", + "malformed", + "transient", + "early-malformed", + "early-transient", + "refreshable", + "unexpired", + "unknown-expiry", + "revoked", + ]) + .await +} + +#[tokio::test] +async fn startup_oauth_recovery() -> anyhow::Result<()> { + run_scenarios(&[ + "startup-malformed", + "startup-transient", + "startup-refreshable", + ]) + .await +} + +async fn run_scenarios(scenarios: &[&str]) -> anyhow::Result<()> { + for scenario in scenarios { + let home = TempDir::new()?; + let output = Command::new(std::env::current_exe()?) + .args(["runtime_oauth_child", "--exact", "--ignored", "--nocapture"]) + .env("CODEX_HOME", home.path()) + .env(SCENARIO_ENV, scenario) + .output() + .await?; + assert!( + output.status.success(), + "{scenario}: {}\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + } + Ok(()) +} + +#[tokio::test] +#[ignore = "spawned by OAuth recovery tests with an isolated CODEX_HOME"] +async fn runtime_oauth_child() -> anyhow::Result<()> { + let scenario = std::env::var(SCENARIO_ENV)?; + let mcp = MockServer::start().await; + let authorization = MockServer::start().await; + let server_url = format!("{}/mcp", mcp.uri()); + let startup = scenario.starts_with("startup-"); + let revoked = scenario == "revoked"; + Mock::given(method("GET")) + .and(path("/.well-known/oauth-authorization-server/mcp")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "issuer": server_url, + "authorization_endpoint": format!("{}/authorize", authorization.uri()), + "token_endpoint": format!("{}/token", authorization.uri()), + }))) + .mount(&mcp) + .await; + Mock::given(method("POST")) + .and(path("/mcp")) + .respond_with(move |request: &Request| { + let body: Value = request.body_json().expect("JSON-RPC request"); + match body["method"].as_str() { + Some("initialize") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", "id": body["id"], + "result": { + "protocolVersion": body["params"]["protocolVersion"], + "capabilities": {"tools": {}}, + "serverInfo": {"name": "runtime-oauth", "version": "1"}, + }, + })), + Some("notifications/initialized") => ResponseTemplate::new(202), + Some("tools/list") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", "id": body["id"], "result": {"tools": []}, + })), + Some("tools/call") if revoked => ResponseTemplate::new(401) + .insert_header("www-authenticate", "Bearer error=\"invalid_token\""), + Some("tools/call") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", "id": body["id"], + "result": {"content": [{"type": "text", "text": "tool ran"}]}, + })), + _ => ResponseTemplate::new(400), + } + }) + .mount(&mcp) + .await; + // A zero lifetime expires after the handshake; 30 seconds exercises the early-refresh + // window while the token remains valid. Neither case needs sleeps or clock changes. + let expires_in = match scenario.as_str() { + "unexpired" | "startup-refreshable" => Some(3600), + "early-malformed" | "early-transient" => Some(30), + "unknown-expiry" => None, + _ => Some(0), + }; + let bootstrap_response = match scenario.as_str() { + "startup-malformed" => ResponseTemplate::new(200).set_body_raw("{", "application/json"), + "startup-transient" => { + ResponseTemplate::new(503).set_body_json(json!({"error": "temporarily_unavailable"})) + } + _ => ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "initial-token", "token_type": "Bearer", + "refresh_token": "refresh-token", "expires_in": expires_in, + })), + }; + let bootstrap = Mock::given(method("POST")) + .and(path("/token")) + .respond_with(bootstrap_response) + .mount_as_scoped(&authorization) + .await; + let mut response = OAuthTokenResponse::new( + AccessToken::new("old-token".to_string()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + if !revoked { + response.set_refresh_token(Some(RefreshToken::new("refresh-token".to_string()))); + } + let tokens = StoredOAuthTokens { + server_name: SERVER_NAME.to_string(), + url: server_url.clone(), + issuer: Some(server_url.clone()), + client_id: "test-client".to_string(), + token_response: WrappedOAuthTokenResponse(response), + expires_at: Some(if revoked { + u64::try_from(SystemTime::now().duration_since(UNIX_EPOCH)?.as_millis())? + 3_600_000 + } else { + 0 + }), + }; + save_oauth_tokens( + SERVER_NAME, + &tokens, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .await?; + let client = RmcpClient::new_streamable_http_client( + SERVER_NAME, + &server_url, + /*bearer_token*/ None, + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await?; + let initialization = initialize_client(&client).await; + if startup { + if scenario == "startup-refreshable" { + initialization?; + let tools = client + .list_tools(/*params*/ None, Some(Duration::from_secs(5))) + .await?; + assert!(tools.tools.is_empty()); + } else { + let error = initialization.expect_err("expired credentials must require reconnect"); + assert!(is_authentication_required_error(&error), "{error:#}"); + } + let requests = mcp.received_requests().await.expect("request recording"); + let rpc_requests = requests + .iter() + .filter(|request| request.method.as_str() == "POST" && request.url.path() == "/mcp") + .collect::>(); + let methods = rpc_requests + .iter() + .map(|request| { + request.body_json::().expect("JSON-RPC request")["method"].clone() + }) + .collect::>(); + assert_eq!( + methods, + if scenario == "startup-refreshable" { + vec![ + json!("initialize"), + json!("notifications/initialized"), + json!("tools/list"), + ] + } else { + Vec::::new() + } + ); + for request in rpc_requests { + assert_eq!( + request + .headers + .get("authorization") + .and_then(|value| value.to_str().ok()), + Some("Bearer initial-token") + ); + } + let after = stored_oauth_credentials( + SERVER_NAME, + &server_url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + )? + .expect("credentials are not deleted"); + if scenario != "startup-refreshable" { + assert_eq!(after, tokens); + } + assert_eq!( + authorization + .received_requests() + .await + .expect("request recording") + .len(), + 1 + ); + return Ok(()); + } + initialization?; + drop(bootstrap); + let mut before = stored_oauth_credentials( + SERVER_NAME, + &server_url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + )? + .expect("saved credentials"); + if scenario == "no-refresh" { + before.token_response.0.set_refresh_token(None); + save_oauth_tokens( + SERVER_NAME, + &before, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + ) + .await?; + } + let refresh_response = match scenario.as_str() { + "refreshable" => ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "renewed-token", "token_type": "Bearer", "expires_in": 3600, + })), + "malformed" | "early-malformed" => { + ResponseTemplate::new(200).set_body_raw("{", "application/json") + } + "transient" | "early-transient" => { + ResponseTemplate::new(503).set_body_json(json!({"error": "temporarily_unavailable"})) + } + _ => ResponseTemplate::new(400).set_body_json(json!({ + "error": "invalid_grant", "error_description": "private provider details", + })), + }; + Mock::given(method("POST")) + .and(path("/token")) + .and(body_string_contains("grant_type=refresh_token")) + .respond_with(refresh_response) + .expect( + if matches!( + scenario.as_str(), + "no-refresh" | "unexpired" | "unknown-expiry" | "revoked" + ) { + 0 + } else { + 1 + }, + ) + .mount(&authorization) + .await; + + let result = client + .call_tool( + "echo".to_string(), + /*arguments*/ None, + /*meta*/ None, + Some(Duration::from_secs(5)), + ) + .await; + match scenario.as_str() { + "no-refresh" | "rejected" | "malformed" | "transient" => { + let mut expected = CallToolResult::error(vec![ContentBlock::text( + "MCP authentication required. Reconnect to continue using this server.", + )]); + expected.meta = Some( + serde_json::Map::from_iter([( + "mcp/www_authenticate".to_string(), + json!("Bearer error=\"invalid_token\""), + )]) + .into(), + ); + assert_eq!(result?, expected); + } + "revoked" => { + let mut expected = + CallToolResult::error(vec![ContentBlock::text("Authentication required")]); + expected.meta = Some( + serde_json::Map::from_iter([( + "mcp/www_authenticate".to_string(), + json!(["Bearer error=\"invalid_token\""]), + )]) + .into(), + ); + assert_eq!(result?, expected); + } + "early-malformed" | "early-transient" => { + let error = result.expect_err("early refresh failures must remain ordinary errors"); + assert!(!is_authentication_required_error(&error)); + } + "refreshable" | "unexpired" | "unknown-expiry" => { + let result = result?; + assert_eq!(result.content, vec![ContentBlock::text("tool ran")]); + assert_eq!(result.meta, None); + } + _ => unreachable!("scenario selected by parent test"), + } + let tool_calls = mcp + .received_requests() + .await + .expect("request recording") + .iter() + .filter(|request| { + request + .body_json::() + .ok() + .is_some_and(|body| body["method"] == "tools/call") + }) + .count(); + assert_eq!( + tool_calls, + usize::from(matches!( + scenario.as_str(), + "refreshable" | "unexpired" | "unknown-expiry" | "revoked" + )) + ); + let after = stored_oauth_credentials( + SERVER_NAME, + &server_url, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + )? + .expect("credentials are not deleted"); + if scenario != "refreshable" { + assert_eq!(before, after); + } + authorization.verify().await; + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/streamable_http_recovery.rs b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs new file mode 100644 index 0000000000000000000000000000000000000000..817587a9b76a16e629d2b86e3a0caf1084b058aa --- /dev/null +++ b/codex-rs/rmcp-client/tests/streamable_http_recovery.rs @@ -0,0 +1,453 @@ +mod streamable_http_test_support; + +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_exec_server::Environment; +use codex_exec_server::ExecServerError; +use codex_exec_server::HttpClient; +use codex_exec_server::HttpRequestParams; +use codex_exec_server::HttpRequestResponse; +use codex_exec_server::HttpResponseBodyStream; +use futures::FutureExt as _; +use futures::future::BoxFuture; +use pretty_assertions::assert_eq; +use rmcp::model::CallToolResult; +use rmcp::model::ContentBlock; +use rmcp::model::MetaObject; +use serde_json::Value; +use serde_json::json; + +use streamable_http_test_support::arm_initialize_post_failure; +use streamable_http_test_support::arm_initialize_post_json_rpc_failure; +use streamable_http_test_support::arm_initialized_notification_post_json_rpc_failure; +use streamable_http_test_support::arm_session_post_failure; +use streamable_http_test_support::arm_session_post_json_rpc_failure; +use streamable_http_test_support::call_echo_tool; +use streamable_http_test_support::create_client; +use streamable_http_test_support::create_client_with_http_client; +use streamable_http_test_support::expected_echo_result; +use streamable_http_test_support::spawn_streamable_http_server; + +const JSON_RPC_INTERNAL_ERROR_CODE: i64 = -32603; +const SIMULATED_NO_RESPONSE_MESSAGE: &str = + "http/request failed: error sending request for url (simulated no response)"; + +#[derive(Clone)] +struct FailFirstInitializeHttpClient { + inner: Arc, + failures_remaining: Arc, + initialize_attempts: Arc, +} + +impl FailFirstInitializeHttpClient { + fn new(inner: Arc, failures_remaining: usize) -> Self { + Self { + inner, + failures_remaining: Arc::new(AtomicUsize::new(failures_remaining)), + initialize_attempts: Arc::new(AtomicUsize::new(0)), + } + } + + fn initialize_attempts(&self) -> usize { + self.initialize_attempts.load(Ordering::SeqCst) + } + + fn fail_next_initialize(&self) { + self.failures_remaining.store(1, Ordering::SeqCst); + } +} + +impl HttpClient for FailFirstInitializeHttpClient { + fn http_request( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result> { + self.inner.http_request(params) + } + + fn http_request_stream( + &self, + params: HttpRequestParams, + ) -> BoxFuture<'_, Result<(HttpRequestResponse, HttpResponseBodyStream), ExecServerError>> { + let inner = Arc::clone(&self.inner); + let failures_remaining = Arc::clone(&self.failures_remaining); + let initialize_attempts = Arc::clone(&self.initialize_attempts); + + async move { + if is_initialize_post(¶ms) { + initialize_attempts.fetch_add(1, Ordering::SeqCst); + if failures_remaining.swap(0, Ordering::SeqCst) > 0 { + return Err(ExecServerError::Server { + code: JSON_RPC_INTERNAL_ERROR_CODE, + message: SIMULATED_NO_RESPONSE_MESSAGE.to_string(), + }); + } + } + + inner.http_request_stream(params).await + } + .boxed() + } +} + +fn is_initialize_post(params: &HttpRequestParams) -> bool { + params.method.eq_ignore_ascii_case("POST") + && params + .body + .as_ref() + .and_then(|body| serde_json::from_slice::(&body.0).ok()) + .and_then(|body| { + body.get("method") + .and_then(Value::as_str) + .map(|method| method == "initialize") + }) + .unwrap_or(false) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_initialize_retries_remote_no_response_error() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let http_client = FailFirstInitializeHttpClient::new( + Environment::default_for_tests().get_http_client(), + /*failures_remaining*/ 1, + ); + + let client = create_client_with_http_client(&base_url, Arc::new(http_client.clone())).await?; + let result = call_echo_tool(&client, "after-init-retry").await?; + + assert_eq!(http_client.initialize_attempts(), 2); + assert_eq!(result, expected_echo_result("after-init-retry")); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_initialize_retries_transient_http_status() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + + arm_initialize_post_failure(&base_url, /*status*/ 502, /*remaining*/ 1).await?; + + let client = create_client(&base_url).await?; + let result = call_echo_tool(&client, "after-status-retry").await?; + + assert_eq!(result, expected_echo_result("after-status-retry")); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_initialize_retries_json_rpc_transient_status() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + + arm_initialize_post_json_rpc_failure(&base_url, /*status*/ 502, /*remaining*/ 1).await?; + + let client = create_client(&base_url).await?; + let result = call_echo_tool(&client, "after-json-status-retry").await?; + + assert_eq!(result, expected_echo_result("after-json-status-retry")); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_retries_initialized_notification_status() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + + arm_initialized_notification_post_json_rpc_failure( + &base_url, /*status*/ 502, /*remaining*/ 1, + ) + .await?; + + let client = create_client(&base_url).await?; + let result = call_echo_tool(&client, "after-notification-status-retry").await?; + + assert_eq!( + result, + expected_echo_result("after-notification-status-retry") + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_tools_list_retries_transient_http_status() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let client = create_client(&base_url).await?; + + let expected = client + .list_tools( + /*params*/ None, + /*timeout*/ Some(Duration::from_secs(5)), + ) + .await?; + arm_session_post_failure( + &base_url, + /*status*/ 502, + /*remaining*/ 1, + /*www_authenticate_headers*/ &[], + ) + .await?; + + let result = client + .list_tools( + /*params*/ None, + /*timeout*/ Some(Duration::from_secs(5)), + ) + .await?; + + assert_eq!(result, expected); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_tools_list_retries_json_rpc_transient_status() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let client = create_client(&base_url).await?; + + let expected = client + .list_tools( + /*params*/ None, + /*timeout*/ Some(Duration::from_secs(5)), + ) + .await?; + arm_session_post_json_rpc_failure(&base_url, /*status*/ 502, /*remaining*/ 1).await?; + + let result = client + .list_tools( + /*params*/ None, + /*timeout*/ Some(Duration::from_secs(5)), + ) + .await?; + + assert_eq!(result, expected); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_404_session_expiry_recovers_and_retries_once() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let client = create_client(&base_url).await?; + + let warmup = call_echo_tool(&client, "warmup").await?; + assert_eq!(warmup, expected_echo_result("warmup")); + + arm_session_post_failure( + &base_url, + /*status*/ 404, + /*remaining*/ 1, + /*www_authenticate_headers*/ &[], + ) + .await?; + + let recovered = call_echo_tool(&client, "recovered").await?; + assert_eq!(recovered, expected_echo_result("recovered")); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_session_recovery_retries_initialize_failure() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let http_client = FailFirstInitializeHttpClient::new( + Environment::default_for_tests().get_http_client(), + /*failures_remaining*/ 0, + ); + let client = create_client_with_http_client(&base_url, Arc::new(http_client.clone())).await?; + + let warmup = call_echo_tool(&client, "warmup").await?; + assert_eq!(warmup, expected_echo_result("warmup")); + + arm_session_post_failure( + &base_url, + /*status*/ 404, + /*remaining*/ 1, + /*www_authenticate_headers*/ &[], + ) + .await?; + http_client.fail_next_initialize(); + + let recovered = call_echo_tool(&client, "recovered-after-retry").await?; + assert_eq!(http_client.initialize_attempts(), 3); + assert_eq!(recovered, expected_echo_result("recovered-after-retry")); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_401_does_not_trigger_recovery() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let client = create_client(&base_url).await?; + + let warmup = call_echo_tool(&client, "warmup").await?; + assert_eq!(warmup, expected_echo_result("warmup")); + + arm_session_post_failure( + &base_url, + /*status*/ 401, + /*remaining*/ 2, + /*www_authenticate_headers*/ &[], + ) + .await?; + + let first_error = call_echo_tool(&client, "unauthorized").await.unwrap_err(); + assert!(first_error.to_string().contains("401")); + + let second_error = call_echo_tool(&client, "still-unauthorized") + .await + .unwrap_err(); + assert!(second_error.to_string().contains("401")); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_401_challenge_is_a_tool_error_without_replay() -> anyhow::Result<()> { + let challenge = r#"Bearer error="invalid_token", resource_metadata="https://example.com/.well-known/oauth-protected-resource""#; + for challenges in [ + vec![challenge], + vec![r#"Basic realm="proxy, login""#, challenge], + ] { + let (_server, base_url) = spawn_streamable_http_server().await?; + let client = create_client(&base_url).await?; + + arm_session_post_failure( + &base_url, + /*status*/ 401, + /*remaining*/ 1, + &challenges, + ) + .await?; + + let result = call_echo_tool(&client, "rejected").await?; + let mut expected = + CallToolResult::error(vec![ContentBlock::text("Authentication required")]); + expected.meta = Some(MetaObject::from(serde_json::Map::from_iter([( + "mcp/www_authenticate".to_string(), + json!([challenges.join(", ")]), + )]))); + assert_eq!(result, expected); + + // A retry inside the failed call would consume the single rejection and return success. + assert_eq!( + call_echo_tool(&client, "next-user-call").await?, + expected_echo_result("next-user-call"), + ); + } + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_403_scope_challenge_returns_insufficient_scope() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let client = create_client(&base_url).await?; + + let warmup = call_echo_tool(&client, "warmup").await?; + assert_eq!(warmup, expected_echo_result("warmup")); + + arm_session_post_failure( + &base_url, + /*status*/ 403, + /*remaining*/ 1, + /*www_authenticate_headers*/ + &[r#"Bearer error="insufficient_scope", scope="files:read files:write""#], + ) + .await?; + + let error = call_echo_tool(&client, "forbidden").await.unwrap_err(); + assert!( + error.to_string().contains("Insufficient scope"), + "expected insufficient-scope transport error, got: {error:#}" + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_403_finds_bearer_challenge_in_later_header_value() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let client = create_client(&base_url).await?; + + let warmup = call_echo_tool(&client, "warmup").await?; + assert_eq!(warmup, expected_echo_result("warmup")); + + arm_session_post_failure( + &base_url, + /*status*/ 403, + /*remaining*/ 1, + /*www_authenticate_headers*/ + &[ + r#"Basic realm="example""#, + r#"Bearer error="insufficient_scope", scope="files:read""#, + ], + ) + .await?; + + let error = call_echo_tool(&client, "forbidden").await.unwrap_err(); + assert!( + error.to_string().contains("Insufficient scope"), + "expected insufficient-scope transport error, got: {error:#}" + ); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_404_recovery_only_retries_once() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let client = create_client(&base_url).await?; + + let warmup = call_echo_tool(&client, "warmup").await?; + assert_eq!(warmup, expected_echo_result("warmup")); + + arm_session_post_failure( + &base_url, + /*status*/ 404, + /*remaining*/ 2, + /*www_authenticate_headers*/ &[], + ) + .await?; + + let error = call_echo_tool(&client, "double-404").await.unwrap_err(); + let error_message = error.to_string(); + assert!( + error_message.contains("404") || error_message.contains("session expired"), + "expected session-expiry error, got: {error:#}" + ); + + let recovered = call_echo_tool(&client, "after-double-404").await?; + assert_eq!(recovered, expected_echo_result("after-double-404")); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_non_session_failure_does_not_trigger_recovery() -> anyhow::Result<()> { + let (_server, base_url) = spawn_streamable_http_server().await?; + let client = create_client(&base_url).await?; + + let warmup = call_echo_tool(&client, "warmup").await?; + assert_eq!(warmup, expected_echo_result("warmup")); + + arm_session_post_failure( + &base_url, + /*status*/ 500, + /*remaining*/ 2, + /*www_authenticate_headers*/ &[], + ) + .await?; + + let first_error = call_echo_tool(&client, "server-error").await.unwrap_err(); + assert!(first_error.to_string().contains("500")); + + let second_error = call_echo_tool(&client, "still-server-error") + .await + .unwrap_err(); + assert!(second_error.to_string().contains("500")); + + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/streamable_http_remote.rs b/codex-rs/rmcp-client/tests/streamable_http_remote.rs new file mode 100644 index 0000000000000000000000000000000000000000..7df871976bc3ae13d99122a8cbb2128fd7244dae --- /dev/null +++ b/codex-rs/rmcp-client/tests/streamable_http_remote.rs @@ -0,0 +1,132 @@ +//! Integration coverage for the remote Streamable HTTP RMCP path. +//! +//! These tests exercise the orchestrator-side RMCP adapter against a real +//! `exec-server` process so HTTP requests go through the remote runtime path +//! instead of direct local `reqwest` calls. + +mod streamable_http_test_support; + +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Context; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::McpProtocolMode; +use codex_rmcp_client::RmcpClient; +use futures::FutureExt as _; +use pretty_assertions::assert_eq; +use rmcp::model::ClientCapabilities; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use serde_json::json; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::body_partial_json; +use wiremock::matchers::method; + +use streamable_http_test_support::call_echo_tool; +use streamable_http_test_support::create_remote_client; +use streamable_http_test_support::expected_echo_result; +use streamable_http_test_support::spawn_exec_server; +use streamable_http_test_support::spawn_streamable_http_server; + +/// What this tests: the RMCP remote Streamable HTTP adapter can initialize +/// a server and call a tool while every MCP HTTP request goes through a real +/// exec-server process instead of a direct reqwest transport. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn streamable_http_remote_client_round_trips_through_exec_server() -> anyhow::Result<()> { + // Phase 1: start the MCP Streamable HTTP test server and a local + // exec-server process that will own the HTTP network calls. + let (_server, base_url) = spawn_streamable_http_server().await?; + let exec_server = spawn_exec_server().await?; + + // Phase 2: create and initialize the RMCP client using the executor-backed + // Streamable HTTP transport. + let client = create_remote_client(&base_url, exec_server.client.clone()).await?; + + // Phase 3: prove the initialized client can complete a tool call and + // preserve the normal RMCP response shape. + let result = call_echo_tool(&client, "remote").await?; + assert_eq!(result, expected_echo_result("remote")); + + Ok(()) +} + +/// A timed-out MCP handshake must release the remote executor for later requests. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn streamable_http_handshake_timeout_unblocks_remote_executor() -> anyhow::Result<()> { + stalled_handshake_unblocks_remote_executor(McpProtocolMode::Legacy, "initialize").await +} + +/// MCP 2026 discovery must also release the remote executor when it times out. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn streamable_http_discovery_timeout_unblocks_remote_executor() -> anyhow::Result<()> { + stalled_handshake_unblocks_remote_executor(McpProtocolMode::V20260728, "server/discover").await +} + +async fn stalled_handshake_unblocks_remote_executor( + protocol_mode: McpProtocolMode, + expected_method: &str, +) -> anyhow::Result<()> { + let http_server = MockServer::start().await; + Mock::given(method("POST")) + .and(body_partial_json(json!({ "method": expected_method }))) + .respond_with(ResponseTemplate::new(200).set_delay(Duration::from_secs(10))) + .expect(1) + .mount(&http_server) + .await; + let exec_server = spawn_exec_server().await?; + let client = RmcpClient::new_streamable_http_client_with_protocol_mode( + "stalled-remote-http", + &format!("{}/mcp", http_server.uri()), + Some("test-bearer".to_string()), + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Arc::new(exec_server.client.clone()), + /*auth_provider*/ None, + protocol_mode, + ) + .await?; + let params = InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("codex-test", "0.0.0-test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18); + let error = client + .initialize( + params, + Some(Duration::from_millis(500)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: None, + meta: None, + }) + } + .boxed() + }), + ) + .await + .err() + .context("the stalled MCP handshake must time out")?; + assert!( + error.to_string().contains("timed out"), + "unexpected MCP handshake error: {error}" + ); + + tokio::time::timeout( + Duration::from_secs(2), + exec_server.client.environment_status(), + ) + .await + .context("a timed-out handshake must not leave the serial executor blocked")??; + Ok(()) +} diff --git a/codex-rs/rmcp-client/tests/streamable_http_test_support.rs b/codex-rs/rmcp-client/tests/streamable_http_test_support.rs new file mode 100644 index 0000000000000000000000000000000000000000..e7d6e2c52ffc7c18a6f91d0428f3e3087d8fef07 --- /dev/null +++ b/codex-rs/rmcp-client/tests/streamable_http_test_support.rs @@ -0,0 +1,435 @@ +//! Shared helpers for Streamable HTTP RMCP integration tests. +//! +//! This support module starts the test HTTP server, launches a real +//! `exec-server` when remote coverage is needed, and provides small helpers for +//! creating RMCP clients and asserting round-trip behavior. +//! HTTP server startup uses the address published by the child that owns the listener. + +// This support module is included by multiple integration-test crates. Each +// crate uses a different subset of the helpers, so dead-code warnings would +// otherwise depend on which test file compiled the module. +#![allow(dead_code)] + +use std::net::SocketAddr; +use std::path::Path; +use std::path::PathBuf; +use std::process::Stdio; +use std::sync::Arc; +use std::time::Duration; +use std::time::Instant; + +use anyhow::Context as _; +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_exec_server::ExecServerClient; +use codex_exec_server::HttpClient; +use codex_exec_server::RemoteExecServerConnectArgs; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_rmcp_client::ElicitationAction; +use codex_rmcp_client::ElicitationResponse; +use codex_rmcp_client::RmcpClient; +use codex_utils_cargo_bin::CargoBinError; +use futures::FutureExt as _; +use pretty_assertions::assert_eq; +use rmcp::model::CallToolResult; +use rmcp::model::ClientCapabilities; +use rmcp::model::ElicitationCapability; +use rmcp::model::FormElicitationCapability; +use rmcp::model::Implementation; +use rmcp::model::InitializeRequestParams; +use rmcp::model::ProtocolVersion; +use serde_json::json; +use tempfile::TempDir; +use tokio::io::AsyncBufReadExt; +use tokio::io::BufReader; +use tokio::process::Child; +use tokio::process::Command; +use tokio::time::sleep; + +const SESSION_POST_FAILURE_CONTROL_PATH: &str = "/test/control/session-post-failure"; +const INITIALIZE_POST_FAILURE_CONTROL_PATH: &str = "/test/control/initialize-post-failure"; +const INITIALIZED_NOTIFICATION_POST_FAILURE_CONTROL_PATH: &str = + "/test/control/initialized-notification-post-failure"; + +fn streamable_http_server_bin() -> Result { + codex_utils_cargo_bin::cargo_bin("test_streamable_http_server") +} + +fn init_params() -> InitializeRequestParams { + let mut capabilities = ClientCapabilities::default(); + capabilities.elicitation = + Some(ElicitationCapability::new().with_form(FormElicitationCapability::new())); + InitializeRequestParams::new( + capabilities, + Implementation::new("codex-test", "0.0.0-test").with_title("Codex rmcp recovery test"), + ) + .with_protocol_version(ProtocolVersion::V_2025_06_18) +} + +pub(crate) fn expected_echo_result(message: &str) -> CallToolResult { + let mut result = CallToolResult::success(Vec::new()); + result.result_type = None; + result.structured_content = Some(json!({ + "echo": format!("ECHOING: {message}"), + "env": null, + })); + result +} + +pub(crate) async fn create_client(base_url: &str) -> anyhow::Result { + create_client_with_http_client(base_url, Environment::default_for_tests().get_http_client()) + .await +} + +pub(crate) async fn create_client_with_http_client( + base_url: &str, + http_client: Arc, +) -> anyhow::Result { + let client = RmcpClient::new_streamable_http_client( + "test-streamable-http", + &format!("{base_url}/mcp"), + Some("test-bearer".to_string()), + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + http_client, + /*auth_provider*/ None, + ) + .await?; + + initialize_client(&client).await?; + + Ok(client) +} + +pub(crate) async fn initialize_client(client: &RmcpClient) -> anyhow::Result<()> { + initialize_client_with_timeout(client, Duration::from_secs(/*secs*/ 5)).await +} + +pub(crate) async fn initialize_client_with_timeout( + client: &RmcpClient, + timeout: Duration, +) -> anyhow::Result<()> { + client + .initialize( + init_params(), + Some(timeout), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: None, + }) + } + .boxed() + }), + ) + .await?; + Ok(()) +} + +/// Creates a Streamable HTTP RMCP client that sends traffic through the remote +/// runtime HTTP API. +pub(crate) async fn create_remote_client( + base_url: &str, + http_client: ExecServerClient, +) -> anyhow::Result { + let client = RmcpClient::new_streamable_http_client( + "test-streamable-http-remote", + &format!("{base_url}/mcp"), + Some("test-bearer".to_string()), + /*http_headers*/ None, + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Arc::new(http_client), + /*auth_provider*/ None, + ) + .await?; + + client + .initialize( + init_params(), + Some(Duration::from_secs(5)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + Ok(client) +} + +pub(crate) async fn call_echo_tool( + client: &RmcpClient, + message: &str, +) -> anyhow::Result { + client + .call_tool( + "echo".to_string(), + Some(json!({ "message": message })), + /*meta*/ None, + Some(Duration::from_secs(5)), + ) + .await +} + +fn control_client() -> anyhow::Result { + Ok(codex_http_client::HttpClientBuilder::new().build_direct()?) +} + +pub(crate) async fn arm_session_post_failure( + base_url: &str, + status: u16, + remaining: usize, + www_authenticate_headers: &[&str], +) -> anyhow::Result<()> { + let response = control_client()? + .post(format!("{base_url}{SESSION_POST_FAILURE_CONTROL_PATH}")) + .json(&json!({ + "status": status, + "remaining": remaining, + "www_authenticate_headers": www_authenticate_headers, + })) + .send() + .await?; + + assert_eq!(response.status(), http::StatusCode::NO_CONTENT); + Ok(()) +} + +pub(crate) async fn arm_session_post_json_rpc_failure( + base_url: &str, + status: u16, + remaining: usize, +) -> anyhow::Result<()> { + let response = control_client()? + .post(format!("{base_url}{SESSION_POST_FAILURE_CONTROL_PATH}")) + .json(&json!({ + "status": status, + "remaining": remaining, + "content_type": "application/json", + "body": json!({ + "jsonrpc": "2.0", + "id": 1, + "error": { + "code": -32000, + "message": "transient session failure", + }, + }).to_string(), + })) + .send() + .await?; + + assert_eq!(response.status(), http::StatusCode::NO_CONTENT); + Ok(()) +} + +pub(crate) async fn arm_initialized_notification_post_json_rpc_failure( + base_url: &str, + status: u16, + remaining: usize, +) -> anyhow::Result<()> { + let response = control_client()? + .post(format!( + "{base_url}{INITIALIZED_NOTIFICATION_POST_FAILURE_CONTROL_PATH}" + )) + .json(&json!({ + "status": status, + "remaining": remaining, + "content_type": "application/json", + "body": json!({ + "jsonrpc": "2.0", + "id": 1, + "error": { + "code": -32000, + "message": "transient session failure", + }, + }).to_string(), + })) + .send() + .await?; + + assert_eq!(response.status(), http::StatusCode::NO_CONTENT); + Ok(()) +} + +pub(crate) async fn arm_initialize_post_failure( + base_url: &str, + status: u16, + remaining: usize, +) -> anyhow::Result<()> { + let response = control_client()? + .post(format!("{base_url}{INITIALIZE_POST_FAILURE_CONTROL_PATH}")) + .json(&json!({ + "status": status, + "remaining": remaining, + })) + .send() + .await?; + + assert_eq!(response.status(), http::StatusCode::NO_CONTENT); + Ok(()) +} + +pub(crate) async fn arm_initialize_post_json_rpc_failure( + base_url: &str, + status: u16, + remaining: usize, +) -> anyhow::Result<()> { + let response = control_client()? + .post(format!("{base_url}{INITIALIZE_POST_FAILURE_CONTROL_PATH}")) + .json(&json!({ + "status": status, + "remaining": remaining, + "content_type": "application/json", + "body": json!({ + "jsonrpc": "2.0", + "id": 1, + "error": { + "code": -32000, + "message": "transient initialize failure", + }, + }).to_string(), + })) + .send() + .await?; + + assert_eq!(response.status(), http::StatusCode::NO_CONTENT); + Ok(()) +} + +pub(crate) async fn spawn_streamable_http_server() -> anyhow::Result<(Child, String)> { + let startup_dir = tempfile::tempdir()?; + let bound_addr_path = startup_dir.path().join("bound_addr"); + // Let the child reserve its own port; releasing a parent-owned port before + // spawning can let another test claim it and satisfy the readiness probe. + let mut child = Command::new(streamable_http_server_bin()?) + .kill_on_drop(true) + .env("MCP_STREAMABLE_HTTP_BIND_ADDR", "127.0.0.1:0") + .env("MCP_STREAMABLE_HTTP_BOUND_ADDR_FILE", &bound_addr_path) + .spawn()?; + + let address = wait_for_streamable_http_server( + &mut child, + &bound_addr_path, + Duration::from_secs(/*secs*/ 5), + ) + .await?; + Ok((child, format!("http://{address}"))) +} + +/// Owns the exec-server process used by the remote-client integration test. +pub(crate) struct ExecServerProcess { + _codex_home: TempDir, + child: Child, + pub(crate) client: ExecServerClient, +} + +impl Drop for ExecServerProcess { + /// Stops the local exec-server process best-effort when the test exits. + fn drop(&mut self) { + let _ = self.child.start_kill(); + } +} + +/// Starts a local exec-server and connects an initialized `ExecServerClient`. +pub(crate) async fn spawn_exec_server() -> anyhow::Result { + let codex_home = TempDir::new()?; + let mut child = Command::new(codex_utils_cargo_bin::cargo_bin("codex")?) + .args(["exec-server", "--listen", "ws://127.0.0.1:0"]) + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::inherit()) + .kill_on_drop(true) + .env("CODEX_HOME", codex_home.path()) + .spawn()?; + + let websocket_url = read_exec_server_listen_url(&mut child).await?; + let client = ExecServerClient::connect_websocket(RemoteExecServerConnectArgs::new( + websocket_url, + "rmcp-client-remote-http-test".to_string(), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + )) + .await?; + + Ok(ExecServerProcess { + _codex_home: codex_home, + child, + client, + }) +} + +/// Reads the websocket URL printed by `codex exec-server --listen`. +async fn read_exec_server_listen_url(child: &mut Child) -> anyhow::Result { + let stdout = child + .stdout + .take() + .context("failed to capture exec-server stdout")?; + let mut lines = BufReader::new(stdout).lines(); + let deadline = Instant::now() + Duration::from_secs(10); + + loop { + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + anyhow::bail!("timed out waiting for exec-server listen URL"); + } + + let line = tokio::time::timeout(remaining, lines.next_line()) + .await + .context("timed out waiting for exec-server stdout")?? + .context("exec-server stdout closed before emitting listen URL")?; + let listen_url = line.trim(); + if listen_url.starts_with("ws://") { + return Ok(listen_url.to_string()); + } + } +} + +async fn wait_for_streamable_http_server( + server_child: &mut Child, + bound_addr_path: &Path, + timeout: Duration, +) -> anyhow::Result { + let deadline = Instant::now() + timeout; + + loop { + if let Some(status) = server_child.try_wait()? { + return Err(anyhow::anyhow!( + "streamable HTTP server exited early with status {status}" + )); + } + + if Instant::now() >= deadline { + return Err(anyhow::anyhow!( + "timed out waiting for streamable HTTP server to publish its bound address" + )); + } + + match std::fs::read_to_string(bound_addr_path) { + Ok(contents) => { + // The child creates the file before writing the address into it. + if let Ok(address) = contents.parse() { + return Ok(address); + } + } + Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} + Err(error) => { + return Err(error).context("failed to read streamable HTTP server bound address"); + } + } + + sleep(Duration::from_millis(50)).await; + } +} diff --git a/codex-rs/rmcp-client/tests/streamable_http_user_agent.rs b/codex-rs/rmcp-client/tests/streamable_http_user_agent.rs new file mode 100644 index 0000000000000000000000000000000000000000..45b5d6025cdd629f690a5c728303c6371bf9fe25 --- /dev/null +++ b/codex-rs/rmcp-client/tests/streamable_http_user_agent.rs @@ -0,0 +1,70 @@ +mod streamable_http_test_support; + +use std::collections::HashMap; + +use codex_config::types::AuthKeyringBackendKind; +use codex_config::types::OAuthCredentialsStoreMode; +use codex_exec_server::Environment; +use codex_rmcp_client::RmcpClient; +use serde_json::Value; +use serde_json::json; +use streamable_http_test_support::initialize_client; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::Request; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +async fn streamable_http_requests_preserve_configured_user_agent() -> anyhow::Result<()> { + let custom_user_agent = "custom-agent/9.9"; + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/mcp")) + .and(header("user-agent", custom_user_agent)) + .respond_with(|request: &Request| { + let body: Value = request.body_json().expect("valid JSON-RPC request"); + match body.get("method").and_then(Value::as_str) { + Some("initialize") => ResponseTemplate::new(200).set_body_json(json!({ + "jsonrpc": "2.0", + "id": body.get("id").cloned().unwrap_or(Value::Null), + "result": { + "protocolVersion": "2025-06-18", + "capabilities": {}, + "serverInfo": { + "name": "user-agent-test", + "version": "0.0.0-test", + }, + }, + })), + Some("notifications/initialized") => ResponseTemplate::new(202), + method => ResponseTemplate::new(400) + .set_body_string(format!("unexpected JSON-RPC method: {method:?}")), + } + }) + .expect(2) + .mount(&server) + .await; + + let custom_client = RmcpClient::new_streamable_http_client( + "test-streamable-http", + &format!("{}/mcp", server.uri()), + Some("test-bearer".to_string()), + Some(HashMap::from([( + "user-agent".to_string(), + custom_user_agent.to_string(), + )])), + /*env_http_headers*/ None, + OAuthCredentialsStoreMode::File, + AuthKeyringBackendKind::default(), + Environment::default_for_tests().get_http_client(), + /*auth_provider*/ None, + ) + .await?; + initialize_client(&custom_client).await?; + server.verify().await; + + Ok(()) +} diff --git a/codex-rs/rollout-trace/src/bundle.rs b/codex-rs/rollout-trace/src/bundle.rs new file mode 100644 index 0000000000000000000000000000000000000000..8f67da3b3bfaf3d8b260ed293174043629c5f154 --- /dev/null +++ b/codex-rs/rollout-trace/src/bundle.rs @@ -0,0 +1,49 @@ +//! Trace bundle manifest and local layout constants. + +use serde::Deserialize; +use serde::Serialize; + +use crate::model::AgentThreadId; + +pub(crate) const MANIFEST_FILE_NAME: &str = "manifest.json"; +pub(crate) const RAW_EVENT_LOG_FILE_NAME: &str = "trace.jsonl"; +pub(crate) const PAYLOADS_DIR_NAME: &str = "payloads"; +/// Conventional file name for a reducer-written `RolloutTrace` cache. +pub const REDUCED_STATE_FILE_NAME: &str = "state.json"; +pub(crate) const TRACE_MANIFEST_SCHEMA_VERSION: u32 = 1; +pub(crate) const REDUCED_TRACE_SCHEMA_VERSION: u32 = 1; + +/// Manifest stored at the root of a trace bundle. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub(crate) struct TraceBundleManifest { + pub(crate) schema_version: u32, + pub(crate) trace_id: String, + pub(crate) rollout_id: String, + /// Root thread for the recorded rollout. Replay should fail rather than + /// inventing a placeholder, because every reduced object is scoped back to + /// this thread tree. + pub(crate) root_thread_id: AgentThreadId, + pub(crate) started_at_unix_ms: i64, + pub(crate) raw_event_log: String, + pub(crate) payloads_dir: String, +} + +impl TraceBundleManifest { + /// Builds a manifest that uses the standard local bundle layout. + pub(crate) fn new( + trace_id: String, + rollout_id: String, + root_thread_id: AgentThreadId, + started_at_unix_ms: i64, + ) -> Self { + Self { + schema_version: TRACE_MANIFEST_SCHEMA_VERSION, + trace_id, + rollout_id, + root_thread_id, + started_at_unix_ms, + raw_event_log: RAW_EVENT_LOG_FILE_NAME.to_string(), + payloads_dir: PAYLOADS_DIR_NAME.to_string(), + } + } +} diff --git a/codex-rs/rollout-trace/src/code_cell.rs b/codex-rs/rollout-trace/src/code_cell.rs new file mode 100644 index 0000000000000000000000000000000000000000..5f2603b70adf85a19d69459137f08de4cd7bfeb1 --- /dev/null +++ b/codex-rs/rollout-trace/src/code_cell.rs @@ -0,0 +1,185 @@ +//! Hot-path helpers for recording code-mode runtime cell lifecycles. +//! +//! The public `exec` tool is reduced as a first-class `CodeCell` instead of a +//! generic tool call. This module keeps the runtime response serialization and +//! lifecycle event policy inside the trace crate while core carries a compact, +//! no-op capable handle through execution and waits. + +use std::sync::Arc; + +use codex_code_mode::RuntimeResponse; +use serde::Serialize; +use tracing::warn; + +use crate::model::AgentThreadId; +use crate::model::CodeCellRuntimeStatus; +use crate::model::CodexTurnId; +use crate::model::ModelVisibleCallId; +use crate::payload::RawPayloadKind; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawTraceEventContext; +use crate::raw_event::RawTraceEventPayload; +use crate::writer::TraceWriter; + +/// No-op capable trace handle for one code-mode runtime cell. +#[derive(Clone, Debug)] +pub struct CodeCellTraceContext { + state: CodeCellTraceContextState, +} + +#[derive(Clone, Debug)] +enum CodeCellTraceContextState { + Disabled, + Enabled(EnabledCodeCellTraceContext), +} + +#[derive(Clone, Debug)] +struct EnabledCodeCellTraceContext { + writer: Arc, + thread_id: AgentThreadId, + codex_turn_id: CodexTurnId, + runtime_cell_id: String, +} + +/// Raw code-mode response captured at the runtime boundary. +/// +/// This is not the model-visible custom-tool output. The reducer links that +/// output through `CodeCell.output_item_ids` once the conversation item appears. +/// Keeping the raw runtime payload here preserves stored-value and lifecycle +/// evidence without duplicating the model-facing transcript. +#[derive(Serialize)] +struct CodeCellResponseTracePayload<'a> { + response: &'a RuntimeResponse, +} + +impl CodeCellTraceContext { + /// Builds a context that accepts trace calls and records nothing. + pub(crate) fn disabled() -> Self { + Self { + state: CodeCellTraceContextState::Disabled, + } + } + + /// Builds a context for an already-known code-mode runtime cell. + pub(crate) fn enabled( + writer: Arc, + thread_id: impl Into, + codex_turn_id: impl Into, + runtime_cell_id: impl Into, + ) -> Self { + Self { + state: CodeCellTraceContextState::Enabled(EnabledCodeCellTraceContext { + writer, + thread_id: thread_id.into(), + codex_turn_id: codex_turn_id.into(), + runtime_cell_id: runtime_cell_id.into(), + }), + } + } + + /// Records the parent runtime object before JavaScript can issue nested tool calls. + pub fn record_started( + &self, + model_visible_call_id: impl Into, + source_js: impl Into, + ) { + let CodeCellTraceContextState::Enabled(context) = &self.state else { + return; + }; + append_with_context_best_effort( + context, + RawTraceEventPayload::CodeCellStarted { + runtime_cell_id: context.runtime_cell_id.clone(), + model_visible_call_id: model_visible_call_id.into(), + source_js: source_js.into(), + }, + ); + } + + /// Records the first response returned by the public code-mode `exec` tool. + /// + /// A yielded response returns control to the model while the cell keeps + /// running. Terminal initial responses should be followed by `record_ended` + /// by the caller so the reducer can distinguish model-visible output from + /// runtime completion. + pub fn record_initial_response(&self, response: &RuntimeResponse) { + let CodeCellTraceContextState::Enabled(context) = &self.state else { + return; + }; + append_with_context_best_effort( + context, + RawTraceEventPayload::CodeCellInitialResponse { + runtime_cell_id: context.runtime_cell_id.clone(), + status: code_cell_status_for_runtime_response(response), + response_payload: code_cell_response_payload(context, response), + }, + ); + } + + /// Records the terminal lifecycle point for a code-mode runtime cell. + pub fn record_ended(&self, response: &RuntimeResponse) { + let CodeCellTraceContextState::Enabled(context) = &self.state else { + return; + }; + append_with_context_best_effort( + context, + RawTraceEventPayload::CodeCellEnded { + runtime_cell_id: context.runtime_cell_id.clone(), + status: code_cell_status_for_runtime_response(response), + response_payload: code_cell_response_payload(context, response), + }, + ); + } +} + +fn code_cell_status_for_runtime_response(response: &RuntimeResponse) -> CodeCellRuntimeStatus { + match response { + RuntimeResponse::Yielded { .. } => CodeCellRuntimeStatus::Yielded, + RuntimeResponse::Terminated { .. } => CodeCellRuntimeStatus::Terminated, + RuntimeResponse::Result { error_text, .. } => { + if error_text.is_some() { + CodeCellRuntimeStatus::Failed + } else { + CodeCellRuntimeStatus::Completed + } + } + } +} + +fn code_cell_response_payload( + context: &EnabledCodeCellTraceContext, + response: &RuntimeResponse, +) -> Option { + write_json_payload_best_effort( + &context.writer, + RawPayloadKind::ToolResult, + &CodeCellResponseTracePayload { response }, + ) +} + +fn write_json_payload_best_effort( + writer: &TraceWriter, + kind: RawPayloadKind, + payload: &impl Serialize, +) -> Option { + match writer.write_json_payload(kind, payload) { + Ok(payload_ref) => Some(payload_ref), + Err(err) => { + warn!("failed to write rollout trace payload: {err:#}"); + None + } + } +} + +fn append_with_context_best_effort( + context: &EnabledCodeCellTraceContext, + payload: RawTraceEventPayload, +) { + let event_context = RawTraceEventContext { + thread_id: Some(context.thread_id.clone()), + codex_turn_id: Some(context.codex_turn_id.clone()), + }; + if let Err(err) = context.writer.append_with_context(event_context, payload) { + warn!("failed to append rollout trace event: {err:#}"); + } +} diff --git a/codex-rs/rollout-trace/src/compaction.rs b/codex-rs/rollout-trace/src/compaction.rs new file mode 100644 index 0000000000000000000000000000000000000000..c8f18795abac178da5ab1cb2cb312fed6a319eb3 --- /dev/null +++ b/codex-rs/rollout-trace/src/compaction.rs @@ -0,0 +1,284 @@ +//! Hot-path helpers for recording upstream remote compaction attempts. +//! +//! Remote compaction is a model-facing request with a different semantic role +//! from normal sampling. Keeping the no-op capable trace handle in this crate +//! lets `codex-core` record exact endpoint payloads without owning trace schema +//! details. + +use std::fmt::Display; +use std::sync::Arc; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; + +use codex_protocol::models::ResponseItem; +use serde::Serialize; +use serde_json::Value as JsonValue; +use tracing::warn; + +use crate::inference::trace_response_item_json; +use crate::model::AgentThreadId; +use crate::model::CodexTurnId; +use crate::model::CompactionId; +use crate::model::CompactionRequestId; +use crate::payload::RawPayloadKind; +use crate::raw_event::RawTraceEventContext; +use crate::raw_event::RawTraceEventPayload; +use crate::writer::TraceWriter; + +static NEXT_COMPACTION_REQUEST: AtomicU64 = AtomicU64::new(1); + +/// Turn-local remote compaction tracing context. +/// +/// A compaction can retry its upstream request before installing one checkpoint. The context +/// owns the stable checkpoint ID; each request attempt gets a separate request ID. +#[derive(Clone, Debug)] +pub struct CompactionTraceContext { + state: CompactionTraceContextState, +} + +#[derive(Clone, Debug)] +enum CompactionTraceContextState { + Disabled, + Enabled(EnabledCompactionTraceContext), +} + +#[derive(Clone, Debug)] +struct EnabledCompactionTraceContext { + writer: Arc, + thread_id: AgentThreadId, + codex_turn_id: CodexTurnId, + compaction_id: CompactionId, + model: String, + provider_name: String, +} + +/// One upstream request attempt made while computing a compaction checkpoint. +#[derive(Clone, Debug)] +pub struct CompactionTraceAttempt { + state: CompactionTraceAttemptState, +} + +#[derive(Clone, Debug)] +enum CompactionTraceAttemptState { + Disabled, + Enabled(EnabledCompactionTraceAttempt), +} + +#[derive(Clone, Debug)] +struct EnabledCompactionTraceAttempt { + context: EnabledCompactionTraceContext, + compaction_request_id: CompactionRequestId, +} + +#[derive(Serialize)] +struct TracedCompactionCompleted { + output_items: Vec, +} + +/// History replacement checkpoint persisted when compaction installs new live history. +/// +/// The checkpoint keeps compaction separate from ordinary sampling snapshots: +/// `input_history` is the live thread history selected for compaction, while +/// `replacement_history` is what future prompts may carry after the checkpoint. +#[derive(Serialize)] +pub struct CompactionCheckpointTracePayload<'a> { + pub input_history: &'a [ResponseItem], + pub replacement_history: &'a [ResponseItem], +} + +impl CompactionTraceContext { + /// Builds a context that accepts trace calls and records nothing. + pub fn disabled() -> Self { + Self { + state: CompactionTraceContextState::Disabled, + } + } + + /// Builds an enabled context for upstream attempts that compute one checkpoint. + pub fn enabled( + writer: Arc, + thread_id: AgentThreadId, + codex_turn_id: CodexTurnId, + compaction_id: CompactionId, + model: String, + provider_name: String, + ) -> Self { + Self { + state: CompactionTraceContextState::Enabled(EnabledCompactionTraceContext { + writer, + thread_id, + codex_turn_id, + compaction_id, + model, + provider_name, + }), + } + } + + /// Returns whether this context records compaction traces. + pub fn is_enabled(&self) -> bool { + matches!(self.state, CompactionTraceContextState::Enabled(_)) + } + + /// Starts a new upstream attempt and records the exact compact endpoint request. + pub fn start_attempt(&self, request: &impl Serialize) -> CompactionTraceAttempt { + let CompactionTraceContextState::Enabled(context) = &self.state else { + return CompactionTraceAttempt::disabled(); + }; + + let attempt = CompactionTraceAttempt { + state: CompactionTraceAttemptState::Enabled(EnabledCompactionTraceAttempt { + context: context.clone(), + compaction_request_id: next_compaction_request_id(), + }), + }; + attempt.record_started(request); + attempt + } + + /// Records the point where compacted history becomes the live thread history. + /// + /// The checkpoint belongs to the same semantic compaction lifecycle as the + /// compact endpoint attempts, so the context reuses its stable compaction ID. + pub fn record_installed(&self, checkpoint: &CompactionCheckpointTracePayload<'_>) { + let CompactionTraceContextState::Enabled(context) = &self.state else { + return; + }; + let checkpoint_payload = match context + .writer + .write_json_payload(RawPayloadKind::CompactionCheckpoint, checkpoint) + { + Ok(payload_ref) => payload_ref, + Err(err) => { + warn!("failed to write rollout trace payload: {err:#}"); + return; + } + }; + + let event_context = RawTraceEventContext { + thread_id: Some(context.thread_id.clone()), + codex_turn_id: Some(context.codex_turn_id.clone()), + }; + if let Err(err) = context.writer.append_with_context( + event_context, + RawTraceEventPayload::CompactionInstalled { + compaction_id: context.compaction_id.clone(), + checkpoint_payload, + }, + ) { + warn!("failed to append rollout trace event: {err:#}"); + } + } +} + +impl CompactionTraceAttempt { + /// Builds an attempt that records nothing. + fn disabled() -> Self { + Self { + state: CompactionTraceAttemptState::Disabled, + } + } + + fn record_started(&self, request: &impl Serialize) { + let CompactionTraceAttemptState::Enabled(attempt) = &self.state else { + return; + }; + let Some(request_payload) = write_json_payload_best_effort( + &attempt.context.writer, + RawPayloadKind::CompactionRequest, + request, + ) else { + return; + }; + + append_with_context_best_effort( + &attempt.context, + RawTraceEventPayload::CompactionRequestStarted { + compaction_id: attempt.context.compaction_id.clone(), + compaction_request_id: attempt.compaction_request_id.clone(), + thread_id: attempt.context.thread_id.clone(), + codex_turn_id: attempt.context.codex_turn_id.clone(), + model: attempt.context.model.clone(), + provider_name: attempt.context.provider_name.clone(), + request_payload, + }, + ); + } + + /// Records the non-streaming compact endpoint response payload. + /// + /// Compaction responses use the same response-item preservation rules as + /// inference streams: traces are evidence, while normal ResponseItem + /// serialization is shaped for future request construction. + pub fn record_completed(&self, output_items: &[ResponseItem]) { + let CompactionTraceAttemptState::Enabled(attempt) = &self.state else { + return; + }; + let response_payload = TracedCompactionCompleted { + output_items: output_items.iter().map(trace_response_item_json).collect(), + }; + let Some(response_payload) = write_json_payload_best_effort( + &attempt.context.writer, + RawPayloadKind::CompactionResponse, + &response_payload, + ) else { + return; + }; + + append_with_context_best_effort( + &attempt.context, + RawTraceEventPayload::CompactionRequestCompleted { + compaction_id: attempt.context.compaction_id.clone(), + compaction_request_id: attempt.compaction_request_id.clone(), + response_payload, + }, + ); + } + + /// Records the compact endpoint result without forcing callers to branch on trace events. + pub fn record_result(&self, result: Result<&[ResponseItem], E>) { + match result { + Ok(output_items) => self.record_completed(output_items), + Err(err) => self.record_failed(err), + } + } + + /// Records pre-response failures from the compact endpoint. + pub fn record_failed(&self, error: impl Display) { + let CompactionTraceAttemptState::Enabled(attempt) = &self.state else { + return; + }; + append_with_context_best_effort( + &attempt.context, + RawTraceEventPayload::CompactionRequestFailed { + compaction_id: attempt.context.compaction_id.clone(), + compaction_request_id: attempt.compaction_request_id.clone(), + error: error.to_string(), + }, + ); + } +} + +fn next_compaction_request_id() -> CompactionRequestId { + let ordinal = NEXT_COMPACTION_REQUEST.fetch_add(1, Ordering::Relaxed); + format!("compaction_request:{ordinal}") +} + +fn write_json_payload_best_effort( + writer: &TraceWriter, + kind: RawPayloadKind, + payload: &impl Serialize, +) -> Option { + writer.write_json_payload(kind, payload).ok() +} + +fn append_with_context_best_effort( + context: &EnabledCompactionTraceContext, + payload: RawTraceEventPayload, +) { + let event_context = RawTraceEventContext { + thread_id: Some(context.thread_id.clone()), + codex_turn_id: Some(context.codex_turn_id.clone()), + }; + let _ = context.writer.append_with_context(event_context, payload); +} diff --git a/codex-rs/rollout-trace/src/inference.rs b/codex-rs/rollout-trace/src/inference.rs new file mode 100644 index 0000000000000000000000000000000000000000..7cea4d0f9752abd840ab3a6123962ffd69a64587 --- /dev/null +++ b/codex-rs/rollout-trace/src/inference.rs @@ -0,0 +1,526 @@ +//! Hot-path helpers for recording upstream inference attempts. +//! +//! The model client should not need to know whether rollout tracing is enabled. +//! A disabled context records nothing, which keeps one-shot HTTP calls, +//! WebSocket reuse, and retry/fallback attempts on the same code path. + +use std::fmt::Display; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use codex_protocol::models::ResponseItem; +use codex_protocol::protocol::TokenUsage; +use http::HeaderMap; +use http::HeaderValue; +use serde::Serialize; +use serde_json::Value as JsonValue; +use uuid::Uuid; + +use crate::model::AgentThreadId; +use crate::model::CodexTurnId; +use crate::model::InferenceCallId; +use crate::payload::RawPayloadKind; +use crate::raw_event::RawTraceEventContext; +use crate::raw_event::RawTraceEventPayload; +use crate::writer::TraceWriter; + +const INFERENCE_CALL_ID_HEADER: &str = "x-codex-inference-call-id"; + +/// Turn-local inference tracing context. +/// +/// This is intentionally a no-op capable handle instead of an `Option` at each +/// transport callsite. Whether tracing is enabled is a session concern; retry, +/// fallback, and stream mapping code should always be able to say what happened +/// without first branching on trace availability. +#[derive(Clone, Debug)] +pub struct InferenceTraceContext { + state: InferenceTraceContextState, +} + +#[derive(Clone, Debug)] +enum InferenceTraceContextState { + Disabled, + Enabled(EnabledInferenceTraceContext), +} + +#[derive(Clone, Debug)] +struct EnabledInferenceTraceContext { + writer: Arc, + thread_id: AgentThreadId, + codex_turn_id: CodexTurnId, + model: String, + provider_name: String, +} + +/// One concrete upstream request attempt. +/// +/// A Codex turn can create multiple attempts when auth recovery retries the +/// HTTP request or WebSocket setup falls back to HTTP. Completion is often +/// observed after the client returns the response stream, so the attempt owns +/// the terminal guard that prevents duplicate lifecycle events. +#[derive(Debug)] +pub struct InferenceTraceAttempt { + state: InferenceTraceAttemptState, +} + +#[derive(Debug)] +enum InferenceTraceAttemptState { + Disabled, + Enabled(EnabledInferenceTraceAttempt), +} + +#[derive(Debug)] +struct EnabledInferenceTraceAttempt { + context: EnabledInferenceTraceContext, + inference_call_id: InferenceCallId, + terminal_recorded: AtomicBool, +} + +/// Non-delta response payload saved for completed or interrupted inference streams. +/// +/// We intentionally record completed output items instead of every stream delta +/// here. The raw stream can be added later as a separate payload class; this +/// response summary gives the reducer stable response identity when available +/// plus model-visible output without duplicating high-volume text deltas. +#[derive(Serialize)] +struct TracedResponseStreamOutput<'a> { + response_id: Option<&'a str>, + upstream_request_id: Option<&'a str>, + token_usage: Option<&'a TokenUsage>, + output_items: Vec, +} + +impl InferenceTraceContext { + /// Builds a context that accepts trace calls and records nothing. + pub fn disabled() -> Self { + Self { + state: InferenceTraceContextState::Disabled, + } + } + + /// Builds an enabled context for all upstream attempts made by one Codex turn. + pub fn enabled( + writer: Arc, + thread_id: AgentThreadId, + codex_turn_id: CodexTurnId, + model: String, + provider_name: String, + ) -> Self { + Self { + state: InferenceTraceContextState::Enabled(EnabledInferenceTraceContext { + writer, + thread_id, + codex_turn_id, + model, + provider_name, + }), + } + } + + /// Starts a new attempt after the concrete provider request has been built. + pub fn start_attempt(&self) -> InferenceTraceAttempt { + let InferenceTraceContextState::Enabled(context) = &self.state else { + return InferenceTraceAttempt::disabled(); + }; + + InferenceTraceAttempt { + state: InferenceTraceAttemptState::Enabled(EnabledInferenceTraceAttempt { + context: context.clone(), + inference_call_id: next_inference_call_id(), + terminal_recorded: AtomicBool::new(false), + }), + } + } +} + +impl InferenceTraceAttempt { + /// Builds an attempt that records nothing. + pub fn disabled() -> Self { + Self { + state: InferenceTraceAttemptState::Disabled, + } + } + + fn inference_call_id(&self) -> Option<&str> { + match &self.state { + InferenceTraceAttemptState::Disabled => None, + InferenceTraceAttemptState::Enabled(attempt) => { + Some(attempt.inference_call_id.as_str()) + } + } + } + + /// Adds rollout-trace propagation headers for this attempt when tracing is enabled. + pub fn add_request_headers(&self, headers: &mut HeaderMap) { + let Some(inference_call_id) = self.inference_call_id() else { + return; + }; + let Ok(inference_call_id) = HeaderValue::from_str(inference_call_id) else { + // These IDs are generated internally as UUID strings, so rejection + // should be impossible in practice. Tracing remains best-effort, + // though, and must never make provider requests fail. + return; + }; + + headers.insert(INFERENCE_CALL_ID_HEADER, inference_call_id); + } + + /// Records the request payload replay should treat as the model-visible inference input. + /// + /// This is usually the exact provider request. Callers may instead pass a + /// logical request when the transport omits already-sent input, such as + /// websocket reuse after an untraced warmup response. + pub fn record_started(&self, request: &impl Serialize) { + let InferenceTraceAttemptState::Enabled(attempt) = &self.state else { + return; + }; + let Some(request_payload) = write_json_payload_best_effort( + &attempt.context.writer, + RawPayloadKind::InferenceRequest, + request, + ) else { + return; + }; + + append_with_context_best_effort( + &attempt.context, + RawTraceEventPayload::InferenceStarted { + inference_call_id: attempt.inference_call_id.clone(), + thread_id: attempt.context.thread_id.clone(), + codex_turn_id: attempt.context.codex_turn_id.clone(), + model: attempt.context.model.clone(), + provider_name: attempt.context.provider_name.clone(), + request_payload, + }, + ); + } + + /// Records successful provider completion and serializes the observed output items. + /// + /// Callers pass protocol-native response items so this crate owns the + /// trace-specific serialization rules. That keeps codex-core focused on + /// transport behavior while preserving trace evidence that normal request + /// serialization intentionally omits. + pub fn record_completed( + &self, + response_id: &str, + upstream_request_id: Option<&str>, + token_usage: &Option, + output_items: &[ResponseItem], + ) { + let Some(attempt) = self.take_terminal_attempt() else { + return; + }; + let Some(response_payload) = write_response_payload_best_effort( + attempt, + Some(response_id), + upstream_request_id, + token_usage.as_ref(), + output_items, + ) else { + return; + }; + + append_with_context_best_effort( + &attempt.context, + RawTraceEventPayload::InferenceCompleted { + inference_call_id: attempt.inference_call_id.clone(), + response_id: Some(response_id.to_string()), + upstream_request_id: upstream_request_id.map(str::to_string), + response_payload, + }, + ); + } + + /// Records pre-response and mid-stream failures. + pub fn record_failed( + &self, + error: impl Display, + upstream_request_id: Option<&str>, + output_items: &[ResponseItem], + ) { + let Some(attempt) = self.take_terminal_attempt() else { + return; + }; + let partial_response_payload = if output_items.is_empty() { + None + } else { + write_response_payload_best_effort( + attempt, + /*response_id*/ None, + upstream_request_id, + /*token_usage*/ None, + output_items, + ) + }; + append_with_context_best_effort( + &attempt.context, + RawTraceEventPayload::InferenceFailed { + inference_call_id: attempt.inference_call_id.clone(), + upstream_request_id: upstream_request_id.map(str::to_string), + error: error.to_string(), + partial_response_payload, + }, + ); + } + + /// Records a provider stream that Codex intentionally stopped consuming. + /// + /// This happens when the turn is interrupted or when mailbox delivery + /// preempts the current sampling request. Complete output items observed + /// before that point are retained as partial response evidence. + pub fn record_cancelled( + &self, + reason: impl Display, + upstream_request_id: Option<&str>, + output_items: &[ResponseItem], + ) { + let Some(attempt) = self.take_terminal_attempt() else { + return; + }; + let partial_response_payload = if output_items.is_empty() { + None + } else { + write_response_payload_best_effort( + attempt, + /*response_id*/ None, + upstream_request_id, + /*token_usage*/ None, + output_items, + ) + }; + append_with_context_best_effort( + &attempt.context, + RawTraceEventPayload::InferenceCancelled { + inference_call_id: attempt.inference_call_id.clone(), + upstream_request_id: upstream_request_id.map(str::to_string), + reason: reason.to_string(), + partial_response_payload, + }, + ); + } + + fn take_terminal_attempt(&self) -> Option<&EnabledInferenceTraceAttempt> { + let attempt = match &self.state { + InferenceTraceAttemptState::Disabled => return None, + InferenceTraceAttemptState::Enabled(attempt) => attempt, + }; + if attempt.terminal_recorded.swap(true, Ordering::AcqRel) { + return None; + } + Some(attempt) + } +} + +/// Serializes a response item for trace evidence rather than future request construction. +/// +/// The protocol serializer intentionally omits some readable reasoning content +/// when shaping items for later model requests. Rollout traces need the item as +/// Codex received it, so this helper restores that content in the raw payload. +pub(crate) fn trace_response_item_json(item: &ResponseItem) -> JsonValue { + let mut value = serde_json::to_value(item).unwrap_or_else(|err| { + serde_json::json!({ + "serialization_error": err.to_string(), + }) + }); + + if let ResponseItem::Reasoning { + content: Some(content), + .. + } = item + && let JsonValue::Object(object) = &mut value + { + object.insert( + "content".to_string(), + serde_json::to_value(content).unwrap_or_else(|err| { + serde_json::json!({ + "serialization_error": err.to_string(), + }) + }), + ); + } + + value +} + +fn next_inference_call_id() -> InferenceCallId { + Uuid::new_v4().to_string() +} + +fn write_json_payload_best_effort( + writer: &TraceWriter, + kind: RawPayloadKind, + payload: &impl Serialize, +) -> Option { + writer.write_json_payload(kind, payload).ok() +} + +fn write_response_payload_best_effort( + attempt: &EnabledInferenceTraceAttempt, + response_id: Option<&str>, + upstream_request_id: Option<&str>, + token_usage: Option<&TokenUsage>, + output_items: &[ResponseItem], +) -> Option { + let response_payload = TracedResponseStreamOutput { + response_id, + upstream_request_id, + token_usage, + output_items: output_items.iter().map(trace_response_item_json).collect(), + }; + write_json_payload_best_effort( + &attempt.context.writer, + RawPayloadKind::InferenceResponse, + &response_payload, + ) +} + +fn append_with_context_best_effort( + context: &EnabledInferenceTraceContext, + payload: RawTraceEventPayload, +) { + let event_context = RawTraceEventContext { + thread_id: Some(context.thread_id.clone()), + codex_turn_id: Some(context.codex_turn_id.clone()), + }; + let _ = context.writer.append_with_context(event_context, payload); +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use codex_protocol::ResponseItemId; + use codex_protocol::models::ReasoningItemContent; + use codex_protocol::models::ReasoningItemReasoningSummary; + use pretty_assertions::assert_eq; + use serde_json::json; + use tempfile::TempDir; + + use super::*; + use crate::model::ExecutionStatus; + use crate::replay_bundle; + + #[test] + fn disabled_attempt_adds_no_request_headers() { + let mut headers = HeaderMap::new(); + + InferenceTraceAttempt::disabled().add_request_headers(&mut headers); + + assert!(headers.is_empty()); + } + + #[test] + fn enabled_attempt_adds_inference_request_header() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = Arc::new(TraceWriter::create( + temp.path(), + "trace-1".to_string(), + "rollout-1".to_string(), + "thread-root".to_string(), + )?); + let context = InferenceTraceContext::enabled( + writer, + "thread-root".to_string(), + "turn-1".to_string(), + "gpt-test".to_string(), + "test-provider".to_string(), + ); + let attempt = context.start_attempt(); + let mut headers = HeaderMap::new(); + + attempt.add_request_headers(&mut headers); + + let header = headers + .get(INFERENCE_CALL_ID_HEADER) + .expect("inference header present"); + assert_eq!(Some(header.to_str()?), attempt.inference_call_id()); + assert!(Uuid::parse_str(header.to_str()?).is_ok()); + Ok(()) + } + + #[test] + fn enabled_context_records_replayable_inference_attempt() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = Arc::new(TraceWriter::create( + temp.path(), + "trace-1".to_string(), + "rollout-1".to_string(), + "thread-root".to_string(), + )?); + writer.append(RawTraceEventPayload::ThreadStarted { + thread_id: "thread-root".to_string(), + agent_path: "/root".to_string(), + metadata_payload: None, + })?; + writer.append(RawTraceEventPayload::CodexTurnStarted { + codex_turn_id: "turn-1".to_string(), + thread_id: "thread-root".to_string(), + })?; + let context = InferenceTraceContext::enabled( + writer, + "thread-root".to_string(), + "turn-1".to_string(), + "gpt-test".to_string(), + "test-provider".to_string(), + ); + + let attempt = context.start_attempt(); + attempt.record_started(&json!({ + "model": "gpt-test", + "input": [{ + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "hello"}] + }], + })); + attempt.record_completed("resp-1", Some("req-1"), &None, &[]); + + let rollout = replay_bundle(temp.path())?; + let inference = rollout + .inference_calls + .values() + .next() + .expect("recorded inference call"); + + assert_eq!(rollout.inference_calls.len(), 1); + assert_eq!(inference.thread_id, "thread-root"); + assert_eq!(inference.codex_turn_id, "turn-1"); + assert_eq!(inference.execution.status, ExecutionStatus::Completed); + assert_eq!(inference.upstream_request_id, Some("req-1".to_string())); + assert_eq!(rollout.raw_payloads.len(), 2); + + Ok(()) + } + + #[test] + fn traced_response_item_preserves_reasoning_content_omitted_by_normal_serializer() { + let item = ResponseItem::Reasoning { + id: Some(ResponseItemId::with_suffix("rs", "1")), + summary: vec![ReasoningItemReasoningSummary::SummaryText { + text: "summary".to_string(), + }], + content: Some(vec![ReasoningItemContent::Text { + text: "raw reasoning".to_string(), + }]), + encrypted_content: Some("encoded".to_string()), + internal_chat_message_metadata_passthrough: None, + }; + + let normal = serde_json::to_value(&item).expect("response item serializes"); + let traced = trace_response_item_json(&item); + + assert_eq!(normal.get("content"), None); + assert_eq!( + traced, + json!({ + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "summary"}], + "content": [{"type": "text", "text": "raw reasoning"}], + "encrypted_content": "encoded", + }), + ); + } +} diff --git a/codex-rs/rollout-trace/src/lib.rs b/codex-rs/rollout-trace/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..54094f0525b37f3f4759f64bf3903741cf331f67 --- /dev/null +++ b/codex-rs/rollout-trace/src/lib.rs @@ -0,0 +1,78 @@ +//! Trace bundle format, writer, and reducer for Codex rollouts. +//! +//! This crate owns the trace schema. Hot-path Codex code should depend on the +//! small writer API here; semantic replay and viewer projections stay outside +//! `codex-core`. +//! +//! See `README.md` for the system diagram and reducer model. + +mod bundle; +mod code_cell; +mod compaction; +mod inference; +mod mcp; +mod model; +mod payload; +mod protocol_event; +mod raw_event; +mod reducer; +mod thread; +mod tool_dispatch; +mod writer; + +/// Conventional reduced-state cache name written next to a raw trace bundle. +pub use bundle::REDUCED_STATE_FILE_NAME; +/// No-op-capable handle for recording one code-mode runtime cell. +pub use code_cell::CodeCellTraceContext; +/// Raw checkpoint payload for a remote compaction install event. +pub use compaction::CompactionCheckpointTracePayload; +/// No-op-capable handle for recording remote-compaction requests. +pub use compaction::CompactionTraceAttempt; +/// Shared recorder context for a compaction checkpoint. +pub use compaction::CompactionTraceContext; +/// No-op-capable handle for recording one upstream inference attempt. +pub use inference::InferenceTraceAttempt; +/// Shared recorder context for inference attempts within one Codex turn. +pub use inference::InferenceTraceContext; +/// Trace-owned MCP execution correlation propagated to bridge request metadata. +pub use mcp::McpCallTraceContext; +/// Public reduced trace model returned by replay. +pub use model::*; +/// Stable identifier for one raw payload inside a rollout bundle. +pub use payload::RawPayloadId; +/// Coarse role labels for raw payload files. +pub use payload::RawPayloadKind; +/// Reference to a raw request/response/log payload stored in the bundle. +pub use payload::RawPayloadRef; +/// Monotonic sequence number assigned by the raw trace writer. +pub use raw_event::RawEventSeq; +/// Runtime requester observed before semantic reduction. +pub use raw_event::RawToolCallRequester; +/// One append-only raw trace event from `trace.jsonl`. +pub use raw_event::RawTraceEvent; +/// Event-envelope context supplied by hot-path trace producers. +pub use raw_event::RawTraceEventContext; +/// Typed payload for one raw trace event. +pub use raw_event::RawTraceEventPayload; +/// Replay a raw trace bundle and write/read its reduced `RolloutTrace`. +pub use reducer::replay_bundle; +/// Raw payload captured when a child agent reports completion to its parent. +pub use thread::AgentResultTracePayload; +/// Environment variable that enables local trace-bundle recording. +pub use thread::CODEX_ROLLOUT_TRACE_ROOT_ENV; +/// Raw metadata captured when a thread starts. +pub use thread::ThreadStartedTraceMetadata; +/// No-op-capable handle for recording one thread in a rollout bundle. +pub use thread::ThreadTraceContext; +/// Request data for the canonical Codex tool boundary. +pub use tool_dispatch::ToolDispatchInvocation; +/// Tool input observed at the registry boundary. +pub use tool_dispatch::ToolDispatchPayload; +/// Runtime source that caused a dispatch-level tool call. +pub use tool_dispatch::ToolDispatchRequester; +/// Result data returned from a dispatch-level tool call. +pub use tool_dispatch::ToolDispatchResult; +/// No-op-capable handle for recording one resolved tool dispatch. +pub use tool_dispatch::ToolDispatchTraceContext; +/// Append-only writer used by hot-path Codex instrumentation. +pub use writer::TraceWriter; diff --git a/codex-rs/rollout-trace/src/mcp.rs b/codex-rs/rollout-trace/src/mcp.rs new file mode 100644 index 0000000000000000000000000000000000000000..b8564e1579bf415f998609a8f8eb4feb83ad7082 --- /dev/null +++ b/codex-rs/rollout-trace/src/mcp.rs @@ -0,0 +1,99 @@ +//! Hot-path helpers for correlating concrete MCP executions with rollout traces. +//! +//! Core decides when an MCP request is actually going to execute. The trace +//! crate owns the globally unique ID, the trace event that preserves it in the +//! reduced artifact, and the bridge-private MCP request metadata key. + +use crate::McpCallId; +use serde_json::Value as JsonValue; + +const MCP_CALL_ID_META_KEY: &str = "codex_bridge_mcp_call_id"; + +/// No-op capable handle for one concrete MCP backend call. +#[derive(Clone, Debug)] +pub struct McpCallTraceContext { + mcp_call_id: Option, +} + +impl McpCallTraceContext { + /// Builds a context that records nothing and leaves request metadata unchanged. + pub fn disabled() -> Self { + Self { mcp_call_id: None } + } + + /// Builds the trace handle for one concrete MCP execution. + pub(crate) fn enabled(mcp_call_id: McpCallId) -> Self { + Self { + mcp_call_id: Some(mcp_call_id), + } + } + + /// Returns the trace-owned MCP call ID when rollout tracing is enabled. + pub(crate) fn mcp_call_id(&self) -> Option<&str> { + self.mcp_call_id.as_deref() + } + + /// Adds bridge-private MCP correlation metadata to one outgoing request. + pub fn add_request_meta(&self, meta: Option) -> Option { + let Some(mcp_call_id) = self.mcp_call_id() else { + return meta; + }; + + match meta { + Some(JsonValue::Object(mut map)) => { + map.insert( + MCP_CALL_ID_META_KEY.to_string(), + JsonValue::String(mcp_call_id.to_string()), + ); + Some(JsonValue::Object(map)) + } + None => { + let mut map = serde_json::Map::new(); + map.insert( + MCP_CALL_ID_META_KEY.to_string(), + JsonValue::String(mcp_call_id.to_string()), + ); + Some(JsonValue::Object(map)) + } + // This should never happen but if it does then we'll fallback to + // a noop rather than any breaking behavior. The tracing is best + // effort after all. + Some(_) => meta, + } + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::MCP_CALL_ID_META_KEY; + use super::McpCallTraceContext; + + #[test] + fn disabled_mcp_trace_leaves_request_meta_unchanged() { + let meta = Some(json!({"source": "test"})); + + assert_eq!( + McpCallTraceContext::disabled().add_request_meta(meta.clone()), + meta + ); + } + + #[test] + fn enabled_mcp_trace_adds_bridge_correlation_meta() { + let trace = McpCallTraceContext::enabled("mcp-call-id".to_string()); + let meta = trace + .add_request_meta(Some(json!({"source": "test"}))) + .expect("enabled trace keeps request metadata"); + let object = meta + .as_object() + .expect("MCP request metadata remains an object"); + + assert_eq!(object["source"], json!("test")); + assert_eq!( + object[MCP_CALL_ID_META_KEY], + json!(trace.mcp_call_id().expect("enabled trace has an ID")) + ); + } +} diff --git a/codex-rs/rollout-trace/src/model/conversation.rs b/codex-rs/rollout-trace/src/model/conversation.rs new file mode 100644 index 0000000000000000000000000000000000000000..0cb72e85ef25a1e75dfb418fa1e2390058816f83 --- /dev/null +++ b/codex-rs/rollout-trace/src/model/conversation.rs @@ -0,0 +1,193 @@ +use serde::Deserialize; +use serde::Serialize; + +use crate::payload::RawPayloadId; + +use super::AgentPath; +use super::AgentThreadId; +use super::CodeCellId; +use super::CodexTurnId; +use super::CompactionId; +use super::ConversationItemId; +use super::EdgeId; +use super::InferenceCallId; +use super::ModelVisibleCallId; +use super::ToolCallId; +use super::session::ExecutionWindow; + +/// One logical transcript item or transcript boundary. +/// +/// The reducer builds conversation items primarily from inference request and +/// response payloads. Runtime objects can be listed in `produced_by`, but they +/// must not rewrite what the item body says the model saw. Structural items, +/// such as compaction markers, live in the same ordered list so conversation +/// views can show where the live history changed. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct ConversationItem { + pub item_id: ConversationItemId, + pub thread_id: AgentThreadId, + /// Runtime activation that first introduced this item locally, when known. + pub codex_turn_id: Option, + pub first_seen_at_unix_ms: i64, + pub role: ConversationRole, + /// Codex channel for assistant/tool content, when the item is channel-specific. + pub channel: Option, + pub kind: ConversationItemKind, + /// Routing metadata carried by a Responses `agent_message` item. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub agent_message: Option, + pub body: ConversationBody, + /// Protocol/model `call_id` for function/custom tool call and output items. + pub call_id: Option, + /// Runtime or control-plane objects that caused this conversation item to exist. + pub produced_by: Vec, +} + +/// Sender and destination identities attached to a model-visible agent message. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct AgentMessageMetadata { + /// Agent path that authored the message. + pub author: AgentPath, + /// Agent path that received the message. + pub recipient: AgentPath, +} + +/// Model-visible role assigned to a conversation item. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ConversationRole { + System, + Developer, + User, + Assistant, + Tool, +} + +/// Codex channel for model-visible content. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ConversationChannel { + Analysis, + Commentary, + Final, + /// Remote compaction summaries are reintroduced as assistant summary-channel content. + Summary, +} + +/// Responses item category after normalization into the reduced transcript. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ConversationItemKind { + Message, + Reasoning, + FunctionCall, + FunctionCallOutput, + CustomToolCall, + CustomToolCallOutput, + /// Structural marker inserted where live history was replaced by compaction. + CompactionMarker, +} + +/// Ordered content parts for a reduced conversation item. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct ConversationBody { + /// Renderable model-visible parts. Raw payload refs are used when the bytes + /// are too large or too structured for the normal conversation path. + pub parts: Vec, +} + +/// One model-visible part inside a conversation item. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum ConversationPart { + Text { + text: String, + }, + /// A model-provided summary of content whose full form may also be present. + /// + /// Reasoning summaries are not interchangeable with raw reasoning text: + /// both can be present in one payload, and replay/debug tooling needs to + /// preserve which representation the model actually returned. + Summary { + text: String, + }, + /// Opaque model-visible content that is intentionally not decoded here. + /// + /// Reasoning can be carried as `encrypted_content` with no readable text. + /// Keeping that blob inline makes it part of item identity, unlike a raw + /// payload reference whose ID changes every time the same item is replayed + /// in a later inference request. + Encoded { + label: String, + value: String, + }, + /// Small JSON-ish body represented by a summary plus a raw ref. + Json { + summary: String, + raw_payload_id: RawPayloadId, + }, + Code { + language: String, + source: String, + }, + /// Large or uncommon payload that should be lazy-loaded from details UI. + PayloadRef { + label: String, + raw_payload_id: RawPayloadId, + }, +} + +/// Explanation for where a conversation item came from. +/// +/// This is deliberately plural at the call site: a function output can be both +/// model-visible conversation and the product of a runtime tool call. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum ProducerRef { + UserInput, + Inference { inference_call_id: InferenceCallId }, + Tool { tool_call_id: ToolCallId }, + CodeCell { code_cell_id: CodeCellId }, + InteractionEdge { edge_id: EdgeId }, + Compaction { compaction_id: CompactionId }, + Harness, +} + +/// One outbound inference request and its response metadata. +/// +/// Full upstream request/response bodies live behind raw payload refs. The +/// request/response item ID lists are the reduced, model-visible snapshot. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct InferenceCall { + pub inference_call_id: InferenceCallId, + pub thread_id: AgentThreadId, + pub codex_turn_id: CodexTurnId, + pub execution: ExecutionWindow, + pub model: String, + pub provider_name: String, + /// Responses API response id, used by follow-up `previous_response_id` requests. + pub response_id: Option, + /// Request id returned by HTTP/proxy/engine infrastructure. + pub upstream_request_id: Option, + /// Complete ordered input snapshot sent with this request. + pub request_item_ids: Vec, + /// Ordered output items produced by this response. + pub response_item_ids: Vec, + /// Runtime tool calls whose model-visible call item came from this response. + pub tool_call_ids_started_by_response: Vec, + pub usage: Option, + pub raw_request_payload_id: RawPayloadId, + /// Full upstream response payload. `None` while running or after pre-stream failures. + pub raw_response_payload_id: Option, +} + +/// Token usage summary for one inference call. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TokenUsage { + pub input_tokens: u64, + pub cached_input_tokens: u64, + #[serde(default)] + pub cache_write_input_tokens: u64, + pub output_tokens: u64, + pub reasoning_output_tokens: u64, +} diff --git a/codex-rs/rollout-trace/src/model/mod.rs b/codex-rs/rollout-trace/src/model/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..0166121da9f4582eaeaaf27961ddc7a01f869656 --- /dev/null +++ b/codex-rs/rollout-trace/src/model/mod.rs @@ -0,0 +1,123 @@ +//! Reduced rollout trace model. +//! +//! These types describe the deterministic replay output. They intentionally +//! separate model-visible conversation from runtime/debug objects. + +use std::collections::BTreeMap; + +use serde::Deserialize; +use serde::Serialize; + +use crate::payload::RawPayloadId; +use crate::payload::RawPayloadRef; +mod conversation; +mod runtime; +mod session; + +pub use conversation::*; +pub use runtime::*; +pub use session::*; + +/// Codex conversation/session UUID. +pub type AgentThreadId = String; +/// Stable multi-agent routing path such as `/root` or `/root/search_docs`. +pub type AgentPath = String; +/// Runtime submission/activation UUID. This is not a chat turn. +pub type CodexTurnId = String; +/// Reduced transcript item ID assigned by the trace reducer. +pub type ConversationItemId = String; +/// Local ID for one outbound upstream inference request. +pub type InferenceCallId = String; +/// Globally unique ID for one concrete MCP backend request. +pub type McpCallId = String; +/// Reducer-owned ID for one runtime tool-call object. +pub type ToolCallId = String; +/// Responses `call_id` / custom-tool call ID visible in inference payloads. +pub type ModelVisibleCallId = String; +/// Tool invocation ID assigned inside the code-mode JavaScript runtime. +pub type CodeModeRuntimeToolId = String; +/// Reducer-owned ID for one model-authored `exec` JavaScript cell. +pub type CodeCellId = String; +/// Process/session ID returned by Codex's terminal runtime. +pub type TerminalId = String; +/// Reducer-owned ID for one command/write/poll operation against a terminal. +pub type TerminalOperationId = String; +/// Reducer-owned ID for one installed conversation-history checkpoint. +pub type CompactionId = String; +/// Reducer-owned ID for one upstream request that computes a compaction. +pub type CompactionRequestId = String; +/// Reducer-owned ID for one information-flow edge. +pub type EdgeId = String; +/// Reducer-owned ID for request/log correlation metadata. +pub type CorrelationId = String; + +/// Canonical reduced graph for one Codex rollout. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct RolloutTrace { + pub schema_version: u32, + /// Unique identity for this trace capture. + /// + /// `rollout_id` names the Codex rollout/session being observed. `trace_id` + /// names the diagnostic artifact produced for that rollout, which keeps + /// storage/replay identity separate from the product-level session identity. + pub trace_id: String, + /// CLI-visible rollout/run identity. Higher-level experiment/sample IDs wrap this object. + pub rollout_id: String, + pub started_at_unix_ms: i64, + /// Wall-clock timestamp for terminal rollout status. `None` means running or partial trace. + pub ended_at_unix_ms: Option, + pub status: RolloutStatus, + pub root_thread_id: AgentThreadId, + pub threads: BTreeMap, + pub codex_turns: BTreeMap, + pub conversation_items: BTreeMap, + pub inference_calls: BTreeMap, + /// Model-authored `exec` JavaScript cells keyed by reducer-owned cell ID. + pub code_cells: BTreeMap, + pub tool_calls: BTreeMap, + /// Terminal runtime sessions keyed by process/session ID returned by the runtime. + pub terminal_sessions: BTreeMap, + /// Commands/writes/polls against terminals keyed by reducer-owned operation ID. + pub terminal_operations: BTreeMap, + /// Installed compaction checkpoints keyed by checkpoint ID. + pub compactions: BTreeMap, + /// Upstream remote compaction calls keyed by local request ID. + pub compaction_requests: BTreeMap, + /// Information-flow edges between threads, cells, tools, and runtime resources. + pub interaction_edges: BTreeMap, + /// Raw JSON payloads keyed by raw-payload ID. Most point at files outside this object. + pub raw_payloads: BTreeMap, +} + +impl RolloutTrace { + /// Builds an empty reduced trace that a reducer can populate. + pub(crate) fn new( + schema_version: u32, + trace_id: String, + rollout_id: String, + root_thread_id: AgentThreadId, + started_at_unix_ms: i64, + ) -> Self { + Self { + schema_version, + trace_id, + rollout_id, + started_at_unix_ms, + ended_at_unix_ms: None, + status: RolloutStatus::Running, + root_thread_id, + threads: BTreeMap::new(), + codex_turns: BTreeMap::new(), + conversation_items: BTreeMap::new(), + inference_calls: BTreeMap::new(), + code_cells: BTreeMap::new(), + tool_calls: BTreeMap::new(), + terminal_sessions: BTreeMap::new(), + terminal_operations: BTreeMap::new(), + compactions: BTreeMap::new(), + compaction_requests: BTreeMap::new(), + interaction_edges: BTreeMap::new(), + raw_payloads: BTreeMap::new(), + } + } +} diff --git a/codex-rs/rollout-trace/src/model/runtime.rs b/codex-rs/rollout-trace/src/model/runtime.rs new file mode 100644 index 0000000000000000000000000000000000000000..f0d1d1caf1d1e81dbb519452b906ab76b017cb93 --- /dev/null +++ b/codex-rs/rollout-trace/src/model/runtime.rs @@ -0,0 +1,334 @@ +use serde::Deserialize; +use serde::Serialize; + +use crate::payload::RawPayloadId; +use crate::raw_event::RawEventSeq; + +use super::AgentPath; +use super::AgentThreadId; +use super::CodeCellId; +use super::CodeModeRuntimeToolId; +use super::CodexTurnId; +use super::CompactionId; +use super::CompactionRequestId; +use super::ConversationItemId; +use super::EdgeId; +use super::McpCallId; +use super::ModelVisibleCallId; +use super::TerminalId; +use super::TerminalOperationId; +use super::ToolCallId; +use super::session::ExecutionWindow; + +/// Runtime/debug object for one model-authored `exec` cell. +/// +/// The JavaScript source and custom-tool outputs are still conversation items; +/// this object tracks the code-mode runtime boundary and nested runtime work. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct CodeCell { + /// Reducer-owned graph id derived from the model-visible `exec` call id. + /// Runtime cell ids are stored separately because they are only handles for + /// later waits and nested code-mode tools. + pub code_cell_id: CodeCellId, + pub model_visible_call_id: ModelVisibleCallId, + pub thread_id: AgentThreadId, + pub codex_turn_id: CodexTurnId, + /// Conversation item containing the model-authored JavaScript. + pub source_item_id: ConversationItemId, + pub output_item_ids: Vec, + /// Raw code-mode runtime/session id, useful when matching runtime payloads. + pub runtime_cell_id: Option, + /// Full JS-cell runtime window; yielded cells can outlive the initial custom call. + pub execution: ExecutionWindow, + pub runtime_status: CodeCellRuntimeStatus, + pub initial_response_at_unix_ms: Option, + pub initial_response_seq: Option, + pub yielded_at_unix_ms: Option, + pub yielded_seq: Option, + pub source_js: String, + pub nested_tool_call_ids: Vec, + pub wait_tool_call_ids: Vec, +} + +/// Code-mode runtime lifecycle. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum CodeCellRuntimeStatus { + /// The `exec` request has been accepted but the runtime has not yet started user code. + Starting, + /// Runtime is executing JavaScript and has not yet yielded or terminated. + Running, + /// Initial `exec` returned while JavaScript kept running in the background. + Yielded, + /// Runtime reached a normal terminal result. + Completed, + /// Runtime reached an error terminal result. + Failed, + /// Runtime was explicitly terminated. + Terminated, +} + +/// Installed conversation-history replacement boundary. +/// +/// Duration-bearing upstream requests live in `CompactionRequest`. This object +/// is the checkpoint where replacement history became the live thread history. +/// The boundary marker and the model-visible summary are separate conversation +/// items: the marker says where history was replaced, while the summary is part +/// of `replacement_item_ids` when the compact endpoint returned one. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct Compaction { + pub compaction_id: CompactionId, + pub thread_id: AgentThreadId, + pub codex_turn_id: CodexTurnId, + pub installed_at_unix_ms: i64, + /// Structural conversation item marking where pre-compaction history ended. + pub marker_item_id: ConversationItemId, + /// Upstream compaction request attempts that contributed to this checkpoint. + pub request_ids: Vec, + /// Logical conversation items present immediately before replacement. + pub input_item_ids: Vec, + /// Replacement conversation items installed by the checkpoint. + pub replacement_item_ids: Vec, +} + +/// One upstream remote request made while computing a compaction checkpoint. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct CompactionRequest { + pub compaction_request_id: CompactionRequestId, + pub compaction_id: CompactionId, + pub thread_id: AgentThreadId, + pub codex_turn_id: CodexTurnId, + pub execution: ExecutionWindow, + pub model: String, + pub provider_name: String, + pub raw_request_payload_id: RawPayloadId, + /// Full compaction response payload. `None` while running or after pre-response failures. + pub raw_response_payload_id: Option, +} + +/// Runtime operation requested by the model, a JS code cell, or Codex itself. +/// +/// A `ToolCall` is not a chat transcript row. Model-visible call/output items +/// link to it through `model_visible_*_item_ids`; runtime-only tools can have +/// empty model-visible lists. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct ToolCall { + pub tool_call_id: ToolCallId, + /// Globally unique MCP execution ID, when this tool reached an MCP backend. + pub mcp_call_id: Option, + /// Model-visible protocol call ID, if the model directly requested this tool. + pub model_visible_call_id: Option, + /// Code-mode runtime's internal tool invocation ID, if this call came from JS. + pub code_mode_runtime_tool_id: Option, + pub thread_id: AgentThreadId, + /// Runtime activation that started the tool. Background work may outlive this turn. + pub started_by_codex_turn_id: Option, + pub execution: ExecutionWindow, + pub requester: ToolCallRequester, + pub kind: ToolCallKind, + pub model_visible_call_item_ids: Vec, + pub model_visible_output_item_ids: Vec, + /// Terminal operation started by this tool, when the tool touched a terminal. + pub terminal_operation_id: Option, + pub summary: ToolCallSummary, + /// Original invocation at the Codex tool boundary. + /// + /// Direct model tools store the model's function/custom call payload here. + /// Code-mode nested tools store the JSON call made by model-authored JS. + /// Runtime protocol events are deliberately kept separate below because + /// they describe how Codex executed the request, not what the caller sent. + pub raw_invocation_payload_id: Option, + /// Result returned to the immediate requester. + /// + /// For direct tools this is the tool output item returned to the model; for + /// code-mode nested tools this is the value returned to JavaScript. + pub raw_result_payload_id: Option, + /// Runtime/protocol payloads observed while executing the tool. + /// + /// Examples include exec begin/end, patch begin/end, and MCP begin/end + /// events. Reducers can use these to build richer runtime objects such as + /// terminal operations without overwriting the canonical invocation/result. + pub raw_runtime_payload_ids: Vec, +} + +/// Requester of a runtime tool. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum ToolCallRequester { + Model, + /// Model-authored JavaScript requested the tool through code-mode. + CodeCell { + code_cell_id: CodeCellId, + }, +} + +/// Runtime tool category. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum ToolCallKind { + ExecCommand, + WriteStdin, + ApplyPatch, + Mcp { + server: String, + tool: String, + }, + Web, + ImageGeneration, + SpawnAgent, + AssignAgentTask, + SendMessage, + /// Multi-agent wait operation. Code-mode wait is modeled separately. + WaitAgent, + CloseAgent, + Other { + name: String, + }, +} + +/// Bounded card/list summary for a tool call. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum ToolCallSummary { + /// Tool is summarized by its terminal operation. + Terminal { operation_id: TerminalOperationId }, + Agent { + target_agent_path: AgentPath, + /// Task name/path segment when the operation creates or targets a task. + task_name: Option, + message_preview: String, + }, + WaitAgent { + /// Wait target, when narrower than "any child". + target_agent_path: Option, + timeout_ms: Option, + }, + Generic { + label: String, + input_preview: Option, + output_preview: Option, + }, +} + +/// Reusable terminal process/session returned by the runtime. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct TerminalSession { + pub terminal_id: TerminalId, + pub thread_id: AgentThreadId, + pub created_by_operation_id: TerminalOperationId, + pub operation_ids: Vec, + /// Terminal lifetime. This can outlive the operation that created it. + pub execution: ExecutionWindow, +} + +/// One command/write/poll operation against a terminal session. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct TerminalOperation { + pub operation_id: TerminalOperationId, + /// Runtime terminal/process ID. `None` is legal only while the operation that creates it is starting. + pub terminal_id: Option, + pub tool_call_id: ToolCallId, + pub kind: TerminalOperationKind, + /// Operation execution window. This is not necessarily the terminal session lifetime. + pub execution: ExecutionWindow, + pub request: TerminalRequest, + /// Runtime-observed terminal result. Model-visible output links through observations. + pub result: Option, + pub model_observations: Vec, + pub raw_payload_ids: Vec, +} + +/// Terminal operation category. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum TerminalOperationKind { + ExecCommand, + WriteStdin, +} + +/// Terminal request summary. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum TerminalRequest { + ExecCommand { + command: Vec, + display_command: String, + cwd: String, + yield_time_ms: Option, + max_output_tokens: Option, + }, + /// Request to interact with an existing terminal. + WriteStdin { + /// Bytes/text sent to stdin. Empty string means poll/read without writing bytes. + stdin: String, + yield_time_ms: Option, + max_output_tokens: Option, + }, +} + +/// Terminal result observed by the runtime. +/// +/// This is debugger/runtime output. It is not proof that the model saw the same +/// bytes; link model-visible call/output items through `TerminalModelObservation`. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TerminalResult { + /// Process exit code. `None` if the process is still running or no exit status was produced. + pub exit_code: Option, + pub stdout: String, + pub stderr: String, + /// Tool runtime's formatted caller-facing output, when present. + pub formatted_output: Option, + /// Token count before truncation, when the tool runtime reported it. + pub original_token_count: Option, + /// Streaming chunk ID, when this result was assembled from chunked terminal output. + pub chunk_id: Option, +} + +/// Conversation items that observed a terminal operation. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct TerminalModelObservation { + pub call_item_ids: Vec, + pub output_item_ids: Vec, + pub source: TerminalObservationSource, +} + +/// Source of model-visible terminal observation. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum TerminalObservationSource { + DirectToolCall, + CodeCellOutput, +} + +/// Directed information-flow relationship between trace objects. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct InteractionEdge { + pub edge_id: EdgeId, + pub kind: InteractionEdgeKind, + pub source: TraceAnchor, + pub target: TraceAnchor, + pub started_at_unix_ms: i64, + pub ended_at_unix_ms: Option, + pub carried_item_ids: Vec, + pub carried_raw_payload_ids: Vec, +} + +/// Information-flow edge category. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum InteractionEdgeKind { + SpawnAgent, + AssignAgentTask, + SendMessage, + AgentResult, + CloseAgent, +} + +/// Typed pointer to one stable reduced-trace object. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum TraceAnchor { + ConversationItem { item_id: ConversationItemId }, + ToolCall { tool_call_id: ToolCallId }, + Thread { thread_id: AgentThreadId }, +} diff --git a/codex-rs/rollout-trace/src/model/session.rs b/codex-rs/rollout-trace/src/model/session.rs new file mode 100644 index 0000000000000000000000000000000000000000..fccc386dc9328eb782cb0038192573b3f72362fe --- /dev/null +++ b/codex-rs/rollout-trace/src/model/session.rs @@ -0,0 +1,110 @@ +use serde::Deserialize; +use serde::Serialize; + +use crate::raw_event::RawEventSeq; + +use super::AgentPath; +use super::AgentThreadId; +use super::CodexTurnId; +use super::ConversationItemId; +use super::EdgeId; + +/// Coarse terminal status for the rollout. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum RolloutStatus { + /// Writer has not seen a terminal rollout event. + Running, + /// Rollout ended normally. + Completed, + /// Rollout ended because an operation failed. + Failed, + /// Rollout was cancelled or otherwise stopped before normal completion. + Aborted, +} + +/// One Codex thread/session participating in the rollout. +/// +/// Threads are agents in the multi-agent sense, but the root interactive +/// session is represented by the same object. Runtime objects live in top-level +/// maps and point back to their owning thread; only transcript order is stored +/// here because compaction/reconciliation makes it semantic. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct AgentThread { + pub thread_id: AgentThreadId, + /// Stable routing identity. Viewer/search should prefer this over nickname. + pub agent_path: AgentPath, + /// Presentation hint. It can collide and must not be used as identity. + pub nickname: Option, + pub origin: AgentOrigin, + /// Session lifecycle for this thread. + /// + /// Child threads can end independently from the root rollout, for example + /// after a parent calls `close_agent`. Keeping this on the thread prevents + /// those shutdowns from being mistaken for whole-rollout completion. + pub execution: ExecutionWindow, + /// Configured model presentation hint. Individual inference calls carry the actual upstream model. + pub default_model: Option, + /// Logical conversation items first observed for this thread, in transcript order. + pub conversation_item_ids: Vec, +} + +/// Provenance for a traced Codex thread. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum AgentOrigin { + Root, + Spawned { + parent_thread_id: AgentThreadId, + /// Interaction edge that carried the spawn task. + spawn_edge_id: EdgeId, + /// Stable path segment/task name selected by the parent/tool call. + task_name: String, + /// Selected agent role/type, for example `worker` or `explorer`. + agent_role: String, + }, +} + +/// Runtime interval for a typed trace object. +/// +/// Wall-clock timestamps are for display and latency. Sequence numbers are the +/// causal ordering primitive and should be used to pair observations or break +/// same-millisecond ties. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ExecutionWindow { + pub started_at_unix_ms: i64, + pub started_seq: RawEventSeq, + pub ended_at_unix_ms: Option, + pub ended_seq: Option, + pub status: ExecutionStatus, +} + +/// Coarse lifecycle status for a runtime object. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ExecutionStatus { + /// Object is still live or the trace ended before its terminal event. + Running, + /// Object completed successfully. + Completed, + /// Object reached an error state. + Failed, + /// Object was cancelled by user/policy/runtime before completion. + Cancelled, + /// Object was aborted when its owner/runtime stopped. + Aborted, +} + +/// One activation of the Codex runtime for one thread. +/// +/// A Codex turn groups protocol/runtime work for one thread activation. +/// It is not a user/assistant message pair; conversation belongs in +/// `ConversationItem`. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct CodexTurn { + pub codex_turn_id: CodexTurnId, + pub thread_id: AgentThreadId, + pub execution: ExecutionWindow, + /// Conversation items that directly triggered this activation, when known. + pub input_item_ids: Vec, +} diff --git a/codex-rs/rollout-trace/src/payload.rs b/codex-rs/rollout-trace/src/payload.rs new file mode 100644 index 0000000000000000000000000000000000000000..5efc7dc3eebece309a8a579948dcf23f2754ef28 --- /dev/null +++ b/codex-rs/rollout-trace/src/payload.rs @@ -0,0 +1,49 @@ +//! References to heavyweight trace payloads stored outside the reduced graph. + +use serde::Deserialize; +use serde::Serialize; + +/// Stable identifier for one raw payload inside a rollout bundle. +pub type RawPayloadId = String; + +/// Reference to a raw request/response/log payload. +/// +/// `RolloutTrace` stores these references so normal timeline and conversation +/// rendering does not require the browser or reducer output to inline every +/// upstream request, tool response, or terminal log. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct RawPayloadRef { + pub raw_payload_id: RawPayloadId, + /// Payload role. This lets details UI choose syntax highlighting and labels + /// without opening the payload file first. + pub kind: RawPayloadKind, + /// Path relative to the trace bundle root. + /// + /// The writer always materializes payloads as bundle-local files. Keeping + /// this as a plain path avoids exposing storage modes we do not produce. + pub path: String, +} + +/// Coarse role of a raw payload. +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type", content = "value")] +pub enum RawPayloadKind { + InferenceRequest, + /// Full upstream inference response or non-delta response stream summary. + InferenceResponse, + CompactionRequest, + /// Trace-only checkpoint captured when processed replacement history is installed. + CompactionCheckpoint, + CompactionResponse, + ToolInvocation, + ToolResult, + /// Raw runtime/protocol observation for an executing tool. + ToolRuntimeEvent, + /// Raw terminal runtime event or stream shard. + TerminalRuntimeEvent, + ProtocolEvent, + /// One-shot metadata captured when a Codex session/thread starts. + SessionMetadata, + /// Runtime notification payload carried when a child agent reports back to its parent. + AgentResult, +} diff --git a/codex-rs/rollout-trace/src/protocol_event.rs b/codex-rs/rollout-trace/src/protocol_event.rs new file mode 100644 index 0000000000000000000000000000000000000000..dadc4bb0acf9563cf2c39156ac129ac2fba5a4e0 --- /dev/null +++ b/codex-rs/rollout-trace/src/protocol_event.rs @@ -0,0 +1,560 @@ +//! Mapping from Codex protocol events into raw rollout-trace events. +//! +//! The session layer already emits protocol events for turn lifecycle, terminal +//! sessions, patch application, MCP calls, and collaboration tools. Rollout +//! tracing reuses those observations instead of adding another set of hooks in +//! `codex-core`: this module translates the protocol surface into the smaller +//! trace vocabulary and keeps the mapping isolated inside `codex-rollout-trace`. +//! +//! The long explicit `EventMsg` matches are intentional. Most protocol events +//! are not trace runtime boundaries, but spelling them out makes new protocol +//! variants a compile-time prompt to decide whether the trace should capture +//! them. + +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::ExecCommandBeginEvent; +use codex_protocol::protocol::ExecCommandEndEvent; +use codex_protocol::protocol::ExecCommandSource; +use codex_protocol::protocol::ExecCommandStatus; +use codex_protocol::protocol::McpToolCallBeginEvent; +use codex_protocol::protocol::McpToolCallEndEvent; +use codex_protocol::protocol::PatchApplyBeginEvent; +use codex_protocol::protocol::PatchApplyEndEvent; +use codex_protocol::protocol::PatchApplyStatus; +use codex_protocol::protocol::SubAgentActivityEvent; +use codex_protocol::protocol::SubAgentActivityKind; +use codex_protocol::protocol::TurnAbortReason; +use serde::Serialize; +use std::time::Duration; + +use crate::AgentThreadId; +use crate::CodexTurnId; +use crate::ExecutionStatus; +use crate::RawTraceEventPayload; + +pub(crate) struct CodexTurnTraceEvent { + pub context_turn_id: CodexTurnId, + pub payload: RawTraceEventPayload, +} + +pub(crate) fn codex_turn_trace_event( + thread_id: AgentThreadId, + default_turn_id: &str, + event: &EventMsg, +) -> Option { + match event { + EventMsg::TurnStarted(event) => { + let codex_turn_id = event.turn_id.clone(); + Some(CodexTurnTraceEvent { + context_turn_id: codex_turn_id.clone(), + payload: RawTraceEventPayload::CodexTurnStarted { + codex_turn_id, + thread_id, + }, + }) + } + EventMsg::TurnComplete(event) => { + let codex_turn_id = event.turn_id.clone(); + Some(CodexTurnTraceEvent { + context_turn_id: codex_turn_id.clone(), + payload: RawTraceEventPayload::CodexTurnEnded { + codex_turn_id, + status: ExecutionStatus::Completed, + }, + }) + } + EventMsg::TurnAborted(event) => { + let codex_turn_id = event + .turn_id + .clone() + .unwrap_or_else(|| default_turn_id.to_string()); + Some(CodexTurnTraceEvent { + context_turn_id: codex_turn_id.clone(), + payload: RawTraceEventPayload::CodexTurnEnded { + codex_turn_id, + status: execution_status_for_abort_reason(&event.reason), + }, + }) + } + _ => None, + } +} + +pub(crate) enum ToolRuntimeTraceEvent<'a> { + Started { + tool_call_id: &'a str, + payload: ToolRuntimePayload<'a>, + }, + Ended { + tool_call_id: &'a str, + status: ExecutionStatus, + payload: ToolRuntimePayload<'a>, + }, +} + +/// Borrowed protocol payload that should be persisted as tool runtime data. +/// +/// The trace wants the exact protocol payload shape for E2E debugging, while +/// reducers consume the surrounding typed trace events. This enum lets the +/// recorder serialize the original event by reference, without first cloning it +/// or converting it through `serde_json::Value`. +pub(crate) enum ToolRuntimePayload<'a> { + ExecCommandBegin(&'a ExecCommandBeginEvent), + ExecCommandEnd(&'a ExecCommandEndEvent), + PatchApplyBegin(&'a PatchApplyBeginEvent), + PatchApplyEnd(&'a PatchApplyEndEvent), + McpToolCallBegin(&'a McpToolCallBeginEvent), + McpToolCallEnd(&'a McpToolCallEndEvent), + CollabAgentSpawnBegin(&'a codex_protocol::protocol::CollabAgentSpawnBeginEvent), + CollabAgentSpawnEnd(&'a codex_protocol::protocol::CollabAgentSpawnEndEvent), + CollabAgentInteractionBegin(&'a codex_protocol::protocol::CollabAgentInteractionBeginEvent), + CollabAgentInteractionEnd(&'a codex_protocol::protocol::CollabAgentInteractionEndEvent), + CollabWaitingBegin(&'a codex_protocol::protocol::CollabWaitingBeginEvent), + CollabWaitingEnd(&'a codex_protocol::protocol::CollabWaitingEndEvent), + CollabCloseBegin(&'a codex_protocol::protocol::CollabCloseBeginEvent), + CollabCloseEnd(&'a codex_protocol::protocol::CollabCloseEndEvent), + SubAgentActivity(&'a SubAgentActivityEvent), +} + +impl Serialize for ToolRuntimePayload<'_> { + fn serialize(&self, serializer: S) -> Result + where + S: serde::Serializer, + { + match self { + ToolRuntimePayload::ExecCommandBegin(event) => { + ExecCommandBeginTracePayload::from(*event).serialize(serializer) + } + ToolRuntimePayload::ExecCommandEnd(event) => { + ExecCommandEndTracePayload::from(*event).serialize(serializer) + } + ToolRuntimePayload::PatchApplyBegin(event) => event.serialize(serializer), + ToolRuntimePayload::PatchApplyEnd(event) => event.serialize(serializer), + ToolRuntimePayload::McpToolCallBegin(event) => event.serialize(serializer), + ToolRuntimePayload::McpToolCallEnd(event) => event.serialize(serializer), + ToolRuntimePayload::CollabAgentSpawnBegin(event) => event.serialize(serializer), + ToolRuntimePayload::CollabAgentSpawnEnd(event) => event.serialize(serializer), + ToolRuntimePayload::CollabAgentInteractionBegin(event) => event.serialize(serializer), + ToolRuntimePayload::CollabAgentInteractionEnd(event) => event.serialize(serializer), + ToolRuntimePayload::CollabWaitingBegin(event) => event.serialize(serializer), + ToolRuntimePayload::CollabWaitingEnd(event) => event.serialize(serializer), + ToolRuntimePayload::CollabCloseBegin(event) => event.serialize(serializer), + ToolRuntimePayload::CollabCloseEnd(event) => event.serialize(serializer), + ToolRuntimePayload::SubAgentActivity(event) => event.serialize(serializer), + } + } +} + +/// Rollout-trace representation of an exec begin event. +/// +/// Rollout traces share the rollout compatibility requirement that paths remain path-flavored +/// strings on disk, even though live events carry `PathUri` internally. +#[derive(Serialize)] +struct ExecCommandBeginTracePayload<'a> { + call_id: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + plugin_id: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + script_path: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + process_id: Option<&'a str>, + turn_id: &'a str, + started_at_ms: i64, + command: &'a [String], + cwd: String, + parsed_cmd: &'a [codex_protocol::parse_command::ParsedCommand], + source: ExecCommandSource, + #[serde(skip_serializing_if = "Option::is_none")] + interaction_input: Option<&'a str>, +} + +impl<'a> From<&'a ExecCommandBeginEvent> for ExecCommandBeginTracePayload<'a> { + fn from(event: &'a ExecCommandBeginEvent) -> Self { + let ExecCommandBeginEvent { + call_id, + plugin_id, + script_path, + process_id, + turn_id, + started_at_ms, + command, + cwd, + parsed_cmd, + source, + interaction_input, + } = event; + Self { + call_id, + plugin_id: plugin_id.as_deref(), + script_path: script_path.as_deref(), + process_id: process_id.as_deref(), + turn_id, + started_at_ms: *started_at_ms, + command, + cwd: cwd.inferred_native_path_string(), + parsed_cmd, + source: *source, + interaction_input: interaction_input.as_deref(), + } + } +} + +/// Rollout-trace representation of an exec end event. +/// +/// Like [`ExecCommandBeginTracePayload`], this renders `cwd` as an inferred native path to preserve +/// the on-disk format rather than serializing the internal `PathUri`. +#[derive(Serialize)] +struct ExecCommandEndTracePayload<'a> { + call_id: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + plugin_id: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + script_path: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + process_id: Option<&'a str>, + turn_id: &'a str, + completed_at_ms: i64, + command: &'a [String], + cwd: String, + parsed_cmd: &'a [codex_protocol::parse_command::ParsedCommand], + source: ExecCommandSource, + #[serde(skip_serializing_if = "Option::is_none")] + interaction_input: Option<&'a str>, + stdout: &'a str, + stderr: &'a str, + aggregated_output: &'a str, + exit_code: i32, + duration: Duration, + formatted_output: &'a str, + status: &'a ExecCommandStatus, +} + +impl<'a> From<&'a ExecCommandEndEvent> for ExecCommandEndTracePayload<'a> { + fn from(event: &'a ExecCommandEndEvent) -> Self { + let ExecCommandEndEvent { + call_id, + plugin_id, + script_path, + process_id, + turn_id, + completed_at_ms, + command, + cwd, + parsed_cmd, + source, + interaction_input, + stdout, + stderr, + aggregated_output, + exit_code, + duration, + formatted_output, + status, + } = event; + Self { + call_id, + plugin_id: plugin_id.as_deref(), + script_path: script_path.as_deref(), + process_id: process_id.as_deref(), + turn_id, + completed_at_ms: *completed_at_ms, + command, + cwd: cwd.inferred_native_path_string(), + parsed_cmd, + source: *source, + interaction_input: interaction_input.as_deref(), + stdout, + stderr, + aggregated_output, + exit_code: *exit_code, + duration: *duration, + formatted_output, + status, + } + } +} + +pub(crate) fn tool_runtime_trace_event(event: &EventMsg) -> Option> { + match event { + EventMsg::ExecCommandBegin(event) if event.source != ExecCommandSource::UserShell => { + Some(ToolRuntimeTraceEvent::Started { + tool_call_id: &event.call_id, + payload: ToolRuntimePayload::ExecCommandBegin(event), + }) + } + EventMsg::ExecCommandEnd(event) if event.source != ExecCommandSource::UserShell => { + Some(ToolRuntimeTraceEvent::Ended { + tool_call_id: &event.call_id, + status: event.status.trace_execution_status(), + payload: ToolRuntimePayload::ExecCommandEnd(event), + }) + } + EventMsg::PatchApplyBegin(event) => Some(ToolRuntimeTraceEvent::Started { + tool_call_id: &event.call_id, + payload: ToolRuntimePayload::PatchApplyBegin(event), + }), + EventMsg::PatchApplyEnd(event) => Some(ToolRuntimeTraceEvent::Ended { + tool_call_id: &event.call_id, + status: event.status.trace_execution_status(), + payload: ToolRuntimePayload::PatchApplyEnd(event), + }), + EventMsg::McpToolCallBegin(event) => Some(ToolRuntimeTraceEvent::Started { + tool_call_id: &event.call_id, + payload: ToolRuntimePayload::McpToolCallBegin(event), + }), + EventMsg::McpToolCallEnd(event) => Some(ToolRuntimeTraceEvent::Ended { + tool_call_id: &event.call_id, + status: if event.result.is_ok() { + ExecutionStatus::Completed + } else { + ExecutionStatus::Failed + }, + payload: ToolRuntimePayload::McpToolCallEnd(event), + }), + EventMsg::CollabAgentSpawnBegin(event) => Some(ToolRuntimeTraceEvent::Started { + tool_call_id: &event.call_id, + payload: ToolRuntimePayload::CollabAgentSpawnBegin(event), + }), + EventMsg::CollabAgentSpawnEnd(event) => Some(ToolRuntimeTraceEvent::Ended { + tool_call_id: &event.call_id, + // A spawn end without a child thread id means the runtime boundary + // finished without creating the requested child thread. + status: if event.new_thread_id.is_some() { + ExecutionStatus::Completed + } else { + ExecutionStatus::Failed + }, + payload: ToolRuntimePayload::CollabAgentSpawnEnd(event), + }), + EventMsg::CollabAgentInteractionBegin(event) => Some(ToolRuntimeTraceEvent::Started { + tool_call_id: &event.call_id, + payload: ToolRuntimePayload::CollabAgentInteractionBegin(event), + }), + EventMsg::CollabAgentInteractionEnd(event) => Some(ToolRuntimeTraceEvent::Ended { + tool_call_id: &event.call_id, + status: ExecutionStatus::Completed, + payload: ToolRuntimePayload::CollabAgentInteractionEnd(event), + }), + EventMsg::CollabWaitingBegin(event) => Some(ToolRuntimeTraceEvent::Started { + tool_call_id: &event.call_id, + payload: ToolRuntimePayload::CollabWaitingBegin(event), + }), + EventMsg::CollabWaitingEnd(event) => Some(ToolRuntimeTraceEvent::Ended { + tool_call_id: &event.call_id, + status: ExecutionStatus::Completed, + payload: ToolRuntimePayload::CollabWaitingEnd(event), + }), + EventMsg::CollabCloseBegin(event) => Some(ToolRuntimeTraceEvent::Started { + tool_call_id: &event.call_id, + payload: ToolRuntimePayload::CollabCloseBegin(event), + }), + EventMsg::CollabCloseEnd(event) => Some(ToolRuntimeTraceEvent::Ended { + tool_call_id: &event.call_id, + status: ExecutionStatus::Completed, + payload: ToolRuntimePayload::CollabCloseEnd(event), + }), + EventMsg::SubAgentActivity(event) if event.kind != SubAgentActivityKind::Completed => { + Some(ToolRuntimeTraceEvent::Ended { + tool_call_id: &event.event_id, + status: ExecutionStatus::Completed, + payload: ToolRuntimePayload::SubAgentActivity(event), + }) + } + EventMsg::SubAgentActivity(_) => None, + EventMsg::Error(_) + | EventMsg::Warning(_) + | EventMsg::AuthRecoveryStarted(_) + | EventMsg::AuthRecoveryCompleted(_) + | EventMsg::GuardianWarning(_) + | EventMsg::SafetyBuffering(_) + | EventMsg::RealtimeConversationStarted(_) + | EventMsg::RealtimeConversationRealtime(_) + | EventMsg::RealtimeConversationClosed(_) + | EventMsg::RealtimeConversationSdp(_) + | EventMsg::ModelReroute(_) + | EventMsg::ModelVerification(_) + | EventMsg::TurnModerationMetadata(_) + | EventMsg::ContextCompacted(_) + | EventMsg::ThreadRolledBack(_) + | EventMsg::ThreadGoalUpdated(_) + | EventMsg::ThreadQueueChanged(_) + | EventMsg::TurnStarted(_) + | EventMsg::ThreadSettingsApplied(_) + | EventMsg::TurnComplete(_) + | EventMsg::TokenCount(_) + | EventMsg::AgentMessage(_) + | EventMsg::UserMessage(_) + | EventMsg::AgentReasoning(_) + | EventMsg::AgentReasoningRawContent(_) + | EventMsg::AgentReasoningSectionBreak(_) + | EventMsg::SessionConfigured(_) + | EventMsg::EnvironmentConnected(_) + | EventMsg::EnvironmentDisconnected(_) + | EventMsg::McpStartupUpdate(_) + | EventMsg::McpStartupComplete(_) + | EventMsg::WebSearchBegin(_) + | EventMsg::WebSearchEnd(_) + | EventMsg::ImageGenerationBegin(_) + | EventMsg::ImageGenerationEnd(_) + | EventMsg::ViewImageToolCall(_) + | EventMsg::ExecCommandBegin(_) + | EventMsg::ExecCommandOutputDelta(_) + | EventMsg::TerminalInteraction(_) + | EventMsg::ExecCommandEnd(_) + | EventMsg::ExecApprovalRequest(_) + | EventMsg::RequestPermissions(_) + | EventMsg::RequestUserInput(_) + | EventMsg::DynamicToolCallRequest(_) + | EventMsg::DynamicToolCallResponse(_) + | EventMsg::ElicitationRequest(_) + | EventMsg::ApplyPatchApprovalRequest(_) + | EventMsg::GuardianAssessment(_) + | EventMsg::DeprecationNotice(_) + | EventMsg::StreamError(_) + | EventMsg::PatchApplyUpdated(_) + | EventMsg::TurnDiff(_) + | EventMsg::RealtimeConversationListVoicesResponse(_) + | EventMsg::PlanUpdate(_) + | EventMsg::TurnAborted(_) + | EventMsg::ShutdownComplete + | EventMsg::EnteredReviewMode(_) + | EventMsg::ExitedReviewMode(_) + | EventMsg::RawResponseItem(_) + | EventMsg::RawResponseCompleted(_) + | EventMsg::ItemStarted(_) + | EventMsg::ItemCompleted(_) + | EventMsg::HookStarted(_) + | EventMsg::HookCompleted(_) + | EventMsg::AgentMessageContentDelta(_) + | EventMsg::PlanDelta(_) + | EventMsg::ReasoningContentDelta(_) + | EventMsg::ReasoningRawContentDelta(_) + | EventMsg::CollabResumeBegin(_) + | EventMsg::CollabResumeEnd(_) => None, + } +} + +pub(crate) fn wrapped_protocol_event_type(event: &EventMsg) -> Option<&'static str> { + match event { + EventMsg::SessionConfigured(_) => Some("session_configured"), + EventMsg::TurnStarted(_) => Some("turn_started"), + EventMsg::TurnComplete(_) => Some("turn_complete"), + EventMsg::TurnAborted(_) => Some("turn_aborted"), + EventMsg::ThreadRolledBack(_) => Some("thread_rolled_back"), + EventMsg::Error(_) => Some("error"), + EventMsg::Warning(_) => Some("warning"), + EventMsg::ShutdownComplete => Some("shutdown_complete"), + EventMsg::AuthRecoveryStarted(_) + | EventMsg::AuthRecoveryCompleted(_) + | EventMsg::GuardianWarning(_) + | EventMsg::SafetyBuffering(_) + | EventMsg::RealtimeConversationStarted(_) + | EventMsg::RealtimeConversationRealtime(_) + | EventMsg::RealtimeConversationClosed(_) + | EventMsg::RealtimeConversationSdp(_) + | EventMsg::ModelReroute(_) + | EventMsg::ModelVerification(_) + | EventMsg::TurnModerationMetadata(_) + | EventMsg::ContextCompacted(_) + | EventMsg::ThreadSettingsApplied(_) + | EventMsg::EnvironmentConnected(_) + | EventMsg::EnvironmentDisconnected(_) + | EventMsg::TokenCount(_) + | EventMsg::AgentMessage(_) + | EventMsg::UserMessage(_) + | EventMsg::AgentReasoning(_) + | EventMsg::AgentReasoningRawContent(_) + | EventMsg::AgentReasoningSectionBreak(_) + | EventMsg::ThreadGoalUpdated(_) + | EventMsg::ThreadQueueChanged(_) + | EventMsg::McpStartupUpdate(_) + | EventMsg::McpStartupComplete(_) + | EventMsg::McpToolCallBegin(_) + | EventMsg::McpToolCallEnd(_) + | EventMsg::WebSearchBegin(_) + | EventMsg::WebSearchEnd(_) + | EventMsg::ImageGenerationBegin(_) + | EventMsg::ImageGenerationEnd(_) + | EventMsg::ViewImageToolCall(_) + | EventMsg::ExecCommandBegin(_) + | EventMsg::ExecCommandOutputDelta(_) + | EventMsg::TerminalInteraction(_) + | EventMsg::ExecCommandEnd(_) + | EventMsg::ExecApprovalRequest(_) + | EventMsg::RequestPermissions(_) + | EventMsg::RequestUserInput(_) + | EventMsg::DynamicToolCallRequest(_) + | EventMsg::DynamicToolCallResponse(_) + | EventMsg::ElicitationRequest(_) + | EventMsg::ApplyPatchApprovalRequest(_) + | EventMsg::GuardianAssessment(_) + | EventMsg::DeprecationNotice(_) + | EventMsg::StreamError(_) + | EventMsg::PatchApplyBegin(_) + | EventMsg::PatchApplyUpdated(_) + | EventMsg::PatchApplyEnd(_) + | EventMsg::TurnDiff(_) + | EventMsg::RealtimeConversationListVoicesResponse(_) + | EventMsg::PlanUpdate(_) + | EventMsg::EnteredReviewMode(_) + | EventMsg::ExitedReviewMode(_) + | EventMsg::RawResponseItem(_) + | EventMsg::RawResponseCompleted(_) + | EventMsg::ItemStarted(_) + | EventMsg::ItemCompleted(_) + | EventMsg::HookStarted(_) + | EventMsg::HookCompleted(_) + | EventMsg::AgentMessageContentDelta(_) + | EventMsg::PlanDelta(_) + | EventMsg::ReasoningContentDelta(_) + | EventMsg::ReasoningRawContentDelta(_) + | EventMsg::CollabAgentSpawnBegin(_) + | EventMsg::CollabAgentSpawnEnd(_) + | EventMsg::CollabAgentInteractionBegin(_) + | EventMsg::CollabAgentInteractionEnd(_) + | EventMsg::CollabWaitingBegin(_) + | EventMsg::CollabWaitingEnd(_) + | EventMsg::CollabCloseBegin(_) + | EventMsg::CollabCloseEnd(_) + | EventMsg::CollabResumeBegin(_) + | EventMsg::CollabResumeEnd(_) + | EventMsg::SubAgentActivity(_) => None, + } +} + +trait TraceExecutionStatus { + fn trace_execution_status(&self) -> ExecutionStatus; +} + +impl TraceExecutionStatus for ExecCommandStatus { + fn trace_execution_status(&self) -> ExecutionStatus { + match self { + ExecCommandStatus::Completed => ExecutionStatus::Completed, + ExecCommandStatus::Failed => ExecutionStatus::Failed, + ExecCommandStatus::Declined => ExecutionStatus::Cancelled, + } + } +} + +impl TraceExecutionStatus for PatchApplyStatus { + fn trace_execution_status(&self) -> ExecutionStatus { + match self { + PatchApplyStatus::Completed => ExecutionStatus::Completed, + PatchApplyStatus::Failed => ExecutionStatus::Failed, + PatchApplyStatus::Declined => ExecutionStatus::Cancelled, + } + } +} + +fn execution_status_for_abort_reason(reason: &TurnAbortReason) -> ExecutionStatus { + match reason { + TurnAbortReason::Interrupted + | TurnAbortReason::Replaced + | TurnAbortReason::ReviewEnded + | TurnAbortReason::BudgetLimited => ExecutionStatus::Cancelled, + } +} + +#[cfg(test)] +#[path = "protocol_event_tests.rs"] +mod tests; diff --git a/codex-rs/rollout-trace/src/protocol_event_tests.rs b/codex-rs/rollout-trace/src/protocol_event_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..c7c53c3ef32a0541546b5c9facdc7748a40c63e9 --- /dev/null +++ b/codex-rs/rollout-trace/src/protocol_event_tests.rs @@ -0,0 +1,151 @@ +use codex_protocol::AgentPath; +use codex_protocol::ThreadId; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::ExecCommandBeginEvent; +use codex_protocol::protocol::ExecCommandEndEvent; +use codex_protocol::protocol::ExecCommandSource; +use codex_protocol::protocol::ExecCommandStatus; +use codex_protocol::protocol::SubAgentActivityEvent; +use codex_protocol::protocol::SubAgentActivityKind; +use pretty_assertions::assert_eq; +use serde_json::json; +use std::time::Duration; + +use super::ToolRuntimeTraceEvent; +use super::tool_runtime_trace_event; +use crate::ExecutionStatus; + +#[test] +fn sub_agent_activity_is_a_terminal_tool_runtime_event() -> anyhow::Result<()> { + let agent_thread_id = ThreadId::new(); + let event = EventMsg::SubAgentActivity(SubAgentActivityEvent { + event_id: "call-spawn".to_string(), + occurred_at_ms: 1234, + agent_thread_id, + agent_path: AgentPath::try_from("/root/reviewer").map_err(anyhow::Error::msg)?, + kind: SubAgentActivityKind::Started, + }); + + let Some(ToolRuntimeTraceEvent::Ended { + tool_call_id, + status, + payload, + }) = tool_runtime_trace_event(&event) + else { + panic!("expected terminal tool runtime event"); + }; + + assert_eq!(tool_call_id, "call-spawn"); + assert_eq!(status, ExecutionStatus::Completed); + assert_eq!( + serde_json::to_value(payload)?, + json!({ + "event_id": "call-spawn", + "occurred_at_ms": 1234, + "agent_thread_id": agent_thread_id, + "agent_path": "/root/reviewer", + "kind": "started" + }) + ); + Ok(()) +} + +#[test] +fn completed_sub_agent_activity_is_not_a_tool_runtime_event() -> anyhow::Result<()> { + let event = EventMsg::SubAgentActivity(SubAgentActivityEvent { + event_id: "child-turn-completed".to_string(), + occurred_at_ms: 1234, + agent_thread_id: ThreadId::new(), + agent_path: AgentPath::try_from("/root/reviewer").map_err(anyhow::Error::msg)?, + kind: SubAgentActivityKind::Completed, + }); + + assert!(tool_runtime_trace_event(&event).is_none()); + Ok(()) +} + +#[test] +fn exec_command_trace_payloads_use_inferred_native_cwd() -> anyhow::Result<()> { + // Convention inference depends on the URI spelling, not the test host, so exercise both + // Windows and POSIX paths on every platform. + let begin = EventMsg::ExecCommandBegin(ExecCommandBeginEvent { + call_id: "call-begin".to_string(), + plugin_id: Some("sample@openai-curated".to_string()), + script_path: Some("scripts/run.py".to_string()), + process_id: Some("process-1".to_string()), + turn_id: "turn-1".to_string(), + started_at_ms: 1234, + command: vec!["pwd".to_string()], + cwd: "file:///C:/windows".parse()?, + parsed_cmd: Vec::new(), + source: ExecCommandSource::Agent, + interaction_input: None, + }); + let end = EventMsg::ExecCommandEnd(ExecCommandEndEvent { + call_id: "call-end".to_string(), + plugin_id: Some("sample@openai-curated".to_string()), + script_path: Some("scripts/run.py".to_string()), + process_id: None, + turn_id: "turn-1".to_string(), + completed_at_ms: 2345, + command: vec!["pwd".to_string()], + cwd: "file:///workspace/project".parse()?, + parsed_cmd: Vec::new(), + source: ExecCommandSource::UnifiedExecInteraction, + interaction_input: Some("input".to_string()), + stdout: "output".to_string(), + stderr: String::new(), + aggregated_output: "output".to_string(), + exit_code: 0, + duration: Duration::from_millis(250), + formatted_output: "output".to_string(), + status: ExecCommandStatus::Completed, + }); + + let Some(ToolRuntimeTraceEvent::Started { payload, .. }) = tool_runtime_trace_event(&begin) + else { + panic!("expected started tool runtime event"); + }; + assert_eq!( + serde_json::to_value(payload)?, + json!({ + "call_id": "call-begin", + "plugin_id": "sample@openai-curated", + "script_path": "scripts/run.py", + "process_id": "process-1", + "turn_id": "turn-1", + "started_at_ms": 1234, + "command": ["pwd"], + "cwd": r"C:\windows", + "parsed_cmd": [], + "source": "agent" + }) + ); + + let Some(ToolRuntimeTraceEvent::Ended { payload, .. }) = tool_runtime_trace_event(&end) else { + panic!("expected ended tool runtime event"); + }; + assert_eq!( + serde_json::to_value(payload)?, + json!({ + "call_id": "call-end", + "plugin_id": "sample@openai-curated", + "script_path": "scripts/run.py", + "turn_id": "turn-1", + "completed_at_ms": 2345, + "command": ["pwd"], + "cwd": "/workspace/project", + "parsed_cmd": [], + "source": "unified_exec_interaction", + "interaction_input": "input", + "stdout": "output", + "stderr": "", + "aggregated_output": "output", + "exit_code": 0, + "duration": {"secs": 0, "nanos": 250000000}, + "formatted_output": "output", + "status": "completed" + }) + ); + Ok(()) +} diff --git a/codex-rs/rollout-trace/src/raw_event.rs b/codex-rs/rollout-trace/src/raw_event.rs new file mode 100644 index 0000000000000000000000000000000000000000..bd9010d7bbee7cffdeb6d2ecde8537e2676ade21 --- /dev/null +++ b/codex-rs/rollout-trace/src/raw_event.rs @@ -0,0 +1,312 @@ +//! Append-only raw trace events. + +use crate::model::AgentThreadId; +use crate::model::CodeCellRuntimeStatus; +use crate::model::CodexTurnId; +use crate::model::CompactionId; +use crate::model::CompactionRequestId; +use crate::model::EdgeId; +use crate::model::ExecutionStatus; +use crate::model::InferenceCallId; +use crate::model::McpCallId; +use crate::model::ModelVisibleCallId; +use crate::model::RolloutStatus; +use crate::model::ToolCallId; +use crate::model::ToolCallKind; +use crate::model::ToolCallSummary; +use crate::payload::RawPayloadRef; +use serde::Deserialize; +use serde::Serialize; +use serde_json::Value; + +/// Monotonic sequence number assigned by the raw trace writer. +pub type RawEventSeq = u64; + +/// Current raw event envelope schema version. +pub(crate) const RAW_TRACE_EVENT_SCHEMA_VERSION: u32 = 1; + +/// One append-only raw trace event. +/// +/// Every event uses the same envelope so partial replay and corruption checks +/// can run before the reducer understands the event-specific payload. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct RawTraceEvent { + pub schema_version: u32, + /// Contiguous writer-assigned order inside one rollout event log. + pub seq: RawEventSeq, + /// Unix wall-clock timestamp in milliseconds. Use for display/latency. + pub wall_time_unix_ms: i64, + pub rollout_id: String, + pub thread_id: Option, + pub codex_turn_id: Option, + pub payload: RawTraceEventPayload, +} + +/// Writer-supplied context that appears in the raw event envelope. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct RawTraceEventContext { + pub thread_id: Option, + pub codex_turn_id: Option, +} + +/// Runtime requester as observed at the raw tool boundary. +/// +/// This intentionally uses runtime-local identifiers. The reducer is the only +/// place that maps these handles to graph identities such as `CodeCellId`. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum RawToolCallRequester { + Model, + CodeCell { + /// Runtime-local code-mode cell handle. + runtime_cell_id: String, + }, +} + +/// Typed payload for a raw trace event. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum RawTraceEventPayload { + RolloutStarted { + trace_id: String, + root_thread_id: AgentThreadId, + }, + RolloutEnded { + status: RolloutStatus, + }, + ThreadStarted { + thread_id: AgentThreadId, + /// Stable agent path. + agent_path: String, + metadata_payload: Option, + }, + ThreadEnded { + thread_id: AgentThreadId, + status: RolloutStatus, + }, + CodexTurnStarted { + codex_turn_id: CodexTurnId, + thread_id: AgentThreadId, + }, + CodexTurnEnded { + codex_turn_id: CodexTurnId, + status: ExecutionStatus, + }, + InferenceStarted { + inference_call_id: InferenceCallId, + thread_id: AgentThreadId, + codex_turn_id: CodexTurnId, + model: String, + provider_name: String, + request_payload: RawPayloadRef, + }, + InferenceCompleted { + inference_call_id: InferenceCallId, + /// Responses API `response.id`; used by `previous_response_id`. + response_id: Option, + /// Provider transport request id, such as `x-request-id`. + upstream_request_id: Option, + response_payload: RawPayloadRef, + }, + InferenceFailed { + inference_call_id: InferenceCallId, + /// Provider transport request id, such as `x-request-id`, when the + /// provider returned one before the stream failed. + upstream_request_id: Option, + error: String, + /// Partial response payload, when stream events arrived before failure. + partial_response_payload: Option, + }, + InferenceCancelled { + inference_call_id: InferenceCallId, + /// Provider transport request id, such as `x-request-id`, when observed + /// before Codex stopped consuming the stream. + upstream_request_id: Option, + /// Why Codex stopped consuming the provider stream before a terminal response event. + reason: String, + /// Completed output items observed before cancellation, if any. + partial_response_payload: Option, + }, + ToolCallStarted { + tool_call_id: ToolCallId, + /// Protocol/model call ID when this runtime call came from model output. + model_visible_call_id: Option, + /// Code-mode runtime bridge ID when model-authored code issued this call. + code_mode_runtime_tool_id: Option, + /// Runtime requester that caused this tool lifecycle. + requester: RawToolCallRequester, + kind: ToolCallKind, + summary: ToolCallSummary, + invocation_payload: Option, + }, + /// Bridge correlation UUID assigned only when a tool reaches an MCP backend. + McpToolCallCorrelationAssigned { + tool_call_id: ToolCallId, + mcp_call_id: McpCallId, + }, + ToolCallRuntimeStarted { + tool_call_id: ToolCallId, + /// Runtime/protocol observation for how Codex began executing the tool. + runtime_payload: RawPayloadRef, + }, + ToolCallRuntimeEnded { + tool_call_id: ToolCallId, + status: ExecutionStatus, + /// Runtime/protocol observation for how Codex finished executing the tool. + runtime_payload: RawPayloadRef, + }, + ToolCallEnded { + tool_call_id: ToolCallId, + status: ExecutionStatus, + result_payload: Option, + }, + CodeCellStarted { + /// Runtime-local handle allocated by code mode for waits and nested tools. + runtime_cell_id: String, + /// Custom tool call id on the model-visible `exec` item. + model_visible_call_id: ModelVisibleCallId, + /// JavaScript source after the public `exec` wrapper has been parsed. + source_js: String, + }, + CodeCellInitialResponse { + /// Runtime-local handle, matching `CodeCellStarted`. + runtime_cell_id: String, + status: CodeCellRuntimeStatus, + response_payload: Option, + }, + CodeCellEnded { + /// Runtime-local handle, matching `CodeCellStarted`. + runtime_cell_id: String, + status: CodeCellRuntimeStatus, + response_payload: Option, + }, + CompactionRequestStarted { + compaction_id: CompactionId, + compaction_request_id: CompactionRequestId, + thread_id: AgentThreadId, + codex_turn_id: CodexTurnId, + model: String, + provider_name: String, + request_payload: RawPayloadRef, + }, + CompactionRequestCompleted { + compaction_id: CompactionId, + compaction_request_id: CompactionRequestId, + response_payload: RawPayloadRef, + }, + CompactionRequestFailed { + compaction_id: CompactionId, + compaction_request_id: CompactionRequestId, + error: String, + }, + /// Checkpoint installation event for remote-compacted replacement history. + CompactionInstalled { + compaction_id: CompactionId, + /// Trace-only checkpoint payload. Do not route this through public UI protocol. + checkpoint_payload: RawPayloadRef, + }, + /// Multi-agent v2 child-to-parent completion delivery. + AgentResultObserved { + edge_id: EdgeId, + child_thread_id: AgentThreadId, + child_codex_turn_id: CodexTurnId, + parent_thread_id: AgentThreadId, + message: String, + /// Raw notification payload. This is evidence for the runtime delivery, + /// not the parent-side model-visible item. + carried_payload: Option, + }, + /// Existing UI/protocol event wrapped into trace format. + ProtocolEventObserved { + event_type: String, + event_payload: RawPayloadRef, + }, + /// Structured payload for early instrumentation before a dedicated variant exists. + Other { + kind: String, + summary: String, + payloads: Vec, + /// Small structured metadata. Large data belongs in `payloads`. + metadata: Value, + }, +} + +impl RawTraceEventPayload { + /// Raw payload refs that must exist before this raw event is appended. + pub(crate) fn raw_payload_refs(&self) -> Vec<&RawPayloadRef> { + match self { + RawTraceEventPayload::RolloutStarted { .. } + | RawTraceEventPayload::RolloutEnded { .. } + | RawTraceEventPayload::ThreadEnded { .. } + | RawTraceEventPayload::CodexTurnStarted { .. } + | RawTraceEventPayload::CodexTurnEnded { .. } + | RawTraceEventPayload::CompactionRequestFailed { .. } + | RawTraceEventPayload::CodeCellStarted { .. } + | RawTraceEventPayload::McpToolCallCorrelationAssigned { .. } + | RawTraceEventPayload::AgentResultObserved { + carried_payload: None, + .. + } => Vec::new(), + RawTraceEventPayload::ThreadStarted { + metadata_payload, .. + } => metadata_payload.iter().collect(), + RawTraceEventPayload::InferenceStarted { + request_payload, .. + } + | RawTraceEventPayload::InferenceCompleted { + response_payload: request_payload, + .. + } + | RawTraceEventPayload::CompactionRequestStarted { + request_payload, .. + } + | RawTraceEventPayload::CompactionRequestCompleted { + response_payload: request_payload, + .. + } + | RawTraceEventPayload::CompactionInstalled { + checkpoint_payload: request_payload, + .. + } + | RawTraceEventPayload::ProtocolEventObserved { + event_payload: request_payload, + .. + } => vec![request_payload], + RawTraceEventPayload::InferenceFailed { + partial_response_payload, + .. + } + | RawTraceEventPayload::InferenceCancelled { + partial_response_payload, + .. + } + | RawTraceEventPayload::ToolCallStarted { + invocation_payload: partial_response_payload, + .. + } + | RawTraceEventPayload::ToolCallEnded { + result_payload: partial_response_payload, + .. + } + | RawTraceEventPayload::CodeCellInitialResponse { + response_payload: partial_response_payload, + .. + } + | RawTraceEventPayload::CodeCellEnded { + response_payload: partial_response_payload, + .. + } => partial_response_payload.iter().collect(), + RawTraceEventPayload::AgentResultObserved { + carried_payload: Some(carried_payload), + .. + } => vec![carried_payload], + RawTraceEventPayload::ToolCallRuntimeStarted { + runtime_payload, .. + } + | RawTraceEventPayload::ToolCallRuntimeEnded { + runtime_payload, .. + } => vec![runtime_payload], + RawTraceEventPayload::Other { payloads, .. } => payloads.iter().collect(), + } + } +} diff --git a/codex-rs/rollout-trace/src/reducer/code_cell.rs b/codex-rs/rollout-trace/src/reducer/code_cell.rs new file mode 100644 index 0000000000000000000000000000000000000000..11ad55f48ebd7775f893ba25ead3c5c9b7ee52bb --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/code_cell.rs @@ -0,0 +1,738 @@ +//! Code-mode reduction. +//! +//! A code cell is the runtime parent for model-authored `exec` +//! JavaScript. Nested tools, waits, and terminal operations hang off this +//! object so viewers can inspect runtime work without flattening it into the +//! model-visible conversation. +//! +//! The reducer has to reconcile two clocks: +//! - model-visible items come from inference request/response payloads; +//! - runtime work starts as soon as Codex dispatches the tool. +//! +//! In real traces `CodeCellStarted` can arrive before the inference completion +//! payload that contains the `custom_tool_call` item. We therefore queue starts +//! until their source conversation item exists, then attach runtime edges. + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use serde_json::Value; + +use super::TraceReducer; +use crate::model::CodeCell; +use crate::model::CodeCellId; +use crate::model::CodeCellRuntimeStatus; +use crate::model::ConversationItemKind; +use crate::model::ExecutionStatus; +use crate::model::ExecutionWindow; +use crate::model::ProducerRef; +use crate::model::ToolCallId; +use crate::model::ToolCallRequester; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawEventSeq; +use crate::raw_event::RawToolCallRequester; + +/// Runtime start payload for one model-authored code-mode exec call. +/// +/// The reduced id is already derived from the model-visible call id before this +/// reaches the code-cell reducer, so the reducer can reconcile runtime lifecycle +/// events against a stable graph identity. +pub(super) struct StartedCodeCell { + pub(super) code_cell_id: CodeCellId, + pub(super) runtime_cell_id: String, + pub(super) model_visible_call_id: crate::model::ModelVisibleCallId, + pub(super) source_js: String, +} + +/// Queued code-cell start waiting for its model-visible source item. +/// +/// Code execution can begin before inference stream completion records the +/// custom-tool call item that authored it. This wrapper keeps the original +/// event timing intact until that source item exists. +pub(super) struct PendingCodeCellStart { + pub(super) seq: RawEventSeq, + pub(super) wall_time_unix_ms: i64, + pub(super) thread_id: String, + pub(super) codex_turn_id: Option, + pub(super) started: StartedCodeCell, +} + +/// Lifecycle event observed before a queued code cell has materialized. +/// +/// These events are replayed after the start is resolved so failed or very fast +/// cells do not lose runtime status while preserving source-item ownership. +pub(super) struct PendingCodeCellLifecycleEvent { + pub(super) seq: RawEventSeq, + pub(super) wall_time_unix_ms: i64, + pub(super) kind: PendingCodeCellLifecycleEventKind, +} + +/// Runtime lifecycle transitions that can arrive while a code-cell start is queued. +pub(super) enum PendingCodeCellLifecycleEventKind { + InitialResponse { + runtime_cell_id: String, + status: CodeCellRuntimeStatus, + }, + Ended { + status: CodeCellRuntimeStatus, + }, +} + +impl TraceReducer { + /// Starts a code cell once its model-visible source item exists. + /// + /// Runtime events are allowed to arrive before stream completion has + /// reduced the model output that requested `exec`. Queueing preserves the + /// event order while still requiring every final `CodeCell` to point at the + /// exact conversation item that authored its JavaScript. + pub(super) fn start_or_queue_code_cell(&mut self, pending: PendingCodeCellStart) -> Result<()> { + let code_cell_id = pending.started.code_cell_id.clone(); + if self + .source_item_id_for_pending_code_cell(&pending)? + .is_none() + { + if self.rollout.code_cells.contains_key(&code_cell_id) + || self.pending_code_cell_starts.contains_key(&code_cell_id) + { + bail!("duplicate code cell start for {code_cell_id}"); + } + self.pending_code_cell_starts.insert(code_cell_id, pending); + return Ok(()); + } + + self.start_code_cell(pending) + } + + /// Materializes any queued code-cell starts unlocked by newly reduced conversation items. + /// + /// This is called after inference and compaction conversation reduction, + /// because those are the only paths that create model-visible items today. + pub(super) fn flush_pending_code_cell_starts(&mut self) -> Result<()> { + let mut ready_ids = Vec::new(); + for (code_cell_id, pending) in &self.pending_code_cell_starts { + if self + .source_item_id_for_pending_code_cell(pending)? + .is_some() + { + ready_ids.push(code_cell_id.clone()); + } + } + + for code_cell_id in ready_ids { + let Some(pending) = self.pending_code_cell_starts.remove(&code_cell_id) else { + continue; + }; + self.start_code_cell(pending)?; + } + Ok(()) + } + + /// Inserts the reduced `CodeCell` once source ownership can be proven. + fn start_code_cell(&mut self, pending: PendingCodeCellStart) -> Result<()> { + let PendingCodeCellStart { + seq, + wall_time_unix_ms, + thread_id, + codex_turn_id, + started, + } = pending; + if self.rollout.code_cells.contains_key(&started.code_cell_id) { + bail!("duplicate code cell start for {}", started.code_cell_id); + } + + let Some(codex_turn_id) = codex_turn_id else { + bail!( + "code cell start {} did not include a Codex turn id", + started.code_cell_id + ); + }; + self.validate_code_cell_turn(&thread_id, &codex_turn_id)?; + + let source_item_id = self.source_item_id_for_code_cell_start( + &thread_id, + &started.code_cell_id, + &started.model_visible_call_id, + )?; + let output_item_ids = self.model_visible_code_cell_item_ids( + &thread_id, + &started.model_visible_call_id, + ConversationItemKind::CustomToolCallOutput, + ); + // Runtime events may also have arrived while the start was queued. + // Seed these reverse links from already-reduced tool calls so replay is + // order-insensitive within the known trace causality. + let requester = ToolCallRequester::CodeCell { + code_cell_id: started.code_cell_id.clone(), + }; + let nested_tool_call_ids = self + .rollout + .tool_calls + .values() + .filter(|tool_call| tool_call.requester == requester) + .map(|tool_call| tool_call.tool_call_id.clone()) + .collect(); + + self.rollout.code_cells.insert( + started.code_cell_id.clone(), + CodeCell { + code_cell_id: started.code_cell_id.clone(), + model_visible_call_id: started.model_visible_call_id, + thread_id: thread_id.clone(), + codex_turn_id, + source_item_id, + output_item_ids: output_item_ids.clone(), + runtime_cell_id: Some(started.runtime_cell_id), + execution: ExecutionWindow { + started_at_unix_ms: wall_time_unix_ms, + started_seq: seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + runtime_status: CodeCellRuntimeStatus::Starting, + initial_response_at_unix_ms: None, + initial_response_seq: None, + yielded_at_unix_ms: None, + yielded_seq: None, + source_js: started.source_js, + nested_tool_call_ids, + wait_tool_call_ids: Vec::new(), + }, + ); + + self.thread_mut(&thread_id)?; + + for item_id in output_item_ids { + self.add_code_cell_output_item(&started.code_cell_id, &item_id)?; + } + self.flush_pending_code_cell_lifecycle_events(&started.code_cell_id)?; + + Ok(()) + } + + /// Returns the source item if the model-visible `exec` call has been reduced. + fn source_item_id_for_pending_code_cell( + &self, + pending: &PendingCodeCellStart, + ) -> Result> { + Ok(self + .model_visible_code_cell_item_ids( + &pending.thread_id, + &pending.started.model_visible_call_id, + ConversationItemKind::CustomToolCall, + ) + .into_iter() + .next()) + } + + /// Records the runtime's first response for a code cell, or waits for its source item. + /// + /// Code-mode execution can start and fail before the inference response payload + /// that introduced the model-visible `exec` call has been reduced. In that + /// case the cell start is already pending; keep the lifecycle event beside it + /// instead of weakening the invariant that every reduced cell has a source + /// conversation item. + pub(super) fn record_or_queue_code_cell_initial_response( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + code_cell_id: CodeCellId, + runtime_cell_id: String, + status: CodeCellRuntimeStatus, + ) -> Result<()> { + if !self.rollout.code_cells.contains_key(&code_cell_id) { + if self.pending_code_cell_starts.contains_key(&code_cell_id) { + self.queue_code_cell_lifecycle_event( + code_cell_id, + PendingCodeCellLifecycleEvent { + seq, + wall_time_unix_ms, + kind: PendingCodeCellLifecycleEventKind::InitialResponse { + runtime_cell_id, + status, + }, + }, + ); + return Ok(()); + } + bail!("code cell initial response referenced unknown cell {code_cell_id}"); + } + self.record_code_cell_initial_response( + seq, + wall_time_unix_ms, + code_cell_id, + runtime_cell_id, + status, + ) + } + + fn record_code_cell_initial_response( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + code_cell_id: CodeCellId, + runtime_cell_id: String, + status: CodeCellRuntimeStatus, + ) -> Result<()> { + let Some(cell) = self.rollout.code_cells.get_mut(&code_cell_id) else { + bail!("code cell initial response referenced unknown cell {code_cell_id}"); + }; + + cell.runtime_cell_id = Some(runtime_cell_id); + if cell.initial_response_at_unix_ms.is_none() { + cell.initial_response_at_unix_ms = Some(wall_time_unix_ms); + cell.initial_response_seq = Some(seq); + } + if status == CodeCellRuntimeStatus::Yielded { + cell.yielded_at_unix_ms = Some(wall_time_unix_ms); + cell.yielded_seq = Some(seq); + } + cell.runtime_status = status; + Ok(()) + } + + /// Ends a code cell, or waits until its queued start can materialize. + /// + /// This mirrors `record_or_queue_code_cell_initial_response`: the reducer is + /// strict about unknown cells, but a cell whose start is pending on the + /// model-visible source item is known and just needs its lifecycle replayed + /// after the source item appears. + pub(super) fn end_or_queue_code_cell( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + code_cell_id: CodeCellId, + status: CodeCellRuntimeStatus, + ) -> Result<()> { + if !self.rollout.code_cells.contains_key(&code_cell_id) { + if self.pending_code_cell_starts.contains_key(&code_cell_id) { + self.queue_code_cell_lifecycle_event( + code_cell_id, + PendingCodeCellLifecycleEvent { + seq, + wall_time_unix_ms, + kind: PendingCodeCellLifecycleEventKind::Ended { status }, + }, + ); + return Ok(()); + } + bail!("code cell end referenced unknown cell {code_cell_id}"); + } + self.end_code_cell(seq, wall_time_unix_ms, code_cell_id, status) + } + + fn end_code_cell( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + code_cell_id: CodeCellId, + status: CodeCellRuntimeStatus, + ) -> Result<()> { + let Some(cell) = self.rollout.code_cells.get_mut(&code_cell_id) else { + bail!("code cell end referenced unknown cell {code_cell_id}"); + }; + + if cell.initial_response_at_unix_ms.is_none() { + cell.initial_response_at_unix_ms = Some(wall_time_unix_ms); + cell.initial_response_seq = Some(seq); + } + cell.execution.ended_at_unix_ms = Some(wall_time_unix_ms); + cell.execution.ended_seq = Some(seq); + cell.execution.status = execution_status_for_code_cell(&status); + cell.runtime_status = status; + Ok(()) + } + + /// Closes unfinished code cells when their owning turn is interrupted. + /// + /// A yielded code cell can outlive a completed turn and be resumed by a + /// later `wait`, so normal turn completion must not imply cell completion. + /// Cancellation/failure is different: the model-visible JS frame has been + /// abandoned even if nested terminal work reports late runtime events. In + /// that case leaving the cell `running` makes a completed trace look live. + pub(super) fn terminate_running_code_cells_for_turn_end( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + codex_turn_id: &str, + turn_status: &ExecutionStatus, + ) -> Result<()> { + let runtime_status = match turn_status { + ExecutionStatus::Running | ExecutionStatus::Completed => return Ok(()), + ExecutionStatus::Failed => CodeCellRuntimeStatus::Failed, + ExecutionStatus::Cancelled | ExecutionStatus::Aborted => { + CodeCellRuntimeStatus::Terminated + } + }; + let code_cell_ids: Vec<_> = self + .rollout + .code_cells + .values() + .filter(|cell| { + cell.codex_turn_id == codex_turn_id + && cell.execution.status == ExecutionStatus::Running + }) + .map(|cell| cell.code_cell_id.clone()) + .collect(); + + for code_cell_id in code_cell_ids { + self.end_code_cell(seq, wall_time_unix_ms, code_cell_id, runtime_status.clone())?; + } + Ok(()) + } + + fn queue_code_cell_lifecycle_event( + &mut self, + code_cell_id: CodeCellId, + event: PendingCodeCellLifecycleEvent, + ) { + let events = self + .pending_code_cell_lifecycle_events + .entry(code_cell_id) + .or_default(); + events.push(event); + events.sort_by_key(|event| event.seq); + } + + fn flush_pending_code_cell_lifecycle_events(&mut self, code_cell_id: &str) -> Result<()> { + let Some(events) = self.pending_code_cell_lifecycle_events.remove(code_cell_id) else { + return Ok(()); + }; + for event in events { + match event.kind { + PendingCodeCellLifecycleEventKind::InitialResponse { + runtime_cell_id, + status, + } => self.record_code_cell_initial_response( + event.seq, + event.wall_time_unix_ms, + code_cell_id.to_string(), + runtime_cell_id, + status, + )?, + PendingCodeCellLifecycleEventKind::Ended { status } => self.end_code_cell( + event.seq, + event.wall_time_unix_ms, + code_cell_id.to_string(), + status, + )?, + } + } + Ok(()) + } + + /// Links a nested tool call back to its parent code cell. + /// + /// If the parent cell is still queued, the link is recovered later from already + /// reduced tool calls when the cell materializes. + pub(super) fn link_tool_call_to_code_cell( + &mut self, + tool_call_id: &ToolCallId, + requester: &ToolCallRequester, + ) -> Result<()> { + let ToolCallRequester::CodeCell { code_cell_id } = requester else { + return Ok(()); + }; + let Some(cell) = self.rollout.code_cells.get_mut(code_cell_id) else { + // The cell start may still be queued behind the inference payload + // that contains its model-visible source item. `start_code_cell` + // backfills these already-reduced nested calls once the source + // ownership can be proven. + return Ok(()); + }; + push_unique(&mut cell.nested_tool_call_ids, tool_call_id); + Ok(()) + } + + /// Records that a model-visible wait call is waiting on a runtime code cell. + /// + /// Wait calls are not nested JavaScript tools, so the relationship is inferred + /// from the runtime cell id inside the function arguments. + pub(super) fn link_wait_tool_call_from_request_payload( + &mut self, + thread_id: &str, + tool_call_id: &ToolCallId, + request_payload: Option<&RawPayloadRef>, + ) -> Result<()> { + let Some(request_payload) = request_payload else { + return Ok(()); + }; + let payload = self.read_payload_json(request_payload)?; + if payload.get("tool_name").and_then(Value::as_str) != Some("wait") { + return Ok(()); + } + // `wait` is a normal model-visible function call, not a nested JS tool + // request. The only stable edge back to the code cell is the runtime + // `cell_id` inside the function arguments. + let Some(arguments) = payload + .get("payload") + .and_then(|payload| payload.get("arguments")) + .and_then(Value::as_str) + else { + bail!( + "wait tool request payload {} did not contain function arguments", + request_payload.raw_payload_id + ); + }; + let arguments: Value = serde_json::from_str(arguments).with_context(|| { + format!( + "wait tool request payload {} had invalid JSON arguments", + request_payload.raw_payload_id + ) + })?; + let Some(runtime_cell_id) = arguments.get("cell_id").and_then(Value::as_str) else { + bail!( + "wait tool request payload {} did not contain cell_id", + request_payload.raw_payload_id + ); + }; + let Some(code_cell_id) = + self.code_cell_id_for_runtime_cell_id_if_known(thread_id, runtime_cell_id) + else { + return Ok(()); + }; + let Some(cell) = self.rollout.code_cells.get_mut(&code_cell_id) else { + return Ok(()); + }; + push_unique(&mut cell.wait_tool_call_ids, tool_call_id); + Ok(()) + } + + /// Attaches a later-observed model-visible output item to its code cell. + /// + /// This is used when an inference request carries a custom-tool output after + /// the runtime cell already exists. + pub(super) fn attach_model_visible_code_cell_item( + &mut self, + item_id: &str, + call_id: Option<&str>, + kind: &ConversationItemKind, + ) -> Result<()> { + let Some(call_id) = call_id else { + return Ok(()); + }; + if *kind != ConversationItemKind::CustomToolCallOutput { + return Ok(()); + } + // The output item can be observed after the CodeCell was created, e.g. + // when a later inference request carries the custom-tool result back to + // the model. Add the reverse ProducerRef at that later observation + // point instead of copying runtime bytes into the conversation model. + let code_cell_id = self.reduced_code_cell_id_for_model_visible_call(call_id); + if !self.rollout.code_cells.contains_key(&code_cell_id) { + return Ok(()); + } + self.add_code_cell_output_item(&code_cell_id, item_id) + } + + /// Resolves the owning thread for a code-cell runtime event. + /// + /// Runtime events should carry a thread id, but older/raw paths may only have + /// the turn id. The fallback keeps replay strict while avoiding duplicate logic + /// in every code-cell event arm. + pub(super) fn code_cell_event_thread_id( + &self, + thread_id: Option, + codex_turn_id: Option<&str>, + runtime_cell_id: &str, + event_name: &str, + ) -> Result { + if let Some(thread_id) = thread_id { + return Ok(thread_id); + } + let Some(codex_turn_id) = codex_turn_id else { + bail!("{event_name} {runtime_cell_id} did not include a thread id"); + }; + self.rollout + .codex_turns + .get(codex_turn_id) + .map(|turn| turn.thread_id.clone()) + .with_context(|| { + format!( + "{event_name} {runtime_cell_id} referenced unknown Codex turn {codex_turn_id}" + ) + }) + } + + /// Derives the stable reduced code-cell id from the model-visible exec call id. + pub(super) fn reduced_code_cell_id_for_model_visible_call( + &self, + model_visible_call_id: &str, + ) -> CodeCellId { + // The model-visible `exec` call is the durable source identity. The + // runtime `cell_id` is only a thread-local handle used for later waits + // and nested tool calls. + format!("code_cell:{model_visible_call_id}") + } + + /// Records the thread-local runtime cell id to reduced code-cell id mapping. + /// + /// Runtime ids can repeat across threads, so callers must provide the owning + /// thread id when creating or resolving this bridge. + pub(super) fn record_runtime_code_cell_id( + &mut self, + thread_id: &str, + runtime_cell_id: &str, + code_cell_id: &str, + ) -> Result<()> { + let key = runtime_code_cell_key(thread_id, runtime_cell_id); + if let Some(existing) = self.code_cell_ids_by_runtime.get(&key) { + if existing == code_cell_id { + return Ok(()); + } + bail!( + "runtime code cell {runtime_cell_id} in thread {thread_id} mapped to both \ + {existing} and {code_cell_id}" + ); + } + self.code_cell_ids_by_runtime + .insert(key, code_cell_id.to_string()); + Ok(()) + } + + /// Resolves a runtime cell id to the reduced code-cell id for the given thread. + pub(super) fn code_cell_id_for_runtime_cell_id( + &self, + thread_id: &str, + runtime_cell_id: &str, + event_name: &str, + ) -> Result { + self.code_cell_id_for_runtime_cell_id_if_known(thread_id, runtime_cell_id) + .with_context(|| { + format!( + "{event_name} referenced unknown runtime cell {runtime_cell_id} \ + in thread {thread_id}" + ) + }) + } + + fn code_cell_id_for_runtime_cell_id_if_known( + &self, + thread_id: &str, + runtime_cell_id: &str, + ) -> Option { + self.code_cell_ids_by_runtime + .get(&runtime_code_cell_key(thread_id, runtime_cell_id)) + .cloned() + } + + /// Converts a raw tool requester into the reduced graph requester. + /// + /// Code-mode tool requests arrive with a runtime cell id, so this method is + /// the boundary that turns that runtime handle into a stable code-cell anchor. + pub(super) fn reduce_tool_call_requester( + &self, + thread_id: &str, + requester: RawToolCallRequester, + ) -> Result { + match requester { + RawToolCallRequester::Model => Ok(ToolCallRequester::Model), + RawToolCallRequester::CodeCell { runtime_cell_id } => Ok(ToolCallRequester::CodeCell { + code_cell_id: self.code_cell_id_for_runtime_cell_id( + thread_id, + &runtime_cell_id, + "code-mode nested tool", + )?, + }), + } + } + + fn validate_code_cell_turn(&self, thread_id: &str, codex_turn_id: &str) -> Result<()> { + if !self.rollout.threads.contains_key(thread_id) { + bail!("code cell start referenced unknown thread {thread_id}"); + } + let Some(turn) = self.rollout.codex_turns.get(codex_turn_id) else { + bail!("code cell start referenced unknown Codex turn {codex_turn_id}"); + }; + if turn.thread_id != thread_id { + bail!( + "code cell start used thread {thread_id}, but Codex turn {codex_turn_id} belongs \ + to {}", + turn.thread_id + ); + } + Ok(()) + } + + fn model_visible_code_cell_item_ids( + &self, + thread_id: &str, + call_id: &str, + kind: ConversationItemKind, + ) -> Vec { + self.rollout + .conversation_items + .values() + .filter(|item| { + item.thread_id == thread_id + && item.call_id.as_deref() == Some(call_id) + && item.kind == kind + }) + .map(|item| item.item_id.clone()) + .collect() + } + + fn source_item_id_for_code_cell_start( + &self, + thread_id: &str, + code_cell_id: &str, + model_visible_call_id: &str, + ) -> Result { + self.model_visible_code_cell_item_ids( + thread_id, + model_visible_call_id, + ConversationItemKind::CustomToolCall, + ) + .into_iter() + .next() + .with_context(|| { + format!( + "code cell {code_cell_id} referenced model-visible call {model_visible_call_id}, \ + but no custom tool call item was observed" + ) + }) + } + + fn add_code_cell_output_item(&mut self, code_cell_id: &str, item_id: &str) -> Result<()> { + let Some(cell) = self.rollout.code_cells.get_mut(code_cell_id) else { + bail!("code cell {code_cell_id} disappeared during output linking"); + }; + push_unique(&mut cell.output_item_ids, item_id); + + let Some(item) = self.rollout.conversation_items.get_mut(item_id) else { + bail!("conversation item {item_id} disappeared during code-cell output linking"); + }; + let producer = ProducerRef::CodeCell { + code_cell_id: code_cell_id.to_string(), + }; + if !item.produced_by.contains(&producer) { + item.produced_by.push(producer); + } + Ok(()) + } +} + +fn execution_status_for_code_cell(status: &CodeCellRuntimeStatus) -> ExecutionStatus { + match status { + CodeCellRuntimeStatus::Starting + | CodeCellRuntimeStatus::Running + | CodeCellRuntimeStatus::Yielded => ExecutionStatus::Running, + CodeCellRuntimeStatus::Completed => ExecutionStatus::Completed, + CodeCellRuntimeStatus::Failed => ExecutionStatus::Failed, + CodeCellRuntimeStatus::Terminated => ExecutionStatus::Cancelled, + } +} + +fn push_unique(items: &mut Vec, item_id: &str) { + if !items.iter().any(|existing| existing == item_id) { + items.push(item_id.to_string()); + } +} + +fn runtime_code_cell_key(thread_id: &str, runtime_cell_id: &str) -> (String, String) { + (thread_id.to_string(), runtime_cell_id.to_string()) +} + +#[cfg(test)] +#[path = "code_cell_tests.rs"] +mod tests; diff --git a/codex-rs/rollout-trace/src/reducer/code_cell_tests.rs b/codex-rs/rollout-trace/src/reducer/code_cell_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..1aabdb8438a5aba20244e19d8c972074dffeaac8 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/code_cell_tests.rs @@ -0,0 +1,427 @@ +use pretty_assertions::assert_eq; +use serde_json::json; +use tempfile::TempDir; + +use crate::model::CodeCellRuntimeStatus; +use crate::model::ConversationItemKind; +use crate::model::ExecutionStatus; +use crate::model::ProducerRef; +use crate::model::ToolCallKind; +use crate::model::ToolCallSummary; +use crate::payload::RawPayloadKind; +use crate::raw_event::RawToolCallRequester; +use crate::raw_event::RawTraceEventPayload; +use crate::reducer::test_support::create_started_writer; +use crate::reducer::test_support::message; +use crate::reducer::test_support::start_turn; +use crate::reducer::test_support::start_turn_for_thread; +use crate::reducer::test_support::trace_context; +use crate::reducer::test_support::trace_context_for_thread; +use crate::replay_bundle; + +#[test] +fn code_cell_lifecycle_links_nested_tools_waits_and_outputs() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "count files")] + }), + )?; + writer.append(RawTraceEventPayload::InferenceStarted { + inference_call_id: "inference-1".to_string(), + thread_id: "thread-root".to_string(), + codex_turn_id: "turn-1".to_string(), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + request_payload: request, + })?; + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [{ + "type": "custom_tool_call", + "name": "exec", + "call_id": "call-code", + "input": "text('hi')" + }] + }), + )?; + // Runtime tool dispatch starts before the stream-completion hook has + // reduced the model response that requested `exec`. + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CodeCellStarted { + runtime_cell_id: "1".to_string(), + model_visible_call_id: "call-code".to_string(), + source_js: "text('hi')".to_string(), + }, + )?; + writer.append(RawTraceEventPayload::InferenceCompleted { + inference_call_id: "inference-1".to_string(), + response_id: Some("resp-1".to_string()), + upstream_request_id: None, + response_payload: response, + })?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CodeCellInitialResponse { + runtime_cell_id: "1".to_string(), + status: CodeCellRuntimeStatus::Yielded, + response_payload: None, + }, + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "nested-tool-1".to_string(), + model_visible_call_id: None, + code_mode_runtime_tool_id: Some("tool-1".to_string()), + requester: RawToolCallRequester::CodeCell { + runtime_cell_id: "1".to_string(), + }, + kind: ToolCallKind::ExecCommand, + summary: ToolCallSummary::Generic { + label: "exec_command".to_string(), + input_preview: Some("pwd".to_string()), + output_preview: None, + }, + invocation_payload: None, + }, + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallEnded { + tool_call_id: "nested-tool-1".to_string(), + status: ExecutionStatus::Completed, + result_payload: None, + }, + )?; + + start_turn(&writer, "turn-2")?; + let followup = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "previous_response_id": "resp-1", + "input": [{ + "type": "custom_tool_call_output", + "call_id": "call-code", + "output": "Script running with cell ID 1" + }] + }), + )?; + writer.append(RawTraceEventPayload::InferenceStarted { + inference_call_id: "inference-2".to_string(), + thread_id: "thread-root".to_string(), + codex_turn_id: "turn-2".to_string(), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + request_payload: followup, + })?; + let wait_request = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "wait", + "tool_namespace": null, + "payload": { + "type": "function", + "arguments": "{\"cell_id\":\"1\"}" + } + }), + )?; + writer.append_with_context( + trace_context("turn-2"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "wait-tool-1".to_string(), + model_visible_call_id: Some("wait-call".to_string()), + code_mode_runtime_tool_id: None, + requester: RawToolCallRequester::Model, + kind: ToolCallKind::Other { + name: "wait".to_string(), + }, + summary: ToolCallSummary::Generic { + label: "wait".to_string(), + input_preview: Some("{\"cell_id\":\"1\"}".to_string()), + output_preview: None, + }, + invocation_payload: Some(wait_request), + }, + )?; + writer.append_with_context( + trace_context("turn-2"), + RawTraceEventPayload::CodeCellEnded { + runtime_cell_id: "1".to_string(), + status: CodeCellRuntimeStatus::Completed, + response_payload: None, + }, + )?; + + let rollout = replay_bundle(temp.path())?; + let code_cell_id = test_reduced_code_cell_id("call-code"); + let cell = &rollout.code_cells[&code_cell_id]; + let output_item_id = rollout.inference_calls["inference-2"] + .request_item_ids + .last() + .expect("exec output item"); + + assert_eq!(cell.thread_id, "thread-root"); + assert_eq!(cell.runtime_status, CodeCellRuntimeStatus::Completed); + assert_eq!(cell.execution.status, ExecutionStatus::Completed); + assert_eq!(cell.runtime_cell_id, Some("1".to_string())); + assert_eq!(cell.nested_tool_call_ids, vec!["nested-tool-1"]); + assert_eq!(cell.wait_tool_call_ids, vec!["wait-tool-1"]); + assert_eq!(cell.output_item_ids, vec![output_item_id.clone()]); + assert_eq!( + rollout.conversation_items[output_item_id].produced_by, + vec![ProducerRef::CodeCell { + code_cell_id: code_cell_id.clone(), + }] + ); + assert_eq!( + rollout.conversation_items[&cell.source_item_id].kind, + ConversationItemKind::CustomToolCall, + ); + + Ok(()) +} + +#[test] +fn fast_code_cell_lifecycle_waits_for_source_item() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "count files")] + }), + )?; + writer.append(RawTraceEventPayload::InferenceStarted { + inference_call_id: "inference-1".to_string(), + thread_id: "thread-root".to_string(), + codex_turn_id: "turn-1".to_string(), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + request_payload: request, + })?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CodeCellStarted { + runtime_cell_id: "1".to_string(), + model_visible_call_id: "call-code".to_string(), + source_js: "not valid js".to_string(), + }, + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CodeCellInitialResponse { + runtime_cell_id: "1".to_string(), + status: CodeCellRuntimeStatus::Failed, + response_payload: None, + }, + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CodeCellEnded { + runtime_cell_id: "1".to_string(), + status: CodeCellRuntimeStatus::Failed, + response_payload: None, + }, + )?; + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [{ + "type": "custom_tool_call", + "name": "exec", + "call_id": "call-code", + "input": "not valid js" + }] + }), + )?; + writer.append(RawTraceEventPayload::InferenceCompleted { + inference_call_id: "inference-1".to_string(), + response_id: Some("resp-1".to_string()), + upstream_request_id: None, + response_payload: response, + })?; + + let rollout = replay_bundle(temp.path())?; + let code_cell_id = test_reduced_code_cell_id("call-code"); + let cell = &rollout.code_cells[&code_cell_id]; + + assert_eq!(cell.thread_id, "thread-root"); + assert_eq!(cell.runtime_status, CodeCellRuntimeStatus::Failed); + assert_eq!(cell.execution.status, ExecutionStatus::Failed); + assert_eq!(cell.runtime_cell_id, Some("1".to_string())); + assert_eq!( + rollout.conversation_items[&cell.source_item_id].kind, + ConversationItemKind::CustomToolCall, + ); + + Ok(()) +} + +#[test] +fn cancelled_turn_terminates_unfinished_code_cell() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "count files")] + }), + )?; + writer.append(RawTraceEventPayload::InferenceStarted { + inference_call_id: "inference-1".to_string(), + thread_id: "thread-root".to_string(), + codex_turn_id: "turn-1".to_string(), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + request_payload: request, + })?; + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [{ + "type": "custom_tool_call", + "name": "exec", + "call_id": "call-code", + "input": "await tools.exec_command({cmd: 'slow'});" + }] + }), + )?; + writer.append(RawTraceEventPayload::InferenceCompleted { + inference_call_id: "inference-1".to_string(), + response_id: Some("resp-1".to_string()), + upstream_request_id: None, + response_payload: response, + })?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CodeCellStarted { + runtime_cell_id: "1".to_string(), + model_visible_call_id: "call-code".to_string(), + source_js: "await tools.exec_command({cmd: 'slow'});".to_string(), + }, + )?; + let turn_end = writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CodexTurnEnded { + codex_turn_id: "turn-1".to_string(), + status: ExecutionStatus::Cancelled, + }, + )?; + + let rollout = replay_bundle(temp.path())?; + let code_cell_id = test_reduced_code_cell_id("call-code"); + let cell = &rollout.code_cells[&code_cell_id]; + + assert_eq!(cell.runtime_status, CodeCellRuntimeStatus::Terminated); + assert_eq!(cell.execution.status, ExecutionStatus::Cancelled); + assert_eq!(cell.execution.ended_seq, Some(turn_end.seq)); + + Ok(()) +} + +#[test] +fn runtime_code_cell_ids_can_repeat_across_threads() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + writer.append(RawTraceEventPayload::ThreadStarted { + thread_id: "thread-child".to_string(), + agent_path: "/root/child".to_string(), + metadata_payload: None, + })?; + start_turn_for_thread(&writer, "thread-root", "turn-root")?; + start_turn_for_thread(&writer, "thread-child", "turn-child")?; + + for (thread_id, turn_id, inference_call_id, call_id) in [ + ("thread-root", "turn-root", "inference-root", "call-root"), + ( + "thread-child", + "turn-child", + "inference-child", + "call-child", + ), + ] { + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "run code")] + }), + )?; + writer.append(RawTraceEventPayload::InferenceStarted { + inference_call_id: inference_call_id.to_string(), + thread_id: thread_id.to_string(), + codex_turn_id: turn_id.to_string(), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + request_payload: request, + })?; + writer.append_with_context( + trace_context_for_thread(thread_id, turn_id), + RawTraceEventPayload::CodeCellStarted { + runtime_cell_id: "1".to_string(), + model_visible_call_id: call_id.to_string(), + source_js: "text('hi')".to_string(), + }, + )?; + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": format!("resp-{thread_id}"), + "output_items": [{ + "type": "custom_tool_call", + "name": "exec", + "call_id": call_id, + "input": "text('hi')" + }] + }), + )?; + writer.append(RawTraceEventPayload::InferenceCompleted { + inference_call_id: inference_call_id.to_string(), + response_id: Some(format!("resp-{thread_id}")), + upstream_request_id: None, + response_payload: response, + })?; + writer.append_with_context( + trace_context_for_thread(thread_id, turn_id), + RawTraceEventPayload::CodeCellEnded { + runtime_cell_id: "1".to_string(), + status: CodeCellRuntimeStatus::Completed, + response_payload: None, + }, + )?; + } + + let rollout = replay_bundle(temp.path())?; + let root_cell_id = test_reduced_code_cell_id("call-root"); + let child_cell_id = test_reduced_code_cell_id("call-child"); + + assert_eq!(rollout.code_cells[&root_cell_id].thread_id, "thread-root"); + assert_eq!(rollout.code_cells[&child_cell_id].thread_id, "thread-child"); + assert_eq!( + rollout.code_cells[&root_cell_id].runtime_cell_id, + Some("1".to_string()) + ); + assert_eq!( + rollout.code_cells[&child_cell_id].runtime_cell_id, + Some("1".to_string()) + ); + + Ok(()) +} + +fn test_reduced_code_cell_id(model_visible_call_id: &str) -> String { + format!("code_cell:{model_visible_call_id}") +} diff --git a/codex-rs/rollout-trace/src/reducer/compaction.rs b/codex-rs/rollout-trace/src/reducer/compaction.rs new file mode 100644 index 0000000000000000000000000000000000000000..f29b5dfa82820b61b08907cfed13c52c498948d4 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/compaction.rs @@ -0,0 +1,183 @@ +//! Reducer support for the remote compaction lifecycle. +//! +//! This module owns request/checkpoint bookkeeping. Conversation item reconciliation stays in +//! `conversation` because it depends on the same normalization and reuse invariants as inference +//! requests. + +use anyhow::Result; +use anyhow::bail; + +use super::TraceReducer; +use crate::model::Compaction; +use crate::model::CompactionRequest; +use crate::model::CompactionRequestId; +use crate::model::ExecutionStatus; +use crate::model::ExecutionWindow; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawEventSeq; + +impl TraceReducer { + /// Starts one upstream request attempt for a compaction operation. + pub(super) fn start_compaction_request( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + started: StartedCompactionRequest, + ) -> Result<()> { + if self + .rollout + .compaction_requests + .contains_key(&started.compaction_request_id) + { + bail!( + "duplicate compaction request start for {}", + started.compaction_request_id + ); + } + self.thread_mut(&started.thread_id)?; + let Some(turn) = self.rollout.codex_turns.get(&started.codex_turn_id) else { + bail!( + "compaction request {} referenced unknown codex turn {}", + started.compaction_request_id, + started.codex_turn_id + ); + }; + if turn.thread_id != started.thread_id { + bail!( + "compaction request {} used thread {}, but codex turn {} belongs to {}", + started.compaction_request_id, + started.thread_id, + started.codex_turn_id, + turn.thread_id + ); + } + + self.rollout.compaction_requests.insert( + started.compaction_request_id.clone(), + CompactionRequest { + compaction_request_id: started.compaction_request_id, + compaction_id: started.compaction_id, + thread_id: started.thread_id, + codex_turn_id: started.codex_turn_id, + execution: ExecutionWindow { + started_at_unix_ms: wall_time_unix_ms, + started_seq: seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + model: started.model, + provider_name: started.provider_name, + raw_request_payload_id: started.request_payload.raw_payload_id, + raw_response_payload_id: None, + }, + ); + Ok(()) + } + + /// Completes an upstream compaction request attempt without modifying conversation history. + /// + /// The request/response payloads are evidence for the remote call. The live + /// conversation changes only when a separate install event provides the checkpoint. + pub(super) fn complete_compaction_request( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + compaction_id: String, + compaction_request_id: CompactionRequestId, + status: ExecutionStatus, + response_payload: Option, + ) -> Result<()> { + let Some(request) = self + .rollout + .compaction_requests + .get_mut(&compaction_request_id) + else { + bail!( + "compaction request completion referenced unknown request {compaction_request_id}" + ); + }; + if request.compaction_id != compaction_id { + bail!( + "compaction request {compaction_request_id} completion used compaction {compaction_id}, but start used {}", + request.compaction_id + ); + } + request.execution.ended_at_unix_ms = Some(wall_time_unix_ms); + request.execution.ended_seq = Some(seq); + request.execution.status = status; + request.raw_response_payload_id = response_payload.map(|payload| payload.raw_payload_id); + Ok(()) + } + + /// Installs a compaction checkpoint into the reduced conversation graph. + /// + /// This is the semantic boundary where replacement history becomes the live + /// thread history; request attempts alone do not imply that change. + pub(super) fn reduce_compaction_installed_event( + &mut self, + wall_time_unix_ms: i64, + thread_id: String, + codex_turn_id: String, + compaction_id: String, + checkpoint_payload: RawPayloadRef, + ) -> Result<()> { + if self.rollout.compactions.contains_key(&compaction_id) { + bail!("duplicate compaction install for {compaction_id}"); + } + self.thread_mut(&thread_id)?; + let Some(turn) = self.rollout.codex_turns.get(&codex_turn_id) else { + bail!( + "compaction install {compaction_id} referenced unknown codex turn {codex_turn_id}" + ); + }; + if turn.thread_id != thread_id { + bail!( + "compaction install {compaction_id} used thread {thread_id}, but codex turn {codex_turn_id} belongs to {}", + turn.thread_id + ); + } + let checkpoint = self.reduce_compaction_checkpoint( + wall_time_unix_ms, + &thread_id, + codex_turn_id.as_str(), + &compaction_id, + &checkpoint_payload, + )?; + let request_ids = self + .rollout + .compaction_requests + .values() + .filter(|request| request.compaction_id == compaction_id) + .map(|request| request.compaction_request_id.clone()) + .collect(); + + self.pending_compaction_replacement_item_ids + .insert(thread_id.clone(), checkpoint.replacement_item_ids.clone()); + self.rollout.compactions.insert( + compaction_id.clone(), + Compaction { + compaction_id, + thread_id, + codex_turn_id, + installed_at_unix_ms: wall_time_unix_ms, + marker_item_id: checkpoint.marker_item_id, + request_ids, + input_item_ids: checkpoint.input_item_ids, + replacement_item_ids: checkpoint.replacement_item_ids, + }, + ); + Ok(()) + } +} + +/// Raw compaction-request start fields after dispatch has stripped the event envelope. +pub(super) struct StartedCompactionRequest { + pub(super) compaction_id: String, + pub(super) compaction_request_id: String, + pub(super) thread_id: String, + pub(super) codex_turn_id: String, + pub(super) model: String, + pub(super) provider_name: String, + pub(super) request_payload: RawPayloadRef, +} diff --git a/codex-rs/rollout-trace/src/reducer/conversation.rs b/codex-rs/rollout-trace/src/reducer/conversation.rs new file mode 100644 index 0000000000000000000000000000000000000000..f105ffebd6d6d1752584435ece9a8a844efd6749 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/conversation.rs @@ -0,0 +1,708 @@ +//! Conversation reduction from model-facing payload snapshots. +//! +//! Inference request inputs and response outputs are both part of the logical +//! conversation because they are the payloads exchanged with the model. Runtime +//! observations, such as local tool output, stay outside the transcript until a +//! later model-facing payload carries their content. + +use anyhow::Context; +use anyhow::Ok; +use anyhow::Result; +use anyhow::bail; +use serde_json::Value; + +use self::normalize::NormalizedConversationItem; +use super::TraceReducer; +use crate::model::CompactionId; +use crate::model::ConversationBody; +use crate::model::ConversationItem; +use crate::model::ConversationItemKind; +use crate::model::ConversationPart; +use crate::model::ConversationRole; +use crate::model::InferenceCallId; +use crate::model::ProducerRef; +use crate::payload::RawPayloadRef; + +mod normalize; + +impl TraceReducer { + /// Reduces an inference request input snapshot into model-visible conversation items. + /// + /// Request snapshots are reconciled by position against the previous model-visible + /// snapshot for the thread so repeated history reuses ids while newly inserted + /// items remain distinct. + pub(super) fn reduce_inference_request( + &mut self, + wall_time_unix_ms: i64, + inference_call_id: &InferenceCallId, + thread_id: &str, + codex_turn_id: &str, + request_payload: &RawPayloadRef, + ) -> Result> { + let payload = self.read_payload_json(request_payload)?; + let Some(input) = payload.get("input") else { + bail!( + "inference request payload {} did not contain input", + request_payload.raw_payload_id + ); + }; + let Some(request_items) = input.as_array() else { + bail!( + "inference request payload {} had non-array input", + request_payload.raw_payload_id + ); + }; + + let items = normalize::normalize_model_items(request_items, request_payload)?; + + let previous_response_id = payload.get("previous_response_id").and_then(Value::as_str); + // After compaction, the next full request is compared against the installed replacement + // history, not the pre-compaction prompt. Any repeated developer/context prefix that Codex + // reinjects must therefore become a fresh post-compaction conversation item. + let post_compaction_snapshot = if previous_response_id.is_none() { + self.pending_compaction_replacement_item_ids + .get(thread_id) + .cloned() + } else { + None + }; + let request_item_ids = if let Some(previous_response_id) = previous_response_id { + // Streaming follow-up requests can send only the new input plus a + // `previous_response_id`. The trace model still exposes the full + // model-visible input, so rebuild the omitted prefix from the + // previous request and response before reducing this delta. + let previous_items = self + .rollout + .inference_calls + .values() + .find(|inference| { + inference.thread_id == thread_id + && inference.response_id.as_deref() == Some(previous_response_id) + }) + .map(|inference| { + let mut ids = inference.request_item_ids.clone(); + ids.extend(inference.response_item_ids.clone()); + ids + }); + let Some(mut item_ids) = previous_items else { + bail!( + "incremental inference request {inference_call_id} referenced unknown previous_response_id {previous_response_id}" + ); + }; + let delta_item_ids = self.reconcile_conversation_items( + items, + ReconcileItems { + thread_id, + codex_turn_id, + wall_time_unix_ms, + produced_by: Vec::new(), + start_index: item_ids.len(), + mode: ReconcileMode::AppendOnly, + snapshot_override: None, + }, + )?; + item_ids.extend(delta_item_ids); + item_ids + } else { + self.reconcile_conversation_items( + items, + ReconcileItems { + thread_id, + codex_turn_id, + wall_time_unix_ms, + produced_by: Vec::new(), + start_index: 0, + mode: ReconcileMode::FullSnapshot, + snapshot_override: post_compaction_snapshot.as_deref(), + }, + )? + }; + + self.append_thread_conversation_items(thread_id, &request_item_ids)?; + if post_compaction_snapshot.is_some() { + self.pending_compaction_replacement_item_ids + .remove(thread_id); + } + self.thread_conversation_snapshots + .insert(thread_id.to_string(), request_item_ids.clone()); + Ok(request_item_ids) + } + + /// Reduces an inference response payload into conversation items produced by the call. + pub(super) fn reduce_inference_response( + &mut self, + wall_time_unix_ms: i64, + inference_call_id: &InferenceCallId, + response_payload: &RawPayloadRef, + ) -> Result> { + let payload = self.read_payload_json(response_payload)?; + let Some(output_items) = payload.get("output_items").and_then(Value::as_array) else { + bail!( + "inference response payload {} did not contain output_items", + response_payload.raw_payload_id + ); + }; + + let Some((thread_id, codex_turn_id)) = self + .rollout + .inference_calls + .get(inference_call_id) + .map(|inference| (inference.thread_id.clone(), inference.codex_turn_id.clone())) + else { + bail!("inference response referenced unknown call {inference_call_id}"); + }; + + let items = normalize::normalize_model_items(output_items, response_payload)?; + // Response output is appended immediately: it was produced by the model, + // so it is conversation even before a later request carries it forward. + let append_at = self + .thread_conversation_snapshots + .get(&thread_id) + .map_or(0, Vec::len); + let response_item_ids = self.reconcile_conversation_items( + items, + ReconcileItems { + thread_id: &thread_id, + codex_turn_id: &codex_turn_id, + wall_time_unix_ms, + produced_by: vec![ProducerRef::Inference { + inference_call_id: inference_call_id.clone(), + }], + start_index: append_at, + mode: ReconcileMode::AppendOnly, + snapshot_override: None, + }, + )?; + self.append_thread_conversation_items(&thread_id, &response_item_ids)?; + self.thread_conversation_snapshots + .entry(thread_id) + .or_default() + .extend(response_item_ids.clone()); + + if let Some(usage) = payload + .get("token_usage") + .and_then(normalize::token_usage_from_value) + && let Some(inference) = self.rollout.inference_calls.get_mut(inference_call_id) + { + inference.usage = Some(usage); + } + + Ok(response_item_ids) + } + + fn reconcile_conversation_items( + &mut self, + items: Vec, + context: ReconcileItems<'_>, + ) -> Result> { + let previous_snapshot = context.snapshot_override.map_or_else( + || { + self.thread_conversation_snapshots + .get(context.thread_id) + .cloned() + .unwrap_or_default() + }, + <[_]>::to_vec, + ); + let mut item_ids = Vec::with_capacity(items.len()); + + for (offset, item) in items.into_iter().enumerate() { + let index = context.start_index + offset; + let tool_link_item = item.clone(); + self.ensure_call_id_consistency(context.thread_id, &item)?; + let item_id = if let Some(previous_item_id) = previous_snapshot.get(index) { + if self.item_matches(previous_item_id, &item) { + previous_item_id.clone() + } else if matches!(context.mode, ReconcileMode::FullSnapshot) { + self.find_matching_snapshot_item(&previous_snapshot, &item_ids, &item) + .unwrap_or_else(|| { + self.create_conversation_item( + context.thread_id, + Some(context.codex_turn_id.to_string()), + context.wall_time_unix_ms, + item, + context.produced_by.clone(), + ) + }) + } else { + let codex_turn_id = context.codex_turn_id; + let thread_id = context.thread_id; + bail!( + "model conversation mismatch while reducing turn {codex_turn_id} for \ + thread {thread_id} at item index {index}: existing item \ + {previous_item_id} does not match the current model payload item" + ); + } + } else if matches!(context.mode, ReconcileMode::FullSnapshot) { + self.find_matching_snapshot_item(&previous_snapshot, &item_ids, &item) + .unwrap_or_else(|| { + self.create_conversation_item( + context.thread_id, + Some(context.codex_turn_id.to_string()), + context.wall_time_unix_ms, + item, + context.produced_by.clone(), + ) + }) + } else { + self.create_conversation_item( + context.thread_id, + Some(context.codex_turn_id.to_string()), + context.wall_time_unix_ms, + item, + context.produced_by.clone(), + ) + }; + self.update_conversation_item_from_sighting( + &item_id, + &tool_link_item, + &context.produced_by, + )?; + self.attach_model_visible_tool_item( + &item_id, + tool_link_item.call_id.as_deref(), + &tool_link_item.kind, + )?; + self.attach_model_visible_code_cell_item( + &item_id, + tool_link_item.call_id.as_deref(), + &tool_link_item.kind, + )?; + self.resolve_pending_agent_edges_for_item(&item_id)?; + item_ids.push(item_id); + } + + self.flush_pending_code_cell_starts()?; + Ok(item_ids) + } + + /// Reduces a compaction checkpoint payload into installed replacement history. + /// + /// The returned ids let the compaction reducer record both the boundary marker + /// and the snapshot that future full requests should reconcile against. + pub(super) fn reduce_compaction_checkpoint( + &mut self, + wall_time_unix_ms: i64, + thread_id: &str, + codex_turn_id: &str, + compaction_id: &CompactionId, + checkpoint_payload: &RawPayloadRef, + ) -> Result { + let payload = self.read_payload_json(checkpoint_payload)?; + let input_history = required_array(&payload, "input_history", checkpoint_payload)?; + let replacement_history = + required_array(&payload, "replacement_history", checkpoint_payload)?; + + let input_items = normalize::normalize_model_items(input_history, checkpoint_payload)?; + let replacement_items = + normalize::normalize_model_items(replacement_history, checkpoint_payload)?; + let input_candidates = self + .thread_conversation_snapshots + .get(thread_id) + .cloned() + .unwrap_or_default(); + let input_item_ids = self.reconcile_detached_conversation_items( + input_items, + DetachedReconcileItems { + thread_id, + codex_turn_id, + wall_time_unix_ms, + produced_by: Vec::new(), + candidates: input_candidates, + }, + )?; + // A compaction checkpoint has two transcript effects. First, record the structural + // boundary where old live history ended. Then append the replacement items, including + // the provider-visible summary item if the compact endpoint returned one. + let marker_item_id = self.create_conversation_item( + thread_id, + Some(codex_turn_id.to_string()), + wall_time_unix_ms, + NormalizedConversationItem { + role: ConversationRole::Assistant, + channel: None, + kind: ConversationItemKind::CompactionMarker, + agent_message: None, + // The summary is a separate model/provider-visible item. Keep the marker body + // empty so transcript renderers cannot mistake the boundary for prompt content. + body: ConversationBody { parts: Vec::new() }, + call_id: None, + }, + vec![ProducerRef::Compaction { + compaction_id: compaction_id.clone(), + }], + ); + let replacement_item_ids = self.reconcile_detached_conversation_items( + replacement_items, + DetachedReconcileItems { + thread_id, + codex_turn_id, + wall_time_unix_ms, + produced_by: vec![ProducerRef::Compaction { + compaction_id: compaction_id.clone(), + }], + // Replacement history is a rewrite boundary. Even if the compact endpoint emits + // text that matches old history, the installed item is a new post-compaction + // conversation item and should not reuse a pre-compaction ID. + candidates: Vec::new(), + }, + )?; + self.append_thread_conversation_items(thread_id, &input_item_ids)?; + self.append_thread_conversation_items(thread_id, std::slice::from_ref(&marker_item_id))?; + self.append_thread_conversation_items(thread_id, &replacement_item_ids)?; + Ok(ReducedCompactionCheckpoint { + input_item_ids, + marker_item_id, + replacement_item_ids, + }) + } + + fn reconcile_detached_conversation_items( + &mut self, + items: Vec, + context: DetachedReconcileItems<'_>, + ) -> Result> { + let mut item_ids = Vec::with_capacity(items.len()); + + for item in items { + let tool_link_item = item.clone(); + self.ensure_call_id_consistency(context.thread_id, &item)?; + let item_id = self + .find_matching_snapshot_item(&context.candidates, &item_ids, &item) + .unwrap_or_else(|| { + self.create_conversation_item( + context.thread_id, + Some(context.codex_turn_id.to_string()), + context.wall_time_unix_ms, + item, + context.produced_by.clone(), + ) + }); + self.update_conversation_item_from_sighting( + &item_id, + &tool_link_item, + &context.produced_by, + )?; + self.attach_model_visible_tool_item( + &item_id, + tool_link_item.call_id.as_deref(), + &tool_link_item.kind, + )?; + self.attach_model_visible_code_cell_item( + &item_id, + tool_link_item.call_id.as_deref(), + &tool_link_item.kind, + )?; + self.resolve_pending_agent_edges_for_item(&item_id)?; + item_ids.push(item_id); + } + + self.flush_pending_code_cell_starts()?; + Ok(item_ids) + } + + fn create_conversation_item( + &mut self, + thread_id: &str, + codex_turn_id: Option, + first_seen_at_unix_ms: i64, + item: NormalizedConversationItem, + produced_by: Vec, + ) -> String { + let item_id = self.next_conversation_item_id(); + self.rollout.conversation_items.insert( + item_id.clone(), + ConversationItem { + item_id: item_id.clone(), + thread_id: thread_id.to_string(), + codex_turn_id, + first_seen_at_unix_ms, + role: item.role, + channel: item.channel, + kind: item.kind, + agent_message: item.agent_message, + body: item.body, + call_id: item.call_id, + produced_by, + }, + ); + item_id + } + + fn update_conversation_item_from_sighting( + &mut self, + item_id: &str, + normalized: &NormalizedConversationItem, + produced_by: &[ProducerRef], + ) -> Result<()> { + let Some(item) = self.rollout.conversation_items.get_mut(item_id) else { + bail!("conversation item {item_id} was referenced before it was created"); + }; + + if item.kind == ConversationItemKind::Reasoning { + merge_reasoning_body(&mut item.body, &normalized.body)?; + } + for producer in produced_by { + if !item.produced_by.contains(producer) { + item.produced_by.push(producer.clone()); + } + } + Ok(()) + } + + fn append_thread_conversation_items( + &mut self, + thread_id: &str, + item_ids: &[String], + ) -> Result<()> { + let thread = self.thread_mut(thread_id)?; + for item_id in item_ids { + if !thread.conversation_item_ids.contains(item_id) { + thread.conversation_item_ids.push(item_id.clone()); + } + } + Ok(()) + } + + fn find_matching_snapshot_item( + &self, + previous_snapshot: &[String], + used_item_ids: &[String], + normalized: &NormalizedConversationItem, + ) -> Option { + previous_snapshot + .iter() + .find(|item_id| { + !used_item_ids.contains(item_id) && self.item_matches(item_id, normalized) + }) + .cloned() + } + + fn ensure_call_id_consistency( + &self, + thread_id: &str, + normalized: &NormalizedConversationItem, + ) -> Result<()> { + let Some(call_id) = normalized.call_id.as_deref() else { + return Ok(()); + }; + for item in self.rollout.conversation_items.values() { + if item.thread_id == thread_id + && item.call_id.as_deref() == Some(call_id) + && item.kind == normalized.kind + && !conversation_item_matches(item, normalized) + { + bail!("model-visible call id {call_id} was reused with different content"); + } + } + Ok(()) + } + + fn item_matches(&self, item_id: &str, normalized: &NormalizedConversationItem) -> bool { + let Some(item) = self.rollout.conversation_items.get(item_id) else { + return false; + }; + conversation_item_matches(item, normalized) + } + + fn next_conversation_item_id(&mut self) -> String { + let ordinal = self.next_conversation_item_ordinal; + self.next_conversation_item_ordinal += 1; + format!("conversation_item:{ordinal}") + } +} + +#[derive(Clone, Copy)] +enum ReconcileMode { + /// Full model requests are authoritative snapshots of the live context. The + /// prompt builder can reorder already-observed items or replace history + /// with synthetic summary messages, so item identity is "same content, + /// reused at most once in this snapshot" rather than "same position only". + FullSnapshot, + /// Incremental request deltas and response outputs append to a known prefix. + /// A mismatch at an occupied position means our reconstructed prefix is + /// wrong and should fail replay. + AppendOnly, +} + +struct ReconcileItems<'a> { + thread_id: &'a str, + codex_turn_id: &'a str, + wall_time_unix_ms: i64, + produced_by: Vec, + start_index: usize, + mode: ReconcileMode, + snapshot_override: Option<&'a [String]>, +} + +struct DetachedReconcileItems<'a> { + thread_id: &'a str, + codex_turn_id: &'a str, + wall_time_unix_ms: i64, + produced_by: Vec, + candidates: Vec, +} + +/// Conversation ids produced when a compaction checkpoint is installed. +/// +/// The marker item records the boundary, while replacement items are the live +/// history that subsequent full requests should treat as their baseline. +pub(super) struct ReducedCompactionCheckpoint { + pub(super) input_item_ids: Vec, + pub(super) marker_item_id: String, + pub(super) replacement_item_ids: Vec, +} + +fn required_array<'a>( + payload: &'a Value, + key: &str, + raw_payload: &RawPayloadRef, +) -> Result<&'a Vec> { + payload.get(key).and_then(Value::as_array).with_context(|| { + format!( + "compaction checkpoint payload {} did not contain array {key}", + raw_payload.raw_payload_id + ) + }) +} + +fn conversation_item_matches( + item: &ConversationItem, + normalized: &NormalizedConversationItem, +) -> bool { + let body_matches = if item.kind == ConversationItemKind::Reasoning + && normalized.kind == ConversationItemKind::Reasoning + { + reasoning_body_matches(&item.body, &normalized.body) + } else { + conversation_body_matches(&item.body, &normalized.body) + }; + + item.role == normalized.role + && item.channel == normalized.channel + && item.kind == normalized.kind + && item.agent_message == normalized.agent_message + && body_matches + && item.call_id == normalized.call_id +} + +fn conversation_body_matches(left: &ConversationBody, right: &ConversationBody) -> bool { + left.parts.len() == right.parts.len() + && left + .parts + .iter() + .zip(&right.parts) + .all(|(left, right)| match (left, right) { + ( + ConversationPart::Json { + summary: left_summary, + raw_payload_id: _, + }, + ConversationPart::Json { + summary: right_summary, + raw_payload_id: _, + }, + ) => left_summary == right_summary, + _ => left == right, + }) +} + +fn reasoning_body_matches(left: &ConversationBody, right: &ConversationBody) -> bool { + if conversation_body_matches(left, right) { + return true; + } + + // The Responses API may return readable reasoning on completion, but later + // request snapshots often replay only the encrypted blob. Treat the blob as + // stable model-visible identity and merge readable text as best-effort + // evidence, because request/response serialization can observe different + // readable forms for the same encrypted reasoning item. + let Some(left_encoded) = reasoning_encoded_part(left) else { + return false; + }; + let Some(right_encoded) = reasoning_encoded_part(right) else { + return false; + }; + + left_encoded == right_encoded +} + +fn merge_reasoning_body( + existing: &mut ConversationBody, + incoming: &ConversationBody, +) -> Result<()> { + if conversation_body_matches(existing, incoming) { + return Ok(()); + } + if !reasoning_body_matches(existing, incoming) { + bail!("reasoning item merge attempted with different encrypted_content identity"); + } + + let existing_text_parts = reasoning_text_parts(existing); + let existing_summary_parts = reasoning_summary_parts(existing); + if !existing_text_parts.is_empty() && !existing_summary_parts.is_empty() { + return Ok(()); + } + + let incoming_text_parts = reasoning_text_parts(incoming); + let incoming_summary_parts = reasoning_summary_parts(incoming); + + let text_parts = if !existing_text_parts.is_empty() { + existing_text_parts + } else { + incoming_text_parts + }; + + let summary_parts = if !existing_summary_parts.is_empty() { + existing_summary_parts + } else { + incoming_summary_parts + }; + + // We already know that the encoded part exist (and matches). + let encoded_parts = reasoning_encoded_parts(existing); + + existing.parts = text_parts + .into_iter() + .cloned() + .chain(summary_parts.into_iter().cloned()) + .chain(encoded_parts.into_iter().cloned()) + .collect(); + + Ok(()) +} + +fn reasoning_text_parts(body: &ConversationBody) -> Vec<&ConversationPart> { + body.parts + .iter() + .filter(|part| matches!(part, ConversationPart::Text { .. })) + .collect() +} + +fn reasoning_summary_parts(body: &ConversationBody) -> Vec<&ConversationPart> { + body.parts + .iter() + .filter(|part| matches!(part, ConversationPart::Summary { .. })) + .collect() +} + +fn reasoning_encoded_parts(body: &ConversationBody) -> Vec<&ConversationPart> { + body.parts + .iter() + .filter(|part| matches!(part, ConversationPart::Encoded { .. })) + .collect() +} + +fn reasoning_encoded_part(body: &ConversationBody) -> Option<(&str, &str)> { + body.parts.iter().find_map(|part| { + if let ConversationPart::Encoded { label, value } = part { + Some((label.as_str(), value.as_str())) + } else { + None + } + }) +} + +#[cfg(test)] +#[path = "conversation_tests.rs"] +mod tests; diff --git a/codex-rs/rollout-trace/src/reducer/conversation/normalize.rs b/codex-rs/rollout-trace/src/reducer/conversation/normalize.rs new file mode 100644 index 0000000000000000000000000000000000000000..142c1ea93028805c98b8674cacfbee8cc93d5d9f --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/conversation/normalize.rs @@ -0,0 +1,516 @@ +//! Normalization from Responses-shaped JSON items into conversation item data. + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use codex_protocol::models::AgentMessageInputContent; +use codex_protocol::models::ResponseItem; +use serde_json::Value; + +use crate::model::AgentMessageMetadata; +use crate::model::ConversationBody; +use crate::model::ConversationChannel; +use crate::model::ConversationItemKind; +use crate::model::ConversationPart; +use crate::model::ConversationRole; +use crate::model::TokenUsage; +use crate::payload::RawPayloadRef; + +/// Conversation fields parsed from one Responses item before trace identity. +/// +/// IDs and provenance are assigned after positional reconciliation. Keeping the +/// normalized data separate from `ConversationItem` makes reuse vs insertion a +/// single reducer decision instead of something the parser has to know about. +#[derive(Clone)] +pub(super) struct NormalizedConversationItem { + pub(super) role: ConversationRole, + pub(super) channel: Option, + pub(super) kind: ConversationItemKind, + pub(super) agent_message: Option, + pub(super) body: ConversationBody, + pub(super) call_id: Option, +} + +pub(super) fn normalize_model_items( + items: &[Value], + raw_payload: &RawPayloadRef, +) -> Result> { + let mut normalized_items = Vec::new(); + for item in items { + if item.get("type").and_then(Value::as_str) == Some("additional_tools") { + continue; + } + + let mut model_visible_item = item.clone(); + if let Some(object) = model_visible_item.as_object_mut() { + object.remove("internal_chat_message_metadata_passthrough"); + } + normalized_items.push(normalize_model_item(&model_visible_item, raw_payload)?); + } + Ok(normalized_items) +} + +pub(super) fn token_usage_from_value(value: &Value) -> Option { + Some(TokenUsage { + input_tokens: u64_field(value, "input_tokens")?, + cached_input_tokens: u64_field(value, "cached_input_tokens")?, + cache_write_input_tokens: u64_field(value, "cache_write_input_tokens").unwrap_or(0), + output_tokens: u64_field(value, "output_tokens")?, + reasoning_output_tokens: u64_field(value, "reasoning_output_tokens")?, + }) +} + +fn normalize_model_item( + item: &Value, + raw_payload: &RawPayloadRef, +) -> Result { + let Some(item_type) = item.get("type").and_then(Value::as_str) else { + bail!( + "model item in payload {} did not contain a string type", + raw_payload.raw_payload_id + ); + }; + match item_type { + "message" => normalize_message_item(item, raw_payload), + "agent_message" => normalize_agent_message_item(item, raw_payload), + "reasoning" => normalize_reasoning_item(item, raw_payload), + "function_call" => Ok(NormalizedConversationItem { + role: ConversationRole::Assistant, + channel: Some(ConversationChannel::Commentary), + kind: ConversationItemKind::FunctionCall, + agent_message: None, + body: raw_text_or_json_body(item.get("arguments"), raw_payload), + call_id: item + .get("call_id") + .and_then(Value::as_str) + .map(ToString::to_string), + }), + "function_call_output" => Ok(NormalizedConversationItem { + role: ConversationRole::Tool, + channel: Some(ConversationChannel::Commentary), + kind: ConversationItemKind::FunctionCallOutput, + agent_message: None, + body: tool_output_body(item.get("output"), raw_payload), + call_id: item + .get("call_id") + .and_then(Value::as_str) + .map(ToString::to_string), + }), + "custom_tool_call" => Ok(NormalizedConversationItem { + role: ConversationRole::Assistant, + channel: Some(ConversationChannel::Commentary), + kind: ConversationItemKind::CustomToolCall, + agent_message: None, + body: custom_tool_call_body(item, raw_payload), + call_id: item + .get("call_id") + .and_then(Value::as_str) + .map(ToString::to_string), + }), + "custom_tool_call_output" => Ok(NormalizedConversationItem { + role: ConversationRole::Tool, + channel: Some(ConversationChannel::Commentary), + kind: ConversationItemKind::CustomToolCallOutput, + agent_message: None, + body: tool_output_body(item.get("output"), raw_payload), + call_id: item + .get("call_id") + .and_then(Value::as_str) + .map(ToString::to_string), + }), + "tool_search_call" | "web_search_call" | "image_generation_call" | "local_shell_call" => { + Ok(NormalizedConversationItem { + role: ConversationRole::Assistant, + channel: Some(ConversationChannel::Commentary), + kind: ConversationItemKind::FunctionCall, + agent_message: None, + body: json_body(item, raw_payload), + call_id: item + .get("call_id") + .and_then(Value::as_str) + .map(ToString::to_string), + }) + } + "tool_search_output" | "mcp_tool_call_output" => Ok(NormalizedConversationItem { + role: ConversationRole::Tool, + channel: Some(ConversationChannel::Commentary), + kind: ConversationItemKind::FunctionCallOutput, + agent_message: None, + body: json_body(item, raw_payload), + call_id: item + .get("call_id") + .and_then(Value::as_str) + .map(ToString::to_string), + }), + "compaction" | "compaction_summary" | "context_compaction" => { + Ok(NormalizedConversationItem { + role: ConversationRole::Assistant, + channel: Some(ConversationChannel::Summary), + kind: ConversationItemKind::Message, + agent_message: None, + body: compaction_body(item, raw_payload)?, + call_id: None, + }) + } + _ => bail!( + "unsupported model item type {item_type} in payload {}", + raw_payload.raw_payload_id + ), + } +} + +fn normalize_message_item( + item: &Value, + raw_payload: &RawPayloadRef, +) -> Result { + let Some(role) = item.get("role").and_then(Value::as_str) else { + bail!( + "message item in payload {} did not contain a string role", + raw_payload.raw_payload_id + ); + }; + let Some(role) = role_from_str(role) else { + bail!( + "unsupported message role {role} in payload {}", + raw_payload.raw_payload_id + ); + }; + Ok(NormalizedConversationItem { + role, + channel: item + .get("phase") + .and_then(Value::as_str) + .and_then(channel_from_phase), + kind: ConversationItemKind::Message, + agent_message: None, + body: ConversationBody { + parts: content_parts(item.get("content"), raw_payload), + }, + call_id: None, + }) +} + +fn normalize_agent_message_item( + item: &Value, + raw_payload: &RawPayloadRef, +) -> Result { + let raw_payload_id = &raw_payload.raw_payload_id; + let response_item = + serde_json::from_value::(item.clone()).with_context(|| { + format!("failed to parse agent_message item in payload {raw_payload_id}") + })?; + let ResponseItem::AgentMessage { + author, + recipient, + content, + .. + } = response_item + else { + bail!("item in payload {raw_payload_id} was not an agent_message"); + }; + let parts = content + .into_iter() + .map(|content| match content { + AgentMessageInputContent::InputText { text } => ConversationPart::Text { text }, + AgentMessageInputContent::EncryptedContent { encrypted_content } => { + ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: encrypted_content, + } + } + }) + .collect::>(); + if parts.is_empty() { + bail!("agent_message item in payload {raw_payload_id} contained no content"); + } + + Ok(NormalizedConversationItem { + role: ConversationRole::Assistant, + channel: Some(ConversationChannel::Analysis), + kind: ConversationItemKind::Message, + agent_message: Some(AgentMessageMetadata { author, recipient }), + body: ConversationBody { parts }, + call_id: None, + }) +} + +fn normalize_reasoning_item( + item: &Value, + raw_payload: &RawPayloadRef, +) -> Result { + let mut parts = Vec::new(); + append_reasoning_parts( + item, + "content", + ReasoningPartKind::Content, + raw_payload, + &mut parts, + )?; + append_reasoning_parts( + item, + "summary", + ReasoningPartKind::Summary, + raw_payload, + &mut parts, + )?; + + if let Some(encrypted_content) = item.get("encrypted_content") { + let encrypted_content = match encrypted_content { + Value::Null => None, + Value::String(encrypted_content) => Some(encrypted_content), + _ => { + bail!( + "reasoning item in payload {} had non-string encrypted_content", + raw_payload.raw_payload_id + ); + } + }; + if let Some(encrypted_content) = encrypted_content { + parts.push(ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: encrypted_content.to_string(), + }); + } + } + + if parts.is_empty() { + bail!( + "reasoning item in payload {} contained no content, summary, or encrypted_content", + raw_payload.raw_payload_id + ); + } + + Ok(NormalizedConversationItem { + role: ConversationRole::Assistant, + channel: Some(ConversationChannel::Analysis), + kind: ConversationItemKind::Reasoning, + agent_message: None, + body: ConversationBody { parts }, + call_id: None, + }) +} + +#[derive(Clone, Copy)] +enum ReasoningPartKind { + Content, + Summary, +} + +fn append_reasoning_parts( + item: &Value, + key: &str, + kind: ReasoningPartKind, + raw_payload: &RawPayloadRef, + parts: &mut Vec, +) -> Result<()> { + let Some(items) = item.get(key) else { + return Ok(()); + }; + if matches!((kind, items), (ReasoningPartKind::Content, Value::Null)) { + return Ok(()); + } + let Some(items) = items.as_array() else { + bail!( + "reasoning item in payload {} had non-array {key}", + raw_payload.raw_payload_id + ); + }; + + for content_item in items { + let Some(item_type) = content_item.get("type").and_then(Value::as_str) else { + bail!( + "reasoning item in payload {} had {key} entry without string type", + raw_payload.raw_payload_id + ); + }; + let expected_type = match kind { + ReasoningPartKind::Content => { + if !matches!(item_type, "reasoning_text" | "text") { + bail!( + "reasoning item in payload {} had unsupported content type {item_type}", + raw_payload.raw_payload_id + ); + } + "content" + } + ReasoningPartKind::Summary => { + if item_type != "summary_text" { + bail!( + "reasoning item in payload {} had unsupported summary type {item_type}", + raw_payload.raw_payload_id + ); + } + "summary" + } + }; + + let Some(text) = content_item.get("text").and_then(Value::as_str) else { + bail!( + "reasoning item in payload {} had {expected_type} entry without string text", + raw_payload.raw_payload_id + ); + }; + match kind { + ReasoningPartKind::Content => parts.push(ConversationPart::Text { + text: text.to_string(), + }), + ReasoningPartKind::Summary => parts.push(ConversationPart::Summary { + text: text.to_string(), + }), + } + } + + Ok(()) +} + +fn role_from_str(role: &str) -> Option { + match role { + "system" => Some(ConversationRole::System), + "developer" => Some(ConversationRole::Developer), + "user" => Some(ConversationRole::User), + "assistant" => Some(ConversationRole::Assistant), + "tool" => Some(ConversationRole::Tool), + _ => None, + } +} + +fn channel_from_phase(phase: &str) -> Option { + match phase { + "commentary" => Some(ConversationChannel::Commentary), + "final_answer" => Some(ConversationChannel::Final), + "summary" => Some(ConversationChannel::Summary), + _ => None, + } +} + +fn content_parts(content: Option<&Value>, raw_payload: &RawPayloadRef) -> Vec { + let Some(content) = content.and_then(Value::as_array) else { + return vec![payload_ref_part("content", raw_payload)]; + }; + + let mut parts = Vec::new(); + for part in content { + match part.get("type").and_then(Value::as_str) { + Some("input_text" | "output_text" | "text") => { + if let Some(text) = part.get("text").and_then(Value::as_str) { + parts.push(ConversationPart::Text { + text: text.to_string(), + }); + } + } + Some("input_image") => parts.push(payload_ref_part("input_image", raw_payload)), + Some(other) => parts.push(payload_ref_part(other, raw_payload)), + None => parts.push(payload_ref_part("content", raw_payload)), + } + } + + if parts.is_empty() { + parts.push(payload_ref_part("empty_content", raw_payload)); + } + parts +} + +fn custom_tool_call_body(item: &Value, raw_payload: &RawPayloadRef) -> ConversationBody { + let Some(input) = item.get("input").and_then(Value::as_str) else { + return json_body(item, raw_payload); + }; + if item.get("name").and_then(Value::as_str) == Some("exec") { + ConversationBody { + parts: vec![ConversationPart::Code { + language: "javascript".to_string(), + source: input.to_string(), + }], + } + } else { + ConversationBody { + parts: vec![ConversationPart::Text { + text: input.to_string(), + }], + } + } +} + +fn raw_text_or_json_body(value: Option<&Value>, raw_payload: &RawPayloadRef) -> ConversationBody { + match value { + Some(Value::String(text)) => { + if let Ok(json) = serde_json::from_str::(text) { + json_body(&json, raw_payload) + } else { + ConversationBody { + parts: vec![ConversationPart::Text { text: text.clone() }], + } + } + } + Some(value) => json_body(value, raw_payload), + None => ConversationBody { + parts: vec![payload_ref_part("payload", raw_payload)], + }, + } +} + +fn tool_output_body(output: Option<&Value>, raw_payload: &RawPayloadRef) -> ConversationBody { + match output { + Some(Value::String(text)) => ConversationBody { + parts: vec![ConversationPart::Text { text: text.clone() }], + }, + Some(Value::Array(_)) => ConversationBody { + parts: content_parts(output, raw_payload), + }, + Some(value) => json_body(value, raw_payload), + None => ConversationBody { + parts: vec![payload_ref_part("tool_output", raw_payload)], + }, + } +} + +fn compaction_body(item: &Value, raw_payload: &RawPayloadRef) -> Result { + let Some(encrypted_content) = item.get("encrypted_content").and_then(Value::as_str) else { + bail!( + "compaction item in payload {} did not contain string encrypted_content", + raw_payload.raw_payload_id + ); + }; + // `type: "compaction"` is the remote-compaction summary that later re-enters model requests. + // The structural "history was cut here" marker is inserted separately when the checkpoint is + // installed; payload refs are observation-local, so the encoded summary itself is identity. + Ok(ConversationBody { + parts: vec![ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: encrypted_content.to_string(), + }], + }) +} + +fn json_body(value: &Value, raw_payload: &RawPayloadRef) -> ConversationBody { + ConversationBody { + parts: vec![ConversationPart::Json { + summary: summarize_json(value), + raw_payload_id: raw_payload.raw_payload_id.clone(), + }], + } +} + +fn payload_ref_part(label: &str, raw_payload: &RawPayloadRef) -> ConversationPart { + ConversationPart::PayloadRef { + label: label.to_string(), + raw_payload_id: raw_payload.raw_payload_id.clone(), + } +} + +fn summarize_json(value: &Value) -> String { + const MAX_JSON_SUMMARY_LEN: usize = 240; + let mut summary = + serde_json::to_string(value).unwrap_or_else(|_| "".to_string()); + if summary.len() > MAX_JSON_SUMMARY_LEN { + summary.truncate(MAX_JSON_SUMMARY_LEN); + summary.push_str("..."); + } + summary +} + +fn u64_field(value: &Value, field: &str) -> Option { + value + .get(field) + .and_then(Value::as_i64) + .map(|value| value.max(0) as u64) +} diff --git a/codex-rs/rollout-trace/src/reducer/conversation_tests.rs b/codex-rs/rollout-trace/src/reducer/conversation_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4c4f7ace9727dce4f6122ffc82e7f4c0b47c6868 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/conversation_tests.rs @@ -0,0 +1,1288 @@ +use pretty_assertions::assert_eq; +use serde_json::json; +use tempfile::TempDir; + +use crate::model::AgentMessageMetadata; +use crate::model::ConversationBody; +use crate::model::ConversationChannel; +use crate::model::ConversationItemKind; +use crate::model::ConversationPart; +use crate::model::ConversationRole; +use crate::model::ExecutionStatus; +use crate::model::ProducerRef; +use crate::model::ToolCallKind; +use crate::model::ToolCallSummary; +use crate::payload::RawPayloadKind; +use crate::raw_event::RawTraceEventPayload; +use crate::reducer::test_support::append_inference_completion; +use crate::reducer::test_support::append_inference_start; +use crate::reducer::test_support::create_started_writer; +use crate::reducer::test_support::expect_replay_error; +use crate::reducer::test_support::message; +use crate::reducer::test_support::start_turn; +use crate::reducer::test_support::trace_context; +use crate::replay_bundle; + +#[test] +fn request_snapshots_reuse_history_without_deduping_new_identical_items() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let first_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "ok")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", first_request)?; + start_turn(&writer, "turn-2")?; + + let second_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + message("user", "ok"), + message("assistant", "ack"), + message("user", "ok") + ] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", second_request)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"].request_item_ids; + let second = &rollout.inference_calls["inference-2"].request_item_ids; + + assert_eq!(first.len(), 1); + assert_eq!(second.len(), 3); + assert_eq!(second[0], first[0]); + assert_ne!(second[2], first[0]); + assert_eq!(rollout.conversation_items.len(), 3); + assert_eq!( + rollout.threads["thread-root"].conversation_item_ids, + *second + ); + + Ok(()) +} + +#[test] +fn response_outputs_enter_thread_conversation_on_completion() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "run tests")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "tests passed"}] + } + ] + }), + )?; + append_inference_completion(&writer, "inference-1", "resp-1", response)?; + + let rollout = replay_bundle(temp.path())?; + let inference = &rollout.inference_calls["inference-1"]; + let mut expected_thread_items = inference.request_item_ids.clone(); + expected_thread_items.extend(inference.response_item_ids.clone()); + + assert_eq!(inference.response_item_ids.len(), 1); + assert_eq!( + rollout.threads["thread-root"].conversation_item_ids, + expected_thread_items, + ); + + Ok(()) +} + +#[test] +fn agent_messages_preserve_routing_and_content() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + { + "type": "agent_message", + "author": "/root/worker", + "recipient": "/root", + "content": [{"type": "input_text", "text": "done"}] + }, + { + "type": "agent_message", + "author": "/root", + "recipient": "/root/worker", + "content": [{ + "type": "encrypted_content", + "encrypted_content": "encrypted-task" + }] + } + ] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let rollout = replay_bundle(temp.path())?; + let actual = rollout.inference_calls["inference-1"] + .request_item_ids + .iter() + .map(|item_id| { + let item = &rollout.conversation_items[item_id]; + ( + item.role.clone(), + item.channel.clone(), + item.kind.clone(), + item.agent_message.clone(), + item.body.clone(), + ) + }) + .collect::>(); + + assert_eq!( + actual, + vec![ + ( + ConversationRole::Assistant, + Some(ConversationChannel::Analysis), + ConversationItemKind::Message, + Some(AgentMessageMetadata { + author: "/root/worker".to_string(), + recipient: "/root".to_string(), + }), + ConversationBody { + parts: vec![ConversationPart::Text { + text: "done".to_string(), + }], + }, + ), + ( + ConversationRole::Assistant, + Some(ConversationChannel::Analysis), + ConversationItemKind::Message, + Some(AgentMessageMetadata { + author: "/root".to_string(), + recipient: "/root/worker".to_string(), + }), + ConversationBody { + parts: vec![ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: "encrypted-task".to_string(), + }], + }, + ), + ] + ); + + Ok(()) +} + +#[test] +fn later_full_request_reuses_prior_json_tool_call_by_position() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "run tests")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [{ + "type": "function_call", + "name": "shell", + "arguments": "{\"cmd\":\"cargo test\"}", + "call_id": "call-1" + }] + }), + )?; + append_inference_completion(&writer, "inference-1", "resp-1", response)?; + start_turn(&writer, "turn-2")?; + + let next_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + message("user", "run tests"), + { + "type": "function_call", + "name": "shell", + "arguments": "{\"cmd\":\"cargo test\"}", + "call_id": "call-1" + } + ] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", next_request)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"]; + let second = &rollout.inference_calls["inference-2"]; + + assert_eq!( + second.request_item_ids, + vec![ + first.request_item_ids[0].clone(), + first.response_item_ids[0].clone(), + ], + ); + assert_eq!(rollout.conversation_items.len(), 2); + + Ok(()) +} + +#[test] +fn request_reuses_prior_tool_search_call_with_internal_metadata() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "search")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [{ + "type": "tool_search_call", + "status": "completed", + "call_id": "call-search", + "arguments": { + "query": "spawn subagent launch manage agents report result", + "limit": 10 + }, + "execution": "client" + }] + }), + )?; + append_inference_completion(&writer, "inference-1", "resp-1", response)?; + start_turn(&writer, "turn-2")?; + + let next_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + message("user", "search"), + { + "type": "tool_search_call", + "status": "completed", + "call_id": "call-search", + "arguments": { + "query": "spawn subagent launch manage agents report result", + "limit": 10 + }, + "execution": "client", + "internal_chat_message_metadata_passthrough": { + "turn_id": "turn-1" + } + } + ] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", next_request)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"]; + let second = &rollout.inference_calls["inference-2"]; + + assert_eq!( + second.request_item_ids, + vec![ + first.request_item_ids[0].clone(), + first.response_item_ids[0].clone(), + ], + ); + assert_eq!(rollout.conversation_items.len(), 2); + + Ok(()) +} + +#[test] +fn request_reuses_prior_tool_outputs_with_internal_metadata() -> anyhow::Result<()> { + for item_type in ["tool_search_output", "mcp_tool_call_output"] { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let output = json!({ + "type": item_type, + "status": "completed", + "call_id": "call-search", + "execution": "client", + "tools": [{ + "name": "search", + "internal_chat_message_metadata_passthrough": "model-visible" + }] + }); + let first_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "search"), output.clone()] + }), + )?; + let first_request_payload_id = first_request.raw_payload_id.clone(); + append_inference_start(&writer, "inference-1", "turn-1", first_request)?; + start_turn(&writer, "turn-2")?; + + let mut replayed_output = output.clone(); + replayed_output["internal_chat_message_metadata_passthrough"] = json!({ + "turn_id": "turn-1" + }); + let next_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "search"), replayed_output] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", next_request)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"]; + let second = &rollout.inference_calls["inference-2"]; + + assert_eq!(second.request_item_ids, first.request_item_ids); + assert_eq!( + rollout.conversation_items[&first.request_item_ids[1]].body, + ConversationBody { + parts: vec![ConversationPart::Json { + summary: serde_json::to_string(&output)?, + raw_payload_id: first_request_payload_id, + }], + }, + ); + assert_eq!(rollout.conversation_items.len(), 2); + } + + Ok(()) +} + +#[test] +fn tool_output_call_id_reuse_with_different_nested_metadata_is_reducer_error() -> anyhow::Result<()> +{ + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [{ + "type": "tool_search_output", + "status": "completed", + "call_id": "call-search", + "execution": "client", + "tools": [{ + "name": "search", + "internal_chat_message_metadata_passthrough": "first" + }] + }] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + start_turn(&writer, "turn-2")?; + + let conflicting_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [{ + "type": "tool_search_output", + "status": "completed", + "call_id": "call-search", + "execution": "client", + "tools": [{ + "name": "search", + "internal_chat_message_metadata_passthrough": "different" + }], + "internal_chat_message_metadata_passthrough": { + "turn_id": "turn-1" + } + }] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", conflicting_request)?; + + expect_replay_error( + &temp, + "model-visible call id call-search was reused with different content", + ) +} + +#[test] +fn incremental_request_carries_prior_request_and_response_items_forward() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "run tests")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "token_usage": { + "input_tokens": 10, + "cached_input_tokens": 1, + "cache_write_input_tokens": 3, + "output_tokens": 5, + "reasoning_output_tokens": 2, + "total_tokens": 15 + }, + "output_items": [ + { + "type": "function_call", + "name": "shell", + "arguments": "{\"cmd\":\"cargo test\"}", + "call_id": "call-1" + } + ] + }), + )?; + append_inference_completion(&writer, "inference-1", "resp-1", response)?; + start_turn(&writer, "turn-2")?; + + let incremental_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "type": "response.create", + "previous_response_id": "resp-1", + "input": [ + { + "type": "function_call_output", + "call_id": "call-1", + "output": "tests passed" + } + ] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", incremental_request)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"]; + let second = &rollout.inference_calls["inference-2"]; + + assert_eq!(first.response_item_ids.len(), 1); + assert_eq!( + second.request_item_ids, + vec![ + first.request_item_ids[0].clone(), + first.response_item_ids[0].clone(), + rollout.threads["thread-root"].conversation_item_ids[2].clone(), + ], + ); + assert_eq!( + rollout.threads["thread-root"].conversation_item_ids, + second.request_item_ids, + ); + assert_eq!( + first.usage.as_ref().map(|usage| usage.input_tokens), + Some(10), + ); + assert_eq!( + first + .usage + .as_ref() + .map(|usage| usage.cache_write_input_tokens), + Some(3), + ); + + Ok(()) +} + +#[test] +fn full_request_snapshot_can_reorder_existing_items_and_insert_summary() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + message("developer", "follow the repo rules"), + message("user", "count files") + ] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + start_turn(&writer, "turn-2")?; + + let compacted_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + message("user", "count files"), + message("user", "summary from a compacted prior attempt"), + message("developer", "follow the repo rules") + ] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", compacted_request)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"].request_item_ids; + let second = &rollout.inference_calls["inference-2"].request_item_ids; + + assert_eq!(second[0], first[1]); + assert_eq!(second[2], first[0]); + assert_ne!(second[1], first[0]); + assert_ne!(second[1], first[1]); + assert_eq!(rollout.conversation_items.len(), 3); + + Ok(()) +} + +#[test] +fn reasoning_body_preserves_text_summary_and_encoded_content() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "think visibly")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [{ + "type": "reasoning", + "content": [{"type": "reasoning_text", "text": "raw reasoning"}], + "summary": [{"type": "summary_text", "text": "brief summary"}], + "encrypted_content": "encoded-reasoning" + }] + }), + )?; + append_inference_completion(&writer, "inference-1", "resp-1", response)?; + + let rollout = replay_bundle(temp.path())?; + let reasoning_item_id = &rollout.inference_calls["inference-1"].response_item_ids[0]; + + assert_eq!( + rollout.conversation_items[reasoning_item_id].body.parts, + vec![ + ConversationPart::Text { + text: "raw reasoning".to_string(), + }, + ConversationPart::Summary { + text: "brief summary".to_string(), + }, + ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: "encoded-reasoning".to_string(), + }, + ], + ); + + Ok(()) +} + +#[test] +fn encrypted_reasoning_reuses_response_item_in_later_request() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let user = message("user", "count files"); + let function_call = json!({ + "type": "function_call", + "name": "shell", + "arguments": "{\"cmd\":\"find . -maxdepth 1 -type f | wc -l\"}", + "call_id": "call-1" + }); + let encrypted_reasoning = json!({ + "type": "reasoning", + "summary": [], + "encrypted_content": "encoded-reasoning" + }); + let readable_reasoning = json!({ + "type": "reasoning", + "content": [{"type": "text", "text": "need count"}], + "summary": [], + "encrypted_content": "encoded-reasoning" + }); + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [user] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [ + readable_reasoning, + function_call + ] + }), + )?; + append_inference_completion(&writer, "inference-1", "resp-1", response)?; + start_turn(&writer, "turn-2")?; + + let followup = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + user, + encrypted_reasoning, + function_call, + { + "type": "function_call_output", + "call_id": "call-1", + "output": "31\n" + } + ] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", followup)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"]; + let second = &rollout.inference_calls["inference-2"]; + let output_item_id = rollout.threads["thread-root"].conversation_item_ids[3].clone(); + + assert_eq!( + second.request_item_ids, + vec![ + first.request_item_ids[0].clone(), + first.response_item_ids[0].clone(), + first.response_item_ids[1].clone(), + output_item_id, + ], + ); + assert_eq!( + rollout.conversation_items[&first.response_item_ids[0]] + .body + .parts, + vec![ + ConversationPart::Text { + text: "need count".to_string(), + }, + ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: "encoded-reasoning".to_string(), + }, + ], + ); + assert_eq!(rollout.conversation_items.len(), 4); + assert_eq!( + rollout.threads["thread-root"].conversation_item_ids, + second.request_item_ids, + ); + + Ok(()) +} + +#[test] +fn encrypted_reasoning_upgrades_when_later_sighting_has_more_readable_body() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let user = message("user", "count files"); + // Both sightings carry the same encrypted identity, but each has a + // different kind of readable evidence. The reducer should keep both + // observations because text and summary are complementary. + let text_only_reasoning = json!({ + "type": "reasoning", + "content": [{"type": "text", "text": "need count"}], + "summary": [], + "encrypted_content": "encoded-reasoning" + }); + let summary_only_reasoning = json!({ + "type": "reasoning", + "summary": [{"type": "summary_text", "text": "counting files"}], + "encrypted_content": "encoded-reasoning" + }); + + let first_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [user, text_only_reasoning] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", first_request)?; + start_turn(&writer, "turn-2")?; + + let second_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [user, summary_only_reasoning] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", second_request)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"]; + let second = &rollout.inference_calls["inference-2"]; + let reasoning_item_id = &first.request_item_ids[1]; + + // The reducer should keep one conversation item and merge the missing + // readable kind without treating the second sighting as a conflict. + assert_eq!(&second.request_item_ids[1], reasoning_item_id); + assert_eq!( + rollout.conversation_items[reasoning_item_id].body.parts, + vec![ + ConversationPart::Text { + text: "need count".to_string(), + }, + ConversationPart::Summary { + text: "counting files".to_string(), + }, + ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: "encoded-reasoning".to_string(), + }, + ], + ); + assert_eq!(rollout.conversation_items.len(), 2); + + Ok(()) +} + +#[test] +fn same_encrypted_reasoning_with_different_text_reuses_first_readable_body() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let user = message("user", "count files"); + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [user] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + // The response is the first readable observation for this encrypted + // reasoning blob, so it is the body later conflicting sightings must not + // overwrite. + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [{ + "type": "reasoning", + "content": [{"type": "text", "text": "first text"}], + "summary": [], + "encrypted_content": "encoded-reasoning" + }] + }), + )?; + append_inference_completion(&writer, "inference-1", "resp-1", response)?; + start_turn(&writer, "turn-2")?; + + let conflicting_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + user, + { + "type": "reasoning", + "content": [{"type": "text", "text": "different text"}], + "summary": [], + "encrypted_content": "encoded-reasoning" + } + ] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", conflicting_request)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"]; + let second = &rollout.inference_calls["inference-2"]; + let reasoning_item_id = &first.response_item_ids[0]; + + // Same encrypted identity still reuses the response item, but conflicting + // readable text is not a safe upgrade. + assert_eq!( + second.request_item_ids, + vec![first.request_item_ids[0].clone(), reasoning_item_id.clone(),], + ); + assert_eq!( + rollout.conversation_items[reasoning_item_id].body.parts, + vec![ + ConversationPart::Text { + text: "first text".to_string(), + }, + ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: "encoded-reasoning".to_string(), + }, + ], + ); + assert_eq!(rollout.conversation_items.len(), 2); + + Ok(()) +} + +#[test] +fn model_visible_call_id_reuse_with_different_content_is_reducer_error() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [{ + "type": "function_call", + "name": "shell", + "arguments": "{\"cmd\":\"cargo test\"}", + "call_id": "call-1" + }] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + start_turn(&writer, "turn-2")?; + + let conflicting_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [{ + "type": "function_call", + "name": "shell", + "arguments": "{\"cmd\":\"cargo check\"}", + "call_id": "call-1" + }] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", conflicting_request)?; + + expect_replay_error( + &temp, + "model-visible call id call-1 was reused with different content", + ) +} + +#[test] +fn unsupported_model_item_is_reducer_error() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + { + "type": "new_unhandled_model_item", + "payload": "must not be silently skipped" + } + ] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + expect_replay_error( + &temp, + "unsupported model item type new_unhandled_model_item", + ) +} + +#[test] +fn additional_tools_are_excluded_from_request_conversation() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [ + { + "type": "additional_tools", + "role": "developer", + "tools": [{ + "type": "function", + "name": "lookup", + "parameters": {"type": "object", "properties": {}} + }] + }, + message("user", "find it") + ] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let rollout = replay_bundle(temp.path())?; + let request_item_ids = &rollout.inference_calls["inference-1"].request_item_ids; + + assert_eq!(request_item_ids.len(), 1); + assert_eq!(rollout.conversation_items.len(), 1); + assert_eq!( + rollout.conversation_items[&request_item_ids[0]].body, + ConversationBody { + parts: vec![ConversationPart::Text { + text: "find it".to_string(), + }], + } + ); + + Ok(()) +} + +#[test] +fn missing_request_input_is_reducer_error() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "model": "gpt-test" + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + expect_replay_error(&temp, "did not contain input") +} + +#[test] +fn unknown_previous_response_id_is_reducer_error() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "previous_response_id": "resp-missing", + "input": [message("user", "still here")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + expect_replay_error(&temp, "unknown previous_response_id resp-missing") +} + +#[test] +fn compaction_boundary_repeats_prefix_and_reuses_replacement_items() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let developer = message("developer", "follow repo rules"); + let user = message("user", "count files"); + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [developer, user] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let summary = message("user", "summary from compacted history"); + let compaction_summary = json!({ + "type": "compaction", + "encrypted_content": "encrypted-summary", + }); + let checkpoint = writer.write_json_payload( + RawPayloadKind::CompactionCheckpoint, + &json!({ + "input_history": [developer, user], + "replacement_history": [user, summary, compaction_summary] + }), + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CompactionInstalled { + compaction_id: "compaction-1".to_string(), + checkpoint_payload: checkpoint, + }, + )?; + + start_turn(&writer, "turn-2")?; + let post_compaction_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [developer, user, summary, compaction_summary] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", post_compaction_request)?; + + let rollout = replay_bundle(temp.path())?; + let first = &rollout.inference_calls["inference-1"]; + let second = &rollout.inference_calls["inference-2"]; + let compaction = &rollout.compactions["compaction-1"]; + + assert_eq!(compaction.input_item_ids, first.request_item_ids); + assert_eq!(second.request_item_ids.len(), 4); + assert_eq!( + &second.request_item_ids[1..], + compaction.replacement_item_ids.as_slice() + ); + let marker = &rollout.conversation_items[&compaction.marker_item_id]; + assert_eq!(marker.kind, ConversationItemKind::CompactionMarker); + assert_eq!(marker.body.parts, Vec::::new()); + assert_eq!( + marker.produced_by, + vec![ProducerRef::Compaction { + compaction_id: "compaction-1".to_string() + }], + ); + assert_ne!(second.request_item_ids[0], first.request_item_ids[0]); + assert_ne!( + compaction.replacement_item_ids[0], + first.request_item_ids[1] + ); + assert_eq!( + rollout.conversation_items[&compaction.replacement_item_ids[0]].produced_by, + vec![ProducerRef::Compaction { + compaction_id: "compaction-1".to_string() + }], + ); + assert_eq!( + rollout.conversation_items[&compaction.replacement_item_ids[1]].produced_by, + vec![ProducerRef::Compaction { + compaction_id: "compaction-1".to_string() + }], + ); + assert_eq!( + rollout.conversation_items[&compaction.replacement_item_ids[2]].channel, + Some(ConversationChannel::Summary), + ); + assert_eq!( + rollout.conversation_items[&compaction.replacement_item_ids[2]].kind, + ConversationItemKind::Message, + ); + assert_eq!( + rollout.conversation_items[&compaction.replacement_item_ids[2]] + .body + .parts, + vec![ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: "encrypted-summary".to_string(), + }], + ); + + Ok(()) +} + +#[test] +fn context_compaction_boundary_repeats_prefix_and_reuses_replacement_items() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let developer = message("developer", "follow repo rules"); + let user = message("user", "count files"); + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [developer, user] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let summary = message("user", "summary from compacted history"); + let compaction_summary = json!({ + "type": "context_compaction", + "encrypted_content": "encrypted-summary", + }); + let checkpoint = writer.write_json_payload( + RawPayloadKind::CompactionCheckpoint, + &json!({ + "input_history": [developer, user], + "replacement_history": [user, summary, compaction_summary] + }), + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CompactionInstalled { + compaction_id: "compaction-1".to_string(), + checkpoint_payload: checkpoint, + }, + )?; + + start_turn(&writer, "turn-2")?; + let post_compaction_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [developer, user, summary, compaction_summary] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", post_compaction_request)?; + + let rollout = replay_bundle(temp.path())?; + let compaction = &rollout.compactions["compaction-1"]; + + assert_eq!( + rollout.conversation_items[&compaction.replacement_item_ids[2]].channel, + Some(ConversationChannel::Summary), + ); + assert_eq!( + rollout.conversation_items[&compaction.replacement_item_ids[2]].kind, + ConversationItemKind::Message, + ); + assert_eq!( + rollout.conversation_items[&compaction.replacement_item_ids[2]] + .body + .parts, + vec![ConversationPart::Encoded { + label: "encrypted_content".to_string(), + value: "encrypted-summary".to_string(), + }], + ); + + Ok(()) +} + +#[test] +fn tool_call_links_model_call_and_followup_output_items() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "run tests")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [{ + "type": "function_call", + "name": "exec_command", + "arguments": "{\"cmd\":\"cargo test\"}", + "call_id": "call-1" + }] + }), + )?; + append_inference_completion(&writer, "inference-1", "resp-1", response)?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "tool-1".to_string(), + model_visible_call_id: Some("call-1".to_string()), + code_mode_runtime_tool_id: None, + requester: crate::raw_event::RawToolCallRequester::Model, + kind: ToolCallKind::ExecCommand, + summary: ToolCallSummary::Generic { + label: "exec_command".to_string(), + input_preview: Some("cargo test".to_string()), + output_preview: None, + }, + invocation_payload: None, + }, + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallEnded { + tool_call_id: "tool-1".to_string(), + status: ExecutionStatus::Completed, + result_payload: None, + }, + )?; + + start_turn(&writer, "turn-2")?; + let followup = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "previous_response_id": "resp-1", + "input": [{ + "type": "function_call_output", + "call_id": "call-1", + "output": "tests passed" + }] + }), + )?; + append_inference_start(&writer, "inference-2", "turn-2", followup)?; + + let rollout = replay_bundle(temp.path())?; + let first_inference = &rollout.inference_calls["inference-1"]; + let second_inference = &rollout.inference_calls["inference-2"]; + let tool_call = &rollout.tool_calls["tool-1"]; + let output_item_id = second_inference + .request_item_ids + .last() + .expect("follow-up output item"); + + assert_eq!( + first_inference.tool_call_ids_started_by_response, + vec!["tool-1".to_string()], + ); + assert_eq!( + tool_call.model_visible_call_item_ids, + first_inference.response_item_ids, + ); + assert_eq!( + tool_call.model_visible_output_item_ids, + vec![output_item_id.clone()], + ); + assert_eq!( + rollout.conversation_items[output_item_id].produced_by, + vec![ProducerRef::Tool { + tool_call_id: "tool-1".to_string(), + }], + ); + + Ok(()) +} + +#[test] +fn inference_start_rejects_unknown_codex_turn() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "hello")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-missing", request)?; + + expect_replay_error(&temp, "referenced unknown codex turn turn-missing") +} diff --git a/codex-rs/rollout-trace/src/reducer/inference.rs b/codex-rs/rollout-trace/src/reducer/inference.rs new file mode 100644 index 0000000000000000000000000000000000000000..5becb59ca01fc2d188d371457fb53265e5eb93f8 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/inference.rs @@ -0,0 +1,231 @@ +//! Inference call lifecycle reduction. +//! +//! Conversation request/response normalization lives in the conversation module; +//! this module owns the runtime envelope around those model-facing payloads. + +use anyhow::Result; +use anyhow::bail; + +use super::TraceReducer; +use crate::model::ExecutionStatus; +use crate::model::ExecutionWindow; +use crate::model::InferenceCall; +use crate::model::InferenceCallId; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawEventSeq; +use crate::raw_event::RawTraceEventPayload; + +/// Raw inference-start fields after dispatch has stripped the common event envelope. +/// +/// Keeping this as one argument prevents callsites from passing a long list of +/// adjacent strings whose ordering is easy to mix up. +pub(super) struct StartedInferenceCall { + pub(super) inference_call_id: InferenceCallId, + pub(super) thread_id: String, + pub(super) codex_turn_id: String, + pub(super) model: String, + pub(super) provider_name: String, + pub(super) request_payload: RawPayloadRef, +} + +impl TraceReducer { + /// Starts an inference call and reduces its request payload into conversation items. + /// + /// Requests are model-visible transcript evidence, so the inference object is only + /// inserted after the request snapshot has been normalized and linked to the turn. + pub(super) fn start_inference_call( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + started: StartedInferenceCall, + ) -> Result<()> { + if self + .rollout + .inference_calls + .contains_key(&started.inference_call_id) + { + bail!( + "duplicate inference start for {}", + started.inference_call_id + ); + } + + let inference_call_id = started.inference_call_id.clone(); + let thread_id = started.thread_id.clone(); + let codex_turn_id = started.codex_turn_id.clone(); + let request_payload = started.request_payload.clone(); + let Some(turn) = self.rollout.codex_turns.get(&codex_turn_id) else { + bail!( + "inference start {inference_call_id} referenced unknown codex turn {codex_turn_id}" + ); + }; + if turn.thread_id != thread_id { + bail!( + "inference start {inference_call_id} used thread {thread_id}, \ + but codex turn {codex_turn_id} belongs to {}", + turn.thread_id + ); + } + + let request_item_ids = self.reduce_inference_request( + wall_time_unix_ms, + &inference_call_id, + &thread_id, + &codex_turn_id, + &request_payload, + )?; + + self.thread_mut(&thread_id)?; + + self.rollout.inference_calls.insert( + inference_call_id.clone(), + InferenceCall { + inference_call_id, + thread_id, + codex_turn_id, + execution: ExecutionWindow { + started_at_unix_ms: wall_time_unix_ms, + started_seq: seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + model: started.model, + provider_name: started.provider_name, + response_id: None, + upstream_request_id: None, + request_item_ids, + response_item_ids: Vec::new(), + tool_call_ids_started_by_response: Vec::new(), + usage: None, + raw_request_payload_id: started.request_payload.raw_payload_id, + raw_response_payload_id: None, + }, + ); + Ok(()) + } + + /// Closes any inference streams that are still live when the owning turn ends. + /// + /// Normal completion events close the active inference before the turn ends. + /// If a call is still `Running`, Codex stopped observing that provider stream + /// earlier and the reduced graph should not present it as live. + pub(super) fn close_running_inference_calls_for_turn_end( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + codex_turn_id: &str, + turn_status: &ExecutionStatus, + ) { + let inference_status = match turn_status { + ExecutionStatus::Running => return, + ExecutionStatus::Completed | ExecutionStatus::Cancelled => ExecutionStatus::Cancelled, + ExecutionStatus::Failed => ExecutionStatus::Failed, + ExecutionStatus::Aborted => ExecutionStatus::Aborted, + }; + for inference in self.rollout.inference_calls.values_mut() { + if inference.codex_turn_id == codex_turn_id + && inference.execution.status == ExecutionStatus::Running + { + inference.execution.ended_at_unix_ms = Some(wall_time_unix_ms); + inference.execution.ended_seq = Some(seq); + inference.execution.status = inference_status.clone(); + } + } + } + + /// Completes an inference call and, when present, reduces response output items. + pub(super) fn complete_inference_call( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + payload: RawTraceEventPayload, + ) -> Result<()> { + let (inference_call_id, status, response_id, upstream_request_id, response_payload) = + match payload { + RawTraceEventPayload::InferenceCompleted { + inference_call_id, + response_id, + upstream_request_id, + response_payload, + } => ( + inference_call_id, + ExecutionStatus::Completed, + response_id, + upstream_request_id, + Some(response_payload), + ), + RawTraceEventPayload::InferenceFailed { + inference_call_id, + upstream_request_id, + partial_response_payload, + .. + } => ( + inference_call_id, + ExecutionStatus::Failed, + None, + upstream_request_id, + partial_response_payload, + ), + RawTraceEventPayload::InferenceCancelled { + inference_call_id, + upstream_request_id, + partial_response_payload, + .. + } => ( + inference_call_id, + ExecutionStatus::Cancelled, + None, + upstream_request_id, + partial_response_payload, + ), + _ => bail!("complete_inference_call received a non-terminal inference event"), + }; + + if !self + .rollout + .inference_calls + .contains_key(&inference_call_id) + { + bail!("inference completion referenced unknown call {inference_call_id}"); + } + + let response_item_ids = response_payload + .as_ref() + .map(|payload| { + self.reduce_inference_response(wall_time_unix_ms, &inference_call_id, payload) + }) + .transpose()?; + { + let Some(inference) = self.rollout.inference_calls.get_mut(&inference_call_id) else { + bail!("inference call {inference_call_id} disappeared during response reduction"); + }; + inference.response_id = response_id; + // Turn-end cleanup can close a stream before the async mapper observes + // cancellation. Preserve that terminal status while still retaining any + // late partial response evidence from the mapper. + if inference.execution.status == ExecutionStatus::Running { + inference.execution.ended_at_unix_ms = Some(wall_time_unix_ms); + inference.execution.ended_seq = Some(seq); + inference.execution.status = status; + } + // Turn-end cleanup can mark an inference terminal before the stream + // mapper records its late partial payload. Keep the server request + // id from that late payload even when the status is already closed. + if let Some(upstream_request_id) = upstream_request_id { + inference.upstream_request_id = Some(upstream_request_id); + } + if let Some(response_payload) = response_payload { + inference.raw_response_payload_id = Some(response_payload.raw_payload_id); + } + if let Some(response_item_ids) = response_item_ids { + inference.response_item_ids = response_item_ids; + } + } + Ok(()) + } +} + +#[cfg(test)] +#[path = "inference_tests.rs"] +mod tests; diff --git a/codex-rs/rollout-trace/src/reducer/inference_tests.rs b/codex-rs/rollout-trace/src/reducer/inference_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..77e0f39a9a35e0f87f120f9c58103f0ba4228622 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/inference_tests.rs @@ -0,0 +1,158 @@ +use pretty_assertions::assert_eq; +use serde_json::json; +use tempfile::TempDir; + +use crate::model::ConversationItemKind; +use crate::model::ExecutionStatus; +use crate::payload::RawPayloadKind; +use crate::raw_event::RawTraceEventPayload; +use crate::reducer::test_support::append_inference_start; +use crate::reducer::test_support::create_started_writer; +use crate::reducer::test_support::message; +use crate::reducer::test_support::start_turn; +use crate::replay_bundle; + +#[test] +fn cancelled_inference_reduces_partial_response_items() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "draft")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + + let partial_response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": null, + "token_usage": null, + "output_items": [{ + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "partial"}] + }] + }), + )?; + writer.append(RawTraceEventPayload::InferenceCancelled { + inference_call_id: "inference-1".to_string(), + upstream_request_id: Some("req-cancelled".to_string()), + reason: "test interruption".to_string(), + partial_response_payload: Some(partial_response), + })?; + + let rollout = replay_bundle(temp.path())?; + let inference = &rollout.inference_calls["inference-1"]; + let response_item_id = &inference.response_item_ids[0]; + + assert_eq!(inference.execution.status, ExecutionStatus::Cancelled); + assert_eq!( + inference.upstream_request_id, + Some("req-cancelled".to_string()), + ); + assert_eq!(inference.response_item_ids.len(), 1); + assert_eq!( + rollout.conversation_items[response_item_id].kind, + ConversationItemKind::Message, + ); + assert_eq!( + rollout.conversation_items[response_item_id].produced_by, + vec![crate::model::ProducerRef::Inference { + inference_call_id: "inference-1".to_string(), + }], + ); + + Ok(()) +} + +#[test] +fn cancelled_turn_closes_running_inference_call() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "wait")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + let turn_end = writer.append(RawTraceEventPayload::CodexTurnEnded { + codex_turn_id: "turn-1".to_string(), + status: ExecutionStatus::Cancelled, + })?; + + let rollout = replay_bundle(temp.path())?; + let inference = &rollout.inference_calls["inference-1"]; + + assert_eq!(inference.execution.status, ExecutionStatus::Cancelled); + assert_eq!(inference.execution.ended_seq, Some(turn_end.seq)); + + Ok(()) +} + +#[test] +fn late_cancelled_inference_preserves_turn_end_status() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "interrupt")] + }), + )?; + append_inference_start(&writer, "inference-1", "turn-1", request)?; + let turn_end = writer.append(RawTraceEventPayload::CodexTurnEnded { + codex_turn_id: "turn-1".to_string(), + status: ExecutionStatus::Failed, + })?; + + let partial_response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": null, + "token_usage": null, + "output_items": [{ + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "late partial"}] + }] + }), + )?; + writer.append(RawTraceEventPayload::InferenceCancelled { + inference_call_id: "inference-1".to_string(), + upstream_request_id: Some("req-late-cancelled".to_string()), + reason: "stream mapper noticed cancellation after turn end".to_string(), + partial_response_payload: Some(partial_response.clone()), + })?; + + let rollout = replay_bundle(temp.path())?; + let inference = &rollout.inference_calls["inference-1"]; + assert_eq!(inference.execution.status, ExecutionStatus::Failed); + assert_eq!(inference.execution.ended_seq, Some(turn_end.seq)); + assert_eq!( + inference.raw_response_payload_id, + Some(partial_response.raw_payload_id), + ); + assert_eq!( + inference.upstream_request_id, + Some("req-late-cancelled".to_string()), + ); + assert_eq!(inference.response_item_ids.len(), 1); + let response_item_id = &inference.response_item_ids[0]; + assert_eq!( + rollout.conversation_items[response_item_id].body.parts, + vec![crate::model::ConversationPart::Text { + text: "late partial".to_string(), + }], + ); + + Ok(()) +} diff --git a/codex-rs/rollout-trace/src/reducer/mod.rs b/codex-rs/rollout-trace/src/reducer/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..4a908b6e57ed3470db9ff99b73b8c59f01117892 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/mod.rs @@ -0,0 +1,487 @@ +//! Deterministic replay from raw trace events to `RolloutTrace`. + +use std::collections::BTreeMap; +use std::fs::File; +use std::io::BufRead; +use std::io::BufReader; +use std::path::Path; +use std::path::PathBuf; + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use serde_json::Value; + +use crate::bundle::MANIFEST_FILE_NAME; +use crate::bundle::RAW_EVENT_LOG_FILE_NAME; +use crate::bundle::REDUCED_TRACE_SCHEMA_VERSION; +use crate::bundle::TraceBundleManifest; +use crate::model::ExecutionStatus; +use crate::model::RolloutTrace; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawTraceEvent; +use crate::raw_event::RawTraceEventPayload; + +mod code_cell; +mod compaction; +mod conversation; +mod inference; +#[cfg(test)] +pub(crate) mod test_support; +mod thread; +mod tool; + +use self::code_cell::PendingCodeCellLifecycleEvent; +use self::code_cell::PendingCodeCellStart; +use self::code_cell::StartedCodeCell; +use self::compaction::StartedCompactionRequest; +use self::inference::StartedInferenceCall; +use self::tool::ObservedAgentResultEdge; +use self::tool::PendingAgentInteractionEdge; +use self::tool::ToolCallStarted; + +/// Replays a local trace bundle into a reduced rollout graph. +pub fn replay_bundle(bundle_dir: impl AsRef) -> Result { + let bundle_dir = bundle_dir.as_ref(); + let manifest: TraceBundleManifest = + serde_json::from_reader(File::open(bundle_dir.join(MANIFEST_FILE_NAME))?) + .with_context(|| format!("read {}", bundle_dir.join(MANIFEST_FILE_NAME).display()))?; + let mut reducer = TraceReducer { + rollout: RolloutTrace::new( + REDUCED_TRACE_SCHEMA_VERSION, + manifest.trace_id, + manifest.rollout_id, + manifest.root_thread_id, + manifest.started_at_unix_ms, + ), + bundle_dir: bundle_dir.to_path_buf(), + next_conversation_item_ordinal: 1, + next_terminal_operation_ordinal: 1, + thread_conversation_snapshots: BTreeMap::new(), + pending_compaction_replacement_item_ids: BTreeMap::new(), + code_cell_ids_by_runtime: BTreeMap::new(), + pending_code_cell_starts: BTreeMap::new(), + pending_code_cell_lifecycle_events: BTreeMap::new(), + pending_agent_interaction_edges: Vec::new(), + }; + + let event_log_path = bundle_dir.join(RAW_EVENT_LOG_FILE_NAME); + let event_log = File::open(&event_log_path) + .with_context(|| format!("open trace event log {}", event_log_path.display()))?; + for (line_index, line) in BufReader::new(event_log).lines().enumerate() { + let line = line.with_context(|| format!("read trace event line {}", line_index + 1))?; + if line.trim().is_empty() { + continue; + } + let event: RawTraceEvent = serde_json::from_str(&line) + .with_context(|| format!("parse trace event line {}", line_index + 1))?; + reducer.apply_event(event)?; + } + // Spawn edges prefer the child task message as their target, but a child can + // fail before that message is ever reduced. Only after replaying the whole + // bundle do we know which spawn deliveries need the child-thread fallback. + reducer.resolve_pending_spawn_edge_fallbacks()?; + + Ok(reducer.rollout) +} + +struct TraceReducer { + rollout: RolloutTrace, + bundle_dir: PathBuf, + next_conversation_item_ordinal: u64, + next_terminal_operation_ordinal: u64, + /// Last model-visible conversation snapshot per thread. + /// + /// Requests and responses both advance this sequence because both are + /// model-facing payloads. Repeated request snapshots reuse item IDs only + /// when the same normalized item appears at the same position; identical + /// content at a new position must remain a distinct conversation item. + thread_conversation_snapshots: BTreeMap>, + /// Replacement snapshot installed by compaction but not yet seen in a sampling request. + /// + /// The first full request after compaction should compare against the installed replacement + /// history, not against the pre-compaction request. That keeps repeated prefix/context messages + /// as fresh post-compaction conversation items while still reusing the summary/replacement + /// items that actually became live history. + pending_compaction_replacement_item_ids: BTreeMap>, + /// Runtime cell ids indexed by thread-local code-mode handle. + /// + /// Reduced `CodeCellId`s are based on the model-visible `exec` call id + /// because that is the durable source identity. Runtime lifecycle, nested + /// tools, and `wait` calls arrive with the runtime-local `cell_id`, so this + /// index is the one intentional bridge between those namespaces. + code_cell_ids_by_runtime: BTreeMap<(String, String), String>, + /// Code-cell starts whose model-visible `custom_tool_call` item has not + /// been reduced yet. + /// + /// Core begins executing tools before the stream-completion hook records + /// the response payload that requested them. Queueing keeps replay strict + /// about eventual source-item ownership without requiring trace producers + /// to reorder runtime events behind inference completion. + pending_code_cell_starts: BTreeMap, + /// Initial/end events that arrived while the matching start was queued. + /// + /// Fast cells can return before the inference response payload that proves + /// the model-visible `exec` source item has been reduced. The start remains + /// queued for ownership validation; these lifecycle events wait with it and + /// are replayed in raw sequence order once the cell materializes. + pending_code_cell_lifecycle_events: BTreeMap>, + /// Multi-agent deliveries whose recipient-side transcript item has not been observed yet. + /// + /// V2 agent tools enqueue mailbox messages in the target thread. The trace event for the + /// sending tool arrives before the recipient inference request materializes that mailbox item + /// as a `ConversationItem`, so the reducer keeps the delivery edge pending until it can point + /// at the exact model-visible item instead of a coarse thread. + pending_agent_interaction_edges: Vec, +} + +impl TraceReducer { + fn read_payload_json(&self, payload: &RawPayloadRef) -> Result { + // Reducers keep raw bodies out of the graph, but typed replay sometimes + // needs a small subset of fields to build semantic objects. + let payload_path = self.bundle_dir.join(&payload.path); + let file = File::open(&payload_path) + .with_context(|| format!("open payload {}", payload.raw_payload_id))?; + serde_json::from_reader(file) + .with_context(|| format!("parse payload {}", payload.raw_payload_id)) + } + + fn apply_event(&mut self, event: RawTraceEvent) -> Result<()> { + // Raw payload refs are reducer-wide evidence, not owned by a single + // semantic arm. Keep this bookkeeping separate so typed reduction can + // stay strict without duplicating payload insertion in every case. + for payload in event.payload.raw_payload_refs() { + self.insert_raw_payload(payload); + } + + match event.payload { + RawTraceEventPayload::RolloutStarted { + trace_id, + root_thread_id, + } => { + self.rollout.trace_id = trace_id; + self.rollout.root_thread_id = root_thread_id; + } + RawTraceEventPayload::RolloutEnded { status } => { + self.rollout.status = status; + self.rollout.ended_at_unix_ms = Some(event.wall_time_unix_ms); + } + RawTraceEventPayload::ThreadStarted { + thread_id, + agent_path, + metadata_payload, + } => { + self.start_thread( + event.seq, + event.wall_time_unix_ms, + thread_id, + agent_path, + metadata_payload, + )?; + } + RawTraceEventPayload::ThreadEnded { thread_id, status } => { + self.end_thread(event.seq, event.wall_time_unix_ms, thread_id, status)?; + } + RawTraceEventPayload::CodexTurnStarted { + codex_turn_id, + thread_id, + } => { + self.start_codex_turn( + event.seq, + event.wall_time_unix_ms, + codex_turn_id, + thread_id, + )?; + } + RawTraceEventPayload::CodexTurnEnded { + codex_turn_id, + status, + } => { + self.end_codex_turn( + event.seq, + event.wall_time_unix_ms, + event.thread_id, + codex_turn_id, + status, + )?; + } + RawTraceEventPayload::InferenceStarted { + inference_call_id, + thread_id, + codex_turn_id, + model, + provider_name, + request_payload, + } => { + self.start_inference_call( + event.seq, + event.wall_time_unix_ms, + StartedInferenceCall { + inference_call_id, + thread_id, + codex_turn_id, + model, + provider_name, + request_payload, + }, + )?; + } + payload @ (RawTraceEventPayload::InferenceCompleted { .. } + | RawTraceEventPayload::InferenceFailed { .. } + | RawTraceEventPayload::InferenceCancelled { .. }) => { + self.complete_inference_call(event.seq, event.wall_time_unix_ms, payload)?; + } + RawTraceEventPayload::ProtocolEventObserved { .. } => { + // Protocol wrappers are raw debug breadcrumbs. Typed hooks own + // the reduced graph, so these payload refs are retained without + // creating semantic objects. + } + RawTraceEventPayload::ToolCallStarted { + tool_call_id, + model_visible_call_id, + code_mode_runtime_tool_id, + requester, + kind, + summary, + invocation_payload, + } => { + self.start_tool_call( + event.seq, + event.wall_time_unix_ms, + event.thread_id, + event.codex_turn_id, + ToolCallStarted { + tool_call_id, + model_visible_call_id, + code_mode_runtime_tool_id, + requester, + kind, + summary, + invocation_payload, + }, + )?; + } + RawTraceEventPayload::McpToolCallCorrelationAssigned { + tool_call_id, + mcp_call_id, + } => { + self.assign_mcp_tool_call_correlation(tool_call_id, mcp_call_id)?; + } + RawTraceEventPayload::ToolCallRuntimeStarted { + tool_call_id, + runtime_payload, + } => { + self.start_tool_runtime_observation( + event.seq, + event.wall_time_unix_ms, + tool_call_id, + runtime_payload, + )?; + } + RawTraceEventPayload::ToolCallRuntimeEnded { + tool_call_id, + status, + runtime_payload, + } => { + self.end_tool_runtime_observation( + event.seq, + event.wall_time_unix_ms, + tool_call_id, + status, + runtime_payload, + )?; + } + RawTraceEventPayload::ToolCallEnded { + tool_call_id, + status, + result_payload, + } => { + self.end_tool_call( + event.seq, + event.wall_time_unix_ms, + tool_call_id, + status, + result_payload, + )?; + } + RawTraceEventPayload::CodeCellStarted { + runtime_cell_id, + model_visible_call_id, + source_js, + } => { + let thread_id = self.code_cell_event_thread_id( + event.thread_id, + event.codex_turn_id.as_deref(), + &runtime_cell_id, + "code cell start", + )?; + let reduced_code_cell_id = + self.reduced_code_cell_id_for_model_visible_call(&model_visible_call_id); + self.record_runtime_code_cell_id( + &thread_id, + &runtime_cell_id, + &reduced_code_cell_id, + )?; + self.start_or_queue_code_cell(PendingCodeCellStart { + seq: event.seq, + wall_time_unix_ms: event.wall_time_unix_ms, + thread_id, + codex_turn_id: event.codex_turn_id, + started: StartedCodeCell { + code_cell_id: reduced_code_cell_id, + runtime_cell_id, + model_visible_call_id, + source_js, + }, + })?; + } + RawTraceEventPayload::CodeCellInitialResponse { + runtime_cell_id, + status, + .. + } => { + let thread_id = self.code_cell_event_thread_id( + event.thread_id, + event.codex_turn_id.as_deref(), + &runtime_cell_id, + "code cell initial response", + )?; + let code_cell_id = self.code_cell_id_for_runtime_cell_id( + &thread_id, + &runtime_cell_id, + "code cell initial response", + )?; + self.record_or_queue_code_cell_initial_response( + event.seq, + event.wall_time_unix_ms, + code_cell_id, + runtime_cell_id, + status, + )?; + } + RawTraceEventPayload::CodeCellEnded { + runtime_cell_id, + status, + .. + } => { + let thread_id = self.code_cell_event_thread_id( + event.thread_id, + event.codex_turn_id.as_deref(), + &runtime_cell_id, + "code cell end", + )?; + let code_cell_id = self.code_cell_id_for_runtime_cell_id( + &thread_id, + &runtime_cell_id, + "code cell end", + )?; + self.end_or_queue_code_cell( + event.seq, + event.wall_time_unix_ms, + code_cell_id, + status, + )?; + } + RawTraceEventPayload::CompactionRequestStarted { + compaction_id, + compaction_request_id, + thread_id, + codex_turn_id, + model, + provider_name, + request_payload, + } => { + self.start_compaction_request( + event.seq, + event.wall_time_unix_ms, + StartedCompactionRequest { + compaction_id, + compaction_request_id, + thread_id, + codex_turn_id, + model, + provider_name, + request_payload, + }, + )?; + } + RawTraceEventPayload::CompactionRequestCompleted { + compaction_id, + compaction_request_id, + response_payload, + } => { + self.complete_compaction_request( + event.seq, + event.wall_time_unix_ms, + compaction_id, + compaction_request_id, + ExecutionStatus::Completed, + Some(response_payload), + )?; + } + RawTraceEventPayload::CompactionRequestFailed { + compaction_id, + compaction_request_id, + .. + } => { + self.complete_compaction_request( + event.seq, + event.wall_time_unix_ms, + compaction_id, + compaction_request_id, + ExecutionStatus::Failed, + /*response_payload*/ None, + )?; + } + RawTraceEventPayload::CompactionInstalled { + compaction_id, + checkpoint_payload, + } => { + let Some(thread_id) = event.thread_id else { + bail!("compaction installed event {compaction_id} did not include a thread id"); + }; + let Some(codex_turn_id) = event.codex_turn_id else { + bail!( + "compaction installed event {compaction_id} did not include a codex turn id" + ); + }; + self.reduce_compaction_installed_event( + event.wall_time_unix_ms, + thread_id, + codex_turn_id, + compaction_id, + checkpoint_payload, + )?; + } + RawTraceEventPayload::AgentResultObserved { + edge_id, + child_thread_id, + child_codex_turn_id, + parent_thread_id, + message, + carried_payload, + } => { + self.queue_agent_result_interaction_edge(ObservedAgentResultEdge { + wall_time_unix_ms: event.wall_time_unix_ms, + edge_id, + child_thread_id, + child_codex_turn_id, + parent_thread_id, + message, + carried_payload, + })?; + } + RawTraceEventPayload::Other { .. } => { + bail!("raw trace event has no reducer implementation"); + } + } + + Ok(()) + } + + fn insert_raw_payload(&mut self, payload: &RawPayloadRef) { + self.rollout + .raw_payloads + .insert(payload.raw_payload_id.clone(), payload.clone()); + } +} diff --git a/codex-rs/rollout-trace/src/reducer/test_support.rs b/codex-rs/rollout-trace/src/reducer/test_support.rs new file mode 100644 index 0000000000000000000000000000000000000000..bd12e9d6522d0e46b1a0d2a63fb13393bb038d32 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/test_support.rs @@ -0,0 +1,202 @@ +//! Shared reducer test fixtures. +//! +//! These helpers only write common trace scaffolding. Scenario-specific event +//! sequences stay in each test so the behavior under test remains visible. + +use serde_json::json; +use tempfile::TempDir; + +use crate::model::ToolCallSummary; +use crate::payload::RawPayloadKind; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawTraceEventContext; +use crate::raw_event::RawTraceEventPayload; +use crate::replay_bundle; +use crate::writer::TraceWriter; + +pub(crate) const ROOT_THREAD_ID: &str = "thread-root"; +pub(crate) const AGENT_ROOT_THREAD_ID: &str = "019d0000-0000-7000-8000-000000000001"; + +pub(crate) fn message(role: &str, text: &str) -> serde_json::Value { + json!({ + "type": "message", + "role": role, + "content": [{"type": "input_text", "text": text}] + }) +} + +pub(crate) fn generic_summary(label: &str) -> ToolCallSummary { + ToolCallSummary::Generic { + label: label.to_string(), + input_preview: None, + output_preview: None, + } +} + +pub(crate) fn create_started_writer(temp: &TempDir) -> anyhow::Result { + create_started_writer_for_thread(temp, ROOT_THREAD_ID, "/root") +} + +pub(crate) fn create_started_agent_writer(temp: &TempDir) -> anyhow::Result { + create_started_writer_for_thread(temp, AGENT_ROOT_THREAD_ID, "/root") +} + +pub(crate) fn create_started_writer_for_thread( + temp: &TempDir, + thread_id: &str, + agent_path: &str, +) -> anyhow::Result { + let writer = TraceWriter::create( + temp.path(), + "trace-1".to_string(), + "rollout-1".to_string(), + thread_id.to_string(), + )?; + start_thread(&writer, thread_id, agent_path)?; + Ok(writer) +} + +pub(crate) fn start_thread( + writer: &TraceWriter, + thread_id: &str, + agent_path: &str, +) -> anyhow::Result<()> { + writer.append(RawTraceEventPayload::ThreadStarted { + thread_id: thread_id.to_string(), + agent_path: agent_path.to_string(), + metadata_payload: None, + })?; + Ok(()) +} + +pub(crate) fn start_turn(writer: &TraceWriter, turn_id: &str) -> anyhow::Result<()> { + start_turn_for_thread(writer, ROOT_THREAD_ID, turn_id) +} + +pub(crate) fn start_agent_turn(writer: &TraceWriter, turn_id: &str) -> anyhow::Result<()> { + start_turn_for_thread(writer, AGENT_ROOT_THREAD_ID, turn_id) +} + +pub(crate) fn start_turn_for_thread( + writer: &TraceWriter, + thread_id: &str, + turn_id: &str, +) -> anyhow::Result<()> { + writer.append(RawTraceEventPayload::CodexTurnStarted { + codex_turn_id: turn_id.to_string(), + thread_id: thread_id.to_string(), + })?; + Ok(()) +} + +pub(crate) fn trace_context(turn_id: &str) -> RawTraceEventContext { + trace_context_for_thread(ROOT_THREAD_ID, turn_id) +} + +pub(crate) fn trace_context_for_agent(turn_id: &str) -> RawTraceEventContext { + trace_context_for_thread(AGENT_ROOT_THREAD_ID, turn_id) +} + +pub(crate) fn trace_context_for_thread(thread_id: &str, turn_id: &str) -> RawTraceEventContext { + RawTraceEventContext { + thread_id: Some(thread_id.to_string()), + codex_turn_id: Some(turn_id.to_string()), + } +} + +pub(crate) fn append_inference_start( + writer: &TraceWriter, + inference_call_id: &str, + codex_turn_id: &str, + request_payload: RawPayloadRef, +) -> anyhow::Result<()> { + append_inference_start_for_thread( + writer, + ROOT_THREAD_ID, + codex_turn_id, + inference_call_id, + request_payload, + ) +} + +pub(crate) fn append_inference_start_for_thread( + writer: &TraceWriter, + thread_id: &str, + codex_turn_id: &str, + inference_call_id: &str, + request_payload: RawPayloadRef, +) -> anyhow::Result<()> { + writer.append(RawTraceEventPayload::InferenceStarted { + inference_call_id: inference_call_id.to_string(), + thread_id: thread_id.to_string(), + codex_turn_id: codex_turn_id.to_string(), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + request_payload, + })?; + Ok(()) +} + +pub(crate) fn append_inference_completion( + writer: &TraceWriter, + inference_call_id: &str, + response_id: &str, + response_payload: RawPayloadRef, +) -> anyhow::Result<()> { + writer.append(RawTraceEventPayload::InferenceCompleted { + inference_call_id: inference_call_id.to_string(), + response_id: Some(response_id.to_string()), + upstream_request_id: None, + response_payload, + })?; + Ok(()) +} + +pub(crate) fn append_inference_request( + writer: &TraceWriter, + thread_id: &str, + turn_id: &str, + inference_id: &str, + input: Vec, +) -> anyhow::Result<()> { + let request = + writer.write_json_payload(RawPayloadKind::InferenceRequest, &json!({ "input": input }))?; + append_inference_start_for_thread(writer, thread_id, turn_id, inference_id, request) +} + +pub(crate) fn append_completed_inference( + writer: &TraceWriter, + thread_id: &str, + turn_id: &str, + inference_id: &str, + input: Vec, + output_items: Vec, +) -> anyhow::Result<()> { + append_inference_request(writer, thread_id, turn_id, inference_id, input)?; + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": format!("resp-{inference_id}"), + "output_items": output_items, + }), + )?; + writer.append_with_context( + trace_context_for_thread(thread_id, turn_id), + RawTraceEventPayload::InferenceCompleted { + inference_call_id: inference_id.to_string(), + response_id: Some(format!("resp-{inference_id}")), + upstream_request_id: None, + response_payload: response, + }, + )?; + Ok(()) +} + +pub(crate) fn expect_replay_error(temp: &TempDir, expected: &str) -> anyhow::Result<()> { + let Err(err) = replay_bundle(temp.path()) else { + panic!("expected replay error containing {expected}"); + }; + let message = err.to_string(); + assert!(message.contains(expected), "unexpected error: {message}"); + Ok(()) +} diff --git a/codex-rs/rollout-trace/src/reducer/thread.rs b/codex-rs/rollout-trace/src/reducer/thread.rs new file mode 100644 index 0000000000000000000000000000000000000000..4ef39ee56309a09070aa0dc698b03948eafa9b18 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/thread.rs @@ -0,0 +1,270 @@ +//! Thread and turn reduction. +//! +//! Threads are the container that every other reducer module links into. This +//! module owns the identity metadata parsing as well, so the central dispatcher +//! does not need to know the shape of multi-agent session-source payloads. + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use serde::Deserialize; +use serde_json::Value; + +use super::TraceReducer; +use super::tool::spawn_edge_id; +use crate::model::AgentOrigin; +use crate::model::AgentThread; +use crate::model::CodexTurn; +use crate::model::CodexTurnId; +use crate::model::ExecutionStatus; +use crate::model::ExecutionWindow; +use crate::model::RolloutStatus; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawEventSeq; + +impl TraceReducer { + /// Inserts a thread and derives its multi-agent identity from optional metadata. + /// + /// The raw event carries a denormalized agent path; when v2 subagent metadata is + /// present, that metadata is authoritative because it also drives spawn edges and task names. + pub(super) fn start_thread( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + thread_id: String, + agent_path: String, + metadata_payload: Option, + ) -> Result<()> { + if self.rollout.threads.contains_key(&thread_id) { + bail!("duplicate thread start for {thread_id}"); + } + + let metadata = metadata_payload + .as_ref() + .map(|payload| self.thread_started_metadata(payload)) + .transpose()?; + let spawn = metadata + .as_ref() + .and_then(ThreadStartedMetadata::thread_spawn); + // The v2 SessionSource is the authoritative child identity record. + // Prefer its nested agent_path over the denormalized event field so + // task derivation and the spawn edge are based on the same metadata. + let agent_path = spawn + .as_ref() + .and_then(|spawn| spawn.agent_path.clone()) + .or_else(|| { + metadata + .as_ref() + .and_then(|metadata| metadata.agent_path.clone()) + }) + .unwrap_or(agent_path); + let nickname = metadata + .as_ref() + .and_then(|metadata| metadata.nickname.clone()); + let default_model = metadata + .as_ref() + .and_then(|metadata| metadata.model.clone()); + let origin = if let Some(spawn) = spawn { + let edge_id = spawn_edge_id(&spawn.parent_thread_id, &thread_id); + let task_name = spawn + .task_name + .clone() + .unwrap_or_else(|| task_name_from_agent_path(&agent_path)); + let agent_role = spawn.agent_role.clone().unwrap_or_default(); + + AgentOrigin::Spawned { + parent_thread_id: spawn.parent_thread_id, + spawn_edge_id: edge_id, + task_name, + agent_role, + } + } else { + AgentOrigin::Root + }; + + self.rollout.threads.insert( + thread_id.clone(), + AgentThread { + thread_id, + agent_path, + nickname, + origin, + execution: ExecutionWindow { + started_at_unix_ms: wall_time_unix_ms, + started_seq: seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + default_model, + conversation_item_ids: Vec::new(), + }, + ); + Ok(()) + } + + /// Marks a thread terminal without treating child shutdown as rollout completion. + pub(super) fn end_thread( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + thread_id: String, + status: RolloutStatus, + ) -> Result<()> { + let thread = self.thread_mut(&thread_id)?; + thread.execution.ended_at_unix_ms = Some(wall_time_unix_ms); + thread.execution.ended_seq = Some(seq); + thread.execution.status = match status { + RolloutStatus::Running => ExecutionStatus::Running, + RolloutStatus::Completed => ExecutionStatus::Completed, + RolloutStatus::Failed => ExecutionStatus::Failed, + RolloutStatus::Aborted => ExecutionStatus::Aborted, + }; + Ok(()) + } + + /// Starts a Codex turn inside an existing thread. + pub(super) fn start_codex_turn( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + codex_turn_id: CodexTurnId, + thread_id: String, + ) -> Result<()> { + if self.rollout.codex_turns.contains_key(&codex_turn_id) { + bail!("duplicate codex turn start for {codex_turn_id}"); + } + + self.thread_mut(&thread_id)?; + + self.rollout.codex_turns.insert( + codex_turn_id.clone(), + CodexTurn { + codex_turn_id, + thread_id, + execution: ExecutionWindow { + started_at_unix_ms: wall_time_unix_ms, + started_seq: seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + input_item_ids: Vec::new(), + }, + ); + Ok(()) + } + + /// Marks a Codex turn terminal and validates any thread id carried by the raw event. + pub(super) fn end_codex_turn( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + thread_id: Option, + codex_turn_id: CodexTurnId, + status: ExecutionStatus, + ) -> Result<()> { + if let Some(event_thread_id) = thread_id.as_deref() + && let Some(turn) = self.rollout.codex_turns.get(&codex_turn_id) + && turn.thread_id != event_thread_id + { + bail!( + "codex turn end for {codex_turn_id} used thread {event_thread_id}, \ + but the turn belongs to {}", + turn.thread_id + ); + } + + let Some(turn) = self.rollout.codex_turns.get_mut(&codex_turn_id) else { + bail!("codex turn end referenced unknown turn {codex_turn_id}"); + }; + turn.execution.ended_at_unix_ms = Some(wall_time_unix_ms); + turn.execution.ended_seq = Some(seq); + turn.execution.status = status.clone(); + self.terminate_running_code_cells_for_turn_end( + seq, + wall_time_unix_ms, + &codex_turn_id, + &status, + )?; + self.close_running_inference_calls_for_turn_end( + seq, + wall_time_unix_ms, + &codex_turn_id, + &status, + ); + Ok(()) + } + + /// Returns a mutable thread or reports a reducer error tied to the unknown id. + pub(super) fn thread_mut(&mut self, thread_id: &str) -> Result<&mut AgentThread> { + self.rollout + .threads + .get_mut(thread_id) + .with_context(|| format!("trace event referenced unknown thread {thread_id}")) + } + + fn thread_started_metadata( + &self, + metadata_payload: &RawPayloadRef, + ) -> Result { + let value = self.read_payload_json(metadata_payload)?; + serde_json::from_value(value) + .with_context(|| format!("parse thread metadata {}", metadata_payload.raw_payload_id)) + } +} + +#[derive(Deserialize)] +struct ThreadStartedMetadata { + agent_path: Option, + task_name: Option, + nickname: Option, + agent_role: Option, + model: Option, + session_source: Option, +} + +impl ThreadStartedMetadata { + fn thread_spawn(&self) -> Option { + let spawn = self + .session_source + .as_ref()? + .get("subagent")? + .get("thread_spawn")?; + let agent_path = spawn + .get("agent_path") + .and_then(Value::as_str) + .map(str::to_string) + .or_else(|| self.agent_path.clone()); + Some(ThreadSpawnMetadata { + parent_thread_id: spawn.get("parent_thread_id")?.as_str()?.to_string(), + agent_path: agent_path.clone(), + task_name: spawn + .get("task_name") + .and_then(Value::as_str) + .map(str::to_string) + .or_else(|| self.task_name.clone()) + .or_else(|| agent_path.as_deref().map(task_name_from_agent_path)), + agent_role: spawn + .get("agent_role") + .and_then(Value::as_str) + .map(str::to_string) + .or_else(|| self.agent_role.clone()), + }) + } +} + +struct ThreadSpawnMetadata { + parent_thread_id: String, + agent_path: Option, + task_name: Option, + agent_role: Option, +} + +fn task_name_from_agent_path(agent_path: &str) -> String { + agent_path + .rsplit('/') + .find(|segment| !segment.is_empty()) + .unwrap_or(agent_path) + .to_string() +} diff --git a/codex-rs/rollout-trace/src/reducer/tool.rs b/codex-rs/rollout-trace/src/reducer/tool.rs new file mode 100644 index 0000000000000000000000000000000000000000..43845ea93918627f8919322b3823038bc10819fb --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/tool.rs @@ -0,0 +1,517 @@ +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; + +use super::TraceReducer; +use crate::model::CodeModeRuntimeToolId; +use crate::model::ConversationItemKind; +use crate::model::ExecutionStatus; +use crate::model::ExecutionWindow; +use crate::model::McpCallId; +use crate::model::ModelVisibleCallId; +use crate::model::ProducerRef; +use crate::model::ToolCall; +use crate::model::ToolCallId; +use crate::model::ToolCallKind; +use crate::model::ToolCallSummary; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawEventSeq; +use crate::raw_event::RawToolCallRequester; + +mod agents; +mod terminal; + +pub(super) use agents::ObservedAgentResultEdge; +pub(super) use agents::PendingAgentInteractionEdge; +pub(super) use agents::spawn_edge_id; + +/// Raw tool-start fields after dispatch has stripped the common event envelope. +/// +/// Tool starts carry several optional identity namespaces: model-visible calls, +/// code-mode runtime tools, and canonical invocation payloads. Grouping them keeps +/// the reducer callsite readable and avoids positional argument mistakes. +pub(super) struct ToolCallStarted { + pub(super) tool_call_id: ToolCallId, + pub(super) model_visible_call_id: Option, + pub(super) code_mode_runtime_tool_id: Option, + pub(super) requester: RawToolCallRequester, + pub(super) kind: ToolCallKind, + pub(super) summary: ToolCallSummary, + pub(super) invocation_payload: Option, +} + +impl TraceReducer { + /// Starts a tool call and links it to model-visible items or runtime parents when available. + /// + /// Some tools also create richer domain objects, such as terminal operations, from + /// the same invocation payload. The generic ToolCall remains the common index. + pub(super) fn start_tool_call( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + thread_id: Option, + codex_turn_id: Option, + started: ToolCallStarted, + ) -> Result<()> { + let tool_call_id = started.tool_call_id.clone(); + if self.rollout.tool_calls.contains_key(&tool_call_id) { + bail!("duplicate tool call start for {tool_call_id}"); + } + self.ensure_unique_model_visible_tool_call( + started.model_visible_call_id.as_deref(), + &tool_call_id, + )?; + + let thread_id = self.tool_thread_id(thread_id, codex_turn_id.as_deref())?; + self.validate_tool_turn(&thread_id, codex_turn_id.as_deref())?; + + let model_visible_call_id = started.model_visible_call_id.clone(); + let requester = self.reduce_tool_call_requester(&thread_id, started.requester.clone())?; + let model_visible_call_item_ids = model_visible_call_id + .as_deref() + .map(|call_id| { + self.model_visible_tool_item_ids( + &thread_id, + call_id, + &[ + ConversationItemKind::FunctionCall, + ConversationItemKind::CustomToolCall, + ], + ) + }) + .unwrap_or_default(); + let model_visible_output_item_ids = model_visible_call_id + .as_deref() + .map(|call_id| { + self.model_visible_tool_item_ids( + &thread_id, + call_id, + &[ + ConversationItemKind::FunctionCallOutput, + ConversationItemKind::CustomToolCallOutput, + ], + ) + }) + .unwrap_or_default(); + + self.thread_mut(&thread_id)?; + + // Some terminal-like tools, notably write_stdin, do not emit a richer + // runtime begin event. For those tools the canonical invocation is the + // only place to recover the terminal/session join key. + let terminal_operation_id = self.start_terminal_operation_from_invocation( + seq, + wall_time_unix_ms, + &thread_id, + &tool_call_id, + &started.kind, + started.invocation_payload.as_ref(), + )?; + // Terminal-backed tools should render through the richer terminal + // operation instead of the generic tool summary captured by producers. + let summary = terminal_operation_id + .as_ref() + .map(|operation_id| ToolCallSummary::Terminal { + operation_id: operation_id.clone(), + }) + .unwrap_or(started.summary); + let raw_invocation_payload_id = started + .invocation_payload + .as_ref() + .map(|payload| payload.raw_payload_id.clone()); + self.link_wait_tool_call_from_request_payload( + &thread_id, + &tool_call_id, + started.invocation_payload.as_ref(), + )?; + + self.rollout.tool_calls.insert( + tool_call_id.clone(), + ToolCall { + tool_call_id: tool_call_id.clone(), + mcp_call_id: None, + model_visible_call_id, + code_mode_runtime_tool_id: started.code_mode_runtime_tool_id, + thread_id, + started_by_codex_turn_id: codex_turn_id, + execution: ExecutionWindow { + started_at_unix_ms: wall_time_unix_ms, + started_seq: seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + requester: requester.clone(), + kind: started.kind, + model_visible_call_item_ids, + model_visible_output_item_ids: Vec::new(), + terminal_operation_id, + summary, + raw_invocation_payload_id, + raw_result_payload_id: None, + raw_runtime_payload_ids: Vec::new(), + }, + ); + + self.link_tool_call_to_code_cell(&tool_call_id, &requester)?; + self.link_tool_to_inference_response(&tool_call_id); + // Output items need the reverse ProducerRef edge as well, so attach + // them after insertion through the same helper used by the transcript + // reducer when the output is observed after the tool start. + for item_id in model_visible_output_item_ids { + self.add_tool_output_item(&tool_call_id, &item_id)?; + } + // The call/output items may have been observed before this tool start. + // Re-sync after insertion so terminal observations get both directions + // of the model-visible link. + self.sync_terminal_model_observation(&tool_call_id)?; + Ok(()) + } + + /// Attaches the bridge-visible MCP UUID after the generic tool call exists. + pub(super) fn assign_mcp_tool_call_correlation( + &mut self, + tool_call_id: ToolCallId, + mcp_call_id: McpCallId, + ) -> Result<()> { + let Some(tool_call) = self.rollout.tool_calls.get_mut(&tool_call_id) else { + bail!("MCP correlation referenced unknown tool call {tool_call_id}"); + }; + if tool_call.mcp_call_id.replace(mcp_call_id).is_some() { + bail!("duplicate MCP correlation for tool call {tool_call_id}"); + } + Ok(()) + } + + /// Completes the canonical tool call and any terminal operation driven by dispatch output. + /// + /// Protocol-backed terminal tools end from runtime events; direct tools + /// may only have the canonical result payload, so this method handles both paths. + pub(super) fn end_tool_call( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + tool_call_id: ToolCallId, + status: ExecutionStatus, + result_payload: Option, + ) -> Result<()> { + let (terminal_operation_id, thread_id, end_terminal_from_result) = { + let Some(tool_call) = self.rollout.tool_calls.get_mut(&tool_call_id) else { + bail!("tool call end referenced unknown call {tool_call_id}"); + }; + tool_call.execution.ended_at_unix_ms = Some(wall_time_unix_ms); + tool_call.execution.ended_seq = Some(seq); + tool_call.execution.status = status.clone(); + tool_call.raw_result_payload_id = result_payload + .as_ref() + .map(|payload| payload.raw_payload_id.clone()); + ( + tool_call.terminal_operation_id.clone(), + tool_call.thread_id.clone(), + // Protocol-backed tools end terminal operations from + // runtime observations. Dispatch result payloads are still kept + // on ToolCall, but they are caller-facing and may be transformed + // relative to the raw terminal output. + tool_call.raw_runtime_payload_ids.is_empty(), + ) + }; + if end_terminal_from_result && let Some(operation_id) = terminal_operation_id { + self.end_terminal_operation( + seq, + wall_time_unix_ms, + &thread_id, + &operation_id, + status, + result_payload.as_ref(), + )?; + } + self.attach_agent_interaction_tool_result(&tool_call_id, result_payload.as_ref())?; + Ok(()) + } + + /// Records a runtime-begin observation for an already started tool call. + /// + /// Runtime observations enrich the generic tool with protocol facts and may + /// create domain-specific children such as terminal operations or agent edges. + pub(super) fn start_tool_runtime_observation( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + tool_call_id: ToolCallId, + runtime_payload: RawPayloadRef, + ) -> Result<()> { + let (thread_id, _requester, kind, existing_terminal_operation_id) = { + let Some(tool_call) = self.rollout.tool_calls.get_mut(&tool_call_id) else { + bail!("tool runtime start referenced unknown call {tool_call_id}"); + }; + push_unique( + &mut tool_call.raw_runtime_payload_ids, + &runtime_payload.raw_payload_id, + ); + ( + tool_call.thread_id.clone(), + tool_call.requester.clone(), + tool_call.kind.clone(), + tool_call.terminal_operation_id.clone(), + ) + }; + if existing_terminal_operation_id.is_some() + && matches!(kind, ToolCallKind::ExecCommand | ToolCallKind::WriteStdin) + { + bail!("tool runtime start would create a second terminal operation for {tool_call_id}"); + } + + // Protocol begin events carry runtime facts such as process ids and + // cwd. These facts should create terminal rows, but they must not + // replace the canonical invocation payload captured at dispatch. + let terminal_operation_id = self.start_terminal_operation_from_runtime( + seq, + wall_time_unix_ms, + &thread_id, + &tool_call_id, + &kind, + &runtime_payload, + )?; + + if let Some(operation_id) = &terminal_operation_id { + let Some(tool_call) = self.rollout.tool_calls.get_mut(&tool_call_id) else { + bail!("tool call {tool_call_id} disappeared during runtime start reduction"); + }; + if tool_call.terminal_operation_id.is_none() { + tool_call.terminal_operation_id = Some(operation_id.clone()); + tool_call.summary = ToolCallSummary::Terminal { + operation_id: operation_id.clone(), + }; + } + } + + if terminal_operation_id.is_some() { + self.sync_terminal_model_observation(&tool_call_id)?; + } + self.start_agent_interaction_from_runtime(&tool_call_id, &runtime_payload)?; + Ok(()) + } + + /// Records a runtime-end observation for an already started tool call. + pub(super) fn end_tool_runtime_observation( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + tool_call_id: ToolCallId, + status: ExecutionStatus, + runtime_payload: RawPayloadRef, + ) -> Result<()> { + let (thread_id, terminal_operation_id) = { + let Some(tool_call) = self.rollout.tool_calls.get_mut(&tool_call_id) else { + bail!("tool runtime end referenced unknown call {tool_call_id}"); + }; + push_unique( + &mut tool_call.raw_runtime_payload_ids, + &runtime_payload.raw_payload_id, + ); + ( + tool_call.thread_id.clone(), + tool_call.terminal_operation_id.clone(), + ) + }; + + if let Some(operation_id) = terminal_operation_id { + self.end_terminal_operation( + seq, + wall_time_unix_ms, + &thread_id, + &operation_id, + status, + Some(&runtime_payload), + )?; + } + self.end_agent_interaction_from_runtime( + wall_time_unix_ms, + &tool_call_id, + &runtime_payload, + )?; + Ok(()) + } + + /// Attaches a conversation item observed after the tool call was reduced. + /// + /// Inference request/response ordering can expose call/output items after the + /// runtime tool object exists, so transcript reduction calls back here to add + /// reverse links without duplicating matching logic. + pub(super) fn attach_model_visible_tool_item( + &mut self, + item_id: &str, + call_id: Option<&str>, + kind: &ConversationItemKind, + ) -> Result<()> { + let Some(call_id) = call_id else { + return Ok(()); + }; + match kind { + ConversationItemKind::FunctionCall | ConversationItemKind::CustomToolCall => { + if let Some(tool_call_id) = self.single_tool_for_model_visible_call(call_id)? { + self.add_tool_call_item(&tool_call_id, item_id)?; + self.link_tool_to_inference_response(&tool_call_id); + self.sync_terminal_model_observation(&tool_call_id)?; + } + } + ConversationItemKind::FunctionCallOutput + | ConversationItemKind::CustomToolCallOutput => { + if let Some(tool_call_id) = self.single_tool_for_model_visible_call(call_id)? { + self.add_tool_output_item(&tool_call_id, item_id)?; + self.sync_terminal_model_observation(&tool_call_id)?; + } + } + ConversationItemKind::Message + | ConversationItemKind::Reasoning + | ConversationItemKind::CompactionMarker => {} + } + Ok(()) + } + + fn tool_thread_id( + &self, + thread_id: Option, + codex_turn_id: Option<&str>, + ) -> Result { + if let Some(thread_id) = thread_id { + return Ok(thread_id); + } + let Some(codex_turn_id) = codex_turn_id else { + bail!("tool call start did not include thread or Codex turn context"); + }; + self.rollout + .codex_turns + .get(codex_turn_id) + .map(|turn| turn.thread_id.clone()) + .with_context(|| { + format!("tool call start referenced unknown Codex turn {codex_turn_id}") + }) + } + + fn validate_tool_turn(&self, thread_id: &str, codex_turn_id: Option<&str>) -> Result<()> { + if !self.rollout.threads.contains_key(thread_id) { + bail!("tool call start referenced unknown thread {thread_id}"); + } + if let Some(codex_turn_id) = codex_turn_id { + let Some(turn) = self.rollout.codex_turns.get(codex_turn_id) else { + bail!("tool call start referenced unknown Codex turn {codex_turn_id}"); + }; + if turn.thread_id != thread_id { + bail!( + "tool call start used thread {thread_id}, but Codex turn {codex_turn_id} \ + belongs to {}", + turn.thread_id + ); + } + } + Ok(()) + } + + fn ensure_unique_model_visible_tool_call( + &self, + model_visible_call_id: Option<&str>, + tool_call_id: &str, + ) -> Result<()> { + let Some(model_visible_call_id) = model_visible_call_id else { + return Ok(()); + }; + if let Some(existing) = self.single_tool_for_model_visible_call(model_visible_call_id)? + && existing != tool_call_id + { + bail!("duplicate tool call for model-visible call id {model_visible_call_id}"); + } + Ok(()) + } + + fn single_tool_for_model_visible_call( + &self, + model_visible_call_id: &str, + ) -> Result> { + let mut matching = self + .rollout + .tool_calls + .values() + .filter(|tool| tool.model_visible_call_id.as_deref() == Some(model_visible_call_id)) + .map(|tool| tool.tool_call_id.clone()); + let first = matching.next(); + if matching.next().is_some() { + bail!("multiple tool calls matched model-visible call id {model_visible_call_id}"); + } + Ok(first) + } + + fn model_visible_tool_item_ids( + &self, + thread_id: &str, + call_id: &str, + kinds: &[ConversationItemKind], + ) -> Vec { + self.rollout + .conversation_items + .values() + .filter(|item| { + item.thread_id == thread_id + && item.call_id.as_deref() == Some(call_id) + && kinds.contains(&item.kind) + }) + .map(|item| item.item_id.clone()) + .collect::>() + } + + fn add_tool_call_item(&mut self, tool_call_id: &str, item_id: &str) -> Result<()> { + let Some(tool_call) = self.rollout.tool_calls.get_mut(tool_call_id) else { + bail!("tool call {tool_call_id} disappeared during conversation linking"); + }; + push_unique(&mut tool_call.model_visible_call_item_ids, item_id); + Ok(()) + } + + fn add_tool_output_item(&mut self, tool_call_id: &str, item_id: &str) -> Result<()> { + let Some(tool_call) = self.rollout.tool_calls.get_mut(tool_call_id) else { + bail!("tool call {tool_call_id} disappeared during output linking"); + }; + push_unique(&mut tool_call.model_visible_output_item_ids, item_id); + + let Some(item) = self.rollout.conversation_items.get_mut(item_id) else { + bail!("conversation item {item_id} disappeared during output linking"); + }; + let producer = ProducerRef::Tool { + tool_call_id: tool_call_id.to_string(), + }; + if !item.produced_by.contains(&producer) { + item.produced_by.push(producer); + } + Ok(()) + } + + fn link_tool_to_inference_response(&mut self, tool_call_id: &str) { + let Some(tool_call) = self.rollout.tool_calls.get(tool_call_id) else { + return; + }; + let call_item_ids = tool_call.model_visible_call_item_ids.clone(); + if call_item_ids.is_empty() { + return; + } + for inference in self.rollout.inference_calls.values_mut() { + if inference + .response_item_ids + .iter() + .any(|item_id| call_item_ids.contains(item_id)) + && !inference + .tool_call_ids_started_by_response + .contains(&tool_call_id.to_string()) + { + inference + .tool_call_ids_started_by_response + .push(tool_call_id.to_string()); + } + } + } +} + +fn push_unique(items: &mut Vec, item_id: &str) { + if !items.iter().any(|existing| existing == item_id) { + items.push(item_id.to_string()); + } +} diff --git a/codex-rs/rollout-trace/src/reducer/tool/agents.rs b/codex-rs/rollout-trace/src/reducer/tool/agents.rs new file mode 100644 index 0000000000000000000000000000000000000000..7c71ddef7d350c377877ec9cd91c3adf346617eb --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/tool/agents.rs @@ -0,0 +1,810 @@ +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use codex_protocol::protocol::CollabAgentInteractionBeginEvent; +use codex_protocol::protocol::CollabAgentInteractionEndEvent; +use codex_protocol::protocol::CollabAgentSpawnEndEvent; +use codex_protocol::protocol::CollabCloseBeginEvent; +use codex_protocol::protocol::CollabCloseEndEvent; +use codex_protocol::protocol::InterAgentCommunication; +use codex_protocol::protocol::SubAgentActivityEvent; +use codex_protocol::protocol::SubAgentActivityKind; +use serde::Deserialize; +use serde_json::Value; + +use super::super::TraceReducer; +use crate::model::ConversationItem; +use crate::model::ConversationItemKind; +use crate::model::ConversationPart; +use crate::model::ConversationRole; +use crate::model::InteractionEdge; +use crate::model::InteractionEdgeKind; +use crate::model::ToolCallKind; +use crate::model::TraceAnchor; +use crate::payload::RawPayloadRef; + +/// Agent delivery edge waiting for the recipient-side conversation item. +/// +/// Multi-agent v2 records the sender tool before the target thread necessarily +/// includes the delivered mailbox message in a model-visible request. The edge +/// stays pending so it can target that exact conversation item when possible. +pub(in crate::reducer) struct PendingAgentInteractionEdge { + pub(in crate::reducer) edge_id: String, + pub(in crate::reducer) kind: InteractionEdgeKind, + pub(in crate::reducer) source: TraceAnchor, + pub(in crate::reducer) target_thread_id: String, + pub(in crate::reducer) message_author: String, + pub(in crate::reducer) message_content: String, + /// Spawn-only fallback for children that fail before their task message is model-visible. + pub(in crate::reducer) unresolved_spawn_thread_id: Option, + pub(in crate::reducer) started_at_unix_ms: i64, + pub(in crate::reducer) ended_at_unix_ms: Option, + pub(in crate::reducer) carried_raw_payload_ids: Vec, +} + +/// Typed reducer input for a multi-agent v2 child completion notification. +/// +/// Child results are observed outside the normal tool lifecycle, but they still +/// carry a parent-thread notification. This wrapper keeps the dispatcher from +/// passing a positional bundle of thread and turn ids. +pub(in crate::reducer) struct ObservedAgentResultEdge { + pub(in crate::reducer) wall_time_unix_ms: i64, + pub(in crate::reducer) edge_id: String, + pub(in crate::reducer) child_thread_id: String, + pub(in crate::reducer) child_codex_turn_id: String, + pub(in crate::reducer) parent_thread_id: String, + pub(in crate::reducer) message: String, + pub(in crate::reducer) carried_payload: Option, +} + +#[derive(Deserialize)] +struct AgentMessageInvocationArgs { + message: String, +} + +/// Builds the stable edge id for the spawn relationship between two threads. +pub(in crate::reducer) fn spawn_edge_id(parent_thread_id: &str, child_thread_id: &str) -> String { + format!("edge:spawn:{parent_thread_id}:{child_thread_id}") +} + +impl TraceReducer { + /// Starts a multi-agent edge from a runtime begin payload, when the tool kind supports one. + pub(super) fn start_agent_interaction_from_runtime( + &mut self, + tool_call_id: &str, + runtime_payload: &RawPayloadRef, + ) -> Result<()> { + let kind = self + .rollout + .tool_calls + .get(tool_call_id) + .with_context(|| format!("agent edge referenced unknown tool call {tool_call_id}"))? + .kind + .clone(); + match kind { + ToolCallKind::AssignAgentTask => { + let payload: CollabAgentInteractionBeginEvent = + serde_json::from_value(self.read_payload_json(runtime_payload)?)?; + self.queue_message_agent_interaction( + tool_call_id, + InteractionEdgeKind::AssignAgentTask, + payload.receiver_thread_id.to_string(), + payload.prompt, + /*ended_at_unix_ms*/ None, + ) + } + ToolCallKind::SendMessage => { + let payload: CollabAgentInteractionBeginEvent = + serde_json::from_value(self.read_payload_json(runtime_payload)?)?; + self.queue_message_agent_interaction( + tool_call_id, + InteractionEdgeKind::SendMessage, + payload.receiver_thread_id.to_string(), + payload.prompt, + /*ended_at_unix_ms*/ None, + ) + } + ToolCallKind::CloseAgent => { + let payload: CollabCloseBeginEvent = + serde_json::from_value(self.read_payload_json(runtime_payload)?)?; + self.upsert_close_agent_interaction( + tool_call_id, + payload.receiver_thread_id.to_string(), + /*ended_at_unix_ms*/ None, + ) + } + ToolCallKind::ExecCommand + | ToolCallKind::WriteStdin + | ToolCallKind::ApplyPatch + | ToolCallKind::Mcp { .. } + | ToolCallKind::Web + | ToolCallKind::ImageGeneration + | ToolCallKind::SpawnAgent + | ToolCallKind::WaitAgent + | ToolCallKind::Other { .. } => Ok(()), + } + } + + /// Ends or enriches a multi-agent edge from a runtime end payload. + pub(super) fn end_agent_interaction_from_runtime( + &mut self, + wall_time_unix_ms: i64, + tool_call_id: &str, + runtime_payload: &RawPayloadRef, + ) -> Result<()> { + let kind = self.rollout.tool_calls[tool_call_id].kind.clone(); + let runtime_payload_json = self.read_payload_json(runtime_payload)?; + if runtime_payload_json.get("agent_thread_id").is_some() { + let payload: SubAgentActivityEvent = serde_json::from_value(runtime_payload_json)?; + return self.end_sub_agent_activity(wall_time_unix_ms, tool_call_id, &kind, &payload); + } + match kind { + ToolCallKind::SpawnAgent => { + let payload: CollabAgentSpawnEndEvent = + serde_json::from_value(runtime_payload_json)?; + self.end_spawn_agent_interaction(wall_time_unix_ms, tool_call_id, &payload) + } + ToolCallKind::AssignAgentTask => { + let payload: CollabAgentInteractionEndEvent = + serde_json::from_value(runtime_payload_json)?; + self.end_message_agent_interaction( + wall_time_unix_ms, + tool_call_id, + InteractionEdgeKind::AssignAgentTask, + &payload, + ) + } + ToolCallKind::SendMessage => { + let payload: CollabAgentInteractionEndEvent = + serde_json::from_value(runtime_payload_json)?; + self.end_message_agent_interaction( + wall_time_unix_ms, + tool_call_id, + InteractionEdgeKind::SendMessage, + &payload, + ) + } + ToolCallKind::CloseAgent => { + let payload: CollabCloseEndEvent = serde_json::from_value(runtime_payload_json)?; + self.upsert_close_agent_interaction( + tool_call_id, + payload.receiver_thread_id.to_string(), + Some(wall_time_unix_ms), + ) + } + ToolCallKind::ExecCommand + | ToolCallKind::WriteStdin + | ToolCallKind::ApplyPatch + | ToolCallKind::Mcp { .. } + | ToolCallKind::Web + | ToolCallKind::ImageGeneration + | ToolCallKind::WaitAgent + | ToolCallKind::Other { .. } => Ok(()), + } + } + + fn end_sub_agent_activity( + &mut self, + wall_time_unix_ms: i64, + tool_call_id: &str, + tool_kind: &ToolCallKind, + payload: &SubAgentActivityEvent, + ) -> Result<()> { + let target_thread_id = payload.agent_thread_id.to_string(); + match (tool_kind, &payload.kind) { + (ToolCallKind::SpawnAgent, SubAgentActivityKind::Started) => { + let parent_thread_id = self + .rollout + .tool_calls + .get(tool_call_id) + .with_context(|| { + format!("agent activity referenced unknown tool call {tool_call_id}") + })? + .thread_id + .clone(); + self.queue_sub_agent_activity_message_edge( + wall_time_unix_ms, + tool_call_id, + spawn_edge_id(&parent_thread_id, &target_thread_id), + InteractionEdgeKind::SpawnAgent, + target_thread_id.clone(), + Some(target_thread_id), + ) + } + (ToolCallKind::AssignAgentTask, SubAgentActivityKind::Interacted) => self + .queue_sub_agent_activity_message_edge( + wall_time_unix_ms, + tool_call_id, + tool_edge_id(tool_call_id), + InteractionEdgeKind::AssignAgentTask, + target_thread_id, + /*unresolved_spawn_thread_id*/ None, + ), + (ToolCallKind::SendMessage, SubAgentActivityKind::Interacted) => self + .queue_sub_agent_activity_message_edge( + wall_time_unix_ms, + tool_call_id, + tool_edge_id(tool_call_id), + InteractionEdgeKind::SendMessage, + target_thread_id, + /*unresolved_spawn_thread_id*/ None, + ), + (ToolCallKind::CloseAgent, SubAgentActivityKind::Interrupted) => self + .upsert_close_agent_interaction( + tool_call_id, + target_thread_id, + Some(wall_time_unix_ms), + ), + (_, SubAgentActivityKind::Completed) => Ok(()), + _ => bail!( + "sub-agent activity {:?} does not match tool call kind {tool_kind:?}", + payload.kind + ), + } + } + + fn queue_sub_agent_activity_message_edge( + &mut self, + wall_time_unix_ms: i64, + tool_call_id: &str, + edge_id: String, + edge_kind: InteractionEdgeKind, + target_thread_id: String, + unresolved_spawn_thread_id: Option, + ) -> Result<()> { + let tool_call = self.rollout.tool_calls.get(tool_call_id).with_context(|| { + format!("agent activity referenced unknown tool call {tool_call_id}") + })?; + let started_at_unix_ms = tool_call.execution.started_at_unix_ms; + let message_author = self.agent_path_for_thread(&tool_call.thread_id)?; + let message_content = self.agent_message_content_from_invocation(tool_call_id)?; + let carried_raw_payload_ids = self.agent_tool_payload_ids(tool_call_id)?; + self.queue_or_resolve_agent_interaction_edge(PendingAgentInteractionEdge { + edge_id, + kind: edge_kind, + source: TraceAnchor::ToolCall { + tool_call_id: tool_call_id.to_string(), + }, + target_thread_id, + message_author, + message_content, + unresolved_spawn_thread_id, + started_at_unix_ms, + ended_at_unix_ms: Some(wall_time_unix_ms), + carried_raw_payload_ids, + }) + } + + fn agent_message_content_from_invocation(&self, tool_call_id: &str) -> Result { + let tool_call = self.rollout.tool_calls.get(tool_call_id).with_context(|| { + format!("agent activity referenced unknown tool call {tool_call_id}") + })?; + let invocation_payload_id = tool_call + .raw_invocation_payload_id + .as_deref() + .with_context(|| { + format!("agent activity tool call {tool_call_id} missing invocation payload") + })?; + let invocation_payload = self + .rollout + .raw_payloads + .get(invocation_payload_id) + .with_context(|| { + format!( + "agent activity tool call {tool_call_id} referenced missing invocation payload {invocation_payload_id}" + ) + })?; + let invocation = self.read_payload_json(invocation_payload)?; + let arguments = invocation + .get("payload") + .and_then(|payload| payload.get("arguments")) + .and_then(Value::as_str) + .with_context(|| { + format!("agent activity tool call {tool_call_id} missing function arguments") + })?; + let args: AgentMessageInvocationArgs = serde_json::from_str(arguments) + .with_context(|| format!("parse agent activity tool call {tool_call_id} arguments"))?; + Ok(args.message) + } + + /// Adds the canonical tool result payload to an already reduced multi-agent edge. + pub(super) fn attach_agent_interaction_tool_result( + &mut self, + tool_call_id: &str, + result_payload: Option<&RawPayloadRef>, + ) -> Result<()> { + let Some(result_payload) = result_payload else { + return Ok(()); + }; + if let Some(edge) = self + .rollout + .interaction_edges + .values_mut() + .find(|edge| tool_call_source_matches(&edge.source, tool_call_id)) + { + push_unique( + &mut edge.carried_raw_payload_ids, + &result_payload.raw_payload_id, + ); + return Ok(()); + } + + // Agent delivery edges intentionally wait for the recipient-side + // conversation item. Tool end can arrive before that item is + // reduced, so preserve the response payload on the pending edge rather + // than dropping evidence until the delivery materializes. + if let Some(pending) = self + .pending_agent_interaction_edges + .iter_mut() + .find(|pending| tool_call_source_matches(&pending.source, tool_call_id)) + { + push_unique( + &mut pending.carried_raw_payload_ids, + &result_payload.raw_payload_id, + ); + } + Ok(()) + } + + fn end_spawn_agent_interaction( + &mut self, + wall_time_unix_ms: i64, + tool_call_id: &str, + payload: &CollabAgentSpawnEndEvent, + ) -> Result<()> { + let Some(child_thread_id) = payload.new_thread_id else { + return Ok(()); + }; + let tool_call = &self.rollout.tool_calls[tool_call_id]; + let child_thread_id = child_thread_id.to_string(); + let edge_id = spawn_edge_id(&payload.sender_thread_id.to_string(), &child_thread_id); + let message_author = self.agent_path_for_thread(&tool_call.thread_id)?; + + self.queue_or_resolve_agent_interaction_edge(PendingAgentInteractionEdge { + edge_id, + kind: InteractionEdgeKind::SpawnAgent, + source: TraceAnchor::ToolCall { + tool_call_id: tool_call_id.to_string(), + }, + target_thread_id: child_thread_id.clone(), + message_author, + message_content: payload.prompt.clone(), + unresolved_spawn_thread_id: Some(child_thread_id), + started_at_unix_ms: tool_call.execution.started_at_unix_ms, + ended_at_unix_ms: Some(wall_time_unix_ms), + carried_raw_payload_ids: self.agent_tool_payload_ids(tool_call_id)?, + }) + } + + fn end_message_agent_interaction( + &mut self, + wall_time_unix_ms: i64, + tool_call_id: &str, + edge_kind: InteractionEdgeKind, + payload: &CollabAgentInteractionEndEvent, + ) -> Result<()> { + self.queue_message_agent_interaction( + tool_call_id, + edge_kind, + payload.receiver_thread_id.to_string(), + payload.prompt.clone(), + Some(wall_time_unix_ms), + ) + } + + fn queue_message_agent_interaction( + &mut self, + tool_call_id: &str, + kind: InteractionEdgeKind, + target_thread_id: String, + message_content: String, + ended_at_unix_ms: Option, + ) -> Result<()> { + let tool_call = &self.rollout.tool_calls[tool_call_id]; + let message_author = self.agent_path_for_thread(&tool_call.thread_id)?; + self.queue_or_resolve_agent_interaction_edge(PendingAgentInteractionEdge { + edge_id: tool_edge_id(tool_call_id), + kind, + source: TraceAnchor::ToolCall { + tool_call_id: tool_call_id.to_string(), + }, + target_thread_id, + message_author, + message_content, + unresolved_spawn_thread_id: None, + started_at_unix_ms: tool_call.execution.started_at_unix_ms, + ended_at_unix_ms, + carried_raw_payload_ids: self.agent_tool_payload_ids(tool_call_id)?, + }) + } + + fn agent_tool_payload_ids(&self, tool_call_id: &str) -> Result> { + let tool_call = + self.rollout.tool_calls.get(tool_call_id).with_context(|| { + format!("agent edge referenced unknown tool call {tool_call_id}") + })?; + let mut payload_ids = Vec::new(); + if let Some(payload_id) = &tool_call.raw_invocation_payload_id { + push_unique(&mut payload_ids, payload_id); + } + for payload_id in &tool_call.raw_runtime_payload_ids { + push_unique(&mut payload_ids, payload_id); + } + if let Some(payload_id) = &tool_call.raw_result_payload_id { + push_unique(&mut payload_ids, payload_id); + } + Ok(payload_ids) + } + + fn upsert_close_agent_interaction( + &mut self, + tool_call_id: &str, + target_thread_id: String, + ended_at_unix_ms: Option, + ) -> Result<()> { + if !self.rollout.threads.contains_key(&target_thread_id) { + // A failed close can name a thread that never participated in this + // trace. Keep that evidence on the ToolCall raw payloads rather + // than creating an anchor to a non-existent reduced object. + return Ok(()); + } + let started_at_unix_ms = self + .rollout + .tool_calls + .get(tool_call_id) + .with_context(|| format!("close edge referenced unknown tool call {tool_call_id}"))? + .execution + .started_at_unix_ms; + let carried_raw_payload_ids = self.agent_tool_payload_ids(tool_call_id)?; + self.upsert_interaction_edge(InteractionEdge { + edge_id: tool_edge_id(tool_call_id), + kind: InteractionEdgeKind::CloseAgent, + source: TraceAnchor::ToolCall { + tool_call_id: tool_call_id.to_string(), + }, + target: TraceAnchor::Thread { + thread_id: target_thread_id, + }, + started_at_unix_ms, + ended_at_unix_ms, + carried_item_ids: Vec::new(), + carried_raw_payload_ids, + }) + } + + /// Queues or resolves the edge from a child completion to its parent notification. + pub(in crate::reducer) fn queue_agent_result_interaction_edge( + &mut self, + observed: ObservedAgentResultEdge, + ) -> Result<()> { + let message_author = self.agent_path_for_thread(&observed.child_thread_id)?; + let source = if let Some(source_item_id) = self.latest_assistant_message_item_for_turn( + &observed.child_thread_id, + &observed.child_codex_turn_id, + ) { + TraceAnchor::ConversationItem { + item_id: source_item_id, + } + } else { + // Child completion is delivered from AgentStatus, not from transcript + // content. Failed or cancelled children can therefore notify the parent + // without producing a final assistant message. Anchor those edges to + // the child thread so the trace keeps the valid delivery instead of + // inventing a missing conversation item. + TraceAnchor::Thread { + thread_id: observed.child_thread_id, + } + }; + + self.queue_or_resolve_agent_interaction_edge(PendingAgentInteractionEdge { + edge_id: observed.edge_id, + kind: InteractionEdgeKind::AgentResult, + source, + target_thread_id: observed.parent_thread_id, + message_author, + message_content: observed.message, + unresolved_spawn_thread_id: None, + started_at_unix_ms: observed.wall_time_unix_ms, + ended_at_unix_ms: Some(observed.wall_time_unix_ms), + carried_raw_payload_ids: observed + .carried_payload + .map(|payload| vec![payload.raw_payload_id]) + .unwrap_or_default(), + }) + } + + /// Resolves pending agent edges whose target is the newly reduced conversation item. + pub(in crate::reducer) fn resolve_pending_agent_edges_for_item( + &mut self, + item_id: &str, + ) -> Result<()> { + if self.is_interaction_edge_target_item(item_id) { + return Ok(()); + } + let Some((thread_id, message_author, message_content)) = + self.inter_agent_message_item(item_id) + else { + return Ok(()); + }; + let Some(pending_index) = self + .pending_agent_interaction_edges + .iter() + .position(|pending| { + pending.target_thread_id == thread_id + && pending.message_author == message_author + && pending.message_content == message_content + }) + else { + return Ok(()); + }; + let pending = self.pending_agent_interaction_edges.remove(pending_index); + self.upsert_agent_interaction_edge_for_item(pending, item_id.to_string()) + } + + fn queue_or_resolve_agent_interaction_edge( + &mut self, + pending: PendingAgentInteractionEdge, + ) -> Result<()> { + if let Some(item_id) = self.find_unlinked_inter_agent_message_item( + &pending.target_thread_id, + &pending.message_author, + &pending.message_content, + ) { + return self.upsert_agent_interaction_edge_for_item(pending, item_id); + } + + if let Some(existing) = self + .pending_agent_interaction_edges + .iter_mut() + .find(|existing| existing.edge_id == pending.edge_id) + { + if existing.kind != pending.kind + || existing.source != pending.source + || existing.target_thread_id != pending.target_thread_id + || existing.message_author != pending.message_author + || existing.message_content != pending.message_content + || existing.unresolved_spawn_thread_id != pending.unresolved_spawn_thread_id + { + bail!( + "pending interaction edge {} was observed with conflicting delivery data", + pending.edge_id + ); + } + existing.started_at_unix_ms = + existing.started_at_unix_ms.min(pending.started_at_unix_ms); + existing.ended_at_unix_ms = match (existing.ended_at_unix_ms, pending.ended_at_unix_ms) + { + (Some(existing_ended), Some(pending_ended)) => { + Some(existing_ended.max(pending_ended)) + } + (None, ended) | (ended, None) => ended, + }; + extend_unique( + &mut existing.carried_raw_payload_ids, + pending.carried_raw_payload_ids, + ); + return Ok(()); + } + + self.pending_agent_interaction_edges.push(pending); + Ok(()) + } + + /// Materializes unresolved spawn edges that have a valid child-thread fallback target. + pub(in crate::reducer) fn resolve_pending_spawn_edge_fallbacks(&mut self) -> Result<()> { + let pending_edges = std::mem::take(&mut self.pending_agent_interaction_edges); + for pending in pending_edges { + let Some(child_thread_id) = pending.unresolved_spawn_thread_id else { + continue; + }; + if pending.kind != InteractionEdgeKind::SpawnAgent { + bail!( + "non-spawn interaction edge {} carried a spawn fallback target", + pending.edge_id + ); + } + if !self.rollout.threads.contains_key(&child_thread_id) { + continue; + } + + // Spawn normally resolves to the child task message because that is + // where the delegated work first becomes model-visible. A child can + // fail before that transcript item exists, but the spawned thread is + // still real and the spawning tool still created it. Preserve that + // relationship with the thread fallback instead of dropping the edge. + self.upsert_interaction_edge(InteractionEdge { + edge_id: pending.edge_id, + kind: pending.kind, + source: pending.source, + target: TraceAnchor::Thread { + thread_id: child_thread_id, + }, + started_at_unix_ms: pending.started_at_unix_ms, + ended_at_unix_ms: pending.ended_at_unix_ms, + carried_item_ids: Vec::new(), + carried_raw_payload_ids: pending.carried_raw_payload_ids, + })?; + } + Ok(()) + } + + fn upsert_agent_interaction_edge_for_item( + &mut self, + pending: PendingAgentInteractionEdge, + target_item_id: String, + ) -> Result<()> { + self.upsert_interaction_edge(InteractionEdge { + edge_id: pending.edge_id, + kind: pending.kind, + source: pending.source, + target: TraceAnchor::ConversationItem { + item_id: target_item_id.clone(), + }, + started_at_unix_ms: pending.started_at_unix_ms, + ended_at_unix_ms: pending.ended_at_unix_ms, + carried_item_ids: vec![target_item_id], + carried_raw_payload_ids: pending.carried_raw_payload_ids, + }) + } + + fn upsert_interaction_edge(&mut self, edge: InteractionEdge) -> Result<()> { + if let Some(existing) = self.rollout.interaction_edges.get_mut(&edge.edge_id) { + if existing.kind != edge.kind + || existing.source != edge.source + || existing.target != edge.target + { + bail!( + "interaction edge {} was observed with conflicting endpoints", + edge.edge_id + ); + } + existing.started_at_unix_ms = existing.started_at_unix_ms.min(edge.started_at_unix_ms); + existing.ended_at_unix_ms = match (existing.ended_at_unix_ms, edge.ended_at_unix_ms) { + (Some(existing_ended), Some(edge_ended)) => Some(existing_ended.max(edge_ended)), + (None, ended) | (ended, None) => ended, + }; + extend_unique(&mut existing.carried_item_ids, edge.carried_item_ids); + extend_unique( + &mut existing.carried_raw_payload_ids, + edge.carried_raw_payload_ids, + ); + return Ok(()); + } + + self.rollout + .interaction_edges + .insert(edge.edge_id.clone(), edge); + Ok(()) + } + + fn find_unlinked_inter_agent_message_item( + &self, + thread_id: &str, + message_author: &str, + message_content: &str, + ) -> Option { + self.rollout + .threads + .get(thread_id)? + .conversation_item_ids + .iter() + .find(|item_id| { + !self.is_interaction_edge_target_item(item_id) + && self + .inter_agent_message_item(item_id) + .is_some_and(|(_, author, content)| { + author == message_author && content == message_content + }) + }) + .cloned() + } + + fn inter_agent_message_item(&self, item_id: &str) -> Option<(String, String, String)> { + let item = self.rollout.conversation_items.get(item_id)?; + let (author_agent_path, recipient_agent_path, message_content) = + inter_agent_message_fields(item)?; + let thread = self.rollout.threads.get(&item.thread_id)?; + if recipient_agent_path != thread.agent_path { + return None; + } + Some((item.thread_id.clone(), author_agent_path, message_content)) + } + + fn agent_path_for_thread(&self, thread_id: &str) -> Result { + self.rollout + .threads + .get(thread_id) + .map(|thread| thread.agent_path.clone()) + .with_context(|| format!("agent edge referenced unknown thread {thread_id}")) + } + + fn is_interaction_edge_target_item(&self, item_id: &str) -> bool { + self.rollout + .interaction_edges + .values() + .any(|edge| matches!(&edge.target, TraceAnchor::ConversationItem { item_id: target } if target == item_id)) + } + + fn latest_assistant_message_item_for_turn( + &self, + thread_id: &str, + codex_turn_id: &str, + ) -> Option { + self.rollout + .conversation_items + .values() + .filter(|item| { + item.thread_id == thread_id + && item.codex_turn_id.as_deref() == Some(codex_turn_id) + && item.role == ConversationRole::Assistant + && item.kind == ConversationItemKind::Message + && item.agent_message.is_none() + }) + .max_by_key(|item| item.first_seen_at_unix_ms) + .map(|item| item.item_id.clone()) + } +} + +fn extend_unique(items: &mut Vec, new_items: Vec) { + for item in new_items { + if !items.iter().any(|existing| existing == &item) { + items.push(item); + } + } +} + +fn tool_edge_id(tool_call_id: &str) -> String { + format!("edge:tool:{tool_call_id}") +} + +fn tool_call_source_matches(anchor: &TraceAnchor, tool_call_id: &str) -> bool { + matches!(anchor, TraceAnchor::ToolCall { tool_call_id: source } if source == tool_call_id) +} + +fn push_unique(items: &mut Vec, item: &str) { + if !items.iter().any(|existing| existing == item) { + items.push(item.to_string()); + } +} + +fn inter_agent_message_fields(item: &ConversationItem) -> Option<(String, String, String)> { + if item.role != ConversationRole::Assistant || item.kind != ConversationItemKind::Message { + return None; + } + if let Some(agent_message) = &item.agent_message { + let message_content = match item.body.parts.as_slice() { + [ConversationPart::Text { text }] => text, + [ConversationPart::Encoded { label, value }] if label == "encrypted_content" => value, + [ + ConversationPart::Text { .. }, + ConversationPart::Encoded { label, value }, + ] if label == "encrypted_content" => value, + _ => return None, + }; + return Some(( + agent_message.author.clone(), + agent_message.recipient.clone(), + message_content.clone(), + )); + } + + // Older traces store multi-agent v2 deliveries as assistant messages whose + // text is serialized `InterAgentCommunication`. Treat only that exact + // transport shape as an edge target; ordinary assistant JSON must not be + // mistaken for cross-thread delivery. + let [ConversationPart::Text { text }] = item.body.parts.as_slice() else { + return None; + }; + let communication = serde_json::from_str::(text).ok()?; + Some(( + communication.author.to_string(), + communication.recipient.to_string(), + communication + .encrypted_content + .unwrap_or(communication.content), + )) +} + +#[cfg(test)] +#[path = "agents_tests.rs"] +mod tests; diff --git a/codex-rs/rollout-trace/src/reducer/tool/agents_tests.rs b/codex-rs/rollout-trace/src/reducer/tool/agents_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..f5a4648f8021e9cd7e0ffe35aecae11a6d8ebc93 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/tool/agents_tests.rs @@ -0,0 +1,991 @@ +use pretty_assertions::assert_eq; +use serde_json::json; +use tempfile::TempDir; + +use crate::model::AgentOrigin; +use crate::model::ExecutionStatus; +use crate::model::InteractionEdgeKind; +use crate::model::RolloutStatus; +use crate::model::ToolCallKind; +use crate::model::ToolCallSummary; +use crate::model::TraceAnchor; +use crate::payload::RawPayloadKind; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawToolCallRequester; +use crate::raw_event::RawTraceEventPayload; +use crate::reducer::test_support::append_completed_inference; +use crate::reducer::test_support::append_inference_request; +use crate::reducer::test_support::create_started_agent_writer; +use crate::reducer::test_support::message; +use crate::reducer::test_support::start_agent_turn; +use crate::reducer::test_support::start_thread; +use crate::reducer::test_support::start_turn_for_thread; +use crate::reducer::test_support::trace_context_for_agent; +use crate::reducer::test_support::trace_context_for_thread; +use crate::replay_bundle; +use crate::writer::TraceWriter; + +#[test] +fn child_thread_metadata_creates_spawn_origin_without_delivery_edge() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = TraceWriter::create( + temp.path(), + "trace-1".to_string(), + "rollout-1".to_string(), + "019d0000-0000-7000-8000-000000000002".to_string(), + )?; + let metadata = writer.write_json_payload( + RawPayloadKind::SessionMetadata, + &json!({ + "nickname": "James", + "agent_role": "explorer", + "task_name": "repo_file_counter", + "model": "gpt-test", + "session_source": { + "subagent": { + "thread_spawn": { + "parent_thread_id": "019d0000-0000-7000-8000-000000000001", + "agent_path": "/root/repo_file_counter", + "agent_nickname": "James", + "agent_role": "explorer" + } + } + } + }), + )?; + writer.append(RawTraceEventPayload::ThreadStarted { + thread_id: "019d0000-0000-7000-8000-000000000002".to_string(), + agent_path: "/root/repo_file_counter".to_string(), + metadata_payload: Some(metadata), + })?; + + let replayed = replay_bundle(temp.path())?; + let thread = &replayed.threads["019d0000-0000-7000-8000-000000000002"]; + assert_eq!(thread.nickname, Some("James".to_string())); + assert_eq!(thread.default_model, Some("gpt-test".to_string())); + assert_eq!( + thread.origin, + AgentOrigin::Spawned { + parent_thread_id: "019d0000-0000-7000-8000-000000000001".to_string(), + spawn_edge_id: "edge:spawn:019d0000-0000-7000-8000-000000000001:019d0000-0000-7000-8000-000000000002".to_string(), + task_name: "repo_file_counter".to_string(), + agent_role: "explorer".to_string(), + } + ); + assert!( + !replayed.interaction_edges.contains_key( + "edge:spawn:019d0000-0000-7000-8000-000000000001:019d0000-0000-7000-8000-000000000002" + ), + "spawn metadata identifies the child, but the delivery edge waits for the recipient \ + conversation item" + ); + + Ok(()) +} + +#[test] +fn spawn_runtime_payload_targets_delivered_child_message() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_agent_writer(&temp)?; + start_agent_turn(&writer, "turn-1")?; + + let spawn_payloads = append_spawn_agent_tool_lifecycle(&writer, "turn-1")?; + + // Then record the child-side model-visible task message. This is the + // preferred target because it pinpoints where the delegated work entered + // the child timeline. + start_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "/root/repo_file_counter", + )?; + start_turn_for_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "turn-child-1", + )?; + let delivered = inter_agent_message( + "/root", + "/root/repo_file_counter", + "count", + /*trigger_turn*/ true, + ); + append_inference_request( + &writer, + "019d0000-0000-7000-8000-000000000002", + "turn-child-1", + "inference-child-1", + vec![message("assistant", &delivered)], + )?; + + let replayed = replay_bundle(temp.path())?; + let edge = &replayed.interaction_edges["edge:spawn:019d0000-0000-7000-8000-000000000001:019d0000-0000-7000-8000-000000000002"]; + assert_eq!(edge.kind, InteractionEdgeKind::SpawnAgent); + assert_eq!( + edge.source, + TraceAnchor::ToolCall { + tool_call_id: "call-spawn".to_string() + } + ); + let target_item_id = target_conversation_item_id(&edge.target); + assert_eq!(edge.carried_item_ids, vec![target_item_id.clone()]); + assert_eq!( + replayed.conversation_items[target_item_id].thread_id, + "019d0000-0000-7000-8000-000000000002" + ); + assert_eq!( + edge.carried_raw_payload_ids, + vec![ + spawn_payloads.invocation.raw_payload_id, + spawn_payloads.begin.raw_payload_id, + spawn_payloads.end.raw_payload_id, + spawn_payloads.result.raw_payload_id, + ] + ); + + Ok(()) +} + +#[test] +fn spawn_runtime_payload_falls_back_to_child_thread_without_delivery_item() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_agent_writer(&temp)?; + start_agent_turn(&writer, "turn-1")?; + let spawn_payloads = append_spawn_agent_tool_lifecycle(&writer, "turn-1")?; + + // Deliberately start the child thread without appending an inference + // request containing the inter-agent task message. This reproduces the + // failure path where the child aborts before the reducer can target the + // precise child-side ConversationItem. + start_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "/root/repo_file_counter", + )?; + + let replayed = replay_bundle(temp.path())?; + let edge = &replayed.interaction_edges["edge:spawn:019d0000-0000-7000-8000-000000000001:019d0000-0000-7000-8000-000000000002"]; + assert_eq!(edge.kind, InteractionEdgeKind::SpawnAgent); + assert_eq!( + edge.source, + TraceAnchor::ToolCall { + tool_call_id: "call-spawn".to_string() + } + ); + assert_eq!( + edge.target, + TraceAnchor::Thread { + thread_id: "019d0000-0000-7000-8000-000000000002".to_string() + } + ); + // No transcript item carried the task, so the fallback edge should not + // claim one. The raw payloads still preserve the tool evidence. + assert!(edge.carried_item_ids.is_empty()); + assert_eq!( + edge.carried_raw_payload_ids, + vec![ + spawn_payloads.invocation.raw_payload_id, + spawn_payloads.begin.raw_payload_id, + spawn_payloads.end.raw_payload_id, + spawn_payloads.result.raw_payload_id, + ] + ); + + Ok(()) +} + +#[test] +fn sub_agent_started_activity_creates_spawn_edge() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_agent_writer(&temp)?; + start_agent_turn(&writer, "turn-1")?; + let child_thread_id = "019d0000-0000-7000-8000-000000000002"; + let invocation_payload = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "spawn_agent", + "payload": { + "type": "function", + "arguments": "{\"message\":\"review this\",\"task_name\":\"reviewer\"}" + } + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "call-spawn-v2".to_string(), + model_visible_call_id: Some("call-spawn-v2".to_string()), + code_mode_runtime_tool_id: None, + requester: RawToolCallRequester::Model, + kind: ToolCallKind::SpawnAgent, + summary: ToolCallSummary::Generic { + label: "spawn_agent".to_string(), + input_preview: None, + output_preview: None, + }, + invocation_payload: Some(invocation_payload.clone()), + }, + )?; + let activity_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "event_id": "call-spawn-v2", + "occurred_at_ms": 1234, + "agent_thread_id": child_thread_id, + "agent_path": "/root/reviewer", + "kind": "started" + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallRuntimeEnded { + tool_call_id: "call-spawn-v2".to_string(), + status: ExecutionStatus::Completed, + runtime_payload: activity_payload.clone(), + }, + )?; + start_thread(&writer, child_thread_id, "/root/reviewer")?; + start_turn_for_thread(&writer, child_thread_id, "turn-child-1")?; + append_inference_request( + &writer, + child_thread_id, + "turn-child-1", + "inference-child-1", + vec![json!({ + "type": "agent_message", + "author": "/root", + "recipient": "/root/reviewer", + "content": [ + { + "type": "input_text", + "text": "Message Type: NEW_TASK\nTask name: /root/reviewer\nSender: /root\nPayload:\n" + }, + {"type": "encrypted_content", "encrypted_content": "review this"} + ] + })], + )?; + + let replayed = replay_bundle(temp.path())?; + let edge_id = format!("edge:spawn:019d0000-0000-7000-8000-000000000001:{child_thread_id}"); + let edge = &replayed.interaction_edges[&edge_id]; + assert_eq!(edge.kind, InteractionEdgeKind::SpawnAgent); + let target_item_id = target_conversation_item_id(&edge.target); + assert_eq!(edge.carried_item_ids, vec![target_item_id.clone()]); + assert_eq!( + replayed.conversation_items[target_item_id].thread_id, + child_thread_id + ); + assert_eq!( + edge.carried_raw_payload_ids, + vec![ + invocation_payload.raw_payload_id, + activity_payload.raw_payload_id, + ] + ); + Ok(()) +} + +#[test] +fn send_message_runtime_payload_targets_delivered_child_message() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_agent_writer(&temp)?; + start_agent_turn(&writer, "turn-1")?; + let invocation_payload = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "send_message", + "payload": { + "type": "function", + "arguments": "{\"target\":\"/root/child\",\"message\":\"hello\"}" + } + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "call-send".to_string(), + model_visible_call_id: Some("call-send".to_string()), + code_mode_runtime_tool_id: None, + requester: RawToolCallRequester::Model, + kind: ToolCallKind::SendMessage, + summary: ToolCallSummary::Generic { + label: "send_message".to_string(), + input_preview: None, + output_preview: None, + }, + invocation_payload: Some(invocation_payload), + }, + )?; + let begin_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "call-send", + "sender_thread_id": "019d0000-0000-7000-8000-000000000001", + "receiver_thread_id": "019d0000-0000-7000-8000-000000000002", + "prompt": "hello", + "status": "running" + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallRuntimeStarted { + tool_call_id: "call-send".to_string(), + runtime_payload: begin_payload, + }, + )?; + let end_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "call-send", + "sender_thread_id": "019d0000-0000-7000-8000-000000000001", + "receiver_thread_id": "019d0000-0000-7000-8000-000000000002", + "prompt": "hello", + "status": "running" + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallRuntimeEnded { + tool_call_id: "call-send".to_string(), + status: ExecutionStatus::Completed, + runtime_payload: end_payload, + }, + )?; + start_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "/root/child", + )?; + start_turn_for_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "turn-child-1", + )?; + let delivered = + inter_agent_message("/root", "/root/child", "hello", /*trigger_turn*/ false); + append_inference_request( + &writer, + "019d0000-0000-7000-8000-000000000002", + "turn-child-1", + "inference-child-1", + vec![message("assistant", &delivered)], + )?; + + let replayed = replay_bundle(temp.path())?; + let edge = &replayed.interaction_edges["edge:tool:call-send"]; + assert_eq!(edge.kind, InteractionEdgeKind::SendMessage); + assert_eq!( + edge.source, + TraceAnchor::ToolCall { + tool_call_id: "call-send".to_string() + } + ); + let target_item_id = target_conversation_item_id(&edge.target); + assert_eq!(edge.carried_item_ids, vec![target_item_id.clone()]); + assert_eq!( + replayed.conversation_items[target_item_id].thread_id, + "019d0000-0000-7000-8000-000000000002" + ); + assert!(edge.ended_at_unix_ms.is_some()); + + Ok(()) +} + +#[test] +fn send_message_activity_targets_delivered_child_message() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_agent_writer(&temp)?; + start_agent_turn(&writer, "turn-1")?; + let child_thread_id = "019d0000-0000-7000-8000-000000000002"; + let invocation_payload = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "send_message", + "payload": { + "type": "function", + "arguments": "{\"target\":\"/root/child\",\"message\":\"hello again\"}" + } + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "call-send-v2".to_string(), + model_visible_call_id: Some("call-send-v2".to_string()), + code_mode_runtime_tool_id: None, + requester: RawToolCallRequester::Model, + kind: ToolCallKind::SendMessage, + summary: ToolCallSummary::Generic { + label: "send_message".to_string(), + input_preview: None, + output_preview: None, + }, + invocation_payload: Some(invocation_payload.clone()), + }, + )?; + let activity_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "event_id": "call-send-v2", + "occurred_at_ms": 1234, + "agent_thread_id": child_thread_id, + "agent_path": "/root/child", + "kind": "interacted" + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallRuntimeEnded { + tool_call_id: "call-send-v2".to_string(), + status: ExecutionStatus::Completed, + runtime_payload: activity_payload.clone(), + }, + )?; + start_thread(&writer, child_thread_id, "/root/child")?; + start_turn_for_thread(&writer, child_thread_id, "turn-child-1")?; + let delivered = inter_agent_message( + "/root", + "/root/child", + "hello again", + /*trigger_turn*/ false, + ); + append_inference_request( + &writer, + child_thread_id, + "turn-child-1", + "inference-child-1", + vec![message("assistant", &delivered)], + )?; + + let replayed = replay_bundle(temp.path())?; + let edge = &replayed.interaction_edges["edge:tool:call-send-v2"]; + assert_eq!(edge.kind, InteractionEdgeKind::SendMessage); + let target_item_id = target_conversation_item_id(&edge.target); + assert_eq!(edge.carried_item_ids, vec![target_item_id.clone()]); + assert_eq!( + replayed.conversation_items[target_item_id].thread_id, + child_thread_id + ); + assert_eq!( + edge.carried_raw_payload_ids, + vec![ + invocation_payload.raw_payload_id, + activity_payload.raw_payload_id, + ] + ); + + Ok(()) +} + +#[test] +fn followup_activity_targets_delivered_child_message() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_agent_writer(&temp)?; + start_agent_turn(&writer, "turn-1")?; + let child_thread_id = "019d0000-0000-7000-8000-000000000002"; + let invocation_payload = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "followup_task", + "payload": { + "type": "function", + "arguments": "{\"target\":\"/root/child\",\"message\":\"continue\"}" + } + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "call-followup-v2".to_string(), + model_visible_call_id: Some("call-followup-v2".to_string()), + code_mode_runtime_tool_id: None, + requester: RawToolCallRequester::Model, + kind: ToolCallKind::AssignAgentTask, + summary: ToolCallSummary::Generic { + label: "followup_task".to_string(), + input_preview: None, + output_preview: None, + }, + invocation_payload: Some(invocation_payload.clone()), + }, + )?; + let activity_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "event_id": "call-followup-v2", + "occurred_at_ms": 1234, + "agent_thread_id": child_thread_id, + "agent_path": "/root/child", + "kind": "interacted" + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallRuntimeEnded { + tool_call_id: "call-followup-v2".to_string(), + status: ExecutionStatus::Completed, + runtime_payload: activity_payload.clone(), + }, + )?; + start_thread(&writer, child_thread_id, "/root/child")?; + start_turn_for_thread(&writer, child_thread_id, "turn-child-1")?; + let delivered = inter_agent_message( + "/root", + "/root/child", + "continue", + /*trigger_turn*/ true, + ); + append_inference_request( + &writer, + child_thread_id, + "turn-child-1", + "inference-child-1", + vec![message("assistant", &delivered)], + )?; + + let replayed = replay_bundle(temp.path())?; + let edge = &replayed.interaction_edges["edge:tool:call-followup-v2"]; + assert_eq!(edge.kind, InteractionEdgeKind::AssignAgentTask); + let target_item_id = target_conversation_item_id(&edge.target); + assert_eq!(edge.carried_item_ids, vec![target_item_id.clone()]); + assert_eq!( + replayed.conversation_items[target_item_id].thread_id, + child_thread_id + ); + assert_eq!( + edge.carried_raw_payload_ids, + vec![ + invocation_payload.raw_payload_id, + activity_payload.raw_payload_id, + ] + ); + + Ok(()) +} + +#[test] +fn close_agent_runtime_payload_targets_thread() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_agent_writer(&temp)?; + start_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "/root/child", + )?; + start_agent_turn(&writer, "turn-1")?; + let invocation_payload = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "close_agent", + "payload": { + "type": "function", + "arguments": r#"{"target":"/root/child"}"# + } + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "call-close".to_string(), + model_visible_call_id: Some("call-close".to_string()), + code_mode_runtime_tool_id: None, + requester: RawToolCallRequester::Model, + kind: ToolCallKind::CloseAgent, + summary: ToolCallSummary::Generic { + label: "close_agent".to_string(), + input_preview: None, + output_preview: None, + }, + invocation_payload: Some(invocation_payload.clone()), + }, + )?; + let begin_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "call-close", + "sender_thread_id": "019d0000-0000-7000-8000-000000000001", + "receiver_thread_id": "019d0000-0000-7000-8000-000000000002" + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallRuntimeStarted { + tool_call_id: "call-close".to_string(), + runtime_payload: begin_payload.clone(), + }, + )?; + let end_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "call-close", + "sender_thread_id": "019d0000-0000-7000-8000-000000000001", + "receiver_thread_id": "019d0000-0000-7000-8000-000000000002", + "receiver_agent_nickname": "Scout", + "receiver_agent_role": "explorer", + "status": "running" + }), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallRuntimeEnded { + tool_call_id: "call-close".to_string(), + status: ExecutionStatus::Completed, + runtime_payload: end_payload.clone(), + }, + )?; + let result_payload = writer.write_json_payload( + RawPayloadKind::ToolResult, + &json!({"previous_status": "running"}), + )?; + writer.append_with_context( + trace_context_for_agent("turn-1"), + RawTraceEventPayload::ToolCallEnded { + tool_call_id: "call-close".to_string(), + status: ExecutionStatus::Completed, + result_payload: Some(result_payload.clone()), + }, + )?; + writer.append(RawTraceEventPayload::ThreadEnded { + thread_id: "019d0000-0000-7000-8000-000000000002".to_string(), + status: RolloutStatus::Completed, + })?; + + let replayed = replay_bundle(temp.path())?; + let edge = &replayed.interaction_edges["edge:tool:call-close"]; + assert_eq!(edge.kind, InteractionEdgeKind::CloseAgent); + assert_eq!( + edge.source, + TraceAnchor::ToolCall { + tool_call_id: "call-close".to_string() + } + ); + assert_eq!( + edge.target, + TraceAnchor::Thread { + thread_id: "019d0000-0000-7000-8000-000000000002".to_string() + } + ); + assert!(edge.carried_item_ids.is_empty()); + assert_eq!( + edge.carried_raw_payload_ids, + vec![ + invocation_payload.raw_payload_id, + begin_payload.raw_payload_id, + end_payload.raw_payload_id, + result_payload.raw_payload_id, + ] + ); + assert_eq!( + replayed.threads["019d0000-0000-7000-8000-000000000002"] + .execution + .status, + ExecutionStatus::Completed + ); + assert_eq!(replayed.status, RolloutStatus::Running); + + Ok(()) +} + +#[test] +fn agent_result_edge_links_child_result_to_parent_notification() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_agent_writer(&temp)?; + start_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "/root/child", + )?; + start_turn_for_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "turn-child-1", + )?; + append_completed_inference( + &writer, + "019d0000-0000-7000-8000-000000000002", + "turn-child-1", + "inference-child-1", + vec![message("assistant", "task")], + vec![message("assistant", "done")], + )?; + + let notification = "{\"agent_path\":\"/root/child\",\"status\":{\"completed\":\"done\"}}"; + let carried_payload = writer.write_json_payload( + RawPayloadKind::AgentResult, + &json!({ + "child_agent_path": "/root/child", + "message": notification, + "status": {"completed": "done"} + }), + )?; + writer.append_with_context( + trace_context_for_thread("019d0000-0000-7000-8000-000000000002", "turn-child-1"), + RawTraceEventPayload::AgentResultObserved { + edge_id: "edge:agent_result:thread-child:turn-child-1:thread-root".to_string(), + child_thread_id: "019d0000-0000-7000-8000-000000000002".to_string(), + child_codex_turn_id: "turn-child-1".to_string(), + parent_thread_id: "019d0000-0000-7000-8000-000000000001".to_string(), + message: notification.to_string(), + carried_payload: Some(carried_payload.clone()), + }, + )?; + + start_agent_turn(&writer, "turn-root-1")?; + let delivered = inter_agent_message( + "/root/child", + "/root", + notification, + /*trigger_turn*/ false, + ); + append_inference_request( + &writer, + "019d0000-0000-7000-8000-000000000001", + "turn-root-1", + "inference-root-1", + vec![message("assistant", &delivered)], + )?; + + let replayed = replay_bundle(temp.path())?; + let edge = + &replayed.interaction_edges["edge:agent_result:thread-child:turn-child-1:thread-root"]; + assert_eq!(edge.kind, InteractionEdgeKind::AgentResult); + let TraceAnchor::ConversationItem { + item_id: source_item_id, + } = &edge.source + else { + panic!("expected child result conversation item source"); + }; + assert_eq!( + text_body(&replayed.conversation_items[source_item_id]), + "done" + ); + let target_item_id = target_conversation_item_id(&edge.target); + assert_eq!( + replayed.conversation_items[target_item_id].thread_id, + "019d0000-0000-7000-8000-000000000001" + ); + assert_eq!(edge.carried_item_ids, vec![target_item_id.clone()]); + assert_eq!( + edge.carried_raw_payload_ids, + vec![carried_payload.raw_payload_id] + ); + + Ok(()) +} + +#[test] +fn agent_result_edge_falls_back_to_child_thread_without_result_message() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_agent_writer(&temp)?; + + // The child received its task but produced no assistant output. Failed + // child tasks can still notify the parent through AgentStatus, so the + // inbound task must not be mistaken for the child's result. + start_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "/root/child", + )?; + start_turn_for_thread( + &writer, + "019d0000-0000-7000-8000-000000000002", + "turn-child-1", + )?; + append_inference_request( + &writer, + "019d0000-0000-7000-8000-000000000002", + "turn-child-1", + "inference-child-1", + vec![json!({ + "type": "agent_message", + "author": "/root", + "recipient": "/root/child", + "content": [{"type": "input_text", "text": "do the task"}] + })], + )?; + + let notification = r#"{"agent_path":"/root/child","status":{"failed":"boom"}}"#; + let carried_payload = writer.write_json_payload( + RawPayloadKind::AgentResult, + &json!({ + "child_agent_path": "/root/child", + "message": notification, + "status": {"failed": "boom"} + }), + )?; + writer.append_with_context( + trace_context_for_thread("019d0000-0000-7000-8000-000000000002", "turn-child-1"), + RawTraceEventPayload::AgentResultObserved { + edge_id: "edge:agent_result:thread-child:turn-child-1:thread-root".to_string(), + child_thread_id: "019d0000-0000-7000-8000-000000000002".to_string(), + child_codex_turn_id: "turn-child-1".to_string(), + parent_thread_id: "019d0000-0000-7000-8000-000000000001".to_string(), + message: notification.to_string(), + carried_payload: Some(carried_payload.clone()), + }, + )?; + + // The parent does receive the failure notification as a model-visible + // mailbox item. The target should remain that precise parent-side + // ConversationItem even though the source falls back to the child thread. + start_agent_turn(&writer, "turn-root-1")?; + let delivered = inter_agent_message( + "/root/child", + "/root", + notification, + /*trigger_turn*/ false, + ); + append_inference_request( + &writer, + "019d0000-0000-7000-8000-000000000001", + "turn-root-1", + "inference-root-1", + vec![message("assistant", &delivered)], + )?; + + let replayed = replay_bundle(temp.path())?; + let edge = + &replayed.interaction_edges["edge:agent_result:thread-child:turn-child-1:thread-root"]; + assert_eq!(edge.kind, InteractionEdgeKind::AgentResult); + assert_eq!( + edge.source, + TraceAnchor::Thread { + thread_id: "019d0000-0000-7000-8000-000000000002".to_string(), + } + ); + let target_item_id = target_conversation_item_id(&edge.target); + assert_eq!( + replayed.conversation_items[target_item_id].thread_id, + "019d0000-0000-7000-8000-000000000001" + ); + assert_eq!(edge.carried_item_ids, vec![target_item_id.clone()]); + assert_eq!( + edge.carried_raw_payload_ids, + vec![carried_payload.raw_payload_id] + ); + + Ok(()) +} + +struct SpawnAgentToolPayloads { + invocation: RawPayloadRef, + begin: RawPayloadRef, + end: RawPayloadRef, + result: RawPayloadRef, +} + +fn append_spawn_agent_tool_lifecycle( + writer: &TraceWriter, + turn_id: &str, +) -> anyhow::Result { + // Keep the parent-side tool lifecycle in one place so the spawn tests can + // focus on the child-side event that decides the edge target. + let invocation = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "spawn_agent", + "payload": { + "type": "function", + "arguments": r#"{"task_name":"repo_file_counter","message":"count"}"# + } + }), + )?; + writer.append_with_context( + trace_context_for_agent(turn_id), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "call-spawn".to_string(), + model_visible_call_id: Some("call-spawn".to_string()), + code_mode_runtime_tool_id: None, + requester: RawToolCallRequester::Model, + kind: ToolCallKind::SpawnAgent, + summary: ToolCallSummary::Generic { + label: "spawn_agent".to_string(), + input_preview: None, + output_preview: None, + }, + invocation_payload: Some(invocation.clone()), + }, + )?; + + let begin = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "call-spawn", + "sender_thread_id": "019d0000-0000-7000-8000-000000000001", + "prompt": "count" + }), + )?; + writer.append_with_context( + trace_context_for_agent(turn_id), + RawTraceEventPayload::ToolCallRuntimeStarted { + tool_call_id: "call-spawn".to_string(), + runtime_payload: begin.clone(), + }, + )?; + + let end = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "call-spawn", + "sender_thread_id": "019d0000-0000-7000-8000-000000000001", + "new_thread_id": "019d0000-0000-7000-8000-000000000002", + "prompt": "count", + "model": "gpt-test", + "reasoning_effort": "medium", + "status": "running" + }), + )?; + writer.append_with_context( + trace_context_for_agent(turn_id), + RawTraceEventPayload::ToolCallRuntimeEnded { + tool_call_id: "call-spawn".to_string(), + status: ExecutionStatus::Completed, + runtime_payload: end.clone(), + }, + )?; + + let result = writer.write_json_payload( + RawPayloadKind::ToolResult, + &json!({"task_name": "/root/repo_file_counter"}), + )?; + writer.append_with_context( + trace_context_for_agent(turn_id), + RawTraceEventPayload::ToolCallEnded { + tool_call_id: "call-spawn".to_string(), + status: ExecutionStatus::Completed, + result_payload: Some(result.clone()), + }, + )?; + + Ok(SpawnAgentToolPayloads { + invocation, + begin, + end, + result, + }) +} + +fn inter_agent_message(author: &str, recipient: &str, content: &str, trigger_turn: bool) -> String { + json!({ + "author": author, + "recipient": recipient, + "other_recipients": [], + "content": content, + "trigger_turn": trigger_turn, + }) + .to_string() +} + +fn target_conversation_item_id(anchor: &TraceAnchor) -> &String { + let TraceAnchor::ConversationItem { item_id } = anchor else { + panic!("expected conversation item target"); + }; + item_id +} + +fn text_body(item: &crate::model::ConversationItem) -> &str { + let [crate::model::ConversationPart::Text { text }] = item.body.parts.as_slice() else { + panic!("expected single text part"); + }; + text +} diff --git a/codex-rs/rollout-trace/src/reducer/tool/terminal.rs b/codex-rs/rollout-trace/src/reducer/tool/terminal.rs new file mode 100644 index 0000000000000000000000000000000000000000..cdd28ec5da1e1ff38d2564092f5b3ff87b7c3fc1 --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/tool/terminal.rs @@ -0,0 +1,606 @@ +//! Terminal reduction for exec-like tool calls. +//! +//! The raw trace records terminal activity as normal tool lifecycle events. +//! Protocol-backed exec events carry `ExecCommand*` payloads with the richest +//! runtime details. Direct tools without protocol observations, such as +//! `write_stdin`, can still form a terminal row from the canonical dispatch +//! invocation/result payloads when those payloads carry the session join key. + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use serde::Deserialize; +use serde_json::Value as JsonValue; + +use super::push_unique; +use crate::model::ExecutionStatus; +use crate::model::ExecutionWindow; +use crate::model::TerminalModelObservation; +use crate::model::TerminalObservationSource; +use crate::model::TerminalOperation; +use crate::model::TerminalOperationId; +use crate::model::TerminalOperationKind; +use crate::model::TerminalRequest; +use crate::model::TerminalResult; +use crate::model::TerminalSession; +use crate::model::ToolCallKind; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawEventSeq; +use crate::reducer::TraceReducer; + +impl TraceReducer { + /// Starts a terminal operation from a canonical dispatch invocation payload. + /// + /// This is currently needed for direct tools such as write-stdin that do not + /// emit a richer protocol runtime-begin event with the terminal join key. + pub(in crate::reducer) fn start_terminal_operation_from_invocation( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + thread_id: &str, + tool_call_id: &str, + kind: &ToolCallKind, + invocation_payload: Option<&RawPayloadRef>, + ) -> Result> { + if !matches!(kind, ToolCallKind::WriteStdin) { + return Ok(None); + } + let operation_kind = TerminalOperationKind::WriteStdin; + let Some(invocation_payload) = invocation_payload else { + // Payload writes are best-effort in the live recorder. If the + // canonical invocation is missing, keep the ToolCall but avoid + // fabricating a lossy terminal row. + return Ok(None); + }; + + let payload = self.read_payload_json(invocation_payload)?; + let request = parse_dispatch_terminal_request(payload).with_context(|| { + format!( + "parse terminal invocation payload {} as dispatch payload", + invocation_payload.raw_payload_id + ) + })?; + self.insert_terminal_operation(TerminalOperationStart { + seq, + wall_time_unix_ms, + thread_id, + tool_call_id, + operation_kind, + raw_payload: invocation_payload, + request, + }) + } + + /// Starts a terminal operation from a protocol runtime-begin payload. + pub(in crate::reducer) fn start_terminal_operation_from_runtime( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + thread_id: &str, + tool_call_id: &str, + kind: &ToolCallKind, + runtime_payload: &RawPayloadRef, + ) -> Result> { + let Some(operation_kind) = terminal_operation_kind(kind) else { + return Ok(None); + }; + + let payload = self.read_payload_json(runtime_payload)?; + let payload: ExecCommandBeginPayload = + serde_json::from_value(payload).with_context(|| { + format!( + "parse terminal runtime start payload {}", + runtime_payload.raw_payload_id + ) + })?; + let request = parse_protocol_terminal_request(payload, &operation_kind); + self.insert_terminal_operation(TerminalOperationStart { + seq, + wall_time_unix_ms, + thread_id, + tool_call_id, + operation_kind, + raw_payload: runtime_payload, + request, + }) + } + + fn insert_terminal_operation( + &mut self, + start: TerminalOperationStart<'_>, + ) -> Result> { + let operation_id = self.next_terminal_operation_id(); + let ParsedTerminalRequest { + terminal_id, + request, + } = start.request; + + self.rollout.terminal_operations.insert( + operation_id.clone(), + TerminalOperation { + operation_id: operation_id.clone(), + terminal_id: terminal_id.clone(), + tool_call_id: start.tool_call_id.to_string(), + kind: start.operation_kind, + execution: ExecutionWindow { + started_at_unix_ms: start.wall_time_unix_ms, + started_seq: start.seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + request, + result: None, + model_observations: Vec::new(), + raw_payload_ids: vec![start.raw_payload.raw_payload_id.clone()], + }, + ); + + if let Some(terminal_id) = terminal_id { + self.ensure_terminal_session( + start.thread_id, + &terminal_id, + &operation_id, + start.wall_time_unix_ms, + start.seq, + )?; + } + + Ok(Some(operation_id)) + } + + /// Completes the terminal operation associated with a tool call, if one exists. + /// + /// Non-terminal tools flow through the same generic tool lifecycle, so callers + /// may invoke this unconditionally and receive Ok for unrelated tool kinds. + pub(in crate::reducer) fn end_terminal_operation( + &mut self, + seq: RawEventSeq, + wall_time_unix_ms: i64, + thread_id: &str, + operation_id: &str, + status: ExecutionStatus, + response_payload: Option<&RawPayloadRef>, + ) -> Result<()> { + let Some(operation_kind) = self + .rollout + .terminal_operations + .get(operation_id) + .map(|operation| operation.kind.clone()) + else { + bail!("terminal end referenced unknown operation {operation_id}"); + }; + let response = response_payload + .map(|payload| { + let value = self.read_payload_json(payload)?; + let response = parse_terminal_response_payload( + value, + &operation_kind, + &payload.raw_payload_id, + )?; + Ok::<_, anyhow::Error>((payload.raw_payload_id.clone(), response)) + }) + .transpose()?; + + let (terminal_id, started_at_unix_ms, started_seq) = { + let Some(operation) = self.rollout.terminal_operations.get_mut(operation_id) else { + bail!("terminal end referenced unknown operation {operation_id}"); + }; + operation.execution.ended_at_unix_ms = Some(wall_time_unix_ms); + operation.execution.ended_seq = Some(seq); + operation.execution.status = status; + + if let Some((raw_payload_id, response)) = response { + push_unique(&mut operation.raw_payload_ids, &raw_payload_id); + // If begin and end both report a process id they must name the + // same terminal. If begin omitted it, the end event completes + // the session join key for this operation. + match (&operation.terminal_id, response.terminal_id.as_deref()) { + (Some(existing), Some(process_id)) if existing != process_id => { + bail!( + "terminal operation {operation_id} changed process id from \ + {existing} to {process_id}" + ); + } + (None, Some(process_id)) => { + operation.terminal_id = Some(process_id.to_string()); + } + (Some(_), Some(_)) | (Some(_), None) | (None, None) => {} + } + operation.result = Some(response.result); + } + + ( + operation.terminal_id.clone(), + operation.execution.started_at_unix_ms, + operation.execution.started_seq, + ) + }; + + if let Some(terminal_id) = terminal_id { + self.ensure_terminal_session( + thread_id, + &terminal_id, + operation_id, + started_at_unix_ms, + started_seq, + )?; + } + + Ok(()) + } + + fn ensure_terminal_session( + &mut self, + thread_id: &str, + terminal_id: &str, + operation_id: &str, + started_at_unix_ms: i64, + started_seq: RawEventSeq, + ) -> Result<()> { + if !self.rollout.terminal_sessions.contains_key(terminal_id) { + self.rollout.terminal_sessions.insert( + terminal_id.to_string(), + TerminalSession { + terminal_id: terminal_id.to_string(), + thread_id: thread_id.to_string(), + created_by_operation_id: operation_id.to_string(), + operation_ids: Vec::new(), + execution: ExecutionWindow { + started_at_unix_ms, + started_seq, + // Current raw events do not report a terminal/session + // shutdown boundary, so the session remains open even + // after individual operations complete. + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + }, + ); + } + + let Some(session) = self.rollout.terminal_sessions.get_mut(terminal_id) else { + bail!("terminal session {terminal_id} disappeared during reduction"); + }; + if session.thread_id != thread_id { + bail!( + "terminal session {terminal_id} belongs to thread {}, not {thread_id}", + session.thread_id + ); + } + push_unique(&mut session.operation_ids, operation_id); + Ok(()) + } + + /// Mirrors model-visible tool items onto the terminal observation view. + /// + /// Runtime terminal rows are useful on their own, but the model-visible call + /// and output item ids let viewers jump between transcript and terminal timelines. + pub(in crate::reducer) fn sync_terminal_model_observation( + &mut self, + tool_call_id: &str, + ) -> Result<()> { + let Some(tool_call) = self.rollout.tool_calls.get(tool_call_id) else { + bail!("tool call {tool_call_id} disappeared during terminal observation linking"); + }; + let Some(operation_id) = tool_call.terminal_operation_id.clone() else { + return Ok(()); + }; + let call_item_ids = tool_call.model_visible_call_item_ids.clone(); + let output_item_ids = tool_call.model_visible_output_item_ids.clone(); + if call_item_ids.is_empty() && output_item_ids.is_empty() { + return Ok(()); + } + + let Some(operation) = self.rollout.terminal_operations.get_mut(&operation_id) else { + bail!("terminal operation {operation_id} disappeared during observation linking"); + }; + // A terminal result and a model-visible tool output are intentionally + // separate: the former is what the runtime saw, the latter is what later + // inference payloads prove was shown back to the model. + if let Some(observation) = operation + .model_observations + .iter_mut() + .find(|observation| observation.source == TerminalObservationSource::DirectToolCall) + { + observation.call_item_ids = call_item_ids; + observation.output_item_ids = output_item_ids; + } else { + operation.model_observations.push(TerminalModelObservation { + call_item_ids, + output_item_ids, + source: TerminalObservationSource::DirectToolCall, + }); + } + Ok(()) + } + + fn next_terminal_operation_id(&mut self) -> TerminalOperationId { + let ordinal = self.next_terminal_operation_ordinal; + self.next_terminal_operation_ordinal += 1; + format!("terminal_operation:{ordinal}") + } +} + +fn terminal_operation_kind(kind: &ToolCallKind) -> Option { + match kind { + ToolCallKind::ExecCommand => Some(TerminalOperationKind::ExecCommand), + ToolCallKind::WriteStdin => Some(TerminalOperationKind::WriteStdin), + ToolCallKind::ApplyPatch + | ToolCallKind::Mcp { .. } + | ToolCallKind::Web + | ToolCallKind::ImageGeneration + | ToolCallKind::SpawnAgent + | ToolCallKind::AssignAgentTask + | ToolCallKind::SendMessage + | ToolCallKind::WaitAgent + | ToolCallKind::CloseAgent + | ToolCallKind::Other { .. } => None, + } +} + +struct TerminalOperationStart<'a> { + seq: RawEventSeq, + wall_time_unix_ms: i64, + thread_id: &'a str, + tool_call_id: &'a str, + operation_kind: TerminalOperationKind, + raw_payload: &'a RawPayloadRef, + request: ParsedTerminalRequest, +} + +struct ParsedTerminalRequest { + terminal_id: Option, + request: TerminalRequest, +} + +struct ParsedTerminalResponse { + terminal_id: Option, + result: TerminalResult, +} + +fn parse_protocol_terminal_request( + payload: ExecCommandBeginPayload, + operation_kind: &TerminalOperationKind, +) -> ParsedTerminalRequest { + // Startup/poll paths usually include a process id at begin time, but plain + // exec starts may only learn it in the matching end event. + let terminal_id = payload.process_id.clone(); + let request = match operation_kind { + TerminalOperationKind::ExecCommand => TerminalRequest::ExecCommand { + display_command: payload.command.join(" "), + command: payload.command, + cwd: payload.cwd, + yield_time_ms: None, + max_output_tokens: None, + }, + TerminalOperationKind::WriteStdin => TerminalRequest::WriteStdin { + stdin: payload.interaction_input.unwrap_or_default(), + yield_time_ms: None, + max_output_tokens: None, + }, + }; + ParsedTerminalRequest { + terminal_id, + request, + } +} + +fn parse_dispatch_terminal_request(value: JsonValue) -> Result { + let payload: DispatchedToolTraceRequestPayload = serde_json::from_value(value)?; + if payload.tool_name != "write_stdin" { + bail!( + "dispatch terminal request is for {}, not write_stdin", + payload.tool_name + ); + } + if payload.payload.kind != "function" { + bail!( + "write_stdin dispatch payload used unsupported {} payload", + payload.payload.kind + ); + } + let arguments = payload + .payload + .arguments + .context("write_stdin dispatch payload omitted function arguments")?; + let args: DispatchedWriteStdinArgs = serde_json::from_str(&arguments) + .context("parse write_stdin dispatch function arguments")?; + let terminal_id = terminal_id_from_json(&args.session_id) + .context("write_stdin dispatch payload omitted session_id")?; + + Ok(ParsedTerminalRequest { + terminal_id: Some(terminal_id), + request: TerminalRequest::WriteStdin { + stdin: args.chars, + yield_time_ms: args.yield_time_ms, + max_output_tokens: args.max_output_tokens, + }, + }) +} + +fn parse_terminal_response_payload( + value: JsonValue, + operation_kind: &TerminalOperationKind, + raw_payload_id: &str, +) -> Result { + match operation_kind { + TerminalOperationKind::ExecCommand => { + let payload = serde_json::from_value::(value) + .with_context(|| format!("parse exec terminal response {raw_payload_id}"))?; + Ok(parse_protocol_terminal_response(payload)) + } + TerminalOperationKind::WriteStdin => { + match serde_json::from_value::(value.clone()) { + Ok(payload) => Ok(parse_protocol_terminal_response(payload)), + Err(protocol_err) => parse_dispatch_terminal_response(value).with_context(|| { + format!( + "parse write_stdin terminal response {raw_payload_id} as protocol payload \ + ({protocol_err}) or dispatch payload" + ) + }), + } + } + } +} + +fn parse_protocol_terminal_response(payload: ExecCommandEndPayload) -> ParsedTerminalResponse { + ParsedTerminalResponse { + terminal_id: payload.process_id, + result: TerminalResult { + exit_code: Some(payload.exit_code), + stdout: payload.stdout, + stderr: payload.stderr, + formatted_output: Some(payload.formatted_output), + original_token_count: None, + chunk_id: None, + }, + } +} + +fn parse_dispatch_terminal_response(value: JsonValue) -> Result { + let payload: DispatchedToolTraceResponsePayload = serde_json::from_value(value)?; + let result = match payload { + DispatchedToolTraceResponsePayload::DirectResponse { response_item } => { + let output = response_item + .get("output") + .and_then(json_text_content) + .unwrap_or_else(|| response_item.to_string()); + TerminalResult { + exit_code: None, + stdout: output.clone(), + stderr: String::new(), + formatted_output: Some(output), + original_token_count: None, + chunk_id: None, + } + } + DispatchedToolTraceResponsePayload::CodeModeResponse { value } => { + // Code-mode returns the JavaScript-facing tool value, not the text + // shown to the model. For write_stdin that value is the structured + // unified-exec result, so keep ToolCall.raw_result_payload_id as the + // raw boundary while projecting terminal-specific fields here. + parse_code_mode_exec_result(value) + } + DispatchedToolTraceResponsePayload::Error { error } => TerminalResult { + exit_code: None, + stdout: String::new(), + stderr: error.clone(), + formatted_output: Some(error), + original_token_count: None, + chunk_id: None, + }, + }; + Ok(ParsedTerminalResponse { + terminal_id: None, + result, + }) +} + +fn parse_code_mode_exec_result(value: JsonValue) -> TerminalResult { + match serde_json::from_value::(value.clone()) { + Ok(result) => TerminalResult { + exit_code: result.exit_code, + stdout: result.output.clone(), + stderr: String::new(), + formatted_output: Some(result.output), + original_token_count: result.original_token_count, + chunk_id: result.chunk_id, + }, + Err(_) => { + let output = json_text_content(&value).unwrap_or_else(|| value.to_string()); + TerminalResult { + exit_code: None, + stdout: output.clone(), + stderr: String::new(), + formatted_output: Some(output), + original_token_count: None, + chunk_id: None, + } + } + } +} + +fn json_text_content(value: &JsonValue) -> Option { + match value { + JsonValue::String(text) => Some(text.clone()), + JsonValue::Array(items) => { + let text = items + .iter() + .filter_map(|item| item.get("text").and_then(JsonValue::as_str)) + .collect::>() + .join("\n"); + (!text.is_empty()).then_some(text) + } + JsonValue::Null => None, + other => Some(other.to_string()), + } +} + +fn terminal_id_from_json(value: &JsonValue) -> Option { + match value { + JsonValue::String(value) if !value.is_empty() => Some(value.clone()), + JsonValue::Number(value) => Some(value.to_string()), + _ => None, + } +} + +#[derive(Deserialize)] +struct ExecCommandBeginPayload { + process_id: Option, + command: Vec, + cwd: String, + interaction_input: Option, +} + +#[derive(Deserialize)] +struct ExecCommandEndPayload { + process_id: Option, + stdout: String, + stderr: String, + exit_code: i32, + formatted_output: String, +} + +#[derive(Deserialize)] +struct DispatchedToolTraceRequestPayload { + tool_name: String, + payload: DispatchedToolPayload, +} + +#[derive(Deserialize)] +struct DispatchedToolPayload { + #[serde(rename = "type")] + kind: String, + arguments: Option, +} + +#[derive(Deserialize)] +struct DispatchedWriteStdinArgs { + session_id: JsonValue, + #[serde(default)] + chars: String, + yield_time_ms: Option, + max_output_tokens: Option, +} + +#[derive(Deserialize)] +#[serde(rename_all = "snake_case", tag = "type")] +enum DispatchedToolTraceResponsePayload { + DirectResponse { response_item: JsonValue }, + CodeModeResponse { value: JsonValue }, + Error { error: String }, +} + +#[derive(Deserialize)] +struct CodeModeExecResult { + chunk_id: Option, + exit_code: Option, + original_token_count: Option, + output: String, +} + +#[cfg(test)] +#[path = "terminal_tests.rs"] +mod tests; diff --git a/codex-rs/rollout-trace/src/reducer/tool/terminal_tests.rs b/codex-rs/rollout-trace/src/reducer/tool/terminal_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..e51e54a740f6ffd79d740cc25d831ead8c7579ba --- /dev/null +++ b/codex-rs/rollout-trace/src/reducer/tool/terminal_tests.rs @@ -0,0 +1,581 @@ +use pretty_assertions::assert_eq; +use serde_json::json; +use tempfile::TempDir; + +use crate::model::ExecutionStatus; +use crate::model::ExecutionWindow; +use crate::model::TerminalModelObservation; +use crate::model::TerminalObservationSource; +use crate::model::TerminalOperation; +use crate::model::TerminalOperationKind; +use crate::model::TerminalRequest; +use crate::model::TerminalResult; +use crate::model::TerminalSession; +use crate::model::ToolCallKind; +use crate::model::ToolCallSummary; +use crate::payload::RawPayloadKind; +use crate::raw_event::RawTraceEventPayload; +use crate::reducer::test_support::create_started_writer; +use crate::reducer::test_support::generic_summary; +use crate::reducer::test_support::message; +use crate::reducer::test_support::start_turn; +use crate::reducer::test_support::trace_context; +use crate::replay_bundle; +use crate::writer::TraceWriter; + +#[test] +fn exec_tool_reduces_to_terminal_operation_and_session() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + append_inference_with_tool_call(&writer)?; + + let invocation_payload = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "exec_command", + "tool_namespace": null, + "payload": { + "type": "function", + "arguments": "{\"cmd\":\"cargo test\"}" + } + }), + )?; + let invocation_payload_id = invocation_payload.raw_payload_id.clone(); + let _tool_start = writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "tool-1".to_string(), + model_visible_call_id: Some("call-1".to_string()), + code_mode_runtime_tool_id: None, + requester: crate::raw_event::RawToolCallRequester::Model, + kind: ToolCallKind::ExecCommand, + summary: generic_summary("exec_command"), + invocation_payload: Some(invocation_payload), + }, + )?; + + let runtime_start_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "tool-1", + "turn_id": "turn-1", + "command": ["cargo", "test"], + "cwd": "/repo" + }), + )?; + let runtime_start_payload_id = runtime_start_payload.raw_payload_id.clone(); + let runtime_start = writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallRuntimeStarted { + tool_call_id: "tool-1".to_string(), + runtime_payload: runtime_start_payload, + }, + )?; + + let runtime_end_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "tool-1", + "process_id": "pty-1", + "turn_id": "turn-1", + "command": ["cargo", "test"], + "cwd": "/repo", + "stdout": "ok\n", + "stderr": "", + "exit_code": 0, + "formatted_output": "ok\n", + "status": "completed" + }), + )?; + let runtime_end_payload_id = runtime_end_payload.raw_payload_id.clone(); + let runtime_end = writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallRuntimeEnded { + tool_call_id: "tool-1".to_string(), + status: ExecutionStatus::Completed, + runtime_payload: runtime_end_payload, + }, + )?; + + let result_payload = writer.write_json_payload( + RawPayloadKind::ToolResult, + &json!({ + "type": "direct_response", + "response_item": { + "type": "function_call_output", + "call_id": "call-1", + "output": "ok\n" + } + }), + )?; + let result_payload_id = result_payload.raw_payload_id.clone(); + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallEnded { + tool_call_id: "tool-1".to_string(), + status: ExecutionStatus::Completed, + result_payload: Some(result_payload), + }, + )?; + + start_turn(&writer, "turn-2")?; + append_followup_with_tool_output(&writer)?; + + let rollout = replay_bundle(temp.path())?; + let operation_id = "terminal_operation:1".to_string(); + let output_item_id = rollout.inference_calls["inference-2"] + .request_item_ids + .last() + .expect("tool output item") + .clone(); + + assert_eq!( + rollout.tool_calls["tool-1"].terminal_operation_id, + Some(operation_id.clone()), + ); + assert_eq!( + rollout.tool_calls["tool-1"].raw_invocation_payload_id, + Some(invocation_payload_id), + ); + assert_eq!( + rollout.tool_calls["tool-1"].raw_result_payload_id, + Some(result_payload_id), + ); + assert_eq!( + rollout.tool_calls["tool-1"].raw_runtime_payload_ids, + vec![ + runtime_start_payload_id.clone(), + runtime_end_payload_id.clone() + ], + ); + assert_eq!( + rollout.tool_calls["tool-1"].summary, + ToolCallSummary::Terminal { + operation_id: operation_id.clone(), + }, + ); + assert_eq!( + rollout.terminal_operations[&operation_id], + TerminalOperation { + operation_id: operation_id.clone(), + terminal_id: Some("pty-1".to_string()), + tool_call_id: "tool-1".to_string(), + kind: TerminalOperationKind::ExecCommand, + execution: ExecutionWindow { + started_at_unix_ms: runtime_start.wall_time_unix_ms, + started_seq: runtime_start.seq, + ended_at_unix_ms: Some(runtime_end.wall_time_unix_ms), + ended_seq: Some(runtime_end.seq), + status: ExecutionStatus::Completed, + }, + request: TerminalRequest::ExecCommand { + command: vec!["cargo".to_string(), "test".to_string()], + display_command: "cargo test".to_string(), + cwd: "/repo".to_string(), + yield_time_ms: None, + max_output_tokens: None, + }, + result: Some(TerminalResult { + exit_code: Some(0), + stdout: "ok\n".to_string(), + stderr: String::new(), + formatted_output: Some("ok\n".to_string()), + original_token_count: None, + chunk_id: None, + }), + model_observations: vec![TerminalModelObservation { + call_item_ids: rollout.inference_calls["inference-1"] + .response_item_ids + .clone(), + output_item_ids: vec![output_item_id], + source: TerminalObservationSource::DirectToolCall, + }], + raw_payload_ids: vec![runtime_start_payload_id, runtime_end_payload_id], + }, + ); + assert_eq!( + rollout.terminal_sessions["pty-1"], + TerminalSession { + terminal_id: "pty-1".to_string(), + thread_id: "thread-root".to_string(), + created_by_operation_id: operation_id.clone(), + operation_ids: vec![operation_id], + execution: ExecutionWindow { + started_at_unix_ms: runtime_start.wall_time_unix_ms, + started_seq: runtime_start.seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + }, + ); + + Ok(()) +} + +#[test] +fn write_stdin_operation_reuses_existing_terminal_session() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let startup_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "tool-start", + "process_id": "pty-1", + "turn_id": "turn-1", + "command": ["bash"], + "cwd": "/repo" + }), + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "tool-start".to_string(), + model_visible_call_id: None, + code_mode_runtime_tool_id: None, + requester: crate::raw_event::RawToolCallRequester::Model, + kind: ToolCallKind::ExecCommand, + summary: generic_summary("exec_command"), + invocation_payload: None, + }, + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallRuntimeStarted { + tool_call_id: "tool-start".to_string(), + runtime_payload: startup_payload, + }, + )?; + + let stdin_payload = writer.write_json_payload( + RawPayloadKind::ToolRuntimeEvent, + &json!({ + "call_id": "tool-stdin", + "process_id": "pty-1", + "turn_id": "turn-1", + "command": ["bash"], + "cwd": "/repo", + "interaction_input": "echo hi\n" + }), + )?; + let _stdin_start = writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "tool-stdin".to_string(), + model_visible_call_id: None, + code_mode_runtime_tool_id: None, + requester: crate::raw_event::RawToolCallRequester::Model, + kind: ToolCallKind::WriteStdin, + summary: generic_summary("write_stdin"), + invocation_payload: None, + }, + )?; + let stdin_runtime_start = writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallRuntimeStarted { + tool_call_id: "tool-stdin".to_string(), + runtime_payload: stdin_payload, + }, + )?; + + let rollout = replay_bundle(temp.path())?; + let startup_operation_id = "terminal_operation:1".to_string(); + let stdin_operation_id = "terminal_operation:2".to_string(); + + assert_eq!( + rollout.terminal_sessions["pty-1"].operation_ids, + vec![startup_operation_id, stdin_operation_id.clone()], + ); + assert_eq!( + rollout.terminal_operations[&stdin_operation_id], + TerminalOperation { + operation_id: stdin_operation_id.clone(), + terminal_id: Some("pty-1".to_string()), + tool_call_id: "tool-stdin".to_string(), + kind: TerminalOperationKind::WriteStdin, + execution: ExecutionWindow { + started_at_unix_ms: stdin_runtime_start.wall_time_unix_ms, + started_seq: stdin_runtime_start.seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + request: TerminalRequest::WriteStdin { + stdin: "echo hi\n".to_string(), + yield_time_ms: None, + max_output_tokens: None, + }, + result: None, + model_observations: Vec::new(), + raw_payload_ids: vec!["raw_payload:2".to_string()], + }, + ); + + Ok(()) +} + +#[test] +fn dispatch_write_stdin_payload_reduces_to_terminal_operation() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request_payload = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "write_stdin", + "tool_namespace": null, + "payload": { + "type": "function", + "arguments": json!({ + "session_id": 123, + "chars": "echo hi\n", + "yield_time_ms": 250, + "max_output_tokens": 2000 + }).to_string() + } + }), + )?; + let request_payload_id = request_payload.raw_payload_id.clone(); + let tool_start = writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "tool-stdin".to_string(), + model_visible_call_id: Some("call-stdin".to_string()), + code_mode_runtime_tool_id: None, + requester: crate::raw_event::RawToolCallRequester::Model, + kind: ToolCallKind::WriteStdin, + summary: generic_summary("write_stdin"), + invocation_payload: Some(request_payload), + }, + )?; + + let response_payload = writer.write_json_payload( + RawPayloadKind::ToolResult, + &json!({ + "type": "direct_response", + "response_item": { + "type": "function_call_output", + "call_id": "call-stdin", + "output": "hi\n" + } + }), + )?; + let response_payload_id = response_payload.raw_payload_id.clone(); + let tool_end = writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallEnded { + tool_call_id: "tool-stdin".to_string(), + status: ExecutionStatus::Completed, + result_payload: Some(response_payload), + }, + )?; + + let rollout = replay_bundle(temp.path())?; + let operation_id = "terminal_operation:1".to_string(); + + assert_eq!( + rollout.tool_calls["tool-stdin"].terminal_operation_id, + Some(operation_id.clone()), + ); + assert_eq!( + rollout.tool_calls["tool-stdin"].summary, + ToolCallSummary::Terminal { + operation_id: operation_id.clone(), + }, + ); + assert_eq!( + rollout.terminal_operations[&operation_id], + TerminalOperation { + operation_id: operation_id.clone(), + terminal_id: Some("123".to_string()), + tool_call_id: "tool-stdin".to_string(), + kind: TerminalOperationKind::WriteStdin, + execution: ExecutionWindow { + started_at_unix_ms: tool_start.wall_time_unix_ms, + started_seq: tool_start.seq, + ended_at_unix_ms: Some(tool_end.wall_time_unix_ms), + ended_seq: Some(tool_end.seq), + status: ExecutionStatus::Completed, + }, + request: TerminalRequest::WriteStdin { + stdin: "echo hi\n".to_string(), + yield_time_ms: Some(250), + max_output_tokens: Some(2000), + }, + result: Some(TerminalResult { + exit_code: None, + stdout: "hi\n".to_string(), + stderr: String::new(), + formatted_output: Some("hi\n".to_string()), + original_token_count: None, + chunk_id: None, + }), + model_observations: Vec::new(), + raw_payload_ids: vec![request_payload_id, response_payload_id], + }, + ); + assert_eq!( + rollout.terminal_sessions["123"], + TerminalSession { + terminal_id: "123".to_string(), + thread_id: "thread-root".to_string(), + created_by_operation_id: operation_id.clone(), + operation_ids: vec![operation_id], + execution: ExecutionWindow { + started_at_unix_ms: tool_start.wall_time_unix_ms, + started_seq: tool_start.seq, + ended_at_unix_ms: None, + ended_seq: None, + status: ExecutionStatus::Running, + }, + }, + ); + + Ok(()) +} + +#[test] +fn code_mode_write_stdin_result_projects_structured_exec_fields() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = create_started_writer(&temp)?; + start_turn(&writer, "turn-1")?; + + let request_payload = writer.write_json_payload( + RawPayloadKind::ToolInvocation, + &json!({ + "tool_name": "write_stdin", + "tool_namespace": null, + "payload": { + "type": "function", + "arguments": json!({ + "session_id": 456, + "chars": "", + "yield_time_ms": 1000, + "max_output_tokens": 4000 + }).to_string() + } + }), + )?; + let response_payload = writer.write_json_payload( + RawPayloadKind::ToolResult, + &json!({ + "type": "code_mode_response", + "value": { + "chunk_id": "abc123", + "wall_time_seconds": 1.25, + "exit_code": 0, + "original_token_count": 3, + "output": "done\n" + } + }), + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::CodeCellStarted { + runtime_cell_id: "cell-1".to_string(), + model_visible_call_id: "call-code".to_string(), + source_js: "await tools.write_stdin({ chars: '' })".to_string(), + }, + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallStarted { + tool_call_id: "tool-stdin".to_string(), + model_visible_call_id: None, + code_mode_runtime_tool_id: Some("runtime-tool-1".to_string()), + requester: crate::raw_event::RawToolCallRequester::CodeCell { + runtime_cell_id: "cell-1".to_string(), + }, + kind: ToolCallKind::WriteStdin, + summary: generic_summary("write_stdin"), + invocation_payload: Some(request_payload), + }, + )?; + writer.append_with_context( + trace_context("turn-1"), + RawTraceEventPayload::ToolCallEnded { + tool_call_id: "tool-stdin".to_string(), + status: ExecutionStatus::Completed, + result_payload: Some(response_payload), + }, + )?; + + let rollout = replay_bundle(temp.path())?; + assert_eq!( + rollout.terminal_operations["terminal_operation:1"].result, + Some(TerminalResult { + exit_code: Some(0), + stdout: "done\n".to_string(), + stderr: String::new(), + formatted_output: Some("done\n".to_string()), + original_token_count: Some(3), + chunk_id: Some("abc123".to_string()), + }), + ); + + Ok(()) +} + +fn append_inference_with_tool_call(writer: &TraceWriter) -> anyhow::Result<()> { + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "input": [message("user", "run tests")] + }), + )?; + writer.append(RawTraceEventPayload::InferenceStarted { + inference_call_id: "inference-1".to_string(), + thread_id: "thread-root".to_string(), + codex_turn_id: "turn-1".to_string(), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + request_payload: request, + })?; + + let response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [{ + "type": "function_call", + "name": "exec_command", + "arguments": "{\"cmd\":\"cargo test\"}", + "call_id": "call-1" + }] + }), + )?; + writer.append(RawTraceEventPayload::InferenceCompleted { + inference_call_id: "inference-1".to_string(), + response_id: Some("resp-1".to_string()), + upstream_request_id: None, + response_payload: response, + })?; + Ok(()) +} + +fn append_followup_with_tool_output(writer: &TraceWriter) -> anyhow::Result<()> { + let request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "previous_response_id": "resp-1", + "input": [{ + "type": "function_call_output", + "call_id": "call-1", + "output": "ok\n" + }] + }), + )?; + writer.append(RawTraceEventPayload::InferenceStarted { + inference_call_id: "inference-2".to_string(), + thread_id: "thread-root".to_string(), + codex_turn_id: "turn-2".to_string(), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + request_payload: request, + })?; + Ok(()) +} diff --git a/codex-rs/rollout-trace/src/thread.rs b/codex-rs/rollout-trace/src/thread.rs new file mode 100644 index 0000000000000000000000000000000000000000..514b8d382bc06840b1d41bdcfd323fa73ee1a5b4 --- /dev/null +++ b/codex-rs/rollout-trace/src/thread.rs @@ -0,0 +1,529 @@ +//! Thread-scoped rollout trace helpers. +//! +//! A rollout bundle can contain a root thread plus spawned child threads. This +//! context owns the stable identity for one thread inside that bundle. Keeping +//! thread-local event methods here avoids repeatedly plumbing `thread_id` +//! through session code. + +use codex_protocol::protocol::AgentStatus; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::SessionSource; +use serde::Serialize; +use std::path::Path; +use std::path::PathBuf; +use std::sync::Arc; +use tracing::debug; +use tracing::warn; +use uuid::Uuid; + +use crate::AgentThreadId; +use crate::CodeCellTraceContext; +use crate::CodexTurnId; +use crate::CompactionId; +use crate::CompactionTraceContext; +use crate::InferenceTraceContext; +use crate::McpCallTraceContext; +use crate::RawPayloadKind; +use crate::RawPayloadRef; +use crate::RawTraceEventContext; +use crate::RawTraceEventPayload; +use crate::RolloutStatus; +use crate::ToolCallId; +use crate::ToolDispatchInvocation; +use crate::ToolDispatchTraceContext; +use crate::TraceWriter; +use crate::protocol_event::codex_turn_trace_event; +use crate::protocol_event::tool_runtime_trace_event; +use crate::protocol_event::wrapped_protocol_event_type; + +/// Environment variable that enables local trace-bundle recording. +/// +/// The value is a root directory. Each independent root session gets one child +/// bundle directory. Spawned child threads share their root session's bundle so +/// one reduced `state.json` describes the whole multi-agent rollout tree. +pub const CODEX_ROLLOUT_TRACE_ROOT_ENV: &str = "CODEX_ROLLOUT_TRACE_ROOT"; + +/// Metadata captured once at thread/session start. +/// +/// This payload is intentionally operational rather than reduced: it is a raw +/// payload that later reducers can mine as the reduced thread model evolves. +#[derive(Serialize)] +pub struct ThreadStartedTraceMetadata { + pub thread_id: String, + pub agent_path: String, + pub task_name: Option, + pub nickname: Option, + pub agent_role: Option, + pub session_source: SessionSource, + pub cwd: std::path::PathBuf, + pub rollout_path: Option, + pub model: String, + pub provider_name: String, + pub approval_policy: String, + pub sandbox_policy: String, +} + +/// Trace-only payload for a child completion notification delivered to its parent. +#[derive(Serialize)] +pub struct AgentResultTracePayload<'a> { + pub child_agent_path: &'a str, + pub message: &'a str, + pub status: &'a AgentStatus, +} + +/// No-op capable trace handle for one thread in a rollout bundle. +#[derive(Clone, Debug)] +pub struct ThreadTraceContext { + state: ThreadTraceContextState, +} + +#[derive(Clone, Debug)] +enum ThreadTraceContextState { + Disabled, + Enabled(EnabledThreadTraceContext), +} + +#[derive(Clone, Debug)] +struct EnabledThreadTraceContext { + writer: Arc, + root_thread_id: AgentThreadId, + thread_id: AgentThreadId, +} + +impl ThreadTraceContext { + /// Builds a context that accepts trace calls and records nothing. + pub fn disabled() -> Self { + Self { + state: ThreadTraceContextState::Disabled, + } + } + + /// Starts a root thread trace from `CODEX_ROLLOUT_TRACE_ROOT`, or disables tracing. + /// + /// Trace startup is best-effort. A tracing failure must not make the Codex + /// session unusable, because traces are diagnostic and can be enabled while + /// debugging unrelated production failures. + pub fn start_root_or_disabled(metadata: ThreadStartedTraceMetadata) -> Self { + let Some(root) = std::env::var_os(CODEX_ROLLOUT_TRACE_ROOT_ENV) else { + return Self::disabled(); + }; + let root = PathBuf::from(root); + match start_root_in_root(root.as_path(), metadata) { + Ok(context) => context, + Err(err) => { + warn!("failed to initialize rollout trace bundle: {err:#}"); + Self::disabled() + } + } + } + + /// Starts a root trace in a known directory. + /// + /// This is public for tests that need replayable trace bundles without + /// mutating process environment. + pub fn start_root_in_root_for_test( + root: &Path, + metadata: ThreadStartedTraceMetadata, + ) -> anyhow::Result { + start_root_in_root(root, metadata) + } + + /// Starts one thread lifecycle inside an existing rollout bundle. + pub(crate) fn start( + writer: Arc, + root_thread_id: AgentThreadId, + metadata: ThreadStartedTraceMetadata, + ) -> Self { + let context = EnabledThreadTraceContext { + writer, + root_thread_id, + thread_id: metadata.thread_id.clone(), + }; + record_thread_started(&context, metadata); + Self { + state: ThreadTraceContextState::Enabled(context), + } + } + + /// Returns whether this handle will write trace events. + /// + /// Most methods have their own disabled fast path. Callers should branch on + /// this only when preparing trace payloads would otherwise clone data the + /// production path needs to move elsewhere. + pub fn is_enabled(&self) -> bool { + matches!(self.state, ThreadTraceContextState::Enabled(_)) + } + + /// Starts a fresh child thread in this context's rollout tree. + /// + /// Callers should use [`ThreadTraceContext::disabled`] for resumed children: + /// reusing the parent trace would emit a duplicate `ThreadStarted` event + /// for an existing thread id and make the bundle unreplayable. + pub fn start_child_thread_trace_or_disabled( + &self, + metadata: ThreadStartedTraceMetadata, + ) -> Self { + match &self.state { + ThreadTraceContextState::Disabled => Self::disabled(), + ThreadTraceContextState::Enabled(context) => Self::start( + Arc::clone(&context.writer), + context.root_thread_id.clone(), + metadata, + ), + } + } + + /// Emits terminal trace events for graceful thread shutdown. + /// + /// Spawned child sessions share their root bundle, so only the root + /// thread end closes the rollout. Child thread ends update the child thread + /// execution state without marking the whole bundle complete. + pub fn record_ended(&self, status: RolloutStatus) { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return; + }; + context.append_best_effort(RawTraceEventPayload::ThreadEnded { + thread_id: context.thread_id.clone(), + status: status.clone(), + }); + if context.thread_id == context.root_thread_id { + context.append_best_effort(RawTraceEventPayload::RolloutEnded { status }); + } + } + + /// Wraps selected protocol events as raw trace breadcrumbs. + /// + /// High-volume stream deltas stay out of this wrapper; typed inference, + /// tool, terminal, and code-mode hooks provide the canonical runtime data. + pub fn record_protocol_event(&self, event: &EventMsg) { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return; + }; + let Some(event_type) = wrapped_protocol_event_type(event) else { + return; + }; + let Some(event_payload) = + context.write_json_payload_best_effort(RawPayloadKind::ProtocolEvent, event) + else { + return; + }; + context.append_best_effort(RawTraceEventPayload::ProtocolEventObserved { + event_type: event_type.to_string(), + event_payload, + }); + } + + /// Emits typed Codex turn lifecycle events from protocol lifecycle events. + pub fn record_codex_turn_event(&self, default_turn_id: &str, event: &EventMsg) { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return; + }; + let Some(trace_event) = + codex_turn_trace_event(context.thread_id.clone(), default_turn_id, event) + else { + return; + }; + context.append_with_context_best_effort( + trace_event.context_turn_id.clone(), + trace_event.payload, + ); + } + + /// Emits typed runtime tool events from existing protocol lifecycle events. + /// + /// These events are runtime observations on an already-dispatched tool. The + /// dispatch trace records the caller-facing boundary; these payloads explain + /// what Codex did while executing that boundary. + pub fn record_tool_call_event(&self, codex_turn_id: impl Into, event: &EventMsg) { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return; + }; + let Some(trace_event) = tool_runtime_trace_event(event) else { + return; + }; + let Some(payload) = context.raw_tool_runtime_payload(trace_event) else { + return; + }; + context.append_with_context_best_effort(codex_turn_id.into(), payload); + } + + /// Emits the v2 child-to-parent completion message as an explicit graph edge. + /// + /// The notification is runtime delivery from a completed child turn into + /// the parent's mailbox, not a tool call executed by the child. Recording it + /// directly preserves timing and source without making the reducer infer + /// the edge from a later parent prompt snapshot. + pub fn record_agent_result_interaction( + &self, + child_codex_turn_id: impl Into, + parent_thread_id: impl Into, + payload: &AgentResultTracePayload<'_>, + ) { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return; + }; + let child_codex_turn_id = child_codex_turn_id.into(); + let parent_thread_id = parent_thread_id.into(); + let carried_payload = + context.write_json_payload_best_effort(RawPayloadKind::AgentResult, payload); + context.append_with_context_best_effort( + child_codex_turn_id.clone(), + RawTraceEventPayload::AgentResultObserved { + edge_id: format!( + "edge:agent_result:{}:{child_codex_turn_id}:{parent_thread_id}", + context.thread_id + ), + child_thread_id: context.thread_id.clone(), + child_codex_turn_id, + parent_thread_id, + message: payload.message.to_string(), + carried_payload, + }, + ); + } + + /// Emits a turn-start lifecycle event. + /// + /// Most production turn lifecycle wiring lives outside this PR layer, but + /// trace-focused integration tests need a small explicit hook so reducer + /// inputs remain valid without exercising the full session loop. + pub fn record_codex_turn_started(&self, codex_turn_id: impl Into) { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return; + }; + let codex_turn_id = codex_turn_id.into(); + context.append_with_context_best_effort( + codex_turn_id.clone(), + RawTraceEventPayload::CodexTurnStarted { + codex_turn_id, + thread_id: context.thread_id.clone(), + }, + ); + } + + /// Starts a first-class code-mode cell lifecycle and returns its trace handle. + pub fn start_code_cell_trace( + &self, + codex_turn_id: impl Into, + runtime_cell_id: impl Into, + model_visible_call_id: impl Into, + source_js: impl Into, + ) -> CodeCellTraceContext { + let context = self.code_cell_trace_context(codex_turn_id, runtime_cell_id); + context.record_started(model_visible_call_id, source_js); + context + } + + /// Builds a trace handle for an already-started code-mode runtime cell. + pub fn code_cell_trace_context( + &self, + codex_turn_id: impl Into, + runtime_cell_id: impl Into, + ) -> CodeCellTraceContext { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return CodeCellTraceContext::disabled(); + }; + CodeCellTraceContext::enabled( + Arc::clone(&context.writer), + context.thread_id.clone(), + codex_turn_id, + runtime_cell_id, + ) + } + + /// Starts one dispatch-level tool lifecycle and returns its trace handle. + /// + /// `invocation` is lazy because adapting core tool objects into trace-owned + /// payloads can clone large arguments. Disabled tracing should not pay that + /// cost on the hot tool-dispatch path. + pub fn start_tool_dispatch_trace( + &self, + invocation: impl FnOnce() -> Option, + ) -> ToolDispatchTraceContext { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return ToolDispatchTraceContext::disabled(); + }; + let Some(invocation) = invocation() else { + return ToolDispatchTraceContext::disabled(); + }; + ToolDispatchTraceContext::start(Arc::clone(&context.writer), invocation) + } + + /// Builds reusable inference trace context for one Codex turn. + /// + /// The returned context is intentionally not "an inference call" yet. + /// Transport code owns retry/fallback attempts and calls `start_attempt` + /// only after it has built the concrete request payload for that attempt. + pub fn inference_trace_context( + &self, + codex_turn_id: impl Into, + model: impl Into, + provider_name: impl Into, + ) -> InferenceTraceContext { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return InferenceTraceContext::disabled(); + }; + InferenceTraceContext::enabled( + Arc::clone(&context.writer), + context.thread_id.clone(), + codex_turn_id.into(), + model.into(), + provider_name.into(), + ) + } + + /// Builds remote-compaction trace context for one checkpoint. + /// + /// Rollout tracing currently has a first-class checkpoint model only for remote compaction. + /// The compact endpoint is a model-facing request whose output replaces live history, so it + /// needs both request/response attempt events and a later checkpoint event when processed + /// replacement history is installed. + pub fn compaction_trace_context( + &self, + codex_turn_id: impl Into, + compaction_id: impl Into, + model: impl Into, + provider_name: impl Into, + ) -> CompactionTraceContext { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return CompactionTraceContext::disabled(); + }; + CompactionTraceContext::enabled( + Arc::clone(&context.writer), + context.thread_id.clone(), + codex_turn_id.into(), + compaction_id.into(), + model.into(), + provider_name.into(), + ) + } + + /// Starts bridge correlation for one concrete MCP backend request. + /// + /// Dispatch-level tool IDs remain compact and UI-friendly. This UUID is + /// deliberately separate: it is only for cross-process log joins where a + /// rollout-local counter would collide across samples. + pub fn start_mcp_call_trace(&self, tool_call_id: impl Into) -> McpCallTraceContext { + let ThreadTraceContextState::Enabled(context) = &self.state else { + return McpCallTraceContext::disabled(); + }; + let mcp_call_id = Uuid::new_v4().to_string(); + let trace = McpCallTraceContext::enabled(mcp_call_id.clone()); + context.append_best_effort(RawTraceEventPayload::McpToolCallCorrelationAssigned { + tool_call_id: tool_call_id.into(), + mcp_call_id, + }); + trace + } +} + +fn start_root_in_root( + root: &Path, + metadata: ThreadStartedTraceMetadata, +) -> anyhow::Result { + let trace_id = Uuid::new_v4().to_string(); + let thread_id = metadata.thread_id.clone(); + let bundle_dir = root.join(format!("trace-{trace_id}-{thread_id}")); + let writer = TraceWriter::create( + &bundle_dir, + trace_id.clone(), + thread_id.clone(), + thread_id.clone(), + )?; + let writer = Arc::new(writer); + + if let Err(err) = writer.append(RawTraceEventPayload::RolloutStarted { + trace_id, + root_thread_id: thread_id.clone(), + }) { + warn!("failed to append rollout trace event: {err:#}"); + } + + debug!("recording rollout trace at {}", bundle_dir.display()); + Ok(ThreadTraceContext::start(writer, thread_id, metadata)) +} + +fn record_thread_started( + context: &EnabledThreadTraceContext, + metadata: ThreadStartedTraceMetadata, +) { + let metadata_payload = + context.write_json_payload_best_effort(RawPayloadKind::SessionMetadata, &metadata); + context.append_best_effort(RawTraceEventPayload::ThreadStarted { + thread_id: metadata.thread_id, + agent_path: metadata.agent_path, + metadata_payload, + }); +} + +impl EnabledThreadTraceContext { + fn write_json_payload_best_effort( + &self, + kind: RawPayloadKind, + payload: &impl Serialize, + ) -> Option { + match self.writer.write_json_payload(kind, payload) { + Ok(payload_ref) => Some(payload_ref), + Err(err) => { + warn!("failed to write rollout trace payload: {err:#}"); + None + } + } + } + + fn raw_tool_runtime_payload( + &self, + trace_event: crate::protocol_event::ToolRuntimeTraceEvent<'_>, + ) -> Option { + match trace_event { + crate::protocol_event::ToolRuntimeTraceEvent::Started { + tool_call_id, + payload, + } => { + let runtime_payload = self + .write_json_payload_best_effort(RawPayloadKind::ToolRuntimeEvent, &payload)?; + Some(RawTraceEventPayload::ToolCallRuntimeStarted { + tool_call_id: tool_call_id.to_string(), + runtime_payload, + }) + } + crate::protocol_event::ToolRuntimeTraceEvent::Ended { + tool_call_id, + status, + payload, + } => { + let runtime_payload = self + .write_json_payload_best_effort(RawPayloadKind::ToolRuntimeEvent, &payload)?; + Some(RawTraceEventPayload::ToolCallRuntimeEnded { + tool_call_id: tool_call_id.to_string(), + status, + runtime_payload, + }) + } + } + } + + fn append_best_effort(&self, payload: RawTraceEventPayload) { + if let Err(err) = self.writer.append(payload) { + warn!("failed to append rollout trace event: {err:#}"); + } + } + + fn append_with_context_best_effort( + &self, + codex_turn_id: CodexTurnId, + payload: RawTraceEventPayload, + ) { + let event_context = RawTraceEventContext { + thread_id: Some(self.thread_id.clone()), + codex_turn_id: Some(codex_turn_id), + }; + if let Err(err) = self.writer.append_with_context(event_context, payload) { + warn!("failed to append rollout trace event: {err:#}"); + } + } +} + +#[cfg(test)] +#[path = "thread_tests.rs"] +mod tests; diff --git a/codex-rs/rollout-trace/src/thread_tests.rs b/codex-rs/rollout-trace/src/thread_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..e0a5547c334ac1ad9d6845336854d5537969196f --- /dev/null +++ b/codex-rs/rollout-trace/src/thread_tests.rs @@ -0,0 +1,252 @@ +use std::cell::Cell; +use std::fs; +use std::path::Path; +use std::path::PathBuf; + +use codex_protocol::AgentPath; +use codex_protocol::ThreadId; +use codex_protocol::protocol::AgentStatus; +use codex_protocol::protocol::EventMsg; +use codex_protocol::protocol::SandboxPolicy; +use codex_protocol::protocol::SessionSource; +use codex_protocol::protocol::SubAgentSource; +use tempfile::TempDir; + +use super::*; +use crate::AgentResultTracePayload; +use crate::CompactionCheckpointTracePayload; +use crate::ExecutionStatus; +use crate::RawTraceEventPayload; +use crate::RolloutStatus; +use crate::replay_bundle; + +#[test] +fn create_in_root_writes_replayable_lifecycle_events() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let thread_id = ThreadId::new(); + let thread_trace = ThreadTraceContext::start_root_in_root_for_test( + temp.path(), + ThreadStartedTraceMetadata { + thread_id: thread_id.to_string(), + agent_path: "/root".to_string(), + task_name: None, + nickname: None, + agent_role: None, + session_source: SessionSource::Exec, + cwd: PathBuf::from("/workspace"), + rollout_path: Some(PathBuf::from("/tmp/rollout.jsonl")), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + approval_policy: "never".to_string(), + sandbox_policy: format!("{:?}", SandboxPolicy::DangerFullAccess), + }, + )?; + + thread_trace.record_ended(RolloutStatus::Completed); + + let bundle_dir = single_bundle_dir(temp.path())?; + let replayed = replay_bundle(&bundle_dir)?; + + assert_eq!(replayed.status, RolloutStatus::Completed); + assert_eq!(replayed.root_thread_id, thread_id.to_string()); + assert_eq!(replayed.threads[&thread_id.to_string()].agent_path, "/root"); + assert_eq!(replayed.raw_payloads.len(), 1); + + Ok(()) +} + +#[test] +fn spawned_thread_start_appends_to_root_bundle() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let root_thread_id = ThreadId::new(); + let child_thread_id = ThreadId::new(); + let root_trace = ThreadTraceContext::start_root_in_root_for_test( + temp.path(), + minimal_metadata(root_thread_id), + )?; + + let child_trace = root_trace.start_child_thread_trace_or_disabled(ThreadStartedTraceMetadata { + thread_id: child_thread_id.to_string(), + agent_path: "/root/repo_file_counter".to_string(), + task_name: Some("repo_file_counter".to_string()), + nickname: Some("Kepler".to_string()), + agent_role: Some("worker".to_string()), + session_source: SessionSource::SubAgent(SubAgentSource::ThreadSpawn { + parent_thread_id: root_thread_id, + depth: 1, + agent_path: Some( + AgentPath::try_from("/root/repo_file_counter").map_err(anyhow::Error::msg)?, + ), + agent_nickname: Some("Kepler".to_string()), + agent_role: Some("worker".to_string()), + }), + cwd: PathBuf::from("/workspace"), + rollout_path: Some(PathBuf::from("/tmp/child-rollout.jsonl")), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + approval_policy: "never".to_string(), + sandbox_policy: format!("{:?}", SandboxPolicy::DangerFullAccess), + }); + child_trace.record_ended(RolloutStatus::Completed); + let bundle_dir = single_bundle_dir(temp.path())?; + let replayed = replay_bundle(&bundle_dir)?; + + assert_eq!(fs::read_dir(temp.path())?.count(), 1); + assert_eq!(replayed.threads.len(), 2); + assert_eq!( + replayed.threads[&child_thread_id.to_string()].agent_path, + "/root/repo_file_counter" + ); + assert_eq!(replayed.status, RolloutStatus::Running); + assert_eq!( + replayed.threads[&child_thread_id.to_string()] + .execution + .status, + ExecutionStatus::Completed + ); + assert_eq!(replayed.raw_payloads.len(), 2); + + Ok(()) +} + +#[test] +fn disabled_thread_context_accepts_trace_calls_without_writing() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let thread_trace = ThreadTraceContext::disabled(); + + thread_trace.record_ended(RolloutStatus::Completed); + thread_trace.record_protocol_event(&EventMsg::ShutdownComplete); + thread_trace.record_codex_turn_event("turn-1", &EventMsg::ShutdownComplete); + thread_trace.record_tool_call_event("turn-1", &EventMsg::ShutdownComplete); + thread_trace.record_agent_result_interaction( + "turn-1", + ThreadId::new(), + &AgentResultTracePayload { + child_agent_path: "/root/child", + message: "done", + status: &AgentStatus::Completed(Some("done".to_string())), + }, + ); + + let inference_trace = + thread_trace.inference_trace_context("turn-1", "gpt-test", "test-provider"); + let inference_attempt = inference_trace.start_attempt(); + inference_attempt.record_started(&serde_json::json!({ "kind": "inference" })); + let token_usage: Option = None; + inference_attempt.record_completed("response-1", Some("req-1"), &token_usage, &[]); + inference_attempt.record_failed("inference failed", /*upstream_request_id*/ None, &[]); + + let compaction_trace = thread_trace.compaction_trace_context( + "turn-1", + "compaction-1", + "gpt-test", + "test-provider", + ); + assert!(!compaction_trace.is_enabled()); + let compaction_attempt = + compaction_trace.start_attempt(&serde_json::json!({ "kind": "compaction" })); + compaction_attempt.record_completed(&[]); + compaction_attempt.record_failed("compaction failed"); + compaction_trace.record_installed(&CompactionCheckpointTracePayload { + input_history: &[], + replacement_history: &[], + }); + + let built_dispatch_invocation = Cell::new(false); + let dispatch_trace = thread_trace.start_tool_dispatch_trace(|| { + built_dispatch_invocation.set(true); + None + }); + assert!(!built_dispatch_invocation.get()); + assert!(!dispatch_trace.is_enabled()); + + assert_eq!(fs::read_dir(temp.path())?.count(), 0); + + Ok(()) +} + +#[test] +fn compaction_contexts_share_identity_across_models() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let thread_id = ThreadId::new(); + let thread_trace = + ThreadTraceContext::start_root_in_root_for_test(temp.path(), minimal_metadata(thread_id))?; + thread_trace.record_codex_turn_started("turn-1"); + + for model in ["gpt-previous", "gpt-selected"] { + let compaction_trace = + thread_trace.compaction_trace_context("turn-1", "compaction-1", model, "test-provider"); + assert!(compaction_trace.is_enabled()); + compaction_trace + .start_attempt(&serde_json::json!({ "model": model })) + .record_failed("test failure"); + } + + let replayed = replay_bundle(&single_bundle_dir(temp.path())?)?; + let mut attempts = replayed + .compaction_requests + .values() + .map(|attempt| (attempt.model.clone(), attempt.compaction_id.clone())) + .collect::>(); + attempts.sort(); + assert_eq!( + attempts, + vec![ + ("gpt-previous".to_string(), "compaction-1".to_string()), + ("gpt-selected".to_string(), "compaction-1".to_string()), + ] + ); + + Ok(()) +} + +#[test] +fn protocol_wrapper_records_selected_events_as_raw_payloads() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let thread_id = ThreadId::new(); + let thread_trace = + ThreadTraceContext::start_root_in_root_for_test(temp.path(), minimal_metadata(thread_id))?; + + thread_trace.record_protocol_event(&EventMsg::ShutdownComplete); + + let event_log = fs::read_to_string(single_bundle_dir(temp.path())?.join("trace.jsonl"))?; + let protocol_event_seen = event_log.lines().any(|line| { + let event: crate::RawTraceEvent = serde_json::from_str(line).expect("raw trace event"); + matches!( + event.payload, + RawTraceEventPayload::ProtocolEventObserved { + event_type, + .. + } if event_type == "shutdown_complete" + ) + }); + + assert!(protocol_event_seen); + Ok(()) +} + +fn minimal_metadata(thread_id: ThreadId) -> ThreadStartedTraceMetadata { + ThreadStartedTraceMetadata { + thread_id: thread_id.to_string(), + agent_path: "/root".to_string(), + task_name: None, + nickname: None, + agent_role: None, + session_source: SessionSource::Exec, + cwd: PathBuf::from("/workspace"), + rollout_path: None, + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + approval_policy: "never".to_string(), + sandbox_policy: "danger-full-access".to_string(), + } +} + +fn single_bundle_dir(root: &Path) -> anyhow::Result { + let mut entries = fs::read_dir(root)? + .map(|entry| entry.map(|entry| entry.path())) + .collect::, _>>()?; + entries.sort(); + assert_eq!(entries.len(), 1); + Ok(entries.remove(0)) +} diff --git a/codex-rs/rollout-trace/src/tool_dispatch.rs b/codex-rs/rollout-trace/src/tool_dispatch.rs new file mode 100644 index 0000000000000000000000000000000000000000..94dcb3750723266491bea1b396e05af76acb82c0 --- /dev/null +++ b/codex-rs/rollout-trace/src/tool_dispatch.rs @@ -0,0 +1,470 @@ +//! Hot-path helpers for recording canonical tool dispatch boundaries. +//! +//! Core owns tool routing and result conversion. The trace crate owns the raw +//! event schema, payload shape, and no-op behavior, so core only adapts its +//! domain objects into the small request/result structs defined here. + +use std::fmt::Display; +use std::sync::Arc; + +use codex_protocol::models::AdditionalPermissionProfile; +use codex_protocol::models::ResponseInputItem; +use codex_protocol::models::SandboxPermissions; +use codex_protocol::models::SearchToolCallParams; +use serde::Serialize; +use serde_json::Value as JsonValue; +use serde_json::json; +use tracing::warn; + +use crate::model::AgentThreadId; +use crate::model::CodeModeRuntimeToolId; +use crate::model::CodexTurnId; +use crate::model::ExecutionStatus; +use crate::model::ModelVisibleCallId; +use crate::model::ToolCallId; +use crate::model::ToolCallKind; +use crate::model::ToolCallSummary; +use crate::payload::RawPayloadKind; +use crate::payload::RawPayloadRef; +use crate::raw_event::RawToolCallRequester; +use crate::raw_event::RawTraceEventContext; +use crate::raw_event::RawTraceEventPayload; +use crate::writer::TraceWriter; + +/// No-op capable trace handle for one resolved tool dispatch. +#[derive(Clone, Debug)] +pub struct ToolDispatchTraceContext { + state: ToolDispatchTraceContextState, +} + +#[derive(Clone, Debug)] +enum ToolDispatchTraceContextState { + Disabled, + Enabled(EnabledToolDispatchTraceContext), +} + +#[derive(Clone, Debug)] +struct EnabledToolDispatchTraceContext { + writer: Arc, + thread_id: AgentThreadId, + codex_turn_id: CodexTurnId, + tool_call_id: ToolCallId, +} + +/// Core-facing request data for the canonical Codex tool boundary. +pub struct ToolDispatchInvocation { + pub thread_id: AgentThreadId, + pub codex_turn_id: CodexTurnId, + pub tool_call_id: ToolCallId, + pub tool_name: String, + pub tool_namespace: Option, + pub requester: ToolDispatchRequester, + pub payload: ToolDispatchPayload, +} + +/// Runtime source that caused a dispatch-level tool call. +pub enum ToolDispatchRequester { + Model { + model_visible_call_id: ModelVisibleCallId, + }, + CodeCell { + runtime_cell_id: String, + runtime_tool_call_id: CodeModeRuntimeToolId, + }, +} + +/// Tool input observed at the registry boundary. +pub enum ToolDispatchPayload { + Function { + arguments: String, + }, + ToolSearch { + arguments: SearchToolCallParams, + }, + Custom { + input: String, + }, + LocalShell { + command: Vec, + workdir: Option, + timeout_ms: Option, + sandbox_permissions: Option, + prefix_rule: Option>, + additional_permissions: Option, + justification: Option, + }, +} + +/// Result data returned from a dispatch-level tool call. +#[derive(Serialize)] +#[serde(rename_all = "snake_case", tag = "type")] +pub enum ToolDispatchResult { + DirectResponse { response_item: ResponseInputItem }, + CodeModeResponse { value: JsonValue }, +} + +/// Raw invocation payload for the canonical Codex tool boundary. +#[derive(Serialize)] +struct DispatchedToolTraceRequest<'a> { + tool_name: &'a str, + tool_namespace: Option<&'a str>, + payload: &'a JsonValue, +} + +/// Raw response payload for dispatch-level tool trace events. +#[derive(Serialize)] +#[serde(rename_all = "snake_case", tag = "type")] +enum DispatchedToolTraceResponse<'a> { + DirectResponse { + response_item: &'a ResponseInputItem, + }, + CodeModeResponse { + value: &'a JsonValue, + }, + Error { + error: String, + }, +} + +impl ToolDispatchTraceContext { + /// Builds a context that accepts trace calls and records nothing. + pub(crate) fn disabled() -> Self { + Self { + state: ToolDispatchTraceContextState::Disabled, + } + } + + /// Returns whether caller-side result conversion would be recorded. + /// + /// Core uses this to avoid formatting or cloning tool outputs when the + /// dispatch lifecycle is suppressed or tracing is disabled. + pub fn is_enabled(&self) -> bool { + matches!(self.state, ToolDispatchTraceContextState::Enabled(_)) + } + + /// Starts one dispatch-level lifecycle and returns the handle for its result. + pub(crate) fn start(writer: Arc, invocation: ToolDispatchInvocation) -> Self { + if suppresses_tool_dispatch_trace(&invocation) { + return Self::disabled(); + } + + let context = EnabledToolDispatchTraceContext { + writer, + thread_id: invocation.thread_id.clone(), + codex_turn_id: invocation.codex_turn_id.clone(), + tool_call_id: invocation.tool_call_id.clone(), + }; + record_started(&context, invocation); + Self { + state: ToolDispatchTraceContextState::Enabled(context), + } + } + + /// Records the caller-facing successful or failed tool result. + pub fn record_completed(&self, status: ExecutionStatus, result: ToolDispatchResult) { + let ToolDispatchTraceContextState::Enabled(context) = &self.state else { + return; + }; + let response = match &result { + ToolDispatchResult::DirectResponse { response_item } => { + DispatchedToolTraceResponse::DirectResponse { response_item } + } + ToolDispatchResult::CodeModeResponse { value } => { + DispatchedToolTraceResponse::CodeModeResponse { value } + } + }; + append_tool_call_ended(context, status, &response); + } + + /// Records a dispatch failure before the tool produced a normal result payload. + pub fn record_failed(&self, error: impl Display) { + let ToolDispatchTraceContextState::Enabled(context) = &self.state else { + return; + }; + append_tool_call_ended( + context, + ExecutionStatus::Failed, + &DispatchedToolTraceResponse::Error { + error: error.to_string(), + }, + ); + } +} + +fn suppresses_tool_dispatch_trace(invocation: &ToolDispatchInvocation) -> bool { + matches!(invocation.payload, ToolDispatchPayload::Custom { .. }) + && invocation.tool_namespace.is_none() + && invocation.tool_name == codex_code_mode::PUBLIC_TOOL_NAME +} + +fn record_started(context: &EnabledToolDispatchTraceContext, invocation: ToolDispatchInvocation) { + let tool_name = invocation.tool_name; + let tool_namespace = invocation.tool_namespace; + let kind = dispatched_tool_kind(&tool_name, &invocation.payload); + let label = dispatched_tool_label(&tool_name, tool_namespace.as_deref(), &invocation.payload); + let input_preview = Some(invocation.payload.log_payload_preview()); + let payload = invocation.payload.into_json_payload(); + let request = DispatchedToolTraceRequest { + tool_name: tool_name.as_str(), + tool_namespace: tool_namespace.as_deref(), + payload: &payload, + }; + let request_payload = + write_json_payload_best_effort(&context.writer, RawPayloadKind::ToolInvocation, &request); + let (model_visible_call_id, code_mode_runtime_tool_id, requester) = + requester_fields(invocation.requester); + + append_with_context_best_effort( + context, + RawTraceEventPayload::ToolCallStarted { + tool_call_id: context.tool_call_id.clone(), + model_visible_call_id, + code_mode_runtime_tool_id, + requester, + kind, + summary: ToolCallSummary::Generic { + label, + input_preview, + output_preview: None, + }, + invocation_payload: request_payload, + }, + ); +} + +fn requester_fields( + requester: ToolDispatchRequester, +) -> ( + Option, + Option, + RawToolCallRequester, +) { + match requester { + ToolDispatchRequester::Model { + model_visible_call_id, + } => ( + Some(model_visible_call_id), + None, + RawToolCallRequester::Model, + ), + ToolDispatchRequester::CodeCell { + runtime_cell_id, + runtime_tool_call_id, + } => ( + None, + Some(runtime_tool_call_id), + RawToolCallRequester::CodeCell { runtime_cell_id }, + ), + } +} + +fn dispatched_tool_kind(tool_name: &str, _payload: &ToolDispatchPayload) -> ToolCallKind { + match tool_name { + "exec_command" | "local_shell" | "shell" | "shell_command" => ToolCallKind::ExecCommand, + "write_stdin" => ToolCallKind::WriteStdin, + "apply_patch" => ToolCallKind::ApplyPatch, + "web_search" | "web_search_preview" => ToolCallKind::Web, + "image_generation" | "image_query" | "imagegen" => ToolCallKind::ImageGeneration, + "spawn_agent" => ToolCallKind::SpawnAgent, + "send_message" => ToolCallKind::SendMessage, + "followup_task" | "assign_task" => ToolCallKind::AssignAgentTask, + "wait_agent" => ToolCallKind::WaitAgent, + "close_agent" | "interrupt_agent" => ToolCallKind::CloseAgent, + other => ToolCallKind::Other { + name: other.to_string(), + }, + } +} + +fn dispatched_tool_label( + tool_name: &str, + tool_namespace: Option<&str>, + _payload: &ToolDispatchPayload, +) -> String { + match tool_namespace { + Some(namespace) => format!("{namespace}.{tool_name}"), + None => tool_name.to_string(), + } +} + +impl ToolDispatchPayload { + fn log_payload_preview(&self) -> String { + match self { + ToolDispatchPayload::Function { arguments } => truncate_preview(arguments), + ToolDispatchPayload::ToolSearch { arguments } => truncate_preview(&arguments.query), + ToolDispatchPayload::Custom { input } => truncate_preview(input), + ToolDispatchPayload::LocalShell { command, .. } => truncate_preview(&command.join(" ")), + } + } + + fn into_json_payload(self) -> JsonValue { + match self { + ToolDispatchPayload::Function { arguments } => json!({ + "type": "function", + "arguments": arguments, + }), + ToolDispatchPayload::ToolSearch { arguments } => json!({ + "type": "tool_search", + "arguments": arguments, + }), + ToolDispatchPayload::Custom { input } => json!({ + "type": "custom", + "input": input, + }), + ToolDispatchPayload::LocalShell { + command, + workdir, + timeout_ms, + sandbox_permissions, + prefix_rule, + additional_permissions, + justification, + } => json!({ + "type": "local_shell", + "command": command, + "workdir": workdir, + "timeout_ms": timeout_ms, + "sandbox_permissions": sandbox_permissions, + "prefix_rule": prefix_rule, + "additional_permissions": additional_permissions, + "justification": justification, + }), + } + } +} + +fn truncate_preview(value: &str) -> String { + const MAX_PREVIEW_CHARS: usize = 160; + let mut chars = value.chars(); + let mut preview = chars.by_ref().take(MAX_PREVIEW_CHARS).collect::(); + if chars.next().is_some() { + preview.push_str("..."); + } + preview +} + +fn append_tool_call_ended( + context: &EnabledToolDispatchTraceContext, + status: ExecutionStatus, + response: &DispatchedToolTraceResponse<'_>, +) { + let response_payload = + write_json_payload_best_effort(&context.writer, RawPayloadKind::ToolResult, response); + append_with_context_best_effort( + context, + RawTraceEventPayload::ToolCallEnded { + tool_call_id: context.tool_call_id.clone(), + status, + result_payload: response_payload, + }, + ); +} + +fn write_json_payload_best_effort( + writer: &TraceWriter, + kind: RawPayloadKind, + payload: &impl Serialize, +) -> Option { + match writer.write_json_payload(kind, payload) { + Ok(payload_ref) => Some(payload_ref), + Err(err) => { + warn!("failed to write rollout trace payload: {err:#}"); + None + } + } +} + +fn append_with_context_best_effort( + context: &EnabledToolDispatchTraceContext, + payload: RawTraceEventPayload, +) { + let event_context = RawTraceEventContext { + thread_id: Some(context.thread_id.clone()), + codex_turn_id: Some(context.codex_turn_id.clone()), + }; + if let Err(err) = context.writer.append_with_context(event_context, payload) { + warn!("failed to append rollout trace event: {err:#}"); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn suppresses_only_noncanonical_dispatch_boundaries() { + assert!(suppresses_tool_dispatch_trace(&invocation( + codex_code_mode::PUBLIC_TOOL_NAME, + /*tool_namespace*/ None, + ToolDispatchRequester::Model { + model_visible_call_id: "call-exec".to_string(), + }, + ToolDispatchPayload::Custom { + input: "1 + 1".to_string(), + }, + ))); + assert!(!suppresses_tool_dispatch_trace(&invocation( + "custom_tool", + /*tool_namespace*/ None, + ToolDispatchRequester::Model { + model_visible_call_id: "call-custom".to_string(), + }, + ToolDispatchPayload::Custom { + input: "payload".to_string(), + }, + ))); + assert!(!suppresses_tool_dispatch_trace(&invocation( + codex_code_mode::PUBLIC_TOOL_NAME, + Some("mcp__server".to_string()), + ToolDispatchRequester::Model { + model_visible_call_id: "call-namespaced".to_string(), + }, + ToolDispatchPayload::Custom { + input: "payload".to_string(), + }, + ))); + } + + #[test] + fn classifies_interrupt_agent_as_close_agent() { + assert_eq!( + dispatched_tool_kind( + "interrupt_agent", + &ToolDispatchPayload::Function { + arguments: r#"{"target":"/root/child"}"#.to_string(), + }, + ), + ToolCallKind::CloseAgent + ); + } + + #[test] + fn classifies_imagegen_as_image_generation() { + assert_eq!( + dispatched_tool_kind( + "imagegen", + &ToolDispatchPayload::Function { + arguments: String::new(), + }, + ), + ToolCallKind::ImageGeneration + ); + } + + fn invocation( + tool_name: &str, + tool_namespace: Option, + requester: ToolDispatchRequester, + payload: ToolDispatchPayload, + ) -> ToolDispatchInvocation { + ToolDispatchInvocation { + thread_id: "thread-1".to_string(), + codex_turn_id: "turn-1".to_string(), + tool_call_id: "tool-call-1".to_string(), + tool_name: tool_name.to_string(), + tool_namespace, + requester, + payload, + } + } +} diff --git a/codex-rs/rollout-trace/src/writer.rs b/codex-rs/rollout-trace/src/writer.rs new file mode 100644 index 0000000000000000000000000000000000000000..925f7bab7a09c83fad9fb8bbd6fd6da68f572992 --- /dev/null +++ b/codex-rs/rollout-trace/src/writer.rs @@ -0,0 +1,265 @@ +//! Hot-path trace bundle writer. + +use std::fs::File; +use std::fs::OpenOptions; +use std::io::BufWriter; +use std::io::Write; +use std::path::Path; +use std::path::PathBuf; +use std::sync::Mutex; +use std::sync::MutexGuard; +use std::sync::PoisonError; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +use anyhow::Context; +use anyhow::Result; +use serde::Serialize; + +use crate::bundle::MANIFEST_FILE_NAME; +use crate::bundle::PAYLOADS_DIR_NAME; +use crate::bundle::RAW_EVENT_LOG_FILE_NAME; +use crate::bundle::TraceBundleManifest; +use crate::model::AgentThreadId; +use crate::payload::RawPayloadKind; +use crate::payload::RawPayloadRef; +use crate::raw_event::RAW_TRACE_EVENT_SCHEMA_VERSION; +use crate::raw_event::RawTraceEvent; +use crate::raw_event::RawTraceEventContext; +use crate::raw_event::RawTraceEventPayload; + +/// Local trace bundle writer. +/// +/// The writer appends raw events and writes payload files. It does not keep a +/// reduced `RolloutTrace` in memory; replay is owned by the reducer. +#[derive(Debug)] +pub struct TraceWriter { + inner: Mutex, +} + +#[derive(Debug)] +struct TraceWriterInner { + manifest: TraceBundleManifest, + payloads_dir: PathBuf, + event_log: BufWriter, + next_seq: u64, + next_payload_ordinal: u64, +} + +impl TraceWriter { + /// Creates a trace bundle directory and writes its manifest. + pub fn create( + bundle_dir: impl AsRef, + trace_id: String, + rollout_id: String, + root_thread_id: AgentThreadId, + ) -> Result { + let bundle_dir = bundle_dir.as_ref().to_path_buf(); + let payloads_dir = bundle_dir.join(PAYLOADS_DIR_NAME); + std::fs::create_dir_all(&payloads_dir) + .with_context(|| format!("create trace payload dir {}", payloads_dir.display()))?; + + let started_at_unix_ms = unix_time_ms(); + let manifest = + TraceBundleManifest::new(trace_id, rollout_id, root_thread_id, started_at_unix_ms); + write_json_file(&bundle_dir.join(MANIFEST_FILE_NAME), &manifest)?; + + let event_log_path = bundle_dir.join(RAW_EVENT_LOG_FILE_NAME); + let event_log = OpenOptions::new() + .create(true) + .append(true) + .open(&event_log_path) + .with_context(|| format!("open trace event log {}", event_log_path.display()))?; + + Ok(Self { + inner: Mutex::new(TraceWriterInner { + manifest, + payloads_dir, + event_log: BufWriter::new(event_log), + next_seq: 1, + next_payload_ordinal: 1, + }), + }) + } + + /// Writes a JSON payload file and returns its reduced-state reference. + pub fn write_json_payload( + &self, + kind: RawPayloadKind, + value: &impl Serialize, + ) -> Result { + let mut inner = self.lock_inner(); + let ordinal = inner.next_payload_ordinal; + inner.next_payload_ordinal += 1; + let raw_payload_id = format!("raw_payload:{ordinal}"); + let relative_path = format!("{PAYLOADS_DIR_NAME}/{ordinal}.json"); + let absolute_path = inner.payloads_dir.join(format!("{ordinal}.json")); + // Payload files are created before the event that references them. A + // replay interrupted after an event is appended should never point at a + // payload file that the writer planned but had not written yet. + write_json_file(&absolute_path, value)?; + Ok(RawPayloadRef { + raw_payload_id, + kind, + path: relative_path, + }) + } + + /// Appends one raw event with no extra envelope context. + pub fn append(&self, payload: RawTraceEventPayload) -> Result { + self.append_with_context(RawTraceEventContext::default(), payload) + } + + /// Appends one raw event with explicit thread/turn context. + pub fn append_with_context( + &self, + context: RawTraceEventContext, + payload: RawTraceEventPayload, + ) -> Result { + let mut inner = self.lock_inner(); + let event = RawTraceEvent { + schema_version: RAW_TRACE_EVENT_SCHEMA_VERSION, + seq: inner.next_seq, + wall_time_unix_ms: unix_time_ms(), + rollout_id: inner.manifest.rollout_id.clone(), + thread_id: context.thread_id, + codex_turn_id: context.codex_turn_id, + payload, + }; + inner.next_seq += 1; + serde_json::to_writer(&mut inner.event_log, &event)?; + inner.event_log.write_all(b"\n")?; + inner.event_log.flush()?; + Ok(event) + } + + fn lock_inner(&self) -> MutexGuard<'_, TraceWriterInner> { + // Preserve the event log after a panic in tracing code. Dropping the + // writer would lose subsequent diagnostic events in exactly the session + // we are trying to debug. + self.inner.lock().unwrap_or_else(PoisonError::into_inner) + } +} + +fn write_json_file(path: &Path, value: &impl Serialize) -> Result<()> { + let file = File::create(path).with_context(|| format!("create {}", path.display()))?; + serde_json::to_writer_pretty(file, value) + .with_context(|| format!("write JSON {}", path.display())) +} + +pub(crate) fn unix_time_ms() -> i64 { + let duration = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default(); + i64::try_from(duration.as_millis()).unwrap_or(i64::MAX) +} + +#[cfg(test)] +mod tests { + use pretty_assertions::assert_eq; + use serde_json::json; + use tempfile::TempDir; + + use crate::model::ExecutionStatus; + use crate::model::RolloutStatus; + use crate::payload::RawPayloadKind; + use crate::raw_event::RawTraceEventPayload; + use crate::replay_bundle; + use crate::writer::TraceWriter; + + #[test] + fn writer_records_payload_refs_and_replays_rollout_status() -> anyhow::Result<()> { + let temp = TempDir::new()?; + let writer = TraceWriter::create( + temp.path(), + "trace-1".to_string(), + "rollout-1".to_string(), + "thread-root".to_string(), + )?; + + writer.append(RawTraceEventPayload::RolloutStarted { + trace_id: "trace-1".to_string(), + root_thread_id: "thread-root".to_string(), + })?; + let metadata_payload = writer.write_json_payload( + RawPayloadKind::ProtocolEvent, + &json!({ + "source": "test", + "model": "gpt-test", + }), + )?; + writer.append(RawTraceEventPayload::ThreadStarted { + thread_id: "thread-root".to_string(), + agent_path: "/root".to_string(), + metadata_payload: Some(metadata_payload.clone()), + })?; + writer.append(RawTraceEventPayload::CodexTurnStarted { + codex_turn_id: "turn-1".to_string(), + thread_id: "thread-root".to_string(), + })?; + let inference_request = writer.write_json_payload( + RawPayloadKind::InferenceRequest, + &json!({ + "model": "gpt-test", + "input": [{ + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "hello"}] + }], + }), + )?; + writer.append(RawTraceEventPayload::InferenceStarted { + inference_call_id: "inference-1".to_string(), + thread_id: "thread-root".to_string(), + codex_turn_id: "turn-1".to_string(), + model: "gpt-test".to_string(), + provider_name: "test-provider".to_string(), + request_payload: inference_request.clone(), + })?; + let inference_response = writer.write_json_payload( + RawPayloadKind::InferenceResponse, + &json!({ + "response_id": "resp-1", + "output_items": [], + }), + )?; + writer.append(RawTraceEventPayload::InferenceCompleted { + inference_call_id: "inference-1".to_string(), + response_id: Some("resp-1".to_string()), + upstream_request_id: Some("req-1".to_string()), + response_payload: inference_response.clone(), + })?; + writer.append(RawTraceEventPayload::CodexTurnEnded { + codex_turn_id: "turn-1".to_string(), + status: ExecutionStatus::Completed, + })?; + writer.append(RawTraceEventPayload::RolloutEnded { + status: RolloutStatus::Completed, + })?; + + let rollout = replay_bundle(temp.path())?; + + assert_eq!(rollout.status, RolloutStatus::Completed); + assert_eq!(rollout.root_thread_id, "thread-root"); + assert_eq!(rollout.threads["thread-root"].agent_path, "/root"); + assert_eq!(rollout.codex_turns["turn-1"].thread_id, "thread-root"); + assert_eq!( + rollout.codex_turns["turn-1"].execution.status, + ExecutionStatus::Completed, + ); + assert_eq!( + rollout.inference_calls["inference-1"].raw_request_payload_id, + inference_request.raw_payload_id, + ); + assert_eq!( + rollout.inference_calls["inference-1"].raw_response_payload_id, + Some(inference_response.raw_payload_id), + ); + assert_eq!( + rollout.raw_payloads[&metadata_payload.raw_payload_id].path, + "payloads/1.json" + ); + + Ok(()) + } +} diff --git a/codex-rs/thread-manager-sample/src/main.rs b/codex-rs/thread-manager-sample/src/main.rs new file mode 100644 index 0000000000000000000000000000000000000000..ca0a84bf5e38c3e6e1dd36e38695c6e7ff039024 --- /dev/null +++ b/codex-rs/thread-manager-sample/src/main.rs @@ -0,0 +1,429 @@ +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::io::IsTerminal; +use std::io::Read; +use std::io::Write; +use std::sync::Arc; + +use anyhow::Context; +use anyhow::bail; +use clap::Parser; +use codex_core_api::AbsolutePathBuf; +use codex_core_api::AltScreenMode; +use codex_core_api::ApprovalsReviewer; +use codex_core_api::Arg0DispatchPaths; +use codex_core_api::AskForApproval; +use codex_core_api::AuthCredentialsStoreMode; +use codex_core_api::AuthManager; +use codex_core_api::AutoCompactTokenLimitScope; +use codex_core_api::CodexAppsToolsCache; +use codex_core_api::CodexHomeUserInstructionsProvider; +use codex_core_api::CodexThread; +use codex_core_api::Config; +use codex_core_api::ConfigLayerStack; +use codex_core_api::Constrained; +use codex_core_api::EnvironmentManager; +use codex_core_api::EventMsg; +use codex_core_api::ExecServerRuntimePaths; +use codex_core_api::ExtensionRegistryBuilder; +use codex_core_api::Features; +use codex_core_api::GhostSnapshotConfig; +use codex_core_api::History; +use codex_core_api::MemoriesConfig; +use codex_core_api::ModelAvailabilityNuxConfig; +use codex_core_api::MultiAgentV2Config; +use codex_core_api::NewThread; +use codex_core_api::Notice; +use codex_core_api::OAuthCredentialsStoreMode; +use codex_core_api::OPENAI_PROVIDER_ID; +use codex_core_api::OtelConfig; +use codex_core_api::PermissionProfile; +use codex_core_api::Permissions; +use codex_core_api::ProjectConfig; +use codex_core_api::RealtimeAudioConfig; +use codex_core_api::RealtimeConfig; +use codex_core_api::SessionPickerViewMode; +use codex_core_api::SessionSource; +use codex_core_api::SqliteConfig; +use codex_core_api::StartIfIdleSubmission; +use codex_core_api::StartThreadOptions; +use codex_core_api::TerminalResizeReflowConfig; +use codex_core_api::ThreadManager; +use codex_core_api::ThreadStoreConfig; +use codex_core_api::ToolSuggestConfig; +use codex_core_api::TuiKeymap; +use codex_core_api::TuiNotificationSettings; +use codex_core_api::TuiPetAnchor; +use codex_core_api::TurnInputRequest; +use codex_core_api::UriBasedFileOpener; +use codex_core_api::UserInput; +use codex_core_api::WebSearchMode; +use codex_core_api::arg0_dispatch_or_else; +use codex_core_api::build_models_manager; +use codex_core_api::built_in_model_providers; +use codex_core_api::find_codex_home; +use codex_core_api::init_state_db; +use codex_core_api::install_image_generation_extension; +use codex_core_api::item_event_to_server_notification; +use codex_core_api::local_agent_graph_store_from_state_db; +use codex_core_api::passthrough_image_store; +use codex_core_api::resolve_installation_id; +use codex_core_api::set_default_originator; +use codex_core_api::thread_store_from_config; + +#[derive(Debug, Parser)] +#[command( + name = "codex-thread-manager-sample", + about = "Run one Codex turn through ThreadManager and print mapped notifications as newline-delimited JSON." +)] +struct Args { + /// Override the model for this run. + #[arg(long, value_name = "MODEL")] + model: Option, + + /// Prompt text. If omitted, the prompt is read from piped stdin. + #[arg(value_name = "PROMPT", num_args = 0.., trailing_var_arg = true)] + prompt: Vec, +} + +fn main() -> anyhow::Result<()> { + arg0_dispatch_or_else(run_main) +} + +async fn run_main(arg0_paths: Arg0DispatchPaths) -> anyhow::Result<()> { + if let Err(err) = set_default_originator("codex_thread_manager_sample".to_string()) { + tracing::warn!("failed to set originator: {err:?}"); + } + + let args = Args::parse(); + let prompt = if args.prompt.is_empty() { + if std::io::stdin().is_terminal() { + bail!("no prompt provided; pass a prompt argument or pipe one into stdin"); + } + + let mut prompt = String::new(); + std::io::stdin() + .read_to_string(&mut prompt) + .context("read prompt from stdin")?; + let prompt = prompt.replace("\r\n", "\n").replace('\r', "\n"); + if prompt.trim().is_empty() { + bail!("no prompt provided via stdin"); + } + prompt + } else { + args.prompt.join(" ") + }; + + let config = new_config(args.model, arg0_paths)?; + let state_db = init_state_db(&config).await; + + let auth_manager = + AuthManager::shared_from_config(&config, /*enable_codex_api_key_env*/ false).await?; + let local_runtime_paths = ExecServerRuntimePaths::from_optional_paths( + config.codex_self_exe.clone(), + config.codex_linux_sandbox_exe.clone(), + )?; + let thread_store = thread_store_from_config(&config, state_db.clone()); + let environment_manager = Arc::new( + EnvironmentManager::from_codex_home( + config.codex_home.clone(), + Some(local_runtime_paths), + config.http_client_factory(), + ) + .await?, + ); + let installation_id = resolve_installation_id(&config.codex_home).await?; + let user_instructions_provider = Arc::new(CodexHomeUserInstructionsProvider::new( + config.codex_home.clone(), + )); + let mut extensions = ExtensionRegistryBuilder::::new(); + install_image_generation_extension(&mut extensions, auth_manager.clone(), |config: &Config| { + Some(config.codex_home.clone()) + }); + let thread_manager = ThreadManager::new( + &config, + Arc::clone(&auth_manager), + build_models_manager(&config, auth_manager), + CodexAppsToolsCache::default(), + SessionSource::Exec, + environment_manager, + Arc::new(extensions.build()), + user_instructions_provider, + /*analytics_events_client*/ None, + passthrough_image_store(), + Arc::clone(&thread_store), + local_agent_graph_store_from_state_db(state_db.as_ref()), + installation_id, + /*attestation_provider*/ None, + /*external_time_provider*/ None, + ); + + let NewThread { + thread_id, thread, .. + } = thread_manager + .start_thread(StartThreadOptions::new(config)) + .await + .context("start Codex thread")?; + + let thread_id_string = thread_id.to_string(); + let turn_output = run_turn(&thread, &thread_id_string, prompt).await; + let shutdown_result = thread.shutdown_and_wait().await; + let _ = thread_manager.remove_thread(&thread_id).await; + + turn_output?; + shutdown_result.context("shut down Codex thread")?; + + Ok(()) +} + +fn new_config(model: Option, arg0_paths: Arg0DispatchPaths) -> anyhow::Result { + let codex_home = find_codex_home().context("find Codex home")?; + let cwd = AbsolutePathBuf::current_dir().context("resolve current directory")?; + let model_provider_id = OPENAI_PROVIDER_ID.to_string(); + let model_providers = built_in_model_providers(/*openai_base_url*/ None); + let model_provider = model_providers + .get(&model_provider_id) + .context("OpenAI model provider should be available")? + .clone(); + + let mut config = Config { + config_layer_stack: ConfigLayerStack::default(), + startup_warnings: Vec::new(), + bypass_hook_trust: false, + model, + service_tier: None, + review_model: None, + model_context_window: None, + model_auto_compact_token_limit: None, + model_auto_compact_token_limit_scope: AutoCompactTokenLimitScope::Total, + model_provider_id, + model_provider, + personality: None, + permissions: Permissions::from_approval_and_profile( + Constrained::allow_any(AskForApproval::Never), + Constrained::allow_any(PermissionProfile::read_only()), + )?, + explicit_permission_profile_mode: false, + custom_permission_profiles: Vec::new(), + approvals_reviewer: ApprovalsReviewer::User, + enforce_residency: Constrained::allow_any(/*initial_value*/ None), + hide_agent_reasoning: false, + show_raw_agent_reasoning: false, + base_instructions: None, + base_instructions_provenance: None, + developer_instructions: None, + guardian_policy_config: None, + guardian_policy_template: None, + include_permissions_instructions: false, + include_apps_instructions: false, + include_collaboration_mode_instructions: false, + include_skill_instructions: false, + skill_max_context_tokens: None, + orchestrator_skills_enabled: false, + orchestrator_mcp_enabled: false, + include_environment_context: false, + compact_prompt: None, + notify: None, + tui_notifications: TuiNotificationSettings::default(), + animations: true, + tui_whimsy: true, + show_tooltips: true, + tui_show_server_version_notice: true, + tui_auto_recap: true, + model_availability_nux: ModelAvailabilityNuxConfig::default(), + tui_alternate_screen: AltScreenMode::Auto, + tui_status_line: None, + tui_status_line_use_colors: true, + tui_terminal_title: None, + tui_theme: None, + tui_raw_output_mode: false, + tui_pet: None, + tui_pet_anchor: TuiPetAnchor::Composer, + terminal_resize_reflow: TerminalResizeReflowConfig::default(), + tui_keymap: TuiKeymap::default(), + tui_session_picker_view: SessionPickerViewMode::Dense, + tui_resume_cwd: None, + tui_vim_mode_default: false, + tui_question_esc_back: true, + cwd: cwd.clone(), + workspace_roots: vec![cwd], + workspace_roots_explicit: false, + cli_auth_credentials_store_mode: AuthCredentialsStoreMode::File, + mcp_servers: Constrained::allow_any(HashMap::new()), + mcp_enterprise_managed_auth: None, + non_prefixed_mcp_tool_servers: None, + mcp_oauth_credentials_store_mode: OAuthCredentialsStoreMode::File, + mcp_oauth_callback_port: None, + mcp_oauth_callback_url: None, + mcp_optional_startup_grace: std::time::Duration::from_secs(1), + model_providers, + project_doc_max_bytes: 32 * 1024, + project_doc_fallback_filenames: Vec::new(), + tool_output_token_limit: None, + agents_enabled: true, + agent_max_threads: Some(6), + agent_default_subagent_model: None, + agent_default_subagent_reasoning_effort: None, + agent_interrupt_message_enabled: false, + agent_max_depth: 1, + agent_roles: BTreeMap::new(), + memories: MemoriesConfig::default(), + sqlite: SqliteConfig::from_sqlite_home(codex_home.clone()), + log_dir: codex_home.join("log").to_path_buf(), + codex_home, + history: History::default(), + ephemeral: true, + extra_config: None, + file_opener: UriBasedFileOpener::VsCode, + codex_self_exe: arg0_paths.codex_self_exe, + codex_linux_sandbox_exe: arg0_paths.codex_linux_sandbox_exe, + main_execve_wrapper_exe: arg0_paths.main_execve_wrapper_exe, + zsh_path: None, + model_reasoning_effort: None, + plan_mode_reasoning_effort: None, + model_reasoning_summary: None, + model_catalog: None, + model_verbosity: None, + chatgpt_base_url: "https://chatgpt.com/backend-api/".to_string(), + respect_system_proxy: false, + apps_mcp_product_sku: None, + responses_api_metadata: BTreeMap::new(), + realtime_audio: RealtimeAudioConfig::default(), + experimental_realtime_ws_base_url: None, + experimental_realtime_webrtc_call_base_url: None, + experimental_realtime_ws_model: None, + realtime: RealtimeConfig::default(), + experimental_realtime_ws_backend_prompt: None, + experimental_realtime_ws_startup_context: None, + experimental_realtime_start_instructions: None, + experimental_thread_store: ThreadStoreConfig::Local, + forced_chatgpt_workspace_id: None, + forced_login_method: None, + web_search_mode: Constrained::allow_any(WebSearchMode::Disabled), + web_search_config: None, + experimental_request_user_input_enabled: true, + update_plan_enabled: true, + tool_registry: Default::default(), + code_mode: Default::default(), + background_terminal_max_timeout: 300_000, + thread_unload_delay: std::time::Duration::from_secs(60), + ghost_snapshot: GhostSnapshotConfig::default(), + multi_agent_v2: MultiAgentV2Config::default(), + max_goal_token_budget: None, + token_budget: None, + token_budget_startup_config: None, + rollout_budget: None, + current_time_reminder: None, + sleep_tool_mode: Default::default(), + features: Default::default(), + suppress_unstable_features_warning: false, + active_project: ProjectConfig { trust_level: None }, + notices: Notice::default(), + check_for_update_on_startup: false, + disable_paste_burst: false, + analytics_enabled: Some(false), + feedback_enabled: false, + tool_suggest: ToolSuggestConfig::default(), + otel: OtelConfig::default(), + }; + config + .features + .set(Features::with_defaults()) + .context("configure default features")?; + Ok(config) +} + +async fn run_turn(thread: &CodexThread, thread_id: &str, prompt: String) -> anyhow::Result<()> { + let submission = thread + .start_turn_if_idle(TurnInputRequest::user_input(vec![UserInput::Text { + text: prompt, + text_elements: Vec::new(), + }])) + .await + .context("submit user input")?; + if let StartIfIdleSubmission::NotSubmitted { reason } = submission { + bail!("turn input was not submitted: {reason:?}"); + } + + let mut current_turn_id: Option = None; + let mut stdout = std::io::stdout().lock(); + loop { + let event = thread.next_event().await.context("read Codex event")?; + let notification = match &event.msg { + EventMsg::TurnStarted(event) => { + current_turn_id = Some(event.turn_id.clone()); + None + } + EventMsg::DynamicToolCallResponse(_) + | EventMsg::McpToolCallBegin(_) + | EventMsg::McpToolCallEnd(_) + | EventMsg::CollabAgentSpawnBegin(_) + | EventMsg::CollabAgentSpawnEnd(_) + | EventMsg::CollabAgentInteractionBegin(_) + | EventMsg::CollabAgentInteractionEnd(_) + | EventMsg::CollabWaitingBegin(_) + | EventMsg::CollabWaitingEnd(_) + | EventMsg::CollabCloseBegin(_) + | EventMsg::CollabCloseEnd(_) + | EventMsg::CollabResumeBegin(_) + | EventMsg::CollabResumeEnd(_) + | EventMsg::SubAgentActivity(_) + | EventMsg::AgentMessageContentDelta(_) + | EventMsg::PlanDelta(_) + | EventMsg::ReasoningContentDelta(_) + | EventMsg::ReasoningRawContentDelta(_) + | EventMsg::AgentReasoningSectionBreak(_) + | EventMsg::ItemStarted(_) + | EventMsg::ItemCompleted(_) + | EventMsg::PatchApplyBegin(_) + | EventMsg::PatchApplyUpdated(_) + | EventMsg::TerminalInteraction(_) + | EventMsg::ExecCommandBegin(_) + | EventMsg::ExecCommandOutputDelta(_) + | EventMsg::ExecCommandEnd(_) => Some(item_event_to_server_notification( + event.msg.clone(), + thread_id, + current_turn_id + .as_deref() + .context("mapped notification arrived before turn started")?, + )), + _ => None, + }; + if let Some(notification) = notification { + serde_json::to_writer(&mut stdout, ¬ification) + .context("serialize mapped notification")?; + stdout + .write_all(b"\n") + .context("write notification newline")?; + stdout.flush().context("flush notification output")?; + } + + match event.msg { + EventMsg::TurnComplete(_) => { + return Ok(()); + } + EventMsg::Error(event) => { + bail!(event.message); + } + EventMsg::TurnAborted(_) => { + bail!("turn aborted"); + } + EventMsg::ExecApprovalRequest(_) => { + bail!("turn requested exec approval"); + } + EventMsg::ApplyPatchApprovalRequest(_) => { + bail!("turn requested patch approval"); + } + EventMsg::RequestPermissions(_) => { + bail!("turn requested permissions"); + } + EventMsg::RequestUserInput(_) => { + bail!("turn requested user input"); + } + EventMsg::DynamicToolCallRequest(_) => { + bail!("turn requested a dynamic tool call"); + } + _ => {} + } + } +} diff --git a/codex-rs/user-verification/src/credential.rs b/codex-rs/user-verification/src/credential.rs new file mode 100644 index 0000000000000000000000000000000000000000..7b2ea925a30e1c7a8a6a7d00b2a9866fe0959c41 --- /dev/null +++ b/codex-rs/user-verification/src/credential.rs @@ -0,0 +1,86 @@ +//! Public credential encoding and results; private keys never cross this boundary. + +use crate::UserVerificationError; +use crate::UserVerificationFailureReason; +use crate::UserVerificationUnavailableReason; +use base64::Engine as _; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use p256::pkcs8::EncodePublicKey as _; +use sha2::Digest as _; +use sha2::Sha256; + +/// Validated display text and exact challenge bytes supplied by the trusted calling UI. +#[derive(Clone)] +pub struct UserVerificationRequest { + pub challenge: Vec, + pub title: String, + pub description: String, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct UserVerificationKeyInfo { + pub credential_id: String, + pub algorithm: String, + /// Unpadded base64url of the P-256 SubjectPublicKeyInfo DER encoding. + pub public_key: String, +} + +impl UserVerificationKeyInfo { + pub fn from_sec1_public_key(bytes: &[u8]) -> Result { + let public_key = p256::PublicKey::from_sec1_bytes(bytes) + .map_err(|_| invalid_public_key())? + .to_public_key_der() + .map_err(|_| invalid_public_key())?; + let bytes = public_key.as_bytes(); + Ok(Self { + credential_id: URL_SAFE_NO_PAD.encode(Sha256::digest(bytes)), + algorithm: "ecdsaP256Sha256X962".to_string(), + public_key: URL_SAFE_NO_PAD.encode(bytes), + }) + } +} + +fn invalid_public_key() -> UserVerificationError { + UserVerificationError::Failed { + reason: UserVerificationFailureReason::ProviderError, + message: "could not encode the user-verification public key".to_string(), + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct UserVerificationKeyCreation { + pub created: bool, + pub credential: UserVerificationKeyInfo, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct UserVerificationKeyDeletion { + pub deleted_credential_id: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct UserVerificationStatus { + pub credential: Option, + pub unavailable_reason: Option, + pub unavailable_message: Option, +} + +#[derive(Clone, PartialEq, Eq)] +pub struct UserVerificationProof { + pub credential_id: String, + /// Unpadded base64url of an ASN.1 DER ECDSA signature over the challenge, hashed once. + pub signature: String, +} + +impl std::fmt::Debug for UserVerificationProof { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("UserVerificationProof") + .field("credential_id", &"[REDACTED]") + .field("signature", &"[REDACTED]") + .finish() + } +} + +#[cfg(test)] +#[path = "credential_tests.rs"] +mod tests; diff --git a/codex-rs/user-verification/src/credential_tests.rs b/codex-rs/user-verification/src/credential_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..4aa3326366e20ccf1fefeb9b1eedbf09de26546b --- /dev/null +++ b/codex-rs/user-verification/src/credential_tests.rs @@ -0,0 +1,43 @@ +//! Checks that the exported credential interoperates with DER-based service verification. + +use super::*; +use p256::ecdsa::Signature; +use p256::ecdsa::SigningKey; +use p256::ecdsa::VerifyingKey; +use p256::ecdsa::signature::Signer as _; +use p256::ecdsa::signature::Verifier as _; +use p256::pkcs8::DecodePublicKey as _; +use pretty_assertions::assert_eq; + +#[test] +fn exported_spki_verifies_the_exact_challenge_signature() { + let key = SigningKey::from_bytes((&[7_u8; 32]).into()).expect("valid signing key"); + let info = UserVerificationKeyInfo::from_sec1_public_key( + key.verifying_key() + .to_encoded_point(/*compress*/ false) + .as_bytes(), + ) + .expect("encode public key"); + let der = URL_SAFE_NO_PAD.decode(&info.public_key).expect("base64url"); + let verifier = VerifyingKey::from_public_key_der(&der).expect("valid SPKI DER"); + let challenge = b"the exact server-issued challenge"; + let signature: Signature = key.sign(challenge); + let signature = Signature::from_der(signature.to_der().as_bytes()).expect("DER signature"); + verifier + .verify(challenge, &signature) + .expect("valid signature"); + assert!(verifier.verify(b"different challenge", &signature).is_err()); + assert_eq!( + info, + UserVerificationKeyInfo { + credential_id: URL_SAFE_NO_PAD.encode(Sha256::digest(&der)), + algorithm: "ecdsaP256Sha256X962".to_string(), + public_key: URL_SAFE_NO_PAD.encode(der), + } + ); +} + +#[test] +fn invalid_curve_point_is_rejected() { + assert!(UserVerificationKeyInfo::from_sec1_public_key(&[4_u8; 65]).is_err()); +} diff --git a/codex-rs/user-verification/src/error.rs b/codex-rs/user-verification/src/error.rs new file mode 100644 index 0000000000000000000000000000000000000000..88f611e1525c84d9157eb07aa142c37c27fb63b2 --- /dev/null +++ b/codex-rs/user-verification/src/error.rs @@ -0,0 +1,42 @@ +//! Stable error categories, with provider diagnostics kept outside public error messages. + +use thiserror::Error; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum UserVerificationUnavailableReason { + CredentialMissing, + BiometricsUnavailable, + ProviderUnavailable, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum UserVerificationCancellationReason { + UserCancelled, + Interrupted, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum UserVerificationFailureReason { + AuthenticationFailed, + Timeout, + ProviderError, +} + +#[derive(Clone, Debug, Error, PartialEq, Eq)] +pub enum UserVerificationError { + #[error("user verification is unavailable: {message}")] + Unavailable { + reason: UserVerificationUnavailableReason, + message: String, + }, + #[error("user verification was cancelled: {message}")] + Cancelled { + reason: UserVerificationCancellationReason, + message: String, + }, + #[error("user verification failed: {message}")] + Failed { + reason: UserVerificationFailureReason, + message: String, + }, +} diff --git a/codex-rs/user-verification/src/guard.rs b/codex-rs/user-verification/src/guard.rs new file mode 100644 index 0000000000000000000000000000000000000000..e020a9dd2425f768fcc576aed9f1a343ebd557c0 --- /dev/null +++ b/codex-rs/user-verification/src/guard.rs @@ -0,0 +1,51 @@ +//! Cancellation and caller-owned identity checks for queued native operations. + +use crate::UserVerificationCancellationReason; +use crate::UserVerificationError; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +/// Invalidates queued work and suppresses late results. Native providers observe cancellation +/// during authentication and request dismissal of their active OS prompt. +#[derive(Clone, Default)] +pub struct UserVerificationRequestGuard { + cancelled: Arc, + activity_check: Option bool + Send + Sync>>, +} + +impl UserVerificationRequestGuard { + /// The callback checks a captured identity; it must be nonblocking and must not prompt. + pub fn with_activity_check(activity_check: impl Fn() -> bool + Send + Sync + 'static) -> Self { + Self { + cancelled: Arc::default(), + activity_check: Some(Arc::new(activity_check)), + } + } + + pub fn cancel(&self) { + self.cancelled.store(/*val*/ true, Ordering::Release); + } + + pub fn is_active(&self) -> bool { + if self.activity_check.as_ref().is_some_and(|check| !check()) { + self.cancel(); + } + !self.cancelled.load(Ordering::Acquire) + } + + pub fn check(&self) -> Result<(), UserVerificationError> { + if self.is_active() { + Ok(()) + } else { + Err(UserVerificationError::Cancelled { + reason: UserVerificationCancellationReason::Interrupted, + message: "the verification operation is no longer active".to_string(), + }) + } + } +} + +#[cfg(test)] +#[path = "guard_tests.rs"] +mod tests; diff --git a/codex-rs/user-verification/src/guard_tests.rs b/codex-rs/user-verification/src/guard_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..920d62edaedac37dc094253a9dda30fdb4a0d1dd --- /dev/null +++ b/codex-rs/user-verification/src/guard_tests.rs @@ -0,0 +1,32 @@ +//! Checks cancellation shared across queued operations and captured-identity callbacks. + +use super::*; +use pretty_assertions::assert_eq; + +#[test] +fn identity_change_permanently_cancels_all_guard_clones() { + let identity_matches = Arc::new(AtomicBool::new(/*v*/ true)); + let activity = Arc::clone(&identity_matches); + let guard = + UserVerificationRequestGuard::with_activity_check(move || activity.load(Ordering::Acquire)); + let queued = guard.clone(); + assert!(queued.check().is_ok()); + identity_matches.store(/*val*/ false, Ordering::Release); + assert_eq!( + queued.check(), + Err(UserVerificationError::Cancelled { + reason: UserVerificationCancellationReason::Interrupted, + message: "the verification operation is no longer active".to_string(), + }) + ); + identity_matches.store(/*val*/ true, Ordering::Release); + assert!(!guard.is_active()); +} + +#[test] +fn cancelling_one_clone_invalidates_queued_work() { + let guard = UserVerificationRequestGuard::default(); + let queued = guard.clone(); + guard.cancel(); + assert!(!queued.is_active()); +} diff --git a/codex-rs/user-verification/src/key_namespace.rs b/codex-rs/user-verification/src/key_namespace.rs new file mode 100644 index 0000000000000000000000000000000000000000..6677147906c51e4a64d003333ec86b3d056f9bb4 --- /dev/null +++ b/codex-rs/user-verification/src/key_namespace.rs @@ -0,0 +1,25 @@ +//! Keychain labels isolate account-user identities within the fixed plugin-service scope. + +use base64::Engine as _; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use sha2::Digest as _; +use sha2::Sha256; + +/// An opaque account-user namespace. App-server authenticates and selects the identity. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct UserVerificationKeyNamespace { + pub(crate) label: String, +} + +impl UserVerificationKeyNamespace { + pub fn new(account_user_id: &str) -> Self { + let identity = URL_SAFE_NO_PAD.encode(Sha256::digest(account_user_id.as_bytes())); + Self { + label: format!("com.openai.codex.user-verification.plugin-service.v1.{identity}"), + } + } +} + +#[cfg(test)] +#[path = "key_namespace_tests.rs"] +mod tests; diff --git a/codex-rs/user-verification/src/key_namespace_tests.rs b/codex-rs/user-verification/src/key_namespace_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..d28d0c5020541e1cda8b33894c83aeb79bf31e4e --- /dev/null +++ b/codex-rs/user-verification/src/key_namespace_tests.rs @@ -0,0 +1,12 @@ +//! Account-user identities share one service scope without exposing raw identifiers. + +use super::*; +use pretty_assertions::assert_eq; + +#[test] +fn namespaces_are_stable_and_separate_accounts() { + let first = UserVerificationKeyNamespace::new("account-user-one"); + assert_eq!(first, UserVerificationKeyNamespace::new("account-user-one")); + assert_ne!(first, UserVerificationKeyNamespace::new("account-user-two")); + assert!(!first.label.contains("account-user-one")); +} diff --git a/codex-rs/user-verification/src/lib.rs b/codex-rs/user-verification/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..7b8b415342a94e393d431d33ba6cdc6f8f4a2ced --- /dev/null +++ b/codex-rs/user-verification/src/lib.rs @@ -0,0 +1,91 @@ +//! Device credentials and signing, independent of RPC routing, UI, and backend registration. + +mod credential; +mod error; +mod guard; +mod key_namespace; +#[cfg(any(target_os = "macos", test))] +mod lifecycle_lock; +#[cfg(any(target_os = "macos", test))] +mod native_operation; +#[cfg(any(target_os = "macos", test))] +mod platform_macos; +#[cfg(not(target_os = "macos"))] +mod unsupported; + +pub use credential::UserVerificationKeyCreation; +pub use credential::UserVerificationKeyDeletion; +pub use credential::UserVerificationKeyInfo; +pub use credential::UserVerificationProof; +pub use credential::UserVerificationRequest; +pub use credential::UserVerificationStatus; +pub use error::UserVerificationCancellationReason; +pub use error::UserVerificationError; +pub use error::UserVerificationFailureReason; +pub use error::UserVerificationUnavailableReason; +pub use guard::UserVerificationRequestGuard; +pub use key_namespace::UserVerificationKeyNamespace; +use std::sync::Arc; + +/// Performs local credential operations for one captured account-user identity. +/// Implementations never perform network registration. Blocking implementations must run off +/// the async executor and check the guard after waiting, before effects, and before returning. +pub trait UserVerificationProvider: Send + Sync { + /// Reads local readiness without creating credentials or prompting for authentication. + fn status( + &self, + guard: &UserVerificationRequestGuard, + ) -> Result; + + /// Creates a protected key only if none exists. Success does not mean server enrollment. + fn ensure_key( + &self, + guard: &UserVerificationRequestGuard, + ) -> Result; + + /// Removes the local key idempotently. Backend revocation belongs to the caller. + fn delete( + &self, + guard: &UserVerificationRequestGuard, + ) -> Result; + + /// Authenticates and signs 1–4096 challenge bytes without interpreting an elicitation. + /// The caller owns approval UI, request correlation, and the captured identity check. + fn verify( + &self, + request: &UserVerificationRequest, + guard: &UserVerificationRequestGuard, + ) -> Result; +} + +/// Reports whether this build contains a native provider, independently of local readiness. +pub fn platform_supported() -> bool { + cfg!(target_os = "macos") +} + +/// Probes biometric hardware without reading credentials, prompting, or checking enrollment. +/// This performs local OS work; async callers should run it off their executor during setup. +pub fn device_supported() -> bool { + #[cfg(target_os = "macos")] + { + platform_macos::device_supported() + } + #[cfg(not(target_os = "macos"))] + { + false + } +} + +pub fn platform_provider( + namespace: UserVerificationKeyNamespace, +) -> Arc { + #[cfg(target_os = "macos")] + { + Arc::new(platform_macos::NativeProvider { namespace }) + } + #[cfg(not(target_os = "macos"))] + { + let _ = namespace.label; + Arc::new(unsupported::UnsupportedProvider) + } +} diff --git a/codex-rs/user-verification/src/lifecycle_lock.rs b/codex-rs/user-verification/src/lifecycle_lock.rs new file mode 100644 index 0000000000000000000000000000000000000000..69614e2fc67f47272197e59ab67ac3bc89f5c6c1 --- /dev/null +++ b/codex-rs/user-verification/src/lifecycle_lock.rs @@ -0,0 +1,86 @@ +//! Serializes operations for one account across local processes, with cancellable lock waits. + +use crate::UserVerificationError; +use crate::UserVerificationFailureReason; +use crate::UserVerificationRequestGuard; +use std::fs; +use std::fs::File; +use std::fs::OpenOptions; +use std::path::Path; +use std::time::Duration; +use std::time::Instant; + +pub(crate) struct LifecycleLock { + _file: File, +} + +impl LifecycleLock { + #[cfg(target_os = "macos")] + pub(crate) fn acquire( + namespace: &crate::UserVerificationKeyNamespace, + guard: &UserVerificationRequestGuard, + ) -> Result { + let directory = dirs::data_local_dir().ok_or_else(|| UserVerificationError::Failed { + reason: UserVerificationFailureReason::ProviderError, + message: "could not locate the credential lock directory".to_string(), + })?; + Self::acquire_at( + &directory + .join("com.openai.codex") + .join("user-verification") + .join(format!("{}.lock", namespace.label)), + Duration::from_secs(/*secs*/ 60), + guard, + ) + } + + fn acquire_at( + path: &Path, + timeout: Duration, + guard: &UserVerificationRequestGuard, + ) -> Result { + guard.check()?; + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).map_err(lock_error)?; + } + let file = OpenOptions::new() + .read(true) + .write(true) + .create(true) + .truncate(false) + .open(path) + .map_err(lock_error)?; + let started = Instant::now(); + loop { + guard.check()?; + match file.try_lock() { + Ok(()) => { + guard.check()?; + return Ok(Self { _file: file }); + } + Err(std::fs::TryLockError::WouldBlock) if started.elapsed() >= timeout => { + return Err(UserVerificationError::Failed { + reason: UserVerificationFailureReason::Timeout, + message: "timed out waiting for another credential operation".to_string(), + }); + } + Err(std::fs::TryLockError::WouldBlock) => { + std::thread::sleep(Duration::from_millis(/*millis*/ 50).min(timeout)); + } + Err(std::fs::TryLockError::Error(error)) => return Err(lock_error(error)), + } + } + } +} + +fn lock_error(error: std::io::Error) -> UserVerificationError { + tracing::warn!(%error, "user-verification lifecycle lock failed"); + UserVerificationError::Failed { + reason: UserVerificationFailureReason::ProviderError, + message: "could not acquire the credential lock".to_string(), + } +} + +#[cfg(test)] +#[path = "lifecycle_lock_tests.rs"] +mod tests; diff --git a/codex-rs/user-verification/src/lifecycle_lock_tests.rs b/codex-rs/user-verification/src/lifecycle_lock_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..c0b024bec66844dd849d43e4dcacd60dbf48bca5 --- /dev/null +++ b/codex-rs/user-verification/src/lifecycle_lock_tests.rs @@ -0,0 +1,60 @@ +//! Exercises real file-lock contention and cancellation without releasing the held lock. + +use super::*; +use crate::UserVerificationCancellationReason; +use pretty_assertions::assert_eq; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; + +#[test] +fn lock_serializes_operations_and_releases_on_drop() { + let directory = tempfile::tempdir().expect("temporary directory"); + let path = directory.path().join("credential.lock"); + let guard = UserVerificationRequestGuard::default(); + let first = LifecycleLock::acquire_at(&path, Duration::ZERO, &guard).expect("first lock"); + assert_eq!( + LifecycleLock::acquire_at(&path, Duration::ZERO, &guard).err(), + Some(UserVerificationError::Failed { + reason: UserVerificationFailureReason::Timeout, + message: "timed out waiting for another credential operation".to_string(), + }) + ); + drop(first); + LifecycleLock::acquire_at(&path, Duration::ZERO, &guard).expect("lock after release"); +} + +#[test] +fn cancelled_waiter_stops_while_another_operation_still_holds_lock() { + let directory = tempfile::tempdir().expect("temporary directory"); + let path = directory.path().join("credential.lock"); + let holder = LifecycleLock::acquire_at( + &path, + Duration::ZERO, + &UserVerificationRequestGuard::default(), + ) + .expect("held lock"); + let (ready_tx, ready_rx) = std::sync::mpsc::channel(); + let checks = AtomicUsize::new(/*v*/ 0); + let guard = UserVerificationRequestGuard::with_activity_check(move || { + if checks.fetch_add(/*val*/ 1, Ordering::Relaxed) == 1 { + ready_tx.send(()).expect("notify lock wait"); + } + true + }); + let queued = guard.clone(); + let worker = std::thread::spawn(move || { + LifecycleLock::acquire_at(&path, Duration::from_secs(/*secs*/ 5), &queued).err() + }); + ready_rx + .recv_timeout(Duration::from_secs(/*secs*/ 5)) + .expect("waiter entered lock loop"); + guard.cancel(); + assert_eq!( + worker.join().expect("waiter stopped"), + Some(UserVerificationError::Cancelled { + reason: UserVerificationCancellationReason::Interrupted, + message: "the verification operation is no longer active".to_string(), + }) + ); + drop(holder); +} diff --git a/codex-rs/user-verification/src/native_operation.rs b/codex-rs/user-verification/src/native_operation.rs new file mode 100644 index 0000000000000000000000000000000000000000..234d25499d000942b8938174cf867809e21aa133 --- /dev/null +++ b/codex-rs/user-verification/src/native_operation.rs @@ -0,0 +1,54 @@ +//! Keeps the authentication context on its owner thread while a native operation blocks. + +use crate::UserVerificationError; +use crate::UserVerificationFailureReason; +use crate::UserVerificationRequestGuard; +use std::sync::mpsc; +use std::time::Duration; + +pub(crate) fn run_with_cancellation( + guard: &UserVerificationRequestGuard, + operation: impl FnOnce() -> Result + Send, + cancel: impl FnOnce(), +) -> Result { + guard.check()?; + let (done_tx, done_rx) = mpsc::sync_channel(1); + std::thread::scope(|scope| { + let worker = std::thread::Builder::new() + .name("user-verification-sign".to_string()) + .spawn_scoped(scope, move || { + let _ = done_tx.send(guard.check().and_then(|()| operation())); + }) + .map_err(|error| { + tracing::warn!(%error, "could not start user-verification signer"); + worker_error() + })?; + let result = loop { + match done_rx.recv_timeout(Duration::from_millis(/*millis*/ 50)) { + Ok(result) => break Some(result), + Err(mpsc::RecvTimeoutError::Timeout) => { + if !guard.is_active() { + cancel(); + break None; + } + } + Err(mpsc::RecvTimeoutError::Disconnected) => break Some(Err(worker_error())), + } + }; + // Join before returning so no signer outlives the context or its lifecycle lock. + worker.join().map_err(|_| worker_error())?; + guard.check()?; + result.unwrap_or_else(|| Err(worker_error())) + }) +} + +fn worker_error() -> UserVerificationError { + UserVerificationError::Failed { + reason: UserVerificationFailureReason::ProviderError, + message: "the user-verification signer could not complete".to_string(), + } +} + +#[cfg(test)] +#[path = "native_operation_tests.rs"] +mod tests; diff --git a/codex-rs/user-verification/src/native_operation_tests.rs b/codex-rs/user-verification/src/native_operation_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..95b30ab2505f7e7be3ab2a8084feb739d8b77e85 --- /dev/null +++ b/codex-rs/user-verification/src/native_operation_tests.rs @@ -0,0 +1,83 @@ +//! Exercises prompt teardown and late-result rejection without requiring biometric hardware. + +use super::*; +use crate::UserVerificationCancellationReason; +use pretty_assertions::assert_eq; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +#[test] +fn cancellation_tears_down_the_prompt_on_the_context_owner_thread() { + let guard = UserVerificationRequestGuard::default(); + let queued = guard.clone(); + let (started_tx, started_rx) = mpsc::channel(); + let (dismissed_tx, dismissed_rx) = mpsc::channel(); + let (exited_tx, exited_rx) = mpsc::channel(); + let owner = std::thread::spawn(move || { + let owner_thread = std::thread::current().id(); + run_with_cancellation( + &queued, + move || { + started_tx.send(()).expect("authentication started"); + dismissed_rx + .recv_timeout(Duration::from_secs(/*secs*/ 5)) + .expect("prompt dismissed before the deadline"); + exited_tx.send(()).expect("signer exited"); + Ok("late signature") + }, + || { + assert_eq!(std::thread::current().id(), owner_thread); + dismissed_tx.send(()).expect("dismiss authentication"); + }, + ) + }); + started_rx + .recv_timeout(Duration::from_secs(/*secs*/ 5)) + .expect("signer is waiting for authentication"); + guard.cancel(); + assert_eq!( + owner.join().expect("owner returned"), + Err(UserVerificationError::Cancelled { + reason: UserVerificationCancellationReason::Interrupted, + message: "the verification operation is no longer active".to_string(), + }) + ); + exited_rx.try_recv().expect("signer exited before return"); +} + +#[test] +fn completed_operation_does_not_cancel_its_context() { + let cancelled = AtomicBool::new(/*v*/ false); + assert_eq!( + run_with_cancellation( + &UserVerificationRequestGuard::default(), + || Ok("signature"), + || cancelled.store(/*val*/ true, Ordering::Release), + ), + Ok("signature") + ); + assert!(!cancelled.load(Ordering::Acquire)); +} + +#[test] +fn already_cancelled_operation_never_starts_a_signer() { + let guard = UserVerificationRequestGuard::default(); + guard.cancel(); + let started = AtomicBool::new(/*v*/ false); + let result = run_with_cancellation( + &guard, + || { + started.store(/*val*/ true, Ordering::Release); + Ok(()) + }, + || panic!("no context needs cancellation"), + ); + assert_eq!( + result, + Err(UserVerificationError::Cancelled { + reason: UserVerificationCancellationReason::Interrupted, + message: "the verification operation is no longer active".to_string(), + }) + ); + assert!(!started.load(Ordering::Acquire)); +} diff --git a/codex-rs/user-verification/src/platform_macos.rs b/codex-rs/user-verification/src/platform_macos.rs new file mode 100644 index 0000000000000000000000000000000000000000..17b76115907d777e2207cf965b19fe25b3e8088a --- /dev/null +++ b/codex-rs/user-verification/src/platform_macos.rs @@ -0,0 +1,12 @@ +//! macOS Security and LocalAuthentication integration. + +mod error; +#[cfg(target_os = "macos")] +mod key_protection; +#[cfg(target_os = "macos")] +mod provider; + +#[cfg(target_os = "macos")] +pub(crate) use provider::NativeProvider; +#[cfg(target_os = "macos")] +pub(crate) use provider::device_supported; diff --git a/codex-rs/user-verification/src/platform_macos/error.rs b/codex-rs/user-verification/src/platform_macos/error.rs new file mode 100644 index 0000000000000000000000000000000000000000..0789d20738d24ada94a3f566e9ff8b108dc2e754 --- /dev/null +++ b/codex-rs/user-verification/src/platform_macos/error.rs @@ -0,0 +1,54 @@ +//! Maps macOS Security and LocalAuthentication failures to stable errors without exposing localized OS diagnostics. + +use crate::UserVerificationCancellationReason; +use crate::UserVerificationError; +use crate::UserVerificationFailureReason; +use crate::UserVerificationUnavailableReason; + +pub(crate) fn classify(domain: &str, code: i64) -> UserVerificationError { + tracing::debug!(domain, code, "native user-verification operation failed"); + match (domain, code) { + ("NSOSStatusErrorDomain", -128) | ("com.apple.LocalAuthentication", -2 | -3) => { + UserVerificationError::Cancelled { + reason: UserVerificationCancellationReason::UserCancelled, + message: "authentication was cancelled".to_string(), + } + } + ("com.apple.LocalAuthentication", -4 | -9) => UserVerificationError::Cancelled { + reason: UserVerificationCancellationReason::Interrupted, + message: "authentication was interrupted".to_string(), + }, + ("NSOSStatusErrorDomain", -25293) | ("com.apple.LocalAuthentication", -1) => { + UserVerificationError::Failed { + reason: UserVerificationFailureReason::AuthenticationFailed, + message: "biometric authentication failed".to_string(), + } + } + ("com.apple.LocalAuthentication", -5 | -6 | -7 | -8 | -12 | -13) => { + UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::BiometricsUnavailable, + message: "biometric authentication is not available right now".to_string(), + } + } + ("com.apple.LocalAuthentication", -1004) => UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::ProviderUnavailable, + message: "authentication UI is not permitted in this process".to_string(), + }, + ("NSOSStatusErrorDomain", -34018) => UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::ProviderUnavailable, + message: "this binary is missing the required keychain entitlements".to_string(), + }, + ("NSOSStatusErrorDomain", -25291 | -25308) => UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::ProviderUnavailable, + message: "the platform keychain is not available right now".to_string(), + }, + _ => UserVerificationError::Failed { + reason: UserVerificationFailureReason::ProviderError, + message: "the platform could not complete user verification".to_string(), + }, + } +} + +#[cfg(test)] +#[path = "error_tests.rs"] +mod tests; diff --git a/codex-rs/user-verification/src/platform_macos/error_tests.rs b/codex-rs/user-verification/src/platform_macos/error_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..d6cca80c11bd3b77b0d2445c0d2a4e54cb583d56 --- /dev/null +++ b/codex-rs/user-verification/src/platform_macos/error_tests.rs @@ -0,0 +1,51 @@ +//! Native failure classification is based on the domain and code, never localized text. + +use super::*; +use pretty_assertions::assert_eq; + +#[test] +fn cancellation_is_distinct_from_authentication_failure() { + assert_eq!( + classify("com.apple.LocalAuthentication", /*code*/ -2), + UserVerificationError::Cancelled { + reason: UserVerificationCancellationReason::UserCancelled, + message: "authentication was cancelled".to_string(), + } + ); + assert_eq!( + classify("NSOSStatusErrorDomain", /*code*/ -25293), + UserVerificationError::Failed { + reason: UserVerificationFailureReason::AuthenticationFailed, + message: "biometric authentication failed".to_string(), + } + ); +} + +#[test] +fn biometric_lockout_and_missing_entitlements_report_distinct_unavailability() { + assert_eq!( + classify("com.apple.LocalAuthentication", /*code*/ -8), + UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::BiometricsUnavailable, + message: "biometric authentication is not available right now".to_string(), + } + ); + assert_eq!( + classify("NSOSStatusErrorDomain", /*code*/ -34018), + UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::ProviderUnavailable, + message: "this binary is missing the required keychain entitlements".to_string(), + } + ); +} + +#[test] +fn unknown_error_domain_cannot_impersonate_cancellation() { + assert_eq!( + classify("unknown", /*code*/ -2), + UserVerificationError::Failed { + reason: UserVerificationFailureReason::ProviderError, + message: "the platform could not complete user verification".to_string(), + } + ); +} diff --git a/codex-rs/user-verification/src/platform_macos/key_protection.rs b/codex-rs/user-verification/src/platform_macos/key_protection.rs new file mode 100644 index 0000000000000000000000000000000000000000..a5b9003f25eaacea80f59f262ec49f346332c1cd --- /dev/null +++ b/codex-rs/user-verification/src/platform_macos/key_protection.rs @@ -0,0 +1,61 @@ +//! Create biometric signing policies and validate Secure Enclave key attributes. + +use super::error; +use crate::UserVerificationError; +use crate::UserVerificationUnavailableReason; +use core_foundation::base::CFEqual; +use core_foundation::base::TCFType as _; +use core_foundation::dictionary::CFDictionary; +use core_foundation::number::CFNumber; +use security_framework::access_control::ProtectionMode; +use security_framework::access_control::SecAccessControl; +use security_framework_sys::access_control::kSecAccessControlBiometryAny; +use security_framework_sys::access_control::kSecAccessControlPrivateKeyUsage; +use security_framework_sys::item::kSecAttrKeyClass; +use security_framework_sys::item::kSecAttrKeyClassPrivate; +use security_framework_sys::item::kSecAttrKeySizeInBits; +use security_framework_sys::item::kSecAttrKeyType; +use security_framework_sys::item::kSecAttrKeyTypeECSECPrimeRandom; +use security_framework_sys::item::kSecAttrTokenID; +use security_framework_sys::item::kSecAttrTokenIDSecureEnclave; + +pub(super) fn access_control() -> Result { + SecAccessControl::create_with_protection( + Some(ProtectionMode::AccessibleWhenUnlockedThisDeviceOnly), + kSecAccessControlBiometryAny | kSecAccessControlPrivateKeyUsage, + ) + .map_err(|err| error::classify("NSOSStatusErrorDomain", i64::from(err.code()))) +} + +pub(super) fn validate(attributes: &CFDictionary) -> Result<(), UserVerificationError> { + let size = CFNumber::from(/*value*/ 256); + // Do not compare opaque SecAccessControl objects: macOS can change their persisted + // representation. Biometric signing is enforced by the policy set at key creation. + // Reuse trusts that keys in this namespace, within the process's entitled keychain + // access groups, were created with that policy; these checks do not revalidate it. + let expected = unsafe { + [ + (kSecAttrTokenID, kSecAttrTokenIDSecureEnclave.cast()), + (kSecAttrKeyClass, kSecAttrKeyClassPrivate.cast()), + (kSecAttrKeyType, kSecAttrKeyTypeECSECPrimeRandom.cast()), + (kSecAttrKeySizeInBits, size.as_CFTypeRef()), + ] + }; + if expected.into_iter().all(|(name, expected)| { + attributes + .find(name.cast()) + .is_some_and(|value| unsafe { CFEqual(*value, expected) != 0 }) + }) { + Ok(()) + } else { + Err(UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::ProviderUnavailable, + message: "the local credential has incompatible key protection; delete it before enrolling again" + .to_string(), + }) + } +} + +#[cfg(test)] +#[path = "key_protection_tests.rs"] +mod tests; diff --git a/codex-rs/user-verification/src/platform_macos/key_protection_tests.rs b/codex-rs/user-verification/src/platform_macos/key_protection_tests.rs new file mode 100644 index 0000000000000000000000000000000000000000..dc44bf2f8ce2acafc4dddb98fa4dcbb29d9064fa --- /dev/null +++ b/codex-rs/user-verification/src/platform_macos/key_protection_tests.rs @@ -0,0 +1,59 @@ +//! Reject keys without the required Secure Enclave key attributes. + +use super::*; +use core_foundation::base::CFType; +use core_foundation::dictionary::CFMutableDictionary; +use core_foundation::string::CFString; +use pretty_assertions::assert_eq; + +fn protected_attributes() -> CFMutableDictionary { + unsafe { + CFMutableDictionary::from_CFType_pairs(&[ + ( + CFString::wrap_under_get_rule(kSecAttrTokenID), + CFString::wrap_under_get_rule(kSecAttrTokenIDSecureEnclave).as_CFType(), + ), + ( + CFString::wrap_under_get_rule(kSecAttrKeyClass), + CFString::wrap_under_get_rule(kSecAttrKeyClassPrivate).as_CFType(), + ), + ( + CFString::wrap_under_get_rule(kSecAttrKeyType), + CFString::wrap_under_get_rule(kSecAttrKeyTypeECSECPrimeRandom).as_CFType(), + ), + ( + CFString::wrap_under_get_rule(kSecAttrKeySizeInBits), + CFNumber::from(/*value*/ 256).as_CFType(), + ), + ]) + } +} + +#[test] +fn reuse_requires_secure_enclave_key_attributes() { + let attributes = protected_attributes(); + assert_eq!(validate(&attributes.to_immutable().to_untyped()), Ok(())); + let alternatives = unsafe { + [ + (kSecAttrTokenID, CFString::new("software").as_CFType()), + (kSecAttrKeyClass, CFString::new("public").as_CFType()), + (kSecAttrKeyType, CFString::new("RSA").as_CFType()), + ( + kSecAttrKeySizeInBits, + CFNumber::from(/*value*/ 384).as_CFType(), + ), + ] + }; + for (attribute, value) in alternatives { + let mut attributes = protected_attributes(); + attributes.set(unsafe { CFString::wrap_under_get_rule(attribute) }, value); + assert_eq!( + validate(&attributes.to_immutable().to_untyped()), + Err(UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::ProviderUnavailable, + message: "the local credential has incompatible key protection; delete it before enrolling again" + .to_string(), + }) + ); + } +} diff --git a/codex-rs/user-verification/src/platform_macos/provider.rs b/codex-rs/user-verification/src/platform_macos/provider.rs new file mode 100644 index 0000000000000000000000000000000000000000..9c532b275f0c29656e31dc1b20b4338b74def98b --- /dev/null +++ b/codex-rs/user-verification/src/platform_macos/provider.rs @@ -0,0 +1,293 @@ +//! Secure Enclave keys; biometric access is enforced by the key, not a separate UI check. + +use super::error; +use super::key_protection; +use crate::UserVerificationError; +use crate::UserVerificationFailureReason; +use crate::UserVerificationKeyCreation; +use crate::UserVerificationKeyDeletion; +use crate::UserVerificationKeyInfo; +use crate::UserVerificationKeyNamespace; +use crate::UserVerificationProof; +use crate::UserVerificationProvider; +use crate::UserVerificationRequest; +use crate::UserVerificationRequestGuard; +use crate::UserVerificationStatus; +use crate::UserVerificationUnavailableReason; +use crate::lifecycle_lock::LifecycleLock; +use crate::native_operation::run_with_cancellation; +use base64::Engine as _; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use core_foundation::base::CFType; +use core_foundation::base::TCFType as _; +use objc2::rc::Retained; +use objc2_foundation::NSString; +use objc2_local_authentication::LABiometryType; +use objc2_local_authentication::LAContext; +use objc2_local_authentication::LAPolicy; +use security_framework::base::Error as SecurityError; +use security_framework::item::ItemClass; +use security_framework::item::ItemSearchOptions; +use security_framework::item::KeyClass; +use security_framework::item::Limit; +use security_framework::item::Location; +use security_framework::item::Reference; +use security_framework::item::SearchResult; +use security_framework::key::Algorithm; +use security_framework::key::GenerateKeyOptions; +use security_framework::key::KeyType; +use security_framework::key::SecKey; +use security_framework::key::Token; +use security_framework_sys::base::errSecItemNotFound; + +pub(crate) struct NativeProvider { + pub(crate) namespace: UserVerificationKeyNamespace, +} + +pub(crate) fn device_supported() -> bool { + let context = noninteractive_context(); + // canEvaluatePolicy populates biometryType even when enrollment or temporary lockout + // prevents authentication. Advertise hardware support independently of current readiness. + unsafe { + let _ = context.canEvaluatePolicy_error(LAPolicy::DeviceOwnerAuthenticationWithBiometrics); + context.biometryType() == LABiometryType::TouchID + } +} + +impl UserVerificationProvider for NativeProvider { + fn status( + &self, + guard: &UserVerificationRequestGuard, + ) -> Result { + let _lock = LifecycleLock::acquire(&self.namespace, guard)?; + let context = noninteractive_context(); + let credential = find_key(&self.namespace.label, &context, KeyUse::Protected)? + .as_ref() + .map(key_info) + .transpose()?; + let unavailable = if credential.is_none() { + Some(( + UserVerificationUnavailableReason::CredentialMissing, + "no user-verification credential has been created".to_string(), + )) + } else { + match check_biometrics(&context) { + Ok(()) => None, + Err(UserVerificationError::Unavailable { reason, message }) => { + Some((reason, message)) + } + Err(error) => return Err(error), + } + }; + let (unavailable_reason, unavailable_message) = match unavailable { + Some((reason, message)) => (Some(reason), Some(message)), + None => (None, None), + }; + guard.check()?; + Ok(UserVerificationStatus { + credential, + unavailable_reason, + unavailable_message, + }) + } + + fn ensure_key( + &self, + guard: &UserVerificationRequestGuard, + ) -> Result { + let _lock = LifecycleLock::acquire(&self.namespace, guard)?; + let context = noninteractive_context(); + let existing = find_key(&self.namespace.label, &context, KeyUse::Protected)?; + guard.check()?; + let (key, created) = match existing { + Some(key) => (key, false), + None => (create_key(&self.namespace.label)?, true), + }; + let credential = key_info(&key)?; + guard.check()?; + Ok(UserVerificationKeyCreation { + created, + credential, + }) + } + + fn delete( + &self, + guard: &UserVerificationRequestGuard, + ) -> Result { + let _lock = LifecycleLock::acquire(&self.namespace, guard)?; + let context = noninteractive_context(); + // The credential ID is optional information for the caller. A colliding key may + // have no P-256 public key, or its lookup may fail because authentication is + // unavailable; neither should prevent removing this namespace's key items. + let credential = find_key(&self.namespace.label, &context, KeyUse::Deletion) + .ok() + .flatten() + .and_then(|key| key_info(&key).ok()); + guard.check()?; + let mut options = ItemSearchOptions::new(); + // GenerateKeyOptions persists both halves on macOS. Remove all key items under + // this exact namespace, including public keys left by an interrupted deletion. + options + .ignore_legacy_keychains() + .class(ItemClass::key()) + .label(&self.namespace.label); + match options.delete() { + Ok(()) => {} + Err(error) if error.code() == errSecItemNotFound => {} + Err(error) => return Err(keychain_error(error)), + } + guard.check()?; + Ok(UserVerificationKeyDeletion { + deleted_credential_id: credential.map(|key| key.credential_id), + }) + } + + fn verify( + &self, + request: &UserVerificationRequest, + guard: &UserVerificationRequestGuard, + ) -> Result { + guard.check()?; + if request.challenge.is_empty() + || request.challenge.len() > 4096 + || request.title.is_empty() + || request.title.len() > 256 + || request.description.len() > 4096 + { + return Err(UserVerificationError::Failed { + reason: UserVerificationFailureReason::ProviderError, + message: "invalid user-verification challenge or display text".to_string(), + }); + } + let _lock = LifecycleLock::acquire(&self.namespace, guard)?; + // A fresh context for every signature prevents authentication reuse between actions. + let context = unsafe { LAContext::new() }; + let reason = if request.description.is_empty() { + request.title.clone() + } else { + format!("{}\n\n{}", request.title, request.description) + }; + // The keychain query retains this context and uses its text for the signing prompt. + unsafe { + context.setLocalizedFallbackTitle(Some(&NSString::from_str(""))); + context.setLocalizedReason(&NSString::from_str(&reason)); + } + check_biometrics(&context)?; + guard.check()?; + let key = + find_key(&self.namespace.label, &context, KeyUse::Protected)?.ok_or_else(|| { + UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::CredentialMissing, + message: "no user-verification credential has been created".to_string(), + } + })?; + let credential = key_info(&key)?; + guard.check()?; + let signature = run_with_cancellation( + guard, + || { + key.create_signature( + Algorithm::ECDSASignatureMessageX962SHA256, + &request.challenge, + ) + .map_err(|error| error::classify(&error.domain().to_string(), error.code() as i64)) + }, + || { + // LAContext remains on this thread. Invalidation cancels its pending + // keychain authentication, allowing the signer and lifecycle lock to exit. + unsafe { context.invalidate() }; + }, + ); + guard.check()?; + Ok(UserVerificationProof { + credential_id: credential.credential_id, + signature: URL_SAFE_NO_PAD.encode(signature?), + }) + } +} + +fn noninteractive_context() -> Retained { + // Each context is confined to this blocking operation and never shared across threads. + unsafe { + let context = LAContext::new(); + context.setInteractionNotAllowed(/*interaction_not_allowed*/ true); + context + } +} + +fn check_biometrics(context: &LAContext) -> Result<(), UserVerificationError> { + // canEvaluatePolicy only checks availability; the protected key triggers authentication. + unsafe { context.canEvaluatePolicy_error(LAPolicy::DeviceOwnerAuthenticationWithBiometrics) } + .map_err(|error| error::classify(&error.domain().to_string(), error.code() as i64)) +} + +enum KeyUse { + Protected, + Deletion, +} + +fn find_key( + label: &str, + context: &LAContext, + key_use: KeyUse, +) -> Result, UserVerificationError> { + // Security accepts an LAContext object as kSecUseAuthenticationContext. Wrapping under + // the get rule retains it, and ItemSearchOptions owns that retain for the query lifetime. + let authentication = unsafe { CFType::wrap_under_get_rule(std::ptr::from_ref(context).cast()) }; + let mut options = ItemSearchOptions::new(); + options + .ignore_legacy_keychains() + .key_class(KeyClass::private()) + .label(label) + .local_authentication_context(Some(authentication)) + .load_refs(true) + .limit(Limit::Max(1)); + match options.search() { + Ok(results) => { + let key = results.into_iter().find_map(|result| match result { + SearchResult::Ref(Reference::Key(key)) => Some(key), + _ => None, + }); + if let Some(key) = &key + && matches!(key_use, KeyUse::Protected) + { + key_protection::validate(&key.attributes())?; + } + Ok(key) + } + Err(error) if error.code() == errSecItemNotFound => Ok(None), + Err(error) => Err(keychain_error(error)), + } +} + +fn create_key(label: &str) -> Result { + let access = key_protection::access_control()?; + let mut options = GenerateKeyOptions::default(); + options + .set_key_type(KeyType::ec_sec_prime_random()) + .set_size_in_bits(256) + .set_label(label) + .set_token(Token::SecureEnclave) + .set_location(Location::DataProtectionKeychain) + .set_access_control(access); + let key = SecKey::new(&options) + .map_err(|error| error::classify(&error.domain().to_string(), error.code() as i64))?; + key_protection::validate(&key.attributes())?; + Ok(key) +} + +fn key_info(key: &SecKey) -> Result { + let bytes = key + .public_key() + .and_then(|key| key.external_representation()) + .ok_or_else(|| UserVerificationError::Failed { + reason: UserVerificationFailureReason::ProviderError, + message: "could not export the user-verification public key".to_string(), + })?; + UserVerificationKeyInfo::from_sec1_public_key(&bytes) +} + +fn keychain_error(error: SecurityError) -> UserVerificationError { + error::classify("NSOSStatusErrorDomain", i64::from(error.code())) +} diff --git a/codex-rs/user-verification/src/unsupported.rs b/codex-rs/user-verification/src/unsupported.rs new file mode 100644 index 0000000000000000000000000000000000000000..200fe7d212e2e0ebad0a1e6ecec18fa628416d2b --- /dev/null +++ b/codex-rs/user-verification/src/unsupported.rs @@ -0,0 +1,53 @@ +//! Typed unavailable behavior until a native provider is available for this platform. + +use crate::*; + +pub(crate) struct UnsupportedProvider; + +fn unavailable() -> UserVerificationError { + UserVerificationError::Unavailable { + reason: UserVerificationUnavailableReason::ProviderUnavailable, + message: "this platform does not support user verification".to_string(), + } +} + +impl UserVerificationProvider for UnsupportedProvider { + fn status( + &self, + guard: &UserVerificationRequestGuard, + ) -> Result { + guard.check()?; + Ok(UserVerificationStatus { + credential: None, + unavailable_reason: Some(UserVerificationUnavailableReason::ProviderUnavailable), + unavailable_message: Some( + "this platform does not support user verification".to_string(), + ), + }) + } + + fn ensure_key( + &self, + guard: &UserVerificationRequestGuard, + ) -> Result { + guard.check()?; + Err(unavailable()) + } + + fn delete( + &self, + guard: &UserVerificationRequestGuard, + ) -> Result { + guard.check()?; + Err(unavailable()) + } + + fn verify( + &self, + _request: &UserVerificationRequest, + guard: &UserVerificationRequestGuard, + ) -> Result { + guard.check()?; + Err(unavailable()) + } +} diff --git a/tools/argument-comment-lint/.cargo/config.toml b/tools/argument-comment-lint/.cargo/config.toml new file mode 100644 index 0000000000000000000000000000000000000000..226eca535bf5f246cbbd0fea86a23034e5b44db1 --- /dev/null +++ b/tools/argument-comment-lint/.cargo/config.toml @@ -0,0 +1,6 @@ +[target.'cfg(all())'] +rustflags = ["-C", "linker=dylint-link"] + +# For Rust versions 1.74.0 and onward, the following alternative can be used +# (see https://github.com/rust-lang/cargo/pull/12535): +# linker = "dylint-link" diff --git a/tools/argument-comment-lint/.gitignore b/tools/argument-comment-lint/.gitignore new file mode 100644 index 0000000000000000000000000000000000000000..ea8c4bf7f35f6f77f75d92ad8ce8349f6e81ddba --- /dev/null +++ b/tools/argument-comment-lint/.gitignore @@ -0,0 +1 @@ +/target diff --git a/tools/argument-comment-lint/BUILD.bazel b/tools/argument-comment-lint/BUILD.bazel new file mode 100644 index 0000000000000000000000000000000000000000..5d2695a80df4462ddf9b90f68865ffa1eb68a4ce --- /dev/null +++ b/tools/argument-comment-lint/BUILD.bazel @@ -0,0 +1,32 @@ +load("@rules_rust//rust:defs.bzl", "rust_binary", "rust_library") + +exports_files(["lint_aspect.bzl"]) + +rust_library( + name = "argument-comment-lint-lib", + srcs = [ + "src/comment_parser.rs", + "src/lib.rs", + ], + crate_features = ["bazel_native"], + crate_name = "argument_comment_lint", + crate_root = "src/lib.rs", + edition = "2024", + tags = ["manual"], + visibility = ["//visibility:public"], + deps = ["@argument_comment_lint_crates//:clippy_utils"], +) + +rust_binary( + name = "argument-comment-lint-driver", + srcs = ["driver.rs"], + crate_name = "argument_comment_lint_driver", + crate_root = "driver.rs", + edition = "2024", + tags = ["manual"], + visibility = ["//visibility:public"], + deps = [ + ":argument-comment-lint-lib", + "@zlib//:z", + ], +) diff --git a/tools/argument-comment-lint/Cargo.lock b/tools/argument-comment-lint/Cargo.lock new file mode 100644 index 0000000000000000000000000000000000000000..1795ff683b86672de2bd2a48a41b4274e3f63d2a --- /dev/null +++ b/tools/argument-comment-lint/Cargo.lock @@ -0,0 +1,1653 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "aho-corasick" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddd31a130427c27518df266943a5308ed92d4b226cc639f5a8f1002816174301" +dependencies = [ + "memchr", +] + +[[package]] +name = "anstream" +version = "0.6.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "43d5b281e737544384e969a5ccad3f1cdd24b48086a0fc1b2a5262a26b8f4f4a" +dependencies = [ + "anstyle", + "anstyle-parse", + "anstyle-query", + "anstyle-wincon", + "colorchoice", + "is_terminal_polyfill", + "utf8parse", +] + +[[package]] +name = "anstyle" +version = "1.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" + +[[package]] +name = "anstyle-parse" +version = "0.2.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4e7644824f0aa2c7b9384579234ef10eb7efb6a0deb83f9630a49594dd9c15c2" +dependencies = [ + "utf8parse", +] + +[[package]] +name = "anstyle-query" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "anstyle-wincon" +version = "3.0.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" +dependencies = [ + "anstyle", + "once_cell_polyfill", + "windows-sys 0.61.2", +] + +[[package]] +name = "anyhow" +version = "1.0.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" + +[[package]] +name = "argument_comment_lint" +version = "0.1.0" +dependencies = [ + "clippy_utils", + "dylint_linting", + "dylint_testing", +] + +[[package]] +name = "arrayvec" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" + +[[package]] +name = "bitflags" +version = "2.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af" + +[[package]] +name = "camino" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e629a66d692cb9ff1a1c664e41771b3dcaf961985a9774c0eb0bd1b51cf60a48" +dependencies = [ + "serde_core", +] + +[[package]] +name = "cargo-platform" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "87a0c0e6148f11f01f32650a2ea02d532b2ad4e81d8bd41e6e565b5adc5e6082" +dependencies = [ + "serde", + "serde_core", +] + +[[package]] +name = "cargo_metadata" +version = "0.23.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ef987d17b0a113becdd19d3d0022d04d7ef41f9efe4f3fb63ac44ba61df3ade9" +dependencies = [ + "camino", + "cargo-platform", + "semver", + "serde", + "serde_json", + "thiserror 2.0.18", +] + +[[package]] +name = "cc" +version = "1.2.56" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "aebf35691d1bfb0ac386a69bac2fde4dd276fb618cf8bf4f5318fe285e821bb2" +dependencies = [ + "find-msvc-tools", + "jobserver", + "libc", + "shlex", +] + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "clippy_utils" +version = "0.1.92" +source = "git+https://github.com/rust-lang/rust-clippy?rev=20ce69b9a63bcd2756cd906fe0964d1e901e042a#20ce69b9a63bcd2756cd906fe0964d1e901e042a" +dependencies = [ + "arrayvec", + "itertools", + "rustc_apfloat", + "serde", +] + +[[package]] +name = "colorchoice" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" + +[[package]] +name = "compiletest_rs" +version = "0.11.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f150fe9105fcd2a57cad53f0c079a24de65195903ef670990f5909f695eac04c" +dependencies = [ + "diff", + "filetime", + "getopts", + "lazy_static", + "libc", + "log", + "miow", + "regex", + "rustfix", + "serde", + "serde_derive", + "serde_json", + "tester", + "windows-sys 0.59.0", +] + +[[package]] +name = "diff" +version = "0.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "56254986775e3233ffa9c4d7d3faaf6d36a2c09d30b20687e9f88bc8bafc16c8" + +[[package]] +name = "dirs-next" +version = "2.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b98cf8ebf19c3d1b223e151f99a4f9f0690dca41414773390fc824184ac833e1" +dependencies = [ + "cfg-if", + "dirs-sys-next", +] + +[[package]] +name = "dirs-sys-next" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ebda144c4fe02d1f7ea1a7d9641b6fc6b580adcfa024ae48797ecdeb6825b4d" +dependencies = [ + "libc", + "redox_users", + "winapi", +] + +[[package]] +name = "displaydoc" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "97369cbbc041bc366949bc74d34658d6cda5621039731c6310521892a3a20ae0" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "dylint" +version = "5.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "aa9a937bab540c0c8cdcbd650a572e6899ef2b6ffbc277d61bd2ae8d17c0edce" +dependencies = [ + "anstyle", + "anyhow", + "cargo_metadata", + "dylint_internal", + "log", + "once_cell", + "semver", + "serde", + "serde_json", + "tempfile", +] + +[[package]] +name = "dylint_internal" +version = "5.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2e11f358a59510be7fa5c4f412729fabbe31a3587a342e4241a6a72020a2a0c5" +dependencies = [ + "anstyle", + "anyhow", + "bitflags", + "cargo_metadata", + "git2", + "home", + "if_chain", + "log", + "regex", + "rustversion", + "serde", + "tar", + "thiserror 2.0.18", + "toml", +] + +[[package]] +name = "dylint_linting" +version = "5.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4588b33aafbd472a6468ad1521d74094faa2bbdb53d534c2d24320300ea94135" +dependencies = [ + "cargo_metadata", + "dylint_internal", + "paste", + "rustversion", + "serde", + "thiserror 2.0.18", + "toml", +] + +[[package]] +name = "dylint_testing" +version = "5.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cdc27f344ddb5488eb16b6e0f8aec889a30fb7e4d135060d336cfa60d1fd671c" +dependencies = [ + "anyhow", + "cargo_metadata", + "compiletest_rs", + "dylint", + "dylint_internal", + "env_logger", + "once_cell", + "regex", + "serde_json", + "tempfile", +] + +[[package]] +name = "either" +version = "1.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719" + +[[package]] +name = "env_filter" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7a1c3cc8e57274ec99de65301228b537f1e4eedc1b8e0f9411c6caac8ae7308f" +dependencies = [ + "log", + "regex", +] + +[[package]] +name = "env_logger" +version = "0.11.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b2daee4ea451f429a58296525ddf28b45a3b64f1acf6587e2067437bb11e218d" +dependencies = [ + "anstream", + "anstyle", + "env_filter", + "jiff", + "log", +] + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "fastrand" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be" + +[[package]] +name = "filetime" +version = "0.2.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f98844151eee8917efc50bd9e8318cb963ae8b297431495d3f758616ea5c57db" +dependencies = [ + "cfg-if", + "libc", + "libredox", +] + +[[package]] +name = "find-msvc-tools" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" + +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + +[[package]] +name = "form_urlencoded" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf" +dependencies = [ + "percent-encoding", +] + +[[package]] +name = "getopts" +version = "0.2.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfe4fbac503b8d1f88e6676011885f34b7174f46e59956bba534ba83abded4df" +dependencies = [ + "unicode-width", +] + +[[package]] +name = "getrandom" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" +dependencies = [ + "cfg-if", + "libc", + "wasi", +] + +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "libc", + "r-efi 5.3.0", + "wasip2", +] + +[[package]] +name = "getrandom" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0de51e6874e94e7bf76d726fc5d13ba782deca734ff60d5bb2fb2607c7406555" +dependencies = [ + "cfg-if", + "libc", + "r-efi 6.0.0", + "wasip2", + "wasip3", +] + +[[package]] +name = "git2" +version = "0.20.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b88256088d75a56f8ecfa070513a775dd9107f6530ef14919dac831af9cfe2b" +dependencies = [ + "bitflags", + "libc", + "libgit2-sys", + "log", + "openssl-probe", + "openssl-sys", + "url", +] + +[[package]] +name = "hashbrown" +version = "0.15.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" +dependencies = [ + "foldhash", +] + +[[package]] +name = "hashbrown" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" + +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + +[[package]] +name = "hermit-abi" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" + +[[package]] +name = "home" +version = "0.5.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3d1354bf6b7235cb4a0576c2619fd4ed18183f689b12b006a0ee7329eeff9a5" +dependencies = [ + "windows-sys 0.52.0", +] + +[[package]] +name = "icu_collections" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4c6b649701667bbe825c3b7e6388cb521c23d88644678e83c0c4d0a621a34b43" +dependencies = [ + "displaydoc", + "potential_utf", + "yoke", + "zerofrom", + "zerovec", +] + +[[package]] +name = "icu_locale_core" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "edba7861004dd3714265b4db54a3c390e880ab658fec5f7db895fae2046b5bb6" +dependencies = [ + "displaydoc", + "litemap", + "tinystr", + "writeable", + "zerovec", +] + +[[package]] +name = "icu_normalizer" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5f6c8828b67bf8908d82127b2054ea1b4427ff0230ee9141c54251934ab1b599" +dependencies = [ + "icu_collections", + "icu_normalizer_data", + "icu_properties", + "icu_provider", + "smallvec", + "zerovec", +] + +[[package]] +name = "icu_normalizer_data" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7aedcccd01fc5fe81e6b489c15b247b8b0690feb23304303a9e560f37efc560a" + +[[package]] +name = "icu_properties" +version = "2.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "020bfc02fe870ec3a66d93e677ccca0562506e5872c650f893269e08615d74ec" +dependencies = [ + "icu_collections", + "icu_locale_core", + "icu_properties_data", + "icu_provider", + "zerotrie", + "zerovec", +] + +[[package]] +name = "icu_properties_data" +version = "2.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "616c294cf8d725c6afcd8f55abc17c56464ef6211f9ed59cccffe534129c77af" + +[[package]] +name = "icu_provider" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85962cf0ce02e1e0a629cc34e7ca3e373ce20dda4c4d7294bbd0bf1fdb59e614" +dependencies = [ + "displaydoc", + "icu_locale_core", + "writeable", + "yoke", + "zerofrom", + "zerotrie", + "zerovec", +] + +[[package]] +name = "id-arena" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d3067d79b975e8844ca9eb072e16b31c3c1c36928edf9c6789548c524d0d954" + +[[package]] +name = "idna" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" +dependencies = [ + "idna_adapter", + "smallvec", + "utf8_iter", +] + +[[package]] +name = "idna_adapter" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3acae9609540aa318d1bc588455225fb2085b9ed0c4f6bd0d9d5bcd86f1a0344" +dependencies = [ + "icu_normalizer", + "icu_properties", +] + +[[package]] +name = "if_chain" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cd62e6b5e86ea8eeeb8db1de02880a6abc01a397b2ebb64b5d74ac255318f5cb" + +[[package]] +name = "indexmap" +version = "2.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7714e70437a7dc3ac8eb7e6f8df75fd8eb422675fc7678aff7364301092b1017" +dependencies = [ + "equivalent", + "hashbrown 0.16.1", + "serde", + "serde_core", +] + +[[package]] +name = "is_terminal_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" + +[[package]] +name = "itertools" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba291022dbbd398a455acf126c1e341954079855bc60dfdda641363bd6922569" +dependencies = [ + "either", +] + +[[package]] +name = "itoa" +version = "1.0.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92ecc6618181def0457392ccd0ee51198e065e016d1d527a7ac1b6dc7c1f09d2" + +[[package]] +name = "jiff" +version = "0.2.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1a3546dc96b6d42c5f24902af9e2538e82e39ad350b0c766eb3fbf2d8f3d8359" +dependencies = [ + "jiff-static", + "log", + "portable-atomic", + "portable-atomic-util", + "serde_core", +] + +[[package]] +name = "jiff-static" +version = "0.2.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2a8c8b344124222efd714b73bb41f8b5120b27a7cc1c75593a6ff768d9d05aa4" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "jobserver" +version = "0.1.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9afb3de4395d6b3e67a780b6de64b51c978ecf11cb9a462c66be7d4ca9039d33" +dependencies = [ + "getrandom 0.3.4", + "libc", +] + +[[package]] +name = "lazy_static" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" + +[[package]] +name = "leb128fmt" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09edd9e8b54e49e587e4f6295a7d29c3ea94d469cb40ab8ca70b288248a81db2" + +[[package]] +name = "libc" +version = "0.2.183" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b5b646652bf6661599e1da8901b3b9522896f01e736bad5f723fe7a3a27f899d" + +[[package]] +name = "libgit2-sys" +version = "0.18.3+1.9.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c9b3acc4b91781bb0b3386669d325163746af5f6e4f73e6d2d630e09a35f3487" +dependencies = [ + "cc", + "libc", + "libssh2-sys", + "libz-sys", + "openssl-sys", + "pkg-config", +] + +[[package]] +name = "libredox" +version = "0.1.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1744e39d1d6a9948f4f388969627434e31128196de472883b39f148769bfe30a" +dependencies = [ + "bitflags", + "libc", + "plain", + "redox_syscall", +] + +[[package]] +name = "libssh2-sys" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "220e4f05ad4a218192533b300327f5150e809b54c4ec83b5a1d91833601811b9" +dependencies = [ + "cc", + "libc", + "libz-sys", + "openssl-sys", + "pkg-config", + "vcpkg", +] + +[[package]] +name = "libz-sys" +version = "1.1.25" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d52f4c29e2a68ac30c9087e1b772dc9f44a2b66ed44edf2266cf2be9b03dafc1" +dependencies = [ + "cc", + "libc", + "pkg-config", + "vcpkg", +] + +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + +[[package]] +name = "litemap" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6373607a59f0be73a39b6fe456b8192fcc3585f602af20751600e974dd455e77" + +[[package]] +name = "log" +version = "0.4.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" + +[[package]] +name = "memchr" +version = "2.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" + +[[package]] +name = "miow" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "536bfad37a309d62069485248eeaba1e8d9853aaf951caaeaed0585a95346f08" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "num_cpus" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" +dependencies = [ + "hermit-abi", + "libc", +] + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "once_cell_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" + +[[package]] +name = "openssl-probe" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d05e27ee213611ffe7d6348b942e8f942b37114c00cc03cec254295a4a17852e" + +[[package]] +name = "openssl-sys" +version = "0.9.112" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57d55af3b3e226502be1526dfdba67ab0e9c96fc293004e79576b2b9edb0dbdb" +dependencies = [ + "cc", + "libc", + "pkg-config", + "vcpkg", +] + +[[package]] +name = "paste" +version = "1.0.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" + +[[package]] +name = "percent-encoding" +version = "2.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" + +[[package]] +name = "pin-project-lite" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" + +[[package]] +name = "pkg-config" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7edddbd0b52d732b21ad9a5fab5c704c14cd949e5e9a1ec5929a24fded1b904c" + +[[package]] +name = "plain" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4596b6d070b27117e987119b4dac604f3c58cfb0b191112e24771b2faeac1a6" + +[[package]] +name = "portable-atomic" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49" + +[[package]] +name = "portable-atomic-util" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7a9db96d7fa8782dd8c15ce32ffe8680bbd1e978a43bf51a34d39483540495f5" +dependencies = [ + "portable-atomic", +] + +[[package]] +name = "potential_utf" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b73949432f5e2a09657003c25bca5e19a0e9c84f8058ca374f49e0ebe605af77" +dependencies = [ + "zerovec", +] + +[[package]] +name = "prettyplease" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" +dependencies = [ + "proc-macro2", + "syn", +] + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "quote" +version = "1.0.45" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41f2619966050689382d2b44f664f4bc593e129785a36d6ee376ddf37259b924" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + +[[package]] +name = "r-efi" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" + +[[package]] +name = "redox_syscall" +version = "0.7.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce70a74e890531977d37e532c34d45e9055d2409ed08ddba14529471ed0be16" +dependencies = [ + "bitflags", +] + +[[package]] +name = "redox_users" +version = "0.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba009ff324d1fc1b900bd1fdb31564febe58a8ccc8a6fdbb93b543d33b13ca43" +dependencies = [ + "getrandom 0.2.17", + "libredox", + "thiserror 1.0.69", +] + +[[package]] +name = "regex" +version = "1.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e10754a14b9137dd7b1e3e5b0493cc9171fdd105e0ab477f51b72e7f3ac0e276" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc897dd8d9e8bd1ed8cdad82b5966c3e0ecae09fb1907d58efaa013543185d0a" + +[[package]] +name = "rustc_apfloat" +version = "0.2.3+llvm-462a31f5a5ab" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "486c2179b4796f65bfe2ee33679acf0927ac83ecf583ad6c91c3b4570911b9ad" +dependencies = [ + "bitflags", + "smallvec", +] + +[[package]] +name = "rustfix" +version = "0.8.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "82fa69b198d894d84e23afde8e9ab2af4400b2cba20d6bf2b428a8b01c222c5a" +dependencies = [ + "serde", + "serde_json", + "thiserror 1.0.69", + "tracing", +] + +[[package]] +name = "rustix" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.61.2", +] + +[[package]] +name = "rustversion" +version = "1.0.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d" + +[[package]] +name = "semver" +version = "1.0.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d767eb0aabc880b29956c35734170f26ed551a859dbd361d140cdbeca61ab1e2" +dependencies = [ + "serde", + "serde_core", +] + +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.149" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + +[[package]] +name = "serde_spanned" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8bbf91e5a4d6315eee45e704372590b30e260ee83af6639d64557f51b067776" +dependencies = [ + "serde_core", +] + +[[package]] +name = "shlex" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0fda2ff0d084019ba4d7c6f371c95d8fd75ce3524c3cb8fb653a3023f6323e64" + +[[package]] +name = "smallvec" +version = "1.15.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" + +[[package]] +name = "stable_deref_trait" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" + +[[package]] +name = "syn" +version = "2.0.117" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e665b8803e7b1d2a727f4023456bbbbe74da67099c585258af0ad9c5013b9b99" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "synstructure" +version = "0.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tar" +version = "0.4.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d863878d212c87a19c1a610eb53bb01fe12951c0501cf5a0d65f724914a667a" +dependencies = [ + "filetime", + "libc", + "xattr", +] + +[[package]] +name = "tempfile" +version = "3.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" +dependencies = [ + "fastrand", + "getrandom 0.4.2", + "once_cell", + "rustix", + "windows-sys 0.61.2", +] + +[[package]] +name = "term" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c59df8ac95d96ff9bede18eb7300b0fda5e5d8d90960e76f8e14ae765eedbf1f" +dependencies = [ + "dirs-next", + "rustversion", + "winapi", +] + +[[package]] +name = "tester" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "89e8bf7e0eb2dd7b4228cc1b6821fc5114cd6841ae59f652a85488c016091e5f" +dependencies = [ + "cfg-if", + "getopts", + "libc", + "num_cpus", + "term", +] + +[[package]] +name = "thiserror" +version = "1.0.69" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6aaf5339b578ea85b50e080feb250a3e8ae8cfcdff9a461c9ec2904bc923f52" +dependencies = [ + "thiserror-impl 1.0.69", +] + +[[package]] +name = "thiserror" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +dependencies = [ + "thiserror-impl 2.0.18", +] + +[[package]] +name = "thiserror-impl" +version = "1.0.69" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tinystr" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42d3e9c45c09de15d06dd8acf5f4e0e399e85927b7f00711024eb7ae10fa4869" +dependencies = [ + "displaydoc", + "zerovec", +] + +[[package]] +name = "toml" +version = "0.9.12+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf92845e79fc2e2def6a5d828f0801e29a2f8acc037becc5ab08595c7d5e9863" +dependencies = [ + "indexmap", + "serde_core", + "serde_spanned", + "toml_datetime", + "toml_parser", + "toml_writer", + "winnow", +] + +[[package]] +name = "toml_datetime" +version = "0.7.5+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92e1cfed4a3038bc5a127e35a2d360f145e1f4b971b551a2ba5fd7aedf7e1347" +dependencies = [ + "serde_core", +] + +[[package]] +name = "toml_parser" +version = "1.0.9+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "702d4415e08923e7e1ef96cd5727c0dfed80b4d2fa25db9647fe5eb6f7c5a4c4" +dependencies = [ + "winnow", +] + +[[package]] +name = "toml_writer" +version = "1.0.6+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab16f14aed21ee8bfd8ec22513f7287cd4a91aa92e44edfe2c17ddd004e92607" + +[[package]] +name = "tracing" +version = "0.1.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" +dependencies = [ + "pin-project-lite", + "tracing-core", +] + +[[package]] +name = "tracing-core" +version = "0.1.36" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" +dependencies = [ + "once_cell", +] + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "unicode-width" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4ac048d71ede7ee76d585517add45da530660ef4390e49b098733c6e897f254" + +[[package]] +name = "unicode-xid" +version = "0.2.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" + +[[package]] +name = "url" +version = "2.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff67a8a4397373c3ef660812acab3268222035010ab8680ec4215f38ba3d0eed" +dependencies = [ + "form_urlencoded", + "idna", + "percent-encoding", + "serde", +] + +[[package]] +name = "utf8_iter" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" + +[[package]] +name = "utf8parse" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" + +[[package]] +name = "vcpkg" +version = "0.2.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" + +[[package]] +name = "wasi" +version = "0.11.1+wasi-snapshot-preview1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" + +[[package]] +name = "wasip2" +version = "1.0.2+wasi-0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9517f9239f02c069db75e65f174b3da828fe5f5b945c4dd26bd25d89c03ebcf5" +dependencies = [ + "wit-bindgen", +] + +[[package]] +name = "wasip3" +version = "0.4.0+wasi-0.3.0-rc-2026-01-06" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5428f8bf88ea5ddc08faddef2ac4a67e390b88186c703ce6dbd955e1c145aca5" +dependencies = [ + "wit-bindgen", +] + +[[package]] +name = "wasm-encoder" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "990065f2fe63003fe337b932cfb5e3b80e0b4d0f5ff650e6985b1048f62c8319" +dependencies = [ + "leb128fmt", + "wasmparser", +] + +[[package]] +name = "wasm-metadata" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bb0e353e6a2fbdc176932bbaab493762eb1255a7900fe0fea1a2f96c296cc909" +dependencies = [ + "anyhow", + "indexmap", + "wasm-encoder", + "wasmparser", +] + +[[package]] +name = "wasmparser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47b807c72e1bac69382b3a6fb3dbe8ea4c0ed87ff5629b8685ae6b9a611028fe" +dependencies = [ + "bitflags", + "hashbrown 0.15.5", + "indexmap", + "semver", +] + +[[package]] +name = "winapi" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c839a674fcd7a98952e593242ea400abe93992746761e38641405d28b00f419" +dependencies = [ + "winapi-i686-pc-windows-gnu", + "winapi-x86_64-pc-windows-gnu", +] + +[[package]] +name = "winapi-i686-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6" + +[[package]] +name = "winapi-x86_64-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f" + +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + +[[package]] +name = "windows-sys" +version = "0.52.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" +dependencies = [ + "windows-targets", +] + +[[package]] +name = "windows-sys" +version = "0.59.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e38bc4d79ed67fd075bcc251a1c39b32a1776bbe92e5bef1f0bf1f8c531853b" +dependencies = [ + "windows-targets", +] + +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_gnullvm", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + +[[package]] +name = "winnow" +version = "0.7.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945" + +[[package]] +name = "wit-bindgen" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7249219f66ced02969388cf2bb044a09756a083d0fab1e566056b04d9fbcaa5" +dependencies = [ + "wit-bindgen-rust-macro", +] + +[[package]] +name = "wit-bindgen-core" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ea61de684c3ea68cb082b7a88508a8b27fcc8b797d738bfc99a82facf1d752dc" +dependencies = [ + "anyhow", + "heck", + "wit-parser", +] + +[[package]] +name = "wit-bindgen-rust" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7c566e0f4b284dd6561c786d9cb0142da491f46a9fbed79ea69cdad5db17f21" +dependencies = [ + "anyhow", + "heck", + "indexmap", + "prettyplease", + "syn", + "wasm-metadata", + "wit-bindgen-core", + "wit-component", +] + +[[package]] +name = "wit-bindgen-rust-macro" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c0f9bfd77e6a48eccf51359e3ae77140a7f50b1e2ebfe62422d8afdaffab17a" +dependencies = [ + "anyhow", + "prettyplease", + "proc-macro2", + "quote", + "syn", + "wit-bindgen-core", + "wit-bindgen-rust", +] + +[[package]] +name = "wit-component" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d66ea20e9553b30172b5e831994e35fbde2d165325bec84fc43dbf6f4eb9cb2" +dependencies = [ + "anyhow", + "bitflags", + "indexmap", + "log", + "serde", + "serde_derive", + "serde_json", + "wasm-encoder", + "wasm-metadata", + "wasmparser", + "wit-parser", +] + +[[package]] +name = "wit-parser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ecc8ac4bc1dc3381b7f59c34f00b67e18f910c2c0f50015669dde7def656a736" +dependencies = [ + "anyhow", + "id-arena", + "indexmap", + "log", + "semver", + "serde", + "serde_derive", + "serde_json", + "unicode-xid", + "wasmparser", +] + +[[package]] +name = "writeable" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9edde0db4769d2dc68579893f2306b26c6ecfbe0ef499b013d731b7b9247e0b9" + +[[package]] +name = "xattr" +version = "1.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156" +dependencies = [ + "libc", + "rustix", +] + +[[package]] +name = "yoke" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72d6e5c6afb84d73944e5cedb052c4680d5657337201555f9f2a16b7406d4954" +dependencies = [ + "stable_deref_trait", + "yoke-derive", + "zerofrom", +] + +[[package]] +name = "yoke-derive" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b659052874eb698efe5b9e8cf382204678a0086ebf46982b79d6ca3182927e5d" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + +[[package]] +name = "zerofrom" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "50cc42e0333e05660c3587f3bf9d0478688e15d870fab3346451ce7f8c9fbea5" +dependencies = [ + "zerofrom-derive", +] + +[[package]] +name = "zerofrom-derive" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d71e5d6e06ab090c67b5e44993ec16b72dcbaabc526db883a360057678b48502" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + +[[package]] +name = "zerotrie" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2a59c17a5562d507e4b54960e8569ebee33bee890c70aa3fe7b97e85a9fd7851" +dependencies = [ + "displaydoc", + "yoke", + "zerofrom", +] + +[[package]] +name = "zerovec" +version = "0.11.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c28719294829477f525be0186d13efa9a3c602f7ec202ca9e353d310fb9a002" +dependencies = [ + "yoke", + "zerofrom", + "zerovec-derive", +] + +[[package]] +name = "zerovec-derive" +version = "0.11.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eadce39539ca5cb3985590102671f2567e659fca9666581ad3411d59207951f3" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "zmij" +version = "1.0.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" diff --git a/tools/argument-comment-lint/Cargo.toml b/tools/argument-comment-lint/Cargo.toml new file mode 100644 index 0000000000000000000000000000000000000000..db02b392f3a5e3f584c0a7e882ae76499af2ab12 --- /dev/null +++ b/tools/argument-comment-lint/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "argument_comment_lint" +version = "0.1.0" +description = "Dylint lints for Rust /*param*/ argument comments" +edition = "2024" +publish = false + +[lib] +crate-type = ["cdylib"] + +[features] +bazel_native = [] + +[dependencies] +clippy_utils = { git = "https://github.com/rust-lang/rust-clippy", rev = "20ce69b9a63bcd2756cd906fe0964d1e901e042a" } +dylint_linting = "5.0.0" + +[dev-dependencies] +dylint_testing = "5.0.0" + +[workspace] + +[package.metadata.rust-analyzer] +rustc_private = true diff --git a/tools/argument-comment-lint/README.md b/tools/argument-comment-lint/README.md new file mode 100644 index 0000000000000000000000000000000000000000..889bae1df231358b342cd8f2fba127b9fe9cdb47 --- /dev/null +++ b/tools/argument-comment-lint/README.md @@ -0,0 +1,175 @@ +# argument-comment-lint + +Isolated [Dylint](https://github.com/trailofbits/dylint) library for enforcing +Rust argument comments in the exact `/*param*/` shape. + +Prefer self-documenting APIs over comment-heavy call sites when possible. If a +call site would otherwise read like `foo(false)` or `bar(None)`, consider an +enum, named helper, newtype, or another idiomatic Rust API shape first, and +use an argument comment only when a smaller compatibility-preserving change is +more appropriate. + +It provides two lints: + +- `argument_comment_mismatch` (`warn` by default): validates that a present + `/*param*/` comment matches the resolved callee parameter name. +- `uncommented_anonymous_literal_argument` (`allow` by default): flags + anonymous literal-like arguments such as `None`, `true`, `false`, and numeric + literals when they do not have a preceding `/*param*/` comment. + +String and char literals are exempt because they are often already +self-descriptive at the callsite. + +The sole non-self method argument is also exempt when the method name matches +the resolved parameter name. For example, `.enabled(false)` is already +self-descriptive when it resolves to `fn enabled(self, enabled: bool)`. An +explicit argument comment is still checked for a mismatch. + +## Behavior + +Given: + +```rust +fn create_openai_url(base_url: Option, retry_count: usize) -> String { + let _ = (base_url, retry_count); + String::new() +} +``` + +This is accepted: + +```rust +create_openai_url(/*base_url*/ None, /*retry_count*/ 3); +``` + +This is warned on by `argument_comment_mismatch`: + +```rust +create_openai_url(/*api_base*/ None, 3); +``` + +This is only warned on when `uncommented_anonymous_literal_argument` is enabled: + +```rust +create_openai_url(None, 3); +``` + +## Development + +Install the required tooling once: + +```bash +cargo install --locked cargo-dylint dylint-link +rustup toolchain install nightly-2025-09-18 \ + --component llvm-tools-preview \ + --component rustc-dev \ + --component rust-src +``` + +Run the lint crate tests: + +```bash +cd tools/argument-comment-lint +cargo test +``` + +GitHub releases also publish a DotSlash file named +`argument-comment-lint` for macOS arm64, Linux arm64, Linux x64, and Windows +x64. The published package contains a small runner executable, a bundled +`cargo-dylint`, and the prebuilt lint library. + +The package is not a full Rust toolchain. Running the prebuilt path still +requires the pinned nightly toolchain to be installed via `rustup`: + +```bash +rustup toolchain install nightly-2025-09-18 \ + --component llvm-tools-preview \ + --component rustc-dev \ + --component rust-src +``` + +The checked-in DotSlash file lives at `tools/argument-comment-lint/argument-comment-lint`. +`run-prebuilt-linter.py` resolves that file via `dotslash` and is the path used by +targeted package runs such as `just argument-comment-lint -p codex-core`. +Repo-wide runs now go through a native Bazel aspect that invokes a custom +`rustc_driver` and reuses Bazel-managed Rust dependency metadata instead of +spawning `cargo dylint` once per crate. The source-build path remains available +in `run.py` for people iterating on the lint crate itself. + +The Unix archive layout is: + +```text +argument-comment-lint/ + bin/ + argument-comment-lint + cargo-dylint + lib/ + libargument_comment_lint@nightly-2025-09-18-.dylib|so +``` + +On Windows the same layout is published as a `.zip`, with `.exe` and `.dll` +filenames instead. + +DotSlash resolves the package entrypoint to `argument-comment-lint/bin/argument-comment-lint` +(or `.exe` on Windows). That runner finds the sibling bundled `cargo-dylint` +binary and the single packaged Dylint library under `lib/`, normalizes the +host-qualified nightly filename to the plain `nightly-2025-09-18` channel when +needed, and then invokes `cargo-dylint dylint --lib-path ` with +the repo's default `DYLINT_RUSTFLAGS` and `CARGO_INCREMENTAL=0` settings. + +The checked-in `run-prebuilt-linter.py` wrapper uses the fetched package +contents directly so the current checked-in alpha artifact works the same way. +It also makes sure the `rustup` shims stay ahead of any direct toolchain +`cargo` binary on `PATH`, and sets `RUSTUP_HOME` from `rustup show home` when +the environment does not already provide it. That extra `RUSTUP_HOME` export is +required for the current Windows Dylint driver path. + +If you are changing the lint crate itself, use the source-build wrapper: + +```bash +./tools/argument-comment-lint/run.py -p codex-core +``` + +Run the lint against `codex-rs` from the repo root: + +```bash +just argument-comment-lint +bazel build --config=argument-comment-lint -- \ + $(./tools/argument-comment-lint/list-bazel-targets.sh) +./tools/argument-comment-lint/run-prebuilt-linter.py -p codex-core +just argument-comment-lint -p codex-core +``` + +If no package selection is provided, `just argument-comment-lint` now defaults +to the Bazel aspect path over `//codex-rs/...`. The Python wrappers remain the +package-scoped escape hatch and still default the underlying Cargo invocation +to `--all-targets` unless you explicitly narrow the target set, so targeted +wrapper runs cover test-only call sites by default. The Bazel entrypoints use +`tools/argument-comment-lint/list-bazel-targets.sh` to add the internal +manual `*-unit-tests-bin` Rust targets explicitly, so inline `#[cfg(test)]` +call sites are covered without pulling in unrelated manual release targets. + +Repo runs also promote `argument_comment_mismatch` and +`uncommented_anonymous_literal_argument` to errors by default: + +```bash +./tools/argument-comment-lint/run-prebuilt-linter.py -p codex-core +``` + +The wrapper does that by setting `DYLINT_RUSTFLAGS`, and it leaves an explicit +existing setting alone. It also defaults `CARGO_INCREMENTAL=0` unless you have +already set it, because the current nightly Dylint flow can otherwise hit a +rustc incremental compilation ICE locally. To override that behavior for an ad +hoc run: + +```bash +DYLINT_RUSTFLAGS="-A argument-comment-mismatch -A uncommented-anonymous-literal-argument" \ +CARGO_INCREMENTAL=1 \ + ./tools/argument-comment-lint/run.py -p codex-core +``` + +To override an explicitly narrow target selection, or to be explicit in scripts: + +```bash +./tools/argument-comment-lint/run-prebuilt-linter.py -p codex-core -- --all-targets +``` diff --git a/tools/argument-comment-lint/argument-comment-lint b/tools/argument-comment-lint/argument-comment-lint new file mode 100644 index 0000000000000000000000000000000000000000..602117e3ce2af40792e31c1d50db872ad18e4160 --- /dev/null +++ b/tools/argument-comment-lint/argument-comment-lint @@ -0,0 +1,79 @@ +#!/usr/bin/env dotslash + +{ + "name": "argument-comment-lint", + "platforms": { + "macos-aarch64": { + "size": 3402747, + "hash": "blake3", + "digest": "a11669d2f184a2c6f226cedce1bf10d1ec478d53413c42fe80d17dd873fdb2d7", + "format": "tar.gz", + "path": "argument-comment-lint/bin/argument-comment-lint", + "providers": [ + { + "url": "https://github.com/openai/codex/releases/download/rust-v0.117.0-alpha.2/argument-comment-lint-aarch64-apple-darwin.tar.gz" + }, + { + "type": "github-release", + "repo": "https://github.com/openai/codex", + "tag": "rust-v0.117.0-alpha.2", + "name": "argument-comment-lint-aarch64-apple-darwin.tar.gz" + } + ] + }, + "linux-x86_64": { + "size": 3869711, + "hash": "blake3", + "digest": "1015f4ba07d57edc5ec79c8f6709ddc1516f64c903e909820437a4b89d8d853a", + "format": "tar.gz", + "path": "argument-comment-lint/bin/argument-comment-lint", + "providers": [ + { + "url": "https://github.com/openai/codex/releases/download/rust-v0.117.0-alpha.2/argument-comment-lint-x86_64-unknown-linux-gnu.tar.gz" + }, + { + "type": "github-release", + "repo": "https://github.com/openai/codex", + "tag": "rust-v0.117.0-alpha.2", + "name": "argument-comment-lint-x86_64-unknown-linux-gnu.tar.gz" + } + ] + }, + "linux-aarch64": { + "size": 3759446, + "hash": "blake3", + "digest": "91f2a31e6390ca728ad09ae1aa6b6f379c67d996efcc22956001df89f068af5b", + "format": "tar.gz", + "path": "argument-comment-lint/bin/argument-comment-lint", + "providers": [ + { + "url": "https://github.com/openai/codex/releases/download/rust-v0.117.0-alpha.2/argument-comment-lint-aarch64-unknown-linux-gnu.tar.gz" + }, + { + "type": "github-release", + "repo": "https://github.com/openai/codex", + "tag": "rust-v0.117.0-alpha.2", + "name": "argument-comment-lint-aarch64-unknown-linux-gnu.tar.gz" + } + ] + }, + "windows-x86_64": { + "size": 3244599, + "hash": "blake3", + "digest": "dc711c6d85b1cabbe52447dda3872deb20c2e64b155da8be0ecb207c7c391683", + "format": "zip", + "path": "argument-comment-lint/bin/argument-comment-lint.exe", + "providers": [ + { + "url": "https://github.com/openai/codex/releases/download/rust-v0.117.0-alpha.2/argument-comment-lint-x86_64-pc-windows-msvc.zip" + }, + { + "type": "github-release", + "repo": "https://github.com/openai/codex", + "tag": "rust-v0.117.0-alpha.2", + "name": "argument-comment-lint-x86_64-pc-windows-msvc.zip" + } + ] + } + } +} diff --git a/tools/argument-comment-lint/driver.rs b/tools/argument-comment-lint/driver.rs new file mode 100644 index 0000000000000000000000000000000000000000..b6417d45fdcbb43b0e858556d71eea1e007ed070 --- /dev/null +++ b/tools/argument-comment-lint/driver.rs @@ -0,0 +1,43 @@ +#![feature(rustc_private)] + +extern crate rustc_driver; +extern crate rustc_interface; + +use std::env; +use std::ffi::OsString; +use std::path::Path; + +fn main() { + let mut callbacks = Callbacks; + let args = rustc_args(env::args_os().skip(1).collect()); + rustc_driver::run_compiler(&args, &mut callbacks); +} + +struct Callbacks; + +impl rustc_driver::Callbacks for Callbacks { + fn config(&mut self, config: &mut rustc_interface::Config) { + let previous = config.register_lints.take(); + config.register_lints = Some(Box::new(move |sess, lint_store| { + if let Some(previous) = &previous { + previous(sess, lint_store); + } + argument_comment_lint::register_lints(sess, lint_store); + })); + } +} + +fn rustc_args(args: Vec) -> Vec { + let mut rustc_args: Vec = args + .into_iter() + .map(|arg| arg.to_string_lossy().into_owned()) + .collect(); + if rustc_args.first().is_none_or(|arg| !is_rustc(arg)) { + rustc_args.insert(0, "rustc".to_string()); + } + rustc_args +} + +fn is_rustc(arg: &str) -> bool { + Path::new(arg).file_stem().and_then(|stem| stem.to_str()) == Some("rustc") +} diff --git a/tools/argument-comment-lint/lint_aspect.bzl b/tools/argument-comment-lint/lint_aspect.bzl new file mode 100644 index 0000000000000000000000000000000000000000..c20a79caa2a710441cccb4e2a884224983404c13 --- /dev/null +++ b/tools/argument-comment-lint/lint_aspect.bzl @@ -0,0 +1,188 @@ +"""Bazel aspect for running argument-comment-lint on Rust targets.""" + +load("@rules_rust//rust:defs.bzl", "rust_common") +load("@rules_rust//rust/private:rust.bzl", "RUSTC_ATTRS") +load( + "@rules_rust//rust/private:rustc.bzl", + "collect_deps", + "collect_inputs", + "construct_arguments", +) +load( + "@rules_rust//rust/private:utils.bzl", + "determine_output_hash", + "find_cc_toolchain", + "find_toolchain", +) + +_STRICT_LINT_FLAGS = [ + "-Dargument-comment-mismatch", + "-Duncommented-anonymous-literal-argument", + "-Aunknown-lints", +] + +def _find_rustc_driver_library(toolchain): + for file in toolchain.rustc_lib.to_list(): + if file.basename.startswith("librustc_driver-") or file.basename.startswith("rustc_driver-"): + return file + return None + +def _prepend_runtime_path(env, key, path, separator): + previous = env.get(key) + env[key] = "{}{}{}".format(path, separator, previous) if previous else path + +def _set_driver_runtime_env(env, toolchain): + driver_library = _find_rustc_driver_library(toolchain) + if not driver_library: + return + + library_dir = driver_library.dirname + if driver_library.basename.endswith(".dll"): + _prepend_runtime_path(env, "PATH", library_dir, ";") + return + + # The lint driver runs in exec configuration. Under remote execution the + # exec OS can differ from the Rust target OS, so populate both Unix loader + # variables from the located driver library instead of keying off target_os. + _prepend_runtime_path(env, "LD_LIBRARY_PATH", library_dir, ":") + _prepend_runtime_path(env, "DYLD_LIBRARY_PATH", library_dir, ":") + +def _get_argument_comment_lint_ready_crate_info(target, aspect_ctx): + if target.label.workspace_root.startswith("external"): + return None + + if aspect_ctx: + ignore_tags = [ + "no_argument_comment_lint", + "no-lint", + "no_lint", + "nolint", + ] + for tag in aspect_ctx.rule.attr.tags: + if tag.replace("-", "_").lower() in ignore_tags: + return None + + if rust_common.crate_info in target: + return target[rust_common.crate_info] + if rust_common.test_crate_info in target: + return target[rust_common.test_crate_info].crate + return None + +def _rust_argument_comment_lint_aspect_impl(target, ctx): + if OutputGroupInfo in target and hasattr(target[OutputGroupInfo], "argument_comment_lint_checks"): + return [] + + crate_info = _get_argument_comment_lint_ready_crate_info(target, ctx) + if not crate_info: + return [] + + toolchain = find_toolchain(ctx) + cc_toolchain, feature_configuration = find_cc_toolchain(ctx) + + dep_info, build_info, _ = collect_deps( + deps = crate_info.deps.to_list(), + proc_macro_deps = crate_info.proc_macro_deps.to_list(), + aliases = crate_info.aliases, + ) + + compile_inputs, out_dir, build_env_files, build_flags_files, linkstamp_outs, ambiguous_libs = collect_inputs( + ctx, + ctx.rule.file, + ctx.rule.files, + depset([]), + toolchain, + cc_toolchain, + feature_configuration, + crate_info, + dep_info, + build_info, + [], + ) + + success_marker = ctx.actions.declare_file( + ctx.label.name + ".argument_comment_lint.ok", + sibling = crate_info.output, + ) + + args, env = construct_arguments( + ctx = ctx, + attr = ctx.rule.attr, + file = ctx.file, + toolchain = toolchain, + tool_file = ctx.executable._driver, + cc_toolchain = cc_toolchain, + feature_configuration = feature_configuration, + crate_info = crate_info, + dep_info = dep_info, + linkstamp_outs = linkstamp_outs, + ambiguous_libs = ambiguous_libs, + output_hash = determine_output_hash(crate_info.root, ctx.label), + rust_flags = [], + out_dir = out_dir, + build_env_files = build_env_files, + build_flags_files = build_flags_files, + emit = ["dep-info", "metadata"], + skip_expanding_rustc_env = True, + ) + + if crate_info.is_test: + args.rustc_flags.add("--test") + + args.process_wrapper_flags.add("--touch-file", success_marker) + args.rustc_flags.add_all(_STRICT_LINT_FLAGS) + _set_driver_runtime_env(env, toolchain) + + driver_runfiles = ctx.attr._driver[DefaultInfo].default_runfiles.files + action_inputs = depset( + transitive = [ + compile_inputs, + driver_runfiles, + toolchain.rustc_lib, + ], + ) + + ctx.actions.run( + executable = toolchain.process_wrapper, + inputs = action_inputs, + outputs = [success_marker], + env = env, + tools = [ctx.executable._driver], + execution_requirements = { + "no-sandbox": "1", + }, + arguments = args.all, + mnemonic = "ArgumentCommentLint", + progress_message = "ArgumentCommentLint %{label}", + toolchain = "@rules_rust//rust:toolchain_type", + ) + + return [OutputGroupInfo(argument_comment_lint_checks = depset([success_marker]))] + +rust_argument_comment_lint_aspect = aspect( + implementation = _rust_argument_comment_lint_aspect_impl, + fragments = ["cpp"], + attrs = { + "_driver": attr.label( + default = Label("//tools/argument-comment-lint:argument-comment-lint-driver"), + executable = True, + cfg = "exec", + ), + } | RUSTC_ATTRS, + toolchains = [ + str(Label("@rules_rust//rust:toolchain_type")), + config_common.toolchain_type("@bazel_tools//tools/cpp:toolchain_type", mandatory = False), + ], + required_providers = [ + [rust_common.crate_info], + [rust_common.test_crate_info], + ], + doc = """\ +Runs argument-comment-lint on Rust targets using Bazel's Rust dependency graph. + +Example: + +```output +$ bazel build --config=argument-comment-lint //codex-rs/... +``` +""", +) diff --git a/tools/argument-comment-lint/list-bazel-targets.sh b/tools/argument-comment-lint/list-bazel-targets.sh new file mode 100644 index 0000000000000000000000000000000000000000..9418f59fbd875da6659210c3b856fa74057e556b --- /dev/null +++ b/tools/argument-comment-lint/list-bazel-targets.sh @@ -0,0 +1,36 @@ +#!/usr/bin/env bash + +set -euo pipefail + +repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)" +cd "${repo_root}" + +# Bazel wildcard builds skip manual targets, which misses the internal +# `*-unit-tests-bin` rust_test targets generated by `codex_rust_crate()`. +# Add only those manual rust_test targets explicitly so inline `#[cfg(test)]` +# call sites are linted without pulling in unrelated manual release targets. +manual_rust_test_targets="$( + ./.github/scripts/run-bazel-query-ci.sh \ + --output=label \ + -- 'kind("rust_test rule", attr(tags, "manual", //codex-rs/...))' +)" +if [[ "${RUNNER_OS:-}" != "Windows" ]]; then + manual_rust_test_targets="$(printf '%s\n' "${manual_rust_test_targets}" | grep -v -- '-windows-cross-bin$' || true)" +fi + +# Convert semantic lint opt-outs into negative target patterns so wildcard +# builds do not analyze toolchains used only by those wrappers. +excluded_targets="$( + ./.github/scripts/run-bazel-query-ci.sh \ + --output=label \ + -- 'attr(tags, "no-argument-comment-lint", //codex-rs/...)' +)" + +# The lint configuration does not register the transitioned Windows toolchain. +printf '%s\n' \ + "//codex-rs/..." \ + "-//codex-rs/core/tests/remote_env_windows:smoke-test" +if [[ -n "${excluded_targets}" ]]; then + printf '%s\n' "${excluded_targets}" | sed 's/^/-/' +fi +printf '%s\n' "${manual_rust_test_targets}" diff --git a/tools/argument-comment-lint/run-prebuilt-linter.py b/tools/argument-comment-lint/run-prebuilt-linter.py new file mode 100644 index 0000000000000000000000000000000000000000..31a5c226f361e06b564622fc131d8b970e84fd15 --- /dev/null +++ b/tools/argument-comment-lint/run-prebuilt-linter.py @@ -0,0 +1,48 @@ +#!/usr/bin/env python3 + +from __future__ import annotations + +import os +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +from wrapper_common import ( + build_final_args, + exec_command, + fetch_packaged_entrypoint, + find_packaged_cargo_dylint, + normalize_packaged_library, + parse_wrapper_args, + prefer_rustup_shims, + repo_root, + set_default_lint_env, +) + + +def main() -> "Never": + root = repo_root() + parsed = parse_wrapper_args(sys.argv[1:]) + final_args = build_final_args(parsed, root / "codex-rs" / "Cargo.toml") + + env = os.environ.copy() + prefer_rustup_shims(env) + set_default_lint_env(env) + + package_entrypoint = fetch_packaged_entrypoint( + root / "tools" / "argument-comment-lint" / "argument-comment-lint", + env, + ) + cargo_dylint = find_packaged_cargo_dylint(package_entrypoint) + library_path = normalize_packaged_library(package_entrypoint) + + command = [str(cargo_dylint), "dylint", "--lib-path", str(library_path)] + if not parsed.has_library_selection: + command.append("--all") + command.extend(final_args) + exec_command(command, env) + + +if __name__ == "__main__": + main() diff --git a/tools/argument-comment-lint/run.py b/tools/argument-comment-lint/run.py new file mode 100644 index 0000000000000000000000000000000000000000..49d0417e472ff669e58aca8ec51d98f0bffad9d8 --- /dev/null +++ b/tools/argument-comment-lint/run.py @@ -0,0 +1,43 @@ +#!/usr/bin/env python3 + +from __future__ import annotations + +import os +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +from wrapper_common import ( + build_final_args, + ensure_source_prerequisites, + exec_command, + parse_wrapper_args, + repo_root, + set_default_lint_env, +) + + +def main() -> "Never": + root = repo_root() + parsed = parse_wrapper_args(sys.argv[1:]) + final_args = build_final_args(parsed, root / "codex-rs" / "Cargo.toml") + + env = os.environ.copy() + ensure_source_prerequisites(env) + set_default_lint_env(env) + + command = [ + "cargo", + "dylint", + "--path", + str(root / "tools" / "argument-comment-lint"), + ] + if not parsed.has_library_selection: + command.append("--all") + command.extend(final_args) + exec_command(command, env) + + +if __name__ == "__main__": + main() diff --git a/tools/argument-comment-lint/rust-toolchain b/tools/argument-comment-lint/rust-toolchain new file mode 100644 index 0000000000000000000000000000000000000000..d159253fc0367b25914cb8bbfcaa95f3c5196336 --- /dev/null +++ b/tools/argument-comment-lint/rust-toolchain @@ -0,0 +1,3 @@ +[toolchain] +channel = "nightly-2025-09-18" +components = ["llvm-tools-preview", "rustc-dev", "rust-src"] diff --git a/tools/argument-comment-lint/src/bin/argument-comment-lint.rs b/tools/argument-comment-lint/src/bin/argument-comment-lint.rs new file mode 100644 index 0000000000000000000000000000000000000000..9dfe83b0fbb5330207d97811effb7e53f56666e5 --- /dev/null +++ b/tools/argument-comment-lint/src/bin/argument-comment-lint.rs @@ -0,0 +1,301 @@ +use std::env; +use std::ffi::OsString; +use std::fs; +use std::path::Path; +use std::path::PathBuf; +use std::process::Command; +use std::process::ExitCode; +use std::time::SystemTime; +use std::time::UNIX_EPOCH; + +const STRICT_LINTS: [&str; 2] = [ + "argument-comment-mismatch", + "uncommented-anonymous-literal-argument", +]; + +fn main() -> ExitCode { + match run() { + Ok(code) => code, + Err(err) => { + eprintln!("{err}"); + ExitCode::from(1) + } + } +} + +fn run() -> Result { + let exe_path = + env::current_exe().map_err(|err| format!("failed to locate current executable: {err}"))?; + let bin_dir = exe_path.parent().ok_or_else(|| { + format!( + "failed to locate parent directory for executable {}", + exe_path.display() + ) + })?; + let package_root = bin_dir.parent().ok_or_else(|| { + format!( + "failed to locate package root for executable {}", + exe_path.display() + ) + })?; + let cargo_dylint = bin_dir.join(cargo_dylint_binary_name()); + let library_dir = package_root.join("lib"); + let library_path = prepare_library_path_for_dylint(&find_bundled_library(&library_dir)?)?; + + ensure_exists(&cargo_dylint, "bundled cargo-dylint executable")?; + ensure_exists( + &library_dir, + "bundled argument-comment lint library directory", + )?; + + let args: Vec = env::args_os().skip(1).collect(); + let mut command = Command::new(&cargo_dylint); + command.arg("dylint"); + command.arg("--lib-path").arg(&library_path); + if !has_library_selection(&args) { + command.arg("--all"); + } + command.args(&args); + set_default_env(&mut command)?; + + let status = command + .status() + .map_err(|err| format!("failed to execute {}: {err}", cargo_dylint.display()))?; + Ok(exit_code_from_status(status.code())) +} + +fn has_library_selection(args: &[OsString]) -> bool { + let mut expect_value = false; + for arg in args { + if expect_value { + return true; + } + + match arg.to_string_lossy().as_ref() { + "--" => break, + "--lib" | "--lib-path" => { + expect_value = true; + } + "--lib=" | "--lib-path=" => return true, + value if value.starts_with("--lib=") || value.starts_with("--lib-path=") => { + return true; + } + _ => {} + } + } + + false +} + +fn set_default_env(command: &mut Command) -> Result<(), String> { + if let Some(flags) = env::var_os("DYLINT_RUSTFLAGS") { + let mut flags = flags.to_string_lossy().to_string(); + for strict_lint in STRICT_LINTS { + append_flag_if_missing(&mut flags, &format!("-D {strict_lint}")); + } + append_flag_if_missing(&mut flags, "-A unknown_lints"); + command.env("DYLINT_RUSTFLAGS", flags); + } else { + command.env("DYLINT_RUSTFLAGS", strict_rustflags()); + } + + if env::var_os("CARGO_INCREMENTAL").is_none() { + command.env("CARGO_INCREMENTAL", "0"); + } + + if env::var_os("RUSTUP_HOME").is_none() + && let Some(rustup_home) = infer_rustup_home()? + { + command.env("RUSTUP_HOME", rustup_home); + } + + Ok(()) +} + +fn strict_rustflags() -> String { + let strict_flags = STRICT_LINTS + .iter() + .map(|lint| format!("-D {lint}")) + .collect::>() + .join(" "); + format!("{strict_flags} -A unknown_lints") +} + +fn append_flag_if_missing(flags: &mut String, flag: &str) { + if flags.contains(flag) { + return; + } + + if !flags.is_empty() { + flags.push(' '); + } + flags.push_str(flag); +} + +fn cargo_dylint_binary_name() -> &'static str { + if cfg!(windows) { + "cargo-dylint.exe" + } else { + "cargo-dylint" + } +} + +fn infer_rustup_home() -> Result, String> { + let output = Command::new("rustup") + .args(["show", "home"]) + .output() + .map_err(|err| format!("failed to query rustup home via `rustup show home`: {err}"))?; + if !output.status.success() { + return Err(format!( + "`rustup show home` failed: {}", + String::from_utf8_lossy(&output.stderr).trim() + )); + } + + let home = String::from_utf8(output.stdout) + .map_err(|err| format!("`rustup show home` returned invalid UTF-8: {err}"))?; + let home = home.trim(); + if home.is_empty() { + Ok(None) + } else { + Ok(Some(OsString::from(home))) + } +} + +fn ensure_exists(path: &Path, label: &str) -> Result<(), String> { + if path.exists() { + Ok(()) + } else { + Err(format!("{label} not found at {}", path.display())) + } +} + +fn find_bundled_library(library_dir: &Path) -> Result { + let entries = fs::read_dir(library_dir).map_err(|err| { + format!( + "failed to read bundled library directory {}: {err}", + library_dir.display() + ) + })?; + let mut candidates = entries + .filter_map(Result::ok) + .map(|entry| entry.path()) + .filter(|path| path.is_file()) + .filter(|path| { + path.file_name() + .map(|name| name.to_string_lossy().contains('@')) + .unwrap_or(false) + }); + + let Some(first) = candidates.next() else { + return Err(format!( + "no packaged Dylint library found in {}", + library_dir.display() + )); + }; + if candidates.next().is_some() { + return Err(format!( + "expected exactly one packaged Dylint library in {}", + library_dir.display() + )); + } + + Ok(first) +} + +fn prepare_library_path_for_dylint(library_path: &Path) -> Result { + let Some(normalized_filename) = normalize_nightly_library_filename(library_path) else { + return Ok(library_path.to_path_buf()); + }; + + let temp_dir = env::temp_dir().join(format!( + "argument-comment-lint-{}-{}", + std::process::id(), + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map_err(|err| format!("failed to compute timestamp for temp dir: {err}"))? + .as_nanos() + )); + fs::create_dir_all(&temp_dir).map_err(|err| { + format!( + "failed to create temporary directory {}: {err}", + temp_dir.display() + ) + })?; + let normalized_path = temp_dir.join(normalized_filename); + fs::copy(library_path, &normalized_path).map_err(|err| { + format!( + "failed to copy packaged library {} to {}: {err}", + library_path.display(), + normalized_path.display() + ) + })?; + Ok(normalized_path) +} + +fn normalize_nightly_library_filename(library_path: &Path) -> Option { + let stem = library_path.file_stem()?.to_string_lossy(); + let extension = library_path.extension()?.to_string_lossy(); + let (lib_name, toolchain) = stem.rsplit_once('@')?; + let normalized_toolchain = normalize_nightly_toolchain(toolchain)?; + Some(format!("{lib_name}@{normalized_toolchain}.{extension}")) +} + +fn normalize_nightly_toolchain(toolchain: &str) -> Option { + let parts: Vec<_> = toolchain.split('-').collect(); + if parts.len() > 4 + && parts[0] == "nightly" + && parts[1].len() == 4 + && parts[2].len() == 2 + && parts[3].len() == 2 + && parts[1..4] + .iter() + .all(|part| part.chars().all(|ch| ch.is_ascii_digit())) + { + Some(format!("nightly-{}-{}-{}", parts[1], parts[2], parts[3])) + } else { + None + } +} + +fn exit_code_from_status(code: Option) -> ExitCode { + code.and_then(|value| u8::try_from(value).ok()) + .map_or_else(|| ExitCode::from(1), ExitCode::from) +} + +#[cfg(test)] +mod tests { + use super::normalize_nightly_library_filename; + use super::strict_rustflags; + use std::path::Path; + + #[test] + fn strips_host_triple_from_nightly_filename() { + assert_eq!( + normalize_nightly_library_filename(Path::new( + "libargument_comment_lint@nightly-2025-09-18-aarch64-apple-darwin.dylib" + )), + Some(String::from( + "libargument_comment_lint@nightly-2025-09-18.dylib" + )) + ); + } + + #[test] + fn leaves_unqualified_nightly_filename_alone() { + assert_eq!( + normalize_nightly_library_filename(Path::new( + "libargument_comment_lint@nightly-2025-09-18.dylib" + )), + None + ); + } + + #[test] + fn strict_rustflags_promotes_both_enforced_lints() { + assert_eq!( + strict_rustflags(), + "-D argument-comment-mismatch -D uncommented-anonymous-literal-argument -A unknown_lints" + ); + } +} diff --git a/tools/argument-comment-lint/src/comment_parser.rs b/tools/argument-comment-lint/src/comment_parser.rs new file mode 100644 index 0000000000000000000000000000000000000000..7ea66c195b7b4d50a4fb2c4ec7d73c73ea09e683 --- /dev/null +++ b/tools/argument-comment-lint/src/comment_parser.rs @@ -0,0 +1,63 @@ +pub fn parse_argument_comment(text: &str) -> Option<&str> { + let trimmed = text.trim_end(); + let comment_start = trimmed.rfind("/*")?; + let comment = &trimmed[comment_start..]; + let name = comment.strip_prefix("/*")?.strip_suffix("*/")?; + is_identifier(name).then_some(name) +} + +pub fn parse_argument_comment_prefix(text: &str) -> Option<&str> { + let trimmed = text.trim_start(); + let comment = trimmed.strip_prefix("/*")?; + let (name, _) = comment.split_once("*/")?; + is_identifier(name).then_some(name) +} + +fn is_identifier(text: &str) -> bool { + let mut chars = text.chars(); + let Some(first) = chars.next() else { + return false; + }; + if !(first == '_' || first.is_ascii_alphabetic()) { + return false; + } + chars.all(|ch| ch == '_' || ch.is_ascii_alphanumeric()) +} + +#[cfg(test)] +mod tests { + use super::parse_argument_comment; + use super::parse_argument_comment_prefix; + + #[test] + fn parses_trailing_comment() { + assert_eq!(parse_argument_comment("(/*base_url*/ "), Some("base_url")); + assert_eq!( + parse_argument_comment(", /*timeout_ms*/ "), + Some("timeout_ms") + ); + assert_eq!( + parse_argument_comment(".method::(/*base_url*/ "), + Some("base_url") + ); + } + + #[test] + fn rejects_non_matching_shapes() { + assert_eq!(parse_argument_comment("(\n"), None); + assert_eq!(parse_argument_comment("(/* base_url*/ "), None); + assert_eq!(parse_argument_comment("(/*base_url */ "), None); + assert_eq!(parse_argument_comment("(/*base_url=*/ "), None); + assert_eq!(parse_argument_comment("(/*1base_url*/ "), None); + assert_eq!(parse_argument_comment_prefix("/*env=*/ None"), None); + } + + #[test] + fn parses_prefix_comment() { + assert_eq!(parse_argument_comment_prefix("/*env*/ None"), Some("env")); + assert_eq!( + parse_argument_comment_prefix("\n /*retry_count*/ 3"), + Some("retry_count") + ); + } +} diff --git a/tools/argument-comment-lint/src/lib.rs b/tools/argument-comment-lint/src/lib.rs new file mode 100644 index 0000000000000000000000000000000000000000..d14d39b3a50c331fcf97474de7aa62bca3740c7d --- /dev/null +++ b/tools/argument-comment-lint/src/lib.rs @@ -0,0 +1,296 @@ +#![feature(rustc_private)] + +mod comment_parser; + +extern crate rustc_ast; +extern crate rustc_errors; +extern crate rustc_hir; +extern crate rustc_lint; +extern crate rustc_middle; +extern crate rustc_session; +extern crate rustc_span; + +use clippy_utils::diagnostics::span_lint_and_help; +use clippy_utils::diagnostics::span_lint_and_sugg; +use clippy_utils::fn_def_id; +use clippy_utils::is_res_lang_ctor; +use clippy_utils::peel_blocks; +use clippy_utils::source::snippet; +use rustc_ast::LitKind; +use rustc_errors::Applicability; +use rustc_hir::Expr; +use rustc_hir::ExprKind; +use rustc_hir::LangItem; +use rustc_hir::UnOp; +use rustc_hir::def::DefKind; +use rustc_lint::LateContext; +use rustc_lint::LateLintPass; +use rustc_span::BytePos; +use rustc_span::Span; + +use crate::comment_parser::parse_argument_comment; +use crate::comment_parser::parse_argument_comment_prefix; + +#[cfg(not(feature = "bazel_native"))] +dylint_linting::dylint_library!(); + +#[unsafe(no_mangle)] +pub fn register_lints(_sess: &rustc_session::Session, lint_store: &mut rustc_lint::LintStore) { + lint_store.register_lints(&[ + ARGUMENT_COMMENT_MISMATCH, + UNCOMMENTED_ANONYMOUS_LITERAL_ARGUMENT, + ]); + lint_store.register_late_pass(|_| Box::new(ArgumentCommentLint)); +} + +rustc_session::declare_lint! { + /// ### What it does + /// + /// Checks `/*param*/` argument comments and verifies that the comment + /// matches the resolved callee parameter name. + /// + /// ### Why is this bad? + /// + /// A mismatched comment is worse than no comment because it actively + /// misleads the reader. + /// + /// ### Known problems + /// + /// This lint only runs when the callee resolves to a concrete function or + /// method with available parameter names. + /// + /// ### Example + /// + /// ```rust + /// fn create_openai_url(base_url: Option) -> String { + /// String::new() + /// } + /// + /// create_openai_url(/*api_base*/ None); + /// ``` + /// + /// Use instead: + /// + /// ```rust + /// fn create_openai_url(base_url: Option) -> String { + /// String::new() + /// } + /// + /// create_openai_url(/*base_url*/ None); + /// ``` + pub ARGUMENT_COMMENT_MISMATCH, + Warn, + "argument comment does not match the resolved parameter name" +} + +rustc_session::declare_lint! { + /// ### What it does + /// + /// Requires a `/*param*/` comment before anonymous literal-like + /// arguments such as `None`, booleans, and numeric literals. + /// A method's sole non-self argument is exempt when its name matches the + /// method name. + /// + /// ### Why is this bad? + /// + /// Bare literal-like arguments make call sites harder to read because the + /// meaning of the value is hidden in the callee signature. + /// + /// ### Known problems + /// + /// This lint is opinionated, so it is `allow` by default. + /// + /// ### Example + /// + /// ```rust + /// fn create_openai_url(base_url: Option) -> String { + /// String::new() + /// } + /// + /// create_openai_url(None); + /// ``` + /// + /// Use instead: + /// + /// ```rust + /// fn create_openai_url(base_url: Option) -> String { + /// String::new() + /// } + /// + /// create_openai_url(/*base_url*/ None); + /// ``` + pub UNCOMMENTED_ANONYMOUS_LITERAL_ARGUMENT, + Allow, + "anonymous literal-like argument is missing a `/*param*/` comment" +} + +#[derive(Default)] +pub struct ArgumentCommentLint; + +enum CallKind { + Function, + Method { name: String }, +} + +rustc_session::impl_lint_pass!( + ArgumentCommentLint => [ARGUMENT_COMMENT_MISMATCH, UNCOMMENTED_ANONYMOUS_LITERAL_ARGUMENT] +); + +impl<'tcx> LateLintPass<'tcx> for ArgumentCommentLint { + fn check_expr(&mut self, cx: &LateContext<'tcx>, expr: &'tcx Expr<'tcx>) { + if expr.span.from_expansion() { + return; + } + + match expr.kind { + ExprKind::Call(callee, args) => { + self.check_call(cx, expr, callee.span, args, CallKind::Function); + } + ExprKind::MethodCall(method, receiver, args, _) => { + self.check_call( + cx, + expr, + receiver.span, + args, + CallKind::Method { + name: method.ident.name.to_string(), + }, + ); + } + _ => {} + } + } +} + +impl ArgumentCommentLint { + fn check_call<'tcx>( + &self, + cx: &LateContext<'tcx>, + call: &'tcx Expr<'tcx>, + first_gap_anchor: Span, + args: &'tcx [Expr<'tcx>], + call_kind: CallKind, + ) { + let Some(def_id) = fn_def_id(cx, call) else { + return; + }; + if !def_id.is_local() && !is_workspace_crate_name(cx.tcx.crate_name(def_id.krate).as_str()) + { + return; + } + if !matches!(cx.tcx.def_kind(def_id), DefKind::Fn | DefKind::AssocFn) { + return; + } + + // Method parameter lists include `self`, which is not present in `args`. + let (parameter_offset, method_name) = match &call_kind { + CallKind::Function => (0, None), + CallKind::Method { name } => (1, Some(name.as_str())), + }; + let parameter_names: Vec<_> = cx.tcx.fn_arg_idents(def_id).iter().copied().collect(); + for (index, arg) in args.iter().enumerate() { + if arg.span.from_expansion() { + continue; + } + + let Some(expected_name) = parameter_names.get(index + parameter_offset) else { + continue; + }; + let Some(expected_name) = expected_name else { + continue; + }; + let expected_name = expected_name.name.to_string(); + if !is_meaningful_parameter_name(&expected_name) { + continue; + } + + let boundary_span = if index == 0 { + first_gap_anchor + } else { + args[index - 1].span + }; + let gap_span = boundary_span.between(arg.span); + let gap_text = snippet(cx, gap_span, ""); + let arg_text = snippet(cx, arg.span, ".."); + let lookbehind_start = BytePos(arg.span.lo().0.saturating_sub(64)); + let lookbehind_text = + snippet(cx, arg.span.shrink_to_lo().with_lo(lookbehind_start), ""); + let argument_comment = parse_argument_comment(gap_text.as_ref()) + .or_else(|| parse_argument_comment(lookbehind_text.as_ref())) + .or_else(|| parse_argument_comment_prefix(arg_text.as_ref())); + + if let Some(actual_name) = argument_comment { + if actual_name != expected_name { + span_lint_and_help( + cx, + ARGUMENT_COMMENT_MISMATCH, + arg.span, + format!( + "argument comment `/*{actual_name}*/` does not match parameter `{expected_name}`" + ), + None, + format!("use `/*{expected_name}*/`"), + ); + } + continue; + } + + // Don't require a clarifying comment for self-documenting arguments whose names + // match the method. + if args.len() == 1 && method_name == Some(expected_name.as_str()) { + continue; + } + + if !is_anonymous_literal_like(cx, arg) { + continue; + } + + span_lint_and_sugg( + cx, + UNCOMMENTED_ANONYMOUS_LITERAL_ARGUMENT, + arg.span, + format!("anonymous literal-like argument for parameter `{expected_name}`"), + "prepend the parameter name comment", + format!("/*{expected_name}*/ {arg_text}"), + Applicability::MachineApplicable, + ); + } + } +} + +fn is_anonymous_literal_like(cx: &LateContext<'_>, expr: &Expr<'_>) -> bool { + let expr = peel_blocks(expr); + match expr.kind { + ExprKind::Lit(lit) => !matches!( + lit.node, + LitKind::Str(..) | LitKind::ByteStr(..) | LitKind::CStr(..) | LitKind::Char(..) + ), + ExprKind::Unary(UnOp::Neg, inner) => matches!(peel_blocks(inner).kind, ExprKind::Lit(_)), + ExprKind::Path(qpath) => { + is_res_lang_ctor(cx, cx.qpath_res(&qpath, expr.hir_id), LangItem::OptionNone) + } + _ => false, + } +} + +fn is_meaningful_parameter_name(name: &str) -> bool { + !name.is_empty() && !name.starts_with('_') +} + +fn is_workspace_crate_name(name: &str) -> bool { + name.starts_with("codex_") || matches!(name, "app_test_support" | "core_test_support") +} + +#[test] +fn ui() { + dylint_testing::ui_test(env!("CARGO_PKG_NAME"), "ui"); +} + +#[test] +fn workspace_crate_filter_accepts_first_party_names_only() { + assert!(is_workspace_crate_name("codex_core")); + assert!(is_workspace_crate_name("codex_tui")); + assert!(is_workspace_crate_name("core_test_support")); + assert!(!is_workspace_crate_name("std")); + assert!(!is_workspace_crate_name("tokio")); +} diff --git a/tools/argument-comment-lint/test_wrapper_common.py b/tools/argument-comment-lint/test_wrapper_common.py new file mode 100644 index 0000000000000000000000000000000000000000..bab6e6f84d3ae937e002c60b505a1399d46448ce --- /dev/null +++ b/tools/argument-comment-lint/test_wrapper_common.py @@ -0,0 +1,136 @@ +#!/usr/bin/env python3 + +from __future__ import annotations + +from pathlib import Path +import sys +import unittest + +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +import wrapper_common + + +class WrapperCommonTest(unittest.TestCase): + def test_defaults_to_workspace_and_all_targets(self) -> None: + parsed = wrapper_common.parse_wrapper_args([]) + final_args = wrapper_common.build_final_args( + parsed, Path("/repo/codex-rs/Cargo.toml") + ) + + self.assertEqual( + final_args, + [ + "--manifest-path", + "/repo/codex-rs/Cargo.toml", + "--workspace", + "--no-deps", + "--", + "--all-targets", + ], + ) + + def test_forwarded_cargo_args_keep_single_separator(self) -> None: + parsed = wrapper_common.parse_wrapper_args( + ["-p", "codex-core", "--", "--tests"] + ) + final_args = wrapper_common.build_final_args( + parsed, Path("/repo/codex-rs/Cargo.toml") + ) + + self.assertEqual( + final_args, + [ + "--manifest-path", + "/repo/codex-rs/Cargo.toml", + "--no-deps", + "-p", + "codex-core", + "--", + "--tests", + ], + ) + + def test_fix_does_not_add_all_targets(self) -> None: + parsed = wrapper_common.parse_wrapper_args(["--fix", "-p", "codex-core"]) + final_args = wrapper_common.build_final_args( + parsed, Path("/repo/codex-rs/Cargo.toml") + ) + + self.assertEqual( + final_args, + [ + "--manifest-path", + "/repo/codex-rs/Cargo.toml", + "--no-deps", + "--fix", + "-p", + "codex-core", + ], + ) + + def test_explicit_manifest_and_workspace_are_preserved(self) -> None: + parsed = wrapper_common.parse_wrapper_args( + [ + "--manifest-path", + "/tmp/custom/Cargo.toml", + "--workspace", + "--no-deps", + "--", + "--bins", + ] + ) + final_args = wrapper_common.build_final_args( + parsed, Path("/repo/codex-rs/Cargo.toml") + ) + + self.assertEqual( + final_args, + [ + "--manifest-path", + "/tmp/custom/Cargo.toml", + "--workspace", + "--no-deps", + "--", + "--bins", + ], + ) + + def test_explicit_package_manifest_does_not_force_workspace(self) -> None: + parsed = wrapper_common.parse_wrapper_args( + [ + "--manifest-path", + "/tmp/custom/Cargo.toml", + ] + ) + final_args = wrapper_common.build_final_args( + parsed, Path("/repo/codex-rs/Cargo.toml") + ) + + self.assertEqual( + final_args, + [ + "--no-deps", + "--manifest-path", + "/tmp/custom/Cargo.toml", + "--", + "--all-targets", + ], + ) + + def test_default_lint_env_promotes_both_strict_lints(self) -> None: + env: dict[str, str] = {} + + wrapper_common.set_default_lint_env(env) + + self.assertEqual( + env["DYLINT_RUSTFLAGS"], + "-D argument-comment-mismatch " + "-D uncommented-anonymous-literal-argument " + "-A unknown_lints", + ) + self.assertEqual(env["CARGO_INCREMENTAL"], "0") + + +if __name__ == "__main__": + unittest.main() diff --git a/tools/argument-comment-lint/ui/allow_char_literals.rs b/tools/argument-comment-lint/ui/allow_char_literals.rs new file mode 100644 index 0000000000000000000000000000000000000000..85ef7d362f306918cb5b9f51941090ee740a2927 --- /dev/null +++ b/tools/argument-comment-lint/ui/allow_char_literals.rs @@ -0,0 +1,9 @@ +#![warn(uncommented_anonymous_literal_argument)] + +fn split_top_level(body: &str, delimiter: char) { + let _ = (body, delimiter); +} + +fn main() { + split_top_level("a|b|c", '|'); +} diff --git a/tools/argument-comment-lint/ui/allow_self_documenting_methods.rs b/tools/argument-comment-lint/ui/allow_self_documenting_methods.rs new file mode 100644 index 0000000000000000000000000000000000000000..d57f7c521da455f8067141294c6a8ace01e74f64 --- /dev/null +++ b/tools/argument-comment-lint/ui/allow_self_documenting_methods.rs @@ -0,0 +1,25 @@ +#![warn(argument_comment_mismatch)] +#![warn(uncommented_anonymous_literal_argument)] + +struct Builder; + +impl Builder { + fn enabled(self, enabled: bool) -> Self { + let _ = enabled; + self + } + + fn retry_count(self, retry_count: usize) -> Self { + let _ = retry_count; + self + } + + fn base_url(self, base_url: Option) -> Self { + let _ = base_url; + self + } +} + +fn main() { + let _ = Builder.enabled(false).retry_count(3).base_url(None); +} diff --git a/tools/argument-comment-lint/ui/allow_string_literals.rs b/tools/argument-comment-lint/ui/allow_string_literals.rs new file mode 100644 index 0000000000000000000000000000000000000000..9d61c4a761ed097dba3ce83aed250c3ded546f36 --- /dev/null +++ b/tools/argument-comment-lint/ui/allow_string_literals.rs @@ -0,0 +1,9 @@ +#![warn(uncommented_anonymous_literal_argument)] + +fn describe(prefix: &str, suffix: &str) { + let _ = (prefix, suffix); +} + +fn main() { + describe("openai", r"https://api.openai.com/v1"); +} diff --git a/tools/argument-comment-lint/ui/comment_matches.rs b/tools/argument-comment-lint/ui/comment_matches.rs new file mode 100644 index 0000000000000000000000000000000000000000..f582fb958f24eccfcd01d601f6559b325624c207 --- /dev/null +++ b/tools/argument-comment-lint/ui/comment_matches.rs @@ -0,0 +1,12 @@ +#![warn(argument_comment_mismatch)] + +fn create_openai_url(base_url: Option, retry_count: usize) -> String { + let _ = (base_url, retry_count); + String::new() +} + +fn main() { + let base_url = Some(String::from("https://api.openai.com")); + create_openai_url(base_url, 3); + create_openai_url(/*base_url*/ None, 3); +} diff --git a/tools/argument-comment-lint/ui/comment_matches_multiline.rs b/tools/argument-comment-lint/ui/comment_matches_multiline.rs new file mode 100644 index 0000000000000000000000000000000000000000..335fffc860053793bedae3517f46427d6c612f8c --- /dev/null +++ b/tools/argument-comment-lint/ui/comment_matches_multiline.rs @@ -0,0 +1,17 @@ +#![warn(argument_comment_mismatch)] +#![warn(uncommented_anonymous_literal_argument)] + +fn run_git_for_stdout(repo_root: &str, args: Vec<&str>, env: Option<&str>) -> String { + let _ = (repo_root, args, env); + String::new() +} + +// Keep the call split across lines so this fixture tests multiline arguments. +#[rustfmt::skip] +fn main() { + let _ = run_git_for_stdout( + "/tmp/repo", + vec!["rev-parse", "HEAD"], + /*env*/ None, + ); +} diff --git a/tools/argument-comment-lint/ui/comment_mismatch.rs b/tools/argument-comment-lint/ui/comment_mismatch.rs new file mode 100644 index 0000000000000000000000000000000000000000..f3ad4db75478f4db40c6ea331e1a06fe78c7229d --- /dev/null +++ b/tools/argument-comment-lint/ui/comment_mismatch.rs @@ -0,0 +1,20 @@ +#![warn(argument_comment_mismatch)] + +fn create_openai_url(base_url: Option) -> String { + let _ = base_url; + String::new() +} + +struct Options; + +impl Options { + fn enabled(self, enabled: bool) -> Self { + let _ = enabled; + self + } +} + +fn main() { + let _ = create_openai_url(/*api_base*/ None); + let _ = Options.enabled(/*value*/ false); +} diff --git a/tools/argument-comment-lint/ui/comment_mismatch.stderr b/tools/argument-comment-lint/ui/comment_mismatch.stderr new file mode 100644 index 0000000000000000000000000000000000000000..058871395a31fb3f28328876db409958ea28111d --- /dev/null +++ b/tools/argument-comment-lint/ui/comment_mismatch.stderr @@ -0,0 +1,23 @@ +warning: argument comment `/*api_base*/` does not match parameter `base_url` + --> $DIR/comment_mismatch.rs:18:44 + | +LL | let _ = create_openai_url(/*api_base*/ None); + | ^^^^ + | + = help: use `/*base_url*/` +note: the lint level is defined here + --> $DIR/comment_mismatch.rs:1:9 + | +LL | #![warn(argument_comment_mismatch)] + | ^^^^^^^^^^^^^^^^^^^^^^^^^ + +warning: argument comment `/*value*/` does not match parameter `enabled` + --> $DIR/comment_mismatch.rs:19:39 + | +LL | let _ = Options.enabled(/*value*/ false); + | ^^^^^ + | + = help: use `/*enabled*/` + +warning: 2 warnings emitted + diff --git a/tools/argument-comment-lint/ui/ignore_external_methods.rs b/tools/argument-comment-lint/ui/ignore_external_methods.rs new file mode 100644 index 0000000000000000000000000000000000000000..44c72707d069f0bf5e9776ae337c234b069b3b3b --- /dev/null +++ b/tools/argument-comment-lint/ui/ignore_external_methods.rs @@ -0,0 +1,9 @@ +#![warn(uncommented_anonymous_literal_argument)] + +fn main() { + let line = "{\"type\":\"response_item\"}"; + let _ = line.starts_with('{'); + let _ = line.find("type"); + let parts = ["type", "response_item"]; + let _ = parts.join("\n"); +} diff --git a/tools/argument-comment-lint/ui/multiple_method_arguments.rs b/tools/argument-comment-lint/ui/multiple_method_arguments.rs new file mode 100644 index 0000000000000000000000000000000000000000..e15f15624300bc09e4e1db165c51d68120cb34f3 --- /dev/null +++ b/tools/argument-comment-lint/ui/multiple_method_arguments.rs @@ -0,0 +1,14 @@ +#![warn(uncommented_anonymous_literal_argument)] + +struct Options; + +impl Options { + fn enabled(self, enabled: bool, retry_count: usize) -> Self { + let _ = (enabled, retry_count); + self + } +} + +fn main() { + let _ = Options.enabled(false, /*retry_count*/ 3); +} diff --git a/tools/argument-comment-lint/ui/multiple_method_arguments.stderr b/tools/argument-comment-lint/ui/multiple_method_arguments.stderr new file mode 100644 index 0000000000000000000000000000000000000000..f9b452e4fce44186dbcd21fd8e07a0d3d52c7a1f --- /dev/null +++ b/tools/argument-comment-lint/ui/multiple_method_arguments.stderr @@ -0,0 +1,14 @@ +warning: anonymous literal-like argument for parameter `enabled` + --> $DIR/multiple_method_arguments.rs:13:29 + | +LL | let _ = Options.enabled(false, /*retry_count*/ 3); + | ^^^^^ help: prepend the parameter name comment: `/*enabled*/ false` + | +note: the lint level is defined here + --> $DIR/multiple_method_arguments.rs:1:9 + | +LL | #![warn(uncommented_anonymous_literal_argument)] + | ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +warning: 1 warning emitted + diff --git a/tools/argument-comment-lint/ui/uncommented_literal.rs b/tools/argument-comment-lint/ui/uncommented_literal.rs new file mode 100644 index 0000000000000000000000000000000000000000..a64804c78094e74646e512cb2815d4592df8b495 --- /dev/null +++ b/tools/argument-comment-lint/ui/uncommented_literal.rs @@ -0,0 +1,18 @@ +#![warn(uncommented_anonymous_literal_argument)] + +struct Client; + +impl Client { + fn set_flag(&self, enabled: bool) {} +} + +fn create_openai_url(base_url: Option, retry_count: usize) -> String { + let _ = (base_url, retry_count); + String::new() +} + +fn main() { + let client = Client; + let _ = create_openai_url(None, 3); + client.set_flag(true); +} diff --git a/tools/argument-comment-lint/ui/uncommented_literal.stderr b/tools/argument-comment-lint/ui/uncommented_literal.stderr new file mode 100644 index 0000000000000000000000000000000000000000..a1060ef6937a7c920b185099c96f1b4f442788a2 --- /dev/null +++ b/tools/argument-comment-lint/ui/uncommented_literal.stderr @@ -0,0 +1,26 @@ +warning: anonymous literal-like argument for parameter `base_url` + --> $DIR/uncommented_literal.rs:16:31 + | +LL | let _ = create_openai_url(None, 3); + | ^^^^ help: prepend the parameter name comment: `/*base_url*/ None` + | +note: the lint level is defined here + --> $DIR/uncommented_literal.rs:1:9 + | +LL | #![warn(uncommented_anonymous_literal_argument)] + | ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +warning: anonymous literal-like argument for parameter `retry_count` + --> $DIR/uncommented_literal.rs:16:37 + | +LL | let _ = create_openai_url(None, 3); + | ^ help: prepend the parameter name comment: `/*retry_count*/ 3` + +warning: anonymous literal-like argument for parameter `enabled` + --> $DIR/uncommented_literal.rs:17:21 + | +LL | client.set_flag(true); + | ^^^^ help: prepend the parameter name comment: `/*enabled*/ true` + +warning: 3 warnings emitted + diff --git a/tools/argument-comment-lint/wrapper_common.py b/tools/argument-comment-lint/wrapper_common.py new file mode 100644 index 0000000000000000000000000000000000000000..382ba7bcf11833695268de797b273f546228ff3d --- /dev/null +++ b/tools/argument-comment-lint/wrapper_common.py @@ -0,0 +1,282 @@ +#!/usr/bin/env python3 + +from __future__ import annotations + +from dataclasses import dataclass +import os +from pathlib import Path +import re +import shlex +import shutil +import subprocess +import sys +import tempfile +from typing import MutableMapping, Sequence + +STRICT_LINTS = [ + "argument-comment-mismatch", + "uncommented-anonymous-literal-argument", +] +NOISE_LINT = "unknown_lints" +TOOLCHAIN_CHANNEL = "nightly-2025-09-18" + +_TARGET_SELECTION_ARGS = { + "--all-targets", + "--lib", + "--bins", + "--tests", + "--examples", + "--benches", + "--doc", +} +_TARGET_SELECTION_PREFIXES = ("--bin=", "--test=", "--example=", "--bench=") +_TARGET_SELECTION_WITH_VALUE = {"--bin", "--test", "--example", "--bench"} +_NIGHTLY_LIBRARY_PATTERN = re.compile(r"^(.+@nightly-[0-9]{4}-[0-9]{2}-[0-9]{2})-.+$") + + +@dataclass +class ParsedWrapperArgs: + lint_args: list[str] + cargo_args: list[str] + has_manifest_path: bool = False + has_package_selection: bool = False + has_no_deps: bool = False + has_library_selection: bool = False + has_cargo_target_selection: bool = False + has_fix: bool = False + + +def repo_root() -> Path: + return Path(__file__).resolve().parents[2] + + +def parse_wrapper_args(argv: Sequence[str]) -> ParsedWrapperArgs: + parsed = ParsedWrapperArgs(lint_args=[], cargo_args=[]) + after_separator = False + expect_value: str | None = None + + for arg in argv: + if after_separator: + parsed.cargo_args.append(arg) + if arg in _TARGET_SELECTION_ARGS or arg in _TARGET_SELECTION_WITH_VALUE: + parsed.has_cargo_target_selection = True + elif arg.startswith(_TARGET_SELECTION_PREFIXES): + parsed.has_cargo_target_selection = True + continue + + if arg == "--": + after_separator = True + continue + + parsed.lint_args.append(arg) + + if expect_value is not None: + if expect_value == "manifest_path": + parsed.has_manifest_path = True + elif expect_value == "package_selection": + parsed.has_package_selection = True + elif expect_value == "library_selection": + parsed.has_library_selection = True + expect_value = None + continue + + if arg == "--manifest-path": + expect_value = "manifest_path" + elif arg.startswith("--manifest-path="): + parsed.has_manifest_path = True + elif arg in {"-p", "--package"}: + expect_value = "package_selection" + elif arg.startswith("--package="): + parsed.has_package_selection = True + elif arg == "--fix": + parsed.has_fix = True + elif arg == "--workspace": + parsed.has_package_selection = True + elif arg == "--no-deps": + parsed.has_no_deps = True + elif arg in {"--lib", "--lib-path"}: + expect_value = "library_selection" + elif arg.startswith("--lib=") or arg.startswith("--lib-path="): + parsed.has_library_selection = True + + return parsed + + +def build_final_args(parsed: ParsedWrapperArgs, manifest_path: Path) -> list[str]: + final_args: list[str] = [] + cargo_args = list(parsed.cargo_args) + + if not parsed.has_manifest_path: + final_args.extend(["--manifest-path", str(manifest_path)]) + if not parsed.has_package_selection and not parsed.has_manifest_path: + final_args.append("--workspace") + if not parsed.has_no_deps: + final_args.append("--no-deps") + if not parsed.has_fix and not parsed.has_cargo_target_selection: + cargo_args.append("--all-targets") + final_args.extend(parsed.lint_args) + if cargo_args: + final_args.extend(["--", *cargo_args]) + return final_args + + +def append_env_flag(env: MutableMapping[str, str], key: str, flag: str) -> None: + value = env.get(key) + if value is None or value == "": + env[key] = flag + return + if flag not in value: + env[key] = f"{value} {flag}" + + +def set_default_lint_env(env: MutableMapping[str, str]) -> None: + for strict_lint in STRICT_LINTS: + append_env_flag(env, "DYLINT_RUSTFLAGS", f"-D {strict_lint}") + append_env_flag(env, "DYLINT_RUSTFLAGS", f"-A {NOISE_LINT}") + if not env.get("CARGO_INCREMENTAL"): + env["CARGO_INCREMENTAL"] = "0" + + +def die(message: str) -> "Never": + print(message, file=sys.stderr) + raise SystemExit(1) + + +def require_command(name: str, install_message: str | None = None) -> str: + executable = shutil.which(name) + if executable is None: + if install_message is None: + die(f"{name} is required but was not found on PATH.") + die(install_message) + return executable + + +def run_capture( + args: Sequence[str], env: MutableMapping[str, str] | None = None +) -> str: + try: + completed = subprocess.run( + list(args), + capture_output=True, + check=True, + env=None if env is None else dict(env), + text=True, + ) + except subprocess.CalledProcessError as error: + command = shlex.join(str(part) for part in error.cmd) + stderr = error.stderr.strip() + stdout = error.stdout.strip() + output = stderr or stdout + if output: + die(f"{command} failed:\n{output}") + die(f"{command} failed with exit code {error.returncode}") + return completed.stdout.strip() + + +def ensure_source_prerequisites(env: MutableMapping[str, str]) -> None: + require_command( + "cargo-dylint", + "argument-comment-lint source wrapper requires cargo-dylint and dylint-link.\n" + "Install them with:\n" + " cargo install --locked cargo-dylint dylint-link", + ) + require_command( + "dylint-link", + "argument-comment-lint source wrapper requires cargo-dylint and dylint-link.\n" + "Install them with:\n" + " cargo install --locked cargo-dylint dylint-link", + ) + require_command( + "rustup", + "argument-comment-lint source wrapper requires rustup.\n" + f"Install the {TOOLCHAIN_CHANNEL} toolchain with:\n" + f" rustup toolchain install {TOOLCHAIN_CHANNEL} \\\n" + " --component llvm-tools-preview \\\n" + " --component rustc-dev \\\n" + " --component rust-src", + ) + toolchains = run_capture(["rustup", "toolchain", "list"], env=env) + if not any(line.startswith(TOOLCHAIN_CHANNEL) for line in toolchains.splitlines()): + die( + "argument-comment-lint source wrapper requires the " + f"{TOOLCHAIN_CHANNEL} toolchain with rustc-dev support.\n" + "Install it with:\n" + f" rustup toolchain install {TOOLCHAIN_CHANNEL} \\\n" + " --component llvm-tools-preview \\\n" + " --component rustc-dev \\\n" + " --component rust-src" + ) + + +def prefer_rustup_shims(env: MutableMapping[str, str]) -> None: + if env.get("CODEX_ARGUMENT_COMMENT_LINT_SKIP_RUSTUP_SHIMS") == "1": + return + + rustup = shutil.which("rustup", path=env.get("PATH")) + if rustup is None: + return + + rustup_bin_dir = str(Path(rustup).resolve().parent) + path_entries = [ + entry + for entry in env.get("PATH", "").split(os.pathsep) + if entry and entry != rustup_bin_dir + ] + env["PATH"] = os.pathsep.join([rustup_bin_dir, *path_entries]) + + if not env.get("RUSTUP_HOME"): + rustup_home = run_capture(["rustup", "show", "home"], env=env) + if rustup_home: + env["RUSTUP_HOME"] = rustup_home + + +def fetch_packaged_entrypoint( + dotslash_manifest: Path, env: MutableMapping[str, str] +) -> Path: + require_command( + "dotslash", + "argument-comment-lint prebuilt wrapper requires dotslash.\n" + "Install dotslash, or use:\n" + " ./tools/argument-comment-lint/run.py ...", + ) + entrypoint = run_capture( + ["dotslash", "--", "fetch", str(dotslash_manifest)], env=env + ) + return Path(entrypoint).resolve() + + +def find_packaged_cargo_dylint(package_entrypoint: Path) -> Path: + bin_dir = package_entrypoint.parent + cargo_dylint = bin_dir / "cargo-dylint" + if not cargo_dylint.is_file(): + cargo_dylint = bin_dir / "cargo-dylint.exe" + if not cargo_dylint.is_file(): + die(f"bundled cargo-dylint executable not found under {bin_dir}") + return cargo_dylint + + +def normalize_packaged_library(package_entrypoint: Path) -> Path: + library_dir = package_entrypoint.parent.parent / "lib" + libraries = sorted(path for path in library_dir.glob("*@*") if path.is_file()) + if not libraries: + die(f"no packaged Dylint library found in {library_dir}") + if len(libraries) != 1: + die(f"expected exactly one packaged Dylint library in {library_dir}") + + library_path = libraries[0] + match = _NIGHTLY_LIBRARY_PATTERN.match(library_path.stem) + if match is None: + return library_path + + temp_dir = Path(tempfile.mkdtemp(prefix="argument-comment-lint.")) + normalized_library_path = temp_dir / f"{match.group(1)}{library_path.suffix}" + shutil.copy2(library_path, normalized_library_path) + return normalized_library_path + + +def exec_command(command: Sequence[str], env: MutableMapping[str, str]) -> "Never": + try: + completed = subprocess.run(list(command), env=dict(env), check=False) + except FileNotFoundError: + die(f"{command[0]} is required but was not found on PATH.") + raise SystemExit(completed.returncode)