Add files using upload-large-folder tool
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +3 -0
- .github/codex-cli-splash.png +3 -0
- codex-cli/bin/codex.js +295 -0
- codex-cli/scripts/README.md +23 -0
- codex-cli/scripts/build_npm_package.py +461 -0
- codex-cli/scripts/init_firewall.sh +115 -0
- codex-cli/scripts/run_in_container.sh +95 -0
- codex-rs/agent-identity/src/lib.rs +1000 -0
- codex-rs/ansi-escape/src/lib.rs +58 -0
- codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst +3 -0
- codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst +3 -0
- codex-rs/app-server-transport/src/connection_auth.rs +51 -0
- codex-rs/app-server-transport/src/daemon_recovery.rs +82 -0
- codex-rs/app-server-transport/src/daemon_shutdown.rs +53 -0
- codex-rs/app-server-transport/src/daemon_shutdown_tests.rs +22 -0
- codex-rs/app-server-transport/src/lib.rs +45 -0
- codex-rs/app-server-transport/src/outgoing_message.rs +59 -0
- codex-rs/app-server-transport/src/transport/auth.rs +751 -0
- codex-rs/app-server-transport/src/transport/mod.rs +602 -0
- codex-rs/app-server-transport/src/transport/remote_control/auth.rs +282 -0
- codex-rs/app-server-transport/src/transport/remote_control/client_tracker.rs +944 -0
- codex-rs/app-server-transport/src/transport/remote_control/clients.rs +303 -0
- codex-rs/app-server-transport/src/transport/remote_control/controller.rs +368 -0
- codex-rs/app-server-transport/src/transport/remote_control/desired_state.rs +155 -0
- codex-rs/app-server-transport/src/transport/remote_control/enroll.rs +770 -0
- codex-rs/app-server-transport/src/transport/remote_control/host_device.rs +74 -0
- codex-rs/app-server-transport/src/transport/remote_control/host_device_tests.rs +38 -0
- codex-rs/app-server-transport/src/transport/remote_control/mod.rs +1002 -0
- codex-rs/app-server-transport/src/transport/remote_control/persistence.rs +165 -0
- codex-rs/app-server-transport/src/transport/remote_control/persistence_tests.rs +51 -0
- codex-rs/app-server-transport/src/transport/remote_control/protocol.rs +401 -0
- codex-rs/app-server-transport/src/transport/remote_control/segment.rs +469 -0
- codex-rs/app-server-transport/src/transport/remote_control/segment_tests.rs +450 -0
- codex-rs/app-server-transport/src/transport/remote_control/server_api.rs +382 -0
- codex-rs/app-server-transport/src/transport/remote_control/server_api_tests.rs +321 -0
- codex-rs/app-server-transport/src/transport/remote_control/tests.rs +0 -0
- codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs +416 -0
- codex-rs/app-server-transport/src/transport/remote_control/tests/pairing_tests.rs +1137 -0
- codex-rs/app-server-transport/src/transport/remote_control/tests/retry_tests.rs +203 -0
- codex-rs/app-server-transport/src/transport/remote_control/websocket.rs +0 -0
- codex-rs/app-server-transport/src/transport/remote_control/websocket_refresh_tests.rs +768 -0
- codex-rs/app-server-transport/src/transport/stdio.rs +194 -0
- codex-rs/app-server-transport/src/transport/unix_socket.rs +265 -0
- codex-rs/app-server-transport/src/transport/unix_socket_tests.rs +340 -0
- codex-rs/app-server-transport/src/transport/websocket.rs +389 -0
- codex-rs/bwrap/src/main.rs +45 -0
- codex-rs/codex-mcp/src/agent_plugin_config.rs +532 -0
- codex-rs/codex-mcp/src/auth_changes.rs +62 -0
- codex-rs/codex-mcp/src/auth_changes_tests.rs +104 -0
- codex-rs/codex-mcp/src/auth_elicitation.rs +435 -0
.gitattributes
CHANGED
|
@@ -1,3 +1,6 @@
|
|
| 1 |
codex-rs/app-server-protocol/schema/** linguist-generated
|
| 2 |
codex-rs/hooks/schema/generated/** linguist-generated
|
| 3 |
third_party/voice/sources.json text eol=lf
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
codex-rs/app-server-protocol/schema/** linguist-generated
|
| 2 |
codex-rs/hooks/schema/generated/** linguist-generated
|
| 3 |
third_party/voice/sources.json text eol=lf
|
| 4 |
+
.github/codex-cli-splash.png filter=lfs diff=lfs merge=lfs -text
|
| 5 |
+
codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst filter=lfs diff=lfs merge=lfs -text
|
| 6 |
+
codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst filter=lfs diff=lfs merge=lfs -text
|
.github/codex-cli-splash.png
ADDED
|
Git LFS Details
|
codex-cli/bin/codex.js
ADDED
|
@@ -0,0 +1,295 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env node
|
| 2 |
+
// Unified entry point for the Codex CLI.
|
| 3 |
+
|
| 4 |
+
import { spawn } from "node:child_process";
|
| 5 |
+
import { existsSync, readFileSync, realpathSync } from "fs";
|
| 6 |
+
import { createRequire } from "node:module";
|
| 7 |
+
import path from "path";
|
| 8 |
+
import { fileURLToPath } from "url";
|
| 9 |
+
|
| 10 |
+
// __dirname equivalent in ESM
|
| 11 |
+
const __filename = fileURLToPath(import.meta.url);
|
| 12 |
+
const __dirname = path.dirname(__filename);
|
| 13 |
+
const require = createRequire(import.meta.url);
|
| 14 |
+
const codexPackageRoot = realpathSync(path.join(__dirname, ".."));
|
| 15 |
+
|
| 16 |
+
const PLATFORM_PACKAGE_BY_TARGET = {
|
| 17 |
+
"x86_64-unknown-linux-musl": "@openai/codex-linux-x64",
|
| 18 |
+
"aarch64-unknown-linux-musl": "@openai/codex-linux-arm64",
|
| 19 |
+
"x86_64-apple-darwin": "@openai/codex-darwin-x64",
|
| 20 |
+
"aarch64-apple-darwin": "@openai/codex-darwin-arm64",
|
| 21 |
+
"x86_64-pc-windows-msvc": "@openai/codex-win32-x64",
|
| 22 |
+
"aarch64-pc-windows-msvc": "@openai/codex-win32-arm64",
|
| 23 |
+
};
|
| 24 |
+
|
| 25 |
+
const { platform, arch } = process;
|
| 26 |
+
|
| 27 |
+
let targetTriple = null;
|
| 28 |
+
switch (platform) {
|
| 29 |
+
case "linux":
|
| 30 |
+
case "android":
|
| 31 |
+
switch (arch) {
|
| 32 |
+
case "x64":
|
| 33 |
+
targetTriple = "x86_64-unknown-linux-musl";
|
| 34 |
+
break;
|
| 35 |
+
case "arm64":
|
| 36 |
+
targetTriple = "aarch64-unknown-linux-musl";
|
| 37 |
+
break;
|
| 38 |
+
default:
|
| 39 |
+
break;
|
| 40 |
+
}
|
| 41 |
+
break;
|
| 42 |
+
case "darwin":
|
| 43 |
+
switch (arch) {
|
| 44 |
+
case "x64":
|
| 45 |
+
targetTriple = "x86_64-apple-darwin";
|
| 46 |
+
break;
|
| 47 |
+
case "arm64":
|
| 48 |
+
targetTriple = "aarch64-apple-darwin";
|
| 49 |
+
break;
|
| 50 |
+
default:
|
| 51 |
+
break;
|
| 52 |
+
}
|
| 53 |
+
break;
|
| 54 |
+
case "win32":
|
| 55 |
+
switch (arch) {
|
| 56 |
+
case "x64":
|
| 57 |
+
targetTriple = "x86_64-pc-windows-msvc";
|
| 58 |
+
break;
|
| 59 |
+
case "arm64":
|
| 60 |
+
targetTriple = "aarch64-pc-windows-msvc";
|
| 61 |
+
break;
|
| 62 |
+
default:
|
| 63 |
+
break;
|
| 64 |
+
}
|
| 65 |
+
break;
|
| 66 |
+
default:
|
| 67 |
+
break;
|
| 68 |
+
}
|
| 69 |
+
|
| 70 |
+
if (!targetTriple) {
|
| 71 |
+
throw new Error(`Unsupported platform: ${platform} (${arch})`);
|
| 72 |
+
}
|
| 73 |
+
|
| 74 |
+
const platformPackage = PLATFORM_PACKAGE_BY_TARGET[targetTriple];
|
| 75 |
+
if (!platformPackage) {
|
| 76 |
+
throw new Error(`Unsupported target triple: ${targetTriple}`);
|
| 77 |
+
}
|
| 78 |
+
|
| 79 |
+
function findCodexExecutable() {
|
| 80 |
+
let vendorRoot;
|
| 81 |
+
try {
|
| 82 |
+
const packageJsonPath = require.resolve(`${platformPackage}/package.json`);
|
| 83 |
+
vendorRoot = path.join(path.dirname(packageJsonPath), "vendor");
|
| 84 |
+
} catch {
|
| 85 |
+
vendorRoot = path.join(__dirname, "..", "vendor");
|
| 86 |
+
}
|
| 87 |
+
|
| 88 |
+
const codexExecutable = path.join(
|
| 89 |
+
vendorRoot,
|
| 90 |
+
targetTriple,
|
| 91 |
+
"bin",
|
| 92 |
+
process.platform === "win32" ? "codex.exe" : "codex",
|
| 93 |
+
);
|
| 94 |
+
if (existsSync(codexExecutable)) {
|
| 95 |
+
return codexExecutable;
|
| 96 |
+
}
|
| 97 |
+
|
| 98 |
+
const packageManager = detectPackageManager();
|
| 99 |
+
const updateCommand =
|
| 100 |
+
packageManager === "bun"
|
| 101 |
+
? "bun install -g @openai/codex@latest"
|
| 102 |
+
: packageManager === "pnpm"
|
| 103 |
+
? "pnpm add -g @openai/codex@latest"
|
| 104 |
+
: packageManager === "vite-plus"
|
| 105 |
+
? "vp install -g @openai/codex@latest"
|
| 106 |
+
: "npm install -g @openai/codex@latest";
|
| 107 |
+
throw new Error(
|
| 108 |
+
`Missing optional dependency ${platformPackage}. Reinstall Codex: ${updateCommand}`,
|
| 109 |
+
);
|
| 110 |
+
}
|
| 111 |
+
|
| 112 |
+
const binaryPath = findCodexExecutable();
|
| 113 |
+
|
| 114 |
+
// Use an asynchronous spawn instead of spawnSync so that Node is able to
|
| 115 |
+
// respond to signals (e.g. Ctrl-C / SIGINT) while the native binary is
|
| 116 |
+
// executing. This allows us to forward those signals to the child process
|
| 117 |
+
// and guarantees that when either the child terminates or the parent
|
| 118 |
+
// receives a fatal signal, both processes exit in a predictable manner.
|
| 119 |
+
|
| 120 |
+
function isPnpmOwnedCodexInstall(nodeModulesDir) {
|
| 121 |
+
if (!existsSync(path.join(nodeModulesDir, ".modules.yaml"))) {
|
| 122 |
+
return false;
|
| 123 |
+
}
|
| 124 |
+
|
| 125 |
+
try {
|
| 126 |
+
return (
|
| 127 |
+
realpathSync(path.join(nodeModulesDir, "@openai", "codex")) ===
|
| 128 |
+
codexPackageRoot
|
| 129 |
+
);
|
| 130 |
+
} catch {
|
| 131 |
+
return false;
|
| 132 |
+
}
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
function isVitePlusOwnedCodexInstall(packagesDir) {
|
| 136 |
+
if (path.basename(packagesDir) !== "packages") {
|
| 137 |
+
return false;
|
| 138 |
+
}
|
| 139 |
+
|
| 140 |
+
try {
|
| 141 |
+
const metadata = JSON.parse(
|
| 142 |
+
readFileSync(path.join(packagesDir, "@openai", "codex.json"), "utf8"),
|
| 143 |
+
);
|
| 144 |
+
if (metadata.name !== "@openai/codex") {
|
| 145 |
+
return false;
|
| 146 |
+
}
|
| 147 |
+
|
| 148 |
+
// Vite+ records the active global installation in packages/@openai/codex.json.
|
| 149 |
+
// Older installs have no ID or append a #-prefixed ID to the package name;
|
| 150 |
+
// newer installs put the ID in a subdirectory of the package prefix.
|
| 151 |
+
const installId = metadata.installId || "";
|
| 152 |
+
const installDir = installId.startsWith("#")
|
| 153 |
+
? path.join(packagesDir, `@openai/codex${installId}`)
|
| 154 |
+
: path.join(packagesDir, "@openai/codex", installId);
|
| 155 |
+
for (const nodeModulesDir of [
|
| 156 |
+
path.join(installDir, "lib", "node_modules"),
|
| 157 |
+
path.join(installDir, "node_modules"),
|
| 158 |
+
]) {
|
| 159 |
+
const packageRoot = path.join(nodeModulesDir, "@openai", "codex");
|
| 160 |
+
if (
|
| 161 |
+
existsSync(packageRoot) &&
|
| 162 |
+
realpathSync(packageRoot) === codexPackageRoot
|
| 163 |
+
) {
|
| 164 |
+
return true;
|
| 165 |
+
}
|
| 166 |
+
}
|
| 167 |
+
} catch {
|
| 168 |
+
// Missing or unreadable ownership metadata must not prevent Codex starting.
|
| 169 |
+
}
|
| 170 |
+
return false;
|
| 171 |
+
}
|
| 172 |
+
|
| 173 |
+
/**
|
| 174 |
+
* Use heuristics to detect the package manager that was used to install Codex
|
| 175 |
+
* in order to give the user a hint about how to update it.
|
| 176 |
+
*/
|
| 177 |
+
function detectPackageManager() {
|
| 178 |
+
// Package-manager ownership metadata can be several parents above the package.
|
| 179 |
+
// Search ancestors of both the canonical package root and lexical entrypoint
|
| 180 |
+
// because the package manager may link either path.
|
| 181 |
+
const entrypointDir = path.dirname(path.resolve(process.argv[1]));
|
| 182 |
+
for (const startDir of new Set([codexPackageRoot, entrypointDir])) {
|
| 183 |
+
const filesystemRoot = path.parse(startDir).root;
|
| 184 |
+
for (
|
| 185 |
+
let currentDir = startDir;
|
| 186 |
+
currentDir !== filesystemRoot;
|
| 187 |
+
currentDir = path.dirname(currentDir)
|
| 188 |
+
) {
|
| 189 |
+
if (isVitePlusOwnedCodexInstall(currentDir)) {
|
| 190 |
+
return "vite-plus";
|
| 191 |
+
}
|
| 192 |
+
if (isPnpmOwnedCodexInstall(path.join(currentDir, "node_modules"))) {
|
| 193 |
+
return "pnpm";
|
| 194 |
+
}
|
| 195 |
+
}
|
| 196 |
+
|
| 197 |
+
if (isPnpmOwnedCodexInstall(path.join(filesystemRoot, "node_modules"))) {
|
| 198 |
+
return "pnpm";
|
| 199 |
+
}
|
| 200 |
+
}
|
| 201 |
+
|
| 202 |
+
const userAgent = process.env.npm_config_user_agent || "";
|
| 203 |
+
if (/\bbun\//.test(userAgent)) {
|
| 204 |
+
return "bun";
|
| 205 |
+
}
|
| 206 |
+
|
| 207 |
+
const execPath = process.env.npm_execpath || "";
|
| 208 |
+
if (execPath.includes("bun")) {
|
| 209 |
+
return "bun";
|
| 210 |
+
}
|
| 211 |
+
|
| 212 |
+
if (
|
| 213 |
+
__dirname.includes(".bun/install/global") ||
|
| 214 |
+
__dirname.includes(".bun\\install\\global")
|
| 215 |
+
) {
|
| 216 |
+
return "bun";
|
| 217 |
+
}
|
| 218 |
+
|
| 219 |
+
return userAgent ? "npm" : null;
|
| 220 |
+
}
|
| 221 |
+
|
| 222 |
+
const packageManager = detectPackageManager();
|
| 223 |
+
const packageManagerEnvVar =
|
| 224 |
+
packageManager === "bun"
|
| 225 |
+
? "CODEX_MANAGED_BY_BUN"
|
| 226 |
+
: packageManager === "pnpm"
|
| 227 |
+
? "CODEX_MANAGED_BY_PNPM"
|
| 228 |
+
: packageManager === "vite-plus"
|
| 229 |
+
? "CODEX_MANAGED_BY_VITE_PLUS"
|
| 230 |
+
: "CODEX_MANAGED_BY_NPM";
|
| 231 |
+
const env = {
|
| 232 |
+
...process.env,
|
| 233 |
+
CODEX_MANAGED_PACKAGE_ROOT: codexPackageRoot,
|
| 234 |
+
};
|
| 235 |
+
delete env.CODEX_MANAGED_BY_NPM;
|
| 236 |
+
delete env.CODEX_MANAGED_BY_BUN;
|
| 237 |
+
delete env.CODEX_MANAGED_BY_PNPM;
|
| 238 |
+
delete env.CODEX_MANAGED_BY_VITE_PLUS;
|
| 239 |
+
env[packageManagerEnvVar] = "1";
|
| 240 |
+
|
| 241 |
+
const child = spawn(binaryPath, process.argv.slice(2), {
|
| 242 |
+
stdio: "inherit",
|
| 243 |
+
env,
|
| 244 |
+
});
|
| 245 |
+
|
| 246 |
+
child.on("error", (err) => {
|
| 247 |
+
// Typically triggered when the binary is missing or not executable.
|
| 248 |
+
// Re-throwing here will terminate the parent with a non-zero exit code
|
| 249 |
+
// while still printing a helpful stack trace.
|
| 250 |
+
// eslint-disable-next-line no-console
|
| 251 |
+
console.error(err);
|
| 252 |
+
process.exit(1);
|
| 253 |
+
});
|
| 254 |
+
|
| 255 |
+
// Forward common termination signals to the child so that it shuts down
|
| 256 |
+
// gracefully. In the handler we temporarily disable the default behavior of
|
| 257 |
+
// exiting immediately; once the child has been signaled we simply wait for
|
| 258 |
+
// its exit event which will in turn terminate the parent (see below).
|
| 259 |
+
const forwardSignal = (signal) => {
|
| 260 |
+
if (child.killed) {
|
| 261 |
+
return;
|
| 262 |
+
}
|
| 263 |
+
try {
|
| 264 |
+
child.kill(signal);
|
| 265 |
+
} catch {
|
| 266 |
+
/* ignore */
|
| 267 |
+
}
|
| 268 |
+
};
|
| 269 |
+
|
| 270 |
+
["SIGINT", "SIGTERM", "SIGHUP"].forEach((sig) => {
|
| 271 |
+
process.on(sig, () => forwardSignal(sig));
|
| 272 |
+
});
|
| 273 |
+
|
| 274 |
+
// When the child exits, mirror its termination reason in the parent so that
|
| 275 |
+
// shell scripts and other tooling observe the correct exit status.
|
| 276 |
+
// Wrap the lifetime of the child process in a Promise so that we can await
|
| 277 |
+
// its termination in a structured way. The Promise resolves with an object
|
| 278 |
+
// describing how the child exited: either via exit code or due to a signal.
|
| 279 |
+
const childResult = await new Promise((resolve) => {
|
| 280 |
+
child.on("exit", (code, signal) => {
|
| 281 |
+
if (signal) {
|
| 282 |
+
resolve({ type: "signal", signal });
|
| 283 |
+
} else {
|
| 284 |
+
resolve({ type: "code", exitCode: code ?? 1 });
|
| 285 |
+
}
|
| 286 |
+
});
|
| 287 |
+
});
|
| 288 |
+
|
| 289 |
+
if (childResult.type === "signal") {
|
| 290 |
+
// Re-emit the same signal so that the parent terminates with the expected
|
| 291 |
+
// semantics (this also sets the correct exit code of 128 + n).
|
| 292 |
+
process.kill(process.pid, childResult.signal);
|
| 293 |
+
} else {
|
| 294 |
+
process.exit(childResult.exitCode);
|
| 295 |
+
}
|
codex-cli/scripts/README.md
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# npm releases
|
| 2 |
+
|
| 3 |
+
Use the staging helper in the repo root to generate npm tarballs for a release. For
|
| 4 |
+
example, to stage the CLI, responses proxy, and SDK packages for version `0.6.0`:
|
| 5 |
+
|
| 6 |
+
```bash
|
| 7 |
+
./scripts/stage_npm_packages.py \
|
| 8 |
+
--release-version 0.6.0 \
|
| 9 |
+
--package codex \
|
| 10 |
+
--package codex-responses-api-proxy \
|
| 11 |
+
--package codex-sdk
|
| 12 |
+
```
|
| 13 |
+
|
| 14 |
+
This downloads the required native package archive artifacts, hydrates `vendor/` for
|
| 15 |
+
each package, and writes tarballs to `dist/npm/`.
|
| 16 |
+
|
| 17 |
+
When `--package codex` is provided, the staging helper builds the lightweight
|
| 18 |
+
`@openai/codex` meta package plus all platform-native `@openai/codex` variants
|
| 19 |
+
that are later published under platform-specific dist-tags.
|
| 20 |
+
|
| 21 |
+
Direct `build_npm_package.py` invocations are still useful for package-specific
|
| 22 |
+
debugging, but native packages expect `--vendor-src` to point at a prehydrated
|
| 23 |
+
`vendor/` tree. Release packaging should use `scripts/stage_npm_packages.py`
|
codex-cli/scripts/build_npm_package.py
ADDED
|
@@ -0,0 +1,461 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Stage and optionally package the @openai/codex npm module."""
|
| 3 |
+
|
| 4 |
+
import argparse
|
| 5 |
+
import json
|
| 6 |
+
import os
|
| 7 |
+
import shutil
|
| 8 |
+
import subprocess
|
| 9 |
+
import sys
|
| 10 |
+
import tempfile
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
|
| 13 |
+
SCRIPT_DIR = Path(__file__).resolve().parent
|
| 14 |
+
CODEX_CLI_ROOT = SCRIPT_DIR.parent
|
| 15 |
+
REPO_ROOT = CODEX_CLI_ROOT.parent
|
| 16 |
+
RESPONSES_API_PROXY_NPM_ROOT = REPO_ROOT / "codex-rs" / "responses-api-proxy" / "npm"
|
| 17 |
+
CODEX_SDK_ROOT = REPO_ROOT / "sdk" / "typescript"
|
| 18 |
+
CODEX_NPM_NAME = "@openai/codex"
|
| 19 |
+
CODEX_PACKAGE_COMPONENT = "codex-package"
|
| 20 |
+
|
| 21 |
+
# `npm_name` is the local optional-dependency alias consumed by `bin/codex.js`.
|
| 22 |
+
# The underlying package published to npm is always `@openai/codex`.
|
| 23 |
+
CODEX_PLATFORM_PACKAGES: dict[str, dict[str, str]] = {
|
| 24 |
+
"codex-linux-x64": {
|
| 25 |
+
"npm_name": "@openai/codex-linux-x64",
|
| 26 |
+
"npm_tag": "linux-x64",
|
| 27 |
+
"target_triple": "x86_64-unknown-linux-musl",
|
| 28 |
+
"os": "linux",
|
| 29 |
+
"cpu": "x64",
|
| 30 |
+
},
|
| 31 |
+
"codex-linux-arm64": {
|
| 32 |
+
"npm_name": "@openai/codex-linux-arm64",
|
| 33 |
+
"npm_tag": "linux-arm64",
|
| 34 |
+
"target_triple": "aarch64-unknown-linux-musl",
|
| 35 |
+
"os": "linux",
|
| 36 |
+
"cpu": "arm64",
|
| 37 |
+
},
|
| 38 |
+
"codex-darwin-x64": {
|
| 39 |
+
"npm_name": "@openai/codex-darwin-x64",
|
| 40 |
+
"npm_tag": "darwin-x64",
|
| 41 |
+
"target_triple": "x86_64-apple-darwin",
|
| 42 |
+
"os": "darwin",
|
| 43 |
+
"cpu": "x64",
|
| 44 |
+
},
|
| 45 |
+
"codex-darwin-arm64": {
|
| 46 |
+
"npm_name": "@openai/codex-darwin-arm64",
|
| 47 |
+
"npm_tag": "darwin-arm64",
|
| 48 |
+
"target_triple": "aarch64-apple-darwin",
|
| 49 |
+
"os": "darwin",
|
| 50 |
+
"cpu": "arm64",
|
| 51 |
+
},
|
| 52 |
+
"codex-win32-x64": {
|
| 53 |
+
"npm_name": "@openai/codex-win32-x64",
|
| 54 |
+
"npm_tag": "win32-x64",
|
| 55 |
+
"target_triple": "x86_64-pc-windows-msvc",
|
| 56 |
+
"os": "win32",
|
| 57 |
+
"cpu": "x64",
|
| 58 |
+
},
|
| 59 |
+
"codex-win32-arm64": {
|
| 60 |
+
"npm_name": "@openai/codex-win32-arm64",
|
| 61 |
+
"npm_tag": "win32-arm64",
|
| 62 |
+
"target_triple": "aarch64-pc-windows-msvc",
|
| 63 |
+
"os": "win32",
|
| 64 |
+
"cpu": "arm64",
|
| 65 |
+
},
|
| 66 |
+
}
|
| 67 |
+
|
| 68 |
+
PACKAGE_EXPANSIONS: dict[str, list[str]] = {
|
| 69 |
+
"codex": ["codex", *CODEX_PLATFORM_PACKAGES],
|
| 70 |
+
}
|
| 71 |
+
|
| 72 |
+
PACKAGE_NATIVE_COMPONENTS: dict[str, list[str]] = {
|
| 73 |
+
"codex": [],
|
| 74 |
+
"codex-linux-x64": [CODEX_PACKAGE_COMPONENT],
|
| 75 |
+
"codex-linux-arm64": [CODEX_PACKAGE_COMPONENT],
|
| 76 |
+
"codex-darwin-x64": [CODEX_PACKAGE_COMPONENT],
|
| 77 |
+
"codex-darwin-arm64": [CODEX_PACKAGE_COMPONENT],
|
| 78 |
+
"codex-win32-x64": [CODEX_PACKAGE_COMPONENT],
|
| 79 |
+
"codex-win32-arm64": [CODEX_PACKAGE_COMPONENT],
|
| 80 |
+
"codex-responses-api-proxy": ["codex-responses-api-proxy"],
|
| 81 |
+
"codex-sdk": [],
|
| 82 |
+
}
|
| 83 |
+
|
| 84 |
+
PACKAGE_TARGET_FILTERS: dict[str, str] = {
|
| 85 |
+
package_name: package_config["target_triple"]
|
| 86 |
+
for package_name, package_config in CODEX_PLATFORM_PACKAGES.items()
|
| 87 |
+
}
|
| 88 |
+
|
| 89 |
+
PACKAGE_CHOICES = tuple(PACKAGE_NATIVE_COMPONENTS)
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
def parse_args() -> argparse.Namespace:
|
| 93 |
+
parser = argparse.ArgumentParser(
|
| 94 |
+
description="Build or stage the Codex CLI npm package."
|
| 95 |
+
)
|
| 96 |
+
parser.add_argument(
|
| 97 |
+
"--package",
|
| 98 |
+
choices=PACKAGE_CHOICES,
|
| 99 |
+
default="codex",
|
| 100 |
+
help="Which npm package to stage (default: codex).",
|
| 101 |
+
)
|
| 102 |
+
parser.add_argument(
|
| 103 |
+
"--version",
|
| 104 |
+
help="Version number to write to package.json inside the staged package.",
|
| 105 |
+
)
|
| 106 |
+
parser.add_argument(
|
| 107 |
+
"--release-version",
|
| 108 |
+
help=("Version to stage for npm release."),
|
| 109 |
+
)
|
| 110 |
+
parser.add_argument(
|
| 111 |
+
"--staging-dir",
|
| 112 |
+
type=Path,
|
| 113 |
+
help=(
|
| 114 |
+
"Directory to stage the package contents. Defaults to a new temporary directory "
|
| 115 |
+
"if omitted. The directory must be empty when provided."
|
| 116 |
+
),
|
| 117 |
+
)
|
| 118 |
+
parser.add_argument(
|
| 119 |
+
"--tmp",
|
| 120 |
+
dest="staging_dir",
|
| 121 |
+
type=Path,
|
| 122 |
+
help=argparse.SUPPRESS,
|
| 123 |
+
)
|
| 124 |
+
parser.add_argument(
|
| 125 |
+
"--pack-output",
|
| 126 |
+
type=Path,
|
| 127 |
+
help="Path where the generated npm tarball should be written.",
|
| 128 |
+
)
|
| 129 |
+
parser.add_argument(
|
| 130 |
+
"--vendor-src",
|
| 131 |
+
type=Path,
|
| 132 |
+
help="Directory containing pre-installed native binaries to bundle (vendor root).",
|
| 133 |
+
)
|
| 134 |
+
return parser.parse_args()
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
def main() -> int:
|
| 138 |
+
args = parse_args()
|
| 139 |
+
|
| 140 |
+
package = args.package
|
| 141 |
+
version = args.version
|
| 142 |
+
release_version = args.release_version
|
| 143 |
+
if release_version:
|
| 144 |
+
if version and version != release_version:
|
| 145 |
+
raise RuntimeError(
|
| 146 |
+
"--version and --release-version must match when both are provided."
|
| 147 |
+
)
|
| 148 |
+
version = release_version
|
| 149 |
+
|
| 150 |
+
if not version:
|
| 151 |
+
raise RuntimeError("Must specify --version or --release-version.")
|
| 152 |
+
|
| 153 |
+
staging_dir, created_temp = prepare_staging_dir(args.staging_dir)
|
| 154 |
+
|
| 155 |
+
try:
|
| 156 |
+
stage_sources(staging_dir, version, package)
|
| 157 |
+
|
| 158 |
+
vendor_src = args.vendor_src.resolve() if args.vendor_src else None
|
| 159 |
+
native_components = PACKAGE_NATIVE_COMPONENTS.get(package, [])
|
| 160 |
+
target_filter = PACKAGE_TARGET_FILTERS.get(package)
|
| 161 |
+
|
| 162 |
+
if native_components:
|
| 163 |
+
if vendor_src is None:
|
| 164 |
+
components_str = ", ".join(native_components)
|
| 165 |
+
raise RuntimeError(
|
| 166 |
+
"Native components "
|
| 167 |
+
f"({components_str}) required for package '{package}'. Provide --vendor-src "
|
| 168 |
+
"pointing to a directory containing pre-installed binaries."
|
| 169 |
+
)
|
| 170 |
+
|
| 171 |
+
copy_native_binaries(
|
| 172 |
+
vendor_src,
|
| 173 |
+
staging_dir,
|
| 174 |
+
native_components,
|
| 175 |
+
target_filter={target_filter} if target_filter else None,
|
| 176 |
+
)
|
| 177 |
+
|
| 178 |
+
if release_version:
|
| 179 |
+
staging_dir_str = str(staging_dir)
|
| 180 |
+
if package == "codex":
|
| 181 |
+
print(
|
| 182 |
+
f"Staged version {version} for release in {staging_dir_str}\n\n"
|
| 183 |
+
"Verify the CLI:\n"
|
| 184 |
+
f" node {staging_dir_str}/bin/codex.js --version\n"
|
| 185 |
+
f" node {staging_dir_str}/bin/codex.js --help\n\n"
|
| 186 |
+
)
|
| 187 |
+
elif package == "codex-responses-api-proxy":
|
| 188 |
+
print(
|
| 189 |
+
f"Staged version {version} for release in {staging_dir_str}\n\n"
|
| 190 |
+
"Verify the responses API proxy:\n"
|
| 191 |
+
f" node {staging_dir_str}/bin/codex-responses-api-proxy.js --help\n\n"
|
| 192 |
+
)
|
| 193 |
+
elif package in CODEX_PLATFORM_PACKAGES:
|
| 194 |
+
print(
|
| 195 |
+
f"Staged version {version} for release in {staging_dir_str}\n\n"
|
| 196 |
+
"Verify native payload contents:\n"
|
| 197 |
+
f" ls {staging_dir_str}/vendor\n\n"
|
| 198 |
+
)
|
| 199 |
+
else:
|
| 200 |
+
print(
|
| 201 |
+
f"Staged version {version} for release in {staging_dir_str}\n\n"
|
| 202 |
+
"Verify the SDK contents:\n"
|
| 203 |
+
f" ls {staging_dir_str}/dist\n"
|
| 204 |
+
" node -e \"import('./dist/index.js').then(() => console.log('ok'))\"\n\n"
|
| 205 |
+
)
|
| 206 |
+
else:
|
| 207 |
+
print(f"Staged package in {staging_dir}")
|
| 208 |
+
|
| 209 |
+
if args.pack_output is not None:
|
| 210 |
+
output_path = run_npm_pack(staging_dir, args.pack_output)
|
| 211 |
+
print(f"npm pack output written to {output_path}")
|
| 212 |
+
finally:
|
| 213 |
+
if created_temp:
|
| 214 |
+
# Preserve the staging directory for further inspection.
|
| 215 |
+
pass
|
| 216 |
+
|
| 217 |
+
return 0
|
| 218 |
+
|
| 219 |
+
|
| 220 |
+
def prepare_staging_dir(staging_dir: Path | None) -> tuple[Path, bool]:
|
| 221 |
+
if staging_dir is not None:
|
| 222 |
+
staging_dir = staging_dir.resolve()
|
| 223 |
+
staging_dir.mkdir(parents=True, exist_ok=True)
|
| 224 |
+
if any(staging_dir.iterdir()):
|
| 225 |
+
raise RuntimeError(f"Staging directory {staging_dir} is not empty.")
|
| 226 |
+
return staging_dir, False
|
| 227 |
+
|
| 228 |
+
temp_dir = Path(tempfile.mkdtemp(prefix="codex-npm-stage-"))
|
| 229 |
+
return temp_dir, True
|
| 230 |
+
|
| 231 |
+
|
| 232 |
+
def stage_sources(staging_dir: Path, version: str, package: str) -> None:
|
| 233 |
+
package_json: dict
|
| 234 |
+
package_json_path: Path | None = None
|
| 235 |
+
|
| 236 |
+
if package == "codex":
|
| 237 |
+
bin_dir = staging_dir / "bin"
|
| 238 |
+
bin_dir.mkdir(parents=True, exist_ok=True)
|
| 239 |
+
shutil.copy2(CODEX_CLI_ROOT / "bin" / "codex.js", bin_dir / "codex.js")
|
| 240 |
+
|
| 241 |
+
readme_src = REPO_ROOT / "README.md"
|
| 242 |
+
if readme_src.exists():
|
| 243 |
+
shutil.copy2(readme_src, staging_dir / "README.md")
|
| 244 |
+
|
| 245 |
+
package_json_path = CODEX_CLI_ROOT / "package.json"
|
| 246 |
+
elif package in CODEX_PLATFORM_PACKAGES:
|
| 247 |
+
platform_package = CODEX_PLATFORM_PACKAGES[package]
|
| 248 |
+
platform_npm_tag = platform_package["npm_tag"]
|
| 249 |
+
platform_version = compute_platform_package_version(version, platform_npm_tag)
|
| 250 |
+
|
| 251 |
+
readme_src = REPO_ROOT / "README.md"
|
| 252 |
+
if readme_src.exists():
|
| 253 |
+
shutil.copy2(readme_src, staging_dir / "README.md")
|
| 254 |
+
|
| 255 |
+
with open(CODEX_CLI_ROOT / "package.json", "r", encoding="utf-8") as fh:
|
| 256 |
+
codex_package_json = json.load(fh)
|
| 257 |
+
|
| 258 |
+
package_json = {
|
| 259 |
+
"name": CODEX_NPM_NAME,
|
| 260 |
+
"version": platform_version,
|
| 261 |
+
"license": codex_package_json.get("license", "Apache-2.0"),
|
| 262 |
+
"os": [platform_package["os"]],
|
| 263 |
+
"cpu": [platform_package["cpu"]],
|
| 264 |
+
"files": ["vendor"],
|
| 265 |
+
"repository": codex_package_json.get("repository"),
|
| 266 |
+
}
|
| 267 |
+
|
| 268 |
+
engines = codex_package_json.get("engines")
|
| 269 |
+
if isinstance(engines, dict):
|
| 270 |
+
package_json["engines"] = engines
|
| 271 |
+
|
| 272 |
+
package_manager = codex_package_json.get("packageManager")
|
| 273 |
+
if isinstance(package_manager, str):
|
| 274 |
+
package_json["packageManager"] = package_manager
|
| 275 |
+
elif package == "codex-responses-api-proxy":
|
| 276 |
+
bin_dir = staging_dir / "bin"
|
| 277 |
+
bin_dir.mkdir(parents=True, exist_ok=True)
|
| 278 |
+
launcher_src = (
|
| 279 |
+
RESPONSES_API_PROXY_NPM_ROOT / "bin" / "codex-responses-api-proxy.js"
|
| 280 |
+
)
|
| 281 |
+
shutil.copy2(launcher_src, bin_dir / "codex-responses-api-proxy.js")
|
| 282 |
+
|
| 283 |
+
readme_src = RESPONSES_API_PROXY_NPM_ROOT / "README.md"
|
| 284 |
+
if readme_src.exists():
|
| 285 |
+
shutil.copy2(readme_src, staging_dir / "README.md")
|
| 286 |
+
|
| 287 |
+
package_json_path = RESPONSES_API_PROXY_NPM_ROOT / "package.json"
|
| 288 |
+
elif package == "codex-sdk":
|
| 289 |
+
package_json_path = CODEX_SDK_ROOT / "package.json"
|
| 290 |
+
stage_codex_sdk_sources(staging_dir)
|
| 291 |
+
else:
|
| 292 |
+
raise RuntimeError(f"Unknown package '{package}'.")
|
| 293 |
+
|
| 294 |
+
if package_json_path is not None:
|
| 295 |
+
with open(package_json_path, "r", encoding="utf-8") as fh:
|
| 296 |
+
package_json = json.load(fh)
|
| 297 |
+
package_json["version"] = version
|
| 298 |
+
|
| 299 |
+
if package == "codex":
|
| 300 |
+
package_json["files"] = ["bin/codex.js"]
|
| 301 |
+
package_json["optionalDependencies"] = {
|
| 302 |
+
CODEX_PLATFORM_PACKAGES[platform_package]["npm_name"]: (
|
| 303 |
+
f"npm:{CODEX_NPM_NAME}@"
|
| 304 |
+
f"{compute_platform_package_version(version, CODEX_PLATFORM_PACKAGES[platform_package]['npm_tag'])}"
|
| 305 |
+
)
|
| 306 |
+
for platform_package in PACKAGE_EXPANSIONS["codex"]
|
| 307 |
+
if platform_package != "codex"
|
| 308 |
+
}
|
| 309 |
+
|
| 310 |
+
elif package == "codex-sdk":
|
| 311 |
+
scripts = package_json.get("scripts")
|
| 312 |
+
if isinstance(scripts, dict):
|
| 313 |
+
scripts.pop("prepare", None)
|
| 314 |
+
|
| 315 |
+
dependencies = package_json.get("dependencies")
|
| 316 |
+
if not isinstance(dependencies, dict):
|
| 317 |
+
dependencies = {}
|
| 318 |
+
dependencies[CODEX_NPM_NAME] = version
|
| 319 |
+
package_json["dependencies"] = dependencies
|
| 320 |
+
|
| 321 |
+
with open(staging_dir / "package.json", "w", encoding="utf-8") as out:
|
| 322 |
+
json.dump(package_json, out, indent=2)
|
| 323 |
+
out.write("\n")
|
| 324 |
+
|
| 325 |
+
|
| 326 |
+
def compute_platform_package_version(version: str, platform_tag: str) -> str:
|
| 327 |
+
# npm forbids republishing the same package name/version, so each
|
| 328 |
+
# platform-specific tarball needs a unique version string.
|
| 329 |
+
return f"{version}-{platform_tag}"
|
| 330 |
+
|
| 331 |
+
|
| 332 |
+
def run_command(cmd: list[str], cwd: Path | None = None) -> None:
|
| 333 |
+
print("+", " ".join(cmd), flush=True)
|
| 334 |
+
subprocess.run(cmd, cwd=cwd, check=True)
|
| 335 |
+
|
| 336 |
+
|
| 337 |
+
def stage_codex_sdk_sources(staging_dir: Path) -> None:
|
| 338 |
+
package_root = CODEX_SDK_ROOT
|
| 339 |
+
|
| 340 |
+
run_command(["pnpm", "install", "--frozen-lockfile"], cwd=package_root)
|
| 341 |
+
run_command(["pnpm", "run", "build"], cwd=package_root)
|
| 342 |
+
|
| 343 |
+
dist_src = package_root / "dist"
|
| 344 |
+
if not dist_src.exists():
|
| 345 |
+
raise RuntimeError("codex-sdk build did not produce a dist directory.")
|
| 346 |
+
|
| 347 |
+
shutil.copytree(dist_src, staging_dir / "dist")
|
| 348 |
+
|
| 349 |
+
readme_src = package_root / "README.md"
|
| 350 |
+
if readme_src.exists():
|
| 351 |
+
shutil.copy2(readme_src, staging_dir / "README.md")
|
| 352 |
+
|
| 353 |
+
license_src = REPO_ROOT / "LICENSE"
|
| 354 |
+
if license_src.exists():
|
| 355 |
+
shutil.copy2(license_src, staging_dir / "LICENSE")
|
| 356 |
+
|
| 357 |
+
|
| 358 |
+
def copy_native_binaries(
|
| 359 |
+
vendor_src: Path,
|
| 360 |
+
staging_dir: Path,
|
| 361 |
+
components: list[str],
|
| 362 |
+
target_filter: set[str] | None = None,
|
| 363 |
+
) -> None:
|
| 364 |
+
vendor_src = vendor_src.resolve()
|
| 365 |
+
if not vendor_src.exists():
|
| 366 |
+
raise RuntimeError(f"Vendor source directory not found: {vendor_src}")
|
| 367 |
+
|
| 368 |
+
components_set = set(components)
|
| 369 |
+
if not components_set:
|
| 370 |
+
return
|
| 371 |
+
|
| 372 |
+
vendor_dest = staging_dir / "vendor"
|
| 373 |
+
if vendor_dest.exists():
|
| 374 |
+
shutil.rmtree(vendor_dest)
|
| 375 |
+
vendor_dest.mkdir(parents=True, exist_ok=True)
|
| 376 |
+
|
| 377 |
+
copied_targets: set[str] = set()
|
| 378 |
+
|
| 379 |
+
for target_dir in vendor_src.iterdir():
|
| 380 |
+
if not target_dir.is_dir():
|
| 381 |
+
continue
|
| 382 |
+
|
| 383 |
+
if target_filter is not None and target_dir.name not in target_filter:
|
| 384 |
+
continue
|
| 385 |
+
|
| 386 |
+
copied_targets.add(target_dir.name)
|
| 387 |
+
|
| 388 |
+
dest_target_dir = vendor_dest / target_dir.name
|
| 389 |
+
|
| 390 |
+
if CODEX_PACKAGE_COMPONENT in components_set:
|
| 391 |
+
if dest_target_dir.exists():
|
| 392 |
+
shutil.rmtree(dest_target_dir)
|
| 393 |
+
shutil.copytree(target_dir, dest_target_dir)
|
| 394 |
+
else:
|
| 395 |
+
dest_target_dir.mkdir(parents=True, exist_ok=True)
|
| 396 |
+
|
| 397 |
+
for component in sorted(components_set - {CODEX_PACKAGE_COMPONENT}):
|
| 398 |
+
src_component_dir = target_dir / component
|
| 399 |
+
if not src_component_dir.exists():
|
| 400 |
+
raise RuntimeError(
|
| 401 |
+
f"Missing native component '{component}' in vendor source: {src_component_dir}"
|
| 402 |
+
)
|
| 403 |
+
|
| 404 |
+
dest_component_dir = dest_target_dir / component
|
| 405 |
+
if dest_component_dir.exists():
|
| 406 |
+
shutil.rmtree(dest_component_dir)
|
| 407 |
+
shutil.copytree(src_component_dir, dest_component_dir)
|
| 408 |
+
|
| 409 |
+
if target_filter is not None:
|
| 410 |
+
missing_targets = sorted(target_filter - copied_targets)
|
| 411 |
+
if missing_targets:
|
| 412 |
+
missing_list = ", ".join(missing_targets)
|
| 413 |
+
raise RuntimeError(
|
| 414 |
+
f"Missing target directories in vendor source: {missing_list}"
|
| 415 |
+
)
|
| 416 |
+
|
| 417 |
+
|
| 418 |
+
def run_npm_pack(staging_dir: Path, output_path: Path) -> Path:
|
| 419 |
+
output_path = output_path.resolve()
|
| 420 |
+
output_path.parent.mkdir(parents=True, exist_ok=True)
|
| 421 |
+
|
| 422 |
+
with tempfile.TemporaryDirectory(prefix="codex-npm-pack-") as pack_dir_str:
|
| 423 |
+
pack_dir = Path(pack_dir_str)
|
| 424 |
+
npm_cache_dir = pack_dir / "npm-cache"
|
| 425 |
+
npm_logs_dir = pack_dir / "npm-logs"
|
| 426 |
+
npm_cache_dir.mkdir()
|
| 427 |
+
npm_logs_dir.mkdir()
|
| 428 |
+
env = os.environ.copy()
|
| 429 |
+
env["NPM_CONFIG_CACHE"] = str(npm_cache_dir)
|
| 430 |
+
env["NPM_CONFIG_LOGS_DIR"] = str(npm_logs_dir)
|
| 431 |
+
stdout = subprocess.check_output(
|
| 432 |
+
["npm", "pack", "--json", "--pack-destination", str(pack_dir)],
|
| 433 |
+
cwd=staging_dir,
|
| 434 |
+
env=env,
|
| 435 |
+
text=True,
|
| 436 |
+
)
|
| 437 |
+
try:
|
| 438 |
+
pack_output = json.loads(stdout)
|
| 439 |
+
except json.JSONDecodeError as exc:
|
| 440 |
+
raise RuntimeError("Failed to parse npm pack output.") from exc
|
| 441 |
+
|
| 442 |
+
if not pack_output:
|
| 443 |
+
raise RuntimeError("npm pack did not produce an output tarball.")
|
| 444 |
+
|
| 445 |
+
tarball_name = pack_output[0].get("filename") or pack_output[0].get("name")
|
| 446 |
+
if not tarball_name:
|
| 447 |
+
raise RuntimeError("Unable to determine npm pack output filename.")
|
| 448 |
+
|
| 449 |
+
tarball_path = pack_dir / tarball_name
|
| 450 |
+
if not tarball_path.exists():
|
| 451 |
+
raise RuntimeError(f"Expected npm pack output not found: {tarball_path}")
|
| 452 |
+
|
| 453 |
+
shutil.move(str(tarball_path), output_path)
|
| 454 |
+
|
| 455 |
+
return output_path
|
| 456 |
+
|
| 457 |
+
|
| 458 |
+
if __name__ == "__main__":
|
| 459 |
+
import sys
|
| 460 |
+
|
| 461 |
+
sys.exit(main())
|
codex-cli/scripts/init_firewall.sh
ADDED
|
@@ -0,0 +1,115 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
set -euo pipefail # Exit on error, undefined vars, and pipeline failures
|
| 3 |
+
IFS=$'\n\t' # Stricter word splitting
|
| 4 |
+
|
| 5 |
+
# Read allowed domains from file
|
| 6 |
+
ALLOWED_DOMAINS_FILE="/etc/codex/allowed_domains.txt"
|
| 7 |
+
if [ -f "$ALLOWED_DOMAINS_FILE" ]; then
|
| 8 |
+
ALLOWED_DOMAINS=()
|
| 9 |
+
while IFS= read -r domain; do
|
| 10 |
+
ALLOWED_DOMAINS+=("$domain")
|
| 11 |
+
done < "$ALLOWED_DOMAINS_FILE"
|
| 12 |
+
echo "Using domains from file: ${ALLOWED_DOMAINS[*]}"
|
| 13 |
+
else
|
| 14 |
+
# Fallback to default domains
|
| 15 |
+
ALLOWED_DOMAINS=("api.openai.com")
|
| 16 |
+
echo "Domains file not found, using default: ${ALLOWED_DOMAINS[*]}"
|
| 17 |
+
fi
|
| 18 |
+
|
| 19 |
+
# Ensure we have at least one domain
|
| 20 |
+
if [ ${#ALLOWED_DOMAINS[@]} -eq 0 ]; then
|
| 21 |
+
echo "ERROR: No allowed domains specified"
|
| 22 |
+
exit 1
|
| 23 |
+
fi
|
| 24 |
+
|
| 25 |
+
# Flush existing rules and delete existing ipsets
|
| 26 |
+
iptables -F
|
| 27 |
+
iptables -X
|
| 28 |
+
iptables -t nat -F
|
| 29 |
+
iptables -t nat -X
|
| 30 |
+
iptables -t mangle -F
|
| 31 |
+
iptables -t mangle -X
|
| 32 |
+
ipset destroy allowed-domains 2>/dev/null || true
|
| 33 |
+
|
| 34 |
+
# First allow DNS and localhost before any restrictions
|
| 35 |
+
# Allow outbound DNS
|
| 36 |
+
iptables -A OUTPUT -p udp --dport 53 -j ACCEPT
|
| 37 |
+
# Allow inbound DNS responses
|
| 38 |
+
iptables -A INPUT -p udp --sport 53 -j ACCEPT
|
| 39 |
+
# Allow localhost
|
| 40 |
+
iptables -A INPUT -i lo -j ACCEPT
|
| 41 |
+
iptables -A OUTPUT -o lo -j ACCEPT
|
| 42 |
+
|
| 43 |
+
# Create ipset with CIDR support
|
| 44 |
+
ipset create allowed-domains hash:net
|
| 45 |
+
|
| 46 |
+
# Resolve and add other allowed domains
|
| 47 |
+
for domain in "${ALLOWED_DOMAINS[@]}"; do
|
| 48 |
+
echo "Resolving $domain..."
|
| 49 |
+
ips=$(dig +short A "$domain")
|
| 50 |
+
if [ -z "$ips" ]; then
|
| 51 |
+
echo "ERROR: Failed to resolve $domain"
|
| 52 |
+
exit 1
|
| 53 |
+
fi
|
| 54 |
+
|
| 55 |
+
while read -r ip; do
|
| 56 |
+
if [[ ! "$ip" =~ ^[0-9]{1,3}\.[0-9]{1,3}\.[0-9]{1,3}\.[0-9]{1,3}$ ]]; then
|
| 57 |
+
echo "ERROR: Invalid IP from DNS for $domain: $ip"
|
| 58 |
+
exit 1
|
| 59 |
+
fi
|
| 60 |
+
echo "Adding $ip for $domain"
|
| 61 |
+
ipset add allowed-domains "$ip"
|
| 62 |
+
done < <(echo "$ips")
|
| 63 |
+
done
|
| 64 |
+
|
| 65 |
+
# Get host IP from default route
|
| 66 |
+
HOST_IP=$(ip route | grep default | cut -d" " -f3)
|
| 67 |
+
if [ -z "$HOST_IP" ]; then
|
| 68 |
+
echo "ERROR: Failed to detect host IP"
|
| 69 |
+
exit 1
|
| 70 |
+
fi
|
| 71 |
+
|
| 72 |
+
HOST_NETWORK=$(echo "$HOST_IP" | sed "s/\.[0-9]*$/.0\/24/")
|
| 73 |
+
echo "Host network detected as: $HOST_NETWORK"
|
| 74 |
+
|
| 75 |
+
# Set up remaining iptables rules
|
| 76 |
+
iptables -A INPUT -s "$HOST_NETWORK" -j ACCEPT
|
| 77 |
+
iptables -A OUTPUT -d "$HOST_NETWORK" -j ACCEPT
|
| 78 |
+
|
| 79 |
+
# Set default policies to DROP first
|
| 80 |
+
iptables -P INPUT DROP
|
| 81 |
+
iptables -P FORWARD DROP
|
| 82 |
+
iptables -P OUTPUT DROP
|
| 83 |
+
|
| 84 |
+
# First allow established connections for already approved traffic
|
| 85 |
+
iptables -A INPUT -m state --state ESTABLISHED,RELATED -j ACCEPT
|
| 86 |
+
iptables -A OUTPUT -m state --state ESTABLISHED,RELATED -j ACCEPT
|
| 87 |
+
|
| 88 |
+
# Then allow only specific outbound traffic to allowed domains
|
| 89 |
+
iptables -A OUTPUT -m set --match-set allowed-domains dst -j ACCEPT
|
| 90 |
+
|
| 91 |
+
# Append final REJECT rules for immediate error responses
|
| 92 |
+
# For TCP traffic, send a TCP reset; for UDP, send ICMP port unreachable.
|
| 93 |
+
iptables -A INPUT -p tcp -j REJECT --reject-with tcp-reset
|
| 94 |
+
iptables -A INPUT -p udp -j REJECT --reject-with icmp-port-unreachable
|
| 95 |
+
iptables -A OUTPUT -p tcp -j REJECT --reject-with tcp-reset
|
| 96 |
+
iptables -A OUTPUT -p udp -j REJECT --reject-with icmp-port-unreachable
|
| 97 |
+
iptables -A FORWARD -p tcp -j REJECT --reject-with tcp-reset
|
| 98 |
+
iptables -A FORWARD -p udp -j REJECT --reject-with icmp-port-unreachable
|
| 99 |
+
|
| 100 |
+
echo "Firewall configuration complete"
|
| 101 |
+
echo "Verifying firewall rules..."
|
| 102 |
+
if curl --connect-timeout 5 https://example.com >/dev/null 2>&1; then
|
| 103 |
+
echo "ERROR: Firewall verification failed - was able to reach https://example.com"
|
| 104 |
+
exit 1
|
| 105 |
+
else
|
| 106 |
+
echo "Firewall verification passed - unable to reach https://example.com as expected"
|
| 107 |
+
fi
|
| 108 |
+
|
| 109 |
+
# Always verify OpenAI API access is working
|
| 110 |
+
if ! curl --connect-timeout 5 https://api.openai.com >/dev/null 2>&1; then
|
| 111 |
+
echo "ERROR: Firewall verification failed - unable to reach https://api.openai.com"
|
| 112 |
+
exit 1
|
| 113 |
+
else
|
| 114 |
+
echo "Firewall verification passed - able to reach https://api.openai.com as expected"
|
| 115 |
+
fi
|
codex-cli/scripts/run_in_container.sh
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
set -e
|
| 3 |
+
|
| 4 |
+
# Usage:
|
| 5 |
+
# ./run_in_container.sh [--work_dir directory] "COMMAND"
|
| 6 |
+
#
|
| 7 |
+
# Examples:
|
| 8 |
+
# ./run_in_container.sh --work_dir project/code "ls -la"
|
| 9 |
+
# ./run_in_container.sh "echo Hello, world!"
|
| 10 |
+
|
| 11 |
+
# Default the work directory to WORKSPACE_ROOT_DIR if not provided.
|
| 12 |
+
WORK_DIR="${WORKSPACE_ROOT_DIR:-$(pwd)}"
|
| 13 |
+
# Default allowed domains - can be overridden with OPENAI_ALLOWED_DOMAINS env var
|
| 14 |
+
OPENAI_ALLOWED_DOMAINS="${OPENAI_ALLOWED_DOMAINS:-api.openai.com}"
|
| 15 |
+
|
| 16 |
+
# Parse optional flag.
|
| 17 |
+
if [ "$1" = "--work_dir" ]; then
|
| 18 |
+
if [ -z "$2" ]; then
|
| 19 |
+
echo "Error: --work_dir flag provided but no directory specified."
|
| 20 |
+
exit 1
|
| 21 |
+
fi
|
| 22 |
+
WORK_DIR="$2"
|
| 23 |
+
shift 2
|
| 24 |
+
fi
|
| 25 |
+
|
| 26 |
+
WORK_DIR=$(realpath "$WORK_DIR")
|
| 27 |
+
|
| 28 |
+
# Generate a unique container name based on the normalized work directory
|
| 29 |
+
CONTAINER_NAME="codex_$(echo "$WORK_DIR" | sed 's/\//_/g' | sed 's/[^a-zA-Z0-9_-]//g')"
|
| 30 |
+
|
| 31 |
+
# Define cleanup to remove the container on script exit, ensuring no leftover containers
|
| 32 |
+
cleanup() {
|
| 33 |
+
docker rm -f "$CONTAINER_NAME" >/dev/null 2>&1 || true
|
| 34 |
+
}
|
| 35 |
+
# Trap EXIT to invoke cleanup regardless of how the script terminates
|
| 36 |
+
trap cleanup EXIT
|
| 37 |
+
|
| 38 |
+
# Ensure a command is provided.
|
| 39 |
+
if [ "$#" -eq 0 ]; then
|
| 40 |
+
echo "Usage: $0 [--work_dir directory] \"COMMAND\""
|
| 41 |
+
exit 1
|
| 42 |
+
fi
|
| 43 |
+
|
| 44 |
+
# Check if WORK_DIR is set.
|
| 45 |
+
if [ -z "$WORK_DIR" ]; then
|
| 46 |
+
echo "Error: No work directory provided and WORKSPACE_ROOT_DIR is not set."
|
| 47 |
+
exit 1
|
| 48 |
+
fi
|
| 49 |
+
|
| 50 |
+
# Verify that OPENAI_ALLOWED_DOMAINS is not empty
|
| 51 |
+
if [ -z "$OPENAI_ALLOWED_DOMAINS" ]; then
|
| 52 |
+
echo "Error: OPENAI_ALLOWED_DOMAINS is empty."
|
| 53 |
+
exit 1
|
| 54 |
+
fi
|
| 55 |
+
|
| 56 |
+
# Kill any existing container for the working directory using cleanup(), centralizing removal logic.
|
| 57 |
+
cleanup
|
| 58 |
+
|
| 59 |
+
# Run the container with the specified directory mounted at the same path inside the container.
|
| 60 |
+
docker run --name "$CONTAINER_NAME" -d \
|
| 61 |
+
-e OPENAI_API_KEY \
|
| 62 |
+
--cap-add=NET_ADMIN \
|
| 63 |
+
--cap-add=NET_RAW \
|
| 64 |
+
-v "$WORK_DIR:/app$WORK_DIR" \
|
| 65 |
+
codex \
|
| 66 |
+
sleep infinity
|
| 67 |
+
|
| 68 |
+
# Write the allowed domains to a file in the container
|
| 69 |
+
docker exec --user root "$CONTAINER_NAME" bash -c "mkdir -p /etc/codex"
|
| 70 |
+
for domain in $OPENAI_ALLOWED_DOMAINS; do
|
| 71 |
+
# Validate domain format to prevent injection
|
| 72 |
+
if [[ ! "$domain" =~ ^[a-zA-Z0-9][a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$ ]]; then
|
| 73 |
+
echo "Error: Invalid domain format: $domain"
|
| 74 |
+
exit 1
|
| 75 |
+
fi
|
| 76 |
+
echo "$domain" | docker exec --user root -i "$CONTAINER_NAME" bash -c "cat >> /etc/codex/allowed_domains.txt"
|
| 77 |
+
done
|
| 78 |
+
|
| 79 |
+
# Set proper permissions on the domains file
|
| 80 |
+
docker exec --user root "$CONTAINER_NAME" bash -c "chmod 444 /etc/codex/allowed_domains.txt && chown root:root /etc/codex/allowed_domains.txt"
|
| 81 |
+
|
| 82 |
+
# Initialize the firewall inside the container as root user
|
| 83 |
+
docker exec --user root "$CONTAINER_NAME" bash -c "/usr/local/bin/init_firewall.sh"
|
| 84 |
+
|
| 85 |
+
# Remove the firewall script after running it
|
| 86 |
+
docker exec --user root "$CONTAINER_NAME" bash -c "rm -f /usr/local/bin/init_firewall.sh"
|
| 87 |
+
|
| 88 |
+
# Execute the provided command in the container, ensuring it runs in the work directory.
|
| 89 |
+
# We use a parameterized bash command to safely handle the command and directory.
|
| 90 |
+
|
| 91 |
+
quoted_args=""
|
| 92 |
+
for arg in "$@"; do
|
| 93 |
+
quoted_args+=" $(printf '%q' "$arg")"
|
| 94 |
+
done
|
| 95 |
+
docker exec -it "$CONTAINER_NAME" bash -c "cd \"/app$WORK_DIR\" && codex --sandbox workspace-write --ask-for-approval on-request ${quoted_args}"
|
codex-rs/agent-identity/src/lib.rs
ADDED
|
@@ -0,0 +1,1000 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::collections::BTreeMap;
|
| 2 |
+
use std::error::Error as StdError;
|
| 3 |
+
use std::fmt;
|
| 4 |
+
use std::time::Duration;
|
| 5 |
+
|
| 6 |
+
use anyhow::Context;
|
| 7 |
+
use anyhow::Result;
|
| 8 |
+
use base64::Engine as _;
|
| 9 |
+
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
|
| 10 |
+
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
| 11 |
+
use chrono::SecondsFormat;
|
| 12 |
+
use chrono::Utc;
|
| 13 |
+
use codex_http_client::HttpClient;
|
| 14 |
+
use codex_http_client::HttpError;
|
| 15 |
+
use codex_protocol::auth::PlanType as AuthPlanType;
|
| 16 |
+
use codex_protocol::protocol::SessionSource;
|
| 17 |
+
use crypto_box::SecretKey as Curve25519SecretKey;
|
| 18 |
+
use ed25519_dalek::Signer as _;
|
| 19 |
+
use ed25519_dalek::SigningKey;
|
| 20 |
+
use ed25519_dalek::VerifyingKey;
|
| 21 |
+
use ed25519_dalek::pkcs8::DecodePrivateKey;
|
| 22 |
+
use ed25519_dalek::pkcs8::EncodePrivateKey;
|
| 23 |
+
use http::StatusCode;
|
| 24 |
+
use jsonwebtoken::Algorithm;
|
| 25 |
+
use jsonwebtoken::DecodingKey;
|
| 26 |
+
use jsonwebtoken::Validation;
|
| 27 |
+
use jsonwebtoken::decode;
|
| 28 |
+
use jsonwebtoken::decode_header;
|
| 29 |
+
use jsonwebtoken::jwk::JwkSet;
|
| 30 |
+
use rand::TryRngCore;
|
| 31 |
+
use rand::rngs::OsRng;
|
| 32 |
+
use serde::Deserialize;
|
| 33 |
+
use serde::Serialize;
|
| 34 |
+
use serde::de::DeserializeOwned;
|
| 35 |
+
use sha2::Digest as _;
|
| 36 |
+
use sha2::Sha512;
|
| 37 |
+
|
| 38 |
+
const AGENT_TASK_REGISTRATION_TIMEOUT: Duration = Duration::from_secs(30);
|
| 39 |
+
const AGENT_IDENTITY_JWKS_TIMEOUT: Duration = Duration::from_secs(10);
|
| 40 |
+
const AGENT_IDENTITY_JWT_AUDIENCE: &str = "codex-app-server";
|
| 41 |
+
const AGENT_IDENTITY_JWT_ISSUER: &str = "https://chatgpt.com/codex-backend/agent-identity";
|
| 42 |
+
const AGENT_REGISTRATION_TIMEOUT: Duration = Duration::from_secs(15);
|
| 43 |
+
const PROD_AGENT_IDENTITY_AUTHAPI_BASE_URL: &str = "https://auth.openai.com/api/accounts";
|
| 44 |
+
const STAGING_AGENT_IDENTITY_AUTHAPI_BASE_URL: &str = "https://auth.api.openai.org/api/accounts";
|
| 45 |
+
const AGENT_IDENTITY_KEY_SEED_BYTES: usize = 64;
|
| 46 |
+
const AGENT_IDENTITY_KEY_DERIVATION_CONTEXT: &[u8] = b"codex-agent-identity-ed25519-v1";
|
| 47 |
+
|
| 48 |
+
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
|
| 49 |
+
pub enum ChatGptEnvironment {
|
| 50 |
+
#[default]
|
| 51 |
+
Production,
|
| 52 |
+
Staging,
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
impl ChatGptEnvironment {
|
| 56 |
+
pub fn from_chatgpt_base_url(chatgpt_base_url: &str) -> Result<Self> {
|
| 57 |
+
match chatgpt_base_url.trim_end_matches('/') {
|
| 58 |
+
"https://chatgpt.com"
|
| 59 |
+
| "https://chatgpt.com/backend-api"
|
| 60 |
+
| "https://chatgpt.com/codex"
|
| 61 |
+
| "https://chatgpt.com/backend-api/codex"
|
| 62 |
+
| "https://chat.openai.com"
|
| 63 |
+
| "https://chat.openai.com/backend-api"
|
| 64 |
+
| "https://chat.openai.com/codex"
|
| 65 |
+
| "https://chat.openai.com/backend-api/codex" => Ok(Self::Production),
|
| 66 |
+
"https://chatgpt-staging.com"
|
| 67 |
+
| "https://chatgpt-staging.com/backend-api"
|
| 68 |
+
| "https://chatgpt-staging.com/codex"
|
| 69 |
+
| "https://chatgpt-staging.com/backend-api/codex" => Ok(Self::Staging),
|
| 70 |
+
_ => anyhow::bail!(
|
| 71 |
+
"Agent Identity only supports production and staging ChatGPT environments"
|
| 72 |
+
),
|
| 73 |
+
}
|
| 74 |
+
}
|
| 75 |
+
|
| 76 |
+
pub fn chatgpt_base_url(self) -> &'static str {
|
| 77 |
+
match self {
|
| 78 |
+
Self::Production => "https://chatgpt.com/backend-api",
|
| 79 |
+
Self::Staging => "https://chatgpt-staging.com/backend-api",
|
| 80 |
+
}
|
| 81 |
+
}
|
| 82 |
+
|
| 83 |
+
pub fn agent_identity_authapi_base_url(self) -> &'static str {
|
| 84 |
+
match self {
|
| 85 |
+
Self::Production => PROD_AGENT_IDENTITY_AUTHAPI_BASE_URL,
|
| 86 |
+
Self::Staging => STAGING_AGENT_IDENTITY_AUTHAPI_BASE_URL,
|
| 87 |
+
}
|
| 88 |
+
}
|
| 89 |
+
}
|
| 90 |
+
|
| 91 |
+
/// Borrowed durable signing material for a registered agent identity.
|
| 92 |
+
///
|
| 93 |
+
/// This intentionally does not include a task id. Task ids are scoped to a
|
| 94 |
+
/// single Codex run, while the agent runtime id and private key are the
|
| 95 |
+
/// reusable identity material used to register and sign that run task.
|
| 96 |
+
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
| 97 |
+
pub struct AgentIdentityKey<'a> {
|
| 98 |
+
pub agent_runtime_id: &'a str,
|
| 99 |
+
pub private_key_pkcs8_base64: &'a str,
|
| 100 |
+
}
|
| 101 |
+
|
| 102 |
+
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
| 103 |
+
pub struct AgentBillOfMaterials {
|
| 104 |
+
pub agent_version: String,
|
| 105 |
+
pub agent_harness_id: String,
|
| 106 |
+
pub running_location: String,
|
| 107 |
+
}
|
| 108 |
+
|
| 109 |
+
pub struct GeneratedAgentKeyMaterial {
|
| 110 |
+
pub private_key_pkcs8_base64: String,
|
| 111 |
+
pub public_key_ssh: String,
|
| 112 |
+
}
|
| 113 |
+
|
| 114 |
+
/// Claims carried by an Agent Identity JWT.
|
| 115 |
+
#[derive(Clone, Debug, Deserialize, PartialEq, Eq)]
|
| 116 |
+
pub struct AgentIdentityJwtClaims {
|
| 117 |
+
pub iss: String,
|
| 118 |
+
pub aud: String,
|
| 119 |
+
pub iat: usize,
|
| 120 |
+
pub exp: usize,
|
| 121 |
+
pub agent_runtime_id: String,
|
| 122 |
+
pub agent_private_key: String,
|
| 123 |
+
pub account_id: String,
|
| 124 |
+
pub chatgpt_user_id: String,
|
| 125 |
+
pub email: Option<String>,
|
| 126 |
+
pub plan_type: AuthPlanType,
|
| 127 |
+
pub chatgpt_account_is_fedramp: bool,
|
| 128 |
+
}
|
| 129 |
+
|
| 130 |
+
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
| 131 |
+
struct AgentAssertionEnvelope {
|
| 132 |
+
agent_runtime_id: String,
|
| 133 |
+
task_id: String,
|
| 134 |
+
timestamp: String,
|
| 135 |
+
signature: String,
|
| 136 |
+
}
|
| 137 |
+
|
| 138 |
+
#[derive(Serialize)]
|
| 139 |
+
struct RegisterTaskRequest {
|
| 140 |
+
timestamp: String,
|
| 141 |
+
signature: String,
|
| 142 |
+
}
|
| 143 |
+
|
| 144 |
+
#[derive(Deserialize)]
|
| 145 |
+
struct RegisterTaskResponse {
|
| 146 |
+
#[serde(default)]
|
| 147 |
+
task_id: Option<String>,
|
| 148 |
+
#[serde(default, rename = "taskId")]
|
| 149 |
+
task_id_camel: Option<String>,
|
| 150 |
+
#[serde(default)]
|
| 151 |
+
encrypted_task_id: Option<String>,
|
| 152 |
+
#[serde(default, rename = "encryptedTaskId")]
|
| 153 |
+
encrypted_task_id_camel: Option<String>,
|
| 154 |
+
}
|
| 155 |
+
|
| 156 |
+
#[derive(Debug, Serialize)]
|
| 157 |
+
struct RegisterAgentRequest {
|
| 158 |
+
abom: AgentBillOfMaterials,
|
| 159 |
+
agent_public_key: String,
|
| 160 |
+
capabilities: Vec<String>,
|
| 161 |
+
ttl: Option<u64>,
|
| 162 |
+
}
|
| 163 |
+
|
| 164 |
+
#[derive(Debug, Deserialize)]
|
| 165 |
+
struct RegisterAgentResponse {
|
| 166 |
+
agent_runtime_id: String,
|
| 167 |
+
}
|
| 168 |
+
|
| 169 |
+
/// HTTP status failure returned by Agent Identity registration endpoints.
|
| 170 |
+
#[derive(Debug)]
|
| 171 |
+
pub struct AgentIdentityRegistrationHttpError {
|
| 172 |
+
operation: &'static str,
|
| 173 |
+
status: StatusCode,
|
| 174 |
+
body: String,
|
| 175 |
+
}
|
| 176 |
+
|
| 177 |
+
impl AgentIdentityRegistrationHttpError {
|
| 178 |
+
fn new(operation: &'static str, status: StatusCode, body: String) -> Self {
|
| 179 |
+
Self {
|
| 180 |
+
operation,
|
| 181 |
+
status,
|
| 182 |
+
body,
|
| 183 |
+
}
|
| 184 |
+
}
|
| 185 |
+
|
| 186 |
+
/// HTTP status returned by the registration endpoint.
|
| 187 |
+
pub fn status(&self) -> StatusCode {
|
| 188 |
+
self.status
|
| 189 |
+
}
|
| 190 |
+
}
|
| 191 |
+
|
| 192 |
+
impl fmt::Display for AgentIdentityRegistrationHttpError {
|
| 193 |
+
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 194 |
+
if self.body.is_empty() {
|
| 195 |
+
write!(f, "{} failed with status {}", self.operation, self.status)
|
| 196 |
+
} else {
|
| 197 |
+
write!(
|
| 198 |
+
f,
|
| 199 |
+
"{} failed with status {}: {}",
|
| 200 |
+
self.operation, self.status, self.body
|
| 201 |
+
)
|
| 202 |
+
}
|
| 203 |
+
}
|
| 204 |
+
}
|
| 205 |
+
|
| 206 |
+
impl StdError for AgentIdentityRegistrationHttpError {}
|
| 207 |
+
|
| 208 |
+
/// Returns whether an Agent Identity registration error is safe to retry.
|
| 209 |
+
pub fn is_retryable_registration_error(error: &anyhow::Error) -> bool {
|
| 210 |
+
error.chain().any(is_retryable_registration_cause)
|
| 211 |
+
}
|
| 212 |
+
|
| 213 |
+
fn is_retryable_registration_cause(cause: &(dyn StdError + 'static)) -> bool {
|
| 214 |
+
if let Some(error) = cause.downcast_ref::<AgentIdentityRegistrationHttpError>() {
|
| 215 |
+
return is_retryable_registration_status(error.status());
|
| 216 |
+
}
|
| 217 |
+
|
| 218 |
+
if let Some(error) = cause.downcast_ref::<HttpError>() {
|
| 219 |
+
if let Some(status) = error.status() {
|
| 220 |
+
return is_retryable_registration_status(status);
|
| 221 |
+
}
|
| 222 |
+
return error.is_timeout() || error.is_connect() || error.is_request();
|
| 223 |
+
}
|
| 224 |
+
|
| 225 |
+
false
|
| 226 |
+
}
|
| 227 |
+
|
| 228 |
+
fn is_retryable_registration_status(status: StatusCode) -> bool {
|
| 229 |
+
status == StatusCode::TOO_MANY_REQUESTS || status.is_server_error()
|
| 230 |
+
}
|
| 231 |
+
|
| 232 |
+
pub fn authorization_header_for_agent_task(
|
| 233 |
+
key: AgentIdentityKey<'_>,
|
| 234 |
+
task_id: &str,
|
| 235 |
+
) -> Result<String> {
|
| 236 |
+
let timestamp = Utc::now().to_rfc3339_opts(SecondsFormat::Secs, true);
|
| 237 |
+
let envelope = AgentAssertionEnvelope {
|
| 238 |
+
agent_runtime_id: key.agent_runtime_id.to_string(),
|
| 239 |
+
task_id: task_id.to_string(),
|
| 240 |
+
timestamp: timestamp.clone(),
|
| 241 |
+
signature: sign_agent_assertion_payload(key, task_id, ×tamp)?,
|
| 242 |
+
};
|
| 243 |
+
let serialized_assertion = serialize_agent_assertion(&envelope)?;
|
| 244 |
+
Ok(format!("AgentAssertion {serialized_assertion}"))
|
| 245 |
+
}
|
| 246 |
+
|
| 247 |
+
pub async fn fetch_agent_identity_jwks(
|
| 248 |
+
client: &HttpClient,
|
| 249 |
+
agent_identity_jwt_base_url: &str,
|
| 250 |
+
) -> Result<JwkSet> {
|
| 251 |
+
let response = client
|
| 252 |
+
.get(agent_identity_jwks_url(agent_identity_jwt_base_url))
|
| 253 |
+
.timeout(AGENT_IDENTITY_JWKS_TIMEOUT)
|
| 254 |
+
.send()
|
| 255 |
+
.await
|
| 256 |
+
.context("failed to request agent identity JWKS")?
|
| 257 |
+
.error_for_status()
|
| 258 |
+
.context("agent identity JWKS endpoint returned an error")?;
|
| 259 |
+
|
| 260 |
+
response
|
| 261 |
+
.json()
|
| 262 |
+
.await
|
| 263 |
+
.context("failed to decode agent identity JWKS")
|
| 264 |
+
}
|
| 265 |
+
|
| 266 |
+
pub fn decode_agent_identity_jwt(
|
| 267 |
+
jwt: &str,
|
| 268 |
+
jwks: Option<&JwkSet>,
|
| 269 |
+
) -> Result<AgentIdentityJwtClaims> {
|
| 270 |
+
let Some(jwks) = jwks else {
|
| 271 |
+
return decode_agent_identity_jwt_payload(jwt);
|
| 272 |
+
};
|
| 273 |
+
|
| 274 |
+
let header = decode_header(jwt).context("failed to decode agent identity JWT header")?;
|
| 275 |
+
let kid = header
|
| 276 |
+
.kid
|
| 277 |
+
.context("agent identity JWT header does not include a kid")?;
|
| 278 |
+
let jwk = jwks
|
| 279 |
+
.find(&kid)
|
| 280 |
+
.with_context(|| format!("agent identity JWT kid {kid} is not trusted"))?;
|
| 281 |
+
let decoding_key = DecodingKey::from_jwk(jwk).context("failed to build JWT decoding key")?;
|
| 282 |
+
let mut validation = Validation::new(Algorithm::RS256);
|
| 283 |
+
validation.set_audience(&[AGENT_IDENTITY_JWT_AUDIENCE]);
|
| 284 |
+
validation.set_issuer(&[AGENT_IDENTITY_JWT_ISSUER]);
|
| 285 |
+
validation.required_spec_claims.insert("iss".to_string());
|
| 286 |
+
validation.required_spec_claims.insert("aud".to_string());
|
| 287 |
+
decode::<AgentIdentityJwtClaims>(jwt, &decoding_key, &validation)
|
| 288 |
+
.map(|data| data.claims)
|
| 289 |
+
.context("failed to verify agent identity JWT")
|
| 290 |
+
}
|
| 291 |
+
|
| 292 |
+
fn decode_agent_identity_jwt_payload<T: DeserializeOwned>(jwt: &str) -> Result<T> {
|
| 293 |
+
let mut parts = jwt.split('.');
|
| 294 |
+
let (_header_b64, payload_b64, _sig_b64) = match (parts.next(), parts.next(), parts.next()) {
|
| 295 |
+
(Some(h), Some(p), Some(s)) if !h.is_empty() && !p.is_empty() && !s.is_empty() => (h, p, s),
|
| 296 |
+
_ => anyhow::bail!("invalid agent identity JWT format"),
|
| 297 |
+
};
|
| 298 |
+
anyhow::ensure!(parts.next().is_none(), "invalid agent identity JWT format");
|
| 299 |
+
|
| 300 |
+
let payload_bytes = URL_SAFE_NO_PAD
|
| 301 |
+
.decode(payload_b64)
|
| 302 |
+
.context("agent identity JWT payload is not valid base64url")?;
|
| 303 |
+
serde_json::from_slice(&payload_bytes).context("agent identity JWT payload is not valid JSON")
|
| 304 |
+
}
|
| 305 |
+
|
| 306 |
+
pub fn sign_task_registration_payload(
|
| 307 |
+
key: AgentIdentityKey<'_>,
|
| 308 |
+
timestamp: &str,
|
| 309 |
+
) -> Result<String> {
|
| 310 |
+
let signing_key = signing_key_from_private_key_pkcs8_base64(key.private_key_pkcs8_base64)?;
|
| 311 |
+
let payload = format!("{}:{timestamp}", key.agent_runtime_id);
|
| 312 |
+
Ok(BASE64_STANDARD.encode(signing_key.sign(payload.as_bytes()).to_bytes()))
|
| 313 |
+
}
|
| 314 |
+
|
| 315 |
+
pub async fn register_agent_task(
|
| 316 |
+
client: &HttpClient,
|
| 317 |
+
agent_identity_authapi_base_url: &str,
|
| 318 |
+
key: AgentIdentityKey<'_>,
|
| 319 |
+
) -> Result<String> {
|
| 320 |
+
let timestamp = Utc::now().to_rfc3339_opts(SecondsFormat::Secs, true);
|
| 321 |
+
let request = RegisterTaskRequest {
|
| 322 |
+
signature: sign_task_registration_payload(key, ×tamp)?,
|
| 323 |
+
timestamp,
|
| 324 |
+
};
|
| 325 |
+
let url = agent_task_registration_url(agent_identity_authapi_base_url, key.agent_runtime_id);
|
| 326 |
+
|
| 327 |
+
let response = client
|
| 328 |
+
.post(url)
|
| 329 |
+
.timeout(AGENT_TASK_REGISTRATION_TIMEOUT)
|
| 330 |
+
.json(&request)
|
| 331 |
+
.send()
|
| 332 |
+
.await
|
| 333 |
+
.context("failed to register agent task")?;
|
| 334 |
+
if !response.status().is_success() {
|
| 335 |
+
let status = response.status();
|
| 336 |
+
let body = response.text().await.unwrap_or_default();
|
| 337 |
+
let body = if body.len() > 512 {
|
| 338 |
+
format!("{}...", body.chars().take(512).collect::<String>())
|
| 339 |
+
} else {
|
| 340 |
+
body
|
| 341 |
+
};
|
| 342 |
+
return Err(AgentIdentityRegistrationHttpError::new(
|
| 343 |
+
"agent task registration",
|
| 344 |
+
status,
|
| 345 |
+
body,
|
| 346 |
+
)
|
| 347 |
+
.into());
|
| 348 |
+
}
|
| 349 |
+
|
| 350 |
+
let response = response
|
| 351 |
+
.json()
|
| 352 |
+
.await
|
| 353 |
+
.context("failed to decode agent task registration response")?;
|
| 354 |
+
|
| 355 |
+
task_id_from_register_task_response(key, response)
|
| 356 |
+
}
|
| 357 |
+
|
| 358 |
+
pub async fn register_agent_identity(
|
| 359 |
+
client: &HttpClient,
|
| 360 |
+
agent_identity_authapi_base_url: &str,
|
| 361 |
+
access_token: &str,
|
| 362 |
+
is_fedramp_account: bool,
|
| 363 |
+
key_material: &GeneratedAgentKeyMaterial,
|
| 364 |
+
abom: AgentBillOfMaterials,
|
| 365 |
+
capabilities: Vec<String>,
|
| 366 |
+
) -> Result<String> {
|
| 367 |
+
let url = agent_registration_url(agent_identity_authapi_base_url);
|
| 368 |
+
let request = RegisterAgentRequest {
|
| 369 |
+
abom,
|
| 370 |
+
agent_public_key: key_material.public_key_ssh.clone(),
|
| 371 |
+
capabilities,
|
| 372 |
+
ttl: None,
|
| 373 |
+
};
|
| 374 |
+
|
| 375 |
+
let mut request_builder = client
|
| 376 |
+
.post(&url)
|
| 377 |
+
.bearer_auth(access_token)
|
| 378 |
+
.json(&request)
|
| 379 |
+
.timeout(AGENT_REGISTRATION_TIMEOUT);
|
| 380 |
+
if is_fedramp_account {
|
| 381 |
+
request_builder = request_builder.header("X-OpenAI-Fedramp", "true");
|
| 382 |
+
}
|
| 383 |
+
|
| 384 |
+
let response = request_builder
|
| 385 |
+
.send()
|
| 386 |
+
.await
|
| 387 |
+
.with_context(|| format!("failed to send agent identity registration request to {url}"))?
|
| 388 |
+
.error_for_status()
|
| 389 |
+
.with_context(|| format!("agent identity registration failed for {url}"))?
|
| 390 |
+
.json::<RegisterAgentResponse>()
|
| 391 |
+
.await
|
| 392 |
+
.with_context(|| format!("failed to parse agent identity response from {url}"))?;
|
| 393 |
+
|
| 394 |
+
Ok(response.agent_runtime_id)
|
| 395 |
+
}
|
| 396 |
+
|
| 397 |
+
fn task_id_from_register_task_response(
|
| 398 |
+
key: AgentIdentityKey<'_>,
|
| 399 |
+
response: RegisterTaskResponse,
|
| 400 |
+
) -> Result<String> {
|
| 401 |
+
if let Some(task_id) = response.task_id.or(response.task_id_camel) {
|
| 402 |
+
return Ok(task_id);
|
| 403 |
+
}
|
| 404 |
+
let encrypted_task_id = response
|
| 405 |
+
.encrypted_task_id
|
| 406 |
+
.or(response.encrypted_task_id_camel)
|
| 407 |
+
.context("agent task registration response omitted task id")?;
|
| 408 |
+
decrypt_task_id_response(key, &encrypted_task_id)
|
| 409 |
+
}
|
| 410 |
+
|
| 411 |
+
pub fn decrypt_task_id_response(
|
| 412 |
+
key: AgentIdentityKey<'_>,
|
| 413 |
+
encrypted_task_id: &str,
|
| 414 |
+
) -> Result<String> {
|
| 415 |
+
let signing_key = signing_key_from_private_key_pkcs8_base64(key.private_key_pkcs8_base64)?;
|
| 416 |
+
let ciphertext = BASE64_STANDARD
|
| 417 |
+
.decode(encrypted_task_id)
|
| 418 |
+
.context("encrypted task id is not valid base64")?;
|
| 419 |
+
let plaintext = curve25519_secret_key_from_signing_key(&signing_key)
|
| 420 |
+
.unseal(&ciphertext)
|
| 421 |
+
.map_err(|_| anyhow::anyhow!("failed to decrypt encrypted task id"))?;
|
| 422 |
+
String::from_utf8(plaintext).context("decrypted task id is not valid UTF-8")
|
| 423 |
+
}
|
| 424 |
+
|
| 425 |
+
pub fn generate_agent_key_material() -> Result<GeneratedAgentKeyMaterial> {
|
| 426 |
+
let mut seed_material = [0u8; AGENT_IDENTITY_KEY_SEED_BYTES];
|
| 427 |
+
OsRng
|
| 428 |
+
.try_fill_bytes(&mut seed_material)
|
| 429 |
+
.context("failed to generate agent identity private key seed material")?;
|
| 430 |
+
// Ed25519 stores a 32-byte seed, so derive it from all sampled seed material.
|
| 431 |
+
let mut digest = Sha512::new();
|
| 432 |
+
digest.update(AGENT_IDENTITY_KEY_DERIVATION_CONTEXT);
|
| 433 |
+
digest.update(seed_material);
|
| 434 |
+
let digest = digest.finalize();
|
| 435 |
+
let mut secret_key_bytes = [0u8; 32];
|
| 436 |
+
secret_key_bytes.copy_from_slice(&digest[..32]);
|
| 437 |
+
let signing_key = SigningKey::from_bytes(&secret_key_bytes);
|
| 438 |
+
let private_key_pkcs8 = signing_key
|
| 439 |
+
.to_pkcs8_der()
|
| 440 |
+
.context("failed to encode agent identity private key as PKCS#8")?;
|
| 441 |
+
|
| 442 |
+
Ok(GeneratedAgentKeyMaterial {
|
| 443 |
+
private_key_pkcs8_base64: BASE64_STANDARD.encode(private_key_pkcs8.as_bytes()),
|
| 444 |
+
public_key_ssh: encode_ssh_ed25519_public_key(&signing_key.verifying_key()),
|
| 445 |
+
})
|
| 446 |
+
}
|
| 447 |
+
|
| 448 |
+
pub fn public_key_ssh_from_private_key_pkcs8_base64(
|
| 449 |
+
private_key_pkcs8_base64: &str,
|
| 450 |
+
) -> Result<String> {
|
| 451 |
+
let signing_key = signing_key_from_private_key_pkcs8_base64(private_key_pkcs8_base64)?;
|
| 452 |
+
Ok(encode_ssh_ed25519_public_key(&signing_key.verifying_key()))
|
| 453 |
+
}
|
| 454 |
+
|
| 455 |
+
pub fn verifying_key_from_private_key_pkcs8_base64(
|
| 456 |
+
private_key_pkcs8_base64: &str,
|
| 457 |
+
) -> Result<VerifyingKey> {
|
| 458 |
+
let signing_key = signing_key_from_private_key_pkcs8_base64(private_key_pkcs8_base64)?;
|
| 459 |
+
Ok(signing_key.verifying_key())
|
| 460 |
+
}
|
| 461 |
+
|
| 462 |
+
pub fn curve25519_secret_key_from_private_key_pkcs8_base64(
|
| 463 |
+
private_key_pkcs8_base64: &str,
|
| 464 |
+
) -> Result<Curve25519SecretKey> {
|
| 465 |
+
let signing_key = signing_key_from_private_key_pkcs8_base64(private_key_pkcs8_base64)?;
|
| 466 |
+
Ok(curve25519_secret_key_from_signing_key(&signing_key))
|
| 467 |
+
}
|
| 468 |
+
|
| 469 |
+
pub fn agent_registration_url(agent_identity_authapi_base_url: &str) -> String {
|
| 470 |
+
agent_identity_authapi_url(agent_identity_authapi_base_url, "/v1/agent/register")
|
| 471 |
+
}
|
| 472 |
+
|
| 473 |
+
pub fn agent_task_registration_url(
|
| 474 |
+
agent_identity_authapi_base_url: &str,
|
| 475 |
+
agent_runtime_id: &str,
|
| 476 |
+
) -> String {
|
| 477 |
+
agent_identity_authapi_url(
|
| 478 |
+
agent_identity_authapi_base_url,
|
| 479 |
+
&format!("/v1/agent/{agent_runtime_id}/task/register"),
|
| 480 |
+
)
|
| 481 |
+
}
|
| 482 |
+
|
| 483 |
+
pub fn agent_identity_jwks_url(agent_identity_jwt_base_url: &str) -> String {
|
| 484 |
+
let trimmed = agent_identity_jwt_base_url.trim_end_matches('/');
|
| 485 |
+
if trimmed.contains("/backend-api") {
|
| 486 |
+
format!("{trimmed}/wham/agent-identities/jwks")
|
| 487 |
+
} else {
|
| 488 |
+
format!("{trimmed}/agent-identities/jwks")
|
| 489 |
+
}
|
| 490 |
+
}
|
| 491 |
+
|
| 492 |
+
fn agent_identity_authapi_url(agent_identity_authapi_base_url: &str, api_path: &str) -> String {
|
| 493 |
+
let base_url = agent_identity_authapi_base_url.trim_end_matches('/');
|
| 494 |
+
format!("{base_url}{api_path}")
|
| 495 |
+
}
|
| 496 |
+
|
| 497 |
+
pub fn build_abom(session_source: SessionSource) -> AgentBillOfMaterials {
|
| 498 |
+
AgentBillOfMaterials {
|
| 499 |
+
agent_version: env!("CARGO_PKG_VERSION").to_string(),
|
| 500 |
+
agent_harness_id: match &session_source {
|
| 501 |
+
SessionSource::VSCode => "codex-app".to_string(),
|
| 502 |
+
SessionSource::Cli
|
| 503 |
+
| SessionSource::Exec
|
| 504 |
+
| SessionSource::Mcp
|
| 505 |
+
| SessionSource::Custom(_)
|
| 506 |
+
| SessionSource::Internal(_)
|
| 507 |
+
| SessionSource::SubAgent(_)
|
| 508 |
+
| SessionSource::Unknown => "codex-cli".to_string(),
|
| 509 |
+
},
|
| 510 |
+
running_location: format!("{}-{}", session_source, std::env::consts::OS),
|
| 511 |
+
}
|
| 512 |
+
}
|
| 513 |
+
|
| 514 |
+
pub fn encode_ssh_ed25519_public_key(verifying_key: &VerifyingKey) -> String {
|
| 515 |
+
let mut blob = Vec::with_capacity(4 + 11 + 4 + 32);
|
| 516 |
+
append_ssh_string(&mut blob, b"ssh-ed25519");
|
| 517 |
+
append_ssh_string(&mut blob, verifying_key.as_bytes());
|
| 518 |
+
format!("ssh-ed25519 {}", BASE64_STANDARD.encode(blob))
|
| 519 |
+
}
|
| 520 |
+
|
| 521 |
+
fn sign_agent_assertion_payload(
|
| 522 |
+
key: AgentIdentityKey<'_>,
|
| 523 |
+
task_id: &str,
|
| 524 |
+
timestamp: &str,
|
| 525 |
+
) -> Result<String> {
|
| 526 |
+
let signing_key = signing_key_from_private_key_pkcs8_base64(key.private_key_pkcs8_base64)?;
|
| 527 |
+
let payload = format!("{}:{task_id}:{timestamp}", key.agent_runtime_id);
|
| 528 |
+
Ok(BASE64_STANDARD.encode(signing_key.sign(payload.as_bytes()).to_bytes()))
|
| 529 |
+
}
|
| 530 |
+
|
| 531 |
+
fn serialize_agent_assertion(envelope: &AgentAssertionEnvelope) -> Result<String> {
|
| 532 |
+
let payload = serde_json::to_vec(&BTreeMap::from([
|
| 533 |
+
("agent_runtime_id", envelope.agent_runtime_id.as_str()),
|
| 534 |
+
("signature", envelope.signature.as_str()),
|
| 535 |
+
("task_id", envelope.task_id.as_str()),
|
| 536 |
+
("timestamp", envelope.timestamp.as_str()),
|
| 537 |
+
]))
|
| 538 |
+
.context("failed to serialize agent assertion envelope")?;
|
| 539 |
+
Ok(URL_SAFE_NO_PAD.encode(payload))
|
| 540 |
+
}
|
| 541 |
+
|
| 542 |
+
fn curve25519_secret_key_from_signing_key(signing_key: &SigningKey) -> Curve25519SecretKey {
|
| 543 |
+
let digest = Sha512::digest(signing_key.to_bytes());
|
| 544 |
+
let mut secret_key = [0u8; 32];
|
| 545 |
+
secret_key.copy_from_slice(&digest[..32]);
|
| 546 |
+
secret_key[0] &= 248;
|
| 547 |
+
secret_key[31] &= 127;
|
| 548 |
+
secret_key[31] |= 64;
|
| 549 |
+
Curve25519SecretKey::from(secret_key)
|
| 550 |
+
}
|
| 551 |
+
|
| 552 |
+
fn append_ssh_string(buf: &mut Vec<u8>, value: &[u8]) {
|
| 553 |
+
buf.extend_from_slice(&(value.len() as u32).to_be_bytes());
|
| 554 |
+
buf.extend_from_slice(value);
|
| 555 |
+
}
|
| 556 |
+
|
| 557 |
+
fn signing_key_from_private_key_pkcs8_base64(private_key_pkcs8_base64: &str) -> Result<SigningKey> {
|
| 558 |
+
let private_key = BASE64_STANDARD
|
| 559 |
+
.decode(private_key_pkcs8_base64)
|
| 560 |
+
.context("stored agent identity private key is not valid base64")?;
|
| 561 |
+
SigningKey::from_pkcs8_der(&private_key)
|
| 562 |
+
.context("stored agent identity private key is not valid PKCS#8")
|
| 563 |
+
}
|
| 564 |
+
|
| 565 |
+
#[cfg(test)]
|
| 566 |
+
mod tests {
|
| 567 |
+
use base64::Engine as _;
|
| 568 |
+
use ed25519_dalek::Signature;
|
| 569 |
+
use ed25519_dalek::Verifier as _;
|
| 570 |
+
use jsonwebtoken::EncodingKey;
|
| 571 |
+
use jsonwebtoken::Header;
|
| 572 |
+
use pretty_assertions::assert_eq;
|
| 573 |
+
|
| 574 |
+
use codex_protocol::auth::KnownPlan;
|
| 575 |
+
|
| 576 |
+
use super::*;
|
| 577 |
+
|
| 578 |
+
#[test]
|
| 579 |
+
fn register_task_request_uses_single_run_task_shape() {
|
| 580 |
+
let request = RegisterTaskRequest {
|
| 581 |
+
timestamp: "2026-04-23T00:00:00Z".to_string(),
|
| 582 |
+
signature: "signature".to_string(),
|
| 583 |
+
};
|
| 584 |
+
|
| 585 |
+
let serialized = serde_json::to_value(request).expect("serialize request");
|
| 586 |
+
|
| 587 |
+
assert_eq!(
|
| 588 |
+
serialized,
|
| 589 |
+
serde_json::json!({
|
| 590 |
+
"timestamp": "2026-04-23T00:00:00Z",
|
| 591 |
+
"signature": "signature",
|
| 592 |
+
})
|
| 593 |
+
);
|
| 594 |
+
}
|
| 595 |
+
|
| 596 |
+
#[test]
|
| 597 |
+
fn authorization_header_for_agent_task_serializes_signed_agent_assertion() {
|
| 598 |
+
let signing_key = SigningKey::from_bytes(&[7u8; 32]);
|
| 599 |
+
let private_key = signing_key
|
| 600 |
+
.to_pkcs8_der()
|
| 601 |
+
.expect("encode test key material");
|
| 602 |
+
let key = AgentIdentityKey {
|
| 603 |
+
agent_runtime_id: "agent-123",
|
| 604 |
+
private_key_pkcs8_base64: &BASE64_STANDARD.encode(private_key.as_bytes()),
|
| 605 |
+
};
|
| 606 |
+
|
| 607 |
+
let header = authorization_header_for_agent_task(key, "task-123")
|
| 608 |
+
.expect("build agent assertion header");
|
| 609 |
+
let token = header
|
| 610 |
+
.strip_prefix("AgentAssertion ")
|
| 611 |
+
.expect("agent assertion scheme");
|
| 612 |
+
let payload = URL_SAFE_NO_PAD
|
| 613 |
+
.decode(token)
|
| 614 |
+
.expect("valid base64url payload");
|
| 615 |
+
let envelope: AgentAssertionEnvelope =
|
| 616 |
+
serde_json::from_slice(&payload).expect("valid assertion envelope");
|
| 617 |
+
|
| 618 |
+
assert_eq!(
|
| 619 |
+
envelope,
|
| 620 |
+
AgentAssertionEnvelope {
|
| 621 |
+
agent_runtime_id: "agent-123".to_string(),
|
| 622 |
+
task_id: "task-123".to_string(),
|
| 623 |
+
timestamp: envelope.timestamp.clone(),
|
| 624 |
+
signature: envelope.signature.clone(),
|
| 625 |
+
}
|
| 626 |
+
);
|
| 627 |
+
let signature_bytes = BASE64_STANDARD
|
| 628 |
+
.decode(&envelope.signature)
|
| 629 |
+
.expect("valid base64 signature");
|
| 630 |
+
let signature = Signature::from_slice(&signature_bytes).expect("valid signature bytes");
|
| 631 |
+
signing_key
|
| 632 |
+
.verifying_key()
|
| 633 |
+
.verify(
|
| 634 |
+
format!(
|
| 635 |
+
"{}:{}:{}",
|
| 636 |
+
envelope.agent_runtime_id, envelope.task_id, envelope.timestamp
|
| 637 |
+
)
|
| 638 |
+
.as_bytes(),
|
| 639 |
+
&signature,
|
| 640 |
+
)
|
| 641 |
+
.expect("signature should verify");
|
| 642 |
+
}
|
| 643 |
+
|
| 644 |
+
#[test]
|
| 645 |
+
fn decode_agent_identity_jwt_reads_claims() {
|
| 646 |
+
let jwt = jwt_with_payload(serde_json::json!({
|
| 647 |
+
"iss": AGENT_IDENTITY_JWT_ISSUER,
|
| 648 |
+
"aud": AGENT_IDENTITY_JWT_AUDIENCE,
|
| 649 |
+
"iat": 1_700_000_000usize,
|
| 650 |
+
"exp": 4_000_000_000usize,
|
| 651 |
+
"agent_runtime_id": "agent-runtime-id",
|
| 652 |
+
"agent_private_key": "private-key",
|
| 653 |
+
"account_id": "account-id",
|
| 654 |
+
"chatgpt_user_id": "user-id",
|
| 655 |
+
"email": "user@example.com",
|
| 656 |
+
"plan_type": "pro",
|
| 657 |
+
"chatgpt_account_is_fedramp": false,
|
| 658 |
+
}));
|
| 659 |
+
|
| 660 |
+
let claims = decode_agent_identity_jwt(&jwt, /*jwks*/ None).expect("JWT should decode");
|
| 661 |
+
|
| 662 |
+
assert_eq!(
|
| 663 |
+
claims,
|
| 664 |
+
AgentIdentityJwtClaims {
|
| 665 |
+
iss: AGENT_IDENTITY_JWT_ISSUER.to_string(),
|
| 666 |
+
aud: AGENT_IDENTITY_JWT_AUDIENCE.to_string(),
|
| 667 |
+
iat: 1_700_000_000,
|
| 668 |
+
exp: 4_000_000_000,
|
| 669 |
+
agent_runtime_id: "agent-runtime-id".to_string(),
|
| 670 |
+
agent_private_key: "private-key".to_string(),
|
| 671 |
+
account_id: "account-id".to_string(),
|
| 672 |
+
chatgpt_user_id: "user-id".to_string(),
|
| 673 |
+
email: Some("user@example.com".to_string()),
|
| 674 |
+
plan_type: AuthPlanType::Known(KnownPlan::Pro),
|
| 675 |
+
chatgpt_account_is_fedramp: false,
|
| 676 |
+
}
|
| 677 |
+
);
|
| 678 |
+
}
|
| 679 |
+
|
| 680 |
+
#[test]
|
| 681 |
+
fn decode_agent_identity_jwt_accepts_missing_email() {
|
| 682 |
+
let jwt = jwt_with_payload(serde_json::json!({
|
| 683 |
+
"iss": AGENT_IDENTITY_JWT_ISSUER,
|
| 684 |
+
"aud": AGENT_IDENTITY_JWT_AUDIENCE,
|
| 685 |
+
"iat": 1_700_000_000usize,
|
| 686 |
+
"exp": 4_000_000_000usize,
|
| 687 |
+
"agent_runtime_id": "agent-runtime-id",
|
| 688 |
+
"agent_private_key": "private-key",
|
| 689 |
+
"account_id": "account-id",
|
| 690 |
+
"chatgpt_user_id": "user-id",
|
| 691 |
+
"plan_type": "pro",
|
| 692 |
+
"chatgpt_account_is_fedramp": false,
|
| 693 |
+
}));
|
| 694 |
+
|
| 695 |
+
let claims = decode_agent_identity_jwt(&jwt, /*jwks*/ None).expect("JWT should decode");
|
| 696 |
+
|
| 697 |
+
assert_eq!(claims.email, None);
|
| 698 |
+
}
|
| 699 |
+
|
| 700 |
+
#[test]
|
| 701 |
+
fn decode_agent_identity_jwt_maps_raw_plan_aliases() {
|
| 702 |
+
let jwt = jwt_with_payload(serde_json::json!({
|
| 703 |
+
"iss": AGENT_IDENTITY_JWT_ISSUER,
|
| 704 |
+
"aud": AGENT_IDENTITY_JWT_AUDIENCE,
|
| 705 |
+
"iat": 1_700_000_000usize,
|
| 706 |
+
"exp": 4_000_000_000usize,
|
| 707 |
+
"agent_runtime_id": "agent-runtime-id",
|
| 708 |
+
"agent_private_key": "private-key",
|
| 709 |
+
"account_id": "account-id",
|
| 710 |
+
"chatgpt_user_id": "user-id",
|
| 711 |
+
"email": "user@example.com",
|
| 712 |
+
"plan_type": "hc",
|
| 713 |
+
"chatgpt_account_is_fedramp": false,
|
| 714 |
+
}));
|
| 715 |
+
|
| 716 |
+
let claims = decode_agent_identity_jwt(&jwt, /*jwks*/ None).expect("JWT should decode");
|
| 717 |
+
|
| 718 |
+
assert_eq!(claims.plan_type, AuthPlanType::Known(KnownPlan::Enterprise));
|
| 719 |
+
}
|
| 720 |
+
|
| 721 |
+
#[test]
|
| 722 |
+
fn decode_agent_identity_jwt_verifies_when_jwks_is_present() {
|
| 723 |
+
let jwks = test_jwks("test-key");
|
| 724 |
+
let claims = AgentIdentityJwtClaims {
|
| 725 |
+
iss: AGENT_IDENTITY_JWT_ISSUER.to_string(),
|
| 726 |
+
aud: AGENT_IDENTITY_JWT_AUDIENCE.to_string(),
|
| 727 |
+
iat: 1_700_000_000,
|
| 728 |
+
exp: 4_000_000_000,
|
| 729 |
+
agent_runtime_id: "agent-runtime-id".to_string(),
|
| 730 |
+
agent_private_key: "private-key".to_string(),
|
| 731 |
+
account_id: "account-id".to_string(),
|
| 732 |
+
chatgpt_user_id: "user-id".to_string(),
|
| 733 |
+
email: Some("user@example.com".to_string()),
|
| 734 |
+
plan_type: AuthPlanType::Known(KnownPlan::Pro),
|
| 735 |
+
chatgpt_account_is_fedramp: false,
|
| 736 |
+
};
|
| 737 |
+
let jwt = jsonwebtoken::encode(
|
| 738 |
+
&test_jwt_header("test-key"),
|
| 739 |
+
&serde_json::json!({
|
| 740 |
+
"iss": claims.iss,
|
| 741 |
+
"aud": claims.aud,
|
| 742 |
+
"iat": claims.iat,
|
| 743 |
+
"exp": claims.exp,
|
| 744 |
+
"agent_runtime_id": claims.agent_runtime_id,
|
| 745 |
+
"agent_private_key": claims.agent_private_key,
|
| 746 |
+
"account_id": claims.account_id,
|
| 747 |
+
"chatgpt_user_id": claims.chatgpt_user_id,
|
| 748 |
+
"email": claims.email,
|
| 749 |
+
"plan_type": "pro",
|
| 750 |
+
"chatgpt_account_is_fedramp": claims.chatgpt_account_is_fedramp,
|
| 751 |
+
}),
|
| 752 |
+
&test_rsa_encoding_key(),
|
| 753 |
+
)
|
| 754 |
+
.expect("JWT should encode");
|
| 755 |
+
|
| 756 |
+
let expected_claims = AgentIdentityJwtClaims {
|
| 757 |
+
iss: AGENT_IDENTITY_JWT_ISSUER.to_string(),
|
| 758 |
+
aud: AGENT_IDENTITY_JWT_AUDIENCE.to_string(),
|
| 759 |
+
iat: 1_700_000_000,
|
| 760 |
+
exp: 4_000_000_000,
|
| 761 |
+
agent_runtime_id: "agent-runtime-id".to_string(),
|
| 762 |
+
agent_private_key: "private-key".to_string(),
|
| 763 |
+
account_id: "account-id".to_string(),
|
| 764 |
+
chatgpt_user_id: "user-id".to_string(),
|
| 765 |
+
email: Some("user@example.com".to_string()),
|
| 766 |
+
plan_type: AuthPlanType::Known(KnownPlan::Pro),
|
| 767 |
+
chatgpt_account_is_fedramp: false,
|
| 768 |
+
};
|
| 769 |
+
assert_eq!(
|
| 770 |
+
decode_agent_identity_jwt(&jwt, Some(&jwks)).expect("JWT should verify"),
|
| 771 |
+
expected_claims
|
| 772 |
+
);
|
| 773 |
+
}
|
| 774 |
+
|
| 775 |
+
#[test]
|
| 776 |
+
fn decode_agent_identity_jwt_rejects_untrusted_kid() {
|
| 777 |
+
let jwks = test_jwks("other-key");
|
| 778 |
+
|
| 779 |
+
let jwt = jsonwebtoken::encode(
|
| 780 |
+
&test_jwt_header("test-key"),
|
| 781 |
+
&serde_json::json!({
|
| 782 |
+
"iss": AGENT_IDENTITY_JWT_ISSUER,
|
| 783 |
+
"aud": AGENT_IDENTITY_JWT_AUDIENCE,
|
| 784 |
+
"iat": 1_700_000_000,
|
| 785 |
+
"exp": 4_000_000_000usize,
|
| 786 |
+
"agent_runtime_id": "agent-runtime-id",
|
| 787 |
+
"agent_private_key": "private-key",
|
| 788 |
+
"account_id": "account-id",
|
| 789 |
+
"chatgpt_user_id": "user-id",
|
| 790 |
+
"email": "user@example.com",
|
| 791 |
+
"plan_type": "pro",
|
| 792 |
+
"chatgpt_account_is_fedramp": false,
|
| 793 |
+
}),
|
| 794 |
+
&test_rsa_encoding_key(),
|
| 795 |
+
)
|
| 796 |
+
.expect("JWT should encode");
|
| 797 |
+
|
| 798 |
+
decode_agent_identity_jwt(&jwt, Some(&jwks)).expect_err("JWT should not verify");
|
| 799 |
+
}
|
| 800 |
+
|
| 801 |
+
#[test]
|
| 802 |
+
fn decode_agent_identity_jwt_requires_issuer_and_audience() {
|
| 803 |
+
let jwks = test_jwks("test-key");
|
| 804 |
+
let jwt = jsonwebtoken::encode(
|
| 805 |
+
&test_jwt_header("test-key"),
|
| 806 |
+
&serde_json::json!({
|
| 807 |
+
"iat": 1_700_000_000,
|
| 808 |
+
"exp": 4_000_000_000usize,
|
| 809 |
+
"agent_runtime_id": "agent-runtime-id",
|
| 810 |
+
"agent_private_key": "private-key",
|
| 811 |
+
"account_id": "account-id",
|
| 812 |
+
"chatgpt_user_id": "user-id",
|
| 813 |
+
"email": "user@example.com",
|
| 814 |
+
"plan_type": "pro",
|
| 815 |
+
"chatgpt_account_is_fedramp": false,
|
| 816 |
+
}),
|
| 817 |
+
&test_rsa_encoding_key(),
|
| 818 |
+
)
|
| 819 |
+
.expect("JWT should encode");
|
| 820 |
+
|
| 821 |
+
decode_agent_identity_jwt(&jwt, Some(&jwks)).expect_err("JWT should not verify");
|
| 822 |
+
}
|
| 823 |
+
|
| 824 |
+
fn test_jwt_header(kid: &str) -> Header {
|
| 825 |
+
let mut header = Header::new(Algorithm::RS256);
|
| 826 |
+
header.kid = Some(kid.to_string());
|
| 827 |
+
header
|
| 828 |
+
}
|
| 829 |
+
|
| 830 |
+
fn test_rsa_encoding_key() -> EncodingKey {
|
| 831 |
+
EncodingKey::from_rsa_pem(
|
| 832 |
+
br#"-----BEGIN PRIVATE KEY-----
|
| 833 |
+
MIIEvgIBADANBgkqhkiG9w0BAQEFAASCBKgwggSkAgEAAoIBAQDWpAXYypOsYAwO
|
| 834 |
+
bvBduMk/mxaoYDze0AZSzaSzLuIlcsl2EKDgC3AabhIWXh/qTGEJLOU3VB1e5mO9
|
| 835 |
+
FPbBlmIZSL3FQTbyt/hYutPFKfCou5PLmScw/TzILS3/RhT8UY9kxxZvXiEbTki9
|
| 836 |
+
mvxRuZFpVqDFJHwfitIjKZGhXDCYVKurPTrxetYZJg0h8sQBLKjkZ0BqqaTUkAsg
|
| 837 |
+
0eBgZAlXEzG3By8PGhUqYLt6W1Q3KYw0FmGy/gTyzH1g0ukGgSJvOd8SkNT8MbOs
|
| 838 |
+
zl5kKxDNqpuEE6UZ3jbuJ+5382d31w+rOAJRzbf7QVdI9+luCSwJcDACYPQ4WNBa
|
| 839 |
+
uCpV0ovpAgMBAAECggEAVu84LwZdqYN9XpswX8VoPYrjMm9IODapWQBRpQFoNyK2
|
| 840 |
+
1ksF3bjEPvA2Azk8U/l7k+vLKw22l6lY3EyRZPcz5GnB8xLm3ogE3mtNOp4yCyVu
|
| 841 |
+
RxhQ91aaN7mU17/a4BdorLi2LYVCg3zBmYociD1Q2AluNGsCmwPu+K7tfR2J0Sg8
|
| 842 |
+
NjqiTbDG1XDpR/icwgC9t6vh8lZpCHDhF4tbQfLLVLeA/OdcuzXDyMCXbmdVIdBQ
|
| 843 |
+
rm4aIFmr2e1/2ctTbCg85S6AGFTH+pSLjrwTzyvf+F6NW5uNjLQAQLFj+EznBDxj
|
| 844 |
+
Xdx90cySrjsKK6PVWQF4RiTvkSW8eWL7R6B2FZbGwQKBgQDuVQRj72hWloR7mbEL
|
| 845 |
+
aUEEv3pIXTMXWEsoMBNczos/1L1RnAN1AI44TurznasPZAWvQj+kVbLDR+TAeZrL
|
| 846 |
+
iA8HIWswQUI18hFmgKzSkwIXGtubcKVrgsKeS4lMDKCM/Ef6WAYdeq6ronoY5lCN
|
| 847 |
+
YrJFmGp81W5zcV7lyiycgbSiGwKBgQDmjWYf6pZjrK7Z+OJ3X1AZfi2vss15SCvL
|
| 848 |
+
3fPgzIDbViztpGyQhc3DQZIsBNIu0xZp/veGce9TEeTds2ro9NfdJFeou8+fC7Pq
|
| 849 |
+
sOsM3amGFFi+ZW/9BWyjZEM88bgWWAjqLHbpfHDxjAf5CSxddqxgHlbP0Ytyb1Vg
|
| 850 |
+
gmPDn9YKSwKBgQDbTi3hC35WFuDHn0/zcSHcDZmnFuOZeqyFyV83yfMGhGrEuqvP
|
| 851 |
+
sPgtRikajJ3IZsB4WZyYSidZXEFY/0z6NjOl2xF38MTNQPbT/FmK1q1Yt2UWrlv5
|
| 852 |
+
BvSwlk87RG9D7C0LZo4R+D7cPoDdgqjiwMvMEIkEX5zn641oI1ZTmWKuuwKBgQCD
|
| 853 |
+
KF+3unnRvHRAVoFnTZbA2fJdqMeRvogD04GhGlYX8V9f1hFY6nXTJaNlXVzA/J8c
|
| 854 |
+
r8ra9kgjJuPfZ+ljG58OFFW2DRohLcQtuHYPfK6rMzoFHqnl9EcIcMp7ijuionR3
|
| 855 |
+
29HOJFgQYgxLFXfit9d6WugiE+BTupiEbckZif13HwKBgE/lAlkVHP6YahOO2Ljc
|
| 856 |
+
J1bwkqKZTB5dHolX9A58e/xXnfZ5P8f3Z83+Izap3FwqQulk7b1WO1MQcHuVg2NN
|
| 857 |
+
5da0D4h2rYOXnbYIg0BVu4spQbaM6ewsp66b8+MzLOBvj8SzWdt1Oyw0q/MRyQAR
|
| 858 |
+
8U4M2TSWCKUY/A6sT4W8+mT9
|
| 859 |
+
-----END PRIVATE KEY-----"#,
|
| 860 |
+
)
|
| 861 |
+
.expect("test RSA key should parse")
|
| 862 |
+
}
|
| 863 |
+
|
| 864 |
+
fn test_jwks(kid: &str) -> jsonwebtoken::jwk::JwkSet {
|
| 865 |
+
serde_json::from_value(serde_json::json!({
|
| 866 |
+
"keys": [{
|
| 867 |
+
"kty": "RSA",
|
| 868 |
+
"kid": kid,
|
| 869 |
+
"use": "sig",
|
| 870 |
+
"alg": "RS256",
|
| 871 |
+
"n": "1qQF2MqTrGAMDm7wXbjJP5sWqGA83tAGUs2ksy7iJXLJdhCg4AtwGm4SFl4f6kxhCSzlN1QdXuZjvRT2wZZiGUi9xUE28rf4WLrTxSnwqLuTy5knMP08yC0t_0YU_FGPZMcWb14hG05IvZr8UbmRaVagxSR8H4rSIymRoVwwmFSrqz068XrWGSYNIfLEASyo5GdAaqmk1JALINHgYGQJVxMxtwcvDxoVKmC7eltUNymMNBZhsv4E8sx9YNLpBoEibznfEpDU_DGzrM5eZCsQzaqbhBOlGd427ifud_Nnd9cPqzgCUc23-0FXSPfpbgksCXAwAmD0OFjQWrgqVdKL6Q",
|
| 872 |
+
"e": "AQAB",
|
| 873 |
+
}]
|
| 874 |
+
}))
|
| 875 |
+
.expect("test JWKS should parse")
|
| 876 |
+
}
|
| 877 |
+
|
| 878 |
+
#[test]
|
| 879 |
+
fn chatgpt_environment_maps_known_urls_to_authapi() -> anyhow::Result<()> {
|
| 880 |
+
assert_eq!(
|
| 881 |
+
ChatGptEnvironment::from_chatgpt_base_url("https://chatgpt.com/backend-api/codex")?,
|
| 882 |
+
ChatGptEnvironment::Production
|
| 883 |
+
);
|
| 884 |
+
assert_eq!(
|
| 885 |
+
ChatGptEnvironment::Production.agent_identity_authapi_base_url(),
|
| 886 |
+
"https://auth.openai.com/api/accounts"
|
| 887 |
+
);
|
| 888 |
+
assert_eq!(
|
| 889 |
+
ChatGptEnvironment::from_chatgpt_base_url("https://chatgpt-staging.com/backend-api")?,
|
| 890 |
+
ChatGptEnvironment::Staging
|
| 891 |
+
);
|
| 892 |
+
assert_eq!(
|
| 893 |
+
ChatGptEnvironment::Staging.agent_identity_authapi_base_url(),
|
| 894 |
+
"https://auth.api.openai.org/api/accounts"
|
| 895 |
+
);
|
| 896 |
+
Ok(())
|
| 897 |
+
}
|
| 898 |
+
|
| 899 |
+
#[test]
|
| 900 |
+
fn chatgpt_environment_rejects_custom_urls() {
|
| 901 |
+
assert!(ChatGptEnvironment::from_chatgpt_base_url("http://localhost:8080").is_err(),);
|
| 902 |
+
}
|
| 903 |
+
|
| 904 |
+
#[test]
|
| 905 |
+
fn agent_registration_url_appends_to_authapi_base_url() {
|
| 906 |
+
assert_eq!(
|
| 907 |
+
agent_registration_url("https://auth.openai.com/api/accounts"),
|
| 908 |
+
"https://auth.openai.com/api/accounts/v1/agent/register"
|
| 909 |
+
);
|
| 910 |
+
assert_eq!(
|
| 911 |
+
agent_registration_url("http://localhost:8080"),
|
| 912 |
+
"http://localhost:8080/v1/agent/register"
|
| 913 |
+
);
|
| 914 |
+
assert_eq!(
|
| 915 |
+
agent_registration_url("http://localhost:8080/backend-api"),
|
| 916 |
+
"http://localhost:8080/backend-api/v1/agent/register"
|
| 917 |
+
);
|
| 918 |
+
}
|
| 919 |
+
|
| 920 |
+
#[test]
|
| 921 |
+
fn agent_task_registration_url_appends_to_authapi_base_url() {
|
| 922 |
+
assert_eq!(
|
| 923 |
+
agent_task_registration_url("https://auth.openai.com/api/accounts", "agent-runtime-id"),
|
| 924 |
+
"https://auth.openai.com/api/accounts/v1/agent/agent-runtime-id/task/register"
|
| 925 |
+
);
|
| 926 |
+
assert_eq!(
|
| 927 |
+
agent_task_registration_url(
|
| 928 |
+
"https://auth.openai.com/api/accounts/",
|
| 929 |
+
"agent-runtime-id"
|
| 930 |
+
),
|
| 931 |
+
"https://auth.openai.com/api/accounts/v1/agent/agent-runtime-id/task/register"
|
| 932 |
+
);
|
| 933 |
+
assert_eq!(
|
| 934 |
+
agent_task_registration_url("http://localhost:8080", "agent-runtime-id"),
|
| 935 |
+
"http://localhost:8080/v1/agent/agent-runtime-id/task/register"
|
| 936 |
+
);
|
| 937 |
+
}
|
| 938 |
+
|
| 939 |
+
#[test]
|
| 940 |
+
fn retryable_registration_error_accepts_429_and_5xx() {
|
| 941 |
+
let too_many_requests = anyhow::Error::new(AgentIdentityRegistrationHttpError::new(
|
| 942 |
+
"agent registration",
|
| 943 |
+
StatusCode::TOO_MANY_REQUESTS,
|
| 944 |
+
"rate limited".to_string(),
|
| 945 |
+
));
|
| 946 |
+
let unavailable = anyhow::Error::new(AgentIdentityRegistrationHttpError::new(
|
| 947 |
+
"agent registration",
|
| 948 |
+
StatusCode::SERVICE_UNAVAILABLE,
|
| 949 |
+
"try later".to_string(),
|
| 950 |
+
));
|
| 951 |
+
|
| 952 |
+
assert!(is_retryable_registration_error(&too_many_requests));
|
| 953 |
+
assert!(is_retryable_registration_error(&unavailable));
|
| 954 |
+
}
|
| 955 |
+
|
| 956 |
+
#[test]
|
| 957 |
+
fn retryable_registration_error_rejects_hard_failures() {
|
| 958 |
+
let forbidden = anyhow::Error::new(AgentIdentityRegistrationHttpError::new(
|
| 959 |
+
"agent registration",
|
| 960 |
+
StatusCode::FORBIDDEN,
|
| 961 |
+
"not allowed".to_string(),
|
| 962 |
+
));
|
| 963 |
+
let malformed = anyhow::anyhow!("failed to sign registration request");
|
| 964 |
+
|
| 965 |
+
assert!(!is_retryable_registration_error(&forbidden));
|
| 966 |
+
assert!(!is_retryable_registration_error(&malformed));
|
| 967 |
+
}
|
| 968 |
+
|
| 969 |
+
#[test]
|
| 970 |
+
fn agent_identity_jwks_url_uses_agent_identity_jwt_route() {
|
| 971 |
+
assert_eq!(
|
| 972 |
+
agent_identity_jwks_url("https://chatgpt.com/backend-api"),
|
| 973 |
+
"https://chatgpt.com/backend-api/wham/agent-identities/jwks"
|
| 974 |
+
);
|
| 975 |
+
assert_eq!(
|
| 976 |
+
agent_identity_jwks_url("https://chatgpt.com/backend-api/"),
|
| 977 |
+
"https://chatgpt.com/backend-api/wham/agent-identities/jwks"
|
| 978 |
+
);
|
| 979 |
+
}
|
| 980 |
+
|
| 981 |
+
#[test]
|
| 982 |
+
fn agent_identity_jwks_url_uses_jwt_issuer_base_url() {
|
| 983 |
+
assert_eq!(
|
| 984 |
+
agent_identity_jwks_url("http://localhost:8080/api/codex"),
|
| 985 |
+
"http://localhost:8080/api/codex/agent-identities/jwks"
|
| 986 |
+
);
|
| 987 |
+
assert_eq!(
|
| 988 |
+
agent_identity_jwks_url("http://localhost:8080/api/codex/"),
|
| 989 |
+
"http://localhost:8080/api/codex/agent-identities/jwks"
|
| 990 |
+
);
|
| 991 |
+
}
|
| 992 |
+
|
| 993 |
+
fn jwt_with_payload(payload: serde_json::Value) -> String {
|
| 994 |
+
let encode = |bytes: &[u8]| URL_SAFE_NO_PAD.encode(bytes);
|
| 995 |
+
let header_b64 = encode(br#"{"alg":"none","typ":"JWT"}"#);
|
| 996 |
+
let payload_b64 = encode(&serde_json::to_vec(&payload).expect("payload should serialize"));
|
| 997 |
+
let signature_b64 = encode(b"sig");
|
| 998 |
+
format!("{header_b64}.{payload_b64}.{signature_b64}")
|
| 999 |
+
}
|
| 1000 |
+
}
|
codex-rs/ansi-escape/src/lib.rs
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use ansi_to_tui::Error;
|
| 2 |
+
use ansi_to_tui::IntoText;
|
| 3 |
+
use ratatui::text::Line;
|
| 4 |
+
use ratatui::text::Text;
|
| 5 |
+
|
| 6 |
+
// Expand tabs in a best-effort way for transcript rendering.
|
| 7 |
+
// Tabs can interact poorly with left-gutter prefixes in our TUI and CLI
|
| 8 |
+
// transcript views (e.g., `nl` separates line numbers from content with a tab).
|
| 9 |
+
// Replacing tabs with spaces avoids odd visual artifacts without changing
|
| 10 |
+
// semantics for our use cases.
|
| 11 |
+
fn expand_tabs(s: &str) -> std::borrow::Cow<'_, str> {
|
| 12 |
+
if s.contains('\t') {
|
| 13 |
+
// Keep it simple: replace each tab with 4 spaces.
|
| 14 |
+
// We do not try to align to tab stops since most usages (like `nl`)
|
| 15 |
+
// look acceptable with a fixed substitution and this avoids stateful math
|
| 16 |
+
// across spans.
|
| 17 |
+
std::borrow::Cow::Owned(s.replace('\t', " "))
|
| 18 |
+
} else {
|
| 19 |
+
std::borrow::Cow::Borrowed(s)
|
| 20 |
+
}
|
| 21 |
+
}
|
| 22 |
+
|
| 23 |
+
/// This function should be used when the contents of `s` are expected to match
|
| 24 |
+
/// a single line. If multiple lines are found, a warning is logged and only the
|
| 25 |
+
/// first line is returned.
|
| 26 |
+
pub fn ansi_escape_line(s: &str) -> Line<'static> {
|
| 27 |
+
// Normalize tabs to spaces to avoid odd gutter collisions in transcript mode.
|
| 28 |
+
let s = expand_tabs(s);
|
| 29 |
+
let text = ansi_escape(&s);
|
| 30 |
+
match text.lines.as_slice() {
|
| 31 |
+
[] => "".into(),
|
| 32 |
+
[only] => only.clone(),
|
| 33 |
+
[first, rest @ ..] => {
|
| 34 |
+
tracing::warn!("ansi_escape_line: expected a single line, got {first:?} and {rest:?}");
|
| 35 |
+
first.clone()
|
| 36 |
+
}
|
| 37 |
+
}
|
| 38 |
+
}
|
| 39 |
+
|
| 40 |
+
pub fn ansi_escape(s: &str) -> Text<'static> {
|
| 41 |
+
// to_text() claims to be faster, but introduces complex lifetime issues
|
| 42 |
+
// such that it's not worth it.
|
| 43 |
+
match s.into_text() {
|
| 44 |
+
Ok(text) => text,
|
| 45 |
+
Err(err) => match err {
|
| 46 |
+
Error::NomError(message) => {
|
| 47 |
+
tracing::error!(
|
| 48 |
+
"ansi_to_tui NomError docs claim should never happen when parsing `{s}`: {message}"
|
| 49 |
+
);
|
| 50 |
+
panic!();
|
| 51 |
+
}
|
| 52 |
+
Error::Utf8Error(utf8error) => {
|
| 53 |
+
tracing::error!("Utf8Error: {utf8error}");
|
| 54 |
+
panic!();
|
| 55 |
+
}
|
| 56 |
+
},
|
| 57 |
+
}
|
| 58 |
+
}
|
codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:d6fb74c9af068055ac5e369ce116c27299e46ad2a0487541a7abbd52761f9832
|
| 3 |
+
size 157496
|
codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:be380114e24a7541e07df1af9b1d2d10f8749f1d97516496a7e9c59bd395c29d
|
| 3 |
+
size 151569
|
codex-rs/app-server-transport/src/connection_auth.rs
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Binds a transport connection to the authentication owner that established it.
|
| 2 |
+
//! Owner revisions invalidate queued work even before transport closure is delivered.
|
| 3 |
+
|
| 4 |
+
use codex_login::AuthChangeState;
|
| 5 |
+
use codex_login::AuthManager;
|
| 6 |
+
use std::io;
|
| 7 |
+
use tokio::sync::watch;
|
| 8 |
+
|
| 9 |
+
#[derive(Clone, Debug)]
|
| 10 |
+
pub struct ConnectionAuth {
|
| 11 |
+
changes: watch::Receiver<AuthChangeState>,
|
| 12 |
+
owner_generation: u64,
|
| 13 |
+
}
|
| 14 |
+
|
| 15 |
+
impl ConnectionAuth {
|
| 16 |
+
pub(crate) fn capture(auth_manager: &AuthManager) -> Self {
|
| 17 |
+
let changes = auth_manager.auth_change_state_receiver();
|
| 18 |
+
let owner_generation = changes.borrow().owner_generation;
|
| 19 |
+
Self::new(changes, owner_generation)
|
| 20 |
+
}
|
| 21 |
+
|
| 22 |
+
pub(crate) fn ensure_current(&self) -> io::Result<()> {
|
| 23 |
+
if self.is_current() {
|
| 24 |
+
Ok(())
|
| 25 |
+
} else {
|
| 26 |
+
Err(io::Error::new(
|
| 27 |
+
io::ErrorKind::Interrupted,
|
| 28 |
+
"remote control authentication changed",
|
| 29 |
+
))
|
| 30 |
+
}
|
| 31 |
+
}
|
| 32 |
+
|
| 33 |
+
pub(crate) fn new(changes: watch::Receiver<AuthChangeState>, owner_generation: u64) -> Self {
|
| 34 |
+
Self {
|
| 35 |
+
changes,
|
| 36 |
+
owner_generation,
|
| 37 |
+
}
|
| 38 |
+
}
|
| 39 |
+
|
| 40 |
+
pub fn is_current(&self) -> bool {
|
| 41 |
+
self.changes.borrow().owner_generation == self.owner_generation
|
| 42 |
+
&& self.changes.has_changed().is_ok()
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
pub(crate) async fn invalidated(&self) {
|
| 46 |
+
let mut changes = self.changes.clone();
|
| 47 |
+
let _ = changes
|
| 48 |
+
.wait_for(|state| state.owner_generation != self.owner_generation)
|
| 49 |
+
.await;
|
| 50 |
+
}
|
| 51 |
+
}
|
codex-rs/app-server-transport/src/daemon_recovery.rs
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Shared on-disk candidate set for managed daemon restarts.
|
| 2 |
+
|
| 3 |
+
use std::collections::BTreeMap;
|
| 4 |
+
use std::collections::BTreeSet;
|
| 5 |
+
use std::io;
|
| 6 |
+
use std::path::Path;
|
| 7 |
+
|
| 8 |
+
use codex_core::path_utils::write_atomically;
|
| 9 |
+
use serde::Deserialize;
|
| 10 |
+
use serde::Serialize;
|
| 11 |
+
|
| 12 |
+
#[derive(Debug, Default, Deserialize, Serialize, PartialEq, Eq)]
|
| 13 |
+
pub struct RecoverySnapshot {
|
| 14 |
+
#[serde(skip)]
|
| 15 |
+
pub loaded: BTreeSet<String>,
|
| 16 |
+
pub interrupted: BTreeMap<String, InterruptedTurn>,
|
| 17 |
+
}
|
| 18 |
+
|
| 19 |
+
#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
|
| 20 |
+
pub struct InterruptedTurn {
|
| 21 |
+
pub turn_id: String,
|
| 22 |
+
pub output_schema: Option<serde_json::Value>,
|
| 23 |
+
pub service_tier: Option<String>,
|
| 24 |
+
pub cyber_access_program: Option<codex_protocol::turn_input::CyberAccessProgram>,
|
| 25 |
+
/// Only local execution with thread-owned configuration can continue automatically.
|
| 26 |
+
/// Older snapshots without this identity are reloaded without continuation.
|
| 27 |
+
pub local_environment: Option<codex_app_server_protocol::ThreadEnvironment>,
|
| 28 |
+
}
|
| 29 |
+
|
| 30 |
+
// Old servers accept the array and skip this non-thread entry during best-effort
|
| 31 |
+
// restoration. Keeping metadata in the same atomic file avoids stale sidecars.
|
| 32 |
+
const INTERRUPTION_PREFIX: &str = "codex-interrupted-v1:";
|
| 33 |
+
|
| 34 |
+
pub fn read_snapshot(path: &Path) -> io::Result<RecoverySnapshot> {
|
| 35 |
+
let mut loaded: BTreeSet<String> = match std::fs::read(path) {
|
| 36 |
+
Ok(contents) => serde_json::from_slice(&contents).map_err(io::Error::other)?,
|
| 37 |
+
Err(err) if err.kind() == io::ErrorKind::NotFound => return Ok(RecoverySnapshot::default()),
|
| 38 |
+
Err(err) => return Err(err),
|
| 39 |
+
};
|
| 40 |
+
let mut snapshot = RecoverySnapshot::default();
|
| 41 |
+
loaded.retain(|entry| {
|
| 42 |
+
if let Some(metadata) = entry.strip_prefix(INTERRUPTION_PREFIX) {
|
| 43 |
+
if let Ok(saved) = serde_json::from_str::<RecoverySnapshot>(metadata) {
|
| 44 |
+
snapshot = saved;
|
| 45 |
+
}
|
| 46 |
+
false
|
| 47 |
+
} else {
|
| 48 |
+
true
|
| 49 |
+
}
|
| 50 |
+
});
|
| 51 |
+
snapshot.interrupted.retain(|id, _| loaded.contains(id));
|
| 52 |
+
snapshot.loaded = loaded;
|
| 53 |
+
Ok(snapshot)
|
| 54 |
+
}
|
| 55 |
+
|
| 56 |
+
pub fn read_candidates(path: &Path) -> io::Result<BTreeSet<String>> {
|
| 57 |
+
Ok(read_snapshot(path)?.loaded)
|
| 58 |
+
}
|
| 59 |
+
|
| 60 |
+
pub fn write_candidates(path: &Path, candidates: &BTreeSet<String>) -> io::Result<()> {
|
| 61 |
+
write_snapshot(
|
| 62 |
+
path,
|
| 63 |
+
&RecoverySnapshot {
|
| 64 |
+
loaded: candidates.clone(),
|
| 65 |
+
..Default::default()
|
| 66 |
+
},
|
| 67 |
+
)
|
| 68 |
+
}
|
| 69 |
+
|
| 70 |
+
pub fn write_snapshot(path: &Path, snapshot: &RecoverySnapshot) -> io::Result<()> {
|
| 71 |
+
let mut saved = snapshot.loaded.clone();
|
| 72 |
+
if !snapshot.interrupted.is_empty() {
|
| 73 |
+
saved.insert(format!(
|
| 74 |
+
"{INTERRUPTION_PREFIX}{}",
|
| 75 |
+
serde_json::to_string(snapshot).map_err(io::Error::other)?
|
| 76 |
+
));
|
| 77 |
+
}
|
| 78 |
+
write_atomically(
|
| 79 |
+
path,
|
| 80 |
+
&serde_json::to_string(&saved).map_err(io::Error::other)?,
|
| 81 |
+
)
|
| 82 |
+
}
|
codex-rs/app-server-transport/src/daemon_shutdown.rs
ADDED
|
@@ -0,0 +1,53 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! The detached Windows updater has no control socket, so its owner requests
|
| 2 |
+
//! termination through a file in its private state directory. Managed app-server
|
| 3 |
+
//! shutdown instead uses the protected local socket.
|
| 4 |
+
|
| 5 |
+
use std::io;
|
| 6 |
+
use std::io::Read;
|
| 7 |
+
use std::path::Path;
|
| 8 |
+
#[cfg(windows)]
|
| 9 |
+
use std::path::PathBuf;
|
| 10 |
+
#[cfg(windows)]
|
| 11 |
+
use std::time::Duration;
|
| 12 |
+
|
| 13 |
+
#[cfg(windows)]
|
| 14 |
+
pub const DAEMON_SHUTDOWN_FILE_ENV: &str = "CODEX_DAEMON_SHUTDOWN_FILE";
|
| 15 |
+
|
| 16 |
+
/// Waits for and consumes one updater shutdown request.
|
| 17 |
+
#[cfg(windows)]
|
| 18 |
+
pub async fn daemon_shutdown_signal() -> io::Result<()> {
|
| 19 |
+
let Some(path) = std::env::var_os(DAEMON_SHUTDOWN_FILE_ENV).map(PathBuf::from) else {
|
| 20 |
+
return std::future::pending().await;
|
| 21 |
+
};
|
| 22 |
+
loop {
|
| 23 |
+
if take_shutdown_request(&path, std::process::id())? {
|
| 24 |
+
return Ok(());
|
| 25 |
+
}
|
| 26 |
+
tokio::time::sleep(Duration::from_millis(50)).await;
|
| 27 |
+
}
|
| 28 |
+
}
|
| 29 |
+
|
| 30 |
+
fn take_shutdown_request(path: &Path, pid: u32) -> io::Result<bool> {
|
| 31 |
+
let file = match std::fs::File::open(path) {
|
| 32 |
+
Ok(file) => file,
|
| 33 |
+
Err(err) if err.kind() == io::ErrorKind::NotFound => return Ok(false),
|
| 34 |
+
Err(err) => return Err(err),
|
| 35 |
+
};
|
| 36 |
+
// Descendant app-servers may inherit the control path. Only the intended
|
| 37 |
+
// process can consume a request. Bound reads to a u32 PID plus one extra byte.
|
| 38 |
+
let mut contents = Vec::new();
|
| 39 |
+
file.take(/*limit*/ 11).read_to_end(&mut contents)?;
|
| 40 |
+
if contents != pid.to_string().as_bytes() {
|
| 41 |
+
return Ok(false);
|
| 42 |
+
}
|
| 43 |
+
// Synchronous consumption cannot lose a request to select cancellation.
|
| 44 |
+
match std::fs::remove_file(path) {
|
| 45 |
+
Ok(()) => Ok(true),
|
| 46 |
+
Err(err) if err.kind() == io::ErrorKind::NotFound => Ok(false),
|
| 47 |
+
Err(err) => Err(err),
|
| 48 |
+
}
|
| 49 |
+
}
|
| 50 |
+
|
| 51 |
+
#[cfg(test)]
|
| 52 |
+
#[path = "daemon_shutdown_tests.rs"]
|
| 53 |
+
mod tests;
|
codex-rs/app-server-transport/src/daemon_shutdown_tests.rs
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::take_shutdown_request;
|
| 2 |
+
|
| 3 |
+
#[test]
|
| 4 |
+
fn inherited_control_path_cannot_consume_another_process_request() {
|
| 5 |
+
let directory = tempfile::tempdir().expect("directory");
|
| 6 |
+
let request = directory.path().join("shutdown");
|
| 7 |
+
std::fs::write(&request, "1234").expect("request");
|
| 8 |
+
assert!(!take_shutdown_request(&request, /*pid*/ 5678).expect("descendant probe"));
|
| 9 |
+
assert!(take_shutdown_request(&request, /*pid*/ 1234).expect("parent probe"));
|
| 10 |
+
assert!(!take_shutdown_request(&request, /*pid*/ 1234).expect("already consumed"));
|
| 11 |
+
}
|
| 12 |
+
|
| 13 |
+
#[test]
|
| 14 |
+
fn incomplete_or_invalid_requests_are_not_consumed() {
|
| 15 |
+
let directory = tempfile::tempdir().expect("directory");
|
| 16 |
+
let request = directory.path().join("shutdown");
|
| 17 |
+
for contents in [b"".as_slice(), b"123", b"1234junk", &[0xff; 20]] {
|
| 18 |
+
std::fs::write(&request, contents).expect("request");
|
| 19 |
+
assert!(!take_shutdown_request(&request, /*pid*/ 1234).expect("probe"));
|
| 20 |
+
assert!(request.exists());
|
| 21 |
+
}
|
| 22 |
+
}
|
codex-rs/app-server-transport/src/lib.rs
ADDED
|
@@ -0,0 +1,45 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
pub mod daemon_recovery;
|
| 2 |
+
#[cfg(any(windows, test))]
|
| 3 |
+
mod daemon_shutdown;
|
| 4 |
+
#[cfg(windows)]
|
| 5 |
+
pub use daemon_shutdown::DAEMON_SHUTDOWN_FILE_ENV;
|
| 6 |
+
#[cfg(windows)]
|
| 7 |
+
pub use daemon_shutdown::daemon_shutdown_signal;
|
| 8 |
+
/// Only managed app-server launches accept the local socket shutdown request.
|
| 9 |
+
pub const DAEMON_SHUTDOWN_SOCKET_ENV: &str = "CODEX_DAEMON_SHUTDOWN_SOCKET";
|
| 10 |
+
mod connection_auth;
|
| 11 |
+
mod outgoing_message;
|
| 12 |
+
mod transport;
|
| 13 |
+
|
| 14 |
+
pub use connection_auth::ConnectionAuth;
|
| 15 |
+
pub use outgoing_message::ConnectionId;
|
| 16 |
+
pub use outgoing_message::OutgoingError;
|
| 17 |
+
pub use outgoing_message::OutgoingMessage;
|
| 18 |
+
pub use outgoing_message::OutgoingResponse;
|
| 19 |
+
pub use outgoing_message::QueuedOutgoingMessage;
|
| 20 |
+
pub use transport::AppServerStartupLock;
|
| 21 |
+
pub use transport::AppServerTransport;
|
| 22 |
+
pub use transport::AppServerTransportParseError;
|
| 23 |
+
pub use transport::CHANNEL_CAPACITY;
|
| 24 |
+
pub use transport::ConnectionOrigin;
|
| 25 |
+
pub use transport::DaemonShutdownAccess;
|
| 26 |
+
pub use transport::REMOTE_CONTROL_DISABLED_ENV_VAR;
|
| 27 |
+
pub use transport::RemoteControlDisabledByRequirements;
|
| 28 |
+
pub use transport::RemoteControlEnableError;
|
| 29 |
+
pub use transport::RemoteControlHandle;
|
| 30 |
+
pub use transport::RemoteControlPolicy;
|
| 31 |
+
pub use transport::RemoteControlStartConfig;
|
| 32 |
+
pub use transport::RemoteControlStartupMode;
|
| 33 |
+
pub use transport::RemoteControlUnavailable;
|
| 34 |
+
pub use transport::TransportEvent;
|
| 35 |
+
pub use transport::acquire_app_server_startup_lock;
|
| 36 |
+
pub use transport::app_server_control_socket_path;
|
| 37 |
+
pub use transport::app_server_startup_lock_path;
|
| 38 |
+
pub use transport::auth;
|
| 39 |
+
pub use transport::daemon_recovery_file_path;
|
| 40 |
+
pub use transport::prepare_control_socket_path;
|
| 41 |
+
pub use transport::start_control_socket_acceptor;
|
| 42 |
+
pub use transport::start_remote_control;
|
| 43 |
+
pub use transport::start_stdio_connection;
|
| 44 |
+
pub use transport::start_websocket_acceptor;
|
| 45 |
+
pub use transport::take_remote_control_disabled_env;
|
codex-rs/app-server-transport/src/outgoing_message.rs
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use std::fmt;
|
| 2 |
+
|
| 3 |
+
use codex_app_server_protocol::ClientResponsePayload;
|
| 4 |
+
use codex_app_server_protocol::JSONRPCErrorError;
|
| 5 |
+
use codex_app_server_protocol::RequestId;
|
| 6 |
+
use codex_app_server_protocol::ServerNotificationEnvelope;
|
| 7 |
+
use codex_app_server_protocol::ServerRequest;
|
| 8 |
+
use serde::Serialize;
|
| 9 |
+
use tokio::sync::oneshot;
|
| 10 |
+
|
| 11 |
+
/// Stable identifier for a transport connection.
|
| 12 |
+
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
|
| 13 |
+
pub struct ConnectionId(pub u64);
|
| 14 |
+
|
| 15 |
+
impl fmt::Display for ConnectionId {
|
| 16 |
+
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 17 |
+
write!(f, "{}", self.0)
|
| 18 |
+
}
|
| 19 |
+
}
|
| 20 |
+
|
| 21 |
+
/// Outgoing message from the server to the client.
|
| 22 |
+
#[derive(Debug, Clone, Serialize)]
|
| 23 |
+
#[serde(untagged)]
|
| 24 |
+
#[allow(clippy::large_enum_variant)]
|
| 25 |
+
pub enum OutgoingMessage {
|
| 26 |
+
Request(ServerRequest),
|
| 27 |
+
/// AppServerNotification is specific to the case where this is run as an
|
| 28 |
+
/// "app server" as opposed to an MCP server.
|
| 29 |
+
AppServerNotification(ServerNotificationEnvelope),
|
| 30 |
+
Response(OutgoingResponse),
|
| 31 |
+
Error(OutgoingError),
|
| 32 |
+
}
|
| 33 |
+
|
| 34 |
+
#[derive(Debug, Clone, Serialize)]
|
| 35 |
+
pub struct OutgoingResponse {
|
| 36 |
+
pub id: RequestId,
|
| 37 |
+
pub result: Box<ClientResponsePayload>,
|
| 38 |
+
}
|
| 39 |
+
|
| 40 |
+
#[derive(Debug, Clone, PartialEq, Serialize)]
|
| 41 |
+
pub struct OutgoingError {
|
| 42 |
+
pub error: JSONRPCErrorError,
|
| 43 |
+
pub id: RequestId,
|
| 44 |
+
}
|
| 45 |
+
|
| 46 |
+
#[derive(Debug)]
|
| 47 |
+
pub struct QueuedOutgoingMessage {
|
| 48 |
+
pub message: OutgoingMessage,
|
| 49 |
+
pub write_complete_tx: Option<oneshot::Sender<()>>,
|
| 50 |
+
}
|
| 51 |
+
|
| 52 |
+
impl QueuedOutgoingMessage {
|
| 53 |
+
pub fn new(message: OutgoingMessage) -> Self {
|
| 54 |
+
Self {
|
| 55 |
+
message,
|
| 56 |
+
write_complete_tx: None,
|
| 57 |
+
}
|
| 58 |
+
}
|
| 59 |
+
}
|
codex-rs/app-server-transport/src/transport/auth.rs
ADDED
|
@@ -0,0 +1,751 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use anyhow::Context;
|
| 2 |
+
use axum::http::HeaderMap;
|
| 3 |
+
use axum::http::StatusCode;
|
| 4 |
+
use axum::http::header::AUTHORIZATION;
|
| 5 |
+
use clap::Args;
|
| 6 |
+
use clap::ValueEnum;
|
| 7 |
+
use codex_utils_absolute_path::AbsolutePathBuf;
|
| 8 |
+
use constant_time_eq::constant_time_eq_32;
|
| 9 |
+
use jsonwebtoken::Algorithm;
|
| 10 |
+
use jsonwebtoken::DecodingKey;
|
| 11 |
+
use jsonwebtoken::Validation;
|
| 12 |
+
use jsonwebtoken::decode;
|
| 13 |
+
use serde::Deserialize;
|
| 14 |
+
use sha2::Digest;
|
| 15 |
+
use sha2::Sha256;
|
| 16 |
+
use std::io;
|
| 17 |
+
use std::io::ErrorKind;
|
| 18 |
+
use std::net::SocketAddr;
|
| 19 |
+
use std::path::Path;
|
| 20 |
+
use std::path::PathBuf;
|
| 21 |
+
use time::OffsetDateTime;
|
| 22 |
+
|
| 23 |
+
const DEFAULT_MAX_CLOCK_SKEW_SECONDS: u64 = 30;
|
| 24 |
+
const MIN_SIGNED_BEARER_SECRET_BYTES: usize = 32;
|
| 25 |
+
const INVALID_AUTHORIZATION_HEADER_MESSAGE: &str = "invalid authorization header";
|
| 26 |
+
|
| 27 |
+
#[derive(Debug, Clone, Default, PartialEq, Eq, Args)]
|
| 28 |
+
pub struct AppServerWebsocketAuthArgs {
|
| 29 |
+
/// Websocket auth mode for non-loopback listeners.
|
| 30 |
+
#[arg(long = "ws-auth", value_name = "MODE", value_enum)]
|
| 31 |
+
pub ws_auth: Option<WebsocketAuthCliMode>,
|
| 32 |
+
|
| 33 |
+
/// Absolute path to the capability-token file.
|
| 34 |
+
#[arg(long = "ws-token-file", value_name = "PATH")]
|
| 35 |
+
pub ws_token_file: Option<PathBuf>,
|
| 36 |
+
|
| 37 |
+
/// Hex-encoded SHA-256 digest of the capability token.
|
| 38 |
+
#[arg(long = "ws-token-sha256", value_name = "HEX")]
|
| 39 |
+
pub ws_token_sha256: Option<String>,
|
| 40 |
+
|
| 41 |
+
/// Absolute path to the shared secret file for signed JWT bearer tokens.
|
| 42 |
+
#[arg(long = "ws-shared-secret-file", value_name = "PATH")]
|
| 43 |
+
pub ws_shared_secret_file: Option<PathBuf>,
|
| 44 |
+
|
| 45 |
+
/// Expected issuer for signed JWT bearer tokens.
|
| 46 |
+
#[arg(long = "ws-issuer", value_name = "ISSUER")]
|
| 47 |
+
pub ws_issuer: Option<String>,
|
| 48 |
+
|
| 49 |
+
/// Expected audience for signed JWT bearer tokens.
|
| 50 |
+
#[arg(long = "ws-audience", value_name = "AUDIENCE")]
|
| 51 |
+
pub ws_audience: Option<String>,
|
| 52 |
+
|
| 53 |
+
/// Maximum clock skew when validating signed JWT bearer tokens.
|
| 54 |
+
#[arg(long = "ws-max-clock-skew-seconds", value_name = "SECONDS")]
|
| 55 |
+
pub ws_max_clock_skew_seconds: Option<u64>,
|
| 56 |
+
}
|
| 57 |
+
|
| 58 |
+
#[derive(Debug, Clone, Copy, PartialEq, Eq, ValueEnum)]
|
| 59 |
+
pub enum WebsocketAuthCliMode {
|
| 60 |
+
CapabilityToken,
|
| 61 |
+
SignedBearerToken,
|
| 62 |
+
}
|
| 63 |
+
|
| 64 |
+
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
| 65 |
+
pub struct AppServerWebsocketAuthSettings {
|
| 66 |
+
pub config: Option<AppServerWebsocketAuthConfig>,
|
| 67 |
+
}
|
| 68 |
+
|
| 69 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 70 |
+
pub enum AppServerWebsocketAuthConfig {
|
| 71 |
+
CapabilityToken {
|
| 72 |
+
source: AppServerWebsocketCapabilityTokenSource,
|
| 73 |
+
},
|
| 74 |
+
SignedBearerToken {
|
| 75 |
+
shared_secret_file: AbsolutePathBuf,
|
| 76 |
+
issuer: Option<String>,
|
| 77 |
+
audience: Option<String>,
|
| 78 |
+
max_clock_skew_seconds: u64,
|
| 79 |
+
},
|
| 80 |
+
}
|
| 81 |
+
|
| 82 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 83 |
+
pub enum AppServerWebsocketCapabilityTokenSource {
|
| 84 |
+
TokenFile { token_file: AbsolutePathBuf },
|
| 85 |
+
TokenSha256 { token_sha256: [u8; 32] },
|
| 86 |
+
}
|
| 87 |
+
|
| 88 |
+
#[derive(Clone, Debug, Default)]
|
| 89 |
+
pub struct WebsocketAuthPolicy {
|
| 90 |
+
pub(crate) mode: Option<WebsocketAuthMode>,
|
| 91 |
+
}
|
| 92 |
+
|
| 93 |
+
#[derive(Clone, Debug)]
|
| 94 |
+
pub(crate) enum WebsocketAuthMode {
|
| 95 |
+
CapabilityToken {
|
| 96 |
+
token_sha256: [u8; 32],
|
| 97 |
+
},
|
| 98 |
+
SignedBearerToken {
|
| 99 |
+
shared_secret: Vec<u8>,
|
| 100 |
+
issuer: Option<String>,
|
| 101 |
+
audience: Option<String>,
|
| 102 |
+
max_clock_skew_seconds: i64,
|
| 103 |
+
},
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
#[derive(Debug)]
|
| 107 |
+
pub(crate) struct WebsocketAuthError {
|
| 108 |
+
status_code: StatusCode,
|
| 109 |
+
message: &'static str,
|
| 110 |
+
}
|
| 111 |
+
|
| 112 |
+
#[derive(Deserialize)]
|
| 113 |
+
struct JwtClaims {
|
| 114 |
+
exp: i64,
|
| 115 |
+
nbf: Option<i64>,
|
| 116 |
+
iss: Option<String>,
|
| 117 |
+
aud: Option<JwtAudienceClaim>,
|
| 118 |
+
}
|
| 119 |
+
|
| 120 |
+
#[derive(Deserialize)]
|
| 121 |
+
#[serde(untagged)]
|
| 122 |
+
enum JwtAudienceClaim {
|
| 123 |
+
Single(String),
|
| 124 |
+
Multiple(Vec<String>),
|
| 125 |
+
}
|
| 126 |
+
|
| 127 |
+
impl WebsocketAuthError {
|
| 128 |
+
pub(crate) fn status_code(&self) -> StatusCode {
|
| 129 |
+
self.status_code
|
| 130 |
+
}
|
| 131 |
+
|
| 132 |
+
pub(crate) fn message(&self) -> &'static str {
|
| 133 |
+
self.message
|
| 134 |
+
}
|
| 135 |
+
}
|
| 136 |
+
|
| 137 |
+
impl AppServerWebsocketAuthArgs {
|
| 138 |
+
pub fn try_into_settings(self) -> anyhow::Result<AppServerWebsocketAuthSettings> {
|
| 139 |
+
let normalize = |value: Option<String>| {
|
| 140 |
+
value.and_then(|value| {
|
| 141 |
+
let trimmed = value.trim();
|
| 142 |
+
(!trimmed.is_empty()).then(|| trimmed.to_string())
|
| 143 |
+
})
|
| 144 |
+
};
|
| 145 |
+
|
| 146 |
+
let config = match self.ws_auth {
|
| 147 |
+
Some(WebsocketAuthCliMode::CapabilityToken) => {
|
| 148 |
+
if self.ws_shared_secret_file.is_some()
|
| 149 |
+
|| self.ws_issuer.is_some()
|
| 150 |
+
|| self.ws_audience.is_some()
|
| 151 |
+
|| self.ws_max_clock_skew_seconds.is_some()
|
| 152 |
+
{
|
| 153 |
+
anyhow::bail!(
|
| 154 |
+
"`--ws-shared-secret-file`, `--ws-issuer`, `--ws-audience`, and `--ws-max-clock-skew-seconds` require `--ws-auth signed-bearer-token`"
|
| 155 |
+
);
|
| 156 |
+
}
|
| 157 |
+
let source = match (self.ws_token_file, self.ws_token_sha256) {
|
| 158 |
+
(Some(_), Some(_)) => {
|
| 159 |
+
anyhow::bail!(
|
| 160 |
+
"`--ws-token-file` and `--ws-token-sha256` are mutually exclusive"
|
| 161 |
+
);
|
| 162 |
+
}
|
| 163 |
+
(Some(token_file), None) => {
|
| 164 |
+
AppServerWebsocketCapabilityTokenSource::TokenFile {
|
| 165 |
+
token_file: absolute_path_arg("--ws-token-file", token_file)?,
|
| 166 |
+
}
|
| 167 |
+
}
|
| 168 |
+
(None, Some(token_sha256)) => {
|
| 169 |
+
AppServerWebsocketCapabilityTokenSource::TokenSha256 {
|
| 170 |
+
token_sha256: sha256_digest_arg("--ws-token-sha256", &token_sha256)?,
|
| 171 |
+
}
|
| 172 |
+
}
|
| 173 |
+
(None, None) => {
|
| 174 |
+
anyhow::bail!(
|
| 175 |
+
"`--ws-token-file` or `--ws-token-sha256` is required when `--ws-auth capability-token` is set"
|
| 176 |
+
);
|
| 177 |
+
}
|
| 178 |
+
};
|
| 179 |
+
Some(AppServerWebsocketAuthConfig::CapabilityToken { source })
|
| 180 |
+
}
|
| 181 |
+
Some(WebsocketAuthCliMode::SignedBearerToken) => {
|
| 182 |
+
if self.ws_token_file.is_some() || self.ws_token_sha256.is_some() {
|
| 183 |
+
anyhow::bail!(
|
| 184 |
+
"`--ws-token-file` and `--ws-token-sha256` require `--ws-auth capability-token`, not `signed-bearer-token`"
|
| 185 |
+
);
|
| 186 |
+
}
|
| 187 |
+
let shared_secret_file = self.ws_shared_secret_file.context(
|
| 188 |
+
"`--ws-shared-secret-file` is required when `--ws-auth signed-bearer-token` is set",
|
| 189 |
+
)?;
|
| 190 |
+
Some(AppServerWebsocketAuthConfig::SignedBearerToken {
|
| 191 |
+
shared_secret_file: absolute_path_arg(
|
| 192 |
+
"--ws-shared-secret-file",
|
| 193 |
+
shared_secret_file,
|
| 194 |
+
)?,
|
| 195 |
+
issuer: normalize(self.ws_issuer),
|
| 196 |
+
audience: normalize(self.ws_audience),
|
| 197 |
+
max_clock_skew_seconds: self
|
| 198 |
+
.ws_max_clock_skew_seconds
|
| 199 |
+
.unwrap_or(DEFAULT_MAX_CLOCK_SKEW_SECONDS),
|
| 200 |
+
})
|
| 201 |
+
}
|
| 202 |
+
None => {
|
| 203 |
+
if self.ws_token_file.is_some()
|
| 204 |
+
|| self.ws_token_sha256.is_some()
|
| 205 |
+
|| self.ws_shared_secret_file.is_some()
|
| 206 |
+
|| self.ws_issuer.is_some()
|
| 207 |
+
|| self.ws_audience.is_some()
|
| 208 |
+
|| self.ws_max_clock_skew_seconds.is_some()
|
| 209 |
+
{
|
| 210 |
+
anyhow::bail!(
|
| 211 |
+
"websocket auth flags require `--ws-auth capability-token` or `--ws-auth signed-bearer-token`"
|
| 212 |
+
);
|
| 213 |
+
}
|
| 214 |
+
None
|
| 215 |
+
}
|
| 216 |
+
};
|
| 217 |
+
|
| 218 |
+
Ok(AppServerWebsocketAuthSettings { config })
|
| 219 |
+
}
|
| 220 |
+
}
|
| 221 |
+
|
| 222 |
+
pub fn policy_from_settings(
|
| 223 |
+
settings: &AppServerWebsocketAuthSettings,
|
| 224 |
+
) -> io::Result<WebsocketAuthPolicy> {
|
| 225 |
+
let mode = match settings.config.as_ref() {
|
| 226 |
+
Some(AppServerWebsocketAuthConfig::CapabilityToken { source }) => match source {
|
| 227 |
+
AppServerWebsocketCapabilityTokenSource::TokenFile { token_file } => {
|
| 228 |
+
let token = read_trimmed_secret(token_file.as_ref())?;
|
| 229 |
+
Some(WebsocketAuthMode::CapabilityToken {
|
| 230 |
+
token_sha256: sha256_digest(token.as_bytes()),
|
| 231 |
+
})
|
| 232 |
+
}
|
| 233 |
+
AppServerWebsocketCapabilityTokenSource::TokenSha256 { token_sha256 } => {
|
| 234 |
+
Some(WebsocketAuthMode::CapabilityToken {
|
| 235 |
+
token_sha256: *token_sha256,
|
| 236 |
+
})
|
| 237 |
+
}
|
| 238 |
+
},
|
| 239 |
+
Some(AppServerWebsocketAuthConfig::SignedBearerToken {
|
| 240 |
+
shared_secret_file,
|
| 241 |
+
issuer,
|
| 242 |
+
audience,
|
| 243 |
+
max_clock_skew_seconds,
|
| 244 |
+
}) => {
|
| 245 |
+
let shared_secret = read_trimmed_secret(shared_secret_file.as_ref())?.into_bytes();
|
| 246 |
+
validate_signed_bearer_secret(shared_secret_file.as_ref(), &shared_secret)?;
|
| 247 |
+
let max_clock_skew_seconds = i64::try_from(*max_clock_skew_seconds).map_err(|_| {
|
| 248 |
+
io::Error::new(
|
| 249 |
+
ErrorKind::InvalidInput,
|
| 250 |
+
"websocket auth clock skew must fit in a signed 64-bit integer",
|
| 251 |
+
)
|
| 252 |
+
})?;
|
| 253 |
+
Some(WebsocketAuthMode::SignedBearerToken {
|
| 254 |
+
shared_secret,
|
| 255 |
+
issuer: issuer.clone(),
|
| 256 |
+
audience: audience.clone(),
|
| 257 |
+
max_clock_skew_seconds,
|
| 258 |
+
})
|
| 259 |
+
}
|
| 260 |
+
None => None,
|
| 261 |
+
};
|
| 262 |
+
|
| 263 |
+
Ok(WebsocketAuthPolicy { mode })
|
| 264 |
+
}
|
| 265 |
+
|
| 266 |
+
pub(crate) fn is_unauthenticated_non_loopback_listener(
|
| 267 |
+
bind_address: SocketAddr,
|
| 268 |
+
policy: &WebsocketAuthPolicy,
|
| 269 |
+
) -> bool {
|
| 270 |
+
!bind_address.ip().is_loopback() && policy.mode.is_none()
|
| 271 |
+
}
|
| 272 |
+
|
| 273 |
+
pub(crate) fn authorize_upgrade(
|
| 274 |
+
headers: &HeaderMap,
|
| 275 |
+
policy: &WebsocketAuthPolicy,
|
| 276 |
+
) -> Result<(), WebsocketAuthError> {
|
| 277 |
+
let Some(mode) = policy.mode.as_ref() else {
|
| 278 |
+
return Ok(());
|
| 279 |
+
};
|
| 280 |
+
|
| 281 |
+
let token = bearer_token_from_headers(headers)?;
|
| 282 |
+
match mode {
|
| 283 |
+
WebsocketAuthMode::CapabilityToken { token_sha256 } => {
|
| 284 |
+
let actual_sha256 = sha256_digest(token.as_bytes());
|
| 285 |
+
if constant_time_eq_32(token_sha256, &actual_sha256) {
|
| 286 |
+
Ok(())
|
| 287 |
+
} else {
|
| 288 |
+
Err(unauthorized("invalid websocket bearer token"))
|
| 289 |
+
}
|
| 290 |
+
}
|
| 291 |
+
WebsocketAuthMode::SignedBearerToken {
|
| 292 |
+
shared_secret,
|
| 293 |
+
issuer,
|
| 294 |
+
audience,
|
| 295 |
+
max_clock_skew_seconds,
|
| 296 |
+
} => verify_signed_bearer_token(
|
| 297 |
+
token,
|
| 298 |
+
shared_secret,
|
| 299 |
+
issuer.as_deref(),
|
| 300 |
+
audience.as_deref(),
|
| 301 |
+
*max_clock_skew_seconds,
|
| 302 |
+
),
|
| 303 |
+
}
|
| 304 |
+
}
|
| 305 |
+
|
| 306 |
+
fn verify_signed_bearer_token(
|
| 307 |
+
token: &str,
|
| 308 |
+
shared_secret: &[u8],
|
| 309 |
+
issuer: Option<&str>,
|
| 310 |
+
audience: Option<&str>,
|
| 311 |
+
max_clock_skew_seconds: i64,
|
| 312 |
+
) -> Result<(), WebsocketAuthError> {
|
| 313 |
+
let claims = decode_jwt_claims(token, shared_secret)?;
|
| 314 |
+
validate_jwt_claims(&claims, issuer, audience, max_clock_skew_seconds)
|
| 315 |
+
}
|
| 316 |
+
|
| 317 |
+
fn decode_jwt_claims(token: &str, shared_secret: &[u8]) -> Result<JwtClaims, WebsocketAuthError> {
|
| 318 |
+
let mut validation = Validation::new(Algorithm::HS256);
|
| 319 |
+
validation.required_spec_claims.clear();
|
| 320 |
+
validation.validate_exp = false;
|
| 321 |
+
validation.validate_nbf = false;
|
| 322 |
+
validation.validate_aud = false;
|
| 323 |
+
|
| 324 |
+
decode::<JwtClaims>(token, &DecodingKey::from_secret(shared_secret), &validation)
|
| 325 |
+
.map(|token_data| token_data.claims)
|
| 326 |
+
.map_err(|_| unauthorized("invalid websocket jwt"))
|
| 327 |
+
}
|
| 328 |
+
|
| 329 |
+
fn validate_jwt_claims(
|
| 330 |
+
claims: &JwtClaims,
|
| 331 |
+
issuer: Option<&str>,
|
| 332 |
+
audience: Option<&str>,
|
| 333 |
+
max_clock_skew_seconds: i64,
|
| 334 |
+
) -> Result<(), WebsocketAuthError> {
|
| 335 |
+
let now = OffsetDateTime::now_utc().unix_timestamp();
|
| 336 |
+
if now > claims.exp.saturating_add(max_clock_skew_seconds) {
|
| 337 |
+
return Err(unauthorized("expired websocket jwt"));
|
| 338 |
+
}
|
| 339 |
+
if let Some(nbf) = claims.nbf
|
| 340 |
+
&& now < nbf.saturating_sub(max_clock_skew_seconds)
|
| 341 |
+
{
|
| 342 |
+
return Err(unauthorized("websocket jwt is not valid yet"));
|
| 343 |
+
}
|
| 344 |
+
if let Some(expected_issuer) = issuer
|
| 345 |
+
&& claims.iss.as_deref() != Some(expected_issuer)
|
| 346 |
+
{
|
| 347 |
+
return Err(unauthorized("websocket jwt issuer mismatch"));
|
| 348 |
+
}
|
| 349 |
+
if let Some(expected_audience) = audience
|
| 350 |
+
&& !audience_matches(claims.aud.as_ref(), expected_audience)
|
| 351 |
+
{
|
| 352 |
+
return Err(unauthorized("websocket jwt audience mismatch"));
|
| 353 |
+
}
|
| 354 |
+
|
| 355 |
+
Ok(())
|
| 356 |
+
}
|
| 357 |
+
|
| 358 |
+
fn audience_matches(audience: Option<&JwtAudienceClaim>, expected_audience: &str) -> bool {
|
| 359 |
+
match audience {
|
| 360 |
+
Some(JwtAudienceClaim::Single(actual)) => actual == expected_audience,
|
| 361 |
+
Some(JwtAudienceClaim::Multiple(actual)) => {
|
| 362 |
+
actual.iter().any(|audience| audience == expected_audience)
|
| 363 |
+
}
|
| 364 |
+
None => false,
|
| 365 |
+
}
|
| 366 |
+
}
|
| 367 |
+
|
| 368 |
+
fn bearer_token_from_headers(headers: &HeaderMap) -> Result<&str, WebsocketAuthError> {
|
| 369 |
+
let raw_header = headers
|
| 370 |
+
.get(AUTHORIZATION)
|
| 371 |
+
.ok_or_else(|| unauthorized("missing websocket bearer token"))?;
|
| 372 |
+
let header = raw_header
|
| 373 |
+
.to_str()
|
| 374 |
+
.map_err(|_| unauthorized(INVALID_AUTHORIZATION_HEADER_MESSAGE))?;
|
| 375 |
+
let Some((scheme, token)) = header.split_once(' ') else {
|
| 376 |
+
return Err(unauthorized(INVALID_AUTHORIZATION_HEADER_MESSAGE));
|
| 377 |
+
};
|
| 378 |
+
if !scheme.eq_ignore_ascii_case("Bearer") {
|
| 379 |
+
return Err(unauthorized(INVALID_AUTHORIZATION_HEADER_MESSAGE));
|
| 380 |
+
}
|
| 381 |
+
let token = token.trim();
|
| 382 |
+
if token.is_empty() {
|
| 383 |
+
return Err(unauthorized(INVALID_AUTHORIZATION_HEADER_MESSAGE));
|
| 384 |
+
}
|
| 385 |
+
Ok(token)
|
| 386 |
+
}
|
| 387 |
+
|
| 388 |
+
fn validate_signed_bearer_secret(path: &Path, shared_secret: &[u8]) -> io::Result<()> {
|
| 389 |
+
if shared_secret.len() < MIN_SIGNED_BEARER_SECRET_BYTES {
|
| 390 |
+
return Err(io::Error::new(
|
| 391 |
+
ErrorKind::InvalidInput,
|
| 392 |
+
format!(
|
| 393 |
+
"signed websocket bearer secret {} must be at least {MIN_SIGNED_BEARER_SECRET_BYTES} bytes",
|
| 394 |
+
path.display()
|
| 395 |
+
),
|
| 396 |
+
));
|
| 397 |
+
}
|
| 398 |
+
Ok(())
|
| 399 |
+
}
|
| 400 |
+
|
| 401 |
+
fn read_trimmed_secret(path: &std::path::Path) -> io::Result<String> {
|
| 402 |
+
let raw = std::fs::read_to_string(path).map_err(|err| {
|
| 403 |
+
io::Error::new(
|
| 404 |
+
err.kind(),
|
| 405 |
+
format!(
|
| 406 |
+
"failed to read websocket auth secret {}: {err}",
|
| 407 |
+
path.display()
|
| 408 |
+
),
|
| 409 |
+
)
|
| 410 |
+
})?;
|
| 411 |
+
let trimmed = raw.trim();
|
| 412 |
+
if trimmed.is_empty() {
|
| 413 |
+
return Err(io::Error::new(
|
| 414 |
+
ErrorKind::InvalidInput,
|
| 415 |
+
format!("websocket auth secret {} must not be empty", path.display()),
|
| 416 |
+
));
|
| 417 |
+
}
|
| 418 |
+
Ok(trimmed.to_string())
|
| 419 |
+
}
|
| 420 |
+
|
| 421 |
+
fn absolute_path_arg(flag_name: &str, path: PathBuf) -> anyhow::Result<AbsolutePathBuf> {
|
| 422 |
+
AbsolutePathBuf::try_from(path).with_context(|| format!("{flag_name} must be an absolute path"))
|
| 423 |
+
}
|
| 424 |
+
|
| 425 |
+
fn sha256_digest_arg(flag_name: &str, value: &str) -> anyhow::Result<[u8; 32]> {
|
| 426 |
+
let trimmed = value.trim();
|
| 427 |
+
if trimmed.len() != 64 {
|
| 428 |
+
anyhow::bail!("{flag_name} must be a 64-character hex SHA-256 digest");
|
| 429 |
+
}
|
| 430 |
+
|
| 431 |
+
let mut digest = [0u8; 32];
|
| 432 |
+
for (index, pair) in trimmed.as_bytes().chunks_exact(2).enumerate() {
|
| 433 |
+
let high = hex_nibble(flag_name, pair[0])?;
|
| 434 |
+
let low = hex_nibble(flag_name, pair[1])?;
|
| 435 |
+
digest[index] = (high << 4) | low;
|
| 436 |
+
}
|
| 437 |
+
Ok(digest)
|
| 438 |
+
}
|
| 439 |
+
|
| 440 |
+
fn hex_nibble(flag_name: &str, byte: u8) -> anyhow::Result<u8> {
|
| 441 |
+
match byte {
|
| 442 |
+
b'0'..=b'9' => Ok(byte - b'0'),
|
| 443 |
+
b'a'..=b'f' => Ok(byte - b'a' + 10),
|
| 444 |
+
b'A'..=b'F' => Ok(byte - b'A' + 10),
|
| 445 |
+
_ => anyhow::bail!("{flag_name} must be a 64-character hex SHA-256 digest"),
|
| 446 |
+
}
|
| 447 |
+
}
|
| 448 |
+
|
| 449 |
+
fn sha256_digest(input: &[u8]) -> [u8; 32] {
|
| 450 |
+
let mut digest = [0u8; 32];
|
| 451 |
+
digest.copy_from_slice(&Sha256::digest(input));
|
| 452 |
+
digest
|
| 453 |
+
}
|
| 454 |
+
|
| 455 |
+
fn unauthorized(message: &'static str) -> WebsocketAuthError {
|
| 456 |
+
WebsocketAuthError {
|
| 457 |
+
status_code: StatusCode::UNAUTHORIZED,
|
| 458 |
+
message,
|
| 459 |
+
}
|
| 460 |
+
}
|
| 461 |
+
|
| 462 |
+
#[cfg(test)]
|
| 463 |
+
mod tests {
|
| 464 |
+
use super::*;
|
| 465 |
+
use axum::http::HeaderValue;
|
| 466 |
+
use base64::Engine;
|
| 467 |
+
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
| 468 |
+
use hmac::Hmac;
|
| 469 |
+
use hmac::Mac;
|
| 470 |
+
use serde_json::json;
|
| 471 |
+
|
| 472 |
+
type HmacSha256 = Hmac<Sha256>;
|
| 473 |
+
|
| 474 |
+
fn signed_token(shared_secret: &[u8], claims: serde_json::Value) -> String {
|
| 475 |
+
let header = URL_SAFE_NO_PAD.encode(br#"{"alg":"HS256","typ":"JWT"}"#);
|
| 476 |
+
let claims_segment = URL_SAFE_NO_PAD.encode(serde_json::to_vec(&claims).unwrap());
|
| 477 |
+
let payload = format!("{header}.{claims_segment}");
|
| 478 |
+
let mut mac = HmacSha256::new_from_slice(shared_secret).unwrap();
|
| 479 |
+
mac.update(payload.as_bytes());
|
| 480 |
+
let signature = URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes());
|
| 481 |
+
format!("{payload}.{signature}")
|
| 482 |
+
}
|
| 483 |
+
|
| 484 |
+
#[test]
|
| 485 |
+
fn detects_unauthenticated_non_loopback_listener() {
|
| 486 |
+
let policy = WebsocketAuthPolicy::default();
|
| 487 |
+
assert!(is_unauthenticated_non_loopback_listener(
|
| 488 |
+
"0.0.0.0:8765".parse().unwrap(),
|
| 489 |
+
&policy,
|
| 490 |
+
));
|
| 491 |
+
assert!(!is_unauthenticated_non_loopback_listener(
|
| 492 |
+
"127.0.0.1:8765".parse().unwrap(),
|
| 493 |
+
&policy,
|
| 494 |
+
));
|
| 495 |
+
assert!(!is_unauthenticated_non_loopback_listener(
|
| 496 |
+
"0.0.0.0:8765".parse().unwrap(),
|
| 497 |
+
&WebsocketAuthPolicy {
|
| 498 |
+
mode: Some(WebsocketAuthMode::CapabilityToken {
|
| 499 |
+
token_sha256: [0u8; 32],
|
| 500 |
+
}),
|
| 501 |
+
},
|
| 502 |
+
));
|
| 503 |
+
}
|
| 504 |
+
|
| 505 |
+
#[test]
|
| 506 |
+
fn capability_token_args_require_token_file_or_hash() {
|
| 507 |
+
let err = AppServerWebsocketAuthArgs {
|
| 508 |
+
ws_auth: Some(WebsocketAuthCliMode::CapabilityToken),
|
| 509 |
+
..Default::default()
|
| 510 |
+
}
|
| 511 |
+
.try_into_settings()
|
| 512 |
+
.expect_err("capability-token mode should require a token source");
|
| 513 |
+
assert!(
|
| 514 |
+
err.to_string().contains("--ws-token-file")
|
| 515 |
+
&& err.to_string().contains("--ws-token-sha256"),
|
| 516 |
+
"unexpected error: {err}"
|
| 517 |
+
);
|
| 518 |
+
}
|
| 519 |
+
|
| 520 |
+
#[test]
|
| 521 |
+
fn capability_token_args_accept_token_hash() {
|
| 522 |
+
let settings = AppServerWebsocketAuthArgs {
|
| 523 |
+
ws_auth: Some(WebsocketAuthCliMode::CapabilityToken),
|
| 524 |
+
ws_token_sha256: Some("ab".repeat(32)),
|
| 525 |
+
..Default::default()
|
| 526 |
+
}
|
| 527 |
+
.try_into_settings()
|
| 528 |
+
.expect("capability-token hash args should parse");
|
| 529 |
+
|
| 530 |
+
assert_eq!(
|
| 531 |
+
settings,
|
| 532 |
+
AppServerWebsocketAuthSettings {
|
| 533 |
+
config: Some(AppServerWebsocketAuthConfig::CapabilityToken {
|
| 534 |
+
source: AppServerWebsocketCapabilityTokenSource::TokenSha256 {
|
| 535 |
+
token_sha256: [0xab; 32],
|
| 536 |
+
},
|
| 537 |
+
}),
|
| 538 |
+
}
|
| 539 |
+
);
|
| 540 |
+
}
|
| 541 |
+
|
| 542 |
+
#[test]
|
| 543 |
+
fn capability_token_args_reject_multiple_token_sources() {
|
| 544 |
+
let err = AppServerWebsocketAuthArgs {
|
| 545 |
+
ws_auth: Some(WebsocketAuthCliMode::CapabilityToken),
|
| 546 |
+
ws_token_file: Some(PathBuf::from("/tmp/token")),
|
| 547 |
+
ws_token_sha256: Some("ab".repeat(32)),
|
| 548 |
+
..Default::default()
|
| 549 |
+
}
|
| 550 |
+
.try_into_settings()
|
| 551 |
+
.expect_err("capability-token mode should reject multiple token sources");
|
| 552 |
+
assert!(
|
| 553 |
+
err.to_string().contains("mutually exclusive"),
|
| 554 |
+
"unexpected error: {err}"
|
| 555 |
+
);
|
| 556 |
+
}
|
| 557 |
+
|
| 558 |
+
#[test]
|
| 559 |
+
fn capability_token_args_reject_malformed_token_hash() {
|
| 560 |
+
let err = AppServerWebsocketAuthArgs {
|
| 561 |
+
ws_auth: Some(WebsocketAuthCliMode::CapabilityToken),
|
| 562 |
+
ws_token_sha256: Some("not-a-sha256".to_string()),
|
| 563 |
+
..Default::default()
|
| 564 |
+
}
|
| 565 |
+
.try_into_settings()
|
| 566 |
+
.expect_err("capability-token mode should reject malformed token hashes");
|
| 567 |
+
assert!(
|
| 568 |
+
err.to_string().contains("64-character hex"),
|
| 569 |
+
"unexpected error: {err}"
|
| 570 |
+
);
|
| 571 |
+
}
|
| 572 |
+
|
| 573 |
+
#[test]
|
| 574 |
+
fn capability_token_hash_policy_authorizes_matching_bearer_token() {
|
| 575 |
+
let settings = AppServerWebsocketAuthSettings {
|
| 576 |
+
config: Some(AppServerWebsocketAuthConfig::CapabilityToken {
|
| 577 |
+
source: AppServerWebsocketCapabilityTokenSource::TokenSha256 {
|
| 578 |
+
token_sha256: sha256_digest(b"super-secret-token"),
|
| 579 |
+
},
|
| 580 |
+
}),
|
| 581 |
+
};
|
| 582 |
+
let policy = policy_from_settings(&settings).expect("hash policy should build");
|
| 583 |
+
let mut headers = HeaderMap::new();
|
| 584 |
+
headers.insert(
|
| 585 |
+
AUTHORIZATION,
|
| 586 |
+
HeaderValue::from_static("Bearer super-secret-token"),
|
| 587 |
+
);
|
| 588 |
+
authorize_upgrade(&headers, &policy).expect("matching token should authorize");
|
| 589 |
+
|
| 590 |
+
headers.insert(
|
| 591 |
+
AUTHORIZATION,
|
| 592 |
+
HeaderValue::from_static("Bearer wrong-token"),
|
| 593 |
+
);
|
| 594 |
+
let err = authorize_upgrade(&headers, &policy).expect_err("wrong token should fail");
|
| 595 |
+
assert_eq!(err.status_code(), StatusCode::UNAUTHORIZED);
|
| 596 |
+
}
|
| 597 |
+
|
| 598 |
+
#[test]
|
| 599 |
+
fn signed_bearer_args_require_mode_when_mode_specific_flags_are_set() {
|
| 600 |
+
let err = AppServerWebsocketAuthArgs {
|
| 601 |
+
ws_shared_secret_file: Some(PathBuf::from("/tmp/secret")),
|
| 602 |
+
..Default::default()
|
| 603 |
+
}
|
| 604 |
+
.try_into_settings()
|
| 605 |
+
.expect_err("mode-specific flags should require --ws-auth");
|
| 606 |
+
assert!(
|
| 607 |
+
err.to_string().contains("websocket auth flags require"),
|
| 608 |
+
"unexpected error: {err}"
|
| 609 |
+
);
|
| 610 |
+
}
|
| 611 |
+
|
| 612 |
+
#[test]
|
| 613 |
+
fn signed_bearer_args_default_clock_skew_and_trim_optional_claims() {
|
| 614 |
+
let settings = AppServerWebsocketAuthArgs {
|
| 615 |
+
ws_auth: Some(WebsocketAuthCliMode::SignedBearerToken),
|
| 616 |
+
ws_shared_secret_file: Some(PathBuf::from("/tmp/secret")),
|
| 617 |
+
ws_issuer: Some(" issuer ".to_string()),
|
| 618 |
+
ws_audience: Some(" ".to_string()),
|
| 619 |
+
..Default::default()
|
| 620 |
+
}
|
| 621 |
+
.try_into_settings()
|
| 622 |
+
.expect("signed bearer args should parse");
|
| 623 |
+
|
| 624 |
+
assert_eq!(
|
| 625 |
+
settings,
|
| 626 |
+
AppServerWebsocketAuthSettings {
|
| 627 |
+
config: Some(AppServerWebsocketAuthConfig::SignedBearerToken {
|
| 628 |
+
shared_secret_file: AbsolutePathBuf::from_absolute_path("/tmp/secret")
|
| 629 |
+
.expect("absolute path"),
|
| 630 |
+
issuer: Some("issuer".to_string()),
|
| 631 |
+
audience: None,
|
| 632 |
+
max_clock_skew_seconds: DEFAULT_MAX_CLOCK_SKEW_SECONDS,
|
| 633 |
+
}),
|
| 634 |
+
}
|
| 635 |
+
);
|
| 636 |
+
}
|
| 637 |
+
|
| 638 |
+
#[test]
|
| 639 |
+
fn signed_bearer_token_verification_rejects_tampering() {
|
| 640 |
+
let shared_secret = b"0123456789abcdef0123456789abcdef";
|
| 641 |
+
let token = signed_token(
|
| 642 |
+
shared_secret,
|
| 643 |
+
json!({
|
| 644 |
+
"exp": OffsetDateTime::now_utc().unix_timestamp() + 60,
|
| 645 |
+
}),
|
| 646 |
+
);
|
| 647 |
+
let tampered = token.replace(".eyJleHAi", ".eyJleHBi");
|
| 648 |
+
let err = verify_signed_bearer_token(
|
| 649 |
+
&tampered,
|
| 650 |
+
shared_secret,
|
| 651 |
+
/*issuer*/ None,
|
| 652 |
+
/*audience*/ None,
|
| 653 |
+
/*max_clock_skew_seconds*/ 30,
|
| 654 |
+
)
|
| 655 |
+
.expect_err("tampered jwt should fail");
|
| 656 |
+
assert_eq!(err.status_code(), StatusCode::UNAUTHORIZED);
|
| 657 |
+
}
|
| 658 |
+
|
| 659 |
+
#[test]
|
| 660 |
+
fn signed_bearer_token_verification_accepts_valid_token() {
|
| 661 |
+
let shared_secret = b"0123456789abcdef0123456789abcdef";
|
| 662 |
+
let token = signed_token(
|
| 663 |
+
shared_secret,
|
| 664 |
+
json!({
|
| 665 |
+
"exp": OffsetDateTime::now_utc().unix_timestamp() + 60,
|
| 666 |
+
"iss": "issuer",
|
| 667 |
+
"aud": "audience",
|
| 668 |
+
}),
|
| 669 |
+
);
|
| 670 |
+
verify_signed_bearer_token(
|
| 671 |
+
&token,
|
| 672 |
+
shared_secret,
|
| 673 |
+
Some("issuer"),
|
| 674 |
+
Some("audience"),
|
| 675 |
+
/*max_clock_skew_seconds*/ 30,
|
| 676 |
+
)
|
| 677 |
+
.expect("valid signed token should verify");
|
| 678 |
+
}
|
| 679 |
+
|
| 680 |
+
#[test]
|
| 681 |
+
fn signed_bearer_token_verification_accepts_multiple_audiences() {
|
| 682 |
+
let shared_secret = b"0123456789abcdef0123456789abcdef";
|
| 683 |
+
let token = signed_token(
|
| 684 |
+
shared_secret,
|
| 685 |
+
json!({
|
| 686 |
+
"exp": OffsetDateTime::now_utc().unix_timestamp() + 60,
|
| 687 |
+
"aud": ["other-audience", "audience"],
|
| 688 |
+
}),
|
| 689 |
+
);
|
| 690 |
+
verify_signed_bearer_token(
|
| 691 |
+
&token,
|
| 692 |
+
shared_secret,
|
| 693 |
+
/*issuer*/ None,
|
| 694 |
+
Some("audience"),
|
| 695 |
+
/*max_clock_skew_seconds*/ 30,
|
| 696 |
+
)
|
| 697 |
+
.expect("jwt audience arrays should verify");
|
| 698 |
+
}
|
| 699 |
+
|
| 700 |
+
#[test]
|
| 701 |
+
fn signed_bearer_token_verification_rejects_alg_none_tokens() {
|
| 702 |
+
let claims_segment = URL_SAFE_NO_PAD.encode(
|
| 703 |
+
serde_json::to_vec(&json!({
|
| 704 |
+
"exp": OffsetDateTime::now_utc().unix_timestamp() + 60,
|
| 705 |
+
}))
|
| 706 |
+
.unwrap(),
|
| 707 |
+
);
|
| 708 |
+
let header_segment = URL_SAFE_NO_PAD.encode(br#"{"alg":"none","typ":"JWT"}"#);
|
| 709 |
+
let token = format!("{header_segment}.{claims_segment}.");
|
| 710 |
+
let err = verify_signed_bearer_token(
|
| 711 |
+
&token,
|
| 712 |
+
b"0123456789abcdef0123456789abcdef",
|
| 713 |
+
/*issuer*/ None,
|
| 714 |
+
/*audience*/ None,
|
| 715 |
+
/*max_clock_skew_seconds*/ 30,
|
| 716 |
+
)
|
| 717 |
+
.expect_err("alg=none jwt should be rejected");
|
| 718 |
+
assert_eq!(err.status_code(), StatusCode::UNAUTHORIZED);
|
| 719 |
+
}
|
| 720 |
+
|
| 721 |
+
#[test]
|
| 722 |
+
fn signed_bearer_token_verification_rejects_missing_exp() {
|
| 723 |
+
let shared_secret = b"0123456789abcdef0123456789abcdef";
|
| 724 |
+
let token = signed_token(
|
| 725 |
+
shared_secret,
|
| 726 |
+
json!({
|
| 727 |
+
"iss": "issuer",
|
| 728 |
+
}),
|
| 729 |
+
);
|
| 730 |
+
let err = verify_signed_bearer_token(
|
| 731 |
+
&token,
|
| 732 |
+
shared_secret,
|
| 733 |
+
/*issuer*/ None,
|
| 734 |
+
/*audience*/ None,
|
| 735 |
+
/*max_clock_skew_seconds*/ 30,
|
| 736 |
+
)
|
| 737 |
+
.expect_err("jwt without exp should be rejected");
|
| 738 |
+
assert_eq!(err.status_code(), StatusCode::UNAUTHORIZED);
|
| 739 |
+
}
|
| 740 |
+
|
| 741 |
+
#[test]
|
| 742 |
+
fn validate_signed_bearer_secret_rejects_short_secret() {
|
| 743 |
+
let err = validate_signed_bearer_secret(Path::new("/tmp/secret"), b"too-short")
|
| 744 |
+
.expect_err("short shared secret should be rejected");
|
| 745 |
+
assert_eq!(err.kind(), ErrorKind::InvalidInput);
|
| 746 |
+
assert!(
|
| 747 |
+
err.to_string().contains("must be at least 32 bytes"),
|
| 748 |
+
"unexpected error: {err}"
|
| 749 |
+
);
|
| 750 |
+
}
|
| 751 |
+
}
|
codex-rs/app-server-transport/src/transport/mod.rs
ADDED
|
@@ -0,0 +1,602 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
pub mod auth;
|
| 2 |
+
|
| 3 |
+
use crate::outgoing_message::ConnectionId;
|
| 4 |
+
use crate::outgoing_message::OutgoingError;
|
| 5 |
+
use crate::outgoing_message::OutgoingMessage;
|
| 6 |
+
use crate::outgoing_message::QueuedOutgoingMessage;
|
| 7 |
+
use codex_app_server_protocol::JSONRPCErrorError;
|
| 8 |
+
use codex_app_server_protocol::JSONRPCMessage;
|
| 9 |
+
use codex_app_server_protocol::RequestId;
|
| 10 |
+
use codex_core::config::find_codex_home;
|
| 11 |
+
use codex_utils_absolute_path::AbsolutePathBuf;
|
| 12 |
+
use std::net::SocketAddr;
|
| 13 |
+
use std::path::Path;
|
| 14 |
+
use std::path::PathBuf;
|
| 15 |
+
use std::str::FromStr;
|
| 16 |
+
use std::sync::atomic::AtomicU64;
|
| 17 |
+
use std::sync::atomic::Ordering;
|
| 18 |
+
use tokio::sync::mpsc;
|
| 19 |
+
use tokio_util::sync::CancellationToken;
|
| 20 |
+
use tracing::error;
|
| 21 |
+
use tracing::warn;
|
| 22 |
+
|
| 23 |
+
/// Size of the bounded channels used to communicate between tasks. The value
|
| 24 |
+
/// is a balance between throughput and memory usage - 128 messages should be
|
| 25 |
+
/// plenty for an interactive CLI.
|
| 26 |
+
pub const CHANNEL_CAPACITY: usize = 128;
|
| 27 |
+
|
| 28 |
+
mod remote_control;
|
| 29 |
+
mod stdio;
|
| 30 |
+
mod unix_socket;
|
| 31 |
+
#[cfg(test)]
|
| 32 |
+
mod unix_socket_tests;
|
| 33 |
+
mod websocket;
|
| 34 |
+
|
| 35 |
+
pub use remote_control::REMOTE_CONTROL_DISABLED_ENV_VAR;
|
| 36 |
+
pub use remote_control::RemoteControlDisabledByRequirements;
|
| 37 |
+
pub use remote_control::RemoteControlEnableError;
|
| 38 |
+
pub use remote_control::RemoteControlHandle;
|
| 39 |
+
pub use remote_control::RemoteControlPolicy;
|
| 40 |
+
pub use remote_control::RemoteControlStartConfig;
|
| 41 |
+
pub use remote_control::RemoteControlStartupMode;
|
| 42 |
+
pub use remote_control::RemoteControlUnavailable;
|
| 43 |
+
pub use remote_control::start_remote_control;
|
| 44 |
+
pub use remote_control::take_remote_control_disabled_env;
|
| 45 |
+
pub use stdio::start_stdio_connection;
|
| 46 |
+
pub use unix_socket::AppServerStartupLock;
|
| 47 |
+
pub use unix_socket::DaemonShutdownAccess;
|
| 48 |
+
pub use unix_socket::acquire_app_server_startup_lock;
|
| 49 |
+
pub use unix_socket::prepare_control_socket_path;
|
| 50 |
+
pub use unix_socket::start_control_socket_acceptor;
|
| 51 |
+
pub use websocket::start_websocket_acceptor;
|
| 52 |
+
|
| 53 |
+
const INTERNAL_ERROR_CODE: i64 = -32603;
|
| 54 |
+
const OVERLOADED_ERROR_CODE: i64 = -32001;
|
| 55 |
+
|
| 56 |
+
const APP_SERVER_CONTROL_SOCKET_DIR_NAME: &str = "app-server-control";
|
| 57 |
+
const APP_SERVER_CONTROL_SOCKET_FILE_NAME: &str = "app-server-control.sock";
|
| 58 |
+
const APP_SERVER_STARTUP_LOCK_FILE_NAME: &str = "app-server-startup.lock";
|
| 59 |
+
const DAEMON_RECOVERY_FILE_NAME: &str = "loaded-threads.json";
|
| 60 |
+
|
| 61 |
+
pub fn daemon_recovery_file_path(codex_home: &Path) -> PathBuf {
|
| 62 |
+
codex_home
|
| 63 |
+
.join("app-server-daemon")
|
| 64 |
+
.join(DAEMON_RECOVERY_FILE_NAME)
|
| 65 |
+
}
|
| 66 |
+
|
| 67 |
+
pub fn app_server_control_socket_path(codex_home: &Path) -> std::io::Result<AbsolutePathBuf> {
|
| 68 |
+
AbsolutePathBuf::from_absolute_path(
|
| 69 |
+
codex_home
|
| 70 |
+
.join(APP_SERVER_CONTROL_SOCKET_DIR_NAME)
|
| 71 |
+
.join(APP_SERVER_CONTROL_SOCKET_FILE_NAME),
|
| 72 |
+
)
|
| 73 |
+
}
|
| 74 |
+
|
| 75 |
+
pub fn app_server_startup_lock_path(codex_home: &Path) -> std::io::Result<AbsolutePathBuf> {
|
| 76 |
+
AbsolutePathBuf::from_absolute_path(
|
| 77 |
+
codex_home
|
| 78 |
+
.join(APP_SERVER_CONTROL_SOCKET_DIR_NAME)
|
| 79 |
+
.join(APP_SERVER_STARTUP_LOCK_FILE_NAME),
|
| 80 |
+
)
|
| 81 |
+
}
|
| 82 |
+
|
| 83 |
+
#[derive(Clone, Debug, Eq, PartialEq)]
|
| 84 |
+
pub enum AppServerTransport {
|
| 85 |
+
Stdio,
|
| 86 |
+
UnixSocket { socket_path: AbsolutePathBuf },
|
| 87 |
+
WebSocket { bind_address: SocketAddr },
|
| 88 |
+
Off,
|
| 89 |
+
}
|
| 90 |
+
|
| 91 |
+
#[derive(Debug, Clone, Eq, PartialEq)]
|
| 92 |
+
pub enum AppServerTransportParseError {
|
| 93 |
+
UnsupportedListenUrl(String),
|
| 94 |
+
InvalidUnixSocketPath { listen_url: String, message: String },
|
| 95 |
+
InvalidWebSocketListenUrl(String),
|
| 96 |
+
}
|
| 97 |
+
|
| 98 |
+
impl std::fmt::Display for AppServerTransportParseError {
|
| 99 |
+
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
| 100 |
+
match self {
|
| 101 |
+
AppServerTransportParseError::UnsupportedListenUrl(listen_url) => write!(
|
| 102 |
+
f,
|
| 103 |
+
"unsupported --listen URL `{listen_url}`; expected `stdio://`, `unix://`, `unix://PATH`, `ws://IP:PORT`, or `off`"
|
| 104 |
+
),
|
| 105 |
+
AppServerTransportParseError::InvalidUnixSocketPath {
|
| 106 |
+
listen_url,
|
| 107 |
+
message,
|
| 108 |
+
} => write!(
|
| 109 |
+
f,
|
| 110 |
+
"invalid unix socket --listen URL `{listen_url}`; failed to resolve socket path: {message}"
|
| 111 |
+
),
|
| 112 |
+
AppServerTransportParseError::InvalidWebSocketListenUrl(listen_url) => write!(
|
| 113 |
+
f,
|
| 114 |
+
"invalid websocket --listen URL `{listen_url}`; expected `ws://IP:PORT`"
|
| 115 |
+
),
|
| 116 |
+
}
|
| 117 |
+
}
|
| 118 |
+
}
|
| 119 |
+
|
| 120 |
+
impl std::error::Error for AppServerTransportParseError {}
|
| 121 |
+
|
| 122 |
+
impl AppServerTransport {
|
| 123 |
+
pub const DEFAULT_LISTEN_URL: &'static str = "stdio://";
|
| 124 |
+
|
| 125 |
+
pub fn from_listen_url(listen_url: &str) -> Result<Self, AppServerTransportParseError> {
|
| 126 |
+
if listen_url == Self::DEFAULT_LISTEN_URL {
|
| 127 |
+
return Ok(Self::Stdio);
|
| 128 |
+
}
|
| 129 |
+
|
| 130 |
+
if let Some(raw_socket_path) = listen_url.strip_prefix("unix://") {
|
| 131 |
+
let socket_path = if raw_socket_path.is_empty() {
|
| 132 |
+
let codex_home = find_codex_home().map_err(|err| {
|
| 133 |
+
AppServerTransportParseError::InvalidUnixSocketPath {
|
| 134 |
+
listen_url: listen_url.to_string(),
|
| 135 |
+
message: format!("failed to resolve CODEX_HOME: {err}"),
|
| 136 |
+
}
|
| 137 |
+
})?;
|
| 138 |
+
app_server_control_socket_path(&codex_home).map_err(|err| {
|
| 139 |
+
AppServerTransportParseError::InvalidUnixSocketPath {
|
| 140 |
+
listen_url: listen_url.to_string(),
|
| 141 |
+
message: err.to_string(),
|
| 142 |
+
}
|
| 143 |
+
})?
|
| 144 |
+
} else {
|
| 145 |
+
AbsolutePathBuf::relative_to_current_dir(raw_socket_path).map_err(|err| {
|
| 146 |
+
AppServerTransportParseError::InvalidUnixSocketPath {
|
| 147 |
+
listen_url: listen_url.to_string(),
|
| 148 |
+
message: err.to_string(),
|
| 149 |
+
}
|
| 150 |
+
})?
|
| 151 |
+
};
|
| 152 |
+
return Ok(Self::UnixSocket { socket_path });
|
| 153 |
+
}
|
| 154 |
+
|
| 155 |
+
if listen_url == "off" {
|
| 156 |
+
return Ok(Self::Off);
|
| 157 |
+
}
|
| 158 |
+
|
| 159 |
+
if let Some(socket_addr) = listen_url.strip_prefix("ws://") {
|
| 160 |
+
let bind_address = socket_addr.parse::<SocketAddr>().map_err(|_| {
|
| 161 |
+
AppServerTransportParseError::InvalidWebSocketListenUrl(listen_url.to_string())
|
| 162 |
+
})?;
|
| 163 |
+
return Ok(Self::WebSocket { bind_address });
|
| 164 |
+
}
|
| 165 |
+
|
| 166 |
+
Err(AppServerTransportParseError::UnsupportedListenUrl(
|
| 167 |
+
listen_url.to_string(),
|
| 168 |
+
))
|
| 169 |
+
}
|
| 170 |
+
}
|
| 171 |
+
|
| 172 |
+
impl FromStr for AppServerTransport {
|
| 173 |
+
type Err = AppServerTransportParseError;
|
| 174 |
+
|
| 175 |
+
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
| 176 |
+
Self::from_listen_url(s)
|
| 177 |
+
}
|
| 178 |
+
}
|
| 179 |
+
|
| 180 |
+
#[derive(Debug)]
|
| 181 |
+
pub enum TransportEvent {
|
| 182 |
+
/// Accepted on the managed local control socket, outside JSON-RPC.
|
| 183 |
+
DaemonShutdown,
|
| 184 |
+
ConnectionOpened {
|
| 185 |
+
connection_id: ConnectionId,
|
| 186 |
+
origin: ConnectionOrigin,
|
| 187 |
+
auth: Option<crate::ConnectionAuth>,
|
| 188 |
+
writer: mpsc::Sender<QueuedOutgoingMessage>,
|
| 189 |
+
disconnect_sender: Option<CancellationToken>,
|
| 190 |
+
},
|
| 191 |
+
ConnectionClosed {
|
| 192 |
+
connection_id: ConnectionId,
|
| 193 |
+
},
|
| 194 |
+
IncomingMessage {
|
| 195 |
+
connection_id: ConnectionId,
|
| 196 |
+
message: JSONRPCMessage,
|
| 197 |
+
},
|
| 198 |
+
}
|
| 199 |
+
|
| 200 |
+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
| 201 |
+
pub enum ConnectionOrigin {
|
| 202 |
+
Stdio,
|
| 203 |
+
InProcess,
|
| 204 |
+
WebSocket,
|
| 205 |
+
RemoteControl,
|
| 206 |
+
}
|
| 207 |
+
|
| 208 |
+
static CONNECTION_ID_COUNTER: AtomicU64 = AtomicU64::new(0);
|
| 209 |
+
|
| 210 |
+
fn next_connection_id() -> ConnectionId {
|
| 211 |
+
ConnectionId(CONNECTION_ID_COUNTER.fetch_add(1, Ordering::Relaxed))
|
| 212 |
+
}
|
| 213 |
+
|
| 214 |
+
async fn forward_incoming_message(
|
| 215 |
+
transport_event_tx: &mpsc::Sender<TransportEvent>,
|
| 216 |
+
writer: &mpsc::Sender<QueuedOutgoingMessage>,
|
| 217 |
+
connection_id: ConnectionId,
|
| 218 |
+
payload: &str,
|
| 219 |
+
) -> bool {
|
| 220 |
+
match serde_json::from_str::<JSONRPCMessage>(payload) {
|
| 221 |
+
Ok(message) => {
|
| 222 |
+
enqueue_incoming_message(transport_event_tx, writer, connection_id, message).await
|
| 223 |
+
}
|
| 224 |
+
Err(err) => {
|
| 225 |
+
error!("Failed to deserialize JSONRPCMessage: {err}");
|
| 226 |
+
true
|
| 227 |
+
}
|
| 228 |
+
}
|
| 229 |
+
}
|
| 230 |
+
|
| 231 |
+
async fn enqueue_incoming_message(
|
| 232 |
+
transport_event_tx: &mpsc::Sender<TransportEvent>,
|
| 233 |
+
writer: &mpsc::Sender<QueuedOutgoingMessage>,
|
| 234 |
+
connection_id: ConnectionId,
|
| 235 |
+
message: JSONRPCMessage,
|
| 236 |
+
) -> bool {
|
| 237 |
+
let event = TransportEvent::IncomingMessage {
|
| 238 |
+
connection_id,
|
| 239 |
+
message,
|
| 240 |
+
};
|
| 241 |
+
match transport_event_tx.try_send(event) {
|
| 242 |
+
Ok(()) => true,
|
| 243 |
+
Err(mpsc::error::TrySendError::Closed(_)) => false,
|
| 244 |
+
Err(mpsc::error::TrySendError::Full(TransportEvent::IncomingMessage {
|
| 245 |
+
connection_id,
|
| 246 |
+
message: JSONRPCMessage::Request(request),
|
| 247 |
+
})) => {
|
| 248 |
+
let overload_error = OutgoingMessage::Error(OutgoingError {
|
| 249 |
+
id: request.id,
|
| 250 |
+
error: JSONRPCErrorError {
|
| 251 |
+
code: OVERLOADED_ERROR_CODE,
|
| 252 |
+
message: "Server overloaded; retry later.".to_string(),
|
| 253 |
+
data: None,
|
| 254 |
+
},
|
| 255 |
+
});
|
| 256 |
+
match writer.try_send(QueuedOutgoingMessage::new(overload_error)) {
|
| 257 |
+
Ok(()) => true,
|
| 258 |
+
Err(mpsc::error::TrySendError::Closed(_)) => false,
|
| 259 |
+
Err(mpsc::error::TrySendError::Full(_overload_error)) => {
|
| 260 |
+
warn!(
|
| 261 |
+
"dropping overload response for connection {:?}: outbound queue is full",
|
| 262 |
+
connection_id
|
| 263 |
+
);
|
| 264 |
+
true
|
| 265 |
+
}
|
| 266 |
+
}
|
| 267 |
+
}
|
| 268 |
+
Err(mpsc::error::TrySendError::Full(event)) => transport_event_tx.send(event).await.is_ok(),
|
| 269 |
+
}
|
| 270 |
+
}
|
| 271 |
+
|
| 272 |
+
fn serialize_outgoing_message(outgoing_message: OutgoingMessage) -> Option<String> {
|
| 273 |
+
match serde_json::to_string(&outgoing_message) {
|
| 274 |
+
Ok(json) => Some(json),
|
| 275 |
+
Err(err) => {
|
| 276 |
+
error!("Failed to serialize JSONRPCMessage: {err}");
|
| 277 |
+
let OutgoingMessage::Response(response) = outgoing_message else {
|
| 278 |
+
return None;
|
| 279 |
+
};
|
| 280 |
+
serde_json::to_string(&response_serialization_error(response.id, err))
|
| 281 |
+
.inspect_err(|err| error!("Failed to serialize JSONRPC error: {err}"))
|
| 282 |
+
.ok()
|
| 283 |
+
}
|
| 284 |
+
}
|
| 285 |
+
}
|
| 286 |
+
|
| 287 |
+
fn response_serialization_error(
|
| 288 |
+
request_id: RequestId,
|
| 289 |
+
err: impl std::fmt::Display,
|
| 290 |
+
) -> OutgoingMessage {
|
| 291 |
+
OutgoingMessage::Error(OutgoingError {
|
| 292 |
+
id: request_id,
|
| 293 |
+
error: JSONRPCErrorError {
|
| 294 |
+
code: INTERNAL_ERROR_CODE,
|
| 295 |
+
message: format!("failed to serialize response: {err}"),
|
| 296 |
+
data: None,
|
| 297 |
+
},
|
| 298 |
+
})
|
| 299 |
+
}
|
| 300 |
+
|
| 301 |
+
#[cfg(test)]
|
| 302 |
+
mod tests {
|
| 303 |
+
use super::*;
|
| 304 |
+
use crate::outgoing_message::OutgoingResponse;
|
| 305 |
+
use codex_app_server_protocol::ClientResponsePayload;
|
| 306 |
+
use codex_app_server_protocol::ConfigWarningNotification;
|
| 307 |
+
use codex_app_server_protocol::JSONRPCNotification;
|
| 308 |
+
use codex_app_server_protocol::JSONRPCRequest;
|
| 309 |
+
use codex_app_server_protocol::JSONRPCResponse;
|
| 310 |
+
use codex_app_server_protocol::RequestId;
|
| 311 |
+
use codex_app_server_protocol::ServerNotification;
|
| 312 |
+
use codex_app_server_protocol::ServerNotificationEnvelope;
|
| 313 |
+
use codex_app_server_protocol::ThreadArchiveResponse;
|
| 314 |
+
use pretty_assertions::assert_eq;
|
| 315 |
+
use serde_json::json;
|
| 316 |
+
use tokio::time::Duration;
|
| 317 |
+
use tokio::time::timeout;
|
| 318 |
+
|
| 319 |
+
#[test]
|
| 320 |
+
fn listen_off_parses_as_off_transport() {
|
| 321 |
+
assert_eq!(
|
| 322 |
+
AppServerTransport::from_listen_url("off"),
|
| 323 |
+
Ok(AppServerTransport::Off)
|
| 324 |
+
);
|
| 325 |
+
}
|
| 326 |
+
|
| 327 |
+
#[test]
|
| 328 |
+
fn serialize_outgoing_message_preserves_wire_shape() {
|
| 329 |
+
let message = OutgoingMessage::AppServerNotification(ServerNotificationEnvelope {
|
| 330 |
+
notification: ServerNotification::ConfigWarning(ConfigWarningNotification {
|
| 331 |
+
summary: "summary".to_string(),
|
| 332 |
+
details: None,
|
| 333 |
+
path: None,
|
| 334 |
+
range: None,
|
| 335 |
+
}),
|
| 336 |
+
emitted_at_ms: Some(1_234),
|
| 337 |
+
});
|
| 338 |
+
|
| 339 |
+
let json = serialize_outgoing_message(message).expect("message should serialize");
|
| 340 |
+
assert_eq!(
|
| 341 |
+
serde_json::from_str::<serde_json::Value>(&json).expect("message should be valid JSON"),
|
| 342 |
+
json!({
|
| 343 |
+
"method": "configWarning",
|
| 344 |
+
"params": {
|
| 345 |
+
"summary": "summary",
|
| 346 |
+
"details": null,
|
| 347 |
+
},
|
| 348 |
+
"emittedAtMs": 1_234,
|
| 349 |
+
})
|
| 350 |
+
);
|
| 351 |
+
}
|
| 352 |
+
|
| 353 |
+
#[test]
|
| 354 |
+
fn serialize_typed_response_preserves_wire_shape() {
|
| 355 |
+
let message = OutgoingMessage::Response(OutgoingResponse {
|
| 356 |
+
id: RequestId::Integer(7),
|
| 357 |
+
result: Box::new(ClientResponsePayload::ThreadArchive(
|
| 358 |
+
ThreadArchiveResponse {},
|
| 359 |
+
)),
|
| 360 |
+
});
|
| 361 |
+
|
| 362 |
+
let json = serialize_outgoing_message(message).expect("message should serialize");
|
| 363 |
+
assert_eq!(
|
| 364 |
+
serde_json::from_str::<serde_json::Value>(&json).expect("message should be valid JSON"),
|
| 365 |
+
json!({ "id": 7, "result": {} })
|
| 366 |
+
);
|
| 367 |
+
}
|
| 368 |
+
|
| 369 |
+
#[cfg(unix)]
|
| 370 |
+
#[test]
|
| 371 |
+
fn serialize_invalid_typed_response_returns_jsonrpc_error() {
|
| 372 |
+
use std::ffi::OsString;
|
| 373 |
+
use std::os::unix::ffi::OsStringExt;
|
| 374 |
+
use std::path::PathBuf;
|
| 375 |
+
|
| 376 |
+
let codex_home =
|
| 377 |
+
AbsolutePathBuf::from_absolute_path(PathBuf::from(OsString::from_vec(vec![
|
| 378 |
+
b'/', b'b', b'a', b'd', 0xff,
|
| 379 |
+
])))
|
| 380 |
+
.expect("non-UTF-8 Unix paths are valid absolute paths");
|
| 381 |
+
let message = OutgoingMessage::Response(OutgoingResponse {
|
| 382 |
+
id: RequestId::Integer(7),
|
| 383 |
+
result: Box::new(ClientResponsePayload::Initialize(
|
| 384 |
+
codex_app_server_protocol::InitializeResponse {
|
| 385 |
+
user_agent: "codex-test-agent".to_string(),
|
| 386 |
+
codex_home,
|
| 387 |
+
platform_family: "unix".to_string(),
|
| 388 |
+
platform_os: "linux".to_string(),
|
| 389 |
+
},
|
| 390 |
+
)),
|
| 391 |
+
});
|
| 392 |
+
|
| 393 |
+
let json = serialize_outgoing_message(message)
|
| 394 |
+
.expect("invalid response should serialize as a JSON-RPC error");
|
| 395 |
+
assert_eq!(
|
| 396 |
+
serde_json::from_str::<serde_json::Value>(&json).expect("message should be valid JSON"),
|
| 397 |
+
json!({
|
| 398 |
+
"id": 7,
|
| 399 |
+
"error": {
|
| 400 |
+
"code": -32603,
|
| 401 |
+
"message": "failed to serialize response: path contains invalid UTF-8 characters",
|
| 402 |
+
}
|
| 403 |
+
})
|
| 404 |
+
);
|
| 405 |
+
}
|
| 406 |
+
|
| 407 |
+
#[tokio::test]
|
| 408 |
+
async fn enqueue_incoming_request_returns_overload_error_when_queue_is_full() {
|
| 409 |
+
let connection_id = ConnectionId(42);
|
| 410 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(1);
|
| 411 |
+
let (writer_tx, mut writer_rx) = mpsc::channel(1);
|
| 412 |
+
|
| 413 |
+
let first_message = JSONRPCMessage::Notification(JSONRPCNotification {
|
| 414 |
+
method: "initialized".to_string(),
|
| 415 |
+
params: None,
|
| 416 |
+
});
|
| 417 |
+
transport_event_tx
|
| 418 |
+
.send(TransportEvent::IncomingMessage {
|
| 419 |
+
connection_id,
|
| 420 |
+
message: first_message.clone(),
|
| 421 |
+
})
|
| 422 |
+
.await
|
| 423 |
+
.expect("queue should accept first message");
|
| 424 |
+
|
| 425 |
+
let request = JSONRPCMessage::Request(JSONRPCRequest {
|
| 426 |
+
id: RequestId::Integer(7),
|
| 427 |
+
method: "config/read".to_string(),
|
| 428 |
+
params: Some(json!({ "includeLayers": false })),
|
| 429 |
+
trace: None,
|
| 430 |
+
});
|
| 431 |
+
assert!(
|
| 432 |
+
enqueue_incoming_message(&transport_event_tx, &writer_tx, connection_id, request).await
|
| 433 |
+
);
|
| 434 |
+
|
| 435 |
+
let queued_event = transport_event_rx
|
| 436 |
+
.recv()
|
| 437 |
+
.await
|
| 438 |
+
.expect("first event should stay queued");
|
| 439 |
+
match queued_event {
|
| 440 |
+
TransportEvent::IncomingMessage {
|
| 441 |
+
connection_id: queued_connection_id,
|
| 442 |
+
message,
|
| 443 |
+
} => {
|
| 444 |
+
assert_eq!(queued_connection_id, connection_id);
|
| 445 |
+
assert_eq!(message, first_message);
|
| 446 |
+
}
|
| 447 |
+
_ => panic!("expected queued incoming message"),
|
| 448 |
+
}
|
| 449 |
+
|
| 450 |
+
let overload = writer_rx
|
| 451 |
+
.recv()
|
| 452 |
+
.await
|
| 453 |
+
.expect("request should receive overload error");
|
| 454 |
+
let overload_json =
|
| 455 |
+
serde_json::to_value(overload.message).expect("serialize overload error");
|
| 456 |
+
assert_eq!(
|
| 457 |
+
overload_json,
|
| 458 |
+
json!({
|
| 459 |
+
"id": 7,
|
| 460 |
+
"error": {
|
| 461 |
+
"code": OVERLOADED_ERROR_CODE,
|
| 462 |
+
"message": "Server overloaded; retry later."
|
| 463 |
+
}
|
| 464 |
+
})
|
| 465 |
+
);
|
| 466 |
+
}
|
| 467 |
+
|
| 468 |
+
#[tokio::test]
|
| 469 |
+
async fn enqueue_incoming_response_waits_instead_of_dropping_when_queue_is_full() {
|
| 470 |
+
let connection_id = ConnectionId(42);
|
| 471 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(1);
|
| 472 |
+
let (writer_tx, _writer_rx) = mpsc::channel(1);
|
| 473 |
+
|
| 474 |
+
let first_message = JSONRPCMessage::Notification(JSONRPCNotification {
|
| 475 |
+
method: "initialized".to_string(),
|
| 476 |
+
params: None,
|
| 477 |
+
});
|
| 478 |
+
transport_event_tx
|
| 479 |
+
.send(TransportEvent::IncomingMessage {
|
| 480 |
+
connection_id,
|
| 481 |
+
message: first_message.clone(),
|
| 482 |
+
})
|
| 483 |
+
.await
|
| 484 |
+
.expect("queue should accept first message");
|
| 485 |
+
|
| 486 |
+
let response = JSONRPCMessage::Response(JSONRPCResponse {
|
| 487 |
+
id: RequestId::Integer(7),
|
| 488 |
+
result: json!({"ok": true}),
|
| 489 |
+
});
|
| 490 |
+
let transport_event_tx_for_enqueue = transport_event_tx.clone();
|
| 491 |
+
let writer_tx_for_enqueue = writer_tx.clone();
|
| 492 |
+
let enqueue_handle = tokio::spawn(async move {
|
| 493 |
+
enqueue_incoming_message(
|
| 494 |
+
&transport_event_tx_for_enqueue,
|
| 495 |
+
&writer_tx_for_enqueue,
|
| 496 |
+
connection_id,
|
| 497 |
+
response,
|
| 498 |
+
)
|
| 499 |
+
.await
|
| 500 |
+
});
|
| 501 |
+
|
| 502 |
+
let queued_event = transport_event_rx
|
| 503 |
+
.recv()
|
| 504 |
+
.await
|
| 505 |
+
.expect("first event should be dequeued");
|
| 506 |
+
match queued_event {
|
| 507 |
+
TransportEvent::IncomingMessage {
|
| 508 |
+
connection_id: queued_connection_id,
|
| 509 |
+
message,
|
| 510 |
+
} => {
|
| 511 |
+
assert_eq!(queued_connection_id, connection_id);
|
| 512 |
+
assert_eq!(message, first_message);
|
| 513 |
+
}
|
| 514 |
+
_ => panic!("expected queued incoming message"),
|
| 515 |
+
}
|
| 516 |
+
|
| 517 |
+
let enqueue_result = enqueue_handle.await.expect("enqueue task should not panic");
|
| 518 |
+
assert!(enqueue_result);
|
| 519 |
+
|
| 520 |
+
let forwarded_event = transport_event_rx
|
| 521 |
+
.recv()
|
| 522 |
+
.await
|
| 523 |
+
.expect("response should be forwarded instead of dropped");
|
| 524 |
+
match forwarded_event {
|
| 525 |
+
TransportEvent::IncomingMessage {
|
| 526 |
+
connection_id: queued_connection_id,
|
| 527 |
+
message: JSONRPCMessage::Response(JSONRPCResponse { id, result }),
|
| 528 |
+
} => {
|
| 529 |
+
assert_eq!(queued_connection_id, connection_id);
|
| 530 |
+
assert_eq!(id, RequestId::Integer(7));
|
| 531 |
+
assert_eq!(result, json!({"ok": true}));
|
| 532 |
+
}
|
| 533 |
+
_ => panic!("expected forwarded response message"),
|
| 534 |
+
}
|
| 535 |
+
}
|
| 536 |
+
|
| 537 |
+
#[tokio::test]
|
| 538 |
+
async fn enqueue_incoming_request_does_not_block_when_writer_queue_is_full() {
|
| 539 |
+
let connection_id = ConnectionId(42);
|
| 540 |
+
let (transport_event_tx, _transport_event_rx) = mpsc::channel(1);
|
| 541 |
+
let (writer_tx, mut writer_rx) = mpsc::channel(1);
|
| 542 |
+
|
| 543 |
+
transport_event_tx
|
| 544 |
+
.send(TransportEvent::IncomingMessage {
|
| 545 |
+
connection_id,
|
| 546 |
+
message: JSONRPCMessage::Notification(JSONRPCNotification {
|
| 547 |
+
method: "initialized".to_string(),
|
| 548 |
+
params: None,
|
| 549 |
+
}),
|
| 550 |
+
})
|
| 551 |
+
.await
|
| 552 |
+
.expect("transport queue should accept first message");
|
| 553 |
+
|
| 554 |
+
writer_tx
|
| 555 |
+
.send(QueuedOutgoingMessage::new(
|
| 556 |
+
OutgoingMessage::AppServerNotification(ServerNotificationEnvelope {
|
| 557 |
+
notification: ServerNotification::ConfigWarning(ConfigWarningNotification {
|
| 558 |
+
summary: "queued".to_string(),
|
| 559 |
+
details: None,
|
| 560 |
+
path: None,
|
| 561 |
+
range: None,
|
| 562 |
+
}),
|
| 563 |
+
emitted_at_ms: Some(1_234),
|
| 564 |
+
}),
|
| 565 |
+
))
|
| 566 |
+
.await
|
| 567 |
+
.expect("writer queue should accept first message");
|
| 568 |
+
|
| 569 |
+
let request = JSONRPCMessage::Request(JSONRPCRequest {
|
| 570 |
+
id: RequestId::Integer(7),
|
| 571 |
+
method: "config/read".to_string(),
|
| 572 |
+
params: Some(json!({ "includeLayers": false })),
|
| 573 |
+
trace: None,
|
| 574 |
+
});
|
| 575 |
+
|
| 576 |
+
let enqueue_result = timeout(
|
| 577 |
+
Duration::from_millis(100),
|
| 578 |
+
enqueue_incoming_message(&transport_event_tx, &writer_tx, connection_id, request),
|
| 579 |
+
)
|
| 580 |
+
.await
|
| 581 |
+
.expect("enqueue should not block while writer queue is full");
|
| 582 |
+
assert!(enqueue_result);
|
| 583 |
+
|
| 584 |
+
let queued_outgoing = writer_rx
|
| 585 |
+
.recv()
|
| 586 |
+
.await
|
| 587 |
+
.expect("writer queue should still contain original message");
|
| 588 |
+
let queued_json =
|
| 589 |
+
serde_json::to_value(queued_outgoing.message).expect("serialize queued message");
|
| 590 |
+
assert_eq!(
|
| 591 |
+
queued_json,
|
| 592 |
+
json!({
|
| 593 |
+
"method": "configWarning",
|
| 594 |
+
"params": {
|
| 595 |
+
"summary": "queued",
|
| 596 |
+
"details": null,
|
| 597 |
+
},
|
| 598 |
+
"emittedAtMs": 1_234,
|
| 599 |
+
})
|
| 600 |
+
);
|
| 601 |
+
}
|
| 602 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/auth.rs
ADDED
|
@@ -0,0 +1,282 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Credentials and recovery bound to one remote-control login lifetime.
|
| 2 |
+
//! A request can refresh credentials, but cannot adopt a replacement authentication owner.
|
| 3 |
+
|
| 4 |
+
use axum::http::HeaderMap;
|
| 5 |
+
use axum::http::HeaderValue;
|
| 6 |
+
use codex_api::SharedAuthProvider;
|
| 7 |
+
use codex_login::AuthManager;
|
| 8 |
+
use codex_login::UnauthorizedRecovery;
|
| 9 |
+
use std::io;
|
| 10 |
+
use std::io::ErrorKind;
|
| 11 |
+
use std::sync::Arc;
|
| 12 |
+
use tokio::sync::watch;
|
| 13 |
+
use tracing::info;
|
| 14 |
+
use tracing::warn;
|
| 15 |
+
|
| 16 |
+
#[derive(Clone)]
|
| 17 |
+
pub(super) struct RemoteControlAuth {
|
| 18 |
+
manager: Arc<AuthManager>,
|
| 19 |
+
pub(super) owner: crate::ConnectionAuth,
|
| 20 |
+
}
|
| 21 |
+
|
| 22 |
+
pub(super) struct RemoteControlRecovery {
|
| 23 |
+
auth: RemoteControlAuth,
|
| 24 |
+
recovery: UnauthorizedRecovery,
|
| 25 |
+
}
|
| 26 |
+
|
| 27 |
+
impl RemoteControlAuth {
|
| 28 |
+
pub(super) fn capture(manager: Arc<AuthManager>) -> (Self, bool) {
|
| 29 |
+
loop {
|
| 30 |
+
let owner = crate::ConnectionAuth::capture(&manager);
|
| 31 |
+
let authenticated = manager
|
| 32 |
+
.auth_cached()
|
| 33 |
+
.is_some_and(|auth| auth.uses_codex_backend() && auth.get_account_id().is_some());
|
| 34 |
+
if owner.is_current() {
|
| 35 |
+
return (Self { manager, owner }, authenticated);
|
| 36 |
+
}
|
| 37 |
+
}
|
| 38 |
+
}
|
| 39 |
+
|
| 40 |
+
pub(super) fn ensure_current(&self) -> io::Result<()> {
|
| 41 |
+
self.owner.ensure_current()
|
| 42 |
+
}
|
| 43 |
+
|
| 44 |
+
pub(super) fn unauthorized_recovery(&self) -> RemoteControlRecovery {
|
| 45 |
+
RemoteControlRecovery {
|
| 46 |
+
auth: self.clone(),
|
| 47 |
+
recovery: self.manager.unauthorized_recovery(),
|
| 48 |
+
}
|
| 49 |
+
}
|
| 50 |
+
|
| 51 |
+
pub(super) fn auth_change_receiver(&self) -> watch::Receiver<u64> {
|
| 52 |
+
self.manager.auth_change_receiver()
|
| 53 |
+
}
|
| 54 |
+
}
|
| 55 |
+
|
| 56 |
+
pub(super) const REMOTE_CONTROL_ACCOUNT_ID_HEADER: &str = "chatgpt-account-id";
|
| 57 |
+
|
| 58 |
+
pub(super) struct RemoteControlConnectionAuth {
|
| 59 |
+
pub(super) auth_provider: SharedAuthProvider,
|
| 60 |
+
pub(super) account_id: String,
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
impl RemoteControlConnectionAuth {
|
| 64 |
+
pub(super) fn request_headers(&self) -> io::Result<HeaderMap> {
|
| 65 |
+
let mut headers = HeaderMap::new();
|
| 66 |
+
self.auth_provider.add_auth_headers(&mut headers);
|
| 67 |
+
headers.insert(
|
| 68 |
+
REMOTE_CONTROL_ACCOUNT_ID_HEADER,
|
| 69 |
+
HeaderValue::from_str(&self.account_id).map_err(|err| {
|
| 70 |
+
io::Error::new(
|
| 71 |
+
ErrorKind::InvalidInput,
|
| 72 |
+
format!("invalid remote control account id header: {err}"),
|
| 73 |
+
)
|
| 74 |
+
})?,
|
| 75 |
+
);
|
| 76 |
+
Ok(headers)
|
| 77 |
+
}
|
| 78 |
+
}
|
| 79 |
+
|
| 80 |
+
pub(super) async fn load_remote_control_auth(
|
| 81 |
+
auth: &RemoteControlAuth,
|
| 82 |
+
) -> io::Result<RemoteControlConnectionAuth> {
|
| 83 |
+
auth.ensure_current()?;
|
| 84 |
+
let credentials = load_auth_manager(&auth.manager).await?;
|
| 85 |
+
auth.ensure_current()?;
|
| 86 |
+
Ok(credentials)
|
| 87 |
+
}
|
| 88 |
+
|
| 89 |
+
async fn load_auth_manager(
|
| 90 |
+
auth_manager: &Arc<AuthManager>,
|
| 91 |
+
) -> io::Result<RemoteControlConnectionAuth> {
|
| 92 |
+
let mut reloaded = false;
|
| 93 |
+
let auth = loop {
|
| 94 |
+
let Some(auth) = auth_manager.auth().await else {
|
| 95 |
+
if reloaded {
|
| 96 |
+
return Err(io::Error::new(
|
| 97 |
+
ErrorKind::PermissionDenied,
|
| 98 |
+
"remote control requires ChatGPT authentication",
|
| 99 |
+
));
|
| 100 |
+
}
|
| 101 |
+
auth_manager.reload().await;
|
| 102 |
+
reloaded = true;
|
| 103 |
+
continue;
|
| 104 |
+
};
|
| 105 |
+
if !auth.uses_codex_backend() {
|
| 106 |
+
break auth;
|
| 107 |
+
}
|
| 108 |
+
if auth.get_account_id().is_none() && !reloaded {
|
| 109 |
+
auth_manager.reload().await;
|
| 110 |
+
reloaded = true;
|
| 111 |
+
continue;
|
| 112 |
+
}
|
| 113 |
+
break auth;
|
| 114 |
+
};
|
| 115 |
+
|
| 116 |
+
if !auth.uses_codex_backend() {
|
| 117 |
+
return Err(io::Error::new(
|
| 118 |
+
ErrorKind::PermissionDenied,
|
| 119 |
+
"remote control requires ChatGPT authentication; API key auth is not supported",
|
| 120 |
+
));
|
| 121 |
+
}
|
| 122 |
+
|
| 123 |
+
Ok(RemoteControlConnectionAuth {
|
| 124 |
+
auth_provider: codex_model_provider::auth_provider_from_auth(&auth),
|
| 125 |
+
account_id: auth.get_account_id().ok_or_else(|| {
|
| 126 |
+
io::Error::new(
|
| 127 |
+
ErrorKind::WouldBlock,
|
| 128 |
+
"remote control enrollment is waiting for a ChatGPT account id",
|
| 129 |
+
)
|
| 130 |
+
})?,
|
| 131 |
+
})
|
| 132 |
+
}
|
| 133 |
+
|
| 134 |
+
pub(super) async fn recover_remote_control_auth(
|
| 135 |
+
recovery: &mut RemoteControlRecovery,
|
| 136 |
+
auth_change_rx: &mut watch::Receiver<u64>,
|
| 137 |
+
) -> bool {
|
| 138 |
+
if recovery.auth.ensure_current().is_err() {
|
| 139 |
+
return false;
|
| 140 |
+
}
|
| 141 |
+
let auth_recovery = &mut recovery.recovery;
|
| 142 |
+
if !auth_recovery.has_next() {
|
| 143 |
+
return false;
|
| 144 |
+
}
|
| 145 |
+
|
| 146 |
+
let mode = auth_recovery.mode_name();
|
| 147 |
+
let step = auth_recovery.step_name();
|
| 148 |
+
let auth_change_revision_before_recovery = *auth_change_rx.borrow();
|
| 149 |
+
match auth_recovery.next().await {
|
| 150 |
+
Ok(step_result) => {
|
| 151 |
+
if recovery.auth.ensure_current().is_err() {
|
| 152 |
+
return false;
|
| 153 |
+
}
|
| 154 |
+
if step_result.auth_state_changed() == Some(true) {
|
| 155 |
+
mark_recovery_auth_change_seen(
|
| 156 |
+
auth_change_rx,
|
| 157 |
+
auth_change_revision_before_recovery,
|
| 158 |
+
);
|
| 159 |
+
}
|
| 160 |
+
info!(
|
| 161 |
+
"remote control auth recovery succeeded: mode={mode}, step={step}, auth_state_changed={:?}",
|
| 162 |
+
step_result.auth_state_changed()
|
| 163 |
+
);
|
| 164 |
+
true
|
| 165 |
+
}
|
| 166 |
+
Err(err) => {
|
| 167 |
+
warn!("remote control auth recovery failed: mode={mode}, step={step}: {err}");
|
| 168 |
+
false
|
| 169 |
+
}
|
| 170 |
+
}
|
| 171 |
+
}
|
| 172 |
+
|
| 173 |
+
pub(super) fn mark_recovery_auth_change_seen(
|
| 174 |
+
auth_change_rx: &mut watch::Receiver<u64>,
|
| 175 |
+
auth_change_revision_before_recovery: u64,
|
| 176 |
+
) {
|
| 177 |
+
let auth_change_revision_after_recovery = *auth_change_rx.borrow();
|
| 178 |
+
if auth_change_revision_after_recovery == auth_change_revision_before_recovery.wrapping_add(1) {
|
| 179 |
+
// Recovery updated the same watch that wakes the outer reconnect
|
| 180 |
+
// loop. Mark only that single revision seen; if more revisions
|
| 181 |
+
// arrived while recovery was in flight, leave them pending so the
|
| 182 |
+
// reconnect loop still reacts to the later external auth change.
|
| 183 |
+
auth_change_rx.borrow_and_update();
|
| 184 |
+
}
|
| 185 |
+
}
|
| 186 |
+
|
| 187 |
+
#[cfg(test)]
|
| 188 |
+
mod tests {
|
| 189 |
+
use super::*;
|
| 190 |
+
use codex_api::AuthProvider;
|
| 191 |
+
use pretty_assertions::assert_eq;
|
| 192 |
+
|
| 193 |
+
#[derive(Debug)]
|
| 194 |
+
struct TestAuthProvider {
|
| 195 |
+
account_ids: Vec<&'static str>,
|
| 196 |
+
}
|
| 197 |
+
|
| 198 |
+
impl AuthProvider for TestAuthProvider {
|
| 199 |
+
fn add_auth_headers(&self, headers: &mut HeaderMap) {
|
| 200 |
+
headers.insert(
|
| 201 |
+
axum::http::header::AUTHORIZATION,
|
| 202 |
+
HeaderValue::from_static("Bearer test-token"),
|
| 203 |
+
);
|
| 204 |
+
headers.insert("x-openai-fedramp", HeaderValue::from_static("true"));
|
| 205 |
+
for account_id in &self.account_ids {
|
| 206 |
+
headers.append("ChatGPT-Account-ID", HeaderValue::from_static(account_id));
|
| 207 |
+
}
|
| 208 |
+
}
|
| 209 |
+
}
|
| 210 |
+
|
| 211 |
+
fn remote_control_auth(
|
| 212 |
+
account_id: &str,
|
| 213 |
+
provider_account_ids: Vec<&'static str>,
|
| 214 |
+
) -> RemoteControlConnectionAuth {
|
| 215 |
+
RemoteControlConnectionAuth {
|
| 216 |
+
auth_provider: Arc::new(TestAuthProvider {
|
| 217 |
+
account_ids: provider_account_ids,
|
| 218 |
+
}),
|
| 219 |
+
account_id: account_id.to_string(),
|
| 220 |
+
}
|
| 221 |
+
}
|
| 222 |
+
|
| 223 |
+
#[test]
|
| 224 |
+
fn request_headers_adds_account_header_when_provider_omits_it() {
|
| 225 |
+
let headers = remote_control_auth("selected-account", Vec::new())
|
| 226 |
+
.request_headers()
|
| 227 |
+
.expect("request headers should build");
|
| 228 |
+
|
| 229 |
+
assert_eq!(
|
| 230 |
+
headers
|
| 231 |
+
.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER)
|
| 232 |
+
.iter()
|
| 233 |
+
.map(|value| value.to_str().expect("account header should be text"))
|
| 234 |
+
.collect::<Vec<_>>(),
|
| 235 |
+
vec!["selected-account"]
|
| 236 |
+
);
|
| 237 |
+
}
|
| 238 |
+
|
| 239 |
+
#[test]
|
| 240 |
+
fn request_headers_replaces_provider_accounts_and_preserves_other_headers() {
|
| 241 |
+
let headers = remote_control_auth(
|
| 242 |
+
"selected-account",
|
| 243 |
+
vec!["provider-account-a", "provider-account-b"],
|
| 244 |
+
)
|
| 245 |
+
.request_headers()
|
| 246 |
+
.expect("request headers should build");
|
| 247 |
+
|
| 248 |
+
assert_eq!(
|
| 249 |
+
headers
|
| 250 |
+
.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER)
|
| 251 |
+
.iter()
|
| 252 |
+
.map(|value| value.to_str().expect("account header should be text"))
|
| 253 |
+
.collect::<Vec<_>>(),
|
| 254 |
+
vec!["selected-account"]
|
| 255 |
+
);
|
| 256 |
+
assert_eq!(
|
| 257 |
+
headers
|
| 258 |
+
.get(axum::http::header::AUTHORIZATION)
|
| 259 |
+
.and_then(|value| value.to_str().ok()),
|
| 260 |
+
Some("Bearer test-token")
|
| 261 |
+
);
|
| 262 |
+
assert_eq!(
|
| 263 |
+
headers
|
| 264 |
+
.get("x-openai-fedramp")
|
| 265 |
+
.and_then(|value| value.to_str().ok()),
|
| 266 |
+
Some("true")
|
| 267 |
+
);
|
| 268 |
+
}
|
| 269 |
+
|
| 270 |
+
#[test]
|
| 271 |
+
fn request_headers_rejects_invalid_account_header_value() {
|
| 272 |
+
let err = remote_control_auth("invalid\naccount", Vec::new())
|
| 273 |
+
.request_headers()
|
| 274 |
+
.expect_err("invalid account header should fail");
|
| 275 |
+
|
| 276 |
+
assert_eq!(err.kind(), ErrorKind::InvalidInput);
|
| 277 |
+
assert!(
|
| 278 |
+
err.to_string()
|
| 279 |
+
.starts_with("invalid remote control account id header:")
|
| 280 |
+
);
|
| 281 |
+
}
|
| 282 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/client_tracker.rs
ADDED
|
@@ -0,0 +1,944 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::CHANNEL_CAPACITY;
|
| 2 |
+
use super::TransportEvent;
|
| 3 |
+
use super::next_connection_id;
|
| 4 |
+
use super::protocol::ClientEnvelope;
|
| 5 |
+
pub use super::protocol::ClientEvent;
|
| 6 |
+
pub use super::protocol::ClientId;
|
| 7 |
+
use super::protocol::PongStatus;
|
| 8 |
+
use super::protocol::ServerEvent;
|
| 9 |
+
use super::protocol::StreamId;
|
| 10 |
+
use crate::outgoing_message::ConnectionId;
|
| 11 |
+
use crate::outgoing_message::QueuedOutgoingMessage;
|
| 12 |
+
use crate::transport::ConnectionOrigin;
|
| 13 |
+
use crate::transport::remote_control::QueuedServerEnvelope;
|
| 14 |
+
use codex_app_server_protocol::JSONRPCMessage;
|
| 15 |
+
use std::collections::HashMap;
|
| 16 |
+
use tokio::sync::mpsc;
|
| 17 |
+
use tokio::sync::watch;
|
| 18 |
+
use tokio::task::JoinHandle;
|
| 19 |
+
use tokio::task::JoinSet;
|
| 20 |
+
use tokio::time::Duration;
|
| 21 |
+
use tokio::time::Instant;
|
| 22 |
+
use tokio::time::timeout;
|
| 23 |
+
use tokio_util::sync::CancellationToken;
|
| 24 |
+
use tracing::info;
|
| 25 |
+
use tracing::warn;
|
| 26 |
+
|
| 27 |
+
const REMOTE_CONTROL_CLIENT_IDLE_TIMEOUT: Duration = Duration::from_secs(10 * 60);
|
| 28 |
+
pub(crate) const REMOTE_CONTROL_IDLE_SWEEP_INTERVAL: Duration = Duration::from_secs(30);
|
| 29 |
+
#[cfg(not(test))]
|
| 30 |
+
const REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT: Duration = Duration::from_secs(5);
|
| 31 |
+
#[cfg(test)]
|
| 32 |
+
const REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT: Duration = Duration::from_millis(10);
|
| 33 |
+
|
| 34 |
+
#[derive(Debug)]
|
| 35 |
+
pub(crate) struct Stopped;
|
| 36 |
+
|
| 37 |
+
struct ClientState {
|
| 38 |
+
connection_id: ConnectionId,
|
| 39 |
+
disconnect_token: CancellationToken,
|
| 40 |
+
last_activity_at: Instant,
|
| 41 |
+
last_inbound_seq_id: Option<u64>,
|
| 42 |
+
status_tx: watch::Sender<PongStatus>,
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
pub(crate) struct ClientTracker {
|
| 46 |
+
pub(super) auth: Option<crate::ConnectionAuth>,
|
| 47 |
+
clients: HashMap<(ClientId, StreamId), ClientState>,
|
| 48 |
+
legacy_stream_ids: HashMap<ClientId, StreamId>,
|
| 49 |
+
join_set: JoinSet<(ClientId, StreamId)>,
|
| 50 |
+
server_event_tx: mpsc::Sender<QueuedServerEnvelope>,
|
| 51 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 52 |
+
shutdown_token: CancellationToken,
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
impl ClientTracker {
|
| 56 |
+
pub(crate) fn new(
|
| 57 |
+
server_event_tx: mpsc::Sender<QueuedServerEnvelope>,
|
| 58 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 59 |
+
shutdown_token: &CancellationToken,
|
| 60 |
+
) -> Self {
|
| 61 |
+
Self {
|
| 62 |
+
auth: None,
|
| 63 |
+
clients: HashMap::new(),
|
| 64 |
+
legacy_stream_ids: HashMap::new(),
|
| 65 |
+
join_set: JoinSet::new(),
|
| 66 |
+
server_event_tx,
|
| 67 |
+
transport_event_tx,
|
| 68 |
+
shutdown_token: shutdown_token.child_token(),
|
| 69 |
+
}
|
| 70 |
+
}
|
| 71 |
+
|
| 72 |
+
pub(crate) async fn bookkeep_join_set(&mut self) -> Option<(ClientId, StreamId)> {
|
| 73 |
+
while let Some(join_result) = self.join_set.join_next().await {
|
| 74 |
+
let Ok(client_key) = join_result else {
|
| 75 |
+
continue;
|
| 76 |
+
};
|
| 77 |
+
return Some(client_key);
|
| 78 |
+
}
|
| 79 |
+
futures::future::pending().await
|
| 80 |
+
}
|
| 81 |
+
|
| 82 |
+
pub(crate) async fn shutdown(&mut self) {
|
| 83 |
+
self.shutdown_token.cancel();
|
| 84 |
+
|
| 85 |
+
while let Some(client_key) = self.clients.keys().next().cloned() {
|
| 86 |
+
let _ = self.close_client(&client_key).await;
|
| 87 |
+
}
|
| 88 |
+
|
| 89 |
+
self.drain_join_set().await;
|
| 90 |
+
}
|
| 91 |
+
|
| 92 |
+
async fn drain_join_set(&mut self) {
|
| 93 |
+
while self.join_set.join_next().await.is_some() {}
|
| 94 |
+
}
|
| 95 |
+
|
| 96 |
+
pub(crate) async fn handle_message(
|
| 97 |
+
&mut self,
|
| 98 |
+
client_envelope: ClientEnvelope,
|
| 99 |
+
) -> Result<(), Stopped> {
|
| 100 |
+
let ClientEnvelope {
|
| 101 |
+
client_id,
|
| 102 |
+
event,
|
| 103 |
+
stream_id,
|
| 104 |
+
seq_id,
|
| 105 |
+
cursor: _,
|
| 106 |
+
} = client_envelope;
|
| 107 |
+
let is_legacy_stream_id = stream_id.is_none();
|
| 108 |
+
let is_initialize = matches!(&event, ClientEvent::ClientMessage { message } if remote_control_message_starts_connection(message));
|
| 109 |
+
let stream_id = match stream_id {
|
| 110 |
+
Some(stream_id) => stream_id,
|
| 111 |
+
None if is_initialize => {
|
| 112 |
+
// TODO(ruslan): delete this fallback once all clients are updated to send stream_id.
|
| 113 |
+
self.legacy_stream_ids
|
| 114 |
+
.remove(&client_id)
|
| 115 |
+
.unwrap_or_else(StreamId::new_random)
|
| 116 |
+
}
|
| 117 |
+
None => self
|
| 118 |
+
.legacy_stream_ids
|
| 119 |
+
.get(&client_id)
|
| 120 |
+
.cloned()
|
| 121 |
+
.unwrap_or_else(|| {
|
| 122 |
+
if matches!(&event, ClientEvent::Ping) {
|
| 123 |
+
StreamId::new_random()
|
| 124 |
+
} else {
|
| 125 |
+
StreamId(String::new())
|
| 126 |
+
}
|
| 127 |
+
}),
|
| 128 |
+
};
|
| 129 |
+
if stream_id.0.is_empty() {
|
| 130 |
+
return Ok(());
|
| 131 |
+
}
|
| 132 |
+
let client_key = (client_id.clone(), stream_id.clone());
|
| 133 |
+
match event {
|
| 134 |
+
ClientEvent::ClientMessage { message } => {
|
| 135 |
+
if let Some(seq_id) = seq_id
|
| 136 |
+
&& let Some(client) = self.clients.get(&client_key)
|
| 137 |
+
&& client
|
| 138 |
+
.last_inbound_seq_id
|
| 139 |
+
.is_some_and(|last_seq_id| last_seq_id >= seq_id)
|
| 140 |
+
&& !is_initialize
|
| 141 |
+
{
|
| 142 |
+
return Ok(());
|
| 143 |
+
}
|
| 144 |
+
|
| 145 |
+
if is_initialize && self.clients.contains_key(&client_key) {
|
| 146 |
+
self.close_client(&client_key).await?;
|
| 147 |
+
}
|
| 148 |
+
|
| 149 |
+
if let Some(connection_id) = self.clients.get_mut(&client_key).map(|client| {
|
| 150 |
+
client.last_activity_at = Instant::now();
|
| 151 |
+
client.connection_id
|
| 152 |
+
}) {
|
| 153 |
+
self.send_transport_event(TransportEvent::IncomingMessage {
|
| 154 |
+
connection_id,
|
| 155 |
+
message,
|
| 156 |
+
})
|
| 157 |
+
.await?;
|
| 158 |
+
self.record_inbound_message_delivery(&client_key, seq_id);
|
| 159 |
+
return Ok(());
|
| 160 |
+
}
|
| 161 |
+
|
| 162 |
+
if !is_initialize {
|
| 163 |
+
return Ok(());
|
| 164 |
+
}
|
| 165 |
+
|
| 166 |
+
let connection_id = next_connection_id();
|
| 167 |
+
let (writer_tx, writer_rx) =
|
| 168 |
+
mpsc::channel::<QueuedOutgoingMessage>(CHANNEL_CAPACITY);
|
| 169 |
+
let disconnect_token = self.shutdown_token.child_token();
|
| 170 |
+
self.send_transport_event(TransportEvent::ConnectionOpened {
|
| 171 |
+
connection_id,
|
| 172 |
+
origin: ConnectionOrigin::RemoteControl,
|
| 173 |
+
auth: self.auth.clone(),
|
| 174 |
+
writer: writer_tx,
|
| 175 |
+
disconnect_sender: Some(disconnect_token.clone()),
|
| 176 |
+
})
|
| 177 |
+
.await?;
|
| 178 |
+
|
| 179 |
+
let (status_tx, status_rx) = watch::channel(PongStatus::Active);
|
| 180 |
+
self.join_set.spawn(Self::run_client_outbound(
|
| 181 |
+
client_id.clone(),
|
| 182 |
+
stream_id.clone(),
|
| 183 |
+
self.server_event_tx.clone(),
|
| 184 |
+
writer_rx,
|
| 185 |
+
status_rx,
|
| 186 |
+
disconnect_token.clone(),
|
| 187 |
+
));
|
| 188 |
+
self.clients.insert(
|
| 189 |
+
client_key.clone(),
|
| 190 |
+
ClientState {
|
| 191 |
+
connection_id,
|
| 192 |
+
disconnect_token,
|
| 193 |
+
last_activity_at: Instant::now(),
|
| 194 |
+
last_inbound_seq_id: None,
|
| 195 |
+
status_tx,
|
| 196 |
+
},
|
| 197 |
+
);
|
| 198 |
+
if is_legacy_stream_id {
|
| 199 |
+
self.legacy_stream_ids.insert(client_id.clone(), stream_id);
|
| 200 |
+
}
|
| 201 |
+
if let Err(err) = self
|
| 202 |
+
.send_transport_event(TransportEvent::IncomingMessage {
|
| 203 |
+
connection_id,
|
| 204 |
+
message,
|
| 205 |
+
})
|
| 206 |
+
.await
|
| 207 |
+
{
|
| 208 |
+
if let Some(client) = self.remove_client(&client_key) {
|
| 209 |
+
client.disconnect_token.cancel();
|
| 210 |
+
// The initialize send already timed out on this queue; preserve close
|
| 211 |
+
// delivery without blocking reconnect on the same backpressure.
|
| 212 |
+
drop(self.spawn_connection_closed(client.connection_id));
|
| 213 |
+
}
|
| 214 |
+
return Err(err);
|
| 215 |
+
}
|
| 216 |
+
if !is_legacy_stream_id {
|
| 217 |
+
self.record_inbound_message_delivery(&client_key, seq_id);
|
| 218 |
+
}
|
| 219 |
+
Ok(())
|
| 220 |
+
}
|
| 221 |
+
ClientEvent::ClientMessageChunk { .. } | ClientEvent::Ack { .. } => Ok(()),
|
| 222 |
+
ClientEvent::Ping => {
|
| 223 |
+
if let Some(client) = self.clients.get_mut(&client_key) {
|
| 224 |
+
client.last_activity_at = Instant::now();
|
| 225 |
+
let _ = client.status_tx.send(PongStatus::Active);
|
| 226 |
+
return Ok(());
|
| 227 |
+
}
|
| 228 |
+
|
| 229 |
+
let server_event_tx = self.server_event_tx.clone();
|
| 230 |
+
tokio::spawn(async move {
|
| 231 |
+
let server_envelope = QueuedServerEnvelope {
|
| 232 |
+
event: ServerEvent::Pong {
|
| 233 |
+
status: PongStatus::Unknown,
|
| 234 |
+
},
|
| 235 |
+
client_id,
|
| 236 |
+
stream_id,
|
| 237 |
+
write_complete_tx: None,
|
| 238 |
+
};
|
| 239 |
+
let _ = server_event_tx.send(server_envelope).await;
|
| 240 |
+
});
|
| 241 |
+
Ok(())
|
| 242 |
+
}
|
| 243 |
+
ClientEvent::ClientClosed => self.close_client(&client_key).await,
|
| 244 |
+
}
|
| 245 |
+
}
|
| 246 |
+
|
| 247 |
+
async fn run_client_outbound(
|
| 248 |
+
client_id: ClientId,
|
| 249 |
+
stream_id: StreamId,
|
| 250 |
+
server_event_tx: mpsc::Sender<QueuedServerEnvelope>,
|
| 251 |
+
mut writer_rx: mpsc::Receiver<QueuedOutgoingMessage>,
|
| 252 |
+
mut status_rx: watch::Receiver<PongStatus>,
|
| 253 |
+
disconnect_token: CancellationToken,
|
| 254 |
+
) -> (ClientId, StreamId) {
|
| 255 |
+
loop {
|
| 256 |
+
let (event, write_complete_tx) = tokio::select! {
|
| 257 |
+
_ = disconnect_token.cancelled() => {
|
| 258 |
+
break;
|
| 259 |
+
}
|
| 260 |
+
queued_message = writer_rx.recv() => {
|
| 261 |
+
let Some(queued_message) = queued_message else {
|
| 262 |
+
break;
|
| 263 |
+
};
|
| 264 |
+
let event = ServerEvent::ServerMessage {
|
| 265 |
+
message: Box::new(queued_message.message),
|
| 266 |
+
};
|
| 267 |
+
(event, queued_message.write_complete_tx)
|
| 268 |
+
}
|
| 269 |
+
changed = status_rx.changed() => {
|
| 270 |
+
if changed.is_err() {
|
| 271 |
+
break;
|
| 272 |
+
}
|
| 273 |
+
let event = ServerEvent::Pong { status: status_rx.borrow().clone() };
|
| 274 |
+
(event, None)
|
| 275 |
+
}
|
| 276 |
+
};
|
| 277 |
+
let send_result = tokio::select! {
|
| 278 |
+
_ = disconnect_token.cancelled() => {
|
| 279 |
+
break;
|
| 280 |
+
}
|
| 281 |
+
send_result = server_event_tx.send(QueuedServerEnvelope {
|
| 282 |
+
event,
|
| 283 |
+
client_id: client_id.clone(),
|
| 284 |
+
stream_id: stream_id.clone(),
|
| 285 |
+
write_complete_tx,
|
| 286 |
+
}) => send_result,
|
| 287 |
+
};
|
| 288 |
+
if send_result.is_err() {
|
| 289 |
+
break;
|
| 290 |
+
}
|
| 291 |
+
}
|
| 292 |
+
(client_id, stream_id)
|
| 293 |
+
}
|
| 294 |
+
|
| 295 |
+
pub(crate) async fn close_expired_clients(
|
| 296 |
+
&mut self,
|
| 297 |
+
) -> Result<Vec<(ClientId, StreamId)>, Stopped> {
|
| 298 |
+
let now = Instant::now();
|
| 299 |
+
let expired_client_ids: Vec<(ClientId, StreamId)> = self
|
| 300 |
+
.clients
|
| 301 |
+
.iter()
|
| 302 |
+
.filter_map(|(client_key, client)| {
|
| 303 |
+
(!remote_control_client_is_alive(client, now)).then_some(client_key.clone())
|
| 304 |
+
})
|
| 305 |
+
.collect();
|
| 306 |
+
for client_key in &expired_client_ids {
|
| 307 |
+
self.close_client(client_key).await?;
|
| 308 |
+
}
|
| 309 |
+
Ok(expired_client_ids)
|
| 310 |
+
}
|
| 311 |
+
|
| 312 |
+
pub(super) async fn close_client(
|
| 313 |
+
&mut self,
|
| 314 |
+
client_key: &(ClientId, StreamId),
|
| 315 |
+
) -> Result<(), Stopped> {
|
| 316 |
+
let Some(client) = self.remove_client(client_key) else {
|
| 317 |
+
return Ok(());
|
| 318 |
+
};
|
| 319 |
+
client.disconnect_token.cancel();
|
| 320 |
+
self.send_transport_event(TransportEvent::ConnectionClosed {
|
| 321 |
+
connection_id: client.connection_id,
|
| 322 |
+
})
|
| 323 |
+
.await
|
| 324 |
+
}
|
| 325 |
+
|
| 326 |
+
fn remove_client(&mut self, client_key: &(ClientId, StreamId)) -> Option<ClientState> {
|
| 327 |
+
let client = self.clients.remove(client_key)?;
|
| 328 |
+
if self
|
| 329 |
+
.legacy_stream_ids
|
| 330 |
+
.get(&client_key.0)
|
| 331 |
+
.is_some_and(|stream_id| stream_id == &client_key.1)
|
| 332 |
+
{
|
| 333 |
+
self.legacy_stream_ids.remove(&client_key.0);
|
| 334 |
+
}
|
| 335 |
+
Some(client)
|
| 336 |
+
}
|
| 337 |
+
|
| 338 |
+
async fn send_transport_event(&self, event: TransportEvent) -> Result<(), Stopped> {
|
| 339 |
+
let event = match event {
|
| 340 |
+
TransportEvent::ConnectionClosed { connection_id } => {
|
| 341 |
+
return self.send_connection_closed(connection_id).await;
|
| 342 |
+
}
|
| 343 |
+
event => event,
|
| 344 |
+
};
|
| 345 |
+
|
| 346 |
+
let event_name = transport_event_name(&event);
|
| 347 |
+
match timeout(
|
| 348 |
+
REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT,
|
| 349 |
+
self.transport_event_tx.send(event),
|
| 350 |
+
)
|
| 351 |
+
.await
|
| 352 |
+
{
|
| 353 |
+
Ok(Ok(())) => Ok(()),
|
| 354 |
+
Ok(Err(_)) => {
|
| 355 |
+
warn!(
|
| 356 |
+
transport_event = event_name,
|
| 357 |
+
"remote control transport event receiver dropped"
|
| 358 |
+
);
|
| 359 |
+
Err(Stopped)
|
| 360 |
+
}
|
| 361 |
+
Err(_) => {
|
| 362 |
+
warn!(
|
| 363 |
+
transport_event = event_name,
|
| 364 |
+
timeout = ?REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT,
|
| 365 |
+
"timed out forwarding remote control transport event"
|
| 366 |
+
);
|
| 367 |
+
Err(Stopped)
|
| 368 |
+
}
|
| 369 |
+
}
|
| 370 |
+
}
|
| 371 |
+
|
| 372 |
+
fn record_inbound_message_delivery(
|
| 373 |
+
&mut self,
|
| 374 |
+
client_key: &(ClientId, StreamId),
|
| 375 |
+
seq_id: Option<u64>,
|
| 376 |
+
) {
|
| 377 |
+
// Timed forwarding can fail, so only dedupe retries after app-server receives it.
|
| 378 |
+
if let Some(seq_id) = seq_id
|
| 379 |
+
&& let Some(client) = self.clients.get_mut(client_key)
|
| 380 |
+
{
|
| 381 |
+
client.last_inbound_seq_id = Some(seq_id);
|
| 382 |
+
}
|
| 383 |
+
}
|
| 384 |
+
|
| 385 |
+
async fn send_connection_closed(&self, connection_id: ConnectionId) -> Result<(), Stopped> {
|
| 386 |
+
// Worker shutdown can abort the caller; detach the cleanup event before awaiting it.
|
| 387 |
+
match self.spawn_connection_closed(connection_id).await {
|
| 388 |
+
Ok(result) => result,
|
| 389 |
+
Err(err) => {
|
| 390 |
+
warn!(
|
| 391 |
+
transport_event = "connection_closed",
|
| 392 |
+
?err,
|
| 393 |
+
"remote control transport event forwarding task failed"
|
| 394 |
+
);
|
| 395 |
+
Err(Stopped)
|
| 396 |
+
}
|
| 397 |
+
}
|
| 398 |
+
}
|
| 399 |
+
|
| 400 |
+
fn spawn_connection_closed(
|
| 401 |
+
&self,
|
| 402 |
+
connection_id: ConnectionId,
|
| 403 |
+
) -> JoinHandle<Result<(), Stopped>> {
|
| 404 |
+
info!(
|
| 405 |
+
connection_id = ?connection_id,
|
| 406 |
+
"forwarding remote control connection closed transport event"
|
| 407 |
+
);
|
| 408 |
+
let transport_event_tx = self.transport_event_tx.clone();
|
| 409 |
+
tokio::spawn(async move {
|
| 410 |
+
transport_event_tx
|
| 411 |
+
.send(TransportEvent::ConnectionClosed { connection_id })
|
| 412 |
+
.await
|
| 413 |
+
.map_err(|_| {
|
| 414 |
+
warn!(
|
| 415 |
+
transport_event = "connection_closed",
|
| 416 |
+
"remote control transport event receiver dropped"
|
| 417 |
+
);
|
| 418 |
+
Stopped
|
| 419 |
+
})
|
| 420 |
+
})
|
| 421 |
+
}
|
| 422 |
+
}
|
| 423 |
+
|
| 424 |
+
fn transport_event_name(event: &TransportEvent) -> &'static str {
|
| 425 |
+
match event {
|
| 426 |
+
TransportEvent::ConnectionOpened { .. } => "connection_opened",
|
| 427 |
+
TransportEvent::ConnectionClosed { .. } => "connection_closed",
|
| 428 |
+
TransportEvent::IncomingMessage { .. } => "incoming_message",
|
| 429 |
+
TransportEvent::DaemonShutdown => "daemon_shutdown",
|
| 430 |
+
}
|
| 431 |
+
}
|
| 432 |
+
|
| 433 |
+
fn remote_control_message_starts_connection(message: &JSONRPCMessage) -> bool {
|
| 434 |
+
matches!(
|
| 435 |
+
message,
|
| 436 |
+
JSONRPCMessage::Request(codex_app_server_protocol::JSONRPCRequest { method, .. })
|
| 437 |
+
if method == "initialize"
|
| 438 |
+
)
|
| 439 |
+
}
|
| 440 |
+
|
| 441 |
+
fn remote_control_client_is_alive(client: &ClientState, now: Instant) -> bool {
|
| 442 |
+
now.duration_since(client.last_activity_at) < REMOTE_CONTROL_CLIENT_IDLE_TIMEOUT
|
| 443 |
+
}
|
| 444 |
+
|
| 445 |
+
#[cfg(test)]
|
| 446 |
+
mod tests {
|
| 447 |
+
use super::*;
|
| 448 |
+
use crate::outgoing_message::OutgoingMessage;
|
| 449 |
+
use crate::transport::remote_control::protocol::ClientEnvelope;
|
| 450 |
+
use crate::transport::remote_control::protocol::ClientEvent;
|
| 451 |
+
use codex_app_server_protocol::ConfigWarningNotification;
|
| 452 |
+
use codex_app_server_protocol::JSONRPCRequest;
|
| 453 |
+
use codex_app_server_protocol::RequestId;
|
| 454 |
+
use codex_app_server_protocol::ServerNotification;
|
| 455 |
+
use codex_app_server_protocol::ServerNotificationEnvelope;
|
| 456 |
+
use pretty_assertions::assert_eq;
|
| 457 |
+
use serde_json::json;
|
| 458 |
+
use tokio::time::timeout;
|
| 459 |
+
|
| 460 |
+
fn initialize_envelope(client_id: &str) -> ClientEnvelope {
|
| 461 |
+
initialize_envelope_with_stream_id(client_id, /*stream_id*/ None)
|
| 462 |
+
}
|
| 463 |
+
|
| 464 |
+
fn initialize_envelope_with_stream_id(
|
| 465 |
+
client_id: &str,
|
| 466 |
+
stream_id: Option<&str>,
|
| 467 |
+
) -> ClientEnvelope {
|
| 468 |
+
ClientEnvelope {
|
| 469 |
+
event: ClientEvent::ClientMessage {
|
| 470 |
+
message: JSONRPCMessage::Request(JSONRPCRequest {
|
| 471 |
+
id: RequestId::Integer(1),
|
| 472 |
+
method: "initialize".to_string(),
|
| 473 |
+
params: Some(json!({
|
| 474 |
+
"clientInfo": {
|
| 475 |
+
"name": "remote-test-client",
|
| 476 |
+
"version": "0.1.0"
|
| 477 |
+
}
|
| 478 |
+
})),
|
| 479 |
+
trace: None,
|
| 480 |
+
}),
|
| 481 |
+
},
|
| 482 |
+
client_id: ClientId(client_id.to_string()),
|
| 483 |
+
stream_id: stream_id.map(|stream_id| StreamId(stream_id.to_string())),
|
| 484 |
+
seq_id: Some(0),
|
| 485 |
+
cursor: None,
|
| 486 |
+
}
|
| 487 |
+
}
|
| 488 |
+
|
| 489 |
+
fn initialized_notification() -> JSONRPCMessage {
|
| 490 |
+
JSONRPCMessage::Notification(codex_app_server_protocol::JSONRPCNotification {
|
| 491 |
+
method: "initialized".to_string(),
|
| 492 |
+
params: None,
|
| 493 |
+
})
|
| 494 |
+
}
|
| 495 |
+
|
| 496 |
+
#[tokio::test]
|
| 497 |
+
async fn cancelled_outbound_task_emits_connection_closed() {
|
| 498 |
+
let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 499 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 500 |
+
let shutdown_token = CancellationToken::new();
|
| 501 |
+
let mut client_tracker =
|
| 502 |
+
ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token);
|
| 503 |
+
|
| 504 |
+
client_tracker
|
| 505 |
+
.handle_message(initialize_envelope("client-1"))
|
| 506 |
+
.await
|
| 507 |
+
.expect("initialize should open client");
|
| 508 |
+
|
| 509 |
+
let (connection_id, disconnect_sender) = match transport_event_rx
|
| 510 |
+
.recv()
|
| 511 |
+
.await
|
| 512 |
+
.expect("connection opened should be sent")
|
| 513 |
+
{
|
| 514 |
+
TransportEvent::ConnectionOpened {
|
| 515 |
+
connection_id,
|
| 516 |
+
disconnect_sender: Some(disconnect_sender),
|
| 517 |
+
..
|
| 518 |
+
} => (connection_id, disconnect_sender),
|
| 519 |
+
other => panic!("expected connection opened, got {other:?}"),
|
| 520 |
+
};
|
| 521 |
+
match transport_event_rx
|
| 522 |
+
.recv()
|
| 523 |
+
.await
|
| 524 |
+
.expect("initialize should be forwarded")
|
| 525 |
+
{
|
| 526 |
+
TransportEvent::IncomingMessage {
|
| 527 |
+
connection_id: incoming_connection_id,
|
| 528 |
+
..
|
| 529 |
+
} => assert_eq!(incoming_connection_id, connection_id),
|
| 530 |
+
other => panic!("expected incoming initialize, got {other:?}"),
|
| 531 |
+
}
|
| 532 |
+
|
| 533 |
+
disconnect_sender.cancel();
|
| 534 |
+
let closed_client_id = timeout(Duration::from_secs(1), client_tracker.bookkeep_join_set())
|
| 535 |
+
.await
|
| 536 |
+
.expect("bookkeeping should process the closed task")
|
| 537 |
+
.expect("closed task should return client id");
|
| 538 |
+
assert_eq!(closed_client_id.0, ClientId("client-1".to_string()));
|
| 539 |
+
client_tracker
|
| 540 |
+
.close_client(&closed_client_id)
|
| 541 |
+
.await
|
| 542 |
+
.expect("closed client should emit connection closed");
|
| 543 |
+
|
| 544 |
+
match transport_event_rx
|
| 545 |
+
.recv()
|
| 546 |
+
.await
|
| 547 |
+
.expect("connection closed should be sent")
|
| 548 |
+
{
|
| 549 |
+
TransportEvent::ConnectionClosed {
|
| 550 |
+
connection_id: closed_connection_id,
|
| 551 |
+
} => assert_eq!(closed_connection_id, connection_id),
|
| 552 |
+
other => panic!("expected connection closed, got {other:?}"),
|
| 553 |
+
}
|
| 554 |
+
}
|
| 555 |
+
|
| 556 |
+
#[tokio::test]
|
| 557 |
+
async fn shutdown_cancels_blocked_outbound_forwarding() {
|
| 558 |
+
let (server_event_tx, _server_event_rx) = mpsc::channel(1);
|
| 559 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 560 |
+
let shutdown_token = CancellationToken::new();
|
| 561 |
+
let mut client_tracker =
|
| 562 |
+
ClientTracker::new(server_event_tx.clone(), transport_event_tx, &shutdown_token);
|
| 563 |
+
|
| 564 |
+
server_event_tx
|
| 565 |
+
.send(QueuedServerEnvelope {
|
| 566 |
+
event: ServerEvent::Pong {
|
| 567 |
+
status: PongStatus::Unknown,
|
| 568 |
+
},
|
| 569 |
+
client_id: ClientId("queued-client".to_string()),
|
| 570 |
+
stream_id: StreamId("queued-stream".to_string()),
|
| 571 |
+
write_complete_tx: None,
|
| 572 |
+
})
|
| 573 |
+
.await
|
| 574 |
+
.expect("server event queue should accept prefill");
|
| 575 |
+
|
| 576 |
+
client_tracker
|
| 577 |
+
.handle_message(initialize_envelope("client-1"))
|
| 578 |
+
.await
|
| 579 |
+
.expect("initialize should open client");
|
| 580 |
+
|
| 581 |
+
let writer = match transport_event_rx
|
| 582 |
+
.recv()
|
| 583 |
+
.await
|
| 584 |
+
.expect("connection opened should be sent")
|
| 585 |
+
{
|
| 586 |
+
TransportEvent::ConnectionOpened { writer, .. } => writer,
|
| 587 |
+
other => panic!("expected connection opened, got {other:?}"),
|
| 588 |
+
};
|
| 589 |
+
let _ = transport_event_rx
|
| 590 |
+
.recv()
|
| 591 |
+
.await
|
| 592 |
+
.expect("initialize should be forwarded");
|
| 593 |
+
|
| 594 |
+
writer
|
| 595 |
+
.send(QueuedOutgoingMessage::new(
|
| 596 |
+
OutgoingMessage::AppServerNotification(ServerNotificationEnvelope {
|
| 597 |
+
notification: ServerNotification::ConfigWarning(ConfigWarningNotification {
|
| 598 |
+
summary: "test".to_string(),
|
| 599 |
+
details: None,
|
| 600 |
+
path: None,
|
| 601 |
+
range: None,
|
| 602 |
+
}),
|
| 603 |
+
emitted_at_ms: Some(1_234),
|
| 604 |
+
}),
|
| 605 |
+
))
|
| 606 |
+
.await
|
| 607 |
+
.expect("writer should accept queued message");
|
| 608 |
+
|
| 609 |
+
timeout(Duration::from_secs(1), client_tracker.shutdown())
|
| 610 |
+
.await
|
| 611 |
+
.expect("shutdown should not hang on blocked server forwarding");
|
| 612 |
+
}
|
| 613 |
+
|
| 614 |
+
#[tokio::test]
|
| 615 |
+
async fn non_close_transport_event_send_times_out_when_queue_stays_full() {
|
| 616 |
+
let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 617 |
+
let (transport_event_tx, _transport_event_rx) = mpsc::channel(1);
|
| 618 |
+
let shutdown_token = CancellationToken::new();
|
| 619 |
+
let client_tracker =
|
| 620 |
+
ClientTracker::new(server_event_tx, transport_event_tx.clone(), &shutdown_token);
|
| 621 |
+
|
| 622 |
+
transport_event_tx
|
| 623 |
+
.send(TransportEvent::ConnectionClosed {
|
| 624 |
+
connection_id: next_connection_id(),
|
| 625 |
+
})
|
| 626 |
+
.await
|
| 627 |
+
.expect("transport event queue should accept prefill");
|
| 628 |
+
|
| 629 |
+
let send_result = client_tracker
|
| 630 |
+
.send_transport_event(TransportEvent::IncomingMessage {
|
| 631 |
+
connection_id: next_connection_id(),
|
| 632 |
+
message: initialized_notification(),
|
| 633 |
+
})
|
| 634 |
+
.await;
|
| 635 |
+
|
| 636 |
+
assert!(send_result.is_err());
|
| 637 |
+
}
|
| 638 |
+
|
| 639 |
+
#[tokio::test]
|
| 640 |
+
async fn incoming_message_timeout_does_not_advance_seq_id() {
|
| 641 |
+
let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 642 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(2);
|
| 643 |
+
let shutdown_token = CancellationToken::new();
|
| 644 |
+
let mut client_tracker =
|
| 645 |
+
ClientTracker::new(server_event_tx, transport_event_tx.clone(), &shutdown_token);
|
| 646 |
+
|
| 647 |
+
client_tracker
|
| 648 |
+
.handle_message(initialize_envelope_with_stream_id(
|
| 649 |
+
"client-1",
|
| 650 |
+
Some("stream-1"),
|
| 651 |
+
))
|
| 652 |
+
.await
|
| 653 |
+
.expect("initialize should open client");
|
| 654 |
+
let connection_id = match transport_event_rx.recv().await.expect("open event") {
|
| 655 |
+
TransportEvent::ConnectionOpened { connection_id, .. } => connection_id,
|
| 656 |
+
other => panic!("expected connection opened, got {other:?}"),
|
| 657 |
+
};
|
| 658 |
+
let _ = transport_event_rx.recv().await.expect("initialize event");
|
| 659 |
+
|
| 660 |
+
for _ in 0..2 {
|
| 661 |
+
transport_event_tx
|
| 662 |
+
.send(TransportEvent::ConnectionClosed {
|
| 663 |
+
connection_id: next_connection_id(),
|
| 664 |
+
})
|
| 665 |
+
.await
|
| 666 |
+
.expect("transport event queue should accept prefill");
|
| 667 |
+
}
|
| 668 |
+
|
| 669 |
+
let retry_envelope = ClientEnvelope {
|
| 670 |
+
event: ClientEvent::ClientMessage {
|
| 671 |
+
message: initialized_notification(),
|
| 672 |
+
},
|
| 673 |
+
client_id: ClientId("client-1".to_string()),
|
| 674 |
+
stream_id: Some(StreamId("stream-1".to_string())),
|
| 675 |
+
seq_id: Some(1),
|
| 676 |
+
cursor: None,
|
| 677 |
+
};
|
| 678 |
+
assert!(
|
| 679 |
+
client_tracker
|
| 680 |
+
.handle_message(retry_envelope.clone())
|
| 681 |
+
.await
|
| 682 |
+
.is_err()
|
| 683 |
+
);
|
| 684 |
+
for _ in 0..2 {
|
| 685 |
+
let _ = transport_event_rx.recv().await.expect("prefilled event");
|
| 686 |
+
}
|
| 687 |
+
|
| 688 |
+
client_tracker
|
| 689 |
+
.handle_message(retry_envelope)
|
| 690 |
+
.await
|
| 691 |
+
.expect("retry should forward after timeout");
|
| 692 |
+
match transport_event_rx.recv().await.expect("retried event") {
|
| 693 |
+
TransportEvent::IncomingMessage {
|
| 694 |
+
connection_id: queued_connection_id,
|
| 695 |
+
..
|
| 696 |
+
} => assert_eq!(queued_connection_id, connection_id),
|
| 697 |
+
other => panic!("expected incoming message, got {other:?}"),
|
| 698 |
+
}
|
| 699 |
+
}
|
| 700 |
+
|
| 701 |
+
#[tokio::test(start_paused = true)]
|
| 702 |
+
async fn initialize_timeout_closes_open_connection() {
|
| 703 |
+
let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 704 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(1);
|
| 705 |
+
let shutdown_token = CancellationToken::new();
|
| 706 |
+
let client_tracker =
|
| 707 |
+
ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token);
|
| 708 |
+
let handle_message = tokio::spawn(async move {
|
| 709 |
+
let mut client_tracker = client_tracker;
|
| 710 |
+
client_tracker
|
| 711 |
+
.handle_message(initialize_envelope_with_stream_id(
|
| 712 |
+
"client-1",
|
| 713 |
+
Some("stream-1"),
|
| 714 |
+
))
|
| 715 |
+
.await
|
| 716 |
+
});
|
| 717 |
+
|
| 718 |
+
tokio::task::yield_now().await;
|
| 719 |
+
tokio::time::advance(
|
| 720 |
+
REMOTE_CONTROL_TRANSPORT_EVENT_SEND_TIMEOUT + Duration::from_millis(1),
|
| 721 |
+
)
|
| 722 |
+
.await;
|
| 723 |
+
|
| 724 |
+
assert!(handle_message.await.expect("handle message task").is_err());
|
| 725 |
+
let connection_id = match transport_event_rx.recv().await.expect("open event") {
|
| 726 |
+
TransportEvent::ConnectionOpened { connection_id, .. } => connection_id,
|
| 727 |
+
other => panic!("expected connection opened, got {other:?}"),
|
| 728 |
+
};
|
| 729 |
+
|
| 730 |
+
match transport_event_rx.recv().await.expect("close event") {
|
| 731 |
+
TransportEvent::ConnectionClosed {
|
| 732 |
+
connection_id: closed_connection_id,
|
| 733 |
+
} => assert_eq!(closed_connection_id, connection_id),
|
| 734 |
+
other => panic!("expected connection closed, got {other:?}"),
|
| 735 |
+
}
|
| 736 |
+
}
|
| 737 |
+
|
| 738 |
+
#[tokio::test]
|
| 739 |
+
async fn close_client_waits_for_transport_event_queue_capacity() {
|
| 740 |
+
let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 741 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(2);
|
| 742 |
+
let shutdown_token = CancellationToken::new();
|
| 743 |
+
let mut client_tracker =
|
| 744 |
+
ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token);
|
| 745 |
+
|
| 746 |
+
client_tracker
|
| 747 |
+
.handle_message(initialize_envelope_with_stream_id(
|
| 748 |
+
"client-1",
|
| 749 |
+
Some("stream-1"),
|
| 750 |
+
))
|
| 751 |
+
.await
|
| 752 |
+
.expect("initialize should open client");
|
| 753 |
+
let connection_id = match transport_event_rx.recv().await.expect("open event") {
|
| 754 |
+
TransportEvent::ConnectionOpened { connection_id, .. } => connection_id,
|
| 755 |
+
other => panic!("expected connection opened, got {other:?}"),
|
| 756 |
+
};
|
| 757 |
+
let _ = transport_event_rx.recv().await.expect("initialize event");
|
| 758 |
+
|
| 759 |
+
for _ in 0..2 {
|
| 760 |
+
client_tracker
|
| 761 |
+
.transport_event_tx
|
| 762 |
+
.send(TransportEvent::IncomingMessage {
|
| 763 |
+
connection_id,
|
| 764 |
+
message: initialized_notification(),
|
| 765 |
+
})
|
| 766 |
+
.await
|
| 767 |
+
.expect("transport event queue should accept prefill");
|
| 768 |
+
}
|
| 769 |
+
|
| 770 |
+
let client_key = (
|
| 771 |
+
ClientId("client-1".to_string()),
|
| 772 |
+
StreamId("stream-1".to_string()),
|
| 773 |
+
);
|
| 774 |
+
let close_client = client_tracker.close_client(&client_key);
|
| 775 |
+
tokio::pin!(close_client);
|
| 776 |
+
assert!(
|
| 777 |
+
timeout(Duration::from_millis(20), &mut close_client)
|
| 778 |
+
.await
|
| 779 |
+
.is_err()
|
| 780 |
+
);
|
| 781 |
+
|
| 782 |
+
for _ in 0..2 {
|
| 783 |
+
match transport_event_rx.recv().await.expect("prefilled event") {
|
| 784 |
+
TransportEvent::IncomingMessage {
|
| 785 |
+
connection_id: queued_connection_id,
|
| 786 |
+
..
|
| 787 |
+
} => assert_eq!(queued_connection_id, connection_id),
|
| 788 |
+
other => panic!("expected incoming message, got {other:?}"),
|
| 789 |
+
}
|
| 790 |
+
}
|
| 791 |
+
|
| 792 |
+
close_client
|
| 793 |
+
.await
|
| 794 |
+
.expect("close should forward after queue drains");
|
| 795 |
+
match transport_event_rx.recv().await.expect("close event") {
|
| 796 |
+
TransportEvent::ConnectionClosed {
|
| 797 |
+
connection_id: closed_connection_id,
|
| 798 |
+
} => assert_eq!(closed_connection_id, connection_id),
|
| 799 |
+
other => panic!("expected connection closed, got {other:?}"),
|
| 800 |
+
}
|
| 801 |
+
}
|
| 802 |
+
|
| 803 |
+
#[tokio::test]
|
| 804 |
+
async fn close_client_keeps_forwarding_after_caller_is_aborted() {
|
| 805 |
+
let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 806 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(2);
|
| 807 |
+
let shutdown_token = CancellationToken::new();
|
| 808 |
+
let mut client_tracker =
|
| 809 |
+
ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token);
|
| 810 |
+
|
| 811 |
+
client_tracker
|
| 812 |
+
.handle_message(initialize_envelope_with_stream_id(
|
| 813 |
+
"client-1",
|
| 814 |
+
Some("stream-1"),
|
| 815 |
+
))
|
| 816 |
+
.await
|
| 817 |
+
.expect("initialize should open client");
|
| 818 |
+
let connection_id = match transport_event_rx.recv().await.expect("open event") {
|
| 819 |
+
TransportEvent::ConnectionOpened { connection_id, .. } => connection_id,
|
| 820 |
+
other => panic!("expected connection opened, got {other:?}"),
|
| 821 |
+
};
|
| 822 |
+
let _ = transport_event_rx.recv().await.expect("initialize event");
|
| 823 |
+
|
| 824 |
+
for _ in 0..2 {
|
| 825 |
+
client_tracker
|
| 826 |
+
.transport_event_tx
|
| 827 |
+
.send(TransportEvent::IncomingMessage {
|
| 828 |
+
connection_id,
|
| 829 |
+
message: initialized_notification(),
|
| 830 |
+
})
|
| 831 |
+
.await
|
| 832 |
+
.expect("transport event queue should accept prefill");
|
| 833 |
+
}
|
| 834 |
+
|
| 835 |
+
let client_key = (
|
| 836 |
+
ClientId("client-1".to_string()),
|
| 837 |
+
StreamId("stream-1".to_string()),
|
| 838 |
+
);
|
| 839 |
+
let mut close_client =
|
| 840 |
+
tokio::spawn(async move { client_tracker.close_client(&client_key).await });
|
| 841 |
+
assert!(
|
| 842 |
+
timeout(Duration::from_millis(20), &mut close_client)
|
| 843 |
+
.await
|
| 844 |
+
.is_err()
|
| 845 |
+
);
|
| 846 |
+
close_client.abort();
|
| 847 |
+
let _ = close_client.await;
|
| 848 |
+
|
| 849 |
+
for _ in 0..2 {
|
| 850 |
+
let _ = transport_event_rx.recv().await.expect("prefilled event");
|
| 851 |
+
}
|
| 852 |
+
match timeout(Duration::from_secs(1), transport_event_rx.recv())
|
| 853 |
+
.await
|
| 854 |
+
.expect("close should be delivered")
|
| 855 |
+
.expect("close event")
|
| 856 |
+
{
|
| 857 |
+
TransportEvent::ConnectionClosed {
|
| 858 |
+
connection_id: closed_connection_id,
|
| 859 |
+
} => assert_eq!(closed_connection_id, connection_id),
|
| 860 |
+
other => panic!("expected connection closed, got {other:?}"),
|
| 861 |
+
}
|
| 862 |
+
}
|
| 863 |
+
|
| 864 |
+
#[tokio::test]
|
| 865 |
+
async fn initialize_with_new_stream_id_opens_new_connection_for_same_client() {
|
| 866 |
+
let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 867 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 868 |
+
let shutdown_token = CancellationToken::new();
|
| 869 |
+
let mut client_tracker =
|
| 870 |
+
ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token);
|
| 871 |
+
|
| 872 |
+
client_tracker
|
| 873 |
+
.handle_message(initialize_envelope_with_stream_id(
|
| 874 |
+
"client-1",
|
| 875 |
+
Some("stream-1"),
|
| 876 |
+
))
|
| 877 |
+
.await
|
| 878 |
+
.expect("first initialize should open client");
|
| 879 |
+
let first_connection_id = match transport_event_rx.recv().await.expect("open event") {
|
| 880 |
+
TransportEvent::ConnectionOpened { connection_id, .. } => connection_id,
|
| 881 |
+
other => panic!("expected connection opened, got {other:?}"),
|
| 882 |
+
};
|
| 883 |
+
let _ = transport_event_rx.recv().await.expect("initialize event");
|
| 884 |
+
|
| 885 |
+
client_tracker
|
| 886 |
+
.handle_message(initialize_envelope_with_stream_id(
|
| 887 |
+
"client-1",
|
| 888 |
+
Some("stream-2"),
|
| 889 |
+
))
|
| 890 |
+
.await
|
| 891 |
+
.expect("second initialize should open client");
|
| 892 |
+
let second_connection_id = match transport_event_rx.recv().await.expect("open event") {
|
| 893 |
+
TransportEvent::ConnectionOpened { connection_id, .. } => connection_id,
|
| 894 |
+
other => panic!("expected connection opened, got {other:?}"),
|
| 895 |
+
};
|
| 896 |
+
|
| 897 |
+
assert_ne!(first_connection_id, second_connection_id);
|
| 898 |
+
}
|
| 899 |
+
|
| 900 |
+
#[tokio::test]
|
| 901 |
+
async fn legacy_initialize_without_stream_id_resets_inbound_seq_id() {
|
| 902 |
+
let (server_event_tx, _server_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 903 |
+
let (transport_event_tx, mut transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 904 |
+
let shutdown_token = CancellationToken::new();
|
| 905 |
+
let mut client_tracker =
|
| 906 |
+
ClientTracker::new(server_event_tx, transport_event_tx, &shutdown_token);
|
| 907 |
+
|
| 908 |
+
client_tracker
|
| 909 |
+
.handle_message(initialize_envelope("client-1"))
|
| 910 |
+
.await
|
| 911 |
+
.expect("initialize should open client");
|
| 912 |
+
let connection_id = match transport_event_rx.recv().await.expect("open event") {
|
| 913 |
+
TransportEvent::ConnectionOpened { connection_id, .. } => connection_id,
|
| 914 |
+
other => panic!("expected connection opened, got {other:?}"),
|
| 915 |
+
};
|
| 916 |
+
let _ = transport_event_rx.recv().await.expect("initialize event");
|
| 917 |
+
|
| 918 |
+
client_tracker
|
| 919 |
+
.handle_message(ClientEnvelope {
|
| 920 |
+
event: ClientEvent::ClientMessage {
|
| 921 |
+
message: JSONRPCMessage::Notification(
|
| 922 |
+
codex_app_server_protocol::JSONRPCNotification {
|
| 923 |
+
method: "initialized".to_string(),
|
| 924 |
+
params: None,
|
| 925 |
+
},
|
| 926 |
+
),
|
| 927 |
+
},
|
| 928 |
+
client_id: ClientId("client-1".to_string()),
|
| 929 |
+
stream_id: None,
|
| 930 |
+
seq_id: Some(0),
|
| 931 |
+
cursor: None,
|
| 932 |
+
})
|
| 933 |
+
.await
|
| 934 |
+
.expect("legacy followup should be forwarded");
|
| 935 |
+
|
| 936 |
+
match transport_event_rx.recv().await.expect("followup event") {
|
| 937 |
+
TransportEvent::IncomingMessage {
|
| 938 |
+
connection_id: incoming_connection_id,
|
| 939 |
+
..
|
| 940 |
+
} => assert_eq!(incoming_connection_id, connection_id),
|
| 941 |
+
other => panic!("expected incoming message, got {other:?}"),
|
| 942 |
+
}
|
| 943 |
+
}
|
| 944 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/clients.rs
ADDED
|
@@ -0,0 +1,303 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::auth::RemoteControlAuth;
|
| 2 |
+
use super::auth::RemoteControlConnectionAuth;
|
| 3 |
+
use super::auth::load_remote_control_auth;
|
| 4 |
+
use super::auth::recover_remote_control_auth;
|
| 5 |
+
use super::enroll::format_headers;
|
| 6 |
+
use super::enroll::preview_remote_control_response_body;
|
| 7 |
+
use super::protocol::normalize_remote_control_base_url;
|
| 8 |
+
use axum::http::HeaderMap;
|
| 9 |
+
use codex_app_server_protocol::RemoteControlClient;
|
| 10 |
+
use codex_app_server_protocol::RemoteControlClientsListOrder;
|
| 11 |
+
use codex_app_server_protocol::RemoteControlClientsListParams;
|
| 12 |
+
use codex_app_server_protocol::RemoteControlClientsListResponse;
|
| 13 |
+
use codex_app_server_protocol::RemoteControlClientsRevokeParams;
|
| 14 |
+
use codex_app_server_protocol::RemoteControlClientsRevokeResponse;
|
| 15 |
+
use codex_login::default_client::create_client_without_request_logging;
|
| 16 |
+
use serde::Deserialize;
|
| 17 |
+
use std::io;
|
| 18 |
+
use std::io::ErrorKind;
|
| 19 |
+
use time::OffsetDateTime;
|
| 20 |
+
use time::format_description::well_known::Rfc3339;
|
| 21 |
+
use url::Url;
|
| 22 |
+
|
| 23 |
+
const REMOTE_CONTROL_CLIENT_MANAGEMENT_TIMEOUT: std::time::Duration =
|
| 24 |
+
std::time::Duration::from_secs(30);
|
| 25 |
+
|
| 26 |
+
#[derive(Debug, Deserialize)]
|
| 27 |
+
struct ListRemoteControlClientsResponse {
|
| 28 |
+
items: Vec<RemoteControlClientResponse>,
|
| 29 |
+
#[serde(default)]
|
| 30 |
+
cursor: Option<String>,
|
| 31 |
+
}
|
| 32 |
+
|
| 33 |
+
#[derive(Debug, Deserialize)]
|
| 34 |
+
struct RemoteControlClientResponse {
|
| 35 |
+
client_id: String,
|
| 36 |
+
#[serde(default)]
|
| 37 |
+
display_name: Option<String>,
|
| 38 |
+
#[serde(default)]
|
| 39 |
+
device_type: Option<String>,
|
| 40 |
+
#[serde(default)]
|
| 41 |
+
platform: Option<String>,
|
| 42 |
+
#[serde(default)]
|
| 43 |
+
os_version: Option<String>,
|
| 44 |
+
#[serde(default)]
|
| 45 |
+
device_model: Option<String>,
|
| 46 |
+
#[serde(default)]
|
| 47 |
+
app_version: Option<String>,
|
| 48 |
+
#[serde(default)]
|
| 49 |
+
last_seen_at: Option<String>,
|
| 50 |
+
}
|
| 51 |
+
|
| 52 |
+
enum ClientManagementRequest<'a> {
|
| 53 |
+
List {
|
| 54 |
+
url: &'a Url,
|
| 55 |
+
params: &'a RemoteControlClientsListParams,
|
| 56 |
+
},
|
| 57 |
+
Revoke {
|
| 58 |
+
url: &'a Url,
|
| 59 |
+
},
|
| 60 |
+
}
|
| 61 |
+
|
| 62 |
+
struct ClientManagementResponse {
|
| 63 |
+
status: axum::http::StatusCode,
|
| 64 |
+
headers: HeaderMap,
|
| 65 |
+
body: Vec<u8>,
|
| 66 |
+
}
|
| 67 |
+
|
| 68 |
+
pub(super) async fn list_remote_control_clients(
|
| 69 |
+
remote_control_url: &str,
|
| 70 |
+
auth_manager: &RemoteControlAuth,
|
| 71 |
+
params: RemoteControlClientsListParams,
|
| 72 |
+
) -> io::Result<RemoteControlClientsListResponse> {
|
| 73 |
+
if params.environment_id.is_empty() {
|
| 74 |
+
return Err(io::Error::new(
|
| 75 |
+
ErrorKind::InvalidInput,
|
| 76 |
+
"remote control client list requires environmentId",
|
| 77 |
+
));
|
| 78 |
+
}
|
| 79 |
+
if params
|
| 80 |
+
.limit
|
| 81 |
+
.is_some_and(|limit| !(1..=100).contains(&limit))
|
| 82 |
+
{
|
| 83 |
+
return Err(io::Error::new(
|
| 84 |
+
ErrorKind::InvalidInput,
|
| 85 |
+
"remote control client list limit must be between 1 and 100",
|
| 86 |
+
));
|
| 87 |
+
}
|
| 88 |
+
let url = environment_clients_url(remote_control_url, ¶ms.environment_id)?;
|
| 89 |
+
let response = send_client_management_request(
|
| 90 |
+
auth_manager,
|
| 91 |
+
ClientManagementRequest::List {
|
| 92 |
+
url: &url,
|
| 93 |
+
params: ¶ms,
|
| 94 |
+
},
|
| 95 |
+
"list remote control clients",
|
| 96 |
+
)
|
| 97 |
+
.await?;
|
| 98 |
+
let ClientManagementResponse {
|
| 99 |
+
status,
|
| 100 |
+
headers,
|
| 101 |
+
body,
|
| 102 |
+
} = response;
|
| 103 |
+
let body_preview = preview_remote_control_response_body(&body);
|
| 104 |
+
ensure_success_response(status, &headers, &url, &body_preview, "client list")?;
|
| 105 |
+
let response = serde_json::from_slice::<ListRemoteControlClientsResponse>(&body).map_err(
|
| 106 |
+
|err| {
|
| 107 |
+
io::Error::other(format!(
|
| 108 |
+
"failed to parse remote control client list response from `{url}`: HTTP {status}, {}, body: {body_preview}, decode error: {err}",
|
| 109 |
+
format_headers(&headers)
|
| 110 |
+
))
|
| 111 |
+
},
|
| 112 |
+
)?;
|
| 113 |
+
Ok(RemoteControlClientsListResponse {
|
| 114 |
+
data: response
|
| 115 |
+
.items
|
| 116 |
+
.into_iter()
|
| 117 |
+
.map(RemoteControlClient::try_from)
|
| 118 |
+
.collect::<io::Result<_>>()?,
|
| 119 |
+
next_cursor: response.cursor,
|
| 120 |
+
})
|
| 121 |
+
}
|
| 122 |
+
|
| 123 |
+
pub(super) async fn revoke_remote_control_client(
|
| 124 |
+
remote_control_url: &str,
|
| 125 |
+
auth_manager: &RemoteControlAuth,
|
| 126 |
+
params: RemoteControlClientsRevokeParams,
|
| 127 |
+
) -> io::Result<RemoteControlClientsRevokeResponse> {
|
| 128 |
+
if params.environment_id.is_empty() {
|
| 129 |
+
return Err(io::Error::new(
|
| 130 |
+
ErrorKind::InvalidInput,
|
| 131 |
+
"remote control client revoke requires environmentId",
|
| 132 |
+
));
|
| 133 |
+
}
|
| 134 |
+
if params.client_id.is_empty() {
|
| 135 |
+
return Err(io::Error::new(
|
| 136 |
+
ErrorKind::InvalidInput,
|
| 137 |
+
"remote control client revoke requires clientId",
|
| 138 |
+
));
|
| 139 |
+
}
|
| 140 |
+
let mut url = environment_clients_url(remote_control_url, ¶ms.environment_id)?;
|
| 141 |
+
url.path_segments_mut()
|
| 142 |
+
.map_err(|()| {
|
| 143 |
+
io::Error::new(
|
| 144 |
+
ErrorKind::InvalidInput,
|
| 145 |
+
"remote control URL cannot be a base",
|
| 146 |
+
)
|
| 147 |
+
})?
|
| 148 |
+
.push(¶ms.client_id);
|
| 149 |
+
let response = send_client_management_request(
|
| 150 |
+
auth_manager,
|
| 151 |
+
ClientManagementRequest::Revoke { url: &url },
|
| 152 |
+
"revoke remote control client",
|
| 153 |
+
)
|
| 154 |
+
.await?;
|
| 155 |
+
let ClientManagementResponse {
|
| 156 |
+
status,
|
| 157 |
+
headers,
|
| 158 |
+
body,
|
| 159 |
+
} = response;
|
| 160 |
+
let body_preview = preview_remote_control_response_body(&body);
|
| 161 |
+
ensure_success_response(status, &headers, &url, &body_preview, "client revoke")?;
|
| 162 |
+
Ok(RemoteControlClientsRevokeResponse {})
|
| 163 |
+
}
|
| 164 |
+
|
| 165 |
+
async fn send_client_management_request(
|
| 166 |
+
auth_manager: &RemoteControlAuth,
|
| 167 |
+
request: ClientManagementRequest<'_>,
|
| 168 |
+
action: &str,
|
| 169 |
+
) -> io::Result<ClientManagementResponse> {
|
| 170 |
+
let mut auth_recovery = auth_manager.unauthorized_recovery();
|
| 171 |
+
let mut auth_change_rx = auth_manager.auth_change_receiver();
|
| 172 |
+
let auth = load_remote_control_auth(auth_manager).await?;
|
| 173 |
+
let response = send_client_management_request_once(&auth, &request, action).await?;
|
| 174 |
+
if response.status.as_u16() != 401
|
| 175 |
+
|| !recover_remote_control_auth(&mut auth_recovery, &mut auth_change_rx).await
|
| 176 |
+
{
|
| 177 |
+
return Ok(response);
|
| 178 |
+
}
|
| 179 |
+
let auth = load_remote_control_auth(auth_manager).await?;
|
| 180 |
+
send_client_management_request_once(&auth, &request, action).await
|
| 181 |
+
}
|
| 182 |
+
|
| 183 |
+
async fn send_client_management_request_once(
|
| 184 |
+
auth: &RemoteControlConnectionAuth,
|
| 185 |
+
request: &ClientManagementRequest<'_>,
|
| 186 |
+
action: &str,
|
| 187 |
+
) -> io::Result<ClientManagementResponse> {
|
| 188 |
+
let client = create_client_without_request_logging();
|
| 189 |
+
let auth_headers = auth.request_headers()?;
|
| 190 |
+
let request = match request {
|
| 191 |
+
ClientManagementRequest::List { url, params } => {
|
| 192 |
+
let mut query = Vec::new();
|
| 193 |
+
if let Some(cursor) = ¶ms.cursor {
|
| 194 |
+
query.push(("cursor", cursor.clone()));
|
| 195 |
+
}
|
| 196 |
+
if let Some(limit) = params.limit {
|
| 197 |
+
query.push(("limit", limit.to_string()));
|
| 198 |
+
}
|
| 199 |
+
if let Some(order) = params.order {
|
| 200 |
+
query.push((
|
| 201 |
+
"order",
|
| 202 |
+
match order {
|
| 203 |
+
RemoteControlClientsListOrder::Asc => "asc",
|
| 204 |
+
RemoteControlClientsListOrder::Desc => "desc",
|
| 205 |
+
}
|
| 206 |
+
.to_string(),
|
| 207 |
+
));
|
| 208 |
+
}
|
| 209 |
+
client.get((*url).clone()).query(&query)
|
| 210 |
+
}
|
| 211 |
+
ClientManagementRequest::Revoke { url } => client.delete((*url).clone()),
|
| 212 |
+
};
|
| 213 |
+
let response = request
|
| 214 |
+
.timeout(REMOTE_CONTROL_CLIENT_MANAGEMENT_TIMEOUT)
|
| 215 |
+
.headers(auth_headers)
|
| 216 |
+
.send()
|
| 217 |
+
.await
|
| 218 |
+
.map_err(|err| io::Error::other(format!("failed to {action}: {err}")))?;
|
| 219 |
+
let headers = response.headers().clone();
|
| 220 |
+
let status = response.status();
|
| 221 |
+
let body = response
|
| 222 |
+
.bytes()
|
| 223 |
+
.await
|
| 224 |
+
.map_err(|err| io::Error::other(format!("failed to read {action} response: {err}")))?
|
| 225 |
+
.to_vec();
|
| 226 |
+
Ok(ClientManagementResponse {
|
| 227 |
+
status,
|
| 228 |
+
headers,
|
| 229 |
+
body,
|
| 230 |
+
})
|
| 231 |
+
}
|
| 232 |
+
|
| 233 |
+
fn ensure_success_response(
|
| 234 |
+
status: axum::http::StatusCode,
|
| 235 |
+
headers: &HeaderMap,
|
| 236 |
+
url: &Url,
|
| 237 |
+
body_preview: &str,
|
| 238 |
+
response_kind: &str,
|
| 239 |
+
) -> io::Result<()> {
|
| 240 |
+
if status.is_success() {
|
| 241 |
+
return Ok(());
|
| 242 |
+
}
|
| 243 |
+
let error_kind = match status.as_u16() {
|
| 244 |
+
400 => ErrorKind::InvalidInput,
|
| 245 |
+
401 | 403 => ErrorKind::PermissionDenied,
|
| 246 |
+
404 => ErrorKind::NotFound,
|
| 247 |
+
_ => ErrorKind::Other,
|
| 248 |
+
};
|
| 249 |
+
Err(io::Error::new(
|
| 250 |
+
error_kind,
|
| 251 |
+
format!(
|
| 252 |
+
"remote control {response_kind} failed at `{url}`: HTTP {status}, {}, body: {body_preview}",
|
| 253 |
+
format_headers(headers)
|
| 254 |
+
),
|
| 255 |
+
))
|
| 256 |
+
}
|
| 257 |
+
|
| 258 |
+
fn environment_clients_url(remote_control_url: &str, environment_id: &str) -> io::Result<Url> {
|
| 259 |
+
let mut url = normalize_remote_control_base_url(remote_control_url)?
|
| 260 |
+
.join("wham/remote/control/environments")
|
| 261 |
+
.map_err(io::Error::other)?;
|
| 262 |
+
url.path_segments_mut()
|
| 263 |
+
.map_err(|()| {
|
| 264 |
+
io::Error::new(
|
| 265 |
+
ErrorKind::InvalidInput,
|
| 266 |
+
"remote control URL cannot be a base",
|
| 267 |
+
)
|
| 268 |
+
})?
|
| 269 |
+
.push(environment_id)
|
| 270 |
+
.push("clients");
|
| 271 |
+
Ok(url)
|
| 272 |
+
}
|
| 273 |
+
|
| 274 |
+
impl TryFrom<RemoteControlClientResponse> for RemoteControlClient {
|
| 275 |
+
type Error = io::Error;
|
| 276 |
+
|
| 277 |
+
fn try_from(client: RemoteControlClientResponse) -> Result<Self, Self::Error> {
|
| 278 |
+
Ok(Self {
|
| 279 |
+
client_id: client.client_id,
|
| 280 |
+
display_name: client.display_name,
|
| 281 |
+
device_type: client.device_type,
|
| 282 |
+
platform: client.platform,
|
| 283 |
+
os_version: client.os_version,
|
| 284 |
+
device_model: client.device_model,
|
| 285 |
+
app_version: client.app_version,
|
| 286 |
+
last_seen_at: client
|
| 287 |
+
.last_seen_at
|
| 288 |
+
.map(|last_seen_at| {
|
| 289 |
+
OffsetDateTime::parse(&last_seen_at, &Rfc3339)
|
| 290 |
+
.map(OffsetDateTime::unix_timestamp)
|
| 291 |
+
.map_err(|err| {
|
| 292 |
+
io::Error::new(
|
| 293 |
+
ErrorKind::InvalidData,
|
| 294 |
+
format!(
|
| 295 |
+
"failed to parse remote control client last_seen_at `{last_seen_at}`: {err}"
|
| 296 |
+
),
|
| 297 |
+
)
|
| 298 |
+
})
|
| 299 |
+
})
|
| 300 |
+
.transpose()?,
|
| 301 |
+
})
|
| 302 |
+
}
|
| 303 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/controller.rs
ADDED
|
@@ -0,0 +1,368 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Owns the current login's relay session and the process-facing handle.
|
| 2 |
+
//! Replacement happens before cleanup; retired sessions cannot publish into their replacements.
|
| 3 |
+
|
| 4 |
+
use super::auth::RemoteControlAuth;
|
| 5 |
+
use super::*;
|
| 6 |
+
use futures::FutureExt;
|
| 7 |
+
use std::panic::AssertUnwindSafe;
|
| 8 |
+
use tokio_util::task::TaskTracker;
|
| 9 |
+
|
| 10 |
+
#[derive(Clone)]
|
| 11 |
+
pub struct RemoteControlHandle {
|
| 12 |
+
pub(super) inner: Arc<RemoteControl>,
|
| 13 |
+
}
|
| 14 |
+
|
| 15 |
+
struct CurrentSession {
|
| 16 |
+
authenticated: bool,
|
| 17 |
+
session: Arc<RemoteControlSession>,
|
| 18 |
+
}
|
| 19 |
+
|
| 20 |
+
pub(super) struct RemoteControl {
|
| 21 |
+
config: RemoteControlStartConfig,
|
| 22 |
+
state_db: Option<Arc<StateRuntime>>,
|
| 23 |
+
auth_manager: Arc<AuthManager>,
|
| 24 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 25 |
+
shutdown: CancellationToken,
|
| 26 |
+
tasks: TaskTracker,
|
| 27 |
+
current: StdMutex<Option<CurrentSession>>,
|
| 28 |
+
startup: RemoteControlDesiredState,
|
| 29 |
+
persistence: RemoteControlPersistence,
|
| 30 |
+
client_name: RemoteControlPairingPersistenceKey,
|
| 31 |
+
requires_client_name: bool,
|
| 32 |
+
status: watch::Sender<RemoteControlStatusChangedNotification>,
|
| 33 |
+
session_changed: watch::Sender<()>,
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
impl RemoteControl {
|
| 37 |
+
pub(super) fn session(&self) -> Arc<RemoteControlSession> {
|
| 38 |
+
let mut current = self
|
| 39 |
+
.current
|
| 40 |
+
.lock()
|
| 41 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 42 |
+
if let Some(current) = current.as_ref()
|
| 43 |
+
&& current.session.auth_manager.owner.is_current()
|
| 44 |
+
{
|
| 45 |
+
return current.session.clone();
|
| 46 |
+
}
|
| 47 |
+
let (auth, authenticated) = RemoteControlAuth::capture(self.auth_manager.clone());
|
| 48 |
+
let desired = match current.take() {
|
| 49 |
+
Some(previous) => {
|
| 50 |
+
previous.session.shutdown_token.cancel();
|
| 51 |
+
if previous.authenticated {
|
| 52 |
+
RemoteControlDesiredState::Disabled
|
| 53 |
+
} else {
|
| 54 |
+
// Startup or an explicit enable while signed out may wait for first login.
|
| 55 |
+
*previous.session.desired_state_tx.borrow()
|
| 56 |
+
}
|
| 57 |
+
}
|
| 58 |
+
None => self.startup,
|
| 59 |
+
};
|
| 60 |
+
let session = self.start_session(auth, desired);
|
| 61 |
+
self.status.send_replace(session.status());
|
| 62 |
+
*current = Some(CurrentSession {
|
| 63 |
+
authenticated,
|
| 64 |
+
session: session.clone(),
|
| 65 |
+
});
|
| 66 |
+
self.session_changed.send_replace(());
|
| 67 |
+
session
|
| 68 |
+
}
|
| 69 |
+
|
| 70 |
+
fn start_session(
|
| 71 |
+
&self,
|
| 72 |
+
auth_manager: RemoteControlAuth,
|
| 73 |
+
desired: RemoteControlDesiredState,
|
| 74 |
+
) -> Arc<RemoteControlSession> {
|
| 75 |
+
let shutdown = self.shutdown.child_token();
|
| 76 |
+
let (desired_state_tx, _) = watch::channel(desired);
|
| 77 |
+
let desired_state_tx = Arc::new(desired_state_tx);
|
| 78 |
+
let current_enrollment =
|
| 79 |
+
Arc::new(RemoteControlEnrollmentState::new(/*enrollment*/ None));
|
| 80 |
+
let server_name = gethostname().to_string_lossy().trim().to_string();
|
| 81 |
+
let (status_tx, _) = watch::channel(RemoteControlStatusChangedNotification {
|
| 82 |
+
status: if desired.is_enabled() {
|
| 83 |
+
RemoteControlConnectionStatus::Connecting
|
| 84 |
+
} else {
|
| 85 |
+
RemoteControlConnectionStatus::Disabled
|
| 86 |
+
},
|
| 87 |
+
server_name: server_name.clone(),
|
| 88 |
+
installation_id: self.config.installation_id.clone(),
|
| 89 |
+
environment_id: None,
|
| 90 |
+
});
|
| 91 |
+
let session = Arc::new(RemoteControlSession {
|
| 92 |
+
policy: self.config.policy,
|
| 93 |
+
shutdown_token: shutdown.clone(),
|
| 94 |
+
desired_state_tx: desired_state_tx.clone(),
|
| 95 |
+
desired_state_rpc_lock: Arc::new(Semaphore::new(1)),
|
| 96 |
+
persistence: self.persistence.clone(),
|
| 97 |
+
status_tx: Arc::new(status_tx.clone()),
|
| 98 |
+
state_db: self.state_db.clone(),
|
| 99 |
+
remote_control_url: self.config.remote_control_url.clone(),
|
| 100 |
+
current_enrollment: current_enrollment.clone(),
|
| 101 |
+
pairing_persistence_key: self.client_name.clone(),
|
| 102 |
+
pairing_persistence_key_required: self.requires_client_name,
|
| 103 |
+
auth_manager: auth_manager.clone(),
|
| 104 |
+
});
|
| 105 |
+
let websocket = RemoteControlWebsocket::new(
|
| 106 |
+
websocket::RemoteControlWebsocketConfig {
|
| 107 |
+
remote_control_url: self.config.remote_control_url.clone(),
|
| 108 |
+
installation_id: self.config.installation_id.clone(),
|
| 109 |
+
remote_control_target: None,
|
| 110 |
+
server_name,
|
| 111 |
+
},
|
| 112 |
+
self.state_db.clone(),
|
| 113 |
+
auth_manager,
|
| 114 |
+
RemoteControlChannels {
|
| 115 |
+
transport_event_tx: self.transport_event_tx.clone(),
|
| 116 |
+
status_publisher: RemoteControlStatusPublisher::new(status_tx),
|
| 117 |
+
current_enrollment,
|
| 118 |
+
pairing_persistence_key: self.client_name.clone(),
|
| 119 |
+
persistence: self.persistence.clone(),
|
| 120 |
+
},
|
| 121 |
+
shutdown.clone(),
|
| 122 |
+
desired_state_tx,
|
| 123 |
+
);
|
| 124 |
+
let client_name_rx = if self.requires_client_name {
|
| 125 |
+
let (tx, rx) = oneshot::channel();
|
| 126 |
+
let mut names = self.client_name.subscribe();
|
| 127 |
+
self.tasks.spawn(async move {
|
| 128 |
+
tokio::select! {
|
| 129 |
+
_ = shutdown.cancelled() => {}
|
| 130 |
+
name = names.wait_for(Option::is_some) => {
|
| 131 |
+
if let Ok(name) = name
|
| 132 |
+
&& let Some(name) = name.as_ref()
|
| 133 |
+
{
|
| 134 |
+
let _ = tx.send(name.clone());
|
| 135 |
+
}
|
| 136 |
+
}
|
| 137 |
+
}
|
| 138 |
+
});
|
| 139 |
+
Some(rx)
|
| 140 |
+
} else {
|
| 141 |
+
None
|
| 142 |
+
};
|
| 143 |
+
let process_shutdown = self.shutdown.clone();
|
| 144 |
+
let failed_session = session.clone();
|
| 145 |
+
self.tasks.spawn(async move {
|
| 146 |
+
if let Err(panic) = AssertUnwindSafe(websocket.run(client_name_rx))
|
| 147 |
+
.catch_unwind()
|
| 148 |
+
.await
|
| 149 |
+
{
|
| 150 |
+
tracing::error!("remote control websocket task panicked");
|
| 151 |
+
failed_session.publish_status(RemoteControlConnectionStatus::Disabled);
|
| 152 |
+
process_shutdown.cancel();
|
| 153 |
+
std::panic::resume_unwind(panic);
|
| 154 |
+
}
|
| 155 |
+
});
|
| 156 |
+
session
|
| 157 |
+
}
|
| 158 |
+
}
|
| 159 |
+
|
| 160 |
+
impl RemoteControlSession {
|
| 161 |
+
async fn run<T>(
|
| 162 |
+
&self,
|
| 163 |
+
operation: impl std::future::Future<Output = io::Result<T>>,
|
| 164 |
+
) -> io::Result<T> {
|
| 165 |
+
self.auth_manager.ensure_current()?;
|
| 166 |
+
tokio::select! {
|
| 167 |
+
biased;
|
| 168 |
+
_ = self.auth_manager.owner.invalidated() => {
|
| 169 |
+
Err(io::Error::new(io::ErrorKind::Interrupted, "remote control authentication changed"))
|
| 170 |
+
}
|
| 171 |
+
result = operation => {
|
| 172 |
+
self.auth_manager.ensure_current()?;
|
| 173 |
+
result
|
| 174 |
+
}
|
| 175 |
+
}
|
| 176 |
+
}
|
| 177 |
+
}
|
| 178 |
+
|
| 179 |
+
impl RemoteControlHandle {
|
| 180 |
+
pub fn ensure_remote_control_allowed(&self) -> Result<(), RemoteControlDisabledByRequirements> {
|
| 181 |
+
self.inner.session().ensure_remote_control_allowed()
|
| 182 |
+
}
|
| 183 |
+
|
| 184 |
+
pub fn status(&self) -> RemoteControlStatusChangedNotification {
|
| 185 |
+
self.inner.session().status()
|
| 186 |
+
}
|
| 187 |
+
|
| 188 |
+
pub fn status_receiver(&self) -> watch::Receiver<RemoteControlStatusChangedNotification> {
|
| 189 |
+
self.inner.session();
|
| 190 |
+
self.inner.status.subscribe()
|
| 191 |
+
}
|
| 192 |
+
|
| 193 |
+
pub fn enable_ephemeral(
|
| 194 |
+
&self,
|
| 195 |
+
) -> Result<RemoteControlStatusChangedNotification, RemoteControlEnableError> {
|
| 196 |
+
self.inner.session().enable_ephemeral()
|
| 197 |
+
}
|
| 198 |
+
|
| 199 |
+
pub async fn disable_ephemeral(&self) -> RemoteControlStatusChangedNotification {
|
| 200 |
+
self.inner.session().disable_ephemeral().await
|
| 201 |
+
}
|
| 202 |
+
|
| 203 |
+
pub async fn enable(
|
| 204 |
+
&self,
|
| 205 |
+
app_server_client_name: Option<&str>,
|
| 206 |
+
) -> io::Result<RemoteControlStatusChangedNotification> {
|
| 207 |
+
let session = self.inner.session();
|
| 208 |
+
session.run(session.enable(app_server_client_name)).await
|
| 209 |
+
}
|
| 210 |
+
|
| 211 |
+
pub async fn disable(
|
| 212 |
+
&self,
|
| 213 |
+
app_server_client_name: Option<&str>,
|
| 214 |
+
) -> io::Result<RemoteControlStatusChangedNotification> {
|
| 215 |
+
let session = self.inner.session();
|
| 216 |
+
session.run(session.disable(app_server_client_name)).await
|
| 217 |
+
}
|
| 218 |
+
|
| 219 |
+
pub async fn resolve_persisted_preference(
|
| 220 |
+
&self,
|
| 221 |
+
app_server_client_name: Option<&str>,
|
| 222 |
+
) -> io::Result<bool> {
|
| 223 |
+
let session = self.inner.session();
|
| 224 |
+
session
|
| 225 |
+
.run(session.resolve_persisted_preference(app_server_client_name))
|
| 226 |
+
.await
|
| 227 |
+
}
|
| 228 |
+
|
| 229 |
+
pub async fn start_pairing(
|
| 230 |
+
&self,
|
| 231 |
+
params: RemoteControlPairingStartParams,
|
| 232 |
+
app_server_client_name: Option<&str>,
|
| 233 |
+
) -> io::Result<RemoteControlPairingStartResponse> {
|
| 234 |
+
let session = self.inner.session();
|
| 235 |
+
session
|
| 236 |
+
.run(session.start_pairing(params, app_server_client_name))
|
| 237 |
+
.await
|
| 238 |
+
}
|
| 239 |
+
|
| 240 |
+
pub async fn pairing_status(
|
| 241 |
+
&self,
|
| 242 |
+
params: RemoteControlPairingStatusParams,
|
| 243 |
+
) -> io::Result<RemoteControlPairingStatusResponse> {
|
| 244 |
+
let session = self.inner.session();
|
| 245 |
+
session.run(session.pairing_status(params)).await
|
| 246 |
+
}
|
| 247 |
+
|
| 248 |
+
pub async fn list_clients(
|
| 249 |
+
&self,
|
| 250 |
+
params: RemoteControlClientsListParams,
|
| 251 |
+
) -> io::Result<RemoteControlClientsListResponse> {
|
| 252 |
+
let session = self.inner.session();
|
| 253 |
+
session.run(session.list_clients(params)).await
|
| 254 |
+
}
|
| 255 |
+
|
| 256 |
+
pub async fn revoke_client(
|
| 257 |
+
&self,
|
| 258 |
+
params: RemoteControlClientsRevokeParams,
|
| 259 |
+
) -> io::Result<RemoteControlClientsRevokeResponse> {
|
| 260 |
+
let session = self.inner.session();
|
| 261 |
+
session.run(session.revoke_client(params)).await
|
| 262 |
+
}
|
| 263 |
+
}
|
| 264 |
+
|
| 265 |
+
pub async fn start_remote_control(
|
| 266 |
+
config: RemoteControlStartConfig,
|
| 267 |
+
state_db: Option<Arc<StateRuntime>>,
|
| 268 |
+
auth_manager: Arc<AuthManager>,
|
| 269 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 270 |
+
shutdown_token: CancellationToken,
|
| 271 |
+
app_server_client_name_rx: Option<oneshot::Receiver<String>>,
|
| 272 |
+
startup_mode: RemoteControlStartupMode,
|
| 273 |
+
) -> io::Result<(JoinHandle<()>, RemoteControlHandle)> {
|
| 274 |
+
let startup =
|
| 275 |
+
if config.policy == RemoteControlPolicy::DisabledByRequirements || state_db.is_none() {
|
| 276 |
+
RemoteControlDesiredState::Disabled
|
| 277 |
+
} else {
|
| 278 |
+
match startup_mode {
|
| 279 |
+
RemoteControlStartupMode::ResolvePersisted => RemoteControlDesiredState::Unknown,
|
| 280 |
+
RemoteControlStartupMode::DisabledEphemeral => RemoteControlDesiredState::Disabled,
|
| 281 |
+
RemoteControlStartupMode::EnabledEphemeral => {
|
| 282 |
+
normalize_remote_control_url(&config.remote_control_url)?;
|
| 283 |
+
RemoteControlDesiredState::Enabled {
|
| 284 |
+
persistence_preference: None,
|
| 285 |
+
}
|
| 286 |
+
}
|
| 287 |
+
}
|
| 288 |
+
};
|
| 289 |
+
let (status, _) = watch::channel(RemoteControlStatusChangedNotification {
|
| 290 |
+
status: RemoteControlConnectionStatus::Disabled,
|
| 291 |
+
server_name: gethostname().to_string_lossy().trim().to_string(),
|
| 292 |
+
installation_id: config.installation_id.clone(),
|
| 293 |
+
environment_id: None,
|
| 294 |
+
});
|
| 295 |
+
let inner = Arc::new(RemoteControl {
|
| 296 |
+
config,
|
| 297 |
+
state_db,
|
| 298 |
+
auth_manager,
|
| 299 |
+
transport_event_tx,
|
| 300 |
+
shutdown: shutdown_token,
|
| 301 |
+
tasks: TaskTracker::new(),
|
| 302 |
+
current: StdMutex::new(None),
|
| 303 |
+
startup,
|
| 304 |
+
persistence: RemoteControlPersistence::default(),
|
| 305 |
+
client_name: watch::channel(None).0,
|
| 306 |
+
requires_client_name: app_server_client_name_rx.is_some(),
|
| 307 |
+
status,
|
| 308 |
+
session_changed: watch::channel(()).0,
|
| 309 |
+
});
|
| 310 |
+
inner.session();
|
| 311 |
+
if let Some(rx) = app_server_client_name_rx {
|
| 312 |
+
let names = inner.client_name.clone();
|
| 313 |
+
let shutdown = inner.shutdown.clone();
|
| 314 |
+
inner.tasks.spawn(async move {
|
| 315 |
+
tokio::select! {
|
| 316 |
+
_ = shutdown.cancelled() => {}
|
| 317 |
+
name = rx => match name {
|
| 318 |
+
Ok(name) => { names.send_replace(Some(name)); }
|
| 319 |
+
Err(_) => shutdown.cancel(),
|
| 320 |
+
}
|
| 321 |
+
}
|
| 322 |
+
});
|
| 323 |
+
}
|
| 324 |
+
let handle = RemoteControlHandle {
|
| 325 |
+
inner: inner.clone(),
|
| 326 |
+
};
|
| 327 |
+
let task = tokio::spawn(async move {
|
| 328 |
+
let mut session_changed = inner.session_changed.subscribe();
|
| 329 |
+
loop {
|
| 330 |
+
// Reconcile on both API entry and notification delivery, so watcher latency is harmless.
|
| 331 |
+
let session = inner.session();
|
| 332 |
+
let mut status = session.status_receiver();
|
| 333 |
+
session_changed.borrow_and_update();
|
| 334 |
+
{
|
| 335 |
+
let current = inner
|
| 336 |
+
.current
|
| 337 |
+
.lock()
|
| 338 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 339 |
+
if current
|
| 340 |
+
.as_ref()
|
| 341 |
+
.is_some_and(|current| Arc::ptr_eq(¤t.session, &session))
|
| 342 |
+
&& session.auth_manager.owner.is_current()
|
| 343 |
+
{
|
| 344 |
+
inner.status.send_if_modified(|current| {
|
| 345 |
+
let next = status.borrow_and_update().clone();
|
| 346 |
+
if *current == next {
|
| 347 |
+
return false;
|
| 348 |
+
}
|
| 349 |
+
*current = next;
|
| 350 |
+
true
|
| 351 |
+
});
|
| 352 |
+
}
|
| 353 |
+
}
|
| 354 |
+
tokio::select! {
|
| 355 |
+
biased;
|
| 356 |
+
_ = inner.shutdown.cancelled() => break,
|
| 357 |
+
_ = session.auth_manager.owner.invalidated() => {}
|
| 358 |
+
_ = session_changed.changed() => {}
|
| 359 |
+
_ = status.changed() => {}
|
| 360 |
+
}
|
| 361 |
+
}
|
| 362 |
+
inner.tasks.close();
|
| 363 |
+
inner.tasks.wait().await;
|
| 364 |
+
inner.persistence.tasks.close();
|
| 365 |
+
inner.persistence.tasks.wait().await;
|
| 366 |
+
});
|
| 367 |
+
Ok((task, handle))
|
| 368 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/desired_state.rs
ADDED
|
@@ -0,0 +1,155 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::RemoteControlEnableError;
|
| 2 |
+
use super::RemoteControlSession;
|
| 3 |
+
use super::RemoteControlUnavailable;
|
| 4 |
+
use super::protocol::normalize_remote_control_url;
|
| 5 |
+
use super::publish_current_enrollment;
|
| 6 |
+
use super::websocket::RemoteControlStatusPublisher;
|
| 7 |
+
use codex_app_server_protocol::RemoteControlStatusChangedNotification;
|
| 8 |
+
use codex_state::RemoteControlEnrollmentRecord;
|
| 9 |
+
use std::io;
|
| 10 |
+
|
| 11 |
+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
| 12 |
+
pub(super) enum RemoteControlDesiredState {
|
| 13 |
+
// `Unknown` exists only on plain startup before auth and enrollment scope resolve. Persisted
|
| 14 |
+
// `1` is `Enabled { persistence_preference: Some(true) }`; `0`, `NULL`, or no row are
|
| 15 |
+
// `Disabled`. Runtime-only enable is `Enabled { persistence_preference: None }`, so new rows
|
| 16 |
+
// keep `NULL`; durable RPC enable uses `Some(true)`, so new rows get `1`. Durable disable writes
|
| 17 |
+
// `0` before entering `Disabled`; runtime-only disable does not write. `Disabled` carries no
|
| 18 |
+
// preference because disabled sessions do not create enrollments.
|
| 19 |
+
Unknown,
|
| 20 |
+
Disabled,
|
| 21 |
+
Enabled {
|
| 22 |
+
persistence_preference: Option<bool>,
|
| 23 |
+
},
|
| 24 |
+
}
|
| 25 |
+
impl RemoteControlDesiredState {
|
| 26 |
+
pub(super) fn is_enabled(self) -> bool {
|
| 27 |
+
matches!(self, Self::Enabled { .. })
|
| 28 |
+
}
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
pub(super) fn desired_state_from_persisted_enrollment(
|
| 32 |
+
enrollment: Option<RemoteControlEnrollmentRecord>,
|
| 33 |
+
) -> RemoteControlDesiredState {
|
| 34 |
+
if enrollment.and_then(|enrollment| enrollment.remote_control_enabled) == Some(true) {
|
| 35 |
+
RemoteControlDesiredState::Enabled {
|
| 36 |
+
persistence_preference: Some(true),
|
| 37 |
+
}
|
| 38 |
+
} else {
|
| 39 |
+
RemoteControlDesiredState::Disabled
|
| 40 |
+
}
|
| 41 |
+
}
|
| 42 |
+
|
| 43 |
+
impl RemoteControlSession {
|
| 44 |
+
pub async fn resolve_persisted_preference(
|
| 45 |
+
&self,
|
| 46 |
+
app_server_client_name: Option<&str>,
|
| 47 |
+
) -> io::Result<bool> {
|
| 48 |
+
if self.ensure_remote_control_allowed().is_err() {
|
| 49 |
+
return Ok(false);
|
| 50 |
+
}
|
| 51 |
+
let _transition = self
|
| 52 |
+
.desired_state_rpc_lock
|
| 53 |
+
.acquire()
|
| 54 |
+
.await
|
| 55 |
+
.unwrap_or_else(|_| unreachable!());
|
| 56 |
+
if !matches!(
|
| 57 |
+
*self.desired_state_tx.borrow(),
|
| 58 |
+
RemoteControlDesiredState::Unknown
|
| 59 |
+
) {
|
| 60 |
+
return Ok(self.desired_state_tx.borrow().is_enabled());
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
let state_db = self
|
| 64 |
+
.state_db
|
| 65 |
+
.as_deref()
|
| 66 |
+
.ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, RemoteControlUnavailable))?;
|
| 67 |
+
let auth = super::auth::load_remote_control_auth(&self.auth_manager).await?;
|
| 68 |
+
let remote_control_target = normalize_remote_control_url(&self.remote_control_url)?;
|
| 69 |
+
let app_server_client_name = self.pairing_persistence_key(app_server_client_name)?;
|
| 70 |
+
let _persistence =
|
| 71 |
+
super::persistence::read_lock(&self.auth_manager, &self.persistence).await?;
|
| 72 |
+
let enrollment = state_db
|
| 73 |
+
.get_remote_control_enrollment(
|
| 74 |
+
&remote_control_target.websocket_url,
|
| 75 |
+
&auth.account_id,
|
| 76 |
+
app_server_client_name.as_deref(),
|
| 77 |
+
)
|
| 78 |
+
.await
|
| 79 |
+
.map_err(io::Error::other)?;
|
| 80 |
+
let desired_state = desired_state_from_persisted_enrollment(enrollment);
|
| 81 |
+
self.desired_state_tx.send_if_modified(|state| {
|
| 82 |
+
if !matches!(*state, RemoteControlDesiredState::Unknown) {
|
| 83 |
+
return false;
|
| 84 |
+
}
|
| 85 |
+
*state = desired_state;
|
| 86 |
+
true
|
| 87 |
+
});
|
| 88 |
+
Ok(self.desired_state_tx.borrow().is_enabled())
|
| 89 |
+
}
|
| 90 |
+
|
| 91 |
+
pub async fn enable(
|
| 92 |
+
&self,
|
| 93 |
+
app_server_client_name: Option<&str>,
|
| 94 |
+
) -> io::Result<RemoteControlStatusChangedNotification> {
|
| 95 |
+
self.ensure_remote_control_allowed()
|
| 96 |
+
.map_err(|err| io::Error::new(io::ErrorKind::PermissionDenied, err))?;
|
| 97 |
+
let _transition = self
|
| 98 |
+
.desired_state_rpc_lock
|
| 99 |
+
.acquire()
|
| 100 |
+
.await
|
| 101 |
+
.unwrap_or_else(|_| unreachable!());
|
| 102 |
+
let state_db = self
|
| 103 |
+
.state_db
|
| 104 |
+
.as_deref()
|
| 105 |
+
.ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, RemoteControlUnavailable))?;
|
| 106 |
+
let mut auth = super::auth::load_remote_control_auth(&self.auth_manager).await?;
|
| 107 |
+
let remote_control_target = normalize_remote_control_url(&self.remote_control_url)?;
|
| 108 |
+
let app_server_client_name = self.pairing_persistence_key(app_server_client_name)?;
|
| 109 |
+
let app_server_client_name = app_server_client_name.as_deref();
|
| 110 |
+
let status = self.status();
|
| 111 |
+
let mut current_enrollment = self.current_enrollment.lock().await;
|
| 112 |
+
let (enrollment, _) = self
|
| 113 |
+
.load_or_enroll_server(
|
| 114 |
+
¤t_enrollment,
|
| 115 |
+
&mut auth,
|
| 116 |
+
&status.installation_id,
|
| 117 |
+
&status.server_name,
|
| 118 |
+
app_server_client_name,
|
| 119 |
+
super::RemoteControlEnrollmentSelection::ReuseOrCreate,
|
| 120 |
+
)
|
| 121 |
+
.await?;
|
| 122 |
+
|
| 123 |
+
let current_auth = super::auth::load_remote_control_auth(&self.auth_manager).await?;
|
| 124 |
+
if current_auth.account_id != auth.account_id {
|
| 125 |
+
return Err(io::Error::new(
|
| 126 |
+
io::ErrorKind::Interrupted,
|
| 127 |
+
"remote control account changed during enrollment",
|
| 128 |
+
));
|
| 129 |
+
}
|
| 130 |
+
|
| 131 |
+
self.set_preference(
|
| 132 |
+
state_db,
|
| 133 |
+
&remote_control_target,
|
| 134 |
+
&auth.account_id,
|
| 135 |
+
app_server_client_name,
|
| 136 |
+
/*enabled*/ true,
|
| 137 |
+
Some(&enrollment),
|
| 138 |
+
)
|
| 139 |
+
.await?;
|
| 140 |
+
publish_current_enrollment(&mut current_enrollment, &enrollment);
|
| 141 |
+
self.enable_with_preference(Some(true)).map_err(|err| {
|
| 142 |
+
let kind = match err {
|
| 143 |
+
RemoteControlEnableError::Unavailable(_) => io::ErrorKind::NotFound,
|
| 144 |
+
RemoteControlEnableError::AuthenticationChanged => io::ErrorKind::Interrupted,
|
| 145 |
+
RemoteControlEnableError::DisabledByRequirements(_) => {
|
| 146 |
+
io::ErrorKind::PermissionDenied
|
| 147 |
+
}
|
| 148 |
+
};
|
| 149 |
+
io::Error::new(kind, err)
|
| 150 |
+
})?;
|
| 151 |
+
RemoteControlStatusPublisher::new(self.status_tx.as_ref().clone())
|
| 152 |
+
.publish_environment_id(Some(enrollment.environment_id));
|
| 153 |
+
Ok(self.status())
|
| 154 |
+
}
|
| 155 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/enroll.rs
ADDED
|
@@ -0,0 +1,770 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::pairing_unavailable_error;
|
| 2 |
+
use super::protocol::RemoteControlPairingStatusRequest;
|
| 3 |
+
use super::protocol::RemoteControlPairingStatusResponse as BackendRemoteControlPairingStatusResponse;
|
| 4 |
+
use super::protocol::RemoteControlTarget;
|
| 5 |
+
use super::protocol::StartRemoteControlPairingRequest;
|
| 6 |
+
use super::protocol::StartRemoteControlPairingResponse;
|
| 7 |
+
use super::server_api::RemoteControlServerRequestError;
|
| 8 |
+
use super::server_api::retry_after_with_jitter;
|
| 9 |
+
use axum::http::HeaderMap;
|
| 10 |
+
use axum::http::StatusCode;
|
| 11 |
+
use codex_app_server_protocol::RemoteControlPairingStartResponse;
|
| 12 |
+
use codex_app_server_protocol::RemoteControlPairingStatusResponse;
|
| 13 |
+
use codex_login::default_client::create_client_without_request_logging;
|
| 14 |
+
use codex_state::RemoteControlEnrollmentRecord;
|
| 15 |
+
use codex_state::StateRuntime;
|
| 16 |
+
use std::io;
|
| 17 |
+
use std::io::ErrorKind;
|
| 18 |
+
use time::OffsetDateTime;
|
| 19 |
+
use time::format_description::well_known::Rfc3339;
|
| 20 |
+
use tracing::info;
|
| 21 |
+
use tracing::warn;
|
| 22 |
+
|
| 23 |
+
const REMOTE_CONTROL_PAIRING_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
|
| 24 |
+
const REMOTE_CONTROL_RESPONSE_BODY_MAX_BYTES: usize = 4096;
|
| 25 |
+
const REMOTE_CONTROL_SERVER_TOKEN_REFRESH_SKEW_SECS: i64 = 5 * 60;
|
| 26 |
+
|
| 27 |
+
const REQUEST_ID_HEADER: &str = "x-request-id";
|
| 28 |
+
const OAI_REQUEST_ID_HEADER: &str = "x-oai-request-id";
|
| 29 |
+
const CF_RAY_HEADER: &str = "cf-ray";
|
| 30 |
+
|
| 31 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 32 |
+
pub(super) struct RemoteControlEnrollment {
|
| 33 |
+
pub(super) remote_control_target: RemoteControlTarget,
|
| 34 |
+
pub(super) account_id: String,
|
| 35 |
+
pub(super) environment_id: String,
|
| 36 |
+
pub(super) server_id: String,
|
| 37 |
+
pub(super) server_name: String,
|
| 38 |
+
pub(super) remote_control_token: Option<String>,
|
| 39 |
+
pub(super) expires_at: Option<OffsetDateTime>,
|
| 40 |
+
pub(super) next_refresh_at: Option<OffsetDateTime>,
|
| 41 |
+
}
|
| 42 |
+
|
| 43 |
+
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
| 44 |
+
pub(super) enum RemoteControlServerTokenRefreshRequirement {
|
| 45 |
+
Required,
|
| 46 |
+
Proactive,
|
| 47 |
+
NotNeeded,
|
| 48 |
+
}
|
| 49 |
+
|
| 50 |
+
impl RemoteControlEnrollment {
|
| 51 |
+
pub(super) async fn start_pairing(
|
| 52 |
+
&self,
|
| 53 |
+
request: StartRemoteControlPairingRequest,
|
| 54 |
+
) -> io::Result<RemoteControlPairingStartResponse> {
|
| 55 |
+
if self.server_token_refresh_requirement()
|
| 56 |
+
== RemoteControlServerTokenRefreshRequirement::Required
|
| 57 |
+
{
|
| 58 |
+
return Err(pairing_unavailable_error());
|
| 59 |
+
}
|
| 60 |
+
let remote_control_token = self
|
| 61 |
+
.remote_control_token
|
| 62 |
+
.as_deref()
|
| 63 |
+
.ok_or_else(pairing_unavailable_error)?;
|
| 64 |
+
|
| 65 |
+
let response = create_client_without_request_logging()
|
| 66 |
+
.post(&self.remote_control_target.pair_url)
|
| 67 |
+
.timeout(REMOTE_CONTROL_PAIRING_TIMEOUT)
|
| 68 |
+
.bearer_auth(remote_control_token)
|
| 69 |
+
.json(&request)
|
| 70 |
+
.send()
|
| 71 |
+
.await
|
| 72 |
+
.map_err(|err| {
|
| 73 |
+
io::Error::other(format!(
|
| 74 |
+
"failed to start remote control pairing at `{}`: {err}",
|
| 75 |
+
self.remote_control_target.pair_url
|
| 76 |
+
))
|
| 77 |
+
})?;
|
| 78 |
+
let headers = response.headers().clone();
|
| 79 |
+
let status = response.status();
|
| 80 |
+
let retry_at = retry_after_with_jitter(&headers, OffsetDateTime::now_utc());
|
| 81 |
+
let body = response.bytes().await.map_err(|err| {
|
| 82 |
+
pairing_response_error(
|
| 83 |
+
format!(
|
| 84 |
+
"failed to read remote control pairing response from `{}`: {err}",
|
| 85 |
+
self.remote_control_target.pair_url
|
| 86 |
+
),
|
| 87 |
+
status,
|
| 88 |
+
retry_at,
|
| 89 |
+
ErrorKind::Other,
|
| 90 |
+
)
|
| 91 |
+
})?;
|
| 92 |
+
let body_preview = preview_remote_control_response_body(&body);
|
| 93 |
+
if !status.is_success() {
|
| 94 |
+
let error_kind = match status.as_u16() {
|
| 95 |
+
401 | 403 => ErrorKind::PermissionDenied,
|
| 96 |
+
404 => ErrorKind::NotFound,
|
| 97 |
+
_ => ErrorKind::Other,
|
| 98 |
+
};
|
| 99 |
+
return Err(pairing_response_error(
|
| 100 |
+
format!(
|
| 101 |
+
"remote control pairing failed at `{}`: HTTP {status}, {}, body: {body_preview}",
|
| 102 |
+
self.remote_control_target.pair_url,
|
| 103 |
+
format_headers(&headers)
|
| 104 |
+
),
|
| 105 |
+
status,
|
| 106 |
+
retry_at,
|
| 107 |
+
error_kind,
|
| 108 |
+
));
|
| 109 |
+
}
|
| 110 |
+
|
| 111 |
+
let pairing = serde_json::from_slice::<StartRemoteControlPairingResponse>(&body).map_err(
|
| 112 |
+
|err| {
|
| 113 |
+
io::Error::other(format!(
|
| 114 |
+
"failed to parse remote control pairing response from `{}`: HTTP {status}, {}, body: {body_preview}, decode error: {err}",
|
| 115 |
+
self.remote_control_target.pair_url,
|
| 116 |
+
format_headers(&headers)
|
| 117 |
+
))
|
| 118 |
+
},
|
| 119 |
+
)?;
|
| 120 |
+
let StartRemoteControlPairingResponse {
|
| 121 |
+
pairing_code,
|
| 122 |
+
manual_pairing_code,
|
| 123 |
+
server_id,
|
| 124 |
+
environment_id,
|
| 125 |
+
expires_at,
|
| 126 |
+
} = pairing;
|
| 127 |
+
if server_id != self.server_id || environment_id != self.environment_id {
|
| 128 |
+
return Err(io::Error::other(format!(
|
| 129 |
+
"remote control pairing returned mismatched enrollment: expected server_id={}, environment_id={}; got server_id={}, environment_id={}",
|
| 130 |
+
self.server_id, self.environment_id, server_id, environment_id
|
| 131 |
+
)));
|
| 132 |
+
}
|
| 133 |
+
let expires_at = OffsetDateTime::parse(&expires_at, &Rfc3339)
|
| 134 |
+
.map_err(|err| {
|
| 135 |
+
io::Error::new(
|
| 136 |
+
ErrorKind::InvalidData,
|
| 137 |
+
format!(
|
| 138 |
+
"failed to parse remote control pairing response from `{}`: HTTP {status}, {}, body: {body_preview}, expires_at parse error: {err}",
|
| 139 |
+
self.remote_control_target.pair_url,
|
| 140 |
+
format_headers(&headers)
|
| 141 |
+
),
|
| 142 |
+
)
|
| 143 |
+
})?
|
| 144 |
+
.unix_timestamp();
|
| 145 |
+
|
| 146 |
+
Ok(RemoteControlPairingStartResponse {
|
| 147 |
+
pairing_code,
|
| 148 |
+
manual_pairing_code,
|
| 149 |
+
environment_id,
|
| 150 |
+
expires_at,
|
| 151 |
+
})
|
| 152 |
+
}
|
| 153 |
+
|
| 154 |
+
pub(super) async fn pairing_status(
|
| 155 |
+
&self,
|
| 156 |
+
request: RemoteControlPairingStatusRequest,
|
| 157 |
+
) -> io::Result<RemoteControlPairingStatusResponse> {
|
| 158 |
+
if self.server_token_refresh_requirement()
|
| 159 |
+
== RemoteControlServerTokenRefreshRequirement::Required
|
| 160 |
+
{
|
| 161 |
+
return Err(pairing_unavailable_error());
|
| 162 |
+
}
|
| 163 |
+
let remote_control_token = self
|
| 164 |
+
.remote_control_token
|
| 165 |
+
.as_deref()
|
| 166 |
+
.ok_or_else(pairing_unavailable_error)?;
|
| 167 |
+
|
| 168 |
+
let response = create_client_without_request_logging()
|
| 169 |
+
.post(&self.remote_control_target.pair_status_url)
|
| 170 |
+
.timeout(REMOTE_CONTROL_PAIRING_TIMEOUT)
|
| 171 |
+
.bearer_auth(remote_control_token)
|
| 172 |
+
.json(&request)
|
| 173 |
+
.send()
|
| 174 |
+
.await
|
| 175 |
+
.map_err(|err| {
|
| 176 |
+
io::Error::other(format!(
|
| 177 |
+
"failed to check remote control pairing status at `{}`: {err}",
|
| 178 |
+
self.remote_control_target.pair_status_url
|
| 179 |
+
))
|
| 180 |
+
})?;
|
| 181 |
+
let headers = response.headers().clone();
|
| 182 |
+
let status = response.status();
|
| 183 |
+
let retry_at = retry_after_with_jitter(&headers, OffsetDateTime::now_utc());
|
| 184 |
+
let body = response.bytes().await.map_err(|err| {
|
| 185 |
+
pairing_response_error(
|
| 186 |
+
format!(
|
| 187 |
+
"failed to read remote control pairing status response from `{}`: {err}",
|
| 188 |
+
self.remote_control_target.pair_status_url
|
| 189 |
+
),
|
| 190 |
+
status,
|
| 191 |
+
retry_at,
|
| 192 |
+
ErrorKind::Other,
|
| 193 |
+
)
|
| 194 |
+
})?;
|
| 195 |
+
let body_preview = preview_remote_control_response_body(&body);
|
| 196 |
+
if !status.is_success() {
|
| 197 |
+
let error_kind = match status.as_u16() {
|
| 198 |
+
401 | 403 => ErrorKind::PermissionDenied,
|
| 199 |
+
404 | 410 => ErrorKind::InvalidInput,
|
| 200 |
+
_ => ErrorKind::Other,
|
| 201 |
+
};
|
| 202 |
+
return Err(pairing_response_error(
|
| 203 |
+
format!(
|
| 204 |
+
"remote control pairing status failed at `{}`: HTTP {status}, {}, body: {body_preview}",
|
| 205 |
+
self.remote_control_target.pair_status_url,
|
| 206 |
+
format_headers(&headers)
|
| 207 |
+
),
|
| 208 |
+
status,
|
| 209 |
+
retry_at,
|
| 210 |
+
error_kind,
|
| 211 |
+
));
|
| 212 |
+
}
|
| 213 |
+
|
| 214 |
+
let response = serde_json::from_slice::<BackendRemoteControlPairingStatusResponse>(&body)
|
| 215 |
+
.map_err(|err| {
|
| 216 |
+
io::Error::other(format!(
|
| 217 |
+
"failed to parse remote control pairing status response from `{}`: HTTP {status}, {}, body: {body_preview}, decode error: {err}",
|
| 218 |
+
self.remote_control_target.pair_status_url,
|
| 219 |
+
format_headers(&headers)
|
| 220 |
+
))
|
| 221 |
+
})?;
|
| 222 |
+
Ok(RemoteControlPairingStatusResponse {
|
| 223 |
+
claimed: response.claimed,
|
| 224 |
+
})
|
| 225 |
+
}
|
| 226 |
+
|
| 227 |
+
pub(super) fn server_token_refresh_requirement(
|
| 228 |
+
&self,
|
| 229 |
+
) -> RemoteControlServerTokenRefreshRequirement {
|
| 230 |
+
self.server_token_refresh_requirement_at(OffsetDateTime::now_utc())
|
| 231 |
+
}
|
| 232 |
+
|
| 233 |
+
pub(super) fn should_refresh_server_token(&self) -> bool {
|
| 234 |
+
self.server_token_refresh_requirement()
|
| 235 |
+
!= RemoteControlServerTokenRefreshRequirement::NotNeeded
|
| 236 |
+
}
|
| 237 |
+
|
| 238 |
+
pub(super) fn server_token_refresh_requirement_at(
|
| 239 |
+
&self,
|
| 240 |
+
now: OffsetDateTime,
|
| 241 |
+
) -> RemoteControlServerTokenRefreshRequirement {
|
| 242 |
+
let Some(expires_at) = self.remote_control_token.as_ref().and(self.expires_at) else {
|
| 243 |
+
return RemoteControlServerTokenRefreshRequirement::Required;
|
| 244 |
+
};
|
| 245 |
+
if expires_at <= now {
|
| 246 |
+
return RemoteControlServerTokenRefreshRequirement::Required;
|
| 247 |
+
}
|
| 248 |
+
if expires_at > now + time::Duration::seconds(REMOTE_CONTROL_SERVER_TOKEN_REFRESH_SKEW_SECS)
|
| 249 |
+
|| self
|
| 250 |
+
.next_refresh_at
|
| 251 |
+
.is_some_and(|next_refresh_at| next_refresh_at > now)
|
| 252 |
+
{
|
| 253 |
+
return RemoteControlServerTokenRefreshRequirement::NotNeeded;
|
| 254 |
+
}
|
| 255 |
+
RemoteControlServerTokenRefreshRequirement::Proactive
|
| 256 |
+
}
|
| 257 |
+
|
| 258 |
+
pub(super) fn clear_server_token(&mut self) {
|
| 259 |
+
self.remote_control_token = None;
|
| 260 |
+
self.expires_at = None;
|
| 261 |
+
}
|
| 262 |
+
}
|
| 263 |
+
|
| 264 |
+
fn pairing_response_error(
|
| 265 |
+
message: String,
|
| 266 |
+
status: StatusCode,
|
| 267 |
+
retry_at: Option<OffsetDateTime>,
|
| 268 |
+
fallback_kind: ErrorKind,
|
| 269 |
+
) -> io::Error {
|
| 270 |
+
if matches!(
|
| 271 |
+
status,
|
| 272 |
+
StatusCode::TOO_MANY_REQUESTS | StatusCode::SERVICE_UNAVAILABLE
|
| 273 |
+
) {
|
| 274 |
+
RemoteControlServerRequestError::io_error(
|
| 275 |
+
message,
|
| 276 |
+
Some(status),
|
| 277 |
+
retry_at,
|
| 278 |
+
/*timed_out*/ false,
|
| 279 |
+
)
|
| 280 |
+
} else {
|
| 281 |
+
io::Error::new(fallback_kind, message)
|
| 282 |
+
}
|
| 283 |
+
}
|
| 284 |
+
|
| 285 |
+
pub(super) async fn load_persisted_remote_control_enrollment(
|
| 286 |
+
state_db: Option<&StateRuntime>,
|
| 287 |
+
remote_control_target: &RemoteControlTarget,
|
| 288 |
+
account_id: &str,
|
| 289 |
+
app_server_client_name: Option<&str>,
|
| 290 |
+
) -> io::Result<Option<RemoteControlEnrollment>> {
|
| 291 |
+
let Some(state_db) = state_db else {
|
| 292 |
+
return Err(io::Error::new(
|
| 293 |
+
ErrorKind::NotFound,
|
| 294 |
+
format!(
|
| 295 |
+
"remote control enrollment cache unavailable because sqlite state db is disabled: websocket_url={}, account_id={}, app_server_client_name={:?}",
|
| 296 |
+
remote_control_target.websocket_url, account_id, app_server_client_name
|
| 297 |
+
),
|
| 298 |
+
));
|
| 299 |
+
};
|
| 300 |
+
let enrollment = match state_db
|
| 301 |
+
.get_remote_control_enrollment(
|
| 302 |
+
&remote_control_target.websocket_url,
|
| 303 |
+
account_id,
|
| 304 |
+
app_server_client_name,
|
| 305 |
+
)
|
| 306 |
+
.await
|
| 307 |
+
{
|
| 308 |
+
Ok(enrollment) => enrollment,
|
| 309 |
+
Err(err) => {
|
| 310 |
+
warn!(
|
| 311 |
+
"failed to load persisted remote control enrollment: websocket_url={}, account_id={}, app_server_client_name={:?}, err={err}",
|
| 312 |
+
remote_control_target.websocket_url, account_id, app_server_client_name
|
| 313 |
+
);
|
| 314 |
+
return Err(io::Error::other(err));
|
| 315 |
+
}
|
| 316 |
+
};
|
| 317 |
+
|
| 318 |
+
match enrollment {
|
| 319 |
+
Some(enrollment) => {
|
| 320 |
+
info!(
|
| 321 |
+
"reusing persisted remote control enrollment: websocket_url={}, account_id={}, app_server_client_name={:?}, server_id={}, environment_id={}",
|
| 322 |
+
remote_control_target.websocket_url,
|
| 323 |
+
account_id,
|
| 324 |
+
app_server_client_name,
|
| 325 |
+
enrollment.server_id,
|
| 326 |
+
enrollment.environment_id
|
| 327 |
+
);
|
| 328 |
+
Ok(Some(RemoteControlEnrollment {
|
| 329 |
+
remote_control_target: remote_control_target.clone(),
|
| 330 |
+
account_id: enrollment.account_id,
|
| 331 |
+
environment_id: enrollment.environment_id,
|
| 332 |
+
server_id: enrollment.server_id,
|
| 333 |
+
server_name: enrollment.server_name,
|
| 334 |
+
remote_control_token: None,
|
| 335 |
+
expires_at: None,
|
| 336 |
+
next_refresh_at: None,
|
| 337 |
+
}))
|
| 338 |
+
}
|
| 339 |
+
None => {
|
| 340 |
+
info!(
|
| 341 |
+
"no persisted remote control enrollment found: websocket_url={}, account_id={}, app_server_client_name={:?}",
|
| 342 |
+
remote_control_target.websocket_url, account_id, app_server_client_name
|
| 343 |
+
);
|
| 344 |
+
Ok(None)
|
| 345 |
+
}
|
| 346 |
+
}
|
| 347 |
+
}
|
| 348 |
+
|
| 349 |
+
pub(super) async fn update_persisted_remote_control_enrollment(
|
| 350 |
+
state_db: Option<&StateRuntime>,
|
| 351 |
+
remote_control_target: &RemoteControlTarget,
|
| 352 |
+
account_id: &str,
|
| 353 |
+
app_server_client_name: Option<&str>,
|
| 354 |
+
enrollment: Option<&RemoteControlEnrollment>,
|
| 355 |
+
remote_control_enabled: Option<bool>,
|
| 356 |
+
) -> io::Result<()> {
|
| 357 |
+
let Some(state_db) = state_db else {
|
| 358 |
+
return Err(io::Error::new(
|
| 359 |
+
ErrorKind::NotFound,
|
| 360 |
+
format!(
|
| 361 |
+
"remote control enrollment persistence unavailable because sqlite state db is disabled: websocket_url={}, account_id={}, app_server_client_name={:?}, has_enrollment={}",
|
| 362 |
+
remote_control_target.websocket_url,
|
| 363 |
+
account_id,
|
| 364 |
+
app_server_client_name,
|
| 365 |
+
enrollment.is_some()
|
| 366 |
+
),
|
| 367 |
+
));
|
| 368 |
+
};
|
| 369 |
+
if let &Some(enrollment) = &enrollment
|
| 370 |
+
&& enrollment.account_id != account_id
|
| 371 |
+
{
|
| 372 |
+
return Err(io::Error::other(format!(
|
| 373 |
+
"enrollment account_id does not match expected account_id `{account_id}`"
|
| 374 |
+
)));
|
| 375 |
+
}
|
| 376 |
+
|
| 377 |
+
if let Some(enrollment) = enrollment {
|
| 378 |
+
state_db
|
| 379 |
+
.upsert_remote_control_enrollment(&RemoteControlEnrollmentRecord {
|
| 380 |
+
websocket_url: remote_control_target.websocket_url.clone(),
|
| 381 |
+
account_id: account_id.to_string(),
|
| 382 |
+
app_server_client_name: app_server_client_name.map(str::to_string),
|
| 383 |
+
server_id: enrollment.server_id.clone(),
|
| 384 |
+
environment_id: enrollment.environment_id.clone(),
|
| 385 |
+
server_name: enrollment.server_name.clone(),
|
| 386 |
+
remote_control_enabled,
|
| 387 |
+
})
|
| 388 |
+
.await
|
| 389 |
+
.map_err(io::Error::other)?;
|
| 390 |
+
info!(
|
| 391 |
+
"persisted remote control enrollment: websocket_url={}, account_id={}, app_server_client_name={:?}, server_id={}, environment_id={}",
|
| 392 |
+
remote_control_target.websocket_url,
|
| 393 |
+
account_id,
|
| 394 |
+
app_server_client_name,
|
| 395 |
+
enrollment.server_id,
|
| 396 |
+
enrollment.environment_id
|
| 397 |
+
);
|
| 398 |
+
Ok(())
|
| 399 |
+
} else {
|
| 400 |
+
let rows_affected = state_db
|
| 401 |
+
.delete_remote_control_enrollment(
|
| 402 |
+
&remote_control_target.websocket_url,
|
| 403 |
+
account_id,
|
| 404 |
+
app_server_client_name,
|
| 405 |
+
)
|
| 406 |
+
.await
|
| 407 |
+
.map_err(io::Error::other)?;
|
| 408 |
+
info!(
|
| 409 |
+
"cleared persisted remote control enrollment: websocket_url={}, account_id={}, app_server_client_name={:?}, rows_affected={rows_affected}",
|
| 410 |
+
remote_control_target.websocket_url, account_id, app_server_client_name
|
| 411 |
+
);
|
| 412 |
+
Ok(())
|
| 413 |
+
}
|
| 414 |
+
}
|
| 415 |
+
|
| 416 |
+
pub(crate) fn preview_remote_control_response_body(body: &[u8]) -> String {
|
| 417 |
+
let body = String::from_utf8_lossy(body);
|
| 418 |
+
let trimmed = body.trim();
|
| 419 |
+
if trimmed.is_empty() {
|
| 420 |
+
return "<empty>".to_string();
|
| 421 |
+
}
|
| 422 |
+
let redacted = redact_remote_control_response_body(trimmed);
|
| 423 |
+
if redacted.len() <= REMOTE_CONTROL_RESPONSE_BODY_MAX_BYTES {
|
| 424 |
+
return redacted;
|
| 425 |
+
}
|
| 426 |
+
|
| 427 |
+
let mut cut = REMOTE_CONTROL_RESPONSE_BODY_MAX_BYTES;
|
| 428 |
+
while !redacted.is_char_boundary(cut) {
|
| 429 |
+
cut = cut.saturating_sub(1);
|
| 430 |
+
}
|
| 431 |
+
let mut truncated = redacted[..cut].to_string();
|
| 432 |
+
truncated.push_str("...");
|
| 433 |
+
truncated
|
| 434 |
+
}
|
| 435 |
+
|
| 436 |
+
fn redact_remote_control_response_body(body: &str) -> String {
|
| 437 |
+
let Ok(mut body_json) = serde_json::from_str::<serde_json::Value>(body) else {
|
| 438 |
+
return body.to_string();
|
| 439 |
+
};
|
| 440 |
+
let Some(body_object) = body_json.as_object_mut() else {
|
| 441 |
+
return body.to_string();
|
| 442 |
+
};
|
| 443 |
+
for sensitive_field in [
|
| 444 |
+
"remote_control_token",
|
| 445 |
+
"pairing_code",
|
| 446 |
+
"manual_pairing_code",
|
| 447 |
+
] {
|
| 448 |
+
if let Some(value) = body_object.get_mut(sensitive_field) {
|
| 449 |
+
*value = serde_json::Value::String("<redacted>".to_string());
|
| 450 |
+
}
|
| 451 |
+
}
|
| 452 |
+
body_json.to_string()
|
| 453 |
+
}
|
| 454 |
+
|
| 455 |
+
pub(crate) fn format_headers(headers: &HeaderMap) -> String {
|
| 456 |
+
let request_id_str = headers
|
| 457 |
+
.get(REQUEST_ID_HEADER)
|
| 458 |
+
.or_else(|| headers.get(OAI_REQUEST_ID_HEADER))
|
| 459 |
+
.map(|value| value.to_str().unwrap_or("<invalid utf-8>").to_owned())
|
| 460 |
+
.unwrap_or_else(|| "<none>".to_owned());
|
| 461 |
+
let cf_ray_str = headers
|
| 462 |
+
.get(CF_RAY_HEADER)
|
| 463 |
+
.map(|value| value.to_str().unwrap_or("<invalid utf-8>").to_owned())
|
| 464 |
+
.unwrap_or_else(|| "<none>".to_owned());
|
| 465 |
+
format!("request-id: {request_id_str}, cf-ray: {cf_ray_str}")
|
| 466 |
+
}
|
| 467 |
+
|
| 468 |
+
#[cfg(test)]
|
| 469 |
+
mod tests {
|
| 470 |
+
use super::*;
|
| 471 |
+
use crate::transport::remote_control::auth::RemoteControlConnectionAuth;
|
| 472 |
+
use crate::transport::remote_control::protocol::normalize_remote_control_url;
|
| 473 |
+
use crate::transport::remote_control::server_api::enroll_remote_control_server;
|
| 474 |
+
use codex_state::StateRuntime;
|
| 475 |
+
use codex_utils_absolute_path::test_support::PathExt;
|
| 476 |
+
use pretty_assertions::assert_eq;
|
| 477 |
+
use serde_json::json;
|
| 478 |
+
use std::sync::Arc;
|
| 479 |
+
use tempfile::TempDir;
|
| 480 |
+
use tokio::io::AsyncBufReadExt;
|
| 481 |
+
use tokio::io::AsyncWriteExt;
|
| 482 |
+
use tokio::io::BufReader;
|
| 483 |
+
use tokio::net::TcpListener;
|
| 484 |
+
use tokio::net::TcpStream;
|
| 485 |
+
use tokio::time::Duration;
|
| 486 |
+
use tokio::time::timeout;
|
| 487 |
+
|
| 488 |
+
async fn remote_control_state_runtime(codex_home: &TempDir) -> Arc<StateRuntime> {
|
| 489 |
+
StateRuntime::init(
|
| 490 |
+
codex_state::SqliteConfig::new_for_testing(codex_home.path().abs()),
|
| 491 |
+
"test-provider".to_string(),
|
| 492 |
+
)
|
| 493 |
+
.await
|
| 494 |
+
.expect("state runtime should initialize")
|
| 495 |
+
}
|
| 496 |
+
|
| 497 |
+
#[test]
|
| 498 |
+
fn preview_remote_control_response_body_redacts_server_token() {
|
| 499 |
+
assert_eq!(
|
| 500 |
+
serde_json::from_str::<serde_json::Value>(&preview_remote_control_response_body(
|
| 501 |
+
br#"{"server_id":"srv_e_test","remote_control_token":"secret","pairing_code":"pairing-code","manual_pairing_code":"ABCD-EFGH"}"#
|
| 502 |
+
))
|
| 503 |
+
.expect("redacted response preview should stay valid json"),
|
| 504 |
+
json!({
|
| 505 |
+
"server_id": "srv_e_test",
|
| 506 |
+
"remote_control_token": "<redacted>",
|
| 507 |
+
"pairing_code": "<redacted>",
|
| 508 |
+
"manual_pairing_code": "<redacted>",
|
| 509 |
+
})
|
| 510 |
+
);
|
| 511 |
+
}
|
| 512 |
+
|
| 513 |
+
#[tokio::test]
|
| 514 |
+
async fn persisted_remote_control_enrollment_round_trips_by_target_and_account() {
|
| 515 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 516 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 517 |
+
let first_target = normalize_remote_control_url("https://chatgpt.com/remote/control")
|
| 518 |
+
.expect("first target should parse");
|
| 519 |
+
let second_target =
|
| 520 |
+
normalize_remote_control_url("https://api.chatgpt-staging.com/other/control")
|
| 521 |
+
.expect("second target should parse");
|
| 522 |
+
let first_enrollment = RemoteControlEnrollment {
|
| 523 |
+
remote_control_target: first_target.clone(),
|
| 524 |
+
account_id: "account-a".to_string(),
|
| 525 |
+
environment_id: "env_first".to_string(),
|
| 526 |
+
server_id: "srv_e_first".to_string(),
|
| 527 |
+
server_name: "first-server".to_string(),
|
| 528 |
+
remote_control_token: None,
|
| 529 |
+
expires_at: None,
|
| 530 |
+
next_refresh_at: None,
|
| 531 |
+
};
|
| 532 |
+
let second_enrollment = RemoteControlEnrollment {
|
| 533 |
+
remote_control_target: second_target.clone(),
|
| 534 |
+
account_id: "account-a".to_string(),
|
| 535 |
+
environment_id: "env_second".to_string(),
|
| 536 |
+
server_id: "srv_e_second".to_string(),
|
| 537 |
+
server_name: "second-server".to_string(),
|
| 538 |
+
remote_control_token: None,
|
| 539 |
+
expires_at: None,
|
| 540 |
+
next_refresh_at: None,
|
| 541 |
+
};
|
| 542 |
+
|
| 543 |
+
update_persisted_remote_control_enrollment(
|
| 544 |
+
Some(state_db.as_ref()),
|
| 545 |
+
&first_target,
|
| 546 |
+
"account-a",
|
| 547 |
+
Some("desktop-client"),
|
| 548 |
+
Some(&first_enrollment),
|
| 549 |
+
/*remote_control_enabled*/ None,
|
| 550 |
+
)
|
| 551 |
+
.await
|
| 552 |
+
.expect("first enrollment should persist");
|
| 553 |
+
update_persisted_remote_control_enrollment(
|
| 554 |
+
Some(state_db.as_ref()),
|
| 555 |
+
&second_target,
|
| 556 |
+
"account-a",
|
| 557 |
+
Some("desktop-client"),
|
| 558 |
+
Some(&second_enrollment),
|
| 559 |
+
/*remote_control_enabled*/ None,
|
| 560 |
+
)
|
| 561 |
+
.await
|
| 562 |
+
.expect("second enrollment should persist");
|
| 563 |
+
|
| 564 |
+
assert_eq!(
|
| 565 |
+
load_persisted_remote_control_enrollment(
|
| 566 |
+
Some(state_db.as_ref()),
|
| 567 |
+
&first_target,
|
| 568 |
+
"account-a",
|
| 569 |
+
Some("desktop-client"),
|
| 570 |
+
)
|
| 571 |
+
.await
|
| 572 |
+
.expect("first enrollment should load"),
|
| 573 |
+
Some(first_enrollment.clone())
|
| 574 |
+
);
|
| 575 |
+
assert_eq!(
|
| 576 |
+
load_persisted_remote_control_enrollment(
|
| 577 |
+
Some(state_db.as_ref()),
|
| 578 |
+
&first_target,
|
| 579 |
+
"account-b",
|
| 580 |
+
Some("desktop-client"),
|
| 581 |
+
)
|
| 582 |
+
.await
|
| 583 |
+
.expect("missing account should load"),
|
| 584 |
+
None
|
| 585 |
+
);
|
| 586 |
+
assert_eq!(
|
| 587 |
+
load_persisted_remote_control_enrollment(
|
| 588 |
+
Some(state_db.as_ref()),
|
| 589 |
+
&second_target,
|
| 590 |
+
"account-a",
|
| 591 |
+
Some("desktop-client"),
|
| 592 |
+
)
|
| 593 |
+
.await
|
| 594 |
+
.expect("second enrollment should load"),
|
| 595 |
+
Some(second_enrollment)
|
| 596 |
+
);
|
| 597 |
+
}
|
| 598 |
+
|
| 599 |
+
#[tokio::test]
|
| 600 |
+
async fn clearing_persisted_remote_control_enrollment_removes_only_matching_entry() {
|
| 601 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 602 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 603 |
+
let first_target = normalize_remote_control_url("https://chatgpt.com/remote/control")
|
| 604 |
+
.expect("first target should parse");
|
| 605 |
+
let second_target =
|
| 606 |
+
normalize_remote_control_url("https://api.chatgpt-staging.com/other/control")
|
| 607 |
+
.expect("second target should parse");
|
| 608 |
+
let first_enrollment = RemoteControlEnrollment {
|
| 609 |
+
remote_control_target: first_target.clone(),
|
| 610 |
+
account_id: "account-a".to_string(),
|
| 611 |
+
environment_id: "env_first".to_string(),
|
| 612 |
+
server_id: "srv_e_first".to_string(),
|
| 613 |
+
server_name: "first-server".to_string(),
|
| 614 |
+
remote_control_token: None,
|
| 615 |
+
expires_at: None,
|
| 616 |
+
next_refresh_at: None,
|
| 617 |
+
};
|
| 618 |
+
let second_enrollment = RemoteControlEnrollment {
|
| 619 |
+
remote_control_target: second_target.clone(),
|
| 620 |
+
account_id: "account-a".to_string(),
|
| 621 |
+
environment_id: "env_second".to_string(),
|
| 622 |
+
server_id: "srv_e_second".to_string(),
|
| 623 |
+
server_name: "second-server".to_string(),
|
| 624 |
+
remote_control_token: None,
|
| 625 |
+
expires_at: None,
|
| 626 |
+
next_refresh_at: None,
|
| 627 |
+
};
|
| 628 |
+
|
| 629 |
+
update_persisted_remote_control_enrollment(
|
| 630 |
+
Some(state_db.as_ref()),
|
| 631 |
+
&first_target,
|
| 632 |
+
"account-a",
|
| 633 |
+
/*app_server_client_name*/ None,
|
| 634 |
+
Some(&first_enrollment),
|
| 635 |
+
/*remote_control_enabled*/ None,
|
| 636 |
+
)
|
| 637 |
+
.await
|
| 638 |
+
.expect("first enrollment should persist");
|
| 639 |
+
update_persisted_remote_control_enrollment(
|
| 640 |
+
Some(state_db.as_ref()),
|
| 641 |
+
&second_target,
|
| 642 |
+
"account-a",
|
| 643 |
+
/*app_server_client_name*/ None,
|
| 644 |
+
Some(&second_enrollment),
|
| 645 |
+
/*remote_control_enabled*/ None,
|
| 646 |
+
)
|
| 647 |
+
.await
|
| 648 |
+
.expect("second enrollment should persist");
|
| 649 |
+
|
| 650 |
+
update_persisted_remote_control_enrollment(
|
| 651 |
+
Some(state_db.as_ref()),
|
| 652 |
+
&first_target,
|
| 653 |
+
"account-a",
|
| 654 |
+
/*app_server_client_name*/ None,
|
| 655 |
+
/*enrollment*/ None,
|
| 656 |
+
/*remote_control_enabled*/ None,
|
| 657 |
+
)
|
| 658 |
+
.await
|
| 659 |
+
.expect("matching enrollment should clear");
|
| 660 |
+
|
| 661 |
+
assert_eq!(
|
| 662 |
+
load_persisted_remote_control_enrollment(
|
| 663 |
+
Some(state_db.as_ref()),
|
| 664 |
+
&first_target,
|
| 665 |
+
"account-a",
|
| 666 |
+
/*app_server_client_name*/ None,
|
| 667 |
+
)
|
| 668 |
+
.await
|
| 669 |
+
.expect("cleared enrollment should load"),
|
| 670 |
+
None
|
| 671 |
+
);
|
| 672 |
+
assert_eq!(
|
| 673 |
+
load_persisted_remote_control_enrollment(
|
| 674 |
+
Some(state_db.as_ref()),
|
| 675 |
+
&second_target,
|
| 676 |
+
"account-a",
|
| 677 |
+
/*app_server_client_name*/ None,
|
| 678 |
+
)
|
| 679 |
+
.await
|
| 680 |
+
.expect("remaining enrollment should load"),
|
| 681 |
+
Some(second_enrollment)
|
| 682 |
+
);
|
| 683 |
+
}
|
| 684 |
+
|
| 685 |
+
#[tokio::test]
|
| 686 |
+
async fn enroll_remote_control_server_parse_failure_includes_response_body() {
|
| 687 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 688 |
+
.await
|
| 689 |
+
.expect("listener should bind");
|
| 690 |
+
let remote_control_url = format!(
|
| 691 |
+
"http://127.0.0.1:{}/backend-api/",
|
| 692 |
+
listener
|
| 693 |
+
.local_addr()
|
| 694 |
+
.expect("listener should have a local addr")
|
| 695 |
+
.port()
|
| 696 |
+
);
|
| 697 |
+
let remote_control_target =
|
| 698 |
+
normalize_remote_control_url(&remote_control_url).expect("target should parse");
|
| 699 |
+
let enroll_url = remote_control_target.enroll_url.clone();
|
| 700 |
+
let response_body = json!({
|
| 701 |
+
"server_id": "srv_e_test",
|
| 702 |
+
"environment_id": "env_test",
|
| 703 |
+
});
|
| 704 |
+
let expected_body = response_body.to_string();
|
| 705 |
+
let server_task = tokio::spawn(async move {
|
| 706 |
+
let stream = accept_http_request(&listener).await;
|
| 707 |
+
respond_with_json(stream, response_body).await;
|
| 708 |
+
});
|
| 709 |
+
|
| 710 |
+
let err = enroll_remote_control_server(
|
| 711 |
+
&remote_control_target,
|
| 712 |
+
&RemoteControlConnectionAuth {
|
| 713 |
+
auth_provider: codex_model_provider::unauthenticated_auth_provider(),
|
| 714 |
+
account_id: "account_id".to_string(),
|
| 715 |
+
},
|
| 716 |
+
"11111111-1111-4111-8111-111111111111",
|
| 717 |
+
"test-server",
|
| 718 |
+
)
|
| 719 |
+
.await
|
| 720 |
+
.expect_err("invalid response should fail to parse");
|
| 721 |
+
|
| 722 |
+
server_task.await.expect("server task should succeed");
|
| 723 |
+
assert_eq!(
|
| 724 |
+
err.to_string(),
|
| 725 |
+
format!(
|
| 726 |
+
"failed to parse remote control server enrollment response from `{enroll_url}`: HTTP 200 OK, request-id: <none>, cf-ray: <none>, body: {expected_body}, decode error: missing field `remote_control_token` at line 1 column {}",
|
| 727 |
+
expected_body.len()
|
| 728 |
+
)
|
| 729 |
+
);
|
| 730 |
+
}
|
| 731 |
+
|
| 732 |
+
async fn accept_http_request(listener: &TcpListener) -> TcpStream {
|
| 733 |
+
let (stream, _) = timeout(Duration::from_secs(5), listener.accept())
|
| 734 |
+
.await
|
| 735 |
+
.expect("HTTP request should arrive in time")
|
| 736 |
+
.expect("listener accept should succeed");
|
| 737 |
+
let mut reader = BufReader::new(stream);
|
| 738 |
+
|
| 739 |
+
let mut request_line = String::new();
|
| 740 |
+
reader
|
| 741 |
+
.read_line(&mut request_line)
|
| 742 |
+
.await
|
| 743 |
+
.expect("request line should read");
|
| 744 |
+
loop {
|
| 745 |
+
let mut line = String::new();
|
| 746 |
+
reader
|
| 747 |
+
.read_line(&mut line)
|
| 748 |
+
.await
|
| 749 |
+
.expect("header line should read");
|
| 750 |
+
if line == "\r\n" {
|
| 751 |
+
break;
|
| 752 |
+
}
|
| 753 |
+
}
|
| 754 |
+
|
| 755 |
+
reader.into_inner()
|
| 756 |
+
}
|
| 757 |
+
|
| 758 |
+
async fn respond_with_json(mut stream: TcpStream, body: serde_json::Value) {
|
| 759 |
+
let body = body.to_string();
|
| 760 |
+
let response = format!(
|
| 761 |
+
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
|
| 762 |
+
body.len()
|
| 763 |
+
);
|
| 764 |
+
stream
|
| 765 |
+
.write_all(response.as_bytes())
|
| 766 |
+
.await
|
| 767 |
+
.expect("response should write");
|
| 768 |
+
stream.flush().await.expect("response should flush");
|
| 769 |
+
}
|
| 770 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/host_device.rs
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#[cfg(any(target_os = "macos", test))]
|
| 2 |
+
use serde::Deserialize;
|
| 3 |
+
|
| 4 |
+
pub(super) const REMOTE_CONTROL_HOST_DEVICE_KIND_HEADER: &str = "x-codex-host-device-kind";
|
| 5 |
+
#[cfg(any(target_os = "macos", test))]
|
| 6 |
+
const MAC_MINI_HOST_DEVICE_KIND: &str = "mac_mini";
|
| 7 |
+
|
| 8 |
+
#[cfg(any(target_os = "macos", test))]
|
| 9 |
+
#[derive(Deserialize)]
|
| 10 |
+
struct MacHardwareProfile {
|
| 11 |
+
#[serde(rename = "SPHardwareDataType")]
|
| 12 |
+
hardware: Vec<MacHardware>,
|
| 13 |
+
}
|
| 14 |
+
|
| 15 |
+
#[cfg(any(target_os = "macos", test))]
|
| 16 |
+
#[derive(Deserialize)]
|
| 17 |
+
struct MacHardware {
|
| 18 |
+
machine_name: String,
|
| 19 |
+
}
|
| 20 |
+
|
| 21 |
+
#[cfg(any(target_os = "macos", test))]
|
| 22 |
+
fn host_device_kind_from_profile(profile: &[u8]) -> serde_json::Result<Option<&'static str>> {
|
| 23 |
+
let profile: MacHardwareProfile = serde_json::from_slice(profile)?;
|
| 24 |
+
Ok(profile
|
| 25 |
+
.hardware
|
| 26 |
+
.first()
|
| 27 |
+
.is_some_and(|hardware| hardware.machine_name == "Mac mini")
|
| 28 |
+
.then_some(MAC_MINI_HOST_DEVICE_KIND))
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
#[cfg(target_os = "macos")]
|
| 32 |
+
pub(super) async fn host_device_kind() -> Option<&'static str> {
|
| 33 |
+
use std::process::Stdio;
|
| 34 |
+
use std::time::Duration;
|
| 35 |
+
use tokio::process::Command;
|
| 36 |
+
use tokio::sync::OnceCell;
|
| 37 |
+
|
| 38 |
+
static HOST_DEVICE_KIND: OnceCell<Option<&'static str>> = OnceCell::const_new();
|
| 39 |
+
|
| 40 |
+
HOST_DEVICE_KIND
|
| 41 |
+
.get_or_try_init(|| async {
|
| 42 |
+
let output = tokio::time::timeout(
|
| 43 |
+
Duration::from_secs(2),
|
| 44 |
+
Command::new("/usr/sbin/system_profiler")
|
| 45 |
+
.args(["-detailLevel", "mini", "SPHardwareDataType", "-json"])
|
| 46 |
+
.stdin(Stdio::null())
|
| 47 |
+
.stderr(Stdio::null())
|
| 48 |
+
.kill_on_drop(true)
|
| 49 |
+
.output(),
|
| 50 |
+
)
|
| 51 |
+
.await
|
| 52 |
+
.map_err(|_| ())?
|
| 53 |
+
.map_err(|_| ())?;
|
| 54 |
+
|
| 55 |
+
if !output.status.success() {
|
| 56 |
+
return Err(());
|
| 57 |
+
}
|
| 58 |
+
|
| 59 |
+
host_device_kind_from_profile(&output.stdout).map_err(|_| ())
|
| 60 |
+
})
|
| 61 |
+
.await
|
| 62 |
+
.ok()
|
| 63 |
+
.copied()
|
| 64 |
+
.flatten()
|
| 65 |
+
}
|
| 66 |
+
|
| 67 |
+
#[cfg(not(target_os = "macos"))]
|
| 68 |
+
pub(super) async fn host_device_kind() -> Option<&'static str> {
|
| 69 |
+
None
|
| 70 |
+
}
|
| 71 |
+
|
| 72 |
+
#[cfg(test)]
|
| 73 |
+
#[path = "host_device_tests.rs"]
|
| 74 |
+
mod tests;
|
codex-rs/app-server-transport/src/transport/remote_control/host_device_tests.rs
ADDED
|
@@ -0,0 +1,38 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::host_device_kind_from_profile;
|
| 2 |
+
use pretty_assertions::assert_eq;
|
| 3 |
+
|
| 4 |
+
#[test]
|
| 5 |
+
fn recognizes_only_the_exact_mac_mini_hardware_name() {
|
| 6 |
+
let mac_mini_profile = br#"{"SPHardwareDataType":[{"machine_name":"Mac mini"}]}"#;
|
| 7 |
+
let macbook_profile = br#"{"SPHardwareDataType":[{"machine_name":"MacBook Pro"}]}"#;
|
| 8 |
+
let misleading_profile = br#"{"SPHardwareDataType":[{"machine_name":"Not a Mac mini"}]}"#;
|
| 9 |
+
|
| 10 |
+
assert_eq!(
|
| 11 |
+
host_device_kind_from_profile(mac_mini_profile).expect("valid Mac mini hardware profile"),
|
| 12 |
+
Some("mac_mini")
|
| 13 |
+
);
|
| 14 |
+
assert_eq!(
|
| 15 |
+
host_device_kind_from_profile(macbook_profile).expect("valid MacBook hardware profile"),
|
| 16 |
+
None
|
| 17 |
+
);
|
| 18 |
+
assert_eq!(
|
| 19 |
+
host_device_kind_from_profile(misleading_profile)
|
| 20 |
+
.expect("valid non-Mac-mini hardware profile"),
|
| 21 |
+
None
|
| 22 |
+
);
|
| 23 |
+
}
|
| 24 |
+
|
| 25 |
+
#[test]
|
| 26 |
+
fn ignores_missing_hardware_profiles() {
|
| 27 |
+
assert_eq!(
|
| 28 |
+
host_device_kind_from_profile(br#"{"SPHardwareDataType":[]}"#)
|
| 29 |
+
.expect("valid empty hardware profile"),
|
| 30 |
+
None
|
| 31 |
+
);
|
| 32 |
+
}
|
| 33 |
+
|
| 34 |
+
#[test]
|
| 35 |
+
fn rejects_malformed_hardware_profiles_so_they_remain_retryable() {
|
| 36 |
+
assert!(host_device_kind_from_profile(br#"{"machine_name":"Mac mini"}"#).is_err());
|
| 37 |
+
assert!(host_device_kind_from_profile(b"not json").is_err());
|
| 38 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/mod.rs
ADDED
|
@@ -0,0 +1,1002 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
mod auth;
|
| 2 |
+
mod client_tracker;
|
| 3 |
+
mod clients;
|
| 4 |
+
mod controller;
|
| 5 |
+
mod persistence;
|
| 6 |
+
pub use controller::RemoteControlHandle;
|
| 7 |
+
pub use controller::start_remote_control;
|
| 8 |
+
mod desired_state;
|
| 9 |
+
mod enroll;
|
| 10 |
+
mod host_device;
|
| 11 |
+
mod protocol;
|
| 12 |
+
mod segment;
|
| 13 |
+
mod server_api;
|
| 14 |
+
mod websocket;
|
| 15 |
+
|
| 16 |
+
use self::auth::load_remote_control_auth;
|
| 17 |
+
use self::auth::recover_remote_control_auth;
|
| 18 |
+
use self::desired_state::RemoteControlDesiredState;
|
| 19 |
+
use self::enroll::RemoteControlEnrollment;
|
| 20 |
+
use self::enroll::load_persisted_remote_control_enrollment;
|
| 21 |
+
use self::persistence::RemoteControlPersistence;
|
| 22 |
+
use self::server_api::enroll_remote_control_server;
|
| 23 |
+
use self::server_api::refresh_remote_control_server;
|
| 24 |
+
use crate::transport::remote_control::websocket::RemoteControlChannels;
|
| 25 |
+
use crate::transport::remote_control::websocket::RemoteControlStatusPublisher;
|
| 26 |
+
use crate::transport::remote_control::websocket::RemoteControlWebsocket;
|
| 27 |
+
|
| 28 |
+
pub use self::protocol::ClientId;
|
| 29 |
+
use self::protocol::RemoteControlPairingStatusCode;
|
| 30 |
+
use self::protocol::ServerEvent;
|
| 31 |
+
use self::protocol::StreamId;
|
| 32 |
+
use self::protocol::normalize_remote_control_url;
|
| 33 |
+
use super::CHANNEL_CAPACITY;
|
| 34 |
+
use super::TransportEvent;
|
| 35 |
+
use super::next_connection_id;
|
| 36 |
+
use codex_app_server_protocol::RemoteControlClientsListParams;
|
| 37 |
+
use codex_app_server_protocol::RemoteControlClientsListResponse;
|
| 38 |
+
use codex_app_server_protocol::RemoteControlClientsRevokeParams;
|
| 39 |
+
use codex_app_server_protocol::RemoteControlClientsRevokeResponse;
|
| 40 |
+
use codex_app_server_protocol::RemoteControlConnectionStatus;
|
| 41 |
+
use codex_app_server_protocol::RemoteControlPairingStartParams;
|
| 42 |
+
use codex_app_server_protocol::RemoteControlPairingStartResponse;
|
| 43 |
+
use codex_app_server_protocol::RemoteControlPairingStatusParams;
|
| 44 |
+
use codex_app_server_protocol::RemoteControlPairingStatusResponse;
|
| 45 |
+
use codex_app_server_protocol::RemoteControlStatusChangedNotification;
|
| 46 |
+
use codex_login::AuthManager;
|
| 47 |
+
use codex_state::StateRuntime;
|
| 48 |
+
use gethostname::gethostname;
|
| 49 |
+
use std::error::Error;
|
| 50 |
+
use std::fmt;
|
| 51 |
+
use std::io;
|
| 52 |
+
use std::ops::Deref;
|
| 53 |
+
use std::ops::DerefMut;
|
| 54 |
+
use std::sync::Arc;
|
| 55 |
+
use std::sync::Mutex as StdMutex;
|
| 56 |
+
use tokio::sync::Semaphore;
|
| 57 |
+
use tokio::sync::SemaphorePermit;
|
| 58 |
+
use tokio::sync::mpsc;
|
| 59 |
+
use tokio::sync::oneshot;
|
| 60 |
+
use tokio::sync::watch;
|
| 61 |
+
use tokio::task::JoinHandle;
|
| 62 |
+
use tokio_util::sync::CancellationToken;
|
| 63 |
+
use tracing::info;
|
| 64 |
+
use tracing::warn;
|
| 65 |
+
|
| 66 |
+
pub struct RemoteControlStartConfig {
|
| 67 |
+
pub remote_control_url: String,
|
| 68 |
+
pub installation_id: String,
|
| 69 |
+
pub policy: RemoteControlPolicy,
|
| 70 |
+
}
|
| 71 |
+
|
| 72 |
+
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
|
| 73 |
+
pub enum RemoteControlPolicy {
|
| 74 |
+
#[default]
|
| 75 |
+
Allowed,
|
| 76 |
+
DisabledByRequirements,
|
| 77 |
+
}
|
| 78 |
+
|
| 79 |
+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
| 80 |
+
pub enum RemoteControlStartupMode {
|
| 81 |
+
ResolvePersisted,
|
| 82 |
+
DisabledEphemeral,
|
| 83 |
+
EnabledEphemeral,
|
| 84 |
+
}
|
| 85 |
+
|
| 86 |
+
/// Internal marker used by the daemon to disable remote control without requiring a new CLI flag.
|
| 87 |
+
pub const REMOTE_CONTROL_DISABLED_ENV_VAR: &str =
|
| 88 |
+
"CODEX_INTERNAL_APP_SERVER_REMOTE_CONTROL_DISABLED";
|
| 89 |
+
|
| 90 |
+
/// Reads and removes the daemon's internal disabled-start marker before worker threads start.
|
| 91 |
+
pub fn take_remote_control_disabled_env() -> bool {
|
| 92 |
+
let disabled =
|
| 93 |
+
std::env::var_os(REMOTE_CONTROL_DISABLED_ENV_VAR).is_some_and(|value| value == "1");
|
| 94 |
+
// SAFETY: app-server calls this synchronously at process startup, before spawning threads.
|
| 95 |
+
unsafe { std::env::remove_var(REMOTE_CONTROL_DISABLED_ENV_VAR) };
|
| 96 |
+
disabled
|
| 97 |
+
}
|
| 98 |
+
|
| 99 |
+
pub(super) struct QueuedServerEnvelope {
|
| 100 |
+
pub(super) event: ServerEvent,
|
| 101 |
+
pub(super) client_id: ClientId,
|
| 102 |
+
pub(super) stream_id: StreamId,
|
| 103 |
+
pub(super) write_complete_tx: Option<oneshot::Sender<()>>,
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
#[derive(Clone)]
|
| 107 |
+
struct RemoteControlSession {
|
| 108 |
+
policy: RemoteControlPolicy,
|
| 109 |
+
shutdown_token: CancellationToken,
|
| 110 |
+
desired_state_tx: Arc<watch::Sender<RemoteControlDesiredState>>,
|
| 111 |
+
desired_state_rpc_lock: Arc<Semaphore>,
|
| 112 |
+
persistence: RemoteControlPersistence,
|
| 113 |
+
status_tx: Arc<watch::Sender<RemoteControlStatusChangedNotification>>,
|
| 114 |
+
state_db: Option<Arc<StateRuntime>>,
|
| 115 |
+
remote_control_url: String,
|
| 116 |
+
current_enrollment: CurrentRemoteControlEnrollment,
|
| 117 |
+
pairing_persistence_key: RemoteControlPairingPersistenceKey,
|
| 118 |
+
pairing_persistence_key_required: bool,
|
| 119 |
+
auth_manager: auth::RemoteControlAuth,
|
| 120 |
+
}
|
| 121 |
+
|
| 122 |
+
// Pairing and websocket connect share one selected server so they cannot enroll or replace
|
| 123 |
+
// different persisted rows while either path is awaiting backend I/O.
|
| 124 |
+
type CurrentRemoteControlEnrollment = Arc<RemoteControlEnrollmentState>;
|
| 125 |
+
type RemoteControlPairingPersistenceKey = watch::Sender<Option<String>>;
|
| 126 |
+
|
| 127 |
+
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
| 128 |
+
enum RemoteControlEnrollmentSelection {
|
| 129 |
+
ReuseOrCreate,
|
| 130 |
+
ReplaceExisting,
|
| 131 |
+
}
|
| 132 |
+
|
| 133 |
+
struct RemoteControlEnrollmentState {
|
| 134 |
+
enrollment: StdMutex<Option<RemoteControlEnrollment>>,
|
| 135 |
+
// Keep an observed server deadline across enrollment replacement and enable transitions.
|
| 136 |
+
retry_at: StdMutex<Option<time::OffsetDateTime>>,
|
| 137 |
+
lock: Semaphore,
|
| 138 |
+
}
|
| 139 |
+
|
| 140 |
+
impl RemoteControlEnrollmentState {
|
| 141 |
+
fn new(enrollment: Option<RemoteControlEnrollment>) -> Self {
|
| 142 |
+
Self {
|
| 143 |
+
enrollment: StdMutex::new(enrollment),
|
| 144 |
+
retry_at: StdMutex::new(None),
|
| 145 |
+
lock: Semaphore::new(1),
|
| 146 |
+
}
|
| 147 |
+
}
|
| 148 |
+
|
| 149 |
+
async fn lock_for_request(&self) -> io::Result<RemoteControlEnrollmentLease<'_>> {
|
| 150 |
+
self.check_retry_after()?;
|
| 151 |
+
let lease = self.lock().await;
|
| 152 |
+
self.check_retry_after()?;
|
| 153 |
+
Ok(lease)
|
| 154 |
+
}
|
| 155 |
+
|
| 156 |
+
// Check immediately before admitting network work. Requests admitted before
|
| 157 |
+
// an overload response publishes its deadline may finish concurrently.
|
| 158 |
+
fn check_retry_after(&self) -> io::Result<()> {
|
| 159 |
+
let retry_at = *self
|
| 160 |
+
.retry_at
|
| 161 |
+
.lock()
|
| 162 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 163 |
+
if let Some(retry_at) = retry_at
|
| 164 |
+
&& retry_at > time::OffsetDateTime::now_utc()
|
| 165 |
+
{
|
| 166 |
+
return Err(server_api::RemoteControlServerRequestError::retry_deferred(
|
| 167 |
+
retry_at,
|
| 168 |
+
));
|
| 169 |
+
}
|
| 170 |
+
Ok(())
|
| 171 |
+
}
|
| 172 |
+
|
| 173 |
+
fn record_retry_after<T>(&self, result: io::Result<T>) -> io::Result<T> {
|
| 174 |
+
if let Err(err) = &result
|
| 175 |
+
&& let Some(retry_at) = server_api::remote_control_retry_at(err)
|
| 176 |
+
{
|
| 177 |
+
let mut current_retry_at = self
|
| 178 |
+
.retry_at
|
| 179 |
+
.lock()
|
| 180 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
| 181 |
+
*current_retry_at = Some(current_retry_at.map_or(retry_at, |at| at.max(retry_at)));
|
| 182 |
+
}
|
| 183 |
+
result
|
| 184 |
+
}
|
| 185 |
+
|
| 186 |
+
async fn lock(&self) -> RemoteControlEnrollmentLease<'_> {
|
| 187 |
+
let permit = match self.lock.acquire().await {
|
| 188 |
+
Ok(permit) => permit,
|
| 189 |
+
Err(_) => unreachable!("remote control enrollment lock should stay open"),
|
| 190 |
+
};
|
| 191 |
+
RemoteControlEnrollmentLease {
|
| 192 |
+
state: self,
|
| 193 |
+
enrollment: self.snapshot(),
|
| 194 |
+
_permit: permit,
|
| 195 |
+
}
|
| 196 |
+
}
|
| 197 |
+
|
| 198 |
+
fn snapshot(&self) -> Option<RemoteControlEnrollment> {
|
| 199 |
+
self.enrollment
|
| 200 |
+
.lock()
|
| 201 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner)
|
| 202 |
+
.clone()
|
| 203 |
+
}
|
| 204 |
+
}
|
| 205 |
+
|
| 206 |
+
struct RemoteControlEnrollmentLease<'a> {
|
| 207 |
+
state: &'a RemoteControlEnrollmentState,
|
| 208 |
+
enrollment: Option<RemoteControlEnrollment>,
|
| 209 |
+
_permit: SemaphorePermit<'a>,
|
| 210 |
+
}
|
| 211 |
+
|
| 212 |
+
impl Deref for RemoteControlEnrollmentLease<'_> {
|
| 213 |
+
type Target = Option<RemoteControlEnrollment>;
|
| 214 |
+
|
| 215 |
+
fn deref(&self) -> &Self::Target {
|
| 216 |
+
&self.enrollment
|
| 217 |
+
}
|
| 218 |
+
}
|
| 219 |
+
|
| 220 |
+
impl DerefMut for RemoteControlEnrollmentLease<'_> {
|
| 221 |
+
fn deref_mut(&mut self) -> &mut Self::Target {
|
| 222 |
+
&mut self.enrollment
|
| 223 |
+
}
|
| 224 |
+
}
|
| 225 |
+
|
| 226 |
+
impl Drop for RemoteControlEnrollmentLease<'_> {
|
| 227 |
+
fn drop(&mut self) {
|
| 228 |
+
*self
|
| 229 |
+
.state
|
| 230 |
+
.enrollment
|
| 231 |
+
.lock()
|
| 232 |
+
.unwrap_or_else(std::sync::PoisonError::into_inner) = self.enrollment.take();
|
| 233 |
+
}
|
| 234 |
+
}
|
| 235 |
+
|
| 236 |
+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
| 237 |
+
pub struct RemoteControlUnavailable;
|
| 238 |
+
|
| 239 |
+
impl fmt::Display for RemoteControlUnavailable {
|
| 240 |
+
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 241 |
+
write!(
|
| 242 |
+
f,
|
| 243 |
+
"remote control cannot be enabled because sqlite state db is unavailable"
|
| 244 |
+
)
|
| 245 |
+
}
|
| 246 |
+
}
|
| 247 |
+
|
| 248 |
+
impl Error for RemoteControlUnavailable {}
|
| 249 |
+
|
| 250 |
+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
| 251 |
+
pub struct RemoteControlDisabledByRequirements;
|
| 252 |
+
|
| 253 |
+
impl fmt::Display for RemoteControlDisabledByRequirements {
|
| 254 |
+
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 255 |
+
write!(f, "remote control is disabled by managed requirements")
|
| 256 |
+
}
|
| 257 |
+
}
|
| 258 |
+
|
| 259 |
+
impl Error for RemoteControlDisabledByRequirements {}
|
| 260 |
+
|
| 261 |
+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
| 262 |
+
pub enum RemoteControlEnableError {
|
| 263 |
+
Unavailable(RemoteControlUnavailable),
|
| 264 |
+
DisabledByRequirements(RemoteControlDisabledByRequirements),
|
| 265 |
+
AuthenticationChanged,
|
| 266 |
+
}
|
| 267 |
+
|
| 268 |
+
impl fmt::Display for RemoteControlEnableError {
|
| 269 |
+
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 270 |
+
match self {
|
| 271 |
+
Self::Unavailable(err) => err.fmt(f),
|
| 272 |
+
Self::DisabledByRequirements(err) => err.fmt(f),
|
| 273 |
+
Self::AuthenticationChanged => f.write_str("remote control authentication changed"),
|
| 274 |
+
}
|
| 275 |
+
}
|
| 276 |
+
}
|
| 277 |
+
|
| 278 |
+
impl Error for RemoteControlEnableError {}
|
| 279 |
+
|
| 280 |
+
impl RemoteControlSession {
|
| 281 |
+
pub fn ensure_remote_control_allowed(&self) -> Result<(), RemoteControlDisabledByRequirements> {
|
| 282 |
+
match self.policy {
|
| 283 |
+
RemoteControlPolicy::Allowed => Ok(()),
|
| 284 |
+
RemoteControlPolicy::DisabledByRequirements => Err(RemoteControlDisabledByRequirements),
|
| 285 |
+
}
|
| 286 |
+
}
|
| 287 |
+
|
| 288 |
+
fn ensure_remote_control_allowed_io(&self) -> io::Result<()> {
|
| 289 |
+
self.ensure_remote_control_allowed()
|
| 290 |
+
.map_err(|err| io::Error::new(io::ErrorKind::PermissionDenied, err))
|
| 291 |
+
}
|
| 292 |
+
|
| 293 |
+
pub fn enable_ephemeral(
|
| 294 |
+
&self,
|
| 295 |
+
) -> Result<RemoteControlStatusChangedNotification, RemoteControlEnableError> {
|
| 296 |
+
self.enable_with_preference(/*persistence_preference*/ None)
|
| 297 |
+
}
|
| 298 |
+
|
| 299 |
+
fn enable_with_preference(
|
| 300 |
+
&self,
|
| 301 |
+
persistence_preference: Option<bool>,
|
| 302 |
+
) -> Result<RemoteControlStatusChangedNotification, RemoteControlEnableError> {
|
| 303 |
+
self.ensure_remote_control_allowed()
|
| 304 |
+
.map_err(RemoteControlEnableError::DisabledByRequirements)?;
|
| 305 |
+
if self.state_db.is_none() {
|
| 306 |
+
warn!("remote control cannot be enabled because sqlite state db is unavailable");
|
| 307 |
+
return Err(RemoteControlEnableError::Unavailable(
|
| 308 |
+
RemoteControlUnavailable,
|
| 309 |
+
));
|
| 310 |
+
}
|
| 311 |
+
|
| 312 |
+
let mut effective_persistence_preference = persistence_preference;
|
| 313 |
+
self.auth_manager
|
| 314 |
+
.ensure_current()
|
| 315 |
+
.map_err(|_| RemoteControlEnableError::AuthenticationChanged)?;
|
| 316 |
+
let desired_state_changed = self.desired_state_tx.send_if_modified(|state| {
|
| 317 |
+
if effective_persistence_preference.is_none()
|
| 318 |
+
&& matches!(
|
| 319 |
+
*state,
|
| 320 |
+
RemoteControlDesiredState::Enabled {
|
| 321 |
+
persistence_preference: Some(true)
|
| 322 |
+
}
|
| 323 |
+
)
|
| 324 |
+
{
|
| 325 |
+
effective_persistence_preference = Some(true);
|
| 326 |
+
}
|
| 327 |
+
let next_state = RemoteControlDesiredState::Enabled {
|
| 328 |
+
persistence_preference: effective_persistence_preference,
|
| 329 |
+
};
|
| 330 |
+
let changed = *state != next_state;
|
| 331 |
+
*state = next_state;
|
| 332 |
+
changed
|
| 333 |
+
});
|
| 334 |
+
|
| 335 |
+
let status = self.status();
|
| 336 |
+
info!(
|
| 337 |
+
desired_state_changed,
|
| 338 |
+
?effective_persistence_preference,
|
| 339 |
+
current_status = ?status.status,
|
| 340 |
+
environment_id = ?status.environment_id,
|
| 341 |
+
installation_id = %status.installation_id,
|
| 342 |
+
server_name = %status.server_name,
|
| 343 |
+
"remote control enable requested"
|
| 344 |
+
);
|
| 345 |
+
if matches!(
|
| 346 |
+
status.status,
|
| 347 |
+
RemoteControlConnectionStatus::Connected | RemoteControlConnectionStatus::Connecting
|
| 348 |
+
) {
|
| 349 |
+
return Ok(status);
|
| 350 |
+
}
|
| 351 |
+
|
| 352 |
+
Ok(self.publish_status(RemoteControlConnectionStatus::Connecting))
|
| 353 |
+
}
|
| 354 |
+
|
| 355 |
+
pub async fn disable(
|
| 356 |
+
&self,
|
| 357 |
+
app_server_client_name: Option<&str>,
|
| 358 |
+
) -> io::Result<RemoteControlStatusChangedNotification> {
|
| 359 |
+
self.ensure_remote_control_allowed_io()?;
|
| 360 |
+
let _transition = self
|
| 361 |
+
.desired_state_rpc_lock
|
| 362 |
+
.acquire()
|
| 363 |
+
.await
|
| 364 |
+
.unwrap_or_else(|_| unreachable!());
|
| 365 |
+
self.persist_preference(
|
| 366 |
+
app_server_client_name,
|
| 367 |
+
/*remote_control_enabled*/ false,
|
| 368 |
+
)
|
| 369 |
+
.await?;
|
| 370 |
+
Ok(self.transition_disabled())
|
| 371 |
+
}
|
| 372 |
+
|
| 373 |
+
pub async fn disable_ephemeral(&self) -> RemoteControlStatusChangedNotification {
|
| 374 |
+
let _transition = self
|
| 375 |
+
.desired_state_rpc_lock
|
| 376 |
+
.acquire()
|
| 377 |
+
.await
|
| 378 |
+
.unwrap_or_else(|_| unreachable!());
|
| 379 |
+
let _persistence = self.persistence.lock().await;
|
| 380 |
+
self.transition_disabled()
|
| 381 |
+
}
|
| 382 |
+
|
| 383 |
+
fn transition_disabled(&self) -> RemoteControlStatusChangedNotification {
|
| 384 |
+
let desired_state_changed = self.desired_state_tx.send_if_modified(|state| {
|
| 385 |
+
let changed = *state != RemoteControlDesiredState::Disabled;
|
| 386 |
+
*state = RemoteControlDesiredState::Disabled;
|
| 387 |
+
changed
|
| 388 |
+
});
|
| 389 |
+
let status = self.status();
|
| 390 |
+
info!(
|
| 391 |
+
desired_state_changed,
|
| 392 |
+
current_status = ?status.status,
|
| 393 |
+
environment_id = ?status.environment_id,
|
| 394 |
+
installation_id = %status.installation_id,
|
| 395 |
+
server_name = %status.server_name,
|
| 396 |
+
"remote control disable requested"
|
| 397 |
+
);
|
| 398 |
+
self.publish_status(RemoteControlConnectionStatus::Disabled)
|
| 399 |
+
}
|
| 400 |
+
|
| 401 |
+
async fn persist_preference(
|
| 402 |
+
&self,
|
| 403 |
+
app_server_client_name: Option<&str>,
|
| 404 |
+
remote_control_enabled: bool,
|
| 405 |
+
) -> io::Result<()> {
|
| 406 |
+
let state_db = self
|
| 407 |
+
.state_db
|
| 408 |
+
.as_deref()
|
| 409 |
+
.ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, RemoteControlUnavailable))?;
|
| 410 |
+
let auth = load_remote_control_auth(&self.auth_manager).await?;
|
| 411 |
+
let remote_control_target = normalize_remote_control_url(&self.remote_control_url)?;
|
| 412 |
+
let app_server_client_name = self.pairing_persistence_key(app_server_client_name)?;
|
| 413 |
+
self.set_preference(
|
| 414 |
+
state_db,
|
| 415 |
+
&remote_control_target,
|
| 416 |
+
&auth.account_id,
|
| 417 |
+
app_server_client_name.as_deref(),
|
| 418 |
+
remote_control_enabled,
|
| 419 |
+
/*fallback_enrollment*/ None,
|
| 420 |
+
)
|
| 421 |
+
.await?;
|
| 422 |
+
Ok(())
|
| 423 |
+
}
|
| 424 |
+
|
| 425 |
+
pub fn status(&self) -> RemoteControlStatusChangedNotification {
|
| 426 |
+
self.status_tx.borrow().clone()
|
| 427 |
+
}
|
| 428 |
+
|
| 429 |
+
pub fn status_receiver(&self) -> watch::Receiver<RemoteControlStatusChangedNotification> {
|
| 430 |
+
self.status_tx.subscribe()
|
| 431 |
+
}
|
| 432 |
+
|
| 433 |
+
pub async fn start_pairing(
|
| 434 |
+
&self,
|
| 435 |
+
params: RemoteControlPairingStartParams,
|
| 436 |
+
app_server_client_name: Option<&str>,
|
| 437 |
+
) -> io::Result<RemoteControlPairingStartResponse> {
|
| 438 |
+
self.ensure_remote_control_allowed_io()?;
|
| 439 |
+
if !self.desired_state_tx.borrow().is_enabled() {
|
| 440 |
+
return Err(Self::pairing_disabled_error());
|
| 441 |
+
}
|
| 442 |
+
let mut current_enrollment = self.current_enrollment.lock_for_request().await?;
|
| 443 |
+
let mut auth = load_remote_control_auth(&self.auth_manager)
|
| 444 |
+
.await
|
| 445 |
+
.map_err(|_| pairing_unavailable_error())?;
|
| 446 |
+
let status = self.status();
|
| 447 |
+
let installation_id = status.installation_id;
|
| 448 |
+
let app_server_client_name = self.pairing_persistence_key(app_server_client_name)?;
|
| 449 |
+
let app_server_client_name = app_server_client_name.as_deref();
|
| 450 |
+
let mut enrollment = self
|
| 451 |
+
.load_or_enroll_pairing_server(
|
| 452 |
+
&mut current_enrollment,
|
| 453 |
+
&mut auth,
|
| 454 |
+
&installation_id,
|
| 455 |
+
&status.server_name,
|
| 456 |
+
app_server_client_name,
|
| 457 |
+
RemoteControlEnrollmentSelection::ReuseOrCreate,
|
| 458 |
+
)
|
| 459 |
+
.await?;
|
| 460 |
+
if enrollment.should_refresh_server_token() {
|
| 461 |
+
let refresh_result = refresh_pairing_enrollment(
|
| 462 |
+
&mut current_enrollment,
|
| 463 |
+
&self.auth_manager,
|
| 464 |
+
&mut auth,
|
| 465 |
+
&installation_id,
|
| 466 |
+
&mut enrollment,
|
| 467 |
+
)
|
| 468 |
+
.await;
|
| 469 |
+
if refresh_result
|
| 470 |
+
.as_ref()
|
| 471 |
+
.is_err_and(|err| err.kind() == io::ErrorKind::NotFound)
|
| 472 |
+
{
|
| 473 |
+
enrollment = self
|
| 474 |
+
.load_or_enroll_pairing_server(
|
| 475 |
+
&mut current_enrollment,
|
| 476 |
+
&mut auth,
|
| 477 |
+
&installation_id,
|
| 478 |
+
&status.server_name,
|
| 479 |
+
app_server_client_name,
|
| 480 |
+
RemoteControlEnrollmentSelection::ReplaceExisting,
|
| 481 |
+
)
|
| 482 |
+
.await?;
|
| 483 |
+
} else {
|
| 484 |
+
refresh_result?;
|
| 485 |
+
}
|
| 486 |
+
}
|
| 487 |
+
let pairing_request = || protocol::StartRemoteControlPairingRequest {
|
| 488 |
+
manual_code: params.manual_code,
|
| 489 |
+
};
|
| 490 |
+
self.current_enrollment.check_retry_after()?;
|
| 491 |
+
let pairing_response = match enrollment.start_pairing(pairing_request()).await {
|
| 492 |
+
Err(err) if err.kind() == io::ErrorKind::PermissionDenied => {
|
| 493 |
+
clear_pairing_server_token(&mut current_enrollment, &mut enrollment)?;
|
| 494 |
+
refresh_pairing_enrollment(
|
| 495 |
+
&mut current_enrollment,
|
| 496 |
+
&self.auth_manager,
|
| 497 |
+
&mut auth,
|
| 498 |
+
&installation_id,
|
| 499 |
+
&mut enrollment,
|
| 500 |
+
)
|
| 501 |
+
.await?;
|
| 502 |
+
self.current_enrollment.check_retry_after()?;
|
| 503 |
+
enrollment.start_pairing(pairing_request()).await
|
| 504 |
+
}
|
| 505 |
+
Err(err) if err.kind() == io::ErrorKind::NotFound => {
|
| 506 |
+
enrollment = self
|
| 507 |
+
.load_or_enroll_pairing_server(
|
| 508 |
+
&mut current_enrollment,
|
| 509 |
+
&mut auth,
|
| 510 |
+
&installation_id,
|
| 511 |
+
&status.server_name,
|
| 512 |
+
app_server_client_name,
|
| 513 |
+
RemoteControlEnrollmentSelection::ReplaceExisting,
|
| 514 |
+
)
|
| 515 |
+
.await?;
|
| 516 |
+
self.current_enrollment.check_retry_after()?;
|
| 517 |
+
enrollment.start_pairing(pairing_request()).await
|
| 518 |
+
}
|
| 519 |
+
pairing_response => pairing_response,
|
| 520 |
+
};
|
| 521 |
+
if let Err(err) = &pairing_response {
|
| 522 |
+
if server_api::remote_control_retry_at(err).is_some() {
|
| 523 |
+
return self.current_enrollment.record_retry_after(pairing_response);
|
| 524 |
+
}
|
| 525 |
+
match err.kind() {
|
| 526 |
+
io::ErrorKind::NotFound => {
|
| 527 |
+
self.load_or_enroll_pairing_server(
|
| 528 |
+
&mut current_enrollment,
|
| 529 |
+
&mut auth,
|
| 530 |
+
&installation_id,
|
| 531 |
+
&status.server_name,
|
| 532 |
+
app_server_client_name,
|
| 533 |
+
RemoteControlEnrollmentSelection::ReplaceExisting,
|
| 534 |
+
)
|
| 535 |
+
.await?;
|
| 536 |
+
return Err(pairing_unavailable_error());
|
| 537 |
+
}
|
| 538 |
+
io::ErrorKind::PermissionDenied => {
|
| 539 |
+
clear_pairing_server_token(&mut current_enrollment, &mut enrollment)?;
|
| 540 |
+
return Err(pairing_unavailable_error());
|
| 541 |
+
}
|
| 542 |
+
_ => {}
|
| 543 |
+
}
|
| 544 |
+
}
|
| 545 |
+
let current_auth = load_remote_control_auth(&self.auth_manager)
|
| 546 |
+
.await
|
| 547 |
+
.map_err(|_| pairing_unavailable_error())?;
|
| 548 |
+
if current_auth.account_id != auth.account_id {
|
| 549 |
+
return Err(pairing_unavailable_error());
|
| 550 |
+
}
|
| 551 |
+
if !self.desired_state_tx.borrow().is_enabled() {
|
| 552 |
+
return Err(Self::pairing_disabled_error());
|
| 553 |
+
}
|
| 554 |
+
pairing_response
|
| 555 |
+
}
|
| 556 |
+
|
| 557 |
+
async fn load_or_enroll_pairing_server(
|
| 558 |
+
&self,
|
| 559 |
+
current_enrollment: &mut Option<RemoteControlEnrollment>,
|
| 560 |
+
auth: &mut auth::RemoteControlConnectionAuth,
|
| 561 |
+
installation_id: &str,
|
| 562 |
+
server_name: &str,
|
| 563 |
+
app_server_client_name: Option<&str>,
|
| 564 |
+
selection: RemoteControlEnrollmentSelection,
|
| 565 |
+
) -> io::Result<RemoteControlEnrollment> {
|
| 566 |
+
let (enrollment, created) = self
|
| 567 |
+
.load_or_enroll_server(
|
| 568 |
+
current_enrollment,
|
| 569 |
+
auth,
|
| 570 |
+
installation_id,
|
| 571 |
+
server_name,
|
| 572 |
+
app_server_client_name,
|
| 573 |
+
selection,
|
| 574 |
+
)
|
| 575 |
+
.await?;
|
| 576 |
+
if !created {
|
| 577 |
+
publish_current_enrollment(current_enrollment, &enrollment);
|
| 578 |
+
return Ok(enrollment);
|
| 579 |
+
}
|
| 580 |
+
|
| 581 |
+
let state_db = self
|
| 582 |
+
.state_db
|
| 583 |
+
.as_deref()
|
| 584 |
+
.ok_or_else(pairing_unavailable_error)?;
|
| 585 |
+
persistence::save_enrollment(
|
| 586 |
+
&self.auth_manager,
|
| 587 |
+
&self.persistence,
|
| 588 |
+
state_db,
|
| 589 |
+
&enrollment,
|
| 590 |
+
app_server_client_name,
|
| 591 |
+
&self.desired_state_tx,
|
| 592 |
+
)
|
| 593 |
+
.await?;
|
| 594 |
+
publish_current_enrollment(current_enrollment, &enrollment);
|
| 595 |
+
Ok(enrollment)
|
| 596 |
+
}
|
| 597 |
+
|
| 598 |
+
async fn load_or_enroll_server(
|
| 599 |
+
&self,
|
| 600 |
+
current_enrollment: &Option<RemoteControlEnrollment>,
|
| 601 |
+
auth: &mut auth::RemoteControlConnectionAuth,
|
| 602 |
+
installation_id: &str,
|
| 603 |
+
server_name: &str,
|
| 604 |
+
app_server_client_name: Option<&str>,
|
| 605 |
+
selection: RemoteControlEnrollmentSelection,
|
| 606 |
+
) -> io::Result<(RemoteControlEnrollment, bool)> {
|
| 607 |
+
let remote_control_target = normalize_remote_control_url(&self.remote_control_url)?;
|
| 608 |
+
match selection {
|
| 609 |
+
RemoteControlEnrollmentSelection::ReuseOrCreate => {
|
| 610 |
+
if let Some(enrollment) = current_enrollment
|
| 611 |
+
.as_ref()
|
| 612 |
+
.filter(|enrollment| enrollment.account_id == auth.account_id)
|
| 613 |
+
.cloned()
|
| 614 |
+
{
|
| 615 |
+
return Ok((enrollment, false));
|
| 616 |
+
}
|
| 617 |
+
|
| 618 |
+
let state_db = self
|
| 619 |
+
.state_db
|
| 620 |
+
.as_deref()
|
| 621 |
+
.ok_or_else(pairing_unavailable_error)?;
|
| 622 |
+
let _persistence =
|
| 623 |
+
persistence::read_lock(&self.auth_manager, &self.persistence).await?;
|
| 624 |
+
if let Some(mut enrollment) = load_persisted_remote_control_enrollment(
|
| 625 |
+
Some(state_db),
|
| 626 |
+
&remote_control_target,
|
| 627 |
+
&auth.account_id,
|
| 628 |
+
app_server_client_name,
|
| 629 |
+
)
|
| 630 |
+
.await?
|
| 631 |
+
{
|
| 632 |
+
enrollment.server_name = server_name.to_string();
|
| 633 |
+
return Ok((enrollment, false));
|
| 634 |
+
}
|
| 635 |
+
}
|
| 636 |
+
RemoteControlEnrollmentSelection::ReplaceExisting => {}
|
| 637 |
+
}
|
| 638 |
+
|
| 639 |
+
self.current_enrollment.check_retry_after()?;
|
| 640 |
+
// Reused enrollments must still reach durable persistence during shutdown.
|
| 641 |
+
let enrollment = tokio::select! {
|
| 642 |
+
biased;
|
| 643 |
+
_ = self.shutdown_token.cancelled() => {
|
| 644 |
+
return Err(io::Error::new(
|
| 645 |
+
io::ErrorKind::Interrupted,
|
| 646 |
+
"remote control is shutting down",
|
| 647 |
+
));
|
| 648 |
+
}
|
| 649 |
+
result = enroll_pairing_server(
|
| 650 |
+
&self.current_enrollment,
|
| 651 |
+
&self.auth_manager,
|
| 652 |
+
auth,
|
| 653 |
+
&remote_control_target,
|
| 654 |
+
installation_id,
|
| 655 |
+
server_name,
|
| 656 |
+
) => self.current_enrollment.record_retry_after(result)?,
|
| 657 |
+
};
|
| 658 |
+
Ok((enrollment, true))
|
| 659 |
+
}
|
| 660 |
+
|
| 661 |
+
fn pairing_persistence_key(
|
| 662 |
+
&self,
|
| 663 |
+
app_server_client_name: Option<&str>,
|
| 664 |
+
) -> io::Result<Option<String>> {
|
| 665 |
+
if self.pairing_persistence_key_required && self.pairing_persistence_key.borrow().is_none()
|
| 666 |
+
{
|
| 667 |
+
let app_server_client_name =
|
| 668 |
+
app_server_client_name.ok_or_else(pairing_unavailable_error)?;
|
| 669 |
+
self.pairing_persistence_key
|
| 670 |
+
.send_replace(Some(app_server_client_name.to_string()));
|
| 671 |
+
}
|
| 672 |
+
Ok(self.pairing_persistence_key.borrow().clone())
|
| 673 |
+
}
|
| 674 |
+
|
| 675 |
+
pub async fn pairing_status(
|
| 676 |
+
&self,
|
| 677 |
+
params: RemoteControlPairingStatusParams,
|
| 678 |
+
) -> io::Result<RemoteControlPairingStatusResponse> {
|
| 679 |
+
self.ensure_remote_control_allowed_io()?;
|
| 680 |
+
if !self.desired_state_tx.borrow().is_enabled() {
|
| 681 |
+
return Err(Self::pairing_disabled_error());
|
| 682 |
+
}
|
| 683 |
+
let mut current_enrollment = self.current_enrollment.lock_for_request().await?;
|
| 684 |
+
let mut auth = load_remote_control_auth(&self.auth_manager)
|
| 685 |
+
.await
|
| 686 |
+
.map_err(|_| pairing_unavailable_error())?;
|
| 687 |
+
let app_server_client_name = self.pairing_persistence_key.borrow().clone();
|
| 688 |
+
let app_server_client_name = app_server_client_name.as_deref();
|
| 689 |
+
let mut enrollment = current_enrollment
|
| 690 |
+
.as_ref()
|
| 691 |
+
.filter(|enrollment| enrollment.account_id == auth.account_id)
|
| 692 |
+
.cloned()
|
| 693 |
+
.ok_or_else(pairing_unavailable_error)?;
|
| 694 |
+
let status = self.status();
|
| 695 |
+
let installation_id = status.installation_id;
|
| 696 |
+
let server_name = status.server_name;
|
| 697 |
+
if enrollment.should_refresh_server_token() {
|
| 698 |
+
let refresh_result = refresh_pairing_enrollment(
|
| 699 |
+
&mut current_enrollment,
|
| 700 |
+
&self.auth_manager,
|
| 701 |
+
&mut auth,
|
| 702 |
+
&installation_id,
|
| 703 |
+
&mut enrollment,
|
| 704 |
+
)
|
| 705 |
+
.await;
|
| 706 |
+
if refresh_result
|
| 707 |
+
.as_ref()
|
| 708 |
+
.is_err_and(|err| err.kind() == io::ErrorKind::NotFound)
|
| 709 |
+
{
|
| 710 |
+
self.load_or_enroll_pairing_server(
|
| 711 |
+
&mut current_enrollment,
|
| 712 |
+
&mut auth,
|
| 713 |
+
&installation_id,
|
| 714 |
+
&server_name,
|
| 715 |
+
app_server_client_name,
|
| 716 |
+
RemoteControlEnrollmentSelection::ReplaceExisting,
|
| 717 |
+
)
|
| 718 |
+
.await?;
|
| 719 |
+
return Err(pairing_unavailable_error());
|
| 720 |
+
}
|
| 721 |
+
refresh_result?;
|
| 722 |
+
}
|
| 723 |
+
let status_code = remote_control_pairing_status_code(¶ms)?;
|
| 724 |
+
let pairing_status_request =
|
| 725 |
+
|| protocol::RemoteControlPairingStatusRequest::from(status_code.clone());
|
| 726 |
+
self.current_enrollment.check_retry_after()?;
|
| 727 |
+
let pairing_status_response =
|
| 728 |
+
match enrollment.pairing_status(pairing_status_request()).await {
|
| 729 |
+
Err(err) if err.kind() == io::ErrorKind::PermissionDenied => {
|
| 730 |
+
clear_pairing_server_token(&mut current_enrollment, &mut enrollment)?;
|
| 731 |
+
refresh_pairing_enrollment(
|
| 732 |
+
&mut current_enrollment,
|
| 733 |
+
&self.auth_manager,
|
| 734 |
+
&mut auth,
|
| 735 |
+
&installation_id,
|
| 736 |
+
&mut enrollment,
|
| 737 |
+
)
|
| 738 |
+
.await?;
|
| 739 |
+
self.current_enrollment.check_retry_after()?;
|
| 740 |
+
enrollment.pairing_status(pairing_status_request()).await
|
| 741 |
+
}
|
| 742 |
+
pairing_status_response => pairing_status_response,
|
| 743 |
+
};
|
| 744 |
+
if let Err(err) = &pairing_status_response {
|
| 745 |
+
if server_api::remote_control_retry_at(err).is_some() {
|
| 746 |
+
return self
|
| 747 |
+
.current_enrollment
|
| 748 |
+
.record_retry_after(pairing_status_response);
|
| 749 |
+
}
|
| 750 |
+
match err.kind() {
|
| 751 |
+
io::ErrorKind::NotFound => {
|
| 752 |
+
self.load_or_enroll_pairing_server(
|
| 753 |
+
&mut current_enrollment,
|
| 754 |
+
&mut auth,
|
| 755 |
+
&installation_id,
|
| 756 |
+
&server_name,
|
| 757 |
+
app_server_client_name,
|
| 758 |
+
RemoteControlEnrollmentSelection::ReplaceExisting,
|
| 759 |
+
)
|
| 760 |
+
.await?;
|
| 761 |
+
return Err(pairing_unavailable_error());
|
| 762 |
+
}
|
| 763 |
+
io::ErrorKind::PermissionDenied => {
|
| 764 |
+
clear_pairing_server_token(&mut current_enrollment, &mut enrollment)?;
|
| 765 |
+
return Err(pairing_unavailable_error());
|
| 766 |
+
}
|
| 767 |
+
_ => {}
|
| 768 |
+
}
|
| 769 |
+
}
|
| 770 |
+
if !self.desired_state_tx.borrow().is_enabled() {
|
| 771 |
+
return Err(Self::pairing_disabled_error());
|
| 772 |
+
}
|
| 773 |
+
let current_auth = load_remote_control_auth(&self.auth_manager)
|
| 774 |
+
.await
|
| 775 |
+
.map_err(|_| pairing_unavailable_error())?;
|
| 776 |
+
if current_auth.account_id != auth.account_id {
|
| 777 |
+
return Err(pairing_unavailable_error());
|
| 778 |
+
}
|
| 779 |
+
pairing_status_response
|
| 780 |
+
}
|
| 781 |
+
|
| 782 |
+
pub async fn list_clients(
|
| 783 |
+
&self,
|
| 784 |
+
params: RemoteControlClientsListParams,
|
| 785 |
+
) -> io::Result<RemoteControlClientsListResponse> {
|
| 786 |
+
self.ensure_remote_control_allowed_io()?;
|
| 787 |
+
clients::list_remote_control_clients(&self.remote_control_url, &self.auth_manager, params)
|
| 788 |
+
.await
|
| 789 |
+
}
|
| 790 |
+
|
| 791 |
+
pub async fn revoke_client(
|
| 792 |
+
&self,
|
| 793 |
+
params: RemoteControlClientsRevokeParams,
|
| 794 |
+
) -> io::Result<RemoteControlClientsRevokeResponse> {
|
| 795 |
+
self.ensure_remote_control_allowed_io()?;
|
| 796 |
+
clients::revoke_remote_control_client(&self.remote_control_url, &self.auth_manager, params)
|
| 797 |
+
.await
|
| 798 |
+
}
|
| 799 |
+
|
| 800 |
+
fn pairing_disabled_error() -> io::Error {
|
| 801 |
+
io::Error::new(
|
| 802 |
+
io::ErrorKind::InvalidInput,
|
| 803 |
+
"remote control pairing requires remote control to be enabled",
|
| 804 |
+
)
|
| 805 |
+
}
|
| 806 |
+
|
| 807 |
+
fn publish_status(
|
| 808 |
+
&self,
|
| 809 |
+
connection_status: RemoteControlConnectionStatus,
|
| 810 |
+
) -> RemoteControlStatusChangedNotification {
|
| 811 |
+
let mut status_change = None;
|
| 812 |
+
self.status_tx.send_if_modified(|status| {
|
| 813 |
+
let next_status =
|
| 814 |
+
remote_control_status_with_connection_status(status, connection_status);
|
| 815 |
+
if *status == next_status {
|
| 816 |
+
return false;
|
| 817 |
+
}
|
| 818 |
+
|
| 819 |
+
status_change = Some((status.clone(), next_status.clone()));
|
| 820 |
+
*status = next_status;
|
| 821 |
+
true
|
| 822 |
+
});
|
| 823 |
+
if let Some((previous_status, next_status)) = status_change {
|
| 824 |
+
info!(
|
| 825 |
+
previous_status = ?previous_status.status,
|
| 826 |
+
next_status = ?next_status.status,
|
| 827 |
+
previous_environment_id = ?previous_status.environment_id,
|
| 828 |
+
next_environment_id = ?next_status.environment_id,
|
| 829 |
+
installation_id = %next_status.installation_id,
|
| 830 |
+
server_name = %next_status.server_name,
|
| 831 |
+
"remote control handle status changed"
|
| 832 |
+
);
|
| 833 |
+
}
|
| 834 |
+
self.status()
|
| 835 |
+
}
|
| 836 |
+
}
|
| 837 |
+
|
| 838 |
+
async fn enroll_pairing_server(
|
| 839 |
+
current_enrollment: &RemoteControlEnrollmentState,
|
| 840 |
+
auth_manager: &auth::RemoteControlAuth,
|
| 841 |
+
auth: &mut auth::RemoteControlConnectionAuth,
|
| 842 |
+
remote_control_target: &protocol::RemoteControlTarget,
|
| 843 |
+
installation_id: &str,
|
| 844 |
+
server_name: &str,
|
| 845 |
+
) -> io::Result<RemoteControlEnrollment> {
|
| 846 |
+
match enroll_remote_control_server(remote_control_target, auth, installation_id, server_name)
|
| 847 |
+
.await
|
| 848 |
+
{
|
| 849 |
+
Ok(enrollment) => return Ok(enrollment),
|
| 850 |
+
Err(err) if err.kind() == io::ErrorKind::PermissionDenied => {
|
| 851 |
+
let mut auth_recovery = auth_manager.unauthorized_recovery();
|
| 852 |
+
let mut auth_change_rx = auth_manager.auth_change_receiver();
|
| 853 |
+
if !recover_remote_control_auth(&mut auth_recovery, &mut auth_change_rx).await {
|
| 854 |
+
return Err(err);
|
| 855 |
+
}
|
| 856 |
+
*auth = load_remote_control_auth(auth_manager)
|
| 857 |
+
.await
|
| 858 |
+
.map_err(|_| pairing_unavailable_error())?;
|
| 859 |
+
}
|
| 860 |
+
Err(err) => return Err(err),
|
| 861 |
+
}
|
| 862 |
+
current_enrollment.check_retry_after()?;
|
| 863 |
+
enroll_remote_control_server(remote_control_target, auth, installation_id, server_name).await
|
| 864 |
+
}
|
| 865 |
+
|
| 866 |
+
fn remote_control_pairing_status_code(
|
| 867 |
+
params: &RemoteControlPairingStatusParams,
|
| 868 |
+
) -> io::Result<RemoteControlPairingStatusCode> {
|
| 869 |
+
match (¶ms.pairing_code, ¶ms.manual_pairing_code) {
|
| 870 |
+
(Some(pairing_code), None) => Ok(RemoteControlPairingStatusCode::PairingCode(
|
| 871 |
+
pairing_code.clone(),
|
| 872 |
+
)),
|
| 873 |
+
(None, Some(manual_pairing_code)) => Ok(RemoteControlPairingStatusCode::ManualPairingCode(
|
| 874 |
+
manual_pairing_code.clone(),
|
| 875 |
+
)),
|
| 876 |
+
(Some(_), Some(_)) => Err(io::Error::new(
|
| 877 |
+
io::ErrorKind::InvalidInput,
|
| 878 |
+
"remote control pairing status accepts either pairingCode or manualPairingCode, not both",
|
| 879 |
+
)),
|
| 880 |
+
(None, None) => Err(io::Error::new(
|
| 881 |
+
io::ErrorKind::InvalidInput,
|
| 882 |
+
"remote control pairing status requires pairingCode or manualPairingCode",
|
| 883 |
+
)),
|
| 884 |
+
}
|
| 885 |
+
}
|
| 886 |
+
|
| 887 |
+
async fn refresh_pairing_enrollment(
|
| 888 |
+
current_enrollment: &mut RemoteControlEnrollmentLease<'_>,
|
| 889 |
+
auth_manager: &auth::RemoteControlAuth,
|
| 890 |
+
auth: &mut auth::RemoteControlConnectionAuth,
|
| 891 |
+
installation_id: &str,
|
| 892 |
+
enrollment: &mut RemoteControlEnrollment,
|
| 893 |
+
) -> io::Result<()> {
|
| 894 |
+
current_enrollment.state.check_retry_after()?;
|
| 895 |
+
let mut refresh_result = refresh_remote_control_server(auth, installation_id, enrollment).await;
|
| 896 |
+
if refresh_result
|
| 897 |
+
.as_ref()
|
| 898 |
+
.is_err_and(|err| err.kind() == io::ErrorKind::PermissionDenied)
|
| 899 |
+
{
|
| 900 |
+
let mut auth_recovery = auth_manager.unauthorized_recovery();
|
| 901 |
+
let mut auth_change_rx = auth_manager.auth_change_receiver();
|
| 902 |
+
if recover_remote_control_auth(&mut auth_recovery, &mut auth_change_rx).await {
|
| 903 |
+
match load_remote_control_auth(auth_manager).await {
|
| 904 |
+
Ok(recovered_auth) if recovered_auth.account_id == enrollment.account_id => {
|
| 905 |
+
*auth = recovered_auth;
|
| 906 |
+
current_enrollment.state.check_retry_after()?;
|
| 907 |
+
refresh_result =
|
| 908 |
+
refresh_remote_control_server(auth, installation_id, enrollment).await;
|
| 909 |
+
}
|
| 910 |
+
Ok(_) | Err(_) => {
|
| 911 |
+
enrollment.clear_server_token();
|
| 912 |
+
refresh_result = Err(pairing_unavailable_error());
|
| 913 |
+
}
|
| 914 |
+
}
|
| 915 |
+
} else {
|
| 916 |
+
enrollment.clear_server_token();
|
| 917 |
+
}
|
| 918 |
+
}
|
| 919 |
+
if refresh_result
|
| 920 |
+
.as_ref()
|
| 921 |
+
.is_err_and(|err| err.kind() == io::ErrorKind::PermissionDenied)
|
| 922 |
+
{
|
| 923 |
+
enrollment.clear_server_token();
|
| 924 |
+
}
|
| 925 |
+
if !replace_current_enrollment(current_enrollment, enrollment) {
|
| 926 |
+
Err(pairing_unavailable_error())
|
| 927 |
+
} else {
|
| 928 |
+
current_enrollment.state.record_retry_after(refresh_result)
|
| 929 |
+
}
|
| 930 |
+
}
|
| 931 |
+
|
| 932 |
+
fn clear_pairing_server_token(
|
| 933 |
+
current_enrollment: &mut Option<RemoteControlEnrollment>,
|
| 934 |
+
enrollment: &mut RemoteControlEnrollment,
|
| 935 |
+
) -> io::Result<()> {
|
| 936 |
+
enrollment.clear_server_token();
|
| 937 |
+
if replace_current_enrollment(current_enrollment, enrollment) {
|
| 938 |
+
Ok(())
|
| 939 |
+
} else {
|
| 940 |
+
Err(pairing_unavailable_error())
|
| 941 |
+
}
|
| 942 |
+
}
|
| 943 |
+
|
| 944 |
+
fn pairing_unavailable_error() -> io::Error {
|
| 945 |
+
io::Error::new(
|
| 946 |
+
io::ErrorKind::InvalidInput,
|
| 947 |
+
"remote control pairing is unavailable until enrollment completes",
|
| 948 |
+
)
|
| 949 |
+
}
|
| 950 |
+
|
| 951 |
+
fn remote_control_status_with_connection_status(
|
| 952 |
+
status: &RemoteControlStatusChangedNotification,
|
| 953 |
+
connection_status: RemoteControlConnectionStatus,
|
| 954 |
+
) -> RemoteControlStatusChangedNotification {
|
| 955 |
+
RemoteControlStatusChangedNotification {
|
| 956 |
+
status: connection_status,
|
| 957 |
+
server_name: status.server_name.clone(),
|
| 958 |
+
installation_id: status.installation_id.clone(),
|
| 959 |
+
environment_id: if connection_status == RemoteControlConnectionStatus::Disabled {
|
| 960 |
+
None
|
| 961 |
+
} else {
|
| 962 |
+
status.environment_id.clone()
|
| 963 |
+
},
|
| 964 |
+
}
|
| 965 |
+
}
|
| 966 |
+
|
| 967 |
+
fn publish_current_enrollment(
|
| 968 |
+
current_enrollment: &mut Option<RemoteControlEnrollment>,
|
| 969 |
+
enrollment: &RemoteControlEnrollment,
|
| 970 |
+
) {
|
| 971 |
+
*current_enrollment = Some(enrollment.clone());
|
| 972 |
+
}
|
| 973 |
+
|
| 974 |
+
fn replace_current_enrollment(
|
| 975 |
+
current_enrollment: &mut Option<RemoteControlEnrollment>,
|
| 976 |
+
enrollment: &RemoteControlEnrollment,
|
| 977 |
+
) -> bool {
|
| 978 |
+
if !current_enrollment
|
| 979 |
+
.as_ref()
|
| 980 |
+
.is_some_and(|current| same_remote_control_enrollment(current, enrollment))
|
| 981 |
+
{
|
| 982 |
+
return false;
|
| 983 |
+
}
|
| 984 |
+
*current_enrollment = Some(enrollment.clone());
|
| 985 |
+
true
|
| 986 |
+
}
|
| 987 |
+
|
| 988 |
+
fn same_remote_control_enrollment(
|
| 989 |
+
left: &RemoteControlEnrollment,
|
| 990 |
+
right: &RemoteControlEnrollment,
|
| 991 |
+
) -> bool {
|
| 992 |
+
// A refresh rotates only the bearer. Pairing remains current while the same persisted server
|
| 993 |
+
// record is still selected for the current account.
|
| 994 |
+
left.account_id == right.account_id
|
| 995 |
+
&& left.server_id == right.server_id
|
| 996 |
+
&& left.environment_id == right.environment_id
|
| 997 |
+
}
|
| 998 |
+
|
| 999 |
+
#[cfg(test)]
|
| 1000 |
+
mod segment_tests;
|
| 1001 |
+
#[cfg(test)]
|
| 1002 |
+
mod tests;
|
codex-rs/app-server-transport/src/transport/remote_control/persistence.rs
ADDED
|
@@ -0,0 +1,165 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Serializes enrollment storage across login sessions.
|
| 2 |
+
//! An admitted SQLite write keeps its permit until completion even if its caller is cancelled.
|
| 3 |
+
//! Process shutdown drains admitted writes after stopping the session workers.
|
| 4 |
+
|
| 5 |
+
use super::RemoteControlSession;
|
| 6 |
+
use super::auth::RemoteControlAuth;
|
| 7 |
+
use super::desired_state::RemoteControlDesiredState;
|
| 8 |
+
use super::enroll::RemoteControlEnrollment;
|
| 9 |
+
use super::enroll::update_persisted_remote_control_enrollment;
|
| 10 |
+
use super::protocol::RemoteControlTarget;
|
| 11 |
+
use codex_state::StateRuntime;
|
| 12 |
+
use std::future::Future;
|
| 13 |
+
use std::io;
|
| 14 |
+
use std::sync::Arc;
|
| 15 |
+
use tokio::sync::Semaphore;
|
| 16 |
+
use tokio::sync::SemaphorePermit;
|
| 17 |
+
use tokio::sync::watch;
|
| 18 |
+
use tokio_util::task::TaskTracker;
|
| 19 |
+
|
| 20 |
+
#[cfg(test)]
|
| 21 |
+
#[path = "persistence_tests.rs"]
|
| 22 |
+
mod tests;
|
| 23 |
+
|
| 24 |
+
#[derive(Clone)]
|
| 25 |
+
pub(super) struct RemoteControlPersistence {
|
| 26 |
+
semaphore: Arc<Semaphore>,
|
| 27 |
+
pub(super) tasks: TaskTracker,
|
| 28 |
+
}
|
| 29 |
+
|
| 30 |
+
impl Default for RemoteControlPersistence {
|
| 31 |
+
fn default() -> Self {
|
| 32 |
+
Self {
|
| 33 |
+
semaphore: Arc::new(Semaphore::new(1)),
|
| 34 |
+
tasks: TaskTracker::new(),
|
| 35 |
+
}
|
| 36 |
+
}
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
impl RemoteControlPersistence {
|
| 40 |
+
pub(super) async fn lock(&self) -> SemaphorePermit<'_> {
|
| 41 |
+
self.semaphore
|
| 42 |
+
.acquire()
|
| 43 |
+
.await
|
| 44 |
+
.unwrap_or_else(|_| unreachable!())
|
| 45 |
+
}
|
| 46 |
+
}
|
| 47 |
+
|
| 48 |
+
pub(super) async fn read_lock<'a>(
|
| 49 |
+
auth: &RemoteControlAuth,
|
| 50 |
+
lock: &'a RemoteControlPersistence,
|
| 51 |
+
) -> io::Result<SemaphorePermit<'a>> {
|
| 52 |
+
let permit = lock.lock().await;
|
| 53 |
+
auth.ensure_current()?;
|
| 54 |
+
Ok(permit)
|
| 55 |
+
}
|
| 56 |
+
|
| 57 |
+
async fn commit<T: Send + 'static>(
|
| 58 |
+
auth: &RemoteControlAuth,
|
| 59 |
+
lock: &RemoteControlPersistence,
|
| 60 |
+
operation: impl Future<Output = io::Result<T>> + Send + 'static,
|
| 61 |
+
) -> io::Result<T> {
|
| 62 |
+
// Register before checking admission. Shutdown either observes this token or rejects us.
|
| 63 |
+
let task = lock.tasks.token();
|
| 64 |
+
if lock.tasks.is_closed() {
|
| 65 |
+
return Err(io::Error::new(
|
| 66 |
+
io::ErrorKind::Interrupted,
|
| 67 |
+
"remote control is shutting down",
|
| 68 |
+
));
|
| 69 |
+
}
|
| 70 |
+
let auth = auth.clone();
|
| 71 |
+
let semaphore = lock.semaphore.clone();
|
| 72 |
+
tokio::spawn(async move {
|
| 73 |
+
let _task = task;
|
| 74 |
+
let _permit = semaphore.acquire_owned().await.map_err(io::Error::other)?;
|
| 75 |
+
auth.ensure_current()?;
|
| 76 |
+
operation.await
|
| 77 |
+
})
|
| 78 |
+
.await
|
| 79 |
+
.map_err(io::Error::other)?
|
| 80 |
+
}
|
| 81 |
+
|
| 82 |
+
pub(super) async fn save_enrollment(
|
| 83 |
+
auth: &RemoteControlAuth,
|
| 84 |
+
lock: &RemoteControlPersistence,
|
| 85 |
+
state_db: &StateRuntime,
|
| 86 |
+
enrollment: &RemoteControlEnrollment,
|
| 87 |
+
client_name: Option<&str>,
|
| 88 |
+
desired: &watch::Sender<RemoteControlDesiredState>,
|
| 89 |
+
) -> io::Result<()> {
|
| 90 |
+
let state_db = state_db.clone();
|
| 91 |
+
let enrollment = enrollment.clone();
|
| 92 |
+
let client_name = client_name.map(str::to_owned);
|
| 93 |
+
let desired = desired.clone();
|
| 94 |
+
commit(auth, lock, async move {
|
| 95 |
+
let preference = match *desired.borrow() {
|
| 96 |
+
RemoteControlDesiredState::Enabled {
|
| 97 |
+
persistence_preference,
|
| 98 |
+
} => persistence_preference,
|
| 99 |
+
RemoteControlDesiredState::Disabled | RemoteControlDesiredState::Unknown => {
|
| 100 |
+
return Err(io::Error::new(
|
| 101 |
+
io::ErrorKind::Interrupted,
|
| 102 |
+
"remote control disabled during enrollment",
|
| 103 |
+
));
|
| 104 |
+
}
|
| 105 |
+
};
|
| 106 |
+
update_persisted_remote_control_enrollment(
|
| 107 |
+
Some(&state_db),
|
| 108 |
+
&enrollment.remote_control_target,
|
| 109 |
+
&enrollment.account_id,
|
| 110 |
+
client_name.as_deref(),
|
| 111 |
+
Some(&enrollment),
|
| 112 |
+
preference,
|
| 113 |
+
)
|
| 114 |
+
.await
|
| 115 |
+
})
|
| 116 |
+
.await
|
| 117 |
+
}
|
| 118 |
+
|
| 119 |
+
impl RemoteControlSession {
|
| 120 |
+
pub(super) async fn set_preference(
|
| 121 |
+
&self,
|
| 122 |
+
state_db: &StateRuntime,
|
| 123 |
+
target: &RemoteControlTarget,
|
| 124 |
+
account_id: &str,
|
| 125 |
+
client_name: Option<&str>,
|
| 126 |
+
enabled: bool,
|
| 127 |
+
fallback_enrollment: Option<&RemoteControlEnrollment>,
|
| 128 |
+
) -> io::Result<()> {
|
| 129 |
+
let state_db = state_db.clone();
|
| 130 |
+
let target = target.clone();
|
| 131 |
+
let account_id = account_id.to_owned();
|
| 132 |
+
let client_name = client_name.map(str::to_owned);
|
| 133 |
+
let enrollment = fallback_enrollment.cloned();
|
| 134 |
+
let desired = self.desired_state_tx.as_ref().clone();
|
| 135 |
+
commit(&self.auth_manager, &self.persistence, async move {
|
| 136 |
+
let updated = state_db
|
| 137 |
+
.set_remote_control_enabled(
|
| 138 |
+
&target.websocket_url,
|
| 139 |
+
&account_id,
|
| 140 |
+
client_name.as_deref(),
|
| 141 |
+
enabled,
|
| 142 |
+
)
|
| 143 |
+
.await
|
| 144 |
+
.map_err(io::Error::other)?;
|
| 145 |
+
if updated == 0
|
| 146 |
+
&& let Some(enrollment) = enrollment
|
| 147 |
+
{
|
| 148 |
+
update_persisted_remote_control_enrollment(
|
| 149 |
+
Some(&state_db),
|
| 150 |
+
&target,
|
| 151 |
+
&account_id,
|
| 152 |
+
client_name.as_deref(),
|
| 153 |
+
Some(&enrollment),
|
| 154 |
+
Some(enabled),
|
| 155 |
+
)
|
| 156 |
+
.await?;
|
| 157 |
+
}
|
| 158 |
+
if !enabled {
|
| 159 |
+
desired.send_replace(RemoteControlDesiredState::Disabled);
|
| 160 |
+
}
|
| 161 |
+
Ok(())
|
| 162 |
+
})
|
| 163 |
+
.await
|
| 164 |
+
}
|
| 165 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/persistence_tests.rs
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Checks the persistence boundary when an operation loses its caller.
|
| 2 |
+
|
| 3 |
+
use super::*;
|
| 4 |
+
use codex_core::test_support::auth_manager_from_auth;
|
| 5 |
+
use codex_login::CodexAuth;
|
| 6 |
+
use futures::poll;
|
| 7 |
+
use pretty_assertions::assert_eq;
|
| 8 |
+
use tokio::sync::oneshot;
|
| 9 |
+
|
| 10 |
+
#[tokio::test]
|
| 11 |
+
async fn cancelled_commit_keeps_its_permit_and_is_drained() -> io::Result<()> {
|
| 12 |
+
let auth = RemoteControlAuth::capture(auth_manager_from_auth(
|
| 13 |
+
CodexAuth::create_dummy_chatgpt_auth_for_testing(),
|
| 14 |
+
))
|
| 15 |
+
.0;
|
| 16 |
+
let persistence = RemoteControlPersistence::default();
|
| 17 |
+
let writes = Arc::new(tokio::sync::Mutex::new(Vec::new()));
|
| 18 |
+
let (started_tx, started_rx) = oneshot::channel();
|
| 19 |
+
let (release_tx, release_rx) = oneshot::channel();
|
| 20 |
+
let old_writes = writes.clone();
|
| 21 |
+
let mut old = Box::pin(commit(&auth, &persistence, async move {
|
| 22 |
+
started_tx.send(()).expect("test receiver is open");
|
| 23 |
+
release_rx.await.map_err(io::Error::other)?;
|
| 24 |
+
old_writes.lock().await.push("old");
|
| 25 |
+
Ok(())
|
| 26 |
+
}));
|
| 27 |
+
assert!(poll!(&mut old).is_pending());
|
| 28 |
+
started_rx.await.map_err(io::Error::other)?;
|
| 29 |
+
drop(old);
|
| 30 |
+
|
| 31 |
+
let new_writes = writes.clone();
|
| 32 |
+
let mut new = Box::pin(commit(&auth, &persistence, async move {
|
| 33 |
+
new_writes.lock().await.push("new");
|
| 34 |
+
Ok(())
|
| 35 |
+
}));
|
| 36 |
+
assert!(poll!(&mut new).is_pending());
|
| 37 |
+
persistence.tasks.close();
|
| 38 |
+
let mut drained = Box::pin(persistence.tasks.wait());
|
| 39 |
+
assert!(poll!(&mut drained).is_pending());
|
| 40 |
+
release_tx.send(()).expect("commit still owns its receiver");
|
| 41 |
+
tokio::time::timeout(std::time::Duration::from_secs(5), drained).await?;
|
| 42 |
+
new.await?;
|
| 43 |
+
assert_eq!(*writes.lock().await, vec!["old", "new"]);
|
| 44 |
+
let error = commit::<()>(&auth, &persistence, async {
|
| 45 |
+
panic!("closed storage must reject new work")
|
| 46 |
+
})
|
| 47 |
+
.await
|
| 48 |
+
.expect_err("shutdown closes admission");
|
| 49 |
+
assert_eq!(error.kind(), io::ErrorKind::Interrupted);
|
| 50 |
+
Ok(())
|
| 51 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/protocol.rs
ADDED
|
@@ -0,0 +1,401 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use crate::outgoing_message::OutgoingMessage;
|
| 2 |
+
use codex_app_server_protocol::JSONRPCMessage;
|
| 3 |
+
use serde::Deserialize;
|
| 4 |
+
use serde::Serialize;
|
| 5 |
+
use std::io;
|
| 6 |
+
use std::io::ErrorKind;
|
| 7 |
+
use url::Host;
|
| 8 |
+
use url::Url;
|
| 9 |
+
|
| 10 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 11 |
+
pub(super) struct RemoteControlTarget {
|
| 12 |
+
pub(super) websocket_url: String,
|
| 13 |
+
pub(super) enroll_url: String,
|
| 14 |
+
pub(super) refresh_url: String,
|
| 15 |
+
pub(super) pair_url: String,
|
| 16 |
+
pub(super) pair_status_url: String,
|
| 17 |
+
}
|
| 18 |
+
|
| 19 |
+
#[derive(Debug, Serialize)]
|
| 20 |
+
pub(super) struct EnrollRemoteServerRequest {
|
| 21 |
+
pub(super) name: String,
|
| 22 |
+
pub(super) os: &'static str,
|
| 23 |
+
pub(super) arch: &'static str,
|
| 24 |
+
pub(super) app_server_version: &'static str,
|
| 25 |
+
pub(super) installation_id: String,
|
| 26 |
+
}
|
| 27 |
+
|
| 28 |
+
#[derive(Debug, Deserialize)]
|
| 29 |
+
pub(super) struct EnrollRemoteServerResponse {
|
| 30 |
+
pub(super) server_id: String,
|
| 31 |
+
pub(super) environment_id: String,
|
| 32 |
+
pub(super) remote_control_token: String,
|
| 33 |
+
pub(super) expires_at: String,
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
#[derive(Debug, Serialize)]
|
| 37 |
+
pub(super) struct RefreshRemoteServerRequest {
|
| 38 |
+
pub(super) server_id: String,
|
| 39 |
+
pub(super) installation_id: String,
|
| 40 |
+
}
|
| 41 |
+
|
| 42 |
+
#[derive(Debug, Serialize)]
|
| 43 |
+
pub(super) struct StartRemoteControlPairingRequest {
|
| 44 |
+
pub(super) manual_code: bool,
|
| 45 |
+
}
|
| 46 |
+
|
| 47 |
+
#[derive(Debug, Deserialize)]
|
| 48 |
+
pub(super) struct StartRemoteControlPairingResponse {
|
| 49 |
+
pub(super) pairing_code: String,
|
| 50 |
+
pub(super) manual_pairing_code: Option<String>,
|
| 51 |
+
pub(super) server_id: String,
|
| 52 |
+
pub(super) environment_id: String,
|
| 53 |
+
pub(super) expires_at: String,
|
| 54 |
+
}
|
| 55 |
+
|
| 56 |
+
#[derive(Debug, Serialize)]
|
| 57 |
+
pub(super) struct RemoteControlPairingStatusRequest {
|
| 58 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 59 |
+
pub(super) pairing_code: Option<String>,
|
| 60 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 61 |
+
pub(super) manual_pairing_code: Option<String>,
|
| 62 |
+
}
|
| 63 |
+
|
| 64 |
+
#[derive(Clone)]
|
| 65 |
+
pub(super) enum RemoteControlPairingStatusCode {
|
| 66 |
+
PairingCode(String),
|
| 67 |
+
ManualPairingCode(String),
|
| 68 |
+
}
|
| 69 |
+
|
| 70 |
+
impl From<RemoteControlPairingStatusCode> for RemoteControlPairingStatusRequest {
|
| 71 |
+
fn from(code: RemoteControlPairingStatusCode) -> Self {
|
| 72 |
+
match code {
|
| 73 |
+
RemoteControlPairingStatusCode::PairingCode(pairing_code) => Self {
|
| 74 |
+
pairing_code: Some(pairing_code),
|
| 75 |
+
manual_pairing_code: None,
|
| 76 |
+
},
|
| 77 |
+
RemoteControlPairingStatusCode::ManualPairingCode(manual_pairing_code) => Self {
|
| 78 |
+
pairing_code: None,
|
| 79 |
+
manual_pairing_code: Some(manual_pairing_code),
|
| 80 |
+
},
|
| 81 |
+
}
|
| 82 |
+
}
|
| 83 |
+
}
|
| 84 |
+
|
| 85 |
+
#[derive(Debug, Deserialize)]
|
| 86 |
+
pub(super) struct RemoteControlPairingStatusResponse {
|
| 87 |
+
pub(super) claimed: bool,
|
| 88 |
+
}
|
| 89 |
+
|
| 90 |
+
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
| 91 |
+
#[serde(transparent)]
|
| 92 |
+
pub struct ClientId(pub String);
|
| 93 |
+
|
| 94 |
+
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
| 95 |
+
#[serde(transparent)]
|
| 96 |
+
pub struct StreamId(pub String);
|
| 97 |
+
|
| 98 |
+
impl StreamId {
|
| 99 |
+
pub fn new_random() -> Self {
|
| 100 |
+
Self(uuid::Uuid::now_v7().to_string())
|
| 101 |
+
}
|
| 102 |
+
}
|
| 103 |
+
|
| 104 |
+
#[derive(Debug, Clone, Serialize, Deserialize)]
|
| 105 |
+
#[serde(tag = "type", rename_all = "snake_case")]
|
| 106 |
+
pub enum ClientEvent {
|
| 107 |
+
ClientMessage {
|
| 108 |
+
message: JSONRPCMessage,
|
| 109 |
+
},
|
| 110 |
+
ClientMessageChunk {
|
| 111 |
+
segment_id: usize,
|
| 112 |
+
segment_count: usize,
|
| 113 |
+
message_size_bytes: usize,
|
| 114 |
+
message_chunk_base64: String,
|
| 115 |
+
},
|
| 116 |
+
/// Backend-generated acknowledgement for all server envelopes addressed to
|
| 117 |
+
/// `client_id` and `stream_id` whose envelope `seq_id` is less than or equal
|
| 118 |
+
/// to this ack's `seq_id`. Chunk acknowledgements carry `segment_id` so the
|
| 119 |
+
/// sender can retain only the still-unacked wire chunks on reconnect.
|
| 120 |
+
Ack {
|
| 121 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 122 |
+
segment_id: Option<usize>,
|
| 123 |
+
},
|
| 124 |
+
Ping,
|
| 125 |
+
ClientClosed,
|
| 126 |
+
}
|
| 127 |
+
|
| 128 |
+
#[derive(Debug, Clone, Serialize, Deserialize)]
|
| 129 |
+
#[serde(rename_all = "snake_case")]
|
| 130 |
+
pub(crate) struct ClientEnvelope {
|
| 131 |
+
#[serde(flatten)]
|
| 132 |
+
pub(crate) event: ClientEvent,
|
| 133 |
+
#[serde(rename = "client_id")]
|
| 134 |
+
pub(crate) client_id: ClientId,
|
| 135 |
+
#[serde(rename = "stream_id", skip_serializing_if = "Option::is_none")]
|
| 136 |
+
pub(crate) stream_id: Option<StreamId>,
|
| 137 |
+
/// For `Ack`, this is the backend-generated per-stream cursor over
|
| 138 |
+
/// `ServerEnvelope.seq_id`.
|
| 139 |
+
#[serde(rename = "seq_id", skip_serializing_if = "Option::is_none")]
|
| 140 |
+
pub(crate) seq_id: Option<u64>,
|
| 141 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 142 |
+
pub(crate) cursor: Option<String>,
|
| 143 |
+
}
|
| 144 |
+
|
| 145 |
+
#[derive(Debug, Clone, Serialize, Deserialize)]
|
| 146 |
+
#[serde(rename_all = "snake_case")]
|
| 147 |
+
pub enum PongStatus {
|
| 148 |
+
Active,
|
| 149 |
+
Unknown,
|
| 150 |
+
}
|
| 151 |
+
|
| 152 |
+
#[derive(Debug, Clone, Serialize)]
|
| 153 |
+
#[serde(tag = "type", rename_all = "snake_case")]
|
| 154 |
+
pub enum ServerEvent {
|
| 155 |
+
ServerMessage {
|
| 156 |
+
message: Box<OutgoingMessage>,
|
| 157 |
+
},
|
| 158 |
+
ServerMessageChunk {
|
| 159 |
+
segment_id: usize,
|
| 160 |
+
segment_count: usize,
|
| 161 |
+
message_size_bytes: usize,
|
| 162 |
+
message_chunk_base64: String,
|
| 163 |
+
},
|
| 164 |
+
#[allow(dead_code)]
|
| 165 |
+
Ack,
|
| 166 |
+
Pong {
|
| 167 |
+
status: PongStatus,
|
| 168 |
+
},
|
| 169 |
+
}
|
| 170 |
+
|
| 171 |
+
impl ServerEvent {
|
| 172 |
+
pub(crate) fn segment_id(&self) -> Option<usize> {
|
| 173 |
+
match self {
|
| 174 |
+
Self::ServerMessageChunk { segment_id, .. } => Some(*segment_id),
|
| 175 |
+
Self::ServerMessage { .. } | Self::Ack | Self::Pong { .. } => None,
|
| 176 |
+
}
|
| 177 |
+
}
|
| 178 |
+
}
|
| 179 |
+
|
| 180 |
+
#[derive(Debug, Clone, Serialize)]
|
| 181 |
+
#[serde(rename_all = "snake_case")]
|
| 182 |
+
pub(crate) struct ServerEnvelope {
|
| 183 |
+
#[serde(flatten)]
|
| 184 |
+
pub(crate) event: ServerEvent,
|
| 185 |
+
#[serde(rename = "client_id")]
|
| 186 |
+
pub(crate) client_id: ClientId,
|
| 187 |
+
#[serde(rename = "stream_id")]
|
| 188 |
+
pub(crate) stream_id: StreamId,
|
| 189 |
+
#[serde(rename = "seq_id")]
|
| 190 |
+
pub(crate) seq_id: u64,
|
| 191 |
+
}
|
| 192 |
+
|
| 193 |
+
fn is_allowed_remote_control_chatgpt_host(host: &Option<Host<&str>>) -> bool {
|
| 194 |
+
let Some(Host::Domain(host)) = *host else {
|
| 195 |
+
return false;
|
| 196 |
+
};
|
| 197 |
+
host == "chatgpt.com"
|
| 198 |
+
|| host == "chatgpt-staging.com"
|
| 199 |
+
|| host.ends_with(".chatgpt.com")
|
| 200 |
+
|| host.ends_with(".chatgpt-staging.com")
|
| 201 |
+
}
|
| 202 |
+
|
| 203 |
+
fn is_localhost(host: &Option<Host<&str>>) -> bool {
|
| 204 |
+
match host {
|
| 205 |
+
Some(Host::Domain("localhost")) => true,
|
| 206 |
+
Some(Host::Ipv4(ip)) => ip.is_loopback(),
|
| 207 |
+
Some(Host::Ipv6(ip)) => ip.is_loopback(),
|
| 208 |
+
_ => false,
|
| 209 |
+
}
|
| 210 |
+
}
|
| 211 |
+
|
| 212 |
+
pub(super) fn normalize_remote_control_url(
|
| 213 |
+
remote_control_url: &str,
|
| 214 |
+
) -> io::Result<RemoteControlTarget> {
|
| 215 |
+
let remote_control_url = normalize_remote_control_base_url(remote_control_url)?;
|
| 216 |
+
let map_url_parse_error = |err: url::ParseError| -> io::Error {
|
| 217 |
+
io::Error::new(
|
| 218 |
+
ErrorKind::InvalidInput,
|
| 219 |
+
format!("invalid remote control URL `{remote_control_url}`: {err}"),
|
| 220 |
+
)
|
| 221 |
+
};
|
| 222 |
+
|
| 223 |
+
let enroll_url = remote_control_url
|
| 224 |
+
.join("wham/remote/control/server/enroll")
|
| 225 |
+
.map_err(map_url_parse_error)?;
|
| 226 |
+
let refresh_url = remote_control_url
|
| 227 |
+
.join("wham/remote/control/server/refresh")
|
| 228 |
+
.map_err(map_url_parse_error)?;
|
| 229 |
+
let pair_url = remote_control_url
|
| 230 |
+
.join("wham/remote/control/server/pair")
|
| 231 |
+
.map_err(map_url_parse_error)?;
|
| 232 |
+
let pair_status_url = remote_control_url
|
| 233 |
+
.join("wham/remote/control/server/pair/status")
|
| 234 |
+
.map_err(map_url_parse_error)?;
|
| 235 |
+
let mut websocket_url = remote_control_url
|
| 236 |
+
.join("wham/remote/control/server")
|
| 237 |
+
.map_err(map_url_parse_error)?;
|
| 238 |
+
websocket_url
|
| 239 |
+
.set_scheme(if enroll_url.scheme() == "https" {
|
| 240 |
+
"wss"
|
| 241 |
+
} else {
|
| 242 |
+
"ws"
|
| 243 |
+
})
|
| 244 |
+
.map_err(|()| {
|
| 245 |
+
io::Error::new(
|
| 246 |
+
ErrorKind::InvalidInput,
|
| 247 |
+
format!("invalid remote control URL `{remote_control_url}`"),
|
| 248 |
+
)
|
| 249 |
+
})?;
|
| 250 |
+
|
| 251 |
+
Ok(RemoteControlTarget {
|
| 252 |
+
websocket_url: websocket_url.to_string(),
|
| 253 |
+
enroll_url: enroll_url.to_string(),
|
| 254 |
+
refresh_url: refresh_url.to_string(),
|
| 255 |
+
pair_url: pair_url.to_string(),
|
| 256 |
+
pair_status_url: pair_status_url.to_string(),
|
| 257 |
+
})
|
| 258 |
+
}
|
| 259 |
+
|
| 260 |
+
pub(super) fn normalize_remote_control_base_url(remote_control_url: &str) -> io::Result<Url> {
|
| 261 |
+
let map_url_parse_error = |err: url::ParseError| -> io::Error {
|
| 262 |
+
io::Error::new(
|
| 263 |
+
ErrorKind::InvalidInput,
|
| 264 |
+
format!("invalid remote control URL `{remote_control_url}`: {err}"),
|
| 265 |
+
)
|
| 266 |
+
};
|
| 267 |
+
let map_scheme_error = |_: ()| -> io::Error {
|
| 268 |
+
io::Error::new(
|
| 269 |
+
ErrorKind::InvalidInput,
|
| 270 |
+
format!(
|
| 271 |
+
"invalid remote control URL `{remote_control_url}`; expected HTTPS URL for chatgpt.com or chatgpt-staging.com, or HTTP/HTTPS URL for localhost"
|
| 272 |
+
),
|
| 273 |
+
)
|
| 274 |
+
};
|
| 275 |
+
|
| 276 |
+
let mut remote_control_url = Url::parse(remote_control_url).map_err(map_url_parse_error)?;
|
| 277 |
+
if !remote_control_url.path().ends_with('/') {
|
| 278 |
+
let normalized_path = format!("{}/", remote_control_url.path());
|
| 279 |
+
remote_control_url.set_path(&normalized_path);
|
| 280 |
+
}
|
| 281 |
+
|
| 282 |
+
let host = remote_control_url.host();
|
| 283 |
+
match remote_control_url.scheme() {
|
| 284 |
+
"https" if is_localhost(&host) || is_allowed_remote_control_chatgpt_host(&host) => {}
|
| 285 |
+
"http" if is_localhost(&host) => {}
|
| 286 |
+
_ => return Err(map_scheme_error(())),
|
| 287 |
+
}
|
| 288 |
+
|
| 289 |
+
Ok(remote_control_url)
|
| 290 |
+
}
|
| 291 |
+
|
| 292 |
+
#[cfg(test)]
|
| 293 |
+
mod tests {
|
| 294 |
+
use super::*;
|
| 295 |
+
use pretty_assertions::assert_eq;
|
| 296 |
+
|
| 297 |
+
#[test]
|
| 298 |
+
fn normalize_remote_control_url_accepts_chatgpt_https_urls() {
|
| 299 |
+
assert_eq!(
|
| 300 |
+
normalize_remote_control_url("https://chatgpt.com/backend-api")
|
| 301 |
+
.expect("chatgpt.com URL should normalize"),
|
| 302 |
+
RemoteControlTarget {
|
| 303 |
+
websocket_url: "wss://chatgpt.com/backend-api/wham/remote/control/server"
|
| 304 |
+
.to_string(),
|
| 305 |
+
enroll_url: "https://chatgpt.com/backend-api/wham/remote/control/server/enroll"
|
| 306 |
+
.to_string(),
|
| 307 |
+
refresh_url: "https://chatgpt.com/backend-api/wham/remote/control/server/refresh"
|
| 308 |
+
.to_string(),
|
| 309 |
+
pair_url: "https://chatgpt.com/backend-api/wham/remote/control/server/pair"
|
| 310 |
+
.to_string(),
|
| 311 |
+
pair_status_url:
|
| 312 |
+
"https://chatgpt.com/backend-api/wham/remote/control/server/pair/status"
|
| 313 |
+
.to_string(),
|
| 314 |
+
}
|
| 315 |
+
);
|
| 316 |
+
assert_eq!(
|
| 317 |
+
normalize_remote_control_url("https://api.chatgpt-staging.com/backend-api")
|
| 318 |
+
.expect("chatgpt-staging.com subdomain URL should normalize"),
|
| 319 |
+
RemoteControlTarget {
|
| 320 |
+
websocket_url:
|
| 321 |
+
"wss://api.chatgpt-staging.com/backend-api/wham/remote/control/server"
|
| 322 |
+
.to_string(),
|
| 323 |
+
enroll_url:
|
| 324 |
+
"https://api.chatgpt-staging.com/backend-api/wham/remote/control/server/enroll"
|
| 325 |
+
.to_string(),
|
| 326 |
+
refresh_url:
|
| 327 |
+
"https://api.chatgpt-staging.com/backend-api/wham/remote/control/server/refresh"
|
| 328 |
+
.to_string(),
|
| 329 |
+
pair_url:
|
| 330 |
+
"https://api.chatgpt-staging.com/backend-api/wham/remote/control/server/pair"
|
| 331 |
+
.to_string(),
|
| 332 |
+
pair_status_url:
|
| 333 |
+
"https://api.chatgpt-staging.com/backend-api/wham/remote/control/server/pair/status"
|
| 334 |
+
.to_string(),
|
| 335 |
+
}
|
| 336 |
+
);
|
| 337 |
+
}
|
| 338 |
+
|
| 339 |
+
#[test]
|
| 340 |
+
fn normalize_remote_control_url_accepts_localhost_urls() {
|
| 341 |
+
assert_eq!(
|
| 342 |
+
normalize_remote_control_url("http://localhost:8080/backend-api")
|
| 343 |
+
.expect("localhost http URL should normalize"),
|
| 344 |
+
RemoteControlTarget {
|
| 345 |
+
websocket_url: "ws://localhost:8080/backend-api/wham/remote/control/server"
|
| 346 |
+
.to_string(),
|
| 347 |
+
enroll_url: "http://localhost:8080/backend-api/wham/remote/control/server/enroll"
|
| 348 |
+
.to_string(),
|
| 349 |
+
refresh_url: "http://localhost:8080/backend-api/wham/remote/control/server/refresh"
|
| 350 |
+
.to_string(),
|
| 351 |
+
pair_url: "http://localhost:8080/backend-api/wham/remote/control/server/pair"
|
| 352 |
+
.to_string(),
|
| 353 |
+
pair_status_url:
|
| 354 |
+
"http://localhost:8080/backend-api/wham/remote/control/server/pair/status"
|
| 355 |
+
.to_string(),
|
| 356 |
+
}
|
| 357 |
+
);
|
| 358 |
+
assert_eq!(
|
| 359 |
+
normalize_remote_control_url("https://localhost:8443/backend-api")
|
| 360 |
+
.expect("localhost https URL should normalize"),
|
| 361 |
+
RemoteControlTarget {
|
| 362 |
+
websocket_url: "wss://localhost:8443/backend-api/wham/remote/control/server"
|
| 363 |
+
.to_string(),
|
| 364 |
+
enroll_url: "https://localhost:8443/backend-api/wham/remote/control/server/enroll"
|
| 365 |
+
.to_string(),
|
| 366 |
+
refresh_url:
|
| 367 |
+
"https://localhost:8443/backend-api/wham/remote/control/server/refresh"
|
| 368 |
+
.to_string(),
|
| 369 |
+
pair_url: "https://localhost:8443/backend-api/wham/remote/control/server/pair"
|
| 370 |
+
.to_string(),
|
| 371 |
+
pair_status_url:
|
| 372 |
+
"https://localhost:8443/backend-api/wham/remote/control/server/pair/status"
|
| 373 |
+
.to_string(),
|
| 374 |
+
}
|
| 375 |
+
);
|
| 376 |
+
}
|
| 377 |
+
|
| 378 |
+
#[test]
|
| 379 |
+
fn normalize_remote_control_url_rejects_unsupported_urls() {
|
| 380 |
+
for remote_control_url in [
|
| 381 |
+
"http://chatgpt.com/backend-api",
|
| 382 |
+
"http://example.com/backend-api",
|
| 383 |
+
"https://example.com/backend-api",
|
| 384 |
+
"https://chat.openai.com/backend-api",
|
| 385 |
+
"https://chatgpt.com.evil.com/backend-api",
|
| 386 |
+
"https://evilchatgpt.com/backend-api",
|
| 387 |
+
"https://foo.localhost/backend-api",
|
| 388 |
+
] {
|
| 389 |
+
let err = normalize_remote_control_url(remote_control_url)
|
| 390 |
+
.expect_err("unsupported URL should be rejected");
|
| 391 |
+
|
| 392 |
+
assert_eq!(err.kind(), ErrorKind::InvalidInput);
|
| 393 |
+
assert_eq!(
|
| 394 |
+
err.to_string(),
|
| 395 |
+
format!(
|
| 396 |
+
"invalid remote control URL `{remote_control_url}`; expected HTTPS URL for chatgpt.com or chatgpt-staging.com, or HTTP/HTTPS URL for localhost"
|
| 397 |
+
)
|
| 398 |
+
);
|
| 399 |
+
}
|
| 400 |
+
}
|
| 401 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/segment.rs
ADDED
|
@@ -0,0 +1,469 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::protocol::ClientEnvelope;
|
| 2 |
+
use super::protocol::ClientEvent;
|
| 3 |
+
use super::protocol::ClientId;
|
| 4 |
+
use super::protocol::ServerEnvelope;
|
| 5 |
+
use super::protocol::ServerEvent;
|
| 6 |
+
use super::protocol::StreamId;
|
| 7 |
+
use crate::outgoing_message::OutgoingMessage;
|
| 8 |
+
use crate::transport::response_serialization_error;
|
| 9 |
+
use base64::DecodeSliceError;
|
| 10 |
+
use base64::Engine;
|
| 11 |
+
use codex_app_server_protocol::JSONRPCMessage;
|
| 12 |
+
use std::collections::HashMap;
|
| 13 |
+
use std::io;
|
| 14 |
+
use std::io::ErrorKind;
|
| 15 |
+
use std::io::Write;
|
| 16 |
+
use tokio::time::Instant;
|
| 17 |
+
use tracing::warn;
|
| 18 |
+
|
| 19 |
+
pub(super) const REMOTE_CONTROL_SEGMENT_TARGET_BYTES: usize = 100 * 1024;
|
| 20 |
+
pub(super) const REMOTE_CONTROL_SEGMENT_MAX_BYTES: usize = 150 * 1024;
|
| 21 |
+
pub(super) const REMOTE_CONTROL_REASSEMBLED_MAX_BYTES: usize = 100 * 1024 * 1024;
|
| 22 |
+
pub(super) const REMOTE_CONTROL_SEGMENT_COUNT_MAX: usize = 1024;
|
| 23 |
+
const REMOTE_CONTROL_SEGMENT_ASSEMBLY_MAX_COUNT: usize = 128;
|
| 24 |
+
|
| 25 |
+
#[derive(Debug)]
|
| 26 |
+
struct ClientSegmentAssembly {
|
| 27 |
+
stream_id: StreamId,
|
| 28 |
+
metadata: ClientSegmentMetadata,
|
| 29 |
+
raw: Vec<u8>,
|
| 30 |
+
next_segment_id: usize,
|
| 31 |
+
last_chunk_seen_at: Instant,
|
| 32 |
+
}
|
| 33 |
+
|
| 34 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 35 |
+
struct ClientSegmentMetadata {
|
| 36 |
+
seq_id: u64,
|
| 37 |
+
segment_count: usize,
|
| 38 |
+
message_size_bytes: usize,
|
| 39 |
+
}
|
| 40 |
+
|
| 41 |
+
#[derive(Default)]
|
| 42 |
+
pub(super) struct ClientSegmentReassembler {
|
| 43 |
+
assemblies: HashMap<ClientId, ClientSegmentAssembly>,
|
| 44 |
+
}
|
| 45 |
+
|
| 46 |
+
pub(super) enum ClientSegmentObservation {
|
| 47 |
+
Forward(Box<ClientEnvelope>),
|
| 48 |
+
Pending,
|
| 49 |
+
Dropped,
|
| 50 |
+
}
|
| 51 |
+
|
| 52 |
+
impl ClientSegmentReassembler {
|
| 53 |
+
pub(super) fn observe(&mut self, envelope: ClientEnvelope) -> ClientSegmentObservation {
|
| 54 |
+
let ClientEvent::ClientMessageChunk {
|
| 55 |
+
segment_id,
|
| 56 |
+
segment_count,
|
| 57 |
+
message_size_bytes,
|
| 58 |
+
message_chunk_base64,
|
| 59 |
+
} = &envelope.event
|
| 60 |
+
else {
|
| 61 |
+
return ClientSegmentObservation::Forward(Box::new(envelope));
|
| 62 |
+
};
|
| 63 |
+
let segment_id = *segment_id;
|
| 64 |
+
let segment_count = *segment_count;
|
| 65 |
+
let message_size_bytes = *message_size_bytes;
|
| 66 |
+
|
| 67 |
+
let Some(metadata) = ClientSegmentMetadata::from_envelope(&envelope) else {
|
| 68 |
+
warn!(
|
| 69 |
+
client_id = envelope.client_id.0.as_str(),
|
| 70 |
+
"dropping segmented remote-control client envelope without seq_id"
|
| 71 |
+
);
|
| 72 |
+
return ClientSegmentObservation::Dropped;
|
| 73 |
+
};
|
| 74 |
+
let Some(stream_id) = envelope.stream_id.clone() else {
|
| 75 |
+
warn!(
|
| 76 |
+
client_id = envelope.client_id.0.as_str(),
|
| 77 |
+
"dropping segmented remote-control client envelope without stream_id"
|
| 78 |
+
);
|
| 79 |
+
return ClientSegmentObservation::Dropped;
|
| 80 |
+
};
|
| 81 |
+
if self.should_ignore_chunk(&envelope.client_id, &stream_id, metadata.seq_id, segment_id) {
|
| 82 |
+
return ClientSegmentObservation::Dropped;
|
| 83 |
+
}
|
| 84 |
+
if segment_count == 0
|
| 85 |
+
|| segment_count > REMOTE_CONTROL_SEGMENT_COUNT_MAX
|
| 86 |
+
|| segment_id >= segment_count
|
| 87 |
+
|| message_size_bytes == 0
|
| 88 |
+
|| message_size_bytes > REMOTE_CONTROL_REASSEMBLED_MAX_BYTES
|
| 89 |
+
|| message_chunk_base64.is_empty()
|
| 90 |
+
{
|
| 91 |
+
warn!(
|
| 92 |
+
client_id = envelope.client_id.0.as_str(),
|
| 93 |
+
"dropping invalid segmented remote-control client envelope"
|
| 94 |
+
);
|
| 95 |
+
self.remove_assembly(&envelope.client_id, &stream_id);
|
| 96 |
+
return ClientSegmentObservation::Dropped;
|
| 97 |
+
}
|
| 98 |
+
|
| 99 |
+
let now = Instant::now();
|
| 100 |
+
match self.assemblies.get(&envelope.client_id) {
|
| 101 |
+
Some(assembly) if assembly.stream_id != stream_id => {
|
| 102 |
+
warn!(
|
| 103 |
+
client_id = envelope.client_id.0.as_str(),
|
| 104 |
+
"resetting segmented remote-control client envelope after stream change"
|
| 105 |
+
);
|
| 106 |
+
self.assemblies.insert(
|
| 107 |
+
envelope.client_id.clone(),
|
| 108 |
+
ClientSegmentAssembly {
|
| 109 |
+
stream_id: stream_id.clone(),
|
| 110 |
+
metadata: metadata.clone(),
|
| 111 |
+
raw: Vec::new(),
|
| 112 |
+
next_segment_id: 0,
|
| 113 |
+
last_chunk_seen_at: now,
|
| 114 |
+
},
|
| 115 |
+
);
|
| 116 |
+
}
|
| 117 |
+
Some(_) => {}
|
| 118 |
+
None => {
|
| 119 |
+
self.evict_assemblies_if_full();
|
| 120 |
+
self.assemblies.insert(
|
| 121 |
+
envelope.client_id.clone(),
|
| 122 |
+
ClientSegmentAssembly {
|
| 123 |
+
stream_id: stream_id.clone(),
|
| 124 |
+
metadata: metadata.clone(),
|
| 125 |
+
raw: Vec::new(),
|
| 126 |
+
next_segment_id: 0,
|
| 127 |
+
last_chunk_seen_at: now,
|
| 128 |
+
},
|
| 129 |
+
);
|
| 130 |
+
}
|
| 131 |
+
}
|
| 132 |
+
let result = {
|
| 133 |
+
let Some(assembly) = self.assemblies.get_mut(&envelope.client_id) else {
|
| 134 |
+
warn!(
|
| 135 |
+
client_id = envelope.client_id.0.as_str(),
|
| 136 |
+
"dropping segmented remote-control client envelope without assembly"
|
| 137 |
+
);
|
| 138 |
+
return ClientSegmentObservation::Dropped;
|
| 139 |
+
};
|
| 140 |
+
if metadata.seq_id < assembly.metadata.seq_id {
|
| 141 |
+
AssemblyUpdate::Ignore
|
| 142 |
+
} else if assembly.metadata != metadata {
|
| 143 |
+
warn!(
|
| 144 |
+
client_id = envelope.client_id.0.as_str(),
|
| 145 |
+
"resetting segmented remote-control client envelope after metadata mismatch"
|
| 146 |
+
);
|
| 147 |
+
AssemblyUpdate::Drop
|
| 148 |
+
} else if segment_id < assembly.next_segment_id {
|
| 149 |
+
AssemblyUpdate::Pending
|
| 150 |
+
} else if segment_id != assembly.next_segment_id {
|
| 151 |
+
warn!(
|
| 152 |
+
client_id = envelope.client_id.0.as_str(),
|
| 153 |
+
"dropping out-of-order segmented remote-control client envelope"
|
| 154 |
+
);
|
| 155 |
+
AssemblyUpdate::Drop
|
| 156 |
+
} else {
|
| 157 |
+
assembly.last_chunk_seen_at = now;
|
| 158 |
+
let chunk_start = assembly.raw.len();
|
| 159 |
+
let decoded_chunk_len = base64::decoded_len_estimate(message_chunk_base64.len());
|
| 160 |
+
let chunk_end = usize::min(
|
| 161 |
+
message_size_bytes,
|
| 162 |
+
chunk_start.saturating_add(decoded_chunk_len),
|
| 163 |
+
);
|
| 164 |
+
assembly.raw.resize(chunk_end, 0);
|
| 165 |
+
match base64::engine::general_purpose::STANDARD.decode_slice(
|
| 166 |
+
message_chunk_base64.as_bytes(),
|
| 167 |
+
&mut assembly.raw[chunk_start..],
|
| 168 |
+
) {
|
| 169 |
+
Ok(decoded_chunk_len) => {
|
| 170 |
+
assembly.raw.truncate(chunk_start + decoded_chunk_len);
|
| 171 |
+
assembly.next_segment_id += 1;
|
| 172 |
+
if assembly.next_segment_id < segment_count {
|
| 173 |
+
AssemblyUpdate::Pending
|
| 174 |
+
} else if assembly.raw.len() != message_size_bytes {
|
| 175 |
+
warn!(
|
| 176 |
+
client_id = envelope.client_id.0.as_str(),
|
| 177 |
+
"dropping reassembled remote-control client envelope with mismatched size"
|
| 178 |
+
);
|
| 179 |
+
AssemblyUpdate::Drop
|
| 180 |
+
} else {
|
| 181 |
+
match serde_json::from_slice::<JSONRPCMessage>(&assembly.raw) {
|
| 182 |
+
Ok(message) => AssemblyUpdate::Complete(message),
|
| 183 |
+
Err(err) => {
|
| 184 |
+
warn!(
|
| 185 |
+
client_id = envelope.client_id.0.as_str(),
|
| 186 |
+
"dropping invalid reassembled remote-control client envelope: {err}"
|
| 187 |
+
);
|
| 188 |
+
AssemblyUpdate::Drop
|
| 189 |
+
}
|
| 190 |
+
}
|
| 191 |
+
}
|
| 192 |
+
}
|
| 193 |
+
Err(DecodeSliceError::OutputSliceTooSmall) => {
|
| 194 |
+
warn!(
|
| 195 |
+
client_id = envelope.client_id.0.as_str(),
|
| 196 |
+
"dropping segmented remote-control client envelope after size overflow"
|
| 197 |
+
);
|
| 198 |
+
AssemblyUpdate::Drop
|
| 199 |
+
}
|
| 200 |
+
Err(err) => {
|
| 201 |
+
warn!(
|
| 202 |
+
client_id = envelope.client_id.0.as_str(),
|
| 203 |
+
"dropping segmented remote-control client envelope with invalid base64: {err}"
|
| 204 |
+
);
|
| 205 |
+
AssemblyUpdate::Drop
|
| 206 |
+
}
|
| 207 |
+
}
|
| 208 |
+
}
|
| 209 |
+
};
|
| 210 |
+
|
| 211 |
+
match result {
|
| 212 |
+
AssemblyUpdate::Pending => ClientSegmentObservation::Pending,
|
| 213 |
+
AssemblyUpdate::Ignore => ClientSegmentObservation::Dropped,
|
| 214 |
+
AssemblyUpdate::Drop => {
|
| 215 |
+
self.remove_assembly(&envelope.client_id, &stream_id);
|
| 216 |
+
ClientSegmentObservation::Dropped
|
| 217 |
+
}
|
| 218 |
+
AssemblyUpdate::Complete(message) => {
|
| 219 |
+
self.remove_assembly(&envelope.client_id, &stream_id);
|
| 220 |
+
ClientSegmentObservation::Forward(Box::new(ClientEnvelope {
|
| 221 |
+
event: ClientEvent::ClientMessage { message },
|
| 222 |
+
..envelope
|
| 223 |
+
}))
|
| 224 |
+
}
|
| 225 |
+
}
|
| 226 |
+
}
|
| 227 |
+
|
| 228 |
+
pub(super) fn invalidate_stream(&mut self, client_id: &ClientId, stream_id: &StreamId) {
|
| 229 |
+
self.remove_assembly(client_id, stream_id);
|
| 230 |
+
}
|
| 231 |
+
|
| 232 |
+
pub(super) fn invalidate_client(&mut self, client_id: &ClientId) {
|
| 233 |
+
self.assemblies.remove(client_id);
|
| 234 |
+
}
|
| 235 |
+
|
| 236 |
+
pub(super) fn should_ignore_chunk(
|
| 237 |
+
&self,
|
| 238 |
+
client_id: &ClientId,
|
| 239 |
+
stream_id: &StreamId,
|
| 240 |
+
seq_id: u64,
|
| 241 |
+
segment_id: usize,
|
| 242 |
+
) -> bool {
|
| 243 |
+
self.assemblies.get(client_id).is_some_and(|assembly| {
|
| 244 |
+
assembly.stream_id == *stream_id
|
| 245 |
+
&& (seq_id < assembly.metadata.seq_id
|
| 246 |
+
|| (seq_id == assembly.metadata.seq_id
|
| 247 |
+
&& segment_id < assembly.next_segment_id))
|
| 248 |
+
})
|
| 249 |
+
}
|
| 250 |
+
|
| 251 |
+
fn remove_assembly(&mut self, client_id: &ClientId, stream_id: &StreamId) {
|
| 252 |
+
if self
|
| 253 |
+
.assemblies
|
| 254 |
+
.get(client_id)
|
| 255 |
+
.is_some_and(|assembly| &assembly.stream_id == stream_id)
|
| 256 |
+
{
|
| 257 |
+
self.assemblies.remove(client_id);
|
| 258 |
+
}
|
| 259 |
+
}
|
| 260 |
+
|
| 261 |
+
fn evict_assemblies_if_full(&mut self) {
|
| 262 |
+
while self.assemblies.len() >= REMOTE_CONTROL_SEGMENT_ASSEMBLY_MAX_COUNT {
|
| 263 |
+
let Some(client_id) = self
|
| 264 |
+
.assemblies
|
| 265 |
+
.iter()
|
| 266 |
+
.min_by_key(|(_, assembly)| assembly.last_chunk_seen_at)
|
| 267 |
+
.map(|(client_id, _)| client_id.clone())
|
| 268 |
+
else {
|
| 269 |
+
return;
|
| 270 |
+
};
|
| 271 |
+
self.assemblies.remove(&client_id);
|
| 272 |
+
}
|
| 273 |
+
}
|
| 274 |
+
}
|
| 275 |
+
|
| 276 |
+
enum AssemblyUpdate {
|
| 277 |
+
Pending,
|
| 278 |
+
Ignore,
|
| 279 |
+
Drop,
|
| 280 |
+
Complete(JSONRPCMessage),
|
| 281 |
+
}
|
| 282 |
+
|
| 283 |
+
impl ClientSegmentMetadata {
|
| 284 |
+
fn from_envelope(envelope: &ClientEnvelope) -> Option<Self> {
|
| 285 |
+
let ClientEvent::ClientMessageChunk {
|
| 286 |
+
segment_count,
|
| 287 |
+
message_size_bytes,
|
| 288 |
+
..
|
| 289 |
+
} = &envelope.event
|
| 290 |
+
else {
|
| 291 |
+
return None;
|
| 292 |
+
};
|
| 293 |
+
Some(Self {
|
| 294 |
+
seq_id: envelope.seq_id?,
|
| 295 |
+
segment_count: *segment_count,
|
| 296 |
+
message_size_bytes: *message_size_bytes,
|
| 297 |
+
})
|
| 298 |
+
}
|
| 299 |
+
}
|
| 300 |
+
|
| 301 |
+
pub(super) fn split_server_envelope_for_transport(
|
| 302 |
+
envelope: ServerEnvelope,
|
| 303 |
+
) -> io::Result<Vec<ServerEnvelope>> {
|
| 304 |
+
if !matches!(envelope.event, ServerEvent::ServerMessage { .. }) {
|
| 305 |
+
return Ok(vec![envelope]);
|
| 306 |
+
}
|
| 307 |
+
|
| 308 |
+
let envelope_size_bytes = match serialized_len(&envelope) {
|
| 309 |
+
Ok(envelope_size_bytes) => envelope_size_bytes,
|
| 310 |
+
Err(err) => {
|
| 311 |
+
let ServerEvent::ServerMessage { message } = envelope.event else {
|
| 312 |
+
unreachable!("server message variant checked above");
|
| 313 |
+
};
|
| 314 |
+
let OutgoingMessage::Response(response) = *message else {
|
| 315 |
+
return Err(err);
|
| 316 |
+
};
|
| 317 |
+
return Ok(vec![ServerEnvelope {
|
| 318 |
+
event: ServerEvent::ServerMessage {
|
| 319 |
+
message: Box::new(response_serialization_error(response.id, err)),
|
| 320 |
+
},
|
| 321 |
+
client_id: envelope.client_id,
|
| 322 |
+
stream_id: envelope.stream_id,
|
| 323 |
+
seq_id: envelope.seq_id,
|
| 324 |
+
}]);
|
| 325 |
+
}
|
| 326 |
+
};
|
| 327 |
+
if envelope_size_bytes <= REMOTE_CONTROL_SEGMENT_MAX_BYTES {
|
| 328 |
+
return Ok(vec![envelope]);
|
| 329 |
+
}
|
| 330 |
+
|
| 331 |
+
let ServerEvent::ServerMessage { message } = envelope.event.clone() else {
|
| 332 |
+
unreachable!("server message variant checked above");
|
| 333 |
+
};
|
| 334 |
+
let raw = serde_json::to_vec(message.as_ref()).map_err(io::Error::other)?;
|
| 335 |
+
let message_size_bytes = raw.len();
|
| 336 |
+
if message_size_bytes > REMOTE_CONTROL_REASSEMBLED_MAX_BYTES {
|
| 337 |
+
warn!("dropping remote-control server envelope that exceeds reassembled size limit");
|
| 338 |
+
return Ok(Vec::new());
|
| 339 |
+
}
|
| 340 |
+
|
| 341 |
+
let minimal_segment_count =
|
| 342 |
+
usize::min(message_size_bytes.max(1), REMOTE_CONTROL_SEGMENT_COUNT_MAX);
|
| 343 |
+
let minimal_chunk = &raw[..usize::min(raw.len(), 1)];
|
| 344 |
+
if serialized_chunk_len(
|
| 345 |
+
&envelope,
|
| 346 |
+
/*segment_id*/ 0,
|
| 347 |
+
minimal_segment_count,
|
| 348 |
+
message_size_bytes,
|
| 349 |
+
minimal_chunk,
|
| 350 |
+
)? > REMOTE_CONTROL_SEGMENT_MAX_BYTES
|
| 351 |
+
{
|
| 352 |
+
warn!("dropping remote-control server envelope that cannot fit within segment size limit");
|
| 353 |
+
return Ok(Vec::new());
|
| 354 |
+
}
|
| 355 |
+
|
| 356 |
+
let mut segment_count = usize::max(
|
| 357 |
+
2,
|
| 358 |
+
message_size_bytes.div_ceil(REMOTE_CONTROL_SEGMENT_TARGET_BYTES),
|
| 359 |
+
);
|
| 360 |
+
loop {
|
| 361 |
+
let chunk_size = usize::max(1, message_size_bytes.div_ceil(segment_count));
|
| 362 |
+
segment_count = message_size_bytes.div_ceil(chunk_size);
|
| 363 |
+
let segments_fit = raw
|
| 364 |
+
.chunks(chunk_size)
|
| 365 |
+
.enumerate()
|
| 366 |
+
.all(|(segment_id, chunk)| {
|
| 367 |
+
serialized_chunk_len(
|
| 368 |
+
&envelope,
|
| 369 |
+
segment_id,
|
| 370 |
+
segment_count,
|
| 371 |
+
message_size_bytes,
|
| 372 |
+
chunk,
|
| 373 |
+
)
|
| 374 |
+
.is_ok_and(|size| size <= REMOTE_CONTROL_SEGMENT_MAX_BYTES)
|
| 375 |
+
});
|
| 376 |
+
if segments_fit {
|
| 377 |
+
return raw
|
| 378 |
+
.chunks(chunk_size)
|
| 379 |
+
.enumerate()
|
| 380 |
+
.map(|(segment_id, chunk)| {
|
| 381 |
+
build_chunk_envelope(
|
| 382 |
+
&envelope,
|
| 383 |
+
segment_id,
|
| 384 |
+
segment_count,
|
| 385 |
+
message_size_bytes,
|
| 386 |
+
chunk,
|
| 387 |
+
)
|
| 388 |
+
})
|
| 389 |
+
.collect();
|
| 390 |
+
}
|
| 391 |
+
if chunk_size == 1 {
|
| 392 |
+
warn!(
|
| 393 |
+
"dropping remote-control server envelope that cannot fit within segment size limit"
|
| 394 |
+
);
|
| 395 |
+
return Ok(Vec::new());
|
| 396 |
+
}
|
| 397 |
+
let next_segment_count = segment_count + 1;
|
| 398 |
+
let next_chunk_size = usize::max(1, message_size_bytes.div_ceil(next_segment_count));
|
| 399 |
+
segment_count = if next_chunk_size == chunk_size {
|
| 400 |
+
message_size_bytes
|
| 401 |
+
} else {
|
| 402 |
+
next_segment_count
|
| 403 |
+
};
|
| 404 |
+
}
|
| 405 |
+
}
|
| 406 |
+
|
| 407 |
+
fn serialized_chunk_len(
|
| 408 |
+
envelope: &ServerEnvelope,
|
| 409 |
+
segment_id: usize,
|
| 410 |
+
segment_count: usize,
|
| 411 |
+
message_size_bytes: usize,
|
| 412 |
+
chunk: &[u8],
|
| 413 |
+
) -> io::Result<usize> {
|
| 414 |
+
serialized_len(&build_chunk_envelope(
|
| 415 |
+
envelope,
|
| 416 |
+
segment_id,
|
| 417 |
+
segment_count,
|
| 418 |
+
message_size_bytes,
|
| 419 |
+
chunk,
|
| 420 |
+
)?)
|
| 421 |
+
}
|
| 422 |
+
|
| 423 |
+
#[derive(Default)]
|
| 424 |
+
struct CountingWriter {
|
| 425 |
+
len: usize,
|
| 426 |
+
}
|
| 427 |
+
|
| 428 |
+
impl Write for CountingWriter {
|
| 429 |
+
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
|
| 430 |
+
self.len += buf.len();
|
| 431 |
+
Ok(buf.len())
|
| 432 |
+
}
|
| 433 |
+
|
| 434 |
+
fn flush(&mut self) -> io::Result<()> {
|
| 435 |
+
Ok(())
|
| 436 |
+
}
|
| 437 |
+
}
|
| 438 |
+
|
| 439 |
+
fn serialized_len(value: &impl serde::Serialize) -> io::Result<usize> {
|
| 440 |
+
let mut writer = CountingWriter::default();
|
| 441 |
+
serde_json::to_writer(&mut writer, value).map_err(io::Error::other)?;
|
| 442 |
+
Ok(writer.len)
|
| 443 |
+
}
|
| 444 |
+
|
| 445 |
+
fn build_chunk_envelope(
|
| 446 |
+
envelope: &ServerEnvelope,
|
| 447 |
+
segment_id: usize,
|
| 448 |
+
segment_count: usize,
|
| 449 |
+
message_size_bytes: usize,
|
| 450 |
+
chunk: &[u8],
|
| 451 |
+
) -> io::Result<ServerEnvelope> {
|
| 452 |
+
if segment_count > REMOTE_CONTROL_SEGMENT_COUNT_MAX {
|
| 453 |
+
return Err(io::Error::new(
|
| 454 |
+
ErrorKind::InvalidData,
|
| 455 |
+
"remote-control segment count exceeds maximum",
|
| 456 |
+
));
|
| 457 |
+
}
|
| 458 |
+
Ok(ServerEnvelope {
|
| 459 |
+
event: ServerEvent::ServerMessageChunk {
|
| 460 |
+
segment_id,
|
| 461 |
+
segment_count,
|
| 462 |
+
message_size_bytes,
|
| 463 |
+
message_chunk_base64: base64::engine::general_purpose::STANDARD.encode(chunk),
|
| 464 |
+
},
|
| 465 |
+
client_id: envelope.client_id.clone(),
|
| 466 |
+
stream_id: envelope.stream_id.clone(),
|
| 467 |
+
seq_id: envelope.seq_id,
|
| 468 |
+
})
|
| 469 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/segment_tests.rs
ADDED
|
@@ -0,0 +1,450 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::protocol::ClientEnvelope;
|
| 2 |
+
use super::protocol::ClientEvent;
|
| 3 |
+
use super::protocol::ClientId;
|
| 4 |
+
use super::protocol::ServerEnvelope;
|
| 5 |
+
use super::protocol::ServerEvent;
|
| 6 |
+
use super::protocol::StreamId;
|
| 7 |
+
use super::segment::ClientSegmentObservation;
|
| 8 |
+
use super::segment::ClientSegmentReassembler;
|
| 9 |
+
use super::segment::REMOTE_CONTROL_SEGMENT_MAX_BYTES;
|
| 10 |
+
use super::segment::split_server_envelope_for_transport;
|
| 11 |
+
use crate::outgoing_message::OutgoingMessage;
|
| 12 |
+
#[cfg(unix)]
|
| 13 |
+
use crate::outgoing_message::OutgoingResponse;
|
| 14 |
+
use base64::Engine;
|
| 15 |
+
#[cfg(unix)]
|
| 16 |
+
use codex_app_server_protocol::ClientResponsePayload;
|
| 17 |
+
use codex_app_server_protocol::ConfigWarningNotification;
|
| 18 |
+
#[cfg(unix)]
|
| 19 |
+
use codex_app_server_protocol::InitializeResponse;
|
| 20 |
+
use codex_app_server_protocol::JSONRPCMessage;
|
| 21 |
+
use codex_app_server_protocol::JSONRPCNotification;
|
| 22 |
+
#[cfg(unix)]
|
| 23 |
+
use codex_app_server_protocol::RequestId;
|
| 24 |
+
use codex_app_server_protocol::ServerNotification;
|
| 25 |
+
use codex_app_server_protocol::ServerNotificationEnvelope;
|
| 26 |
+
#[cfg(unix)]
|
| 27 |
+
use codex_utils_absolute_path::AbsolutePathBuf;
|
| 28 |
+
use pretty_assertions::assert_eq;
|
| 29 |
+
#[cfg(unix)]
|
| 30 |
+
use serde_json::json;
|
| 31 |
+
|
| 32 |
+
#[test]
|
| 33 |
+
fn reassembles_client_message_chunks() {
|
| 34 |
+
let message = JSONRPCMessage::Notification(JSONRPCNotification {
|
| 35 |
+
method: "initialized".to_string(),
|
| 36 |
+
params: None,
|
| 37 |
+
});
|
| 38 |
+
let raw = serde_json::to_vec(&message).expect("message should serialize");
|
| 39 |
+
let split = raw.len() / 2;
|
| 40 |
+
let client_id = ClientId("client-1".to_string());
|
| 41 |
+
let stream_id = Some(StreamId("stream-1".to_string()));
|
| 42 |
+
let mut reassembler = ClientSegmentReassembler::default();
|
| 43 |
+
|
| 44 |
+
assert!(matches!(
|
| 45 |
+
reassembler.observe(chunk_envelope(
|
| 46 |
+
client_id.clone(),
|
| 47 |
+
stream_id.clone(),
|
| 48 |
+
/*seq_id*/ 7,
|
| 49 |
+
/*segment_id*/ 0,
|
| 50 |
+
/*segment_count*/ 2,
|
| 51 |
+
raw.len(),
|
| 52 |
+
&raw[..split],
|
| 53 |
+
)),
|
| 54 |
+
ClientSegmentObservation::Pending
|
| 55 |
+
));
|
| 56 |
+
let reassembled = match reassembler.observe(chunk_envelope(
|
| 57 |
+
client_id.clone(),
|
| 58 |
+
stream_id,
|
| 59 |
+
/*seq_id*/ 7,
|
| 60 |
+
/*segment_id*/ 1,
|
| 61 |
+
/*segment_count*/ 2,
|
| 62 |
+
raw.len(),
|
| 63 |
+
&raw[split..],
|
| 64 |
+
)) {
|
| 65 |
+
ClientSegmentObservation::Forward(reassembled) => *reassembled,
|
| 66 |
+
ClientSegmentObservation::Pending | ClientSegmentObservation::Dropped => {
|
| 67 |
+
panic!("message should reassemble")
|
| 68 |
+
}
|
| 69 |
+
};
|
| 70 |
+
assert_eq!(reassembled.client_id, client_id);
|
| 71 |
+
assert_eq!(
|
| 72 |
+
reassembled.stream_id,
|
| 73 |
+
Some(StreamId("stream-1".to_string()))
|
| 74 |
+
);
|
| 75 |
+
assert_eq!(reassembled.seq_id, Some(7));
|
| 76 |
+
assert_eq!(reassembled.cursor, None);
|
| 77 |
+
match reassembled.event {
|
| 78 |
+
ClientEvent::ClientMessage {
|
| 79 |
+
message: reassembled_message,
|
| 80 |
+
} => assert_eq!(reassembled_message, message),
|
| 81 |
+
other => panic!("expected client message, got {other:?}"),
|
| 82 |
+
}
|
| 83 |
+
}
|
| 84 |
+
|
| 85 |
+
#[test]
|
| 86 |
+
fn splits_large_server_messages_into_wire_chunks() {
|
| 87 |
+
let envelope = ServerEnvelope {
|
| 88 |
+
event: ServerEvent::ServerMessage {
|
| 89 |
+
message: Box::new(OutgoingMessage::AppServerNotification(
|
| 90 |
+
ServerNotificationEnvelope {
|
| 91 |
+
notification: ServerNotification::ConfigWarning(ConfigWarningNotification {
|
| 92 |
+
summary: "x".repeat(REMOTE_CONTROL_SEGMENT_MAX_BYTES),
|
| 93 |
+
details: None,
|
| 94 |
+
path: None,
|
| 95 |
+
range: None,
|
| 96 |
+
}),
|
| 97 |
+
emitted_at_ms: Some(1_234),
|
| 98 |
+
},
|
| 99 |
+
)),
|
| 100 |
+
},
|
| 101 |
+
client_id: ClientId("client-1".to_string()),
|
| 102 |
+
stream_id: StreamId("stream-1".to_string()),
|
| 103 |
+
seq_id: 9,
|
| 104 |
+
};
|
| 105 |
+
|
| 106 |
+
let segments = split_server_envelope_for_transport(envelope).expect("split should succeed");
|
| 107 |
+
|
| 108 |
+
assert!(segments.len() > 1);
|
| 109 |
+
assert!(
|
| 110 |
+
segments
|
| 111 |
+
.iter()
|
| 112 |
+
.all(|segment| matches!(segment.event, ServerEvent::ServerMessageChunk { .. }))
|
| 113 |
+
);
|
| 114 |
+
assert!(segments.iter().all(|segment| segment.seq_id == 9));
|
| 115 |
+
assert!(segments.iter().all(|segment| {
|
| 116 |
+
serde_json::to_vec(segment)
|
| 117 |
+
.expect("segment should serialize")
|
| 118 |
+
.len()
|
| 119 |
+
<= REMOTE_CONTROL_SEGMENT_MAX_BYTES
|
| 120 |
+
}));
|
| 121 |
+
}
|
| 122 |
+
|
| 123 |
+
#[cfg(unix)]
|
| 124 |
+
#[test]
|
| 125 |
+
fn invalid_response_becomes_remote_control_jsonrpc_error() {
|
| 126 |
+
use std::ffi::OsString;
|
| 127 |
+
use std::os::unix::ffi::OsStringExt;
|
| 128 |
+
use std::path::PathBuf;
|
| 129 |
+
|
| 130 |
+
let codex_home = AbsolutePathBuf::from_absolute_path(PathBuf::from(OsString::from_vec(vec![
|
| 131 |
+
b'/', b'b', b'a', b'd', 0xff,
|
| 132 |
+
])))
|
| 133 |
+
.expect("non-UTF-8 Unix paths are valid absolute paths");
|
| 134 |
+
let envelope = ServerEnvelope {
|
| 135 |
+
event: ServerEvent::ServerMessage {
|
| 136 |
+
message: Box::new(OutgoingMessage::Response(OutgoingResponse {
|
| 137 |
+
id: RequestId::Integer(7),
|
| 138 |
+
result: Box::new(ClientResponsePayload::Initialize(InitializeResponse {
|
| 139 |
+
user_agent: "codex-test-agent".to_string(),
|
| 140 |
+
codex_home,
|
| 141 |
+
platform_family: "unix".to_string(),
|
| 142 |
+
platform_os: "linux".to_string(),
|
| 143 |
+
})),
|
| 144 |
+
})),
|
| 145 |
+
},
|
| 146 |
+
client_id: ClientId("client-1".to_string()),
|
| 147 |
+
stream_id: StreamId("stream-1".to_string()),
|
| 148 |
+
seq_id: 9,
|
| 149 |
+
};
|
| 150 |
+
|
| 151 |
+
let envelopes = split_server_envelope_for_transport(envelope)
|
| 152 |
+
.expect("invalid response should become a remote-control JSON-RPC error");
|
| 153 |
+
assert_eq!(
|
| 154 |
+
serde_json::to_value(envelopes).expect("error envelope should serialize"),
|
| 155 |
+
json!([{
|
| 156 |
+
"type": "server_message",
|
| 157 |
+
"client_id": "client-1",
|
| 158 |
+
"stream_id": "stream-1",
|
| 159 |
+
"seq_id": 9,
|
| 160 |
+
"message": {
|
| 161 |
+
"id": 7,
|
| 162 |
+
"error": {
|
| 163 |
+
"code": -32603,
|
| 164 |
+
"message": "failed to serialize response: path contains invalid UTF-8 characters",
|
| 165 |
+
}
|
| 166 |
+
}
|
| 167 |
+
}])
|
| 168 |
+
);
|
| 169 |
+
}
|
| 170 |
+
|
| 171 |
+
#[test]
|
| 172 |
+
fn invalidates_incomplete_stream_assemblies() {
|
| 173 |
+
let message = JSONRPCMessage::Notification(JSONRPCNotification {
|
| 174 |
+
method: "initialized".to_string(),
|
| 175 |
+
params: None,
|
| 176 |
+
});
|
| 177 |
+
let raw = serde_json::to_vec(&message).expect("message should serialize");
|
| 178 |
+
let split = raw.len() / 2;
|
| 179 |
+
let client_id = ClientId("client-1".to_string());
|
| 180 |
+
let stream_id = StreamId("stream-1".to_string());
|
| 181 |
+
let mut reassembler = ClientSegmentReassembler::default();
|
| 182 |
+
|
| 183 |
+
assert!(matches!(
|
| 184 |
+
reassembler.observe(chunk_envelope(
|
| 185 |
+
client_id.clone(),
|
| 186 |
+
Some(stream_id.clone()),
|
| 187 |
+
/*seq_id*/ 7,
|
| 188 |
+
/*segment_id*/ 0,
|
| 189 |
+
/*segment_count*/ 2,
|
| 190 |
+
raw.len(),
|
| 191 |
+
&raw[..split],
|
| 192 |
+
)),
|
| 193 |
+
ClientSegmentObservation::Pending
|
| 194 |
+
));
|
| 195 |
+
reassembler.invalidate_stream(&client_id, &stream_id);
|
| 196 |
+
assert!(matches!(
|
| 197 |
+
reassembler.observe(chunk_envelope(
|
| 198 |
+
client_id,
|
| 199 |
+
Some(stream_id),
|
| 200 |
+
/*seq_id*/ 7,
|
| 201 |
+
/*segment_id*/ 1,
|
| 202 |
+
/*segment_count*/ 2,
|
| 203 |
+
raw.len(),
|
| 204 |
+
&raw[split..],
|
| 205 |
+
)),
|
| 206 |
+
ClientSegmentObservation::Dropped
|
| 207 |
+
));
|
| 208 |
+
}
|
| 209 |
+
|
| 210 |
+
#[test]
|
| 211 |
+
fn resets_incomplete_client_assembly_when_stream_changes() {
|
| 212 |
+
let message = JSONRPCMessage::Notification(JSONRPCNotification {
|
| 213 |
+
method: "initialized".to_string(),
|
| 214 |
+
params: None,
|
| 215 |
+
});
|
| 216 |
+
let raw = serde_json::to_vec(&message).expect("message should serialize");
|
| 217 |
+
let split = raw.len() / 2;
|
| 218 |
+
let client_id = ClientId("client-1".to_string());
|
| 219 |
+
let first_stream_id = StreamId("stream-1".to_string());
|
| 220 |
+
let second_stream_id = StreamId("stream-2".to_string());
|
| 221 |
+
let mut reassembler = ClientSegmentReassembler::default();
|
| 222 |
+
|
| 223 |
+
assert!(matches!(
|
| 224 |
+
reassembler.observe(chunk_envelope(
|
| 225 |
+
client_id.clone(),
|
| 226 |
+
Some(first_stream_id.clone()),
|
| 227 |
+
/*seq_id*/ 7,
|
| 228 |
+
/*segment_id*/ 0,
|
| 229 |
+
/*segment_count*/ 2,
|
| 230 |
+
raw.len(),
|
| 231 |
+
&raw[..split],
|
| 232 |
+
)),
|
| 233 |
+
ClientSegmentObservation::Pending
|
| 234 |
+
));
|
| 235 |
+
assert!(matches!(
|
| 236 |
+
reassembler.observe(chunk_envelope(
|
| 237 |
+
client_id.clone(),
|
| 238 |
+
Some(second_stream_id.clone()),
|
| 239 |
+
/*seq_id*/ 8,
|
| 240 |
+
/*segment_id*/ 0,
|
| 241 |
+
/*segment_count*/ 2,
|
| 242 |
+
raw.len(),
|
| 243 |
+
&raw[..split],
|
| 244 |
+
)),
|
| 245 |
+
ClientSegmentObservation::Pending
|
| 246 |
+
));
|
| 247 |
+
let reassembled = match reassembler.observe(chunk_envelope(
|
| 248 |
+
client_id.clone(),
|
| 249 |
+
Some(second_stream_id),
|
| 250 |
+
/*seq_id*/ 8,
|
| 251 |
+
/*segment_id*/ 1,
|
| 252 |
+
/*segment_count*/ 2,
|
| 253 |
+
raw.len(),
|
| 254 |
+
&raw[split..],
|
| 255 |
+
)) {
|
| 256 |
+
ClientSegmentObservation::Forward(reassembled) => *reassembled,
|
| 257 |
+
ClientSegmentObservation::Pending | ClientSegmentObservation::Dropped => {
|
| 258 |
+
panic!("replacement stream should reassemble")
|
| 259 |
+
}
|
| 260 |
+
};
|
| 261 |
+
assert_eq!(
|
| 262 |
+
reassembled.stream_id,
|
| 263 |
+
Some(StreamId("stream-2".to_string()))
|
| 264 |
+
);
|
| 265 |
+
assert!(matches!(
|
| 266 |
+
reassembler.observe(chunk_envelope(
|
| 267 |
+
client_id,
|
| 268 |
+
Some(first_stream_id),
|
| 269 |
+
/*seq_id*/ 7,
|
| 270 |
+
/*segment_id*/ 1,
|
| 271 |
+
/*segment_count*/ 2,
|
| 272 |
+
raw.len(),
|
| 273 |
+
&raw[split..],
|
| 274 |
+
)),
|
| 275 |
+
ClientSegmentObservation::Dropped
|
| 276 |
+
));
|
| 277 |
+
}
|
| 278 |
+
|
| 279 |
+
#[test]
|
| 280 |
+
fn ignores_stale_chunks_without_dropping_newer_assembly() {
|
| 281 |
+
let message = JSONRPCMessage::Notification(JSONRPCNotification {
|
| 282 |
+
method: "initialized".to_string(),
|
| 283 |
+
params: None,
|
| 284 |
+
});
|
| 285 |
+
let raw = serde_json::to_vec(&message).expect("message should serialize");
|
| 286 |
+
let split = raw.len() / 2;
|
| 287 |
+
let client_id = ClientId("client-1".to_string());
|
| 288 |
+
let stream_id = Some(StreamId("stream-1".to_string()));
|
| 289 |
+
let mut reassembler = ClientSegmentReassembler::default();
|
| 290 |
+
|
| 291 |
+
assert!(matches!(
|
| 292 |
+
reassembler.observe(chunk_envelope(
|
| 293 |
+
client_id.clone(),
|
| 294 |
+
stream_id.clone(),
|
| 295 |
+
/*seq_id*/ 8,
|
| 296 |
+
/*segment_id*/ 0,
|
| 297 |
+
/*segment_count*/ 2,
|
| 298 |
+
raw.len(),
|
| 299 |
+
&raw[..split],
|
| 300 |
+
)),
|
| 301 |
+
ClientSegmentObservation::Pending
|
| 302 |
+
));
|
| 303 |
+
assert!(matches!(
|
| 304 |
+
reassembler.observe(chunk_envelope(
|
| 305 |
+
client_id.clone(),
|
| 306 |
+
stream_id.clone(),
|
| 307 |
+
/*seq_id*/ 7,
|
| 308 |
+
/*segment_id*/ 0,
|
| 309 |
+
/*segment_count*/ 2,
|
| 310 |
+
raw.len(),
|
| 311 |
+
&raw[..split],
|
| 312 |
+
)),
|
| 313 |
+
ClientSegmentObservation::Dropped
|
| 314 |
+
));
|
| 315 |
+
assert!(matches!(
|
| 316 |
+
reassembler.observe(chunk_envelope(
|
| 317 |
+
client_id,
|
| 318 |
+
stream_id,
|
| 319 |
+
/*seq_id*/ 8,
|
| 320 |
+
/*segment_id*/ 1,
|
| 321 |
+
/*segment_count*/ 2,
|
| 322 |
+
raw.len(),
|
| 323 |
+
&raw[split..],
|
| 324 |
+
)),
|
| 325 |
+
ClientSegmentObservation::Forward(_)
|
| 326 |
+
));
|
| 327 |
+
}
|
| 328 |
+
|
| 329 |
+
#[test]
|
| 330 |
+
fn ignores_invalid_stale_chunks_without_dropping_newer_assembly() {
|
| 331 |
+
let message = JSONRPCMessage::Notification(JSONRPCNotification {
|
| 332 |
+
method: "initialized".to_string(),
|
| 333 |
+
params: None,
|
| 334 |
+
});
|
| 335 |
+
let raw = serde_json::to_vec(&message).expect("message should serialize");
|
| 336 |
+
let split = raw.len() / 2;
|
| 337 |
+
let client_id = ClientId("client-1".to_string());
|
| 338 |
+
let stream_id = Some(StreamId("stream-1".to_string()));
|
| 339 |
+
let mut reassembler = ClientSegmentReassembler::default();
|
| 340 |
+
|
| 341 |
+
assert!(matches!(
|
| 342 |
+
reassembler.observe(chunk_envelope(
|
| 343 |
+
client_id.clone(),
|
| 344 |
+
stream_id.clone(),
|
| 345 |
+
/*seq_id*/ 8,
|
| 346 |
+
/*segment_id*/ 0,
|
| 347 |
+
/*segment_count*/ 2,
|
| 348 |
+
raw.len(),
|
| 349 |
+
&raw[..split],
|
| 350 |
+
)),
|
| 351 |
+
ClientSegmentObservation::Pending
|
| 352 |
+
));
|
| 353 |
+
assert!(matches!(
|
| 354 |
+
reassembler.observe(chunk_envelope(
|
| 355 |
+
client_id.clone(),
|
| 356 |
+
stream_id.clone(),
|
| 357 |
+
/*seq_id*/ 7,
|
| 358 |
+
/*segment_id*/ 1,
|
| 359 |
+
/*segment_count*/ 2,
|
| 360 |
+
raw.len(),
|
| 361 |
+
b"",
|
| 362 |
+
)),
|
| 363 |
+
ClientSegmentObservation::Dropped
|
| 364 |
+
));
|
| 365 |
+
assert!(matches!(
|
| 366 |
+
reassembler.observe(chunk_envelope(
|
| 367 |
+
client_id,
|
| 368 |
+
stream_id,
|
| 369 |
+
/*seq_id*/ 8,
|
| 370 |
+
/*segment_id*/ 1,
|
| 371 |
+
/*segment_count*/ 2,
|
| 372 |
+
raw.len(),
|
| 373 |
+
&raw[split..],
|
| 374 |
+
)),
|
| 375 |
+
ClientSegmentObservation::Forward(_)
|
| 376 |
+
));
|
| 377 |
+
}
|
| 378 |
+
|
| 379 |
+
#[test]
|
| 380 |
+
fn ignores_invalid_duplicate_chunks_without_dropping_current_assembly() {
|
| 381 |
+
let message = JSONRPCMessage::Notification(JSONRPCNotification {
|
| 382 |
+
method: "initialized".to_string(),
|
| 383 |
+
params: None,
|
| 384 |
+
});
|
| 385 |
+
let raw = serde_json::to_vec(&message).expect("message should serialize");
|
| 386 |
+
let split = raw.len() / 2;
|
| 387 |
+
let client_id = ClientId("client-1".to_string());
|
| 388 |
+
let stream_id = Some(StreamId("stream-1".to_string()));
|
| 389 |
+
let mut reassembler = ClientSegmentReassembler::default();
|
| 390 |
+
|
| 391 |
+
assert!(matches!(
|
| 392 |
+
reassembler.observe(chunk_envelope(
|
| 393 |
+
client_id.clone(),
|
| 394 |
+
stream_id.clone(),
|
| 395 |
+
/*seq_id*/ 8,
|
| 396 |
+
/*segment_id*/ 0,
|
| 397 |
+
/*segment_count*/ 2,
|
| 398 |
+
raw.len(),
|
| 399 |
+
&raw[..split],
|
| 400 |
+
)),
|
| 401 |
+
ClientSegmentObservation::Pending
|
| 402 |
+
));
|
| 403 |
+
assert!(matches!(
|
| 404 |
+
reassembler.observe(chunk_envelope(
|
| 405 |
+
client_id.clone(),
|
| 406 |
+
stream_id.clone(),
|
| 407 |
+
/*seq_id*/ 8,
|
| 408 |
+
/*segment_id*/ 0,
|
| 409 |
+
/*segment_count*/ 2,
|
| 410 |
+
raw.len(),
|
| 411 |
+
b"",
|
| 412 |
+
)),
|
| 413 |
+
ClientSegmentObservation::Dropped
|
| 414 |
+
));
|
| 415 |
+
assert!(matches!(
|
| 416 |
+
reassembler.observe(chunk_envelope(
|
| 417 |
+
client_id,
|
| 418 |
+
stream_id,
|
| 419 |
+
/*seq_id*/ 8,
|
| 420 |
+
/*segment_id*/ 1,
|
| 421 |
+
/*segment_count*/ 2,
|
| 422 |
+
raw.len(),
|
| 423 |
+
&raw[split..],
|
| 424 |
+
)),
|
| 425 |
+
ClientSegmentObservation::Forward(_)
|
| 426 |
+
));
|
| 427 |
+
}
|
| 428 |
+
|
| 429 |
+
fn chunk_envelope(
|
| 430 |
+
client_id: ClientId,
|
| 431 |
+
stream_id: Option<StreamId>,
|
| 432 |
+
seq_id: u64,
|
| 433 |
+
segment_id: usize,
|
| 434 |
+
segment_count: usize,
|
| 435 |
+
message_size_bytes: usize,
|
| 436 |
+
chunk: &[u8],
|
| 437 |
+
) -> ClientEnvelope {
|
| 438 |
+
ClientEnvelope {
|
| 439 |
+
event: ClientEvent::ClientMessageChunk {
|
| 440 |
+
segment_id,
|
| 441 |
+
segment_count,
|
| 442 |
+
message_size_bytes,
|
| 443 |
+
message_chunk_base64: base64::engine::general_purpose::STANDARD.encode(chunk),
|
| 444 |
+
},
|
| 445 |
+
client_id,
|
| 446 |
+
stream_id,
|
| 447 |
+
seq_id: Some(seq_id),
|
| 448 |
+
cursor: None,
|
| 449 |
+
}
|
| 450 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/server_api.rs
ADDED
|
@@ -0,0 +1,382 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::auth::RemoteControlConnectionAuth;
|
| 2 |
+
use super::enroll::RemoteControlEnrollment;
|
| 3 |
+
use super::enroll::RemoteControlServerTokenRefreshRequirement;
|
| 4 |
+
use super::enroll::format_headers;
|
| 5 |
+
use super::enroll::preview_remote_control_response_body;
|
| 6 |
+
use super::protocol::EnrollRemoteServerRequest;
|
| 7 |
+
use super::protocol::EnrollRemoteServerResponse;
|
| 8 |
+
use super::protocol::RefreshRemoteServerRequest;
|
| 9 |
+
use super::protocol::RemoteControlTarget;
|
| 10 |
+
use axum::http::HeaderMap;
|
| 11 |
+
use axum::http::StatusCode;
|
| 12 |
+
use codex_login::default_client::create_client_without_request_logging;
|
| 13 |
+
use rand::Rng;
|
| 14 |
+
use serde::Serialize;
|
| 15 |
+
use serde::de::DeserializeOwned;
|
| 16 |
+
use std::fmt;
|
| 17 |
+
use std::io;
|
| 18 |
+
use std::io::ErrorKind;
|
| 19 |
+
use std::time::Duration;
|
| 20 |
+
use time::OffsetDateTime;
|
| 21 |
+
use time::format_description::well_known::Rfc3339;
|
| 22 |
+
use tracing::warn;
|
| 23 |
+
|
| 24 |
+
const REMOTE_CONTROL_RETRY_AFTER_JITTER_MAX_MILLIS: u64 = 30_000;
|
| 25 |
+
|
| 26 |
+
const REMOTE_CONTROL_ENROLL_TIMEOUT: Duration = Duration::from_secs(30);
|
| 27 |
+
const REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MIN_SECS: u64 = 24;
|
| 28 |
+
const REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MAX_SECS: u64 = 36;
|
| 29 |
+
|
| 30 |
+
pub(super) const REMOTE_CONTROL_INSTALLATION_ID_HEADER: &str = "x-codex-installation-id";
|
| 31 |
+
|
| 32 |
+
#[derive(Debug)]
|
| 33 |
+
pub(super) struct RemoteControlServerRequestError {
|
| 34 |
+
message: String,
|
| 35 |
+
status: Option<StatusCode>,
|
| 36 |
+
retry_at: Option<OffsetDateTime>,
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
impl RemoteControlServerRequestError {
|
| 40 |
+
pub(super) fn retry_deferred(retry_at: OffsetDateTime) -> io::Error {
|
| 41 |
+
io::Error::new(
|
| 42 |
+
ErrorKind::WouldBlock,
|
| 43 |
+
Self {
|
| 44 |
+
message: format!("remote control retry deferred until {retry_at}"),
|
| 45 |
+
status: None,
|
| 46 |
+
retry_at: Some(retry_at),
|
| 47 |
+
},
|
| 48 |
+
)
|
| 49 |
+
}
|
| 50 |
+
|
| 51 |
+
pub(super) fn io_error(
|
| 52 |
+
message: String,
|
| 53 |
+
status: Option<StatusCode>,
|
| 54 |
+
retry_at: Option<OffsetDateTime>,
|
| 55 |
+
timed_out: bool,
|
| 56 |
+
) -> io::Error {
|
| 57 |
+
let kind = match status {
|
| 58 |
+
Some(StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN) => ErrorKind::PermissionDenied,
|
| 59 |
+
Some(StatusCode::NOT_FOUND) => ErrorKind::NotFound,
|
| 60 |
+
Some(status) if timed_out && !status.is_client_error() => ErrorKind::TimedOut,
|
| 61 |
+
None if timed_out => ErrorKind::TimedOut,
|
| 62 |
+
Some(_) | None => ErrorKind::Other,
|
| 63 |
+
};
|
| 64 |
+
io::Error::new(
|
| 65 |
+
kind,
|
| 66 |
+
Self {
|
| 67 |
+
message,
|
| 68 |
+
status,
|
| 69 |
+
retry_at,
|
| 70 |
+
},
|
| 71 |
+
)
|
| 72 |
+
}
|
| 73 |
+
|
| 74 |
+
fn is_transient(&self, kind: ErrorKind) -> bool {
|
| 75 |
+
kind == ErrorKind::TimedOut
|
| 76 |
+
|| self.status.is_none()
|
| 77 |
+
|| self.status.is_some_and(|status| {
|
| 78 |
+
status == StatusCode::TOO_MANY_REQUESTS || status.is_server_error()
|
| 79 |
+
})
|
| 80 |
+
}
|
| 81 |
+
}
|
| 82 |
+
|
| 83 |
+
impl fmt::Display for RemoteControlServerRequestError {
|
| 84 |
+
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
| 85 |
+
formatter.write_str(&self.message)
|
| 86 |
+
}
|
| 87 |
+
}
|
| 88 |
+
|
| 89 |
+
impl std::error::Error for RemoteControlServerRequestError {}
|
| 90 |
+
|
| 91 |
+
pub(super) async fn enroll_remote_control_server(
|
| 92 |
+
remote_control_target: &RemoteControlTarget,
|
| 93 |
+
auth: &RemoteControlConnectionAuth,
|
| 94 |
+
installation_id: &str,
|
| 95 |
+
server_name: &str,
|
| 96 |
+
) -> io::Result<RemoteControlEnrollment> {
|
| 97 |
+
let enroll_url = &remote_control_target.enroll_url;
|
| 98 |
+
let request = EnrollRemoteServerRequest {
|
| 99 |
+
name: server_name.to_string(),
|
| 100 |
+
os: std::env::consts::OS,
|
| 101 |
+
arch: std::env::consts::ARCH,
|
| 102 |
+
app_server_version: env!("CARGO_PKG_VERSION"),
|
| 103 |
+
installation_id: installation_id.to_string(),
|
| 104 |
+
};
|
| 105 |
+
let enrollment_response = send_remote_control_server_request::<_, EnrollRemoteServerResponse>(
|
| 106 |
+
enroll_url,
|
| 107 |
+
auth,
|
| 108 |
+
installation_id,
|
| 109 |
+
&request,
|
| 110 |
+
"enroll",
|
| 111 |
+
"server enrollment",
|
| 112 |
+
REMOTE_CONTROL_ENROLL_TIMEOUT,
|
| 113 |
+
)
|
| 114 |
+
.await?;
|
| 115 |
+
let mut enrollment = RemoteControlEnrollment {
|
| 116 |
+
remote_control_target: remote_control_target.clone(),
|
| 117 |
+
account_id: auth.account_id.clone(),
|
| 118 |
+
environment_id: enrollment_response.environment_id,
|
| 119 |
+
server_id: enrollment_response.server_id,
|
| 120 |
+
server_name: server_name.to_string(),
|
| 121 |
+
remote_control_token: None,
|
| 122 |
+
expires_at: None,
|
| 123 |
+
next_refresh_at: None,
|
| 124 |
+
};
|
| 125 |
+
update_remote_control_server_token(
|
| 126 |
+
&mut enrollment,
|
| 127 |
+
enroll_url,
|
| 128 |
+
enrollment_response.remote_control_token,
|
| 129 |
+
enrollment_response.expires_at,
|
| 130 |
+
)?;
|
| 131 |
+
Ok(enrollment)
|
| 132 |
+
}
|
| 133 |
+
|
| 134 |
+
pub(super) async fn refresh_remote_control_server(
|
| 135 |
+
auth: &RemoteControlConnectionAuth,
|
| 136 |
+
installation_id: &str,
|
| 137 |
+
enrollment: &mut RemoteControlEnrollment,
|
| 138 |
+
) -> io::Result<()> {
|
| 139 |
+
let now = OffsetDateTime::now_utc();
|
| 140 |
+
let refresh_requirement = enrollment.server_token_refresh_requirement_at(now);
|
| 141 |
+
if refresh_requirement == RemoteControlServerTokenRefreshRequirement::NotNeeded {
|
| 142 |
+
return Ok(());
|
| 143 |
+
}
|
| 144 |
+
if refresh_requirement == RemoteControlServerTokenRefreshRequirement::Required
|
| 145 |
+
&& let Some(next_refresh_at) = enrollment.next_refresh_at
|
| 146 |
+
&& next_refresh_at > now
|
| 147 |
+
{
|
| 148 |
+
return Err(RemoteControlServerRequestError::retry_deferred(
|
| 149 |
+
next_refresh_at,
|
| 150 |
+
));
|
| 151 |
+
}
|
| 152 |
+
let refresh_url = enrollment.remote_control_target.refresh_url.clone();
|
| 153 |
+
let request = RefreshRemoteServerRequest {
|
| 154 |
+
server_id: enrollment.server_id.clone(),
|
| 155 |
+
installation_id: installation_id.to_string(),
|
| 156 |
+
};
|
| 157 |
+
let refreshed = match send_remote_control_server_request::<_, EnrollRemoteServerResponse>(
|
| 158 |
+
&refresh_url,
|
| 159 |
+
auth,
|
| 160 |
+
installation_id,
|
| 161 |
+
&request,
|
| 162 |
+
"refresh",
|
| 163 |
+
"server refresh",
|
| 164 |
+
REMOTE_CONTROL_ENROLL_TIMEOUT,
|
| 165 |
+
)
|
| 166 |
+
.await
|
| 167 |
+
{
|
| 168 |
+
Ok(refreshed) => refreshed,
|
| 169 |
+
Err(err) => {
|
| 170 |
+
let Some(refresh_error) = remote_control_server_request_error(&err) else {
|
| 171 |
+
return Err(err);
|
| 172 |
+
};
|
| 173 |
+
if !refresh_error.is_transient(err.kind()) {
|
| 174 |
+
return Err(err);
|
| 175 |
+
}
|
| 176 |
+
let now = OffsetDateTime::now_utc();
|
| 177 |
+
let refresh_is_required = enrollment.server_token_refresh_requirement_at(now)
|
| 178 |
+
== RemoteControlServerTokenRefreshRequirement::Required;
|
| 179 |
+
let (refresh_delay, next_refresh_at) = refresh_deferral(refresh_error.retry_at, now);
|
| 180 |
+
enrollment.next_refresh_at = Some(next_refresh_at);
|
| 181 |
+
// An explicit server deadline takes precedence over valid-token fallback.
|
| 182 |
+
// Keep the token, but let callers defer new requests until this deadline.
|
| 183 |
+
if refresh_is_required
|
| 184 |
+
|| refresh_error
|
| 185 |
+
.retry_at
|
| 186 |
+
.is_some_and(|retry_at| retry_at > now)
|
| 187 |
+
{
|
| 188 |
+
warn!(
|
| 189 |
+
refresh_url,
|
| 190 |
+
server_id = %enrollment.server_id,
|
| 191 |
+
environment_id = %enrollment.environment_id,
|
| 192 |
+
error = %err,
|
| 193 |
+
?refresh_delay,
|
| 194 |
+
%next_refresh_at,
|
| 195 |
+
"remote control server token refresh failed; deferring new requests"
|
| 196 |
+
);
|
| 197 |
+
return Err(err);
|
| 198 |
+
}
|
| 199 |
+
warn!(
|
| 200 |
+
refresh_url,
|
| 201 |
+
server_id = %enrollment.server_id,
|
| 202 |
+
environment_id = %enrollment.environment_id,
|
| 203 |
+
error = %err,
|
| 204 |
+
?refresh_delay,
|
| 205 |
+
%next_refresh_at,
|
| 206 |
+
"proactive remote control server token refresh failed; continuing with valid token"
|
| 207 |
+
);
|
| 208 |
+
return Ok(());
|
| 209 |
+
}
|
| 210 |
+
};
|
| 211 |
+
if refreshed.server_id != enrollment.server_id
|
| 212 |
+
|| refreshed.environment_id != enrollment.environment_id
|
| 213 |
+
{
|
| 214 |
+
return Err(io::Error::other(format!(
|
| 215 |
+
"remote control server refresh returned mismatched enrollment: expected server_id={}, environment_id={}; got server_id={}, environment_id={}",
|
| 216 |
+
enrollment.server_id,
|
| 217 |
+
enrollment.environment_id,
|
| 218 |
+
refreshed.server_id,
|
| 219 |
+
refreshed.environment_id
|
| 220 |
+
)));
|
| 221 |
+
}
|
| 222 |
+
|
| 223 |
+
update_remote_control_server_token(
|
| 224 |
+
enrollment,
|
| 225 |
+
&refresh_url,
|
| 226 |
+
refreshed.remote_control_token,
|
| 227 |
+
refreshed.expires_at,
|
| 228 |
+
)
|
| 229 |
+
}
|
| 230 |
+
|
| 231 |
+
async fn send_remote_control_server_request<Request, Response>(
|
| 232 |
+
url: &str,
|
| 233 |
+
auth: &RemoteControlConnectionAuth,
|
| 234 |
+
installation_id: &str,
|
| 235 |
+
request: &Request,
|
| 236 |
+
action: &str,
|
| 237 |
+
response_kind: &str,
|
| 238 |
+
timeout: Duration,
|
| 239 |
+
) -> io::Result<Response>
|
| 240 |
+
where
|
| 241 |
+
Request: Serialize,
|
| 242 |
+
Response: DeserializeOwned,
|
| 243 |
+
{
|
| 244 |
+
let client = create_client_without_request_logging();
|
| 245 |
+
let auth_headers = auth.request_headers()?;
|
| 246 |
+
let response = client
|
| 247 |
+
.post(url)
|
| 248 |
+
.timeout(timeout)
|
| 249 |
+
.headers(auth_headers)
|
| 250 |
+
.header(REMOTE_CONTROL_INSTALLATION_ID_HEADER, installation_id)
|
| 251 |
+
.json(request)
|
| 252 |
+
.send()
|
| 253 |
+
.await
|
| 254 |
+
.map_err(|err| {
|
| 255 |
+
let timed_out = err.is_timeout();
|
| 256 |
+
RemoteControlServerRequestError::io_error(
|
| 257 |
+
format!("failed to {action} remote control server at `{url}`: {err}"),
|
| 258 |
+
/*status*/ None,
|
| 259 |
+
/*retry_at*/ None,
|
| 260 |
+
timed_out,
|
| 261 |
+
)
|
| 262 |
+
})?;
|
| 263 |
+
let headers = response.headers().clone();
|
| 264 |
+
let status = response.status();
|
| 265 |
+
let retry_at = retry_after_with_jitter(&headers, OffsetDateTime::now_utc());
|
| 266 |
+
let body = response.bytes().await.map_err(|err| {
|
| 267 |
+
let timed_out = err.is_timeout();
|
| 268 |
+
RemoteControlServerRequestError::io_error(
|
| 269 |
+
format!("failed to read remote control {response_kind} response from `{url}`: {err}"),
|
| 270 |
+
Some(status),
|
| 271 |
+
retry_at,
|
| 272 |
+
timed_out,
|
| 273 |
+
)
|
| 274 |
+
})?;
|
| 275 |
+
let body_preview = preview_remote_control_response_body(&body);
|
| 276 |
+
if !status.is_success() {
|
| 277 |
+
let headers_str = format_headers(&headers);
|
| 278 |
+
return Err(RemoteControlServerRequestError::io_error(
|
| 279 |
+
format!(
|
| 280 |
+
"remote control {response_kind} failed at `{url}`: HTTP {status}, {headers_str}, body: {body_preview}"
|
| 281 |
+
),
|
| 282 |
+
Some(status),
|
| 283 |
+
retry_at,
|
| 284 |
+
/*timed_out*/ false,
|
| 285 |
+
));
|
| 286 |
+
}
|
| 287 |
+
|
| 288 |
+
serde_json::from_slice::<Response>(&body).map_err(|err| {
|
| 289 |
+
let headers_str = format_headers(&headers);
|
| 290 |
+
io::Error::other(format!(
|
| 291 |
+
"failed to parse remote control {response_kind} response from `{url}`: HTTP {status}, {headers_str}, body: {body_preview}, decode error: {err}"
|
| 292 |
+
))
|
| 293 |
+
})
|
| 294 |
+
}
|
| 295 |
+
|
| 296 |
+
fn update_remote_control_server_token(
|
| 297 |
+
enrollment: &mut RemoteControlEnrollment,
|
| 298 |
+
url: &str,
|
| 299 |
+
token: String,
|
| 300 |
+
expires_at: String,
|
| 301 |
+
) -> io::Result<()> {
|
| 302 |
+
let expires_at = OffsetDateTime::parse(&expires_at, &Rfc3339).map_err(|err| {
|
| 303 |
+
io::Error::other(format!(
|
| 304 |
+
"failed to parse remote control server token expiry from `{url}`: {err}"
|
| 305 |
+
))
|
| 306 |
+
})?;
|
| 307 |
+
enrollment.remote_control_token = Some(token);
|
| 308 |
+
enrollment.expires_at = Some(expires_at);
|
| 309 |
+
enrollment.next_refresh_at = None;
|
| 310 |
+
Ok(())
|
| 311 |
+
}
|
| 312 |
+
|
| 313 |
+
fn remote_control_server_request_error(
|
| 314 |
+
err: &io::Error,
|
| 315 |
+
) -> Option<&RemoteControlServerRequestError> {
|
| 316 |
+
err.get_ref()?.downcast_ref()
|
| 317 |
+
}
|
| 318 |
+
|
| 319 |
+
pub(super) fn remote_control_retry_at(err: &io::Error) -> Option<OffsetDateTime> {
|
| 320 |
+
let request_error = remote_control_server_request_error(err)?;
|
| 321 |
+
if !request_error.is_transient(err.kind()) {
|
| 322 |
+
return None;
|
| 323 |
+
}
|
| 324 |
+
request_error.retry_at
|
| 325 |
+
}
|
| 326 |
+
|
| 327 |
+
pub(super) fn remote_control_retry_delay(err: &io::Error) -> Option<Duration> {
|
| 328 |
+
Duration::try_from(remote_control_retry_at(err)? - OffsetDateTime::now_utc()).ok()
|
| 329 |
+
}
|
| 330 |
+
|
| 331 |
+
pub(super) fn retry_after_with_jitter(
|
| 332 |
+
headers: &HeaderMap,
|
| 333 |
+
received_at: OffsetDateTime,
|
| 334 |
+
) -> Option<OffsetDateTime> {
|
| 335 |
+
let retry_at = parse_retry_after(headers, received_at)?;
|
| 336 |
+
// Keep the server's deadline as a lower bound, then spread clients across
|
| 337 |
+
// the next 30 seconds. Sample once per response so retries share a deadline.
|
| 338 |
+
let jitter = time::Duration::milliseconds(
|
| 339 |
+
rand::rng().random_range(0..=REMOTE_CONTROL_RETRY_AFTER_JITTER_MAX_MILLIS) as i64,
|
| 340 |
+
);
|
| 341 |
+
Some(retry_at.checked_add(jitter).unwrap_or(retry_at))
|
| 342 |
+
}
|
| 343 |
+
|
| 344 |
+
fn parse_retry_after(headers: &HeaderMap, received_at: OffsetDateTime) -> Option<OffsetDateTime> {
|
| 345 |
+
let retry_after = headers
|
| 346 |
+
.get(axum::http::header::RETRY_AFTER)?
|
| 347 |
+
.to_str()
|
| 348 |
+
.ok()?;
|
| 349 |
+
let retry_at = if let Ok(seconds) = retry_after.parse::<u64>() {
|
| 350 |
+
let seconds = i64::try_from(seconds).ok()?;
|
| 351 |
+
received_at.checked_add(time::Duration::seconds(seconds))?
|
| 352 |
+
} else {
|
| 353 |
+
OffsetDateTime::from(httpdate::parse_http_date(retry_after).ok()?)
|
| 354 |
+
};
|
| 355 |
+
(retry_at >= received_at).then_some(retry_at)
|
| 356 |
+
}
|
| 357 |
+
|
| 358 |
+
fn refresh_deferral(
|
| 359 |
+
retry_at: Option<OffsetDateTime>,
|
| 360 |
+
now: OffsetDateTime,
|
| 361 |
+
) -> (Duration, OffsetDateTime) {
|
| 362 |
+
if let Some(retry_at) = retry_at
|
| 363 |
+
&& let Ok(delay) = Duration::try_from(retry_at - now)
|
| 364 |
+
&& !delay.is_zero()
|
| 365 |
+
{
|
| 366 |
+
return (delay, retry_at);
|
| 367 |
+
}
|
| 368 |
+
let delay = remote_control_server_token_refresh_backoff();
|
| 369 |
+
let next_refresh_at = now + time::Duration::seconds(delay.as_secs() as i64);
|
| 370 |
+
(delay, next_refresh_at)
|
| 371 |
+
}
|
| 372 |
+
|
| 373 |
+
fn remote_control_server_token_refresh_backoff() -> Duration {
|
| 374 |
+
Duration::from_secs(rand::rng().random_range(
|
| 375 |
+
REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MIN_SECS
|
| 376 |
+
..=REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MAX_SECS,
|
| 377 |
+
))
|
| 378 |
+
}
|
| 379 |
+
|
| 380 |
+
#[cfg(test)]
|
| 381 |
+
#[path = "server_api_tests.rs"]
|
| 382 |
+
mod tests;
|
codex-rs/app-server-transport/src/transport/remote_control/server_api_tests.rs
ADDED
|
@@ -0,0 +1,321 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::*;
|
| 2 |
+
use crate::transport::remote_control::protocol::normalize_remote_control_url;
|
| 3 |
+
use pretty_assertions::assert_eq;
|
| 4 |
+
use serde_json::json;
|
| 5 |
+
use std::time::SystemTime;
|
| 6 |
+
use tokio::io::AsyncWriteExt;
|
| 7 |
+
use tokio::net::TcpListener;
|
| 8 |
+
use tokio::sync::oneshot;
|
| 9 |
+
|
| 10 |
+
const TEST_REQUEST_TIMEOUT: Duration = Duration::from_millis(100);
|
| 11 |
+
|
| 12 |
+
fn auth() -> RemoteControlConnectionAuth {
|
| 13 |
+
RemoteControlConnectionAuth {
|
| 14 |
+
auth_provider: codex_model_provider::unauthenticated_auth_provider(),
|
| 15 |
+
account_id: "account-a".to_string(),
|
| 16 |
+
}
|
| 17 |
+
}
|
| 18 |
+
|
| 19 |
+
fn assert_transient_timeout(err: &io::Error, expected_status: Option<StatusCode>) {
|
| 20 |
+
let request_error = remote_control_server_request_error(err)
|
| 21 |
+
.expect("request error should preserve refresh metadata");
|
| 22 |
+
assert_eq!(
|
| 23 |
+
(
|
| 24 |
+
err.kind(),
|
| 25 |
+
request_error.status,
|
| 26 |
+
request_error.is_transient(err.kind()),
|
| 27 |
+
),
|
| 28 |
+
(ErrorKind::TimedOut, expected_status, true)
|
| 29 |
+
);
|
| 30 |
+
}
|
| 31 |
+
|
| 32 |
+
async fn timed_out_request(partial_response: Option<&'static [u8]>) -> io::Error {
|
| 33 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 34 |
+
.await
|
| 35 |
+
.expect("listener should bind");
|
| 36 |
+
let url = format!(
|
| 37 |
+
"http://{}/backend-api/wham/remote/control/server/refresh",
|
| 38 |
+
listener
|
| 39 |
+
.local_addr()
|
| 40 |
+
.expect("listener should have a local address")
|
| 41 |
+
);
|
| 42 |
+
let (request_done_tx, request_done_rx) = oneshot::channel();
|
| 43 |
+
let server_task = tokio::spawn(async move {
|
| 44 |
+
let (mut stream, _) = listener.accept().await.expect("request should connect");
|
| 45 |
+
if let Some(partial_response) = partial_response {
|
| 46 |
+
stream
|
| 47 |
+
.write_all(partial_response)
|
| 48 |
+
.await
|
| 49 |
+
.expect("partial response should write");
|
| 50 |
+
}
|
| 51 |
+
request_done_rx
|
| 52 |
+
.await
|
| 53 |
+
.expect("test should report request completion");
|
| 54 |
+
});
|
| 55 |
+
|
| 56 |
+
let err = send_remote_control_server_request::<_, serde_json::Value>(
|
| 57 |
+
&url,
|
| 58 |
+
&auth(),
|
| 59 |
+
"installation-id",
|
| 60 |
+
&json!({"server_id": "server-id"}),
|
| 61 |
+
"refresh",
|
| 62 |
+
"server refresh",
|
| 63 |
+
TEST_REQUEST_TIMEOUT,
|
| 64 |
+
)
|
| 65 |
+
.await
|
| 66 |
+
.expect_err("incomplete response should time out");
|
| 67 |
+
request_done_tx
|
| 68 |
+
.send(())
|
| 69 |
+
.expect("server should wait for request completion");
|
| 70 |
+
server_task.await.expect("server task should finish");
|
| 71 |
+
err
|
| 72 |
+
}
|
| 73 |
+
|
| 74 |
+
fn enrollment(now: OffsetDateTime) -> RemoteControlEnrollment {
|
| 75 |
+
RemoteControlEnrollment {
|
| 76 |
+
remote_control_target: normalize_remote_control_url("http://localhost/backend-api/")
|
| 77 |
+
.expect("target should normalize"),
|
| 78 |
+
account_id: "account-a".to_string(),
|
| 79 |
+
environment_id: "env_first".to_string(),
|
| 80 |
+
server_id: "srv_e_first".to_string(),
|
| 81 |
+
server_name: "first-server".to_string(),
|
| 82 |
+
remote_control_token: Some("token".to_string()),
|
| 83 |
+
expires_at: Some(now + time::Duration::seconds(300)),
|
| 84 |
+
next_refresh_at: None,
|
| 85 |
+
}
|
| 86 |
+
}
|
| 87 |
+
|
| 88 |
+
#[test]
|
| 89 |
+
fn remote_control_enrollment_classifies_server_token_refresh_requirement() {
|
| 90 |
+
let now =
|
| 91 |
+
OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse");
|
| 92 |
+
let enrollment = enrollment(now);
|
| 93 |
+
let cases = [
|
| 94 |
+
(
|
| 95 |
+
enrollment.clone(),
|
| 96 |
+
RemoteControlServerTokenRefreshRequirement::Proactive,
|
| 97 |
+
),
|
| 98 |
+
(
|
| 99 |
+
RemoteControlEnrollment {
|
| 100 |
+
expires_at: Some(now + time::Duration::seconds(301)),
|
| 101 |
+
..enrollment.clone()
|
| 102 |
+
},
|
| 103 |
+
RemoteControlServerTokenRefreshRequirement::NotNeeded,
|
| 104 |
+
),
|
| 105 |
+
(
|
| 106 |
+
RemoteControlEnrollment {
|
| 107 |
+
next_refresh_at: Some(now + time::Duration::seconds(30)),
|
| 108 |
+
..enrollment.clone()
|
| 109 |
+
},
|
| 110 |
+
RemoteControlServerTokenRefreshRequirement::NotNeeded,
|
| 111 |
+
),
|
| 112 |
+
(
|
| 113 |
+
RemoteControlEnrollment {
|
| 114 |
+
next_refresh_at: Some(now),
|
| 115 |
+
..enrollment.clone()
|
| 116 |
+
},
|
| 117 |
+
RemoteControlServerTokenRefreshRequirement::Proactive,
|
| 118 |
+
),
|
| 119 |
+
(
|
| 120 |
+
RemoteControlEnrollment {
|
| 121 |
+
remote_control_token: None,
|
| 122 |
+
..enrollment.clone()
|
| 123 |
+
},
|
| 124 |
+
RemoteControlServerTokenRefreshRequirement::Required,
|
| 125 |
+
),
|
| 126 |
+
(
|
| 127 |
+
RemoteControlEnrollment {
|
| 128 |
+
expires_at: None,
|
| 129 |
+
..enrollment.clone()
|
| 130 |
+
},
|
| 131 |
+
RemoteControlServerTokenRefreshRequirement::Required,
|
| 132 |
+
),
|
| 133 |
+
(
|
| 134 |
+
RemoteControlEnrollment {
|
| 135 |
+
expires_at: Some(now),
|
| 136 |
+
next_refresh_at: Some(now + time::Duration::hours(1)),
|
| 137 |
+
..enrollment
|
| 138 |
+
},
|
| 139 |
+
RemoteControlServerTokenRefreshRequirement::Required,
|
| 140 |
+
),
|
| 141 |
+
];
|
| 142 |
+
|
| 143 |
+
for (enrollment, expected) in cases {
|
| 144 |
+
assert_eq!(
|
| 145 |
+
enrollment.server_token_refresh_requirement_at(now),
|
| 146 |
+
expected
|
| 147 |
+
);
|
| 148 |
+
}
|
| 149 |
+
}
|
| 150 |
+
|
| 151 |
+
#[test]
|
| 152 |
+
fn remote_control_server_request_error_classifies_status_before_timeout() {
|
| 153 |
+
let cases = [
|
| 154 |
+
(None, true, ErrorKind::TimedOut, true),
|
| 155 |
+
(Some(StatusCode::OK), true, ErrorKind::TimedOut, true),
|
| 156 |
+
(
|
| 157 |
+
Some(StatusCode::TOO_MANY_REQUESTS),
|
| 158 |
+
false,
|
| 159 |
+
ErrorKind::Other,
|
| 160 |
+
true,
|
| 161 |
+
),
|
| 162 |
+
(Some(StatusCode::BAD_GATEWAY), false, ErrorKind::Other, true),
|
| 163 |
+
(
|
| 164 |
+
Some(StatusCode::UNAUTHORIZED),
|
| 165 |
+
true,
|
| 166 |
+
ErrorKind::PermissionDenied,
|
| 167 |
+
false,
|
| 168 |
+
),
|
| 169 |
+
(
|
| 170 |
+
Some(StatusCode::FORBIDDEN),
|
| 171 |
+
true,
|
| 172 |
+
ErrorKind::PermissionDenied,
|
| 173 |
+
false,
|
| 174 |
+
),
|
| 175 |
+
(
|
| 176 |
+
Some(StatusCode::NOT_FOUND),
|
| 177 |
+
true,
|
| 178 |
+
ErrorKind::NotFound,
|
| 179 |
+
false,
|
| 180 |
+
),
|
| 181 |
+
(Some(StatusCode::BAD_REQUEST), true, ErrorKind::Other, false),
|
| 182 |
+
(None, false, ErrorKind::Other, true),
|
| 183 |
+
];
|
| 184 |
+
|
| 185 |
+
for (status, timed_out, expected_kind, expected_transient) in cases {
|
| 186 |
+
let err = RemoteControlServerRequestError::io_error(
|
| 187 |
+
String::new(),
|
| 188 |
+
status,
|
| 189 |
+
/*retry_at*/ None,
|
| 190 |
+
timed_out,
|
| 191 |
+
);
|
| 192 |
+
let request_error = remote_control_server_request_error(&err)
|
| 193 |
+
.expect("request error should preserve refresh metadata");
|
| 194 |
+
assert_eq!(
|
| 195 |
+
(err.kind(), request_error.is_transient(err.kind())),
|
| 196 |
+
(expected_kind, expected_transient)
|
| 197 |
+
);
|
| 198 |
+
}
|
| 199 |
+
}
|
| 200 |
+
|
| 201 |
+
#[tokio::test]
|
| 202 |
+
async fn request_timeout_before_response_headers_is_transient() {
|
| 203 |
+
let err = timed_out_request(/*partial_response*/ None).await;
|
| 204 |
+
assert_transient_timeout(&err, /*expected_status*/ None);
|
| 205 |
+
}
|
| 206 |
+
|
| 207 |
+
#[tokio::test]
|
| 208 |
+
async fn response_body_timeout_is_transient() {
|
| 209 |
+
let err = timed_out_request(Some(b"HTTP/1.1 200 OK\r\nContent-Length: 20\r\n\r\n{")).await;
|
| 210 |
+
assert_transient_timeout(&err, Some(StatusCode::OK));
|
| 211 |
+
}
|
| 212 |
+
|
| 213 |
+
#[test]
|
| 214 |
+
fn retry_after_supports_delta_seconds_and_http_dates() {
|
| 215 |
+
let now =
|
| 216 |
+
OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse");
|
| 217 |
+
let mut headers = HeaderMap::new();
|
| 218 |
+
headers.insert(
|
| 219 |
+
axum::http::header::RETRY_AFTER,
|
| 220 |
+
axum::http::HeaderValue::from_static("120"),
|
| 221 |
+
);
|
| 222 |
+
assert_eq!(
|
| 223 |
+
parse_retry_after(&headers, now),
|
| 224 |
+
Some(now + time::Duration::seconds(120))
|
| 225 |
+
);
|
| 226 |
+
|
| 227 |
+
let retry_at = now + time::Duration::seconds(90);
|
| 228 |
+
let retry_at_system = SystemTime::UNIX_EPOCH + Duration::from_secs(1_700_000_090);
|
| 229 |
+
headers.insert(
|
| 230 |
+
axum::http::header::RETRY_AFTER,
|
| 231 |
+
httpdate::fmt_http_date(retry_at_system)
|
| 232 |
+
.parse()
|
| 233 |
+
.expect("HTTP date should be a valid header value"),
|
| 234 |
+
);
|
| 235 |
+
assert_eq!(parse_retry_after(&headers, now), Some(retry_at));
|
| 236 |
+
}
|
| 237 |
+
|
| 238 |
+
#[test]
|
| 239 |
+
fn invalid_or_expired_retry_after_uses_bounded_fallback() {
|
| 240 |
+
let now =
|
| 241 |
+
OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse");
|
| 242 |
+
let mut headers = HeaderMap::new();
|
| 243 |
+
headers.insert(
|
| 244 |
+
axum::http::header::RETRY_AFTER,
|
| 245 |
+
axum::http::HeaderValue::from_static("invalid"),
|
| 246 |
+
);
|
| 247 |
+
assert_eq!(parse_retry_after(&headers, now), None);
|
| 248 |
+
|
| 249 |
+
headers.insert(
|
| 250 |
+
axum::http::header::RETRY_AFTER,
|
| 251 |
+
httpdate::fmt_http_date(SystemTime::UNIX_EPOCH + Duration::from_secs(1_699_999_999))
|
| 252 |
+
.parse()
|
| 253 |
+
.expect("HTTP date should be a valid header value"),
|
| 254 |
+
);
|
| 255 |
+
assert_eq!(parse_retry_after(&headers, now), None);
|
| 256 |
+
|
| 257 |
+
let expired_while_reading_body = Some(now + time::Duration::seconds(1));
|
| 258 |
+
for retry_at in [None, expired_while_reading_body] {
|
| 259 |
+
let deferred_at = now + time::Duration::seconds(2);
|
| 260 |
+
let (delay, next_refresh_at) = refresh_deferral(retry_at, deferred_at);
|
| 261 |
+
assert!(
|
| 262 |
+
(Duration::from_secs(REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MIN_SECS)
|
| 263 |
+
..=Duration::from_secs(REMOTE_CONTROL_SERVER_TOKEN_REFRESH_BACKOFF_MAX_SECS,))
|
| 264 |
+
.contains(&delay)
|
| 265 |
+
);
|
| 266 |
+
assert_eq!(
|
| 267 |
+
next_refresh_at,
|
| 268 |
+
deferred_at + time::Duration::seconds(delay.as_secs() as i64)
|
| 269 |
+
);
|
| 270 |
+
}
|
| 271 |
+
}
|
| 272 |
+
|
| 273 |
+
#[test]
|
| 274 |
+
fn http_date_retry_after_preserves_absolute_deadline() {
|
| 275 |
+
let received_at =
|
| 276 |
+
OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse");
|
| 277 |
+
let retry_at = received_at + time::Duration::seconds(120);
|
| 278 |
+
let body_read_at = received_at + time::Duration::seconds(30);
|
| 279 |
+
|
| 280 |
+
assert_eq!(
|
| 281 |
+
refresh_deferral(Some(retry_at), body_read_at),
|
| 282 |
+
(Duration::from_secs(90), retry_at)
|
| 283 |
+
);
|
| 284 |
+
}
|
| 285 |
+
|
| 286 |
+
#[test]
|
| 287 |
+
fn retry_after_jitter_never_shortens_the_server_deadline() {
|
| 288 |
+
let now =
|
| 289 |
+
OffsetDateTime::from_unix_timestamp(1_700_000_000).expect("test timestamp should parse");
|
| 290 |
+
for value in ["120", "Tue, 14 Nov 2023 22:15:20 GMT"] {
|
| 291 |
+
let mut headers = HeaderMap::new();
|
| 292 |
+
headers.insert(
|
| 293 |
+
axum::http::header::RETRY_AFTER,
|
| 294 |
+
value.parse().expect("retry hint should be a valid header"),
|
| 295 |
+
);
|
| 296 |
+
let deadline = retry_after_with_jitter(&headers, now)
|
| 297 |
+
.expect("a valid retry hint should produce a deadline");
|
| 298 |
+
assert!(
|
| 299 |
+
(now + time::Duration::seconds(120)..=now + time::Duration::seconds(150))
|
| 300 |
+
.contains(&deadline)
|
| 301 |
+
);
|
| 302 |
+
// Reading the response body must not start a new relative wait or resample jitter.
|
| 303 |
+
let (delay, deferred_until) =
|
| 304 |
+
refresh_deferral(Some(deadline), now + time::Duration::seconds(30));
|
| 305 |
+
assert_eq!(deferred_until, deadline);
|
| 306 |
+
assert!((Duration::from_secs(90)..=Duration::from_secs(120)).contains(&delay));
|
| 307 |
+
}
|
| 308 |
+
}
|
| 309 |
+
|
| 310 |
+
#[test]
|
| 311 |
+
fn zero_retry_after_preserves_bounded_jitter() {
|
| 312 |
+
let now = OffsetDateTime::UNIX_EPOCH;
|
| 313 |
+
let mut headers = HeaderMap::new();
|
| 314 |
+
headers.insert(
|
| 315 |
+
axum::http::header::RETRY_AFTER,
|
| 316 |
+
axum::http::HeaderValue::from_static("0"),
|
| 317 |
+
);
|
| 318 |
+
let deadline = retry_after_with_jitter(&headers, now)
|
| 319 |
+
.expect("a zero-second retry hint should produce a deadline");
|
| 320 |
+
assert!((now..=now + time::Duration::seconds(30)).contains(&deadline));
|
| 321 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/tests.rs
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs
ADDED
|
@@ -0,0 +1,416 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::super::clients::list_remote_control_clients;
|
| 2 |
+
use super::super::clients::revoke_remote_control_client;
|
| 3 |
+
use super::*;
|
| 4 |
+
use codex_app_server_protocol::RemoteControlClient;
|
| 5 |
+
use codex_app_server_protocol::RemoteControlClientsListOrder;
|
| 6 |
+
use codex_app_server_protocol::RemoteControlClientsListParams;
|
| 7 |
+
use codex_app_server_protocol::RemoteControlClientsListResponse;
|
| 8 |
+
use codex_app_server_protocol::RemoteControlClientsRevokeParams;
|
| 9 |
+
use codex_app_server_protocol::RemoteControlClientsRevokeResponse;
|
| 10 |
+
use codex_login::AuthKeyringBackendKind;
|
| 11 |
+
use pretty_assertions::assert_eq;
|
| 12 |
+
|
| 13 |
+
fn client_management_handle(
|
| 14 |
+
remote_control_url: String,
|
| 15 |
+
auth_manager: Arc<AuthManager>,
|
| 16 |
+
) -> RemoteControlSession {
|
| 17 |
+
let desired_state_tx = watch::channel(RemoteControlDesiredState::Disabled).0;
|
| 18 |
+
let (status_tx, _status_rx) = watch::channel(RemoteControlStatusChangedNotification {
|
| 19 |
+
status: RemoteControlConnectionStatus::Disabled,
|
| 20 |
+
server_name: test_server_name(),
|
| 21 |
+
installation_id: TEST_INSTALLATION_ID.to_string(),
|
| 22 |
+
environment_id: None,
|
| 23 |
+
});
|
| 24 |
+
RemoteControlSession {
|
| 25 |
+
policy: RemoteControlPolicy::Allowed,
|
| 26 |
+
shutdown_token: CancellationToken::new(),
|
| 27 |
+
desired_state_tx: Arc::new(desired_state_tx),
|
| 28 |
+
desired_state_rpc_lock: Arc::new(Semaphore::new(1)),
|
| 29 |
+
persistence: RemoteControlPersistence::default(),
|
| 30 |
+
status_tx: Arc::new(status_tx),
|
| 31 |
+
state_db: None,
|
| 32 |
+
remote_control_url,
|
| 33 |
+
current_enrollment: Arc::new(RemoteControlEnrollmentState::new(/*enrollment*/ None)),
|
| 34 |
+
pairing_persistence_key: watch::channel(None).0,
|
| 35 |
+
pairing_persistence_key_required: false,
|
| 36 |
+
auth_manager: auth::RemoteControlAuth::capture(auth_manager).0,
|
| 37 |
+
}
|
| 38 |
+
}
|
| 39 |
+
|
| 40 |
+
fn empty_client_list() -> serde_json::Value {
|
| 41 |
+
json!({
|
| 42 |
+
"items": [],
|
| 43 |
+
"cursor": null,
|
| 44 |
+
})
|
| 45 |
+
}
|
| 46 |
+
|
| 47 |
+
#[tokio::test]
|
| 48 |
+
async fn remote_control_handle_lists_clients_while_disabled() {
|
| 49 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 50 |
+
.await
|
| 51 |
+
.expect("listener should bind");
|
| 52 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 53 |
+
let server_task = tokio::spawn(async move {
|
| 54 |
+
let request = accept_http_request(&listener).await;
|
| 55 |
+
assert_eq!(
|
| 56 |
+
request.request_line,
|
| 57 |
+
"GET /backend-api/wham/remote/control/environments/env%20%2F%3F/clients?cursor=cursor+%2F%3F&limit=10&order=asc HTTP/1.1"
|
| 58 |
+
);
|
| 59 |
+
assert_eq!(
|
| 60 |
+
request.headers.get("authorization"),
|
| 61 |
+
Some(&"Bearer Access Token".to_string())
|
| 62 |
+
);
|
| 63 |
+
assert_eq!(
|
| 64 |
+
request.headers.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER),
|
| 65 |
+
vec!["account_id"]
|
| 66 |
+
);
|
| 67 |
+
respond_with_json(
|
| 68 |
+
request.stream,
|
| 69 |
+
json!({
|
| 70 |
+
"items": [{
|
| 71 |
+
"client_id": "client-123",
|
| 72 |
+
"account_user_id": "user-123",
|
| 73 |
+
"enrollment_status": "enrolled_device_key",
|
| 74 |
+
"display_name": "Anton Phone",
|
| 75 |
+
"device_type": "phone",
|
| 76 |
+
"platform": "ios",
|
| 77 |
+
"os_version": "19.0",
|
| 78 |
+
"device_model": "iPhone",
|
| 79 |
+
"app_version": "1.2.3",
|
| 80 |
+
"last_seen_at": "2026-03-05T07:00:00Z",
|
| 81 |
+
"last_seen_city": "San Francisco",
|
| 82 |
+
}],
|
| 83 |
+
"cursor": "next-cursor",
|
| 84 |
+
}),
|
| 85 |
+
)
|
| 86 |
+
.await;
|
| 87 |
+
});
|
| 88 |
+
let handle = client_management_handle(remote_control_url, remote_control_auth_manager());
|
| 89 |
+
|
| 90 |
+
let response = handle
|
| 91 |
+
.list_clients(RemoteControlClientsListParams {
|
| 92 |
+
environment_id: "env /?".to_string(),
|
| 93 |
+
cursor: Some("cursor /?".to_string()),
|
| 94 |
+
limit: Some(10),
|
| 95 |
+
order: Some(RemoteControlClientsListOrder::Asc),
|
| 96 |
+
})
|
| 97 |
+
.await
|
| 98 |
+
.expect("client list should succeed while remote control is disabled");
|
| 99 |
+
server_task.await.expect("server task should finish");
|
| 100 |
+
|
| 101 |
+
assert_eq!(
|
| 102 |
+
response,
|
| 103 |
+
RemoteControlClientsListResponse {
|
| 104 |
+
data: vec![RemoteControlClient {
|
| 105 |
+
client_id: "client-123".to_string(),
|
| 106 |
+
display_name: Some("Anton Phone".to_string()),
|
| 107 |
+
device_type: Some("phone".to_string()),
|
| 108 |
+
platform: Some("ios".to_string()),
|
| 109 |
+
os_version: Some("19.0".to_string()),
|
| 110 |
+
device_model: Some("iPhone".to_string()),
|
| 111 |
+
app_version: Some("1.2.3".to_string()),
|
| 112 |
+
last_seen_at: Some(1_772_694_000),
|
| 113 |
+
}],
|
| 114 |
+
next_cursor: Some("next-cursor".to_string()),
|
| 115 |
+
}
|
| 116 |
+
);
|
| 117 |
+
}
|
| 118 |
+
|
| 119 |
+
#[tokio::test]
|
| 120 |
+
async fn remote_control_handle_revokes_client_while_disabled() {
|
| 121 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 122 |
+
.await
|
| 123 |
+
.expect("listener should bind");
|
| 124 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 125 |
+
let server_task = tokio::spawn(async move {
|
| 126 |
+
let request = accept_http_request(&listener).await;
|
| 127 |
+
assert_eq!(
|
| 128 |
+
request.request_line,
|
| 129 |
+
"DELETE /backend-api/wham/remote/control/environments/env%20%2F%3F/clients/client%20%2F%3F HTTP/1.1"
|
| 130 |
+
);
|
| 131 |
+
assert_eq!(
|
| 132 |
+
request.headers.get("authorization"),
|
| 133 |
+
Some(&"Bearer Access Token".to_string())
|
| 134 |
+
);
|
| 135 |
+
assert_eq!(
|
| 136 |
+
request.headers.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER),
|
| 137 |
+
vec!["account_id"]
|
| 138 |
+
);
|
| 139 |
+
respond_with_status(request.stream, "204 No Content", "").await;
|
| 140 |
+
});
|
| 141 |
+
let handle = client_management_handle(remote_control_url, remote_control_auth_manager());
|
| 142 |
+
|
| 143 |
+
let response = handle
|
| 144 |
+
.revoke_client(RemoteControlClientsRevokeParams {
|
| 145 |
+
environment_id: "env /?".to_string(),
|
| 146 |
+
client_id: "client /?".to_string(),
|
| 147 |
+
})
|
| 148 |
+
.await
|
| 149 |
+
.expect("client revoke should succeed while remote control is disabled");
|
| 150 |
+
server_task.await.expect("server task should finish");
|
| 151 |
+
|
| 152 |
+
assert_eq!(response, RemoteControlClientsRevokeResponse {});
|
| 153 |
+
}
|
| 154 |
+
|
| 155 |
+
#[tokio::test]
|
| 156 |
+
async fn list_remote_control_clients_recovers_auth_after_unauthorized() {
|
| 157 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 158 |
+
.await
|
| 159 |
+
.expect("listener should bind");
|
| 160 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 161 |
+
let server_task = tokio::spawn(async move {
|
| 162 |
+
let stale_request = accept_http_request(&listener).await;
|
| 163 |
+
assert_eq!(
|
| 164 |
+
stale_request.headers.get("authorization"),
|
| 165 |
+
Some(&"Bearer stale-token".to_string())
|
| 166 |
+
);
|
| 167 |
+
assert_eq!(
|
| 168 |
+
stale_request
|
| 169 |
+
.headers
|
| 170 |
+
.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER),
|
| 171 |
+
vec!["account_id"]
|
| 172 |
+
);
|
| 173 |
+
respond_with_status(stale_request.stream, "401 Unauthorized", "").await;
|
| 174 |
+
|
| 175 |
+
let recovered_request = accept_http_request(&listener).await;
|
| 176 |
+
assert_eq!(
|
| 177 |
+
recovered_request.headers.get("authorization"),
|
| 178 |
+
Some(&"Bearer fresh-token".to_string())
|
| 179 |
+
);
|
| 180 |
+
assert_eq!(
|
| 181 |
+
recovered_request
|
| 182 |
+
.headers
|
| 183 |
+
.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER),
|
| 184 |
+
vec!["account_id"]
|
| 185 |
+
);
|
| 186 |
+
respond_with_json(recovered_request.stream, empty_client_list()).await;
|
| 187 |
+
});
|
| 188 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 189 |
+
let mut stale_auth = remote_control_auth_dot_json(Some("account_id"));
|
| 190 |
+
stale_auth
|
| 191 |
+
.tokens
|
| 192 |
+
.as_mut()
|
| 193 |
+
.expect("stale auth should include tokens")
|
| 194 |
+
.access_token = "stale-token".to_string();
|
| 195 |
+
save_auth(
|
| 196 |
+
codex_home.path(),
|
| 197 |
+
&stale_auth,
|
| 198 |
+
AuthCredentialsStoreMode::File,
|
| 199 |
+
AuthKeyringBackendKind::default(),
|
| 200 |
+
)
|
| 201 |
+
.expect("stale auth should save");
|
| 202 |
+
let auth_manager = AuthManager::shared(
|
| 203 |
+
codex_home.path().to_path_buf(),
|
| 204 |
+
/*enable_codex_api_key_env*/ false,
|
| 205 |
+
AuthCredentialsStoreMode::File,
|
| 206 |
+
/*forced_chatgpt_workspace_id*/ None,
|
| 207 |
+
/*chatgpt_base_url*/ None,
|
| 208 |
+
AuthKeyringBackendKind::default(),
|
| 209 |
+
codex_login::test_support::transport_default_auth_route_config(),
|
| 210 |
+
)
|
| 211 |
+
.await;
|
| 212 |
+
let mut fresh_auth = remote_control_auth_dot_json(Some("account_id"));
|
| 213 |
+
fresh_auth
|
| 214 |
+
.tokens
|
| 215 |
+
.as_mut()
|
| 216 |
+
.expect("fresh auth should include tokens")
|
| 217 |
+
.access_token = "fresh-token".to_string();
|
| 218 |
+
save_auth(
|
| 219 |
+
codex_home.path(),
|
| 220 |
+
&fresh_auth,
|
| 221 |
+
AuthCredentialsStoreMode::File,
|
| 222 |
+
AuthKeyringBackendKind::default(),
|
| 223 |
+
)
|
| 224 |
+
.expect("fresh auth should save");
|
| 225 |
+
|
| 226 |
+
let response = list_remote_control_clients(
|
| 227 |
+
&remote_control_url,
|
| 228 |
+
&auth::RemoteControlAuth::capture(auth_manager.clone()).0,
|
| 229 |
+
RemoteControlClientsListParams {
|
| 230 |
+
environment_id: "env-123".to_string(),
|
| 231 |
+
..Default::default()
|
| 232 |
+
},
|
| 233 |
+
)
|
| 234 |
+
.await
|
| 235 |
+
.expect("client list should recover auth");
|
| 236 |
+
server_task.await.expect("server task should finish");
|
| 237 |
+
|
| 238 |
+
assert_eq!(
|
| 239 |
+
response,
|
| 240 |
+
RemoteControlClientsListResponse {
|
| 241 |
+
data: Vec::new(),
|
| 242 |
+
next_cursor: None,
|
| 243 |
+
}
|
| 244 |
+
);
|
| 245 |
+
}
|
| 246 |
+
|
| 247 |
+
#[tokio::test]
|
| 248 |
+
async fn list_remote_control_clients_retries_unauthorized_only_once() {
|
| 249 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 250 |
+
.await
|
| 251 |
+
.expect("listener should bind");
|
| 252 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 253 |
+
let server_task = tokio::spawn(async move {
|
| 254 |
+
let stale_request = accept_http_request(&listener).await;
|
| 255 |
+
assert_eq!(
|
| 256 |
+
stale_request.headers.get("authorization"),
|
| 257 |
+
Some(&"Bearer stale-token".to_string())
|
| 258 |
+
);
|
| 259 |
+
assert_eq!(
|
| 260 |
+
stale_request
|
| 261 |
+
.headers
|
| 262 |
+
.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER),
|
| 263 |
+
vec!["account_id"]
|
| 264 |
+
);
|
| 265 |
+
respond_with_status(stale_request.stream, "401 Unauthorized", "").await;
|
| 266 |
+
|
| 267 |
+
let recovered_request = accept_http_request(&listener).await;
|
| 268 |
+
assert_eq!(
|
| 269 |
+
recovered_request.headers.get("authorization"),
|
| 270 |
+
Some(&"Bearer fresh-token".to_string())
|
| 271 |
+
);
|
| 272 |
+
assert_eq!(
|
| 273 |
+
recovered_request
|
| 274 |
+
.headers
|
| 275 |
+
.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER),
|
| 276 |
+
vec!["account_id"]
|
| 277 |
+
);
|
| 278 |
+
respond_with_status(recovered_request.stream, "401 Unauthorized", "").await;
|
| 279 |
+
|
| 280 |
+
assert!(
|
| 281 |
+
timeout(Duration::from_millis(100), accept_http_request(&listener))
|
| 282 |
+
.await
|
| 283 |
+
.is_err()
|
| 284 |
+
);
|
| 285 |
+
});
|
| 286 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 287 |
+
let mut stale_auth = remote_control_auth_dot_json(Some("account_id"));
|
| 288 |
+
stale_auth
|
| 289 |
+
.tokens
|
| 290 |
+
.as_mut()
|
| 291 |
+
.expect("stale auth should include tokens")
|
| 292 |
+
.access_token = "stale-token".to_string();
|
| 293 |
+
save_auth(
|
| 294 |
+
codex_home.path(),
|
| 295 |
+
&stale_auth,
|
| 296 |
+
AuthCredentialsStoreMode::File,
|
| 297 |
+
AuthKeyringBackendKind::default(),
|
| 298 |
+
)
|
| 299 |
+
.expect("stale auth should save");
|
| 300 |
+
let auth_manager = AuthManager::shared(
|
| 301 |
+
codex_home.path().to_path_buf(),
|
| 302 |
+
/*enable_codex_api_key_env*/ false,
|
| 303 |
+
AuthCredentialsStoreMode::File,
|
| 304 |
+
/*forced_chatgpt_workspace_id*/ None,
|
| 305 |
+
/*chatgpt_base_url*/ None,
|
| 306 |
+
AuthKeyringBackendKind::default(),
|
| 307 |
+
codex_login::test_support::transport_default_auth_route_config(),
|
| 308 |
+
)
|
| 309 |
+
.await;
|
| 310 |
+
let mut fresh_auth = remote_control_auth_dot_json(Some("account_id"));
|
| 311 |
+
fresh_auth
|
| 312 |
+
.tokens
|
| 313 |
+
.as_mut()
|
| 314 |
+
.expect("fresh auth should include tokens")
|
| 315 |
+
.access_token = "fresh-token".to_string();
|
| 316 |
+
save_auth(
|
| 317 |
+
codex_home.path(),
|
| 318 |
+
&fresh_auth,
|
| 319 |
+
AuthCredentialsStoreMode::File,
|
| 320 |
+
AuthKeyringBackendKind::default(),
|
| 321 |
+
)
|
| 322 |
+
.expect("fresh auth should save");
|
| 323 |
+
|
| 324 |
+
let err = list_remote_control_clients(
|
| 325 |
+
&remote_control_url,
|
| 326 |
+
&auth::RemoteControlAuth::capture(auth_manager.clone()).0,
|
| 327 |
+
RemoteControlClientsListParams {
|
| 328 |
+
environment_id: "env-123".to_string(),
|
| 329 |
+
..Default::default()
|
| 330 |
+
},
|
| 331 |
+
)
|
| 332 |
+
.await
|
| 333 |
+
.expect_err("second unauthorized response should fail");
|
| 334 |
+
server_task.await.expect("server task should finish");
|
| 335 |
+
|
| 336 |
+
assert_eq!(err.kind(), std::io::ErrorKind::PermissionDenied);
|
| 337 |
+
}
|
| 338 |
+
|
| 339 |
+
#[tokio::test]
|
| 340 |
+
async fn revoke_remote_control_client_does_not_retry_forbidden() {
|
| 341 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 342 |
+
.await
|
| 343 |
+
.expect("listener should bind");
|
| 344 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 345 |
+
let server_task = tokio::spawn(async move {
|
| 346 |
+
let request = accept_http_request(&listener).await;
|
| 347 |
+
assert_eq!(
|
| 348 |
+
request.headers.get("authorization"),
|
| 349 |
+
Some(&"Bearer Access Token".to_string())
|
| 350 |
+
);
|
| 351 |
+
assert_eq!(
|
| 352 |
+
request.headers.get_all(REMOTE_CONTROL_ACCOUNT_ID_HEADER),
|
| 353 |
+
vec!["account_id"]
|
| 354 |
+
);
|
| 355 |
+
respond_with_status_and_headers(
|
| 356 |
+
request.stream,
|
| 357 |
+
"403 Forbidden",
|
| 358 |
+
&[("x-request-id", "request-123"), ("cf-ray", "ray-123")],
|
| 359 |
+
"forbidden",
|
| 360 |
+
)
|
| 361 |
+
.await;
|
| 362 |
+
});
|
| 363 |
+
|
| 364 |
+
let err = revoke_remote_control_client(
|
| 365 |
+
&remote_control_url,
|
| 366 |
+
&auth::RemoteControlAuth::capture(remote_control_auth_manager()).0,
|
| 367 |
+
RemoteControlClientsRevokeParams {
|
| 368 |
+
environment_id: "env-123".to_string(),
|
| 369 |
+
client_id: "client-123".to_string(),
|
| 370 |
+
},
|
| 371 |
+
)
|
| 372 |
+
.await
|
| 373 |
+
.expect_err("forbidden revoke should fail");
|
| 374 |
+
server_task.await.expect("server task should finish");
|
| 375 |
+
|
| 376 |
+
assert_eq!(err.kind(), std::io::ErrorKind::PermissionDenied);
|
| 377 |
+
assert_eq!(
|
| 378 |
+
err.to_string(),
|
| 379 |
+
format!(
|
| 380 |
+
"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"
|
| 381 |
+
)
|
| 382 |
+
);
|
| 383 |
+
}
|
| 384 |
+
|
| 385 |
+
#[tokio::test]
|
| 386 |
+
async fn list_remote_control_clients_preserves_decode_error_context() {
|
| 387 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 388 |
+
.await
|
| 389 |
+
.expect("listener should bind");
|
| 390 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 391 |
+
let server_task = tokio::spawn(async move {
|
| 392 |
+
let request = accept_http_request(&listener).await;
|
| 393 |
+
respond_with_status(request.stream, "200 OK", "{").await;
|
| 394 |
+
});
|
| 395 |
+
|
| 396 |
+
let err = list_remote_control_clients(
|
| 397 |
+
&remote_control_url,
|
| 398 |
+
&auth::RemoteControlAuth::capture(remote_control_auth_manager()).0,
|
| 399 |
+
RemoteControlClientsListParams {
|
| 400 |
+
environment_id: "env-123".to_string(),
|
| 401 |
+
..Default::default()
|
| 402 |
+
},
|
| 403 |
+
)
|
| 404 |
+
.await
|
| 405 |
+
.expect_err("malformed client list should fail");
|
| 406 |
+
server_task.await.expect("server task should finish");
|
| 407 |
+
|
| 408 |
+
assert!(
|
| 409 |
+
err.to_string().contains(
|
| 410 |
+
"failed to parse remote control client list response from `http://127.0.0.1:"
|
| 411 |
+
)
|
| 412 |
+
);
|
| 413 |
+
assert!(err.to_string().contains("HTTP 200 OK"));
|
| 414 |
+
assert!(err.to_string().contains("body: {"));
|
| 415 |
+
assert!(err.to_string().contains("decode error:"));
|
| 416 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/tests/pairing_tests.rs
ADDED
|
@@ -0,0 +1,1137 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::super::protocol::RemoteControlPairingStatusRequest;
|
| 2 |
+
use super::super::protocol::StartRemoteControlPairingRequest;
|
| 3 |
+
use super::*;
|
| 4 |
+
use codex_login::AuthKeyringBackendKind;
|
| 5 |
+
use pretty_assertions::assert_eq;
|
| 6 |
+
use std::io;
|
| 7 |
+
|
| 8 |
+
fn remote_control_enrollment(
|
| 9 |
+
remote_control_url: &str,
|
| 10 |
+
remote_control_token: &str,
|
| 11 |
+
) -> RemoteControlEnrollment {
|
| 12 |
+
RemoteControlEnrollment {
|
| 13 |
+
remote_control_target: normalize_remote_control_url(remote_control_url)
|
| 14 |
+
.expect("target should normalize"),
|
| 15 |
+
account_id: "account-id".to_string(),
|
| 16 |
+
environment_id: "environment-id".to_string(),
|
| 17 |
+
server_id: "server-id".to_string(),
|
| 18 |
+
server_name: "server-name".to_string(),
|
| 19 |
+
remote_control_token: Some(remote_control_token.to_string()),
|
| 20 |
+
expires_at: Some(
|
| 21 |
+
OffsetDateTime::from_unix_timestamp(33_336_362_096)
|
| 22 |
+
.expect("future timestamp should parse"),
|
| 23 |
+
),
|
| 24 |
+
next_refresh_at: None,
|
| 25 |
+
}
|
| 26 |
+
}
|
| 27 |
+
|
| 28 |
+
async fn auth_manager_with_replacement(
|
| 29 |
+
codex_home: &TempDir,
|
| 30 |
+
replacement_account_id: &str,
|
| 31 |
+
) -> Arc<AuthManager> {
|
| 32 |
+
let mut stale_auth = remote_control_auth_dot_json(Some("account_id"));
|
| 33 |
+
stale_auth
|
| 34 |
+
.tokens
|
| 35 |
+
.as_mut()
|
| 36 |
+
.expect("stale auth should include tokens")
|
| 37 |
+
.access_token = "stale-token".to_string();
|
| 38 |
+
save_auth(
|
| 39 |
+
codex_home.path(),
|
| 40 |
+
&stale_auth,
|
| 41 |
+
AuthCredentialsStoreMode::File,
|
| 42 |
+
AuthKeyringBackendKind::default(),
|
| 43 |
+
)
|
| 44 |
+
.expect("stale auth should save");
|
| 45 |
+
let auth_manager = AuthManager::shared(
|
| 46 |
+
codex_home.path().to_path_buf(),
|
| 47 |
+
/*enable_codex_api_key_env*/ false,
|
| 48 |
+
AuthCredentialsStoreMode::File,
|
| 49 |
+
/*forced_chatgpt_workspace_id*/ None,
|
| 50 |
+
/*chatgpt_base_url*/ None,
|
| 51 |
+
AuthKeyringBackendKind::default(),
|
| 52 |
+
codex_login::test_support::transport_default_auth_route_config(),
|
| 53 |
+
)
|
| 54 |
+
.await;
|
| 55 |
+
let mut replacement_auth = remote_control_auth_dot_json(Some(replacement_account_id));
|
| 56 |
+
replacement_auth
|
| 57 |
+
.tokens
|
| 58 |
+
.as_mut()
|
| 59 |
+
.expect("replacement auth should include tokens")
|
| 60 |
+
.access_token = "fresh-token".to_string();
|
| 61 |
+
save_auth(
|
| 62 |
+
codex_home.path(),
|
| 63 |
+
&replacement_auth,
|
| 64 |
+
AuthCredentialsStoreMode::File,
|
| 65 |
+
AuthKeyringBackendKind::default(),
|
| 66 |
+
)
|
| 67 |
+
.expect("replacement auth should save");
|
| 68 |
+
auth_manager
|
| 69 |
+
}
|
| 70 |
+
|
| 71 |
+
fn pairing_response_json(server_id: &str, environment_id: &str) -> serde_json::Value {
|
| 72 |
+
json!({
|
| 73 |
+
"pairing_code": "pairing-code",
|
| 74 |
+
"manual_pairing_code": "ABCD-EFGH",
|
| 75 |
+
"server_id": server_id,
|
| 76 |
+
"environment_id": environment_id,
|
| 77 |
+
"expires_at": "3026-05-22T12:34:56Z",
|
| 78 |
+
})
|
| 79 |
+
}
|
| 80 |
+
|
| 81 |
+
fn pairing_response(environment_id: &str) -> RemoteControlPairingStartResponse {
|
| 82 |
+
RemoteControlPairingStartResponse {
|
| 83 |
+
pairing_code: "pairing-code".to_string(),
|
| 84 |
+
manual_pairing_code: Some("ABCD-EFGH".to_string()),
|
| 85 |
+
environment_id: environment_id.to_string(),
|
| 86 |
+
expires_at: 33_336_362_096,
|
| 87 |
+
}
|
| 88 |
+
}
|
| 89 |
+
|
| 90 |
+
async fn pairing_error(status: &'static str, body: &'static str) -> (String, String) {
|
| 91 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 92 |
+
.await
|
| 93 |
+
.expect("listener should bind");
|
| 94 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 95 |
+
let expected_pair_url = normalize_remote_control_url(&remote_control_url)
|
| 96 |
+
.expect("target should normalize")
|
| 97 |
+
.pair_url;
|
| 98 |
+
let server_task = tokio::spawn(async move {
|
| 99 |
+
let pairing_request = accept_http_request(&listener).await;
|
| 100 |
+
respond_with_status_and_headers(
|
| 101 |
+
pairing_request.stream,
|
| 102 |
+
status,
|
| 103 |
+
&[("x-request-id", "request-123"), ("cf-ray", "ray-123")],
|
| 104 |
+
body,
|
| 105 |
+
)
|
| 106 |
+
.await;
|
| 107 |
+
});
|
| 108 |
+
|
| 109 |
+
let err = remote_control_enrollment(&remote_control_url, "remote-control-token")
|
| 110 |
+
.start_pairing(StartRemoteControlPairingRequest { manual_code: false })
|
| 111 |
+
.await
|
| 112 |
+
.expect_err("pairing should fail");
|
| 113 |
+
server_task.await.expect("server task should finish");
|
| 114 |
+
(err.to_string(), expected_pair_url)
|
| 115 |
+
}
|
| 116 |
+
|
| 117 |
+
async fn pairing_response_error(body: serde_json::Value) -> String {
|
| 118 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 119 |
+
.await
|
| 120 |
+
.expect("listener should bind");
|
| 121 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 122 |
+
let server_task = tokio::spawn(async move {
|
| 123 |
+
let pairing_request = accept_http_request(&listener).await;
|
| 124 |
+
respond_with_json(pairing_request.stream, body).await;
|
| 125 |
+
});
|
| 126 |
+
|
| 127 |
+
let err = remote_control_enrollment(&remote_control_url, "remote-control-token")
|
| 128 |
+
.start_pairing(StartRemoteControlPairingRequest { manual_code: false })
|
| 129 |
+
.await
|
| 130 |
+
.expect_err("pairing should fail");
|
| 131 |
+
server_task.await.expect("server task should finish");
|
| 132 |
+
err.to_string()
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
async fn pairing_status_error(status: &'static str, body: &'static str) -> (io::Error, String) {
|
| 136 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 137 |
+
.await
|
| 138 |
+
.expect("listener should bind");
|
| 139 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 140 |
+
let expected_status_url = normalize_remote_control_url(&remote_control_url)
|
| 141 |
+
.expect("target should normalize")
|
| 142 |
+
.pair_status_url;
|
| 143 |
+
let server_task = tokio::spawn(async move {
|
| 144 |
+
let status_request = accept_http_request(&listener).await;
|
| 145 |
+
respond_with_status_and_headers(
|
| 146 |
+
status_request.stream,
|
| 147 |
+
status,
|
| 148 |
+
&[("x-request-id", "request-123"), ("cf-ray", "ray-123")],
|
| 149 |
+
body,
|
| 150 |
+
)
|
| 151 |
+
.await;
|
| 152 |
+
});
|
| 153 |
+
|
| 154 |
+
let err = remote_control_enrollment(&remote_control_url, "remote-control-token")
|
| 155 |
+
.pairing_status(RemoteControlPairingStatusRequest {
|
| 156 |
+
pairing_code: Some("pairing-code".to_string()),
|
| 157 |
+
manual_pairing_code: None,
|
| 158 |
+
})
|
| 159 |
+
.await
|
| 160 |
+
.expect_err("pairing status should fail");
|
| 161 |
+
server_task.await.expect("server task should finish");
|
| 162 |
+
(err, expected_status_url)
|
| 163 |
+
}
|
| 164 |
+
|
| 165 |
+
#[tokio::test]
|
| 166 |
+
async fn remote_control_handle_starts_pairing_before_websocket_connects() {
|
| 167 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 168 |
+
.await
|
| 169 |
+
.expect("listener should bind");
|
| 170 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 171 |
+
let server_task = tokio::spawn(async move {
|
| 172 |
+
let refresh_request = accept_http_request(&listener).await;
|
| 173 |
+
assert_eq!(
|
| 174 |
+
refresh_request.request_line,
|
| 175 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 176 |
+
);
|
| 177 |
+
assert_eq!(
|
| 178 |
+
serde_json::from_str::<serde_json::Value>(&refresh_request.body)
|
| 179 |
+
.expect("refresh request body should deserialize"),
|
| 180 |
+
json!({
|
| 181 |
+
"server_id": "srv_e_test",
|
| 182 |
+
"installation_id": TEST_INSTALLATION_ID,
|
| 183 |
+
})
|
| 184 |
+
);
|
| 185 |
+
respond_with_json(
|
| 186 |
+
refresh_request.stream,
|
| 187 |
+
remote_control_server_token_response(
|
| 188 |
+
"srv_e_test",
|
| 189 |
+
"env_test",
|
| 190 |
+
TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN,
|
| 191 |
+
),
|
| 192 |
+
)
|
| 193 |
+
.await;
|
| 194 |
+
|
| 195 |
+
let pairing_request = accept_http_request(&listener).await;
|
| 196 |
+
assert_eq!(
|
| 197 |
+
pairing_request.request_line,
|
| 198 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 199 |
+
);
|
| 200 |
+
assert_eq!(
|
| 201 |
+
pairing_request.headers.get("authorization"),
|
| 202 |
+
Some(&format!(
|
| 203 |
+
"Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}"
|
| 204 |
+
))
|
| 205 |
+
);
|
| 206 |
+
assert_eq!(
|
| 207 |
+
serde_json::from_str::<serde_json::Value>(&pairing_request.body)
|
| 208 |
+
.expect("pairing request body should deserialize"),
|
| 209 |
+
json!({ "manual_code": true })
|
| 210 |
+
);
|
| 211 |
+
respond_with_json(
|
| 212 |
+
pairing_request.stream,
|
| 213 |
+
pairing_response_json("srv_e_test", "env_test"),
|
| 214 |
+
)
|
| 215 |
+
.await;
|
| 216 |
+
});
|
| 217 |
+
let remote_handle = remote_control_handle_with_current_enrollment(
|
| 218 |
+
&remote_control_url,
|
| 219 |
+
remote_control_auth_manager(),
|
| 220 |
+
);
|
| 221 |
+
remote_handle
|
| 222 |
+
.current_enrollment
|
| 223 |
+
.lock()
|
| 224 |
+
.await
|
| 225 |
+
.as_mut()
|
| 226 |
+
.expect("current enrollment should exist")
|
| 227 |
+
.expires_at = Some(OffsetDateTime::now_utc() + time::Duration::seconds(29));
|
| 228 |
+
|
| 229 |
+
let response = remote_handle
|
| 230 |
+
.start_pairing(
|
| 231 |
+
RemoteControlPairingStartParams { manual_code: true },
|
| 232 |
+
/*app_server_client_name*/ None,
|
| 233 |
+
)
|
| 234 |
+
.await
|
| 235 |
+
.expect("pairing should use the current server before websocket connect");
|
| 236 |
+
server_task.await.expect("server task should finish");
|
| 237 |
+
|
| 238 |
+
assert_eq!(response, pairing_response("env_test"));
|
| 239 |
+
}
|
| 240 |
+
|
| 241 |
+
#[tokio::test]
|
| 242 |
+
async fn proactive_refresh_rate_limit_uses_valid_token_for_pairing() {
|
| 243 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 244 |
+
.await
|
| 245 |
+
.expect("listener should bind");
|
| 246 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 247 |
+
let server_task = tokio::spawn(async move {
|
| 248 |
+
let refresh_request = accept_http_request(&listener).await;
|
| 249 |
+
assert_eq!(
|
| 250 |
+
refresh_request.request_line,
|
| 251 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 252 |
+
);
|
| 253 |
+
respond_with_status_and_headers(
|
| 254 |
+
refresh_request.stream,
|
| 255 |
+
"429 Too Many Requests",
|
| 256 |
+
&[],
|
| 257 |
+
"rate limited",
|
| 258 |
+
)
|
| 259 |
+
.await;
|
| 260 |
+
|
| 261 |
+
let pairing_request = accept_http_request(&listener).await;
|
| 262 |
+
assert_eq!(
|
| 263 |
+
pairing_request.request_line,
|
| 264 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 265 |
+
);
|
| 266 |
+
assert_eq!(
|
| 267 |
+
pairing_request.headers.get("authorization"),
|
| 268 |
+
Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}"))
|
| 269 |
+
);
|
| 270 |
+
respond_with_json(
|
| 271 |
+
pairing_request.stream,
|
| 272 |
+
pairing_response_json("srv_e_test", "env_test"),
|
| 273 |
+
)
|
| 274 |
+
.await;
|
| 275 |
+
});
|
| 276 |
+
let remote_handle = remote_control_handle_with_current_enrollment(
|
| 277 |
+
&remote_control_url,
|
| 278 |
+
remote_control_auth_manager(),
|
| 279 |
+
);
|
| 280 |
+
remote_handle
|
| 281 |
+
.current_enrollment
|
| 282 |
+
.lock()
|
| 283 |
+
.await
|
| 284 |
+
.as_mut()
|
| 285 |
+
.expect("current enrollment should exist")
|
| 286 |
+
.expires_at = Some(OffsetDateTime::now_utc() + time::Duration::minutes(4));
|
| 287 |
+
|
| 288 |
+
let response = remote_handle
|
| 289 |
+
.start_pairing(
|
| 290 |
+
RemoteControlPairingStartParams::default(),
|
| 291 |
+
/*app_server_client_name*/ None,
|
| 292 |
+
)
|
| 293 |
+
.await
|
| 294 |
+
.expect("valid token should allow pairing after proactive refresh failure");
|
| 295 |
+
server_task.await.expect("server task should finish");
|
| 296 |
+
|
| 297 |
+
assert_eq!(response, pairing_response("env_test"));
|
| 298 |
+
assert!(
|
| 299 |
+
remote_handle
|
| 300 |
+
.current_enrollment
|
| 301 |
+
.snapshot()
|
| 302 |
+
.and_then(|enrollment| enrollment.next_refresh_at)
|
| 303 |
+
.is_some()
|
| 304 |
+
);
|
| 305 |
+
}
|
| 306 |
+
|
| 307 |
+
#[tokio::test]
|
| 308 |
+
async fn required_refresh_deadline_blocks_pairing_without_request() {
|
| 309 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 310 |
+
.await
|
| 311 |
+
.expect("listener should bind");
|
| 312 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 313 |
+
let server_task = tokio::spawn(async move {
|
| 314 |
+
let refresh_request = accept_http_request(&listener).await;
|
| 315 |
+
assert_eq!(
|
| 316 |
+
refresh_request.request_line,
|
| 317 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 318 |
+
);
|
| 319 |
+
respond_with_status_and_headers(
|
| 320 |
+
refresh_request.stream,
|
| 321 |
+
"502 Bad Gateway",
|
| 322 |
+
&[("retry-after", "120")],
|
| 323 |
+
"upstream unavailable",
|
| 324 |
+
)
|
| 325 |
+
.await;
|
| 326 |
+
listener
|
| 327 |
+
});
|
| 328 |
+
let remote_handle = remote_control_handle_with_current_enrollment(
|
| 329 |
+
&remote_control_url,
|
| 330 |
+
remote_control_auth_manager(),
|
| 331 |
+
);
|
| 332 |
+
remote_handle
|
| 333 |
+
.current_enrollment
|
| 334 |
+
.lock()
|
| 335 |
+
.await
|
| 336 |
+
.as_mut()
|
| 337 |
+
.expect("current enrollment should exist")
|
| 338 |
+
.expires_at = Some(OffsetDateTime::now_utc() - time::Duration::seconds(1));
|
| 339 |
+
|
| 340 |
+
let refresh_err = remote_handle
|
| 341 |
+
.start_pairing(
|
| 342 |
+
RemoteControlPairingStartParams::default(),
|
| 343 |
+
/*app_server_client_name*/ None,
|
| 344 |
+
)
|
| 345 |
+
.await
|
| 346 |
+
.expect_err("required refresh failure should block pairing");
|
| 347 |
+
let listener = server_task.await.expect("server task should finish");
|
| 348 |
+
let next_refresh_at = remote_handle
|
| 349 |
+
.current_enrollment
|
| 350 |
+
.snapshot()
|
| 351 |
+
.and_then(|enrollment| enrollment.next_refresh_at)
|
| 352 |
+
.expect("required pairing refresh should preserve the retry deadline");
|
| 353 |
+
let deferred_err = remote_handle
|
| 354 |
+
.start_pairing(
|
| 355 |
+
RemoteControlPairingStartParams::default(),
|
| 356 |
+
/*app_server_client_name*/ None,
|
| 357 |
+
)
|
| 358 |
+
.await
|
| 359 |
+
.expect_err("required refresh deadline should block pairing");
|
| 360 |
+
|
| 361 |
+
assert!(refresh_err.to_string().contains("HTTP 502 Bad Gateway"));
|
| 362 |
+
assert_eq!(deferred_err.kind(), io::ErrorKind::WouldBlock);
|
| 363 |
+
assert!(
|
| 364 |
+
deferred_err
|
| 365 |
+
.to_string()
|
| 366 |
+
.contains(&next_refresh_at.to_string())
|
| 367 |
+
);
|
| 368 |
+
timeout(Duration::from_millis(100), listener.accept())
|
| 369 |
+
.await
|
| 370 |
+
.expect_err("pairing should not issue a request before the refresh deadline");
|
| 371 |
+
}
|
| 372 |
+
|
| 373 |
+
#[tokio::test]
|
| 374 |
+
async fn remote_control_pairing_status_returns_pending() {
|
| 375 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 376 |
+
.await
|
| 377 |
+
.expect("listener should bind");
|
| 378 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 379 |
+
let server_task = tokio::spawn(async move {
|
| 380 |
+
let status_request = accept_http_request(&listener).await;
|
| 381 |
+
assert_eq!(
|
| 382 |
+
status_request.request_line,
|
| 383 |
+
"POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1"
|
| 384 |
+
);
|
| 385 |
+
assert_eq!(
|
| 386 |
+
status_request.headers.get("authorization"),
|
| 387 |
+
Some(&"Bearer remote-control-token".to_string())
|
| 388 |
+
);
|
| 389 |
+
assert_eq!(
|
| 390 |
+
serde_json::from_str::<serde_json::Value>(&status_request.body)
|
| 391 |
+
.expect("status request body should deserialize"),
|
| 392 |
+
json!({ "pairing_code": "pairing-code" })
|
| 393 |
+
);
|
| 394 |
+
respond_with_json(status_request.stream, json!({ "claimed": false })).await;
|
| 395 |
+
});
|
| 396 |
+
|
| 397 |
+
let response = remote_control_enrollment(&remote_control_url, "remote-control-token")
|
| 398 |
+
.pairing_status(RemoteControlPairingStatusRequest {
|
| 399 |
+
pairing_code: Some("pairing-code".to_string()),
|
| 400 |
+
manual_pairing_code: None,
|
| 401 |
+
})
|
| 402 |
+
.await
|
| 403 |
+
.expect("pairing status should succeed");
|
| 404 |
+
server_task.await.expect("server task should finish");
|
| 405 |
+
|
| 406 |
+
assert!(!response.claimed);
|
| 407 |
+
}
|
| 408 |
+
|
| 409 |
+
#[tokio::test]
|
| 410 |
+
async fn remote_control_pairing_status_accepts_manual_pairing_code() {
|
| 411 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 412 |
+
.await
|
| 413 |
+
.expect("listener should bind");
|
| 414 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 415 |
+
let server_task = tokio::spawn(async move {
|
| 416 |
+
let status_request = accept_http_request(&listener).await;
|
| 417 |
+
assert_eq!(
|
| 418 |
+
status_request.request_line,
|
| 419 |
+
"POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1"
|
| 420 |
+
);
|
| 421 |
+
assert_eq!(
|
| 422 |
+
serde_json::from_str::<serde_json::Value>(&status_request.body)
|
| 423 |
+
.expect("status request body should deserialize"),
|
| 424 |
+
json!({ "manual_pairing_code": "ABCD-EFGH" })
|
| 425 |
+
);
|
| 426 |
+
respond_with_json(status_request.stream, json!({ "claimed": false })).await;
|
| 427 |
+
});
|
| 428 |
+
|
| 429 |
+
let response = remote_control_enrollment(&remote_control_url, "remote-control-token")
|
| 430 |
+
.pairing_status(RemoteControlPairingStatusRequest {
|
| 431 |
+
pairing_code: None,
|
| 432 |
+
manual_pairing_code: Some("ABCD-EFGH".to_string()),
|
| 433 |
+
})
|
| 434 |
+
.await
|
| 435 |
+
.expect("pairing status should succeed");
|
| 436 |
+
server_task.await.expect("server task should finish");
|
| 437 |
+
|
| 438 |
+
assert!(!response.claimed);
|
| 439 |
+
}
|
| 440 |
+
|
| 441 |
+
#[tokio::test]
|
| 442 |
+
async fn remote_control_pairing_status_returns_claimed() {
|
| 443 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 444 |
+
.await
|
| 445 |
+
.expect("listener should bind");
|
| 446 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 447 |
+
let server_task = tokio::spawn(async move {
|
| 448 |
+
let status_request = accept_http_request(&listener).await;
|
| 449 |
+
assert_eq!(
|
| 450 |
+
status_request.request_line,
|
| 451 |
+
"POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1"
|
| 452 |
+
);
|
| 453 |
+
respond_with_json(status_request.stream, json!({ "claimed": true })).await;
|
| 454 |
+
});
|
| 455 |
+
|
| 456 |
+
let response = remote_control_enrollment(&remote_control_url, "remote-control-token")
|
| 457 |
+
.pairing_status(RemoteControlPairingStatusRequest {
|
| 458 |
+
pairing_code: Some("pairing-code".to_string()),
|
| 459 |
+
manual_pairing_code: None,
|
| 460 |
+
})
|
| 461 |
+
.await
|
| 462 |
+
.expect("pairing status should succeed");
|
| 463 |
+
server_task.await.expect("server task should finish");
|
| 464 |
+
|
| 465 |
+
assert!(response.claimed);
|
| 466 |
+
}
|
| 467 |
+
|
| 468 |
+
#[tokio::test]
|
| 469 |
+
async fn remote_control_handle_refreshes_after_pairing_status_auth_failure() {
|
| 470 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 471 |
+
.await
|
| 472 |
+
.expect("listener should bind");
|
| 473 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 474 |
+
let server_task = tokio::spawn(async move {
|
| 475 |
+
let stale_status_request = accept_http_request(&listener).await;
|
| 476 |
+
assert_eq!(
|
| 477 |
+
stale_status_request.request_line,
|
| 478 |
+
"POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1"
|
| 479 |
+
);
|
| 480 |
+
assert_eq!(
|
| 481 |
+
stale_status_request.headers.get("authorization"),
|
| 482 |
+
Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}"))
|
| 483 |
+
);
|
| 484 |
+
respond_with_status(stale_status_request.stream, "401 Unauthorized", "").await;
|
| 485 |
+
|
| 486 |
+
let refresh_request = accept_http_request(&listener).await;
|
| 487 |
+
assert_eq!(
|
| 488 |
+
refresh_request.request_line,
|
| 489 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 490 |
+
);
|
| 491 |
+
respond_with_json(
|
| 492 |
+
refresh_request.stream,
|
| 493 |
+
remote_control_server_token_response(
|
| 494 |
+
"srv_e_test",
|
| 495 |
+
"env_test",
|
| 496 |
+
TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN,
|
| 497 |
+
),
|
| 498 |
+
)
|
| 499 |
+
.await;
|
| 500 |
+
|
| 501 |
+
let refreshed_status_request = accept_http_request(&listener).await;
|
| 502 |
+
assert_eq!(
|
| 503 |
+
refreshed_status_request.request_line,
|
| 504 |
+
"POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1"
|
| 505 |
+
);
|
| 506 |
+
assert_eq!(
|
| 507 |
+
refreshed_status_request.headers.get("authorization"),
|
| 508 |
+
Some(&format!(
|
| 509 |
+
"Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}"
|
| 510 |
+
))
|
| 511 |
+
);
|
| 512 |
+
respond_with_json(refreshed_status_request.stream, json!({ "claimed": true })).await;
|
| 513 |
+
});
|
| 514 |
+
let remote_handle = remote_control_handle_with_current_enrollment(
|
| 515 |
+
&remote_control_url,
|
| 516 |
+
remote_control_auth_manager(),
|
| 517 |
+
);
|
| 518 |
+
|
| 519 |
+
let response = remote_handle
|
| 520 |
+
.pairing_status(RemoteControlPairingStatusParams {
|
| 521 |
+
pairing_code: Some("pairing-code".to_string()),
|
| 522 |
+
manual_pairing_code: None,
|
| 523 |
+
})
|
| 524 |
+
.await
|
| 525 |
+
.expect("pairing status should refresh after server token auth failure");
|
| 526 |
+
server_task.await.expect("server task should finish");
|
| 527 |
+
|
| 528 |
+
assert!(response.claimed);
|
| 529 |
+
}
|
| 530 |
+
|
| 531 |
+
#[tokio::test]
|
| 532 |
+
async fn remote_control_pairing_status_maps_user_actionable_backend_errors() {
|
| 533 |
+
for (status, expected_kind) in [
|
| 534 |
+
("403 Forbidden", io::ErrorKind::PermissionDenied),
|
| 535 |
+
("404 Not Found", io::ErrorKind::InvalidInput),
|
| 536 |
+
("410 Gone", io::ErrorKind::InvalidInput),
|
| 537 |
+
] {
|
| 538 |
+
let (err, _expected_status_url) = pairing_status_error(status, "not available").await;
|
| 539 |
+
assert_eq!(err.kind(), expected_kind);
|
| 540 |
+
}
|
| 541 |
+
}
|
| 542 |
+
|
| 543 |
+
#[tokio::test]
|
| 544 |
+
async fn remote_control_pairing_status_preserves_decode_error_context() {
|
| 545 |
+
let (err, expected_status_url) = pairing_status_error("200 OK", "{").await;
|
| 546 |
+
let err = err.to_string();
|
| 547 |
+
|
| 548 |
+
assert!(err.contains(&format!(
|
| 549 |
+
"failed to parse remote control pairing status response from `{expected_status_url}`: HTTP 200 OK"
|
| 550 |
+
)));
|
| 551 |
+
assert!(err.contains("request-id: request-123"));
|
| 552 |
+
assert!(err.contains("cf-ray: ray-123"));
|
| 553 |
+
assert!(err.contains("body: {"));
|
| 554 |
+
assert!(err.contains("decode error:"));
|
| 555 |
+
}
|
| 556 |
+
|
| 557 |
+
#[tokio::test]
|
| 558 |
+
async fn remote_control_handle_refreshes_after_pairing_auth_failure() {
|
| 559 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 560 |
+
.await
|
| 561 |
+
.expect("listener should bind");
|
| 562 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 563 |
+
let server_task = tokio::spawn(async move {
|
| 564 |
+
let stale_pairing_request = accept_http_request(&listener).await;
|
| 565 |
+
assert_eq!(
|
| 566 |
+
stale_pairing_request.request_line,
|
| 567 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 568 |
+
);
|
| 569 |
+
assert_eq!(
|
| 570 |
+
stale_pairing_request.headers.get("authorization"),
|
| 571 |
+
Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}"))
|
| 572 |
+
);
|
| 573 |
+
respond_with_status(stale_pairing_request.stream, "401 Unauthorized", "").await;
|
| 574 |
+
|
| 575 |
+
let refresh_request = accept_http_request(&listener).await;
|
| 576 |
+
assert_eq!(
|
| 577 |
+
refresh_request.request_line,
|
| 578 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 579 |
+
);
|
| 580 |
+
assert_eq!(
|
| 581 |
+
refresh_request.headers.get("authorization"),
|
| 582 |
+
Some(&"Bearer Access Token".to_string())
|
| 583 |
+
);
|
| 584 |
+
respond_with_json(
|
| 585 |
+
refresh_request.stream,
|
| 586 |
+
remote_control_server_token_response(
|
| 587 |
+
"srv_e_test",
|
| 588 |
+
"env_test",
|
| 589 |
+
TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN,
|
| 590 |
+
),
|
| 591 |
+
)
|
| 592 |
+
.await;
|
| 593 |
+
|
| 594 |
+
let refreshed_pairing_request = accept_http_request(&listener).await;
|
| 595 |
+
assert_eq!(
|
| 596 |
+
refreshed_pairing_request.request_line,
|
| 597 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 598 |
+
);
|
| 599 |
+
assert_eq!(
|
| 600 |
+
refreshed_pairing_request.headers.get("authorization"),
|
| 601 |
+
Some(&format!(
|
| 602 |
+
"Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}"
|
| 603 |
+
))
|
| 604 |
+
);
|
| 605 |
+
respond_with_json(
|
| 606 |
+
refreshed_pairing_request.stream,
|
| 607 |
+
pairing_response_json("srv_e_test", "env_test"),
|
| 608 |
+
)
|
| 609 |
+
.await;
|
| 610 |
+
});
|
| 611 |
+
let remote_handle = remote_control_handle_with_current_enrollment(
|
| 612 |
+
&remote_control_url,
|
| 613 |
+
remote_control_auth_manager(),
|
| 614 |
+
);
|
| 615 |
+
|
| 616 |
+
let response = remote_handle
|
| 617 |
+
.start_pairing(
|
| 618 |
+
RemoteControlPairingStartParams::default(),
|
| 619 |
+
/*app_server_client_name*/ None,
|
| 620 |
+
)
|
| 621 |
+
.await
|
| 622 |
+
.expect("pairing should refresh after server token auth failure");
|
| 623 |
+
server_task.await.expect("server task should finish");
|
| 624 |
+
|
| 625 |
+
assert_eq!(response, pairing_response("env_test"));
|
| 626 |
+
}
|
| 627 |
+
|
| 628 |
+
#[tokio::test]
|
| 629 |
+
async fn pairing_auth_failure_preserves_refresh_deadline() {
|
| 630 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 631 |
+
.await
|
| 632 |
+
.expect("listener should bind");
|
| 633 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 634 |
+
let server_task = tokio::spawn(async move {
|
| 635 |
+
let pairing_request = accept_http_request(&listener).await;
|
| 636 |
+
assert_eq!(
|
| 637 |
+
pairing_request.request_line,
|
| 638 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 639 |
+
);
|
| 640 |
+
respond_with_status(pairing_request.stream, "401 Unauthorized", "").await;
|
| 641 |
+
});
|
| 642 |
+
let remote_handle = remote_control_handle_with_current_enrollment(
|
| 643 |
+
&remote_control_url,
|
| 644 |
+
remote_control_auth_manager(),
|
| 645 |
+
);
|
| 646 |
+
let next_refresh_at = OffsetDateTime::now_utc() + time::Duration::minutes(2);
|
| 647 |
+
remote_handle
|
| 648 |
+
.current_enrollment
|
| 649 |
+
.lock()
|
| 650 |
+
.await
|
| 651 |
+
.as_mut()
|
| 652 |
+
.expect("current enrollment should exist")
|
| 653 |
+
.next_refresh_at = Some(next_refresh_at);
|
| 654 |
+
let mut expected_enrollment = remote_handle
|
| 655 |
+
.current_enrollment
|
| 656 |
+
.snapshot()
|
| 657 |
+
.expect("current enrollment should exist");
|
| 658 |
+
expected_enrollment.clear_server_token();
|
| 659 |
+
|
| 660 |
+
let err = remote_handle
|
| 661 |
+
.start_pairing(
|
| 662 |
+
RemoteControlPairingStartParams::default(),
|
| 663 |
+
/*app_server_client_name*/ None,
|
| 664 |
+
)
|
| 665 |
+
.await
|
| 666 |
+
.expect_err("refresh deadline should throttle recovery after token rejection");
|
| 667 |
+
server_task.await.expect("server task should finish");
|
| 668 |
+
|
| 669 |
+
assert_eq!(err.kind(), io::ErrorKind::WouldBlock);
|
| 670 |
+
assert_eq!(
|
| 671 |
+
remote_handle.current_enrollment.snapshot(),
|
| 672 |
+
Some(expected_enrollment)
|
| 673 |
+
);
|
| 674 |
+
}
|
| 675 |
+
|
| 676 |
+
#[tokio::test]
|
| 677 |
+
async fn remote_control_handle_recovers_auth_before_refreshing_pairing() {
|
| 678 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 679 |
+
.await
|
| 680 |
+
.expect("listener should bind");
|
| 681 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 682 |
+
let server_task = tokio::spawn(async move {
|
| 683 |
+
let stale_refresh_request = accept_http_request(&listener).await;
|
| 684 |
+
assert_eq!(
|
| 685 |
+
stale_refresh_request.request_line,
|
| 686 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 687 |
+
);
|
| 688 |
+
assert_eq!(
|
| 689 |
+
stale_refresh_request.headers.get("authorization"),
|
| 690 |
+
Some(&"Bearer stale-token".to_string())
|
| 691 |
+
);
|
| 692 |
+
respond_with_status(stale_refresh_request.stream, "401 Unauthorized", "").await;
|
| 693 |
+
|
| 694 |
+
let recovered_refresh_request = accept_http_request(&listener).await;
|
| 695 |
+
assert_eq!(
|
| 696 |
+
recovered_refresh_request.request_line,
|
| 697 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 698 |
+
);
|
| 699 |
+
assert_eq!(
|
| 700 |
+
recovered_refresh_request.headers.get("authorization"),
|
| 701 |
+
Some(&"Bearer fresh-token".to_string())
|
| 702 |
+
);
|
| 703 |
+
respond_with_json(
|
| 704 |
+
recovered_refresh_request.stream,
|
| 705 |
+
remote_control_server_token_response(
|
| 706 |
+
"srv_e_test",
|
| 707 |
+
"env_test",
|
| 708 |
+
TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN,
|
| 709 |
+
),
|
| 710 |
+
)
|
| 711 |
+
.await;
|
| 712 |
+
|
| 713 |
+
let pairing_request = accept_http_request(&listener).await;
|
| 714 |
+
assert_eq!(
|
| 715 |
+
pairing_request.request_line,
|
| 716 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 717 |
+
);
|
| 718 |
+
assert_eq!(
|
| 719 |
+
pairing_request.headers.get("authorization"),
|
| 720 |
+
Some(&format!(
|
| 721 |
+
"Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}"
|
| 722 |
+
))
|
| 723 |
+
);
|
| 724 |
+
respond_with_json(
|
| 725 |
+
pairing_request.stream,
|
| 726 |
+
pairing_response_json("srv_e_test", "env_test"),
|
| 727 |
+
)
|
| 728 |
+
.await;
|
| 729 |
+
});
|
| 730 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 731 |
+
let auth_manager = auth_manager_with_replacement(&codex_home, "account_id").await;
|
| 732 |
+
let remote_handle =
|
| 733 |
+
remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager);
|
| 734 |
+
remote_handle
|
| 735 |
+
.current_enrollment
|
| 736 |
+
.lock()
|
| 737 |
+
.await
|
| 738 |
+
.as_mut()
|
| 739 |
+
.expect("current enrollment should exist")
|
| 740 |
+
.expires_at = Some(OffsetDateTime::now_utc() + time::Duration::seconds(29));
|
| 741 |
+
|
| 742 |
+
let response = remote_handle
|
| 743 |
+
.start_pairing(
|
| 744 |
+
RemoteControlPairingStartParams::default(),
|
| 745 |
+
/*app_server_client_name*/ None,
|
| 746 |
+
)
|
| 747 |
+
.await
|
| 748 |
+
.expect("pairing should refresh after auth recovery");
|
| 749 |
+
server_task.await.expect("server task should finish");
|
| 750 |
+
|
| 751 |
+
assert_eq!(response, pairing_response("env_test"));
|
| 752 |
+
}
|
| 753 |
+
|
| 754 |
+
#[tokio::test]
|
| 755 |
+
async fn pairing_publishes_refresh_deferral_after_auth_recovery() {
|
| 756 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 757 |
+
.await
|
| 758 |
+
.expect("listener should bind");
|
| 759 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 760 |
+
let server_task = tokio::spawn(async move {
|
| 761 |
+
let stale_refresh_request = accept_http_request(&listener).await;
|
| 762 |
+
assert_eq!(
|
| 763 |
+
stale_refresh_request.headers.get("authorization"),
|
| 764 |
+
Some(&"Bearer stale-token".to_string())
|
| 765 |
+
);
|
| 766 |
+
respond_with_status(stale_refresh_request.stream, "401 Unauthorized", "").await;
|
| 767 |
+
|
| 768 |
+
let recovered_refresh_request = accept_http_request(&listener).await;
|
| 769 |
+
assert_eq!(
|
| 770 |
+
recovered_refresh_request.headers.get("authorization"),
|
| 771 |
+
Some(&"Bearer fresh-token".to_string())
|
| 772 |
+
);
|
| 773 |
+
let response_started_at = OffsetDateTime::now_utc();
|
| 774 |
+
respond_with_status_and_headers(
|
| 775 |
+
recovered_refresh_request.stream,
|
| 776 |
+
"502 Bad Gateway",
|
| 777 |
+
&[("retry-after", "120")],
|
| 778 |
+
"upstream unavailable",
|
| 779 |
+
)
|
| 780 |
+
.await;
|
| 781 |
+
response_started_at
|
| 782 |
+
});
|
| 783 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 784 |
+
let auth_manager = auth_manager_with_replacement(&codex_home, "account_id").await;
|
| 785 |
+
let remote_handle =
|
| 786 |
+
remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager);
|
| 787 |
+
remote_handle
|
| 788 |
+
.current_enrollment
|
| 789 |
+
.lock()
|
| 790 |
+
.await
|
| 791 |
+
.as_mut()
|
| 792 |
+
.expect("current enrollment should exist")
|
| 793 |
+
.expires_at = Some(OffsetDateTime::now_utc() - time::Duration::seconds(1));
|
| 794 |
+
|
| 795 |
+
let refresh_err = remote_handle
|
| 796 |
+
.start_pairing(
|
| 797 |
+
RemoteControlPairingStartParams::default(),
|
| 798 |
+
/*app_server_client_name*/ None,
|
| 799 |
+
)
|
| 800 |
+
.await
|
| 801 |
+
.expect_err("required refresh should remain strict after auth recovery");
|
| 802 |
+
let refresh_completed_at = OffsetDateTime::now_utc();
|
| 803 |
+
let deferred_err = remote_handle
|
| 804 |
+
.start_pairing(
|
| 805 |
+
RemoteControlPairingStartParams::default(),
|
| 806 |
+
/*app_server_client_name*/ None,
|
| 807 |
+
)
|
| 808 |
+
.await
|
| 809 |
+
.expect_err("published deadline should throttle the next pairing refresh");
|
| 810 |
+
let response_started_at = server_task.await.expect("server task should finish");
|
| 811 |
+
|
| 812 |
+
assert!(refresh_err.to_string().contains("HTTP 502 Bad Gateway"));
|
| 813 |
+
assert_eq!(deferred_err.kind(), io::ErrorKind::WouldBlock);
|
| 814 |
+
let next_refresh_at = remote_handle
|
| 815 |
+
.current_enrollment
|
| 816 |
+
.snapshot()
|
| 817 |
+
.and_then(|enrollment| enrollment.next_refresh_at)
|
| 818 |
+
.expect("required refresh failure should publish its retry deadline");
|
| 819 |
+
assert!(
|
| 820 |
+
(response_started_at + time::Duration::seconds(120)
|
| 821 |
+
..=refresh_completed_at + time::Duration::seconds(150))
|
| 822 |
+
.contains(&next_refresh_at)
|
| 823 |
+
);
|
| 824 |
+
}
|
| 825 |
+
|
| 826 |
+
#[tokio::test]
|
| 827 |
+
async fn pairing_auth_recovery_failure_publishes_cleared_server_token() {
|
| 828 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 829 |
+
.await
|
| 830 |
+
.expect("listener should bind");
|
| 831 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 832 |
+
let server_task = tokio::spawn(async move {
|
| 833 |
+
let stale_refresh_request = accept_http_request(&listener).await;
|
| 834 |
+
assert_eq!(
|
| 835 |
+
stale_refresh_request.request_line,
|
| 836 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 837 |
+
);
|
| 838 |
+
assert_eq!(
|
| 839 |
+
stale_refresh_request.headers.get("authorization"),
|
| 840 |
+
Some(&"Bearer stale-token".to_string())
|
| 841 |
+
);
|
| 842 |
+
respond_with_status(stale_refresh_request.stream, "401 Unauthorized", "").await;
|
| 843 |
+
});
|
| 844 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 845 |
+
let auth_manager = auth_manager_with_replacement(&codex_home, "different_account_id").await;
|
| 846 |
+
let remote_handle =
|
| 847 |
+
remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager);
|
| 848 |
+
remote_handle
|
| 849 |
+
.current_enrollment
|
| 850 |
+
.lock()
|
| 851 |
+
.await
|
| 852 |
+
.as_mut()
|
| 853 |
+
.expect("current enrollment should exist")
|
| 854 |
+
.expires_at = Some(OffsetDateTime::now_utc() + time::Duration::seconds(29));
|
| 855 |
+
let mut expected_enrollment = remote_handle
|
| 856 |
+
.current_enrollment
|
| 857 |
+
.snapshot()
|
| 858 |
+
.expect("current enrollment should exist");
|
| 859 |
+
expected_enrollment.clear_server_token();
|
| 860 |
+
|
| 861 |
+
let err = remote_handle
|
| 862 |
+
.start_pairing(
|
| 863 |
+
RemoteControlPairingStartParams::default(),
|
| 864 |
+
/*app_server_client_name*/ None,
|
| 865 |
+
)
|
| 866 |
+
.await
|
| 867 |
+
.expect_err("pairing should fail after auth changes account");
|
| 868 |
+
server_task.await.expect("server task should finish");
|
| 869 |
+
|
| 870 |
+
assert_eq!(err.kind(), io::ErrorKind::PermissionDenied);
|
| 871 |
+
assert_eq!(
|
| 872 |
+
remote_handle.current_enrollment.snapshot(),
|
| 873 |
+
Some(expected_enrollment)
|
| 874 |
+
);
|
| 875 |
+
}
|
| 876 |
+
|
| 877 |
+
#[tokio::test]
|
| 878 |
+
async fn start_remote_control_pairing_preserves_backend_error_context() {
|
| 879 |
+
let (err, expected_pair_url) =
|
| 880 |
+
pairing_error("503 Service Unavailable", "pairing unavailable").await;
|
| 881 |
+
|
| 882 |
+
assert_eq!(
|
| 883 |
+
err,
|
| 884 |
+
format!(
|
| 885 |
+
"remote control pairing failed at `{expected_pair_url}`: HTTP 503 Service Unavailable, request-id: request-123, cf-ray: ray-123, body: pairing unavailable"
|
| 886 |
+
)
|
| 887 |
+
);
|
| 888 |
+
}
|
| 889 |
+
|
| 890 |
+
#[tokio::test]
|
| 891 |
+
async fn start_remote_control_pairing_preserves_decode_error_context() {
|
| 892 |
+
let (err, expected_pair_url) = pairing_error("200 OK", "{").await;
|
| 893 |
+
assert!(err.contains(&format!(
|
| 894 |
+
"failed to parse remote control pairing response from `{expected_pair_url}`: HTTP 200 OK"
|
| 895 |
+
)));
|
| 896 |
+
assert!(err.contains("request-id: request-123"));
|
| 897 |
+
assert!(err.contains("cf-ray: ray-123"));
|
| 898 |
+
assert!(err.contains("body: {"));
|
| 899 |
+
assert!(err.contains("decode error:"));
|
| 900 |
+
}
|
| 901 |
+
|
| 902 |
+
#[tokio::test]
|
| 903 |
+
async fn start_remote_control_pairing_rejects_mismatched_backend_enrollment() {
|
| 904 |
+
assert_eq!(
|
| 905 |
+
pairing_response_error(json!({
|
| 906 |
+
"pairing_code": "pairing-code",
|
| 907 |
+
"manual_pairing_code": "ABCD-EFGH",
|
| 908 |
+
"server_id": "other-server-id",
|
| 909 |
+
"environment_id": "other-environment-id",
|
| 910 |
+
"expires_at": "3026-05-22T12:34:56Z",
|
| 911 |
+
}))
|
| 912 |
+
.await,
|
| 913 |
+
"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"
|
| 914 |
+
);
|
| 915 |
+
}
|
| 916 |
+
|
| 917 |
+
#[tokio::test]
|
| 918 |
+
async fn start_remote_control_pairing_preserves_expiry_parse_error_context() {
|
| 919 |
+
let err = pairing_response_error(json!({
|
| 920 |
+
"pairing_code": "pairing-code",
|
| 921 |
+
"manual_pairing_code": "ABCD-EFGH",
|
| 922 |
+
"server_id": "server-id",
|
| 923 |
+
"environment_id": "environment-id",
|
| 924 |
+
"expires_at": "not-a-timestamp",
|
| 925 |
+
}))
|
| 926 |
+
.await;
|
| 927 |
+
|
| 928 |
+
assert!(err.contains("failed to parse remote control pairing response"));
|
| 929 |
+
assert!(err.contains("HTTP 200 OK"));
|
| 930 |
+
assert!(err.contains("request-id: <none>"));
|
| 931 |
+
assert!(err.contains("cf-ray: <none>"));
|
| 932 |
+
assert!(err.contains("\"expires_at\":\"not-a-timestamp\""));
|
| 933 |
+
assert!(err.contains("expires_at parse error:"));
|
| 934 |
+
}
|
| 935 |
+
|
| 936 |
+
#[tokio::test]
|
| 937 |
+
async fn remote_control_handle_disable_keeps_current_enrollment() {
|
| 938 |
+
let remote_handle = remote_control_handle_with_current_enrollment(
|
| 939 |
+
TEST_REMOTE_CONTROL_URL,
|
| 940 |
+
remote_control_auth_manager(),
|
| 941 |
+
);
|
| 942 |
+
|
| 943 |
+
remote_handle
|
| 944 |
+
.desired_state_tx
|
| 945 |
+
.send_replace(RemoteControlDesiredState::Disabled);
|
| 946 |
+
assert!(
|
| 947 |
+
remote_handle.current_enrollment.lock().await.is_some(),
|
| 948 |
+
"disabled remote control should keep the selected pairing server"
|
| 949 |
+
);
|
| 950 |
+
}
|
| 951 |
+
|
| 952 |
+
#[tokio::test]
|
| 953 |
+
async fn remote_control_handle_reenrolls_after_stale_pairing_enrollment() {
|
| 954 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 955 |
+
.await
|
| 956 |
+
.expect("listener should bind");
|
| 957 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 958 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 959 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 960 |
+
let mut remote_handle = remote_control_handle_with_current_enrollment(
|
| 961 |
+
&remote_control_url,
|
| 962 |
+
remote_control_auth_manager_with_home(&codex_home),
|
| 963 |
+
);
|
| 964 |
+
remote_handle.state_db = Some(state_db.clone());
|
| 965 |
+
let stale_enrollment = remote_handle
|
| 966 |
+
.current_enrollment
|
| 967 |
+
.lock()
|
| 968 |
+
.await
|
| 969 |
+
.clone()
|
| 970 |
+
.expect("current enrollment should exist");
|
| 971 |
+
let remote_control_target = stale_enrollment.remote_control_target.clone();
|
| 972 |
+
let refreshed_enrollment = RemoteControlEnrollment {
|
| 973 |
+
remote_control_target: remote_control_target.clone(),
|
| 974 |
+
account_id: "account_id".to_string(),
|
| 975 |
+
environment_id: "env_refreshed".to_string(),
|
| 976 |
+
server_id: "srv_e_refreshed".to_string(),
|
| 977 |
+
server_name: test_server_name(),
|
| 978 |
+
remote_control_token: None,
|
| 979 |
+
expires_at: None,
|
| 980 |
+
next_refresh_at: None,
|
| 981 |
+
};
|
| 982 |
+
update_persisted_remote_control_enrollment(
|
| 983 |
+
Some(state_db.as_ref()),
|
| 984 |
+
&remote_control_target,
|
| 985 |
+
"account_id",
|
| 986 |
+
/*app_server_client_name*/ None,
|
| 987 |
+
Some(&stale_enrollment),
|
| 988 |
+
/*remote_control_enabled*/ Some(true),
|
| 989 |
+
)
|
| 990 |
+
.await
|
| 991 |
+
.expect("stale enrollment should save");
|
| 992 |
+
remote_handle
|
| 993 |
+
.desired_state_tx
|
| 994 |
+
.send_replace(RemoteControlDesiredState::Enabled {
|
| 995 |
+
persistence_preference: Some(true),
|
| 996 |
+
});
|
| 997 |
+
let server_refreshed_enrollment = refreshed_enrollment.clone();
|
| 998 |
+
let server_task = tokio::spawn(async move {
|
| 999 |
+
let stale_pairing_request = accept_http_request(&listener).await;
|
| 1000 |
+
assert_eq!(
|
| 1001 |
+
stale_pairing_request.request_line,
|
| 1002 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 1003 |
+
);
|
| 1004 |
+
assert_eq!(
|
| 1005 |
+
stale_pairing_request.headers.get("authorization"),
|
| 1006 |
+
Some(&format!("Bearer {TEST_REMOTE_CONTROL_SERVER_TOKEN}"))
|
| 1007 |
+
);
|
| 1008 |
+
respond_with_status(stale_pairing_request.stream, "404 Not Found", "").await;
|
| 1009 |
+
|
| 1010 |
+
let enroll_request = accept_http_request(&listener).await;
|
| 1011 |
+
assert_eq!(
|
| 1012 |
+
enroll_request.request_line,
|
| 1013 |
+
"POST /backend-api/wham/remote/control/server/enroll HTTP/1.1"
|
| 1014 |
+
);
|
| 1015 |
+
respond_with_json(
|
| 1016 |
+
enroll_request.stream,
|
| 1017 |
+
remote_control_server_token_response(
|
| 1018 |
+
&server_refreshed_enrollment.server_id,
|
| 1019 |
+
&server_refreshed_enrollment.environment_id,
|
| 1020 |
+
TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN,
|
| 1021 |
+
),
|
| 1022 |
+
)
|
| 1023 |
+
.await;
|
| 1024 |
+
|
| 1025 |
+
let refreshed_pairing_request = accept_http_request(&listener).await;
|
| 1026 |
+
assert_eq!(
|
| 1027 |
+
refreshed_pairing_request.request_line,
|
| 1028 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 1029 |
+
);
|
| 1030 |
+
assert_eq!(
|
| 1031 |
+
refreshed_pairing_request.headers.get("authorization"),
|
| 1032 |
+
Some(&format!(
|
| 1033 |
+
"Bearer {TEST_REFRESHED_REMOTE_CONTROL_SERVER_TOKEN}"
|
| 1034 |
+
))
|
| 1035 |
+
);
|
| 1036 |
+
respond_with_json(
|
| 1037 |
+
refreshed_pairing_request.stream,
|
| 1038 |
+
pairing_response_json(
|
| 1039 |
+
&server_refreshed_enrollment.server_id,
|
| 1040 |
+
&server_refreshed_enrollment.environment_id,
|
| 1041 |
+
),
|
| 1042 |
+
)
|
| 1043 |
+
.await;
|
| 1044 |
+
});
|
| 1045 |
+
let response = remote_handle
|
| 1046 |
+
.start_pairing(
|
| 1047 |
+
RemoteControlPairingStartParams::default(),
|
| 1048 |
+
/*app_server_client_name*/ None,
|
| 1049 |
+
)
|
| 1050 |
+
.await
|
| 1051 |
+
.expect("pairing should re-enroll after stale enrollment");
|
| 1052 |
+
server_task.await.expect("server task should finish");
|
| 1053 |
+
|
| 1054 |
+
assert_eq!(response, pairing_response("env_refreshed"));
|
| 1055 |
+
assert_eq!(
|
| 1056 |
+
state_db
|
| 1057 |
+
.get_remote_control_enrollment(
|
| 1058 |
+
&remote_control_target.websocket_url,
|
| 1059 |
+
"account_id",
|
| 1060 |
+
/*app_server_client_name*/ None,
|
| 1061 |
+
)
|
| 1062 |
+
.await
|
| 1063 |
+
.expect("refreshed enrollment should load")
|
| 1064 |
+
.expect("refreshed enrollment should exist")
|
| 1065 |
+
.remote_control_enabled,
|
| 1066 |
+
Some(true)
|
| 1067 |
+
);
|
| 1068 |
+
}
|
| 1069 |
+
|
| 1070 |
+
#[tokio::test]
|
| 1071 |
+
async fn remote_control_handle_discards_pairing_response_after_auth_change() {
|
| 1072 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 1073 |
+
.await
|
| 1074 |
+
.expect("listener should bind");
|
| 1075 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 1076 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 1077 |
+
save_auth(
|
| 1078 |
+
codex_home.path(),
|
| 1079 |
+
&remote_control_auth_dot_json(Some("account_id")),
|
| 1080 |
+
AuthCredentialsStoreMode::File,
|
| 1081 |
+
AuthKeyringBackendKind::default(),
|
| 1082 |
+
)
|
| 1083 |
+
.expect("initial auth should save");
|
| 1084 |
+
let auth_manager = AuthManager::shared(
|
| 1085 |
+
codex_home.path().to_path_buf(),
|
| 1086 |
+
/*enable_codex_api_key_env*/ false,
|
| 1087 |
+
AuthCredentialsStoreMode::File,
|
| 1088 |
+
/*forced_chatgpt_workspace_id*/ None,
|
| 1089 |
+
/*chatgpt_base_url*/ None,
|
| 1090 |
+
AuthKeyringBackendKind::default(),
|
| 1091 |
+
codex_login::test_support::transport_default_auth_route_config(),
|
| 1092 |
+
)
|
| 1093 |
+
.await;
|
| 1094 |
+
let remote_handle =
|
| 1095 |
+
remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager.clone());
|
| 1096 |
+
let pairing_task = tokio::spawn({
|
| 1097 |
+
let remote_handle = remote_handle.clone();
|
| 1098 |
+
async move {
|
| 1099 |
+
remote_handle
|
| 1100 |
+
.start_pairing(
|
| 1101 |
+
RemoteControlPairingStartParams::default(),
|
| 1102 |
+
/*app_server_client_name*/ None,
|
| 1103 |
+
)
|
| 1104 |
+
.await
|
| 1105 |
+
}
|
| 1106 |
+
});
|
| 1107 |
+
|
| 1108 |
+
let pairing_request = accept_http_request(&listener).await;
|
| 1109 |
+
save_auth(
|
| 1110 |
+
codex_home.path(),
|
| 1111 |
+
&remote_control_auth_dot_json(Some("next_account_id")),
|
| 1112 |
+
AuthCredentialsStoreMode::File,
|
| 1113 |
+
AuthKeyringBackendKind::default(),
|
| 1114 |
+
)
|
| 1115 |
+
.expect("next auth should save");
|
| 1116 |
+
auth_manager.reload().await;
|
| 1117 |
+
respond_with_json(
|
| 1118 |
+
pairing_request.stream,
|
| 1119 |
+
json!({
|
| 1120 |
+
"pairing_code": "stale-pairing-code",
|
| 1121 |
+
"manual_pairing_code": "ABCD-EFGH",
|
| 1122 |
+
"server_id": "srv_e_test",
|
| 1123 |
+
"environment_id": "env_test",
|
| 1124 |
+
"expires_at": "3026-05-22T12:34:56Z",
|
| 1125 |
+
}),
|
| 1126 |
+
)
|
| 1127 |
+
.await;
|
| 1128 |
+
|
| 1129 |
+
assert_eq!(
|
| 1130 |
+
pairing_task
|
| 1131 |
+
.await
|
| 1132 |
+
.expect("pairing task should join")
|
| 1133 |
+
.expect_err("stale pairing response should be discarded")
|
| 1134 |
+
.to_string(),
|
| 1135 |
+
"remote control pairing is unavailable until enrollment completes"
|
| 1136 |
+
);
|
| 1137 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/tests/retry_tests.rs
ADDED
|
@@ -0,0 +1,203 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Exercises overload retry deadlines over HTTP and WebSocket connections.
|
| 2 |
+
//! Auth reloads must respect the deadline; shutdown must remain prompt.
|
| 3 |
+
|
| 4 |
+
use super::*;
|
| 5 |
+
use pretty_assertions::assert_eq;
|
| 6 |
+
|
| 7 |
+
#[tokio::test]
|
| 8 |
+
async fn rate_limited_enrollment_respects_retry_after() {
|
| 9 |
+
assert_overload_retry_after("429 Too Many Requests", /*reject_enrollment*/ true).await;
|
| 10 |
+
}
|
| 11 |
+
|
| 12 |
+
#[tokio::test]
|
| 13 |
+
async fn rate_limited_websocket_respects_retry_after() {
|
| 14 |
+
assert_overload_retry_after("429 Too Many Requests", /*reject_enrollment*/ false).await;
|
| 15 |
+
}
|
| 16 |
+
|
| 17 |
+
#[tokio::test]
|
| 18 |
+
async fn enrollment_resumes_after_retry_after() {
|
| 19 |
+
assert_overload_retry_after("503 Service Unavailable", /*reject_enrollment*/ true).await;
|
| 20 |
+
}
|
| 21 |
+
|
| 22 |
+
#[tokio::test]
|
| 23 |
+
async fn websocket_resumes_after_retry_after() {
|
| 24 |
+
assert_overload_retry_after("503 Service Unavailable", /*reject_enrollment*/ false).await;
|
| 25 |
+
}
|
| 26 |
+
|
| 27 |
+
async fn assert_overload_retry_after(status: &str, reject_enrollment: bool) {
|
| 28 |
+
let verify_recovery = status == "503 Service Unavailable";
|
| 29 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 30 |
+
.await
|
| 31 |
+
.expect("listener should bind");
|
| 32 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 33 |
+
let (transport_event_tx, _transport_event_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 34 |
+
let shutdown_token = CancellationToken::new();
|
| 35 |
+
let auth_manager = remote_control_auth_manager_with_home(&codex_home);
|
| 36 |
+
let mut initial_auth = remote_control_auth_dot_json(Some("account_id"));
|
| 37 |
+
initial_auth
|
| 38 |
+
.tokens
|
| 39 |
+
.as_mut()
|
| 40 |
+
.expect("fixture should contain tokens")
|
| 41 |
+
.access_token = "Initial Access Token".to_string();
|
| 42 |
+
save_auth(
|
| 43 |
+
codex_home.path(),
|
| 44 |
+
&initial_auth,
|
| 45 |
+
AuthCredentialsStoreMode::File,
|
| 46 |
+
AuthKeyringBackendKind::default(),
|
| 47 |
+
)
|
| 48 |
+
.expect("initial credentials should save");
|
| 49 |
+
auth_manager.reload().await;
|
| 50 |
+
let (remote_task, remote_handle) = start_remote_control(
|
| 51 |
+
RemoteControlStartConfig {
|
| 52 |
+
remote_control_url: remote_control_url_for_listener(&listener),
|
| 53 |
+
installation_id: TEST_INSTALLATION_ID.to_string(),
|
| 54 |
+
policy: RemoteControlPolicy::Allowed,
|
| 55 |
+
},
|
| 56 |
+
Some(remote_control_state_runtime(&codex_home).await),
|
| 57 |
+
auth_manager.clone(),
|
| 58 |
+
transport_event_tx,
|
| 59 |
+
shutdown_token.clone(),
|
| 60 |
+
/*app_server_client_name_rx*/ None,
|
| 61 |
+
RemoteControlStartupMode::EnabledEphemeral,
|
| 62 |
+
)
|
| 63 |
+
.await
|
| 64 |
+
.expect("remote control should start");
|
| 65 |
+
let mut status_rx = remote_handle.status_receiver();
|
| 66 |
+
let mut rejected_request = accept_http_request(&listener).await;
|
| 67 |
+
assert_eq!(
|
| 68 |
+
rejected_request.request_line,
|
| 69 |
+
"POST /backend-api/wham/remote/control/server/enroll HTTP/1.1"
|
| 70 |
+
);
|
| 71 |
+
if !reject_enrollment {
|
| 72 |
+
respond_with_json(
|
| 73 |
+
rejected_request.stream,
|
| 74 |
+
remote_control_server_token_response(
|
| 75 |
+
"srv_e_test",
|
| 76 |
+
"env_test",
|
| 77 |
+
TEST_REMOTE_CONTROL_SERVER_TOKEN,
|
| 78 |
+
),
|
| 79 |
+
)
|
| 80 |
+
.await;
|
| 81 |
+
rejected_request = accept_http_request(&listener).await;
|
| 82 |
+
assert_eq!(
|
| 83 |
+
rejected_request.request_line,
|
| 84 |
+
"GET /backend-api/wham/remote/control/server HTTP/1.1"
|
| 85 |
+
);
|
| 86 |
+
}
|
| 87 |
+
let response_started_at = std::time::Instant::now();
|
| 88 |
+
respond_with_status_and_headers(
|
| 89 |
+
rejected_request.stream,
|
| 90 |
+
status,
|
| 91 |
+
&[("Retry-After", if verify_recovery { "3" } else { "120" })],
|
| 92 |
+
"overloaded",
|
| 93 |
+
)
|
| 94 |
+
.await;
|
| 95 |
+
timeout(
|
| 96 |
+
Duration::from_secs(5),
|
| 97 |
+
status_rx.wait_for(|status| status.status == RemoteControlConnectionStatus::Errored),
|
| 98 |
+
)
|
| 99 |
+
.await
|
| 100 |
+
.expect("the overload response should be processed")
|
| 101 |
+
.expect("the status channel should remain open");
|
| 102 |
+
|
| 103 |
+
let auth_changes = auth_manager.auth_change_receiver();
|
| 104 |
+
save_auth(
|
| 105 |
+
codex_home.path(),
|
| 106 |
+
&remote_control_auth_dot_json(Some("account_id")),
|
| 107 |
+
AuthCredentialsStoreMode::File,
|
| 108 |
+
AuthKeyringBackendKind::default(),
|
| 109 |
+
)
|
| 110 |
+
.expect("updated credentials should save");
|
| 111 |
+
auth_manager.reload().await;
|
| 112 |
+
assert!(
|
| 113 |
+
auth_changes
|
| 114 |
+
.has_changed()
|
| 115 |
+
.expect("auth watch should remain open")
|
| 116 |
+
);
|
| 117 |
+
|
| 118 |
+
if !verify_recovery {
|
| 119 |
+
let pairing_error = timeout(
|
| 120 |
+
Duration::from_secs(1),
|
| 121 |
+
remote_handle.start_pairing(
|
| 122 |
+
RemoteControlPairingStartParams::default(),
|
| 123 |
+
/*app_server_client_name*/ None,
|
| 124 |
+
),
|
| 125 |
+
)
|
| 126 |
+
.await
|
| 127 |
+
.expect("pairing should report the active retry delay promptly")
|
| 128 |
+
.expect_err("pairing must respect the shared server retry delay");
|
| 129 |
+
let retry_delay = server_api::remote_control_retry_delay(&pairing_error)
|
| 130 |
+
.expect("pairing should preserve the overload retry deadline");
|
| 131 |
+
assert!(
|
| 132 |
+
retry_delay >= Duration::from_secs(120).saturating_sub(response_started_at.elapsed()),
|
| 133 |
+
"pairing must preserve the full remaining Retry-After interval"
|
| 134 |
+
);
|
| 135 |
+
assert!(
|
| 136 |
+
timeout(Duration::from_secs(2), listener.accept())
|
| 137 |
+
.await
|
| 138 |
+
.is_err(),
|
| 139 |
+
"no enrollment, refresh or handshake should retry before Retry-After ({status})"
|
| 140 |
+
);
|
| 141 |
+
remote_handle.disable_ephemeral().await;
|
| 142 |
+
assert!(
|
| 143 |
+
timeout(Duration::from_secs(2), listener.accept())
|
| 144 |
+
.await
|
| 145 |
+
.is_err(),
|
| 146 |
+
"disabling remote control must stop connection attempts"
|
| 147 |
+
);
|
| 148 |
+
remote_handle
|
| 149 |
+
.enable_ephemeral()
|
| 150 |
+
.expect("remote control should enable again");
|
| 151 |
+
assert!(
|
| 152 |
+
timeout(Duration::from_secs(2), listener.accept())
|
| 153 |
+
.await
|
| 154 |
+
.is_err(),
|
| 155 |
+
"re-enabling remote control must preserve the server retry delay"
|
| 156 |
+
);
|
| 157 |
+
}
|
| 158 |
+
if !reject_enrollment {
|
| 159 |
+
assert_eq!(
|
| 160 |
+
remote_handle
|
| 161 |
+
.inner
|
| 162 |
+
.session()
|
| 163 |
+
.current_enrollment
|
| 164 |
+
.snapshot()
|
| 165 |
+
.and_then(|enrollment| enrollment.remote_control_token),
|
| 166 |
+
Some(TEST_REMOTE_CONTROL_SERVER_TOKEN.to_string()),
|
| 167 |
+
"overload must preserve the enrolled token"
|
| 168 |
+
);
|
| 169 |
+
}
|
| 170 |
+
if verify_recovery {
|
| 171 |
+
remote_handle.disable_ephemeral().await;
|
| 172 |
+
remote_handle
|
| 173 |
+
.enable_ephemeral()
|
| 174 |
+
.expect("remote control should enable again");
|
| 175 |
+
let (stream, _) = timeout(Duration::from_secs(45), listener.accept())
|
| 176 |
+
.await
|
| 177 |
+
.expect("retry should resume after the server delay and bounded jitter")
|
| 178 |
+
.expect("retry connection should succeed");
|
| 179 |
+
assert!(
|
| 180 |
+
response_started_at.elapsed() >= Duration::from_secs(3),
|
| 181 |
+
"the retry must wait for the full Retry-After interval"
|
| 182 |
+
);
|
| 183 |
+
let mut reader = BufReader::new(stream);
|
| 184 |
+
let mut request_line = String::new();
|
| 185 |
+
timeout(Duration::from_secs(5), reader.read_line(&mut request_line))
|
| 186 |
+
.await
|
| 187 |
+
.expect("retry should send the request promptly")
|
| 188 |
+
.expect("retry request should read");
|
| 189 |
+
assert_eq!(
|
| 190 |
+
request_line.trim_end(),
|
| 191 |
+
if reject_enrollment {
|
| 192 |
+
"POST /backend-api/wham/remote/control/server/enroll HTTP/1.1"
|
| 193 |
+
} else {
|
| 194 |
+
"GET /backend-api/wham/remote/control/server HTTP/1.1"
|
| 195 |
+
}
|
| 196 |
+
);
|
| 197 |
+
}
|
| 198 |
+
shutdown_token.cancel();
|
| 199 |
+
timeout(Duration::from_secs(1), remote_task)
|
| 200 |
+
.await
|
| 201 |
+
.expect("shutdown must interrupt the server retry delay")
|
| 202 |
+
.expect("remote task should finish");
|
| 203 |
+
}
|
codex-rs/app-server-transport/src/transport/remote_control/websocket.rs
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
codex-rs/app-server-transport/src/transport/remote_control/websocket_refresh_tests.rs
ADDED
|
@@ -0,0 +1,768 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::tests::TEST_HTTP_ACCEPT_TIMEOUT;
|
| 2 |
+
use super::tests::TEST_INSTALLATION_ID;
|
| 3 |
+
use super::tests::TEST_REMOTE_CONTROL_SERVER_TOKEN;
|
| 4 |
+
use super::tests::accept_http_request;
|
| 5 |
+
use super::tests::enabled_desired_state_sender;
|
| 6 |
+
use super::tests::remote_control_auth_dot_json;
|
| 7 |
+
use super::tests::remote_control_auth_manager;
|
| 8 |
+
use super::tests::remote_control_enrollment;
|
| 9 |
+
use super::tests::remote_control_state_runtime;
|
| 10 |
+
use super::tests::remote_control_status_channel;
|
| 11 |
+
use super::tests::remote_control_url_for_listener;
|
| 12 |
+
use super::tests::respond_with_status_and_headers;
|
| 13 |
+
use super::tests::test_current_enrollment;
|
| 14 |
+
use super::*;
|
| 15 |
+
use crate::transport::remote_control::protocol::normalize_remote_control_url;
|
| 16 |
+
use crate::transport::remote_control::server_api::remote_control_retry_at;
|
| 17 |
+
use crate::transport::remote_control::tests::remote_control_handle_with_current_enrollment;
|
| 18 |
+
use codex_app_server_protocol::RemoteControlPairingStartParams;
|
| 19 |
+
use codex_app_server_protocol::RemoteControlPairingStatusParams;
|
| 20 |
+
use codex_config::types::AuthCredentialsStoreMode;
|
| 21 |
+
use codex_login::AuthKeyringBackendKind;
|
| 22 |
+
use codex_login::AuthManager;
|
| 23 |
+
use codex_login::save_auth;
|
| 24 |
+
use pretty_assertions::assert_eq;
|
| 25 |
+
use tempfile::TempDir;
|
| 26 |
+
use tokio::io::AsyncWriteExt;
|
| 27 |
+
use tokio::net::TcpListener;
|
| 28 |
+
use tokio::net::TcpStream;
|
| 29 |
+
use tokio::time::Duration;
|
| 30 |
+
use tokio::time::timeout;
|
| 31 |
+
use tokio_tungstenite::WebSocketStream;
|
| 32 |
+
use tokio_tungstenite::accept_async;
|
| 33 |
+
|
| 34 |
+
async fn connect_test_websocket(
|
| 35 |
+
remote_control_target: &RemoteControlTarget,
|
| 36 |
+
state_db: &StateRuntime,
|
| 37 |
+
auth_manager: &Arc<AuthManager>,
|
| 38 |
+
current_enrollment: &CurrentRemoteControlEnrollment,
|
| 39 |
+
) -> io::Result<()> {
|
| 40 |
+
let session_auth = RemoteControlAuth::capture(auth_manager.clone()).0;
|
| 41 |
+
let mut auth_recovery = session_auth.unauthorized_recovery();
|
| 42 |
+
let mut auth_change_rx = auth_manager.auth_change_receiver();
|
| 43 |
+
let (status_publisher, _) = remote_control_status_channel();
|
| 44 |
+
let desired_state_tx = enabled_desired_state_sender();
|
| 45 |
+
let persistence = RemoteControlPersistence::default();
|
| 46 |
+
connect_remote_control_websocket(
|
| 47 |
+
remote_control_target,
|
| 48 |
+
Some(state_db),
|
| 49 |
+
RemoteControlAuthContext {
|
| 50 |
+
auth_manager: &session_auth,
|
| 51 |
+
auth_recovery: &mut auth_recovery,
|
| 52 |
+
auth_change_rx: &mut auth_change_rx,
|
| 53 |
+
},
|
| 54 |
+
current_enrollment,
|
| 55 |
+
RemoteControlConnectOptions {
|
| 56 |
+
installation_id: TEST_INSTALLATION_ID,
|
| 57 |
+
server_name: "test-server",
|
| 58 |
+
subscribe_cursor: None,
|
| 59 |
+
app_server_client_name: None,
|
| 60 |
+
desired_state_tx: &desired_state_tx,
|
| 61 |
+
persistence: &persistence,
|
| 62 |
+
},
|
| 63 |
+
&status_publisher,
|
| 64 |
+
)
|
| 65 |
+
.await
|
| 66 |
+
.map(|_| ())
|
| 67 |
+
}
|
| 68 |
+
|
| 69 |
+
#[tokio::test]
|
| 70 |
+
async fn proactive_refresh_failure_uses_valid_token_for_websocket_connect() {
|
| 71 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 72 |
+
.await
|
| 73 |
+
.expect("listener should bind");
|
| 74 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 75 |
+
let remote_control_target =
|
| 76 |
+
normalize_remote_control_url(&remote_control_url).expect("target should parse");
|
| 77 |
+
let server_task = tokio::spawn(async move {
|
| 78 |
+
let (stream, request_line) = accept_http_request(&listener).await;
|
| 79 |
+
assert_eq!(
|
| 80 |
+
request_line,
|
| 81 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 82 |
+
);
|
| 83 |
+
respond_with_status_and_headers(stream, "502 Bad Gateway", &[], "upstream unavailable")
|
| 84 |
+
.await;
|
| 85 |
+
accept_test_websocket(&listener).await
|
| 86 |
+
});
|
| 87 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 88 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 89 |
+
let auth_manager = remote_control_auth_manager();
|
| 90 |
+
let mut enrollment = remote_control_enrollment(Some(TEST_REMOTE_CONTROL_SERVER_TOKEN));
|
| 91 |
+
enrollment.expires_at = Some(time::OffsetDateTime::now_utc() + time::Duration::minutes(4));
|
| 92 |
+
let current_enrollment = test_current_enrollment(Some(enrollment));
|
| 93 |
+
|
| 94 |
+
let refresh_started_at = time::OffsetDateTime::now_utc();
|
| 95 |
+
connect_test_websocket(
|
| 96 |
+
&remote_control_target,
|
| 97 |
+
state_db.as_ref(),
|
| 98 |
+
&auth_manager,
|
| 99 |
+
¤t_enrollment,
|
| 100 |
+
)
|
| 101 |
+
.await
|
| 102 |
+
.expect("valid token should allow websocket connect after proactive refresh failure");
|
| 103 |
+
let refresh_completed_at = time::OffsetDateTime::now_utc();
|
| 104 |
+
let server_websocket = server_task.await.expect("server task should succeed");
|
| 105 |
+
|
| 106 |
+
let enrollment = current_enrollment
|
| 107 |
+
.lock()
|
| 108 |
+
.await
|
| 109 |
+
.clone()
|
| 110 |
+
.expect("enrollment should remain available");
|
| 111 |
+
assert_eq!(
|
| 112 |
+
enrollment.remote_control_token.as_deref(),
|
| 113 |
+
Some(TEST_REMOTE_CONTROL_SERVER_TOKEN)
|
| 114 |
+
);
|
| 115 |
+
let next_refresh_at = enrollment
|
| 116 |
+
.next_refresh_at
|
| 117 |
+
.expect("transient refresh should set a retry deadline");
|
| 118 |
+
assert!(
|
| 119 |
+
(refresh_started_at + time::Duration::seconds(24)
|
| 120 |
+
..=refresh_completed_at + time::Duration::seconds(36))
|
| 121 |
+
.contains(&next_refresh_at)
|
| 122 |
+
);
|
| 123 |
+
drop(server_websocket);
|
| 124 |
+
}
|
| 125 |
+
|
| 126 |
+
#[tokio::test]
|
| 127 |
+
async fn proactive_refresh_connection_failure_uses_valid_token_for_websocket_connect() {
|
| 128 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 129 |
+
.await
|
| 130 |
+
.expect("listener should bind");
|
| 131 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 132 |
+
let remote_control_target =
|
| 133 |
+
normalize_remote_control_url(&remote_control_url).expect("target should parse");
|
| 134 |
+
let server_task = tokio::spawn(async move {
|
| 135 |
+
let (stream, request_line) = accept_http_request(&listener).await;
|
| 136 |
+
assert_eq!(
|
| 137 |
+
request_line,
|
| 138 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 139 |
+
);
|
| 140 |
+
drop(stream);
|
| 141 |
+
accept_test_websocket(&listener).await
|
| 142 |
+
});
|
| 143 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 144 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 145 |
+
let auth_manager = remote_control_auth_manager();
|
| 146 |
+
let mut enrollment = remote_control_enrollment(Some(TEST_REMOTE_CONTROL_SERVER_TOKEN));
|
| 147 |
+
enrollment.expires_at = Some(time::OffsetDateTime::now_utc() + time::Duration::minutes(4));
|
| 148 |
+
let current_enrollment = test_current_enrollment(Some(enrollment));
|
| 149 |
+
|
| 150 |
+
connect_test_websocket(
|
| 151 |
+
&remote_control_target,
|
| 152 |
+
state_db.as_ref(),
|
| 153 |
+
&auth_manager,
|
| 154 |
+
¤t_enrollment,
|
| 155 |
+
)
|
| 156 |
+
.await
|
| 157 |
+
.expect("valid token should allow websocket connect after refresh connection failure");
|
| 158 |
+
let server_websocket = server_task.await.expect("server task should succeed");
|
| 159 |
+
|
| 160 |
+
assert!(
|
| 161 |
+
current_enrollment
|
| 162 |
+
.snapshot()
|
| 163 |
+
.and_then(|enrollment| enrollment.next_refresh_at)
|
| 164 |
+
.is_some(),
|
| 165 |
+
"connection failure should set a retry deadline"
|
| 166 |
+
);
|
| 167 |
+
drop(server_websocket);
|
| 168 |
+
}
|
| 169 |
+
|
| 170 |
+
#[tokio::test]
|
| 171 |
+
async fn websocket_retry_after_throttles_pairing_refresh() {
|
| 172 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 173 |
+
.await
|
| 174 |
+
.expect("listener should bind");
|
| 175 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 176 |
+
let remote_control_target =
|
| 177 |
+
normalize_remote_control_url(&remote_control_url).expect("target should parse");
|
| 178 |
+
let server_task = tokio::spawn(async move {
|
| 179 |
+
let (stream, request_line) = accept_http_request(&listener).await;
|
| 180 |
+
assert_eq!(
|
| 181 |
+
request_line,
|
| 182 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 183 |
+
);
|
| 184 |
+
respond_with_status_and_headers(
|
| 185 |
+
stream,
|
| 186 |
+
"502 Bad Gateway",
|
| 187 |
+
&[("retry-after", "120")],
|
| 188 |
+
"upstream unavailable",
|
| 189 |
+
)
|
| 190 |
+
.await;
|
| 191 |
+
listener
|
| 192 |
+
});
|
| 193 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 194 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 195 |
+
let auth_manager = remote_control_auth_manager();
|
| 196 |
+
let mut remote_handle =
|
| 197 |
+
remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager.clone());
|
| 198 |
+
remote_handle.state_db = Some(state_db.clone());
|
| 199 |
+
remote_handle
|
| 200 |
+
.current_enrollment
|
| 201 |
+
.lock()
|
| 202 |
+
.await
|
| 203 |
+
.as_mut()
|
| 204 |
+
.expect("current enrollment should exist")
|
| 205 |
+
.expires_at = Some(time::OffsetDateTime::now_utc() + time::Duration::minutes(4));
|
| 206 |
+
let current_enrollment = remote_handle.current_enrollment.clone();
|
| 207 |
+
let refresh_started_at = time::OffsetDateTime::now_utc();
|
| 208 |
+
let refresh_error = connect_test_websocket(
|
| 209 |
+
&remote_control_target,
|
| 210 |
+
state_db.as_ref(),
|
| 211 |
+
&auth_manager,
|
| 212 |
+
¤t_enrollment,
|
| 213 |
+
)
|
| 214 |
+
.await
|
| 215 |
+
.expect_err("an explicit server deadline must defer the handshake even with a valid token");
|
| 216 |
+
let refresh_completed_at = time::OffsetDateTime::now_utc();
|
| 217 |
+
let next_refresh_at = current_enrollment
|
| 218 |
+
.snapshot()
|
| 219 |
+
.and_then(|enrollment| enrollment.next_refresh_at)
|
| 220 |
+
.expect("Retry-After should set a retry deadline");
|
| 221 |
+
assert!(
|
| 222 |
+
(refresh_started_at + time::Duration::seconds(120)
|
| 223 |
+
..=refresh_completed_at + time::Duration::seconds(150))
|
| 224 |
+
.contains(&next_refresh_at)
|
| 225 |
+
);
|
| 226 |
+
|
| 227 |
+
let pairing_error = remote_handle
|
| 228 |
+
.start_pairing(
|
| 229 |
+
RemoteControlPairingStartParams::default(),
|
| 230 |
+
/*app_server_client_name*/ None,
|
| 231 |
+
)
|
| 232 |
+
.await
|
| 233 |
+
.expect_err("the refresh deadline must also defer pairing");
|
| 234 |
+
let listener = server_task.await.expect("server task should succeed");
|
| 235 |
+
assert_eq!(
|
| 236 |
+
remote_control_retry_at(&refresh_error),
|
| 237 |
+
Some(next_refresh_at)
|
| 238 |
+
);
|
| 239 |
+
assert_eq!(
|
| 240 |
+
remote_control_retry_at(&pairing_error),
|
| 241 |
+
Some(next_refresh_at)
|
| 242 |
+
);
|
| 243 |
+
assert_eq!(
|
| 244 |
+
current_enrollment
|
| 245 |
+
.snapshot()
|
| 246 |
+
.and_then(|enrollment| enrollment.remote_control_token),
|
| 247 |
+
Some(TEST_REMOTE_CONTROL_SERVER_TOKEN.to_string())
|
| 248 |
+
);
|
| 249 |
+
timeout(Duration::from_millis(100), listener.accept())
|
| 250 |
+
.await
|
| 251 |
+
.expect_err("no handshake or pairing should bypass the proactive refresh deadline");
|
| 252 |
+
}
|
| 253 |
+
|
| 254 |
+
#[tokio::test]
|
| 255 |
+
async fn pairing_http_date_retry_after_throttles_websocket_refresh() {
|
| 256 |
+
for status in [
|
| 257 |
+
"429 Too Many Requests",
|
| 258 |
+
"503 Service Unavailable",
|
| 259 |
+
"502 Bad Gateway",
|
| 260 |
+
] {
|
| 261 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 262 |
+
.await
|
| 263 |
+
.expect("listener should bind");
|
| 264 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 265 |
+
let remote_control_target =
|
| 266 |
+
normalize_remote_control_url(&remote_control_url).expect("target should parse");
|
| 267 |
+
let retry_after =
|
| 268 |
+
httpdate::fmt_http_date(std::time::SystemTime::now() + Duration::from_secs(120));
|
| 269 |
+
let expected_next_refresh_at = time::OffsetDateTime::from(
|
| 270 |
+
httpdate::parse_http_date(&retry_after).expect("Retry-After date should parse"),
|
| 271 |
+
);
|
| 272 |
+
let server_task = tokio::spawn(async move {
|
| 273 |
+
let (refresh_stream, request_line) = accept_http_request(&listener).await;
|
| 274 |
+
assert_eq!(
|
| 275 |
+
request_line,
|
| 276 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 277 |
+
);
|
| 278 |
+
respond_with_status_and_headers(
|
| 279 |
+
refresh_stream,
|
| 280 |
+
status,
|
| 281 |
+
&[("retry-after", &retry_after)],
|
| 282 |
+
"upstream unavailable",
|
| 283 |
+
)
|
| 284 |
+
.await;
|
| 285 |
+
listener
|
| 286 |
+
});
|
| 287 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 288 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 289 |
+
let auth_manager = remote_control_auth_manager();
|
| 290 |
+
let mut remote_handle = remote_control_handle_with_current_enrollment(
|
| 291 |
+
&remote_control_url,
|
| 292 |
+
auth_manager.clone(),
|
| 293 |
+
);
|
| 294 |
+
remote_handle.state_db = Some(state_db.clone());
|
| 295 |
+
remote_handle
|
| 296 |
+
.current_enrollment
|
| 297 |
+
.lock()
|
| 298 |
+
.await
|
| 299 |
+
.as_mut()
|
| 300 |
+
.expect("current enrollment should exist")
|
| 301 |
+
.expires_at = Some(time::OffsetDateTime::now_utc() + time::Duration::minutes(4));
|
| 302 |
+
let current_enrollment = remote_handle.current_enrollment.clone();
|
| 303 |
+
|
| 304 |
+
let refresh_error = remote_handle
|
| 305 |
+
.start_pairing(
|
| 306 |
+
RemoteControlPairingStartParams::default(),
|
| 307 |
+
/*app_server_client_name*/ None,
|
| 308 |
+
)
|
| 309 |
+
.await
|
| 310 |
+
.expect_err("an explicit server deadline must defer pairing even with a valid token");
|
| 311 |
+
let listener = server_task.await.expect("server task should succeed");
|
| 312 |
+
let retry_at = remote_control_retry_at(&refresh_error)
|
| 313 |
+
.expect("the proactive refresh response should preserve its deadline");
|
| 314 |
+
let enrollment = current_enrollment
|
| 315 |
+
.snapshot()
|
| 316 |
+
.expect("enrollment should remain");
|
| 317 |
+
assert_eq!(enrollment.next_refresh_at, Some(retry_at));
|
| 318 |
+
assert_eq!(
|
| 319 |
+
enrollment.remote_control_token.as_deref(),
|
| 320 |
+
Some(TEST_REMOTE_CONTROL_SERVER_TOKEN)
|
| 321 |
+
);
|
| 322 |
+
assert!(
|
| 323 |
+
(expected_next_refresh_at..=expected_next_refresh_at + time::Duration::seconds(30))
|
| 324 |
+
.contains(&retry_at)
|
| 325 |
+
);
|
| 326 |
+
let connect_error = connect_test_websocket(
|
| 327 |
+
&remote_control_target,
|
| 328 |
+
state_db.as_ref(),
|
| 329 |
+
&auth_manager,
|
| 330 |
+
¤t_enrollment,
|
| 331 |
+
)
|
| 332 |
+
.await
|
| 333 |
+
.expect_err("the refresh deadline must also defer the handshake");
|
| 334 |
+
assert_eq!(remote_control_retry_at(&connect_error), Some(retry_at));
|
| 335 |
+
let pairing_error = remote_handle
|
| 336 |
+
.start_pairing(
|
| 337 |
+
RemoteControlPairingStartParams::default(),
|
| 338 |
+
/*app_server_client_name*/ None,
|
| 339 |
+
)
|
| 340 |
+
.await
|
| 341 |
+
.expect_err("another pairing request must retain the same deadline");
|
| 342 |
+
assert_eq!(remote_control_retry_at(&pairing_error), Some(retry_at));
|
| 343 |
+
timeout(Duration::from_millis(100), listener.accept())
|
| 344 |
+
.await
|
| 345 |
+
.expect_err("no pairing, refresh or handshake should bypass the deadline");
|
| 346 |
+
}
|
| 347 |
+
}
|
| 348 |
+
|
| 349 |
+
#[tokio::test]
|
| 350 |
+
async fn pairing_during_pending_handshake_respects_later_overload() {
|
| 351 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 352 |
+
.await
|
| 353 |
+
.expect("listener should bind");
|
| 354 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 355 |
+
let remote_control_target =
|
| 356 |
+
normalize_remote_control_url(&remote_control_url).expect("target should parse");
|
| 357 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 358 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 359 |
+
let auth_manager = remote_control_auth_manager();
|
| 360 |
+
let mut remote_handle =
|
| 361 |
+
remote_control_handle_with_current_enrollment(&remote_control_url, auth_manager.clone());
|
| 362 |
+
remote_handle.state_db = Some(state_db.clone());
|
| 363 |
+
let current_enrollment = remote_handle.current_enrollment.clone();
|
| 364 |
+
let connect = connect_test_websocket(
|
| 365 |
+
&remote_control_target,
|
| 366 |
+
state_db.as_ref(),
|
| 367 |
+
&auth_manager,
|
| 368 |
+
¤t_enrollment,
|
| 369 |
+
);
|
| 370 |
+
tokio::pin!(connect);
|
| 371 |
+
let (handshake_stream, request_line) = tokio::select! {
|
| 372 |
+
result = &mut connect => panic!("handshake should wait for its response: {result:?}"),
|
| 373 |
+
request = accept_http_request(&listener) => request,
|
| 374 |
+
};
|
| 375 |
+
assert_eq!(
|
| 376 |
+
request_line,
|
| 377 |
+
"GET /backend-api/wham/remote/control/server HTTP/1.1"
|
| 378 |
+
);
|
| 379 |
+
|
| 380 |
+
let pairing = remote_handle.start_pairing(
|
| 381 |
+
RemoteControlPairingStartParams::default(),
|
| 382 |
+
/*app_server_client_name*/ None,
|
| 383 |
+
);
|
| 384 |
+
tokio::pin!(pairing);
|
| 385 |
+
let (pairing_stream, request_line) = tokio::select! {
|
| 386 |
+
result = &mut pairing => panic!("pairing should wait for its response: {result:?}"),
|
| 387 |
+
request = accept_http_request(&listener) => request,
|
| 388 |
+
};
|
| 389 |
+
assert_eq!(
|
| 390 |
+
request_line,
|
| 391 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 392 |
+
);
|
| 393 |
+
respond_with_status_and_headers(
|
| 394 |
+
pairing_stream,
|
| 395 |
+
"200 OK",
|
| 396 |
+
&[],
|
| 397 |
+
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"}"#,
|
| 398 |
+
)
|
| 399 |
+
.await;
|
| 400 |
+
timeout(Duration::from_secs(1), pairing)
|
| 401 |
+
.await
|
| 402 |
+
.expect("pairing must stay responsive while the handshake is pending")
|
| 403 |
+
.expect("pairing should succeed before any overload is observed");
|
| 404 |
+
respond_with_status_and_headers(
|
| 405 |
+
handshake_stream,
|
| 406 |
+
"503 Service Unavailable",
|
| 407 |
+
&[("Retry-After", "120")],
|
| 408 |
+
"overloaded",
|
| 409 |
+
)
|
| 410 |
+
.await;
|
| 411 |
+
let connect_error = connect
|
| 412 |
+
.await
|
| 413 |
+
.expect_err("the handshake should report overload");
|
| 414 |
+
let retry_at = remote_control_retry_at(&connect_error)
|
| 415 |
+
.expect("the handshake should preserve its retry deadline");
|
| 416 |
+
let pairing_error = timeout(
|
| 417 |
+
Duration::from_secs(1),
|
| 418 |
+
remote_handle.start_pairing(
|
| 419 |
+
RemoteControlPairingStartParams::default(),
|
| 420 |
+
/*app_server_client_name*/ None,
|
| 421 |
+
),
|
| 422 |
+
)
|
| 423 |
+
.await
|
| 424 |
+
.expect("new pairing should report the deadline promptly")
|
| 425 |
+
.expect_err("new pairing must honor the handshake deadline");
|
| 426 |
+
assert_eq!(remote_control_retry_at(&pairing_error), Some(retry_at));
|
| 427 |
+
timeout(Duration::from_millis(100), listener.accept())
|
| 428 |
+
.await
|
| 429 |
+
.expect_err("new pairing must not bypass the newly recorded deadline");
|
| 430 |
+
}
|
| 431 |
+
|
| 432 |
+
#[tokio::test]
|
| 433 |
+
async fn pairing_auth_recovery_respects_concurrent_handshake_overload() {
|
| 434 |
+
for (pairing_status, auth_endpoint) in
|
| 435 |
+
[("401 Unauthorized", "refresh"), ("404 Not Found", "enroll")]
|
| 436 |
+
{
|
| 437 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 438 |
+
.await
|
| 439 |
+
.expect("listener should bind");
|
| 440 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 441 |
+
let remote_control_target =
|
| 442 |
+
normalize_remote_control_url(&remote_control_url).expect("target should parse");
|
| 443 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 444 |
+
save_auth(
|
| 445 |
+
codex_home.path(),
|
| 446 |
+
&remote_control_auth_dot_json("stale-token"),
|
| 447 |
+
AuthCredentialsStoreMode::File,
|
| 448 |
+
AuthKeyringBackendKind::default(),
|
| 449 |
+
)
|
| 450 |
+
.expect("stale auth should save");
|
| 451 |
+
let auth_manager = AuthManager::shared(
|
| 452 |
+
codex_home.path().to_path_buf(),
|
| 453 |
+
/*enable_codex_api_key_env*/ false,
|
| 454 |
+
AuthCredentialsStoreMode::File,
|
| 455 |
+
/*forced_chatgpt_workspace_id*/ None,
|
| 456 |
+
/*chatgpt_base_url*/ None,
|
| 457 |
+
AuthKeyringBackendKind::default(),
|
| 458 |
+
codex_login::test_support::transport_default_auth_route_config(),
|
| 459 |
+
)
|
| 460 |
+
.await;
|
| 461 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 462 |
+
let mut remote_handle = remote_control_handle_with_current_enrollment(
|
| 463 |
+
&remote_control_url,
|
| 464 |
+
auth_manager.clone(),
|
| 465 |
+
);
|
| 466 |
+
remote_handle.state_db = Some(state_db.clone());
|
| 467 |
+
let current_enrollment = remote_handle.current_enrollment.clone();
|
| 468 |
+
let connect = connect_test_websocket(
|
| 469 |
+
&remote_control_target,
|
| 470 |
+
state_db.as_ref(),
|
| 471 |
+
&auth_manager,
|
| 472 |
+
¤t_enrollment,
|
| 473 |
+
);
|
| 474 |
+
tokio::pin!(connect);
|
| 475 |
+
let (handshake_stream, request_line) = tokio::select! {
|
| 476 |
+
result = &mut connect => panic!("handshake should wait for its response: {result:?}"),
|
| 477 |
+
request = accept_http_request(&listener) => request,
|
| 478 |
+
};
|
| 479 |
+
assert_eq!(
|
| 480 |
+
request_line,
|
| 481 |
+
"GET /backend-api/wham/remote/control/server HTTP/1.1"
|
| 482 |
+
);
|
| 483 |
+
|
| 484 |
+
let pairing = remote_handle.start_pairing(
|
| 485 |
+
RemoteControlPairingStartParams::default(),
|
| 486 |
+
/*app_server_client_name*/ None,
|
| 487 |
+
);
|
| 488 |
+
tokio::pin!(pairing);
|
| 489 |
+
let (pairing_stream, request_line) = tokio::select! {
|
| 490 |
+
result = &mut pairing => panic!("pairing should wait for its response: {result:?}"),
|
| 491 |
+
request = accept_http_request(&listener) => request,
|
| 492 |
+
};
|
| 493 |
+
assert_eq!(
|
| 494 |
+
request_line,
|
| 495 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 496 |
+
);
|
| 497 |
+
respond_with_status_and_headers(pairing_stream, pairing_status, &[], "retry auth").await;
|
| 498 |
+
let (auth_stream, request_line) = tokio::select! {
|
| 499 |
+
result = &mut pairing => panic!("auth should wait for its response: {result:?}"),
|
| 500 |
+
request = accept_http_request(&listener) => request,
|
| 501 |
+
};
|
| 502 |
+
assert_eq!(
|
| 503 |
+
request_line,
|
| 504 |
+
format!("POST /backend-api/wham/remote/control/server/{auth_endpoint} HTTP/1.1")
|
| 505 |
+
);
|
| 506 |
+
|
| 507 |
+
respond_with_status_and_headers(
|
| 508 |
+
handshake_stream,
|
| 509 |
+
"503 Service Unavailable",
|
| 510 |
+
&[("Retry-After", "120")],
|
| 511 |
+
"overloaded",
|
| 512 |
+
)
|
| 513 |
+
.await;
|
| 514 |
+
let connect_error = connect
|
| 515 |
+
.await
|
| 516 |
+
.expect_err("the handshake should report overload");
|
| 517 |
+
let retry_at = remote_control_retry_at(&connect_error)
|
| 518 |
+
.expect("the handshake should preserve its retry deadline");
|
| 519 |
+
save_auth(
|
| 520 |
+
codex_home.path(),
|
| 521 |
+
&remote_control_auth_dot_json("fresh-token"),
|
| 522 |
+
AuthCredentialsStoreMode::File,
|
| 523 |
+
AuthKeyringBackendKind::default(),
|
| 524 |
+
)
|
| 525 |
+
.expect("replacement auth should save");
|
| 526 |
+
respond_with_status_and_headers(auth_stream, "401 Unauthorized", &[], "stale auth").await;
|
| 527 |
+
let pairing_error = timeout(TEST_HTTP_ACCEPT_TIMEOUT, pairing)
|
| 528 |
+
.await
|
| 529 |
+
.expect("auth recovery should report the deadline promptly")
|
| 530 |
+
.expect_err("auth recovery must honor the handshake deadline");
|
| 531 |
+
assert_eq!(remote_control_retry_at(&pairing_error), Some(retry_at));
|
| 532 |
+
timeout(Duration::from_millis(100), listener.accept())
|
| 533 |
+
.await
|
| 534 |
+
.expect_err("auth recovery must not send another request during overload");
|
| 535 |
+
}
|
| 536 |
+
}
|
| 537 |
+
|
| 538 |
+
#[tokio::test]
|
| 539 |
+
async fn pairing_overload_defers_pairing_and_websocket_requests() {
|
| 540 |
+
for status in ["429 Too Many Requests", "503 Service Unavailable"] {
|
| 541 |
+
for check_status in [false, true] {
|
| 542 |
+
for incomplete_body in [false, true] {
|
| 543 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 544 |
+
.await
|
| 545 |
+
.expect("listener should bind");
|
| 546 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 547 |
+
let remote_control_target =
|
| 548 |
+
normalize_remote_control_url(&remote_control_url).expect("target should parse");
|
| 549 |
+
let server_task = tokio::spawn(async move {
|
| 550 |
+
let (mut stream, request_line) = accept_http_request(&listener).await;
|
| 551 |
+
assert_eq!(
|
| 552 |
+
request_line,
|
| 553 |
+
if check_status {
|
| 554 |
+
"POST /backend-api/wham/remote/control/server/pair/status HTTP/1.1"
|
| 555 |
+
} else {
|
| 556 |
+
"POST /backend-api/wham/remote/control/server/pair HTTP/1.1"
|
| 557 |
+
}
|
| 558 |
+
);
|
| 559 |
+
if incomplete_body {
|
| 560 |
+
stream.write_all(format!(
|
| 561 |
+
"HTTP/1.1 {status}\r\nContent-Length: 100\r\nRetry-After: 120\r\nConnection: close\r\n\r\npartial"
|
| 562 |
+
).as_bytes()).await.expect("partial response should send");
|
| 563 |
+
stream.shutdown().await.expect("response should close");
|
| 564 |
+
} else {
|
| 565 |
+
respond_with_status_and_headers(
|
| 566 |
+
stream,
|
| 567 |
+
status,
|
| 568 |
+
&[("Retry-After", "120")],
|
| 569 |
+
"overloaded",
|
| 570 |
+
)
|
| 571 |
+
.await;
|
| 572 |
+
}
|
| 573 |
+
listener
|
| 574 |
+
});
|
| 575 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 576 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 577 |
+
let auth_manager = remote_control_auth_manager();
|
| 578 |
+
let mut remote_handle = remote_control_handle_with_current_enrollment(
|
| 579 |
+
&remote_control_url,
|
| 580 |
+
auth_manager.clone(),
|
| 581 |
+
);
|
| 582 |
+
remote_handle.state_db = Some(state_db.clone());
|
| 583 |
+
let current_enrollment = remote_handle.current_enrollment.clone();
|
| 584 |
+
let status_params = || RemoteControlPairingStatusParams {
|
| 585 |
+
pairing_code: Some("pairing-code".to_string()),
|
| 586 |
+
manual_pairing_code: None,
|
| 587 |
+
};
|
| 588 |
+
let response = if check_status {
|
| 589 |
+
remote_handle
|
| 590 |
+
.pairing_status(status_params())
|
| 591 |
+
.await
|
| 592 |
+
.map(|_| ())
|
| 593 |
+
} else {
|
| 594 |
+
remote_handle
|
| 595 |
+
.start_pairing(
|
| 596 |
+
RemoteControlPairingStartParams::default(),
|
| 597 |
+
/*app_server_client_name*/ None,
|
| 598 |
+
)
|
| 599 |
+
.await
|
| 600 |
+
.map(|_| ())
|
| 601 |
+
};
|
| 602 |
+
let error = response.expect_err("pairing should report overload");
|
| 603 |
+
let listener = server_task.await.expect("server task should succeed");
|
| 604 |
+
let retry_at = remote_control_retry_at(&error)
|
| 605 |
+
.expect("even a partial response must preserve the retry deadline");
|
| 606 |
+
let pairing_error = remote_handle
|
| 607 |
+
.start_pairing(
|
| 608 |
+
RemoteControlPairingStartParams::default(),
|
| 609 |
+
/*app_server_client_name*/ None,
|
| 610 |
+
)
|
| 611 |
+
.await
|
| 612 |
+
.expect_err("pairing must wait for the server deadline");
|
| 613 |
+
let status_error = remote_handle
|
| 614 |
+
.pairing_status(status_params())
|
| 615 |
+
.await
|
| 616 |
+
.expect_err("pairing status must wait for the server deadline");
|
| 617 |
+
let connect_error = connect_test_websocket(
|
| 618 |
+
&remote_control_target,
|
| 619 |
+
state_db.as_ref(),
|
| 620 |
+
&auth_manager,
|
| 621 |
+
¤t_enrollment,
|
| 622 |
+
)
|
| 623 |
+
.await
|
| 624 |
+
.expect_err("the handshake must wait for the same deadline");
|
| 625 |
+
for error in [pairing_error, status_error, connect_error] {
|
| 626 |
+
assert_eq!(remote_control_retry_at(&error), Some(retry_at));
|
| 627 |
+
}
|
| 628 |
+
assert_eq!(
|
| 629 |
+
current_enrollment
|
| 630 |
+
.snapshot()
|
| 631 |
+
.and_then(|enrollment| enrollment.remote_control_token),
|
| 632 |
+
Some(TEST_REMOTE_CONTROL_SERVER_TOKEN.to_string()),
|
| 633 |
+
);
|
| 634 |
+
timeout(Duration::from_millis(100), listener.accept())
|
| 635 |
+
.await
|
| 636 |
+
.expect_err("no remote-control request should bypass the pairing deadline");
|
| 637 |
+
}
|
| 638 |
+
}
|
| 639 |
+
}
|
| 640 |
+
}
|
| 641 |
+
|
| 642 |
+
async fn assert_refresh_failure_blocks_websocket(
|
| 643 |
+
expires_in: time::Duration,
|
| 644 |
+
response_delay: Duration,
|
| 645 |
+
) {
|
| 646 |
+
let listener = TcpListener::bind("127.0.0.1:0")
|
| 647 |
+
.await
|
| 648 |
+
.expect("listener should bind");
|
| 649 |
+
let remote_control_url = remote_control_url_for_listener(&listener);
|
| 650 |
+
let remote_control_target =
|
| 651 |
+
normalize_remote_control_url(&remote_control_url).expect("target should parse");
|
| 652 |
+
let (connects_done_tx, connects_done_rx) = oneshot::channel();
|
| 653 |
+
let server_task = tokio::spawn(async move {
|
| 654 |
+
let (stream, request_line) = accept_http_request(&listener).await;
|
| 655 |
+
assert_eq!(
|
| 656 |
+
request_line,
|
| 657 |
+
"POST /backend-api/wham/remote/control/server/refresh HTTP/1.1"
|
| 658 |
+
);
|
| 659 |
+
tokio::time::sleep(response_delay).await;
|
| 660 |
+
respond_with_status_and_headers(
|
| 661 |
+
stream,
|
| 662 |
+
"502 Bad Gateway",
|
| 663 |
+
&[("retry-after", "120")],
|
| 664 |
+
"upstream unavailable",
|
| 665 |
+
)
|
| 666 |
+
.await;
|
| 667 |
+
assert_no_connection_until_connect_finishes(&listener, connects_done_rx).await;
|
| 668 |
+
});
|
| 669 |
+
let codex_home = TempDir::new().expect("temp dir should create");
|
| 670 |
+
let state_db = remote_control_state_runtime(&codex_home).await;
|
| 671 |
+
let auth_manager = remote_control_auth_manager();
|
| 672 |
+
let mut enrollment = remote_control_enrollment(Some(TEST_REMOTE_CONTROL_SERVER_TOKEN));
|
| 673 |
+
enrollment.expires_at = Some(time::OffsetDateTime::now_utc() + expires_in);
|
| 674 |
+
let current_enrollment = test_current_enrollment(Some(enrollment));
|
| 675 |
+
|
| 676 |
+
let refresh_started_at = time::OffsetDateTime::now_utc();
|
| 677 |
+
let refresh_err = connect_test_websocket(
|
| 678 |
+
&remote_control_target,
|
| 679 |
+
state_db.as_ref(),
|
| 680 |
+
&auth_manager,
|
| 681 |
+
¤t_enrollment,
|
| 682 |
+
)
|
| 683 |
+
.await
|
| 684 |
+
.expect_err("required refresh failure should block websocket connect");
|
| 685 |
+
let refresh_completed_at = time::OffsetDateTime::now_utc();
|
| 686 |
+
let deferred_err = connect_test_websocket(
|
| 687 |
+
&remote_control_target,
|
| 688 |
+
state_db.as_ref(),
|
| 689 |
+
&auth_manager,
|
| 690 |
+
¤t_enrollment,
|
| 691 |
+
)
|
| 692 |
+
.await
|
| 693 |
+
.expect_err("required refresh deadline should block websocket reconnect");
|
| 694 |
+
connects_done_tx
|
| 695 |
+
.send(())
|
| 696 |
+
.expect("server should wait for connect attempts to finish");
|
| 697 |
+
|
| 698 |
+
server_task.await.expect("server task should succeed");
|
| 699 |
+
assert!(refresh_err.to_string().contains("HTTP 502 Bad Gateway"));
|
| 700 |
+
assert_eq!(deferred_err.kind(), io::ErrorKind::WouldBlock);
|
| 701 |
+
let retry_at = remote_control_retry_at(&refresh_err)
|
| 702 |
+
.expect("the overload response should retain its retry deadline");
|
| 703 |
+
assert_eq!(remote_control_retry_at(&deferred_err), Some(retry_at));
|
| 704 |
+
let next_refresh_at = current_enrollment
|
| 705 |
+
.snapshot()
|
| 706 |
+
.and_then(|enrollment| enrollment.next_refresh_at)
|
| 707 |
+
.expect("required refresh failure should set a retry deadline");
|
| 708 |
+
assert!(
|
| 709 |
+
(refresh_started_at + time::Duration::seconds(120)
|
| 710 |
+
..=refresh_completed_at + time::Duration::seconds(150))
|
| 711 |
+
.contains(&next_refresh_at)
|
| 712 |
+
);
|
| 713 |
+
}
|
| 714 |
+
|
| 715 |
+
#[tokio::test]
|
| 716 |
+
async fn expired_token_refresh_failure_throttles_reconnect_without_websocket() {
|
| 717 |
+
assert_refresh_failure_blocks_websocket(-time::Duration::seconds(1), Duration::ZERO).await;
|
| 718 |
+
}
|
| 719 |
+
|
| 720 |
+
#[tokio::test]
|
| 721 |
+
async fn token_expiring_during_refresh_failure_throttles_reconnect_without_websocket() {
|
| 722 |
+
assert_refresh_failure_blocks_websocket(
|
| 723 |
+
time::Duration::seconds(1),
|
| 724 |
+
Duration::from_millis(1_200),
|
| 725 |
+
)
|
| 726 |
+
.await;
|
| 727 |
+
}
|
| 728 |
+
|
| 729 |
+
#[tokio::test]
|
| 730 |
+
async fn websocket_auth_failure_does_not_clear_rotated_server_token() {
|
| 731 |
+
let attempted_enrollment = remote_control_enrollment(Some("old-token"));
|
| 732 |
+
let mut rotated_enrollment = attempted_enrollment.clone();
|
| 733 |
+
rotated_enrollment.remote_control_token = Some("new-token".to_string());
|
| 734 |
+
rotated_enrollment.expires_at =
|
| 735 |
+
Some(time::OffsetDateTime::now_utc() + time::Duration::hours(1));
|
| 736 |
+
let current_enrollment = test_current_enrollment(Some(rotated_enrollment.clone()));
|
| 737 |
+
|
| 738 |
+
clear_remote_control_server_token_if_matches(¤t_enrollment, &attempted_enrollment)
|
| 739 |
+
.await
|
| 740 |
+
.expect("matching enrollment identity should remain available");
|
| 741 |
+
|
| 742 |
+
assert_eq!(current_enrollment.snapshot(), Some(rotated_enrollment));
|
| 743 |
+
}
|
| 744 |
+
|
| 745 |
+
async fn accept_test_websocket(listener: &TcpListener) -> WebSocketStream<TcpStream> {
|
| 746 |
+
let (stream, _) = timeout(TEST_HTTP_ACCEPT_TIMEOUT, listener.accept())
|
| 747 |
+
.await
|
| 748 |
+
.expect("websocket request should arrive in time")
|
| 749 |
+
.expect("listener accept should succeed");
|
| 750 |
+
accept_async(stream)
|
| 751 |
+
.await
|
| 752 |
+
.expect("websocket handshake should succeed")
|
| 753 |
+
}
|
| 754 |
+
|
| 755 |
+
async fn assert_no_connection_until_connect_finishes(
|
| 756 |
+
listener: &TcpListener,
|
| 757 |
+
mut connect_done_rx: oneshot::Receiver<()>,
|
| 758 |
+
) {
|
| 759 |
+
tokio::select! {
|
| 760 |
+
accepted = listener.accept() => {
|
| 761 |
+
accepted.expect("unexpected websocket connection should be accepted");
|
| 762 |
+
panic!("required refresh failure must not proceed to websocket connect");
|
| 763 |
+
}
|
| 764 |
+
connect_done = &mut connect_done_rx => {
|
| 765 |
+
connect_done.expect("connect completion should be reported");
|
| 766 |
+
}
|
| 767 |
+
}
|
| 768 |
+
}
|
codex-rs/app-server-transport/src/transport/stdio.rs
ADDED
|
@@ -0,0 +1,194 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Stdio transport with EOF cleanup and Unix SIGTERM shutdown. Dedicated I/O
|
| 2 |
+
//! threads let SIGTERM abandon blocked pipes without holding the runtime open.
|
| 3 |
+
//! A dedicated signal thread arms the shared EOF/SIGTERM deadline even when
|
| 4 |
+
//! logging or the Tokio runtime is blocked.
|
| 5 |
+
|
| 6 |
+
use super::CHANNEL_CAPACITY;
|
| 7 |
+
use super::ConnectionOrigin;
|
| 8 |
+
use super::TransportEvent;
|
| 9 |
+
use super::forward_incoming_message;
|
| 10 |
+
use super::next_connection_id;
|
| 11 |
+
use super::serialize_outgoing_message;
|
| 12 |
+
use crate::outgoing_message::QueuedOutgoingMessage;
|
| 13 |
+
use codex_app_server_protocol::InitializeParams;
|
| 14 |
+
use codex_app_server_protocol::JSONRPCMessage;
|
| 15 |
+
use codex_app_server_protocol::JSONRPCRequest;
|
| 16 |
+
use std::io::BufRead;
|
| 17 |
+
use std::io::ErrorKind;
|
| 18 |
+
use std::io::Result as IoResult;
|
| 19 |
+
use std::io::Write;
|
| 20 |
+
use tokio::sync::mpsc;
|
| 21 |
+
use tokio::sync::oneshot;
|
| 22 |
+
use tokio::task::JoinHandle;
|
| 23 |
+
use tokio_util::sync::CancellationToken;
|
| 24 |
+
use tracing::debug;
|
| 25 |
+
use tracing::error;
|
| 26 |
+
use tracing::info;
|
| 27 |
+
|
| 28 |
+
pub async fn start_stdio_connection(
|
| 29 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 30 |
+
initialize_client_name_tx: oneshot::Sender<String>,
|
| 31 |
+
install_shutdown_signal_handler: bool,
|
| 32 |
+
) -> IoResult<JoinHandle<()>> {
|
| 33 |
+
let shutdown_signal = CancellationToken::new();
|
| 34 |
+
#[cfg(unix)]
|
| 35 |
+
if install_shutdown_signal_handler {
|
| 36 |
+
use signal_hook::consts::SIGTERM;
|
| 37 |
+
use signal_hook::iterator::Signals;
|
| 38 |
+
|
| 39 |
+
// Register before accepting requests. Receiving SIGTERM and arming the
|
| 40 |
+
// watchdog must not depend on Tokio or synchronous transport logging.
|
| 41 |
+
let mut signals = Signals::new([SIGTERM])?;
|
| 42 |
+
let shutdown_signal = shutdown_signal.clone();
|
| 43 |
+
std::thread::Builder::new()
|
| 44 |
+
.name("app-server-signal".to_string())
|
| 45 |
+
.spawn(move || {
|
| 46 |
+
if signals.forever().next().is_some() {
|
| 47 |
+
start_shutdown_watchdog();
|
| 48 |
+
shutdown_signal.cancel();
|
| 49 |
+
}
|
| 50 |
+
})?;
|
| 51 |
+
}
|
| 52 |
+
let connection_id = next_connection_id();
|
| 53 |
+
let (writer_tx, mut writer_rx) = mpsc::channel::<QueuedOutgoingMessage>(CHANNEL_CAPACITY);
|
| 54 |
+
let writer_tx_for_reader = writer_tx.clone();
|
| 55 |
+
transport_event_tx
|
| 56 |
+
.send(TransportEvent::ConnectionOpened {
|
| 57 |
+
connection_id,
|
| 58 |
+
origin: ConnectionOrigin::Stdio,
|
| 59 |
+
auth: None,
|
| 60 |
+
writer: writer_tx,
|
| 61 |
+
disconnect_sender: None,
|
| 62 |
+
})
|
| 63 |
+
.await
|
| 64 |
+
.map_err(|_| std::io::Error::new(ErrorKind::BrokenPipe, "processor unavailable"))?;
|
| 65 |
+
|
| 66 |
+
// Tokio's stdin uses an uncancellable blocking task, which keeps its runtime
|
| 67 |
+
// alive while the client leaves stdin open. This process-owned thread may
|
| 68 |
+
// remain blocked until process exit, but is not joined by the runtime.
|
| 69 |
+
let (stdin_tx, mut stdin_rx) = mpsc::channel(/*buffer*/ 1);
|
| 70 |
+
std::thread::Builder::new()
|
| 71 |
+
.name("app-server-stdin".to_string())
|
| 72 |
+
.spawn(move || {
|
| 73 |
+
for line in std::io::stdin().lock().lines() {
|
| 74 |
+
if stdin_tx.blocking_send(line).is_err() {
|
| 75 |
+
break;
|
| 76 |
+
}
|
| 77 |
+
}
|
| 78 |
+
})?;
|
| 79 |
+
|
| 80 |
+
// Keep stdout's blocking writes off Tokio's pool too. The forwarding future
|
| 81 |
+
// owns writer_rx so cancelling it also releases producers stuck on a full queue.
|
| 82 |
+
let (stdout_tx, mut stdout_rx) =
|
| 83 |
+
mpsc::channel::<(String, oneshot::Sender<()>)>(/*buffer*/ 1);
|
| 84 |
+
std::thread::Builder::new()
|
| 85 |
+
.name("app-server-stdout".to_string())
|
| 86 |
+
.spawn(move || {
|
| 87 |
+
let mut stdout = std::io::stdout().lock();
|
| 88 |
+
while let Some((json, written_tx)) = stdout_rx.blocking_recv() {
|
| 89 |
+
if let Err(err) = stdout.write_all(json.as_bytes()) {
|
| 90 |
+
error!("Failed to write to stdout: {err}");
|
| 91 |
+
break;
|
| 92 |
+
}
|
| 93 |
+
let _ = written_tx.send(());
|
| 94 |
+
}
|
| 95 |
+
})?;
|
| 96 |
+
|
| 97 |
+
let transport_event_tx_for_reader = transport_event_tx.clone();
|
| 98 |
+
let read_messages = async move {
|
| 99 |
+
let mut initialize_client_name_tx = Some(initialize_client_name_tx);
|
| 100 |
+
while let Some(line) = stdin_rx.recv().await {
|
| 101 |
+
let line = match line {
|
| 102 |
+
Ok(line) => line,
|
| 103 |
+
Err(err) => {
|
| 104 |
+
error!("Failed reading stdin: {err}");
|
| 105 |
+
break;
|
| 106 |
+
}
|
| 107 |
+
};
|
| 108 |
+
if let Some(client_name) = stdio_initialize_client_name(&line)
|
| 109 |
+
&& let Some(initialize_client_name_tx) = initialize_client_name_tx.take()
|
| 110 |
+
{
|
| 111 |
+
let _ = initialize_client_name_tx.send(client_name);
|
| 112 |
+
}
|
| 113 |
+
if !forward_incoming_message(
|
| 114 |
+
&transport_event_tx_for_reader,
|
| 115 |
+
&writer_tx_for_reader,
|
| 116 |
+
connection_id,
|
| 117 |
+
&line,
|
| 118 |
+
)
|
| 119 |
+
.await
|
| 120 |
+
{
|
| 121 |
+
break;
|
| 122 |
+
}
|
| 123 |
+
}
|
| 124 |
+
|
| 125 |
+
// EOF can finish the transport before RPC or runtime cleanup. Start
|
| 126 |
+
// the same process deadline even if no SIGTERM arrives.
|
| 127 |
+
if cfg!(unix) && install_shutdown_signal_handler {
|
| 128 |
+
start_shutdown_watchdog();
|
| 129 |
+
}
|
| 130 |
+
let _ = transport_event_tx_for_reader
|
| 131 |
+
.send(TransportEvent::ConnectionClosed { connection_id })
|
| 132 |
+
.await;
|
| 133 |
+
debug!("stdin reader finished (EOF)");
|
| 134 |
+
};
|
| 135 |
+
|
| 136 |
+
let write_messages = async move {
|
| 137 |
+
while let Some(queued_message) = writer_rx.recv().await {
|
| 138 |
+
let Some(mut json) = serialize_outgoing_message(queued_message.message) else {
|
| 139 |
+
continue;
|
| 140 |
+
};
|
| 141 |
+
json.push('\n');
|
| 142 |
+
let (written_tx, written_rx) = oneshot::channel();
|
| 143 |
+
if stdout_tx.send((json, written_tx)).await.is_err() || written_rx.await.is_err() {
|
| 144 |
+
break;
|
| 145 |
+
}
|
| 146 |
+
if let Some(write_complete_tx) = queued_message.write_complete_tx {
|
| 147 |
+
let _ = write_complete_tx.send(());
|
| 148 |
+
}
|
| 149 |
+
}
|
| 150 |
+
info!("stdout writer exited (channel closed)");
|
| 151 |
+
};
|
| 152 |
+
|
| 153 |
+
Ok(tokio::spawn(async move {
|
| 154 |
+
tokio::select! {
|
| 155 |
+
_ = shutdown_signal.cancelled() => {
|
| 156 |
+
// Cancelling both forwarding futures drops their queues before
|
| 157 |
+
// connection teardown, including when EOF already began draining.
|
| 158 |
+
info!("SIGTERM received; closing stdio connection (45s shutdown deadline)");
|
| 159 |
+
let _ = transport_event_tx
|
| 160 |
+
.send(TransportEvent::ConnectionClosed { connection_id })
|
| 161 |
+
.await;
|
| 162 |
+
}
|
| 163 |
+
_ = async move { tokio::join!(read_messages, write_messages); } => {}
|
| 164 |
+
}
|
| 165 |
+
}))
|
| 166 |
+
}
|
| 167 |
+
|
| 168 |
+
fn start_shutdown_watchdog() {
|
| 169 |
+
// EOF and SIGTERM can both start cleanup; keep the first process deadline.
|
| 170 |
+
static STARTED: std::sync::Once = std::sync::Once::new();
|
| 171 |
+
STARTED.call_once(|| {
|
| 172 |
+
std::thread::Builder::new()
|
| 173 |
+
.name("app-server-shutdown".to_string())
|
| 174 |
+
.spawn(|| {
|
| 175 |
+
// Allow the processor's 30s RPC drain, then bound even Tokio's
|
| 176 |
+
// runtime teardown. Do not log here: stderr may also be blocked.
|
| 177 |
+
std::thread::sleep(std::time::Duration::from_secs(45));
|
| 178 |
+
std::process::exit(/*code*/ 1);
|
| 179 |
+
})
|
| 180 |
+
.unwrap_or_else(|_| std::process::exit(/*code*/ 1));
|
| 181 |
+
});
|
| 182 |
+
}
|
| 183 |
+
|
| 184 |
+
fn stdio_initialize_client_name(line: &str) -> Option<String> {
|
| 185 |
+
let message = serde_json::from_str::<JSONRPCMessage>(line).ok()?;
|
| 186 |
+
let JSONRPCMessage::Request(JSONRPCRequest { method, params, .. }) = message else {
|
| 187 |
+
return None;
|
| 188 |
+
};
|
| 189 |
+
if method != "initialize" {
|
| 190 |
+
return None;
|
| 191 |
+
}
|
| 192 |
+
let params = serde_json::from_value::<InitializeParams>(params?).ok()?;
|
| 193 |
+
Some(params.client_info.name)
|
| 194 |
+
}
|
codex-rs/app-server-transport/src/transport/unix_socket.rs
ADDED
|
@@ -0,0 +1,265 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Control socket startup, guarded rendezvous paths, and WebSocket acceptance.
|
| 2 |
+
|
| 3 |
+
use std::fs::OpenOptions;
|
| 4 |
+
use std::io::ErrorKind;
|
| 5 |
+
use std::io::Result as IoResult;
|
| 6 |
+
use std::path::Path;
|
| 7 |
+
|
| 8 |
+
use super::TransportEvent;
|
| 9 |
+
use crate::transport::websocket::run_websocket_connection;
|
| 10 |
+
use codex_uds::UnixListener;
|
| 11 |
+
use codex_uds::UnixStream;
|
| 12 |
+
use codex_utils_absolute_path::AbsolutePathBuf;
|
| 13 |
+
use futures::SinkExt;
|
| 14 |
+
use futures::StreamExt;
|
| 15 |
+
use tokio::sync::mpsc;
|
| 16 |
+
use tokio::task::JoinHandle;
|
| 17 |
+
use tokio::time::Duration;
|
| 18 |
+
use tokio_tungstenite::accept_hdr_async;
|
| 19 |
+
use tokio_tungstenite::tungstenite::Message;
|
| 20 |
+
use tokio_tungstenite::tungstenite::http::Response;
|
| 21 |
+
use tokio_tungstenite::tungstenite::http::StatusCode;
|
| 22 |
+
use tokio_util::sync::CancellationToken;
|
| 23 |
+
use tracing::error;
|
| 24 |
+
use tracing::info;
|
| 25 |
+
use tracing::warn;
|
| 26 |
+
|
| 27 |
+
#[cfg(unix)]
|
| 28 |
+
const CONTROL_SOCKET_MODE: u32 = 0o600;
|
| 29 |
+
|
| 30 |
+
#[derive(Clone, Copy)]
|
| 31 |
+
pub enum DaemonShutdownAccess {
|
| 32 |
+
Disabled,
|
| 33 |
+
Managed,
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
pub async fn start_control_socket_acceptor(
|
| 37 |
+
socket_path: AbsolutePathBuf,
|
| 38 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 39 |
+
shutdown_token: CancellationToken,
|
| 40 |
+
daemon_shutdown_access: DaemonShutdownAccess,
|
| 41 |
+
) -> IoResult<JoinHandle<()>> {
|
| 42 |
+
#[cfg(windows)]
|
| 43 |
+
let (socket_path, directory_guard) = {
|
| 44 |
+
if let Some(parent) = socket_path.as_path().parent() {
|
| 45 |
+
codex_uds::prepare_private_socket_directory(parent).await?;
|
| 46 |
+
}
|
| 47 |
+
let (path, guard) = codex_uds::validate_private_socket_path(socket_path.as_path())?;
|
| 48 |
+
(AbsolutePathBuf::from_absolute_path_checked(path)?, guard)
|
| 49 |
+
};
|
| 50 |
+
prepare_control_socket_path(socket_path.as_path()).await?;
|
| 51 |
+
let listener = UnixListener::bind(socket_path.as_path()).await?;
|
| 52 |
+
let socket_guard = ControlSocketFileGuard {
|
| 53 |
+
socket_path,
|
| 54 |
+
#[cfg(windows)]
|
| 55 |
+
_directory_guard: directory_guard,
|
| 56 |
+
};
|
| 57 |
+
set_control_socket_permissions(socket_guard.socket_path.as_path()).await?;
|
| 58 |
+
info!(
|
| 59 |
+
socket_path = %socket_guard.socket_path.display(),
|
| 60 |
+
"app-server control socket listening"
|
| 61 |
+
);
|
| 62 |
+
|
| 63 |
+
Ok(tokio::spawn(run_control_socket_acceptor(
|
| 64 |
+
listener,
|
| 65 |
+
transport_event_tx,
|
| 66 |
+
shutdown_token,
|
| 67 |
+
socket_guard,
|
| 68 |
+
daemon_shutdown_access,
|
| 69 |
+
)))
|
| 70 |
+
}
|
| 71 |
+
|
| 72 |
+
async fn run_control_socket_acceptor(
|
| 73 |
+
mut listener: UnixListener,
|
| 74 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 75 |
+
shutdown_token: CancellationToken,
|
| 76 |
+
socket_guard: ControlSocketFileGuard,
|
| 77 |
+
daemon_shutdown_access: DaemonShutdownAccess,
|
| 78 |
+
) {
|
| 79 |
+
let _socket_guard = socket_guard;
|
| 80 |
+
loop {
|
| 81 |
+
let stream = tokio::select! {
|
| 82 |
+
_ = shutdown_token.cancelled() => {
|
| 83 |
+
break;
|
| 84 |
+
}
|
| 85 |
+
result = listener.accept() => {
|
| 86 |
+
match result {
|
| 87 |
+
Ok(stream) => stream,
|
| 88 |
+
Err(err) => {
|
| 89 |
+
if matches!(
|
| 90 |
+
err.kind(),
|
| 91 |
+
ErrorKind::ConnectionAborted | ErrorKind::ConnectionReset | ErrorKind::Interrupted
|
| 92 |
+
) {
|
| 93 |
+
warn!("recoverable control socket accept error: {err}");
|
| 94 |
+
continue;
|
| 95 |
+
}
|
| 96 |
+
error!("control socket accept error: {err}");
|
| 97 |
+
tokio::time::sleep(Duration::from_secs(1)).await;
|
| 98 |
+
continue;
|
| 99 |
+
}
|
| 100 |
+
}
|
| 101 |
+
}
|
| 102 |
+
};
|
| 103 |
+
|
| 104 |
+
let transport_event_tx = transport_event_tx.clone();
|
| 105 |
+
tokio::spawn(async move {
|
| 106 |
+
let mut shutdown_request = false;
|
| 107 |
+
let websocket_stream = match accept_hdr_async(
|
| 108 |
+
stream,
|
| 109 |
+
|request: &tokio_tungstenite::tungstenite::handshake::server::Request, response| {
|
| 110 |
+
if request.uri().path() == "/daemon/shutdown" {
|
| 111 |
+
if !matches!(daemon_shutdown_access, DaemonShutdownAccess::Managed) {
|
| 112 |
+
let mut rejection = Response::new(Some("unmanaged server".to_string()));
|
| 113 |
+
*rejection.status_mut() = StatusCode::FORBIDDEN;
|
| 114 |
+
return Err(rejection);
|
| 115 |
+
}
|
| 116 |
+
shutdown_request = true;
|
| 117 |
+
}
|
| 118 |
+
Ok(response)
|
| 119 |
+
},
|
| 120 |
+
)
|
| 121 |
+
.await
|
| 122 |
+
{
|
| 123 |
+
Ok(websocket_stream) => websocket_stream,
|
| 124 |
+
Err(err) => {
|
| 125 |
+
warn!("failed to upgrade control socket websocket connection: {err}");
|
| 126 |
+
return;
|
| 127 |
+
}
|
| 128 |
+
};
|
| 129 |
+
if shutdown_request {
|
| 130 |
+
run_daemon_shutdown(websocket_stream, transport_event_tx).await;
|
| 131 |
+
return;
|
| 132 |
+
}
|
| 133 |
+
let (websocket_writer, websocket_reader) = websocket_stream.split();
|
| 134 |
+
run_websocket_connection(websocket_writer, websocket_reader, transport_event_tx).await;
|
| 135 |
+
});
|
| 136 |
+
}
|
| 137 |
+
info!("control socket acceptor shutting down");
|
| 138 |
+
}
|
| 139 |
+
|
| 140 |
+
async fn run_daemon_shutdown(
|
| 141 |
+
mut websocket: tokio_tungstenite::WebSocketStream<UnixStream>,
|
| 142 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 143 |
+
) {
|
| 144 |
+
let pid = std::process::id().to_string();
|
| 145 |
+
if !matches!(websocket.next().await, Some(Ok(Message::Text(request))) if request == pid) {
|
| 146 |
+
return;
|
| 147 |
+
}
|
| 148 |
+
if websocket.send(Message::Text(pid.into())).await.is_err() {
|
| 149 |
+
return;
|
| 150 |
+
}
|
| 151 |
+
// Let the manager receive the acknowledgment before the main loop closes connections.
|
| 152 |
+
let _ = tokio::time::timeout(Duration::from_secs(2), websocket.next()).await;
|
| 153 |
+
let _ = transport_event_tx
|
| 154 |
+
.send(TransportEvent::DaemonShutdown)
|
| 155 |
+
.await;
|
| 156 |
+
}
|
| 157 |
+
|
| 158 |
+
pub async fn prepare_control_socket_path(socket_path: &Path) -> IoResult<()> {
|
| 159 |
+
if let Some(parent) = socket_path.parent() {
|
| 160 |
+
codex_uds::prepare_private_socket_directory(parent).await?;
|
| 161 |
+
}
|
| 162 |
+
|
| 163 |
+
#[cfg(windows)]
|
| 164 |
+
let (socket_path, _directory_guard) = codex_uds::validate_private_socket_path(socket_path)?;
|
| 165 |
+
#[cfg(windows)]
|
| 166 |
+
let socket_path = AbsolutePathBuf::from_absolute_path_checked(socket_path)?;
|
| 167 |
+
#[cfg(windows)]
|
| 168 |
+
let socket_path = socket_path.as_path();
|
| 169 |
+
|
| 170 |
+
match UnixStream::connect(socket_path).await {
|
| 171 |
+
Ok(_stream) => {
|
| 172 |
+
return Err(std::io::Error::new(
|
| 173 |
+
ErrorKind::AddrInUse,
|
| 174 |
+
format!(
|
| 175 |
+
"app-server control socket is already in use at {}",
|
| 176 |
+
socket_path.display()
|
| 177 |
+
),
|
| 178 |
+
));
|
| 179 |
+
}
|
| 180 |
+
Err(err) if err.kind() == ErrorKind::NotFound => return Ok(()),
|
| 181 |
+
Err(err) if err.kind() == ErrorKind::ConnectionRefused => {}
|
| 182 |
+
Err(err) => {
|
| 183 |
+
if !socket_path.exists() {
|
| 184 |
+
return Ok(());
|
| 185 |
+
}
|
| 186 |
+
return Err(err);
|
| 187 |
+
}
|
| 188 |
+
}
|
| 189 |
+
|
| 190 |
+
if !socket_path.try_exists()? {
|
| 191 |
+
return Ok(());
|
| 192 |
+
}
|
| 193 |
+
|
| 194 |
+
if !codex_uds::is_stale_socket_path(socket_path).await? {
|
| 195 |
+
return Err(std::io::Error::new(
|
| 196 |
+
ErrorKind::AlreadyExists,
|
| 197 |
+
format!(
|
| 198 |
+
"app-server control socket path exists and is not a socket: {}",
|
| 199 |
+
socket_path.display()
|
| 200 |
+
),
|
| 201 |
+
));
|
| 202 |
+
}
|
| 203 |
+
tokio::fs::remove_file(socket_path).await
|
| 204 |
+
}
|
| 205 |
+
|
| 206 |
+
pub struct AppServerStartupLock {
|
| 207 |
+
_file: std::fs::File,
|
| 208 |
+
}
|
| 209 |
+
|
| 210 |
+
pub async fn acquire_app_server_startup_lock(
|
| 211 |
+
startup_lock_path: AbsolutePathBuf,
|
| 212 |
+
) -> IoResult<AppServerStartupLock> {
|
| 213 |
+
if let Some(parent) = startup_lock_path.as_path().parent() {
|
| 214 |
+
codex_uds::prepare_private_socket_directory(parent).await?;
|
| 215 |
+
}
|
| 216 |
+
tokio::task::spawn_blocking(move || {
|
| 217 |
+
let file = OpenOptions::new()
|
| 218 |
+
.create(true)
|
| 219 |
+
.truncate(false)
|
| 220 |
+
.read(true)
|
| 221 |
+
.write(true)
|
| 222 |
+
.open(startup_lock_path.as_path())?;
|
| 223 |
+
file.lock()?;
|
| 224 |
+
Ok(AppServerStartupLock { _file: file })
|
| 225 |
+
})
|
| 226 |
+
.await
|
| 227 |
+
.map_err(|err| std::io::Error::other(format!("startup lock task failed: {err}")))?
|
| 228 |
+
}
|
| 229 |
+
|
| 230 |
+
#[cfg(unix)]
|
| 231 |
+
async fn set_control_socket_permissions(socket_path: &Path) -> IoResult<()> {
|
| 232 |
+
use std::os::unix::fs::PermissionsExt;
|
| 233 |
+
|
| 234 |
+
tokio::fs::set_permissions(
|
| 235 |
+
socket_path,
|
| 236 |
+
std::fs::Permissions::from_mode(CONTROL_SOCKET_MODE),
|
| 237 |
+
)
|
| 238 |
+
.await
|
| 239 |
+
}
|
| 240 |
+
|
| 241 |
+
#[cfg(not(unix))]
|
| 242 |
+
async fn set_control_socket_permissions(_socket_path: &Path) -> IoResult<()> {
|
| 243 |
+
Ok(())
|
| 244 |
+
}
|
| 245 |
+
|
| 246 |
+
struct ControlSocketFileGuard {
|
| 247 |
+
socket_path: AbsolutePathBuf,
|
| 248 |
+
// Keep the directory pinned until after the socket file is removed in Drop.
|
| 249 |
+
#[cfg(windows)]
|
| 250 |
+
_directory_guard: std::os::windows::io::OwnedHandle,
|
| 251 |
+
}
|
| 252 |
+
|
| 253 |
+
impl Drop for ControlSocketFileGuard {
|
| 254 |
+
fn drop(&mut self) {
|
| 255 |
+
if let Err(err) = std::fs::remove_file(self.socket_path.as_path())
|
| 256 |
+
&& err.kind() != ErrorKind::NotFound
|
| 257 |
+
{
|
| 258 |
+
warn!(
|
| 259 |
+
socket_path = %self.socket_path.display(),
|
| 260 |
+
%err,
|
| 261 |
+
"failed to remove app-server control socket file"
|
| 262 |
+
);
|
| 263 |
+
}
|
| 264 |
+
}
|
| 265 |
+
}
|
codex-rs/app-server-transport/src/transport/unix_socket_tests.rs
ADDED
|
@@ -0,0 +1,340 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::AppServerTransport;
|
| 2 |
+
use super::CHANNEL_CAPACITY;
|
| 3 |
+
use super::DaemonShutdownAccess;
|
| 4 |
+
use super::TransportEvent;
|
| 5 |
+
use super::acquire_app_server_startup_lock;
|
| 6 |
+
use super::app_server_control_socket_path;
|
| 7 |
+
use super::start_control_socket_acceptor;
|
| 8 |
+
use codex_app_server_protocol::JSONRPCMessage;
|
| 9 |
+
use codex_app_server_protocol::JSONRPCNotification;
|
| 10 |
+
use codex_core::config::find_codex_home;
|
| 11 |
+
use codex_uds::UnixStream;
|
| 12 |
+
use codex_utils_absolute_path::AbsolutePathBuf;
|
| 13 |
+
use futures::SinkExt;
|
| 14 |
+
use futures::StreamExt;
|
| 15 |
+
use pretty_assertions::assert_eq;
|
| 16 |
+
use std::io::Result as IoResult;
|
| 17 |
+
use std::path::Path;
|
| 18 |
+
use tokio::sync::mpsc;
|
| 19 |
+
use tokio::time::Duration;
|
| 20 |
+
use tokio::time::timeout;
|
| 21 |
+
use tokio_tungstenite::client_async;
|
| 22 |
+
use tokio_tungstenite::tungstenite::Bytes;
|
| 23 |
+
use tokio_tungstenite::tungstenite::Message as WebSocketMessage;
|
| 24 |
+
use tokio_util::sync::CancellationToken;
|
| 25 |
+
|
| 26 |
+
#[test]
|
| 27 |
+
fn listen_unix_socket_parses_as_unix_socket_transport() {
|
| 28 |
+
assert_eq!(
|
| 29 |
+
AppServerTransport::from_listen_url("unix://"),
|
| 30 |
+
Ok(AppServerTransport::UnixSocket {
|
| 31 |
+
socket_path: default_control_socket_path()
|
| 32 |
+
})
|
| 33 |
+
);
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
#[test]
|
| 37 |
+
fn listen_unix_socket_accepts_absolute_custom_path() {
|
| 38 |
+
assert_eq!(
|
| 39 |
+
AppServerTransport::from_listen_url("unix:///tmp/codex.sock"),
|
| 40 |
+
Ok(AppServerTransport::UnixSocket {
|
| 41 |
+
socket_path: absolute_path("/tmp/codex.sock")
|
| 42 |
+
})
|
| 43 |
+
);
|
| 44 |
+
}
|
| 45 |
+
|
| 46 |
+
#[test]
|
| 47 |
+
fn listen_unix_socket_accepts_relative_custom_path() {
|
| 48 |
+
assert_eq!(
|
| 49 |
+
AppServerTransport::from_listen_url("unix://codex.sock"),
|
| 50 |
+
Ok(AppServerTransport::UnixSocket {
|
| 51 |
+
socket_path: AbsolutePathBuf::relative_to_current_dir("codex.sock")
|
| 52 |
+
.expect("relative path should resolve")
|
| 53 |
+
})
|
| 54 |
+
);
|
| 55 |
+
}
|
| 56 |
+
|
| 57 |
+
#[tokio::test]
|
| 58 |
+
async fn control_socket_acceptor_upgrades_and_forwards_websocket_text_messages_and_pings() {
|
| 59 |
+
let temp_dir = tempfile::TempDir::new().expect("temp dir");
|
| 60 |
+
let socket_path = test_socket_path(temp_dir.path());
|
| 61 |
+
let (transport_event_tx, mut transport_event_rx) =
|
| 62 |
+
mpsc::channel::<TransportEvent>(CHANNEL_CAPACITY);
|
| 63 |
+
let shutdown_token = CancellationToken::new();
|
| 64 |
+
let accept_handle = start_control_socket_acceptor(
|
| 65 |
+
socket_path.clone(),
|
| 66 |
+
transport_event_tx,
|
| 67 |
+
shutdown_token.clone(),
|
| 68 |
+
DaemonShutdownAccess::Disabled,
|
| 69 |
+
)
|
| 70 |
+
.await
|
| 71 |
+
.expect("control socket acceptor should start");
|
| 72 |
+
|
| 73 |
+
let stream = connect_to_socket(socket_path.as_path())
|
| 74 |
+
.await
|
| 75 |
+
.expect("client should connect");
|
| 76 |
+
let (mut websocket, response) = client_async("ws://localhost/rpc", stream)
|
| 77 |
+
.await
|
| 78 |
+
.expect("websocket upgrade should complete");
|
| 79 |
+
assert_eq!(response.status().as_u16(), 101);
|
| 80 |
+
|
| 81 |
+
let opened = timeout(Duration::from_secs(1), transport_event_rx.recv())
|
| 82 |
+
.await
|
| 83 |
+
.expect("connection opened event should arrive")
|
| 84 |
+
.expect("connection opened event");
|
| 85 |
+
let connection_id = match opened {
|
| 86 |
+
TransportEvent::ConnectionOpened { connection_id, .. } => connection_id,
|
| 87 |
+
_ => panic!("expected connection opened event"),
|
| 88 |
+
};
|
| 89 |
+
|
| 90 |
+
let notification = JSONRPCMessage::Notification(JSONRPCNotification {
|
| 91 |
+
method: "initialized".to_string(),
|
| 92 |
+
params: None,
|
| 93 |
+
});
|
| 94 |
+
websocket
|
| 95 |
+
.send(WebSocketMessage::Text(
|
| 96 |
+
serde_json::to_string(¬ification)
|
| 97 |
+
.expect("notification should serialize")
|
| 98 |
+
.into(),
|
| 99 |
+
))
|
| 100 |
+
.await
|
| 101 |
+
.expect("notification should send");
|
| 102 |
+
|
| 103 |
+
let incoming = timeout(Duration::from_secs(1), transport_event_rx.recv())
|
| 104 |
+
.await
|
| 105 |
+
.expect("incoming message event should arrive")
|
| 106 |
+
.expect("incoming message event");
|
| 107 |
+
assert_eq!(
|
| 108 |
+
match incoming {
|
| 109 |
+
TransportEvent::IncomingMessage {
|
| 110 |
+
connection_id: incoming_connection_id,
|
| 111 |
+
message,
|
| 112 |
+
} => (incoming_connection_id, message),
|
| 113 |
+
_ => panic!("expected incoming message event"),
|
| 114 |
+
},
|
| 115 |
+
(connection_id, notification)
|
| 116 |
+
);
|
| 117 |
+
|
| 118 |
+
websocket
|
| 119 |
+
.send(WebSocketMessage::Ping(Bytes::from_static(b"check")))
|
| 120 |
+
.await
|
| 121 |
+
.expect("ping should send");
|
| 122 |
+
let pong = timeout(Duration::from_secs(1), websocket.next())
|
| 123 |
+
.await
|
| 124 |
+
.expect("pong should arrive")
|
| 125 |
+
.expect("pong frame")
|
| 126 |
+
.expect("pong should be valid");
|
| 127 |
+
assert_eq!(pong, WebSocketMessage::Pong(Bytes::from_static(b"check")));
|
| 128 |
+
|
| 129 |
+
websocket.close(None).await.expect("close should send");
|
| 130 |
+
let closed = timeout(Duration::from_secs(1), transport_event_rx.recv())
|
| 131 |
+
.await
|
| 132 |
+
.expect("connection closed event should arrive")
|
| 133 |
+
.expect("connection closed event");
|
| 134 |
+
assert!(matches!(
|
| 135 |
+
closed,
|
| 136 |
+
TransportEvent::ConnectionClosed {
|
| 137 |
+
connection_id: closed_connection_id,
|
| 138 |
+
} if closed_connection_id == connection_id
|
| 139 |
+
));
|
| 140 |
+
|
| 141 |
+
shutdown_token.cancel();
|
| 142 |
+
accept_handle.await.expect("acceptor should join");
|
| 143 |
+
assert_socket_path_removed(socket_path.as_path());
|
| 144 |
+
}
|
| 145 |
+
|
| 146 |
+
#[tokio::test]
|
| 147 |
+
async fn shutdown_is_only_accepted_on_managed_local_socket_for_its_own_pid() {
|
| 148 |
+
let temp_dir = tempfile::TempDir::new().expect("temp dir");
|
| 149 |
+
let socket_path = test_socket_path(temp_dir.path());
|
| 150 |
+
let (tx, mut rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 151 |
+
let shutdown = CancellationToken::new();
|
| 152 |
+
let acceptor = start_control_socket_acceptor(
|
| 153 |
+
socket_path.clone(),
|
| 154 |
+
tx,
|
| 155 |
+
shutdown.clone(),
|
| 156 |
+
DaemonShutdownAccess::Disabled,
|
| 157 |
+
)
|
| 158 |
+
.await
|
| 159 |
+
.expect("acceptor");
|
| 160 |
+
|
| 161 |
+
let stream = connect_to_socket(socket_path.as_path())
|
| 162 |
+
.await
|
| 163 |
+
.expect("connect");
|
| 164 |
+
assert!(
|
| 165 |
+
client_async("ws://localhost/daemon/shutdown", stream)
|
| 166 |
+
.await
|
| 167 |
+
.is_err()
|
| 168 |
+
);
|
| 169 |
+
assert!(rx.try_recv().is_err());
|
| 170 |
+
shutdown.cancel();
|
| 171 |
+
acceptor.await.expect("acceptor shutdown");
|
| 172 |
+
|
| 173 |
+
let (tx, mut rx) = mpsc::channel(CHANNEL_CAPACITY);
|
| 174 |
+
let shutdown = CancellationToken::new();
|
| 175 |
+
let acceptor = start_control_socket_acceptor(
|
| 176 |
+
socket_path.clone(),
|
| 177 |
+
tx,
|
| 178 |
+
shutdown.clone(),
|
| 179 |
+
DaemonShutdownAccess::Managed,
|
| 180 |
+
)
|
| 181 |
+
.await
|
| 182 |
+
.expect("managed acceptor");
|
| 183 |
+
let stream = connect_to_socket(socket_path.as_path())
|
| 184 |
+
.await
|
| 185 |
+
.expect("connect");
|
| 186 |
+
let (mut websocket, _) = client_async("ws://localhost/daemon/shutdown", stream)
|
| 187 |
+
.await
|
| 188 |
+
.expect("upgrade");
|
| 189 |
+
websocket
|
| 190 |
+
.send(WebSocketMessage::Text("0".into()))
|
| 191 |
+
.await
|
| 192 |
+
.expect("wrong pid");
|
| 193 |
+
assert!(!matches!(
|
| 194 |
+
websocket.next().await,
|
| 195 |
+
Some(Ok(WebSocketMessage::Text(_)))
|
| 196 |
+
));
|
| 197 |
+
assert!(rx.try_recv().is_err());
|
| 198 |
+
|
| 199 |
+
let stream = connect_to_socket(socket_path.as_path())
|
| 200 |
+
.await
|
| 201 |
+
.expect("connect");
|
| 202 |
+
let (mut websocket, _) = client_async("ws://localhost/daemon/shutdown", stream)
|
| 203 |
+
.await
|
| 204 |
+
.expect("upgrade");
|
| 205 |
+
let pid = std::process::id().to_string();
|
| 206 |
+
websocket
|
| 207 |
+
.send(WebSocketMessage::Text(pid.clone().into()))
|
| 208 |
+
.await
|
| 209 |
+
.expect("request");
|
| 210 |
+
assert_eq!(
|
| 211 |
+
websocket.next().await.expect("ack").expect("ack frame"),
|
| 212 |
+
WebSocketMessage::Text(pid.into())
|
| 213 |
+
);
|
| 214 |
+
assert!(
|
| 215 |
+
rx.try_recv().is_err(),
|
| 216 |
+
"server must wait until the ack is received"
|
| 217 |
+
);
|
| 218 |
+
websocket.close(None).await.expect("confirm receipt");
|
| 219 |
+
assert!(matches!(
|
| 220 |
+
timeout(Duration::from_secs(2), rx.recv()).await,
|
| 221 |
+
Ok(Some(TransportEvent::DaemonShutdown))
|
| 222 |
+
));
|
| 223 |
+
shutdown.cancel();
|
| 224 |
+
acceptor.await.expect("acceptor shutdown");
|
| 225 |
+
}
|
| 226 |
+
|
| 227 |
+
#[tokio::test]
|
| 228 |
+
async fn app_server_startup_lock_serializes_waiters() {
|
| 229 |
+
let temp_dir = tempfile::TempDir::new().expect("temp dir");
|
| 230 |
+
let lock_path = test_startup_lock_path(temp_dir.path());
|
| 231 |
+
let first_lock = acquire_app_server_startup_lock(lock_path.clone())
|
| 232 |
+
.await
|
| 233 |
+
.expect("first startup lock should succeed");
|
| 234 |
+
let mut second_lock = tokio::spawn(acquire_app_server_startup_lock(lock_path));
|
| 235 |
+
|
| 236 |
+
assert!(
|
| 237 |
+
timeout(Duration::from_millis(100), &mut second_lock)
|
| 238 |
+
.await
|
| 239 |
+
.is_err()
|
| 240 |
+
);
|
| 241 |
+
|
| 242 |
+
drop(first_lock);
|
| 243 |
+
second_lock
|
| 244 |
+
.await
|
| 245 |
+
.expect("second startup lock task should join")
|
| 246 |
+
.expect("second startup lock should succeed");
|
| 247 |
+
}
|
| 248 |
+
|
| 249 |
+
#[cfg(unix)]
|
| 250 |
+
#[tokio::test]
|
| 251 |
+
async fn control_socket_file_is_private_after_bind() {
|
| 252 |
+
use std::os::unix::fs::PermissionsExt;
|
| 253 |
+
|
| 254 |
+
let temp_dir = tempfile::TempDir::new().expect("temp dir");
|
| 255 |
+
let socket_path = test_socket_path(temp_dir.path());
|
| 256 |
+
let (transport_event_tx, _transport_event_rx) =
|
| 257 |
+
mpsc::channel::<TransportEvent>(CHANNEL_CAPACITY);
|
| 258 |
+
let shutdown_token = CancellationToken::new();
|
| 259 |
+
let accept_handle = start_control_socket_acceptor(
|
| 260 |
+
socket_path.clone(),
|
| 261 |
+
transport_event_tx,
|
| 262 |
+
shutdown_token.clone(),
|
| 263 |
+
DaemonShutdownAccess::Disabled,
|
| 264 |
+
)
|
| 265 |
+
.await
|
| 266 |
+
.expect("control socket acceptor should start");
|
| 267 |
+
|
| 268 |
+
let metadata = tokio::fs::metadata(socket_path.as_path())
|
| 269 |
+
.await
|
| 270 |
+
.expect("socket metadata should exist");
|
| 271 |
+
assert_eq!(metadata.permissions().mode() & 0o777, 0o600);
|
| 272 |
+
|
| 273 |
+
shutdown_token.cancel();
|
| 274 |
+
accept_handle.await.expect("acceptor should join");
|
| 275 |
+
}
|
| 276 |
+
|
| 277 |
+
#[cfg(windows)]
|
| 278 |
+
#[tokio::test]
|
| 279 |
+
async fn control_socket_pins_directory_until_shutdown() {
|
| 280 |
+
let temp_dir = tempfile::TempDir::new().expect("temp dir");
|
| 281 |
+
let socket_path = test_socket_path(temp_dir.path());
|
| 282 |
+
let directory = socket_path.as_path().parent().unwrap();
|
| 283 |
+
let moved = temp_dir.path().join("moved");
|
| 284 |
+
let (tx, _rx) = mpsc::channel::<TransportEvent>(CHANNEL_CAPACITY);
|
| 285 |
+
let shutdown = CancellationToken::new();
|
| 286 |
+
let acceptor = start_control_socket_acceptor(
|
| 287 |
+
socket_path.clone(),
|
| 288 |
+
tx,
|
| 289 |
+
shutdown.clone(),
|
| 290 |
+
DaemonShutdownAccess::Disabled,
|
| 291 |
+
)
|
| 292 |
+
.await
|
| 293 |
+
.expect("acceptor");
|
| 294 |
+
assert!(std::fs::rename(directory, &moved).is_err());
|
| 295 |
+
shutdown.cancel();
|
| 296 |
+
acceptor.await.expect("shutdown");
|
| 297 |
+
std::fs::rename(directory, moved).expect("directory unpinned after cleanup");
|
| 298 |
+
}
|
| 299 |
+
|
| 300 |
+
fn absolute_path(path: &str) -> AbsolutePathBuf {
|
| 301 |
+
AbsolutePathBuf::from_absolute_path(path).expect("absolute path")
|
| 302 |
+
}
|
| 303 |
+
|
| 304 |
+
fn default_control_socket_path() -> AbsolutePathBuf {
|
| 305 |
+
let codex_home = find_codex_home().expect("codex home");
|
| 306 |
+
app_server_control_socket_path(&codex_home).expect("default control socket path")
|
| 307 |
+
}
|
| 308 |
+
|
| 309 |
+
fn test_socket_path(temp_dir: &Path) -> AbsolutePathBuf {
|
| 310 |
+
AbsolutePathBuf::from_absolute_path(
|
| 311 |
+
temp_dir
|
| 312 |
+
.join("app-server-control")
|
| 313 |
+
.join("app-server-control.sock"),
|
| 314 |
+
)
|
| 315 |
+
.expect("socket path should resolve")
|
| 316 |
+
}
|
| 317 |
+
|
| 318 |
+
fn test_startup_lock_path(temp_dir: &Path) -> AbsolutePathBuf {
|
| 319 |
+
AbsolutePathBuf::from_absolute_path(
|
| 320 |
+
temp_dir
|
| 321 |
+
.join("app-server-control")
|
| 322 |
+
.join("app-server-startup.lock"),
|
| 323 |
+
)
|
| 324 |
+
.expect("startup lock path should resolve")
|
| 325 |
+
}
|
| 326 |
+
|
| 327 |
+
async fn connect_to_socket(socket_path: &Path) -> IoResult<UnixStream> {
|
| 328 |
+
UnixStream::connect(socket_path).await
|
| 329 |
+
}
|
| 330 |
+
|
| 331 |
+
#[cfg(unix)]
|
| 332 |
+
fn assert_socket_path_removed(socket_path: &Path) {
|
| 333 |
+
assert!(!socket_path.exists());
|
| 334 |
+
}
|
| 335 |
+
|
| 336 |
+
#[cfg(windows)]
|
| 337 |
+
fn assert_socket_path_removed(_socket_path: &Path) {
|
| 338 |
+
// uds_windows uses a regular filesystem path as its rendezvous point,
|
| 339 |
+
// but there is no Unix socket filesystem node to assert on.
|
| 340 |
+
}
|
codex-rs/app-server-transport/src/transport/websocket.rs
ADDED
|
@@ -0,0 +1,389 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::CHANNEL_CAPACITY;
|
| 2 |
+
use super::ConnectionOrigin;
|
| 3 |
+
use super::TransportEvent;
|
| 4 |
+
use super::auth::WebsocketAuthPolicy;
|
| 5 |
+
use super::auth::authorize_upgrade;
|
| 6 |
+
use super::auth::is_unauthenticated_non_loopback_listener;
|
| 7 |
+
use super::forward_incoming_message;
|
| 8 |
+
use super::next_connection_id;
|
| 9 |
+
use super::serialize_outgoing_message;
|
| 10 |
+
use crate::outgoing_message::ConnectionId;
|
| 11 |
+
use crate::outgoing_message::QueuedOutgoingMessage;
|
| 12 |
+
use axum::Router;
|
| 13 |
+
use axum::body::Body;
|
| 14 |
+
use axum::body::Bytes;
|
| 15 |
+
use axum::extract::ConnectInfo;
|
| 16 |
+
use axum::extract::State;
|
| 17 |
+
use axum::extract::ws::Message as AxumWebSocketMessage;
|
| 18 |
+
use axum::extract::ws::WebSocketUpgrade;
|
| 19 |
+
use axum::http::HeaderMap;
|
| 20 |
+
use axum::http::Request;
|
| 21 |
+
use axum::http::StatusCode;
|
| 22 |
+
use axum::http::header::ORIGIN;
|
| 23 |
+
use axum::middleware;
|
| 24 |
+
use axum::middleware::Next;
|
| 25 |
+
use axum::response::IntoResponse;
|
| 26 |
+
use axum::response::Response;
|
| 27 |
+
use axum::routing::any;
|
| 28 |
+
use axum::routing::get;
|
| 29 |
+
use futures::SinkExt;
|
| 30 |
+
use futures::StreamExt;
|
| 31 |
+
use owo_colors::OwoColorize;
|
| 32 |
+
use owo_colors::Stream;
|
| 33 |
+
use owo_colors::Style;
|
| 34 |
+
use std::io::Result as IoResult;
|
| 35 |
+
use std::net::SocketAddr;
|
| 36 |
+
use std::sync::Arc;
|
| 37 |
+
use tokio::net::TcpListener;
|
| 38 |
+
use tokio::sync::mpsc;
|
| 39 |
+
use tokio::task::JoinHandle;
|
| 40 |
+
use tokio_tungstenite::tungstenite::Message as TungsteniteWebSocketMessage;
|
| 41 |
+
use tokio_util::sync::CancellationToken;
|
| 42 |
+
use tracing::error;
|
| 43 |
+
use tracing::info;
|
| 44 |
+
use tracing::warn;
|
| 45 |
+
|
| 46 |
+
/// WebSocket clients can briefly lag behind normal turn output bursts while the
|
| 47 |
+
/// writer task is healthy, so give them more headroom than internal channels.
|
| 48 |
+
const WEBSOCKET_OUTBOUND_CHANNEL_CAPACITY: usize = 32 * 1024;
|
| 49 |
+
const _: () = assert!(WEBSOCKET_OUTBOUND_CHANNEL_CAPACITY > CHANNEL_CAPACITY);
|
| 50 |
+
|
| 51 |
+
fn colorize(text: &str, style: Style) -> String {
|
| 52 |
+
text.if_supports_color(Stream::Stderr, |value| value.style(style))
|
| 53 |
+
.to_string()
|
| 54 |
+
}
|
| 55 |
+
|
| 56 |
+
#[allow(clippy::print_stderr)]
|
| 57 |
+
fn print_websocket_startup_banner(addr: SocketAddr) {
|
| 58 |
+
let title = colorize("codex app-server (WebSockets)", Style::new().bold().cyan());
|
| 59 |
+
let listening_label = colorize("listening on:", Style::new().dimmed());
|
| 60 |
+
let listen_url = colorize(&format!("ws://{addr}"), Style::new().green());
|
| 61 |
+
let ready_label = colorize("readyz:", Style::new().dimmed());
|
| 62 |
+
let ready_url = colorize(&format!("http://{addr}/readyz"), Style::new().green());
|
| 63 |
+
let health_label = colorize("healthz:", Style::new().dimmed());
|
| 64 |
+
let health_url = colorize(&format!("http://{addr}/healthz"), Style::new().green());
|
| 65 |
+
let note_label = colorize("note:", Style::new().dimmed());
|
| 66 |
+
eprintln!("{title}");
|
| 67 |
+
eprintln!(" {listening_label} {listen_url}");
|
| 68 |
+
eprintln!(" {ready_label} {ready_url}");
|
| 69 |
+
eprintln!(" {health_label} {health_url}");
|
| 70 |
+
if addr.ip().is_loopback() {
|
| 71 |
+
eprintln!(
|
| 72 |
+
" {note_label} binds localhost only (use SSH port-forwarding for remote access)"
|
| 73 |
+
);
|
| 74 |
+
} else {
|
| 75 |
+
eprintln!(" {note_label} websocket auth is required for non-localhost listeners");
|
| 76 |
+
}
|
| 77 |
+
}
|
| 78 |
+
|
| 79 |
+
#[derive(Clone)]
|
| 80 |
+
struct WebSocketListenerState {
|
| 81 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 82 |
+
auth_policy: Arc<WebsocketAuthPolicy>,
|
| 83 |
+
}
|
| 84 |
+
|
| 85 |
+
async fn health_check_handler() -> StatusCode {
|
| 86 |
+
StatusCode::OK
|
| 87 |
+
}
|
| 88 |
+
|
| 89 |
+
async fn reject_requests_with_origin_header(
|
| 90 |
+
request: Request<Body>,
|
| 91 |
+
next: Next,
|
| 92 |
+
) -> Result<Response, StatusCode> {
|
| 93 |
+
if request.headers().contains_key(ORIGIN) {
|
| 94 |
+
warn!(
|
| 95 |
+
method = %request.method(),
|
| 96 |
+
uri = %request.uri(),
|
| 97 |
+
"rejecting websocket listener request with Origin header"
|
| 98 |
+
);
|
| 99 |
+
Err(StatusCode::FORBIDDEN)
|
| 100 |
+
} else {
|
| 101 |
+
Ok(next.run(request).await)
|
| 102 |
+
}
|
| 103 |
+
}
|
| 104 |
+
|
| 105 |
+
async fn websocket_upgrade_handler(
|
| 106 |
+
websocket: WebSocketUpgrade,
|
| 107 |
+
ConnectInfo(peer_addr): ConnectInfo<SocketAddr>,
|
| 108 |
+
State(state): State<WebSocketListenerState>,
|
| 109 |
+
headers: HeaderMap,
|
| 110 |
+
) -> impl IntoResponse {
|
| 111 |
+
if let Err(err) = authorize_upgrade(&headers, state.auth_policy.as_ref()) {
|
| 112 |
+
warn!(
|
| 113 |
+
%peer_addr,
|
| 114 |
+
message = err.message(),
|
| 115 |
+
"rejecting websocket client during upgrade"
|
| 116 |
+
);
|
| 117 |
+
return (err.status_code(), err.message()).into_response();
|
| 118 |
+
}
|
| 119 |
+
info!(%peer_addr, "websocket client connected");
|
| 120 |
+
websocket
|
| 121 |
+
.on_upgrade(move |stream| async move {
|
| 122 |
+
let (websocket_writer, websocket_reader) = stream.split();
|
| 123 |
+
run_websocket_connection(websocket_writer, websocket_reader, state.transport_event_tx)
|
| 124 |
+
.await;
|
| 125 |
+
})
|
| 126 |
+
.into_response()
|
| 127 |
+
}
|
| 128 |
+
|
| 129 |
+
pub async fn start_websocket_acceptor(
|
| 130 |
+
bind_address: SocketAddr,
|
| 131 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 132 |
+
shutdown_token: CancellationToken,
|
| 133 |
+
auth_policy: WebsocketAuthPolicy,
|
| 134 |
+
) -> IoResult<JoinHandle<()>> {
|
| 135 |
+
if is_unauthenticated_non_loopback_listener(bind_address, &auth_policy) {
|
| 136 |
+
return Err(std::io::Error::new(
|
| 137 |
+
std::io::ErrorKind::InvalidInput,
|
| 138 |
+
format!(
|
| 139 |
+
"refusing to start non-loopback websocket listener {bind_address} without auth; configure `--ws-auth capability-token` or `--ws-auth signed-bearer-token`"
|
| 140 |
+
),
|
| 141 |
+
));
|
| 142 |
+
}
|
| 143 |
+
let listener = TcpListener::bind(bind_address).await?;
|
| 144 |
+
let local_addr = listener.local_addr()?;
|
| 145 |
+
print_websocket_startup_banner(local_addr);
|
| 146 |
+
info!("app-server websocket listening on ws://{local_addr}");
|
| 147 |
+
|
| 148 |
+
let router = Router::new()
|
| 149 |
+
.route("/readyz", get(health_check_handler))
|
| 150 |
+
.route("/healthz", get(health_check_handler))
|
| 151 |
+
.fallback(any(websocket_upgrade_handler))
|
| 152 |
+
.layer(middleware::from_fn(reject_requests_with_origin_header))
|
| 153 |
+
.with_state(WebSocketListenerState {
|
| 154 |
+
transport_event_tx,
|
| 155 |
+
auth_policy: Arc::new(auth_policy),
|
| 156 |
+
});
|
| 157 |
+
let server = axum::serve(
|
| 158 |
+
listener,
|
| 159 |
+
router.into_make_service_with_connect_info::<SocketAddr>(),
|
| 160 |
+
)
|
| 161 |
+
.with_graceful_shutdown(async move {
|
| 162 |
+
shutdown_token.cancelled().await;
|
| 163 |
+
});
|
| 164 |
+
Ok(tokio::spawn(async move {
|
| 165 |
+
if let Err(err) = server.await {
|
| 166 |
+
error!("websocket acceptor failed: {err}");
|
| 167 |
+
}
|
| 168 |
+
info!("websocket acceptor shutting down");
|
| 169 |
+
}))
|
| 170 |
+
}
|
| 171 |
+
|
| 172 |
+
pub(crate) async fn run_websocket_connection<M, SinkError, StreamError>(
|
| 173 |
+
websocket_writer: impl futures::sink::Sink<M, Error = SinkError> + Send + 'static,
|
| 174 |
+
websocket_reader: impl futures::stream::Stream<Item = Result<M, StreamError>> + Send + 'static,
|
| 175 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 176 |
+
) where
|
| 177 |
+
M: AppServerWebSocketMessage + Send + 'static,
|
| 178 |
+
SinkError: Send + 'static,
|
| 179 |
+
StreamError: std::fmt::Display + Send + 'static,
|
| 180 |
+
{
|
| 181 |
+
let connection_id = next_connection_id();
|
| 182 |
+
let (writer_tx, writer_rx) =
|
| 183 |
+
mpsc::channel::<QueuedOutgoingMessage>(WEBSOCKET_OUTBOUND_CHANNEL_CAPACITY);
|
| 184 |
+
let writer_tx_for_reader = writer_tx.clone();
|
| 185 |
+
let disconnect_token = CancellationToken::new();
|
| 186 |
+
if transport_event_tx
|
| 187 |
+
.send(TransportEvent::ConnectionOpened {
|
| 188 |
+
connection_id,
|
| 189 |
+
origin: ConnectionOrigin::WebSocket,
|
| 190 |
+
auth: None,
|
| 191 |
+
writer: writer_tx,
|
| 192 |
+
disconnect_sender: Some(disconnect_token.clone()),
|
| 193 |
+
})
|
| 194 |
+
.await
|
| 195 |
+
.is_err()
|
| 196 |
+
{
|
| 197 |
+
return;
|
| 198 |
+
}
|
| 199 |
+
|
| 200 |
+
let (writer_control_tx, writer_control_rx) = mpsc::channel::<M>(CHANNEL_CAPACITY);
|
| 201 |
+
let mut outbound_task = tokio::spawn(run_websocket_outbound_loop(
|
| 202 |
+
websocket_writer,
|
| 203 |
+
writer_rx,
|
| 204 |
+
writer_control_rx,
|
| 205 |
+
disconnect_token.clone(),
|
| 206 |
+
));
|
| 207 |
+
let mut inbound_task = tokio::spawn(run_websocket_inbound_loop(
|
| 208 |
+
websocket_reader,
|
| 209 |
+
transport_event_tx.clone(),
|
| 210 |
+
writer_tx_for_reader,
|
| 211 |
+
writer_control_tx,
|
| 212 |
+
connection_id,
|
| 213 |
+
disconnect_token.clone(),
|
| 214 |
+
));
|
| 215 |
+
|
| 216 |
+
tokio::select! {
|
| 217 |
+
_ = &mut outbound_task => {
|
| 218 |
+
disconnect_token.cancel();
|
| 219 |
+
inbound_task.abort();
|
| 220 |
+
}
|
| 221 |
+
_ = &mut inbound_task => {
|
| 222 |
+
disconnect_token.cancel();
|
| 223 |
+
outbound_task.abort();
|
| 224 |
+
}
|
| 225 |
+
}
|
| 226 |
+
|
| 227 |
+
let _ = transport_event_tx
|
| 228 |
+
.send(TransportEvent::ConnectionClosed { connection_id })
|
| 229 |
+
.await;
|
| 230 |
+
}
|
| 231 |
+
|
| 232 |
+
pub(crate) enum IncomingWebSocketMessage {
|
| 233 |
+
Text(String),
|
| 234 |
+
Binary,
|
| 235 |
+
Ping(Bytes),
|
| 236 |
+
Pong,
|
| 237 |
+
Close,
|
| 238 |
+
}
|
| 239 |
+
|
| 240 |
+
/// Converts concrete WebSocket message types into the small message surface the
|
| 241 |
+
/// app-server transport needs, and constructs the only outbound frames it
|
| 242 |
+
/// sends directly.
|
| 243 |
+
pub(crate) trait AppServerWebSocketMessage: Sized {
|
| 244 |
+
fn text(text: String) -> Self;
|
| 245 |
+
fn pong(payload: Bytes) -> Self;
|
| 246 |
+
fn into_incoming(self) -> Option<IncomingWebSocketMessage>;
|
| 247 |
+
}
|
| 248 |
+
|
| 249 |
+
impl AppServerWebSocketMessage for AxumWebSocketMessage {
|
| 250 |
+
fn text(text: String) -> Self {
|
| 251 |
+
Self::Text(text.into())
|
| 252 |
+
}
|
| 253 |
+
|
| 254 |
+
fn pong(payload: Bytes) -> Self {
|
| 255 |
+
Self::Pong(payload)
|
| 256 |
+
}
|
| 257 |
+
|
| 258 |
+
fn into_incoming(self) -> Option<IncomingWebSocketMessage> {
|
| 259 |
+
Some(match self {
|
| 260 |
+
Self::Text(text) => IncomingWebSocketMessage::Text(text.to_string()),
|
| 261 |
+
Self::Binary(_) => IncomingWebSocketMessage::Binary,
|
| 262 |
+
Self::Ping(payload) => IncomingWebSocketMessage::Ping(payload),
|
| 263 |
+
Self::Pong(_) => IncomingWebSocketMessage::Pong,
|
| 264 |
+
Self::Close(_) => IncomingWebSocketMessage::Close,
|
| 265 |
+
})
|
| 266 |
+
}
|
| 267 |
+
}
|
| 268 |
+
|
| 269 |
+
impl AppServerWebSocketMessage for TungsteniteWebSocketMessage {
|
| 270 |
+
fn text(text: String) -> Self {
|
| 271 |
+
Self::Text(text.into())
|
| 272 |
+
}
|
| 273 |
+
|
| 274 |
+
fn pong(payload: Bytes) -> Self {
|
| 275 |
+
Self::Pong(payload)
|
| 276 |
+
}
|
| 277 |
+
|
| 278 |
+
fn into_incoming(self) -> Option<IncomingWebSocketMessage> {
|
| 279 |
+
Some(match self {
|
| 280 |
+
Self::Text(text) => IncomingWebSocketMessage::Text(text.to_string()),
|
| 281 |
+
Self::Binary(_) => IncomingWebSocketMessage::Binary,
|
| 282 |
+
Self::Ping(payload) => IncomingWebSocketMessage::Ping(payload),
|
| 283 |
+
Self::Pong(_) => IncomingWebSocketMessage::Pong,
|
| 284 |
+
Self::Close(_) => IncomingWebSocketMessage::Close,
|
| 285 |
+
Self::Frame(_) => return None,
|
| 286 |
+
})
|
| 287 |
+
}
|
| 288 |
+
}
|
| 289 |
+
|
| 290 |
+
async fn run_websocket_outbound_loop<M, SinkError>(
|
| 291 |
+
websocket_writer: impl futures::sink::Sink<M, Error = SinkError> + Send + 'static,
|
| 292 |
+
mut writer_rx: mpsc::Receiver<QueuedOutgoingMessage>,
|
| 293 |
+
mut writer_control_rx: mpsc::Receiver<M>,
|
| 294 |
+
disconnect_token: CancellationToken,
|
| 295 |
+
) where
|
| 296 |
+
M: AppServerWebSocketMessage + Send + 'static,
|
| 297 |
+
SinkError: Send + 'static,
|
| 298 |
+
{
|
| 299 |
+
tokio::pin!(websocket_writer);
|
| 300 |
+
loop {
|
| 301 |
+
tokio::select! {
|
| 302 |
+
_ = disconnect_token.cancelled() => {
|
| 303 |
+
break;
|
| 304 |
+
}
|
| 305 |
+
message = writer_control_rx.recv() => {
|
| 306 |
+
let Some(message) = message else {
|
| 307 |
+
break;
|
| 308 |
+
};
|
| 309 |
+
if websocket_writer.send(message).await.is_err() {
|
| 310 |
+
break;
|
| 311 |
+
}
|
| 312 |
+
}
|
| 313 |
+
queued_message = writer_rx.recv() => {
|
| 314 |
+
let Some(queued_message) = queued_message else {
|
| 315 |
+
break;
|
| 316 |
+
};
|
| 317 |
+
let Some(json) = serialize_outgoing_message(queued_message.message) else {
|
| 318 |
+
continue;
|
| 319 |
+
};
|
| 320 |
+
if websocket_writer.send(M::text(json)).await.is_err() {
|
| 321 |
+
break;
|
| 322 |
+
}
|
| 323 |
+
if let Some(write_complete_tx) = queued_message.write_complete_tx {
|
| 324 |
+
let _ = write_complete_tx.send(());
|
| 325 |
+
}
|
| 326 |
+
}
|
| 327 |
+
}
|
| 328 |
+
}
|
| 329 |
+
}
|
| 330 |
+
|
| 331 |
+
async fn run_websocket_inbound_loop<M, StreamError>(
|
| 332 |
+
websocket_reader: impl futures::stream::Stream<Item = Result<M, StreamError>> + Send + 'static,
|
| 333 |
+
transport_event_tx: mpsc::Sender<TransportEvent>,
|
| 334 |
+
writer_tx_for_reader: mpsc::Sender<QueuedOutgoingMessage>,
|
| 335 |
+
writer_control_tx: mpsc::Sender<M>,
|
| 336 |
+
connection_id: ConnectionId,
|
| 337 |
+
disconnect_token: CancellationToken,
|
| 338 |
+
) where
|
| 339 |
+
M: AppServerWebSocketMessage + Send + 'static,
|
| 340 |
+
StreamError: std::fmt::Display + Send + 'static,
|
| 341 |
+
{
|
| 342 |
+
tokio::pin!(websocket_reader);
|
| 343 |
+
loop {
|
| 344 |
+
tokio::select! {
|
| 345 |
+
_ = disconnect_token.cancelled() => {
|
| 346 |
+
break;
|
| 347 |
+
}
|
| 348 |
+
incoming_message = websocket_reader.next() => {
|
| 349 |
+
match incoming_message {
|
| 350 |
+
Some(Ok(message)) => match message.into_incoming() {
|
| 351 |
+
Some(IncomingWebSocketMessage::Text(text))
|
| 352 |
+
if !forward_incoming_message(
|
| 353 |
+
&transport_event_tx,
|
| 354 |
+
&writer_tx_for_reader,
|
| 355 |
+
connection_id,
|
| 356 |
+
&text,
|
| 357 |
+
)
|
| 358 |
+
.await
|
| 359 |
+
=> {
|
| 360 |
+
break;
|
| 361 |
+
}
|
| 362 |
+
Some(IncomingWebSocketMessage::Text(_)) => {}
|
| 363 |
+
Some(IncomingWebSocketMessage::Ping(payload)) => {
|
| 364 |
+
match writer_control_tx.try_send(M::pong(payload)) {
|
| 365 |
+
Ok(()) => {}
|
| 366 |
+
Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => break,
|
| 367 |
+
Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => {
|
| 368 |
+
warn!("websocket control queue full while replying to ping; closing connection");
|
| 369 |
+
break;
|
| 370 |
+
}
|
| 371 |
+
}
|
| 372 |
+
}
|
| 373 |
+
Some(IncomingWebSocketMessage::Pong) => {}
|
| 374 |
+
Some(IncomingWebSocketMessage::Close) => break,
|
| 375 |
+
Some(IncomingWebSocketMessage::Binary) => {
|
| 376 |
+
warn!("dropping unsupported binary websocket message");
|
| 377 |
+
}
|
| 378 |
+
None => {}
|
| 379 |
+
},
|
| 380 |
+
None => break,
|
| 381 |
+
Some(Err(err)) => {
|
| 382 |
+
warn!("websocket receive error: {err}");
|
| 383 |
+
break;
|
| 384 |
+
}
|
| 385 |
+
}
|
| 386 |
+
}
|
| 387 |
+
}
|
| 388 |
+
}
|
| 389 |
+
}
|
codex-rs/bwrap/src/main.rs
ADDED
|
@@ -0,0 +1,45 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#[cfg(all(target_os = "linux", bwrap_available))]
|
| 2 |
+
fn main() {
|
| 3 |
+
use std::ffi::CStr;
|
| 4 |
+
use std::ffi::CString;
|
| 5 |
+
use std::os::raw::c_char;
|
| 6 |
+
use std::os::unix::ffi::OsStrExt;
|
| 7 |
+
|
| 8 |
+
unsafe extern "C" {
|
| 9 |
+
fn bwrap_main(argc: libc::c_int, argv: *const *const c_char) -> libc::c_int;
|
| 10 |
+
}
|
| 11 |
+
|
| 12 |
+
let cstrings = std::env::args_os()
|
| 13 |
+
.map(|arg| {
|
| 14 |
+
CString::new(arg.as_os_str().as_bytes())
|
| 15 |
+
.unwrap_or_else(|err| panic!("failed to convert argv to CString: {err}"))
|
| 16 |
+
})
|
| 17 |
+
.collect::<Vec<_>>();
|
| 18 |
+
let mut argv_ptrs = cstrings
|
| 19 |
+
.iter()
|
| 20 |
+
.map(CString::as_c_str)
|
| 21 |
+
.map(CStr::as_ptr)
|
| 22 |
+
.collect::<Vec<*const c_char>>();
|
| 23 |
+
argv_ptrs.push(std::ptr::null());
|
| 24 |
+
|
| 25 |
+
// SAFETY: We provide a null-terminated argv vector whose pointers remain
|
| 26 |
+
// valid for the duration of the call.
|
| 27 |
+
let exit_code = unsafe { bwrap_main(cstrings.len() as libc::c_int, argv_ptrs.as_ptr()) };
|
| 28 |
+
std::process::exit(exit_code);
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
#[cfg(all(target_os = "linux", not(bwrap_available)))]
|
| 32 |
+
fn main() {
|
| 33 |
+
panic!(
|
| 34 |
+
r#"bubblewrap is not available in this build.
|
| 35 |
+
Notes:
|
| 36 |
+
- ensure the target OS is Linux
|
| 37 |
+
- libcap headers must be available via pkg-config
|
| 38 |
+
- bubblewrap sources expected at codex-rs/vendor/bubblewrap (default)"#
|
| 39 |
+
);
|
| 40 |
+
}
|
| 41 |
+
|
| 42 |
+
#[cfg(not(target_os = "linux"))]
|
| 43 |
+
fn main() {
|
| 44 |
+
panic!("bwrap is only supported on Linux");
|
| 45 |
+
}
|
codex-rs/codex-mcp/src/agent_plugin_config.rs
ADDED
|
@@ -0,0 +1,532 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::PluginMcpConfigParseOutcome;
|
| 2 |
+
use super::PluginMcpServerParseError;
|
| 3 |
+
use codex_config::McpServerConfig;
|
| 4 |
+
use serde::Deserialize;
|
| 5 |
+
use serde_json::Map as JsonMap;
|
| 6 |
+
use serde_json::Value as JsonValue;
|
| 7 |
+
use std::collections::BTreeMap;
|
| 8 |
+
use std::ffi::OsString;
|
| 9 |
+
use std::path::Path;
|
| 10 |
+
use std::path::PathBuf;
|
| 11 |
+
use url::Host;
|
| 12 |
+
|
| 13 |
+
// Published Agent Plugins v1 MCP schema:
|
| 14 |
+
// https://github.com/agentplugins/agent-plugins-spec/blob/main/schemas/1.0.0/mcp.schema.json
|
| 15 |
+
const AGENT_PLUGIN_MCP_SCHEMA_URI: &str = "https://agent-plugins.org/schemas/1.0.0/mcp.schema.json";
|
| 16 |
+
const SUPPORTED_AGENT_PLUGIN_MCP_SCHEMA_URIS: &[&str] = &[AGENT_PLUGIN_MCP_SCHEMA_URI];
|
| 17 |
+
const PLUGIN_ROOT_VARIABLE: &str = "PLUGIN_ROOT";
|
| 18 |
+
const PLUGIN_DATA_VARIABLE: &str = "PLUGIN_DATA";
|
| 19 |
+
const CLIENT_OWNED_HTTP_HEADERS: &[&str] = &[
|
| 20 |
+
"accept",
|
| 21 |
+
"authorization",
|
| 22 |
+
"connection",
|
| 23 |
+
"content-encoding",
|
| 24 |
+
"content-length",
|
| 25 |
+
"content-type",
|
| 26 |
+
"host",
|
| 27 |
+
"last-event-id",
|
| 28 |
+
"mcp-protocol-version",
|
| 29 |
+
"mcp-session-id",
|
| 30 |
+
"proxy-authorization",
|
| 31 |
+
"te",
|
| 32 |
+
"trailer",
|
| 33 |
+
"transfer-encoding",
|
| 34 |
+
"upgrade",
|
| 35 |
+
"user-agent",
|
| 36 |
+
];
|
| 37 |
+
|
| 38 |
+
#[derive(Debug, Deserialize)]
|
| 39 |
+
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
| 40 |
+
struct AgentPluginMcpFile {
|
| 41 |
+
#[serde(rename = "$schema")]
|
| 42 |
+
schema: String,
|
| 43 |
+
mcp_servers: BTreeMap<String, JsonValue>,
|
| 44 |
+
}
|
| 45 |
+
|
| 46 |
+
#[derive(Debug, Deserialize)]
|
| 47 |
+
#[serde(tag = "type", deny_unknown_fields)]
|
| 48 |
+
enum AgentPluginMcpServer {
|
| 49 |
+
#[serde(rename = "stdio")]
|
| 50 |
+
Stdio {
|
| 51 |
+
command: String,
|
| 52 |
+
#[serde(default)]
|
| 53 |
+
args: Vec<String>,
|
| 54 |
+
#[serde(default)]
|
| 55 |
+
env: BTreeMap<String, String>,
|
| 56 |
+
cwd: Option<String>,
|
| 57 |
+
},
|
| 58 |
+
#[serde(rename = "streamable-http")]
|
| 59 |
+
StreamableHttp {
|
| 60 |
+
url: String,
|
| 61 |
+
headers: Option<BTreeMap<String, String>>,
|
| 62 |
+
},
|
| 63 |
+
#[serde(rename = "sse")]
|
| 64 |
+
Sse {
|
| 65 |
+
#[serde(rename = "url")]
|
| 66 |
+
_url: String,
|
| 67 |
+
#[serde(rename = "headers")]
|
| 68 |
+
_headers: Option<BTreeMap<String, String>>,
|
| 69 |
+
},
|
| 70 |
+
}
|
| 71 |
+
|
| 72 |
+
/// Translates an Agent Plugins `mcp.json` into Codex MCP configuration.
|
| 73 |
+
pub fn parse_agent_plugin_mcp_config(
|
| 74 |
+
plugin_root: &Path,
|
| 75 |
+
plugin_data_root: &Path,
|
| 76 |
+
contents: &str,
|
| 77 |
+
) -> Result<PluginMcpConfigParseOutcome, serde_json::Error> {
|
| 78 |
+
parse_agent_plugin_mcp_config_from(contents, plugin_root, plugin_data_root)
|
| 79 |
+
}
|
| 80 |
+
|
| 81 |
+
fn parse_agent_plugin_mcp_config_from(
|
| 82 |
+
contents: &str,
|
| 83 |
+
plugin_root: &Path,
|
| 84 |
+
plugin_data_root: &Path,
|
| 85 |
+
) -> Result<PluginMcpConfigParseOutcome, serde_json::Error> {
|
| 86 |
+
let AgentPluginMcpFile {
|
| 87 |
+
schema,
|
| 88 |
+
mcp_servers,
|
| 89 |
+
} = serde_json::from_str(contents)?;
|
| 90 |
+
if !SUPPORTED_AGENT_PLUGIN_MCP_SCHEMA_URIS.contains(&schema.as_str()) {
|
| 91 |
+
return Err(plugin_mcp_json_error(format!(
|
| 92 |
+
"unsupported Agent Plugins MCP schema `{schema}`; supported schemas: {}",
|
| 93 |
+
SUPPORTED_AGENT_PLUGIN_MCP_SCHEMA_URIS.join(", ")
|
| 94 |
+
)));
|
| 95 |
+
}
|
| 96 |
+
|
| 97 |
+
let mut outcome = PluginMcpConfigParseOutcome::default();
|
| 98 |
+
for (name, value) in mcp_servers {
|
| 99 |
+
match normalize_agent_plugin_mcp_server(value, plugin_root, plugin_data_root) {
|
| 100 |
+
Ok(config) => {
|
| 101 |
+
outcome.servers.insert(name, config);
|
| 102 |
+
}
|
| 103 |
+
Err(message) => outcome
|
| 104 |
+
.errors
|
| 105 |
+
.push(PluginMcpServerParseError { name, message }),
|
| 106 |
+
}
|
| 107 |
+
}
|
| 108 |
+
Ok(outcome)
|
| 109 |
+
}
|
| 110 |
+
|
| 111 |
+
fn normalize_agent_plugin_mcp_server(
|
| 112 |
+
value: JsonValue,
|
| 113 |
+
plugin_root: &Path,
|
| 114 |
+
plugin_data_root: &Path,
|
| 115 |
+
) -> Result<McpServerConfig, String> {
|
| 116 |
+
let object = value
|
| 117 |
+
.as_object()
|
| 118 |
+
.ok_or_else(|| "Agent Plugins MCP server must be an object".to_string())?;
|
| 119 |
+
match object.get("type").and_then(JsonValue::as_str) {
|
| 120 |
+
Some("stdio") => reject_explicit_null(object, "cwd")?,
|
| 121 |
+
Some("streamable-http" | "sse") => reject_explicit_null(object, "headers")?,
|
| 122 |
+
_ => {}
|
| 123 |
+
}
|
| 124 |
+
let server =
|
| 125 |
+
serde_json::from_value::<AgentPluginMcpServer>(value).map_err(|err| err.to_string())?;
|
| 126 |
+
let object = match server {
|
| 127 |
+
AgentPluginMcpServer::Stdio {
|
| 128 |
+
command,
|
| 129 |
+
args,
|
| 130 |
+
env,
|
| 131 |
+
cwd,
|
| 132 |
+
} => normalize_agent_plugin_stdio_server(
|
| 133 |
+
command,
|
| 134 |
+
args,
|
| 135 |
+
env,
|
| 136 |
+
cwd,
|
| 137 |
+
plugin_root,
|
| 138 |
+
plugin_data_root,
|
| 139 |
+
)?,
|
| 140 |
+
AgentPluginMcpServer::StreamableHttp { url, headers } => {
|
| 141 |
+
normalize_agent_plugin_http_server(url, headers)?
|
| 142 |
+
}
|
| 143 |
+
AgentPluginMcpServer::Sse { .. } => {
|
| 144 |
+
return Err("Agent Plugins legacy SSE transport is not supported by Codex".to_string());
|
| 145 |
+
}
|
| 146 |
+
};
|
| 147 |
+
serde_json::from_value(JsonValue::Object(object)).map_err(|err| err.to_string())
|
| 148 |
+
}
|
| 149 |
+
|
| 150 |
+
fn normalize_agent_plugin_stdio_server(
|
| 151 |
+
mut command: String,
|
| 152 |
+
mut args: Vec<String>,
|
| 153 |
+
mut env: BTreeMap<String, String>,
|
| 154 |
+
cwd: Option<String>,
|
| 155 |
+
plugin_root: &Path,
|
| 156 |
+
plugin_data_root: &Path,
|
| 157 |
+
) -> Result<JsonMap<String, JsonValue>, String> {
|
| 158 |
+
#[cfg(windows)]
|
| 159 |
+
let has_windows_path_prefix = matches!(
|
| 160 |
+
Path::new(&command).components().next(),
|
| 161 |
+
Some(std::path::Component::Prefix(_))
|
| 162 |
+
);
|
| 163 |
+
#[cfg(not(windows))]
|
| 164 |
+
let has_windows_path_prefix = false;
|
| 165 |
+
let is_bare_command = !command.is_empty()
|
| 166 |
+
&& !command.contains('/')
|
| 167 |
+
&& !command.contains('\\')
|
| 168 |
+
&& !has_windows_path_prefix;
|
| 169 |
+
let is_plugin_relative_command =
|
| 170 |
+
command.starts_with("./") && is_portable_relative_path(&command);
|
| 171 |
+
if !is_bare_command && !is_plugin_relative_command {
|
| 172 |
+
return Err(
|
| 173 |
+
"Agent Plugins stdio command must be a bare executable name or a contained `./` path"
|
| 174 |
+
.to_string(),
|
| 175 |
+
);
|
| 176 |
+
}
|
| 177 |
+
for reserved in [PLUGIN_ROOT_VARIABLE, PLUGIN_DATA_VARIABLE] {
|
| 178 |
+
if env
|
| 179 |
+
.keys()
|
| 180 |
+
.any(|name| environment_variable_names_match(name, reserved))
|
| 181 |
+
{
|
| 182 |
+
return Err(format!(
|
| 183 |
+
"Agent Plugins stdio `env` cannot override reserved variable `{reserved}`"
|
| 184 |
+
));
|
| 185 |
+
}
|
| 186 |
+
}
|
| 187 |
+
#[cfg(windows)]
|
| 188 |
+
{
|
| 189 |
+
let mut normalized_env = BTreeMap::new();
|
| 190 |
+
for (name, value) in env {
|
| 191 |
+
let normalized_name = name.to_ascii_uppercase();
|
| 192 |
+
if normalized_env.insert(normalized_name, value).is_some() {
|
| 193 |
+
return Err(format!(
|
| 194 |
+
"duplicate case-insensitive Agent Plugins environment variable `{name}`"
|
| 195 |
+
));
|
| 196 |
+
}
|
| 197 |
+
}
|
| 198 |
+
env = normalized_env;
|
| 199 |
+
}
|
| 200 |
+
|
| 201 |
+
let root_path = absolute_plugin_path(plugin_root)?;
|
| 202 |
+
let data_root_path = absolute_plugin_path(plugin_data_root)?;
|
| 203 |
+
let root = host_path_string(&root_path);
|
| 204 |
+
let data_root = host_path_string(&data_root_path);
|
| 205 |
+
if command.starts_with("./") {
|
| 206 |
+
command = host_path_string(&resolve_contained_host_path(
|
| 207 |
+
&command, &root_path, &root_path,
|
| 208 |
+
)?);
|
| 209 |
+
}
|
| 210 |
+
for arg in &mut args {
|
| 211 |
+
*arg = expand_agent_plugin_placeholders(arg, &root, &data_root);
|
| 212 |
+
}
|
| 213 |
+
for value in env.values_mut() {
|
| 214 |
+
*value = expand_agent_plugin_placeholders(value, &root, &data_root);
|
| 215 |
+
}
|
| 216 |
+
|
| 217 |
+
let cwd = cwd.as_deref().unwrap_or("${PLUGIN_ROOT}");
|
| 218 |
+
let Some(cwd_root) = parse_agent_plugin_cwd(cwd) else {
|
| 219 |
+
return Err(
|
| 220 |
+
"Agent Plugins stdio `cwd` must be a contained `./`, `${PLUGIN_ROOT}`, or `${PLUGIN_DATA}` path"
|
| 221 |
+
.to_string(),
|
| 222 |
+
);
|
| 223 |
+
};
|
| 224 |
+
let cwd = expand_agent_plugin_placeholders(cwd, &root, &data_root);
|
| 225 |
+
let cwd_root = match cwd_root {
|
| 226 |
+
AgentPluginCwdRoot::Package => &root_path,
|
| 227 |
+
AgentPluginCwdRoot::Data => &data_root_path,
|
| 228 |
+
};
|
| 229 |
+
env.insert(PLUGIN_ROOT_VARIABLE.to_string(), root);
|
| 230 |
+
env.insert(PLUGIN_DATA_VARIABLE.to_string(), data_root);
|
| 231 |
+
|
| 232 |
+
Ok(JsonMap::from_iter([
|
| 233 |
+
("command".to_string(), JsonValue::String(command)),
|
| 234 |
+
(
|
| 235 |
+
"args".to_string(),
|
| 236 |
+
JsonValue::Array(args.into_iter().map(JsonValue::String).collect()),
|
| 237 |
+
),
|
| 238 |
+
("env".to_string(), string_map_value(env)),
|
| 239 |
+
(
|
| 240 |
+
"cwd".to_string(),
|
| 241 |
+
JsonValue::String(host_path_string(&resolve_contained_host_path(
|
| 242 |
+
&cwd, cwd_root, cwd_root,
|
| 243 |
+
)?)),
|
| 244 |
+
),
|
| 245 |
+
]))
|
| 246 |
+
}
|
| 247 |
+
|
| 248 |
+
fn reject_explicit_null(object: &JsonMap<String, JsonValue>, field: &str) -> Result<(), String> {
|
| 249 |
+
if object.get(field).is_some_and(JsonValue::is_null) {
|
| 250 |
+
return Err(format!(
|
| 251 |
+
"Agent Plugins MCP `{field}` must use its declared type when present"
|
| 252 |
+
));
|
| 253 |
+
}
|
| 254 |
+
Ok(())
|
| 255 |
+
}
|
| 256 |
+
|
| 257 |
+
fn environment_variable_names_match(left: &str, right: &str) -> bool {
|
| 258 |
+
if cfg!(windows) {
|
| 259 |
+
left.eq_ignore_ascii_case(right)
|
| 260 |
+
} else {
|
| 261 |
+
left == right
|
| 262 |
+
}
|
| 263 |
+
}
|
| 264 |
+
|
| 265 |
+
fn normalize_agent_plugin_http_server(
|
| 266 |
+
url: String,
|
| 267 |
+
mut headers: Option<BTreeMap<String, String>>,
|
| 268 |
+
) -> Result<JsonMap<String, JsonValue>, String> {
|
| 269 |
+
validate_agent_plugin_url(&url)?;
|
| 270 |
+
if let Some(configured_headers) = headers.as_mut() {
|
| 271 |
+
validate_agent_plugin_headers(configured_headers)?;
|
| 272 |
+
configured_headers.retain(|name, _| {
|
| 273 |
+
!CLIENT_OWNED_HTTP_HEADERS
|
| 274 |
+
.iter()
|
| 275 |
+
.any(|owned| name.eq_ignore_ascii_case(owned))
|
| 276 |
+
});
|
| 277 |
+
}
|
| 278 |
+
let mut object = JsonMap::from_iter([("url".to_string(), JsonValue::String(url))]);
|
| 279 |
+
if let Some(headers) = headers.filter(|headers| !headers.is_empty()) {
|
| 280 |
+
object.insert("http_headers".to_string(), string_map_value(headers));
|
| 281 |
+
}
|
| 282 |
+
Ok(object)
|
| 283 |
+
}
|
| 284 |
+
|
| 285 |
+
fn validate_agent_plugin_url(raw_url: &str) -> Result<(), String> {
|
| 286 |
+
if raw_url.is_empty() {
|
| 287 |
+
return Err("Agent Plugins HTTP server requires a non-empty `url`".to_string());
|
| 288 |
+
}
|
| 289 |
+
let parsed = url::Url::parse(raw_url)
|
| 290 |
+
.map_err(|err| format!("invalid Agent Plugins MCP URL `{raw_url}`: {err}"))?;
|
| 291 |
+
if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() {
|
| 292 |
+
return Err("Agent Plugins MCP URL must be absolute HTTP or HTTPS".to_string());
|
| 293 |
+
}
|
| 294 |
+
if !parsed.username().is_empty() || parsed.password().is_some() || parsed.fragment().is_some() {
|
| 295 |
+
return Err(
|
| 296 |
+
"Agent Plugins MCP URL must not contain user information or a fragment".to_string(),
|
| 297 |
+
);
|
| 298 |
+
}
|
| 299 |
+
let is_loopback = match parsed.host() {
|
| 300 |
+
Some(Host::Domain(host)) => host == "localhost",
|
| 301 |
+
Some(Host::Ipv4(address)) => address.is_loopback(),
|
| 302 |
+
Some(Host::Ipv6(address)) => address.is_loopback(),
|
| 303 |
+
None => false,
|
| 304 |
+
};
|
| 305 |
+
if parsed.scheme() == "http" && !is_loopback {
|
| 306 |
+
return Err("non-loopback Agent Plugins MCP endpoints must use HTTPS".to_string());
|
| 307 |
+
}
|
| 308 |
+
Ok(())
|
| 309 |
+
}
|
| 310 |
+
|
| 311 |
+
fn validate_agent_plugin_headers(headers: &BTreeMap<String, String>) -> Result<(), String> {
|
| 312 |
+
let mut seen = std::collections::HashSet::new();
|
| 313 |
+
for (name, value) in headers {
|
| 314 |
+
if !seen.insert(name.to_ascii_lowercase()) {
|
| 315 |
+
return Err(format!(
|
| 316 |
+
"duplicate case-insensitive Agent Plugins HTTP header `{name}`"
|
| 317 |
+
));
|
| 318 |
+
}
|
| 319 |
+
if !is_valid_http_header_name(name) {
|
| 320 |
+
return Err(format!("invalid Agent Plugins HTTP header name `{name}`"));
|
| 321 |
+
}
|
| 322 |
+
if value
|
| 323 |
+
.bytes()
|
| 324 |
+
.any(|byte| (byte < 32 && byte != b'\t') || byte == 127)
|
| 325 |
+
{
|
| 326 |
+
return Err(format!(
|
| 327 |
+
"invalid Agent Plugins HTTP header value for `{name}`"
|
| 328 |
+
));
|
| 329 |
+
}
|
| 330 |
+
}
|
| 331 |
+
Ok(())
|
| 332 |
+
}
|
| 333 |
+
|
| 334 |
+
fn string_map_value(values: BTreeMap<String, String>) -> JsonValue {
|
| 335 |
+
JsonValue::Object(
|
| 336 |
+
values
|
| 337 |
+
.into_iter()
|
| 338 |
+
.map(|(name, value)| (name, JsonValue::String(value)))
|
| 339 |
+
.collect(),
|
| 340 |
+
)
|
| 341 |
+
}
|
| 342 |
+
|
| 343 |
+
#[derive(Clone, Copy, Debug)]
|
| 344 |
+
enum AgentPluginCwdRoot {
|
| 345 |
+
Package,
|
| 346 |
+
Data,
|
| 347 |
+
}
|
| 348 |
+
|
| 349 |
+
fn parse_agent_plugin_cwd(value: &str) -> Option<AgentPluginCwdRoot> {
|
| 350 |
+
if value == "./" {
|
| 351 |
+
return Some(AgentPluginCwdRoot::Package);
|
| 352 |
+
}
|
| 353 |
+
if let Some(relative) = value.strip_prefix("./")
|
| 354 |
+
&& is_portable_path_suffix(relative)
|
| 355 |
+
{
|
| 356 |
+
return Some(AgentPluginCwdRoot::Package);
|
| 357 |
+
}
|
| 358 |
+
for (placeholder, root) in [
|
| 359 |
+
("${PLUGIN_ROOT}", AgentPluginCwdRoot::Package),
|
| 360 |
+
("${PLUGIN_DATA}", AgentPluginCwdRoot::Data),
|
| 361 |
+
] {
|
| 362 |
+
if value == placeholder {
|
| 363 |
+
return Some(root);
|
| 364 |
+
}
|
| 365 |
+
if let Some(relative) = value.strip_prefix(&format!("{placeholder}/"))
|
| 366 |
+
&& (relative.is_empty() || is_portable_path_suffix(relative))
|
| 367 |
+
{
|
| 368 |
+
return Some(root);
|
| 369 |
+
}
|
| 370 |
+
}
|
| 371 |
+
None
|
| 372 |
+
}
|
| 373 |
+
|
| 374 |
+
fn expand_agent_plugin_placeholders(value: &str, plugin_root: &str, plugin_data: &str) -> String {
|
| 375 |
+
const ROOT: &str = "${PLUGIN_ROOT}";
|
| 376 |
+
const DATA: &str = "${PLUGIN_DATA}";
|
| 377 |
+
let mut output = String::with_capacity(value.len());
|
| 378 |
+
let mut remaining = value;
|
| 379 |
+
loop {
|
| 380 |
+
let next = match (remaining.find(ROOT), remaining.find(DATA)) {
|
| 381 |
+
(Some(root), Some(data)) if root <= data => Some((root, ROOT, plugin_root)),
|
| 382 |
+
(Some(_), Some(data)) => Some((data, DATA, plugin_data)),
|
| 383 |
+
(Some(root), None) => Some((root, ROOT, plugin_root)),
|
| 384 |
+
(None, Some(data)) => Some((data, DATA, plugin_data)),
|
| 385 |
+
(None, None) => None,
|
| 386 |
+
};
|
| 387 |
+
let Some((index, placeholder, replacement)) = next else {
|
| 388 |
+
output.push_str(remaining);
|
| 389 |
+
break;
|
| 390 |
+
};
|
| 391 |
+
output.push_str(&remaining[..index]);
|
| 392 |
+
output.push_str(replacement);
|
| 393 |
+
remaining = &remaining[index + placeholder.len()..];
|
| 394 |
+
}
|
| 395 |
+
output
|
| 396 |
+
}
|
| 397 |
+
|
| 398 |
+
fn absolute_plugin_path(path: &Path) -> Result<PathBuf, String> {
|
| 399 |
+
let absolute = if path.is_absolute() {
|
| 400 |
+
Ok(path.to_path_buf())
|
| 401 |
+
} else {
|
| 402 |
+
std::env::current_dir()
|
| 403 |
+
.map(|cwd| cwd.join(path))
|
| 404 |
+
.map_err(|err| format!("failed to resolve plugin path: {err}"))
|
| 405 |
+
}?;
|
| 406 |
+
resolve_existing_path_prefix(&absolute)
|
| 407 |
+
}
|
| 408 |
+
|
| 409 |
+
fn resolve_contained_host_path(
|
| 410 |
+
value: &str,
|
| 411 |
+
root: &Path,
|
| 412 |
+
allowed_root: &Path,
|
| 413 |
+
) -> Result<PathBuf, String> {
|
| 414 |
+
let value = Path::new(value);
|
| 415 |
+
let path = if value.is_absolute() {
|
| 416 |
+
value.to_path_buf()
|
| 417 |
+
} else {
|
| 418 |
+
root.join(value)
|
| 419 |
+
};
|
| 420 |
+
let path = resolve_existing_path_prefix(&path)?;
|
| 421 |
+
if !path.starts_with(allowed_root) {
|
| 422 |
+
return Err(format!(
|
| 423 |
+
"expanded path `{}` must remain within `{}`",
|
| 424 |
+
value.display(),
|
| 425 |
+
allowed_root.display()
|
| 426 |
+
));
|
| 427 |
+
}
|
| 428 |
+
Ok(path)
|
| 429 |
+
}
|
| 430 |
+
|
| 431 |
+
fn resolve_existing_path_prefix(path: &Path) -> Result<PathBuf, String> {
|
| 432 |
+
let mut existing = path.to_path_buf();
|
| 433 |
+
let mut missing_components = Vec::<OsString>::new();
|
| 434 |
+
loop {
|
| 435 |
+
match std::fs::canonicalize(&existing) {
|
| 436 |
+
Ok(mut resolved) => {
|
| 437 |
+
for component in missing_components.iter().rev() {
|
| 438 |
+
resolved.push(component);
|
| 439 |
+
}
|
| 440 |
+
return Ok(lexical_normalize(&resolved));
|
| 441 |
+
}
|
| 442 |
+
Err(err) if err.kind() == std::io::ErrorKind::NotFound => {
|
| 443 |
+
if std::fs::symlink_metadata(&existing)
|
| 444 |
+
.is_ok_and(|metadata| metadata.file_type().is_symlink())
|
| 445 |
+
{
|
| 446 |
+
return Err(format!(
|
| 447 |
+
"failed to resolve symlinked path `{}`",
|
| 448 |
+
path.display()
|
| 449 |
+
));
|
| 450 |
+
}
|
| 451 |
+
let Some(component) = existing.components().next_back() else {
|
| 452 |
+
return Err(format!(
|
| 453 |
+
"failed to resolve path `{}`: {err}",
|
| 454 |
+
path.display()
|
| 455 |
+
));
|
| 456 |
+
};
|
| 457 |
+
if matches!(
|
| 458 |
+
component,
|
| 459 |
+
std::path::Component::Prefix(_) | std::path::Component::RootDir
|
| 460 |
+
) {
|
| 461 |
+
return Err(format!(
|
| 462 |
+
"failed to resolve path `{}`: {err}",
|
| 463 |
+
path.display()
|
| 464 |
+
));
|
| 465 |
+
}
|
| 466 |
+
missing_components.push(component.as_os_str().to_os_string());
|
| 467 |
+
if !existing.pop() {
|
| 468 |
+
return Err(format!(
|
| 469 |
+
"failed to resolve path `{}`: {err}",
|
| 470 |
+
path.display()
|
| 471 |
+
));
|
| 472 |
+
}
|
| 473 |
+
}
|
| 474 |
+
Err(err) => {
|
| 475 |
+
return Err(format!(
|
| 476 |
+
"failed to resolve path `{}`: {err}",
|
| 477 |
+
path.display()
|
| 478 |
+
));
|
| 479 |
+
}
|
| 480 |
+
}
|
| 481 |
+
}
|
| 482 |
+
}
|
| 483 |
+
|
| 484 |
+
fn host_path_string(path: &Path) -> String {
|
| 485 |
+
let rendered = path.to_string_lossy();
|
| 486 |
+
#[cfg(windows)]
|
| 487 |
+
if let Some(path) = rendered.strip_prefix(r"\\?\") {
|
| 488 |
+
return path
|
| 489 |
+
.strip_prefix(r"UNC\")
|
| 490 |
+
.map(|path| format!(r"\\{path}"))
|
| 491 |
+
.unwrap_or_else(|| path.to_string());
|
| 492 |
+
}
|
| 493 |
+
rendered.into_owned()
|
| 494 |
+
}
|
| 495 |
+
|
| 496 |
+
fn is_portable_relative_path(value: &str) -> bool {
|
| 497 |
+
value
|
| 498 |
+
.strip_prefix("./")
|
| 499 |
+
.is_some_and(is_portable_path_suffix)
|
| 500 |
+
}
|
| 501 |
+
|
| 502 |
+
fn is_portable_path_suffix(value: &str) -> bool {
|
| 503 |
+
!value.is_empty() && !value.contains('\\')
|
| 504 |
+
}
|
| 505 |
+
|
| 506 |
+
fn is_valid_http_header_name(name: &str) -> bool {
|
| 507 |
+
!name.is_empty()
|
| 508 |
+
&& name
|
| 509 |
+
.bytes()
|
| 510 |
+
.all(|byte| byte.is_ascii_alphanumeric() || b"!#$%&'*+-.^_`|~".contains(&byte))
|
| 511 |
+
}
|
| 512 |
+
|
| 513 |
+
fn lexical_normalize(path: &Path) -> PathBuf {
|
| 514 |
+
let mut normalized = PathBuf::new();
|
| 515 |
+
for component in path.components() {
|
| 516 |
+
match component {
|
| 517 |
+
std::path::Component::CurDir => {}
|
| 518 |
+
std::path::Component::ParentDir => {
|
| 519 |
+
normalized.pop();
|
| 520 |
+
}
|
| 521 |
+
component => normalized.push(component.as_os_str()),
|
| 522 |
+
}
|
| 523 |
+
}
|
| 524 |
+
normalized
|
| 525 |
+
}
|
| 526 |
+
|
| 527 |
+
fn plugin_mcp_json_error(message: impl Into<String>) -> serde_json::Error {
|
| 528 |
+
serde_json::Error::io(std::io::Error::new(
|
| 529 |
+
std::io::ErrorKind::InvalidData,
|
| 530 |
+
message.into(),
|
| 531 |
+
))
|
| 532 |
+
}
|
codex-rs/codex-mcp/src/auth_changes.rs
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Forwards auth invalidations without credentials. The managed client owns the watcher.
|
| 2 |
+
|
| 3 |
+
use std::sync::Arc;
|
| 4 |
+
use std::time::Duration;
|
| 5 |
+
|
| 6 |
+
use anyhow::Result;
|
| 7 |
+
use codex_login::AuthChangeState;
|
| 8 |
+
use codex_rmcp_client::RmcpClient;
|
| 9 |
+
use rmcp::model::ServerCapabilities;
|
| 10 |
+
use serde_json::json;
|
| 11 |
+
use tokio::sync::watch;
|
| 12 |
+
use tokio_util::task::AbortOnDropHandle;
|
| 13 |
+
|
| 14 |
+
pub(crate) const CAPABILITY: &str = "codex/auth-change";
|
| 15 |
+
const NOTIFICATION: &str = "notifications/codex/authChanged";
|
| 16 |
+
const SEND_TIMEOUT: Duration = Duration::from_secs(5);
|
| 17 |
+
|
| 18 |
+
pub(crate) async fn start(
|
| 19 |
+
client: Arc<RmcpClient>,
|
| 20 |
+
capabilities: &ServerCapabilities,
|
| 21 |
+
changes: Option<watch::Receiver<AuthChangeState>>,
|
| 22 |
+
) -> Result<Option<Arc<AbortOnDropHandle<()>>>> {
|
| 23 |
+
let Some(mut changes) = changes.filter(|_| {
|
| 24 |
+
capabilities
|
| 25 |
+
.experimental
|
| 26 |
+
.as_ref()
|
| 27 |
+
.is_some_and(|capabilities| capabilities.contains_key(CAPABILITY))
|
| 28 |
+
}) else {
|
| 29 |
+
return Ok(None);
|
| 30 |
+
};
|
| 31 |
+
|
| 32 |
+
notify(&client, &mut changes).await?;
|
| 33 |
+
let task = tokio::spawn(async move {
|
| 34 |
+
while changes.changed().await.is_ok() {
|
| 35 |
+
if notify(&client, &mut changes).await.is_err() {
|
| 36 |
+
tracing::warn!("MCP auth invalidation delivery failed; closing connection");
|
| 37 |
+
client.shutdown().await;
|
| 38 |
+
break;
|
| 39 |
+
}
|
| 40 |
+
}
|
| 41 |
+
});
|
| 42 |
+
Ok(Some(Arc::new(AbortOnDropHandle::new(task))))
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
async fn notify(client: &RmcpClient, changes: &mut watch::Receiver<AuthChangeState>) -> Result<()> {
|
| 46 |
+
let state = *changes.borrow_and_update();
|
| 47 |
+
tokio::time::timeout(
|
| 48 |
+
SEND_TIMEOUT,
|
| 49 |
+
client.send_custom_notification(
|
| 50 |
+
NOTIFICATION,
|
| 51 |
+
Some(json!({
|
| 52 |
+
"generation": state.generation,
|
| 53 |
+
"ownerGeneration": state.owner_generation,
|
| 54 |
+
})),
|
| 55 |
+
),
|
| 56 |
+
)
|
| 57 |
+
.await?
|
| 58 |
+
}
|
| 59 |
+
|
| 60 |
+
#[cfg(test)]
|
| 61 |
+
#[path = "auth_changes_tests.rs"]
|
| 62 |
+
mod tests;
|
codex-rs/codex-mcp/src/auth_changes_tests.rs
ADDED
|
@@ -0,0 +1,104 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
use super::*;
|
| 2 |
+
use codex_rmcp_client::InProcessTransportFactory;
|
| 3 |
+
use futures::FutureExt;
|
| 4 |
+
use futures::future::BoxFuture;
|
| 5 |
+
use pretty_assertions::assert_eq;
|
| 6 |
+
use rmcp::ServiceExt;
|
| 7 |
+
use rmcp::model::ClientCapabilities;
|
| 8 |
+
use rmcp::model::CustomNotification;
|
| 9 |
+
use rmcp::model::Implementation;
|
| 10 |
+
use rmcp::model::InitializeRequestParams;
|
| 11 |
+
use tokio::sync::mpsc;
|
| 12 |
+
use tokio::time::timeout;
|
| 13 |
+
|
| 14 |
+
#[derive(Clone)]
|
| 15 |
+
struct NotificationServer(mpsc::Sender<serde_json::Value>);
|
| 16 |
+
|
| 17 |
+
impl rmcp::ServerHandler for NotificationServer {
|
| 18 |
+
async fn on_custom_notification(
|
| 19 |
+
&self,
|
| 20 |
+
notification: CustomNotification,
|
| 21 |
+
_context: rmcp::service::NotificationContext<rmcp::RoleServer>,
|
| 22 |
+
) {
|
| 23 |
+
self.0.send(json!(notification)).await.unwrap();
|
| 24 |
+
}
|
| 25 |
+
}
|
| 26 |
+
|
| 27 |
+
impl InProcessTransportFactory for NotificationServer {
|
| 28 |
+
fn open(&self) -> BoxFuture<'static, std::io::Result<tokio::io::DuplexStream>> {
|
| 29 |
+
let server = self.clone();
|
| 30 |
+
async move {
|
| 31 |
+
let (client, transport) = tokio::io::duplex(/*max_buf_size*/ 4096);
|
| 32 |
+
tokio::spawn(async move {
|
| 33 |
+
let service = server.serve(transport).await.unwrap();
|
| 34 |
+
service.waiting().await.unwrap();
|
| 35 |
+
});
|
| 36 |
+
Ok(client)
|
| 37 |
+
}
|
| 38 |
+
.boxed()
|
| 39 |
+
}
|
| 40 |
+
}
|
| 41 |
+
|
| 42 |
+
#[tokio::test]
|
| 43 |
+
async fn auth_notifications_require_opt_in_and_follow_client_lifetime() -> Result<()> {
|
| 44 |
+
let (notifications, mut received) = mpsc::channel(/*buffer*/ 8);
|
| 45 |
+
let client = Arc::new(
|
| 46 |
+
RmcpClient::new_in_process_client(Arc::new(NotificationServer(notifications))).await?,
|
| 47 |
+
);
|
| 48 |
+
client
|
| 49 |
+
.initialize(
|
| 50 |
+
InitializeRequestParams::new(
|
| 51 |
+
ClientCapabilities::default(),
|
| 52 |
+
Implementation::new("test", "1"),
|
| 53 |
+
),
|
| 54 |
+
Some(SEND_TIMEOUT),
|
| 55 |
+
Box::new(|_, _| async { anyhow::bail!("unexpected elicitation") }.boxed()),
|
| 56 |
+
)
|
| 57 |
+
.await?;
|
| 58 |
+
let (changes, receiver) = watch::channel(AuthChangeState::default());
|
| 59 |
+
let mut capabilities = ServerCapabilities::default();
|
| 60 |
+
assert!(
|
| 61 |
+
start(Arc::clone(&client), &capabilities, Some(receiver.clone()))
|
| 62 |
+
.await?
|
| 63 |
+
.is_none()
|
| 64 |
+
);
|
| 65 |
+
capabilities.experimental = Some([(CAPABILITY.to_string(), Default::default())].into());
|
| 66 |
+
assert!(
|
| 67 |
+
start(Arc::clone(&client), &capabilities, /*changes*/ None)
|
| 68 |
+
.await?
|
| 69 |
+
.is_none()
|
| 70 |
+
);
|
| 71 |
+
assert_eq!(received.try_recv(), Err(mpsc::error::TryRecvError::Empty));
|
| 72 |
+
let watcher = start(Arc::clone(&client), &capabilities, Some(receiver))
|
| 73 |
+
.await?
|
| 74 |
+
.unwrap();
|
| 75 |
+
assert_eq!(
|
| 76 |
+
timeout(SEND_TIMEOUT, received.recv()).await?,
|
| 77 |
+
Some(
|
| 78 |
+
json!({"method": NOTIFICATION, "params": {"_meta": {}, "generation": 0, "ownerGeneration": 0}})
|
| 79 |
+
),
|
| 80 |
+
);
|
| 81 |
+
changes.send_modify(|state| state.generation += 1);
|
| 82 |
+
assert_eq!(
|
| 83 |
+
timeout(SEND_TIMEOUT, received.recv()).await?,
|
| 84 |
+
Some(
|
| 85 |
+
json!({"method": NOTIFICATION, "params": {"_meta": {}, "generation": 1, "ownerGeneration": 0}})
|
| 86 |
+
),
|
| 87 |
+
);
|
| 88 |
+
for _ in 0..2 {
|
| 89 |
+
changes.send_modify(|state| {
|
| 90 |
+
state.generation += 1;
|
| 91 |
+
state.owner_generation += 1;
|
| 92 |
+
});
|
| 93 |
+
}
|
| 94 |
+
assert_eq!(
|
| 95 |
+
timeout(SEND_TIMEOUT, received.recv()).await?,
|
| 96 |
+
Some(
|
| 97 |
+
json!({"method": NOTIFICATION, "params": {"_meta": {}, "generation": 3, "ownerGeneration": 2}})
|
| 98 |
+
),
|
| 99 |
+
);
|
| 100 |
+
drop(watcher);
|
| 101 |
+
timeout(SEND_TIMEOUT, changes.closed()).await?;
|
| 102 |
+
client.shutdown().await;
|
| 103 |
+
Ok(())
|
| 104 |
+
}
|
codex-rs/codex-mcp/src/auth_elicitation.rs
ADDED
|
@@ -0,0 +1,435 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
//! Auth elicitation helpers.
|
| 2 |
+
//!
|
| 3 |
+
//! This module owns protocol-neutral auth elicitation parsing and payload shaping.
|
| 4 |
+
//! Session orchestration stays in `codex-core`.
|
| 5 |
+
|
| 6 |
+
use codex_protocol::mcp::CallToolResult;
|
| 7 |
+
use serde::Serialize;
|
| 8 |
+
|
| 9 |
+
pub const MCP_TOOL_CODEX_APPS_META_KEY: &str = "_codex_apps";
|
| 10 |
+
pub const CONNECTOR_AUTH_FAILURE_META_KEY: &str = "connector_auth_failure";
|
| 11 |
+
pub const CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY: &str = "is_auth_failure";
|
| 12 |
+
pub const CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY: &str = "auth_reason";
|
| 13 |
+
pub const CONNECTOR_AUTH_FAILURE_CONNECTOR_ID_KEY: &str = "connector_id";
|
| 14 |
+
pub const CONNECTOR_AUTH_FAILURE_LINK_ID_KEY: &str = "link_id";
|
| 15 |
+
pub const CONNECTOR_AUTH_FAILURE_ERROR_CODE_KEY: &str = "error_code";
|
| 16 |
+
pub const CONNECTOR_AUTH_FAILURE_ERROR_HTTP_STATUS_CODE_KEY: &str = "error_http_status_code";
|
| 17 |
+
pub const CONNECTOR_AUTH_FAILURE_ERROR_ACTION_KEY: &str = "error_action";
|
| 18 |
+
|
| 19 |
+
#[derive(Debug, Clone, PartialEq, Eq)]
|
| 20 |
+
pub struct CodexAppsConnectorAuthFailure {
|
| 21 |
+
pub connector_id: String,
|
| 22 |
+
pub connector_name: String,
|
| 23 |
+
pub install_url: String,
|
| 24 |
+
pub auth_reason: Option<String>,
|
| 25 |
+
pub link_id: Option<String>,
|
| 26 |
+
pub error_code: Option<String>,
|
| 27 |
+
pub error_http_status_code: Option<i64>,
|
| 28 |
+
pub error_action: Option<String>,
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
#[derive(Debug, Clone, PartialEq)]
|
| 32 |
+
pub struct CodexAppsAuthElicitation {
|
| 33 |
+
pub meta: serde_json::Value,
|
| 34 |
+
pub message: String,
|
| 35 |
+
pub url: String,
|
| 36 |
+
pub elicitation_id: String,
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
#[derive(Debug, Clone, PartialEq)]
|
| 40 |
+
pub struct CodexAppsAuthElicitationPlan {
|
| 41 |
+
pub auth_failure: CodexAppsConnectorAuthFailure,
|
| 42 |
+
pub elicitation: CodexAppsAuthElicitation,
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
#[derive(Serialize)]
|
| 46 |
+
struct CodexAppsConnectorAuthFailureMeta<'a> {
|
| 47 |
+
is_auth_failure: bool,
|
| 48 |
+
connector_id: &'a str,
|
| 49 |
+
connector_name: &'a str,
|
| 50 |
+
install_url: &'a str,
|
| 51 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 52 |
+
auth_reason: Option<&'a str>,
|
| 53 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 54 |
+
link_id: Option<&'a str>,
|
| 55 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 56 |
+
error_code: Option<&'a str>,
|
| 57 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 58 |
+
error_http_status_code: Option<i64>,
|
| 59 |
+
#[serde(skip_serializing_if = "Option::is_none")]
|
| 60 |
+
error_action: Option<&'a str>,
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
pub fn connector_auth_failure_from_tool_result(
|
| 64 |
+
result: &CallToolResult,
|
| 65 |
+
connector_id: Option<&str>,
|
| 66 |
+
connector_name: Option<&str>,
|
| 67 |
+
install_url: Option<String>,
|
| 68 |
+
) -> Option<CodexAppsConnectorAuthFailure> {
|
| 69 |
+
let connector_id = connector_id
|
| 70 |
+
.map(str::trim)
|
| 71 |
+
.filter(|connector_id| !connector_id.is_empty())?;
|
| 72 |
+
let auth_failure = connector_auth_failure_metadata(result, connector_id)?;
|
| 73 |
+
let connector_name = connector_name
|
| 74 |
+
.map(str::trim)
|
| 75 |
+
.filter(|name| !name.is_empty())
|
| 76 |
+
.unwrap_or(connector_id)
|
| 77 |
+
.to_string();
|
| 78 |
+
|
| 79 |
+
Some(CodexAppsConnectorAuthFailure {
|
| 80 |
+
connector_id: connector_id.to_string(),
|
| 81 |
+
connector_name,
|
| 82 |
+
install_url: install_url?,
|
| 83 |
+
auth_reason: string_auth_failure_field(
|
| 84 |
+
auth_failure,
|
| 85 |
+
CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY,
|
| 86 |
+
),
|
| 87 |
+
link_id: string_auth_failure_field(auth_failure, CONNECTOR_AUTH_FAILURE_LINK_ID_KEY),
|
| 88 |
+
error_code: string_auth_failure_field(auth_failure, CONNECTOR_AUTH_FAILURE_ERROR_CODE_KEY),
|
| 89 |
+
error_http_status_code: auth_failure
|
| 90 |
+
.get(CONNECTOR_AUTH_FAILURE_ERROR_HTTP_STATUS_CODE_KEY)
|
| 91 |
+
.and_then(serde_json::Value::as_i64),
|
| 92 |
+
error_action: string_auth_failure_field(
|
| 93 |
+
auth_failure,
|
| 94 |
+
CONNECTOR_AUTH_FAILURE_ERROR_ACTION_KEY,
|
| 95 |
+
),
|
| 96 |
+
})
|
| 97 |
+
}
|
| 98 |
+
|
| 99 |
+
pub fn is_connector_auth_failure_from_tool_result(
|
| 100 |
+
result: &CallToolResult,
|
| 101 |
+
connector_id: Option<&str>,
|
| 102 |
+
) -> bool {
|
| 103 |
+
connector_id
|
| 104 |
+
.map(str::trim)
|
| 105 |
+
.filter(|connector_id| !connector_id.is_empty())
|
| 106 |
+
.is_some_and(|connector_id| connector_auth_failure_metadata(result, connector_id).is_some())
|
| 107 |
+
}
|
| 108 |
+
|
| 109 |
+
fn connector_auth_failure_metadata<'a>(
|
| 110 |
+
result: &'a CallToolResult,
|
| 111 |
+
connector_id: &str,
|
| 112 |
+
) -> Option<&'a serde_json::Map<String, serde_json::Value>> {
|
| 113 |
+
if result.is_error != Some(true) {
|
| 114 |
+
return None;
|
| 115 |
+
}
|
| 116 |
+
|
| 117 |
+
let auth_failure = result
|
| 118 |
+
.meta
|
| 119 |
+
.as_ref()?
|
| 120 |
+
.as_object()?
|
| 121 |
+
.get(MCP_TOOL_CODEX_APPS_META_KEY)?
|
| 122 |
+
.as_object()?
|
| 123 |
+
.get(CONNECTOR_AUTH_FAILURE_META_KEY)?
|
| 124 |
+
.as_object()?;
|
| 125 |
+
if auth_failure
|
| 126 |
+
.get(CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY)
|
| 127 |
+
.and_then(serde_json::Value::as_bool)
|
| 128 |
+
!= Some(true)
|
| 129 |
+
{
|
| 130 |
+
return None;
|
| 131 |
+
}
|
| 132 |
+
if let Some(auth_failure_connector_id) =
|
| 133 |
+
string_auth_failure_field(auth_failure, CONNECTOR_AUTH_FAILURE_CONNECTOR_ID_KEY)
|
| 134 |
+
&& auth_failure_connector_id != connector_id
|
| 135 |
+
{
|
| 136 |
+
return None;
|
| 137 |
+
}
|
| 138 |
+
|
| 139 |
+
Some(auth_failure)
|
| 140 |
+
}
|
| 141 |
+
|
| 142 |
+
pub fn build_auth_elicitation_plan(
|
| 143 |
+
call_id: &str,
|
| 144 |
+
result: &CallToolResult,
|
| 145 |
+
connector_id: Option<&str>,
|
| 146 |
+
connector_name: Option<&str>,
|
| 147 |
+
install_url: Option<String>,
|
| 148 |
+
) -> Option<CodexAppsAuthElicitationPlan> {
|
| 149 |
+
let auth_failure =
|
| 150 |
+
connector_auth_failure_from_tool_result(result, connector_id, connector_name, install_url)?;
|
| 151 |
+
let elicitation = build_auth_elicitation(call_id, &auth_failure);
|
| 152 |
+
Some(CodexAppsAuthElicitationPlan {
|
| 153 |
+
auth_failure,
|
| 154 |
+
elicitation,
|
| 155 |
+
})
|
| 156 |
+
}
|
| 157 |
+
|
| 158 |
+
pub fn build_auth_elicitation(
|
| 159 |
+
call_id: &str,
|
| 160 |
+
auth_failure: &CodexAppsConnectorAuthFailure,
|
| 161 |
+
) -> CodexAppsAuthElicitation {
|
| 162 |
+
CodexAppsAuthElicitation {
|
| 163 |
+
meta: serde_json::json!({
|
| 164 |
+
MCP_TOOL_CODEX_APPS_META_KEY: {
|
| 165 |
+
CONNECTOR_AUTH_FAILURE_META_KEY: CodexAppsConnectorAuthFailureMeta {
|
| 166 |
+
is_auth_failure: true,
|
| 167 |
+
connector_id: &auth_failure.connector_id,
|
| 168 |
+
connector_name: &auth_failure.connector_name,
|
| 169 |
+
install_url: &auth_failure.install_url,
|
| 170 |
+
auth_reason: auth_failure.auth_reason.as_deref(),
|
| 171 |
+
link_id: auth_failure.link_id.as_deref(),
|
| 172 |
+
error_code: auth_failure.error_code.as_deref(),
|
| 173 |
+
error_http_status_code: auth_failure.error_http_status_code,
|
| 174 |
+
error_action: auth_failure.error_action.as_deref(),
|
| 175 |
+
},
|
| 176 |
+
},
|
| 177 |
+
}),
|
| 178 |
+
message: auth_elicitation_message(auth_failure),
|
| 179 |
+
url: auth_failure.install_url.clone(),
|
| 180 |
+
elicitation_id: auth_elicitation_id(call_id),
|
| 181 |
+
}
|
| 182 |
+
}
|
| 183 |
+
|
| 184 |
+
pub fn auth_elicitation_completed_result(
|
| 185 |
+
auth_failure: &CodexAppsConnectorAuthFailure,
|
| 186 |
+
meta: Option<serde_json::Value>,
|
| 187 |
+
) -> CallToolResult {
|
| 188 |
+
CallToolResult {
|
| 189 |
+
content: vec![serde_json::json!({
|
| 190 |
+
"type": "text",
|
| 191 |
+
"text": format!(
|
| 192 |
+
"Authentication for {} was requested and accepted. Retry this tool call now.",
|
| 193 |
+
auth_failure.connector_name
|
| 194 |
+
),
|
| 195 |
+
})],
|
| 196 |
+
structured_content: None,
|
| 197 |
+
is_error: Some(true),
|
| 198 |
+
meta,
|
| 199 |
+
}
|
| 200 |
+
}
|
| 201 |
+
|
| 202 |
+
pub fn auth_elicitation_id(call_id: &str) -> String {
|
| 203 |
+
format!("codex_apps_auth_{call_id}")
|
| 204 |
+
}
|
| 205 |
+
|
| 206 |
+
fn string_auth_failure_field(
|
| 207 |
+
auth_failure: &serde_json::Map<String, serde_json::Value>,
|
| 208 |
+
key: &str,
|
| 209 |
+
) -> Option<String> {
|
| 210 |
+
auth_failure
|
| 211 |
+
.get(key)
|
| 212 |
+
.and_then(serde_json::Value::as_str)
|
| 213 |
+
.map(str::trim)
|
| 214 |
+
.filter(|value| !value.is_empty())
|
| 215 |
+
.map(ToString::to_string)
|
| 216 |
+
}
|
| 217 |
+
|
| 218 |
+
fn auth_elicitation_message(auth_failure: &CodexAppsConnectorAuthFailure) -> String {
|
| 219 |
+
match auth_failure.auth_reason.as_deref() {
|
| 220 |
+
Some("oauth_upgrade_required") => format!(
|
| 221 |
+
"Reconnect {} on ChatGPT to grant the permissions needed for this request.",
|
| 222 |
+
auth_failure.connector_name
|
| 223 |
+
),
|
| 224 |
+
Some("reauthentication_required") => format!(
|
| 225 |
+
"Reconnect {} on ChatGPT to restore access for this request.",
|
| 226 |
+
auth_failure.connector_name
|
| 227 |
+
),
|
| 228 |
+
Some("missing_link") => format!(
|
| 229 |
+
"Sign in to {} on ChatGPT to use it in Codex.",
|
| 230 |
+
auth_failure.connector_name
|
| 231 |
+
),
|
| 232 |
+
_ => format!(
|
| 233 |
+
"Sign in to {} on ChatGPT to continue.",
|
| 234 |
+
auth_failure.connector_name
|
| 235 |
+
),
|
| 236 |
+
}
|
| 237 |
+
}
|
| 238 |
+
|
| 239 |
+
#[cfg(test)]
|
| 240 |
+
mod tests {
|
| 241 |
+
use super::*;
|
| 242 |
+
use pretty_assertions::assert_eq;
|
| 243 |
+
|
| 244 |
+
fn auth_failure_result() -> CallToolResult {
|
| 245 |
+
CallToolResult {
|
| 246 |
+
content: vec![serde_json::json!({
|
| 247 |
+
"type": "text",
|
| 248 |
+
"text": "Connector reauthentication required",
|
| 249 |
+
})],
|
| 250 |
+
structured_content: None,
|
| 251 |
+
is_error: Some(true),
|
| 252 |
+
meta: Some(serde_json::json!({
|
| 253 |
+
MCP_TOOL_CODEX_APPS_META_KEY: {
|
| 254 |
+
CONNECTOR_AUTH_FAILURE_META_KEY: {
|
| 255 |
+
CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY: true,
|
| 256 |
+
CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY: "reauthentication_required",
|
| 257 |
+
CONNECTOR_AUTH_FAILURE_CONNECTOR_ID_KEY: "connector_calendar",
|
| 258 |
+
"connector_name": "Untrusted Calendar",
|
| 259 |
+
CONNECTOR_AUTH_FAILURE_LINK_ID_KEY: "link_123",
|
| 260 |
+
CONNECTOR_AUTH_FAILURE_ERROR_CODE_KEY: "UNAUTHORIZED",
|
| 261 |
+
CONNECTOR_AUTH_FAILURE_ERROR_HTTP_STATUS_CODE_KEY: 401,
|
| 262 |
+
CONNECTOR_AUTH_FAILURE_ERROR_ACTION_KEY: "TRIGGER_REAUTHENTICATION",
|
| 263 |
+
},
|
| 264 |
+
},
|
| 265 |
+
})),
|
| 266 |
+
}
|
| 267 |
+
}
|
| 268 |
+
|
| 269 |
+
#[test]
|
| 270 |
+
fn parses_auth_failure_from_trusted_connector_metadata() {
|
| 271 |
+
assert_eq!(
|
| 272 |
+
connector_auth_failure_from_tool_result(
|
| 273 |
+
&auth_failure_result(),
|
| 274 |
+
Some("connector_calendar"),
|
| 275 |
+
Some("Google Calendar"),
|
| 276 |
+
Some("https://chatgpt.com/apps/google-calendar/connector_calendar".to_string()),
|
| 277 |
+
),
|
| 278 |
+
Some(CodexAppsConnectorAuthFailure {
|
| 279 |
+
connector_id: "connector_calendar".to_string(),
|
| 280 |
+
connector_name: "Google Calendar".to_string(),
|
| 281 |
+
install_url: "https://chatgpt.com/apps/google-calendar/connector_calendar"
|
| 282 |
+
.to_string(),
|
| 283 |
+
auth_reason: Some("reauthentication_required".to_string()),
|
| 284 |
+
link_id: Some("link_123".to_string()),
|
| 285 |
+
error_code: Some("UNAUTHORIZED".to_string()),
|
| 286 |
+
error_http_status_code: Some(401),
|
| 287 |
+
error_action: Some("TRIGGER_REAUTHENTICATION".to_string()),
|
| 288 |
+
})
|
| 289 |
+
);
|
| 290 |
+
}
|
| 291 |
+
|
| 292 |
+
#[test]
|
| 293 |
+
fn rejects_missing_or_mismatched_connector_ids() {
|
| 294 |
+
assert_eq!(
|
| 295 |
+
connector_auth_failure_from_tool_result(
|
| 296 |
+
&auth_failure_result(),
|
| 297 |
+
/*connector_id*/ None,
|
| 298 |
+
Some("Google Calendar"),
|
| 299 |
+
Some("https://chatgpt.com/apps/google-calendar/connector_calendar".to_string()),
|
| 300 |
+
),
|
| 301 |
+
None
|
| 302 |
+
);
|
| 303 |
+
assert_eq!(
|
| 304 |
+
connector_auth_failure_from_tool_result(
|
| 305 |
+
&auth_failure_result(),
|
| 306 |
+
Some("connector_drive"),
|
| 307 |
+
Some("Google Drive"),
|
| 308 |
+
Some("https://chatgpt.com/apps/google-drive/connector_drive".to_string()),
|
| 309 |
+
),
|
| 310 |
+
None
|
| 311 |
+
);
|
| 312 |
+
}
|
| 313 |
+
|
| 314 |
+
#[test]
|
| 315 |
+
fn detects_auth_failure_without_an_install_url() {
|
| 316 |
+
let result = auth_failure_result();
|
| 317 |
+
|
| 318 |
+
assert_eq!(
|
| 319 |
+
is_connector_auth_failure_from_tool_result(&result, Some("connector_calendar")),
|
| 320 |
+
true
|
| 321 |
+
);
|
| 322 |
+
assert_eq!(
|
| 323 |
+
connector_auth_failure_from_tool_result(
|
| 324 |
+
&result,
|
| 325 |
+
Some("connector_calendar"),
|
| 326 |
+
Some("Google Calendar"),
|
| 327 |
+
/*install_url*/ None,
|
| 328 |
+
),
|
| 329 |
+
None
|
| 330 |
+
);
|
| 331 |
+
}
|
| 332 |
+
|
| 333 |
+
#[test]
|
| 334 |
+
fn auth_failure_detection_requires_trusted_connector_identity_and_auth_flag() {
|
| 335 |
+
let result = auth_failure_result();
|
| 336 |
+
assert_eq!(
|
| 337 |
+
is_connector_auth_failure_from_tool_result(&result, /*connector_id*/ None),
|
| 338 |
+
false
|
| 339 |
+
);
|
| 340 |
+
assert_eq!(
|
| 341 |
+
is_connector_auth_failure_from_tool_result(&result, Some("connector_drive")),
|
| 342 |
+
false
|
| 343 |
+
);
|
| 344 |
+
|
| 345 |
+
let mut ordinary_error = result.clone();
|
| 346 |
+
ordinary_error.meta.as_mut().expect("auth metadata")[MCP_TOOL_CODEX_APPS_META_KEY]
|
| 347 |
+
[CONNECTOR_AUTH_FAILURE_META_KEY][CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY] =
|
| 348 |
+
serde_json::Value::Bool(false);
|
| 349 |
+
assert_eq!(
|
| 350 |
+
is_connector_auth_failure_from_tool_result(&ordinary_error, Some("connector_calendar"),),
|
| 351 |
+
false
|
| 352 |
+
);
|
| 353 |
+
|
| 354 |
+
let mut successful_result = result;
|
| 355 |
+
successful_result.is_error = Some(false);
|
| 356 |
+
assert_eq!(
|
| 357 |
+
is_connector_auth_failure_from_tool_result(
|
| 358 |
+
&successful_result,
|
| 359 |
+
Some("connector_calendar"),
|
| 360 |
+
),
|
| 361 |
+
false
|
| 362 |
+
);
|
| 363 |
+
}
|
| 364 |
+
|
| 365 |
+
#[test]
|
| 366 |
+
fn detects_each_supported_connector_auth_reason() {
|
| 367 |
+
for auth_reason in [
|
| 368 |
+
"missing_link",
|
| 369 |
+
"oauth_upgrade_required",
|
| 370 |
+
"reauthentication_required",
|
| 371 |
+
] {
|
| 372 |
+
let mut result = auth_failure_result();
|
| 373 |
+
result.meta.as_mut().expect("auth metadata")[MCP_TOOL_CODEX_APPS_META_KEY]
|
| 374 |
+
[CONNECTOR_AUTH_FAILURE_META_KEY][CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY] =
|
| 375 |
+
serde_json::Value::String(auth_reason.to_string());
|
| 376 |
+
|
| 377 |
+
assert_eq!(
|
| 378 |
+
is_connector_auth_failure_from_tool_result(&result, Some("connector_calendar")),
|
| 379 |
+
true
|
| 380 |
+
);
|
| 381 |
+
}
|
| 382 |
+
}
|
| 383 |
+
|
| 384 |
+
#[test]
|
| 385 |
+
fn builds_url_elicitation_payload() {
|
| 386 |
+
let auth_failure = connector_auth_failure_from_tool_result(
|
| 387 |
+
&auth_failure_result(),
|
| 388 |
+
Some("connector_calendar"),
|
| 389 |
+
Some("Google Calendar"),
|
| 390 |
+
Some("https://chatgpt.com/apps/google-calendar/connector_calendar".to_string()),
|
| 391 |
+
)
|
| 392 |
+
.expect("auth failure");
|
| 393 |
+
|
| 394 |
+
assert_eq!(
|
| 395 |
+
build_auth_elicitation("call_123", &auth_failure),
|
| 396 |
+
CodexAppsAuthElicitation {
|
| 397 |
+
meta: serde_json::json!({
|
| 398 |
+
MCP_TOOL_CODEX_APPS_META_KEY: {
|
| 399 |
+
CONNECTOR_AUTH_FAILURE_META_KEY: {
|
| 400 |
+
CONNECTOR_AUTH_FAILURE_IS_AUTH_FAILURE_KEY: true,
|
| 401 |
+
CONNECTOR_AUTH_FAILURE_CONNECTOR_ID_KEY: "connector_calendar",
|
| 402 |
+
"connector_name": "Google Calendar",
|
| 403 |
+
"install_url":
|
| 404 |
+
"https://chatgpt.com/apps/google-calendar/connector_calendar",
|
| 405 |
+
CONNECTOR_AUTH_FAILURE_AUTH_REASON_KEY: "reauthentication_required",
|
| 406 |
+
CONNECTOR_AUTH_FAILURE_LINK_ID_KEY: "link_123",
|
| 407 |
+
CONNECTOR_AUTH_FAILURE_ERROR_CODE_KEY: "UNAUTHORIZED",
|
| 408 |
+
CONNECTOR_AUTH_FAILURE_ERROR_HTTP_STATUS_CODE_KEY: 401,
|
| 409 |
+
CONNECTOR_AUTH_FAILURE_ERROR_ACTION_KEY: "TRIGGER_REAUTHENTICATION",
|
| 410 |
+
},
|
| 411 |
+
},
|
| 412 |
+
}),
|
| 413 |
+
message: "Reconnect Google Calendar on ChatGPT to restore access for this request."
|
| 414 |
+
.to_string(),
|
| 415 |
+
url: "https://chatgpt.com/apps/google-calendar/connector_calendar".to_string(),
|
| 416 |
+
elicitation_id: "codex_apps_auth_call_123".to_string(),
|
| 417 |
+
}
|
| 418 |
+
);
|
| 419 |
+
}
|
| 420 |
+
|
| 421 |
+
#[test]
|
| 422 |
+
fn builds_auth_elicitation_plan() {
|
| 423 |
+
let plan = build_auth_elicitation_plan(
|
| 424 |
+
"call_123",
|
| 425 |
+
&auth_failure_result(),
|
| 426 |
+
Some("connector_calendar"),
|
| 427 |
+
Some("Google Calendar"),
|
| 428 |
+
Some("https://chatgpt.com/apps/google-calendar/connector_calendar".to_string()),
|
| 429 |
+
)
|
| 430 |
+
.expect("auth elicitation plan");
|
| 431 |
+
|
| 432 |
+
assert_eq!(plan.auth_failure.connector_name, "Google Calendar");
|
| 433 |
+
assert_eq!(plan.elicitation.elicitation_id, "codex_apps_auth_call_123");
|
| 434 |
+
}
|
| 435 |
+
}
|