fix(proxy): re-encode video id with routed model_id

This commit is contained in:
Yash Raj Pandey 2026-07-16 12:05:45 -04:00
parent cd63c7e5a7
commit 7a0c9c4cba
2 changed files with 127 additions and 5 deletions

View file

@ -25,11 +25,36 @@ from litellm.proxy.video_endpoints.utils import (
from litellm.types.videos.utils import (
decode_character_id_with_provider,
decode_video_id_with_provider,
encode_video_id_with_provider,
)
router: Final = APIRouter()
def _hidden_param(response: object, key: str) -> str | None:
"""Read one string entry from a response's _hidden_params, if it has any."""
params: Final = getattr(response, "_hidden_params", None)
if isinstance(params, dict):
value: Final = params.get(key)
if isinstance(value, str):
return value
return None
def _apply_routed_video_id(response: object, fallback_model: str | None) -> None:
"""Re-encode the returned video id so it carries the provider and routed model_id."""
decoded: Final = decode_video_id_with_provider(response["id"] if isinstance(response, dict) else response.id)
encoded: Final = encode_video_id_with_provider(
video_id=decoded.get("video_id", ""),
provider=_hidden_param(response, "custom_llm_provider") or decoded.get("custom_llm_provider") or "openai",
model_id=_hidden_param(response, "model_id") or fallback_model,
)
if isinstance(response, dict):
response["id"] = encoded
else:
response.id = encoded
@router.post(
"/v1/videos",
dependencies=[Depends(user_api_key_auth)],
@ -89,7 +114,7 @@ async def video_generation(
# Process request using ProxyBaseLLMRequestProcessing
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response: Final = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -107,6 +132,7 @@ async def video_generation(
user_api_base=user_api_base,
version=version,
)
_apply_routed_video_id(response, data.get("model"))
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
@ -114,6 +140,8 @@ async def video_generation(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
else:
return response
@router.get(
@ -478,7 +506,7 @@ async def video_remix(
# Process request using ProxyBaseLLMRequestProcessing
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response: Final = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -496,6 +524,7 @@ async def video_remix(
user_api_base=user_api_base,
version=version,
)
_apply_routed_video_id(response, data.get("model"))
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
@ -503,6 +532,8 @@ async def video_remix(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
else:
return response
@router.post(
@ -789,7 +820,7 @@ async def video_edit(
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response: Final = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -807,6 +838,7 @@ async def video_edit(
user_api_base=user_api_base,
version=version,
)
_apply_routed_video_id(response, data.get("model"))
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
@ -814,6 +846,8 @@ async def video_edit(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
else:
return response
@router.post(
@ -884,7 +918,7 @@ async def video_extension(
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response: Final = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -902,6 +936,7 @@ async def video_extension(
user_api_base=user_api_base,
version=version,
)
_apply_routed_video_id(response, data.get("model"))
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
@ -909,3 +944,5 @@ async def video_extension(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
else:
return response

View file

@ -41,7 +41,9 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.utils import ProxyLogging
from litellm.router import Router
from litellm.types.videos.main import VideoObject
from litellm.types.videos.utils import (
decode_video_id_with_provider,
encode_character_id_with_provider,
encode_video_id_with_provider,
)
@ -67,7 +69,7 @@ AZURE_CHARACTER_ID = encode_character_id_with_provider(
RESOLVED_MODELS: Dict[str, str] = {VIDEO_MODEL_ID: "azure-sora"}
# Sentinel propagated by base_process for the passthrough endpoints.
SENTINEL = object()
SENTINEL = VideoObject(id="video_raw", object="video", status="processing")
class FakeRequest:
@ -111,6 +113,7 @@ class Harness:
@pytest.fixture
def harness():
SENTINEL.id = "video_raw"
logging = MagicMock(spec=ProxyLogging)
router = MagicMock(spec=Router)
@ -229,6 +232,26 @@ async def test_generation__route_type_data_and_no_provider_default(harness):
harness.batch_to_bytesio.assert_not_called()
@pytest.mark.asyncio
async def test_generation__reencodes_id_with_model_id(harness):
response = VideoObject(
id=encode_video_id_with_provider("video_raw", "openai", None),
object="video",
status="processing",
)
response._hidden_params = {
"custom_llm_provider": "openai",
"model_id": VIDEO_MODEL_ID,
}
harness.base_process.return_value = response
resp = await call_generation(harness, body={"model": "sora-2"})
decoded_video_id = decode_video_id_with_provider(resp.id)
assert decoded_video_id["model_id"] == VIDEO_MODEL_ID
assert decoded_video_id["video_id"] == "video_raw"
@pytest.mark.asyncio
async def test_generation__input_reference_attached(harness):
body = {"model": "sora-2", "prompt": "a sunset"}
@ -400,6 +423,28 @@ async def test_edit__extracts_nested_video_id_full_contract(harness):
}
@pytest.mark.asyncio
async def test_edit__reencodes_id_with_model_id(harness):
response = VideoObject(
id=encode_video_id_with_provider("video_raw", "openai", None),
object="video",
status="processing",
)
response._hidden_params = {
"custom_llm_provider": "openai",
"model_id": VIDEO_MODEL_ID,
}
harness.base_process.return_value = response
resp = await call_edit(
harness, body={"prompt": "brighter", "video": {"id": "video_plain"}}
)
decoded_video_id = decode_video_id_with_provider(resp.id)
assert decoded_video_id["model_id"] == VIDEO_MODEL_ID
assert decoded_video_id["video_id"] == "video_raw"
@pytest.mark.asyncio
async def test_edit__provider_from_body_data_for_plain_id(harness):
"""For a plain id, get_custom_provider_from_data (run for real) pulls the
@ -543,6 +588,26 @@ async def test_remix__model_encoded_id_full_contract(harness):
}
@pytest.mark.asyncio
async def test_remix__reencodes_id_with_model_id(harness):
response = VideoObject(
id=encode_video_id_with_provider("video_raw", "openai", None),
object="video",
status="processing",
)
response._hidden_params = {
"custom_llm_provider": "openai",
"model_id": VIDEO_MODEL_ID,
}
harness.base_process.return_value = response
resp = await call_remix(harness, "video_plain", body={"prompt": "new colors"})
decoded_video_id = decode_video_id_with_provider(resp.id)
assert decoded_video_id["model_id"] == VIDEO_MODEL_ID
assert decoded_video_id["video_id"] == "video_raw"
@pytest.mark.asyncio
async def test_remix__provider_from_body_data_not_request_body_reader(harness):
"""remix resolves the provider from data.get('custom_llm_provider'), never
@ -701,3 +766,23 @@ async def test_extension__extracts_nested_video_id_full_contract(harness):
"custom_llm_provider": "azure",
"model": "azure-sora",
}
@pytest.mark.asyncio
async def test_extension__reencodes_id_with_model_id(harness):
response = VideoObject(
id=encode_video_id_with_provider("video_raw", "openai", None),
object="video",
status="processing",
)
response._hidden_params = {
"custom_llm_provider": "openai",
"model_id": VIDEO_MODEL_ID,
}
harness.base_process.return_value = response
resp = await call_extension(harness, body={"prompt": "continue"})
decoded_video_id = decode_video_id_with_provider(resp.id)
assert decoded_video_id["model_id"] == VIDEO_MODEL_ID
assert decoded_video_id["video_id"] == "video_raw"