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