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
1 change: 1 addition & 0 deletions src/entrypoints/serve_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -247,6 +247,7 @@ async def _serve(args, workspace: str, token: str,
workspace=workspace,
manager=manager,
spawn_agent=spawn,
agent_config=agent_config,
protocol_version=PROTOCOL_VERSION,
)
app = build_app(state)
Expand Down
31 changes: 25 additions & 6 deletions src/server/desktop_gateway_methods.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,10 +85,11 @@ def __init__(self, session_id: str, state: DesktopServeState) -> None:

# ── lifecycle ────────────────────────────────────────────────────────────

async def start(self, cwd: str) -> None:
async def start(self, cwd: str, spawn: Any = None) -> None:
# spawn's third arg is a permission-mode override, NOT a resume id;
# resuming a stored session is a post-init `resume` control request.
self.agent = await self.state.spawn_agent(self.session_id, cwd, None)
spawn = spawn or self.state.spawn_agent
self.agent = await spawn(self.session_id, cwd, None)
self.pump_task = asyncio.create_task(
self._pump(), name=f"desktop-session-{self.session_id}"
)
Expand Down Expand Up @@ -364,6 +365,13 @@ def _catalog_from_config() -> dict[str, Any]:
}


def _clean(value: Any) -> str | None:
"""A non-empty trimmed string, or None (empty selections → base config)."""
if isinstance(value, str) and value.strip():
return value.strip()
return None


def _suggestion_scope(suggestion: dict[str, Any]) -> str:
"""Bucket a permission suggestion as a session or always/persistent grant."""
destination = str(suggestion.get("destination") or "").lower()
Expand Down Expand Up @@ -424,7 +432,8 @@ def _session(self, params: dict[str, Any]) -> DesktopSession:
raise ValueError(f"unknown session: {session_id or '<missing>'}")
return session

async def _create(self, cwd: str | None, resume: str | None) -> DesktopSession:
async def _create(self, cwd: str | None, resume: str | None,
params: dict[str, Any] | None = None) -> DesktopSession:
manager = self.state.manager
workspace = cwd or self.state.workspace
if resume and resume in self.state.sessions:
Expand All @@ -436,8 +445,18 @@ async def _create(self, cwd: str | None, resume: str | None) -> DesktopSession:
session = DesktopSession(session_id, self.state)
session.sockets.add(self.websocket)
self.state.sessions[session_id] = session
# Honor the composer's provider/model/effort selection at spawn time, so
# a session can use a working provider even when the config default is
# broken (expired subscription, missing key). Empty values → the base
# config's default provider.
params = params or {}
spawn = self.state.spawn_for(
_clean(params.get("provider")),
_clean(params.get("model")),
_clean(params.get("reasoning_effort") or params.get("effort")),
)
try:
await session.start(workspace)
await session.start(workspace, spawn=spawn)
except Exception:
self.state.sessions.pop(session_id, None)
raise
Expand All @@ -460,7 +479,7 @@ async def _create(self, cwd: str | None, resume: str | None) -> DesktopSession:
# ── methods ──────────────────────────────────────────────────────────────

