From d1597444d02a25aa720dba70fdafb58c977ecc4f Mon Sep 17 00:00:00 2001 From: Luxferre Date: Sun, 16 Aug 2026 08:12:49 +0300 Subject: [PATCH] improved compacting algo and token visibility --- README.md | 49 +++++++--- main.go | 272 ++++++++++++++++++++++++++++++++++++++------------- main_test.go | 218 +++++++++++++++++++++++++++++++++++------ 3 files changed, 429 insertions(+), 110 deletions(-) diff --git a/README.md b/README.md index 7a8e76a..d684e17 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ ## About -Bantam is a minimalist, dependency-free AI agent specification with reference implementations in **Go** (`main.go` + `term_*.go`, module `code.luxferre.top/luxferre/bantam`) and **Perl 5** as **MicroBantam** (`mb`, under 100 SLOC). It provides an agentic loop capable of autonomous tool execution, shell interaction, real-time response streaming, markdown terminal rendering with box-drawing tables, Fibonacci backoff network resilience, and subagent delegation using any OpenAI-compatible completions API. +Bantam is a minimalist, dependency-free AI agent specification with reference implementations in **Go** (`main.go` + `term_*.go`, module `code.luxferre.top/luxferre/bantam`) and **Perl 5** as **MicroBantam** (`mb`, under 100 SLOC). It provides an agentic loop capable of autonomous tool execution, direct shell interaction, real-time response streaming, markdown terminal rendering with box-drawing tables, Fibonacci backoff network resilience, context window auto-discovery, token usage tracking with prompt cache breakdowns, prefix-cache-friendly conversation compaction, and subagent delegation using any OpenAI-compatible completions API. The entire philosophy of Bantam is built upon two principles: @@ -44,6 +44,7 @@ All implementations read the same `model.cfg` and `system.txt` from the current temperature=0.7 api_key=your_api_key_here stream=true + context_window=200000 ``` 2. Interactive mode: @@ -52,19 +53,31 @@ All implementations read the same `model.cfg` and `system.txt` from the current ./mb # MicroBantam (Perl 5) ``` - In interactive mode, prompts can span multiple lines: press **Ctrl+J** to insert a real line break (the cursor moves to the next line), then **Enter** to submit the whole multi-line prompt. The Go port ships its own raw-mode line editor (arrow keys move the cursor, Up/Down browse history, Backspace edits, Ctrl+C/Ctrl+D exit), so this works everywhere without dependencies. + In interactive mode, prompts can span multiple lines: press **Ctrl+J** to insert a real line break (the cursor moves to the next line), then **Enter** to submit the whole multi-line prompt. The Go port ships its own raw-mode line editor (arrow keys move the cursor, Up/Down browse history, Backspace edits, Ctrl+C clears line / interrupts in-flight run, Ctrl+D exits), working everywhere without third-party dependencies. + + After every interaction, Bantam displays token usage (prompt tokens, cached/uncached breakdown when supported by the provider, completion tokens, and context window utilization): + ```text + [tokens: 1420 prompt (1000 cached, 420 uncached) + 85 completion | context: 1420/200000 (0.7%)] + ``` Sessions are saved under `~/.bantam/sessions/` and can be managed with these commands: - `/save` — save the entire conversation to a new session file (auto-id like `20260808-190038`) and generate its summary - `/list` — list saved sessions (newest first) with their ids, timestamps, message counts and summaries - `/load ` — load a saved session (exact id or unique prefix) and continue from there - - `/compact` — summarize the conversation with the LLM and compact the context down to just the system message plus the summary + - `/compact` — compact context down to the system message and a concise summary using the LLM; the compaction prompt is appended directly to the existing message prefix to guarantee a 100% prompt cache hit - `/cfg [val]` — inspect or update a configuration parameter in `model.cfg` live - `!` — execute a shell command directly through `shell_exec` without adding the result to the conversation context (Go port) - `/help` — show all supported commands - `/clear` — reset the conversation to just the system prompt - `/quit` — exit + When context window usage reaches **60% or higher**, Bantam automatically offers to compact the conversation: + ```text + Context usage is at 62.4% (124800 / 200000 tokens). Compact conversation? [Y/n]: + ``` + + Pressing **Ctrl+C** during an active inference run or long-running shell execution cleanly interrupts the turn without appending partial or malformed responses to the conversation context. + The current conversation is also **auto-saved** to `~/.bantam/sessions/autosave.json` after every turn, on `/clear`, `/load`, `/compact`, and on exit — so you can always `/load autosave` to resume where you left off. 3. File input mode: @@ -79,27 +92,32 @@ Using these rules, everyone can build their own copy of Bantam from scratch in l ### High-Level Overview -1. **Initialization**: Read `system.txt` and `model.cfg`. Prepare an array of messages starting with the system prompt `{"role": "system", "content": system_prompt}`. -2. **Input Processing**: Take user prompt (via command-line file parameter or interactive stdin), append `{"role": "user", "content": prompt}`, and invoke `AL(cfg, messages)`. +1. **Initialization**: Read `system.txt` and `model.cfg`. Discover context window size from the `/models` endpoint (or fallback to `context_window` from `model.cfg` or 200,000 tokens). Prepare an array of messages starting with the system prompt `{"role": "system", "content": system_prompt}`. +2. **Input Processing**: Take user prompt (via command-line file parameter or interactive stdin). If prefixed with `!`, execute the command directly via `shell_exec` without appending to conversation context. Otherwise, append `{"role": "user", "content": prompt}`, and invoke `AL(cfg, messages)`. 3. **Agentic Loop (`AL`)**: - - Send `messages` and tool definitions to the OpenAI-compatible `/chat/completions` API endpoint with custom `User-Agent` headers. + - Send `messages` and tool definitions to the OpenAI-compatible `/chat/completions` API endpoint with custom `User-Agent` headers and `stream_options: {"include_usage": true}`. + - Support context cancellation (e.g. on `SIGINT` / Ctrl+C) to cleanly abort in-flight requests without appending incomplete messages. - On network or HTTP failure, retry using Fibonacci backoff delays (`1s, 1s, 2s, 3s, 5s, 8s, 13s, 21s, 34s`). - - If `stream=true`, parse SSE data chunks (`data: {...}`) in real-time to stream reasoning content (`reasoning_content`) and response text directly to stdout, bracketing the reasoning block with `--- reasoning start ---` / `--- reasoning end ---` markers. - - Reconstruct the assistant message. If `tool_calls` exist, trace the call (`[tool call: name(args)]`), validate JSON arguments, execute the requested tool (`shell_exec` or `run_subagent`), trace the result (`[tool result: name]`), append the tool response `{"role": "tool", "tool_call_id": id, "content": result}`, and repeat the loop. - - If no tool calls remain or `max_al_iterations` is reached, return the updated messages list. + - If `stream=true`, parse SSE data chunks (`data: {...}`) in real-time to stream reasoning content (`reasoning_content`) and response text directly to stdout, bracketing the reasoning block with `--- reasoning start ---` / `--- reasoning end ---` markers, rendering Markdown and tables constrained to terminal width. + - Reconstruct the assistant message and track usage tokens (`prompt_tokens`, `completion_tokens`, cached tokens). If `tool_calls` exist, trace the call (`[tool call: name(args)]`), validate JSON arguments, execute the requested tool (`shell_exec` or `run_subagent`), trace the result (`[tool result: name]`), append the tool response `{"role": "tool", "tool_call_id": id, "content": result}`, and repeat the loop. + - If no tool calls remain or `max_al_iterations` is reached, return the updated messages list and turn usage stats. +4. **Post-Turn Reporting & Compaction**: + - Display token usage and context window percentage. + - If context usage is >= 60%, prompt user to compact. + - Compaction appends `"You are now acting as a compaction engine. Summarize the preceding conversation concisely but completely..."` as a user message to the conversation, invokes the LLM (ensuring zero prompt cache misses), and resets the conversation to the system prompt and the resulting summary. ### Main program 1. Read system prompt from `system.txt` (default if missing). -2. Read model parameters from `model.cfg` (`key=value` format). +2. Read model parameters from `model.cfg` (`key=value` format) and discover context window size. 3. Prepare a new message list with the system prompt (`role: "system"`). -4. Read the first command-line parameter. If non-empty, read user prompt from the specified file. If prefixed with `!`, execute the shell command directly via `shell_exec` and exit. Otherwise, append to `messages` (`role: "user"`), run `AL(cfg, messages)`, and exit. -5. Read user prompt from standard input (with `readline` line editing and history in `~/.bantam_history`; **Ctrl+J** inserts a real newline into the line being edited). If equal to `/quit` or EOF, exit. If equal to `/clear`, reset `messages` to step 3 and return to step 5. If equal to `/save`, write the whole `messages` array to `~/.bantam/sessions/.json` (with an auto-generated summary) and return to step 5. If equal to `/list`, print saved sessions and their summaries and return to step 5. If starting with `/load`, replace `messages` with the saved session's messages (by exact id or unique prefix) and return to step 5. If equal to `/compact`, ask the LLM to summarize the conversation, replace `messages` with `[system, summary-user-message]`, and return to step 5. If starting with `/cfg`, display the current value (`/cfg `) or update `model.cfg` live (`/cfg `) and return to step 5. If starting with `!`, execute the command directly via `shell_exec` without adding the result to `messages` and return to step 5. If equal to `/help`, print the command list and return to step 5. After every user turn and on exit, auto-save `messages` to `~/.bantam/sessions/autosave.json`. -6. Append user prompt to `messages` (`role: "user"`), run `AL(cfg, messages)`, and go to step 5. +4. Read the first command-line parameter. If non-empty, read user prompt from the specified file. If prefixed with `!`, execute the shell command directly via `shell_exec` and exit. Otherwise, append to `messages` (`role: "user"`), run `AL(cfg, messages)`, display token usage, and exit. +5. Read user prompt from standard input (with `readline` line editing and history in `~/.bantam_history`; **Ctrl+J** inserts a real newline into the line being edited). If equal to `/quit` or EOF, exit. If equal to `/clear`, reset `messages` to step 3 and return to step 5. If equal to `/save`, write the whole `messages` array to `~/.bantam/sessions/.json` (with an auto-generated summary) and return to step 5. If equal to `/list`, print saved sessions and their summaries and return to step 5. If starting with `/load`, replace `messages` with the saved session's messages (by exact id or unique prefix) and return to step 5. If equal to `/compact`, ask the LLM to summarize the conversation by appending the compaction prompt to preserve KV cache, replace `messages` with `[system, summary-user-message]`, and return to step 5. If starting with `/cfg`, display the current value (`/cfg `) or update `model.cfg` live (`/cfg `) and return to step 5. If starting with `!`, execute the command directly via `shell_exec` without adding the result to `messages` and return to step 5. If equal to `/help`, print the command list and return to step 5. After every user turn and on exit, auto-save `messages` to `~/.bantam/sessions/autosave.json`. +6. Append user prompt to `messages` (`role: "user"`), run `AL(cfg, messages)`, display token usage, check 60% context threshold for auto-compaction, and go to step 5. ### Agentic loop (`AL(cfg, messages)`) function -1. Call OpenAI-compatible Completions API (`POST {endpoint}/chat/completions`) forwarding relevant parameters from `cfg` (`model`, `temperature`, `stream`, `reasoning_effort`, etc., excluding internal agent configs like `endpoint`, `api_key`, `timeout`, `shell_timeout`, `max_al_iterations`, `color`) and optional `api_key` bearer header. +1. Call OpenAI-compatible Completions API (`POST {endpoint}/chat/completions`) forwarding relevant parameters from `cfg` (`model`, `temperature`, `stream`, `reasoning_effort`, etc., excluding internal agent configs like `endpoint`, `api_key`, `timeout`, `shell_timeout`, `max_al_iterations`, `color`, `context_window`) and optional `api_key` bearer header. - Set custom `User-Agent` header (`Mozilla/5.0 (compatible; Bantam/1.0)`) to avoid gateway 403 blocks. - Retry network/HTTP errors with Fibonacci backoff delays (`1s, 1s, 2s, 3s, 5s, 8s, 13s, 21s, 34s`). - If `stream=true`, parse SSE stream (`data: {...}`) for real-time reasoning and text output, bracketing reasoning with `--- reasoning start ---` / `--- reasoning end ---` markers. @@ -111,7 +129,7 @@ Using these rules, everyone can build their own copy of Bantam from scratch in l - Output a trace log of the result (`[tool result: name]`). - Append tool result message (`role: "tool"`, `tool_call_id`, `content`: result string) to `messages`. - Loop back to step 1. -4. If no pending tool calls (or if `max_al_iterations` is reached), stop and return `messages`. +4. If no pending tool calls (or if `max_al_iterations` is reached), stop and return `messages` and accumulated usage. If the API rejects the request with an `Invalid assistant message: content or tool_calls must be set` error (usually caused by a previously cut-off stream that left an empty assistant message in the session), all implementations strip the last `assistant`-role message from the session and retry the call. @@ -128,6 +146,7 @@ If the API rejects the request with an `Invalid assistant message: content or to - `timeout` (HTTP timeout in seconds for LLM API calls, default 300; in the Go port it bounds connection setup and time-to-first-byte, so long streaming responses are not cut off mid-stream) - `shell_timeout` (timeout in seconds for `shell_exec` commands, default 120) - `max_al_iterations` (max tool-call loop iterations per `AL()` invocation, default 1000) +- `context_window` (context window size in tokens, auto-discovered from `/models` API if available, fallback to this setting, default 200000) The Go port's built-in editor tracks the cursor with its own column math (terminal auto-wrap aware) and redraws from the first line of the buffer, so wrapped input stays clean at any terminal width. diff --git a/main.go b/main.go index 6dd67f0..498e302 100644 --- a/main.go +++ b/main.go @@ -44,10 +44,11 @@ type Cfg struct { MaxALIterations int Stream bool Color string + ContextWindow int Raw map[string]string } -var defCfg = Cfg{"https://opencode.ai/zen/v1", "big-pickle", "-", 0.7, 300, 120, 1000, true, "auto", nil} +var defCfg = Cfg{"https://opencode.ai/zen/v1", "big-pickle", "-", 0.7, 300, 120, 1000, true, "auto", 200000, nil} func atoiD(s string, d int) int { if v, e := strconv.Atoi(strings.TrimSpace(s)); e == nil { @@ -56,6 +57,50 @@ func atoiD(s string, d int) int { return d } +func queryModelsContextWindow(cfg *Cfg) int { + client := &http.Client{Timeout: 3 * time.Second} + req, err := http.NewRequest("GET", strings.TrimRight(cfg.Endpoint, "/")+"/models", nil) + if err != nil { return 0 } + req.Header.Set("User-Agent", "Mozilla/5.0 (compatible; Bantam/1.0)") + if cfg.APIKey != "" && cfg.APIKey != "-" { req.Header.Set("Authorization", "Bearer "+cfg.APIKey) } + resp, err := client.Do(req) + if err != nil || resp.StatusCode >= 400 { return 0 } + defer resp.Body.Close() + var res struct { + Data []map[string]any `json:"data"` + Models []map[string]any `json:"models"` + } + if json.NewDecoder(resp.Body).Decode(&res) != nil { return 0 } + list := res.Data + if len(list) == 0 { list = res.Models } + for _, item := range list { + id, _ := item["id"].(string) + if id == cfg.Model || strings.EqualFold(id, cfg.Model) { + for _, key := range []string{"context_window", "context_length", "max_context_length", "max_model_len", "context_size", "max_tokens", "max_input_tokens"} { + if val, ok := item[key]; ok { + switch v := val.(type) { + case float64: + if v > 0 { return int(v) } + case string: + if n := atoiD(v, 0); n > 0 { return n } + } + } + } + } + } + return 0 +} + +func fetchContextWindow(cfg *Cfg) int { + if cw := queryModelsContextWindow(cfg); cw > 0 { + return cw + } + if v, ok := cfg.Raw["context_window"]; ok { + return atoiD(v, 200000) + } + return 200000 +} + func getCfg(path string) Cfg { cfg := defCfg cfg.Raw = map[string]string{ @@ -63,6 +108,7 @@ func getCfg(path string) Cfg { "api_key": cfg.APIKey, "stream": strconv.FormatBool(cfg.Stream), "color": cfg.Color, "timeout": strconv.Itoa(cfg.Timeout), "shell_timeout": strconv.Itoa(cfg.ShellTimeout), "max_al_iterations": strconv.Itoa(cfg.MaxALIterations), + "context_window": strconv.Itoa(cfg.ContextWindow), } if d, err := os.ReadFile(path); err == nil { for _, ln := range strings.Split(string(d), "\n") { @@ -81,6 +127,7 @@ func getCfg(path string) Cfg { case "max_al_iterations": cfg.MaxALIterations = atoiD(v, cfg.MaxALIterations) case "stream": cfg.Stream = v == "true" || v == "1" || v == "yes" case "color": cfg.Color = v + case "context_window": cfg.ContextWindow = atoiD(v, cfg.ContextWindow) } } } @@ -590,6 +637,48 @@ var TOOLS = []map[string]any{ func strp(s string) *string { return &s } +type Usage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + PromptTokensDetails struct { + CachedTokens int `json:"cached_tokens"` + } `json:"prompt_tokens_details"` + CachedTokens int `json:"cached_tokens"` +} + +func (u Usage) Cached() int { + if u.PromptTokensDetails.CachedTokens > 0 { return u.PromptTokensDetails.CachedTokens } + return u.CachedTokens +} + +func estTokens(msgs []Message) int { + chars := 0 + for _, m := range msgs { + if m.Content != nil { chars += len(*m.Content) } + chars += len(m.ReasoningContent) + for _, tc := range m.ToolCalls { + chars += len(tc.Function.Name) + len(tc.Function.Arguments) + } + } + if chars == 0 { return 0 } + t := chars / 4 + if t == 0 { t = 1 } + return t +} + +func formatUsage(u Usage, cw int) string { + if cw <= 0 { cw = 200000 } + pct := float64(u.PromptTokens) * 100.0 / float64(cw) + cached := u.Cached() + if cached > 0 { + uncached := u.PromptTokens - cached + if uncached < 0 { uncached = 0 } + return fmt.Sprintf("[tokens: %d prompt (%d cached, %d uncached) + %d completion | context: %d/%d (%.1f%%)]", u.PromptTokens, cached, uncached, u.CompletionTokens, u.PromptTokens, cw, pct) + } + return fmt.Sprintf("[tokens: %d prompt + %d completion | context: %d/%d (%.1f%%)]", u.PromptTokens, u.CompletionTokens, u.PromptTokens, cw, pct) +} + type streamDelta struct { Choices []struct { Delta struct { @@ -606,6 +695,7 @@ type streamDelta struct { } `json:"tool_calls"` } `json:"delta"` } `json:"choices"` + Usage *Usage `json:"usage"` } func sanitizeMessages(msgs []Message) { @@ -629,14 +719,15 @@ func isInvalidAssistantErr(err error) bool { return strings.Contains(s, "Invalid assistant message") || strings.Contains(s, "content or tool_calls must be set") } -func llm(ctx context.Context, cfg *Cfg, msgs []Message, tools []map[string]any) (Message, error) { - if err := ctx.Err(); err != nil { return Message{}, err } +func llm(ctx context.Context, cfg *Cfg, msgs []Message, tools []map[string]any) (Message, Usage, error) { + if err := ctx.Err(); err != nil { return Message{}, Usage{}, err } sanitizeMessages(msgs) p := map[string]any{"model": cfg.Model, "temperature": cfg.Temperature, "messages": msgs, "stream": cfg.Stream} if tools != nil { p["tools"] = tools } + if cfg.Stream { p["stream_options"] = map[string]any{"include_usage": true} } for k, v := range cfg.Raw { switch k { - case "endpoint", "api_key", "timeout", "shell_timeout", "max_al_iterations", "color": + case "endpoint", "api_key", "timeout", "shell_timeout", "max_al_iterations", "color", "context_window": continue default: var jv any @@ -660,7 +751,7 @@ func llm(ctx context.Context, cfg *Cfg, msgs []Message, tools []map[string]any) for i := 0; i <= len(fib); i++ { if err := ctx.Err(); err != nil { if COL { fmt.Print("\r\033[K") } - return Message{}, err + return Message{}, Usage{}, err } if COL { fmt.Print("\r" + pend) } else { fmt.Println(pend) } req, _ := http.NewRequestWithContext(ctx, "POST", strings.TrimRight(cfg.Endpoint, "/")+"/chat/completions", bytes.NewReader(body)) @@ -681,23 +772,23 @@ func llm(ctx context.Context, cfg *Cfg, msgs []Message, tools []map[string]any) if err == nil { break } if COL { fmt.Print("\r\033[K") } if errors.Is(err, context.Canceled) || ctx.Err() != nil { - return Message{}, ctx.Err() + return Message{}, Usage{}, ctx.Err() } if is4xxClientErr { - return Message{}, err + return Message{}, Usage{}, err } if i < len(fib) { fmt.Println(c(fmt.Sprintf("[network error: %v, retrying in %ds...]", err, fib[i]), 31)) select { case <-ctx.Done(): - return Message{}, ctx.Err() + return Message{}, Usage{}, ctx.Err() case <-time.After(time.Duration(fib[i]) * time.Second): } } } if err != nil { if COL { fmt.Print("\r\033[K") } - return Message{}, err + return Message{}, Usage{}, err } defer resp.Body.Close() if COL { fmt.Print("\r\033[K") } @@ -709,25 +800,39 @@ func llm(ctx context.Context, cfg *Cfg, msgs []Message, tools []map[string]any) Reasoning string `json:"reasoning"` } `json:"message"` } `json:"choices"` + Usage Usage `json:"usage"` } if err := json.NewDecoder(resp.Body).Decode(&cr); err != nil { - if errors.Is(err, context.Canceled) || ctx.Err() != nil { return Message{}, ctx.Err() } - return Message{}, err + if errors.Is(err, context.Canceled) || ctx.Err() != nil { return Message{}, Usage{}, ctx.Err() } + return Message{}, Usage{}, err } - if len(cr.Choices) == 0 { return Message{}, errors.New("empty choices in LLM response") } + if len(cr.Choices) == 0 { return Message{}, Usage{}, errors.New("empty choices in LLM response") } m := cr.Choices[0].Message.Message if m.ReasoningContent == "" { m.ReasoningContent = cr.Choices[0].Message.Reasoning } - return m, nil + u := cr.Usage + if u.PromptTokens == 0 { + u.PromptTokens = estTokens(msgs) + u.CompletionTokens = estTokens([]Message{m}) + u.TotalTokens = u.PromptTokens + u.CompletionTokens + } + return m, u, nil } - return parseStream(ctx, resp.Body) + m, u, err := parseStream(ctx, resp.Body) + if err == nil && u.PromptTokens == 0 { + u.PromptTokens = estTokens(msgs) + u.CompletionTokens = estTokens([]Message{m}) + u.TotalTokens = u.PromptTokens + u.CompletionTokens + } + return m, u, err } -func parseStream(ctx context.Context, r io.Reader) (Message, error) { +func parseStream(ctx context.Context, r io.Reader) (Message, Usage, error) { var content, reas string var rh, ch bool var lineBuf string var mdSt mdState var tblBuf []string + var lastUsage Usage tcs := map[int]*ToolCall{} var order []int @@ -741,13 +846,17 @@ func parseStream(ctx context.Context, r io.Reader) (Message, error) { sc := bufio.NewScanner(r) sc.Buffer(make([]byte, 0, 64*1024), 4*1024*1024) for sc.Scan() { - if err := ctx.Err(); err != nil { return Message{}, err } + if err := ctx.Err(); err != nil { return Message{}, lastUsage, err } ln := strings.TrimSpace(sc.Text()) if !strings.HasPrefix(ln, "data:") { continue } data := strings.TrimSpace(ln[5:]) if data == "[DONE]" { break } var d streamDelta - if json.Unmarshal([]byte(data), &d) != nil || len(d.Choices) == 0 { continue } + if json.Unmarshal([]byte(data), &d) != nil { continue } + if d.Usage != nil && (d.Usage.PromptTokens > 0 || d.Usage.TotalTokens > 0) { + lastUsage = *d.Usage + } + if len(d.Choices) == 0 { continue } dl := d.Choices[0].Delta rc := dl.ReasoningContent if rc == "" { rc = dl.Reasoning } @@ -792,7 +901,7 @@ func parseStream(ctx context.Context, r io.Reader) (Message, error) { if tc.Function.Arguments != "" { t.Function.Arguments += tc.Function.Arguments } } } - if err := ctx.Err(); err != nil { return Message{}, err } + if err := ctx.Err(); err != nil { return Message{}, lastUsage, err } flushTable() if lineBuf != "" { if !mdSt.inCode && isTableLine(lineBuf) { @@ -812,7 +921,10 @@ func parseStream(ctx context.Context, r io.Reader) (Message, error) { m.ToolCalls = make([]ToolCall, 0, len(order)) for _, idx := range order { m.ToolCalls = append(m.ToolCalls, *tcs[idx]) } } - return m, sc.Err() + if lastUsage.TotalTokens == 0 && lastUsage.PromptTokens > 0 { + lastUsage.TotalTokens = lastUsage.PromptTokens + lastUsage.CompletionTokens + } + return m, lastUsage, sc.Err() } func shell(ctx context.Context, cmd string, timeout int) string { @@ -848,14 +960,15 @@ func last(msgs []Message) string { return "" } -func AL(ctx context.Context, cfg *Cfg, msgs []Message, sp string, depth int) ([]Message, error) { +func AL(ctx context.Context, cfg *Cfg, msgs []Message, sp string, depth int) ([]Message, Usage, error) { done := false + var turnUsage Usage for i := 0; i < cfg.MaxALIterations && !done; i++ { - if err := ctx.Err(); err != nil { return msgs, err } - m, err := llm(ctx, cfg, msgs, TOOLS) + if err := ctx.Err(); err != nil { return msgs, turnUsage, err } + m, u, err := llm(ctx, cfg, msgs, TOOLS) if err != nil { if errors.Is(err, context.Canceled) || ctx.Err() != nil { - return msgs, err + return msgs, turnUsage, err } if isInvalidAssistantErr(err) { stripped := false @@ -869,8 +982,12 @@ func AL(ctx context.Context, cfg *Cfg, msgs []Message, sp string, depth int) ([] } if stripped { continue } } - return msgs, err + return msgs, turnUsage, err } + turnUsage.PromptTokens = u.PromptTokens + turnUsage.CompletionTokens += u.CompletionTokens + turnUsage.TotalTokens += u.TotalTokens + if u.Cached() > 0 { turnUsage.CachedTokens = u.Cached() } for j := range m.ToolCalls { tc := &m.ToolCalls[j] astr := tc.Function.Arguments @@ -889,7 +1006,7 @@ func AL(ctx context.Context, cfg *Cfg, msgs []Message, sp string, depth int) ([] } if len(m.ToolCalls) == 0 { done = true; break } for _, tc := range m.ToolCalls { - if err := ctx.Err(); err != nil { return msgs, err } + if err := ctx.Err(); err != nil { return msgs, turnUsage, err } fn, astr := tc.Function.Name, tc.Function.Arguments fmt.Println(c(fmt.Sprintf("[tool call: %s(%s)]", fn, astr), 33)) res, sty := "", 2 @@ -901,7 +1018,7 @@ func AL(ctx context.Context, cfg *Cfg, msgs []Message, sp string, depth int) ([] case "shell_exec": cmd, _ := a["command"].(string) res = shell(ctx, cmd, cfg.ShellTimeout) - if err := ctx.Err(); err != nil { return msgs, err } + if err := ctx.Err(); err != nil { return msgs, turnUsage, err } case "run_subagent": pr, _ := a["prompt"].(string) if depth >= MAX_DEPTH { @@ -911,12 +1028,14 @@ func AL(ctx context.Context, cfg *Cfg, msgs []Message, sp string, depth int) ([] {Role: "system", Content: strp(sp + "\n\nImportant: this is a child agent")}, {Role: "user", Content: strp(pr)}, } - if subr, err := AL(ctx, cfg, subMsgs, sp, depth+1); err != nil { + if subr, subu, err := AL(ctx, cfg, subMsgs, sp, depth+1); err != nil { if errors.Is(err, context.Canceled) || ctx.Err() != nil { - return msgs, err + return msgs, turnUsage, err } res, sty = "[subagent error: "+err.Error()+"]", 31 } else { + turnUsage.CompletionTokens += subu.CompletionTokens + turnUsage.TotalTokens += subu.TotalTokens res, sty = last(subr), 2 } } @@ -931,7 +1050,7 @@ func AL(ctx context.Context, cfg *Cfg, msgs []Message, sp string, depth int) ([] if !done { msgs = append(msgs, Message{Role: "assistant", Content: strp(fmt.Sprintf("[max AL iterations (%d) reached]", cfg.MaxALIterations))}) } - return msgs, nil + return msgs, turnUsage, nil } func homeDir() string { @@ -1028,46 +1147,38 @@ func autosave(msgs []Message) { os.WriteFile(filepath.Join(sdir(), "autosave.json"), b, 0644) } -func summarize(ctx context.Context, cfg *Cfg, msgs []Message) (string, error) { - var sb strings.Builder - for _, m := range msgs { - if m.Role == "system" { continue } - ct := "" - if m.Content != nil { ct = *m.Content } - if ct == "" && len(m.ToolCalls) > 0 { - jc := make([]map[string]any, 0, len(m.ToolCalls)) - for _, tc := range m.ToolCalls { - jc = append(jc, map[string]any{"function": map[string]any{"name": tc.Function.Name, "arguments": tc.Function.Arguments}}) - } - b, _ := json.Marshal(jc) - ct = string(b) - } - if ct == "" { continue } - if len(ct) > 4000 { ct = ct[:4000] + "...[truncated]" } - sb.WriteString(m.Role + ": " + ct + "\n\n") - } - if sb.Len() == 0 { return "", errors.New("no conversation to summarize") } - joined := sb.String() - if len(joined) > 100000 { joined = joined[len(joined)-100000:] + "\n...[earlier parts truncated]" } - sys := "You are a conversation summarizer for an AI agent's context window. Summarize concisely but completely, preserving all important facts, decisions, code, errors, and the current task state, so the agent can continue the work without the original messages. Output only the summary." - cc := *cfg - cc.Stream = false - m, err := llm(ctx, &cc, []Message{{Role: "system", Content: strp(sys)}, {Role: "user", Content: strp("Summarize this conversation:\n\n" + joined)}}, nil) - if err != nil { return "", err } - s := "" - if m.Content != nil { s = *m.Content } - if s == "" { s = m.ReasoningContent } - if strings.TrimSpace(s) == "" { return "", errors.New("LLM returned an empty summary") } - return strings.TrimSpace(s), nil -} +const compactionPrompt = "You are now acting as a compaction engine. Summarize the preceding conversation concisely but completely, preserving all important facts, decisions, code snippets, tool outputs, errors, and current task state so work can seamlessly continue. Output only the summary." func compact(ctx context.Context, cfg *Cfg, msgs []Message) ([]Message, string, error) { if len(msgs) == 0 || msgs[0].Role != "system" { return msgs, "", errors.New("session has no system message") } - s, err := summarize(ctx, cfg, msgs) + if len(msgs) <= 1 { + return msgs, "", errors.New("nothing to compact") + } + cMsgs := append(append([]Message{}, msgs...), Message{ + Role: "user", + Content: strp(compactionPrompt), + }) + cc := *cfg + cc.Stream = false + m, _, err := llm(ctx, &cc, cMsgs, nil) if err != nil { return msgs, "", err } - return []Message{{Role: "system", Content: msgs[0].Content}, {Role: "user", Content: strp("Summary of the previous conversation:\n" + s + "\n\nPlease continue from here.")}}, s, nil + s := "" + if m.Content != nil { s = *m.Content } + if s == "" { s = m.ReasoningContent } + s = strings.TrimSpace(s) + if s == "" { return msgs, "", errors.New("LLM returned an empty summary") } + newMsgs := []Message{ + {Role: "system", Content: msgs[0].Content}, + {Role: "user", Content: strp("Summary of the previous conversation:\n" + s + "\n\nPlease continue from here.")}, + } + return newMsgs, s, nil +} + +func summarize(ctx context.Context, cfg *Cfg, msgs []Message) (string, error) { + _, s, err := compact(ctx, cfg, msgs) + return s, err } func loadHistory() { @@ -1270,6 +1381,7 @@ func readLine(prompt string) (string, bool) { func main() { sp := prompt("system.txt") cfg := getCfg("model.cfg") + cfg.ContextWindow = fetchContextWindow(&cfg) COL = col(cfg) stdin = bufio.NewReader(os.Stdin) msgs := []Message{{Role: "system", Content: strp(sp)}} @@ -1295,7 +1407,8 @@ func main() { } msgs = append(msgs, Message{Role: "user", Content: strp(u)}) sigCtx, cancel := signal.NotifyContext(context.Background(), os.Interrupt) - msgs, err = AL(sigCtx, &cfg, msgs, sp, 0) + var usg Usage + msgs, usg, err = AL(sigCtx, &cfg, msgs, sp, 0) cancel() if err != nil { if errors.Is(err, context.Canceled) || sigCtx.Err() != nil { @@ -1305,12 +1418,13 @@ func main() { } os.Exit(1) } + fmt.Println(c(formatUsage(usg, cfg.ContextWindow), 2)) autosave(msgs) return } loadHistory() fmt.Println(c("Bantam Agent ready", 1, 32) + c(" (Ctrl+J = new line)", 2)) - fmt.Println(c(fmt.Sprintf("endpoint: %s model: %s temp: %v", cfg.Endpoint, cfg.Model, cfg.Temperature), 2)) + fmt.Println(c(fmt.Sprintf("endpoint: %s model: %s temp: %v context: %d", cfg.Endpoint, cfg.Model, cfg.Temperature, cfg.ContextWindow), 2)) for { u, ok := readLine(c("> ", 1, 36)) if !ok { @@ -1397,6 +1511,9 @@ func main() { continue } cfg = getCfg("model.cfg") + if k == "model" || k == "endpoint" || k == "api_key" { + cfg.ContextWindow = fetchContextWindow(&cfg) + } COL = col(cfg) fmt.Println(c(fmt.Sprintf("[config updated: %s=%s]", k, v), 32)) } else { @@ -1424,7 +1541,7 @@ func main() { turnMsgs := append([]Message{}, msgs...) turnMsgs = append(turnMsgs, Message{Role: "user", Content: strp(u)}) sigCtx, cancel := signal.NotifyContext(context.Background(), os.Interrupt) - resMsgs, err := AL(sigCtx, &cfg, turnMsgs, sp, 0) + resMsgs, usg, err := AL(sigCtx, &cfg, turnMsgs, sp, 0) cancel() if err != nil { if errors.Is(err, context.Canceled) || sigCtx.Err() != nil { @@ -1436,6 +1553,29 @@ func main() { } msgs = resMsgs autosave(msgs) + fmt.Println(c(formatUsage(usg, cfg.ContextWindow), 2)) + pct := float64(usg.PromptTokens) * 100.0 / float64(cfg.ContextWindow) + if pct >= 60.0 && len(msgs) > 1 { + fmt.Print(c(fmt.Sprintf("Context usage is at %.1f%% (%d / %d tokens). Compact conversation? [Y/n]: ", pct, usg.PromptTokens, cfg.ContextWindow), 33)) + ans, ok := readPlain("") + if ok { + ans = strings.TrimSpace(strings.ToLower(ans)) + if ans == "" || ans == "y" || ans == "yes" { + fmt.Println(c("[compacting conversation...]", 33)) + sigCtx, cancel := signal.NotifyContext(context.Background(), os.Interrupt) + nm, sm, err := compact(sigCtx, &cfg, msgs) + cancel() + if err != nil { + fmt.Println(c("[compact failed: "+err.Error()+"]", 31)) + } else { + msgs = nm + autosave(msgs) + fmt.Println(c(fmt.Sprintf("[compacted to %d messages]", len(msgs)), 32)) + fmt.Println(c("--- summary ---", 33) + "\n" + c(sm, 2)) + } + } + } + } } done: autosave(msgs) diff --git a/main_test.go b/main_test.go index a871474..c309e32 100644 --- a/main_test.go +++ b/main_test.go @@ -59,7 +59,7 @@ func TestParseStreamReasoningNoDuplication(t *testing.T) { `data: [DONE]`, }, "\n") - msg, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) + msg, _, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) if err != nil { t.Fatalf("unexpected parseStream error: %v", err) } @@ -91,7 +91,7 @@ func TestParseStreamReasoningAlias(t *testing.T) { `data: [DONE]`, }, "\n") - msg, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) + msg, _, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) if err != nil { t.Fatalf("unexpected parseStream error: %v", err) } @@ -109,7 +109,7 @@ func TestParseStreamContentOnly(t *testing.T) { `data: [DONE]`, }, "\n") - msg, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) + msg, _, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) if err != nil { t.Fatalf("unexpected parseStream error: %v", err) } @@ -130,7 +130,7 @@ func TestParseStreamReasoningOnly(t *testing.T) { `data: [DONE]`, }, "\n") - msg, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) + msg, _, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) if err != nil { t.Fatalf("unexpected parseStream error: %v", err) } @@ -144,7 +144,7 @@ func TestParseStreamReasoningOnly(t *testing.T) { func TestParseStreamEmpty(t *testing.T) { for _, in := range []string{"", "\n\n", "event: message\n\n"} { - msg, err := parseStream(context.Background(), strings.NewReader(in)) + msg, _, err := parseStream(context.Background(), strings.NewReader(in)) if err != nil { t.Fatalf("unexpected parseStream error for input %q: %v", in, err) } @@ -162,7 +162,7 @@ func TestParseStreamToolCallSplitAcrossChunks(t *testing.T) { `data: [DONE]`, }, "\n") - msg, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) + msg, _, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) if err != nil { t.Fatalf("unexpected parseStream error: %v", err) } @@ -188,7 +188,7 @@ func TestParseStreamMultipleToolCallsKeepFirstAppearanceOrder(t *testing.T) { `data: [DONE]`, }, "\n") - msg, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) + msg, _, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) if err != nil { t.Fatalf("unexpected parseStream error: %v", err) } @@ -210,7 +210,7 @@ func TestParseStreamJunkAndNoChoicesIgnored(t *testing.T) { `data: [DONE]`, }, "\n") - msg, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) + msg, _, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) if err != nil { t.Fatalf("unexpected parseStream error: %v", err) } @@ -226,7 +226,7 @@ func TestParseStreamReasoningAfterContent(t *testing.T) { `data: [DONE]`, }, "\n") - msg, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) + msg, _, err := parseStream(context.Background(), bytes.NewBufferString(sseData)) if err != nil { t.Fatalf("unexpected parseStream error: %v", err) } @@ -707,11 +707,11 @@ func TestAutosave(t *testing.T) { // ---------- summarize / compact (error paths only, no network) ---------- func TestSummarizeEmpty(t *testing.T) { - if _, err := summarize(context.Background(), &Cfg{}, nil); err == nil || !strings.Contains(err.Error(), "no conversation") { - t.Errorf("expected no-conversation error, got %v", err) + if _, err := summarize(context.Background(), &Cfg{}, nil); err == nil || !strings.Contains(err.Error(), "no system message") { + t.Errorf("expected no-system-message error, got %v", err) } - if _, err := summarize(context.Background(), &Cfg{}, []Message{{Role: "system", Content: strp("sys")}}); err == nil { - t.Errorf("expected error for system-only conversation") + if _, err := summarize(context.Background(), &Cfg{}, []Message{{Role: "system", Content: strp("sys")}}); err == nil || !strings.Contains(err.Error(), "nothing to compact") { + t.Errorf("expected nothing-to-compact error for system-only conversation, got %v", err) } } @@ -724,12 +724,19 @@ func TestCompactNoSystem(t *testing.T) { t.Errorf("expected original messages on error") } msgs2, _, err2 := compact(context.Background(), &Cfg{}, []Message{{Role: "user", Content: strp("x")}}) - if err2 == nil { - t.Errorf("expected error when first message is not system") + if err2 == nil || !strings.Contains(err2.Error(), "no system message") { + t.Errorf("expected error when first message is not system, got %v", err2) } if len(msgs2) != 1 { t.Errorf("expected original messages returned, got %d", len(msgs2)) } + msgs3, _, err3 := compact(context.Background(), &Cfg{}, []Message{{Role: "system", Content: strp("sys")}}) + if err3 == nil || !strings.Contains(err3.Error(), "nothing to compact") { + t.Errorf("expected nothing to compact error, got %v", err3) + } + if len(msgs3) != 1 { + t.Errorf("expected original messages returned, got %d", len(msgs3)) + } } // ---------- history ---------- @@ -943,7 +950,7 @@ func TestLLMNonStreamingAndHeaders(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = false cfg.APIKey = "secret" - m, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, TOOLS) + m, _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, TOOLS) if err != nil { t.Fatalf("llm: %v", err) } @@ -973,7 +980,7 @@ func TestLLMNoAuthHeaderWhenNoKey(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = false cfg.APIKey = "-" - if _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil); err != nil { + if _, _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil); err != nil { t.Fatalf("llm: %v", err) } if gotAuth != "" { @@ -994,7 +1001,7 @@ func TestLLMStreaming(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = true cfg.APIKey = "-" - m, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) + m, _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) if err != nil { t.Fatalf("llm: %v", err) } @@ -1018,7 +1025,7 @@ func TestLLM4xxReturnsImmediately(t *testing.T) { cfg.Stream = false cfg.APIKey = "-" start := time.Now() - _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) + _, _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) if err == nil || !strings.Contains(err.Error(), "Invalid assistant message") { t.Fatalf("expected 400 error, got %v", err) } @@ -1044,7 +1051,7 @@ func TestLLMRetriesOn5xxThenSucceeds(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = false cfg.APIKey = "-" - m, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) + m, _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) if err != nil { t.Fatalf("llm after retries: %v", err) } @@ -1066,7 +1073,7 @@ func TestLLMEmptyChoices(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = false cfg.APIKey = "-" - if _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil); err == nil || !strings.Contains(err.Error(), "empty choices") { + if _, _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil); err == nil || !strings.Contains(err.Error(), "empty choices") { t.Errorf("expected empty choices error, got %v", err) } } @@ -1093,7 +1100,7 @@ func TestALToolLoop(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = false cfg.APIKey = "-" - msgs, err := AL(context.Background(), &cfg, []Message{{Role: "system", Content: strp("sys")}, {Role: "user", Content: strp("run")}}, "sys", 0) + msgs, _, err := AL(context.Background(), &cfg, []Message{{Role: "system", Content: strp("sys")}, {Role: "user", Content: strp("run")}}, "sys", 0) if err != nil { t.Fatalf("AL: %v", err) } @@ -1155,7 +1162,7 @@ func TestALRunSubagent(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = false cfg.APIKey = "-" - msgs, err := AL(context.Background(), &cfg, []Message{{Role: "system", Content: strp("sys")}, {Role: "user", Content: strp("parent task")}}, "sys", 0) + msgs, _, err := AL(context.Background(), &cfg, []Message{{Role: "system", Content: strp("sys")}, {Role: "user", Content: strp("parent task")}}, "sys", 0) if err != nil { t.Fatalf("AL: %v", err) } @@ -1198,7 +1205,7 @@ func TestALSubagentDepthLimit(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = false cfg.APIKey = "-" - msgs, err := AL(context.Background(), &cfg, []Message{{Role: "user", Content: strp("go")}}, "sys", MAX_DEPTH) + msgs, _, err := AL(context.Background(), &cfg, []Message{{Role: "user", Content: strp("go")}}, "sys", MAX_DEPTH) if err != nil { t.Fatalf("AL: %v", err) } @@ -1245,7 +1252,7 @@ func TestALStripsInvalidAssistantAndRetries(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = false cfg.APIKey = "-" - msgs, err := AL(context.Background(), &cfg, []Message{{Role: "system", Content: strp("sys")}, {Role: "user", Content: strp("go")}}, "sys", 0) + msgs, _, err := AL(context.Background(), &cfg, []Message{{Role: "system", Content: strp("sys")}, {Role: "user", Content: strp("go")}}, "sys", 0) if err != nil { t.Fatalf("AL: %v", err) } @@ -1269,7 +1276,8 @@ func TestSummarizeHappyPath(t *testing.T) { cfg.Endpoint = srv.URL cfg.Stream = false cfg.APIKey = "-" - s, err := summarize(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hello world")}}) + orig := []Message{{Role: "system", Content: strp("sys")}, {Role: "user", Content: strp("hello world")}} + s, err := summarize(context.Background(), &cfg, orig) if err != nil { t.Fatalf("summarize: %v", err) } @@ -1279,7 +1287,13 @@ func TestSummarizeHappyPath(t *testing.T) { } func TestCompactHappyPath(t *testing.T) { + var receivedMessages []Message srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req struct { + Messages []Message `json:"messages"` + } + json.NewDecoder(r.Body).Decode(&req) + receivedMessages = req.Messages w.Write([]byte(`{"choices":[{"message":{"role":"assistant","content":"the summary"}}]}`)) })) defer srv.Close() @@ -1302,6 +1316,19 @@ func TestCompactHappyPath(t *testing.T) { if msgs[1].Role != "user" || msgs[1].Content == nil || !strings.Contains(*msgs[1].Content, "the summary") { t.Errorf("continuation message = %+v", msgs[1]) } + // Verify request sent to LLM contains the original conversation prefix plus compaction prompt + if len(receivedMessages) != 3 { + t.Fatalf("expected 3 messages sent to LLM, got %d", len(receivedMessages)) + } + if receivedMessages[0].Role != "system" || *receivedMessages[0].Content != "sys" { + t.Errorf("message 0 mismatch: %+v", receivedMessages[0]) + } + if receivedMessages[1].Role != "user" || *receivedMessages[1].Content != "hello world" { + t.Errorf("message 1 mismatch: %+v", receivedMessages[1]) + } + if receivedMessages[2].Role != "user" || !strings.Contains(*receivedMessages[2].Content, "compaction engine") { + t.Errorf("message 2 mismatch (expected compaction prompt): %+v", receivedMessages[2]) + } } // ---------- setCfg and LLM parameter forwarding ---------- @@ -1352,7 +1379,7 @@ func TestLLMForwardsRelevantParameters(t *testing.T) { }, "\n")) cfg := getCfg(cfgFile) - _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) + _, _, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) if err != nil { t.Fatalf("llm: %v", err) } @@ -1644,7 +1671,7 @@ func TestALContextCancellation(t *testing.T) { origMsgs := []Message{{Role: "system", Content: strp("sys")}, {Role: "user", Content: strp("hello")}} inputMsgs := append([]Message{}, origMsgs...) - msgs, err := AL(ctx, &cfg, inputMsgs, "sys", 0) + msgs, _, err := AL(ctx, &cfg, inputMsgs, "sys", 0) if err == nil { t.Fatalf("expected context cancellation error, got nil") } @@ -1682,7 +1709,7 @@ func TestLLMContextCancellation(t *testing.T) { cancel() }() - _, err := llm(ctx, &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) + _, _, err := llm(ctx, &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) if err == nil { t.Fatalf("expected context cancellation error, got nil") } @@ -1691,6 +1718,139 @@ func TestLLMContextCancellation(t *testing.T) { } } +func TestTokenCounterAndUsageNonStreaming(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Write([]byte(`{ + "choices":[{"message":{"role":"assistant","content":"hello world"}}], + "usage":{ + "prompt_tokens": 120, + "completion_tokens": 30, + "total_tokens": 150, + "prompt_tokens_details": {"cached_tokens": 80} + } + }`)) + })) + defer srv.Close() + + cfg := defCfg + cfg.Endpoint = srv.URL + cfg.Stream = false + cfg.APIKey = "-" + + m, u, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) + if err != nil { + t.Fatalf("llm: %v", err) + } + if m.Content == nil || *m.Content != "hello world" { + t.Errorf("unexpected content: %+v", m.Content) + } + if u.PromptTokens != 120 || u.CompletionTokens != 30 || u.TotalTokens != 150 { + t.Errorf("unexpected usage: %+v", u) + } + if u.Cached() != 80 { + t.Errorf("expected cached tokens 80, got %d", u.Cached()) + } +} + +func TestTokenCounterAndUsageStreaming(t *testing.T) { + sseData := strings.Join([]string{ + `data: {"choices":[{"delta":{"content":"streaming "}}]}`, + `data: {"choices":[{"delta":{"content":"response"}}]}`, + `data: {"choices":[],"usage":{"prompt_tokens":250,"completion_tokens":45,"total_tokens":295,"cached_tokens":100}}`, + `data: [DONE]`, + }, "\n") + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + w.Write([]byte(sseData)) + })) + defer srv.Close() + + cfg := defCfg + cfg.Endpoint = srv.URL + cfg.Stream = true + cfg.APIKey = "-" + + m, u, err := llm(context.Background(), &cfg, []Message{{Role: "user", Content: strp("hi")}}, nil) + if err != nil { + t.Fatalf("llm streaming: %v", err) + } + if m.Content == nil || *m.Content != "streaming response" { + t.Errorf("unexpected content: %+v", m.Content) + } + if u.PromptTokens != 250 || u.CompletionTokens != 45 || u.TotalTokens != 295 { + t.Errorf("unexpected usage: %+v", u) + } + if u.Cached() != 100 { + t.Errorf("expected cached tokens 100, got %d", u.Cached()) + } +} + +func TestFormatUsage(t *testing.T) { + // With cached tokens + u1 := Usage{PromptTokens: 1000, CompletionTokens: 200, TotalTokens: 1200, CachedTokens: 800} + s1 := formatUsage(u1, 200000) + if s1 != "[tokens: 1000 prompt (800 cached, 200 uncached) + 200 completion | context: 1000/200000 (0.5%)]" { + t.Errorf("formatUsage u1 = %q", s1) + } + + // Without cached tokens + u2 := Usage{PromptTokens: 120000, CompletionTokens: 500, TotalTokens: 120500} + s2 := formatUsage(u2, 200000) + if s2 != "[tokens: 120000 prompt + 500 completion | context: 120000/200000 (60.0%)]" { + t.Errorf("formatUsage u2 = %q", s2) + } +} + +func TestFetchContextWindowFromModelsAPI(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/models" { + http.NotFound(w, r) + return + } + w.Write([]byte(`{ + "data": [ + {"id": "other-model", "context_window": 32000}, + {"id": "my-target-model", "max_context_length": 131072} + ] + }`)) + })) + defer srv.Close() + + cfg := defCfg + cfg.Endpoint = srv.URL + cfg.Model = "my-target-model" + cfg.APIKey = "-" + + cw := fetchContextWindow(&cfg) + if cw != 131072 { + t.Errorf("expected context window 131072 from /models API, got %d", cw) + } +} + +func TestFetchContextWindowFallbackConfig(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Error(w, "server error", 500) + })) + defer srv.Close() + + // Case 1: Config specifies context_window + cfg1 := defCfg + cfg1.Endpoint = srv.URL + cfg1.Raw = map[string]string{"context_window": "65536"} + if cw := fetchContextWindow(&cfg1); cw != 65536 { + t.Errorf("expected fallback to raw context_window 65536, got %d", cw) + } + + // Case 2: Config does not specify context_window -> default 200000 + cfg2 := defCfg + cfg2.Endpoint = srv.URL + cfg2.Raw = map[string]string{} + if cw := fetchContextWindow(&cfg2); cw != 200000 { + t.Errorf("expected default 200000, got %d", cw) + } +} +