769 lines
26 KiB
Python
769 lines
26 KiB
Python
from __future__ import annotations
|
|
|
|
import base64
|
|
import binascii
|
|
import ipaddress
|
|
import socket
|
|
from dataclasses import dataclass, field
|
|
from time import perf_counter
|
|
from typing import Any, Callable, Mapping, Sequence
|
|
from urllib.parse import urljoin, urlsplit
|
|
|
|
import requests
|
|
from django.conf import settings
|
|
from rest_framework import status
|
|
|
|
from apps.ai.aliases import (
|
|
AliasResolutionError,
|
|
ModelCapabilityError,
|
|
REQUIRED_CAPABILITIES,
|
|
resolve_model_alias,
|
|
)
|
|
from apps.ai.providers import (
|
|
AiCapabilityError,
|
|
AiProviderError,
|
|
MultimodalImage,
|
|
get_provider,
|
|
)
|
|
from apps.billing.models import CallRecord, normalize_resolution
|
|
from apps.billing.pricing import NoPricingRuleError, calculate_points_cost
|
|
from apps.billing.services import (
|
|
InsufficientPointsError,
|
|
precharge_call,
|
|
refund_call_points,
|
|
mark_call_success,
|
|
)
|
|
|
|
from apps.moderation.services import moderate_prompt
|
|
|
|
from .errors import api_error
|
|
from .storage import save_generated_image
|
|
|
|
|
|
class ApiRequestError(ValueError):
|
|
def __init__(self, code: str, message: str, http_status: int):
|
|
self.code = code
|
|
self.message = message
|
|
self.http_status = http_status
|
|
super().__init__(message)
|
|
|
|
def as_response_data(self) -> dict:
|
|
return api_error(self.code, self.message)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ImageInput:
|
|
data: bytes
|
|
mime_type: str = "image/png"
|
|
filename: str = "image.png"
|
|
|
|
|
|
ImageUrlBuilder = Callable[[str], str]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class GenerationInput:
|
|
user: Any
|
|
api_key: Any
|
|
operation_type: str
|
|
prompt: str
|
|
alias: str | None = None
|
|
resolution: str = "1K"
|
|
parameters: Mapping[str, Any] = field(default_factory=dict)
|
|
image_url: str = ""
|
|
image_base64: str = ""
|
|
aspect_ratio: str = "1:1"
|
|
images: tuple[Mapping[str, Any], ...] = field(default_factory=tuple)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class PreparedGeneration:
|
|
user: Any
|
|
api_key: Any
|
|
operation_type: str
|
|
prompt: str
|
|
alias: str
|
|
resolution: str
|
|
parameters: dict[str, Any]
|
|
image_input: ImageInput | None
|
|
image_inputs: tuple[ImageInput, ...]
|
|
model_alias: Any
|
|
resolved_model: Any
|
|
provider: Any
|
|
points_cost: int
|
|
aspect_ratio: str = "1:1"
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class PrechargedGeneration:
|
|
prepared: PreparedGeneration
|
|
call_record: CallRecord
|
|
points_cost: int
|
|
points_balance_after_charge: int
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class GenerationResult:
|
|
operation_type: str
|
|
alias: str
|
|
model_used: str
|
|
points_cost: int
|
|
points_balance: int
|
|
call_record: CallRecord
|
|
titles: tuple[str, ...] = field(default_factory=tuple)
|
|
image_url: str = ""
|
|
text: str = ""
|
|
|
|
def as_response_data(self) -> dict[str, Any]:
|
|
common = {
|
|
"alias": self.alias,
|
|
"model_used": self.model_used,
|
|
"points_cost": self.points_cost,
|
|
"points_balance": self.points_balance,
|
|
"call_id": self.call_record.id,
|
|
}
|
|
if self.operation_type == CallRecord.OperationType.TITLE:
|
|
return {"titles": list(self.titles), **common}
|
|
if self.operation_type == CallRecord.OperationType.VISION:
|
|
return {"text": self.text, **common}
|
|
return {"image_url": self.image_url, **common}
|
|
|
|
|
|
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:
|
|
result = run_synchronous_generation(
|
|
GenerationInput(
|
|
user=user,
|
|
api_key=api_key,
|
|
operation_type=CallRecord.OperationType.TITLE,
|
|
prompt=request_data["prompt"],
|
|
alias=request_data.get("model") or None,
|
|
resolution=request_data.get("resolution") or "1K",
|
|
parameters=dict(request_data.get("parameters") or {}),
|
|
image_url=str(request_data.get("image_url") or ""),
|
|
image_base64=str(request_data.get("image_base64") or ""),
|
|
)
|
|
)
|
|
return result.as_response_data()
|
|
|
|
|
|
def generate_image_response(
|
|
*,
|
|
user,
|
|
api_key,
|
|
request_data: Mapping[str, Any],
|
|
image_url_builder: ImageUrlBuilder | None = None,
|
|
) -> dict:
|
|
result = run_synchronous_generation(
|
|
GenerationInput(
|
|
user=user,
|
|
api_key=api_key,
|
|
operation_type=CallRecord.OperationType.IMAGE,
|
|
prompt=request_data["prompt"],
|
|
alias=request_data.get("model") or None,
|
|
resolution=request_data.get("resolution") or "1K",
|
|
parameters=dict(request_data.get("parameters") or {}),
|
|
image_url=str(request_data.get("image_url") or ""),
|
|
image_base64=str(request_data.get("image_base64") or ""),
|
|
aspect_ratio=request_data.get("aspect_ratio") or "1:1",
|
|
),
|
|
image_url_builder=image_url_builder,
|
|
)
|
|
return result.as_response_data()
|
|
|
|
|
|
def analyze_images_response(*, user, api_key, request_data: Mapping[str, Any]) -> dict:
|
|
result = run_synchronous_generation(
|
|
GenerationInput(
|
|
user=user,
|
|
api_key=api_key,
|
|
operation_type=CallRecord.OperationType.VISION,
|
|
prompt=request_data["prompt"],
|
|
alias=request_data.get("model") or None,
|
|
parameters=dict(request_data.get("parameters") or {}),
|
|
images=tuple(dict(item) for item in request_data.get("images") or ()),
|
|
)
|
|
)
|
|
return result.as_response_data()
|
|
|
|
|
|
def run_synchronous_generation(
|
|
generation_input: GenerationInput,
|
|
*,
|
|
image_url_builder: ImageUrlBuilder | None = None,
|
|
) -> GenerationResult:
|
|
prepared = prepare_generation(generation_input)
|
|
precharged = precharge_generation(prepared)
|
|
return execute_precharged_generation(
|
|
precharged,
|
|
image_url_builder=image_url_builder,
|
|
)
|
|
|
|
|
|
def prepare_generation(generation_input: GenerationInput) -> PreparedGeneration:
|
|
operation_type = normalize_operation_type(generation_input.operation_type)
|
|
resolution = (
|
|
""
|
|
if operation_type == CallRecord.OperationType.VISION
|
|
else normalize_resolution(generation_input.resolution or "1K") or "1K"
|
|
)
|
|
parameters = dict(generation_input.parameters or {})
|
|
prompt = str(generation_input.prompt or "")
|
|
|
|
moderate_prompt_or_raise(
|
|
user=generation_input.user,
|
|
api_key=generation_input.api_key,
|
|
prompt=prompt,
|
|
)
|
|
if operation_type == CallRecord.OperationType.VISION:
|
|
image_input = None
|
|
image_inputs = load_vision_image_inputs(generation_input.images)
|
|
else:
|
|
image_input = load_image_input(
|
|
{
|
|
"image_base64": generation_input.image_base64,
|
|
"image_url": generation_input.image_url,
|
|
}
|
|
)
|
|
image_inputs = ()
|
|
|
|
model_alias = resolve_model_alias_or_raise(operation_type, generation_input.alias)
|
|
resolved_model = resolved_model_or_raise(model_alias)
|
|
provider = provider_or_raise(resolved_model)
|
|
ensure_provider_supports(provider, operation_type)
|
|
points_cost = calculate_points_cost_or_raise(
|
|
operation_type,
|
|
model_alias.alias,
|
|
resolution,
|
|
)
|
|
|
|
return PreparedGeneration(
|
|
user=generation_input.user,
|
|
api_key=generation_input.api_key,
|
|
operation_type=operation_type,
|
|
prompt=prompt,
|
|
alias=model_alias.alias,
|
|
resolution=resolution,
|
|
parameters=parameters,
|
|
image_input=image_input,
|
|
image_inputs=image_inputs,
|
|
model_alias=model_alias,
|
|
resolved_model=resolved_model,
|
|
provider=provider,
|
|
points_cost=points_cost,
|
|
aspect_ratio=generation_input.aspect_ratio or "1:1",
|
|
)
|
|
|
|
|
|
def precharge_generation(prepared: PreparedGeneration) -> PrechargedGeneration:
|
|
charge = precharge_or_raise(
|
|
user=prepared.user,
|
|
api_key=prepared.api_key,
|
|
operation_type=prepared.operation_type,
|
|
alias=prepared.alias,
|
|
model_used=prepared.resolved_model.model,
|
|
resolution=prepared.resolution,
|
|
prompt=prepared.prompt,
|
|
points_cost=prepared.points_cost,
|
|
)
|
|
return PrechargedGeneration(
|
|
prepared=prepared,
|
|
call_record=charge.call_record,
|
|
points_cost=charge.points_cost,
|
|
points_balance_after_charge=charge.balance_after,
|
|
)
|
|
|
|
|
|
def execute_precharged_generation(
|
|
precharged: PrechargedGeneration,
|
|
*,
|
|
image_url_builder: ImageUrlBuilder | None = None,
|
|
refund_on_failure: bool = True,
|
|
) -> GenerationResult:
|
|
prepared = precharged.prepared
|
|
try:
|
|
started = perf_counter()
|
|
if prepared.operation_type == CallRecord.OperationType.TITLE:
|
|
result = execute_title_generation(precharged, started)
|
|
elif prepared.operation_type == CallRecord.OperationType.VISION:
|
|
result = execute_vision_generation(precharged, started)
|
|
else:
|
|
result = execute_image_generation(
|
|
precharged,
|
|
started,
|
|
image_url_builder=image_url_builder,
|
|
)
|
|
except AiCapabilityError as exc:
|
|
if refund_on_failure:
|
|
refund_call_points(
|
|
precharged.call_record,
|
|
error_message=str(exc),
|
|
reason=f"Provider rejected the {prepared.operation_type} request.",
|
|
)
|
|
raise ApiRequestError(
|
|
"bad_request",
|
|
"请求参数不支持当前模型",
|
|
status.HTTP_400_BAD_REQUEST,
|
|
) from exc
|
|
except Exception as exc:
|
|
if refund_on_failure:
|
|
refund_call_points(
|
|
precharged.call_record,
|
|
error_message=str(exc),
|
|
reason=f"Upstream {prepared.operation_type} generation failed.",
|
|
)
|
|
raise upstream_error(exc) from exc
|
|
|
|
return result
|
|
|
|
|
|
def execute_title_generation(
|
|
precharged: PrechargedGeneration,
|
|
started: float,
|
|
) -> GenerationResult:
|
|
prepared = precharged.prepared
|
|
image_input = prepared.image_input
|
|
generation = prepared.provider.generate_text(
|
|
prepared.prompt,
|
|
prepared.resolved_model,
|
|
image=image_input.data if image_input else None,
|
|
image_mime_type=image_input.mime_type if image_input else "image/png",
|
|
resolution=prepared.resolution,
|
|
parameters=prepared.parameters,
|
|
)
|
|
latency_ms = elapsed_ms(started)
|
|
titles = list(generation.titles or ())
|
|
if not titles and generation.text:
|
|
titles = [generation.text]
|
|
result_summary = summarize_titles(titles, generation.text)
|
|
call_record = mark_call_success(
|
|
precharged.call_record,
|
|
result_summary=result_summary,
|
|
upstream_latency_ms=latency_ms,
|
|
)
|
|
return GenerationResult(
|
|
operation_type=prepared.operation_type,
|
|
alias=prepared.alias,
|
|
model_used=generation.model_used,
|
|
points_cost=precharged.points_cost,
|
|
points_balance=precharged.points_balance_after_charge,
|
|
call_record=call_record,
|
|
titles=tuple(titles),
|
|
)
|
|
|
|
|
|
def execute_image_generation(
|
|
precharged: PrechargedGeneration,
|
|
started: float,
|
|
*,
|
|
image_url_builder: ImageUrlBuilder | None = None,
|
|
) -> GenerationResult:
|
|
prepared = precharged.prepared
|
|
image_input = prepared.image_input
|
|
generation = prepared.provider.generate_image(
|
|
prepared.prompt,
|
|
prepared.resolved_model,
|
|
image=image_input.data if image_input else None,
|
|
image_mime_type=image_input.mime_type if image_input else "image/png",
|
|
image_filename=image_input.filename if image_input else "image.png",
|
|
resolution=prepared.resolution,
|
|
aspect_ratio=prepared.aspect_ratio,
|
|
parameters=prepared.parameters,
|
|
)
|
|
latency_ms = elapsed_ms(started)
|
|
image_url = save_generated_image(
|
|
generation.image,
|
|
url_builder=image_url_builder,
|
|
)
|
|
call_record = mark_call_success(
|
|
precharged.call_record,
|
|
result_ref=image_url,
|
|
result_summary=f"image_bytes={len(generation.image)}",
|
|
upstream_latency_ms=latency_ms,
|
|
)
|
|
return GenerationResult(
|
|
operation_type=prepared.operation_type,
|
|
alias=prepared.alias,
|
|
model_used=generation.model_used,
|
|
points_cost=precharged.points_cost,
|
|
points_balance=precharged.points_balance_after_charge,
|
|
call_record=call_record,
|
|
image_url=image_url,
|
|
)
|
|
|
|
|
|
def execute_vision_generation(
|
|
precharged: PrechargedGeneration,
|
|
started: float,
|
|
) -> GenerationResult:
|
|
prepared = precharged.prepared
|
|
generation = prepared.provider.analyze_images(
|
|
prepared.prompt,
|
|
prepared.resolved_model,
|
|
images=tuple(
|
|
MultimodalImage(data=image.data, mime_type=image.mime_type)
|
|
for image in prepared.image_inputs
|
|
),
|
|
parameters=prepared.parameters,
|
|
)
|
|
latency_ms = elapsed_ms(started)
|
|
text = str(generation.text or "").strip()
|
|
call_record = mark_call_success(
|
|
precharged.call_record,
|
|
result_summary=summarize_text(text),
|
|
upstream_latency_ms=latency_ms,
|
|
)
|
|
return GenerationResult(
|
|
operation_type=prepared.operation_type,
|
|
alias=prepared.alias,
|
|
model_used=generation.model_used,
|
|
points_cost=precharged.points_cost,
|
|
points_balance=precharged.points_balance_after_charge,
|
|
call_record=call_record,
|
|
text=text,
|
|
)
|
|
|
|
|
|
def normalize_operation_type(operation_type: str) -> str:
|
|
normalized = str(operation_type or "").strip()
|
|
if normalized not in {
|
|
CallRecord.OperationType.TITLE,
|
|
CallRecord.OperationType.IMAGE,
|
|
CallRecord.OperationType.VISION,
|
|
}:
|
|
raise ValueError(f"Unsupported generation operation type: {operation_type}")
|
|
return normalized
|
|
|
|
|
|
def moderate_prompt_or_raise(*, user, api_key, prompt: str) -> None:
|
|
outcome = moderate_prompt(user=user, api_key=api_key, prompt=prompt)
|
|
if outcome.blocked:
|
|
raise ApiRequestError("content_blocked", "输入内容未通过安全审核", status.HTTP_400_BAD_REQUEST)
|
|
|
|
|
|
def resolve_model_alias_or_raise(operation_type: str, alias: str | None):
|
|
try:
|
|
return resolve_model_alias(operation_type, alias)
|
|
except ModelCapabilityError as exc:
|
|
raise ApiRequestError(
|
|
"model_not_allowed",
|
|
"该模型不支持此操作",
|
|
status.HTTP_400_BAD_REQUEST,
|
|
) from exc
|
|
except AliasResolutionError as exc:
|
|
raise ApiRequestError(
|
|
"model_not_allowed",
|
|
"模型别名不可用或不支持此操作",
|
|
status.HTTP_400_BAD_REQUEST,
|
|
) from exc
|
|
|
|
|
|
def ensure_provider_supports(provider, operation_type: str) -> None:
|
|
required_capabilities = REQUIRED_CAPABILITIES[operation_type]
|
|
if not required_capabilities.issubset(provider.capabilities()):
|
|
raise ApiRequestError(
|
|
"model_not_allowed",
|
|
"该模型不支持此操作",
|
|
status.HTTP_400_BAD_REQUEST,
|
|
)
|
|
|
|
|
|
def resolved_model_or_raise(model_alias):
|
|
try:
|
|
return model_alias.ai_model.to_resolved_model()
|
|
except Exception as exc:
|
|
raise ApiRequestError(
|
|
"upstream_error",
|
|
"上游模型配置不可用",
|
|
status.HTTP_502_BAD_GATEWAY,
|
|
) from exc
|
|
|
|
|
|
def provider_or_raise(resolved_model):
|
|
try:
|
|
return get_provider(resolved_model.api_type, resolved_model.url)
|
|
except AiProviderError as exc:
|
|
raise ApiRequestError(
|
|
"upstream_error",
|
|
"上游模型配置不可用",
|
|
status.HTTP_502_BAD_GATEWAY,
|
|
) from exc
|
|
|
|
|
|
def calculate_points_cost_or_raise(
|
|
operation_type: str,
|
|
alias: str,
|
|
resolution: str,
|
|
) -> int:
|
|
try:
|
|
return calculate_points_cost(operation_type, alias, resolution)
|
|
except NoPricingRuleError as exc:
|
|
raise ApiRequestError(
|
|
"no_pricing_rule",
|
|
"未配置对应计费规则",
|
|
status.HTTP_400_BAD_REQUEST,
|
|
) from exc
|
|
|
|
|
|
def precharge_or_raise(**kwargs):
|
|
try:
|
|
return precharge_call(**kwargs)
|
|
except InsufficientPointsError as exc:
|
|
raise ApiRequestError(
|
|
"insufficient_points",
|
|
"点数不足,请先充值",
|
|
status.HTTP_402_PAYMENT_REQUIRED,
|
|
) from exc
|
|
|
|
|
|
def upstream_error(exc: Exception) -> ApiRequestError:
|
|
if isinstance(exc, ApiRequestError):
|
|
return exc
|
|
if isinstance(exc, requests.Timeout):
|
|
return ApiRequestError("upstream_timeout", "上游 AI 调用超时,已退回点数", status.HTTP_502_BAD_GATEWAY)
|
|
if isinstance(exc, AiProviderError | requests.RequestException | OSError):
|
|
return ApiRequestError("upstream_error", "上游 AI 调用失败,已退回点数", status.HTTP_502_BAD_GATEWAY)
|
|
return ApiRequestError("upstream_error", "生成失败,已退回点数", status.HTTP_502_BAD_GATEWAY)
|
|
|
|
|
|
def load_image_input(data: Mapping[str, Any]) -> ImageInput | None:
|
|
raw_base64 = str(data.get("image_base64") or "").strip()
|
|
if raw_base64:
|
|
return decode_image_input(raw_base64)
|
|
|
|
image_url = str(data.get("image_url") or "").strip()
|
|
if image_url:
|
|
return download_image_input(image_url)
|
|
|
|
return None
|
|
|
|
|
|
def load_vision_image_inputs(
|
|
items: Sequence[Mapping[str, Any]],
|
|
) -> tuple[ImageInput, ...]:
|
|
max_images = max(1, int(getattr(settings, "VISION_MAX_IMAGES", 8)))
|
|
max_image_bytes = max(
|
|
1,
|
|
int(getattr(settings, "VISION_MAX_IMAGE_BYTES", 10 * 1024 * 1024)),
|
|
)
|
|
max_total_bytes = max(
|
|
1,
|
|
int(getattr(settings, "VISION_MAX_TOTAL_BYTES", 32 * 1024 * 1024)),
|
|
)
|
|
if not items:
|
|
raise ApiRequestError("bad_request", "images 至少需要一张图片", status.HTTP_400_BAD_REQUEST)
|
|
if len(items) > max_images:
|
|
raise ApiRequestError(
|
|
"bad_request",
|
|
f"单次最多上传 {max_images} 张图片",
|
|
status.HTTP_400_BAD_REQUEST,
|
|
)
|
|
|
|
image_inputs = []
|
|
total_bytes = 0
|
|
for item in items:
|
|
raw_base64 = str(item.get("image_base64") or "").strip()
|
|
image_url = str(item.get("image_url") or "").strip()
|
|
if bool(raw_base64) == bool(image_url):
|
|
raise ApiRequestError(
|
|
"bad_request",
|
|
"每张图片必须且只能提供 image_url 或 image_base64",
|
|
status.HTTP_400_BAD_REQUEST,
|
|
)
|
|
image_input = (
|
|
decode_image_input(raw_base64, max_bytes=max_image_bytes)
|
|
if raw_base64
|
|
else download_image_input(image_url, max_bytes=max_image_bytes)
|
|
)
|
|
if not image_input.mime_type.lower().startswith("image/"):
|
|
raise ApiRequestError("bad_request", "图片格式无效", status.HTTP_400_BAD_REQUEST)
|
|
total_bytes += len(image_input.data)
|
|
if total_bytes > max_total_bytes:
|
|
raise ApiRequestError(
|
|
"bad_request",
|
|
"图片总大小超过限制",
|
|
status.HTTP_400_BAD_REQUEST,
|
|
)
|
|
image_inputs.append(image_input)
|
|
return tuple(image_inputs)
|
|
|
|
|
|
def decode_image_input(value: str, *, max_bytes: int | None = None) -> ImageInput:
|
|
mime_type = "image/png"
|
|
encoded = value
|
|
if value.startswith("data:"):
|
|
if ";base64," not in value:
|
|
raise ApiRequestError("bad_request", "image_base64 格式无效", status.HTTP_400_BAD_REQUEST)
|
|
prefix, encoded = value.split(",", 1)
|
|
mime_type = prefix[len("data:") :].split(";", 1)[0] or mime_type
|
|
if max_bytes is not None and len(encoded) > ((max_bytes + 2) // 3) * 4:
|
|
raise ApiRequestError("bad_request", "图片过大", status.HTTP_400_BAD_REQUEST)
|
|
try:
|
|
image = base64.b64decode(encoded, validate=True)
|
|
except (binascii.Error, ValueError) as exc:
|
|
raise ApiRequestError("bad_request", "image_base64 格式无效", status.HTTP_400_BAD_REQUEST) from exc
|
|
if not image:
|
|
raise ApiRequestError("bad_request", "image_base64 不能为空", status.HTTP_400_BAD_REQUEST)
|
|
if max_bytes is not None and len(image) > max_bytes:
|
|
raise ApiRequestError("bad_request", "图片过大", status.HTTP_400_BAD_REQUEST)
|
|
return ImageInput(data=image, mime_type=mime_type, filename=filename_for_mime(mime_type))
|
|
|
|
|
|
def download_image_input(url: str, *, max_bytes: int | None = None) -> 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, max_bytes=max_bytes)
|
|
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:
|
|
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
|
|
|
|
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, *, max_bytes: int | None = None) -> bytes:
|
|
byte_limit = max(
|
|
1,
|
|
int(
|
|
max_bytes
|
|
if max_bytes is not None
|
|
else getattr(settings, "IMAGE_URL_MAX_BYTES", 10 * 1024 * 1024)
|
|
),
|
|
)
|
|
content_length = response.headers.get("Content-Length")
|
|
if content_length:
|
|
try:
|
|
if int(content_length) > byte_limit:
|
|
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 > byte_limit:
|
|
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",
|
|
"image/png": "png",
|
|
"image/webp": "webp",
|
|
}.get(mime_type, "png")
|
|
return f"input.{extension}"
|
|
|
|
|
|
def summarize_titles(titles: list[str], text: str) -> str:
|
|
summary = " | ".join(titles[:3]) or str(text or "")
|
|
return summary[:500]
|
|
|
|
|
|
def summarize_text(text: str) -> str:
|
|
return str(text or "").strip()[:500]
|
|
|
|
|
|
def elapsed_ms(started: float) -> int:
|
|
return int((perf_counter() - started) * 1000)
|