472 lines
16 KiB
Python
472 lines
16 KiB
Python
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 (
|
|
AliasResolutionError,
|
|
ModelCapabilityError,
|
|
REQUIRED_CAPABILITIES,
|
|
resolve_model_alias,
|
|
)
|
|
from apps.ai.providers import AiCapabilityError, AiProviderError, 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"
|
|
|
|
|
|
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
|
|
resolution = normalize_resolution(request_data.get("resolution") or "1K") or "1K"
|
|
parameters = dict(request_data.get("parameters") or {})
|
|
moderate_prompt_or_raise(user=user, api_key=api_key, prompt=prompt)
|
|
image_input = load_image_input(request_data)
|
|
|
|
model_alias = resolve_model_alias_or_raise(CallRecord.OperationType.TITLE, alias)
|
|
resolved_model = resolved_model_or_raise(model_alias)
|
|
provider = provider_or_raise(resolved_model)
|
|
ensure_provider_supports(provider, CallRecord.OperationType.TITLE)
|
|
points_cost = calculate_points_cost_or_raise(
|
|
CallRecord.OperationType.TITLE,
|
|
model_alias.alias,
|
|
resolution,
|
|
)
|
|
|
|
charge = precharge_or_raise(
|
|
user=user,
|
|
api_key=api_key,
|
|
operation_type=CallRecord.OperationType.TITLE,
|
|
alias=model_alias.alias,
|
|
model_used=resolved_model.model,
|
|
resolution=resolution,
|
|
prompt=prompt,
|
|
points_cost=points_cost,
|
|
)
|
|
|
|
try:
|
|
started = perf_counter()
|
|
generation = provider.generate_text(
|
|
prompt,
|
|
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=resolution,
|
|
parameters=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(
|
|
charge.call_record,
|
|
result_summary=result_summary,
|
|
upstream_latency_ms=latency_ms,
|
|
)
|
|
except AiCapabilityError as exc:
|
|
refund_call_points(
|
|
charge.call_record,
|
|
error_message=str(exc),
|
|
reason="Provider rejected the title request.",
|
|
)
|
|
raise ApiRequestError("bad_request", "请求参数不支持当前模型", status.HTTP_400_BAD_REQUEST) from exc
|
|
except Exception as exc:
|
|
refund_call_points(
|
|
charge.call_record,
|
|
error_message=str(exc),
|
|
reason="Upstream title generation failed.",
|
|
)
|
|
raise upstream_error(exc) from exc
|
|
|
|
return {
|
|
"titles": titles,
|
|
"alias": model_alias.alias,
|
|
"model_used": generation.model_used,
|
|
"points_cost": points_cost,
|
|
"points_balance": charge.balance_after,
|
|
"call_id": call_record.id,
|
|
}
|
|
|
|
|
|
def generate_image_response(*, user, api_key, request, request_data: Mapping[str, Any]) -> dict:
|
|
prompt = request_data["prompt"]
|
|
alias = request_data.get("model") or None
|
|
resolution = normalize_resolution(request_data.get("resolution") or "1K") or "1K"
|
|
aspect_ratio = request_data.get("aspect_ratio") or "1:1"
|
|
parameters = dict(request_data.get("parameters") or {})
|
|
moderate_prompt_or_raise(user=user, api_key=api_key, prompt=prompt)
|
|
image_input = load_image_input(request_data)
|
|
|
|
model_alias = resolve_model_alias_or_raise(CallRecord.OperationType.IMAGE, alias)
|
|
resolved_model = resolved_model_or_raise(model_alias)
|
|
provider = provider_or_raise(resolved_model)
|
|
ensure_provider_supports(provider, CallRecord.OperationType.IMAGE)
|
|
points_cost = calculate_points_cost_or_raise(
|
|
CallRecord.OperationType.IMAGE,
|
|
model_alias.alias,
|
|
resolution,
|
|
)
|
|
|
|
charge = precharge_or_raise(
|
|
user=user,
|
|
api_key=api_key,
|
|
operation_type=CallRecord.OperationType.IMAGE,
|
|
alias=model_alias.alias,
|
|
model_used=resolved_model.model,
|
|
resolution=resolution,
|
|
prompt=prompt,
|
|
points_cost=points_cost,
|
|
)
|
|
|
|
try:
|
|
started = perf_counter()
|
|
generation = provider.generate_image(
|
|
prompt,
|
|
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=resolution,
|
|
aspect_ratio=aspect_ratio,
|
|
parameters=parameters,
|
|
)
|
|
latency_ms = elapsed_ms(started)
|
|
image_url = save_generated_image(generation.image, request=request)
|
|
call_record = mark_call_success(
|
|
charge.call_record,
|
|
result_ref=image_url,
|
|
result_summary=f"image_bytes={len(generation.image)}",
|
|
upstream_latency_ms=latency_ms,
|
|
)
|
|
except AiCapabilityError as exc:
|
|
refund_call_points(
|
|
charge.call_record,
|
|
error_message=str(exc),
|
|
reason="Provider rejected the image request.",
|
|
)
|
|
raise ApiRequestError("bad_request", "请求参数不支持当前模型", status.HTTP_400_BAD_REQUEST) from exc
|
|
except Exception as exc:
|
|
refund_call_points(
|
|
charge.call_record,
|
|
error_message=str(exc),
|
|
reason="Upstream image generation failed.",
|
|
)
|
|
raise upstream_error(exc) from exc
|
|
|
|
return {
|
|
"image_url": image_url,
|
|
"alias": model_alias.alias,
|
|
"model_used": generation.model_used,
|
|
"points_cost": points_cost,
|
|
"points_balance": charge.balance_after,
|
|
"call_id": call_record.id,
|
|
}
|
|
|
|
|
|
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_capability = REQUIRED_CAPABILITIES[operation_type]
|
|
if required_capability not in 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 decode_image_input(value: str) -> 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
|
|
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)
|
|
return ImageInput(data=image, mime_type=mime_type, filename=filename_for_mime(mime_type))
|
|
|
|
|
|
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:
|
|
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) -> 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",
|
|
"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 elapsed_ms(started: float) -> int:
|
|
return int((perf_counter() - started) * 1000)
|