Files
cmhub/apps/api/generation.py
T

470 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, 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)