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:
hx 2026-09-27 18:29:06 +08:00
parent b210c6bf4c
commit ab4de2f02a
5 changed files with 58 additions and 44 deletions

View file

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

View file

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

View file

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

View file

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

View file

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