feat(byteplus): add seedream image generation provider

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Mubashir Osmani 2026-07-18 22:19:56 +00:00
parent 1ebf2a78a9
commit 49fcfe503c
10 changed files with 487 additions and 0 deletions

View file

@ -636,6 +636,7 @@ inception_models: Set = set()
hyperbolic_models: Set = set()
black_forest_labs_models: Set = set()
recraft_models: Set = set()
byteplus_models: Set = set()
cometapi_models: Set = set()
oci_models: Set = set()
vercel_ai_gateway_models: Set = set()
@ -857,6 +858,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
nebius_embedding_models.add(key)
elif value.get("litellm_provider") == "aiml":
aiml_models.add(key)
elif value.get("litellm_provider") == "byteplus":
byteplus_models.add(key)
elif value.get("litellm_provider") == "assemblyai":
assemblyai_models.add(key)
elif value.get("litellm_provider") == "jina_ai":
@ -1038,6 +1041,7 @@ model_list = list(
| inception_models
| black_forest_labs_models
| recraft_models
| byteplus_models
| cometapi_models
| oci_models
| heroku_models
@ -1145,6 +1149,7 @@ models_by_provider: dict = {
"hyperbolic": hyperbolic_models,
"black_forest_labs": black_forest_labs_models,
"recraft": recraft_models,
"byteplus": byteplus_models,
"cometapi": cometapi_models,
"oci": oci_models,
"volcengine": volcengine_models,

View file

@ -385,6 +385,7 @@ def image_generation(
#########################################################
elif custom_llm_provider in (
litellm.LlmProviders.RECRAFT,
litellm.LlmProviders.BYTEPLUS,
litellm.LlmProviders.AIML,
litellm.LlmProviders.GEMINI,
litellm.LlmProviders.FAL_AI,

View file

View file

@ -0,0 +1,13 @@
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from .transformation import BytePlusImageGenerationConfig
__all__ = [
"BytePlusImageGenerationConfig",
]
def get_byteplus_image_generation_config(model: str) -> BaseImageGenerationConfig:
return BytePlusImageGenerationConfig()

View file

@ -0,0 +1,142 @@
from typing import TYPE_CHECKING, Any
import httpx
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIImageGenerationOptionalParams,
)
from litellm.types.utils import ImageObject, ImageResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
DEFAULT_API_BASE = "https://ark.ap-southeast.bytepluses.com/api/v3"
IMAGE_GENERATION_ENDPOINT = "images/generations"
class BytePlusImageGenerationConfig(BaseImageGenerationConfig):
"""
Configuration for BytePlus ModelArk (Seedream) image generation.
OpenAI-compatible POST {api_base}/images/generations
https://docs.byteplus.com/en/docs/ModelArk/1541523
"""
def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]:
return ["n", "size", "response_format"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
passthrough = {k: v for k, v in non_default_params.items() if k not in optional_params and k != "n"}
n = non_default_params.get("n")
sequential = (
{
"sequential_image_generation": "auto",
"sequential_image_generation_options": {"max_images": n},
}
if isinstance(n, int) and n > 1
else {}
)
return {**optional_params, **passthrough, **sequential}
def get_complete_url(
self,
api_base: str | None,
api_key: str | None,
model: str,
optional_params: dict,
litellm_params: dict,
stream: bool | None = None,
) -> str:
base_url = api_base or get_secret_str("BYTEPLUS_API_BASE") or get_secret_str("ARK_API_BASE") or DEFAULT_API_BASE
base_url = base_url.rstrip("/")
if base_url.endswith(IMAGE_GENERATION_ENDPOINT):
return base_url
return f"{base_url}/{IMAGE_GENERATION_ENDPOINT}"
def validate_environment(
self,
headers: dict,
model: str,
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
final_api_key = api_key or get_secret_str("BYTEPLUS_API_KEY") or get_secret_str("ARK_API_KEY")
if not final_api_key:
raise ValueError("BYTEPLUS_API_KEY or ARK_API_KEY is not set")
headers["Authorization"] = f"Bearer {final_api_key}"
headers["Content-Type"] = "application/json"
return headers
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
return {"model": model, "prompt": prompt, **optional_params}
def transform_image_generation_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ImageResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: str | None = None,
json_mode: bool | None = None,
) -> ImageResponse:
if raw_response.status_code != 200:
raise self.get_error_class(
error_message=raw_response.text,
status_code=raw_response.status_code,
headers=raw_response.headers,
)
try:
response_data = raw_response.json()
except ValueError as e:
raise self.get_error_class(
error_message=f"Failed to parse BytePlus image generation response: {e}",
status_code=raw_response.status_code,
headers=raw_response.headers,
)
if "error" in response_data and "data" not in response_data:
error = response_data["error"]
raise self.get_error_class(
error_message=str(error.get("message", error) if isinstance(error, dict) else error),
status_code=raw_response.status_code,
headers=raw_response.headers,
)
model_response.data = [
ImageObject(
url=item.get("url"),
b64_json=item.get("b64_json"),
)
for item in response_data.get("data", [])
]
return model_response

View file

@ -13237,6 +13237,46 @@
"/v1/images/generations"
]
},
"byteplus/seedream-5-0-pro-260628": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"byteplus/seedream-5-0-260128": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"byteplus/seedream-5-0-lite-260128": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"byteplus/seedream-4-5-251128": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"byteplus/seedream-4-0-250828": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"databricks/databricks-bge-large-en": {
"input_cost_per_token": 1.0003e-07,
"input_dbu_cost_per_token": 1.429e-06,

View file

@ -3434,6 +3434,7 @@ class LlmProviders(str, Enum):
CURSOR = "cursor"
BEDROCK_MANTLE = "bedrock_mantle"
GDC = "gdc"
BYTEPLUS = "byteplus"
# Create a set of all provider values for quick lookup

View file

@ -8579,6 +8579,12 @@ class ProviderConfigManager:
)
return get_recraft_image_generation_config(model)
elif LlmProviders.BYTEPLUS == provider:
from litellm.llms.byteplus.image_generation import (
get_byteplus_image_generation_config,
)
return get_byteplus_image_generation_config(model)
elif LlmProviders.AIML == provider:
from litellm.llms.aiml.image_generation import (
get_aiml_image_generation_config,

View file

@ -13237,6 +13237,46 @@
"/v1/images/generations"
]
},
"byteplus/seedream-5-0-pro-260628": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"byteplus/seedream-5-0-260128": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"byteplus/seedream-5-0-lite-260128": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"byteplus/seedream-4-5-251128": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"byteplus/seedream-4-0-250828": {
"litellm_provider": "byteplus",
"mode": "image_generation",
"source": "https://docs.byteplus.com/en/docs/ModelArk/1541523",
"supported_endpoints": [
"/v1/images/generations"
]
},
"databricks/databricks-bge-large-en": {
"input_cost_per_token": 1.0003e-07,
"input_dbu_cost_per_token": 1.429e-06,

View file

@ -0,0 +1,239 @@
import os
import sys
from unittest.mock import MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm import get_llm_provider
from litellm.llms.byteplus.image_generation.transformation import (
DEFAULT_API_BASE,
IMAGE_GENERATION_ENDPOINT,
BytePlusImageGenerationConfig,
)
from litellm.types.utils import ImageResponse, LlmProviders
from litellm.utils import ProviderConfigManager
MODULE = "litellm.llms.byteplus.image_generation.transformation.get_secret_str"
class TestBytePlusImageGenerationTransformation:
def setup_method(self):
self.config = BytePlusImageGenerationConfig()
self.model = "seedream-5-0-260128"
self.logging_obj = MagicMock()
def test_provider_routing(self):
model, provider, _, _ = get_llm_provider("byteplus/seedream-5-0-260128")
assert provider == "byteplus"
assert model == "seedream-5-0-260128"
def test_provider_config_registered(self):
config = ProviderConfigManager.get_provider_image_generation_config(
model=self.model,
provider=LlmProviders.BYTEPLUS,
)
assert isinstance(config, BytePlusImageGenerationConfig)
def test_supported_params(self):
assert self.config.get_supported_openai_params(self.model) == ["n", "size", "response_format"]
def test_map_openai_params_passthrough(self):
result = self.config.map_openai_params(
non_default_params={"size": "2048x2048", "response_format": "url", "watermark": False},
optional_params={},
model=self.model,
drop_params=False,
)
assert result == {"size": "2048x2048", "response_format": "url", "watermark": False}
def test_map_openai_params_n_expands_to_sequential(self):
result = self.config.map_openai_params(
non_default_params={"n": 4, "size": "1024x1024"},
optional_params={},
model=self.model,
drop_params=False,
)
assert result["size"] == "1024x1024"
assert result["sequential_image_generation"] == "auto"
assert result["sequential_image_generation_options"] == {"max_images": 4}
assert "n" not in result
def test_map_openai_params_n_one_no_sequential(self):
result = self.config.map_openai_params(
non_default_params={"n": 1},
optional_params={},
model=self.model,
drop_params=False,
)
assert "sequential_image_generation" not in result
assert "n" not in result
@patch(MODULE)
def test_get_complete_url_with_api_base(self, mock_secret):
result = self.config.get_complete_url(
api_base="https://custom.ark.example.com/api/v3",
api_key="k",
model=self.model,
optional_params={},
litellm_params={},
)
assert result == f"https://custom.ark.example.com/api/v3/{IMAGE_GENERATION_ENDPOINT}"
mock_secret.assert_not_called()
@patch(MODULE)
def test_get_complete_url_default(self, mock_secret):
mock_secret.return_value = None
result = self.config.get_complete_url(
api_base=None,
api_key="k",
model=self.model,
optional_params={},
litellm_params={},
)
assert result == f"{DEFAULT_API_BASE}/{IMAGE_GENERATION_ENDPOINT}"
@patch(MODULE)
def test_get_complete_url_no_double_endpoint(self, mock_secret):
mock_secret.return_value = None
base = f"{DEFAULT_API_BASE}/{IMAGE_GENERATION_ENDPOINT}"
result = self.config.get_complete_url(
api_base=base,
api_key="k",
model=self.model,
optional_params={},
litellm_params={},
)
assert result == base
@patch(MODULE)
def test_validate_environment_with_api_key(self, mock_secret):
headers = self.config.validate_environment(
headers={},
model=self.model,
messages=[],
optional_params={},
litellm_params={},
api_key="my_key",
)
assert headers["Authorization"] == "Bearer my_key"
assert headers["Content-Type"] == "application/json"
mock_secret.assert_not_called()
@patch(MODULE)
def test_validate_environment_ark_key_fallback(self, mock_secret):
mock_secret.side_effect = lambda name: "ark_secret" if name == "ARK_API_KEY" else None
headers = self.config.validate_environment(
headers={},
model=self.model,
messages=[],
optional_params={},
litellm_params={},
api_key=None,
)
assert headers["Authorization"] == "Bearer ark_secret"
@patch(MODULE)
def test_validate_environment_no_key_raises(self, mock_secret):
mock_secret.return_value = None
with pytest.raises(ValueError, match="BYTEPLUS_API_KEY or ARK_API_KEY is not set"):
self.config.validate_environment(
headers={},
model=self.model,
messages=[],
optional_params={},
litellm_params={},
api_key=None,
)
def test_transform_request(self):
result = self.config.transform_image_generation_request(
model=self.model,
prompt="a cat surfing",
optional_params={"size": "2048x2048", "watermark": False},
litellm_params={},
headers={},
)
assert result == {
"model": self.model,
"prompt": "a cat surfing",
"size": "2048x2048",
"watermark": False,
}
def test_transform_response_success(self):
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"data": [
{"url": "https://img.example.com/1.png"},
{"b64_json": "abc123"},
]
}
result = self.config.transform_image_generation_response(
model=self.model,
raw_response=mock_response,
model_response=ImageResponse(data=[]),
logging_obj=self.logging_obj,
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 2
assert result.data[0].url == "https://img.example.com/1.png"
assert result.data[0].b64_json is None
assert result.data[1].b64_json == "abc123"
assert result.data[1].url is None
def test_transform_response_non_200_raises(self):
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.text = "unauthorized"
mock_response.headers = {}
with pytest.raises(Exception, match="unauthorized"):
self.config.transform_image_generation_response(
model=self.model,
raw_response=mock_response,
model_response=ImageResponse(data=[]),
logging_obj=self.logging_obj,
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
def test_transform_response_api_error_body_raises(self):
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.return_value = {"error": {"message": "invalid model"}}
with pytest.raises(Exception, match="invalid model"):
self.config.transform_image_generation_response(
model=self.model,
raw_response=mock_response,
model_response=ImageResponse(data=[]),
logging_obj=self.logging_obj,
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
def test_transform_response_json_parse_error_raises(self):
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.side_effect = ValueError("bad json")
with pytest.raises(Exception, match="Failed to parse BytePlus"):
self.config.transform_image_generation_response(
model=self.model,
raw_response=mock_response,
model_response=ImageResponse(data=[]),
logging_obj=self.logging_obj,
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)