返回 DeepSeek-Reasonix
usage_accounting_test.go
根目录 / internal / agent / usage_accounting_test.go
1 package agent
2
3 import (
4 "context"
5 "errors"
6 "io"
7 "net/http"
8 "strings"
9 "testing"
10
11 "reasonix/internal/event"
12 "reasonix/internal/provider"
13 "reasonix/internal/tool"
14 )
15
16 type accountingRoundTripFunc func(*http.Request) (*http.Response, error)
17
18 func (f accountingRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
19 return f(req)
20 }
21
22 type failedRequestProvider struct{}
23
24 func (failedRequestProvider) Name() string { return "failed-request" }
25
26 func (failedRequestProvider) Stream(ctx context.Context, _ provider.Request) (<-chan provider.Chunk, error) {
27 requestCtx := provider.WithRequestAttemptCounter(ctx)
28 client := &http.Client{Transport: accountingRoundTripFunc(func(*http.Request) (*http.Response, error) {
29 return &http.Response{
30 StatusCode: http.StatusBadRequest,
31 Header: make(http.Header),
32 Body: io.NopCloser(strings.NewReader("bad request")),
33 }, nil
34 })}
35 _, err := provider.SendWithRetry(requestCtx, client, provider.SendOptions{Provider: "failed-request"}, func(reqCtx context.Context) (*http.Request, error) {
36 return http.NewRequestWithContext(reqCtx, http.MethodPost, "https://example.invalid", nil)
37 })
38 return nil, err
39 }
40
41 func TestMergeStreamUsageCountsProviderRequests(t *testing.T) {
42 first := &provider.Usage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15, CacheWriteTokens: 2, CacheWriteBilledTokens: 2.5, RequestCount: 1}
43 retry := &provider.Usage{PromptTokens: 20, CompletionTokens: 8, TotalTokens: 28, CacheWriteTokens: 3, CacheWriteBilledTokens: 6, RequestCount: 1}
44 got := mergeStreamUsage(first, retry)
45 if got == nil || got.TotalTokens != 43 || got.RequestCount != 2 || got.CompletionTokens != 13 {
46 t.Fatalf("merged usage = %+v, want total=43 requests=2 completion=13", got)
47 }
48 // Billable PromptTokens align with summed cache hit+miss.
49 if got.CacheMissTokens != 30 || got.PromptTokens != 30 {
50 t.Fatalf("billable input = prompt %d miss %d, want 30/30", got.PromptTokens, got.CacheMissTokens)
51 }
52 if got.CacheWriteTokens != 5 || got.CacheWriteBilledTokens != 8.5 {
53 t.Fatalf("merged cache writes = raw %d billed %v, want 5/8.5", got.CacheWriteTokens, got.CacheWriteBilledTokens)
54 }
55
56 third := &provider.Usage{PromptTokens: 1, CompletionTokens: 1, TotalTokens: 2, RequestCount: 1}
57 got = mergeStreamUsage(got, third)
58 if got.RequestCount != 3 {
59 t.Fatalf("nested merged request count = %d, want 3", got.RequestCount)
60 }
61
62 got = mergeStreamUsage(nil, retry)
63 if got == nil || got.TotalTokens != retry.TotalTokens || got.RequestCount != 1 {
64 t.Fatalf("missing first usage = %+v, want retry tokens and 1 request", got)
65 }
66 got = mergeStreamUsage(first, nil)
67 if got == nil || got.TotalTokens != first.TotalTokens || got.RequestCount != 1 {
68 t.Fatalf("missing retry usage = %+v, want first tokens and 1 request", got)
69 }
70
71 requestOnly := &provider.Usage{RequestCount: 3}
72 got = mergeStreamUsage(first, requestOnly)
73 if got == nil || got.RequestCount != 4 {
74 t.Fatalf("request-only retry usage = %+v, want 4 requests", got)
75 }
76 }
77
78 func TestFinalizeSamplingUsageKeepsLatestPromptContext(t *testing.T) {
79 billable := &provider.Usage{
80 PromptTokens: 90000, CompletionTokens: 30, TotalTokens: 90030,
81 CacheMissTokens: 90000, RequestCount: 3,
82 }
83 latest := &provider.Usage{PromptTokens: 30000, CompletionTokens: 10, TotalTokens: 30010, CacheMissTokens: 30000, RequestCount: 1}
84 got := finalizeSamplingUsage(billable, latest)
85 if got == nil || got.PromptTokens != 90000 {
86 t.Fatalf("prompt tokens = %+v, want billable total 90000", got)
87 }
88 if got.ContextPromptTokens != 30000 || got.ContextCompletionTokens != 10 {
89 t.Fatalf("context shape = prompt %d completion %d, want latest 30000/10", got.ContextPromptTokens, got.ContextCompletionTokens)
90 }
91 if got.ContextFillTokens() != 30000 {
92 t.Fatalf("ContextFillTokens = %d, want 30000", got.ContextFillTokens())
93 }
94 completionOnly := &provider.Usage{PromptTokens: 500, ContextCompletionTokens: 20}
95 if fill := completionOnly.ContextFillTokens(); fill != 500 {
96 t.Fatalf("completion-only ContextFillTokens = %d, want prompt fallback 500", fill)
97 }
98 if got.CompletionTokens != 30 || got.RequestCount != 3 {
99 t.Fatalf("billable fields = %+v, want summed completion/requests", got)
100 }
101 // lastUsage stores the latest attempt wholesale (prompt+completion of that
102 // request), never the billable aggregate.
103 if latest.PromptTokens != 30000 || latest.CompletionTokens != 10 {
104 t.Fatalf("latest attempt shape mutated: %+v", latest)
105 }
106 }
107
108 func TestMergeSamplingUsageKeepsBillableTokensAcrossRequestOnlyAttempt(t *testing.T) {
109 first := &provider.Usage{
110 PromptTokens: 100, CompletionTokens: 0, TotalTokens: 100,
111 CacheMissTokens: 100, RequestCount: 1,
112 }
113 second := &provider.Usage{RequestCount: 1}
114 got := mergeSamplingUsage(first, second)
115 if got.PromptTokens != 100 || got.TotalTokens != 100 || got.RequestCount != 2 {
116 t.Fatalf("merged billable = %+v, want first tokens + 2 requests", got)
117 }
118 final := finalizeSamplingUsage(got, second)
119 if final == nil || final.PromptTokens != 100 {
120 t.Fatalf("final usage = %+v, want billable prompt 100", final)
121 }
122 }
123
124 func TestEstimateFailedAttemptUsageIncludesArgChars(t *testing.T) {
125 frozen := samplingRequest{
126 req: provider.Request{Messages: []provider.Message{{Role: provider.RoleUser, Content: "write a large file"}}},
127 }
128 // ~8KB of streamed tool args with no terminal usage.
129 result := streamedTurn{
130 maxArgChars: 8192,
131 err: &provider.StreamInterruptedError{Err: io.ErrUnexpectedEOF, Reason: provider.StreamInterruptPrematureEOF},
132 interrupted: true,
133 }
134 got := estimateFailedAttemptUsage(nil, frozen, result, 1)
135 if got == nil || !got.Estimated {
136 t.Fatalf("usage = %+v, want estimated failed-attempt record", got)
137 }
138 argTokens := (8192 + 3) / 4
139 if got.CompletionTokens < argTokens {
140 t.Fatalf("completion tokens = %d, want at least arg estimate %d", got.CompletionTokens, argTokens)
141 }
142 if got.PromptTokens <= 0 {
143 t.Fatalf("prompt tokens = %d, want request input estimate", got.PromptTokens)
144 }
145 }
146
147 func TestEstimateFailedAttemptUsageSkipsZeroHTTPLocalFailure(t *testing.T) {
148 frozen := samplingRequest{
149 req: provider.Request{Messages: []provider.Message{{Role: provider.RoleUser, Content: "hi"}}},
150 }
151 result := streamedTurn{
152 err: errors.New("local request validation failed"),
153 }
154 // No HTTP request and no speculative output: do not invent billable usage.
155 got := estimateFailedAttemptUsage(nil, frozen, result, 0)
156 if got != nil {
157 t.Fatalf("pre-body local reject usage = %+v, want nil (no invented billable tokens)", got)
158 }
159 first := &provider.Usage{PromptTokens: 100, TotalTokens: 100, CacheMissTokens: 100, RequestCount: 1}
160 merged := mergeSamplingUsage(first, got)
161 if merged == nil || merged.PromptTokens != 100 || merged.RequestCount != 1 {
162 t.Fatalf("merged after local reject = %+v, want first attempt only", merged)
163 }
164 }
165
166 func TestStreamReturnsRequestOnlyUsageOnProviderFailure(t *testing.T) {
167 var events []event.Event
168 sink := event.FuncSink(func(e event.Event) { events = append(events, e) })
169 a := New(failedRequestProvider{}, tool.NewRegistry(), NewSession(""), Options{ModelRef: "failed/model"}, sink)
170
171 st := a.stream(context.Background(), 1, sink)
172 if st.err == nil {
173 t.Fatal("expected provider failure")
174 }
175 if st.usage == nil || st.usage.TotalTokens != 0 || st.usage.RequestCount != 1 {
176 t.Fatalf("failed stream usage = %+v, want tokens=0 requests=1", st.usage)
177 }
178 a.emitTurnUsage(st.usage, nil)
179 if len(events) != 1 || events[0].Kind != event.Usage || events[0].Usage.RequestCount != 1 {
180 t.Fatalf("request-only usage event = %+v", events)
181 }
182 }
183
184 func TestTaskUsageModelRefUsesCanonicalRuntimeIdentity(t *testing.T) {
185 task := (&TaskTool{baseModel: "deepseek/deepseek-v4-pro"}).WithTranscriptIdentityResolver(
186 func(modelRef, effort string) (string, string) {
187 if modelRef == "flash" {
188 return "deepseek/deepseek-v4-flash", effort
189 }
190 return "deepseek/deepseek-v4-pro", effort
191 },
192 )
193 if got := task.usageModelRef("flash", "high"); got != "deepseek/deepseek-v4-flash" {
194 t.Fatalf("alias usage model = %q", got)
195 }
196 if got := task.usageModelRef("", ""); got != "deepseek/deepseek-v4-pro" {
197 t.Fatalf("inherited usage model = %q", got)
198 }
199 }
200
200 lines GO