diff --git a/internal/acp/client.go b/internal/acp/client.go index 5ef360b..c049a0e 100644 --- a/internal/acp/client.go +++ b/internal/acp/client.go @@ -31,6 +31,8 @@ type Client struct { sessionID string transport *transport history *History + streamChs []chan<- StreamEvent + notifyWG sync.WaitGroup } // NewClient creates an ACP client for the given workspace. @@ -159,6 +161,73 @@ func (c *Client) Prompt(ctx context.Context, content string) (PromptResult, erro return result, nil } +// sendStream writes ev to out without blocking. It is used for both chunk +// notifications and terminal events so that a slow consumer cannot stall the +// transport goroutine. +func (c *Client) sendStream(out chan<- StreamEvent, ev StreamEvent) { + select { + case out <- ev: + default: + } +} + +// AddStream registers a new streaming consumer. It must be paired with RemoveStream. +func (c *Client) AddStream(ch chan<- StreamEvent) int { + c.mu.Lock() + defer c.mu.Unlock() + c.streamChs = append(c.streamChs, ch) + return len(c.streamChs) +} + +// RemoveStream unregisters a streaming consumer previously added with AddStream. +func (c *Client) RemoveStream(ch chan<- StreamEvent) { + c.mu.Lock() + defer c.mu.Unlock() + for i, existing := range c.streamChs { + if existing == ch { + c.streamChs[i] = c.streamChs[len(c.streamChs)-1] + c.streamChs = c.streamChs[:len(c.streamChs)-1] + return + } + } +} + +// Stream sends a user prompt and forwards streamed chunks to out as they arrive. +// It emits exactly one terminal event, either complete or error, before returning. +func (c *Client) Stream(ctx context.Context, content string, out chan<- StreamEvent) error { + c.mu.RLock() + if !c.ready { + c.mu.RUnlock() + c.sendStream(out, StreamEvent{Type: "error", Error: "client not ready"}) + return util.New(util.CodeConflict, "client not ready") + } + sessionID := c.sessionID + tr := c.transport + c.mu.RUnlock() + + c.history.Add("user", content) + + c.AddStream(out) + defer c.RemoveStream(out) + + params := SessionPromptParams{ + SessionID: sessionID, + Prompt: []PromptMessage{{Type: "text", Text: content}}, + } + var result PromptResult + if err := tr.call(ctx, "session/prompt", params, &result); err != nil { + c.sendStream(out, StreamEvent{Type: "error", Error: err.Error()}) + return util.Wrap(util.CodeInternal, "session/prompt failed", err) + } + + // Wait for any chunk notifications that were already in flight to be + // delivered before we emit the terminal complete event. + c.notifyWG.Wait() + + c.sendStream(out, StreamEvent{Type: "complete", StopReason: result.StopReason}) + return nil +} + // Cancel sends a best-effort session/cancel notification. func (c *Client) Cancel() error { c.mu.RLock() @@ -230,26 +299,45 @@ func (c *Client) handleNotification(method string, params json.RawMessage) { c.lg.Debug("acp drop notification", "workspace_id", c.workspaceID, "method", method) return } + + c.notifyWG.Add(1) + defer c.notifyWG.Done() + var up SessionUpdateParams if err := json.Unmarshal(params, &up); err != nil { c.lg.Error("acp invalid session/update", "workspace_id", c.workspaceID, "error", err) return } + + var role string switch up.Update.SessionUpdate { case "agent_message_chunk": - if up.Update.Content.Type == "text" { - c.history.AppendOrCreate("agent", up.Update.MessageID, up.Update.Content.Text) - } else { - c.lg.Debug("acp drop non-text agent chunk", "workspace_id", c.workspaceID, "type", up.Update.Content.Type) - } + role = "agent" case "user_message_chunk": - if up.Update.Content.Type == "text" { - c.history.AppendOrCreate("user", up.Update.MessageID, up.Update.Content.Text) - } else { - c.lg.Debug("acp drop non-text user chunk", "workspace_id", c.workspaceID, "type", up.Update.Content.Type) - } + role = "user" default: c.lg.Debug("acp drop update", "workspace_id", c.workspaceID, "type", up.Update.SessionUpdate) + return + } + + if up.Update.Content.Type != "text" { + c.lg.Debug("acp drop non-text chunk", "workspace_id", c.workspaceID, "type", up.Update.Content.Type) + return + } + + c.history.AppendOrCreate(role, up.Update.MessageID, up.Update.Content.Text) + + ev := StreamEvent{ + Type: "chunk", + MessageID: up.Update.MessageID, + Text: up.Update.Content.Text, + } + c.mu.RLock() + streams := make([]chan<- StreamEvent, len(c.streamChs)) + copy(streams, c.streamChs) + c.mu.RUnlock() + for _, ch := range streams { + c.sendStream(ch, ev) } } diff --git a/internal/acp/client_test.go b/internal/acp/client_test.go index fd1408a..fdc9706 100644 --- a/internal/acp/client_test.go +++ b/internal/acp/client_test.go @@ -2,6 +2,7 @@ package acp import ( "bytes" + "context" "encoding/json" "log/slog" "os" @@ -351,3 +352,74 @@ func TestHistoryAppendOrCreateGroupsByMessageID(t *testing.T) { t.Fatalf("after clear message = %+v, want text=\"after clear\" messageId=msg-1", msgs[0]) } } + +func TestClientStreamEmitsChunkAndComplete(t *testing.T) { + c := &Client{ + workspaceID: "ws", + fs: newStubFS(), + history: NewHistory(), + lg: slog.Default(), + } + stdin := &recordingWriteCloser{} + sub := newFakeSubscription() + tr := newTransport(stdin, sub, slog.Default()) + tr.onNotification = c.handleNotification + c.transport = tr + c.ready = true + c.sessionID = "s1" + tr.start() + defer tr.Close() + + out := make(chan StreamEvent, 16) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + go func() { + for strings.TrimSpace(stdin.String()) == "" { + time.Sleep(5 * time.Millisecond) + } + chunk1 := `{"jsonrpc":"2.0","method":"session/update","params":{"sessionId":"s1","update":{"sessionUpdate":"agent_message_chunk","messageId":"m1","content":{"type":"text","text":"hello"}}}}` + chunk2 := `{"jsonrpc":"2.0","method":"session/update","params":{"sessionId":"s1","update":{"sessionUpdate":"agent_message_chunk","messageId":"m1","content":{"type":"text","text":" world"}}}}` + response := `{"jsonrpc":"2.0","id":1,"result":{"stopReason":"end_turn"}}` + sub.send([]byte(chunk1 + "\n")) + sub.send([]byte(chunk2 + "\n")) + sub.send([]byte(response + "\n")) + }() + + if err := c.Stream(ctx, "hi", out); err != nil { + t.Fatalf("Stream error: %v", err) + } + close(out) + + var chunks int + var complete *StreamEvent + for ev := range out { + switch ev.Type { + case "chunk": + chunks++ + if ev.MessageID != "m1" { + t.Fatalf("chunk messageId = %q, want m1", ev.MessageID) + } + case "complete": + if complete != nil { + t.Fatal("received more than one complete event") + } + cp := ev + complete = &cp + case "error": + t.Fatalf("unexpected error event: %q", ev.Error) + default: + t.Fatalf("unexpected event type %q", ev.Type) + } + } + + if chunks != 2 { + t.Fatalf("chunks = %d, want 2", chunks) + } + if complete == nil { + t.Fatal("expected complete event") + } + if complete.StopReason != "end_turn" { + t.Fatalf("stopReason = %q, want end_turn", complete.StopReason) + } +} diff --git a/internal/acp/messages.go b/internal/acp/messages.go index 0c8b62c..530ed38 100644 --- a/internal/acp/messages.go +++ b/internal/acp/messages.go @@ -19,6 +19,17 @@ type rpcError struct { Data any `json:"data,omitempty"` } +// StreamEvent is the wire shape of the streaming protocol. The Type field +// is one of "chunk", "complete", "error"; the other fields are populated +// based on Type. +type StreamEvent struct { + Type string `json:"type"` + MessageID string `json:"messageId,omitempty"` + Text string `json:"text,omitempty"` + StopReason string `json:"stopReason,omitempty"` + Error string `json:"error,omitempty"` +} + // InitializeParams is sent once to negotiate the ACP session. type InitializeParams struct { ProtocolVersion int `json:"protocolVersion"` @@ -91,10 +102,10 @@ type SessionUpdateParams struct { // Update carries a single streamed update from the agent. type Update struct { - SessionUpdate string `json:"sessionUpdate"` - MessageID string `json:"messageId,omitempty"` + SessionUpdate string `json:"sessionUpdate"` + MessageID string `json:"messageId,omitempty"` Content ContentBlock `json:"content,omitempty"` - ToolCallID string `json:"toolCallId,omitempty"` + ToolCallID string `json:"toolCallId,omitempty"` } // ContentBlock is a content payload in session/update notifications. diff --git a/internal/acp/service.go b/internal/acp/service.go index d06daac..0211d62 100644 --- a/internal/acp/service.go +++ b/internal/acp/service.go @@ -16,6 +16,7 @@ type Service interface { Status(workspaceID string) (StatusResponse, error) History(workspaceID string) (HistoryResponse, error) Prompt(workspaceID string, content string) (PromptResponse, error) + Stream(workspaceID string, content string, out chan<- StreamEvent) error Cancel(workspaceID string) error EnsureReady(workspaceID string) error } @@ -186,3 +187,24 @@ func (s *localService) Cancel(workspaceID string) error { } return c.Cancel() } + +// Stream sends a user prompt and streams response chunks to out. +func (s *localService) Stream(workspaceID string, content string, out chan<- StreamEvent) error { + if _, err := s.workspaces.Get(workspaceID); err != nil { + return err + } + c, err := s.clientFor(workspaceID) + if err != nil { + return err + } + lock := s.workspaceLock(workspaceID) + lock.Lock() + defer lock.Unlock() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) + defer cancel() + if err := c.Initialize(ctx); err != nil { + return err + } + return c.Stream(ctx, content, out) +} diff --git a/internal/acp/transport.go b/internal/acp/transport.go index c610893..0914ed8 100644 --- a/internal/acp/transport.go +++ b/internal/acp/transport.go @@ -95,7 +95,7 @@ func (t *transport) dispatchLoop() { case msg.ID == nil && msg.Method != "": // Notification from the agent. if t.onNotification != nil { - go t.onNotification(msg.Method, msg.Params) + t.onNotification(msg.Method, msg.Params) } case msg.ID != nil && (msg.Result != nil || msg.Error != nil): // Response to a client-initiated request. diff --git a/internal/api/acp_handler.go b/internal/api/acp_handler.go index 5f43074..24abbab 100644 --- a/internal/api/acp_handler.go +++ b/internal/api/acp_handler.go @@ -1,12 +1,16 @@ package api import ( + "encoding/json" "net/http" + "sync" + "time" "codespace/internal/model" "codespace/internal/service" "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" ) type acpHandler struct { @@ -67,3 +71,108 @@ func (h *acpHandler) cancel(c *gin.Context) { } c.Status(http.StatusNoContent) } + +// stream upgrades the connection to a WebSocket and streams the prompt response. +// +// Protocol: +// +// Client -> Server: {"type":"prompt","content":"..."} +// Server -> Client: zero or more {"type":"chunk","messageId":"...","text":"..."} +// Server -> Client: exactly one {"type":"complete","stopReason":"..."} or +// {"type":"error","error":"..."} +// +// Closing the WebSocket from the client side cancels the in-flight prompt. +func (h *acpHandler) stream(c *gin.Context) { + id := c.Param("id") + + upgrader := websocket.Upgrader{ + CheckOrigin: func(r *http.Request) bool { return true }, + ReadBufferSize: 4096, + WriteBufferSize: 4096, + } + + conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) + if err != nil { + return + } + + done := make(chan struct{}) + var closeOnce sync.Once + closeAll := func() { + closeOnce.Do(func() { + conn.Close() + close(done) + }) + } + + conn.SetReadLimit(64 << 10) + conn.SetReadDeadline(time.Now().Add(60 * time.Second)) + conn.SetPongHandler(func(string) error { + conn.SetReadDeadline(time.Now().Add(60 * time.Second)) + return nil + }) + + // Read the single prompt request from the client. + _, data, err := conn.ReadMessage() + if err != nil { + closeAll() + return + } + + var req model.AcpStreamRequest + if err := json.Unmarshal(data, &req); err != nil || req.Type != "prompt" || req.Content == "" { + _ = conn.WriteJSON(model.AcpStreamEvent{Type: "error", Error: "expected {type: prompt, content: ...}"}) + closeAll() + return + } + + out := make(chan model.AcpStreamEvent, 32) + + // Run the prompt in the background and close the event channel when done. + go func() { + defer close(out) + if err := h.svc.Stream(id, req.Content, out); err != nil { + select { + case out <- model.AcpStreamEvent{Type: "error", Error: err.Error()}: + default: + } + } + }() + + // Write events and pings to the WebSocket. + go func() { + ticker := time.NewTicker(30 * time.Second) + defer ticker.Stop() + for { + select { + case <-done: + return + case ev, ok := <-out: + if !ok { + conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + closeAll() + return + } + if err := conn.WriteJSON(ev); err != nil { + closeAll() + return + } + case <-ticker.C: + if err := conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(5*time.Second)); err != nil { + closeAll() + return + } + } + } + }() + + // Wait for the client to close the connection and cancel any in-flight prompt. + for { + _, _, err := conn.ReadMessage() + if err != nil { + h.svc.Cancel(id) + closeAll() + break + } + } +} diff --git a/internal/api/acp_handler_test.go b/internal/api/acp_handler_test.go new file mode 100644 index 0000000..2ac2748 --- /dev/null +++ b/internal/api/acp_handler_test.go @@ -0,0 +1,112 @@ +package api + +import ( + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "codespace/internal/model" + "codespace/internal/process" + "codespace/internal/service" + "codespace/internal/shell" + "codespace/internal/workspace" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" +) + +func TestAcpStreamHandlerRoutes(t *testing.T) { + tmpDir := t.TempDir() + + opencodePath := filepath.Join(tmpDir, "opencode") + script := []byte(`#!/bin/sh +while IFS= read -r line; do + id=$(printf '%s' "$line" | grep -o '"id":[0-9]*' | cut -d: -f2) + case "$line" in + *initialize*) + printf '{"jsonrpc":"2.0","id":%s,"result":{"protocolVersion":1,"agentCapabilities":{},"agentInfo":{}}}\n' "$id" + ;; + *session/new*) + printf '{"jsonrpc":"2.0","id":%s,"result":{"sessionId":"sess-1"}}\n' "$id" + ;; + *session/prompt*) + printf '{"jsonrpc":"2.0","method":"session/update","params":{"sessionId":"sess-1","update":{"sessionUpdate":"agent_message_chunk","messageId":"m1","content":{"type":"text","text":"hello"}}}}\n' + printf '{"jsonrpc":"2.0","method":"session/update","params":{"sessionId":"sess-1","update":{"sessionUpdate":"agent_message_chunk","messageId":"m1","content":{"type":"text","text":" world"}}}}\n' + printf '{"jsonrpc":"2.0","id":%s,"result":{"stopReason":"end_turn"}}\n' "$id" + ;; + esac +done +`) + if err := os.WriteFile(opencodePath, script, 0o755); err != nil { + t.Fatalf("write fake opencode: %v", err) + } + + wsRoot := filepath.Join(tmpDir, "workspaces") + wsMgr := workspace.NewLocalManager(wsRoot) + procMgr := process.NewManager(opencodePath, nil) + shellMgr := shell.NewManager("bash", []string{"-i"}) + wsSvc := service.NewWorkspaceService(wsMgr, procMgr, shellMgr, nil) + fileSvc := service.NewFileService(wsMgr, 1<<20) + procSvc := service.NewProcessService(wsMgr, procMgr, nil) + shellSvc := service.NewShellService(wsMgr, shellMgr, nil) + acpSvc := service.NewAcpService(procMgr, wsMgr, nil) + + r := NewRouter(wsSvc, fileSvc, procSvc, shellSvc, acpSvc, nil, gin.TestMode) + srv := httptest.NewServer(r) + defer srv.Close() + + ws, err := wsSvc.Create("streamws") + if err != nil { + t.Fatalf("create workspace: %v", err) + } + t.Cleanup(func() { + _ = wsSvc.Delete(ws.ID) + }) + + wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/api/workspaces/" + ws.ID + "/acp/stream" + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer conn.Close() + + if err := conn.WriteJSON(model.AcpStreamRequest{Type: "prompt", Content: "hello"}); err != nil { + t.Fatalf("write prompt request: %v", err) + } + + if err := conn.SetReadDeadline(time.Now().Add(3 * time.Second)); err != nil { + t.Fatalf("set read deadline: %v", err) + } + + var chunks int + var complete bool + for { + var ev model.AcpStreamEvent + if err := conn.ReadJSON(&ev); err != nil { + t.Fatalf("read event: %v", err) + } + switch ev.Type { + case "chunk": + chunks++ + case "complete": + complete = true + case "error": + t.Fatalf("unexpected error event: %s", ev.Error) + default: + t.Fatalf("unexpected event type %q", ev.Type) + } + if complete { + break + } + } + + if chunks != 2 { + t.Fatalf("chunks = %d, want 2", chunks) + } + if !complete { + t.Fatal("expected complete event") + } +} diff --git a/internal/api/router.go b/internal/api/router.go index 0013631..4023d73 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -55,6 +55,7 @@ func NewRouter(workspaces *service.WorkspaceService, files *service.FileService, api.GET("/workspaces/:id/acp/history", acpHandler.history) api.POST("/workspaces/:id/acp/prompt", acpHandler.prompt) api.POST("/workspaces/:id/acp/cancel", acpHandler.cancel) + api.GET("/workspaces/:id/acp/stream", acpHandler.stream) return r } diff --git a/internal/model/acp.go b/internal/model/acp.go index 16c8151..f90c08c 100644 --- a/internal/model/acp.go +++ b/internal/model/acp.go @@ -37,3 +37,18 @@ type AcpPromptResponse struct { StopReason string `json:"stopReason"` Text string `json:"text"` } + +// AcpStreamRequest is the first client message on the ACP streaming WebSocket. +type AcpStreamRequest struct { + Type string `json:"type"` + Content string `json:"content"` +} + +// AcpStreamEvent is a server-to-client streaming event. +type AcpStreamEvent struct { + Type string `json:"type"` + MessageID string `json:"messageId,omitempty"` + Text string `json:"text,omitempty"` + StopReason string `json:"stopReason,omitempty"` + Error string `json:"error,omitempty"` +} diff --git a/internal/service/acp_service.go b/internal/service/acp_service.go index d5d2135..c465784 100644 --- a/internal/service/acp_service.go +++ b/internal/service/acp_service.go @@ -78,6 +78,35 @@ func (s *AcpService) Cancel(workspaceID string) error { return nil } +// Stream sends a prompt to the agent and forwards streamed events to out. +func (s *AcpService) Stream(workspaceID string, content string, out chan<- model.AcpStreamEvent) error { + acpOut := make(chan acp.StreamEvent, 32) + + var streamErr error + done := make(chan struct{}) + go func() { + defer close(done) + streamErr = s.svc.Stream(workspaceID, content, acpOut) + close(acpOut) + }() + + for ev := range acpOut { + out <- model.AcpStreamEvent{ + Type: ev.Type, + MessageID: ev.MessageID, + Text: ev.Text, + StopReason: ev.StopReason, + Error: ev.Error, + } + } + <-done + + if streamErr != nil { + s.lg.Error("acp stream failed", "workspace_id", workspaceID, "error", streamErr) + } + return streamErr +} + // EnsureReady ensures the ACP process is started and initialized. func (s *AcpService) EnsureReady(workspaceID string) error { if err := s.svc.EnsureReady(workspaceID); err != nil {