mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-07-30 09:56:05 +00:00
* feat(gateway): cache-aware cost accounting + /api/console observability endpoints - Capture prompt-cache hits (usage_metadata.input_token_details.cache_read) in RunJournal and SubagentTokenCollector as a sparse cache_read_tokens key in token_usage_by_model (JSON field — no schema migration; legacy bucket shapes unchanged) - New read-only /api/console router: GET /stats (headline counters), GET /runs (cross-thread paginated history joined with thread titles), GET /usage (zero-filled daily token series + per-model breakdown); user-scoped, 503 on the memory database backend - Optional models[*].pricing (currency, input_per_million, output_per_million, input_cache_hit_per_million) powers real spend estimation; cache-hit input tokens are billed at the hit price (omitted hit price falls back to the miss price as a conservative upper bound); unpriced models yield cost: null - create_chat_model strips the presentation-only pricing block so it never reaches the provider client (unknown kwargs are forwarded into the completion payload and break live calls) - Tests: console router SQLite round-trips, journal/collector cache capture incl. a DeepSeek raw-usage pin test, factory strip regression Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> * refactor: address review feedback on cost sum and sparse cache_read_tokens - console.py: replace the walrus-in-generator total-cost sum with an explicit loop (review noted the multi-line form reads ambiguously) - token_collector.py: omit cache_read_tokens from usage records when the provider reported no cache hits, matching the journal's sparse per-model bucket shape; absent is treated as 0 downstream - add a regression test pinning the sparse record shape Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> --------- Co-authored-by: coffeeFish <codeingforcoffee@users.noreply.github.com> Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
84 lines
3.4 KiB
Python
84 lines
3.4 KiB
Python
"""Callback handler that collects LLM token usage within a subagent.
|
|
|
|
Each subagent execution creates its own collector. After the subagent
|
|
finishes, the collected records are transferred to the parent RunJournal
|
|
via :meth:`RunJournal.record_external_llm_usage_records`.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from collections.abc import Mapping
|
|
from typing import Any
|
|
|
|
from langchain_core.callbacks import BaseCallbackHandler
|
|
|
|
|
|
class SubagentTokenCollector(BaseCallbackHandler):
|
|
"""Lightweight callback handler that collects LLM token usage within a subagent."""
|
|
|
|
def __init__(self, caller: str):
|
|
super().__init__()
|
|
self.caller = caller
|
|
self._records: list[dict[str, int | str | None]] = []
|
|
self._counted_run_ids: set[str] = set()
|
|
|
|
def on_llm_end(
|
|
self,
|
|
response: Any,
|
|
*,
|
|
run_id: Any,
|
|
tags: list[str] | None = None,
|
|
**kwargs: Any,
|
|
) -> None:
|
|
rid = str(run_id)
|
|
if rid in self._counted_run_ids:
|
|
return
|
|
|
|
for generation in response.generations:
|
|
for gen in generation:
|
|
if not hasattr(gen, "message"):
|
|
continue
|
|
usage = getattr(gen.message, "usage_metadata", None)
|
|
usage_dict = dict(usage) if usage else {}
|
|
input_tk = usage_dict.get("input_tokens", 0) or 0
|
|
output_tk = usage_dict.get("output_tokens", 0) or 0
|
|
total_tk = usage_dict.get("total_tokens", 0) or 0
|
|
if total_tk <= 0:
|
|
total_tk = input_tk + output_tk
|
|
if total_tk <= 0:
|
|
continue
|
|
# Prompt-cache hits (needed for cache-aware cost accounting)
|
|
details = usage_dict.get("input_token_details") or {}
|
|
cache_read_tk = 0
|
|
if isinstance(details, Mapping):
|
|
try:
|
|
cache_read_tk = max(int(details.get("cache_read") or 0), 0)
|
|
except (TypeError, ValueError):
|
|
cache_read_tk = 0
|
|
# Capture the model that actually produced this response so the
|
|
# parent journal can bucket tokens by real model rather than the
|
|
# lead agent's resolved model
|
|
response_metadata = getattr(gen.message, "response_metadata", None) or {}
|
|
model_name: str | None = None
|
|
if isinstance(response_metadata, Mapping):
|
|
model_name = response_metadata.get("model_name") or response_metadata.get("model")
|
|
self._counted_run_ids.add(rid)
|
|
record: dict[str, int | str | None] = {
|
|
"source_run_id": rid,
|
|
"caller": self.caller,
|
|
"model_name": model_name,
|
|
"input_tokens": input_tk,
|
|
"output_tokens": output_tk,
|
|
"total_tokens": total_tk,
|
|
}
|
|
# Sparse, matching the journal's per-model buckets: the key is
|
|
# only present when the provider actually reported cache hits.
|
|
if cache_read_tk > 0:
|
|
record["cache_read_tokens"] = cache_read_tk
|
|
self._records.append(record)
|
|
return
|
|
|
|
def snapshot_records(self) -> list[dict[str, int | str | None]]:
|
|
"""Return a copy of the accumulated usage records."""
|
|
return list(self._records)
|