import uuid import base64 import json import tempfile from decimal import Decimal from pathlib import Path from unittest.mock import patch import requests from cryptography.fernet import Fernet from django.contrib.auth import get_user_model from django.core.cache import cache from django.test import TestCase, override_settings from django.urls import path from django.utils import timezone from rest_framework.response import Response from rest_framework.test import APIClient from rest_framework.views import APIView from apps.api.authentication import ApiKeyAuthentication 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, AiProviderError, ImageGenerationResult, TextGenerationResult, ) from apps.billing.models import ( CallRecord, ExchangeRate, PointsLedger, PricingRule, RechargeOrder, ) from apps.billing.payment_gateways import ( build_mock_alipay_signature, build_mock_body_signature, ) from apps.billing.services import RechargePayment 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, } ) class DefaultAuthProbeView(APIView): def get(self, request): return Response({"ok": True}) urlpatterns = [ path("api/test-auth/", AuthenticatedEchoView.as_view()), path("api/default-auth/", DefaultAuthProbeView.as_view()), ] @override_settings(ROOT_URLCONF=__name__) class ApiKeyAuthenticationTests(TestCase): url = "/api/test-auth/" def setUp(self): cache.clear() 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_global_drf_default_does_not_accept_web_session_authentication(self): self.client.force_login(self.user) response = self.client.get("/api/default-auth/") self.assertEqual(response.status_code, 403) 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") @override_settings(API_AUTH_FAILURE_THROTTLE_RATE="1/min") def test_invalid_api_key_failures_are_throttled_by_ip(self): first = self.client.get( self.url, **self.auth_header("sk_cmhub_invalid"), REMOTE_ADDR="198.51.100.21", ) second = self.client.get( self.url, **self.auth_header("sk_cmhub_invalid"), REMOTE_ADDR="198.51.100.21", ) self.assertEqual(first.status_code, 401) self.assertEqual(second.status_code, 429) self.assertEqual(second.data["error"]["code"], "rate_limited") 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 BalanceApiTests(TestCase): url = "/api/v1/balance" def setUp(self): suffix = uuid.uuid4().hex[:8] self.user = get_user_model().objects.create_user( username=f"balance-user-{suffix}", email=f"balance-user-{suffix}@example.com", password="password", ) self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="balance") 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_balance_returns_wallet_balance_matching_ledger_sum(self): UserWallet.objects.create(user=self.user, points_balance=100) PointsLedger.objects.create( user=self.user, change_type=PointsLedger.ChangeType.RECHARGE, points_delta=120, balance_after=120, ref_order_id=1, ) call = CallRecord.objects.create( user=self.user, api_key=self.api_key, operation_type=CallRecord.OperationType.TITLE, alias="title-standard", model_used="gpt-5.5", points_cost=20, status=CallRecord.Status.SUCCESS, ) PointsLedger.objects.create( user=self.user, change_type=PointsLedger.ChangeType.CONSUME, points_delta=-20, balance_after=100, ref_call=call, ) response = self.client.get(self.url, **self.auth_header()) ledger_sum = sum( PointsLedger.objects.filter(user=self.user).values_list( "points_delta", flat=True, ) ) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["user"], self.user.username) self.assertEqual(response.data["points_balance"], 100) self.assertEqual(response.data["points_balance"], ledger_sum) def test_balance_returns_zero_without_creating_missing_wallet(self): response = self.client.get(self.url, **self.auth_header()) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["points_balance"], 0) self.assertFalse(UserWallet.objects.filter(user=self.user).exists()) def test_balance_does_not_accept_web_session_without_api_key(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 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", ) class RechargeCallbackApiTests(TestCase): wechat_url = "/api/v1/recharge/callback/wechat" alipay_url = "/api/v1/recharge/callback/alipay" def setUp(self): suffix = uuid.uuid4().hex[:8] self.user = get_user_model().objects.create_user( username=f"recharge-user-{suffix}", email=f"recharge-user-{suffix}@example.com", password="password", ) self.wallet = UserWallet.objects.create(user=self.user, points_balance=100) self.client = APIClient(enforce_csrf_checks=True) def create_order( self, *, amount="20.00", points_granted=200, pay_method=RechargeOrder.PayMethod.WEIXIN, ) -> RechargeOrder: return RechargeOrder.objects.create( user=self.user, order_no=f"R{uuid.uuid4().hex[:12]}", amount_money=Decimal(amount), pay_method=pay_method, exchange_rate=Decimal("10.0000"), points_granted=points_granted, ) def signed_wechat_body(self, order, *, total_cents=2000): payload = { "event_type": "TRANSACTION.SUCCESS", "resource": { "trade_state": "SUCCESS", "out_trade_no": order.order_no, "transaction_id": "wx-txn-001", "success_time": "2026-07-03T00:00:00+08:00", "amount": {"total": total_cents}, }, } body = json.dumps(payload, separators=(",", ":")).encode("utf-8") return body, build_mock_body_signature(body) def signed_alipay_payload(self, order, *, total_amount="20.00"): payload = { "trade_status": "TRADE_SUCCESS", "out_trade_no": order.order_no, "trade_no": "ali-txn-001", "total_amount": total_amount, "gmt_payment": "2026-07-03 00:00:00", } payload["sign"] = build_mock_alipay_signature(payload) return payload def test_wechat_callback_credits_once_and_is_csrf_exempt(self): order = self.create_order(amount="20.00", points_granted=200) body, signature = self.signed_wechat_body(order) first = self.client.post( self.wechat_url, data=body, content_type="application/json", HTTP_WECHATPAY_SIGNATURE=signature, ) second = self.client.post( self.wechat_url, data=body, content_type="application/json", HTTP_WECHATPAY_SIGNATURE=signature, ) self.assertEqual(first.status_code, 200) self.assertEqual(second.status_code, 200) self.assertEqual(first.data["code"], "SUCCESS") self.wallet.refresh_from_db() order.refresh_from_db() self.assertEqual(self.wallet.points_balance, 300) self.assertEqual(order.status, RechargeOrder.Status.PAID) self.assertEqual(order.payment_txn_no, "wx-txn-001") self.assertEqual( PointsLedger.objects.filter( user=self.user, ref_order_id=order.id, change_type=PointsLedger.ChangeType.RECHARGE, ).count(), 1, ) def test_wechat_callback_rejects_bad_signature_without_crediting(self): order = self.create_order(amount="20.00", points_granted=200) body, _signature = self.signed_wechat_body(order) response = self.client.post( self.wechat_url, data=body, content_type="application/json", HTTP_WECHATPAY_SIGNATURE="bad-signature", ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "signature_invalid") self.wallet.refresh_from_db() order.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) self.assertEqual(order.status, RechargeOrder.Status.PENDING) self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists()) def test_wechat_callback_rejects_amount_mismatch_without_crediting(self): order = self.create_order(amount="20.00", points_granted=200) body, signature = self.signed_wechat_body(order, total_cents=1999) response = self.client.post( self.wechat_url, data=body, content_type="application/json", HTTP_WECHATPAY_SIGNATURE=signature, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "amount_mismatch") self.wallet.refresh_from_db() order.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) self.assertEqual(order.status, RechargeOrder.Status.PENDING) self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists()) def test_alipay_callback_credits_once_returns_success_and_is_csrf_exempt(self): order = self.create_order( amount="20.00", points_granted=200, pay_method=RechargeOrder.PayMethod.ALIPAY, ) payload = self.signed_alipay_payload(order) first = self.client.post(self.alipay_url, data=payload) second = self.client.post(self.alipay_url, data=payload) self.assertEqual(first.status_code, 200) self.assertEqual(second.status_code, 200) self.assertEqual(first.content, b"success") self.wallet.refresh_from_db() order.refresh_from_db() self.assertEqual(self.wallet.points_balance, 300) self.assertEqual(order.status, RechargeOrder.Status.PAID) self.assertEqual(order.payment_txn_no, "ali-txn-001") self.assertEqual( PointsLedger.objects.filter( ref_order_id=order.id, change_type=PointsLedger.ChangeType.RECHARGE, ).count(), 1, ) def test_alipay_callback_rejects_bad_signature_without_crediting(self): order = self.create_order( amount="20.00", points_granted=200, pay_method=RechargeOrder.PayMethod.ALIPAY, ) payload = self.signed_alipay_payload(order) payload["sign"] = "bad-signature" response = self.client.post(self.alipay_url, data=payload) self.assertEqual(response.status_code, 400) self.assertEqual(response.content, b"fail") self.wallet.refresh_from_db() order.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) self.assertEqual(order.status, RechargeOrder.Status.PENDING) self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists()) @override_settings( PAYMENT_CALLBACK_MODE="mock", PAYMENT_MOCK_CALLBACK_SECRET="test-payment-callback-secret", PAYMENT_QR_EXPIRES_MINUTES=15, ) class RechargeCreateStatusApiTests(TestCase): create_url = "/api/v1/recharge/create" status_url = "/api/v1/recharge/status" def setUp(self): suffix = uuid.uuid4().hex[:8] self.user = get_user_model().objects.create_user( username=f"recharge-create-{suffix}", email=f"recharge-create-{suffix}@example.com", password="password", ) self.other_user = get_user_model().objects.create_user( username=f"recharge-other-{suffix}", email=f"recharge-other-{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="recharge") self.exchange_rate = ExchangeRate.objects.create( currency="CNY", points_per_unit=Decimal("10.0000"), effective_from=timezone.now(), ) self.client = APIClient() def create_order( self, *, user=None, amount="20.00", points_granted=200, pay_method=RechargeOrder.PayMethod.WEIXIN, ): return RechargeOrder.objects.create( user=user or self.user, order_no=f"R{uuid.uuid4().hex[:12]}", amount_money=Decimal(amount), pay_method=pay_method, exchange_rate=Decimal("10.0000"), points_granted=points_granted, code_url=f"mockpay://{pay_method}/existing", ) def test_recharge_create_requires_web_session_not_api_key(self): response = self.client.post( self.create_url, {"amount": "20.00", "pay_method": "weixin"}, format="json", HTTP_AUTHORIZATION=f"Bearer {self.raw_key}", ) self.assertEqual(response.status_code, 403) self.assertFalse(RechargeOrder.objects.filter(user=self.user).exists()) def test_recharge_create_is_session_authenticated_and_locks_quote(self): self.client.force_login(self.user) response = self.client.post( self.create_url, {"amount": "20.00", "pay_method": "weixin"}, format="json", ) self.assertEqual(response.status_code, 201) self.assertEqual(response.data["amount"], "20.00") self.assertEqual(response.data["exchange_rate"], "10.0000") self.assertEqual(response.data["points_granted"], 200) self.assertEqual(response.data["pay_method"], RechargeOrder.PayMethod.WEIXIN) self.assertEqual(response.data["status"], RechargeOrder.Status.PENDING) self.assertTrue(response.data["code_url"].startswith("weixin://wxpay/cmhub-mock")) self.assertIsNotNone(response.data["expires_at"]) order = RechargeOrder.objects.get(order_no=response.data["order_no"]) self.assertEqual(order.user, self.user) self.assertEqual(order.exchange_rate, Decimal("10.0000")) self.assertEqual(order.points_granted, 200) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists()) @override_settings(RECHARGE_MAX_AMOUNT_CNY="100.00") def test_recharge_create_rejects_amount_above_configured_maximum(self): self.client.force_login(self.user) response = self.client.post( self.create_url, {"amount": "100.01", "pay_method": "weixin"}, format="json", ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") self.assertFalse(RechargeOrder.objects.filter(user=self.user).exists()) def test_recharge_create_supports_alipay_mock_qr_code(self): self.client.force_login(self.user) response = self.client.post( self.create_url, {"amount": "30.00", "pay_method": "alipay"}, format="json", ) self.assertEqual(response.status_code, 201) self.assertEqual(response.data["pay_method"], RechargeOrder.PayMethod.ALIPAY) self.assertTrue(response.data["code_url"].startswith("https://qr.alipay.com/cmhub-mock")) def test_recharge_create_enforces_csrf_for_real_session_clients(self): csrf_client = APIClient(enforce_csrf_checks=True) csrf_client.force_login(self.user) response = csrf_client.post( self.create_url, {"amount": "20.00", "pay_method": "weixin"}, format="json", ) self.assertEqual(response.status_code, 403) self.assertFalse(RechargeOrder.objects.filter(user=self.user).exists()) def test_recharge_status_returns_pending_order_for_owner_only(self): order = self.create_order() self.client.force_login(self.user) response = self.client.get(self.status_url, {"order_no": order.order_no}) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["order_no"], order.order_no) self.assertEqual(response.data["status"], RechargeOrder.Status.PENDING) self.client.force_login(self.other_user) denied = self.client.get(self.status_url, {"order_no": order.order_no}) self.assertEqual(denied.status_code, 404) self.assertEqual(denied.data["error"]["code"], "order_not_found") def test_recharge_status_active_query_can_apply_paid_order_once(self): order = self.create_order(amount="20.00", points_granted=200) self.client.force_login(self.user) def fake_query(queried_order): return RechargePayment( order_no=queried_order.order_no, pay_method=queried_order.pay_method, amount=queried_order.amount_money, transaction_id="queried-txn-001", paid_at=timezone.now(), ) with patch("apps.api.views.query_payment_order", side_effect=fake_query) as query: first = self.client.get(self.status_url, {"order_no": order.order_no}) second = self.client.get(self.status_url, {"order_no": order.order_no}) self.assertEqual(first.status_code, 200) self.assertEqual(second.status_code, 200) self.assertEqual(first.data["status"], RechargeOrder.Status.PAID) self.assertEqual(second.data["status"], RechargeOrder.Status.PAID) self.assertEqual(query.call_count, 1) self.wallet.refresh_from_db() order.refresh_from_db() self.assertEqual(self.wallet.points_balance, 300) self.assertEqual(order.payment_txn_no, "queried-txn-001") self.assertEqual( PointsLedger.objects.filter( ref_order_id=order.id, change_type=PointsLedger.ChangeType.RECHARGE, ).count(), 1, ) 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"}, ) class FakeImageUrlResponse: def __init__(self, *, status_code=200, headers=None, chunks=()): self.status_code = status_code self.headers = headers or {} self._chunks = list(chunks) self.closed = False def raise_for_status(self): if self.status_code >= 400: raise requests.HTTPError("image_url request failed", response=self) def iter_content(self, chunk_size=1): for chunk in self._chunks: yield chunk def close(self): self.closed = True def dns_result(address: str): return [(None, None, None, "", (address, 443))] @override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY) class GenerateApiTests(TestCase): def setUp(self): cache.clear() 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 assert_generation_not_charged(self): self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) self.assertFalse(CallRecord.objects.filter(user=self.user).exists()) self.assertFalse(PointsLedger.objects.filter(user=self.user).exists()) 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_generate_image_downloads_safe_image_url(self): response = FakeImageUrlResponse( headers={"Content-Type": "image/jpeg"}, chunks=[b"remote-image"], ) with ( patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("93.184.216.34")), patch("apps.api.generation.requests.Session.get", return_value=response), ): result = self.post_with_provider( "/api/v1/generate/image", { "prompt": "生成图片", "model": self.image_alias, "image_url": "https://safe.example.com/input.jpg", "resolution": "1K", "aspect_ratio": "1:1", }, ) self.assertEqual(result.status_code, 200) self.assertEqual(self.provider.image_calls[0]["image"], b"remote-image") self.assertEqual(self.provider.image_calls[0]["image_mime_type"], "image/jpeg") self.assertTrue(response.closed) def test_image_url_rejects_loopback_address_without_fetch_or_charge(self): with ( patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("127.0.0.1")), patch("apps.api.generation.requests.Session.get") as image_get, ): response = self.post_with_provider( "/api/v1/generate/image", { "prompt": "生成图片", "model": self.image_alias, "image_url": "http://127.0.0.1/private.png", "resolution": "1K", }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") image_get.assert_not_called() self.assert_generation_not_charged() def test_image_url_rejects_cloud_metadata_address_without_fetch_or_charge(self): with ( patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("169.254.169.254")), patch("apps.api.generation.requests.Session.get") as image_get, ): response = self.post_with_provider( "/api/v1/generate/title", { "prompt": "生成标题", "model": self.title_alias, "image_url": "http://169.254.169.254/latest/meta-data/", }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") image_get.assert_not_called() self.assert_generation_not_charged() def test_image_url_rejects_redirect_to_private_address_without_charge(self): def fake_getaddrinfo(host, port, *args, **kwargs): if host == "safe.example.com": return dns_result("93.184.216.34") return dns_result("127.0.0.1") redirect = FakeImageUrlResponse( status_code=302, headers={"Location": "http://127.0.0.1/private.png"}, ) with ( patch("apps.api.generation.socket.getaddrinfo", side_effect=fake_getaddrinfo), patch("apps.api.generation.requests.Session.get", return_value=redirect) as image_get, ): response = self.post_with_provider( "/api/v1/generate/image", { "prompt": "生成图片", "model": self.image_alias, "image_url": "https://safe.example.com/input.png", "resolution": "1K", }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") self.assertEqual(image_get.call_count, 1) self.assertTrue(redirect.closed) self.assert_generation_not_charged() @override_settings(IMAGE_URL_MAX_BYTES=4) def test_image_url_rejects_oversized_response_without_charge(self): oversized = FakeImageUrlResponse( headers={"Content-Type": "image/png"}, chunks=[b"1234", b"5"], ) with ( patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("93.184.216.34")), patch("apps.api.generation.requests.Session.get", return_value=oversized), ): response = self.post_with_provider( "/api/v1/generate/image", { "prompt": "生成图片", "model": self.image_alias, "image_url": "https://safe.example.com/input.png", "resolution": "1K", }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") self.assertTrue(oversized.closed) self.assert_generation_not_charged() @override_settings(API_GENERATE_THROTTLE_RATE="1/min") def test_generate_endpoint_is_throttled_by_api_key_without_extra_charge(self): first = self.post_with_provider( "/api/v1/generate/title", {"prompt": "生成标题", "model": self.title_alias}, ) second = self.post_with_provider( "/api/v1/generate/title", {"prompt": "生成标题", "model": self.title_alias}, ) self.assertEqual(first.status_code, 200) self.assertEqual(second.status_code, 429) self.assertEqual(second.data["error"]["code"], "rate_limited") self.assertEqual(len(self.provider.text_calls), 1) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 98) self.assertEqual(CallRecord.objects.filter(user=self.user).count(), 1) 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, )