diff --git a/docs/architecture/session-isolation-analysis.md b/docs/architecture/session-isolation-analysis.md new file mode 100644 index 000000000..5edae895a --- /dev/null +++ b/docs/architecture/session-isolation-analysis.md @@ -0,0 +1,493 @@ +# 多窗口、多 Agent 与子 Agent 执行下的会话隔离分析 + +## 概述 + +本文分析当前 OpenCode + AgentPool 多 Agent、多窗口架构下,上下文是否会发生混乱,并说明当前的会话隔离、父子会话隔离以及运行时状态隔离是如何工作的。 + +这是一份**分析文档**,不是决策文档。对应的架构提案与迁移方案单独记录在: + +- `docs/rfcs/draft/RFC-0023-session-runtime-hard-isolation.md` + +--- + +## 执行摘要 + +当前系统更准确的描述是:**共享进程中的逻辑隔离执行模型**。 + +在大多数常规场景下,上下文通常不会混乱,因为实现中持续依赖以下机制: + +- 以 `session_id` 作为主隔离键 +- 以 `parent_id` / `parent_session_id` 表示父子链路 +- 每个 session 独立的 turn 锁 +- 每个 session 独立的 `MessageHistory` +- 每个 session 独立的 input provider +- 不可变的 `RunSnapshot` +- 面向子 agent 的事件包装与转发 + +但这种隔离还不是完全结构性的硬隔离。部分可变状态仍然位于 agent 实例级或进程级,因此系统正确性依赖于多个机制协同成立,而不是依赖一个天然的强所有权边界。 + +当前最大的未解决风险是:**同一个底层 agent 实例被多个 session 并发复用**。 + +--- + +## 分析范围 + +本文评估以下边界上的隔离情况: + +1. 会话与会话之间的隔离 +2. 多窗口之间的隔离 +3. 主 agent 与子 agent 之间的隔离 +4. agent 与 agent 之间的隔离 +5. 后台 worker / task 的隔离 +6. 传输层事件的隔离 +7. 运行时状态的归属边界 + +--- + +## 当前隔离模型 + +## 1. 会话身份与持久化 + +当前的 session 模型是显式且可持久化的。 + +关键观察: + +- `SessionData` 持有 `session_id`、`agent_name`、`pool_id`、`project_id`、`parent_id`、`cwd`、`agent_type` 以及 metadata。 +- session store 以 `session_id` 为核心索引或查询条件,也可以附加 `pool_id`、`agent_name`、`parent_id` 等过滤条件。 +- 子会话链路通过 `parent_id` 被显式持久化。 + +### 评估 + +这一层相对稳固。持久化模型本身看起来并不是上下文混乱的主要来源。这里的 session 是显式身份,而不是推导出来的隐式身份。 + +--- + +## 2. 服务端运行时隔离 + +OpenCode server 使用一个共享的 `ServerState`,但大多数可变状态桶都按 `session_id` 分区。 + +典型例子包括: + +- `sessions[session_id]` +- `messages[session_id]` +- `todos[session_id]` +- `session_conversations[session_id]` +- `input_providers[session_id]` +- `session_locks[session_id]` +- `pending_async_prompts[session_id]` + +### 评估 + +这意味着运行时是**共享的,但按 session 切分管理**。它不是每个窗口一个独立 server 对象,也不是每个 session 一个独立 server 实例。隔离是通过同一进程内的键控状态容器实现的。 + +这种方式在组织与路由上是有效的,但它仍然属于逻辑隔离,而不是独立进程级隔离。 + +--- + +## 3. 同一 Session 内的并发保护 + +同一个 session 的 turn 通过 `get_session_lock(session_id)` 串行化。 + +在 `message_routes.py` 中,服务端会: + +1. 先创建 user message +2. 再把它追加到 `state.messages[session_id]` +3. 获取该 session 的独立锁 +4. 在这把锁内完成剩余 turn 处理 + +### 评估 + +这是当前设计里最强的一层保护之一。 + +它能避免: + +- 同一 session 内的消息交叉插入 +- 同一 session 内 assistant 响应重叠 +- 同一 session 的 queued / active 状态切换不一致 + +如果没有这把锁,同 session 上下文混乱的概率会高很多。 + +--- + +## 4. 基于 Snapshot 的单次运行隔离 + +在真正开始流式运行之前,服务端会捕获一个 `RunSnapshot`,其中包含当前 session 作用域下的运行时值,例如: + +- `session_id` +- `conversation` +- `input_provider` +- model / mode 信息 + +随后运行过程会依赖这个 snapshot,而不是在执行过程中反复读取共享对象上的实时字段。 + +### 评估 + +这是当前防止运行中跨 session 污染的关键机制。 + +它降低了这样一种风险:某个 session 在一个上下文中启动运行,随后却在执行中意外读到了另一个 session 的实时 agent 状态。 + +不过,当前 snapshot 的捕获仍然发生在“共享 agent 临时重绑定到目标 session 之后”。这意味着 snapshot 本身具有保护作用,但其下层所有权模型依然建立在共享实例之上。 + +--- + +## 5. 父 / 子 Agent 隔离 + +子 agent 的执行使用的是显式 child session,而不是直接复用父 session。 + +当前行为是: + +- 每次委派 task 或 worker run 时都会分配新的 `child_session_id` +- 子 session 通过 `parent_session_id` / `parent_id` 关联父 session +- 子消息写入 `messages[child_session_id]` +- 父 session 中只保留一个通过 metadata 指向 child session 的 tool part + +### 评估 + +这是一个合理的设计。 + +它意味着: + +- 父 session 仍然是协调面 +- 子 session 仍然是执行面与转录面 +- 父子默认不会共享同一份可变 conversation transcript + +这是一种比较强的链路隔离形式。 + +--- + +## 6. 多窗口隔离 + +当前多窗口隔离依赖的是:**全局事件广播 + 按 session 感知的客户端路由**。 + +服务端行为: + +- 所有 SSE subscribers 维护在一个全局集合中 +- `broadcast_event()` 将事件广播给所有订阅者 +- 事件本身携带 `sessionId`,或者服务端根据事件结构推导 `sessionId` + +客户端行为: + +- 每个窗口按 `sessionId` 对事件进行路由或过滤 + +### 评估 + +这属于**UI 层的逻辑隔离**,而不是传输层硬隔离。 + +该模型成立的前提是: + +1. 所有 session 绑定事件都带有正确的 session 身份 +2. 客户端过滤逻辑是正确的 + +因此,多窗口正确性是真实存在的,但它部分依赖客户端实现。 + +--- + +## 7. Agent 级运行时状态归属 + +这里是当前架构张力最明显的地方。 + +代码已经明确警告:同一个共享 agent 实例上的并发执行并不安全。如下字段仍然归属于 agent 实例本身: + +- `_active_run_ctx` +- `_iteration_task` +- `conversation` +- `internal_fs` + +### 评估 + +这是当前最大的残余风险区域。 + +当前设计通过以下机制组合,避免了很多实际错误: + +- session lock +- agent lock +- snapshot capture +- session-scoped message history + +但真正执行 run 的底层对象本身,仍然不是完全 session-native 的。 + +换句话说: + +> 隔离模型说的是“每个 session 一份”,但一部分运行时所有权仍然在说“每个 agent 实例一份”。 + +这两者之间的不一致,就是当前最核心的未解决问题。 + +--- + +## 已经隔离得比较好的部分 + +下面这些部分目前相对稳固: + +### Session 身份 +- 显式 `session_id` +- 显式 `parent_id` +- 存储以 session 身份为键 + +### 同 Session 时序 +- 每 session 独立 turn 锁 +- 按 session 排空 queued async prompts + +### 子 Session 链路 +- child session 显式创建 +- child messages 单独存储 +- 父 session 通过 metadata 引用子 session,而不是共享 transcript + +### 运行时事件标注 +- 事件携带 `sessionId` +- 父子事件路径是显式的 +- 子 agent 事件转发中存在 loop detection + +--- + +## 仅做到逻辑隔离、尚未做到硬隔离的部分 + +### SSE 投递 +- 先广播给所有订阅者 +- 再由 `sessionId` 过滤 + +### Agent 运行时状态 +- 共享 agent 实例仍然持有部分可变状态 +- 当前正确性依赖“短暂绑定、快速 snapshot、之后尽量不读 live state” + +### Worker 历史继承 +- 某些 worker 路径仍会临时替换另一个运行时的 history,之后再恢复 + +### 临时文件系统状态 +- `internal_fs` 是 agent-instance scoped,而不是严格的 session-runtime scoped + +--- + +## 当前未解决的问题 + +## P0:共享 Agent 实例并发复用 + +### 问题 + +同一个 agent 实例仍可能被多个 session 复用,而部分运行时状态仍然归属于实例本身。 + +### 为什么重要 + +如果某条执行路径绕过了 snapshot discipline,或者削弱了它,这里就是最容易发生跨 session 污染的地方。 + +### 可能表现 + +- 中断到了错误的 run +- 活跃 iteration 跟踪被覆盖 +- 读取到了错误的实时 conversation 状态 +- session 本地的临时输出发生混用 + +--- + +## P0:Worker History 覆盖 / 恢复模式 + +### 问题 + +某些 worker 流程会临时把 worker 的 conversation/history 设置为父级 history,执行完后再恢复。 + +### 为什么重要 + +这是一种可变共享状态模式。在完全串行的流程中它可以工作,但相比 copy-based 或 runtime-bound 的继承方式,它更脆弱。 + +--- + +## P1:全局 SSE Fan-Out + +### 问题 + +session 绑定事件目前仍然会被广播给所有 subscribers。 + +### 为什么重要 + +这会提高对客户端过滤的依赖,也使多窗口场景下的传输层隔离叙事更弱。 + +--- + +## P1:`internal_fs` 不是按 Session 持有 + +### 问题 + +临时输出与中间数据仍可能归属于 agent 实例,而不是归属于独立的 session runtime。 + +### 为什么重要 + +这会削弱后台任务、调试产物和工具输出的隔离保证。 + +--- + +## P2:Child Context 生命周期管理 + +### 问题 + +child event context 当前虽然能正确缓存和更新,但它的生命周期清理没有 session 身份模型那样明确。 + +### 为什么重要 + +这更像是完整性和可维护性问题,而不是立即会造成 session 混乱的问题,但长期可能累积出陈旧的 child runtime 状态。 + +--- + +## 风险分级汇总 + +| 优先级 | 问题 | 风险类型 | 摘要 | +|--------|------|----------|------| +| P0 | 共享 agent 实例复用 | 运行时正确性 | session 模型与 runtime ownership 之间最大的结构性错位 | +| P0 | worker history 覆盖/恢复 | 可变共享状态 | 正确性依赖严格的串行时序 | +| P1 | 全局 SSE fan-out | 传输层隔离 | UI 正确性依赖客户端路由/过滤 | +| P1 | agent-scoped `internal_fs` | 临时状态隔离 | 工具与任务输出并没有强 session 绑定 | +| P2 | child context 生命周期 | 状态卫生 | 更容易造成陈旧状态,而不是直接造成跨 session 污染 | + +--- + +## 为什么系统今天通常仍然是可用的 + +尽管存在上述未解决问题,系统在大多数场景下仍然表现正确,因为有多层保护在同时生效: + +1. session ID 是显式且稳定的 +2. 同 session turn 被串行化 +3. 存在每 session 独立 message history +4. 子 agent 执行会创建 child session,而不是直接复用父 session +5. 长时间运行依赖 `RunSnapshot` +6. 消息与 part 的持久化始终按 session 落盘 + +这套组合在常规流程里已经足够强。 + +当前的架构问题并不是“实现已经混乱”。真正的问题是:**最强的正确性保证依赖多个约定同时成立**。 + +--- + +## 建议方向 + +建议的长期方向是: + +1. 引入显式 `SessionRuntime` +2. 将 conversation、active run context、iteration task、internal filesystem 等状态都迁入其中 +3. 让 worker 和 subagent 通过 copy 或 runtime binding 继承上下文,而不是通过临时 mutation 继承 +4. 将 SSE 投递从全局 fan-out 改为按 session 感知的服务端路由 + +这一路径的详细设计见: + +- `docs/rfcs/draft/RFC-0023-session-runtime-hard-isolation.md` + +--- + +## 当前状态时序图 + +```mermaid +sequenceDiagram + autonumber + participant W1 as Window A + participant W2 as Window B + participant API as OpenCode API + participant ST as ServerState + participant SL as session_lock(session_id) + participant AL as agent_lock + participant AG as Shared Agent Instance + participant SNAP as RunSnapshot + participant EP as EventProcessor + participant CS as Child Session + participant SSE as SSE Broadcast + participant UI as Client Router(by sessionId) + + W1->>API: POST /session/S1/message + API->>ST: append user message to messages[S1] + API->>SSE: broadcast MessageUpdated(sessionId=S1) + + API->>SL: acquire lock(S1) + SL-->>API: granted + + API->>ST: mark session S1 busy + API->>SSE: broadcast SessionStatus(sessionId=S1,busy) + + API->>AL: acquire agent_lock + AL-->>API: granted + API->>AG: bind session_id=S1, input_provider=S1 + API->>ST: snapshot_for_session(S1) + ST-->>SNAP: {session_id, conversation[S1], input_provider[S1], mode, model} + API->>AL: release agent_lock + + API->>AG: run_stream(..., snapshot=S1) + AG-->>EP: stream events(session S1) + + alt tool spawns subagent + AG-->>EP: SpawnSessionStart(child=C1,parent=S1) + EP->>ST: ensure_session(C1,parent=S1) + EP->>CS: create child messages in messages[C1] + EP->>SSE: broadcast child events(sessionId=C1) + EP->>SSE: broadcast parent ToolPart(metadata.sessionId=C1) + end + + EP->>SSE: broadcast parts/messages/status + SSE-->>W1: all events + SSE-->>W2: all events + W1->>UI: filter sessionId=S1/C1 + W2->>UI: filter own sessionId + + API->>ST: persist final assistant message for S1 + API->>ST: mark S1 idle + API->>SSE: broadcast SessionIdle(sessionId=S1) + API->>SL: release lock(S1) +``` + +--- + +## 目标状态时序图 + +```mermaid +sequenceDiagram + autonumber + participant UI as Client Window + participant API as OpenCode API + participant ST as ServerState + participant SR as SessionRuntime(S1) + participant AR as Runtime-bound Agent Execution + participant AG as Agent Definition + participant CRT as SessionRuntime(C1) + participant SSE as Session-aware SSE Router + + UI->>API: POST /session/S1/message + API->>ST: get_or_create SessionRuntime(S1) + ST-->>SR: runtime for S1 + + API->>SR: acquire runtime lock + API->>SR: append user message + API->>SSE: emit only to subscribers(S1) + + API->>AR: run_stream(agent=AG, runtime=SR) + + alt spawn child task + AR->>ST: create_child_runtime(C1,parent=S1) + ST-->>CRT: runtime for C1 + AR->>CRT: execute child task + CRT-->>SSE: child events only to subscribers(C1) + AR-->>SSE: parent tool updates only to subscribers(S1) + end + + AR-->>SSE: session-bound events only to subscribers(S1) + API->>SR: finalize runtime state + API->>SSE: emit SessionIdle(S1) + API->>SR: release runtime lock +``` + +--- + +## 结论 + +当前架构**并不是根本性混乱的**,但也**还没有达到完全硬隔离**。 + +当前实现里最强的部分是: + +- 显式 session identity +- 每 session turn 串行化 +- child session lineage +- 基于 snapshot 的运行中隔离 + +当前实现里最弱的部分是: + +- 共享 agent 实例上的 runtime ownership +- 基于临时 mutation 的 worker 继承方式 +- 全局 SSE fan-out +- agent-scoped 的临时文件系统状态 + +因此,更准确的结论是: + +> 当前系统之所以大多数时候可靠,是因为多层逻辑隔离机制协同工作;而它剩余的架构性风险,主要集中在运行时状态仍由共享对象持有这一点上。 diff --git a/docs/rca/opencode-1.4.4-compat-and-queue-unblock.md b/docs/rca/opencode-1.4.4-compat-and-queue-unblock.md new file mode 100644 index 000000000..fd2cc34bc --- /dev/null +++ b/docs/rca/opencode-1.4.4-compat-and-queue-unblock.md @@ -0,0 +1,192 @@ +# RCA: OpenCode 1.4.4+ compatibility and queued prompt unblock + +## Summary + +This issue was caused by **three separate but related problems**: + +1. **Protocol compatibility gaps for OpenCode 1.4.4+** + - missing OTLP compatibility endpoints + - missing `/global/*` compatibility routes + +2. **Incorrect completion signaling for queued prompts** + - the server only emitted completion when the **entire async queue** drained + - the client expected a **per-turn completion signal** (`session.idle`) + +3. **A route-owned timeout on synchronous `/message` processing** + - the server wrapped the whole streamed turn in a hard `asyncio.timeout(...)` + - long silent waits (especially question / permission waits) were treated as route failures + +Together, these caused: +- `405 Method Not Allowed` during OTLP export +- startup/runtime compatibility issues with newer OpenCode clients +- delayed exit after the first interaction +- queued prompts appearing stuck until much later +- `500 Internal Server Error` when a sync OpenCode turn stayed silent for too long +- follow-up `/question/{id}/reply` requests returning `404` after the server had already torn down the pending question + +--- + +## Symptoms + +- OpenCode `1.4.3` worked normally +- OpenCode `1.4.4+` showed failures +- logs included: + - `Failed to export metrics batch code: 405, reason: Method Not Allowed` +- after the first interaction: + - the client could not exit promptly + - later prompts were queued + - the previous turn was not considered finished in real time +- after a longer silent wait during `/message` processing: + - the route failed with `TimeoutError` + - the agent stream saw `CancelledError` + - later question replies could hit `404 Not Found` + +--- + +## Root Cause + +### 1. Missing OTLP endpoints +Newer OpenCode clients send: +- `POST /v1/metrics` +- `POST /v1/traces` +- `POST /v1/logs` + +These requests were not explicitly handled and fell through to a catch-all route that only supported `GET/HEAD/OPTIONS`, resulting in `405`. + +### 2. Missing global compatibility routes +The server did not provide: +- `GET /global/config` +- `PATCH /global/config` +- `POST /global/dispose` +- `POST /global/upgrade` + +This created compatibility issues for newer client lifecycle flows. + +### 3. Completion signaling was too coarse +Queued async prompts only triggered `session.idle` when the **whole queue** drained. + +However, the client depends on `session.idle` as the signal that: +- the current turn is done +- input can be unblocked +- the session can continue or exit + +As a result, a completed queued turn could still appear unfinished to the client. + +### 4. Busy-session enqueue did not guarantee worker startup +If an async prompt was enqueued while a synchronous `/message` turn was already running, the queue item could exist without a guaranteed worker handoff immediately after the sync turn completed. + +### 5. `/message` used request/response timeout semantics for an event-driven interaction +OpenCode session turns are not purely request/response. + +During a single `/message` turn, the server may legitimately spend a long time with no streamed model output while it is: +- waiting for a tool approval +- waiting for a question reply +- waiting for other user-driven side-channel events + +The server wrapped the entire sync turn in a hard timeout: + +- `async with asyncio.timeout(STREAM_TIMEOUT_SECONDS)` + +When that timeout fired, it cancelled the active stream, which produced the observed chain: + +- `TimeoutError` at the route level +- `CancelledError` inside the agent's event queue wait +- `500` returned from `POST /session/{id}/message` +- cleanup of pending question state, causing later `/question/{id}/reply` to return `404` + +The key mistake was treating "no stream output yet" as equivalent to "the turn is broken". In this protocol, a silent turn can still be healthy because progress may be happening through separate events. + +--- + +## Fix + +### OTLP compatibility +Added minimal compatibility sinks: +- `POST /v1/metrics` +- `POST /v1/traces` +- `POST /v1/logs` + +These return success and intentionally discard payloads. + +### Global compatibility routes +Added minimal safe compatibility routes: +- `GET /global/config` +- `PATCH /global/config` +- `POST /global/dispose` +- `POST /global/upgrade` + +Behavior is intentionally minimal: +- config routes reuse existing config behavior +- dispose/upgrade are safe stubs with no destructive side effects + +### Queued prompt unblock +Updated queue/session lifecycle handling so that: +- each queued async turn emits a **per-turn** completion signal +- full `mark_session_idle()` only happens when the queue is truly empty +- sync `/message` completion immediately hands off to queued async work +- async enqueue always ensures a worker exists + +This was implemented by: +- adding queue/worker helper methods in `ServerState` +- adding `emit_session_turn_complete(session_id)` +- updating async queue draining logic in `message_routes.py` +- updating sync-to-async handoff behavior + +### Long-wait sync turn handling +Updated synchronous `/message` processing so that it no longer applies a route-owned hard timeout to the entire agent stream. + +Behavior now: +- the sync turn remains alive as long as the underlying work is still active +- question / permission flows can complete through the normal event endpoints +- the route no longer converts a legitimate silent wait into `CancelledError` + `TimeoutError` + `500` + +This was implemented by: +- removing the `asyncio.timeout(...)` wrapper around `adapter.process_stream(iterator)` in `message_routes.py` +- preserving the existing event-driven lifecycle instead of forcing request-timeout semantics onto it + +--- + +## Validation + +Targeted regression coverage was added for: +- OTLP compatibility endpoints +- `/global/*` compatibility routes +- queued async prompt worker startup +- per-turn `session.idle` signaling +- long-running sync `/message` turns that stay silent before resuming +- concurrency behavior +- SSE/global event compliance +- OpenCode storage/project persistence + +### Result +- **256 passed** +- **6 skipped** + +--- + +## Key Lessons + +1. **End of streamed output is not the same as turn completion** + - clients often depend on explicit lifecycle events, not just stream exhaustion + +2. **Queue completion and turn completion are different concepts** + - signaling only when the whole queue drains is too coarse for interactive clients + +3. **A silent turn is not necessarily a hung turn** + - OpenCode uses side-channel events (`question`, permission, SSE lifecycle) during active work + - route-level hard timeouts can break valid interactions by destroying state the client still depends on + +4. **Compatibility fixes should start with minimal safe behavior** + - protocol recovery first, full functionality later + +5. **Do not trust delegated implementation blindly** + - route-level compatibility fixes must be verified by reading actual code and running focused tests + +--- + +## Related commits + +- `5537d2d7f` `fix(opencode-storage): add project persistence support for OpenCode storage` +- `6de476f92` `fix(opencode-server): add OTLP compatibility sinks for 1.4.4+` +- `e28090479` `fix(opencode-server): add global compatibility routes for newer clients` +- `c28e2f3f1` `fix(opencode-server): unblock queued prompts on per-turn completion` diff --git a/docs/rfcs/draft/RFC-0022-opencode-v144-global-event-protocol.md b/docs/rfcs/accepted/RFC-0022-opencode-v144-global-event-protocol.md similarity index 95% rename from docs/rfcs/draft/RFC-0022-opencode-v144-global-event-protocol.md rename to docs/rfcs/accepted/RFC-0022-opencode-v144-global-event-protocol.md index 6ea962dc8..18b56596a 100644 --- a/docs/rfcs/draft/RFC-0022-opencode-v144-global-event-protocol.md +++ b/docs/rfcs/accepted/RFC-0022-opencode-v144-global-event-protocol.md @@ -1,12 +1,12 @@ --- rfc_id: RFC-0022 title: OpenCode v1.4.4+ GlobalEvent Protocol Support -status: DRAFT +status: ACCEPTED author: yuchen.liu reviewers: [] created: 2026-04-15 -last_updated: 2026-04-15 -decision_date: +last_updated: 2026-04-16 +decision_date: 2026-04-16 related_prds: [] related_rfcs: - RFC-0013-subagent-event-unification.md @@ -635,33 +635,46 @@ No new endpoints. Changes are limited to the SSE data format: --- -## Decision Record +## Implementation Status + +> Implemented 2026-04-16. Key deviations from the original design: + +| # | Deviation | Rationale | +|---|-----------|-----------| +| 1 | `GlobalEventFactory` placed in `global_routes.py` instead of `state.py` | Avoids circular imports — `state.py` cannot import from `models/events.py` without creating a dependency cycle | +| 2 | `wrap()` returns `str` (JSON) instead of `GlobalEvent` model instance | The SSE generator needs serialized JSON strings; returning the model would require the caller to serialize, adding unnecessary coupling | +| 3 | Reuses `_serialize_event()` for payload generation | The existing `_serialize_event()` already handles event serialization correctly; duplicating that logic in the factory would violate DRY | +| 4 | Uses `json.dumps(ensure_ascii=False)` instead of `model_dump_json()` | `model_dump_json()` escapes non-ASCII characters by default; `ensure_ascii=False` preserves Unicode content in event payloads | + +--- -> Complete this section after RFC review is concluded. +## Decision Record ### Decision -**Status**: PENDING +**Status**: ACCEPTED -**Date**: +**Date**: 2026-04-16 -**Approvers**: +**Approvers**: yuchen.liu ### Decision Summary -[To be completed after review] +Accepted Option 2 (Formal GlobalEvent Model) with pragmatic implementation deviations documented above. ### Key Discussion Points -[To be completed after review] +- Circular import issue required moving factory out of `state.py` +- Serialization strategy prioritized simplicity and DRY over strict model-driven design ### Conditions of Approval -[To be completed after review] +- Must not break the existing `/event` endpoint +- Must pass manual testing with OpenCode v1.4.4+ TUI ### Dissenting Opinions -[To be completed after review] +None --- diff --git a/docs/rfcs/draft/RFC-0023-session-runtime-hard-isolation.md b/docs/rfcs/draft/RFC-0023-session-runtime-hard-isolation.md new file mode 100644 index 000000000..6618e74f2 --- /dev/null +++ b/docs/rfcs/draft/RFC-0023-session-runtime-hard-isolation.md @@ -0,0 +1,614 @@ +--- +rfc_id: RFC-0023 +title: Session Runtime Hard Isolation for Multi-Window and Subagent Execution +status: DRAFT +author: Hephaestus +reviewers: [] +created: 2026-04-20 +last_updated: 2026-04-20 +decision_date: +related_prds: [] +related_rfcs: + - RFC-0014-spawn-session-events.md + - RFC-0021-agent-concurrent-execution-safety.md +--- + +# RFC-0023: Session Runtime Hard Isolation for Multi-Window and Subagent Execution + +## Overview + +This RFC proposes moving session-sensitive runtime state out of shared agent instances and into explicit per-session runtime containers. The goal is to strengthen isolation across concurrent sessions, multiple OpenCode windows, parent/child subagent execution, and background worker flows without changing the external session model. + +The current implementation already uses `session_id`, `parent_id`, per-session locks, and `RunSnapshot` to achieve logical isolation. However, some mutable state remains owned by the agent instance or by process-wide transport structures. This RFC evaluates three architectural options and recommends a staged move to `SessionRuntime` plus session-aware server-side event routing. + +## Table of Contents + +- [Background & Context](#background--context) +- [Problem Statement](#problem-statement) +- [Goals & Non-Goals](#goals--non-goals) +- [Evaluation Criteria](#evaluation-criteria) +- [Options Analysis](#options-analysis) +- [Recommendation](#recommendation) +- [Technical Design](#technical-design) +- [Security Considerations](#security-considerations) +- [Implementation Plan](#implementation-plan) +- [Open Questions](#open-questions) +- [Decision Record](#decision-record) +- [References](#references) + +--- + +## Background & Context + +### Current State + +AgentPool currently runs OpenCode sessions inside a shared server process. Runtime state in `ServerState` is heavily partitioned by `session_id`, including message lists, todo state, conversation caches, input providers, and per-session locks. Subagent execution creates explicit child sessions linked via `parent_id` and event metadata. + +When a message is processed, OpenCode acquires a per-session lock, applies short-lived agent mutations under `agent_lock`, captures a `RunSnapshot`, and then streams the run using snapshot-derived `session_id`, `conversation`, and `input_provider`. This reduces cross-session contamination during ordinary execution. + +At the same time, several mutable fields remain instance-scoped on the agent or process-scoped on the server: + +- `BaseAgent._active_run_ctx` +- `NativeAgent._iteration_task` +- `BaseAgent.conversation` +- `BaseAgent.internal_fs` +- global SSE subscriber fan-out in `ServerState.broadcast_event()` + +### Historical Context + +Two earlier RFCs are directly relevant: + +1. **RFC-0014** introduced `SpawnSessionStart` and formalized child-session creation for subagent flows. +2. **RFC-0021** addressed concurrent execution safety at the per-run context level and documented that some shared instance state was unsafe under concurrent execution. + +This RFC builds on those changes. It does not replace session tracking or subagent lineage. Instead, it hardens the ownership boundary of runtime state. + +### Glossary + +| Term | Definition | +|------|------------| +| Agent Definition | A reusable agent object that holds configuration, model capabilities, tools, and shared static dependencies | +| Session Runtime | A proposed per-session state container that owns conversation, execution state, and temporary resources | +| Logical Isolation | Correctness achieved by IDs, routing, locks, and conventions inside shared processes or objects | +| Hard Isolation | Correctness achieved by moving mutable state ownership to independently scoped runtime containers | +| Child Session | A subagent session with its own `session_id` and a `parent_id` referring to the spawning session | + +--- + +## Problem Statement + +### The Problem + +The current architecture isolates most state by `session_id`, but it still relies on shared agent instances and shared server transport for part of its correctness. This creates a gap between the intended isolation model and the actual ownership model of mutable runtime state. + +The result is not an immediate correctness failure in every case. The result is a system where cross-session correctness depends on multiple cooperating mechanisms remaining aligned: + +- per-session locks +- short `agent_lock` critical sections +- consistent use of `RunSnapshot` +- careful event tagging with `sessionId` +- client-side filtering of globally broadcast events + +### Evidence + +The following observations are verified from the current implementation: + +- `message_routes.py` serializes turns per session with `get_session_lock(session_id)`. +- `state.py` creates per-session `MessageHistory` and per-session `OpenCodeInputProvider` instances. +- `snapshot_for_session()` still binds `resolved.session_id` and `resolved._input_provider` on the shared agent before capturing the snapshot. +- `base_agent.py` logs that concurrent runs on a shared agent instance are not safe because `_active_run_ctx` is single-session state. +- `native_agent/agent.py` documents the same constraint for `_iteration_task`. +- `workers.py` temporarily replaces worker history in some execution modes, then restores it later. +- `broadcast_event()` fans events out to all SSE subscribers, while `global_routes.py` relies on `sessionId` extraction and client-side routing to associate events with windows. + +### Impact of Inaction + +If the ownership model remains unchanged: + +- **Cost**: new features must continue threading through special-case locking, snapshotting, and restore logic +- **Risk**: new call paths may accidentally read or mutate instance-scoped state during concurrent multi-session runs +- **Opportunity**: server-side selective event delivery, session-local debug tooling, and stronger background-task isolation remain harder to implement + +The main concern is architectural fragility rather than a single isolated bug. The current design works under disciplined usage, but it is easier to regress than a model where runtime state ownership is explicit. + +--- + +## Goals & Non-Goals + +### Goals (In Scope) + +1. Move session-sensitive mutable runtime state to an explicit per-session container. +2. Preserve the existing external session model (`session_id`, `parent_id`, child session lineage). +3. Eliminate the need to bind shared agent instance state before snapshot capture. +4. Replace client-only event isolation with server-side session-aware event delivery. +5. Remove execution paths that temporarily overwrite another runtime's conversation/history. + +### Non-Goals (Out of Scope) + +1. Redesigning the AgentPool storage model or changing session primary keys. +2. Replacing the existing parent/child session UI model in OpenCode. +3. Solving distributed execution, thread safety, or cross-process scheduling. +4. Changing model provider semantics beyond what is required for session-local execution configuration. +5. Rewriting all agent backends at once; the design must support staged migration. + +### Success Criteria + +How will we know this RFC achieved its goals? + +- [ ] Concurrent sessions using the same agent definition no longer share conversation, internal filesystem, or run-task ownership. +- [ ] Worker and subagent flows do not mutate another runtime's history in place. +- [ ] Session-bound SSE events are delivered server-side only to subscribers for that session, plus optional global subscribers. +- [ ] `snapshot_for_session()` no longer mutates shared agent instance session fields before producing a run snapshot. +- [ ] Parent and child sessions still preserve lineage and UI behavior after migration. + +--- + +## Evaluation Criteria + +The following criteria will be used to objectively evaluate each option: + +| Criterion | Weight | Description | Minimum Threshold | +|-----------|--------|-------------|-------------------| +| Isolation Strength | High | Degree to which mutable runtime state is owned per session rather than shared | Session-local ownership for conversation, run context, and temporary files | +| Backward Compatibility | High | Ability to preserve session APIs, storage shape, and current client behavior | No breaking OpenCode route or storage changes | +| Implementation Risk | Medium | Probability of regressions during migration | Can be staged with compatibility shims | +| Operational Simplicity | Medium | Ease of debugging, tracing, and canceling session-specific work | Per-session ownership must improve observability | +| Performance | Medium | Impact on throughput, latency, and resource use | No requirement for global serialization across sessions | +| Extensibility | Medium | Ability to support future subagent, background worker, and window-routing features | New features should not rely on shared instance mutation | + +--- + +## Options Analysis + +### Option 1: Keep Shared Agent Instances and Add More Guards + +**Description** + +Retain the current ownership model and continue improving correctness with additional locks, invariants, assertions, and helper wrappers. Under this option, `ServerState` remains the primary partitioning layer, while agent instances continue to hold some mutable execution state. + +**Advantages** + +- Smallest implementation delta from the current codebase. +- Keeps current agent backend interfaces largely unchanged. +- Can reduce some short-term risk by tightening discipline around known unsafe paths. + +**Disadvantages** + +- Does not change the underlying ownership mismatch between sessions and shared instance state. +- Leaves correctness dependent on all future call paths honoring locking and snapshot discipline. +- Preserves client-side dependence on filtering globally broadcast session events. + +**Evaluation Against Criteria** + +| Criterion | Rating | Notes | +|-----------|--------|-------| +| Isolation Strength | Low | State remains partially instance-scoped | +| Backward Compatibility | High | Minimal interface change | +| Implementation Risk | Medium | Low code churn, but continued hidden coupling | +| Operational Simplicity | Medium | Debugging remains split across shared instance and session buckets | +| Performance | High | No structural overhead beyond more checks | +| Extensibility | Low | Future features still depend on shared mutable state conventions | + +**Effort Estimate** + +- Complexity: Low +- Resources: 1 engineer, 3-5 days +- Dependencies: auditing unsafe paths, adding tests and assertions + +**Risk Assessment** + +| Risk | Likelihood | Impact | Mitigation | +|------|------------|--------|------------| +| Residual cross-session coupling remains | High | High | Add more assertions and tests | +| Future regressions reintroduce unsafe reads | Medium | High | Code review checklists and runtime warnings | + +--- + +### Option 2: Introduce Per-Session Runtime Containers (Recommended) + +**Description** + +Introduce a first-class `SessionRuntime` object that owns all session-sensitive mutable execution state. Agent objects become reusable definitions that execute against a supplied runtime rather than storing session state internally. + +Representative fields for `SessionRuntime`: + +```python +@dataclass +class SessionRuntime: + session_id: str + parent_session_id: str | None + conversation: MessageHistory + input_provider: OpenCodeInputProvider + internal_fs: IsolatedMemoryFileSystem + active_run_ctx: AgentRunContext | None + iteration_task: asyncio.Task[Any] | None + lock: asyncio.Lock + queued_prompts: list[QueuedAsyncPrompt] + status: SessionStatus + child_sessions: set[str] +``` + +**Advantages** + +- Aligns ownership boundaries with the isolation model already implied by `session_id`. +- Removes the need to bind shared agent state before snapshot capture. +- Provides a clean home for session-local filesystem, cancellation, async queue, and runtime metrics. + +**Disadvantages** + +- Requires interface changes in agents, workers, and protocol adapters. +- Introduces a migration period where compatibility shims may coexist with the new runtime model. +- Some backends may need adapter code if they currently assume instance-owned state. + +**Evaluation Against Criteria** + +| Criterion | Rating | Notes | +|-----------|--------|-------| +| Isolation Strength | High | Runtime ownership becomes session-local | +| Backward Compatibility | Medium | External APIs can remain stable, internal APIs will change | +| Implementation Risk | Medium | Migration must be staged carefully | +| Operational Simplicity | High | Session ownership becomes explicit and easier to inspect | +| Performance | Medium | Modest extra runtime objects; avoids global serialization | +| Extensibility | High | Child sessions, tracing, and routing all become easier to extend | + +**Effort Estimate** + +- Complexity: High +- Resources: 1-2 engineers, 2-4 weeks +- Dependencies: runtime abstraction, backend shims, focused regression tests + +**Risk Assessment** + +| Risk | Likelihood | Impact | Mitigation | +|------|------------|--------|------------| +| Migration complexity causes temporary dual-state confusion | Medium | High | Add deprecation shims and ownership assertions | +| Backend-specific assumptions break during migration | Medium | Medium | Use adapters and backend-by-backend rollout | + +--- + +### Option 3: Clone Agent Instances Per Session + +**Description** + +Create a dedicated agent instance per session or per run. Instead of separating definition from runtime state, this option duplicates the agent object so mutable state is no longer shared. + +**Advantages** + +- Stronger isolation than the current design without redesigning internal ownership as deeply as Option 2. +- Can reduce contention around instance-scoped fields by construction. +- May be easier for some backends that strongly assume instance-local execution state. + +**Disadvantages** + +- Duplicates tools, providers, caches, and model-adapter state that may not need duplication. +- Increases memory and startup overhead proportional to session count. +- Leaves ambiguity around which state should be shared versus copied. + +**Evaluation Against Criteria** + +| Criterion | Rating | Notes | +|-----------|--------|-------| +| Isolation Strength | Medium | Stronger than today, but state-sharing policy remains implicit | +| Backward Compatibility | Medium | External APIs can stay stable, but lifecycle semantics change | +| Implementation Risk | Medium | Fewer interface changes than Option 2, but more lifecycle complexity | +| Operational Simplicity | Medium | Easier reasoning per instance, harder resource accounting | +| Performance | Low | More memory and setup work per session | +| Extensibility | Medium | Helps isolation, but does not create a reusable runtime abstraction | + +**Effort Estimate** + +- Complexity: Medium +- Resources: 1-2 engineers, 1-3 weeks +- Dependencies: instance factory logic, cache and provider lifecycle decisions + +**Risk Assessment** + +| Risk | Likelihood | Impact | Mitigation | +|------|------------|--------|------------| +| Session count multiplies resource use | Medium | Medium | Pool clones or add eviction policies | +| Shared-vs-copied state remains inconsistent | Medium | High | Define explicit clone semantics for each field | + +--- + +### Options Comparison Summary + +| Criterion | Option 1 | Option 2 | Option 3 | +|-----------|----------|----------|----------| +| Isolation Strength | Low | High | Medium | +| Backward Compatibility | High | Medium | Medium | +| Implementation Risk | Medium | Medium | Medium | +| Operational Simplicity | Medium | High | Medium | +| Performance | High | Medium | Low | +| Extensibility | Low | High | Medium | +| **Overall** | Short-term mitigation | Best long-term fit | Partial structural fix | + +--- + +## Recommendation + +### Recommended Option + +**Option 2: Introduce Per-Session Runtime Containers** + +### Justification + +Option 2 scores highest against the criteria that matter most for this problem: isolation strength, operational clarity, and future extensibility. The current system already models sessions explicitly in storage, runtime buckets, and event lineage. A `SessionRuntime` design extends that same model to mutable execution ownership. + +Option 1 is suitable as a short-term mitigation but does not resolve the architectural mismatch. Option 3 improves isolation but still leaves key ownership questions implicit and may duplicate more state than necessary. Based on the evidence in the current codebase, Option 2 is the smallest change that resolves the structural issue rather than only reducing exposure to it. + +### Accepted Trade-offs + +1. **Internal API churn**: Acceptable because the current API already passes snapshot and history objects explicitly, which creates a natural migration path toward runtime-based execution. +2. **Staged migration complexity**: Acceptable because compatibility shims can preserve external behavior while ownership is moved incrementally. + +### Conditions + +- The migration must preserve `session_id` / `parent_id` semantics and OpenCode route compatibility. +- Ownership changes must be introduced with regression tests covering multi-window, child-session, cancellation, and worker flows. +- Session-aware SSE routing should ship behind a compatibility flag until clients are validated. + +--- + +## Technical Design + +> Note: This section contains a preliminary design suitable for RFC review. Exact field names and adapters may change during implementation. + +### Architecture Overview + +The proposed design separates agent definition from session runtime. + +```text +┌────────────────────┐ +│ Client Window │ +└─────────┬──────────┘ + │ POST /session/{id}/message + ▼ +┌────────────────────┐ +│ OpenCode API │ +└─────────┬──────────┘ + │ get_or_create_runtime(session_id) + ▼ +┌────────────────────┐ ┌────────────────────┐ +│ ServerState │───────▶│ SessionRuntime │ +│ session_runtimes │ │ conversation │ +│ session_subscribers│ │ input_provider │ +└─────────┬──────────┘ │ internal_fs │ + │ │ active_run_ctx │ + │ run(agent, runtime)│ iteration_task │ + ▼ └─────────┬──────────┘ +┌────────────────────┐ │ +│ Agent Definition │◀─────────────────┘ +│ tools / model cfg │ +└────────────────────┘ +``` + +### Proposed Data Model + +New server-owned container: + +```python +@dataclass +class SessionRuntime: + session_id: str + parent_session_id: str | None + conversation: MessageHistory + input_provider: OpenCodeInputProvider + internal_fs: IsolatedMemoryFileSystem + active_run_ctx: AgentRunContext | None = None + iteration_task: asyncio.Task[Any] | None = None + lock: asyncio.Lock = field(default_factory=asyncio.Lock) + queued_prompts: list[QueuedAsyncPrompt] = field(default_factory=list) + status: SessionStatus = field(default_factory=lambda: SessionStatus(type="idle")) + child_sessions: set[str] = field(default_factory=set) +``` + +New `ServerState` fields: + +```python +session_runtimes: dict[str, SessionRuntime] +global_subscribers: list[Queue[Event]] +session_subscribers: dict[str, list[Queue[Event]]] +``` + +### Execution Model + +The runtime becomes the primary execution context: + +```python +async def run_stream( + self, + *prompts: PromptCompatible, + runtime: SessionRuntime, + snapshot: RunSnapshot | None = None, + execution_config: ExecutionConfig | None = None, +) -> AsyncIterator[RichAgentStreamEvent[Any]]: + ... +``` + +Rules: + +1. `run_stream()` must not write `self.session_id`, `self._input_provider`, `self._active_run_ctx`, `self._iteration_task`, `self.conversation`, or `self.internal_fs` as part of ordinary session execution. +2. Interrupt, cancellation, and async queue ownership move to `SessionRuntime`. +3. Snapshots are generated from `SessionRuntime`, not by mutating the shared agent first. + +### Snapshot Model + +`snapshot_for_session()` changes from bind-then-capture to runtime-derived capture: + +```python +RunSnapshot( + session_id=runtime.session_id, + input_provider=runtime.input_provider, + conversation=runtime.conversation, + model_name=effective_model_name, + mode_name=effective_mode_name, +) +``` + +### Worker and Subagent Inheritance + +Worker and subagent execution must not overwrite another runtime's history in place. Instead, inheritance becomes explicit: + +```python +history_copy = parent_runtime.conversation.copy() +child_runtime = create_child_runtime(child_session_id, parent_session_id=parent_runtime.session_id) +await child_agent.run_stream(..., runtime=child_runtime, inherited_history=history_copy) +``` + +Recommended inheritance modes: + +- `NONE` +- `COPY_PARENT_VISIBLE_HISTORY` +- `COPY_PARENT_COMPACTED_HISTORY` +- `LINK_READONLY_PARENT_CONTEXT` (optional future extension) + +### Server-Side Event Routing + +Current behavior broadcasts every event to every subscriber. Proposed behavior: + +1. Global events go only to `global_subscribers`. +2. Session-bound events go to `session_subscribers[session_id]` and optionally to global subscribers. +3. Event extraction rules remain the same, but routing is enforced server-side rather than delegated entirely to the client. + +### Child Session Lifecycle + +Child session lifecycle becomes explicit and runtime-backed: + +1. `SpawnSessionStart` -> create child `SessionRuntime` +2. wrapped child events -> update child runtime and parent tool part +3. `StreamCompleteEvent` / cancellation / error -> finalize child runtime and parent state +4. runtime cleanup -> release child runtime resources when safe + +### Sequence Diagram + +```mermaid +sequenceDiagram + autonumber + participant UI as Client Window + participant API as OpenCode API + participant ST as ServerState + participant RT as SessionRuntime(S1) + participant AG as Agent Definition + participant RUN as Runtime-bound Run + participant CRT as Child SessionRuntime(C1) + participant SSE as Session-aware SSE + + UI->>API: POST /session/S1/message + API->>ST: get_or_create_runtime(S1) + ST-->>RT: SessionRuntime(S1) + + API->>RT: acquire runtime.lock + API->>RT: append user message + API->>SSE: emit MessageUpdated to subscribers(S1) + + API->>RUN: run_stream(agent=AG, runtime=RT) + RUN-->>SSE: stream session-bound events to S1 subscribers + + alt spawn child task + RUN->>ST: create_child_runtime(C1,parent=S1) + ST-->>CRT: SessionRuntime(C1) + RUN->>CRT: execute child task + CRT-->>SSE: child events to subscribers(C1) + RUN-->>SSE: parent tool update(metadata.sessionId=C1) to subscribers(S1) + end + + RUN->>RT: finalize conversation, status, and filesystem outputs + API->>SSE: emit SessionIdle(S1) + API->>RT: release runtime.lock +``` + +--- + +## Security Considerations + +This RFC is primarily about runtime correctness, but it has security and privacy implications. + +1. **Event Delivery Scope**: server-side session-aware routing reduces unnecessary exposure of session-bound events to unrelated subscribers. +2. **Temporary Output Isolation**: per-session internal filesystems reduce accidental cross-session access to task output, debug artifacts, or tool-generated files. +3. **Cancellation Scope**: session-owned run state makes it easier to ensure interrupts and background-task cancellation only affect the intended session. +4. **Auditability**: explicit runtime ownership makes it easier to reason about which component wrote or read session-local data. + +This RFC does not replace authentication or authorization. If clients with different trust boundaries share one server process, session-scoped event routing should be treated as necessary but not sufficient protection. + +--- + +## Implementation Plan + +### Phase 1: Safety Improvements Before Structural Migration + +1. Add stronger guards against concurrent reuse of shared agent instance execution state. +2. Remove worker paths that overwrite another runtime's history in place. +3. Namespace temporary filesystem output by session and task ID even before `internal_fs` is fully moved. + +### Phase 2: Introduce `SessionRuntime` + +1. Add `SessionRuntime` and `ServerState.session_runtimes`. +2. Move conversation, input provider, async prompt queue, active run state, and temporary filesystem ownership into the runtime. +3. Preserve compatibility shims for older internal call paths during migration. + +### Phase 3: Runtime-Native Execution + +1. Update agent execution to accept a runtime object directly. +2. Change snapshot generation to derive from runtime, not from shared instance mutation. +3. Update interrupt and cancel flows to locate active work through the runtime. + +### Phase 4: Session-Aware Event Delivery + +1. Add global vs session-scoped subscriber registries. +2. Deliver session-bound events only to matching subscribers. +3. Keep compatibility mode available during rollout. + +### Rollback Strategy + +- Keep session-aware routing behind a feature flag until validated with the current OpenCode client. +- Maintain a compatibility shim for old execution entry points until tests confirm runtime-native paths are stable. +- Roll back phase-by-phase rather than in one large revert; each phase should remain independently deployable. + +--- + +## Open Questions + +1. Should model/mode selection become runtime-local execution configuration, or should some backends use lightweight clone-on-run behavior? +2. Which existing agent backends depend most heavily on instance-owned state and need adapters first? +3. Should child runtime cleanup be reference-counted, event-driven, or explicit on session close? +4. Do any current OpenCode clients depend on receiving globally broadcast session events beyond their active session? +5. Should `LINK_READONLY_PARENT_CONTEXT` be supported initially, or deferred until after runtime ownership is stable? + +--- + +## Decision Record + +**Current Status**: DRAFT + +### Draft Decision + +No final decision has been made yet. This RFC recommends Option 2 for review because it provides the strongest alignment between the current session model and runtime state ownership. + +### Review Focus + +Reviewers are asked to comment on: + +1. Whether `SessionRuntime` is the right ownership boundary. +2. Whether any backend requires a different migration path. +3. Whether server-side event routing should be introduced in the same RFC or split into a follow-up RFC. + +### Approval Conditions + +- Agreement on the ownership boundary for session-sensitive state. +- Agreement on phased rollout and compatibility strategy. +- Agreement on the minimum regression test matrix for concurrent sessions, subagents, and multi-window routing. + +--- + +## References + +- `src/agentpool_server/opencode_server/state.py` +- `src/agentpool_server/opencode_server/routes/message_routes.py` +- `src/agentpool_server/opencode_server/routes/global_routes.py` +- `src/agentpool_server/opencode_server/event_processor.py` +- `src/agentpool/agents/base_agent.py` +- `src/agentpool/agents/native_agent/agent.py` +- `src/agentpool_toolsets/builtin/workers.py` +- `src/agentpool_toolsets/builtin/subagent_tools.py` +- `docs/rfcs/accepted/RFC-0014-spawn-session-events.md` +- `docs/rfcs/accepted/RFC-0021-agent-concurrent-execution-safety.md` diff --git a/scripts/diagnose_skills.py b/scripts/diagnose_skills.py new file mode 100644 index 000000000..40cfd258b --- /dev/null +++ b/scripts/diagnose_skills.py @@ -0,0 +1,371 @@ +#!/usr/bin/env python3 +"""Skills TUI display diagnostic script. + +Checks every stage of the skill discovery pipeline to identify +why skills may not appear in the TUI. + +Usage: + uv run python scripts/diagnose_skills.py [config.yml] +""" + +from __future__ import annotations + +import asyncio +import sys +from pathlib import Path +from typing import Any + +# Add src to path +sys.path.insert(0, str(Path(__file__).parent.parent / "src")) + + +def print_header(title: str) -> None: + print(f"\n{'=' * 60}") + print(f" {title}") + print(f"{'=' * 60}") + + +def print_result(name: str, passed: bool, detail: str = "") -> None: + icon = "✅" if passed else "❌" + print(f" {icon} {name}") + if detail: + print(f" {detail}") + + +async def diagnose(config_path: str | None = None) -> None: + """Run full diagnostic pipeline.""" + from upathtools import UPath, to_upath + + # ============================================================ + # Stage 1: Filesystem Discovery + # ============================================================ + print_header("Stage 1: Filesystem Discovery") + + from agentpool_config.skills import DEFAULT_SKILLS_PATHS + + # Check default skill directories + for dp in DEFAULT_SKILLS_PATHS: + resolved = to_upath(dp).expanduser() + exists = resolved.exists() + print_result( + f"Default path: {dp} (resolved: {resolved})", + exists, + "Directory exists" if exists else "Directory NOT found", + ) + if exists: + # Check for SKILL.md in subdirectories + subdirs = [p for p in resolved.iterdir() if p.is_dir()] + if subdirs: + for subdir in subdirs: + skill_md = subdir / "SKILL.md" + has_skill_md = skill_md.exists() + print_result( + f" Subdir: {subdir.name}/", + has_skill_md, + "SKILL.md found" if has_skill_md else "SKILL.md NOT found", + ) + if has_skill_md: + # Try to parse frontmatter + try: + content = skill_md.read_text() + if content.startswith("---"): + end = content.find("---", 3) + if end != -1: + import yaml + + frontmatter = yaml.safe_load(content[3:end]) + has_name = "name" in (frontmatter or {}) + has_desc = "description" in (frontmatter or {}) + print_result( + " Frontmatter has 'name'", + has_name, + f"name={frontmatter.get('name')!r}" + if has_name + else "MISSING", + ) + print_result( + " Frontmatter has 'description'", + has_desc, + f"description={frontmatter.get('description')!r}" + if has_desc + else "MISSING", + ) + # Check for unknown keys + from agentpool.skills.skill import SkillMetadata + + try: + SkillMetadata.model_validate(frontmatter) + print_result( + " Frontmatter validation", + True, + "Passes strict validation", + ) + except Exception as e: + print_result( + " Frontmatter validation", False, f"FAILED: {e}" + ) + else: + print_result( + " Frontmatter parsing", False, "No closing --- found" + ) + else: + print_result( + " Frontmatter parsing", + False, + "No YAML frontmatter (must start with ---)", + ) + except Exception as e: + print_result(" Frontmatter parsing", False, str(e)) + else: + print_result( + f" No subdirectories in {resolved}", False, "Skills must be in SUBDIRECTORIES" + ) + + # Check custom paths from config + if config_path: + print("\n --- Custom paths from config ---") + try: + from agentpool.models.manifest import AgentsManifest + + manifest = AgentsManifest.from_file(config_path) + if manifest.skills: + for p in manifest.skills.paths: + resolved = to_upath(p).expanduser() + exists = resolved.exists() + print_result( + f"Custom path: {p} (resolved: {resolved})", + exists, + "Directory exists" if exists else "Directory NOT found", + ) + print_result( + "include_default=True", + manifest.skills.include_default, + "Default paths will also be searched" + if manifest.skills.include_default + else "Default paths will NOT be searched", + ) + else: + print_result("Skills config in manifest", False, "No skills section in config") + except Exception as e: + print_result("Config loading", False, str(e)) + + # ============================================================ + # Stage 2: SkillsManager / SkillsRegistry + # ============================================================ + print_header("Stage 2: SkillsManager / SkillsRegistry") + + from agentpool.skills.manager import SkillsManager + + try: + config_file_path = to_upath(config_path) if config_path else None + from agentpool_config.skills import SkillsConfig + + skills_config: SkillsConfig | None = None + if config_path: + from agentpool.models.manifest import AgentsManifest + + manifest = AgentsManifest.from_file(config_path) + skills_config = manifest.skills + + manager = SkillsManager( + name="diagnostic", + config=skills_config, + config_file_path=config_file_path, + ) + async with manager: + skill_names = manager.registry.list_items() + print_result( + "SkillsManager discovered skills", + len(skill_names) > 0, + f"Found {len(skill_names)} skills: {skill_names}" + if skill_names + else "NO skills found!", + ) + + # Check resource_provider + try: + rp = manager.resource_provider + print_result("ResourceProvider available", True, f"type={type(rp).__name__}") + except RuntimeError as e: + print_result("ResourceProvider available", False, str(e)) + + except Exception as e: + print_result("SkillsManager initialization", False, str(e)) + + # ============================================================ + # Stage 3: AgentPool initialization (if config provided) + # ============================================================ + if config_path: + print_header("Stage 3: AgentPool Initialization") + + try: + from agentpool.delegation import AgentPool + + async with AgentPool(config_path) as pool: + # Check skill_commands + sc = pool.skill_commands + print_result( + "pool.skill_commands is not None", + sc is not None, + f"type={type(sc).__name__}" + if sc + else "This is the problem! Bridge won't be created.", + ) + + if sc is not None: + cmd_names = list(sc._items.keys()) + print_result( + "skill_commands has entries", + len(cmd_names) > 0, + f"Commands: {cmd_names}" if cmd_names else "NO commands registered!", + ) + + # Check skill_provider + sp = pool.skill_provider + print_result( + "pool.skill_provider is not None", + sp is not None, + f"type={type(sp).__name__}" if sp else "Skill provider not initialized!", + ) + + if sp is not None: + try: + provider_skills = await sp.get_skills() + print_result( + "skill_provider returns skills", + len(provider_skills) > 0, + f"Found {len(provider_skills)} skills: {[s.name for s in provider_skills]}", + ) + except Exception as e: + print_result("skill_provider.get_skills()", False, str(e)) + + # Check skills registry + skills = pool.skills + skill_list = skills.list_skills() + print_result( + "pool.skills has skills", + len(skill_list) > 0, + f"Found {len(skill_list)} skills: {[s.name for s in skill_list]}", + ) + + except Exception as e: + print_result("AgentPool initialization", False, str(e)) + import traceback + + traceback.print_exc() + + # ============================================================ + # Stage 4: OpenCode Server Bridge (simulated) + # ============================================================ + print_header("Stage 4: OpenCode Server Bridge Check") + + if config_path: + try: + from agentpool.delegation import AgentPool + + async with AgentPool(config_path) as pool: + sc = pool.skill_commands + if sc is not None: + # Simulate what server.py does + from agentpool_server.opencode_server.skill_bridge import OpenCodeSkillBridge + + bridge = OpenCodeSkillBridge(skill_provider=pool.skill_provider) + sc.on_command_change(bridge.handle_change) + + commands = bridge.get_commands() + print_result( + "Bridge has commands after subscription", + len(commands) > 0, + f"Commands: {[c.name for c in commands]}" + if commands + else "NO commands in bridge!", + ) + + skill_commands = bridge.get_skill_commands() + print_result( + "Bridge has skill_commands", + len(skill_commands) > 0, + f"Skills: {[c.name for c in skill_commands]}" + if skill_commands + else "NO skill commands!", + ) + else: + print_result( + "Bridge creation", + False, + "skill_commands is None - bridge will NOT be created! This is the root cause.", + ) + except Exception as e: + print_result("Bridge simulation", False, str(e)) + else: + print(" ⚠️ Skipped (no config path provided)") + + # ============================================================ + # Stage 5: HTTP API Check (if server is running) + # ============================================================ + print_header("Stage 5: HTTP API Check") + + import urllib.request + import json + + base_url = "http://127.0.0.1:4096" + for endpoint in ["/skill", "/command"]: + try: + req = urllib.request.Request(f"{base_url}{endpoint}") + with urllib.request.urlopen(req, timeout=3) as resp: + data = json.loads(resp.read()) + if isinstance(data, list): + print_result( + f"GET {endpoint}", + len(data) > 0, + f"Returns {len(data)} items: {[d.get('name', '?') for d in data]}", + ) + else: + print_result(f"GET {endpoint}", True, f"Returns: {data}") + except Exception as e: + print_result(f"GET {endpoint}", False, f"Server not reachable or error: {e}") + + # ============================================================ + # Summary + # ============================================================ + print_header("Summary & Common Fixes") + + print(""" +Common fixes for skills not showing in TUI: + +1. SKILL.md must be in a SUBDIRECTORY of the skills dir: + ✅ ~/.claude/skills/my-skill/SKILL.md + ❌ ~/.claude/skills/SKILL.md + +2. SKILL.md must have valid YAML frontmatter: + --- + name: my-skill + description: What this skill does + --- + (No unknown keys allowed! Only: name, description) + +3. Config must include the skills directory: + skills: + paths: + - ./my-skills + include_default: true + +4. Relative paths resolve against the CONFIG FILE location, not CWD. + +5. Check server logs with: + OBSERVABILITY_ENABLED=true agentpool serve-opencode config.yml + +6. If pool.skill_commands is None, the OpenCode bridge is never created, + which means /command endpoint won't return skill commands. +""") + + +if __name__ == "__main__": + config = sys.argv[1] if len(sys.argv) > 1 else None + if config: + print(f"Using config: {config}") + else: + print("No config path provided. Some checks will be skipped.") + print("Usage: uv run python scripts/diagnose_skills.py [config.yml]") + + asyncio.run(diagnose(config)) diff --git a/src/agentpool/agents/acp_agent/acp_agent.py b/src/agentpool/agents/acp_agent/acp_agent.py index 933a42d52..a97cd5aad 100644 --- a/src/agentpool/agents/acp_agent/acp_agent.py +++ b/src/agentpool/agents/acp_agent/acp_agent.py @@ -80,7 +80,7 @@ from acp.schema.capabilities import AgentCapabilities from acp.schema.mcp import McpServer from agentpool.agents.acp_agent.client_handler import ACPClientHandler - from agentpool.agents.context import AgentRunContext + from agentpool.agents.context import AgentRunContext, RunSnapshot from agentpool.agents.events import RichAgentStreamEvent from agentpool.agents.modes import ModeCategory from agentpool.common_types import AnyEventHandlerType @@ -414,6 +414,7 @@ async def _stream_events( # noqa: PLR0915 deps: TDeps | None = None, wait_for_connections: bool | None = None, store_history: bool = True, + snapshot: RunSnapshot | None = None, ) -> AsyncIterator[RichAgentStreamEvent[str]]: from agentpool.agents.acp_agent.acp_converters import ( convert_to_acp_content, @@ -426,6 +427,7 @@ async def _stream_events( # noqa: PLR0915 if not self._api or not self._sdk_session_id or not self._state: raise AgentNotInitializedError + effective_session_id = snapshot.session_id if snapshot else self.session_id run_id = str(uuid.uuid4()) self._state.clear() model_messages: list[ModelResponse | ModelRequest] = [] @@ -434,9 +436,9 @@ async def _stream_events( # noqa: PLR0915 current_response_parts: list[TextPart | ThinkingPart | ToolCallPart] = [] text_chunks: list[str] = [] - assert self.session_id is not None + assert effective_session_id is not None yield RunStartedEvent( - session_id=self.session_id, + session_id=effective_session_id, run_id=run_id, agent_name=self.name, parent_session_id=parent_session_id, @@ -514,7 +516,7 @@ async def poll_acp_events() -> AsyncIterator[RichAgentStreamEvent[str]]: role="assistant", name=self.name, message_id=message_id or str(uuid.uuid4()), - session_id=self.session_id, + session_id=effective_session_id, parent_id=user_msg.message_id, model_name=self.model_name, messages=model_messages, @@ -551,7 +553,7 @@ async def poll_acp_events() -> AsyncIterator[RichAgentStreamEvent[str]]: role="assistant", name=self.name, message_id=message_id or str(uuid.uuid4()), - session_id=self.session_id, + session_id=effective_session_id, parent_id=user_msg.message_id, model_name=self.model_name, messages=model_messages, diff --git a/src/agentpool/agents/agui_agent/agui_agent.py b/src/agentpool/agents/agui_agent/agui_agent.py index dde92a9a7..08f6b08b9 100644 --- a/src/agentpool/agents/agui_agent/agui_agent.py +++ b/src/agentpool/agents/agui_agent/agui_agent.py @@ -55,7 +55,7 @@ from slashed import BaseCommand from tokonomics.model_discovery.model_info import ModelInfo - from agentpool.agents.context import AgentRunContext + from agentpool.agents.context import AgentRunContext, RunSnapshot from agentpool.agents.events import RichAgentStreamEvent from agentpool.agents.modes import ModeCategory from agentpool.common_types import AnyEventHandlerType, StrPath, ToolType @@ -295,6 +295,7 @@ async def _stream_events( # noqa: PLR0915 deps: TDeps | None = None, wait_for_connections: bool | None = None, store_history: bool = True, + snapshot: RunSnapshot | None = None, ) -> AsyncIterator[RichAgentStreamEvent[str]]: from ag_ui.core import RunAgentInput, UserMessage @@ -310,9 +311,11 @@ async def _stream_events( # noqa: PLR0915 if not self._client: raise AgentNotInitializedError + effective_session_id = snapshot.session_id if snapshot else self.session_id + # Set thread_id from session_id (needed for AG-UI protocol) if self._sdk_session_id is None: - self._sdk_session_id = self.session_id + self._sdk_session_id = effective_session_id run_id = str(uuid4()) # New run ID for each run # Track messages in pydantic-ai format: ModelRequest -> ModelResponse -> ModelRequest... @@ -322,8 +325,8 @@ async def _stream_events( # noqa: PLR0915 initial_request = ModelRequest(parts=[UserPromptPart(content=prompts)]) model_messages.append(initial_request) response_parts: list[TextPart | ThinkingPart | ToolCallPart] = [] - assert self.session_id is not None # Initialized by BaseAgent.run_stream() - thread_id = self._sdk_session_id or self.session_id + assert effective_session_id is not None # Initialized by BaseAgent.run_stream() + thread_id = self._sdk_session_id or effective_session_id yield RunStartedEvent( session_id=thread_id, run_id=run_id, @@ -350,7 +353,7 @@ async def _stream_events( # noqa: PLR0915 break request_data = RunAgentInput( - thread_id=self._sdk_session_id or self.session_id, + thread_id=self._sdk_session_id or effective_session_id, run_id=run_id, state={}, messages=messages, @@ -436,7 +439,7 @@ async def _stream_events( # noqa: PLR0915 role="assistant", name=self.name, message_id=message_id or str(uuid4()), - session_id=self.session_id, + session_id=effective_session_id, parent_id=user_msg.message_id, messages=model_messages, finish_reason="stop", @@ -464,7 +467,7 @@ async def _stream_events( # noqa: PLR0915 role="assistant", name=self.name, message_id=message_id or str(uuid4()), - session_id=self.session_id, + session_id=effective_session_id, parent_id=user_msg.message_id, messages=model_messages, usage=usage, diff --git a/src/agentpool/agents/base_agent.py b/src/agentpool/agents/base_agent.py index 6c5152bd1..ba224df60 100644 --- a/src/agentpool/agents/base_agent.py +++ b/src/agentpool/agents/base_agent.py @@ -44,6 +44,7 @@ from upathtools.filesystems import OverlayFileSystem from acp.schema import AvailableCommandsUpdate + from agentpool.agents.context import RunSnapshot from agentpool.agents.events import ( CommandCompleteEvent, CommandOutputEvent, @@ -614,6 +615,7 @@ async def run_stream( wait_for_connections: bool | None = None, deps: TDeps | None = None, event_handlers: Sequence[AnyEventHandlerType] | None = None, + snapshot: RunSnapshot | None = None, ) -> AsyncIterator[RichAgentStreamEvent[TResult]]: """Run agent with streaming output. @@ -633,25 +635,13 @@ async def run_stream( wait_for_connections: Whether to wait for connected agents deps: Optional dependencies event_handlers: Optional event handlers + snapshot: Optional per-run snapshot for concurrent session isolation Yields: Stream events during execution """ from agentpool.utils.identifiers import generate_session_id - # Initialize session_id once for the entire run (including queued prompts) - if self.session_id is None: - self.session_id = session_id or generate_session_id() - self.parent_session_id = parent_session_id - user_prompts = [str(p) for p in prompts if isinstance(p, str)] - initial_prompt = user_prompts[-1] if user_prompts else None - await self.log_session( - initial_prompt, model=self.model_name, parent_session_id=self.parent_session_id - ) - elif session_id and self.session_id != session_id: - self.session_id = session_id - self.parent_session_id = parent_session_id - # Create per-run context for state isolation run_ctx = AgentRunContext(deps=deps) # Reset cancellation state and track current task @@ -660,6 +650,25 @@ async def run_stream( run_ctx.current_task = asyncio.current_task() # Track the stream task so interrupt() can cancel it even without run_ctx self._current_stream_task = run_ctx.current_task + + # Initialize session_id once for the entire run (including queued prompts) + if snapshot is not None: + # Snapshot is the source of truth — do NOT mutate self.session_id + run_ctx.snapshot = snapshot + run_ctx.session_id = snapshot.session_id + else: + run_ctx.snapshot = None + if self.session_id is None: + self.session_id = session_id or generate_session_id() + self.parent_session_id = parent_session_id + user_prompts = [str(p) for p in prompts if isinstance(p, str)] + initial_prompt = user_prompts[-1] if user_prompts else None + await self.log_session( + initial_prompt, model=self.model_name, parent_session_id=self.parent_session_id + ) + elif session_id and self.session_id != session_id: + self.session_id = session_id + self.parent_session_id = parent_session_id # Store run_ctx as instance variable so interrupt() can find it # from a different task (ContextVar is task-scoped and returns None # when read from outside the run_stream task). @@ -698,6 +707,7 @@ async def run_stream( wait_for_connections=wait_for_connections, deps=deps, event_handlers=event_handlers, + snapshot=snapshot, ): yield event @@ -724,6 +734,7 @@ async def _run_stream_once( wait_for_connections: bool | None = None, deps: TDeps | None = None, event_handlers: Sequence[AnyEventHandlerType] | None = None, + snapshot: RunSnapshot | None = None, ) -> AsyncIterator[RichAgentStreamEvent[TResult]]: """Process a single prompt group with streaming output. @@ -743,6 +754,7 @@ async def _run_stream_once( wait_for_connections: Whether to wait for connected agents deps: Optional dependencies event_handlers: Optional event handlers + snapshot: Optional per-run snapshot for concurrent session isolation Yields: Stream events during execution @@ -766,7 +778,7 @@ async def _run_stream_once( user_msg = ChatMessage.user_prompt( message=converted_prompts, parent_id=effective_parent_id, - session_id=self.session_id, + session_id=snapshot.session_id if snapshot else self.session_id, ) # Resolve event handlers @@ -800,7 +812,7 @@ async def _run_stream_once( prompt=user_msg.content if isinstance(user_msg.content, str) else str(user_msg.content), - session_id=self.session_id, + session_id=snapshot.session_id if snapshot else self.session_id, ) if pre_run_result.get("decision") == "deny": reason = pre_run_result.get("reason", "Blocked by pre-run hook") @@ -821,6 +833,7 @@ async def _run_stream_once( input_provider=input_provider, wait_for_connections=wait_for_connections, deps=deps, + snapshot=snapshot, ): await resolved_handler(context, event) yield event @@ -848,7 +861,7 @@ async def _run_stream_once( agent_name=self.name, prompt=prompt_str, result=final_message.content, - session_id=self.session_id, + session_id=snapshot.session_id if snapshot else self.session_id, ) # Emit signal (always - for event handlers) @@ -1001,6 +1014,7 @@ def _stream_events( deps: TDeps | None = None, wait_for_connections: bool | None = None, store_history: bool = True, + snapshot: RunSnapshot | None = None, ) -> AsyncIterator[RichAgentStreamEvent[TResult]]: """Agent-specific streaming implementation. @@ -1021,6 +1035,7 @@ def _stream_events( deps: Optional dependencies wait_for_connections: Whether to wait for connected agents store_history: Whether to store in history + snapshot: Optional per-run snapshot for concurrent session isolation Yields: Stream events during execution @@ -1076,7 +1091,11 @@ def is_cancelled(self) -> bool: ) return self._cancelled or background_cancelled - async def interrupt(self, run_ctx: AgentRunContext | None = None) -> None: + async def interrupt( + self, + run_ctx: AgentRunContext | None = None, + session_id: str | None = None, + ) -> None: """Interrupt the currently running stream. Sets the cancelled flag, calls subclass-specific _interrupt(), @@ -1086,9 +1105,29 @@ async def interrupt(self, run_ctx: AgentRunContext | None = None) -> None: falls back to the active run_ctx stored by run_stream() so that run_ctx.cancelled is set and the streaming loop can exit. + When session_id is provided, performs a targeted interrupt for that + specific session only, without setting self._cancelled (which would + kill all sessions). + Args: run_ctx: Optional per-run context for the stream to interrupt + session_id: Optional session ID for targeted session-scoped interrupt """ + if session_id is not None: + # Targeted interrupt — only cancel this session's context + if run_ctx is not None: + run_ctx.cancelled = True + effective_run_ctx = run_ctx or self._active_run_ctx + if effective_run_ctx: + effective_run_ctx.cancelled = True + await self._interrupt(effective_run_ctx) + await self.interrupted.emit(self.InterruptEvent(agent_name=self.name)) + logger.info( + "Agent interrupted (session-scoped)", agent=self.name, session_id=session_id + ) + return + + # Legacy global interrupt self._cancelled = True # When no run_ctx is provided, try the active per-run context # stored by run_stream() as _active_run_ctx. We can't use @@ -1146,6 +1185,7 @@ async def run( input_provider: InputProvider | None = None, event_handlers: Sequence[AnyEventHandlerType] | None = None, wait_for_connections: bool | None = None, + snapshot: RunSnapshot | None = None, ) -> ChatMessage[TResult]: """Run agent with prompt and get response. @@ -1167,6 +1207,7 @@ async def run( input_provider: Optional input provider for the agent event_handlers: Optional event handlers for this run (overrides agent's handlers) wait_for_connections: Whether to wait for connected agents to complete + snapshot: Optional per-run snapshot for concurrent session isolation Returns: ChatMessage containing response and run information @@ -1189,6 +1230,7 @@ async def run( input_provider=input_provider, event_handlers=event_handlers, wait_for_connections=wait_for_connections, + snapshot=snapshot, ): if isinstance(event, StreamCompleteEvent): final_message = event.message diff --git a/src/agentpool/agents/claude_code_agent/claude_code_agent.py b/src/agentpool/agents/claude_code_agent/claude_code_agent.py index 552a66019..948ecad1b 100644 --- a/src/agentpool/agents/claude_code_agent/claude_code_agent.py +++ b/src/agentpool/agents/claude_code_agent/claude_code_agent.py @@ -148,6 +148,7 @@ ClaudeCodeServerInfo, ) from agentpool.agents.context import AgentRunContext + from agentpool.agents.context import RunSnapshot from agentpool.agents.events import RichAgentStreamEvent from agentpool.agents.modes import ModeCategory from agentpool.common_types import AnyEventHandlerType, StrPath @@ -851,6 +852,7 @@ async def _stream_events( # noqa: PLR0915 deps: TDeps | None = None, wait_for_connections: bool | None = None, store_history: bool = True, + snapshot: RunSnapshot | None = None, ) -> AsyncIterator[RichAgentStreamEvent[TResult]]: from clawd_code_sdk import ( AssistantMessage, @@ -887,8 +889,10 @@ async def _stream_events( # noqa: PLR0915 prompt_text = " ".join(str(p) for p in prompts) run_id = str(uuid.uuid4()) assert self.session_id is not None # Initialized by BaseAgent.run_stream() + effective_session_id = snapshot.session_id if snapshot else self.session_id + assert effective_session_id is not None yield RunStartedEvent( - session_id=self.session_id, + session_id=effective_session_id, run_id=run_id, agent_name=self.name, parent_session_id=parent_session_id, @@ -1199,9 +1203,9 @@ async def _stream_events( # noqa: PLR0915 role="assistant", name=self.name, message_id=message_id or str(uuid.uuid4()), - session_id=self.session_id, + session_id=effective_session_id, parent_id=user_msg.message_id, - model_name=resolved_model or self.model_name, + model_name=resolved_model or (snapshot.model_name if snapshot else self.model_name), messages=model_messages, finish_reason="stop", metadata=metadata, @@ -1272,9 +1276,9 @@ async def _stream_events( # noqa: PLR0915 role="assistant", name=self.name, message_id=message_id or str(uuid.uuid4()), - session_id=self.session_id, + session_id=effective_session_id, parent_id=user_msg.message_id, - model_name=resolved_model or self.model_name, + model_name=resolved_model or (snapshot.model_name if snapshot else self.model_name), messages=model_messages, cost_info=cost_info, usage=request_usage or RequestUsage(), diff --git a/src/agentpool/agents/claude_code_agent/converters.py b/src/agentpool/agents/claude_code_agent/converters.py index 45807b1db..e9f63eba3 100644 --- a/src/agentpool/agents/claude_code_agent/converters.py +++ b/src/agentpool/agents/claude_code_agent/converters.py @@ -284,8 +284,7 @@ def _convert_edit_result(result: EditOutput) -> EditMetadata: additions, deletions = _count_diff_changes(structured_patch) filediff = FileDiff( file=file_path, - before=original_file or "", - after=after_content or "", + patch=diff, additions=additions, deletions=deletions, ) diff --git a/src/agentpool/agents/codex_agent/codex_agent.py b/src/agentpool/agents/codex_agent/codex_agent.py index 291ae7356..038ef52d5 100644 --- a/src/agentpool/agents/codex_agent/codex_agent.py +++ b/src/agentpool/agents/codex_agent/codex_agent.py @@ -42,6 +42,7 @@ from tokonomics.model_discovery.model_info import ModelInfo from agentpool.agents.context import AgentRunContext + from agentpool.agents.context import RunSnapshot from agentpool.agents.events import RichAgentStreamEvent from agentpool.agents.modes import ModeCategory from agentpool.common_types import AnyEventHandlerType, MCPServerStatus, StrPath @@ -353,6 +354,7 @@ async def _stream_events( # noqa: PLR0915 deps: TDeps | None = None, wait_for_connections: bool | None = None, store_history: bool = True, + snapshot: RunSnapshot | None = None, ) -> AsyncIterator[RichAgentStreamEvent[OutputDataT]]: """Stream events from Codex turn execution.""" from codex_adapter.events import ( @@ -370,7 +372,7 @@ async def _stream_events( # noqa: PLR0915 # Generate IDs if not provided run_id = str(uuid4()) final_message_id = message_id or str(uuid4()) - final_session_id = session_id or self.session_id + final_session_id = snapshot.session_id if snapshot else (session_id or self.session_id) # Ensure session_id is set (should always be from base class) if final_session_id is None: raise ValueError("session_id must be set") diff --git a/src/agentpool/agents/context.py b/src/agentpool/agents/context.py index 20e2cf1c5..3bd416fba 100644 --- a/src/agentpool/agents/context.py +++ b/src/agentpool/agents/context.py @@ -60,6 +60,9 @@ class AgentRunContext: session_id: str = field(default_factory=lambda: uuid.uuid4().hex) """Unique identifier for this run session.""" + snapshot: RunSnapshot | None = None + """Per-run snapshot captured under agent_lock for concurrent session isolation.""" + deps: Any = None """Optional dependencies passed to the run.""" @@ -67,6 +70,23 @@ class AgentRunContext: """Timestamp when the run started (for metrics).""" +@dataclass(kw_only=True) +class RunSnapshot: + """Immutable per-run state captured from the shared agent under a short lock. + + Once created, this snapshot is the source of truth for the entire run. + In-flight runs MUST NOT read live singleton fields (self.session_id, etc.) + -- they read from this snapshot instead. + """ + + session_id: str + input_provider: Any = None + conversation: Any = None # MessageHistory instance + model_name: str | None = None + mode_name: str | None = None + parent_session_id: str | None = None + + @dataclass(kw_only=True) class AgentContext[TDeps = Any](NodeContext[TDeps]): """Runtime context for agent execution. diff --git a/src/agentpool/agents/native_agent/agent.py b/src/agentpool/agents/native_agent/agent.py index 626fc1ae3..ec22edb4a 100644 --- a/src/agentpool/agents/native_agent/agent.py +++ b/src/agentpool/agents/native_agent/agent.py @@ -46,6 +46,7 @@ from upathtools import JoinablePathLike from agentpool.agents.context import AgentRunContext + from agentpool.agents.context import RunSnapshot from agentpool.agents.events import RichAgentStreamEvent from agentpool.agents.modes import ModeCategory from agentpool.common_types import ( @@ -826,6 +827,7 @@ async def _stream_events( # noqa: PLR0915 input_provider: InputProvider | None = None, wait_for_connections: bool | None = None, deps: TDeps | None = None, + snapshot: RunSnapshot | None = None, ) -> AsyncIterator[RichAgentStreamEvent[OutputDataT]]: from pydantic_graph import End @@ -836,8 +838,10 @@ async def _stream_events( # noqa: PLR0915 start_time = time.perf_counter() history_list = message_history.get_history() assert self.session_id is not None # Initialized by BaseAgent.run_stream() + effective_session_id = snapshot.session_id if snapshot else self.session_id + assert effective_session_id is not None yield RunStartedEvent( - session_id=self.session_id, + session_id=effective_session_id, run_id=run_id, agent_name=self.name, parent_session_id=parent_session_id, @@ -905,7 +909,7 @@ async def agent_iteration_task() -> None: role="assistant", name=self.name, message_id=message_id, - session_id=self.session_id, + session_id=effective_session_id, parent_id=user_msg.message_id, response_time=response_time, finish_reason="stop", @@ -916,7 +920,7 @@ async def agent_iteration_task() -> None: agent_run.result, agent_name=self.name, message_id=message_id, - session_id=self.session_id, + session_id=effective_session_id, parent_id=user_msg.message_id, response_time=time.perf_counter() - start_time, metadata=None, diff --git a/src/agentpool/storage/manager.py b/src/agentpool/storage/manager.py index ebae99ac6..35d5a4dd0 100644 --- a/src/agentpool/storage/manager.py +++ b/src/agentpool/storage/manager.py @@ -827,10 +827,11 @@ def get_project_provider(self) -> StorageProvider: Raises: RuntimeError: If no capable provider found. """ - if self.providers: - return self.providers[0] + for provider in self.providers: + if provider.can_store_projects: + return provider - raise RuntimeError("No provider found that supports project storage") + raise RuntimeError("No storage provider supports project storage") @method_spawner async def save_project(self, project: ProjectData) -> None: diff --git a/src/agentpool_server/opencode_server/input_provider.py b/src/agentpool_server/opencode_server/input_provider.py index 9380178a8..3e98f4e78 100644 --- a/src/agentpool_server/opencode_server/input_provider.py +++ b/src/agentpool_server/opencode_server/input_provider.py @@ -66,6 +66,11 @@ def _generate_permission_id(self) -> str: self._id_counter += 1 return f"perm_{self._id_counter}_{int(__import__('time').time() * 1000)}" + def _generate_question_id(self) -> str: + """Generate a unique question ID.""" + self._id_counter += 1 + return f"que_{self._id_counter}_{int(__import__('time').time() * 1000)}" + async def get_tool_confirmation( self, context: AgentContext[Any], @@ -193,6 +198,17 @@ def resolve_permission(self, permission_id: str, response: PermissionReply) -> b ) return True + def has_pending_permission(self, permission_id: str) -> bool: + """Check whether a specific permission request is pending. + + Args: + permission_id: The permission request ID to look up + + Returns: + True if the permission is pending, False otherwise + """ + return permission_id in self._pending_permissions + def get_pending_permissions(self) -> list[PermissionAskedProperties]: """Get all pending permission requests. @@ -292,7 +308,7 @@ async def _handle_single_enum( return types.ElicitResult(action="decline") # Extract descriptions if available (custom x-option-descriptions field) descriptions = schema.get("x-option-descriptions", {}) - question_id = self._generate_permission_id() # Reuse ID generator + question_id = self._generate_question_id() opts = [ QuestionOption(label=str(val), description=descriptions.get(str(val), "")) for val in enum_values @@ -478,7 +494,7 @@ async def _handle_multi_question( logger.warning("No valid questions could be created from object schema") return types.ElicitResult(action="decline") - question_id = self._generate_permission_id() + question_id = self._generate_question_id() # Create future to wait for answers future: asyncio.Future[list[list[str]]] = asyncio.get_event_loop().create_future() diff --git a/src/agentpool_server/opencode_server/models/__init__.py b/src/agentpool_server/opencode_server/models/__init__.py index 929f6db2e..081fbfc8d 100644 --- a/src/agentpool_server/opencode_server/models/__init__.py +++ b/src/agentpool_server/opencode_server/models/__init__.py @@ -18,11 +18,14 @@ from agentpool_server.opencode_server.models.app import ( App, AppTimeInfo, + DiagnosticResponse, + DisposeResponse, HealthResponse, PathInfo, Project, ProjectTime, ProjectUpdateRequest, + UpgradeResponse, VcsInfo, ) from agentpool_server.opencode_server.models.provider import ( @@ -125,6 +128,8 @@ ProviderAuthAuthorization, ProviderAuthMethod, SkillInfo, + WorkspaceConnectionStatus, + WorkspaceInfo, WorktreeCreateRequest, WorktreeInfo, WorktreeRemoveRequest, @@ -146,6 +151,7 @@ CommandExecutedEvent, Event, FileEditedEvent, + GlobalEvent, QuestionRepliedEvent, QuestionRejectedEvent, LspStatus, @@ -228,6 +234,8 @@ "ContextOverflowErrorData", "Diagnostic", "DiagnosticRange", + "DiagnosticResponse", + "DisposeResponse", "Event", "FileContent", "FileDiff", @@ -240,6 +248,7 @@ "FileWatcherUpdatedEvent", "FindMatch", "FormatterStatus", + "GlobalEvent", "HealthResponse", "LogRequest", "LspStatus", @@ -375,9 +384,12 @@ "TuiSessionSelectEvent", "UnknownError", "UnknownErrorData", + "UpgradeResponse", "UserMessage", "VcsBranchUpdatedEvent", "VcsInfo", + "WorkspaceConnectionStatus", + "WorkspaceInfo", "WorktreeCreateRequest", "WorktreeInfo", "WorktreeRemoveRequest", diff --git a/src/agentpool_server/opencode_server/models/agent.py b/src/agentpool_server/opencode_server/models/agent.py index 819a37790..3c81ffb4f 100644 --- a/src/agentpool_server/opencode_server/models/agent.py +++ b/src/agentpool_server/opencode_server/models/agent.py @@ -139,6 +139,47 @@ class WorktreeResetRequest(OpenCodeBaseModel): """Worktree directory path to reset.""" +WorkspaceConnectionState = Literal["connected", "connecting", "disconnected", "error"] + + +class WorkspaceInfo(OpenCodeBaseModel): + """Workspace information matching OpenCode's experimental workspace API.""" + + id: str + """Stable workspace identifier used by the TUI.""" + + type: str = "local" + """Workspace adaptor type.""" + + name: str + """Human-readable workspace name.""" + + branch: str | None = None + """Active VCS branch if known.""" + + directory: str | None = None + """Absolute workspace directory path.""" + + extra: object | None = None + """Adaptor-specific metadata.""" + + project_id: str + """Project identifier owning this workspace.""" + + +class WorkspaceConnectionStatus(OpenCodeBaseModel): + """Workspace connection status for OpenCode TUI bootstrap.""" + + workspace_id: str + """Workspace identifier corresponding to ``WorkspaceInfo.id``.""" + + status: WorkspaceConnectionState = "connected" + """Current connectivity status.""" + + error: str | None = None + """Optional connection error message.""" + + class AuthInfo(OpenCodeBaseModel): """Authentication credential info.""" diff --git a/src/agentpool_server/opencode_server/models/app.py b/src/agentpool_server/opencode_server/models/app.py index b6d584f5e..d9673651b 100644 --- a/src/agentpool_server/opencode_server/models/app.py +++ b/src/agentpool_server/opencode_server/models/app.py @@ -27,6 +27,22 @@ class HealthResponse(OpenCodeBaseModel): version: str +class DiagnosticResponse(OpenCodeBaseModel): + """Response for /global/diagnostic endpoint.""" + + directory: str | None = None + """Working directory of the server.""" + + project: str + """Project identifier computed from the working directory.""" + + subscribers: int + """Current number of SSE event subscribers.""" + + server_version: str + """Server version string.""" + + class PathInfo(OpenCodeBaseModel): """Path information for the OpenCode instance. @@ -102,6 +118,27 @@ class VcsInfo(OpenCodeBaseModel): commit: str | None = None +class DisposeResponse(OpenCodeBaseModel): + """Response for /global/dispose endpoint (OpenCode 1.4.4+ compat). + + Minimal stub: acknowledges the request without actually shutting down. + """ + + success: bool = True + message: str = "dispose acknowledged (no-op)" + + +class UpgradeResponse(OpenCodeBaseModel): + """Response for /global/upgrade endpoint (OpenCode 1.4.4+ compat). + + Minimal stub: indicates no upgrade was performed. + """ + + success: bool = True + message: str = "upgrade not supported (stub)" + upgraded: bool = False + + class ProjectUpdateRequest(OpenCodeBaseModel): """Request to update project metadata.""" diff --git a/src/agentpool_server/opencode_server/models/common.py b/src/agentpool_server/opencode_server/models/common.py index 6616ae60f..ff56ac14e 100644 --- a/src/agentpool_server/opencode_server/models/common.py +++ b/src/agentpool_server/opencode_server/models/common.py @@ -43,10 +43,20 @@ class TimeStartEnd(OpenCodeBaseModel): class ModelRef(OpenCodeBaseModel): - """Reference to a provider model (provider_id + model_id).""" + """Reference to a provider model with optional variant. - provider_id: str - model_id: str + OpenCode v1.4.0+ nests ``variant`` inside the ``model`` object: + ``{ providerID?, modelID?, variant? }``. + + All fields are optional to support partial references — e.g. a message + that only specifies a ``variant`` (thinking effort level) without + changing the provider or model. + """ + + provider_id: str | None = None + model_id: str | None = None + variant: str | None = None + """Reasoning/thinking variant for this model (e.g. 'low', 'medium', 'high', 'max').""" class TokenCache(OpenCodeBaseModel): @@ -95,11 +105,14 @@ class TextSpan(OpenCodeBaseModel): class FileDiff(OpenCodeBaseModel): - """A file diff entry.""" + """A file diff entry. + + Matches the OpenCode v1.4.0+ SnapshotFileDiff schema: + ``{ file, patch, additions, deletions, status? }`` + """ file: str - before: str - after: str + patch: str | None = None additions: int deletions: int status: FileDiffStatus | None = None @@ -119,8 +132,7 @@ def from_file_change(cls, change: FileChange) -> Self: status = None return cls( file=change.path, - before=change.old_content or "", - after=change.new_content or "", + patch=diff_text, additions=diff_text.count("\n+"), deletions=diff_text.count("\n-"), status=status, diff --git a/src/agentpool_server/opencode_server/models/events.py b/src/agentpool_server/opencode_server/models/events.py index 708caad4b..7f1cb6841 100644 --- a/src/agentpool_server/opencode_server/models/events.py +++ b/src/agentpool_server/opencode_server/models/events.py @@ -1050,3 +1050,25 @@ def create( | TuiToastShowEvent | TuiSessionSelectEvent ) + + +class GlobalEvent(OpenCodeBaseModel): + """SSE envelope for OpenCode v1.4.4+ global event routing. + + Not an event type itself — wraps Event instances with routing metadata. + """ + + directory: str + """Working directory used for event routing in multi-directory servers.""" + + project: str | None = None + """Project identifier for event routing (git root commit SHA or 'global').""" + + workspace: str | None = None + """Workspace identifier for event routing. + + Omitted for single-directory servers. + """ + + payload: dict[str, Any] + """The wrapped event data.""" diff --git a/src/agentpool_server/opencode_server/models/message.py b/src/agentpool_server/opencode_server/models/message.py index 3c194045c..ecafbc6f6 100644 --- a/src/agentpool_server/opencode_server/models/message.py +++ b/src/agentpool_server/opencode_server/models/message.py @@ -4,7 +4,7 @@ from typing import Any, Literal, Self -from pydantic import Field +from pydantic import Field, model_validator from agentpool.utils import identifiers as identifier from agentpool_server.opencode_server.models.base import OpenCodeBaseModel @@ -71,6 +71,27 @@ class OutputFormatJsonSchema(OpenCodeBaseModel): OutputFormat = OutputFormatText | OutputFormatJsonSchema +def _migrate_variant_into_model_dict(data: dict[str, Any]) -> dict[str, Any]: + """Backward-compat: move top-level ``variant`` into ``model.variant``. + + OpenCode v1.4.0+ nests ``variant`` inside the ``model`` object. + Older clients may still send ``variant`` at the top level. + """ + variant = data.get("variant") + if variant is None: + return data + model = data.get("model") + if model is None: + # No model object — create one with just the variant + data["model"] = {"variant": variant} + elif isinstance(model, dict) and "variant" not in model: + # Model exists but has no variant — add it + model["variant"] = variant + # Remove top-level variant so it doesn't shadow the nested one + data.pop("variant", None) + return data + + class UserMessage(OpenCodeBaseModel): """User message.""" @@ -84,7 +105,14 @@ class UserMessage(OpenCodeBaseModel): summary: MessageSummary | None = None system: str | None = None tools: dict[str, bool] | None = None - variant: str | None = None + + @model_validator(mode="before") + @classmethod + def _migrate_variant_to_model(cls, data: Any) -> Any: + """Backward-compat: move top-level ``variant`` into ``model.variant``.""" + if not isinstance(data, dict): + return data + return _migrate_variant_into_model_dict(data) # --- Assistant message error types --- @@ -231,7 +259,23 @@ class AssistantMessage(OpenCodeBaseModel): summary: bool | None = None finish: str | None = None structured: Any | None = None - variant: str | None = None + + @model_validator(mode="before") + @classmethod + def _migrate_variant_to_model(cls, data: Any) -> Any: + """Backward-compat: move top-level ``variant`` into model context. + + For AssistantMessage, variant is informational only (no model object). + We keep the variant in a private field for backward compat but the + canonical location is on the associated UserMessage's model.variant. + """ + if not isinstance(data, dict): + return data + # For assistant messages, variant was informational only. + # Remove from top-level to match OpenCode v1.4.0+ schema. + # We don't have a model object on AssistantMessage to nest it in. + data.pop("variant", None) + return data class MessageWithParts(OpenCodeBaseModel): @@ -515,12 +559,14 @@ class MessageRequest(OpenCodeBaseModel): no_reply: bool | None = None system: str | None = None tools: dict[str, bool] | None = None - variant: str | None = None - """Reasoning/thinking variant for this message. - Maps to the model's variants (e.g., 'low', 'medium', 'high', 'max'). - When set, the agent will use this thinking effort level for the response. - """ + @model_validator(mode="before") + @classmethod + def _migrate_variant_to_model(cls, data: Any) -> Any: + """Backward-compat: move top-level ``variant`` into ``model.variant``.""" + if not isinstance(data, dict): + return data + return _migrate_variant_into_model_dict(data) class ShellRequest(OpenCodeBaseModel): diff --git a/src/agentpool_server/opencode_server/models/tool_metadata.py b/src/agentpool_server/opencode_server/models/tool_metadata.py index 9a1b62adb..0f5b07e83 100644 --- a/src/agentpool_server/opencode_server/models/tool_metadata.py +++ b/src/agentpool_server/opencode_server/models/tool_metadata.py @@ -60,12 +60,12 @@ class LSPDiagnostic(TypedDict): class FileDiff(TypedDict): """A file diff entry (mirrors ``Snapshot.FileDiff`` in opencode). - Used by edit and apply_patch tools. + Matches the OpenCode v1.4.0+ SnapshotFileDiff schema: + ``{ file, patch, additions, deletions, status? }`` """ file: str - before: str - after: str + patch: str additions: int deletions: int status: NotRequired[FileDiffStatus] @@ -125,7 +125,7 @@ class EditMetadata(TruncationFields): diff: str """Unified diff string.""" filediff: FileDiff - """Structured before/after content with change counts.""" + """Structured patch content with change counts.""" diagnostics: LSPDiagnosticsMap """LSP diagnostics keyed by normalized file path.""" diff --git a/src/agentpool_server/opencode_server/routes/agent_routes.py b/src/agentpool_server/opencode_server/routes/agent_routes.py index f60ac4328..ea54740b4 100644 --- a/src/agentpool_server/opencode_server/routes/agent_routes.py +++ b/src/agentpool_server/opencode_server/routes/agent_routes.py @@ -2,6 +2,7 @@ from __future__ import annotations +from pathlib import Path import re from typing import Any @@ -33,11 +34,14 @@ ProviderAuthMethod, Session, SkillInfo, + WorkspaceConnectionStatus, + WorkspaceInfo, WorktreeCreateRequest, WorktreeInfo, WorktreeRemoveRequest, WorktreeResetRequest, ) +from agentpool_storage.opencode_provider import helpers router = APIRouter(tags=["agent"]) @@ -46,6 +50,33 @@ logger = get_logger(__name__) +def _build_workspace_info(state: StateDep) -> WorkspaceInfo: + """Build the singleton local workspace description for OpenCode clients.""" + directory = state.base_path + project_id = helpers.compute_project_id(directory) + workspace_id = f"wrk_{project_id[:12]}" + + return WorkspaceInfo( + id=workspace_id, + type="local", + name=Path(directory).name, + branch=None, + directory=directory, + extra=None, + project_id=project_id, + ) + + +def _build_workspace_status(state: StateDep) -> WorkspaceConnectionStatus: + """Build the singleton local workspace connection status.""" + workspace = _build_workspace_info(state) + return WorkspaceConnectionStatus( + workspace_id=workspace.id, + status="connected", + error=None, + ) + + def _extract_hints(template: str | None) -> list[str]: """Extract input hints from a command template. @@ -534,6 +565,33 @@ async def list_sessions_global( return sessions +@router.get("/experimental/workspace") +async def list_workspaces( + state: StateDep, + directory: str | None = None, + workspace: str | None = None, +) -> list[WorkspaceInfo]: + """List workspaces for the current project. + + AgentPool currently exposes a single local workspace rooted at the attached + server working directory. Query parameters are accepted for OpenCode SDK + compatibility but do not alter the singleton response. + """ + _ = directory, workspace + return [_build_workspace_info(state)] + + +@router.get("/experimental/workspace/status") +async def get_workspace_status( + state: StateDep, + directory: str | None = None, + workspace: str | None = None, +) -> list[WorkspaceConnectionStatus]: + """Return connection status for the singleton local workspace.""" + _ = directory, workspace + return [_build_workspace_status(state)] + + @router.get("/experimental/tool/ids") async def list_tool_ids(state: StateDep) -> list[str]: """List all available tool IDs. diff --git a/src/agentpool_server/opencode_server/routes/app_routes.py b/src/agentpool_server/opencode_server/routes/app_routes.py index 3ce717567..ec7b67416 100644 --- a/src/agentpool_server/opencode_server/routes/app_routes.py +++ b/src/agentpool_server/opencode_server/routes/app_routes.py @@ -62,7 +62,12 @@ def _project_data_to_response(data: ProjectData) -> Project: async def _get_current_project(state: StateDep) -> ProjectData: - """Get or create the current project from storage.""" + """Get or create the current project from storage. + + The returned ``ProjectData`` carries ``project_id`` from + ``generate_project_id()`` (a SHA1 of the worktree path — AgentPool's + internal identifier). + """ project_store = ProjectStore(state.storage) return await project_store.get_or_create(state.working_dir) diff --git a/src/agentpool_server/opencode_server/routes/config_routes.py b/src/agentpool_server/opencode_server/routes/config_routes.py index c1f7f49f6..ef2def44e 100644 --- a/src/agentpool_server/opencode_server/routes/config_routes.py +++ b/src/agentpool_server/opencode_server/routes/config_routes.py @@ -376,6 +376,18 @@ async def _get_variants_from_agent(agent: object) -> dict[str, dict[str, object] return {} +@router.get("/global/config") +async def get_global_config(state: StateDep) -> Config: + """Get server configuration (global alias for OpenCode 1.4.4+ compat).""" + return await get_config(state) + + +@router.patch("/global/config") +async def update_global_config(state: StateDep, config_update: Config) -> Config: + """Update server configuration (global alias for OpenCode 1.4.4+ compat).""" + return await update_config(state, config_update) + + @router.get("/config/providers") async def get_providers(state: StateDep) -> ProvidersResponse: """Get available providers and models from agent.""" diff --git a/src/agentpool_server/opencode_server/routes/global_routes.py b/src/agentpool_server/opencode_server/routes/global_routes.py index f9f610ec9..6c25c346a 100644 --- a/src/agentpool_server/opencode_server/routes/global_routes.py +++ b/src/agentpool_server/opencode_server/routes/global_routes.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import contextlib import json from typing import TYPE_CHECKING, Any @@ -11,31 +12,62 @@ from agentpool import log from agentpool_server.opencode_server.dependencies import StateDep -from agentpool_server.opencode_server.models import Event, HealthResponse # noqa: TC001 +from agentpool_server.opencode_server.models import GlobalEvent, HealthResponse +from agentpool_server.opencode_server.models.app import ( + DiagnosticResponse, + DisposeResponse, + UpgradeResponse, +) from agentpool_server.opencode_server.models.events import ( + CommandExecutedEvent, + FileEditedEvent, + FileWatcherUpdatedEvent, + LspClientDiagnosticsEvent, + LspUpdatedEvent, + McpToolsChangedEvent, MessageRemovedEvent, + MessageUpdatedEvent, + PartDeltaEvent, PartRemovedEvent, PartUpdatedEvent, PermissionRequestEvent, PermissionResolvedEvent, + PermissionUpdatedEvent, + ProjectUpdatedEvent, + PtyCreatedEvent, + PtyDeletedEvent, + PtyExitedEvent, + PtyUpdatedEvent, QuestionAskedEvent, QuestionRejectedEvent, QuestionRepliedEvent, ServerConnectedEvent, + ServerHeartbeatEvent, SessionCompactedEvent, SessionCreatedEvent, SessionDeletedEvent, + SessionDiffEvent, SessionErrorEvent, SessionIdleEvent, SessionStatusEvent, SessionUpdatedEvent, TodoUpdatedEvent, + TuiCommandExecuteEvent, + TuiPromptAppendEvent, + TuiSessionSelectEvent, + TuiToastShowEvent, + VcsBranchUpdatedEvent, +) +from agentpool_server.opencode_server.routes.routing import ( + RoutingCheckResponse, + tui_event_filter, ) if TYPE_CHECKING: from collections.abc import AsyncGenerator + from agentpool_server.opencode_server.models import Event from agentpool_server.opencode_server.state import ServerState @@ -51,8 +83,54 @@ async def get_health() -> HealthResponse: return HealthResponse(healthy=True, version=VERSION) +@router.get("/global/diagnostic") +async def get_diagnostic(state: StateDep) -> DiagnosticResponse: + """Get server diagnostic information. + + Returns directory, project, subscriber count, and server version. + """ + if state.working_dir is None: + return DiagnosticResponse( + directory=None, + project="", + subscribers=len(state.event_subscribers), + server_version=VERSION, + ) + + factory = state.get_event_factory() + return DiagnosticResponse( + directory=state.working_dir, + project=factory._project, + subscribers=len(state.event_subscribers), + server_version=VERSION, + ) + + +@router.post("/global/dispose") +async def post_global_dispose() -> DisposeResponse: + """Acknowledge OpenCode dispose requests without stopping the server.""" + return DisposeResponse(message="dispose acknowledged (no-op)") + + +@router.post("/global/upgrade") +async def post_global_upgrade() -> UpgradeResponse: + """Acknowledge OpenCode upgrade requests without performing an upgrade.""" + return UpgradeResponse(message="upgrade not supported (stub)") + + def _extract_session_id(event: Event) -> str | None: # noqa: PLR0911 - """Extract session_id from various event types.""" + """Extract session ID from various event types. + + Uses pattern matching to access session_id from four different + property structures: + - properties.session_id (most events) + - properties.info.id (SessionCreated/Updated events) + - properties.info.session_id (MessageUpdatedEvent) + - properties.part.session_id (PartUpdatedEvent) + + Unrecognized event types trigger a warning log and return None, + since some events genuinely have no session association. + """ match event: # Events with properties.session_id directly case SessionDeletedEvent(properties=props): @@ -81,6 +159,39 @@ def _extract_session_id(event: Event) -> str | None: # noqa: PLR0911 return props.session_id case SessionErrorEvent(properties=props): return props.session_id + case SessionDiffEvent(properties=props): + return props.session_id + case PartDeltaEvent(properties=props): + return props.session_id + case PermissionUpdatedEvent(properties=props): + return props.session_id + case CommandExecutedEvent(properties=props): + return props.session_id + case TuiSessionSelectEvent(properties=props): + return props.session_id + + # Events with no session association (explicitly listed to avoid + # spurious warnings; these events are broadcast globally and are + # not tied to any particular session). + case ( + ServerHeartbeatEvent() + | ServerConnectedEvent() + | FileWatcherUpdatedEvent() + | FileEditedEvent() + | McpToolsChangedEvent() + | PtyCreatedEvent() + | PtyUpdatedEvent() + | PtyExitedEvent() + | PtyDeletedEvent() + | LspUpdatedEvent() + | LspClientDiagnosticsEvent() + | ProjectUpdatedEvent() + | VcsBranchUpdatedEvent() + | TuiPromptAppendEvent() + | TuiCommandExecuteEvent() + | TuiToastShowEvent() + ): + return None # Events with properties.info.id (Session has id field) case SessionCreatedEvent(properties=props): @@ -88,28 +199,96 @@ def _extract_session_id(event: Event) -> str | None: # noqa: PLR0911 case SessionUpdatedEvent(properties=props): return props.info.id + # Events with properties.info.session_id (MessageInfo has session_id field) + case MessageUpdatedEvent(properties=props): + return props.info.session_id + # Events with properties.part.session_id (Part has session_id field) case PartUpdatedEvent(properties=props): return props.part.session_id - # Events without session_id return None case _: + logger.warning("Unhandled event type in _extract_session_id: %s", type(event).__name__) return None -def _serialize_event(event: Event, wrap_payload: bool = False) -> str: - """Serialize event, optionally wrapping in payload structure. +class GlobalEventFactory: + """Creates GlobalEvent envelope JSON from Event instances. - Uses ensure_ascii=False to preserve Unicode characters (Chinese, emoji, etc.) - in the JSON output instead of escaping them as \\uXXXX sequences. + Stored on ServerState since directory/project don't change during + the server's lifetime. Created lazily on first access. """ - event_data = event.model_dump(by_alias=True, exclude_none=True) - # Add sessionId at top level if available (for subagent session tracking) + def __init__(self, directory: str, project: str, workspace: str | None = None) -> None: + """Initialize with directory and project routing metadata. + + Args: + directory: Working directory for event routing + project: Project identifier for event routing + workspace: Workspace identifier for TUI workspace routing + """ + self._directory = directory + self._project = project + self._workspace = workspace + + def wrap(self, event: Event) -> str: + """Wrap an Event in a GlobalEvent envelope JSON string. + + Args: + event: The event to wrap + + Returns: + JSON string with directory, project, workspace, and payload keys. + """ + payload = _event_to_dict(event) + envelope: dict[str, Any] = { + "directory": self._directory, + "project": self._project, + "payload": payload, + } + if self._workspace is not None: + envelope["workspace"] = self._workspace + return json.dumps(envelope, ensure_ascii=False) + + +def _event_to_dict(event: Event) -> dict[str, Any]: + """Convert an Event to a dict with sessionId injected at top level. + + This is the dict-building half of serialization; the caller decides + whether to wrap it in a payload envelope and when to call json.dumps. + + Injects sessionId (lowercase 'd') at the top level for subagent + session tracking, separate from the alias-converted sessionID that + appears inside properties. + + Args: + event: The event to convert + + Returns: + Dict with the event data and optional sessionId field. + """ + event_data = event.model_dump(by_alias=True, exclude_none=True) session_id = _extract_session_id(event) if session_id is not None: event_data["sessionId"] = session_id + return event_data + + +def _serialize_event(event: Event, wrap_payload: bool = False) -> str: + r"""Serialize event, optionally wrapping in payload structure. + + Thin convenience wrapper around _event_to_dict + json.dumps. + Uses ensure_ascii=False to preserve Unicode characters (Chinese, emoji, etc.) + in the JSON output instead of escaping them as \uXXXX sequences. + Args: + event: The event to serialize + wrap_payload: Whether to wrap in a {"payload": ...} structure + + Returns: + JSON string of the serialized event data. + """ + event_data = _event_to_dict(event) if wrap_payload: return json.dumps({"payload": event_data}, ensure_ascii=False) return json.dumps(event_data, ensure_ascii=False) @@ -118,13 +297,40 @@ def _serialize_event(event: Event, wrap_payload: bool = False) -> str: async def _event_generator( state: ServerState, *, wrap_payload: bool = False ) -> AsyncGenerator[dict[str, Any]]: - """Generate SSE events.""" - queue: asyncio.Queue[Event] = asyncio.Queue() + """Generate SSE events for connected clients. + + Registers a subscriber queue, sends an initial connected event, + then streams subsequent events from the broadcast system. + + When wrap_payload is True, session-scoped events are wrapped in a + GlobalEvent envelope via the factory. Global server lifecycle events + still use a top-level ``payload`` wrapper, but omit directory/project + metadata to match OpenCode's `/global/event` contract. + + Subscriber lifecycle: + 1. Queue appended to state.event_subscribers + 2. If this is the first subscriber, triggers on_first_subscriber + callback (e.g., for update check) + 3. Streams events until client disconnects + 4. Finally block removes queue from subscribers (suppresses + ValueError if already removed by broadcast_event error handler) + + Args: + state: The server state holding subscribers and event factory + wrap_payload: Whether to wrap events in GlobalEvent envelopes + """ + factory = state.get_event_factory() if wrap_payload else None + queue: asyncio.Queue[Event] = asyncio.Queue(maxsize=100) state.event_subscribers.append(queue) subscriber_count = len(state.event_subscribers) logger.info("SSE: New client connected (total subscribers: %s)", subscriber_count) - # Trigger first subscriber callback if this is the first connection + # Trigger first subscriber callback if this is the first connection. + # Race condition analysis: This is safe because: + # 1. The append (line above) and len check happen in the same async frame + # (no await between them), so no other coroutine can interleave. + # 2. The _first_subscriber_triggered flag prevents double-firing even if + # a subscriber disconnects and reconnects rapidly. if ( subscriber_count == 1 and not state._first_subscriber_triggered @@ -134,19 +340,37 @@ async def _event_generator( state.create_background_task(state.on_first_subscriber(), name="on_first_subscriber") try: - # Send initial connected event + # Send initial connected event with payload wrapper on /global/event, + # but without directory/project metadata. connected = ServerConnectedEvent() data = _serialize_event(connected, wrap_payload=wrap_payload) logger.info("SSE: Sending connected event", data=data) yield {"data": data} # Stream events while True: - event = await queue.get() - data = _serialize_event(event, wrap_payload=wrap_payload) + try: + event = await asyncio.wait_for(queue.get(), timeout=10.0) + except TimeoutError: + # No events for 10s — send heartbeat to keep connection alive + heartbeat = ServerHeartbeatEvent() + data = _serialize_event(heartbeat, wrap_payload=wrap_payload) + yield {"data": data} + continue + if factory is not None and not isinstance( + event, ServerHeartbeatEvent | ServerConnectedEvent + ): + data = factory.wrap(event) + elif wrap_payload: + data = _serialize_event(event, wrap_payload=True) + else: + data = _serialize_event(event) logger.info("SSE: Sending event", event_type=event.type) yield {"data": data} finally: - state.event_subscribers.remove(queue) + # Use safe removal: broadcast_event may have already removed this queue + # due to error handling. Using discard-style pattern to avoid ValueError. + with contextlib.suppress(ValueError): + state.event_subscribers.remove(queue) logger.info("SSE: Client disconnected", remaining_subscribers=len(state.event_subscribers)) @@ -160,3 +384,36 @@ async def get_global_events(state: StateDep) -> EventSourceResponse: async def get_events(state: StateDep) -> EventSourceResponse: """Get events as SSE stream (no payload wrapper).""" return EventSourceResponse(_event_generator(state, wrap_payload=False), sep="\n") + + +@router.get("/global/routing-check", response_model=RoutingCheckResponse) +async def get_routing_check( + state: StateDep, + directory: str, + workspace: str | None = None, + current_workspace: str | None = None, + project_directory: str | None = None, +) -> RoutingCheckResponse: + """Check whether an event would pass the OpenCode TUI routing filter. + + Diagnostic endpoint that constructs a synthetic GlobalEvent with the + given directory/workspace and runs it through the 4-rule TUI event + routing filter. Returns whether the event would pass and why. + + Args: + state: Server state (injected dependency). + directory: The event's directory field. + workspace: The event's workspace field (optional). + current_workspace: The TUI's active workspace for rule 3 filtering. + project_directory: The project directory to match against + (defaults to state.base_path). + + Returns: + RoutingCheckResponse with would_pass and reason fields. + """ + effective_project_dir = project_directory if project_directory is not None else state.base_path + event = GlobalEvent(directory=directory, workspace=workspace, payload={}) + would_pass, reason = tui_event_filter( + event, effective_project_dir, current_workspace=current_workspace + ) + return RoutingCheckResponse(would_pass=would_pass, reason=reason) diff --git a/src/agentpool_server/opencode_server/routes/message_routes.py b/src/agentpool_server/opencode_server/routes/message_routes.py index 56365cadb..3dc7b765e 100644 --- a/src/agentpool_server/opencode_server/routes/message_routes.py +++ b/src/agentpool_server/opencode_server/routes/message_routes.py @@ -33,7 +33,6 @@ Part, PartRemovedEvent, PartUpdatedEvent, - SessionIdleEvent, SessionStatus, SessionStatusEvent, StepStartPart, @@ -45,6 +44,7 @@ UserMessage, ) from agentpool_server.opencode_server.routes.session_routes import get_or_load_session +from agentpool_server.opencode_server.state import QueuedAsyncPrompt from agentpool_server.opencode_server.stream_adapter import OpenCodeStreamAdapter @@ -256,7 +256,6 @@ async def _process_message( time=TimeCreated.now(), agent=request.agent or "default", model=request.model, - variant=request.variant, ) user_msg_with_parts = MessageWithParts(info=user_message) @@ -303,17 +302,26 @@ async def _process_message_locked( # noqa: PLR0915 state: StateDep, user_msg_id: str, user_msg_with_parts: MessageWithParts, + *, + mark_busy: bool = True, + mark_idle: bool = True, ) -> MessageWithParts: """Actual agent processing logic (called within lock). Args: + session_id: Session receiving the message. + request: Request payload containing the user's parts and agent/model choice. + state: Shared OpenCode server state. user_msg_id: ID of already-created user message user_msg_with_parts: The user message with parts (already broadcast) + mark_busy: Whether to emit a busy transition before processing. + mark_idle: Whether to emit an idle transition when processing completes. """ # --- Mark session busy --- - busy = SessionStatus(type="busy") - state.session_status[session_id] = busy - await state.broadcast_event(SessionStatusEvent.create(session_id, busy)) + if mark_busy: + busy = SessionStatus(type="busy") + state.session_status[session_id] = busy + await state.broadcast_event(SessionStatusEvent.create(session_id, busy)) # --- Extract user prompt --- user_prompt = await extract_user_prompt_from_parts( request.parts, @@ -352,64 +360,71 @@ async def _process_message_locked( # noqa: PLR0915 step_start = StepStartPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) assistant_msg_with_parts.parts.append(step_start) await state.broadcast_event(PartUpdatedEvent.create(step_start)) - # --- Resolve agent and variant --- - agent = state.agent - if request.agent and state.agent.agent_pool is not None: - agent = state.agent.agent_pool.all_agents.get(request.agent, state.agent) - if request.variant: - with contextlib.suppress(Exception): - await agent.set_mode(request.variant, category_id="thought_level") - - # Handle model selection if requested - original_model: str | None = None - if request.model and request.model.model_id and request.model.provider_id: - provider_id = request.model.provider_id - model_id = request.model.model_id - - # Strategy: First try to use model_id as a variant name - # OpenCode TUI sends variant names as model_id (e.g., "ack-dev", "qwen35") - # The provider_id is the first part of the identifier (e.g., "openai-chat") - requested_model = model_id # Try variant name first - - logger.info(f"Model selection requested: provider={provider_id}, model_id={model_id}") + # --- Short critical section: bind agent to session, apply mutations, then capture snapshot --- + from agentpool.agents.context import RunSnapshot - try: - available_models = await agent.get_available_models() - is_valid = False - - # Check 1: Is model_id a variant name in manifest? - if state.pool and model_id in state.pool.manifest.model_variants: - is_valid = True - logger.info(f"Model {model_id} found as variant name in manifest") - # Check 2: Is it in tokonomics models? - elif available_models: - valid_ids = [m.id_override if m.id_override else m.id for m in available_models] - # Try both "provider:model" format and just model_id - full_id = f"{provider_id}:{model_id}" - if full_id in valid_ids: - is_valid = True - requested_model = full_id - logger.info(f"Model {full_id} found in tokonomics models") - elif model_id in valid_ids: + snapshot: RunSnapshot | None = None + async with state.agent_lock: + agent = state.agent + if request.agent and state.agent.agent_pool is not None: + agent = state.agent.agent_pool.all_agents.get(request.agent, state.agent) + + # Apply variant/mode under the lock + original_variant: str | None = None + request_variant = request.model.variant if request.model else None + if request_variant: + # Save current thought_level mode before mutation so we can restore after the turn + current_variant: str | None = None + try: + modes = await agent.get_modes() + for cat in modes: + if cat.category == "thought_level" and cat.current_mode_id: + current_variant = cat.current_mode_id + break + except Exception: # noqa: BLE001 + pass + try: + await agent.set_mode(request_variant, category_id="thought_level") + original_variant = current_variant # Only save for restore after successful mutation + except ValueError: + logger.debug("Variant mode not applicable", variant=request_variant) + + # Handle model selection if requested + original_model: str | None = None + if request.model and request.model.model_id and request.model.provider_id: + provider_id = request.model.provider_id + model_id = request.model.model_id + requested_model = model_id + + logger.info(f"Model selection requested: provider={provider_id}, model_id={model_id}") + + try: + available_models = await agent.get_available_models() + is_valid = False + + if state.pool and model_id in state.pool.manifest.model_variants: is_valid = True - logger.info(f"Model {model_id} found in tokonomics models") - - if is_valid: - # Store original model to restore later - original_model = agent.model_name - logger.info(f"Switching model from {original_model} to {requested_model}") - await agent.set_model(requested_model) - logger.info("Switched to requested model", model=requested_model) - else: - logger.warning(f"Model {model_id} (provider: {provider_id}) is not valid") - if state.pool: - logger.warning( - f"Available model_variants: {list(state.pool.manifest.model_variants.keys())}" - ) - except Exception as e: # noqa: BLE001 - # Broad catch: agents differ on how they signal unsupported/invalid model switching. - # Keep behavior stable for OpenCode (see PR #10 review iterations). - logger.warning(f"Failed to switch model: {e}") + elif available_models: + valid_ids = [m.id_override if m.id_override else m.id for m in available_models] + full_id = f"{provider_id}:{model_id}" + if full_id in valid_ids: + is_valid = True + requested_model = full_id + elif model_id in valid_ids: + is_valid = True + + if is_valid: + original_model = agent.model_name + logger.info(f"Switching model from {original_model} to {requested_model}") + await agent.set_model(requested_model) + else: + logger.warning(f"Model {model_id} (provider: {provider_id}) is not valid") + except ValueError as e: + logger.warning(f"Failed to switch model: {e}") + + # Capture snapshot AFTER mutations so mode_name/model_name are current + snapshot = await state.snapshot_for_session(session_id, agent=agent) + # --- End short critical section --- # --- Stream via adapter --- adapter = OpenCodeStreamAdapter( @@ -422,9 +437,14 @@ async def _process_message_locked( # noqa: PLR0915 ) response_time: int | None = None - cancelled = False try: - iterator = agent.run_stream(*user_prompt, session_id=session_id) + iterator = agent.run_stream( + *user_prompt, + session_id=snapshot.session_id, + message_history=snapshot.conversation, + input_provider=snapshot.input_provider, + snapshot=snapshot, + ) async for oc_event in adapter.process_stream(iterator): await state.broadcast_event(oc_event) @@ -446,7 +466,6 @@ async def _process_message_locked( # noqa: PLR0915 except asyncio.CancelledError: # User cancelled the request (e.g., pressed ESC) logger.info("Request cancelled by user", session_id=session_id) - cancelled = True # Finalize the assistant message with aborted state. # This mirrors upstream OpenCode's cleanup() in processor.ts:518 @@ -465,19 +484,13 @@ async def _process_message_locked( # noqa: PLR0915 await state.broadcast_event(MessageUpdatedEvent.create(updated_assistant)) await persist_message_to_storage(state, assistant_msg_with_parts, session_id) - # Add the aborted assistant message to the agent's in-memory conversation. - # Without this, the agent's conversation.chat_messages only has the user - # message (added by _run_stream_once at base_agent.py:784) but not the - # assistant response. On the next message, get_or_load_session() skips - # reloading because agent.session_id matches, so the LLM receives - # incomplete history — it doesn't know it already (partially) responded. - # - # NOTE: This mutates the shared agent's conversation. In the current - # single-session server architecture this is safe, but if multi-session - # support is added, the agent instance should be per-session to avoid - # history contamination between concurrent sessions. - chat_msg = opencode_to_chat_message(assistant_msg_with_parts, session_id=session_id) - agent.conversation.add_chat_messages([chat_msg], extend_last=True) + # Add the aborted assistant message to the snapshot's conversation. + # Using snapshot.conversation instead of agent.conversation avoids + # cross-session contamination when multiple sessions run concurrently. + chat_msg = opencode_to_chat_message( + assistant_msg_with_parts, session_id=snapshot.session_id + ) + snapshot.conversation.add_chat_messages([chat_msg], extend_last=True) finally: # Restore original model if we changed it if original_model is not None: @@ -485,12 +498,24 @@ async def _process_message_locked( # noqa: PLR0915 await agent.set_model(original_model) logger.info("Restored original model", model=original_model) + # Restore original variant/mode if we changed it + if original_variant is not None: + with contextlib.suppress(Exception): + await agent.set_mode(original_variant, category_id="thought_level") + logger.info("Restored original variant mode", variant=original_variant) + # --- Mark session idle --- - # Always set session to idle, even if processing failed or was cancelled - status = SessionStatus(type="idle") - state.session_status[session_id] = status - await state.broadcast_event(SessionStatusEvent.create(session_id, status)) - await state.broadcast_event(SessionIdleEvent.create(session_id)) + # The async prompt worker owns session idling while it drains queued work. + if mark_idle: + has_queued = state.has_pending_async_prompts(session_id) + if has_queued: + # Queued work exists — skip idle to avoid idle→busy flicker. + # If no worker is running, start one; if one is already + # running, it will drain the queue and mark idle when done. + if not state.has_session_background_task(session_id): + await _ensure_async_prompt_worker(session_id, state, mark_busy=True) + else: + await state.mark_session_idle(session_id) # --- Update session timestamp --- if response_time is not None: session = state.sessions[session_id] @@ -502,6 +527,71 @@ async def _process_message_locked( # noqa: PLR0915 return assistant_msg_with_parts +async def _ensure_async_prompt_worker( + session_id: str, + state: StateDep, + *, + mark_busy: bool, +) -> None: + """Start the per-session async prompt worker when queued work exists.""" + if not state.has_pending_async_prompts(session_id): + return + if state.has_session_background_task(session_id): + return + + if mark_busy: + busy = SessionStatus(type="busy") + state.session_status[session_id] = busy + await state.broadcast_event(SessionStatusEvent.create(session_id, busy)) + + state.create_background_task( + _run_async_prompt_queue(session_id, state), + name=f"process_message_{session_id}", + ) + + +async def _run_async_prompt_queue(session_id: str, state: StateDep) -> None: + """Drain queued async prompts for a session in FIFO order.""" + lock = state.get_session_lock(session_id) + try: + while True: + async with lock: + queued_prompt = state.pop_next_async_prompt(session_id) + if queued_prompt is None: + await state.mark_session_idle(session_id) + return + + try: + await _process_message_locked( + session_id, + queued_prompt.request, + state, + queued_prompt.user_msg_id, + queued_prompt.user_msg_with_parts, + mark_busy=False, + mark_idle=False, + ) + except asyncio.CancelledError: + raise + except Exception: + logger.exception("Async prompt processing failed", session_id=session_id) + # Continue draining — don't kill the queue + + if state.has_pending_async_prompts(session_id): + await state.emit_session_turn_complete(session_id) + continue + + await state.mark_session_idle(session_id) + return + except asyncio.CancelledError: + logger.info("Async prompt worker cancelled", session_id=session_id) + raise + except Exception: + logger.exception("Async prompt worker failed catastrophically", session_id=session_id) + await state.mark_session_idle(session_id) + raise + + @router.post("/message") async def send_message( session_id: str, @@ -524,8 +614,8 @@ async def send_message_async(session_id: str, request: MessageRequest, state: St """Send a message asynchronously without waiting for response. Starts the agent processing in the background and returns immediately. - If the session is busy, the message is queued using agent.queue() and - will be processed after the current run completes. + If the session is busy, the message is queued in server state and + processed after the current run completes. Client should listen to SSE events to get updates. @@ -543,7 +633,6 @@ async def send_message_async(session_id: str, request: MessageRequest, state: St time=TimeCreated.now(), agent=request.agent or "default", model=request.model, - variant=request.variant, ) user_msg_with_parts = MessageWithParts(info=user_message) @@ -576,33 +665,29 @@ async def send_message_async(session_id: str, request: MessageRequest, state: St await persist_message_to_storage(state, user_msg_with_parts, session_id) await state.broadcast_event(MessageUpdatedEvent.create(user_message)) - # 2. Extract user prompt for queuing/processing - user_prompt = await extract_user_prompt_from_parts( - request.parts, - fs=state.fs, - tools=state.agent.tools, - ) - - # 3. Check if session is busy - current_status = state.session_status.get(session_id) - is_busy = current_status is not None and current_status.type == "busy" + # 2. Atomically queue work, then start a single per-session worker if needed. + lock = state.get_session_lock(session_id) + async with lock: + state.enqueue_async_prompt( + session_id, + QueuedAsyncPrompt( + request=request, + user_msg_id=user_msg_id, + user_msg_with_parts=user_msg_with_parts, + ), + ) - if is_busy: - # Session is busy → queue the prompt using agent.queue_prompt() - # The agent will automatically process queued prompts after current run - logger.info("Session busy, queuing prompt via agent.queue_prompt()", session_id=session_id) - agent = state.agent - if request.agent and state.agent.agent_pool is not None: - agent = state.agent.agent_pool.all_agents.get(request.agent, state.agent) - agent.queue_prompt(user_prompt) - return + current_status = state.session_status.get(session_id) + mark_busy = current_status is None or current_status.type != "busy" + if not mark_busy: + logger.info( + "Session became busy before async dispatch, keeping prompt in server queue", + session_id=session_id, + ) + else: + logger.info("Session idle, starting background task", session_id=session_id) - # 4. Session is idle → start background task to process - logger.info("Session idle, starting background task", session_id=session_id) - state.create_background_task( - _process_message_locked(session_id, request, state, user_msg_id, user_msg_with_parts), - name=f"process_message_{session_id}", - ) + await _ensure_async_prompt_worker(session_id, state, mark_busy=mark_busy) @router.get("/message/{message_id}") diff --git a/src/agentpool_server/opencode_server/routes/permission_routes.py b/src/agentpool_server/opencode_server/routes/permission_routes.py index 5ea2d48b1..0a7bd6e0c 100644 --- a/src/agentpool_server/opencode_server/routes/permission_routes.py +++ b/src/agentpool_server/opencode_server/routes/permission_routes.py @@ -46,7 +46,7 @@ async def reply_to_permission( # Find which session has this permission request for session_id, input_provider in state.input_providers.items(): # Check if this permission belongs to this session - if permission_id not in input_provider._pending_permissions: + if not input_provider.has_pending_permission(permission_id): continue # Resolve the permission resolved = input_provider.resolve_permission(permission_id, body.reply) diff --git a/src/agentpool_server/opencode_server/routes/question_routes.py b/src/agentpool_server/opencode_server/routes/question_routes.py index 773021610..788f40334 100644 --- a/src/agentpool_server/opencode_server/routes/question_routes.py +++ b/src/agentpool_server/opencode_server/routes/question_routes.py @@ -7,6 +7,7 @@ from agentpool_server.opencode_server.dependencies import StateDep from agentpool_server.opencode_server.input_provider import OpenCodeInputProvider from agentpool_server.opencode_server.models import ( + PermissionResolvedEvent, QuestionRejectedEvent, QuestionRepliedEvent, QuestionReply, @@ -17,6 +18,31 @@ router = APIRouter(prefix="/question", tags=["question"]) +def _find_permission_provider( + state: StateDep, + permission_id: str, +) -> tuple[str, OpenCodeInputProvider] | None: + for session_id, input_provider in state.input_providers.items(): + if input_provider.has_pending_permission(permission_id): + return session_id, input_provider + return None + + +def _extract_permission_reply(reply: QuestionReply) -> str | None: + if len(reply.answers) != 1: + return None + selected_answers = reply.answers[0] + if len(selected_answers) != 1: + return None + + selected_reply = selected_answers[0] + match selected_reply: + case "once" | "always" | "reject": + return selected_reply + case _: + return None + + @router.get("/", response_model=list[QuestionRequest]) async def list_questions(state: StateDep) -> list[QuestionRequest]: """List all pending question requests. @@ -50,7 +76,26 @@ async def reply_to_question(requestID: str, reply: QuestionReply, state: StateDe """ pending = state.pending_questions.get(requestID) if not pending: - raise HTTPException(status_code=404, detail="Question request not found") + permission_target = _find_permission_provider(state, requestID) + if permission_target is None: + raise HTTPException(status_code=404, detail="Question request not found") + + session_id, provider = permission_target + permission_reply = _extract_permission_reply(reply) + if permission_reply is None: + raise HTTPException(status_code=400, detail="Invalid permission reply") + + if not provider.resolve_permission(requestID, permission_reply): + raise HTTPException(status_code=404, detail="Permission not found or already resolved") + + event = PermissionResolvedEvent.create( + session_id=session_id, + request_id=requestID, + reply=permission_reply, + ) + await state.broadcast_event(event) + return True + session_id = pending.session_id provider = state.input_providers.get(session_id) if not isinstance(provider, OpenCodeInputProvider): @@ -86,7 +131,21 @@ async def reject_question(requestID: str, state: StateDep) -> bool: # noqa: N80 """ pending = state.pending_questions.get(requestID) if not pending: - raise HTTPException(status_code=404, detail="Question request not found") + permission_target = _find_permission_provider(state, requestID) + if permission_target is None: + raise HTTPException(status_code=404, detail="Question request not found") + + session_id, provider = permission_target + if not provider.resolve_permission(requestID, "reject"): + raise HTTPException(status_code=404, detail="Permission not found or already resolved") + + event = PermissionResolvedEvent.create( + session_id=session_id, + request_id=requestID, + reply="reject", + ) + await state.broadcast_event(event) + return True # Cancel the future if not pending.future.done(): pending.future.cancel() diff --git a/src/agentpool_server/opencode_server/routes/routing.py b/src/agentpool_server/opencode_server/routes/routing.py new file mode 100644 index 000000000..3da93f8bc --- /dev/null +++ b/src/agentpool_server/opencode_server/routes/routing.py @@ -0,0 +1,80 @@ +"""TUI event routing filter — diagnostic re-implementation of OpenCode TUI routing. + +Re-implements the 4-rule event routing filter from the OpenCode TUI's +event.ts as a pure Python function for diagnostic purposes. This allows +debugging why events may be dropped by the TUI's client-side filter. + +The 4 rules (evaluated in order): +1. **Sync events always dropped**: if `payload.type == "sync"` → False +2. **Global directory always passes**: if `directory == "global"` → True +3. **Workspace filtering (if active)**: if `current_workspace` is set → + `event.workspace == current_workspace` +4. **Directory must match exactly**: `event.directory == project_directory` +""" + +from __future__ import annotations + +from typing import Literal + +from agentpool_server.opencode_server.models.base import OpenCodeBaseModel +from agentpool_server.opencode_server.models.events import GlobalEvent + + +RoutingReason = Literal[ + "sync_dropped", + "global_directory", + "workspace_match", + "workspace_mismatch", + "directory_match", + "directory_mismatch", +] +"""Reason strings explaining why an event passes or fails the routing filter.""" + + +class RoutingCheckResponse(OpenCodeBaseModel): + """Response for GET /global/routing-check endpoint.""" + + would_pass: bool + """Whether the event would pass the TUI routing filter.""" + + reason: RoutingReason + """Explanation for the routing decision.""" + + +def tui_event_filter( + event: GlobalEvent, + project_directory: str, + current_workspace: str | None = None, +) -> tuple[bool, RoutingReason]: + """Re-implements OpenCode TUI event routing filter for diagnostic purposes. + + Evaluates the 4-rule filter in priority order and returns both + the pass/fail result and the reason for the decision. + + Args: + event: The GlobalEvent to check against the routing filter. + project_directory: The server's working directory for directory matching. + current_workspace: The TUI's active workspace, or None if not set. + + Returns: + A tuple of (would_pass, reason) where reason explains which rule + determined the outcome. + """ + # Rule 1: sync events always dropped + if event.payload.get("type") == "sync": + return (False, "sync_dropped") + + # Rule 2: global directory always passes (except sync, handled above) + if event.directory == "global": + return (True, "global_directory") + + # Rule 3: workspace filtering (if workspace active) + if current_workspace is not None: + if event.workspace == current_workspace: + return (True, "workspace_match") + return (False, "workspace_mismatch") + + # Rule 4: directory must match exactly (string comparison, no normalization) + if event.directory == project_directory: + return (True, "directory_match") + return (False, "directory_mismatch") diff --git a/src/agentpool_server/opencode_server/routes/session_routes.py b/src/agentpool_server/opencode_server/routes/session_routes.py index 21db3dc13..04faa6793 100644 --- a/src/agentpool_server/opencode_server/routes/session_routes.py +++ b/src/agentpool_server/opencode_server/routes/session_routes.py @@ -44,6 +44,7 @@ SessionDeletedEvent, SessionDiffEvent, SessionForkRequest, + SessionIdleEvent, SessionInitRequest, SessionRevert, SessionShare, @@ -217,102 +218,110 @@ async def _execute_slashed_command( state.messages[session_id].append(message_with_parts) await state.broadcast_event(MessageUpdatedEvent.create(assistant_message)) - # Mark session as busy - state.session_status[session_id] = SessionStatus(type="busy") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="busy"))) - - # Add step-start part to indicate command is running - part_id = identifier.ascending("part") - step_start = StepStartPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) - message_with_parts.parts.append(step_start) - await state.broadcast_event(PartUpdatedEvent.create(step_start)) - - # Parse arguments - args = request.arguments.split() if request.arguments else [] - - # Create command context with output capture - output_capture = _CommandOutputCapture() - cmd_ctx = CommandContext( - output=output_capture, - data=state.agent.get_context(), - command_store=state.command_store, - ) - - # Execute command try: - await command.execute(cmd_ctx, args, {}) - except Exception as e: - # Mark session as idle before raising - state.session_status[session_id] = SessionStatus(type="idle") + # Mark session as busy + state.session_status[session_id] = SessionStatus(type="busy") await state.broadcast_event( - SessionStatusEvent.create(session_id, SessionStatus(type="idle")) + SessionStatusEvent.create(session_id, SessionStatus(type="busy")) ) - raise HTTPException(status_code=500, detail=f"Command execution failed: {e}") from e - # Get command output - output_text = str(output_capture) if output_capture else "Command executed" + # Add step-start part to indicate command is running + part_id = identifier.ascending("part") + step_start = StepStartPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) + message_with_parts.parts.append(step_start) + await state.broadcast_event(PartUpdatedEvent.create(step_start)) + + # Parse arguments + args = request.arguments.split() if request.arguments else [] + + # Create command context with output capture + output_capture = _CommandOutputCapture() + cmd_ctx = CommandContext( + output=output_capture, + data=state.agent.get_context(), + command_store=state.command_store, + ) - # Create text part with output - text_part = TextPart( - id=identifier.ascending("part"), - message_id=assistant_msg_id, - session_id=session_id, - text=output_text, - ) - message_with_parts.parts.append(text_part) - await state.broadcast_event(PartUpdatedEvent.create(text_part)) + # Execute command + try: + await command.execute(cmd_ctx, args, {}) + except Exception as e: + raise HTTPException(status_code=500, detail=f"Command execution failed: {e}") from e - # Run agent to process the loaded skill context - try: - # Create adapter to stream agent events through existing text_part - adapter = OpenCodeStreamAdapter( - state=state, - session_id=session_id, - assistant_msg_id=assistant_msg_id, - assistant_msg=message_with_parts, - working_dir=state.working_dir, - ) - # Build prompt including user arguments - user_request = request.arguments if request.arguments else "请使用已加载的 skill context" - agent_prompt = f"用户执行了命令 '{request.command}' 并说: {user_request}\n\n请使用已加载的 skill context 来回答用户的请求。" + # Get command output + output_text = str(output_capture) if output_capture else "Command executed" - # Run agent with prompt to use the skill context - iterator = state.agent.run_stream( - agent_prompt, + # Create text part with output + text_part = TextPart( + id=identifier.ascending("part"), + message_id=assistant_msg_id, session_id=session_id, + text=output_text, ) - async for oc_event in adapter.process_stream(iterator): - await state.broadcast_event(oc_event) - # Append adapter's response to text_part - if adapter.response_text: - text_part.text = f"{output_text}\n\n{adapter.response_text}" - await state.broadcast_event(PartUpdatedEvent.create(text_part)) - except Exception: # noqa: BLE001 - # Command already executed, ignore agent errors - pass + message_with_parts.parts.append(text_part) + await state.broadcast_event(PartUpdatedEvent.create(text_part)) - # Add step-finish part to indicate command completed - step_finish = StepFinishPart( - id=identifier.ascending("part"), - message_id=assistant_msg_id, - session_id=session_id, - ) - message_with_parts.parts.append(step_finish) - await state.broadcast_event(PartUpdatedEvent.create(step_finish)) + # Run agent to process the loaded skill context + try: + # Create adapter to stream agent events through existing text_part + adapter = OpenCodeStreamAdapter( + state=state, + session_id=session_id, + assistant_msg_id=assistant_msg_id, + assistant_msg=message_with_parts, + working_dir=state.working_dir, + ) + # Build prompt including user arguments + user_request = ( + request.arguments if request.arguments else "请使用已加载的 skill context" + ) + agent_prompt = ( + f"用户执行了命令 '{request.command}' 并说: {user_request}\n\n" + "请使用已加载的 skill context 来回答用户的请求。" + ) - # Mark session as idle - state.session_status[session_id] = SessionStatus(type="idle") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="idle"))) + # Capture snapshot under short agent_lock critical section + async with state.agent_lock: + snapshot = await state.snapshot_for_session(session_id) + + # Run agent with prompt to use the skill context + iterator = state.agent.run_stream( + agent_prompt, + session_id=snapshot.session_id, + message_history=snapshot.conversation, + input_provider=snapshot.input_provider, + snapshot=snapshot, + ) + async for oc_event in adapter.process_stream(iterator): + await state.broadcast_event(oc_event) + # Append adapter's response to text_part + if adapter.response_text: + text_part.text = f"{output_text}\n\n{adapter.response_text}" + await state.broadcast_event(PartUpdatedEvent.create(text_part)) + except Exception: # noqa: BLE001 + # Command already executed, ignore agent errors + pass - # Broadcast command.executed event - await state.broadcast_event( - CommandExecutedEvent.create( - name=request.command, - session_id=session_id, - arguments=request.arguments or "", + # Add step-finish part to indicate command completed + step_finish = StepFinishPart( + id=identifier.ascending("part"), message_id=assistant_msg_id, + session_id=session_id, ) - ) + message_with_parts.parts.append(step_finish) + await state.broadcast_event(PartUpdatedEvent.create(step_finish)) + + # Broadcast command.executed event + await state.broadcast_event( + CommandExecutedEvent.create( + name=request.command, + session_id=session_id, + arguments=request.arguments or "", + message_id=assistant_msg_id, + ) + ) + finally: + await state.mark_session_idle(session_id) return message_with_parts @@ -378,112 +387,122 @@ async def _execute_skill_command( {args} """ - # Mark session as busy - state.session_status[session_id] = SessionStatus(type="busy") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="busy"))) - - # Load session into agent to ensure conversation history is restored - # This ensures agent sees all previous messages during this run - await state.agent.load_session(session_id) - - # Create USER message (not assistant!) - user_msg_id = identifier.ascending("message") - user_message = UserMessage( - id=user_msg_id, - session_id=session_id, - role="user", - time=TimeCreated.now(), - agent=request.agent or "default", - ) - user_part_id = identifier.ascending("part") - user_msg_with_parts = MessageWithParts( - info=user_message, - parts=[ - TextPart( - id=user_part_id, message_id=user_msg_id, session_id=session_id, text=user_prompt - ) - ], - ) - - # Store and broadcast user message - state.messages[session_id].append(user_msg_with_parts) - await state.broadcast_event(PartUpdatedEvent.create(user_msg_with_parts.parts[0])) - await state.broadcast_event(MessageUpdatedEvent.create(user_message)) + try: + # Mark session as busy + state.session_status[session_id] = SessionStatus(type="busy") + await state.broadcast_event( + SessionStatusEvent.create(session_id, SessionStatus(type="busy")) + ) - # Create assistant message (for response) - assistant_msg_id = identifier.ascending("message") - assistant_message = AssistantMessage( - id=assistant_msg_id, - session_id=session_id, - parent_id=user_msg_id, - model_id=request.model or "default", - provider_id="opencode", - mode="command", - agent=request.agent or "default", - path=MessagePath(cwd=state.working_dir, root=state.working_dir), - time=MessageTime(created=now_ms()), - ) - message_with_parts = MessageWithParts(info=assistant_message, parts=[]) - state.messages[session_id].append(message_with_parts) - await state.broadcast_event(MessageUpdatedEvent.create(assistant_message)) + # Load session and capture snapshot under short agent_lock critical section. + # This ensures load_session (which mutates agent.conversation) and the + # snapshot capture see a consistent agent state. + async with state.agent_lock: + await state.agent.load_session(session_id) + snapshot = await state.snapshot_for_session(session_id) + + # Create USER message (not assistant!) + user_msg_id = identifier.ascending("message") + user_message = UserMessage( + id=user_msg_id, + session_id=session_id, + role="user", + time=TimeCreated.now(), + agent=request.agent or "default", + ) + user_part_id = identifier.ascending("part") + user_msg_with_parts = MessageWithParts( + info=user_message, + parts=[ + TextPart( + id=user_part_id, message_id=user_msg_id, session_id=session_id, text=user_prompt + ) + ], + ) - # Add step-start part - step_start = StepStartPart( - id=identifier.ascending("part"), - message_id=assistant_msg_id, - session_id=session_id, - ) - message_with_parts.parts.append(step_start) - await state.broadcast_event(PartUpdatedEvent.create(step_start)) + # Store and broadcast user message + state.messages[session_id].append(user_msg_with_parts) + await state.broadcast_event(PartUpdatedEvent.create(user_msg_with_parts.parts[0])) + await state.broadcast_event(MessageUpdatedEvent.create(user_message)) - # Run agent with the user message context - try: - adapter = OpenCodeStreamAdapter( - state=state, + # Create assistant message (for response) + assistant_msg_id = identifier.ascending("message") + assistant_message = AssistantMessage( + id=assistant_msg_id, session_id=session_id, - assistant_msg_id=assistant_msg_id, - assistant_msg=message_with_parts, - working_dir=state.working_dir, + parent_id=user_msg_id, + model_id=request.model or "default", + provider_id="opencode", + mode="command", + agent=request.agent or "default", + path=MessagePath(cwd=state.working_dir, root=state.working_dir), + time=MessageTime(created=now_ms()), ) + message_with_parts = MessageWithParts(info=assistant_message, parts=[]) + state.messages[session_id].append(message_with_parts) + await state.broadcast_event(MessageUpdatedEvent.create(assistant_message)) - # Run agent with the user prompt - iterator = state.agent.run_stream(user_prompt, session_id=session_id) - async for oc_event in adapter.process_stream(iterator): - await state.broadcast_event(oc_event) - - except Exception as e: - error_text = f"Error: {e}" - text_part = TextPart( + # Add step-start part + step_start = StepStartPart( id=identifier.ascending("part"), message_id=assistant_msg_id, session_id=session_id, - text=error_text, ) - message_with_parts.parts.append(text_part) - await state.broadcast_event(PartUpdatedEvent.create(text_part)) + message_with_parts.parts.append(step_start) + await state.broadcast_event(PartUpdatedEvent.create(step_start)) - # Add step-finish part - step_finish = StepFinishPart( - id=identifier.ascending("part"), - message_id=assistant_msg_id, - session_id=session_id, - ) - message_with_parts.parts.append(step_finish) - await state.broadcast_event(PartUpdatedEvent.create(step_finish)) + # Run agent with the user message context + try: + adapter = OpenCodeStreamAdapter( + state=state, + session_id=session_id, + assistant_msg_id=assistant_msg_id, + assistant_msg=message_with_parts, + working_dir=state.working_dir, + ) - # Mark session as idle - state.session_status[session_id] = SessionStatus(type="idle") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="idle"))) + # Run agent with the user prompt using snapshot + iterator = state.agent.run_stream( + user_prompt, + session_id=snapshot.session_id, + message_history=snapshot.conversation, + input_provider=snapshot.input_provider, + snapshot=snapshot, + ) + async for oc_event in adapter.process_stream(iterator): + await state.broadcast_event(oc_event) + + except Exception as e: + error_text = f"Error: {e}" + text_part = TextPart( + id=identifier.ascending("part"), + message_id=assistant_msg_id, + session_id=session_id, + text=error_text, + ) + message_with_parts.parts.append(text_part) + await state.broadcast_event(PartUpdatedEvent.create(text_part)) - # Broadcast command.executed event - await state.broadcast_event( - CommandExecutedEvent.create( - name=request.command, - session_id=session_id, - arguments=request.arguments or "", + # Add step-finish part + step_finish = StepFinishPart( + id=identifier.ascending("part"), message_id=assistant_msg_id, + session_id=session_id, ) - ) + message_with_parts.parts.append(step_finish) + await state.broadcast_event(PartUpdatedEvent.create(step_finish)) + + # Broadcast command.executed event + await state.broadcast_event( + CommandExecutedEvent.create( + name=request.command, + session_id=session_id, + arguments=request.arguments or "", + message_id=assistant_msg_id, + ) + ) + finally: + await state.mark_session_idle(session_id) return message_with_parts @@ -545,9 +564,10 @@ async def get_or_load_session(state: ServerState, session_id: str) -> Session | session = session_data_to_opencode(data) # Cache the session state.sessions[session_id] = session + state.ensure_runtime_session_state(session_id) # Initialize runtime state if session_id not in state.session_status: - state.session_status[session_id] = SessionStatus(type="idle") + await state.mark_session_idle(session_id) # For subagent sessions with existing in-memory messages, preserve them # Subagent messages are streamed in real-time and may not be persisted yet @@ -573,10 +593,11 @@ async def get_or_load_session(state: ServerState, session_id: str) -> Session | if session_id not in state.input_providers: input_provider = OpenCodeInputProvider(state, session_id) state.input_providers[session_id] = input_provider - # Set input provider on agent to ensure correct session routing - state.agent._input_provider = state.input_providers[session_id] - # Update agent's session_id to track which session is loaded - state.agent.session_id = session_id + # Bind agent to this session under agent_lock so that a concurrent turn's + # snapshot capture sees a consistent (input_provider, session_id) pair. + async with state.agent_lock: + state.agent._input_provider = state.input_providers[session_id] + state.agent.session_id = session_id return session @@ -653,19 +674,20 @@ async def create_session(state: StateDep, request: SessionCreateRequest | None = # Cache in memory state.sessions[session_id] = session state.messages[session_id] = [] - state.session_status[session_id] = SessionStatus(type="idle") + await state.mark_session_idle(session_id) state.todos[session_id] = [] # Create input provider for this session input_provider = OpenCodeInputProvider(state, session_id) state.input_providers[session_id] = input_provider - # Set input provider on agent - state.agent._input_provider = input_provider - # Clear agent's conversation for the new session - # Agent is shared across sessions, so we need to clear its conversation state - if hasattr(state.agent, "conversation") and state.agent.conversation: - state.agent.conversation.chat_messages.clear() - # Update agent's session_id to the new session - state.agent.session_id = session_id + # Bind agent to this session under lock to avoid racing with snapshot_for_session() + async with state.agent_lock: + state.agent._input_provider = input_provider + # Clear agent's conversation for the new session + # Agent is shared across sessions, so we need to clear its conversation state + if hasattr(state.agent, "conversation") and state.agent.conversation: + state.agent.conversation.chat_messages.clear() + # Update agent's session_id to the new session + state.agent.session_id = session_id await state.broadcast_event(SessionCreatedEvent.create(session)) return session @@ -817,9 +839,15 @@ async def abort_session(session_id: str, state: StateDep) -> bool: if session is None: raise HTTPException(status_code=404, detail="Session not found") - # Interrupt the agent to cancel any ongoing stream + # Stop any in-flight prompt worker for this session before interrupting the agent. + await state.cancel_session_background_tasks(session_id) + + # Cancel the specific session's run without affecting other sessions + state.cancel_session_run(session_id) + + # Interrupt the agent with session-scoped targeting try: - await state.agent.interrupt() + await state.agent.interrupt(session_id=session_id) # Give a moment for the cancellation to propagate await asyncio.sleep(0.1) except Exception: # noqa: BLE001 @@ -828,6 +856,7 @@ async def abort_session(session_id: str, state: StateDep) -> bool: # Update and broadcast session status to notify clients state.session_status[session_id] = SessionStatus(type="idle") await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="idle"))) + await state.broadcast_event(SessionIdleEvent.create(session_id)) return True @@ -901,7 +930,7 @@ async def fork_session( # noqa: D417 await state.pool.sessions.store.save(session_data) # Cache in memory state.sessions[new_session_id] = forked_session - state.session_status[new_session_id] = SessionStatus(type="idle") + await state.mark_session_idle(new_session_id) state.todos[new_session_id] = [] # Copy messages to the new session (with updated session_id references) copied_messages: list[MessageWithParts] = [] @@ -1007,15 +1036,38 @@ async def init_session( # noqa: D417 # Agent doesn't support model selection, ignore pass + # Save current thought_level mode to restore after run + original_variant: str | None = None + try: + modes = await agent.get_modes() + for cat in modes: + if cat.category == "thought_level" and cat.current_mode_id: + original_variant = cat.current_mode_id + break + except Exception: # noqa: BLE001 + pass + # Run the agent in the background async def run_init() -> None: try: - await agent.run(init_prompt) + async with state.agent_lock: + snapshot = await state.snapshot_for_session(session_id) + await agent.run( + init_prompt, + session_id=snapshot.session_id, + message_history=snapshot.conversation, + input_provider=snapshot.input_provider, + snapshot=snapshot, + ) finally: # Restore original model if we changed it if original_model is not None: with contextlib.suppress(Exception): await agent.set_model(original_model) + # Restore original variant/mode if it was saved + if original_variant is not None: + with contextlib.suppress(Exception): + await agent.set_mode(original_variant, category_id="thought_level") state.create_background_task(run_init(), name=f"init_{session_id}") @@ -1096,49 +1148,51 @@ async def run_shell_command( state.messages[session_id].append(assistant_msg_with_parts) # Broadcast message created await state.broadcast_event(MessageUpdatedEvent.create(assistant_message)) - # Mark session as busy - state.session_status[session_id] = SessionStatus(type="busy") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="busy"))) - # Add step-start part - part_id = identifier.ascending("part") - step_start = StepStartPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) - assistant_msg_with_parts.parts.append(step_start) - await state.broadcast_event(PartUpdatedEvent.create(step_start)) - # Execute the command - output_text = "" - success = False try: - result = await state.agent.env.execute_command(request.command) - success = result.success - if success: - output_text = str(result.result) if result.result else "" - else: - output_text = f"Error: {result.error}" if result.error else "Command failed" - except Exception as e: # noqa: BLE001 - output_text = f"Error executing command: {e}" - - response_time = now_ms() - # Create text part with output - text_part = TextPart( - id=identifier.ascending("part"), - message_id=assistant_msg_id, - session_id=session_id, - text=f"$ {request.command}\n{output_text}", - ) - assistant_msg_with_parts.parts.append(text_part) - await state.broadcast_event(PartUpdatedEvent.create(text_part)) - part_id = identifier.ascending("part") - step_finish = StepFinishPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) - assistant_msg_with_parts.parts.append(step_finish) - await state.broadcast_event(PartUpdatedEvent.create(step_finish)) - # Update message with completion time - time_ = MessageTime(created=now, completed=response_time) - updated_assistant = assistant_message.model_copy(update={"time": time_}) - assistant_msg_with_parts.info = updated_assistant - await state.broadcast_event(MessageUpdatedEvent.create(updated_assistant)) - # Mark session as idle - state.session_status[session_id] = SessionStatus(type="idle") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="idle"))) + # Mark session as busy + state.session_status[session_id] = SessionStatus(type="busy") + await state.broadcast_event( + SessionStatusEvent.create(session_id, SessionStatus(type="busy")) + ) + # Add step-start part + part_id = identifier.ascending("part") + step_start = StepStartPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) + assistant_msg_with_parts.parts.append(step_start) + await state.broadcast_event(PartUpdatedEvent.create(step_start)) + # Execute the command + output_text = "" + success = False + try: + result = await state.agent.env.execute_command(request.command) + success = result.success + if success: + output_text = str(result.result) if result.result else "" + else: + output_text = f"Error: {result.error}" if result.error else "Command failed" + except Exception as e: # noqa: BLE001 + output_text = f"Error executing command: {e}" + + response_time = now_ms() + # Create text part with output + text_part = TextPart( + id=identifier.ascending("part"), + message_id=assistant_msg_id, + session_id=session_id, + text=f"$ {request.command}\n{output_text}", + ) + assistant_msg_with_parts.parts.append(text_part) + await state.broadcast_event(PartUpdatedEvent.create(text_part)) + part_id = identifier.ascending("part") + step_finish = StepFinishPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) + assistant_msg_with_parts.parts.append(step_finish) + await state.broadcast_event(PartUpdatedEvent.create(step_finish)) + # Update message with completion time + time_ = MessageTime(created=now, completed=response_time) + updated_assistant = assistant_message.model_copy(update={"time": time_}) + assistant_msg_with_parts.info = updated_assistant + await state.broadcast_event(MessageUpdatedEvent.create(updated_assistant)) + finally: + await state.mark_session_idle(session_id) return assistant_msg_with_parts @@ -1254,132 +1308,144 @@ async def summarize_session( # noqa: PLR0915 state.messages[session_id].append(assistant_msg_with_parts) # Broadcast message created await state.broadcast_event(MessageUpdatedEvent.create(assistant_message)) - # Mark session as busy - state.session_status[session_id] = SessionStatus(type="busy") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="busy"))) - # Add step-start part - part_id = identifier.ascending("part") - step_start = StepStartPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) - assistant_msg_with_parts.parts.append(step_start) - await state.broadcast_event(PartUpdatedEvent.create(step_start)) - # Step 1: Stream LLM summary generation FIRST (while we have full history) - # The LLM sees the complete conversation and generates a continuation prompt. - response_text = "" - usage = None - cost = 0.0 - text_part: TextPart | None = None try: - # Stream events from the agent with the summarization prompt - # This runs with FULL history - the summary is based on complete context - async for event in state.agent.run_stream(SUMMARIZE_PROMPT): - match event: - # Text streaming start - case PartStartEvent(part=PydanticTextPart(content=delta)): - response_text = delta - text_part = TextPart( - id=identifier.ascending("part"), - message_id=assistant_msg_id, - session_id=session_id, - text=delta, - ) - assistant_msg_with_parts.parts.append(text_part) - await state.broadcast_event(PartUpdatedEvent.create(text_part)) - - # Text streaming delta - case PydanticPartDeltaEvent(delta=TextPartDelta(content_delta=delta)) if delta: - response_text += delta - if text_part is not None: + # Mark session as busy + state.session_status[session_id] = SessionStatus(type="busy") + await state.broadcast_event( + SessionStatusEvent.create(session_id, SessionStatus(type="busy")) + ) + # Add step-start part + part_id = identifier.ascending("part") + step_start = StepStartPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) + assistant_msg_with_parts.parts.append(step_start) + await state.broadcast_event(PartUpdatedEvent.create(step_start)) + # Step 1: Stream LLM summary generation FIRST (while we have full history) + # The LLM sees the complete conversation and generates a continuation prompt. + response_text = "" + usage = None + cost = 0.0 + text_part: TextPart | None = None + try: + # Capture snapshot under short agent_lock critical section + async with state.agent_lock: + snapshot = await state.snapshot_for_session(session_id) + + # Stream events from the agent with the summarization prompt + # This runs with FULL history - the summary is based on complete context + async for event in state.agent.run_stream( + SUMMARIZE_PROMPT, + session_id=snapshot.session_id, + message_history=snapshot.conversation, + input_provider=snapshot.input_provider, + snapshot=snapshot, + ): + match event: + # Text streaming start + case PartStartEvent(part=PydanticTextPart(content=delta)): + response_text = delta text_part = TextPart( - id=text_part.id, + id=identifier.ascending("part"), message_id=assistant_msg_id, session_id=session_id, - text=response_text, + text=delta, ) - # Update in parts list - for i, p in enumerate(assistant_msg_with_parts.parts): - if isinstance(p, TextPart) and p.id == text_part.id: - assistant_msg_with_parts.parts[i] = text_part - break - await state.broadcast_event( - PartDeltaEvent.create( - session_id=session_id, + assistant_msg_with_parts.parts.append(text_part) + await state.broadcast_event(PartUpdatedEvent.create(text_part)) + + # Text streaming delta + case PydanticPartDeltaEvent(delta=TextPartDelta(content_delta=delta)) if delta: + response_text += delta + if text_part is not None: + text_part = TextPart( + id=text_part.id, message_id=assistant_msg_id, - part_id=text_part.id, - delta=delta, + session_id=session_id, + text=response_text, + ) + # Update in parts list + for i, p in enumerate(assistant_msg_with_parts.parts): + if isinstance(p, TextPart) and p.id == text_part.id: + assistant_msg_with_parts.parts[i] = text_part + break + await state.broadcast_event( + PartDeltaEvent.create( + session_id=session_id, + message_id=assistant_msg_id, + part_id=text_part.id, + delta=delta, + ) ) - ) - # Stream complete - extract token usage - case StreamCompleteEvent(message=msg) if msg and msg.usage: - usage = msg.usage - cost = float(msg.cost_info.total_cost) if msg.cost_info else 0 + # Stream complete - extract token usage + case StreamCompleteEvent(message=msg) if msg and msg.usage: + usage = msg.usage + cost = float(msg.cost_info.total_cost) if msg.cost_info else 0 - except Exception as e: # noqa: BLE001 - response_text = f"Error generating summary: {e}" + except Exception as e: # noqa: BLE001 + response_text = f"Error generating summary: {e}" - response_time = now_ms() - # Create/update text part with final response - if text_part is None: - text_part = TextPart( + response_time = now_ms() + # Create/update text part with final response + if text_part is None: + text_part = TextPart( + id=identifier.ascending("part"), + message_id=assistant_msg_id, + session_id=session_id, + text=response_text, + ) + assistant_msg_with_parts.parts.append(text_part) + await state.broadcast_event(PartUpdatedEvent.create(text_part)) + + # Step 2: Run compaction pipeline AFTER summary is generated + # The summary was generated with full context. Now we compact the history. + # Final state will be: [compacted history] + [summary message] + # The compacted history becomes the cached prefix for future LLM calls. + try: + # Get the compaction pipeline from the agent pool configuration + pipeline = None + if state.agent.agent_pool is not None: + pipeline = state.agent.agent_pool.compaction_pipeline + if pipeline is None: + # Fall back to a default summarizing pipeline + pipeline = summarizing_context() + + # Apply the compaction pipeline (modifies snapshot.conversation in place) + await compact_conversation(pipeline, snapshot.conversation) + # Persist compacted messages to storage, replacing the old ones + if state.storage is not None: + compacted_history = snapshot.conversation.get_history() + await state.storage.replace_conversation_messages(session_id, compacted_history) + # Update in-memory OpenCode messages list with compacted versions + # Keep only the summary message we just created + state.messages[session_id] = [assistant_msg_with_parts] + + except Exception: # noqa: BLE001 + # Compaction failure is not fatal - we still have the summary + pass + tokens = Tokens.from_pydantic_ai(usage) if usage else Tokens() + # Add step-finish part + step_finish = StepFinishPart( id=identifier.ascending("part"), message_id=assistant_msg_id, session_id=session_id, - text=response_text, + tokens=tokens, + cost=cost, ) - assistant_msg_with_parts.parts.append(text_part) - await state.broadcast_event(PartUpdatedEvent.create(text_part)) - - # Step 2: Run compaction pipeline AFTER summary is generated - # The summary was generated with full context. Now we compact the history. - # Final state will be: [compacted history] + [summary message] - # The compacted history becomes the cached prefix for future LLM calls. - try: - # Get the compaction pipeline from the agent pool configuration - pipeline = None - if state.agent.agent_pool is not None: - pipeline = state.agent.agent_pool.compaction_pipeline - if pipeline is None: - # Fall back to a default summarizing pipeline - pipeline = summarizing_context() - - # Apply the compaction pipeline (modifies agent.conversation in place) - await compact_conversation(pipeline, state.agent.conversation) - # Persist compacted messages to storage, replacing the old ones - if state.storage is not None: - compacted_history = state.agent.conversation.get_history() - await state.storage.replace_conversation_messages(session_id, compacted_history) - # Update in-memory OpenCode messages list with compacted versions - # Keep only the summary message we just created - state.messages[session_id] = [assistant_msg_with_parts] - - except Exception: # noqa: BLE001 - # Compaction failure is not fatal - we still have the summary - pass - tokens = Tokens.from_pydantic_ai(usage) if usage else Tokens() - # Add step-finish part - step_finish = StepFinishPart( - id=identifier.ascending("part"), - message_id=assistant_msg_id, - session_id=session_id, - tokens=tokens, - cost=cost, - ) - assistant_msg_with_parts.parts.append(step_finish) - await state.broadcast_event(PartUpdatedEvent.create(step_finish)) - # Update message with completion time and tokens - msg_time = MessageTime(created=now, completed=response_time) - update = {"time": msg_time, "tokens": tokens, "cost": cost} - updated_assistant = assistant_message.model_copy(update=update) - assistant_msg_with_parts.info = updated_assistant - await state.broadcast_event(MessageUpdatedEvent.create(updated_assistant)) - # Mark session as idle - state.session_status[session_id] = SessionStatus(type="idle") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="idle"))) - - # Broadcast session.diff event after summarization - file_ops = state.pool.file_ops - diffs = [FileDiff.from_file_change(change) for change in file_ops.changes] - await state.broadcast_event(SessionDiffEvent.create(session_id, diffs)) + assistant_msg_with_parts.parts.append(step_finish) + await state.broadcast_event(PartUpdatedEvent.create(step_finish)) + # Update message with completion time and tokens + msg_time = MessageTime(created=now, completed=response_time) + update = {"time": msg_time, "tokens": tokens, "cost": cost} + updated_assistant = assistant_message.model_copy(update=update) + assistant_msg_with_parts.info = updated_assistant + await state.broadcast_event(MessageUpdatedEvent.create(updated_assistant)) + + # Broadcast session.diff event after summarization + file_ops = state.pool.file_ops + diffs = [FileDiff.from_file_change(change) for change in file_ops.changes] + await state.broadcast_event(SessionDiffEvent.create(session_id, diffs)) + finally: + await state.mark_session_idle(session_id) return assistant_msg_with_parts @@ -1674,65 +1740,67 @@ async def execute_command( # noqa: PLR0915 assistant_msg_with_parts = MessageWithParts(info=assistant_message, parts=[]) state.messages[session_id].append(assistant_msg_with_parts) await state.broadcast_event(MessageUpdatedEvent.create(assistant_message)) - # Mark session as busy - state.session_status[session_id] = SessionStatus(type="busy") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="busy"))) - # Add step-start part - part_id = identifier.ascending("part") - step_start = StepStartPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) - assistant_msg_with_parts.parts.append(step_start) - await state.broadcast_event(PartUpdatedEvent.create(step_start)) - - # Get prompt content and execute through the agent try: - prompt_parts = await prompt.get_components(arguments) - # Extract text content from parts - prompt_texts = [] - for part in prompt_parts: - if hasattr(part, "content"): - content = part.content - if isinstance(content, str): - prompt_texts.append(content) - elif isinstance(content, list): - # Handle Sequence[UserContent] - for item in content: - if isinstance(item, FileUrl): - prompt_texts.append(item.url) - elif isinstance(item, str): - prompt_texts.append(item) - prompt_text = "\n".join(prompt_texts) - # Run the expanded prompt through the agent - result = await state.agent.run(prompt_text) - output_text = str(result.data) - - except Exception as e: # noqa: BLE001 - output_text = f"Error executing command: {e}" + # Mark session as busy + state.session_status[session_id] = SessionStatus(type="busy") + await state.broadcast_event( + SessionStatusEvent.create(session_id, SessionStatus(type="busy")) + ) + # Add step-start part + part_id = identifier.ascending("part") + step_start = StepStartPart(id=part_id, message_id=assistant_msg_id, session_id=session_id) + assistant_msg_with_parts.parts.append(step_start) + await state.broadcast_event(PartUpdatedEvent.create(step_start)) - response_time = now_ms() - # Create text part with output - text_part = TextPart( - id=identifier.ascending("part"), - message_id=assistant_msg_id, - session_id=session_id, - text=output_text, - ) - assistant_msg_with_parts.parts.append(text_part) - await state.broadcast_event(PartUpdatedEvent.create(text_part)) - step_finish = StepFinishPart( - id=identifier.ascending("part"), - message_id=assistant_msg_id, - session_id=session_id, - ) - assistant_msg_with_parts.parts.append(step_finish) - await state.broadcast_event(PartUpdatedEvent.create(step_finish)) - # Update message with completion time - time_ = MessageTime(created=now, completed=response_time) - updated_assistant = assistant_message.model_copy(update={"time": time_}) - assistant_msg_with_parts.info = updated_assistant - await state.broadcast_event(MessageUpdatedEvent.create(updated_assistant)) - # Mark session as idle - state.session_status[session_id] = SessionStatus(type="idle") - await state.broadcast_event(SessionStatusEvent.create(session_id, SessionStatus(type="idle"))) + # Get prompt content and execute through the agent + try: + prompt_parts = await prompt.get_components(arguments) + # Extract text content from parts + prompt_texts = [] + for part in prompt_parts: + if hasattr(part, "content"): + content = part.content + if isinstance(content, str): + prompt_texts.append(content) + elif isinstance(content, list): + # Handle Sequence[UserContent] + for item in content: + if isinstance(item, FileUrl): + prompt_texts.append(item.url) + elif isinstance(item, str): + prompt_texts.append(item) + prompt_text = "\n".join(prompt_texts) + # Run the expanded prompt through the agent + result = await state.agent.run(prompt_text) + output_text = str(result.data) + + except Exception as e: # noqa: BLE001 + output_text = f"Error executing command: {e}" + + response_time = now_ms() + # Create text part with output + text_part = TextPart( + id=identifier.ascending("part"), + message_id=assistant_msg_id, + session_id=session_id, + text=output_text, + ) + assistant_msg_with_parts.parts.append(text_part) + await state.broadcast_event(PartUpdatedEvent.create(text_part)) + step_finish = StepFinishPart( + id=identifier.ascending("part"), + message_id=assistant_msg_id, + session_id=session_id, + ) + assistant_msg_with_parts.parts.append(step_finish) + await state.broadcast_event(PartUpdatedEvent.create(step_finish)) + # Update message with completion time + time_ = MessageTime(created=now, completed=response_time) + updated_assistant = assistant_message.model_copy(update={"time": time_}) + assistant_msg_with_parts.info = updated_assistant + await state.broadcast_event(MessageUpdatedEvent.create(updated_assistant)) + finally: + await state.mark_session_idle(session_id) # Broadcast command.executed event await state.broadcast_event( diff --git a/src/agentpool_server/opencode_server/server.py b/src/agentpool_server/opencode_server/server.py index 1eee9256b..d277cdcaf 100644 --- a/src/agentpool_server/opencode_server/server.py +++ b/src/agentpool_server/opencode_server/server.py @@ -11,7 +11,7 @@ from pathlib import Path from typing import TYPE_CHECKING, Any -from fastapi import FastAPI, Request # noqa: TC002 +from fastapi import FastAPI, Request from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse, RedirectResponse, Response @@ -344,6 +344,24 @@ async def get_doc() -> RedirectResponse: """Redirect to OpenAPI docs.""" return RedirectResponse(url="/docs") + # OTLP telemetry sink endpoints (compatibility for OpenCode 1.4.4+) + # Must be registered BEFORE the catch-all proxy so POST /v1/metrics etc. + # don't fall through to a GET/HEAD/OPTIONS-only route (→ 405). + @app.post("/v1/metrics") + async def otlp_metrics(request: Request) -> Response: + """Accept OTLP metrics payloads and discard them.""" + return Response(status_code=204) + + @app.post("/v1/traces") + async def otlp_traces(request: Request) -> Response: + """Accept OTLP traces payloads and discard them.""" + return Response(status_code=204) + + @app.post("/v1/logs") + async def otlp_logs(request: Request) -> Response: + """Accept OTLP logs payloads and discard them.""" + return Response(status_code=204) + # Proxy catch-all for OpenCode's hosted web UI # This must be registered LAST so it doesn't catch API routes @app.api_route("/{path:path}", methods=["GET", "HEAD", "OPTIONS"]) diff --git a/src/agentpool_server/opencode_server/state.py b/src/agentpool_server/opencode_server/state.py index 31ff86a08..fa906c49c 100644 --- a/src/agentpool_server/opencode_server/state.py +++ b/src/agentpool_server/opencode_server/state.py @@ -4,35 +4,41 @@ import asyncio from collections.abc import Callable, Coroutine +import contextlib from dataclasses import dataclass, field from pathlib import Path import time from typing import TYPE_CHECKING, Any +from agentpool import log from agentpool.diagnostics.lsp_manager import LSPManager -from agentpool.storage import StorageManager from agentpool.utils.time_utils import now_ms +from agentpool_server.opencode_server.models import SessionStatus from agentpool_server.opencode_server.provider_auth import create_default_auth_service from agentpool_storage.opencode_provider import helpers +logger = log.get_logger(__name__) + if TYPE_CHECKING: from fsspec.asyn import AsyncFileSystem from slashed import CommandStore from agentpool.agents.base_agent import BaseAgent from agentpool.delegation import AgentPool + from agentpool.storage import StorageManager from agentpool_server.opencode_server.input_provider import OpenCodeInputProvider from agentpool_server.opencode_server.models import ( Config, Event, + MessageRequest, MessageWithParts, QuestionInfo, Session, - SessionStatus, Todo, ) from agentpool_server.opencode_server.models.question import QuestionToolInfo + from agentpool_server.opencode_server.routes.global_routes import GlobalEventFactory # Type alias for async callback OnFirstSubscriberCallback = Callable[[], Coroutine[Any, Any, None]] @@ -55,6 +61,15 @@ class PendingQuestion: """Optional tool context.""" +@dataclass +class QueuedAsyncPrompt: + """Queued async prompt work owned by the OpenCode server.""" + + request: MessageRequest + user_msg_id: str + user_msg_with_parts: MessageWithParts + + @dataclass class ServerState: """Shared state for the OpenCode server. @@ -76,6 +91,8 @@ class ServerState: # Per-session locks for concurrent message handling # Ensures messages to the same session are processed sequentially session_locks: dict[str, asyncio.Lock] = field(default_factory=dict) + # Global lock for binding agent to a session (short critical section) + agent_lock: asyncio.Lock = field(default_factory=asyncio.Lock) # Message storage (session_id -> messages) # Runtime cache - messages are also persisted via pool.storage messages: dict[str, list[MessageWithParts]] = field(default_factory=dict) @@ -87,15 +104,20 @@ class ServerState: todos: dict[str, list[Todo]] = field(default_factory=dict) # Input providers for permission handling (session_id -> provider) input_providers: dict[str, OpenCodeInputProvider] = field(default_factory=dict) + # Per-session conversation histories for snapshot-based execution + session_conversations: dict[str, Any] = field(default_factory=dict) # Question storage (question_id -> pending question info) pending_questions: dict[str, PendingQuestion] = field(default_factory=dict) # SSE event subscribers event_subscribers: list[asyncio.Queue[Event]] = field(default_factory=list) + _event_factory: GlobalEventFactory | None = field(default=None, repr=False) # Callback for first subscriber connection (e.g., for update check) on_first_subscriber: OnFirstSubscriberCallback | None = None _first_subscriber_triggered: bool = field(default=False, repr=False) # Background tasks (for cleanup on shutdown) background_tasks: set[asyncio.Task[Any]] = field(default_factory=set) + # Per-session async prompt queue owned by the server runtime. + pending_async_prompts: dict[str, list[QueuedAsyncPrompt]] = field(default_factory=dict) # Event managers for subagent event routing (session_id -> event_manager) event_managers: dict[str, Any] = field(default_factory=dict) # Provider authentication service @@ -109,6 +131,36 @@ def __post_init__(self) -> None: """Initialize derived state.""" self.lsp_manager = LSPManager(env=self.agent.env) self.lsp_manager.register_defaults() + self._active_run_tasks: dict[str, asyncio.Task[Any]] = {} + + def get_event_factory(self) -> GlobalEventFactory: + """Get or lazily create the GlobalEventFactory for event wrapping. + + The factory is created on first access using the working directory + and computed project ID, then cached for the server's lifetime. + Imports GlobalEventFactory locally to avoid circular imports. + """ + from agentpool_server.opencode_server.routes.global_routes import GlobalEventFactory + + if self._event_factory is None: + directory = self.base_path + self._event_factory = GlobalEventFactory( + directory=directory, + project=helpers.compute_project_id(directory), + workspace=None, + ) + return self._event_factory + + def ensure_runtime_session_state(self, session_id: str) -> None: + """Ensure in-memory runtime buckets exist for a session. + + This is used both for brand-new sessions and for sessions reloaded from + persisted storage after a server restart. Cold-start recovery should not + depend on individual routes remembering to initialize each bucket. + """ + self.messages.setdefault(session_id, []) + self.reverted_messages.setdefault(session_id, []) + self.todos.setdefault(session_id, []) @property def fs(self) -> AsyncFileSystem: @@ -117,9 +169,14 @@ def fs(self) -> AsyncFileSystem: @property def base_path(self) -> str: - """Get the resolved root directory for file operations.""" - raw_path = self.agent.env.cwd or self.working_dir - return str(Path(raw_path).resolve()) + """Get the resolved OpenCode project root for routing and file operations. + + OpenCode routes SSE events against the server/project directory the client + attached to, not an agent-specific execution sandbox. Agent execution + environments may override `env.cwd` for tool isolation, but routing + metadata must remain anchored to the server's configured `working_dir`. + """ + return str(Path(self.working_dir).resolve()) @property def is_local_fs(self) -> bool: @@ -153,6 +210,64 @@ def get_session_lock(self, session_id: str) -> asyncio.Lock: self.session_locks[session_id] = asyncio.Lock() return self.session_locks[session_id] + def get_session_conversation(self, session_id: str) -> Any: + """Get or create a per-session MessageHistory.""" + from agentpool.messaging import MessageHistory + + if session_id not in self.session_conversations: + self.session_conversations[session_id] = MessageHistory() + return self.session_conversations[session_id] + + async def snapshot_for_session(self, session_id: str, *, agent: BaseAgent | None = None) -> Any: + """Capture per-run state from the shared agent. + + Must be called while holding agent_lock. Binds the agent to the session + and captures all mutable state into an immutable snapshot. + + Args: + session_id: The session to snapshot. + agent: Resolved agent to bind. Falls back to ``self.agent`` when + ``None`` so existing callers are unaffected. + """ + from agentpool.agents.context import RunSnapshot + from agentpool_server.opencode_server.input_provider import OpenCodeInputProvider + + resolved = agent or self.agent + + if session_id not in self.input_providers: + self.input_providers[session_id] = OpenCodeInputProvider(self, session_id) + + resolved._input_provider = self.input_providers[session_id] + resolved.session_id = session_id + + conversation = self.get_session_conversation(session_id) + model_name = resolved.model_name + mode_name = getattr(resolved, "_current_mode", None) + + return RunSnapshot( + session_id=session_id, + input_provider=self.input_providers[session_id], + conversation=conversation, + model_name=model_name, + mode_name=mode_name, + ) + + def register_active_run(self, session_id: str, task: asyncio.Task[Any]) -> None: + """Register an active run task for a session.""" + self._active_run_tasks[session_id] = task + + def unregister_active_run(self, session_id: str) -> None: + """Unregister an active run task for a session.""" + self._active_run_tasks.pop(session_id, None) + + def cancel_session_run(self, session_id: str) -> bool: + """Cancel a specific session's active run without affecting other sessions.""" + task = self._active_run_tasks.get(session_id) + if task and not task.done(): + task.cancel() + return True + return False + @property def storage(self) -> StorageManager: """Get the storage manager from the agent's pool. @@ -173,6 +288,45 @@ def create_background_task(self, coro: Any, *, name: str | None = None) -> async task.add_done_callback(self.background_tasks.discard) return task + def enqueue_async_prompt(self, session_id: str, queued_prompt: QueuedAsyncPrompt) -> None: + """Append async prompt work to a session-owned queue.""" + self.pending_async_prompts.setdefault(session_id, []).append(queued_prompt) + + def pop_next_async_prompt(self, session_id: str) -> QueuedAsyncPrompt | None: + """Pop the next queued async prompt for a session, if any.""" + queue = self.pending_async_prompts.get(session_id) + if not queue: + return None + queued_prompt = queue.pop(0) + if not queue: + self.pending_async_prompts.pop(session_id, None) + return queued_prompt + + def clear_pending_async_prompts(self, session_id: str) -> None: + """Drop queued async prompt work for a session.""" + self.pending_async_prompts.pop(session_id, None) + + def has_pending_async_prompts(self, session_id: str) -> bool: + """Return whether a session currently has queued async prompt work.""" + return bool(self.pending_async_prompts.get(session_id)) + + def has_session_background_task(self, session_id: str) -> bool: + """Return whether a per-session prompt worker is already running.""" + task_name = f"process_message_{session_id}" + return any( + task.get_name() == task_name and not task.done() for task in self.background_tasks + ) + + async def cancel_session_background_tasks(self, session_id: str) -> None: + """Cancel background tasks associated with a session.""" + task_name = f"process_message_{session_id}" + tasks = [task for task in self.background_tasks if task.get_name() == task_name] + self.clear_pending_async_prompts(session_id) + for task in tasks: + task.cancel() + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + async def cleanup_tasks(self) -> None: """Cancel and wait for all background tasks.""" for task in self.background_tasks: @@ -182,10 +336,46 @@ async def cleanup_tasks(self) -> None: self.background_tasks.clear() async def broadcast_event(self, event: Event) -> None: - """Broadcast an event to all SSE subscribers.""" - # print(f"Broadcasting event: {event.type} to {len(self.event_subscribers)} subscribers") - for queue in self.event_subscribers: - await queue.put(event) + """Broadcast an event to all SSE subscribers. + + Isolates failures: if one subscriber's queue raises, + other subscribers still receive the event. + + Uses put_nowait() instead of await queue.put() to avoid blocking + the broadcaster when a subscriber's queue is full. Iterates over + a copy of event_subscribers to avoid mutation during iteration + (subscribers can be removed by the _event_generator finally block + or by error handling below). + """ + for queue in list(self.event_subscribers): # iterate copy to avoid mutation + try: + queue.put_nowait(event) + except asyncio.QueueFull: + logger.warning("SSE subscriber queue full, dropping event") + except Exception: # noqa: BLE001 + logger.warning("SSE subscriber queue error, removing subscriber") + with contextlib.suppress(ValueError): + self.event_subscribers.remove(queue) + + async def mark_session_idle(self, session_id: str) -> None: + """Mark a session idle and broadcast the matching status events.""" + from agentpool_server.opencode_server.models import SessionIdleEvent, SessionStatusEvent + + status = SessionStatus(type="idle") + self.session_status[session_id] = status + await self.broadcast_event(SessionStatusEvent.create(session_id, status)) + await self.broadcast_event(SessionIdleEvent.create(session_id)) + + async def emit_session_turn_complete(self, session_id: str) -> None: + """Broadcast the per-turn completion signal without changing busy state. + + OpenCode clients still use ``session.idle`` as an end-of-turn marker. + For queued async prompts we need that signal after each finished turn, + even while the server-owned queue still has follow-up work to process. + """ + from agentpool_server.opencode_server.models import SessionIdleEvent + + await self.broadcast_event(SessionIdleEvent.create(session_id)) async def ensure_session( self, @@ -215,7 +405,6 @@ async def ensure_session( from agentpool_server.opencode_server.models import ( Session, SessionCreatedEvent, - SessionStatus, TimeCreatedUpdated, ) @@ -234,17 +423,23 @@ async def ensure_session( # Persist to storage id_ = self.pool.manifest.config_file_path session_data = opencode_to_session_data(session, agent_name=self.agent.name, pool_id=id_) - await self.pool.storage.save_session(session_data) + if self.pool.sessions.store: + await self.pool.sessions.store.save(session_data) + else: + await self.pool.storage.save_session(session_data) # Cache in memory self.sessions[session_id] = session - self.messages[session_id] = [] - self.session_status[session_id] = SessionStatus(type="idle") - self.todos[session_id] = [] + self.ensure_runtime_session_state(session_id) + await self.mark_session_idle(session_id) # Create input provider for this session input_provider = OpenCodeInputProvider(self, session_id) self.input_providers[session_id] = input_provider + # Bind agent to this session under lock to avoid racing with snapshot_for_session() + async with self.agent_lock: + self.agent._input_provider = input_provider + self.agent.session_id = session_id await self.broadcast_event(SessionCreatedEvent.create(session)) diff --git a/src/agentpool_server/opencode_server/stream_adapter.py b/src/agentpool_server/opencode_server/stream_adapter.py index 4543033b5..cd95dd4d8 100644 --- a/src/agentpool_server/opencode_server/stream_adapter.py +++ b/src/agentpool_server/opencode_server/stream_adapter.py @@ -80,6 +80,8 @@ class OpenCodeStreamAdapter: processor: EventProcessor = field(default_factory=EventProcessor, init=False) main_context: EventProcessorContext = field(init=False) _cost_info: Any = field(default=None, init=False) + _step_finish_emitted: bool = field(default=False, init=False) + """Tracks whether StepFinishPart was already emitted by _process_stream_complete.""" def __post_init__(self) -> None: self.main_context = EventProcessorContext( @@ -147,6 +149,11 @@ async def process_stream( try: async for event in stream: async for oc_event in self.processor.process(event, self.main_context): + # Track if StepFinishPart was emitted by _process_stream_complete + if isinstance(oc_event, PartUpdatedEvent) and isinstance( + oc_event.properties.part, StepFinishPart + ): + self._step_finish_emitted = True yield oc_event except asyncio.CancelledError: # Stream was cancelled by user - this is expected behavior @@ -176,7 +183,8 @@ def finalize(self) -> Iterator[Event]: """Yield final events after the stream has ended. Produces the final text part update (or creates one if text was never - streamed), the step-finish part, and the final text timing update. + streamed), the step-finish part (if not already emitted by + _process_stream_complete), and the final text timing update. """ response_time = now_ms() start = self.main_context.stream_start_ms @@ -204,7 +212,11 @@ def finalize(self) -> Iterator[Event]: ) self.assistant_msg.update_part(final_text_part) - # Step finish + # Step finish — skip if already emitted by _process_stream_complete + # (StreamCompleteEvent handler in EventProcessor also emits StepFinishPart) + if self._step_finish_emitted: + return + cache = TokenCache(read=0, write=0) tokens = Tokens( cache=cache, diff --git a/src/agentpool_storage/base.py b/src/agentpool_storage/base.py index 187aaddad..b399a3467 100644 --- a/src/agentpool_storage/base.py +++ b/src/agentpool_storage/base.py @@ -53,6 +53,9 @@ class StorageProvider: can_load_history: bool = False """Whether this provider supports loading history.""" + can_store_projects: bool = False + """Whether this provider supports project storage.""" + def __init__(self, config: BaseStorageProviderConfig) -> None: super().__init__() self.config = config diff --git a/src/agentpool_storage/file_provider/provider.py b/src/agentpool_storage/file_provider/provider.py index f960b344a..23fcbc97d 100644 --- a/src/agentpool_storage/file_provider/provider.py +++ b/src/agentpool_storage/file_provider/provider.py @@ -99,6 +99,7 @@ class FileProvider(StorageProvider): """ can_load_history = True + can_store_projects = True def __init__(self, config: FileStorageConfig | None = None) -> None: """Initialize file provider. diff --git a/src/agentpool_storage/memory_provider/provider.py b/src/agentpool_storage/memory_provider/provider.py index 30607670d..b9950a884 100644 --- a/src/agentpool_storage/memory_provider/provider.py +++ b/src/agentpool_storage/memory_provider/provider.py @@ -25,6 +25,7 @@ class MemoryStorageProvider(StorageProvider): """In-memory storage provider for testing.""" can_load_history = True + can_store_projects = True def __init__(self, config: MemoryStorageConfig | None = None) -> None: super().__init__(config or MemoryStorageConfig()) diff --git a/src/agentpool_storage/opencode_provider/helpers.py b/src/agentpool_storage/opencode_provider/helpers.py index 7bed03f87..43a508bf3 100644 --- a/src/agentpool_storage/opencode_provider/helpers.py +++ b/src/agentpool_storage/opencode_provider/helpers.py @@ -15,10 +15,12 @@ from pydantic import TypeAdapter from pydantic_ai import ( BinaryContent, + FileUrl, ModelRequest, ModelResponse, RequestUsage, RunUsage, + TextContent, TextPart as PydanticTextPart, ThinkingPart, ToolCallPart, @@ -28,6 +30,7 @@ from agentpool.log import get_logger from agentpool.messaging import ChatMessage, TokenCost +from agentpool.utils import identifiers as identifier from agentpool.utils.pydantic_ai_helpers import to_user_content from agentpool.utils.time_utils import ms_to_datetime from agentpool_server.opencode_server.models.message import ( @@ -47,6 +50,7 @@ if TYPE_CHECKING: + from collections.abc import Sequence from datetime import datetime from pydantic_ai.messages import UserContent @@ -61,8 +65,21 @@ def compute_project_id(directory: str) -> str: """Compute OpenCode project ID from directory. - OpenCode uses the root commit SHA1 of the git repository as the project ID. - If not in a git repository, returns 'global'. + This is the **OpenCode external contract** for project identification. + OpenCode uses the root commit SHA1 of the git repository as the project ID, + which determines the session directory structure: + ``storage/session/{project_id}/``. This ID is stable across different + machines that share the same git history. + + If not in a git repository, returns ``'global'``. + + !!! note "Dual ID scheme" + This is DIFFERENT from ``generate_project_id()`` in + ``agentpool_storage.project_store``, which hashes the canonical worktree + path for AgentPool-internal ``ProjectData`` identification. The two IDs + serve distinct purposes and are not interchangeable: + - ``compute_project_id`` — git root commit SHA1 (OpenCode session layout) + - ``generate_project_id`` — path SHA1 (AgentPool project registry) Args: directory: Project directory path @@ -182,6 +199,82 @@ def extract_text_content(parts: list[Part]) -> str: return "\n".join(text_segments) +def convert_user_content_to_parts( + content: str | Sequence[UserContent], + message_id: str, + session_id: str, + part_counter_start: int, +) -> list[TextPart]: + """Convert user content from pydantic-ai UserPromptPart to OpenCode TextParts. + + Handles both simple string content and structured UserContent lists. + Non-text content types (BinaryContent, FileUrl, etc.) are skipped with a + warning since the OpenCode storage provider only persists text parts. + + Args: + content: String or list of UserContent items from a UserPromptPart + message_id: The parent message ID + session_id: The parent session ID + part_counter_start: Starting counter for generating part IDs + + Returns: + List of TextPart models for text content items + """ + if isinstance(content, str): + part_id = identifier.ascending("part") + return [ + TextPart( + id=part_id, + message_id=message_id, + session_id=session_id, + text=content, + ) + ] + + parts: list[TextPart] = [] + for item in content: + match item: + case str(): + part_id = identifier.ascending("part") + parts.append( + TextPart( + id=part_id, + message_id=message_id, + session_id=session_id, + text=item, + ) + ) + case TextContent(): + part_id = identifier.ascending("part") + parts.append( + TextPart( + id=part_id, + message_id=message_id, + session_id=session_id, + text=item.content, + ) + ) + case BinaryContent(): + logger.warning( + "Skipping BinaryContent in user message", + media_type=item.media_type, + message_id=message_id, + ) + case FileUrl(): + logger.warning( + "Skipping FileUrl in user message", + url=item.url, + message_id=message_id, + ) + case _: + logger.warning( + "Skipping unsupported UserContent type", + content_type=type(item).__name__, + message_id=message_id, + ) + return parts + + def _build_user_pydantic_messages( parts: list[Part], timestamp: datetime, diff --git a/src/agentpool_storage/opencode_provider/provider.py b/src/agentpool_storage/opencode_provider/provider.py index a47beb246..f3e8d7a2e 100644 --- a/src/agentpool_storage/opencode_provider/provider.py +++ b/src/agentpool_storage/opencode_provider/provider.py @@ -33,6 +33,7 @@ ) from agentpool.log import get_logger +from agentpool.sessions.models import ProjectData, SessionData from agentpool.utils.pydantic_ai_helpers import safe_args_as_dict from agentpool.utils.time_utils import datetime_to_ms, get_now, ms_to_datetime from agentpool_config.storage import OpenCodeStorageConfig @@ -47,11 +48,11 @@ from agentpool_storage.base import StorageProvider from agentpool_storage.models import ConversationData as ConvData, TokenUsage from agentpool_storage.opencode_provider import helpers +from agentpool_storage.project_store import generate_project_id if TYPE_CHECKING: from agentpool.messaging import ChatMessage, TokenCost - from agentpool.sessions.models import SessionData from agentpool_config.session import SessionQuery from agentpool_storage.models import QueryFilters, StatsFilters @@ -127,6 +128,7 @@ class OpenCodeStorageProvider(StorageProvider): """ can_load_history = True + can_store_projects = True def __init__(self, config: OpenCodeStorageConfig | None = None) -> None: """Initialize OpenCode storage provider.""" @@ -142,6 +144,8 @@ def __init__(self, config: OpenCodeStorageConfig | None = None) -> None: self.sessions_path = self.base_path / "session" self.messages_path = self.base_path / "message" self.parts_path = self.base_path / "part" + self.projects_path = self.base_path / "project" + self.projects_path.mkdir(parents=True, exist_ok=True) def _list_sessions(self, project_id: str | None = None) -> list[tuple[str, Path]]: """List all sessions, optionally filtered by project.""" @@ -259,8 +263,8 @@ async def _write_message( # noqa: PLR0915 session_id=session_id, parent_id=parent_id or "", model_id=model or "", - provider_id="", # TODO: get from somewhere - path=MessagePath(cwd="", root=""), # TODO: get real paths + provider_id=model.split(":")[0] if model else "agentpool", + path=MessagePath(cwd=str(self.base_path), root=str(self.base_path)), time=MessageTime(created=now_ms), tokens=Tokens( input=cost_info.token_usage.input_tokens if cost_info else 0, @@ -797,6 +801,145 @@ async def fork_conversation( (self.parts_path / new_session_id).mkdir(parents=True, exist_ok=True) return fork_point_id + # Project storage methods + + async def save_project(self, project: ProjectData) -> None: + """Save or update a project. + + Writes project data as a JSON file at projects_path/{project_id}.json. + + Args: + project: Project data to persist + """ + project_file = self.projects_path / f"{project.project_id}.json" + data = project.model_dump(mode="json") + project_file.write_text(anyenv.dump_json(data, indent=True), encoding="utf-8") + logger.debug("Saved project", project_id=project.project_id) + + async def get_project(self, project_id: str) -> ProjectData | None: + """Get a project by ID. + + Args: + project_id: Project identifier + + Returns: + Project data if found, None otherwise + """ + project_file = self.projects_path / f"{project_id}.json" + if not project_file.exists(): + return None + try: + content = project_file.read_text(encoding="utf-8") + data = anyenv.load_json(content, return_type=dict) + return ProjectData.model_validate(data) + except (anyenv.JsonLoadError, Exception) as e: # noqa: BLE001 + logger.warning("Failed to read project file", path=str(project_file), error=str(e)) + return None + + async def get_project_by_worktree(self, worktree: str) -> ProjectData | None: + """Get a project by worktree path. + + Resolves the worktree path before comparing to handle symlink differences. + Uses generate_project_id for O(1) lookup instead of scanning all project files. + + Args: + worktree: Absolute path to the project worktree + + Returns: + Project data if found, None otherwise + """ + project_id = generate_project_id(worktree) + project_file = self.projects_path / f"{project_id}.json" + if not project_file.exists(): + return None + try: + content = project_file.read_text(encoding="utf-8") + data = anyenv.load_json(content, return_type=dict) + project = ProjectData.model_validate(data) + # Verify the worktree matches (safety check for hash collisions / stale data) + if project.worktree and str(Path(project.worktree).resolve()) == str( + Path(worktree).resolve() + ): + return project + except (anyenv.JsonLoadError, Exception): # noqa: BLE001 + pass + return None + + async def get_project_by_name(self, name: str) -> ProjectData | None: + """Get a project by friendly name. + + Args: + name: Project name + + Returns: + Project data if found, None otherwise + """ + for project_file in self.projects_path.glob("*.json"): + try: + content = project_file.read_text(encoding="utf-8") + data = anyenv.load_json(content, return_type=dict) + project = ProjectData.model_validate(data) + if project.name == name: + return project + except (anyenv.JsonLoadError, Exception): # noqa: BLE001 + continue + return None + + async def list_projects(self, limit: int | None = None) -> list[ProjectData]: + """List all projects, ordered by last_active descending. + + Args: + limit: Maximum number of projects to return + + Returns: + List of project data objects sorted by last_active descending + """ + projects: list[ProjectData] = [] + for project_file in self.projects_path.glob("*.json"): + try: + content = project_file.read_text(encoding="utf-8") + data = anyenv.load_json(content, return_type=dict) + project = ProjectData.model_validate(data) + projects.append(project) + except (anyenv.JsonLoadError, Exception) as e: # noqa: BLE001 + logger.warning( + "Skipping corrupted project file", + path=str(project_file), + error=str(e), + ) + continue + projects.sort(key=lambda p: p.last_active, reverse=True) + if limit is not None: + projects = projects[:limit] + return projects + + async def delete_project(self, project_id: str) -> bool: + """Delete a project. + + Args: + project_id: Project identifier + + Returns: + True if project was deleted, False if not found + """ + project_file = self.projects_path / f"{project_id}.json" + if not project_file.exists(): + return False + project_file.unlink() + logger.debug("Deleted project", project_id=project_id) + return True + + async def touch_project(self, project_id: str) -> None: + """Update project's last_active timestamp. + + Args: + project_id: Project identifier + """ + project = await self.get_project(project_id) + if project is not None: + updated = project.touch() + await self.save_project(updated) + # Session persistence methods (required by StorageProvider base class) async def load_session(self, session_id: str) -> SessionData | None: @@ -810,8 +953,6 @@ async def load_session(self, session_id: str) -> SessionData | None: Returns: SessionData if session was found and loaded, None otherwise """ - from agentpool.sessions.models import SessionData - # Find session file session_path = next( (p for sid, p in self._list_sessions() if sid == session_id), @@ -946,6 +1087,213 @@ async def delete_session(self, session_id: str) -> bool: return True + async def update_session_title(self, session_id: str, title: str) -> None: + """Update the title of a conversation. + + Finds the session JSON file and updates its title field. + + Args: + session_id: ID of the conversation to update + title: New title for the conversation + """ + session_path = next( + (p for sid, p in self._list_sessions() if sid == session_id), + None, + ) + if not session_path: + logger.warning("Session not found for title update", session_id=session_id) + return + + oc_session = helpers.read_session(session_path) + if not oc_session: + logger.warning("Failed to read session for title update", session_id=session_id) + return + + oc_session.title = title + oc_session.time.updated = datetime_to_ms(get_now()) + dct = oc_session.model_dump(by_alias=True) + session_path.write_text(anyenv.dump_json(dct, indent=True), encoding="utf-8") + + async def update_sdk_session_id(self, session_id: str, sdk_session_id: str) -> None: + """Update the external SDK session ID for a session. + + Stores the SDK session ID in the session JSON's metadata.sdk_session_id field. + Creates the metadata dict if it doesn't exist. + + Args: + session_id: Internal session identifier + sdk_session_id: External SDK session ID + """ + session_path = next( + (p for sid, p in self._list_sessions() if sid == session_id), + None, + ) + if not session_path: + logger.warning("Session not found for SDK session ID update", session_id=session_id) + return + + try: + content = session_path.read_text(encoding="utf-8") + data = anyenv.load_json(content, return_type=dict) + except anyenv.JsonLoadError as e: + logger.warning( + "Failed to read session for SDK session ID update", + session_id=session_id, + error=str(e), + ) + return + + metadata = data.get("metadata", {}) + metadata["sdk_session_id"] = sdk_session_id + data["metadata"] = metadata + data["time"] = data.get("time", {}) + data["time"]["updated"] = datetime_to_ms(get_now()) + session_path.write_text(anyenv.dump_json(data, indent=True), encoding="utf-8") + + async def delete_session_messages(self, session_id: str) -> int: + """Delete all messages for a session. + + Removes message JSON files and their associated part files from disk. + + Args: + session_id: ID of the conversation to clear + + Returns: + Number of messages deleted + """ + msg_dir = self.messages_path / session_id + if not msg_dir.exists(): + return 0 + + # Collect message IDs before deleting so we can clean up parts + message_files = list(msg_dir.glob("*.json")) + message_ids = [f.stem for f in message_files] + + # Delete part files for each message + for message_id in message_ids: + parts_dir = self.parts_path / message_id + if parts_dir.exists(): + for part_file in parts_dir.glob("*.json"): + part_file.unlink(missing_ok=True) + # Remove empty parts directory + if not any(parts_dir.iterdir()): + parts_dir.rmdir() + + # Delete message files + for msg_file in message_files: + msg_file.unlink(missing_ok=True) + + # Remove empty message directory + if msg_dir.exists() and not any(msg_dir.iterdir()): + msg_dir.rmdir() + + return len(message_files) + + async def get_filtered_conversations( + self, + agent_name: str | None = None, + period: str | None = None, + since: datetime | None = None, + query: str | None = None, + model: str | None = None, + limit: int | None = None, + *, + compact: bool = False, + include_tokens: bool = False, + ) -> list[ConvData]: + """Get filtered conversations with formatted output. + + Iterates all session JSON files, applies filters, and returns matching + ConversationData objects. + + Args: + agent_name: Filter by agent name + period: Time period to include (e.g. "1h", "2d") + since: Only show conversations after this time + query: Search in message content + model: Filter by model used + limit: Maximum number of conversations + compact: Only show first/last message of each conversation + include_tokens: Include token usage statistics + """ + from agentpool.utils.parse_time import parse_time_period + from agentpool.utils.time_utils import get_now + + cutoff: datetime | None = None + if period: + cutoff = get_now() - parse_time_period(period) + elif since: + cutoff = since + + result: list[ConvData] = [] + for session_id, session_path in self._list_sessions(): + session = helpers.read_session(session_path) + if not session: + continue + + # Date filter + session_created = ms_to_datetime(session.time.created) + if cutoff and session_created < cutoff: + continue + + oc_messages = self._read_messages(session_id) + if not oc_messages: + continue + + # Read parts for all messages + msg_parts_map = {oc_msg.id: self._read_parts(oc_msg.id) for oc_msg in oc_messages} + + # Convert messages + chat_messages: list[ChatMessage[str]] = [] + total_tokens = 0 + for oc_msg in oc_messages: + parts = msg_parts_map.get(oc_msg.id, []) + chat_msg = helpers.to_chat_message(msg=oc_msg, parts=parts) + chat_messages.append(chat_msg) + + if isinstance(oc_msg, AssistantMessage) and oc_msg.tokens: + total_tokens += oc_msg.tokens.input + oc_msg.tokens.output + + if not chat_messages: + continue + + # Agent filter + if agent_name and not any(m.name == agent_name for m in chat_messages): + continue + + # Content query filter + if query and not any(query in m.content for m in chat_messages): + continue + + # Model filter + if model and not any( + isinstance(oc_msg, AssistantMessage) and oc_msg.model_id == model + for oc_msg in oc_messages + ): + continue + + usage = TokenUsage(total=total_tokens, prompt=0, completion=0) if total_tokens else None + + # Compact mode: only first and last message + filtered_messages = chat_messages + if compact and len(chat_messages) > 2: + filtered_messages = [chat_messages[0], chat_messages[-1]] + + conv_data = ConvData( + id=session_id, + agent=chat_messages[0].name or "opencode", + title=session.title, + start_time=session_created.isoformat(), + messages=filtered_messages, + token_usage=usage if include_tokens else None, + ) + result.append(conv_data) + + if limit and len(result) >= limit: + break + + return result + async def list_session_ids( self, *, @@ -962,7 +1310,7 @@ async def list_session_ids( List of session IDs """ session_ids: list[str] = [] - for session_id, session_path in self._list_sessions(): + for session_id, _session_path in self._list_sessions(): # Check agent filter if specified # Note: OpenCode session files don't store agent name directly, # so we need to check the first message's agent diff --git a/src/agentpool_storage/project_store.py b/src/agentpool_storage/project_store.py index 595c17a7f..e606b1178 100644 --- a/src/agentpool_storage/project_store.py +++ b/src/agentpool_storage/project_store.py @@ -88,6 +88,20 @@ def detect_project_root(cwd: str) -> tuple[str, str | None]: def generate_project_id(worktree: str) -> str: """Generate stable hash of canonical worktree path. + This is the **AgentPool internal contract** for project identification. + The ID is a SHA1 hash of the resolved (canonical) worktree path, used by + AgentPool to identify projects in storage (``ProjectData.project_id``). + Because it is derived from the filesystem path, it is machine-specific — + the same project cloned to different locations will have different IDs. + + !!! note "Dual ID scheme" + This is DIFFERENT from ``compute_project_id()`` in + ``agentpool_storage.opencode_provider.helpers``, which derives the ID + from the git root commit SHA1 for OpenCode session directory layout. + The two IDs serve distinct purposes and are not interchangeable: + - ``generate_project_id`` — path SHA1 (AgentPool project registry) + - ``compute_project_id`` — git root commit SHA1 (OpenCode session layout) + Args: worktree: Absolute path to project root diff --git a/src/agentpool_storage/sql_provider/sql_provider.py b/src/agentpool_storage/sql_provider/sql_provider.py index 3dad13873..a32ceaf3c 100644 --- a/src/agentpool_storage/sql_provider/sql_provider.py +++ b/src/agentpool_storage/sql_provider/sql_provider.py @@ -64,6 +64,7 @@ class SQLModelProvider(StorageProvider): """ can_load_history = True + can_store_projects = True def __init__(self, config: SQLStorageConfig | None = None) -> None: """Initialize provider with async database engine. diff --git a/tests/agents/test_run_snapshot_passthrough.py b/tests/agents/test_run_snapshot_passthrough.py new file mode 100644 index 000000000..9c0cff03f --- /dev/null +++ b/tests/agents/test_run_snapshot_passthrough.py @@ -0,0 +1,94 @@ +"""Regression test: BaseAgent.run() must accept snapshot kwarg. + +init_session() in the OpenCode server calls agent.run(..., snapshot=snapshot). +Before the fix, BaseAgent.run() lacked the `snapshot` parameter, causing +``TypeError: unexpected keyword argument 'snapshot'`` at runtime. + +This test ensures the parameter is present and passed through to run_stream(). +""" + +from __future__ import annotations + +import inspect +from typing import TYPE_CHECKING + +import pytest + +from agentpool.agents.base_agent import BaseAgent +from agentpool.agents.context import RunSnapshot + + +if TYPE_CHECKING: + pass + + +def test_run_accepts_snapshot_parameter() -> None: + """BaseAgent.run() must accept a snapshot keyword argument. + + This is the direct regression test: if the parameter is missing from + the signature, calling agent.run(snapshot=...) would raise TypeError. + """ + sig = inspect.signature(BaseAgent.run) + assert "snapshot" in sig.parameters, ( + "BaseAgent.run() must accept 'snapshot' parameter for init_session compatibility" + ) + param = sig.parameters["snapshot"] + assert param.default is not inspect.Parameter.empty, ( + "snapshot parameter must have a default value for backward compatibility" + ) + annotation = param.annotation + # Accept both "RunSnapshot | None" and "Optional[RunSnapshot]" string forms + assert "RunSnapshot" in str(annotation), ( + f"snapshot parameter must be typed as RunSnapshot | None, got: {annotation}" + ) + + +def test_run_snapshot_default_is_none() -> None: + """snapshot parameter defaults to None for backward compatibility.""" + sig = inspect.signature(BaseAgent.run) + param = sig.parameters["snapshot"] + assert param.default is None, ( + f"snapshot default must be None for backward compat, got: {param.default}" + ) + + +def test_run_signature_matches_run_stream_for_snapshot() -> None: + """Both run() and run_stream() must expose the same snapshot parameter type.""" + run_sig = inspect.signature(BaseAgent.run) + stream_sig = inspect.signature(BaseAgent.run_stream) + run_snapshot = run_sig.parameters["snapshot"] + stream_snapshot = stream_sig.parameters["snapshot"] + assert str(run_snapshot.annotation) == str(stream_snapshot.annotation), ( + f"run() snapshot annotation ({run_snapshot.annotation}) must match " + f"run_stream() snapshot annotation ({stream_snapshot.annotation})" + ) + + +@pytest.mark.asyncio +async def test_run_delegates_snapshot_to_run_stream() -> None: + """BaseAgent.run() must pass snapshot through to run_stream(). + + Verifies the internal delegation: when run() is called with snapshot, + it must forward that value to run_stream() rather than dropping it. + """ + # We inspect the source to verify the delegation — this is more reliable + # than trying to instantiate a full agent for a signature passthrough test. + source = inspect.getsource(BaseAgent.run) + # The run() method should pass snapshot=snapshot in the run_stream() call + assert "snapshot=snapshot" in source, ( + "BaseAgent.run() must delegate snapshot to run_stream() via 'snapshot=snapshot'" + ) + + +def test_run_snapshot_type_annotation() -> None: + """snapshot parameter must accept RunSnapshot instances.""" + # Verify RunSnapshot is constructible and the type is correct + snapshot = RunSnapshot(session_id="test-session-123") + assert snapshot.session_id == "test-session-123" + # The parameter annotation should accept this type + sig = inspect.signature(BaseAgent.run) + param = sig.parameters["snapshot"] + # Check "None" appears in the annotation (optional parameter) + assert "None" in str(param.annotation), ( + f"snapshot must be optional (None allowed), got: {param.annotation}" + ) diff --git a/tests/servers/opencode_server/conftest.py b/tests/servers/opencode_server/conftest.py index cb898143c..159c82f86 100644 --- a/tests/servers/opencode_server/conftest.py +++ b/tests/servers/opencode_server/conftest.py @@ -11,6 +11,8 @@ from __future__ import annotations import asyncio +import contextlib +import json from pathlib import Path import tempfile from typing import TYPE_CHECKING, Any @@ -30,8 +32,9 @@ from agentpool_server.opencode_server.models import Session from agentpool_server.opencode_server.models.common import TimeCreatedUpdated from agentpool_server.opencode_server.routes import agent_router, file_router, session_router +from agentpool_server.opencode_server.routes.global_routes import router as global_router +from agentpool_server.opencode_server.routes.message_routes import router as message_router from agentpool_server.opencode_server.state import ServerState -from agentpool_storage.memory_provider.provider import MemoryStorageProvider if TYPE_CHECKING: @@ -145,6 +148,9 @@ def mock_pool( # Sessions store must use AsyncMock for awaitable save/delete operations pool.sessions = Mock() pool.sessions.store = AsyncMock() + pool.sessions.store.save = AsyncMock() + pool.sessions.store.delete = AsyncMock() + pool.sessions.store.list_sessions = AsyncMock(return_value=[]) return pool @@ -210,8 +216,10 @@ def app(server_state: ServerState) -> FastAPI: """Create a FastAPI app with all routes for testing.""" app = FastAPI() app.include_router(session_router) + app.include_router(message_router) app.include_router(file_router) app.include_router(agent_router) + app.include_router(global_router) app.dependency_overrides[get_state] = lambda: server_state return app @@ -271,6 +279,79 @@ async def capturing_broadcast(event: Any) -> None: return capture +# ============================================================================= +# SSE Stream Fixtures +# ============================================================================= + + +class SSEStream: + r"""Async helper for consuming SSE events from the /global/event endpoint. + + Connects via httpx streaming, parses ``data: {json}\n\n`` lines, + and exposes parsed events through an async queue. + """ + + def __init__(self, client: AsyncClient) -> None: + self._client = client + self._queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue() + self._task: asyncio.Task[None] | None = None + + async def connect(self) -> None: + """Connect to SSE endpoint and start consuming events.""" + self._task = asyncio.create_task(self._consume()) + + async def _consume(self) -> None: + """Background task that reads SSE events and puts them in queue.""" + async with self._client.stream("GET", "/global/event") as response: + async for line in response.aiter_lines(): + if line.startswith("data: "): + event_data = json.loads(line[6:]) + await self._queue.put(event_data) + elif line.startswith(": "): + continue # SSE comment / keepalive + + async def next_event(self, timeout: float = 5.0) -> dict[str, Any]: + """Get next parsed SSE event with timeout.""" + return await asyncio.wait_for(self._queue.get(), timeout=timeout) + + async def aclose(self) -> None: + """Close the SSE stream.""" + if self._task: + self._task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await self._task + + +@pytest.fixture +async def global_event_stream(async_client: AsyncClient) -> AsyncIterator[SSEStream]: + """Create an SSE stream consumer for /global/event endpoint. + + Automatically connects and consumes the initial ``server.connected`` + event before yielding. + """ + stream = SSEStream(async_client) + await stream.connect() + # Consume the initial server.connected event + connected = await stream.next_event(timeout=5.0) + assert connected.get("type") == "server.connected" + yield stream + await stream.aclose() + + +def parse_sse_event(line: str) -> dict[str, Any]: + """Parse a single SSE data line into a dict. + + Args: + line: Raw SSE line, e.g. ``data: {"type": "server.connected"}`` + + Returns: + Parsed JSON dict from the data payload. + """ + if line.startswith("data: "): + return json.loads(line[6:]) + return json.loads(line) + + # ============================================================================= # Session Factory Fixtures # ============================================================================= diff --git a/tests/servers/opencode_server/test_concurrent_messages.py b/tests/servers/opencode_server/test_concurrent_messages.py index 26991bfb6..e442fd359 100644 --- a/tests/servers/opencode_server/test_concurrent_messages.py +++ b/tests/servers/opencode_server/test_concurrent_messages.py @@ -14,21 +14,17 @@ from agentpool_server.opencode_server.models import ( MessageRequest, - Session, - SessionStatus, TextPartInput, ) -from agentpool_server.opencode_server.models.common import TimeCreatedUpdated from agentpool_server.opencode_server.models.message import UserMessage +from agentpool_server.opencode_server.models.parts import TextPart from agentpool_server.opencode_server.routes.message_routes import _process_message from agentpool_server.opencode_server.state import ServerState -from agentpool.utils.time_utils import now_ms + if TYPE_CHECKING: from collections.abc import AsyncIterator - from agentpool.agents.base_agent import BaseAgent - class SlowAgentMock: """Mock agent that simulates slow processing to expose concurrency issues.""" @@ -45,6 +41,10 @@ def __init__(self, delay: float = 0.5) -> None: self._input_provider = None self.model_name = "test-model" self.session_id: str | None = None + # Snapshot tracking: maps session_id -> model_name captured from snapshot + self.snapshot_session_ids: dict[str, str] = {} + # Model name captured from snapshot at run_stream call time + self.model_names_at_call: dict[str, str | None] = {} async def set_model(self, model: str) -> None: """Mock set_model method.""" @@ -52,7 +52,7 @@ async def set_model(self, model: str) -> None: async def set_mode(self, mode: str, category_id: str | None = None) -> None: """Mock set_mode method.""" - pass + return async def get_available_models(self): """Mock get_available_models method.""" @@ -67,10 +67,17 @@ def run_stream( self, user_prompt: Any, session_id: str | None = None, + **kwargs: Any, ) -> AsyncIterator[Any]: """Simulate slow processing with concurrent run detection.""" self.run_stream_call_count += 1 + # Capture snapshot info for verification in tests + snapshot = kwargs.get("snapshot") + if snapshot is not None and session_id is not None: + self.snapshot_session_ids[session_id] = snapshot.session_id + self.model_names_at_call[session_id] = snapshot.model_name + # Check if another run is already active for this session if session_id in self.active_runs: raise RuntimeError( @@ -86,7 +93,7 @@ async def stream() -> AsyncIterator[Any]: await asyncio.sleep(self.delay) # Yield a simple text event - from agentpool.agents.events import TextContentItem, StreamCompleteEvent + from agentpool.agents.events import StreamCompleteEvent, TextContentItem from agentpool.messaging import ChatMessage yield TextContentItem(text=f"Response for {session_id}") @@ -101,6 +108,7 @@ async def stream() -> AsyncIterator[Any]: def slow_mock_agent(): """Create a slow mock agent for testing concurrency.""" agent = SlowAgentMock(delay=0.3) + saved_sessions: dict[str, Any] = {} # Set up pool mock with async storage methods pool = Mock() @@ -110,7 +118,12 @@ def slow_mock_agent(): # Storage needs to be properly mocked with async methods storage = Mock() - storage.save_session = AsyncMock() + + async def save_session(session_data: Any) -> None: + saved_sessions[session_data.session_id] = session_data + + storage.save_session = AsyncMock(side_effect=save_session) + storage.log_session = AsyncMock() storage.log_message = AsyncMock() pool.storage = storage @@ -118,6 +131,9 @@ def slow_mock_agent(): pool.todos.on_change = None pool.skill_commands = None + pool.sessions = Mock() + pool.sessions.store = None + # CRITICAL: all_agents must return a real dict to avoid Mock issues pool.all_agents = {agent.name: agent} @@ -134,6 +150,15 @@ def slow_mock_agent(): # Set up storage agent.storage = storage + conversation = Mock() + conversation.chat_messages = [] + agent.conversation = conversation + + async def load_session(session_id: str) -> Any: + return saved_sessions.get(session_id) + + agent.load_session = AsyncMock(side_effect=load_session) + return agent @@ -356,3 +381,306 @@ async def send_message_with_content(content: str, msg_id: str): # Verify the agent was called 3 times agent_mock = cast(SlowAgentMock, state.agent) assert agent_mock.run_stream_call_count == 3 + + @pytest.mark.asyncio + async def test_two_sessions_back_to_back_b_not_blocked( + self, + concurrent_test_state: ServerState, + sample_message_request: MessageRequest, + ) -> None: + """Test that Session B can start while Session A is still running. + + When two sessions start their turns close together, Session B should NOT + wait for Session A's full turn to finish. The per-session locking only + serializes messages within the same session; different sessions run in + parallel. This is verified by checking that the total wall time for both + sessions is less than 2 * agent_delay (which would be the sequential time). + """ + state = concurrent_test_state + session_id_a = "test-session-back-to-back-a" + session_id_b = "test-session-back-to-back-b" + + # Create both sessions + await state.ensure_session(session_id_a) + await state.ensure_session(session_id_b) + + agent_mock = cast(SlowAgentMock, state.agent) + agent_delay = agent_mock.delay + + start = asyncio.get_event_loop().time() + + # Start Session A processing + task_a = asyncio.create_task(_process_message(session_id_a, sample_message_request, state)) + + # After a short delay, start Session B processing + await asyncio.sleep(0.05) + task_b = asyncio.create_task(_process_message(session_id_b, sample_message_request, state)) + + # Wait for both to complete + await asyncio.gather(task_a, task_b) + + elapsed = asyncio.get_event_loop().time() - start + + # Both sessions should have their messages (2 each: user + assistant) + assert len(state.messages[session_id_a]) == 2 + assert len(state.messages[session_id_b]) == 2 + + # Agent was called twice (once per session) + assert agent_mock.run_stream_call_count == 2 + + # Timing assertion: if they ran sequentially, total would be >= 2 * delay. + # Since they ran concurrently, total should be < 2 * delay. + # We use a generous margin to avoid flaky CI failures. + assert elapsed < 2 * agent_delay, ( + f"Sessions did not run concurrently: elapsed={elapsed:.3f}s " + f"but sequential would be ~{2 * agent_delay:.3f}s" + ) + + @pytest.mark.asyncio + async def test_in_flight_uses_snapshot_not_live_fields( + self, + concurrent_test_state: ServerState, + sample_message_request: MessageRequest, + ) -> None: + """Test that an in-flight run uses captured snapshot values, not live fields. + + When Session A is mid-run and the shared agent's model_name is mutated + (simulating Session B binding to a different model), Session A's ongoing + run must continue using the snapshot it captured at the start — not the + new live value. This is the core guarantee of the RunSnapshot mechanism. + """ + state = concurrent_test_state + session_id_a = "test-session-snapshot-a" + + # Create session A + await state.ensure_session(session_id_a) + + agent_mock = cast(SlowAgentMock, state.agent) + + # Start Session A processing (model_name is "test-model" at this point) + task_a = asyncio.create_task(_process_message(session_id_a, sample_message_request, state)) + + # Wait briefly for Session A to have captured its snapshot + await asyncio.sleep(0.05) + + # Mutate the live agent model_name (simulating Session B binding) + agent_mock.model_name = "different-model" + + # Wait for Session A to finish + await task_a + + # Verify Session A used the original snapshot model, not the mutated live value + assert agent_mock.model_names_at_call.get(session_id_a) == "test-model", ( + f"Session A should have seen snapshot model_name='test-model', " + f"but got '{agent_mock.model_names_at_call.get(session_id_a)}'" + ) + + # The live value should still be the mutated one + assert agent_mock.model_name == "different-model" + + # Session A should have completed successfully + assert len(state.messages[session_id_a]) == 2 + + @pytest.mark.asyncio + async def test_read_only_route_during_active_turn( + self, + concurrent_test_state: ServerState, + sample_message_request: MessageRequest, + ) -> None: + """Test that a read-only route returns immediately during an active turn. + + When Session A has an in-flight agent run, calling get_or_load_session + for a different, already-cached Session B should return without waiting + for A's turn to complete. This verifies that read-only (cache-hit) paths + do not acquire agent_lock or block on in-flight turns. + """ + state = concurrent_test_state + session_id_a = "test-session-readonly-a" + session_id_b = "test-session-readonly-b" + + # Create both sessions and ensure B is cached + await state.ensure_session(session_id_a) + await state.ensure_session(session_id_b) + + # Ensure B has messages so it passes the agent_has_correct_session check + from agentpool_server.opencode_server.models import ( + AssistantMessage, + MessagePath, + MessageTime, + MessageWithParts, + TextPart, + ) + from agentpool.utils import identifiers as identifier + + msg_id = identifier.ascending("message") + part_id = identifier.ascending("part") + assistant_msg = AssistantMessage( + id=msg_id, + session_id=session_id_b, + parent_id="", + model_id="test", + provider_id="test", + mode="test", + agent="default", + path=MessagePath(cwd=state.working_dir, root=state.working_dir), + time=MessageTime(created=0), + ) + state.messages[session_id_b] = [ + MessageWithParts( + info=assistant_msg, + parts=[TextPart(id=part_id, message_id=msg_id, session_id=session_id_b, text="hi")], + ) + ] + + # Make agent report session_id_b so the fast-path triggers + state.agent.session_id = session_id_b + + # Start a slow agent turn on Session A + task_a = asyncio.create_task(_process_message(session_id_a, sample_message_request, state)) + + # Wait briefly for A's turn to start + await asyncio.sleep(0.05) + + # Now call get_or_load_session for B — this should return immediately + # because B is cached and agent has correct session loaded (fast path) + start = asyncio.get_event_loop().time() + from agentpool_server.opencode_server.routes.session_routes import get_or_load_session + + session_b = await get_or_load_session(state, session_id_b) + elapsed = asyncio.get_event_loop().time() - start + + # The read should have returned very quickly (< 0.1s), not waited for A + assert elapsed < 0.1, ( + f"get_or_load_session for cached session B took {elapsed:.3f}s, " + f"should return immediately from cache" + ) + assert session_b is not None + assert session_b.id == session_id_b + + # Clean up: wait for A to finish + await task_a + + @pytest.mark.asyncio + async def test_conversation_isolation_concurrent_sessions( + self, + concurrent_test_state: ServerState, + ) -> None: + """Test that messages from one session never leak into another. + + Two sessions run concurrently; after both finish, messages from Session A + must not appear in Session B's `state.messages` and vice versa. Also + verifies that per-session `MessageHistory` instances are separate objects. + """ + state = concurrent_test_state + session_id_a = "test-session-isolation-a" + session_id_b = "test-session-isolation-b" + + # Create both sessions + await state.ensure_session(session_id_a) + await state.ensure_session(session_id_b) + + # Build unique requests so we can distinguish message content + req_a = MessageRequest( + parts=[TextPartInput(text="Message for Session A")], + agent="default", + ) + req_b = MessageRequest( + parts=[TextPartInput(text="Message for Session B")], + agent="default", + ) + + # Process both sessions concurrently + await asyncio.gather( + _process_message(session_id_a, req_a, state), + _process_message(session_id_b, req_b, state), + ) + + # Each session should have exactly 2 messages (user + assistant) + assert len(state.messages[session_id_a]) == 2 + assert len(state.messages[session_id_b]) == 2 + + # Verify content isolation: Session A's messages must not appear in B + texts_b = [ + part.text + for msg in state.messages[session_id_b] + for part in msg.parts + if isinstance(part, TextPart) + ] + assert not any("Session A" in t for t in texts_b), ( + "Session A content leaked into Session B's messages" + ) + + # Verify content isolation: Session B's messages must not appear in A + texts_a = [ + part.text + for msg in state.messages[session_id_a] + for part in msg.parts + if isinstance(part, TextPart) + ] + assert not any("Session B" in t for t in texts_a), ( + "Session B content leaked into Session A's messages" + ) + + # Per-session MessageHistory instances must be separate objects + assert ( + state.session_conversations[session_id_a] + is not state.session_conversations[session_id_b] + ), "Session A and B share the same MessageHistory instance" + + @pytest.mark.asyncio + async def test_interrupt_session_a_does_not_cancel_session_b( + self, + concurrent_test_state: ServerState, + sample_message_request: MessageRequest, + ) -> None: + """Test that cancelling Session A's run does not affect Session B. + + Two sessions start overlapping runs; `cancel_session_run` on Session A + cancels A but Session B completes successfully with all messages intact. + """ + state = concurrent_test_state + session_id_a = "test-session-interrupt-a" + session_id_b = "test-session-interrupt-b" + + # Create both sessions + await state.ensure_session(session_id_a) + await state.ensure_session(session_id_b) + + # Start Session A processing + task_a = asyncio.create_task(_process_message(session_id_a, sample_message_request, state)) + + # Wait for A to have started (snapshot captured + task registered) + await asyncio.sleep(0.05) + + # Register A's task so cancel_session_run can find it + state.register_active_run(session_id_a, task_a) + + # Start Session B processing + task_b = asyncio.create_task(_process_message(session_id_b, sample_message_request, state)) + + # Wait for B to have started + await asyncio.sleep(0.05) + + # Cancel only Session A + cancelled = state.cancel_session_run(session_id_a) + assert cancelled, "cancel_session_run should return True for active Session A" + + # Session A should be done (cancelled or completed after handling CancelledError) + # _process_message_locked catches CancelledError internally, but the + # task may still surface it depending on timing. Either way, it must be done. + try: + await task_a + except asyncio.CancelledError: + pass + assert task_a.done() + + # Wait for Session B to complete (with timeout) + await asyncio.wait_for(task_b, timeout=2.0) + + # Session B should have completed successfully + assert len(state.messages[session_id_b]) == 2, ( + "Session B should have 2 messages (user + assistant) after completing" + ) + + # Session B was NOT cancelled by Session A's interrupt + assert not task_b.cancelled(), "Session B should not have been cancelled" diff --git a/tests/servers/opencode_server/test_diagnostic.py b/tests/servers/opencode_server/test_diagnostic.py new file mode 100644 index 000000000..89a72ab2f --- /dev/null +++ b/tests/servers/opencode_server/test_diagnostic.py @@ -0,0 +1,162 @@ +"""Tests for GET /global/diagnostic endpoint.""" + +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING, Any +from unittest.mock import Mock + +import pytest + +from agentpool_server.opencode_server.routes.global_routes import ( + GlobalEventFactory, + VERSION, +) +from agentpool_server.opencode_server.state import ServerState + + +if TYPE_CHECKING: + from httpx import AsyncClient + + +# ============================================================================= +# _MockState for unit tests (reuses pattern from test_global_event.py) +# ============================================================================= + + +class _MockState: + """Minimal ServerState-like object for diagnostic endpoint tests.""" + + def __init__(self, working_dir: str | None = "/tmp/test_wd") -> None: + self.working_dir = working_dir + self.event_subscribers: list[asyncio.Queue[Any]] = [] + self._event_factory: GlobalEventFactory | None = None + + def get_event_factory(self) -> GlobalEventFactory: + if self._event_factory is None: + from agentpool_storage.opencode_provider import helpers + + self._event_factory = GlobalEventFactory( + directory=self.working_dir or "", + project=helpers.compute_project_id(self.working_dir or ""), + ) + return self._event_factory + + +# ============================================================================= +# Integration tests using async_client (real FastAPI app) +# ============================================================================= + + +@pytest.mark.anyio +async def test_diagnostic_returns_200(async_client: AsyncClient) -> None: + """GET /global/diagnostic returns 200 status code.""" + response = await async_client.get("/global/diagnostic") + assert response.status_code == 200 + + +@pytest.mark.anyio +async def test_diagnostic_has_required_fields(async_client: AsyncClient) -> None: + """GET /global/diagnostic returns JSON with directory, project, subscribers, serverVersion.""" + response = await async_client.get("/global/diagnostic") + data = response.json() + assert "directory" in data + assert "project" in data + assert "subscribers" in data + assert "serverVersion" in data + + +@pytest.mark.anyio +async def test_diagnostic_directory_matches_working_dir( + async_client: AsyncClient, + server_state: ServerState, +) -> None: + """Directory field equals state.working_dir.""" + response = await async_client.get("/global/diagnostic") + data = response.json() + assert data["directory"] == server_state.working_dir + + +@pytest.mark.anyio +async def test_diagnostic_subscribers_is_non_negative_integer( + async_client: AsyncClient, +) -> None: + """Subscribers field is an integer >= 0.""" + response = await async_client.get("/global/diagnostic") + data = response.json() + assert isinstance(data["subscribers"], int) + assert data["subscribers"] >= 0 + + +@pytest.mark.anyio +async def test_diagnostic_server_version_matches_constant( + async_client: AsyncClient, +) -> None: + """ServerVersion field matches the VERSION constant from global_routes.""" + response = await async_client.get("/global/diagnostic") + data = response.json() + assert data["serverVersion"] == VERSION + + +@pytest.mark.anyio +async def test_diagnostic_project_is_computed(async_client: AsyncClient) -> None: + """Project field is computed via compute_project_id.""" + from agentpool_storage.opencode_provider.helpers import compute_project_id + + response = await async_client.get("/global/diagnostic") + data = response.json() + # project should be a non-empty string + assert isinstance(data["project"], str) + assert len(data["project"]) > 0 + + +# ============================================================================= +# Edge case: working_dir=None +# ============================================================================= + + +def _make_state_with_none_working_dir() -> ServerState: + """Create a ServerState with working_dir=None for edge-case testing.""" + mock_env = Mock() + mock_env.get_fs = Mock(return_value=Mock()) + mock_agent = Mock() + mock_agent.env = mock_env + return ServerState(working_dir=None, agent=mock_agent) # type: ignore[arg-type] + + +@pytest.mark.anyio +async def test_diagnostic_working_dir_none_returns_directory_null() -> None: + """When working_dir is None, directory is null (not crashing).""" + from fastapi import FastAPI + from httpx import ASGITransport, AsyncClient + + from agentpool_server.opencode_server.dependencies import get_state + from agentpool_server.opencode_server.routes.global_routes import router as global_router + + state = _make_state_with_none_working_dir() + app = FastAPI() + app.include_router(global_router) + app.dependency_overrides[get_state] = lambda: state + + transport = ASGITransport(app=app) + async with AsyncClient(transport=transport, base_url="http://test") as ac: + response = await ac.get("/global/diagnostic") + assert response.status_code == 200 + data = response.json() + assert data["directory"] is None + + +# ============================================================================= +# Subscriber count reflects real subscribers +# ============================================================================= + + +@pytest.mark.anyio +async def test_diagnostic_subscribers_reflects_real_count( + async_client: AsyncClient, + server_state: ServerState, +) -> None: + """Subscribers count equals len(state.event_subscribers).""" + response = await async_client.get("/global/diagnostic") + data = response.json() + assert data["subscribers"] == len(server_state.event_subscribers) diff --git a/tests/servers/opencode_server/test_ensure_session.py b/tests/servers/opencode_server/test_ensure_session.py index 881c08335..e8069cd7c 100644 --- a/tests/servers/opencode_server/test_ensure_session.py +++ b/tests/servers/opencode_server/test_ensure_session.py @@ -9,6 +9,8 @@ from agentpool.agents.base_agent import BaseAgent from agentpool_server.opencode_server.models import ( Session, + SessionIdleEvent, + SessionStatusEvent, TimeCreatedUpdated, ) from agentpool_server.opencode_server.state import ServerState @@ -21,6 +23,7 @@ def create_mock_agent() -> MagicMock: agent.agent_pool = MagicMock() agent.agent_pool.manifest.config_file_path = "test_config.yml" agent.agent_pool.storage.save_session = AsyncMock() + agent.agent_pool.sessions.store = None agent.env = MagicMock() agent.env.cwd = "/test/dir" return agent @@ -148,6 +151,34 @@ async def test_ensure_session_caches_in_memory(mock_state: ServerState) -> None: assert mock_state.input_providers[session_id] is mock_provider +@pytest.mark.asyncio +async def test_ensure_session_broadcasts_idle_events(mock_state: ServerState) -> None: + """Test that ensure_session broadcasts both status and idle events.""" + session_id = "test_session_idle_event" + + with ( + patch("agentpool_server.opencode_server.converters.opencode_to_session_data"), + patch("agentpool_server.opencode_server.input_provider.OpenCodeInputProvider"), + patch.object(mock_state, "broadcast_event", new=AsyncMock()) as mock_broadcast, + ): + await mock_state.ensure_session(session_id) + + status_events = [ + call.args[0] + for call in mock_broadcast.await_args_list + if isinstance(call.args[0], SessionStatusEvent) + ] + idle_events = [ + call.args[0] + for call in mock_broadcast.await_args_list + if isinstance(call.args[0], SessionIdleEvent) + ] + assert len(status_events) == 1 + assert len(idle_events) == 1 + assert status_events[0].properties.status.type == "idle" + assert idle_events[0].properties.session_id == session_id + + @pytest.mark.asyncio async def test_ensure_session_creates_input_provider(mock_state: ServerState) -> None: """Test that ensure_session creates and stores an OpenCodeInputProvider.""" diff --git a/tests/servers/opencode_server/test_file_diff_models.py b/tests/servers/opencode_server/test_file_diff_models.py new file mode 100644 index 000000000..d6514e2ac --- /dev/null +++ b/tests/servers/opencode_server/test_file_diff_models.py @@ -0,0 +1,156 @@ +"""Tests for FileDiff model alignment with OpenCode v1.4.0+ SnapshotFileDiff schema.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from agentpool_server.opencode_server.models.common import FileDiff, FileDiffStatus + + +# Minimal stand-in for FileChange used by from_file_change() +@dataclass +class _StubFileChange: + path: str + old_content: str | None + new_content: str | None + operation: str + + def to_unified_diff(self) -> str: + """Return a minimal unified diff string.""" + from agentpool.utils.diffs import compute_unified_diff + + return compute_unified_diff( + self.old_content or "", + self.new_content or "", + fromfile=f"a/{self.path}", + tofile=f"b/{self.path}", + ) + + +def test_filediff_schema_has_patch_no_before_after(): + """FileDiff must serialize with file/patch/additions/deletions/status — no before/after.""" + diff = FileDiff( + file="src/main.py", + patch="--- a/src/main.py\n+++ b/src/main.py\n@@ -1 +1 @@\n-old\n+new\n", + additions=1, + deletions=1, + status="modified", + ) + data = diff.model_dump() + assert "file" in data + assert "patch" in data + assert "additions" in data + assert "deletions" in data + assert "status" in data + # Must NOT contain before/after/to/from keys + assert "before" not in data + assert "after" not in data + assert "to" not in data + assert "from" not in data + + +def test_filediff_schema_camelCase_serialization(): + """FileDiff camelCase serialization must not leak before/after.""" + diff = FileDiff( + file="app.ts", + patch="patch content", + additions=5, + deletions=3, + status="added", + ) + data = diff.model_dump(by_alias=True) + assert "file" in data + assert "patch" in data + assert "before" not in data + assert "after" not in data + + +def test_filediff_from_file_change_populates_patch(): + """from_file_change() must store unified diff in 'patch' field.""" + change = _StubFileChange( + path="hello.txt", + old_content="hello world\n", + new_content="hello universe\n", + operation="edit", + ) + diff = FileDiff.from_file_change(change) + assert diff.patch is not None + assert len(diff.patch) > 0 + # The patch should contain unified diff markers + assert "---" in diff.patch + assert "+++" in diff.patch + assert diff.file == "hello.txt" + assert diff.status == "modified" + # Must NOT have before/after + assert not hasattr(diff, "before") + assert not hasattr(diff, "after") + + +def test_filediff_from_file_change_create_operation(): + """from_file_change() with 'create' operation must set status='added'.""" + change = _StubFileChange( + path="new_file.py", + old_content=None, + new_content="print('hello')\n", + operation="create", + ) + diff = FileDiff.from_file_change(change) + assert diff.status == "added" + assert diff.patch is not None + + +def test_filediff_from_file_change_delete_operation(): + """from_file_change() with 'delete' operation must set status='deleted'.""" + change = _StubFileChange( + path="old_file.py", + old_content="print('bye')\n", + new_content=None, + operation="delete", + ) + diff = FileDiff.from_file_change(change) + assert diff.status == "deleted" + assert diff.patch is not None + + +def test_filediff_from_file_change_write_operation(): + """from_file_change() with 'write' operation must set status='modified'.""" + change = _StubFileChange( + path="config.json", + old_content='{"key": "old"}\n', + new_content='{"key": "new"}\n', + operation="write", + ) + diff = FileDiff.from_file_change(change) + assert diff.status == "modified" + + +def test_filediff_additions_deletions_count(): + """additions and deletions must be counted from the unified diff.""" + change = _StubFileChange( + path="test.txt", + old_content="line1\nline2\nline3\n", + new_content="line1\nmodified2\nline3\nadded4\n", + operation="edit", + ) + diff = FileDiff.from_file_change(change) + assert diff.additions > 0 + assert diff.deletions > 0 + + +def test_filediff_patch_default_none(): + """FileDiff constructed without patch must default to None.""" + diff = FileDiff(file="empty.txt", additions=0, deletions=0) + assert diff.patch is None + + +def test_filediff_status_optional(): + """FileDiff status must be optional (defaults to None).""" + diff = FileDiff(file="test.py", patch="some patch", additions=1, deletions=0) + assert diff.status is None + + +def test_filediff_status_literal_values(): + """FileDiff status must accept only 'added', 'deleted', 'modified'.""" + for status_val in ("added", "deleted", "modified"): + diff = FileDiff(file="f.py", patch="p", additions=0, deletions=0, status=status_val) + assert diff.status == status_val diff --git a/tests/servers/opencode_server/test_global_compat_routes.py b/tests/servers/opencode_server/test_global_compat_routes.py new file mode 100644 index 000000000..321080f38 --- /dev/null +++ b/tests/servers/opencode_server/test_global_compat_routes.py @@ -0,0 +1,153 @@ +"""Tests for OpenCode 1.4.4+ global compatibility routes. + +Covers: +- GET /global/config (delegates to /config) +- PATCH /global/config (delegates to /config) +- POST /global/dispose (stub no-op) +- POST /global/upgrade (stub no-op) +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING +from unittest.mock import AsyncMock, Mock + +from fastapi import FastAPI +from fastapi.testclient import TestClient +import pytest + +from agentpool.models.manifest import AgentsManifest +from agentpool.storage import StorageManager +from agentpool.utils.streams import FileOpsTracker +from agentpool.utils.todos import TodoTracker +from agentpool_server.opencode_server.dependencies import get_state +from agentpool_server.opencode_server.routes.config_routes import router as config_router +from agentpool_server.opencode_server.routes.global_routes import router as global_router +from agentpool_server.opencode_server.state import ServerState + + +if TYPE_CHECKING: + from pathlib import Path + + +@pytest.fixture +def _server_state(tmp_path: Path) -> ServerState: + """Build a ServerState with a mock agent for config route tests.""" + from agentpool_config.storage import MemoryStorageConfig, StorageConfig + + storage_manager = StorageManager(config=StorageConfig(providers=[MemoryStorageConfig()])) + file_ops = FileOpsTracker() + todos = TodoTracker() + manifest = AgentsManifest(config_file_path="/tmp/test-pool") + + pool = Mock() + pool.storage = storage_manager + pool.file_ops = file_ops + pool.todos = todos + pool.manifest = manifest + + env = Mock() + env.cwd = str(tmp_path) + + agent = Mock() + agent.name = "test-agent" + agent.env = env + agent._input_provider = None + agent.agent_pool = pool + agent.storage = storage_manager + agent.get_available_models = AsyncMock(return_value=[]) + + return ServerState(working_dir=str(tmp_path), agent=agent) + + +@pytest.fixture +def client(_server_state: ServerState) -> TestClient: + """Create a test client with both global and config routers.""" + app = FastAPI() + app.include_router(config_router) + app.include_router(global_router) + app.dependency_overrides[get_state] = lambda: _server_state + return TestClient(app) + + +class TestGlobalConfigRoutes: + """Tests for GET/PATCH /global/config.""" + + def test_get_global_config_returns_config(self, client: TestClient) -> None: + """GET /global/config should return a Config object.""" + resp = client.get("/global/config") + assert resp.status_code == 200 + data = resp.json() + # Config should have at least keybinds and watcher fields + assert "keybinds" in data or "model" in data + + def test_get_global_config_matches_get_config(self, client: TestClient) -> None: + """GET /global/config should return the same data as GET /config.""" + global_resp = client.get("/global/config") + config_resp = client.get("/config") + assert global_resp.status_code == 200 + assert config_resp.status_code == 200 + assert global_resp.json() == config_resp.json() + + def test_patch_global_config_updates_model(self, client: TestClient) -> None: + """PATCH /global/config should update config fields.""" + # First, get current config + get_resp = client.get("/global/config") + assert get_resp.status_code == 200 + + # Patch the theme + patch_resp = client.patch("/global/config", json={"theme": "dark"}) + assert patch_resp.status_code == 200 + data = patch_resp.json() + assert data.get("theme") == "dark" + + def test_patch_global_config_matches_patch_config(self, client: TestClient) -> None: + """PATCH /global/config should behave identically to PATCH /config.""" + # Set via /config + r1 = client.patch("/config", json={"theme": "light"}) + assert r1.status_code == 200 + + # Set via /global/config + r2 = client.patch("/global/config", json={"theme": "dark"}) + assert r2.status_code == 200 + + # Both should update the same underlying state + final = client.get("/global/config") + assert final.json().get("theme") == "dark" + + +class TestGlobalDisposeRoute: + """Tests for POST /global/dispose.""" + + def test_global_dispose_returns_success(self, client: TestClient) -> None: + """POST /global/dispose should return a success stub response.""" + resp = client.post("/global/dispose") + assert resp.status_code == 200 + data = resp.json() + assert data["success"] is True + assert "no-op" in data["message"] + + def test_global_dispose_does_not_crash_server(self, client: TestClient) -> None: + """Server should still respond after /global/dispose.""" + client.post("/global/dispose") + # Server should still work + resp = client.get("/global/health") + assert resp.status_code == 200 + + +class TestGlobalUpgradeRoute: + """Tests for POST /global/upgrade.""" + + def test_global_upgrade_returns_stub(self, client: TestClient) -> None: + """POST /global/upgrade should return a stub response.""" + resp = client.post("/global/upgrade") + assert resp.status_code == 200 + data = resp.json() + assert data["success"] is True + assert data["upgraded"] is False + + def test_global_upgrade_does_not_crash_server(self, client: TestClient) -> None: + """Server should still respond after /global/upgrade.""" + client.post("/global/upgrade") + resp = client.get("/global/health") + assert resp.status_code == 200 diff --git a/tests/servers/opencode_server/test_global_event.py b/tests/servers/opencode_server/test_global_event.py new file mode 100644 index 000000000..b6b78f9f4 --- /dev/null +++ b/tests/servers/opencode_server/test_global_event.py @@ -0,0 +1,1464 @@ +"""Tests for _serialize_event, GlobalEvent model, and GlobalEventFactory.""" + +from __future__ import annotations + +import asyncio +import contextlib +import json +from pathlib import Path +from typing import TYPE_CHECKING, Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from agentpool_server.opencode_server.models import GlobalEvent +from agentpool_server.opencode_server.models.app import ProjectTime +from agentpool_server.opencode_server.models.common import TimeCreated +from agentpool_server.opencode_server.models.events import ( + CommandExecutedEvent, + Event, + FileEditedEvent, + FileWatcherUpdatedEvent, + LspClientDiagnosticsEvent, + LspUpdatedEvent, + McpToolsChangedEvent, + MessageRemovedEvent, + MessageUpdatedEvent, + PartDeltaEvent, + PartRemovedEvent, + PartUpdatedEvent, + PermissionRequestEvent, + PermissionResolvedEvent, + PermissionUpdatedEvent, + Project, + ProjectUpdatedEvent, + PtyCreatedEvent, + PtyDeletedEvent, + PtyExitedEvent, + PtyUpdatedEvent, + QuestionAskedEvent, + QuestionRejectedEvent, + QuestionRepliedEvent, + ServerConnectedEvent, + ServerHeartbeatEvent, + SessionCompactedEvent, + SessionCreatedEvent, + SessionDeletedEvent, + SessionDiffEvent, + SessionErrorEvent, + SessionIdleEvent, + SessionStatusEvent, + SessionUpdatedEvent, + Todo, + TodoUpdatedEvent, + TuiCommandExecuteEvent, + TuiPromptAppendEvent, + TuiSessionSelectEvent, + TuiToastShowEvent, + VcsBranchUpdatedEvent, +) +from agentpool_server.opencode_server.models.message import ( + UserMessage, +) +from agentpool_server.opencode_server.models.parts import Part, TextPart # noqa: TC001 +from agentpool_server.opencode_server.models.pty import PtyInfo +from agentpool_server.opencode_server.models.question import ( + QuestionInfo, + QuestionOption, +) +from agentpool_server.opencode_server.models.session import ( + Session, + TimeCreatedUpdated, +) +from agentpool_server.opencode_server.routes.global_routes import ( + GlobalEventFactory, + _event_generator, + _extract_session_id, + _serialize_event, +) +from agentpool_server.opencode_server.state import ServerState + + +if TYPE_CHECKING: + from httpx import AsyncClient + + +# ============================================================================= +# _serialize_event baseline tests +# ============================================================================= + + +def test_serialize_event_session_id_injection() -> None: + """SessionId is injected at top level when event has a session.""" + event = SessionStatusEvent.create(session_id="abc", status_type="busy") + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert data["sessionId"] == "abc" + + +def test_serialize_event_no_session_id() -> None: + """No sessionId key when event has no session.""" + event = ServerConnectedEvent() + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert "sessionId" not in data + + +def test_serialize_event_wrap_payload_true() -> None: + """wrap_payload=True nests event data under 'payload' key.""" + event = ServerConnectedEvent() + result = _serialize_event(event, wrap_payload=True) + data = json.loads(result) + assert "payload" in data + assert data["payload"]["type"] == "server.connected" + + +def test_serialize_event_unicode_preserved() -> None: + r"""Unicode characters are preserved (not \uXXXX escaped).""" + event = SessionStatusEvent.create(session_id="你好", status_type="idle") + result = _serialize_event(event, wrap_payload=False) + assert "你好" in result + assert "\\u" not in result + + +def test_serialize_event_camel_case_aliases() -> None: + """Model fields use camelCase aliases in serialized output.""" + event = SessionStatusEvent.create(session_id="abc", status_type="busy") + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + # session_id → sessionID via convert() alias generator + props = data["properties"] + assert "sessionID" in props + + +# ============================================================================= +# GlobalEvent model and GlobalEventFactory tests +# ============================================================================= + + +def test_global_event_model_construction() -> None: + """GlobalEvent stores directory, project, and payload correctly.""" + payload = {"type": "test"} + event = GlobalEvent(directory="/tmp/test", project="abc123", payload=payload) + dumped = event.model_dump(by_alias=True, exclude_none=True) + assert dumped["directory"] == "/tmp/test" + assert dumped["project"] == "abc123" + assert dumped["payload"] == payload + + +def test_global_event_workspace_omitted_when_none() -> None: + """Workspace is excluded from output when not provided.""" + event = GlobalEvent(directory="/tmp/test", project="abc123", payload={}) + dumped = event.model_dump(by_alias=True, exclude_none=True) + assert "workspace" not in dumped + + +def test_global_event_factory_wrap() -> None: + """Factory.wrap() produces JSON with directory, project, payload.""" + factory = GlobalEventFactory(directory="/tmp", project="abc") + event = SessionStatusEvent.create(session_id="sid1", status_type="idle") + result = factory.wrap(event) + data = json.loads(result) + assert data["directory"] == "/tmp" + assert data["project"] == "abc" + assert isinstance(data["payload"], dict) + + +def test_global_event_factory_session_id_in_payload() -> None: + """SessionId is injected inside payload by GlobalEventFactory.wrap.""" + factory = GlobalEventFactory(directory="/tmp", project="abc") + event = SessionStatusEvent.create(session_id="sid1", status_type="idle") + result = factory.wrap(event) + data = json.loads(result) + assert data["payload"]["sessionId"] == "sid1" + + +def test_global_event_factory_unicode_preserved() -> None: + """Factory.wrap() preserves Unicode characters in output.""" + factory = GlobalEventFactory(directory="/tmp", project="abc") + event = SessionStatusEvent.create(session_id="会话1", status_type="busy") + result = factory.wrap(event) + assert "会话1" in result + assert "\\u" not in result + + +def test_global_event_model_project_global() -> None: + """GlobalEvent with project='global' preserves the value.""" + event = GlobalEvent(directory="/tmp", project="global", payload={}) + dumped = event.model_dump(by_alias=True, exclude_none=True) + assert dumped["project"] == "global" + + +def test_global_event_factory_wrap_returns_string() -> None: + """Factory.wrap() returns a str (JSON string).""" + factory = GlobalEventFactory(directory="/tmp", project="abc") + event = ServerHeartbeatEvent() + result = factory.wrap(event) + assert isinstance(result, str) + + +# ============================================================================= +# _event_generator integration tests +# ============================================================================= + + +class _MockState: + """Minimal ServerState-like object for _event_generator tests.""" + + def __init__(self, working_dir: str = "/tmp/test_wd") -> None: + self.working_dir = working_dir + self.event_subscribers: list[asyncio.Queue[Event]] = [] + self._event_factory: GlobalEventFactory | None = None + self._first_subscriber_triggered = False + self.on_first_subscriber: Any = None + + def get_event_factory(self) -> GlobalEventFactory: + if self._event_factory is None: + from agentpool_storage.opencode_provider import helpers + + self._event_factory = GlobalEventFactory( + directory=self.working_dir, + project=helpers.compute_project_id(self.working_dir), + ) + return self._event_factory + + def create_background_task(self, coro: Any, name: str = "") -> asyncio.Task[Any]: + return asyncio.ensure_future(coro) + + +async def _collect_events( + state: _MockState, + wrap_payload: bool, + events_to_send: list[Event], +) -> list[dict[str, Any]]: + """Collect SSE items from _event_generator with given events.""" + results: list[dict[str, Any]] = [] + gen = _event_generator(state, wrap_payload=wrap_payload) + # Get the initial connected event + item = await gen.__anext__() + results.append(json.loads(item["data"])) + # Send additional events through the queue + queue = state.event_subscribers[-1] + for event in events_to_send: + await queue.put(event) + item = await gen.__anext__() + results.append(json.loads(item["data"])) + return results + + +@pytest.mark.anyio +async def test_global_event_server_connected_is_payload_wrapped() -> None: + """First event from /global/event keeps only the payload wrapper.""" + state = _MockState() + events = await _collect_events(state, wrap_payload=True, events_to_send=[]) + # Only the connected event + assert len(events) == 1 + connected = events[0] + assert connected["payload"]["type"] == "server.connected" + assert "directory" not in connected + assert "project" not in connected + + +@pytest.mark.anyio +async def test_global_event_wraps_regular_events_in_envelope() -> None: + """/global/event wraps SessionStatusEvent in GlobalEvent envelope.""" + state = _MockState() + session_evt = SessionStatusEvent.create(session_id="s1", status_type="busy") + events = await _collect_events(state, wrap_payload=True, events_to_send=[session_evt]) + assert len(events) == 2 + wrapped = events[1] + assert "directory" in wrapped + assert "project" in wrapped + assert "payload" in wrapped + assert wrapped["payload"]["type"] == "session.status" + + +@pytest.mark.anyio +async def test_global_event_heartbeat_is_payload_wrapped() -> None: + """/global/event keeps only the payload wrapper for heartbeat.""" + state = _MockState() + hb = ServerHeartbeatEvent() + events = await _collect_events(state, wrap_payload=True, events_to_send=[hb]) + assert len(events) == 2 + heartbeat = events[1] + assert heartbeat["payload"]["type"] == "server.heartbeat" + assert "directory" not in heartbeat + assert "project" not in heartbeat + + +@pytest.mark.anyio +async def test_event_endpoint_all_events_are_bare() -> None: + """/event sends connected, heartbeat, and session events all as bare.""" + state = _MockState() + hb = ServerHeartbeatEvent() + session_evt = SessionStatusEvent.create(session_id="s2", status_type="idle") + events = await _collect_events(state, wrap_payload=False, events_to_send=[hb, session_evt]) + assert len(events) == 3 + for evt in events: + # None should have envelope wrapper keys + assert "directory" not in evt + assert "project" not in evt + assert "payload" not in evt + assert events[0]["type"] == "server.connected" + assert events[1]["type"] == "server.heartbeat" + assert events[2]["type"] == "session.status" + + +@pytest.mark.anyio +async def test_global_events_have_no_session_id() -> None: + """ServerConnectedEvent and ServerHeartbeatEvent lack sessionId.""" + state = _MockState() + hb = ServerHeartbeatEvent() + events = await _collect_events(state, wrap_payload=True, events_to_send=[hb]) + assert "sessionId" not in events[0] # top-level envelope + assert "sessionId" not in events[0]["payload"] # server.connected payload + assert "sessionId" not in events[1] # top-level envelope + assert "sessionId" not in events[1]["payload"] # server.heartbeat payload + + +@pytest.mark.anyio +async def test_global_event_directory_matches_working_dir() -> None: + """Envelope directory field matches the server working directory.""" + wd = "/custom/working/dir" + state = _MockState(working_dir=wd) + session_evt = SessionStatusEvent.create(session_id="s3", status_type="retry") + events = await _collect_events(state, wrap_payload=True, events_to_send=[session_evt]) + wrapped = events[1] + assert wrapped["directory"] == wd + + +@pytest.mark.anyio +async def test_multiple_events_maintain_correct_wrapping() -> None: + """Sequence of wrapped events all have correct format.""" + state = _MockState() + session_evt = SessionStatusEvent.create(session_id="s4", status_type="busy") + hb = ServerHeartbeatEvent() + session_evt2 = SessionStatusEvent.create(session_id="s5", status_type="idle") + events = await _collect_events( + state, + wrap_payload=True, + events_to_send=[session_evt, hb, session_evt2], + ) + assert len(events) == 4 + # [0] connected — payload wrapped, no routing metadata + assert events[0]["payload"]["type"] == "server.connected" + assert "directory" not in events[0] + # [1] session status — wrapped + assert "payload" in events[1] + assert events[1]["payload"]["type"] == "session.status" + # [2] heartbeat — payload wrapped, no routing metadata + assert events[2]["payload"]["type"] == "server.heartbeat" + assert "directory" not in events[2] + # [3] session status — wrapped + assert "payload" in events[3] + assert events[3]["payload"]["type"] == "session.status" + + +# ============================================================================= +# /event endpoint backward compatibility tests +# ============================================================================= + + +@pytest.mark.anyio +async def test_event_endpoint_no_global_event_fields() -> None: + """wrap_payload=False events have no directory/project/workspace.""" + state = _MockState() + session_evt = SessionStatusEvent.create(session_id="bc1", status_type="busy") + events = await _collect_events(state, wrap_payload=False, events_to_send=[session_evt]) + for evt in events: + assert "directory" not in evt + assert "project" not in evt + assert "workspace" not in evt + + +@pytest.mark.anyio +async def test_event_endpoint_no_payload_wrapper() -> None: + """No payload wrapper key; event data is at top level.""" + state = _MockState() + session_evt = SessionStatusEvent.create(session_id="bc2", status_type="idle") + events = await _collect_events(state, wrap_payload=False, events_to_send=[session_evt]) + session_data = events[1] + assert "payload" not in session_data + # Event fields directly at top level + assert session_data["type"] == "session.status" + + +@pytest.mark.anyio +async def test_event_endpoint_session_id_at_top_level() -> None: + """SessionId present at top level for session events.""" + state = _MockState() + session_evt = SessionStatusEvent.create(session_id="bc3", status_type="busy") + events = await _collect_events(state, wrap_payload=False, events_to_send=[session_evt]) + session_data = events[1] + assert session_data["sessionId"] == "bc3" + + +@pytest.mark.anyio +async def test_event_endpoint_unicode_preserved() -> None: + r"""Unicode characters not escaped as \uXXXX in /event output.""" + state = _MockState() + session_evt = SessionStatusEvent.create(session_id="会话测试", status_type="idle") + gen = _event_generator(state, wrap_payload=False) + # Consume connected event + await gen.__anext__() + # Send unicode session event + queue = state.event_subscribers[-1] + await queue.put(session_evt) + item = await gen.__anext__() + raw_data = item["data"] + assert "会话测试" in raw_data + assert "\\u" not in raw_data + + +# ============================================================================= +# SSE integration tests for /global/event +# ============================================================================= + + +async def _collect_real_events( + state: ServerState, + wrap_payload: bool, + events_to_send: list[Event], +) -> list[dict[str, Any]]: + """Collect SSE items from _event_generator with real ServerState.""" + results: list[dict[str, Any]] = [] + gen = _event_generator(state, wrap_payload=wrap_payload) + # Get the initial connected event + item = await gen.__anext__() + results.append(json.loads(item["data"])) + # Send additional events through the real broadcast system + for event in events_to_send: + await state.broadcast_event(event) + # Yield control so the event can propagate through subscriber queues + await asyncio.sleep(0.01) + item = await gen.__anext__() + results.append(json.loads(item["data"])) + return results + + +@pytest.mark.integration +@pytest.mark.anyio +async def test_global_event_integration_envelope_fields( + server_state: ServerState, +) -> None: + """Test /global/event returns SSE with GlobalEvent envelope.""" + event = SessionStatusEvent.create(session_id="s1", status_type="busy") + results = await _collect_real_events(server_state, wrap_payload=True, events_to_send=[event]) + assert len(results) == 2 + received = results[1] + assert "directory" in received + assert "project" in received + assert "payload" in received + + +@pytest.mark.integration +@pytest.mark.anyio +async def test_global_event_integration_directory_matches_working_dir( + server_state: ServerState, +) -> None: + """Test directory field matches the resolved server working directory.""" + event = SessionStatusEvent.create(session_id="s2", status_type="idle") + results = await _collect_real_events(server_state, wrap_payload=True, events_to_send=[event]) + received = results[1] + assert received["directory"] == str(Path(server_state.working_dir).resolve()) + + +@pytest.mark.integration +@pytest.mark.anyio +async def test_global_event_integration_project_computed( + server_state: ServerState, +) -> None: + """Test project field is computed via compute_project_id.""" + from agentpool_storage.opencode_provider.helpers import compute_project_id + + event = SessionStatusEvent.create(session_id="s3", status_type="busy") + results = await _collect_real_events(server_state, wrap_payload=True, events_to_send=[event]) + received = results[1] + expected_project = compute_project_id(server_state.working_dir) + assert received["project"] == expected_project + + +@pytest.mark.integration +@pytest.mark.anyio +async def test_global_event_integration_workspace_absent( + server_state: ServerState, +) -> None: + """Test workspace field is omitted for single-directory routing.""" + event = SessionStatusEvent.create(session_id="s4", status_type="retry") + results = await _collect_real_events(server_state, wrap_payload=True, events_to_send=[event]) + received = results[1] + assert "workspace" not in received + + +@pytest.mark.integration +@pytest.mark.anyio +async def test_global_event_routing_ignores_agent_execution_cwd( + server_state: ServerState, +) -> None: + """Routing metadata stays anchored to server working_dir, not agent env.cwd.""" + server_state.agent.env.cwd = "/tmp/non-exists-dir" + + event = SessionStatusEvent.create(session_id="s5", status_type="busy") + results = await _collect_real_events(server_state, wrap_payload=True, events_to_send=[event]) + received = results[1] + + assert received["directory"] == str(Path(server_state.working_dir).resolve()) + assert "workspace" not in received + + +@pytest.mark.integration +@pytest.mark.anyio +async def test_global_event_integration_session_id_in_payload( + server_state: ServerState, +) -> None: + """Test sessionId injection at top level of GlobalEvent payload.""" + event = SessionStatusEvent.create(session_id="injected-sid", status_type="busy") + results = await _collect_real_events(server_state, wrap_payload=True, events_to_send=[event]) + received = results[1] + payload = received["payload"] + assert payload["sessionId"] == "injected-sid" + + +@pytest.mark.integration +@pytest.mark.anyio +async def test_global_event_integration_unicode_preserved( + server_state: ServerState, +) -> None: + r"""Test unicode characters preserved in SSE output (not \uXXXX escaped).""" + event = SessionStatusEvent.create(session_id="会话测试", status_type="idle") + results = await _collect_real_events(server_state, wrap_payload=True, events_to_send=[event]) + received = results[1] + payload = received["payload"] + assert payload["sessionId"] == "会话测试" + + +# ============================================================================= +# on_first_subscriber callback tests +# ============================================================================= + + +@pytest.mark.anyio +async def test_on_first_subscriber_fires_once() -> None: + """Callback fires exactly once on first subscriber.""" + state = _MockState() + callback = AsyncMock() + state.on_first_subscriber = callback + + gen = _event_generator(state, wrap_payload=False) + await gen.__anext__() # consume connected event + + assert state._first_subscriber_triggered is True + await asyncio.sleep(0.05) + callback.assert_called_once() + + +@pytest.mark.anyio +async def test_on_first_subscriber_does_not_fire_on_second_subscriber() -> None: + """Callback does not fire again on second subscriber.""" + state = _MockState() + callback = AsyncMock() + state.on_first_subscriber = callback + + gen1 = _event_generator(state, wrap_payload=False) + await gen1.__anext__() # consume connected event + + await asyncio.sleep(0.05) + assert callback.call_count == 1 + + gen2 = _event_generator(state, wrap_payload=False) + await gen2.__anext__() # consume connected event + + await asyncio.sleep(0.05) + # Callback should still have been called only once + callback.assert_called_once() + + +@pytest.mark.anyio +async def test_on_first_subscriber_flag_set_and_stays_true() -> None: + """First subscriber flag is set to True after first subscriber and stays True.""" + state = _MockState() + callback = AsyncMock() + state.on_first_subscriber = callback + + assert state._first_subscriber_triggered is False + + gen1 = _event_generator(state, wrap_payload=False) + await gen1.__anext__() # consume connected event + + assert state._first_subscriber_triggered is True + + gen2 = _event_generator(state, wrap_payload=False) + await gen2.__anext__() # consume connected event + + # Flag must remain True, never reset + assert state._first_subscriber_triggered is True + + +@pytest.mark.anyio +async def test_on_first_subscriber_fires_before_events_delivered() -> None: + """Callback fires before the generator yields any events beyond connected.""" + state = _MockState() + callback = AsyncMock() + state.on_first_subscriber = callback + + gen = _event_generator(state, wrap_payload=False) + # Consuming the connected event should have already triggered the callback + await gen.__anext__() + + # The flag is set synchronously before yielding the connected event + assert state._first_subscriber_triggered is True + await asyncio.sleep(0.05) + # The background task created by the callback should have been scheduled + callback.assert_called_once() + + +@pytest.mark.anyio +async def test_on_first_subscriber_no_callback_set() -> None: + """No callback invocation when on_first_subscriber is None.""" + state = _MockState() + # on_first_subscriber is None by default + assert state.on_first_subscriber is None + + gen = _event_generator(state, wrap_payload=False) + await gen.__anext__() # consume connected event + + # Flag should not be set because there is no callback + assert state._first_subscriber_triggered is False + + +# ============================================================================= +# Client disconnect cleanup tests +# ============================================================================= + + +@pytest.mark.anyio +async def test_disconnect_queue_removed_from_subscribers() -> None: + """Queue is removed from event_subscribers when client disconnects.""" + state = _MockState() + gen = _event_generator(state, wrap_payload=False) + await gen.__anext__() # consume connected event + assert len(state.event_subscribers) == 1 + + await gen.aclose() + assert len(state.event_subscribers) == 0 + + +@pytest.mark.anyio +async def test_disconnect_events_not_delivered() -> None: + """After disconnect, broadcast_event does not deliver to disconnected client.""" + state = _MockState() + gen1 = _event_generator(state, wrap_payload=False) + await gen1.__anext__() # consume connected event + # Add a second subscriber that stays connected to verify isolation + gen2 = _event_generator(state, wrap_payload=False) + await gen2.__anext__() # consume connected event + assert len(state.event_subscribers) == 2 + + queue1 = state.event_subscribers[0] + queue2 = state.event_subscribers[1] + + # Disconnect first client + await gen1.aclose() + assert len(state.event_subscribers) == 1 + assert queue1 not in state.event_subscribers + assert queue2 in state.event_subscribers + + # Put an event directly — only queue2 should receive it + event = SessionStatusEvent.create(session_id="disc1", status_type="busy") + await queue2.put(event) + item = await gen2.__anext__() + data = json.loads(item["data"]) + assert data["type"] == "session.status" + + +@pytest.mark.anyio +async def test_disconnect_finally_block_executes() -> None: + """The finally block in _event_generator runs on disconnect, removing the queue.""" + state = _MockState() + gen = _event_generator(state, wrap_payload=False) + await gen.__anext__() # consume connected event + queue_before = state.event_subscribers[-1] + assert queue_before in state.event_subscribers + + await gen.aclose() + + # The finally block removed the queue + assert queue_before not in state.event_subscribers + assert len(state.event_subscribers) == 0 + + +@pytest.mark.anyio +async def test_disconnect_abrupt_cleanup() -> None: + """Abrupt disconnect (task cancellation) still triggers cleanup.""" + state = _MockState() + + async def consume() -> None: + gen = _event_generator(state, wrap_payload=False) + with contextlib.suppress(StopAsyncIteration): + async for _ in gen: + pass + + task = asyncio.create_task(consume()) + # Let the generator start and consume the connected event + await asyncio.sleep(0.05) + assert len(state.event_subscribers) == 1 + + task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await task + + # Cleanup should have removed the subscriber queue + assert len(state.event_subscribers) == 0 + + +@pytest.mark.anyio +async def test_disconnect_no_memory_leak() -> None: + """Multiple connect/disconnect cycles do not cause subscriber list growth.""" + state = _MockState() + + for _ in range(5): + gen = _event_generator(state, wrap_payload=False) + await gen.__anext__() # consume connected event + assert len(state.event_subscribers) == 1 + await gen.aclose() + assert len(state.event_subscribers) == 0 + + # After 5 cycles, no leaked subscribers + assert len(state.event_subscribers) == 0 + + +# ============================================================================= +# _extract_session_id exhaustiveness tests +# ============================================================================= + + +def _make_session(session_id: str = "test-sid") -> Session: + """Create a minimal Session for event construction.""" + return Session( + id=session_id, + project_id="proj1", + directory="/tmp", + title="Test", + time=TimeCreatedUpdated(created=0, updated=0), + ) + + +def _make_part(session_id: str = "test-sid") -> Part: + """Create a minimal Part for event construction.""" + return TextPart( + id="part1", + message_id="msg1", + session_id=session_id, + text="hello", + ) + + +# All 22 handled event types with their constructors +_HANDLED_EVENT_FACTORIES: list[tuple[str, type]] = [ + ("session.deleted", SessionDeletedEvent), + ("session.status", SessionStatusEvent), + ("session.idle", SessionIdleEvent), + ("session.compacted", SessionCompactedEvent), + ("message.removed", MessageRemovedEvent), + ("message.part.removed", PartRemovedEvent), + ("permission.asked", PermissionRequestEvent), + ("permission.replied", PermissionResolvedEvent), + ("question.asked", QuestionAskedEvent), + ("question.replied", QuestionRepliedEvent), + ("question.rejected", QuestionRejectedEvent), + ("todo.updated", TodoUpdatedEvent), + ("session.error", SessionErrorEvent), + ("session.created", SessionCreatedEvent), + ("session.updated", SessionUpdatedEvent), + ("message.part.updated", PartUpdatedEvent), + ("session.diff", SessionDiffEvent), + ("message.part.delta", PartDeltaEvent), + ("permission.updated", PermissionUpdatedEvent), + ("command.executed", CommandExecutedEvent), + ("tui.session.select", TuiSessionSelectEvent), + ("message.updated", MessageUpdatedEvent), +] + + +def _build_handled_event(event_type: type) -> Event: # noqa: PLR0911 + """Build a handled event with session_id='abc' using the appropriate constructor.""" + sid = "abc" + if event_type is SessionDeletedEvent: + return SessionDeletedEvent.create(session_id=sid) + if event_type is SessionStatusEvent: + return SessionStatusEvent.create(session_id=sid, status_type="busy") + if event_type is SessionIdleEvent: + return SessionIdleEvent.create(session_id=sid) + if event_type is SessionCompactedEvent: + return SessionCompactedEvent.create(session_id=sid) + if event_type is MessageRemovedEvent: + return MessageRemovedEvent.create(session_id=sid, message_id="m1") + if event_type is PartRemovedEvent: + return PartRemovedEvent.create(session_id=sid, message_id="m1", part_id="p1") + if event_type is PermissionRequestEvent: + return PermissionRequestEvent.create( + session_id=sid, + permission_id="perm1", + tool_name="bash", + args_preview="ls", + message="Allow?", + ) + if event_type is PermissionResolvedEvent: + return PermissionResolvedEvent.create( + session_id=sid, + request_id="perm1", + reply="once", + ) + if event_type is QuestionAskedEvent: + return QuestionAskedEvent.create( + request_id="q1", + session_id=sid, + questions=[ + QuestionInfo( + question="Continue?", + header="Confirm", + options=[QuestionOption(label="Yes", description="Proceed")], + ) + ], + ) + if event_type is QuestionRepliedEvent: + return QuestionRepliedEvent.create( + session_id=sid, + request_id="q1", + answers=[["Yes"]], + ) + if event_type is QuestionRejectedEvent: + return QuestionRejectedEvent.create( + session_id=sid, + request_id="q1", + ) + if event_type is TodoUpdatedEvent: + return TodoUpdatedEvent.create(session_id=sid, todos=[]) + if event_type is SessionErrorEvent: + return SessionErrorEvent.create(session_id=sid, error_name="TestError") + if event_type is SessionCreatedEvent: + return SessionCreatedEvent.create(session=_make_session(sid)) + if event_type is SessionUpdatedEvent: + return SessionUpdatedEvent.create(session=_make_session(sid)) + if event_type is PartUpdatedEvent: + return PartUpdatedEvent.create(part=_make_part(sid)) + if event_type is SessionDiffEvent: + return SessionDiffEvent.create(session_id=sid, diff=[]) + if event_type is PartDeltaEvent: + return PartDeltaEvent.create(session_id=sid, message_id="m1", part_id="p1", delta="hi") + if event_type is PermissionUpdatedEvent: + return PermissionUpdatedEvent.create( + session_id=sid, + permission_id="perm1", + tool_name="bash", + patterns=["bash: *"], + metadata={}, + ) + if event_type is CommandExecutedEvent: + return CommandExecutedEvent.create( + name="test", + session_id=sid, + arguments="", + message_id="m1", + ) + if event_type is TuiSessionSelectEvent: + return TuiSessionSelectEvent.create(session_id=sid) + if event_type is MessageUpdatedEvent: + return MessageUpdatedEvent.create( + message=UserMessage( + id="m1", + session_id=sid, + time=TimeCreated(created=0), + ), + ) + msg = f"Unhandled event type in test helper: {event_type}" + raise ValueError(msg) + + +@pytest.mark.parametrize( + ("event_type_name", "event_type"), + [(name, cls) for name, cls in _HANDLED_EVENT_FACTORIES], + ids=[name for name, _ in _HANDLED_EVENT_FACTORIES], +) +def test_extract_session_id_handled_events( + event_type_name: str, + event_type: type, +) -> None: + """All 22 handled event types extract sessionId correctly.""" + event = _build_handled_event(event_type) + result = _extract_session_id(event) + assert result == "abc", f"Expected 'abc' for {event_type_name}, got {result!r}" + + +def test_extract_session_id_session_error_nullable() -> None: + """SessionErrorEvent with None session_id returns None.""" + event = SessionErrorEvent.create(session_id=None, error_name="TestError") + result = _extract_session_id(event) + assert result is None + + +def test_extract_session_id_no_session_events_return_none() -> None: + """Explicitly-listed no-session event types return None without warnings.""" + no_session_events: list[Event] = [ + ServerConnectedEvent(), + ServerHeartbeatEvent(), + FileEditedEvent.create(file="/tmp/test.py"), + FileWatcherUpdatedEvent.create(file="/tmp/test.py", event="change"), + McpToolsChangedEvent.create(server="test_server"), + PtyCreatedEvent.create( + info=PtyInfo( + id="p1", + title="test", + command="echo", + args=[], + cwd="/tmp", + status="running", + pid=1234, + ), + ), + PtyUpdatedEvent.create( + info=PtyInfo( + id="p1", + title="test", + command="echo", + args=[], + cwd="/tmp", + status="running", + pid=1234, + ), + ), + PtyExitedEvent.create(pty_id="p1", exit_code=0), + PtyDeletedEvent.create(pty_id="p1"), + LspUpdatedEvent(), + LspClientDiagnosticsEvent.create(server_id="s1", path="/tmp"), + ProjectUpdatedEvent.create( + project=Project(id="test", worktree="/tmp", time=ProjectTime(created=0)), + ), + VcsBranchUpdatedEvent.create(branch="main"), + TuiPromptAppendEvent.create(text="hello"), + TuiCommandExecuteEvent.create(command="test"), + TuiToastShowEvent.create(message="hi"), + ] + for event in no_session_events: + result = _extract_session_id(event) + assert result is None, f"Expected None for {type(event).__name__}, got {result!r}" + + +def test_extract_session_id_no_warning_for_no_session_events( + caplog: pytest.LogCaptureFixture, +) -> None: + """No warning logged for explicitly-listed no-session event types.""" + event = FileEditedEvent.create(file="/tmp/test.py") + with caplog.at_level("WARNING"): + result = _extract_session_id(event) + assert result is None + assert "Unhandled event type in _extract_session_id" not in caplog.text + + +def test_extract_session_id_warning_for_unknown_event_type( + caplog: pytest.LogCaptureFixture, +) -> None: + """Warning is logged when an unknown event type hits the wildcard case. + + Simulates a future event type being added to the Event union but not + yet covered in _extract_session_id. The exhaustiveness test catches + this at the type level; this test verifies the runtime warning. + """ + # Use a plain MagicMock (no spec) to simulate an unrecognized event. + # A spec-less mock won't match any of the pattern-matching cases + # and will fall through to the `case _:` wildcard. + mock_event = MagicMock() + mock_event.__class__.__name__ = "FutureUnknownEvent" + with caplog.at_level("WARNING"): + result = _extract_session_id(mock_event) # type: ignore[arg-type] + assert result is None + assert "Unhandled event type in _extract_session_id" in caplog.text + assert "FutureUnknownEvent" in caplog.text + + +def test_extract_session_id_no_warning_for_handled(caplog: pytest.LogCaptureFixture) -> None: + """No warning logged for handled event types.""" + event = SessionStatusEvent.create(session_id="no-warn", status_type="idle") + with caplog.at_level("WARNING"): + _extract_session_id(event) + assert "Unhandled event type in _extract_session_id" not in caplog.text + + +def test_extract_session_id_exhaustiveness() -> None: + """All Event union members are either handled or explicitly documented as no-session. + + Catches future regressions: if a new event type with session_id is added + to the Event union but not to _extract_session_id, this test fails. + """ + # Event types handled by _extract_session_id match cases + handled_types: set[type] = { + SessionDeletedEvent, + SessionStatusEvent, + SessionIdleEvent, + SessionCompactedEvent, + MessageRemovedEvent, + PartRemovedEvent, + PermissionRequestEvent, + PermissionResolvedEvent, + QuestionAskedEvent, + QuestionRepliedEvent, + QuestionRejectedEvent, + TodoUpdatedEvent, + SessionErrorEvent, + SessionCreatedEvent, + SessionUpdatedEvent, + PartUpdatedEvent, + SessionDiffEvent, + PartDeltaEvent, + PermissionUpdatedEvent, + CommandExecutedEvent, + TuiSessionSelectEvent, + MessageUpdatedEvent, + } + + # Event types that genuinely have no session association + # (no session_id field anywhere in their properties) + no_session_types: set[type] = { + ServerConnectedEvent, + ServerHeartbeatEvent, + FileWatcherUpdatedEvent, + FileEditedEvent, + McpToolsChangedEvent, + VcsBranchUpdatedEvent, + TuiPromptAppendEvent, + TuiCommandExecuteEvent, + TuiToastShowEvent, + ProjectUpdatedEvent, + LspUpdatedEvent, + LspClientDiagnosticsEvent, + PtyCreatedEvent, + PtyUpdatedEvent, + PtyExitedEvent, + PtyDeletedEvent, + } + + # Event types with session_id that are NOT handled (known gaps) + known_gap_types: set[type] = set() + + expected = handled_types | no_session_types | known_gap_types + + # Get all members of the Event union + event_union_args: set[type] = set(Event.__args__) + + # Every union member must be accounted for + missing = event_union_args - expected + assert not missing, ( + f"New event types not covered by _extract_session_id: " + f"{sorted(t.__name__ for t in missing)}. " + f"Add them to handled_types, no_session_types, or known_gap_types." + ) + + # No extra types that aren't in the union + extra = expected - event_union_args + assert not extra, ( + f"Types listed in test but not in Event union: {sorted(t.__name__ for t in extra)}" + ) + + # Known gaps should be documented — if they're fixed, move them to handled + if known_gap_types: + gap_names = sorted(t.__name__ for t in known_gap_types) + # This assertion always passes but documents the known gaps + assert True, f"Known gap types with session_id not handled: {gap_names}" + + +# ============================================================================= +# Concurrent subscriber tests +# ============================================================================= + + +@pytest.mark.anyio +async def test_concurrent_two_subscribers_both_receive_events() -> None: + """Two SSE clients both receive a broadcast event.""" + state = _MockState() + + gen1 = _event_generator(state, wrap_payload=True) + gen2 = _event_generator(state, wrap_payload=True) + + # Consume initial connected events + await gen1.__anext__() + await gen2.__anext__() + + assert len(state.event_subscribers) == 2 + + # Broadcast event to both subscribers via their queues + event = SessionStatusEvent.create(session_id="s_concurrent", status_type="busy") + queue1 = state.event_subscribers[0] + queue2 = state.event_subscribers[1] + await queue1.put(event) + await queue2.put(event) + + item1 = await gen1.__anext__() + item2 = await gen2.__anext__() + + data1 = json.loads(item1["data"]) + data2 = json.loads(item2["data"]) + + assert data1["payload"]["type"] == "session.status" + assert data2["payload"]["type"] == "session.status" + assert data1["payload"]["sessionId"] == "s_concurrent" + assert data2["payload"]["sessionId"] == "s_concurrent" + + +@pytest.mark.anyio +async def test_concurrent_subscribers_receive_same_content() -> None: + """Both subscribers get identical GlobalEvent envelopes.""" + state = _MockState() + + gen1 = _event_generator(state, wrap_payload=True) + gen2 = _event_generator(state, wrap_payload=True) + + await gen1.__anext__() + await gen2.__anext__() + + event = SessionStatusEvent.create(session_id="s_same", status_type="idle") + queue1 = state.event_subscribers[0] + queue2 = state.event_subscribers[1] + await queue1.put(event) + await queue2.put(event) + + item1 = await gen1.__anext__() + item2 = await gen2.__anext__() + + # Both envelopes must have identical directory, project, and payload + data1 = json.loads(item1["data"]) + data2 = json.loads(item2["data"]) + + assert data1["directory"] == data2["directory"] + assert data1["project"] == data2["project"] + assert data1["payload"] == data2["payload"] + + +@pytest.mark.anyio +async def test_concurrent_event_ordering_preserved() -> None: + """Broadcast 3 events; each subscriber receives them in order.""" + state = _MockState() + + gen1 = _event_generator(state, wrap_payload=True) + gen2 = _event_generator(state, wrap_payload=True) + + await gen1.__anext__() + await gen2.__anext__() + + queue1 = state.event_subscribers[0] + queue2 = state.event_subscribers[1] + + events = [ + SessionStatusEvent.create(session_id="ord1", status_type="busy"), + SessionStatusEvent.create(session_id="ord2", status_type="idle"), + SessionStatusEvent.create(session_id="ord3", status_type="retry"), + ] + + for ev in events: + await queue1.put(ev) + await queue2.put(ev) + + # Collect all 3 events from each subscriber + received1 = [json.loads((await gen1.__anext__())["data"]) for _ in range(3)] + received2 = [json.loads((await gen2.__anext__())["data"]) for _ in range(3)] + + expected_order = ["ord1", "ord2", "ord3"] + ids1 = [r["payload"]["sessionId"] for r in received1] + ids2 = [r["payload"]["sessionId"] for r in received2] + + assert ids1 == expected_order + assert ids2 == expected_order + + +@pytest.mark.anyio +async def test_concurrent_subscriber_receives_after_another_disconnects() -> None: + """Subscriber B still receives events after subscriber A disconnects.""" + state = _MockState() + + gen_a = _event_generator(state, wrap_payload=True) + gen_b = _event_generator(state, wrap_payload=True) + + await gen_a.__anext__() + await gen_b.__anext__() + + assert len(state.event_subscribers) == 2 + + # Disconnect subscriber A + await gen_a.aclose() + assert len(state.event_subscribers) == 1 + + # Send event only to remaining subscriber B's queue + event = SessionStatusEvent.create(session_id="s_survive", status_type="busy") + queue_b = state.event_subscribers[0] + await queue_b.put(event) + + item_b = await gen_b.__anext__() + data_b = json.loads(item_b["data"]) + + assert data_b["payload"]["type"] == "session.status" + assert data_b["payload"]["sessionId"] == "s_survive" + + +@pytest.mark.anyio +async def test_concurrent_all_get_server_connected() -> None: + """Each subscriber gets the initial payload-wrapped server.connected event.""" + state = _MockState() + + gen1 = _event_generator(state, wrap_payload=True) + gen2 = _event_generator(state, wrap_payload=True) + gen3 = _event_generator(state, wrap_payload=True) + + item1 = await gen1.__anext__() + item2 = await gen2.__anext__() + item3 = await gen3.__anext__() + + for item in [item1, item2, item3]: + data = json.loads(item["data"]) + assert data["payload"]["type"] == "server.connected" + assert "directory" not in data + assert "project" not in data + assert "payload" in data + + +# ============================================================================= +# ServerState.broadcast_event direct tests +# ============================================================================= + + +def _make_broadcast_state() -> ServerState: + """Create a ServerState with a minimal mock agent for broadcast_event tests.""" + from unittest.mock import Mock + + mock_env = Mock() + mock_env.get_fs = Mock(return_value=Mock()) + mock_agent = Mock() + mock_agent.env = mock_env + return ServerState(working_dir="/test", agent=mock_agent) + + +@pytest.mark.anyio +async def test_broadcast_event_single_subscriber() -> None: + """Broadcast delivers event to one subscriber queue.""" + state = _make_broadcast_state() + queue: asyncio.Queue[Event] = asyncio.Queue() + state.event_subscribers.append(queue) + + event = SessionStatusEvent.create(session_id="abc", status_type="busy") + await state.broadcast_event(event) + + received = queue.get_nowait() + assert received is event + + +@pytest.mark.anyio +async def test_broadcast_event_multiple_subscribers() -> None: + """Broadcast delivers event to all subscriber queues.""" + state = _make_broadcast_state() + queue1: asyncio.Queue[Event] = asyncio.Queue() + queue2: asyncio.Queue[Event] = asyncio.Queue() + queue3: asyncio.Queue[Event] = asyncio.Queue() + state.event_subscribers.extend([queue1, queue2, queue3]) + + event = SessionStatusEvent.create(session_id="abc", status_type="busy") + await state.broadcast_event(event) + + assert queue1.get_nowait() is event + assert queue2.get_nowait() is event + assert queue3.get_nowait() is event + + +@pytest.mark.anyio +async def test_broadcast_event_no_subscribers() -> None: + """Broadcast with no subscribers does not raise.""" + state = _make_broadcast_state() + assert state.event_subscribers == [] + + event = SessionStatusEvent.create(session_id="abc", status_type="busy") + await state.broadcast_event(event) # Should not raise + + +@pytest.mark.anyio +async def test_broadcast_event_exception_isolation() -> None: + """Subscriber whose queue raises is removed; other subscribers still receive.""" + state = _make_broadcast_state() + + good_queue: asyncio.Queue[Event] = asyncio.Queue() + state.event_subscribers.append(good_queue) + + # Create a mock queue that raises on put_nowait + bad_queue = MagicMock(spec=asyncio.Queue) + bad_queue.put_nowait.side_effect = RuntimeError("queue broken") + state.event_subscribers.append(bad_queue) + + event = SessionStatusEvent.create(session_id="abc", status_type="busy") + await state.broadcast_event(event) + + # Good queue should still have received the event + assert good_queue.get_nowait() is event + # Bad queue should have been removed from subscribers + assert bad_queue not in state.event_subscribers + assert good_queue in state.event_subscribers + + +@pytest.mark.anyio +async def test_broadcast_event_queue_full_dropped() -> None: + """Full queue has event dropped; other subscribers still receive.""" + state = _make_broadcast_state() + + # Queue with maxsize=1, already full + full_queue: asyncio.Queue[Event] = asyncio.Queue(maxsize=1) + full_queue.put_nowait(ServerHeartbeatEvent()) # Fill the queue + state.event_subscribers.append(full_queue) + + good_queue: asyncio.Queue[Event] = asyncio.Queue() + state.event_subscribers.append(good_queue) + + event = SessionStatusEvent.create(session_id="abc", status_type="busy") + await state.broadcast_event(event) + + # Full queue should still have only the original item (event dropped) + assert full_queue.qsize() == 1 + assert not isinstance(full_queue.get_nowait(), SessionStatusEvent) + + # Good queue should have received the event + assert good_queue.get_nowait() is event + + +# ============================================================================= +# /global/health endpoint tests +# ============================================================================= + + +@pytest.mark.anyio +async def test_global_health_endpoint(async_client: AsyncClient) -> None: + """GET /global/health returns 200 with HealthResponse body.""" + response = await async_client.get("/global/health") + assert response.status_code == 200 + data = response.json() + assert data["healthy"] is True + assert "version" in data + + +@pytest.mark.anyio +async def test_global_health_endpoint_fields(async_client: AsyncClient) -> None: + """GET /global/health returns correct healthy and version fields.""" + from agentpool_server.opencode_server.routes.global_routes import VERSION + + response = await async_client.get("/global/health") + data = response.json() + assert data["healthy"] is True + assert data["version"] == VERSION + + +# ============================================================================= +# GlobalEvent edge case tests +# ============================================================================= + + +def test_global_event_large_payload() -> None: + r"""Large payload (100KB+) serializes and deserializes correctly. + + Uses a TodoUpdatedEvent with a very long todo content string to produce + a payload exceeding 100KB. Verifies round-trip correctness via json.loads. + """ + large_text = "A" * 100_000 # 100KB string + todo = Todo(id="t1", content=large_text, status="pending", priority="high") + event = TodoUpdatedEvent.create(session_id="large-sid", todos=[todo]) + factory = GlobalEventFactory(directory="/tmp", project="abc") + result = factory.wrap(event) + data = json.loads(result) + # Payload contains the full large text + assert data["payload"]["properties"]["todos"][0]["content"] == large_text + # Round-trip: re-serialize and re-parse + round_tripped = json.loads(json.dumps(data, ensure_ascii=False)) + assert round_tripped["payload"]["properties"]["todos"][0]["content"] == large_text + + +def test_global_event_special_characters() -> None: + r"""Special characters (quotes, backslashes, control chars, emojis) preserved. + + Verifies that characters like '"', '\\', '\n', '\t', and emoji are correctly + serialized and deserialized through _serialize_event and GlobalEventFactory. + """ + special_sid = 'sid-with-"quotes"-and-\\backslash\\' + event = SessionStatusEvent.create(session_id=special_sid, status_type="busy") + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + # sessionId at top level + assert data["sessionId"] == special_sid + # sessionID inside properties (alias-converted) + assert data["properties"]["sessionID"] == special_sid + + # Also test via factory.wrap with emoji and control chars in Todo content + emoji_text = "Hello 🔥🚀 world!\n\ttabbed line" + todo = Todo(id="t2", content=emoji_text, status="in_progress", priority="medium") + event2 = TodoUpdatedEvent.create(session_id="emoji-sid", todos=[todo]) + factory = GlobalEventFactory(directory="/tmp", project="abc") + result2 = factory.wrap(event2) + data2 = json.loads(result2) + assert data2["payload"]["properties"]["todos"][0]["content"] == emoji_text + + +def test_global_event_workspace_none_omitted() -> None: + """Workspace=None is excluded from GlobalEvent via exclude_none=True.""" + event = GlobalEvent(directory="/tmp/test", project="abc123", payload={"type": "test"}) + dumped = event.model_dump(by_alias=True, exclude_none=True) + assert "workspace" not in dumped + assert dumped["directory"] == "/tmp/test" + assert dumped["project"] == "abc123" + + # Also verify JSON serialization round-trip + json_str = json.dumps(dumped, ensure_ascii=False) + parsed = json.loads(json_str) + assert "workspace" not in parsed + + +def test_global_event_empty_string_fields() -> None: + r"""Empty string sessionId is preserved (not treated as None). + + An empty string is a valid value and should not be excluded by + exclude_none=True (which only drops None, not empty strings). + """ + event = SessionStatusEvent.create(session_id="", status_type="idle") + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + # sessionId injected at top level + assert data["sessionId"] == "" + # sessionID inside properties + assert data["properties"]["sessionID"] == "" + + # Via factory.wrap + factory = GlobalEventFactory(directory="/tmp", project="abc") + wrapped = factory.wrap(event) + wrapped_data = json.loads(wrapped) + assert wrapped_data["payload"]["sessionId"] == "" + + +def test_global_event_unicode_multibyte() -> None: + r"""CJK characters and multi-byte emoji sequences are preserved. + + Verifies that characters like 日本語 (Japanese) and complex emoji + sequences like 👨‍👩‍👧‍👦 (family emoji, ZWJ sequence) survive + serialization round-trip without \uXXXX escaping. + """ + cjk_sid = "会話-日本語-テスト" + event = SessionStatusEvent.create(session_id=cjk_sid, status_type="busy") + result = _serialize_event(event, wrap_payload=False) + # Raw string must contain CJK chars, not \uXXXX escapes + assert "日本語" in result + assert "テスト" in result + assert "\\u" not in result + + # Round-trip via json.loads + data = json.loads(result) + assert data["sessionId"] == cjk_sid + + # ZWJ emoji in todo content + family_emoji = "👨‍👩‍👧‍👦" + todo = Todo(id="t3", content=f"Family: {family_emoji}", status="completed", priority="low") + event2 = TodoUpdatedEvent.create(session_id=cjk_sid, todos=[todo]) + factory = GlobalEventFactory(directory="/tmp", project="abc") + wrapped = factory.wrap(event2) + assert family_emoji in wrapped + assert "\\u" not in wrapped + wrapped_data = json.loads(wrapped) + assert wrapped_data["payload"]["properties"]["todos"][0]["content"] == f"Family: {family_emoji}" diff --git a/tests/servers/opencode_server/test_message_models.py b/tests/servers/opencode_server/test_message_models.py index b83ea7cb0..342c56a3e 100644 --- a/tests/servers/opencode_server/test_message_models.py +++ b/tests/servers/opencode_server/test_message_models.py @@ -1,15 +1,25 @@ """Tests for OpenCode message models. -Tests the MessageWithParts class and related models. +Tests the MessageWithParts class, UserMessage model nesting, and +backward-compatible variant deserialization. """ from __future__ import annotations +import json + from agentpool_server.opencode_server.models import ( MessagePath, + MessageRequest, MessageTime, MessageWithParts, + ModelRef, TimeCreated, + UserMessage, +) +from agentpool_server.opencode_server.models.message import ( + AssistantMessage, + TextPartInput, ) @@ -157,3 +167,221 @@ def test_add_multiple_parts(self): msg.add_text_part("Processing...") assert len(msg.parts) == 2 + + +class TestModelRefVariant: + """Tests for ModelRef with optional variant field.""" + + def test_model_ref_without_variant(self): + """ModelRef should work without variant (backward compat).""" + ref = ModelRef(provider_id="openai", model_id="gpt-4o") + assert ref.provider_id == "openai" + assert ref.model_id == "gpt-4o" + assert ref.variant is None + + def test_model_ref_with_variant(self): + """ModelRef should accept variant field.""" + ref = ModelRef(provider_id="openai", model_id="gpt-4o", variant="high") + assert ref.variant == "high" + + def test_model_ref_with_only_variant(self): + """ModelRef should accept variant-only (no provider/model).""" + ref = ModelRef(variant="medium") + assert ref.variant == "medium" + assert ref.provider_id is None + assert ref.model_id is None + + def test_model_ref_serialization_with_variant(self): + """ModelRef should serialize variant as camelCase under by_alias.""" + ref = ModelRef(provider_id="openai", model_id="gpt-4o", variant="medium") + data = ref.model_dump(by_alias=True, exclude_none=True) + assert data == {"providerID": "openai", "modelID": "gpt-4o", "variant": "medium"} + + def test_model_ref_serialization_without_variant(self): + """ModelRef without variant should not include variant in output.""" + ref = ModelRef(provider_id="openai", model_id="gpt-4o") + data = ref.model_dump(by_alias=True, exclude_none=True) + assert "variant" not in data + assert data == {"providerID": "openai", "modelID": "gpt-4o"} + + def test_model_ref_serialization_variant_only(self): + """ModelRef with only variant should only include variant in output.""" + ref = ModelRef(variant="low") + data = ref.model_dump(by_alias=True, exclude_none=True) + assert data == {"variant": "low"} + + +class TestUserMessageVariantNesting: + """Tests for UserMessage with variant nested under model.""" + + def test_user_message_with_model_variant(self): + """UserMessage should accept variant via model object.""" + msg = UserMessage( + id="msg-1", + session_id="session-1", + time=TimeCreated(created=1234567890), + model=ModelRef(provider_id="openai", model_id="gpt-4o", variant="high"), + ) + assert msg.model is not None + assert msg.model.variant == "high" + + def test_user_message_without_model(self): + """UserMessage should work without model at all.""" + msg = UserMessage( + id="msg-1", + session_id="session-1", + time=TimeCreated(created=1234567890), + ) + assert msg.model is None + + def test_user_message_no_top_level_variant_in_output(self): + """Serialized UserMessage should NOT have top-level variant.""" + msg = UserMessage( + id="msg-1", + session_id="session-1", + time=TimeCreated(created=1234567890), + model=ModelRef(provider_id="openai", model_id="gpt-4o", variant="high"), + ) + data = msg.model_dump(by_alias=True, exclude_none=True) + assert "variant" not in data + assert data["model"] == {"providerID": "openai", "modelID": "gpt-4o", "variant": "high"} + + def test_backward_compat_top_level_variant_no_model(self): + """Old JSON with top-level variant and no model should deserialize. + + The variant should be migrated into a model object with just variant. + """ + msg = UserMessage.model_validate({ + "id": "msg-1", + "sessionID": "session-1", + "time": {"created": 1234567890}, + "variant": "medium", + }) + assert msg.model is not None + assert msg.model.variant == "medium" + assert msg.model.provider_id is None + assert msg.model.model_id is None + + def test_backward_compat_top_level_variant_with_existing_model(self): + """Old JSON with top-level variant AND model should merge variant into model.""" + msg = UserMessage.model_validate({ + "id": "msg-1", + "sessionID": "session-1", + "time": {"created": 1234567890}, + "model": {"providerID": "openai", "modelID": "gpt-4o"}, + "variant": "low", + }) + assert msg.model is not None + assert msg.model.provider_id == "openai" + assert msg.model.model_id == "gpt-4o" + assert msg.model.variant == "low" + + def test_new_format_variant_in_model_no_migration(self): + """New JSON with variant inside model should work without migration.""" + msg = UserMessage.model_validate({ + "id": "msg-1", + "sessionID": "session-1", + "time": {"created": 1234567890}, + "model": {"providerID": "openai", "modelID": "gpt-4o", "variant": "max"}, + }) + assert msg.model is not None + assert msg.model.variant == "max" + assert msg.model.provider_id == "openai" + assert msg.model.model_id == "gpt-4o" + + def test_no_variant_no_model(self): + """JSON without variant or model should work.""" + msg = UserMessage.model_validate({ + "id": "msg-1", + "sessionID": "session-1", + "time": {"created": 1234567890}, + }) + assert msg.model is None + + def test_variant_not_in_serialized_output(self): + """After migration, top-level variant should NOT appear in serialized output.""" + msg = UserMessage.model_validate({ + "id": "msg-1", + "sessionID": "session-1", + "time": {"created": 1234567890}, + "variant": "medium", + }) + data = msg.model_dump(by_alias=True, exclude_none=True) + assert "variant" not in data + assert data["model"]["variant"] == "medium" + + +class TestMessageRequestVariantNesting: + """Tests for MessageRequest with variant nested under model.""" + + def test_message_request_with_model_variant(self): + """MessageRequest should accept variant via model object.""" + req = MessageRequest( + parts=[TextPartInput(text="hello")], + model=ModelRef(provider_id="openai", model_id="gpt-4o", variant="high"), + ) + assert req.model is not None + assert req.model.variant == "high" + + def test_backward_compat_top_level_variant_no_model(self): + """Old JSON with top-level variant should migrate into model.""" + req = MessageRequest.model_validate({ + "parts": [{"type": "text", "text": "hello"}], + "variant": "medium", + }) + assert req.model is not None + assert req.model.variant == "medium" + + def test_backward_compat_top_level_variant_with_model(self): + """Old JSON with both top-level variant and model should merge.""" + req = MessageRequest.model_validate({ + "parts": [{"type": "text", "text": "hello"}], + "model": {"providerID": "openai", "modelID": "gpt-4o"}, + "variant": "low", + }) + assert req.model is not None + assert req.model.variant == "low" + assert req.model.provider_id == "openai" + + def test_no_variant_in_serialized_output(self): + """MessageRequest should not have top-level variant in output.""" + req = MessageRequest( + parts=[TextPartInput(text="hello")], + model=ModelRef(provider_id="openai", model_id="gpt-4o", variant="high"), + ) + data = req.model_dump(by_alias=True, exclude_none=True) + assert "variant" not in data + assert data["model"]["variant"] == "high" + + +class TestAssistantMessageVariantMigration: + """Tests for AssistantMessage variant removal.""" + + def test_assistant_message_ignores_top_level_variant(self): + """Old JSON with top-level variant should silently drop it.""" + msg = AssistantMessage.model_validate({ + "id": "msg-1", + "sessionID": "session-1", + "parentID": "parent-1", + "modelID": "gpt-4o", + "providerID": "openai", + "path": {"cwd": "/test", "root": "/test"}, + "time": {"created": 1234567890}, + "variant": "high", + }) + # AssistantMessage no longer has a variant field + assert not hasattr(msg, "variant") + + def test_assistant_message_serialization_no_variant(self): + """AssistantMessage should not include variant in output.""" + msg = AssistantMessage( + id="msg-1", + session_id="session-1", + parent_id="parent-1", + model_id="gpt-4o", + provider_id="openai", + path=MessagePath(cwd="/test", root="/test"), + time=MessageTime(created=1234567890), + ) + data = msg.model_dump(by_alias=True, exclude_none=True) + assert "variant" not in data diff --git a/tests/servers/opencode_server/test_message_timeout.py b/tests/servers/opencode_server/test_message_timeout.py new file mode 100644 index 000000000..dbbe8d8e3 --- /dev/null +++ b/tests/servers/opencode_server/test_message_timeout.py @@ -0,0 +1,80 @@ +"""Regression tests for long-running sync message handling.""" + +from __future__ import annotations + +import asyncio + +from pydantic_ai import RequestUsage +import pytest + +from agentpool_server.opencode_server.models import MessageRequest, TextPartInput +from agentpool_server.opencode_server.routes import message_routes + + +class _DelayedAdapter: + """Test adapter that blocks until the test releases it.""" + + gate: asyncio.Event + started: asyncio.Event + + def __init__(self, **_: object) -> None: + self.response_text = "Delayed reply" + self.usage = RequestUsage(input_tokens=0, output_tokens=0) + self.cost_info = None + + async def process_stream(self, stream): + self.started.set() + await self.gate.wait() + if False: + yield stream + + def finalize(self): + return iter(()) + + +@pytest.mark.asyncio +async def test_sync_message_does_not_use_route_timeout( + async_client, + server_state, + event_capture, + monkeypatch, +) -> None: + """Long-silent sync turns should stay alive until the server-side work finishes.""" + response = await async_client.post("/session", json={"title": "Delayed Reply"}) + session_id = response.json()["id"] + + gate = asyncio.Event() + started = asyncio.Event() + _DelayedAdapter.gate = gate + _DelayedAdapter.started = started + + async def silent_stream(): + await gate.wait() + if False: + yield None + + def fail_if_timeout_used(*args: object, **kwargs: object): + msg = "sync /message must not wrap agent streams in a route-owned timeout" + raise AssertionError(msg) + + monkeypatch.setattr(message_routes.asyncio, "timeout", fail_if_timeout_used) + monkeypatch.setattr(message_routes, "OpenCodeStreamAdapter", _DelayedAdapter) + server_state.agent.run_stream = lambda *args, **kwargs: silent_stream() + + request = MessageRequest(parts=[TextPartInput(text="hello")], agent="default") + request_task = asyncio.create_task( + async_client.post(f"/session/{session_id}/message", json=request.model_dump(mode="json")) + ) + + await asyncio.wait_for(started.wait(), timeout=1.0) + await asyncio.sleep(0.05) + + assert not request_task.done() + assert server_state.session_status[session_id].type == "busy" + + gate.set() + result = await asyncio.wait_for(request_task, timeout=1.0) + + assert result.status_code == 200 + assert server_state.session_status[session_id].type == "idle" + assert event_capture.get_events_by_type("session.error") == [] diff --git a/tests/servers/opencode_server/test_otlp_compat.py b/tests/servers/opencode_server/test_otlp_compat.py new file mode 100644 index 000000000..dc885ad4d --- /dev/null +++ b/tests/servers/opencode_server/test_otlp_compat.py @@ -0,0 +1,74 @@ +"""Tests for OTLP telemetry compatibility routes. + +OpenCode clients >= 1.4.4 may POST telemetry to /v1/metrics, /v1/traces, +and /v1/logs. These endpoints must return a success status instead of 405. +""" + +from __future__ import annotations + +from fastapi import FastAPI, Request +from fastapi.responses import Response +from fastapi.testclient import TestClient +import pytest + + +@pytest.fixture +def otlp_app() -> FastAPI: + """Create a minimal FastAPI app with only the OTLP sink routes.""" + app = FastAPI() + + @app.post("/v1/metrics") + async def otlp_metrics(request: Request) -> Response: + return Response(status_code=204) + + @app.post("/v1/traces") + async def otlp_traces(request: Request) -> Response: + return Response(status_code=204) + + @app.post("/v1/logs") + async def otlp_logs(request: Request) -> Response: + return Response(status_code=204) + + # Catch-all that only accepts GET (mirrors production server) + @app.api_route("/{path:path}", methods=["GET", "HEAD", "OPTIONS"]) + async def catch_all(request: Request, path: str) -> Response: + return Response(status_code=404) + + return app + + +@pytest.fixture +def otlp_client(otlp_app: FastAPI) -> TestClient: + """Create a test client for the OTLP app.""" + return TestClient(otlp_app) + + +@pytest.mark.parametrize( + "path", + ["/v1/metrics", "/v1/traces", "/v1/logs"], +) +def test_otlp_post_returns_success(otlp_client: TestClient, path: str) -> None: + """POST to OTLP endpoints should return 204, not 405.""" + response = otlp_client.post(path, content=b"") + assert response.status_code == 204 + + +@pytest.mark.parametrize( + "path", + ["/v1/metrics", "/v1/traces", "/v1/logs"], +) +def test_otlp_post_with_body_returns_success(otlp_client: TestClient, path: str) -> None: + """POST with a body should also succeed — payloads are discarded.""" + response = otlp_client.post(path, content=b"some telemetry data") + assert response.status_code == 204 + + +@pytest.mark.parametrize( + "path", + ["/v1/metrics", "/v1/traces", "/v1/logs"], +) +def test_otlp_get_does_not_return_post_content(otlp_client: TestClient, path: str) -> None: + """GET on OTLP endpoints should not hit the POST handler (falls to catch-all).""" + response = otlp_client.get(path) + # The POST endpoint returns 204; GET must NOT return 204. + assert response.status_code != 204 diff --git a/tests/servers/opencode_server/test_prompt_async.py b/tests/servers/opencode_server/test_prompt_async.py new file mode 100644 index 000000000..95bd86e83 --- /dev/null +++ b/tests/servers/opencode_server/test_prompt_async.py @@ -0,0 +1,753 @@ +"""Regression tests for async prompt handling in OpenCode server.""" + +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING +from unittest.mock import AsyncMock, Mock + +import pytest + +from agentpool_server.opencode_server.models import MessageRequest, TextPartInput +from agentpool_server.opencode_server.models.common import TimeCreated +from agentpool_server.opencode_server.models.events import SessionIdleEvent, SessionStatusEvent +from agentpool_server.opencode_server.models.message import MessageWithParts, UserMessage +from agentpool_server.opencode_server.routes import message_routes + + +if TYPE_CHECKING: + from collections.abc import Awaitable + + +def async_mock_return_value(value): + """Create an async function that returns *value*, suitable for monkeypatching.""" + + async def _mock(*args, **kwargs): + return value + + return _mock + + +class TestPromptAsync: + """Tests for `/prompt_async` session serialization.""" + + @pytest.mark.asyncio + async def test_prompt_async_marks_busy_before_scheduling( + self, + async_client, + server_state, + ) -> None: + """The first async prompt should lock the session before scheduling work.""" + response = await async_client.post("/session", json={"title": "Async Lock"}) + session_id = response.json()["id"] + + background_calls: list[str | None] = [] + + def fake_create_background_task(coro, *, name=None): + background_calls.append(name) + coro.close() + task = Mock() + task.get_name.return_value = name + task.done.return_value = False + server_state.background_tasks.add(task) + return task + + server_state.create_background_task = Mock(side_effect=fake_create_background_task) + + request = MessageRequest( + parts=[TextPartInput(text="first")], + agent="default", + message_id="msg-1", + ) + response = await async_client.post( + f"/session/{session_id}/prompt_async", + json=request.model_dump(mode="json"), + ) + assert response.status_code == 204 + assert server_state.session_status[session_id].type == "busy" + assert server_state.create_background_task.call_count == 1 + + second_request = MessageRequest( + parts=[TextPartInput(text="second")], + agent="default", + message_id="msg-2", + ) + response = await async_client.post( + f"/session/{session_id}/prompt_async", + json=second_request.model_dump(mode="json"), + ) + assert response.status_code == 204 + assert server_state.create_background_task.call_count == 1 + assert background_calls == [f"process_message_{session_id}"] + assert len(server_state.pending_async_prompts[session_id]) == 2 + + @pytest.mark.asyncio + async def test_prompt_async_drains_server_queue_in_order( + self, + async_client, + server_state, + monkeypatch, + ) -> None: + """Queued async prompts should be processed FIFO by one background worker.""" + response = await async_client.post("/session", json={"title": "Async Queue"}) + session_id = response.json()["id"] + + processed: list[str] = [] + drained = asyncio.Event() + + async def fake_process_message_locked( + session_id: str, + request: MessageRequest, + state, + user_msg_id: str, + user_msg_with_parts, + *, + mark_busy: bool = True, + mark_idle: bool = True, + ): + processed.append(request.parts[0].text) + if len(processed) == 2: + drained.set() + return user_msg_with_parts + + monkeypatch.setattr( + message_routes, + "_process_message_locked", + fake_process_message_locked, + ) + + first_request = MessageRequest( + parts=[TextPartInput(text="first")], + agent="default", + message_id="msg-1", + ) + second_request = MessageRequest( + parts=[TextPartInput(text="second")], + agent="default", + message_id="msg-2", + ) + + first_response = await async_client.post( + f"/session/{session_id}/prompt_async", + json=first_request.model_dump(mode="json"), + ) + second_response = await async_client.post( + f"/session/{session_id}/prompt_async", + json=second_request.model_dump(mode="json"), + ) + + assert first_response.status_code == 204 + assert second_response.status_code == 204 + + await asyncio.wait_for(drained.wait(), timeout=1.0) + await asyncio.sleep(0) + + assert processed == ["first", "second"] + assert session_id not in server_state.pending_async_prompts + assert server_state.session_status[session_id].type == "idle" + + @pytest.mark.asyncio + async def test_prompt_async_emits_turn_complete_between_queued_prompts( + self, + server_state, + monkeypatch, + ) -> None: + """Queued prompts should emit a turn-complete idle signal between turns.""" + session = await server_state.ensure_session("async-turn-complete") + session_id = session.id + + event_types: list[str] = [] + + original_broadcast = server_state.broadcast_event + + async def tracking_broadcast(event) -> None: + if isinstance(event, SessionStatusEvent): + event_types.append(f"status:{event.properties.status.type}") + elif isinstance(event, SessionIdleEvent): + event_types.append("session.idle") + await original_broadcast(event) + + server_state.broadcast_event = tracking_broadcast # type: ignore[method-assign] + + for idx in range(2): + request = MessageRequest( + parts=[TextPartInput(text=f"prompt-{idx}")], + agent="default", + message_id=f"msg-{idx}", + ) + queued_user = UserMessage( + id=f"msg-{idx}", + session_id=session_id, + time=TimeCreated(created=idx), + agent="default", + model=None, + ) + server_state.enqueue_async_prompt( + session_id, + message_routes.QueuedAsyncPrompt( + request=request, + user_msg_id=f"msg-{idx}", + user_msg_with_parts=MessageWithParts(info=queued_user), + ), + ) + + server_state.session_status[session_id] = message_routes.SessionStatus(type="busy") + + async def fake_process_message_locked( + session_id: str, + request: MessageRequest, + state, + user_msg_id: str, + user_msg_with_parts, + *, + mark_busy: bool = True, + mark_idle: bool = True, + ): + return user_msg_with_parts + + monkeypatch.setattr(message_routes, "_process_message_locked", fake_process_message_locked) + + await message_routes._run_async_prompt_queue(session_id, server_state) + + assert event_types.count("session.idle") == 2 + assert event_types == ["session.idle", "status:idle", "session.idle"] + + @pytest.mark.asyncio + async def test_ensure_async_prompt_worker_starts_worker_for_queued_prompts( + self, + server_state, + ) -> None: + """Queued async prompts should start a worker when a turn hands off.""" + session = await server_state.ensure_session("sync-handoff") + session_id = session.id + + request = MessageRequest( + parts=[TextPartInput(text="queued")], + agent="default", + message_id="queued-msg", + ) + queued_user = UserMessage( + id="queued-msg", + session_id=session_id, + time=TimeCreated(created=0), + agent="default", + model=None, + ) + server_state.enqueue_async_prompt( + session_id, + message_routes.QueuedAsyncPrompt( + request=request, + user_msg_id="queued-msg", + user_msg_with_parts=MessageWithParts(info=queued_user), + ), + ) + + started_workers: list[str | None] = [] + + def fake_create_background_task(coro: Awaitable[object], *, name: str | None = None): + started_workers.append(name) + coro.close() + return Mock() + + server_state.create_background_task = fake_create_background_task # type: ignore[method-assign] + + await message_routes._ensure_async_prompt_worker(session_id, server_state, mark_busy=True) + + assert started_workers == [f"process_message_{session_id}"] + assert server_state.session_status[session_id].type == "busy" + + @pytest.mark.asyncio + async def test_handoff_skips_idle_when_async_prompts_queued( + self, + server_state, + monkeypatch, + ) -> None: + """When mark_idle=True and async prompts are queued with no worker, + skip idle→busy flicker.""" + session = await server_state.ensure_session("handoff-flicker") + session_id = session.id + + # Enqueue a pending async prompt so has_pending_async_prompts returns True. + request = MessageRequest( + parts=[TextPartInput(text="queued")], + agent="default", + message_id="msg-queued", + ) + queued_user = UserMessage( + id="msg-queued", + session_id=session_id, + time=TimeCreated(created=0), + agent="default", + model=None, + ) + server_state.enqueue_async_prompt( + session_id, + message_routes.QueuedAsyncPrompt( + request=request, + user_msg_id="msg-queued", + user_msg_with_parts=MessageWithParts(info=queued_user), + ), + ) + + # Track status/idle events. + event_types: list[str] = [] + original_broadcast = server_state.broadcast_event + + async def tracking_broadcast(event) -> None: + if isinstance(event, SessionStatusEvent): + event_types.append(f"status:{event.properties.status.type}") + elif isinstance(event, SessionIdleEvent): + event_types.append("session.idle") + await original_broadcast(event) + + server_state.broadcast_event = tracking_broadcast # type: ignore[method-assign] + + user_msg = UserMessage( + id="msg-handoff", + session_id=session_id, + time=TimeCreated(created=0), + agent="default", + model=None, + ) + msg_with_parts = MessageWithParts(info=user_msg) + + # Mock the agent's run_stream to yield nothing (empty response). + async def empty_run_stream(*args, **kwargs): + return + yield # noqa: unreachable — makes this an async generator + + server_state.agent.run_stream = empty_run_stream # type: ignore[assignment] + + # Mock extract_user_prompt_from_parts to return a simple text prompt. + monkeypatch.setattr( + message_routes, + "extract_user_prompt_from_parts", + async_mock_return_value(["hello"]), + ) + + # Prevent background tasks (title gen, async prompt worker) from running. + worker_names: list[str | None] = [] + + def fake_create_background_task(coro, *, name=None): + worker_names.append(name) + coro.close() + return Mock() + + server_state.create_background_task = fake_create_background_task # type: ignore[method-assign] + + assert not server_state.has_session_background_task(session_id) + + # Call the REAL _process_message_locked — the handoff logic at + # lines 487-497 will execute against actual state methods. + await message_routes._process_message_locked( + session_id, + request, + server_state, + "msg-handoff", + msg_with_parts, + mark_busy=True, + mark_idle=True, + ) + + # No status:idle should appear — the real handoff code skipped + # mark_session_idle and went straight to _ensure_async_prompt_worker. + assert "status:idle" not in event_types, f"Expected no status:idle but got {event_types}" + # The async prompt worker should have been started. + assert f"process_message_{session_id}" in worker_names + + @pytest.mark.asyncio + async def test_handoff_emits_idle_when_no_async_prompts_queued( + self, + server_state, + monkeypatch, + ) -> None: + """When mark_idle=True and no async prompts are queued, idle is emitted normally.""" + session = await server_state.ensure_session("handoff-idle-normal") + session_id = session.id + + event_types: list[str] = [] + original_broadcast = server_state.broadcast_event + + async def tracking_broadcast(event) -> None: + if isinstance(event, SessionStatusEvent): + event_types.append(f"status:{event.properties.status.type}") + elif isinstance(event, SessionIdleEvent): + event_types.append("session.idle") + await original_broadcast(event) + + server_state.broadcast_event = tracking_broadcast # type: ignore[method-assign] + + request = MessageRequest( + parts=[TextPartInput(text="hello")], + agent="default", + message_id="msg-idle", + ) + user_msg = UserMessage( + id="msg-idle", + session_id=session_id, + time=TimeCreated(created=0), + agent="default", + model=None, + ) + msg_with_parts = MessageWithParts(info=user_msg) + + # Mock the agent's run_stream to yield nothing (empty response). + async def empty_run_stream(*args, **kwargs): + return + yield # noqa: unreachable — makes this an async generator + + server_state.agent.run_stream = empty_run_stream # type: ignore[assignment] + + # Mock extract_user_prompt_from_parts to return a simple text prompt. + monkeypatch.setattr( + message_routes, + "extract_user_prompt_from_parts", + async_mock_return_value(["hello"]), + ) + + # Prevent background tasks from running. + worker_names: list[str | None] = [] + + def fake_create_background_task(coro, *, name=None): + worker_names.append(name) + coro.close() + return Mock() + + server_state.create_background_task = fake_create_background_task # type: ignore[method-assign] + + assert not server_state.has_pending_async_prompts(session_id) + + # Call the REAL _process_message_locked. + await message_routes._process_message_locked( + session_id, + request, + server_state, + "msg-idle", + msg_with_parts, + mark_busy=True, + mark_idle=True, + ) + + # The real mark_session_idle emits SessionStatusEvent(idle) then + # SessionIdleEvent, so the full sequence is: + # status:busy (mark_busy) → status:idle (mark_session_idle) → session.idle + assert event_types == ["status:busy", "status:idle", "session.idle"] + # No async prompt worker started — nothing was queued. + assert f"process_message_{session_id}" not in worker_names + + @pytest.mark.asyncio + async def test_snapshot_binds_resolved_agent( + self, + server_state, + ) -> None: + """snapshot_for_session(agent=...) must bind the resolved agent, not self.agent.""" + session_id = "snapshot-resolved" + + # Create an alternate agent that differs from the default. + alt_agent = Mock() + alt_agent.model_name = "alt-model-v2" + alt_agent._input_provider = None + alt_agent._current_mode = "reasoning" + + # Ensure the default agent also has model_name so the comparison is meaningful. + server_state.agent.model_name = "default-model" + + snapshot = await server_state.snapshot_for_session(session_id, agent=alt_agent) + + # The snapshot must reflect the alternate agent, not the default. + assert snapshot.model_name == "alt-model-v2" + assert snapshot.mode_name == "reasoning" + # The alternate agent must have received the input provider and session id. + assert alt_agent._input_provider is server_state.input_providers[session_id] + assert alt_agent.session_id == session_id + # The default agent must NOT have been touched. + assert server_state.agent._input_provider is None + + @pytest.mark.asyncio + async def test_snapshot_default_agent_unchanged( + self, + server_state, + ) -> None: + """snapshot_for_session(session_id) without agent= still binds self.agent.""" + session_id = "snapshot-default" + + server_state.agent.model_name = "default-model" + + snapshot = await server_state.snapshot_for_session(session_id) + + # The snapshot must reflect the default agent. + assert snapshot.model_name == "default-model" + # The default agent must have received the input provider and session id. + assert server_state.agent._input_provider is server_state.input_providers[session_id] + assert server_state.agent.session_id == session_id + + @pytest.mark.asyncio + async def test_snapshot_mode_name_after_mutation( + self, + server_state, + monkeypatch, + ) -> None: + """Snapshot must reflect mode_name AFTER set_mode mutation, not before.""" + from agentpool_server.opencode_server.models.common import ModelRef + + session = await server_state.ensure_session("snapshot-mode-mutation") + session_id = session.id + + # Set up agent with initial mode + server_state.agent.model_name = "test-model" + server_state.agent._current_mode = "low" + + # Make set_mode actually mutate _current_mode on the mock agent + original_set_mode = server_state.agent.set_mode + + async def fake_set_mode(variant: str, **kwargs: object) -> None: + server_state.agent._current_mode = variant + + server_state.agent.set_mode = fake_set_mode + server_state.agent.get_available_models = AsyncMock(return_value=[]) + + # Mock run_stream and extract to allow _process_message_locked to complete + async def empty_run_stream(*args, **kwargs): + return + yield # noqa: unreachable — makes this an async generator + + server_state.agent.run_stream = empty_run_stream # type: ignore[assignment] + + monkeypatch.setattr( + message_routes, + "extract_user_prompt_from_parts", + async_mock_return_value(["hello"]), + ) + + # Prevent background tasks from running + def fake_create_background_task(coro, *, name=None): + coro.close() + return Mock() + + server_state.create_background_task = fake_create_background_task # type: ignore[method-assign] + + request = MessageRequest( + parts=[TextPartInput(text="hello")], + agent=None, + message_id="msg-mutation", + model=ModelRef(provider_id="test", model_id="test-model", variant="high"), + ) + user_msg = UserMessage( + id="msg-mutation", + session_id=session_id, + time=TimeCreated(created=0), + agent="default", + model=request.model, + ) + msg_with_parts = MessageWithParts(info=user_msg) + + # We need to capture the snapshot that _process_message_locked uses. + # Intercept snapshot_for_session to capture the returned snapshot. + captured_snapshot = None + original_snapshot = server_state.snapshot_for_session + + async def capturing_snapshot(*args, **kwargs): + nonlocal captured_snapshot + snap = await original_snapshot(*args, **kwargs) + captured_snapshot = snap + return snap + + server_state.snapshot_for_session = capturing_snapshot # type: ignore[method-assign] + + await message_routes._process_message_locked( + session_id, + request, + server_state, + "msg-mutation", + msg_with_parts, + mark_busy=True, + mark_idle=True, + ) + + # The snapshot must reflect the post-mutation mode_name + assert captured_snapshot is not None + assert captured_snapshot.mode_name == "high", ( + f"Expected mode_name='high' but got '{captured_snapshot.mode_name}'" + ) + + @pytest.mark.asyncio + async def test_async_queue_survives_prompt_failure( + self, + server_state, + monkeypatch, + ) -> None: + """A single prompt failure should not kill remaining queued prompts.""" + session = await server_state.ensure_session("queue-resilience") + session_id = session.id + + processed: list[str] = [] + call_count = 0 + + async def fake_process_message_locked( + session_id: str, + request: MessageRequest, + state, + user_msg_id: str, + user_msg_with_parts, + *, + mark_busy: bool = True, + mark_idle: bool = True, + ): + nonlocal call_count + call_count += 1 + text = request.parts[0].text + # Make the 2nd call fail + if call_count == 2: + msg = "Simulated prompt failure" + raise RuntimeError(msg) + processed.append(text) + return user_msg_with_parts + + monkeypatch.setattr( + message_routes, + "_process_message_locked", + fake_process_message_locked, + ) + + # Enqueue 3 prompts + for idx in range(3): + request = MessageRequest( + parts=[TextPartInput(text=f"prompt-{idx}")], + agent="default", + message_id=f"msg-{idx}", + ) + queued_user = UserMessage( + id=f"msg-{idx}", + session_id=session_id, + time=TimeCreated(created=idx), + agent="default", + model=None, + ) + server_state.enqueue_async_prompt( + session_id, + message_routes.QueuedAsyncPrompt( + request=request, + user_msg_id=f"msg-{idx}", + user_msg_with_parts=MessageWithParts(info=queued_user), + ), + ) + + server_state.session_status[session_id] = message_routes.SessionStatus(type="busy") + + # Run the queue worker + await message_routes._run_async_prompt_queue(session_id, server_state) + + # Prompts 1 and 3 should have been processed (2nd failed but queue continued) + assert processed == ["prompt-0", "prompt-2"], f"Expected ['prompt-0', 'prompt-2'] but got {processed}" + # Session should end idle + assert server_state.session_status[session_id].type == "idle" + # No orphaned prompts + assert not server_state.has_pending_async_prompts(session_id) + + @pytest.mark.asyncio + async def test_handoff_no_idle_when_worker_running( + self, + server_state, + monkeypatch, + ) -> None: + """When has_queued=True AND has_worker=True, no idle status is emitted.""" + session = await server_state.ensure_session("handoff-worker-running") + session_id = session.id + + # Enqueue a pending async prompt so has_pending_async_prompts returns True. + request = MessageRequest( + parts=[TextPartInput(text="queued")], + agent="default", + message_id="msg-queued", + ) + queued_user = UserMessage( + id="msg-queued", + session_id=session_id, + time=TimeCreated(created=0), + agent="default", + model=None, + ) + server_state.enqueue_async_prompt( + session_id, + message_routes.QueuedAsyncPrompt( + request=request, + user_msg_id="msg-queued", + user_msg_with_parts=MessageWithParts(info=queued_user), + ), + ) + + # Track status/idle events. + event_types: list[str] = [] + original_broadcast = server_state.broadcast_event + + async def tracking_broadcast(event) -> None: + if isinstance(event, SessionStatusEvent): + event_types.append(f"status:{event.properties.status.type}") + elif isinstance(event, SessionIdleEvent): + event_types.append("session.idle") + await original_broadcast(event) + + server_state.broadcast_event = tracking_broadcast # type: ignore[method-assign] + + user_msg = UserMessage( + id="msg-handoff", + session_id=session_id, + time=TimeCreated(created=0), + agent="default", + model=None, + ) + msg_with_parts = MessageWithParts(info=user_msg) + + # Mock the agent's run_stream to yield nothing (empty response). + async def empty_run_stream(*args, **kwargs): + return + yield # noqa: unreachable — makes this an async generator + + server_state.agent.run_stream = empty_run_stream # type: ignore[assignment] + + # Mock extract_user_prompt_from_parts to return a simple text prompt. + monkeypatch.setattr( + message_routes, + "extract_user_prompt_from_parts", + async_mock_return_value(["hello"]), + ) + + # Simulate a running worker by creating a mock task with the expected name. + worker_names: list[str | None] = [] + + def fake_create_background_task(coro, *, name=None): + worker_names.append(name) + coro.close() + task = Mock() + task.get_name.return_value = name + task.done.return_value = False + server_state.background_tasks.add(task) + return task + + server_state.create_background_task = fake_create_background_task # type: ignore[method-assign] + + # Simulate that a worker is already running + mock_task = Mock() + mock_task.get_name.return_value = f"process_message_{session_id}" + mock_task.done.return_value = False + server_state.background_tasks.add(mock_task) + + assert server_state.has_session_background_task(session_id) + + # Call the REAL _process_message_locked — the handoff logic will + # execute with has_queued=True AND has_worker=True. + await message_routes._process_message_locked( + session_id, + request, + server_state, + "msg-handoff", + msg_with_parts, + mark_busy=True, + mark_idle=True, + ) + + # No idle should appear — the existing worker will drain the queue + assert "status:idle" not in event_types, f"Expected no status:idle but got {event_types}" + assert "session.idle" not in event_types, f"Expected no session.idle but got {event_types}" + # No new worker should be started since one is already running + assert f"process_message_{session_id}" not in worker_names diff --git a/tests/servers/opencode_server/test_question_integration.py b/tests/servers/opencode_server/test_question_integration.py index 4277fddc7..36a811ff9 100644 --- a/tests/servers/opencode_server/test_question_integration.py +++ b/tests/servers/opencode_server/test_question_integration.py @@ -3,12 +3,21 @@ from __future__ import annotations import asyncio -from unittest.mock import Mock +from unittest.mock import AsyncMock, Mock from mcp import types import pytest -from agentpool_server.opencode_server.input_provider import OpenCodeInputProvider +from agentpool_server.opencode_server.input_provider import OpenCodeInputProvider, PendingPermission +from agentpool_server.opencode_server.models import ( + PermissionRequestEvent, + PermissionResolvedEvent, + QuestionReply, +) +from agentpool_server.opencode_server.routes.question_routes import ( + reject_question, + reply_to_question, +) from agentpool_server.opencode_server.state import ServerState @@ -261,8 +270,97 @@ async def test_multi_question_cancellation(): assert isinstance(result, types.ElicitResult) assert result.action == "cancel" - # Clean up if still present - assert question_id not in state.pending_questions + +async def test_question_reply_can_resolve_permission_request(): + """Permission replies routed through /question should still resolve.""" + mock_agent = Mock() + mock_agent.agent_pool = None + state = ServerState(working_dir="/tmp", agent=mock_agent) + provider = OpenCodeInputProvider(state=state, session_id="test_session") + state.input_providers["test_session"] = provider + state.broadcast_event = AsyncMock() + + permission_id = "perm_1_1776434635956" + future = asyncio.get_running_loop().create_future() + provider._pending_permissions[permission_id] = PendingPermission( + permission_id=permission_id, + tool_name="bash", + args={"command": "echo test"}, + future=future, + ) + + result = await reply_to_question( + permission_id, + QuestionReply(answers=[["once"]]), + state, + ) + + assert result is True + assert future.done() + assert future.result() == "once" + assert state.broadcast_event.await_count == 1 + event = state.broadcast_event.await_args.args[0] + assert isinstance(event, PermissionResolvedEvent) + assert event.properties.request_id == permission_id + assert event.properties.reply == "once" + + +async def test_question_reject_can_resolve_permission_request(): + """Permission rejects routed through /question should still resolve.""" + mock_agent = Mock() + mock_agent.agent_pool = None + state = ServerState(working_dir="/tmp", agent=mock_agent) + provider = OpenCodeInputProvider(state=state, session_id="test_session") + state.input_providers["test_session"] = provider + state.broadcast_event = AsyncMock() + + permission_id = "perm_2_1776434635957" + future = asyncio.get_running_loop().create_future() + provider._pending_permissions[permission_id] = PendingPermission( + permission_id=permission_id, + tool_name="bash", + args={"command": "echo reject"}, + future=future, + ) + + result = await reject_question(permission_id, state) + + assert result is True + assert future.done() + assert future.result() == "reject" + assert state.broadcast_event.await_count == 1 + event = state.broadcast_event.await_args.args[0] + assert isinstance(event, PermissionResolvedEvent) + assert event.properties.request_id == permission_id + assert event.properties.reply == "reject" + + +async def test_permission_request_uses_permission_prefix(): + """Permission requests should keep the permission ID namespace.""" + mock_agent = Mock() + mock_agent.agent_pool = None + state = ServerState(working_dir="/tmp", agent=mock_agent) + provider = OpenCodeInputProvider(state=state, session_id="test_session") + state.broadcast_event = AsyncMock() + + context = Mock() + context.tool_name = "bash" + context.tool_input = {"command": "echo test"} + context.tool_call_id = "call-123" + + task = asyncio.create_task(provider.get_tool_confirmation(context)) + await asyncio.sleep(0.1) + + assert state.broadcast_event.await_count == 1 + event = state.broadcast_event.await_args.args[0] + assert isinstance(event, PermissionRequestEvent) + assert event.properties.id.startswith("perm_") + + resolved = provider.resolve_permission(event.properties.id, "once") + assert resolved is True + + result = await task + assert result == "allow" async def test_multi_question_partial_answers(): @@ -347,6 +445,7 @@ async def test_multi_question_rfc0010_backward_compat(): assert len(state.pending_questions) == 1 question_id = next(iter(state.pending_questions.keys())) pending = state.pending_questions[question_id] + assert question_id.startswith("que_") # Single question in multi-question format assert len(pending.questions) == 1 diff --git a/tests/servers/opencode_server/test_question_routes.py b/tests/servers/opencode_server/test_question_routes.py new file mode 100644 index 000000000..004340b58 --- /dev/null +++ b/tests/servers/opencode_server/test_question_routes.py @@ -0,0 +1,121 @@ +"""Tests for question_routes permission lookup using public API.""" + +from __future__ import annotations + +import asyncio +from unittest.mock import AsyncMock, Mock + +import pytest + +from agentpool_server.opencode_server.input_provider import OpenCodeInputProvider, PendingPermission +from agentpool_server.opencode_server.routes.question_routes import _find_permission_provider +from agentpool_server.opencode_server.state import ServerState + + +def _make_state_with_providers( + providers: dict[str, OpenCodeInputProvider], +) -> ServerState: + """Create a ServerState with pre-populated input_providers.""" + mock_agent = Mock() + mock_agent.agent_pool = None + state = ServerState(working_dir="/tmp", agent=mock_agent) + state.input_providers = providers + return state + + +async def test_find_permission_provider_finds_matching_permission(): + """_find_permission_provider should find provider with matching pending permission.""" + mock_agent = Mock() + mock_agent.agent_pool = None + state = ServerState(working_dir="/tmp", agent=mock_agent) + provider = OpenCodeInputProvider(state=state, session_id="sess_1") + state.input_providers["sess_1"] = provider + + # Add a pending permission + permission_id = "perm_1_1234" + future = asyncio.get_running_loop().create_future() + provider._pending_permissions[permission_id] = PendingPermission( + permission_id=permission_id, + tool_name="bash", + args={"command": "echo test"}, + future=future, + ) + + result = _find_permission_provider(state, permission_id) + assert result is not None + found_session_id, found_provider = result + assert found_session_id == "sess_1" + assert found_provider is provider + + +async def test_find_permission_provider_returns_none_for_unknown_permission(): + """_find_permission_provider should return None when no provider has the permission.""" + mock_agent = Mock() + mock_agent.agent_pool = None + state = ServerState(working_dir="/tmp", agent=mock_agent) + provider = OpenCodeInputProvider(state=state, session_id="sess_1") + state.input_providers["sess_1"] = provider + + result = _find_permission_provider(state, "nonexistent_perm") + assert result is None + + +async def test_find_permission_provider_uses_public_api(): + """Verify that _find_permission_provider delegates to has_pending_permission(). + + This test ensures the function uses the public method rather than + directly accessing _pending_permissions, by checking the behavior + matches what has_pending_permission() would return. + """ + mock_agent = Mock() + mock_agent.agent_pool = None + state = ServerState(working_dir="/tmp", agent=mock_agent) + provider = OpenCodeInputProvider(state=state, session_id="sess_1") + state.input_providers["sess_1"] = provider + + permission_id = "perm_2_5678" + future = asyncio.get_running_loop().create_future() + provider._pending_permissions[permission_id] = PendingPermission( + permission_id=permission_id, + tool_name="bash", + args={"command": "ls"}, + future=future, + ) + + # Verify has_pending_permission works (public API) + assert provider.has_pending_permission(permission_id) is True + assert provider.has_pending_permission("nonexistent") is False + + # Verify _find_permission_provider returns the same result as has_pending_permission + result = _find_permission_provider(state, permission_id) + assert result is not None + + result_missing = _find_permission_provider(state, "nonexistent") + assert result_missing is None + + +async def test_find_permission_provider_multiple_providers(): + """_find_permission_provider should find the correct provider among multiple.""" + mock_agent = Mock() + mock_agent.agent_pool = None + state = ServerState(working_dir="/tmp", agent=mock_agent) + + provider_a = OpenCodeInputProvider(state=state, session_id="sess_a") + provider_b = OpenCodeInputProvider(state=state, session_id="sess_b") + state.input_providers["sess_a"] = provider_a + state.input_providers["sess_b"] = provider_b + + permission_id = "perm_3_9999" + future = asyncio.get_running_loop().create_future() + provider_b._pending_permissions[permission_id] = PendingPermission( + permission_id=permission_id, + tool_name="bash", + args={"command": "rm -rf /tmp/test"}, + future=future, + ) + + result = _find_permission_provider(state, permission_id) + assert result is not None + found_session_id, found_provider = result + assert found_session_id == "sess_b" + assert found_provider is provider_b diff --git a/tests/servers/opencode_server/test_restart_recovery.py b/tests/servers/opencode_server/test_restart_recovery.py new file mode 100644 index 000000000..684ed636b --- /dev/null +++ b/tests/servers/opencode_server/test_restart_recovery.py @@ -0,0 +1,103 @@ +"""Regression tests for cold-start recovery after server restart.""" + +from __future__ import annotations + +from datetime import UTC, datetime +import json +from typing import TYPE_CHECKING +from unittest.mock import Mock + +import pytest + +from agentpool.sessions.models import SessionData +from agentpool_server.opencode_server.models import SessionIdleEvent +from agentpool_server.opencode_server.routes.session_routes import get_or_load_session +from agentpool_server.opencode_server.state import ServerState +from agentpool_storage.opencode_provider import helpers + + +if TYPE_CHECKING: + from pathlib import Path + + +class TestRestartRecovery: + """Verify persisted sessions recover correctly after a fresh server start.""" + + async def test_get_or_load_session_restores_runtime_state_after_restart( + self, + server_state: ServerState, + tmp_project_dir: Path, + event_capture, + ): + """Cold-start loading should rebuild all runtime buckets for a persisted session.""" + session_id = "restart-session" + now = datetime.now(UTC) + session_data = SessionData( + session_id=session_id, + agent_name="test-agent", + cwd=str(tmp_project_dir), + created_at=now, + last_active=now, + metadata={"title": "Recovered Session"}, + ) + + server_state.agent.session_id = None + server_state.agent._input_provider = None + server_state.agent.conversation = Mock() + server_state.agent.conversation.chat_messages = [] + + async def mock_load_session(sid: str) -> SessionData | None: + if sid == session_id: + return session_data + return None + + server_state.agent.load_session = mock_load_session # type: ignore[method-assign] + + loaded_session = await get_or_load_session(server_state, session_id) + + assert loaded_session is not None + assert loaded_session.id == session_id + assert loaded_session.directory == str(tmp_project_dir) + assert server_state.agent.session_id == session_id + assert server_state.agent._input_provider is server_state.input_providers[session_id] + assert session_id in server_state.sessions + assert server_state.messages[session_id] == [] + assert server_state.reverted_messages[session_id] == [] + assert server_state.todos[session_id] == [] + assert server_state.session_status[session_id].type == "idle" + + status_events = [ + event + for event in event_capture.get_events_by_type("session.status") + if event.properties.session_id == session_id + ] + idle_events = [ + event + for event in event_capture.get_events_by_type("session.idle") + if event.properties.session_id == session_id + ] + assert status_events + assert idle_events + + def test_event_factory_uses_resolved_directory_for_restart_routing( + self, + tmp_project_dir: Path, + mock_agent: Mock, + ): + """Global event routing metadata should stay stable across restart path aliases.""" + nested = tmp_project_dir / "nested" + nested.mkdir() + aliased_working_dir = str(nested / "..") + + state = ServerState(working_dir=aliased_working_dir, agent=mock_agent) + payload = json.loads(state.get_event_factory().wrap(SessionIdleEvent.create("sess-1"))) + + resolved_dir = str(tmp_project_dir.resolve()) + assert payload["directory"] == resolved_dir + assert payload["project"] == helpers.compute_project_id(resolved_dir) + assert payload["payload"]["type"] == "session.idle" + assert payload["payload"]["sessionId"] == "sess-1" + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/tests/servers/opencode_server/test_session_history_loading.py b/tests/servers/opencode_server/test_session_history_loading.py index 35f8b60a0..fdd10d6ba 100644 --- a/tests/servers/opencode_server/test_session_history_loading.py +++ b/tests/servers/opencode_server/test_session_history_loading.py @@ -7,8 +7,8 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any -from unittest.mock import AsyncMock, Mock +from typing import TYPE_CHECKING +from unittest.mock import Mock import pytest @@ -32,6 +32,7 @@ async def test_session_switch_reloads_history( async_client: AsyncClient, server_state: ServerState, tmp_project_dir: Path, + event_capture, ): """When switching sessions, agent should reload conversation history.""" # Setup: Create session A @@ -88,6 +89,18 @@ async def mock_load_session(sid: str) -> SessionData | None: assert loaded_session is not None assert loaded_session.id == session_a_id assert server_state.agent.session_id == session_a_id + status_events = [ + event + for event in event_capture.get_events_by_type("session.status") + if event.properties.session_id == session_a_id + ] + idle_events = [ + event + for event in event_capture.get_events_by_type("session.idle") + if event.properties.session_id == session_a_id + ] + assert status_events + assert idle_events async def test_cached_session_with_wrong_agent_session_gets_reloaded( self, @@ -154,6 +167,7 @@ async def mock_load_session(sid: str) -> SessionData | None: loaded_session = await get_or_load_session(server_state, session_a_id) # VERIFY: load_session should have been called + assert loaded_session is not None assert load_session_called assert server_state.agent.session_id == session_a_id diff --git a/tests/servers/opencode_server/test_session_lifecycle.py b/tests/servers/opencode_server/test_session_lifecycle.py index f390863de..3f2576137 100644 --- a/tests/servers/opencode_server/test_session_lifecycle.py +++ b/tests/servers/opencode_server/test_session_lifecycle.py @@ -12,13 +12,19 @@ from __future__ import annotations +import asyncio from datetime import UTC, datetime from pathlib import Path from typing import TYPE_CHECKING +from unittest.mock import AsyncMock from agentpool.sessions.models import SessionData from agentpool_server.opencode_server.models import SessionStatus -from agentpool_server.opencode_server.models.events import SessionCreatedEvent +from agentpool_server.opencode_server.models.events import ( + SessionCreatedEvent, + SessionIdleEvent, + SessionStatusEvent, +) if TYPE_CHECKING: @@ -57,6 +63,13 @@ async def test_should_emit_session_created_event_when_session_is_created( assert event.properties.info.id == session_data["id"] assert event.properties.info.title == session_data["title"] assert event.properties.info.project_id == "global" # Non-git directory returns "global" + status_events = event_capture.get_events_by_type("session.status") + idle_events = event_capture.get_events_by_type("session.idle") + assert len(status_events) == 1 + assert len(idle_events) == 1 + assert isinstance(status_events[0], SessionStatusEvent) + assert isinstance(idle_events[0], SessionIdleEvent) + assert status_events[0].properties.status.type == "idle" async def test_session_created_event_should_be_emitted_before_session_updated( self, @@ -99,6 +112,7 @@ async def test_create_session_returns_valid_session( self, async_client: AsyncClient, tmp_project_dir: Path, + event_capture: EventCapture, ): """Creating a session should return a valid session object.""" response = await async_client.post("/session", json={"title": "My Session"}) @@ -114,6 +128,10 @@ async def test_create_session_returns_valid_session( assert "time" in session assert "created" in session["time"] assert "updated" in session["time"] + status_events = event_capture.get_events_by_type("session.status") + idle_events = event_capture.get_events_by_type("session.idle") + assert len(status_events) == 1 + assert len(idle_events) == 1 async def test_create_session_with_parent_id(self, async_client: AsyncClient): """Creating a session with parent_id should set the parent.""" @@ -310,6 +328,36 @@ async def test_abort_session(self, async_client: AsyncClient, server_state: Serv assert abort_response.json() is True assert server_state.session_status[session_id].type == "idle" + async def test_abort_session_cancels_prompt_background_task( + self, + async_client: AsyncClient, + server_state: ServerState, + ): + """Aborting should cancel the in-flight prompt worker for the session.""" + response = await async_client.post("/session", json={"title": "Abort Session"}) + session_id = response.json()["id"] + server_state.agent.interrupt = AsyncMock() + + started = asyncio.Event() + release = asyncio.Event() + + async def background_worker() -> None: + started.set() + await release.wait() + + task_name = f"process_message_{session_id}" + background_task = server_state.create_background_task(background_worker(), name=task_name) + await started.wait() + + abort_response = await async_client.post(f"/session/{session_id}/abort") + assert abort_response.status_code == 200 + assert abort_response.json() is True + + assert background_task.cancelled() + assert task_name not in {task.get_name() for task in server_state.background_tasks} + assert server_state.session_status[session_id].type == "idle" + server_state.agent.interrupt.assert_awaited_once() + async def test_abort_nonexistent_session_returns_404(self, async_client: AsyncClient): """Aborting a non-existent session should return 404.""" response = await async_client.post("/session/nonexistent-id/abort") @@ -342,6 +390,10 @@ async def test_fork_session_creates_new_session_with_parent( fork_event = created_events[-1] assert fork_event.properties.info.id == forked["id"] assert fork_event.properties.info.parent_id == original_id # Python attr + status_events = event_capture.get_events_by_type("session.status") + idle_events = event_capture.get_events_by_type("session.idle") + assert status_events[-1].properties.session_id == forked["id"] + assert idle_events[-1].properties.session_id == forked["id"] async def test_fork_nonexistent_session_returns_404(self, async_client: AsyncClient): """Forking a non-existent session should return 404.""" diff --git a/tests/servers/opencode_server/test_shell.py b/tests/servers/opencode_server/test_shell.py index 972c9e2a4..b8d57b1b7 100644 --- a/tests/servers/opencode_server/test_shell.py +++ b/tests/servers/opencode_server/test_shell.py @@ -8,8 +8,14 @@ from __future__ import annotations +import asyncio from unittest.mock import AsyncMock, Mock +import pytest + +from agentpool_server.opencode_server.models import ShellRequest +from agentpool_server.opencode_server.routes.session_routes import run_shell_command + class TestShellBasic: """Basic shell command execution tests.""" @@ -238,6 +244,30 @@ async def test_session_returns_to_idle_after_execution( # Check final session status assert server_state.session_status[session_id].type == "idle" + async def test_cancelled_shell_command_still_unlocks_session( + self, + async_client, + server_state, + event_capture, + ): + """A cancelled shell command should still broadcast idle state.""" + response = await async_client.post("/session", json={"title": "Cancel Shell"}) + session_id = response.json()["id"] + server_state.agent.env.execute_command = AsyncMock(side_effect=asyncio.CancelledError) + + with pytest.raises(asyncio.CancelledError): + await run_shell_command( + session_id, + ShellRequest(agent="test", command="echo test"), + server_state, + ) + + assert server_state.session_status[session_id].type == "idle" + status_events = event_capture.get_events_by_type("session.status") + idle_events = event_capture.get_events_by_type("session.idle") + assert status_events[-1].properties.status.type == "idle" + assert idle_events[-1].properties.session_id == session_id + class TestShellMessageStructure: """Tests for shell command message/part structure.""" diff --git a/tests/servers/opencode_server/test_sse_compliance.py b/tests/servers/opencode_server/test_sse_compliance.py new file mode 100644 index 000000000..094dde1f3 --- /dev/null +++ b/tests/servers/opencode_server/test_sse_compliance.py @@ -0,0 +1,1145 @@ +"""End-to-end SSE protocol compliance tests. + +Validates that the SSE event stream conforms to the OpenCode TUI protocol: +- /global/event emits events that pass the TUI routing filter +- Directory in GlobalEvent envelope uses the configured server working_dir +- Workspace is omitted for single-directory routing +- server.connected and server.heartbeat keep a payload wrapper on /global/event +- All other events are wrapped in GlobalEvent with correct directory +- sessionId appears at top level of payload for ALL event types +- UserMessage events use nested model.variant format +- FileDiff events use v1.4.0+ schema (patch, not before/after) +- /event endpoint still works for backward compatibility (raw events, no envelope) +- PartDeltaEvent streaming (most critical for "no response" symptom) +""" + +from __future__ import annotations + +import asyncio +import json +from typing import TYPE_CHECKING, Any + +import pytest + +from agentpool_server.opencode_server.models import GlobalEvent +from agentpool_server.opencode_server.models.common import ( + FileDiff, + ModelRef, + TimeCreated, +) +from agentpool_server.opencode_server.models.events import ( + CommandExecutedEvent, + MessageUpdatedEvent, + PartDeltaEvent, + PermissionRequestEvent, + PermissionResolvedEvent, + ServerConnectedEvent, + ServerHeartbeatEvent, + SessionCompactedEvent, + SessionCreatedEvent, + SessionDeletedEvent, + SessionDiffEvent, + SessionStatusEvent, + TodoUpdatedEvent, + TuiSessionSelectEvent, +) +from agentpool_server.opencode_server.models.message import UserMessage +from agentpool_server.opencode_server.models.parts import TextPart +from agentpool_server.opencode_server.models.session import ( + Session, + TimeCreatedUpdated as SessionTimeCreatedUpdated, +) +from agentpool_server.opencode_server.routes.global_routes import ( + GlobalEventFactory, + _event_generator, + _extract_session_id, + _serialize_event, +) +from agentpool_server.opencode_server.routes.routing import tui_event_filter + + +if TYPE_CHECKING: + from agentpool_server.opencode_server.models.events import Event + from agentpool_server.opencode_server.models.parts import Part + + +# ============================================================================= +# Test helpers (reusing _MockState / _collect_events pattern from test_global_event) +# ============================================================================= + + +class _MockState: + """Minimal ServerState-like object for _event_generator tests.""" + + def __init__(self, working_dir: str = "/tmp/test_wd") -> None: + self.working_dir = working_dir + self.event_subscribers: list[asyncio.Queue[Event]] = [] + self._event_factory: GlobalEventFactory | None = None + self._first_subscriber_triggered = False + self.on_first_subscriber: Any = None + + def get_event_factory(self) -> GlobalEventFactory: + if self._event_factory is None: + from agentpool_storage.opencode_provider import helpers + + directory = self.working_dir + self._event_factory = GlobalEventFactory( + directory=directory, + project=helpers.compute_project_id(directory), + ) + return self._event_factory + + def create_background_task(self, coro: Any, name: str = "") -> asyncio.Task[Any]: + return asyncio.ensure_future(coro) + + +async def _collect_events( + state: _MockState, + wrap_payload: bool, + events_to_send: list[Event], +) -> list[dict[str, Any]]: + """Collect SSE items from _event_generator with given events.""" + results: list[dict[str, Any]] = [] + gen = _event_generator(state, wrap_payload=wrap_payload) + # Get the initial connected event + item = await gen.__anext__() + results.append(json.loads(item["data"])) + # Send additional events through the queue + queue = state.event_subscribers[-1] + for event in events_to_send: + await queue.put(event) + item = await gen.__anext__() + results.append(json.loads(item["data"])) + return results + + +def _make_session(session_id: str = "test-sid") -> Session: + """Create a minimal Session for event construction.""" + return Session( + id=session_id, + project_id="proj1", + directory="/tmp", + title="Test", + time=SessionTimeCreatedUpdated(created=0, updated=0), + ) + + +def _make_part(session_id: str = "test-sid") -> Part: + """Create a minimal Part for event construction.""" + return TextPart( + id="part1", + message_id="msg1", + session_id=session_id, + text="hello", + ) + + +# ============================================================================= +# 1. TUI routing filter compliance — /global/event events pass the filter +# ============================================================================= + + +@pytest.mark.anyio +async def test_tui_filter_session_status_event_passes() -> None: + """SessionStatusEvent in GlobalEvent envelope passes tui_event_filter.""" + wd = "/tmp/compliance_test" + state = _MockState(working_dir=wd) + event = SessionStatusEvent.create(session_id="s1", status_type="busy") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + # Build a GlobalEvent from the wire data for filter testing + ge = GlobalEvent( + directory=wrapped["directory"], + project=wrapped.get("project"), + workspace=wrapped.get("workspace"), + payload=wrapped["payload"], + ) + passes, reason = tui_event_filter(ge, state.working_dir) + assert passes, f"SessionStatusEvent should pass filter, got reason={reason}" + + +@pytest.mark.anyio +async def test_tui_filter_part_delta_event_passes() -> None: + """PartDeltaEvent in GlobalEvent envelope passes tui_event_filter.""" + wd = "/tmp/compliance_delta" + state = _MockState(working_dir=wd) + event = PartDeltaEvent.create(session_id="s2", message_id="m1", part_id="p1", delta="hello") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + ge = GlobalEvent( + directory=wrapped["directory"], + project=wrapped.get("project"), + workspace=wrapped.get("workspace"), + payload=wrapped["payload"], + ) + passes, reason = tui_event_filter(ge, state.working_dir) + assert passes, f"PartDeltaEvent should pass filter, got reason={reason}" + + +@pytest.mark.anyio +async def test_tui_filter_all_session_events_pass() -> None: + """Multiple session-scoped events in envelopes all pass tui_event_filter.""" + wd = "/tmp/compliance_multi" + state = _MockState(working_dir=wd) + + session_events: list[Event] = [ + SessionStatusEvent.create(session_id="s1", status_type="busy"), + SessionCompactedEvent.create(session_id="s1"), + SessionDeletedEvent.create(session_id="s1"), + TodoUpdatedEvent.create(session_id="s1", todos=[]), + PartDeltaEvent.create(session_id="s1", message_id="m1", part_id="p1", delta="x"), + ] + events = await _collect_events(state, wrap_payload=True, events_to_send=session_events) + + for i, raw in enumerate(events[1:], start=0): + ge = GlobalEvent( + directory=raw["directory"], + project=raw.get("project"), + workspace=raw.get("workspace"), + payload=raw["payload"], + ) + passes, reason = tui_event_filter(ge, state.working_dir) + assert passes, f"Event {i} ({raw['payload']['type']}) should pass, reason={reason}" + + +# ============================================================================= +# 2. Directory/workspace routing metadata in GlobalEvent envelope +# ============================================================================= + + +@pytest.mark.anyio +async def test_envelope_directory_matches_working_dir() -> None: + """Envelope directory field matches the server working directory.""" + wd = "/custom/exact/working/dir/../dir" + state = _MockState(working_dir=wd) + event = SessionStatusEvent.create(session_id="dir1", status_type="busy") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + assert wrapped["directory"] == state.working_dir + + +@pytest.mark.anyio +async def test_envelope_directory_has_no_trailing_slash() -> None: + """Directory emits without a trailing slash.""" + wd = "/tmp/no_trailing_slash" + state = _MockState(working_dir=wd) + event = PartDeltaEvent.create(session_id="dir2", message_id="m1", part_id="p1", delta="hi") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + assert wrapped["directory"] == state.working_dir + assert not wrapped["directory"].endswith("/") + + +@pytest.mark.anyio +async def test_envelope_workspace_is_omitted() -> None: + """Envelope omits workspace for single-directory routing.""" + wd = "/tmp/workspace_meta/../workspace_meta" + state = _MockState(working_dir=wd) + event = SessionStatusEvent.create(session_id="dir3", status_type="busy") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + assert "workspace" not in wrapped + + +@pytest.mark.anyio +async def test_envelope_directory_different_from_mismatched_path() -> None: + """TUI filter rejects event whose directory doesn't match project_directory.""" + wd = "/correct/dir" + state = _MockState(working_dir=wd) + event = SessionStatusEvent.create(session_id="mismatch1", status_type="idle") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + ge = GlobalEvent( + directory=wrapped["directory"], + project=wrapped.get("project"), + workspace=wrapped.get("workspace"), + payload=wrapped["payload"], + ) + # Should pass with correct directory + passes, _ = tui_event_filter(ge, wd) + assert passes + # Should fail with wrong directory + passes2, reason2 = tui_event_filter(ge, "/wrong/dir") + assert not passes2 + assert reason2 == "directory_mismatch" + + +# ============================================================================= +# 3. server.connected and server.heartbeat keep only a payload wrapper on /global/event +# ============================================================================= + + +@pytest.mark.anyio +async def test_server_connected_is_payload_wrapped() -> None: + """Initial server.connected keeps only the payload wrapper.""" + state = _MockState() + events = await _collect_events(state, wrap_payload=True, events_to_send=[]) + assert len(events) == 1 + connected = events[0] + assert connected["payload"]["type"] == "server.connected" + assert "directory" not in connected + assert "project" not in connected + + +@pytest.mark.anyio +async def test_server_heartbeat_is_payload_wrapped() -> None: + """ServerHeartbeatEvent keeps only the payload wrapper.""" + state = _MockState() + hb = ServerHeartbeatEvent() + events = await _collect_events(state, wrap_payload=True, events_to_send=[hb]) + assert len(events) == 2 + heartbeat = events[1] + assert heartbeat["payload"]["type"] == "server.heartbeat" + assert "directory" not in heartbeat + assert "project" not in heartbeat + + +@pytest.mark.anyio +async def test_server_events_have_no_session_id() -> None: + """server.connected and server.heartbeat lack sessionId at top level.""" + state = _MockState() + hb = ServerHeartbeatEvent() + events = await _collect_events(state, wrap_payload=True, events_to_send=[hb]) + assert "sessionId" not in events[0] # envelope + assert "sessionId" not in events[0]["payload"] # server.connected payload + assert "sessionId" not in events[1] # envelope + assert "sessionId" not in events[1]["payload"] # server.heartbeat payload + + +# ============================================================================= +# 4. All other events are wrapped in GlobalEvent with correct directory +# ============================================================================= + + +@pytest.mark.anyio +async def test_session_created_is_wrapped() -> None: + """SessionCreatedEvent is wrapped in GlobalEvent envelope.""" + state = _MockState(working_dir="/wrap/test") + event = SessionCreatedEvent.create(session=_make_session("wrap1")) + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + assert "directory" in wrapped + assert "project" in wrapped + assert "payload" in wrapped + assert wrapped["directory"] == state.working_dir + assert "workspace" not in wrapped + assert wrapped["payload"]["type"] == "session.created" + + +@pytest.mark.anyio +async def test_session_status_is_wrapped() -> None: + """SessionStatusEvent is wrapped in GlobalEvent envelope.""" + state = _MockState(working_dir="/wrap/test2") + event = SessionStatusEvent.create(session_id="wrap2", status_type="idle") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + assert "directory" in wrapped + assert "payload" in wrapped + assert wrapped["directory"] == state.working_dir + assert "workspace" not in wrapped + + +@pytest.mark.anyio +async def test_session_compacted_is_wrapped() -> None: + """SessionCompactedEvent is wrapped in GlobalEvent envelope.""" + state = _MockState(working_dir="/wrap/compacted") + event = SessionCompactedEvent.create(session_id="wrap3") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + assert "directory" in wrapped + assert "payload" in wrapped + assert wrapped["directory"] == state.working_dir + assert "workspace" not in wrapped + + +@pytest.mark.anyio +async def test_tui_session_select_is_wrapped() -> None: + """TuiSessionSelectEvent is wrapped in GlobalEvent envelope.""" + state = _MockState(working_dir="/wrap/select") + event = TuiSessionSelectEvent.create(session_id="wrap4") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + assert "directory" in wrapped + assert "payload" in wrapped + assert wrapped["directory"] == state.working_dir + assert "workspace" not in wrapped + + +# ============================================================================= +# 5. sessionId appears at top level of payload for ALL event types +# ============================================================================= + + +@pytest.mark.anyio +async def test_session_id_in_payload_session_status() -> None: + """SessionStatusEvent payload has sessionId at top level.""" + event = SessionStatusEvent.create(session_id="top1", status_type="busy") + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert data["sessionId"] == "top1" + + +@pytest.mark.anyio +async def test_session_id_in_payload_part_delta() -> None: + """PartDeltaEvent payload has sessionId at top level (critical for streaming).""" + event = PartDeltaEvent.create( + session_id="top2", message_id="m1", part_id="p1", delta="streaming text" + ) + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert data["sessionId"] == "top2" + + +@pytest.mark.anyio +async def test_session_id_in_payload_session_created() -> None: + """SessionCreatedEvent payload has sessionId at top level.""" + event = SessionCreatedEvent.create(session=_make_session("top3")) + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert data["sessionId"] == "top3" + + +@pytest.mark.anyio +async def test_session_id_in_payload_message_updated() -> None: + """MessageUpdatedEvent payload has sessionId at top level.""" + msg = UserMessage(id="m1", session_id="top4", time=TimeCreated(created=0)) + event = MessageUpdatedEvent.create(message=msg) + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert data["sessionId"] == "top4" + + +@pytest.mark.anyio +async def test_session_id_in_payload_command_executed() -> None: + """CommandExecutedEvent payload has sessionId at top level.""" + event = CommandExecutedEvent.create( + name="test", session_id="top5", arguments="", message_id="m1" + ) + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert data["sessionId"] == "top5" + + +@pytest.mark.anyio +async def test_session_id_in_payload_permission_events() -> None: + """PermissionRequestEvent and PermissionResolvedEvent payloads have sessionId.""" + req = PermissionRequestEvent.create( + session_id="top6", + permission_id="p1", + tool_name="bash", + args_preview="ls", + message="Allow?", + ) + resolved = PermissionResolvedEvent.create( + session_id="top6", + request_id="p1", + reply="once", + ) + for event in (req, resolved): + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert data["sessionId"] == "top6" + + +@pytest.mark.anyio +async def test_session_id_absent_for_server_events() -> None: + """ServerConnectedEvent and ServerHeartbeatEvent have no sessionId.""" + for event in (ServerConnectedEvent(), ServerHeartbeatEvent()): + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert "sessionId" not in data + + +@pytest.mark.anyio +async def test_session_id_in_global_event_payload_part_delta() -> None: + """PartDeltaEvent in GlobalEvent envelope has sessionId inside payload.""" + state = _MockState() + event = PartDeltaEvent.create(session_id="env1", message_id="m1", part_id="p1", delta="x") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + wrapped = events[1] + assert wrapped["payload"]["sessionId"] == "env1" + + +# ============================================================================= +# 6. UserMessage events use nested model.variant format +# ============================================================================= + + +def test_user_message_variant_nested_in_model() -> None: + """UserMessage with model.variant serializes variant inside model object.""" + msg = UserMessage( + id="m1", + session_id="s1", + time=TimeCreated(created=0), + model=ModelRef(provider_id="openai", model_id="gpt-4o", variant="high"), + ) + data = msg.model_dump(by_alias=True, exclude_none=True) + assert "model" in data + assert isinstance(data["model"], dict) + assert data["model"]["variant"] == "high" + # variant should NOT appear at top level + assert "variant" not in data + + +def test_user_message_variant_not_at_top_level() -> None: + """UserMessage serialization does NOT put variant at top level.""" + msg = UserMessage( + id="m2", + session_id="s2", + time=TimeCreated(created=0), + model=ModelRef(variant="medium"), + ) + data = msg.model_dump(by_alias=True, exclude_none=True) + assert "variant" not in data + assert data["model"]["variant"] == "medium" + + +def test_user_message_variant_migration_from_top_level() -> None: + """UserMessage validator migrates top-level variant into model.variant.""" + # Simulate old client sending variant at top level + msg = UserMessage.model_validate({ + "id": "m3", + "session_id": "s3", + "time": {"created": 0}, + "variant": "low", + }) + data = msg.model_dump(by_alias=True, exclude_none=True) + # variant should now be inside model + assert "variant" not in data + assert data["model"]["variant"] == "low" + + +def test_user_message_no_variant_no_model() -> None: + """UserMessage without variant or model has neither in output.""" + msg = UserMessage( + id="m4", + session_id="s4", + time=TimeCreated(created=0), + ) + data = msg.model_dump(by_alias=True, exclude_none=True) + assert "variant" not in data + assert "model" not in data + + +def test_message_updated_event_carries_nested_variant() -> None: + """MessageUpdatedEvent with UserMessage preserves model.variant nesting.""" + msg = UserMessage( + id="m5", + session_id="s5", + time=TimeCreated(created=0), + model=ModelRef(variant="max"), + ) + event = MessageUpdatedEvent.create(message=msg) + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + # sessionId at top level + assert data["sessionId"] == "s5" + # variant nested inside model in properties + info = data["properties"]["info"] + assert "variant" not in info # no top-level variant + assert info["model"]["variant"] == "max" + + +# ============================================================================= +# 7. FileDiff events use v1.4.0+ schema (patch field, not before/after) +# ============================================================================= + + +def test_file_diff_has_patch_field() -> None: + """FileDiff model serializes with 'patch' field.""" + diff = FileDiff( + file="src/main.py", + patch="@@ -1,3 +1,4 @@\n-old line\n+new line\n+added line", + additions=2, + deletions=1, + status="modified", + ) + data = diff.model_dump(by_alias=True, exclude_none=True) + assert "patch" in data + assert data["patch"] == "@@ -1,3 +1,4 @@\n-old line\n+new line\n+added line" + + +def test_file_diff_no_before_after_fields() -> None: + """FileDiff serialization does NOT contain 'before' or 'after' fields.""" + diff = FileDiff( + file="src/util.py", + patch="some patch text", + additions=1, + deletions=0, + status="added", + ) + data = diff.model_dump(by_alias=True, exclude_none=True) + assert "before" not in data + assert "after" not in data + + +def test_file_diff_patch_none_excluded() -> None: + """FileDiff with patch=None excludes patch from output (exclude_none).""" + diff = FileDiff( + file="README.md", + additions=0, + deletions=0, + ) + data = diff.model_dump(by_alias=True, exclude_none=True) + assert "patch" not in data # None excluded + assert "before" not in data + assert "after" not in data + + +def test_session_diff_event_serializes_patch_not_before_after() -> None: + """SessionDiffEvent wraps FileDiff objects that have 'patch' not 'before'/'after'.""" + diff = FileDiff( + file="app.py", + patch="@@ -1 +1 @@\n-old\n+new", + additions=1, + deletions=1, + status="modified", + ) + event = SessionDiffEvent.create(session_id="diff1", diff=[diff]) + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + + diff_entries = data["properties"]["diff"] + assert len(diff_entries) == 1 + entry = diff_entries[0] + assert "patch" in entry + assert entry["patch"] == "@@ -1 +1 @@\n-old\n+new" + assert "before" not in entry + assert "after" not in entry + + +def test_session_diff_in_global_event_has_patch() -> None: + """SessionDiffEvent in GlobalEvent envelope serializes FileDiff with patch.""" + factory = GlobalEventFactory(directory="/tmp", project="abc") + diff = FileDiff( + file="config.yml", + patch="@@ -1,2 +1,3 @@\n-key: old\n+key: new\n+key2: val", + additions=2, + deletions=1, + status="modified", + ) + event = SessionDiffEvent.create(session_id="envdiff1", diff=[diff]) + result = factory.wrap(event) + data = json.loads(result) + + payload = data["payload"] + entry = payload["properties"]["diff"][0] + assert "patch" in entry + assert "before" not in entry + assert "after" not in entry + + +# ============================================================================= +# 8. /event endpoint backward compatibility (raw events, no envelope) +# ============================================================================= + + +@pytest.mark.anyio +async def test_event_endpoint_raw_events_no_envelope() -> None: + """/event (wrap_payload=False) sends events without GlobalEvent wrapper.""" + state = _MockState() + event = SessionStatusEvent.create(session_id="bc1", status_type="busy") + events = await _collect_events(state, wrap_payload=False, events_to_send=[event]) + + for evt in events: + assert "directory" not in evt + assert "project" not in evt + assert "payload" not in evt + + +@pytest.mark.anyio +async def test_event_endpoint_session_id_at_top_level() -> None: + """/event endpoint puts sessionId at top level for session events.""" + state = _MockState() + event = SessionStatusEvent.create(session_id="bc2", status_type="idle") + events = await _collect_events(state, wrap_payload=False, events_to_send=[event]) + session_data = events[1] + assert session_data["sessionId"] == "bc2" + + +@pytest.mark.anyio +async def test_event_endpoint_part_delta_raw() -> None: + """/event sends PartDeltaEvent as raw JSON with sessionId at top level.""" + state = _MockState() + event = PartDeltaEvent.create( + session_id="bc3", message_id="m1", part_id="p1", delta="hello world" + ) + events = await _collect_events(state, wrap_payload=False, events_to_send=[event]) + delta_data = events[1] + assert delta_data["type"] == "message.part.delta" + assert delta_data["sessionId"] == "bc3" + assert delta_data["properties"]["delta"] == "hello world" + + +@pytest.mark.anyio +async def test_event_endpoint_message_updated_raw() -> None: + """/event sends MessageUpdatedEvent as raw JSON with sessionId at top level.""" + state = _MockState() + msg = UserMessage(id="m1", session_id="bc4", time=TimeCreated(created=0)) + event = MessageUpdatedEvent.create(message=msg) + events = await _collect_events(state, wrap_payload=False, events_to_send=[event]) + msg_data = events[1] + assert msg_data["type"] == "message.updated" + assert msg_data["sessionId"] == "bc4" + + +# ============================================================================= +# 9. PartDeltaEvent streaming — the most critical event for "no response" symptom +# ============================================================================= + + +@pytest.mark.anyio +async def test_part_delta_session_id_at_top_level() -> None: + """PartDeltaEvent has sessionId injected at top level (was missing = no response).""" + event = PartDeltaEvent.create( + session_id="delta1", message_id="m1", part_id="p1", delta="text chunk" + ) + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert data["sessionId"] == "delta1" + + +@pytest.mark.anyio +async def test_part_delta_session_id_in_global_event_payload() -> None: + """PartDeltaEvent wrapped in GlobalEvent has sessionId inside payload.""" + state = _MockState(working_dir="/delta/test") + event = PartDeltaEvent.create( + session_id="delta2", message_id="m1", part_id="p1", delta="streaming" + ) + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + assert "directory" in wrapped + assert wrapped["payload"]["sessionId"] == "delta2" + assert wrapped["payload"]["type"] == "message.part.delta" + + +@pytest.mark.anyio +async def test_part_delta_extract_session_id() -> None: + """_extract_session_id correctly extracts sessionId from PartDeltaEvent.""" + event = PartDeltaEvent.create(session_id="delta3", message_id="m1", part_id="p1", delta="x") + sid = _extract_session_id(event) + assert sid == "delta3" + + +@pytest.mark.anyio +async def test_part_delta_tui_filter_passes() -> None: + """PartDeltaEvent in GlobalEvent passes the TUI routing filter.""" + wd = "/delta/filter" + state = _MockState(working_dir=wd) + event = PartDeltaEvent.create(session_id="delta4", message_id="m1", part_id="p1", delta="x") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + ge = GlobalEvent( + directory=wrapped["directory"], + project=wrapped.get("project"), + workspace=wrapped.get("workspace"), + payload=wrapped["payload"], + ) + passes, reason = tui_event_filter(ge, wd) + assert passes, f"PartDeltaEvent should pass filter, reason={reason}" + + +@pytest.mark.anyio +async def test_part_delta_multiple_deltas_stream() -> None: + """Sequence of PartDeltaEvents (simulating streaming) all have correct sessionId.""" + state = _MockState(working_dir="/delta/stream") + deltas = [ + PartDeltaEvent.create( + session_id="stream1", message_id="m1", part_id="p1", delta=f"chunk{i}" + ) + for i in range(5) + ] + events = await _collect_events(state, wrap_payload=True, events_to_send=deltas) + + # First event is payload-wrapped server.connected, then 5 wrapped deltas + for i, wrapped in enumerate(events[1:], start=0): + assert wrapped["payload"]["sessionId"] == "stream1" + assert wrapped["payload"]["properties"]["delta"] == f"chunk{i}" + + +@pytest.mark.anyio +async def test_part_delta_properties_fields() -> None: + """PartDeltaEvent payload has correct properties fields.""" + event = PartDeltaEvent.create( + session_id="delta5", + message_id="msg_abc", + part_id="part_xyz", + delta="hello text", + field="text", + ) + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert data["type"] == "message.part.delta" + assert data["sessionId"] == "delta5" + props = data["properties"] + assert props["sessionID"] == "delta5" # camelCase alias + assert props["messageID"] == "msg_abc" + assert props["partID"] == "part_xyz" + assert props["field"] == "text" + assert props["delta"] == "hello text" + + +# ============================================================================= +# Cross-cutting: mixed event sequences with correct wrapping +# ============================================================================= + + +@pytest.mark.anyio +async def test_mixed_events_correct_wrapping_sequence() -> None: + """Sequence of wrapped events all have correct format.""" + wd = "/mixed/test" + state = _MockState(working_dir=wd) + + events_to_send: list[Event] = [ + SessionStatusEvent.create(session_id="mix1", status_type="busy"), + ServerHeartbeatEvent(), + PartDeltaEvent.create(session_id="mix2", message_id="m1", part_id="p1", delta="chunk"), + ServerHeartbeatEvent(), + SessionCompactedEvent.create(session_id="mix3"), + ] + events = await _collect_events(state, wrap_payload=True, events_to_send=events_to_send) + + # [0] server.connected — payload wrapped, no routing metadata + assert events[0]["payload"]["type"] == "server.connected" + assert "directory" not in events[0] + + # [1] session.status — wrapped + assert "payload" in events[1] + assert events[1]["payload"]["type"] == "session.status" + assert events[1]["payload"]["sessionId"] == "mix1" + assert events[1]["directory"] == state.working_dir + assert "workspace" not in events[1] + + # [2] server.heartbeat — payload wrapped, no routing metadata + assert events[2]["payload"]["type"] == "server.heartbeat" + assert "directory" not in events[2] + + # [3] part.delta — wrapped (CRITICAL) + assert "payload" in events[3] + assert events[3]["payload"]["type"] == "message.part.delta" + assert events[3]["payload"]["sessionId"] == "mix2" + assert events[3]["directory"] == state.working_dir + assert "workspace" not in events[3] + + # [4] server.heartbeat — payload wrapped, no routing metadata + assert events[4]["payload"]["type"] == "server.heartbeat" + assert "directory" not in events[4] + + # [5] session.compacted — wrapped + assert "payload" in events[5] + assert events[5]["payload"]["type"] == "session.compacted" + assert events[5]["payload"]["sessionId"] == "mix3" + assert events[5]["directory"] == state.working_dir + assert "workspace" not in events[5] + + +@pytest.mark.anyio +async def test_workspace_mode_routing_falls_back_to_directory() -> None: + """Wrapped events fall back to directory routing when workspace is omitted.""" + wd = "/workspace/mode/test" + state = _MockState(working_dir=wd) + event = PartDeltaEvent.create(session_id="ws1", message_id="m1", part_id="p1", delta="x") + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + wrapped = events[1] + ge = GlobalEvent( + directory=wrapped["directory"], + project=wrapped.get("project"), + workspace=wrapped.get("workspace"), + payload=wrapped["payload"], + ) + passes, reason = tui_event_filter( + ge, + wd, + ) + assert passes, f"Directory-routed event should pass, got reason={reason}" + + +# ============================================================================= +# Cross-cutting: sessionId consistency between top level and properties +# ============================================================================= + + +def test_session_id_consistency_session_status() -> None: + """SessionId at top level matches sessionID inside properties (alias).""" + event = SessionStatusEvent.create(session_id="consist1", status_type="busy") + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + top_level_sid = data["sessionId"] + props_sid = data["properties"]["sessionID"] # camelCase alias + assert top_level_sid == props_sid == "consist1" + + +def test_session_id_consistency_part_delta() -> None: + """PartDeltaEvent: sessionId at top level matches sessionID in properties.""" + event = PartDeltaEvent.create(session_id="consist2", message_id="m1", part_id="p1", delta="x") + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + top_level_sid = data["sessionId"] + props_sid = data["properties"]["sessionID"] + assert top_level_sid == props_sid == "consist2" + + +def test_session_id_consistency_command_executed() -> None: + """CommandExecutedEvent: sessionId at top level matches sessionID in properties.""" + event = CommandExecutedEvent.create( + name="test", session_id="consist3", arguments="", message_id="m1" + ) + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + top_level_sid = data["sessionId"] + props_sid = data["properties"]["sessionID"] + assert top_level_sid == props_sid == "consist3" + + +# ============================================================================= +# Exhaustive: all handled event types produce sessionId at top level +# ============================================================================= + + +# Re-use the same event factories from test_global_event.py +_ALL_HANDLED_EVENTS_WITH_SID: list[tuple[str, Event]] = [ + ("session.deleted", SessionDeletedEvent.create(session_id="ex1")), + ("session.status", SessionStatusEvent.create(session_id="ex2", status_type="busy")), + ("session.idle", SessionCompactedEvent.create(session_id="ex3")), + ("session.compacted", SessionCompactedEvent.create(session_id="ex4")), + ("message.removed", SessionDeletedEvent.create(session_id="ex5")), + ( + "message.part.delta", + PartDeltaEvent.create(session_id="ex6", message_id="m1", part_id="p1", delta="x"), + ), + ( + "message.updated", + MessageUpdatedEvent.create( + message=UserMessage(id="m1", session_id="ex7", time=TimeCreated(created=0)) + ), + ), + ("session.created", SessionCreatedEvent.create(session=_make_session("ex8"))), + ("session.diff", SessionDiffEvent.create(session_id="ex9", diff=[])), + ( + "command.executed", + CommandExecutedEvent.create(name="test", session_id="ex10", arguments="", message_id="m1"), + ), + ("tui.session.select", TuiSessionSelectEvent.create(session_id="ex11")), +] + + +@pytest.mark.parametrize( + ("event_type_name", "event"), + [(name, evt) for name, evt in _ALL_HANDLED_EVENTS_WITH_SID], + ids=[name for name, _ in _ALL_HANDLED_EVENTS_WITH_SID], +) +def test_all_session_events_have_session_id_at_top_level( + event_type_name: str, + event: Event, +) -> None: + """All session-scoped events produce sessionId at top level in serialized output.""" + result = _serialize_event(event, wrap_payload=False) + data = json.loads(result) + assert "sessionId" in data, f"{event_type_name} missing sessionId at top level" + assert data["sessionId"] is not None, f"{event_type_name} has null sessionId" + + +# ============================================================================= +# GlobalEventFactory.wrap() produces correct envelope structure +# ============================================================================= + + +def test_factory_wrap_envelope_structure() -> None: + """GlobalEventFactory.wrap() produces {directory, project, payload} structure.""" + factory = GlobalEventFactory(directory="/factory/test", project="proj123") + event = PartDeltaEvent.create(session_id="fac1", message_id="m1", part_id="p1", delta="x") + result = factory.wrap(event) + data = json.loads(result) + + # Top-level keys + assert set(data.keys()) >= {"directory", "project", "payload"} + assert data["directory"] == "/factory/test" + assert data["project"] == "proj123" + + # Payload structure + payload = data["payload"] + assert payload["type"] == "message.part.delta" + assert payload["sessionId"] == "fac1" + + +def test_factory_wrap_session_created_has_session_id() -> None: + """Factory.wrap(SessionCreatedEvent) has sessionId in payload.""" + factory = GlobalEventFactory(directory="/tmp", project="abc") + event = SessionCreatedEvent.create(session=_make_session("fac2")) + result = factory.wrap(event) + data = json.loads(result) + assert data["payload"]["sessionId"] == "fac2" + + +def test_factory_wrap_message_updated_has_session_id() -> None: + """Factory.wrap(MessageUpdatedEvent) has sessionId in payload.""" + factory = GlobalEventFactory(directory="/tmp", project="abc") + msg = UserMessage(id="m1", session_id="fac3", time=TimeCreated(created=0)) + event = MessageUpdatedEvent.create(message=msg) + result = factory.wrap(event) + data = json.loads(result) + assert data["payload"]["sessionId"] == "fac3" + + +# ============================================================================= +# Unicode preservation through the full SSE pipeline +# ============================================================================= + + +@pytest.mark.anyio +async def test_unicode_part_delta_in_global_event() -> None: + r"""PartDeltaEvent with CJK text preserves Unicode in GlobalEvent envelope.""" + state = _MockState(working_dir="/unicode/test") + event = PartDeltaEvent.create( + session_id="unicode1", message_id="m1", part_id="p1", delta="你好世界 🌍" + ) + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + wrapped = events[1] + assert wrapped["payload"]["properties"]["delta"] == "你好世界 🌍" + assert wrapped["payload"]["sessionId"] == "unicode1" + + +def test_unicode_file_diff_patch_preserved() -> None: + r"""FileDiff with Unicode in patch preserves characters (not \uXXXX).""" + diff = FileDiff( + file="中文文件.py", + patch="@@ -1 +1 @@\n-旧代码\n+新代码 🔥", + additions=1, + deletions=1, + status="modified", + ) + event = SessionDiffEvent.create(session_id="udiff1", diff=[diff]) + result = _serialize_event(event, wrap_payload=False) + assert "中文文件" in result + assert "新代码" in result + assert "\\u" not in result + + +# ============================================================================= +# Bounded subscriber queue tests — verifying maxsize=100 does not affect payloads +# ============================================================================= + + +@pytest.mark.anyio +async def test_subscriber_queue_is_bounded() -> None: + """Subscriber queues created by _event_generator are bounded at maxsize=100.""" + state = _MockState() + gen = _event_generator(state, wrap_payload=False) + await gen.__anext__() # consume connected event + + queue = state.event_subscribers[-1] + assert queue.maxsize == 100 + + await gen.aclose() + + +@pytest.mark.anyio +async def test_bounded_queue_event_endpoint_payload_compatibility() -> None: + """Bounded queue does not change /event raw payload format.""" + state = _MockState() + event = PartDeltaEvent.create( + session_id="bq1", message_id="m1", part_id="p1", delta="bounded queue test" + ) + events = await _collect_events(state, wrap_payload=False, events_to_send=[event]) + + # Verify queue is bounded + queue = state.event_subscribers[-1] + assert queue.maxsize == 100 + + # Verify payload compatibility — raw /event format unchanged + connected = events[0] + assert connected["type"] == "server.connected" + assert "directory" not in connected + assert "payload" not in connected + + delta_data = events[1] + assert delta_data["type"] == "message.part.delta" + assert delta_data["sessionId"] == "bq1" + assert delta_data["properties"]["delta"] == "bounded queue test" + assert "directory" not in delta_data + assert "payload" not in delta_data + + +@pytest.mark.anyio +async def test_bounded_queue_global_event_endpoint_payload_compatibility() -> None: + """Bounded queue does not change /global/event wrapped payload format.""" + wd = "/bounded/queue/test" + state = _MockState(working_dir=wd) + event = PartDeltaEvent.create( + session_id="bq2", message_id="m1", part_id="p1", delta="global bounded test" + ) + events = await _collect_events(state, wrap_payload=True, events_to_send=[event]) + + # Verify queue is bounded + queue = state.event_subscribers[-1] + assert queue.maxsize == 100 + + # Verify payload compatibility — /global/event envelope format unchanged + connected = events[0] + assert connected["payload"]["type"] == "server.connected" + assert "directory" not in connected + + wrapped = events[1] + assert wrapped["directory"] == wd + assert "project" in wrapped + assert "workspace" not in wrapped + assert wrapped["payload"]["type"] == "message.part.delta" + assert wrapped["payload"]["sessionId"] == "bq2" + assert wrapped["payload"]["properties"]["delta"] == "global bounded test" + + +@pytest.mark.anyio +async def test_bounded_queue_backpressure_drops_without_affecting_others() -> None: + """When one subscriber's bounded queue is full, events are dropped for that + subscriber but other subscribers with available capacity still receive them. + """ + state = _MockState() + + gen1 = _event_generator(state, wrap_payload=True) + gen2 = _event_generator(state, wrap_payload=True) + await gen1.__anext__() + await gen2.__anext__() + + queue1 = state.event_subscribers[0] + queue2 = state.event_subscribers[1] + assert queue1.maxsize == 100 + assert queue2.maxsize == 100 + + # Fill queue1 to capacity + for _ in range(100): + queue1.put_nowait(ServerHeartbeatEvent()) + + # Broadcast an event — queue1 is full (event dropped by put_nowait), + # queue2 receives it via direct put + event = SessionStatusEvent.create(session_id="bp1", status_type="busy") + # Simulate broadcast_event behavior: put_nowait to each queue + for q in list(state.event_subscribers): + try: + q.put_nowait(event) + except asyncio.QueueFull: + pass # drop-on-backpressure + + # queue1 still has exactly 100 items (the heartbeats), event was dropped + assert queue1.qsize() == 100 + + # queue2 received the event + item2 = await gen2.__anext__() + data2 = json.loads(item2["data"]) + assert data2["payload"]["type"] == "session.status" + assert data2["payload"]["sessionId"] == "bp1" + + await gen1.aclose() + await gen2.aclose() diff --git a/tests/servers/opencode_server/test_tui_routing.py b/tests/servers/opencode_server/test_tui_routing.py new file mode 100644 index 000000000..c13db2b03 --- /dev/null +++ b/tests/servers/opencode_server/test_tui_routing.py @@ -0,0 +1,466 @@ +"""Tests for TUI event routing filter and /global/routing-check endpoint. + +Covers all 4 rules of the OpenCode TUI routing filter: +1. Sync events always dropped +2. Global directory always passes (except sync) +3. Workspace filtering (if active) +4. Directory must match exactly (string comparison, no normalization) + +Plus edge cases and the HTTP endpoint integration. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +import pytest + +from agentpool_server.opencode_server.models import GlobalEvent +from agentpool_server.opencode_server.routes.routing import ( + RoutingCheckResponse, + tui_event_filter, +) + + +if TYPE_CHECKING: + from httpx import AsyncClient + + from agentpool_server.opencode_server.state import ServerState + + +# ============================================================================= +# Helper +# ============================================================================= + + +def _make_event( + directory: str = "/project", + workspace: str | None = None, + payload_type: str | None = None, +) -> GlobalEvent: + """Create a GlobalEvent for testing.""" + payload: dict[str, Any] = {} + if payload_type is not None: + payload["type"] = payload_type + return GlobalEvent(directory=directory, workspace=workspace, payload=payload) + + +# ============================================================================= +# Rule 1: sync events always dropped +# ============================================================================= + + +def test_sync_event_dropped() -> None: + """Sync events are always dropped regardless of other conditions.""" + event = _make_event(directory="global", payload_type="sync") + passed, reason = tui_event_filter(event, "/project") + assert passed is False + assert reason == "sync_dropped" + + +def test_sync_event_dropped_even_with_matching_directory() -> None: + """Sync event with matching directory is still dropped.""" + event = _make_event(directory="/project", payload_type="sync") + passed, reason = tui_event_filter(event, "/project") + assert passed is False + assert reason == "sync_dropped" + + +def test_sync_event_dropped_with_matching_workspace() -> None: + """Sync event with matching workspace is still dropped.""" + event = _make_event(directory="/project", workspace="ws1", payload_type="sync") + passed, reason = tui_event_filter(event, "/project", current_workspace="ws1") + assert passed is False + assert reason == "sync_dropped" + + +def test_non_sync_payload_type_not_dropped() -> None: + """Non-sync payload types are not affected by rule 1.""" + event = _make_event(directory="/project", payload_type="session.status") + passed, reason = tui_event_filter(event, "/project") + assert passed is True + assert reason == "directory_match" + + +def test_no_payload_type_not_dropped() -> None: + """Event without a payload type is not affected by rule 1.""" + event = _make_event(directory="/project") + passed, reason = tui_event_filter(event, "/project") + assert passed is True + assert reason == "directory_match" + + +# ============================================================================= +# Rule 2: global directory always passes (except sync) +# ============================================================================= + + +def test_global_directory_passes() -> None: + """Event with directory='global' always passes (non-sync).""" + event = _make_event(directory="global") + passed, reason = tui_event_filter(event, "/project") + assert passed is True + assert reason == "global_directory" + + +def test_global_directory_passes_regardless_of_project_directory() -> None: + """Global directory passes even when project_directory differs.""" + event = _make_event(directory="global") + passed, reason = tui_event_filter(event, "/completely/different") + assert passed is True + assert reason == "global_directory" + + +def test_global_directory_passes_with_workspace_active() -> None: + """Global directory passes even when workspace filtering is active.""" + event = _make_event(directory="global", workspace="other-ws") + passed, reason = tui_event_filter(event, "/project", current_workspace="ws1") + assert passed is True + assert reason == "global_directory" + + +def test_global_directory_with_sync_still_dropped() -> None: + """Sync event with global directory is dropped (rule 1 takes precedence).""" + event = _make_event(directory="global", payload_type="sync") + passed, reason = tui_event_filter(event, "/project") + assert passed is False + assert reason == "sync_dropped" + + +# ============================================================================= +# Rule 3: workspace filtering (if active) +# ============================================================================= + + +def test_workspace_match_passes() -> None: + """Event workspace matches current_workspace → passes.""" + event = _make_event(directory="/project", workspace="ws1") + passed, reason = tui_event_filter(event, "/project", current_workspace="ws1") + assert passed is True + assert reason == "workspace_match" + + +def test_workspace_mismatch_fails() -> None: + """Event workspace doesn't match current_workspace → fails.""" + event = _make_event(directory="/project", workspace="ws1") + passed, reason = tui_event_filter(event, "/project", current_workspace="ws2") + assert passed is False + assert reason == "workspace_mismatch" + + +def test_workspace_match_passes_even_with_wrong_directory() -> None: + """Workspace match takes priority over directory mismatch.""" + event = _make_event(directory="/other", workspace="ws1") + passed, reason = tui_event_filter(event, "/project", current_workspace="ws1") + assert passed is True + assert reason == "workspace_match" + + +def test_workspace_mismatch_even_with_matching_directory() -> None: + """Workspace mismatch fails even when directory matches.""" + event = _make_event(directory="/project", workspace="ws1") + passed, reason = tui_event_filter(event, "/project", current_workspace="ws2") + assert passed is False + assert reason == "workspace_mismatch" + + +def test_workspace_none_event_with_active_workspace_fails() -> None: + """Event with workspace=None fails when current_workspace is set.""" + event = _make_event(directory="/project", workspace=None) + passed, reason = tui_event_filter(event, "/project", current_workspace="ws1") + assert passed is False + assert reason == "workspace_mismatch" + + +def test_empty_string_workspace_vs_none_mismatch() -> None: + """Empty string workspace != None current_workspace.""" + event = _make_event(directory="/project", workspace="") + passed, reason = tui_event_filter(event, "/project", current_workspace=None) + assert passed is True + assert reason == "directory_match" + + +def test_workspace_none_current_none_falls_through() -> None: + """Both workspace and current_workspace None falls through to directory check.""" + event = _make_event(directory="/project", workspace=None) + passed, reason = tui_event_filter(event, "/project", current_workspace=None) + assert passed is True + assert reason == "directory_match" + + +# ============================================================================= +# Rule 4: directory must match exactly +# ============================================================================= + + +def test_directory_match_passes() -> None: + """Exact directory match → passes.""" + event = _make_event(directory="/project") + passed, reason = tui_event_filter(event, "/project") + assert passed is True + assert reason == "directory_match" + + +def test_directory_mismatch_fails() -> None: + """Directory mismatch → fails.""" + event = _make_event(directory="/other") + passed, reason = tui_event_filter(event, "/project") + assert passed is False + assert reason == "directory_mismatch" + + +def test_directory_trailing_slash_not_normalized() -> None: + """Trailing slash vs no trailing slash → NOT equal (no normalization).""" + event = _make_event(directory="/project/") + passed, reason = tui_event_filter(event, "/project") + assert passed is False + assert reason == "directory_mismatch" + + +def test_directory_no_trailing_vs_trailing_not_normalized() -> None: + """No trailing slash vs trailing slash → NOT equal (no normalization).""" + event = _make_event(directory="/project") + passed, reason = tui_event_filter(event, "/project/") + assert passed is False + assert reason == "directory_mismatch" + + +def test_directory_case_sensitive() -> None: + """Directory comparison is case-sensitive.""" + event = _make_event(directory="/Project") + passed, reason = tui_event_filter(event, "/project") + assert passed is False + assert reason == "directory_mismatch" + + +def test_directory_empty_string_mismatch() -> None: + """Empty string directory doesn't match a real path.""" + event = _make_event(directory="") + passed, reason = tui_event_filter(event, "/project") + assert passed is False + assert reason == "directory_mismatch" + + +def test_directory_match_with_spaces() -> None: + """Directory with spaces matches exactly.""" + event = _make_event(directory="/path with spaces/project") + passed, reason = tui_event_filter(event, "/path with spaces/project") + assert passed is True + assert reason == "directory_match" + + +# ============================================================================= +# Rule priority / interaction edge cases +# ============================================================================= + + +def test_rule_priority_sync_over_global() -> None: + """Rule 1 (sync) takes precedence over rule 2 (global directory).""" + event = _make_event(directory="global", payload_type="sync") + passed, reason = tui_event_filter(event, "/project") + assert passed is False + assert reason == "sync_dropped" + + +def test_rule_priority_global_over_workspace() -> None: + """Rule 2 (global directory) takes precedence over rule 3 (workspace).""" + event = _make_event(directory="global", workspace="wrong-ws") + passed, reason = tui_event_filter(event, "/project", current_workspace="ws1") + assert passed is True + assert reason == "global_directory" + + +def test_rule_priority_workspace_over_directory() -> None: + """Rule 3 (workspace) takes precedence over rule 4 (directory).""" + event = _make_event(directory="/wrong", workspace="ws1") + passed, reason = tui_event_filter(event, "/project", current_workspace="ws1") + assert passed is True + assert reason == "workspace_match" + + +def test_full_path_cjk_directory() -> None: + """CJK characters in directory match exactly.""" + event = _make_event(directory="/项目/代码") + passed, reason = tui_event_filter(event, "/项目/代码") + assert passed is True + assert reason == "directory_match" + + +# ============================================================================= +# RoutingCheckResponse model tests +# ============================================================================= + + +def test_routing_check_response_serialization() -> None: + """RoutingCheckResponse serializes correctly with camelCase aliases.""" + response = RoutingCheckResponse(would_pass=True, reason="directory_match") + dumped = response.model_dump(by_alias=True, exclude_none=True) + assert dumped["wouldPass"] is True + assert dumped["reason"] == "directory_match" + + +def test_routing_check_response_all_reasons() -> None: + """All valid reason values produce valid RoutingCheckResponse.""" + valid_reasons = [ + "sync_dropped", + "global_directory", + "workspace_match", + "workspace_mismatch", + "directory_match", + "directory_mismatch", + ] + for reason in valid_reasons: + response = RoutingCheckResponse(would_pass=True, reason=reason) + assert response.reason == reason + + +# ============================================================================= +# /global/routing-check HTTP endpoint tests +# ============================================================================= + + +@pytest.mark.anyio +async def test_routing_check_directory_match(async_client: AsyncClient) -> None: + """Directory match returns would_pass=True, reason=directory_match.""" + response = await async_client.get( + "/global/routing-check", + params={"directory": "/my/project", "project_directory": "/my/project"}, + ) + assert response.status_code == 200 + data = response.json() + assert data["wouldPass"] is True + assert data["reason"] == "directory_match" + + +@pytest.mark.anyio +async def test_routing_check_directory_mismatch(async_client: AsyncClient) -> None: + """Directory mismatch returns would_pass=False, reason=directory_mismatch.""" + response = await async_client.get( + "/global/routing-check", + params={"directory": "/different/path"}, + ) + assert response.status_code == 200 + data = response.json() + assert data["wouldPass"] is False + assert data["reason"] == "directory_mismatch" + + +@pytest.mark.anyio +async def test_routing_check_global_directory(async_client: AsyncClient) -> None: + """Global directory returns would_pass=True, reason=global_directory.""" + response = await async_client.get( + "/global/routing-check", + params={"directory": "global"}, + ) + assert response.status_code == 200 + data = response.json() + assert data["wouldPass"] is True + assert data["reason"] == "global_directory" + + +@pytest.mark.anyio +async def test_routing_check_sync_dropped(async_client: AsyncClient) -> None: + """Sync event with global directory is dropped by rule 1. + + Note: the endpoint constructs a GlobalEvent with an empty payload {}, + so we can't directly test sync_dropped via the endpoint (no payload.type + parameter is exposed). This test verifies the general behavior. + """ + # The endpoint doesn't support payload_type param, so we test the + # pure function directly for sync_dropped and verify the endpoint + # works for the other cases + event = _make_event(directory="global", payload_type="sync") + passed, reason = tui_event_filter(event, "/project") + assert passed is False + assert reason == "sync_dropped" + + +@pytest.mark.anyio +async def test_routing_check_workspace_match(async_client: AsyncClient) -> None: + """Workspace match with current_workspace set returns would_pass=True.""" + response = await async_client.get( + "/global/routing-check", + params={ + "directory": "/project", + "workspace": "ws1", + "current_workspace": "ws1", + }, + ) + assert response.status_code == 200 + data = response.json() + assert data["wouldPass"] is True + assert data["reason"] == "workspace_match" + + +@pytest.mark.anyio +async def test_routing_check_workspace_mismatch(async_client: AsyncClient) -> None: + """Workspace mismatch with current_workspace set returns would_pass=False.""" + response = await async_client.get( + "/global/routing-check", + params={ + "directory": "/project", + "workspace": "ws1", + "current_workspace": "ws2", + }, + ) + assert response.status_code == 200 + data = response.json() + assert data["wouldPass"] is False + assert data["reason"] == "workspace_mismatch" + + +@pytest.mark.anyio +async def test_routing_check_custom_project_directory(async_client: AsyncClient) -> None: + """Custom project_directory parameter overrides state.working_dir.""" + response = await async_client.get( + "/global/routing-check", + params={ + "directory": "/custom/project", + "project_directory": "/custom/project", + }, + ) + assert response.status_code == 200 + data = response.json() + assert data["wouldPass"] is True + assert data["reason"] == "directory_match" + + +@pytest.mark.anyio +async def test_routing_check_default_project_directory(async_client: AsyncClient) -> None: + """Without project_directory param, uses server's base_path (resolved). + + Since the server's base_path is a temp directory, we use a known + different directory to verify it fails (directory_mismatch), which + confirms the default is being used rather than matching any value. + """ + response = await async_client.get( + "/global/routing-check", + params={"directory": "/definitely/not/the/working/dir"}, + ) + assert response.status_code == 200 + data = response.json() + assert data["wouldPass"] is False + assert data["reason"] == "directory_mismatch" + + +@pytest.mark.anyio +async def test_routing_check_default_uses_base_path( + async_client: AsyncClient, + server_state: ServerState, +) -> None: + """Routing-check default is state.base_path, not raw working_dir. + + The endpoint should use the resolved canonical path (base_path) as + the default project directory, matching how get_event_factory() + already uses self.base_path for directory normalization. Verify by + sending directory=base_path without an explicit project_directory + override — the event must pass the directory-match rule. + """ + response = await async_client.get( + "/global/routing-check", + params={"directory": server_state.base_path}, + ) + assert response.status_code == 200 + data = response.json() + assert data["wouldPass"] is True + assert data["reason"] == "directory_match" diff --git a/tests/servers/opencode_server/test_workspace_routes.py b/tests/servers/opencode_server/test_workspace_routes.py new file mode 100644 index 000000000..cf81a0506 --- /dev/null +++ b/tests/servers/opencode_server/test_workspace_routes.py @@ -0,0 +1,82 @@ +"""Tests for experimental workspace compatibility routes.""" + +from __future__ import annotations + +from pathlib import Path +from typing import TYPE_CHECKING + +import pytest + +from agentpool_storage.opencode_provider import helpers + + +if TYPE_CHECKING: + from httpx import AsyncClient + + from agentpool_server.opencode_server.state import ServerState + + +pytestmark = pytest.mark.asyncio + + +async def test_list_workspaces_returns_singleton_local_workspace( + async_client: AsyncClient, + server_state: ServerState, +) -> None: + """Workspace list returns the singleton local workspace expected by OpenCode.""" + response = await async_client.get("/experimental/workspace") + + assert response.status_code == 200 + data = response.json() + assert isinstance(data, list) + assert len(data) == 1 + + workspace = data[0] + expected_directory = server_state.base_path + expected_project_id = helpers.compute_project_id(expected_directory) + + assert workspace["id"] == f"wrk_{expected_project_id[:12]}" + assert workspace["type"] == "local" + assert workspace["name"] == Path(expected_directory).name + assert workspace["branch"] is None + assert workspace["directory"] == expected_directory + assert workspace["extra"] is None + assert workspace["projectID"] == expected_project_id + + +async def test_workspace_status_returns_array_with_matching_workspace_id( + async_client: AsyncClient, +) -> None: + """Workspace status returns an array compatible with TUI bootstrap mapping.""" + workspace_response = await async_client.get("/experimental/workspace") + status_response = await async_client.get("/experimental/workspace/status") + + workspace_data = workspace_response.json() + status_data = status_response.json() + + assert workspace_response.status_code == 200 + assert status_response.status_code == 200 + assert isinstance(status_data, list) + assert len(status_data) == 1 + assert status_data[0]["workspaceID"] == workspace_data[0]["id"] + assert status_data[0]["status"] == "connected" + assert status_data[0]["error"] is None + + +async def test_workspace_routes_accept_sdk_query_params( + async_client: AsyncClient, + server_state: ServerState, +) -> None: + """Workspace routes stay JSON-shaped when SDK sends directory/workspace params.""" + params = { + "directory": server_state.base_path, + "workspace": "wrk_local_override", + } + + list_response = await async_client.get("/experimental/workspace", params=params) + status_response = await async_client.get("/experimental/workspace/status", params=params) + + assert list_response.status_code == 200 + assert status_response.status_code == 200 + assert isinstance(list_response.json(), list) + assert isinstance(status_response.json(), list) diff --git a/tests/sessions/test_opencode_helpers.py b/tests/sessions/test_opencode_helpers.py new file mode 100644 index 000000000..b2065a359 --- /dev/null +++ b/tests/sessions/test_opencode_helpers.py @@ -0,0 +1,148 @@ +"""Tests for OpenCode storage provider helper functions. + +Covers: +- convert_user_content_to_parts() with various UserContent types +""" + +from __future__ import annotations + +from pydantic_ai import BinaryContent, ImageUrl, TextContent +import pytest + +from agentpool_storage.opencode_provider.helpers import convert_user_content_to_parts + + +@pytest.fixture +def message_id() -> str: + """Provide a stable message ID for tests.""" + return "msg_test123" + + +@pytest.fixture +def session_id() -> str: + """Provide a stable session ID for tests.""" + return "ses_test123" + + +def test_convert_str_content_returns_single_text_part( + message_id: str, + session_id: str, +) -> None: + """Simple string input produces exactly one TextPart.""" + result = convert_user_content_to_parts( + content="Hello, world!", + message_id=message_id, + session_id=session_id, + part_counter_start=0, + ) + + assert len(result) == 1 + part = result[0] + assert part.text == "Hello, world!" + assert part.message_id == message_id + assert part.session_id == session_id + assert part.id.startswith("prt_") + + +def test_convert_list_with_str_items_returns_text_parts( + message_id: str, + session_id: str, +) -> None: + """List of str items produces one TextPart per item.""" + content: list[str] = ["Hello", "World"] + result = convert_user_content_to_parts( + content=content, + message_id=message_id, + session_id=session_id, + part_counter_start=0, + ) + + assert len(result) == 2 + assert result[0].text == "Hello" + assert result[1].text == "World" + # Each part gets a unique ID + assert result[0].id != result[1].id + + +def test_convert_list_with_binary_content_skips_with_warning( + message_id: str, + session_id: str, +) -> None: + """BinaryContent items are skipped (with warning), no TextPart produced.""" + binary = BinaryContent(data=b"\x89PNG", media_type="image/png") + result = convert_user_content_to_parts( + content=[binary], + message_id=message_id, + session_id=session_id, + part_counter_start=0, + ) + + assert result == [] + + +def test_convert_list_with_file_url_skips_with_warning( + message_id: str, + session_id: str, +) -> None: + """FileUrl items (e.g. ImageUrl) are skipped (with warning), no TextPart produced.""" + url = ImageUrl(url="https://example.com/image.png") + result = convert_user_content_to_parts( + content=[url], + message_id=message_id, + session_id=session_id, + part_counter_start=0, + ) + + assert result == [] + + +def test_convert_empty_list_returns_empty_list( + message_id: str, + session_id: str, +) -> None: + """Empty list input returns empty list.""" + result = convert_user_content_to_parts( + content=[], + message_id=message_id, + session_id=session_id, + part_counter_start=0, + ) + + assert result == [] + + +def test_convert_mixed_content_produces_only_text_parts( + message_id: str, + session_id: str, +) -> None: + """Mixed str and binary content: only text items become TextParts.""" + binary = BinaryContent(data=b"\x00\x01", media_type="application/octet-stream") + url = ImageUrl(url="https://example.com/doc.pdf") + content = ["First text", binary, "Second text", url] + result = convert_user_content_to_parts( + content=content, + message_id=message_id, + session_id=session_id, + part_counter_start=0, + ) + + assert len(result) == 2 + assert result[0].text == "First text" + assert result[1].text == "Second text" + + +def test_convert_text_content_produces_text_part( + message_id: str, + session_id: str, +) -> None: + """TextContent items produce TextParts with their content.""" + tc = TextContent(content="Structured text") + result = convert_user_content_to_parts( + content=[tc], + message_id=message_id, + session_id=session_id, + part_counter_start=0, + ) + + assert len(result) == 1 + assert result[0].text == "Structured text" diff --git a/tests/sessions/test_opencode_integration.py b/tests/sessions/test_opencode_integration.py new file mode 100644 index 000000000..9b3ccbecb --- /dev/null +++ b/tests/sessions/test_opencode_integration.py @@ -0,0 +1,172 @@ +"""Integration tests for OpenCode storage persistence and dual-ID scheme. + +Covers: +- Project and session persistence across provider restart +- Accessibility of both project ID schemes (compute_project_id / generate_project_id) +""" + +from __future__ import annotations + +from pathlib import Path +import subprocess +import tempfile + +import pytest + +from agentpool.sessions.models import ProjectData, SessionData +from agentpool.utils.identifiers import ascending +from agentpool_config.storage import OpenCodeStorageConfig +from agentpool_storage.opencode_provider import OpenCodeStorageProvider +from agentpool_storage.opencode_provider.helpers import compute_project_id +from agentpool_storage.project_store import generate_project_id + + +@pytest.fixture +async def provider(): + """Create an OpenCode provider with temp directory.""" + with tempfile.TemporaryDirectory() as tmpdir: + config = OpenCodeStorageConfig(path=tmpdir) + prov = OpenCodeStorageProvider(config) + async with prov: + yield prov + + +def _init_git_repo(directory: str) -> None: + """Initialize a minimal git repo so compute_project_id returns a commit SHA.""" + subprocess.run(["git", "init"], cwd=directory, capture_output=True, check=True) + subprocess.run( + ["git", "config", "user.email", "test@test.com"], + cwd=directory, + capture_output=True, + check=True, + ) + subprocess.run( + ["git", "config", "user.name", "Test"], + cwd=directory, + capture_output=True, + check=True, + ) + # Create an initial commit so there's a root commit SHA + dummy = Path(directory) / "README.md" + dummy.write_text("init", encoding="utf-8") + subprocess.run(["git", "add", "."], cwd=directory, capture_output=True, check=True) + subprocess.run( + ["git", "commit", "-m", "init"], + cwd=directory, + capture_output=True, + check=True, + ) + + +async def test_project_session_persistence_across_restart( + provider: OpenCodeStorageProvider, +) -> None: + """Project and sessions survive closing and re-opening the provider.""" + base_path = provider.base_path + + # Create project data + worktree = str(Path(base_path) / "my_project") + Path(worktree).mkdir(parents=True, exist_ok=True) + _init_git_repo(worktree) + + project_id = generate_project_id(worktree) + project = ProjectData( + project_id=project_id, + worktree=worktree, + name="test-project", + vcs="git", + ) + await provider.save_project(project) + + # Create session data for this project + session_id = ascending("session") + session = SessionData( + session_id=session_id, + agent_name="test-agent", + project_id=project_id, + cwd=worktree, + ) + await provider.save_session(session) + + # Verify within the same provider instance + loaded_project = await provider.get_project(project_id) + assert loaded_project is not None + assert loaded_project.project_id == project_id + assert loaded_project.worktree == worktree + assert loaded_project.name == "test-project" + + loaded_sessions = await provider.list_session_ids() + assert session_id in loaded_sessions + + # Close the provider (simulate restart) + await provider.__aexit__(None, None, None) + + # Create a NEW provider instance with the same base path + config = OpenCodeStorageConfig(path=str(base_path)) + new_provider = OpenCodeStorageProvider(config) + async with new_provider: + # Verify project is recovered + recovered_project = await new_provider.get_project(project_id) + assert recovered_project is not None + assert recovered_project.project_id == project_id + assert recovered_project.worktree == worktree + assert recovered_project.name == "test-project" + + # Verify sessions persist + recovered_sessions = await new_provider.list_session_ids() + assert session_id in recovered_sessions + + # Verify get_project_by_worktree returns the project + by_worktree = await new_provider.get_project_by_worktree(worktree) + assert by_worktree is not None + assert by_worktree.project_id == project_id + + +async def test_both_project_ids_accessible() -> None: + """Both compute_project_id and generate_project_id return valid values. + + The two functions use different algorithms and serve different purposes: + - compute_project_id: git root commit SHA1 (OpenCode session layout) + - generate_project_id: path SHA1 (AgentPool project registry) + + This test documents that both are accessible and produce distinct values + for the same directory. + """ + with tempfile.TemporaryDirectory() as tmpdir: + project_dir = Path(tmpdir) / "dual_id_project" + project_dir.mkdir(parents=True, exist_ok=True) + _init_git_repo(str(project_dir)) + + opencode_id = compute_project_id(str(project_dir)) + agentpool_id = generate_project_id(str(project_dir)) + + # Both should return valid hex strings (or "global" for compute_project_id) + assert len(opencode_id) >= 1 + assert len(agentpool_id) == 40 # SHA1 hex digest + + # With a git repo, compute_project_id returns the root commit SHA + assert opencode_id != "global" + assert len(opencode_id) == 40 # commit SHA1 hex digest + + # They are different algorithms — expect different values + assert opencode_id != agentpool_id, ( + "compute_project_id (git root SHA) and generate_project_id (path SHA) " + "use different algorithms and should produce different values" + ) + + +async def test_compute_project_id_without_git() -> None: + """compute_project_id returns 'global' when not in a git repo.""" + with tempfile.TemporaryDirectory() as tmpdir: + # No git init — should return "global" + result = compute_project_id(tmpdir) + assert result == "global" + + +async def test_generate_project_id_deterministic() -> None: + """generate_project_id returns the same value for the same path.""" + with tempfile.TemporaryDirectory() as tmpdir: + id_1 = generate_project_id(tmpdir) + id_2 = generate_project_id(tmpdir) + assert id_1 == id_2 + assert len(id_1) == 40 diff --git a/tests/sessions/test_opencode_project_methods.py b/tests/sessions/test_opencode_project_methods.py new file mode 100644 index 000000000..efc3d48cc --- /dev/null +++ b/tests/sessions/test_opencode_project_methods.py @@ -0,0 +1,245 @@ +"""Tests for OpenCodeStorageProvider project storage methods. + +Covers all 7 project methods: +- save_project, get_project, get_project_by_worktree, get_project_by_name +- list_projects, delete_project, touch_project +""" + +from __future__ import annotations + +from datetime import UTC, datetime +from pathlib import Path +import tempfile + +import pytest + +from agentpool.sessions.models import ProjectData +from agentpool_config.storage import OpenCodeStorageConfig +from agentpool_storage.opencode_provider import OpenCodeStorageProvider + + +@pytest.fixture +async def provider(): + """Create an OpenCode provider with temp directory.""" + with tempfile.TemporaryDirectory() as tmpdir: + config = OpenCodeStorageConfig(path=tmpdir) + prov = OpenCodeStorageProvider(config) + async with prov: + yield prov + + +def _make_project( + *, + project_id: str = "test_project_001", + worktree: str = "/tmp/test-project", + name: str | None = "test-project", + vcs: str | None = "git", + config_path: str | None = None, + settings: dict | None = None, + last_active: datetime | None = None, +) -> ProjectData: + """Helper to create ProjectData instances for tests.""" + return ProjectData( + project_id=project_id, + worktree=worktree, + name=name, + vcs=vcs, + config_path=config_path, + settings=settings or {}, + last_active=last_active or datetime.now(UTC), + ) + + +async def test_save_and_get_project(provider: OpenCodeStorageProvider) -> None: + """Test saving a project and retrieving it by ID.""" + project = _make_project( + project_id="proj_001", + worktree="/home/user/project-a", + name="project-a", + vcs="git", + config_path="/home/user/project-a/.agentpool.yml", + settings={"model": "openai:gpt-4o"}, + ) + + await provider.save_project(project) + + result = await provider.get_project("proj_001") + + assert result is not None + assert result.project_id == "proj_001" + assert result.worktree == "/home/user/project-a" + assert result.name == "project-a" + assert result.vcs == "git" + assert result.config_path == "/home/user/project-a/.agentpool.yml" + assert result.settings == {"model": "openai:gpt-4o"} + + +async def test_get_project_not_found(provider: OpenCodeStorageProvider) -> None: + """Test that getting a nonexistent project returns None.""" + result = await provider.get_project("nonexistent_id") + assert result is None + + +async def test_get_project_by_worktree(provider: OpenCodeStorageProvider) -> None: + """Test finding a project by worktree path with path resolution.""" + worktree = str(Path("/tmp/test-worktree-project").resolve()) + project = _make_project( + project_id="proj_worktree", + worktree=worktree, + ) + await provider.save_project(project) + + # Search by the same resolved path + result = await provider.get_project_by_worktree(worktree) + assert result is not None + assert result.project_id == "proj_worktree" + + # Search with unresolved path — should still resolve and match + unresolved = "/tmp/test-worktree-project" + result2 = await provider.get_project_by_worktree(unresolved) + assert result2 is not None + assert result2.project_id == "proj_worktree" + + +async def test_get_project_by_name(provider: OpenCodeStorageProvider) -> None: + """Test finding a project by its friendly name.""" + project = _make_project( + project_id="proj_named", + name="my-special-project", + ) + await provider.save_project(project) + + result = await provider.get_project_by_name("my-special-project") + assert result is not None + assert result.project_id == "proj_named" + + # Nonexistent name returns None + result2 = await provider.get_project_by_name("nonexistent-name") + assert result2 is None + + +async def test_list_projects_sorted_by_last_active(provider: OpenCodeStorageProvider) -> None: + """Test that list_projects returns items sorted by last_active descending.""" + project_a = _make_project( + project_id="proj_a", + name="alpha", + last_active=datetime(2025, 1, 1, tzinfo=UTC), + ) + project_b = _make_project( + project_id="proj_b", + name="beta", + last_active=datetime(2025, 6, 15, tzinfo=UTC), + ) + project_c = _make_project( + project_id="proj_c", + name="gamma", + last_active=datetime(2025, 3, 10, tzinfo=UTC), + ) + + await provider.save_project(project_a) + await provider.save_project(project_b) + await provider.save_project(project_c) + + result = await provider.list_projects() + + assert len(result) == 3 + # Sorted by last_active descending: beta (June), gamma (March), alpha (Jan) + assert result[0].project_id == "proj_b" + assert result[1].project_id == "proj_c" + assert result[2].project_id == "proj_a" + + +async def test_list_projects_with_limit(provider: OpenCodeStorageProvider) -> None: + """Test that the limit parameter works correctly.""" + project_a = _make_project( + project_id="proj_limit_a", + last_active=datetime(2025, 1, 1, tzinfo=UTC), + ) + project_b = _make_project( + project_id="proj_limit_b", + last_active=datetime(2025, 6, 15, tzinfo=UTC), + ) + project_c = _make_project( + project_id="proj_limit_c", + last_active=datetime(2025, 3, 10, tzinfo=UTC), + ) + + await provider.save_project(project_a) + await provider.save_project(project_b) + await provider.save_project(project_c) + + result = await provider.list_projects(limit=2) + assert len(result) == 2 + # Should be the two most recently active + assert result[0].project_id == "proj_limit_b" + assert result[1].project_id == "proj_limit_c" + + +async def test_delete_project(provider: OpenCodeStorageProvider) -> None: + """Test deleting a project removes the file and returns True.""" + project = _make_project(project_id="proj_delete_me") + await provider.save_project(project) + + # Verify it exists + result = await provider.get_project("proj_delete_me") + assert result is not None + + # Delete it + deleted = await provider.delete_project("proj_delete_me") + assert deleted is True + + # Verify it's gone + result2 = await provider.get_project("proj_delete_me") + assert result2 is None + + # JSON file should be removed + project_file = provider.projects_path / "proj_delete_me.json" + assert not project_file.exists() + + +async def test_delete_project_not_found(provider: OpenCodeStorageProvider) -> None: + """Test deleting a nonexistent project returns False.""" + deleted = await provider.delete_project("nonexistent_project") + assert deleted is False + + +async def test_touch_project_updates_timestamp(provider: OpenCodeStorageProvider) -> None: + """Test that touch_project updates the last_active timestamp.""" + original_time = datetime(2020, 1, 1, 0, 0, 0, tzinfo=UTC) + project = _make_project( + project_id="proj_touch", + last_active=original_time, + ) + await provider.save_project(project) + + # Verify original timestamp + result = await provider.get_project("proj_touch") + assert result is not None + assert result.last_active.year == 2020 + + # Touch it + await provider.touch_project("proj_touch") + + # Verify timestamp was updated + result2 = await provider.get_project("proj_touch") + assert result2 is not None + assert result2.last_active > original_time + + +async def test_list_projects_handles_corrupted_file(provider: OpenCodeStorageProvider) -> None: + """Test that corrupted JSON files are skipped without failing the whole listing.""" + project = _make_project( + project_id="proj_good", + name="good-project", + ) + await provider.save_project(project) + + # Write a corrupted file directly + corrupted_file = provider.projects_path / "proj_bad.json" + corrupted_file.write_text("{ this is not valid json }", encoding="utf-8") + + result = await provider.list_projects() + + # Should still return the good project + assert len(result) == 1 + assert result[0].project_id == "proj_good" diff --git a/tests/sessions/test_opencode_session_methods.py b/tests/sessions/test_opencode_session_methods.py new file mode 100644 index 000000000..3684591e4 --- /dev/null +++ b/tests/sessions/test_opencode_session_methods.py @@ -0,0 +1,277 @@ +"""Tests for OpenCodeStorageProvider session methods. + +Covers: +- update_session_title, update_sdk_session_id +- delete_session_messages, get_filtered_conversations +""" + +from __future__ import annotations + +from datetime import UTC, datetime +import tempfile +from typing import TYPE_CHECKING + + +if TYPE_CHECKING: + from pathlib import Path + +import anyenv +import pytest + +from agentpool_config.storage import OpenCodeStorageConfig +from agentpool_server.opencode_server.models import Session, TimeCreatedUpdated +from agentpool_storage.opencode_provider import OpenCodeStorageProvider +from agentpool_storage.opencode_provider.helpers import compute_project_id + + +@pytest.fixture +async def provider(): + """Create an OpenCode provider with temp directory.""" + with tempfile.TemporaryDirectory() as tmpdir: + config = OpenCodeStorageConfig(path=tmpdir) + prov = OpenCodeStorageProvider(config) + async with prov: + yield prov + + +def _write_session_json( + provider: OpenCodeStorageProvider, + session_id: str, + *, + title: str = "Test Session", + project_id: str | None = None, +) -> Path: + """Write a session JSON file directly to the provider's sessions_path.""" + pid = project_id or compute_project_id(str(provider.base_path)) + project_dir = provider.sessions_path / pid + project_dir.mkdir(parents=True, exist_ok=True) + + now_ms = int(datetime.now(UTC).timestamp() * 1000) + session = Session( + id=session_id, + project_id=pid, + directory=str(provider.base_path), + title=title, + time=TimeCreatedUpdated(created=now_ms, updated=now_ms), + ) + session_path = project_dir / f"{session_id}.json" + dct = session.model_dump(by_alias=True) + session_path.write_text(anyenv.dump_json(dct, indent=True), encoding="utf-8") + return session_path + + +def _write_message_json( + provider: OpenCodeStorageProvider, + session_id: str, + message_id: str, + role: str = "user", +) -> Path: + """Write a minimal message JSON file and return its path.""" + msg_dir = provider.messages_path / session_id + msg_dir.mkdir(parents=True, exist_ok=True) + + now_ms = int(datetime.now(UTC).timestamp() * 1000) + data = { + "id": message_id, + "sessionID": session_id, + "role": role, + "time": {"created": now_ms}, + } + msg_path = msg_dir / f"{message_id}.json" + msg_path.write_text(anyenv.dump_json(data, indent=True), encoding="utf-8") + return msg_path + + +def _write_part_json( + provider: OpenCodeStorageProvider, + message_id: str, + part_id: str, +) -> Path: + """Write a minimal part JSON file and return its path.""" + parts_dir = provider.parts_path / message_id + parts_dir.mkdir(parents=True, exist_ok=True) + + now_ms = int(datetime.now(UTC).timestamp() * 1000) + data = { + "id": part_id, + "messageID": message_id, + "type": "text", + "text": "hello", + "time": {"start": now_ms}, + } + part_path = parts_dir / f"{part_id}.json" + part_path.write_text(anyenv.dump_json(data, indent=True), encoding="utf-8") + return part_path + + +# --- update_session_title --- + + +async def test_update_session_title_persists(provider: OpenCodeStorageProvider) -> None: + """Test that update_session_title modifies the session JSON on disk.""" + session_path = _write_session_json(provider, "sess_title_001", title="Old Title") + + await provider.update_session_title("sess_title_001", "New Title") + + # Read raw JSON from disk to verify title was updated + content = session_path.read_text(encoding="utf-8") + data = anyenv.load_json(content, return_type=dict) + assert data["title"] == "New Title" + + +async def test_update_session_title_not_found(provider: OpenCodeStorageProvider) -> None: + """Test that updating a nonexistent session logs a warning but does not raise.""" + # Should not raise + await provider.update_session_title("nonexistent_session", "Title") + + +# --- update_sdk_session_id --- + + +async def test_update_sdk_session_id_persists(provider: OpenCodeStorageProvider) -> None: + """Test that update_sdk_session_id adds metadata.sdk_session_id to the JSON.""" + session_path = _write_session_json(provider, "sess_sdk_001") + + await provider.update_sdk_session_id("sess_sdk_001", "sdk-sess-abc123") + + # Read raw JSON to verify metadata.sdk_session_id was written + content = session_path.read_text(encoding="utf-8") + data = anyenv.load_json(content, return_type=dict) + assert data["metadata"]["sdk_session_id"] == "sdk-sess-abc123" + + +async def test_update_sdk_session_id_not_found(provider: OpenCodeStorageProvider) -> None: + """Test that updating a nonexistent session logs a warning but does not raise.""" + await provider.update_sdk_session_id("nonexistent_session", "sdk-id") + + +# --- delete_session_messages --- + + +async def test_delete_session_messages_removes_files(provider: OpenCodeStorageProvider) -> None: + """Test that delete_session_messages removes message JSON and part JSON files.""" + _write_session_json(provider, "sess_del_001") + _write_message_json(provider, "sess_del_001", "msg_001") + _write_message_json(provider, "sess_del_001", "msg_002") + _write_part_json(provider, "msg_001", "msg_001-0") + _write_part_json(provider, "msg_002", "msg_002-0") + + # Verify files exist + assert (provider.messages_path / "sess_del_001" / "msg_001.json").exists() + assert (provider.messages_path / "sess_del_001" / "msg_002.json").exists() + assert (provider.parts_path / "msg_001" / "msg_001-0.json").exists() + assert (provider.parts_path / "msg_002" / "msg_002-0.json").exists() + + count = await provider.delete_session_messages("sess_del_001") + + assert count == 2 + # Message files should be gone + assert not (provider.messages_path / "sess_del_001" / "msg_001.json").exists() + assert not (provider.messages_path / "sess_del_001" / "msg_002.json").exists() + # Part files should be gone + assert not (provider.parts_path / "msg_001" / "msg_001-0.json").exists() + assert not (provider.parts_path / "msg_002" / "msg_002-0.json").exists() + + +async def test_delete_session_messages_nonexistent(provider: OpenCodeStorageProvider) -> None: + """Test that deleting messages for a nonexistent session returns 0 gracefully.""" + count = await provider.delete_session_messages("nonexistent_session") + assert count == 0 + + +# --- get_filtered_conversations --- + + +async def test_get_filtered_conversations_by_name(provider: OpenCodeStorageProvider) -> None: + """Test that filtering by agent_name only returns matching sessions.""" + # Create two sessions with messages + _write_session_json(provider, "sess_filter_001", title="Session Alpha") + _write_session_json(provider, "sess_filter_002", title="Session Beta") + + # Create messages with different roles (agent_name comes from message name) + msg_dir_1 = provider.messages_path / "sess_filter_001" + msg_dir_1.mkdir(parents=True, exist_ok=True) + now_ms = int(datetime.now(UTC).timestamp() * 1000) + + # Message from "agent_alpha" + alpha_msg = { + "id": "msg_alpha", + "sessionID": "sess_filter_001", + "role": "assistant", + "parentID": "", + "modelID": "gpt-4o", + "providerID": "", + "path": {"cwd": "", "root": ""}, + "time": {"created": now_ms}, + "tokens": {"input": 10, "output": 20, "cache": {"read": 0, "write": 0}}, + "cost": 0.01, + "finish": "stop", + } + (msg_dir_1 / "msg_alpha.json").write_text( + anyenv.dump_json(alpha_msg, indent=True), encoding="utf-8" + ) + + # Part for alpha message + parts_dir = provider.parts_path / "msg_alpha" + parts_dir.mkdir(parents=True, exist_ok=True) + alpha_part = { + "id": "msg_alpha-0", + "messageID": "msg_alpha", + "sessionID": "sess_filter_001", + "type": "text", + "text": "Hello from alpha", + "time": {"start": now_ms}, + } + (parts_dir / "msg_alpha-0.json").write_text( + anyenv.dump_json(alpha_part, indent=True), encoding="utf-8" + ) + + # Second session with "agent_beta" + msg_dir_2 = provider.messages_path / "sess_filter_002" + msg_dir_2.mkdir(parents=True, exist_ok=True) + + beta_msg = { + "id": "msg_beta", + "sessionID": "sess_filter_002", + "role": "assistant", + "parentID": "", + "modelID": "claude-sonnet", + "providerID": "", + "path": {"cwd": "", "root": ""}, + "time": {"created": now_ms}, + "tokens": {"input": 15, "output": 25, "cache": {"read": 0, "write": 0}}, + "cost": 0.02, + "finish": "stop", + } + (msg_dir_2 / "msg_beta.json").write_text( + anyenv.dump_json(beta_msg, indent=True), encoding="utf-8" + ) + + beta_parts_dir = provider.parts_path / "msg_beta" + beta_parts_dir.mkdir(parents=True, exist_ok=True) + beta_part = { + "id": "msg_beta-0", + "messageID": "msg_beta", + "sessionID": "sess_filter_002", + "type": "text", + "text": "Hello from beta", + "time": {"start": now_ms}, + } + (beta_parts_dir / "msg_beta-0.json").write_text( + anyenv.dump_json(beta_part, indent=True), encoding="utf-8" + ) + + # Filter by content query - only alpha matches + results = await provider.get_filtered_conversations(query="alpha") + assert len(results) == 1 + assert results[0]["id"] == "sess_filter_001" + + # No filter - should return both + all_results = await provider.get_filtered_conversations() + assert len(all_results) == 2 + + +async def test_get_filtered_conversations_empty(provider: OpenCodeStorageProvider) -> None: + """Test that no sessions returns an empty list.""" + results = await provider.get_filtered_conversations() + assert results == [] diff --git a/tests/sessions/test_project_provider_selection.py b/tests/sessions/test_project_provider_selection.py new file mode 100644 index 000000000..502f345f4 --- /dev/null +++ b/tests/sessions/test_project_provider_selection.py @@ -0,0 +1,68 @@ +"""Tests for StorageManager.get_project_provider() capability-based selection.""" + +from __future__ import annotations + +import pytest + +from agentpool.storage.manager import StorageManager +from agentpool_config.storage import MemoryStorageConfig, StorageConfig +from agentpool_storage.memory_provider import MemoryStorageProvider +from agentpool_storage.opencode_provider import OpenCodeStorageProvider + + +def _make_manager_with_providers(*capable: bool) -> StorageManager: + """Create a StorageManager with mock providers. + + Args: + capable: Whether each provider has can_store_projects=True + """ + config = StorageConfig() + manager = StorageManager(config) + manager.providers = [] + for is_capable in capable: + provider = MemoryStorageProvider(MemoryStorageConfig()) + provider.can_store_projects = is_capable + manager.providers.append(provider) + return manager + + +def test_get_project_provider_returns_first_capable() -> None: + """First capable provider is returned.""" + manager = _make_manager_with_providers(False, True, True) + provider = manager.get_project_provider() + assert provider is manager.providers[1] + assert provider.can_store_projects is True + + +def test_get_project_provider_skips_incapable() -> None: + """Incapable providers are skipped, capable one selected.""" + manager = _make_manager_with_providers(False, False, True) + provider = manager.get_project_provider() + assert provider is manager.providers[2] + assert provider.can_store_projects is True + + +def test_get_project_provider_raises_when_none_capable() -> None: + """RuntimeError raised when no provider supports project storage.""" + manager = _make_manager_with_providers(False, False) + with pytest.raises(RuntimeError, match="No storage provider supports project storage"): + manager.get_project_provider() + + +def test_get_project_provider_raises_when_no_providers() -> None: + """RuntimeError raised when provider list is empty.""" + manager = _make_manager_with_providers() + with pytest.raises(RuntimeError, match="No storage provider supports project storage"): + manager.get_project_provider() + + +def test_opencode_provider_is_capable() -> None: + """OpenCodeStorageProvider has can_store_projects=True after implementing project methods.""" + provider = OpenCodeStorageProvider() + assert provider.can_store_projects is True + + +def test_memory_provider_is_capable_by_default() -> None: + """MemoryStorageProvider has can_store_projects=True by default.""" + provider = MemoryStorageProvider(MemoryStorageConfig()) + assert provider.can_store_projects is True diff --git a/tests/storage/__init__.py b/tests/storage/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/storage/test_opencode_provider.py b/tests/storage/test_opencode_provider.py new file mode 100644 index 000000000..6926cbb62 --- /dev/null +++ b/tests/storage/test_opencode_provider.py @@ -0,0 +1,174 @@ +"""Tests for OpenCodeStorageProvider encapsulation and path correctness. + +Covers: +- get_project_by_worktree uses direct O(1) lookup via generate_project_id +- _write_message uses self.base_path instead of Path.cwd() +""" + +from __future__ import annotations + +from pathlib import Path + +import anyenv +import pytest +from pydantic_ai.messages import ModelRequest, ModelResponse, UserPromptPart + +from agentpool.sessions.models import ProjectData +from agentpool_storage.opencode_provider.provider import OpenCodeStorageProvider +from agentpool_storage.project_store import generate_project_id + + +@pytest.fixture +def provider(tmp_path: Path) -> OpenCodeStorageProvider: + """Create an OpenCodeStorageProvider backed by a temporary directory.""" + from agentpool_config.storage import OpenCodeStorageConfig + + config = OpenCodeStorageConfig(path=str(tmp_path / "storage")) + return OpenCodeStorageProvider(config=config) + + +async def test_get_project_by_worktree_finds_existing_project(provider: OpenCodeStorageProvider): + """get_project_by_worktree should find a project by its worktree path using direct lookup.""" + worktree = str(Path("/tmp/test-project-worktree").resolve()) + project_id = generate_project_id(worktree) + project = ProjectData(project_id=project_id, worktree=worktree, name="test-project") + await provider.save_project(project) + + result = await provider.get_project_by_worktree(worktree) + assert result is not None + assert result.project_id == project_id + assert result.worktree == worktree + assert result.name == "test-project" + + +async def test_get_project_by_worktree_returns_none_for_missing(provider: OpenCodeStorageProvider): + """get_project_by_worktree should return None when no project file exists.""" + result = await provider.get_project_by_worktree("/nonexistent/path") + assert result is None + + +async def test_get_project_by_worktree_direct_lookup_not_scan(provider: OpenCodeStorageProvider, tmp_path: Path): + """Verify get_project_by_worktree uses direct file lookup, not O(N) scan. + + Creates multiple project files and verifies that only the target project + is found by worktree lookup. The method should read exactly one file + (the one computed by generate_project_id), not iterate all project files. + """ + # Create several projects with different worktrees + for i in range(5): + worktree = str(Path(f"/tmp/worktree-{i}").resolve()) + pid = generate_project_id(worktree) + project = ProjectData(project_id=pid, worktree=worktree, name=f"project-{i}") + await provider.save_project(project) + + # Look up a specific project + target_worktree = str(Path("/tmp/worktree-3").resolve()) + result = await provider.get_project_by_worktree(target_worktree) + assert result is not None + assert result.name == "project-3" + + # Verify that a non-existent worktree returns None + # (even though other project files exist in the same directory) + result_missing = await provider.get_project_by_worktree("/tmp/nonexistent-worktree") + assert result_missing is None + + +async def test_get_project_by_worktree_worktree_mismatch_returns_none(provider: OpenCodeStorageProvider): + """If the stored project's worktree doesn't match, should return None. + + This guards against hash collisions or stale/corrupted data. + """ + worktree = "/tmp/real-worktree" + project_id = generate_project_id(worktree) + # Manually create a project file with wrong worktree (simulating stale data) + project = ProjectData(project_id=project_id, worktree="/tmp/different-worktree", name="stale") + # Write it directly to the expected file location + project_file = provider.projects_path / f"{project_id}.json" + data = project.model_dump(mode="json") + project_file.write_text(anyenv.dump_json(data, indent=True), encoding="utf-8") + + # Lookup by the original worktree should return None because the stored + # worktree doesn't match (safety verification) + result = await provider.get_project_by_worktree(worktree) + assert result is None + + +async def test_write_message_uses_base_path_not_cwd(provider: OpenCodeStorageProvider, tmp_path: Path): + """_write_message should use self.base_path for MessagePath, not Path.cwd(). + + This ensures that when the process CWD differs from the provider's base_path, + the stored message still references the correct directory. + """ + from pydantic_ai.messages import ModelRequest, ModelResponse, UserPromptPart + + session_id = "test-session-base-path" + message_id = "msg-basepath-001" + + # The provider's base_path should NOT be Path.cwd() in this test + assert provider.base_path != Path.cwd() or True # base_path is tmp_path based + + # Write a user message + model_messages: list[ModelRequest | ModelResponse] = [ + ModelRequest(parts=[UserPromptPart(content="Hello, base_path test")]), + ] + + await provider._write_message( + message_id=message_id, + session_id=session_id, + role="assistant", + model_messages=model_messages, + model="test:model", + ) + + # Read the message file and verify the path fields + msg_file = provider.messages_path / session_id / f"{message_id}.json" + assert msg_file.exists(), f"Message file not found at {msg_file}" + + content = msg_file.read_text(encoding="utf-8") + data = anyenv.load_json(content, return_type=dict) + + # The 'path' field should contain the base_path, not the process CWD + path_data = data.get("path", {}) + cwd = path_data.get("cwd", "") + root = path_data.get("root", "") + + expected_base = str(provider.base_path) + assert cwd == expected_base, f"Expected cwd={expected_base!r}, got {cwd!r}" + assert root == expected_base, f"Expected root={expected_base!r}, got {root!r}" + + +async def test_write_message_base_path_differs_from_process_cwd(provider: OpenCodeStorageProvider, tmp_path: Path): + """When provider.base_path differs from process CWD, messages should use base_path. + + This is a more targeted test ensuring the provider doesn't accidentally + fall back to Path.cwd(). + """ + from pydantic_ai.messages import ModelRequest, ModelResponse, UserPromptPart + + # Provider's base_path is under tmp_path, which should differ from cwd + provider_base = str(provider.base_path) + process_cwd = str(Path.cwd()) + + # If they happen to be the same, this test can't distinguish the fix + # but the other test still validates the content + if provider_base == process_cwd: + pytest.skip("Provider base_path happens to equal process CWD; test cannot distinguish") + + session_id = "test-session-cwd-diff" + message_id = "msg-cwd-diff-001" + + await provider._write_message( + message_id=message_id, + session_id=session_id, + role="assistant", + model_messages=[ModelRequest(parts=[UserPromptPart(content="test")])], + model="test:model", + ) + + msg_file = provider.messages_path / session_id / f"{message_id}.json" + data = anyenv.load_json(msg_file.read_text(encoding="utf-8"), return_type=dict) + path_data = data.get("path", {}) + + # Must use base_path, NOT process cwd + assert path_data.get("cwd") == provider_base + assert path_data.get("cwd") != process_cwd