Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion backend/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,7 +133,7 @@ def get_image_provider_config(cls, provider_name: str = None):
)

provider_type = provider_config.get('type', provider_name)
if provider_type in ['openai', 'openai_compatible', 'image_api']:
if provider_type in ['openai', 'openai_compatible', 'image_api', 'atlascloud_image']:
if not provider_config.get('base_url'):
logger.error(f"服务商 [{provider_name}] 类型为 {provider_type},但未配置 base_url")
raise ValueError(
Expand Down
163 changes: 163 additions & 0 deletions backend/generators/atlascloud_image.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,163 @@
"""Atlas Cloud Media API image generator."""
import base64
import logging
import time
from typing import Any, Dict, Optional

import requests

from .base import ImageGeneratorBase

logger = logging.getLogger(__name__)


class AtlasCloudImageGenerator(ImageGeneratorBase):
"""Generate images through the Atlas Cloud asynchronous Media API."""

DEFAULT_BASE_URL = "https://api.atlascloud.ai/api/v1"
DEFAULT_MODEL = "bytedance/seedream-v5.0-lite"
DEFAULT_SIZE = "1728*2304"
DEFAULT_OUTPUT_FORMAT = "png"

def __init__(self, config: Dict[str, Any]):
super().__init__(config)
self.base_url = (config.get("base_url") or self.DEFAULT_BASE_URL).rstrip("/")
self.model = config.get("model") or self.DEFAULT_MODEL
self.size = config.get("size") or config.get("image_size") or self.DEFAULT_SIZE
self.output_format = config.get("output_format") or self.DEFAULT_OUTPUT_FORMAT
self.poll_interval_seconds = _non_negative_float(
config.get("poll_interval_seconds"),
default=3.0,
)
self.max_poll_attempts = _positive_int(
config.get("max_poll_attempts"),
default=120,
)
self.session = requests.Session()

def validate_config(self) -> bool:
if not self.api_key:
raise ValueError(
"Atlas Cloud API Key 未配置。\n"
"解决方案:在系统设置页面编辑该服务商,填写 API Key"
)
return True

def generate_image(
self,
prompt: str,
model: Optional[str] = None,
**kwargs,
) -> bytes:
"""Submit an Atlas Cloud image task, poll for completion, then download the image."""
self.validate_config()

request_model = model or self.model
payload = {
"model": request_model,
"prompt": prompt,
"size": kwargs.get("size") or self.size,
"output_format": kwargs.get("output_format") or self.output_format,
"enable_base64_output": bool(kwargs.get("enable_base64_output", False)),
}

logger.info(f"Atlas Cloud 生成图片: model={request_model}, size={payload['size']}")
prediction = self._submit_task(payload)
outputs = self._poll_outputs(prediction["id"])
return self._read_output(outputs[0])

def _submit_task(self, payload: Dict[str, Any]) -> Dict[str, Any]:
response = self.session.post(
f"{self.base_url}/model/generateImage",
headers=self._headers(),
json=payload,
timeout=60,
)
if response.status_code != 200:
raise Exception(f"Atlas Cloud 提交任务失败: HTTP {response.status_code}: {response.text[:500]}")

data = _unwrap_data(response.json())
request_id = data.get("id") or data.get("request_id")
if not request_id:
raise Exception(f"Atlas Cloud 提交响应缺少任务 ID: {str(data)[:500]}")
return {"id": request_id}

def _poll_outputs(self, request_id: str) -> list[str]:
last_status = ""
for attempt in range(self.max_poll_attempts):
response = self._get_prediction_response(request_id)
if response.status_code != 200:
raise Exception(f"Atlas Cloud 查询任务失败: HTTP {response.status_code}: {response.text[:500]}")

data = _unwrap_data(response.json())
last_status = str(data.get("status") or "").lower()
outputs = data.get("outputs") or data.get("output") or []
if isinstance(outputs, str):
outputs = [outputs]

if last_status in {"completed", "succeeded", "success"} and outputs:
return outputs
if last_status in {"failed", "canceled", "cancelled"}:
raise Exception(f"Atlas Cloud 图片任务失败: {str(data)[:500]}")

if attempt < self.max_poll_attempts - 1:
time.sleep(self.poll_interval_seconds)

raise Exception(f"Atlas Cloud 图片任务超时,最后状态: {last_status or 'unknown'}")

def _get_prediction_response(self, request_id: str):
response = self.session.get(
f"{self.base_url}/model/result/{request_id}",
headers=self._headers(),
timeout=60,
)
if response.status_code != 404:
return response
return self.session.get(
f"{self.base_url}/model/prediction/{request_id}",
headers=self._headers(),
timeout=60,
)

def _read_output(self, output: str) -> bytes:
if output.startswith("data:image"):
return base64.b64decode(output.split(",", 1)[1])
if _looks_like_base64(output):
return base64.b64decode(output)

response = self.session.get(output, timeout=60)
if response.status_code != 200:
raise Exception(f"Atlas Cloud 图片下载失败: HTTP {response.status_code}")
return response.content

def _headers(self) -> Dict[str, str]:
return {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}


def _unwrap_data(payload: Dict[str, Any]) -> Dict[str, Any]:
data = payload.get("data")
return data if isinstance(data, dict) else payload


def _looks_like_base64(value: str) -> bool:
stripped = value.strip()
return bool(stripped) and not stripped.startswith(("http://", "https://"))


def _positive_int(value: Any, default: int) -> int:
try:
parsed = int(value)
return parsed if parsed > 0 else default
except (TypeError, ValueError):
return default


def _non_negative_float(value: Any, default: float) -> float:
try:
parsed = float(value)
return parsed if parsed >= 0 else default
except (TypeError, ValueError):
return default
2 changes: 2 additions & 0 deletions backend/generators/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from .google_genai import GoogleGenAIGenerator
from .openai_compatible import OpenAICompatibleGenerator
from .image_api import ImageApiGenerator
from .atlascloud_image import AtlasCloudImageGenerator


class ImageGeneratorFactory:
Expand All @@ -15,6 +16,7 @@ class ImageGeneratorFactory:
'openai': OpenAICompatibleGenerator,
'openai_compatible': OpenAICompatibleGenerator,
'image_api': ImageApiGenerator,
'atlascloud_image': AtlasCloudImageGenerator,
}

