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/cmd/olliesrv/internal/backend"
|
||||||
"ollie/paths"
|
"ollie/paths"
|
||||||
|
|
||||||
"gopkg.in/yaml.v3"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|
@ -391,9 +389,9 @@ func (s *History) stripCold(ctx context.Context, b backend.Backend) {
|
||||||
|
|
||||||
// Collect cold-zone tool results exceeding 200 chars.
|
// Collect cold-zone tool results exceeding 200 chars.
|
||||||
type item struct {
|
type item struct {
|
||||||
id int
|
id int
|
||||||
idx int // index into s.messages
|
idx int // index into s.messages
|
||||||
orig string
|
orig string
|
||||||
}
|
}
|
||||||
var items []item
|
var items []item
|
||||||
for i := range s.messages {
|
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.
|
// 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.
|
// Priority: agent config > OLLIE_COMPACTION_MODEL > backends.conf compactionModel > backends.conf model > current model.
|
||||||
// Returns "" if no override is configured (use the session's current model).
|
|
||||||
func resolveCompactionModel(cfgModel string, b backend.Backend) string {
|
func resolveCompactionModel(cfgModel string, b backend.Backend) string {
|
||||||
if cfgModel != "" {
|
if cfgModel != "" {
|
||||||
return cfgModel
|
return cfgModel
|
||||||
|
|
@ -451,34 +448,11 @@ func resolveCompactionModel(cfgModel string, b backend.Backend) string {
|
||||||
if env := os.Getenv("OLLIE_COMPACTION_MODEL"); env != "" {
|
if env := os.Getenv("OLLIE_COMPACTION_MODEL"); env != "" {
|
||||||
return env
|
return env
|
||||||
}
|
}
|
||||||
cfg := loadModelsConfig()
|
|
||||||
if b != nil {
|
if b != nil {
|
||||||
if m, ok := cfg.Compaction[b.Name()]; ok {
|
if model := backend.CompactionModel(b.Name()); model != "" {
|
||||||
return m
|
return model
|
||||||
}
|
}
|
||||||
|
return b.Model()
|
||||||
}
|
}
|
||||||
return ""
|
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++
|
step++
|
||||||
}
|
}
|
||||||
|
|
@ -210,20 +214,23 @@ func (ag *Agent) run(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
// autoCompact triggers context compaction if above threshold.
|
// 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 {
|
if ctx.Err() != nil || ag.history == nil {
|
||||||
return
|
return 0, nil
|
||||||
}
|
}
|
||||||
limit := ag.autoCompactLimit(ctx)
|
limit := ag.autoCompactLimit(ctx)
|
||||||
if limit <= 0 || ag.history.estimateTokens() < limit {
|
if limit <= 0 || ag.history.estimateTokens() < limit {
|
||||||
return
|
return 0, nil
|
||||||
}
|
}
|
||||||
ag.emit(Event{Role: "info", Content: "auto-compacting context...\n"})
|
ag.emit(Event{Role: "info", Content: "auto-compacting context...\n"})
|
||||||
ag.SetState("compacting")
|
ag.SetState("compacting")
|
||||||
if _, err := ag.runCompact(ctx, "auto"); err != nil {
|
n, err := ag.runCompact(ctx, "auto")
|
||||||
panic(fmt.Sprintf("mid-turn auto-compact: %v", err))
|
if err != nil {
|
||||||
|
return n, fmt.Errorf("auto-compact: %w", err)
|
||||||
}
|
}
|
||||||
ag.SetState("thinking")
|
ag.SetState("thinking")
|
||||||
|
return n, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ── streamResponse ──────────────────────────────────────────────────────────
|
// ── 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%.
|
// Warn once when context usage crosses 60%; compact at 75%.
|
||||||
if ag.history != nil {
|
if ag.history != nil {
|
||||||
tokens := ag.history.estimateTokens()
|
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.emit(Event{Role: "info", Content: "auto-compacting context...\n"})
|
||||||
ag.SetState("compacting")
|
ag.SetState("compacting")
|
||||||
if _, err := ag.runCompact(ctx, "auto"); err != nil {
|
if _, err := ag.runCompact(actCtx, "auto"); err != nil {
|
||||||
panic(fmt.Sprintf("auto-compact: %v", err))
|
ag.log.Error("auto-compact failed: %v", err)
|
||||||
}
|
}
|
||||||
ag.SetState("thinking")
|
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)
|
ctxLen := ag.runtime.Backend.ContextLength(ctx)
|
||||||
if ctxLen <= 0 {
|
if ctxLen <= 0 {
|
||||||
ctxLen = defaultContextLength
|
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).
|
// Strip cold-zone tool results (one batched LLM call per turn).
|
||||||
if ag.history != nil {
|
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.
|
// 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.history.messages = ag.history.messages[:preTurnLen]
|
||||||
ag.emit(Event{Role: "info", Content: "context overflow — compacting and retrying...\n"})
|
ag.emit(Event{Role: "info", Content: "context overflow — compacting and retrying...\n"})
|
||||||
ag.SetState("compacting")
|
ag.SetState("compacting")
|
||||||
if _, cerr := ag.runCompact(ctx, "overflow"); cerr != nil {
|
if _, cerr := ag.runCompact(actCtx, "overflow"); cerr != nil {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
ag.history.appendUserMessage(input)
|
ag.history.appendUserMessage(input)
|
||||||
|
|
|
||||||
|
|
@ -10,10 +10,11 @@ import (
|
||||||
|
|
||||||
// BackendConfig holds the per-backend configuration from backends.conf.
|
// BackendConfig holds the per-backend configuration from backends.conf.
|
||||||
type BackendConfig struct {
|
type BackendConfig struct {
|
||||||
Key string // API key / token
|
Key string // API key / token
|
||||||
URL string // base URL override (empty = use backend's default)
|
URL string // base URL override (empty = use backend's default)
|
||||||
Model string // default model override
|
Model string // default model override
|
||||||
Host string // for ollama: host address
|
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.
|
// 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")
|
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.
|
// parseConfigFile parses an INI-style config file.
|
||||||
// Format:
|
// Format:
|
||||||
//
|
//
|
||||||
|
|
@ -89,6 +102,8 @@ func parseConfigFile(path string) configFile {
|
||||||
bc.URL = v
|
bc.URL = v
|
||||||
case "model":
|
case "model":
|
||||||
bc.Model = v
|
bc.Model = v
|
||||||
|
case "compactionModel", "compaction_model", "compaction-model":
|
||||||
|
bc.CompactionModel = v
|
||||||
case "host":
|
case "host":
|
||||||
bc.Host = v
|
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
|
# key = API key or token
|
||||||
# url = base URL override (optional; each backend has a sensible default)
|
# url = base URL override (optional; each backend has a sensible default)
|
||||||
# model = default model (optional; overridden by session model= parameter)
|
# model = default model (optional; overridden by session model= parameter)
|
||||||
|
# compactionModel = model used for history compaction (optional; defaults to model)
|
||||||
# host = host address (ollama only)
|
# host = host address (ollama only)
|
||||||
|
|
||||||
backend = openrouter
|
backend = openrouter
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue