update aliases

This commit is contained in:
9qeklajc
2025-12-23 23:43:16 +01:00
parent b01c7b2e56
commit a6d0bd1a19
6 changed files with 98 additions and 115 deletions
+7 -6
View File
@@ -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
+84 -106
View File
@@ -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(
+2
View File
@@ -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")
+1
View File
@@ -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:
+1 -1
View File
@@ -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]:
+3 -2
View File
@@ -354,8 +354,9 @@ export class AdminService {
original: data.pricing,
converted: payload.pricing,
});
const model = await apiClient.patch<AdminModel>(
`/admin/api/upstream-providers/${providerId}/models/${encodeURIComponent(modelId)}`,
// Use the same POST endpoint for both create and update (upsert)
const model = await apiClient.post<AdminModel>(
`/admin/api/upstream-providers/${providerId}/models`,
payload
);
return {