ollie/cmd/olliesrv/internal/agent/cancellation_test.go

143 lines
4.1 KiB
Go

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")
}
}