diff --git a/litellm/llms/xai/image_edit/transformation.py b/litellm/llms/xai/image_edit/transformation.py index b20382a7acf..bdb7ae35510 100644 --- a/litellm/llms/xai/image_edit/transformation.py +++ b/litellm/llms/xai/image_edit/transformation.py @@ -1,6 +1,6 @@ import base64 from io import BufferedReader, BytesIO -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable import httpx from httpx._types import RequestFiles @@ -32,6 +32,11 @@ _SIZE_TO_ASPECT_RATIO: Final = { # mutable-ok: provider JSON body and base-clas _XAI_NATIVE_PARAMS: Final = frozenset({"aspect_ratio", "n", "resolution"}) +@runtime_checkable +class _Readable(Protocol): + def read(self) -> bytes | str: ... + + def _read_seekable(image: BytesIO | BufferedReader) -> bytes: current_pos: Final = image.tell() image.seek(0) @@ -222,11 +227,12 @@ class XAIImageEditConfig(BaseImageEditConfig): if isinstance(image, str): return {"url": image} # mutable-ok: provider JSON body and base-class dict signature if isinstance(image, dict): - if image.get("url"): - return {"url": str(image["url"])} # mutable-ok: provider JSON body and base-class dict signature - if image.get("file_id"): - file_id: Final = str(image["file_id"]) - return {"file_id": file_id} # mutable-ok: provider JSON body and base-class dict signature + url: Final = image.get("url") + if url: + return {"url": str(url)} # mutable-ok: provider JSON body and base-class dict signature + file_id: Final = image.get("file_id") + if file_id: + return {"file_id": str(file_id)} # mutable-ok: provider JSON body and base-class dict signature mime: Final = ImageEditRequestUtils.get_image_content_type(image) encoded: Final = base64.b64encode(self._read_all_bytes(image)).decode("utf-8") @@ -239,7 +245,7 @@ class XAIImageEditConfig(BaseImageEditConfig): return bytes(image) if isinstance(image, (BytesIO, BufferedReader)): return _read_seekable(image) - if hasattr(image, "read"): + if isinstance(image, _Readable): raw: Final = image.read() if isinstance(raw, str): return raw.encode("utf-8") diff --git a/litellm/llms/xai/videos/transformation.py b/litellm/llms/xai/videos/transformation.py index 5b4cfb94fb4..4ad01b4ea2c 100644 --- a/litellm/llms/xai/videos/transformation.py +++ b/litellm/llms/xai/videos/transformation.py @@ -11,8 +11,6 @@ from litellm.litellm_core_utils.url_utils import async_safe_get, encode_url_path from litellm.llms.base_llm.videos.transformation import BaseVideoConfig from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, - HTTPHandler, - _get_httpx_client, get_async_httpx_client, ) from litellm.llms.xai.common_utils import XAIModelInfo @@ -248,7 +246,6 @@ class XAIVideoConfig(BaseVideoConfig): video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, model) usage_body: Final = usage if isinstance(usage, dict) else None video_obj.usage = usage_body or {} # mutable-ok: provider JSON body and base-class dict signature - video_obj._hidden_params["video_url"] = None return video_obj def _video_resource_url(self, api_base: str, video_id: str) -> str: @@ -342,8 +339,7 @@ class XAIVideoConfig(BaseVideoConfig): url: Final = self._video_cdn_url(raw_response) if url is None: return raw_response.content - httpx_client: Final[HTTPHandler] = _get_httpx_client() - video_response: Final = safe_get(httpx_client, url) + video_response: Final = safe_get(litellm.module_level_client, url) video_response.raise_for_status() return video_response.content diff --git a/litellm/proxy/video_endpoints/utils.py b/litellm/proxy/video_endpoints/utils.py index 6e305bdff2a..b76a20d2642 100644 --- a/litellm/proxy/video_endpoints/utils.py +++ b/litellm/proxy/video_endpoints/utils.py @@ -1,5 +1,5 @@ from collections.abc import Mapping, Sequence -from typing import Final, Protocol +from typing import Any, Final, Protocol import orjson @@ -71,7 +71,7 @@ def get_custom_provider_from_data(data: Mapping[str, object]) -> str | None: return None -def encode_character_id_in_response(response: object, custom_llm_provider: str, model_id: str | None) -> object: +def encode_character_id_in_response(response: Any, custom_llm_provider: str, model_id: str | None) -> Any: if isinstance(response, dict) and response.get("id"): response["id"] = encode_character_id_with_provider( character_id=response["id"], diff --git a/tests/unit/llms/xai/test_xai_image_edit.py b/tests/unit/llms/xai/test_xai_image_edit.py index 575cbf459b6..3e996b94f4b 100644 --- a/tests/unit/llms/xai/test_xai_image_edit.py +++ b/tests/unit/llms/xai/test_xai_image_edit.py @@ -1,3 +1,4 @@ +import base64 from unittest.mock import MagicMock, patch import httpx @@ -134,6 +135,24 @@ def test_transform_bytes_to_data_uri_and_response(): assert response.data[0].url == "https://imgen.x.ai/edited.jpeg" +def test_transform_duck_typed_reader_to_data_uri(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + class Reader: + def read(self) -> bytes: + return b"\x89PNG\r\n\x1a\nfakepng" + + data, _ = XAIImageEditConfig().transform_image_edit_request( + model="grok-imagine-image", + prompt="make it night", + image=Reader(), + image_edit_optional_request_params={}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert data["image"]["url"] == "data:image/png;base64," + base64.b64encode(b"\x89PNG\r\n\x1a\nfakepng").decode() + + def test_transform_http_url_passthrough(): from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig diff --git a/tests/unit/llms/xai/test_xai_video_generation.py b/tests/unit/llms/xai/test_xai_video_generation.py index 261ce2e88c0..446b6fe47e2 100644 --- a/tests/unit/llms/xai/test_xai_video_generation.py +++ b/tests/unit/llms/xai/test_xai_video_generation.py @@ -196,13 +196,13 @@ def test_content_response_fetches_cdn_via_shared_client(): video_resp = MagicMock() video_resp.content = b"mp4-bytes" video_resp.raise_for_status.return_value = None - with patch( - "litellm.llms.xai.videos.transformation._get_httpx_client", - return_value=cdn, - ), patch( - "litellm.llms.xai.videos.transformation.safe_get", - return_value=video_resp, - ) as safe_get: + with ( + patch("litellm.module_level_client", cdn), + patch( + "litellm.llms.xai.videos.transformation.safe_get", + return_value=video_resp, + ) as safe_get, + ): body = XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock()) assert body == b"mp4-bytes" safe_get.assert_called_once_with(cdn, "https://vidgen.x.ai/x.mp4") @@ -219,20 +219,20 @@ async def test_async_content_response_does_not_use_sync_client(): video_resp = MagicMock() video_resp.content = b"async-mp4" video_resp.raise_for_status.return_value = None - with patch( - "litellm.llms.xai.videos.transformation.get_async_httpx_client", - return_value=async_client, - ), patch( - "litellm.llms.xai.videos.transformation._get_httpx_client", - ) as sync_client, patch( - "litellm.llms.xai.videos.transformation.async_safe_get", - new=AsyncMock(return_value=video_resp), - ) as async_safe_get: - body = await XAIVideoConfig().async_transform_video_content_response( - status, logging_obj=MagicMock() - ) + with ( + patch( + "litellm.llms.xai.videos.transformation.get_async_httpx_client", + return_value=async_client, + ), + patch("litellm.module_level_client") as sync_client, + patch( + "litellm.llms.xai.videos.transformation.async_safe_get", + new=AsyncMock(return_value=video_resp), + ) as async_safe_get, + ): + body = await XAIVideoConfig().async_transform_video_content_response(status, logging_obj=MagicMock()) assert body == b"async-mp4" - sync_client.assert_not_called() + assert sync_client.method_calls == [] async_safe_get.assert_awaited_once_with(async_client, "https://vidgen.x.ai/x.mp4") @@ -265,10 +265,7 @@ def test_content_response_rejects_internal_cdn_host(): def boom(*args, **kwargs): raise AssertionError("unsafe CDN fetch must not run") - with patch( - "litellm.llms.xai.videos.transformation._get_httpx_client", - return_value=MagicMock(get=boom), - ): + with patch("litellm.module_level_client", MagicMock(get=boom)): with pytest.raises((SSRFError, ValueError)): XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock()) @@ -309,9 +306,7 @@ async def test_async_content_response_rejects_internal_cdn_host(): return_value=MagicMock(get=boom), ): with pytest.raises((SSRFError, ValueError)): - await XAIVideoConfig().async_transform_video_content_response( - status, logging_obj=MagicMock() - ) + await XAIVideoConfig().async_transform_video_content_response(status, logging_obj=MagicMock()) @pytest.mark.asyncio @@ -329,9 +324,7 @@ async def test_async_content_response_fetches_public_cdn_via_async_safe_get(): new=AsyncMock(return_value=video_resp), create=True, ) as async_safe_get: - body = await XAIVideoConfig().async_transform_video_content_response( - status, logging_obj=MagicMock() - ) + body = await XAIVideoConfig().async_transform_video_content_response(status, logging_obj=MagicMock()) assert body == b"async-safe-mp4" async_safe_get.assert_awaited() assert async_safe_get.call_args.args[1] == "https://vidgen.x.ai/x.mp4"