fix(proxy): type the new pipeline tests and keep tag values out of the deferral warning

Every test this PR adds now annotates its fixture and parametrize
parameters. The submit-time warning for a tag-matched deferred policy
names only the policies, since a wildcard attachment pattern would let
caller-provided tag text reach the log.
This commit is contained in:
mateo-berri 2026-09-08 19:59:09 -07:00
parent 0c58346ba9
commit 94f9230d13
3 changed files with 155 additions and 109 deletions

View file

@ -620,19 +620,19 @@ def _defer_post_call_pipelines(
"retrieval re-matches only the key, team, and model scopes, so a tag carried in the request body "
"does not govern the completed response: %s",
response.id,
", ".join(f"{policy_name} ({source})" for policy_name, source in tag_matched),
", ".join(tag_matched),
)
_withdraw_deferred_claims(data, deferred)
def _tag_matched_deferrals(
data: Mapping[str, object], deferred: Sequence[tuple[str, "GuardrailPipeline"]]
) -> tuple[tuple[str, str], ...]:
) -> tuple[str, ...]:
sources: Final = _policy_state_metadata(data).get("policy_sources")
if not isinstance(sources, dict):
return ()
return tuple(
(policy_name, str(sources[policy_name]))
policy_name
for policy_name, _pipeline in deferred
if policy_name in sources and "tag:" in str(sources[policy_name])
)

View file

