Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
94 changes: 63 additions & 31 deletions cc_remote/claude_service/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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")
Expand Down
7 changes: 3 additions & 4 deletions cc_remote/codex_app_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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


Expand Down Expand Up @@ -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"]
Expand Down
5 changes: 2 additions & 3 deletions cc_remote/codex_desktop.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
33 changes: 22 additions & 11 deletions cc_remote/relay/forward.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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
Expand All @@ -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):
Expand Down Expand Up @@ -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")
23 changes: 18 additions & 5 deletions cc_remote/relay/pairing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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)

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Loading
Loading