fix(proxy): hand the embeddings failure hook the post-setup request data

This commit is contained in:
mateo-berri 2026-08-18 14:49:42 -07:00
parent 3a4d3a01af
commit e9355a7fe9
2 changed files with 45 additions and 10 deletions

View file

@ -10197,11 +10197,9 @@ async def embeddings(
"""
global proxy_logging_obj
data: Any = {}
data: Final = await _read_request_body(request=request)
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
# Use shared request body reading helper (same as chat/completions)
data = await _read_request_body(request=request)
### HANDLE TOKEN ARRAY INPUT DECODING ###
# This must happen BEFORE base_process_llm_request() since it modifies the input
router_model_names: Final = llm_router.model_names if llm_router is not None else []
@ -10245,10 +10243,6 @@ async def embeddings(
if hasattr(user_api_key_dict, "agent_id") and user_api_key_dict.agent_id is not None:
data["metadata"]["agent_id"] = user_api_key_dict.agent_id
# Use unified request processor (same as chat/completions and responses)
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
# Process the request with all optimizations (shared sessions, network tuning, etc.)
response: Final = await base_llm_response_processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
@ -10270,8 +10264,6 @@ async def embeddings(
return response
except Exception as e:
# Use unified error handler
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
raise await base_llm_response_processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,

View file

@ -11098,3 +11098,46 @@ async def test_moderations_reraises_proxy_exception_unwrapped():
assert exc_info.value.code == "400"
assert exc_info.value.param == "metadata"
mock_logging.post_call_failure_hook.assert_awaited_once()
class TestEmbeddingsFailureHookRequestData:
@pytest.mark.asyncio
async def test_failure_hook_gets_post_setup_data_with_logging_obj(self):
"""Request setup replaces the processor's data dict (adding the logging
object the failure hook needs to lift token usage from); the embeddings
exception handler must pass that replaced dict, not the raw request body
dict it was rebuilt from."""
from litellm.proxy._types import ProxyException
captured = {}
logging_obj_sentinel = MagicMock()
async def fake_process(self, **kwargs):
self.data = {**self.data, "litellm_logging_obj": logging_obj_sentinel}
captured["processor_data"] = self.data
raise RuntimeError("provider timeout")
with (
patch.object(
proxy_server_module,
"_read_request_body",
new=AsyncMock(return_value={"model": "my-embed", "input": "hello"}),
),
patch.object(
proxy_server_module.ProxyBaseLLMRequestProcessing,
"base_process_llm_request",
new=fake_process,
),
patch.object(proxy_server_module, "proxy_logging_obj") as mock_logging,
):
mock_logging.post_call_failure_hook = AsyncMock(return_value=None)
with pytest.raises(ProxyException):
await proxy_server_module.embeddings(
request=MagicMock(),
fastapi_response=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(),
)
hook_request_data = mock_logging.post_call_failure_hook.await_args.kwargs["request_data"]
assert hook_request_data is captured["processor_data"]
assert hook_request_data["litellm_logging_obj"] is logging_obj_sentinel