diff --git a/agent/loop.go b/agent/loop.go index 8223d30..02790d7 100644 --- a/agent/loop.go +++ b/agent/loop.go @@ -30,14 +30,11 @@ type Config struct { } type Loop struct { - cfg Config - lastUsage backend.Usage - streamed bool - skippedCalls map[string]bool + cfg Config } func New(cfg Config) *Loop { - return &Loop{cfg: cfg, skippedCalls: make(map[string]bool)} + return &Loop{cfg: cfg} } func (l *Loop) Run(ctx context.Context, state State) error { @@ -46,91 +43,83 @@ func (l *Loop) Run(ctx context.Context, state State) error { maxSteps = 1 } - sb, streaming := l.cfg.Backend.(backend.StreamingBackend) - for step := range maxSteps { - l.streamed = false history := state.History() if l.cfg.SystemPrompt != "" { history = append([]backend.Message{{Role: "system", Content: l.cfg.SystemPrompt}}, history...) } - var resp *backend.Response + // Stream the assistant's response. + ch, err := l.cfg.Backend.ChatStream(ctx, l.cfg.Model, history, l.cfg.Tools) + if err != nil { + return fmt.Errorf("step %d: %w", step, err) + } - if streaming { - msg, err := l.runStreamStep(ctx, sb, l.cfg.Model, history, l.cfg.Tools) - if err != nil { - return fmt.Errorf("step %d decide: %w", step, err) + var content strings.Builder + var toolCalls []backend.ToolCall + var usage backend.Usage + var done bool + + for ev := range ch { + if ev.Content != "" { + content.WriteString(ev.Content) + l.emit(OutputMsg{Role: "assistant", Content: ev.Content}) } - resp = &backend.Response{ - Message: msg, - StopReason: decideStopReason(msg), - Usage: l.lastUsage, - } - } else { - var err error - resp, err = l.cfg.Backend.Chat(ctx, l.cfg.Model, history, l.cfg.Tools) - if err != nil { - return fmt.Errorf("step %d decide: %w", step, err) + toolCalls = append(toolCalls, ev.ToolCalls...) + if ev.Done { + usage = ev.Usage + done = true + break } } - // Act: execute tool calls. + if !done { + return fmt.Errorf("step %d: stream ended without done event", step) + } + + // Announce and execute tool calls. + msg := backend.Message{Role: "assistant", Content: content.String(), ToolCalls: toolCalls} var results []ToolResult - for _, tc := range resp.Message.ToolCalls { - // In the non-streaming path, emit the call announcement here. - // In the streaming path this was already emitted inside - // runStreamStep (once all arguments were accumulated), so skip it. - if !l.skippedCalls[tc.Name] { - l.emit(OutputMsg{Role: "call", Name: tc.Name, Content: string(tc.Arguments)}) - } - var content string + for _, tc := range toolCalls { + l.emit(OutputMsg{Role: "call", Name: tc.Name, Content: string(tc.Arguments)}) + + var result string var isErr bool - if l.cfg.Exec != nil { out, err := l.cfg.Exec(tc.Name, tc.Arguments) if err != nil { - content = fmt.Sprintf("error: %v", err) + result = fmt.Sprintf("error: %v", err) isErr = true } else { - content = out + result = out } } else { - content = "error: no tool executor configured" + result = "error: no tool executor configured" isErr = true } results = append(results, ToolResult{ ToolCallID: tc.ID, Name: tc.Name, - Content: content, + Content: result, IsError: isErr, }) - - l.emit(OutputMsg{Role: "tool", Name: tc.Name, Content: content}) + l.emit(OutputMsg{Role: "tool", Name: tc.Name, Content: result}) } - // Only emit the full assistant text if we did NOT stream it. - if !l.streamed && resp.Message.Content != "" { - l.emit(OutputMsg{Role: "assistant", Content: resp.Message.Content}) + // Emit usage when we have real token counts. + if usage.InputTokens > 0 || usage.OutputTokens > 0 { + l.emit(OutputMsg{Role: "usage", Usage: usage}) } - // Emit usage only when there are actual token counts (skips - // intermediate steps and avoids [↑0 ↓0 tokens] noise). - if resp.Usage.InputTokens > 0 || resp.Usage.OutputTokens > 0 { - l.emit(OutputMsg{ - Role: "usage", - Usage: resp.Usage, - }) - } - - if err := state.Update(resp.Message, results); err != nil { + if err := state.Update(msg, results); err != nil { return fmt.Errorf("step %d update: %w", step, err) } - if l.shouldStop(resp, step, maxSteps) { - if resp.StopReason == "stop" { + // Stop when the model has nothing more to call. + if len(toolCalls) == 0 || step >= maxSteps-1 { + if len(toolCalls) == 0 { if err := state.MarkComplete(); err != nil { return fmt.Errorf("mark complete: %w", err) } @@ -139,76 +128,11 @@ func (l *Loop) Run(ctx context.Context, state State) error { } } - l.skippedCalls = make(map[string]bool) return nil } -func (l *Loop) runStreamStep( - ctx context.Context, - sb backend.StreamingBackend, - model string, - messages []backend.Message, - tools []backend.Tool, -) (msg backend.Message, err error) { - ch, err := sb.ChatStream(ctx, model, messages, tools) - if err != nil { - return msg, err - } - - l.skippedCalls = make(map[string]bool) - l.streamed = true - var content strings.Builder - var accumulatedTcs []backend.ToolCall - - for ev := range ch { - if ev.Content != "" { - content.WriteString(ev.Content) - l.emit(OutputMsg{Role: "assistant", Content: ev.Content}) - } - - for _, tc := range ev.ToolCalls { - accumulatedTcs = append(accumulatedTcs, tc) - } - - if ev.Done { - l.lastUsage = ev.Usage - msg.Content = content.String() - msg.Role = "assistant" - msg.ToolCalls = accumulatedTcs - // Emit 'call' now that arguments are fully accumulated, and - // mark each as skipped so the Act loop does not re-emit it. - for _, tc := range accumulatedTcs { - l.emit(OutputMsg{Role: "call", Name: tc.Name, Content: string(tc.Arguments)}) - l.skippedCalls[tc.Name] = true - } - return msg, nil - } - } - - // Stream closed without a done event — treat as a transient error so the - // caller does not save partial/corrupt state. Any tool calls that arrived - // before the drop are attached to the message for diagnostic use, but - // because we return a non-nil error the outer Run loop will NOT call - // state.Update, keeping the session history clean for a retry. - msg.Content = content.String() - msg.Role = "assistant" - msg.ToolCalls = accumulatedTcs - return msg, fmt.Errorf("stream ended without done event") -} - -func decideStopReason(m backend.Message) string { - if len(m.ToolCalls) > 0 { - return "tool_calls" - } - return "stop" -} - func (l *Loop) emit(msg OutputMsg) { if l.cfg.Output != nil { l.cfg.Output(msg) } } - -func (l *Loop) shouldStop(resp *backend.Response, step, maxSteps int) bool { - return resp.StopReason == "stop" || step >= maxSteps-1 -} diff --git a/backend/backend.go b/backend/backend.go index eec22cb..ae5b7ce 100644 --- a/backend/backend.go +++ b/backend/backend.go @@ -10,9 +10,9 @@ import ( // Message is a single conversation turn. type Message struct { - Role string `json:"role"` // "system" | "user" | "assistant" | "tool" + Role string `json:"role"` // "system" | "user" | "assistant" | "tool" Content string `json:"content"` - ToolCalls []ToolCall `json:"tool_calls,omitempty"` // set by assistant when calling tools + ToolCalls []ToolCall `json:"tool_calls,omitempty"` // set by assistant when calling tools ToolCallID string `json:"tool_call_id,omitempty"` // set on role=tool replies (required by OpenAI) } @@ -36,36 +36,21 @@ type Usage struct { OutputTokens int } -// Response is the model's reply for one Chat call. -type Response struct { - Message Message - StopReason string // "stop" | "tool_calls" | "length" | ... - Usage Usage -} - -// StreamEvent is a single increment from a streaming Chat call. +// StreamEvent is a single increment from a streaming chat call. // Content is an incremental text delta (append, not replace). -// ToolCalls are complete calls ready for execution when non-empty. -// The final event has Done==true and the fully assembled Message. +// ToolCalls accumulates complete calls; they may arrive on any event. +// The final event has Done==true. type StreamEvent struct { Content string // incremental text delta (may be "") - ToolCalls []ToolCall // complete tool calls, if any assembled this tick + ToolCalls []ToolCall // complete tool calls assembled so far Done bool - StopReason string // meaningful when Done==true - Usage Usage // meaningful only when Done==true + StopReason string // meaningful when Done==true + Usage Usage // meaningful only when Done==true } // Backend is the interface all LLM providers must implement. +// Streaming is the only supported mode; backends that wrap blocking APIs +// should implement ChatStream as a single-event stream. type Backend interface { - // Chat sends messages to the model and returns its response. - // tools may be nil for plain completion requests. - Chat(ctx context.Context, model string, messages []Message, tools []Tool) (*Response, error) -} - -// StreamingBackend is an optional interface; backends that support streaming -// should implement it. Use a type assertion to check. -type StreamingBackend interface { - // ChatStream returns a channel of incremental events. The channel is - // closed after the final Done event. ChatStream(ctx context.Context, model string, messages []Message, tools []Tool) (<-chan StreamEvent, error) } diff --git a/backend/ollama.go b/backend/ollama.go index 7c2822d..bb5bbb6 100644 --- a/backend/ollama.go +++ b/backend/ollama.go @@ -68,7 +68,7 @@ type ollamaChatResponse struct { // -- implementation -- -func (b *OllamaBackend) doChat(ctx context.Context, model string, messages []Message, tools []Tool, stream bool) (*http.Response, error) { +func (b *OllamaBackend) ChatStream(ctx context.Context, model string, messages []Message, tools []Tool) (<-chan StreamEvent, error) { wireMessages := make([]ollamaMessage, len(messages)) for i, m := range messages { wireMessages[i] = ollamaMessage{Role: m.Role, Content: m.Content} @@ -91,14 +91,12 @@ func (b *OllamaBackend) doChat(ctx context.Context, model string, messages []Mes }) } - req := ollamaChatRequest{ + data, err := json.Marshal(ollamaChatRequest{ Model: model, Messages: wireMessages, Tools: wireTools, - Stream: stream, - } - - data, err := json.Marshal(req) + Stream: true, + }) if err != nil { return nil, err } @@ -113,57 +111,12 @@ func (b *OllamaBackend) doChat(ctx context.Context, model string, messages []Mes if err != nil { return nil, err } - if resp.StatusCode != http.StatusOK { body, _ := io.ReadAll(resp.Body) resp.Body.Close() return nil, fmt.Errorf("ollama HTTP %d: %s", resp.StatusCode, body) } - return resp, nil -} - -func (b *OllamaBackend) Chat(ctx context.Context, model string, messages []Message, tools []Tool) (*Response, error) { - resp, err := b.doChat(ctx, model, messages, tools, false) - if err != nil { - return nil, err - } - defer resp.Body.Close() - - var wire ollamaChatResponse - if err := json.NewDecoder(resp.Body).Decode(&wire); err != nil { - return nil, err - } - - msg := Message{Role: wire.Message.Role, Content: wire.Message.Content} - for _, tc := range wire.Message.ToolCalls { - msg.ToolCalls = append(msg.ToolCalls, ToolCall{ - Name: tc.Function.Name, - Arguments: tc.Function.Arguments, - }) - } - - stopReason := "stop" - if wire.DoneReason != "" { - stopReason = wire.DoneReason - } - if len(msg.ToolCalls) > 0 { - stopReason = "tool_calls" - } - - return &Response{ - Message: msg, - StopReason: stopReason, - Usage: Usage{InputTokens: wire.PromptEvalCount, OutputTokens: wire.EvalCount}, - }, nil -} - -func (b *OllamaBackend) ChatStream(ctx context.Context, model string, messages []Message, tools []Tool) (<-chan StreamEvent, error) { - resp, err := b.doChat(ctx, model, messages, tools, true) - if err != nil { - return nil, err - } - ch := make(chan StreamEvent, 8) go func() { @@ -198,22 +151,18 @@ func (b *OllamaBackend) ChatStream(ctx context.Context, model string, messages [ return } - event := StreamEvent{} + ev := StreamEvent{} if wire.Message.Content != "" { - event.Content = wire.Message.Content + ev.Content = wire.Message.Content } - if len(wire.Message.ToolCalls) > 0 { - for _, tc := range wire.Message.ToolCalls { - event.ToolCalls = append(event.ToolCalls, ToolCall{ - Name: tc.Function.Name, - Arguments: tc.Function.Arguments, - }) - } - // Ollama puts tool calls alongside content in the same delta. - // Emit one event with both content and tool calls. + for _, tc := range wire.Message.ToolCalls { + ev.ToolCalls = append(ev.ToolCalls, ToolCall{ + Name: tc.Function.Name, + Arguments: tc.Function.Arguments, + }) } - if event.Content != "" || len(event.ToolCalls) > 0 { - ch <- event + if ev.Content != "" || len(ev.ToolCalls) > 0 { + ch <- ev } } if err := scanner.Err(); err != nil { diff --git a/backend/openai.go b/backend/openai.go index 9043e7e..7c67d6c 100644 --- a/backend/openai.go +++ b/backend/openai.go @@ -63,10 +63,10 @@ type openAIStreamOptions struct { } type openAIChatRequest struct { - Model string `json:"model"` - Messages []openAIMessage `json:"messages"` - Tools []openAITool `json:"tools,omitempty"` - Stream bool `json:"stream"` + Model string `json:"model"` + Messages []openAIMessage `json:"messages"` + Tools []openAITool `json:"tools,omitempty"` + Stream bool `json:"stream"` StreamOptions *openAIStreamOptions `json:"stream_options,omitempty"` } @@ -75,23 +75,12 @@ type openAIUsage struct { CompletionTokens int `json:"completion_tokens"` } -type openAIChatResponse struct { - Choices []openAIChoice `json:"choices"` - Usage openAIUsage `json:"usage"` -} - type openAIDelta struct { Role string `json:"role"` Content string `json:"content"` ToolCalls []openAIToolCall `json:"tool_calls,omitempty"` } -type openAIChoice struct { - Delta openAIDelta `json:"delta,omitempty"` - Message openAIMessage `json:"message,omitempty"` - FinishReason string `json:"finish_reason"` -} - type openAIStreamChoice struct { Index int `json:"index"` Delta openAIDelta `json:"delta"` @@ -105,7 +94,7 @@ type openAIStreamResponse struct { // -- implementation -- -func (b *OpenAIBackend) doChat(ctx context.Context, model string, messages []Message, tools []Tool, stream bool) (*http.Response, error) { +func (b *OpenAIBackend) ChatStream(ctx context.Context, model string, messages []Message, tools []Tool) (<-chan StreamEvent, error) { wireMessages := make([]openAIMessage, len(messages)) for i, m := range messages { wireMessages[i] = openAIMessage{ @@ -138,16 +127,11 @@ func (b *OpenAIBackend) doChat(ctx context.Context, model string, messages []Mes } req := openAIChatRequest{ - Model: model, - Messages: wireMessages, - Tools: wireTools, - Stream: stream, - } - - // Ask OpenAI to include token usage in the stream. Without this flag - // the API omits usage entirely from SSE chunks. - if stream { - req.StreamOptions = &openAIStreamOptions{IncludeUsage: true} + Model: model, + Messages: wireMessages, + Tools: wireTools, + Stream: true, + StreamOptions: &openAIStreamOptions{IncludeUsage: true}, } data, err := json.Marshal(req) @@ -168,56 +152,12 @@ func (b *OpenAIBackend) doChat(ctx context.Context, model string, messages []Mes if err != nil { return nil, err } - if resp.StatusCode != http.StatusOK { body, _ := io.ReadAll(resp.Body) resp.Body.Close() return nil, fmt.Errorf("openai HTTP %d: %s", resp.StatusCode, body) } - return resp, nil -} - -func (b *OpenAIBackend) Chat(ctx context.Context, model string, messages []Message, tools []Tool) (*Response, error) { - resp, err := b.doChat(ctx, model, messages, tools, false) - if err != nil { - return nil, err - } - defer resp.Body.Close() - - var wire openAIChatResponse - if err := json.NewDecoder(resp.Body).Decode(&wire); err != nil { - return nil, err - } - - if len(wire.Choices) == 0 { - return nil, fmt.Errorf("openai: empty choices in response") - } - - choice := wire.Choices[0] - msg := Message{Role: choice.Message.Role, Content: choice.Message.Content} - for _, tc := range choice.Message.ToolCalls { - args := json.RawMessage(tc.Function.Arguments) - msg.ToolCalls = append(msg.ToolCalls, ToolCall{ - ID: tc.ID, - Name: tc.Function.Name, - Arguments: args, - }) - } - - return &Response{ - Message: msg, - StopReason: choice.FinishReason, - Usage: Usage{InputTokens: wire.Usage.PromptTokens, OutputTokens: wire.Usage.CompletionTokens}, - }, nil -} - -func (b *OpenAIBackend) ChatStream(ctx context.Context, model string, messages []Message, tools []Tool) (<-chan StreamEvent, error) { - resp, err := b.doChat(ctx, model, messages, tools, true) - if err != nil { - return nil, err - } - ch := make(chan StreamEvent, 8) go func() { @@ -233,8 +173,7 @@ func (b *OpenAIBackend) ChatStream(ctx context.Context, model string, messages [ } accum := make(map[int]*tcAccum) - // finishReason and accumulated usage are tracked separately because - // OpenAI sends usage in a trailing chunk *after* the finish_reason chunk. + // finishReason and usage arrive in separate trailing chunks. var finishReason string var usage Usage @@ -244,8 +183,6 @@ func (b *OpenAIBackend) ChatStream(ctx context.Context, model string, messages [ continue } if line == "data: [DONE]" { - // Stream finished. Emit the Done event with whatever usage - // we have accumulated (may arrive before or after this sentinel). ev := StreamEvent{ Done: true, StopReason: finishReason, @@ -264,16 +201,13 @@ func (b *OpenAIBackend) ChatStream(ctx context.Context, model string, messages [ if !strings.HasPrefix(line, "data: ") { continue } - payload := line[6:] // strip "data: " var wire openAIStreamResponse - if err := json.Unmarshal([]byte(payload), &wire); err != nil { + if err := json.Unmarshal([]byte(line[6:]), &wire); err != nil { ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream decode: %v", err)} return } - // Accumulate usage from every chunk; it will be non-zero only on - // the trailing usage chunk that OpenAI sends after finish_reason. if wire.Usage.PromptTokens > 0 { usage.InputTokens = wire.Usage.PromptTokens } @@ -282,23 +216,17 @@ func (b *OpenAIBackend) ChatStream(ctx context.Context, model string, messages [ } if len(wire.Choices) == 0 { - // Choices-less chunk (e.g. the trailing usage-only chunk). continue } choice := wire.Choices[0] - - // Record finish reason when it arrives. if choice.FinishReason != "" { finishReason = choice.FinishReason } - - // Text content delta. if choice.Delta.Content != "" { ch <- StreamEvent{Content: choice.Delta.Content} } - // Tool call fragments — accumulate, do not emit yet. for _, dtc := range choice.Delta.ToolCalls { idx := dtc.Index if _, ok := accum[idx]; !ok {