Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions agent/app/routers/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
should_emit_langgraph_event,
thinking_delta_from_cumulative,
)
from app.usage import UsageCollector

router = APIRouter(prefix="/agent", tags=["agent"])
logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -73,13 +74,15 @@ async def _stream_agent(request: ChatRequest) -> AsyncGenerator[str, None]:
open_tools: dict[str, dict[str, Any]] = {}
thinking_acc = ""
emitted_tool_ids: set[str] = set()
usage = UsageCollector(default_model=request.config.model)
try:
agent = _build_request_agent(request)
async for event in agent.astream_events(
{"messages": _request_to_agent_messages(request)},
config=_thread_config(request),
version="v2",
):
usage.observe(event)
if not should_emit_langgraph_event(event, agent_mode=request.agent_mode):
continue
for sse_data in _format_sse_events(
Expand All @@ -95,6 +98,8 @@ async def _stream_agent(request: ChatRequest) -> AsyncGenerator[str, None]:
yield sse_data
for frame in _close_open_tools(open_tools, reason="Stream ended without tool_end"):
yield frame
if usage_payload := usage.event_payload():
yield _sse("usage", usage_payload)
yield _sse("done", {"finished": True})
except Exception as exc:
tb = traceback.format_exc()
Expand All @@ -107,6 +112,8 @@ async def _stream_agent(request: ChatRequest) -> AsyncGenerator[str, None]:
)
for frame in _close_open_tools(open_tools, reason=str(exc) or "stream error"):
yield frame
if usage_payload := usage.event_payload():
yield _sse("usage", usage_payload)
yield _sse("error", {"message": str(exc)})


Expand Down
261 changes: 261 additions & 0 deletions agent/app/usage.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,261 @@
"""Normalize LangChain/LangGraph token metadata for the desktop usage ledger."""

from __future__ import annotations

from dataclasses import dataclass, field
from typing import Any

USAGE_SSE_SCHEMA_VERSION = 1


@dataclass(slots=True)
class UsageMeasurement:
run_id: str
model: str | None = None
input_tokens: int | None = None
output_tokens: int | None = None
total_tokens: int | None = None
cache_read_tokens: int | None = None
cache_creation_tokens: int | None = None
reasoning_tokens: int | None = None
source: str = "provider_reported"
provider_metadata: dict[str, Any] = field(default_factory=dict)

def completeness(self) -> int:
return sum(
value is not None
for value in (
self.input_tokens,
self.output_tokens,
self.total_tokens,
self.cache_read_tokens,
self.cache_creation_tokens,
self.reasoning_tokens,
)
)

def to_dict(self) -> dict[str, Any]:
return {
"run_id": self.run_id,
"model": self.model,
"input_tokens": self.input_tokens,
"output_tokens": self.output_tokens,
"total_tokens": self.total_tokens,
"cache_read_tokens": self.cache_read_tokens,
"cache_creation_tokens": self.cache_creation_tokens,
"reasoning_tokens": self.reasoning_tokens,
"source": self.source,
"provider_metadata": self.provider_metadata,
}


class UsageCollector:
"""Upsert one most-complete measurement per model run_id."""

def __init__(self, default_model: str | None = None) -> None:
self._default_model = default_model
self._measurements: dict[str, UsageMeasurement] = {}

def observe(self, event: dict[str, Any]) -> None:
kind = str(event.get("event") or "")
if kind not in {
"on_chat_model_stream",
"on_chat_model_end",
"on_llm_stream",
"on_llm_end",
}:
return
run_id = str(event.get("run_id") or event.get("id") or "").strip()
if not run_id:
return

data = event.get("data") or {}
candidates = [
_get(data, "chunk"),
_get(data, "output"),
data,
]
usage, usage_keys = _best_usage(candidates)
if not usage:
return

model = _extract_model(event, candidates) or self._default_model
incoming = UsageMeasurement(
run_id=run_id,
model=model,
input_tokens=_token(usage, "input_tokens", "prompt_tokens"),
output_tokens=_token(usage, "output_tokens", "completion_tokens"),
total_tokens=_token(usage, "total_tokens"),
cache_read_tokens=_token(
usage,
"cache_read_tokens",
"cached_input_tokens",
"cache_read_input_tokens",
),
cache_creation_tokens=_token(
usage,
"cache_creation_tokens",
"cache_creation_input_tokens",
"cache_write_tokens",
),
reasoning_tokens=_token(usage, "reasoning_tokens"),
provider_metadata={"raw_usage_keys": sorted(usage_keys)},
)
if incoming.completeness() == 0:
return
current = self._measurements.get(run_id)
self._measurements[run_id] = _merge(current, incoming)

def event_payload(self) -> dict[str, Any] | None:
if not self._measurements:
return None
return {
"schema_version": USAGE_SSE_SCHEMA_VERSION,
"measurements": [
measurement.to_dict()
for measurement in self._measurements.values()
],
}


def _merge(
current: UsageMeasurement | None,
incoming: UsageMeasurement,
) -> UsageMeasurement:
if current is None:
return incoming
# End events commonly add totals/cache fields. Merge field-wise rather than
# summing: stream and end are two views of the same run, not two calls.
for name in (
"input_tokens",
"output_tokens",
"total_tokens",
"cache_read_tokens",
"cache_creation_tokens",
"reasoning_tokens",
):
value = getattr(incoming, name)
if value is not None:
setattr(current, name, value)
if incoming.model:
current.model = incoming.model
keys = set(current.provider_metadata.get("raw_usage_keys", []))
keys.update(incoming.provider_metadata.get("raw_usage_keys", []))
current.provider_metadata["raw_usage_keys"] = sorted(keys)
return current


