Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 50 additions & 3 deletions app/chat/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
GET /api/conversations/<id> conversation with its messages
POST /api/conversations/<id>/messages send a message (LLM reply)
DELETE /api/conversations/<id> 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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)


Expand Down Expand Up @@ -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,
}
)
1 change: 1 addition & 0 deletions app/services/providers/anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)

Expand Down
68 changes: 63 additions & 5 deletions app/services/providers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand All @@ -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."""
Expand Down
13 changes: 10 additions & 3 deletions app/services/providers/mock.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
"""

Expand All @@ -19,6 +19,7 @@
message_role,
prepare_messages,
)
from app.services.token_usage import messages_text, usage_from_text


class MockProvider(LLMProvider):
Expand Down Expand Up @@ -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,
)

Expand All @@ -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)
Expand Down
12 changes: 12 additions & 0 deletions app/services/providers/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
)

Expand All @@ -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: "):
Expand All @@ -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 = ""
Expand All @@ -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"),
}
2 changes: 1 addition & 1 deletion app/services/token_usage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
"""

Expand Down