diff --git a/litellm/__init__.py b/litellm/__init__.py index 6e2a03b7c7c..c0bd073e54b 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() +pruna_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") == "pruna": + pruna_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 + | pruna_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, + "pruna": pruna_models, "cometapi": cometapi_models, "oci": oci_models, "volcengine": volcengine_models, diff --git a/litellm/images/main.py b/litellm/images/main.py index 17ea9aa177b..c52825b4f71 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.PRUNA, litellm.LlmProviders.AIML, litellm.LlmProviders.GEMINI, litellm.LlmProviders.FAL_AI, diff --git a/litellm/llms/pruna/__init__.py b/litellm/llms/pruna/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/pruna/image_generation/__init__.py b/litellm/llms/pruna/image_generation/__init__.py new file mode 100644 index 00000000000..d79a4515bc6 --- /dev/null +++ b/litellm/llms/pruna/image_generation/__init__.py @@ -0,0 +1,13 @@ +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) + +from .transformation import PrunaImageGenerationConfig + +__all__ = [ + "PrunaImageGenerationConfig", +] + + +def get_pruna_image_generation_config(model: str) -> BaseImageGenerationConfig: + return PrunaImageGenerationConfig() diff --git a/litellm/llms/pruna/image_generation/transformation.py b/litellm/llms/pruna/image_generation/transformation.py new file mode 100644 index 00000000000..c34e2cdcb9f --- /dev/null +++ b/litellm/llms/pruna/image_generation/transformation.py @@ -0,0 +1,148 @@ +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://api.pruna.ai" +PREDICTIONS_ENDPOINT = "v1/predictions" + + +class PrunaImageGenerationConfig(BaseImageGenerationConfig): + """ + Configuration for Pruna AI image generation. + + Pruna is not OpenAI-compatible. The model is passed via a `Model` header, + auth via an `apikey` header, and `Try-Sync: true` returns the result inline + within 60 seconds instead of an async prediction id + https://docs.api.pruna.ai/guides/models/p-image + """ + + def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]: + return ["size"] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + size = non_default_params.get("size") + dims = self._size_to_dimensions(size) + passthrough = { + k: v for k, v in non_default_params.items() if k not in optional_params and k not in ("n", "size") + } + return {**optional_params, **passthrough, **dims} + + def _size_to_dimensions(self, size: object) -> dict: + if not isinstance(size, str) or "x" not in size: + return {} + width, _, height = size.partition("x") + if not width.isdigit() or not height.isdigit(): + return {} + return {"width": int(width), "height": int(height), "aspect_ratio": "custom"} + + 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("PRUNA_API_BASE") or DEFAULT_API_BASE).rstrip("/") + if base_url.endswith(PREDICTIONS_ENDPOINT): + return base_url + return f"{base_url}/{PREDICTIONS_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("PRUNA_API_KEY") + if not final_api_key: + raise ValueError("PRUNA_API_KEY is not set") + headers["apikey"] = final_api_key + headers["Model"] = model + headers["Try-Sync"] = "true" + 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 {"input": {"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 Pruna image generation response: {e}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + generation_url = response_data.get("generation_url") + if response_data.get("status") != "succeeded" or not generation_url: + raise self.get_error_class( + error_message=(f"Pruna synchronous generation did not complete in time; response: {response_data}"), + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + model_response.data = [ImageObject(url=self._absolute_url(raw_response, generation_url))] + return model_response + + def _absolute_url(self, raw_response: httpx.Response, generation_url: str) -> str: + if generation_url.startswith("http"): + return generation_url + request_url = raw_response.request.url + return f"{request_url.scheme}://{request_url.host}{generation_url}" diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index dedb9bbf40a..638a2ba967e 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -32088,6 +32088,15 @@ "/v1/ocr" ] }, + "pruna/p-image": { + "litellm_provider": "pruna", + "mode": "image_generation", + "output_cost_per_image": 0.005, + "source": "https://docs.api.pruna.ai/guides/models/p-image", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, "recraft/recraftv2": { "litellm_provider": "recraft", "mode": "image_generation", diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 88b3a39844f..1d831c2a9a5 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" + PRUNA = "pruna" # Create a set of all provider values for quick lookup diff --git a/litellm/utils.py b/litellm/utils.py index 0636d3683b7..995c46fdf7d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8579,6 +8579,12 @@ class ProviderConfigManager: ) return get_recraft_image_generation_config(model) + elif LlmProviders.PRUNA == provider: + from litellm.llms.pruna.image_generation import ( + get_pruna_image_generation_config, + ) + + return get_pruna_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..24edcf9aaf0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -32179,6 +32179,15 @@ "/v1/ocr" ] }, + "pruna/p-image": { + "litellm_provider": "pruna", + "mode": "image_generation", + "output_cost_per_image": 0.005, + "source": "https://docs.api.pruna.ai/guides/models/p-image", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, "recraft/recraftv2": { "litellm_provider": "recraft", "mode": "image_generation", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 65db63dc045..df3ac2c7da5 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2087,6 +2087,22 @@ "interactions": true } }, + "pruna": { + "display_name": "Pruna AI (`pruna`)", + "url": "https://docs.litellm.ai/docs/providers/pruna", + "endpoints": { + "chat_completions": false, + "messages": false, + "responses": false, + "embeddings": false, + "image_generations": true, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false + } + }, "recraft": { "display_name": "Recraft (`recraft`)", "url": "https://docs.litellm.ai/docs/providers/recraft", diff --git a/tests/test_litellm/llms/pruna/image_generation/test_pruna_image_gen_transformation.py b/tests/test_litellm/llms/pruna/image_generation/test_pruna_image_gen_transformation.py new file mode 100644 index 00000000000..8e491caabbe --- /dev/null +++ b/tests/test_litellm/llms/pruna/image_generation/test_pruna_image_gen_transformation.py @@ -0,0 +1,253 @@ +import os +import sys +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +import litellm +from litellm import get_llm_provider +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.pruna.image_generation.transformation import ( + DEFAULT_API_BASE, + PREDICTIONS_ENDPOINT, + PrunaImageGenerationConfig, +) +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + +MODULE = "litellm.llms.pruna.image_generation.transformation.get_secret_str" + + +def _response(status_code: int, *, json_body=None, text_body=None) -> httpx.Response: + request = httpx.Request("POST", "https://api.pruna.ai/v1/predictions") + if json_body is not None: + return httpx.Response(status_code, json=json_body, request=request) + return httpx.Response(status_code, content=text_body, request=request) + + +class TestPrunaImageGenerationTransformation: + def setup_method(self): + self.config = PrunaImageGenerationConfig() + self.model = "p-image" + self.logging_obj = MagicMock() + + def test_provider_routing(self): + model, provider, _, _ = get_llm_provider("pruna/p-image") + assert provider == "pruna" + assert model == "p-image" + + def test_provider_config_registered(self): + config = ProviderConfigManager.get_provider_image_generation_config( + model=self.model, + provider=LlmProviders.PRUNA, + ) + assert isinstance(config, PrunaImageGenerationConfig) + + def test_supported_params(self): + assert self.config.get_supported_openai_params(self.model) == ["size"] + + def test_map_openai_params_size_to_custom_dimensions(self): + result = self.config.map_openai_params( + non_default_params={"size": "1024x768"}, + optional_params={}, + model=self.model, + drop_params=False, + ) + assert result == {"width": 1024, "height": 768, "aspect_ratio": "custom"} + + def test_map_openai_params_passthrough_and_drops_n(self): + result = self.config.map_openai_params( + non_default_params={"n": 3, "aspect_ratio": "16:9", "seed": 7}, + optional_params={}, + model=self.model, + drop_params=False, + ) + assert result == {"aspect_ratio": "16:9", "seed": 7} + + def test_map_openai_params_ignores_invalid_size(self): + result = self.config.map_openai_params( + non_default_params={"size": "auto"}, + optional_params={}, + model=self.model, + drop_params=False, + ) + assert result == {} + + @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}/{PREDICTIONS_ENDPOINT}" + + @patch(MODULE) + def test_get_complete_url_no_double_endpoint(self, mock_secret): + mock_secret.return_value = None + base = f"{DEFAULT_API_BASE}/{PREDICTIONS_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_sets_pruna_headers(self, mock_secret): + headers = self.config.validate_environment( + headers={}, + model=self.model, + messages=[], + optional_params={}, + litellm_params={}, + api_key="my_key", + ) + assert headers["apikey"] == "my_key" + assert headers["Model"] == "p-image" + assert headers["Try-Sync"] == "true" + assert headers["Content-Type"] == "application/json" + mock_secret.assert_not_called() + + @patch(MODULE) + def test_validate_environment_env_fallback(self, mock_secret): + mock_secret.return_value = "env_key" + headers = self.config.validate_environment( + headers={}, + model=self.model, + messages=[], + optional_params={}, + litellm_params={}, + ) + assert headers["apikey"] == "env_key" + + @patch(MODULE) + def test_validate_environment_missing_key_raises(self, mock_secret): + mock_secret.return_value = None + with pytest.raises(ValueError, match="PRUNA_API_KEY is not set"): + self.config.validate_environment( + headers={}, + model=self.model, + messages=[], + optional_params={}, + litellm_params={}, + ) + + def test_transform_request_wraps_input(self): + result = self.config.transform_image_generation_request( + model=self.model, + prompt="a lion at sunset", + optional_params={"aspect_ratio": "16:9", "seed": 7}, + litellm_params={}, + headers={}, + ) + assert result == {"input": {"prompt": "a lion at sunset", "aspect_ratio": "16:9", "seed": 7}} + + def test_transform_response_builds_absolute_url(self): + raw = _response( + 200, + json_body={"status": "succeeded", "generation_url": "/v1/predictions/delivery/abc/output.jpg"}, + ) + result = self.config.transform_image_generation_response( + model=self.model, + raw_response=raw, + model_response=litellm.ImageResponse(), + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + assert [img.url for img in result.data] == ["https://api.pruna.ai/v1/predictions/delivery/abc/output.jpg"] + + def test_transform_response_keeps_absolute_url(self): + raw = _response( + 200, + json_body={"status": "succeeded", "generation_url": "https://cdn.pruna.ai/out.jpg"}, + ) + result = self.config.transform_image_generation_response( + model=self.model, + raw_response=raw, + model_response=litellm.ImageResponse(), + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + assert [img.url for img in result.data] == ["https://cdn.pruna.ai/out.jpg"] + + def test_transform_response_async_not_completed_raises(self): + raw = _response( + 200, + json_body={"id": "abc", "get_url": "https://api.pruna.ai/v1/predictions/status/abc"}, + ) + with pytest.raises(BaseLLMException): + self.config.transform_image_generation_response( + model=self.model, + raw_response=raw, + model_response=litellm.ImageResponse(), + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + def test_transform_response_non_200_raises(self): + raw = _response(401, text_body=b"unauthorized") + with pytest.raises(BaseLLMException): + self.config.transform_image_generation_response( + model=self.model, + raw_response=raw, + model_response=litellm.ImageResponse(), + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + def test_transform_response_json_parse_error_raises(self): + raw = _response(200, text_body=b"not json") + with pytest.raises(BaseLLMException): + self.config.transform_image_generation_response( + model=self.model, + raw_response=raw, + model_response=litellm.ImageResponse(), + logging_obj=self.logging_obj, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + def test_image_generation_dispatches_to_pruna_handler(self): + fake_response = MagicMock() + + with patch.object( + litellm.images.main.llm_http_handler, + "image_generation_handler", + return_value=fake_response, + ) as mock_handler: + result = litellm.image_generation( + model="pruna/p-image", + prompt="a lion at sunset", + api_key="sk-test", + ) + + assert result is fake_response + mock_handler.assert_called_once() + kwargs = mock_handler.call_args.kwargs + assert kwargs["custom_llm_provider"] == "pruna" + assert kwargs["model"] == "p-image" + assert kwargs["prompt"] == "a lion at sunset" + assert isinstance(kwargs["image_generation_provider_config"], PrunaImageGenerationConfig)