mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: session detail API (#23203)
This commit is contained in:
@@ -1097,6 +1097,287 @@ func AIBridgeToolUsage(usage database.AIBridgeToolUsage) codersdk.AIBridgeToolUs
|
||||
}
|
||||
}
|
||||
|
||||
// AIBridgeSessionThreads converts session metadata and thread interceptions
|
||||
// into the threads response. It groups interceptions into threads, builds
|
||||
// agentic actions from tool usages and model thoughts, and aggregates
|
||||
// token usage with metadata.
|
||||
func AIBridgeSessionThreads(
|
||||
session database.ListAIBridgeSessionsRow,
|
||||
interceptions []database.ListAIBridgeSessionThreadsRow,
|
||||
tokenUsages []database.AIBridgeTokenUsage,
|
||||
toolUsages []database.AIBridgeToolUsage,
|
||||
userPrompts []database.AIBridgeUserPrompt,
|
||||
modelThoughts []database.AIBridgeModelThought,
|
||||
) codersdk.AIBridgeSessionThreadsResponse {
|
||||
// Index subresources by interception ID.
|
||||
tokensByInterception := make(map[uuid.UUID][]database.AIBridgeTokenUsage, len(interceptions))
|
||||
for _, tu := range tokenUsages {
|
||||
tokensByInterception[tu.InterceptionID] = append(tokensByInterception[tu.InterceptionID], tu)
|
||||
}
|
||||
toolsByInterception := make(map[uuid.UUID][]database.AIBridgeToolUsage, len(interceptions))
|
||||
for _, tu := range toolUsages {
|
||||
toolsByInterception[tu.InterceptionID] = append(toolsByInterception[tu.InterceptionID], tu)
|
||||
}
|
||||
promptsByInterception := make(map[uuid.UUID][]database.AIBridgeUserPrompt, len(interceptions))
|
||||
for _, up := range userPrompts {
|
||||
promptsByInterception[up.InterceptionID] = append(promptsByInterception[up.InterceptionID], up)
|
||||
}
|
||||
thoughtsByInterception := make(map[uuid.UUID][]database.AIBridgeModelThought, len(interceptions))
|
||||
for _, mt := range modelThoughts {
|
||||
thoughtsByInterception[mt.InterceptionID] = append(thoughtsByInterception[mt.InterceptionID], mt)
|
||||
}
|
||||
|
||||
// Group interceptions by thread_id, preserving the order returned by the
|
||||
// SQL query.
|
||||
interceptionsByThread := make(map[uuid.UUID][]database.AIBridgeInterception, len(interceptions))
|
||||
var threadIDs []uuid.UUID
|
||||
for _, row := range interceptions {
|
||||
if _, ok := interceptionsByThread[row.ThreadID]; !ok {
|
||||
threadIDs = append(threadIDs, row.ThreadID)
|
||||
}
|
||||
interceptionsByThread[row.ThreadID] = append(interceptionsByThread[row.ThreadID], row.AIBridgeInterception)
|
||||
}
|
||||
|
||||
// Build threads and track page time bounds.
|
||||
threads := make([]codersdk.AIBridgeThread, 0, len(threadIDs))
|
||||
var pageStartedAt, pageEndedAt *time.Time
|
||||
for _, threadID := range threadIDs {
|
||||
intcs := interceptionsByThread[threadID]
|
||||
thread := buildAIBridgeThread(threadID, intcs, tokensByInterception, toolsByInterception, promptsByInterception, thoughtsByInterception)
|
||||
for _, intc := range intcs {
|
||||
if pageStartedAt == nil || intc.StartedAt.Before(*pageStartedAt) {
|
||||
t := intc.StartedAt
|
||||
pageStartedAt = &t
|
||||
}
|
||||
if intc.EndedAt.Valid {
|
||||
if pageEndedAt == nil || intc.EndedAt.Time.After(*pageEndedAt) {
|
||||
t := intc.EndedAt.Time
|
||||
pageEndedAt = &t
|
||||
}
|
||||
}
|
||||
}
|
||||
threads = append(threads, thread)
|
||||
}
|
||||
|
||||
// Aggregate session-level token usage metadata from all token
|
||||
// usages in the session (not just the page).
|
||||
sessionTokenMeta := aggregateTokenMetadata(tokenUsages)
|
||||
|
||||
resp := codersdk.AIBridgeSessionThreadsResponse{
|
||||
ID: session.SessionID,
|
||||
Initiator: MinimalUserFromVisibleUser(database.VisibleUser{
|
||||
ID: session.UserID,
|
||||
Username: session.UserUsername,
|
||||
Name: session.UserName,
|
||||
AvatarURL: session.UserAvatarUrl,
|
||||
}),
|
||||
Providers: session.Providers,
|
||||
Models: session.Models,
|
||||
Metadata: jsonOrEmptyMap(pqtype.NullRawMessage{RawMessage: session.Metadata, Valid: len(session.Metadata) > 0}),
|
||||
StartedAt: session.StartedAt,
|
||||
PageStartedAt: pageStartedAt,
|
||||
PageEndedAt: pageEndedAt,
|
||||
TokenUsageSummary: codersdk.AIBridgeSessionThreadsTokenUsage{
|
||||
InputTokens: session.InputTokens,
|
||||
OutputTokens: session.OutputTokens,
|
||||
Metadata: sessionTokenMeta,
|
||||
},
|
||||
Threads: threads,
|
||||
}
|
||||
if resp.Providers == nil {
|
||||
resp.Providers = []string{}
|
||||
}
|
||||
if resp.Models == nil {
|
||||
resp.Models = []string{}
|
||||
}
|
||||
if session.Client != "" {
|
||||
resp.Client = &session.Client
|
||||
}
|
||||
if !session.EndedAt.IsZero() {
|
||||
resp.EndedAt = &session.EndedAt
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
func buildAIBridgeThread(
|
||||
threadID uuid.UUID,
|
||||
interceptions []database.AIBridgeInterception,
|
||||
tokensByInterception map[uuid.UUID][]database.AIBridgeTokenUsage,
|
||||
toolsByInterception map[uuid.UUID][]database.AIBridgeToolUsage,
|
||||
promptsByInterception map[uuid.UUID][]database.AIBridgeUserPrompt,
|
||||
thoughtsByInterception map[uuid.UUID][]database.AIBridgeModelThought,
|
||||
) codersdk.AIBridgeThread {
|
||||
// Find the root interception (where id == threadID) to get the
|
||||
// thread prompt and model.
|
||||
var rootIntc *database.AIBridgeInterception
|
||||
for i := range interceptions {
|
||||
if interceptions[i].ID == threadID {
|
||||
rootIntc = &interceptions[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
// Fallback to first interception if root not found.
|
||||
if rootIntc == nil && len(interceptions) > 0 {
|
||||
rootIntc = &interceptions[0]
|
||||
}
|
||||
|
||||
thread := codersdk.AIBridgeThread{
|
||||
ID: threadID,
|
||||
}
|
||||
if rootIntc != nil {
|
||||
thread.Model = rootIntc.Model
|
||||
thread.Provider = rootIntc.Provider
|
||||
// Get first user prompt from root interception.
|
||||
// A thread can only have one prompt, by definition, since we currently
|
||||
// only store the last prompt observed in an interception.
|
||||
if prompts := promptsByInterception[rootIntc.ID]; len(prompts) > 0 {
|
||||
thread.Prompt = &prompts[0].Prompt
|
||||
}
|
||||
}
|
||||
|
||||
// Compute thread time bounds from interceptions.
|
||||
for _, intc := range interceptions {
|
||||
if thread.StartedAt.IsZero() || intc.StartedAt.Before(thread.StartedAt) {
|
||||
thread.StartedAt = intc.StartedAt
|
||||
}
|
||||
if intc.EndedAt.Valid {
|
||||
if thread.EndedAt == nil || intc.EndedAt.Time.After(*thread.EndedAt) {
|
||||
t := intc.EndedAt.Time
|
||||
thread.EndedAt = &t
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Build agentic actions grouped by interception. Each interception that
|
||||
// has tool calls produces one action with all its tool calls, thinking
|
||||
// blocks, and token usage.
|
||||
var actions []codersdk.AIBridgeAgenticAction
|
||||
for _, intc := range interceptions {
|
||||
tools := toolsByInterception[intc.ID]
|
||||
if len(tools) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
// Thinking blocks for this interception.
|
||||
thoughts := thoughtsByInterception[intc.ID]
|
||||
thinking := make([]codersdk.AIBridgeModelThought, 0, len(thoughts))
|
||||
for _, mt := range thoughts {
|
||||
thinking = append(thinking, codersdk.AIBridgeModelThought{
|
||||
Text: mt.Content,
|
||||
})
|
||||
}
|
||||
|
||||
// Token usage for the interception.
|
||||
actionTokenUsage := aggregateTokenUsage(tokensByInterception[intc.ID])
|
||||
|
||||
// Build tool call list.
|
||||
toolCalls := make([]codersdk.AIBridgeToolCall, 0, len(tools))
|
||||
for _, tu := range tools {
|
||||
toolCalls = append(toolCalls, codersdk.AIBridgeToolCall{
|
||||
ID: tu.ID,
|
||||
InterceptionID: tu.InterceptionID,
|
||||
ProviderResponseID: tu.ProviderResponseID,
|
||||
ServerURL: tu.ServerUrl.String,
|
||||
Tool: tu.Tool,
|
||||
Injected: tu.Injected,
|
||||
Input: tu.Input,
|
||||
Metadata: jsonOrEmptyMap(tu.Metadata),
|
||||
CreatedAt: tu.CreatedAt,
|
||||
})
|
||||
}
|
||||
|
||||
actions = append(actions, codersdk.AIBridgeAgenticAction{
|
||||
Model: intc.Model,
|
||||
TokenUsage: actionTokenUsage,
|
||||
Thinking: thinking,
|
||||
ToolCalls: toolCalls,
|
||||
})
|
||||
}
|
||||
|
||||
if actions == nil {
|
||||
// Make an empty slice so we don't serialize `null`.
|
||||
actions = make([]codersdk.AIBridgeAgenticAction, 0)
|
||||
}
|
||||
|
||||
thread.AgenticActions = actions
|
||||
|
||||
// Aggregate thread-level token usage.
|
||||
var threadTokens []database.AIBridgeTokenUsage
|
||||
for _, intc := range interceptions {
|
||||
threadTokens = append(threadTokens, tokensByInterception[intc.ID]...)
|
||||
}
|
||||
thread.TokenUsage = aggregateTokenUsage(threadTokens)
|
||||
|
||||
return thread
|
||||
}
|
||||
|
||||
// aggregateTokenUsage sums token usage rows and aggregates metadata.
|
||||
func aggregateTokenUsage(tokens []database.AIBridgeTokenUsage) codersdk.AIBridgeSessionThreadsTokenUsage {
|
||||
var inputTokens, outputTokens int64
|
||||
for _, tu := range tokens {
|
||||
inputTokens += tu.InputTokens
|
||||
outputTokens += tu.OutputTokens
|
||||
// TODO: once https://github.com/coder/aibridge/issues/150 lands we
|
||||
// should aggregate the other token types.
|
||||
}
|
||||
return codersdk.AIBridgeSessionThreadsTokenUsage{
|
||||
InputTokens: inputTokens,
|
||||
OutputTokens: outputTokens,
|
||||
Metadata: aggregateTokenMetadata(tokens),
|
||||
}
|
||||
}
|
||||
|
||||
// aggregateTokenMetadata sums all numeric values from the metadata
|
||||
// JSONB across the given token usage rows by key. Nested objects are
|
||||
// flattened using dot-notation (e.g. {"cache": {"read_tokens": 10}}
|
||||
// becomes "cache.read_tokens"). Non-numeric leaves (strings,
|
||||
// booleans, arrays, nulls) are silently skipped.
|
||||
func aggregateTokenMetadata(tokens []database.AIBridgeTokenUsage) map[string]any {
|
||||
sums := make(map[string]int64)
|
||||
for _, tu := range tokens {
|
||||
if !tu.Metadata.Valid || len(tu.Metadata.RawMessage) == 0 {
|
||||
continue
|
||||
}
|
||||
var m map[string]json.RawMessage
|
||||
if err := json.Unmarshal(tu.Metadata.RawMessage, &m); err != nil {
|
||||
continue
|
||||
}
|
||||
flattenAndSum(sums, "", m)
|
||||
}
|
||||
result := make(map[string]any, len(sums))
|
||||
for k, v := range sums {
|
||||
result[k] = v
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// flattenAndSum recursively walks a JSON object and sums all numeric
|
||||
// leaf values into sums, using dot-separated keys for nested objects.
|
||||
func flattenAndSum(sums map[string]int64, prefix string, m map[string]json.RawMessage) {
|
||||
for k, raw := range m {
|
||||
key := k
|
||||
if prefix != "" {
|
||||
key = prefix + "." + k
|
||||
}
|
||||
|
||||
// Try as a number first.
|
||||
var n json.Number
|
||||
if err := json.Unmarshal(raw, &n); err == nil {
|
||||
if v, err := n.Int64(); err == nil {
|
||||
sums[key] += v
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Try as a nested object.
|
||||
var nested map[string]json.RawMessage
|
||||
if err := json.Unmarshal(raw, &nested); err == nil {
|
||||
flattenAndSum(sums, key, nested)
|
||||
}
|
||||
// Arrays, strings, booleans, nulls are skipped.
|
||||
}
|
||||
}
|
||||
|
||||
func InvalidatedPresets(invalidatedPresets []database.UpdatePresetsLastInvalidatedAtRow) []codersdk.InvalidatedPreset {
|
||||
var presets []codersdk.InvalidatedPreset
|
||||
for _, p := range invalidatedPresets {
|
||||
|
||||
@@ -0,0 +1,308 @@
|
||||
package db2sdk
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
)
|
||||
|
||||
func TestAggregateTokenMetadata(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("empty_input", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
result := aggregateTokenMetadata(nil)
|
||||
require.Empty(t, result)
|
||||
})
|
||||
|
||||
t.Run("sums_across_rows", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{"cache_read_tokens":100,"reasoning_tokens":50}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{"cache_read_tokens":200,"reasoning_tokens":75}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenMetadata(tokens)
|
||||
require.Equal(t, int64(300), result["cache_read_tokens"])
|
||||
require.Equal(t, int64(125), result["reasoning_tokens"])
|
||||
require.Len(t, result, 2)
|
||||
})
|
||||
|
||||
t.Run("skips_null_and_invalid_metadata", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{Valid: false},
|
||||
},
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: nil,
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{"tokens":42}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenMetadata(tokens)
|
||||
require.Equal(t, int64(42), result["tokens"])
|
||||
require.Len(t, result, 1)
|
||||
})
|
||||
|
||||
t.Run("skips_non_integer_values", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
// Float values fail json.Number.Int64(), so they
|
||||
// are silently dropped.
|
||||
RawMessage: json.RawMessage(`{"good":10,"fractional":1.5}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenMetadata(tokens)
|
||||
require.Equal(t, int64(10), result["good"])
|
||||
_, hasFractional := result["fractional"]
|
||||
require.False(t, hasFractional)
|
||||
})
|
||||
|
||||
t.Run("skips_malformed_json", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`not json`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{"tokens":5}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenMetadata(tokens)
|
||||
// The malformed row is skipped, the valid one is counted.
|
||||
require.Equal(t, int64(5), result["tokens"])
|
||||
require.Len(t, result, 1)
|
||||
})
|
||||
|
||||
t.Run("flattens_nested_objects", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{
|
||||
"cache_read_tokens": 100,
|
||||
"cache": {"creation_tokens": 40, "read_tokens": 60},
|
||||
"reasoning_tokens": 50,
|
||||
"tags": ["a", "b"]
|
||||
}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{
|
||||
"cache_read_tokens": 200,
|
||||
"cache": {"creation_tokens": 10}
|
||||
}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenMetadata(tokens)
|
||||
require.Equal(t, int64(300), result["cache_read_tokens"])
|
||||
require.Equal(t, int64(50), result["reasoning_tokens"])
|
||||
require.Equal(t, int64(50), result["cache.creation_tokens"])
|
||||
require.Equal(t, int64(60), result["cache.read_tokens"])
|
||||
// Arrays are skipped.
|
||||
_, hasTags := result["tags"]
|
||||
require.False(t, hasTags)
|
||||
require.Len(t, result, 4)
|
||||
})
|
||||
|
||||
t.Run("flattens_deeply_nested_objects", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{
|
||||
"provider": {
|
||||
"anthropic": {"cache_creation_tokens": 100, "cache_read_tokens": 200},
|
||||
"openai": {"reasoning_tokens": 50}
|
||||
},
|
||||
"total": 500
|
||||
}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenMetadata(tokens)
|
||||
require.Equal(t, int64(100), result["provider.anthropic.cache_creation_tokens"])
|
||||
require.Equal(t, int64(200), result["provider.anthropic.cache_read_tokens"])
|
||||
require.Equal(t, int64(50), result["provider.openai.reasoning_tokens"])
|
||||
require.Equal(t, int64(500), result["total"])
|
||||
require.Len(t, result, 4)
|
||||
})
|
||||
|
||||
// Real-world provider metadata shapes from
|
||||
// https://github.com/coder/aibridge/issues/150.
|
||||
t.Run("aggregates_real_provider_metadata", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
// Anthropic-style: cache fields are top-level.
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 23490
|
||||
}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
// OpenAI-style: cache fields are nested inside
|
||||
// input_tokens_details.
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{
|
||||
"input_tokens_details": {"cached_tokens": 11904}
|
||||
}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
// Second Anthropic row to verify summing.
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{
|
||||
"cache_creation_input_tokens": 500,
|
||||
"cache_read_input_tokens": 10000
|
||||
}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenMetadata(tokens)
|
||||
// Anthropic fields are summed across two rows.
|
||||
require.Equal(t, int64(500), result["cache_creation_input_tokens"])
|
||||
require.Equal(t, int64(33490), result["cache_read_input_tokens"])
|
||||
// OpenAI nested field is flattened with dot notation.
|
||||
require.Equal(t, int64(11904), result["input_tokens_details.cached_tokens"])
|
||||
require.Len(t, result, 3)
|
||||
})
|
||||
|
||||
t.Run("skips_string_boolean_null_values", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
ID: uuid.New(),
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{"tokens":10,"name":"test","enabled":true,"nothing":null}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenMetadata(tokens)
|
||||
require.Equal(t, int64(10), result["tokens"])
|
||||
require.Len(t, result, 1)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAggregateTokenUsage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("empty_input", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
result := aggregateTokenUsage(nil)
|
||||
require.Equal(t, int64(0), result.InputTokens)
|
||||
require.Equal(t, int64(0), result.OutputTokens)
|
||||
require.Empty(t, result.Metadata)
|
||||
})
|
||||
|
||||
t.Run("sums_tokens_and_metadata", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
ID: uuid.New(),
|
||||
InputTokens: 100,
|
||||
OutputTokens: 50,
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{"reasoning_tokens":20}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: uuid.New(),
|
||||
InputTokens: 200,
|
||||
OutputTokens: 75,
|
||||
Metadata: pqtype.NullRawMessage{
|
||||
RawMessage: json.RawMessage(`{"reasoning_tokens":30}`),
|
||||
Valid: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenUsage(tokens)
|
||||
require.Equal(t, int64(300), result.InputTokens)
|
||||
require.Equal(t, int64(125), result.OutputTokens)
|
||||
require.Equal(t, int64(50), result.Metadata["reasoning_tokens"])
|
||||
})
|
||||
|
||||
t.Run("handles_rows_without_metadata", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tokens := []database.AIBridgeTokenUsage{
|
||||
{
|
||||
ID: uuid.New(),
|
||||
InputTokens: 500,
|
||||
OutputTokens: 200,
|
||||
Metadata: pqtype.NullRawMessage{Valid: false},
|
||||
},
|
||||
}
|
||||
|
||||
result := aggregateTokenUsage(tokens)
|
||||
require.Equal(t, int64(500), result.InputTokens)
|
||||
require.Equal(t, int64(200), result.OutputTokens)
|
||||
require.Empty(t, result.Metadata)
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user