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)