diff --git a/README.md b/README.md index 6d770677..d2f9957d 100644 --- a/README.md +++ b/README.md @@ -44,6 +44,7 @@ by another machine on your network, or by a provider. - **Third-Party Providers** — integrate OpenAI, DeepSeek, MiMo, Kimi, BigModel, Qianfan, MiniMax, OpenRouter, and any OpenAI-compatible API - **Coding Agents** — one-click config for Claude Code, Codex, Pi, OpenCode, and Open Code Review +- **Context compression** — optionally shrinks the tool output inside agent requests before any model sees it (Settings → Context compression). Safe mode removes only redundancy and keeps every distinct line; aggressive mode also samples long logs, arrays and search results. File reads and source code are never changed. On captured Claude Code / Codex traffic tool output drops 6% (safe) to 17% (aggressive), search results 24–56%; savings per request show in Observability - **AI Applications** — one-click setup for Claude Code, OpenCode, Open Code Review, Codex, Codex App, ZCode, Pi, OpenClaw, CSGClaw, Dify, and AnythingLLM ### Dataset Support diff --git a/internal/config/config.go b/internal/config/config.go index dac92126..068103a0 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -150,6 +150,14 @@ type InferenceConfig struct { // global default decide how many slots a load gets. LlamaNumParallel int `json:"llama_num_parallel,omitempty"` + // ContextCompression is how far the gateway shrinks the tool output + // inside agent requests before a model sees them: ContextCompressionOff + // (the default when empty), ContextCompressionSafe or + // ContextCompressionAggressive. It is one process-wide policy because the + // same agent traffic reaches local, cluster and provider models alike, and + // no existing setting describes request rewriting. + ContextCompression string `json:"context_compression,omitempty"` + // Models holds the per-model load options keyed by model ID. They live in // the app config rather than in the model directory so that re-downloading // a model keeps its settings. @@ -507,6 +515,34 @@ func Load() (*Config, error) { return globalConfig, loadErr } +const ( + ContextCompressionOff = "off" + ContextCompressionSafe = "safe" + ContextCompressionAggressive = "aggressive" +) + +// NormalizeContextCompression maps a stored or requested mode to one of the +// ContextCompression constants; anything unknown means off. +func NormalizeContextCompression(value string) string { + switch strings.ToLower(strings.TrimSpace(value)) { + case ContextCompressionSafe: + return ContextCompressionSafe + case ContextCompressionAggressive: + return ContextCompressionAggressive + default: + return ContextCompressionOff + } +} + +// IsContextCompressionMode reports whether value names a mode. +func IsContextCompressionMode(value string) bool { + switch strings.ToLower(strings.TrimSpace(value)) { + case ContextCompressionOff, ContextCompressionSafe, ContextCompressionAggressive: + return true + } + return false +} + func NormalizeMarketplaceModelSource(value string) string { switch strings.ToLower(strings.TrimSpace(value)) { case "huggingface": diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 440a7025..2f8d8421 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -323,6 +323,49 @@ func TestInferenceConfigPersistsAcrossSaveAndLoad(t *testing.T) { } } +func TestContextCompressionPersistsAndDefaultsOff(t *testing.T) { + var legacy Config + if err := json.Unmarshal([]byte(`{"inference":{"llama_num_parallel":2}}`), &legacy); err != nil { + t.Fatal(err) + } + if got := NormalizeContextCompression(legacy.Inference.ContextCompression); got != ContextCompressionOff { + t.Fatalf("legacy config mode = %q, want off", got) + } + + home := t.TempDir() + t.Setenv("HOME", home) + t.Setenv("USERPROFILE", home) + clearCloudServiceEnv(t) + Reset() + t.Cleanup(Reset) + + cfg, err := Load() + if err != nil { + t.Fatal(err) + } + cfg.Inference.ContextCompression = ContextCompressionAggressive + if err := Save(cfg); err != nil { + t.Fatal(err) + } + Reset() + loaded, err := Load() + if err != nil { + t.Fatal(err) + } + if loaded.Inference.ContextCompression != ContextCompressionAggressive { + t.Fatalf("mode after reload = %q", loaded.Inference.ContextCompression) + } + + for value, want := range map[string]string{"": "off", "OFF": "off", " Safe ": "safe", "aggressive": "aggressive", "max": "off"} { + if got := NormalizeContextCompression(value); got != want { + t.Errorf("NormalizeContextCompression(%q) = %q, want %q", value, got, want) + } + } + if IsContextCompressionMode("max") || !IsContextCompressionMode("Safe") { + t.Fatal("IsContextCompressionMode accepted or rejected the wrong value") + } +} + func TestMarketplaceModelSourcePersistsAndDefaults(t *testing.T) { home := t.TempDir() t.Setenv("HOME", home) diff --git a/internal/ctxcompress/corpus_test.go b/internal/ctxcompress/corpus_test.go new file mode 100644 index 00000000..8f62a293 --- /dev/null +++ b/internal/ctxcompress/corpus_test.go @@ -0,0 +1,84 @@ +package ctxcompress + +import ( + "bufio" + "encoding/json" + "os" + "path/filepath" + "testing" +) + +// TestCorpus replays captured traffic through both modes and writes the +// results next to the corpus for token counting. It runs only when +// CTXCOMPRESS_CORPUS_DIR names a directory holding corpus_blocks.jsonl +// ({"tool","text"} per line) and corpus_requests.jsonl ({"id","protocol", +// "body"} per line), such as one exported from the observability store. +func TestCorpus(t *testing.T) { + dir := os.Getenv("CTXCOMPRESS_CORPUS_DIR") + if dir == "" { + t.Skip("CTXCOMPRESS_CORPUS_DIR not set") + } + modes := map[string]Options{"safe": {}, "aggressive": {Sample: true}} + + blocksOut, err := os.Create(filepath.Join(dir, "result_blocks.jsonl")) + if err != nil { + t.Fatal(err) + } + defer blocksOut.Close() + encoder := json.NewEncoder(blocksOut) + forEachLine(t, filepath.Join(dir, "corpus_blocks.jsonl"), func(line []byte) { + var block struct{ Tool, Text string } + if err := json.Unmarshal(line, &block); err != nil { + t.Fatal(err) + } + for mode, opts := range modes { + out, kind := Text(block.Text, block.Tool, opts) + if again, _ := Text(out, block.Tool, opts); again != out { + t.Errorf("%s: not idempotent for a %s block from %s (%d → %d → %d bytes)", mode, kind, block.Tool, len(block.Text), len(out), len(again)) + } + _ = encoder.Encode(map[string]any{"mode": mode, "tool": block.Tool, "kind": kind, "before": block.Text, "after": out}) + } + }) + + requestsOut, err := os.Create(filepath.Join(dir, "result_requests.jsonl")) + if err != nil { + t.Fatal(err) + } + defer requestsOut.Close() + encoder = json.NewEncoder(requestsOut) + forEachLine(t, filepath.Join(dir, "corpus_requests.jsonl"), func(line []byte) { + var request struct { + ID string + Protocol Protocol + Body string + } + if err := json.Unmarshal(line, &request); err != nil { + t.Fatal(err) + } + for mode, opts := range modes { + result, err := CompressRequest(request.Protocol, []byte(request.Body), opts) + if err != nil { + t.Errorf("%s: %v", request.ID, err) + continue + } + _ = encoder.Encode(map[string]any{"mode": mode, "id": request.ID, "protocol": request.Protocol, "after": string(result.Body), "stats": result.Stats}) + } + }) +} + +func forEachLine(t *testing.T, path string, fn func([]byte)) { + t.Helper() + file, err := os.Open(path) + if err != nil { + t.Fatal(err) + } + defer file.Close() + scanner := bufio.NewScanner(file) + scanner.Buffer(make([]byte, 0, 1<<20), 64<<20) + for scanner.Scan() { + fn(scanner.Bytes()) + } + if err := scanner.Err(); err != nil { + t.Fatal(err) + } +} diff --git a/internal/ctxcompress/json.go b/internal/ctxcompress/json.go new file mode 100644 index 00000000..756be12c --- /dev/null +++ b/internal/ctxcompress/json.go @@ -0,0 +1,327 @@ +package ctxcompress + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "sort" + "strings" +) + +// jsonValue is a decoded JSON value that remembers object key order, so a +// rewrite changes only what it means to change. +type jsonValue struct { + kind byte // 'o' object, 'a' array, 's' string, 'n' number, 'b' bool, 'z' null + keys []string + fields []*jsonValue + items []*jsonValue + str string + raw string // number literal, or "true"/"false"/"null" +} + +func parseJSON(text string) (*jsonValue, error) { + decoder := json.NewDecoder(strings.NewReader(text)) + decoder.UseNumber() + value, err := decodeJSONValue(decoder) + if err != nil { + return nil, err + } + if _, err := decoder.Token(); err != io.EOF { + return nil, fmt.Errorf("trailing data after JSON value") + } + return value, nil +} + +func decodeJSONValue(decoder *json.Decoder) (*jsonValue, error) { + token, err := decoder.Token() + if err != nil { + return nil, err + } + switch t := token.(type) { + case json.Delim: + switch t { + case '{': + value := &jsonValue{kind: 'o'} + for decoder.More() { + keyToken, err := decoder.Token() + if err != nil { + return nil, err + } + key, _ := keyToken.(string) + field, err := decodeJSONValue(decoder) + if err != nil { + return nil, err + } + value.keys = append(value.keys, key) + value.fields = append(value.fields, field) + } + _, err := decoder.Token() + return value, err + case '[': + value := &jsonValue{kind: 'a'} + for decoder.More() { + item, err := decodeJSONValue(decoder) + if err != nil { + return nil, err + } + value.items = append(value.items, item) + } + _, err := decoder.Token() + return value, err + } + case string: + return &jsonValue{kind: 's', str: t}, nil + case json.Number: + return &jsonValue{kind: 'n', raw: t.String()}, nil + case bool: + return &jsonValue{kind: 'b', raw: fmt.Sprint(t)}, nil + case nil: + return &jsonValue{kind: 'z', raw: "null"}, nil + } + return nil, fmt.Errorf("unexpected JSON token %v", token) +} + +func (v *jsonValue) field(key string) *jsonValue { + for i, k := range v.keys { + if k == key { + return v.fields[i] + } + } + return nil +} + +func (v *jsonValue) encode(buf *bytes.Buffer) { + switch v.kind { + case 'o': + buf.WriteByte('{') + for i, key := range v.keys { + if i > 0 { + buf.WriteByte(',') + } + writeJSONString(buf, key) + buf.WriteByte(':') + v.fields[i].encode(buf) + } + buf.WriteByte('}') + case 'a': + buf.WriteByte('[') + for i, item := range v.items { + if i > 0 { + buf.WriteByte(',') + } + item.encode(buf) + } + buf.WriteByte(']') + case 's': + writeJSONString(buf, v.str) + default: + buf.WriteString(v.raw) + } +} + +func (v *jsonValue) compact() string { + var buf bytes.Buffer + v.encode(&buf) + return buf.String() +} + +func writeJSONString(buf *bytes.Buffer, s string) { + encoder := json.NewEncoder(buf) + encoder.SetEscapeHTML(false) + _ = encoder.Encode(s) + buf.Truncate(buf.Len() - 1) // Encode appends a newline. +} + +// compressJSON handles a tool result that is one JSON document. It always +// minifies. A top-level array of records is rendered as a table with the +// field names written once, Headroom's CSV-with-schema form; with sampling +// enabled, long arrays anywhere in the document keep only a representative +// subset first. +func compressJSON(text string, opts Options) (string, bool) { + trimmed := strings.TrimSpace(text) + if trimmed == "" || (trimmed[0] != '{' && trimmed[0] != '[') { + return "", false + } + value, err := parseJSON(trimmed) + if err != nil { + return "", false + } + if value.kind != 'a' { + if opts.Sample { + sampleJSON(value) + } + return value.compact(), true + } + + // A top-level array keeps its omitted-items note outside the table. + items, total := value.items, len(value.items) + if opts.Sample { + for _, item := range items { + sampleJSON(item) + } + items, total = sampleArray(items) + } + note := "" + if len(items) < total { + note = fmt.Sprintf(omittedItemsFormat, total-len(items), total) + } + if table, ok := renderTable(items); ok { + if note != "" { + table += "\n" + note + } + return table, true + } + if note != "" { + items = append(items, &jsonValue{kind: 's', str: note}) + } + return (&jsonValue{kind: 'a', items: items}).compact(), true +} + +// isTableHeader recognises a rendered table, so that compressing it a +// second time leaves it alone. +func isTableHeader(line string) bool { + if !strings.HasPrefix(line, "[") { + return false + } + end := strings.Index(line, "]{") + if end < 2 || !strings.HasSuffix(line, "}") { + return false + } + for _, c := range line[1:end] { + if c < '0' || c > '9' { + return false + } + } + return true +} + +const ( + tableMinRows = 3 + tableCoreFrequency = 0.8 + tableCoreRatio = 0.6 +) + +// renderTable writes an array of objects as +// +// [N]{name:type,name:type?,...} +// cell,cell,... +// +// with one row per item. It declines when the items are not objects or do +// not share most of their keys, since a sparse table saves nothing. +func renderTable(items []*jsonValue) (string, bool) { + if len(items) < tableMinRows { + return "", false + } + frequency := map[string]int{} + var order []string + for _, item := range items { + if item.kind != 'o' { + return "", false + } + for _, key := range item.keys { + if key == "" || strings.ContainsAny(key, ",:{}[]\"\n\r") { + return "", false + } + if frequency[key] == 0 { + order = append(order, key) + } + frequency[key]++ + } + } + if len(order) == 0 { + return "", false + } + threshold := int(float64(len(items))*tableCoreFrequency + 0.999) + core := 0 + for _, key := range order { + if frequency[key] >= threshold { + core++ + } + } + if float64(core) < float64(len(order))*tableCoreRatio { + return "", false + } + // Frequent columns first; among equals, the order keys first appeared in. + columns := append([]string(nil), order...) + sort.SliceStable(columns, func(i, j int) bool { return frequency[columns[i]] > frequency[columns[j]] }) + + var buf bytes.Buffer + fmt.Fprintf(&buf, "[%d]{", len(items)) + for i, column := range columns { + if i > 0 { + buf.WriteByte(',') + } + buf.WriteString(column) + buf.WriteByte(':') + buf.WriteString(columnType(items, column)) + if columnNullable(items, column) { + buf.WriteByte('?') + } + } + buf.WriteString("}") + for _, item := range items { + buf.WriteByte('\n') + for i, column := range columns { + if i > 0 { + buf.WriteByte(',') + } + if cell := item.field(column); cell != nil { + buf.WriteString(tableCell(cell)) + } + } + } + return buf.String(), true +} + +func columnType(items []*jsonValue, column string) string { + kind := byte(0) + for _, item := range items { + cell := item.field(column) + if cell == nil || cell.kind == 'z' { + continue + } + if kind == 0 { + kind = cell.kind + } else if kind != cell.kind { + return "json" + } + } + switch kind { + case 'n': + return "number" + case 'b': + return "bool" + case 'o', 'a': + return "json" + default: + return "string" + } +} + +func columnNullable(items []*jsonValue, column string) bool { + for _, item := range items { + if cell := item.field(column); cell == nil || cell.kind == 'z' { + return true + } + } + return false +} + +// tableCell renders one value CSV-style: bare when unambiguous, quoted with +// doubled quotes otherwise. Nested objects and arrays stay compact JSON. +func tableCell(cell *jsonValue) string { + var text string + switch cell.kind { + case 's': + text = cell.str + if text != "" && text != "null" && !strings.ContainsAny(text, ",\"\n\r") { + return text + } + case 'o', 'a': + text = cell.compact() + default: + return cell.raw + } + return `"` + strings.ReplaceAll(text, `"`, `""`) + `"` +} diff --git a/internal/ctxcompress/lines.go b/internal/ctxcompress/lines.go new file mode 100644 index 00000000..c4e6a59f --- /dev/null +++ b/internal/ctxcompress/lines.go @@ -0,0 +1,370 @@ +package ctxcompress + +import ( + "fmt" + "regexp" + "strconv" + "strings" + "unicode/utf8" +) + +// ansiEscape matches terminal colour and cursor sequences (CSI and OSC), +// which cost tokens and mean nothing to a model. +var ansiEscape = regexp.MustCompile(`\x1b\[[0-9;?]*[ -/]*[@-~]|\x1b\][^\x07\x1b]*(?:\x07|\x1b\\)`) + +// normalizeLines splits command output into lines the way a terminal would +// show them: escape sequences dropped, a carriage-return progress bar reduced +// to its final state, and trailing blanks removed. +func normalizeLines(text string) []string { + if strings.Contains(text, "\x1b") { + text = ansiEscape.ReplaceAllString(text, "") + } + lines := strings.Split(text, "\n") + for i, line := range lines { + if strings.Contains(line, "\r") { + line = strings.TrimRight(line, "\r") + if idx := strings.LastIndex(line, "\r"); idx >= 0 { + line = line[idx+1:] + } + } + lines[i] = strings.TrimRight(line, " \t") + } + return lines +} + +// trimLongLines cuts every line longer than maxBytes down to its start and +// end. Such lines are minified bundles, base64 or binary matches that a model +// cannot use in full anyway. +func trimLongLines(lines []string, maxBytes int) []string { + out := lines + copied := false + for i, line := range lines { + if len(line) <= maxBytes { + continue + } + if !copied { + out = append([]string(nil), lines...) + copied = true + } + out[i] = trimLine(line, maxBytes) + } + return out +} + +func trimLine(line string, maxBytes int) string { + head := cutRunes(line, maxBytes/2, false) + tail := cutRunes(line, maxBytes/4, true) + omitted := len(line) - len(head) - len(tail) + return fmt.Sprintf("%s …[%d bytes omitted]… %s", head, omitted, tail) +} + +// cutRunes returns at most n bytes from the start (or, with fromEnd, the end) +// of s without splitting a UTF-8 sequence. +func cutRunes(s string, n int, fromEnd bool) string { + if n >= len(s) { + return s + } + if !fromEnd { + for n > 0 && !utf8.RuneStart(s[n]) { + n-- + } + return s[:n] + } + start := len(s) - n + for start < len(s) && !utf8.RuneStart(s[start]) { + start++ + } + return s[start:] +} + +// foldRepeats replaces three or more identical consecutive lines with one +// copy and a count. +func foldRepeats(lines []string) []string { + out := make([]string, 0, len(lines)) + for i := 0; i < len(lines); { + j := i + 1 + for j < len(lines) && lines[j] == lines[i] { + j++ + } + run := j - i + if run >= 3 && strings.TrimSpace(lines[i]) != "" { + out = append(out, lines[i], fmt.Sprintf("[… previous line repeated %d more times]", run-1)) + } else { + out = append(out, lines[i:j]...) + } + i = j + } + return out +} + +// foldBlankRuns keeps at most one empty line in a row. +func foldBlankRuns(lines []string) []string { + out := make([]string, 0, len(lines)) + for i, line := range lines { + if line == "" && i > 0 && lines[i-1] == "" { + continue + } + out = append(out, line) + } + return out +} + +const ( + similarRunMin = 6 + similarRunKeep = 2 +) + +// foldSimilar shortens runs of consecutive lines that differ only in their +// numbers — progress output, download counters, polling loops — to the first +// and last two with a count in between. It drops the numbers of the folded +// lines, so only aggressive mode uses it. +func foldSimilar(lines []string) []string { + out := make([]string, 0, len(lines)) + for i := 0; i < len(lines); { + template := lineTemplate(lines[i]) + j := i + 1 + for j < len(lines) && template != "" && lineTemplate(lines[j]) == template { + j++ + } + run := j - i + if run >= similarRunMin { + out = append(out, lines[i:i+similarRunKeep]...) + out = append(out, fmt.Sprintf("[… %d similar lines omitted]", run-2*similarRunKeep)) + out = append(out, lines[j-similarRunKeep:j]...) + } else { + out = append(out, lines[i:j]...) + } + i = j + } + return out +} + +var ( + templateUUID = regexp.MustCompile(`[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}`) + templateHex = regexp.MustCompile(`\b(?:0x)?[0-9a-fA-F]{8,}\b`) + templateNumber = regexp.MustCompile(`\d+(?:[.,:]\d+)*`) +) + +// lineTemplate is the line with its variable parts masked. Blank lines, +// lines without any number and lines reporting an error have no template, so +// only numeric churn folds and every error line survives on its own. +func lineTemplate(line string) string { + if strings.TrimSpace(line) == "" || !strings.ContainsAny(line, "0123456789") || isErrorLine(line) { + return "" + } + masked := templateUUID.ReplaceAllString(line, "U") + masked = templateHex.ReplaceAllString(masked, "H") + masked = templateNumber.ReplaceAllString(masked, "0") + if !strings.ContainsFunc(masked, func(r rune) bool { return r != '0' && r != 'U' && r != 'H' && r != ' ' }) { + // A line that is nothing but numbers is data, not churn. + return "" + } + return masked +} + +// errorKeywords mark the lines of command output a model most needs: they +// survive head/tail sampling wherever they sit. +var errorKeywords = []string{ + "error", "fail", "fatal", "panic", "exception", "traceback", "warn", + "denied", "refused", "timeout", "timed out", "cannot", "can't", "unable", + "not found", "no such", "undefined", "unexpected", "invalid", "segfault", + "assert", "abort", "critical", "missing", "conflict", "rejected", "错误", "失败", +} + +func isErrorLine(line string) bool { + lower := strings.ToLower(line) + for _, keyword := range errorKeywords { + if strings.Contains(lower, keyword) { + return true + } + } + return false +} + +// sampleLines keeps the head and tail of output longer than maxLines, plus +// the lines around every error in between, within the same line budget. +func sampleLines(lines []string, maxLines int) []string { + // Prefix headers are added after sampling, so they do not count. + if len(lines) <= maxLines || len(lines)-countFactoredHeaders(lines) <= maxLines { + return lines + } + head := maxLines / 4 + tail := maxLines * 3 / 8 + budget := maxLines - head - tail - 1 + + middle := lines[head : len(lines)-tail] + keep := make([]bool, len(middle)) + for i, line := range middle { + if !isErrorLine(line) { + continue + } + for k := max(0, i-1); k <= min(len(middle)-1, i+1); k++ { + keep[k] = true + } + } + + out := append([]string(nil), lines[:head]...) + omitted := 0 + flush := func() { + if omitted > 0 { + out = append(out, fmt.Sprintf("[… %d lines omitted]", omitted)) + omitted = 0 + budget-- + } + } + for i, line := range middle { + // Each kept line costs one slot and may need a marker before it. + if keep[i] && budget >= 2 { + flush() + out = append(out, line) + budget-- + continue + } + omitted++ + } + flush() + return append(out, lines[len(lines)-tail:]...) +} + +func compressLog(lines []string, opts Options) string { + lines = trimLongLines(lines, opts.MaxLineBytes) + lines = foldBlankRuns(lines) + lines = foldRepeats(lines) + if opts.Sample { + // Folding lines that differ only in their numbers loses those + // numbers, so it is a sampling step, not a redundancy one. + lines = foldSimilar(lines) + lines = sampleLines(lines, opts.MaxLines) + lines = factorPrefixes(lines) + } + return joinLines(lines) +} + +// compressProse handles fetched pages and other free text, where long lines +// are paragraphs rather than noise, so only far longer ones are cut. +func compressProse(lines []string, opts Options) string { + lines = trimLongLines(lines, opts.MaxLineBytes*4) + lines = foldBlankRuns(lines) + lines = foldRepeats(lines) + return joinLines(lines) +} + +// looksLikeDiff spots unified diffs and git patches. +func looksLikeDiff(lines []string) bool { + hunks := 0 + for _, line := range lines { + if strings.HasPrefix(line, "diff --git ") { + return true + } + if strings.HasPrefix(line, "@@ -") && strings.Contains(line, " @@") { + hunks++ + } + } + return hunks > 0 && containsPrefixPair(lines, "--- ", "+++ ") +} + +func containsPrefixPair(lines []string, first, second string) bool { + for i := 0; i+1 < len(lines); i++ { + if strings.HasPrefix(lines[i], first) && strings.HasPrefix(lines[i+1], second) { + return true + } + } + return false +} + +// looksLikeLog separates line-oriented command output from prose: many +// lines, and short ones on average. +func looksLikeLog(lines []string) bool { + nonEmpty, total := 0, 0 + for _, line := range lines { + if line == "" { + continue + } + nonEmpty++ + total += len(line) + } + return nonEmpty >= 8 && total/nonEmpty <= 240 +} + +const ( + prefixRunMin = 5 + prefixMinBytes = 16 +) + +const prefixHeaderFormat = "[the next %d lines each start with: %s]" + +var prefixHeader = regexp.MustCompile(`^\[the next (\d+) lines each start with: ".*"\]$`) + +// factorPrefixes writes a prefix shared by a run of lines — the job and +// step columns of a CI log, a deep directory in a file list — once above +// the run instead of on every line. The prefix always ends at a space, tab +// or slash so the remainder reads naturally. +func factorPrefixes(lines []string) []string { + out := make([]string, 0, len(lines)) + for i := 0; i < len(lines); { + // A run factored by an earlier pass stays as it is. + if m := prefixHeader.FindStringSubmatch(lines[i]); m != nil { + n, _ := strconv.Atoi(m[1]) + end := min(len(lines), i+1+n) + out = append(out, lines[i:end]...) + i = end + continue + } + prefix, end := "", i+1 + if lines[i] != "" { + common := lines[i] + for j := i + 1; j < len(lines); j++ { + next := commonPrefix(common, lines[j]) + if cutAtSeparator(next) == "" || len(cutAtSeparator(next)) < prefixMinBytes { + break + } + common, end = next, j+1 + } + prefix = cutAtSeparator(common) + } + if end-i < prefixRunMin || len(prefix) < prefixMinBytes { + out = append(out, lines[i]) + i++ + continue + } + out = append(out, fmt.Sprintf(prefixHeaderFormat, end-i, strconv.Quote(prefix))) + for _, line := range lines[i:end] { + out = append(out, line[len(prefix):]) + } + i = end + } + return out +} + +func hasFactoredPrefix(lines []string) bool { + return countFactoredHeaders(lines) > 0 +} + +func countFactoredHeaders(lines []string) int { + n := 0 + for _, line := range lines { + if strings.HasPrefix(line, "[the next ") && prefixHeader.MatchString(line) { + n++ + } + } + return n +} + +func commonPrefix(a, b string) string { + n := min(len(a), len(b)) + i := 0 + for i < n && a[i] == b[i] { + i++ + } + return a[:i] +} + +// cutAtSeparator shortens a common prefix to end just after its last space, +// tab or slash. +func cutAtSeparator(prefix string) string { + idx := strings.LastIndexAny(prefix, " \t/\\") + if idx < 0 { + return "" + } + return prefix[:idx+1] +} diff --git a/internal/ctxcompress/request.go b/internal/ctxcompress/request.go new file mode 100644 index 00000000..1048a6b9 --- /dev/null +++ b/internal/ctxcompress/request.go @@ -0,0 +1,219 @@ +package ctxcompress + +import ( + "bytes" + "encoding/json" + "fmt" +) + +// Protocol names the request body shape a gateway endpoint accepts. +type Protocol string + +const ( + ProtocolAnthropic Protocol = "anthropic" + ProtocolOpenAIChat Protocol = "openai" + ProtocolResponses Protocol = "responses" +) + +// Result describes what CompressRequest did to one request body. +type Result struct { + // Body is the rewritten request, or the original bytes when nothing was + // compressed. + Body []byte + // Changed reports whether Body differs from the input. + Changed bool + Stats Stats +} + +// CompressRequest rewrites the tool results inside a chat request body. +// +// Only tool output is touched: system prompts, tool definitions, user and +// assistant turns pass through unchanged. Every tool result in the history is +// compressed, not just the newest one, and the rewrite is a pure function of +// the result's own text. A client resends the original history on each turn, +// so the same result compresses to the same bytes every time and the prefix a +// provider or llama-server caches stays identical from one request to the +// next. +func CompressRequest(protocol Protocol, body []byte, opts Options) (Result, error) { + result := Result{Body: body} + decoder := json.NewDecoder(bytes.NewReader(body)) + decoder.UseNumber() + var root map[string]any + if err := decoder.Decode(&root); err != nil { + return result, fmt.Errorf("decoding request body: %w", err) + } + + walker := requestWalker{opts: opts.withDefaults()} + switch protocol { + case ProtocolAnthropic: + walker.anthropic(root) + case ProtocolOpenAIChat: + walker.openAIChat(root) + case ProtocolResponses: + walker.responses(root) + default: + return result, fmt.Errorf("unsupported protocol %q", protocol) + } + result.Stats = walker.stats + if !walker.changed { + return result, nil + } + + var out bytes.Buffer + encoder := json.NewEncoder(&out) + encoder.SetEscapeHTML(false) + if err := encoder.Encode(root); err != nil { + return result, fmt.Errorf("encoding request body: %w", err) + } + result.Body = bytes.TrimRight(out.Bytes(), "\n") + result.Changed = true + return result, nil +} + +type requestWalker struct { + opts Options + stats Stats + changed bool + // seen maps each tool result already walked to the call that produced + // it, for pointing a repeat back at the first copy. + seen map[string]string +} + +// duplicateFormat replaces a tool result identical to an earlier one. Only +// later copies are replaced, so the decision never depends on what comes +// after a result and the prefix before it stays cacheable. +const duplicateFormat = "[Identical to the earlier output of tool call %s.]" + +// dedupable reports whether a repeat of this kind may be replaced by a +// pointer to its first copy. File contents, source code and diffs never +// are: an agent editing a file works from the read it looks at last, and +// must find the text there rather than a reference to an older turn. +func dedupable(kind Kind) bool { + switch kind { + case KindJSON, KindSearch, KindLog, KindText: + return true + } + return false +} + +// compress runs one tool result's text through Text and records the outcome. +func (w *requestWalker) compress(text, toolName, callID string) (string, bool) { + out, kind := Text(text, toolName, w.opts) + if dedupable(kind) { + if first, ok := w.seen[text]; ok && first != "" && first != callID { + out = fmt.Sprintf(duplicateFormat, first) + w.stats.add(KindDuplicate, text, out) + w.changed = true + return out, true + } + if w.seen == nil { + w.seen = map[string]string{} + } + if _, ok := w.seen[text]; !ok { + w.seen[text] = callID + } + } + changed := out != text + w.stats.add(kind, text, out) + if changed { + w.changed = true + } + return out, changed +} + +// anthropic handles /v1/messages: tool_result blocks inside user messages, +// whose content is either a string or a list of text/image blocks. +func (w *requestWalker) anthropic(root map[string]any) { + messages, _ := root["messages"].([]any) + names := map[string]string{} + for _, rawMessage := range messages { + message, _ := rawMessage.(map[string]any) + blocks, _ := message["content"].([]any) + for _, rawBlock := range blocks { + block, _ := rawBlock.(map[string]any) + switch block["type"] { + case "tool_use", "server_tool_use": + if id, ok := block["id"].(string); ok { + names[id], _ = block["name"].(string) + } + case "tool_result": + if isError, _ := block["is_error"].(bool); isError { + continue + } + id, _ := block["tool_use_id"].(string) + w.rewriteContent(block, "content", names[id], id) + } + } + } +} + +// openAIChat handles /v1/chat/completions: messages with role "tool". +func (w *requestWalker) openAIChat(root map[string]any) { + messages, _ := root["messages"].([]any) + names := map[string]string{} + for _, rawMessage := range messages { + message, _ := rawMessage.(map[string]any) + if calls, ok := message["tool_calls"].([]any); ok { + for _, rawCall := range calls { + call, _ := rawCall.(map[string]any) + function, _ := call["function"].(map[string]any) + if id, ok := call["id"].(string); ok { + names[id], _ = function["name"].(string) + } + } + } + if message["role"] != "tool" { + continue + } + id, _ := message["tool_call_id"].(string) + name := names[id] + if name == "" { + name, _ = message["name"].(string) + } + w.rewriteContent(message, "content", name, id) + } +} + +// responses handles /v1/responses: function_call_output and +// custom_tool_call_output items in the input list. +func (w *requestWalker) responses(root map[string]any) { + items, _ := root["input"].([]any) + names := map[string]string{} + for _, rawItem := range items { + item, _ := rawItem.(map[string]any) + switch item["type"] { + case "function_call", "custom_tool_call", "local_shell_call": + if id, ok := item["call_id"].(string); ok { + names[id], _ = item["name"].(string) + } + case "function_call_output", "custom_tool_call_output": + id, _ := item["call_id"].(string) + w.rewriteContent(item, "output", names[id], id) + } + } +} + +// rewriteContent compresses holder[key] when it is a string, or each text +// part when it is a list of content parts. Non-text parts are left alone. +func (w *requestWalker) rewriteContent(holder map[string]any, key, toolName, callID string) { + switch value := holder[key].(type) { + case string: + if out, changed := w.compress(value, toolName, callID); changed { + holder[key] = out + } + case []any: + for _, rawPart := range value { + part, _ := rawPart.(map[string]any) + switch part["type"] { + case "text", "input_text", "output_text": + text, ok := part["text"].(string) + if !ok { + continue + } + if out, changed := w.compress(text, toolName, callID); changed { + part["text"] = out + } + } + } + } +} diff --git a/internal/ctxcompress/request_test.go b/internal/ctxcompress/request_test.go new file mode 100644 index 00000000..76ae95e6 --- /dev/null +++ b/internal/ctxcompress/request_test.go @@ -0,0 +1,199 @@ +package ctxcompress + +import ( + "encoding/json" + "fmt" + "strings" + "testing" +) + +func grepOutput(files, matches int) string { + var b strings.Builder + for f := 0; f < files; f++ { + for m := 1; m <= matches; m++ { + fmt.Fprintf(&b, "internal/pkg/file%d.go:%d:\tvalue := compute(%d)\n", f, m*3, m) + } + } + return b.String() +} + +func mustJSON(t *testing.T, v any) []byte { + t.Helper() + raw, err := json.Marshal(v) + if err != nil { + t.Fatal(err) + } + return raw +} + +func TestCompressRequestAnthropic(t *testing.T) { + search := grepOutput(3, 10) + fileText := strings.Repeat("same line\n", 100) + body := mustJSON(t, map[string]any{ + "model": "claude", + "system": []any{map[string]any{"type": "text", "text": search, "cache_control": map[string]any{"type": "ephemeral"}}}, + "messages": []any{ + map[string]any{"role": "user", "content": search}, + map[string]any{"role": "assistant", "content": []any{ + map[string]any{"type": "tool_use", "id": "toolu_1", "name": "Grep", "input": map[string]any{}}, + map[string]any{"type": "tool_use", "id": "toolu_2", "name": "Read", "input": map[string]any{}}, + map[string]any{"type": "tool_use", "id": "toolu_3", "name": "Bash", "input": map[string]any{}}, + map[string]any{"type": "tool_use", "id": "toolu_4", "name": "Bash", "input": map[string]any{}}, + }}, + map[string]any{"role": "user", "content": []any{ + map[string]any{"type": "tool_result", "tool_use_id": "toolu_1", "content": search, "cache_control": map[string]any{"type": "ephemeral"}}, + map[string]any{"type": "tool_result", "tool_use_id": "toolu_2", "content": []any{map[string]any{"type": "text", "text": fileText}}}, + map[string]any{"type": "tool_result", "tool_use_id": "toolu_3", "content": search, "is_error": true}, + map[string]any{"type": "tool_result", "tool_use_id": "toolu_4", "content": search}, + }}, + }, + }) + + result, err := CompressRequest(ProtocolAnthropic, body, Options{}) + if err != nil { + t.Fatal(err) + } + if !result.Changed || result.Stats.Blocks != 3 || result.Stats.Compressed != 2 { + t.Fatalf("stats = %+v", result.Stats) + } + if result.Stats.ByKind[KindDuplicate].Compressed != 1 || result.Stats.TokensSaved <= 0 { + t.Fatalf("by kind = %+v", result.Stats.ByKind) + } + + var got struct { + System []map[string]any `json:"system"` + Messages []struct { + Content any `json:"content"` + } `json:"messages"` + } + if err := json.Unmarshal(result.Body, &got); err != nil { + t.Fatal(err) + } + if got.System[0]["text"] != search || got.Messages[0].Content != search { + t.Fatal("system prompt or user text was rewritten") + } + results := got.Messages[2].Content.([]any) + grep := results[0].(map[string]any) + if text := grep["content"].(string); !strings.HasPrefix(text, "internal/pkg/file0.go\n 3:") { + t.Fatalf("grep result not grouped: %q", text[:60]) + } + if grep["cache_control"] == nil { + t.Fatal("cache_control marker was dropped") + } + read := results[1].(map[string]any)["content"].([]any)[0].(map[string]any) + if read["text"] != fileText { + t.Fatal("Read result was rewritten") + } + if results[2].(map[string]any)["content"] != search { + t.Fatal("error result was rewritten") + } + if dup := results[3].(map[string]any)["content"]; dup != "[Identical to the earlier output of tool call toolu_1.]" { + t.Fatalf("duplicate = %v", dup) + } + + again, err := CompressRequest(ProtocolAnthropic, body, Options{}) + if err != nil || string(again.Body) != string(result.Body) { + t.Fatal("compressing the same request twice gave different bytes") + } +} + +func TestCompressRequestOpenAIChatAndResponses(t *testing.T) { + search := grepOutput(2, 10) + chat := mustJSON(t, map[string]any{ + "model": "gpt", + "messages": []any{ + map[string]any{"role": "assistant", "tool_calls": []any{ + map[string]any{"id": "call_1", "type": "function", "function": map[string]any{"name": "grep", "arguments": "{}"}}, + }}, + map[string]any{"role": "tool", "tool_call_id": "call_1", "content": search}, + }, + }) + result, err := CompressRequest(ProtocolOpenAIChat, chat, Options{}) + if err != nil || !result.Changed || result.Stats.Compressed != 1 { + t.Fatalf("chat: %+v %v", result.Stats, err) + } + if strings.Contains(string(result.Body), `internal/pkg/file0.go:3:`) { + t.Fatal("chat tool message not compressed") + } + + responses := mustJSON(t, map[string]any{ + "model": "gpt", + "input": []any{ + map[string]any{"type": "message", "role": "user", "content": search}, + map[string]any{"type": "function_call", "call_id": "call_9", "name": "shell", "arguments": "{}"}, + map[string]any{"type": "function_call_output", "call_id": "call_9", "output": search}, + }, + }) + result, err = CompressRequest(ProtocolResponses, responses, Options{}) + if err != nil || result.Stats.Blocks != 1 || result.Stats.Compressed != 1 { + t.Fatalf("responses: %+v %v", result.Stats, err) + } + var got struct { + Input []map[string]any `json:"input"` + } + if err := json.Unmarshal(result.Body, &got); err != nil { + t.Fatal(err) + } + if got.Input[0]["content"] != search || got.Input[2]["output"] == search { + t.Fatal("responses: wrong item rewritten") + } +} + +func TestCompressRequestUnchangedKeepsOriginalBytes(t *testing.T) { + body := []byte(`{"model":"m", "messages":[{"role":"user","content":"hi"}], "big": 12345678901234567890}`) + result, err := CompressRequest(ProtocolOpenAIChat, body, Options{}) + if err != nil || result.Changed || string(result.Body) != string(body) { + t.Fatalf("unchanged request was rewritten: %s %v", result.Body, err) + } + if _, err := CompressRequest(ProtocolAnthropic, []byte("not json"), Options{}); err == nil { + t.Fatal("invalid JSON accepted") + } +} + +func TestCompressRequestPreservesLargeNumbers(t *testing.T) { + body := mustJSON(t, map[string]any{ + "messages": []any{map[string]any{"role": "tool", "tool_call_id": "c", "content": grepOutput(2, 10)}}, + }) + body = []byte(strings.Replace(string(body), `{"messages"`, `{"seed":12345678901234567890,"temperature":0.10,"messages"`, 1)) + result, err := CompressRequest(ProtocolOpenAIChat, body, Options{}) + if err != nil || !result.Changed { + t.Fatalf("%v %+v", err, result.Stats) + } + if !strings.Contains(string(result.Body), `"seed":12345678901234567890`) || !strings.Contains(string(result.Body), `"temperature":0.10`) { + t.Fatalf("numbers changed: %s", result.Body[:80]) + } +} + +func TestCompressRequestNeverDedupesFileContent(t *testing.T) { + var read, code strings.Builder + for i := 1; i <= 60; i++ { + fmt.Fprintf(&read, "%d\tline %d of the file\n", i, i) + } + for i := 0; i < 30; i++ { + fmt.Fprintf(&code, "func f%d() {\n\treturn\n}\n", i) + } + result := func(id, text string) map[string]any { + return map[string]any{"type": "tool_result", "tool_use_id": id, "content": text} + } + body := mustJSON(t, map[string]any{"messages": []any{ + map[string]any{"role": "assistant", "content": []any{ + map[string]any{"type": "tool_use", "id": "t1", "name": "Read", "input": map[string]any{}}, + map[string]any{"type": "tool_use", "id": "t2", "name": "Read", "input": map[string]any{}}, + map[string]any{"type": "tool_use", "id": "t3", "name": "Bash", "input": map[string]any{}}, + map[string]any{"type": "tool_use", "id": "t4", "name": "Bash", "input": map[string]any{}}, + map[string]any{"type": "tool_use", "id": "t5", "name": "Bash", "input": map[string]any{}}, + }}, + map[string]any{"role": "user", "content": []any{ + result("t1", read.String()), result("t2", read.String()), + result("t3", read.String()), // the same numbered content through a shell + result("t4", code.String()), result("t5", code.String()), + }}, + }}) + out, err := CompressRequest(ProtocolAnthropic, body, Options{Sample: true}) + if err != nil { + t.Fatal(err) + } + if out.Changed || strings.Contains(string(out.Body), "Identical to the earlier output") { + t.Fatalf("repeated file content was replaced: %s", out.Body) + } +} diff --git a/internal/ctxcompress/sample.go b/internal/ctxcompress/sample.go new file mode 100644 index 00000000..415aa890 --- /dev/null +++ b/internal/ctxcompress/sample.go @@ -0,0 +1,214 @@ +package ctxcompress + +import ( + "fmt" + "math" + "sort" + "strings" +) + +// Array sampling follows Headroom's SmartCrusher: keep the boundaries, every +// item that reports an error or looks structurally unusual, and an even +// spread of the rest. +const ( + sampleMinItems = 30 + sampleKeep = 15 + sampleFirst = 5 + sampleLast = 3 +) + +// jsonErrorKeywords is SmartCrusher's list: an item whose compact JSON +// contains one of them is always kept. +var jsonErrorKeywords = []string{ + "error", "exception", "failed", "failure", "critical", "fatal", "crash", + "panic", "abort", "timeout", "denied", "rejected", +} + +const omittedItemsFormat = "[… %d of %d items omitted]" + +// sampleJSON shortens every long array in the document in place, replacing +// the dropped items with one marker string at the end. +func sampleJSON(value *jsonValue) { + switch value.kind { + case 'o': + for _, field := range value.fields { + sampleJSON(field) + } + case 'a': + for _, item := range value.items { + sampleJSON(item) + } + if kept, total := sampleArray(value.items); len(kept) < total { + value.items = append(kept, &jsonValue{kind: 's', str: fmt.Sprintf(omittedItemsFormat, total-len(kept), total)}) + } + } +} + +// sampleArray returns the items to keep, in their original order. +func sampleArray(items []*jsonValue) ([]*jsonValue, int) { + n := len(items) + if n < sampleMinItems { + return items, n + } + encoded := make([]string, n) + for i, item := range items { + encoded[i] = item.compact() + } + + keep := map[int]bool{} + for i := 0; i < sampleFirst; i++ { + keep[i] = true + } + for i := n - sampleLast; i < n; i++ { + keep[i] = true + } + for i, text := range encoded { + lower := strings.ToLower(text) + for _, keyword := range jsonErrorKeywords { + if strings.Contains(lower, keyword) { + keep[i] = true + break + } + } + } + for _, i := range structuralOutliers(items) { + keep[i] = true + } + for _, i := range lengthOutliers(encoded) { + keep[i] = true + } + + // Spread the remaining budget evenly, skipping duplicates of kept items. + seen := map[string]bool{} + for i := range keep { + seen[encoded[i]] = true + } + if remaining := sampleKeep - len(keep); remaining > 0 { + step := max(1, n/(remaining+1)) + for i := step; i < n && remaining > 0; i += step { + if keep[i] || seen[encoded[i]] { + continue + } + keep[i] = true + seen[encoded[i]] = true + remaining-- + } + } + + indices := make([]int, 0, len(keep)) + for i := range keep { + indices = append(indices, i) + } + sort.Ints(indices) + kept := make([]*jsonValue, 0, len(indices)) + for _, i := range indices { + kept = append(kept, items[i]) + } + return kept, n +} + +// structuralOutliers finds objects carrying a field that fewer than a fifth +// of the items have, and objects whose value in a common field is one of its +// rare values (a "failed" status among "ok"s). +func structuralOutliers(items []*jsonValue) []int { + n := len(items) + frequency := map[string]int{} + for _, item := range items { + if item.kind != 'o' { + return nil + } + for _, key := range item.keys { + frequency[key]++ + } + } + var out []int + for i, item := range items { + for _, key := range item.keys { + if frequency[key]*5 < n { + out = append(out, i) + break + } + } + } + + var common []string + for key, count := range frequency { + if count*5 >= n*4 { + common = append(common, key) + } + } + sort.Strings(common) + for _, key := range common { + counts := map[string]int{} + for _, item := range items { + if cell := item.field(key); cell != nil && cell.kind != 'o' && cell.kind != 'a' { + counts[cell.compact()]++ + } + } + if len(counts) < 2 || len(counts) > 50 { + continue + } + type valueCount struct { + value string + count int + } + ranked := make([]valueCount, 0, len(counts)) + total := 0 + for value, count := range counts { + ranked = append(ranked, valueCount{value, count}) + total += count + } + sort.Slice(ranked, func(i, j int) bool { + if ranked[i].count != ranked[j].count { + return ranked[i].count > ranked[j].count + } + return ranked[i].value < ranked[j].value + }) + // The smallest set of values covering 80% of items is "normal". + covered, top := 0, 0 + for top < len(ranked) && covered*5 < total*4 { + covered += ranked[top].count + top++ + } + if top > 5 || top == len(ranked) { + continue + } + normal := map[string]bool{} + for _, entry := range ranked[:top] { + normal[entry.value] = true + } + for i, item := range items { + if cell := item.field(key); cell != nil && cell.kind != 'o' && cell.kind != 'a' && !normal[cell.compact()] { + out = append(out, i) + } + } + } + return out +} + +// lengthOutliers finds items whose encoded size is more than two standard +// deviations from the mean — the one record with a stack trace in it. +func lengthOutliers(encoded []string) []int { + n := float64(len(encoded)) + mean := 0.0 + for _, text := range encoded { + mean += float64(len(text)) + } + mean /= n + variance := 0.0 + for _, text := range encoded { + d := float64(len(text)) - mean + variance += d * d + } + sigma := math.Sqrt(variance / n) + if sigma == 0 { + return nil + } + var out []int + for i, text := range encoded { + if math.Abs(float64(len(text))-mean) > 2*sigma { + out = append(out, i) + } + } + return out +} diff --git a/internal/ctxcompress/search.go b/internal/ctxcompress/search.go new file mode 100644 index 00000000..62d9144e --- /dev/null +++ b/internal/ctxcompress/search.go @@ -0,0 +1,181 @@ +package ctxcompress + +import ( + "fmt" + "regexp" + "strings" +) + +// searchLine matches grep/ripgrep output: path, line number, then ':' for a +// match or '-' for a context line. The path must look like one (a slash or +// an extension) so timestamps such as "12:30:45" are not mistaken for it. +var searchLine = regexp.MustCompile(`^((?:[A-Za-z]:[\\/])?[^\s:]*[/\\.][^\s:]*)([:-])(\d+)([:-])(.*)$`) + +// searchPathOnly matches `grep` output without line numbers. The path needs +// a directory separator so "key: value" lines are not taken for matches. +var searchPathOnly = regexp.MustCompile(`^((?:[A-Za-z]:[\\/])?[^\s:]*[/\\][^\s:]*[^\s:/\\]):(.*)$`) + +type searchMatch struct { + path string + // rest is everything after the path, starting with the line number when + // there is one ("12:content"), otherwise the content itself. + rest string + ok bool +} + +func parseSearchLine(line string) searchMatch { + if m := searchLine.FindStringSubmatch(line); m != nil && m[2] == m[4] { + return searchMatch{path: m[1], rest: m[3] + m[4] + m[5], ok: true} + } + if m := searchPathOnly.FindStringSubmatch(line); m != nil { + return searchMatch{path: m[1], rest: m[2], ok: true} + } + return searchMatch{} +} + +// looksLikeSearch wants most lines to parse as matches and at least one file +// to appear more than once, since grouping is what saves the bytes. +func looksLikeSearch(lines []string) bool { + nonEmpty, parsed, repeats := 0, 0, 0 + previous := "" + for _, line := range lines { + if strings.TrimSpace(line) == "" || line == "--" { + continue + } + nonEmpty++ + match := parseSearchLine(line) + if !match.ok { + continue + } + parsed++ + if match.path == previous { + repeats++ + } + previous = match.path + } + return parsed >= 5 && parsed*10 >= nonEmpty*6 && repeats > 0 +} + +// groupHeader is the file name compressSearch prints above a group. +var groupHeader = regexp.MustCompile(`^(?:[A-Za-z]:[\\/])?[^\s:]*[/\\.][^\s:]*$`) + +// looksLikeGroupedSearch recognises compressSearch output — file names each +// followed by indented matches — so a second pass leaves it alone rather than +// treating it as a log. +func looksLikeGroupedSearch(lines []string) bool { + nonEmpty, grouped, ungrouped := 0, 0, 0 + inGroup := false + for i, line := range lines { + if strings.TrimSpace(line) == "" { + inGroup = false + continue + } + nonEmpty++ + switch { + case inGroup && strings.HasPrefix(line, " "): + grouped++ + case groupHeader.MatchString(line) && i+1 < len(lines) && strings.HasPrefix(lines[i+1], " "): + inGroup = true + grouped++ + default: + inGroup = false + if parseSearchLine(line).ok { + ungrouped++ + } + } + } + return grouped >= 3 && (grouped+ungrouped)*2 >= nonEmpty +} + +const ( + searchMaxLineDivisor = 4 + searchPerFileKeep = 12 + searchPerFileFirst = 4 + searchPerFileLast = 2 +) + +// compressSearch prints each file's name once above its matches instead of +// on every line, and cuts over-long matches. With sampling it +// also keeps at most searchPerFileKeep matches per file, preferring the first, +// the last and those mentioning errors. +func compressSearch(lines []string, opts Options) string { + // A match longer than a normal line is usually minified code or a binary + // file; sampling cuts those shorter than safe mode does. + maxLine := opts.MaxLineBytes + if opts.Sample { + maxLine /= searchMaxLineDivisor + } + lines = trimLongLines(lines, maxLine) + out := make([]string, 0, len(lines)) + for i := 0; i < len(lines); { + first := parseSearchLine(lines[i]) + if !first.ok { + out = append(out, lines[i]) + i++ + continue + } + var group []string + j := i + for ; j < len(lines); j++ { + match := parseSearchLine(lines[j]) + if !match.ok || match.path != first.path { + break + } + group = append(group, match.rest) + } + if len(group) == 1 { + out = append(out, lines[i]) + } else { + if opts.Sample { + group = sampleSearchGroup(group) + } + out = append(out, first.path) + for _, rest := range group { + out = append(out, " "+rest) + } + } + i = j + } + return joinLines(foldRepeats(out)) +} + +func sampleSearchGroup(group []string) []string { + if len(group) <= searchPerFileKeep { + return group + } + keep := make([]bool, len(group)) + kept := 0 + mark := func(i int) { + if !keep[i] { + keep[i] = true + kept++ + } + } + for i := 0; i < searchPerFileFirst; i++ { + mark(i) + } + for i := len(group) - searchPerFileLast; i < len(group); i++ { + mark(i) + } + for i, rest := range group { + if kept >= searchPerFileKeep { + break + } + if isErrorLine(rest) { + mark(i) + } + } + if remaining := searchPerFileKeep - kept; remaining > 0 { + step := max(1, len(group)/(remaining+1)) + for i := step; i < len(group) && kept < searchPerFileKeep; i += step { + mark(i) + } + } + out := make([]string, 0, kept+1) + for i, rest := range group { + if keep[i] { + out = append(out, rest) + } + } + return append(out, fmt.Sprintf("[… %d more matches in this file]", len(group)-kept)) +} diff --git a/internal/ctxcompress/stats.go b/internal/ctxcompress/stats.go new file mode 100644 index 00000000..74f1a620 --- /dev/null +++ b/internal/ctxcompress/stats.go @@ -0,0 +1,79 @@ +package ctxcompress + +// Kind is the content type Text detected for a block and routed it by. +type Kind string + +const ( + KindSkipped Kind = "skipped" + KindFileRead Kind = "file_read" + KindCode Kind = "code" + KindJSON Kind = "json" + KindSearch Kind = "search" + KindDiff Kind = "diff" + KindLog Kind = "log" + KindText Kind = "text" + // KindDuplicate is a tool result replaced by a pointer to an identical + // earlier one. + KindDuplicate Kind = "duplicate" +) + +// KindStats counts the tool results of one kind. +type KindStats struct { + Blocks int `json:"blocks"` + Compressed int `json:"compressed"` + BytesBefore int `json:"bytes_before"` + BytesAfter int `json:"bytes_after"` + // TokensSaved is an estimate; see EstimateTokens. + TokensSaved int `json:"tokens_saved"` +} + +// Stats totals what CompressRequest did to the tool results of one request. +type Stats struct { + Blocks int `json:"blocks"` + Compressed int `json:"compressed"` + BytesBefore int `json:"bytes_before"` + BytesAfter int `json:"bytes_after"` + TokensSaved int `json:"tokens_saved"` + ByKind map[Kind]KindStats `json:"by_kind,omitempty"` +} + +func (s *Stats) add(kind Kind, before, after string) { + changed := before != after + saved := 0 + if changed { + saved = EstimateTokens(before) - EstimateTokens(after) + } + s.Blocks++ + s.BytesBefore += len(before) + s.BytesAfter += len(after) + s.TokensSaved += saved + if s.ByKind == nil { + s.ByKind = map[Kind]KindStats{} + } + entry := s.ByKind[kind] + entry.Blocks++ + entry.BytesBefore += len(before) + entry.BytesAfter += len(after) + entry.TokensSaved += saved + if changed { + s.Compressed++ + entry.Compressed++ + } + s.ByKind[kind] = entry +} + +// EstimateTokens approximates a tokenizer without loading one: 3.5 ASCII +// characters or 1.2 characters of other scripts per token. On captured +// coding-agent traffic it is within 2% of o200k counts overall, and for text +// that is mostly Chinese. +func EstimateTokens(text string) int { + ascii, other := 0, 0 + for _, r := range text { + if r < 0x80 { + ascii++ + } else { + other++ + } + } + return (ascii*12 + other*35) / 42 +} diff --git a/internal/ctxcompress/text.go b/internal/ctxcompress/text.go new file mode 100644 index 00000000..a052c5c3 --- /dev/null +++ b/internal/ctxcompress/text.go @@ -0,0 +1,199 @@ +// Package ctxcompress shrinks the tool output inside LLM chat requests before +// they are sent to a model. +// +// It is a model-free port of the deterministic parts of Headroom +// (github.com/headroomlabs-ai/headroom): content is routed by type, JSON is +// minified, runs of identical lines are counted, over-long lines are cut and +// search results are grouped by file. With Options.Sample it also drops +// content: long arrays, logs and search results keep a representative subset, +// and lines differing only in their numbers are folded. Source files a tool +// read back are never changed, because a coding agent edits them by exact +// string match. +// +// Every transform is deterministic and idempotent: the same text always +// compresses to the same bytes, and compressing the output again changes +// nothing. The first keeps prompt caches warm across turns; the second makes a +// request that passes through two gateways, such as a LAN cluster hop, safe. +package ctxcompress + +import ( + "strings" + "unicode/utf8" +) + +// Options tunes Text. The zero value selects the defaults. +type Options struct { + // MinBytes is the smallest tool result worth compressing. + MinBytes int + // MaxLineBytes is the length above which a single line is cut down. + MaxLineBytes int + // MaxLines is the line count above which sampled command output keeps + // only its head, its tail and the lines that report errors. On captured + // agent traffic 150 keeps every error line and still halves long build + // and test logs. + MaxLines int + // Sample enables the lossy transforms: long JSON arrays, long command + // output and long search results keep a representative subset instead of + // everything. Without it only redundancy is removed. + Sample bool +} + +const ( + defaultMinBytes = 512 + defaultMaxLineBytes = 2000 + defaultMaxLines = 150 +) + +func (o Options) withDefaults() Options { + if o.MinBytes <= 0 { + o.MinBytes = defaultMinBytes + } + if o.MaxLineBytes <= 0 { + o.MaxLineBytes = defaultMaxLineBytes + } + if o.MaxLines <= 0 { + o.MaxLines = defaultMaxLines + } + return o +} + +// Text compresses one tool result. toolName is the name of the tool that +// produced it when the request says so, and only steers detection. +func Text(text, toolName string, opts Options) (string, Kind) { + opts = opts.withDefaults() + if len(text) < opts.MinBytes || !utf8.ValidString(text) { + return text, KindSkipped + } + if isFileReadTool(toolName) || looksLikeFileRead(text) { + return text, KindFileRead + } + if first, _, _ := strings.Cut(text, "\n"); isTableHeader(first) { + return text, KindJSON // already compressed + } + if out, ok := compressJSON(text, opts); ok { + return keepSmaller(text, out), KindJSON + } + + lines := normalizeLines(text) + switch { + case hasFactoredPrefix(lines): + // Only compressLog writes prefix headers, so this is its output. + return keepSmaller(text, compressLog(lines, opts)), KindLog + case looksLikeDiff(lines): + // Every line of a diff is load-bearing for a review or a patch, so it + // only loses what normalizeLines strips and over-long lines. + return keepSmaller(text, joinLines(trimLongLines(lines, opts.MaxLineBytes))), KindDiff + case looksLikeGroupedSearch(lines): + return text, KindSearch // already compressed + case looksLikeSearch(lines): + return keepSmaller(text, compressSearch(lines, opts)), KindSearch + case looksLikeCode(lines): + return text, KindCode + case looksLikeLog(lines): + return keepSmaller(text, compressLog(lines, opts)), KindLog + default: + return keepSmaller(text, compressProse(lines, opts)), KindText + } +} + +// keepSmaller returns the original unless the rewrite is strictly shorter, +// so a transform that happens to add bytes never ships. +func keepSmaller(original, rewritten string) string { + if len(rewritten) < len(original) { + return rewritten + } + return original +} + +// fileReadTools are the tool names coding agents use to read a file back. +// Their output is quoted verbatim in later edits, so it is never rewritten. +var fileReadTools = map[string]bool{ + "read": true, "read_file": true, "readfile": true, "view": true, + "view_file": true, "open_file": true, "cat": true, "str_replace_editor": true, + "str_replace_based_edit_tool": true, "notebookread": true, "read_many_files": true, + "edit": true, "write": true, "multiedit": true, "apply_patch": true, + "notebookedit": true, "str_replace": true, +} + +func isFileReadTool(name string) bool { + return fileReadTools[strings.ToLower(strings.TrimSpace(name))] +} + +// looksLikeFileRead spots `cat -n` style output — a line number, then a tab +// or an arrow — which is how agents show file contents whatever the tool is +// called. +func looksLikeFileRead(text string) bool { + total, numbered := 0, 0 + for _, line := range firstLines(text, 30) { + if strings.TrimSpace(line) == "" { + continue + } + total++ + if isNumberedLine(line) { + numbered++ + } + } + return total >= 3 && numbered*10 >= total*8 +} + +func isNumberedLine(line string) bool { + trimmed := strings.TrimLeft(line, " ") + digits := 0 + for digits < len(trimmed) && trimmed[digits] >= '0' && trimmed[digits] <= '9' { + digits++ + } + if digits == 0 || digits == len(trimmed) { + return digits > 0 + } + rest := trimmed[digits:] + return rest[0] == '\t' || strings.HasPrefix(rest, "→") || strings.HasPrefix(rest, "│") +} + +func firstLines(text string, n int) []string { + lines := make([]string, 0, n) + for len(lines) < n && text != "" { + line, rest, found := strings.Cut(text, "\n") + lines = append(lines, line) + if !found { + break + } + text = rest + } + return lines +} + +func joinLines(lines []string) string { + return strings.Join(lines, "\n") +} + +var codeKeywords = []string{ + "func ", "def ", "class ", "import ", "from ", "return", "if ", "for ", "while ", + "const ", "let ", "var ", "package ", "#include", "public ", "private ", "static ", + "export ", "type ", "interface ", "struct ", "fn ", "impl ", "use ", "async ", "await ", + "try", "catch", "else", "switch ", "case ", "//", "/*", "* ", "# ", "@", +} + +// looksLikeCode spots a source file printed without line numbers, as by +// `cat` or `sed -n` through a shell tool. Such output is protected like a +// file read: an agent copies it into edits byte for byte. +func looksLikeCode(lines []string) bool { + nonEmpty, code := 0, 0 + for _, line := range lines { + trimmed := strings.TrimSpace(line) + if trimmed == "" { + continue + } + nonEmpty++ + if strings.ContainsAny(trimmed[len(trimmed)-1:], "{};:") { + code++ + continue + } + for _, keyword := range codeKeywords { + if strings.HasPrefix(trimmed, keyword) { + code++ + break + } + } + } + return nonEmpty >= 8 && code*2 >= nonEmpty +} diff --git a/internal/ctxcompress/text_test.go b/internal/ctxcompress/text_test.go new file mode 100644 index 00000000..1a94611e --- /dev/null +++ b/internal/ctxcompress/text_test.go @@ -0,0 +1,269 @@ +package ctxcompress + +import ( + "fmt" + "strings" + "testing" +) + +// assertIdempotent checks that compressing the output again changes nothing. +func assertIdempotent(t *testing.T, out, tool string, opts Options) { + t.Helper() + if again, _ := Text(out, tool, opts); again != out { + t.Fatalf("second pass changed the output:\nfirst:\n%s\nsecond:\n%s", out, again) + } +} + +func TestTextLeavesSmallAndFileContentAlone(t *testing.T) { + small := "ok\nok\nok\n" + if out, kind := Text(small, "Bash", Options{}); out != small || kind != KindSkipped { + t.Fatalf("small result changed: %q (%s)", out, kind) + } + + var numbered strings.Builder + for i := 1; i <= 200; i++ { + fmt.Fprintf(&numbered, "%d\tline \n", i) + } + if out, kind := Text(numbered.String(), "Bash", Options{}); out != numbered.String() || kind != KindFileRead { + t.Fatalf("numbered file content was rewritten (%s)", kind) + } + + repeated := strings.Repeat("same line\n", 200) + if out, kind := Text(repeated, "Read", Options{}); out != repeated || kind != KindFileRead { + t.Fatalf("output of the Read tool was rewritten (%s)", kind) + } + + var code strings.Builder + for i := 0; i < 40; i++ { + fmt.Fprintf(&code, "func f%d() {\n\treturn\n}\n\n", i) + } + if out, kind := Text(code.String(), "Bash", Options{}); out != code.String() || kind != KindCode { + t.Fatalf("source code printed by a shell was rewritten (%s)", kind) + } +} + +func TestTextJSONMinifiesAndRendersTables(t *testing.T) { + var rows []string + for i := 0; i < 20; i++ { + rows = append(rows, fmt.Sprintf(` {"id": %d, "name": "item %d", "ok": true, "note": "a, b"}`, i, i)) + } + input := "[\n" + strings.Join(rows, ",\n") + "\n]" + out, kind := Text(input, "", Options{}) + if kind != KindJSON { + t.Fatalf("kind = %s", kind) + } + lines := strings.Split(out, "\n") + if lines[0] != "[20]{id:number,name:string,ok:bool,note:string}" { + t.Fatalf("header = %q", lines[0]) + } + if lines[1] != `0,item 0,true,"a, b"` || len(lines) != 21 { + t.Fatalf("rows = %q (%d lines)", lines[1], len(lines)) + } + assertIdempotent(t, out, "", Options{}) + + object := "{\n \"b\": 1,\n \"a\": [1, 2, 3],\n \"text\": \"
\"" + strings.Repeat(" ", 600) + "\n}" + out, _ = Text(object, "", Options{}) + if out != `{"b":1,"a":[1,2,3],"text":"
"}` {
+ t.Fatalf("object not minified in key order: %s", out)
+ }
+}
+
+func TestTextJSONSamplingKeepsErrorsAndOutliers(t *testing.T) {
+ var rows []string
+ for i := 0; i < 100; i++ {
+ status := "ok"
+ switch i {
+ case 37:
+ status = "error: disk full"
+ case 61:
+ status = "degraded"
+ }
+ rows = append(rows, fmt.Sprintf(`{"id":%d,"status":%q}`, i, status))
+ }
+ input := "[" + strings.Join(rows, ",") + "]"
+
+ safe, _ := Text(input, "", Options{})
+ if strings.Count(safe, "\n") != 100 {
+ t.Fatalf("safe mode dropped rows:\n%s", safe)
+ }
+
+ opts := Options{Sample: true}
+ out, _ := Text(input, "", opts)
+ for _, want := range []string{"\n0,ok", "\n37,error: disk full", "\n61,degraded", "\n99,ok", "items omitted]"} {
+ if !strings.Contains(out, want) {
+ t.Fatalf("sampled output misses %q:\n%s", want, out)
+ }
+ }
+ if rows := strings.Count(out, "\n"); rows > 25 {
+ t.Fatalf("sampled output kept %d rows", rows)
+ }
+ assertIdempotent(t, out, "", opts)
+}
+
+func TestTextLogFoldsRepeatsAndNumericChurn(t *testing.T) {
+ var log strings.Builder
+ log.WriteString("\x1b[32mstarting\x1b[0m \n")
+ for i := 0; i < 50; i++ {
+ fmt.Fprintf(&log, "downloading chunk %d of 50 (%d%%)\n", i+1, i*2)
+ }
+ for i := 0; i < 20; i++ {
+ log.WriteString("waiting for server\n")
+ }
+ log.WriteString("progress 10%\rprogress 50%\rprogress 100%\n")
+ log.WriteString("done\n")
+
+ // Safe mode only counts identical lines; every numbered line survives.
+ safe, kind := Text(log.String(), "Bash", Options{})
+ if kind != KindLog {
+ t.Fatalf("kind = %s", kind)
+ }
+ for _, want := range []string{
+ "starting\n",
+ "downloading chunk 17 of 50 (32%)\n",
+ "waiting for server\n[… previous line repeated 19 more times]\n",
+ "\nprogress 100%\ndone",
+ } {
+ if !strings.Contains(safe, want) {
+ t.Fatalf("safe output misses %q:\n%s", want, safe)
+ }
+ }
+ if strings.Contains(safe, "similar lines omitted") || strings.Contains(safe, "\x1b") {
+ t.Fatalf("safe mode folded numeric lines or kept escapes:\n%q", safe)
+ }
+ assertIdempotent(t, safe, "Bash", Options{})
+
+ // Aggressive mode also folds lines that differ only in their numbers.
+ opts := Options{Sample: true}
+ out, _ := Text(log.String(), "Bash", opts)
+ want := "downloading chunk 1 of 50 (0%)\ndownloading chunk 2 of 50 (2%)\n[… 46 similar lines omitted]\ndownloading chunk 49 of 50 (96%)\ndownloading chunk 50 of 50 (98%)\n"
+ if !strings.Contains(out, want) {
+ t.Fatalf("aggressive output misses the folded run:\n%s", out)
+ }
+ assertIdempotent(t, out, "Bash", opts)
+}
+
+// Lines that differ only in numbers can still each carry information: test
+// failures with their values, numbered files in a listing.
+func TestTextKeepsNumberedErrorsAndListings(t *testing.T) {
+ var log strings.Builder
+ log.WriteString("running tests in package github.com/opencsgs/csglite/internal/example\n")
+ for i := 1; i <= 8; i++ {
+ fmt.Fprintf(&log, "error: test %d failed: expected %d got %d\n", i, i, i+1)
+ }
+ for i := 1; i <= 8; i++ {
+ fmt.Fprintf(&log, "internal/server/src/file%d.go\n", i)
+ }
+ log.WriteString("done\n")
+
+ safe, _ := Text(log.String(), "Bash", Options{})
+ for i := 1; i <= 8; i++ {
+ for _, want := range []string{
+ fmt.Sprintf("error: test %d failed: expected %d got %d", i, i, i+1),
+ fmt.Sprintf("internal/server/src/file%d.go", i),
+ } {
+ if !strings.Contains(safe, want) {
+ t.Fatalf("safe mode dropped %q:\n%s", want, safe)
+ }
+ }
+ }
+
+ aggressive, _ := Text(log.String(), "Bash", Options{Sample: true})
+ for i := 1; i <= 8; i++ {
+ if want := fmt.Sprintf("error: test %d failed: expected %d got %d", i, i, i+1); !strings.Contains(aggressive, want) {
+ t.Fatalf("aggressive mode dropped error line %q:\n%s", want, aggressive)
+ }
+ }
+}
+
+func TestTextLogSamplingKeepsErrorsWithinBudget(t *testing.T) {
+ var log strings.Builder
+ for i := 0; i < 1000; i++ {
+ switch i {
+ case 500:
+ log.WriteString("FAIL: TestSomething expected 3 got 4\n")
+ default:
+ fmt.Fprintf(&log, "step %s completed\n", strings.Repeat("x", i%7+1))
+ }
+ }
+ input := log.String()
+ if out, _ := Text(input, "Bash", Options{}); strings.Contains(out, "lines omitted") {
+ t.Fatal("safe mode sampled lines")
+ }
+
+ opts := Options{Sample: true, MaxLines: 100}
+ out, _ := Text(input, "Bash", opts)
+ lines := strings.Split(out, "\n")
+ if len(lines) > 100 {
+ t.Fatalf("kept %d lines, budget 100", len(lines))
+ }
+ if !strings.Contains(out, "FAIL: TestSomething") || !strings.Contains(out, "lines omitted]") {
+ t.Fatalf("error line or marker missing:\n%s", out)
+ }
+ assertIdempotent(t, out, "Bash", opts)
+}
+
+func TestTextSearchGroupsByFile(t *testing.T) {
+ var grep strings.Builder
+ for i := 1; i <= 8; i++ {
+ fmt.Fprintf(&grep, "internal/server/routes.go:%d:\tmux.HandleFunc(%d)\n", i*10, i)
+ }
+ grep.WriteString("internal/server/auth.go:12:func auth() {}\n")
+ fmt.Fprintf(&grep, "bin/csghub-lite:%s\n", strings.Repeat("\x01garbage", 400))
+ for i := 1; i <= 3; i++ {
+ fmt.Fprintf(&grep, "web/src/api/client.ts:%d: const x = %d;\n", i, i)
+ }
+
+ out, kind := Text(grep.String(), "Grep", Options{})
+ if kind != KindSearch {
+ t.Fatalf("kind = %s", kind)
+ }
+ for _, want := range []string{
+ "internal/server/routes.go\n 10:\tmux.HandleFunc(1)\n 20:\tmux.HandleFunc(2)\n",
+ "\ninternal/server/auth.go:12:func auth() {}\n",
+ "bytes omitted]…",
+ "\nweb/src/api/client.ts\n 1: const x = 1;\n",
+ } {
+ if !strings.Contains(out, want) {
+ t.Fatalf("output misses %q:\n%s", want, out)
+ }
+ }
+ assertIdempotent(t, out, "Grep", Options{})
+}
+
+func TestTextAggressiveFactorsSharedPrefixes(t *testing.T) {
+ var log strings.Builder
+ for i := 0; i < 40; i++ {
+ fmt.Fprintf(&log, "build (windows-latest)\tRun tests\tok github.com/opencsgs/csglite/internal/pkg%c\n", 'a'+i%26)
+ }
+ opts := Options{Sample: true}
+ out, _ := Text(log.String(), "Bash", opts)
+ if !strings.HasPrefix(out, `[the next 40 lines each start with: "build (windows-latest)\tRun tests\tok github.com/opencsgs/csglite/internal/"]`) {
+ t.Fatalf("prefix not factored:\n%s", out)
+ }
+ assertIdempotent(t, out, "Bash", opts)
+ if safe, _ := Text(log.String(), "Bash", Options{}); strings.Contains(safe, "[the next") {
+ t.Fatal("safe mode factored prefixes")
+ }
+}
+
+func TestTrimLineKeepsUTF8Valid(t *testing.T) {
+ line := strings.Repeat("中文", 2000)
+ out := trimLine(line, 1001)
+ if !strings.Contains(out, "bytes omitted") || !strings.HasPrefix(out, "中文") {
+ t.Fatalf("trimmed line = %q", out[:40])
+ }
+ for _, r := range out {
+ if r == '\uFFFD' {
+ t.Fatal("trimLine split a UTF-8 sequence")
+ }
+ }
+}
+
+func TestEstimateTokens(t *testing.T) {
+ if got := EstimateTokens(strings.Repeat("abcdefg", 100)); got != 200 {
+ t.Fatalf("ASCII estimate = %d, want 200", got)
+ }
+ if got := EstimateTokens(strings.Repeat("中", 120)); got != 100 {
+ t.Fatalf("CJK estimate = %d, want 100", got)
+ }
+}
diff --git a/internal/observability/store.go b/internal/observability/store.go
index 4eddc0c0..2aa54b5d 100644
--- a/internal/observability/store.go
+++ b/internal/observability/store.go
@@ -82,6 +82,15 @@ type RequestRecord struct {
ResponseBody string
RequestBodyTruncated bool
ResponseBodyTruncated bool
+ // Context compression: the mode that applied, the tool results the
+ // request carried and how many were rewritten, their size before and
+ // after, and an estimate of the input tokens that saved.
+ ContextCompressionMode string
+ ContextBlocks int64
+ ContextCompressedBlocks int64
+ ContextBytesBefore int64
+ ContextBytesAfter int64
+ ContextTokensSaved int64
}
type RequestFilter struct {
@@ -102,6 +111,9 @@ type RequestSummary struct {
Failed int64
TotalTokens int64
AverageLatency float64
+ // ContextTokensSaved is the estimated input tokens context compression
+ // removed across the matching requests.
+ ContextTokensSaved int64
}
type RequestPage struct {
@@ -227,7 +239,13 @@ CREATE TABLE IF NOT EXISTS requests (
response_body BLOB,
request_body_truncated INTEGER NOT NULL DEFAULT 0,
response_body_truncated INTEGER NOT NULL DEFAULT 0,
- usage_reconciled INTEGER NOT NULL DEFAULT 0
+ usage_reconciled INTEGER NOT NULL DEFAULT 0,
+ context_compression_mode TEXT NOT NULL DEFAULT '',
+ context_blocks INTEGER NOT NULL DEFAULT 0,
+ context_compressed_blocks INTEGER NOT NULL DEFAULT 0,
+ context_bytes_before INTEGER NOT NULL DEFAULT 0,
+ context_bytes_after INTEGER NOT NULL DEFAULT 0,
+ context_tokens_saved INTEGER NOT NULL DEFAULT 0
);
CREATE INDEX IF NOT EXISTS idx_requests_started_at ON requests(started_at DESC);
CREATE INDEX IF NOT EXISTS idx_requests_trace_id ON requests(trace_id, started_at);
@@ -272,6 +290,12 @@ CREATE INDEX IF NOT EXISTS idx_requests_pool_id ON requests(pool_id, pool_name);
{"estimated_cost", "REAL NOT NULL DEFAULT 0"},
{"cost_currency", "TEXT NOT NULL DEFAULT ''"},
{"cost_known", "INTEGER NOT NULL DEFAULT 0"},
+ {"context_compression_mode", "TEXT NOT NULL DEFAULT ''"},
+ {"context_blocks", "INTEGER NOT NULL DEFAULT 0"},
+ {"context_compressed_blocks", "INTEGER NOT NULL DEFAULT 0"},
+ {"context_bytes_before", "INTEGER NOT NULL DEFAULT 0"},
+ {"context_bytes_after", "INTEGER NOT NULL DEFAULT 0"},
+ {"context_tokens_saved", "INTEGER NOT NULL DEFAULT 0"},
} {
if err := s.addColumnIfMissing("requests", migration.name, migration.definition); err != nil {
return err
@@ -345,6 +369,8 @@ func (s *Store) Add(ctx context.Context, record RequestRecord) error {
record.CacheReadInputTokens, record.CacheCreationTokens, record.CacheEligibleTokens,
record.FirstTokenLatencyMS, record.ErrorMessage, requestBody, responseBody,
boolInt(record.RequestBodyTruncated), boolInt(record.ResponseBodyTruncated), 1,
+ record.ContextCompressionMode, record.ContextBlocks, record.ContextCompressedBlocks,
+ record.ContextBytesBefore, record.ContextBytesAfter, record.ContextTokensSaved,
}
query := `
INSERT OR REPLACE INTO requests (
@@ -358,7 +384,9 @@ INSERT OR REPLACE INTO requests (
fallback_count, limited_count, input_tokens, output_tokens, duration_ms,
cache_read_input_tokens, cache_creation_input_tokens, cache_eligible_input_tokens,
first_token_latency_ms, error_message, request_body, response_body,
- request_body_truncated, response_body_truncated, usage_reconciled
+ request_body_truncated, response_body_truncated, usage_reconciled,
+ context_compression_mode, context_blocks, context_compressed_blocks,
+ context_bytes_before, context_bytes_after, context_tokens_saved
) VALUES (` + strings.TrimSuffix(strings.Repeat("?,", len(args)), ",") + `)`
_, err = s.db.ExecContext(ctx, query, args...)
if err != nil {
@@ -452,10 +480,11 @@ SELECT COUNT(*),
COALESCE(SUM(CASE WHEN status = 'completed' THEN 1 ELSE 0 END), 0),
COALESCE(SUM(CASE WHEN status != 'completed' THEN 1 ELSE 0 END), 0),
COALESCE(SUM(MAX(input_tokens, cache_eligible_input_tokens) + output_tokens), 0),
- COALESCE(AVG(duration_ms), 0)
+ COALESCE(AVG(duration_ms), 0),
+ COALESCE(SUM(context_tokens_saved), 0)
FROM requests`+where, args...).Scan(
&page.Summary.Requests, &page.Summary.Succeeded, &page.Summary.Failed,
- &page.Summary.TotalTokens, &page.Summary.AverageLatency,
+ &page.Summary.TotalTokens, &page.Summary.AverageLatency, &page.Summary.ContextTokensSaved,
); err != nil {
return page, fmt.Errorf("summarizing observability requests: %w", err)
}
@@ -919,7 +948,9 @@ semantic_fallback, semantic_fallback_reason, price_input_per_million, price_outp
estimated_cost, cost_currency, cost_known, fallback_count, limited_count, input_tokens,
output_tokens, duration_ms, cache_read_input_tokens, cache_creation_input_tokens,
cache_eligible_input_tokens, first_token_latency_ms, error_message,
-request_body_truncated, response_body_truncated` + bodyColumns + ` FROM requests`
+request_body_truncated, response_body_truncated, context_compression_mode, context_blocks,
+context_compressed_blocks, context_bytes_before, context_bytes_after,
+context_tokens_saved` + bodyColumns + ` FROM requests`
}
type scanner interface {
@@ -945,6 +976,8 @@ func scanRequest(row scanner, includeBodies bool) (RequestRecord, error) {
&record.LimitedCount, &record.InputTokens, &record.OutputTokens, &record.DurationMS,
&record.CacheReadInputTokens, &record.CacheCreationTokens, &record.CacheEligibleTokens,
&record.FirstTokenLatencyMS, &record.ErrorMessage, &requestTruncated, &responseTruncated,
+ &record.ContextCompressionMode, &record.ContextBlocks, &record.ContextCompressedBlocks,
+ &record.ContextBytesBefore, &record.ContextBytesAfter, &record.ContextTokensSaved,
}
var requestBody, responseBody []byte
if includeBodies {
diff --git a/internal/server/context_compression.go b/internal/server/context_compression.go
new file mode 100644
index 00000000..324f298c
--- /dev/null
+++ b/internal/server/context_compression.go
@@ -0,0 +1,80 @@
+package server
+
+import (
+ "bytes"
+ "fmt"
+ "io"
+ "log"
+ "net/http"
+ "strconv"
+
+ "github.com/opencsgs/csglite/internal/config"
+ "github.com/opencsgs/csglite/internal/ctxcompress"
+)
+
+// contextCompressionHeader reports on the response what was done to the
+// request, so a client or a curl session can see the effect directly.
+const contextCompressionHeader = "X-CSGLite-Context-Compression"
+
+// contextCompressionOptions returns the compressor settings for the current
+// mode, and false when compression is off.
+func (s *Server) contextCompressionOptions() (string, ctxcompress.Options, bool) {
+ mode := config.NormalizeContextCompression(s.cfg.Inference.ContextCompression)
+ switch mode {
+ case config.ContextCompressionSafe:
+ return mode, ctxcompress.Options{}, true
+ case config.ContextCompressionAggressive:
+ return mode, ctxcompress.Options{Sample: true}, true
+ }
+ return mode, ctxcompress.Options{}, false
+}
+
+// withContextCompression rewrites the tool output in a chat request body
+// before next sees it. It wraps only the routes clients call; requests a
+// cluster member forwards arrive on the peer handler and were compressed
+// once already on the node that received them.
+//
+// The observability store keeps the body as the client sent it, and records
+// what compression saved beside it.
+func (s *Server) withContextCompression(protocol ctxcompress.Protocol, next http.HandlerFunc) http.HandlerFunc {
+ return func(w http.ResponseWriter, r *http.Request) {
+ mode, opts, enabled := s.contextCompressionOptions()
+ if !enabled || r.Body == nil {
+ next(w, r)
+ return
+ }
+ body, err := io.ReadAll(r.Body)
+ _ = r.Body.Close()
+ if err != nil {
+ // Let the handler report the broken body as it always has.
+ r.Body = io.NopCloser(io.MultiReader(bytes.NewReader(body), errorReader{err}))
+ next(w, r)
+ return
+ }
+ result, err := ctxcompress.CompressRequest(protocol, body, opts)
+ if err != nil {
+ // Not JSON the compressor understands; the handler decides.
+ r.Body = io.NopCloser(bytes.NewReader(body))
+ next(w, r)
+ return
+ }
+ if result.Stats.Blocks > 0 {
+ observationFromContext(r.Context()).setContextCompression(mode, result.Stats)
+ w.Header().Set(contextCompressionHeader, fmt.Sprintf("mode=%s; blocks=%d; compressed=%d; bytes=%d->%d; tokens_saved=%d",
+ mode, result.Stats.Blocks, result.Stats.Compressed, result.Stats.BytesBefore, result.Stats.BytesAfter, result.Stats.TokensSaved))
+ }
+ if result.Changed {
+ log.Printf("CONTEXT COMPRESSION %s %s: %d of %d tool results, %d -> %d bytes, ~%d tokens saved",
+ mode, r.URL.Path, result.Stats.Compressed, result.Stats.Blocks,
+ result.Stats.BytesBefore, result.Stats.BytesAfter, result.Stats.TokensSaved)
+ }
+ r.Body = io.NopCloser(bytes.NewReader(result.Body))
+ r.ContentLength = int64(len(result.Body))
+ r.Header.Set("Content-Length", strconv.Itoa(len(result.Body)))
+ next(w, r)
+ }
+}
+
+type errorReader struct{ err error }
+
+func (e errorReader) Read([]byte) (int, error) { return 0, e.err }
diff --git a/internal/server/context_compression_test.go b/internal/server/context_compression_test.go
new file mode 100644
index 00000000..542f1272
--- /dev/null
+++ b/internal/server/context_compression_test.go
@@ -0,0 +1,246 @@
+package server
+
+import (
+ "bytes"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+
+ "github.com/opencsgs/csglite/internal/config"
+ "github.com/opencsgs/csglite/internal/ctxcompress"
+ "github.com/opencsgs/csglite/internal/observability"
+ "github.com/opencsgs/csglite/pkg/api"
+)
+
+func contextCompressionTestBody(t *testing.T) (string, string) {
+ t.Helper()
+ var grep strings.Builder
+ for i := 1; i <= 30; i++ {
+ fmt.Fprintf(&grep, "internal/server/routes.go:%d:\tmux.HandleFunc(%d)\n", i, i)
+ }
+ body, err := json.Marshal(map[string]any{
+ "model": "test/model",
+ "messages": []any{
+ map[string]any{"role": "assistant", "tool_calls": []any{
+ map[string]any{"id": "call_1", "type": "function", "function": map[string]any{"name": "grep", "arguments": "{}"}},
+ }},
+ map[string]any{"role": "tool", "tool_call_id": "call_1", "content": grep.String()},
+ },
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ return string(body), grep.String()
+}
+
+func TestWithContextCompressionRewritesToolOutput(t *testing.T) {
+ s := newTestServer(t)
+ s.cfg.Inference.ContextCompression = config.ContextCompressionSafe
+ var seen []byte
+ var seenLength int64
+ handler := s.observabilityMiddleware(s.withContextCompression(ctxcompress.ProtocolOpenAIChat, func(w http.ResponseWriter, r *http.Request) {
+ seen, _ = io.ReadAll(r.Body)
+ seenLength = r.ContentLength
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write([]byte(`{"usage":{"prompt_tokens":10,"completion_tokens":1}}`))
+ }))
+ body, grep := contextCompressionTestBody(t)
+ req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(body))
+ recorder := httptest.NewRecorder()
+ handler.ServeHTTP(recorder, req)
+
+ if strings.Contains(string(seen), "routes.go:1:") || !strings.Contains(string(seen), `internal/server/routes.go\n 1:`) {
+ t.Fatalf("handler saw an uncompressed body: %s", seen)
+ }
+ if seenLength != int64(len(seen)) {
+ t.Fatalf("ContentLength = %d, body is %d bytes", seenLength, len(seen))
+ }
+ header := recorder.Header().Get(contextCompressionHeader)
+ if !strings.HasPrefix(header, "mode=safe; blocks=1; compressed=1; bytes=") {
+ t.Fatalf("%s = %q", contextCompressionHeader, header)
+ }
+
+ s.observabilityMu.RLock()
+ page, err := s.observability.ListRequests(req.Context(), observability.RequestFilter{})
+ s.observabilityMu.RUnlock()
+ if err != nil || page.Total != 1 {
+ t.Fatalf("captured %d requests: %v", page.Total, err)
+ }
+ record := page.Items[0]
+ if record.ContextCompressionMode != "safe" || record.ContextBlocks != 1 || record.ContextCompressedBlocks != 1 ||
+ record.ContextBytesBefore != int64(len(grep)) || record.ContextBytesAfter >= record.ContextBytesBefore ||
+ record.ContextTokensSaved <= 0 {
+ t.Fatalf("compression not recorded: %+v", record)
+ }
+ if page.Summary.ContextTokensSaved != record.ContextTokensSaved {
+ t.Fatalf("summary saved %d tokens, record %d", page.Summary.ContextTokensSaved, record.ContextTokensSaved)
+ }
+ s.observabilityMu.RLock()
+ detail, err := s.observability.GetRequest(req.Context(), record.ID)
+ s.observabilityMu.RUnlock()
+ if err != nil || !strings.Contains(detail.RequestBody, "routes.go:1:") {
+ t.Fatalf("stored request body is not the one the client sent: %v", err)
+ }
+ if resp := observabilityRequestResponse(detail); resp.ContextCompression == nil || resp.ContextCompression.Mode != "safe" {
+ t.Fatalf("API response context_compression = %+v", resp.ContextCompression)
+ }
+}
+
+func TestWithContextCompressionOffPassesBodyThrough(t *testing.T) {
+ s := newTestServer(t)
+ body, _ := contextCompressionTestBody(t)
+ for _, mode := range []string{"", config.ContextCompressionOff, "unknown"} {
+ s.cfg.Inference.ContextCompression = mode
+ var seen []byte
+ handler := s.withContextCompression(ctxcompress.ProtocolOpenAIChat, func(w http.ResponseWriter, r *http.Request) {
+ seen, _ = io.ReadAll(r.Body)
+ })
+ recorder := httptest.NewRecorder()
+ handler(recorder, httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(body)))
+ if string(seen) != body || recorder.Header().Get(contextCompressionHeader) != "" {
+ t.Fatalf("mode %q changed the request", mode)
+ }
+ }
+
+ // A body the compressor cannot parse reaches the handler untouched.
+ s.cfg.Inference.ContextCompression = config.ContextCompressionAggressive
+ var seen []byte
+ handler := s.withContextCompression(ctxcompress.ProtocolAnthropic, func(w http.ResponseWriter, r *http.Request) {
+ seen, _ = io.ReadAll(r.Body)
+ })
+ handler(httptest.NewRecorder(), httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader("{broken")))
+ if string(seen) != "{broken" {
+ t.Fatalf("broken body = %q", seen)
+ }
+}
+
+func TestHandleSettingsUpdatesContextCompression(t *testing.T) {
+ home := t.TempDir()
+ t.Setenv("HOME", home)
+ t.Setenv("USERPROFILE", home)
+ config.Reset()
+ s := newTestServer(t)
+
+ update := func(mode string) *httptest.ResponseRecorder {
+ body, err := json.Marshal(api.SettingsUpdateRequest{ContextCompression: &mode})
+ if err != nil {
+ t.Fatal(err)
+ }
+ w := httptest.NewRecorder()
+ s.handleSettingsUpdate(w, httptest.NewRequest(http.MethodPost, "/api/settings", bytes.NewReader(body)))
+ return w
+ }
+
+ w := update("Aggressive")
+ if w.Code != http.StatusOK || s.cfg.Inference.ContextCompression != config.ContextCompressionAggressive {
+ t.Fatalf("status %d, stored %q", w.Code, s.cfg.Inference.ContextCompression)
+ }
+ var resp api.SettingsResponse
+ if err := json.NewDecoder(w.Body).Decode(&resp); err != nil || resp.ContextCompression != "aggressive" {
+ t.Fatalf("response context_compression = %q (%v)", resp.ContextCompression, err)
+ }
+
+ if w := update("everything"); w.Code != http.StatusBadRequest {
+ t.Fatalf("invalid mode status = %d", w.Code)
+ }
+ if s.cfg.Inference.ContextCompression != config.ContextCompressionAggressive {
+ t.Fatal("invalid mode changed the setting")
+ }
+
+ if w := update("off"); w.Code != http.StatusOK || s.cfg.Inference.ContextCompression != "" {
+ t.Fatalf("off stored as %q", s.cfg.Inference.ContextCompression)
+ }
+ if got := s.settingsResponse().ContextCompression; got != "off" {
+ t.Fatalf("settings report %q, want off", got)
+ }
+}
+
+// TestContextCompressionEndToEnd sends agent requests in all three protocols
+// through the full router to a fake provider and checks what the provider
+// receives.
+func TestContextCompressionEndToEnd(t *testing.T) {
+ s := newTestServer(t)
+ var received []map[string]any
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ var payload map[string]any
+ if err := json.NewDecoder(r.Body).Decode(&payload); err != nil {
+ t.Errorf("decode upstream body: %v", err)
+ }
+ received = append(received, payload)
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write([]byte(`{"id":"chatcmpl-1","object":"chat.completion","model":"up-model","choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"usage":{"prompt_tokens":5,"completion_tokens":1}}`))
+ }))
+ defer upstream.Close()
+ if err := config.SaveProviders([]config.ThirdPartyProvider{
+ {ID: "up", Name: "Up", BaseURL: upstream.URL + "/v1", APIKey: "key", Enabled: true},
+ }); err != nil {
+ t.Fatal(err)
+ }
+
+ _, grep := contextCompressionTestBody(t)
+ tool := map[string]any{"type": "function", "function": map[string]any{"name": "grep", "parameters": map[string]any{"type": "object"}}}
+ requests := map[string]map[string]any{
+ "/v1/chat/completions": {
+ "model": "up-model", "source": "provider:up", "tools": []any{tool},
+ "messages": []any{
+ map[string]any{"role": "user", "content": "find handlers"},
+ map[string]any{"role": "assistant", "tool_calls": []any{map[string]any{"id": "call_1", "type": "function", "function": map[string]any{"name": "grep", "arguments": "{}"}}}},
+ map[string]any{"role": "tool", "tool_call_id": "call_1", "content": grep},
+ },
+ },
+ "/v1/messages": {
+ "model": "up-model", "source": "provider:up", "max_tokens": 100,
+ "tools": []any{map[string]any{"name": "Grep", "input_schema": map[string]any{"type": "object"}}},
+ "messages": []any{
+ map[string]any{"role": "user", "content": "find handlers"},
+ map[string]any{"role": "assistant", "content": []any{map[string]any{"type": "tool_use", "id": "toolu_1", "name": "Grep", "input": map[string]any{}}}},
+ map[string]any{"role": "user", "content": []any{map[string]any{"type": "tool_result", "tool_use_id": "toolu_1", "content": grep}}},
+ },
+ },
+ "/v1/responses": {
+ "model": "up-model", "source": "provider:up",
+ "tools": []any{map[string]any{"type": "function", "name": "shell", "parameters": map[string]any{"type": "object"}}},
+ "input": []any{
+ map[string]any{"type": "message", "role": "user", "content": "find handlers"},
+ map[string]any{"type": "function_call", "call_id": "call_9", "name": "shell", "arguments": "{}"},
+ map[string]any{"type": "function_call_output", "call_id": "call_9", "output": grep},
+ },
+ },
+ }
+
+ for _, mode := range []string{config.ContextCompressionOff, config.ContextCompressionSafe} {
+ s.cfg.Inference.ContextCompression = mode
+ for path, payload := range requests {
+ received = nil
+ body, err := json.Marshal(payload)
+ if err != nil {
+ t.Fatal(err)
+ }
+ req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(body))
+ req.Header.Set("Content-Type", "application/json")
+ w := httptest.NewRecorder()
+ s.routes().ServeHTTP(w, req)
+ if w.Code != http.StatusOK || len(received) != 1 {
+ t.Fatalf("%s %s: status %d, upstream calls %d, body %s", mode, path, w.Code, len(received), w.Body.String())
+ }
+ sent, _ := json.Marshal(received[0]["messages"])
+ grouped := strings.Contains(string(sent), `internal/server/routes.go\n 1:`)
+ original := strings.Contains(string(sent), `routes.go:1:\tmux`)
+ header := w.Header().Get(contextCompressionHeader)
+ switch mode {
+ case config.ContextCompressionSafe:
+ if !grouped || original || !strings.HasPrefix(header, "mode=safe; blocks=1; compressed=1") {
+ t.Fatalf("%s: provider got uncompressed tool output (header %q): %s", path, header, sent)
+ }
+ default:
+ if grouped || !original || header != "" {
+ t.Fatalf("%s: compression ran while off (header %q): %s", path, header, sent)
+ }
+ }
+ }
+ }
+}
diff --git a/internal/server/handlers_observability.go b/internal/server/handlers_observability.go
index a16f2f1f..e8cebfb0 100644
--- a/internal/server/handlers_observability.go
+++ b/internal/server/handlers_observability.go
@@ -42,6 +42,8 @@ func (s *Server) handleObservabilityRequests(w http.ResponseWriter, r *http.Requ
Failed: page.Summary.Failed,
TotalTokens: page.Summary.TotalTokens,
AverageLatency: page.Summary.AverageLatency,
+
+ ContextTokensSaved: page.Summary.ContextTokensSaved,
},
}
for _, record := range page.Items {
@@ -284,6 +286,21 @@ func observabilityRequestResponse(record observability.RequestRecord) api.Observ
ResponseBody: record.ResponseBody,
RequestBodyTruncated: record.RequestBodyTruncated,
ResponseBodyTruncated: record.ResponseBodyTruncated,
+ ContextCompression: observabilityContextCompression(record),
+ }
+}
+
+func observabilityContextCompression(record observability.RequestRecord) *api.ObservabilityContextCompression {
+ if record.ContextCompressionMode == "" {
+ return nil
+ }
+ return &api.ObservabilityContextCompression{
+ Mode: record.ContextCompressionMode,
+ Blocks: record.ContextBlocks,
+ Compressed: record.ContextCompressedBlocks,
+ BytesBefore: record.ContextBytesBefore,
+ BytesAfter: record.ContextBytesAfter,
+ TokensSaved: record.ContextTokensSaved,
}
}
diff --git a/internal/server/handlers_system.go b/internal/server/handlers_system.go
index b6f4c511..f9ed09d1 100644
--- a/internal/server/handlers_system.go
+++ b/internal/server/handlers_system.go
@@ -169,6 +169,18 @@ func (s *Server) handleSettingsUpdate(w http.ResponseWriter, r *http.Request) {
s.cfg.Inference.LlamaNumParallel = numParallel
configUpdated = true
}
+ if req.ContextCompression != nil {
+ if !config.IsContextCompressionMode(*req.ContextCompression) {
+ writeError(w, http.StatusBadRequest, `context_compression must be "off", "safe" or "aggressive"`)
+ return
+ }
+ mode := config.NormalizeContextCompression(*req.ContextCompression)
+ if mode == config.ContextCompressionOff {
+ mode = "" // off is the default, so it is not written out
+ }
+ s.cfg.Inference.ContextCompression = mode
+ configUpdated = true
+ }
if req.ServerURL != nil {
serverURL := strings.TrimSpace(*req.ServerURL)
if serverURL == "" {
@@ -347,6 +359,7 @@ func currentSettingsResponse(cfg *config.Config, version string) api.SettingsRes
},
LlamaUseModelMaxCtx: inference.UseModelMaxCtxByDefault(cfg.Inference.LlamaUseModelMaxCtx),
LlamaNumParallel: inference.ResolveNumParallel(cfg.Inference.LlamaNumParallel),
+ ContextCompression: config.NormalizeContextCompression(cfg.Inference.ContextCompression),
HiddenNavItems: append([]string{}, cfg.HiddenNavItems...),
}
}
diff --git a/internal/server/observability.go b/internal/server/observability.go
index cc3a8308..b5a40baa 100644
--- a/internal/server/observability.go
+++ b/internal/server/observability.go
@@ -18,6 +18,7 @@ import (
"github.com/opencsgs/csglite/internal/config"
"github.com/opencsgs/csglite/internal/correlation"
+ "github.com/opencsgs/csglite/internal/ctxcompress"
"github.com/opencsgs/csglite/internal/inference"
"github.com/opencsgs/csglite/internal/observability"
routerprofile "github.com/opencsgs/semantic-router"
@@ -42,6 +43,7 @@ type observationMetadata struct {
pool *apiUsagePoolMetadata
inputTokens int64
outputTokens int64
+ compression observationCompression
}
type observationMetadataSnapshot struct {
@@ -52,6 +54,33 @@ type observationMetadataSnapshot struct {
pool *apiUsagePoolMetadata
inputTokens int64
outputTokens int64
+ compression observationCompression
+}
+
+// observationCompression is what context compression did to the request.
+type observationCompression struct {
+ mode string
+ blocks int64
+ compressed int64
+ bytesBefore int64
+ bytesAfter int64
+ tokensSaved int64
+}
+
+func (m *observationMetadata) setContextCompression(mode string, stats ctxcompress.Stats) {
+ if m == nil {
+ return
+ }
+ m.mu.Lock()
+ defer m.mu.Unlock()
+ m.compression = observationCompression{
+ mode: mode,
+ blocks: int64(stats.Blocks),
+ compressed: int64(stats.Compressed),
+ bytesBefore: int64(stats.BytesBefore),
+ bytesAfter: int64(stats.BytesAfter),
+ tokensSaved: int64(stats.TokensSaved),
+ }
}
func observationFromContext(ctx context.Context) *observationMetadata {
@@ -90,6 +119,7 @@ func (m *observationMetadata) snapshot() observationMetadataSnapshot {
sourceName: m.sourceName,
inputTokens: m.inputTokens,
outputTokens: m.outputTokens,
+ compression: m.compression,
}
if m.pool != nil {
pool := *m.pool
@@ -283,6 +313,13 @@ func (s *Server) observabilityMiddleware(next http.Handler) http.Handler {
ResponseBody: string(responseBody),
RequestBodyTruncated: requestTruncated,
ResponseBodyTruncated: ow.truncated,
+
+ ContextCompressionMode: snapshot.compression.mode,
+ ContextBlocks: snapshot.compression.blocks,
+ ContextCompressedBlocks: snapshot.compression.compressed,
+ ContextBytesBefore: snapshot.compression.bytesBefore,
+ ContextBytesAfter: snapshot.compression.bytesAfter,
+ ContextTokensSaved: snapshot.compression.tokensSaved,
}
if !ow.firstWrite.IsZero() {
record.FirstTokenLatencyMS = ow.firstWrite.Sub(startedAt).Milliseconds()
diff --git a/internal/server/provider_routes.go b/internal/server/provider_routes.go
index 8eaae5ea..81cf3b0a 100644
--- a/internal/server/provider_routes.go
+++ b/internal/server/provider_routes.go
@@ -9,6 +9,7 @@ import (
"github.com/opencsgs/csglite/ee/cluster"
"github.com/opencsgs/csglite/internal/config"
+ "github.com/opencsgs/csglite/internal/ctxcompress"
"github.com/opencsgs/csglite/internal/inference"
"github.com/opencsgs/csglite/pkg/api"
)
@@ -28,13 +29,13 @@ func (s *Server) registerProviderInferenceRoutes(mux *http.ServeMux) {
register("GET /providers/{providerID}/v1/models", s.handleModels)
register("GET /providers/{providerID}/v1/responses", s.handleOpenAIResponsesUnsupported)
- register("POST /providers/{providerID}/v1/chat/completions", s.handleOpenAIChatCompletions)
+ register("POST /providers/{providerID}/v1/chat/completions", s.withContextCompression(ctxcompress.ProtocolOpenAIChat, s.handleOpenAIChatCompletions))
register("POST /providers/{providerID}/v1/embeddings", s.handleOpenAIEmbeddings)
register("POST /providers/{providerID}/v1/images/generations", s.handleOpenAIImagesGenerations)
register("POST /providers/{providerID}/v1/images/edits", s.handleOpenAIImagesEdits)
register("POST /providers/{providerID}/v1/audio/transcriptions", s.handleOpenAIAudioTranscriptions)
- register("POST /providers/{providerID}/v1/responses", s.handleOpenAIResponses)
- register("POST /providers/{providerID}/v1/messages", s.handleAnthropicMessages)
+ register("POST /providers/{providerID}/v1/responses", s.withContextCompression(ctxcompress.ProtocolResponses, s.handleOpenAIResponses))
+ register("POST /providers/{providerID}/v1/messages", s.withContextCompression(ctxcompress.ProtocolAnthropic, s.handleAnthropicMessages))
register("POST /providers/{providerID}/v1/messages/count_tokens", s.handleAnthropicCountTokens)
notFound := func(w http.ResponseWriter, _ *http.Request) {
diff --git a/internal/server/routes.go b/internal/server/routes.go
index a6bc8da7..8f3754e3 100644
--- a/internal/server/routes.go
+++ b/internal/server/routes.go
@@ -4,6 +4,7 @@ import (
"net/http"
"github.com/opencsgs/csglite/ee/cluster"
+ "github.com/opencsgs/csglite/internal/ctxcompress"
)
func (s *Server) routes() http.Handler {
@@ -59,7 +60,7 @@ func (s *Server) routes() http.Handler {
mux.HandleFunc("DELETE /api/datasets/pull/partial", s.handlePartialDatasetPullDelete)
mux.HandleFunc("DELETE /api/datasets/delete", s.handleDatasetDelete)
- mux.HandleFunc("POST /v1/chat/completions", s.handleOpenAIChatCompletions)
+ mux.HandleFunc("POST /v1/chat/completions", s.withContextCompression(ctxcompress.ProtocolOpenAIChat, s.handleOpenAIChatCompletions))
mux.HandleFunc("POST /v1/embeddings", s.handleOpenAIEmbeddings)
mux.HandleFunc("POST /v1/images/generations", s.handleOpenAIImagesGenerations)
mux.HandleFunc("POST /v1/images/edits", s.handleOpenAIImagesEdits)
@@ -77,12 +78,12 @@ func (s *Server) routes() http.Handler {
mux.HandleFunc("DELETE /api/images/jobs/{jobID}", s.handleImageGenerationJobCancel)
mux.HandleFunc("GET /v1/models", s.handleModels)
mux.HandleFunc("GET /v1/responses", s.handleOpenAIResponsesUnsupported)
- mux.HandleFunc("POST /v1/responses", s.handleOpenAIResponses)
- mux.HandleFunc("POST /v1/messages", s.handleAnthropicMessages)
+ mux.HandleFunc("POST /v1/responses", s.withContextCompression(ctxcompress.ProtocolResponses, s.handleOpenAIResponses))
+ mux.HandleFunc("POST /v1/messages", s.withContextCompression(ctxcompress.ProtocolAnthropic, s.handleAnthropicMessages))
mux.HandleFunc("POST /v1/messages/count_tokens", s.handleAnthropicCountTokens)
- mux.HandleFunc("POST /anthropic/messages", s.handleAnthropicMessages)
+ mux.HandleFunc("POST /anthropic/messages", s.withContextCompression(ctxcompress.ProtocolAnthropic, s.handleAnthropicMessages))
mux.HandleFunc("POST /anthropic/messages/count_tokens", s.handleAnthropicCountTokens)
- mux.HandleFunc("POST /anthropic/v1/messages", s.handleAnthropicMessages)
+ mux.HandleFunc("POST /anthropic/v1/messages", s.withContextCompression(ctxcompress.ProtocolAnthropic, s.handleAnthropicMessages))
mux.HandleFunc("POST /anthropic/v1/messages/count_tokens", s.handleAnthropicCountTokens)
s.registerProviderInferenceRoutes(mux)
@@ -244,17 +245,17 @@ func (s *Server) externalAPIRoutes() http.Handler {
mux.HandleFunc("GET /v1/models", s.handleModels)
mux.HandleFunc("GET /v1/responses", s.handleOpenAIResponsesUnsupported)
- mux.HandleFunc("POST /v1/chat/completions", s.handleOpenAIChatCompletions)
+ mux.HandleFunc("POST /v1/chat/completions", s.withContextCompression(ctxcompress.ProtocolOpenAIChat, s.handleOpenAIChatCompletions))
mux.HandleFunc("POST /v1/embeddings", s.handleOpenAIEmbeddings)
mux.HandleFunc("POST /v1/images/generations", s.handleOpenAIImagesGenerations)
mux.HandleFunc("POST /v1/images/edits", s.handleOpenAIImagesEdits)
mux.HandleFunc("POST /v1/audio/transcriptions", s.handleOpenAIAudioTranscriptions)
- mux.HandleFunc("POST /v1/responses", s.handleOpenAIResponses)
- mux.HandleFunc("POST /v1/messages", s.handleAnthropicMessages)
+ mux.HandleFunc("POST /v1/responses", s.withContextCompression(ctxcompress.ProtocolResponses, s.handleOpenAIResponses))
+ mux.HandleFunc("POST /v1/messages", s.withContextCompression(ctxcompress.ProtocolAnthropic, s.handleAnthropicMessages))
mux.HandleFunc("POST /v1/messages/count_tokens", s.handleAnthropicCountTokens)
- mux.HandleFunc("POST /anthropic/messages", s.handleAnthropicMessages)
+ mux.HandleFunc("POST /anthropic/messages", s.withContextCompression(ctxcompress.ProtocolAnthropic, s.handleAnthropicMessages))
mux.HandleFunc("POST /anthropic/messages/count_tokens", s.handleAnthropicCountTokens)
- mux.HandleFunc("POST /anthropic/v1/messages", s.handleAnthropicMessages)
+ mux.HandleFunc("POST /anthropic/v1/messages", s.withContextCompression(ctxcompress.ProtocolAnthropic, s.handleAnthropicMessages))
mux.HandleFunc("POST /anthropic/v1/messages/count_tokens", s.handleAnthropicCountTokens)
s.registerProviderInferenceRoutes(mux)
diff --git a/internal/server/static/openapi/local-api.json b/internal/server/static/openapi/local-api.json
index 95c72a85..c7158c3d 100644
--- a/internal/server/static/openapi/local-api.json
+++ b/internal/server/static/openapi/local-api.json
@@ -9115,6 +9115,7 @@
"autostart",
"desktop_mode",
"llama_use_model_max_ctx",
+ "context_compression",
"web_search",
"observability",
"hidden_nav_items",
@@ -9197,6 +9198,11 @@
"type": "boolean",
"description": "Whether local models use their native maximum context when a request does not explicitly set num_ctx. CSGHUB_LITE_LLAMA_USE_MODEL_MAX_CTX takes precedence when set."
},
+ "context_compression": {
+ "type": "string",
+ "enum": ["off", "safe", "aggressive"],
+ "description": "How the gateway shrinks tool output (tool_result blocks, role=tool messages, function_call_output items) in /v1/messages, /v1/chat/completions and /v1/responses requests before a model sees them. off leaves requests untouched. safe removes only redundancy and keeps every distinct line: JSON is minified and arrays of records become a header plus rows, search results are grouped by file, runs of identical lines are counted, terminal escapes are dropped, lines over 2000 bytes are cut, and a log, search or JSON result identical to an earlier one points back to it. aggressive also drops content: long JSON arrays, command output over 150 lines and search results over 12 matches per file are sampled, and lines differing only in their numbers are folded; every error line, boundaries and outliers are kept. File reads and source code are never changed. Affected responses carry an X-CSGLite-Context-Compression header."
+ },
"web_search": {
"$ref": "#/components/schemas/WebSearchSettings"
},
@@ -9310,6 +9316,11 @@
"type": "boolean",
"description": "Use each local model's native maximum context when requests omit num_ctx."
},
+ "context_compression": {
+ "type": "string",
+ "enum": ["off", "safe", "aggressive"],
+ "description": "Context compression mode for agent requests; see SettingsResponse.context_compression."
+ },
"web_search": {
"$ref": "#/components/schemas/WebSearchSettings"
},
@@ -9830,7 +9841,21 @@
"request_body": {"type": "string"},
"response_body": {"type": "string"},
"request_body_truncated": {"type": "boolean"},
- "response_body_truncated": {"type": "boolean"}
+ "response_body_truncated": {"type": "boolean"},
+ "context_compression": {"$ref": "#/components/schemas/ObservabilityContextCompression"}
+ }
+ },
+ "ObservabilityContextCompression": {
+ "type": "object",
+ "description": "What context compression did to the request's tool results. Present only when compression was on and the request carried tool results. request_body keeps the body as the client sent it.",
+ "required": ["mode", "blocks", "compressed", "bytes_before", "bytes_after", "tokens_saved"],
+ "properties": {
+ "mode": {"type": "string", "enum": ["safe", "aggressive"]},
+ "blocks": {"type": "integer", "format": "int64", "description": "Tool results in the request."},
+ "compressed": {"type": "integer", "format": "int64", "description": "Tool results that were rewritten."},
+ "bytes_before": {"type": "integer", "format": "int64"},
+ "bytes_after": {"type": "integer", "format": "int64"},
+ "tokens_saved": {"type": "integer", "format": "int64", "description": "Estimated input tokens removed."}
}
},
"ObservabilityRequestSummary": {
@@ -9840,7 +9865,8 @@
"succeeded": {"type": "integer", "format": "int64"},
"failed": {"type": "integer", "format": "int64"},
"total_tokens": {"type": "integer", "format": "int64"},
- "average_latency_ms": {"type": "number", "format": "double"}
+ "average_latency_ms": {"type": "number", "format": "double"},
+ "context_tokens_saved": {"type": "integer", "format": "int64", "description": "Estimated input tokens context compression removed across the matching requests."}
}
},
"ObservabilityTrace": {
diff --git a/openapi/local-api.json b/openapi/local-api.json
index 95c72a85..c7158c3d 100644
--- a/openapi/local-api.json
+++ b/openapi/local-api.json
@@ -9115,6 +9115,7 @@
"autostart",
"desktop_mode",
"llama_use_model_max_ctx",
+ "context_compression",
"web_search",
"observability",
"hidden_nav_items",
@@ -9197,6 +9198,11 @@
"type": "boolean",
"description": "Whether local models use their native maximum context when a request does not explicitly set num_ctx. CSGHUB_LITE_LLAMA_USE_MODEL_MAX_CTX takes precedence when set."
},
+ "context_compression": {
+ "type": "string",
+ "enum": ["off", "safe", "aggressive"],
+ "description": "How the gateway shrinks tool output (tool_result blocks, role=tool messages, function_call_output items) in /v1/messages, /v1/chat/completions and /v1/responses requests before a model sees them. off leaves requests untouched. safe removes only redundancy and keeps every distinct line: JSON is minified and arrays of records become a header plus rows, search results are grouped by file, runs of identical lines are counted, terminal escapes are dropped, lines over 2000 bytes are cut, and a log, search or JSON result identical to an earlier one points back to it. aggressive also drops content: long JSON arrays, command output over 150 lines and search results over 12 matches per file are sampled, and lines differing only in their numbers are folded; every error line, boundaries and outliers are kept. File reads and source code are never changed. Affected responses carry an X-CSGLite-Context-Compression header."
+ },
"web_search": {
"$ref": "#/components/schemas/WebSearchSettings"
},
@@ -9310,6 +9316,11 @@
"type": "boolean",
"description": "Use each local model's native maximum context when requests omit num_ctx."
},
+ "context_compression": {
+ "type": "string",
+ "enum": ["off", "safe", "aggressive"],
+ "description": "Context compression mode for agent requests; see SettingsResponse.context_compression."
+ },
"web_search": {
"$ref": "#/components/schemas/WebSearchSettings"
},
@@ -9830,7 +9841,21 @@
"request_body": {"type": "string"},
"response_body": {"type": "string"},
"request_body_truncated": {"type": "boolean"},
- "response_body_truncated": {"type": "boolean"}
+ "response_body_truncated": {"type": "boolean"},
+ "context_compression": {"$ref": "#/components/schemas/ObservabilityContextCompression"}
+ }
+ },
+ "ObservabilityContextCompression": {
+ "type": "object",
+ "description": "What context compression did to the request's tool results. Present only when compression was on and the request carried tool results. request_body keeps the body as the client sent it.",
+ "required": ["mode", "blocks", "compressed", "bytes_before", "bytes_after", "tokens_saved"],
+ "properties": {
+ "mode": {"type": "string", "enum": ["safe", "aggressive"]},
+ "blocks": {"type": "integer", "format": "int64", "description": "Tool results in the request."},
+ "compressed": {"type": "integer", "format": "int64", "description": "Tool results that were rewritten."},
+ "bytes_before": {"type": "integer", "format": "int64"},
+ "bytes_after": {"type": "integer", "format": "int64"},
+ "tokens_saved": {"type": "integer", "format": "int64", "description": "Estimated input tokens removed."}
}
},
"ObservabilityRequestSummary": {
@@ -9840,7 +9865,8 @@
"succeeded": {"type": "integer", "format": "int64"},
"failed": {"type": "integer", "format": "int64"},
"total_tokens": {"type": "integer", "format": "int64"},
- "average_latency_ms": {"type": "number", "format": "double"}
+ "average_latency_ms": {"type": "number", "format": "double"},
+ "context_tokens_saved": {"type": "integer", "format": "int64", "description": "Estimated input tokens context compression removed across the matching requests."}
}
},
"ObservabilityTrace": {
diff --git a/pkg/api/types.go b/pkg/api/types.go
index 8ba835f1..21c253ce 100644
--- a/pkg/api/types.go
+++ b/pkg/api/types.go
@@ -548,7 +548,9 @@ type SettingsResponse struct {
Observability ObservabilitySettings `json:"observability"`
LlamaUseModelMaxCtx bool `json:"llama_use_model_max_ctx"`
LlamaNumParallel int `json:"llama_num_parallel"`
- HiddenNavItems []string `json:"hidden_nav_items"`
+ // ContextCompression is "off", "safe" or "aggressive".
+ ContextCompression string `json:"context_compression"`
+ HiddenNavItems []string `json:"hidden_nav_items"`
// Edition is "Enterprise" while a license is in effect, otherwise "Community".
Edition string `json:"edition"`
// LicenseStatus is the license.Status string; "none" when no license is installed.
@@ -640,6 +642,7 @@ type SettingsUpdateRequest struct {
Observability *ObservabilitySettings `json:"observability,omitempty"`
LlamaUseModelMaxCtx *bool `json:"llama_use_model_max_ctx,omitempty"`
LlamaNumParallel *int `json:"llama_num_parallel,omitempty"`
+ ContextCompression *string `json:"context_compression,omitempty"`
}
type ObservabilitySettings struct {
@@ -709,6 +712,20 @@ type ObservabilityRequest struct {
ResponseBody string `json:"response_body,omitempty"`
RequestBodyTruncated bool `json:"request_body_truncated"`
ResponseBodyTruncated bool `json:"response_body_truncated"`
+ // ContextCompression is set when context compression was on for the
+ // request and it carried tool results.
+ ContextCompression *ObservabilityContextCompression `json:"context_compression,omitempty"`
+}
+
+// ObservabilityContextCompression is what context compression did to one
+// request's tool results. TokensSaved is an estimate.
+type ObservabilityContextCompression struct {
+ Mode string `json:"mode"`
+ Blocks int64 `json:"blocks"`
+ Compressed int64 `json:"compressed"`
+ BytesBefore int64 `json:"bytes_before"`
+ BytesAfter int64 `json:"bytes_after"`
+ TokensSaved int64 `json:"tokens_saved"`
}
type ObservabilityRequestSummary struct {
@@ -717,6 +734,9 @@ type ObservabilityRequestSummary struct {
Failed int64 `json:"failed"`
TotalTokens int64 `json:"total_tokens"`
AverageLatency float64 `json:"average_latency_ms"`
+ // ContextTokensSaved estimates the input tokens context compression
+ // removed across the matching requests.
+ ContextTokensSaved int64 `json:"context_tokens_saved"`
}
type ObservabilityRequestListResponse struct {
diff --git a/web/src/api/client.ts b/web/src/api/client.ts
index 6c7472d6..f24647b2 100644
--- a/web/src/api/client.ts
+++ b/web/src/api/client.ts
@@ -308,6 +308,9 @@ export interface SystemInfo {
gpu_shared_memory: boolean;
}
+/** How the gateway shrinks tool output in agent requests. */
+export type ContextCompressionMode = "off" | "safe" | "aggressive";
+
export interface AppSettings {
version: string;
llama_server_version?: string;
@@ -331,6 +334,7 @@ export interface AppSettings {
autostart: boolean;
llama_use_model_max_ctx: boolean;
llama_num_parallel: number;
+ context_compression: ContextCompressionMode;
web_search: WebSearchSettings;
observability: ObservabilitySettings;
hidden_nav_items: string[];
@@ -755,6 +759,18 @@ export interface ObservabilityRequest {
response_body?: string;
request_body_truncated: boolean;
response_body_truncated: boolean;
+ context_compression?: ObservabilityContextCompression;
+}
+
+/** What context compression did to one request's tool results. */
+export interface ObservabilityContextCompression {
+ mode: ContextCompressionMode;
+ blocks: number;
+ compressed: number;
+ bytes_before: number;
+ bytes_after: number;
+ /** Estimated input tokens removed. */
+ tokens_saved: number;
}
export interface ObservabilityRequestSummary {
@@ -763,6 +779,7 @@ export interface ObservabilityRequestSummary {
failed: number;
total_tokens: number;
average_latency_ms: number;
+ context_tokens_saved: number;
}
export interface ObservabilityRequestListResponse {
@@ -1717,6 +1734,7 @@ export async function saveSettings(patch: {
autostart?: boolean;
llama_use_model_max_ctx?: boolean;
llama_num_parallel?: number;
+ context_compression?: ContextCompressionMode;
web_search?: WebSearchSettings;
observability?: ObservabilitySettings;
}): Promise {t("settings.contextCompressionDesc")}