fix: aggregate token usage across native tool-call loops for analytics (#206)

A streaming reply that uses native/server-side tool calls makes several
LLM calls for one visible message, but only the final call's usage was
persisted — so the built-in analytics undercounted tokens for tool-heavy
turns and outlet() filters couldn't see the earlier calls.

Persist every call's usage object instead of overwriting. `usage` becomes
a list (one object per LLM call) for tool-loop turns, or stays a single
object otherwise:

- middleware: collect each call's usage during the tool loop and persist
  the list (or the single object) as the message's usage.
- analytics: _token_columns sums the per-call list in SQL (json_each /
  json_array_elements) when usage is an array, else reads the object —
  unchanged for normal/legacy rows. Computed on demand; nothing pre-summed.
- chat_messages: get_usage and the upsert merge handle list-valued usage.
- frontend: the usage tooltip (already gated on message.done) renders the
  last call's object, so the visible message still shows only the final
  call; earlier tool-turns persist but aren't displayed.

outlet() filters receive the full per-call list (via the message and
ctx['assistant_message']) for custom usage logging/tracking.
This commit is contained in:
Classic298 2026-06-03 13:19:14 +02:00 • committed by GitHub
parent 1a97751e37
commit 9f091e6ed2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 103 additions and 23 deletions

View file

@ -15,6 +15,7 @@ from sqlalchemy import (
Index,
Integer,
Text,
case,
cast,
delete,
func,
@ -45,28 +46,67 @@ def _normalize_timestamp(timestamp: int) -> float:
return timestamp
def get_usage(data: dict) -> Optional[dict]:
"""Extract and normalize usage from message data."""
def get_usage(data: dict):
"""Extract and normalize usage from message data.
A native tool-call loop persists one usage object per LLM call, so usage
may be a *list* of objects; a normal turn is a single object. Returns the
same shape it was given (list or dict), normalized, or None.
"""
usage = data.get('usage') or (data.get('info') or {}).get('usage')
return normalize_usage(usage) if usage else None
if not usage:
return None
if isinstance(usage, list):
return [normalize_usage(u) for u in usage if u]
return normalize_usage(usage)
def _token_columns(dialect: str):
"""Return (input_tokens, output_tokens) SQL column expressions.
Falls back to OpenAI-style keys (prompt_tokens / completion_tokens)
when the normalized keys are absent.
``usage`` is normally a single object, but a native tool-call loop stores
a list with one usage object per LLM call. When it's a list, sum the
tokens across the calls so analytics counts the true total; when it's an
object, read it directly. Falls back to OpenAI-style keys (prompt_tokens /
completion_tokens) when the normalized keys are absent.
"""
if dialect == 'sqlite':
extract = lambda key: cast(func.json_extract(ChatMessage.usage, f'$.{key}'), Integer)
scalar = lambda key: cast(func.json_extract(ChatMessage.usage, f'$.{key}'), Integer)
is_list = func.json_type(ChatMessage.usage) == 'array'
def list_sum(norm, fallback):
elem = func.json_each(ChatMessage.usage).table_valued('value', joins_implicitly=True)
token = func.coalesce(
cast(func.json_extract(elem.c.value, f'$.{norm}'), Integer),
cast(func.json_extract(elem.c.value, f'$.{fallback}'), Integer),
)
return select(func.coalesce(func.sum(token), 0)).scalar_subquery()
elif dialect == 'postgresql':
extract = lambda key: cast(func.json_extract_path_text(ChatMessage.usage, key), Integer)
scalar = lambda key: cast(func.json_extract_path_text(ChatMessage.usage, key), Integer)
is_list = func.json_typeof(ChatMessage.usage) == 'array'
def list_sum(norm, fallback):
elem = func.json_array_elements(ChatMessage.usage).table_valued('value', joins_implicitly=True)
token = func.coalesce(
cast(func.json_extract_path_text(elem.c.value, norm), Integer),
cast(func.json_extract_path_text(elem.c.value, fallback), Integer),
)
return select(func.coalesce(func.sum(token), 0)).scalar_subquery()
else:
raise NotImplementedError(f'Unsupported dialect: {dialect}')
def column(norm, fallback):
# Per row: sum the per-call list, or read the single object.
return case(
(is_list, list_sum(norm, fallback)),
else_=func.coalesce(scalar(norm), scalar(fallback)),
)
return (
func.coalesce(extract('input_tokens'), extract('prompt_tokens')),
func.coalesce(extract('output_tokens'), extract('completion_tokens')),
column('input_tokens', 'prompt_tokens'),
column('output_tokens', 'completion_tokens'),
)
@ -140,7 +180,7 @@ class ChatMessageModel(BaseModel):
done: bool = True
status_history: Optional[list] = None
error: Optional[dict | str] = None
usage: Optional[dict] = None
usage: Optional[dict | list] = None # list = one object per call (tool loops)
created_at: int
updated_at: int
@ -194,11 +234,16 @@ class ChatMessageTable:
existing.error = data.get('error')
# Extract and normalize usage
usage = get_usage(data)
if usage:
# Deep-merge: preserve existing keys not present in new data
# This prevents background tasks (follow-ups, title, tags)
# from accidentally clearing the primary response's token counts
existing.usage = {**(existing.usage or {}), **usage}
if usage is not None:
if isinstance(usage, dict) and not isinstance(existing.usage, list):
# Deep-merge: preserve existing keys not present in new data
# This prevents background tasks (follow-ups, title, tags)
# from accidentally clearing the primary response's token counts
existing.usage = {**(existing.usage or {}), **usage}
else:
# Native tool loop persists a list (one object per LLM
# call) — store it as-is rather than dict-merging.
existing.usage = usage
existing.updated_at = now
await db.commit()
await db.refresh(existing)

View file

@ -3864,6 +3864,12 @@ async def streaming_chat_response_handler(response, ctx):
output = []
usage = None
# One usage object per LLM call. A native tool-call loop makes
# several calls for a single visible message; we keep every call's
# object so the persisted `usage` is the full list (the UI renders
# only the last; analytics sums the list). `usage` above stays the
# last call, used for the live per-chunk events during streaming.
usage_calls = []
prior_output = []
last_response_id = None
@ -3921,10 +3927,15 @@ async def streaming_chat_response_handler(response, ctx):
async def stream_body_handler(response, form_data):
nonlocal content
nonlocal usage
nonlocal usage_calls
nonlocal output
nonlocal prior_output
nonlocal last_response_id
# Usage reported by *this* LLM call (one segment of a native
# tool-call loop). Appended to usage_calls when the call ends.
segment_usage = None
response_tool_calls = []
delta_count = 0
@ -4054,6 +4065,7 @@ async def streaming_chat_response_handler(response, ctx):
if response_metadata.get('usage'):
response_metadata['usage'] = normalize_usage(response_metadata['usage'])
usage = response_metadata['usage']
segment_usage = usage
processed_data.update(response_metadata)
processed_data.pop('done', None)
@ -4073,6 +4085,7 @@ async def streaming_chat_response_handler(response, ctx):
raw_usage.update(data.get('timings', {})) # llama.cpp
if raw_usage:
usage = normalize_usage(raw_usage)
segment_usage = usage
await event_emitter(
{
'type': 'chat:completion',
@ -4505,6 +4518,11 @@ async def streaming_chat_response_handler(response, ctx):
if responses_api_tool_calls:
tool_calls.append(_split_tool_calls(responses_api_tool_calls))
# Record this call's usage. Across a tool loop this builds
# the list of every call's usage object that gets persisted.
if segment_usage:
usage_calls.append(segment_usage)
try:
await stream_body_handler(response, form_data)
finally:
@ -5111,12 +5129,19 @@ async def streaming_chat_response_handler(response, ctx):
if not metadata.get('chat_id', '').startswith('channel:')
else ''
)
# Persist the full per-call usage: a list when the turn made
# several LLM calls (native tool loop), or the single object
# otherwise. The UI renders only the last call; analytics sums
# the list; outlet() filters can read every call.
usage_final = usage_calls if len(usage_calls) > 1 else (usage_calls[0] if usage_calls else None)
data = {
'done': True,
'content': serialize_output(output),
'output': output,
'title': title,
**({'usage': usage} if usage else {}),
**({'usage': usage_final} if usage_final else {}),
}
if not metadata.get('chat_id', '').startswith('channel:'):
@ -5129,14 +5154,14 @@ async def streaming_chat_response_handler(response, ctx):
'done': True,
'content': serialize_output(output),
'output': output,
**({'usage': usage} if usage else {}),
**({'usage': usage_final} if usage_final else {}),
},
)
elif usage:
elif usage_final:
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{'done': True, 'usage': usage},
{'done': True, 'usage': usage_final},
)
else:
await Chats.upsert_message_to_chat_by_id_and_message_id(
@ -5171,7 +5196,9 @@ async def streaming_chat_response_handler(response, ctx):
ctx['assistant_message'] = {
'content': serialize_output(output),
'output': output,
**({'usage': usage} if usage else {}),
# Full per-call usage so outlet() filters can read every
# tool-turn (a list for tool loops, else a single object).
**({'usage': usage_final} if usage_final else {}),
}
await outlet_filter_handler(ctx)
await background_tasks_handler(ctx)

View file

@ -113,6 +113,9 @@
load_duration?: number;
usage?: unknown;
};
// Single object for a normal turn, or one object per LLM call for a
// native tool-call loop (the UI shows only the last call).
usage?: Record<string, any> | Record<string, any>[];
annotation?: { type: string; rating: number };
}
@ -136,6 +139,11 @@
}
}
// A native tool-call loop persists `usage` as a list (one object per LLM
// call); only the final call is the visible answer, so show that one. A
// normal turn keeps `usage` as a single object.
$: usageInfo = Array.isArray(message?.usage) ? message.usage.at(-1) : message?.usage;
export let siblings;
export let setInputText: Function = () => {};
@ -1142,11 +1150,11 @@
</Tooltip>
{/if}
{#if message.usage}
{#if usageInfo}
<Tooltip
content={message.usage
content={usageInfo
? `<pre>${sanitizeResponseContent(
JSON.stringify(message.usage, null, 2)
JSON.stringify(usageInfo, null, 2)
.replace(/"([^(")"]+)":/g, '$1:')
.slice(1, -1)
.split('\n')