@classmethod
Expand Down
63 changes: 61 additions & 2 deletions backend/routes/config_routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,7 +132,7 @@ def test_connection():
测试服务商连接

请求体:
- type: 服务商类型(google_genai/google_gemini/openai_compatible/image_api)
- type: 服务商类型(google_genai/google_gemini/openai_compatible/image_api/atlascloud_image
- provider_name: 服务商名称(用于从配置读取 API Key)
- api_key: API Key(可选,若不提供则从配置读取)
- base_url: Base URL(可选)
Expand All @@ -153,7 +153,7 @@ def test_connection():
return api_error_response(
validation_error("缺少 type 参数", "请选择服务商类型后再测试连接。")
)
if provider_type not in ['google_genai', 'google_gemini', 'openai_compatible', 'image_api']:
if provider_type not in ['google_genai', 'google_gemini', 'openai_compatible', 'image_api', 'atlascloud_image']:
return api_error_response(
validation_error(f"不支持的类型: {provider_type}", "请选择正确的服务商类型后再测试连接。")
)
Expand Down Expand Up @@ -324,6 +324,9 @@ def _test_provider_connection(provider_type: str, config: dict) -> dict:
elif provider_type == 'image_api':
return _test_image_api(config)

elif provider_type == 'atlascloud_image':
return _test_atlascloud_image(config)

else:
raise ValueError(f"不支持的类型: {provider_type}")

Expand Down Expand Up @@ -440,6 +443,62 @@ def _test_image_api(config: dict) -> dict:
raise Exception(f"HTTP {response.status_code}: {response.text[:200]}")


def _test_atlascloud_image(config: dict) -> dict:
"""测试 Atlas Cloud Media API 连接和模型目录。"""
import requests

base_url = (config.get('base_url') or 'https://api.atlascloud.ai/api/v1').rstrip('/')
model = config.get('model') or 'bytedance/seedream-v5.0-lite'

response = requests.get(
f"{base_url}/models",
headers={'Authorization': f"Bearer {config['api_key']}"},
timeout=30,
)
if response.status_code != 200:
raise Exception(f"HTTP {response.status_code}: {response.text[:200]}")

try:
payload = response.json()
except Exception as exc:
raise Exception(f"Atlas Cloud 模型目录响应不是合法 JSON: {response.text[:500]}") from exc

model_ids = _extract_model_ids(payload)
if model_ids and model not in model_ids:
return {
"success": True,
"warning": True,
"status": "warning",
"message": f"Atlas Cloud 连接成功,但模型目录中未找到 {model}。请确认模型名是否可用。"
}

return {
"success": True,
"message": "Atlas Cloud 连接成功!模型目录可访问,图片任务提交将在生成时执行。"
}


def _extract_model_ids(payload) -> set[str]:
data = payload.get('data') if isinstance(payload, dict) else payload
if isinstance(data, dict):
for key in ['models', 'items', 'list']:
if isinstance(data.get(key), list):
data = data[key]
break
if not isinstance(data, list):
return set()

model_ids = set()
for item in data:
if isinstance(item, dict):
model_id = item.get('model') or item.get('id') or item.get('name')
if isinstance(model_id, str):
model_ids.add(model_id)
elif isinstance(item, str):
model_ids.add(item)
return model_ids


def _test_openai_chat_completion(config: dict, test_prompt: str) -> LlmSmokeResult:
"""用当前配置发送一次真实 OpenAI-compatible LLM 请求。"""
import requests
Expand Down
8 changes: 8 additions & 0 deletions backend/services/image.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,6 +214,14 @@ def _generate_single_image(
model=self.provider_config.get('model', 'nano-banana-2'),
reference_images=reference_images if reference_images else None,
)
elif self.provider_config.get('type') == 'atlascloud_image':
logger.debug(" 使用 Atlas Cloud 图片生成器")
image_data = self.generator.generate_image(
prompt=prompt,
size=self.provider_config.get('size', '1728*2304'),
output_format=self.provider_config.get('output_format', 'png'),
model=self.provider_config.get('model', 'bytedance/seedream-v5.0-lite'),
)
else:
logger.debug(f" 使用 OpenAI 兼容生成器")
image_data = self.generator.generate_image(
Expand Down
12 changes: 12 additions & 0 deletions image_providers.yaml.example
Original file line number Diff line number Diff line change
Expand Up @@ -27,3 +27,15 @@ providers:
base_url: https://your-api-endpoint.com
model: dall-e-3
high_concurrency: false

# Atlas Cloud Media API(异步图片生成)
atlascloud:
type: atlascloud_image
api_key: ak-xxxxxxxxxxxxxxxxxxxx
base_url: https://api.atlascloud.ai/api/v1
model: bytedance/seedream-v5.0-lite
size: "1728*2304"
output_format: png
poll_interval_seconds: 3
max_poll_attempts: 120
high_concurrency: false
Loading