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: 2 additions & 0 deletions invokeai/app/api/dependencies.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from invokeai.app.services.external_generation.external_generation_default import ExternalGenerationService
from invokeai.app.services.external_generation.providers import (
AlibabaCloudProvider,
AtlasCloudProvider,
GeminiProvider,
OpenAIProvider,
SeedreamProvider,
Expand Down Expand Up @@ -188,6 +189,7 @@ def initialize(
external_generation = ExternalGenerationService(
providers={
AlibabaCloudProvider.provider_id: AlibabaCloudProvider(app_config=configuration, logger=logger),
AtlasCloudProvider.provider_id: AtlasCloudProvider(app_config=configuration, logger=logger),
GeminiProvider.provider_id: GeminiProvider(app_config=configuration, logger=logger),
OpenAIProvider.provider_id: OpenAIProvider(app_config=configuration, logger=logger),
SeedreamProvider.provider_id: SeedreamProvider(app_config=configuration, logger=logger),
Expand Down
1 change: 1 addition & 0 deletions invokeai/app/api/routers/app_info.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,7 @@ class ExternalProviderConfigModel(BaseModel):

EXTERNAL_PROVIDER_FIELDS: dict[str, tuple[str, str]] = {
"alibabacloud": ("external_alibabacloud_api_key", "external_alibabacloud_base_url"),
"atlascloud": ("external_atlascloud_api_key", "external_atlascloud_base_url"),
"gemini": ("external_gemini_api_key", "external_gemini_base_url"),
"openai": ("external_openai_api_key", "external_openai_base_url"),
"seedream": ("external_seedream_api_key", "external_seedream_base_url"),
Expand Down
28 changes: 28 additions & 0 deletions invokeai/app/invocations/external_image_generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -349,3 +349,31 @@ class AlibabaCloudImageGenerationInvocation(BaseExternalImageGenerationInvocatio
ui_model_format=[ModelFormat.ExternalApi],
ui_model_provider_id=["alibabacloud"],
)


@invocation(
"atlascloud_image_generation",
title="Atlas Cloud Image Generation",
tags=["external", "generation", "atlascloud"],
category="image",
version="1.0.0",
)
class AtlasCloudImageGenerationInvocation(BaseExternalImageGenerationInvocation):
"""Generate images through the Atlas Cloud asynchronous media API."""

provider_id = "atlascloud"

model: ModelIdentifierField = InputField(
description=FieldDescriptions.main_model,
ui_model_base=[BaseModelType.External],
ui_model_type=[ModelType.ExternalImageGenerator],
ui_model_format=[ModelFormat.ExternalApi],
ui_model_provider_id=["atlascloud"],
)

mode: ExternalGenerationMode = InputField(default="txt2img", description="Generation mode.", ui_hidden=True)
init_image: ImageField | None = InputField(
default=None, description="Init image for img2img/inpaint", ui_hidden=True
)
mask_image: ImageField | None = InputField(default=None, description="Mask image for inpaint", ui_hidden=True)
reference_images: list[ImageField] = InputField(default=[], description="Reference images", ui_hidden=True)
10 changes: 10 additions & 0 deletions invokeai/app/services/config/config_default.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,8 @@
EXTERNAL_PROVIDER_CONFIG_FIELDS = (
"external_alibabacloud_api_key",
"external_alibabacloud_base_url",
"external_atlascloud_api_key",
"external_atlascloud_base_url",
"external_gemini_api_key",
"external_gemini_base_url",
"external_openai_api_key",
Expand Down Expand Up @@ -133,6 +135,8 @@ class InvokeAIAppConfig(BaseSettings):
strict_password_checking: Enforce strict password requirements. When True, passwords must contain uppercase, lowercase, and numbers. When False (default), any password is accepted but its strength (weak/moderate/strong) is reported to the user.
external_alibabacloud_api_key: API key for Alibaba Cloud DashScope image generation.
external_alibabacloud_base_url: Base URL override for Alibaba Cloud DashScope image generation.
external_atlascloud_api_key: API key for Atlas Cloud image generation.
external_atlascloud_base_url: Base URL override for Atlas Cloud image generation.
external_gemini_api_key: API key for Gemini image generation.
external_openai_api_key: API key for OpenAI image generation.
external_gemini_base_url: Base URL override for Gemini image generation.
Expand Down Expand Up @@ -250,6 +254,12 @@ class InvokeAIAppConfig(BaseSettings):
external_alibabacloud_base_url: Optional[str] = Field(
default=None, description="Base URL override for Alibaba Cloud DashScope image generation."
)
external_atlascloud_api_key: Optional[str] = Field(
default=None, description="API key for Atlas Cloud image generation."
)
external_atlascloud_base_url: Optional[str] = Field(
default=None, description="Base URL override for Atlas Cloud image generation."
)
external_gemini_api_key: Optional[str] = Field(default=None, description="API key for Gemini image generation.")
external_openai_api_key: Optional[str] = Field(default=None, description="API key for OpenAI image generation.")
external_gemini_base_url: Optional[str] = Field(
Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from invokeai.app.services.external_generation.providers.alibabacloud import AlibabaCloudProvider
from invokeai.app.services.external_generation.providers.atlascloud import AtlasCloudProvider
from invokeai.app.services.external_generation.providers.gemini import GeminiProvider
from invokeai.app.services.external_generation.providers.openai import OpenAIProvider
from invokeai.app.services.external_generation.providers.seedream import SeedreamProvider

__all__ = ["AlibabaCloudProvider", "GeminiProvider", "OpenAIProvider", "SeedreamProvider"]
__all__ = ["AlibabaCloudProvider", "AtlasCloudProvider", "GeminiProvider", "OpenAIProvider", "SeedreamProvider"]
210 changes: 210 additions & 0 deletions invokeai/app/services/external_generation/providers/atlascloud.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,210 @@
from __future__ import annotations

import io
import time
from typing import Any

import requests
from PIL import Image
from PIL.Image import Image as PILImageType

from invokeai.app.services.external_generation.errors import (
ExternalProviderRateLimitError,
ExternalProviderRequestError,
)
from invokeai.app.services.external_generation.external_generation_base import ExternalProvider
from invokeai.app.services.external_generation.external_generation_common import (
ExternalGeneratedImage,
ExternalGenerationRequest,
ExternalGenerationResult,
)

_DEFAULT_BASE_URL = "https://api.atlascloud.ai"
_REQUEST_TIMEOUT = 30
_POLL_INTERVAL = 3.0
_POLL_TIMEOUT = 300.0
_DOWNLOAD_TIMEOUT = 60
_DOWNLOAD_MAX_BYTES = 32 * 1024 * 1024
_SUCCESS_STATUSES = {"completed", "succeeded"}
_FAILURE_STATUSES = {"canceled", "cancelled", "failed"}


class AtlasCloudProvider(ExternalProvider):
provider_id = "atlascloud"

def is_configured(self) -> bool:
return bool(self._app_config.external_atlascloud_api_key)

def generate(self, request: ExternalGenerationRequest) -> ExternalGenerationResult:
api_key = self._app_config.external_atlascloud_api_key
if not api_key:
raise ExternalProviderRequestError("Atlas Cloud API key is not configured")

base_url = (self._app_config.external_atlascloud_base_url or _DEFAULT_BASE_URL).rstrip("/")
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
}
payload: dict[str, object] = {
"model": request.model.provider_model_id,
"prompt": request.prompt,
"size": f"{request.width}*{request.height}",
"num_images": request.num_images,
}
if request.seed is not None:
payload["seed"] = request.seed

submit_url = f"{base_url}/api/v1/model/generateImage"
try:
response = requests.post(submit_url, headers=headers, json=payload, timeout=_REQUEST_TIMEOUT)
except requests.RequestException as exc:
raise ExternalProviderRequestError(f"Atlas Cloud image submission failed: {exc}") from exc

self._raise_for_error(response, "image submission")
prediction = self._response_data(response, "image submission")
prediction_id = prediction.get("id")
if not isinstance(prediction_id, str) or not prediction_id:
raise ExternalProviderRequestError("Atlas Cloud image submission response missing prediction id")

poll_url = self._get_poll_url(prediction, base_url, prediction_id)
completed = self._poll_prediction(poll_url, headers, prediction_id)
output_urls = completed.get("output", completed.get("outputs"))
if not isinstance(output_urls, list):
raise ExternalProviderRequestError("Atlas Cloud completed prediction contained no image outputs")

images: list[ExternalGeneratedImage] = []
for output_url in output_urls:
if isinstance(output_url, str) and output_url:
images.append(ExternalGeneratedImage(image=self._download_image(output_url), seed=request.seed))

if not images:
raise ExternalProviderRequestError("Atlas Cloud completed prediction contained no downloadable images")

return ExternalGenerationResult(
images=images,
seed_used=request.seed,
provider_request_id=prediction_id,
provider_metadata={
"model": request.model.provider_model_id,
"status": str(completed.get("status", "succeeded")),
},
)

def _poll_prediction(
self,
poll_url: str,
headers: dict[str, str],
prediction_id: str,
) -> dict[str, Any]:
started_at = time.monotonic()

while True:
if time.monotonic() - started_at > _POLL_TIMEOUT:
raise ExternalProviderRequestError(
f"Atlas Cloud prediction {prediction_id} timed out after {_POLL_TIMEOUT:.0f}s"
)

try:
response = requests.get(poll_url, headers=headers, timeout=_REQUEST_TIMEOUT)
except requests.RequestException as exc:
raise ExternalProviderRequestError(f"Atlas Cloud prediction polling failed: {exc}") from exc

self._raise_for_error(response, "prediction polling")
prediction = self._response_data(response, "prediction polling")
status = str(prediction.get("status", "")).lower()

if status in _SUCCESS_STATUSES:
return prediction
if status in _FAILURE_STATUSES:
detail = prediction.get("error") or prediction.get("logs") or "Unknown provider error"
raise ExternalProviderRequestError(f"Atlas Cloud prediction {prediction_id} failed: {detail}")

self._logger.debug("Atlas Cloud prediction %s status: %s", prediction_id, status or "unknown")
time.sleep(_POLL_INTERVAL)

@staticmethod
def _get_poll_url(prediction: dict[str, Any], base_url: str, prediction_id: str) -> str:
urls = prediction.get("urls")
if isinstance(urls, dict):
result_url = urls.get("result")
if isinstance(result_url, str) and result_url:
if result_url.startswith("/"):
return f"{base_url}{result_url}"
return result_url
return f"{base_url}/api/v1/model/result/{prediction_id}"

@staticmethod
def _response_data(response: requests.Response, operation: str) -> dict[str, Any]:
try:
payload = response.json()
except ValueError as exc:
raise ExternalProviderRequestError(f"Atlas Cloud {operation} returned invalid JSON") from exc
if not isinstance(payload, dict):
raise ExternalProviderRequestError(f"Atlas Cloud {operation} response was not a JSON object")

if "data" in payload:
code = payload.get("code")
if code not in (None, 0, 200):
detail = payload.get("message") or payload.get("msg") or "Unknown provider error"
raise ExternalProviderRequestError(f"Atlas Cloud {operation} failed: {detail}")
data = payload.get("data")
if not isinstance(data, dict):
raise ExternalProviderRequestError(f"Atlas Cloud {operation} response missing data")
return data
return payload

@staticmethod
def _raise_for_error(response: requests.Response, operation: str) -> None:
if response.ok:
return
if response.status_code == 429:
retry_after = _parse_retry_after(response.headers.get("Retry-After"))
raise ExternalProviderRateLimitError(
f"Atlas Cloud rate limit exceeded during {operation}",
retry_after=retry_after,
)
raise ExternalProviderRequestError(
f"Atlas Cloud {operation} failed with status {response.status_code}: {response.text}"
)

def _download_image(self, url: str) -> PILImageType:
try:
response = requests.get(url, timeout=_DOWNLOAD_TIMEOUT, stream=True)
except requests.RequestException as exc:
raise ExternalProviderRequestError(f"Failed to download image from Atlas Cloud: {exc}") from exc

with response:
if not response.ok:
raise ExternalProviderRequestError(
f"Failed to download image from Atlas Cloud (status {response.status_code})"
)

content_length = response.headers.get("Content-Length")
if content_length is not None:
try:
if int(content_length) > _DOWNLOAD_MAX_BYTES:
raise ExternalProviderRequestError(f"Atlas Cloud image exceeds {_DOWNLOAD_MAX_BYTES} byte cap")
except ValueError:
pass

buffer = bytearray()
for chunk in response.iter_content(chunk_size=64 * 1024):
if not chunk:
continue
buffer.extend(chunk)
if len(buffer) > _DOWNLOAD_MAX_BYTES:
raise ExternalProviderRequestError(f"Atlas Cloud image exceeds {_DOWNLOAD_MAX_BYTES} byte cap")

try:
return Image.open(io.BytesIO(bytes(buffer))).convert("RGB")
except Exception as exc:
raise ExternalProviderRequestError("Atlas Cloud output was not a valid image") from exc


def _parse_retry_after(value: str | None) -> float | None:
if not value:
return None
try:
return float(value)
except ValueError:
return None
18 changes: 18 additions & 0 deletions invokeai/backend/model_manager/starter_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -1799,6 +1799,23 @@ def _gemini_3_resolution_presets(
prompts=[{"name": "reference_images"}], image=[{"name": "dimensions"}]
)

atlascloud_flux_schnell = StarterModel(
name="Atlas Cloud FLUX.1 Schnell",
base=BaseModelType.External,
source="external://atlascloud/black-forest-labs/flux-schnell",
description="FLUX.1 Schnell text-to-image generation through the Atlas Cloud asynchronous media API. Requires a configured Atlas Cloud API key and may incur provider usage costs.",
type=ModelType.ExternalImageGenerator,
format=ModelFormat.ExternalApi,
capabilities=ExternalModelCapabilities(
modes=["txt2img"],
supports_negative_prompt=False,
supports_seed=True,
max_images_per_request=4,
),
default_settings=ExternalApiModelDefaultSettings(width=1024, height=1024, num_images=1),
panel_schema=ExternalModelPanelSchema(image=[{"name": "dimensions"}]),
)

openai_gpt_image_2 = StarterModel(
name="GPT Image 2",
base=BaseModelType.External,
Expand Down Expand Up @@ -2242,6 +2259,7 @@ def _gemini_3_resolution_presets(
gemini_flash_image,
gemini_pro_image_preview,
gemini_3_1_flash_image_preview,
atlascloud_flux_schnell,
openai_gpt_image_2,
openai_gpt_image_1_5,
openai_gpt_image_1,
Expand Down
Loading