1488 lines
57 KiB
Python
1488 lines
57 KiB
Python
import uuid
|
|
import base64
|
|
import json
|
|
import tempfile
|
|
from decimal import Decimal
|
|
from pathlib import Path
|
|
from unittest.mock import patch
|
|
|
|
import requests
|
|
from cryptography.fernet import Fernet
|
|
from django.conf import settings
|
|
from django.contrib.auth import get_user_model
|
|
from django.core.cache import cache
|
|
from django.test import TestCase, override_settings
|
|
from django.urls import path
|
|
from django.utils import timezone
|
|
from rest_framework.response import Response
|
|
from rest_framework.test import APIClient
|
|
from rest_framework.views import APIView
|
|
|
|
from apps.api.authentication import ApiKeyAuthentication
|
|
from apps.api.throttles import GenerateRateThrottle
|
|
from apps.api.views import ClientLatestReleaseView, ExternalApiView, ModelsView
|
|
from apps.ai.models import AiModel, ModelAlias
|
|
from apps.ai.providers import (
|
|
AiCapabilityError,
|
|
AiProviderError,
|
|
ImageGenerationResult,
|
|
TextGenerationResult,
|
|
)
|
|
from apps.billing.models import (
|
|
CallRecord,
|
|
ExchangeRate,
|
|
PointsLedger,
|
|
PricingRule,
|
|
RechargeOrder,
|
|
)
|
|
from apps.billing.payment_gateways import (
|
|
build_mock_alipay_signature,
|
|
build_mock_body_signature,
|
|
)
|
|
from apps.billing.services import RechargePayment
|
|
from apps.moderation.models import SensitiveWord
|
|
from apps.moderation.providers.keyword import reset_keyword_matcher_cache
|
|
from apps.portal.models import DownloadRelease
|
|
from apps.users.models import ApiKey
|
|
from apps.users.models import UserWallet
|
|
|
|
|
|
TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii")
|
|
|
|
|
|
class AuthenticatedEchoView(ExternalApiView):
|
|
def get(self, request):
|
|
return Response(
|
|
{
|
|
"user_id": request.user.id,
|
|
"api_key_id": request.auth.id,
|
|
}
|
|
)
|
|
|
|
|
|
class DefaultAuthProbeView(APIView):
|
|
def get(self, request):
|
|
return Response({"ok": True})
|
|
|
|
|
|
urlpatterns = [
|
|
path("api/test-auth/", AuthenticatedEchoView.as_view()),
|
|
path("api/default-auth/", DefaultAuthProbeView.as_view()),
|
|
]
|
|
|
|
|
|
@override_settings(ROOT_URLCONF=__name__)
|
|
class ApiKeyAuthenticationTests(TestCase):
|
|
url = "/api/test-auth/"
|
|
|
|
def setUp(self):
|
|
cache.clear()
|
|
suffix = uuid.uuid4().hex[:8]
|
|
self.user = get_user_model().objects.create_user(
|
|
username=f"api-user-{suffix}",
|
|
email=f"api-user-{suffix}@example.com",
|
|
password="password",
|
|
)
|
|
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="test")
|
|
self.client = APIClient()
|
|
|
|
def auth_header(self, raw_key: str | None = None) -> dict:
|
|
return {"HTTP_AUTHORIZATION": f"Bearer {raw_key or self.raw_key}"}
|
|
|
|
def test_external_api_view_only_uses_api_key_authentication(self):
|
|
self.assertEqual(AuthenticatedEchoView.authentication_classes, (ApiKeyAuthentication,))
|
|
|
|
def test_global_drf_default_does_not_accept_web_session_authentication(self):
|
|
self.client.force_login(self.user)
|
|
|
|
response = self.client.get("/api/default-auth/")
|
|
|
|
self.assertEqual(response.status_code, 403)
|
|
|
|
def test_valid_bearer_key_authenticates_user_and_api_key(self):
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["user_id"], self.user.id)
|
|
self.assertEqual(response.data["api_key_id"], self.api_key.id)
|
|
|
|
self.api_key.refresh_from_db()
|
|
self.assertIsNotNone(self.api_key.last_used_at)
|
|
|
|
def test_missing_api_key_returns_401(self):
|
|
response = self.client.get(self.url)
|
|
|
|
self.assertEqual(response.status_code, 401)
|
|
self.assertEqual(response["WWW-Authenticate"], "Bearer")
|
|
self.assertEqual(response.data["error"]["code"], "unauthorized")
|
|
|
|
def test_invalid_api_key_returns_401(self):
|
|
response = self.client.get(self.url, **self.auth_header("sk_cmhub_invalid"))
|
|
|
|
self.assertEqual(response.status_code, 401)
|
|
self.assertEqual(response["WWW-Authenticate"], "Bearer")
|
|
self.assertEqual(response.data["error"]["code"], "unauthorized")
|
|
|
|
@override_settings(API_AUTH_FAILURE_THROTTLE_RATE="1/min")
|
|
def test_invalid_api_key_failures_are_throttled_by_ip(self):
|
|
first = self.client.get(
|
|
self.url,
|
|
**self.auth_header("sk_cmhub_invalid"),
|
|
REMOTE_ADDR="198.51.100.21",
|
|
)
|
|
second = self.client.get(
|
|
self.url,
|
|
**self.auth_header("sk_cmhub_invalid"),
|
|
REMOTE_ADDR="198.51.100.21",
|
|
)
|
|
|
|
self.assertEqual(first.status_code, 401)
|
|
self.assertEqual(second.status_code, 429)
|
|
self.assertEqual(second.data["error"]["code"], "rate_limited")
|
|
|
|
def test_malformed_authorization_header_returns_401(self):
|
|
response = self.client.get(self.url, HTTP_AUTHORIZATION=f"Token {self.raw_key}")
|
|
|
|
self.assertEqual(response.status_code, 401)
|
|
self.assertEqual(response.data["error"]["code"], "unauthorized")
|
|
|
|
def test_revoked_api_key_returns_403(self):
|
|
self.api_key.status = ApiKey.Status.REVOKED
|
|
self.api_key.save(update_fields=("status", "updated_at"))
|
|
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
self.assertEqual(response.status_code, 403)
|
|
self.assertEqual(response.data["error"]["code"], "account_disabled")
|
|
|
|
def test_disabled_user_returns_403(self):
|
|
self.user.status = self.user.Status.DISABLED
|
|
self.user.save(update_fields=("status",))
|
|
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
self.assertEqual(response.status_code, 403)
|
|
self.assertEqual(response.data["error"]["code"], "account_disabled")
|
|
|
|
def test_web_session_login_is_not_accepted_for_external_api(self):
|
|
self.client.force_login(self.user)
|
|
|
|
response = self.client.get(self.url)
|
|
|
|
self.assertEqual(response.status_code, 401)
|
|
self.assertEqual(response.data["error"]["code"], "unauthorized")
|
|
|
|
|
|
class BalanceApiTests(TestCase):
|
|
url = "/api/v1/balance"
|
|
|
|
def setUp(self):
|
|
suffix = uuid.uuid4().hex[:8]
|
|
self.user = get_user_model().objects.create_user(
|
|
username=f"balance-user-{suffix}",
|
|
email=f"balance-user-{suffix}@example.com",
|
|
password="password",
|
|
first_name="主账号",
|
|
)
|
|
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="balance")
|
|
self.client = APIClient()
|
|
|
|
def auth_header(self, raw_key: str | None = None) -> dict:
|
|
return {"HTTP_AUTHORIZATION": f"Bearer {raw_key or self.raw_key}"}
|
|
|
|
def test_balance_returns_wallet_balance_matching_ledger_sum(self):
|
|
UserWallet.objects.create(user=self.user, points_balance=100)
|
|
PointsLedger.objects.create(
|
|
user=self.user,
|
|
change_type=PointsLedger.ChangeType.RECHARGE,
|
|
points_delta=120,
|
|
balance_after=120,
|
|
ref_order_id=1,
|
|
)
|
|
call = CallRecord.objects.create(
|
|
user=self.user,
|
|
api_key=self.api_key,
|
|
operation_type=CallRecord.OperationType.TITLE,
|
|
alias="title-standard",
|
|
model_used="gpt-5.5",
|
|
points_cost=20,
|
|
status=CallRecord.Status.SUCCESS,
|
|
)
|
|
PointsLedger.objects.create(
|
|
user=self.user,
|
|
change_type=PointsLedger.ChangeType.CONSUME,
|
|
points_delta=-20,
|
|
balance_after=100,
|
|
ref_call=call,
|
|
)
|
|
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
ledger_sum = sum(
|
|
PointsLedger.objects.filter(user=self.user).values_list(
|
|
"points_delta",
|
|
flat=True,
|
|
)
|
|
)
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["user"], self.user.username)
|
|
self.assertEqual(response.data["points_balance"], 100)
|
|
self.assertEqual(response.data["points_balance"], ledger_sum)
|
|
self.assertEqual(
|
|
response.data["account"],
|
|
{
|
|
"username": self.user.username,
|
|
"display_name": "主账号",
|
|
},
|
|
)
|
|
self.assertNotIn("email", response.data["account"])
|
|
self.assertNotIn("id", response.data["account"])
|
|
|
|
def test_balance_returns_zero_without_creating_missing_wallet(self):
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["user"], self.user.username)
|
|
self.assertEqual(response.data["points_balance"], 0)
|
|
self.assertEqual(response.data["account"]["username"], self.user.username)
|
|
self.assertEqual(response.data["account"]["display_name"], "主账号")
|
|
self.assertFalse(UserWallet.objects.filter(user=self.user).exists())
|
|
|
|
def test_balance_does_not_accept_web_session_without_api_key(self):
|
|
self.client.force_login(self.user)
|
|
|
|
response = self.client.get(self.url)
|
|
|
|
self.assertEqual(response.status_code, 401)
|
|
self.assertEqual(response.data["error"]["code"], "unauthorized")
|
|
|
|
|
|
class ModelsCatalogApiTests(TestCase):
|
|
url = "/api/v1/models"
|
|
|
|
def setUp(self):
|
|
cache.clear()
|
|
suffix = uuid.uuid4().hex[:8]
|
|
self.user = get_user_model().objects.create_user(
|
|
username=f"models-user-{suffix}",
|
|
email=f"models-user-{suffix}@example.com",
|
|
password="password",
|
|
)
|
|
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="models")
|
|
self.client = APIClient()
|
|
|
|
def auth_header(self, raw_key: str | None = None) -> dict:
|
|
return {"HTTP_AUTHORIZATION": f"Bearer {raw_key or self.raw_key}"}
|
|
|
|
def create_alias(
|
|
self,
|
|
*,
|
|
alias: str,
|
|
operation_type: str = ModelAlias.OperationType.TITLE,
|
|
capabilities: list[str] | None = None,
|
|
api_type: str = AiModel.ApiType.CHAT,
|
|
url: str = "https://provider-secret.example/v1/chat/completions",
|
|
model_sku: str = "secret-sku-gpt-5.5",
|
|
model_active: bool = True,
|
|
alias_active: bool = True,
|
|
) -> ModelAlias:
|
|
ai_model = AiModel.objects.create(
|
|
name=f"{alias}-{uuid.uuid4().hex[:8]}",
|
|
url=url,
|
|
model=model_sku,
|
|
api_type=api_type,
|
|
api_key_encrypted="encrypted-provider-key",
|
|
capabilities=capabilities if capabilities is not None else ["text"],
|
|
extra_body={"internal": "provider-extra-secret"},
|
|
is_active=model_active,
|
|
)
|
|
return ModelAlias.objects.create(
|
|
alias=alias,
|
|
operation_type=operation_type,
|
|
ai_model=ai_model,
|
|
is_active=alias_active,
|
|
)
|
|
|
|
def test_models_returns_public_alias_catalog_without_internal_fields(self):
|
|
title_alias = self.create_alias(alias="title-standard", capabilities=["text"])
|
|
image_alias = self.create_alias(
|
|
alias="image-edit",
|
|
operation_type=ModelAlias.OperationType.IMAGE,
|
|
capabilities=["image", "vision"],
|
|
api_type=AiModel.ApiType.IMAGES_EDITS,
|
|
url="https://provider-secret.example/v1/images/edits",
|
|
model_sku="secret-sku-image-2",
|
|
)
|
|
PricingRule.objects.create(
|
|
operation_type=title_alias.operation_type,
|
|
alias=title_alias.alias,
|
|
resolution="",
|
|
points_cost=2,
|
|
)
|
|
PricingRule.objects.create(
|
|
operation_type=image_alias.operation_type,
|
|
alias=image_alias.alias,
|
|
resolution="",
|
|
points_cost=10,
|
|
)
|
|
PricingRule.objects.create(
|
|
operation_type=image_alias.operation_type,
|
|
alias=image_alias.alias,
|
|
resolution="1k",
|
|
points_cost=12,
|
|
)
|
|
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertNotIn(GenerateRateThrottle, ModelsView.throttle_classes)
|
|
models = {item["alias"]: item for item in response.data["models"]}
|
|
self.assertEqual(set(models), {"title-standard", "image-edit"})
|
|
self.assertEqual(
|
|
set(models["title-standard"]),
|
|
{
|
|
"alias",
|
|
"operation_type",
|
|
"capabilities",
|
|
"requires_image",
|
|
"pricing_status",
|
|
"prices",
|
|
},
|
|
)
|
|
self.assertEqual(models["title-standard"]["operation_type"], "title")
|
|
self.assertEqual(models["title-standard"]["capabilities"], ["text"])
|
|
self.assertFalse(models["title-standard"]["requires_image"])
|
|
self.assertEqual(models["title-standard"]["pricing_status"], "priced")
|
|
self.assertEqual(
|
|
models["title-standard"]["prices"],
|
|
[{"resolution": "default", "points_cost": 2}],
|
|
)
|
|
self.assertEqual(models["image-edit"]["capabilities"], ["image", "vision"])
|
|
self.assertTrue(models["image-edit"]["requires_image"])
|
|
self.assertEqual(
|
|
models["image-edit"]["prices"],
|
|
[
|
|
{"resolution": "default", "points_cost": 10},
|
|
{"resolution": "1K", "points_cost": 12},
|
|
],
|
|
)
|
|
response_body = json.dumps(response.data, ensure_ascii=False)
|
|
self.assertNotIn("secret-sku", response_body)
|
|
self.assertNotIn("provider-secret.example", response_body)
|
|
self.assertNotIn("encrypted-provider-key", response_body)
|
|
self.assertNotIn("provider-extra-secret", response_body)
|
|
self.assertNotIn("api_key", response_body)
|
|
self.assertNotIn("api_key_encrypted", response_body)
|
|
self.assertNotIn("extra_body", response_body)
|
|
self.assertNotIn("url", response_body)
|
|
self.assertNotIn("model_used", response_body)
|
|
|
|
def test_models_rejects_missing_invalid_and_session_only_authentication(self):
|
|
missing = self.client.get(self.url)
|
|
invalid = self.client.get(self.url, **self.auth_header("sk_cmhub_invalid"))
|
|
self.client.force_login(self.user)
|
|
session_only = self.client.get(self.url)
|
|
|
|
self.assertEqual(missing.status_code, 401)
|
|
self.assertEqual(missing.data["error"]["code"], "unauthorized")
|
|
self.assertEqual(invalid.status_code, 401)
|
|
self.assertEqual(invalid.data["error"]["code"], "unauthorized")
|
|
self.assertEqual(session_only.status_code, 401)
|
|
self.assertEqual(session_only.data["error"]["code"], "unauthorized")
|
|
|
|
def test_models_only_lists_callable_active_aliases_and_allows_unpriced_alias(self):
|
|
self.create_alias(alias="title-unpriced", capabilities=["text"])
|
|
self.create_alias(alias="title-inactive-alias", alias_active=False)
|
|
self.create_alias(alias="title-inactive-model", model_active=False)
|
|
self.create_alias(alias="title-wrong-capability", capabilities=["image"])
|
|
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(len(response.data["models"]), 1)
|
|
item = response.data["models"][0]
|
|
self.assertEqual(item["alias"], "title-unpriced")
|
|
self.assertEqual(item["pricing_status"], "unpriced")
|
|
self.assertEqual(item["prices"], [])
|
|
|
|
|
|
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 客户端",
|
|
) -> 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,
|
|
)
|
|
|
|
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",
|
|
)
|
|
|
|
response = self.client.get(self.url)
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertNotIn(GenerateRateThrottle, ClientLatestReleaseView.throttle_classes)
|
|
self.assertEqual(response.data["platform"], "windows")
|
|
self.assertEqual(
|
|
set(response.data["release"]),
|
|
{"version", "download_url", "sha256", "release_notes", "published_at"},
|
|
)
|
|
self.assertEqual(response.data["release"]["version"], "1.2.3")
|
|
self.assertEqual(
|
|
response.data["release"]["download_url"],
|
|
"https://download.example.com/cmhub-1.2.3.exe",
|
|
)
|
|
self.assertEqual(response.data["release"]["sha256"], "b" * 64)
|
|
self.assertEqual(
|
|
response.data["release"]["release_notes"],
|
|
"修复下载入口并补充 SHA256",
|
|
)
|
|
self.assertEqual(
|
|
response.data["release"]["published_at"],
|
|
timezone.localtime(release.updated_at).isoformat(),
|
|
)
|
|
|
|
def test_latest_release_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"], "暂未发布")
|
|
|
|
def test_latest_release_rejects_invalid_platform(self):
|
|
response = self.client.get(f"{self.url}?platform=android")
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "bad_request")
|
|
|
|
def test_latest_release_supports_non_windows_platform(self):
|
|
self.create_release(
|
|
platform=DownloadRelease.Platform.MACOS,
|
|
version="2.0.0",
|
|
external_url="https://download.example.com/cmhub-2.0.0.dmg",
|
|
sha256="c" * 64,
|
|
release_notes="macOS 客户端",
|
|
)
|
|
|
|
response = self.client.get(f"{self.url}?platform=macos")
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["platform"], "macos")
|
|
self.assertEqual(response.data["release"]["version"], "2.0.0")
|
|
|
|
def test_latest_release_response_does_not_expose_internal_fields(self):
|
|
self.create_release(
|
|
external_url="",
|
|
file_name="downloads/cmhub-desktop-1.0.0.exe",
|
|
)
|
|
|
|
response = self.client.get(self.url, secure=True)
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(set(response.data), {"platform", "release"})
|
|
self.assertEqual(
|
|
set(response.data["release"]),
|
|
{"version", "download_url", "sha256", "release_notes", "published_at"},
|
|
)
|
|
response_body = json.dumps(response.data, ensure_ascii=False)
|
|
for forbidden in (
|
|
"id",
|
|
"is_current",
|
|
"created_at",
|
|
"updated_at",
|
|
"MEDIA_ROOT",
|
|
str(settings.MEDIA_ROOT),
|
|
"user",
|
|
"email",
|
|
"api_key",
|
|
"api_key_encrypted",
|
|
"model_used",
|
|
):
|
|
self.assertNotIn(forbidden, response_body)
|
|
|
|
|
|
@override_settings(
|
|
PAYMENT_CALLBACK_MODE="mock",
|
|
PAYMENT_MOCK_CALLBACK_SECRET="test-payment-callback-secret",
|
|
)
|
|
class RechargeCallbackApiTests(TestCase):
|
|
wechat_url = "/api/v1/recharge/callback/wechat"
|
|
alipay_url = "/api/v1/recharge/callback/alipay"
|
|
|
|
def setUp(self):
|
|
suffix = uuid.uuid4().hex[:8]
|
|
self.user = get_user_model().objects.create_user(
|
|
username=f"recharge-user-{suffix}",
|
|
email=f"recharge-user-{suffix}@example.com",
|
|
password="password",
|
|
)
|
|
self.wallet = UserWallet.objects.create(user=self.user, points_balance=100)
|
|
self.client = APIClient(enforce_csrf_checks=True)
|
|
|
|
def create_order(
|
|
self,
|
|
*,
|
|
amount="20.00",
|
|
points_granted=200,
|
|
pay_method=RechargeOrder.PayMethod.WEIXIN,
|
|
) -> RechargeOrder:
|
|
return RechargeOrder.objects.create(
|
|
user=self.user,
|
|
order_no=f"R{uuid.uuid4().hex[:12]}",
|
|
amount_money=Decimal(amount),
|
|
pay_method=pay_method,
|
|
exchange_rate=Decimal("10.0000"),
|
|
points_granted=points_granted,
|
|
)
|
|
|
|
def signed_wechat_body(self, order, *, total_cents=2000):
|
|
payload = {
|
|
"event_type": "TRANSACTION.SUCCESS",
|
|
"resource": {
|
|
"trade_state": "SUCCESS",
|
|
"out_trade_no": order.order_no,
|
|
"transaction_id": "wx-txn-001",
|
|
"success_time": "2026-07-03T00:00:00+08:00",
|
|
"amount": {"total": total_cents},
|
|
},
|
|
}
|
|
body = json.dumps(payload, separators=(",", ":")).encode("utf-8")
|
|
return body, build_mock_body_signature(body)
|
|
|
|
def signed_alipay_payload(self, order, *, total_amount="20.00"):
|
|
payload = {
|
|
"trade_status": "TRADE_SUCCESS",
|
|
"out_trade_no": order.order_no,
|
|
"trade_no": "ali-txn-001",
|
|
"total_amount": total_amount,
|
|
"gmt_payment": "2026-07-03 00:00:00",
|
|
}
|
|
payload["sign"] = build_mock_alipay_signature(payload)
|
|
return payload
|
|
|
|
def test_wechat_callback_credits_once_and_is_csrf_exempt(self):
|
|
order = self.create_order(amount="20.00", points_granted=200)
|
|
body, signature = self.signed_wechat_body(order)
|
|
|
|
first = self.client.post(
|
|
self.wechat_url,
|
|
data=body,
|
|
content_type="application/json",
|
|
HTTP_WECHATPAY_SIGNATURE=signature,
|
|
)
|
|
second = self.client.post(
|
|
self.wechat_url,
|
|
data=body,
|
|
content_type="application/json",
|
|
HTTP_WECHATPAY_SIGNATURE=signature,
|
|
)
|
|
|
|
self.assertEqual(first.status_code, 200)
|
|
self.assertEqual(second.status_code, 200)
|
|
self.assertEqual(first.data["code"], "SUCCESS")
|
|
self.wallet.refresh_from_db()
|
|
order.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 300)
|
|
self.assertEqual(order.status, RechargeOrder.Status.PAID)
|
|
self.assertEqual(order.payment_txn_no, "wx-txn-001")
|
|
self.assertEqual(
|
|
PointsLedger.objects.filter(
|
|
user=self.user,
|
|
ref_order_id=order.id,
|
|
change_type=PointsLedger.ChangeType.RECHARGE,
|
|
).count(),
|
|
1,
|
|
)
|
|
|
|
def test_wechat_callback_rejects_bad_signature_without_crediting(self):
|
|
order = self.create_order(amount="20.00", points_granted=200)
|
|
body, _signature = self.signed_wechat_body(order)
|
|
|
|
response = self.client.post(
|
|
self.wechat_url,
|
|
data=body,
|
|
content_type="application/json",
|
|
HTTP_WECHATPAY_SIGNATURE="bad-signature",
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "signature_invalid")
|
|
self.wallet.refresh_from_db()
|
|
order.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 100)
|
|
self.assertEqual(order.status, RechargeOrder.Status.PENDING)
|
|
self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists())
|
|
|
|
def test_wechat_callback_rejects_amount_mismatch_without_crediting(self):
|
|
order = self.create_order(amount="20.00", points_granted=200)
|
|
body, signature = self.signed_wechat_body(order, total_cents=1999)
|
|
|
|
response = self.client.post(
|
|
self.wechat_url,
|
|
data=body,
|
|
content_type="application/json",
|
|
HTTP_WECHATPAY_SIGNATURE=signature,
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "amount_mismatch")
|
|
self.wallet.refresh_from_db()
|
|
order.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 100)
|
|
self.assertEqual(order.status, RechargeOrder.Status.PENDING)
|
|
self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists())
|
|
|
|
def test_alipay_callback_credits_once_returns_success_and_is_csrf_exempt(self):
|
|
order = self.create_order(
|
|
amount="20.00",
|
|
points_granted=200,
|
|
pay_method=RechargeOrder.PayMethod.ALIPAY,
|
|
)
|
|
payload = self.signed_alipay_payload(order)
|
|
|
|
first = self.client.post(self.alipay_url, data=payload)
|
|
second = self.client.post(self.alipay_url, data=payload)
|
|
|
|
self.assertEqual(first.status_code, 200)
|
|
self.assertEqual(second.status_code, 200)
|
|
self.assertEqual(first.content, b"success")
|
|
self.wallet.refresh_from_db()
|
|
order.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 300)
|
|
self.assertEqual(order.status, RechargeOrder.Status.PAID)
|
|
self.assertEqual(order.payment_txn_no, "ali-txn-001")
|
|
self.assertEqual(
|
|
PointsLedger.objects.filter(
|
|
ref_order_id=order.id,
|
|
change_type=PointsLedger.ChangeType.RECHARGE,
|
|
).count(),
|
|
1,
|
|
)
|
|
|
|
def test_alipay_callback_rejects_bad_signature_without_crediting(self):
|
|
order = self.create_order(
|
|
amount="20.00",
|
|
points_granted=200,
|
|
pay_method=RechargeOrder.PayMethod.ALIPAY,
|
|
)
|
|
payload = self.signed_alipay_payload(order)
|
|
payload["sign"] = "bad-signature"
|
|
|
|
response = self.client.post(self.alipay_url, data=payload)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.content, b"fail")
|
|
self.wallet.refresh_from_db()
|
|
order.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 100)
|
|
self.assertEqual(order.status, RechargeOrder.Status.PENDING)
|
|
self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists())
|
|
|
|
|
|
@override_settings(
|
|
PAYMENT_CALLBACK_MODE="mock",
|
|
PAYMENT_MOCK_CALLBACK_SECRET="test-payment-callback-secret",
|
|
PAYMENT_QR_EXPIRES_MINUTES=15,
|
|
)
|
|
class RechargeCreateStatusApiTests(TestCase):
|
|
create_url = "/api/v1/recharge/create"
|
|
status_url = "/api/v1/recharge/status"
|
|
|
|
def setUp(self):
|
|
suffix = uuid.uuid4().hex[:8]
|
|
self.user = get_user_model().objects.create_user(
|
|
username=f"recharge-create-{suffix}",
|
|
email=f"recharge-create-{suffix}@example.com",
|
|
password="password",
|
|
)
|
|
self.other_user = get_user_model().objects.create_user(
|
|
username=f"recharge-other-{suffix}",
|
|
email=f"recharge-other-{suffix}@example.com",
|
|
password="password",
|
|
)
|
|
self.wallet = UserWallet.objects.create(user=self.user, points_balance=100)
|
|
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="recharge")
|
|
self.exchange_rate = ExchangeRate.objects.create(
|
|
currency="CNY",
|
|
points_per_unit=Decimal("10.0000"),
|
|
effective_from=timezone.now(),
|
|
)
|
|
self.client = APIClient()
|
|
|
|
def create_order(
|
|
self,
|
|
*,
|
|
user=None,
|
|
amount="20.00",
|
|
points_granted=200,
|
|
pay_method=RechargeOrder.PayMethod.WEIXIN,
|
|
):
|
|
return RechargeOrder.objects.create(
|
|
user=user or self.user,
|
|
order_no=f"R{uuid.uuid4().hex[:12]}",
|
|
amount_money=Decimal(amount),
|
|
pay_method=pay_method,
|
|
exchange_rate=Decimal("10.0000"),
|
|
points_granted=points_granted,
|
|
code_url=f"mockpay://{pay_method}/existing",
|
|
)
|
|
|
|
def test_recharge_create_requires_web_session_not_api_key(self):
|
|
response = self.client.post(
|
|
self.create_url,
|
|
{"amount": "20.00", "pay_method": "weixin"},
|
|
format="json",
|
|
HTTP_AUTHORIZATION=f"Bearer {self.raw_key}",
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 403)
|
|
self.assertFalse(RechargeOrder.objects.filter(user=self.user).exists())
|
|
|
|
def test_recharge_create_is_session_authenticated_and_locks_quote(self):
|
|
self.client.force_login(self.user)
|
|
|
|
response = self.client.post(
|
|
self.create_url,
|
|
{"amount": "20.00", "pay_method": "weixin"},
|
|
format="json",
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 201)
|
|
self.assertEqual(response.data["amount"], "20.00")
|
|
self.assertEqual(response.data["exchange_rate"], "10.0000")
|
|
self.assertEqual(response.data["points_granted"], 200)
|
|
self.assertEqual(response.data["pay_method"], RechargeOrder.PayMethod.WEIXIN)
|
|
self.assertEqual(response.data["status"], RechargeOrder.Status.PENDING)
|
|
self.assertTrue(response.data["code_url"].startswith("weixin://wxpay/cmhub-mock"))
|
|
self.assertIsNotNone(response.data["expires_at"])
|
|
|
|
order = RechargeOrder.objects.get(order_no=response.data["order_no"])
|
|
self.assertEqual(order.user, self.user)
|
|
self.assertEqual(order.exchange_rate, Decimal("10.0000"))
|
|
self.assertEqual(order.points_granted, 200)
|
|
self.wallet.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 100)
|
|
self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists())
|
|
|
|
@override_settings(RECHARGE_MAX_AMOUNT_CNY="100.00")
|
|
def test_recharge_create_rejects_amount_above_configured_maximum(self):
|
|
self.client.force_login(self.user)
|
|
|
|
response = self.client.post(
|
|
self.create_url,
|
|
{"amount": "100.01", "pay_method": "weixin"},
|
|
format="json",
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "bad_request")
|
|
self.assertFalse(RechargeOrder.objects.filter(user=self.user).exists())
|
|
|
|
def test_recharge_create_supports_alipay_mock_qr_code(self):
|
|
self.client.force_login(self.user)
|
|
|
|
response = self.client.post(
|
|
self.create_url,
|
|
{"amount": "30.00", "pay_method": "alipay"},
|
|
format="json",
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 201)
|
|
self.assertEqual(response.data["pay_method"], RechargeOrder.PayMethod.ALIPAY)
|
|
self.assertTrue(response.data["code_url"].startswith("https://qr.alipay.com/cmhub-mock"))
|
|
|
|
def test_recharge_create_enforces_csrf_for_real_session_clients(self):
|
|
csrf_client = APIClient(enforce_csrf_checks=True)
|
|
csrf_client.force_login(self.user)
|
|
|
|
response = csrf_client.post(
|
|
self.create_url,
|
|
{"amount": "20.00", "pay_method": "weixin"},
|
|
format="json",
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 403)
|
|
self.assertFalse(RechargeOrder.objects.filter(user=self.user).exists())
|
|
|
|
def test_recharge_status_returns_pending_order_for_owner_only(self):
|
|
order = self.create_order()
|
|
self.client.force_login(self.user)
|
|
|
|
response = self.client.get(self.status_url, {"order_no": order.order_no})
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["order_no"], order.order_no)
|
|
self.assertEqual(response.data["status"], RechargeOrder.Status.PENDING)
|
|
|
|
self.client.force_login(self.other_user)
|
|
denied = self.client.get(self.status_url, {"order_no": order.order_no})
|
|
|
|
self.assertEqual(denied.status_code, 404)
|
|
self.assertEqual(denied.data["error"]["code"], "order_not_found")
|
|
|
|
def test_recharge_status_active_query_can_apply_paid_order_once(self):
|
|
order = self.create_order(amount="20.00", points_granted=200)
|
|
self.client.force_login(self.user)
|
|
|
|
def fake_query(queried_order):
|
|
return RechargePayment(
|
|
order_no=queried_order.order_no,
|
|
pay_method=queried_order.pay_method,
|
|
amount=queried_order.amount_money,
|
|
transaction_id="queried-txn-001",
|
|
paid_at=timezone.now(),
|
|
)
|
|
|
|
with patch("apps.api.views.query_payment_order", side_effect=fake_query) as query:
|
|
first = self.client.get(self.status_url, {"order_no": order.order_no})
|
|
second = self.client.get(self.status_url, {"order_no": order.order_no})
|
|
|
|
self.assertEqual(first.status_code, 200)
|
|
self.assertEqual(second.status_code, 200)
|
|
self.assertEqual(first.data["status"], RechargeOrder.Status.PAID)
|
|
self.assertEqual(second.data["status"], RechargeOrder.Status.PAID)
|
|
self.assertEqual(query.call_count, 1)
|
|
self.wallet.refresh_from_db()
|
|
order.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 300)
|
|
self.assertEqual(order.payment_txn_no, "queried-txn-001")
|
|
self.assertEqual(
|
|
PointsLedger.objects.filter(
|
|
ref_order_id=order.id,
|
|
change_type=PointsLedger.ChangeType.RECHARGE,
|
|
).count(),
|
|
1,
|
|
)
|
|
|
|
def test_recharge_status_keeps_pending_when_active_query_raises_unexpected_error(self):
|
|
order = self.create_order(amount="20.00", points_granted=200)
|
|
self.client.force_login(self.user)
|
|
|
|
with patch("apps.api.views.query_payment_order", side_effect=RuntimeError("gateway down")):
|
|
response = self.client.get(self.status_url, {"order_no": order.order_no})
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["order_no"], order.order_no)
|
|
self.assertEqual(response.data["status"], RechargeOrder.Status.PENDING)
|
|
self.wallet.refresh_from_db()
|
|
order.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 100)
|
|
self.assertEqual(order.status, RechargeOrder.Status.PENDING)
|
|
self.assertFalse(
|
|
PointsLedger.objects.filter(
|
|
ref_order_id=order.id,
|
|
change_type=PointsLedger.ChangeType.RECHARGE,
|
|
).exists()
|
|
)
|
|
|
|
|
|
class FakeGenerationProvider:
|
|
def __init__(self, *, capabilities=None):
|
|
self._capabilities = set(capabilities or {"text", "image", "vision"})
|
|
self.text_calls = []
|
|
self.image_calls = []
|
|
self.text_error = None
|
|
self.image_error = None
|
|
|
|
def capabilities(self):
|
|
return set(self._capabilities)
|
|
|
|
def generate_text(self, prompt, model, **kwargs):
|
|
self.text_calls.append({"prompt": prompt, "model": model, **kwargs})
|
|
if self.text_error is not None:
|
|
raise self.text_error
|
|
return TextGenerationResult(
|
|
text="测试标题一",
|
|
titles=("测试标题一", "测试标题二"),
|
|
model_used=model.model,
|
|
raw={"secret": "SECRET_RAW_SHOULD_NOT_BE_STORED"},
|
|
)
|
|
|
|
def generate_image(self, prompt, model, **kwargs):
|
|
self.image_calls.append({"prompt": prompt, "model": model, **kwargs})
|
|
if self.image_error is not None:
|
|
raise self.image_error
|
|
return ImageGenerationResult(
|
|
image=b"generated-image-bytes",
|
|
model_used=model.model,
|
|
raw={"b64_json": "SECRET_RAW_SHOULD_NOT_BE_STORED"},
|
|
)
|
|
|
|
|
|
class FakeImageUrlResponse:
|
|
def __init__(self, *, status_code=200, headers=None, chunks=()):
|
|
self.status_code = status_code
|
|
self.headers = headers or {}
|
|
self._chunks = list(chunks)
|
|
self.closed = False
|
|
|
|
def raise_for_status(self):
|
|
if self.status_code >= 400:
|
|
raise requests.HTTPError("image_url request failed", response=self)
|
|
|
|
def iter_content(self, chunk_size=1):
|
|
for chunk in self._chunks:
|
|
yield chunk
|
|
|
|
def close(self):
|
|
self.closed = True
|
|
|
|
|
|
def dns_result(address: str):
|
|
return [(None, None, None, "", (address, 443))]
|
|
|
|
|
|
@override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY)
|
|
class GenerateApiTests(TestCase):
|
|
def setUp(self):
|
|
cache.clear()
|
|
reset_keyword_matcher_cache()
|
|
suffix = uuid.uuid4().hex[:8]
|
|
self.media_dir = tempfile.TemporaryDirectory()
|
|
self.addCleanup(self.media_dir.cleanup)
|
|
self.addCleanup(reset_keyword_matcher_cache)
|
|
media_override = override_settings(
|
|
MEDIA_ROOT=self.media_dir.name,
|
|
MEDIA_URL="/media/",
|
|
)
|
|
media_override.enable()
|
|
self.addCleanup(media_override.disable)
|
|
|
|
self.user = get_user_model().objects.create_user(
|
|
username=f"generate-user-{suffix}",
|
|
email=f"generate-user-{suffix}@example.com",
|
|
password="password",
|
|
)
|
|
self.wallet = UserWallet.objects.create(user=self.user, points_balance=100)
|
|
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="generate")
|
|
self.client = APIClient()
|
|
self.provider = FakeGenerationProvider()
|
|
|
|
self.title_model = self.create_ai_model(
|
|
name=f"title-model-{suffix}",
|
|
model=f"gpt-title-{suffix}",
|
|
capabilities=["text", "vision"],
|
|
)
|
|
self.image_model = self.create_ai_model(
|
|
name=f"image-model-{suffix}",
|
|
model=f"gpt-image-{suffix}",
|
|
capabilities=["image", "vision"],
|
|
)
|
|
self.title_alias = f"title-standard-{suffix}"
|
|
self.image_alias = f"image-hd-{suffix}"
|
|
ModelAlias.objects.create(
|
|
operation_type=ModelAlias.OperationType.TITLE,
|
|
alias=self.title_alias,
|
|
ai_model=self.title_model,
|
|
is_default=True,
|
|
)
|
|
ModelAlias.objects.create(
|
|
operation_type=ModelAlias.OperationType.IMAGE,
|
|
alias=self.image_alias,
|
|
ai_model=self.image_model,
|
|
is_default=True,
|
|
)
|
|
PricingRule.objects.create(
|
|
operation_type=CallRecord.OperationType.TITLE,
|
|
alias=self.title_alias,
|
|
points_cost=2,
|
|
)
|
|
PricingRule.objects.create(
|
|
operation_type=CallRecord.OperationType.IMAGE,
|
|
alias=self.image_alias,
|
|
resolution="1K",
|
|
points_cost=10,
|
|
)
|
|
|
|
def create_ai_model(self, *, name, model, capabilities):
|
|
ai_model = AiModel(
|
|
name=name,
|
|
url="https://api.example.test/v1",
|
|
model=model,
|
|
api_type=AiModel.ApiType.CHAT,
|
|
capabilities=capabilities,
|
|
)
|
|
ai_model.set_api_key("sk-test-secret")
|
|
ai_model.save()
|
|
return ai_model
|
|
|
|
def auth_header(self) -> dict:
|
|
return {"HTTP_AUTHORIZATION": f"Bearer {self.raw_key}"}
|
|
|
|
def post_with_provider(self, path, payload, provider=None):
|
|
with patch("apps.api.generation.get_provider", return_value=provider or self.provider):
|
|
return self.client.post(path, payload, format="json", **self.auth_header())
|
|
|
|
def assert_generation_not_charged(self):
|
|
self.wallet.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 100)
|
|
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
|
|
self.assertFalse(PointsLedger.objects.filter(user=self.user).exists())
|
|
|
|
@override_settings(
|
|
MODERATION_ENABLED=True,
|
|
MODERATION_PROVIDER="keyword",
|
|
MODERATION_CACHE_VERSION_KEY="test:api:moderation:sensitive_words:version",
|
|
)
|
|
def test_blocked_prompt_returns_content_blocked_before_image_download_or_charge(self):
|
|
SensitiveWord.objects.create(word="敏感词", category="policy")
|
|
|
|
with (
|
|
patch("apps.api.generation.socket.getaddrinfo") as dns_lookup,
|
|
patch("apps.api.generation.requests.Session.get") as image_get,
|
|
):
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/image",
|
|
{
|
|
"prompt": "请生成敏-感\u200b 词图片",
|
|
"model": self.image_alias,
|
|
"image_url": "https://safe.example.com/input.jpg",
|
|
"resolution": "1K",
|
|
"aspect_ratio": "1:1",
|
|
},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "content_blocked")
|
|
dns_lookup.assert_not_called()
|
|
image_get.assert_not_called()
|
|
self.assertEqual(self.provider.image_calls, [])
|
|
self.assert_generation_not_charged()
|
|
|
|
@override_settings(
|
|
MODERATION_ENABLED=False,
|
|
MODERATION_PROVIDER="keyword",
|
|
MODERATION_CACHE_VERSION_KEY="test:api:moderation:sensitive_words:version",
|
|
)
|
|
def test_disabled_moderation_does_not_block_matching_prompt(self):
|
|
SensitiveWord.objects.create(word="敏感词")
|
|
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/title",
|
|
{"prompt": "敏感词", "model": self.title_alias},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["points_cost"], 2)
|
|
self.assertEqual(len(self.provider.text_calls), 1)
|
|
|
|
def test_generate_title_uses_default_alias_charges_points_and_writes_call_record(self):
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/title",
|
|
{
|
|
"prompt": "生成标题",
|
|
"resolution": "1k",
|
|
"parameters": {"temperature": 0.2, "model": "bad-overridden-model"},
|
|
},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["titles"], ["测试标题一", "测试标题二"])
|
|
self.assertEqual(response.data["alias"], self.title_alias)
|
|
self.assertEqual(response.data["model_used"], self.title_model.model)
|
|
self.assertEqual(response.data["points_cost"], 2)
|
|
self.assertEqual(response.data["points_balance"], 98)
|
|
self.assertEqual(len(self.provider.text_calls), 1)
|
|
self.assertEqual(self.provider.text_calls[0]["model"].model, self.title_model.model)
|
|
|
|
self.wallet.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 98)
|
|
call = CallRecord.objects.get(pk=response.data["call_id"])
|
|
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
|
|
self.assertEqual(call.api_key, self.api_key)
|
|
self.assertEqual(call.alias, self.title_alias)
|
|
self.assertEqual(call.model_used, self.title_model.model)
|
|
self.assertEqual(call.resolution, "1K")
|
|
self.assertNotIn("SECRET_RAW", call.result_summary)
|
|
self.assertEqual(
|
|
PointsLedger.objects.filter(
|
|
user=self.user,
|
|
change_type=PointsLedger.ChangeType.CONSUME,
|
|
).count(),
|
|
1,
|
|
)
|
|
|
|
def test_generate_image_stores_file_returns_url_and_does_not_store_raw_base64(self):
|
|
encoded = base64.b64encode(b"input-image").decode("ascii")
|
|
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/image",
|
|
{
|
|
"prompt": "生成图片",
|
|
"model": self.image_alias,
|
|
"image_base64": f"data:image/png;base64,{encoded}",
|
|
"resolution": "1K",
|
|
"aspect_ratio": "1:1",
|
|
},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["alias"], self.image_alias)
|
|
self.assertEqual(response.data["model_used"], self.image_model.model)
|
|
self.assertEqual(response.data["points_cost"], 10)
|
|
self.assertEqual(response.data["points_balance"], 90)
|
|
self.assertTrue(response.data["image_url"].startswith("http://testserver/media/"))
|
|
self.assertEqual(self.provider.image_calls[0]["image"], b"input-image")
|
|
|
|
media_relative_path = response.data["image_url"].split("/media/", 1)[1]
|
|
self.assertTrue((Path(self.media_dir.name) / media_relative_path).exists())
|
|
|
|
call = CallRecord.objects.get(pk=response.data["call_id"])
|
|
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
|
|
self.assertEqual(call.result_ref, response.data["image_url"])
|
|
self.assertEqual(call.result_summary, "image_bytes=21")
|
|
self.assertNotIn("SECRET_RAW", call.result_ref + call.result_summary)
|
|
|
|
def test_generate_image_downloads_safe_image_url(self):
|
|
response = FakeImageUrlResponse(
|
|
headers={"Content-Type": "image/jpeg"},
|
|
chunks=[b"remote-image"],
|
|
)
|
|
|
|
with (
|
|
patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("93.184.216.34")),
|
|
patch("apps.api.generation.requests.Session.get", return_value=response),
|
|
):
|
|
result = self.post_with_provider(
|
|
"/api/v1/generate/image",
|
|
{
|
|
"prompt": "生成图片",
|
|
"model": self.image_alias,
|
|
"image_url": "https://safe.example.com/input.jpg",
|
|
"resolution": "1K",
|
|
"aspect_ratio": "1:1",
|
|
},
|
|
)
|
|
|
|
self.assertEqual(result.status_code, 200)
|
|
self.assertEqual(self.provider.image_calls[0]["image"], b"remote-image")
|
|
self.assertEqual(self.provider.image_calls[0]["image_mime_type"], "image/jpeg")
|
|
self.assertTrue(response.closed)
|
|
|
|
def test_image_url_rejects_loopback_address_without_fetch_or_charge(self):
|
|
with (
|
|
patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("127.0.0.1")),
|
|
patch("apps.api.generation.requests.Session.get") as image_get,
|
|
):
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/image",
|
|
{
|
|
"prompt": "生成图片",
|
|
"model": self.image_alias,
|
|
"image_url": "http://127.0.0.1/private.png",
|
|
"resolution": "1K",
|
|
},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "bad_request")
|
|
image_get.assert_not_called()
|
|
self.assert_generation_not_charged()
|
|
|
|
def test_image_url_rejects_cloud_metadata_address_without_fetch_or_charge(self):
|
|
with (
|
|
patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("169.254.169.254")),
|
|
patch("apps.api.generation.requests.Session.get") as image_get,
|
|
):
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/title",
|
|
{
|
|
"prompt": "生成标题",
|
|
"model": self.title_alias,
|
|
"image_url": "http://169.254.169.254/latest/meta-data/",
|
|
},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "bad_request")
|
|
image_get.assert_not_called()
|
|
self.assert_generation_not_charged()
|
|
|
|
def test_image_url_rejects_redirect_to_private_address_without_charge(self):
|
|
def fake_getaddrinfo(host, port, *args, **kwargs):
|
|
if host == "safe.example.com":
|
|
return dns_result("93.184.216.34")
|
|
return dns_result("127.0.0.1")
|
|
|
|
redirect = FakeImageUrlResponse(
|
|
status_code=302,
|
|
headers={"Location": "http://127.0.0.1/private.png"},
|
|
)
|
|
|
|
with (
|
|
patch("apps.api.generation.socket.getaddrinfo", side_effect=fake_getaddrinfo),
|
|
patch("apps.api.generation.requests.Session.get", return_value=redirect) as image_get,
|
|
):
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/image",
|
|
{
|
|
"prompt": "生成图片",
|
|
"model": self.image_alias,
|
|
"image_url": "https://safe.example.com/input.png",
|
|
"resolution": "1K",
|
|
},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "bad_request")
|
|
self.assertEqual(image_get.call_count, 1)
|
|
self.assertTrue(redirect.closed)
|
|
self.assert_generation_not_charged()
|
|
|
|
@override_settings(IMAGE_URL_MAX_BYTES=4)
|
|
def test_image_url_rejects_oversized_response_without_charge(self):
|
|
oversized = FakeImageUrlResponse(
|
|
headers={"Content-Type": "image/png"},
|
|
chunks=[b"1234", b"5"],
|
|
)
|
|
|
|
with (
|
|
patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("93.184.216.34")),
|
|
patch("apps.api.generation.requests.Session.get", return_value=oversized),
|
|
):
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/image",
|
|
{
|
|
"prompt": "生成图片",
|
|
"model": self.image_alias,
|
|
"image_url": "https://safe.example.com/input.png",
|
|
"resolution": "1K",
|
|
},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "bad_request")
|
|
self.assertTrue(oversized.closed)
|
|
self.assert_generation_not_charged()
|
|
|
|
@override_settings(API_GENERATE_THROTTLE_RATE="1/min")
|
|
def test_generate_endpoint_is_throttled_by_api_key_without_extra_charge(self):
|
|
first = self.post_with_provider(
|
|
"/api/v1/generate/title",
|
|
{"prompt": "生成标题", "model": self.title_alias},
|
|
)
|
|
second = self.post_with_provider(
|
|
"/api/v1/generate/title",
|
|
{"prompt": "生成标题", "model": self.title_alias},
|
|
)
|
|
|
|
self.assertEqual(first.status_code, 200)
|
|
self.assertEqual(second.status_code, 429)
|
|
self.assertEqual(second.data["error"]["code"], "rate_limited")
|
|
self.assertEqual(len(self.provider.text_calls), 1)
|
|
self.wallet.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 98)
|
|
self.assertEqual(CallRecord.objects.filter(user=self.user).count(), 1)
|
|
|
|
def test_insufficient_points_returns_402_without_calling_provider_or_writing_call(self):
|
|
self.wallet.points_balance = 1
|
|
self.wallet.save(update_fields=("points_balance", "updated_at"))
|
|
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/title",
|
|
{"prompt": "生成标题", "model": self.title_alias},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 402)
|
|
self.assertEqual(response.data["error"]["code"], "insufficient_points")
|
|
self.assertEqual(self.provider.text_calls, [])
|
|
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
|
|
self.assertFalse(PointsLedger.objects.filter(user=self.user).exists())
|
|
self.wallet.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 1)
|
|
|
|
def test_missing_pricing_rule_returns_400_without_charging(self):
|
|
PricingRule.objects.filter(alias=self.title_alias).delete()
|
|
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/title",
|
|
{"prompt": "生成标题", "model": self.title_alias},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "no_pricing_rule")
|
|
self.assertEqual(self.provider.text_calls, [])
|
|
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
|
|
self.wallet.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 100)
|
|
|
|
def test_alias_capability_mismatch_returns_model_not_allowed_without_charging(self):
|
|
bad_alias = f"bad-title-{uuid.uuid4().hex[:8]}"
|
|
ModelAlias.objects.create(
|
|
operation_type=ModelAlias.OperationType.TITLE,
|
|
alias=bad_alias,
|
|
ai_model=self.image_model,
|
|
)
|
|
PricingRule.objects.create(
|
|
operation_type=CallRecord.OperationType.TITLE,
|
|
alias=bad_alias,
|
|
points_cost=2,
|
|
)
|
|
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/title",
|
|
{"prompt": "生成标题", "model": bad_alias},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "model_not_allowed")
|
|
self.assertEqual(self.provider.text_calls, [])
|
|
self.assertFalse(CallRecord.objects.filter(alias=bad_alias).exists())
|
|
|
|
def test_provider_capability_mismatch_returns_model_not_allowed_before_charging(self):
|
|
text_only_provider = FakeGenerationProvider(capabilities={"text"})
|
|
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/image",
|
|
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
|
|
provider=text_only_provider,
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "model_not_allowed")
|
|
self.assertEqual(text_only_provider.image_calls, [])
|
|
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
|
|
|
|
def test_upstream_failure_refunds_precharged_points_and_marks_call_failed(self):
|
|
self.provider.text_error = AiProviderError("provider timeout")
|
|
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/title",
|
|
{"prompt": "生成标题", "model": self.title_alias},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 502)
|
|
self.assertEqual(response.data["error"]["code"], "upstream_error")
|
|
self.wallet.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 100)
|
|
|
|
call = CallRecord.objects.get(user=self.user)
|
|
self.assertEqual(call.status, CallRecord.Status.FAILED)
|
|
self.assertIn("provider timeout", call.error_message)
|
|
self.assertEqual(
|
|
PointsLedger.objects.filter(
|
|
ref_call=call,
|
|
change_type=PointsLedger.ChangeType.CONSUME,
|
|
).count(),
|
|
1,
|
|
)
|
|
self.assertEqual(
|
|
PointsLedger.objects.filter(
|
|
ref_call=call,
|
|
change_type=PointsLedger.ChangeType.REFUND,
|
|
).count(),
|
|
1,
|
|
)
|
|
|
|
def test_provider_capability_error_returns_400_and_refunds_points(self):
|
|
self.provider.image_error = AiCapabilityError("input image is required")
|
|
|
|
response = self.post_with_provider(
|
|
"/api/v1/generate/image",
|
|
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 400)
|
|
self.assertEqual(response.data["error"]["code"], "bad_request")
|
|
self.wallet.refresh_from_db()
|
|
self.assertEqual(self.wallet.points_balance, 100)
|
|
call = CallRecord.objects.get(user=self.user)
|
|
self.assertEqual(call.status, CallRecord.Status.FAILED)
|
|
self.assertEqual(
|
|
PointsLedger.objects.filter(
|
|
ref_call=call,
|
|
change_type=PointsLedger.ChangeType.REFUND,
|
|
).count(),
|
|
1,
|
|
)
|