"""Cost tracking and calculation for evaluation runs.""" from typing import Any from pydantic import BaseModel, Field class ModelPricing(BaseModel): """Pricing for a model (per 1M tokens).""" model_id: str prompt_cost_per_1m: float = Field(description="Cost per 1M prompt tokens in USD") completion_cost_per_1m: float = Field(description="Cost per 1M completion tokens in USD") class TokenUsage(BaseModel): """Token usage for a single API call.""" prompt_tokens: int = 0 completion_tokens: int = 0 total_tokens: int = 0 class CostBreakdown(BaseModel): """Cost breakdown for a turn, case, or run.""" prompt_tokens: int = 0 completion_tokens: int = 0 total_tokens: int = 0 cost_usd: float = 0.0 # Default pricing for common models (per 1M tokens in USD) DEFAULT_PRICING: dict[str, ModelPricing] = { "gpt-4o": ModelPricing(model_id="gpt-4o", prompt_cost_per_1m=5.0, completion_cost_per_1m=15.0), "gpt-4o-mini": ModelPricing(model_id="gpt-4o-mini", prompt_cost_per_1m=0.15, completion_cost_per_1m=0.60), "gpt-4-turbo": ModelPricing(model_id="gpt-4-turbo", prompt_cost_per_1m=10.0, completion_cost_per_1m=30.0), "gpt-3.5-turbo": ModelPricing(model_id="gpt-3.5-turbo", prompt_cost_per_1m=0.50, completion_cost_per_1m=1.50), "claude-3-5-sonnet": ModelPricing(model_id="claude-3-5-sonnet", prompt_cost_per_1m=3.0, completion_cost_per_1m=15.0), "claude-3-haiku": ModelPricing(model_id="claude-3-haiku", prompt_cost_per_1m=0.25, completion_cost_per_1m=1.25), } def calculate_cost( prompt_tokens: int, completion_tokens: int, pricing: ModelPricing, ) -> float: """Calculate cost in USD for given token usage and pricing. Args: prompt_tokens: Number of prompt tokens completion_tokens: Number of completion tokens pricing: Model pricing configuration Returns: Cost in USD """ prompt_cost = (prompt_tokens / 1_000_000) * pricing.prompt_cost_per_1m completion_cost = (completion_tokens / 1_000_000) * pricing.completion_cost_per_1m return prompt_cost + completion_cost def aggregate_token_usage(turns: list[dict[str, Any]]) -> TokenUsage: """Aggregate token usage from a list of turns. Args: turns: List of turn dicts with optional token fields Returns: Aggregated token usage """ prompt_tokens = sum(t.get("prompt_tokens") or 0 for t in turns) completion_tokens = sum(t.get("completion_tokens") or 0 for t in turns) return TokenUsage( prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=prompt_tokens + completion_tokens, ) def calculate_turn_cost( turn: dict[str, Any], pricing: ModelPricing, ) -> CostBreakdown: """Calculate cost for a single turn. Args: turn: Turn dict with optional token fields pricing: Model pricing configuration Returns: Cost breakdown for the turn """ prompt_tokens = turn.get("prompt_tokens") or 0 completion_tokens = turn.get("completion_tokens") or 0 total_tokens = prompt_tokens + completion_tokens cost_usd = calculate_cost(prompt_tokens, completion_tokens, pricing) return CostBreakdown( prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=total_tokens, cost_usd=cost_usd, ) def calculate_case_cost( turns: list[dict[str, Any]], pricing: ModelPricing, ) -> CostBreakdown: """Calculate cost for a case (multiple turns). Args: turns: List of turn dicts pricing: Model pricing configuration Returns: Aggregated cost breakdown for the case """ usage = aggregate_token_usage(turns) cost_usd = calculate_cost(usage.prompt_tokens, usage.completion_tokens, pricing) return CostBreakdown( prompt_tokens=usage.prompt_tokens, completion_tokens=usage.completion_tokens, total_tokens=usage.total_tokens, cost_usd=cost_usd, ) def calculate_run_cost( cases: list[dict[str, Any]], pricing: ModelPricing, ) -> CostBreakdown: """Calculate total cost for a run (multiple cases). Args: cases: List of case dicts, each with a 'turns' field pricing: Model pricing configuration Returns: Aggregated cost breakdown for the run """ all_turns = [] for case in cases: all_turns.extend(case.get("turns", [])) usage = aggregate_token_usage(all_turns) cost_usd = calculate_cost(usage.prompt_tokens, usage.completion_tokens, pricing) return CostBreakdown( prompt_tokens=usage.prompt_tokens, completion_tokens=usage.completion_tokens, total_tokens=usage.total_tokens, cost_usd=cost_usd, )