import uuid import base64 import tempfile from pathlib import Path from unittest.mock import patch from cryptography.fernet import Fernet from django.contrib.auth import get_user_model from django.test import TestCase, override_settings from django.urls import path from rest_framework.response import Response from rest_framework.test import APIClient from apps.api.authentication import ApiKeyAuthentication from apps.api.views import ExternalApiView from apps.ai.models import AiModel, ModelAlias from apps.ai.providers import ( AiCapabilityError, AiProviderError, ImageGenerationResult, TextGenerationResult, ) from apps.billing.models import CallRecord, PointsLedger, PricingRule from apps.users.models import ApiKey from apps.users.models import UserWallet TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii") class AuthenticatedEchoView(ExternalApiView): def get(self, request): return Response( { "user_id": request.user.id, "api_key_id": request.auth.id, } ) urlpatterns = [ path("api/test-auth/", AuthenticatedEchoView.as_view()), ] @override_settings(ROOT_URLCONF=__name__) class ApiKeyAuthenticationTests(TestCase): url = "/api/test-auth/" def setUp(self): suffix = uuid.uuid4().hex[:8] self.user = get_user_model().objects.create_user( username=f"api-user-{suffix}", email=f"api-user-{suffix}@example.com", password="password", ) self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="test") 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 test_external_api_view_only_uses_api_key_authentication(self): self.assertEqual(AuthenticatedEchoView.authentication_classes, (ApiKeyAuthentication,)) def test_valid_bearer_key_authenticates_user_and_api_key(self): response = self.client.get(self.url, **self.auth_header()) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["user_id"], self.user.id) self.assertEqual(response.data["api_key_id"], self.api_key.id) self.api_key.refresh_from_db() self.assertIsNotNone(self.api_key.last_used_at) def test_missing_api_key_returns_401(self): response = self.client.get(self.url) self.assertEqual(response.status_code, 401) self.assertEqual(response["WWW-Authenticate"], "Bearer") self.assertEqual(response.data["error"]["code"], "unauthorized") def test_invalid_api_key_returns_401(self): response = self.client.get(self.url, **self.auth_header("sk_cmhub_invalid")) self.assertEqual(response.status_code, 401) self.assertEqual(response["WWW-Authenticate"], "Bearer") self.assertEqual(response.data["error"]["code"], "unauthorized") def test_malformed_authorization_header_returns_401(self): response = self.client.get(self.url, HTTP_AUTHORIZATION=f"Token {self.raw_key}") self.assertEqual(response.status_code, 401) self.assertEqual(response.data["error"]["code"], "unauthorized") def test_revoked_api_key_returns_403(self): self.api_key.status = ApiKey.Status.REVOKED self.api_key.save(update_fields=("status", "updated_at")) response = self.client.get(self.url, **self.auth_header()) self.assertEqual(response.status_code, 403) self.assertEqual(response.data["error"]["code"], "account_disabled") def test_disabled_user_returns_403(self): self.user.status = self.user.Status.DISABLED self.user.save(update_fields=("status",)) response = self.client.get(self.url, **self.auth_header()) self.assertEqual(response.status_code, 403) self.assertEqual(response.data["error"]["code"], "account_disabled") def test_web_session_login_is_not_accepted_for_external_api(self): self.client.force_login(self.user) response = self.client.get(self.url) self.assertEqual(response.status_code, 401) self.assertEqual(response.data["error"]["code"], "unauthorized") class FakeGenerationProvider: def __init__(self, *, capabilities=None): self._capabilities = set(capabilities or {"text", "image", "vision"}) self.text_calls = [] self.image_calls = [] self.text_error = None self.image_error = None def capabilities(self): return set(self._capabilities) def generate_text(self, prompt, model, **kwargs): self.text_calls.append({"prompt": prompt, "model": model, **kwargs}) if self.text_error is not None: raise self.text_error return TextGenerationResult( text="测试标题一", titles=("测试标题一", "测试标题二"), model_used=model.model, raw={"secret": "SECRET_RAW_SHOULD_NOT_BE_STORED"}, ) def generate_image(self, prompt, model, **kwargs): self.image_calls.append({"prompt": prompt, "model": model, **kwargs}) if self.image_error is not None: raise self.image_error return ImageGenerationResult( image=b"generated-image-bytes", model_used=model.model, raw={"b64_json": "SECRET_RAW_SHOULD_NOT_BE_STORED"}, ) @override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY) class GenerateApiTests(TestCase): def setUp(self): suffix = uuid.uuid4().hex[:8] self.media_dir = tempfile.TemporaryDirectory() self.addCleanup(self.media_dir.cleanup) media_override = override_settings( MEDIA_ROOT=self.media_dir.name, MEDIA_URL="/media/", ) media_override.enable() self.addCleanup(media_override.disable) self.user = get_user_model().objects.create_user( username=f"generate-user-{suffix}", email=f"generate-user-{suffix}@example.com", password="password", ) self.wallet = UserWallet.objects.create(user=self.user, points_balance=100) self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="generate") self.client = APIClient() self.provider = FakeGenerationProvider() self.title_model = self.create_ai_model( name=f"title-model-{suffix}", model=f"gpt-title-{suffix}", capabilities=["text", "vision"], ) self.image_model = self.create_ai_model( name=f"image-model-{suffix}", model=f"gpt-image-{suffix}", capabilities=["image", "vision"], ) self.title_alias = f"title-standard-{suffix}" self.image_alias = f"image-hd-{suffix}" ModelAlias.objects.create( operation_type=ModelAlias.OperationType.TITLE, alias=self.title_alias, ai_model=self.title_model, is_default=True, ) ModelAlias.objects.create( operation_type=ModelAlias.OperationType.IMAGE, alias=self.image_alias, ai_model=self.image_model, is_default=True, ) PricingRule.objects.create( operation_type=CallRecord.OperationType.TITLE, alias=self.title_alias, points_cost=2, ) PricingRule.objects.create( operation_type=CallRecord.OperationType.IMAGE, alias=self.image_alias, resolution="1K", points_cost=10, ) def create_ai_model(self, *, name, model, capabilities): ai_model = AiModel( name=name, url="https://api.example.test/v1", model=model, api_type=AiModel.ApiType.CHAT, capabilities=capabilities, ) ai_model.set_api_key("sk-test-secret") ai_model.save() return ai_model def auth_header(self) -> dict: return {"HTTP_AUTHORIZATION": f"Bearer {self.raw_key}"} def post_with_provider(self, path, payload, provider=None): with patch("apps.api.generation.get_provider", return_value=provider or self.provider): return self.client.post(path, payload, format="json", **self.auth_header()) def test_generate_title_uses_default_alias_charges_points_and_writes_call_record(self): response = self.post_with_provider( "/api/v1/generate/title", { "prompt": "生成标题", "resolution": "1k", "parameters": {"temperature": 0.2, "model": "bad-overridden-model"}, }, ) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["titles"], ["测试标题一", "测试标题二"]) self.assertEqual(response.data["alias"], self.title_alias) self.assertEqual(response.data["model_used"], self.title_model.model) self.assertEqual(response.data["points_cost"], 2) self.assertEqual(response.data["points_balance"], 98) self.assertEqual(len(self.provider.text_calls), 1) self.assertEqual(self.provider.text_calls[0]["model"].model, self.title_model.model) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 98) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual(call.status, CallRecord.Status.SUCCESS) self.assertEqual(call.api_key, self.api_key) self.assertEqual(call.alias, self.title_alias) self.assertEqual(call.model_used, self.title_model.model) self.assertEqual(call.resolution, "1K") self.assertNotIn("SECRET_RAW", call.result_summary) self.assertEqual( PointsLedger.objects.filter( user=self.user, change_type=PointsLedger.ChangeType.CONSUME, ).count(), 1, ) def test_generate_image_stores_file_returns_url_and_does_not_store_raw_base64(self): encoded = base64.b64encode(b"input-image").decode("ascii") response = self.post_with_provider( "/api/v1/generate/image", { "prompt": "生成图片", "model": self.image_alias, "image_base64": f"data:image/png;base64,{encoded}", "resolution": "1K", "aspect_ratio": "1:1", }, ) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["alias"], self.image_alias) self.assertEqual(response.data["model_used"], self.image_model.model) self.assertEqual(response.data["points_cost"], 10) self.assertEqual(response.data["points_balance"], 90) self.assertTrue(response.data["image_url"].startswith("http://testserver/media/")) self.assertEqual(self.provider.image_calls[0]["image"], b"input-image") media_relative_path = response.data["image_url"].split("/media/", 1)[1] self.assertTrue((Path(self.media_dir.name) / media_relative_path).exists()) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual(call.status, CallRecord.Status.SUCCESS) self.assertEqual(call.result_ref, response.data["image_url"]) self.assertEqual(call.result_summary, "image_bytes=21") self.assertNotIn("SECRET_RAW", call.result_ref + call.result_summary) def test_insufficient_points_returns_402_without_calling_provider_or_writing_call(self): self.wallet.points_balance = 1 self.wallet.save(update_fields=("points_balance", "updated_at")) response = self.post_with_provider( "/api/v1/generate/title", {"prompt": "生成标题", "model": self.title_alias}, ) self.assertEqual(response.status_code, 402) self.assertEqual(response.data["error"]["code"], "insufficient_points") self.assertEqual(self.provider.text_calls, []) self.assertFalse(CallRecord.objects.filter(user=self.user).exists()) self.assertFalse(PointsLedger.objects.filter(user=self.user).exists()) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 1) def test_missing_pricing_rule_returns_400_without_charging(self): PricingRule.objects.filter(alias=self.title_alias).delete() response = self.post_with_provider( "/api/v1/generate/title", {"prompt": "生成标题", "model": self.title_alias}, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "no_pricing_rule") self.assertEqual(self.provider.text_calls, []) self.assertFalse(CallRecord.objects.filter(user=self.user).exists()) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) def test_alias_capability_mismatch_returns_model_not_allowed_without_charging(self): bad_alias = f"bad-title-{uuid.uuid4().hex[:8]}" ModelAlias.objects.create( operation_type=ModelAlias.OperationType.TITLE, alias=bad_alias, ai_model=self.image_model, ) PricingRule.objects.create( operation_type=CallRecord.OperationType.TITLE, alias=bad_alias, points_cost=2, ) response = self.post_with_provider( "/api/v1/generate/title", {"prompt": "生成标题", "model": bad_alias}, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "model_not_allowed") self.assertEqual(self.provider.text_calls, []) self.assertFalse(CallRecord.objects.filter(alias=bad_alias).exists()) def test_provider_capability_mismatch_returns_model_not_allowed_before_charging(self): text_only_provider = FakeGenerationProvider(capabilities={"text"}) response = self.post_with_provider( "/api/v1/generate/image", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, provider=text_only_provider, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "model_not_allowed") self.assertEqual(text_only_provider.image_calls, []) self.assertFalse(CallRecord.objects.filter(user=self.user).exists()) def test_upstream_failure_refunds_precharged_points_and_marks_call_failed(self): self.provider.text_error = AiProviderError("provider timeout") response = self.post_with_provider( "/api/v1/generate/title", {"prompt": "生成标题", "model": self.title_alias}, ) self.assertEqual(response.status_code, 502) self.assertEqual(response.data["error"]["code"], "upstream_error") self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) call = CallRecord.objects.get(user=self.user) self.assertEqual(call.status, CallRecord.Status.FAILED) self.assertIn("provider timeout", call.error_message) self.assertEqual( PointsLedger.objects.filter( ref_call=call, change_type=PointsLedger.ChangeType.CONSUME, ).count(), 1, ) self.assertEqual( PointsLedger.objects.filter( ref_call=call, change_type=PointsLedger.ChangeType.REFUND, ).count(), 1, ) def test_provider_capability_error_returns_400_and_refunds_points(self): self.provider.image_error = AiCapabilityError("input image is required") response = self.post_with_provider( "/api/v1/generate/image", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) call = CallRecord.objects.get(user=self.user) self.assertEqual(call.status, CallRecord.Status.FAILED) self.assertEqual( PointsLedger.objects.filter( ref_call=call, change_type=PointsLedger.ChangeType.REFUND, ).count(), 1, )