First Commit
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
web / web (push) Waiting to run
docs / links (push) Canceled after 0s
detections / detections (push) Canceled after 0s
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
web / web (push) Waiting to run
docs / links (push) Canceled after 0s
detections / detections (push) Canceled after 0s
This commit is contained in:
@@ -0,0 +1,112 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user