Compare commits
7 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 47671f3098 | |||
| 28195519d1 | |||
| 425f6ee5ee | |||
| c8d8329dbb | |||
| dabf5bfecc | |||
| 9abd5c3b5f | |||
| 3d403bd685 |
@@ -8,6 +8,7 @@ model/provider switching, and local tools.
|
||||
|
||||
- TUI-first chat experience with a light theme by default.
|
||||
- OpenAI-compatible `/v1/chat/completions` provider support.
|
||||
- OpenAI Responses API (`/v1/responses`) provider support.
|
||||
- Multiple providers from `~/.agentu/config.yaml`.
|
||||
- Session-only provider, model, and thinking-level switching.
|
||||
- Read-only project tools enabled by default.
|
||||
@@ -23,6 +24,12 @@ export AGENTU_API_KEY='your-api-key'
|
||||
go run .
|
||||
```
|
||||
|
||||
Use dark mode when your terminal is dark:
|
||||
|
||||
```sh
|
||||
go run . --theme dark
|
||||
```
|
||||
|
||||
Use `--yolo` only when you want agentu to write files or run shell commands:
|
||||
|
||||
```sh
|
||||
@@ -42,6 +49,7 @@ Use `--config <path>` to load another file.
|
||||
```yaml
|
||||
providers:
|
||||
loveuer:
|
||||
api_type: chat
|
||||
base_url: https://ai.loveuer.com
|
||||
api_key: ${AGENTU_API_KEY}
|
||||
models:
|
||||
@@ -58,6 +66,20 @@ Provider names are the keys under `providers`. At startup, agentu selects the
|
||||
first provider by name. Use `/model provider <name>` to switch providers for the
|
||||
current session.
|
||||
|
||||
`api_type` selects which API protocol the provider speaks:
|
||||
|
||||
```text
|
||||
chat OpenAI-compatible /v1/chat/completions (default)
|
||||
responses OpenAI Responses API /v1/responses
|
||||
```
|
||||
|
||||
The Responses API client converts the conversation history to Responses input
|
||||
items (`user`/`assistant` messages, `function_call`, and
|
||||
`function_call_output`), sends the system prompt as the top-level
|
||||
`instructions` field, streams `response.output_text.delta` events, and supports
|
||||
function calling through `response.output_item.added`,
|
||||
`response.function_call_arguments.delta`, and `response.output_item.done`.
|
||||
|
||||
`models` is required. `/model model <name>` is restricted to that list. Each
|
||||
entry is an object with:
|
||||
|
||||
@@ -77,7 +99,11 @@ none, middle, high, xhigh, max
|
||||
```
|
||||
|
||||
When configured or changed with `/model thinking ...`, the value is sent as a
|
||||
top-level OpenAI-compatible request field named `thinking`.
|
||||
top-level request field. The default field depends on `api_type`: `chat` uses
|
||||
`thinking`, while `responses` uses the official `reasoning` field, sent as
|
||||
`{"effort": ...}` with agentu's `middle` level mapped to the official `medium`
|
||||
value (`none`, `high`, `xhigh`, `max` pass through). Override with
|
||||
`thinking_param` if a provider expects a different field name.
|
||||
|
||||
## CLI Flags
|
||||
|
||||
@@ -94,22 +120,25 @@ go run . [flags]
|
||||
## TUI Controls
|
||||
|
||||
- `enter` sends the message.
|
||||
- `ctrl+j` or `alt+enter` inserts a newline.
|
||||
- `ctrl+j` inserts a newline.
|
||||
- `up` recalls the previous input; `ctrl+p` / `ctrl+n` navigate input history.
|
||||
- The input starts at one line and grows up to six lines.
|
||||
- `page up` / `page down` scrolls the transcript.
|
||||
- `esc` cancels a running response.
|
||||
- Type `/` to show available commands.
|
||||
- In `--yolo` mode, type `! <command>` to run a shell command from the configured working directory.
|
||||
|
||||
Slash commands:
|
||||
|
||||
- `/clear` resets conversation history.
|
||||
- `/compact` summarizes older conversation history and keeps recent context.
|
||||
- `/compact` summarizes older conversation history and keeps the newest ~20% of context, with a small recent-message floor.
|
||||
- `/status` shows changed files with `git status --short`.
|
||||
- `/diff [--stat] [path...]` shows the unstaged git diff.
|
||||
- `/test [command...]` runs a project test command; defaults to `go test ./...` for Go modules or `npm test` for Node projects.
|
||||
- `/model` shows current provider, model, thinking, and switch usage.
|
||||
- `/model provider <name>` switches provider for the current session.
|
||||
- `/model model <name>` or `/model <name>` switches model for the current session.
|
||||
In the TUI, type `/model ` or `/model model ` to choose from configured model IDs in the candidate palette.
|
||||
- `/model thinking <none|middle|high|xhigh|max>` changes thinking for the current session.
|
||||
- `/sessions` lists saved sessions.
|
||||
- `/resume [id]` resumes by ID or opens an interactive picker.
|
||||
@@ -125,6 +154,9 @@ Read-only tools are available by default:
|
||||
- `file_list`
|
||||
- `file_search` with optional `context_lines` and `max_results`
|
||||
- `code_symbols` for Go packages, imports, types, functions, and methods
|
||||
- `git_status` shows changed files
|
||||
- `git_diff` shows unstaged or staged diff output
|
||||
- `git_log` shows recent commit history
|
||||
|
||||
Web tools are also available in read-only mode:
|
||||
|
||||
@@ -133,8 +165,10 @@ Web tools are also available in read-only mode:
|
||||
|
||||
`--yolo` additionally enables:
|
||||
|
||||
- `file_write`
|
||||
- `shell_run`
|
||||
- `file_write` for creating or replacing whole files
|
||||
- `file_edit` for surgical line-range or pattern edits to existing files
|
||||
- `shell_run` for non-interactive shell commands; output streams live into the
|
||||
TUI while the command runs instead of appearing only after completion
|
||||
|
||||
Use `--yolo` for prompts that ask agentu to run or inspect local commands, for
|
||||
example:
|
||||
@@ -143,6 +177,8 @@ example:
|
||||
请帮我运行 `go test ./...` 并总结结果
|
||||
```
|
||||
|
||||
|
||||
|
||||
## Development
|
||||
|
||||
Run the verification suite:
|
||||
@@ -159,5 +195,8 @@ Local config and build outputs are intentionally ignored by git.
|
||||
|
||||
- Provider/model/thinking changes are session-only; agentu does not write them
|
||||
back to `~/.agentu/config.yaml`.
|
||||
- `api_type` is per provider and cannot be changed at runtime; use
|
||||
`/model provider <name>` to switch to a provider configured with a different
|
||||
API type.
|
||||
- The config schema is intentionally strict during MVP development. Deprecated
|
||||
fields such as `provider` or `active_provider` are rejected.
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
providers:
|
||||
loveuer:
|
||||
# API protocol: "chat" (OpenAI-compatible /v1/chat/completions, default)
|
||||
# or "responses" (OpenAI Responses API /v1/responses).
|
||||
api_type: chat
|
||||
base_url: https://ai.loveuer.com
|
||||
api_key: ${AGENTU_API_KEY}
|
||||
models:
|
||||
|
||||
+197
-22
@@ -7,18 +7,24 @@ import (
|
||||
"io"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"agentu/pkg/llm"
|
||||
"agentu/internal/tools"
|
||||
"agentu/pkg/llm"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultMaxToolRounds = 50
|
||||
doomLoopThreshold = 3
|
||||
compactRecentMessages = 6
|
||||
compactMinRecentMessages = 6
|
||||
compactRecentRetentionPercent = 20
|
||||
summaryMessageLimit = 8000
|
||||
contextSummaryPrefix = "[context summary]"
|
||||
toolLogArgumentLimit = 400
|
||||
toolLogOutputPreviewBytes = 4096
|
||||
toolLogOutputMarker = "[output]"
|
||||
toolStreamLogLimit = tools.MaxToolOutputBytes
|
||||
|
||||
summarizerSystemPrompt = "You are a conversation summarizer. Summarize the following conversation history into a concise but detailed summary. Preserve:\n" +
|
||||
"- Key facts, decisions, and conclusions\n" +
|
||||
@@ -29,6 +35,11 @@ const (
|
||||
"Output only the summary, no preamble."
|
||||
)
|
||||
|
||||
type toolExecutionResult struct {
|
||||
call llm.ToolCall
|
||||
output string
|
||||
}
|
||||
|
||||
type Agent struct {
|
||||
provider llm.Provider
|
||||
providerName string
|
||||
@@ -173,6 +184,61 @@ func (a *Agent) SetLastUsage(usage *llm.Usage) {
|
||||
a.lastUsage = ©
|
||||
}
|
||||
|
||||
// compactRetainedTokenBudget returns the number of estimated tokens to retain
|
||||
// based on compactRecentRetentionPercent, using integer ceiling division.
|
||||
func compactRetainedTokenBudget(totalTokens int) int {
|
||||
if totalTokens <= 0 {
|
||||
return 1
|
||||
}
|
||||
return (totalTokens*compactRecentRetentionPercent + 99) / 100
|
||||
}
|
||||
|
||||
// compactTurnBoundary walks backward from start until it finds a user message
|
||||
// or reaches prefixEnd, ensuring the retained suffix starts at a user turn.
|
||||
func compactTurnBoundary(messages []llm.Message, prefixEnd int, start int) int {
|
||||
for start > prefixEnd && messages[start].Role != llm.RoleUser {
|
||||
start--
|
||||
}
|
||||
return start
|
||||
}
|
||||
|
||||
// compactRecentStart returns the index where the retained recent suffix begins.
|
||||
// It balances the 20% token-budget ratio with the compactMinRecentMessages floor
|
||||
// and aligns to a user-turn boundary.
|
||||
func compactRecentStart(messages []llm.Message, prefixEnd int) int {
|
||||
compactableCount := len(messages) - prefixEnd
|
||||
if compactableCount <= compactMinRecentMessages {
|
||||
return len(messages)
|
||||
}
|
||||
|
||||
compactable := messages[prefixEnd:]
|
||||
totalTokens := llm.EstimateTokens(compactable, nil)
|
||||
retainTokens := compactRetainedTokenBudget(totalTokens)
|
||||
|
||||
// Walk backward from the end to find the earliest suffix fitting the budget.
|
||||
ratioStart := len(messages) - 1
|
||||
for i := len(messages) - 1; i >= prefixEnd; i-- {
|
||||
if llm.EstimateTokens(messages[i:], nil) <= retainTokens {
|
||||
ratioStart = i
|
||||
} else {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// Enforce the minimum recent-message floor.
|
||||
minStart := len(messages) - compactMinRecentMessages
|
||||
if minStart < prefixEnd {
|
||||
minStart = prefixEnd
|
||||
}
|
||||
// Choose the earlier start so the floor is never weakened.
|
||||
start := minStart
|
||||
if ratioStart < start {
|
||||
start = ratioStart
|
||||
}
|
||||
|
||||
return compactTurnBoundary(messages, prefixEnd, start)
|
||||
}
|
||||
|
||||
// Compact summarizes older conversation messages while preserving the system
|
||||
// prompt and recent messages intact.
|
||||
func (a *Agent) Compact(ctx context.Context) error {
|
||||
@@ -187,12 +253,8 @@ func (a *Agent) Compact(ctx context.Context) error {
|
||||
if a.messages[0].Role == llm.RoleSystem {
|
||||
prefixEnd = 1
|
||||
}
|
||||
if len(a.messages)-prefixEnd <= compactRecentMessages {
|
||||
return nil
|
||||
}
|
||||
|
||||
recentStart := len(a.messages) - compactRecentMessages
|
||||
if recentStart <= prefixEnd {
|
||||
recentStart := compactRecentStart(a.messages, prefixEnd)
|
||||
if recentStart <= prefixEnd || recentStart >= len(a.messages) {
|
||||
return nil
|
||||
}
|
||||
middle := append([]llm.Message(nil), a.messages[prefixEnd:recentStart]...)
|
||||
@@ -364,14 +426,15 @@ func (a *Agent) RunTurn(ctx context.Context, input string, out io.Writer, logs i
|
||||
}
|
||||
}
|
||||
|
||||
for _, call := range calls {
|
||||
result := a.executeTool(ctx, call, logs)
|
||||
results := a.executeTools(ctx, calls, logs)
|
||||
for _, result := range results {
|
||||
a.messages = append(a.messages, llm.Message{
|
||||
Role: llm.RoleTool,
|
||||
ToolCallID: call.ID,
|
||||
Content: result,
|
||||
ToolCallID: result.call.ID,
|
||||
Content: result.output,
|
||||
})
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
return fmt.Errorf("tool round limit reached after %d rounds; the task may be too complex for a single turn", a.maxToolRounds)
|
||||
@@ -433,39 +496,151 @@ func (a *Agent) toolDefinitions() []llm.Tool {
|
||||
return a.toolRegistry.Definitions()
|
||||
}
|
||||
|
||||
func (a *Agent) executeTool(ctx context.Context, call llm.ToolCall, logs io.Writer) string {
|
||||
func (a *Agent) executeTool(ctx context.Context, call llm.ToolCall, logs io.Writer) (string, bool) {
|
||||
name := call.Function.Name
|
||||
if logs != nil {
|
||||
fmt.Fprintf(logs, "\n[tool] %s %s\n", name, compact(call.Function.Arguments, 400))
|
||||
}
|
||||
|
||||
if a.toolRegistry == nil {
|
||||
return tools.FormatError(fmt.Errorf("tool registry is disabled"))
|
||||
return tools.FormatError(fmt.Errorf("tool registry is disabled")), false
|
||||
}
|
||||
tool, ok := a.toolRegistry.Get(name)
|
||||
if !ok {
|
||||
return tools.FormatError(fmt.Errorf("unknown tool: %s", name))
|
||||
return tools.FormatError(fmt.Errorf("unknown tool: %s", name)), false
|
||||
}
|
||||
|
||||
args := strings.TrimSpace(call.Function.Arguments)
|
||||
if args == "" {
|
||||
args = "{}"
|
||||
}
|
||||
raw := json.RawMessage(args)
|
||||
|
||||
toolCtx, cancel := context.WithTimeout(ctx, a.toolTimeout)
|
||||
defer cancel()
|
||||
output, err := tool.Execute(toolCtx, json.RawMessage(args))
|
||||
|
||||
if streaming, ok := tool.(tools.StreamingTool); ok {
|
||||
return a.executeStreamingTool(streaming, toolCtx, raw, name, args, logs)
|
||||
}
|
||||
|
||||
output, err := tool.Execute(toolCtx, raw)
|
||||
if err != nil {
|
||||
if output != "" {
|
||||
return output + "\n" + tools.FormatError(err)
|
||||
return output + "\n" + tools.FormatError(err), false
|
||||
}
|
||||
return tools.FormatError(err)
|
||||
return tools.FormatError(err), false
|
||||
}
|
||||
return output, false
|
||||
}
|
||||
|
||||
func (a *Agent) executeStreamingTool(streaming tools.StreamingTool, ctx context.Context, raw json.RawMessage, name, args string, logs io.Writer) (string, bool) {
|
||||
var streamed int
|
||||
emit := func(chunk string) error {
|
||||
if chunk == "" || logs == nil {
|
||||
return nil
|
||||
}
|
||||
if streamed >= toolStreamLogLimit {
|
||||
return nil
|
||||
}
|
||||
if len(chunk) > toolStreamLogLimit-streamed {
|
||||
chunk = chunk[:toolStreamLogLimit-streamed]
|
||||
streamed = toolStreamLogLimit
|
||||
_, err := fmt.Fprintf(logs, "\n[tool] %s %s\n%s\n%s\n[streaming output truncated after %d bytes]", name, compact(args, toolLogArgumentLimit), toolLogOutputMarker, chunk, toolStreamLogLimit)
|
||||
return err
|
||||
}
|
||||
streamed += len(chunk)
|
||||
_, err := fmt.Fprintf(logs, "\n[tool] %s %s\n%s\n%s", name, compact(args, toolLogArgumentLimit), toolLogOutputMarker, chunk)
|
||||
return err
|
||||
}
|
||||
|
||||
output, err := streaming.ExecuteStream(ctx, raw, emit)
|
||||
if logs != nil {
|
||||
fmt.Fprintf(logs, "[tool] %s done\n", name)
|
||||
if err != nil {
|
||||
_ = emit("\n" + tools.FormatError(err))
|
||||
} else if streamed == 0 {
|
||||
_ = emit("(no output)")
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
if output != "" {
|
||||
return output + "\n" + tools.FormatError(err), true
|
||||
}
|
||||
return tools.FormatError(err), true
|
||||
}
|
||||
return output, true
|
||||
}
|
||||
|
||||
func (a *Agent) hasMutatingCall(calls []llm.ToolCall) bool {
|
||||
if a.toolRegistry == nil {
|
||||
return false
|
||||
}
|
||||
for _, call := range calls {
|
||||
if tool, ok := a.toolRegistry.Get(call.Function.Name); ok {
|
||||
if tools.IsMutating(tool) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (a *Agent) executeTools(ctx context.Context, calls []llm.ToolCall, logs io.Writer) []toolExecutionResult {
|
||||
results := make([]toolExecutionResult, len(calls))
|
||||
|
||||
// Log all tool starts sequentially so TUI gets ordered placeholders.
|
||||
if logs != nil {
|
||||
for _, call := range calls {
|
||||
logToolStart(logs, call)
|
||||
}
|
||||
}
|
||||
|
||||
if a.hasMutatingCall(calls) {
|
||||
// Sequential: avoid racing mutating tools.
|
||||
for i, call := range calls {
|
||||
output, streamed := a.executeTool(ctx, call, logs)
|
||||
results[i] = toolExecutionResult{call: call, output: output}
|
||||
if logs != nil && !streamed {
|
||||
logToolOutput(logs, call, output)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
// Concurrent: independent read-only tools.
|
||||
var mu sync.Mutex
|
||||
var wg sync.WaitGroup
|
||||
for i, call := range calls {
|
||||
wg.Add(1)
|
||||
go func(idx int, c llm.ToolCall) {
|
||||
defer wg.Done()
|
||||
output, streamed := a.executeTool(ctx, c, logs)
|
||||
results[idx] = toolExecutionResult{call: c, output: output}
|
||||
if logs != nil && !streamed {
|
||||
mu.Lock()
|
||||
logToolOutput(logs, c, output)
|
||||
mu.Unlock()
|
||||
}
|
||||
}(i, call)
|
||||
}
|
||||
wg.Wait()
|
||||
return results
|
||||
}
|
||||
|
||||
func logToolStart(logs io.Writer, call llm.ToolCall) {
|
||||
fmt.Fprintf(logs, "\n[tool] %s %s\n", call.Function.Name, compact(call.Function.Arguments, toolLogArgumentLimit))
|
||||
}
|
||||
|
||||
func logToolOutput(logs io.Writer, call llm.ToolCall, output string) {
|
||||
fmt.Fprintf(logs, "\n[tool] %s %s\n%s\n%s\n", call.Function.Name, compact(call.Function.Arguments, toolLogArgumentLimit), toolLogOutputMarker, toolOutputPreview(output))
|
||||
}
|
||||
|
||||
func toolOutputPreview(output string) string {
|
||||
output = strings.TrimSpace(output)
|
||||
if output == "" {
|
||||
return "(no output)"
|
||||
}
|
||||
if len(output) <= toolLogOutputPreviewBytes {
|
||||
return output
|
||||
}
|
||||
return output[:toolLogOutputPreviewBytes] + fmt.Sprintf("\n[truncated: showing first %d bytes; omitted %d bytes]", toolLogOutputPreviewBytes, len(output)-toolLogOutputPreviewBytes)
|
||||
}
|
||||
|
||||
type toolCallBuilder struct {
|
||||
index int
|
||||
|
||||
@@ -7,8 +7,8 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"agentu/pkg/llm"
|
||||
"agentu/internal/tools"
|
||||
"agentu/pkg/llm"
|
||||
)
|
||||
|
||||
type fakeProvider struct {
|
||||
@@ -87,6 +87,86 @@ func TestAgentRunsToolLoop(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
type streamToolProvider struct {
|
||||
calls int
|
||||
t *testing.T
|
||||
}
|
||||
|
||||
func (p *streamToolProvider) ChatStream(ctx context.Context, req llm.ChatRequest, emit func(llm.StreamEvent) error) error {
|
||||
p.calls++
|
||||
switch p.calls {
|
||||
case 1:
|
||||
return emit(llm.StreamEvent{ToolCalls: []llm.ToolCallDelta{
|
||||
{Index: 0, ID: "call_1", Type: "function", Name: "test_stream", Arguments: `{}`},
|
||||
}})
|
||||
case 2:
|
||||
last := req.Messages[len(req.Messages)-1]
|
||||
if last.Role != llm.RoleTool || last.Content != "full result" {
|
||||
p.t.Fatalf("last message = %#v", last)
|
||||
}
|
||||
return emit(llm.StreamEvent{Content: "done"})
|
||||
default:
|
||||
p.t.Fatalf("unexpected call count %d", p.calls)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
type streamEchoTool struct{}
|
||||
|
||||
func (streamEchoTool) Definition() llm.Tool {
|
||||
return llm.Tool{
|
||||
Type: "function",
|
||||
Function: llm.ToolFunction{
|
||||
Name: "test_stream",
|
||||
Description: "streaming test",
|
||||
Parameters: json.RawMessage(`{"type":"object"}`),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (streamEchoTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
||||
return "full result", nil
|
||||
}
|
||||
|
||||
func (streamEchoTool) ExecuteStream(ctx context.Context, raw json.RawMessage, emit func(string) error) (string, error) {
|
||||
if err := emit("one\n"); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := emit("two\n"); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "full result", nil
|
||||
}
|
||||
|
||||
func TestAgentStreamsToolOutputToLogs(t *testing.T) {
|
||||
provider := &streamToolProvider{t: t}
|
||||
a := New(Options{
|
||||
Provider: provider,
|
||||
Model: "test",
|
||||
SystemPrompt: "system",
|
||||
ToolRegistry: tools.NewRegistry(streamEchoTool{}),
|
||||
})
|
||||
|
||||
var out strings.Builder
|
||||
var logs strings.Builder
|
||||
if err := a.RunTurn(context.Background(), "stream", &out, &logs); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if out.String() != "done" {
|
||||
t.Fatalf("out = %q", out.String())
|
||||
}
|
||||
logText := logs.String()
|
||||
if strings.Count(logText, "[output]") != 2 {
|
||||
t.Fatalf("logs = %q", logText)
|
||||
}
|
||||
if !strings.Contains(logText, "one\n") || !strings.Contains(logText, "two\n") {
|
||||
t.Fatalf("logs missing streamed chunks: %q", logText)
|
||||
}
|
||||
if strings.Contains(logText, "[output]\nfull result") {
|
||||
t.Fatalf("logs contain duplicate full output: %q", logText)
|
||||
}
|
||||
}
|
||||
|
||||
type contentProvider struct {
|
||||
called bool
|
||||
text string
|
||||
@@ -145,6 +225,9 @@ func TestAgentCompact(t *testing.T) {
|
||||
if messages[len(messages)-1].Content != "recent-answer-2" {
|
||||
t.Fatalf("last recent message = %#v", messages[len(messages)-1])
|
||||
}
|
||||
if messages[2].Content != "recent-user-0" {
|
||||
t.Fatalf("first retained recent message = %#v", messages[2])
|
||||
}
|
||||
if usage := a.LastUsage(); usage == nil || usage.TotalTokens != 12 {
|
||||
t.Fatalf("usage = %#v", usage)
|
||||
}
|
||||
@@ -217,6 +300,48 @@ func longConversation() []llm.Message {
|
||||
return messages
|
||||
}
|
||||
|
||||
func manyTurnConversation(turns int) []llm.Message {
|
||||
messages := []llm.Message{{Role: llm.RoleSystem, Content: "system"}}
|
||||
for i := 0; i < turns; i++ {
|
||||
messages = append(messages,
|
||||
llm.Message{Role: llm.RoleUser, Content: fmt.Sprintf("user-%02d", i)},
|
||||
llm.Message{Role: llm.RoleAssistant, Content: fmt.Sprintf("assistant-%02d", i)},
|
||||
)
|
||||
}
|
||||
return messages
|
||||
}
|
||||
|
||||
func TestAgentCompactKeepsTwentyPercentRecentContext(t *testing.T) {
|
||||
provider := &compactProvider{}
|
||||
a := New(Options{Provider: provider, Model: "test", SystemPrompt: "system"})
|
||||
a.SetMessages(manyTurnConversation(25))
|
||||
|
||||
if err := a.Compact(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
messages := a.Messages()
|
||||
// 1 system + 1 summary + 10 retained recent messages (5 user/assistant turns)
|
||||
if len(messages) != 12 {
|
||||
t.Fatalf("messages = %d, want 12; contents: %v", len(messages), joinMessageContents(messages))
|
||||
}
|
||||
if provider.summaryCalls != 1 {
|
||||
t.Fatalf("summaryCalls = %d, want 1", provider.summaryCalls)
|
||||
}
|
||||
if messages[2].Content != "user-20" {
|
||||
t.Fatalf("first retained message = %#v, want user-20", messages[2])
|
||||
}
|
||||
if messages[len(messages)-1].Content != "assistant-24" {
|
||||
t.Fatalf("last retained message = %#v, want assistant-24", messages[len(messages)-1])
|
||||
}
|
||||
if !strings.Contains(provider.lastSummary, "user-19") {
|
||||
t.Fatalf("summary should contain user-19 (part of summarized 80%%): %q", provider.lastSummary)
|
||||
}
|
||||
if strings.Contains(provider.lastSummary, "user-20") {
|
||||
t.Fatalf("summary should not contain user-20 (part of retained 20%%): %q", provider.lastSummary)
|
||||
}
|
||||
}
|
||||
|
||||
func joinMessageContents(messages []llm.Message) string {
|
||||
var b strings.Builder
|
||||
for _, msg := range messages {
|
||||
@@ -340,3 +465,174 @@ func TestAgentHardLimitStopsVaryingCalls(t *testing.T) {
|
||||
t.Fatalf("calls = %d, want %d", provider.calls, limit)
|
||||
}
|
||||
}
|
||||
|
||||
// slowTool blocks on a channel before returning.
|
||||
type slowTool struct {
|
||||
name string
|
||||
started chan<- string
|
||||
release <-chan struct{}
|
||||
}
|
||||
|
||||
func (t *slowTool) Definition() llm.Tool {
|
||||
return llm.Tool{
|
||||
Type: "function",
|
||||
Function: llm.ToolFunction{
|
||||
Name: t.name,
|
||||
Description: "slow test tool",
|
||||
Parameters: json.RawMessage(`{"type":"object","properties":{"text":{"type":"string"}}}`),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (t *slowTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
||||
var args struct {
|
||||
Text string `json:"text"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &args); err != nil {
|
||||
return "", err
|
||||
}
|
||||
t.started <- t.name
|
||||
<-t.release
|
||||
return args.Text, nil
|
||||
}
|
||||
|
||||
// mutatingSlowTool is like slowTool but implements MutatingTool.
|
||||
type mutatingSlowTool struct {
|
||||
slowTool
|
||||
}
|
||||
|
||||
func (*mutatingSlowTool) Mutates() bool { return true }
|
||||
|
||||
func TestAgentExecutesIndependentToolCallsInParallel(t *testing.T) {
|
||||
started := make(chan string, 2)
|
||||
release := make(chan struct{})
|
||||
|
||||
tool1 := &slowTool{name: "slow_one", started: started, release: release}
|
||||
tool2 := &slowTool{name: "slow_two", started: started, release: release}
|
||||
|
||||
provider := ¶llelProvider{
|
||||
responses: [][]llm.ToolCallDelta{
|
||||
{
|
||||
{Index: 0, ID: "call_1", Type: "function", Name: "slow_one", Arguments: `{"text":"one"}`},
|
||||
{Index: 1, ID: "call_2", Type: "function", Name: "slow_two", Arguments: `{"text":"two"}`},
|
||||
},
|
||||
},
|
||||
finalContent: "done",
|
||||
}
|
||||
|
||||
a := New(Options{
|
||||
Provider: provider,
|
||||
Model: "test",
|
||||
SystemPrompt: "system",
|
||||
ToolRegistry: tools.NewRegistry(tool1, tool2),
|
||||
})
|
||||
|
||||
var out strings.Builder
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
done <- a.RunTurn(context.Background(), "run both", &out, nil)
|
||||
}()
|
||||
|
||||
// Both tools must start before we release them.
|
||||
names := make(map[string]bool)
|
||||
for i := 0; i < 2; i++ {
|
||||
select {
|
||||
case name := <-started:
|
||||
names[name] = true
|
||||
case <-make(chan struct{}):
|
||||
// Use a timer in real code, but for test simplicity we rely on deadlock detection.
|
||||
t.Fatal("timed out waiting for tool to start")
|
||||
}
|
||||
}
|
||||
if !names["slow_one"] || !names["slow_two"] {
|
||||
t.Fatalf("expected both tools to start, got %v", names)
|
||||
}
|
||||
close(release)
|
||||
|
||||
if err := <-done; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if out.String() != "done" {
|
||||
t.Fatalf("out = %q", out.String())
|
||||
}
|
||||
// Verify tool results were in order.
|
||||
msgs := a.Messages()
|
||||
var toolMsgs []string
|
||||
for _, m := range msgs {
|
||||
if m.Role == llm.RoleTool {
|
||||
toolMsgs = append(toolMsgs, m.Content)
|
||||
}
|
||||
}
|
||||
if len(toolMsgs) != 2 || toolMsgs[0] != "one" || toolMsgs[1] != "two" {
|
||||
t.Fatalf("tool messages = %v, want [one two]", toolMsgs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentRunsMutatingToolCallsSequentially(t *testing.T) {
|
||||
started := make(chan string, 2)
|
||||
release := make(chan struct{})
|
||||
|
||||
tool1 := &mutatingSlowTool{slowTool{name: "mut_one", started: started, release: release}}
|
||||
tool2 := &mutatingSlowTool{slowTool{name: "mut_two", started: started, release: release}}
|
||||
|
||||
provider := ¶llelProvider{
|
||||
responses: [][]llm.ToolCallDelta{
|
||||
{
|
||||
{Index: 0, ID: "call_1", Type: "function", Name: "mut_one", Arguments: `{"text":"one"}`},
|
||||
{Index: 1, ID: "call_2", Type: "function", Name: "mut_two", Arguments: `{"text":"two"}`},
|
||||
},
|
||||
},
|
||||
finalContent: "done",
|
||||
}
|
||||
|
||||
a := New(Options{
|
||||
Provider: provider,
|
||||
Model: "test",
|
||||
SystemPrompt: "system",
|
||||
ToolRegistry: tools.NewRegistry(tool1, tool2),
|
||||
})
|
||||
|
||||
var out strings.Builder
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
done <- a.RunTurn(context.Background(), "run both", &out, nil)
|
||||
}()
|
||||
|
||||
// Only the first tool should start.
|
||||
select {
|
||||
case name := <-started:
|
||||
if name != "mut_one" {
|
||||
t.Fatalf("expected mut_one to start first, got %q", name)
|
||||
}
|
||||
case <-make(chan struct{}):
|
||||
t.Fatal("timed out waiting for first tool to start")
|
||||
}
|
||||
|
||||
// Second tool should NOT have started yet.
|
||||
select {
|
||||
case name := <-started:
|
||||
t.Fatalf("second tool started too early: %q", name)
|
||||
default:
|
||||
// Expected: second tool hasn't started.
|
||||
}
|
||||
|
||||
close(release)
|
||||
if err := <-done; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// parallelProvider emits a batch of tool calls on the first request, then returns finalContent.
|
||||
type parallelProvider struct {
|
||||
responses [][]llm.ToolCallDelta
|
||||
finalContent string
|
||||
calls int
|
||||
}
|
||||
|
||||
func (p *parallelProvider) ChatStream(ctx context.Context, req llm.ChatRequest, emit func(llm.StreamEvent) error) error {
|
||||
p.calls++
|
||||
if p.calls <= len(p.responses) {
|
||||
return emit(llm.StreamEvent{ToolCalls: p.responses[p.calls-1]})
|
||||
}
|
||||
return emit(llm.StreamEvent{Content: p.finalContent})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,171 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"agentu/internal/agent"
|
||||
"agentu/internal/session"
|
||||
"agentu/internal/tools"
|
||||
"agentu/internal/tui"
|
||||
"agentu/pkg/config"
|
||||
"agentu/pkg/llm"
|
||||
)
|
||||
|
||||
type cliOptions struct {
|
||||
configPath string
|
||||
yolo bool
|
||||
plain bool
|
||||
themeName string
|
||||
resumeID string
|
||||
}
|
||||
|
||||
func Run() error {
|
||||
opts := parseCLIOptions()
|
||||
|
||||
themeMode, err := tui.ParseThemeMode(opts.themeName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cfg, err := loadConfig(opts.configPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
providerName, providerConfig, provider := defaultProvider(cfg, http.DefaultClient)
|
||||
assistant := newAssistant(cfg, providerName, providerConfig, provider, opts.yolo, http.DefaultClient)
|
||||
modelManager, err := newSessionManager(cfg, assistant, http.DefaultClient)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := initializeSession(modelManager, opts.resumeID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt)
|
||||
defer stop()
|
||||
|
||||
return runInterface(ctx, assistant, modelManager, providerConfig, cfg.Agent.WorkingDir, opts, themeMode)
|
||||
}
|
||||
|
||||
func parseCLIOptions() cliOptions {
|
||||
configPath := flag.String("config", config.DefaultPath, "path to YAML config")
|
||||
yolo := flag.Bool("yolo", false, "also enable automatic file writes and shell execution")
|
||||
plain := flag.Bool("plain", false, "use the simple line-based REPL instead of the TUI")
|
||||
themeName := flag.String("theme", "light", "TUI theme: light or dark")
|
||||
resumeID := flag.String("resume", "", "resume a specific session by ID")
|
||||
flag.Parse()
|
||||
|
||||
return cliOptions{
|
||||
configPath: *configPath,
|
||||
yolo: *yolo,
|
||||
plain: *plain,
|
||||
themeName: *themeName,
|
||||
resumeID: *resumeID,
|
||||
}
|
||||
}
|
||||
|
||||
func loadConfig(configPath string) (*config.Config, error) {
|
||||
cfg, err := config.Load(configPath)
|
||||
if err == nil {
|
||||
return cfg, nil
|
||||
}
|
||||
if !errors.Is(err, os.ErrNotExist) {
|
||||
return nil, err
|
||||
}
|
||||
if bootstrapErr := bootstrapConfig(configPath); bootstrapErr != nil {
|
||||
return nil, bootstrapErr
|
||||
}
|
||||
return nil, fmt.Errorf("created default config at %s; edit it to set your API key and provider, then restart agentu", configPath)
|
||||
}
|
||||
|
||||
func defaultProvider(cfg *config.Config, httpClient *http.Client) (string, config.ProviderConfig, llm.Provider) {
|
||||
providerName := cfg.DefaultProviderName()
|
||||
providerConfig := cfg.Providers[providerName]
|
||||
provider := newProviderClient(providerConfig, httpClient)
|
||||
return providerName, providerConfig, provider
|
||||
}
|
||||
|
||||
func newProviderClient(providerConfig config.ProviderConfig, httpClient *http.Client) llm.Provider {
|
||||
if providerConfig.APIType == config.APITypeResponses {
|
||||
return llm.NewOpenAIResponsesClient(providerConfig.BaseURL, providerConfig.APIKey, httpClient)
|
||||
}
|
||||
return llm.NewOpenAICompatibleClient(providerConfig.BaseURL, providerConfig.APIKey, httpClient)
|
||||
}
|
||||
|
||||
func newAssistant(cfg *config.Config, providerName string, providerConfig config.ProviderConfig, provider llm.Provider, yolo bool, httpClient *http.Client) *agent.Agent {
|
||||
return agent.New(agent.Options{
|
||||
Provider: provider,
|
||||
ProviderName: providerName,
|
||||
Model: providerConfig.Model,
|
||||
Thinking: providerConfig.Thinking,
|
||||
ThinkingParam: providerConfig.ThinkingParam,
|
||||
ThinkingEnabled: providerConfig.ThinkingConfigured,
|
||||
SystemPrompt: systemPrompt(cfg, yolo),
|
||||
ToolRegistry: newToolRegistry(cfg.Agent.WorkingDir, yolo, httpClient),
|
||||
ToolTimeout: cfg.Agent.ToolTimeout,
|
||||
MaxContextTokens: maxContextTokens(cfg, providerConfig),
|
||||
})
|
||||
}
|
||||
|
||||
func systemPrompt(cfg *config.Config, yolo bool) string {
|
||||
instructions := tools.AgentInstructions(cfg.Agent.WorkingDir, yolo)
|
||||
return strings.TrimSpace(cfg.Agent.SystemPrompt + "\n\n" + instructions)
|
||||
}
|
||||
|
||||
func maxContextTokens(cfg *config.Config, providerConfig config.ProviderConfig) int {
|
||||
if cfg.Agent.MaxContextTokens != 0 {
|
||||
return cfg.Agent.MaxContextTokens
|
||||
}
|
||||
return providerConfig.ContextTokens
|
||||
}
|
||||
|
||||
func newToolRegistry(workingDir string, yolo bool, httpClient *http.Client) *tools.Registry {
|
||||
var registry *tools.Registry
|
||||
if yolo {
|
||||
registry = tools.Builtins(workingDir)
|
||||
} else {
|
||||
registry = tools.ReadOnlyBuiltins(workingDir)
|
||||
}
|
||||
registry.Register(tools.NewBuiltinSearchTool(httpClient))
|
||||
registry.Register(tools.NewBuiltinFetchTool(httpClient))
|
||||
return registry
|
||||
}
|
||||
|
||||
func newSessionManager(cfg *config.Config, assistant *agent.Agent, httpClient *http.Client) (*session.Manager, error) {
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("resolve home dir: %w", err)
|
||||
}
|
||||
sessionsDir := filepath.Join(home, ".agentu", "sessions")
|
||||
store, err := session.NewStore(sessionsDir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("init session store: %w", err)
|
||||
}
|
||||
return session.NewManager(cfg, assistant, httpClient, store)
|
||||
}
|
||||
|
||||
func runInterface(ctx context.Context, assistant *agent.Agent, modelManager *session.Manager, providerConfig config.ProviderConfig, workingDir string, opts cliOptions, themeMode tui.ThemeMode) error {
|
||||
if opts.plain {
|
||||
return repl(ctx, assistant, modelManager)
|
||||
}
|
||||
|
||||
err := tui.Run(ctx, assistant, tui.Options{
|
||||
ModelName: providerConfig.Model,
|
||||
Yolo: opts.yolo,
|
||||
ThemeMode: themeMode,
|
||||
ModelManager: modelManager,
|
||||
WorkingDir: workingDir,
|
||||
})
|
||||
printExitSession(modelManager)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"agentu/pkg/config"
|
||||
)
|
||||
|
||||
const defaultConfigContent = `# agentu configuration
|
||||
# Replace the placeholder values below with your actual provider settings.
|
||||
# See agentu.example.yaml for the full reference.
|
||||
providers:
|
||||
default:
|
||||
base_url: https://api.openai.com
|
||||
api_key: YOUR_API_KEY_HERE
|
||||
models:
|
||||
- id: gpt-4o
|
||||
context: 128000
|
||||
- id: gpt-4o-mini
|
||||
context: 128000
|
||||
`
|
||||
|
||||
func bootstrapConfig(configPath string) error {
|
||||
resolved, err := config.ResolvePath(configPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve config path: %w", err)
|
||||
}
|
||||
dir := filepath.Dir(resolved)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return fmt.Errorf("create config directory: %w", err)
|
||||
}
|
||||
if err := os.WriteFile(resolved, []byte(defaultConfigContent), 0o644); err != nil {
|
||||
return fmt.Errorf("write default config: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package main
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -0,0 +1,60 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"agentu/internal/agent"
|
||||
"agentu/internal/session"
|
||||
)
|
||||
|
||||
func repl(ctx context.Context, assistant *agent.Agent, modelManager *session.Manager) error {
|
||||
scanner := bufio.NewScanner(os.Stdin)
|
||||
scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024)
|
||||
|
||||
fmt.Println("agentu")
|
||||
fmt.Println("Type /exit to quit, /clear to reset context, /compact to compress context.")
|
||||
for {
|
||||
fmt.Print("> ")
|
||||
if !scanner.Scan() {
|
||||
if err := scanner.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Println()
|
||||
printExitSession(modelManager)
|
||||
return nil
|
||||
}
|
||||
|
||||
input := strings.TrimSpace(scanner.Text())
|
||||
switch input {
|
||||
case "":
|
||||
continue
|
||||
case "/exit", "/quit":
|
||||
printExitSession(modelManager)
|
||||
return nil
|
||||
case "/clear":
|
||||
assistant.Clear()
|
||||
fmt.Println("context cleared")
|
||||
continue
|
||||
case "/compact":
|
||||
if err := assistant.Compact(ctx); err != nil {
|
||||
fmt.Fprintln(os.Stderr, "compact error:", err)
|
||||
} else {
|
||||
fmt.Println("context compacted")
|
||||
if modelManager != nil {
|
||||
_ = modelManager.Save()
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if err := assistant.RunTurn(ctx, input, os.Stdout, os.Stderr); err != nil {
|
||||
fmt.Fprintln(os.Stderr, "\nerror:", err)
|
||||
}
|
||||
_ = modelManager.Save()
|
||||
fmt.Println()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"agentu/internal/session"
|
||||
)
|
||||
|
||||
func initializeSession(modelManager *session.Manager, resumeID string) error {
|
||||
if resumeID != "" {
|
||||
if err := modelManager.Resume(resumeID); err != nil {
|
||||
return fmt.Errorf("resume session: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if _, err := modelManager.NewSession(); err != nil {
|
||||
return fmt.Errorf("start new session: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func printExitSession(m *session.Manager) {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
_ = m.Save()
|
||||
id := m.CurrentSessionID()
|
||||
name := m.CurrentSessionName()
|
||||
if id != "" {
|
||||
if name != "" && name != id {
|
||||
fmt.Fprintf(os.Stderr, "session: %s (%s)\n", id, name)
|
||||
} else {
|
||||
fmt.Fprintf(os.Stderr, "session: %s\n", id)
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "resume: agentu --resume %s\n", id)
|
||||
}
|
||||
}
|
||||
@@ -68,11 +68,29 @@ func (m *Manager) Resume(id string) error {
|
||||
return err
|
||||
}
|
||||
m.currentSession = sess
|
||||
m.restoreRuntime(sess)
|
||||
m.agent.SetMessages(sess.Messages)
|
||||
m.agent.SetLastUsage(sess.LastUsage)
|
||||
return nil
|
||||
}
|
||||
|
||||
// restoreRuntime switches the manager and agent back to the provider, model,
|
||||
// and thinking level recorded on a resumed session. Values that are no longer
|
||||
// valid for the current config are skipped and the defaults remain active.
|
||||
func (m *Manager) restoreRuntime(sess *Session) {
|
||||
if sess.Provider != "" && sess.Provider != m.providerName {
|
||||
if _, ok := m.cfg.Providers[sess.Provider]; ok {
|
||||
_, _ = m.SetProvider(sess.Provider)
|
||||
}
|
||||
}
|
||||
if sess.Model != "" && sess.Model != m.provider.Model {
|
||||
_, _ = m.SetModel(sess.Model)
|
||||
}
|
||||
if sess.Thinking != "" && sess.Thinking != m.provider.Thinking {
|
||||
_, _ = m.SetThinking(sess.Thinking)
|
||||
}
|
||||
}
|
||||
|
||||
// Save persists the current session to disk.
|
||||
func (m *Manager) Save() error {
|
||||
if m.store == nil || m.currentSession == nil {
|
||||
@@ -82,6 +100,8 @@ func (m *Manager) Save() error {
|
||||
m.currentSession.LastUsage = m.agent.LastUsage()
|
||||
m.currentSession.Provider = m.providerName
|
||||
m.currentSession.Model = m.provider.Model
|
||||
m.currentSession.APIType = m.provider.APIType
|
||||
m.currentSession.Thinking = m.provider.Thinking
|
||||
// Auto-name if empty
|
||||
if m.currentSession.Name == "" {
|
||||
m.currentSession.Name = DefaultName(m.currentSession.Messages)
|
||||
@@ -150,6 +170,8 @@ func (m *Manager) NewSession() (string, error) {
|
||||
m.currentSession.Messages = msgs
|
||||
m.currentSession.Provider = m.providerName
|
||||
m.currentSession.Model = m.provider.Model
|
||||
m.currentSession.APIType = m.provider.APIType
|
||||
m.currentSession.Thinking = m.provider.Thinking
|
||||
if m.currentSession.Name == "" {
|
||||
m.currentSession.Name = DefaultName(msgs)
|
||||
}
|
||||
@@ -204,10 +226,18 @@ func (m *Manager) CurrentThinking() string {
|
||||
return m.provider.Thinking
|
||||
}
|
||||
|
||||
func (m *Manager) CurrentAPIType() string {
|
||||
if m.provider.APIType == "" {
|
||||
return config.APITypeChat
|
||||
}
|
||||
return m.provider.APIType
|
||||
}
|
||||
|
||||
func (m *Manager) Info() string {
|
||||
var b strings.Builder
|
||||
fmt.Fprintf(&b, "provider: %s\n", m.CurrentProvider())
|
||||
fmt.Fprintf(&b, "model: %s\n", m.CurrentModel())
|
||||
fmt.Fprintf(&b, "api: %s\n", m.CurrentAPIType())
|
||||
fmt.Fprintf(&b, "thinking: %s\n", m.CurrentThinking())
|
||||
fmt.Fprintf(&b, "providers: %s\n", strings.Join(m.providerNames(), ", "))
|
||||
fmt.Fprintf(&b, "models: %s\n", strings.Join(m.provider.Models, ", "))
|
||||
@@ -223,7 +253,7 @@ func (m *Manager) SetProvider(name string) (string, error) {
|
||||
}
|
||||
m.providerName = name
|
||||
m.provider = provider
|
||||
m.syncAgent(llm.NewOpenAICompatibleClient(provider.BaseURL, provider.APIKey, m.httpClient))
|
||||
m.syncAgent(providerClient(provider, m.httpClient))
|
||||
return "Switched provider.\n" + m.summary(), nil
|
||||
}
|
||||
|
||||
@@ -249,7 +279,7 @@ func (m *Manager) SetThinking(level string) (string, error) {
|
||||
m.provider.Thinking = level
|
||||
m.provider.ThinkingConfigured = true
|
||||
if m.provider.ThinkingParam == "" {
|
||||
m.provider.ThinkingParam = "thinking"
|
||||
m.provider.ThinkingParam = config.DefaultThinkingParam(m.provider.APIType)
|
||||
}
|
||||
m.cfg.Providers[m.providerName] = m.provider
|
||||
m.syncAgent(nil)
|
||||
@@ -257,7 +287,14 @@ func (m *Manager) SetThinking(level string) (string, error) {
|
||||
}
|
||||
|
||||
func (m *Manager) summary() string {
|
||||
return fmt.Sprintf("provider: %s\nmodel: %s\nthinking: %s", m.CurrentProvider(), m.CurrentModel(), m.CurrentThinking())
|
||||
return fmt.Sprintf("provider: %s\nmodel: %s\napi: %s\nthinking: %s", m.CurrentProvider(), m.CurrentModel(), m.CurrentAPIType(), m.CurrentThinking())
|
||||
}
|
||||
|
||||
func providerClient(provider config.ProviderConfig, httpClient *http.Client) llm.Provider {
|
||||
if provider.APIType == config.APITypeResponses {
|
||||
return llm.NewOpenAIResponsesClient(provider.BaseURL, provider.APIKey, httpClient)
|
||||
}
|
||||
return llm.NewOpenAICompatibleClient(provider.BaseURL, provider.APIKey, httpClient)
|
||||
}
|
||||
|
||||
func (m *Manager) maxContextTokens() int {
|
||||
@@ -288,9 +325,13 @@ func (m *Manager) providerNames() []string {
|
||||
return names
|
||||
}
|
||||
|
||||
func (m *Manager) AvailableModels() []string {
|
||||
return append([]string(nil), m.provider.Models...)
|
||||
}
|
||||
|
||||
func contains(values []string, value string) bool {
|
||||
for _, item := range values {
|
||||
if item == value {
|
||||
for _, v := range values {
|
||||
if v == value {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,6 +2,8 @@ package session
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -109,6 +111,76 @@ func TestManagerRejectsUnknownModelAndThinking(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerSwitchesToResponsesProvider(t *testing.T) {
|
||||
var hitPath string
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hitPath = r.URL.Path
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = w.Write([]byte(`data: {"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"ok"}` + "\n\n"))
|
||||
_, _ = w.Write([]byte(`data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}` + "\n\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
cfg := &config.Config{
|
||||
Providers: map[string]config.ProviderConfig{
|
||||
"chat": {
|
||||
Name: "chat",
|
||||
APIType: config.APITypeChat,
|
||||
BaseURL: "https://chat.example.com",
|
||||
APIKey: "sk-test",
|
||||
Model: "model-a",
|
||||
Models: []string{"model-a"},
|
||||
},
|
||||
"responses": {
|
||||
Name: "responses",
|
||||
APIType: config.APITypeResponses,
|
||||
BaseURL: server.URL,
|
||||
APIKey: "sk-test",
|
||||
Model: "model-b",
|
||||
Models: []string{"model-b"},
|
||||
},
|
||||
},
|
||||
}
|
||||
assistant := agent.New(agent.Options{Provider: &recordingProvider{}, Model: "model-a"})
|
||||
store, err := NewStore(filepath.Join(t.TempDir(), "sessions"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
manager, err := NewManager(cfg, assistant, server.Client(), store)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if _, err := manager.SetProvider("responses"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if manager.CurrentAPIType() != config.APITypeResponses {
|
||||
t.Fatalf("api type = %q", manager.CurrentAPIType())
|
||||
}
|
||||
var out strings.Builder
|
||||
if err := assistant.RunTurn(context.Background(), "hi", &out, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if hitPath != "/v1/responses" {
|
||||
t.Fatalf("path = %q, want /v1/responses", hitPath)
|
||||
}
|
||||
if out.String() != "ok" {
|
||||
t.Fatalf("output = %q", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerExposesModelCommandCandidates(t *testing.T) {
|
||||
manager, _, _ := testManager(t)
|
||||
models := manager.AvailableModels()
|
||||
if len(models) != 2 || models[0] != "model-a" || models[1] != "model-b" {
|
||||
t.Fatalf("AvailableModels() = %v", models)
|
||||
}
|
||||
models[0] = "mutated"
|
||||
if got := manager.AvailableModels()[0]; got != "model-a" {
|
||||
t.Fatalf("mutation leaked: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerSaveAndResume(t *testing.T) {
|
||||
manager, assistant, _ := testManager(t)
|
||||
|
||||
@@ -331,3 +403,132 @@ func TestManagerUsesModelContextForAgentCompaction(t *testing.T) {
|
||||
t.Fatalf("answerCalls = %d, want 1", provider.answerCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerResumeRestoresRuntime(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = w.Write([]byte(`data: {"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"ok"}` + "\n\n"))
|
||||
_, _ = w.Write([]byte(`data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}` + "\n\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
cfg := &config.Config{
|
||||
Providers: map[string]config.ProviderConfig{
|
||||
"chat": {
|
||||
Name: "chat",
|
||||
APIType: config.APITypeChat,
|
||||
BaseURL: "https://chat.example.com",
|
||||
APIKey: "sk-test",
|
||||
Model: "model-a",
|
||||
Models: []string{"model-a"},
|
||||
Thinking: "none",
|
||||
ThinkingParam: "thinking",
|
||||
ThinkingConfigured: true,
|
||||
},
|
||||
"resp": {
|
||||
Name: "resp",
|
||||
APIType: config.APITypeResponses,
|
||||
BaseURL: server.URL,
|
||||
APIKey: "sk-test",
|
||||
Model: "m1",
|
||||
Models: []string{"m1", "m2"},
|
||||
Thinking: "none",
|
||||
ThinkingParam: "reasoning",
|
||||
ThinkingConfigured: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
assistant := agent.New(agent.Options{
|
||||
Provider: &recordingProvider{},
|
||||
Model: "model-a",
|
||||
Thinking: "none",
|
||||
ThinkingParam: "thinking",
|
||||
ThinkingEnabled: true,
|
||||
})
|
||||
store, err := NewStore(filepath.Join(t.TempDir(), "sessions"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
manager, err := NewManager(cfg, assistant, server.Client(), store)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Build a session on the responses provider with a custom model and thinking.
|
||||
manager.newSession()
|
||||
if _, err := manager.SetProvider("resp"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := manager.SetModel("m2"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := manager.SetThinking("xhigh"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := manager.Save(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
id := manager.CurrentSessionID()
|
||||
|
||||
// Drift away from the saved runtime.
|
||||
if _, err := manager.SetProvider("chat"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := manager.SetModel("model-a"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := manager.SetThinking("none"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if err := manager.Resume(id); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := manager.CurrentProvider(); got != "resp" {
|
||||
t.Fatalf("provider = %q, want resp", got)
|
||||
}
|
||||
if got := manager.CurrentModel(); got != "m2" {
|
||||
t.Fatalf("model = %q, want m2", got)
|
||||
}
|
||||
if got := manager.CurrentThinking(); got != "xhigh" {
|
||||
t.Fatalf("thinking = %q, want xhigh", got)
|
||||
}
|
||||
if got := manager.CurrentAPIType(); got != config.APITypeResponses {
|
||||
t.Fatalf("api type = %q, want responses", got)
|
||||
}
|
||||
runtime := assistant.RuntimeInfo()
|
||||
if runtime.ProviderName != "resp" || runtime.Model != "m2" || runtime.Thinking != "xhigh" {
|
||||
t.Fatalf("agent runtime = %#v", runtime)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerResumeFallsBackWhenRuntimeMissing(t *testing.T) {
|
||||
manager, assistant, _ := testManager(t)
|
||||
|
||||
sess := &Session{
|
||||
ID: GenerateID(),
|
||||
Provider: "ghost",
|
||||
Model: "ghost-model",
|
||||
Thinking: "xhigh",
|
||||
Messages: []llm.Message{{Role: llm.RoleUser, Content: "hi"}},
|
||||
}
|
||||
if err := manager.store.Save(sess); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := manager.Resume(sess.ID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := manager.CurrentProvider(); got != "one" {
|
||||
t.Fatalf("provider = %q, want default one", got)
|
||||
}
|
||||
if got := manager.CurrentModel(); got != "model-a" {
|
||||
t.Fatalf("model = %q, want default model-a", got)
|
||||
}
|
||||
if got := manager.CurrentThinking(); got != "xhigh" {
|
||||
t.Fatalf("thinking = %q, want restored xhigh", got)
|
||||
}
|
||||
runtime := assistant.RuntimeInfo()
|
||||
if runtime.ProviderName != "one" || runtime.Model != "model-a" {
|
||||
t.Fatalf("agent runtime = %#v", runtime)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -22,6 +22,8 @@ type Session struct {
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
Provider string `json:"provider"`
|
||||
Model string `json:"model"`
|
||||
APIType string `json:"api_type,omitempty"`
|
||||
Thinking string `json:"thinking,omitempty"`
|
||||
Messages []llm.Message `json:"messages"`
|
||||
LastUsage *llm.Usage `json:"last_usage,omitempty"`
|
||||
}
|
||||
|
||||
@@ -81,6 +81,7 @@ type FileWriteTool struct {
|
||||
func NewFileWriteTool(workingDir string) *FileWriteTool {
|
||||
return &FileWriteTool{workingDir: workingDir}
|
||||
}
|
||||
func (*FileWriteTool) Mutates() bool { return true }
|
||||
|
||||
func (t *FileWriteTool) Definition() llm.Tool {
|
||||
return llm.Tool{
|
||||
|
||||
@@ -0,0 +1,178 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"agentu/pkg/llm"
|
||||
)
|
||||
|
||||
type FileEditTool struct {
|
||||
workingDir string
|
||||
}
|
||||
|
||||
func NewFileEditTool(workingDir string) *FileEditTool {
|
||||
return &FileEditTool{workingDir: workingDir}
|
||||
}
|
||||
|
||||
func (*FileEditTool) Mutates() bool { return true }
|
||||
|
||||
func (t *FileEditTool) Definition() llm.Tool {
|
||||
return llm.Tool{
|
||||
Type: "function",
|
||||
Function: llm.ToolFunction{
|
||||
Name: "file_edit",
|
||||
Description: "Edit an existing UTF-8 text file by replacing an inclusive 1-based line range or an exact old text pattern.",
|
||||
Parameters: JSONSchema(`{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {"type": "string", "description": "Existing file path relative to the configured agent working directory."},
|
||||
"start_line": {"type": "integer", "description": "1-based first line for a line-range edit."},
|
||||
"end_line": {"type": "integer", "description": "Optional 1-based final line; defaults to start_line when omitted."},
|
||||
"old": {"type": "string", "description": "Exact text to replace for a pattern edit; not regex."},
|
||||
"new": {"type": "string", "description": "Replacement text; may be empty string for deletion."},
|
||||
"replace_all": {"type": "boolean", "description": "When true, replace every exact old occurrence; otherwise require exactly one occurrence."}
|
||||
},
|
||||
"required": ["path", "new"],
|
||||
"additionalProperties": false
|
||||
}`),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
type fileEditArgs struct {
|
||||
Path string `json:"path"`
|
||||
StartLine int `json:"start_line"`
|
||||
EndLine int `json:"end_line"`
|
||||
Old string `json:"old"`
|
||||
New *string `json:"new"`
|
||||
ReplaceAll bool `json:"replace_all"`
|
||||
}
|
||||
|
||||
func (t *FileEditTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
||||
var args fileEditArgs
|
||||
if err := json.Unmarshal(raw, &args); err != nil {
|
||||
return "", fmt.Errorf("parse arguments: %w", err)
|
||||
}
|
||||
if strings.TrimSpace(args.Path) == "" {
|
||||
return "", errors.New("path is required")
|
||||
}
|
||||
if args.New == nil {
|
||||
return "", errors.New("new is required")
|
||||
}
|
||||
path, err := resolvePath(t.workingDir, args.Path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("cannot access %s: %w", args.Path, err)
|
||||
}
|
||||
if info.IsDir() {
|
||||
return "", fmt.Errorf("%s is a directory, not a file", args.Path)
|
||||
}
|
||||
perm := info.Mode().Perm()
|
||||
|
||||
hasLine := args.StartLine > 0 || args.EndLine > 0
|
||||
hasOld := strings.TrimSpace(args.Old) != ""
|
||||
|
||||
if hasLine && hasOld {
|
||||
return "", errors.New("provide either start_line/end_line or old, not both")
|
||||
}
|
||||
if !hasLine && !hasOld {
|
||||
return "", errors.New("provide either start_line/end_line or old")
|
||||
}
|
||||
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
content := string(data)
|
||||
|
||||
var result string
|
||||
if hasOld {
|
||||
result, err = editByPattern(content, args.Path, args.Old, *args.New, args.ReplaceAll)
|
||||
} else {
|
||||
result, err = editByLineRange(content, args.Path, args.StartLine, args.EndLine, *args.New)
|
||||
}
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := os.WriteFile(path, []byte(result), perm); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
switch {
|
||||
case hasOld:
|
||||
count := strings.Count(content, args.Old)
|
||||
if args.ReplaceAll {
|
||||
return fmt.Sprintf("edited %s: replaced %d occurrences", args.Path, count), nil
|
||||
}
|
||||
return fmt.Sprintf("edited %s: replaced 1 occurrence", args.Path), nil
|
||||
default:
|
||||
endLine := args.EndLine
|
||||
if endLine == 0 {
|
||||
endLine = args.StartLine
|
||||
}
|
||||
return fmt.Sprintf("edited %s: replaced lines %d-%d", args.Path, args.StartLine, endLine), nil
|
||||
}
|
||||
}
|
||||
|
||||
func editByLineRange(content, displayPath string, startLine, endLine int, replacement string) (string, error) {
|
||||
if startLine < 1 {
|
||||
return "", errors.New("start_line must be >= 1")
|
||||
}
|
||||
if endLine == 0 {
|
||||
endLine = startLine
|
||||
}
|
||||
if endLine < startLine {
|
||||
return "", fmt.Errorf("end_line (%d) must be >= start_line (%d)", endLine, startLine)
|
||||
}
|
||||
|
||||
lines := strings.Split(content, "\n")
|
||||
if startLine > len(lines) {
|
||||
return "", fmt.Errorf("line range %d-%d is outside file with %d lines", startLine, endLine, len(lines))
|
||||
}
|
||||
if endLine > len(lines) {
|
||||
return "", fmt.Errorf("line range %d-%d is outside file with %d lines", startLine, endLine, len(lines))
|
||||
}
|
||||
|
||||
var b strings.Builder
|
||||
b.WriteString(strings.Join(lines[:startLine-1], "\n"))
|
||||
if startLine > 1 {
|
||||
b.WriteString("\n")
|
||||
}
|
||||
b.WriteString(replacement)
|
||||
if endLine < len(lines) {
|
||||
if !strings.HasSuffix(replacement, "\n") && replacement != "" {
|
||||
b.WriteString("\n")
|
||||
}
|
||||
b.WriteString(strings.Join(lines[endLine:], "\n"))
|
||||
}
|
||||
|
||||
return b.String(), nil
|
||||
}
|
||||
|
||||
func editByPattern(content, displayPath, old, replacement string, replaceAll bool) (string, error) {
|
||||
count := strings.Count(content, old)
|
||||
if count == 0 {
|
||||
return "", errors.New("old text not found")
|
||||
}
|
||||
if !replaceAll && count != 1 {
|
||||
return "", fmt.Errorf("old text matched %d occurrences; set replace_all=true or provide a more specific old value", count)
|
||||
}
|
||||
return strings.ReplaceAll(content, old, replacement), nil
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestFileEditReplacesLineRange(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "notes.txt")
|
||||
if err := os.WriteFile(path, []byte("one\ntwo\nthree\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
edit := NewFileEditTool(dir)
|
||||
out, err := edit.Execute(context.Background(), json.RawMessage(`{"path":"notes.txt","start_line":2,"end_line":2,"new":"TWO\n"}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(out, "replaced lines 2-2") {
|
||||
t.Fatalf("output = %q", out)
|
||||
}
|
||||
got, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(got) != "one\nTWO\nthree\n" {
|
||||
t.Fatalf("file content = %q", string(got))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileEditReplacesExactPattern(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "notes.txt")
|
||||
if err := os.WriteFile(path, []byte("alpha beta alpha"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
edit := NewFileEditTool(dir)
|
||||
out, err := edit.Execute(context.Background(), json.RawMessage(`{"path":"notes.txt","old":"beta","new":"BETA"}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(out, "replaced 1 occurrence") {
|
||||
t.Fatalf("output = %q", out)
|
||||
}
|
||||
got, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(got) != "alpha BETA alpha" {
|
||||
t.Fatalf("file content = %q", string(got))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileEditRejectsAmbiguousPattern(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "notes.txt")
|
||||
if err := os.WriteFile(path, []byte("x x"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
edit := NewFileEditTool(dir)
|
||||
_, err := edit.Execute(context.Background(), json.RawMessage(`{"path":"notes.txt","old":"x","new":"y"}`))
|
||||
if err == nil {
|
||||
t.Fatal("expected error for ambiguous pattern")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "matched 2 occurrences") {
|
||||
t.Fatalf("error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileEditReplaceAll(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "notes.txt")
|
||||
if err := os.WriteFile(path, []byte("x x"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
edit := NewFileEditTool(dir)
|
||||
out, err := edit.Execute(context.Background(), json.RawMessage(`{"path":"notes.txt","old":"x","new":"y","replace_all":true}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(out, "replaced 2 occurrences") {
|
||||
t.Fatalf("output = %q", out)
|
||||
}
|
||||
got, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(got) != "y y" {
|
||||
t.Fatalf("file content = %q", string(got))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileEditRejectsEscapingPath(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
edit := NewFileEditTool(dir)
|
||||
_, err := edit.Execute(context.Background(), json.RawMessage(`{"path":"../outside.txt","old":"x","new":"y"}`))
|
||||
if err == nil {
|
||||
t.Fatal("expected error for escaping path")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "escapes") {
|
||||
t.Fatalf("error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileEditRejectsMissingNew(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "notes.txt")
|
||||
if err := os.WriteFile(path, []byte("hello"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
edit := NewFileEditTool(dir)
|
||||
_, err := edit.Execute(context.Background(), json.RawMessage(`{"path":"notes.txt","old":"hello"}`))
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing new")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "new is required") {
|
||||
t.Fatalf("error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileEditRejectsNonExistentFile(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
edit := NewFileEditTool(dir)
|
||||
_, err := edit.Execute(context.Background(), json.RawMessage(`{"path":"missing.txt","old":"x","new":"y"}`))
|
||||
if err == nil {
|
||||
t.Fatal("expected error for non-existent file")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,246 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"agentu/pkg/llm"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultGitLogMaxCount = 20
|
||||
maxGitLogMaxCount = 100
|
||||
)
|
||||
|
||||
func runGit(ctx context.Context, workingDir string, args ...string) (string, error) {
|
||||
gitArgs := append([]string{"-C", workingDir}, args...)
|
||||
cmd := exec.CommandContext(ctx, "git", gitArgs...)
|
||||
output, err := cmd.CombinedOutput()
|
||||
text := FormatResult(string(output))
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "context deadline exceeded") || strings.Contains(err.Error(), "signal: killed") {
|
||||
if text != "" {
|
||||
return text, fmt.Errorf("git command timed out: %w", err)
|
||||
}
|
||||
return "", fmt.Errorf("git command timed out: %w", err)
|
||||
}
|
||||
if text != "" {
|
||||
return text, fmt.Errorf("git command failed: %w", err)
|
||||
}
|
||||
return "", fmt.Errorf("git command failed: %w", err)
|
||||
}
|
||||
return text, nil
|
||||
}
|
||||
|
||||
func gitPathspec(baseDir, path string) (string, error) {
|
||||
if strings.TrimSpace(path) == "" {
|
||||
return "", nil
|
||||
}
|
||||
resolved, err := resolvePath(baseDir, path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
rel, err := filepath.Rel(baseDir, resolved)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.ToSlash(filepath.Clean(rel)), nil
|
||||
}
|
||||
|
||||
func normalizeGitLogMaxCount(value int) int {
|
||||
if value <= 0 {
|
||||
return defaultGitLogMaxCount
|
||||
}
|
||||
if value > maxGitLogMaxCount {
|
||||
return maxGitLogMaxCount
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
// git_status
|
||||
|
||||
type GitStatusTool struct {
|
||||
workingDir string
|
||||
}
|
||||
|
||||
func NewGitStatusTool(workingDir string) *GitStatusTool {
|
||||
return &GitStatusTool{workingDir: workingDir}
|
||||
}
|
||||
|
||||
func (t *GitStatusTool) Definition() llm.Tool {
|
||||
return llm.Tool{
|
||||
Type: "function",
|
||||
Function: llm.ToolFunction{
|
||||
Name: "git_status",
|
||||
Description: "Show changed files using git status --short.",
|
||||
Parameters: JSONSchema(`{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {"type": "string", "description": "Optional path to limit status to. Relative paths are resolved from the configured agent working directory."}
|
||||
},
|
||||
"additionalProperties": false
|
||||
}`),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
type gitStatusArgs struct {
|
||||
Path string `json:"path"`
|
||||
}
|
||||
|
||||
func (t *GitStatusTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
||||
var args gitStatusArgs
|
||||
if len(raw) > 0 {
|
||||
if err := json.Unmarshal(raw, &args); err != nil {
|
||||
return "", fmt.Errorf("parse arguments: %w", err)
|
||||
}
|
||||
}
|
||||
gitArgs := []string{"status", "--short"}
|
||||
if args.Path != "" {
|
||||
pathspec, err := gitPathspec(t.workingDir, args.Path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
gitArgs = append(gitArgs, "--", pathspec)
|
||||
}
|
||||
text, err := runGit(ctx, t.workingDir, gitArgs...)
|
||||
if err != nil {
|
||||
return text, err
|
||||
}
|
||||
if strings.TrimSpace(text) == "" {
|
||||
return "Working tree clean.", nil
|
||||
}
|
||||
return "Changed files:\n" + text, nil
|
||||
}
|
||||
|
||||
// git_diff
|
||||
|
||||
type GitDiffTool struct {
|
||||
workingDir string
|
||||
}
|
||||
|
||||
func NewGitDiffTool(workingDir string) *GitDiffTool {
|
||||
return &GitDiffTool{workingDir: workingDir}
|
||||
}
|
||||
|
||||
func (t *GitDiffTool) Definition() llm.Tool {
|
||||
return llm.Tool{
|
||||
Type: "function",
|
||||
Function: llm.ToolFunction{
|
||||
Name: "git_diff",
|
||||
Description: "Show unstaged or staged git diff output, optionally as --stat.",
|
||||
Parameters: JSONSchema(`{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {"type": "string", "description": "Optional path to limit diff to. Relative paths are resolved from the configured agent working directory."},
|
||||
"staged": {"type": "boolean", "description": "Show staged diff instead of unstaged."},
|
||||
"stat": {"type": "boolean", "description": "Show diffstat summary instead of full diff."}
|
||||
},
|
||||
"additionalProperties": false
|
||||
}`),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
type gitDiffArgs struct {
|
||||
Path string `json:"path"`
|
||||
Staged bool `json:"staged"`
|
||||
Stat bool `json:"stat"`
|
||||
}
|
||||
|
||||
func (t *GitDiffTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
||||
var args gitDiffArgs
|
||||
if len(raw) > 0 {
|
||||
if err := json.Unmarshal(raw, &args); err != nil {
|
||||
return "", fmt.Errorf("parse arguments: %w", err)
|
||||
}
|
||||
}
|
||||
gitArgs := []string{"diff"}
|
||||
if args.Staged {
|
||||
gitArgs = append(gitArgs, "--cached")
|
||||
}
|
||||
if args.Stat {
|
||||
gitArgs = append(gitArgs, "--stat")
|
||||
}
|
||||
if args.Path != "" {
|
||||
pathspec, err := gitPathspec(t.workingDir, args.Path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
gitArgs = append(gitArgs, "--", pathspec)
|
||||
}
|
||||
text, err := runGit(ctx, t.workingDir, gitArgs...)
|
||||
if err != nil {
|
||||
return text, err
|
||||
}
|
||||
if strings.TrimSpace(text) == "" {
|
||||
if args.Staged {
|
||||
return "No staged diff.", nil
|
||||
}
|
||||
return "No unstaged diff.", nil
|
||||
}
|
||||
return text, nil
|
||||
}
|
||||
|
||||
// git_log
|
||||
|
||||
type GitLogTool struct {
|
||||
workingDir string
|
||||
}
|
||||
|
||||
func NewGitLogTool(workingDir string) *GitLogTool {
|
||||
return &GitLogTool{workingDir: workingDir}
|
||||
}
|
||||
|
||||
func (t *GitLogTool) Definition() llm.Tool {
|
||||
return llm.Tool{
|
||||
Type: "function",
|
||||
Function: llm.ToolFunction{
|
||||
Name: "git_log",
|
||||
Description: "Show recent git commits with git log --oneline --decorate.",
|
||||
Parameters: JSONSchema(`{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {"type": "string", "description": "Optional path to limit log to. Relative paths are resolved from the configured agent working directory."},
|
||||
"max_count": {"type": "integer", "description": "Maximum number of commits to show. Defaults to 20, capped at 100."}
|
||||
},
|
||||
"additionalProperties": false
|
||||
}`),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
type gitLogArgs struct {
|
||||
Path string `json:"path"`
|
||||
MaxCount int `json:"max_count"`
|
||||
}
|
||||
|
||||
func (t *GitLogTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
||||
var args gitLogArgs
|
||||
if len(raw) > 0 {
|
||||
if err := json.Unmarshal(raw, &args); err != nil {
|
||||
return "", fmt.Errorf("parse arguments: %w", err)
|
||||
}
|
||||
}
|
||||
maxCount := normalizeGitLogMaxCount(args.MaxCount)
|
||||
gitArgs := []string{"log", "--oneline", "--decorate", fmt.Sprintf("-n%d", maxCount)}
|
||||
if args.Path != "" {
|
||||
pathspec, err := gitPathspec(t.workingDir, args.Path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
gitArgs = append(gitArgs, "--", pathspec)
|
||||
}
|
||||
text, err := runGit(ctx, t.workingDir, gitArgs...)
|
||||
if err != nil {
|
||||
return text, err
|
||||
}
|
||||
if strings.TrimSpace(text) == "" {
|
||||
return "No git history.", nil
|
||||
}
|
||||
return text, nil
|
||||
}
|
||||
@@ -0,0 +1,142 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func requireGit(t *testing.T) {
|
||||
t.Helper()
|
||||
if _, err := exec.LookPath("git"); err != nil {
|
||||
t.Skip("git not installed")
|
||||
}
|
||||
}
|
||||
|
||||
func runGitForTest(t *testing.T, dir string, args ...string) {
|
||||
t.Helper()
|
||||
gitArgs := append([]string{"-C", dir}, args...)
|
||||
cmd := exec.Command("git", gitArgs...)
|
||||
out, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
t.Fatalf("git %v failed: %v\n%s", args, err, out)
|
||||
}
|
||||
}
|
||||
|
||||
func initTestRepo(t *testing.T) string {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
runGitForTest(t, dir, "init")
|
||||
runGitForTest(t, dir, "-c", "user.email=test@example.com", "-c", "user.name=Test", "commit", "--allow-empty", "-m", "initial")
|
||||
return dir
|
||||
}
|
||||
|
||||
func TestGitStatusToolShowsChangedFiles(t *testing.T) {
|
||||
requireGit(t)
|
||||
dir := initTestRepo(t)
|
||||
|
||||
path := filepath.Join(dir, "tracked.txt")
|
||||
if err := os.WriteFile(path, []byte("hello"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
runGitForTest(t, dir, "add", "tracked.txt")
|
||||
runGitForTest(t, dir, "-c", "user.email=test@example.com", "-c", "user.name=Test", "commit", "-m", "add tracked")
|
||||
|
||||
if err := os.WriteFile(path, []byte("world"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tool := NewGitStatusTool(dir)
|
||||
out, err := tool.Execute(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(out, "Changed files:") {
|
||||
t.Fatalf("output missing 'Changed files:': %q", out)
|
||||
}
|
||||
if !strings.Contains(out, "M tracked.txt") {
|
||||
t.Fatalf("output missing 'M tracked.txt': %q", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGitDiffToolShowsUnstagedDiff(t *testing.T) {
|
||||
requireGit(t)
|
||||
dir := initTestRepo(t)
|
||||
|
||||
path := filepath.Join(dir, "tracked.txt")
|
||||
if err := os.WriteFile(path, []byte("hello"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
runGitForTest(t, dir, "add", "tracked.txt")
|
||||
runGitForTest(t, dir, "-c", "user.email=test@example.com", "-c", "user.name=Test", "commit", "-m", "add tracked")
|
||||
|
||||
if err := os.WriteFile(path, []byte("world"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tool := NewGitDiffTool(dir)
|
||||
out, err := tool.Execute(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(out, "diff --git") {
|
||||
t.Fatalf("output missing 'diff --git': %q", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGitDiffToolShowsStat(t *testing.T) {
|
||||
requireGit(t)
|
||||
dir := initTestRepo(t)
|
||||
|
||||
path := filepath.Join(dir, "tracked.txt")
|
||||
if err := os.WriteFile(path, []byte("hello"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
runGitForTest(t, dir, "add", "tracked.txt")
|
||||
runGitForTest(t, dir, "-c", "user.email=test@example.com", "-c", "user.name=Test", "commit", "-m", "add tracked")
|
||||
|
||||
if err := os.WriteFile(path, []byte("world"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tool := NewGitDiffTool(dir)
|
||||
out, err := tool.Execute(context.Background(), json.RawMessage(`{"stat":true}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(out, "tracked.txt") {
|
||||
t.Fatalf("output missing 'tracked.txt': %q", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGitLogToolShowsCommits(t *testing.T) {
|
||||
requireGit(t)
|
||||
dir := initTestRepo(t)
|
||||
|
||||
tool := NewGitLogTool(dir)
|
||||
out, err := tool.Execute(context.Background(), json.RawMessage(`{"max_count":1}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(out, "initial") {
|
||||
t.Fatalf("output missing 'initial': %q", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGitToolsRejectEscapingPath(t *testing.T) {
|
||||
requireGit(t)
|
||||
dir := initTestRepo(t)
|
||||
|
||||
tool := NewGitStatusTool(dir)
|
||||
_, err := tool.Execute(context.Background(), json.RawMessage(`{"path":"../outside"}`))
|
||||
if err == nil {
|
||||
t.Fatal("expected error for escaping path")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "escapes") {
|
||||
t.Fatalf("error = %v", err)
|
||||
}
|
||||
}
|
||||
+41
-5
@@ -1,6 +1,7 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
@@ -8,6 +9,7 @@ import (
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"agentu/pkg/llm"
|
||||
)
|
||||
@@ -19,6 +21,7 @@ type ShellRunTool struct {
|
||||
func NewShellRunTool(workingDir string) *ShellRunTool {
|
||||
return &ShellRunTool{workingDir: workingDir}
|
||||
}
|
||||
func (*ShellRunTool) Mutates() bool { return true }
|
||||
|
||||
func (t *ShellRunTool) Definition() llm.Tool {
|
||||
return llm.Tool{
|
||||
@@ -45,6 +48,10 @@ type shellRunArgs struct {
|
||||
}
|
||||
|
||||
func (t *ShellRunTool) Execute(ctx context.Context, raw json.RawMessage) (string, error) {
|
||||
return t.ExecuteStream(ctx, raw, nil)
|
||||
}
|
||||
|
||||
func (t *ShellRunTool) ExecuteStream(ctx context.Context, raw json.RawMessage, emit func(string) error) (string, error) {
|
||||
var args shellRunArgs
|
||||
if err := json.Unmarshal(raw, &args); err != nil {
|
||||
return "", fmt.Errorf("parse arguments: %w", err)
|
||||
@@ -65,16 +72,45 @@ func (t *ShellRunTool) Execute(ctx context.Context, raw json.RawMessage) (string
|
||||
cmd := exec.CommandContext(ctx, shell, "-lc", args.Command)
|
||||
cmd.Dir = workdir
|
||||
|
||||
output, err := cmd.CombinedOutput()
|
||||
text := FormatResult(string(output))
|
||||
var mu sync.Mutex
|
||||
var combined bytes.Buffer
|
||||
writer := &shellStreamWriter{mu: &mu, buf: &combined, emit: emit}
|
||||
cmd.Stdout = writer
|
||||
cmd.Stderr = writer
|
||||
|
||||
runErr := cmd.Run()
|
||||
text := FormatResult(combined.String())
|
||||
if errors.Is(ctx.Err(), context.DeadlineExceeded) {
|
||||
return text, fmt.Errorf("command timed out: %w", ctx.Err())
|
||||
}
|
||||
if err != nil {
|
||||
if runErr != nil {
|
||||
if text != "" {
|
||||
return text, fmt.Errorf("command failed: %w", err)
|
||||
return text, fmt.Errorf("command failed: %w", runErr)
|
||||
}
|
||||
return "", fmt.Errorf("command failed: %w", err)
|
||||
return "", fmt.Errorf("command failed: %w", runErr)
|
||||
}
|
||||
return text, nil
|
||||
}
|
||||
|
||||
// shellStreamWriter captures command output for the final result while
|
||||
// forwarding every write to the streaming callback as it happens.
|
||||
type shellStreamWriter struct {
|
||||
mu *sync.Mutex
|
||||
buf *bytes.Buffer
|
||||
emit func(string) error
|
||||
}
|
||||
|
||||
func (w *shellStreamWriter) Write(p []byte) (int, error) {
|
||||
w.mu.Lock()
|
||||
n, writeErr := w.buf.Write(p)
|
||||
w.mu.Unlock()
|
||||
if writeErr != nil {
|
||||
return n, writeErr
|
||||
}
|
||||
if w.emit != nil {
|
||||
if err := w.emit(string(p)); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestShellRunStreamsOutputWhileRunning(t *testing.T) {
|
||||
tool := NewShellRunTool(t.TempDir())
|
||||
var mu sync.Mutex
|
||||
var chunks []string
|
||||
output, err := tool.ExecuteStream(context.Background(), json.RawMessage(`{"command":"printf 'a\\n'; sleep 0.05; printf 'b\\n'"}`), func(chunk string) error {
|
||||
mu.Lock()
|
||||
chunks = append(chunks, chunk)
|
||||
mu.Unlock()
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(output, "a") || !strings.Contains(output, "b") {
|
||||
t.Fatalf("output = %q", output)
|
||||
}
|
||||
mu.Lock()
|
||||
joined := strings.Join(chunks, "")
|
||||
mu.Unlock()
|
||||
if len(chunks) == 0 {
|
||||
t.Fatal("expected at least one streamed chunk")
|
||||
}
|
||||
if !strings.Contains(joined, "a") || !strings.Contains(joined, "b") {
|
||||
t.Fatalf("streamed chunks = %q", joined)
|
||||
}
|
||||
}
|
||||
|
||||
func TestShellRunExecuteCollectsFullOutput(t *testing.T) {
|
||||
tool := NewShellRunTool(t.TempDir())
|
||||
output, err := tool.Execute(context.Background(), json.RawMessage(`{"command":"printf 'hello world'"}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.TrimSpace(output) != "hello world" {
|
||||
t.Fatalf("output = %q", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestShellRunStreamsStderrAndStdout(t *testing.T) {
|
||||
tool := NewShellRunTool(t.TempDir())
|
||||
var mu sync.Mutex
|
||||
var chunks []string
|
||||
output, err := tool.ExecuteStream(context.Background(), json.RawMessage(`{"command":"echo out; echo err >&2"}`), func(chunk string) error {
|
||||
mu.Lock()
|
||||
chunks = append(chunks, chunk)
|
||||
mu.Unlock()
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(output, "out") || !strings.Contains(output, "err") {
|
||||
t.Fatalf("output = %q", output)
|
||||
}
|
||||
mu.Lock()
|
||||
joined := strings.Join(chunks, "")
|
||||
mu.Unlock()
|
||||
if !strings.Contains(joined, "out") || !strings.Contains(joined, "err") {
|
||||
t.Fatalf("streamed chunks = %q", joined)
|
||||
}
|
||||
}
|
||||
+36
-3
@@ -19,6 +19,27 @@ type Tool interface {
|
||||
Execute(ctx context.Context, args json.RawMessage) (string, error)
|
||||
}
|
||||
|
||||
// StreamingTool is optionally implemented by tools that can report partial
|
||||
// output while they execute. emit receives chunks of the tool's output as
|
||||
// they are produced; the returned string is still the complete result that
|
||||
// gets sent back to the model.
|
||||
type StreamingTool interface {
|
||||
Tool
|
||||
ExecuteStream(ctx context.Context, args json.RawMessage, emit func(string) error) (string, error)
|
||||
}
|
||||
|
||||
// MutatingTool is optionally implemented by tools that modify external state.
|
||||
// It is used by the agent to determine whether tool calls can run in parallel.
|
||||
type MutatingTool interface {
|
||||
Mutates() bool
|
||||
}
|
||||
|
||||
// IsMutating reports whether tool implements MutatingTool and returns true.
|
||||
func IsMutating(tool Tool) bool {
|
||||
mutating, ok := tool.(MutatingTool)
|
||||
return ok && mutating.Mutates()
|
||||
}
|
||||
|
||||
type Registry struct {
|
||||
tools map[string]Tool
|
||||
primary []string
|
||||
@@ -66,30 +87,39 @@ func ReadOnlyBuiltins(workingDir string) *Registry {
|
||||
fileList := NewFileListTool(workingDir)
|
||||
fileSearch := NewFileSearchTool(workingDir)
|
||||
codeSymbols := NewCodeSymbolsTool(workingDir)
|
||||
gitStatus := NewGitStatusTool(workingDir)
|
||||
gitDiff := NewGitDiffTool(workingDir)
|
||||
gitLog := NewGitLogTool(workingDir)
|
||||
|
||||
registry := NewRegistry(fileRead, fileList, fileSearch, codeSymbols)
|
||||
registry := NewRegistry(fileRead, fileList, fileSearch, codeSymbols, gitStatus, gitDiff, gitLog)
|
||||
registry.RegisterAlias("file.read", fileRead)
|
||||
registry.RegisterAlias("file.list", fileList)
|
||||
registry.RegisterAlias("file.search", fileSearch)
|
||||
registry.RegisterAlias("code.symbols", codeSymbols)
|
||||
registry.RegisterAlias("git.status", gitStatus)
|
||||
registry.RegisterAlias("git.diff", gitDiff)
|
||||
registry.RegisterAlias("git.log", gitLog)
|
||||
return registry
|
||||
}
|
||||
|
||||
func Builtins(workingDir string) *Registry {
|
||||
registry := ReadOnlyBuiltins(workingDir)
|
||||
fileWrite := NewFileWriteTool(workingDir)
|
||||
fileEdit := NewFileEditTool(workingDir)
|
||||
shellRun := NewShellRunTool(workingDir)
|
||||
registry.Register(fileWrite)
|
||||
registry.Register(fileEdit)
|
||||
registry.Register(shellRun)
|
||||
registry.RegisterAlias("file.write", fileWrite)
|
||||
registry.RegisterAlias("file.edit", fileEdit)
|
||||
registry.RegisterAlias("shell.run", shellRun)
|
||||
return registry
|
||||
}
|
||||
|
||||
func AgentInstructions(workingDir string, yolo bool) string {
|
||||
searchInstructions := webSearchInstructions()
|
||||
readOnlyTools := "file_list, file_read, file_search, code_symbols, web_search, web_fetch"
|
||||
yoloTools := "file_list, file_read, file_search, code_symbols, file_write, shell_run, web_search, web_fetch"
|
||||
readOnlyTools := "file_list, file_read, file_search, code_symbols, git_status, git_diff, git_log, web_search, web_fetch"
|
||||
yoloTools := "file_list, file_read, file_search, code_symbols, file_write, file_edit, shell_run, git_status, git_diff, git_log, web_search, web_fetch"
|
||||
if yolo {
|
||||
return fmt.Sprintf(`Local project context:
|
||||
- The project working directory is %q.
|
||||
@@ -99,7 +129,9 @@ func AgentInstructions(workingDir string, yolo bool) string {
|
||||
- When the user asks to inspect, list, search, or read project files, use file_list, file_search, and file_read instead of saying you cannot access local files.
|
||||
- When the user asks about code structure, symbols, functions, methods, types, or definitions in Go files, use code_symbols before reading whole files.
|
||||
- When the user asks to run, execute, test, check, debug, or inspect a shell command, use shell_run. This includes local commands written in backticks, such as go, make, git, pwd, ls, ps, and similar terminal commands.
|
||||
- When editing existing files, prefer file_edit for small or localized changes; use file_write only for creating files, replacing whole files intentionally, or appending.
|
||||
- Do not use file tools as a substitute for a command execution request.
|
||||
- When the user asks about changed files, diffs, or git history, use git_status, git_diff, or git_log instead of shell_run.
|
||||
- Do not start interactive or long-running sessions. Prefer bounded, non-interactive local commands with clear completion conditions.
|
||||
- %s
|
||||
- Tool paths should normally be relative to the project working directory.
|
||||
@@ -113,6 +145,7 @@ func AgentInstructions(workingDir string, yolo bool) string {
|
||||
- When the user asks to inspect, list, search, or read project files, use file_list, file_search, and file_read instead of saying you cannot access local files.
|
||||
- When the user asks about code structure, symbols, functions, methods, types, or definitions in Go files, use code_symbols before reading whole files.
|
||||
- Shell command execution is disabled in this mode. When the user asks to run, execute, test, check, debug, or inspect a shell command, do not use file tools as a substitute; explain that they need to restart agentu with --yolo to enable shell_run.
|
||||
- When the user asks about changed files, diffs, or git history, use git_status, git_diff, or git_log instead of shell_run.
|
||||
- %s
|
||||
- Tool paths should normally be relative to the project working directory.
|
||||
- Available tools: %s.`, workingDir, searchInstructions, readOnlyTools)
|
||||
|
||||
@@ -104,12 +104,12 @@ func TestReadOnlyBuiltinsExposeOpenAICompatibleNamesAndAliases(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
for _, want := range []string{"file_list", "file_read", "file_search", "code_symbols"} {
|
||||
for _, want := range []string{"file_list", "file_read", "file_search", "code_symbols", "git_status", "git_diff", "git_log"} {
|
||||
if !slices.Contains(names, want) {
|
||||
t.Fatalf("definitions missing %s: %#v", want, names)
|
||||
}
|
||||
}
|
||||
for _, alias := range []string{"file.list", "file.read", "file.search", "code.symbols"} {
|
||||
for _, alias := range []string{"file.list", "file.read", "file.search", "code.symbols", "git.status", "git.diff", "git.log"} {
|
||||
if _, ok := registry.Get(alias); !ok {
|
||||
t.Fatalf("alias missing: %s", alias)
|
||||
}
|
||||
@@ -117,6 +117,9 @@ func TestReadOnlyBuiltinsExposeOpenAICompatibleNamesAndAliases(t *testing.T) {
|
||||
if _, ok := registry.Get("file_write"); ok {
|
||||
t.Fatal("read-only registry should not expose file_write")
|
||||
}
|
||||
if _, ok := registry.Get("file_edit"); ok {
|
||||
t.Fatal("read-only registry should not expose file_edit")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodeSymbolsToolListsGoSymbols(t *testing.T) {
|
||||
@@ -170,7 +173,7 @@ func TestRegistryCanExposeWebSearch(t *testing.T) {
|
||||
|
||||
func TestBuiltinsExposeYoloTools(t *testing.T) {
|
||||
registry := Builtins(t.TempDir())
|
||||
for _, name := range []string{"file_write", "shell_run", "file.write", "shell.run"} {
|
||||
for _, name := range []string{"file_write", "file_edit", "shell_run", "file.write", "file.edit", "shell.run"} {
|
||||
if _, ok := registry.Get(name); !ok {
|
||||
t.Fatalf("tool missing: %s", name)
|
||||
}
|
||||
@@ -179,14 +182,14 @@ func TestBuiltinsExposeYoloTools(t *testing.T) {
|
||||
|
||||
func TestAgentInstructionsExplainCommandModes(t *testing.T) {
|
||||
readOnly := AgentInstructions("/tmp/project", false)
|
||||
for _, want := range []string{"Shell command execution is disabled", "--yolo", "Do not call tools for greetings", "Do not explore the project preemptively", "Do not call file_list just to discover context"} {
|
||||
for _, want := range []string{"Shell command execution is disabled", "--yolo", "Do not call tools for greetings", "Do not explore the project preemptively", "Do not call file_list just to discover context", "git_status, git_diff, git_log"} {
|
||||
if !strings.Contains(readOnly, want) {
|
||||
t.Fatalf("read-only instructions missing %q:\n%s", want, readOnly)
|
||||
}
|
||||
}
|
||||
|
||||
yolo := AgentInstructions("/tmp/project", true)
|
||||
for _, want := range []string{"use shell_run", "commands written in backticks", "Do not start interactive", "non-interactive local commands", "Do not call tools for greetings"} {
|
||||
for _, want := range []string{"use shell_run", "commands written in backticks", "Do not start interactive", "non-interactive local commands", "Do not call tools for greetings", "prefer file_edit"} {
|
||||
if !strings.Contains(yolo, want) {
|
||||
t.Fatalf("yolo instructions missing %q:\n%s", want, yolo)
|
||||
}
|
||||
@@ -195,7 +198,7 @@ func TestAgentInstructionsExplainCommandModes(t *testing.T) {
|
||||
|
||||
func TestAgentInstructionsExplainWebSearch(t *testing.T) {
|
||||
withSearch := AgentInstructions("/tmp/project", false)
|
||||
for _, want := range []string{"use web_search", "Use web_fetch", "latest information", "source URLs", "Available tools: file_list, file_read, file_search, code_symbols, web_search, web_fetch"} {
|
||||
for _, want := range []string{"use web_search", "Use web_fetch", "latest information", "source URLs", "Available tools: file_list, file_read, file_search, code_symbols, git_status, git_diff, git_log, web_search, web_fetch"} {
|
||||
if !strings.Contains(withSearch, want) {
|
||||
t.Fatalf("web search instructions missing %q:\n%s", want, withSearch)
|
||||
}
|
||||
@@ -257,3 +260,16 @@ func TestResolvePathRejectsDotDotThatEscapes(t *testing.T) {
|
||||
t.Fatal("expected error for path escaping via ..")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsMutating(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if IsMutating(NewFileReadTool(dir)) {
|
||||
t.Fatal("FileReadTool should not be mutating")
|
||||
}
|
||||
if !IsMutating(NewFileWriteTool(dir)) {
|
||||
t.Fatal("FileWriteTool should be mutating")
|
||||
}
|
||||
if !IsMutating(NewShellRunTool(dir)) {
|
||||
t.Fatal("ShellRunTool should be mutating")
|
||||
}
|
||||
}
|
||||
|
||||
+527
-98
@@ -50,6 +50,7 @@ type ModelManager interface {
|
||||
Rename(name string) (string, error)
|
||||
Sessions() (string, error)
|
||||
NewSession() (string, error)
|
||||
AvailableModels() []string
|
||||
}
|
||||
|
||||
func Run(ctx context.Context, assistant *agent.Agent, opts Options) error {
|
||||
@@ -96,6 +97,7 @@ type model struct {
|
||||
inputHistory []string
|
||||
historyIndex int
|
||||
draftInput string
|
||||
commandCursor int
|
||||
|
||||
running bool
|
||||
cancel context.CancelFunc
|
||||
@@ -133,6 +135,13 @@ type commandMatch struct {
|
||||
score int
|
||||
}
|
||||
|
||||
type commandSuggestion struct {
|
||||
label string
|
||||
description string
|
||||
value string
|
||||
executable bool
|
||||
}
|
||||
|
||||
var slashCommands = []slashCommand{
|
||||
{Name: "/clear", Description: "reset context"},
|
||||
{Name: "/compact", Description: "compress context"},
|
||||
@@ -145,7 +154,7 @@ var slashCommands = []slashCommand{
|
||||
{Name: "/new", Description: "new session"},
|
||||
{Name: "/status", Description: "show changed files"},
|
||||
{Name: "/diff", Description: "show unstaged diff"},
|
||||
{Name: "/test", Description: "run project tests"},
|
||||
{Name: "/theme", Description: "switch theme light|dark"},
|
||||
}
|
||||
|
||||
func newModel(ctx context.Context, assistant *agent.Agent, opts Options) model {
|
||||
@@ -157,7 +166,7 @@ func newModel(ctx context.Context, assistant *agent.Agent, opts Options) model {
|
||||
st := th.styles()
|
||||
|
||||
input := textarea.New()
|
||||
input.Placeholder = "Message agentu..."
|
||||
input.Placeholder = ""
|
||||
input.Prompt = ""
|
||||
input.ShowLineNumbers = false
|
||||
input.MaxHeight = maxInputLines
|
||||
@@ -197,8 +206,6 @@ func (m model) Init() tea.Cmd {
|
||||
m.syncMessages()
|
||||
return textarea.Blink
|
||||
}
|
||||
|
||||
// syncMessages rebuilds the TUI display messages from the agent's conversation history.
|
||||
func (m *model) syncMessages() {
|
||||
if m.agent == nil {
|
||||
return
|
||||
@@ -246,6 +253,9 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
if m.picking {
|
||||
return m.updatePicker(msg)
|
||||
}
|
||||
if handled, next, cmd := m.updateCommandSuggestions(msg); handled {
|
||||
return next, cmd
|
||||
}
|
||||
switch msg.String() {
|
||||
case "ctrl+c":
|
||||
if m.cancel != nil {
|
||||
@@ -263,6 +273,29 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
return m, nil
|
||||
}
|
||||
return m.submit()
|
||||
case "up":
|
||||
if m.running {
|
||||
return m, nil
|
||||
}
|
||||
if m.input.Line() == 0 {
|
||||
m.previousInputHistory()
|
||||
m.afterInputChanged()
|
||||
return m, nil
|
||||
}
|
||||
case "down":
|
||||
if m.running {
|
||||
return m, nil
|
||||
}
|
||||
value := m.input.Value()
|
||||
lineCount := 1
|
||||
if value != "" {
|
||||
lineCount = len(strings.Split(value, "\n"))
|
||||
}
|
||||
if m.input.Line() >= lineCount-1 {
|
||||
m.nextInputHistory()
|
||||
m.afterInputChanged()
|
||||
return m, nil
|
||||
}
|
||||
case "ctrl+p":
|
||||
if m.running {
|
||||
return m, nil
|
||||
@@ -277,7 +310,7 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
m.nextInputHistory()
|
||||
m.afterInputChanged()
|
||||
return m, nil
|
||||
case "alt+enter", "ctrl+j":
|
||||
case "ctrl+j":
|
||||
return m.insertInputNewline()
|
||||
}
|
||||
|
||||
@@ -290,8 +323,10 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
case toolLogMsg:
|
||||
text := cleanToolLog(string(msg))
|
||||
if text != "" {
|
||||
if !m.mergeToolLog(text) {
|
||||
m.messages = append(m.messages, message{role: roleTool, content: text})
|
||||
m.status = text
|
||||
}
|
||||
m.status = toolStatusText(text)
|
||||
m.refreshViewport(true)
|
||||
}
|
||||
return m, m.waitForEvent()
|
||||
@@ -336,7 +371,6 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
if !m.running {
|
||||
nextInput, cmd := m.input.Update(msg)
|
||||
m.input = nextInput
|
||||
m.afterInputChanged()
|
||||
cmds = append(cmds, cmd)
|
||||
}
|
||||
|
||||
@@ -354,7 +388,7 @@ func (m model) View() string {
|
||||
|
||||
if m.picking {
|
||||
parts := []string{m.viewport.View()}
|
||||
parts = append(parts, m.pickerView(), m.modelMetaView(), m.footerView())
|
||||
parts = append(parts, m.pickerView(), m.footerView())
|
||||
body := lipgloss.JoinVertical(lipgloss.Left, parts...)
|
||||
return m.styles.App.Width(m.width).Height(m.height).Render(body)
|
||||
}
|
||||
@@ -363,7 +397,7 @@ func (m model) View() string {
|
||||
if activity := m.activityView(); activity != "" {
|
||||
parts = append(parts, activity)
|
||||
}
|
||||
parts = append(parts, m.inputView(), m.modelMetaView(), m.footerView())
|
||||
parts = append(parts, m.inputView(), m.footerView())
|
||||
body := lipgloss.JoinVertical(lipgloss.Left, parts...)
|
||||
return m.styles.App.Width(m.width).Height(m.height).Render(body)
|
||||
}
|
||||
@@ -387,7 +421,11 @@ func (m *model) submit() (tea.Model, tea.Cmd) {
|
||||
m.status = "ready"
|
||||
m.refreshViewport(true)
|
||||
return *m, nil
|
||||
case "/model", "/compact":
|
||||
case "/model", "/compact", "/theme":
|
||||
return m.runLocalCommand(input)
|
||||
}
|
||||
|
||||
if strings.HasPrefix(input, "!") {
|
||||
return m.runLocalCommand(input)
|
||||
}
|
||||
|
||||
@@ -436,6 +474,16 @@ func (m *model) insertInputNewline() (tea.Model, tea.Cmd) {
|
||||
m.historyIndex = -1
|
||||
m.draftInput = ""
|
||||
m.afterInputChanged()
|
||||
// The textarea's repositionView was called inside Update with the old
|
||||
// viewport height, so the YOffset may be stale after syncInputHeight
|
||||
// changed the height. Send up then down to force repositionView to
|
||||
// recalculate with the new height.
|
||||
up := tea.KeyMsg{Type: tea.KeyUp}
|
||||
vp, _ := m.input.Update(up)
|
||||
m.input = vp
|
||||
down := tea.KeyMsg{Type: tea.KeyDown}
|
||||
vp, _ = m.input.Update(down)
|
||||
m.input = vp
|
||||
return *m, cmd
|
||||
}
|
||||
|
||||
@@ -465,7 +513,7 @@ func (m *model) appendAssistantChunk(chunk string) {
|
||||
|
||||
func (m *model) resize() {
|
||||
width := max(40, m.width)
|
||||
inputWidth := max(20, width-m.styles.InputBox.GetHorizontalFrameSize())
|
||||
inputWidth := max(20, width-m.styles.InputBox.GetHorizontalFrameSize()-m.inputModeWidth()-1)
|
||||
m.input.SetWidth(inputWidth)
|
||||
m.syncInputHeight()
|
||||
|
||||
@@ -574,6 +622,7 @@ func roleLabel(r role) string {
|
||||
type toolLogParts struct {
|
||||
name string
|
||||
args string
|
||||
output string
|
||||
}
|
||||
|
||||
func parseToolLog(content string) toolLogParts {
|
||||
@@ -581,14 +630,21 @@ func parseToolLog(content string) toolLogParts {
|
||||
if content == "" {
|
||||
return toolLogParts{}
|
||||
}
|
||||
idx := strings.IndexFunc(content, unicode.IsSpace)
|
||||
beforeOutput, output, hasOutput := strings.Cut(content, "\n[output]\n")
|
||||
parts := toolLogParts{}
|
||||
if hasOutput {
|
||||
parts.output = strings.TrimSpace(output)
|
||||
}
|
||||
firstLine := strings.SplitN(beforeOutput, "\n", 2)[0]
|
||||
firstLine = strings.TrimSpace(firstLine)
|
||||
idx := strings.IndexFunc(firstLine, unicode.IsSpace)
|
||||
if idx < 0 {
|
||||
return toolLogParts{name: content}
|
||||
}
|
||||
return toolLogParts{
|
||||
name: strings.TrimSpace(content[:idx]),
|
||||
args: strings.TrimSpace(content[idx+1:]),
|
||||
parts.name = firstLine
|
||||
return parts
|
||||
}
|
||||
parts.name = strings.TrimSpace(firstLine[:idx])
|
||||
parts.args = strings.TrimSpace(firstLine[idx+1:])
|
||||
return parts
|
||||
}
|
||||
|
||||
func (m model) toolMessageView(content string, width int) string {
|
||||
@@ -599,82 +655,96 @@ func (m model) toolMessageView(content string, width int) string {
|
||||
args := fitLine(parts.args, max(1, width-4))
|
||||
lines = append(lines, m.styles.ToolValue.Render(args))
|
||||
}
|
||||
if parts.output != "" {
|
||||
lines = append(lines, m.styles.ToolKey.Render("output"))
|
||||
previewLines := toolOutputPreviewLines(parts.output, max(1, width-4))
|
||||
for _, line := range previewLines {
|
||||
lines = append(lines, m.styles.ToolValue.Render(line))
|
||||
}
|
||||
}
|
||||
bodyContent := strings.Join(lines, "\n")
|
||||
bodyWidth := messageBodyWidth(bodyContent, width, m.styles.ToolMessage)
|
||||
return m.styles.ToolMessage.Width(bodyWidth).Render(bodyContent)
|
||||
}
|
||||
|
||||
const maxToolOutputPreviewLines = 8
|
||||
|
||||
func toolStatusText(content string) string {
|
||||
content = strings.TrimSpace(content)
|
||||
if content == "" {
|
||||
return ""
|
||||
}
|
||||
lines := strings.SplitN(content, "\n", 2)
|
||||
return lines[0]
|
||||
}
|
||||
|
||||
func (m *model) mergeToolLog(text string) bool {
|
||||
incoming := parseToolLog(text)
|
||||
if incoming.output == "" {
|
||||
return false
|
||||
}
|
||||
incomingHeader, _, ok := strings.Cut(text, "\n[output]\n")
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
for i := len(m.messages) - 1; i >= 0; i-- {
|
||||
if m.messages[i].role != roleTool {
|
||||
continue
|
||||
}
|
||||
existing := parseToolLog(m.messages[i].content)
|
||||
if existing.name != incoming.name || existing.args != incoming.args {
|
||||
continue
|
||||
}
|
||||
if !strings.Contains(m.messages[i].content, "\n[output]\n") {
|
||||
// Placeholder card from the tool-start log; first chunk fills it.
|
||||
m.messages[i].content = text
|
||||
} else {
|
||||
// Streaming tools send incremental chunks; append the raw delta
|
||||
// so line boundaries are preserved inside the same card.
|
||||
m.messages[i].content += strings.TrimPrefix(text, incomingHeader+"\n[output]\n")
|
||||
}
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func toolOutputPreviewLines(output string, width int) []string {
|
||||
lines := strings.Split(output, "\n")
|
||||
if len(lines) > maxToolOutputPreviewLines {
|
||||
lines = lines[:maxToolOutputPreviewLines]
|
||||
lines = append(lines, fmt.Sprintf("[truncated: showing first %d lines]", maxToolOutputPreviewLines))
|
||||
}
|
||||
result := make([]string, 0, len(lines))
|
||||
for _, line := range lines {
|
||||
result = append(result, fitLine(line, width))
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
type metaPill struct {
|
||||
text string
|
||||
accent bool
|
||||
}
|
||||
|
||||
func (m model) renderMetaPills(parts []metaPill) string {
|
||||
if len(parts) == 0 {
|
||||
return ""
|
||||
func (m model) inputModeView() string {
|
||||
if strings.HasPrefix(strings.TrimSpace(m.input.Value()), "!") {
|
||||
return m.styles.InputModeCommand.Render(iconCommand)
|
||||
}
|
||||
|
||||
plain := make([]string, len(parts))
|
||||
for i, part := range parts {
|
||||
plain[i] = part.text
|
||||
}
|
||||
|
||||
width := max(1, m.width-2)
|
||||
visible := len(plain)
|
||||
for visible > 1 && lipgloss.Width(strings.Join(plain[:visible], " ")) > width {
|
||||
visible--
|
||||
}
|
||||
if lipgloss.Width(strings.Join(plain[:visible], " ")) > width {
|
||||
return m.styles.PillAccent.Render(fitLine(plain[0], width))
|
||||
}
|
||||
|
||||
rendered := make([]string, 0, visible)
|
||||
for i := 0; i < visible; i++ {
|
||||
style := m.styles.Pill
|
||||
if parts[i].accent {
|
||||
style = m.styles.PillAccent
|
||||
}
|
||||
rendered = append(rendered, style.Render(parts[i].text))
|
||||
}
|
||||
return strings.Join(rendered, " ")
|
||||
}
|
||||
|
||||
func (m model) modelMetaView() string {
|
||||
mode := "read-only tools"
|
||||
if m.yolo {
|
||||
mode = "yolo tools"
|
||||
return m.styles.InputModeYolo.Render("Y")
|
||||
}
|
||||
return m.styles.InputModeReadonly.Render("R")
|
||||
}
|
||||
|
||||
modelName := m.modelName
|
||||
if m.models != nil {
|
||||
if currentModel := m.models.CurrentModel(); currentModel != "" {
|
||||
modelName = currentModel
|
||||
}
|
||||
}
|
||||
|
||||
parts := []metaPill{{text: iconModel + " " + modelName, accent: true}}
|
||||
if m.models != nil {
|
||||
if provider := m.models.CurrentProvider(); provider != "" {
|
||||
parts = append(parts, metaPill{text: "provider " + provider})
|
||||
}
|
||||
if currentThinking := m.models.CurrentThinking(); currentThinking != "" && currentThinking != "unset" {
|
||||
parts = append(parts, metaPill{text: "thinking " + currentThinking})
|
||||
}
|
||||
if sessionName := m.models.CurrentSessionName(); sessionName != "" {
|
||||
parts = append(parts, metaPill{text: iconSession + " " + sessionName})
|
||||
}
|
||||
}
|
||||
if m.contextUsage != "" {
|
||||
parts = append(parts, metaPill{text: iconContext + " " + m.contextUsage})
|
||||
}
|
||||
parts = append(parts, metaPill{text: mode}, metaPill{text: "theme " + string(m.theme.Mode)})
|
||||
|
||||
return m.styles.ModelMeta.Width(max(1, m.width)).Render(m.renderMetaPills(parts))
|
||||
func (m model) inputModeWidth() int {
|
||||
return lipgloss.Width(m.inputModeView())
|
||||
}
|
||||
|
||||
func (m model) inputView() string {
|
||||
palette := m.commandPaletteView()
|
||||
box := m.styles.InputBox.Width(max(20, m.width-2)).Render(m.input.View())
|
||||
gap := lipgloss.NewStyle().Background(m.styles.InputBox.GetBackground()).Render(" ")
|
||||
content := lipgloss.JoinHorizontal(lipgloss.Top, m.inputModeView(), gap, m.input.View())
|
||||
box := m.styles.InputBox.Width(max(20, m.width)).Render(content)
|
||||
if palette != "" {
|
||||
return lipgloss.JoinVertical(lipgloss.Left, palette, box)
|
||||
}
|
||||
@@ -692,13 +762,85 @@ func (m model) activityView() string {
|
||||
line := m.styles.StatusInfo.Render(iconWorking) + " Working " + m.spinner.View() + " " + status + " · esc cancels"
|
||||
return m.styles.Activity.Render(fitLine(line, max(1, m.width-2)))
|
||||
}
|
||||
|
||||
func (m model) footerView() string {
|
||||
hint := "enter send · ctrl+j newline · alt+enter newline · ctrl+p/ctrl+n history · / commands"
|
||||
hint := ""
|
||||
if m.running {
|
||||
hint = "esc cancel · ctrl+c quit"
|
||||
}
|
||||
return m.styles.Footer.Width(max(1, m.width)).Render(fitLine(hint, max(1, m.width-2)))
|
||||
if m.showCommandPalette() {
|
||||
if sug, ok := m.selectedCommandSuggestion(); ok && sug.executable {
|
||||
hint = "↑/↓ select · enter apply · tab fill"
|
||||
}
|
||||
}
|
||||
|
||||
inner := max(1, m.width-2) // Footer padding(0,1)
|
||||
pills := m.modelMetaPillsClamped(inner / 3)
|
||||
pillsW := lipgloss.Width(pills)
|
||||
|
||||
var line string
|
||||
if pills != "" {
|
||||
remain := max(1, inner-pillsW-1)
|
||||
line = pills + " " + fitLine(hint, remain)
|
||||
} else {
|
||||
line = fitLine(hint, inner)
|
||||
}
|
||||
return m.styles.Footer.Width(max(1, m.width)).Render(line)
|
||||
}
|
||||
|
||||
func (m model) modelMetaPillsClamped(maxW int) string {
|
||||
modelName := m.modelName
|
||||
if m.models != nil {
|
||||
if currentModel := m.models.CurrentModel(); currentModel != "" {
|
||||
modelName = currentModel
|
||||
}
|
||||
}
|
||||
|
||||
parts := []metaPill{{text: iconModel + " " + modelName, accent: true}}
|
||||
if m.models != nil {
|
||||
if provider := m.models.CurrentProvider(); provider != "" {
|
||||
parts = append(parts, metaPill{text: provider})
|
||||
}
|
||||
if currentThinking := m.models.CurrentThinking(); currentThinking != "" && currentThinking != "unset" {
|
||||
parts = append(parts, metaPill{text: currentThinking})
|
||||
}
|
||||
if sessionName := m.models.CurrentSessionName(); sessionName != "" {
|
||||
parts = append(parts, metaPill{text: iconSession + " " + sessionName})
|
||||
}
|
||||
}
|
||||
if m.contextUsage != "" {
|
||||
parts = append(parts, metaPill{text: iconContext + " " + m.contextUsage})
|
||||
}
|
||||
|
||||
return m.renderMetaPillsClamped(parts, maxW)
|
||||
}
|
||||
|
||||
func (m model) renderMetaPillsClamped(parts []metaPill, width int) string {
|
||||
if len(parts) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
plain := make([]string, len(parts))
|
||||
for i, part := range parts {
|
||||
plain[i] = part.text
|
||||
}
|
||||
|
||||
visible := len(plain)
|
||||
for visible > 1 && lipgloss.Width(strings.Join(plain[:visible], " ")) > width {
|
||||
visible--
|
||||
}
|
||||
if visible > 0 && lipgloss.Width(strings.Join(plain[:visible], " ")) > width {
|
||||
return ""
|
||||
}
|
||||
|
||||
rendered := make([]string, 0, visible)
|
||||
for i := 0; i < visible; i++ {
|
||||
style := m.styles.Pill
|
||||
if parts[i].accent {
|
||||
style = m.styles.PillAccent
|
||||
}
|
||||
rendered = append(rendered, style.Render(parts[i].text))
|
||||
}
|
||||
return strings.Join(rendered, " ")
|
||||
}
|
||||
|
||||
type eventWriter struct {
|
||||
@@ -742,6 +884,7 @@ func isViewportKey(key string) bool {
|
||||
}
|
||||
|
||||
func (m *model) afterInputChanged() {
|
||||
m.commandCursor = 0
|
||||
m.syncInputHeight()
|
||||
if m.width > 0 && m.height > 0 {
|
||||
m.resize()
|
||||
@@ -823,18 +966,37 @@ func (m model) inputBlockHeight() int {
|
||||
}
|
||||
|
||||
func (m model) commandPaletteHeight() int {
|
||||
matches := commandMatches(commandPaletteQuery(m.input.Value()))
|
||||
count := len(matches)
|
||||
if count > maxCommandPaletteItems {
|
||||
count = maxCommandPaletteItems
|
||||
sugs := m.commandSuggestions()
|
||||
total := len(sugs)
|
||||
if !m.showCommandPalette() {
|
||||
return 0
|
||||
}
|
||||
count := min(total, maxCommandPaletteItems)
|
||||
if count == 0 {
|
||||
count = 1
|
||||
}
|
||||
return 1 + count + m.styles.Palette.GetVerticalFrameSize()
|
||||
scrollUp := 0
|
||||
scrollDown := 0
|
||||
if total > maxCommandPaletteItems {
|
||||
visible := min(total, maxCommandPaletteItems)
|
||||
scroll := max(0, m.commandCursor-visible+1)
|
||||
if m.commandCursor < scroll {
|
||||
scroll = m.commandCursor
|
||||
}
|
||||
if scroll > 0 {
|
||||
scrollUp = 1
|
||||
}
|
||||
if scroll+visible < total {
|
||||
scrollDown = 1
|
||||
}
|
||||
}
|
||||
return 1 + scrollUp + count + scrollDown + m.styles.Palette.GetVerticalFrameSize()
|
||||
}
|
||||
|
||||
func (m model) showCommandPalette() bool {
|
||||
if len(m.commandSuggestions()) > 0 {
|
||||
return true
|
||||
}
|
||||
value := strings.TrimSpace(m.input.Value())
|
||||
return strings.HasPrefix(value, "/") && !strings.Contains(value, " ")
|
||||
}
|
||||
@@ -843,18 +1005,42 @@ func (m model) commandPaletteView() string {
|
||||
if !m.showCommandPalette() {
|
||||
return ""
|
||||
}
|
||||
matches := commandMatches(commandPaletteQuery(m.input.Value()))
|
||||
sugs := m.commandSuggestions()
|
||||
if len(sugs) > 0 {
|
||||
header := iconCommand + " Commands"
|
||||
if sugs[0].executable {
|
||||
header = iconCommand + " Model candidates"
|
||||
}
|
||||
lines := []string{m.styles.Header.Render(header)}
|
||||
m.clampCommandCursor()
|
||||
total := len(sugs)
|
||||
visible := min(total, maxCommandPaletteItems)
|
||||
scroll := max(0, m.commandCursor-visible+1)
|
||||
if m.commandCursor < scroll {
|
||||
scroll = m.commandCursor
|
||||
}
|
||||
if scroll > 0 {
|
||||
lines = append(lines, m.styles.Muted.Render(" ↑ more"))
|
||||
}
|
||||
for i := scroll; i < scroll+visible && i < total; i++ {
|
||||
s := sugs[i]
|
||||
prefix := " "
|
||||
if i == m.commandCursor {
|
||||
prefix = "▸ "
|
||||
}
|
||||
lines = append(lines, fmt.Sprintf("%s%s %s", prefix, m.styles.PaletteMatch.Render(s.label), m.styles.Muted.Render(s.description)))
|
||||
}
|
||||
if scroll+visible < total {
|
||||
lines = append(lines, m.styles.Muted.Render(" ↓ more"))
|
||||
}
|
||||
return m.styles.Palette.Width(max(20, m.width-2)).Render(strings.Join(lines, "\n"))
|
||||
}
|
||||
value := strings.TrimSpace(m.input.Value())
|
||||
if strings.Contains(value, " ") {
|
||||
return ""
|
||||
}
|
||||
lines := []string{m.styles.Header.Render(iconCommand + " Commands")}
|
||||
if len(matches) == 0 {
|
||||
lines = append(lines, m.styles.Muted.Render("No matching commands"))
|
||||
} else {
|
||||
for i, match := range matches {
|
||||
if i >= maxCommandPaletteItems {
|
||||
break
|
||||
}
|
||||
lines = append(lines, fmt.Sprintf("%s %s", m.styles.PaletteMatch.Render(match.command.Name), m.styles.Muted.Render(match.command.Description)))
|
||||
}
|
||||
}
|
||||
return m.styles.Palette.Width(max(20, m.width-2)).Render(strings.Join(lines, "\n"))
|
||||
}
|
||||
|
||||
@@ -932,6 +1118,162 @@ func subsequenceScore(candidate string, query string, base int) (int, bool) {
|
||||
return base + gaps, true
|
||||
}
|
||||
|
||||
func (m model) commandSuggestions() []commandSuggestion {
|
||||
input := m.input.Value()
|
||||
if sugs := m.modelCommandSuggestions(input); len(sugs) > 0 {
|
||||
return sugs
|
||||
}
|
||||
value := strings.TrimSpace(input)
|
||||
if strings.HasPrefix(value, "/") && !strings.Contains(value, " ") {
|
||||
return slashCommandSuggestions(commandPaletteQuery(input))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func slashCommandSuggestions(query string) []commandSuggestion {
|
||||
matches := commandMatches(query)
|
||||
sugs := make([]commandSuggestion, 0, len(matches))
|
||||
for _, m := range matches {
|
||||
sugs = append(sugs, commandSuggestion{
|
||||
label: m.command.Name,
|
||||
description: m.command.Description,
|
||||
value: m.command.Name,
|
||||
})
|
||||
}
|
||||
return sugs
|
||||
}
|
||||
|
||||
func splitModelCommandTail(input string) (query string, ok bool) {
|
||||
fields := strings.Fields(input)
|
||||
if len(fields) == 0 || fields[0] != "/model" {
|
||||
return "", false
|
||||
}
|
||||
hasTrailingSpace := strings.TrimRight(input, " \t") != input
|
||||
switch len(fields) {
|
||||
case 1:
|
||||
if !hasTrailingSpace {
|
||||
return "", false
|
||||
}
|
||||
return "", true
|
||||
case 2:
|
||||
if fields[1] == "provider" || fields[1] == "thinking" {
|
||||
return "", false
|
||||
}
|
||||
if fields[1] == "model" {
|
||||
if !hasTrailingSpace {
|
||||
return "", false
|
||||
}
|
||||
return "", true
|
||||
}
|
||||
if hasTrailingSpace {
|
||||
return fields[1], true
|
||||
}
|
||||
return fields[1], true
|
||||
case 3:
|
||||
if fields[1] != "model" {
|
||||
return "", false
|
||||
}
|
||||
if hasTrailingSpace {
|
||||
return fields[2], true
|
||||
}
|
||||
return fields[2], true
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
|
||||
func (m model) modelCommandSuggestions(input string) []commandSuggestion {
|
||||
if m.models == nil {
|
||||
return nil
|
||||
}
|
||||
query, ok := splitModelCommandTail(input)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
available := m.models.AvailableModels()
|
||||
if len(available) == 0 {
|
||||
return nil
|
||||
}
|
||||
current := m.models.CurrentModel()
|
||||
candidates := make([]commandSuggestion, 0, len(available))
|
||||
for _, id := range available {
|
||||
desc := "switch model"
|
||||
if id == current {
|
||||
desc = "current model"
|
||||
}
|
||||
candidates = append(candidates, commandSuggestion{
|
||||
label: id,
|
||||
description: desc,
|
||||
value: "/model model " + id,
|
||||
executable: true,
|
||||
})
|
||||
}
|
||||
return scoredCommandSuggestions(candidates, query)
|
||||
}
|
||||
|
||||
func scoredCommandSuggestions(candidates []commandSuggestion, query string) []commandSuggestion {
|
||||
if query == "" {
|
||||
return candidates
|
||||
}
|
||||
type scored struct {
|
||||
sug commandSuggestion
|
||||
score int
|
||||
idx int
|
||||
}
|
||||
var kept []scored
|
||||
for i, c := range candidates {
|
||||
if s, ok := scoreCommandSuggestion(c.label, query); ok {
|
||||
kept = append(kept, scored{sug: c, score: s, idx: i})
|
||||
}
|
||||
}
|
||||
sort.Slice(kept, func(i, j int) bool {
|
||||
if kept[i].score != kept[j].score {
|
||||
return kept[i].score < kept[j].score
|
||||
}
|
||||
return kept[i].idx < kept[j].idx
|
||||
})
|
||||
out := make([]commandSuggestion, 0, len(kept))
|
||||
for _, s := range kept {
|
||||
out = append(out, s.sug)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func scoreCommandSuggestion(label string, query string) (int, bool) {
|
||||
query = strings.ToLower(strings.TrimSpace(query))
|
||||
name := strings.ToLower(label)
|
||||
if query == "" || query == name {
|
||||
return 0, true
|
||||
}
|
||||
if strings.HasPrefix(name, query) {
|
||||
return 100 + len(name) - len(query), true
|
||||
}
|
||||
return subsequenceScore(name, query, 200)
|
||||
}
|
||||
|
||||
func (m model) selectedCommandSuggestion() (commandSuggestion, bool) {
|
||||
sugs := m.commandSuggestions()
|
||||
m.clampCommandCursor()
|
||||
if m.commandCursor >= 0 && m.commandCursor < len(sugs) {
|
||||
return sugs[m.commandCursor], true
|
||||
}
|
||||
return commandSuggestion{}, false
|
||||
}
|
||||
|
||||
func (m *model) clampCommandCursor() {
|
||||
sugs := m.commandSuggestions()
|
||||
if len(sugs) == 0 {
|
||||
m.commandCursor = 0
|
||||
return
|
||||
}
|
||||
if m.commandCursor < 0 {
|
||||
m.commandCursor = 0
|
||||
}
|
||||
if m.commandCursor >= len(sugs) {
|
||||
m.commandCursor = len(sugs) - 1
|
||||
}
|
||||
}
|
||||
|
||||
func (m model) slashSuggestionCount() int {
|
||||
return m.commandPaletteHeight()
|
||||
}
|
||||
@@ -948,11 +1290,7 @@ func (m model) modelInfo() string {
|
||||
if m.models != nil {
|
||||
return m.models.Info()
|
||||
}
|
||||
mode := "read-only tools"
|
||||
if m.yolo {
|
||||
mode = "yolo tools"
|
||||
}
|
||||
return fmt.Sprintf("model: %s\nmode: %s\ntheme: %s", m.modelName, mode, m.theme.Mode)
|
||||
return m.modelName
|
||||
}
|
||||
|
||||
func (m *model) refreshContextUsage() {
|
||||
@@ -1031,6 +1369,9 @@ func fitLine(text string, width int) string {
|
||||
}
|
||||
|
||||
func (m *model) runLocalCommand(input string) (tea.Model, tea.Cmd) {
|
||||
if strings.HasPrefix(input, "/theme") {
|
||||
return m.handleThemeCommand(input)
|
||||
}
|
||||
output, err := m.handleLocalCommand(input)
|
||||
if err == errOpenPicker {
|
||||
m.openSessionPicker()
|
||||
@@ -1055,7 +1396,60 @@ func (m *model) runLocalCommand(input string) (tea.Model, tea.Cmd) {
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
func (m *model) handleThemeCommand(input string) (tea.Model, tea.Cmd) {
|
||||
fields := strings.Fields(input)
|
||||
if len(fields) < 2 {
|
||||
m.input.SetValue("")
|
||||
m.afterInputChanged()
|
||||
m.messages = append(m.messages, message{
|
||||
role: roleSystem, content: "theme: " + string(m.theme.Mode) + "\nusage: /theme light|dark",
|
||||
})
|
||||
m.refreshViewport(true)
|
||||
return *m, nil
|
||||
}
|
||||
mode, err := ParseThemeMode(fields[1])
|
||||
if err != nil {
|
||||
m.input.SetValue("")
|
||||
m.afterInputChanged()
|
||||
m.messages = append(m.messages, message{role: roleError, content: err.Error()})
|
||||
m.refreshViewport(true)
|
||||
return *m, nil
|
||||
}
|
||||
m.theme = themeForMode(mode)
|
||||
m.styles = m.theme.styles()
|
||||
applyTextareaTheme(&m.input, m.theme)
|
||||
m.input.SetValue("")
|
||||
m.afterInputChanged()
|
||||
m.messages = append(m.messages, message{role: roleSystem, content: "theme: " + string(mode)})
|
||||
m.refreshViewport(true)
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
func (m model) handleBangCommand(input string) (string, error) {
|
||||
command := strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(input), "!"))
|
||||
if command == "" {
|
||||
return "", fmt.Errorf("usage: ! <command>")
|
||||
}
|
||||
if !m.yolo {
|
||||
return "", fmt.Errorf("shell command input requires --yolo")
|
||||
}
|
||||
out, err := m.runShellCommand(command)
|
||||
if err != nil {
|
||||
if strings.TrimSpace(out) != "" {
|
||||
return "", fmt.Errorf("%s\n%w", strings.TrimRight(out, "\n"), err)
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
if strings.TrimSpace(out) == "" {
|
||||
return fmt.Sprintf("$ %s\n(no output)", command), nil
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (m model) handleLocalCommand(input string) (string, error) {
|
||||
if strings.HasPrefix(strings.TrimSpace(input), "!") {
|
||||
return m.handleBangCommand(input)
|
||||
}
|
||||
fields := strings.Fields(input)
|
||||
if len(fields) == 0 {
|
||||
return "", fmt.Errorf("empty command")
|
||||
@@ -1251,6 +1645,41 @@ func (m *model) updatePicker(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
func (m *model) updateCommandSuggestions(msg tea.KeyMsg) (bool, tea.Model, tea.Cmd) {
|
||||
if m.running || len(m.commandSuggestions()) == 0 {
|
||||
return false, *m, nil
|
||||
}
|
||||
m.clampCommandCursor()
|
||||
switch msg.String() {
|
||||
case "up":
|
||||
if m.commandCursor > 0 {
|
||||
m.commandCursor--
|
||||
}
|
||||
return true, *m, nil
|
||||
case "down":
|
||||
if m.commandCursor < len(m.commandSuggestions())-1 {
|
||||
m.commandCursor++
|
||||
}
|
||||
return true, *m, nil
|
||||
case "tab":
|
||||
if sug, ok := m.selectedCommandSuggestion(); ok {
|
||||
m.input.SetValue(sug.value)
|
||||
m.afterInputChanged()
|
||||
}
|
||||
return true, *m, nil
|
||||
case "enter":
|
||||
if sug, ok := m.selectedCommandSuggestion(); ok && sug.executable {
|
||||
m.rememberInput(sug.value)
|
||||
m.input.SetValue(sug.value)
|
||||
m.afterInputChanged()
|
||||
next, cmd := m.runLocalCommand(sug.value)
|
||||
return true, next, cmd
|
||||
}
|
||||
return false, *m, nil
|
||||
}
|
||||
return false, *m, nil
|
||||
}
|
||||
|
||||
func (m model) pickerView() string {
|
||||
if !m.picking || len(m.pickerItems) == 0 {
|
||||
return ""
|
||||
|
||||
+383
-23
@@ -21,7 +21,7 @@ func TestModelRendersChatShell(t *testing.T) {
|
||||
m = next.(model)
|
||||
rendered := m.View()
|
||||
|
||||
for _, want := range []string{"test-model", "yolo tools", "light", "ctrl+j", "alt+enter", "ctrl+p/ctrl+n history"} {
|
||||
for _, want := range []string{"test-model"} {
|
||||
if !strings.Contains(rendered, want) {
|
||||
t.Fatalf("rendered view missing %q:\n%s", want, rendered)
|
||||
}
|
||||
@@ -37,18 +37,16 @@ func TestModelRendersChatShell(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestInputTextDoesNotPaintBackground(t *testing.T) {
|
||||
func TestInputTextBackgroundMatchesInputBox(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
|
||||
for name, color := range map[string]any{
|
||||
"input box": m.styles.InputBox.GetBackground(),
|
||||
want := lightTheme().InputBg
|
||||
for name, color := range map[string]lipgloss.TerminalColor{
|
||||
"focused base": m.input.FocusedStyle.Base.GetBackground(),
|
||||
"focused cursor line": m.input.FocusedStyle.CursorLine.GetBackground(),
|
||||
"blurred base": m.input.BlurredStyle.Base.GetBackground(),
|
||||
"blurred cursor line": m.input.BlurredStyle.CursorLine.GetBackground(),
|
||||
} {
|
||||
if _, ok := color.(lipgloss.NoColor); !ok {
|
||||
t.Fatalf("%s should not paint input text background", name)
|
||||
}
|
||||
requireStyleColor(t, name, color, want)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -58,12 +56,12 @@ func TestInputBoxIsRoomier(t *testing.T) {
|
||||
m = next.(model)
|
||||
|
||||
if got := m.styles.InputBox.GetVerticalFrameSize(); got != 2 {
|
||||
t.Fatalf("input box vertical frame should stay compact, got %d", got)
|
||||
t.Fatalf("input box vertical frame should be 2 with vertical padding, got %d", got)
|
||||
}
|
||||
if got := m.styles.InputBox.GetHorizontalFrameSize(); got < 6 {
|
||||
t.Fatalf("input box horizontal frame should include roomier padding, got %d", got)
|
||||
if got := m.styles.InputBox.GetHorizontalFrameSize(); got != 3 {
|
||||
t.Fatalf("input box horizontal frame should be 3 (asymmetric padding), got %d", got)
|
||||
}
|
||||
if got, want := m.input.Width(), 100-m.styles.InputBox.GetHorizontalFrameSize(); got != want {
|
||||
if got, want := m.input.Width(), 100-m.styles.InputBox.GetHorizontalFrameSize()-m.inputModeWidth()-1; got != want {
|
||||
t.Fatalf("input content width = %d, want %d", got, want)
|
||||
}
|
||||
rendered := m.inputView()
|
||||
@@ -78,13 +76,8 @@ func TestModelMetaRendersBelowInput(t *testing.T) {
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
rendered := next.(model).View()
|
||||
|
||||
inputIndex := strings.Index(rendered, "Message agentu")
|
||||
modelIndex := strings.Index(rendered, "test-model")
|
||||
if inputIndex < 0 || modelIndex < 0 {
|
||||
t.Fatalf("rendered view missing input or model meta:\n%s", rendered)
|
||||
}
|
||||
if modelIndex < inputIndex {
|
||||
t.Fatalf("model meta should render below input:\n%s", rendered)
|
||||
if !strings.Contains(rendered, "test-model") {
|
||||
t.Fatalf("rendered view missing %q:\n%s", "test-model", rendered)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -93,7 +86,7 @@ func TestModelMetaRendersContextUsage(t *testing.T) {
|
||||
assistant.SetMessages([]llm.Message{{Role: llm.RoleUser, Content: strings.Repeat("x", 800)}})
|
||||
m := newModel(context.Background(), assistant, Options{ModelName: "test-model"})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 120, Height: 30})
|
||||
rendered := next.(model).modelMetaView()
|
||||
rendered := next.(model).footerView()
|
||||
|
||||
for _, want := range []string{"ctx ~", "/1k", "%"} {
|
||||
if !strings.Contains(rendered, want) {
|
||||
@@ -336,6 +329,23 @@ func TestInputHistoryNavigation(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpArrowRecallsPreviousInput(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
|
||||
m.rememberInput("first")
|
||||
m.rememberInput("second")
|
||||
|
||||
next, _ := m.Update(tea.KeyMsg{Type: tea.KeyUp})
|
||||
m = next.(model)
|
||||
if got := m.input.Value(); got != "second" {
|
||||
t.Fatalf("after first up input = %q, want %q", got, "second")
|
||||
}
|
||||
next, _ = m.Update(tea.KeyMsg{Type: tea.KeyUp})
|
||||
m = next.(model)
|
||||
if got := m.input.Value(); got != "first" {
|
||||
t.Fatalf("after second up input = %q, want %q", got, "first")
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolMessageRendersStructuredCard(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
|
||||
m.messages = append(m.messages, message{role: roleTool, content: `file_read {"path":"README.md"}`})
|
||||
@@ -348,6 +358,65 @@ func TestToolMessageRendersStructuredCard(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolMessageRendersOutputPreview(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
|
||||
m.messages = append(m.messages, message{role: roleTool, content: "file_read {\"path\":\"README.md\"}\n[output]\nfirst line\nsecond line"})
|
||||
|
||||
plain := stripANSI(m.renderMessages())
|
||||
for _, want := range []string{iconTool + " file_read", "README.md", "output", "first line", "second line"} {
|
||||
if !strings.Contains(plain, want) {
|
||||
t.Fatalf("tool card missing %q:\n%s", want, plain)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolLogOutputUpdatesExistingToolCard(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
|
||||
m.messages = append(m.messages, message{role: roleTool, content: `file_read {"path":"README.md"}`})
|
||||
|
||||
updated := m.mergeToolLog("file_read {\"path\":\"README.md\"}\n[output]\nhello")
|
||||
if !updated {
|
||||
t.Fatal("expected mergeToolLog to return true")
|
||||
}
|
||||
if len(m.messages) != 1 {
|
||||
t.Fatalf("messages count = %d, want 1", len(m.messages))
|
||||
}
|
||||
plain := stripANSI(m.renderMessages())
|
||||
if !strings.Contains(plain, "hello") {
|
||||
t.Fatalf("rendered output missing 'hello':\n%s", plain)
|
||||
}
|
||||
|
||||
status := toolStatusText("file_read {\"path\":\"README.md\"}\n[output]\nhello world")
|
||||
if strings.Contains(status, "hello") {
|
||||
t.Fatalf("status should not contain output: %q", status)
|
||||
}
|
||||
if !strings.Contains(status, "file_read") {
|
||||
t.Fatalf("status missing tool name: %q", status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolLogStreamingChunksAppendToSameCard(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
|
||||
m.messages = append(m.messages, message{role: roleTool, content: `shell_run echo hi`})
|
||||
|
||||
if !m.mergeToolLog("shell_run echo hi\n[output]\none\n") {
|
||||
t.Fatal("expected first mergeToolLog to return true")
|
||||
}
|
||||
if !m.mergeToolLog("shell_run echo hi\n[output]\ntwo\n") {
|
||||
t.Fatal("expected second mergeToolLog to return true")
|
||||
}
|
||||
if len(m.messages) != 1 {
|
||||
t.Fatalf("messages count = %d, want 1", len(m.messages))
|
||||
}
|
||||
plain := stripANSI(m.renderMessages())
|
||||
if !strings.Contains(plain, "one") || !strings.Contains(plain, "two") {
|
||||
t.Fatalf("rendered output missing streamed lines:\n%s", plain)
|
||||
}
|
||||
if strings.Contains(plain, "onetwo") {
|
||||
t.Fatalf("streamed lines were concatenated:\n%s", plain)
|
||||
}
|
||||
}
|
||||
|
||||
type tuiCompactProvider struct{}
|
||||
|
||||
func (tuiCompactProvider) ChatStream(ctx context.Context, req llm.ChatRequest, emit func(llm.StreamEvent) error) error {
|
||||
@@ -443,16 +512,60 @@ func runInDir(t *testing.T, dir string, name string, args ...string) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBangShellCommandRunsFromInput(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model", Yolo: true, WorkingDir: dir})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
m.input.SetValue("! printf hello")
|
||||
next, _ = m.submit()
|
||||
m = next.(model)
|
||||
if m.input.Value() != "" {
|
||||
t.Fatalf("input not cleared after bang command: %q", m.input.Value())
|
||||
}
|
||||
rendered := stripANSI(m.renderMessages())
|
||||
if !strings.Contains(rendered, "$ printf hello") {
|
||||
t.Fatalf("rendered output missing shell echo:\n%s", rendered)
|
||||
}
|
||||
if !strings.Contains(rendered, "hello") {
|
||||
t.Fatalf("rendered output missing 'hello':\n%s", rendered)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBangShellCommandRequiresYolo(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model", Yolo: false})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
m.input.SetValue("! printf hello")
|
||||
next, _ = m.submit()
|
||||
m = next.(model)
|
||||
rendered := stripANSI(m.renderMessages())
|
||||
if !strings.Contains(rendered, "shell command input requires --yolo") {
|
||||
t.Fatalf("rendered output missing yolo requirement:\n%s", rendered)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBangShellCommandRejectsEmptyCommand(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model", Yolo: true})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
m.input.SetValue("!")
|
||||
next, _ = m.submit()
|
||||
m = next.(model)
|
||||
rendered := stripANSI(m.renderMessages())
|
||||
if !strings.Contains(rendered, "usage: ! <command>") {
|
||||
t.Fatalf("rendered output missing usage hint:\n%s", rendered)
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelCommand(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model", Yolo: true})
|
||||
m.input.SetValue("/model")
|
||||
|
||||
next, _ := m.submit()
|
||||
rendered := next.(model).renderMessages()
|
||||
for _, want := range []string{"model: test-model", "mode: yolo tools", "theme: light"} {
|
||||
if !strings.Contains(rendered, want) {
|
||||
t.Fatalf("model command output missing %q:\n%s", want, rendered)
|
||||
}
|
||||
if !strings.Contains(rendered, "test-model") {
|
||||
t.Fatalf("model command output missing %q:\n%s", "test-model", rendered)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -471,11 +584,124 @@ func TestModelThinkingCommand(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelCommandPaletteShowsConfiguredModels(t *testing.T) {
|
||||
manager := &fakeModelManager{model: "model-a", models: []string{"model-a", "model-b"}}
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "model-a", ModelManager: manager})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
m.input.SetValue("/model ")
|
||||
m.afterInputChanged()
|
||||
plain := stripANSI(m.View())
|
||||
for _, want := range []string{iconCommand + " Model candidates", "model-a", "current model", "model-b", "switch model"} {
|
||||
if !strings.Contains(plain, want) {
|
||||
t.Fatalf("rendered palette missing %q:\n%s", want, plain)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelCommandPaletteSelectsConfiguredModel(t *testing.T) {
|
||||
manager := &fakeModelManager{model: "model-a", models: []string{"model-a", "model-b"}}
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "model-a", ModelManager: manager})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
m.input.SetValue("/model ")
|
||||
m.afterInputChanged()
|
||||
next, _ = m.Update(tea.KeyMsg{Type: tea.KeyDown})
|
||||
m = next.(model)
|
||||
next, _ = m.Update(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
m = next.(model)
|
||||
if manager.model != "model-b" {
|
||||
t.Fatalf("model = %q", manager.model)
|
||||
}
|
||||
if got := m.input.Value(); got != "" {
|
||||
t.Fatalf("input = %q", got)
|
||||
}
|
||||
rendered := stripANSI(m.renderMessages())
|
||||
if !strings.Contains(rendered, "model: model-b") {
|
||||
t.Fatalf("rendered output missing model switch:\n%s", rendered)
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelCommandPaletteFiltersBareModelPartial(t *testing.T) {
|
||||
manager := &fakeModelManager{model: "claude-sonnet", models: []string{"claude-sonnet", "deepseek-v4-flash"}}
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "claude-sonnet", ModelManager: manager})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
m.input.SetValue("/model deep")
|
||||
m.afterInputChanged()
|
||||
palette := stripANSI(m.commandPaletteView())
|
||||
if !strings.Contains(palette, "deepseek-v4-flash") {
|
||||
t.Fatalf("expected deepseek-v4-flash in palette:\n%s", palette)
|
||||
}
|
||||
if strings.Contains(palette, "claude-sonnet") {
|
||||
t.Fatalf("unexpected claude-sonnet in palette:\n%s", palette)
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelCommandPaletteFiltersExplicitModelSubcommand(t *testing.T) {
|
||||
manager := &fakeModelManager{model: "claude-sonnet", models: []string{"claude-sonnet", "deepseek-v4-flash"}}
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "claude-sonnet", ModelManager: manager})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
m.input.SetValue("/model model deep")
|
||||
m.afterInputChanged()
|
||||
palette := stripANSI(m.commandPaletteView())
|
||||
if !strings.Contains(palette, "deepseek-v4-flash") {
|
||||
t.Fatalf("expected deepseek-v4-flash in palette:\n%s", palette)
|
||||
}
|
||||
if strings.Contains(palette, "claude-sonnet") {
|
||||
t.Fatalf("unexpected claude-sonnet in palette:\n%s", palette)
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelCommandPaletteTabFillsConfiguredModel(t *testing.T) {
|
||||
manager := &fakeModelManager{model: "claude-sonnet", models: []string{"claude-sonnet", "deepseek-v4-flash"}}
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "claude-sonnet", ModelManager: manager})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
m.input.SetValue("/model deep")
|
||||
m.afterInputChanged()
|
||||
next, _ = m.Update(tea.KeyMsg{Type: tea.KeyTab})
|
||||
m = next.(model)
|
||||
if got := m.input.Value(); got != "/model model deepseek-v4-flash" {
|
||||
t.Fatalf("input = %q", got)
|
||||
}
|
||||
if manager.model != "claude-sonnet" {
|
||||
t.Fatalf("model should not change on tab: %q", manager.model)
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelCommandPaletteIgnoresProviderAndThinkingSubcommands(t *testing.T) {
|
||||
manager := &fakeModelManager{model: "model-a", models: []string{"model-a", "model-b"}}
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "model-a", ModelManager: manager})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
|
||||
m.input.SetValue("/model provider ")
|
||||
m.afterInputChanged()
|
||||
if sugs := m.commandSuggestions(); len(sugs) != 0 {
|
||||
t.Fatalf("provider suggestions should be empty, got %d", len(sugs))
|
||||
}
|
||||
if view := m.commandPaletteView(); view != "" {
|
||||
t.Fatalf("provider palette should be empty:\n%s", view)
|
||||
}
|
||||
|
||||
m.input.SetValue("/model thinking ")
|
||||
m.afterInputChanged()
|
||||
if sugs := m.commandSuggestions(); len(sugs) != 0 {
|
||||
t.Fatalf("thinking suggestions should be empty, got %d", len(sugs))
|
||||
}
|
||||
if view := m.commandPaletteView(); view != "" {
|
||||
t.Fatalf("thinking palette should be empty:\n%s", view)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeModelManager struct {
|
||||
provider string
|
||||
model string
|
||||
thinking string
|
||||
sessionName string
|
||||
models []string
|
||||
saved bool
|
||||
}
|
||||
|
||||
@@ -546,6 +772,16 @@ func (f *fakeModelManager) NewSession() (string, error) {
|
||||
return "New session.", nil
|
||||
}
|
||||
|
||||
func (f *fakeModelManager) AvailableModels() []string {
|
||||
if len(f.models) > 0 {
|
||||
return append([]string(nil), f.models...)
|
||||
}
|
||||
if f.model != "" {
|
||||
return []string{f.model}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestParseThemeMode(t *testing.T) {
|
||||
for input, want := range map[string]ThemeMode{
|
||||
"": ThemeLight,
|
||||
@@ -567,6 +803,130 @@ func TestParseThemeMode(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func requireStyleColor(t *testing.T, name string, got lipgloss.TerminalColor, want string) {
|
||||
t.Helper()
|
||||
color, ok := got.(lipgloss.Color)
|
||||
if !ok {
|
||||
t.Fatalf("%s = %T, want lipgloss.Color", name, got)
|
||||
}
|
||||
if string(color) != want {
|
||||
t.Fatalf("%s = %q, want %q", name, string(color), want)
|
||||
}
|
||||
}
|
||||
|
||||
func requireNoStyleColor(t *testing.T, name string, got lipgloss.TerminalColor) {
|
||||
t.Helper()
|
||||
if _, ok := got.(lipgloss.NoColor); !ok {
|
||||
t.Fatalf("%s should not paint a color, got %T", name, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDarkThemePaintsAppBackground(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model", ThemeMode: ThemeDark})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
rendered := stripANSI(m.View())
|
||||
_ = rendered
|
||||
if m.theme.Mode != ThemeDark {
|
||||
t.Fatalf("theme.Mode = %q, want %q", m.theme.Mode, ThemeDark)
|
||||
}
|
||||
want := darkTheme()
|
||||
requireStyleColor(t, "App.Background", m.styles.App.GetBackground(), want.Background)
|
||||
requireStyleColor(t, "App.Foreground", m.styles.App.GetForeground(), want.Text)
|
||||
}
|
||||
|
||||
func TestDarkThemeKeepsComponentBackgroundsLightweight(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model", ThemeMode: ThemeDark})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
requireNoStyleColor(t, "Viewport.Background", m.styles.Viewport.GetBackground())
|
||||
requireStyleColor(t, "InputBox.Background", m.styles.InputBox.GetBackground(), darkTheme().InputBg)
|
||||
requireNoStyleColor(t, "Activity.Background", m.styles.Activity.GetBackground())
|
||||
m.running = true
|
||||
m.status = "thinking..."
|
||||
m.resize()
|
||||
if backgroundPattern.MatchString(m.activityView()) {
|
||||
t.Fatalf("activity view should not paint a background:\n%q", m.activityView())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDarkThemeStylesUseDarkPalette(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model", ThemeMode: ThemeDark})
|
||||
want := darkTheme()
|
||||
requireStyleColor(t, "Header.Foreground", m.styles.Header.GetForeground(), want.Title)
|
||||
requireStyleColor(t, "Muted.Foreground", m.styles.Muted.GetForeground(), want.Muted)
|
||||
requireStyleColor(t, "Pill.Background", m.styles.Pill.GetBackground(), want.SurfaceAlt)
|
||||
requireStyleColor(t, "input.FocusedStyle.Base.Foreground", m.input.FocusedStyle.Base.GetForeground(), want.Text)
|
||||
requireStyleColor(t, "input.FocusedStyle.Placeholder.Foreground", m.input.FocusedStyle.Placeholder.GetForeground(), want.Muted)
|
||||
requireStyleColor(t, "InputModeYolo.Foreground", m.styles.InputModeYolo.GetForeground(), want.Error)
|
||||
requireStyleColor(t, "InputModeCommand.Foreground", m.styles.InputModeCommand.GetForeground(), want.User)
|
||||
}
|
||||
|
||||
func TestInputModeSymbolReflectsYoloAndBangCommand(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model", Yolo: true})
|
||||
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
|
||||
m = next.(model)
|
||||
|
||||
rendered := stripANSI(m.inputView())
|
||||
if !strings.Contains(rendered, "Y") {
|
||||
t.Fatalf("YOLO input view should contain Y symbol:\n%s", rendered)
|
||||
}
|
||||
|
||||
m.input.SetValue("! pwd")
|
||||
m.afterInputChanged()
|
||||
rendered = stripANSI(m.inputView())
|
||||
if !strings.Contains(rendered, iconCommand) {
|
||||
t.Fatalf("bang input view should contain command icon:\n%s", rendered)
|
||||
}
|
||||
if strings.Contains(rendered, "Y") {
|
||||
t.Fatalf("bang input view should not contain Y symbol:\n%s", rendered)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInputModeYoloUsesWarningColor(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model", Yolo: true})
|
||||
requireStyleColor(t, "InputModeYolo.Foreground", m.styles.InputModeYolo.GetForeground(), lightTheme().Error)
|
||||
}
|
||||
|
||||
func TestThemeSlashCommandSwitchesAtRuntime(t *testing.T) {
|
||||
m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
|
||||
if m.theme.Mode != ThemeLight {
|
||||
t.Fatalf("initial theme = %q, want light", m.theme.Mode)
|
||||
}
|
||||
requireNoStyleColor(t, "initial App.Background", m.styles.App.GetBackground())
|
||||
|
||||
// Switch to dark
|
||||
m.input.SetValue("/theme dark")
|
||||
next, _ := m.submit()
|
||||
m = next.(model)
|
||||
if m.theme.Mode != ThemeDark {
|
||||
t.Fatalf("after /theme dark: Mode = %q, want dark", m.theme.Mode)
|
||||
}
|
||||
requireStyleColor(t, "dark App.Background", m.styles.App.GetBackground(), darkTheme().Background)
|
||||
last := m.messages[len(m.messages)-1]
|
||||
if !strings.Contains(last.content, "theme: dark") {
|
||||
t.Fatalf("expected confirmation message, got %q", last.content)
|
||||
}
|
||||
|
||||
// Switch back to light
|
||||
m.input.SetValue("/theme light")
|
||||
next, _ = m.submit()
|
||||
m = next.(model)
|
||||
if m.theme.Mode != ThemeLight {
|
||||
t.Fatalf("after /theme light: Mode = %q, want light", m.theme.Mode)
|
||||
}
|
||||
requireNoStyleColor(t, "restored App.Background", m.styles.App.GetBackground())
|
||||
|
||||
// No arg shows usage
|
||||
m.input.SetValue("/theme")
|
||||
next, _ = m.submit()
|
||||
m = next.(model)
|
||||
last = m.messages[len(m.messages)-1]
|
||||
if !strings.Contains(last.content, "usage: /theme") {
|
||||
t.Fatalf("expected usage hint, got %q", last.content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanToolLog(t *testing.T) {
|
||||
got := cleanToolLog("\n[tool] shell_run {\"command\":\"pwd\"}\n")
|
||||
want := `shell_run {"command":"pwd"}`
|
||||
|
||||
+40
-15
@@ -43,6 +43,7 @@ type theme struct {
|
||||
Background string
|
||||
Surface string
|
||||
SurfaceAlt string
|
||||
InputBg string
|
||||
Border string
|
||||
Text string
|
||||
Muted string
|
||||
@@ -60,11 +61,13 @@ type styles struct {
|
||||
Muted lipgloss.Style
|
||||
Activity lipgloss.Style
|
||||
Header lipgloss.Style
|
||||
ModelMeta lipgloss.Style
|
||||
Footer lipgloss.Style
|
||||
Viewport lipgloss.Style
|
||||
InputBox lipgloss.Style
|
||||
CommandHint lipgloss.Style
|
||||
InputModeReadonly lipgloss.Style
|
||||
InputModeYolo lipgloss.Style
|
||||
InputModeCommand lipgloss.Style
|
||||
StatusReady lipgloss.Style
|
||||
StatusWarn lipgloss.Style
|
||||
StatusError lipgloss.Style
|
||||
@@ -98,8 +101,8 @@ func lightTheme() theme {
|
||||
return theme{
|
||||
Mode: ThemeLight,
|
||||
Background: "#F8FAFC",
|
||||
Surface: "#FFFFFF",
|
||||
SurfaceAlt: "#EEF2F7",
|
||||
InputBg: "#E4E4E4",
|
||||
Border: "#CBD5E1",
|
||||
Text: "#111827",
|
||||
Muted: "#64748B",
|
||||
@@ -116,8 +119,8 @@ func darkTheme() theme {
|
||||
return theme{
|
||||
Mode: ThemeDark,
|
||||
Background: "#111827",
|
||||
Surface: "#1F2937",
|
||||
SurfaceAlt: "#0F172A",
|
||||
InputBg: "#383838",
|
||||
Border: "#475569",
|
||||
Text: "#E5E7EB",
|
||||
Muted: "#94A3B8",
|
||||
@@ -130,10 +133,17 @@ func darkTheme() theme {
|
||||
}
|
||||
}
|
||||
|
||||
func appStyle(t theme) lipgloss.Style {
|
||||
style := lipgloss.NewStyle().Foreground(lipgloss.Color(t.Text))
|
||||
if t.Mode == ThemeDark {
|
||||
style = style.Background(lipgloss.Color(t.Background))
|
||||
}
|
||||
return style
|
||||
}
|
||||
|
||||
func (t theme) styles() styles {
|
||||
return styles{
|
||||
App: lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(t.Text)),
|
||||
App: appStyle(t),
|
||||
|
||||
Muted: lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(t.Muted)),
|
||||
@@ -146,10 +156,6 @@ func (t theme) styles() styles {
|
||||
Bold(true).
|
||||
Foreground(lipgloss.Color(t.Title)),
|
||||
|
||||
ModelMeta: lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(t.Muted)).
|
||||
Padding(0, 1),
|
||||
|
||||
Footer: lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(t.Muted)).
|
||||
Padding(0, 1),
|
||||
@@ -159,14 +165,31 @@ func (t theme) styles() styles {
|
||||
|
||||
InputBox: lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(t.Text)).
|
||||
Border(lipgloss.ThickBorder()).
|
||||
BorderForeground(lipgloss.Color(t.Border)).
|
||||
Padding(0, 2),
|
||||
Background(lipgloss.Color(t.InputBg)).
|
||||
Padding(1, 2, 1, 1),
|
||||
|
||||
CommandHint: lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(t.Muted)).
|
||||
Padding(0, 1),
|
||||
|
||||
InputModeReadonly: lipgloss.NewStyle().
|
||||
Bold(true).
|
||||
Foreground(lipgloss.Color(t.Muted)).
|
||||
Background(lipgloss.Color(t.InputBg)).
|
||||
Padding(0, 1),
|
||||
|
||||
InputModeYolo: lipgloss.NewStyle().
|
||||
Bold(true).
|
||||
Foreground(lipgloss.Color(t.Error)).
|
||||
Background(lipgloss.Color(t.InputBg)).
|
||||
Padding(0, 1),
|
||||
|
||||
InputModeCommand: lipgloss.NewStyle().
|
||||
Bold(true).
|
||||
Foreground(lipgloss.Color(t.User)).
|
||||
Background(lipgloss.Color(t.InputBg)).
|
||||
Padding(0, 1),
|
||||
|
||||
StatusReady: lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(t.Assistant)),
|
||||
|
||||
@@ -251,8 +274,10 @@ func bubbleMessageStyle(textColor, backgroundColor string) lipgloss.Style {
|
||||
func applyTextareaTheme(input *textarea.Model, t theme) {
|
||||
focused := textarea.Style{
|
||||
Base: lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(t.Text)),
|
||||
CursorLine: lipgloss.NewStyle(),
|
||||
Foreground(lipgloss.Color(t.Text)).
|
||||
Background(lipgloss.Color(t.InputBg)),
|
||||
CursorLine: lipgloss.NewStyle().
|
||||
Background(lipgloss.Color(t.InputBg)),
|
||||
Placeholder: lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(t.Muted)),
|
||||
Prompt: lipgloss.NewStyle().
|
||||
@@ -263,7 +288,7 @@ func applyTextareaTheme(input *textarea.Model, t theme) {
|
||||
Foreground(lipgloss.Color(t.Surface)),
|
||||
}
|
||||
blurred := focused
|
||||
blurred.CursorLine = lipgloss.NewStyle()
|
||||
blurred.CursorLine = lipgloss.NewStyle().Background(lipgloss.Color(t.InputBg))
|
||||
|
||||
input.FocusedStyle = focused
|
||||
input.BlurredStyle = blurred
|
||||
|
||||
@@ -1,228 +1,15 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"agentu/internal/agent"
|
||||
"agentu/internal/session"
|
||||
"agentu/internal/tools"
|
||||
"agentu/internal/tui"
|
||||
"agentu/pkg/config"
|
||||
"agentu/pkg/llm"
|
||||
"agentu/internal/app"
|
||||
)
|
||||
|
||||
const defaultConfigContent = `# agentu configuration
|
||||
# Replace the placeholder values below with your actual provider settings.
|
||||
# See agentu.example.yaml for the full reference.
|
||||
providers:
|
||||
default:
|
||||
base_url: https://api.openai.com
|
||||
api_key: YOUR_API_KEY_HERE
|
||||
models:
|
||||
- id: gpt-4o
|
||||
context: 128000
|
||||
- id: gpt-4o-mini
|
||||
context: 128000
|
||||
`
|
||||
|
||||
func main() {
|
||||
if err := run(); err != nil {
|
||||
if err := app.Run(); err != nil {
|
||||
fmt.Fprintln(os.Stderr, "agentu:", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
func run() error {
|
||||
configPath := flag.String("config", config.DefaultPath, "path to YAML config")
|
||||
yolo := flag.Bool("yolo", false, "also enable automatic file writes and shell execution")
|
||||
plain := flag.Bool("plain", false, "use the simple line-based REPL instead of the TUI")
|
||||
themeName := flag.String("theme", "light", "TUI theme: light or dark")
|
||||
resumeID := flag.String("resume", "", "resume a specific session by ID")
|
||||
flag.Parse()
|
||||
|
||||
themeMode, err := tui.ParseThemeMode(*themeName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cfg, err := config.Load(*configPath)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
if bootstrapErr := bootstrapConfig(*configPath); bootstrapErr != nil {
|
||||
return bootstrapErr
|
||||
}
|
||||
return fmt.Errorf("created default config at %s; edit it to set your API key and provider, then restart agentu", *configPath)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
providerName := cfg.DefaultProviderName()
|
||||
providerConfig := cfg.Providers[providerName]
|
||||
provider := llm.NewOpenAICompatibleClient(providerConfig.BaseURL, providerConfig.APIKey, http.DefaultClient)
|
||||
|
||||
var registry *tools.Registry
|
||||
if *yolo {
|
||||
registry = tools.Builtins(cfg.Agent.WorkingDir)
|
||||
} else {
|
||||
registry = tools.ReadOnlyBuiltins(cfg.Agent.WorkingDir)
|
||||
}
|
||||
registry.Register(tools.NewBuiltinSearchTool(http.DefaultClient))
|
||||
registry.Register(tools.NewBuiltinFetchTool(http.DefaultClient))
|
||||
systemPrompt := strings.TrimSpace(cfg.Agent.SystemPrompt + "\n\n" + tools.AgentInstructions(cfg.Agent.WorkingDir, *yolo))
|
||||
maxContextTokens := cfg.Agent.MaxContextTokens
|
||||
if maxContextTokens == 0 {
|
||||
maxContextTokens = providerConfig.ContextTokens
|
||||
}
|
||||
|
||||
assistant := agent.New(agent.Options{
|
||||
Provider: provider,
|
||||
ProviderName: providerName,
|
||||
Model: providerConfig.Model,
|
||||
Thinking: providerConfig.Thinking,
|
||||
ThinkingParam: providerConfig.ThinkingParam,
|
||||
ThinkingEnabled: providerConfig.ThinkingConfigured,
|
||||
SystemPrompt: systemPrompt,
|
||||
ToolRegistry: registry,
|
||||
ToolTimeout: cfg.Agent.ToolTimeout,
|
||||
MaxContextTokens: maxContextTokens,
|
||||
})
|
||||
|
||||
// Initialize session store
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve home dir: %w", err)
|
||||
}
|
||||
sessionsDir := filepath.Join(home, ".agentu", "sessions")
|
||||
store, err := session.NewStore(sessionsDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("init session store: %w", err)
|
||||
}
|
||||
|
||||
modelManager, err := session.NewManager(cfg, assistant, http.DefaultClient, store)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := initializeSession(modelManager, *resumeID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt)
|
||||
defer stop()
|
||||
|
||||
if !*plain {
|
||||
err := tui.Run(ctx, assistant, tui.Options{
|
||||
ModelName: providerConfig.Model,
|
||||
Yolo: *yolo,
|
||||
ThemeMode: themeMode,
|
||||
ModelManager: modelManager,
|
||||
WorkingDir: cfg.Agent.WorkingDir,
|
||||
})
|
||||
printExitSession(modelManager)
|
||||
return err
|
||||
}
|
||||
|
||||
return repl(ctx, assistant, modelManager)
|
||||
}
|
||||
|
||||
func initializeSession(modelManager *session.Manager, resumeID string) error {
|
||||
if resumeID != "" {
|
||||
if err := modelManager.Resume(resumeID); err != nil {
|
||||
return fmt.Errorf("resume session: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if _, err := modelManager.NewSession(); err != nil {
|
||||
return fmt.Errorf("start new session: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func bootstrapConfig(configPath string) error {
|
||||
resolved, err := config.ResolvePath(configPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve config path: %w", err)
|
||||
}
|
||||
dir := filepath.Dir(resolved)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return fmt.Errorf("create config directory: %w", err)
|
||||
}
|
||||
if err := os.WriteFile(resolved, []byte(defaultConfigContent), 0o644); err != nil {
|
||||
return fmt.Errorf("write default config: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func repl(ctx context.Context, assistant *agent.Agent, modelManager *session.Manager) error {
|
||||
scanner := bufio.NewScanner(os.Stdin)
|
||||
scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024)
|
||||
|
||||
fmt.Println("agentu")
|
||||
fmt.Println("Type /exit to quit, /clear to reset context, /compact to compress context.")
|
||||
for {
|
||||
fmt.Print("> ")
|
||||
if !scanner.Scan() {
|
||||
if err := scanner.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Println()
|
||||
printExitSession(modelManager)
|
||||
return nil
|
||||
}
|
||||
|
||||
input := strings.TrimSpace(scanner.Text())
|
||||
switch input {
|
||||
case "":
|
||||
continue
|
||||
case "/exit", "/quit":
|
||||
printExitSession(modelManager)
|
||||
return nil
|
||||
case "/clear":
|
||||
assistant.Clear()
|
||||
fmt.Println("context cleared")
|
||||
continue
|
||||
case "/compact":
|
||||
if err := assistant.Compact(ctx); err != nil {
|
||||
fmt.Fprintln(os.Stderr, "compact error:", err)
|
||||
} else {
|
||||
fmt.Println("context compacted")
|
||||
if modelManager != nil {
|
||||
_ = modelManager.Save()
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if err := assistant.RunTurn(ctx, input, os.Stdout, os.Stderr); err != nil {
|
||||
fmt.Fprintln(os.Stderr, "\nerror:", err)
|
||||
}
|
||||
_ = modelManager.Save()
|
||||
fmt.Println()
|
||||
}
|
||||
}
|
||||
|
||||
func printExitSession(m *session.Manager) {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
_ = m.Save()
|
||||
id := m.CurrentSessionID()
|
||||
name := m.CurrentSessionName()
|
||||
if id != "" {
|
||||
if name != "" && name != id {
|
||||
fmt.Fprintf(os.Stderr, "session: %s (%s)\n", id, name)
|
||||
} else {
|
||||
fmt.Fprintf(os.Stderr, "session: %s\n", id)
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "resume: agentu --resume %s\n", id)
|
||||
}
|
||||
}
|
||||
|
||||
+45
-1
@@ -14,6 +14,11 @@ import (
|
||||
|
||||
const DefaultPath = "~/.agentu/config.yaml"
|
||||
|
||||
const (
|
||||
APITypeChat = "chat"
|
||||
APITypeResponses = "responses"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
Providers map[string]ProviderConfig `yaml:"providers"`
|
||||
Agent AgentConfig `yaml:"agent"`
|
||||
@@ -21,6 +26,7 @@ type Config struct {
|
||||
|
||||
type ProviderConfig struct {
|
||||
Name string
|
||||
APIType string
|
||||
BaseURL string
|
||||
Model string
|
||||
Models []string
|
||||
@@ -56,6 +62,7 @@ type rawConfig struct {
|
||||
}
|
||||
|
||||
type rawProviderConfig struct {
|
||||
APIType string `yaml:"api_type"`
|
||||
BaseURL string `yaml:"base_url"`
|
||||
Models []rawModelConfig `yaml:"models"`
|
||||
APIKey string `yaml:"api_key"`
|
||||
@@ -153,6 +160,9 @@ func (c *Config) Validate() error {
|
||||
for name, provider := range c.Providers {
|
||||
var missing []string
|
||||
prefix := "providers." + name
|
||||
if !IsAPIType(provider.APIType) {
|
||||
return fmt.Errorf("%s.api_type must be one of: %s", prefix, strings.Join(APITypes(), ", "))
|
||||
}
|
||||
if strings.TrimSpace(provider.BaseURL) == "" {
|
||||
missing = append(missing, prefix+".base_url")
|
||||
}
|
||||
@@ -220,6 +230,7 @@ func normalizeProvider(name string, raw rawProviderConfig) ProviderConfig {
|
||||
models, modelConfigs := normalizeModelConfigs(raw.Models)
|
||||
provider := ProviderConfig{
|
||||
Name: name,
|
||||
APIType: normalizeAPIType(raw.APIType),
|
||||
BaseURL: raw.BaseURL,
|
||||
Models: models,
|
||||
ModelConfigs: modelConfigs,
|
||||
@@ -230,11 +241,44 @@ func normalizeProvider(name string, raw rawProviderConfig) ProviderConfig {
|
||||
}
|
||||
applyModelDefaults(&provider)
|
||||
if provider.ThinkingConfigured && provider.ThinkingParam == "" {
|
||||
provider.ThinkingParam = "thinking"
|
||||
provider.ThinkingParam = DefaultThinkingParam(provider.APIType)
|
||||
}
|
||||
return provider
|
||||
}
|
||||
|
||||
func normalizeAPIType(value string) string {
|
||||
value = strings.ToLower(strings.TrimSpace(value))
|
||||
if value == "" {
|
||||
return APITypeChat
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func APITypes() []string {
|
||||
return []string{APITypeChat, APITypeResponses}
|
||||
}
|
||||
|
||||
func IsAPIType(value string) bool {
|
||||
value = strings.ToLower(strings.TrimSpace(value))
|
||||
for _, apiType := range APITypes() {
|
||||
if value == apiType {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// DefaultThinkingParam returns the default request field used for the
|
||||
// thinking level. Chat-completions providers historically use a top-level
|
||||
// "thinking" field; the OpenAI Responses API uses "reasoning" with an
|
||||
// {"effort": ...} value.
|
||||
func DefaultThinkingParam(apiType string) string {
|
||||
if apiType == APITypeResponses {
|
||||
return "reasoning"
|
||||
}
|
||||
return "thinking"
|
||||
}
|
||||
|
||||
func normalizeModelConfigs(rawModels []rawModelConfig) ([]string, map[string]ModelConfig) {
|
||||
models := make([]string, 0, len(rawModels))
|
||||
configs := make(map[string]ModelConfig, len(rawModels))
|
||||
|
||||
@@ -370,3 +370,110 @@ providers:
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDefaultsAPITypeToChat(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv("AGENTU_TEST_KEY", "sk-test")
|
||||
path := filepath.Join(dir, "agentu.yaml")
|
||||
if err := os.WriteFile(path, []byte(`
|
||||
providers:
|
||||
test:
|
||||
base_url: https://example.com
|
||||
api_key: ${AGENTU_TEST_KEY}
|
||||
models:
|
||||
- id: test-model
|
||||
`), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
cfg, err := Load(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := cfg.Providers["test"].APIType; got != APITypeChat {
|
||||
t.Fatalf("api type = %q, want %q", got, APITypeChat)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadParsesResponsesAPIType(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv("AGENTU_TEST_KEY", "sk-test")
|
||||
path := filepath.Join(dir, "agentu.yaml")
|
||||
if err := os.WriteFile(path, []byte(`
|
||||
providers:
|
||||
openai:
|
||||
api_type: RESPONSES
|
||||
base_url: https://api.openai.com
|
||||
api_key: ${AGENTU_TEST_KEY}
|
||||
models:
|
||||
- id: gpt-test
|
||||
`), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
cfg, err := Load(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := cfg.Providers["openai"].APIType; got != APITypeResponses {
|
||||
t.Fatalf("api type = %q, want %q", got, APITypeResponses)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsInvalidAPIType(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv("AGENTU_TEST_KEY", "sk-test")
|
||||
path := filepath.Join(dir, "agentu.yaml")
|
||||
if err := os.WriteFile(path, []byte(`
|
||||
providers:
|
||||
test:
|
||||
api_type: completions
|
||||
base_url: https://example.com
|
||||
api_key: ${AGENTU_TEST_KEY}
|
||||
models:
|
||||
- id: test-model
|
||||
`), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err := Load(path)
|
||||
if err == nil || !strings.Contains(err.Error(), "api_type") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDefaultsThinkingParamByAPIType(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv("AGENTU_TEST_KEY", "sk-test")
|
||||
path := filepath.Join(dir, "agentu.yaml")
|
||||
if err := os.WriteFile(path, []byte(`
|
||||
providers:
|
||||
chat:
|
||||
api_type: chat
|
||||
base_url: https://chat.example.com
|
||||
api_key: ${AGENTU_TEST_KEY}
|
||||
models:
|
||||
- id: chat-model
|
||||
thinking: none
|
||||
responses:
|
||||
api_type: responses
|
||||
base_url: https://responses.example.com
|
||||
api_key: ${AGENTU_TEST_KEY}
|
||||
models:
|
||||
- id: responses-model
|
||||
thinking: high
|
||||
`), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
cfg, err := Load(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := cfg.Providers["chat"].ThinkingParam; got != "thinking" {
|
||||
t.Fatalf("chat thinking param = %q, want %q", got, "thinking")
|
||||
}
|
||||
if got := cfg.Providers["responses"].ThinkingParam; got != "reasoning" {
|
||||
t.Fatalf("responses thinking param = %q, want %q", got, "reasoning")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,481 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const responsesEndpoint = "/v1/responses"
|
||||
|
||||
// OpenAIResponsesClient implements llm.Provider against the OpenAI Responses
|
||||
// API (/v1/responses). It maps the agent's chat-style conversation history to
|
||||
// Responses input items (messages, function_call, function_call_output) and
|
||||
// converts Responses streaming events back into the shared StreamEvent shape.
|
||||
type OpenAIResponsesClient struct {
|
||||
baseURL string
|
||||
apiKey string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewOpenAIResponsesClient(baseURL, apiKey string, httpClient *http.Client) *OpenAIResponsesClient {
|
||||
if httpClient == nil {
|
||||
httpClient = http.DefaultClient
|
||||
}
|
||||
return &OpenAIResponsesClient{
|
||||
baseURL: strings.TrimRight(baseURL, "/"),
|
||||
apiKey: apiKey,
|
||||
httpClient: httpClient,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *OpenAIResponsesClient) ChatStream(ctx context.Context, req ChatRequest, emit func(StreamEvent) error) error {
|
||||
body, err := json.Marshal(buildResponsesRequest(req))
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal responses request: %w", err)
|
||||
}
|
||||
state := &responsesStreamState{}
|
||||
return postResponsesSSE(ctx, c.baseURL, c.apiKey, c.httpClient, body, func(payload string) error {
|
||||
return state.handle(payload, emit)
|
||||
})
|
||||
}
|
||||
|
||||
func postResponsesSSE(ctx context.Context, baseURL, apiKey string, httpClient *http.Client, body []byte, handle func(payload string) error) error {
|
||||
var lastErr error
|
||||
for attempt := 1; attempt <= maxChatAttempts; attempt++ {
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, baseURL+responsesEndpoint, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return fmt.Errorf("create responses request: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
httpReq.Header.Set("Accept", "text/event-stream")
|
||||
if apiKey != "" {
|
||||
httpReq.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
}
|
||||
|
||||
resp, err := httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
lastErr = fmt.Errorf("send responses request: %w", err)
|
||||
if !shouldRetryRequest(ctx, attempt, 0) {
|
||||
return lastErr
|
||||
}
|
||||
if waitErr := waitBeforeRetry(ctx); waitErr != nil {
|
||||
return waitErr
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
limit := io.LimitReader(resp.Body, 4096)
|
||||
data, _ := io.ReadAll(limit)
|
||||
_ = resp.Body.Close()
|
||||
lastErr = fmt.Errorf("responses request failed: status=%d body=%s", resp.StatusCode, strings.TrimSpace(string(data)))
|
||||
if !shouldRetryRequest(ctx, attempt, resp.StatusCode) {
|
||||
return lastErr
|
||||
}
|
||||
if waitErr := waitBeforeRetry(ctx); waitErr != nil {
|
||||
return waitErr
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
defer resp.Body.Close()
|
||||
return readResponsesSSE(resp.Body, handle)
|
||||
}
|
||||
return lastErr
|
||||
}
|
||||
|
||||
func readResponsesSSE(r io.Reader, handle func(payload string) error) error {
|
||||
reader := bufio.NewReader(r)
|
||||
for {
|
||||
line, err := reader.ReadString('\n')
|
||||
if err != nil && !errors.Is(err, io.EOF) {
|
||||
return fmt.Errorf("read responses stream: %w", err)
|
||||
}
|
||||
|
||||
line = strings.TrimSpace(line)
|
||||
if strings.HasPrefix(line, "data:") {
|
||||
payload := strings.TrimSpace(strings.TrimPrefix(line, "data:"))
|
||||
if payload == "[DONE]" {
|
||||
return nil
|
||||
}
|
||||
if handleErr := handle(payload); handleErr != nil {
|
||||
return handleErr
|
||||
}
|
||||
}
|
||||
|
||||
if errors.Is(err, io.EOF) {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// buildResponsesRequest converts a ChatRequest into the /v1/responses body.
|
||||
func buildResponsesRequest(req ChatRequest) map[string]any {
|
||||
payload := map[string]any{
|
||||
"model": req.Model,
|
||||
"input": buildResponsesInput(req.Messages),
|
||||
"stream": true,
|
||||
}
|
||||
if instructions := responsesInstructions(req.Messages); instructions != "" {
|
||||
payload["instructions"] = instructions
|
||||
}
|
||||
if len(req.Tools) > 0 {
|
||||
payload["tools"] = buildResponsesTools(req.Tools)
|
||||
}
|
||||
if req.ToolChoice != "" {
|
||||
payload["tool_choice"] = req.ToolChoice
|
||||
}
|
||||
for key, value := range req.Extra {
|
||||
if _, exists := payload[key]; exists {
|
||||
continue
|
||||
}
|
||||
if key == "reasoning" {
|
||||
if level, ok := value.(string); ok {
|
||||
payload[key] = map[string]any{"effort": mapReasoningEffort(level)}
|
||||
continue
|
||||
}
|
||||
}
|
||||
payload[key] = value
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
// buildResponsesInput maps chat-style messages to Responses input items.
|
||||
// Assistant tool calls become function_call items and tool results become
|
||||
// function_call_output items, preserving their position in the conversation.
|
||||
func buildResponsesInput(messages []Message) []any {
|
||||
items := make([]any, 0, len(messages)+4)
|
||||
for _, msg := range messages {
|
||||
switch msg.Role {
|
||||
case RoleSystem:
|
||||
// System instructions are sent as the top-level "instructions"
|
||||
// field and must not be duplicated inside the input array.
|
||||
continue
|
||||
case RoleUser:
|
||||
items = append(items, responsesMessageInput{Role: "user", Content: msg.Content})
|
||||
case RoleAssistant:
|
||||
if msg.Content != "" {
|
||||
items = append(items, responsesMessageInput{Role: "assistant", Content: msg.Content})
|
||||
}
|
||||
for _, call := range msg.ToolCalls {
|
||||
items = append(items, responsesFunctionCallInput{
|
||||
Type: "function_call",
|
||||
CallID: call.ID,
|
||||
Name: call.Function.Name,
|
||||
Arguments: call.Function.Arguments,
|
||||
})
|
||||
}
|
||||
case RoleTool:
|
||||
items = append(items, responsesFunctionCallOutput{
|
||||
Type: "function_call_output",
|
||||
CallID: msg.ToolCallID,
|
||||
Output: msg.Content,
|
||||
})
|
||||
}
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
// responsesInstructions extracts system messages into a single instructions
|
||||
// string, which is the OpenAI Responses API way to pass developer/system
|
||||
// guidance (and the most prompt-cache friendly).
|
||||
func responsesInstructions(messages []Message) string {
|
||||
var instructions []string
|
||||
for _, msg := range messages {
|
||||
if msg.Role != RoleSystem {
|
||||
continue
|
||||
}
|
||||
if text := strings.TrimSpace(msg.Content); text != "" {
|
||||
instructions = append(instructions, text)
|
||||
}
|
||||
}
|
||||
return strings.Join(instructions, "\n\n")
|
||||
}
|
||||
|
||||
func buildResponsesTools(tools []Tool) []responsesTool {
|
||||
out := make([]responsesTool, 0, len(tools))
|
||||
for _, tool := range tools {
|
||||
out = append(out, responsesTool{
|
||||
Type: "function",
|
||||
Name: tool.Function.Name,
|
||||
Description: tool.Function.Description,
|
||||
Parameters: tool.Function.Parameters,
|
||||
})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func mapReasoningEffort(level string) string {
|
||||
level = strings.ToLower(strings.TrimSpace(level))
|
||||
if level == "middle" {
|
||||
return "medium"
|
||||
}
|
||||
return level
|
||||
}
|
||||
|
||||
type responsesMessageInput struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
type responsesFunctionCallInput struct {
|
||||
Type string `json:"type"`
|
||||
CallID string `json:"call_id"`
|
||||
Name string `json:"name"`
|
||||
Arguments string `json:"arguments"`
|
||||
}
|
||||
|
||||
type responsesFunctionCallOutput struct {
|
||||
Type string `json:"type"`
|
||||
CallID string `json:"call_id"`
|
||||
Output string `json:"output"`
|
||||
}
|
||||
|
||||
type responsesTool struct {
|
||||
Type string `json:"type"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Parameters json.RawMessage `json:"parameters,omitempty"`
|
||||
}
|
||||
|
||||
// responsesStreamState tracks in-flight output items while parsing SSE events.
|
||||
type responsesStreamState struct {
|
||||
calls map[int]*responsesCallBuilder
|
||||
done bool
|
||||
}
|
||||
|
||||
type responsesCallBuilder struct {
|
||||
index int
|
||||
id string
|
||||
callType string
|
||||
name string
|
||||
arguments strings.Builder
|
||||
finalized bool
|
||||
}
|
||||
|
||||
func (s *responsesStreamState) handle(payload string, emit func(StreamEvent) error) error {
|
||||
if s.done {
|
||||
return nil
|
||||
}
|
||||
var event responsesStreamEvent
|
||||
if err := json.Unmarshal([]byte(payload), &event); err != nil {
|
||||
return fmt.Errorf("parse responses stream payload: %w", err)
|
||||
}
|
||||
|
||||
switch event.Type {
|
||||
case "error":
|
||||
s.done = true
|
||||
return event.err()
|
||||
case "response.failed":
|
||||
s.done = true
|
||||
if event.Response != nil && event.Response.Error != nil {
|
||||
return fmt.Errorf("provider error: %s: %s", event.Response.Error.Code, event.Response.Error.Message)
|
||||
}
|
||||
return fmt.Errorf("provider error: response failed")
|
||||
case "response.incomplete":
|
||||
s.done = true
|
||||
return s.finish(emit, "length", nil)
|
||||
case "response.completed":
|
||||
s.done = true
|
||||
var usage *Usage
|
||||
if event.Response != nil && event.Response.Usage != nil {
|
||||
usage = &Usage{
|
||||
PromptTokens: event.Response.Usage.InputTokens,
|
||||
CompletionTokens: event.Response.Usage.OutputTokens,
|
||||
TotalTokens: event.Response.Usage.TotalTokens,
|
||||
}
|
||||
}
|
||||
return s.finish(emit, "stop", usage)
|
||||
case "response.output_text.delta":
|
||||
if event.Delta != "" {
|
||||
return emit(StreamEvent{Content: event.Delta})
|
||||
}
|
||||
case "response.output_item.added":
|
||||
if event.Item != nil && event.Item.Type == "function_call" {
|
||||
s.builder(event.OutputIndex).applyAdded(*event.Item)
|
||||
}
|
||||
case "response.function_call_arguments.delta":
|
||||
if event.Delta != "" {
|
||||
s.builder(event.OutputIndex).arguments.WriteString(event.Delta)
|
||||
}
|
||||
case "response.function_call_arguments.done":
|
||||
builder := s.builder(event.OutputIndex)
|
||||
if event.Name != "" {
|
||||
builder.name = event.Name
|
||||
}
|
||||
if event.Arguments != "" {
|
||||
builder.arguments.Reset()
|
||||
builder.arguments.WriteString(event.Arguments)
|
||||
}
|
||||
case "response.output_item.done":
|
||||
if event.Item != nil && event.Item.Type == "function_call" {
|
||||
builder := s.builder(event.OutputIndex)
|
||||
builder.applyDone(*event.Item)
|
||||
return s.emitCall(builder, emit)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// finish flushes any not-yet-finalized function calls and reports the
|
||||
// terminal finish reason.
|
||||
func (s *responsesStreamState) finish(emit func(StreamEvent) error, finishReason string, usage *Usage) error {
|
||||
if len(s.calls) == 0 {
|
||||
return emit(StreamEvent{FinishReason: finishReason, Usage: usage})
|
||||
}
|
||||
indexes := make([]int, 0, len(s.calls))
|
||||
for index := range s.calls {
|
||||
indexes = append(indexes, index)
|
||||
}
|
||||
sort.Ints(indexes)
|
||||
var deltas []ToolCallDelta
|
||||
for _, index := range indexes {
|
||||
builder := s.calls[index]
|
||||
if builder.finalized {
|
||||
continue
|
||||
}
|
||||
builder.finalized = true
|
||||
deltas = append(deltas, builder.delta())
|
||||
}
|
||||
if len(deltas) > 0 {
|
||||
if err := emit(StreamEvent{ToolCalls: deltas}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return emit(StreamEvent{FinishReason: finishReason, Usage: usage})
|
||||
}
|
||||
|
||||
func (s *responsesStreamState) emitCall(builder *responsesCallBuilder, emit func(StreamEvent) error) error {
|
||||
if builder.finalized {
|
||||
return nil
|
||||
}
|
||||
builder.finalized = true
|
||||
return emit(StreamEvent{ToolCalls: []ToolCallDelta{builder.delta()}})
|
||||
}
|
||||
|
||||
func (s *responsesStreamState) builder(index int) *responsesCallBuilder {
|
||||
if s.calls == nil {
|
||||
s.calls = make(map[int]*responsesCallBuilder)
|
||||
}
|
||||
builder := s.calls[index]
|
||||
if builder == nil {
|
||||
builder = &responsesCallBuilder{index: index}
|
||||
s.calls[index] = builder
|
||||
}
|
||||
return builder
|
||||
}
|
||||
|
||||
func (b *responsesCallBuilder) applyAdded(item responsesStreamItem) {
|
||||
if b.id == "" {
|
||||
b.id = item.CallID
|
||||
if b.id == "" {
|
||||
b.id = item.ID
|
||||
}
|
||||
}
|
||||
if b.name == "" {
|
||||
b.name = item.Name
|
||||
}
|
||||
if item.Arguments != "" {
|
||||
b.arguments.Reset()
|
||||
b.arguments.WriteString(item.Arguments)
|
||||
}
|
||||
}
|
||||
|
||||
func (b *responsesCallBuilder) applyDone(item responsesStreamItem) {
|
||||
if item.CallID != "" {
|
||||
b.id = item.CallID
|
||||
} else if item.ID != "" {
|
||||
b.id = item.ID
|
||||
}
|
||||
if item.Name != "" {
|
||||
b.name = item.Name
|
||||
}
|
||||
if item.Arguments != "" {
|
||||
b.arguments.Reset()
|
||||
b.arguments.WriteString(item.Arguments)
|
||||
}
|
||||
}
|
||||
|
||||
func (b *responsesCallBuilder) delta() ToolCallDelta {
|
||||
id := b.id
|
||||
if id == "" {
|
||||
id = fmt.Sprintf("call_%d", b.index)
|
||||
}
|
||||
callType := b.callType
|
||||
if callType == "" {
|
||||
callType = "function"
|
||||
}
|
||||
return ToolCallDelta{
|
||||
Index: b.index,
|
||||
ID: id,
|
||||
Type: callType,
|
||||
Name: b.name,
|
||||
Arguments: b.arguments.String(),
|
||||
}
|
||||
}
|
||||
|
||||
type responsesStreamEvent struct {
|
||||
Type string `json:"type"`
|
||||
OutputIndex int `json:"output_index"`
|
||||
Delta string `json:"delta"`
|
||||
Arguments string `json:"arguments"`
|
||||
Name string `json:"name"`
|
||||
Code string `json:"code"`
|
||||
Message string `json:"message"`
|
||||
Param string `json:"param"`
|
||||
Item *responsesStreamItem `json:"item"`
|
||||
Response *responsesStreamResponse `json:"response"`
|
||||
Error *responsesStreamError `json:"error"`
|
||||
}
|
||||
|
||||
type responsesStreamItem struct {
|
||||
Type string `json:"type"`
|
||||
ID string `json:"id"`
|
||||
CallID string `json:"call_id"`
|
||||
Name string `json:"name"`
|
||||
Arguments string `json:"arguments"`
|
||||
}
|
||||
|
||||
type responsesStreamResponse struct {
|
||||
Status string `json:"status"`
|
||||
Error *responsesStreamError `json:"error"`
|
||||
Usage *responsesStreamUsage `json:"usage"`
|
||||
}
|
||||
|
||||
type responsesStreamUsage struct {
|
||||
InputTokens int `json:"input_tokens"`
|
||||
OutputTokens int `json:"output_tokens"`
|
||||
TotalTokens int `json:"total_tokens"`
|
||||
}
|
||||
|
||||
type responsesStreamError struct {
|
||||
Type string `json:"type"`
|
||||
Code string `json:"code"`
|
||||
Message string `json:"message"`
|
||||
Param string `json:"param"`
|
||||
}
|
||||
|
||||
func (e responsesStreamEvent) err() error {
|
||||
code, message, param := e.Code, e.Message, e.Param
|
||||
if code == "" && message == "" && e.Error != nil {
|
||||
code, message, param = e.Error.Code, e.Error.Message, e.Error.Param
|
||||
}
|
||||
if message == "" {
|
||||
if param != "" {
|
||||
return fmt.Errorf("provider error: %s", param)
|
||||
}
|
||||
return fmt.Errorf("provider error")
|
||||
}
|
||||
if code != "" {
|
||||
return fmt.Errorf("provider error: %s: %s", code, message)
|
||||
}
|
||||
return fmt.Errorf("provider error: %s", message)
|
||||
}
|
||||
@@ -0,0 +1,337 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestBuildResponsesRequest(t *testing.T) {
|
||||
req := ChatRequest{
|
||||
Model: "gpt-test",
|
||||
Messages: []Message{
|
||||
{Role: RoleSystem, Content: "be concise"},
|
||||
{Role: RoleUser, Content: "hi"},
|
||||
{
|
||||
Role: RoleAssistant,
|
||||
Content: "let me check",
|
||||
ToolCalls: []ToolCall{{
|
||||
ID: "call_1",
|
||||
Type: "function",
|
||||
Function: FunctionCall{
|
||||
Name: "file_read",
|
||||
Arguments: `{"path":"README.md"}`,
|
||||
},
|
||||
}},
|
||||
},
|
||||
{Role: RoleTool, ToolCallID: "call_1", Content: "contents"},
|
||||
{Role: RoleAssistant, Content: "done"},
|
||||
},
|
||||
Tools: []Tool{{
|
||||
Type: "function",
|
||||
Function: ToolFunction{
|
||||
Name: "file_read",
|
||||
Description: "Read a local file",
|
||||
Parameters: json.RawMessage(`{"type":"object","properties":{}}`),
|
||||
},
|
||||
}},
|
||||
ToolChoice: "auto",
|
||||
Extra: map[string]any{
|
||||
"reasoning": "middle",
|
||||
"thinking": "high",
|
||||
"temperature": 0.2,
|
||||
},
|
||||
}
|
||||
|
||||
payload := buildResponsesRequest(req)
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var raw map[string]any
|
||||
if err := json.Unmarshal(data, &raw); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if raw["model"] != "gpt-test" || raw["stream"] != true {
|
||||
t.Fatalf("model/stream = %#v / %#v", raw["model"], raw["stream"])
|
||||
}
|
||||
if raw["instructions"] != "be concise" {
|
||||
t.Fatalf("instructions = %#v", raw["instructions"])
|
||||
}
|
||||
if raw["tool_choice"] != "auto" {
|
||||
t.Fatalf("tool_choice = %#v", raw["tool_choice"])
|
||||
}
|
||||
|
||||
input, ok := raw["input"].([]any)
|
||||
if !ok || len(input) != 5 {
|
||||
t.Fatalf("input = %#v", raw["input"])
|
||||
}
|
||||
checkInputItem := func(index int, wantType, wantKey, wantValue string) {
|
||||
t.Helper()
|
||||
item, ok := input[index].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("input[%d] = %#v", index, input[index])
|
||||
}
|
||||
if item[wantKey] != wantValue {
|
||||
t.Fatalf("input[%d].%s = %#v, want %q", index, wantKey, item[wantKey], wantValue)
|
||||
}
|
||||
if wantType != "" && item["type"] != wantType {
|
||||
t.Fatalf("input[%d].type = %#v, want %q", index, item["type"], wantType)
|
||||
}
|
||||
}
|
||||
checkInputItem(0, "", "role", "user")
|
||||
checkInputItem(1, "", "role", "assistant")
|
||||
checkInputItem(2, "function_call", "call_id", "call_1")
|
||||
checkInputItem(3, "function_call_output", "call_id", "call_1")
|
||||
checkInputItem(4, "", "role", "assistant")
|
||||
if item := input[2].(map[string]any); item["name"] != "file_read" || item["arguments"] != `{"path":"README.md"}` {
|
||||
t.Fatalf("function_call item = %#v", item)
|
||||
}
|
||||
if item := input[3].(map[string]any); item["output"] != "contents" {
|
||||
t.Fatalf("function_call_output item = %#v", item)
|
||||
}
|
||||
|
||||
tools, ok := raw["tools"].([]any)
|
||||
if !ok || len(tools) != 1 {
|
||||
t.Fatalf("tools = %#v", raw["tools"])
|
||||
}
|
||||
tool := tools[0].(map[string]any)
|
||||
if tool["type"] != "function" || tool["name"] != "file_read" || tool["description"] != "Read a local file" {
|
||||
t.Fatalf("tool = %#v", tool)
|
||||
}
|
||||
if _, ok := tool["parameters"].(map[string]any); !ok {
|
||||
t.Fatalf("tool parameters = %#v", tool["parameters"])
|
||||
}
|
||||
|
||||
reasoning, ok := raw["reasoning"].(map[string]any)
|
||||
if !ok || reasoning["effort"] != "medium" {
|
||||
t.Fatalf("reasoning = %#v", raw["reasoning"])
|
||||
}
|
||||
if raw["thinking"] != "high" {
|
||||
t.Fatalf("thinking = %#v", raw["thinking"])
|
||||
}
|
||||
if raw["temperature"] != 0.2 {
|
||||
t.Fatalf("temperature = %#v", raw["temperature"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestResponsesInstructionsCombinesSystemMessages(t *testing.T) {
|
||||
payload := buildResponsesRequest(ChatRequest{
|
||||
Model: "gpt-test",
|
||||
Messages: []Message{
|
||||
{Role: RoleSystem, Content: "first instruction"},
|
||||
{Role: RoleUser, Content: "hi"},
|
||||
{Role: RoleSystem, Content: "second instruction"},
|
||||
},
|
||||
})
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var raw map[string]any
|
||||
if err := json.Unmarshal(data, &raw); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if raw["instructions"] != "first instruction\n\nsecond instruction" {
|
||||
t.Fatalf("instructions = %#v", raw["instructions"])
|
||||
}
|
||||
input, ok := raw["input"].([]any)
|
||||
if !ok || len(input) != 1 {
|
||||
t.Fatalf("input = %#v", raw["input"])
|
||||
}
|
||||
item, ok := input[0].(map[string]any)
|
||||
if !ok || item["role"] != "user" {
|
||||
t.Fatalf("input[0] = %#v", input[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesClientStreamsContentAndToolCalls(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/v1/responses" {
|
||||
t.Fatalf("path = %s", r.URL.Path)
|
||||
}
|
||||
if got := r.Header.Get("Authorization"); got != "Bearer sk-test" {
|
||||
t.Fatalf("authorization = %q", got)
|
||||
}
|
||||
body := mustReadBody(t, r)
|
||||
var raw map[string]any
|
||||
if err := json.NewDecoder(strings.NewReader(body)).Decode(&raw); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if raw["stream"] != true {
|
||||
t.Fatal("request did not enable stream")
|
||||
}
|
||||
if raw["tool_choice"] != "auto" {
|
||||
t.Fatalf("tool_choice = %#v", raw["tool_choice"])
|
||||
}
|
||||
input, ok := raw["input"].([]any)
|
||||
if !ok || len(input) == 0 {
|
||||
t.Fatalf("input = %#v", raw["input"])
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
events := []string{
|
||||
`{"type":"response.output_item.added","output_index":0,"item":{"type":"message","id":"msg_1","role":"assistant","content":[]}}`,
|
||||
`{"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"Hel"}`,
|
||||
`{"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"lo"}`,
|
||||
`{"type":"response.output_item.added","output_index":1,"item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"file_read","arguments":""}}`,
|
||||
`{"type":"response.function_call_arguments.delta","item_id":"fc_1","output_index":1,"delta":"{\"path\""}`,
|
||||
`{"type":"response.function_call_arguments.delta","item_id":"fc_1","output_index":1,"delta":":\"README.md\"}"}`,
|
||||
`{"type":"response.function_call_arguments.done","item_id":"fc_1","output_index":1,"name":"file_read","arguments":"{\"path\":\"README.md\"}"}`,
|
||||
`{"type":"response.output_item.done","output_index":1,"item":{"type":"function_call","id":"fc_1","call_id":"call_1","name":"file_read","arguments":"{\"path\":\"README.md\"}"}}`,
|
||||
`{"type":"response.output_item.done","output_index":0,"item":{"type":"message","id":"msg_1","role":"assistant","content":[{"type":"output_text","text":"Hello"}]}}`,
|
||||
`{"type":"response.completed","response":{"id":"resp_1","status":"completed","usage":{"input_tokens":10,"output_tokens":20,"total_tokens":30}}}`,
|
||||
}
|
||||
for _, event := range events {
|
||||
_, _ = w.Write([]byte("data: " + event + "\n\n"))
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client())
|
||||
var content strings.Builder
|
||||
var calls []ToolCallDelta
|
||||
var finishes []string
|
||||
var usage *Usage
|
||||
err := client.ChatStream(context.Background(), ChatRequest{
|
||||
Model: "gpt-test",
|
||||
Messages: []Message{{Role: RoleUser, Content: "hi"}},
|
||||
ToolChoice: "auto",
|
||||
Tools: []Tool{{
|
||||
Type: "function",
|
||||
Function: ToolFunction{
|
||||
Name: "file_read",
|
||||
Description: "Read a local file",
|
||||
Parameters: json.RawMessage(`{"type":"object"}`),
|
||||
},
|
||||
}},
|
||||
}, func(event StreamEvent) error {
|
||||
content.WriteString(event.Content)
|
||||
calls = append(calls, event.ToolCalls...)
|
||||
if event.Usage != nil {
|
||||
usage = event.Usage
|
||||
}
|
||||
if event.FinishReason != "" {
|
||||
finishes = append(finishes, event.FinishReason)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if content.String() != "Hello" {
|
||||
t.Fatalf("content = %q", content.String())
|
||||
}
|
||||
if len(calls) != 1 {
|
||||
t.Fatalf("calls = %#v", calls)
|
||||
}
|
||||
call := calls[0]
|
||||
if call.Index != 1 || call.ID != "call_1" || call.Name != "file_read" {
|
||||
t.Fatalf("call = %#v", call)
|
||||
}
|
||||
if call.Arguments != `{"path":"README.md"}` {
|
||||
t.Fatalf("arguments = %q", call.Arguments)
|
||||
}
|
||||
if len(finishes) != 1 || finishes[0] != "stop" {
|
||||
t.Fatalf("finishes = %#v", finishes)
|
||||
}
|
||||
if usage == nil || usage.PromptTokens != 10 || usage.CompletionTokens != 20 || usage.TotalTokens != 30 {
|
||||
t.Fatalf("usage = %#v", usage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesClientErrorEvent(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = w.Write([]byte(`data: {"type":"error","code":"invalid_request_error","message":"bad request"}` + "\n\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client())
|
||||
err := client.ChatStream(context.Background(), ChatRequest{Model: "test"}, func(StreamEvent) error {
|
||||
return nil
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "bad request") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesClientFailedEvent(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = w.Write([]byte(`data: {"type":"response.failed","response":{"id":"resp_1","status":"failed","error":{"code":"server_error","message":"boom"}}}` + "\n\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client())
|
||||
err := client.ChatStream(context.Background(), ChatRequest{Model: "test"}, func(StreamEvent) error {
|
||||
return nil
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "boom") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesClientIncomplete(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = w.Write([]byte(`data: {"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"partial"}` + "\n\n"))
|
||||
_, _ = w.Write([]byte(`data: {"type":"response.incomplete","response":{"id":"resp_1","status":"incomplete","incomplete_details":{"reason":"max_tokens"}}}` + "\n\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client())
|
||||
var content strings.Builder
|
||||
var finishes []string
|
||||
err := client.ChatStream(context.Background(), ChatRequest{Model: "test"}, func(event StreamEvent) error {
|
||||
content.WriteString(event.Content)
|
||||
if event.FinishReason != "" {
|
||||
finishes = append(finishes, event.FinishReason)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if content.String() != "partial" {
|
||||
t.Fatalf("content = %q", content.String())
|
||||
}
|
||||
if len(finishes) != 1 || finishes[0] != "length" {
|
||||
t.Fatalf("finishes = %#v", finishes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesClientRetriesTransientHTTPError(t *testing.T) {
|
||||
attempts := 0
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
attempts++
|
||||
if attempts == 1 {
|
||||
http.Error(w, "temporary", http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = w.Write([]byte(`data: {"type":"response.output_text.delta","item_id":"msg_1","output_index":0,"content_index":0,"delta":"ok"}` + "\n\n"))
|
||||
_, _ = w.Write([]byte(`data: {"type":"response.completed","response":{"id":"resp_1","status":"completed"}}` + "\n\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewOpenAIResponsesClient(server.URL, "sk-test", server.Client())
|
||||
var content strings.Builder
|
||||
err := client.ChatStream(context.Background(), ChatRequest{Model: "test"}, func(event StreamEvent) error {
|
||||
content.WriteString(event.Content)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if attempts != 2 {
|
||||
t.Fatalf("attempts = %d, want 2", attempts)
|
||||
}
|
||||
if content.String() != "ok" {
|
||||
t.Fatalf("content = %q", content.String())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user