feat: harden phase 3 api security
This commit is contained in:
@@ -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"))
|
||||
|
||||
Reference in New Issue
Block a user