feat(pruna): add p-image image generation support

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:51:35 +00:00
parent 1ebf2a78a9
commit 2e20a427c9
11 changed files with 461 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()
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,

View file

@ -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,

View file

View file

@ -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()

View file

@ -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}"

View file

@ -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",

View file

@ -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

View file

@ -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,

View file

@ -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",

View file

@ -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",

View file

@ -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)