Files
cmhub/apps/api/tests.py
T

3399 lines
135 KiB
Python
Raw Normal View History

2026-07-02 17:40:26 +08:00
import uuid
2026-07-02 22:41:37 +08:00
import base64
2026-07-09 11:57:06 +08:00
import io
2026-07-03 09:07:21 +08:00
import json
2026-07-02 22:41:37 +08:00
import tempfile
2026-07-08 22:08:48 +08:00
from datetime import timedelta
2026-07-03 09:07:21 +08:00
from decimal import Decimal
2026-07-02 22:41:37 +08:00
from pathlib import Path
from unittest.mock import patch
2026-07-02 09:07:15 +08:00
2026-07-03 10:34:37 +08:00
import requests
2026-07-02 22:41:37 +08:00
from cryptography.fernet import Fernet
2026-07-07 08:33:49 +08:00
from django.conf import settings
2026-07-08 15:38:59 +08:00
from django.contrib import admin
2026-07-02 17:40:26 +08:00
from django.contrib.auth import get_user_model
2026-07-03 10:34:37 +08:00
from django.core.cache import cache
2026-07-20 08:58:34 +08:00
from django.core.files.base import ContentFile
2026-07-09 11:57:06 +08:00
from django.core.management import call_command
2026-07-02 17:40:26 +08:00
from django.test import TestCase, override_settings
2026-07-20 08:58:34 +08:00
from django.urls import path, reverse
2026-07-03 09:34:07 +08:00
from django.utils import timezone
2026-07-02 17:40:26 +08:00
from rest_framework.response import Response
from rest_framework.test import APIClient
2026-07-03 10:34:37 +08:00
from rest_framework.views import APIView
2026-07-02 17:40:26 +08:00
from apps.api.authentication import ApiKeyAuthentication
2026-07-08 21:24:59 +08:00
from apps.api.generation import (
ApiRequestError,
GenerationInput,
execute_precharged_generation,
precharge_generation,
prepare_generation,
run_synchronous_generation,
)
2026-07-08 22:08:48 +08:00
from apps.api.image_tasks import (
claim_next_image_task,
reap_stale_image_tasks,
run_image_generation_task,
)
2026-07-17 15:55:38 +08:00
from apps.api.models import ImageGenerationTask, ImageGenerationTaskInput
2026-07-04 10:14:16 +08:00
from apps.api.throttles import GenerateRateThrottle
2026-07-07 08:33:49 +08:00
from apps.api.views import ClientLatestReleaseView, ExternalApiView, ModelsView
2026-07-02 22:41:37 +08:00
from apps.ai.models import AiModel, ModelAlias
from apps.ai.providers import (
AiCapabilityError,
AiProviderError,
ImageGenerationResult,
TextGenerationResult,
)
2026-07-03 09:34:07 +08:00
from apps.billing.models import (
CallRecord,
ExchangeRate,
PointsLedger,
PricingRule,
RechargeOrder,
)
2026-07-03 09:07:21 +08:00
from apps.billing.payment_gateways import (
build_mock_alipay_signature,
build_mock_body_signature,
2026-07-21 11:52:49 +08:00
PaymentOrderCode,
2026-07-03 09:07:21 +08:00
)
2026-07-03 09:34:07 +08:00
from apps.billing.services import RechargePayment
from apps.moderation.models import SensitiveWord
from apps.moderation.providers.keyword import reset_keyword_matcher_cache
2026-07-07 08:33:49 +08:00
from apps.portal.models import DownloadRelease
2026-07-22 15:13:28 +08:00
from apps.licensing.models import (
ClientDevice,
SoftwareEntitlement,
SoftwareOrder,
SoftwarePlan,
)
from apps.licensing.services import (
create_software_order,
grant_software_entitlement,
register_device,
)
2026-07-02 17:40:26 +08:00
from apps.users.models import ApiKey
2026-07-02 22:41:37 +08:00
from apps.users.models import UserWallet
TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii")
2026-07-02 17:40:26 +08:00
class AuthenticatedEchoView(ExternalApiView):
def get(self, request):
return Response(
{
"user_id": request.user.id,
"api_key_id": request.auth.id,
}
)
2026-07-03 10:34:37 +08:00
class DefaultAuthProbeView(APIView):
def get(self, request):
return Response({"ok": True})
2026-07-02 17:40:26 +08:00
urlpatterns = [
path("api/test-auth/", AuthenticatedEchoView.as_view()),
2026-07-03 10:34:37 +08:00
path("api/default-auth/", DefaultAuthProbeView.as_view()),
2026-07-02 17:40:26 +08:00
]
@override_settings(ROOT_URLCONF=__name__)
class ApiKeyAuthenticationTests(TestCase):
url = "/api/test-auth/"
def setUp(self):
2026-07-03 10:34:37 +08:00
cache.clear()
2026-07-02 17:40:26 +08:00
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,))
2026-07-03 10:34:37 +08:00
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)
2026-07-02 17:40:26 +08:00
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")
2026-07-03 10:34:37 +08:00
@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")
2026-07-02 17:40:26 +08:00
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")
2026-07-02 22:41:37 +08:00
2026-07-03 08:36:30 +08:00
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="主账号",
2026-07-03 08:36:30 +08:00
)
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"])
2026-07-03 08:36:30 +08:00
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)
2026-07-03 08:36:30 +08:00
self.assertEqual(response.data["points_balance"], 0)
self.assertEqual(response.data["account"]["username"], self.user.username)
self.assertEqual(response.data["account"]["display_name"], "主账号")
2026-07-03 08:36:30 +08:00
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")
2026-07-04 10:14:16 +08:00
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",
)
2026-07-16 14:13:19 +08:00
vision_alias = self.create_alias(
alias="vision-standard",
operation_type=ModelAlias.OperationType.VISION,
capabilities=["text", "vision"],
model_sku="secret-sku-vision",
)
2026-07-04 10:14:16 +08:00
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,
)
2026-07-16 14:13:19 +08:00
PricingRule.objects.create(
operation_type=vision_alias.operation_type,
alias=vision_alias.alias,
resolution="",
points_cost=3,
)
2026-07-04 10:14:16 +08:00
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"]}
2026-07-16 14:13:19 +08:00
self.assertEqual(set(models), {"title-standard", "image-edit", "vision-standard"})
2026-07-04 10:14:16 +08:00
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},
],
)
2026-07-16 14:13:19 +08:00
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}],
)
2026-07-04 10:14:16 +08:00
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"], [])
2026-07-07 08:33:49 +08:00
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 客户端",
2026-07-08 15:38:59 +08:00
force_update: bool = False,
2026-07-13 17:04:59 +08:00
size_bytes: int | None = None,
2026-07-07 08:33:49 +08:00
) -> 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,
2026-07-08 15:38:59 +08:00
force_update=force_update,
2026-07-13 17:04:59 +08:00
size_bytes=size_bytes,
2026-07-07 08:33:49 +08:00
)
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",
2026-07-13 17:04:59 +08:00
size_bytes=18_765_432,
2026-07-07 08:33:49 +08:00
)
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"]),
2026-07-08 15:38:59 +08:00
{
"version",
"download_url",
"sha256",
"release_notes",
"force_update",
2026-07-13 17:04:59 +08:00
"size_bytes",
2026-07-08 15:38:59 +08:00
"published_at",
},
2026-07-07 08:33:49 +08:00
)
self.assertEqual(response.data["release"]["version"], "1.2.3")
2026-07-08 15:38:59 +08:00
self.assertFalse(response.data["release"]["force_update"])
2026-07-13 17:04:59 +08:00
self.assertEqual(response.data["release"]["size_bytes"], 18_765_432)
2026-07-07 08:33:49 +08:00
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(),
)
2026-07-08 15:38:59 +08:00
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"])
2026-07-13 17:04:59 +08:00
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"])
2026-07-07 08:33:49 +08:00
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"], "暂未发布")
2026-07-08 15:38:59 +08:00
self.assertNotIn("force_update", response.data)
2026-07-13 17:04:59 +08:00
self.assertNotIn("size_bytes", response.data)
2026-07-07 08:33:49 +08:00
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"]),
2026-07-08 15:38:59 +08:00
{
"version",
"download_url",
"sha256",
"release_notes",
"force_update",
2026-07-13 17:04:59 +08:00
"size_bytes",
2026-07-08 15:38:59 +08:00
"published_at",
},
2026-07-07 08:33:49 +08:00
)
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)
2026-07-13 17:04:59 +08:00
def test_download_release_admin_exposes_release_metadata_fields(self):
2026-07-08 15:38:59 +08:00
registered_admin = admin.site._registry[DownloadRelease]
self.assertIn("force_update", registered_admin.list_display)
self.assertIn("force_update", registered_admin.list_filter)
2026-07-13 17:04:59 +08:00
self.assertIn("size_bytes", registered_admin.list_display)
2026-07-08 15:38:59 +08:00
version_fields = registered_admin.fieldsets[0][1]["fields"]
self.assertIn("force_update", version_fields)
2026-07-13 17:04:59 +08:00
download_fields = registered_admin.fieldsets[1][1]["fields"]
self.assertIn("size_bytes", download_fields)
2026-07-08 15:38:59 +08:00
2026-07-07 08:33:49 +08:00
2026-07-20 08:58:34 +08:00
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)
2026-07-20 09:03:23 +08:00
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],
)
2026-07-20 08:58:34 +08:00
2026-07-03 09:07:21 +08:00
@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())
2026-07-03 09:34:07 +08:00
@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())
2026-07-03 10:34:37 +08:00
@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())
2026-07-03 09:34:07 +08:00
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,
)
2026-07-06 08:56:30 +08:00
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()
)
2026-07-03 09:34:07 +08:00
2026-07-02 22:41:37 +08:00
class FakeGenerationProvider:
def __init__(self, *, capabilities=None):
self._capabilities = set(capabilities or {"text", "image", "vision"})
self.text_calls = []
self.image_calls = []
2026-07-16 14:13:19 +08:00
self.vision_calls = []
2026-07-02 22:41:37 +08:00
self.text_error = None
self.image_error = None
2026-07-16 14:13:19 +08:00
self.vision_error = None
2026-07-02 22:41:37 +08:00
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"},
)
2026-07-16 14:13:19 +08:00
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"},
)
2026-07-02 22:41:37 +08:00
2026-07-03 10:34:37 +08:00
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))]
2026-07-02 22:41:37 +08:00
@override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY)
class GenerateApiTests(TestCase):
def setUp(self):
2026-07-03 10:34:37 +08:00
cache.clear()
reset_keyword_matcher_cache()
2026-07-02 22:41:37 +08:00
suffix = uuid.uuid4().hex[:8]
self.media_dir = tempfile.TemporaryDirectory()
self.addCleanup(self.media_dir.cleanup)
self.addCleanup(reset_keyword_matcher_cache)
2026-07-02 22:41:37 +08:00
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"],
)
2026-07-16 14:13:19 +08:00
self.vision_model = self.create_ai_model(
name=f"vision-model-{suffix}",
model=f"gpt-vision-{suffix}",
capabilities=["text", "vision"],
)
2026-07-02 22:41:37 +08:00
self.title_alias = f"title-standard-{suffix}"
self.image_alias = f"image-hd-{suffix}"
2026-07-16 14:13:19 +08:00
self.vision_alias = f"vision-standard-{suffix}"
2026-07-02 22:41:37 +08:00
ModelAlias.objects.create(
operation_type=ModelAlias.OperationType.TITLE,
alias=self.title_alias,
ai_model=self.title_model,
is_default=True,
)
2026-07-16 14:13:19 +08:00
ModelAlias.objects.create(
operation_type=ModelAlias.OperationType.VISION,
alias=self.vision_alias,
ai_model=self.vision_model,
is_default=True,
)
2026-07-02 22:41:37 +08:00
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,
)
2026-07-16 14:13:19 +08:00
PricingRule.objects.create(
operation_type=CallRecord.OperationType.VISION,
alias=self.vision_alias,
points_cost=3,
)
2026-07-02 22:41:37 +08:00
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
2026-07-08 22:08:48 +08:00
def post_with_provider(self, path, payload, provider=None, **extra):
2026-07-02 22:41:37 +08:00
with patch("apps.api.generation.get_provider", return_value=provider or self.provider):
2026-07-08 22:08:48 +08:00
return self.client.post(path, payload, format="json", **self.auth_header(), **extra)
2026-07-02 22:41:37 +08:00
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,
)
2026-07-22 15:13:28 +08:00
@override_settings(
CMSHOPEE_SUBSCRIPTION_MODE="open",
PUBLIC_BASE_URL="https://cm.example.test",
)
def test_cmshopee_open_mode_returns_compatible_access_without_writing_entitlement(self):
response = self.client.get(
"/api/v1/cmshopee/subscription/status",
**self.auth_header(),
)
self.assertEqual(response.status_code, 200)
data = response.json()
self.assertEqual(data["product_code"], "cmshopee")
self.assertEqual(data["status"], "active")
self.assertTrue(data["allowed"])
self.assertIsNone(data["code"])
self.assertEqual(data["account"]["username"], self.user.username)
self.assertEqual(data["plan"]["code"], "development-open")
self.assertEqual(data["plan"]["display_name"], "开发测试长期会员")
self.assertEqual(data["expires_at"], data["plan"]["expires_at"])
self.assertIsNone(data["grace_expires_at"])
self.assertEqual(data["manage_url"], "https://cm.example.test/subscription")
self.assertEqual(data["access_source"], "open_mode")
self.assertEqual(data["entitlement_status"], "required")
self.assertTrue(data["notice_id"])
self.assertFalse(SoftwareEntitlement.objects.filter(user=self.user).exists())
@override_settings(CMSHOPEE_SUBSCRIPTION_MODE="shadow")
def test_cmshopee_shadow_mode_allows_missing_entitlement_and_reports_real_status(self):
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.assertEqual(data["access_source"], "shadow_fallback")
self.assertEqual(data["entitlement_status"], "required")
self.assertEqual(data["plan"]["code"], "shadow-fallback")
@override_settings(CMSHOPEE_SUBSCRIPTION_MODE="open")
def test_cmshopee_open_mode_submit_does_not_create_entitlement(self):
response = self.post_with_provider(
"/api/v1/cmshopee/generate/title",
{"prompt": "开发测试标题", "model": self.title_alias},
)
self.assertEqual(response.status_code, 200)
self.assertEqual(len(self.provider.text_calls), 1)
self.assertFalse(SoftwareEntitlement.objects.filter(user=self.user).exists())
@override_settings(CMSHOPEE_SUBSCRIPTION_MODE="shadow")
def test_cmshopee_shadow_submit_logs_real_entitlement_status(self):
with self.assertLogs("cmhub.licensing.authorization", level="INFO") as captured:
response = self.post_with_provider(
"/api/v1/cmshopee/generate/title",
{"prompt": "影子模式标题", "model": self.title_alias},
)
self.assertEqual(response.status_code, 200)
events = [
getattr(record, "cmshopee_authorization", None)
for record in captured.records
if getattr(record, "cmshopee_authorization", None)
]
self.assertEqual(len(events), 1)
self.assertEqual(events[0]["subscription_mode"], "shadow")
self.assertEqual(events[0]["access_source"], "shadow_fallback")
self.assertEqual(events[0]["entitlement_status"], "required")
self.assertTrue(events[0]["would_reject"])
self.assertEqual(events[0]["would_reject_code"], "subscription_required")
self.assertTrue(events[0]["final_allowed"])
@override_settings(CMSHOPEE_SUBSCRIPTION_MODE="enforce")
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)
2026-07-22 15:13:28 +08:00
data = response.json()
self.assertEqual(data["status"], "required")
self.assertFalse(data["allowed"])
self.assertEqual(data["code"], "subscription_required")
self.assertIsNone(data["plan"])
self.assertEqual(data["entitlement_status"], "required")
self.assertEqual(data["access_source"], "entitlement")
@override_settings(
CMSHOPEE_SUBSCRIPTION_MODE="",
CMSHOPEE_SUBSCRIPTION_ENFORCEMENT=True,
)
def test_cmshopee_legacy_enforcement_setting_remains_supported(self):
response = self.client.get(
"/api/v1/cmshopee/subscription/status",
**self.auth_header(),
)
2026-07-22 15:13:28 +08:00
self.assertEqual(response.status_code, 200)
self.assertFalse(response.json()["allowed"])
self.assertEqual(response.json()["code"], "subscription_required")
@override_settings(CMSHOPEE_SUBSCRIPTION_MODE="enforce")
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"])
2026-07-22 15:13:28 +08:00
self.assertEqual(data["plan"]["display_name"], entitlement.plan_name)
self.assertEqual(data["plan"]["name"], entitlement.plan_name)
2026-07-22 15:13:28 +08:00
self.assertEqual(data["plan"]["code"], f"plan-{entitlement.source_plan_id}")
self.assertEqual(data["expires_at"], entitlement.expires_at.isoformat())
self.assertEqual(
data["grace_expires_at"],
entitlement.grace_expires_at.isoformat(),
)
self.assertEqual(data["access_source"], "entitlement")
self.assertEqual(data["entitlement_status"], "active")
2026-07-22 15:13:28 +08:00
@override_settings(CMSHOPEE_SUBSCRIPTION_MODE="enforce")
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")
2026-07-22 15:13:28 +08:00
self.assertEqual(data["entitlement_status"], "expired")
2026-07-22 15:13:28 +08:00
@override_settings(CMSHOPEE_SUBSCRIPTION_MODE="enforce")
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()
2026-07-22 15:13:28 +08:00
@override_settings(CMSHOPEE_SUBSCRIPTION_MODE="enforce")
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)
2026-07-22 15:13:28 +08:00
@override_settings(CMSHOPEE_SUBSCRIPTION_MODE="enforce")
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)
2026-07-08 22:44:17 +08:00
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",
2026-07-08 22:44:17 +08:00
"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)
2026-07-03 10:34:37 +08:00
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)
2026-07-02 22:41:37 +08:00
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()
2026-07-16 14:13:19 +08:00
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,
)
2026-07-02 22:41:37 +08:00
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)
2026-07-17 15:55:38 +08:00
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()
2026-07-08 22:44:17 +08:00
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"])
2026-07-08 22:44:17 +08:00
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)
2026-07-08 22:44:17 +08:00
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)
2026-07-08 22:08:48 +08:00
@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)
2026-07-09 14:46:09 +08:00
self.assertEqual(response.data["attempt_count"], 0)
self.assertEqual(response.data["max_attempts"], 3)
self.assertIsNone(response.data["next_attempt_at"])
2026-07-08 22:08:48 +08:00
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)
2026-07-09 14:46:09 +08:00
self.assertEqual(poll.data["attempt_count"], 1)
self.assertEqual(poll.data["max_attempts"], 3)
self.assertIsNone(poll.data["next_attempt_at"])
2026-07-08 22:08:48 +08:00
self.assertEqual(poll.data["result"]["image_url"], task.result_url)
self.assertEqual(repeat.data["result"]["image_url"], task.result_url)
2026-07-17 15:55:38 +08:00
@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"])
2026-07-08 22:08:48 +08:00
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")
2026-07-09 14:46:09 +08:00
@override_settings(IMAGE_TASK_MAX_RETRIES=0)
2026-07-08 22:08:48 +08:00
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")
2026-07-09 14:46:09 +08:00
self.assertEqual(poll.data["attempt_count"], 1)
self.assertEqual(poll.data["max_attempts"], 1)
self.assertIsNone(poll.data["next_attempt_at"])
2026-07-08 22:08:48 +08:00
2026-07-09 14:46:09 +08:00
@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)
2026-07-09 11:57:06 +08:00
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)
2026-07-09 14:46:09 +08:00
self.assertIn("attempt=1", output)
self.assertIn("max_attempts=1", output)
self.assertIn("retrying=false", output)
2026-07-09 11:57:06 +08:00
self.assertIn("error_code=upstream_timeout", output)
self.assertRegex(output, r"duration_ms=\d+")
self.assertNotIn("生成图片", output)
2026-07-09 14:46:09 +08:00
@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)
2026-07-08 22:08:48 +08:00
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,
)
2026-07-08 21:24:59 +08:00
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)
2026-07-03 10:34:37 +08:00
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)
2026-07-02 22:41:37 +08:00
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(
2026-07-08 20:14:33 +08:00
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)
2026-07-08 21:24:59 +08:00
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,
)
2026-07-08 20:14:33 +08:00
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(
2026-07-02 22:41:37 +08:00
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,
)
2026-07-21 11:52:49 +08:00
@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)