mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-06 20:19:03 +00:00
fix(gateway): attribute token usage to actual models (#3658)
* fix(gateway): attribute token usage to actual models
Capture per-call model names from LLM response metadata for lead, middleware, and subagent calls.
Persist a per-run token_usage_by_model breakdown and aggregate by that map in both SQL and memory stores, with legacy fallback to the run-level model_name for older rows.
Add regression coverage for by_model totals, caller consistency, active progress snapshots, store parity, and SubagentTokenCollector model propagation.
* fix(gateway): harden by-model token aggregation
Use usage.get("total_tokens", 0) when reducing per-model token usage maps so aggregation tolerates partially written or manually edited JSON blobs without changing behavior for journal-written rows.
* docs(gateway): clarify by-model run count semantics
Document that by_model[*].runs counts the number of runs in which a model appeared, so multi-model runs can increment multiple model buckets.
This commit is contained in:
parent
8cde305fe4
commit
e7a03e5243
@ -79,7 +79,10 @@ class RunResponse(BaseModel):
|
|||||||
|
|
||||||
class ThreadTokenUsageModelBreakdown(BaseModel):
|
class ThreadTokenUsageModelBreakdown(BaseModel):
|
||||||
tokens: int = 0
|
tokens: int = 0
|
||||||
runs: int = 0
|
runs: int = Field(
|
||||||
|
default=0,
|
||||||
|
description="Number of runs in which this model appeared; counts are non-exclusive for runs that used multiple models.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ThreadTokenUsageCallerBreakdown(BaseModel):
|
class ThreadTokenUsageCallerBreakdown(BaseModel):
|
||||||
|
|||||||
@ -39,6 +39,8 @@ class RunRow(Base):
|
|||||||
lead_agent_tokens: Mapped[int] = mapped_column(default=0)
|
lead_agent_tokens: Mapped[int] = mapped_column(default=0)
|
||||||
subagent_tokens: Mapped[int] = mapped_column(default=0)
|
subagent_tokens: Mapped[int] = mapped_column(default=0)
|
||||||
middleware_tokens: Mapped[int] = mapped_column(default=0)
|
middleware_tokens: Mapped[int] = mapped_column(default=0)
|
||||||
|
# Per-model token breakdown
|
||||||
|
token_usage_by_model: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||||
|
|
||||||
# Follow-up association
|
# Follow-up association
|
||||||
follow_up_to_run_id: Mapped[str | None] = mapped_column(String(64))
|
follow_up_to_run_id: Mapped[str | None] = mapped_column(String(64))
|
||||||
|
|||||||
@ -11,7 +11,7 @@ import json
|
|||||||
from datetime import UTC, datetime
|
from datetime import UTC, datetime
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from sqlalchemy import func, select, update
|
from sqlalchemy import select, update
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
||||||
|
|
||||||
from deerflow.persistence.run.model import RunRow
|
from deerflow.persistence.run.model import RunRow
|
||||||
@ -230,6 +230,7 @@ class RunRepository(RunStore):
|
|||||||
lead_agent_tokens: int = 0,
|
lead_agent_tokens: int = 0,
|
||||||
subagent_tokens: int = 0,
|
subagent_tokens: int = 0,
|
||||||
middleware_tokens: int = 0,
|
middleware_tokens: int = 0,
|
||||||
|
token_usage_by_model: dict[str, dict[str, int]] | None = None,
|
||||||
message_count: int = 0,
|
message_count: int = 0,
|
||||||
last_ai_message: str | None = None,
|
last_ai_message: str | None = None,
|
||||||
first_human_message: str | None = None,
|
first_human_message: str | None = None,
|
||||||
@ -248,6 +249,7 @@ class RunRepository(RunStore):
|
|||||||
"lead_agent_tokens": lead_agent_tokens,
|
"lead_agent_tokens": lead_agent_tokens,
|
||||||
"subagent_tokens": subagent_tokens,
|
"subagent_tokens": subagent_tokens,
|
||||||
"middleware_tokens": middleware_tokens,
|
"middleware_tokens": middleware_tokens,
|
||||||
|
"token_usage_by_model": self._safe_json(token_usage_by_model) or {},
|
||||||
"message_count": message_count,
|
"message_count": message_count,
|
||||||
"updated_at": datetime.now(UTC),
|
"updated_at": datetime.now(UTC),
|
||||||
}
|
}
|
||||||
@ -273,6 +275,7 @@ class RunRepository(RunStore):
|
|||||||
lead_agent_tokens: int | None = None,
|
lead_agent_tokens: int | None = None,
|
||||||
subagent_tokens: int | None = None,
|
subagent_tokens: int | None = None,
|
||||||
middleware_tokens: int | None = None,
|
middleware_tokens: int | None = None,
|
||||||
|
token_usage_by_model: dict[str, dict[str, int]] | None = None,
|
||||||
message_count: int | None = None,
|
message_count: int | None = None,
|
||||||
last_ai_message: str | None = None,
|
last_ai_message: str | None = None,
|
||||||
first_human_message: str | None = None,
|
first_human_message: str | None = None,
|
||||||
@ -292,6 +295,8 @@ class RunRepository(RunStore):
|
|||||||
for key, value in optional_counters.items():
|
for key, value in optional_counters.items():
|
||||||
if value is not None:
|
if value is not None:
|
||||||
values[key] = value
|
values[key] = value
|
||||||
|
if token_usage_by_model is not None:
|
||||||
|
values["token_usage_by_model"] = self._safe_json(token_usage_by_model) or {}
|
||||||
if last_ai_message is not None:
|
if last_ai_message is not None:
|
||||||
values["last_ai_message"] = last_ai_message[:2000]
|
values["last_ai_message"] = last_ai_message[:2000]
|
||||||
if first_human_message is not None:
|
if first_human_message is not None:
|
||||||
@ -301,26 +306,33 @@ class RunRepository(RunStore):
|
|||||||
await session.commit()
|
await session.commit()
|
||||||
|
|
||||||
async def aggregate_tokens_by_thread(self, thread_id: str, *, include_active: bool = False) -> dict[str, Any]:
|
async def aggregate_tokens_by_thread(self, thread_id: str, *, include_active: bool = False) -> dict[str, Any]:
|
||||||
"""Aggregate token usage via a single SQL GROUP BY query."""
|
"""Aggregate token usage for a thread.
|
||||||
|
|
||||||
|
``by_model`` is reduced in Python from each row's ``token_usage_by_model``
|
||||||
|
JSON column so subagent / middleware tokens land on the model that
|
||||||
|
actually produced them (issue #3645). Rows written before that column
|
||||||
|
existed fall back to ``RunRow.model_name`` + ``RunRow.total_tokens``,
|
||||||
|
preserving the legacy lead-only behavior instead of dropping the data.
|
||||||
|
|
||||||
|
Headline totals (``total_tokens``, ``total_input_tokens``,
|
||||||
|
``total_output_tokens``) and the ``by_caller`` bucket are summed from
|
||||||
|
their own columns and are therefore unaffected by the JSON column being
|
||||||
|
empty.
|
||||||
|
"""
|
||||||
statuses = ("success", "error", "running") if include_active else ("success", "error")
|
statuses = ("success", "error", "running") if include_active else ("success", "error")
|
||||||
_completed = RunRow.status.in_(statuses)
|
_completed = RunRow.status.in_(statuses)
|
||||||
_thread = RunRow.thread_id == thread_id
|
_thread = RunRow.thread_id == thread_id
|
||||||
model_name = func.coalesce(RunRow.model_name, "unknown")
|
|
||||||
|
|
||||||
stmt = (
|
stmt = select(
|
||||||
select(
|
RunRow.model_name,
|
||||||
model_name.label("model"),
|
RunRow.total_tokens,
|
||||||
func.count().label("runs"),
|
RunRow.total_input_tokens,
|
||||||
func.coalesce(func.sum(RunRow.total_tokens), 0).label("total_tokens"),
|
RunRow.total_output_tokens,
|
||||||
func.coalesce(func.sum(RunRow.total_input_tokens), 0).label("total_input_tokens"),
|
RunRow.lead_agent_tokens,
|
||||||
func.coalesce(func.sum(RunRow.total_output_tokens), 0).label("total_output_tokens"),
|
RunRow.subagent_tokens,
|
||||||
func.coalesce(func.sum(RunRow.lead_agent_tokens), 0).label("lead_agent"),
|
RunRow.middleware_tokens,
|
||||||
func.coalesce(func.sum(RunRow.subagent_tokens), 0).label("subagent"),
|
RunRow.token_usage_by_model,
|
||||||
func.coalesce(func.sum(RunRow.middleware_tokens), 0).label("middleware"),
|
).where(_thread, _completed)
|
||||||
)
|
|
||||||
.where(_thread, _completed)
|
|
||||||
.group_by(model_name)
|
|
||||||
)
|
|
||||||
|
|
||||||
async with self._sf() as session:
|
async with self._sf() as session:
|
||||||
rows = (await session.execute(stmt)).all()
|
rows = (await session.execute(stmt)).all()
|
||||||
@ -329,14 +341,28 @@ class RunRepository(RunStore):
|
|||||||
lead_agent = subagent = middleware = 0
|
lead_agent = subagent = middleware = 0
|
||||||
by_model: dict[str, dict] = {}
|
by_model: dict[str, dict] = {}
|
||||||
for r in rows:
|
for r in rows:
|
||||||
by_model[r.model] = {"tokens": r.total_tokens, "runs": r.runs}
|
total_runs += 1
|
||||||
total_tokens += r.total_tokens
|
total_tokens += r.total_tokens
|
||||||
total_input += r.total_input_tokens
|
total_input += r.total_input_tokens
|
||||||
total_output += r.total_output_tokens
|
total_output += r.total_output_tokens
|
||||||
total_runs += r.runs
|
lead_agent += r.lead_agent_tokens
|
||||||
lead_agent += r.lead_agent
|
subagent += r.subagent_tokens
|
||||||
subagent += r.subagent
|
middleware += r.middleware_tokens
|
||||||
middleware += r.middleware
|
|
||||||
|
# ``or {}`` covers rows written before ``token_usage_by_model``
|
||||||
|
# existed (the column is NULL on a manual ALTER ADD COLUMN without
|
||||||
|
# backfill); fresh rows always carry the journal-produced dict.
|
||||||
|
usage_by_model = r.token_usage_by_model or {}
|
||||||
|
if usage_by_model:
|
||||||
|
for model, usage in usage_by_model.items():
|
||||||
|
entry = by_model.setdefault(model, {"tokens": 0, "runs": 0})
|
||||||
|
entry["tokens"] += usage.get("total_tokens", 0)
|
||||||
|
entry["runs"] += 1
|
||||||
|
else:
|
||||||
|
model = r.model_name or "unknown"
|
||||||
|
entry = by_model.setdefault(model, {"tokens": 0, "runs": 0})
|
||||||
|
entry["tokens"] += r.total_tokens
|
||||||
|
entry["runs"] += 1
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"total_tokens": total_tokens,
|
"total_tokens": total_tokens,
|
||||||
|
|||||||
@ -77,6 +77,9 @@ class RunJournal(BaseCallbackHandler):
|
|||||||
self._subagent_tokens = 0
|
self._subagent_tokens = 0
|
||||||
self._middleware_tokens = 0
|
self._middleware_tokens = 0
|
||||||
|
|
||||||
|
# Per-model token accumulator
|
||||||
|
self._tokens_by_model: dict[str, dict[str, int]] = {}
|
||||||
|
|
||||||
# Dedup: LangChain may fire on_llm_end multiple times for the same run_id
|
# Dedup: LangChain may fire on_llm_end multiple times for the same run_id
|
||||||
self._counted_llm_run_ids: set[str] = set()
|
self._counted_llm_run_ids: set[str] = set()
|
||||||
self._counted_external_source_ids: set[str] = set()
|
self._counted_external_source_ids: set[str] = set()
|
||||||
@ -327,6 +330,13 @@ class RunJournal(BaseCallbackHandler):
|
|||||||
else:
|
else:
|
||||||
self._lead_agent_tokens += total_tk
|
self._lead_agent_tokens += total_tk
|
||||||
|
|
||||||
|
# Per-model bucket
|
||||||
|
response_metadata = getattr(message, "response_metadata", None) or {}
|
||||||
|
per_call_model: str | None = None
|
||||||
|
if isinstance(response_metadata, Mapping):
|
||||||
|
per_call_model = response_metadata.get("model_name") or response_metadata.get("model")
|
||||||
|
self._record_model_usage(per_call_model, input_tk, output_tk, total_tk)
|
||||||
|
|
||||||
self._schedule_progress_flush()
|
self._schedule_progress_flush()
|
||||||
|
|
||||||
if messages:
|
if messages:
|
||||||
@ -435,17 +445,42 @@ class RunJournal(BaseCallbackHandler):
|
|||||||
# themselves.
|
# themselves.
|
||||||
return "lead_agent"
|
return "lead_agent"
|
||||||
|
|
||||||
|
def _record_model_usage(
|
||||||
|
self,
|
||||||
|
model_name: str | None,
|
||||||
|
input_tokens: int,
|
||||||
|
output_tokens: int,
|
||||||
|
total_tokens: int,
|
||||||
|
) -> None:
|
||||||
|
"""Add a single LLM call's token usage to the per-model accumulator.
|
||||||
|
|
||||||
|
Missing / empty ``model_name`` collapses into a shared ``"unknown"``
|
||||||
|
bucket so the breakdown stays usable when a provider doesn't surface
|
||||||
|
``response_metadata.model_name``.
|
||||||
|
"""
|
||||||
|
if total_tokens <= 0:
|
||||||
|
return
|
||||||
|
bucket = self._tokens_by_model.setdefault(
|
||||||
|
model_name or "unknown",
|
||||||
|
{"input_tokens": 0, "output_tokens": 0, "total_tokens": 0},
|
||||||
|
)
|
||||||
|
bucket["input_tokens"] += int(input_tokens or 0)
|
||||||
|
bucket["output_tokens"] += int(output_tokens or 0)
|
||||||
|
bucket["total_tokens"] += int(total_tokens)
|
||||||
|
|
||||||
# -- Public methods (called by worker) --
|
# -- Public methods (called by worker) --
|
||||||
|
|
||||||
def record_external_llm_usage_records(
|
def record_external_llm_usage_records(
|
||||||
self,
|
self,
|
||||||
records: list[dict[str, int | str]],
|
records: list[dict[str, int | str | None]],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Record token usage from external sources (e.g., subagents).
|
"""Record token usage from external sources (e.g., subagents).
|
||||||
|
|
||||||
Each record should contain:
|
Each record should contain:
|
||||||
source_run_id: Unique identifier to prevent double-counting
|
source_run_id: Unique identifier to prevent double-counting
|
||||||
caller: Caller tag (e.g. "subagent:general-purpose")
|
caller: Caller tag (e.g. "subagent:general-purpose")
|
||||||
|
model_name: Real per-call model name (str or None; falls back to
|
||||||
|
``"unknown"`` bucket when missing)
|
||||||
input_tokens: Input token count
|
input_tokens: Input token count
|
||||||
output_tokens: Output token count
|
output_tokens: Output token count
|
||||||
total_tokens: Total token count (computed from input+output if 0/missing)
|
total_tokens: Total token count (computed from input+output if 0/missing)
|
||||||
@ -467,9 +502,12 @@ class RunJournal(BaseCallbackHandler):
|
|||||||
if total_tk <= 0:
|
if total_tk <= 0:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
input_tk = record.get("input_tokens", 0) or 0
|
||||||
|
output_tk = record.get("output_tokens", 0) or 0
|
||||||
|
|
||||||
self._counted_external_source_ids.add(source_id)
|
self._counted_external_source_ids.add(source_id)
|
||||||
self._total_input_tokens += record.get("input_tokens", 0) or 0
|
self._total_input_tokens += input_tk
|
||||||
self._total_output_tokens += record.get("output_tokens", 0) or 0
|
self._total_output_tokens += output_tk
|
||||||
self._total_tokens += total_tk
|
self._total_tokens += total_tk
|
||||||
|
|
||||||
caller = str(record.get("caller", ""))
|
caller = str(record.get("caller", ""))
|
||||||
@ -480,6 +518,8 @@ class RunJournal(BaseCallbackHandler):
|
|||||||
else:
|
else:
|
||||||
self._lead_agent_tokens += total_tk
|
self._lead_agent_tokens += total_tk
|
||||||
|
|
||||||
|
self._record_model_usage(record.get("model_name"), input_tk, output_tk, total_tk)
|
||||||
|
|
||||||
self._schedule_progress_flush()
|
self._schedule_progress_flush()
|
||||||
|
|
||||||
def set_first_human_message(self, content: str) -> None:
|
def set_first_human_message(self, content: str) -> None:
|
||||||
@ -590,6 +630,7 @@ class RunJournal(BaseCallbackHandler):
|
|||||||
"lead_agent_tokens": self._lead_agent_tokens,
|
"lead_agent_tokens": self._lead_agent_tokens,
|
||||||
"subagent_tokens": self._subagent_tokens,
|
"subagent_tokens": self._subagent_tokens,
|
||||||
"middleware_tokens": self._middleware_tokens,
|
"middleware_tokens": self._middleware_tokens,
|
||||||
|
"token_usage_by_model": {model: dict(usage) for model, usage in self._tokens_by_model.items()},
|
||||||
"message_count": self._msg_count,
|
"message_count": self._msg_count,
|
||||||
"last_ai_message": self._last_ai_msg,
|
"last_ai_message": self._last_ai_msg,
|
||||||
"first_human_message": self._first_human_msg,
|
"first_human_message": self._first_human_msg,
|
||||||
|
|||||||
@ -99,6 +99,8 @@ class RunRecord:
|
|||||||
lead_agent_tokens: int = 0
|
lead_agent_tokens: int = 0
|
||||||
subagent_tokens: int = 0
|
subagent_tokens: int = 0
|
||||||
middleware_tokens: int = 0
|
middleware_tokens: int = 0
|
||||||
|
# Per-model token breakdown
|
||||||
|
token_usage_by_model: dict[str, dict[str, int]] = field(default_factory=dict)
|
||||||
message_count: int = 0
|
message_count: int = 0
|
||||||
last_ai_message: str | None = None
|
last_ai_message: str | None = None
|
||||||
first_human_message: str | None = None
|
first_human_message: str | None = None
|
||||||
@ -291,6 +293,7 @@ class RunManager:
|
|||||||
lead_agent_tokens=row.get("lead_agent_tokens") or 0,
|
lead_agent_tokens=row.get("lead_agent_tokens") or 0,
|
||||||
subagent_tokens=row.get("subagent_tokens") or 0,
|
subagent_tokens=row.get("subagent_tokens") or 0,
|
||||||
middleware_tokens=row.get("middleware_tokens") or 0,
|
middleware_tokens=row.get("middleware_tokens") or 0,
|
||||||
|
token_usage_by_model=row.get("token_usage_by_model") or {},
|
||||||
message_count=row.get("message_count") or 0,
|
message_count=row.get("message_count") or 0,
|
||||||
last_ai_message=row.get("last_ai_message"),
|
last_ai_message=row.get("last_ai_message"),
|
||||||
first_human_message=row.get("first_human_message"),
|
first_human_message=row.get("first_human_message"),
|
||||||
|
|||||||
@ -93,6 +93,7 @@ class RunStore(abc.ABC):
|
|||||||
lead_agent_tokens: int = 0,
|
lead_agent_tokens: int = 0,
|
||||||
subagent_tokens: int = 0,
|
subagent_tokens: int = 0,
|
||||||
middleware_tokens: int = 0,
|
middleware_tokens: int = 0,
|
||||||
|
token_usage_by_model: dict[str, dict[str, int]] | None = None,
|
||||||
message_count: int = 0,
|
message_count: int = 0,
|
||||||
last_ai_message: str | None = None,
|
last_ai_message: str | None = None,
|
||||||
first_human_message: str | None = None,
|
first_human_message: str | None = None,
|
||||||
@ -115,6 +116,7 @@ class RunStore(abc.ABC):
|
|||||||
lead_agent_tokens: int | None = None,
|
lead_agent_tokens: int | None = None,
|
||||||
subagent_tokens: int | None = None,
|
subagent_tokens: int | None = None,
|
||||||
middleware_tokens: int | None = None,
|
middleware_tokens: int | None = None,
|
||||||
|
token_usage_by_model: dict[str, dict[str, int]] | None = None,
|
||||||
message_count: int | None = None,
|
message_count: int | None = None,
|
||||||
last_ai_message: str | None = None,
|
last_ai_message: str | None = None,
|
||||||
first_human_message: str | None = None,
|
first_human_message: str | None = None,
|
||||||
|
|||||||
@ -110,10 +110,21 @@ class MemoryRunStore(RunStore):
|
|||||||
completed = [r for r in self._runs.values() if r["thread_id"] == thread_id and r.get("status") in statuses]
|
completed = [r for r in self._runs.values() if r["thread_id"] == thread_id and r.get("status") in statuses]
|
||||||
by_model: dict[str, dict] = {}
|
by_model: dict[str, dict] = {}
|
||||||
for r in completed:
|
for r in completed:
|
||||||
model = r.get("model_name") or "unknown"
|
usage_by_model = r.get("token_usage_by_model") or {}
|
||||||
entry = by_model.setdefault(model, {"tokens": 0, "runs": 0})
|
if usage_by_model:
|
||||||
entry["tokens"] += r.get("total_tokens", 0)
|
for model, usage in usage_by_model.items():
|
||||||
entry["runs"] += 1
|
entry = by_model.setdefault(model, {"tokens": 0, "runs": 0})
|
||||||
|
entry["tokens"] += usage.get("total_tokens", 0)
|
||||||
|
entry["runs"] += 1
|
||||||
|
else:
|
||||||
|
# Fallback for rows written before per-model accounting landed:
|
||||||
|
# attribute the whole run to its single ``model_name``. Keeps
|
||||||
|
# the legacy lead-only behavior for old data instead of
|
||||||
|
# silently dropping it.
|
||||||
|
model = r.get("model_name") or "unknown"
|
||||||
|
entry = by_model.setdefault(model, {"tokens": 0, "runs": 0})
|
||||||
|
entry["tokens"] += r.get("total_tokens", 0)
|
||||||
|
entry["runs"] += 1
|
||||||
return {
|
return {
|
||||||
"total_tokens": sum(r.get("total_tokens", 0) for r in completed),
|
"total_tokens": sum(r.get("total_tokens", 0) for r in completed),
|
||||||
"total_input_tokens": sum(r.get("total_input_tokens", 0) for r in completed),
|
"total_input_tokens": sum(r.get("total_input_tokens", 0) for r in completed),
|
||||||
|
|||||||
@ -89,7 +89,7 @@ class SubagentResult:
|
|||||||
started_at: datetime | None = None
|
started_at: datetime | None = None
|
||||||
completed_at: datetime | None = None
|
completed_at: datetime | None = None
|
||||||
ai_messages: list[dict[str, Any]] | None = None
|
ai_messages: list[dict[str, Any]] | None = None
|
||||||
token_usage_records: list[dict[str, int | str]] = field(default_factory=list)
|
token_usage_records: list[dict[str, int | str | None]] = field(default_factory=list)
|
||||||
usage_reported: bool = False
|
usage_reported: bool = False
|
||||||
cancel_event: threading.Event = field(default_factory=threading.Event, repr=False)
|
cancel_event: threading.Event = field(default_factory=threading.Event, repr=False)
|
||||||
_state_lock: threading.Lock = field(default_factory=threading.Lock, init=False, repr=False)
|
_state_lock: threading.Lock = field(default_factory=threading.Lock, init=False, repr=False)
|
||||||
@ -107,7 +107,7 @@ class SubagentResult:
|
|||||||
error: str | None = None,
|
error: str | None = None,
|
||||||
completed_at: datetime | None = None,
|
completed_at: datetime | None = None,
|
||||||
ai_messages: list[dict[str, Any]] | None = None,
|
ai_messages: list[dict[str, Any]] | None = None,
|
||||||
token_usage_records: list[dict[str, int | str]] | None = None,
|
token_usage_records: list[dict[str, int | str | None]] | None = None,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Set a terminal status exactly once.
|
"""Set a terminal status exactly once.
|
||||||
|
|
||||||
|
|||||||
@ -7,6 +7,7 @@ via :meth:`RunJournal.record_external_llm_usage_records`.
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Mapping
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from langchain_core.callbacks import BaseCallbackHandler
|
from langchain_core.callbacks import BaseCallbackHandler
|
||||||
@ -18,7 +19,7 @@ class SubagentTokenCollector(BaseCallbackHandler):
|
|||||||
def __init__(self, caller: str):
|
def __init__(self, caller: str):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.caller = caller
|
self.caller = caller
|
||||||
self._records: list[dict[str, int | str]] = []
|
self._records: list[dict[str, int | str | None]] = []
|
||||||
self._counted_run_ids: set[str] = set()
|
self._counted_run_ids: set[str] = set()
|
||||||
|
|
||||||
def on_llm_end(
|
def on_llm_end(
|
||||||
@ -46,11 +47,19 @@ class SubagentTokenCollector(BaseCallbackHandler):
|
|||||||
total_tk = input_tk + output_tk
|
total_tk = input_tk + output_tk
|
||||||
if total_tk <= 0:
|
if total_tk <= 0:
|
||||||
continue
|
continue
|
||||||
|
# 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)
|
self._counted_run_ids.add(rid)
|
||||||
self._records.append(
|
self._records.append(
|
||||||
{
|
{
|
||||||
"source_run_id": rid,
|
"source_run_id": rid,
|
||||||
"caller": self.caller,
|
"caller": self.caller,
|
||||||
|
"model_name": model_name,
|
||||||
"input_tokens": input_tk,
|
"input_tokens": input_tk,
|
||||||
"output_tokens": output_tk,
|
"output_tokens": output_tk,
|
||||||
"total_tokens": total_tk,
|
"total_tokens": total_tk,
|
||||||
@ -58,6 +67,6 @@ class SubagentTokenCollector(BaseCallbackHandler):
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
def snapshot_records(self) -> list[dict[str, int | str]]:
|
def snapshot_records(self) -> list[dict[str, int | str | None]]:
|
||||||
"""Return a copy of the accumulated usage records."""
|
"""Return a copy of the accumulated usage records."""
|
||||||
return list(self._records)
|
return list(self._records)
|
||||||
|
|||||||
@ -3,8 +3,6 @@
|
|||||||
Uses a temp SQLite DB to test ORM-backed CRUD operations.
|
Uses a temp SQLite DB to test ORM-backed CRUD operations.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import re
|
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from sqlalchemy.dialects import postgresql
|
from sqlalchemy.dialects import postgresql
|
||||||
|
|
||||||
@ -451,7 +449,11 @@ class TestRunRepository:
|
|||||||
await _cleanup()
|
await _cleanup()
|
||||||
|
|
||||||
@pytest.mark.anyio
|
@pytest.mark.anyio
|
||||||
async def test_aggregate_tokens_by_thread_reuses_shared_model_name_expression(self):
|
async def test_aggregate_tokens_by_thread_returns_zeros_when_no_rows(self):
|
||||||
|
"""Empty thread aggregates to all-zero totals, no model buckets, and a
|
||||||
|
single query — replaces the older test that pinned the now-removed
|
||||||
|
``GROUP BY coalesce(model_name)`` shape (issue #3645 reduces by_model
|
||||||
|
in Python from each row's per-model JSON column instead)."""
|
||||||
captured = []
|
captured = []
|
||||||
|
|
||||||
class FakeResult:
|
class FakeResult:
|
||||||
@ -483,17 +485,44 @@ class TestRunRepository:
|
|||||||
}
|
}
|
||||||
assert len(captured) == 1
|
assert len(captured) == 1
|
||||||
|
|
||||||
stmt = captured[0]
|
@pytest.mark.anyio
|
||||||
compiled_sql = str(stmt.compile(dialect=postgresql.dialect()))
|
async def test_aggregate_tokens_by_thread_compiles_on_postgres_dialect(self):
|
||||||
select_sql, group_by_sql = compiled_sql.split(" GROUP BY ", maxsplit=1)
|
"""Compile-smoke the new SELECT on the postgres dialect.
|
||||||
model_expr_pattern = r"coalesce\(runs\.model_name, %\(([^)]+)\)s\)"
|
|
||||||
|
|
||||||
select_match = re.search(model_expr_pattern + r" AS model", select_sql)
|
The project ships both SQLite and Postgres backends. The new aggregation
|
||||||
group_by_match = re.fullmatch(model_expr_pattern, group_by_sql.strip())
|
projects ``RunRow.token_usage_by_model`` (a JSON column) directly into
|
||||||
|
the row set instead of grouping on a scalar, so the SQL needs to compile
|
||||||
|
cleanly under PG's JSON/JSONB binding too. Pins:
|
||||||
|
* the JSON column is selected by name (PG would otherwise need a
|
||||||
|
``::jsonb`` cast or coalesce around it)
|
||||||
|
* there is no GROUP BY / aggregate function left (the per-model
|
||||||
|
reduction now happens in Python — see issue #3645)
|
||||||
|
"""
|
||||||
|
|
||||||
assert select_match is not None
|
captured = []
|
||||||
assert group_by_match is not None
|
|
||||||
assert select_match.group(1) == group_by_match.group(1)
|
class FakeResult:
|
||||||
|
def all(self):
|
||||||
|
return []
|
||||||
|
|
||||||
|
class FakeSession:
|
||||||
|
async def execute(self, stmt):
|
||||||
|
captured.append(stmt)
|
||||||
|
return FakeResult()
|
||||||
|
|
||||||
|
class FakeSessionContext:
|
||||||
|
async def __aenter__(self):
|
||||||
|
return FakeSession()
|
||||||
|
|
||||||
|
async def __aexit__(self, exc_type, exc, tb):
|
||||||
|
return None
|
||||||
|
|
||||||
|
repo = RunRepository(lambda: FakeSessionContext())
|
||||||
|
await repo.aggregate_tokens_by_thread("t1")
|
||||||
|
|
||||||
|
compiled = str(captured[0].compile(dialect=postgresql.dialect()))
|
||||||
|
assert "token_usage_by_model" in compiled
|
||||||
|
assert "GROUP BY" not in compiled.upper()
|
||||||
|
|
||||||
@pytest.mark.anyio
|
@pytest.mark.anyio
|
||||||
async def test_run_manager_hydrates_store_only_run_from_sql(self, tmp_path):
|
async def test_run_manager_hydrates_store_only_run_from_sql(self, tmp_path):
|
||||||
|
|||||||
@ -6,11 +6,12 @@ from uuid import uuid4
|
|||||||
from deerflow.subagents.token_collector import SubagentTokenCollector
|
from deerflow.subagents.token_collector import SubagentTokenCollector
|
||||||
|
|
||||||
|
|
||||||
def _make_llm_response(content="Hello", usage=None):
|
def _make_llm_response(content="Hello", usage=None, response_metadata=None):
|
||||||
"""Create a mock LLM response with a message."""
|
"""Create a mock LLM response with a message."""
|
||||||
msg = MagicMock()
|
msg = MagicMock()
|
||||||
msg.content = content
|
msg.content = content
|
||||||
msg.usage_metadata = usage
|
msg.usage_metadata = usage
|
||||||
|
msg.response_metadata = response_metadata or {}
|
||||||
|
|
||||||
gen = MagicMock()
|
gen = MagicMock()
|
||||||
gen.message = msg
|
gen.message = msg
|
||||||
@ -50,6 +51,32 @@ class TestSubagentTokenCollector:
|
|||||||
assert records[0]["total_tokens"] == 150
|
assert records[0]["total_tokens"] == 150
|
||||||
assert "source_run_id" in records[0]
|
assert "source_run_id" in records[0]
|
||||||
|
|
||||||
|
def test_collects_model_name_from_response_metadata(self):
|
||||||
|
collector = SubagentTokenCollector(caller="subagent:test")
|
||||||
|
usage = {"input_tokens": 100, "output_tokens": 50, "total_tokens": 150}
|
||||||
|
collector.on_llm_end(
|
||||||
|
_make_llm_response("Hi", usage=usage, response_metadata={"model_name": "subagent-model"}),
|
||||||
|
run_id=uuid4(),
|
||||||
|
)
|
||||||
|
|
||||||
|
records = collector.snapshot_records()
|
||||||
|
|
||||||
|
assert len(records) == 1
|
||||||
|
assert records[0]["model_name"] == "subagent-model"
|
||||||
|
|
||||||
|
def test_collects_model_name_from_response_metadata_model_fallback(self):
|
||||||
|
collector = SubagentTokenCollector(caller="subagent:test")
|
||||||
|
usage = {"input_tokens": 100, "output_tokens": 50, "total_tokens": 150}
|
||||||
|
collector.on_llm_end(
|
||||||
|
_make_llm_response("Hi", usage=usage, response_metadata={"model": "provider-model"}),
|
||||||
|
run_id=uuid4(),
|
||||||
|
)
|
||||||
|
|
||||||
|
records = collector.snapshot_records()
|
||||||
|
|
||||||
|
assert len(records) == 1
|
||||||
|
assert records[0]["model_name"] == "provider-model"
|
||||||
|
|
||||||
def test_total_tokens_zero_uses_input_plus_output(self):
|
def test_total_tokens_zero_uses_input_plus_output(self):
|
||||||
collector = SubagentTokenCollector(caller="subagent:test")
|
collector = SubagentTokenCollector(caller="subagent:test")
|
||||||
usage = {"input_tokens": 200, "output_tokens": 100, "total_tokens": 0}
|
usage = {"input_tokens": 200, "output_tokens": 100, "total_tokens": 0}
|
||||||
|
|||||||
431
backend/tests/test_token_usage_by_model.py
Normal file
431
backend/tests/test_token_usage_by_model.py
Normal file
@ -0,0 +1,431 @@
|
|||||||
|
"""Per-model token usage regression tests (issue #3645).
|
||||||
|
|
||||||
|
Covers the full path that powers ``GET /api/threads/{id}/token-usage``'s
|
||||||
|
``by_model`` field:
|
||||||
|
|
||||||
|
* ``RunJournal`` capturing each LLM call's real ``response_metadata.model_name``
|
||||||
|
for both the lead agent / middleware path (``on_llm_end``) and the subagent
|
||||||
|
external-records path (``record_external_llm_usage_records``).
|
||||||
|
* ``RunJournal.get_completion_data`` exposing the per-model breakdown so it can
|
||||||
|
be threaded into the run store on completion.
|
||||||
|
* ``MemoryRunStore`` and ``RunRepository`` (SQLAlchemy) returning the same
|
||||||
|
``by_model`` shape from ``aggregate_tokens_by_thread``, with the invariant
|
||||||
|
``sum(by_model[*].tokens) == total_tokens``.
|
||||||
|
* Legacy rows written before this fix (``token_usage_by_model`` empty) falling
|
||||||
|
back to the old ``model_name + total_tokens`` attribution instead of being
|
||||||
|
silently dropped.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from deerflow.persistence.run import RunRepository
|
||||||
|
from deerflow.runtime.events.store.memory import MemoryRunEventStore
|
||||||
|
from deerflow.runtime.journal import RunJournal
|
||||||
|
from deerflow.runtime.runs.store.memory import MemoryRunStore
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Test doubles
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _make_llm_response(*, usage: dict | None, model_name: str | None = "lead-model"):
|
||||||
|
"""Build a minimal LLM response carrying the bits journal/collector read."""
|
||||||
|
msg = MagicMock()
|
||||||
|
msg.type = "ai"
|
||||||
|
msg.content = ""
|
||||||
|
msg.id = f"msg-{id(msg)}"
|
||||||
|
msg.tool_calls = []
|
||||||
|
msg.invalid_tool_calls = []
|
||||||
|
msg.response_metadata = {} if model_name is None else {"model_name": model_name}
|
||||||
|
msg.usage_metadata = usage
|
||||||
|
msg.additional_kwargs = {}
|
||||||
|
msg.name = None
|
||||||
|
msg.model_dump.return_value = {
|
||||||
|
"content": "",
|
||||||
|
"additional_kwargs": {},
|
||||||
|
"response_metadata": msg.response_metadata,
|
||||||
|
"type": "ai",
|
||||||
|
"name": None,
|
||||||
|
"id": msg.id,
|
||||||
|
"tool_calls": [],
|
||||||
|
"invalid_tool_calls": [],
|
||||||
|
"usage_metadata": usage,
|
||||||
|
}
|
||||||
|
gen = MagicMock()
|
||||||
|
gen.message = msg
|
||||||
|
response = MagicMock()
|
||||||
|
response.generations = [[gen]]
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
def _journal() -> RunJournal:
|
||||||
|
return RunJournal("r1", "t1", MemoryRunEventStore(), flush_threshold=100)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# RunJournal: per-call model accounting
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class TestJournalByModel:
|
||||||
|
def test_lead_agent_call_lands_on_real_model(self) -> None:
|
||||||
|
j = _journal()
|
||||||
|
j.on_llm_end(
|
||||||
|
_make_llm_response(usage={"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, model_name="lead-model"),
|
||||||
|
run_id=uuid4(),
|
||||||
|
parent_run_id=None,
|
||||||
|
tags=["lead_agent"],
|
||||||
|
)
|
||||||
|
data = j.get_completion_data()
|
||||||
|
assert data["token_usage_by_model"] == {
|
||||||
|
"lead-model": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||||
|
}
|
||||||
|
assert data["lead_agent_tokens"] == 15
|
||||||
|
assert data["total_tokens"] == 15
|
||||||
|
|
||||||
|
def test_middleware_call_lands_on_its_own_model(self) -> None:
|
||||||
|
"""A middleware (e.g. title/summarization) on a different model gets its own bucket."""
|
||||||
|
j = _journal()
|
||||||
|
j.on_llm_end(
|
||||||
|
_make_llm_response(usage={"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, model_name="lead-model"),
|
||||||
|
run_id=uuid4(),
|
||||||
|
parent_run_id=None,
|
||||||
|
tags=["lead_agent"],
|
||||||
|
)
|
||||||
|
j.on_llm_end(
|
||||||
|
_make_llm_response(usage={"input_tokens": 4, "output_tokens": 1, "total_tokens": 5}, model_name="title-model"),
|
||||||
|
run_id=uuid4(),
|
||||||
|
parent_run_id=None,
|
||||||
|
tags=["middleware:title"],
|
||||||
|
)
|
||||||
|
data = j.get_completion_data()
|
||||||
|
assert data["token_usage_by_model"] == {
|
||||||
|
"lead-model": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||||
|
"title-model": {"input_tokens": 4, "output_tokens": 1, "total_tokens": 5},
|
||||||
|
}
|
||||||
|
assert data["lead_agent_tokens"] == 15
|
||||||
|
assert data["middleware_tokens"] == 5
|
||||||
|
|
||||||
|
def test_missing_model_name_falls_back_to_unknown(self) -> None:
|
||||||
|
j = _journal()
|
||||||
|
j.on_llm_end(
|
||||||
|
_make_llm_response(usage={"input_tokens": 3, "output_tokens": 2, "total_tokens": 5}, model_name=None),
|
||||||
|
run_id=uuid4(),
|
||||||
|
parent_run_id=None,
|
||||||
|
tags=["lead_agent"],
|
||||||
|
)
|
||||||
|
data = j.get_completion_data()
|
||||||
|
assert data["token_usage_by_model"] == {
|
||||||
|
"unknown": {"input_tokens": 3, "output_tokens": 2, "total_tokens": 5},
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_same_model_aggregates_across_calls(self) -> None:
|
||||||
|
j = _journal()
|
||||||
|
for _ in range(2):
|
||||||
|
j.on_llm_end(
|
||||||
|
_make_llm_response(usage={"input_tokens": 7, "output_tokens": 3, "total_tokens": 10}, model_name="lead-model"),
|
||||||
|
run_id=uuid4(),
|
||||||
|
parent_run_id=None,
|
||||||
|
tags=["lead_agent"],
|
||||||
|
)
|
||||||
|
data = j.get_completion_data()
|
||||||
|
assert data["token_usage_by_model"] == {
|
||||||
|
"lead-model": {"input_tokens": 14, "output_tokens": 6, "total_tokens": 20},
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_subagent_external_records_attribute_to_real_model(self) -> None:
|
||||||
|
"""The fix's headline behavior: subagent on a different model no longer
|
||||||
|
steals tokens from the lead model bucket."""
|
||||||
|
j = _journal()
|
||||||
|
# Lead emits 10 tokens on lead-model.
|
||||||
|
j.on_llm_end(
|
||||||
|
_make_llm_response(usage={"input_tokens": 6, "output_tokens": 4, "total_tokens": 10}, model_name="lead-model"),
|
||||||
|
run_id=uuid4(),
|
||||||
|
parent_run_id=None,
|
||||||
|
tags=["lead_agent"],
|
||||||
|
)
|
||||||
|
# Subagent ran on subagent-model and reports 25 tokens via the
|
||||||
|
# external-records bridge (the path SubagentTokenCollector uses).
|
||||||
|
j.record_external_llm_usage_records(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"source_run_id": "sub-1",
|
||||||
|
"caller": "subagent:general-purpose",
|
||||||
|
"model_name": "subagent-model",
|
||||||
|
"input_tokens": 15,
|
||||||
|
"output_tokens": 10,
|
||||||
|
"total_tokens": 25,
|
||||||
|
},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
data = j.get_completion_data()
|
||||||
|
assert data["token_usage_by_model"] == {
|
||||||
|
"lead-model": {"input_tokens": 6, "output_tokens": 4, "total_tokens": 10},
|
||||||
|
"subagent-model": {"input_tokens": 15, "output_tokens": 10, "total_tokens": 25},
|
||||||
|
}
|
||||||
|
assert data["total_tokens"] == 35
|
||||||
|
# by_caller stays accurate too.
|
||||||
|
assert data["lead_agent_tokens"] == 10
|
||||||
|
assert data["subagent_tokens"] == 25
|
||||||
|
# Invariant the issue calls out: by_model sums to total_tokens.
|
||||||
|
assert sum(b["total_tokens"] for b in data["token_usage_by_model"].values()) == data["total_tokens"]
|
||||||
|
|
||||||
|
def test_subagent_record_without_model_falls_back_to_unknown(self) -> None:
|
||||||
|
j = _journal()
|
||||||
|
j.record_external_llm_usage_records(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"source_run_id": "sub-1",
|
||||||
|
"caller": "subagent:bash",
|
||||||
|
"input_tokens": 5,
|
||||||
|
"output_tokens": 2,
|
||||||
|
"total_tokens": 7,
|
||||||
|
},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
data = j.get_completion_data()
|
||||||
|
assert data["token_usage_by_model"] == {
|
||||||
|
"unknown": {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7},
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_on_llm_end_dedup_does_not_double_count_model(self) -> None:
|
||||||
|
j = _journal()
|
||||||
|
rid = uuid4()
|
||||||
|
usage = {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}
|
||||||
|
j.on_llm_end(_make_llm_response(usage=usage, model_name="lead-model"), run_id=rid, parent_run_id=None, tags=["lead_agent"])
|
||||||
|
# Same langchain run_id firing twice (real callbacks do this) must
|
||||||
|
# not inflate either total_tokens or the per-model bucket.
|
||||||
|
j.on_llm_end(_make_llm_response(usage=usage, model_name="lead-model"), run_id=rid, parent_run_id=None, tags=["lead_agent"])
|
||||||
|
data = j.get_completion_data()
|
||||||
|
assert data["total_tokens"] == 15
|
||||||
|
assert data["token_usage_by_model"] == {
|
||||||
|
"lead-model": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_external_records_dedup_does_not_double_count_model(self) -> None:
|
||||||
|
j = _journal()
|
||||||
|
record = {
|
||||||
|
"source_run_id": "sub-1",
|
||||||
|
"caller": "subagent:general-purpose",
|
||||||
|
"model_name": "subagent-model",
|
||||||
|
"input_tokens": 15,
|
||||||
|
"output_tokens": 10,
|
||||||
|
"total_tokens": 25,
|
||||||
|
}
|
||||||
|
j.record_external_llm_usage_records([record])
|
||||||
|
j.record_external_llm_usage_records([record])
|
||||||
|
data = j.get_completion_data()
|
||||||
|
assert data["subagent_tokens"] == 25
|
||||||
|
assert data["token_usage_by_model"] == {
|
||||||
|
"subagent-model": {"input_tokens": 15, "output_tokens": 10, "total_tokens": 25},
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_track_tokens_disabled_keeps_by_model_empty(self) -> None:
|
||||||
|
store = MemoryRunEventStore()
|
||||||
|
j = RunJournal("r1", "t1", store, track_token_usage=False, flush_threshold=100)
|
||||||
|
j.on_llm_end(
|
||||||
|
_make_llm_response(usage={"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, model_name="lead-model"),
|
||||||
|
run_id=uuid4(),
|
||||||
|
parent_run_id=None,
|
||||||
|
tags=["lead_agent"],
|
||||||
|
)
|
||||||
|
j.record_external_llm_usage_records(
|
||||||
|
[{"source_run_id": "sub", "caller": "subagent:x", "model_name": "sub-model", "input_tokens": 1, "output_tokens": 1, "total_tokens": 2}],
|
||||||
|
)
|
||||||
|
data = j.get_completion_data()
|
||||||
|
assert data["token_usage_by_model"] == {}
|
||||||
|
assert data["total_tokens"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Store aggregation: invariants and parity across MemoryRunStore + RunRepository
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
_THREAD = "thread-by-model"
|
||||||
|
|
||||||
|
|
||||||
|
def _completed_run(
|
||||||
|
run_id: str,
|
||||||
|
*,
|
||||||
|
model_name: str | None,
|
||||||
|
total_tokens: int,
|
||||||
|
lead: int = 0,
|
||||||
|
sub: int = 0,
|
||||||
|
mw: int = 0,
|
||||||
|
by_model: dict | None = None,
|
||||||
|
) -> dict:
|
||||||
|
"""Shape that both stores accept for completion writes (kwargs to update_run_completion)."""
|
||||||
|
return {
|
||||||
|
"run_id": run_id,
|
||||||
|
"model_name": model_name,
|
||||||
|
"completion": {
|
||||||
|
"status": "success",
|
||||||
|
"total_input_tokens": 0,
|
||||||
|
"total_output_tokens": 0,
|
||||||
|
"total_tokens": total_tokens,
|
||||||
|
"llm_call_count": 1,
|
||||||
|
"lead_agent_tokens": lead,
|
||||||
|
"subagent_tokens": sub,
|
||||||
|
"middleware_tokens": mw,
|
||||||
|
"token_usage_by_model": by_model or {},
|
||||||
|
"message_count": 0,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed_run(store, *, run_id: str, model_name: str | None, completion: dict) -> None:
|
||||||
|
await store.put(run_id, thread_id=_THREAD, status="pending", model_name=model_name)
|
||||||
|
await store.update_run_completion(run_id, **completion)
|
||||||
|
|
||||||
|
|
||||||
|
_RUN_FIXTURES = [
|
||||||
|
# 1. Run where subagent and middleware ran on different models than lead.
|
||||||
|
_completed_run(
|
||||||
|
"run-1",
|
||||||
|
model_name="lead-model",
|
||||||
|
total_tokens=300,
|
||||||
|
lead=100,
|
||||||
|
sub=150,
|
||||||
|
mw=50,
|
||||||
|
by_model={
|
||||||
|
"lead-model": {"input_tokens": 60, "output_tokens": 40, "total_tokens": 100},
|
||||||
|
"subagent-model": {"input_tokens": 90, "output_tokens": 60, "total_tokens": 150},
|
||||||
|
"middleware-model": {"input_tokens": 30, "output_tokens": 20, "total_tokens": 50},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
# 2. Another run, lead on a *different* lead model — exercises multi-run merge.
|
||||||
|
_completed_run(
|
||||||
|
"run-2",
|
||||||
|
model_name="lead-model-b",
|
||||||
|
total_tokens=80,
|
||||||
|
lead=80,
|
||||||
|
by_model={
|
||||||
|
"lead-model-b": {"input_tokens": 50, "output_tokens": 30, "total_tokens": 80},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
# 3. Legacy row written before this fix: empty token_usage_by_model. Must
|
||||||
|
# fall back to (model_name, total_tokens) instead of disappearing from
|
||||||
|
# by_model entirely.
|
||||||
|
_completed_run(
|
||||||
|
"run-3",
|
||||||
|
model_name="legacy-model",
|
||||||
|
total_tokens=42,
|
||||||
|
lead=42,
|
||||||
|
by_model={},
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed_all(store) -> None:
|
||||||
|
for fix in _RUN_FIXTURES:
|
||||||
|
await _seed_run(store, run_id=fix["run_id"], model_name=fix["model_name"], completion=fix["completion"])
|
||||||
|
|
||||||
|
|
||||||
|
def _assert_aggregate_shape(agg: dict) -> None:
|
||||||
|
"""Pin the contract that powers /api/threads/{id}/token-usage."""
|
||||||
|
# The headline totals stay the simple SUMs.
|
||||||
|
assert agg["total_tokens"] == 300 + 80 + 42
|
||||||
|
assert agg["total_runs"] == 3
|
||||||
|
assert agg["by_caller"] == {
|
||||||
|
"lead_agent": 100 + 80 + 42,
|
||||||
|
"subagent": 150,
|
||||||
|
"middleware": 50,
|
||||||
|
}
|
||||||
|
# The core fix: subagent / middleware models show up in by_model with their
|
||||||
|
# real tokens; the lead-model bucket is NOT inflated with subagent tokens.
|
||||||
|
assert agg["by_model"]["lead-model"] == {"tokens": 100, "runs": 1}
|
||||||
|
assert agg["by_model"]["subagent-model"] == {"tokens": 150, "runs": 1}
|
||||||
|
assert agg["by_model"]["middleware-model"] == {"tokens": 50, "runs": 1}
|
||||||
|
assert agg["by_model"]["lead-model-b"] == {"tokens": 80, "runs": 1}
|
||||||
|
# Legacy fallback path — empty token_usage_by_model maps to the row's
|
||||||
|
# ``model_name`` with the full total_tokens.
|
||||||
|
assert agg["by_model"]["legacy-model"] == {"tokens": 42, "runs": 1}
|
||||||
|
# Invariant from issue #3645.
|
||||||
|
assert sum(b["tokens"] for b in agg["by_model"].values()) == agg["total_tokens"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_memory_store_by_model_invariant_and_fallback():
|
||||||
|
store = MemoryRunStore()
|
||||||
|
await _seed_all(store)
|
||||||
|
agg = await store.aggregate_tokens_by_thread(_THREAD)
|
||||||
|
_assert_aggregate_shape(agg)
|
||||||
|
|
||||||
|
|
||||||
|
async def _make_sql_repo(tmp_path):
|
||||||
|
from deerflow.persistence.engine import get_session_factory, init_engine
|
||||||
|
|
||||||
|
url = f"sqlite+aiosqlite:///{tmp_path / 'by-model.db'}"
|
||||||
|
await init_engine("sqlite", url=url, sqlite_dir=str(tmp_path))
|
||||||
|
return RunRepository(get_session_factory())
|
||||||
|
|
||||||
|
|
||||||
|
async def _close_sql_engine() -> None:
|
||||||
|
from deerflow.persistence.engine import close_engine
|
||||||
|
|
||||||
|
await close_engine()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_sql_store_by_model_invariant_and_fallback(tmp_path):
|
||||||
|
repo = await _make_sql_repo(tmp_path)
|
||||||
|
try:
|
||||||
|
await _seed_all(repo)
|
||||||
|
agg = await repo.aggregate_tokens_by_thread(_THREAD)
|
||||||
|
_assert_aggregate_shape(agg)
|
||||||
|
finally:
|
||||||
|
await _close_sql_engine()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_memory_and_sql_stores_agree(tmp_path):
|
||||||
|
"""Memory and SQL stores must return byte-identical aggregations so
|
||||||
|
behavior does not silently diverge based on database.backend choice."""
|
||||||
|
mem = MemoryRunStore()
|
||||||
|
sql = await _make_sql_repo(tmp_path)
|
||||||
|
try:
|
||||||
|
await _seed_all(mem)
|
||||||
|
await _seed_all(sql)
|
||||||
|
mem_agg = await mem.aggregate_tokens_by_thread(_THREAD)
|
||||||
|
sql_agg = await sql.aggregate_tokens_by_thread(_THREAD)
|
||||||
|
assert mem_agg == sql_agg
|
||||||
|
finally:
|
||||||
|
await _close_sql_engine()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_include_active_picks_up_running_progress_snapshot(tmp_path):
|
||||||
|
"""``update_run_progress`` must persist ``token_usage_by_model`` so the
|
||||||
|
``include_active=true`` view of /token-usage reflects in-flight tokens."""
|
||||||
|
repo = await _make_sql_repo(tmp_path)
|
||||||
|
try:
|
||||||
|
await repo.put("run-active", thread_id=_THREAD, status="pending")
|
||||||
|
# Transition to running so update_run_progress' status guard fires.
|
||||||
|
await repo.update_status("run-active", "running")
|
||||||
|
await repo.update_run_progress(
|
||||||
|
"run-active",
|
||||||
|
total_tokens=70,
|
||||||
|
total_input_tokens=40,
|
||||||
|
total_output_tokens=30,
|
||||||
|
lead_agent_tokens=70,
|
||||||
|
token_usage_by_model={
|
||||||
|
"lead-model": {"input_tokens": 40, "output_tokens": 30, "total_tokens": 70},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
# Default (completed-only) excludes running runs.
|
||||||
|
completed_only = await repo.aggregate_tokens_by_thread(_THREAD)
|
||||||
|
assert completed_only["total_runs"] == 0
|
||||||
|
assert completed_only["by_model"] == {}
|
||||||
|
|
||||||
|
active = await repo.aggregate_tokens_by_thread(_THREAD, include_active=True)
|
||||||
|
assert active["total_runs"] == 1
|
||||||
|
assert active["by_model"] == {"lead-model": {"tokens": 70, "runs": 1}}
|
||||||
|
assert active["total_tokens"] == 70
|
||||||
|
finally:
|
||||||
|
await _close_sql_engine()
|
||||||
Loading…
x
Reference in New Issue
Block a user