mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(policy_engine): propagate post_call pipeline replacement responses to the client
This commit is contained in:
parent
dfcea2c186
commit
e6edd62f5d
4 changed files with 112 additions and 22 deletions
|
|
@ -108,8 +108,10 @@ class PipelineExecutor:
|
|||
action,
|
||||
)
|
||||
|
||||
# Forward modified data to next step if pass_data is True
|
||||
if step.pass_data and modified_data is not None:
|
||||
# Forward modified data to the next step if pass_data is True;
|
||||
# post_call response replacements always chain, matching the flat
|
||||
# callback loop where each hook sees the previous hook's response
|
||||
if modified_data is not None and (step.pass_data or mode == "post_call"):
|
||||
working_data = {**working_data, **modified_data}
|
||||
|
||||
# Handle terminal actions
|
||||
|
|
@ -227,11 +229,14 @@ class PipelineExecutor:
|
|||
# same contract as run_in_parallel/scan_raw_request elsewhere: any
|
||||
# data it returned is discarded, since applying it on top of the
|
||||
# raw snapshot would silently undo whatever an earlier step in
|
||||
# this pipeline already did.
|
||||
modified_data = None
|
||||
if response is not None and isinstance(response, dict) and not scans_raw_request:
|
||||
modified_data = response
|
||||
return ("pass", modified_data, None, None)
|
||||
# this pipeline already did. A post_call hook's non-None return is
|
||||
# a replacement response (the flat callback-loop contract), carried
|
||||
# under the same "response" key the step input uses.
|
||||
if response is None or scans_raw_request:
|
||||
return ("pass", None, None, None)
|
||||
if mode == "post_call":
|
||||
return ("pass", {"response": response}, None, None)
|
||||
return ("pass", response if isinstance(response, dict) else None, None, None)
|
||||
|
||||
except Exception as e:
|
||||
if CustomGuardrail._is_guardrail_intervention(e):
|
||||
|
|
|
|||
|
|
@ -1603,7 +1603,7 @@ class ProxyLogging:
|
|||
event_hook: str,
|
||||
raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data
|
||||
response: LLMResponseTypes | None = None,
|
||||
) -> dict:
|
||||
) -> tuple[dict, LLMResponseTypes | None]:
|
||||
"""
|
||||
Execute guardrail pipelines if any are configured for this request.
|
||||
|
||||
|
|
@ -1615,18 +1615,21 @@ class ProxyLogging:
|
|||
``scan_raw_request`` evaluates the pristine request, not whatever an
|
||||
earlier ``pass_data`` step in the same pipeline already rewrote.
|
||||
|
||||
Returns the (possibly modified) data dict.
|
||||
Returns the (possibly modified) data dict, plus the replacement
|
||||
response when a post_call pipeline step returned one (None when the
|
||||
response is unchanged), matching the flat callback-loop contract.
|
||||
"""
|
||||
pipelines: Final = _policy_pipelines(data)
|
||||
if not pipelines:
|
||||
return data
|
||||
|
||||
step_input: Final = {**data, "response": response} if response is not None else data
|
||||
return data, None
|
||||
|
||||
current_response = response # rebind-ok: chains each pipeline's replacement response into the next
|
||||
for policy_name, pipeline in pipelines:
|
||||
if pipeline.mode != event_hook:
|
||||
continue
|
||||
|
||||
step_input: dict = {**data, "response": current_response} if current_response is not None else data
|
||||
|
||||
result: PipelineExecutionResult = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
|
|
@ -1641,10 +1644,13 @@ class ProxyLogging:
|
|||
result=result,
|
||||
data=data,
|
||||
policy_name=policy_name,
|
||||
original_response=response,
|
||||
original_response=current_response,
|
||||
)
|
||||
|
||||
return data
|
||||
if current_response is not None and result.modified_data is not None:
|
||||
current_response = result.modified_data.get("response", current_response)
|
||||
|
||||
return data, current_response if current_response is not response else None
|
||||
|
||||
@staticmethod
|
||||
def _handle_pipeline_result(
|
||||
|
|
@ -1657,9 +1663,9 @@ class ProxyLogging:
|
|||
Handle a PipelineExecutionResult — allow, block, or modify_response.
|
||||
|
||||
Returns data dict if allowed, raises on block/modify_response.
|
||||
``original_response`` is set on the post_call path, where allowed
|
||||
modifications land on the response object in place, so the request
|
||||
payload (already sent upstream) is left untouched.
|
||||
``original_response`` is set on the post_call path, where the request
|
||||
payload (already sent upstream) must stay untouched; a replacement
|
||||
response carried in ``modified_data`` is adopted by the caller.
|
||||
"""
|
||||
if result.terminal_action == "allow":
|
||||
if result.modified_data is not None and original_response is None:
|
||||
|
|
@ -1822,7 +1828,7 @@ class ProxyLogging:
|
|||
_raise_for_streaming_post_call_pipelines(data)
|
||||
|
||||
# Execute guardrail pipelines before the normal callback loop
|
||||
data = await self._maybe_execute_pipelines(
|
||||
data, _ = await self._maybe_execute_pipelines(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
|
|
@ -2809,13 +2815,15 @@ class ProxyLogging:
|
|||
from litellm.proxy.proxy_server import llm_router
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
await self._maybe_execute_pipelines(
|
||||
_, pipeline_response = await self._maybe_execute_pipelines(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=getattr(data.get("litellm_logging_obj"), "call_type", None) or "acompletion",
|
||||
event_hook="post_call",
|
||||
response=response,
|
||||
)
|
||||
if pipeline_response is not None:
|
||||
response = pipeline_response # rebind-ok: adopt the pipeline's replacement response, same contract as the callback loops below
|
||||
|
||||
guardrail_callbacks: Final[list[CustomGuardrail]] = []
|
||||
other_callbacks: Final[list[CustomLogger]] = []
|
||||
|
|
|
|||
|
|
@ -326,13 +326,14 @@ def test_process_guardrail_metadata_invalid_data_raises(proxy_logging):
|
|||
@pytest.mark.asyncio
|
||||
async def test_maybe_execute_pipelines_no_pipelines_returns_data(proxy_logging, make_user_api_key_auth):
|
||||
data = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1}
|
||||
out = await proxy_logging._maybe_execute_pipelines(
|
||||
out, replacement = await proxy_logging._maybe_execute_pipelines(
|
||||
data=data,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
call_type="completion",
|
||||
event_hook="pre_call",
|
||||
)
|
||||
assert out == {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1}
|
||||
assert replacement is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -344,7 +345,7 @@ async def test_maybe_execute_pipelines_skips_pipelines_with_other_mode(proxy_log
|
|||
monkeypatch.setattr(
|
||||
"litellm.proxy.policy_engine.pipeline_executor.PipelineExecutor.execute_steps", executed
|
||||
)
|
||||
out = await proxy_logging._maybe_execute_pipelines(
|
||||
out, replacement = await proxy_logging._maybe_execute_pipelines(
|
||||
data=data,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
call_type="completion",
|
||||
|
|
@ -352,6 +353,7 @@ async def test_maybe_execute_pipelines_skips_pipelines_with_other_mode(proxy_log
|
|||
)
|
||||
executed.assert_not_called()
|
||||
assert out is data
|
||||
assert replacement is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -949,6 +951,81 @@ async def test_post_call_pipeline_pass_runs_once_and_leaves_request_data_untouch
|
|||
assert "guardrails" not in data["metadata"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_pipeline_replacement_response_reaches_caller(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
masked = litellm.ModelResponse()
|
||||
|
||||
class MaskingGuardrail(CustomGuardrail):
|
||||
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
|
||||
return masked
|
||||
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"callbacks",
|
||||
[MaskingGuardrail(guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False)],
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
|
||||
data = _post_call_pipeline_data()
|
||||
|
||||
out = await proxy_logging.post_call_success_hook(
|
||||
data=data, response=litellm.ModelResponse(), user_api_key_dict=make_user_api_key_auth()
|
||||
)
|
||||
|
||||
assert out is masked
|
||||
assert "response" not in data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_pipeline_replacement_chains_to_next_step_without_pass_data(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
masked = litellm.ModelResponse()
|
||||
seen: Dict[str, Any] = {}
|
||||
|
||||
class MaskingGuardrail(CustomGuardrail):
|
||||
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
|
||||
return masked
|
||||
|
||||
class RecordingGuardrail(CustomGuardrail):
|
||||
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
|
||||
seen["response"] = response
|
||||
return None
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="post_call",
|
||||
steps=[
|
||||
PipelineStep(guardrail="gr-mask", on_pass="next", on_fail="block"),
|
||||
PipelineStep(guardrail="gr-audit", on_pass="allow", on_fail="block"),
|
||||
],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"callbacks",
|
||||
[
|
||||
MaskingGuardrail(guardrail_name="gr-mask", event_hook=GuardrailEventHooks.post_call, default_on=False),
|
||||
RecordingGuardrail(guardrail_name="gr-audit", event_hook=GuardrailEventHooks.post_call, default_on=False),
|
||||
],
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
|
||||
data = {
|
||||
"model": "m",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"metadata": {
|
||||
"_guardrail_pipelines": [("response-governance", pipeline)],
|
||||
"_pipeline_managed_guardrails": {"gr-mask", "gr-audit"},
|
||||
},
|
||||
}
|
||||
|
||||
out = await proxy_logging.post_call_success_hook(
|
||||
data=data, response=litellm.ModelResponse(), user_api_key_dict=make_user_api_key_auth()
|
||||
)
|
||||
|
||||
assert out is masked
|
||||
assert seen["response"] is masked
|
||||
|
||||
|
||||
def test_handle_pipeline_result_modify_response_carries_original_response():
|
||||
result = MagicMock()
|
||||
result.terminal_action = "modify_response"
|
||||
|
|
|
|||
|
|
@ -660,7 +660,7 @@ async def test_scan_raw_request_snapshot_taken_before_pipelines(
|
|||
for msg in data.get("messages", []):
|
||||
if "SECRET" in msg.get("content", ""):
|
||||
msg["content"] = msg["content"].replace("SECRET", "[REDACTED]")
|
||||
return data
|
||||
return data, None
|
||||
|
||||
monkeypatch.setattr(ProxyLogging, "_maybe_execute_pipelines", fake_pipelines)
|
||||
monkeypatch.setattr(litellm, "callbacks", [_BlockOnSecretGuardrail(scan_raw_request=True)])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue