mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): make async_pre_call_hook str returns reach the caller as completion content
process_pre_call_hook_response only routed "completion"/"text_completion" onto the RejectedRequestError path, but /v1/chat/completions and /v1/completions run with call_type "acompletion"/"atext_completion", so a documented str return 400'd instead. RejectedRequestError now keeps the unprefixed text in raw_message and the proxy uses it for completion content, so callers no longer see "litellm.BadRequestError: litellm.RejectedRequestError: ..." in the message body. The streaming branch also crashed with AttributeError because post_call_failure_hook pops litellm_logging_obj before CustomStreamWrapper is built; capture it first. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
23de7a15d9
commit
544524ab8e
5 changed files with 40 additions and 5 deletions
|
|
@ -553,6 +553,7 @@ class RejectedRequestError(BadRequestError): # type: ignore
|
|||
litellm_debug_info: Optional[str] = None,
|
||||
):
|
||||
self.status_code = 400
|
||||
self.raw_message = message
|
||||
self.message = "litellm.RejectedRequestError: {}".format(message)
|
||||
self.model = model
|
||||
self.llm_provider = llm_provider
|
||||
|
|
|
|||
|
|
@ -9172,13 +9172,14 @@ async def chat_completion(
|
|||
return _chat_response
|
||||
except RejectedRequestError as e:
|
||||
_data = e.request_data
|
||||
_logging_obj = _data.get("litellm_logging_obj")
|
||||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=e,
|
||||
request_data=_data,
|
||||
)
|
||||
_chat_response = litellm.ModelResponse()
|
||||
_chat_response.choices[0].message.content = e.message # type: ignore
|
||||
_chat_response.choices[0].message.content = e.raw_message # type: ignore
|
||||
|
||||
if data.get("stream", None) is not None and data["stream"] is True:
|
||||
_iterator = litellm.utils.ModelResponseIterator(model_response=_chat_response, convert_to_delta=True)
|
||||
|
|
@ -9186,7 +9187,7 @@ async def chat_completion(
|
|||
completion_stream=_iterator,
|
||||
model=data.get("model", ""),
|
||||
custom_llm_provider="cached_response",
|
||||
logging_obj=_data.get("litellm_logging_obj", None),
|
||||
logging_obj=_logging_obj,
|
||||
)
|
||||
selected_data_generator = select_data_generator(
|
||||
response=_streaming_response,
|
||||
|
|
@ -9340,7 +9341,7 @@ async def completion(
|
|||
total_tokens=0,
|
||||
)
|
||||
_chat_response.usage = _usage # type: ignore
|
||||
_chat_response.choices[0].message.content = e.message # type: ignore
|
||||
_chat_response.choices[0].message.content = e.raw_message # type: ignore
|
||||
_iterator = litellm.utils.ModelResponseIterator(model_response=_chat_response, convert_to_delta=True)
|
||||
_streaming_response = litellm.TextCompletionStreamWrapper(
|
||||
completion_stream=_iterator,
|
||||
|
|
@ -9362,7 +9363,7 @@ async def completion(
|
|||
)
|
||||
else:
|
||||
_response = litellm.TextCompletionResponse()
|
||||
_response.choices[0].text = e.message
|
||||
_response.choices[0].text = e.raw_message
|
||||
return _response
|
||||
except Exception as e:
|
||||
await proxy_logging_obj.post_call_failure_hook(
|
||||
|
|
|
|||
|
|
@ -189,6 +189,16 @@ unified_guardrail = UnifiedLLMGuardrails()
|
|||
|
||||
NON_OPENAI_STREAM_GUARDRAIL_TRANSLATION_CALL_TYPES: "frozenset[CallTypes]" = frozenset({CallTypes.anthropic_messages})
|
||||
|
||||
REJECTED_REQUEST_SUPPORTED_CALL_TYPES: "frozenset[str]" = frozenset(
|
||||
call_type.value
|
||||
for call_type in (
|
||||
CallTypes.completion,
|
||||
CallTypes.acompletion,
|
||||
CallTypes.text_completion,
|
||||
CallTypes.atext_completion,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def print_verbose(print_statement):
|
||||
"""
|
||||
|
|
@ -920,7 +930,7 @@ class ProxyLogging:
|
|||
if isinstance(response, dict):
|
||||
return response
|
||||
if isinstance(response, str):
|
||||
if call_type in ["completion", "text_completion"]:
|
||||
if call_type in REJECTED_REQUEST_SUPPORTED_CALL_TYPES:
|
||||
raise RejectedRequestError(
|
||||
message=response,
|
||||
model=data.get("model", ""),
|
||||
|
|
|
|||
|
|
@ -46,6 +46,18 @@ async def test_process_pre_call_hook_response_string_completion_raises_rejected(
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("call_type", ["completion", "acompletion", "text_completion", "atext_completion"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_pre_call_hook_response_string_completion_call_types_raise_rejected(proxy_logging, call_type):
|
||||
with pytest.raises(RejectedRequestError) as info:
|
||||
await proxy_logging.process_pre_call_hook_response(
|
||||
response="rejected",
|
||||
data={"model": "m"},
|
||||
call_type=call_type,
|
||||
)
|
||||
assert info.value.raw_message == "rejected"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_pre_call_hook_response_string_other_call_type_raises_http(proxy_logging):
|
||||
with pytest.raises(HTTPException) as info:
|
||||
|
|
|
|||
11
tests/test_litellm/test_rejected_request_error_message.py
Normal file
11
tests/test_litellm/test_rejected_request_error_message.py
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
"""Regression: ``RejectedRequestError`` keeps the unprefixed text for completion content."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from litellm.exceptions import RejectedRequestError
|
||||
|
||||
|
||||
def test_rejected_request_error_keeps_raw_message():
|
||||
err = RejectedRequestError(message="366", model="gpt-4o", llm_provider="", request_data={"model": "gpt-4o"})
|
||||
assert err.raw_message == "366"
|
||||
assert "RejectedRequestError" in err.message
|
||||
Loading…
Add table
Reference in a new issue