From a6d0bd1a1958698af1816ec9ece6b422d138d58e Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 23 Dec 2025 23:43:16 +0100 Subject: [PATCH] update aliases --- routstr/algorithm.py | 13 +-- routstr/core/admin.py | 190 ++++++++++++++++------------------- routstr/core/db.py | 2 + routstr/payment/models.py | 1 + routstr/proxy.py | 2 +- ui/lib/api/services/admin.ts | 5 +- 6 files changed, 98 insertions(+), 115 deletions(-) diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 515c2576..537f07f5 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -206,19 +206,20 @@ def create_model_mappings( alias: str, model: "Model", provider: "BaseUpstreamProvider" ) -> None: """Set alias to model/provider if not set or if new model is preferred.""" - existing_model = model_instances.get(alias) + alias_lower = alias.lower() + existing_model = model_instances.get(alias_lower) if not existing_model: # No existing mapping, set it - model_instances[alias] = model - provider_map[alias] = provider + model_instances[alias_lower] = model + provider_map[alias_lower] = provider else: # Check if candidate should replace existing - existing_provider = provider_map[alias] + existing_provider = provider_map[alias_lower] if should_prefer_model( model, provider, existing_model, existing_provider, alias ): - model_instances[alias] = model - provider_map[alias] = provider + model_instances[alias_lower] = model + provider_map[alias_lower] = provider def process_provider_models( upstream: "BaseUpstreamProvider", is_openrouter: bool = False diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 29ac8319..308153b7 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1493,21 +1493,11 @@ class ModelCreate(BaseModel): per_request_limits: dict[str, object] | None = None top_provider: dict[str, object] | None = None upstream_provider_id: int | None = None + canonical_slug: str | None = None + alias_ids: list[str] | None = None enabled: bool = True -class ModelUpdate(BaseModel): - id: str - name: str - description: str - created: int - context_length: int - architecture: dict[str, object] - pricing: dict[str, object] - per_request_limits: dict[str, object] | None = None - top_provider: dict[str, object] | None = None - upstream_provider_id: int | None = None - enabled: bool = True @admin_router.get("/models", response_class=HTMLResponse) @@ -2416,44 +2406,88 @@ async def admin_upstream_providers(request: Request) -> str: "/api/upstream-providers/{provider_id}/models", dependencies=[Depends(require_admin_api)], ) -async def create_provider_model( +async def upsert_provider_model( provider_id: int, payload: ModelCreate ) -> dict[str, object]: + print(payload) + logger.info(f"UPSERT_PROVIDER_MODEL called: provider_id={provider_id}, model_id={payload.id}") async with create_session() as session: provider = await session.get(UpstreamProviderRow, provider_id) if not provider: raise HTTPException(status_code=404, detail="Provider not found") - exists = await session.get(ModelRow, (payload.id, provider_id)) - if exists: - raise HTTPException( - status_code=409, - detail="Model with this ID already exists for this provider", - ) + # Try to get existing model + existing_row = await session.get(ModelRow, (payload.id, provider_id)) - row = ModelRow( - id=payload.id, - name=payload.name, - description=payload.description, - created=int(payload.created), - context_length=int(payload.context_length), - architecture=json.dumps(payload.architecture), - pricing=json.dumps(payload.pricing), - sats_pricing=None, - per_request_limits=( + if existing_row: + # Update existing model + logger.info(f"Updating existing model: {payload.id}") + existing_row.name = payload.name + existing_row.description = payload.description + existing_row.created = int(payload.created) + existing_row.context_length = int(payload.context_length) + existing_row.architecture = json.dumps(payload.architecture) + existing_row.pricing = json.dumps(payload.pricing) + existing_row.sats_pricing = None + existing_row.per_request_limits = ( json.dumps(payload.per_request_limits) if payload.per_request_limits is not None else None - ), - top_provider=( + ) + existing_row.top_provider = ( json.dumps(payload.top_provider) if payload.top_provider else None - ), - upstream_provider_id=provider_id, - enabled=payload.enabled, - ) - session.add(row) - await session.commit() - await session.refresh(row) + ) + existing_row.canonical_slug = payload.canonical_slug + existing_row.alias_ids = ( + json.dumps(payload.alias_ids) if payload.alias_ids else None + ) + was_disabled = not existing_row.enabled + existing_row.enabled = payload.enabled + + session.add(existing_row) + await session.commit() + await session.refresh(existing_row) + row = existing_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": payload.id, "error": str(e)}, + ) + else: + # Create new model + logger.info(f"Creating new model: {payload.id}") + row = ModelRow( + id=payload.id, + name=payload.name, + description=payload.description, + created=int(payload.created), + context_length=int(payload.context_length), + architecture=json.dumps(payload.architecture), + pricing=json.dumps(payload.pricing), + sats_pricing=None, + per_request_limits=( + json.dumps(payload.per_request_limits) + if payload.per_request_limits is not None + else None + ), + top_provider=( + json.dumps(payload.top_provider) if payload.top_provider else None + ), + canonical_slug=payload.canonical_slug, + alias_ids=( + json.dumps(payload.alias_ids) if payload.alias_ids else None + ), + upstream_provider_id=provider_id, + enabled=payload.enabled, + ) + session.add(row) + await session.commit() + await session.refresh(row) await refresh_model_maps() return _row_to_model( @@ -2461,6 +2495,18 @@ async def create_provider_model( ).dict() # type: ignore +@admin_router.patch( + "/api/upstream-providers/{provider_id}/models/{model_id:path}", + dependencies=[Depends(require_admin_api)], +) +async def update_provider_model_legacy( + provider_id: int, model_id: str, payload: ModelCreate +) -> dict[str, object]: + """Legacy PATCH endpoint - redirects to upsert POST endpoint for backward compatibility.""" + logger.info(f"LEGACY_PATCH_UPDATE called: provider_id={provider_id}, model_id={model_id}") + return await upsert_provider_model(provider_id, payload) + + @admin_router.get( "/api/upstream-providers/{provider_id}/models/{model_id:path}", dependencies=[Depends(require_admin_api)], @@ -2481,74 +2527,6 @@ async def get_provider_model(provider_id: int, model_id: str) -> dict[str, objec ).dict() # type: ignore -@admin_router.patch( - "/api/upstream-providers/{provider_id}/models/{model_id:path}", - dependencies=[Depends(require_admin_api)], -) -async def update_provider_model( - provider_id: int, model_id: str, payload: ModelUpdate -) -> dict[str, object]: - if payload.id != model_id: - raise HTTPException(status_code=400, detail="Path id does not match payload id") - - async with create_session() as session: - provider = await session.get(UpstreamProviderRow, provider_id) - if not provider: - raise HTTPException(status_code=404, detail="Provider not found") - - row = await session.get(ModelRow, (model_id, provider_id)) - if not row: - raise HTTPException( - status_code=404, detail="Model not found for this provider" - ) - - row.name = payload.name - row.description = payload.description - row.created = int(payload.created) - row.context_length = int(payload.context_length) - row.architecture = json.dumps(payload.architecture) - row.pricing = json.dumps(payload.pricing) - row.sats_pricing = None - row.per_request_limits = ( - json.dumps(payload.per_request_limits) - if payload.per_request_limits is not None - else None - ) - 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 - ).dict() # type: ignore - - -@admin_router.put( - "/api/upstream-providers/{provider_id}/models/{model_id:path}", - dependencies=[Depends(require_admin_api)], -) -async def update_provider_model_put( - provider_id: int, model_id: str, payload: ModelUpdate -) -> dict[str, object]: - return await update_provider_model(provider_id, model_id, payload) @admin_router.delete( diff --git a/routstr/core/db.py b/routstr/core/db.py index f40212b3..a67b5923 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -68,6 +68,8 @@ class ModelRow(SQLModel, table=True): # type: ignore sats_pricing: str | None = Field(default=None) per_request_limits: str | None = Field(default=None) top_provider: str | None = Field(default=None) + canonical_slug: str | None = Field(default=None, description="Canonical model slug") + alias_ids: str | None = Field(default=None, description="JSON array of model alias IDs") enabled: bool = Field(default=True, description="Whether this model is enabled") upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models") diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 06799705..2d2b7b17 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -185,6 +185,7 @@ def _row_to_model( enabled=row.enabled, upstream_provider_id=row.upstream_provider_id, canonical_slug=getattr(row, "canonical_slug", None), + alias_ids=json.loads(row.alias_ids) if row.alias_ids else None, ) if apply_provider_fee: diff --git a/routstr/proxy.py b/routstr/proxy.py index edceaaa6..0fa48f4f 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -70,7 +70,7 @@ def get_model_instance(model_id: str) -> Model | None: def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None: """Get UpstreamProvider for model ID from global cache.""" - return _provider_map.get(model_id) + return _provider_map.get(model_id.lower()) def get_unique_models() -> list[Model]: diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index 39269ee6..c1c11d52 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -354,8 +354,9 @@ export class AdminService { original: data.pricing, converted: payload.pricing, }); - const model = await apiClient.patch( - `/admin/api/upstream-providers/${providerId}/models/${encodeURIComponent(modelId)}`, + // Use the same POST endpoint for both create and update (upsert) + const model = await apiClient.post( + `/admin/api/upstream-providers/${providerId}/models`, payload ); return {