Add streaming support for bot output and tool execution
- backend: StreamingBackend interface + StreamEvent type - backend/ollama: ChatStream with line-by-line JSON-LD decoder - backend/openai: ChatStream with SSE delta parser + tool arg accumulation - agent/loop: runStreamStep reads deltas, emits OutputMsg per chunk, deduplicates call/tool display between stream and outer loop - main: agentMsg tea type streams incremental updates to the TUI; plain string buf replaces strings.Builder (avoids copy-by-value panic); drainAgent cancels in-flight goroutines cleanly on user interrupt
This commit is contained in:
parent
c31d157756
commit
d60a1b9215
128
agent/loop.go
128
agent/loop.go
|
|
@ -4,72 +4,86 @@ import (
|
|||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"ollie/backend"
|
||||
)
|
||||
|
||||
// ToolExecutor runs a named tool with the given JSON arguments and returns
|
||||
// its output. Implemented by the caller using ollie/exec.
|
||||
type ToolExecutor func(name string, args json.RawMessage) (string, error)
|
||||
type OutputFn func(msg OutputMsg)
|
||||
|
||||
// OutputMsg is emitted by the loop for each visible event: assistant replies,
|
||||
// tool calls, tool results, errors, and token usage.
|
||||
type OutputMsg struct {
|
||||
Role string // "assistant" | "call" | "tool" | "error" | "usage"
|
||||
Name string // tool name, for Role=="call"/"tool"
|
||||
Role string
|
||||
Name string
|
||||
Content string
|
||||
}
|
||||
|
||||
// Config holds everything the loop needs to run.
|
||||
type Config struct {
|
||||
Backend backend.Backend
|
||||
Model string
|
||||
Tools []backend.Tool
|
||||
Exec ToolExecutor
|
||||
MaxSteps int // 0 → default 1
|
||||
Output func(msg OutputMsg) // nil → discard
|
||||
SystemPrompt string // prepended as system message if non-empty
|
||||
MaxSteps int
|
||||
Output OutputFn
|
||||
SystemPrompt string
|
||||
}
|
||||
|
||||
// Loop implements observe → decide → act → update → terminate.
|
||||
type Loop struct {
|
||||
cfg Config
|
||||
cfg Config
|
||||
lastUsage backend.Usage
|
||||
streamed bool
|
||||
skippedCalls map[string]bool
|
||||
}
|
||||
|
||||
// New creates a Loop from the given Config.
|
||||
func New(cfg Config) *Loop {
|
||||
return &Loop{cfg: cfg}
|
||||
return &Loop{cfg: cfg, skippedCalls: make(map[string]bool)}
|
||||
}
|
||||
|
||||
// Run executes the agent loop against state until the goal is complete or
|
||||
// MaxSteps is exhausted.
|
||||
func (l *Loop) Run(ctx context.Context, state State) error {
|
||||
maxSteps := l.cfg.MaxSteps
|
||||
if maxSteps <= 0 {
|
||||
maxSteps = 1
|
||||
}
|
||||
|
||||
for step := range maxSteps {
|
||||
// 1. Observe: history is the context; already up to date in state.
|
||||
sb, streaming := l.cfg.Backend.(backend.StreamingBackend)
|
||||
|
||||
// 2. Decide: call the model.
|
||||
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...)
|
||||
}
|
||||
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)
|
||||
|
||||
var resp *backend.Response
|
||||
|
||||
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)
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
// 3. Act: execute any tool calls.
|
||||
// Act: execute tool calls.
|
||||
var results []ToolResult
|
||||
for _, tc := range resp.Message.ToolCalls {
|
||||
if !l.skippedCalls[tc.Name] {
|
||||
l.emit(OutputMsg{Role: "call", Name: tc.Name, Content: string(tc.Arguments)})
|
||||
}
|
||||
|
||||
var content string
|
||||
var isErr bool
|
||||
|
||||
l.emit(OutputMsg{Role: "call", Name: tc.Name, Content: string(tc.Arguments)})
|
||||
|
||||
if l.cfg.Exec != nil {
|
||||
out, err := l.cfg.Exec(tc.Name, tc.Arguments)
|
||||
if err != nil {
|
||||
|
|
@ -93,7 +107,7 @@ func (l *Loop) Run(ctx context.Context, state State) error {
|
|||
l.emit(OutputMsg{Role: "tool", Name: tc.Name, Content: content})
|
||||
}
|
||||
|
||||
if len(resp.Message.ToolCalls) == 0 && resp.Message.Content != "" {
|
||||
if !l.streamed && resp.Message.Content != "" {
|
||||
l.emit(OutputMsg{Role: "assistant", Content: resp.Message.Content})
|
||||
}
|
||||
|
||||
|
|
@ -102,12 +116,10 @@ func (l *Loop) Run(ctx context.Context, state State) error {
|
|||
Content: fmt.Sprintf("↑%d ↓%d tokens", resp.Usage.InputTokens, resp.Usage.OutputTokens),
|
||||
})
|
||||
|
||||
// 4. Update: append assistant message and tool results to state.
|
||||
if err := state.Update(resp.Message, results); err != nil {
|
||||
return fmt.Errorf("step %d update: %w", step, err)
|
||||
}
|
||||
|
||||
// 5. Terminate?
|
||||
if l.shouldStop(resp, step, maxSteps) {
|
||||
if resp.StopReason == "stop" {
|
||||
if err := state.MarkComplete(); err != nil {
|
||||
|
|
@ -118,19 +130,59 @@ func (l *Loop) Run(ctx context.Context, state State) error {
|
|||
}
|
||||
}
|
||||
|
||||
l.skippedCalls = make(map[string]bool)
|
||||
l.streamed = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (l *Loop) shouldStop(resp *backend.Response, step, maxSteps int) bool {
|
||||
// Always stop at maxSteps.
|
||||
if step >= maxSteps-1 {
|
||||
return true
|
||||
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
|
||||
}
|
||||
// Stop when the model is done and has no pending tool calls.
|
||||
if resp.StopReason == "stop" && len(resp.Message.ToolCalls) == 0 {
|
||||
return true
|
||||
|
||||
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 {
|
||||
l.skippedCalls[tc.Name] = true
|
||||
accumulatedTcs = append(accumulatedTcs, tc)
|
||||
}
|
||||
|
||||
if ev.Done {
|
||||
l.lastUsage = ev.Usage
|
||||
msg.Content = content.String()
|
||||
msg.Role = "assistant"
|
||||
msg.ToolCalls = accumulatedTcs
|
||||
return msg, nil
|
||||
}
|
||||
}
|
||||
return false
|
||||
|
||||
l.emit(OutputMsg{Role: "error", Content: "stream ended without done event"})
|
||||
msg.Content = content.String()
|
||||
msg.Role = "assistant"
|
||||
return msg, nil
|
||||
}
|
||||
|
||||
func decideStopReason(m backend.Message) string {
|
||||
if len(m.ToolCalls) > 0 {
|
||||
return "tool_calls"
|
||||
}
|
||||
return "stop"
|
||||
}
|
||||
|
||||
func (l *Loop) emit(msg OutputMsg) {
|
||||
|
|
@ -138,3 +190,7 @@ func (l *Loop) emit(msg OutputMsg) {
|
|||
l.cfg.Output(msg)
|
||||
}
|
||||
}
|
||||
|
||||
func (l *Loop) shouldStop(resp *backend.Response, step, maxSteps int) bool {
|
||||
return resp.StopReason == "stop" || step >= maxSteps-1
|
||||
}
|
||||
|
|
|
|||
|
|
@ -43,9 +43,29 @@ type Response struct {
|
|||
Usage Usage
|
||||
}
|
||||
|
||||
// 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.
|
||||
type StreamEvent struct {
|
||||
Content string // incremental text delta (may be "")
|
||||
ToolCalls []ToolCall // complete tool calls, if any assembled this tick
|
||||
Done bool
|
||||
StopReason string // meaningful when Done==true
|
||||
Usage Usage // meaningful only when Done==true
|
||||
}
|
||||
|
||||
// Backend is the interface all LLM providers must implement.
|
||||
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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
package backend
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
|
|
@ -62,11 +63,12 @@ type ollamaChatResponse struct {
|
|||
Done bool `json:"done"`
|
||||
PromptEvalCount int `json:"prompt_eval_count"`
|
||||
EvalCount int `json:"eval_count"`
|
||||
DoneReason string `json:"done_reason"`
|
||||
}
|
||||
|
||||
// -- implementation --
|
||||
|
||||
func (b *OllamaBackend) Chat(ctx context.Context, model string, messages []Message, tools []Tool) (*Response, error) {
|
||||
func (b *OllamaBackend) doChat(ctx context.Context, model string, messages []Message, tools []Tool, stream bool) (*http.Response, error) {
|
||||
wireMessages := make([]ollamaMessage, len(messages))
|
||||
for i, m := range messages {
|
||||
wireMessages[i] = ollamaMessage{Role: m.Role, Content: m.Content}
|
||||
|
|
@ -93,7 +95,7 @@ func (b *OllamaBackend) Chat(ctx context.Context, model string, messages []Messa
|
|||
Model: model,
|
||||
Messages: wireMessages,
|
||||
Tools: wireTools,
|
||||
Stream: false,
|
||||
Stream: stream,
|
||||
}
|
||||
|
||||
data, err := json.Marshal(req)
|
||||
|
|
@ -111,13 +113,23 @@ func (b *OllamaBackend) Chat(ctx context.Context, model string, messages []Messa
|
|||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
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
|
||||
|
|
@ -132,6 +144,9 @@ func (b *OllamaBackend) Chat(ctx context.Context, model string, messages []Messa
|
|||
}
|
||||
|
||||
stopReason := "stop"
|
||||
if wire.DoneReason != "" {
|
||||
stopReason = wire.DoneReason
|
||||
}
|
||||
if len(msg.ToolCalls) > 0 {
|
||||
stopReason = "tool_calls"
|
||||
}
|
||||
|
|
@ -142,3 +157,69 @@ func (b *OllamaBackend) Chat(ctx context.Context, model string, messages []Messa
|
|||
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() {
|
||||
defer close(ch)
|
||||
defer resp.Body.Close()
|
||||
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024)
|
||||
|
||||
for scanner.Scan() {
|
||||
line := scanner.Bytes()
|
||||
if len(line) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
var wire ollamaChatResponse
|
||||
if err := json.Unmarshal(line, &wire); err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream decode: %v", err)}
|
||||
return
|
||||
}
|
||||
|
||||
if wire.Done {
|
||||
stopReason := "stop"
|
||||
if wire.DoneReason != "" {
|
||||
stopReason = wire.DoneReason
|
||||
}
|
||||
ch <- StreamEvent{
|
||||
Done: true,
|
||||
StopReason: stopReason,
|
||||
Usage: Usage{InputTokens: wire.PromptEvalCount, OutputTokens: wire.EvalCount},
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
event := StreamEvent{}
|
||||
if wire.Message.Content != "" {
|
||||
event.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.
|
||||
}
|
||||
if event.Content != "" || len(event.ToolCalls) > 0 {
|
||||
ch <- event
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream read: %v", err)}
|
||||
}
|
||||
}()
|
||||
|
||||
return ch, nil
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,12 +1,14 @@
|
|||
package backend
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// OpenAIBackend speaks the OpenAI /v1/chat/completions wire format.
|
||||
|
|
@ -27,16 +29,17 @@ func NewOpenAI(baseURL, apiKey string) *OpenAIBackend {
|
|||
// -- wire types --
|
||||
|
||||
type openAIMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content,omitempty"`
|
||||
ToolCalls []openAIToolCall `json:"tool_calls,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content,omitempty"`
|
||||
ToolCalls []openAIToolCall `json:"tool_calls,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
}
|
||||
|
||||
type openAIToolCall struct {
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"`
|
||||
Function openAIFunctionCall `json:"function"`
|
||||
ID string `json:"id"`
|
||||
Index int `json:"index,omitempty"`
|
||||
Type string `json:"type"`
|
||||
Function openAIFunctionCall `json:"function"`
|
||||
}
|
||||
|
||||
type openAIFunctionCall struct {
|
||||
|
|
@ -72,14 +75,32 @@ type openAIChatResponse struct {
|
|||
Usage openAIUsage `json:"usage"`
|
||||
}
|
||||
|
||||
type openAIDelta struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
ToolCalls []openAIToolCall `json:"tool_calls,omitempty"`
|
||||
}
|
||||
|
||||
type openAIChoice struct {
|
||||
Message openAIMessage `json:"message"`
|
||||
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"`
|
||||
FinishReason string `json:"finish_reason"`
|
||||
}
|
||||
|
||||
type openAIStreamResponse struct {
|
||||
Choices []openAIStreamChoice `json:"choices"`
|
||||
Usage openAIUsage `json:"usage"`
|
||||
}
|
||||
|
||||
// -- implementation --
|
||||
|
||||
func (b *OpenAIBackend) Chat(ctx context.Context, model string, messages []Message, tools []Tool) (*Response, error) {
|
||||
func (b *OpenAIBackend) doChat(ctx context.Context, model string, messages []Message, tools []Tool, stream bool) (*http.Response, error) {
|
||||
wireMessages := make([]openAIMessage, len(messages))
|
||||
for i, m := range messages {
|
||||
wireMessages[i] = openAIMessage{
|
||||
|
|
@ -115,7 +136,7 @@ func (b *OpenAIBackend) Chat(ctx context.Context, model string, messages []Messa
|
|||
Model: model,
|
||||
Messages: wireMessages,
|
||||
Tools: wireTools,
|
||||
Stream: false,
|
||||
Stream: stream,
|
||||
}
|
||||
|
||||
data, err := json.Marshal(req)
|
||||
|
|
@ -136,13 +157,23 @@ func (b *OpenAIBackend) Chat(ctx context.Context, model string, messages []Messa
|
|||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
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
|
||||
|
|
@ -155,7 +186,6 @@ func (b *OpenAIBackend) Chat(ctx context.Context, model string, messages []Messa
|
|||
choice := wire.Choices[0]
|
||||
msg := Message{Role: choice.Message.Role, Content: choice.Message.Content}
|
||||
for _, tc := range choice.Message.ToolCalls {
|
||||
// Arguments arrive as a JSON string; convert to RawMessage.
|
||||
args := json.RawMessage(tc.Function.Arguments)
|
||||
msg.ToolCalls = append(msg.ToolCalls, ToolCall{
|
||||
ID: tc.ID,
|
||||
|
|
@ -170,3 +200,97 @@ func (b *OpenAIBackend) Chat(ctx context.Context, model string, messages []Messa
|
|||
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() {
|
||||
defer close(ch)
|
||||
defer resp.Body.Close()
|
||||
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024)
|
||||
|
||||
// Accumulate tool call fragments by index (OpenAI sends args incrementally).
|
||||
type tcAccum struct {
|
||||
id, name, args string
|
||||
}
|
||||
accum := make(map[int]*tcAccum)
|
||||
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if line == "" || line == "data: [DONE]" {
|
||||
continue
|
||||
}
|
||||
if !strings.HasPrefix(line, "data: ") {
|
||||
continue
|
||||
}
|
||||
payload := line[6:] // strip "data: "
|
||||
|
||||
var wire openAIStreamResponse
|
||||
if err := json.Unmarshal([]byte(payload), &wire); err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream decode: %v", err)}
|
||||
return
|
||||
}
|
||||
|
||||
if len(wire.Choices) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
ev := StreamEvent{}
|
||||
choice := wire.Choices[0]
|
||||
|
||||
// Text content delta.
|
||||
if choice.Delta.Content != "" {
|
||||
ev.Content = choice.Delta.Content
|
||||
}
|
||||
|
||||
// Tool call fragments.
|
||||
for _, dtc := range choice.Delta.ToolCalls {
|
||||
idx := dtc.Index
|
||||
if _, ok := accum[idx]; !ok {
|
||||
accum[idx] = &tcAccum{id: dtc.ID, name: dtc.Function.Name}
|
||||
}
|
||||
a := accum[idx]
|
||||
if dtc.ID != "" {
|
||||
a.id = dtc.ID
|
||||
}
|
||||
if dtc.Function.Name != "" {
|
||||
a.name = dtc.Function.Name
|
||||
}
|
||||
a.args += dtc.Function.Arguments
|
||||
}
|
||||
|
||||
// If finish_reason is set, flush assembled tool calls.
|
||||
if choice.FinishReason != "" {
|
||||
for _, a := range accum {
|
||||
ev.ToolCalls = append(ev.ToolCalls, ToolCall{
|
||||
ID: a.id,
|
||||
Name: a.name,
|
||||
Arguments: json.RawMessage(a.args),
|
||||
})
|
||||
}
|
||||
ev.StopReason = choice.FinishReason
|
||||
ev.Done = true
|
||||
ev.Usage = Usage{
|
||||
InputTokens: wire.Usage.PromptTokens,
|
||||
OutputTokens: wire.Usage.CompletionTokens,
|
||||
}
|
||||
}
|
||||
|
||||
if ev.Content != "" || len(ev.ToolCalls) > 0 || ev.Done {
|
||||
ch <- ev
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("stream read: %v", err)}
|
||||
}
|
||||
}()
|
||||
|
||||
return ch, nil
|
||||
}
|
||||
|
|
|
|||
276
main.go
276
main.go
|
|
@ -1,6 +1,7 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
|
@ -37,14 +38,14 @@ Do not restate the task. Do not hedge. Do not self-congratulate.
|
|||
Use execute_code whenever the task requires:
|
||||
- running shell commands or scripts
|
||||
- reading or writing files
|
||||
- making network requests
|
||||
making network requests
|
||||
- any information you cannot reliably state from memory
|
||||
|
||||
Do not answer from memory when you can verify with execute_code.
|
||||
|
||||
Do not describe what you would run or show code blocks. Call execute_code
|
||||
directly. If you find yourself writing a markdown code block, stop and make
|
||||
the tool call instead.
|
||||
Do not describe what you would run or show code blocks.
|
||||
Call execute_code directly.
|
||||
If you find yourself writing a markdown code block, stop and make the tool call instead.
|
||||
|
||||
## execute_code — how to call it
|
||||
|
||||
|
|
@ -93,12 +94,9 @@ waiting for user input:
|
|||
Do not stop after planning. Do not ask for confirmation. Work through each
|
||||
step in sequence. Revise the plan if a step fails or reveals new information.`
|
||||
|
||||
// executeCodeTool is the single built-in tool exposed to the model.
|
||||
var executeCodeTool = backend.Tool{
|
||||
Name: "execute_code",
|
||||
Description: "Execute shell code or a named tool script in a sandboxed environment. " +
|
||||
"Use 'code' for inline bash, 'tool'+'args' for a named script, " +
|
||||
"or 'pipe' for a sequence of {tool, args} steps.",
|
||||
Name: "execute_code",
|
||||
Description: "Execute shell code or a named tool script in a sandboxed environment. Use 'code' for inline bash, 'tool'+'args' for a named script, or 'pipe' for a sequence of {tool, args} steps.",
|
||||
Parameters: json.RawMessage(`{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
|
|
@ -113,14 +111,28 @@ var executeCodeTool = backend.Tool{
|
|||
}`),
|
||||
}
|
||||
|
||||
// NOTE: buf and display lines use plain strings; bubbletea copies the model
|
||||
// by value on every Update(), so strings.Builder and other copy-sensitive
|
||||
// types will panic.
|
||||
type model struct {
|
||||
textarea textarea.Model
|
||||
viewport viewport.Model
|
||||
session *agent.Session
|
||||
loopcfg agent.Config
|
||||
display []string
|
||||
ready bool
|
||||
hooks map[string]string
|
||||
textarea textarea.Model
|
||||
viewport viewport.Model
|
||||
session *agent.Session
|
||||
loopcfg agent.Config
|
||||
display []string
|
||||
buf string // live streaming assistant text
|
||||
hooks map[string]string
|
||||
ready bool
|
||||
agentCh chan tea.Msg
|
||||
cancel context.CancelFunc
|
||||
doneCh chan struct{}
|
||||
}
|
||||
|
||||
type agentMsg struct {
|
||||
role string
|
||||
content string
|
||||
name string
|
||||
done bool
|
||||
}
|
||||
|
||||
func main() {
|
||||
|
|
@ -145,7 +157,6 @@ func main() {
|
|||
startup = append(startup, fmt.Sprintf("config: %v", err))
|
||||
}
|
||||
|
||||
// Connect MCP servers.
|
||||
mcpExec := tools.NewExecutor()
|
||||
if cfg != nil {
|
||||
for name, serverCfg := range cfg.MCPServers {
|
||||
|
|
@ -253,15 +264,96 @@ func main() {
|
|||
}
|
||||
}
|
||||
|
||||
// -- tea model --
|
||||
|
||||
func (m model) Init() tea.Cmd {
|
||||
return textarea.Blink
|
||||
}
|
||||
|
||||
type responseMsg struct {
|
||||
display []string
|
||||
session *agent.Session
|
||||
// renderDisplay produces the full viewport text. Each element of m.display
|
||||
// is considered a single line; the current streaming buffer (m.buf) is the
|
||||
// last (incomplete) line.
|
||||
func (m model) renderDisplay() string {
|
||||
var b strings.Builder
|
||||
for i, line := range m.display {
|
||||
if i > 0 {
|
||||
b.WriteByte('\n')
|
||||
}
|
||||
b.WriteString(line)
|
||||
}
|
||||
if m.buf != "" {
|
||||
if len(m.display) > 0 {
|
||||
b.WriteByte('\n')
|
||||
}
|
||||
b.WriteString("Bot: " + m.buf)
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func (m *model) refreshView() {
|
||||
m.viewport.SetContent(wordWrap(m.renderDisplay(), m.viewport.Width))
|
||||
m.viewport.GotoBottom()
|
||||
}
|
||||
|
||||
// finalizeBuf commits the in-progress streaming text as a new display line.
|
||||
func (m *model) finalizeBuf() {
|
||||
if m.buf != "" {
|
||||
m.display = append(m.display, "Bot: "+m.buf)
|
||||
m.buf = ""
|
||||
}
|
||||
}
|
||||
|
||||
// apply writes one agent output event into the model.
|
||||
func (m *model) apply(am agentMsg) {
|
||||
switch am.role {
|
||||
case "assistant":
|
||||
// Streaming text delta — accumulate in the live buffer.
|
||||
m.buf += am.content
|
||||
|
||||
case "call":
|
||||
// Tool call: finalize any pending bot text, then add a new line.
|
||||
m.finalizeBuf()
|
||||
args := squashWhitespace(am.content)
|
||||
if len(args) > 500 {
|
||||
args = args[:500] + "..."
|
||||
}
|
||||
m.display = append(m.display, fmt.Sprintf("-> %s(%s)", am.name, args))
|
||||
|
||||
case "tool":
|
||||
s := squashWhitespace(am.content)
|
||||
if len(s) > 500 {
|
||||
s = s[:500] + "..."
|
||||
}
|
||||
m.display = append(m.display, "= "+s)
|
||||
|
||||
case "error":
|
||||
m.finalizeBuf()
|
||||
m.display = append(m.display, "Error: "+am.content)
|
||||
|
||||
case "usage":
|
||||
m.finalizeBuf()
|
||||
m.display = append(m.display, "["+am.content+"]")
|
||||
}
|
||||
}
|
||||
|
||||
// drainAgent cancels the in-flight goroutine and drains remaining events
|
||||
// so the display is in a consistent state.
|
||||
func (m *model) drainAgent() {
|
||||
if m.cancel == nil {
|
||||
return
|
||||
}
|
||||
m.cancel()
|
||||
m.cancel = nil
|
||||
if m.doneCh != nil {
|
||||
<-m.doneCh // goroutine closed the channel
|
||||
m.doneCh = nil
|
||||
}
|
||||
if m.agentCh != nil {
|
||||
for msg := range m.agentCh {
|
||||
if am, ok := msg.(agentMsg); ok {
|
||||
m.apply(am)
|
||||
}
|
||||
}
|
||||
m.agentCh = nil
|
||||
}
|
||||
}
|
||||
|
||||
func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
|
|
@ -272,15 +364,18 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
|||
m.viewport.Width = msg.Width
|
||||
m.viewport.Height = msg.Height - 5
|
||||
m.textarea.SetWidth(msg.Width)
|
||||
m.viewport.SetContent(m.renderDisplay())
|
||||
m.refreshView()
|
||||
m.ready = true
|
||||
|
||||
case responseMsg:
|
||||
m.display = msg.display
|
||||
m.session = msg.session
|
||||
m.viewport.SetContent(m.renderDisplay())
|
||||
m.viewport.GotoBottom()
|
||||
return m, nil
|
||||
case agentMsg:
|
||||
m.apply(msg)
|
||||
m.refreshView()
|
||||
if msg.done {
|
||||
m.agentCh = nil
|
||||
return m, nil
|
||||
}
|
||||
// Chain to the next goroutine message.
|
||||
return m, func() tea.Msg { return <-m.agentCh }
|
||||
|
||||
case tea.KeyMsg:
|
||||
switch msg.String() {
|
||||
|
|
@ -291,16 +386,29 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
|||
if input == "" {
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// Cancel any running agent to avoid concurrent session mutation.
|
||||
m.drainAgent()
|
||||
m.finalizeBuf()
|
||||
|
||||
m.display = append(m.display, "You: "+input)
|
||||
m.viewport.SetContent(m.renderDisplay())
|
||||
m.viewport.GotoBottom()
|
||||
m.refreshView()
|
||||
m.textarea.Reset()
|
||||
|
||||
if hook := m.hooks["userPromptSubmit"]; hook != "" {
|
||||
exec.Command("sh", "-c", hook).Run()
|
||||
}
|
||||
|
||||
return m, m.runLoop(input)
|
||||
if m.session == nil {
|
||||
m.session = agent.NewSession(input)
|
||||
} else {
|
||||
m.session.AppendUserMessage(input)
|
||||
}
|
||||
ch, cancel, doneCh := m.startAgent(m.session)
|
||||
m.agentCh = ch
|
||||
m.cancel = cancel
|
||||
m.doneCh = doneCh
|
||||
return m, func() tea.Msg { return <-ch }
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -310,47 +418,46 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
|||
return m, tea.Batch(cmd, vpCmd)
|
||||
}
|
||||
|
||||
func (m model) runLoop(input string) tea.Cmd {
|
||||
// startAgent launches the loop in a goroutine, wiring its output to ch.
|
||||
// cancel stops the loop; doneCh is closed when the goroutine exits.
|
||||
func (m model) startAgent(session *agent.Session) (chan tea.Msg, context.CancelFunc, chan struct{}) {
|
||||
loopcfg := m.loopcfg
|
||||
session := m.session
|
||||
hooks := m.hooks
|
||||
display := append([]string{}, m.display...)
|
||||
ch := make(chan tea.Msg, 64)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
doneCh := make(chan struct{})
|
||||
|
||||
return func() tea.Msg {
|
||||
if session == nil {
|
||||
session = agent.NewSession(input)
|
||||
} else {
|
||||
session.AppendUserMessage(input)
|
||||
}
|
||||
go func() {
|
||||
defer close(ch)
|
||||
defer close(doneCh)
|
||||
|
||||
var lines []string
|
||||
loopcfg.Output = func(msg agent.OutputMsg) {
|
||||
switch msg.Role {
|
||||
case "assistant":
|
||||
lines = append(lines, "Bot: "+msg.Content)
|
||||
case "call":
|
||||
lines = append(lines, fmt.Sprintf("→ %s(%s)", msg.Name, msg.Content))
|
||||
case "tool":
|
||||
lines = append(lines, fmt.Sprintf(" = %s", msg.Content))
|
||||
case "usage":
|
||||
lines = append(lines, "["+msg.Content+"]")
|
||||
case "error":
|
||||
lines = append(lines, "Error: "+msg.Content)
|
||||
loopcfg.Output = func(em agent.OutputMsg) {
|
||||
select {
|
||||
case ch <- agentMsg{role: em.Role, content: em.Content, name: em.Name}:
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if err := agent.New(loopcfg).Run(context.Background(), session); err != nil {
|
||||
display = append(display, "Error: "+err.Error())
|
||||
} else {
|
||||
display = append(display, lines...)
|
||||
if err := agent.New(loopcfg).Run(ctx, session); err != nil {
|
||||
select {
|
||||
case ch <- agentMsg{role: "error", content: err.Error()}:
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if hook := hooks["stop"]; hook != "" {
|
||||
exec.Command("sh", "-c", hook).Run()
|
||||
}
|
||||
|
||||
return responseMsg{display: display, session: session}
|
||||
}
|
||||
select {
|
||||
case ch <- agentMsg{done: true}:
|
||||
case <-ctx.Done():
|
||||
}
|
||||
}()
|
||||
|
||||
return ch, cancel, doneCh
|
||||
}
|
||||
|
||||
func (m model) View() string {
|
||||
|
|
@ -360,57 +467,54 @@ func (m model) View() string {
|
|||
return m.viewport.View() + "\n" + m.textarea.View()
|
||||
}
|
||||
|
||||
// -- display helpers --
|
||||
|
||||
// renderDisplay word-wraps each display line to the viewport width and joins them.
|
||||
func (m model) renderDisplay() string {
|
||||
w := m.viewport.Width
|
||||
var buf strings.Builder
|
||||
for i, line := range m.display {
|
||||
if i > 0 {
|
||||
buf.WriteByte('\n')
|
||||
}
|
||||
buf.WriteString(wordWrap(line, w))
|
||||
// compactJSON removes whitespace between JSON tokens.
|
||||
func compactJSON(s string) string {
|
||||
var buf bytes.Buffer
|
||||
if err := json.Compact(&buf, []byte(s)); err != nil {
|
||||
return squashWhitespace(s)
|
||||
}
|
||||
return buf.String()
|
||||
return squashWhitespace(buf.String())
|
||||
}
|
||||
|
||||
// wordWrap wraps s at word boundaries so no line exceeds width columns.
|
||||
// squashWhitespace collapses all runs of whitespace into single spaces.
|
||||
func squashWhitespace(s string) string {
|
||||
return strings.Join(strings.Fields(s), " ")
|
||||
}
|
||||
|
||||
// wordWrap wraps text at word boundaries so no line exceeds width.
|
||||
func wordWrap(s string, width int) string {
|
||||
if width <= 0 {
|
||||
return s
|
||||
}
|
||||
var out strings.Builder
|
||||
for i, raw := range strings.Split(s, "\n") {
|
||||
if i > 0 {
|
||||
lines := strings.Split(s, "\n")
|
||||
for _, line := range lines {
|
||||
if out.Len() > 0 {
|
||||
out.WriteByte('\n')
|
||||
}
|
||||
col := 0
|
||||
first := true
|
||||
for _, word := range strings.Fields(raw) {
|
||||
col, first := 0, true
|
||||
for _, word := range strings.Fields(line) {
|
||||
wl := len(word)
|
||||
switch {
|
||||
case first:
|
||||
out.WriteString(word)
|
||||
col = wl
|
||||
first = false
|
||||
case col+1+wl > width:
|
||||
case col+wl+1 > width:
|
||||
out.WriteByte('\n')
|
||||
out.WriteString(word)
|
||||
col = wl
|
||||
first = true // next word goes at BOL
|
||||
default:
|
||||
out.WriteByte(' ')
|
||||
out.WriteString(word)
|
||||
col += 1 + wl
|
||||
col += wl + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
return out.String()
|
||||
}
|
||||
|
||||
// -- tool helpers --
|
||||
|
||||
// mcpToolsToBackend converts MCP tool descriptors to backend.Tool entries.
|
||||
func mcpToolsToBackend(mcpTools []tools.ToolInfo) []backend.Tool {
|
||||
out := make([]backend.Tool, len(mcpTools))
|
||||
for i, t := range mcpTools {
|
||||
|
|
@ -423,7 +527,6 @@ func mcpToolsToBackend(mcpTools []tools.ToolInfo) []backend.Tool {
|
|||
return out
|
||||
}
|
||||
|
||||
// extractMCPText unwraps {"content":[{"type":"text","text":"..."}]}.
|
||||
func extractMCPText(raw json.RawMessage) string {
|
||||
var result struct {
|
||||
Content []struct {
|
||||
|
|
@ -443,7 +546,6 @@ func extractMCPText(raw json.RawMessage) string {
|
|||
return strings.Join(parts, "\n")
|
||||
}
|
||||
|
||||
// dispatchBuiltinExec handles execute_code natively.
|
||||
func dispatchBuiltinExec(e *execpkg.Executor, args json.RawMessage) (string, error) {
|
||||
var a struct {
|
||||
Code string `json:"code"`
|
||||
|
|
@ -457,10 +559,8 @@ func dispatchBuiltinExec(e *execpkg.Executor, args json.RawMessage) (string, err
|
|||
if err := json.Unmarshal(args, &a); err != nil {
|
||||
return "", fmt.Errorf("execute_code: bad args: %w", err)
|
||||
}
|
||||
|
||||
code := a.Code
|
||||
trusted := false
|
||||
|
||||
switch {
|
||||
case len(a.Pipe) > 0:
|
||||
var err error
|
||||
|
|
@ -483,7 +583,6 @@ func dispatchBuiltinExec(e *execpkg.Executor, args json.RawMessage) (string, err
|
|||
code = fmt.Sprintf("set -- %s\n%s", strings.Join(escaped, " "), code)
|
||||
}
|
||||
}
|
||||
|
||||
if code == "" {
|
||||
return "", fmt.Errorf("execute_code: one of 'code', 'tool', or 'pipe' is required")
|
||||
}
|
||||
|
|
@ -496,6 +595,5 @@ func dispatchBuiltinExec(e *execpkg.Executor, args json.RawMessage) (string, err
|
|||
if a.Sandbox == "" {
|
||||
a.Sandbox = "default"
|
||||
}
|
||||
|
||||
return e.Execute(code, a.Language, a.Timeout, a.Sandbox, trusted)
|
||||
}
|
||||
|
|
|
|||
Reference in New Issue