mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(grayswan): resolve post-call surface from request route before call_type
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
cd0e9cee7c
commit
e5fac49e7b
2 changed files with 36 additions and 2 deletions
|
|
@ -615,10 +615,20 @@ class GraySwanGuardrail(CustomGuardrail):
|
|||
Returns the scoped context messages plus the request's tool definitions,
|
||||
or ``((), None)`` when the request surface cannot be resolved.
|
||||
"""
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
|
||||
call_type: Final = (logging_obj.call_type if logging_obj is not None else None) or getattr(
|
||||
request_data.get("litellm_logging_obj"), "call_type", None
|
||||
litellm_metadata: Final = request_data.get("litellm_metadata")
|
||||
request_route: Final = (
|
||||
litellm_metadata.get("user_api_key_request_route") if isinstance(litellm_metadata, Mapping) else None
|
||||
)
|
||||
route_call_types: Final = (
|
||||
get_call_types_for_route(request_route) if isinstance(request_route, str) else None
|
||||
)
|
||||
call_type: Final = (
|
||||
(route_call_types[0].value if route_call_types else None)
|
||||
or (logging_obj.call_type if logging_obj is not None else None)
|
||||
or getattr(request_data.get("litellm_logging_obj"), "call_type", None)
|
||||
)
|
||||
if not isinstance(call_type, str):
|
||||
return (), None
|
||||
|
|
|
|||
|
|
@ -747,6 +747,30 @@ async def test_post_call_scan_only_tool_results_scopes_context_and_tools() -> No
|
|||
assert "tools" not in payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_prefers_request_route_over_logging_call_type() -> None:
|
||||
guardrail = _post_call_guardrail()
|
||||
client = _CapturingClient()
|
||||
guardrail.async_handler = client
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["response text"]},
|
||||
request_data={
|
||||
**_REQUEST_DATA,
|
||||
"litellm_metadata": {"user_api_key_request_route": "/v1/chat/completions"},
|
||||
},
|
||||
input_type="response",
|
||||
logging_obj=_LoggingObj("responses"),
|
||||
)
|
||||
|
||||
payload = client.calls[0]["json"]
|
||||
assert list(payload["messages"]) == [
|
||||
*_REQUEST_DATA["messages"],
|
||||
{"role": "assistant", "content": "response text"},
|
||||
]
|
||||
assert list(payload["tools"]) == _REQUEST_DATA["tools"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_unresolvable_call_type_sends_response_only() -> None:
|
||||
guardrail = _post_call_guardrail()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue