diff --git a/cc_remote/claude_service/server.py b/cc_remote/claude_service/server.py index 27904b6..eabab75 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/codex_app_tools.py b/cc_remote/codex_app_tools.py index 5743ea3..fe216d0 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 caae0f5..97771f4 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/relay/forward.py b/cc_remote/relay/forward.py index 14806ea..c6a9aa2 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 cdf4010..df66769 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 0000000..0d8e6a9 --- /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 dfbb15e..56e9229 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 @@ -6940,6 +6943,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 +6981,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( @@ -9026,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() @@ -11460,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: @@ -18524,9 +18558,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 @@ -35272,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)) @@ -35290,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 @@ -36423,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( @@ -36436,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 @@ -36595,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: @@ -39081,6 +39133,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 122a6c4..e2e5c45 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_claude_btw_identity.py b/tests/test_claude_btw_identity.py index 741521d..6cf69b6 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 1b6fec1..42f3b59 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 0000000..e28567f --- /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_codex_app_tools.py b/tests/test_codex_app_tools.py index 2e5d680..8fca659 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 a02b8d3..b4719c7 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_command_scheduling.py b/tests/test_command_scheduling.py new file mode 100644 index 0000000..7da6662 --- /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 259e757..50598d0 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/tests/test_wrapper_core_fixes.py b/tests/test_wrapper_core_fixes.py index 94576dc..0a59b03 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/components/AgentDetailPanel.tsx b/web/src/components/AgentDetailPanel.tsx index f578173..aa3531a 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 && (
这个协作代理暂时没有可展示的过程。
)} diff --git a/web/src/components/BtwPanel.tsx b/web/src/components/BtwPanel.tsx index 25fb45d..f6af533 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" &&