fix(usage): stop double counting scheduled runs; attribute workflow node usage

sum_tokens_in_range now skips the scheduler's per-run rollup rows, whose
tokens are already on the run's per-call rows, so the per-agent 24h limit
and the admin total no longer count scheduled spend twice.

Workflow node LLMs carry the workflow agent's id, so their usage rows are
attributed to the agent instead of landing with a user id only.

usage_totals returns a user's tokens and cost since a window start, split
by interactive and agent-key traffic.
This commit is contained in:
Alex committed 2026-09-21 12:11:11 +01:00
1 parent 43fad2a865
commit 6c42139224
4 files changed
+97 -3

No files matched your search

@@ -393,6 +393,8 @@ class WorkflowEngine:
"prompt": node_prompt,
"chat_history": self.agent.chat_history,
"decoded_token": self.agent.decoded_token,
# Attributes the node's token usage to the workflow agent.
"agent_id": getattr(self.agent, "agent_id", None),
"json_schema": node_json_schema,
"retrieved_docs": node_docs,
# A template that interpolates the documents itself already carries
+34 -3
View File
@@ -118,9 +118,13 @@ class TokenUsageRepository:
user_id: Optional[str] = None,
api_key: Optional[str] = None,
) -> int:
"""Total (prompt + generated) tokens in the given time range."""
clauses = ["timestamp >= :start", "timestamp <= :end"]
params: dict = {"start": start, "end": end}
"""Total (prompt + generated) tokens in the given time range.
Run-level rollup rows (``ROLLUP_SOURCES``) are excluded: their tokens
are already counted on the run's per-call rows.
"""
clauses = ["timestamp >= :start", "timestamp <= :end", "source <> ALL(:rollup_sources)"]
params: dict = {"start": start, "end": end, "rollup_sources": list(self.ROLLUP_SOURCES)}
if user_id is not None:
clauses.append("user_id = :user_id")
params["user_id"] = user_id
@@ -134,6 +138,33 @@ class TokenUsageRepository:
)
return result.scalar()
def usage_totals(self, *, user_id: str, start: datetime, bucket: str = "all") -> tuple[int, float]:
"""Return ``(tokens, cost_usd)`` a user has consumed since ``start``.
Args:
user_id: The billable user (auth ``sub``).
start: Inclusive window start.
bucket: ``all``, ``direct`` (rows without an agent key) or
``agent`` (rows with one).
Rollup rows are excluded; side-channel calls count, they are real spend.
"""
clauses = ["user_id = :user_id", "timestamp >= :start", "source <> ALL(:rollup_sources)"]
if bucket == "direct":
clauses.append("api_key IS NULL")
elif bucket == "agent":
clauses.append("api_key IS NOT NULL")
elif bucket != "all":
raise ValueError(f"unknown usage bucket: {bucket!r}")
row = self._conn.execute(
text(
"SELECT COALESCE(SUM(prompt_tokens + generated_tokens), 0), COALESCE(SUM(cost), 0) "
f"FROM token_usage WHERE {' AND '.join(clauses)}"
),
{"user_id": user_id, "start": start, "rollup_sources": list(self.ROLLUP_SOURCES)},
).one()
return int(row[0]), float(row[1])
# Token usage written outside a user-initiated request (conversation
# title generation, history compression, RAG question condensing,
# provider fallback). Mirrors the exclusion list in ``count_in_range``.
+21
View File
@@ -214,6 +214,27 @@ class TestWorkflowEngineAgenticNode:
assert engine.state["node_agent_agentic_output"] == "agentic answer"
assert engine.state["result"] == "agentic answer"
def test_node_usage_is_attributed_to_the_workflow_agent(self, monkeypatch):
engine = create_engine()
engine.agent.agent_id = "11111111-1111-1111-1111-111111111111"
node = create_agent_node(node_id="agent_attr", agent_type="classic")
captured: Dict[str, Any] = {}
def capture_create(**kwargs):
captured.update(kwargs)
return StubNodeAgent([{"answer": "ok"}])
monkeypatch.setattr(WorkflowNodeAgentFactory, "create", staticmethod(capture_create))
monkeypatch.setattr(
"docsgpt.core.model_utils.get_api_key_for_provider",
lambda _provider: None,
)
list(engine._execute_agent_node(node))
assert captured["agent_id"] == "11111111-1111-1111-1111-111111111111"
def test_agentic_node_passes_retriever_config(self, monkeypatch):
engine = create_engine()
# The node-source authorization gate is exercised separately;
@@ -45,6 +45,46 @@ class TestInsert:
assert [float(c) for c in costs] == [0.0, 0.00012345]
class TestRollupExclusion:
def test_sum_tokens_ignores_schedule_rollups(self, pg_conn):
repo = _repo(pg_conn)
repo.insert(user_id="u-roll", prompt_tokens=10, generated_tokens=5, source="agent_stream")
repo.insert(user_id="u-roll", prompt_tokens=10, generated_tokens=5, source="schedule")
total = repo.sum_tokens_in_range(
start=_now() - timedelta(minutes=1), end=_now() + timedelta(minutes=1), user_id="u-roll"
)
assert total == 15
class TestUsageTotals:
def _seed(self, repo):
repo.insert(user_id="u-tot", prompt_tokens=100, generated_tokens=10, cost=0.5)
repo.insert(user_id="u-tot", api_key="k", prompt_tokens=20, generated_tokens=2, cost=0.25)
repo.insert(user_id="u-tot", prompt_tokens=7, generated_tokens=0, cost=0.125, source="title")
repo.insert(user_id="u-tot", prompt_tokens=999, generated_tokens=0, source="schedule")
repo.insert(user_id="u-other", prompt_tokens=999, generated_tokens=0, cost=9)
repo.insert(
user_id="u-tot", prompt_tokens=999, generated_tokens=0, cost=9,
timestamp=_now() - timedelta(days=40),
)
@pytest.mark.parametrize(
"bucket, expected",
[("all", (139, 0.875)), ("direct", (117, 0.625)), ("agent", (22, 0.25))],
)
def test_totals_per_bucket(self, pg_conn, bucket, expected):
repo = _repo(pg_conn)
self._seed(repo)
assert repo.usage_totals(user_id="u-tot", start=_now() - timedelta(days=1), bucket=bucket) == expected
def test_no_usage_is_zero(self, pg_conn):
assert _repo(pg_conn).usage_totals(user_id="nobody", start=_now() - timedelta(days=1)) == (0, 0.0)
def test_unknown_bucket_rejected(self, pg_conn):
with pytest.raises(ValueError):
_repo(pg_conn).usage_totals(user_id="u", start=_now(), bucket="nope")
class TestReassignApiKey:
def test_rewrites_and_preserves_rate_limit_window(self, pg_conn):
# Rotating an agent key must carry the running 24h usage window over