mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
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:
parent
0c58346ba9
commit
94f9230d13
3 changed files with 155 additions and 109 deletions
|
|
@ -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])
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue