diff --git a/routstr/core/admin.py b/routstr/core/admin.py index ec3cc0e7..bc44f369 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -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], } diff --git a/routstr/proxy.py b/routstr/proxy.py index d153ad01..8eece67b 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -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,