import uuid import base64 import io import json import tempfile from datetime import timedelta from decimal import Decimal from pathlib import Path from unittest.mock import patch import requests from cryptography.fernet import Fernet from django.conf import settings from django.contrib import admin from django.contrib.auth import get_user_model from django.core.cache import cache from django.core.files.base import ContentFile from django.core.management import call_command from django.test import TestCase, override_settings from django.urls import path, reverse 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.generation import ( ApiRequestError, GenerationInput, execute_precharged_generation, precharge_generation, prepare_generation, run_synchronous_generation, ) from apps.api.image_tasks import ( claim_next_image_task, reap_stale_image_tasks, run_image_generation_task, ) from apps.api.models import ImageGenerationTask, ImageGenerationTaskInput from apps.api.throttles import GenerateRateThrottle from apps.api.views import ClientLatestReleaseView, 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, PaymentOrderCode, ) from apps.billing.services import RechargePayment from apps.moderation.models import SensitiveWord from apps.moderation.providers.keyword import reset_keyword_matcher_cache from apps.portal.models import DownloadRelease from apps.licensing.models import ClientDevice, SoftwareOrder, SoftwarePlan from apps.licensing.services import ( create_software_order, grant_software_entitlement, register_device, ) 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", first_name="主账号", ) 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) self.assertEqual( response.data["account"], { "username": self.user.username, "display_name": "主账号", }, ) self.assertNotIn("email", response.data["account"]) self.assertNotIn("id", response.data["account"]) 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["user"], self.user.username) self.assertEqual(response.data["points_balance"], 0) self.assertEqual(response.data["account"]["username"], self.user.username) self.assertEqual(response.data["account"]["display_name"], "主账号") 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", ) vision_alias = self.create_alias( alias="vision-standard", operation_type=ModelAlias.OperationType.VISION, capabilities=["text", "vision"], model_sku="secret-sku-vision", ) 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, ) PricingRule.objects.create( operation_type=vision_alias.operation_type, alias=vision_alias.alias, resolution="", points_cost=3, ) 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", "vision-standard"}) 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}, ], ) self.assertEqual(models["vision-standard"]["operation_type"], "vision") self.assertEqual(models["vision-standard"]["capabilities"], ["text", "vision"]) self.assertTrue(models["vision-standard"]["requires_image"]) self.assertEqual(models["vision-standard"]["pricing_status"], "priced") self.assertEqual( models["vision-standard"]["prices"], [{"resolution": "default", "points_cost": 3}], ) 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"], []) class ClientLatestReleaseApiTests(TestCase): url = "/api/v1/client/releases/latest" def setUp(self): cache.clear() self.client = APIClient() def create_release( self, *, platform: str = DownloadRelease.Platform.WINDOWS, version: str = "1.0.0", is_current: bool = True, external_url: str = "https://download.example.com/cmhub-desktop.exe", file_name: str = "", sha256: str = "a" * 64, release_notes: str = "首版 Windows 客户端", force_update: bool = False, size_bytes: int | None = None, ) -> DownloadRelease: return DownloadRelease.objects.create( platform=platform, version=version, is_current=is_current, external_url=external_url, file=file_name, sha256=sha256, release_notes=release_notes, force_update=force_update, size_bytes=size_bytes, ) def test_latest_release_is_public_without_api_key_and_returns_current_release(self): release = self.create_release( version="1.2.3", external_url="https://download.example.com/cmhub-1.2.3.exe", sha256="b" * 64, release_notes="修复下载入口并补充 SHA256", size_bytes=18_765_432, ) response = self.client.get(self.url) self.assertEqual(response.status_code, 200) self.assertNotIn(GenerateRateThrottle, ClientLatestReleaseView.throttle_classes) self.assertEqual(response.data["platform"], "windows") self.assertEqual( set(response.data["release"]), { "version", "download_url", "sha256", "release_notes", "force_update", "size_bytes", "published_at", }, ) self.assertEqual(response.data["release"]["version"], "1.2.3") self.assertFalse(response.data["release"]["force_update"]) self.assertEqual(response.data["release"]["size_bytes"], 18_765_432) self.assertEqual( response.data["release"]["download_url"], "https://download.example.com/cmhub-1.2.3.exe", ) self.assertEqual(response.data["release"]["sha256"], "b" * 64) self.assertEqual( response.data["release"]["release_notes"], "修复下载入口并补充 SHA256", ) self.assertEqual( response.data["release"]["published_at"], timezone.localtime(release.updated_at).isoformat(), ) def test_latest_release_returns_force_update_true(self): self.create_release( version="0.1.1", external_url="https://download.example.com/cmhub-0.1.1.zip", release_notes="优化了ai模块的生图的功能", force_update=True, ) response = self.client.get(f"{self.url}?platform=windows") self.assertEqual(response.status_code, 200) self.assertEqual(response.data["platform"], "windows") self.assertEqual(response.data["release"]["version"], "0.1.1") self.assertEqual( response.data["release"]["download_url"], "https://download.example.com/cmhub-0.1.1.zip", ) self.assertEqual( response.data["release"]["release_notes"], "优化了ai模块的生图的功能", ) self.assertTrue(response.data["release"]["force_update"]) def test_latest_release_returns_null_size_bytes_when_not_configured(self): self.create_release(size_bytes=None) response = self.client.get(self.url) self.assertEqual(response.status_code, 200) self.assertIsNone(response.data["release"]["size_bytes"]) def test_latest_release_ignores_web_session_and_does_not_return_user_data(self): user = get_user_model().objects.create_user( username="release-session-user", email="release-session-user@example.com", password="password", ) self.create_release() self.client.force_login(user) response = self.client.get(self.url, HTTP_AUTHORIZATION="Bearer sk_cmhub_invalid") self.assertEqual(response.status_code, 200) response_body = json.dumps(response.data, ensure_ascii=False) self.assertNotIn(user.username, response_body) self.assertNotIn(user.email, response_body) self.assertNotIn("api_key", response_body) self.assertNotIn("key_hash", response_body) def test_latest_release_builds_absolute_file_url(self): self.create_release( external_url="", file_name="downloads/cmhub-desktop-1.0.0.exe", ) response = self.client.get(self.url, secure=True) self.assertEqual(response.status_code, 200) self.assertEqual( response.data["release"]["download_url"], "https://testserver/media/downloads/cmhub-desktop-1.0.0.exe", ) def test_latest_release_prefers_external_url_over_uploaded_file(self): self.create_release( external_url="https://cdn.example.com/cmhub-desktop-1.0.0.exe", file_name="downloads/local-secret-name.exe", ) response = self.client.get(self.url, secure=True) self.assertEqual(response.status_code, 200) self.assertEqual( response.data["release"]["download_url"], "https://cdn.example.com/cmhub-desktop-1.0.0.exe", ) response_body = json.dumps(response.data, ensure_ascii=False) self.assertNotIn("local-secret-name.exe", response_body) self.assertNotIn(str(settings.MEDIA_ROOT), response_body) def test_latest_release_returns_unpublished_when_no_current_release(self): self.create_release(version="0.9.0", is_current=False) response = self.client.get(f"{self.url}?platform=windows") self.assertEqual(response.status_code, 200) self.assertEqual( response.data, { "platform": "windows", "release": None, "message": "暂未发布", }, ) def test_latest_release_returns_unpublished_when_current_release_has_no_download_url(self): self.create_release(external_url="", file_name="") response = self.client.get(self.url) self.assertEqual(response.status_code, 200) self.assertIsNone(response.data["release"]) self.assertEqual(response.data["message"], "暂未发布") self.assertNotIn("force_update", response.data) self.assertNotIn("size_bytes", response.data) def test_latest_release_rejects_invalid_platform(self): response = self.client.get(f"{self.url}?platform=android") self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") def test_latest_release_supports_non_windows_platform(self): self.create_release( platform=DownloadRelease.Platform.MACOS, version="2.0.0", external_url="https://download.example.com/cmhub-2.0.0.dmg", sha256="c" * 64, release_notes="macOS 客户端", ) response = self.client.get(f"{self.url}?platform=macos") self.assertEqual(response.status_code, 200) self.assertEqual(response.data["platform"], "macos") self.assertEqual(response.data["release"]["version"], "2.0.0") def test_latest_release_response_does_not_expose_internal_fields(self): self.create_release( external_url="", file_name="downloads/cmhub-desktop-1.0.0.exe", ) response = self.client.get(self.url, secure=True) self.assertEqual(response.status_code, 200) self.assertEqual(set(response.data), {"platform", "release"}) self.assertEqual( set(response.data["release"]), { "version", "download_url", "sha256", "release_notes", "force_update", "size_bytes", "published_at", }, ) response_body = json.dumps(response.data, ensure_ascii=False) for forbidden in ( "id", "is_current", "created_at", "updated_at", "MEDIA_ROOT", str(settings.MEDIA_ROOT), "user", "email", "api_key", "api_key_encrypted", "model_used", ): self.assertNotIn(forbidden, response_body) def test_download_release_admin_exposes_release_metadata_fields(self): registered_admin = admin.site._registry[DownloadRelease] self.assertIn("force_update", registered_admin.list_display) self.assertIn("force_update", registered_admin.list_filter) self.assertIn("size_bytes", registered_admin.list_display) version_fields = registered_admin.fieldsets[0][1]["fields"] self.assertIn("force_update", version_fields) download_fields = registered_admin.fieldsets[1][1]["fields"] self.assertIn("size_bytes", download_fields) class ImageGenerationTaskAdminTests(TestCase): def setUp(self): self.media_dir = tempfile.TemporaryDirectory() self.addCleanup(self.media_dir.cleanup) self.media_override = override_settings(MEDIA_ROOT=self.media_dir.name) self.media_override.enable() self.addCleanup(self.media_override.disable) suffix = uuid.uuid4().hex[:8] user_model = get_user_model() self.admin_user = user_model.objects.create_superuser( username=f"image-task-admin-{suffix}", email=f"image-task-admin-{suffix}@example.com", password="test-password", ) self.user = user_model.objects.create_user( username=f"image-task-user-{suffix}", email=f"image-task-user-{suffix}@example.com", password="test-password", ) self.api_key, _raw_key = ApiKey.create_for_user(self.user, name="image-task-admin-test") self.client.force_login(self.admin_user) def create_task(self, *, status=ImageGenerationTask.Status.SUCCEEDED, result_url=""): call_record = CallRecord.objects.create( user=self.user, api_key=self.api_key, operation_type=CallRecord.OperationType.IMAGE, alias="image-standard", model_used="test-image-model", resolution="1K", prompt="生成商品主图", points_cost=10, status=CallRecord.Status.SUCCESS, ) return ImageGenerationTask.objects.create( user=self.user, api_key=self.api_key, call_record=call_record, status=status, request_hash=uuid.uuid4().hex + uuid.uuid4().hex, result_url=result_url, points_balance_after_charge=90, ) def change_url(self, task): return reverse("admin:api_imagegenerationtask_change", args=(task.pk,)) def add_input_image(self, task, *, ordinal, filename): task_input = ImageGenerationTaskInput( task=task, ordinal=ordinal, mime_type="image/png", filename=filename, ) task_input.image.save(filename, ContentFile(b"test-image"), save=True) return task_input def test_change_view_shows_ordered_input_and_result_thumbnails_with_modal_preview(self): task = self.create_task(result_url="https://images.example.test/generated.png") main = self.add_input_image(task, ordinal=0, filename="main.png") reference = self.add_input_image(task, ordinal=1, filename="reference.png") response = self.client.get(self.change_url(task)) self.assertEqual(response.status_code, 200) self.assertContains(response, "图片预览") self.assertContains(response, "主图") self.assertContains(response, "参考图 1") self.assertContains(response, "生成结果") self.assertContains(response, main.image.url) self.assertContains(response, reference.image.url) self.assertContains(response, task.result_url) self.assertContains(response, "data-image-preview-dialog") self.assertContains(response, "image-task-gallery.js") self.assertContains(response, "双击查看大图") def test_change_view_uses_legacy_single_input_image_as_main_image(self): task = self.create_task() task.input_image.save("legacy-main.png", ContentFile(b"legacy-image"), save=True) response = self.client.get(self.change_url(task)) self.assertEqual(response.status_code, 200) self.assertContains(response, "主图") self.assertContains(response, task.input_image.url) self.assertNotContains(response, "参考图 1") def test_change_view_handles_task_without_images_or_result(self): task = self.create_task(status=ImageGenerationTask.Status.QUEUED) response = self.client.get(self.change_url(task)) self.assertEqual(response.status_code, 200) self.assertNotContains(response, 'id="image-task-gallery-title"') def test_change_view_requires_staff_access(self): task = self.create_task() self.client.force_login(self.user) response = self.client.get(self.change_url(task)) self.assertEqual(response.status_code, 302) def test_changelist_filters_new_and_legacy_single_image_tasks(self): new_single = self.create_task() self.add_input_image(new_single, ordinal=0, filename="new-single.png") legacy_single = self.create_task() legacy_single.input_image.save("legacy-single.png", ContentFile(b"legacy-image"), save=True) multiple = self.create_task() self.add_input_image(multiple, ordinal=0, filename="multiple-main.png") self.add_input_image(multiple, ordinal=1, filename="multiple-reference.png") no_input = self.create_task() response = self.client.get( reverse("admin:api_imagegenerationtask_changelist"), {"input_image_type": "single"}, ) self.assertEqual(response.status_code, 200) self.assertContains(response, "输入图片类型") self.assertContains(response, "单图生图") self.assertContains(response, "多图生图") self.assertContains(response, str(new_single.task_id)) self.assertContains(response, str(legacy_single.task_id)) self.assertNotContains(response, str(multiple.task_id)) self.assertNotContains(response, str(no_input.task_id)) self.assertCountEqual( response.context["cl"].queryset.values_list("pk", flat=True), [new_single.pk, legacy_single.pk], ) def test_changelist_filters_multiple_image_tasks_and_composes_with_status_filter(self): succeeded_multiple = self.create_task(status=ImageGenerationTask.Status.SUCCEEDED) self.add_input_image(succeeded_multiple, ordinal=0, filename="succeeded-main.png") self.add_input_image(succeeded_multiple, ordinal=1, filename="succeeded-reference.png") failed_multiple = self.create_task(status=ImageGenerationTask.Status.FAILED) self.add_input_image(failed_multiple, ordinal=0, filename="failed-main.png") self.add_input_image(failed_multiple, ordinal=1, filename="failed-reference.png") single = self.create_task() self.add_input_image(single, ordinal=0, filename="single.png") response = self.client.get( reverse("admin:api_imagegenerationtask_changelist"), { "input_image_type": "multiple", "status__exact": ImageGenerationTask.Status.SUCCEEDED, }, ) self.assertEqual(response.status_code, 200) self.assertContains(response, str(succeeded_multiple.task_id)) self.assertNotContains(response, str(failed_multiple.task_id)) self.assertNotContains(response, str(single.task_id)) self.assertEqual( list(response.context["cl"].queryset.values_list("pk", flat=True)), [succeeded_multiple.pk], ) @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, ) def test_recharge_status_keeps_pending_when_active_query_raises_unexpected_error(self): order = self.create_order(amount="20.00", points_granted=200) self.client.force_login(self.user) with patch("apps.api.views.query_payment_order", side_effect=RuntimeError("gateway down")): 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.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, change_type=PointsLedger.ChangeType.RECHARGE, ).exists() ) class FakeGenerationProvider: def __init__(self, *, capabilities=None): self._capabilities = set(capabilities or {"text", "image", "vision"}) self.text_calls = [] self.image_calls = [] self.vision_calls = [] self.text_error = None self.image_error = None self.vision_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"}, ) def analyze_images(self, prompt, model, **kwargs): self.vision_calls.append({"prompt": prompt, "model": model, **kwargs}) if self.vision_error is not None: raise self.vision_error return TextGenerationResult( text="第一张展示商品正面。\n第二张展示商品细节。", titles=(), model_used=model.model, raw={"secret": "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() reset_keyword_matcher_cache() suffix = uuid.uuid4().hex[:8] self.media_dir = tempfile.TemporaryDirectory() self.addCleanup(self.media_dir.cleanup) self.addCleanup(reset_keyword_matcher_cache) 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.vision_model = self.create_ai_model( name=f"vision-model-{suffix}", model=f"gpt-vision-{suffix}", capabilities=["text", "vision"], ) self.title_alias = f"title-standard-{suffix}" self.image_alias = f"image-hd-{suffix}" self.vision_alias = f"vision-standard-{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.VISION, alias=self.vision_alias, ai_model=self.vision_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, ) PricingRule.objects.create( operation_type=CallRecord.OperationType.VISION, alias=self.vision_alias, points_cost=3, ) 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 register_device_session(self, *, user=None, api_key=None): user = user or self.user api_key = api_key or self.api_key result = register_device( user=user, api_key=api_key, product_code="cmshopee", device_id_version="v1", device_id=f"test-device-{uuid.uuid4().hex}", public_key=f"test-public-key-{uuid.uuid4().hex}", platform="windows", client_version="0.1.0", ) return result.device, result.session_token def post_with_provider(self, path, payload, provider=None, **extra): with patch("apps.api.generation.get_provider", return_value=provider or self.provider): return self.client.post(path, payload, format="json", **self.auth_header(), **extra) def create_cmshopee_entitlement(self, *, starts_at=None): plan = SoftwarePlan.objects.create( product_code=ClientDevice.ProductCode.CMSHOPEE, name="虾皮圈月度订阅", duration_days=30, price=Decimal("19.90"), device_limit=1, grace_days=3, ) return grant_software_entitlement( user=self.user, plan=plan, reason="API 订阅授权测试", starts_at=starts_at, ) def test_cmshopee_subscription_status_is_required_without_entitlement(self): response = self.client.get( "/api/v1/cmshopee/subscription/status", **self.auth_header(), ) self.assertEqual(response.status_code, 200) self.assertEqual( response.json(), { "product_code": "cmshopee", "status": "required", "allowed": False, "code": "subscription_required", "plan": None, }, ) def test_cmshopee_subscription_status_is_active_without_device_session(self): entitlement = self.create_cmshopee_entitlement() response = self.client.get( "/api/v1/cmshopee/subscription/status", **self.auth_header(), ) self.assertEqual(response.status_code, 200) data = response.json() self.assertEqual(data["status"], "active") self.assertTrue(data["allowed"]) self.assertIsNone(data["code"]) self.assertEqual(data["plan"]["name"], entitlement.plan_name) def test_cmshopee_subscription_status_is_expired_after_grace_period(self): self.create_cmshopee_entitlement( starts_at=timezone.now() - timedelta(days=40), ) response = self.client.get( "/api/v1/cmshopee/subscription/status", **self.auth_header(), ) self.assertEqual(response.status_code, 200) data = response.json() self.assertEqual(data["status"], "expired") self.assertFalse(data["allowed"]) self.assertEqual(data["code"], "subscription_expired") @override_settings(CMSHOPEE_SUBSCRIPTION_ENFORCEMENT=True) def test_cmshopee_enforcement_rejects_without_subscription(self): response = self.post_with_provider( "/api/v1/cmshopee/generate/title", {"prompt": "生成一个商品标题", "model": self.title_alias}, ) self.assertEqual(response.status_code, 403) self.assertEqual(response.json()["error"]["code"], "subscription_required") self.assertEqual(self.provider.text_calls, []) self.assert_generation_not_charged() @override_settings(CMSHOPEE_SUBSCRIPTION_ENFORCEMENT=True) def test_cmshopee_account_subscription_allows_multiple_device_contexts(self): self.create_cmshopee_entitlement() first_response = self.post_with_provider( "/api/v1/cmshopee/generate/title", {"prompt": "生成一个商品标题", "model": self.title_alias}, ) second_response = self.post_with_provider( "/api/v1/cmshopee/generate/title", {"prompt": "再生成一个商品标题", "model": self.title_alias}, HTTP_X_DEVICE_SESSION="stale-device-session", ) self.assertEqual(first_response.status_code, 200) self.assertEqual(second_response.status_code, 200) self.assertEqual(len(self.provider.text_calls), 2) @override_settings(CMSHOPEE_SUBSCRIPTION_ENFORCEMENT=True) def test_generic_generation_remains_available_without_subscription(self): response = self.post_with_provider( "/api/v1/generate/title", {"prompt": "通用接口标题", "model": self.title_alias}, ) self.assertEqual(response.status_code, 200) self.assertEqual(len(self.provider.text_calls), 1) def telemetry_event_from_logs(self, captured): events = [ getattr(record, "generation_route_usage", None) for record in captured.records if getattr(record, "generation_route_usage", None) ] self.assertEqual(len(events), 1) return events[0] def assert_generation_telemetry_is_safe(self, event, *, payload=None): self.assertEqual( set(event), { "event", "route_type", "api_key_id", "api_key_prefix", "user_id", "product_code", "client_device_id", "device_session_present", "client_version", "alias", "status", "latency_ms", "error_code", "http_status", }, ) serialized = json.dumps(event, ensure_ascii=False) self.assertNotIn(self.raw_key, serialized) self.assertNotIn("SECRET_RAW", serialized) self.assertNotIn("prompt", serialized) self.assertNotIn("image_base64", serialized) if payload: self.assertNotIn(str(payload.get("prompt") or ""), serialized) encoded_image = str(payload.get("image_base64") or "") if encoded_image: self.assertNotIn(encoded_image, serialized) 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()) @override_settings( MODERATION_ENABLED=True, MODERATION_PROVIDER="keyword", MODERATION_CACHE_VERSION_KEY="test:api:moderation:sensitive_words:version", ) def test_blocked_prompt_returns_content_blocked_before_image_download_or_charge(self): SensitiveWord.objects.create(word="敏感词", category="policy") with ( patch("apps.api.generation.socket.getaddrinfo") as dns_lookup, patch("apps.api.generation.requests.Session.get") as image_get, ): response = self.post_with_provider( "/api/v1/generate/image", { "prompt": "请生成敏-感\u200b 词图片", "model": self.image_alias, "image_url": "https://safe.example.com/input.jpg", "resolution": "1K", "aspect_ratio": "1:1", }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "content_blocked") dns_lookup.assert_not_called() image_get.assert_not_called() self.assertEqual(self.provider.image_calls, []) self.assert_generation_not_charged() @override_settings( MODERATION_ENABLED=False, MODERATION_PROVIDER="keyword", MODERATION_CACHE_VERSION_KEY="test:api:moderation:sensitive_words:version", ) def test_disabled_moderation_does_not_block_matching_prompt(self): SensitiveWord.objects.create(word="敏感词") response = self.post_with_provider( "/api/v1/generate/title", {"prompt": "敏感词", "model": self.title_alias}, ) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["points_cost"], 2) self.assertEqual(len(self.provider.text_calls), 1) 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_generation_endpoints_link_call_records_to_valid_device_session(self): device, session_token = self.register_device_session() device_header = {"HTTP_X_DEVICE_SESSION": session_token} encoded = base64.b64encode(b"device-linked-image").decode("ascii") title = self.post_with_provider( "/api/v1/generate/title", {"prompt": "生成标题", "model": self.title_alias}, **device_header, ) vision = self.post_with_provider( "/api/v1/analyze/images", { "prompt": "理解商品图", "model": self.vision_alias, "images": [{"image_base64": encoded}], }, **device_header, ) image = self.post_with_provider( "/api/v1/generate/image", { "prompt": "生成图片", "model": self.image_alias, "image_base64": encoded, }, **device_header, ) task_submit = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "异步生成图片", "model": self.image_alias}, **device_header, ) for response in (title, vision, image): self.assertEqual(response.status_code, 200) self.assertEqual( CallRecord.objects.get(pk=response.data["call_id"]).client_device_id, device.id, ) self.assertEqual(task_submit.status_code, 202) task = ImageGenerationTask.objects.get(task_id=task_submit.data["task_id"]) self.assertEqual(task.call_record.client_device_id, device.id) poll = self.client.get( f"/api/v1/generate/image/tasks/{task.task_id}", **self.auth_header(), ) self.assertEqual(poll.status_code, 200) def test_generation_without_device_session_remains_legacy_compatible(self): response = self.post_with_provider( "/api/v1/generate/title", {"prompt": "无设备头标题", "model": self.title_alias}, ) self.assertEqual(response.status_code, 200) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertIsNone(call.client_device_id) self.assertEqual(response.data["points_cost"], 2) self.assertEqual(response.data["points_balance"], 98) def test_invalid_or_cross_user_device_session_is_rejected_before_charge(self): invalid = self.post_with_provider( "/api/v1/generate/title", {"prompt": "无效会话", "model": self.title_alias}, HTTP_X_DEVICE_SESSION="dvs_cmhub_invalid", ) self.assertEqual(invalid.status_code, 401) self.assertEqual(invalid.data["error"]["code"], "device_session_invalid") self.assert_generation_not_charged() other_user = get_user_model().objects.create_user( username=f"other-device-user-{uuid.uuid4().hex[:8]}", email=f"other-device-user-{uuid.uuid4().hex[:8]}@example.com", password="password", ) other_key, _raw_other_key = ApiKey.create_for_user(other_user, name="other-device") _device, other_session_token = self.register_device_session( user=other_user, api_key=other_key, ) cross_user = self.post_with_provider( "/api/v1/generate/title", {"prompt": "跨账号会话", "model": self.title_alias}, HTTP_X_DEVICE_SESSION=other_session_token, ) self.assertEqual(cross_user.status_code, 403) self.assertEqual(cross_user.data["error"]["code"], "device_mismatch") self.assert_generation_not_charged() def test_analyze_images_supports_single_image_with_explicit_alias(self): encoded = base64.b64encode(b"single-image").decode("ascii") response = self.post_with_provider( "/api/v1/analyze/images", { "prompt": "描述这张商品图", "model": self.vision_alias, "images": [{"image_base64": encoded}], }, ) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["alias"], self.vision_alias) self.assertEqual(response.data["points_cost"], 3) self.assertEqual(len(self.provider.vision_calls), 1) self.assertEqual( [image.data for image in self.provider.vision_calls[0]["images"]], [b"single-image"], ) def test_analyze_images_supports_ordered_mixed_sources_and_charges_once(self): first = base64.b64encode(b"first-image").decode("ascii") response_from_url = FakeImageUrlResponse( headers={"Content-Type": "image/jpeg"}, chunks=(b"second-", b"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_from_url, ), ): response = self.post_with_provider( "/api/v1/analyze/images", { "prompt": "比较两张商品图", "images": [ {"image_base64": f"data:image/png;base64,{first}"}, {"image_url": "https://images.example.test/detail.jpg"}, ], "parameters": {"temperature": 0.2}, }, ) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["text"], "第一张展示商品正面。\n第二张展示商品细节。") self.assertEqual(response.data["alias"], self.vision_alias) self.assertEqual(response.data["model_used"], self.vision_model.model) self.assertEqual(response.data["points_cost"], 3) self.assertEqual(response.data["points_balance"], 97) self.assertEqual(len(self.provider.vision_calls), 1) images = self.provider.vision_calls[0]["images"] self.assertEqual([image.data for image in images], [b"first-image", b"second-image"]) self.assertEqual( [image.mime_type for image in images], ["image/png", "image/jpeg"], ) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 97) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual(call.operation_type, CallRecord.OperationType.VISION) self.assertEqual(call.status, CallRecord.Status.SUCCESS) self.assertEqual(call.resolution, "") self.assertEqual(call.result_summary, response.data["text"]) self.assertNotIn("first-image", call.result_summary) self.assertNotIn("SECRET_RAW", call.result_summary) self.assertEqual( PointsLedger.objects.filter( ref_call=call, change_type=PointsLedger.ChangeType.CONSUME, ).count(), 1, ) def test_analyze_images_requires_nonempty_exclusive_image_sources(self): empty = self.post_with_provider( "/api/v1/analyze/images", {"prompt": "分析图片", "images": []}, ) both = self.post_with_provider( "/api/v1/analyze/images", { "prompt": "分析图片", "images": [ { "image_url": "https://images.example.test/input.jpg", "image_base64": "aW1hZ2U=", } ], }, ) self.assertEqual(empty.status_code, 400) self.assertEqual(empty.data["error"]["code"], "bad_request") self.assertEqual(both.status_code, 400) self.assertEqual(both.data["error"]["code"], "bad_request") self.assertEqual(self.provider.vision_calls, []) self.assert_generation_not_charged() @override_settings(VISION_MAX_IMAGES=1) def test_analyze_images_rejects_too_many_images_before_charge(self): encoded = base64.b64encode(b"image").decode("ascii") response = self.post_with_provider( "/api/v1/analyze/images", { "prompt": "分析图片", "images": [ {"image_base64": encoded}, {"image_base64": encoded}, ], }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") self.assert_generation_not_charged() @override_settings(VISION_MAX_IMAGE_BYTES=3, VISION_MAX_TOTAL_BYTES=10) def test_analyze_images_rejects_oversized_single_image_before_charge(self): encoded = base64.b64encode(b"four").decode("ascii") response = self.post_with_provider( "/api/v1/analyze/images", {"prompt": "分析图片", "images": [{"image_base64": encoded}]}, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") self.assert_generation_not_charged() @override_settings(VISION_MAX_IMAGE_BYTES=10, VISION_MAX_TOTAL_BYTES=5) def test_analyze_images_rejects_oversized_total_before_charge(self): encoded = base64.b64encode(b"abc").decode("ascii") response = self.post_with_provider( "/api/v1/analyze/images", { "prompt": "分析图片", "images": [ {"image_base64": encoded}, {"image_base64": encoded}, ], }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") self.assert_generation_not_charged() def test_analyze_images_rejects_private_image_url_before_charge(self): with patch( "apps.api.generation.socket.getaddrinfo", return_value=dns_result("127.0.0.1"), ): response = self.post_with_provider( "/api/v1/analyze/images", { "prompt": "分析图片", "images": [{"image_url": "http://internal.example.test/input.jpg"}], }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "bad_request") self.assertEqual(self.provider.vision_calls, []) self.assert_generation_not_charged() def test_analyze_images_requires_api_key(self): encoded = base64.b64encode(b"image").decode("ascii") response = self.client.post( "/api/v1/analyze/images", {"prompt": "分析图片", "images": [{"image_base64": encoded}]}, format="json", ) self.assertEqual(response.status_code, 401) self.assertEqual(response.data["error"]["code"], "unauthorized") self.assert_generation_not_charged() @override_settings( MODERATION_ENABLED=True, MODERATION_PROVIDER="keyword", MODERATION_CACHE_VERSION_KEY="test:api:vision:moderation:version", ) def test_analyze_images_blocks_prompt_before_loading_images_or_charge(self): SensitiveWord.objects.create(word="敏感词", category="policy") with ( patch("apps.api.generation.decode_image_input") as decode_image, patch("apps.api.generation.download_image_input") as download_image, ): response = self.post_with_provider( "/api/v1/analyze/images", { "prompt": "分析敏-感​词图片", "images": [ {"image_url": "https://images.example.test/input.jpg"} ], }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "content_blocked") decode_image.assert_not_called() download_image.assert_not_called() self.assert_generation_not_charged() def test_analyze_images_rejects_model_or_provider_without_text_vision(self): bad_alias = f"vision-without-text-{uuid.uuid4().hex[:8]}" ModelAlias.objects.create( operation_type=ModelAlias.OperationType.VISION, alias=bad_alias, ai_model=self.image_model, ) encoded = base64.b64encode(b"image").decode("ascii") payload = { "prompt": "分析图片", "model": bad_alias, "images": [{"image_base64": encoded}], } model_rejected = self.post_with_provider("/api/v1/analyze/images", payload) provider_rejected = self.post_with_provider( "/api/v1/analyze/images", {**payload, "model": self.vision_alias}, provider=FakeGenerationProvider(capabilities={"vision"}), ) self.assertEqual(model_rejected.status_code, 400) self.assertEqual(model_rejected.data["error"]["code"], "model_not_allowed") self.assertEqual(provider_rejected.status_code, 400) self.assertEqual(provider_rejected.data["error"]["code"], "model_not_allowed") self.assert_generation_not_charged() def test_analyze_images_upstream_failure_refunds_once(self): encoded = base64.b64encode(b"image").decode("ascii") self.provider.vision_error = requests.Timeout("vision timeout") response = self.post_with_provider( "/api/v1/analyze/images", { "prompt": "分析图片", "model": self.vision_alias, "images": [{"image_base64": encoded}], }, ) self.assertEqual(response.status_code, 502) self.assertEqual(response.data["error"]["code"], "upstream_timeout") self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) call = CallRecord.objects.get(operation_type=CallRecord.OperationType.VISION) self.assertEqual(call.status, CallRecord.Status.FAILED) 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_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_accepts_ordered_images_and_injects_role_rules(self): first = base64.b64encode(b"main-image").decode("ascii") second = base64.b64encode(b"reference-image").decode("ascii") response = self.post_with_provider( "/api/v1/generate/image", { "prompt": "生成新的商品主图", "model": self.image_alias, "images": [ {"image_base64": f"data:image/jpeg;base64,{first}"}, {"image_base64": f"data:image/png;base64,{second}"}, ], }, ) self.assertEqual(response.status_code, 200) self.assertEqual(response.data["points_cost"], 10) self.assertEqual(response.data["points_balance"], 90) provider_call = self.provider.image_calls[0] self.assertEqual(provider_call["image"], b"main-image") self.assertEqual( [image.data for image in provider_call["images"]], [b"main-image", b"reference-image"], ) self.assertIn("第 1 张图片是主商品图", provider_call["prompt"]) self.assertIn("第 2 张及之后的图片仅作为", provider_call["prompt"]) self.assertIn("生成新的商品主图", provider_call["prompt"]) def test_generate_image_accepts_mixed_base64_and_url_images_in_order(self): encoded = base64.b64encode(b"main-image").decode("ascii") downloaded = FakeImageUrlResponse( headers={"Content-Type": "image/jpeg"}, chunks=(b"reference-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=downloaded, ), ): response = self.post_with_provider( "/api/v1/generate/image", { "prompt": "生成新的商品主图", "model": self.image_alias, "images": [ {"image_base64": f"data:image/png;base64,{encoded}"}, {"image_url": "https://images.example.test/reference.jpg"}, ], }, ) self.assertEqual(response.status_code, 200) provider_call = self.provider.image_calls[0] self.assertEqual( [image.data for image in provider_call["images"]], [b"main-image", b"reference-image"], ) self.assertEqual( [image.mime_type for image in provider_call["images"]], ["image/png", "image/jpeg"], ) def test_generate_image_rejects_mixed_legacy_and_images_inputs_without_charge(self): encoded = base64.b64encode(b"main-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}", "images": [{"image_base64": f"data:image/png;base64,{encoded}"}], }, ) self.assertEqual(response.status_code, 400) self.assertEqual(self.provider.image_calls, []) self.assert_generation_not_charged() @override_settings(IMAGE_MAX_INPUT_IMAGES=1) def test_generate_image_rejects_too_many_input_images_without_charge(self): encoded = base64.b64encode(b"input-image").decode("ascii") response = self.post_with_provider( "/api/v1/generate/image", { "prompt": "生成图片", "model": self.image_alias, "images": [ {"image_base64": encoded}, {"image_base64": encoded}, ], }, ) self.assertEqual(response.status_code, 400) self.assertEqual(self.provider.image_calls, []) self.assert_generation_not_charged() def test_sync_image_usage_telemetry_logs_safe_client_version_and_key_identity(self): encoded = base64.b64encode(b"input-image").decode("ascii") payload = { "prompt": "生成图片遥测测试", "model": self.image_alias, "image_base64": f"data:image/png;base64,{encoded}", "resolution": "1K", "aspect_ratio": "1:1", } with self.assertLogs("cmhub.api.generation_usage", level="INFO") as captured: response = self.post_with_provider( "/api/v1/generate/image", payload, HTTP_X_CLIENT_VERSION="0.1.1", ) self.assertEqual(response.status_code, 200) event = self.telemetry_event_from_logs(captured) self.assertEqual(event["event"], "generation_route_usage") self.assertEqual(event["route_type"], "sync") self.assertEqual(event["api_key_id"], self.api_key.id) self.assertEqual(event["api_key_prefix"], self.api_key.key_prefix) self.assertEqual(event["user_id"], self.user.id) self.assertEqual(event["product_code"], "") self.assertIsNone(event["client_device_id"]) self.assertFalse(event["device_session_present"]) self.assertEqual(event["client_version"], "0.1.1") self.assertEqual(event["alias"], self.image_alias) self.assertEqual(event["status"], "success") self.assertEqual(event["error_code"], "") self.assertEqual(event["http_status"], 200) self.assertIsInstance(event["latency_ms"], int) self.assertGreaterEqual(event["latency_ms"], 0) self.assert_generation_telemetry_is_safe(event, payload=payload) def test_sync_image_usage_telemetry_records_only_safe_device_metadata(self): device, session_token = self.register_device_session() encoded = base64.b64encode(b"telemetry-device-image").decode("ascii") payload = { "prompt": "设备遥测图片生成", "model": self.image_alias, "image_base64": encoded, } with self.assertLogs("cmhub.api.generation_usage", level="INFO") as captured: response = self.post_with_provider( "/api/v1/generate/image", payload, HTTP_X_DEVICE_SESSION=session_token, HTTP_X_CLIENT_VERSION="0.1.3", ) self.assertEqual(response.status_code, 200) event = self.telemetry_event_from_logs(captured) self.assertEqual(event["product_code"], "cmshopee") self.assertEqual(event["client_device_id"], device.id) self.assertTrue(event["device_session_present"]) self.assertNotIn(session_token, json.dumps(event, ensure_ascii=False)) self.assert_generation_telemetry_is_safe(event, payload=payload) def test_async_image_submit_usage_telemetry_logs_safe_success_event(self): payload = { "prompt": "生成异步图片遥测测试", "model": self.image_alias, "resolution": "1K", } with self.assertLogs("cmhub.api.generation_usage", level="INFO") as captured: response = self.post_with_provider( "/api/v1/generate/image/tasks", payload, HTTP_X_CLIENT_VERSION="0.1.2", ) self.assertEqual(response.status_code, 202) event = self.telemetry_event_from_logs(captured) self.assertEqual(event["route_type"], "async") self.assertEqual(event["api_key_id"], self.api_key.id) self.assertEqual(event["api_key_prefix"], self.api_key.key_prefix) self.assertEqual(event["user_id"], self.user.id) self.assertEqual(event["client_version"], "0.1.2") self.assertEqual(event["alias"], self.image_alias) self.assertEqual(event["status"], "success") self.assertEqual(event["error_code"], "") self.assertEqual(event["http_status"], 202) self.assert_generation_telemetry_is_safe(event, payload=payload) def test_async_image_submit_usage_telemetry_logs_error_code_without_sensitive_data(self): self.wallet.points_balance = 1 self.wallet.save(update_fields=("points_balance", "updated_at")) payload = { "prompt": "余额不足遥测测试", "model": self.image_alias, "resolution": "1K", } with self.assertLogs("cmhub.api.generation_usage", level="INFO") as captured: response = self.post_with_provider( "/api/v1/generate/image/tasks", payload, HTTP_X_CLIENT_VERSION="0.1.3", ) self.assertEqual(response.status_code, 402) event = self.telemetry_event_from_logs(captured) self.assertEqual(event["route_type"], "async") self.assertEqual(event["status"], "error") self.assertEqual(event["error_code"], "insufficient_points") self.assertEqual(event["http_status"], 402) self.assertEqual(event["client_version"], "0.1.3") self.assertEqual(event["alias"], self.image_alias) self.assert_generation_telemetry_is_safe(event, payload=payload) @override_settings( MODERATION_ENABLED=True, MODERATION_PROVIDER="keyword", MODERATION_CACHE_VERSION_KEY="test:api:moderation:sensitive_words:version", ) def test_async_image_blocked_prompt_creates_no_task_or_charge(self): SensitiveWord.objects.create(word="敏感词", category="policy") with ( patch("apps.api.generation.socket.getaddrinfo") as dns_lookup, patch("apps.api.generation.requests.Session.get") as image_get, ): response = self.post_with_provider( "/api/v1/generate/image/tasks", { "prompt": "请生成敏-感\u200b 词图片", "model": self.image_alias, "image_url": "https://safe.example.com/input.jpg", "resolution": "1K", "aspect_ratio": "1:1", }, ) self.assertEqual(response.status_code, 400) self.assertEqual(response.data["error"]["code"], "content_blocked") dns_lookup.assert_not_called() image_get.assert_not_called() self.assertFalse(ImageGenerationTask.objects.exists()) self.assert_generation_not_charged() def test_async_image_insufficient_points_returns_402_without_task(self): self.wallet.points_balance = 1 self.wallet.save(update_fields=("points_balance", "updated_at")) response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) self.assertEqual(response.status_code, 402) self.assertEqual(response.data["error"]["code"], "insufficient_points") self.assertFalse(ImageGenerationTask.objects.exists()) self.assertEqual(self.provider.image_calls, []) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 1) self.assertFalse(CallRecord.objects.filter(user=self.user).exists()) def test_async_image_idempotency_reuses_task_and_rejects_conflict(self): payload = {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"} first = self.post_with_provider( "/api/v1/generate/image/tasks", payload, HTTP_IDEMPOTENCY_KEY="image-job-001", ) second = self.post_with_provider( "/api/v1/generate/image/tasks", payload, HTTP_IDEMPOTENCY_KEY="image-job-001", ) conflict = self.post_with_provider( "/api/v1/generate/image/tasks", {**payload, "prompt": "生成另一张图片"}, HTTP_IDEMPOTENCY_KEY="image-job-001", ) self.assertEqual(first.status_code, 202) self.assertEqual(second.status_code, 202) self.assertEqual(first.data["task_id"], second.data["task_id"]) self.assertEqual(conflict.status_code, 409) self.assertEqual(conflict.data["error"]["code"], "idempotency_conflict") self.assertEqual(ImageGenerationTask.objects.count(), 1) self.assertEqual(CallRecord.objects.filter(user=self.user).count(), 1) self.assertEqual( PointsLedger.objects.filter( user=self.user, change_type=PointsLedger.ChangeType.CONSUME, ).count(), 1, ) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 90) self.assertEqual(self.provider.image_calls, []) @override_settings(MEDIA_PUBLIC_BASE_URL="https://cm.example.test") def test_async_image_worker_success_and_poll_are_idempotent(self): response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) self.assertEqual(response.status_code, 202) self.assertEqual(response.data["status"], ImageGenerationTask.Status.QUEUED) self.assertEqual(response.data["points_balance"], 90) self.assertEqual(response.data["attempt_count"], 0) self.assertEqual(response.data["max_attempts"], 3) self.assertIsNone(response.data["next_attempt_at"]) self.assertEqual(self.provider.image_calls, []) with patch("apps.api.generation.get_provider", return_value=self.provider): claimed = claim_next_image_task("worker-a") self.assertIsNotNone(claimed) task = run_image_generation_task(claimed, worker_id="worker-a") self.assertEqual(task.status, ImageGenerationTask.Status.SUCCEEDED) self.assertTrue(task.result_url.startswith("https://cm.example.test/media/")) self.assertEqual(len(self.provider.image_calls), 1) poll = self.client.get( f"/api/v1/generate/image/tasks/{response.data['task_id']}", **self.auth_header(), ) repeat = self.client.get( f"/api/v1/generate/image/tasks/{response.data['task_id']}", **self.auth_header(), ) self.assertEqual(poll.status_code, 200) self.assertEqual(poll.data["status"], ImageGenerationTask.Status.SUCCEEDED) self.assertEqual(poll.data["attempt_count"], 1) self.assertEqual(poll.data["max_attempts"], 3) self.assertIsNone(poll.data["next_attempt_at"]) self.assertEqual(poll.data["result"]["image_url"], task.result_url) self.assertEqual(repeat.data["result"]["image_url"], task.result_url) @override_settings(MEDIA_PUBLIC_BASE_URL="https://cm.example.test") def test_async_image_task_stores_and_restores_ordered_inputs(self): first = base64.b64encode(b"main-image").decode("ascii") second = base64.b64encode(b"reference-image").decode("ascii") response = self.post_with_provider( "/api/v1/generate/image/tasks", { "prompt": "生成新的商品主图", "model": self.image_alias, "images": [ {"image_base64": f"data:image/jpeg;base64,{first}"}, {"image_base64": f"data:image/png;base64,{second}"}, ], }, ) self.assertEqual(response.status_code, 202) task = ImageGenerationTask.objects.get(task_id=response.data["task_id"]) stored_inputs = list(task.input_images.order_by("ordinal")) self.assertEqual(len(stored_inputs), 2) self.assertEqual([item.ordinal for item in stored_inputs], [0, 1]) self.assertFalse(bool(task.input_image)) self.assertFalse(ImageGenerationTaskInput.objects.filter(task=task, image__isnull=True).exists()) serialized = json.dumps(task.request_payload, ensure_ascii=False) self.assertNotIn(first, serialized) self.assertNotIn(second, serialized) with patch("apps.api.generation.get_provider", return_value=self.provider): claimed = claim_next_image_task("worker-multi") completed = run_image_generation_task(claimed, worker_id="worker-multi") self.assertEqual(completed.status, ImageGenerationTask.Status.SUCCEEDED) provider_call = self.provider.image_calls[0] self.assertEqual( [image.data for image in provider_call["images"]], [b"main-image", b"reference-image"], ) self.assertIn("第 1 张图片是主商品图", provider_call["prompt"]) def test_async_image_poll_rejects_cross_user_access(self): response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) other_user = get_user_model().objects.create_user( username=f"other-{uuid.uuid4().hex[:8]}", email=f"other-{uuid.uuid4().hex[:8]}@example.com", password="password", ) _other_key, other_raw_key = ApiKey.create_for_user(other_user, name="other") denied = self.client.get( f"/api/v1/generate/image/tasks/{response.data['task_id']}", HTTP_AUTHORIZATION=f"Bearer {other_raw_key}", ) self.assertEqual(denied.status_code, 404) self.assertEqual(denied.data["error"]["code"], "task_not_found") @override_settings(IMAGE_TASK_MAX_RETRIES=0) def test_async_image_worker_failure_refunds_precharged_points(self): self.provider.image_error = requests.Timeout("image upstream deadline exceeded") response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) with patch("apps.api.generation.get_provider", return_value=self.provider): task = run_image_generation_task( claim_next_image_task("worker-failure"), worker_id="worker-failure", ) self.assertEqual(task.status, ImageGenerationTask.Status.FAILED) self.assertEqual(task.error_code, "upstream_timeout") self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual(call.status, CallRecord.Status.FAILED) self.assertEqual( PointsLedger.objects.filter( ref_call=call, change_type=PointsLedger.ChangeType.REFUND, ).count(), 1, ) poll = self.client.get( f"/api/v1/generate/image/tasks/{response.data['task_id']}", **self.auth_header(), ) self.assertEqual(poll.data["status"], ImageGenerationTask.Status.FAILED) self.assertEqual(poll.data["error"]["code"], "upstream_timeout") self.assertEqual(poll.data["attempt_count"], 1) self.assertEqual(poll.data["max_attempts"], 1) self.assertIsNone(poll.data["next_attempt_at"]) @override_settings(IMAGE_TASK_RETRY_BACKOFF_SECONDS="60,120") def test_async_image_retryable_timeout_requeues_without_refund_and_respects_backoff(self): self.provider.image_error = requests.Timeout("image upstream deadline exceeded") response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) with patch("apps.api.generation.get_provider", return_value=self.provider): task = run_image_generation_task( claim_next_image_task("worker-retry"), worker_id="worker-retry", ) self.assertEqual(task.status, ImageGenerationTask.Status.QUEUED) self.assertEqual(task.attempt_count, 1) self.assertEqual(task.error_code, "upstream_timeout") self.assertEqual(task.error_message, "上游 AI 调用超时,稍后自动重试") self.assertIsNotNone(task.next_attempt_at) self.assertGreater(task.next_attempt_at, timezone.now()) self.assertIsNone(claim_next_image_task("worker-too-soon")) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 90) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual(call.status, CallRecord.Status.PENDING) self.assertEqual( PointsLedger.objects.filter( ref_call=call, change_type=PointsLedger.ChangeType.REFUND, ).count(), 0, ) poll = self.client.get( f"/api/v1/generate/image/tasks/{response.data['task_id']}", **self.auth_header(), ) self.assertEqual(poll.status_code, 200) self.assertEqual(poll.data["status"], ImageGenerationTask.Status.QUEUED) self.assertEqual(poll.data["attempt_count"], 1) self.assertEqual(poll.data["max_attempts"], 3) self.assertIsNotNone(poll.data["next_attempt_at"]) @override_settings( MEDIA_PUBLIC_BASE_URL="https://cm.example.test", IMAGE_TASK_RETRY_BACKOFF_SECONDS="0,0", ) def test_async_image_retryable_timeouts_then_success_charges_once(self): response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) with patch("apps.api.generation.get_provider", return_value=self.provider): self.provider.image_error = requests.Timeout("first timeout") first = run_image_generation_task( claim_next_image_task("worker-retry-1"), worker_id="worker-retry-1", ) ImageGenerationTask.objects.filter(pk=first.pk).update( next_attempt_at=timezone.now() - timedelta(seconds=1) ) self.provider.image_error = requests.Timeout("second timeout") second = run_image_generation_task( claim_next_image_task("worker-retry-2"), worker_id="worker-retry-2", ) ImageGenerationTask.objects.filter(pk=second.pk).update( next_attempt_at=timezone.now() - timedelta(seconds=1) ) self.provider.image_error = None succeeded = run_image_generation_task( claim_next_image_task("worker-retry-3"), worker_id="worker-retry-3", ) self.assertEqual(first.status, ImageGenerationTask.Status.QUEUED) self.assertEqual(second.status, ImageGenerationTask.Status.QUEUED) self.assertEqual(succeeded.status, ImageGenerationTask.Status.SUCCEEDED) self.assertEqual(succeeded.attempt_count, 3) self.assertTrue(succeeded.result_url.startswith("https://cm.example.test/media/")) self.assertEqual(len(self.provider.image_calls), 3) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 90) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual(call.status, CallRecord.Status.SUCCESS) 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(), 0, ) @override_settings(IMAGE_TASK_RETRY_BACKOFF_SECONDS="0,0") def test_async_image_retryable_timeouts_final_failure_refunds_once(self): self.provider.image_error = requests.Timeout("image upstream deadline exceeded") response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) with patch("apps.api.generation.get_provider", return_value=self.provider): first = run_image_generation_task( claim_next_image_task("worker-final-1"), worker_id="worker-final-1", ) ImageGenerationTask.objects.filter(pk=first.pk).update( next_attempt_at=timezone.now() - timedelta(seconds=1) ) second = run_image_generation_task( claim_next_image_task("worker-final-2"), worker_id="worker-final-2", ) ImageGenerationTask.objects.filter(pk=second.pk).update( next_attempt_at=timezone.now() - timedelta(seconds=1) ) failed = run_image_generation_task( claim_next_image_task("worker-final-3"), worker_id="worker-final-3", ) self.assertEqual(failed.status, ImageGenerationTask.Status.FAILED) self.assertEqual(failed.error_code, "upstream_timeout") self.assertEqual(failed.attempt_count, 3) self.assertIsNone(failed.next_attempt_at) self.assertEqual(len(self.provider.image_calls), 3) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual(call.status, CallRecord.Status.FAILED) 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_async_image_non_retryable_provider_error_fails_immediately_and_refunds(self): self.provider.image_error = AiCapabilityError("input image is required") response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) with patch("apps.api.generation.get_provider", return_value=self.provider): failed = run_image_generation_task( claim_next_image_task("worker-no-retry"), worker_id="worker-no-retry", ) self.assertEqual(failed.status, ImageGenerationTask.Status.FAILED) self.assertEqual(failed.error_code, "bad_request") self.assertEqual(failed.attempt_count, 1) self.assertIsNone(failed.next_attempt_at) self.assertEqual(len(self.provider.image_calls), 1) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual(call.status, CallRecord.Status.FAILED) self.assertEqual( PointsLedger.objects.filter( ref_call=call, change_type=PointsLedger.ChangeType.REFUND, ).count(), 1, ) @override_settings(IMAGE_TASK_MAX_RETRIES=0) def test_run_image_tasks_logs_failed_task_alias_error_and_duration(self): self.provider.image_error = requests.Timeout("image upstream deadline exceeded") response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) out = io.StringIO() with patch("apps.api.generation.get_provider", return_value=self.provider): call_command( "run_image_tasks", "--once", "--worker-id", "worker-log", stdout=out, ) task = ImageGenerationTask.objects.get(task_id=response.data["task_id"]) output = out.getvalue() self.assertEqual(task.status, ImageGenerationTask.Status.FAILED) self.assertIn("event=image_task_processed", output) self.assertIn(f"task_id={task.task_id}", output) self.assertIn(f"alias={self.image_alias}", output) self.assertIn("status=failed", output) self.assertIn("attempt=1", output) self.assertIn("max_attempts=1", output) self.assertIn("retrying=false", output) self.assertIn("error_code=upstream_timeout", output) self.assertRegex(output, r"duration_ms=\d+") self.assertNotIn("生成图片", output) @override_settings(IMAGE_TASK_RETRY_BACKOFF_SECONDS="60,120") def test_run_image_tasks_logs_retrying_task_attempt_fields_without_sensitive_data(self): self.provider.image_error = requests.Timeout("image upstream deadline exceeded") response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) out = io.StringIO() with patch("apps.api.generation.get_provider", return_value=self.provider): call_command( "run_image_tasks", "--once", "--worker-id", "worker-log-retry", stdout=out, ) task = ImageGenerationTask.objects.get(task_id=response.data["task_id"]) output = out.getvalue() self.assertEqual(task.status, ImageGenerationTask.Status.QUEUED) self.assertIn("event=image_task_processed", output) self.assertIn(f"task_id={task.task_id}", output) self.assertIn(f"alias={self.image_alias}", output) self.assertIn("status=queued", output) self.assertIn("attempt=1", output) self.assertIn("max_attempts=3", output) self.assertIn("retrying=true", output) self.assertIn("next_attempt_at=", output) self.assertIn("error_code=upstream_timeout", output) self.assertRegex(output, r"duration_ms=\d+") self.assertNotIn("生成图片", output) self.assertNotIn(self.raw_key, output) def test_async_image_reaper_fails_stale_running_task_and_refunds(self): response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) claimed = claim_next_image_task("worker-crash") stale_at = timezone.now() - timedelta(seconds=5) ImageGenerationTask.objects.filter(pk=claimed.pk).update( lease_expires_at=stale_at, heartbeat_at=stale_at, ) reaped = reap_stale_image_tasks(now=timezone.now()) task = ImageGenerationTask.objects.get(pk=claimed.pk) self.assertEqual(reaped, 1) self.assertEqual(task.status, ImageGenerationTask.Status.FAILED) self.assertEqual(task.error_code, "task_timeout") self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual( PointsLedger.objects.filter( ref_call=call, change_type=PointsLedger.ChangeType.REFUND, ).count(), 1, ) @override_settings(MEDIA_PUBLIC_BASE_URL="https://cm.example.test") def test_async_image_duplicate_worker_does_not_double_charge_or_refund(self): response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) with patch("apps.api.generation.get_provider", return_value=self.provider): task = run_image_generation_task( claim_next_image_task("worker-a"), worker_id="worker-a", ) duplicate = run_image_generation_task(task, worker_id="worker-b") self.assertEqual(duplicate.status, ImageGenerationTask.Status.SUCCEEDED) self.assertEqual(duplicate.result_url, task.result_url) self.assertEqual(len(self.provider.image_calls), 1) call = CallRecord.objects.get(pk=response.data["call_id"]) 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(), 0, ) def test_async_image_late_worker_after_reaper_cannot_flip_failed_task(self): response = self.post_with_provider( "/api/v1/generate/image/tasks", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) claimed = claim_next_image_task("worker-late") stale_at = timezone.now() - timedelta(seconds=5) ImageGenerationTask.objects.filter(pk=claimed.pk).update( lease_expires_at=stale_at, heartbeat_at=stale_at, ) reap_stale_image_tasks(now=timezone.now()) with patch("apps.api.generation.get_provider", return_value=self.provider): late = run_image_generation_task(claimed, worker_id="worker-late") self.assertEqual(late.status, ImageGenerationTask.Status.FAILED) self.assertEqual(late.result_url, "") self.assertEqual(self.provider.image_calls, []) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) call = CallRecord.objects.get(pk=response.data["call_id"]) self.assertEqual(call.status, CallRecord.Status.FAILED) self.assertEqual( PointsLedger.objects.filter( ref_call=call, change_type=PointsLedger.ChangeType.REFUND, ).count(), 1, ) def test_generation_core_saves_image_with_url_builder_without_request(self): encoded = base64.b64encode(b"input-image").decode("ascii") with patch("apps.api.generation.get_provider", return_value=self.provider): result = run_synchronous_generation( GenerationInput( user=self.user, api_key=self.api_key, operation_type=CallRecord.OperationType.IMAGE, prompt="生成图片", alias=self.image_alias, resolution="1K", image_base64=f"data:image/png;base64,{encoded}", ), image_url_builder=lambda url: f"https://cdn.example.test{url}", ) self.assertEqual(result.operation_type, CallRecord.OperationType.IMAGE) self.assertTrue(result.image_url.startswith("https://cdn.example.test/media/")) self.assertEqual(result.as_response_data()["image_url"], result.image_url) self.assertEqual(self.provider.image_calls[0]["image"], b"input-image") call = CallRecord.objects.get(pk=result.call_record.id) self.assertEqual(call.status, CallRecord.Status.SUCCESS) self.assertEqual(call.result_ref, result.image_url) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 90) 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_precharged_generation_stage_refunds_on_upstream_failure(self): self.provider.text_error = AiProviderError("provider timeout") with patch("apps.api.generation.get_provider", return_value=self.provider): prepared = prepare_generation( GenerationInput( user=self.user, api_key=self.api_key, operation_type=CallRecord.OperationType.TITLE, prompt="生成标题", alias=self.title_alias, resolution="1K", ) ) precharged = precharge_generation(prepared) self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 98) self.assertEqual(precharged.call_record.status, CallRecord.Status.PENDING) with self.assertRaises(ApiRequestError) as captured: execute_precharged_generation(precharged) self.assertEqual(captured.exception.code, "upstream_error") self.wallet.refresh_from_db() self.assertEqual(self.wallet.points_balance, 100) call = CallRecord.objects.get(pk=precharged.call_record.id) 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_image_upstream_timeout_refunds_precharged_points_and_marks_call_failed(self): self.provider.image_error = requests.Timeout("image upstream deadline exceeded") response = self.post_with_provider( "/api/v1/generate/image", {"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"}, ) self.assertEqual(response.status_code, 502) self.assertEqual(response.data["error"]["code"], "upstream_timeout") 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("image upstream deadline exceeded", 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, ) @override_settings( PAYMENT_CALLBACK_MODE="mock", PAYMENT_MOCK_CALLBACK_SECRET="test-payment-callback-secret", ) class SoftwareOrderCallbackApiTests(TestCase): callback_url = "/api/v1/software-orders/callback/wechat" def setUp(self): self.user = get_user_model().objects.create_user( username="software-callback-user", email="software-callback@example.com", password="test-password", ) self.plan = SoftwarePlan.objects.create( product_code=ClientDevice.ProductCode.CMSHOPEE, name="软件月度套餐", duration_days=30, price=Decimal("19.90"), device_limit=1, ) self.order = create_software_order( user=self.user, plan=self.plan, pay_method=SoftwareOrder.PayMethod.WEIXIN, payment_order_func=lambda _order: PaymentOrderCode( code_url="weixin://software-order-test", expires_at=timezone.now() + timedelta(minutes=10), ), ) def signed_body(self, *, amount_cents=1990, transaction_id="wx-software-callback-001"): payload = { "event_type": "TRANSACTION.SUCCESS", "resource": { "trade_state": "SUCCESS", "out_trade_no": self.order.order_no, "transaction_id": transaction_id, "success_time": "2026-07-21T12:00:00+08:00", "amount": {"total": amount_cents}, }, } body = json.dumps(payload, separators=(",", ":")).encode("utf-8") return body, build_mock_body_signature(body) def test_callback_fulfills_once_without_writing_points_ledger(self): body, signature = self.signed_body() first = self.client.post( self.callback_url, data=body, content_type="application/json", HTTP_WECHATPAY_SIGNATURE=signature, ) second = self.client.post( self.callback_url, data=body, content_type="application/json", HTTP_WECHATPAY_SIGNATURE=signature, ) self.assertEqual(first.status_code, 200) self.assertEqual(second.status_code, 200) self.order.refresh_from_db() self.assertEqual(self.order.status, SoftwareOrder.Status.PAID) self.assertIsNotNone(self.order.entitlement_id) self.assertEqual(PointsLedger.objects.filter(user=self.user).count(), 0) def test_callback_rejects_amount_mismatch_without_fulfilling(self): body, signature = self.signed_body(amount_cents=1989) response = self.client.post( self.callback_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.order.refresh_from_db() self.assertEqual(self.order.status, SoftwareOrder.Status.PENDING) self.assertEqual(PointsLedger.objects.filter(user=self.user).count(), 0)