fix models filtering

This commit is contained in:
9qeklajc
2025-11-02 23:42:03 +01:00
parent a3d9022e2b
commit 4859dbd163
2 changed files with 23 additions and 8 deletions
+21 -4
View File
@@ -2516,12 +2516,24 @@ async def update_provider_model(
row.top_provider = (
json.dumps(payload.top_provider) if payload.top_provider else None
)
was_disabled = not row.enabled
row.enabled = payload.enabled
session.add(row)
await session.commit()
await session.refresh(row)
if was_disabled and payload.enabled:
from ..payment.models import _cleanup_enabled_models_once
try:
await _cleanup_enabled_models_once()
except Exception as e:
logger.warning(
f"Failed to run model cleanup after enabling: {e}",
extra={"model_id": model_id, "error": str(e)},
)
await refresh_model_maps()
return _row_to_model(
row, apply_provider_fee=True, provider_fee=provider.provider_fee
@@ -2784,13 +2796,13 @@ async def get_provider_models(provider_id: int) -> dict[str, object]:
session=session, upstream_id=provider_id, include_disabled=True
)
remote_models = []
upstream_models = []
upstream_instance = _instantiate_provider(provider)
if upstream_instance:
try:
raw_models = await upstream_instance.fetch_models()
remote_models = [
upstream_instance._apply_provider_fee_to_model(m).dict()
upstream_models = [
upstream_instance._apply_provider_fee_to_model(m)
for m in raw_models
]
except Exception as e:
@@ -2798,6 +2810,11 @@ async def get_provider_models(provider_id: int) -> dict[str, object]:
f"Failed to fetch models from {provider.provider_type}: {e}"
)
db_model_ids = {model.id for model in db_models}
filtered_remote_models = [
m for m in upstream_models if m.name not in db_model_ids
]
return {
"provider": {
"id": provider.id,
@@ -2805,7 +2822,7 @@ async def get_provider_models(provider_id: int) -> dict[str, object]:
"base_url": provider.base_url,
},
"db_models": [m.dict() for m in db_models],
"remote_models": remote_models,
"remote_models": [m.dict() for m in filtered_remote_models],
}
+2 -4
View File
@@ -83,7 +83,7 @@ async def refresh_model_maps() -> None:
# Gather database overrides and disabled models
async with create_session() as session:
result = await session.exec(select(ModelRow).where(ModelRow.enabled))
result = await session.exec(select(ModelRow).where(ModelRow.enabled == True))
override_rows = result.all()
provider_result = await session.exec(select(UpstreamProviderRow))
@@ -100,13 +100,11 @@ async def refresh_model_maps() -> None:
if row.upstream_provider_id is not None
}
# Get all disabled model IDs from database to filter them out
disabled_result = await session.exec(
select(ModelRow.id).where(not col(ModelRow.enabled))
select(ModelRow.id).where(ModelRow.enabled == False)
)
disabled_model_ids = {row for row in disabled_result.all()}
# Use the algorithm to create optimal mappings
_model_instances, _provider_map, _unique_models = create_model_mappings(
upstreams=_upstreams,
overrides_by_id=overrides_by_id,