This commit is contained in:
9qeklajc
2026-02-15 17:09:46 +01:00
parent b3c5e4cbf6
commit 21f421b212
2 changed files with 24 additions and 14 deletions
+21 -10
View File
@@ -456,18 +456,18 @@ async def batch_override_provider_models(
logger.info(
f"BATCH_OVERRIDE called: provider_id={provider_id}, count={len(payload.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")
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))
if existing_row:
# Update existing
existing_row.name = model_data.name
@@ -483,7 +483,9 @@ async def batch_override_provider_models(
else None
)
existing_row.top_provider = (
json.dumps(model_data.top_provider) if model_data.top_provider else None
json.dumps(model_data.top_provider)
if model_data.top_provider
else None
)
existing_row.canonical_slug = model_data.canonical_slug
existing_row.alias_ids = (
@@ -508,23 +510,32 @@ async def batch_override_provider_models(
else None
),
top_provider=(
json.dumps(model_data.top_provider) if model_data.top_provider else None
json.dumps(model_data.top_provider)
if model_data.top_provider
else None
),
canonical_slug=model_data.canonical_slug,
alias_ids=(
json.dumps(model_data.alias_ids) if model_data.alias_ids else None
json.dumps(model_data.alias_ids)
if model_data.alias_ids
else None
),
upstream_provider_id=provider_id,
enabled=model_data.enabled,
)
session.add(row)
overridden_count += 1
await session.commit()
await refresh_model_maps()
return {"ok": True, "count": overridden_count, "message": f"Successfully batch overridden {overridden_count} models"}
return {
"ok": True,
"count": overridden_count,
"message": f"Successfully batch overridden {overridden_count} models",
}
class UpstreamProviderCreate(BaseModel):
provider_type: str
+3 -4
View File
@@ -1,8 +1,7 @@
from typing import Any
import pytest
from httpx import AsyncClient
from typing import Any
from sqlmodel import select
from routstr.core.db import ApiKey
@pytest.mark.integration
@@ -18,7 +17,7 @@ async def test_wallet_info_returns_child_keys(
response = await authenticated_client.get("/v1/wallet/info")
assert response.status_code == 200
parent_data = response.json()
parent_api_key = parent_data["api_key"]
parent_data["api_key"]
# 2. Create child keys for this parent
# We need to use the parent's authentication for this