diff --git a/src/plyngent/agent/usage.py b/src/plyngent/agent/usage.py index 0b8048f..ab9913b 100644 --- a/src/plyngent/agent/usage.py +++ b/src/plyngent/agent/usage.py @@ -1,34 +1,74 @@ from __future__ import annotations -from typing import cast +from typing import TYPE_CHECKING, cast from msgspec import UNSET, Struct +from plyngent.agent.budget import estimate_message_chars, estimate_messages_chars + +if TYPE_CHECKING: + from collections.abc import Sequence + + from plyngent.lmproto.openai_compatible.model import AnyChatMessage, AssistantChatMessage + +# Rough OpenAI-style heuristic: ~4 characters per token (not model-accurate). +DEFAULT_CHARS_PER_TOKEN = 4.0 + class TokenUsage(Struct, omit_defaults=True): - """Token counts from a provider ``usage`` object (OpenAI-compatible).""" + """Token counts from API ``usage`` or a char-based estimate.""" prompt_tokens: int = 0 completion_tokens: int = 0 total_tokens: int = 0 + # "api" | "estimate" | "mixed" (session totals combining both) + source: str = "api" def add(self, other: TokenUsage) -> TokenUsage: + if self.is_zero(): + return other + if other.is_zero(): + return self + source = self.source if self.source == other.source else "mixed" return TokenUsage( prompt_tokens=self.prompt_tokens + other.prompt_tokens, completion_tokens=self.completion_tokens + other.completion_tokens, total_tokens=self.total_tokens + other.total_tokens, + source=source, ) def is_zero(self) -> bool: return self.prompt_tokens == 0 and self.completion_tokens == 0 and self.total_tokens == 0 def format_line(self) -> str: + tag = "" + if self.source == "estimate": + tag = " (est)" + elif self.source == "mixed": + tag = " (api+est)" return ( f"tokens prompt={self.prompt_tokens} completion={self.completion_tokens} " - f"total={self.total_tokens}" + f"total={self.total_tokens}{tag}" ) +def chars_to_tokens(chars: int, *, chars_per_token: float = DEFAULT_CHARS_PER_TOKEN) -> int: + """Convert character count to a rough token estimate (ceiling, min 0).""" + if chars <= 0 or chars_per_token <= 0: + return 0 + # Ceiling division without float drift for large values + return max(0, int((chars + chars_per_token - 1e-9) // chars_per_token)) + + +def estimate_tokens_from_chars( + chars: int, + *, + chars_per_token: float = DEFAULT_CHARS_PER_TOKEN, +) -> int: + """Alias for :func:`chars_to_tokens` (public name for fallback counters).""" + return chars_to_tokens(chars, chars_per_token=chars_per_token) + + def _as_nonneg_int(value: object) -> int: if isinstance(value, bool): return 0 @@ -53,4 +93,48 @@ def token_usage_from_api(usage: object) -> TokenUsage | None: total = prompt + completion if prompt == 0 and completion == 0 and total == 0: return None - return TokenUsage(prompt_tokens=prompt, completion_tokens=completion, total_tokens=total) + return TokenUsage( + prompt_tokens=prompt, + completion_tokens=completion, + total_tokens=total, + source="api", + ) + + +def estimate_token_usage( + prompt_messages: Sequence[AnyChatMessage], + assistant: AssistantChatMessage | None = None, + *, + chars_per_token: float = DEFAULT_CHARS_PER_TOKEN, +) -> TokenUsage: + """Char-based fallback when the provider does not report ``usage``.""" + prompt_chars = estimate_messages_chars(prompt_messages) + completion_chars = 0 + if assistant is not None: + completion_chars = estimate_message_chars(assistant) + prompt_tokens = chars_to_tokens(prompt_chars, chars_per_token=chars_per_token) + completion_tokens = chars_to_tokens(completion_chars, chars_per_token=chars_per_token) + return TokenUsage( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + source="estimate", + ) + + +def resolve_round_usage( + api_usage: object, + prompt_messages: Sequence[AnyChatMessage], + assistant: AssistantChatMessage, + *, + chars_per_token: float = DEFAULT_CHARS_PER_TOKEN, +) -> TokenUsage: + """Prefer API usage; otherwise estimate from message characters.""" + parsed = token_usage_from_api(api_usage) + if parsed is not None: + return parsed + return estimate_token_usage( + prompt_messages, + assistant, + chars_per_token=chars_per_token, + ) diff --git a/tests/test_agent/test_usage.py b/tests/test_agent/test_usage.py index 06dd42d..a1bb34f 100644 --- a/tests/test_agent/test_usage.py +++ b/tests/test_agent/test_usage.py @@ -2,17 +2,34 @@ from __future__ import annotations from msgspec import UNSET -from plyngent.agent.usage import TokenUsage, token_usage_from_api +from plyngent.agent.usage import ( + TokenUsage, + chars_to_tokens, + estimate_token_usage, + resolve_round_usage, + token_usage_from_api, +) +from plyngent.lmproto.openai_compatible.model import ( + AssistantChatMessage, + UserChatMessage, +) def test_token_usage_add() -> None: - a = TokenUsage(prompt_tokens=10, completion_tokens=5, total_tokens=15) - b = TokenUsage(prompt_tokens=3, completion_tokens=2, total_tokens=5) + a = TokenUsage(prompt_tokens=10, completion_tokens=5, total_tokens=15, source="api") + b = TokenUsage(prompt_tokens=3, completion_tokens=2, total_tokens=5, source="api") c = a.add(b) assert c.prompt_tokens == 13 assert c.completion_tokens == 7 assert c.total_tokens == 20 - assert a.prompt_tokens == 10 # immutable-ish via new struct + assert c.source == "api" + assert a.prompt_tokens == 10 + + +def test_token_usage_add_mixed_source() -> None: + a = TokenUsage(prompt_tokens=1, completion_tokens=0, total_tokens=1, source="api") + b = TokenUsage(prompt_tokens=2, completion_tokens=0, total_tokens=2, source="estimate") + assert a.add(b).source == "mixed" def test_token_usage_from_api() -> None: @@ -24,6 +41,7 @@ def test_token_usage_from_api() -> None: assert u.prompt_tokens == 11 assert u.completion_tokens == 4 assert u.total_tokens == 15 + assert u.source == "api" def test_token_usage_from_api_infers_total() -> None: @@ -32,8 +50,50 @@ def test_token_usage_from_api_infers_total() -> None: assert u.total_tokens == 5 -def test_format_line() -> None: - line = TokenUsage(prompt_tokens=1, completion_tokens=2, total_tokens=3).format_line() +def test_chars_to_tokens() -> None: + assert chars_to_tokens(0) == 0 + assert chars_to_tokens(1) == 1 + assert chars_to_tokens(4) == 1 + assert chars_to_tokens(5) == 2 + assert chars_to_tokens(8) == 2 + assert chars_to_tokens(9) == 3 + + +def test_estimate_token_usage() -> None: + # 8 chars prompt → 2 tokens at 4 cpt; 4 chars completion → 1 token + usage = estimate_token_usage( + [UserChatMessage(content="12345678")], + AssistantChatMessage(content="abcd"), + ) + assert usage.source == "estimate" + assert usage.prompt_tokens == 2 + assert usage.completion_tokens == 1 + assert usage.total_tokens == 3 + + +def test_resolve_round_usage_prefers_api() -> None: + usage = resolve_round_usage( + {"prompt_tokens": 100, "completion_tokens": 10, "total_tokens": 110}, + [UserChatMessage(content="x")], + AssistantChatMessage(content="y"), + ) + assert usage.source == "api" + assert usage.prompt_tokens == 100 + + +def test_resolve_round_usage_falls_back_to_estimate() -> None: + usage = resolve_round_usage( + None, + [UserChatMessage(content="12345678")], + AssistantChatMessage(content="abcd"), + ) + assert usage.source == "estimate" + assert usage.total_tokens == 3 + + +def test_format_line_marks_estimate() -> None: + line = TokenUsage(prompt_tokens=1, completion_tokens=2, total_tokens=3, source="estimate").format_line() assert "prompt=1" in line - assert "completion=2" in line - assert "total=3" in line + assert "(est)" in line + mixed = TokenUsage(prompt_tokens=1, completion_tokens=0, total_tokens=1, source="mixed").format_line() + assert "(api+est)" in mixed