@ -1,5 +1,5 @@
import logging
from collections.abc import Mapping
from collections.abc import Iterator, Mapping
import pytest
@ -60,7 +60,7 @@ def _pipeline_policy(guardrail: str, mode: str = "post_call") -> dict[str, objec
@pytest.fixture
def policy_engine():
def policy_engine() -> Iterator[None]:
policy_registry = get_policy_registry()
attachment_registry = get_attachment_registry()
policy_registry.load_policies(
@ -97,7 +97,7 @@ def _attached_pipelines(data: Mapping[str, object]) -> tuple[tuple[str, str], ..
)
def test_attaches_model_scoped_post_call_pipeline_to_retrieval(policy_engine):
def test_attaches_model_scoped_post_call_pipeline_to_retrieval(policy_engine: None) -> None:
data = _retrieval_data(GOVERNED_MODEL_ID)
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=_router())
@ -111,7 +111,7 @@ def test_attaches_model_scoped_post_call_pipeline_to_retrieval(policy_engine):
assert "guardrails" not in data["litellm_metadata"]
def test_key_and_team_context_also_governs_retrieval(policy_engine):
def test_key_and_team_context_also_governs_retrieval(policy_engine: None) -> None:
data = _retrieval_data(UNGOVERNED_MODEL_ID)
attach_post_call_pipelines_to_retrieval(
@ -121,7 +121,7 @@ def test_key_and_team_context_also_governs_retrieval(policy_engine):
assert _attached_pipelines(data) == (("team-governance", "team-word-filter"),)
def test_tag_attached_policy_is_not_re_matched_when_the_retrieval_carries_no_tag(policy_engine):
def test_tag_attached_policy_is_not_re_matched_when_the_retrieval_carries_no_tag(policy_engine: None) -> None:
data = _retrieval_data(GOVERNED_MODEL_ID)
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=_router())
@ -129,7 +129,7 @@ def test_tag_attached_policy_is_not_re_matched_when_the_retrieval_carries_no_tag
assert _attached_pipelines(data) == (("response-governance", "output-word-filter"),)
def test_tag_attached_policy_governs_a_retrieval_whose_metadata_carries_the_tag(policy_engine):
def test_tag_attached_policy_governs_a_retrieval_whose_metadata_carries_the_tag(policy_engine: None) -> None:
data: dict[str, object] = {
"response_id": _encoded_response_id(UNGOVERNED_MODEL_ID),
"litellm_metadata": {"tags": ["governed"]},
@ -141,7 +141,7 @@ def test_tag_attached_policy_governs_a_retrieval_whose_metadata_carries_the_tag(
assert data["litellm_metadata"]["policy_sources"] == {"tag-governance": "tag:governed"}
def test_retrieval_of_an_ungoverned_model_attaches_nothing(policy_engine):
def test_retrieval_of_an_ungoverned_model_attaches_nothing(policy_engine: None) -> None:
data = _retrieval_data(UNGOVERNED_MODEL_ID)
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=_router())
@ -149,7 +149,7 @@ def test_retrieval_of_an_ungoverned_model_attaches_nothing(policy_engine):
assert data == _retrieval_data(UNGOVERNED_MODEL_ID)
def test_already_attached_policy_is_not_attached_twice(policy_engine):
def test_already_attached_policy_is_not_attached_twice(policy_engine: None) -> None:
data = _retrieval_data(GOVERNED_MODEL_ID)
router = _router()
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=router)
@ -168,7 +168,9 @@ def _hidden_submit_model_warnings(caplog: pytest.LogCaptureFixture) -> list[str]
]
def test_wildcard_deployment_attaches_nothing_for_the_submitted_model_and_warns(policy_engine, caplog):
def test_wildcard_deployment_attaches_nothing_for_the_submitted_model_and_warns(
policy_engine: None, caplog: pytest.LogCaptureFixture
) -> None:
data = _retrieval_data(WILDCARD_MODEL_ID)
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
@ -176,11 +178,14 @@ def test_wildcard_deployment_attaches_nothing_for_the_submitted_model_and_warns(
assert data == _retrieval_data(WILDCARD_MODEL_ID)
assert [
"as model group openai/* (a wildcard deployment)" in message for message in _hidden_submit_model_warnings(caplog)
"as model group openai/* (a wildcard deployment)" in message
for message in _hidden_submit_model_warnings(caplog)
] == [True]
def test_aliased_model_group_still_attaches_its_own_policies_and_warns(policy_engine, caplog):
def test_aliased_model_group_still_attaches_its_own_policies_and_warns(
policy_engine: None, caplog: pytest.LogCaptureFixture
) -> None:
data = _retrieval_data(GOVERNED_MODEL_ID)
router = _router({"gpt-mini": GOVERNED_MODEL_GROUP, "gpt-hidden": {"model": GOVERNED_MODEL_GROUP, "hidden": True}})
@ -194,7 +199,9 @@ def test_aliased_model_group_still_attaches_its_own_policies_and_warns(policy_en
] == [True]
def test_plain_model_group_retrieval_does_not_warn_about_the_submitted_model(policy_engine, caplog):
def test_plain_model_group_retrieval_does_not_warn_about_the_submitted_model(
policy_engine: None, caplog: pytest.LogCaptureFixture
) -> None:
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
attach_post_call_pipelines_to_retrieval(
data=_retrieval_data(GOVERNED_MODEL_ID),
@ -209,7 +216,8 @@ def _ungoverned_retrieval_warnings(caplog: pytest.LogCaptureFixture) -> list[str
return [
record.getMessage()
for record in caplog.records
if record.levelno == logging.WARNING and "retrieved without its post_call policy pipelines" in record.getMessage()
if record.levelno == logging.WARNING
and "retrieved without its post_call policy pipelines" in record.getMessage()
]
@ -221,7 +229,9 @@ def _ungoverned_retrieval_warnings(caplog: pytest.LogCaptureFixture) -> list[str
(None, "response id names no deployment"),
],
)
def test_unresolvable_response_id_attaches_nothing_and_warns(policy_engine, caplog, response_id, reason):
def test_unresolvable_response_id_attaches_nothing_and_warns(
policy_engine: None, caplog: pytest.LogCaptureFixture, response_id: str, reason: str
) -> None:
data = {"response_id": response_id, "litellm_metadata": {}}
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
@ -231,7 +241,7 @@ def test_unresolvable_response_id_attaches_nothing_and_warns(policy_engine, capl
assert [message.endswith(f"({reason})") for message in _ungoverned_retrieval_warnings(caplog)] == [True]
def test_without_a_router_attaches_nothing_and_warns(policy_engine, caplog):
def test_without_a_router_attaches_nothing_and_warns(policy_engine: None, caplog: pytest.LogCaptureFixture) -> None:
data = _retrieval_data(GOVERNED_MODEL_ID)
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
@ -241,7 +251,7 @@ def test_without_a_router_attaches_nothing_and_warns(policy_engine, caplog):
assert [message.endswith("(no router)") for message in _ungoverned_retrieval_warnings(caplog)] == [True]
def test_without_policy_engine_attaches_nothing_quietly(caplog):
def test_without_policy_engine_attaches_nothing_quietly(caplog: pytest.LogCaptureFixture) -> None:
get_policy_registry().clear()
data = _retrieval_data(GOVERNED_MODEL_ID)

View file

@ -12,6 +12,7 @@ from __future__ import annotations
import asyncio
import json
import logging
from collections.abc import Iterator
from typing import Any, Callable, Dict, List
from unittest.mock import AsyncMock, MagicMock, patch
@ -155,9 +156,7 @@ async def test_execute_guardrail_hook_unknown_hook_type_raises(proxy_logging, ma
@pytest.mark.asyncio
async def test_execute_guardrail_with_load_balancing_routes_through_router(
proxy_logging, make_user_api_key_auth
):
async def test_execute_guardrail_with_load_balancing_routes_through_router(proxy_logging, make_user_api_key_auth):
cb = _make_guardrail()
router = MagicMock()
router.get_available_guardrail = MagicMock(return_value={"callback": cb})
@ -173,9 +172,7 @@ async def test_execute_guardrail_with_load_balancing_routes_through_router(
@pytest.mark.asyncio
async def test_execute_guardrail_with_load_balancing_router_none_raises(
proxy_logging, make_user_api_key_auth
):
async def test_execute_guardrail_with_load_balancing_router_none_raises(proxy_logging, make_user_api_key_auth):
with patch("litellm.proxy.proxy_server.llm_router", None):
with pytest.raises(ValueError, match="Router not initialized"):
await proxy_logging._execute_guardrail_with_load_balancing(
@ -188,9 +185,7 @@ async def test_execute_guardrail_with_load_balancing_router_none_raises(
@pytest.mark.asyncio
async def test_execute_guardrail_with_load_balancing_no_callback_raises(
proxy_logging, make_user_api_key_auth
):
async def test_execute_guardrail_with_load_balancing_no_callback_raises(proxy_logging, make_user_api_key_auth):
router = MagicMock()
router.get_available_guardrail = MagicMock(return_value={"callback": None})
with patch("litellm.proxy.proxy_server.llm_router", router):
@ -210,9 +205,7 @@ async def test_execute_guardrail_with_load_balancing_no_callback_raises(
@pytest.mark.asyncio
async def test_process_guardrail_callback_skipped_when_should_run_false(
proxy_logging, make_user_api_key_auth
):
async def test_process_guardrail_callback_skipped_when_should_run_false(proxy_logging, make_user_api_key_auth):
cb = _make_guardrail()
cb.should_run_guardrail = MagicMock(return_value=False)
out = await proxy_logging._process_guardrail_callback(
@ -226,9 +219,7 @@ async def test_process_guardrail_callback_skipped_when_should_run_false(
@pytest.mark.asyncio
async def test_process_guardrail_callback_returns_data_on_success(
proxy_logging, make_user_api_key_auth, monkeypatch
):
async def test_process_guardrail_callback_returns_data_on_success(proxy_logging, make_user_api_key_auth, monkeypatch):
cb = _make_guardrail()
cb.should_run_guardrail = MagicMock(return_value=True)
proxy_logging._should_use_guardrail_load_balancing = MagicMock(return_value=False)
@ -343,14 +334,14 @@ async def test_maybe_execute_pipelines_no_pipelines_returns_data(proxy_logging,
@pytest.mark.asyncio
async def test_maybe_execute_pipelines_skips_pipelines_with_other_mode(proxy_logging, make_user_api_key_auth, monkeypatch):
async def test_maybe_execute_pipelines_skips_pipelines_with_other_mode(
proxy_logging, make_user_api_key_auth, monkeypatch
):
pipeline = MagicMock()
pipeline.mode = "post_call" # not pre_call
data = {"metadata": {"_guardrail_pipelines": [("p1", pipeline)]}, "model": "m", "messages": []}
executed = MagicMock()
monkeypatch.setattr(
"litellm.proxy.policy_engine.pipeline_executor.PipelineExecutor.execute_steps", executed
)
monkeypatch.setattr("litellm.proxy.policy_engine.pipeline_executor.PipelineExecutor.execute_steps", executed)
out, replacement = await proxy_logging._maybe_execute_pipelines(
data=data,
user_api_key_dict=make_user_api_key_auth(),
@ -538,9 +529,7 @@ def test_handle_pipeline_result_block_enriches_with_guardrail_name_and_mode():
litellm.callbacks = [cb]
try:
with pytest.raises(HTTPException) as info:
ProxyLogging._handle_pipeline_result(
result=result, data={"model": "m"}, policy_name="p"
)
ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p")
finally:
litellm.callbacks = saved
@ -652,9 +641,7 @@ async def test_run_guardrail_with_metrics_records_error_and_enriches(monkeypatch
monkeypatch.setattr(litellm, "callbacks", [prom])
with pytest.raises(HTTPException):
await ProxyLogging._run_guardrail_with_metrics(
callback=cb, coro=task(), hook_type="post_call"
)
await ProxyLogging._run_guardrail_with_metrics(callback=cb, coro=task(), hook_type="post_call")
assert detail["guardrail_name"] == "presidio"
recorded = prom._record_guardrail_metrics.call_args.kwargs
@ -682,9 +669,7 @@ def _moderation_guardrail() -> MagicMock:
@pytest.mark.asyncio
async def test_during_call_hook_records_latency_metric(
proxy_logging, make_user_api_key_auth, monkeypatch
):
async def test_during_call_hook_records_latency_metric(proxy_logging, make_user_api_key_auth, monkeypatch):
cb = _moderation_guardrail()
prom = _prometheus_callback()
monkeypatch.setattr(litellm, "callbacks", [prom, cb])
@ -703,9 +688,7 @@ async def test_during_call_hook_records_latency_metric(
@pytest.mark.asyncio
async def test_post_call_success_hook_records_latency_metric(
proxy_logging, make_user_api_key_auth, monkeypatch
):
async def test_post_call_success_hook_records_latency_metric(proxy_logging, make_user_api_key_auth, monkeypatch):
cb = _moderation_guardrail()
prom = _prometheus_callback()
monkeypatch.setattr(litellm, "callbacks", [prom, cb])
@ -733,9 +716,7 @@ async def test_post_call_success_hook_records_latency_metric(
async def test_process_prompt_template_no_op_when_no_prompt_spec(proxy_logging, monkeypatch):
from litellm.proxy.prompts import prompt_registry
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: None
)
monkeypatch.setattr(prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: None)
data: Dict[str, Any] = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1}
await proxy_logging._process_prompt_template(
data=data,
@ -760,9 +741,7 @@ async def test_process_prompt_template_applies_when_spec_resolves(proxy_logging,
"get_prompt_callback_for_prompt",
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec
)
monkeypatch.setattr(prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec)
logging_obj = MagicMock()
logging_obj.async_get_chat_completion_prompt = AsyncMock(
@ -810,9 +789,7 @@ async def test_process_prompt_template_async_get_prompt_error_raises(proxy_loggi
"get_prompt_callback_for_prompt",
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec
)
monkeypatch.setattr(prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec)
logging_obj = MagicMock()
logging_obj.async_get_chat_completion_prompt = AsyncMock(side_effect=RuntimeError("bad prompt"))
with pytest.raises(RuntimeError):
@ -913,9 +890,7 @@ async def test_process_prompt_template_aresponses_swaps_model_and_merges_input(p
"get_prompt_callback_for_prompt",
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec
)
monkeypatch.setattr(prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec)
logging_obj = MagicMock()
logging_obj.async_get_chat_completion_prompt = AsyncMock(
@ -1124,9 +1099,7 @@ async def test_pre_call_hook_still_runs_guardrail_managed_only_by_post_call_pipe
},
}
await proxy_logging.pre_call_hook(
user_api_key_dict=make_user_api_key_auth(), data=data, call_type="completion"
)
await proxy_logging.pre_call_hook(user_api_key_dict=make_user_api_key_auth(), data=data, call_type="completion")
assert seen["count"] == 1
@ -1289,11 +1262,7 @@ async def test_post_call_pipeline_block_keeps_guardrail_metadata_writes(
monkeypatch.setattr(
litellm,
"callbacks",
[
BlockingWriterGuardrail(
guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False
)
],
[BlockingWriterGuardrail(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()
@ -1379,9 +1348,7 @@ async def test_pre_call_pipeline_managed_parallel_guardrail_runs_exactly_once(
},
}
await proxy_logging.pre_call_hook(
user_api_key_dict=make_user_api_key_auth(), data=data, call_type="completion"
)
await proxy_logging.pre_call_hook(user_api_key_dict=make_user_api_key_auth(), data=data, call_type="completion")
assert seen["count"] == 1
@ -1433,14 +1400,20 @@ def _output_blocking_callbacks(seen: dict[str, object]) -> list[CustomGuardrail]
seen["response"] = response
raise HTTPException(status_code=400, detail={"error": "output blocked"})
return [OutputBlockingGuardrail(guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False)]
return [
OutputBlockingGuardrail(guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False)
]
@pytest.mark.asyncio
@pytest.mark.parametrize("pending_status", ["queued", "in_progress"])
async def test_post_call_success_hook_waits_for_pending_background_response_before_running_pipeline(
proxy_logging, make_user_api_key_auth, monkeypatch, caplog, pending_status
):
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
pending_status: str,
) -> None:
seen: dict[str, object] = {}
monkeypatch.setattr(litellm, "callbacks", _output_blocking_callbacks(seen))
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
@ -1464,8 +1437,11 @@ async def test_post_call_success_hook_waits_for_pending_background_response_befo
@pytest.mark.asyncio
@pytest.mark.parametrize("final_status", ["completed", "incomplete"])
async def test_post_call_success_hook_runs_pipeline_on_retrieved_background_response(
proxy_logging, make_user_api_key_auth, monkeypatch, final_status
):
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
final_status: str,
) -> None:
seen: dict[str, object] = {}
monkeypatch.setattr(litellm, "callbacks", _output_blocking_callbacks(seen))
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
@ -1486,7 +1462,9 @@ def _output_passing_callbacks() -> list[CustomGuardrail]:
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
return response
return [OutputPassingGuardrail(guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False)]
return [
OutputPassingGuardrail(guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False)
]
def _claimed_post_call_pipeline_data(
@ -1519,7 +1497,7 @@ def _claimed_post_call_pipeline_data(
@pytest.fixture
def clear_policy_registry():
def clear_policy_registry() -> Iterator[None]:
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
yield
@ -1528,8 +1506,11 @@ def clear_policy_registry():
@pytest.mark.asyncio
async def test_pending_background_response_withdraws_the_deferred_policy_claims(
proxy_logging, make_user_api_key_auth, monkeypatch, clear_policy_registry
):
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
clear_policy_registry: None,
) -> None:
monkeypatch.setattr(litellm, "callbacks", _output_blocking_callbacks({}))
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _claimed_post_call_pipeline_data("response-governance")
@ -1546,8 +1527,12 @@ async def test_pending_background_response_withdraws_the_deferred_policy_claims(
@pytest.mark.asyncio
async def test_pending_background_response_warns_when_the_deferred_policy_was_matched_through_a_tag(
proxy_logging, make_user_api_key_auth, monkeypatch, clear_policy_registry, caplog
):
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
clear_policy_registry: None,
caplog: pytest.LogCaptureFixture,
) -> None:
monkeypatch.setattr(litellm, "callbacks", _output_blocking_callbacks({}))
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _claimed_post_call_pipeline_data("response-governance", policy_source="tag:governed+model:m")
@ -1559,19 +1544,21 @@ async def test_pending_background_response_warns_when_the_deferred_policy_was_ma
assert out.status == "queued"
assert "policy_sources" not in data["metadata"]
assert [
message for message in _warnings(caplog) if "response-governance (tag:governed+model:m)" in message
] == [
assert [message for message in _warnings(caplog) if "through a request tag" in message] == [
"Policy engine: background response resp_bg matched post_call policies through a request tag at submit; "
"retrieval re-matches only the key, team, and model scopes, so a tag carried in the request body "
"does not govern the completed response: response-governance (tag:governed+model:m)"
"does not govern the completed response: response-governance"
]
@pytest.mark.asyncio
async def test_pending_background_response_matched_through_its_model_does_not_warn(
proxy_logging, make_user_api_key_auth, monkeypatch, clear_policy_registry, caplog
):
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
clear_policy_registry: None,
caplog: pytest.LogCaptureFixture,
) -> None:
monkeypatch.setattr(litellm, "callbacks", _output_blocking_callbacks({}))
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
@ -1587,12 +1574,17 @@ async def test_pending_background_response_matched_through_its_model_does_not_wa
@pytest.mark.asyncio
async def test_pending_background_response_keeps_the_claim_of_a_policy_that_runs_outside_its_pipeline(
proxy_logging, make_user_api_key_auth, monkeypatch, clear_policy_registry
):
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
clear_policy_registry: None,
) -> None:
monkeypatch.setattr(litellm, "callbacks", _output_blocking_callbacks({}))
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _claimed_post_call_pipeline_data(
"input-and-output-governance", "response-governance", extra_guardrails={"input-and-output-governance": ["gr-pre"]}
"input-and-output-governance",
"response-governance",
extra_guardrails={"input-and-output-governance": ["gr-pre"]},
)
await proxy_logging.post_call_success_hook(
@ -1606,8 +1598,11 @@ async def test_pending_background_response_keeps_the_claim_of_a_policy_that_runs
@pytest.mark.asyncio
async def test_pending_background_response_keeps_the_claim_of_a_default_on_guardrail_that_ran_pre_call(
proxy_logging, make_user_api_key_auth, monkeypatch, clear_policy_registry
):
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
clear_policy_registry: None,
) -> None:
class DualStageGuardrail(CustomGuardrail):
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
return response
@ -1630,8 +1625,11 @@ async def test_pending_background_response_keeps_the_claim_of_a_default_on_guard
@pytest.mark.asyncio
async def test_retrieved_background_response_keeps_the_policy_claim_once_its_pipeline_ran(
proxy_logging, make_user_api_key_auth, monkeypatch, clear_policy_registry
):
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
clear_policy_registry: None,
) -> None:
monkeypatch.setattr(litellm, "callbacks", _output_passing_callbacks())
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _claimed_post_call_pipeline_data("response-governance")
@ -1647,8 +1645,11 @@ async def test_retrieved_background_response_keeps_the_policy_claim_once_its_pip
@pytest.mark.asyncio
async def test_pre_call_hook_stays_quiet_on_background_request_with_post_call_pipeline(
proxy_logging, make_user_api_key_auth, monkeypatch, caplog
):
proxy_logging: ProxyLogging,
make_user_api_key_auth: Callable[..., UserAPIKeyAuth],
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
monkeypatch.setattr(litellm, "callbacks", [])
data = _post_call_pipeline_data(background=True)
@ -1709,7 +1710,9 @@ def test_streamable_post_call_pipelines_keeps_supported_and_drops_unsupported(
steps=[PipelineStep(guardrail="gr-post", on_fail="next"), PipelineStep(guardrail="gr-native", on_fail="block")],
)
pre_call = GuardrailPipeline(mode="pre_call", steps=[PipelineStep(guardrail="gr-native", on_fail="block")])
data = {"metadata": {"_guardrail_pipelines": [("governed", governed), ("ungoverned", ungoverned), ("req", pre_call)]}}
data = {
"metadata": {"_guardrail_pipelines": [("governed", governed), ("ungoverned", ungoverned), ("req", pre_call)]}
}
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
streamable = _streamable_post_call_pipelines(data, make_user_api_key_auth(request_route="/v1/chat/completions"))
@ -2009,7 +2012,9 @@ def _rewriting_stream_guardrail(transform: Callable[[Dict[str, Any]], Dict[str,
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
return {**inputs, **transform(inputs)}
return RewritingStreamGuardrail(guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False)
return RewritingStreamGuardrail(
guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=False
)
def _tool_call_stream_chunks() -> List[Any]:
@ -2188,11 +2193,38 @@ async def test_streaming_iterator_hook_pipeline_releases_originals_on_unresolvab
def _anthropic_sse_chunks() -> List[bytes]:
events = [
("message_start", {"type": "message_start", "message": {"id": "msg_1", "type": "message", "role": "assistant", "model": "m", "content": [], "stop_reason": None, "usage": {"input_tokens": 1, "output_tokens": 0}}}),
("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}),
("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hello world"}}),
(
"message_start",
{
"type": "message_start",
"message": {
"id": "msg_1",
"type": "message",
"role": "assistant",
"model": "m",
"content": [],
"stop_reason": None,
"usage": {"input_tokens": 1, "output_tokens": 0},
},
},
),
(
"content_block_start",
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
),
(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hello world"}},
),
("content_block_stop", {"type": "content_block_stop", "index": 0}),
("message_delta", {"type": "message_delta", "delta": {"stop_reason": "end_turn", "stop_sequence": None}, "usage": {"output_tokens": 2}}),
(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 2},
},
),
("message_stop", {"type": "message_stop"}),
]
return [f"event: {name}\ndata: {json.dumps(payload)}\n\n".encode() for name, payload in events]
@ -2259,7 +2291,13 @@ async def test_streaming_iterator_hook_pipeline_delivers_text_rewrite_on_anthrop
assert "hello [MASKED]" in raw
assert "hello world" not in raw
assert raw.count("event: content_block_delta") == 1
for expected_event in ("message_start", "content_block_start", "content_block_stop", "message_delta", "message_stop"):
for expected_event in (
"message_start",
"content_block_start",
"content_block_stop",
"message_delta",
"message_stop",
):
assert f"event: {expected_event}" in raw
@ -2332,9 +2370,7 @@ async def test_per_chunk_streaming_hook_skips_pipeline_managed_guardrail(
managed = UnifiedRecordingGuardrail(
guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=True
)
free = RecordingGuardrail(
guardrail_name="gr-free", event_hook=GuardrailEventHooks.post_call, default_on=True
)
free = RecordingGuardrail(guardrail_name="gr-free", event_hook=GuardrailEventHooks.post_call, default_on=True)
monkeypatch.setattr(litellm, "callbacks", [managed, free])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _post_call_pipeline_data(stream=True)