mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-09 02:54:37 +00:00
fix models filtering
This commit is contained in:
+21
-4
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user