async def session_create(self, params: dict[str, Any]) -> dict[str, Any]:
session = await self._create(params.get("cwd"), None)
session = await self._create(params.get("cwd"), None, params)
return {
"session_id": session.session_id,
"stored_session_id": session.session_id,
Expand All @@ -469,7 +488,7 @@ async def session_create(self, params: dict[str, Any]) -> dict[str, Any]:

async def session_resume(self, params: dict[str, Any]) -> dict[str, Any]:
wanted = str(params.get("session_id") or "") or None
session = await self._create(params.get("cwd"), wanted)
session = await self._create(params.get("cwd"), wanted, params)
response: dict[str, Any] = {
"session_id": session.session_id,
"stored_session_id": wanted or session.session_id,
Expand Down
28 changes: 28 additions & 0 deletions src/server/desktop_serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,11 +45,39 @@ class DesktopServeState:
manager: Any
spawn_agent: Callable[..., Awaitable[Any]]
protocol_version: str
# The base AgentServerConfig every session inherits. A session.create that
# names a provider/model/effort (the composer's selection) is spawned from
# a per-session copy of this — otherwise every session would boot the
# default provider and a bad default (e.g. an expired Claude subscription)
# would make the app unusable even after switching models.
agent_config: Any = None
# session_id -> live DesktopSession (created lazily by the gateway).
sessions: dict[str, Any] = field(default_factory=dict)
# Saved-transcript dir override (tests); default resolves per request.
sessions_dir: Path | None = None

def spawn_for(self, provider: str | None, model: str | None,
effort: str | None) -> Callable[..., Awaitable[Any]]:
"""A spawn closure honoring a session's provider/model/effort override.

Falls back to the shared ``spawn_agent`` when nothing is overridden or
no base config is available (tests inject a bare spawn).
"""
if self.agent_config is None or not (provider or model or effort):
return self.spawn_agent
import dataclasses

from src.server.agent_server import make_spawn_agent

overrides: dict[str, Any] = {}
if provider:
overrides["provider_name"] = provider
if model:
overrides["model"] = model
if effort:
overrides["effort"] = effort
return make_spawn_agent(dataclasses.replace(self.agent_config, **overrides))

def saved_sessions_dir(self) -> Path:
if self.sessions_dir is not None:
return self.sessions_dir
Expand Down
50 changes: 50 additions & 0 deletions tests/server/test_desktop_gateway.py
Original file line number Diff line number Diff line change
Expand Up @@ -280,6 +280,56 @@ def reply_frames():
assert reply["response"]["updatedInput"] == {"command": "rm -rf /tmp/x"}


def test_session_create_honors_provider_override(tmp_path: Path) -> None:
"""A composer selection (provider/model) must reach the spawn, so a session
can use a working provider even when the config default is broken."""
import dataclasses

from src.server.agent_server import AgentServerConfig

spawned_with: list = []

def make_spawn(config):
async def spawn(session_id, cwd, resume):
spawned_with.append((config.provider_name, config.model, config.effort))
agent = FakeAgent()
return agent
return spawn

base = AgentServerConfig(provider_name="anthropic", model="claude-sonnet-4-6")
state, _ = _fake_state(tmp_path)
state.agent_config = base
state.spawn_agent = make_spawn(base)
# Real make_spawn_agent is heavy; substitute ours for the override path.
import src.server.desktop_serve as serve_mod
import src.server.agent_server as agent_mod
orig = agent_mod.make_spawn_agent
agent_mod.make_spawn_agent = make_spawn
try:
with TestClient(build_app(state)) as client, _connect(client) as ws:
ws.receive_json()
events: list[dict] = []
_rpc(ws, 1, "session.create",
{"cwd": "/tmp", "source": "desktop", "provider": "deepseek",
"model": "deepseek-v4-flash", "reasoning_effort": "high"})
_drain_for_response(ws, 1, events)
finally:
agent_mod.make_spawn_agent = orig

assert spawned_with[-1] == ("deepseek", "deepseek-v4-flash", "high")


def test_session_create_without_override_uses_base_spawn(tmp_path: Path) -> None:
state, agents = _fake_state(tmp_path)
# No agent_config → spawn_for returns the shared spawn_agent unchanged.
with TestClient(build_app(state)) as client, _connect(client) as ws:
ws.receive_json()
events: list[dict] = []
_rpc(ws, 1, "session.create", {"cwd": "/tmp", "source": "desktop"})
_drain_for_response(ws, 1, events)
assert len(agents) == 1 # the base spawn was used


def test_resume_hydrates_saved_transcript(tmp_path: Path) -> None:
state, agents = _fake_state(tmp_path)
sessions_dir = tmp_path / "saved"
Expand Down
Loading