fix(proxy): sync normalized call_type into model_call_details for proxy-only errors

This commit is contained in:
Alexey
2026-03-18 23:49:07 +03:00
parent b00096f2a0
commit 71b687e00a
2 changed files with 78 additions and 4 deletions
+9 -4
View File
@@ -1863,28 +1863,33 @@ class ProxyLogging:
)
input: Union[list, str, dict] = ""
normalized_call_type: Optional[str] = None
if "messages" in request_data and isinstance(
request_data["messages"], list
):
input = request_data["messages"]
litellm_logging_obj.model_call_details["messages"] = input
if litellm_logging_obj.call_type != CallTypes.pass_through.value:
litellm_logging_obj.call_type = CallTypes.acompletion.value
normalized_call_type = CallTypes.acompletion.value
elif "prompt" in request_data and isinstance(request_data["prompt"], str):
input = request_data["prompt"]
litellm_logging_obj.model_call_details["prompt"] = input
if litellm_logging_obj.call_type != CallTypes.pass_through.value:
litellm_logging_obj.call_type = CallTypes.atext_completion.value
normalized_call_type = CallTypes.atext_completion.value
elif "input" in request_data and isinstance(request_data["input"], list):
input = request_data["input"]
litellm_logging_obj.model_call_details["input"] = input
if litellm_logging_obj.call_type != CallTypes.pass_through.value:
litellm_logging_obj.call_type = CallTypes.aembedding.value
normalized_call_type = CallTypes.aembedding.value
if normalized_call_type is not None:
litellm_logging_obj.call_type = normalized_call_type
litellm_logging_obj.model_call_details["call_type"] = (
normalized_call_type
)
# Pass-through endpoints are logged via the callback loop's
# async_post_call_failure_hook — skip pre_call and failure handlers.
if litellm_logging_obj.call_type == CallTypes.pass_through.value:
return
litellm_logging_obj.pre_call(
input=input,
api_key="",
@@ -2278,6 +2278,75 @@ async def test_post_call_failure_hook_auth_error_llm_api_route():
mock_handle_logging.assert_called_once()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"request_data, route, expected_call_type",
[
(
{"model": "bad-model", "messages": [{"role": "user", "content": "hello"}]},
"/v1/chat/completions",
"acompletion",
),
(
{"model": "bad-model", "prompt": "hello"},
"/v1/completions",
"atext_completion",
),
(
{"model": "bad-model", "input": ["hello"]},
"/v1/embeddings",
"aembedding",
),
],
)
async def test_handle_logging_proxy_only_error_syncs_normalized_call_type(
request_data, route, expected_call_type
):
from fastapi import HTTPException
from litellm.caching.caching import DualCache
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.utils import ProxyLogging
cache = DualCache()
proxy_logging = ProxyLogging(user_api_key_cache=cache)
captured_logging_obj = {}
original_function_setup = litellm.utils.function_setup
def _capture_function_setup(*args, **kwargs):
logging_obj, data = original_function_setup(*args, **kwargs)
captured_logging_obj["logging_obj"] = logging_obj
return logging_obj, data
with patch(
"litellm.proxy.utils.litellm.utils.function_setup",
side_effect=_capture_function_setup,
), patch.object(
Logging, "async_failure_handler", new=AsyncMock(return_value=None)
), patch.object(
Logging, "failure_handler", return_value=None
), patch(
"litellm.proxy.utils.threading.Thread"
) as mock_thread:
mock_thread.return_value.start = Mock()
await proxy_logging._handle_logging_proxy_only_error(
request_data=request_data,
user_api_key_dict=UserAPIKeyAuth(
api_key="test_key",
user_id="test_user",
token="test_token",
request_route=route,
),
route=route,
original_exception=HTTPException(status_code=400, detail="bad request"),
)
logging_obj = captured_logging_obj["logging_obj"]
assert logging_obj.call_type == expected_call_type
assert logging_obj.model_call_details["call_type"] == expected_call_type
@pytest.mark.asyncio
async def test_during_call_hook_parallel_execution():
"""