493 lines
19 KiB
Python
493 lines
19 KiB
Python
import uuid
|
|
import base64
|
|
import tempfile
|
|
from pathlib import Path
|
|
from unittest.mock import patch
|
|
|
|
from cryptography.fernet import Fernet
|
|
from django.contrib.auth import get_user_model
|
|
from django.test import TestCase, override_settings
|
|
from django.urls import path
|
|
from rest_framework.response import Response
|
|
from rest_framework.test import APIClient
|
|
|
|
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, PointsLedger, PricingRule
|
|
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,
|
|
}
|
|
)
|
|
|
|
|
|
urlpatterns = [
|
|
path("api/test-auth/", AuthenticatedEchoView.as_view()),
|
|
]
|
|
|
|
|
|
@override_settings(ROOT_URLCONF=__name__)
|
|
class ApiKeyAuthenticationTests(TestCase):
|
|
url = "/api/test-auth/"
|
|
|
|
def setUp(self):
|
|
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_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")
|
|
|
|
def test_malformed_authorization_header_returns_401(self):
|
|
response = self.client.get(self.url, HTTP_AUTHORIZATION=f"Token {self.raw_key}")
|
|
|
|
self.assertEqual(response.status_code, 401)
|
|
self.assertEqual(response.data["error"]["code"], "unauthorized")
|
|
|
|
def test_revoked_api_key_returns_403(self):
|
|
self.api_key.status = ApiKey.Status.REVOKED
|
|
self.api_key.save(update_fields=("status", "updated_at"))
|
|
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
self.assertEqual(response.status_code, 403)
|
|
self.assertEqual(response.data["error"]["code"], "account_disabled")
|
|
|
|
def test_disabled_user_returns_403(self):
|
|
self.user.status = self.user.Status.DISABLED
|
|
self.user.save(update_fields=("status",))
|
|
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
self.assertEqual(response.status_code, 403)
|
|
self.assertEqual(response.data["error"]["code"], "account_disabled")
|
|
|
|
def test_web_session_login_is_not_accepted_for_external_api(self):
|
|
self.client.force_login(self.user)
|
|
|
|
response = self.client.get(self.url)
|
|
|
|
self.assertEqual(response.status_code, 401)
|
|
self.assertEqual(response.data["error"]["code"], "unauthorized")
|
|
|
|
|
|
class BalanceApiTests(TestCase):
|
|
url = "/api/v1/balance"
|
|
|
|
def setUp(self):
|
|
suffix = uuid.uuid4().hex[:8]
|
|
self.user = get_user_model().objects.create_user(
|
|
username=f"balance-user-{suffix}",
|
|
email=f"balance-user-{suffix}@example.com",
|
|
password="password",
|
|
)
|
|
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="balance")
|
|
self.client = APIClient()
|
|
|
|
def auth_header(self, raw_key: str | None = None) -> dict:
|
|
return {"HTTP_AUTHORIZATION": f"Bearer {raw_key or self.raw_key}"}
|
|
|
|
def test_balance_returns_wallet_balance_matching_ledger_sum(self):
|
|
UserWallet.objects.create(user=self.user, points_balance=100)
|
|
PointsLedger.objects.create(
|
|
user=self.user,
|
|
change_type=PointsLedger.ChangeType.RECHARGE,
|
|
points_delta=120,
|
|
balance_after=120,
|
|
ref_order_id=1,
|
|
)
|
|
call = CallRecord.objects.create(
|
|
user=self.user,
|
|
api_key=self.api_key,
|
|
operation_type=CallRecord.OperationType.TITLE,
|
|
alias="title-standard",
|
|
model_used="gpt-5.5",
|
|
points_cost=20,
|
|
status=CallRecord.Status.SUCCESS,
|
|
)
|
|
PointsLedger.objects.create(
|
|
user=self.user,
|
|
change_type=PointsLedger.ChangeType.CONSUME,
|
|
points_delta=-20,
|
|
balance_after=100,
|
|
ref_call=call,
|
|
)
|
|
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
ledger_sum = sum(
|
|
PointsLedger.objects.filter(user=self.user).values_list(
|
|
"points_delta",
|
|
flat=True,
|
|
)
|
|
)
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["user"], self.user.username)
|
|
self.assertEqual(response.data["points_balance"], 100)
|
|
self.assertEqual(response.data["points_balance"], ledger_sum)
|
|
|
|
def test_balance_returns_zero_without_creating_missing_wallet(self):
|
|
response = self.client.get(self.url, **self.auth_header())
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(response.data["points_balance"], 0)
|
|
self.assertFalse(UserWallet.objects.filter(user=self.user).exists())
|
|
|
|
def test_balance_does_not_accept_web_session_without_api_key(self):
|
|
self.client.force_login(self.user)
|
|
|
|
response = self.client.get(self.url)
|
|
|
|
self.assertEqual(response.status_code, 401)
|
|
self.assertEqual(response.data["error"]["code"], "unauthorized")
|
|
|
|
|
|
class 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"},
|
|
)
|
|
|
|
|
|
@override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY)
|
|
class GenerateApiTests(TestCase):
|
|
def setUp(self):
|
|
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 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_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,
|
|
)
|