diff --git a/litellm/__init__.py b/litellm/__init__.py index 6e2a03b7c7c..5ebf3bf5e14 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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, diff --git a/litellm/images/main.py b/litellm/images/main.py index 17ea9aa177b..57b5497b177 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -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, diff --git a/litellm/llms/byteplus/__init__.py b/litellm/llms/byteplus/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/byteplus/image_generation/__init__.py b/litellm/llms/byteplus/image_generation/__init__.py new file mode 100644 index 00000000000..e8458f6851a --- /dev/null +++ b/litellm/llms/byteplus/image_generation/__init__.py @@ -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() diff --git a/litellm/llms/byteplus/image_generation/transformation.py b/litellm/llms/byteplus/image_generation/transformation.py new file mode 100644 index 00000000000..5cfd5950fb1 --- /dev/null +++ b/litellm/llms/byteplus/image_generation/transformation.py @@ -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 diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index dedb9bbf40a..811285ae19b 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 88b3a39844f..af350bfcb76 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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 diff --git a/litellm/utils.py b/litellm/utils.py index 0636d3683b7..f242e55b6df 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index e10dde793d1..60771a2c639 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/tests/test_litellm/llms/byteplus/image_generation/test_byteplus_image_gen_transformation.py b/tests/test_litellm/llms/byteplus/image_generation/test_byteplus_image_gen_transformation.py new file mode 100644 index 00000000000..b0ae27c5cf9 --- /dev/null +++ b/tests/test_litellm/llms/byteplus/image_generation/test_byteplus_image_gen_transformation.py @@ -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, + )