SaylorTwift HF Staff commited on
Commit
17f328f
·
verified ·
1 Parent(s): 52a9af3

Add files using upload-large-folder tool

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +3 -0
  2. .github/codex-cli-splash.png +3 -0
  3. codex-cli/bin/codex.js +295 -0
  4. codex-cli/scripts/README.md +23 -0
  5. codex-cli/scripts/build_npm_package.py +461 -0
  6. codex-cli/scripts/init_firewall.sh +115 -0
  7. codex-cli/scripts/run_in_container.sh +95 -0
  8. codex-rs/agent-identity/src/lib.rs +1000 -0
  9. codex-rs/ansi-escape/src/lib.rs +58 -0
  10. codex-rs/app-server-protocol/schema/precomputed/app-server-exports-experimental.json.zst +3 -0
  11. codex-rs/app-server-protocol/schema/precomputed/app-server-exports-stable.json.zst +3 -0
  12. codex-rs/app-server-transport/src/connection_auth.rs +51 -0
  13. codex-rs/app-server-transport/src/daemon_recovery.rs +82 -0
  14. codex-rs/app-server-transport/src/daemon_shutdown.rs +53 -0
  15. codex-rs/app-server-transport/src/daemon_shutdown_tests.rs +22 -0
  16. codex-rs/app-server-transport/src/lib.rs +45 -0
  17. codex-rs/app-server-transport/src/outgoing_message.rs +59 -0
  18. codex-rs/app-server-transport/src/transport/auth.rs +751 -0
  19. codex-rs/app-server-transport/src/transport/mod.rs +602 -0
  20. codex-rs/app-server-transport/src/transport/remote_control/auth.rs +282 -0
  21. codex-rs/app-server-transport/src/transport/remote_control/client_tracker.rs +944 -0
  22. codex-rs/app-server-transport/src/transport/remote_control/clients.rs +303 -0
  23. codex-rs/app-server-transport/src/transport/remote_control/controller.rs +368 -0
  24. codex-rs/app-server-transport/src/transport/remote_control/desired_state.rs +155 -0
  25. codex-rs/app-server-transport/src/transport/remote_control/enroll.rs +770 -0
  26. codex-rs/app-server-transport/src/transport/remote_control/host_device.rs +74 -0
  27. codex-rs/app-server-transport/src/transport/remote_control/host_device_tests.rs +38 -0
  28. codex-rs/app-server-transport/src/transport/remote_control/mod.rs +1002 -0
  29. codex-rs/app-server-transport/src/transport/remote_control/persistence.rs +165 -0
  30. codex-rs/app-server-transport/src/transport/remote_control/persistence_tests.rs +51 -0
  31. codex-rs/app-server-transport/src/transport/remote_control/protocol.rs +401 -0
  32. codex-rs/app-server-transport/src/transport/remote_control/segment.rs +469 -0
  33. codex-rs/app-server-transport/src/transport/remote_control/segment_tests.rs +450 -0
  34. codex-rs/app-server-transport/src/transport/remote_control/server_api.rs +382 -0
  35. codex-rs/app-server-transport/src/transport/remote_control/server_api_tests.rs +321 -0
  36. codex-rs/app-server-transport/src/transport/remote_control/tests.rs +0 -0
  37. codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs +416 -0
  38. codex-rs/app-server-transport/src/transport/remote_control/tests/pairing_tests.rs +1137 -0
  39. codex-rs/app-server-transport/src/transport/remote_control/tests/retry_tests.rs +203 -0
  40. codex-rs/app-server-transport/src/transport/remote_control/websocket.rs +0 -0
  41. codex-rs/app-server-transport/src/transport/remote_control/websocket_refresh_tests.rs +768 -0
  42. codex-rs/app-server-transport/src/transport/stdio.rs +194 -0
  43. codex-rs/app-server-transport/src/transport/unix_socket.rs +265 -0
  44. codex-rs/app-server-transport/src/transport/unix_socket_tests.rs +340 -0
  45. codex-rs/app-server-transport/src/transport/websocket.rs +389 -0
  46. codex-rs/bwrap/src/main.rs +45 -0
  47. codex-rs/codex-mcp/src/agent_plugin_config.rs +532 -0
  48. codex-rs/codex-mcp/src/auth_changes.rs +62 -0
  49. codex-rs/codex-mcp/src/auth_changes_tests.rs +104 -0
  50. 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

  • SHA256: 15b86fa9c0790a779ecdd84f1e7dee029ab79bfb093bd3c4876998696925b013
  • Pointer size: 131 Bytes
  • Size of remote file: 838 kB
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, &timestamp)?,
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, &timestamp)?,
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, &params.environment_id)?;
89
+ let response = send_client_management_request(
90
+ auth_manager,
91
+ ClientManagementRequest::List {
92
+ url: &url,
93
+ params: &params,
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, &params.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(&params.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) = &params.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(&current.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
+ &current_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(&params)?;
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 (&params.pairing_code, &params.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
+ &current_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
+ &current_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
+ &current_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
+ &current_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
+ &current_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
+ &current_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
+ &current_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
+ &current_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
+ &current_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(&current_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(&notification)
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
+ }