mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
fix(replay): rerun model mapping on replay target
This commit is contained in:
@@ -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,15 +1939,14 @@ 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",
|
||||
source_model,
|
||||
target_provider.name or str(target_provider.id),
|
||||
str(getattr(target_endpoint, "id", "") or "unknown"),
|
||||
target_api_format or "unknown",
|
||||
)
|
||||
logger.debug(
|
||||
"[replay] No explicit model mapping for '{}' on provider '{}' (endpoint={}, api_format={}); "
|
||||
"forwarding original source model",
|
||||
source_model,
|
||||
target_provider.name or str(target_provider.id),
|
||||
str(getattr(target_endpoint, "id", "") or "unknown"),
|
||||
target_api_format or "unknown",
|
||||
)
|
||||
|
||||
# Keep replay aligned with the normal request path: if no global-model mapping exists,
|
||||
# forward the original source model name and let the target provider validate it.
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user