diff --git a/app/chat/api.py b/app/chat/api.py index 4b127ea..20c5ca8 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 @@ -85,6 +86,22 @@ 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: + if message.total_tokens is None: + continue + 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, + "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 +157,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 +224,42 @@ 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 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" 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/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..2f97a54 100644 --- a/app/services/providers/base.py +++ b/app/services/providers/base.py @@ -80,6 +80,48 @@ 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.""" @@ -88,14 +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 + 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 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 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.""" diff --git a/app/services/providers/mock.py b/app/services/providers/mock.py index de307b6..5b532ba 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 messages_text, usage_from_text class MockProvider(LLMProvider): @@ -54,9 +55,14 @@ def chat( if self.delay: time.sleep(self.delay) prepared = prepare_messages(messages, supports_vision=self.supports_vision) + content = self._respond(prepared) + usage = usage_from_text(messages_text(prepared), content) return ProviderResponse( - content=self._respond(prepared), + 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, ) @@ -67,7 +73,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. """