def _best_usage(candidates: list[Any]) -> tuple[dict[str, Any], set[str]]:
best: dict[str, Any] = {}
best_keys: set[str] = set()
for candidate in candidates:
for mapping in _usage_mappings(candidate):
keys = set(mapping)
score = sum(
_token(mapping, key) is not None
for key in (
"input_tokens",
"prompt_tokens",
"output_tokens",
"completion_tokens",
"total_tokens",
"cache_read_tokens",
"cached_input_tokens",
"cache_creation_tokens",
"reasoning_tokens",
)
)
if score > sum(value is not None for value in best.values()):
best = mapping
best_keys = keys
elif score > 0:
for key, value in mapping.items():
if value is not None:
best[key] = value
best_keys.update(keys)
return best, best_keys


def _usage_mappings(value: Any) -> list[dict[str, Any]]:
if value is None:
return []
mappings: list[dict[str, Any]] = []
for key in ("usage_metadata", "usage", "token_usage"):
raw = _get(value, key)
mapping = _as_mapping(raw)
if mapping:
mappings.append(mapping)
response_metadata = _get(value, "response_metadata")
if response_metadata:
for key in ("usage", "token_usage"):
mapping = _as_mapping(_get(response_metadata, key))
if mapping:
mappings.append(mapping)
llm_output = _get(value, "llm_output")
if llm_output:
for key in ("usage", "token_usage"):
mapping = _as_mapping(_get(llm_output, key))
if mapping:
mappings.append(mapping)
if isinstance(value, (list, tuple)):
for item in value:
mappings.extend(_usage_mappings(item))
return mappings


def _extract_model(event: dict[str, Any], candidates: list[Any]) -> str | None:
metadata = event.get("metadata") or {}
for value in (
_get(metadata, "ls_model_name"),
_get(metadata, "model"),
*(
candidate_model
for candidate in candidates
for candidate_model in (
_get(candidate, "model"),
_get(candidate, "model_name"),
_get(_get(candidate, "response_metadata"), "model_name"),
_get(_get(candidate, "response_metadata"), "model"),
)
),
):
if value:
return str(value)
return None


def _token(mapping: dict[str, Any], *keys: str) -> int | None:
for key in keys:
value = mapping.get(key)
if value is None:
continue
try:
token = int(value)
except (TypeError, ValueError):
continue
if token >= 0:
return token
return None


def _get(value: Any, key: str) -> Any:
if value is None:
return None
if isinstance(value, dict):
return value.get(key)
return getattr(value, key, None)


def _as_mapping(value: Any) -> dict[str, Any]:
if isinstance(value, dict):
return value
if value is None:
return {}
if hasattr(value, "model_dump"):
dumped = value.model_dump()
return dumped if isinstance(dumped, dict) else {}
return {
key: getattr(value, key)
for key in dir(value)
if not key.startswith("_") and not callable(getattr(value, key, None))
}
4 changes: 4 additions & 0 deletions agent/tests/test_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,10 +223,12 @@ def test_thread_config_raises_langgraph_recursion_limit():
async def test_agent_stream_emits_token_and_done(client):
class Chunk:
content = "Hi"
usage_metadata = {"input_tokens": 3, "output_tokens": 1, "total_tokens": 4}

async def fake_events(*_args, **_kwargs):
yield {
"event": "on_chat_model_stream",
"run_id": "run-usage",
"data": {"chunk": Chunk()},
}

Expand All @@ -246,6 +248,8 @@ async def fake_events(*_args, **_kwargs):
assert response.status_code == 200
body = response.text
assert "event: token" in body
assert "event: usage" in body
assert body.index("event: usage") < body.index("event: done")
assert "event: done" in body


Expand Down
73 changes: 73 additions & 0 deletions agent/tests/test_usage.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
from types import SimpleNamespace

from app.usage import UsageCollector


def test_collector_dedupes_stream_and_end_by_run_id() -> None:
collector = UsageCollector(default_model="fallback")
collector.observe(
{
"event": "on_chat_model_stream",
"run_id": "run-1",
"data": {"chunk": SimpleNamespace(usage_metadata={"input_tokens": 10})},
}
)
collector.observe(
{
"event": "on_chat_model_end",
"run_id": "run-1",
"metadata": {"ls_model_name": "actual-model"},
"data": {
"output": SimpleNamespace(
usage_metadata={
"input_tokens": 10,
"output_tokens": 4,
"total_tokens": 14,
}
)
},
}
)
payload = collector.event_payload()
assert payload is not None
assert payload["schema_version"] == 1
assert len(payload["measurements"]) == 1
assert payload["measurements"][0]["total_tokens"] == 14
assert payload["measurements"][0]["model"] == "actual-model"


def test_collector_preserves_distinct_models_and_cache_details() -> None:
collector = UsageCollector()
for run_id, model, total in (("a", "main", 5), ("b", "subagent", 7)):
collector.observe(
{
"event": "on_chat_model_end",
"run_id": run_id,
"metadata": {"ls_model_name": model},
"data": {
"output": SimpleNamespace(
response_metadata={
"token_usage": {
"prompt_tokens": total - 2,
"completion_tokens": 2,
"total_tokens": total,
"cached_input_tokens": 1,
}
}
)
},
}
)
payload = collector.event_payload()
assert payload is not None
assert [item["model"] for item in payload["measurements"]] == ["main", "subagent"]
assert payload["measurements"][0]["cache_read_tokens"] == 1


def test_collector_ignores_tools_and_missing_usage_metadata() -> None:
collector = UsageCollector()
collector.observe({"event": "on_tool_end", "run_id": "tool", "data": {}})
collector.observe(
{"event": "on_chat_model_end", "run_id": "model", "data": {"output": {}}}
)
assert collector.event_payload() is None
Loading
Loading