fix cancellation and configure compaction model
This commit is contained in:
parent
159eb9c857
commit
e4c24e8a5f
|
|
@ -0,0 +1,142 @@
|
|||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"ollie/cmd/olliesrv/internal/backend"
|
||||
olog "ollie/log"
|
||||
)
|
||||
|
||||
// noOpLogger is a logger that discards all output.
|
||||
var noOpLogger = olog.NewWriter("test", olog.LevelError+1, nil, nil)
|
||||
|
||||
// blockingBackend implements backend.Backend but blocks on ChatStream until ctx is cancelled.
|
||||
type blockingBackend struct {
|
||||
chatStreamBlocked atomic.Bool
|
||||
model string
|
||||
ctxLength int
|
||||
}
|
||||
|
||||
func (b *blockingBackend) ChatStream(ctx context.Context, messages []backend.Message, tools []backend.Tool, params backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
||||
b.chatStreamBlocked.Store(true)
|
||||
ch := make(chan backend.StreamEvent, 1)
|
||||
go func() {
|
||||
defer close(ch)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
}()
|
||||
return ch, nil
|
||||
}
|
||||
func (b *blockingBackend) Name() string { return "test" }
|
||||
func (b *blockingBackend) DefaultModel() string { return "test-model" }
|
||||
func (b *blockingBackend) Model() string { return b.model }
|
||||
func (b *blockingBackend) SetModel(model string) { b.model = model }
|
||||
func (b *blockingBackend) ContextLength(_ context.Context) int { return b.ctxLength }
|
||||
func (b *blockingBackend) Models(_ context.Context) []string { return []string{"test-model"} }
|
||||
|
||||
// TestStripColdCancellation verifies that /stop (Interrupt) cancels stripCold's
|
||||
// blocking LLM call promptly instead of waiting for it to complete.
|
||||
func TestStripColdCancellation(t *testing.T) {
|
||||
b := &blockingBackend{model: "test", ctxLength: 128000}
|
||||
h := newHistory("test input")
|
||||
|
||||
// Add enough messages so stripCold has work to do.
|
||||
// stripCold processes messages before index len(messages)-hotTailSize (8).
|
||||
// Need >8 cold-zone tool results (>200 chars each).
|
||||
for i := 0; i < 10; i++ {
|
||||
h.messages = append(h.messages, backend.Message{
|
||||
Role: "tool",
|
||||
Content: fmt.Sprintf("tool-%d: %s", i, strings.Repeat("x", 300)),
|
||||
})
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
h.stripCold(ctx, b)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Wait for the backend to be blocked.
|
||||
timeout := time.After(2 * time.Second)
|
||||
for !b.chatStreamBlocked.Load() {
|
||||
select {
|
||||
case <-done:
|
||||
t.Fatal("stripCold exited before backend was blocked")
|
||||
case <-timeout:
|
||||
t.Fatal("timed out waiting for backend to block")
|
||||
case <-time.After(10 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
// Cancel via the same mechanism as Interrupt — this should unblock stripCold.
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
// Good — stripCold returned promptly after cancellation.
|
||||
case <-time.After(1 * time.Second):
|
||||
t.Fatal("stripCold did not return within 1s of context cancellation")
|
||||
}
|
||||
}
|
||||
|
||||
// TestPreRunCompactionCancellation verifies that /stop cancels pre-run auto-compaction
|
||||
// which uses runCompact → history.compact → ChatStream.
|
||||
func TestPreRunCompactionCancellation(t *testing.T) {
|
||||
b := &blockingBackend{model: "test", ctxLength: 128000}
|
||||
|
||||
// Create an agent with enough history to trigger compaction (>hotTailSize+warmIndexSize=18).
|
||||
ag := &Agent{
|
||||
runtime: &Runtime{
|
||||
Backend: b,
|
||||
},
|
||||
history: newHistory("test"),
|
||||
log: noOpLogger,
|
||||
}
|
||||
|
||||
// Add enough messages to exceed the compact threshold.
|
||||
for i := 0; i < 30; i++ {
|
||||
ag.history.appendUserMessage(strings.Repeat("x", 500))
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
_, err := ag.runCompact(ctx, "test")
|
||||
if err != nil && !strings.Contains(err.Error(), "auto-compact") {
|
||||
t.Logf("runCompact error (expected): %v", err)
|
||||
}
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Wait for the backend to be blocked.
|
||||
timeout := time.After(2 * time.Second)
|
||||
for !b.chatStreamBlocked.Load() {
|
||||
select {
|
||||
case <-done:
|
||||
t.Fatal("runCompact exited before backend was blocked")
|
||||
case <-timeout:
|
||||
t.Fatal("timed out waiting for backend to block")
|
||||
case <-time.After(10 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
// Cancel — this should unblock the compaction.
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
// Good — runCompact returned promptly after cancellation.
|
||||
case <-time.After(1 * time.Second):
|
||||
t.Fatal("runCompact did not return within 1s of context cancellation")
|
||||
}
|
||||
}
|
||||
|
|
@ -9,8 +9,6 @@ import (
|
|||
|
||||
"ollie/cmd/olliesrv/internal/backend"
|
||||
"ollie/paths"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
const (
|
||||
|
|
@ -391,9 +389,9 @@ func (s *History) stripCold(ctx context.Context, b backend.Backend) {
|
|||
|
||||
// Collect cold-zone tool results exceeding 200 chars.
|
||||
type item struct {
|
||||
id int
|
||||
idx int // index into s.messages
|
||||
orig string
|
||||
id int
|
||||
idx int // index into s.messages
|
||||
orig string
|
||||
}
|
||||
var items []item
|
||||
for i := range s.messages {
|
||||
|
|
@ -442,8 +440,7 @@ func (s *History) stripCold(ctx context.Context, b backend.Backend) {
|
|||
}
|
||||
|
||||
// resolveCompactionModel returns the model to use for compaction.
|
||||
// Priority: agent config > OLLIE_COMPACTION_MODEL env > per-backend default from models.yaml > session's current model.
|
||||
// Returns "" if no override is configured (use the session's current model).
|
||||
// Priority: agent config > OLLIE_COMPACTION_MODEL > backends.conf compactionModel > backends.conf model > current model.
|
||||
func resolveCompactionModel(cfgModel string, b backend.Backend) string {
|
||||
if cfgModel != "" {
|
||||
return cfgModel
|
||||
|
|
@ -451,34 +448,11 @@ func resolveCompactionModel(cfgModel string, b backend.Backend) string {
|
|||
if env := os.Getenv("OLLIE_COMPACTION_MODEL"); env != "" {
|
||||
return env
|
||||
}
|
||||
cfg := loadModelsConfig()
|
||||
if b != nil {
|
||||
if m, ok := cfg.Compaction[b.Name()]; ok {
|
||||
return m
|
||||
if model := backend.CompactionModel(b.Name()); model != "" {
|
||||
return model
|
||||
}
|
||||
return b.Model()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// modelsConfig holds the parsed contents of ~/.config/ollie/models.yaml.
|
||||
type modelsConfig struct {
|
||||
Compaction map[string]string `yaml:"compaction"`
|
||||
Completion struct {
|
||||
Model string `yaml:"model"`
|
||||
Backend string `yaml:"backend"`
|
||||
} `yaml:"completion"`
|
||||
}
|
||||
|
||||
// loadModelsConfig reads and parses the models config file.
|
||||
// Returns zero value if the file doesn't exist or is invalid.
|
||||
func loadModelsConfig() modelsConfig {
|
||||
data, err := os.ReadFile(paths.CfgDir() + "/models.yaml")
|
||||
if err != nil {
|
||||
return modelsConfig{}
|
||||
}
|
||||
var cfg modelsConfig
|
||||
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
||||
return modelsConfig{}
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
|
|
|
|||
|
|
@ -201,7 +201,11 @@ func (ag *Agent) run(ctx context.Context) error {
|
|||
}
|
||||
}
|
||||
|
||||
ag.autoCompact(ctx)
|
||||
_, compactErr := ag.autoCompact(ctx)
|
||||
if compactErr != nil {
|
||||
ag.emit(Event{Role: "error", Content: compactErr.Error()})
|
||||
break
|
||||
}
|
||||
|
||||
step++
|
||||
}
|
||||
|
|
@ -210,20 +214,23 @@ func (ag *Agent) run(ctx context.Context) error {
|
|||
}
|
||||
|
||||
// autoCompact triggers context compaction if above threshold.
|
||||
func (ag *Agent) autoCompact(ctx context.Context) {
|
||||
// Returns the number of messages compacted, or an error if compaction failed.
|
||||
func (ag *Agent) autoCompact(ctx context.Context) (int, error) {
|
||||
if ctx.Err() != nil || ag.history == nil {
|
||||
return
|
||||
return 0, nil
|
||||
}
|
||||
limit := ag.autoCompactLimit(ctx)
|
||||
if limit <= 0 || ag.history.estimateTokens() < limit {
|
||||
return
|
||||
return 0, nil
|
||||
}
|
||||
ag.emit(Event{Role: "info", Content: "auto-compacting context...\n"})
|
||||
ag.SetState("compacting")
|
||||
if _, err := ag.runCompact(ctx, "auto"); err != nil {
|
||||
panic(fmt.Sprintf("mid-turn auto-compact: %v", err))
|
||||
n, err := ag.runCompact(ctx, "auto")
|
||||
if err != nil {
|
||||
return n, fmt.Errorf("auto-compact: %w", err)
|
||||
}
|
||||
ag.SetState("thinking")
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// ── streamResponse ──────────────────────────────────────────────────────────
|
||||
|
|
|
|||
|
|
@ -180,14 +180,14 @@ func (ag *Agent) executeTurn(ctx context.Context, input string) string {
|
|||
// Warn once when context usage crosses 60%; compact at 75%.
|
||||
if ag.history != nil {
|
||||
tokens := ag.history.estimateTokens()
|
||||
if compactLimit := ag.autoCompactLimit(ctx); compactLimit > 0 && tokens >= compactLimit {
|
||||
if compactLimit := ag.autoCompactLimit(actCtx); compactLimit > 0 && tokens >= compactLimit {
|
||||
ag.emit(Event{Role: "info", Content: "auto-compacting context...\n"})
|
||||
ag.SetState("compacting")
|
||||
if _, err := ag.runCompact(ctx, "auto"); err != nil {
|
||||
panic(fmt.Sprintf("auto-compact: %v", err))
|
||||
if _, err := ag.runCompact(actCtx, "auto"); err != nil {
|
||||
ag.log.Error("auto-compact failed: %v", err)
|
||||
}
|
||||
ag.SetState("thinking")
|
||||
} else if warnLimit := ag.autoWarnLimit(ctx); warnLimit > 0 && tokens >= warnLimit && !ag.warnedContext {
|
||||
} else if warnLimit := ag.autoWarnLimit(actCtx); warnLimit > 0 && tokens >= warnLimit && !ag.warnedContext {
|
||||
ctxLen := ag.runtime.Backend.ContextLength(ctx)
|
||||
if ctxLen <= 0 {
|
||||
ctxLen = defaultContextLength
|
||||
|
|
@ -204,7 +204,7 @@ func (ag *Agent) executeTurn(ctx context.Context, input string) string {
|
|||
|
||||
// Strip cold-zone tool results (one batched LLM call per turn).
|
||||
if ag.history != nil {
|
||||
ag.history.stripCold(ctx, ag.runtime.Backend)
|
||||
ag.history.stripCold(actCtx, ag.runtime.Backend)
|
||||
}
|
||||
|
||||
// Inject the preamble system message into history before running.
|
||||
|
|
@ -234,7 +234,7 @@ func (ag *Agent) executeTurn(ctx context.Context, input string) string {
|
|||
ag.history.messages = ag.history.messages[:preTurnLen]
|
||||
ag.emit(Event{Role: "info", Content: "context overflow — compacting and retrying...\n"})
|
||||
ag.SetState("compacting")
|
||||
if _, cerr := ag.runCompact(ctx, "overflow"); cerr != nil {
|
||||
if _, cerr := ag.runCompact(actCtx, "overflow"); cerr != nil {
|
||||
break
|
||||
}
|
||||
ag.history.appendUserMessage(input)
|
||||
|
|
|
|||
|
|
@ -10,10 +10,11 @@ import (
|
|||
|
||||
// BackendConfig holds the per-backend configuration from backends.conf.
|
||||
type BackendConfig struct {
|
||||
Key string // API key / token
|
||||
URL string // base URL override (empty = use backend's default)
|
||||
Model string // default model override
|
||||
Host string // for ollama: host address
|
||||
Key string // API key / token
|
||||
URL string // base URL override (empty = use backend's default)
|
||||
Model string // default model override
|
||||
CompactionModel string // model used for history compaction (empty = use Model)
|
||||
Host string // for ollama: host address
|
||||
}
|
||||
|
||||
// configFile is the parsed contents of backends.conf: section name → key/value pairs.
|
||||
|
|
@ -28,6 +29,18 @@ func loadConfig() configFile {
|
|||
return parseConfigFile(paths.CfgDir() + "/backends.conf")
|
||||
}
|
||||
|
||||
// CompactionModel returns the configured compaction model for a backend.
|
||||
// It falls back to that backend's configured default model when no dedicated
|
||||
// compaction model is set.
|
||||
func CompactionModel(name string) string {
|
||||
cfg := loadConfig()
|
||||
bc := cfg.Backends[name]
|
||||
if bc.CompactionModel != "" {
|
||||
return bc.CompactionModel
|
||||
}
|
||||
return bc.Model
|
||||
}
|
||||
|
||||
// parseConfigFile parses an INI-style config file.
|
||||
// Format:
|
||||
//
|
||||
|
|
@ -89,6 +102,8 @@ func parseConfigFile(path string) configFile {
|
|||
bc.URL = v
|
||||
case "model":
|
||||
bc.Model = v
|
||||
case "compactionModel", "compaction_model", "compaction-model":
|
||||
bc.CompactionModel = v
|
||||
case "host":
|
||||
bc.Host = v
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,26 @@
|
|||
package backend
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseConfigFileCompactionModel(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "backends.conf")
|
||||
content := "backend = openrouter\n\n[openrouter]\nmodel = primary\ncompactionModel = compact\n\n[ollama]\nmodel = local\n"
|
||||
if err := os.WriteFile(path, []byte(content), 0600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
cfg := parseConfigFile(path)
|
||||
if got := cfg.Backends["openrouter"].CompactionModel; got != "compact" {
|
||||
t.Fatalf("openrouter compactionModel = %q, want %q", got, "compact")
|
||||
}
|
||||
if got := cfg.Backends["ollama"].CompactionModel; got != "" {
|
||||
t.Fatalf("ollama compactionModel = %q, want empty", got)
|
||||
}
|
||||
if got := cfg.Backends["ollama"].Model; got != "local" {
|
||||
t.Fatalf("ollama model = %q, want %q", got, "local")
|
||||
}
|
||||
}
|
||||
|
|
@ -9,6 +9,7 @@
|
|||
# key = API key or token
|
||||
# url = base URL override (optional; each backend has a sensible default)
|
||||
# model = default model (optional; overridden by session model= parameter)
|
||||
# compactionModel = model used for history compaction (optional; defaults to model)
|
||||
# host = host address (ollama only)
|
||||
|
||||
backend = openrouter
|
||||
|
|
|
|||
Loading…
Reference in New Issue