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
+8 -6
View File
@@ -5,6 +5,7 @@ from rest_framework.authentication import BaseAuthentication, get_authorization_
from rest_framework.exceptions import AuthenticationFailed, PermissionDenied
from apps.api.errors import api_error
from apps.api.throttles import throttle_api_auth_failure
from apps.users.models import ApiKey
@@ -19,24 +20,24 @@ class ApiKeyAuthentication(BaseAuthentication):
try:
header = raw_header.decode("utf-8")
except UnicodeError as exc:
raise self.authentication_failed() from exc
raise self.authentication_failed(request) from exc
parts = header.split()
if len(parts) != 2 or parts[0].lower() != self.keyword.lower():
raise self.authentication_failed()
raise self.authentication_failed(request)
raw_key = parts[1]
if not raw_key:
raise self.authentication_failed()
raise self.authentication_failed(request)
key_hash = ApiKey.hash_key(raw_key)
try:
api_key = ApiKey.objects.select_related("user").get(key_hash=key_hash)
except ApiKey.DoesNotExist as exc:
raise self.authentication_failed() from exc
raise self.authentication_failed(request) from exc
if not api_key.matches_key(raw_key):
raise self.authentication_failed()
raise self.authentication_failed(request)
if not api_key.is_active_key:
raise PermissionDenied(
@@ -57,5 +58,6 @@ class ApiKeyAuthentication(BaseAuthentication):
return self.keyword
@staticmethod
def authentication_failed() -> AuthenticationFailed:
def authentication_failed(request) -> AuthenticationFailed:
throttle_api_auth_failure(request)
return AuthenticationFailed(api_error("unauthorized", "缺失或无效 API Key"))
+13
View File
@@ -0,0 +1,13 @@
from __future__ import annotations
from rest_framework.exceptions import Throttled
from rest_framework.views import exception_handler
from apps.api.errors import api_error
def api_exception_handler(exc, context):
response = exception_handler(exc, context)
if response is not None and isinstance(exc, Throttled):
response.data = api_error("rate_limited", "请求过于频繁,请稍后再试")
return response
+127 -12
View File
@@ -2,11 +2,15 @@ from __future__ import annotations
import base64
import binascii
import ipaddress
import socket
from dataclasses import dataclass
from time import perf_counter
from typing import Any, Mapping
from urllib.parse import urljoin, urlsplit
import requests
from django.conf import settings
from rest_framework import status
from apps.ai.aliases import (
@@ -47,6 +51,10 @@ class ImageInput:
filename: str = "image.png"
IMAGE_URL_ALLOWED_SCHEMES = {"http", "https"}
IMAGE_URL_CHUNK_SIZE = 64 * 1024
def generate_title_response(*, user, api_key, request_data: Mapping[str, Any]) -> dict:
prompt = request_data["prompt"]
alias = request_data.get("model") or None
@@ -309,23 +317,130 @@ def decode_image_input(value: str) -> ImageInput:
def download_image_input(url: str) -> ImageInput:
session = requests.Session()
session.trust_env = False
current_url = validated_image_url(url)
max_redirects = max(0, int(getattr(settings, "IMAGE_URL_MAX_REDIRECTS", 3)))
for redirect_count in range(max_redirects + 1):
try:
response = session.get(
current_url,
allow_redirects=False,
stream=True,
timeout=(
int(getattr(settings, "IMAGE_URL_CONNECT_TIMEOUT_SECONDS", 10)),
int(getattr(settings, "IMAGE_URL_READ_TIMEOUT_SECONDS", 60)),
),
)
except requests.RequestException as exc:
raise ApiRequestError("bad_request", "无法读取 image_url", status.HTTP_400_BAD_REQUEST) from exc
try:
if is_redirect_response(response):
if redirect_count >= max_redirects:
raise ApiRequestError("bad_request", "image_url 重定向次数过多", status.HTTP_400_BAD_REQUEST)
location = response.headers.get("Location", "")
if not location:
raise ApiRequestError("bad_request", "image_url 重定向无效", status.HTTP_400_BAD_REQUEST)
current_url = validated_image_url(urljoin(current_url, location))
continue
response.raise_for_status()
content_type = response.headers.get("Content-Type", "image/png").split(";", 1)[0].strip().lower()
if not content_type.startswith("image/"):
raise ApiRequestError("bad_request", "image_url 不是图片资源", status.HTTP_400_BAD_REQUEST)
image = read_limited_image_response(response)
if not image:
raise ApiRequestError("bad_request", "image_url 图片内容为空", status.HTTP_400_BAD_REQUEST)
return ImageInput(
data=image,
mime_type=content_type,
filename=filename_for_mime(content_type),
)
except requests.RequestException as exc:
raise ApiRequestError("bad_request", "无法读取 image_url", status.HTTP_400_BAD_REQUEST) from exc
finally:
close = getattr(response, "close", None)
if close:
close()
raise ApiRequestError("bad_request", "image_url 重定向次数过多", status.HTTP_400_BAD_REQUEST)
def validated_image_url(url: str) -> str:
try:
response = session.get(url, timeout=(10, 60))
response.raise_for_status()
except requests.RequestException as exc:
parsed = urlsplit(url)
port = parsed.port
except ValueError as exc:
raise ApiRequestError("bad_request", "image_url 地址无效", status.HTTP_400_BAD_REQUEST) from exc
scheme = parsed.scheme.lower()
if scheme not in IMAGE_URL_ALLOWED_SCHEMES or not parsed.hostname:
raise ApiRequestError("bad_request", "image_url 地址不允许", status.HTTP_400_BAD_REQUEST)
default_port = 443 if scheme == "https" else 80
validate_image_url_host(parsed.hostname, port or default_port)
return parsed.geturl()
def validate_image_url_host(hostname: str, port: int) -> None:
try:
resolved = socket.getaddrinfo(hostname, port, type=socket.SOCK_STREAM)
except socket.gaierror as exc:
raise ApiRequestError("bad_request", "无法读取 image_url", status.HTTP_400_BAD_REQUEST) from exc
content_type = response.headers.get("Content-Type", "image/png").split(";", 1)[0]
if not content_type.startswith("image/"):
raise ApiRequestError("bad_request", "image_url 不是图片资源", status.HTTP_400_BAD_REQUEST)
if not response.content:
raise ApiRequestError("bad_request", "image_url 图片内容为空", status.HTTP_400_BAD_REQUEST)
return ImageInput(
data=response.content,
mime_type=content_type,
filename=filename_for_mime(content_type),
addresses = {item[4][0] for item in resolved if item and item[4]}
if not addresses:
raise ApiRequestError("bad_request", "无法读取 image_url", status.HTTP_400_BAD_REQUEST)
for address in addresses:
if image_url_ip_is_blocked(address):
raise ApiRequestError("bad_request", "image_url 地址不允许", status.HTTP_400_BAD_REQUEST)
def image_url_ip_is_blocked(address: str) -> bool:
try:
ip = ipaddress.ip_address(address)
except ValueError:
return True
if ip.version == 6 and ip.ipv4_mapped is not None:
ip = ip.ipv4_mapped
return (
not ip.is_global
or ip.is_private
or ip.is_loopback
or ip.is_link_local
or ip.is_reserved
or ip.is_multicast
or ip.is_unspecified
)
def is_redirect_response(response) -> bool:
return 300 <= int(getattr(response, "status_code", 0)) < 400
def read_limited_image_response(response) -> bytes:
max_bytes = max(1, int(getattr(settings, "IMAGE_URL_MAX_BYTES", 10 * 1024 * 1024)))
content_length = response.headers.get("Content-Length")
if content_length:
try:
if int(content_length) > max_bytes:
raise ApiRequestError("bad_request", "image_url 图片过大", status.HTTP_400_BAD_REQUEST)
except ValueError:
pass
chunks = []
total = 0
for chunk in response.iter_content(chunk_size=IMAGE_URL_CHUNK_SIZE):
if not chunk:
continue
total += len(chunk)
if total > max_bytes:
raise ApiRequestError("bad_request", "image_url 图片过大", status.HTTP_400_BAD_REQUEST)
chunks.append(chunk)
return b"".join(chunks)
def filename_for_mime(mime_type: str) -> str:
extension = {
"image/jpeg": "jpg",
+7
View File
@@ -1,5 +1,6 @@
from decimal import Decimal
from django.conf import settings
from rest_framework import serializers
from apps.billing.models import RechargeOrder
@@ -57,6 +58,12 @@ class RechargeCreateRequestSerializer(serializers.Serializer):
)
pay_method = serializers.ChoiceField(choices=RechargeOrder.PayMethod.values)
def validate_amount(self, value):
max_amount = Decimal(str(settings.RECHARGE_MAX_AMOUNT_CNY))
if value > max_amount:
raise serializers.ValidationError(f"单笔充值金额不能超过 {max_amount:.2f} CNY")
return value
class RechargeStatusRequestSerializer(serializers.Serializer):
order_no = serializers.CharField(
+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"))
+46
View File
@@ -0,0 +1,46 @@
from __future__ import annotations
from django.conf import settings
from rest_framework.exceptions import Throttled
from rest_framework.throttling import SimpleRateThrottle
class SettingsRateThrottle(SimpleRateThrottle):
setting_name = ""
default_rate = ""
def get_rate(self):
return getattr(settings, self.setting_name, self.default_rate)
class GenerateRateThrottle(SettingsRateThrottle):
scope = "generate"
setting_name = "API_GENERATE_THROTTLE_RATE"
default_rate = "60/min"
def get_cache_key(self, request, view):
if request.auth is not None:
ident = f"key:{request.auth.pk}"
elif request.user and request.user.is_authenticated:
ident = f"user:{request.user.pk}"
else:
ident = self.get_ident(request)
return self.cache_format % {"scope": self.scope, "ident": ident}
class ApiAuthFailureThrottle(SettingsRateThrottle):
scope = "api_auth_failure"
setting_name = "API_AUTH_FAILURE_THROTTLE_RATE"
default_rate = "30/min"
def get_cache_key(self, request, view):
return self.cache_format % {
"scope": self.scope,
"ident": self.get_ident(request),
}
def throttle_api_auth_failure(request) -> None:
throttle = ApiAuthFailureThrottle()
if not throttle.allow_request(request, None):
raise Throttled(wait=throttle.wait())
+6
View File
@@ -24,6 +24,7 @@ from apps.api.serializers import (
RechargeCreateRequestSerializer,
RechargeStatusRequestSerializer,
)
from apps.api.throttles import GenerateRateThrottle, throttle_api_auth_failure
from apps.billing.models import RechargeOrder
from apps.billing.payment_gateways import (
PaymentOrderCreateError,
@@ -56,11 +57,14 @@ class ExternalApiView(APIView):
def permission_denied(self, request, message=None, code=None):
if request.authenticators and not request.successful_authenticator:
throttle_api_auth_failure(request)
raise AuthenticationFailed(api_error("unauthorized", "缺失或无效 API Key"))
super().permission_denied(request, message=message, code=code)
class GenerateTitleView(ExternalApiView):
throttle_classes = (GenerateRateThrottle,)
def post(self, request):
serializer = GenerateTitleRequestSerializer(data=request.data)
if not serializer.is_valid():
@@ -80,6 +84,8 @@ class GenerateTitleView(ExternalApiView):
class GenerateImageView(ExternalApiView):
throttle_classes = (GenerateRateThrottle,)
def post(self, request):
serializer = GenerateImageRequestSerializer(data=request.data)
if not serializer.is_valid():