mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Merge pull request #514 from Routstr/add-provider-indentifier
add provider slug
This commit is contained in:
@@ -0,0 +1,88 @@
|
||||
"""add slug to upstream_providers
|
||||
|
||||
Revision ID: c6d7e8f9a0b1
|
||||
Revises: b5e7c9d1f3a2
|
||||
Create Date: 2026-06-29 00:00:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
from routstr.core.provider_slugs import provider_slug_base, provider_slug_candidate
|
||||
|
||||
revision = "c6d7e8f9a0b1"
|
||||
down_revision = "b5e7c9d1f3a2"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _allocate_backfill_slug(provider_type: str, reserved_slugs: set[str]) -> str:
|
||||
base = provider_slug_base(provider_type)
|
||||
suffix_number = 1
|
||||
while True:
|
||||
candidate = provider_slug_candidate(base, suffix_number)
|
||||
if candidate not in reserved_slugs:
|
||||
reserved_slugs.add(candidate)
|
||||
return candidate
|
||||
suffix_number += 1
|
||||
|
||||
|
||||
def _backfill_provider_slugs(conn: sa.Connection) -> None:
|
||||
existing_rows = conn.execute(
|
||||
sa.text(
|
||||
"SELECT slug FROM upstream_providers "
|
||||
"WHERE slug IS NOT NULL AND slug != ''"
|
||||
)
|
||||
)
|
||||
reserved_slugs = {str(row.slug).lower() for row in existing_rows}
|
||||
|
||||
rows_to_backfill = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id, provider_type FROM upstream_providers "
|
||||
"WHERE slug IS NULL OR slug = '' "
|
||||
"ORDER BY id"
|
||||
)
|
||||
)
|
||||
for row in rows_to_backfill:
|
||||
slug = _allocate_backfill_slug(str(row.provider_type), reserved_slugs)
|
||||
conn.execute(
|
||||
sa.text("UPDATE upstream_providers SET slug = :slug WHERE id = :id"),
|
||||
{"slug": slug, "id": row.id},
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
columns = {c["name"] for c in inspector.get_columns("upstream_providers")}
|
||||
|
||||
if "slug" not in columns:
|
||||
op.add_column(
|
||||
"upstream_providers",
|
||||
sa.Column("slug", sa.String(), nullable=True),
|
||||
)
|
||||
|
||||
_backfill_provider_slugs(conn)
|
||||
|
||||
existing_indexes = {idx["name"] for idx in inspector.get_indexes("upstream_providers")}
|
||||
if "ix_upstream_providers_slug" not in existing_indexes:
|
||||
op.create_index(
|
||||
"ix_upstream_providers_slug",
|
||||
"upstream_providers",
|
||||
["slug"],
|
||||
unique=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
existing_indexes = {idx["name"] for idx in inspector.get_indexes("upstream_providers")}
|
||||
if "ix_upstream_providers_slug" in existing_indexes:
|
||||
op.drop_index("ix_upstream_providers_slug", table_name="upstream_providers")
|
||||
|
||||
columns = {c["name"] for c in inspector.get_columns("upstream_providers")}
|
||||
if "slug" in columns:
|
||||
op.drop_column("upstream_providers", "slug")
|
||||
+219
-127
@@ -1,5 +1,6 @@
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import secrets
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
@@ -8,6 +9,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, RootModel
|
||||
from pydantic.v1 import ValidationError as PydanticValidationError
|
||||
from sqlmodel import select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from ..payment.models import _row_to_model, list_models
|
||||
from ..proxy import refresh_model_maps, reinitialize_upstreams
|
||||
@@ -29,6 +31,7 @@ from .db import (
|
||||
)
|
||||
from .log_manager import log_manager
|
||||
from .logging import get_logger
|
||||
from .provider_slugs import allocate_unique_provider_slug
|
||||
from .settings import SettingsService, settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
@@ -456,19 +459,18 @@ class ModelCreate(BaseModel):
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def upsert_provider_model(
|
||||
provider_id: int, payload: ModelCreate
|
||||
provider_id: str, 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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
provider_pk = _provider_pk(provider)
|
||||
|
||||
# Try to get existing model
|
||||
existing_row = await session.get(ModelRow, (payload.id, provider_id))
|
||||
existing_row = await session.get(ModelRow, (payload.id, provider_pk))
|
||||
|
||||
if existing_row:
|
||||
# Update existing model
|
||||
@@ -524,7 +526,7 @@ async def upsert_provider_model(
|
||||
alias_ids=(
|
||||
json.dumps(payload.alias_ids) if payload.alias_ids else None
|
||||
),
|
||||
upstream_provider_id=provider_id,
|
||||
upstream_provider_id=provider_pk,
|
||||
enabled=payload.enabled,
|
||||
forwarded_model_id=payload.forwarded_model_id or payload.id,
|
||||
)
|
||||
@@ -543,7 +545,7 @@ async def upsert_provider_model(
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def update_provider_model_legacy(
|
||||
provider_id: int, model_id: str, payload: ModelCreate
|
||||
provider_id: str, model_id: str, payload: ModelCreate
|
||||
) -> dict[str, object]:
|
||||
"""Legacy PATCH endpoint - redirects to upsert POST endpoint for backward compatibility."""
|
||||
logger.info(
|
||||
@@ -556,13 +558,12 @@ async def update_provider_model_legacy(
|
||||
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def get_provider_model(provider_id: int, model_id: str) -> dict[str, object]:
|
||||
async def get_provider_model(provider_id: str, model_id: str) -> dict[str, object]:
|
||||
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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
provider_pk = _provider_pk(provider)
|
||||
|
||||
row = await session.get(ModelRow, (model_id, provider_id))
|
||||
row = await session.get(ModelRow, (model_id, provider_pk))
|
||||
if not row:
|
||||
raise HTTPException(
|
||||
status_code=404, detail="Model not found for this provider"
|
||||
@@ -576,9 +577,11 @@ async def get_provider_model(provider_id: int, model_id: str) -> dict[str, objec
|
||||
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def delete_provider_model(provider_id: int, model_id: str) -> dict[str, object]:
|
||||
async def delete_provider_model(provider_id: str, model_id: str) -> dict[str, object]:
|
||||
async with create_session() as session:
|
||||
row = await session.get(ModelRow, (model_id, provider_id))
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
provider_pk = _provider_pk(provider)
|
||||
row = await session.get(ModelRow, (model_id, provider_pk))
|
||||
if not row:
|
||||
raise HTTPException(
|
||||
status_code=404, detail="Model not found for this provider"
|
||||
@@ -593,10 +596,12 @@ async def delete_provider_model(provider_id: int, model_id: str) -> dict[str, ob
|
||||
"/api/upstream-providers/{provider_id}/models",
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def delete_all_provider_models(provider_id: int) -> dict[str, object]:
|
||||
async def delete_all_provider_models(provider_id: str) -> dict[str, object]:
|
||||
async with create_session() as session:
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
provider_pk = _provider_pk(provider)
|
||||
result = await session.exec(
|
||||
select(ModelRow).where(ModelRow.upstream_provider_id == provider_id)
|
||||
select(ModelRow).where(ModelRow.upstream_provider_id == provider_pk)
|
||||
) # type: ignore
|
||||
rows = result.all()
|
||||
for row in rows:
|
||||
@@ -615,7 +620,7 @@ class BatchOverrideRequest(BaseModel):
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def batch_override_provider_models(
|
||||
provider_id: int, payload: BatchOverrideRequest
|
||||
provider_id: str, payload: BatchOverrideRequest
|
||||
) -> dict[str, object]:
|
||||
"""Batch override models for a specific provider."""
|
||||
logger.info(
|
||||
@@ -623,15 +628,14 @@ async def batch_override_provider_models(
|
||||
)
|
||||
|
||||
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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
provider_pk = _provider_pk(provider)
|
||||
|
||||
overridden_count = 0
|
||||
|
||||
for model_data in payload.models:
|
||||
# Try to get existing model regardless of whether it's enabled or not
|
||||
existing_row = await session.get(ModelRow, (model_data.id, provider_id))
|
||||
existing_row = await session.get(ModelRow, (model_data.id, provider_pk))
|
||||
|
||||
if existing_row:
|
||||
# Update existing
|
||||
@@ -685,7 +689,7 @@ async def batch_override_provider_models(
|
||||
if model_data.alias_ids
|
||||
else None
|
||||
),
|
||||
upstream_provider_id=provider_id,
|
||||
upstream_provider_id=provider_pk,
|
||||
enabled=model_data.enabled,
|
||||
)
|
||||
session.add(row)
|
||||
@@ -702,6 +706,85 @@ async def batch_override_provider_models(
|
||||
}
|
||||
|
||||
|
||||
_SLUG_PATTERN = re.compile(r"^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$")
|
||||
|
||||
|
||||
def _validate_slug(value: str) -> str:
|
||||
candidate = value.strip().lower()
|
||||
if not _SLUG_PATTERN.fullmatch(candidate):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"slug must be 3-64 chars, lowercase letters/digits/hyphens, "
|
||||
"and may not start or end with a hyphen"
|
||||
),
|
||||
)
|
||||
if candidate.isdigit():
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="slug must not be all digits",
|
||||
)
|
||||
return candidate
|
||||
|
||||
|
||||
async def _ensure_unique_slug(
|
||||
session: AsyncSession, slug: str, exclude_id: int | None = None
|
||||
) -> None:
|
||||
stmt = select(UpstreamProviderRow).where(UpstreamProviderRow.slug == slug)
|
||||
result = await session.exec(stmt)
|
||||
existing = result.first()
|
||||
if existing and existing.id != exclude_id:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail="Provider with this slug already exists",
|
||||
)
|
||||
|
||||
|
||||
async def _get_upstream_provider_by_ref(
|
||||
session: AsyncSession, provider_ref: str
|
||||
) -> UpstreamProviderRow:
|
||||
if provider_ref.isdigit():
|
||||
provider = await session.get(UpstreamProviderRow, int(provider_ref))
|
||||
else:
|
||||
slug = _validate_slug(provider_ref)
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(UpstreamProviderRow.slug == slug)
|
||||
)
|
||||
provider = result.first()
|
||||
|
||||
if not provider:
|
||||
raise HTTPException(status_code=404, detail="Provider not found")
|
||||
return provider
|
||||
|
||||
|
||||
def _provider_pk(provider: UpstreamProviderRow) -> int:
|
||||
if provider.id is None:
|
||||
raise HTTPException(status_code=500, detail="Provider has no database id")
|
||||
return provider.id
|
||||
|
||||
|
||||
def _serialize_provider(
|
||||
provider: UpstreamProviderRow, redact_api_key: bool = True
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"id": provider.id,
|
||||
"slug": provider.slug,
|
||||
"provider_type": provider.provider_type,
|
||||
"base_url": provider.base_url,
|
||||
"api_key": "[REDACTED]"
|
||||
if (redact_api_key and provider.api_key)
|
||||
else provider.api_key
|
||||
if not redact_api_key
|
||||
else "",
|
||||
"api_version": provider.api_version,
|
||||
"enabled": provider.enabled,
|
||||
"provider_fee": provider.provider_fee,
|
||||
"provider_settings": json.loads(provider.provider_settings)
|
||||
if provider.provider_settings
|
||||
else None,
|
||||
}
|
||||
|
||||
|
||||
class UpstreamProviderCreate(BaseModel):
|
||||
provider_type: str
|
||||
base_url: str
|
||||
@@ -710,6 +793,7 @@ class UpstreamProviderCreate(BaseModel):
|
||||
enabled: bool = True
|
||||
provider_fee: float = 1.01
|
||||
provider_settings: dict | None = None
|
||||
slug: str | None = None
|
||||
|
||||
|
||||
class UpstreamProviderUpdate(BaseModel):
|
||||
@@ -720,6 +804,50 @@ class UpstreamProviderUpdate(BaseModel):
|
||||
enabled: bool | None = None
|
||||
provider_fee: float | None = None
|
||||
provider_settings: dict | None = None
|
||||
slug: str | None = None
|
||||
|
||||
|
||||
class UpstreamProviderUpdateBySlug(BaseModel):
|
||||
slug: str
|
||||
new_slug: str | None = None
|
||||
provider_type: str | None = None
|
||||
base_url: str | None = None
|
||||
api_key: str | None = None
|
||||
api_version: str | None = None
|
||||
enabled: bool | None = None
|
||||
provider_fee: float | None = None
|
||||
provider_settings: dict | None = None
|
||||
|
||||
|
||||
async def _apply_provider_update(
|
||||
session: AsyncSession,
|
||||
provider: UpstreamProviderRow,
|
||||
payload: UpstreamProviderUpdate,
|
||||
new_slug: str | None = None,
|
||||
) -> None:
|
||||
if new_slug is not None:
|
||||
validated = _validate_slug(new_slug)
|
||||
await _ensure_unique_slug(session, validated, exclude_id=provider.id)
|
||||
provider.slug = validated
|
||||
|
||||
if payload.provider_type is not None:
|
||||
provider.provider_type = payload.provider_type
|
||||
if payload.base_url is not None:
|
||||
provider.base_url = payload.base_url
|
||||
if payload.api_key is not None:
|
||||
provider.api_key = payload.api_key
|
||||
if payload.api_version is not None:
|
||||
provider.api_version = payload.api_version
|
||||
if payload.enabled is not None:
|
||||
provider.enabled = payload.enabled
|
||||
if payload.provider_fee is not None:
|
||||
provider.provider_fee = payload.provider_fee
|
||||
if payload.provider_settings is not None:
|
||||
provider.provider_settings = json.dumps(payload.provider_settings)
|
||||
|
||||
session.add(provider)
|
||||
await session.commit()
|
||||
await session.refresh(provider)
|
||||
|
||||
|
||||
@admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)])
|
||||
@@ -727,21 +855,7 @@ async def get_upstream_providers() -> list[dict[str, object]]:
|
||||
async with create_session() as session:
|
||||
result = await session.exec(select(UpstreamProviderRow))
|
||||
providers = result.all()
|
||||
return [
|
||||
{
|
||||
"id": p.id,
|
||||
"provider_type": p.provider_type,
|
||||
"base_url": p.base_url,
|
||||
"api_key": "[REDACTED]" if p.api_key else "",
|
||||
"api_version": p.api_version,
|
||||
"enabled": p.enabled,
|
||||
"provider_fee": p.provider_fee,
|
||||
"provider_settings": json.loads(p.provider_settings)
|
||||
if p.provider_settings
|
||||
else None,
|
||||
}
|
||||
for p in providers
|
||||
]
|
||||
return [_serialize_provider(p) for p in providers]
|
||||
|
||||
|
||||
@admin_router.post("/api/upstream-providers", dependencies=[Depends(require_admin_api)])
|
||||
@@ -761,7 +875,14 @@ async def create_upstream_provider(
|
||||
detail="Provider with this base URL and API key already exists",
|
||||
)
|
||||
|
||||
if payload.slug:
|
||||
slug = _validate_slug(payload.slug)
|
||||
await _ensure_unique_slug(session, slug)
|
||||
else:
|
||||
slug = await allocate_unique_provider_slug(session, payload.provider_type)
|
||||
|
||||
provider = UpstreamProviderRow(
|
||||
slug=slug,
|
||||
provider_type=payload.provider_type,
|
||||
base_url=payload.base_url,
|
||||
api_key=payload.api_key,
|
||||
@@ -778,99 +899,81 @@ async def create_upstream_provider(
|
||||
|
||||
await reinitialize_upstreams()
|
||||
await refresh_model_maps()
|
||||
return {
|
||||
"id": provider.id,
|
||||
"provider_type": provider.provider_type,
|
||||
"base_url": provider.base_url,
|
||||
"api_key": "[REDACTED]",
|
||||
"api_version": provider.api_version,
|
||||
"enabled": provider.enabled,
|
||||
"provider_fee": provider.provider_fee,
|
||||
"provider_settings": payload.provider_settings,
|
||||
}
|
||||
return _serialize_provider(provider)
|
||||
|
||||
|
||||
@admin_router.get(
|
||||
"/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)]
|
||||
)
|
||||
async def get_upstream_provider(provider_id: int) -> dict[str, object]:
|
||||
async def get_upstream_provider(provider_id: str) -> dict[str, object]:
|
||||
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")
|
||||
return {
|
||||
"id": provider.id,
|
||||
"provider_type": provider.provider_type,
|
||||
"base_url": provider.base_url,
|
||||
"api_key": "[REDACTED]" if provider.api_key else "",
|
||||
"api_version": provider.api_version,
|
||||
"enabled": provider.enabled,
|
||||
"provider_fee": provider.provider_fee,
|
||||
"provider_settings": json.loads(provider.provider_settings)
|
||||
if provider.provider_settings
|
||||
else None,
|
||||
}
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
return _serialize_provider(provider)
|
||||
|
||||
|
||||
@admin_router.patch(
|
||||
"/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)]
|
||||
)
|
||||
async def update_upstream_provider(
|
||||
provider_id: int, payload: UpstreamProviderUpdate
|
||||
provider_id: str, payload: UpstreamProviderUpdate
|
||||
) -> dict[str, object]:
|
||||
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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
|
||||
if payload.provider_type is not None:
|
||||
provider.provider_type = payload.provider_type
|
||||
if payload.base_url is not None:
|
||||
provider.base_url = payload.base_url
|
||||
if payload.api_key is not None:
|
||||
provider.api_key = payload.api_key
|
||||
if payload.api_version is not None:
|
||||
provider.api_version = payload.api_version
|
||||
if payload.enabled is not None:
|
||||
provider.enabled = payload.enabled
|
||||
if payload.provider_fee is not None:
|
||||
provider.provider_fee = payload.provider_fee
|
||||
if payload.provider_settings is not None:
|
||||
provider.provider_settings = json.dumps(payload.provider_settings)
|
||||
|
||||
session.add(provider)
|
||||
await session.commit()
|
||||
await session.refresh(provider)
|
||||
await _apply_provider_update(session, provider, payload, new_slug=payload.slug)
|
||||
|
||||
await reinitialize_upstreams()
|
||||
await refresh_model_maps()
|
||||
return {
|
||||
"id": provider.id,
|
||||
"provider_type": provider.provider_type,
|
||||
"base_url": provider.base_url,
|
||||
"api_key": "[REDACTED]",
|
||||
"api_version": provider.api_version,
|
||||
"enabled": provider.enabled,
|
||||
"provider_fee": provider.provider_fee,
|
||||
"provider_settings": json.loads(provider.provider_settings)
|
||||
if provider.provider_settings
|
||||
else None,
|
||||
}
|
||||
return _serialize_provider(provider)
|
||||
|
||||
|
||||
@admin_router.patch(
|
||||
"/api/upstream-providers", dependencies=[Depends(require_admin_api)]
|
||||
)
|
||||
async def update_upstream_provider_by_slug(
|
||||
payload: UpstreamProviderUpdateBySlug,
|
||||
) -> dict[str, object]:
|
||||
lookup = _validate_slug(payload.slug)
|
||||
async with create_session() as session:
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(
|
||||
UpstreamProviderRow.slug == lookup
|
||||
)
|
||||
)
|
||||
provider = result.first()
|
||||
if not provider:
|
||||
raise HTTPException(status_code=404, detail="Provider not found")
|
||||
|
||||
update_payload = UpstreamProviderUpdate(
|
||||
provider_type=payload.provider_type,
|
||||
base_url=payload.base_url,
|
||||
api_key=payload.api_key,
|
||||
api_version=payload.api_version,
|
||||
enabled=payload.enabled,
|
||||
provider_fee=payload.provider_fee,
|
||||
provider_settings=payload.provider_settings,
|
||||
)
|
||||
await _apply_provider_update(
|
||||
session, provider, update_payload, new_slug=payload.new_slug
|
||||
)
|
||||
|
||||
await reinitialize_upstreams()
|
||||
await refresh_model_maps()
|
||||
return _serialize_provider(provider)
|
||||
|
||||
|
||||
@admin_router.delete(
|
||||
"/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)]
|
||||
)
|
||||
async def delete_upstream_provider(provider_id: int) -> dict[str, object]:
|
||||
async def delete_upstream_provider(provider_id: str) -> dict[str, object]:
|
||||
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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
deleted_id = _provider_pk(provider)
|
||||
await session.delete(provider)
|
||||
await session.commit()
|
||||
await reinitialize_upstreams()
|
||||
await refresh_model_maps()
|
||||
return {"ok": True, "deleted_id": provider_id}
|
||||
return {"ok": True, "deleted_id": deleted_id}
|
||||
|
||||
|
||||
@admin_router.get("/api/provider-types", dependencies=[Depends(require_admin_api)])
|
||||
@@ -885,17 +988,16 @@ async def get_provider_types() -> list[dict[str, object]]:
|
||||
"/api/upstream-providers/{provider_id}/models",
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def get_provider_models(provider_id: int) -> dict[str, object]:
|
||||
async def get_provider_models(provider_id: str) -> dict[str, object]:
|
||||
from ..upstream.helpers import _instantiate_provider
|
||||
|
||||
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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
provider_pk = _provider_pk(provider)
|
||||
|
||||
db_models = await list_models(
|
||||
session=session,
|
||||
upstream_id=provider_id,
|
||||
upstream_id=provider_pk,
|
||||
include_disabled=True,
|
||||
apply_fees=False,
|
||||
)
|
||||
@@ -985,13 +1087,11 @@ class TopupTokenRequest(BaseModel):
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def topup_provider_with_token(
|
||||
provider_id: int, payload: TopupTokenRequest
|
||||
provider_id: str, payload: TopupTokenRequest
|
||||
) -> dict:
|
||||
"""Redeem a Cashu token for an upstream provider."""
|
||||
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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -1022,15 +1122,13 @@ async def topup_provider_with_token(
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def initiate_provider_topup(
|
||||
provider_id: int, payload: TopupRequest
|
||||
provider_id: str, payload: TopupRequest
|
||||
) -> dict[str, object]:
|
||||
"""Initiate a Lightning Network top-up for the upstream provider account."""
|
||||
from ..upstream.helpers import _instantiate_provider
|
||||
|
||||
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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
|
||||
try:
|
||||
logger.info(
|
||||
@@ -1150,15 +1248,13 @@ async def initiate_provider_topup(
|
||||
"/api/upstream-providers/{provider_id}/topup/{invoice_id}/status",
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def check_topup_status(provider_id: int, invoice_id: str) -> dict[str, object]:
|
||||
async def check_topup_status(provider_id: str, invoice_id: str) -> dict[str, object]:
|
||||
"""Check the status of a Lightning Network top-up invoice."""
|
||||
from ..upstream.helpers import _instantiate_provider
|
||||
from ..upstream.ppqai import PPQAIUpstreamProvider
|
||||
|
||||
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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
|
||||
# For Routstr providers, proxy the status check
|
||||
if provider.provider_type == "routstr":
|
||||
@@ -1205,14 +1301,12 @@ async def check_topup_status(provider_id: int, invoice_id: str) -> dict[str, obj
|
||||
"/api/upstream-providers/{provider_id}/balance",
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def get_provider_balance(provider_id: int) -> dict[str, object]:
|
||||
async def get_provider_balance(provider_id: str) -> dict[str, object]:
|
||||
"""Get the current balance for an upstream provider account."""
|
||||
from ..upstream.helpers import _instantiate_provider
|
||||
|
||||
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")
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
|
||||
# For Routstr providers, proxy the balance check
|
||||
if provider.provider_type == "routstr":
|
||||
@@ -1591,15 +1685,13 @@ async def get_lightning_invoices_api(
|
||||
"/api/upstream-providers/{provider_id}/routstr/refund",
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def refund_routstr_provider_balance(provider_id: int) -> dict[str, object]:
|
||||
async def refund_routstr_provider_balance(provider_id: str) -> dict[str, object]:
|
||||
"""Refund balance from an upstream Routstr provider back to the local wallet."""
|
||||
from ..upstream.helpers import _instantiate_provider
|
||||
from ..upstream.routstr import RoutstrUpstreamProvider
|
||||
|
||||
async with create_session() as session:
|
||||
provider_row = await session.get(UpstreamProviderRow, provider_id)
|
||||
if not provider_row:
|
||||
raise HTTPException(status_code=404, detail="Provider not found")
|
||||
provider_row = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
|
||||
if provider_row.provider_type != "routstr":
|
||||
raise HTTPException(
|
||||
|
||||
@@ -319,6 +319,12 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
|
||||
),
|
||||
)
|
||||
id: int | None = Field(default=None, primary_key=True)
|
||||
slug: str | None = Field(
|
||||
default=None,
|
||||
unique=True,
|
||||
index=True,
|
||||
description="Stable external slug used for updates via API key.",
|
||||
)
|
||||
provider_type: str = Field(
|
||||
description="Provider type: custom, openai, anthropic, azure, openrouter, etc."
|
||||
)
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from itertools import count
|
||||
from typing import Collection
|
||||
|
||||
from sqlmodel import select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from .db import UpstreamProviderRow
|
||||
|
||||
_SLUG_BASE_PATTERN = re.compile(r"[^a-z0-9]+")
|
||||
_MAX_SLUG_LENGTH = 64
|
||||
|
||||
|
||||
def provider_slug_base(provider_type: str) -> str:
|
||||
"""Return a deterministic slug base for a provider type."""
|
||||
base = _SLUG_BASE_PATTERN.sub("-", provider_type.lower()).strip("-")
|
||||
if not base:
|
||||
base = "provider"
|
||||
elif base.isdigit():
|
||||
base = f"provider-{base}"
|
||||
elif len(base) < 3:
|
||||
base = f"{base}-provider"
|
||||
|
||||
if len(base) > _MAX_SLUG_LENGTH:
|
||||
base = base[:_MAX_SLUG_LENGTH].rstrip("-") or "provider"
|
||||
return base
|
||||
|
||||
|
||||
def provider_slug_candidate(base: str, suffix_number: int) -> str:
|
||||
if suffix_number == 1:
|
||||
return base
|
||||
|
||||
suffix = f"-{suffix_number}"
|
||||
max_base_length = _MAX_SLUG_LENGTH - len(suffix)
|
||||
return f"{base[:max_base_length].rstrip('-')}{suffix}"
|
||||
|
||||
|
||||
async def allocate_unique_provider_slug(
|
||||
session: AsyncSession,
|
||||
provider_type: str,
|
||||
reserved_slugs: Collection[str] = (),
|
||||
) -> str:
|
||||
"""Allocate a stable, deterministic provider slug.
|
||||
|
||||
The first provider of a type gets ``openai``; later collisions get
|
||||
``openai-2``, ``openai-3``, etc. ``reserved_slugs`` covers rows staged in
|
||||
memory but not flushed yet, such as settings/env seeding.
|
||||
"""
|
||||
base = provider_slug_base(provider_type)
|
||||
reserved = {slug.lower() for slug in reserved_slugs}
|
||||
|
||||
for suffix_number in count(1):
|
||||
candidate = provider_slug_candidate(base, suffix_number)
|
||||
if candidate in reserved:
|
||||
continue
|
||||
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(UpstreamProviderRow.slug == candidate)
|
||||
)
|
||||
if result.first() is None:
|
||||
return candidate
|
||||
|
||||
raise RuntimeError("unreachable")
|
||||
@@ -12,6 +12,7 @@ from sqlmodel import select
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session
|
||||
from ..core.provider_slugs import allocate_unique_provider_slug
|
||||
from ..payment.models import Model
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
@@ -250,6 +251,7 @@ async def _seed_providers_from_settings(
|
||||
|
||||
providers_to_add: list[UpstreamProviderRow] = []
|
||||
seeded_provider_keys: set[tuple[str, str]] = set()
|
||||
reserved_slugs: set[str] = set()
|
||||
|
||||
provider_classes_by_type = {
|
||||
cls.provider_type: cls
|
||||
@@ -279,8 +281,13 @@ async def _seed_providers_from_settings(
|
||||
)
|
||||
)
|
||||
if not result.first():
|
||||
slug = await allocate_unique_provider_slug(
|
||||
session, provider_type, reserved_slugs
|
||||
)
|
||||
reserved_slugs.add(slug)
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
slug=slug,
|
||||
provider_type=provider_type,
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
@@ -299,8 +306,13 @@ async def _seed_providers_from_settings(
|
||||
)
|
||||
)
|
||||
if not result.first():
|
||||
slug = await allocate_unique_provider_slug(
|
||||
session, "ollama", reserved_slugs
|
||||
)
|
||||
reserved_slugs.add(slug)
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
slug=slug,
|
||||
provider_type="ollama",
|
||||
base_url=ollama_base_url,
|
||||
api_key=ollama_api_key,
|
||||
@@ -320,8 +332,13 @@ async def _seed_providers_from_settings(
|
||||
)
|
||||
)
|
||||
if not result.first():
|
||||
slug = await allocate_unique_provider_slug(
|
||||
session, "azure", reserved_slugs
|
||||
)
|
||||
reserved_slugs.add(slug)
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
slug=slug,
|
||||
provider_type="azure",
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
@@ -342,8 +359,13 @@ async def _seed_providers_from_settings(
|
||||
)
|
||||
)
|
||||
if not result.first():
|
||||
slug = await allocate_unique_provider_slug(
|
||||
session, "custom", reserved_slugs
|
||||
)
|
||||
reserved_slugs.add(slug)
|
||||
providers_to_add.append(
|
||||
UpstreamProviderRow(
|
||||
slug=slug,
|
||||
provider_type="custom",
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
@@ -356,7 +378,7 @@ async def _seed_providers_from_settings(
|
||||
session.add(provider)
|
||||
logger.info(
|
||||
f"Seeding {provider.provider_type} provider", # type: ignore[str-format]
|
||||
extra={"base_url": provider.base_url},
|
||||
extra={"base_url": provider.base_url, "slug": provider.slug},
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
_MIGRATION_PATH = (
|
||||
Path(__file__).resolve().parents[2]
|
||||
/ "migrations"
|
||||
/ "versions"
|
||||
/ "c6d7e8f9a0b1_add_slug_to_upstream_providers.py"
|
||||
)
|
||||
_spec = importlib.util.spec_from_file_location("provider_slug_migration", _MIGRATION_PATH)
|
||||
assert _spec is not None and _spec.loader is not None
|
||||
migration = importlib.util.module_from_spec(_spec)
|
||||
_spec.loader.exec_module(migration)
|
||||
|
||||
|
||||
def test_slug_migration_backfill_uses_api_safe_deterministic_slugs() -> None:
|
||||
engine = sa.create_engine("sqlite:///:memory:")
|
||||
with engine.begin() as conn:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"CREATE TABLE upstream_providers ("
|
||||
"id INTEGER PRIMARY KEY, "
|
||||
"provider_type VARCHAR NOT NULL, "
|
||||
"slug VARCHAR NULL"
|
||||
")"
|
||||
)
|
||||
)
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"INSERT INTO upstream_providers (id, provider_type, slug) VALUES "
|
||||
"(1, 'OpenAI Compatible', NULL), "
|
||||
"(2, 'OpenAI Compatible', ''), "
|
||||
"(3, '123', NULL), "
|
||||
"(4, 'x', NULL), "
|
||||
"(5, 'anthropic', 'anthropic')"
|
||||
)
|
||||
)
|
||||
|
||||
migration._backfill_provider_slugs(conn)
|
||||
|
||||
rows = conn.execute(
|
||||
sa.text("SELECT id, slug FROM upstream_providers ORDER BY id")
|
||||
).all()
|
||||
|
||||
assert rows == [
|
||||
(1, "openai-compatible"),
|
||||
(2, "openai-compatible-2"),
|
||||
(3, "provider-123"),
|
||||
(4, "x-provider"),
|
||||
(5, "anthropic"),
|
||||
]
|
||||
@@ -0,0 +1,144 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import SQLModel, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.admin import _get_upstream_provider_by_ref
|
||||
from routstr.core.db import UpstreamProviderRow
|
||||
from routstr.core.provider_slugs import (
|
||||
allocate_unique_provider_slug,
|
||||
provider_slug_base,
|
||||
)
|
||||
from routstr.upstream.helpers import _seed_providers_from_settings
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_allocate_unique_provider_slug_is_deterministic_with_suffixes() -> None:
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
async with AsyncSession(engine) as session:
|
||||
session.add(
|
||||
UpstreamProviderRow(
|
||||
slug="openai",
|
||||
provider_type="openai",
|
||||
base_url="https://api.openai.com/v1",
|
||||
api_key="key-1",
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
assert await allocate_unique_provider_slug(session, "openai") == "openai-2"
|
||||
assert (
|
||||
await allocate_unique_provider_slug(session, "openai", {"openai-2"})
|
||||
== "openai-3"
|
||||
)
|
||||
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def test_provider_slug_base_sanitizes_provider_type() -> None:
|
||||
assert provider_slug_base("OpenAI Compatible") == "openai-compatible"
|
||||
assert provider_slug_base("!!!") == "provider"
|
||||
assert provider_slug_base("AI") == "ai-provider"
|
||||
assert provider_slug_base("123") == "provider-123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_ref_lookup_accepts_existing_numeric_ids_and_slugs() -> None:
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
async with AsyncSession(engine) as session:
|
||||
provider = UpstreamProviderRow(
|
||||
slug="openai",
|
||||
provider_type="openai",
|
||||
base_url="https://api.openai.com/v1",
|
||||
api_key="key-1",
|
||||
)
|
||||
session.add(provider)
|
||||
await session.commit()
|
||||
await session.refresh(provider)
|
||||
|
||||
by_id = await _get_upstream_provider_by_ref(session, str(provider.id))
|
||||
by_slug = await _get_upstream_provider_by_ref(session, "openai")
|
||||
|
||||
assert by_id.id == provider.id
|
||||
assert by_slug.id == provider.id
|
||||
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_seed_providers_from_settings_sets_deterministic_slug(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "seeded-openai-key")
|
||||
|
||||
class SettingsStub:
|
||||
chat_completions_api_version: str | None = None
|
||||
upstream_base_url: str | None = None
|
||||
upstream_api_key: str = ""
|
||||
|
||||
async with AsyncSession(engine) as session:
|
||||
await _seed_providers_from_settings(session, SettingsStub()) # type: ignore[arg-type]
|
||||
await session.commit()
|
||||
|
||||
result = await session.exec(select(UpstreamProviderRow))
|
||||
providers: list[UpstreamProviderRow] = list(result.all())
|
||||
|
||||
assert [(p.provider_type, p.slug) for p in providers] == [("openai", "openai")]
|
||||
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_seed_providers_from_settings_keeps_slug_stable_on_reseed(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "seeded-openai-key")
|
||||
|
||||
class SettingsStub:
|
||||
chat_completions_api_version: str | None = None
|
||||
upstream_base_url: str | None = None
|
||||
upstream_api_key: str = ""
|
||||
|
||||
async with AsyncSession(engine) as session:
|
||||
session.add(
|
||||
UpstreamProviderRow(
|
||||
slug="openai",
|
||||
provider_type="openai",
|
||||
base_url="https://example.invalid/v1",
|
||||
api_key="other-key",
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
await _seed_providers_from_settings(session, SettingsStub()) # type: ignore[arg-type]
|
||||
await session.commit()
|
||||
await _seed_providers_from_settings(session, SettingsStub()) # type: ignore[arg-type]
|
||||
await session.commit()
|
||||
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).order_by(UpstreamProviderRow.slug)
|
||||
)
|
||||
providers: list[UpstreamProviderRow] = list(result.all())
|
||||
|
||||
assert [(p.provider_type, p.slug) for p in providers] == [
|
||||
("openai", "openai"),
|
||||
("openai", "openai-2"),
|
||||
]
|
||||
|
||||
await engine.dispose()
|
||||
@@ -118,6 +118,26 @@ export function ProviderFormFields({
|
||||
/>
|
||||
)}
|
||||
|
||||
<div className='grid gap-2'>
|
||||
<Label htmlFor={`${idPrefix}slug`}>
|
||||
Slug {mode === 'create' ? '(optional, auto-generated)' : ''}
|
||||
</Label>
|
||||
<Input
|
||||
id={`${idPrefix}slug`}
|
||||
value={formData.slug || ''}
|
||||
onChange={(e) =>
|
||||
setFormData((prev) => ({
|
||||
...prev,
|
||||
slug: e.target.value || undefined,
|
||||
}))
|
||||
}
|
||||
placeholder='e.g. openai-prod'
|
||||
/>
|
||||
<p className='text-muted-foreground text-xs'>
|
||||
Stable external key used to update this provider via the admin API.
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div className='grid gap-2'>
|
||||
<Label htmlFor={`${idPrefix}base_url`}>Base URL</Label>
|
||||
<Input
|
||||
|
||||
@@ -14,6 +14,7 @@ export const ProviderTypeSchema = z.object({
|
||||
|
||||
export const UpstreamProviderSchema = z.object({
|
||||
id: z.number(),
|
||||
slug: z.string().nullable().optional(),
|
||||
provider_type: z.string(),
|
||||
base_url: z.string(),
|
||||
api_key: z.string().optional(),
|
||||
@@ -31,6 +32,7 @@ export const CreateUpstreamProviderSchema = z.object({
|
||||
enabled: z.boolean().default(true),
|
||||
provider_fee: z.number().optional(),
|
||||
provider_settings: z.record(z.string(), z.any()).nullable().optional(),
|
||||
slug: z.string().optional(),
|
||||
});
|
||||
|
||||
export const UpdateUpstreamProviderSchema = z.object({
|
||||
@@ -41,6 +43,7 @@ export const UpdateUpstreamProviderSchema = z.object({
|
||||
enabled: z.boolean().optional(),
|
||||
provider_fee: z.number().optional(),
|
||||
provider_settings: z.record(z.string(), z.any()).nullable().optional(),
|
||||
slug: z.string().optional(),
|
||||
});
|
||||
|
||||
export const AdminModelPricingSchema = z.object({
|
||||
|
||||
Reference in New Issue
Block a user