mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): re-encode video id with routed model_id
This commit is contained in:
parent
cd63c7e5a7
commit
7a0c9c4cba
2 changed files with 127 additions and 5 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue