fix(policy_engine): propagate post_call pipeline replacement responses to the client

This commit is contained in:
mateo-berri 2026-08-28 17:40:48 -07:00
parent dfcea2c186
commit e6edd62f5d
4 changed files with 112 additions and 22 deletions

View file

@ -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):

View file

@ -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]] = []

View file

@ -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"

View file

@ -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)])