mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix: withdraw policy header claims while a background response is pending and log ungoverned retrievals
This commit is contained in:
parent
5386ba37ff
commit
91091fd93e
4 changed files with 205 additions and 21 deletions
|
|
@ -1,6 +1,7 @@
|
|||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias
|
||||
from typing import TYPE_CHECKING, Final, Literal, TypeAlias
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
|
|
@ -29,14 +30,23 @@ PolicyPipelines: TypeAlias = tuple[tuple[str, GuardrailPipeline], ...]
|
|||
_POLICY_PIPELINES_ADAPTER: Final = TypeAdapter(PolicyPipelines)
|
||||
|
||||
|
||||
def _model_group_for_response_id(response_id: object, llm_router: "Router | None") -> str | None:
|
||||
if llm_router is None or not isinstance(response_id, str):
|
||||
return None
|
||||
model_id: Final = ResponsesAPIRequestUtils.get_model_id_from_response_id(response_id)
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class UngovernedRetrieval:
|
||||
reason: Literal["no router", "response id names no deployment", "deployment no longer in the router"]
|
||||
|
||||
|
||||
def _model_group_for_response_id(response_id: object, llm_router: "Router | None") -> str | UngovernedRetrieval:
|
||||
if llm_router is None:
|
||||
return UngovernedRetrieval("no router")
|
||||
model_id: Final = (
|
||||
ResponsesAPIRequestUtils.get_model_id_from_response_id(response_id) if isinstance(response_id, str) else None
|
||||
)
|
||||
if model_id is None:
|
||||
return None
|
||||
return UngovernedRetrieval("response id names no deployment")
|
||||
deployment: Final = llm_router.get_deployment(model_id)
|
||||
return deployment.model_name if deployment is not None else None
|
||||
if deployment is None:
|
||||
return UngovernedRetrieval("deployment no longer in the router")
|
||||
return deployment.model_name
|
||||
|
||||
|
||||
def _retrieval_context(
|
||||
|
|
@ -78,7 +88,12 @@ def attach_post_call_pipelines_to_retrieval(
|
|||
if not get_policy_registry().is_initialized():
|
||||
return
|
||||
model_group: Final = _model_group_for_response_id(data.get("response_id"), llm_router)
|
||||
if model_group is None:
|
||||
if isinstance(model_group, UngovernedRetrieval):
|
||||
verbose_proxy_logger.warning(
|
||||
"Policy engine: background response %s is retrieved without its post_call policy pipelines (%s)",
|
||||
data.get("response_id"),
|
||||
model_group.reason,
|
||||
)
|
||||
return
|
||||
context: Final = _retrieval_context(data, user_api_key_dict, model_group)
|
||||
post_call_pipelines, policy_sources = _post_call_pipelines_for_context(context)
|
||||
|
|
|
|||
|
|
@ -93,6 +93,7 @@ from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
|||
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
coerce_token_limit,
|
||||
get_or_create_metadata_bucket,
|
||||
independent_snapshot,
|
||||
is_expected_client_error,
|
||||
)
|
||||
|
|
@ -156,6 +157,8 @@ from litellm.proxy.hooks.sensitive_data_routing import (
|
|||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_guardrails_from_auth_metadata
|
||||
from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at
|
||||
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
from litellm.repositories.table_repositories import (
|
||||
|
|
@ -533,16 +536,65 @@ def _is_pending_background_response(response: LLMResponseTypes) -> bool:
|
|||
return isinstance(response, ResponsesAPIResponse) and response.status in _PENDING_BACKGROUND_RESPONSE_STATUSES
|
||||
|
||||
|
||||
def _log_deferred_post_call_pipelines(data: Mapping[str, object], response: ResponsesAPIResponse) -> None:
|
||||
policy_names: Final = tuple(policy_name for policy_name, _pipeline in _post_call_pipelines(data))
|
||||
if not policy_names:
|
||||
def _guardrails_outside_pipeline(policy_name: str, pipeline: "GuardrailPipeline") -> frozenset[str]:
|
||||
resolved: Final = PolicyResolver.resolve_policy_guardrails(
|
||||
policy_name=policy_name, policies=get_policy_registry().get_all_policies()
|
||||
)
|
||||
return frozenset(resolved.guardrails) - frozenset(step.guardrail for step in pipeline.steps)
|
||||
|
||||
|
||||
def _without_names(
|
||||
bucket: dict, # mutable-ok: the applied_* header slots live in the request-state dict every hook writes in place
|
||||
slot: str,
|
||||
names: frozenset[str],
|
||||
) -> None:
|
||||
claimed: Final = bucket.get(slot)
|
||||
if not isinstance(claimed, list):
|
||||
return
|
||||
remaining: Final = [name for name in claimed if name not in names]
|
||||
if remaining:
|
||||
bucket[slot] = remaining
|
||||
else:
|
||||
bucket.pop(slot)
|
||||
|
||||
|
||||
def _withdraw_deferred_claims(
|
||||
data: dict, # mutable-ok: same request-payload shape as post_call_success_hook's data
|
||||
deferred: Sequence[tuple[str, "GuardrailPipeline"]],
|
||||
) -> None:
|
||||
outside_by_policy: Final = MappingProxyType(
|
||||
{policy_name: _guardrails_outside_pipeline(policy_name, pipeline) for policy_name, pipeline in deferred}
|
||||
)
|
||||
running_elsewhere: Final = _pipeline_managed_guardrail_names(data, "pre_call").union(*outside_by_policy.values())
|
||||
withdrawn_policies: Final = frozenset(name for name, outside in outside_by_policy.items() if not outside)
|
||||
withdrawn_guardrails: Final = _pipeline_step_guardrail_names(deferred) - running_elsewhere
|
||||
_, bucket = get_or_create_metadata_bucket(data)
|
||||
_without_names(bucket, "applied_policies", withdrawn_policies)
|
||||
_without_names(bucket, "applied_guardrails", withdrawn_guardrails)
|
||||
sources: Final = bucket.get("policy_sources")
|
||||
if not isinstance(sources, dict):
|
||||
return
|
||||
remaining_sources: Final = {name: reason for name, reason in sources.items() if name not in withdrawn_policies}
|
||||
if remaining_sources:
|
||||
bucket["policy_sources"] = remaining_sources
|
||||
else:
|
||||
bucket.pop("policy_sources")
|
||||
|
||||
|
||||
def _defer_post_call_pipelines(
|
||||
data: dict, # mutable-ok: same request-payload shape as post_call_success_hook's data
|
||||
response: ResponsesAPIResponse,
|
||||
) -> None:
|
||||
deferred: Final = _post_call_pipelines(data)
|
||||
if not deferred:
|
||||
return
|
||||
verbose_proxy_logger.debug(
|
||||
"Post_call guardrail pipelines wait for background response %s (status=%s) to be retrieved complete: %s",
|
||||
response.id,
|
||||
response.status,
|
||||
", ".join(policy_names),
|
||||
", ".join(policy_name for policy_name, _pipeline in deferred),
|
||||
)
|
||||
_withdraw_deferred_claims(data, deferred)
|
||||
|
||||
|
||||
def _pipeline_is_streamable(policy_name: str, pipeline: "GuardrailPipeline") -> bool:
|
||||
|
|
@ -2939,7 +2991,7 @@ class ProxyLogging:
|
|||
response: LLMResponseTypes,
|
||||
) -> LLMResponseTypes | None:
|
||||
if _is_pending_background_response(response):
|
||||
_log_deferred_post_call_pipelines(data, response)
|
||||
_defer_post_call_pipelines(data, response)
|
||||
return None
|
||||
_, pipeline_response = await self._maybe_execute_pipelines(
|
||||
data=data,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
import logging
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -128,30 +130,48 @@ def test_already_attached_policy_is_not_attached_twice(policy_engine):
|
|||
assert data["litellm_metadata"]["applied_policies"] == ["response-governance"]
|
||||
|
||||
|
||||
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()
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"response_id",
|
||||
["resp_plain_upstream_id", _encoded_response_id("deployment-missing-from-router"), None],
|
||||
("response_id", "reason"),
|
||||
[
|
||||
("resp_plain_upstream_id", "response id names no deployment"),
|
||||
(_encoded_response_id("deployment-missing-from-router"), "deployment no longer in the router"),
|
||||
(None, "response id names no deployment"),
|
||||
],
|
||||
)
|
||||
def test_unresolvable_response_id_attaches_nothing(policy_engine, response_id):
|
||||
def test_unresolvable_response_id_attaches_nothing_and_warns(policy_engine, caplog, response_id, reason):
|
||||
data = {"response_id": response_id, "litellm_metadata": {}}
|
||||
|
||||
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=_router())
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=_router())
|
||||
|
||||
assert data == {"response_id": response_id, "litellm_metadata": {}}
|
||||
assert [message.endswith(f"({reason})") for message in _ungoverned_retrieval_warnings(caplog)] == [True]
|
||||
|
||||
|
||||
def test_without_a_router_attaches_nothing(policy_engine):
|
||||
def test_without_a_router_attaches_nothing_and_warns(policy_engine, caplog):
|
||||
data = _retrieval_data(GOVERNED_MODEL_ID)
|
||||
|
||||
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=None)
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=None)
|
||||
|
||||
assert data == _retrieval_data(GOVERNED_MODEL_ID)
|
||||
assert [message.endswith("(no router)") for message in _ungoverned_retrieval_warnings(caplog)] == [True]
|
||||
|
||||
|
||||
def test_without_policy_engine_attaches_nothing():
|
||||
def test_without_policy_engine_attaches_nothing_quietly(caplog):
|
||||
get_policy_registry().clear()
|
||||
data = _retrieval_data(GOVERNED_MODEL_ID)
|
||||
|
||||
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=_router())
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
attach_post_call_pipelines_to_retrieval(data=data, user_api_key_dict=UserAPIKeyAuth(), llm_router=None)
|
||||
|
||||
assert data == _retrieval_data(GOVERNED_MODEL_ID)
|
||||
assert _ungoverned_retrieval_warnings(caplog) == []
|
||||
|
|
|
|||
|
|
@ -1481,6 +1481,103 @@ async def test_post_call_success_hook_runs_pipeline_on_retrieved_background_resp
|
|||
assert seen["response"] is response
|
||||
|
||||
|
||||
def _output_passing_callbacks() -> list[CustomGuardrail]:
|
||||
class OutputPassingGuardrail(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)]
|
||||
|
||||
|
||||
def _claimed_post_call_pipeline_data(*policy_names: str, extra_guardrails: dict[str, list[str]] | None = None):
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
|
||||
step = {"guardrail": "gr-post", "on_pass": "allow", "on_fail": "block"}
|
||||
get_policy_registry().load_policies(
|
||||
{
|
||||
policy_name: {
|
||||
"guardrails": {"add": ["gr-post", *(extra_guardrails or {}).get(policy_name, [])]},
|
||||
"pipeline": {"mode": "post_call", "steps": [step]},
|
||||
}
|
||||
for policy_name in policy_names
|
||||
}
|
||||
)
|
||||
pipeline = GuardrailPipeline(mode="post_call", steps=[PipelineStep(**step)])
|
||||
return {
|
||||
"model": "m",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"metadata": {
|
||||
"_guardrail_pipelines": [(policy_name, pipeline) for policy_name in policy_names],
|
||||
"_pipeline_managed_guardrails": {"gr-post"},
|
||||
"applied_policies": list(policy_names),
|
||||
"applied_guardrails": ["gr-post", *(g for gs in (extra_guardrails or {}).values() for g in gs)],
|
||||
"policy_sources": {policy_name: "model:m" for policy_name in policy_names},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def clear_policy_registry():
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
|
||||
yield
|
||||
get_policy_registry().clear()
|
||||
|
||||
|
||||
@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
|
||||
):
|
||||
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")
|
||||
|
||||
out = await proxy_logging.post_call_success_hook(
|
||||
data=data, response=_background_response("queued"), user_api_key_dict=make_user_api_key_auth()
|
||||
)
|
||||
|
||||
assert out.status == "queued"
|
||||
assert "applied_policies" not in data["metadata"]
|
||||
assert "policy_sources" not in data["metadata"]
|
||||
assert "applied_guardrails" not in data["metadata"]
|
||||
|
||||
|
||||
@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
|
||||
):
|
||||
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"]}
|
||||
)
|
||||
|
||||
await proxy_logging.post_call_success_hook(
|
||||
data=data, response=_background_response("in_progress"), user_api_key_dict=make_user_api_key_auth()
|
||||
)
|
||||
|
||||
assert data["metadata"]["applied_policies"] == ["input-and-output-governance"]
|
||||
assert data["metadata"]["applied_guardrails"] == ["gr-pre"]
|
||||
assert data["metadata"]["policy_sources"] == {"input-and-output-governance": "model:m"}
|
||||
|
||||
|
||||
@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
|
||||
):
|
||||
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")
|
||||
|
||||
await proxy_logging.post_call_success_hook(
|
||||
data=data, response=_background_response("completed", text="fine"), user_api_key_dict=make_user_api_key_auth()
|
||||
)
|
||||
|
||||
assert data["metadata"]["applied_policies"] == ["response-governance"]
|
||||
assert data["metadata"]["policy_sources"] == {"response-governance": "model:m"}
|
||||
assert data["metadata"]["applied_guardrails"] == ["gr-post"]
|
||||
|
||||
|
||||
@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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue