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] IMAGE_INPUT_ROLE_INSTRUCTIONS = ( "图片角色规则(必须遵守,优先于用户关于图片角色的要求):\n" "- 第 1 张图片是主商品图,必须优先保留其商品主体、外观和关键细节。\n" "- 第 2 张及之后的图片仅作为风格、构图、场景或排版参考," "不得用参考图商品替换主图商品。" ) @dataclass(frozen=True) class GenerationInput: user: Any api_key: Any operation_type: str prompt: str client_device: Any | None = None 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 client_device: Any | None 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, client_device=None, request_data: Mapping[str, Any]) -> dict: result = run_synchronous_generation( GenerationInput( user=user, api_key=api_key, client_device=client_device, 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, client_device=None, request_data: Mapping[str, Any], image_url_builder: ImageUrlBuilder | None = None, ) -> dict: result = run_synchronous_generation( GenerationInput( user=user, api_key=api_key, client_device=client_device, 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", images=tuple(dict(item) for item in request_data.get("images") or ()), ), image_url_builder=image_url_builder, ) return result.as_response_data() def analyze_images_response(*, user, api_key, client_device=None, request_data: Mapping[str, Any]) -> dict: result = run_synchronous_generation( GenerationInput( user=user, api_key=api_key, client_device=client_device, 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) elif operation_type == CallRecord.OperationType.IMAGE: image_inputs = load_image_generation_inputs( generation_input.images, image_base64=generation_input.image_base64, image_url=generation_input.image_url, ) image_input = image_inputs[0] if image_inputs else None 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, client_device=generation_input.client_device, 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, client_device=prepared.client_device, 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 image_inputs = prepared.image_inputs generation = prepared.provider.generate_image( image_generation_prompt(prepared.prompt, len(image_inputs)), 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", images=tuple( MultimodalImage( data=item.data, mime_type=item.mime_type, filename=item.filename, ) for item in image_inputs ), 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: return load_image_input_with_limit(data) def load_image_input_with_limit( data: Mapping[str, Any], *, max_bytes: int | None = None, ) -> ImageInput | None: raw_base64 = str(data.get("image_base64") or "").strip() if raw_base64: return decode_image_input(raw_base64, max_bytes=max_bytes) image_url = str(data.get("image_url") or "").strip() if image_url: return download_image_input(image_url, max_bytes=max_bytes) return None def load_vision_image_inputs( items: Sequence[Mapping[str, Any]], ) -> tuple[ImageInput, ...]: return load_ordered_image_inputs( items, 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)), ), ) def load_image_generation_inputs( items: Sequence[Mapping[str, Any]], *, image_base64: str = "", image_url: str = "", ) -> tuple[ImageInput, ...]: max_image_bytes = max( 1, int(getattr(settings, "IMAGE_MAX_INPUT_IMAGE_BYTES", 10 * 1024 * 1024)), ) if items: return load_ordered_image_inputs( items, max_images=max(1, int(getattr(settings, "IMAGE_MAX_INPUT_IMAGES", 8))), max_image_bytes=max_image_bytes, max_total_bytes=max( 1, int(getattr(settings, "IMAGE_MAX_INPUT_TOTAL_BYTES", 32 * 1024 * 1024)), ), ) image_input = load_image_input_with_limit( {"image_base64": image_base64, "image_url": image_url}, max_bytes=max_image_bytes, ) return (image_input,) if image_input is not None else () def load_ordered_image_inputs( items: Sequence[Mapping[str, Any]], *, max_images: int, max_image_bytes: int, max_total_bytes: int, ) -> tuple[ImageInput, ...]: 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 image_generation_prompt(prompt: str, image_count: int) -> str: if image_count < 1: return prompt return f"{prompt}\n\n{IMAGE_INPUT_ROLE_INSTRUCTIONS}" 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=effective_image_url_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 effective_image_url_max_bytes(max_bytes: int | None) -> int: url_limit = max(1, int(getattr(settings, "IMAGE_URL_MAX_BYTES", 10 * 1024 * 1024))) if max_bytes is None: return url_limit return min(url_limit, max(1, int(max_bytes))) 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)