feat: harden phase 3 api security

This commit is contained in:
QiuSW
2026-07-03 10:34:37 +08:00
parent 1c5aa7845d
commit 0168aa30ba
19 changed files with 540 additions and 46 deletions
+219
View File
@@ -6,13 +6,16 @@ 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
@@ -52,8 +55,14 @@ class AuthenticatedEchoView(ExternalApiView):
)
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()),
]
@@ -62,6 +71,7 @@ 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}",
@@ -77,6 +87,13 @@ class ApiKeyAuthenticationTests(TestCase):
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())
@@ -101,6 +118,23 @@ class ApiKeyAuthenticationTests(TestCase):
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}")
@@ -468,6 +502,20 @@ class RechargeCreateStatusApiTests(TestCase):
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)
@@ -578,9 +626,33 @@ class FakeGenerationProvider:
)
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)
@@ -656,6 +728,12 @@ class GenerateApiTests(TestCase):
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",
@@ -723,6 +801,147 @@ class GenerateApiTests(TestCase):
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"))