mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
core/agent: char-based token estimate fallback for usage
When API usage is missing, estimate prompt/completion tokens at ~4 chars per token and mark source as estimate (or mixed when combined).
This commit is contained in:
@@ -1,34 +1,74 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import cast
|
from typing import TYPE_CHECKING, cast
|
||||||
|
|
||||||
from msgspec import UNSET, Struct
|
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):
|
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
|
prompt_tokens: int = 0
|
||||||
completion_tokens: int = 0
|
completion_tokens: int = 0
|
||||||
total_tokens: int = 0
|
total_tokens: int = 0
|
||||||
|
# "api" | "estimate" | "mixed" (session totals combining both)
|
||||||
|
source: str = "api"
|
||||||
|
|
||||||
def add(self, other: TokenUsage) -> TokenUsage:
|
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(
|
return TokenUsage(
|
||||||
prompt_tokens=self.prompt_tokens + other.prompt_tokens,
|
prompt_tokens=self.prompt_tokens + other.prompt_tokens,
|
||||||
completion_tokens=self.completion_tokens + other.completion_tokens,
|
completion_tokens=self.completion_tokens + other.completion_tokens,
|
||||||
total_tokens=self.total_tokens + other.total_tokens,
|
total_tokens=self.total_tokens + other.total_tokens,
|
||||||
|
source=source,
|
||||||
)
|
)
|
||||||
|
|
||||||
def is_zero(self) -> bool:
|
def is_zero(self) -> bool:
|
||||||
return self.prompt_tokens == 0 and self.completion_tokens == 0 and self.total_tokens == 0
|
return self.prompt_tokens == 0 and self.completion_tokens == 0 and self.total_tokens == 0
|
||||||
|
|
||||||
def format_line(self) -> str:
|
def format_line(self) -> str:
|
||||||
|
tag = ""
|
||||||
|
if self.source == "estimate":
|
||||||
|
tag = " (est)"
|
||||||
|
elif self.source == "mixed":
|
||||||
|
tag = " (api+est)"
|
||||||
return (
|
return (
|
||||||
f"tokens prompt={self.prompt_tokens} completion={self.completion_tokens} "
|
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:
|
def _as_nonneg_int(value: object) -> int:
|
||||||
if isinstance(value, bool):
|
if isinstance(value, bool):
|
||||||
return 0
|
return 0
|
||||||
@@ -53,4 +93,48 @@ def token_usage_from_api(usage: object) -> TokenUsage | None:
|
|||||||
total = prompt + completion
|
total = prompt + completion
|
||||||
if prompt == 0 and completion == 0 and total == 0:
|
if prompt == 0 and completion == 0 and total == 0:
|
||||||
return None
|
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,
|
||||||
|
)
|
||||||
|
|||||||
@@ -2,17 +2,34 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from msgspec import UNSET
|
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:
|
def test_token_usage_add() -> None:
|
||||||
a = TokenUsage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
|
a = TokenUsage(prompt_tokens=10, completion_tokens=5, total_tokens=15, source="api")
|
||||||
b = TokenUsage(prompt_tokens=3, completion_tokens=2, total_tokens=5)
|
b = TokenUsage(prompt_tokens=3, completion_tokens=2, total_tokens=5, source="api")
|
||||||
c = a.add(b)
|
c = a.add(b)
|
||||||
assert c.prompt_tokens == 13
|
assert c.prompt_tokens == 13
|
||||||
assert c.completion_tokens == 7
|
assert c.completion_tokens == 7
|
||||||
assert c.total_tokens == 20
|
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:
|
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.prompt_tokens == 11
|
||||||
assert u.completion_tokens == 4
|
assert u.completion_tokens == 4
|
||||||
assert u.total_tokens == 15
|
assert u.total_tokens == 15
|
||||||
|
assert u.source == "api"
|
||||||
|
|
||||||
|
|
||||||
def test_token_usage_from_api_infers_total() -> None:
|
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
|
assert u.total_tokens == 5
|
||||||
|
|
||||||
|
|
||||||
def test_format_line() -> None:
|
def test_chars_to_tokens() -> None:
|
||||||
line = TokenUsage(prompt_tokens=1, completion_tokens=2, total_tokens=3).format_line()
|
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 "prompt=1" in line
|
||||||
assert "completion=2" in line
|
assert "(est)" in line
|
||||||
assert "total=3" in line
|
mixed = TokenUsage(prompt_tokens=1, completion_tokens=0, total_tokens=1, source="mixed").format_line()
|
||||||
|
assert "(api+est)" in mixed
|
||||||
|
|||||||
Reference in New Issue
Block a user