sink.go

  1package acp
  2
  3import (
  4	"context"
  5	"log/slog"
  6
  7	"github.com/charmbracelet/crush/internal/message"
  8	"github.com/charmbracelet/crush/internal/permission"
  9	"github.com/charmbracelet/crush/internal/pubsub"
 10	"github.com/coder/acp-go-sdk"
 11)
 12
 13// Sink receives events from Crush's pubsub system and translates them to ACP
 14// session updates.
 15type Sink struct {
 16	ctx       context.Context
 17	cancel    context.CancelFunc
 18	conn      *acp.AgentSideConnection
 19	sessionID string
 20
 21	// Track text deltas per message to avoid re-sending content.
 22	textOffsets      map[string]int
 23	reasoningOffsets map[string]int
 24}
 25
 26// NewSink creates a new event sink for the given session.
 27func NewSink(ctx context.Context, conn *acp.AgentSideConnection, sessionID string) *Sink {
 28	sinkCtx, cancel := context.WithCancel(ctx)
 29	return &Sink{
 30		ctx:              sinkCtx,
 31		cancel:           cancel,
 32		conn:             conn,
 33		sessionID:        sessionID,
 34		textOffsets:      make(map[string]int),
 35		reasoningOffsets: make(map[string]int),
 36	}
 37}
 38
 39// Start subscribes to messages and permissions, forwarding events to ACP.
 40func (s *Sink) Start(messages message.Service, permissions permission.Service) {
 41	// Subscribe to message events.
 42	go func() {
 43		msgCh := messages.Subscribe(s.ctx)
 44		for {
 45			select {
 46			case event, ok := <-msgCh:
 47				if !ok {
 48					return
 49				}
 50				s.HandleMessage(event)
 51			case <-s.ctx.Done():
 52				return
 53			}
 54		}
 55	}()
 56
 57	// Subscribe to permission events.
 58	go func() {
 59		permCh := permissions.Subscribe(s.ctx)
 60		for {
 61			select {
 62			case event, ok := <-permCh:
 63				if !ok {
 64					return
 65				}
 66				s.HandlePermission(event.Payload, permissions)
 67			case <-s.ctx.Done():
 68				return
 69			}
 70		}
 71	}()
 72}
 73
 74// Stop cancels the sink's subscriptions.
 75func (s *Sink) Stop() {
 76	s.cancel()
 77}
 78
 79// HandleMessage translates a Crush message event to ACP session updates.
 80func (s *Sink) HandleMessage(event pubsub.Event[message.Message]) {
 81	msg := event.Payload
 82
 83	// Only handle messages for our session.
 84	if msg.SessionID != s.sessionID {
 85		return
 86	}
 87
 88	for _, part := range msg.Parts {
 89		update := s.translatePart(msg.ID, msg.Role, part)
 90		if update == nil {
 91			continue
 92		}
 93
 94		if err := s.conn.SessionUpdate(s.ctx, acp.SessionNotification{
 95			SessionId: acp.SessionId(s.sessionID),
 96			Update:    *update,
 97		}); err != nil {
 98			slog.Error("Failed to send session update", "error", err)
 99		}
100	}
101}
102
103// HandlePermission translates a permission request to an ACP permission request.
104func (s *Sink) HandlePermission(req permission.PermissionRequest, permissions permission.Service) {
105	// Only handle permissions for our session.
106	if req.SessionID != s.sessionID {
107		return
108	}
109
110	slog.Debug("ACP permission request", "tool", req.ToolName, "action", req.Action)
111
112	resp, err := s.conn.RequestPermission(s.ctx, acp.RequestPermissionRequest{
113		SessionId: acp.SessionId(s.sessionID),
114		ToolCall: acp.RequestPermissionToolCall{
115			ToolCallId: acp.ToolCallId(req.ToolCallID),
116			Title:      acp.Ptr(req.Description),
117			Kind:       acp.Ptr(acp.ToolKindEdit),
118			Status:     acp.Ptr(acp.ToolCallStatusPending),
119			Locations:  []acp.ToolCallLocation{{Path: req.Path}},
120			RawInput:   req.Params,
121		},
122		Options: []acp.PermissionOption{
123			{Kind: acp.PermissionOptionKindAllowOnce, Name: "Allow", OptionId: "allow"},
124			{Kind: acp.PermissionOptionKindAllowAlways, Name: "Allow always", OptionId: "allow_always"},
125			{Kind: acp.PermissionOptionKindRejectOnce, Name: "Deny", OptionId: "deny"},
126		},
127	})
128	if err != nil {
129		slog.Error("Failed to request permission", "error", err)
130		permissions.Deny(req)
131		return
132	}
133
134	if resp.Outcome.Cancelled != nil {
135		permissions.Deny(req)
136		return
137	}
138
139	if resp.Outcome.Selected != nil {
140		switch string(resp.Outcome.Selected.OptionId) {
141		case "allow":
142			permissions.Grant(req)
143		case "allow_always":
144			permissions.GrantPersistent(req)
145		default:
146			permissions.Deny(req)
147		}
148	}
149}
150
151// translatePart converts a message part to an ACP session update.
152func (s *Sink) translatePart(msgID string, role message.MessageRole, part message.ContentPart) *acp.SessionUpdate {
153	switch p := part.(type) {
154	case message.TextContent:
155		return s.translateText(msgID, role, p)
156
157	case message.ReasoningContent:
158		return s.translateReasoning(msgID, p)
159
160	case message.ToolCall:
161		return s.translateToolCall(p)
162
163	case message.ToolResult:
164		return s.translateToolResult(p)
165
166	case message.Finish:
167		// Reset offsets on message finish.
168		delete(s.textOffsets, msgID)
169		delete(s.reasoningOffsets, msgID)
170		return nil
171
172	default:
173		return nil
174	}
175}
176
177func (s *Sink) translateText(msgID string, role message.MessageRole, text message.TextContent) *acp.SessionUpdate {
178	// Skip user messages - the client already knows what it sent via the
179	// prompt request.
180	if role != message.Assistant {
181		return nil
182	}
183
184	offset := s.textOffsets[msgID]
185	if len(text.Text) <= offset {
186		return nil
187	}
188
189	delta := text.Text[offset:]
190	s.textOffsets[msgID] = len(text.Text)
191
192	if delta == "" {
193		return nil
194	}
195
196	update := acp.UpdateAgentMessageText(delta)
197	return &update
198}
199
200func (s *Sink) translateReasoning(msgID string, reasoning message.ReasoningContent) *acp.SessionUpdate {
201	offset := s.reasoningOffsets[msgID]
202	if len(reasoning.Thinking) <= offset {
203		return nil
204	}
205
206	delta := reasoning.Thinking[offset:]
207	s.reasoningOffsets[msgID] = len(reasoning.Thinking)
208
209	if delta == "" {
210		return nil
211	}
212
213	update := acp.UpdateAgentThoughtText(delta)
214	return &update
215}
216
217func (s *Sink) translateToolCall(tc message.ToolCall) *acp.SessionUpdate {
218	if !tc.Finished {
219		update := acp.StartToolCall(
220			acp.ToolCallId(tc.ID),
221			tc.Name,
222			acp.WithStartStatus(acp.ToolCallStatusPending),
223		)
224		return &update
225	}
226
227	update := acp.UpdateToolCall(
228		acp.ToolCallId(tc.ID),
229		acp.WithUpdateStatus(acp.ToolCallStatusInProgress),
230	)
231	return &update
232}
233
234func (s *Sink) translateToolResult(tr message.ToolResult) *acp.SessionUpdate {
235	status := acp.ToolCallStatusCompleted
236	if tr.IsError {
237		status = acp.ToolCallStatusFailed
238	}
239
240	update := acp.UpdateToolCall(
241		acp.ToolCallId(tr.ToolCallID),
242		acp.WithUpdateStatus(status),
243		acp.WithUpdateContent([]acp.ToolCallContent{
244			acp.ToolContent(acp.TextBlock(tr.Content)),
245		}),
246	)
247	return &update
248}