mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-04 16:13:47 +00:00
122 lines
4.2 KiB
Go
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
|
|
}
|