3 Commits

Author SHA1 Message Date
loveuer 425f6ee5ee feat(llm): support OpenAI Responses API via per-provider api_type 2026-08-13 11:46:23 +08:00
loveuer c8d8329dbb feat(tui): codex-style input, history, bang commands, single-line status bar
- Arrow-key input history (up/down with multiline awareness)
- Bang shell commands via ! prefix (requires --yolo)
- Codex-style input bar with mode indicator (R/Y/⌘)
- Command palette virtual scrolling
- Multiline input viewport fix
- Merged model pills + footer into single status line
- Removed shift+enter (wait for TUI library support)
- Removed default footer hints
2026-06-25 20:40:50 -07:00
loveuer dabf5bfecc compact: retain newest ~20% of context; refactor startup into internal/app
- Replace fixed 6-message compact retention with ratio-based logic:
  retain ~20% of compactable estimated tokens, with a 6-message floor
  and user-turn boundary alignment
- Add compactRetainedTokenBudget, compactRecentStart, compactTurnBoundary
  helpers in internal/agent/agent.go
- Add TestAgentCompactKeepsTwentyPercentRecentContext
- Move startup/session/repl/bootstrap logic from root into internal/app
  so main.go is a thin binary entry point (15 lines)
- Update /compact docs in README.md
2026-06-24 23:59:21 -07:00
19 changed files with 2115 additions and 401 deletions
+27 -3
View File
@@ -8,6 +8,7 @@ model/provider switching, and local tools.
- TUI-first chat experience with a light theme by default. - TUI-first chat experience with a light theme by default.
- OpenAI-compatible `/v1/chat/completions` provider support. - OpenAI-compatible `/v1/chat/completions` provider support.
- OpenAI Responses API (`/v1/responses`) provider support.
- Multiple providers from `~/.agentu/config.yaml`. - Multiple providers from `~/.agentu/config.yaml`.
- Session-only provider, model, and thinking-level switching. - Session-only provider, model, and thinking-level switching.
- Read-only project tools enabled by default. - Read-only project tools enabled by default.
@@ -48,6 +49,7 @@ Use `--config <path>` to load another file.
```yaml ```yaml
providers: providers:
loveuer: loveuer:
api_type: chat
base_url: https://ai.loveuer.com base_url: https://ai.loveuer.com
api_key: ${AGENTU_API_KEY} api_key: ${AGENTU_API_KEY}
models: models:
@@ -64,6 +66,19 @@ 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 first provider by name. Use `/model provider <name>` to switch providers for the
current session. 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`/`system` messages, `function_call`, and
`function_call_output`), 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 `models` is required. `/model model <name>` is restricted to that list. Each
entry is an object with: entry is an object with:
@@ -83,7 +98,10 @@ none, middle, high, xhigh, max
``` ```
When configured or changed with `/model thinking ...`, the value is sent as a 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 named `thinking`. For OpenAI reasoning models on the
Responses API, set `thinking_param: reasoning` so the level is sent as the
official `reasoning: {"effort": ...}` field; agentu maps its `middle` level to
the official `medium` value and passes `none`, `high`, `xhigh`, `max` through.
## CLI Flags ## CLI Flags
@@ -100,22 +118,25 @@ go run . [flags]
## TUI Controls ## TUI Controls
- `enter` sends the message. - `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. - The input starts at one line and grows up to six lines.
- `page up` / `page down` scrolls the transcript. - `page up` / `page down` scrolls the transcript.
- `esc` cancels a running response. - `esc` cancels a running response.
- Type `/` to show available commands. - Type `/` to show available commands.
- In `--yolo` mode, type `! <command>` to run a shell command from the configured working directory.
Slash commands: Slash commands:
- `/clear` resets conversation history. - `/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`. - `/status` shows changed files with `git status --short`.
- `/diff [--stat] [path...]` shows the unstaged git diff. - `/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. - `/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` shows current provider, model, thinking, and switch usage.
- `/model provider <name>` switches provider for the current session. - `/model provider <name>` switches provider for the current session.
- `/model model <name>` or `/model <name>` switches model 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. - `/model thinking <none|middle|high|xhigh|max>` changes thinking for the current session.
- `/sessions` lists saved sessions. - `/sessions` lists saved sessions.
- `/resume [id]` resumes by ID or opens an interactive picker. - `/resume [id]` resumes by ID or opens an interactive picker.
@@ -171,5 +192,8 @@ Local config and build outputs are intentionally ignored by git.
- Provider/model/thinking changes are session-only; agentu does not write them - Provider/model/thinking changes are session-only; agentu does not write them
back to `~/.agentu/config.yaml`. 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 - The config schema is intentionally strict during MVP development. Deprecated
fields such as `provider` or `active_provider` are rejected. fields such as `provider` or `active_provider` are rejected.
+3
View File
@@ -1,5 +1,8 @@
providers: providers:
loveuer: 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 base_url: https://ai.loveuer.com
api_key: ${AGENTU_API_KEY} api_key: ${AGENTU_API_KEY}
models: models:
+59 -7
View File
@@ -17,7 +17,8 @@ import (
const ( const (
defaultMaxToolRounds = 50 defaultMaxToolRounds = 50
doomLoopThreshold = 3 doomLoopThreshold = 3
compactRecentMessages = 6 compactMinRecentMessages = 6
compactRecentRetentionPercent = 20
summaryMessageLimit = 8000 summaryMessageLimit = 8000
contextSummaryPrefix = "[context summary]" contextSummaryPrefix = "[context summary]"
toolLogArgumentLimit = 400 toolLogArgumentLimit = 400
@@ -181,6 +182,61 @@ func (a *Agent) SetLastUsage(usage *llm.Usage) {
copy := *usage copy := *usage
a.lastUsage = &copy a.lastUsage = &copy
} }
// 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 // Compact summarizes older conversation messages while preserving the system
// prompt and recent messages intact. // prompt and recent messages intact.
@@ -196,12 +252,8 @@ func (a *Agent) Compact(ctx context.Context) error {
if a.messages[0].Role == llm.RoleSystem { if a.messages[0].Role == llm.RoleSystem {
prefixEnd = 1 prefixEnd = 1
} }
if len(a.messages)-prefixEnd <= compactRecentMessages { recentStart := compactRecentStart(a.messages, prefixEnd)
return nil if recentStart <= prefixEnd || recentStart >= len(a.messages) {
}
recentStart := len(a.messages) - compactRecentMessages
if recentStart <= prefixEnd {
return nil return nil
} }
middle := append([]llm.Message(nil), a.messages[prefixEnd:recentStart]...) middle := append([]llm.Message(nil), a.messages[prefixEnd:recentStart]...)
+45
View File
@@ -145,6 +145,9 @@ func TestAgentCompact(t *testing.T) {
if messages[len(messages)-1].Content != "recent-answer-2" { if messages[len(messages)-1].Content != "recent-answer-2" {
t.Fatalf("last recent message = %#v", messages[len(messages)-1]) 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 { if usage := a.LastUsage(); usage == nil || usage.TotalTokens != 12 {
t.Fatalf("usage = %#v", usage) t.Fatalf("usage = %#v", usage)
} }
@@ -217,6 +220,48 @@ func longConversation() []llm.Message {
return messages 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 { func joinMessageContents(messages []llm.Message) string {
var b strings.Builder var b strings.Builder
for _, msg := range messages { for _, msg := range messages {
+171
View File
@@ -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
}
+38
View File
@@ -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 -1
View File
@@ -1,4 +1,4 @@
package main package app
import ( import (
"context" "context"
+60
View File
@@ -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()
}
}
+38
View File
@@ -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)
}
}
+23 -4
View File
@@ -204,10 +204,18 @@ func (m *Manager) CurrentThinking() string {
return m.provider.Thinking 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 { func (m *Manager) Info() string {
var b strings.Builder var b strings.Builder
fmt.Fprintf(&b, "provider: %s\n", m.CurrentProvider()) fmt.Fprintf(&b, "provider: %s\n", m.CurrentProvider())
fmt.Fprintf(&b, "model: %s\n", m.CurrentModel()) 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, "thinking: %s\n", m.CurrentThinking())
fmt.Fprintf(&b, "providers: %s\n", strings.Join(m.providerNames(), ", ")) fmt.Fprintf(&b, "providers: %s\n", strings.Join(m.providerNames(), ", "))
fmt.Fprintf(&b, "models: %s\n", strings.Join(m.provider.Models, ", ")) fmt.Fprintf(&b, "models: %s\n", strings.Join(m.provider.Models, ", "))
@@ -223,7 +231,7 @@ func (m *Manager) SetProvider(name string) (string, error) {
} }
m.providerName = name m.providerName = name
m.provider = provider 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 return "Switched provider.\n" + m.summary(), nil
} }
@@ -257,7 +265,14 @@ func (m *Manager) SetThinking(level string) (string, error) {
} }
func (m *Manager) summary() string { 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 { func (m *Manager) maxContextTokens() int {
@@ -288,9 +303,13 @@ func (m *Manager) providerNames() []string {
return names return names
} }
func (m *Manager) AvailableModels() []string {
return append([]string(nil), m.provider.Models...)
}
func contains(values []string, value string) bool { func contains(values []string, value string) bool {
for _, item := range values { for _, v := range values {
if item == value { if v == value {
return true return true
} }
} }
+72
View File
@@ -2,6 +2,8 @@ package session
import ( import (
"context" "context"
"net/http"
"net/http/httptest"
"path/filepath" "path/filepath"
"strings" "strings"
"testing" "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) { func TestManagerSaveAndResume(t *testing.T) {
manager, assistant, _ := testManager(t) manager, assistant, _ := testManager(t)
+415 -89
View File
@@ -50,6 +50,7 @@ type ModelManager interface {
Rename(name string) (string, error) Rename(name string) (string, error)
Sessions() (string, error) Sessions() (string, error)
NewSession() (string, error) NewSession() (string, error)
AvailableModels() []string
} }
func Run(ctx context.Context, assistant *agent.Agent, opts Options) error { func Run(ctx context.Context, assistant *agent.Agent, opts Options) error {
@@ -96,6 +97,7 @@ type model struct {
inputHistory []string inputHistory []string
historyIndex int historyIndex int
draftInput string draftInput string
commandCursor int
running bool running bool
cancel context.CancelFunc cancel context.CancelFunc
@@ -133,6 +135,13 @@ type commandMatch struct {
score int score int
} }
type commandSuggestion struct {
label string
description string
value string
executable bool
}
var slashCommands = []slashCommand{ var slashCommands = []slashCommand{
{Name: "/clear", Description: "reset context"}, {Name: "/clear", Description: "reset context"},
{Name: "/compact", Description: "compress context"}, {Name: "/compact", Description: "compress context"},
@@ -157,7 +166,7 @@ func newModel(ctx context.Context, assistant *agent.Agent, opts Options) model {
st := th.styles() st := th.styles()
input := textarea.New() input := textarea.New()
input.Placeholder = "Message agentu..." input.Placeholder = ""
input.Prompt = "" input.Prompt = ""
input.ShowLineNumbers = false input.ShowLineNumbers = false
input.MaxHeight = maxInputLines input.MaxHeight = maxInputLines
@@ -197,8 +206,6 @@ func (m model) Init() tea.Cmd {
m.syncMessages() m.syncMessages()
return textarea.Blink return textarea.Blink
} }
// syncMessages rebuilds the TUI display messages from the agent's conversation history.
func (m *model) syncMessages() { func (m *model) syncMessages() {
if m.agent == nil { if m.agent == nil {
return return
@@ -246,6 +253,9 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
if m.picking { if m.picking {
return m.updatePicker(msg) return m.updatePicker(msg)
} }
if handled, next, cmd := m.updateCommandSuggestions(msg); handled {
return next, cmd
}
switch msg.String() { switch msg.String() {
case "ctrl+c": case "ctrl+c":
if m.cancel != nil { if m.cancel != nil {
@@ -263,6 +273,29 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
return m, nil return m, nil
} }
return m.submit() 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": case "ctrl+p":
if m.running { if m.running {
return m, nil return m, nil
@@ -277,7 +310,7 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
m.nextInputHistory() m.nextInputHistory()
m.afterInputChanged() m.afterInputChanged()
return m, nil return m, nil
case "alt+enter", "ctrl+j": case "ctrl+j":
return m.insertInputNewline() return m.insertInputNewline()
} }
@@ -338,7 +371,6 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
if !m.running { if !m.running {
nextInput, cmd := m.input.Update(msg) nextInput, cmd := m.input.Update(msg)
m.input = nextInput m.input = nextInput
m.afterInputChanged()
cmds = append(cmds, cmd) cmds = append(cmds, cmd)
} }
@@ -356,7 +388,7 @@ func (m model) View() string {
if m.picking { if m.picking {
parts := []string{m.viewport.View()} 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...) body := lipgloss.JoinVertical(lipgloss.Left, parts...)
return m.styles.App.Width(m.width).Height(m.height).Render(body) return m.styles.App.Width(m.width).Height(m.height).Render(body)
} }
@@ -365,7 +397,7 @@ func (m model) View() string {
if activity := m.activityView(); activity != "" { if activity := m.activityView(); activity != "" {
parts = append(parts, 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...) body := lipgloss.JoinVertical(lipgloss.Left, parts...)
return m.styles.App.Width(m.width).Height(m.height).Render(body) return m.styles.App.Width(m.width).Height(m.height).Render(body)
} }
@@ -393,6 +425,10 @@ func (m *model) submit() (tea.Model, tea.Cmd) {
return m.runLocalCommand(input) return m.runLocalCommand(input)
} }
if strings.HasPrefix(input, "!") {
return m.runLocalCommand(input)
}
if strings.HasPrefix(input, "/") { if strings.HasPrefix(input, "/") {
return m.runLocalCommand(input) return m.runLocalCommand(input)
} }
@@ -438,6 +474,16 @@ func (m *model) insertInputNewline() (tea.Model, tea.Cmd) {
m.historyIndex = -1 m.historyIndex = -1
m.draftInput = "" m.draftInput = ""
m.afterInputChanged() 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 return *m, cmd
} }
@@ -467,7 +513,7 @@ func (m *model) appendAssistantChunk(chunk string) {
func (m *model) resize() { func (m *model) resize() {
width := max(40, m.width) 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.input.SetWidth(inputWidth)
m.syncInputHeight() m.syncInputHeight()
@@ -668,72 +714,25 @@ type metaPill struct {
accent bool accent bool
} }
func (m model) renderMetaPills(parts []metaPill) string { func (m model) inputModeView() string {
if len(parts) == 0 { if strings.HasPrefix(strings.TrimSpace(m.input.Value()), "!") {
return "" return m.styles.InputModeCommand.Render(iconCommand)
} }
if m.yolo {
plain := make([]string, len(parts)) return m.styles.InputModeYolo.Render("Y")
for i, part := range parts {
plain[i] = part.text
} }
return m.styles.InputModeReadonly.Render("R")
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 { func (m model) inputModeWidth() int {
mode := "read-only tools" return lipgloss.Width(m.inputModeView())
if m.yolo {
mode = "yolo tools"
}
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) inputView() string { func (m model) inputView() string {
palette := m.commandPaletteView() 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 != "" { if palette != "" {
return lipgloss.JoinVertical(lipgloss.Left, palette, box) return lipgloss.JoinVertical(lipgloss.Left, palette, box)
} }
@@ -751,13 +750,85 @@ func (m model) activityView() string {
line := m.styles.StatusInfo.Render(iconWorking) + " Working " + m.spinner.View() + " " + status + " · esc cancels" 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))) return m.styles.Activity.Render(fitLine(line, max(1, m.width-2)))
} }
func (m model) footerView() string { func (m model) footerView() string {
hint := "enter send · ctrl+j newline · alt+enter newline · ctrl+p/ctrl+n history · / commands" hint := ""
if m.running { if m.running {
hint = "esc cancel · ctrl+c quit" 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 { type eventWriter struct {
@@ -801,6 +872,7 @@ func isViewportKey(key string) bool {
} }
func (m *model) afterInputChanged() { func (m *model) afterInputChanged() {
m.commandCursor = 0
m.syncInputHeight() m.syncInputHeight()
if m.width > 0 && m.height > 0 { if m.width > 0 && m.height > 0 {
m.resize() m.resize()
@@ -882,18 +954,37 @@ func (m model) inputBlockHeight() int {
} }
func (m model) commandPaletteHeight() int { func (m model) commandPaletteHeight() int {
matches := commandMatches(commandPaletteQuery(m.input.Value())) sugs := m.commandSuggestions()
count := len(matches) total := len(sugs)
if count > maxCommandPaletteItems { if !m.showCommandPalette() {
count = maxCommandPaletteItems return 0
} }
count := min(total, maxCommandPaletteItems)
if count == 0 { if count == 0 {
count = 1 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 { func (m model) showCommandPalette() bool {
if len(m.commandSuggestions()) > 0 {
return true
}
value := strings.TrimSpace(m.input.Value()) value := strings.TrimSpace(m.input.Value())
return strings.HasPrefix(value, "/") && !strings.Contains(value, " ") return strings.HasPrefix(value, "/") && !strings.Contains(value, " ")
} }
@@ -902,18 +993,42 @@ func (m model) commandPaletteView() string {
if !m.showCommandPalette() { if !m.showCommandPalette() {
return "" 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")} lines := []string{m.styles.Header.Render(iconCommand + " Commands")}
if len(matches) == 0 {
lines = append(lines, m.styles.Muted.Render("No matching commands")) 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")) return m.styles.Palette.Width(max(20, m.width-2)).Render(strings.Join(lines, "\n"))
} }
@@ -991,6 +1106,162 @@ func subsequenceScore(candidate string, query string, base int) (int, bool) {
return base + gaps, true 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 { func (m model) slashSuggestionCount() int {
return m.commandPaletteHeight() return m.commandPaletteHeight()
} }
@@ -1007,11 +1278,7 @@ func (m model) modelInfo() string {
if m.models != nil { if m.models != nil {
return m.models.Info() return m.models.Info()
} }
mode := "read-only tools" return m.modelName
if m.yolo {
mode = "yolo tools"
}
return fmt.Sprintf("model: %s\nmode: %s\ntheme: %s", m.modelName, mode, m.theme.Mode)
} }
func (m *model) refreshContextUsage() { func (m *model) refreshContextUsage() {
@@ -1146,7 +1413,31 @@ func (m *model) handleThemeCommand(input string) (tea.Model, tea.Cmd) {
return *m, nil 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) { func (m model) handleLocalCommand(input string) (string, error) {
if strings.HasPrefix(strings.TrimSpace(input), "!") {
return m.handleBangCommand(input)
}
fields := strings.Fields(input) fields := strings.Fields(input)
if len(fields) == 0 { if len(fields) == 0 {
return "", fmt.Errorf("empty command") return "", fmt.Errorf("empty command")
@@ -1342,6 +1633,41 @@ func (m *model) updatePicker(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
return *m, nil 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 { func (m model) pickerView() string {
if !m.picking || len(m.pickerItems) == 0 { if !m.picking || len(m.pickerItems) == 0 {
return "" return ""
+227 -35
View File
@@ -21,7 +21,7 @@ func TestModelRendersChatShell(t *testing.T) {
m = next.(model) m = next.(model)
rendered := m.View() 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) { if !strings.Contains(rendered, want) {
t.Fatalf("rendered view missing %q:\n%s", want, rendered) 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"}) m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
for name, color := range map[string]any{ want := lightTheme().InputBg
"input box": m.styles.InputBox.GetBackground(), for name, color := range map[string]lipgloss.TerminalColor{
"focused base": m.input.FocusedStyle.Base.GetBackground(), "focused base": m.input.FocusedStyle.Base.GetBackground(),
"focused cursor line": m.input.FocusedStyle.CursorLine.GetBackground(), "focused cursor line": m.input.FocusedStyle.CursorLine.GetBackground(),
"blurred base": m.input.BlurredStyle.Base.GetBackground(), "blurred base": m.input.BlurredStyle.Base.GetBackground(),
"blurred cursor line": m.input.BlurredStyle.CursorLine.GetBackground(), "blurred cursor line": m.input.BlurredStyle.CursorLine.GetBackground(),
} { } {
if _, ok := color.(lipgloss.NoColor); !ok { requireStyleColor(t, name, color, want)
t.Fatalf("%s should not paint input text background", name)
}
} }
} }
@@ -58,12 +56,12 @@ func TestInputBoxIsRoomier(t *testing.T) {
m = next.(model) m = next.(model)
if got := m.styles.InputBox.GetVerticalFrameSize(); got != 2 { 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 { if got := m.styles.InputBox.GetHorizontalFrameSize(); got != 3 {
t.Fatalf("input box horizontal frame should include roomier padding, got %d", got) 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) t.Fatalf("input content width = %d, want %d", got, want)
} }
rendered := m.inputView() rendered := m.inputView()
@@ -78,13 +76,8 @@ func TestModelMetaRendersBelowInput(t *testing.T) {
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30}) next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
rendered := next.(model).View() rendered := next.(model).View()
inputIndex := strings.Index(rendered, "Message agentu") if !strings.Contains(rendered, "test-model") {
modelIndex := strings.Index(rendered, "test-model") t.Fatalf("rendered view missing %q:\n%s", "test-model", rendered)
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)
} }
} }
@@ -93,7 +86,7 @@ func TestModelMetaRendersContextUsage(t *testing.T) {
assistant.SetMessages([]llm.Message{{Role: llm.RoleUser, Content: strings.Repeat("x", 800)}}) assistant.SetMessages([]llm.Message{{Role: llm.RoleUser, Content: strings.Repeat("x", 800)}})
m := newModel(context.Background(), assistant, Options{ModelName: "test-model"}) m := newModel(context.Background(), assistant, Options{ModelName: "test-model"})
next, _ := m.Update(tea.WindowSizeMsg{Width: 120, Height: 30}) next, _ := m.Update(tea.WindowSizeMsg{Width: 120, Height: 30})
rendered := next.(model).modelMetaView() rendered := next.(model).footerView()
for _, want := range []string{"ctx ~", "/1k", "%"} { for _, want := range []string{"ctx ~", "/1k", "%"} {
if !strings.Contains(rendered, want) { 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) { func TestToolMessageRendersStructuredCard(t *testing.T) {
m := newModel(context.Background(), nil, Options{ModelName: "test-model"}) m := newModel(context.Background(), nil, Options{ModelName: "test-model"})
m.messages = append(m.messages, message{role: roleTool, content: `file_read {"path":"README.md"}`}) m.messages = append(m.messages, message{role: roleTool, content: `file_read {"path":"README.md"}`})
@@ -480,16 +490,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) { func TestModelCommand(t *testing.T) {
m := newModel(context.Background(), nil, Options{ModelName: "test-model", Yolo: true}) m := newModel(context.Background(), nil, Options{ModelName: "test-model", Yolo: true})
m.input.SetValue("/model") m.input.SetValue("/model")
next, _ := m.submit() next, _ := m.submit()
rendered := next.(model).renderMessages() rendered := next.(model).renderMessages()
for _, want := range []string{"model: test-model", "mode: yolo tools", "theme: light"} { if !strings.Contains(rendered, "test-model") {
if !strings.Contains(rendered, want) { t.Fatalf("model command output missing %q:\n%s", "test-model", rendered)
t.Fatalf("model command output missing %q:\n%s", want, rendered)
}
} }
} }
@@ -508,11 +562,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 { type fakeModelManager struct {
provider string provider string
model string model string
thinking string thinking string
sessionName string sessionName string
models []string
saved bool saved bool
} }
@@ -583,6 +750,16 @@ func (f *fakeModelManager) NewSession() (string, error) {
return "New session.", nil 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) { func TestParseThemeMode(t *testing.T) {
for input, want := range map[string]ThemeMode{ for input, want := range map[string]ThemeMode{
"": ThemeLight, "": ThemeLight,
@@ -627,9 +804,7 @@ func TestDarkThemePaintsAppBackground(t *testing.T) {
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30}) next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
m = next.(model) m = next.(model)
rendered := stripANSI(m.View()) rendered := stripANSI(m.View())
if !strings.Contains(rendered, "theme dark") { _ = rendered
t.Fatalf("rendered view missing 'theme dark':\n%s", rendered)
}
if m.theme.Mode != ThemeDark { if m.theme.Mode != ThemeDark {
t.Fatalf("theme.Mode = %q, want %q", m.theme.Mode, ThemeDark) t.Fatalf("theme.Mode = %q, want %q", m.theme.Mode, ThemeDark)
} }
@@ -643,8 +818,7 @@ func TestDarkThemeKeepsComponentBackgroundsLightweight(t *testing.T) {
next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30}) next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
m = next.(model) m = next.(model)
requireNoStyleColor(t, "Viewport.Background", m.styles.Viewport.GetBackground()) requireNoStyleColor(t, "Viewport.Background", m.styles.Viewport.GetBackground())
requireNoStyleColor(t, "InputBox.Background", m.styles.InputBox.GetBackground()) requireStyleColor(t, "InputBox.Background", m.styles.InputBox.GetBackground(), darkTheme().InputBg)
requireNoStyleColor(t, "AssistantMsg.Background", m.styles.AssistantMsg.GetBackground())
requireNoStyleColor(t, "Activity.Background", m.styles.Activity.GetBackground()) requireNoStyleColor(t, "Activity.Background", m.styles.Activity.GetBackground())
m.running = true m.running = true
m.status = "thinking..." m.status = "thinking..."
@@ -662,16 +836,34 @@ func TestDarkThemeStylesUseDarkPalette(t *testing.T) {
requireStyleColor(t, "Pill.Background", m.styles.Pill.GetBackground(), want.SurfaceAlt) 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.Base.Foreground", m.input.FocusedStyle.Base.GetForeground(), want.Text)
requireStyleColor(t, "input.FocusedStyle.Placeholder.Foreground", m.input.FocusedStyle.Placeholder.GetForeground(), want.Muted) 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 TestModelCommandShowsDarkTheme(t *testing.T) { func TestInputModeSymbolReflectsYoloAndBangCommand(t *testing.T) {
m := newModel(context.Background(), nil, Options{ModelName: "test-model", ThemeMode: ThemeDark}) m := newModel(context.Background(), nil, Options{ModelName: "test-model", Yolo: true})
m.input.SetValue("/model") next, _ := m.Update(tea.WindowSizeMsg{Width: 100, Height: 30})
next, _ := m.submit() m = next.(model)
rendered := next.(model).renderMessages()
if !strings.Contains(rendered, "theme: dark") { rendered := stripANSI(m.inputView())
t.Fatalf("model command output missing 'theme: dark':\n%s", rendered) 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) { func TestThemeSlashCommandSwitchesAtRuntime(t *testing.T) {
+31 -13
View File
@@ -43,6 +43,7 @@ type theme struct {
Background string Background string
Surface string Surface string
SurfaceAlt string SurfaceAlt string
InputBg string
Border string Border string
Text string Text string
Muted string Muted string
@@ -60,11 +61,13 @@ type styles struct {
Muted lipgloss.Style Muted lipgloss.Style
Activity lipgloss.Style Activity lipgloss.Style
Header lipgloss.Style Header lipgloss.Style
ModelMeta lipgloss.Style
Footer lipgloss.Style Footer lipgloss.Style
Viewport lipgloss.Style Viewport lipgloss.Style
InputBox lipgloss.Style InputBox lipgloss.Style
CommandHint lipgloss.Style CommandHint lipgloss.Style
InputModeReadonly lipgloss.Style
InputModeYolo lipgloss.Style
InputModeCommand lipgloss.Style
StatusReady lipgloss.Style StatusReady lipgloss.Style
StatusWarn lipgloss.Style StatusWarn lipgloss.Style
StatusError lipgloss.Style StatusError lipgloss.Style
@@ -98,8 +101,8 @@ func lightTheme() theme {
return theme{ return theme{
Mode: ThemeLight, Mode: ThemeLight,
Background: "#F8FAFC", Background: "#F8FAFC",
Surface: "#FFFFFF",
SurfaceAlt: "#EEF2F7", SurfaceAlt: "#EEF2F7",
InputBg: "#E4E4E4",
Border: "#CBD5E1", Border: "#CBD5E1",
Text: "#111827", Text: "#111827",
Muted: "#64748B", Muted: "#64748B",
@@ -116,8 +119,8 @@ func darkTheme() theme {
return theme{ return theme{
Mode: ThemeDark, Mode: ThemeDark,
Background: "#111827", Background: "#111827",
Surface: "#1F2937",
SurfaceAlt: "#0F172A", SurfaceAlt: "#0F172A",
InputBg: "#383838",
Border: "#475569", Border: "#475569",
Text: "#E5E7EB", Text: "#E5E7EB",
Muted: "#94A3B8", Muted: "#94A3B8",
@@ -153,10 +156,6 @@ func (t theme) styles() styles {
Bold(true). Bold(true).
Foreground(lipgloss.Color(t.Title)), Foreground(lipgloss.Color(t.Title)),
ModelMeta: lipgloss.NewStyle().
Foreground(lipgloss.Color(t.Muted)).
Padding(0, 1),
Footer: lipgloss.NewStyle(). Footer: lipgloss.NewStyle().
Foreground(lipgloss.Color(t.Muted)). Foreground(lipgloss.Color(t.Muted)).
Padding(0, 1), Padding(0, 1),
@@ -166,14 +165,31 @@ func (t theme) styles() styles {
InputBox: lipgloss.NewStyle(). InputBox: lipgloss.NewStyle().
Foreground(lipgloss.Color(t.Text)). Foreground(lipgloss.Color(t.Text)).
Border(lipgloss.ThickBorder()). Background(lipgloss.Color(t.InputBg)).
BorderForeground(lipgloss.Color(t.Border)). Padding(1, 2, 1, 1),
Padding(0, 2),
CommandHint: lipgloss.NewStyle(). CommandHint: lipgloss.NewStyle().
Foreground(lipgloss.Color(t.Muted)). Foreground(lipgloss.Color(t.Muted)).
Padding(0, 1), 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(). StatusReady: lipgloss.NewStyle().
Foreground(lipgloss.Color(t.Assistant)), Foreground(lipgloss.Color(t.Assistant)),
@@ -258,8 +274,10 @@ func bubbleMessageStyle(textColor, backgroundColor string) lipgloss.Style {
func applyTextareaTheme(input *textarea.Model, t theme) { func applyTextareaTheme(input *textarea.Model, t theme) {
focused := textarea.Style{ focused := textarea.Style{
Base: lipgloss.NewStyle(). Base: lipgloss.NewStyle().
Foreground(lipgloss.Color(t.Text)), Foreground(lipgloss.Color(t.Text)).
CursorLine: lipgloss.NewStyle(), Background(lipgloss.Color(t.InputBg)),
CursorLine: lipgloss.NewStyle().
Background(lipgloss.Color(t.InputBg)),
Placeholder: lipgloss.NewStyle(). Placeholder: lipgloss.NewStyle().
Foreground(lipgloss.Color(t.Muted)), Foreground(lipgloss.Color(t.Muted)),
Prompt: lipgloss.NewStyle(). Prompt: lipgloss.NewStyle().
@@ -270,7 +288,7 @@ func applyTextareaTheme(input *textarea.Model, t theme) {
Foreground(lipgloss.Color(t.Surface)), Foreground(lipgloss.Color(t.Surface)),
} }
blurred := focused blurred := focused
blurred.CursorLine = lipgloss.NewStyle() blurred.CursorLine = lipgloss.NewStyle().Background(lipgloss.Color(t.InputBg))
input.FocusedStyle = focused input.FocusedStyle = focused
input.BlurredStyle = blurred input.BlurredStyle = blurred
+2 -215
View File
@@ -1,228 +1,15 @@
package main package main
import ( import (
"bufio"
"context"
"errors"
"flag"
"fmt" "fmt"
"net/http"
"os" "os"
"os/signal"
"path/filepath"
"strings"
"agentu/internal/agent" "agentu/internal/app"
"agentu/internal/session"
"agentu/internal/tools"
"agentu/internal/tui"
"agentu/pkg/config"
"agentu/pkg/llm"
) )
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() { func main() {
if err := run(); err != nil { if err := app.Run(); err != nil {
fmt.Fprintln(os.Stderr, "agentu:", err) fmt.Fprintln(os.Stderr, "agentu:", err)
os.Exit(1) 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)
}
}
+33
View File
@@ -14,6 +14,11 @@ import (
const DefaultPath = "~/.agentu/config.yaml" const DefaultPath = "~/.agentu/config.yaml"
const (
APITypeChat = "chat"
APITypeResponses = "responses"
)
type Config struct { type Config struct {
Providers map[string]ProviderConfig `yaml:"providers"` Providers map[string]ProviderConfig `yaml:"providers"`
Agent AgentConfig `yaml:"agent"` Agent AgentConfig `yaml:"agent"`
@@ -21,6 +26,7 @@ type Config struct {
type ProviderConfig struct { type ProviderConfig struct {
Name string Name string
APIType string
BaseURL string BaseURL string
Model string Model string
Models []string Models []string
@@ -56,6 +62,7 @@ type rawConfig struct {
} }
type rawProviderConfig struct { type rawProviderConfig struct {
APIType string `yaml:"api_type"`
BaseURL string `yaml:"base_url"` BaseURL string `yaml:"base_url"`
Models []rawModelConfig `yaml:"models"` Models []rawModelConfig `yaml:"models"`
APIKey string `yaml:"api_key"` APIKey string `yaml:"api_key"`
@@ -153,6 +160,9 @@ func (c *Config) Validate() error {
for name, provider := range c.Providers { for name, provider := range c.Providers {
var missing []string var missing []string
prefix := "providers." + name 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) == "" { if strings.TrimSpace(provider.BaseURL) == "" {
missing = append(missing, prefix+".base_url") missing = append(missing, prefix+".base_url")
} }
@@ -220,6 +230,7 @@ func normalizeProvider(name string, raw rawProviderConfig) ProviderConfig {
models, modelConfigs := normalizeModelConfigs(raw.Models) models, modelConfigs := normalizeModelConfigs(raw.Models)
provider := ProviderConfig{ provider := ProviderConfig{
Name: name, Name: name,
APIType: normalizeAPIType(raw.APIType),
BaseURL: raw.BaseURL, BaseURL: raw.BaseURL,
Models: models, Models: models,
ModelConfigs: modelConfigs, ModelConfigs: modelConfigs,
@@ -235,6 +246,28 @@ func normalizeProvider(name string, raw rawProviderConfig) ProviderConfig {
return provider 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
}
func normalizeModelConfigs(rawModels []rawModelConfig) ([]string, map[string]ModelConfig) { func normalizeModelConfigs(rawModels []rawModelConfig) ([]string, map[string]ModelConfig) {
models := make([]string, 0, len(rawModels)) models := make([]string, 0, len(rawModels))
configs := make(map[string]ModelConfig, len(rawModels)) configs := make(map[string]ModelConfig, len(rawModels))
+71
View File
@@ -370,3 +370,74 @@ providers:
t.Fatalf("err = %v", err) 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)
}
}
+460
View File
@@ -0,0 +1,460 @@
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 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:
items = append(items, responsesMessageInput{Role: "system", Content: msg.Content})
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
}
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)
}
+305
View File
@@ -0,0 +1,305 @@
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["tool_choice"] != "auto" {
t.Fatalf("tool_choice = %#v", raw["tool_choice"])
}
input, ok := raw["input"].([]any)
if !ok || len(input) != 6 {
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", "system")
checkInputItem(1, "", "role", "user")
checkInputItem(2, "", "role", "assistant")
checkInputItem(3, "function_call", "call_id", "call_1")
checkInputItem(4, "function_call_output", "call_id", "call_1")
checkInputItem(5, "", "role", "assistant")
if item := input[3].(map[string]any); item["name"] != "file_read" || item["arguments"] != `{"path":"README.md"}` {
t.Fatalf("function_call item = %#v", item)
}
if item := input[4].(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 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())
}
}