Files
cmhub/apps/api/generation.py
T

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)