Files
artex/agent/capture_usage_test.go
dela 0335d572de
ci / go (push) Waiting to run
ci / go-db (agent) (push) Waiting to run
ci / go-db (config) (push) Waiting to run
ci / go-db (db) (push) Waiting to run
ci / go-db (evidence) (push) Waiting to run
ci / go-db (llmrec) (push) Waiting to run
ci / go-db (server) (push) Waiting to run
detections / detections (push) Waiting to run
web / web (push) Waiting to run
docs / links (push) Canceled after 0s
First Commit
2026-10-09 08:38:16 +08:00

113 lines
3.8 KiB
Go

package agent
import (
"context"
"errors"
"iter"
"testing"
"github.com/Autumn-27/artex/db"
"github.com/Autumn-27/norma/agentcore"
"github.com/Autumn-27/norma/llm"
)
type captureUsageProvider struct {
stream func(context.Context, func(llm.StreamEvent, error) bool)
}
func (p captureUsageProvider) Stream(ctx context.Context, _ llm.CompletionRequest) iter.Seq2[llm.StreamEvent, error] {
return func(yield func(llm.StreamEvent, error) bool) { p.stream(ctx, yield) }
}
func (p captureUsageProvider) Complete(ctx context.Context, req llm.CompletionRequest) (llm.Message, string, llm.Usage, error) {
acc := llm.NewAccumulator()
for ev, err := range p.Stream(ctx, req) {
if err != nil {
return llm.Message{}, "", llm.Usage{}, err
}
acc.Add(ev)
}
return acc.Message(), acc.StopReason, acc.Usage, nil
}
func TestCaptureRunPersistsUsageOnProviderFailure(t *testing.T) {
wantErr := errors.New("provider failed after reporting usage")
provider := captureUsageProvider{stream: func(_ context.Context, yield func(llm.StreamEvent, error) bool) {
// 遵守迭代器协议:yield 返回 false 后立即停止,不再调用它。
if !yield(llm.StreamEvent{Type: llm.SEMessageStart, Usage: llm.Usage{InputTokens: 11, CacheReadTokens: 3}}, nil) {
return
}
if !yield(llm.StreamEvent{Type: llm.SETextDelta, Text: "partial"}, nil) {
return
}
if !yield(llm.StreamEvent{Type: llm.SEMessageDelta, Usage: llm.Usage{OutputTokens: 7, CacheWriteTokens: 2}}, nil) {
return
}
yield(llm.StreamEvent{}, wantErr)
}}
var activities []db.Activity
_, _, err := captureRun(context.Background(), agentcore.Options{Provider: provider, MaxTurns: 1}, "test", func(a db.Activity) {
activities = append(activities, a)
})
if !errors.Is(err, wantErr) {
t.Fatalf("captureRun error=%v, want %v", err, wantErr)
}
assertCapturedResultUsage(t, activities, 11, 7, 3, 2)
}
func TestCaptureRunPersistsUsageOnCancellation(t *testing.T) {
started := make(chan struct{})
provider := captureUsageProvider{stream: func(ctx context.Context, yield func(llm.StreamEvent, error) bool) {
// 遵守迭代器协议:yield 返回 false 后立即停止,不再调用它。
if !yield(llm.StreamEvent{Type: llm.SEMessageStart, Usage: llm.Usage{InputTokens: 13, CacheReadTokens: 5}}, nil) {
return
}
if !yield(llm.StreamEvent{Type: llm.SETextDelta, Text: "partial"}, nil) {
return
}
if !yield(llm.StreamEvent{Type: llm.SEMessageDelta, Usage: llm.Usage{OutputTokens: 9, CacheWriteTokens: 4}}, nil) {
return
}
close(started)
<-ctx.Done()
yield(llm.StreamEvent{}, ctx.Err())
}}
ctx, cancel := context.WithCancel(context.Background())
var activities []db.Activity
done := make(chan error, 1)
go func() {
_, _, err := captureRun(ctx, agentcore.Options{Provider: provider, MaxTurns: 1}, "test", func(a db.Activity) {
activities = append(activities, a)
})
done <- err
}()
<-started
cancel()
if err := <-done; !errors.Is(err, context.Canceled) {
t.Fatalf("captureRun error=%v, want context canceled", err)
}
assertCapturedResultUsage(t, activities, 13, 9, 5, 4)
}
func assertCapturedResultUsage(t *testing.T, activities []db.Activity, input, output, read, write int) {
t.Helper()
results := 0
for _, activity := range activities {
if activity.Kind != "result" {
continue
}
results++
if activity.InputTokens == nil || *activity.InputTokens != input ||
activity.OutputTokens == nil || *activity.OutputTokens != output ||
activity.CacheReadTokens == nil || *activity.CacheReadTokens != read ||
activity.CacheWriteTokens == nil || *activity.CacheWriteTokens != write {
t.Fatalf("result usage=%+v, want input=%d output=%d read=%d write=%d", activity, input, output, read, write)
}
}
if results != 1 {
t.Fatalf("result activity count=%d, activities=%+v", results, activities)
}
}