mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
0722d6efdc
commit
09022930a3
6 changed files with 65 additions and 86 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 #
|
||||
# =========================================================================== #
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue