mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(guardrails): scope the logging_only response scan once
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
87263cefca
commit
af312dc8d7
4 changed files with 45 additions and 19 deletions
|
|
@ -960,12 +960,14 @@ class CustomGuardrail(CustomLogger):
|
||||||
|
|
||||||
def _chat_shaped_request(
|
def _chat_shaped_request(
|
||||||
self,
|
self,
|
||||||
scratch_request: dict, # mutable-ok: CustomLogger.async_logging_hook contract
|
scratch_request: Mapping[str, object],
|
||||||
translation: "BaseTranslation",
|
translation: "BaseTranslation",
|
||||||
) -> dict: # mutable-ok: BaseTranslation.process_output_response contract
|
) -> dict[str, object]: # mutable-ok: BaseTranslation.process_output_response contract
|
||||||
"""The logged request in OpenAI chat shape, for an output scan whose translation differs from the input's."""
|
"""The logged request in OpenAI chat shape, for an output scan whose translation differs from the input's."""
|
||||||
context: Final = translation.request_scan_context(scratch_request, self)
|
messages, tools = translation.chat_shaped_request_conversation(
|
||||||
return {**scratch_request, "messages": list(context.structured_messages), "tools": list(context.tools)}
|
dict(scratch_request) # mutable-ok: BaseTranslation.chat_shaped_request_conversation requires a dict
|
||||||
|
)
|
||||||
|
return {**scratch_request, "messages": list(messages), "tools": list(tools)}
|
||||||
|
|
||||||
def supports_scan_only_tool_results(self) -> bool:
|
def supports_scan_only_tool_results(self) -> bool:
|
||||||
"""Whether this guardrail can scan tool-result content.
|
"""Whether this guardrail can scan tool-result content.
|
||||||
|
|
|
||||||
|
|
@ -528,23 +528,26 @@ class AnthropicMessagesHandler(BaseTranslation):
|
||||||
)
|
)
|
||||||
return result if result else None
|
return result if result else None
|
||||||
|
|
||||||
def request_scan_context(self, data: dict, guardrail_to_apply: "CustomGuardrail") -> RequestScanContext:
|
def chat_shaped_request_conversation(
|
||||||
|
self, data: dict
|
||||||
|
) -> tuple[tuple[AllMessageValues, ...], tuple[ChatCompletionToolParam, ...]]:
|
||||||
if data.get("messages") is None:
|
if data.get("messages") is None:
|
||||||
return RequestScanContext()
|
return (), ()
|
||||||
translated: Final = self._translate_to_openai(
|
translated: Final = self._translate_to_openai(
|
||||||
{key: value for key, value in data.items() if key != "system"} # mutable-ok: API message payload
|
{key: value for key, value in data.items() if key != "system"} # mutable-ok: API message payload
|
||||||
)
|
)
|
||||||
hoisted_system_message: Final = (
|
hoisted_system_message: Final = self._hoisted_top_level_system_message(data)
|
||||||
None
|
messages: Final = (
|
||||||
if effective_skip_system_message_for_guardrail(guardrail_to_apply)
|
*(() if hoisted_system_message is None else (hoisted_system_message,)),
|
||||||
else self._hoisted_top_level_system_message(data)
|
*translated["messages"],
|
||||||
)
|
|
||||||
return RequestScanContext.scoped(
|
|
||||||
(*(() if hoisted_system_message is None else (hoisted_system_message,)), *translated["messages"]),
|
|
||||||
tuple(tool for tool in translated.get("tools") or () if not is_provider_native_tool_dict(tool)),
|
|
||||||
guardrail_to_apply,
|
|
||||||
skip_system=False,
|
|
||||||
)
|
)
|
||||||
|
tools: Final = tuple(tool for tool in translated.get("tools") or () if not is_provider_native_tool_dict(tool))
|
||||||
|
return messages, tools
|
||||||
|
|
||||||
|
def request_scan_context(self, data: dict, guardrail_to_apply: "CustomGuardrail") -> RequestScanContext:
|
||||||
|
if data.get("messages") is None:
|
||||||
|
return RequestScanContext()
|
||||||
|
return RequestScanContext.scoped(*self.chat_shaped_request_conversation(data), guardrail_to_apply)
|
||||||
|
|
||||||
async def process_input_messages(
|
async def process_input_messages(
|
||||||
self,
|
self,
|
||||||
|
|
|
||||||
|
|
@ -298,11 +298,16 @@ class BaseTranslation(ABC):
|
||||||
"""
|
"""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def chat_shaped_request_conversation(
|
||||||
|
self, data: dict
|
||||||
|
) -> tuple[tuple["AllMessageValues", ...], tuple["ChatCompletionToolParam", ...]]:
|
||||||
|
"""The full, unscoped request turns and tool definitions in OpenAI chat shape."""
|
||||||
|
return tuple(self.get_structured_messages(data) or ()), tuple(data.get("tools") or ())
|
||||||
|
|
||||||
def request_scan_context(self, data: dict, guardrail_to_apply: "CustomGuardrail") -> RequestScanContext:
|
def request_scan_context(self, data: dict, guardrail_to_apply: "CustomGuardrail") -> RequestScanContext:
|
||||||
"""Override wherever ``process_input_messages`` scopes or translates the request differently."""
|
"""Override wherever ``process_input_messages`` scopes or translates the request differently."""
|
||||||
return RequestScanContext.scoped(
|
messages, tools = self.chat_shaped_request_conversation(data)
|
||||||
self.get_structured_messages(data) or (), data.get("tools") or (), guardrail_to_apply
|
return RequestScanContext.scoped(messages, tools, guardrail_to_apply)
|
||||||
)
|
|
||||||
|
|
||||||
def with_response_context(
|
def with_response_context(
|
||||||
self,
|
self,
|
||||||
|
|
|
||||||
|
|
@ -2699,6 +2699,22 @@ class TestLoggingOnlyApplyGuardrail:
|
||||||
("response", [*expected_request, {"role": "assistant", "content": "general kenobi"}], expected_tools),
|
("response", [*expected_request, {"role": "assistant", "content": "general kenobi"}], expected_tools),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_anthropic_messages_response_scan_keeps_reply_when_scoping_empties_request(self):
|
||||||
|
class _ContextObserver(_ApplyOnlyObserver):
|
||||||
|
@log_guardrail_information
|
||||||
|
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
|
||||||
|
self.calls.append((input_type, inputs.get("structured_messages"), inputs.get("tools")))
|
||||||
|
return inputs
|
||||||
|
|
||||||
|
guardrail = _ContextObserver()
|
||||||
|
guardrail.scan_only_tool_results = True
|
||||||
|
kwargs, response = _logged_call([{"role": "user", "content": "What is the capital of France?"}])
|
||||||
|
|
||||||
|
await guardrail.async_logging_hook(kwargs, response, CallTypes.anthropic_messages.value)
|
||||||
|
|
||||||
|
assert guardrail.calls == [("response", [{"role": "assistant", "content": "general kenobi"}], None)]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_async_success_handler_records_verdict_in_standard_logging_object(self):
|
async def test_async_success_handler_records_verdict_in_standard_logging_object(self):
|
||||||
import datetime as dt
|
import datetime as dt
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue