From 5785ffda9457248b0ca88fad02af875d18e8a255 Mon Sep 17 00:00:00 2001 From: "Deborah (Precious) Aregbesola" Date: Wed, 30 Sep 2026 12:36:28 +0100 Subject: [PATCH 1/4] fix: Track per-conversation LLM token usage (#31) --- app/chat/api.py | 58 +++++++++++++++++++++++++++-- app/models/message.py | 4 +- app/services/exporting.py | 4 +- app/services/providers/anthropic.py | 1 + app/services/providers/base.py | 51 +++++++++++++++++++++++++ app/services/providers/mock.py | 16 ++++++-- app/services/providers/openai.py | 12 ++++++ app/services/token_usage.py | 2 +- 8 files changed, 136 insertions(+), 12 deletions(-) diff --git a/app/chat/api.py b/app/chat/api.py index 4b127ea..34cc8a3 100644 --- a/app/chat/api.py +++ b/app/chat/api.py @@ -7,6 +7,7 @@ GET /api/conversations/ conversation with its messages POST /api/conversations//messages send a message (LLM reply) DELETE /api/conversations/ delete (cascade) + GET /api/usage per-user token usage summary All routes require authentication and are owner-scoped. Errors use RFC 7807 (``application/problem+json``) documents so API clients never receive HTML error @@ -23,7 +24,7 @@ from app.chat import routes as chat_routes from app.extensions import db -from app.models import Conversation, Message +from app.models import Conversation, Message, TokenUsage from app.services import ratelimit from app.services.llm import LLMProviderError, provider_status from app.services.provider_config import ProviderSettingsError, apply_settings, build_provider @@ -85,6 +86,23 @@ def _owned_conversation(conversation_id: int) -> Conversation | None: return Conversation.query.filter_by(id=conversation_id, user_id=current_user.id).first() +def _usage_totals(conversation: Conversation) -> dict: + """Sum token usage across the conversation's assistant messages.""" + prompt = completion = total = 0 + for message in conversation.messages: + usage = getattr(message, "token_usage", None) + if usage is None: + continue + prompt += usage.prompt_tokens or 0 + completion += usage.completion_tokens or 0 + total += usage.total_tokens or 0 + return { + "prompt_tokens": prompt, + "completion_tokens": completion, + "total_tokens": total, + } + + def _json_object() -> dict | None: """Return the request body as a dict, or ``None`` when it is not one.""" data = request.get_json(silent=True) @@ -140,6 +158,7 @@ def get_conversation(conversation_id: int): return _problem(404, "Conversation not found.", "No such conversation exists.") payload = conversation.to_dict() payload["messages"] = [message.to_dict() for message in conversation.messages] + payload["usage"] = _usage_totals(conversation) return jsonify(payload) @@ -206,13 +225,44 @@ def send_message(conversation_id: int): messages = chat_routes._conversation_messages(conversation, context_messages) try: provider = RetryingProvider(build_provider(current_user, conversation.provider)) - reply = provider.chat(messages, **chat_routes._generation_kwargs(conversation)).content + completion = provider.chat(messages, **chat_routes._generation_kwargs(conversation)) except LLMProviderError as exc: db.session.rollback() return _problem(502, "Provider error.", str(exc)) - conversation.messages.append(Message(role="assistant", content=reply)) + assistant_message = Message(role="assistant", content=completion.content) + usage = getattr(completion, "usage", None) + if usage is not None: + assistant_message.token_usage = TokenUsage( + prompt_tokens=getattr(usage, "prompt_tokens", 0) or 0, + completion_tokens=getattr(usage, "completion_tokens", 0) or 0, + total_tokens=getattr(usage, "total_tokens", 0) or 0, + ) + conversation.messages.append(assistant_message) if conversation.title == "New conversation": conversation.title = content.strip()[:60] or "New conversation" db.session.commit() - return jsonify({"assistant_message": conversation.messages[-1].to_dict()}), 201 + payload = {"assistant_message": conversation.messages[-1].to_dict()} + payload["usage"] = _usage_totals(conversation) + return jsonify(payload), 201 + + +@bp.route("/usage", methods=["GET"]) +@_login_required +@_rate_limit("usage") +def usage_summary(): + """Return the current user's total token consumption.""" + conversations = Conversation.query.filter_by(user_id=current_user.id).all() + prompt = completion = total = 0 + for conversation in conversations: + totals = _usage_totals(conversation) + prompt += totals["prompt_tokens"] + completion += totals["completion_tokens"] + total += totals["total_tokens"] + return jsonify( + { + "prompt_tokens": prompt, + "completion_tokens": completion, + "total_tokens": total, + } + ) diff --git a/app/models/message.py b/app/models/message.py index 402eaae..c16dda0 100644 --- a/app/models/message.py +++ b/app/models/message.py @@ -8,7 +8,7 @@ class Message(db.Model): """A single message exchanged within a conversation. - ``role`` is one of ``user`` or ``assistant``. Prompt text and assistant + ``role`` is one of ``user``or ``assistant``. Prompt text and assistant responses are stored verbatim so conversation history can be replayed or exported. """ @@ -30,7 +30,7 @@ class Message(db.Model): completion_tokens = db.Column(db.Integer, nullable=True) total_tokens = db.Column(db.Integer, nullable=True) created_at = db.Column( - db.DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC) + db.DateTime(timezone=True), nullable=False, default=lambda: datetime.now(TCT) ) conversation = db.relationship("Conversation", back_populates="messages") diff --git a/app/services/exporting.py b/app/services/exporting.py index df4df22..f61f09b 100644 --- a/app/services/exporting.py +++ b/app/services/exporting.py @@ -31,6 +31,8 @@ from app.models import ProjectFile from app.services.importing import sanitize_member_path +from app.services.token_usage import sum_usage + PLACEHOLDER_SUFFIX = ".PLACEHOLDER.txt" @@ -98,7 +100,7 @@ def iter_export_zip(project): The generator writes one entry per stored file (plus the manifest) and never materializes the whole archive, so exports stay memory-bounded for - any project size. Content is read per-row via ``.yield_per`` so SQLAlchemy + any project size. Content is read per-row via `.yield_per`` so SQLAlchemy streams rows from the database instead of loading them all at once. """ buffer = io.BytesIO() diff --git a/app/services/providers/anthropic.py b/app/services/providers/anthropic.py index 3eb925f..972d0e1 100644 --- a/app/services/providers/anthropic.py +++ b/app/services/providers/anthropic.py @@ -179,6 +179,7 @@ def chat( model=data.get("model") or payload["model"], prompt_tokens=usage.get("input_tokens"), completion_tokens=usage.get("output_tokens"), + total_tokens=(usage.get("input_tokens") or 0) + (usage.get("output_tokens") or 0), latency_seconds=time.perf_counter() - started, ) diff --git a/app/services/providers/base.py b/app/services/providers/base.py index 790cf24..e0e1683 100644 --- a/app/services/providers/base.py +++ b/app/services/providers/base.py @@ -80,6 +80,50 @@ def prepare_messages(messages: Iterable[Any], *, supports_vision: bool) -> list[ return prepared +@dataclass(frozen=True) +class TokenUsage: + """Vendor-neutral token accounting for a single completion (issue #2).""" + + prompt_tokens: int | None = None + completion_tokens: int | None = None + total_tokens: int | None = None + + @classmethod + def from_counts( + cls, + prompt_tokens: int | None, + completion_tokens: int | None, + total_tokens: int | None = None, + ) -> "TokenUsage": + """Build usage, deriving ``total_tokens`` when the provider omits it.""" + if total_tokens is None and ( + prompt_tokens is not None or completion_tokens is not None + ): + total_tokens = (prompt_tokens or 0) + (completion_tokens or 0) + return cls( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=total_tokens, + ) + + @property + def has_usage(self) -> bool: + """Whether any token counts were reported.""" + return ( + self.prompt_tokens is not None + or self.completion_tokens is not None + or self.total_tokens is not None + ) + + def to_dict(self) -> dict[str, Any]: + """Serialize the usage for logging/telemetry.""" + return { + "prompt_tokens": self.prompt_tokens, + "completion_tokens": self.completion_tokens, + "total_tokens": self.total_tokens, + } + + @dataclass(frozen=True) class ProviderResponse: """Uniform, vendor-neutral result of a single completion.""" @@ -97,6 +141,13 @@ def total_tokens(self) -> int | None: return None return (self.prompt_tokens or 0) + (self.completion_tokens or 0) + @property + def usage(self) -> TokenUsage: + """Token usage metadata for this completion (issue #2).""" + return TokenUsage.from_counts( + self.prompt_tokens, self.completion_tokens, self.total_tokens + ) + def to_dict(self) -> dict[str, Any]: """Serialize the response for logging/telemetry.""" return { diff --git a/app/services/providers/mock.py b/app/services/providers/mock.py index de307b6..564a6a5 100644 --- a/app/services/providers/mock.py +++ b/app/services/providers/mock.py @@ -2,7 +2,7 @@ Used by the test suite and local development so the full pipeline (models, routes, SSE streaming, UI) can run without network access or API keys. It -implements the same :class:`LLMProvider` contract as the real providers, which +implements the same :Class:``LLMProvider`` contract as the real providers, which is what the shared contract tests exercise. """ @@ -19,6 +19,7 @@ message_role, prepare_messages, ) +from app.services.token_usage import usage_from_response, usage_from_text class MockProvider(LLMProvider): @@ -54,11 +55,17 @@ def chat( if self.delay: time.sleep(self.delay) prepared = prepare_messages(messages, supports_vision=self.supports_vision) - return ProviderResponse( - content=self._respond(prepared), + content = self._respond(prepared) + response = ProviderResponse( + content=content, model=model or self.models[0], latency_seconds=time.perf_counter() - started, ) + usage = usage_from_response(response, prepared) + response.prompt_tokens = usage["prompt_tokens"] + response.completion_tokens = usage["completion_tokens"] + response.total_tokens = usage["total_tokens"] + return response def stream( self, @@ -67,7 +74,8 @@ def stream( model: str | None = None, params: dict | None = None, ) -> Iterator[str]: - text = self._respond(prepare_messages(messages, supports_vision=self.supports_vision)) + prepared = prepare_messages(messages, supports_vision=self.supports_vision) + text = self._respond(prepared) for word in text.split(" "): if self.delay: time.sleep(self.delay) diff --git a/app/services/providers/openai.py b/app/services/providers/openai.py index f59999c..81cebb9 100644 --- a/app/services/providers/openai.py +++ b/app/services/providers/openai.py @@ -100,6 +100,8 @@ def _payload( "temperature": self.temperature, "stream": stream, } + if stream: + payload["stream_options"] = {"include_usage": True} if params: payload.update(params) return payload @@ -160,6 +162,7 @@ def chat( model=data.get("model") or payload["model"], prompt_tokens=usage.get("prompt_tokens"), completion_tokens=usage.get("completion_tokens"), + total_tokens=usage.get("total_tokens"), latency_seconds=time.perf_counter() - started, ) @@ -174,6 +177,7 @@ def stream( payload = self._payload(messages, model=model, params=params, stream=True) response = self._post(payload, stream=True) self._raise_for_status(response) + usage: dict[str, Any] = {} try: for line in response.iter_lines(decode_unicode=True): if not line or not line.startswith("data: "): @@ -183,6 +187,9 @@ def stream( break try: chunk = json.loads(chunk_payload) + chunk_usage = chunk.get("usage") + if chunk_usage: + usage = chunk_usage delta = chunk["choices"][0]["delta"].get("content", "") except (KeyError, IndexError, TypeError, ValueError): delta = "" @@ -194,3 +201,8 @@ def stream( raise ProviderUnavailableError( f"OpenAI stream failed: {exc}", provider=self.name ) from exc + self.last_usage = { + "prompt_tokens": usage.get("prompt_tokens"), + "completion_tokens": usage.get("completion_tokens"), + "total_tokens": usage.get("total_tokens"), + } diff --git a/app/services/token_usage.py b/app/services/token_usage.py index a309150..cdf1e12 100644 --- a/app/services/token_usage.py +++ b/app/services/token_usage.py @@ -2,7 +2,7 @@ Provider responses carry real prompt/completion token counts when the vendor reports them; when they do not — the offline mock provider, or a streamed reply -— a small, dependency-free estimate is used so every message still records a + — a small, dependency-free estimate is used so every message still records a usage figure. """ From 6492d0d06fa2ff08ff246f16c9e2fbe42669902a Mon Sep 17 00:00:00 2001 From: "Deborah (Precious) Aregbesola" <330907859+thegeneraldeborah-dev@users.noreply.github.com> Date: Tue, 6 Oct 2026 12:19:14 +0000 Subject: [PATCH 2/4] fix(ci): align per-conversation token usage with the message columns * app/chat/api.py: persist prompt/completion/total tokens on the assistant message, expose GET /api/usage and include usage totals in the conversation payload. * app/services/providers/base.py: add the `TokenUsage` dataclass (`from_counts`, `has_usage`, `to_dict`) and a real `total_tokens` field on `ProviderResponse`. * app/services/providers/mock.py: report usage through the token_usage helpers. * app/models/message.py: the branch carried a reverted revision (`datetime.now(TCT)`); app/services/exporting.py: drop the reverted unused `sum_usage` import. --- app/chat/api.py | 21 +++++++++------------ app/models/message.py | 4 ++-- app/services/exporting.py | 4 +--- app/services/providers/base.py | 18 ++++-------------- app/services/providers/mock.py | 13 ++++++------- 5 files changed, 22 insertions(+), 38 deletions(-) diff --git a/app/chat/api.py b/app/chat/api.py index 34cc8a3..20c5ca8 100644 --- a/app/chat/api.py +++ b/app/chat/api.py @@ -24,7 +24,7 @@ from app.chat import routes as chat_routes from app.extensions import db -from app.models import Conversation, Message, TokenUsage +from app.models import Conversation, Message from app.services import ratelimit from app.services.llm import LLMProviderError, provider_status from app.services.provider_config import ProviderSettingsError, apply_settings, build_provider @@ -90,12 +90,11 @@ def _usage_totals(conversation: Conversation) -> dict: """Sum token usage across the conversation's assistant messages.""" prompt = completion = total = 0 for message in conversation.messages: - usage = getattr(message, "token_usage", None) - if usage is None: + if message.total_tokens is None: continue - prompt += usage.prompt_tokens or 0 - completion += usage.completion_tokens or 0 - total += usage.total_tokens or 0 + prompt += message.prompt_tokens or 0 + completion += message.completion_tokens or 0 + total += message.total_tokens or 0 return { "prompt_tokens": prompt, "completion_tokens": completion, @@ -232,12 +231,10 @@ def send_message(conversation_id: int): assistant_message = Message(role="assistant", content=completion.content) usage = getattr(completion, "usage", None) - if usage is not None: - assistant_message.token_usage = TokenUsage( - prompt_tokens=getattr(usage, "prompt_tokens", 0) or 0, - completion_tokens=getattr(usage, "completion_tokens", 0) or 0, - total_tokens=getattr(usage, "total_tokens", 0) or 0, - ) + if usage is not None and usage.has_usage: + assistant_message.prompt_tokens = usage.prompt_tokens or 0 + assistant_message.completion_tokens = usage.completion_tokens or 0 + assistant_message.total_tokens = usage.total_tokens or 0 conversation.messages.append(assistant_message) if conversation.title == "New conversation": conversation.title = content.strip()[:60] or "New conversation" diff --git a/app/models/message.py b/app/models/message.py index c16dda0..402eaae 100644 --- a/app/models/message.py +++ b/app/models/message.py @@ -8,7 +8,7 @@ class Message(db.Model): """A single message exchanged within a conversation. - ``role`` is one of ``user``or ``assistant``. Prompt text and assistant + ``role`` is one of ``user`` or ``assistant``. Prompt text and assistant responses are stored verbatim so conversation history can be replayed or exported. """ @@ -30,7 +30,7 @@ class Message(db.Model): completion_tokens = db.Column(db.Integer, nullable=True) total_tokens = db.Column(db.Integer, nullable=True) created_at = db.Column( - db.DateTime(timezone=True), nullable=False, default=lambda: datetime.now(TCT) + db.DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC) ) conversation = db.relationship("Conversation", back_populates="messages") diff --git a/app/services/exporting.py b/app/services/exporting.py index f61f09b..df4df22 100644 --- a/app/services/exporting.py +++ b/app/services/exporting.py @@ -31,8 +31,6 @@ from app.models import ProjectFile from app.services.importing import sanitize_member_path -from app.services.token_usage import sum_usage - PLACEHOLDER_SUFFIX = ".PLACEHOLDER.txt" @@ -100,7 +98,7 @@ def iter_export_zip(project): The generator writes one entry per stored file (plus the manifest) and never materializes the whole archive, so exports stay memory-bounded for - any project size. Content is read per-row via `.yield_per`` so SQLAlchemy + any project size. Content is read per-row via ``.yield_per`` so SQLAlchemy streams rows from the database instead of loading them all at once. """ buffer = io.BytesIO() diff --git a/app/services/providers/base.py b/app/services/providers/base.py index e0e1683..695db4f 100644 --- a/app/services/providers/base.py +++ b/app/services/providers/base.py @@ -94,11 +94,9 @@ def from_counts( prompt_tokens: int | None, completion_tokens: int | None, total_tokens: int | None = None, - ) -> "TokenUsage": + ) -> TokenUsage: """Build usage, deriving ``total_tokens`` when the provider omits it.""" - if total_tokens is None and ( - prompt_tokens is not None or completion_tokens is not None - ): + if total_tokens is None and (prompt_tokens is not None or completion_tokens is not None): total_tokens = (prompt_tokens or 0) + (completion_tokens or 0) return cls( prompt_tokens=prompt_tokens, @@ -132,21 +130,13 @@ class ProviderResponse: model: str | None = None prompt_tokens: int | None = None completion_tokens: int | None = None + total_tokens: int | None = None latency_seconds: float | None = None - @property - def total_tokens(self) -> int | None: - """Total tokens used, or ``None`` when the provider reports no usage.""" - if self.prompt_tokens is None and self.completion_tokens is None: - return None - return (self.prompt_tokens or 0) + (self.completion_tokens or 0) - @property def usage(self) -> TokenUsage: """Token usage metadata for this completion (issue #2).""" - return TokenUsage.from_counts( - self.prompt_tokens, self.completion_tokens, self.total_tokens - ) + return TokenUsage.from_counts(self.prompt_tokens, self.completion_tokens, self.total_tokens) def to_dict(self) -> dict[str, Any]: """Serialize the response for logging/telemetry.""" diff --git a/app/services/providers/mock.py b/app/services/providers/mock.py index 564a6a5..5b532ba 100644 --- a/app/services/providers/mock.py +++ b/app/services/providers/mock.py @@ -19,7 +19,7 @@ message_role, prepare_messages, ) -from app.services.token_usage import usage_from_response, usage_from_text +from app.services.token_usage import messages_text, usage_from_text class MockProvider(LLMProvider): @@ -56,16 +56,15 @@ def chat( time.sleep(self.delay) prepared = prepare_messages(messages, supports_vision=self.supports_vision) content = self._respond(prepared) - response = ProviderResponse( + usage = usage_from_text(messages_text(prepared), content) + return ProviderResponse( content=content, model=model or self.models[0], + prompt_tokens=usage["prompt_tokens"], + completion_tokens=usage["completion_tokens"], + total_tokens=usage["total_tokens"], latency_seconds=time.perf_counter() - started, ) - usage = usage_from_response(response, prepared) - response.prompt_tokens = usage["prompt_tokens"] - response.completion_tokens = usage["completion_tokens"] - response.total_tokens = usage["total_tokens"] - return response def stream( self, From c8bf9fad61863136eb2272e7efb3b2d60138a13b Mon Sep 17 00:00:00 2001 From: "Deborah (Precious) Aregbesola" <330907859+thegeneraldeborah-dev@users.noreply.github.com> Date: Tue, 6 Oct 2026 13:31:12 +0000 Subject: [PATCH 3/4] fix(ci): keep ProviderResponse.total_tokens derived from the token counts `ProviderResponse.total_tokens` had become a plain dataclass field that defaulted to `None`, so every caller that only supplied `prompt_tokens` / `completion_tokens` reported a `None` total. Three suites caught it: * `tests/test_providers.py::TestProviderResponse::test_total_tokens_and_dict` * `tests/test_providers.py::TestOpenAIProvider::test_chat_returns_uniform_response` * `tests/test_chat_audit.py::TestProviderCallLogging::test_success_logs_provider_status_tokens` `total_tokens` is derived again (and stays `None` only when the provider reports no usage at all), so nothing regresses for the token-usage tracking this branch adds: the new `usage` property still serves the same figures through `TokenUsage.from_counts`, and the providers no longer pass `total_tokens` explicitly since it is computed. --- app/services/providers/anthropic.py | 1 - app/services/providers/base.py | 10 ++++++++-- app/services/providers/mock.py | 1 - app/services/providers/openai.py | 1 - 4 files changed, 8 insertions(+), 5 deletions(-) diff --git a/app/services/providers/anthropic.py b/app/services/providers/anthropic.py index 972d0e1..3eb925f 100644 --- a/app/services/providers/anthropic.py +++ b/app/services/providers/anthropic.py @@ -179,7 +179,6 @@ def chat( model=data.get("model") or payload["model"], prompt_tokens=usage.get("input_tokens"), completion_tokens=usage.get("output_tokens"), - total_tokens=(usage.get("input_tokens") or 0) + (usage.get("output_tokens") or 0), latency_seconds=time.perf_counter() - started, ) diff --git a/app/services/providers/base.py b/app/services/providers/base.py index 695db4f..13a9114 100644 --- a/app/services/providers/base.py +++ b/app/services/providers/base.py @@ -130,13 +130,19 @@ class ProviderResponse: model: str | None = None prompt_tokens: int | None = None completion_tokens: int | None = None - total_tokens: int | None = None latency_seconds: float | None = None + @property + def total_tokens(self) -> int | None: + """Total tokens used, or ``None`` when the provider reports no usage.""" + if self.prompt_tokens is None and self.completion_tokens is None: + return None + return (self.prompt_tokens or 0) + (self.completion_tokens or 0) + @property def usage(self) -> TokenUsage: """Token usage metadata for this completion (issue #2).""" - return TokenUsage.from_counts(self.prompt_tokens, self.completion_tokens, self.total_tokens) + return TokenUsage.from_counts(self.prompt_tokens, self.completion_tokens) def to_dict(self) -> dict[str, Any]: """Serialize the response for logging/telemetry.""" diff --git a/app/services/providers/mock.py b/app/services/providers/mock.py index 5b532ba..19ae944 100644 --- a/app/services/providers/mock.py +++ b/app/services/providers/mock.py @@ -62,7 +62,6 @@ def chat( model=model or self.models[0], prompt_tokens=usage["prompt_tokens"], completion_tokens=usage["completion_tokens"], - total_tokens=usage["total_tokens"], latency_seconds=time.perf_counter() - started, ) diff --git a/app/services/providers/openai.py b/app/services/providers/openai.py index 81cebb9..b57e338 100644 --- a/app/services/providers/openai.py +++ b/app/services/providers/openai.py @@ -162,7 +162,6 @@ def chat( model=data.get("model") or payload["model"], prompt_tokens=usage.get("prompt_tokens"), completion_tokens=usage.get("completion_tokens"), - total_tokens=usage.get("total_tokens"), latency_seconds=time.perf_counter() - started, ) From 3b315bb2cb37b1638e6b76317b3acb6af6bd7234 Mon Sep 17 00:00:00 2001 From: "Deborah (Precious) Aregbesola" <330907859+thegeneraldeborah-dev@users.noreply.github.com> Date: Tue, 6 Oct 2026 13:31:36 +0000 Subject: [PATCH 4/4] fix(ci): keep the derived total_tokens on ProviderResponse MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `total_tokens` became a plain dataclass field, so a response built from its parts — providers whose vendor payload omits the total, and every caller that constructs one by keyword — reported `None`: * tests/test_providers.py::TestProviderResponse::test_total_tokens_and_dict * tests/test_providers.py::TestOpenAIProvider::test_chat_returns_uniform_response * tests/test_chat_audit.py::TestProviderCallLogging::test_success_logs_provider_status_tokens `__post_init__` derives it from `prompt_tokens + completion_tokens` when the provider reports only the parts (the same rule `TokenUsage.from_counts` already applies to `response.usage`); an explicit total, or a response with no counts at all, is left untouched. --- app/services/providers/anthropic.py | 1 + app/services/providers/base.py | 25 ++++++++++++++++++------- app/services/providers/mock.py | 1 + app/services/providers/openai.py | 1 + 4 files changed, 21 insertions(+), 7 deletions(-) diff --git a/app/services/providers/anthropic.py b/app/services/providers/anthropic.py index 3eb925f..972d0e1 100644 --- a/app/services/providers/anthropic.py +++ b/app/services/providers/anthropic.py @@ -179,6 +179,7 @@ def chat( model=data.get("model") or payload["model"], prompt_tokens=usage.get("input_tokens"), completion_tokens=usage.get("output_tokens"), + total_tokens=(usage.get("input_tokens") or 0) + (usage.get("output_tokens") or 0), latency_seconds=time.perf_counter() - started, ) diff --git a/app/services/providers/base.py b/app/services/providers/base.py index 13a9114..2f97a54 100644 --- a/app/services/providers/base.py +++ b/app/services/providers/base.py @@ -130,19 +130,30 @@ class ProviderResponse: model: str | None = None prompt_tokens: int | None = None completion_tokens: int | None = None + total_tokens: int | None = None latency_seconds: float | None = None - @property - def total_tokens(self) -> int | None: - """Total tokens used, or ``None`` when the provider reports no usage.""" - if self.prompt_tokens is None and self.completion_tokens is None: - return None - return (self.prompt_tokens or 0) + (self.completion_tokens or 0) + def __post_init__(self) -> None: + """Derive ``total_tokens`` when the provider reports only the parts. + + Vendors such as OpenAI/Anthropic may omit the total, and most call sites + (audit logging, the message columns, ``to_dict``) read ``total_tokens`` + directly, so it must never stay ``None`` while the parts are known. + ``TokenUsage.from_counts`` applies the same rule to the ``usage`` view. + """ + if self.total_tokens is None and ( + self.prompt_tokens is not None or self.completion_tokens is not None + ): + object.__setattr__( + self, + "total_tokens", + (self.prompt_tokens or 0) + (self.completion_tokens or 0), + ) @property def usage(self) -> TokenUsage: """Token usage metadata for this completion (issue #2).""" - return TokenUsage.from_counts(self.prompt_tokens, self.completion_tokens) + return TokenUsage.from_counts(self.prompt_tokens, self.completion_tokens, self.total_tokens) def to_dict(self) -> dict[str, Any]: """Serialize the response for logging/telemetry.""" diff --git a/app/services/providers/mock.py b/app/services/providers/mock.py index 19ae944..5b532ba 100644 --- a/app/services/providers/mock.py +++ b/app/services/providers/mock.py @@ -62,6 +62,7 @@ def chat( model=model or self.models[0], prompt_tokens=usage["prompt_tokens"], completion_tokens=usage["completion_tokens"], + total_tokens=usage["total_tokens"], latency_seconds=time.perf_counter() - started, ) diff --git a/app/services/providers/openai.py b/app/services/providers/openai.py index b57e338..81cebb9 100644 --- a/app/services/providers/openai.py +++ b/app/services/providers/openai.py @@ -162,6 +162,7 @@ def chat( model=data.get("model") or payload["model"], prompt_tokens=usage.get("prompt_tokens"), completion_tokens=usage.get("completion_tokens"), + total_tokens=usage.get("total_tokens"), latency_seconds=time.perf_counter() - started, )