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

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 

12from typing import Optional 

13 

14 

15class ModelFamily(Enum): 

16 

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

18 

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" 

28 

29 

30@dataclass 

31class TokenCount: 

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

33 

34 prompt_tokens: int = 0 

35 completion_tokens: int = 0 

36 total_tokens: int = 0 

37 model: str = "" 

38 

39 

40@dataclass 

41class CostEstimate: 

42 """Estimated cost for token usage.""" 

43 

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 

49 

50 

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} 

73 

74 

75class TokenCounter: 

76 """ 

77 Model-aware token counting and cost estimation. 

78 

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

80 character-based approximation for other models. 

81 

82 Example:: 

83 

84 counter = TokenCounter() 

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

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

87 """ 

88 

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 } 

101 

102 def __init__(self): 

103 self._tiktoken_available = self._try_load_tiktoken() 

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

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

106 

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 

114 

115 def _get_encoder(self, model: str): 

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

117 if not self._tiktoken_available: 

118 return None 

119 

120 if model in self._encoders: 

121 return self._encoders[model] 

122 

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 

132 

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

134 """ 

135 Count tokens in text for a specific model. 

136 

137 Args: 

138 text: The text to count tokens for. 

139 model: Model identifier string. 

140 

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) 

146 

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

152 

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 

160 

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. 

166 

167 Args: 

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

169 model: Model identifier. 

170 

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 

180 

181 result = TokenCount( 

182 prompt_tokens=total, 

183 total_tokens=total, 

184 model=model, 

185 ) 

186 self._usage_log.append(result) 

187 return result 

188 

189 def estimate_cost( 

190 self, token_count: TokenCount, model: Optional[str] = None, 

191 ) -> CostEstimate: 

192 """ 

193 Estimate USD cost from token usage. 

194 

195 Args: 

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

197 model: Override model for pricing lookup. 

198 

199 Returns: 

200 CostEstimate with total USD cost. 

201 """ 

202 m = model or token_count.model 

203 pricing = self._get_pricing(m) 

204 

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

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

207 

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 ) 

214 

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] 

219 

220 # Try prefix match 

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

222 if model.startswith(key): 

223 return pricing 

224 

225 # Default: conservative estimate 

226 return (1.00, 3.00) 

227 

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 

247 

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 ) 

257 

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) 

266 

267 def reset_usage(self) -> None: 

268 self._usage_log.clear() 

269 

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

278 

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"