feat: add public model catalog discovery

This commit is contained in:
QiuSW
2026-07-04 10:14:16 +08:00
parent 3a661afaa5
commit 898329e047
16 changed files with 533 additions and 14 deletions
+151 -1
View File
@@ -18,7 +18,8 @@ from rest_framework.test import APIClient
from rest_framework.views import APIView
from apps.api.authentication import ApiKeyAuthentication
from apps.api.views import ExternalApiView
from apps.api.throttles import GenerateRateThrottle
from apps.api.views import ExternalApiView, ModelsView
from apps.ai.models import AiModel, ModelAlias
from apps.ai.providers import (
AiCapabilityError,
@@ -239,6 +240,155 @@ class BalanceApiTests(TestCase):
self.assertEqual(response.data["error"]["code"], "unauthorized")
class ModelsCatalogApiTests(TestCase):
url = "/api/v1/models"
def setUp(self):
cache.clear()
suffix = uuid.uuid4().hex[:8]
self.user = get_user_model().objects.create_user(
username=f"models-user-{suffix}",
email=f"models-user-{suffix}@example.com",
password="password",
)
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="models")
self.client = APIClient()
def auth_header(self, raw_key: str | None = None) -> dict:
return {"HTTP_AUTHORIZATION": f"Bearer {raw_key or self.raw_key}"}
def create_alias(
self,
*,
alias: str,
operation_type: str = ModelAlias.OperationType.TITLE,
capabilities: list[str] | None = None,
api_type: str = AiModel.ApiType.CHAT,
url: str = "https://provider-secret.example/v1/chat/completions",
model_sku: str = "secret-sku-gpt-5.5",
model_active: bool = True,
alias_active: bool = True,
) -> ModelAlias:
ai_model = AiModel.objects.create(
name=f"{alias}-{uuid.uuid4().hex[:8]}",
url=url,
model=model_sku,
api_type=api_type,
api_key_encrypted="encrypted-provider-key",
capabilities=capabilities if capabilities is not None else ["text"],
extra_body={"internal": "provider-extra-secret"},
is_active=model_active,
)
return ModelAlias.objects.create(
alias=alias,
operation_type=operation_type,
ai_model=ai_model,
is_active=alias_active,
)
def test_models_returns_public_alias_catalog_without_internal_fields(self):
title_alias = self.create_alias(alias="title-standard", capabilities=["text"])
image_alias = self.create_alias(
alias="image-edit",
operation_type=ModelAlias.OperationType.IMAGE,
capabilities=["image", "vision"],
api_type=AiModel.ApiType.IMAGES_EDITS,
url="https://provider-secret.example/v1/images/edits",
model_sku="secret-sku-image-2",
)
PricingRule.objects.create(
operation_type=title_alias.operation_type,
alias=title_alias.alias,
resolution="",
points_cost=2,
)
PricingRule.objects.create(
operation_type=image_alias.operation_type,
alias=image_alias.alias,
resolution="",
points_cost=10,
)
PricingRule.objects.create(
operation_type=image_alias.operation_type,
alias=image_alias.alias,
resolution="1k",
points_cost=12,
)
response = self.client.get(self.url, **self.auth_header())
self.assertEqual(response.status_code, 200)
self.assertNotIn(GenerateRateThrottle, ModelsView.throttle_classes)
models = {item["alias"]: item for item in response.data["models"]}
self.assertEqual(set(models), {"title-standard", "image-edit"})
self.assertEqual(
set(models["title-standard"]),
{
"alias",
"operation_type",
"capabilities",
"requires_image",
"pricing_status",
"prices",
},
)
self.assertEqual(models["title-standard"]["operation_type"], "title")
self.assertEqual(models["title-standard"]["capabilities"], ["text"])
self.assertFalse(models["title-standard"]["requires_image"])
self.assertEqual(models["title-standard"]["pricing_status"], "priced")
self.assertEqual(
models["title-standard"]["prices"],
[{"resolution": "default", "points_cost": 2}],
)
self.assertEqual(models["image-edit"]["capabilities"], ["image", "vision"])
self.assertTrue(models["image-edit"]["requires_image"])
self.assertEqual(
models["image-edit"]["prices"],
[
{"resolution": "default", "points_cost": 10},
{"resolution": "1K", "points_cost": 12},
],
)
response_body = json.dumps(response.data, ensure_ascii=False)
self.assertNotIn("secret-sku", response_body)
self.assertNotIn("provider-secret.example", response_body)
self.assertNotIn("encrypted-provider-key", response_body)
self.assertNotIn("provider-extra-secret", response_body)
self.assertNotIn("api_key", response_body)
self.assertNotIn("api_key_encrypted", response_body)
self.assertNotIn("extra_body", response_body)
self.assertNotIn("url", response_body)
self.assertNotIn("model_used", response_body)
def test_models_rejects_missing_invalid_and_session_only_authentication(self):
missing = self.client.get(self.url)
invalid = self.client.get(self.url, **self.auth_header("sk_cmhub_invalid"))
self.client.force_login(self.user)
session_only = self.client.get(self.url)
self.assertEqual(missing.status_code, 401)
self.assertEqual(missing.data["error"]["code"], "unauthorized")
self.assertEqual(invalid.status_code, 401)
self.assertEqual(invalid.data["error"]["code"], "unauthorized")
self.assertEqual(session_only.status_code, 401)
self.assertEqual(session_only.data["error"]["code"], "unauthorized")
def test_models_only_lists_callable_active_aliases_and_allows_unpriced_alias(self):
self.create_alias(alias="title-unpriced", capabilities=["text"])
self.create_alias(alias="title-inactive-alias", alias_active=False)
self.create_alias(alias="title-inactive-model", model_active=False)
self.create_alias(alias="title-wrong-capability", capabilities=["image"])
response = self.client.get(self.url, **self.auth_header())
self.assertEqual(response.status_code, 200)
self.assertEqual(len(response.data["models"]), 1)
item = response.data["models"][0]
self.assertEqual(item["alias"], "title-unpriced")
self.assertEqual(item["pricing_status"], "unpriced")
self.assertEqual(item["prices"], [])
@override_settings(
PAYMENT_CALLBACK_MODE="mock",
PAYMENT_MOCK_CALLBACK_SECRET="test-payment-callback-secret",