agent.go

  1package acp
  2
  3import (
  4	"context"
  5	"errors"
  6	"fmt"
  7	"log/slog"
  8	"strings"
  9
 10	"charm.land/fantasy"
 11	"github.com/charmbracelet/crush/internal/app"
 12	"github.com/charmbracelet/crush/internal/config"
 13	"github.com/charmbracelet/crush/internal/csync"
 14	"github.com/charmbracelet/crush/internal/message"
 15	"github.com/charmbracelet/crush/internal/permission"
 16	"github.com/coder/acp-go-sdk"
 17)
 18
 19// Agent implements the acp.Agent interface to handle ACP protocol methods.
 20type Agent struct {
 21	app   *app.App
 22	conn  *acp.AgentSideConnection
 23	sinks *csync.Map[string, *Sink]
 24}
 25
 26// Compile-time interface checks.
 27var (
 28	_ acp.Agent             = (*Agent)(nil)
 29	_ acp.AgentLoader       = (*Agent)(nil)
 30	_ acp.AgentExperimental = (*Agent)(nil)
 31)
 32
 33// NewAgent creates a new ACP agent backed by a Crush app instance.
 34func NewAgent(app *app.App) *Agent {
 35	return &Agent{
 36		app:   app,
 37		sinks: csync.NewMap[string, *Sink](),
 38	}
 39}
 40
 41// SetAgentConnection stores the connection for sending notifications.
 42func (a *Agent) SetAgentConnection(conn *acp.AgentSideConnection) {
 43	a.conn = conn
 44}
 45
 46// Initialize handles the ACP initialize request.
 47func (a *Agent) Initialize(ctx context.Context, params acp.InitializeRequest) (acp.InitializeResponse, error) {
 48	slog.Debug("ACP Initialize", "protocol_version", params.ProtocolVersion)
 49	return acp.InitializeResponse{
 50		ProtocolVersion: acp.ProtocolVersionNumber,
 51		AgentCapabilities: acp.AgentCapabilities{
 52			LoadSession: true,
 53			McpCapabilities: acp.McpCapabilities{
 54				Http: false,
 55				Sse:  false,
 56			},
 57			PromptCapabilities: acp.PromptCapabilities{
 58				EmbeddedContext: true,
 59				Audio:           false,
 60				Image:           false,
 61			},
 62		},
 63	}, nil
 64}
 65
 66// Authenticate handles authentication requests (stub for local stdio).
 67func (a *Agent) Authenticate(ctx context.Context, params acp.AuthenticateRequest) (acp.AuthenticateResponse, error) {
 68	slog.Debug("ACP Authenticate")
 69	return acp.AuthenticateResponse{}, nil
 70}
 71
 72// NewSession creates a new Crush session.
 73func (a *Agent) NewSession(ctx context.Context, params acp.NewSessionRequest) (acp.NewSessionResponse, error) {
 74	slog.Info("ACP NewSession", "cwd", params.Cwd)
 75
 76	sess, err := a.app.Sessions.Create(ctx, "ACP Session")
 77	if err != nil {
 78		return acp.NewSessionResponse{}, err
 79	}
 80
 81	// Create and start the event sink to stream updates to this session.
 82	// Use a background context since the sink needs to outlive the NewSession
 83	// request.
 84	sink := NewSink(context.Background(), a.conn, sess.ID)
 85	sink.Start(a.app.Messages, a.app.Permissions, a.app.Sessions)
 86	a.sinks.Set(sess.ID, sink)
 87
 88	return acp.NewSessionResponse{
 89		SessionId: acp.SessionId(sess.ID),
 90		Models:    a.buildSessionModelState(),
 91	}, nil
 92}
 93
 94// LoadSession loads an existing session to resume a previous conversation.
 95func (a *Agent) LoadSession(ctx context.Context, params acp.LoadSessionRequest) (acp.LoadSessionResponse, error) {
 96	sessionID := string(params.SessionId)
 97	slog.Info("ACP LoadSession", "session_id", sessionID)
 98
 99	// Verify the session exists.
100	session, err := a.app.Sessions.Get(ctx, sessionID)
101	if err != nil {
102		return acp.LoadSessionResponse{}, err
103	}
104
105	// Create and start the event sink for future updates.
106	sink := NewSink(context.Background(), a.conn, session.ID)
107	sink.Start(a.app.Messages, a.app.Permissions, a.app.Sessions)
108	a.sinks.Set(session.ID, sink)
109
110	// Load and replay historical messages to the client.
111	messages, err := a.app.Messages.List(ctx, sessionID)
112	if err != nil {
113		return acp.LoadSessionResponse{}, err
114	}
115
116	for _, msg := range messages {
117		if err := a.replayMessage(ctx, sessionID, msg); err != nil {
118			slog.Error("Failed to replay message", "message_id", msg.ID, "error", err)
119		}
120	}
121
122	return acp.LoadSessionResponse{
123		Models: a.buildSessionModelState(),
124	}, nil
125}
126
127// SetSessionMode handles mode switching (stub - Crush doesn't have modes yet).
128func (a *Agent) SetSessionMode(ctx context.Context, params acp.SetSessionModeRequest) (acp.SetSessionModeResponse, error) {
129	slog.Debug("ACP SetSessionMode", "mode_id", params.ModeId)
130	return acp.SetSessionModeResponse{}, nil
131}
132
133// SetSessionModel handles model switching by parsing the model ID and updating
134// the agent's active model.
135func (a *Agent) SetSessionModel(ctx context.Context, params acp.SetSessionModelRequest) (acp.SetSessionModelResponse, error) {
136	slog.Info("ACP SetSessionModel", "session_id", params.SessionId, "model_id", params.ModelId)
137
138	// Parse model ID (format: "provider:model").
139	parts := strings.SplitN(string(params.ModelId), ":", 2)
140	if len(parts) != 2 || parts[0] == "" || parts[1] == "" {
141		return acp.SetSessionModelResponse{}, fmt.Errorf("invalid model ID format %q: expected provider:model", params.ModelId)
142	}
143	providerID, modelID := parts[0], parts[1]
144
145	// Validate that the model exists.
146	cfg := config.Get()
147	if cfg.GetModel(providerID, modelID) == nil {
148		return acp.SetSessionModelResponse{}, fmt.Errorf("model %q not found for provider %q", modelID, providerID)
149	}
150
151	// Check if the agent is busy.
152	if a.app.AgentCoordinator.IsBusy() {
153		return acp.SetSessionModelResponse{}, fmt.Errorf("agent is busy, cannot switch models")
154	}
155
156	// Update the preferred model in config.
157	selectedModel := config.SelectedModel{
158		Provider: providerID,
159		Model:    modelID,
160	}
161	if err := cfg.UpdatePreferredModel(config.SelectedModelTypeLarge, selectedModel); err != nil {
162		return acp.SetSessionModelResponse{}, fmt.Errorf("failed to update preferred model: %w", err)
163	}
164
165	// Apply the model change to the agent.
166	if err := a.app.UpdateAgentModel(ctx); err != nil {
167		return acp.SetSessionModelResponse{}, fmt.Errorf("failed to apply model change: %w", err)
168	}
169
170	slog.Info("ACP SetSessionModel completed", "provider", providerID, "model", modelID)
171	return acp.SetSessionModelResponse{}, nil
172}
173
174// Prompt handles a prompt request by running the agent.
175func (a *Agent) Prompt(ctx context.Context, params acp.PromptRequest) (acp.PromptResponse, error) {
176	slog.Info("ACP Prompt", "session_id", params.SessionId)
177
178	// Extract text from content blocks.
179	var prompt string
180	for _, block := range params.Prompt {
181		if block.Text != nil {
182			prompt += block.Text.Text
183		}
184	}
185
186	if prompt == "" {
187		return acp.PromptResponse{StopReason: acp.StopReasonEndTurn}, nil
188	}
189
190	// Run the agent.
191	result, err := a.app.AgentCoordinator.Run(ctx, string(params.SessionId), prompt)
192	if err != nil {
193		// Permission denial is a normal user choice, not an error.
194		if errors.Is(err, permission.ErrorPermissionDenied) {
195			return acp.PromptResponse{StopReason: acp.StopReasonRefusal}, nil
196		}
197		// Context cancellation means the user cancelled the request.
198		if errors.Is(err, context.Canceled) {
199			return acp.PromptResponse{StopReason: acp.StopReasonCancelled}, nil
200		}
201		// Other errors are actual errors.
202		return acp.PromptResponse{StopReason: acp.StopReasonEndTurn}, err
203	}
204
205	// Map the agent's finish reason to an ACP stop reason.
206	if result != nil && result.Response.FinishReason == fantasy.FinishReasonLength {
207		return acp.PromptResponse{StopReason: acp.StopReasonMaxTokens}, nil
208	}
209
210	return acp.PromptResponse{StopReason: acp.StopReasonEndTurn}, nil
211}
212
213// Cancel handles cancellation of an in-flight prompt.
214func (a *Agent) Cancel(ctx context.Context, params acp.CancelNotification) error {
215	slog.Info("ACP Cancel", "session_id", params.SessionId)
216	a.app.AgentCoordinator.Cancel(string(params.SessionId))
217	return nil
218}
219
220// replayMessage sends a historical message to the client via session updates.
221func (a *Agent) replayMessage(ctx context.Context, sessionID string, msg message.Message) error {
222	for _, part := range msg.Parts {
223		update := a.translateHistoryPart(msg.Role, part)
224		if update == nil {
225			continue
226		}
227
228		if err := a.conn.SessionUpdate(ctx, acp.SessionNotification{
229			SessionId: acp.SessionId(sessionID),
230			Update:    *update,
231		}); err != nil {
232			return err
233		}
234	}
235	return nil
236}
237
238// translateHistoryPart converts a message part to an ACP session update for
239// history replay. Unlike streaming updates, this sends full content rather
240// than deltas.
241func (a *Agent) translateHistoryPart(role message.MessageRole, part message.ContentPart) *acp.SessionUpdate {
242	switch p := part.(type) {
243	case message.TextContent:
244		if p.Text == "" {
245			return nil
246		}
247		var update acp.SessionUpdate
248		if role == message.User {
249			update = acp.UpdateUserMessageText(p.Text)
250		} else {
251			update = acp.UpdateAgentMessageText(p.Text)
252		}
253		return &update
254
255	case message.ReasoningContent:
256		if p.Thinking == "" {
257			return nil
258		}
259		update := acp.UpdateAgentThoughtText(p.Thinking)
260		return &update
261
262	case message.ToolCall:
263		// For history replay, send the tool call as completed with full input.
264		opts := []acp.ToolCallStartOpt{
265			acp.WithStartStatus(acp.ToolCallStatusCompleted),
266			acp.WithStartKind(toolKind(p.Name)),
267		}
268		if input := parseToolInput(p.Input); input != nil {
269			if input.Path != "" {
270				opts = append(opts, acp.WithStartLocations([]acp.ToolCallLocation{{Path: input.Path}}))
271			}
272			opts = append(opts, acp.WithStartRawInput(input.Raw))
273		}
274		title := p.Name
275		if input := parseToolInput(p.Input); input != nil && input.Title != "" {
276			title = input.Title
277		}
278		update := acp.StartToolCall(acp.ToolCallId(p.ID), title, opts...)
279		return &update
280
281	case message.ToolResult:
282		status := acp.ToolCallStatusCompleted
283		if p.IsError {
284			status = acp.ToolCallStatusFailed
285		}
286		content := []acp.ToolCallContent{acp.ToolContent(acp.TextBlock(p.Content))}
287		update := acp.UpdateToolCall(
288			acp.ToolCallId(p.ToolCallID),
289			acp.WithUpdateStatus(status),
290			acp.WithUpdateContent(content),
291		)
292		return &update
293
294	default:
295		return nil
296	}
297}
298
299// buildSessionModelState constructs the model state for session responses,
300// listing all available models and the currently selected one.
301func (a *Agent) buildSessionModelState() *acp.SessionModelState {
302	cfg := config.Get()
303	if cfg == nil {
304		return nil
305	}
306
307	var availableModels []acp.ModelInfo
308	for providerID, providerConfig := range cfg.Providers.Seq2() {
309		if providerConfig.Disable {
310			continue
311		}
312		providerName := providerConfig.Name
313		if providerName == "" {
314			providerName = providerID
315		}
316		for _, model := range providerConfig.Models {
317			modelID := acp.ModelId(providerID + ":" + model.ID)
318			modelName := model.Name
319			if modelName == "" {
320				modelName = model.ID
321			}
322			availableModels = append(availableModels, acp.ModelInfo{
323				ModelId: modelID,
324				Name:    providerName + " / " + modelName,
325			})
326		}
327	}
328
329	// Get current model.
330	currentModel := cfg.Models[config.SelectedModelTypeLarge]
331	currentModelID := acp.ModelId(currentModel.Provider + ":" + currentModel.Model)
332
333	return &acp.SessionModelState{
334		AvailableModels: availableModels,
335		CurrentModelId:  currentModelID,
336	}
337}