mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
Merge pull request #558 from devdeepsarkar/refactor-model-resolution
refactor: extract litellm model resolution to shared utility
This commit is contained in:
commit
ec7d0065cc
5 changed files with 127 additions and 158 deletions
|
|
@ -16,6 +16,7 @@ from dataclasses import asdict, dataclass, field
|
|||
from datetime import datetime, timedelta
|
||||
|
||||
from headroom import paths as _paths
|
||||
from headroom.pricing.litellm_pricing import resolve_litellm_model
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
|
@ -62,38 +63,6 @@ try:
|
|||
except ImportError:
|
||||
_LITELLM_AVAILABLE = False
|
||||
|
||||
# Cache resolved model names (e.g. "claude-opus-4-6" → "anthropic/claude-opus-4-6")
|
||||
_resolved_model_cache: dict[str, str] = {}
|
||||
|
||||
|
||||
def _resolve_model(model: str) -> str:
|
||||
"""Resolve to a model name LiteLLM recognises, adding provider prefix if needed.
|
||||
|
||||
TODO: Duplicated with CostTracker._resolve_litellm_model in proxy/server.py.
|
||||
Extract to shared utility.
|
||||
"""
|
||||
if model in _resolved_model_cache:
|
||||
return _resolved_model_cache[model]
|
||||
|
||||
if not _LITELLM_AVAILABLE:
|
||||
_resolved_model_cache[model] = model
|
||||
return model
|
||||
|
||||
# Try as-is
|
||||
if model in _litellm.model_cost:
|
||||
_resolved_model_cache[model] = model
|
||||
return model
|
||||
|
||||
# Try provider prefixes
|
||||
for prefix in ("anthropic/", "openai/", "google/", "mistral/", "deepseek/"):
|
||||
prefixed = f"{prefix}{model}"
|
||||
if prefixed in _litellm.model_cost:
|
||||
_resolved_model_cache[model] = prefixed
|
||||
return prefixed
|
||||
|
||||
_resolved_model_cache[model] = model
|
||||
return model
|
||||
|
||||
|
||||
def _litellm_cost(
|
||||
model: str,
|
||||
|
|
@ -107,7 +76,7 @@ def _litellm_cost(
|
|||
"""
|
||||
if not _LITELLM_AVAILABLE:
|
||||
return None
|
||||
resolved = _resolve_model(model)
|
||||
resolved = resolve_litellm_model(model)
|
||||
try:
|
||||
input_cost, _ = _litellm.cost_per_token(
|
||||
model=resolved,
|
||||
|
|
@ -125,7 +94,7 @@ def _get_list_price(model: str) -> float | None:
|
|||
"""Get list input price per 1M tokens."""
|
||||
if not _LITELLM_AVAILABLE:
|
||||
return None
|
||||
resolved = _resolve_model(model)
|
||||
resolved = resolve_litellm_model(model)
|
||||
info = _litellm.model_cost.get(resolved, {})
|
||||
cost_per_token = info.get("input_cost_per_token")
|
||||
return cost_per_token * 1_000_000 if cost_per_token else None
|
||||
|
|
|
|||
|
|
@ -43,6 +43,50 @@ _MODEL_ALIASES: dict[str, str] = {
|
|||
"claude-3-sonnet-20240229": "claude-3-haiku-20240307",
|
||||
}
|
||||
|
||||
_resolved_model_cache: dict[str, str] = {}
|
||||
|
||||
|
||||
def resolve_litellm_model(model: str) -> str:
|
||||
"""Resolve model name to one LiteLLM recognizes, adding provider prefix if needed.
|
||||
Results are cached per model name to avoid blocking the event loop
|
||||
with repeated synchronous litellm lookups.
|
||||
"""
|
||||
if model in _resolved_model_cache:
|
||||
return _resolved_model_cache[model]
|
||||
resolved = _resolve_litellm_model_uncached(model)
|
||||
_resolved_model_cache[model] = resolved
|
||||
return resolved
|
||||
|
||||
|
||||
def _resolve_litellm_model_uncached(model: str) -> str:
|
||||
"""Uncached resolution — called once per unique model name."""
|
||||
if not LITELLM_AVAILABLE:
|
||||
return model
|
||||
# Try as-is first
|
||||
try:
|
||||
litellm.cost_per_token(model=model, prompt_tokens=1, completion_tokens=0)
|
||||
return model
|
||||
except Exception:
|
||||
pass
|
||||
# Try with provider prefix
|
||||
prefixes = {
|
||||
"claude-": "anthropic/",
|
||||
"gpt-": "openai/",
|
||||
"o1-": "openai/",
|
||||
"o3-": "openai/",
|
||||
"o4-": "openai/",
|
||||
"gemini-": "google/",
|
||||
}
|
||||
for pattern, prefix in prefixes.items():
|
||||
if model.startswith(pattern):
|
||||
prefixed = f"{prefix}{model}"
|
||||
try:
|
||||
litellm.cost_per_token(model=prefixed, prompt_tokens=1, completion_tokens=0)
|
||||
return prefixed
|
||||
except Exception:
|
||||
break
|
||||
return model
|
||||
|
||||
|
||||
@dataclass
|
||||
class LiteLLMModelPricing:
|
||||
|
|
|
|||
|
|
@ -566,59 +566,6 @@ class CostTracker:
|
|||
self._api_cache_write_1h_by_model.clear()
|
||||
self._api_uncached_by_model.clear()
|
||||
|
||||
# Cache resolved model names to avoid repeated litellm lookups.
|
||||
# This is critical: litellm.cost_per_token() is synchronous and can block
|
||||
# the async event loop if it triggers I/O (lazy model info download).
|
||||
_resolved_model_cache: dict[str, str] = {}
|
||||
|
||||
@classmethod
|
||||
def _resolve_litellm_model(cls, model: str) -> str:
|
||||
"""Resolve model name to one LiteLLM recognizes, adding provider prefix if needed.
|
||||
|
||||
Results are cached per model name to avoid blocking the event loop
|
||||
with repeated synchronous litellm lookups.
|
||||
"""
|
||||
if model in cls._resolved_model_cache:
|
||||
return cls._resolved_model_cache[model]
|
||||
|
||||
resolved = cls._resolve_litellm_model_uncached(model)
|
||||
cls._resolved_model_cache[model] = resolved
|
||||
return resolved
|
||||
|
||||
@staticmethod
|
||||
def _resolve_litellm_model_uncached(model: str) -> str:
|
||||
"""Uncached resolution — called once per unique model name."""
|
||||
litellm = _get_litellm_module()
|
||||
if litellm is None:
|
||||
return model
|
||||
|
||||
# Try as-is first
|
||||
try:
|
||||
litellm.cost_per_token(model=model, prompt_tokens=1, completion_tokens=0)
|
||||
return model
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Try with provider prefix
|
||||
prefixes = {
|
||||
"claude-": "anthropic/",
|
||||
"gpt-": "openai/",
|
||||
"o1-": "openai/",
|
||||
"o3-": "openai/",
|
||||
"o4-": "openai/",
|
||||
"gemini-": "google/",
|
||||
}
|
||||
for pattern, prefix in prefixes.items():
|
||||
if model.startswith(pattern):
|
||||
prefixed = f"{prefix}{model}"
|
||||
try:
|
||||
litellm.cost_per_token(model=prefixed, prompt_tokens=1, completion_tokens=0)
|
||||
return prefixed
|
||||
except Exception:
|
||||
break
|
||||
|
||||
return model
|
||||
|
||||
def estimate_cost(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -645,7 +592,9 @@ class CostTracker:
|
|||
return None
|
||||
|
||||
try:
|
||||
resolved_model = self._resolve_litellm_model(model)
|
||||
from headroom.pricing.litellm_pricing import resolve_litellm_model
|
||||
|
||||
resolved_model = resolve_litellm_model(model)
|
||||
|
||||
# litellm.cost_per_token handles all token types natively:
|
||||
# prompt_tokens at input rate, cache_read at ~10%, cache_creation at ~125%
|
||||
|
|
@ -753,7 +702,9 @@ class CostTracker:
|
|||
if litellm is None:
|
||||
return None
|
||||
try:
|
||||
resolved = self._resolve_litellm_model(model)
|
||||
from headroom.pricing.litellm_pricing import resolve_litellm_model
|
||||
|
||||
resolved = resolve_litellm_model(model)
|
||||
info = litellm.model_cost.get(resolved, {})
|
||||
cost_per_token = info.get("input_cost_per_token")
|
||||
return cost_per_token * 1_000_000 if cost_per_token else None
|
||||
|
|
@ -770,7 +721,9 @@ class CostTracker:
|
|||
if litellm is None:
|
||||
return None
|
||||
try:
|
||||
resolved = self._resolve_litellm_model(model)
|
||||
from headroom.pricing.litellm_pricing import resolve_litellm_model
|
||||
|
||||
resolved = resolve_litellm_model(model)
|
||||
info = litellm.model_cost.get(resolved, {})
|
||||
uncached = info.get("input_cost_per_token")
|
||||
if not uncached:
|
||||
|
|
|
|||
|
|
@ -38,7 +38,9 @@ def test_savings_at_list_price():
|
|||
# Savings should be 100k tokens * list input price (NOT affected by cache mix)
|
||||
import litellm
|
||||
|
||||
resolved = ct._resolve_litellm_model(model)
|
||||
from headroom.pricing.litellm_pricing import resolve_litellm_model
|
||||
|
||||
resolved = resolve_litellm_model(model)
|
||||
info = litellm.model_cost.get(resolved, {})
|
||||
list_price = info.get("input_cost_per_token", 0)
|
||||
|
||||
|
|
|
|||
|
|
@ -24,39 +24,39 @@ class TestModelResolutionCaching:
|
|||
|
||||
def setup_method(self):
|
||||
"""Clear the cache before each test."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
CostTracker._resolved_model_cache.clear()
|
||||
lp._resolved_model_cache.clear()
|
||||
|
||||
def test_cache_returns_same_result_on_second_call(self):
|
||||
"""First call resolves, second call returns cached value without calling litellm."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
with patch.object(
|
||||
CostTracker, "_resolve_litellm_model_uncached", return_value="anthropic/claude-opus-4-6"
|
||||
with patch(
|
||||
"headroom.pricing.litellm_pricing._resolve_litellm_model_uncached",
|
||||
return_value="anthropic/claude-opus-4-6",
|
||||
) as mock_uncached:
|
||||
# First call — should invoke uncached resolution
|
||||
result1 = CostTracker._resolve_litellm_model("claude-opus-4-6")
|
||||
result1 = lp.resolve_litellm_model("claude-opus-4-6")
|
||||
assert result1 == "anthropic/claude-opus-4-6"
|
||||
assert mock_uncached.call_count == 1
|
||||
|
||||
# Second call — should use cache, NOT call uncached again
|
||||
result2 = CostTracker._resolve_litellm_model("claude-opus-4-6")
|
||||
result2 = lp.resolve_litellm_model("claude-opus-4-6")
|
||||
assert result2 == "anthropic/claude-opus-4-6"
|
||||
assert mock_uncached.call_count == 1 # Still 1, not 2
|
||||
|
||||
def test_cache_is_per_model_name(self):
|
||||
"""Different model names get separate cache entries."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
with patch.object(
|
||||
CostTracker,
|
||||
"_resolve_litellm_model_uncached",
|
||||
with patch(
|
||||
"headroom.pricing.litellm_pricing._resolve_litellm_model_uncached",
|
||||
side_effect=lambda m: f"resolved/{m}",
|
||||
) as mock_uncached:
|
||||
result1 = CostTracker._resolve_litellm_model("gpt-4o")
|
||||
result2 = CostTracker._resolve_litellm_model("claude-opus-4-6")
|
||||
result3 = CostTracker._resolve_litellm_model("gpt-4o") # cached
|
||||
result1 = lp.resolve_litellm_model("gpt-4o")
|
||||
result2 = lp.resolve_litellm_model("claude-opus-4-6")
|
||||
result3 = lp.resolve_litellm_model("gpt-4o") # cached
|
||||
|
||||
assert result1 == "resolved/gpt-4o"
|
||||
assert result2 == "resolved/claude-opus-4-6"
|
||||
|
|
@ -65,14 +65,14 @@ class TestModelResolutionCaching:
|
|||
|
||||
def test_cached_call_is_fast(self):
|
||||
"""Cached resolution should be sub-millisecond (dict lookup)."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
# Pre-populate cache
|
||||
CostTracker._resolved_model_cache["test-model"] = "resolved/test-model"
|
||||
lp._resolved_model_cache["test-model"] = "resolved/test-model"
|
||||
|
||||
start = time.perf_counter()
|
||||
for _ in range(10_000):
|
||||
CostTracker._resolve_litellm_model("test-model")
|
||||
lp.resolve_litellm_model("test-model")
|
||||
elapsed_ms = (time.perf_counter() - start) * 1000
|
||||
|
||||
# 10k lookups should take < 50ms (dict lookup is ~0.001ms each)
|
||||
|
|
@ -80,11 +80,11 @@ class TestModelResolutionCaching:
|
|||
|
||||
def test_uncached_adds_provider_prefix_for_claude(self):
|
||||
"""_resolve_litellm_model_uncached tries provider prefix for claude- models."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
with (
|
||||
patch("headroom.proxy.cost.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.proxy.cost.litellm") as mock_litellm,
|
||||
patch("headroom.pricing.litellm_pricing.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.pricing.litellm_pricing.litellm") as mock_litellm,
|
||||
):
|
||||
# First call (bare name) fails, second call (prefixed) succeeds
|
||||
mock_litellm.cost_per_token.side_effect = [
|
||||
|
|
@ -92,91 +92,89 @@ class TestModelResolutionCaching:
|
|||
(0.001, 0.002), # "anthropic/claude-opus-4-6"
|
||||
]
|
||||
|
||||
result = CostTracker._resolve_litellm_model_uncached("claude-opus-4-6")
|
||||
result = lp._resolve_litellm_model_uncached("claude-opus-4-6")
|
||||
assert result == "anthropic/claude-opus-4-6"
|
||||
|
||||
def test_uncached_adds_provider_prefix_for_gpt(self):
|
||||
"""_resolve_litellm_model_uncached tries provider prefix for gpt- models."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
with (
|
||||
patch("headroom.proxy.cost.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.proxy.cost.litellm") as mock_litellm,
|
||||
patch("headroom.pricing.litellm_pricing.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.pricing.litellm_pricing.litellm") as mock_litellm,
|
||||
):
|
||||
mock_litellm.cost_per_token.side_effect = [
|
||||
Exception("Unknown model"),
|
||||
(0.001, 0.002),
|
||||
]
|
||||
|
||||
result = CostTracker._resolve_litellm_model_uncached("gpt-4o")
|
||||
result = lp._resolve_litellm_model_uncached("gpt-4o")
|
||||
assert result == "openai/gpt-4o"
|
||||
|
||||
def test_uncached_adds_provider_prefix_for_gemini(self):
|
||||
"""_resolve_litellm_model_uncached tries provider prefix for gemini- models."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
with (
|
||||
patch("headroom.proxy.cost.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.proxy.cost.litellm") as mock_litellm,
|
||||
patch("headroom.pricing.litellm_pricing.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.pricing.litellm_pricing.litellm") as mock_litellm,
|
||||
):
|
||||
mock_litellm.cost_per_token.side_effect = [
|
||||
Exception("Unknown model"),
|
||||
(0.001, 0.002),
|
||||
]
|
||||
|
||||
result = CostTracker._resolve_litellm_model_uncached("gemini-1.5-pro")
|
||||
result = lp._resolve_litellm_model_uncached("gemini-1.5-pro")
|
||||
assert result == "google/gemini-1.5-pro"
|
||||
|
||||
def test_uncached_returns_original_when_both_fail(self):
|
||||
"""If both bare and prefixed lookups fail, return original model name."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
with (
|
||||
patch("headroom.proxy.cost.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.proxy.cost.litellm") as mock_litellm,
|
||||
patch("headroom.pricing.litellm_pricing.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.pricing.litellm_pricing.litellm") as mock_litellm,
|
||||
):
|
||||
mock_litellm.cost_per_token.side_effect = Exception("Unknown model")
|
||||
|
||||
result = CostTracker._resolve_litellm_model_uncached("totally-unknown-model-xyz")
|
||||
result = lp._resolve_litellm_model_uncached("totally-unknown-model-xyz")
|
||||
assert result == "totally-unknown-model-xyz"
|
||||
|
||||
def test_uncached_returns_original_when_litellm_unavailable(self):
|
||||
"""When litellm is not available, return model as-is."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
with patch("headroom.proxy.cost.LITELLM_AVAILABLE", False):
|
||||
result = CostTracker._resolve_litellm_model_uncached("claude-opus-4-6")
|
||||
with patch("headroom.pricing.litellm_pricing.LITELLM_AVAILABLE", False):
|
||||
result = lp._resolve_litellm_model_uncached("claude-opus-4-6")
|
||||
assert result == "claude-opus-4-6"
|
||||
|
||||
def test_uncached_returns_bare_when_it_works(self):
|
||||
"""If bare model name works, don't add prefix."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
with (
|
||||
patch("headroom.proxy.cost.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.proxy.cost.litellm") as mock_litellm,
|
||||
patch("headroom.pricing.litellm_pricing.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.pricing.litellm_pricing.litellm") as mock_litellm,
|
||||
):
|
||||
mock_litellm.cost_per_token.return_value = (0.001, 0.002)
|
||||
|
||||
result = CostTracker._resolve_litellm_model_uncached("claude-3-5-sonnet-20241022")
|
||||
result = lp._resolve_litellm_model_uncached("claude-3-5-sonnet-20241022")
|
||||
assert result == "claude-3-5-sonnet-20241022"
|
||||
|
||||
def test_cache_is_class_level_shared_across_instances(self):
|
||||
"""Cache is shared across CostTracker instances (class variable)."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
tracker1 = CostTracker()
|
||||
tracker2 = CostTracker()
|
||||
|
||||
with patch.object(
|
||||
CostTracker, "_resolve_litellm_model_uncached", return_value="resolved/model-a"
|
||||
with patch(
|
||||
"headroom.pricing.litellm_pricing._resolve_litellm_model_uncached",
|
||||
return_value="resolved/model-a",
|
||||
) as mock_uncached:
|
||||
# Resolve via instance 1
|
||||
result1 = tracker1._resolve_litellm_model("model-a")
|
||||
# Resolve
|
||||
result1 = lp.resolve_litellm_model("model-a")
|
||||
assert mock_uncached.call_count == 1
|
||||
|
||||
# Instance 2 should get cached result
|
||||
result2 = tracker2._resolve_litellm_model("model-a")
|
||||
# Second call should get cached result
|
||||
result2 = lp.resolve_litellm_model("model-a")
|
||||
assert mock_uncached.call_count == 1 # Not called again
|
||||
assert result1 == result2
|
||||
|
||||
|
|
@ -403,14 +401,14 @@ class TestConcurrentSessionSafety:
|
|||
"""Test that multiple concurrent sessions don't interfere with each other."""
|
||||
|
||||
def setup_method(self):
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
CostTracker._resolved_model_cache.clear()
|
||||
lp._resolved_model_cache.clear()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_model_resolution_is_safe(self):
|
||||
"""Multiple concurrent tasks resolving the same model should all get correct result."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
call_count = 0
|
||||
|
||||
|
|
@ -420,13 +418,13 @@ class TestConcurrentSessionSafety:
|
|||
# Simulate the slow litellm lookup
|
||||
return f"resolved/{model}"
|
||||
|
||||
with patch.object(
|
||||
CostTracker, "_resolve_litellm_model_uncached", side_effect=slow_uncached
|
||||
with patch(
|
||||
"headroom.pricing.litellm_pricing._resolve_litellm_model_uncached",
|
||||
side_effect=slow_uncached,
|
||||
):
|
||||
# Launch 50 concurrent resolution tasks for the same model
|
||||
tasks = [
|
||||
asyncio.to_thread(CostTracker._resolve_litellm_model, "claude-opus-4-6")
|
||||
for _ in range(50)
|
||||
asyncio.to_thread(lp.resolve_litellm_model, "claude-opus-4-6") for _ in range(50)
|
||||
]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
||||
|
|
@ -438,17 +436,16 @@ class TestConcurrentSessionSafety:
|
|||
@pytest.mark.asyncio
|
||||
async def test_concurrent_resolution_different_models(self):
|
||||
"""Concurrent resolution of different models should each resolve independently."""
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
models = ["gpt-4o", "claude-opus-4-6", "gemini-1.5-pro", "gpt-4o-mini"]
|
||||
|
||||
with patch.object(
|
||||
CostTracker,
|
||||
"_resolve_litellm_model_uncached",
|
||||
with patch(
|
||||
"headroom.pricing.litellm_pricing._resolve_litellm_model_uncached",
|
||||
side_effect=lambda m: f"resolved/{m}",
|
||||
):
|
||||
tasks = [
|
||||
asyncio.to_thread(CostTracker._resolve_litellm_model, model)
|
||||
asyncio.to_thread(lp.resolve_litellm_model, model)
|
||||
for model in models * 10 # 40 tasks total
|
||||
]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
|
@ -458,7 +455,7 @@ class TestConcurrentSessionSafety:
|
|||
assert results[i] == f"resolved/{model}"
|
||||
|
||||
# Cache should have exactly 4 entries
|
||||
assert len(CostTracker._resolved_model_cache) == 4
|
||||
assert len(lp._resolved_model_cache) == 4
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_streaming_errors_are_independent(self):
|
||||
|
|
@ -508,19 +505,23 @@ class TestConcurrentSessionSafety:
|
|||
@pytest.mark.asyncio
|
||||
async def test_estimate_cost_concurrent_with_caching(self):
|
||||
"""Multiple concurrent estimate_cost calls should not block each other."""
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
from headroom.proxy.server import CostTracker
|
||||
|
||||
tracker = CostTracker()
|
||||
|
||||
# Pre-populate cache to simulate steady-state
|
||||
CostTracker._resolved_model_cache["gpt-4o"] = "openai/gpt-4o"
|
||||
lp._resolved_model_cache["gpt-4o"] = "openai/gpt-4o"
|
||||
|
||||
with (
|
||||
patch("headroom.proxy.cost.LITELLM_AVAILABLE", True),
|
||||
patch("headroom.proxy.cost.litellm") as mock_litellm,
|
||||
patch("headroom.pricing.litellm_pricing.litellm") as mock_litellm,
|
||||
patch("headroom.proxy.cost.litellm") as mock_cost_litellm,
|
||||
):
|
||||
mock_litellm.cost_per_token.return_value = (0.001, 0.002)
|
||||
mock_litellm.get_model_info.return_value = {}
|
||||
mock_cost_litellm.cost_per_token.return_value = (0.001, 0.002)
|
||||
mock_cost_litellm.get_model_info.return_value = {}
|
||||
|
||||
start = time.perf_counter()
|
||||
tasks = [
|
||||
|
|
@ -544,9 +545,9 @@ class TestCostTrackingAccuracy:
|
|||
"""Test that cost calculations don't double-count cache tokens."""
|
||||
|
||||
def setup_method(self):
|
||||
from headroom.proxy.server import CostTracker
|
||||
import headroom.pricing.litellm_pricing as lp
|
||||
|
||||
CostTracker._resolved_model_cache.clear()
|
||||
lp._resolved_model_cache.clear()
|
||||
|
||||
def test_estimate_cost_separates_input_and_cache(self):
|
||||
"""Input tokens and cache tokens should be billed separately, not double-counted."""
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue