feat: add image upstream deadline
This commit is contained in:
@@ -20,6 +20,7 @@ from .utils import (
|
||||
extract_image_from_response,
|
||||
extract_text_from_response,
|
||||
extract_titles_from_response,
|
||||
image_request_timeout,
|
||||
image_bytes_to_data_url,
|
||||
normalize_api_url,
|
||||
request_timeout,
|
||||
@@ -79,6 +80,16 @@ class BaseHttpProvider:
|
||||
def _read_timeout(self, model: ResolvedModel, resolution: str) -> int:
|
||||
return self._timeout(model, resolution)[1]
|
||||
|
||||
def _image_timeout(self, model: ResolvedModel, resolution: str) -> tuple[int, int]:
|
||||
return image_request_timeout(
|
||||
model.connect_timeout_seconds,
|
||||
model.timeout_seconds,
|
||||
resolution,
|
||||
)
|
||||
|
||||
def _image_read_timeout(self, model: ResolvedModel, resolution: str) -> int:
|
||||
return self._image_timeout(model, resolution)[1]
|
||||
|
||||
|
||||
class ChatCompletionsProvider(BaseHttpProvider):
|
||||
def capabilities(self) -> set[str]:
|
||||
@@ -142,14 +153,14 @@ class ChatCompletionsProvider(BaseHttpProvider):
|
||||
url,
|
||||
headers=self._headers(model, json=True),
|
||||
json=payload,
|
||||
timeout=self._timeout(model, resolution),
|
||||
timeout=self._image_timeout(model, resolution),
|
||||
)
|
||||
response.raise_for_status()
|
||||
raw = response.json()
|
||||
image_bytes = extract_image_from_response(
|
||||
raw,
|
||||
session=self.session,
|
||||
timeout=self._read_timeout(model, resolution),
|
||||
timeout=self._image_read_timeout(model, resolution),
|
||||
)
|
||||
if not image_bytes:
|
||||
raise AiResponseParseError("AI response did not contain an image")
|
||||
@@ -217,14 +228,14 @@ class GeminiProvider(ChatCompletionsProvider):
|
||||
url,
|
||||
headers=self._headers(model, json=True),
|
||||
json=payload,
|
||||
timeout=self._timeout(model, resolution),
|
||||
timeout=self._image_timeout(model, resolution),
|
||||
)
|
||||
response.raise_for_status()
|
||||
raw = response.json()
|
||||
image_bytes = extract_image_from_response(
|
||||
raw,
|
||||
session=self.session,
|
||||
timeout=self._read_timeout(model, resolution),
|
||||
timeout=self._image_read_timeout(model, resolution),
|
||||
)
|
||||
if not image_bytes:
|
||||
raise AiResponseParseError("AI response did not contain an image")
|
||||
@@ -266,14 +277,14 @@ class ImagesGenerationProvider(BaseHttpProvider):
|
||||
url,
|
||||
headers=self._headers(model, json=True),
|
||||
json=payload,
|
||||
timeout=self._timeout(model, resolution),
|
||||
timeout=self._image_timeout(model, resolution),
|
||||
)
|
||||
response.raise_for_status()
|
||||
raw = response.json()
|
||||
image_bytes = extract_image_from_response(
|
||||
raw,
|
||||
session=self.session,
|
||||
timeout=self._read_timeout(model, resolution),
|
||||
timeout=self._image_read_timeout(model, resolution),
|
||||
)
|
||||
if not image_bytes:
|
||||
raise AiResponseParseError("AI response did not contain an image")
|
||||
@@ -317,14 +328,14 @@ class ImagesEditsProvider(BaseHttpProvider):
|
||||
headers=self._headers(model),
|
||||
data=data,
|
||||
files=files,
|
||||
timeout=self._timeout(model, resolution),
|
||||
timeout=self._image_timeout(model, resolution),
|
||||
)
|
||||
response.raise_for_status()
|
||||
raw = response.json()
|
||||
image_bytes = extract_image_from_response(
|
||||
raw,
|
||||
session=self.session,
|
||||
timeout=self._read_timeout(model, resolution),
|
||||
timeout=self._image_read_timeout(model, resolution),
|
||||
)
|
||||
if not image_bytes:
|
||||
raise AiResponseParseError("AI response did not contain an image")
|
||||
|
||||
@@ -8,6 +8,8 @@ from typing import Any, Iterable
|
||||
from urllib.parse import urljoin, urlparse
|
||||
|
||||
import requests
|
||||
from django.conf import settings
|
||||
from django.core.exceptions import ImproperlyConfigured
|
||||
|
||||
from .base import AiProviderConfigError, AiResponseParseError
|
||||
|
||||
@@ -26,6 +28,7 @@ SUPPORTED_API_TYPES = {
|
||||
}
|
||||
|
||||
RESOLUTION_TIMEOUTS = {"512": 180, "1K": 240, "2K": 360, "4K": 600}
|
||||
DEFAULT_IMAGE_UPSTREAM_DEADLINE_SECONDS = 180
|
||||
BASE64_KEYS = {"image_base64", "base64", "b64_json", "data"}
|
||||
IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp", ".gif"}
|
||||
|
||||
@@ -44,6 +47,34 @@ def request_timeout(connect_timeout: int, read_timeout: int, resolution: str) ->
|
||||
return connect_timeout, resolved_read_timeout
|
||||
|
||||
|
||||
def image_request_timeout(connect_timeout: int, read_timeout: int, resolution: str) -> tuple[int, int]:
|
||||
resolved_connect_timeout, resolved_read_timeout = request_timeout(
|
||||
connect_timeout,
|
||||
read_timeout,
|
||||
resolution,
|
||||
)
|
||||
return resolved_connect_timeout, cap_image_read_timeout(resolved_read_timeout)
|
||||
|
||||
|
||||
def cap_image_read_timeout(read_timeout: int) -> int:
|
||||
deadline = image_upstream_deadline_seconds()
|
||||
if deadline <= 0:
|
||||
return read_timeout
|
||||
return min(read_timeout, deadline)
|
||||
|
||||
|
||||
def image_upstream_deadline_seconds() -> int:
|
||||
try:
|
||||
value = getattr(
|
||||
settings,
|
||||
"AI_IMAGE_UPSTREAM_DEADLINE_SECONDS",
|
||||
DEFAULT_IMAGE_UPSTREAM_DEADLINE_SECONDS,
|
||||
)
|
||||
except ImproperlyConfigured:
|
||||
value = DEFAULT_IMAGE_UPSTREAM_DEADLINE_SECONDS
|
||||
return int(value)
|
||||
|
||||
|
||||
def detect_api_type(url: str, api_type: str = API_AUTO) -> str:
|
||||
if api_type and api_type != API_AUTO:
|
||||
if api_type not in SUPPORTED_API_TYPES:
|
||||
|
||||
+64
-1
@@ -13,7 +13,7 @@ from apps.ai.importers import import_ai_models_config
|
||||
from apps.ai.models import AiConfigAuditLog, AiModel, ModelAlias
|
||||
from apps.ai.providers import AiCapabilityError, ResolvedModel, get_provider, resolve_api_type
|
||||
from apps.ai.providers.openai_compatible import ChatCompletionsProvider, ImagesEditsProvider
|
||||
from apps.ai.providers.utils import resolution_to_size
|
||||
from apps.ai.providers.utils import image_request_timeout, resolution_to_size
|
||||
|
||||
|
||||
TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii")
|
||||
@@ -60,6 +60,11 @@ class ProviderUtilsTests(SimpleTestCase):
|
||||
self.assertEqual(resolution_to_size("1k"), "1024x1024")
|
||||
self.assertEqual(resolution_to_size("512px"), "512x512")
|
||||
|
||||
@override_settings(AI_IMAGE_UPSTREAM_DEADLINE_SECONDS=180)
|
||||
def test_image_request_timeout_caps_read_timeout_to_deadline(self):
|
||||
self.assertEqual(image_request_timeout(30, 0, "4K"), (30, 180))
|
||||
self.assertEqual(image_request_timeout(30, 120, "4K"), (30, 120))
|
||||
|
||||
|
||||
class ChatCompletionsProviderTests(SimpleTestCase):
|
||||
def test_generate_text_builds_chat_payload_and_cleans_titles(self):
|
||||
@@ -192,6 +197,64 @@ class ChatCompletionsProviderTests(SimpleTestCase):
|
||||
self.assertEqual(content[0], {"type": "text", "text": "Generate product image"})
|
||||
self.assertTrue(content[1]["image_url"]["url"].startswith("data:image/jpeg;base64,"))
|
||||
|
||||
@override_settings(AI_IMAGE_UPSTREAM_DEADLINE_SECONDS=180)
|
||||
def test_generate_image_caps_post_and_download_timeouts(self):
|
||||
session = FakeSession(
|
||||
FakeResponse(
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"content": [
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://cdn.example.com/out.png"},
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
),
|
||||
FakeResponse({}, content=b"generated-image"),
|
||||
)
|
||||
provider = ChatCompletionsProvider(session=session)
|
||||
model = ResolvedModel(
|
||||
name="Slow image",
|
||||
url="https://api.vectorengine.ai/v1/chat/completions",
|
||||
model="image-model",
|
||||
api_key="test-key",
|
||||
api_type="chat",
|
||||
timeout_seconds=0,
|
||||
connect_timeout_seconds=30,
|
||||
)
|
||||
|
||||
result = provider.generate_image("Generate product image", model, resolution="4K")
|
||||
|
||||
self.assertEqual(result.image, b"generated-image")
|
||||
self.assertEqual(session.posts[0]["timeout"], (30, 180))
|
||||
self.assertEqual(session.gets[0]["timeout"], 180)
|
||||
|
||||
@override_settings(AI_IMAGE_UPSTREAM_DEADLINE_SECONDS=180)
|
||||
def test_generate_text_does_not_use_image_deadline(self):
|
||||
session = FakeSession(
|
||||
FakeResponse({"choices": [{"message": {"content": "1. Red Dress"}}]})
|
||||
)
|
||||
provider = ChatCompletionsProvider(session=session)
|
||||
model = ResolvedModel(
|
||||
name="Slow text",
|
||||
url="https://api.vectorengine.ai/v1/chat/completions",
|
||||
model="text-model",
|
||||
api_key="test-key",
|
||||
api_type="chat",
|
||||
timeout_seconds=0,
|
||||
connect_timeout_seconds=30,
|
||||
)
|
||||
|
||||
provider.generate_text("Generate titles", model, resolution="4K")
|
||||
|
||||
self.assertEqual(session.posts[0]["timeout"], (30, 600))
|
||||
|
||||
|
||||
class ImagesEditsProviderTests(SimpleTestCase):
|
||||
def test_generate_image_builds_multipart_request_and_parses_base64(self):
|
||||
|
||||
@@ -290,6 +290,8 @@ def precharge_or_raise(**kwargs):
|
||||
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)
|
||||
|
||||
@@ -1514,6 +1514,37 @@ class GenerateApiTests(TestCase):
|
||||
1,
|
||||
)
|
||||
|
||||
def test_image_upstream_timeout_refunds_precharged_points_and_marks_call_failed(self):
|
||||
self.provider.image_error = requests.Timeout("image upstream deadline exceeded")
|
||||
|
||||
response = self.post_with_provider(
|
||||
"/api/v1/generate/image",
|
||||
{"prompt": "生成图片", "model": self.image_alias, "resolution": "1K"},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 502)
|
||||
self.assertEqual(response.data["error"]["code"], "upstream_timeout")
|
||||
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("image upstream deadline exceeded", 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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user