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