feat: add generation API endpoints

This commit is contained in:
QiuSW
2026-07-02 22:41:37 +08:00
parent 47e6aed7a8
commit 0f7798ec2f
17 changed files with 851 additions and 19 deletions
+8 -3
View File
@@ -23,8 +23,8 @@ REQUIRED_CAPABILITIES = {
}
def resolve_alias(operation_type: str, alias: str | None = None) -> ResolvedModel:
"""Resolve an external capability alias to a provider-ready model config."""
def resolve_model_alias(operation_type: str, alias: str | None = None) -> ModelAlias:
"""Resolve an external capability alias to an active ModelAlias row."""
required_capability = REQUIRED_CAPABILITIES.get(operation_type)
if required_capability is None:
raise AliasResolutionError(f"unsupported operation_type: {operation_type}")
@@ -52,4 +52,9 @@ def resolve_alias(operation_type: str, alias: str | None = None) -> ResolvedMode
f"alias {model_alias.alias} maps to model {ai_model.name} without "
f"{required_capability} capability"
)
return ai_model.to_resolved_model()
return model_alias
def resolve_alias(operation_type: str, alias: str | None = None) -> ResolvedModel:
"""Resolve an external capability alias to a provider-ready model config."""
return resolve_model_alias(operation_type, alias).ai_model.to_resolved_model()
+344
View File
@@ -0,0 +1,344 @@
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)
+45
View File
@@ -0,0 +1,45 @@
from rest_framework import serializers
class GenerateTitleRequestSerializer(serializers.Serializer):
prompt = serializers.CharField(trim_whitespace=True, allow_blank=False)
model = serializers.CharField(
required=False,
allow_blank=True,
trim_whitespace=True,
max_length=64,
)
image_url = serializers.URLField(required=False, allow_blank=True)
image_base64 = serializers.CharField(required=False, allow_blank=True)
resolution = serializers.CharField(
required=False,
allow_blank=True,
trim_whitespace=True,
default="1K",
)
parameters = serializers.DictField(required=False, default=dict)
class GenerateImageRequestSerializer(serializers.Serializer):
prompt = serializers.CharField(trim_whitespace=True, allow_blank=False)
model = serializers.CharField(
required=False,
allow_blank=True,
trim_whitespace=True,
max_length=64,
)
image_url = serializers.URLField(required=False, allow_blank=True)
image_base64 = serializers.CharField(required=False, allow_blank=True)
resolution = serializers.CharField(
required=False,
allow_blank=True,
trim_whitespace=True,
default="1K",
)
aspect_ratio = serializers.CharField(
required=False,
allow_blank=True,
trim_whitespace=True,
default="1:1",
)
parameters = serializers.DictField(required=False, default=dict)
+21
View File
@@ -0,0 +1,21 @@
from __future__ import annotations
from uuid import uuid4
from django.core.files.base import ContentFile
from django.core.files.storage import default_storage
from django.utils import timezone
def save_generated_image(image: bytes, *, request=None) -> str:
today = timezone.now()
path = (
"generated/images/"
f"{today:%Y/%m/%d}/"
f"{uuid4().hex}.png"
)
saved_path = default_storage.save(path, ContentFile(image))
url = default_storage.url(saved_path)
if request is not None and url.startswith("/"):
return request.build_absolute_uri(url)
return url
+318
View File
@@ -1,5 +1,10 @@
import uuid
import base64
import tempfile
from pathlib import Path
from unittest.mock import patch
from cryptography.fernet import Fernet
from django.contrib.auth import get_user_model
from django.test import TestCase, override_settings
from django.urls import path
@@ -8,7 +13,19 @@ from rest_framework.test import APIClient
from apps.api.authentication import ApiKeyAuthentication
from apps.api.views import ExternalApiView
from apps.ai.models import AiModel, ModelAlias
from apps.ai.providers import (
AiCapabilityError,
AiProviderError,
ImageGenerationResult,
TextGenerationResult,
)
from apps.billing.models import CallRecord, PointsLedger, PricingRule
from apps.users.models import ApiKey
from apps.users.models import UserWallet
TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii")
class AuthenticatedEchoView(ExternalApiView):
@@ -101,3 +118,304 @@ class ApiKeyAuthenticationTests(TestCase):
self.assertEqual(response.status_code, 401)
self.assertEqual(response.data["error"]["code"], "unauthorized")
class FakeGenerationProvider:
def __init__(self, *, capabilities=None):
self._capabilities = set(capabilities or {"text", "image", "vision"})
self.text_calls = []
self.image_calls = []
self.text_error = None
self.image_error = None
def capabilities(self):
return set(self._capabilities)
def generate_text(self, prompt, model, **kwargs):
self.text_calls.append({"prompt": prompt, "model": model, **kwargs})
if self.text_error is not None:
raise self.text_error
return TextGenerationResult(
text="测试标题一",
titles=("测试标题一", "测试标题二"),
model_used=model.model,
raw={"secret": "SECRET_RAW_SHOULD_NOT_BE_STORED"},
)
def generate_image(self, prompt, model, **kwargs):
self.image_calls.append({"prompt": prompt, "model": model, **kwargs})
if self.image_error is not None:
raise self.image_error
return ImageGenerationResult(
image=b"generated-image-bytes",
model_used=model.model,
raw={"b64_json": "SECRET_RAW_SHOULD_NOT_BE_STORED"},
)
@override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY)
class GenerateApiTests(TestCase):
def setUp(self):
suffix = uuid.uuid4().hex[:8]
self.media_dir = tempfile.TemporaryDirectory()
self.addCleanup(self.media_dir.cleanup)
media_override = override_settings(
MEDIA_ROOT=self.media_dir.name,
MEDIA_URL="/media/",
)
media_override.enable()
self.addCleanup(media_override.disable)
self.user = get_user_model().objects.create_user(
username=f"generate-user-{suffix}",
email=f"generate-user-{suffix}@example.com",
password="password",
)
self.wallet = UserWallet.objects.create(user=self.user, points_balance=100)
self.api_key, self.raw_key = ApiKey.create_for_user(self.user, name="generate")
self.client = APIClient()
self.provider = FakeGenerationProvider()
self.title_model = self.create_ai_model(
name=f"title-model-{suffix}",
model=f"gpt-title-{suffix}",
capabilities=["text", "vision"],
)
self.image_model = self.create_ai_model(
name=f"image-model-{suffix}",
model=f"gpt-image-{suffix}",
capabilities=["image", "vision"],
)
self.title_alias = f"title-standard-{suffix}"
self.image_alias = f"image-hd-{suffix}"
ModelAlias.objects.create(
operation_type=ModelAlias.OperationType.TITLE,
alias=self.title_alias,
ai_model=self.title_model,
is_default=True,
)
ModelAlias.objects.create(
operation_type=ModelAlias.OperationType.IMAGE,
alias=self.image_alias,
ai_model=self.image_model,
is_default=True,
)
PricingRule.objects.create(
operation_type=CallRecord.OperationType.TITLE,
alias=self.title_alias,
points_cost=2,
)
PricingRule.objects.create(
operation_type=CallRecord.OperationType.IMAGE,
alias=self.image_alias,
resolution="1K",
points_cost=10,
)
def create_ai_model(self, *, name, model, capabilities):
ai_model = AiModel(
name=name,
url="https://api.example.test/v1",
model=model,
api_type=AiModel.ApiType.CHAT,
capabilities=capabilities,
)
ai_model.set_api_key("sk-test-secret")
ai_model.save()
return ai_model
def auth_header(self) -> dict:
return {"HTTP_AUTHORIZATION": f"Bearer {self.raw_key}"}
def post_with_provider(self, path, payload, provider=None):
with patch("apps.api.generation.get_provider", return_value=provider or self.provider):
return self.client.post(path, payload, format="json", **self.auth_header())
def test_generate_title_uses_default_alias_charges_points_and_writes_call_record(self):
response = self.post_with_provider(
"/api/v1/generate/title",
{
"prompt": "生成标题",
"resolution": "1k",
"parameters": {"temperature": 0.2, "model": "bad-overridden-model"},
},
)
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["titles"], ["测试标题一", "测试标题二"])
self.assertEqual(response.data["alias"], self.title_alias)
self.assertEqual(response.data["model_used"], self.title_model.model)
self.assertEqual(response.data["points_cost"], 2)
self.assertEqual(response.data["points_balance"], 98)
self.assertEqual(len(self.provider.text_calls), 1)
self.assertEqual(self.provider.text_calls[0]["model"].model, self.title_model.model)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 98)
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
self.assertEqual(call.api_key, self.api_key)
self.assertEqual(call.alias, self.title_alias)
self.assertEqual(call.model_used, self.title_model.model)
self.assertEqual(call.resolution, "1K")
self.assertNotIn("SECRET_RAW", call.result_summary)
self.assertEqual(
PointsLedger.objects.filter(
user=self.user,
change_type=PointsLedger.ChangeType.CONSUME,
).count(),
1,
)
def test_generate_image_stores_file_returns_url_and_does_not_store_raw_base64(self):
encoded = base64.b64encode(b"input-image").decode("ascii")
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成图片",
"model": self.image_alias,
"image_base64": f"data:image/png;base64,{encoded}",
"resolution": "1K",
"aspect_ratio": "1:1",
},
)
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["alias"], self.image_alias)
self.assertEqual(response.data["model_used"], self.image_model.model)
self.assertEqual(response.data["points_cost"], 10)
self.assertEqual(response.data["points_balance"], 90)
self.assertTrue(response.data["image_url"].startswith("http://testserver/media/"))
self.assertEqual(self.provider.image_calls[0]["image"], b"input-image")
media_relative_path = response.data["image_url"].split("/media/", 1)[1]
self.assertTrue((Path(self.media_dir.name) / media_relative_path).exists())
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
self.assertEqual(call.result_ref, response.data["image_url"])
self.assertEqual(call.result_summary, "image_bytes=21")
self.assertNotIn("SECRET_RAW", call.result_ref + call.result_summary)
def test_insufficient_points_returns_402_without_calling_provider_or_writing_call(self):
self.wallet.points_balance = 1
self.wallet.save(update_fields=("points_balance", "updated_at"))
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": self.title_alias},
)
self.assertEqual(response.status_code, 402)
self.assertEqual(response.data["error"]["code"], "insufficient_points")
self.assertEqual(self.provider.text_calls, [])
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
self.assertFalse(PointsLedger.objects.filter(user=self.user).exists())
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 1)
def test_missing_pricing_rule_returns_400_without_charging(self):
PricingRule.objects.filter(alias=self.title_alias).delete()
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": self.title_alias},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "no_pricing_rule")
self.assertEqual(self.provider.text_calls, [])
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
def test_alias_capability_mismatch_returns_model_not_allowed_without_charging(self):
bad_alias = f"bad-title-{uuid.uuid4().hex[:8]}"
ModelAlias.objects.create(
operation_type=ModelAlias.OperationType.TITLE,
alias=bad_alias,
ai_model=self.image_model,
)
PricingRule.objects.create(
operation_type=CallRecord.OperationType.TITLE,
alias=bad_alias,
points_cost=2,
)
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": bad_alias},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "model_not_allowed")
self.assertEqual(self.provider.text_calls, [])
self.assertFalse(CallRecord.objects.filter(alias=bad_alias).exists())
def test_provider_capability_mismatch_returns_model_not_allowed_before_charging(self):
text_only_provider = FakeGenerationProvider(capabilities={"text"})
response = self.post_with_provider(
"/api/v1/generate/image",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
provider=text_only_provider,
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "model_not_allowed")
self.assertEqual(text_only_provider.image_calls, [])
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
def test_upstream_failure_refunds_precharged_points_and_marks_call_failed(self):
self.provider.text_error = AiProviderError("provider timeout")
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": self.title_alias},
)
self.assertEqual(response.status_code, 502)
self.assertEqual(response.data["error"]["code"], "upstream_error")
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
call = CallRecord.objects.get(user=self.user)
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_provider_capability_error_returns_400_and_refunds_points(self):
self.provider.image_error = AiCapabilityError("input image is required")
response = self.post_with_provider(
"/api/v1/generate/image",
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "bad_request")
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
call = CallRecord.objects.get(user=self.user)
self.assertEqual(call.status, CallRecord.Status.FAILED)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=call,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)
+8
View File
@@ -0,0 +1,8 @@
from django.urls import path
from .views import GenerateImageView, GenerateTitleView
urlpatterns = [
path("v1/generate/title", GenerateTitleView.as_view(), name="api-generate-title"),
path("v1/generate/image", GenerateImageView.as_view(), name="api-generate-image"),
]
+50
View File
@@ -1,9 +1,20 @@
from rest_framework.exceptions import AuthenticationFailed
from rest_framework.permissions import IsAuthenticated
from rest_framework.response import Response
from rest_framework import status
from rest_framework.views import APIView
from apps.api.authentication import ApiKeyAuthentication
from apps.api.errors import api_error
from apps.api.generation import (
ApiRequestError,
generate_image_response,
generate_title_response,
)
from apps.api.serializers import (
GenerateImageRequestSerializer,
GenerateTitleRequestSerializer,
)
class ExternalApiView(APIView):
@@ -14,3 +25,42 @@ class ExternalApiView(APIView):
if request.authenticators and not request.successful_authenticator:
raise AuthenticationFailed(api_error("unauthorized", "缺失或无效 API Key"))
super().permission_denied(request, message=message, code=code)
class GenerateTitleView(ExternalApiView):
def post(self, request):
serializer = GenerateTitleRequestSerializer(data=request.data)
if not serializer.is_valid():
return Response(
api_error("bad_request", "参数错误"),
status=status.HTTP_400_BAD_REQUEST,
)
try:
data = generate_title_response(
user=request.user,
api_key=request.auth,
request_data=serializer.validated_data,
)
except ApiRequestError as exc:
return Response(exc.as_response_data(), status=exc.http_status)
return Response(data, status=status.HTTP_200_OK)
class GenerateImageView(ExternalApiView):
def post(self, request):
serializer = GenerateImageRequestSerializer(data=request.data)
if not serializer.is_valid():
return Response(
api_error("bad_request", "参数错误"),
status=status.HTTP_400_BAD_REQUEST,
)
try:
data = generate_image_response(
user=request.user,
api_key=request.auth,
request=request,
request_data=serializer.validated_data,
)
except ApiRequestError as exc:
return Response(exc.as_response_data(), status=exc.http_status)
return Response(data, status=status.HTTP_200_OK)