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