fix(replay): rerun model mapping on replay target

This commit is contained in:
fawney19
2026-03-18 23:52:32 +08:00
parent b90d5095f1
commit 8d8cddcef6
2 changed files with 44 additions and 35 deletions

View File

@@ -1920,22 +1920,15 @@ async def _resolve_replay_model_name(
db: Session,
*,
source_model: str,
original_target_model: str | None,
target_provider: Provider,
target_endpoint: ProviderEndpoint,
target_api_key: ProviderAPIKey | None,
same_provider: bool,
same_endpoint: bool,
force_remap: bool,
) -> tuple[str, str]:
"""解析 replay 的最终模型名,并返回 mapping_source。"""
"""按当前 replay 目标重新解析模型名,并返回 mapping_source。"""
from src.services.model.mapper import ModelMapperMiddleware
target_api_format = (getattr(target_endpoint, "api_format", "") or "").strip().lower()
if same_provider and same_endpoint and original_target_model and not force_remap:
return original_target_model, "original_target_model"
mapper = ModelMapperMiddleware(db)
mapping = await mapper.get_mapping(source_model, str(target_provider.id))
@@ -1946,7 +1939,6 @@ async def _resolve_replay_model_name(
)
return mapped_name, "model_mapping"
if not (same_provider and same_endpoint):
logger.debug(
"[replay] No explicit model mapping for '{}' on provider '{}' (endpoint={}, api_format={}); "
"forwarding original source model",
@@ -2161,13 +2153,9 @@ class AdminUsageReplayAdapter(AdminApiAdapter):
resolved_model_name, mapping_source = await _resolve_replay_model_name(
db,
source_model=source_model,
original_target_model=original_target_model,
target_provider=target_provider_obj,
target_endpoint=endpoint,
target_api_key=provider_key,
same_provider=same_provider,
same_endpoint=same_endpoint,
force_remap=bool(override_model),
)
mapping_applied = mapping_source != "none"

View File

@@ -280,18 +280,8 @@ async def test_admin_usage_detail_defers_large_body_columns_when_bodies_excluded
@pytest.mark.asyncio
@pytest.mark.parametrize(
("same_provider", "same_endpoint"),
[
(True, False),
(False, False),
],
ids=["same-provider-cross-endpoint", "cross-provider"],
)
async def test_resolve_replay_model_name_falls_back_to_source_model_when_mapping_missing(
monkeypatch: pytest.MonkeyPatch,
same_provider: bool,
same_endpoint: bool,
) -> None:
class _FakeMapper:
def __init__(self, db: Any) -> None:
@@ -305,14 +295,45 @@ async def test_resolve_replay_model_name_falls_back_to_source_model_when_mapping
resolved_model, mapping_source = await _resolve_replay_model_name(
SimpleNamespace(),
source_model="gpt-4o-mini",
original_target_model="provider-specific-model",
target_provider=SimpleNamespace(id="provider-2", name="OpenAI Compatible"),
target_endpoint=SimpleNamespace(id="endpoint-2", api_format="openai:responses"),
target_api_key=None,
same_provider=same_provider,
same_endpoint=same_endpoint,
force_remap=False,
)
assert resolved_model == "gpt-4o-mini"
assert mapping_source == "none"
@pytest.mark.asyncio
async def test_resolve_replay_model_name_reruns_mapping_for_same_endpoint_replay(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class _FakeModel:
def select_provider_model_name(
self, affinity_key: str | None = None, api_format: str | None = None
) -> str:
assert affinity_key == "key-2"
assert api_format == "openai:responses"
return "provider-model-for-key-2"
class _FakeMapper:
def __init__(self, db: Any) -> None:
self.db = db
async def get_mapping(self, source_model: str, provider_id: str) -> Any:
assert source_model == "gpt-4o-mini"
assert provider_id == "provider-2"
return SimpleNamespace(model=_FakeModel())
monkeypatch.setattr("src.services.model.mapper.ModelMapperMiddleware", _FakeMapper)
resolved_model, mapping_source = await _resolve_replay_model_name(
SimpleNamespace(),
source_model="gpt-4o-mini",
target_provider=SimpleNamespace(id="provider-2", name="OpenAI Compatible"),
target_endpoint=SimpleNamespace(id="endpoint-2", api_format="openai:responses"),
target_api_key=SimpleNamespace(id="key-2"),
)
assert resolved_model == "provider-model-for-key-2"
assert mapping_source == "model_mapping"