Coverage for agentos/cost/token_counter.py: 34%
131 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
1"""
2Token Counter — Model-aware token counting and cost estimation.
4Supports tiktoken-based counting for OpenAI models and approximate
5counting for other providers (Anthropic, Google, local models).
6"""
8from __future__ import annotations
10from dataclasses import dataclass
11from enum import Enum
12from typing import Optional
15class ModelFamily(Enum):
17 """模型系列枚举。"""
19 GPT4 = "gpt-4"
20 GPT4O = "gpt-4o"
21 GPT35 = "gpt-3.5-turbo"
22 CLAUDE3 = "claude-3"
23 CLAUDE35 = "claude-3.5"
24 GEMINI = "gemini"
25 LLAMA = "llama"
26 MIXTRAL = "mixtral"
27 UNKNOWN = "unknown"
30@dataclass
31class TokenCount:
32 """Token counts for a message or conversation."""
34 prompt_tokens: int = 0
35 completion_tokens: int = 0
36 total_tokens: int = 0
37 model: str = ""
40@dataclass
41class CostEstimate:
42 """Estimated cost for token usage."""
44 prompt_cost: float = 0.0
45 completion_cost: float = 0.0
46 total_cost: float = 0.0
47 currency: str = "USD"
48 token_count: Optional[TokenCount] = None
51# Pricing per 1M tokens (input, output) — updated mid-2025
52PRICING_TABLE: dict[str, tuple[float, float]] = {
53 # OpenAI
54 "gpt-4o": (2.50, 10.00),
55 "gpt-4o-mini": (0.15, 0.60),
56 "gpt-4-turbo": (10.00, 30.00),
57 "gpt-4": (30.00, 60.00),
58 "gpt-3.5-turbo": (0.50, 1.50),
59 # Anthropic
60 "claude-3.5-sonnet": (3.00, 15.00),
61 "claude-3-opus": (15.00, 75.00),
62 "claude-3-haiku": (0.25, 1.25),
63 "claude-3-sonnet": (3.00, 15.00),
64 # Google
65 "gemini-1.5-pro": (1.25, 5.00),
66 "gemini-1.5-flash": (0.075, 0.30),
67 "gemini-2.0-flash": (0.10, 0.40),
68 # Open-source (hosted)
69 "llama-3-70b": (0.59, 0.79),
70 "llama-3-8b": (0.06, 0.06),
71 "mixtral-8x7b": (0.24, 0.24),
72}
75class TokenCounter:
76 """
77 Model-aware token counting and cost estimation.
79 Uses tiktoken when available for OpenAI models, falls back to
80 character-based approximation for other models.
82 Example::
84 counter = TokenCounter()
85 tokens = counter.count("Hello, world!", model="gpt-4o")
86 cost = counter.estimate_cost(tokens, model="gpt-4o")
87 """
89 # Characters per token — rough estimates per model family
90 CHARS_PER_TOKEN: dict[ModelFamily, float] = {
91 ModelFamily.GPT4: 3.5,
92 ModelFamily.GPT4O: 3.8,
93 ModelFamily.GPT35: 4.0,
94 ModelFamily.CLAUDE3: 3.2,
95 ModelFamily.CLAUDE35: 3.4,
96 ModelFamily.GEMINI: 3.0,
97 ModelFamily.LLAMA: 3.8,
98 ModelFamily.MIXTRAL: 3.6,
99 ModelFamily.UNKNOWN: 4.0,
100 }
102 def __init__(self):
103 self._tiktoken_available = self._try_load_tiktoken()
104 self._encoders: dict[str, object] = {}
105 self._usage_log: list[TokenCount] = []
107 def _try_load_tiktoken(self) -> bool:
108 try:
109 import tiktoken
110 self._tiktoken = tiktoken
111 return True
112 except ImportError:
113 return False
115 def _get_encoder(self, model: str):
116 """Get tiktoken encoder for model, with caching."""
117 if not self._tiktoken_available:
118 return None
120 if model in self._encoders:
121 return self._encoders[model]
123 try:
124 encoder = self._tiktoken.encoding_for_model(model)
125 except KeyError:
126 try:
127 encoder = self._tiktoken.get_encoding("cl100k_base")
128 except Exception:
129 return None
130 self._encoders[model] = encoder
131 return encoder
133 def count(self, text: str, model: str = "gpt-4o") -> TokenCount:
134 """
135 Count tokens in text for a specific model.
137 Args:
138 text: The text to count tokens for.
139 model: Model identifier string.
141 Returns:
142 TokenCount with prompt_tokens set (single text counts as prompt).
143 """
144 family = self._classify_model(model)
145 encoder = self._get_encoder(model)
147 if encoder:
148 count_val = len(encoder.encode(text))
149 else:
150 chars_per = self.CHARS_PER_TOKEN.get(family, 4.0)
151 count_val = max(1, int(len(text) / chars_per))
153 result = TokenCount(
154 prompt_tokens=count_val,
155 total_tokens=count_val,
156 model=model,
157 )
158 self._usage_log.append(result)
159 return result
161 def count_messages(
162 self, messages: list[dict[str, str]], model: str = "gpt-4o",
163 ) -> TokenCount:
164 """
165 Count tokens for a list of chat messages.
167 Args:
168 messages: List of {"role": "...", "content": "..."} dicts.
169 model: Model identifier.
171 Returns:
172 TokenCount with total prompt tokens.
173 """
174 total = 0
175 for msg in messages:
176 content = msg.get("content", "")
177 # Role overhead: ~4 tokens per message
178 total += 4
179 total += self.count(content, model=model).prompt_tokens
181 result = TokenCount(
182 prompt_tokens=total,
183 total_tokens=total,
184 model=model,
185 )
186 self._usage_log.append(result)
187 return result
189 def estimate_cost(
190 self, token_count: TokenCount, model: Optional[str] = None,
191 ) -> CostEstimate:
192 """
193 Estimate USD cost from token usage.
195 Args:
196 token_count: Token counts from count() or count_messages().
197 model: Override model for pricing lookup.
199 Returns:
200 CostEstimate with total USD cost.
201 """
202 m = model or token_count.model
203 pricing = self._get_pricing(m)
205 prompt_cost = (token_count.prompt_tokens / 1_000_000) * pricing[0]
206 completion_cost = (token_count.completion_tokens / 1_000_000) * pricing[1]
208 return CostEstimate(
209 prompt_cost=prompt_cost,
210 completion_cost=completion_cost,
211 total_cost=prompt_cost + completion_cost,
212 token_count=token_count,
213 )
215 def _get_pricing(self, model: str) -> tuple[float, float]:
216 """Find closest pricing match for model."""
217 if model in PRICING_TABLE:
218 return PRICING_TABLE[model]
220 # Try prefix match
221 for key, pricing in PRICING_TABLE.items():
222 if model.startswith(key):
223 return pricing
225 # Default: conservative estimate
226 return (1.00, 3.00)
228 def _classify_model(self, model: str) -> ModelFamily:
229 model_lower = model.lower()
230 if "gpt-4o" in model_lower:
231 return ModelFamily.GPT4O
232 if "gpt-4" in model_lower:
233 return ModelFamily.GPT4
234 if "gpt-3.5" in model_lower:
235 return ModelFamily.GPT35
236 if "claude-3.5" in model_lower:
237 return ModelFamily.CLAUDE35
238 if "claude-3" in model_lower or "claude" in model_lower:
239 return ModelFamily.CLAUDE3
240 if "gemini" in model_lower:
241 return ModelFamily.GEMINI
242 if "llama" in model_lower:
243 return ModelFamily.LLAMA
244 if "mixtral" in model_lower:
245 return ModelFamily.MIXTRAL
246 return ModelFamily.UNKNOWN
248 def get_total_usage(self) -> TokenCount:
249 """Aggregate all logged usage."""
250 prompt = sum(u.prompt_tokens for u in self._usage_log)
251 completion = sum(u.completion_tokens for u in self._usage_log)
252 return TokenCount(
253 prompt_tokens=prompt,
254 completion_tokens=completion,
255 total_tokens=prompt + completion,
256 )
258 def get_total_cost(self) -> CostEstimate:
259 """Estimate total cost of all logged usage."""
260 total_tokens = self.get_total_usage()
261 total_cost = 0.0
262 for entry in self._usage_log:
263 cost = self.estimate_cost(entry)
264 total_cost += cost.total_cost
265 return CostEstimate(total_cost=total_cost, token_count=total_tokens)
267 def reset_usage(self) -> None:
268 self._usage_log.clear()
270 @staticmethod
271 def format_cost(cost: CostEstimate) -> str:
272 """Human-readable cost string."""
273 if cost.total_cost < 0.01:
274 return f"${cost.total_cost:.6f}"
275 if cost.total_cost < 1.0:
276 return f"${cost.total_cost:.4f}"
277 return f"${cost.total_cost:.2f}"
279 @staticmethod
280 def format_tokens(tokens: TokenCount) -> str:
281 """Human-readable token count string."""
282 if tokens.total_tokens < 1000:
283 return str(tokens.total_tokens)
284 return f"{tokens.total_tokens / 1000:.1f}K"