fix(xai): resolve plain video id providers from the requested model

Video status forced an "openai" default whenever no provider was explicit, so
a plain xAI request id with ?model=xai/grok-imagine-video was sent to OpenAI on
the direct (non-router) path. The proxy now only defaults to openai when no
model is known, and the SDK resolves the provider from the model with
get_llm_provider instead of the grok-imagine name heuristics that lived in both
the proxy and the SDK

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
hx 2026-09-27 18:16:04 +08:00
parent 0722d6efdc
commit 09022930a3
6 changed files with 65 additions and 86 deletions

View file

@ -20,7 +20,6 @@ from litellm.proxy.video_endpoints.utils import (
encode_character_id_in_response,
extract_model_from_target_model_names,
get_custom_provider_from_data,
infer_video_provider_from_model,
resolve_video_request_model,
video_reference_to_id,
)
@ -273,7 +272,7 @@ async def video_status(
if resolved_model:
data["model"] = resolved_model
custom_llm_provider: Final = explicit_provider or infer_video_provider_from_model(resolved_model) or "openai"
custom_llm_provider: Final = explicit_provider or (None if resolved_model else "openai")
if custom_llm_provider:
data["custom_llm_provider"] = custom_llm_provider

View file

@ -10,15 +10,6 @@ class VideoModelIdResolver(Protocol):
def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None: ...
def infer_video_provider_from_model(model: str | None) -> str | None:
if not isinstance(model, str) or not model:
return None
unprefixed: Final = model.split("/", 1)[-1]
if unprefixed.startswith("grok-imagine-video"):
return "xai"
return None
def resolve_video_request_model(
*,
model_id_from_decoded: str | None,

View file

@ -30,51 +30,25 @@ from litellm.videos.utils import VideoGenerationRequestUtils
llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler()
def _litellm_provider_from_cost_entry(info: object) -> str | None:
if isinstance(info, dict):
catalog_provider: Final = info.get("litellm_provider")
if isinstance(catalog_provider, str) and catalog_provider:
return catalog_provider
return None
def _provider_from_prefixed_model(model: str) -> str | None:
if "/" not in model:
def _provider_from_model(model: object) -> str | None:
if not isinstance(model, str) or not model:
return None
try:
_, provider, _, _ = get_llm_provider(model=model)
except Exception:
except litellm.BadRequestError:
return None
return provider
def _custom_llm_provider_from_model(model: str) -> str | None:
def _provider_for_video_id(video_id: str, custom_llm_provider: str | None, model: object) -> str:
return (
_provider_from_prefixed_model(model)
or _litellm_provider_from_cost_entry(litellm.model_cost.get(model))
or _litellm_provider_from_cost_entry(litellm.model_cost.get(f"xai/{model}"))
or ("xai" if model.startswith("grok-imagine-video") else None)
custom_llm_provider
or decode_video_id_with_provider(video_id).get("custom_llm_provider")
or _provider_from_model(model)
or "openai"
)
def _provider_for_video_id(
video_id: str,
custom_llm_provider: str | None,
model: object | None = None,
) -> str:
if custom_llm_provider is not None:
return custom_llm_provider
decoded: Final = decode_video_id_with_provider(video_id)
from_id: Final = decoded.get("custom_llm_provider")
if from_id:
return from_id
if isinstance(model, str) and model:
from_model: Final = _custom_llm_provider_from_model(model)
if from_model:
return from_model
return "openai"
##### Video Generation #######################
@client
async def avideo_generation(

View file

@ -328,35 +328,21 @@ async def test_status__resolve_fail_keeps_decoded_model_id(harness):
@pytest.mark.asyncio
async def test_status__query_model_on_plain_id(harness):
async def test_status__query_model_on_plain_id_leaves_provider_to_the_model(harness):
# an "openai" default here would override the provider the router / SDK derive from the model
await call_status(
harness,
"9b444cea-aaaa-bbbb-cccc-dddddddddddd",
query={"model": "grok-imagine-video-1.5"},
query={"model": "xai/grok-imagine-video-1.5"},
)
harness.resolve_model.assert_not_called()
assert harness.processor_data() == {
"video_id": "9b444cea-aaaa-bbbb-cccc-dddddddddddd",
"custom_llm_provider": "xai",
"model": "grok-imagine-video-1.5",
"model": "xai/grok-imagine-video-1.5",
}
@pytest.mark.asyncio
async def test_status__query_model_grok_imagine_does_not_default_openai_before_inference(
harness,
):
await call_status(
harness,
"video_plain_xai",
query={"model": "grok-imagine-video"},
)
assert harness.processor_data()["custom_llm_provider"] == "xai"
assert harness.processor_data()["model"] == "grok-imagine-video"
# =========================================================================== #
# GET /v1/videos/{video_id}/content - video_content #
# =========================================================================== #

View file

@ -13,13 +13,14 @@ is encode_character_id_with_provider, which runs for real; encoding assertions
are checked by the genuine decode round-trip.
"""
import pytest
from litellm.proxy.video_endpoints.utils import (
encode_character_id_in_response,
extract_model_from_target_model_names,
get_custom_provider_from_data,
infer_video_provider_from_model,
resolve_video_request_model,
video_reference_to_id,
)
@ -74,21 +75,6 @@ def test_resolve_video_request_model__query_model_on_plain_id():
)
@pytest.mark.parametrize(
"model,expected",
[
("grok-imagine-video", "xai"),
("grok-imagine-video-1.5", "xai"),
("xai/grok-imagine-video", "xai"),
("sora-2", None),
(None, None),
("", None),
],
)
def test_infer_video_provider_from_model(model, expected):
assert infer_video_provider_from_model(model) == expected
# =========================================================================== #
# extract_model_from_target_model_names
# =========================================================================== #
@ -164,7 +150,12 @@ def test_provider__falsy_top_level_falls_through_to_extra_body(falsy):
def test_provider__from_extra_body_dict():
assert get_custom_provider_from_data({"extra_body": {"custom_llm_provider": "bedrock"}}) == "bedrock"
assert (
get_custom_provider_from_data(
{"extra_body": {"custom_llm_provider": "bedrock"}}
)
== "bedrock"
)
def test_provider__from_extra_body_json_string():
@ -182,7 +173,10 @@ def test_provider__json_string_parsing_to_non_dict_is_none():
def test_provider__extra_body_provider_not_a_string_is_none():
assert get_custom_provider_from_data({"extra_body": {"custom_llm_provider": 123}}) is None
assert (
get_custom_provider_from_data({"extra_body": {"custom_llm_provider": 123}})
is None
)
@pytest.mark.parametrize(
@ -214,7 +208,9 @@ def test_encode__dict_with_id_mutates_in_place_and_preserves_other_keys():
assert out is response # same dict, mutated in place
assert out["object"] == "character" and out["name"] == "hero"
assert out["id"] == encode_character_id_with_provider("char_raw", "azure", "model-1")
assert out["id"] == encode_character_id_with_provider(
"char_raw", "azure", "model-1"
)
decoded = decode_character_id_with_provider(out["id"])
assert decoded["custom_llm_provider"] == "azure"
assert decoded["model_id"] == "model-1"

View file

@ -39,6 +39,7 @@ import pytest
import litellm
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.videos.main import CharacterObject, VideoObject
from litellm.types.videos.utils import encode_video_id_with_provider
@ -178,15 +179,47 @@ def test_video_content__plain_id_defaults_to_openai(seams):
assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "openai"
def test_video_content__plain_id_with_grok_model_uses_xai(seams):
@pytest.fixture
def real_provider_resolution(seams):
with patch.object(videos_main, "get_llm_provider", get_llm_provider):
yield seams
def test_video_content__plain_id_with_xai_model_uses_xai(real_provider_resolution):
seams = real_provider_resolution
videos_main.video_content(
video_id="9b444cea-aaaa-bbbb-cccc-dddddddddddd",
model="grok-imagine-video-1.5",
model="xai/grok-imagine-video-1.5",
)
assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "xai"
def test_video_status__plain_id_with_xai_model_uses_xai(real_provider_resolution):
seams = real_provider_resolution
videos_main.video_status(
video_id="9b444cea-aaaa-bbbb-cccc-dddddddddddd",
model="xai/grok-imagine-video-1.5",
)
assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "xai"
assert seams.get_config.call_args.kwargs["provider"] == litellm.LlmProviders.XAI
def test_video_status__plain_id_with_unresolvable_model_defaults_to_openai(real_provider_resolution):
seams = real_provider_resolution
videos_main.video_status(video_id="video_plain", model="not-a-known-model")
assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "openai"
def test_video_status__decoded_provider_beats_model(real_provider_resolution):
seams = real_provider_resolution
videos_main.video_status(video_id=AZURE_VIDEO_ID, model="xai/grok-imagine-video-1.5")
assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "azure"
def test_video_remix__dispatch_and_provider_from_id(seams):
result = videos_main.video_remix(video_id=AZURE_VIDEO_ID, prompt="new colors")
@ -387,14 +420,14 @@ async def test_avideo_content__pre_decodes_provider_before_delegating():
@pytest.mark.asyncio
async def test_avideo_content__plain_id_with_grok_model_uses_xai():
async def test_avideo_content__plain_id_with_xai_model_uses_xai():
sentinel = b"mp4-bytes"
with patch.object(
videos_main, "video_content", MagicMock(return_value=sentinel)
) as sync:
result = await videos_main.avideo_content(
video_id="9b444cea-aaaa-bbbb-cccc-dddddddddddd",
model="grok-imagine-video-1.5",
model="xai/grok-imagine-video-1.5",
)
assert result is sentinel