from __future__ import annotations import base64 import binascii from dataclasses import dataclass from time import perf_counter from typing import Any, Mapping import requests 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 .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" 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 {}) 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 {}) 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 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 try: response = session.get(url, timeout=(10, 60)) response.raise_for_status() except requests.RequestException as exc: raise ApiRequestError("bad_request", "无法读取 image_url", status.HTTP_400_BAD_REQUEST) from exc content_type = response.headers.get("Content-Type", "image/png").split(";", 1)[0] if not content_type.startswith("image/"): raise ApiRequestError("bad_request", "image_url 不是图片资源", status.HTTP_400_BAD_REQUEST) if not response.content: raise ApiRequestError("bad_request", "image_url 图片内容为空", status.HTTP_400_BAD_REQUEST) return ImageInput( data=response.content, mime_type=content_type, filename=filename_for_mime(content_type), ) 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)