Files
goclaw/internal/channelmemory/extractor_test.go
T
2026-07-05 22:09:59 +07:00

122 lines
4.2 KiB
Go

package channelmemory
import (
"context"
"testing"
"time"
"github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/store"
)
type fakeExtractionProvider struct {
responses []providers.ChatResponse
requests []providers.ChatRequest
}
func (f *fakeExtractionProvider) Chat(_ context.Context, req providers.ChatRequest) (*providers.ChatResponse, error) {
f.requests = append(f.requests, req)
if len(f.responses) == 0 {
return &providers.ChatResponse{Content: "[]", FinishReason: "stop"}, nil
}
resp := f.responses[0]
f.responses = f.responses[1:]
return &resp, nil
}
func (f *fakeExtractionProvider) ChatStream(_ context.Context, req providers.ChatRequest, _ func(providers.StreamChunk)) (*providers.ChatResponse, error) {
return f.Chat(context.Background(), req)
}
func (f *fakeExtractionProvider) DefaultModel() string { return "fake-model" }
func (f *fakeExtractionProvider) Name() string { return "fake" }
func TestExtractRetriesAfterLengthFinish(t *testing.T) {
provider := &fakeExtractionProvider{responses: []providers.ChatResponse{
{Content: `[{"type":"todos","summary":"Finish`, FinishReason: "length"},
{Content: `[{"type":"todos","summary":"Finish release notes","topics":["release"],"entities":["GoClaw"],"confidence":0.9}]`, FinishReason: "stop"},
}}
items, err := Extract(context.Background(), provider, "fake-model", nil, extractionTestMessages(12), DefaultAllowedTypes)
if err != nil {
t.Fatalf("Extract() error = %v", err)
}
if len(items) != 1 {
t.Fatalf("len(items) = %d, want 1", len(items))
}
if items[0].Summary != "Finish release notes" {
t.Fatalf("summary = %q", items[0].Summary)
}
if len(provider.requests) != 2 {
t.Fatalf("provider calls = %d, want 2", len(provider.requests))
}
if got := provider.requests[0].Options[providers.OptMaxTokens]; got != extractionMaxOutputTokens {
t.Fatalf("first max_tokens = %v, want %d", got, extractionMaxOutputTokens)
}
if got := provider.requests[1].Options[providers.OptMaxTokens]; got != extractionRetryMaxOutputTokens {
t.Fatalf("retry max_tokens = %v, want %d", got, extractionRetryMaxOutputTokens)
}
}
func TestExtractRetriesAfterEmptyResponse(t *testing.T) {
provider := &fakeExtractionProvider{responses: []providers.ChatResponse{
{Content: ``, FinishReason: "stop"},
{Content: `[]`, FinishReason: "stop"},
}}
items, err := Extract(context.Background(), provider, "fake-model", nil, extractionTestMessages(5), DefaultAllowedTypes)
if err != nil {
t.Fatalf("Extract() error = %v", err)
}
if len(items) != 0 {
t.Fatalf("len(items) = %d, want 0", len(items))
}
if len(provider.requests) != 2 {
t.Fatalf("provider calls = %d, want 2", len(provider.requests))
}
}
func TestExtractRetriesAfterUnexpectedJSONEnd(t *testing.T) {
provider := &fakeExtractionProvider{responses: []providers.ChatResponse{
{Content: `[{"type":"projects","summary":"`, FinishReason: "stop"},
{Content: `[{"type":"projects","summary":"Gateway dashboard rollout","topics":["dashboard"],"entities":["Gateway"],"confidence":0.86}]`, FinishReason: "stop"},
}}
items, err := Extract(context.Background(), provider, "fake-model", nil, extractionTestMessages(5), DefaultAllowedTypes)
if err != nil {
t.Fatalf("Extract() error = %v", err)
}
if len(items) != 1 {
t.Fatalf("len(items) = %d, want 1", len(items))
}
if items[0].Type != "projects" {
t.Fatalf("type = %q, want projects", items[0].Type)
}
if len(provider.requests) != 2 {
t.Fatalf("provider calls = %d, want 2", len(provider.requests))
}
}
func TestParseExtractionResponseStripsCodeFence(t *testing.T) {
items, err := parseExtractionResponse("```json\n[]\n```")
if err != nil {
t.Fatalf("parseExtractionResponse() error = %v", err)
}
if len(items) != 0 {
t.Fatalf("len(items) = %d, want 0", len(items))
}
}
func extractionTestMessages(n int) []store.PendingMessage {
messages := make([]store.PendingMessage, 0, n)
base := time.Date(2026, 7, 5, 9, 0, 0, 0, time.UTC)
for i := range n {
messages = append(messages, store.PendingMessage{
Sender: "tester",
Body: "Durable project context that may be extracted into channel memory.",
CreatedAt: base.Add(time.Duration(i) * time.Minute),
})
}
return messages
}