From 4aeafa0acd9984a7582edc821e4965b965308071 Mon Sep 17 00:00:00 2001 From: muggle-stack Date: Thu, 24 Sep 2026 23:53:59 +0800 Subject: [PATCH 1/3] fix(runtime): correct Claude stop recovery and Codex shared sockets - Accept protected native Codex socket aliases in Shared launch and App discovery while preserving account and permission checks. - Keep one owner for Claude Stop recovery and retain the original drain deadline when handing off background work. - Distinguish stop confirmation timeouts from confirmed interruptions in live and historical messages. - Add regression coverage for socket aliases, recovery races, and stop timeout presentation. - Validate the unchanged deployed snapshot with the full local gate: 5036 pytest tests passed, 4 skipped, plus Web, lint, and shell checks. --- cc_remote/codex_app_tools.py | 7 +- cc_remote/codex_desktop.py | 5 +- cc_remote/wrapper/machine.py | 20 ++++- docs/codex-desktop-launcher.md | 7 ++ tests/test_codex_app_tools.py | 50 +++++++++++- tests/test_codex_desktop.py | 16 +++- tests/test_wrapper_core_fixes.py | 105 ++++++++++++++++++++++++++ web/src/problem-presentation.ts | 22 +++--- web/tests/notices-rate-limits.test.ts | 37 +++++++++ 9 files changed, 244 insertions(+), 25 deletions(-) diff --git a/cc_remote/codex_app_tools.py b/cc_remote/codex_app_tools.py index 5743ea39..fe216d07 100644 --- a/cc_remote/codex_app_tools.py +++ b/cc_remote/codex_app_tools.py @@ -19,6 +19,7 @@ import sys from urllib.parse import urlsplit +from cc_remote.wrapper.codex_daemon import socket_identity from cc_remote.wrapper.process_scan import ( ProcessIdentity, process_command, @@ -136,9 +137,7 @@ def _shared_bridge(profile: Path, port: int) -> ProcessIdentity | None: # This supervisor forwards only to this profile's canonical private Unix # socket, never to a custom WebSocket backend or a private stdio server. upstream = profile / "app-server-control/app-server-control.sock" - if upstream.resolve(strict=True) != upstream: - return None - _private_socket(upstream) + socket_identity(str(upstream)) return identity if process_identity(identity.pid) == identity else None @@ -249,7 +248,7 @@ def discover(profile: Path, app: Path, expected_manifest: str | None = None) -> return {"state": "unavailable", "reason": "macos_required"} try: profile = profile.resolve(strict=True) - _private_socket(profile / "app-server-control/app-server-control.sock") + socket_identity(str(profile / "app-server-control/app-server-control.sock")) executable, node, plugin, bundle_id = app_paths(app) if expected_manifest is not None: entry = json.loads((plugin / "desktop-mcp.json").read_text())["mcpServers"]["codex_app"] diff --git a/cc_remote/codex_desktop.py b/cc_remote/codex_desktop.py index caae0f53..97771f4e 100644 --- a/cc_remote/codex_desktop.py +++ b/cc_remote/codex_desktop.py @@ -26,6 +26,7 @@ from websockets.asyncio.client import connect from cc_remote import codex_app_tools as desktop +from cc_remote.wrapper.codex_daemon import socket_identity from cc_remote.wrapper.process_scan import ( ProcessIdentity, process_command, @@ -57,9 +58,7 @@ def private_directory(path: Path) -> None: def daemon_socket(profile: Path) -> Path: path = profile / "app-server-control/app-server-control.sock" try: - if path.resolve(strict=True) != path: - raise ValueError("symlink") - desktop._private_socket(path) + socket_identity(str(path)) except (OSError, ValueError): raise LaunchError("共享 daemon 尚未就绪。请先启动该账号的 cc-remote Wrapper,再点共享入口。") from None return path diff --git a/cc_remote/wrapper/machine.py b/cc_remote/wrapper/machine.py index dfbb15e3..27a366b6 100644 --- a/cc_remote/wrapper/machine.py +++ b/cc_remote/wrapper/machine.py @@ -6940,6 +6940,13 @@ async def _on_claude_message_pump_failure( def _schedule_claude_autonomous_interrupt_watchdog( self, ctx: SessionContext, ) -> None: + # The managed consumer already enforces this Stop's drain deadline. + # Anonymous pre-input activity can coexist with that consumer; it must + # not create a second reconnect owner. Its finalizer hands off any + # still-pending autonomous continuation after releasing ownership. + managed = ctx.turn_task + if managed is not None and not managed.done(): + return current = ctx.claude_autonomous_interrupt_task if current is not None and not current.done(): return @@ -6971,10 +6978,12 @@ async def _watch_claude_autonomous_interrupt( return except asyncio.TimeoutError: pass + managed = ctx.turn_task if ( not self._is_resident_context(ctx) or not self._claude_autonomous_followup_pending(ctx) or ctx.state not in {"interrupting", "draining"} + or (managed is not None and not managed.done()) ): return log.error( @@ -18524,9 +18533,8 @@ async def _handle_interrupt(self, cmd) -> None: ) await self._emit(ctx, StateEvent(state="interrupting")) if self._claude_autonomous_followup_pending(ctx): - # No managed receive_response() consumer owns this continuation. - # Its Result arrives through the background pump, so give Stop its - # own bounded recovery path if that terminal never arrives. + # The scheduler only installs a watchdog without a managed reader; + # otherwise that reader owns Stop until its finalizer hands off. self._schedule_claude_autonomous_interrupt_watchdog(ctx) # A turn can still be reconnecting to apply effort/tier changes and may # not have submitted its query yet. Serialize against that final launch @@ -39081,6 +39089,12 @@ async def freeze_indeterminate_recovery() -> None: await asyncio.gather( codex_restart_watch_task, return_exceptions=True) if not is_codex: + if (ctx.state in {"interrupting", "draining"} + and self._claude_autonomous_followup_pending(ctx)): + # A managed Result can precede an autonomous Result. Keep + # Stop bounded by its original deadline after this reader + # releases the stream; do not leave recovery ownerless. + self._schedule_claude_autonomous_interrupt_watchdog(ctx) # In the managed-finalizer-wins race, the background callback # already attempted these schedulers while state was running. # Retry after final quiescence; both schedulers are idempotent. diff --git a/docs/codex-desktop-launcher.md b/docs/codex-desktop-launcher.md index 122a6c4b..e2e5c457 100644 --- a/docs/codex-desktop-launcher.md +++ b/docs/codex-desktop-launcher.md @@ -14,6 +14,13 @@ App-process-checked loopback WebSocket bridge forwards bytes unchanged to that profile's private Unix control socket. There is no global environment override, model API proxy, App patch, automatic takeover or private stdio fallback. +The launcher and App discovery share the Wrapper's native socket validation. +Both direct private sockets and Codex 0.156's protected socket aliases are +supported. An alias must match the selected account's exact native address hash +under the current user's private daemon directory; arbitrary links, cross-account +targets and sockets with shared permissions remain rejected. The App tools pipe +keeps its separate direct-socket validation. + The [official app-server protocol](https://learn.chatgpt.com/docs/app-server) documents WebSockets over the Unix control socket. WebSocket transport is experimental. The Desktop `CODEX_APP_SERVER_WS_URL` launch override is an diff --git a/tests/test_codex_app_tools.py b/tests/test_codex_app_tools.py index 2e5d6802..8fca6593 100644 --- a/tests/test_codex_app_tools.py +++ b/tests/test_codex_app_tools.py @@ -1,5 +1,7 @@ """Zero-model-turn coverage for the opt-in Desktop tools transport.""" +import hashlib import json +import os from pathlib import Path import shutil import socket @@ -10,6 +12,7 @@ import pytest from cc_remote import codex_app_tools as tools +from cc_remote.wrapper import codex_daemon from cc_remote.wrapper.process_scan import ProcessIdentity @@ -65,15 +68,23 @@ def test_discovery_requires_exact_shared_profile(monkeypatch, tmp_path, home, en "missing_listener", "ambiguous_listener", "foreign_owner", "wrong_peer", "bridge_reused", "app_reused", "missing_daemon", "linked_daemon", ]) -def test_shared_app_requires_profile_bound_bridge_and_both_connection_ends(monkeypatch, failure): +@pytest.mark.parametrize("protected_alias", [False, True]) +def test_shared_app_requires_profile_bound_bridge_and_both_connection_ends(monkeypatch, failure, protected_alias): with tempfile.TemporaryDirectory(prefix="cc-bridge-", dir="/tmp") as root: profile = Path(root).resolve() directory = profile / "app-server-control" directory.mkdir() upstream = directory / "app-server-control.sock" + listener = upstream + if protected_alias: + protected = profile / "p" + protected.mkdir(mode=0o700) + monkeypatch.setattr(codex_daemon, "_protected_socket_directory", lambda _uid: protected) + listener = protected / hashlib.sha256(os.fsencode(upstream)).hexdigest() + upstream.symlink_to(listener) with socket.socket(socket.AF_UNIX) as sock: - sock.bind(str(upstream)) - upstream.chmod(0o600) + sock.bind(str(listener)) + listener.chmod(0o600) if failure == "missing_daemon": upstream.unlink() elif failure == "linked_daemon": @@ -136,6 +147,7 @@ def test_discovery_without_macos_is_optional(monkeypatch, tmp_path): def test_changed_official_policy_cannot_use_old_approval_rules(monkeypatch, tmp_path): monkeypatch.setattr(tools.sys, "platform", "darwin") + monkeypatch.setattr(tools, "socket_identity", lambda _path: (0, 0, 0)) monkeypatch.setattr(tools, "_private_socket", lambda _path: None) monkeypatch.setattr(tools, "app_paths", lambda _app: (tmp_path, tmp_path, tmp_path, "test")) (tmp_path / "desktop-mcp.json").write_text(json.dumps({"mcpServers": {"codex_app": {"tools": {"new_write": {"approval_mode": "prompt"}}}}})) @@ -144,6 +156,7 @@ def test_changed_official_policy_cannot_use_old_approval_rules(monkeypatch, tmp_ def test_discovery_rejects_two_matching_apps_and_pid_reuse(monkeypatch, tmp_path): monkeypatch.setattr(tools.sys, "platform", "darwin") + monkeypatch.setattr(tools, "socket_identity", lambda _path: (0, 0, 0)) monkeypatch.setattr(tools, "app_paths", lambda _app: (tmp_path, tmp_path, tmp_path, "test")) monkeypatch.setattr(tools, "_signed", lambda _path: None) monkeypatch.setattr(tools, "_matching_app_pids", lambda _exe: [41, 42]) @@ -205,6 +218,7 @@ def _ready_app(monkeypatch, tmp_path): node = tmp_path / "node" node.write_text("old runtime") monkeypatch.setattr(tools.sys, "platform", "darwin") + monkeypatch.setattr(tools, "socket_identity", lambda _path: (0, 0, 0)) monkeypatch.setattr(tools, "app_paths", lambda _app: (tmp_path, node, tmp_path, "test")) monkeypatch.setattr(tools, "_signed", lambda _path: None) monkeypatch.setattr(tools, "_matching_app_pids", lambda _exe: [42]) @@ -217,6 +231,36 @@ def _ready_app(monkeypatch, tmp_path): return node +@pytest.mark.parametrize("failure", [None, "cross_account", "shared_directory", "shared_socket"]) +def test_discovery_validates_protected_daemon_alias(monkeypatch, failure): + with tempfile.TemporaryDirectory(prefix="cc-app-", dir="/tmp") as root: + profile = Path(root).resolve() + control = profile / "app-server-control" + control.mkdir(mode=0o700) + address = control / "app-server-control.sock" + protected = profile / "p" + protected.mkdir(mode=0o700) + native_address = address if failure != "cross_account" else profile / "other.sock" + listener = protected / hashlib.sha256(os.fsencode(native_address)).hexdigest() + address.symlink_to(listener) + pipe = profile / "tools.sock" + private_socket = tools._private_socket + _ready_app(monkeypatch, profile) + monkeypatch.setattr(codex_daemon, "_protected_socket_directory", lambda _uid: protected) + monkeypatch.setattr(tools, "socket_identity", codex_daemon.socket_identity) + monkeypatch.setattr(tools, "_private_socket", private_socket) + monkeypatch.setattr(tools, "_pipe_from_open_logs", lambda *_args: pipe) + with socket.socket(socket.AF_UNIX) as daemon, socket.socket(socket.AF_UNIX) as app: + daemon.bind(str(listener)) + app.bind(str(pipe)) + listener.chmod(0o666 if failure == "shared_socket" else 0o600) + pipe.chmod(0o600) + if failure == "shared_directory": + protected.chmod(0o755) + result = tools.discover(profile, profile) + assert result["state"] == ("ready" if failure is None else "unavailable") + + def test_runtime_update_changes_generation_without_app_or_pipe_restart(monkeypatch, tmp_path): node = _ready_app(monkeypatch, tmp_path) before = tools.discover(tmp_path, tmp_path) diff --git a/tests/test_codex_desktop.py b/tests/test_codex_desktop.py index a02b8d34..b4719c73 100644 --- a/tests/test_codex_desktop.py +++ b/tests/test_codex_desktop.py @@ -1,5 +1,6 @@ """No model calls or real App lifecycle changes: desktop launch regressions.""" import asyncio +import hashlib import json import os from pathlib import Path @@ -11,6 +12,7 @@ from websockets.asyncio.server import unix_serve from cc_remote import codex_desktop as launcher +from cc_remote.wrapper import codex_daemon from cc_remote.wrapper.process_scan import ProcessIdentity @@ -108,12 +110,20 @@ def process_identity(_pid): @pytest.mark.asyncio -async def test_real_proxy_roundtrip_and_reject_self_after_preflight(monkeypatch): +@pytest.mark.parametrize("protected_alias", [False, True]) +async def test_real_proxy_roundtrip_and_reject_self_after_preflight(monkeypatch, protected_alias): with tempfile.TemporaryDirectory(prefix="cc-launch-", dir="/tmp") as root: profile = Path(root).resolve() directory = profile / "app-server-control" directory.mkdir(mode=0o700) path = directory / "app-server-control.sock" + listener = path + if protected_alias: + protected = profile / "p" + protected.mkdir(mode=0o700) + monkeypatch.setattr(codex_daemon, "_protected_socket_directory", lambda _uid: protected) + listener = protected / hashlib.sha256(os.fsencode(path)).hexdigest() + path.symlink_to(listener) frames = [] async def echo(ws): @@ -132,8 +142,8 @@ async def command(*args): return f"p{os.getpid()}\nu{os.getuid()}\nn127.0.0.1:{peer_port}->127.0.0.1:{bridge.port}\n" monkeypatch.setattr(launcher, "command", command) - async with unix_serve(echo, str(path)): - path.chmod(0o600) + async with unix_serve(echo, str(listener)): + listener.chmod(0o600) async with await asyncio.start_server(bridge.handle, "127.0.0.1", 0, limit=16384) as server: bridge.port = server.sockets[0].getsockname()[1] await bridge.preflight() diff --git a/tests/test_wrapper_core_fixes.py b/tests/test_wrapper_core_fixes.py index 94576dcc..0a59b035 100644 --- a/tests/test_wrapper_core_fixes.py +++ b/tests/test_wrapper_core_fixes.py @@ -944,6 +944,111 @@ async def run(): asyncio.run(run()) +@pytest.mark.parametrize("watchdog_first", [False, True]) +def test_claude_managed_interrupt_owns_recovery_with_background_activity(watchdog_first): + async def run(): + machine, transport = _mk_machine() + machine.cfg.drain_timeout = 0.03 + ctx = _mk_ctx("claude-stop-owner", "claude-stop-owner") + ctx.engine = "claude" + ctx.state = "running" + ctx.active_msg_id = "stop-message" + machine.sessions[ctx.key] = ctx + + class StalledSdk(_ClaudeStalledSdk): + async def force_reconnect(self, resume_id, cwd, **_kwargs): + self.reconnects += 1 + # Real disconnect/connect yields before lifecycle reset. Both + # expired consumers must not reconnect the same native child. + await asyncio.sleep(0.02) + machine._reset_claude_task_lifecycle(ctx) + + sdk = StalledSdk() + ctx.sdk = sdk + + async def no_external_owner(_sid): + return False + + machine._prime_claude_ownership = no_external_owner + turn = asyncio.create_task(machine._run_turn(ctx, "inspect")) + ctx.turn_task = turn + await asyncio.wait_for(sdk.reader_started.wait(), timeout=1) + ctx.claude_background_followups["background-request"] = "active" + if watchdog_first: + # An existing watchdog can gain a managed owner before its deadline + # when an already accepted steering input completes its handoff. + ctx.turn_task = None + await machine._handle_interrupt(SimpleNamespace(sid=ctx.key)) + watchdog = ctx.claude_autonomous_interrupt_task + ctx.turn_task = turn + await asyncio.wait_for(turn, timeout=1) + if watchdog is not None: + await asyncio.gather(watchdog, return_exceptions=True) + + assert sdk.reconnects == 1 + assert ctx.state == "idle" + assert ctx.claude_autonomous_interrupt_task is None + failures = [msg for msg in transport.sent + if getattr(msg, "code", None) == ERR_DRAIN_TIMEOUT] + assert len(failures) == 1 + assert failures[0].msg_id == "stop-message" + + asyncio.run(run()) + + +def test_claude_interrupt_hands_remaining_background_work_the_original_deadline(): + async def run(): + machine, transport = _mk_machine() + machine.cfg.drain_timeout = 0.1 + ctx = _mk_ctx("claude-stop-handoff", "claude-stop-handoff") + ctx.engine = "claude" + ctx.state = "running" + ctx.active_msg_id = "stop-message" + machine.sessions[ctx.key] = ctx + + class StalledSdk(_ClaudeStalledSdk): + async def force_reconnect(self, resume_id, cwd, **_kwargs): + self.reconnects += 1 + machine._reset_claude_task_lifecycle(ctx) + + sdk = StalledSdk() + sdk.responses = [ResultMessage( + subtype="error_during_execution", duration_ms=1, + duration_api_ms=1, is_error=True, num_turns=1, + session_id=ctx.session_id, + )] + ctx.sdk = sdk + + async def no_external_owner(_sid): + return False + + machine._prime_claude_ownership = no_external_owner + turn = asyncio.create_task(machine._run_turn(ctx, "inspect")) + ctx.turn_task = turn + await asyncio.wait_for(sdk.reader_started.wait(), timeout=1) + ctx.claude_background_followups["background-request"] = "active" + await machine._handle_interrupt(SimpleNamespace(sid=ctx.key)) + deadline = ctx.interrupt_deadline + # A handoff must retain the accepted Stop's deadline, not start a fresh + # wait using the config value at the managed terminal. + machine.cfg.drain_timeout = 100 + sdk.release.set() + await asyncio.wait_for(turn, timeout=1) + + assert ctx.state == "interrupting" + assert ctx.interrupt_deadline == deadline + watchdog = ctx.claude_autonomous_interrupt_task + assert watchdog is not None + await asyncio.wait_for(watchdog, timeout=1) + assert sdk.reconnects == 1 + assert ctx.state == "idle" + assert ctx.claude_autonomous_interrupt_task is None + assert len([msg for msg in transport.sent + if getattr(msg, "code", None) == ERR_DRAIN_TIMEOUT]) == 1 + + asyncio.run(run()) + + def test_codex_silence_emits_no_synthetic_waiting_notice_and_later_completes( monkeypatch): async def run(): diff --git a/web/src/problem-presentation.ts b/web/src/problem-presentation.ts index 27688764..7bac52d3 100644 --- a/web/src/problem-presentation.ts +++ b/web/src/problem-presentation.ts @@ -5,19 +5,21 @@ const INCOMPLETE_CLAUDE_SESSION = "Claude 会话历史不完整,无法恢复;可从会话菜单删除该条目。"; const PROVIDER_AUTH_TURN_FAILURE = "模型服务认证已失效或当前账号无权限,请检查当前服务的凭据或账号权限后重试。"; -const LEGACY_CODEX_AUTH_TURN_FAILURE = - "Codex 登录已失效或当前账号无权限,请重新登录后重试。"; const CODEX_UPDATE_INTERRUPTION = "Codex 自动更新时连接中断,本轮未确认完成。请检查已有结果后继续。"; const CODEX_CONNECTION_INTERRUPTION = "与 Codex 的连接中断,本轮未确认完成。请检查已有结果后继续。"; +const STOP_TIMEOUT_FAILURE = "停止确认超时,本轮未确认完成。"; // Exact authored causes only; never expose an arbitrary upstream diagnostic. const STEER_REJECTION_MESSAGES = /^(?:该会话(?:未启动,无法引导当前任务|当前为只读状态,无法从 Remote 引导)|本次引导已取消。|Claude 当前无法接收引导,本次未发送;请稍后重试或排队。|Codex (?:(?:自动压缩|Review|当前阶段)不支持引导,或任务已经切换;本次未发送。|正在核对当前回合归属,本次引导未发送;请稍后重试。|当前回合不支持引导,请等待后重试。)|消息内容为空,请输入内容或添加附件。|附件不符合要求,请调整后重试。)$/; const CODEX_USAGE_LIMIT_FAILURE = "本轮使用的 Codex 账号额度已用完。可切换账号、补充额度,或等待恢复后重试。"; const CODEX_USAGE_LIMIT_RETRY = /^官方提示可于 ((?!0000)[0-9]{4}-[0-9]{2}-[0-9]{2} [0-9]{2}:[0-9]{2})(设备当地时间)重试。$/; -const CODEX_TRANSPORT_MESSAGES = new Map([ +// Normalize exact legacy/replayed copy before applying the safe-message filter. +const TURN_FAILURE_ALIASES = new Map([ + ["Codex 登录已失效或当前账号无权限,请重新登录后重试。", PROVIDER_AUTH_TURN_FAILURE], + ["停止操作未及时完成,会话正在恢复。", STOP_TIMEOUT_FAILURE], ["Codex 已自动更新,当前回合在更新时中断;为避免重复执行工具," + "本次任务未自动重试。请确认已有结果后重新发送。", CODEX_UPDATE_INTERRUPTION], ["Codex 共享通道意外断开;为避免重复执行工具,本次任务未自动重试。" @@ -34,6 +36,7 @@ const SAFE_TURN_FAILURE_MESSAGES = new Set([ "当前模型繁忙,请稍后重试或切换模型。", CODEX_UPDATE_INTERRUPTION, CODEX_CONNECTION_INTERRUPTION, + STOP_TIMEOUT_FAILURE, "上游模型因安全策略拒绝了本次请求(cyber_policy)。" + "这不是本地权限或网络错误;请核实并说明任务背景与授权范围," + "若属误判请向服务提供方反馈。", @@ -52,11 +55,8 @@ function isUsageLimitFailure(message: string): boolean { function safeTurnFailureMessage(message: string): string | null { const trimmed = message.trim(); - if (trimmed === LEGACY_CODEX_AUTH_TURN_FAILURE) { - return PROVIDER_AUTH_TURN_FAILURE; - } - const transportMessage = CODEX_TRANSPORT_MESSAGES.get(trimmed); - if (transportMessage) return transportMessage; + const alias = TURN_FAILURE_ALIASES.get(trimmed); + if (alias) return alias; return SAFE_TURN_FAILURE_MESSAGES.has(trimmed) || isUsageLimitFailure(trimmed) ? trimmed : null; } @@ -74,6 +74,8 @@ export function presentTurnOutcome( const cause = codexTransportInterruption(message); if (cause === "update") return "Codex 自动升级,本轮中断"; if (cause === "connection") return "连接中断,回复未完成"; + if (outcome === "failed" && message + && safeTurnFailureMessage(message) === STOP_TIMEOUT_FAILURE) return "停止确认超时"; if (outcome === "failed" && message && isUsageLimitFailure(message.trim())) return "账号额度已用完"; return outcome === "interrupted" ? "已打断" : "回复未完成"; } @@ -100,7 +102,7 @@ export function presentTurnProblem(error: Pick): s return "消息内容无法发送,请检查输入或附件后重试。"; } if (error.code === "drain_timeout") { - return "停止操作未及时完成,会话正在恢复。"; + return STOP_TIMEOUT_FAILURE; } if (error.code === "cc_crash") { return safeTurnFailureMessage(error.message) @@ -113,6 +115,8 @@ export function presentCommandProblem( error: Pick, ): string { switch (error.code) { + case "drain_timeout": + return STOP_TIMEOUT_FAILURE; case "wrapper_offline": return "设备正在重新连接…"; case "invalid_cwd": diff --git a/web/tests/notices-rate-limits.test.ts b/web/tests/notices-rate-limits.test.ts index eb7bd133..030f989e 100644 --- a/web/tests/notices-rate-limits.test.ts +++ b/web/tests/notices-rate-limits.test.ts @@ -394,6 +394,43 @@ try { /crash|warning|wrapper|private|traceback|secret/i); const hiddenDiagnostic = "provider crash at /private/token; see wrapper logs"; + const stopTimeoutMessage = "停止确认超时,本轮未确认完成。"; + const stopTimeout = { code: "drain_timeout", message: hiddenDiagnostic }; + assert.equal(presentTurnProblem(stopTimeout), stopTimeoutMessage); + assert.equal(presentCommandProblem(stopTimeout), stopTimeoutMessage); + for (const message of [stopTimeoutMessage, "停止操作未及时完成,会话正在恢复。"]) { + assert.equal(presentHistoricalTurnProblem(message), stopTimeoutMessage, + "stop timeouts must survive live-to-history presentation and old caches"); + assert.equal(presentTurnOutcome("failed", message), "停止确认超时"); + } + const stopSid = "claude-stop-timeout"; + const stopping = { ...initialState, focusedSid: stopSid, + sessions: [{ session_id: stopSid, engine: "claude", space: "code" }], + runtimes: { [stopSid]: { ...createRuntime(), state: "interrupting", + turns: [{ id: "stop-message", prompt: "inspect", blocks: [], done: false }], + } }, + }; + const timedOut = reduce(stopping, { type: "event", event: event({ + type: "error", sid: stopSid, msg_id: "stop-message", ...stopTimeout, + }) }); + const failedStop = timedOut.runtimes[stopSid].turns[0]; + assert.equal(failedStop.done, true); + assert.equal(failedStop.error, stopTimeoutMessage); + assert.notEqual(failedStop.interrupted, true, + "a stop timeout cannot invent a confirmed native interruption"); + const { default: TurnProblem } = await harness.ssrLoadModule("/src/components/TurnProblem.tsx"); + const stopMarkup = renderToStaticMarkup(createElement(TurnProblem, { + message: failedStop.error, continuing: false, + })); + assert.match(stopMarkup, /停止确认超时,本轮未确认完成/); + assert.doesNotMatch(stopMarkup, /该轮未正常结束|正在恢复|private|crash/); + const confirmedStop = reduce(stopping, { type: "event", event: event({ + type: "turn_end", sid: stopSid, turn_id: "stop-message", + result: { subtype: "error_during_execution", is_error: true, duration_ms: 1 }, + }) }).runtimes[stopSid].turns[0]; + assert.equal(confirmedStop.interrupted, true); + assert.equal(confirmedStop.error, undefined); + assert.equal(presentTurnOutcome("interrupted"), "已打断"); const quotaFailure = "本轮使用的 Codex 账号额度已用完。可切换账号、补充额度,或等待恢复后重试。"; const quotaRetry = quotaFailure + "官方提示可于 2026-09-15 09:24(设备当地时间)重试。"; const policyFailure = "上游模型因安全策略拒绝了本次请求(cyber_policy)。" From 285aaa4d966315b55a54f6e52621374b21989224 Mon Sep 17 00:00:00 2001 From: muggle-stack Date: Fri, 25 Sep 2026 16:21:17 +0800 Subject: [PATCH 2/3] fix(web): restore subagent activity animations - Share the animated working sparkle across agent engines. - Show a completion sparkle after successful agent runs. - Keep process disclosures interactive and scoped to each agent. --- web/src/components/AgentDetailPanel.tsx | 23 ++++++++++++++++++----- 1 file changed, 18 insertions(+), 5 deletions(-) diff --git a/web/src/components/AgentDetailPanel.tsx b/web/src/components/AgentDetailPanel.tsx index f5781736..aa3531a3 100644 --- a/web/src/components/AgentDetailPanel.tsx +++ b/web/src/components/AgentDetailPanel.tsx @@ -1,7 +1,8 @@ +import { useState } from "react"; import type { AgentDetailRun } from "../agent-detail"; import type { Engine } from "../protocol"; import { finalTextBlocks, presentableProcessBlocks } from "../process-blocks"; -import { ClaudeWorking, EngineIcon, Icon } from "../icons"; +import { ClaudeSpark, ClaudeWorking, Icon } from "../icons"; import { MessageBlock } from "./MessageBlock"; import { ProcessTimeline } from "./ProcessTimeline"; import { PanelResizer } from "./PanelResizer"; @@ -26,6 +27,7 @@ export function AgentDetailPanel({ run, engine = "claude", canGoBack, onBack, on onOpenAgent: (runId: string, title?: string) => void; onOpenFile?: (path: string, line?: number) => void; }) { + const [processOpen, setProcessOpen] = useState>({}); const process = presentableProcessBlocks(run.blocks, engine); const final = finalTextBlocks(run.blocks); const done = !["running", "pending"].includes(run.status); @@ -67,18 +69,29 @@ export function AgentDetailPanel({ run, engine = "claude", canGoBack, onBack, on )} {process.length > 0 && ( - setProcessOpen((current) => ({ + ...current, [run.runId]: open, + }))} onOpenAgent={onOpenAgent} onOpenFile={onOpenFile} /> )} {final.map((block) => ( ))} - {!done &&
- {engine === "claude" ? : } + {!done &&
+ 子代理处理中
} + {run.status === "succeeded" && run.blocks.length > 0 && ( +
+ +
+ )} {!run.loading && !run.error && run.blocks.length === 0 && (
这个协作代理暂时没有可展示的过程。
)} From 609c73b3e247e5b0c6ee80458c47560aa99713d0 Mon Sep 17 00:00:00 2001 From: muggle-stack Date: Sun, 27 Sep 2026 01:20:30 -0700 Subject: [PATCH 3/3] fix(claude): keep remote sessions responsive during slow operations - Schedule bounded commands per session while preserving lifecycle order and reliable receipts. - Isolate native startup, relay sends and slow client cleanup from unrelated sessions. - Initialize BTW model selection from the parent and unblock the main composer shortcut when settings are closed. - Cover startup, interruption, command ordering and relay isolation regressions. --- cc_remote/claude_service/server.py | 94 ++++-- cc_remote/relay/forward.py | 33 +- cc_remote/relay/pairing.py | 23 +- cc_remote/wrapper/command_scheduler.py | 155 ++++++++++ cc_remote/wrapper/machine.py | 126 +++++--- tests/test_claude_btw_identity.py | 38 ++- tests/test_claude_live_context.py | 46 ++- tests/test_claude_service_startup.py | 196 ++++++++++++ tests/test_command_scheduling.py | 402 +++++++++++++++++++++++++ tests/test_relay_forward.py | 120 ++++++++ web/src/components/BtwPanel.tsx | 12 +- 11 files changed, 1147 insertions(+), 98 deletions(-) create mode 100644 cc_remote/wrapper/command_scheduler.py create mode 100644 tests/test_claude_service_startup.py create mode 100644 tests/test_command_scheduling.py diff --git a/cc_remote/claude_service/server.py b/cc_remote/claude_service/server.py index 27904b60..eabab751 100644 --- a/cc_remote/claude_service/server.py +++ b/cc_remote/claude_service/server.py @@ -489,6 +489,68 @@ def __init__(self, directory: Path, *, factory=None): self.factory = factory self.sessions: dict[str, Session] = {} self.open_lock = asyncio.Lock() + self._opening: dict[str, asyncio.Event] = {} + + @staticmethod + def _check_open_identity(session: Session, identity: dict, owner) -> None: + if any(session.metadata.get(key) != identity.get(key) for key in ( + "profile_root", "session_id", "space", "work_id", "btw", "cwd", + )): + raise PermissionError("Claude session identity mismatch") + if session.controller is not None and session.controller is not owner: + raise ControllerLeaseConflict("Claude service already has a controller") + + async def _open(self, owner, params): + # Only reserve/check identity and capacity under the service-wide lock. + # Native connect can wait on a slow CLI; unrelated sessions stay usable. + async with self.open_lock: + identity = params["metadata"] + session = self.sessions.get(params.get("session")) + if params.get("strict_session") and params.get("session") is not None and session is None: + raise KeyError("Claude service recovery worker no longer exists") + if session is None and identity.get("session_id") and not params.get("fork"): + session = next((item for item in self.sessions.values() if all( + item.metadata.get(key) == identity.get(key) + for key in ("profile_root", "session_id", "space", "work_id", "btw") + )), None) + attached = session is not None + if session is not None: + self._check_open_identity(session, identity, owner) + opening = self._opening.get(session.id) + else: + if len(self.sessions) >= MAX_SESSIONS: + raise RuntimeError("Claude service session capacity reached") + session = Session(self.directory, identity, self.factory) + self.sessions[session.id] = session + opening = self._opening[session.id] = asyncio.Event() + + if not attached: + try: + await session.start(params["options"], params.get("isolated", False)) + session.controller = owner + return {**session.description(), "attached": False} + except BaseException: + # Retain the reservation until startup cleanup finishes. A + # concurrent open must not create a second native writer. + try: + await session.close() + finally: + self.sessions.pop(session.id, None) + raise + finally: + self._opening.pop(session.id, None) + opening.set() + + if opening is not None: + # Cancelling a joining controller only cancels its own wait, never + # the creator's startup. Failure does not silently spawn a new CLI. + await opening.wait() + async with self.open_lock: + if self.sessions.get(session.id) is not session or session.closed: + raise RuntimeError("Claude service session is no longer available") + self._check_open_identity(session, identity, owner) + session.controller = owner + return {**session.description(), "attached": True} async def connection(self, reader, writer) -> None: if not same_user(writer): @@ -546,37 +608,7 @@ async def dispatch(self, owner, request_id, method, params): if method == "list": return [session.description() for session in self.sessions.values() if not session.closed] if method == "open": - async with self.open_lock: - identity = params["metadata"] - session = self.sessions.get(params.get("session")) - if params.get("strict_session") and params.get("session") is not None and session is None: - raise KeyError("Claude service recovery worker no longer exists") - if session is None and identity.get("session_id") and not params.get("fork"): - session = next((item for item in self.sessions.values() if all( - item.metadata.get(key) == identity.get(key) - for key in ("profile_root", "session_id", "space", "work_id", "btw") - )), None) - attached = session is not None - if session is not None: - if any(session.metadata.get(key) != identity.get(key) for key in ( - "profile_root", "session_id", "space", "work_id", "btw", "cwd", - )): - raise PermissionError("Claude session identity mismatch") - if session.controller is not None and session.controller is not owner: - raise ControllerLeaseConflict("Claude service already has a controller") - else: - if len(self.sessions) >= MAX_SESSIONS: - raise RuntimeError("Claude service session capacity reached") - session = Session(self.directory, identity, self.factory) - self.sessions[session.id] = session - try: - await session.start(params["options"], params.get("isolated", False)) - except BaseException: - self.sessions.pop(session.id, None) - await session.close() - raise - session.controller = owner - return {**session.description(), "attached": attached} + return await self._open(owner, params) session = self.sessions[params["session"]] if session.controller is not owner: raise PermissionError("Claude service controller lease is required") diff --git a/cc_remote/relay/forward.py b/cc_remote/relay/forward.py index 14806ea1..c6a9aa2a 100644 --- a/cc_remote/relay/forward.py +++ b/cc_remote/relay/forward.py @@ -16,6 +16,9 @@ from cc_remote.protocol import serialize log = logger("cc_remote.relay.forward") +# Let the WebSocket backend finish its handshake/TCP close deadlines before +# this last-resort guard. Cleanup runs off the forwarding path throughout. +CLIENT_CLOSE_TIMEOUT = 35.0 class SlowClientError(RuntimeError): @@ -37,6 +40,7 @@ def __init__(self, ws, cap: int, client_id: str, self.queue: asyncio.Queue[tuple[str, int]] = asyncio.Queue(maxsize=self.cap) self._queued_bytes = 0 self._sender: Optional[asyncio.Task] = None + self._stop_task: Optional[asyncio.Task] = None self._closed = False @property @@ -50,16 +54,27 @@ def queued_bytes(self) -> int: def start(self) -> None: self._sender = asyncio.create_task(self._run()) + def begin_stop(self, *, code: int | None = None, reason: str = "") -> asyncio.Task: + """Retire this connection synchronously; finish socket cleanup off-path.""" + if self._stop_task is None: + self._closed = True + if self._sender and self._sender is not asyncio.current_task(): + self._sender.cancel() + self._stop_task = asyncio.create_task(self._finish_stop(code, reason)) + return self._stop_task + async def stop(self, *, code: int | None = None, reason: str = "") -> None: - already_closed = self._closed - self._closed = True - if code is not None and not already_closed: + # Cancellation of the reader must not abandon the one close operation. + await asyncio.shield(self.begin_stop(code=code, reason=reason)) + + async def _finish_stop(self, code: int | None, reason: str) -> None: + if code is not None: try: - await self.ws.close(code=code, reason=reason) + async with asyncio.timeout(CLIENT_CLOSE_TIMEOUT): + await self.ws.close(code=code, reason=reason) except Exception: pass - if self._sender and self._sender is not asyncio.current_task(): - self._sender.cancel() + if self._sender: try: await self._sender except (asyncio.CancelledError, Exception): @@ -88,9 +103,5 @@ async def _run(self) -> None: except asyncio.CancelledError: raise except Exception as e: - self._closed = True log.debug("client sender ended", client_id=self.client_id, error=str(e)) - try: - await self.ws.close(code=1011, reason="client sender failed") - except Exception: - pass + self.begin_stop(code=1011, reason="client sender failed") diff --git a/cc_remote/relay/pairing.py b/cc_remote/relay/pairing.py index cdf40104..df66769f 100644 --- a/cc_remote/relay/pairing.py +++ b/cc_remote/relay/pairing.py @@ -18,6 +18,7 @@ import json import re import uuid +from weakref import WeakValueDictionary from typing import Awaitable, Callable, Optional from fastapi import WebSocket, WebSocketDisconnect @@ -72,9 +73,19 @@ def __init__( # Linearizes client generation replacement with client -> wrapper sends. # Once a new generation owns client_id, no old generation can enter a # wrapper send after that point. - self._wrapper_send_lock = asyncio.Lock() + self._wrapper_send_locks: WeakValueDictionary[str, asyncio.Lock] = WeakValueDictionary() self._wrapper_versions: OrderedDict[str, dict] = OrderedDict() + def _wrapper_send_lock_for(self, machine_id: str) -> asyncio.Lock: + # Hold a strong reference through acquire/release (including waiters). + # Idle machine names must not accumulate locks, and a stalled machine's + # network write must never hold another machine's registration/send gate. + lock = self._wrapper_send_locks.get(machine_id) + if lock is None: + lock = asyncio.Lock() + self._wrapper_send_locks[machine_id] = lock + return lock + def remember_wrapper_version(self, machine_id: str, protocol: int, version: str | None) -> None: if not isinstance(version, str) or not re.fullmatch(r"\d{1,6}\.\d{1,6}\.\d{1,6}", version): version = None @@ -466,7 +477,7 @@ async def serve_client(self, ws: WebSocket, msg.route_id = conn.route_id msg.owner_id = conn.owner_id over_capacity = False - async with self._wrapper_send_lock: + async with self._wrapper_send_lock_for(machine_id): async with self._lock: clients = self._clients_for(machine_id) old = clients.get(client_id) @@ -483,7 +494,7 @@ async def serve_client(self, ws: WebSocket, ) return if old is not None and old is not conn: - await old.stop(code=4009, reason="replaced by reconnect") + old.begin_stop(code=4009, reason="replaced by reconnect") log.info("client registered", client_id=client_id, machine_id=machine_id, total=self.client_count) @@ -542,7 +553,7 @@ async def _forward_client_msg( machine_id: str = "default", ) -> bool: """Forward iff ``conn`` still owns client_id at the send linearization point.""" - async with self._wrapper_send_lock: + async with self._wrapper_send_lock_for(machine_id): async with self._lock: current = self._clients_for(machine_id).get(client_id) is conn wrapper = self._wrapper_for(machine_id) @@ -614,4 +625,6 @@ async def _drop_client(self, conn: ClientConn, *, code: int | None = None, if clients.get(conn.client_id) is conn: del clients[conn.client_id] self._prune_clients_for(machine_id, clients) - await conn.stop(code=code, reason=reason) + # Remove the route and stop its sender immediately. Closing a broken + # browser is bounded connection cleanup, never part of wrapper ingress. + conn.begin_stop(code=code, reason=reason) diff --git a/cc_remote/wrapper/command_scheduler.py b/cc_remote/wrapper/command_scheduler.py new file mode 100644 index 00000000..0d8e6a90 --- /dev/null +++ b/cc_remote/wrapper/command_scheduler.py @@ -0,0 +1,155 @@ +"""Bounded command intake with ordered mutations and independent sessions. + +Handlers still own session locks, reliable receipts and native turn lifetimes. +This scheduler only orders command *handlers*, never the turns they launch. +""" +from __future__ import annotations + +import asyncio +from dataclasses import dataclass +from typing import Awaitable, Callable + +from cc_remote.log import logger +from cc_remote.protocol import serialize + +log = logger("cc_remote.wrapper.command_scheduler") + + +@dataclass +class _Pending: + task: asyncio.Task + kind: str + target: str | None + serial: bool + barrier: bool + size: int + + +class CommandScheduler: + def __init__( + self, + process: Callable[[object], Awaitable[None]], + resolve_target: Callable[[str], str], + *, + max_items: int = 128, + max_bytes: int = 32 * 1024 * 1024, + ) -> None: + self._process = process + self._resolve_target = resolve_target + self._max_items = max_items + self._max_bytes = max_bytes + self._pending: dict[tuple[str, str], _Pending] = {} + self._bytes = 0 + self._changed = asyncio.Event() + + def owns_target(self, target: str) -> bool: + """Pin a resident context even before its queued handler takes a lock.""" + target = self._resolve_target(target) + return any( + not pending.task.done() and pending.target is not None + and self._resolve_target(pending.target) == target + for pending in self._pending.values() + ) + + async def submit( + self, command, *, target: str | None = None, + serial: bool = True, barrier: bool = False, urgent: bool = False, + ) -> None: + received = asyncio.get_running_loop().time() + fields = { + "type": command.type, + "sid": target, + "client_id": getattr(command, "client_id", None), + "cmd_id": getattr(command, "cmd_id", None), + "msg_id": getattr(command, "msg_id", None), + } + # Deliberately omit prompts, attachments, SDK stderr and credentials. + log.info("command received", **fields) + key = ( + getattr(command, "client_id", None) or "", + getattr(command, "cmd_id", None) or f"untracked-{id(command)}", + ) + # Include decoded attachments in the budget instead of counting tasks + # alone. A blocked SDK must not turn the transport's bounded inbox into + # an unbounded collection of tasks retaining multi-megabyte prompts. + size = len(serialize(command).encode("utf-8")) + while True: + current = self._pending.get(key) + if current is not None and not current.task.done(): + return + # One oversize item can run alone (the transport already bounds a + # single frame). Keep ordinary overload as backpressure, not drops. + if (len(self._pending) < self._max_items + and (self._bytes + size <= self._max_bytes + or not self._pending)): + break + self._changed.clear() + await self._changed.wait() + + dependencies = [] + for pending in self._pending.values(): + if pending.task.done() or urgent: + continue + same_target = ( + target is not None and pending.target is not None + and self._resolve_target(target) + == self._resolve_target(pending.target) + ) + # Stop may pass an explicitly addressed metadata read, but not the + # preceding resume/query/mutation that makes its target runnable. + # Keep the read alive so it still returns its report and receipt. + if (same_target and command.type == "interrupt" + and pending.kind == "get_context"): + continue + if ((serial and pending.serial) or barrier or pending.barrier + or same_target): + dependencies.append(pending.task) + + async def run() -> None: + for dependency in dependencies: + # Cancelling this command must not cancel an earlier handler + # which another client/session is still waiting for. + await asyncio.shield(dependency) + started = asyncio.get_running_loop().time() + log.info("command dispatch", **fields, + queued_ms=round((started - received) * 1000)) + try: + await self._process(command) + finally: + log.info("command handler finished", **fields, + elapsed_ms=round((asyncio.get_running_loop().time() + - started) * 1000)) + + task = asyncio.create_task(run()) + pending = _Pending(task, command.type, target, serial, barrier, size) + # A completed task's callback may not have run yet. Retire its budget + # before replacing it with a reliable replay of the same command id. + if current is not None: + self._bytes -= current.size + self._pending[key] = pending + self._bytes += size + + def finished(_task: asyncio.Task) -> None: + if self._pending.get(key) is pending: + self._pending.pop(key) + self._bytes -= size + self._changed.set() + + task.add_done_callback(finished) + + async def drain(self) -> None: + """A finite input source ends after its admitted commands complete.""" + while self._pending: + # gather() of already-finished tasks can return synchronously on + # Python 3.13, starving the callbacks that retire their entries. + self._changed.clear() + await self._changed.wait() + + async def close(self) -> None: + tasks = [pending.task for pending in self._pending.values()] + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + self._pending.clear() + self._bytes = 0 + self._changed.set() diff --git a/cc_remote/wrapper/machine.py b/cc_remote/wrapper/machine.py index 27a366b6..56e92299 100644 --- a/cc_remote/wrapper/machine.py +++ b/cc_remote/wrapper/machine.py @@ -399,6 +399,7 @@ process_owner_uid, ) from cc_remote.wrapper.command_router import CommandRouter, UNHANDLED_COMMAND +from cc_remote.wrapper.command_scheduler import CommandScheduler from cc_remote.wrapper.preview_capabilities import ( PREVIEW_PATH_MAX_BYTES, PreviewCapability, @@ -2024,6 +2025,8 @@ def __init__(self, cfg: WrapperConfig, transport: WrapperTransport): self._timed_task_catalog = {} self.transport = transport self._command_router = CommandRouter(self) + self._command_scheduler = CommandScheduler( + self._process_command_safely, self._command_target_identity) self.instance_id = uuid4().hex # Each configured CLAUDE_CONFIG_DIR is an independent account/session # boundary. Keep implicit single-account mode source-compatible and @@ -9035,44 +9038,10 @@ async def run(self) -> None: self._work_schedule_task = asyncio.create_task( self._work_schedule_loop()) async for cmd in self.transport.incoming(): - if cmd.type == "list_sessions": - self._start_session_list_command(cmd) - continue - if (cmd.type in {"get_history", "get_turn_detail", "get_agent_detail", "get_turn_file_changes"} - or (cmd.type == "get_diff" and getattr(cmd, "turn_id", None))): - self._start_history_command(cmd) - continue - if cmd.type == "get_models": - self._start_models_command(cmd) - continue - if cmd.type == "steer": - # turn/steer is a short app-server RPC, but it must not - # monopolize the serial command lane and delay an explicit - # Stop arriving from another client. - self._start_interactive_control_command(cmd) - continue - if cmd.type == "get_permission_profiles": - self._start_models_command(cmd) - continue - if cmd.type in { - "get_status", "consume_rate_limit_reset_credit", - }: - self._start_status_command(cmd) - continue - if cmd.type in { - "get_engine_capabilities", "manage_engine_plugin", - "manage_engine_skill", "manage_engine_hook", - }: - self._start_capabilities_command(cmd) - continue - if cmd.type == "set_model": - ctx = self._ctx_for(getattr(cmd, "sid", None)) - if (ctx is not None and ctx.engine == "claude" - and ctx.space == "code"): - self._start_interactive_control_command(cmd) - continue - await self._process_command_safely(cmd) + await self._dispatch_incoming_command(cmd) + await self._command_scheduler.drain() finally: + await self._command_scheduler.close() models_tasks = list(self._models_command_tasks.values()) for task in models_tasks: task.cancel() @@ -11469,6 +11438,62 @@ async def _send_command_ack(self, client_id: str, cmd_id: str) -> None: to=client_id, )) + def _command_target_identity(self, sid: str) -> str: + # Resolve at admission/comparison time: a running command may still + # carry a tmp key after its context captures the real native identity. + ctx = self._ctx_for(sid) + return (ctx.key if ctx is not None else None) or ( + self._resolve_session_alias(sid) or sid) + + async def _dispatch_incoming_command(self, cmd) -> None: + """Admit commands without waiting for SDK I/O on the receive loop.""" + if cmd.type == "list_sessions": + self._start_session_list_command(cmd) + return + if (cmd.type in {"get_history", "get_turn_detail", "get_agent_detail", "get_turn_file_changes"} + or (cmd.type == "get_diff" and getattr(cmd, "turn_id", None))): + self._start_history_command(cmd) + return + if cmd.type in {"get_models", "get_permission_profiles"}: + self._start_models_command(cmd) + return + if cmd.type == "steer": + self._start_interactive_control_command(cmd) + return + if cmd.type in {"get_status", "consume_rate_limit_reset_credit"}: + self._start_status_command(cmd) + return + if cmd.type in { + "get_engine_capabilities", "manage_engine_plugin", + "manage_engine_skill", "manage_engine_hook", + }: + self._start_capabilities_command(cmd) + return + if cmd.type == "set_model": + ctx = self._ctx_for(getattr(cmd, "sid", None)) + if ctx is not None and ctx.engine == "claude" and ctx.space == "code": + self._start_interactive_control_command(cmd) + return + + target = (getattr(cmd, "session_id", None) + or getattr(cmd, "sid", None)) + urgent = cmd.type in {"ping", "answer_question"} + # Mutations/focus changes retain their single ordered lane. Explicitly + # addressed query/context/stop commands only wait for that session's + # earlier handlers, including a cold SwitchSession. A legacy command + # without sid is a barrier: resolve focus after preceding switches. + # NewSession creates its own context and cannot target the old focus. + creation = cmd.type == "new_session" + if creation: + target = None + serial = not urgent and not ( + target and cmd.type in {"query", "get_context", "interrupt"}) + await self._command_scheduler.submit( + cmd, target=target, serial=serial, + barrier=not target and not creation and not urgent, + urgent=urgent, + ) + async def _process_command_safely(self, cmd) -> None: """Run one command without letting a handler failure stop the loop.""" try: @@ -35280,7 +35305,9 @@ async def reject( # engine must not evict a healthy resident Codex session and then fail. if engine == "claude" and broker_handle is None: try: - SdkHandle.preflight(self.cfg.claude_bin) + # The CLI version probe runs a blocking subprocess. Keep its + # bounded wait off the event loop used by every other session. + await asyncio.to_thread(SdkHandle.preflight, self.cfg.claude_bin) except Exception as exc: log.warning("Claude preflight failed; engine unavailable", error=str(exc)) @@ -35298,6 +35325,7 @@ async def reject( if k != self.focused_sid and c.state == "idle" and not c.btw and not c.queued_queries and not self._query_queue_task_active(c) + and not self._command_scheduler.owns_target(k) and not c.query_lock.locked() and not c.claude_control_persist_lock.locked() and (c.auto_compact_apply_task is None @@ -36431,7 +36459,7 @@ async def _spawn_btw( ERR_AUTH, "Codex Work 会话不属于当前账号,已拒绝打开 btw") if engine != "codex": try: - SdkHandle.preflight(self.cfg.claude_bin) + await asyncio.to_thread(SdkHandle.preflight, self.cfg.claude_bin) except Exception as exc: log.warning("Claude preflight failed for btw", error=str(exc)) raise _BtwSpawnFailure( @@ -36444,6 +36472,7 @@ async def _spawn_btw( if k != self.focused_sid and c.state == "idle" and not c.btw and not c.queued_queries and not self._query_queue_task_active(c) + and not self._command_scheduler.owns_target(k) and not c.query_lock.locked() and not c.claude_control_persist_lock.locked() and (c.auto_compact_apply_task is None @@ -36603,8 +36632,23 @@ async def _spawn_btw( ) from exc ctx.btw_reserved_id = ctx.sdk.fork_session_id = reserved_id try: - await ctx.sdk.connect( - resume_id=parent_id, cwd=parent.cwd, fork=True) + connect_kwargs = { + "resume_id": parent_id, "cwd": parent.cwd, "fork": True, + } + if engine == "claude" and parent_space == "code": + # Resumed/forked Claude children deliberately skip the optional + # startup context probe. Without an explicit launch selection, + # their model remains unknown until the first turn and BTW's + # picker stays on "loading" indefinitely. Apply the parent's + # current selection to the actual child, then publish that same + # value through OpenBtw's owner-only Model event. Work keeps its + # native policy; a transcript observation must not override it. + inherited_model = normalize_claude_model_selection( + _session_model(parent)) + if inherited_model: + ctx.sdk.model = inherited_model + connect_kwargs["model_override"] = inherited_model + await ctx.sdk.connect(**connect_kwargs) if engine == "codex": native_thread_id = getattr(ctx.sdk, "thread_id", None) if not isinstance(native_thread_id, str) or not native_thread_id: diff --git a/tests/test_claude_btw_identity.py b/tests/test_claude_btw_identity.py index 741521dc..6cf69b60 100644 --- a/tests/test_claude_btw_identity.py +++ b/tests/test_claude_btw_identity.py @@ -7,7 +7,7 @@ import pytest -from cc_remote.protocol import CloseBtw, ListSessions, SessionList +from cc_remote.protocol import CloseBtw, ListSessions, OpenBtw, SessionList from cc_remote.wrapper import machine as machine_module from cc_remote.wrapper.sdk import SdkHandle from tests.test_multisession import _mk_ctx, _mk_machine @@ -25,8 +25,10 @@ class Handle(SdkHandle): def preflight(_path): pass - async def connect(self, resume_id=None, cwd=None, fork=False): - self.launch = self._options(resume_id, cwd, fork=fork) + async def connect(self, resume_id=None, cwd=None, fork=False, + model_override=None): + self.launch = self._options( + resume_id, cwd, fork=fork, model_override=model_override) handles.append(self) # Model a native transcript visible as soon as connect starts, # before _spawn_btw has inserted the context into the pool. @@ -86,6 +88,36 @@ async def run(): asyncio.run(run()) +@pytest.mark.parametrize("from_announcement", [False, True]) +@pytest.mark.parametrize(("parent_model", "model"), [ + ("claude-opus-5", "claude-opus-5[1m]"), + ("claude-opus-4-6[1m]", "claude-opus-4-6[1m]"), + ("provider-model", "provider-model"), +]) +def test_btw_launches_with_parent_model_and_reports_it_before_first_query( + setup, from_announcement, parent_model, model): + async def run(): + machine, transport, parent, handles, _deleted = setup + if from_announcement: + parent.announced_model = parent_model + else: + parent.sdk.model = parent_model + await machine._handle_open_btw(OpenBtw( + sid=parent.key, client_id="owner", request_id="open")) + # Exercise real SDK option construction, not a UI-only fallback label. + # No model turn or extra control RPC is needed to report this selection. + assert handles[0].launch.model == model + assert handles[0].model == model + events = [event for event in transport.sent if event.type == "model"] + assert len(events) == 1 + assert events[0].model == model + assert events[0].owner_id == "owner" + assert events[0].sid.startswith("btw-") + assert machine.focused_sid == parent.key + + asyncio.run(run()) + + @pytest.mark.parametrize("captured", [False, True]) def test_close_cleans_reserved_identity_even_without_a_first_turn(setup, captured): async def run(): diff --git a/tests/test_claude_live_context.py b/tests/test_claude_live_context.py index 1b6fec12..42f3b59b 100644 --- a/tests/test_claude_live_context.py +++ b/tests/test_claude_live_context.py @@ -7,7 +7,7 @@ from claude_agent_sdk.types import AssistantMessage, SystemMessage, TextBlock from cc_remote.config import WrapperConfig -from cc_remote.protocol import ContextReport, Error, GetContext +from cc_remote.protocol import ContextReport, Error, GetContext, Interrupt from cc_remote.wrapper.sdk import SdkHandle from tests.test_claude_autocompact import SESSION_ID, _machine_with_sdk @@ -38,6 +38,50 @@ async def _send_control_request(self, request, timeout): return dict(SUMMARY) +def test_stop_reaches_sdk_while_live_context_read_is_pending(): + async def go(): + reading, release, interrupted = (asyncio.Event() for _ in range(3)) + + class SlowSummaryClient(SummaryClient): + async def _send_control_request(self, request, timeout): + reading.set() + await release.wait() + return await super()._send_control_request(request, timeout) + + async def interrupt(self): + interrupted.set() + + sdk = SdkHandle(WrapperConfig()) + sdk.client = SlowSummaryClient() + machine, transport, ctx = _machine_with_sdk(sdk) + ctx.state = "running" + try: + await machine._dispatch_incoming_command(GetContext( + sid=SESSION_ID, refresh=True, cmd_id="context", client_id="browser")) + await asyncio.wait_for(reading.wait(), 0.5) + await machine._dispatch_incoming_command(Interrupt( + sid=SESSION_ID, cmd_id="stop", client_id="browser")) + await asyncio.wait_for(interrupted.wait(), 0.5) + assert not release.is_set() + # Stop must still await the native terminal instead of inventing idle. + assert ctx.state == "interrupting" and ctx.interrupt_event.is_set() + release.set() + await asyncio.wait_for(machine._command_scheduler.drain(), 0.5) + report = next(frame for frame in transport.sent if isinstance(frame, ContextReport)) + assert report.total_tokens == SUMMARY["totalTokens"] + assert report.source == "control" + assert sdk.client.requests == [( + {"subtype": "get_context_usage", "detail": "summary"}, 5.0)] + assert {frame.cmd_id for frame in transport.sent if frame.type == "command_ack"} == { + "context", "stop"} + assert ctx.state == "interrupting" + finally: + release.set() + await machine._command_scheduler.close() + + asyncio.run(go()) + + def test_live_summary_yields_to_a_pending_sdk_control_operation(): async def go(): sdk = SdkHandle(WrapperConfig()) diff --git a/tests/test_claude_service_startup.py b/tests/test_claude_service_startup.py new file mode 100644 index 00000000..e28567fd --- /dev/null +++ b/tests/test_claude_service_startup.py @@ -0,0 +1,196 @@ +"""Slow native startup must isolate sessions without creating duplicate writers.""" + +import asyncio + +import pytest + +from cc_remote.claude_service.server import Service +from cc_remote.claude_service import server as service_module +from cc_remote.claude_service.wire import ControllerLeaseConflict +from tests.test_claude_service import FakeClient + + +def params(sid): + return {"metadata": {"session_id": sid, "space": "code", "cwd": f"/{sid}"}, + "options": {"cwd": f"/{sid}"}} + + +@pytest.mark.asyncio +async def test_slow_start_does_not_block_other_session_open_or_attach(tmp_path): + started, release = asyncio.Event(), asyncio.Event() + + class Client(FakeClient): + async def connect(self): + if self.options.cwd == "/slow": + started.set() + await release.wait() + + service = Service(tmp_path, factory=Client) + owner = object() + existing = await service.dispatch(owner, "existing", "open", params("existing")) + slow = asyncio.create_task(service.dispatch(owner, "slow", "open", params("slow"))) + try: + await asyncio.wait_for(started.wait(), 0.5) + attached = await asyncio.wait_for(service.dispatch( + owner, "attach", "open", params("existing")), 0.5) + fast = await asyncio.wait_for(service.dispatch(owner, "fast", "open", params("fast")), 0.5) + assert attached["attached"] and attached["id"] == existing["id"] + assert not fast["attached"] and not slow.done() + await service.dispatch(owner, "query", "query", { + "session": existing["id"], "prompt": "accepted once", "turn": {"id": "turn"}}) + assert service.sessions[existing["id"]].client.prompts == ["accepted once"] + finally: + release.set() + await asyncio.gather(slow, return_exceptions=True) + for session in service.sessions.values(): + await session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("other_owner", [False, True]) +async def test_concurrent_same_identity_starts_once_and_preserves_lease(tmp_path, other_owner): + started, release = asyncio.Event(), asyncio.Event() + clients = [] + + class Client(FakeClient): + async def connect(self): + clients.append(self) + started.set() + await release.wait() + + service = Service(tmp_path, factory=Client) + owner = object() + first = asyncio.create_task(service.dispatch(owner, "one", "open", params("same"))) + await asyncio.wait_for(started.wait(), 0.5) + second = asyncio.create_task(service.dispatch( + object() if other_owner else owner, "two", "open", params("same"))) + try: + await asyncio.sleep(0) + assert not second.done() and len(clients) == 1 + release.set() + opened = await asyncio.wait_for(first, 0.5) + if other_owner: + with pytest.raises(ControllerLeaseConflict): + await asyncio.wait_for(second, 0.5) + else: + attached = await asyncio.wait_for(second, 0.5) + assert attached["id"] == opened["id"] and attached["attached"] + assert len(clients) == 1 + assert service.sessions[opened["id"]].controller is owner + finally: + release.set() + await asyncio.gather(first, second, return_exceptions=True) + for session in service.sessions.values(): + await session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ["error", "cancel"]) +async def test_failed_start_keeps_identity_reserved_until_cleanup_and_wakes_joiners(tmp_path, failure): + started, release_start, closing, release_close = (asyncio.Event() for _ in range(4)) + clients = [] + + class Client(FakeClient): + async def connect(self): + clients.append(self) + if self.options.cwd == "/broken": + started.set() + await release_start.wait() + raise RuntimeError("deliberate startup failure") + + async def disconnect(self): + if self.options.cwd == "/broken": + closing.set() + await release_close.wait() + await super().disconnect() + + service = Service(tmp_path, factory=Client) + owner = object() + first = asyncio.create_task(service.dispatch(owner, "one", "open", params("broken"))) + await asyncio.wait_for(started.wait(), 0.5) + if failure == "cancel": + first.cancel() + else: + release_start.set() + await asyncio.wait_for(closing.wait(), 0.5) + joining = asyncio.create_task(service.dispatch(owner, "two", "open", params("broken"))) + try: + fast = await asyncio.wait_for(service.dispatch(owner, "fast", "open", params("fast")), 0.5) + await asyncio.sleep(0) + assert not joining.done() + assert len(clients) == 2 and len(service.sessions) == 2 + release_close.set() + with pytest.raises(asyncio.CancelledError if failure == "cancel" else RuntimeError): + await first + with pytest.raises(RuntimeError, match="no longer available"): + await asyncio.wait_for(joining, 0.5) + assert list(service.sessions) == [fast["id"]] + assert len(clients) == 2 and clients[0].closed + assert {path.stem for path in tmp_path.glob("*.sqlite3")} == {fast["id"]} + assert not service._opening + finally: + release_start.set() + release_close.set() + await asyncio.gather(first, joining, return_exceptions=True) + for session in service.sessions.values(): + await session.close() + + +@pytest.mark.asyncio +async def test_cancelling_startup_joiner_does_not_cancel_creator(tmp_path): + started, release = asyncio.Event(), asyncio.Event() + + class Client(FakeClient): + async def connect(self): + started.set() + await release.wait() + + service = Service(tmp_path, factory=Client) + owner = object() + first = asyncio.create_task(service.dispatch(owner, "one", "open", params("same"))) + await asyncio.wait_for(started.wait(), 0.5) + joining = asyncio.create_task(service.dispatch(owner, "two", "open", params("same"))) + try: + await asyncio.sleep(0) + joining.cancel() + with pytest.raises(asyncio.CancelledError): + await joining + assert not first.done() and len(service.sessions) == 1 + release.set() + opened = await asyncio.wait_for(first, 0.5) + session = service.sessions[opened["id"]] + assert session.controller is owner and not session.client.closed + finally: + release.set() + await asyncio.gather(first, joining, return_exceptions=True) + for session in service.sessions.values(): + await session.close() + + +@pytest.mark.asyncio +async def test_starting_session_counts_toward_capacity_and_checks_identity(tmp_path, monkeypatch): + monkeypatch.setattr(service_module, "MAX_SESSIONS", 1) + started, release = asyncio.Event(), asyncio.Event() + + class Client(FakeClient): + async def connect(self): + started.set() + await release.wait() + + service = Service(tmp_path, factory=Client) + owner = object() + first = asyncio.create_task(service.dispatch(owner, "one", "open", params("same"))) + try: + await asyncio.wait_for(started.wait(), 0.5) + with pytest.raises(RuntimeError, match="capacity"): + await asyncio.wait_for(service.dispatch(owner, "two", "open", params("other")), 0.5) + mismatched = params("same") + mismatched["metadata"]["cwd"] = "/different" + with pytest.raises(PermissionError, match="identity mismatch"): + await asyncio.wait_for(service.dispatch(owner, "bad", "open", mismatched), 0.5) + assert not first.done() and len(service.sessions) == 1 + finally: + release.set() + await asyncio.gather(first, return_exceptions=True) + for session in service.sessions.values(): + await session.close() diff --git a/tests/test_command_scheduling.py b/tests/test_command_scheduling.py new file mode 100644 index 00000000..7da66624 --- /dev/null +++ b/tests/test_command_scheduling.py @@ -0,0 +1,402 @@ +"""Command intake stays responsive without reordering a session's lifecycle.""" +from __future__ import annotations + +import asyncio +import threading +from types import SimpleNamespace + +import pytest + +from cc_remote.protocol import ( + AnswerQuestion, GetContext, Interrupt, NewSession, OpenBtw, Ping, Query, SetEffort, + SwitchSession, serialize, +) +from cc_remote.wrapper.command_scheduler import CommandScheduler +from cc_remote.wrapper.sdk import SdkHandle +from tests.test_multisession import _mk_ctx, _mk_machine + + +def _query(sid="b", cmd_id="send"): + return Query(sid=sid, prompt="private prompt", msg_id=f"msg-{cmd_id}", + client_id="browser", cmd_id=cmd_id) + + +@pytest.mark.parametrize("slow_command", [ + GetContext(sid="a", refresh=True), + SwitchSession(session_id="a"), + NewSession(), + _query("a", "slow"), +]) +def test_slow_session_does_not_block_another_sessions_query_stop_or_ping( + slow_command, monkeypatch): + async def run(): + machine, transport = _mk_machine() + a, b = _mk_ctx("a"), _mk_ctx("b") + machine.sessions = {"a": a, "b": b} + machine.focused_sid = "b" + started, release, stopped = (asyncio.Event() for _ in range(3)) + calls = [] + + async def slow(_cmd): + started.set() + await release.wait() + + async def query(cmd): + if cmd.sid == "a": + await slow(cmd) + else: + calls.append("query") + b.state = "running" + + async def native_interrupt(): + assert b.state == "interrupting" + assert b.interrupt_event.is_set() + calls.append("interrupt") + stopped.set() + + b.sdk = SimpleNamespace(interrupt=native_interrupt) + monkeypatch.setattr(machine, "_handle_query", query) + if slow_command.type != "query": + monkeypatch.setattr(machine, f"_handle_{slow_command.type}", slow) + try: + await machine._dispatch_incoming_command(slow_command) + await asyncio.wait_for(started.wait(), 1) + await machine._dispatch_incoming_command(_query()) + await machine._dispatch_incoming_command(Interrupt(sid="b")) + await machine._dispatch_incoming_command(Ping(n=7)) + await asyncio.wait_for(stopped.wait(), 1) + assert calls == ["query", "interrupt"] + assert not release.is_set() + assert any(frame.type == "pong" and frame.n == 7 + for frame in transport.sent) + assert any(frame.type == "command_ack" and frame.cmd_id == "send" + for frame in transport.sent) + finally: + release.set() + await machine._command_scheduler.close() + + asyncio.run(run()) + + +@pytest.mark.parametrize("command", [ + NewSession(), + OpenBtw(sid="a", request_id="open", client_id="browser"), +]) +def test_cli_preflight_does_not_block_another_sessions_stop_or_heartbeat( + command, monkeypatch): + async def run(): + machine, transport = _mk_machine() + a, b = _mk_ctx("a", "a"), _mk_ctx("b", "b") + machine.sessions = {"a": a, "b": b} + machine.focused_sid = "b" + b.state = "running" + started, stopped = asyncio.Event(), asyncio.Event() + release, finished = threading.Event(), threading.Event() + loop = asyncio.get_running_loop() + + def preflight(_path): + loop.call_soon_threadsafe(started.set) + try: + # The timeout makes the old synchronous path fail without + # hanging pytest. A responsive event loop releases us first. + release.wait(2) + raise RuntimeError("deliberate preflight failure") + finally: + finished.set() + + async def native_interrupt(): + stopped.set() + + b.sdk = SimpleNamespace(interrupt=native_interrupt) + monkeypatch.setattr(SdkHandle, "preflight", preflight) + try: + await machine._dispatch_incoming_command(command) + await asyncio.wait_for(started.wait(), 3) + assert not finished.is_set(), "CLI preflight blocked the event loop" + await machine._dispatch_incoming_command(Interrupt(sid="b")) + await machine._dispatch_incoming_command(Ping(n=8)) + await asyncio.wait_for(stopped.wait(), 1) + assert not finished.is_set() + assert any(frame.type == "pong" and frame.n == 8 + for frame in transport.sent) + finally: + release.set() + await asyncio.wait_for(machine._command_scheduler.drain(), 3) + await machine._command_scheduler.close() + # The version gate still rejects startup; no SDK connect/model turn is + # allowed merely because its blocking probe moved off the event loop. + assert finished.is_set() + assert set(machine.sessions) == {"a", "b"} + assert any(frame.type == "error" and frame.code == "cc_crash" + for frame in transport.sent) + + asyncio.run(run()) + + +def test_cold_resume_query_and_stop_preserve_same_session_order(monkeypatch): + async def run(): + machine, _ = _mk_machine() + started, release = asyncio.Event(), asyncio.Event() + calls = [] + + async def switch(cmd): + calls.append("resuming") + started.set() + await release.wait() + machine.sessions[cmd.session_id] = _mk_ctx(cmd.session_id) + machine.focused_sid = cmd.session_id + calls.append("resumed") + + async def query(cmd): + assert machine._ctx_for(cmd.sid) is not None + calls.append("query") + + async def interrupt(_cmd): + calls.append("interrupt") + + monkeypatch.setattr(machine, "_handle_switch_session", switch) + monkeypatch.setattr(machine, "_handle_query", query) + monkeypatch.setattr(machine, "_handle_interrupt", interrupt) + await machine._dispatch_incoming_command(SwitchSession(session_id="cold")) + await started.wait() + await machine._dispatch_incoming_command(_query("cold")) + await machine._dispatch_incoming_command(Interrupt(sid="cold")) + await asyncio.sleep(0) + assert calls == ["resuming"] + release.set() + await asyncio.wait_for(machine._command_scheduler.drain(), 1) + assert calls == ["resuming", "resumed", "query", "interrupt"] + + asyncio.run(run()) + + +def test_legacy_query_resolves_focus_after_switch_and_mutations_stay_ordered( + monkeypatch): + async def run(): + machine, _ = _mk_machine() + machine.sessions = {key: _mk_ctx(key) for key in ("a", "b")} + machine.focused_sid = "a" + started, release = asyncio.Event(), asyncio.Event() + calls = [] + + async def switch(cmd): + started.set() + await release.wait() + machine.focused_sid = cmd.session_id + calls.append(cmd.session_id) + + async def query(cmd): + calls.append(f"query:{machine._ctx_for(cmd.sid).key}") + + monkeypatch.setattr(machine, "_handle_switch_session", switch) + monkeypatch.setattr(machine, "_handle_query", query) + await machine._dispatch_incoming_command(SwitchSession(session_id="b")) + await started.wait() + await machine._dispatch_incoming_command(_query(None)) + await machine._dispatch_incoming_command(SwitchSession(session_id="a")) + await asyncio.sleep(0) + assert calls == [] + release.set() + await asyncio.wait_for(machine._command_scheduler.drain(), 1) + assert calls == ["b", "query:b", "a"] + + asyncio.run(run()) + + +def test_settings_and_query_remain_ordered_on_their_session(monkeypatch): + async def run(): + machine, _ = _mk_machine() + started, release = asyncio.Event(), asyncio.Event() + calls = [] + + async def effort(cmd): + started.set() + await release.wait() + calls.append(cmd.effort) + + async def query(_cmd): + calls.append("query") + + monkeypatch.setattr(machine, "_handle_set_effort", effort) + monkeypatch.setattr(machine, "_handle_query", query) + await machine._dispatch_incoming_command(SetEffort(sid="b", effort="high")) + await started.wait() + await machine._dispatch_incoming_command(_query()) + await asyncio.sleep(0) + assert calls == [] + release.set() + await machine._command_scheduler.drain() + assert calls == ["high", "query"] + + asyncio.run(run()) + + +def test_question_answer_bypasses_the_handler_waiting_for_it(monkeypatch): + async def run(): + machine, _ = _mk_machine() + asked, answered = asyncio.Event(), asyncio.Event() + + async def effort(_cmd): + asked.set() + await answered.wait() + + async def answer(_cmd): + answered.set() + + monkeypatch.setattr(machine, "_handle_set_effort", effort) + monkeypatch.setattr(machine, "_handle_answer_question", answer) + await machine._dispatch_incoming_command(SetEffort(sid="b", effort="high")) + await asked.wait() + await machine._dispatch_incoming_command(AnswerQuestion( + sid="b", ask_id="permission", answer="allow")) + await asyncio.wait_for(machine._command_scheduler.drain(), 1) + assert answered.is_set() + + asyncio.run(run()) + + +def test_waiting_query_is_not_evicted_before_it_can_take_query_lock(monkeypatch): + async def run(): + machine, transport = _mk_machine() + machine.cfg.max_concurrent_sessions = 2 + a, b = _mk_ctx("a"), _mk_ctx("b") + machine.sessions = {"a": a, "b": b} + machine.focused_sid = "a" + monkeypatch.setattr(SdkHandle, "preflight", lambda _path: None) + + async def pending_query(_cmd): + await asyncio.Event().wait() + + monkeypatch.setattr(machine, "_handle_query", pending_query) + await machine._dispatch_incoming_command(_query()) + assert not b.query_lock.locked() + try: + assert await machine._spawn(resume_id=None) is None + assert machine.sessions.get("b") is b + assert any(frame.type == "error" and frame.code == "busy" + for frame in transport.sent) + finally: + await machine._command_scheduler.close() + + asyncio.run(run()) + + +def test_reconnect_retry_coalesces_inflight_and_only_acks_after_acceptance( + monkeypatch): + async def run(): + machine, transport = _mk_machine() + started, release = asyncio.Event(), asyncio.Event() + calls = [] + + async def query(cmd): + calls.append(cmd.msg_id) + started.set() + await release.wait() + + monkeypatch.setattr(machine, "_handle_query", query) + cmd = _query() + await machine._dispatch_incoming_command(cmd) + await started.wait() + await machine._dispatch_incoming_command(cmd.model_copy()) + assert calls == [cmd.msg_id] + assert not transport.sent + release.set() + await machine._command_scheduler.drain() + await machine._dispatch_incoming_command(cmd.model_copy()) + await machine._command_scheduler.drain() + assert calls == [cmd.msg_id] + assert [frame.type for frame in transport.sent] == [ + "command_ack", "command_ack"] + + asyncio.run(run()) + + +def test_rekey_does_not_split_the_command_order_or_residency_pin(monkeypatch): + async def run(): + machine, _ = _mk_machine() + ctx = _mk_ctx("tmp-a") + machine.sessions[ctx.key] = ctx + machine.focused_sid = ctx.key + started, release = asyncio.Event(), asyncio.Event() + calls = [] + + async def query(cmd): + calls.append(cmd.cmd_id) + if cmd.cmd_id == "first": + started.set() + await release.wait() + + monkeypatch.setattr(machine, "_handle_query", query) + await machine._dispatch_incoming_command(_query("tmp-a", "first")) + await started.wait() + await machine._capture_session_id(ctx, "real-a") + assert machine._command_scheduler.owns_target("real-a") + await machine._dispatch_incoming_command(_query("real-a", "second")) + await asyncio.sleep(0) + assert calls == ["first"] + release.set() + await machine._command_scheduler.drain() + assert calls == ["first", "second"] + assert not machine._command_scheduler.owns_target("real-a") + + asyncio.run(run()) + + +@pytest.mark.parametrize("bound", ["items", "bytes"]) +def test_scheduler_applies_backpressure_and_releases_capacity(bound): + async def run(): + started, release = asyncio.Event(), asyncio.Event() + calls = [] + + async def process(cmd): + calls.append(cmd.cmd_id) + if cmd.cmd_id == "first": + started.set() + await release.wait() + + first = _query("a", "first") + scheduler = CommandScheduler( + process, lambda sid: sid, + max_items=1 if bound == "items" else 128, + max_bytes=(len(serialize(first).encode("utf-8")) + if bound == "bytes" else 32 * 1024 * 1024), + ) + await scheduler.submit(first) + await started.wait() + blocked = asyncio.create_task(scheduler.submit(_query("b", "second"))) + await asyncio.sleep(0) + assert not blocked.done() + assert calls == ["first"] + release.set() + await asyncio.wait_for(blocked, 1) + await scheduler.drain() + assert calls == ["first", "second"] + assert scheduler._bytes == 0 + + asyncio.run(run()) + + +def test_shutdown_cancels_waiting_commands_without_launching_them(): + async def run(): + started, cancelled = asyncio.Event(), asyncio.Event() + calls = [] + + async def process(cmd): + calls.append(cmd.cmd_id) + started.set() + try: + await asyncio.Event().wait() + finally: + cancelled.set() + + scheduler = CommandScheduler(process, lambda sid: sid) + await scheduler.submit(_query("a", "first"), target="a") + await started.wait() + await scheduler.submit(_query("a", "second"), target="a") + await scheduler.close() + assert cancelled.is_set() + assert calls == ["first"] + assert not scheduler.owns_target("a") + assert scheduler._bytes == 0 + + asyncio.run(run()) diff --git a/tests/test_relay_forward.py b/tests/test_relay_forward.py index 259e757e..50598d07 100644 --- a/tests/test_relay_forward.py +++ b/tests/test_relay_forward.py @@ -9,6 +9,7 @@ ) from cc_remote.relay.forward import ClientConn, SlowClientError from cc_remote.relay import pairing +from cc_remote.relay import forward from cc_remote.relay.pairing import RelayHub @@ -500,3 +501,122 @@ async def run(): assert "ephemeral-machine" not in hub._machine_clients asyncio.run(run()) + + +@pytest.mark.parametrize("slow_machine", ["default", "slow"]) +def test_blocked_wrapper_does_not_block_another_machines_client_hello(slow_machine): + async def run(): + hub = RelayHub(SimpleNamespace(client_queue_cap=4)) + slow, fast = FakeWs(), ScriptedWs() + if slow_machine == "default": + hub._wrapper_ws = slow + else: + hub._wrappers[slow_machine] = slow + hub._wrappers["fast"] = fast + conn = ClientConn(ScriptedWs(), 4, "phone") + hub._ensure_clients_for(slow_machine)["phone"] = conn + sending = asyncio.create_task(hub._forward_client_msg( + conn, "phone", Query(prompt="one", msg_id="one"), slow_machine)) + await asyncio.sleep(0) + ws = ScriptedWs() + await ws.incoming.put(serialize(Hello(role="client", client_id="phone"))) + serving = asyncio.create_task(hub.serve_client(ws, "fast")) + try: + async with asyncio.timeout(0.5): + while not fast.sent: + await asyncio.sleep(0.001) + assert json.loads(fast.sent[0])["type"] == "hello" + assert not sending.done() + assert hub._clients_for(slow_machine)["phone"] is conn + finally: + slow.block.set() + serving.cancel() + await asyncio.gather(sending, serving, return_exceptions=True) + await conn.stop() + + asyncio.run(run()) + + +@pytest.mark.parametrize("route", ["broadcast", "owner", "direct"]) +def test_stalled_client_close_does_not_block_following_frames(route): + class SlowCloseWs(FakeWs): + def __init__(self): + super().__init__() + self.closing = asyncio.Event() + self.release_close = asyncio.Event() + + async def close(self, **kwargs): + self.closing.set() + await self.release_close.wait() + await super().close(**kwargs) + + async def run(): + hub = RelayHub(SimpleNamespace()) + slow_ws, fast_ws = SlowCloseWs(), ScriptedWs() + slow = ClientConn(slow_ws, 1, "slow", owner_id="owner") + fast = ClientConn(fast_ws, 8, "fast", owner_id="owner") + slow.start() + fast.start() + hub._clients = {"slow": slow, "fast": fast} + await slow.send(Delta(message_id="old", text="in flight")) + await asyncio.sleep(0) + await slow.send(Delta(message_id="old", text="queued")) + routing = ({"owner_id": "owner"} if route == "owner" else + {"to": "slow"} if route == "direct" else {}) + + async def incoming(): + await hub._on_wrapper_msg(Delta(message_id="one", text="first", **routing)) + await hub._on_wrapper_msg(Delta(message_id="two", text="next")) + + task = asyncio.create_task(incoming()) + try: + await asyncio.wait_for(slow_ws.closing.wait(), 0.5) + await asyncio.wait_for(asyncio.shield(task), 0.5) + await asyncio.sleep(0) + assert "slow" not in hub._clients and slow.closed + assert slow._sender.done() + assert json.loads(fast_ws.sent[-1])["message_id"] == "two" + finally: + slow_ws.release_close.set() + await asyncio.gather(task, return_exceptions=True) + await slow.stop() + await fast.stop() + + asyncio.run(run()) + + +def test_client_close_is_bounded_and_survives_reader_cancellation(monkeypatch): + monkeypatch.setattr(forward, "CLIENT_CLOSE_TIMEOUT", 0.02) + + class StuckCloseWs(FakeWs): + def __init__(self): + super().__init__() + self.closing = asyncio.Event() + self.close_cancelled = asyncio.Event() + + async def close(self, **kwargs): + self.closed.append(kwargs) + self.closing.set() + try: + await asyncio.Event().wait() + finally: + self.close_cancelled.set() + + async def run(): + ws = StuckCloseWs() + conn = ClientConn(ws, 1, "stuck") + conn.start() + await conn.send(Delta(message_id="m", text="one")) + await asyncio.sleep(0) + stopping = asyncio.create_task(conn.stop(code=4008, reason="slow client")) + await asyncio.wait_for(ws.closing.wait(), 0.5) + stopping.cancel() + with pytest.raises(asyncio.CancelledError): + await stopping + await asyncio.wait_for(conn.stop(), 0.5) + assert ws.close_cancelled.is_set() and len(ws.closed) == 1 + assert conn._sender.done() and conn.closed + with pytest.raises(ConnectionError): + await conn.send(Delta(message_id="m", text="two")) + + asyncio.run(run()) diff --git a/web/src/components/BtwPanel.tsx b/web/src/components/BtwPanel.tsx index 25fb45df..f6af533a 100644 --- a/web/src/components/BtwPanel.tsx +++ b/web/src/components/BtwPanel.tsx @@ -624,7 +624,8 @@ export function BtwPanel(p: Props) {
BTW 设置 + disabled={inputLocked || busy || !p.sid}>{model?.name + ?? (p.sid && p.rt ? "选择模型" : "模型读取中")} {p.engine === "codex" &&