refactor: extract generation core service

This commit is contained in:
QiuSW
2026-07-08 21:24:59 +08:00
parent 12717d09e8
commit c98f713762
11 changed files with 438 additions and 160 deletions
+301 -142
View File
@@ -4,9 +4,9 @@ import base64
import binascii
import ipaddress
import socket
from dataclasses import dataclass
from dataclasses import dataclass, field
from time import perf_counter
from typing import Any, Mapping
from typing import Any, Callable, Mapping
from urllib.parse import urljoin, urlsplit
import requests
@@ -53,157 +53,316 @@ class ImageInput:
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"
@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
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 = ""
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}
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:
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)
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()
model_alias = resolve_model_alias_or_raise(CallRecord.OperationType.TITLE, alias)
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 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 = 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,
)
image_input = load_image_input(
{
"image_base64": generation_input.image_base64,
"image_url": generation_input.image_url,
}
)
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, CallRecord.OperationType.TITLE)
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,
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,
) -> GenerationResult:
prepared = precharged.prepared
try:
started = perf_counter()
if prepared.operation_type == CallRecord.OperationType.TITLE:
result = execute_title_generation(precharged, started)
else:
result = execute_image_generation(
precharged,
started,
image_url_builder=image_url_builder,
)
except AiCapabilityError as exc:
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:
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 normalize_operation_type(operation_type: str) -> str:
normalized = str(operation_type or "").strip()
if normalized not in {
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,
}
}:
raise ValueError(f"Unsupported generation operation type: {operation_type}")
return normalized
def moderate_prompt_or_raise(*, user, api_key, prompt: str) -> None:
+9 -1
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
from typing import Callable
from uuid import uuid4
from django.core.files.base import ContentFile
@@ -7,7 +8,12 @@ from django.core.files.storage import default_storage
from django.utils import timezone
def save_generated_image(image: bytes, *, request=None) -> str:
def save_generated_image(
image: bytes,
*,
url_builder: Callable[[str], str] | None = None,
request=None,
) -> str:
today = timezone.now()
path = (
"generated/images/"
@@ -16,6 +22,8 @@ def save_generated_image(image: bytes, *, request=None) -> str:
)
saved_path = default_storage.save(path, ContentFile(image))
url = default_storage.url(saved_path)
if url_builder is not None and url.startswith("/"):
return url_builder(url)
if request is not None and url.startswith("/"):
return request.build_absolute_uri(url)
return url
+81
View File
@@ -20,6 +20,14 @@ from rest_framework.test import APIClient
from rest_framework.views import APIView
from apps.api.authentication import ApiKeyAuthentication
from apps.api.generation import (
ApiRequestError,
GenerationInput,
execute_precharged_generation,
precharge_generation,
prepare_generation,
run_synchronous_generation,
)
from apps.api.throttles import GenerateRateThrottle
from apps.api.views import ClientLatestReleaseView, ExternalApiView, ModelsView
from apps.ai.models import AiModel, ModelAlias
@@ -1273,6 +1281,34 @@ class GenerateApiTests(TestCase):
self.assertEqual(call.result_summary, "image_bytes=21")
self.assertNotIn("SECRET_RAW", call.result_ref + call.result_summary)
def test_generation_core_saves_image_with_url_builder_without_request(self):
encoded = base64.b64encode(b"input-image").decode("ascii")
with patch("apps.api.generation.get_provider", return_value=self.provider):
result = run_synchronous_generation(
GenerationInput(
user=self.user,
api_key=self.api_key,
operation_type=CallRecord.OperationType.IMAGE,
prompt="生成图片",
alias=self.image_alias,
resolution="1K",
image_base64=f"data:image/png;base64,{encoded}",
),
image_url_builder=lambda url: f"https://cdn.example.test{url}",
)
self.assertEqual(result.operation_type, CallRecord.OperationType.IMAGE)
self.assertTrue(result.image_url.startswith("https://cdn.example.test/media/"))
self.assertEqual(result.as_response_data()["image_url"], result.image_url)
self.assertEqual(self.provider.image_calls[0]["image"], b"input-image")
call = CallRecord.objects.get(pk=result.call_record.id)
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
self.assertEqual(call.result_ref, result.image_url)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 90)
def test_generate_image_downloads_safe_image_url(self):
response = FakeImageUrlResponse(
headers={"Content-Type": "image/jpeg"},
@@ -1514,6 +1550,51 @@ class GenerateApiTests(TestCase):
1,
)
def test_precharged_generation_stage_refunds_on_upstream_failure(self):
self.provider.text_error = AiProviderError("provider timeout")
with patch("apps.api.generation.get_provider", return_value=self.provider):
prepared = prepare_generation(
GenerationInput(
user=self.user,
api_key=self.api_key,
operation_type=CallRecord.OperationType.TITLE,
prompt="生成标题",
alias=self.title_alias,
resolution="1K",
)
)
precharged = precharge_generation(prepared)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 98)
self.assertEqual(precharged.call_record.status, CallRecord.Status.PENDING)
with self.assertRaises(ApiRequestError) as captured:
execute_precharged_generation(precharged)
self.assertEqual(captured.exception.code, "upstream_error")
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
call = CallRecord.objects.get(pk=precharged.call_record.id)
self.assertEqual(call.status, CallRecord.Status.FAILED)
self.assertIn("provider timeout", call.error_message)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.CONSUME,
).count(),
1,
)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)
def test_image_upstream_timeout_refunds_precharged_points_and_marks_call_failed(self):
self.provider.image_error = requests.Timeout("image upstream deadline exceeded")
+1 -1
View File
@@ -99,8 +99,8 @@ class GenerateImageView(ExternalApiView):
data = generate_image_response(
user=request.user,
api_key=request.auth,
request=request,
request_data=serializer.validated_data,
image_url_builder=request.build_absolute_uri,
)
except ApiRequestError as exc:
return Response(exc.as_response_data(), status=exc.http_status)