mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(xai): stay within the basedpyright budget in imagine transforms
Use the public module-level http client for CDN downloads, drop the dead create-time video_url hidden param, narrow file-like and reference image inputs without subscripting the FileTypes union, and restore the upstream signature of encode_character_id_in_response. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
b210c6bf4c
commit
ab4de2f02a
5 changed files with 58 additions and 44 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue