mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
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:
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
|
||||
|
||||
@@ -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``.
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in new issue
Block a user