Coverage for agentos/cost/token_counter.py: 34%

130 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 12:20 +0800

1""" 

2Token Counter — Model-aware token counting and cost estimation. 

3 

4Supports tiktoken-based counting for OpenAI models and approximate 

5counting for other providers (Anthropic, Google, local models). 

6""" 

7 

8from __future__ import annotations 

9 

10from dataclasses import dataclass 

11from enum import Enum 

12 

13 

14class ModelFamily(Enum): 

15 """模型系列枚举。""" 

16 

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" 

26 

27 

28@dataclass 

29class TokenCount: 

30 """Token counts for a message or conversation.""" 

31 

32 prompt_tokens: int = 0 

33 completion_tokens: int = 0 

34 total_tokens: int = 0 

35 model: str = "" 

36 

37 

38@dataclass 

39class CostEstimate: 

40 """Estimated cost for token usage.""" 

41 

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 

47 

48 

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} 

71 

72 

73class TokenCounter: 

74 """ 

75 Model-aware token counting and cost estimation. 

76 

77 Uses tiktoken when available for OpenAI models, falls back to 

78 character-based approximation for other models. 

79 

80 Example:: 

81 

82 counter = TokenCounter() 

83 tokens = counter.count("Hello, world!", model="gpt-4o") 

84 cost = counter.estimate_cost(tokens, model="gpt-4o") 

85 """ 

86 

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 } 

99 

100 def __init__(self): 

101 self._tiktoken_available = self._try_load_tiktoken() 

102 self._encoders: dict[str, object] = {} 

103 self._usage_log: list[TokenCount] = [] 

104 

105 def _try_load_tiktoken(self) -> bool: 

106 try: 

107 import tiktoken 

108 

109 self._tiktoken = tiktoken 

110 return True 

111 except ImportError: 

112 return False 

113 

114 def _get_encoder(self, model: str): 

115 """Get tiktoken encoder for model, with caching.""" 

116 if not self._tiktoken_available: 

117 return None 

118 

119 if model in self._encoders: 

120 return self._encoders[model] 

121 

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 

131 

132 def count(self, text: str, model: str = "gpt-4o") -> TokenCount: 

133 """ 

134 Count tokens in text for a specific model. 

135 

136 Args: 

137 text: The text to count tokens for. 

138 model: Model identifier string. 

139 

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) 

145 

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)) 

151 

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 

159 

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. 

167 

168 Args: 

169 messages: List of {"role": "...", "content": "..."} dicts. 

170 model: Model identifier. 

171 

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 

181 

182 result = TokenCount( 

183 prompt_tokens=total, 

184 total_tokens=total, 

185 model=model, 

186 ) 

187 self._usage_log.append(result) 

188 return result 

189 

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. 

197 

198 Args: 

199 token_count: Token counts from count() or count_messages(). 

200 model: Override model for pricing lookup. 

201 

202 Returns: 

203 CostEstimate with total USD cost. 

204 """ 

205 m = model or token_count.model 

206 pricing = self._get_pricing(m) 

207 

208 prompt_cost = (token_count.prompt_tokens / 1_000_000) * pricing[0] 

209 completion_cost = (token_count.completion_tokens / 1_000_000) * pricing[1] 

210 

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 ) 

217 

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] 

222 

223 # Try prefix match 

224 for key, pricing in PRICING_TABLE.items(): 

225 if model.startswith(key): 

226 return pricing 

227 

228 # Default: conservative estimate 

229 return (1.00, 3.00) 

230 

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 

250 

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 ) 

260 

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) 

269 

270 def reset_usage(self) -> None: 

271 self._usage_log.clear() 

272 

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}" 

281 

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"