Files
cmhub/apps/api/tests.py
T

1066 lines
40 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.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.views import ExternalApiView
from apps.ai.models import AiModel, ModelAlias
from apps.ai.providers import (
AiCapabilityError,
AiProviderError,
ImageGenerationResult,
TextGenerationResult,
)
from apps.billing.models import (
CallRecord,
ExchangeRate,
PointsLedger,
PricingRule,
RechargeOrder,
)
from apps.billing.payment_gateways import (
build_mock_alipay_signature,
build_mock_body_signature,
)
from apps.billing.services import RechargePayment
from apps.users.models import ApiKey
from apps.users.models import UserWallet
TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii")
class AuthenticatedEchoView(ExternalApiView):
def get(self, request):
return Response(
{
"user_id": request.user.id,
"api_key_id": request.auth.id,
}
)
class DefaultAuthProbeView(APIView):
def get(self, request):
return Response({"ok": True})
urlpatterns = [
path("api/test-auth/", AuthenticatedEchoView.as_view()),
path("api/default-auth/", DefaultAuthProbeView.as_view()),
]
@override_settings(ROOT_URLCONF=__name__)
class ApiKeyAuthenticationTests(TestCase):
url = "/api/test-auth/"
def setUp(self):
cache.clear()
suffix = uuid.uuid4().hex[:8]
self.user = get_user_model().objects.create_user(
username=f"api-user-{suffix}",
email=f"api-user-{suffix}@example.com",
password="password",
)
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="test")
self.client = APIClient()
def auth_header(self, raw_key: str | None = None) -> dict:
return {"HTTP_AUTHORIZATION": f"Bearer {raw_key or self.raw_key}"}
def test_external_api_view_only_uses_api_key_authentication(self):
self.assertEqual(AuthenticatedEchoView.authentication_classes, (ApiKeyAuthentication,))
def test_global_drf_default_does_not_accept_web_session_authentication(self):
self.client.force_login(self.user)
response = self.client.get("/api/default-auth/")
self.assertEqual(response.status_code, 403)
def test_valid_bearer_key_authenticates_user_and_api_key(self):
response = self.client.get(self.url, **self.auth_header())
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["user_id"], self.user.id)
self.assertEqual(response.data["api_key_id"], self.api_key.id)
self.api_key.refresh_from_db()
self.assertIsNotNone(self.api_key.last_used_at)
def test_missing_api_key_returns_401(self):
response = self.client.get(self.url)
self.assertEqual(response.status_code, 401)
self.assertEqual(response["WWW-Authenticate"], "Bearer")
self.assertEqual(response.data["error"]["code"], "unauthorized")
def test_invalid_api_key_returns_401(self):
response = self.client.get(self.url, **self.auth_header("sk_cmhub_invalid"))
self.assertEqual(response.status_code, 401)
self.assertEqual(response["WWW-Authenticate"], "Bearer")
self.assertEqual(response.data["error"]["code"], "unauthorized")
@override_settings(API_AUTH_FAILURE_THROTTLE_RATE="1/min")
def test_invalid_api_key_failures_are_throttled_by_ip(self):
first = self.client.get(
self.url,
**self.auth_header("sk_cmhub_invalid"),
REMOTE_ADDR="198.51.100.21",
)
second = self.client.get(
self.url,
**self.auth_header("sk_cmhub_invalid"),
REMOTE_ADDR="198.51.100.21",
)
self.assertEqual(first.status_code, 401)
self.assertEqual(second.status_code, 429)
self.assertEqual(second.data["error"]["code"], "rate_limited")
def test_malformed_authorization_header_returns_401(self):
response = self.client.get(self.url, HTTP_AUTHORIZATION=f"Token {self.raw_key}")
self.assertEqual(response.status_code, 401)
self.assertEqual(response.data["error"]["code"], "unauthorized")
def test_revoked_api_key_returns_403(self):
self.api_key.status = ApiKey.Status.REVOKED
self.api_key.save(update_fields=("status", "updated_at"))
response = self.client.get(self.url, **self.auth_header())
self.assertEqual(response.status_code, 403)
self.assertEqual(response.data["error"]["code"], "account_disabled")
def test_disabled_user_returns_403(self):
self.user.status = self.user.Status.DISABLED
self.user.save(update_fields=("status",))
response = self.client.get(self.url, **self.auth_header())
self.assertEqual(response.status_code, 403)
self.assertEqual(response.data["error"]["code"], "account_disabled")
def test_web_session_login_is_not_accepted_for_external_api(self):
self.client.force_login(self.user)
response = self.client.get(self.url)
self.assertEqual(response.status_code, 401)
self.assertEqual(response.data["error"]["code"], "unauthorized")
class BalanceApiTests(TestCase):
url = "/api/v1/balance"
def setUp(self):
suffix = uuid.uuid4().hex[:8]
self.user = get_user_model().objects.create_user(
username=f"balance-user-{suffix}",
email=f"balance-user-{suffix}@example.com",
password="password",
)
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="balance")
self.client = APIClient()
def auth_header(self, raw_key: str | None = None) -> dict:
return {"HTTP_AUTHORIZATION": f"Bearer {raw_key or self.raw_key}"}
def test_balance_returns_wallet_balance_matching_ledger_sum(self):
UserWallet.objects.create(user=self.user, points_balance=100)
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.RECHARGE,
points_delta=120,
balance_after=120,
ref_order_id=1,
)
call = CallRecord.objects.create(
user=self.user,
api_key=self.api_key,
operation_type=CallRecord.OperationType.TITLE,
alias="title-standard",
model_used="gpt-5.5",
points_cost=20,
status=CallRecord.Status.SUCCESS,
)
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.CONSUME,
points_delta=-20,
balance_after=100,
ref_call=call,
)
response = self.client.get(self.url, **self.auth_header())
ledger_sum = sum(
PointsLedger.objects.filter(user=self.user).values_list(
"points_delta",
flat=True,
)
)
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["user"], self.user.username)
self.assertEqual(response.data["points_balance"], 100)
self.assertEqual(response.data["points_balance"], ledger_sum)
def test_balance_returns_zero_without_creating_missing_wallet(self):
response = self.client.get(self.url, **self.auth_header())
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["points_balance"], 0)
self.assertFalse(UserWallet.objects.filter(user=self.user).exists())
def test_balance_does_not_accept_web_session_without_api_key(self):
self.client.force_login(self.user)
response = self.client.get(self.url)
self.assertEqual(response.status_code, 401)
self.assertEqual(response.data["error"]["code"], "unauthorized")
@override_settings(
PAYMENT_CALLBACK_MODE="mock",
PAYMENT_MOCK_CALLBACK_SECRET="test-payment-callback-secret",
)
class RechargeCallbackApiTests(TestCase):
wechat_url = "/api/v1/recharge/callback/wechat"
alipay_url = "/api/v1/recharge/callback/alipay"
def setUp(self):
suffix = uuid.uuid4().hex[:8]
self.user = get_user_model().objects.create_user(
username=f"recharge-user-{suffix}",
email=f"recharge-user-{suffix}@example.com",
password="password",
)
self.wallet = UserWallet.objects.create(user=self.user, points_balance=100)
self.client = APIClient(enforce_csrf_checks=True)
def create_order(
self,
*,
amount="20.00",
points_granted=200,
pay_method=RechargeOrder.PayMethod.WEIXIN,
) -> RechargeOrder:
return RechargeOrder.objects.create(
user=self.user,
order_no=f"R{uuid.uuid4().hex[:12]}",
amount_money=Decimal(amount),
pay_method=pay_method,
exchange_rate=Decimal("10.0000"),
points_granted=points_granted,
)
def signed_wechat_body(self, order, *, total_cents=2000):
payload = {
"event_type": "TRANSACTION.SUCCESS",
"resource": {
"trade_state": "SUCCESS",
"out_trade_no": order.order_no,
"transaction_id": "wx-txn-001",
"success_time": "2026-07-03T00:00:00+08:00",
"amount": {"total": total_cents},
},
}
body = json.dumps(payload, separators=(",", ":")).encode("utf-8")
return body, build_mock_body_signature(body)
def signed_alipay_payload(self, order, *, total_amount="20.00"):
payload = {
"trade_status": "TRADE_SUCCESS",
"out_trade_no": order.order_no,
"trade_no": "ali-txn-001",
"total_amount": total_amount,
"gmt_payment": "2026-07-03 00:00:00",
}
payload["sign"] = build_mock_alipay_signature(payload)
return payload
def test_wechat_callback_credits_once_and_is_csrf_exempt(self):
order = self.create_order(amount="20.00", points_granted=200)
body, signature = self.signed_wechat_body(order)
first = self.client.post(
self.wechat_url,
data=body,
content_type="application/json",
HTTP_WECHATPAY_SIGNATURE=signature,
)
second = self.client.post(
self.wechat_url,
data=body,
content_type="application/json",
HTTP_WECHATPAY_SIGNATURE=signature,
)
self.assertEqual(first.status_code, 200)
self.assertEqual(second.status_code, 200)
self.assertEqual(first.data["code"], "SUCCESS")
self.wallet.refresh_from_db()
order.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 300)
self.assertEqual(order.status, RechargeOrder.Status.PAID)
self.assertEqual(order.payment_txn_no, "wx-txn-001")
self.assertEqual(
PointsLedger.objects.filter(
user=self.user,
ref_order_id=order.id,
change_type=PointsLedger.ChangeType.RECHARGE,
).count(),
1,
)
def test_wechat_callback_rejects_bad_signature_without_crediting(self):
order = self.create_order(amount="20.00", points_granted=200)
body, _signature = self.signed_wechat_body(order)
response = self.client.post(
self.wechat_url,
data=body,
content_type="application/json",
HTTP_WECHATPAY_SIGNATURE="bad-signature",
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "signature_invalid")
self.wallet.refresh_from_db()
order.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
self.assertEqual(order.status, RechargeOrder.Status.PENDING)
self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists())
def test_wechat_callback_rejects_amount_mismatch_without_crediting(self):
order = self.create_order(amount="20.00", points_granted=200)
body, signature = self.signed_wechat_body(order, total_cents=1999)
response = self.client.post(
self.wechat_url,
data=body,
content_type="application/json",
HTTP_WECHATPAY_SIGNATURE=signature,
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "amount_mismatch")
self.wallet.refresh_from_db()
order.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
self.assertEqual(order.status, RechargeOrder.Status.PENDING)
self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists())
def test_alipay_callback_credits_once_returns_success_and_is_csrf_exempt(self):
order = self.create_order(
amount="20.00",
points_granted=200,
pay_method=RechargeOrder.PayMethod.ALIPAY,
)
payload = self.signed_alipay_payload(order)
first = self.client.post(self.alipay_url, data=payload)
second = self.client.post(self.alipay_url, data=payload)
self.assertEqual(first.status_code, 200)
self.assertEqual(second.status_code, 200)
self.assertEqual(first.content, b"success")
self.wallet.refresh_from_db()
order.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 300)
self.assertEqual(order.status, RechargeOrder.Status.PAID)
self.assertEqual(order.payment_txn_no, "ali-txn-001")
self.assertEqual(
PointsLedger.objects.filter(
ref_order_id=order.id,
change_type=PointsLedger.ChangeType.RECHARGE,
).count(),
1,
)
def test_alipay_callback_rejects_bad_signature_without_crediting(self):
order = self.create_order(
amount="20.00",
points_granted=200,
pay_method=RechargeOrder.PayMethod.ALIPAY,
)
payload = self.signed_alipay_payload(order)
payload["sign"] = "bad-signature"
response = self.client.post(self.alipay_url, data=payload)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.content, b"fail")
self.wallet.refresh_from_db()
order.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
self.assertEqual(order.status, RechargeOrder.Status.PENDING)
self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists())
@override_settings(
PAYMENT_CALLBACK_MODE="mock",
PAYMENT_MOCK_CALLBACK_SECRET="test-payment-callback-secret",
PAYMENT_QR_EXPIRES_MINUTES=15,
)
class RechargeCreateStatusApiTests(TestCase):
create_url = "/api/v1/recharge/create"
status_url = "/api/v1/recharge/status"
def setUp(self):
suffix = uuid.uuid4().hex[:8]
self.user = get_user_model().objects.create_user(
username=f"recharge-create-{suffix}",
email=f"recharge-create-{suffix}@example.com",
password="password",
)
self.other_user = get_user_model().objects.create_user(
username=f"recharge-other-{suffix}",
email=f"recharge-other-{suffix}@example.com",
password="password",
)
self.wallet = UserWallet.objects.create(user=self.user, points_balance=100)
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="recharge")
self.exchange_rate = ExchangeRate.objects.create(
currency="CNY",
points_per_unit=Decimal("10.0000"),
effective_from=timezone.now(),
)
self.client = APIClient()
def create_order(
self,
*,
user=None,
amount="20.00",
points_granted=200,
pay_method=RechargeOrder.PayMethod.WEIXIN,
):
return RechargeOrder.objects.create(
user=user or self.user,
order_no=f"R{uuid.uuid4().hex[:12]}",
amount_money=Decimal(amount),
pay_method=pay_method,
exchange_rate=Decimal("10.0000"),
points_granted=points_granted,
code_url=f"mockpay://{pay_method}/existing",
)
def test_recharge_create_requires_web_session_not_api_key(self):
response = self.client.post(
self.create_url,
{"amount": "20.00", "pay_method": "weixin"},
format="json",
HTTP_AUTHORIZATION=f"Bearer {self.raw_key}",
)
self.assertEqual(response.status_code, 403)
self.assertFalse(RechargeOrder.objects.filter(user=self.user).exists())
def test_recharge_create_is_session_authenticated_and_locks_quote(self):
self.client.force_login(self.user)
response = self.client.post(
self.create_url,
{"amount": "20.00", "pay_method": "weixin"},
format="json",
)
self.assertEqual(response.status_code, 201)
self.assertEqual(response.data["amount"], "20.00")
self.assertEqual(response.data["exchange_rate"], "10.0000")
self.assertEqual(response.data["points_granted"], 200)
self.assertEqual(response.data["pay_method"], RechargeOrder.PayMethod.WEIXIN)
self.assertEqual(response.data["status"], RechargeOrder.Status.PENDING)
self.assertTrue(response.data["code_url"].startswith("weixin://wxpay/cmhub-mock"))
self.assertIsNotNone(response.data["expires_at"])
order = RechargeOrder.objects.get(order_no=response.data["order_no"])
self.assertEqual(order.user, self.user)
self.assertEqual(order.exchange_rate, Decimal("10.0000"))
self.assertEqual(order.points_granted, 200)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists())
@override_settings(RECHARGE_MAX_AMOUNT_CNY="100.00")
def test_recharge_create_rejects_amount_above_configured_maximum(self):
self.client.force_login(self.user)
response = self.client.post(
self.create_url,
{"amount": "100.01", "pay_method": "weixin"},
format="json",
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "bad_request")
self.assertFalse(RechargeOrder.objects.filter(user=self.user).exists())
def test_recharge_create_supports_alipay_mock_qr_code(self):
self.client.force_login(self.user)
response = self.client.post(
self.create_url,
{"amount": "30.00", "pay_method": "alipay"},
format="json",
)
self.assertEqual(response.status_code, 201)
self.assertEqual(response.data["pay_method"], RechargeOrder.PayMethod.ALIPAY)
self.assertTrue(response.data["code_url"].startswith("https://qr.alipay.com/cmhub-mock"))
def test_recharge_create_enforces_csrf_for_real_session_clients(self):
csrf_client = APIClient(enforce_csrf_checks=True)
csrf_client.force_login(self.user)
response = csrf_client.post(
self.create_url,
{"amount": "20.00", "pay_method": "weixin"},
format="json",
)
self.assertEqual(response.status_code, 403)
self.assertFalse(RechargeOrder.objects.filter(user=self.user).exists())
def test_recharge_status_returns_pending_order_for_owner_only(self):
order = self.create_order()
self.client.force_login(self.user)
response = self.client.get(self.status_url, {"order_no": order.order_no})
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["order_no"], order.order_no)
self.assertEqual(response.data["status"], RechargeOrder.Status.PENDING)
self.client.force_login(self.other_user)
denied = self.client.get(self.status_url, {"order_no": order.order_no})
self.assertEqual(denied.status_code, 404)
self.assertEqual(denied.data["error"]["code"], "order_not_found")
def test_recharge_status_active_query_can_apply_paid_order_once(self):
order = self.create_order(amount="20.00", points_granted=200)
self.client.force_login(self.user)
def fake_query(queried_order):
return RechargePayment(
order_no=queried_order.order_no,
pay_method=queried_order.pay_method,
amount=queried_order.amount_money,
transaction_id="queried-txn-001",
paid_at=timezone.now(),
)
with patch("apps.api.views.query_payment_order", side_effect=fake_query) as query:
first = self.client.get(self.status_url, {"order_no": order.order_no})
second = self.client.get(self.status_url, {"order_no": order.order_no})
self.assertEqual(first.status_code, 200)
self.assertEqual(second.status_code, 200)
self.assertEqual(first.data["status"], RechargeOrder.Status.PAID)
self.assertEqual(second.data["status"], RechargeOrder.Status.PAID)
self.assertEqual(query.call_count, 1)
self.wallet.refresh_from_db()
order.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 300)
self.assertEqual(order.payment_txn_no, "queried-txn-001")
self.assertEqual(
PointsLedger.objects.filter(
ref_order_id=order.id,
change_type=PointsLedger.ChangeType.RECHARGE,
).count(),
1,
)
class FakeGenerationProvider:
def __init__(self, *, capabilities=None):
self._capabilities = set(capabilities or {"text", "image", "vision"})
self.text_calls = []
self.image_calls = []
self.text_error = None
self.image_error = None
def capabilities(self):
return set(self._capabilities)
def generate_text(self, prompt, model, **kwargs):
self.text_calls.append({"prompt": prompt, "model": model, **kwargs})
if self.text_error is not None:
raise self.text_error
return TextGenerationResult(
text="测试标题一",
titles=("测试标题一", "测试标题二"),
model_used=model.model,
raw={"secret": "SECRET_RAW_SHOULD_NOT_BE_STORED"},
)
def generate_image(self, prompt, model, **kwargs):
self.image_calls.append({"prompt": prompt, "model": model, **kwargs})
if self.image_error is not None:
raise self.image_error
return ImageGenerationResult(
image=b"generated-image-bytes",
model_used=model.model,
raw={"b64_json": "SECRET_RAW_SHOULD_NOT_BE_STORED"},
)
class FakeImageUrlResponse:
def __init__(self, *, status_code=200, headers=None, chunks=()):
self.status_code = status_code
self.headers = headers or {}
self._chunks = list(chunks)
self.closed = False
def raise_for_status(self):
if self.status_code >= 400:
raise requests.HTTPError("image_url request failed", response=self)
def iter_content(self, chunk_size=1):
for chunk in self._chunks:
yield chunk
def close(self):
self.closed = True
def dns_result(address: str):
return [(None, None, None, "", (address, 443))]
@override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY)
class GenerateApiTests(TestCase):
def setUp(self):
cache.clear()
suffix = uuid.uuid4().hex[:8]
self.media_dir = tempfile.TemporaryDirectory()
self.addCleanup(self.media_dir.cleanup)
media_override = override_settings(
MEDIA_ROOT=self.media_dir.name,
MEDIA_URL="/media/",
)
media_override.enable()
self.addCleanup(media_override.disable)
self.user = get_user_model().objects.create_user(
username=f"generate-user-{suffix}",
email=f"generate-user-{suffix}@example.com",
password="password",
)
self.wallet = UserWallet.objects.create(user=self.user, points_balance=100)
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="generate")
self.client = APIClient()
self.provider = FakeGenerationProvider()
self.title_model = self.create_ai_model(
name=f"title-model-{suffix}",
model=f"gpt-title-{suffix}",
capabilities=["text", "vision"],
)
self.image_model = self.create_ai_model(
name=f"image-model-{suffix}",
model=f"gpt-image-{suffix}",
capabilities=["image", "vision"],
)
self.title_alias = f"title-standard-{suffix}"
self.image_alias = f"image-hd-{suffix}"
ModelAlias.objects.create(
operation_type=ModelAlias.OperationType.TITLE,
alias=self.title_alias,
ai_model=self.title_model,
is_default=True,
)
ModelAlias.objects.create(
operation_type=ModelAlias.OperationType.IMAGE,
alias=self.image_alias,
ai_model=self.image_model,
is_default=True,
)
PricingRule.objects.create(
operation_type=CallRecord.OperationType.TITLE,
alias=self.title_alias,
points_cost=2,
)
PricingRule.objects.create(
operation_type=CallRecord.OperationType.IMAGE,
alias=self.image_alias,
resolution="1K",
points_cost=10,
)
def create_ai_model(self, *, name, model, capabilities):
ai_model = AiModel(
name=name,
url="https://api.example.test/v1",
model=model,
api_type=AiModel.ApiType.CHAT,
capabilities=capabilities,
)
ai_model.set_api_key("sk-test-secret")
ai_model.save()
return ai_model
def auth_header(self) -> dict:
return {"HTTP_AUTHORIZATION": f"Bearer {self.raw_key}"}
def post_with_provider(self, path, payload, provider=None):
with patch("apps.api.generation.get_provider", return_value=provider or self.provider):
return self.client.post(path, payload, format="json", **self.auth_header())
def assert_generation_not_charged(self):
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
self.assertFalse(PointsLedger.objects.filter(user=self.user).exists())
def test_generate_title_uses_default_alias_charges_points_and_writes_call_record(self):
response = self.post_with_provider(
"/api/v1/generate/title",
{
"prompt": "生成标题",
"resolution": "1k",
"parameters": {"temperature": 0.2, "model": "bad-overridden-model"},
},
)
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["titles"], ["测试标题一", "测试标题二"])
self.assertEqual(response.data["alias"], self.title_alias)
self.assertEqual(response.data["model_used"], self.title_model.model)
self.assertEqual(response.data["points_cost"], 2)
self.assertEqual(response.data["points_balance"], 98)
self.assertEqual(len(self.provider.text_calls), 1)
self.assertEqual(self.provider.text_calls[0]["model"].model, self.title_model.model)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 98)
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
self.assertEqual(call.api_key, self.api_key)
self.assertEqual(call.alias, self.title_alias)
self.assertEqual(call.model_used, self.title_model.model)
self.assertEqual(call.resolution, "1K")
self.assertNotIn("SECRET_RAW", call.result_summary)
self.assertEqual(
PointsLedger.objects.filter(
user=self.user,
change_type=PointsLedger.ChangeType.CONSUME,
).count(),
1,
)
def test_generate_image_stores_file_returns_url_and_does_not_store_raw_base64(self):
encoded = base64.b64encode(b"input-image").decode("ascii")
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成图片",
"model": self.image_alias,
"image_base64": f"data:image/png;base64,{encoded}",
"resolution": "1K",
"aspect_ratio": "1:1",
},
)
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["alias"], self.image_alias)
self.assertEqual(response.data["model_used"], self.image_model.model)
self.assertEqual(response.data["points_cost"], 10)
self.assertEqual(response.data["points_balance"], 90)
self.assertTrue(response.data["image_url"].startswith("http://testserver/media/"))
self.assertEqual(self.provider.image_calls[0]["image"], b"input-image")
media_relative_path = response.data["image_url"].split("/media/", 1)[1]
self.assertTrue((Path(self.media_dir.name) / media_relative_path).exists())
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
self.assertEqual(call.result_ref, response.data["image_url"])
self.assertEqual(call.result_summary, "image_bytes=21")
self.assertNotIn("SECRET_RAW", call.result_ref + call.result_summary)
def test_generate_image_downloads_safe_image_url(self):
response = FakeImageUrlResponse(
headers={"Content-Type": "image/jpeg"},
chunks=[b"remote-image"],
)
with (
patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("93.184.216.34")),
patch("apps.api.generation.requests.Session.get", return_value=response),
):
result = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成图片",
"model": self.image_alias,
"image_url": "https://safe.example.com/input.jpg",
"resolution": "1K",
"aspect_ratio": "1:1",
},
)
self.assertEqual(result.status_code, 200)
self.assertEqual(self.provider.image_calls[0]["image"], b"remote-image")
self.assertEqual(self.provider.image_calls[0]["image_mime_type"], "image/jpeg")
self.assertTrue(response.closed)
def test_image_url_rejects_loopback_address_without_fetch_or_charge(self):
with (
patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("127.0.0.1")),
patch("apps.api.generation.requests.Session.get") as image_get,
):
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成图片",
"model": self.image_alias,
"image_url": "http://127.0.0.1/private.png",
"resolution": "1K",
},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "bad_request")
image_get.assert_not_called()
self.assert_generation_not_charged()
def test_image_url_rejects_cloud_metadata_address_without_fetch_or_charge(self):
with (
patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("169.254.169.254")),
patch("apps.api.generation.requests.Session.get") as image_get,
):
response = self.post_with_provider(
"/api/v1/generate/title",
{
"prompt": "生成标题",
"model": self.title_alias,
"image_url": "http://169.254.169.254/latest/meta-data/",
},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "bad_request")
image_get.assert_not_called()
self.assert_generation_not_charged()
def test_image_url_rejects_redirect_to_private_address_without_charge(self):
def fake_getaddrinfo(host, port, *args, **kwargs):
if host == "safe.example.com":
return dns_result("93.184.216.34")
return dns_result("127.0.0.1")
redirect = FakeImageUrlResponse(
status_code=302,
headers={"Location": "http://127.0.0.1/private.png"},
)
with (
patch("apps.api.generation.socket.getaddrinfo", side_effect=fake_getaddrinfo),
patch("apps.api.generation.requests.Session.get", return_value=redirect) as image_get,
):
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成图片",
"model": self.image_alias,
"image_url": "https://safe.example.com/input.png",
"resolution": "1K",
},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "bad_request")
self.assertEqual(image_get.call_count, 1)
self.assertTrue(redirect.closed)
self.assert_generation_not_charged()
@override_settings(IMAGE_URL_MAX_BYTES=4)
def test_image_url_rejects_oversized_response_without_charge(self):
oversized = FakeImageUrlResponse(
headers={"Content-Type": "image/png"},
chunks=[b"1234", b"5"],
)
with (
patch("apps.api.generation.socket.getaddrinfo", return_value=dns_result("93.184.216.34")),
patch("apps.api.generation.requests.Session.get", return_value=oversized),
):
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成图片",
"model": self.image_alias,
"image_url": "https://safe.example.com/input.png",
"resolution": "1K",
},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "bad_request")
self.assertTrue(oversized.closed)
self.assert_generation_not_charged()
@override_settings(API_GENERATE_THROTTLE_RATE="1/min")
def test_generate_endpoint_is_throttled_by_api_key_without_extra_charge(self):
first = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": self.title_alias},
)
second = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": self.title_alias},
)
self.assertEqual(first.status_code, 200)
self.assertEqual(second.status_code, 429)
self.assertEqual(second.data["error"]["code"], "rate_limited")
self.assertEqual(len(self.provider.text_calls), 1)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 98)
self.assertEqual(CallRecord.objects.filter(user=self.user).count(), 1)
def test_insufficient_points_returns_402_without_calling_provider_or_writing_call(self):
self.wallet.points_balance = 1
self.wallet.save(update_fields=("points_balance", "updated_at"))
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": self.title_alias},
)
self.assertEqual(response.status_code, 402)
self.assertEqual(response.data["error"]["code"], "insufficient_points")
self.assertEqual(self.provider.text_calls, [])
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
self.assertFalse(PointsLedger.objects.filter(user=self.user).exists())
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 1)
def test_missing_pricing_rule_returns_400_without_charging(self):
PricingRule.objects.filter(alias=self.title_alias).delete()
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": self.title_alias},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "no_pricing_rule")
self.assertEqual(self.provider.text_calls, [])
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
def test_alias_capability_mismatch_returns_model_not_allowed_without_charging(self):
bad_alias = f"bad-title-{uuid.uuid4().hex[:8]}"
ModelAlias.objects.create(
operation_type=ModelAlias.OperationType.TITLE,
alias=bad_alias,
ai_model=self.image_model,
)
PricingRule.objects.create(
operation_type=CallRecord.OperationType.TITLE,
alias=bad_alias,
points_cost=2,
)
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": bad_alias},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "model_not_allowed")
self.assertEqual(self.provider.text_calls, [])
self.assertFalse(CallRecord.objects.filter(alias=bad_alias).exists())
def test_provider_capability_mismatch_returns_model_not_allowed_before_charging(self):
text_only_provider = FakeGenerationProvider(capabilities={"text"})
response = self.post_with_provider(
"/api/v1/generate/image",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
provider=text_only_provider,
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "model_not_allowed")
self.assertEqual(text_only_provider.image_calls, [])
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
def test_upstream_failure_refunds_precharged_points_and_marks_call_failed(self):
self.provider.text_error = AiProviderError("provider timeout")
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": self.title_alias},
)
self.assertEqual(response.status_code, 502)
self.assertEqual(response.data["error"]["code"], "upstream_error")
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
call = CallRecord.objects.get(user=self.user)
self.assertEqual(call.status, CallRecord.Status.FAILED)
self.assertIn("provider timeout", call.error_message)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.CONSUME,
).count(),
1,
)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)
def test_provider_capability_error_returns_400_and_refunds_points(self):
self.provider.image_error = AiCapabilityError("input image is required")
response = self.post_with_provider(
"/api/v1/generate/image",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "bad_request")
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
call = CallRecord.objects.get(user=self.user)
self.assertEqual(call.status, CallRecord.Status.FAILED)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)