fix(ehbp): compare resolved upstream model identities

This commit is contained in:
redshift
2026-07-10 15:16:21 +08:00
parent b30346bcb6
commit 8ce91dc2b2
2 changed files with 131 additions and 43 deletions
+40 -27
View File
@@ -50,6 +50,14 @@ _TINFOIL_PROVIDER_TYPE = "tinfoil"
_TINFOIL_ALLOWED_ENCLAVE_HOST_SUFFIX = ".tinfoil.sh"
_TINFOIL_ALLOWED_ENCLAVE_HOSTS = frozenset({"tinfoil.sh"})
def _normalize_upstream_model_id(model_id: str | None) -> str:
"""Normalize casing and whitespace for upstream identity comparisons."""
if not model_id:
return ""
return model_id.strip().lower()
# Headers that must not be forwarded to the upstream enclave.
_PROXY_ONLY_HEADERS = frozenset(
{
@@ -331,50 +339,55 @@ async def _compute_ehbp_actual_cost(
actual_model: str | None = usage_dict.pop("model", None) # type: ignore[arg-type]
pricing_model_id = model_obj.id
expected_upstream_model = model_obj.forwarded_model_id or model_obj.id
# Case-insensitive comparison: ``get_model_instance`` lowercases lookup
# keys, so a casing difference between the header and the configured
# ``forwarded_model_id`` (e.g. ``GLM-5-2`` vs ``glm-5-2``) should not
# be treated as a real mismatch.
if (
actual_model
and actual_model.lower() != expected_upstream_model.lower()
):
expected_identity = _normalize_upstream_model_id(expected_upstream_model)
served_identity = _normalize_upstream_model_id(actual_model)
# Ignore casing and surrounding whitespace when comparing the model
# reported by the enclave with the expected upstream model. Version
# suffixes remain part of the identity because a configured
# ``forwarded_model_id`` may intentionally include one.
if actual_model and served_identity != expected_identity:
from ..proxy import get_model_instance
# ``forwarded_model_id`` values are registered as routable aliases in
# the global model map, so ``get_model_instance`` will find a model
# whose upstream ID matches the actually-served model. It also strips
# date-version suffixes (e.g. ``glm-5-2-20260415`` -> ``glm-5-2``),
# so a resolved model that is actually the *same* as the requested
# one is treated as a non-mismatch.
# the global model map. The resolved object can belong to a different
# provider and therefore have a different client-facing ``id`` while
# still representing the same upstream model.
actual_model_obj = get_model_instance(actual_model)
if actual_model_obj and actual_model_obj.id != model_obj.id:
logger.info(
"EHBP served model differs from requested, using actual "
"model for pricing",
if actual_model_obj is None:
logger.warning(
"EHBP served model not found in registry, falling back "
"to requested model for pricing",
extra={
"requested_model": model_obj.id,
"expected_upstream_model": expected_upstream_model,
"actual_model": actual_model,
},
)
pricing_model_id = actual_model_obj.id
actual_model = None
else:
# Either the served model is not in the registry (unknown), or
# it resolves back to the requested model (e.g. a date-versioned
# alias like ``glm-5-2-20260415``). In both cases use the
# requested model's pricing and do not propagate actual_model.
if actual_model_obj is None:
logger.warning(
"EHBP served model not found in registry, falling back "
"to requested model for pricing",
resolved_upstream_model = (
actual_model_obj.forwarded_model_id or actual_model_obj.id
)
resolved_identity = _normalize_upstream_model_id(
resolved_upstream_model
)
if resolved_identity != expected_identity:
logger.info(
"EHBP served model differs from requested, using actual "
"model for pricing",
extra={
"requested_model": model_obj.id,
"expected_upstream_model": expected_upstream_model,
"actual_model": actual_model,
"resolved_upstream_model": resolved_upstream_model,
},
)
actual_model = None # do not propagate unknown / same model
pricing_model_id = actual_model_obj.id
else:
# A different registry/client alias resolved to the same
# upstream model; retain the requested model's pricing.
actual_model = None
else:
# Models match or no model in header — use requested model's pricing.
actual_model = None
+91 -16
View File
@@ -491,6 +491,8 @@ class TestComputeEhbpActualCost:
model_obj.id = "tinfoil-glm-5-2"
model_obj.forwarded_model_id = "glm-5-2" # lowercase
with patch(
"routstr.proxy.get_model_instance"
) as mock_get_model, patch(
"routstr.upstream.ehbp.calculate_cost",
new_callable=AsyncMock,
) as mock_calc:
@@ -515,11 +517,7 @@ class TestComputeEhbpActualCost:
# No mismatch: requested model pricing used
call_args = mock_calc.call_args
assert call_args[0][0]["model"] == "tinfoil-glm-5-2"
# get_model_instance must not be consulted for a casing-only diff
assert not any(
call[0] == ("GLM-5-2",)
for call in mock_calc.call_args_list
)
mock_get_model.assert_not_called()
@pytest.mark.asyncio
async def test_date_versioned_alias_resolves_to_requested(self) -> None:
@@ -529,15 +527,14 @@ class TestComputeEhbpActualCost:
model_obj.id = "tinfoil-glm-5-2"
model_obj.forwarded_model_id = "glm-5-2"
# get_model_instance strips the date suffix and returns the SAME model
actual_model_obj = MagicMock()
actual_model_obj.id = "tinfoil-glm-5-2" # identical to requested
actual_model_obj.forwarded_model_id = "glm-5-2"
resolved_model_obj = MagicMock()
resolved_model_obj.id = "other-provider-glm-5-2"
resolved_model_obj.forwarded_model_id = "glm-5-2"
with patch(
"routstr.proxy.get_model_instance",
return_value=actual_model_obj,
), patch(
return_value=resolved_model_obj,
) as mock_get_model, patch(
"routstr.upstream.ehbp.calculate_cost",
new_callable=AsyncMock,
) as mock_calc:
@@ -552,16 +549,94 @@ class TestComputeEhbpActualCost:
input_tokens=42,
output_tokens=10,
)
# Tinfoil returns a date-versioned ID
# Tinfoil returns a date-versioned ID with different casing.
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=glm-5-2-20260415",
"prompt=42,completion=10,total=52,model=GLM-5-2-20260415",
model_obj,
100_000,
)
# No mismatch — resolves to the same model
# Registry resolution, rather than unconditional suffix removal,
# establishes that this alias represents the expected model.
assert "actual_model" not in result
call_args = mock_calc.call_args
assert call_args[0][0]["model"] == "tinfoil-glm-5-2"
assert mock_calc.call_args[0][0]["model"] == "tinfoil-glm-5-2"
mock_get_model.assert_called_once_with("GLM-5-2-20260415")
@pytest.mark.asyncio
async def test_configured_date_version_is_preserved_as_identity(self) -> None:
"""A date suffix in forwarded_model_id is meaningful and preserved."""
model_obj = MagicMock()
model_obj.id = "tinfoil-glm-5-2-20260415"
model_obj.forwarded_model_id = "glm-5-2-20260415"
with patch(
"routstr.proxy.get_model_instance"
) as mock_get_model, patch(
"routstr.upstream.ehbp.calculate_cost",
new_callable=AsyncMock,
) as mock_calc:
from routstr.payment.cost_calculation import CostData
mock_calc.return_value = CostData(
base_msats=0,
input_msats=5,
output_msats=10,
total_msats=15,
total_usd=0.0,
input_tokens=42,
output_tokens=10,
)
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=GLM-5-2-20260415",
model_obj,
100_000,
)
assert "actual_model" not in result
assert (
mock_calc.call_args[0][0]["model"]
== "tinfoil-glm-5-2-20260415"
)
mock_get_model.assert_not_called()
@pytest.mark.asyncio
async def test_different_client_alias_same_upstream_identity(self) -> None:
"""A global alias winner from another provider is not a failover when
its forwarded model ID matches the requested upstream identity."""
model_obj = MagicMock()
model_obj.id = "tinfoil-glm-5-2"
model_obj.forwarded_model_id = "glm-5-2"
resolved_model_obj = MagicMock()
resolved_model_obj.id = "other-provider-glm-5-2"
resolved_model_obj.forwarded_model_id = "GLM-5-2"
with patch(
"routstr.proxy.get_model_instance",
return_value=resolved_model_obj,
) as mock_get_model, patch(
"routstr.upstream.ehbp.calculate_cost",
new_callable=AsyncMock,
) as mock_calc:
from routstr.payment.cost_calculation import CostData
mock_calc.return_value = CostData(
base_msats=0,
input_msats=5,
output_msats=10,
total_msats=15,
total_usd=0.0,
input_tokens=42,
output_tokens=10,
)
result = await _compute_ehbp_actual_cost(
"prompt=42,completion=10,total=52,model=provider-alias",
model_obj,
100_000,
)
mock_get_model.assert_called_once_with("provider-alias")
assert "actual_model" not in result
assert mock_calc.call_args[0][0]["model"] == "tinfoil-glm-5-2"
# ---------------------------------------------------------------------------