diff --git a/docsgpt/agents/workflows/workflow_engine.py b/docsgpt/agents/workflows/workflow_engine.py index 54fa6e67..15876c82 100644 --- a/docsgpt/agents/workflows/workflow_engine.py +++ b/docsgpt/agents/workflows/workflow_engine.py @@ -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 diff --git a/docsgpt/storage/db/repositories/token_usage.py b/docsgpt/storage/db/repositories/token_usage.py index 83c6e041..17e546fe 100644 --- a/docsgpt/storage/db/repositories/token_usage.py +++ b/docsgpt/storage/db/repositories/token_usage.py @@ -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``. diff --git a/tests/agents/test_workflow_agent_types.py b/tests/agents/test_workflow_agent_types.py index 41bf7cd4..18bf22d3 100644 --- a/tests/agents/test_workflow_agent_types.py +++ b/tests/agents/test_workflow_agent_types.py @@ -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; diff --git a/tests/storage/db/repositories/test_token_usage.py b/tests/storage/db/repositories/test_token_usage.py index 02b2eb98..9698c610 100644 --- a/tests/storage/db/repositories/test_token_usage.py +++ b/tests/storage/db/repositories/test_token_usage.py @@ -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