diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 2eb9cfb5042..8e71f7a08ff 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -21,6 +21,7 @@ from litellm.types.guardrails import ( DynamicGuardrailParams, GuardrailEventHooks, LitellmParams, + LoggingOnlyScope, Mode, ) from litellm.types.llms.openai import AllMessageValues @@ -175,6 +176,7 @@ class CustomGuardrail(CustomLogger): use_native_lifecycle_hooks: ClassVar[bool] = False records_own_guardrail_information: ClassVar[bool] = False + logging_only_scope: LoggingOnlyScope | None def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks super().__init_subclass__(**kwargs) @@ -246,6 +248,7 @@ class CustomGuardrail(CustomLogger): self.run_in_parallel: bool = run_in_parallel self.scan_raw_request: bool = scan_raw_request self.only_scan_new_messages: bool = only_scan_new_messages + self.logging_only_scope = None if supported_event_hooks: ## validate event_hook is in supported_event_hooks @@ -803,6 +806,13 @@ class CustomGuardrail(CustomLogger): def uses_apply_guardrail_interface(self) -> bool: return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail + @classmethod + def supports_logging_only_scope(cls) -> bool: + return ( + cls.apply_guardrail is not CustomGuardrail.apply_guardrail + and cls.async_logging_hook is CustomGuardrail.async_logging_hook + ) + def _deployment_hook_target(self) -> "CustomLogger": if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks: return self @@ -934,7 +944,7 @@ class CustomGuardrail(CustomLogger): result: object, call_type: str, ) -> tuple[dict, object]: # mutable-ok: CustomLogger.async_logging_hook contract - """logging_only: run apply_guardrail on copies of the logged request/response and record the verdict.""" + """logging_only: scan copies of the logged request and/or response according to logging_only_scope.""" from litellm.llms import get_guardrail_translation_mapping if not self.uses_apply_guardrail_interface(): @@ -981,6 +991,21 @@ class CustomGuardrail(CustomLogger): "standard_logging_object": {**standard_logging_object, "guardrail_information": [*existing, *entries]}, }, result + def _copy_scratch_request_fields( + self, + kwargs: Mapping[str, object], + ) -> tuple[object, object]: + optional_params: Final = kwargs.get("optional_params") + try: + return ( + copy.deepcopy(kwargs.get("messages") or kwargs.get("input")), + copy.deepcopy(optional_params.get("tools") if isinstance(optional_params, Mapping) else None), + ) + except Exception: + if self.logging_only_scope == "output": + return None, None + raise + async def _scan_logged_call( self, kwargs: dict, # mutable-ok: CustomLogger.async_logging_hook contract @@ -989,22 +1014,28 @@ class CustomGuardrail(CustomLogger): output_translation: "BaseTranslation", scratch_metadata: dict, # mutable-ok: apply_guardrail records its verdict into request metadata ) -> None: - optional_params: Final = kwargs.get("optional_params") or {} - scratch_input: Final = copy.deepcopy(kwargs.get("messages") or kwargs.get("input")) + scratch_input, scratch_tools = self._copy_scratch_request_fields(kwargs) scratch_request: Final = { "model": kwargs.get("model"), "messages": scratch_input, "input": scratch_input, - "tools": copy.deepcopy(optional_params.get("tools")), + "tools": scratch_tools, "litellm_call_id": kwargs.get("litellm_call_id"), "metadata": scratch_metadata, } - await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self) - if response is None: + if self.logging_only_scope != "output": + try: + await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self) + except Exception as e: + verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e) + if response is None or self.logging_only_scope == "input": return - await output_translation.process_output_response( - response=copy.deepcopy(response), guardrail_to_apply=self, request_data=scratch_request - ) + try: + await output_translation.process_output_response( + response=copy.deepcopy(response), guardrail_to_apply=self, request_data=scratch_request + ) + except Exception as e: + verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e) def supports_scan_only_tool_results(self) -> bool: """Whether this guardrail can scan tool-result content. diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index fa4b36a03aa..2a4fa1548cd 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -10432,6 +10432,23 @@ "description": "Google Cloud location/region (e.g., us-central1)", "title": "Location" }, + "logging_only_scope": { + "anyOf": [ + { + "enum": [ + "input", + "output", + "both" + ], + "type": "string" + }, + { + "type": "null" + } + ], + "description": "which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking.", + "title": "Logging Only Scope" + }, "mask_request_content": { "anyOf": [ { @@ -11587,6 +11604,77 @@ "title": "GuardrailSubmissionSummary", "type": "object" }, + "GuardrailUIAddGuardrailSettings": { + "properties": { + "content_filter_settings": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "Content Filter Settings" + }, + "pii_entity_categories": { + "items": { + "$ref": "#/components/schemas/PiiEntityCategoryMap" + }, + "title": "Pii Entity Categories", + "type": "array" + }, + "providers_without_directional_logging_only_scope": { + "items": { + "type": "string" + }, + "title": "Providers Without Directional Logging Only Scope", + "type": "array" + }, + "supported_actions": { + "items": { + "type": "string" + }, + "title": "Supported Actions", + "type": "array" + }, + "supported_entities": { + "items": { + "type": "string" + }, + "title": "Supported Entities", + "type": "array" + }, + "supported_modes": { + "items": { + "type": "string" + }, + "title": "Supported Modes", + "type": "array" + }, + "supported_modes_by_provider": { + "additionalProperties": { + "items": { + "type": "string" + }, + "type": "array" + }, + "title": "Supported Modes By Provider", + "type": "object" + } + }, + "required": [ + "supported_entities", + "supported_actions", + "supported_modes", + "supported_modes_by_provider", + "providers_without_directional_logging_only_scope", + "pii_entity_categories" + ], + "title": "GuardrailUIAddGuardrailSettings", + "type": "object" + }, "HTTPValidationError": { "properties": { "detail": { @@ -12682,6 +12770,23 @@ "description": "Google Cloud location/region (e.g., us-central1)", "title": "Location" }, + "logging_only_scope": { + "anyOf": [ + { + "enum": [ + "input", + "output", + "both" + ], + "type": "string" + }, + { + "type": "null" + } + ], + "description": "which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking.", + "title": "Logging Only Scope" + }, "mask": { "anyOf": [ { @@ -13742,6 +13847,27 @@ "title": "PiiAction", "type": "string" }, + "PiiEntityCategoryMap": { + "properties": { + "category": { + "title": "Category", + "type": "string" + }, + "entities": { + "items": { + "type": "string" + }, + "title": "Entities", + "type": "array" + } + }, + "required": [ + "category", + "entities" + ], + "title": "PiiEntityCategoryMap", + "type": "object" + }, "PiiEntityType": { "enum": [ "CREDIT_CARD", @@ -15183,7 +15309,9 @@ "200": { "content": { "application/json": { - "schema": {} + "schema": { + "$ref": "#/components/schemas/GuardrailUIAddGuardrailSettings" + } } }, "description": "Successful Response" diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index acad9403ed4..e51de58f3c7 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -33,7 +33,11 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import ( build_sandbox_globals, compile_sandboxed, ) -from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry +from litellm.proxy.guardrails.guardrail_registry import ( + GuardrailRegistry, + _configured_event_hooks, + parse_tolerant_litellm_params, +) from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.repositories.prisma_protocols import TableActions @@ -404,7 +408,11 @@ async def create_guardrail( guardrail_id: Final = result.get("guardrail_id", "Unknown") try: - IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(guardrail=cast(Guardrail, result), source="db") + IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail( + guardrail=cast(Guardrail, result), + source="db", + reject_invalid_logging_only_scope=True, + ) verbose_proxy_logger.info( "Immediate sync: Successfully initialized guardrail '%s' (ID: %s)", guardrail_name, guardrail_id ) @@ -527,7 +535,10 @@ async def update_guardrail( guardrail_name: Final = result.get("guardrail_name", "Unknown") try: - IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=cast(Guardrail, result)) + IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( + guardrail=cast(Guardrail, result), + reject_invalid_logging_only_scope=True, + ) verbose_proxy_logger.info( "Immediate sync: Successfully updated guardrail '%s' (ID: %s)", guardrail_name, guardrail_id ) @@ -1206,19 +1217,33 @@ async def patch_guardrail( # Update litellm_params if default_on is provided or pii_entities_config is provided existing_litellm_params: Final = _as_str_object_mapping(dict(existing_guardrail.get("litellm_params", {}))) - litellm_params = LitellmParams(**existing_litellm_params) - if request.litellm_params is not None: - requested_litellm_params: Final = request.litellm_params.model_dump(exclude_unset=True) - litellm_params_dict: Final = litellm_params.model_dump(exclude_unset=True) - litellm_params_dict.update(requested_litellm_params) - merged_litellm_params: Final = _as_str_object_mapping(litellm_params_dict) - try: - litellm_params = LitellmParams(**merged_litellm_params) - except ValidationError as validation_error: - raise HTTPException( - status_code=422, - detail=f"Invalid guardrail configuration, update rejected: {validation_error}", - ) from validation_error + current_litellm_params: Final = parse_tolerant_litellm_params( + existing_litellm_params, + existing_guardrail.get("guardrail_name") or "Unknown", + ) + requested_litellm_params: Final = ( + request.litellm_params.model_dump(exclude_unset=True) if request.litellm_params is not None else {} + ) + merged_litellm_params: Final = _as_str_object_mapping( + {**current_litellm_params.model_dump(exclude_unset=True), **requested_litellm_params} + ) + try: + parsed_litellm_params: Final = LitellmParams(**merged_litellm_params) + except ValidationError as validation_error: + raise HTTPException( + status_code=422, + detail=f"Invalid guardrail configuration, update rejected: {validation_error}", + ) from validation_error + clear_stored_scope: Final = ( + "logging_only_scope" not in requested_litellm_params + and parsed_litellm_params.logging_only_scope is not None + and GuardrailEventHooks.logging_only.value not in _configured_event_hooks(parsed_litellm_params.mode) + ) + litellm_params: Final = ( + LitellmParams(**{**merged_litellm_params, "logging_only_scope": None}) + if clear_stored_scope + else parsed_litellm_params + ) # Update guardrail_info if provided guardrail_info: Final = ( @@ -1247,6 +1272,7 @@ async def patch_guardrail( try: IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( guardrail=guardrail, + reject_invalid_logging_only_scope="logging_only_scope" in requested_litellm_params, ) verbose_proxy_logger.info( "Immediate sync: Successfully updated guardrail '%s' (ID: %s)", guardrail_name, guardrail_id @@ -1260,15 +1286,7 @@ async def patch_guardrail( # the caller instead of a misleading 200. await GUARDRAIL_REGISTRY.update_guardrail_in_db( guardrail_id=guardrail_id, - guardrail=Guardrail( - guardrail_id=guardrail_id, - guardrail_name=existing_guardrail.get("guardrail_name") or "", - litellm_params=LitellmParams(**existing_litellm_params), - guardrail_info=existing_guardrail.get( - "guardrail_info", - {}, # mutable-ok: Guardrail's own constructor takes a plain dict - ), - ), + guardrail=existing_guardrail, prisma_client=prisma_client, ) raise HTTPException( @@ -1391,7 +1409,7 @@ async def get_guardrail_info(guardrail_id: str): tags=["Guardrails"], dependencies=[Depends(user_api_key_auth)], ) -async def get_guardrail_ui_settings(): +async def get_guardrail_ui_settings() -> GuardrailUIAddGuardrailSettings: """ Get the UI settings for the guardrails @@ -1425,12 +1443,18 @@ async def get_guardrail_ui_settings(): # above; it only runs on pre_call. {SupportedGuardrailIntegrations.HIDE_SECRETS.value: [GuardrailEventHooks.pre_call.value]} ) + providers_without_directional_logging_only_scope: Final = [ + provider + for provider, guardrail_class in guardrail_class_registry.items() + if not guardrail_class.supports_logging_only_scope() + ] return GuardrailUIAddGuardrailSettings( supported_entities=[entity.value for entity in PiiEntityType], supported_actions=[action.value for action in PiiAction], supported_modes=[mode.value for mode in GuardrailEventHooks], supported_modes_by_provider=supported_modes_by_provider, + providers_without_directional_logging_only_scope=providers_without_directional_logging_only_scope, pii_entity_categories=category_maps, content_filter_settings={ "prebuilt_patterns": get_pattern_metadata(), diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 0dc50cd6196..e4f00cdf80b 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -54,6 +54,7 @@ from .guardrail_hooks.llm_as_a_judge import ( initialize_guardrail as initialize_llm_as_a_judge, ) from .guardrail_initializers import ( + _configured_event_hooks, initialize_bedrock, initialize_hide_secrets, initialize_lakera, @@ -436,9 +437,46 @@ def _as_callback_tuple( return (initialized,) -def _configure_callback_scoping( +def _logging_only_scope_error( custom_guardrail_callback: CustomGuardrail, guardrail_name: str, litellm_params: LitellmParams +) -> str | None: + logging_only_scope: Final = litellm_params.logging_only_scope + if logging_only_scope is not None and GuardrailEventHooks.logging_only.value not in _configured_event_hooks( + litellm_params.mode + ): + return ( + f"Guardrail {guardrail_name}: logging_only_scope is set, but mode does not include logging_only, " + "so it would never apply. Add logging_only to mode or remove logging_only_scope." + ) + if logging_only_scope in ("input", "output") and not custom_guardrail_callback.supports_logging_only_scope(): + return ( + f"Guardrail {guardrail_name}: logging_only_scope={logging_only_scope!r} is not supported by this " + "guardrail, whose logging_only hook scans on its own. Remove logging_only_scope." + ) + return None + + +def _configure_callback_scoping( + custom_guardrail_callback: CustomGuardrail, + guardrail_name: str, + litellm_params: LitellmParams, + *, + reject_invalid_logging_only_scope: bool = False, ) -> None: + logging_only_scope: Final = litellm_params.logging_only_scope + logging_only_scope_error: Final = _logging_only_scope_error( + custom_guardrail_callback, guardrail_name, litellm_params + ) + if logging_only_scope_error is not None: + if reject_invalid_logging_only_scope: + raise ValueError(logging_only_scope_error) + verbose_proxy_logger.error( + "%s Ignoring logging_only_scope; the guardrail keeps its configured mode.", + logging_only_scope_error.replace("\r", "").replace("\n", ""), + ) + custom_guardrail_callback.logging_only_scope = None + else: + custom_guardrail_callback.logging_only_scope = logging_only_scope for scoping_param in ( "skip_system_message_in_guardrail", "skip_tool_message_in_guardrail", @@ -461,6 +499,24 @@ def _configure_callback_scoping( _apply_configured_bool_overrides(custom_guardrail_callback, litellm_params) +def parse_tolerant_litellm_params( + litellm_params_data: Mapping[str, object], + guardrail_name: str, +) -> LitellmParams: + try: + return LitellmParams(**litellm_params_data) + except ValidationError as validation_error: + if any(tuple(error["loc"]) != ("logging_only_scope",) for error in validation_error.errors()): + raise + verbose_proxy_logger.error( + "Guardrail %s: logging_only_scope=%r is not one of 'input', 'output' or 'both'. " + "Ignoring logging_only_scope; the guardrail keeps its configured mode.", + guardrail_name.replace("\r", "").replace("\n", ""), + str(litellm_params_data.get("logging_only_scope")).replace("\r", "").replace("\n", "")[:100], + ) + return LitellmParams(**{**litellm_params_data, "logging_only_scope": None}) + + class InMemoryGuardrailHandler: """ Class that handles initializing guardrails and adding them to the CallbackManager @@ -497,6 +553,8 @@ class InMemoryGuardrailHandler: config_file_path: str | None = None, llm_router: Optional["Router"] = None, source: Literal["db", "config"] = "config", + *, + reject_invalid_logging_only_scope: bool = False, ) -> Guardrail | None: """ Initialize a guardrail from a dictionary and add it to the litellm callback manager @@ -517,7 +575,10 @@ class InMemoryGuardrailHandler: verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data) if isinstance(litellm_params_data, dict): - litellm_params = LitellmParams(**litellm_params_data) + if reject_invalid_logging_only_scope: + litellm_params = LitellmParams(**litellm_params_data) + else: + litellm_params = parse_tolerant_litellm_params(litellm_params_data, guardrail["guardrail_name"]) else: litellm_params = litellm_params_data @@ -543,8 +604,18 @@ class InMemoryGuardrailHandler: config_file_path=config_file_path, llm_router=llm_router, ) - for custom_guardrail_callback in created_callbacks: - _configure_callback_scoping(custom_guardrail_callback, guardrail["guardrail_name"], litellm_params) + try: + for custom_guardrail_callback in created_callbacks: + _configure_callback_scoping( + custom_guardrail_callback, + guardrail["guardrail_name"], + litellm_params, + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, + ) + except Exception: + for custom_guardrail_callback in created_callbacks: + litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback) + raise parsed_guardrail: Final = Guardrail( guardrail_id=guardrail.get("guardrail_id"), @@ -652,6 +723,8 @@ class InMemoryGuardrailHandler: guardrail_id: str, guardrail: Guardrail, source: Literal["db", "config"] = "db", + *, + reject_invalid_logging_only_scope: bool = False, ) -> None: """ Update a guardrail in memory: a changed name or litellm_params rebuilds the @@ -659,8 +732,12 @@ class InMemoryGuardrailHandler: previous instance and raises), anything else only refreshes the stored row """ updated_guardrail: Final = cast(Guardrail, {**guardrail, "guardrail_id": guardrail_id}) - if self._has_guardrail_params_changed(guardrail_id, updated_guardrail): - self.reinitialize_guardrail(guardrail=updated_guardrail, source=source) + if reject_invalid_logging_only_scope or self._has_guardrail_params_changed(guardrail_id, updated_guardrail): + self.reinitialize_guardrail( + guardrail=updated_guardrail, + source=source, + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, + ) return self.IN_MEMORY_GUARDRAILS[guardrail_id] = updated_guardrail self._sources[guardrail_id] = source @@ -747,6 +824,7 @@ class InMemoryGuardrailHandler: @staticmethod def _normalize_litellm_params_for_comparison( params: LitellmParams | Mapping[str, object] | None, + guardrail_name: str, ) -> Mapping[str, object] | None: """ Render litellm_params to a canonical dict so an in-memory LitellmParams and @@ -763,7 +841,7 @@ class InMemoryGuardrailHandler: return params.model_dump() if isinstance(params, dict): try: - return LitellmParams(**params).model_dump() + return parse_tolerant_litellm_params(params, guardrail_name).model_dump() except ValidationError as e: verbose_proxy_logger.warning( "Could not normalize guardrail litellm_params for comparison; treating the guardrail as changed. Error: %s", @@ -786,8 +864,12 @@ class InMemoryGuardrailHandler: return True # Compare litellm_params - existing_dict: Final = self._normalize_litellm_params_for_comparison(existing.get("litellm_params")) - new_dict: Final = self._normalize_litellm_params_for_comparison(new_guardrail.get("litellm_params")) + existing_dict: Final = self._normalize_litellm_params_for_comparison( + existing.get("litellm_params"), existing.get("guardrail_name", "Unknown") + ) + new_dict: Final = self._normalize_litellm_params_for_comparison( + new_guardrail.get("litellm_params"), new_guardrail.get("guardrail_name", "Unknown") + ) # Compare and identify specific differences changed_fields = {} @@ -813,6 +895,8 @@ class InMemoryGuardrailHandler: guardrail: Guardrail, config_file_path: str | None = None, source: Literal["db", "config"] = "config", + *, + reject_invalid_logging_only_scope: bool = False, ) -> Guardrail | None: """ Force re-initialization of a guardrail even if it exists in memory. @@ -842,7 +926,12 @@ class InMemoryGuardrailHandler: # instance instead of leaving the guardrail silently removed: a guardrail # that was enforcing must never fail open because an update was bad. try: - return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source) + return self.initialize_guardrail( + guardrail=guardrail, + config_file_path=config_file_path, + source=source, + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, + ) except Exception as init_error: if previous_guardrail is not None: verbose_proxy_logger.exception( @@ -857,7 +946,13 @@ class InMemoryGuardrailHandler: verbose_proxy_logger.exception("Restoring previous guardrail %s also failed", guardrail_id) raise ValueError(f"Guardrail initialization failed: {init_error}") from init_error - def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None: + def sync_guardrail_from_db( + self, + guardrail: Guardrail, + config_file_path: str | None = None, + *, + reject_invalid_logging_only_scope: bool = False, + ) -> Guardrail | None: """ Sync a guardrail from DB - initializes if new, re-initializes if changed. This is the method to call during DB polling. @@ -867,7 +962,7 @@ class InMemoryGuardrailHandler: verbose_proxy_logger.error("Cannot sync guardrail without guardrail_id") return None - if self._has_guardrail_params_changed(guardrail_id, guardrail): + if reject_invalid_logging_only_scope or self._has_guardrail_params_changed(guardrail_id, guardrail): guardrail_name: Final = guardrail.get("guardrail_name", "Unknown") verbose_proxy_logger.info( "Guardrail '%s' (ID: %s) params changed, re-initializing...", guardrail_name, guardrail_id @@ -876,6 +971,7 @@ class InMemoryGuardrailHandler: guardrail=guardrail, config_file_path=config_file_path, source="db", + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, ) # Params unchanged but the entry is still DB-backed; make sure the diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 46026c12d24..000c602b45d 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -899,6 +899,8 @@ class ContentFilterConfigModel(BaseModel): MCP_SECURITY_ON_VIOLATION: Final = frozenset({"block", "alert"}) +LoggingOnlyScope = Literal["input", "output", "both"] + class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch update guardrails api_key: str | None = Field(default=None, description="API key for the guardrail service") @@ -1140,6 +1142,14 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up ), ) + logging_only_scope: LoggingOnlyScope | None = Field( + default=None, + description=( + "which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' " + "(default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking." + ), + ) + @field_validator( "mode", "default_action", @@ -1308,6 +1318,7 @@ class GuardrailUIAddGuardrailSettings(BaseModel): supported_actions: list[str] supported_modes: list[str] supported_modes_by_provider: dict[str, list[str]] + providers_without_directional_logging_only_scope: list[str] pii_entity_categories: list[PiiEntityCategoryMap] content_filter_settings: dict[str, object] | None = None diff --git a/tests/integration/observability/_logging_only_scope_support.py b/tests/integration/observability/_logging_only_scope_support.py new file mode 100644 index 00000000000..75cacf22636 --- /dev/null +++ b/tests/integration/observability/_logging_only_scope_support.py @@ -0,0 +1,1068 @@ +from __future__ import annotations + +import asyncio +import base64 +import binascii +import json +import os +import re +import uuid +from collections.abc import Callable, Iterator, Mapping +from contextlib import contextmanager +from contextvars import ContextVar +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal + +import httpx +import pytest +import yaml +from anthropic import Anthropic, AsyncAnthropic +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, Scenario, eventually, object_value +from integration._support.database import read_rows, write_rows +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request, Wire +from integration._support.wire import wire_server as _wire_server +from openai import AsyncOpenAI, OpenAI +from openai.types.chat import ChatCompletionChunk +from pydantic import JsonValue, TypeAdapter + +from tests.integration.cost_calculation.cost_tracking_case import JsonResponse, SseResponse + +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + +Endpoint = Literal["chat", "messages", "responses"] + +ClientKind = Literal["openai_sync", "openai_async", "anthropic_sync", "anthropic_async", "httpx"] + +Direction = Literal["request", "response"] + +BASE_DEFAULT_NORMAL_DIRECTIONS: Final[Mapping[tuple[Endpoint, bool], tuple[Direction, ...]]] = MappingProxyType( + { + ("chat", False): ("request", "response"), + ("chat", True): ("request", "response"), + ("messages", False): ("request", "response"), + ("messages", True): ("request", "response"), + ("responses", False): ("request", "response"), + ("responses", True): ("request", "response"), + } +) + +BASE_DEFAULT_CACHE_HIT_DIRECTIONS: Final[Mapping[Endpoint, tuple[Direction, ...]]] = MappingProxyType( + { + "chat": ("request", "response"), + "messages": ("request", "response"), + "responses": ("request", "response"), + } +) + +BASE_DEFAULT_UPSTREAM_FAILURE_DIRECTIONS: Final[Mapping[Endpoint, tuple[Direction, ...]]] = MappingProxyType( + { + "chat": (), + "messages": (), + "responses": (), + } +) + +_AUDIT_RESPONSE_IDS: Final[ContextVar[tuple[str, ...]]] = ContextVar("audit_response_ids", default=()) + +_AUDIT_POLICY_REQUEST_COUNT: Final[ContextVar[int]] = ContextVar("audit_policy_request_count", default=0) + +_AUDIT_UPSTREAM_REQUEST_COUNT: Final[ContextVar[int]] = ContextVar("audit_upstream_request_count", default=0) + + +@dataclass(frozen=True, slots=True) +class CallerResult: + status: int + body: dict[str, JsonValue] + response_id: str + text: str + + +@dataclass(frozen=True, slots=True) +class ChaosCall: + index: int + endpoint: Endpoint + client_kind: ClientKind + model: str + stream: bool + prompt: str + call_id: str + + +@dataclass(frozen=True, slots=True) +class ChaosDeployment: + model_name: str + model: str + api_base: str + + +def _chaos_models(scenario: Scenario, marker: str) -> tuple[ChaosDeployment, ...]: + endpoints: Final[tuple[Endpoint, Endpoint, Endpoint]] = ("chat", "messages", "responses") + specs: Final = tuple((endpoint, False) for endpoint in endpoints) + tuple( + (endpoint, True) for endpoint in endpoints + ) + handles: Final = tuple( + register_scenario( + f"{marker}-{endpoint}-{'stream' if stream else 'complete'}", + _provider_response(endpoint, marker, f"synthetic K response {marker}", stream), + ) + for endpoint, stream in specs + ) + for handle in handles: + scenario.cleanups.callback(delete_scenario, handle) + + return tuple( + ChaosDeployment( + model_name=f"integration-{marker}-{endpoint}-{'stream' if stream else 'complete'}", + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=handle.api_base() if endpoint == "messages" else f"{handle.api_base()}/v1", + ) + for (endpoint, stream), handle in zip(specs, handles) + ) + + +def _chaos_model_list(deployments: tuple[ChaosDeployment, ...]) -> tuple[dict[str, JsonValue], ...]: + return tuple( + { + "model_name": deployment.model_name, + "litellm_params": { + "model": deployment.model, + "api_base": deployment.api_base, + "api_key": "synthetic-provider-key", + }, + } + for deployment in deployments + ) + + +def _chaos_control_configuration( + tmp_path: Path, + identity: str, + model_list: tuple[dict[str, JsonValue], ...], +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["model_list"] = list(model_list) + config["guardrails"] = [] + path: Final = tmp_path / f"{identity}-models.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _chaos_calls(deployments: tuple[ChaosDeployment, ...], marker: str) -> tuple[ChaosCall, ...]: + endpoints: Final[tuple[Endpoint, Endpoint, Endpoint]] = ("chat", "messages", "responses") + + def one(index: int) -> ChaosCall: + endpoint_index: Final = index % 3 + endpoint: Final = endpoints[endpoint_index] + stream: Final = (index // 3) % 2 == 1 + client_kind: Final[ClientKind] = ( + "openai_async" + if endpoint == "chat" and stream + else "openai_sync" + if endpoint in ("chat", "responses") + else "anthropic_async" + if stream + else "anthropic_sync" + ) + return ChaosCall( + index=index, + endpoint=endpoint, + client_kind=client_kind, + model=deployments[endpoint_index + 3 * int(stream)].model_name, + stream=stream, + prompt=f"synthetic K burst {marker}-{index}", + call_id=f"{marker}-k-{index}", + ) + + return tuple(one(index) for index in range(30)) + + +def _chaos_spend_minimums( + models: tuple[str, ...], calls: tuple[ChaosCall, ...], baseline_count: int +) -> tuple[int, ...]: + return tuple(baseline_count + sum(call.model == model for call in calls) for model in models) + + +def _proxy_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def _is_base_audit_leg() -> bool: + leg: Final = os.environ.get("LITELLM_LOGGING_ONLY_SCOPE_AUDIT_LEG", "head") + assert leg in ("base", "head"), leg + return leg == "base" + + +def _record_response_id(response_id: str) -> None: + response_ids: Final = _AUDIT_RESPONSE_IDS.get() + _record_response_ids((response_id,) if response_id not in response_ids else ()) + + +def _record_response_ids(response_ids: tuple[str, ...]) -> None: + current: Final = _AUDIT_RESPONSE_IDS.get() + _AUDIT_RESPONSE_IDS.set(tuple(dict.fromkeys((*current, *response_ids)))) + + +def _record_policy_request_count(count: int) -> None: + _AUDIT_POLICY_REQUEST_COUNT.set(_AUDIT_POLICY_REQUEST_COUNT.get() + count) + + +def _record_upstream_request_count(count: int) -> None: + _AUDIT_UPSTREAM_REQUEST_COUNT.set(_AUDIT_UPSTREAM_REQUEST_COUNT.get() + count) + + +@contextmanager +def wire_server(respond: Callable[[Request], Reply], port: int = 0, *, policy_edge: bool = True) -> Iterator[Wire]: + received: Final[SimpleQueue[Request]] = SimpleQueue() + + def record(request: Request) -> Reply: + received.put(request) + return respond(request) + + try: + with _wire_server(record, port=port) as server: + yield server + finally: + if policy_edge: + _record_policy_request_count(received.qsize()) + + +def _directions_for_scope(base_default: tuple[Direction, ...], scope: str | None) -> tuple[Direction, ...]: + if scope is None or scope == "both": + return base_default + selected_direction: Final = "request" if scope == "input" else "response" + return tuple(direction for direction in base_default if direction == selected_direction) + + +def _directions_for_audit_leg(base_default: tuple[Direction, ...], scope: str | None) -> tuple[Direction, ...]: + if _is_base_audit_leg(): + return base_default + return _directions_for_scope(base_default, scope) + + +def _response_text(endpoint: Endpoint, body: Mapping[str, JsonValue]) -> str: + if endpoint == "chat": + choices: Final = body.get("choices") + assert isinstance(choices, list) and choices, body + message: Final = object_value(object_value(choices[0])["message"]) + return str(message["content"]) + if endpoint == "messages": + content: Final = body.get("content") + assert isinstance(content, list), body + return "".join( + str(object_value(block)["text"]) + for block in content + if isinstance(block, dict) and isinstance(block.get("text"), str) + ) + output: Final = body.get("output") + assert isinstance(output, list), body + return "".join( + _response_text_from_blocks(object_value(item).get("content")) + for item in output + if isinstance(item, dict) and object_value(item).get("type") == "message" + ) + + +def _response_text_from_blocks(value: JsonValue | None) -> str: + if not isinstance(value, list): + return "" + return "".join( + str(object_value(block)["text"]) + for block in value + if isinstance(block, dict) and isinstance(block.get("text"), str) + ) + + +def _chat_chunk_text(chunk: ChatCompletionChunk) -> str: + return "".join(choice.delta.content for choice in chunk.choices if isinstance(choice.delta.content, str)) + + +def _caller_result(endpoint: Endpoint, status: int, body: Mapping[str, JsonValue]) -> CallerResult: + response_id: Final = body.get("id") + assert isinstance(response_id, str), body + _record_response_id(response_id) + normalized: Final = JSON_OBJECT.validate_python(dict(body)) + return CallerResult(status, normalized, response_id, _response_text(endpoint, normalized)) + + +def _stream_result(endpoint: Endpoint, response_id: str, text: str) -> CallerResult: + _record_response_id(response_id) + body: Final = JSON_OBJECT.validate_python({"id": response_id, "text": text}) + return CallerResult(200, body, response_id, text) + + +def _response_body_without_ids(value: JsonValue) -> JsonValue: + if isinstance(value, dict): + return {key: _response_body_without_ids(item) for key, item in value.items() if key != "id"} + if isinstance(value, list): + return [_response_body_without_ids(item) for item in value] + return value + + +def _call_sync( + client_kind: ClientKind, + endpoint: Endpoint, + proxy_url: str, + key: str, + model: str, + prompt: str, + stream: bool, + call_id: str, +) -> CallerResult: + if client_kind == "httpx": + body: Final[dict[str, JsonValue]] = {"model": model, "messages": [{"role": "user", "content": prompt}]} + with httpx.Client(base_url=proxy_url, timeout=30, trust_env=False) as client: + response: Final = client.post( + "/v1/chat/completions", + json=body, + headers={"Authorization": f"Bearer {key}", "x-litellm-call-id": call_id}, + ) + return _caller_result(endpoint, response.status_code, JSON_OBJECT.validate_json(response.content)) + if client_kind == "anthropic_sync": + with Anthropic( + base_url=proxy_url, + api_key=key, + max_retries=0, + http_client=httpx.Client(timeout=30, trust_env=False), + ) as client: + if stream: + with client.messages.stream( + model=model, + max_tokens=32, + messages=[{"role": "user", "content": prompt}], + extra_headers={"x-litellm-call-id": call_id}, + ) as stream_response: + message: Final = stream_response.get_final_message() + return _stream_result( + endpoint, + message.id, + "".join(block.text for block in message.content if block.type == "text"), + ) + message: Final = client.messages.create( + model=model, + max_tokens=32, + messages=[{"role": "user", "content": prompt}], + extra_headers={"x-litellm-call-id": call_id}, + ) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(message.model_dump(mode="json"))) + assert client_kind == "openai_sync", client_kind + with OpenAI( + base_url=f"{proxy_url}/v1", + api_key=key, + max_retries=0, + http_client=httpx.Client(timeout=30, trust_env=False), + ) as client: + headers: Final = {"x-litellm-call-id": call_id} + if endpoint == "chat": + if stream: + chunks: Final = tuple( + client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + stream=True, + extra_headers=headers, + ) + ) + response_id: Final = chunks[0].id + text: Final = "".join(_chat_chunk_text(chunk) for chunk in chunks) + return _stream_result(endpoint, response_id, text) + response: Final = client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + extra_headers=headers, + ) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(response.model_dump(mode="json"))) + assert endpoint == "responses", endpoint + if stream: + events: Final = tuple( + client.responses.create(model=model, input=prompt, stream=True, extra_headers=headers) + ) + completed: Final = next(event.response for event in events if event.type == "response.completed") + body: Final = JSON_OBJECT.validate_python(completed.model_dump(mode="json")) + return _stream_result(endpoint, str(body["id"]), _response_text(endpoint, body)) + completion: Final = client.responses.create(model=model, input=prompt, extra_headers=headers) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(completion.model_dump(mode="json"))) + + +async def _call_async( + client_kind: ClientKind, + endpoint: Endpoint, + proxy_url: str, + key: str, + model: str, + prompt: str, + stream: bool, + call_id: str, +) -> CallerResult: + if client_kind == "anthropic_async": + async with AsyncAnthropic( + base_url=proxy_url, + api_key=key, + max_retries=0, + http_client=httpx.AsyncClient(timeout=30, trust_env=False), + ) as client: + if stream: + async with client.messages.stream( + model=model, + max_tokens=32, + messages=[{"role": "user", "content": prompt}], + extra_headers={"x-litellm-call-id": call_id}, + ) as stream_response: + message: Final = await stream_response.get_final_message() + return _stream_result( + endpoint, + message.id, + "".join(block.text for block in message.content if block.type == "text"), + ) + message: Final = await client.messages.create( + model=model, + max_tokens=32, + messages=[{"role": "user", "content": prompt}], + extra_headers={"x-litellm-call-id": call_id}, + ) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(message.model_dump(mode="json"))) + assert client_kind == "openai_async", client_kind + async with AsyncOpenAI( + base_url=f"{proxy_url}/v1", + api_key=key, + max_retries=0, + http_client=httpx.AsyncClient(timeout=30, trust_env=False), + ) as client: + headers: Final = {"x-litellm-call-id": call_id} + if endpoint == "chat": + assert stream, "The audit only uses the async OpenAI chat client for streaming rows" + stream_response: Final = await client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + stream=True, + extra_headers=headers, + ) + chunks: Final = tuple([chunk async for chunk in stream_response]) + response_id: Final = chunks[0].id + text: Final = "".join( + choice.delta.content + for chunk in chunks + for choice in chunk.choices + if isinstance(choice.delta.content, str) + ) + return _stream_result(endpoint, response_id, text) + assert endpoint == "responses", endpoint + if stream: + responses_stream: Final = await client.responses.create( + model=model, input=prompt, stream=True, extra_headers=headers + ) + events: Final = tuple([event async for event in responses_stream]) + completed: Final = next(event.response for event in events if event.type == "response.completed") + body: Final = JSON_OBJECT.validate_python(completed.model_dump(mode="json")) + return _stream_result(endpoint, str(body["id"]), _response_text(endpoint, body)) + completion: Final = await client.responses.create(model=model, input=prompt, extra_headers=headers) + return _caller_result(endpoint, 200, JSON_OBJECT.validate_python(completion.model_dump(mode="json"))) + + +def _call_client( + client_kind: ClientKind, + endpoint: Endpoint, + gateway: Gateway, + model: str, + prompt: str, + stream: bool, + call_id: str, +) -> CallerResult: + if client_kind in ("openai_async", "anthropic_async"): + return _record_caller_result( + asyncio.run( + _call_async(client_kind, endpoint, _proxy_url(gateway), gateway.key, model, prompt, stream, call_id) + ) + ) + return _record_caller_result( + _call_sync(client_kind, endpoint, _proxy_url(gateway), gateway.key, model, prompt, stream, call_id) + ) + + +def _record_caller_result(result: CallerResult) -> CallerResult: + _record_response_id(result.response_id) + return result + + +def _call_cache_client(endpoint: Endpoint, gateway: Gateway, model: str, prompt: str, call_id: str) -> CallerResult: + path: Final = { + "chat": "/v1/chat/completions", + "messages": "/v1/messages", + "responses": "/v1/responses", + }[endpoint] + body: Final = { + "chat": {"model": model, "messages": [{"role": "user", "content": prompt}]}, + "messages": {"model": model, "max_tokens": 32, "messages": [{"role": "user", "content": prompt}]}, + "responses": {"model": model, "input": prompt}, + }[endpoint] + with httpx.Client(timeout=30, trust_env=False) as client: + response: Final = client.post( + f"{_proxy_url(gateway)}{path}", + json=body, + headers={"Authorization": f"Bearer {gateway.key}", "x-litellm-call-id": call_id}, + ) + assert response.status_code == 200, response.text + return _caller_result( + endpoint, + response.status_code, + JSON_OBJECT.validate_json(response.content), + ) + + +def _provider_response(endpoint: Endpoint, _scenario_id: str, reply: str, stream: bool) -> JsonResponse | SseResponse: + response_id: Final = { + "chat": "chatcmpl-$UNIQUE_ID", + "messages": "msg_$UNIQUE_ID", + "responses": "resp_$UNIQUE_ID", + }[endpoint] + if endpoint == "chat": + if stream: + return SseResponse( + content_type="text/event-stream", + frames=( + f"data: {json.dumps({'id': response_id, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4o-mini', 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': reply}, 'finish_reason': None}]})}", + f"data: {json.dumps({'id': response_id, 'object': 'chat.completion.chunk', 'created': 1, 'model': 'gpt-4o-mini', 'choices': [{'index': 0, 'delta': {}, 'finish_reason': 'stop'}]})}", + "data: [DONE]", + ), + ) + return JsonResponse( + content_type="application/json", + body={ + "id": response_id, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": reply}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14}, + }, + ) + if endpoint == "messages": + if stream: + return SseResponse( + content_type="text/event-stream", + frames=( + f"event: message_start\ndata: {json.dumps({'type': 'message_start', 'message': {'id': response_id, 'type': 'message', 'role': 'assistant', 'content': [], 'model': 'claude-3-7-sonnet-20250219', 'stop_reason': None, 'stop_sequence': None, 'usage': {'input_tokens': 9, 'output_tokens': 0}}})}", + f"event: content_block_start\ndata: {json.dumps({'type': 'content_block_start', 'index': 0, 'content_block': {'type': 'text', 'text': ''}})}", + f"event: content_block_delta\ndata: {json.dumps({'type': 'content_block_delta', 'index': 0, 'delta': {'type': 'text_delta', 'text': reply}})}", + f"event: content_block_stop\ndata: {json.dumps({'type': 'content_block_stop', 'index': 0})}", + f"event: message_delta\ndata: {json.dumps({'type': 'message_delta', 'delta': {'stop_reason': 'end_turn', 'stop_sequence': None}, 'usage': {'output_tokens': 5}})}", + f"event: message_stop\ndata: {json.dumps({'type': 'message_stop'})}", + ), + ) + return JsonResponse( + content_type="application/json", + body={ + "id": response_id, + "type": "message", + "role": "assistant", + "model": "claude-3-7-sonnet-20250219", + "content": [{"type": "text", "text": reply}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 9, "output_tokens": 5}, + }, + ) + if stream: + completed: Final = { + "id": response_id, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "id": "msg_$UNIQUE_ID", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": reply, "annotations": []}], + } + ], + "usage": {"input_tokens": 9, "output_tokens": 5, "total_tokens": 14}, + } + return SseResponse( + content_type="text/event-stream", + frames=( + f"data: {json.dumps({'type': 'response.created', 'response': {'id': response_id, 'object': 'response', 'created_at': 1, 'status': 'in_progress', 'model': 'gpt-4.1-mini', 'output': []}})}", + f"data: {json.dumps({'type': 'response.output_text.delta', 'item_id': 'msg_$UNIQUE_ID', 'output_index': 0, 'content_index': 0, 'delta': reply})}", + f"data: {json.dumps({'type': 'response.completed', 'response': completed})}", + ), + ) + return JsonResponse( + content_type="application/json", + body={ + "id": response_id, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "id": "msg_$UNIQUE_ID", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": reply, "annotations": []}], + } + ], + "usage": {"input_tokens": 9, "output_tokens": 5, "total_tokens": 14}, + }, + ) + + +def _configuration( + tmp_path: Path, + identity: str, + policy_url: str, + scope: str | None, + *, + include_scope: bool = True, + default_on: bool = True, + mode: str | list[str] = "logging_only", + cache: bool = False, + num_retries: int | None = None, + model_list: tuple[dict[str, JsonValue], ...] = (), +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = cache + if num_retries is not None: + config["litellm_settings"]["num_retries"] = num_retries + config["model_list"] = list(model_list) + params: Final = { + "guardrail": "generic_guardrail_api", + "mode": mode, + "default_on": default_on, + "api_base": policy_url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + **({"logging_only_scope": scope} if include_scope else {}), + } + config["guardrails"] = [{"guardrail_name": identity, "litellm_params": params}] + scope_name: Final = scope if scope is not None else "unset" + path: Final = tmp_path / f"{identity}-{scope_name}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _content_filter_configuration(tmp_path: Path, identity: str, scope: str, blocked_word: str) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "litellm_content_filter", + "mode": "logging_only", + "logging_only_scope": scope, + "default_on": True, + "blocked_words": [{"keyword": blocked_word, "action": "BLOCK"}], + }, + } + ] + path: Final = tmp_path / f"{identity}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _presidio_configuration( + tmp_path: Path, + identity: str, + analyzer_api_base: str, + anonymizer_api_base: str, + scope: str | None, +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + params: Final = { + "guardrail": "presidio", + "mode": "logging_only", + "default_on": True, + "presidio_analyzer_api_base": analyzer_api_base, + "presidio_anonymizer_api_base": anonymizer_api_base, + "pii_entities_config": {"PERSON": "MASK"}, + **({"logging_only_scope": scope} if scope is not None else {}), + } + config["guardrails"] = [{"guardrail_name": identity, "litellm_params": params}] + scope_name: Final = scope if scope is not None else "unset" + path: Final = tmp_path / f"{identity}-{scope_name}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _empty_proxy_configuration(tmp_path: Path, identity: str, reload_seconds: int = 30) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [] + config["general_settings"]["proxy_config_reload_interval_seconds"] = reload_seconds + path: Final = tmp_path / f"{identity}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _model_armor_configuration(tmp_path: Path, identity: str, api_endpoint: str, token_uri: str, scope: str) -> Path: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + private_key_pem: Final = private_key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ).decode() + credentials: Final = { + "type": "service_account", + "project_id": "synthetic-model-armor-project", + "private_key_id": "synthetic-key-id", + "private_key": private_key_pem, + "client_email": "integration-model-armor@synthetic-project.iam.gserviceaccount.com", + "client_id": "123456789012345678901", + "token_uri": token_uri, + "auth_uri": "https://accounts.google.com/o/oauth2/auth", + "auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs", + "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/integration", + } + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "model_armor", + "mode": "logging_only", + "logging_only_scope": scope, + "default_on": True, + "template_id": "synthetic-template", + "project_id": "synthetic-model-armor-project", + "location": "us-central1", + "credentials": json.dumps(credentials), + "api_endpoint": api_endpoint, + }, + } + ] + path: Final = tmp_path / f"{identity}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _insert_database_guardrail( + identity: str, + policy_url: str, + scope: str | None, + *, + mode: str = "pre_call", + default_on: bool = True, +) -> None: + params: Final = { + "guardrail": "generic_guardrail_api", + "mode": mode, + "default_on": default_on, + "api_base": policy_url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + **({"logging_only_scope": scope} if scope is not None else {}), + } + write_rows( + 'INSERT INTO "LiteLLM_GuardrailsTable" ' + "(guardrail_id, guardrail_name, litellm_params, guardrail_info, updated_at) " + "VALUES (%s, %s, %s::jsonb, %s::jsonb, NOW())", + (str(uuid.uuid5(uuid.NAMESPACE_URL, identity)), identity, json.dumps(params), "{}"), + ) + + +@contextmanager +def _database_guardrail( + identity: str, + policy_url: str, + scope: str | None, + *, + mode: str = "pre_call", + default_on: bool = True, +) -> Iterator[None]: + _insert_database_guardrail(identity, policy_url, scope, mode=mode, default_on=default_on) + try: + yield + finally: + _delete_database_guardrail(identity) + + +def _delete_database_guardrail(identity: str) -> None: + write_rows('DELETE FROM "LiteLLM_GuardrailsTable" WHERE guardrail_name=%s', (identity,)) + + +def _post_guardrail_body( + identity: str, + provider: str, + mode: str | list[str], + api_base: str, + scope: JsonValue, + include_scope: bool = True, +) -> dict[str, JsonValue]: + params: Final = { + "guardrail": provider, + "mode": mode, + "default_on": True, + **( + { + "presidio_analyzer_api_base": api_base, + "presidio_anonymizer_api_base": api_base, + "pii_entities_config": {"PERSON": "MASK"}, + } + if provider == "presidio" + else {"api_base": api_base, "api_key": "synthetic-guardrail-key"} + ), + **({"extra_headers": ["x-litellm-call-id"]} if provider == "generic_guardrail_api" else {}), + **({"logging_only_scope": scope} if include_scope else {}), + } + return { + "guardrail": { + "guardrail_name": identity, + "litellm_params": params, + "guardrail_info": {"description": "phase-12 logging scope audit"}, + } + } + + +def _create_guardrail(candidate: Gateway, identity: str, params: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + guardrail_params: Final = { + **params, + **({"extra_headers": ["x-litellm-call-id"]} if params.get("guardrail") == "generic_guardrail_api" else {}), + } + response: Final = candidate.request( + "POST", + "/guardrails", + { + "guardrail": { + "guardrail_name": identity, + "litellm_params": guardrail_params, + "guardrail_info": {"description": "phase-12 logging scope audit"}, + } + }, + ) + assert response.status_code == 200, response.text + return JSON_OBJECT.validate_json(response.content) + + +def _management_guardrail_rows(identity: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + object_value(row) + for row in read_rows( + "SELECT guardrail_id, guardrail_name, litellm_params, guardrail_info " + 'FROM "LiteLLM_GuardrailsTable" WHERE guardrail_name=%s', + (identity,), + ) + ) + + +def _drain_upstream(upstream_url: str) -> tuple[dict[str, JsonValue], ...]: + response: Final = httpx.get(f"{upstream_url.rstrip('/')}/__observations", trust_env=False, timeout=15) + response.raise_for_status() + requests: Final = object_value(JSON_OBJECT.validate_python(response.json())).get("requests") + assert isinstance(requests, list), response.text + _record_upstream_request_count(len(requests)) + return tuple(object_value(request) for request in requests) + + +def _json_contains_exact_string(value: JsonValue, expected: str) -> bool: + if isinstance(value, str): + return value == expected + if isinstance(value, list): + return any(_json_contains_exact_string(item, expected) for item in value) + if isinstance(value, dict): + return any(_json_contains_exact_string(item, expected) for item in value.values()) + return False + + +@pytest.fixture(autouse=True) +def _record_audit_properties( + request: pytest.FixtureRequest, + record_property: Callable[[str, object], None], + gateway: Gateway, +) -> Iterator[None]: + response_ids_token: Final = _AUDIT_RESPONSE_IDS.set(()) + policy_count_token: Final = _AUDIT_POLICY_REQUEST_COUNT.set(0) + upstream_count_token: Final = _AUDIT_UPSTREAM_REQUEST_COUNT.set(0) + try: + response: Final = httpx.get(f"{gateway.upstream_url.rstrip('/')}/__observations", trust_env=False, timeout=15) + response.raise_for_status() + yield + finally: + node_id: Final = request.node.nodeid + inventory_ids: Final = re.findall(r"[A-Z]{1,2}\d+", node_id) + record_property("node_id", node_id) + record_property("inventory_id", inventory_ids[0] if inventory_ids else "support") + record_property("response_ids", ",".join(_AUDIT_RESPONSE_IDS.get())) + record_property("policy_edge_request_count", str(_AUDIT_POLICY_REQUEST_COUNT.get())) + record_property("upstream_request_count", str(_AUDIT_UPSTREAM_REQUEST_COUNT.get())) + _AUDIT_RESPONSE_IDS.reset(response_ids_token) + _AUDIT_POLICY_REQUEST_COUNT.reset(policy_count_token) + _AUDIT_UPSTREAM_REQUEST_COUNT.reset(upstream_count_token) + + +def _policy_call_id_matches(payload: Mapping[str, JsonValue], call_id: str) -> bool: + actual: Final = payload.get("litellm_call_id") + if actual == call_id: + return True + headers: Final = payload.get("request_headers") + return isinstance(headers, dict) and any( + key.lower() == "x-litellm-call-id" and value == call_id for key, value in headers.items() + ) + + +def _policy_call_id(payload: Mapping[str, JsonValue]) -> str | None: + actual: Final = payload.get("litellm_call_id") + if isinstance(actual, str): + return actual + headers: Final = payload.get("request_headers") + if not isinstance(headers, dict): + return None + return next( + (value for key, value in headers.items() if key.lower() == "x-litellm-call-id" and isinstance(value, str)), + None, + ) + + +def _chat_request(gateway: Gateway, model: str, prompt: str, call_id: str) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + + +def _direction(payload: Mapping[str, JsonValue]) -> str: + value: Final = payload.get("input_type") + assert value in ("request", "response"), payload + return str(value) + + +def _cache_hit(value: JsonValue) -> bool: + return value is True or value == "True" + + +def _guardrail_mode_values(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, list): + return tuple(str(mode) for mode in value) + return (str(value),) + + +def _guardrail_mode_status_pairs( + entries: tuple[dict[str, JsonValue], ...], +) -> tuple[tuple[tuple[str, ...], str], ...]: + return tuple((_guardrail_mode_values(entry["guardrail_mode"]), str(entry["guardrail_status"])) for entry in entries) + + +def _spend_row_for_response_id(response_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (response_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + row: Final = object_value(rows[0]) + _record_response_ids((str(row["request_id"]),)) + return row + + +def _spend_rows(model: str, minimum: int) -> tuple[dict[str, JsonValue], ...]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, metadata, cache_hit FROM "LiteLLM_SpendLogs" WHERE model_group=%s', + (model,), + ), + lambda rows: len(rows) >= minimum, + seconds=70, + ) + response_ids: Final = tuple(str(row["request_id"]) for row in rows) + _record_response_ids(response_ids) + return rows + + +def _spend_rows_for_calls( + models: tuple[str, ...], + expected: tuple[tuple[str, str, str], ...], + *, + tolerate_missing: bool = False, +) -> tuple[dict[str, JsonValue], ...]: + model_placeholders: Final = ", ".join("%s" for _ in models) + query: Final = ( + "SELECT model_group, request_id, metadata, cache_hit " + f'FROM "LiteLLM_SpendLogs" WHERE model_group IN ({model_placeholders})' + ) + rows: Final = eventually( + lambda: tuple(read_rows(query, models)), + lambda values: all( + len(_spend_rows_matching_call(values, model, call_id)) == 1 for model, _, call_id in expected + ), + seconds=45 if tolerate_missing else 70, + return_last_on_timeout=tolerate_missing, + ) + _record_response_ids(tuple(response_id for _, response_id, _ in expected)) + return rows + + +def _spend_rows_matching_call( + rows: tuple[dict[str, JsonValue], ...], + model: str, + call_id: str, +) -> tuple[dict[str, JsonValue], ...]: + return tuple( + row + for row in rows + if row["model_group"] == model and object_value(row["metadata"]).get("litellm_call_id") == call_id + ) + + +def _spend_row_for_call_id(call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (call_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert str(rows[0]["request_id"]) == call_id, rows + row: Final = object_value(rows[0]) + _record_response_ids((str(row["request_id"]),)) + return row + + +def _guardrail_entries(row: Mapping[str, JsonValue]) -> tuple[dict[str, JsonValue], ...]: + metadata: Final = object_value(row["metadata"]) + entries: Final = metadata.get("guardrail_information") + if not isinstance(entries, list): + return () + return tuple(object_value(entry) for entry in entries) + + +def _response_id_matches(endpoint: Endpoint, request_id: str, response_id: str, scenario_id: str) -> bool: + if request_id == response_id: + return True + if endpoint != "responses" or not request_id.startswith("resp_"): + return False + encoded: Final = request_id.removeprefix("resp_") + padding: Final = "=" * (-len(encoded) % 4) + try: + decoded: Final = base64.urlsafe_b64decode(encoded + padding).decode("utf-8") + except (binascii.Error, UnicodeDecodeError): + return False + return response_id in decoded or scenario_id in decoded + + +def _assert_response_id(endpoint: Endpoint, request_id: str, response_id: str, scenario_id: str) -> None: + assert _response_id_matches(endpoint, request_id, response_id, scenario_id), ( + endpoint, + request_id, + response_id, + scenario_id, + ) diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index d377afb206c..7f7870c9fb5 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -4,6 +4,7 @@ import signal import socket import uuid from concurrent.futures import ThreadPoolExecutor +from datetime import datetime, timezone from pathlib import Path from typing import Final @@ -1796,3 +1797,160 @@ def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway, for response in responses: assert response.status_code == 200, response.text assert response.headers["content-type"].startswith("text/event-stream"), response.text + + +@pytest.mark.parametrize( + ("logging_only_scope", "scanned_directions"), + (("input", ("request",)), ("output", ("response",)), ("both", ("request", "response"))), +) +def test_logging_only_scope_observes_only_the_configured_direction_without_blocking( + gateway: Gateway, tmp_path: Path, logging_only_scope: str, scanned_directions: tuple[str, ...] +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + prompt: Final = "synthetic observed prompt " + identity + reply: Final = "synthetic observed reply " + identity + texts_by_direction: Final = {"request": [prompt], "response": [reply]} + + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api" + return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic observed denial"}).encode()) + + def provider(request: Request) -> Reply: + assert request.target == "/v1/chat/completions" + assert json.loads(request.body)["messages"] == [{"role": "user", "content": prompt}] + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": reply}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14}, + } + ).encode() + ) + + with wire_server(guardrail) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "logging_only_scope": logging_only_scope, + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + } + ] + path: Final = tmp_path / "logging-only-scope.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(api_base=upstream.url + "/v1") + response: Final = candidate.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": prompt}]} + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == reply, response.text + assert len(upstream.drain()) == 1 + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda values: len(values) == 1, + seconds=70, + ) + scans: Final = tuple(json.loads(scan.body) for scan in policy.drain()) + assert [(scan["input_type"], scan["texts"]) for scan in scans] == [ + (direction, texts_by_direction[direction]) for direction in scanned_directions + ], scans + entries: Final = object_value(rows[0]["metadata"])["guardrail_information"] + assert isinstance(entries, list), rows[0] + assert [ + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in map(object_value, entries) + ] == [(identity, "logging_only", "guardrail_intervened")] * len(scanned_directions), entries + today: Final = datetime.now(timezone.utc).date().isoformat() + guardrail_id: Final = next( + object_value(row)["guardrail_id"] + for row in candidate.get("/v2/guardrails/list")["guardrails"] + if object_value(row)["guardrail_name"] == identity + ) + detail: Final = eventually( + lambda: candidate.request( + "GET", + f"/guardrails/usage/detail/{guardrail_id}", + params={"start_date": today, "end_date": today}, + ).json(), + lambda body: body["requestsEvaluated"] >= len(scanned_directions), + seconds=30, + return_last_on_timeout=True, + ) + assert detail["requestsEvaluated"] == len(scanned_directions), detail + + +@pytest.mark.parametrize("logging_only_scope", ("input", "Input")) +def test_logging_only_scope_literal_or_mode_mismatch_is_ignored_at_load_and_keeps_blocking( + gateway: Gateway, tmp_path: Path, logging_only_scope: str +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + prompt: Final = "synthetic invalid-scope prompt pineapple " + identity + + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api" + return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic policy denial"}).encode()) + + def provider(request: Request) -> Reply: + assert request.target == "/v1/chat/completions" + assert json.loads(request.body)["messages"] == [{"role": "user", "content": prompt}] + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "unchanged provider reply"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14}, + } + ).encode() + ) + + with wire_server(guardrail) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "logging_only_scope": logging_only_scope, + "default_on": True, + "blocked_words": [{"keyword": "pineapple", "action": "BLOCK"}], + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + } + ] + path: Final = tmp_path / "invalid-scope-pre-call.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(api_base=upstream.url + "/v1") + response: Final = candidate.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": prompt}]} + ) + assert response.status_code == 400, response.text + assert "synthetic policy denial" in response.text, response.text + assert len(policy.drain()) == 1 + assert len(upstream.drain()) == 0 + guardrails: Final = candidate.get("/v2/guardrails/list")["guardrails"] + assert any(object_value(row)["guardrail_name"] == identity for row in guardrails), guardrails diff --git a/tests/integration/observability/test_logging_only_scope_chaos.py b/tests/integration/observability/test_logging_only_scope_chaos.py new file mode 100644 index 00000000000..d70073a5071 --- /dev/null +++ b/tests/integration/observability/test_logging_only_scope_chaos.py @@ -0,0 +1,762 @@ +from __future__ import annotations + +import signal +import socket +import threading +import time +import uuid +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from itertools import accumulate, repeat +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import httpx +import psutil +import pytest +from _logging_only_scope_support import ( + JSON_OBJECT, + CallerResult, + ChaosCall, + _assert_response_id, + _call_client, + _chaos_calls, + _chaos_control_configuration, + _chaos_model_list, + _chaos_models, + _chaos_spend_minimums, + _configuration, + _direction, + _directions_for_audit_leg, + _drain_upstream, + _guardrail_entries, + _is_base_audit_leg, + _json_contains_exact_string, + _policy_call_id_matches, + _spend_rows, + _spend_rows_for_calls, + _spend_rows_matching_call, + wire_server, +) +from _logging_only_scope_support import ( + _record_audit_properties as _record_audit_properties, +) +from anthropic import APIConnectionError as AnthropicAPIConnectionError +from integration._support.client import Gateway, eventually +from integration._support.process import owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request +from openai import APIConnectionError as OpenAIAPIConnectionError +from pydantic import JsonValue + + +def test_K1_policy_edge_restart_mid_burst_keeps_output_observation_fail_open(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-k1-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + with socket.socket() as reservation: + reservation.bind(("127.0.0.1", 0)) + policy_port: Final = reservation.getsockname()[1] + with gateway.scenario() as scenario: + deployments: Final = _chaos_models(scenario, marker) + models: Final = tuple(deployment.model_name for deployment in deployments) + model_list: Final = _chaos_model_list(deployments) + chaos_calls: Final = _chaos_calls(deployments, marker) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list) + with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy: + controls: Final = tuple( + _call_client( + call.client_kind, + call.endpoint, + control_proxy, + call.model, + call.prompt, + call.stream, + f"{call.call_id}-control", + ) + for call in chaos_calls + ) + control_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(control_upstream) == 30, control_upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in control_upstream) + for call in chaos_calls + ) + == (1,) * 30 + ), control_upstream + config: Final = _configuration( + tmp_path, + identity, + f"http://127.0.0.1:{policy_port}", + "output", + model_list=model_list, + ) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate: + starts: Final = tuple(threading.Event() for _ in range(3)) + + def run(index: int) -> tuple[int, CallerResult]: + call: Final = chaos_calls[index] + phase: Final = index // 10 + assert starts[phase].wait(timeout=90), (index, phase) + return index, _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + + with ThreadPoolExecutor(max_workers=30) as pool: + futures: Final = tuple(pool.submit(run, index) for index in range(30)) + try: + with wire_server(policy, port=policy_port) as initial_edge: + starts[0].set() + first: Final = tuple(futures[index].result(timeout=90) for index in range(10)) + eventually( + lambda: initial_edge.received.qsize(), + lambda count: count == 10 * len(expected_directions), + seconds=30, + ) + tuple( + _spend_rows(model, minimum) + for model, minimum in zip(models, _chaos_spend_minimums(models, chaos_calls[:10], 5)) + ) + starts[1].set() + middle: Final = tuple(futures[index].result(timeout=90) for index in range(10, 20)) + tuple( + _spend_rows(model, minimum) + for model, minimum in zip(models, _chaos_spend_minimums(models, chaos_calls[:20], 5)) + ) + with wire_server(policy, port=policy_port) as recovered_edge: + starts[2].set() + recovered: Final = tuple(futures[index].result(timeout=90) for index in range(20, 30)) + eventually( + lambda: recovered_edge.received.qsize(), + lambda count: count == 10 * len(expected_directions), + seconds=30, + ) + tuple(_spend_rows(model, 10) for model in models) + finally: + for start in starts: + start.set() + results: Final = first + middle + recovered + assert tuple(index for index, _ in results) == tuple(range(30)), results + assert all(result.status == controls[index].status for index, result in results), results + assert all(result.text == controls[index].text for index, result in results), results + candidate_ids: Final = tuple(result.response_id for _, result in results) + assert len(set(candidate_ids)) == 30, candidate_ids + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 30, observed_upstream + assert ( + tuple( + sum( + _json_contains_exact_string(observation["body"], call.prompt) + for observation in observed_upstream + ) + for call in chaos_calls + ) + == (1,) * 30 + ), observed_upstream + expected_success_ids: Final = frozenset( + call.call_id for call in chaos_calls if call.index < 10 or call.index >= 20 + ) + edge_payloads: Final = tuple( + JSON_OBJECT.validate_json(call.body) for call in initial_edge.drain() + recovered_edge.drain() + ) + assert len(edge_payloads) == 20 * len(expected_directions), edge_payloads + successful_calls: Final = tuple(call for call in chaos_calls if call.call_id in expected_success_ids) + for call in successful_calls: + payloads_for_call: Final = tuple( + payload for payload in edge_payloads if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple( + sorted(expected_directions) + ), ( + call.call_id, + payloads_for_call, + ) + assert all( + payload["texts"] + == ([call.prompt] if _direction(payload) == "request" else [results[call.index][1].text]) + for payload in payloads_for_call + ), payloads_for_call + rows: Final = _spend_rows_for_calls( + models, + tuple( + ( + chaos_calls[index].model, + response_id, + chaos_calls[index].call_id, + ) + for index, response_id in enumerate(candidate_ids) + ), + ) + for index, call in enumerate(chaos_calls): + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + assert len(matching_rows) == 1, (call.call_id, matching_rows) + row: Final = matching_rows[0] + expected_status: Final = "guardrail_failed_to_respond" if 10 <= index < 20 else "success" + entries: Final = _guardrail_entries(row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries + ) == tuple((identity, "logging_only", expected_status) for _ in expected_directions), (index, entries) + + +def test_K2_policy_edge_delay_does_not_delay_concurrent_callers(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-k2-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + + def policy(_request: Request) -> Reply: + time.sleep(2) + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + deployments: Final = _chaos_models(scenario, marker) + models: Final = tuple(deployment.model_name for deployment in deployments) + model_list: Final = _chaos_model_list(deployments) + chaos_calls: Final = _chaos_calls(deployments, marker) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list) + with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy: + controls: Final = tuple( + _call_client( + call.client_kind, + call.endpoint, + control_proxy, + call.model, + call.prompt, + call.stream, + f"{call.call_id}-control", + ) + for call in chaos_calls + ) + control_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(control_upstream) == 30, control_upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in control_upstream) + for call in chaos_calls + ) + == (1,) * 30 + ), control_upstream + with wire_server(policy) as edge: + config: Final = _configuration(tmp_path, identity, edge.url, "output", model_list=model_list) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate: + + def run(call: ChaosCall) -> tuple[int, CallerResult, float]: + started: Final = time.monotonic() + result: Final = _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + return call.index, result, time.monotonic() - started + + with ThreadPoolExecutor(max_workers=30) as pool: + results: Final = tuple(pool.map(run, chaos_calls)) + assert tuple(index for index, _, _ in results) == tuple(range(30)), results + assert all(result.status == controls[index].status for index, result, _ in results), results + assert all(result.text == controls[index].text for index, result, _ in results), results + assert all(duration < 2 for _, _, duration in results), results + eventually( + lambda: edge.received.qsize(), + lambda count: count == 30 * len(expected_directions), + seconds=30, + ) + edge_calls: Final = edge.drain() + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge_calls) + for call in chaos_calls: + payloads_for_call: Final = tuple( + payload for payload in payloads if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple( + sorted(expected_directions) + ), ( + call, + payloads_for_call, + ) + assert all( + payload["texts"] + == ([call.prompt] if _direction(payload) == "request" else [controls[call.index].text]) + for payload in payloads_for_call + ), payloads_for_call + response_ids: Final = tuple(result.response_id for _, result, _ in results) + assert len(set(response_ids)) == 30, response_ids + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 30, upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream) + for call in chaos_calls + ) + == (1,) * 30 + ), upstream + rows: Final = _spend_rows_for_calls( + models, + tuple( + (call.model, result.response_id, call.call_id) + for call, (_, result, _) in zip(chaos_calls, results) + ), + ) + for call, (_, result, _) in zip(chaos_calls, results): + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + assert len(matching_rows) == 1, (call, result.response_id, matching_rows) + row: Final = matching_rows[0] + entries: Final = _guardrail_entries(row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + + +@pytest.mark.timeout(180) +def test_K3_two_worker_sigkill_checks_post_kill_spend_rows( + gateway: Gateway, + tmp_path: Path, + record_property: Callable[[str, object], None], +) -> None: + identity: Final = f"logging-scope-k3-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + scan_started: Final = threading.Event() + release_scans: Final = threading.Event() + + def policy(_request: Request) -> Reply: + scan_started.set() + assert release_scans.wait(timeout=60), identity + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + deployments: Final = _chaos_models(scenario, marker) + models: Final = tuple(deployment.model_name for deployment in deployments) + model_list: Final = _chaos_model_list(deployments) + calls: Final = _chaos_calls(deployments, marker) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + expected_entries: Final = tuple((identity, "logging_only", "success") for _ in expected_directions) + control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list) + with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy: + controls: Final = tuple( + _call_client( + call.client_kind, + call.endpoint, + control_proxy, + call.model, + call.prompt, + call.stream, + f"{call.call_id}-control", + ) + for call in calls + ) + assert len(_drain_upstream(gateway.upstream_url)) == 30 + with wire_server(policy) as edge: + config: Final = _configuration(tmp_path, identity, edge.url, "output", model_list=model_list) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + root: Final = psutil.Process(owned.process.pid) + workers: Final = eventually( + lambda: tuple( + child + for child in root.children(recursive=True) + if any("spawn_main" in part for part in child.cmdline()) + ), + lambda children: len(children) == 2, + seconds=30, + ) + + def run(call: ChaosCall) -> tuple[int, CallerResult | None, str | None]: + try: + result: Final = _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + return call.index, result, None + except (OpenAIAPIConnectionError, AnthropicAPIConnectionError, httpx.RemoteProtocolError) as error: + return call.index, None, str(error) + + with ThreadPoolExecutor(max_workers=30) as pool: + futures: Final = tuple(pool.submit(run, call) for call in calls) + try: + assert eventually(lambda: scan_started.is_set(), bool, seconds=30) + eventually(lambda: edge.received.qsize(), lambda count: count >= 5, seconds=30) + workers[0].send_signal(signal.SIGKILL) + killed_workers, surviving_workers = psutil.wait_procs((workers[0],), timeout=10) + assert len(killed_workers) == 1 and not surviving_workers, ( + killed_workers, + surviving_workers, + ) + finally: + release_scans.set() + outcomes: Final = tuple(future.result(timeout=90) for future in futures) + assert owned.process.poll() is None, "Proxy supervisor exited after a worker was killed" + successful: Final = tuple( + (calls[index], result) for index, result, error in outcomes if result is not None and error is None + ) + assert successful, outcomes + assert all( + result.status == controls[call.index].status and result.text == controls[call.index].text + for call, result in successful + ), successful + pre_kill_expected: Final = tuple( + (call.model, result.response_id, call.call_id) for call, result in successful + ) + pre_kill_rows: Final = _spend_rows_for_calls( + models, + pre_kill_expected, + tolerate_missing=True, + ) + pre_kill_rows_by_call: Final = tuple( + (call, _spend_rows_matching_call(pre_kill_rows, call.model, call.call_id)) for call in calls + ) + assert all(len(rows) <= 1 for _, rows in pre_kill_rows_by_call), pre_kill_rows_by_call + pre_kill_missing_rows: Final = sum(not rows for _, rows in pre_kill_rows_by_call) + record_property("k3_pre_kill_missing_spend_rows", pre_kill_missing_rows) + for call, matching_rows in pre_kill_rows_by_call: + if not matching_rows: + continue + entries: Final = _guardrail_entries(matching_rows[0]) + assert tuple( + sorted( + ( + entry["guardrail_name"], + entry["guardrail_mode"], + entry["guardrail_status"], + ) + for entry in entries + ) + ) == tuple(sorted(expected_entries)), (call, entries) + + post_kill_templates: Final = calls[:6] + post_kill_calls: Final = tuple( + ChaosCall( + index=call.index, + endpoint=call.endpoint, + client_kind=call.client_kind, + model=call.model, + stream=call.stream, + prompt=f"synthetic K post-kill burst {marker}-{call.index}", + call_id=f"{marker}-k-post-kill-{call.index}", + ) + for call in post_kill_templates + ) + + def run_post_kill(call: ChaosCall) -> CallerResult: + return _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + + with ThreadPoolExecutor(max_workers=len(post_kill_calls)) as pool: + post_kill_futures: Final = tuple(pool.submit(run_post_kill, call) for call in post_kill_calls) + post_kill_results: Final = tuple(future.result(timeout=90) for future in post_kill_futures) + assert all( + result.status == controls[call.index].status and result.text == controls[call.index].text + for call, result in zip(post_kill_calls, post_kill_results) + ), post_kill_results + served: Final = successful + tuple(zip(post_kill_calls, post_kill_results)) + response_ids: Final = tuple(result.response_id for _, result in served) + assert len(set(response_ids)) == len(response_ids), response_ids + served_calls: Final = tuple(call for call, _ in served) + requested_call_ids: Final = frozenset(call.call_id for call in calls + post_kill_calls) + upstream: Final = _drain_upstream(gateway.upstream_url) + assert all( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream) == 1 + for call in served_calls + ), upstream + + def accumulate_policy_payloads( + collected: tuple[dict[str, JsonValue], ...], _: None + ) -> tuple[dict[str, JsonValue], ...]: + return collected + tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain()) + + payload_batches: Final = accumulate(repeat(None), accumulate_policy_payloads, initial=()) + + def has_expected_post_kill_scans(collected: tuple[dict[str, JsonValue], ...], call: ChaosCall) -> bool: + payloads_for_call: Final = tuple( + payload for payload in collected if _policy_call_id_matches(payload, call.call_id) + ) + return all( + sum(_direction(payload) == direction for payload in payloads_for_call) + >= expected_directions.count(direction) + for direction in expected_directions + ) + + payloads: Final = eventually( + lambda: next(payload_batches), + lambda collected: all(has_expected_post_kill_scans(collected, call) for call in post_kill_calls), + seconds=30, + ) + assert all( + any(_policy_call_id_matches(payload, call_id) for call_id in requested_call_ids) + and _direction(payload) in expected_directions + for payload in payloads + ), payloads + for call, result in zip(post_kill_calls, post_kill_results): + payloads_for_call: Final = tuple( + payload for payload in payloads if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple( + sorted(expected_directions) + ), (call, payloads_for_call) + assert all( + payload["texts"] + == ([call.prompt] if _direction(payload) == "request" else [controls[call.index].text]) + for payload in payloads_for_call + ), payloads_for_call + + post_kill_expected: Final = tuple( + (call.model, result.response_id, call.call_id) + for call, result in zip(post_kill_calls, post_kill_results) + ) + rows: Final = _spend_rows_for_calls( + models, + post_kill_expected, + tolerate_missing=True, + ) + all_candidate_calls: Final = calls + post_kill_calls + rows_by_call: Final = tuple( + (call, _spend_rows_matching_call(rows, call.model, call.call_id)) for call in all_candidate_calls + ) + assert all(len(matching_rows) <= 1 for _, matching_rows in rows_by_call), rows_by_call + for call, matching_rows in rows_by_call: + if not matching_rows: + continue + entries: Final = _guardrail_entries(matching_rows[0]) + assert tuple( + sorted( + ( + entry["guardrail_name"], + entry["guardrail_mode"], + entry["guardrail_status"], + ) + for entry in entries + ) + ) == tuple(sorted(expected_entries)), (call, entries) + for call, result in successful: + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + if matching_rows: + _assert_response_id( + call.endpoint, + str(matching_rows[0]["request_id"]), + result.response_id, + marker if call.endpoint == "responses" else call.call_id, + ) + for call, result in zip(post_kill_calls, post_kill_results): + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + assert len(matching_rows) == 1, (result.response_id, matching_rows) + _assert_response_id( + call.endpoint, + str(matching_rows[0]["request_id"]), + result.response_id, + marker if call.endpoint == "responses" else call.call_id, + ) + + +@pytest.mark.timeout(180) +def test_K4_proxy_restart_after_fifteen_responses_records_lost_ids( + gateway: Gateway, + tmp_path: Path, + record_property: Callable[[str, object], None], +) -> None: + identity: Final = f"logging-scope-k4-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + deployments: Final = _chaos_models(scenario, marker) + models: Final = tuple(deployment.model_name for deployment in deployments) + model_list: Final = _chaos_model_list(deployments) + calls: Final = _chaos_calls(deployments, marker) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + control_config: Final = _chaos_control_configuration(tmp_path, identity, model_list) + with owned_proxy(gateway, tmp_path, {}, config=control_config, workers=2) as control_proxy: + controls: Final = tuple( + _call_client( + call.client_kind, + call.endpoint, + control_proxy, + call.model, + call.prompt, + call.stream, + f"{call.call_id}-control", + ) + for call in calls + ) + control_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(control_upstream) == 30, control_upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in control_upstream) + for call in calls + ) + == (1,) * 30 + ), control_upstream + with wire_server(policy) as edge: + config: Final = _configuration(tmp_path, identity, edge.url, "output", model_list=model_list) + restart_gate: Final = threading.Event() + second_wave_ready: Final = threading.Event() + second_wave_barrier: Final = threading.Barrier(15, action=second_wave_ready.set) + restarted_gateways: Final[SimpleQueue[Gateway]] = SimpleQueue() + + def gateway_for_call(call: ChaosCall, first_gateway: Gateway) -> Gateway: + if call.index < 15: + return first_gateway + second_wave_barrier.wait(timeout=90) + assert restart_gate.wait(timeout=90) + return restarted_gateways.get() + + def run(call: ChaosCall, first_gateway: Gateway) -> tuple[int, CallerResult]: + candidate: Final = gateway_for_call(call, first_gateway) + result: Final = _call_client( + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + return call.index, result + + with ThreadPoolExecutor(max_workers=30) as pool: + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as first_proxy: + first_proxy_port: Final = first_proxy.gateway.client.base_url.port + assert first_proxy_port is not None + futures: Final = tuple(pool.submit(run, call, first_proxy.gateway) for call in calls) + first_results: Final = tuple(futures[index].result(timeout=90)[1] for index in range(15)) + assert all( + result.status == controls[call.index].status and result.text == controls[call.index].text + for call, result in zip(calls[:15], first_results) + ), first_results + assert eventually(lambda: second_wave_ready.is_set(), bool, seconds=30) + eventually( + lambda: edge.received.qsize(), + lambda count: count == 15 * len(expected_directions), + seconds=30, + ) + first_wave_expected: Final = tuple( + (call.model, result.response_id, call.call_id) + for call, result in zip(calls[:15], first_results) + ) + first_wave_rows: Final = _spend_rows_for_calls( + models, + first_wave_expected, + tolerate_missing=True, + ) + assert all( + len(_spend_rows_matching_call(first_wave_rows, model, call_id)) <= 1 + for model, _, call_id in first_wave_expected + ), first_wave_rows + first_wave_present_response_ids: Final = frozenset( + response_id + for model, response_id, call_id in first_wave_expected + if len(_spend_rows_matching_call(first_wave_rows, model, call_id)) == 1 + ) + first_wave_lost_response_ids: Final = ( + frozenset(response_id for _, response_id, _ in first_wave_expected) + - first_wave_present_response_ids + ) + record_property( + "K4_PRE_RESTART_LOST_RESPONSE_IDS", + tuple(sorted(first_wave_lost_response_ids)), + ) + record_property( + f"K4_PRE_RESTART_LOST_ROW_COUNT_{'base' if _is_base_audit_leg() else 'head'}", + len(first_wave_lost_response_ids), + ) + assert first_proxy.process.poll() is not None, first_proxy.process.pid + eventually( + lambda: tuple( + connection + for connection in psutil.net_connections(kind="tcp") + if connection.status == psutil.CONN_LISTEN and connection.laddr.port == first_proxy_port + ), + lambda listeners: not listeners, + seconds=30, + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as restarted_proxy: + for _ in range(15): + restarted_gateways.put(restarted_proxy.gateway) + restart_gate.set() + second_results: Final = tuple(futures[index].result(timeout=90)[1] for index in range(15, 30)) + assert all( + result.status == controls[call.index].status and result.text == controls[call.index].text + for call, result in zip(calls[15:], second_results) + ), second_results + eventually( + lambda: edge.received.qsize(), + lambda count: count == 30 * len(expected_directions), + seconds=30, + ) + results: Final = first_results + second_results + expected_response_ids: Final = frozenset(result.response_id for result in results) + assert len(expected_response_ids) == 30, results + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 30, upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream) + for call in calls + ) + == (1,) * 30 + ), upstream + edge_payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain()) + assert len(edge_payloads) == 30 * len(expected_directions), edge_payloads + for call, result in zip(calls, results): + payloads_for_call: Final = tuple( + payload for payload in edge_payloads if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple(sorted(_direction(payload) for payload in payloads_for_call)) == tuple( + sorted(expected_directions) + ), ( + call, + payloads_for_call, + ) + assert all( + payload["texts"] == ([call.prompt] if _direction(payload) == "request" else [result.text]) + for payload in payloads_for_call + ), payloads_for_call + post_restart_expected: Final = tuple( + (call.model, result.response_id, call.call_id) for call, result in zip(calls[15:], second_results) + ) + rows: Final = _spend_rows_for_calls( + models, + post_restart_expected, + tolerate_missing=True, + ) + record_property("K4_RESPONSE_IDS", tuple(sorted(expected_response_ids))) + for call, result in zip(calls, results): + matching_rows: Final = _spend_rows_matching_call(rows, call.model, call.call_id) + assert len(matching_rows) <= 1, (result.response_id, matching_rows) + if call.index >= 15: + assert len(matching_rows) == 1, (call.call_id, result.response_id, matching_rows) + if matching_rows: + entries: Final = _guardrail_entries(matching_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), ( + call.call_id, + entries, + ) diff --git a/tests/integration/observability/test_logging_only_scope_config.py b/tests/integration/observability/test_logging_only_scope_config.py new file mode 100644 index 00000000000..9ef53da04b4 --- /dev/null +++ b/tests/integration/observability/test_logging_only_scope_config.py @@ -0,0 +1,1516 @@ +from __future__ import annotations + +import json +import threading +import uuid +from collections.abc import Callable +from pathlib import Path +from typing import Final +from urllib.parse import parse_qs + +import pytest +from _logging_only_scope_support import ( + JSON_OBJECT, + Direction, + _assert_response_id, + _call_client, + _configuration, + _content_filter_configuration, + _create_guardrail, + _delete_database_guardrail, + _direction, + _directions_for_audit_leg, + _drain_upstream, + _empty_proxy_configuration, + _guardrail_entries, + _guardrail_mode_values, + _insert_database_guardrail, + _is_base_audit_leg, + _management_guardrail_rows, + _model_armor_configuration, + _policy_call_id, + _policy_call_id_matches, + _post_guardrail_body, + _presidio_configuration, + _provider_response, + _response_body_without_ids, + _spend_row_for_call_id, + _spend_rows, + wire_server, +) +from _logging_only_scope_support import ( + _record_audit_properties as _record_audit_properties, +) +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.database import read_rows, write_rows +from integration._support.process import owned_proxy, owned_proxy_process +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request +from pydantic import JsonValue + + +@pytest.mark.parametrize( + ("row_id", "scope", "blocked_side"), + ( + pytest.param("F1", "output", "response", id="F1-native-filter-response"), + pytest.param("F2", "input", "request", id="F2-native-filter-request"), + ), +) +def test_native_content_filter_scope_logs_without_blocking( + gateway: Gateway, tmp_path: Path, row_id: str, scope: str, blocked_side: str +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + blocked_word: Final = f"pineapple{uuid.uuid4().hex[:6]}" + prompt: Final = ( + f"synthetic {blocked_word} request {identity}" + if blocked_side == "request" + else f"synthetic clean request {identity}" + ) + reply: Final = ( + f"synthetic {blocked_word} response {identity}" + if blocked_side == "response" + else f"synthetic clean response {identity}" + ) + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + config: Final = _content_filter_configuration(tmp_path, identity, scope, blocked_word) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = _call_client( + "openai_sync", "chat", gateway, model, prompt, False, f"{scenario_id}-baseline" + ) + assert baseline.status == 200 and baseline.text == reply, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + guarded: Final = _call_client( + "openai_sync", "chat", candidate, model, prompt, False, f"{scenario_id}-guarded" + ) + assert (guarded.status, _response_body_without_ids(guarded.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (guarded, baseline) + assert guarded.text == reply, guarded + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 1, observed_upstream + assert prompt in json.dumps(observed_upstream[0]["body"]), observed_upstream + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id("chat", str(guarded_rows[0]["request_id"]), guarded.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), scope) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries + ) == tuple( + ( + identity, + "logging_only", + "guardrail_intervened" if direction == blocked_side else "success", + ) + for direction in expected_directions + ), entries + intervened_entries: Final = tuple( + entry for entry in entries if entry["guardrail_status"] == "guardrail_intervened" + ) + assert len(intervened_entries) == 1, entries + assert intervened_entries[0]["guardrail_response"] is not None, entries + assert intervened_entries[0]["guardrail_response"] == "REDACTED_BY_LITELM", entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope"), + ( + pytest.param("F3", "output", id="F3-model-armor-response-only"), + pytest.param("F4", "input", id="F4-model-armor-request-only"), + ), +) +def test_model_armor_directional_scope_uses_real_service_account_oauth( + gateway: Gateway, + tmp_path: Path, + row_id: str, + scope: str, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic Model Armor prompt {identity}" + reply: Final = f"synthetic Model Armor response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + expected_directions: Final = _directions_for_audit_leg(("request", "response"), scope) + directions_to_collect: Final = _directions_for_audit_leg(("request", "response"), scope) + expected_fields: Final = tuple( + "userPromptData" if direction == "request" else "modelResponseData" for direction in expected_directions + ) + fields_to_collect: Final = tuple( + "userPromptData" if direction == "request" else "modelResponseData" for direction in directions_to_collect + ) + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def oauth(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/token", request + form: Final = parse_qs(request.body.decode()) + assert form.get("grant_type") == ["urn:ietf:params:oauth:grant-type:jwt-bearer"], form + assertion: Final = form.get("assertion") + assert assertion is not None and len(assertion[0].split(".")) == 3, form + return Reply(body=b'{"access_token":"synthetic-model-armor-token","expires_in":3600,"token_type":"Bearer"}') + + def model_armor(request: Request) -> Reply: + assert request.method == "POST", request + assert request.headers.get("authorization") == "Bearer synthetic-model-armor-token", request.headers + payload: Final = JSON_OBJECT.validate_json(request.body) + assert len(payload) == 1, payload + field: Final = next(iter(payload)) + assert field in ("userPromptData", "modelResponseData"), payload + expected_text: Final = prompt if field == "userPromptData" else reply + assert payload[field] == {"text": expected_text}, payload + return Reply(body=b'{"sanitizationResult":{"filterMatchState":"NO_MATCH_FOUND"}}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + baseline: Final = _call_client( + "openai_sync", "chat", gateway, model, prompt, False, f"{scenario_id}-baseline" + ) + assert baseline.status == 200 and baseline.text == reply, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(oauth, policy_edge=False) as token_edge, wire_server(model_armor) as armor_edge: + config: Final = _model_armor_configuration( + tmp_path, identity, armor_edge.url, token_edge.url + "/token", scope + ) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = _call_client( + "openai_sync", "chat", candidate, model, prompt, False, f"{scenario_id}-candidate" + ) + assert (result.status, _response_body_without_ids(result.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (result, baseline) + assert result.text == reply, result + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + eventually( + lambda: armor_edge.received.qsize(), + lambda count: count >= len(fields_to_collect), + seconds=30, + ) + eventually( + lambda: token_edge.received.qsize(), + lambda count: count >= 1, + seconds=30, + ) + armor_calls: Final = armor_edge.drain() + assert len(armor_calls) == len(fields_to_collect), armor_calls + armor_payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in armor_calls) + observed_fields: Final = tuple(next(iter(payload)) for payload in armor_payloads) + assert tuple(sorted(observed_fields)) == tuple(sorted(fields_to_collect)), armor_calls + assert all( + call.target.endswith( + ":sanitizeUserPrompt" if field == "userPromptData" else ":sanitizeModelResponse" + ) + for call, field in zip(armor_calls, observed_fields) + ), armor_calls + token_calls: Final = token_edge.drain() + assert token_calls and all(call.target == "/token" for call in token_calls), token_calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id("chat", str(guarded_rows[0]["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + observed_entries: Final = tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) + expected_entries: Final = tuple((identity, "logging_only", "success") for _ in expected_directions) + assert ( + tuple(sorted(observed_fields)), + tuple(sorted(observed_entries)), + ) == (tuple(sorted(expected_fields)), tuple(sorted(expected_entries))), ( + armor_calls, + entries, + rows, + ) + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope"), + ( + pytest.param("F5", "input", id="F5-presidio-input-scope-ignored"), + pytest.param("F6", "both", id="F6-presidio-both-scope-ignored"), + ), +) +def test_presidio_scope_matches_no_scope_behavior(gateway: Gateway, tmp_path: Path, row_id: str, scope: str) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + person: Final = f"synthetic Person {identity}" + prompt: Final = f"synthetic Presidio prompt {person}" + reply: Final = f"synthetic Presidio response {person}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def analyzer(response_seen: threading.Event) -> Callable[[Request], Reply]: + def handle(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/analyze", request + payload: Final = JSON_OBJECT.validate_json(request.body) + text: Final = payload["text"] + assert isinstance(text, str), payload + if reply in text: + response_seen.set() + start: Final = text.index(person) + return Reply( + body=json.dumps( + [{"entity_type": "PERSON", "start": start, "end": start + len(person), "score": 0.99}] + ).encode() + ) + + return handle + + def anonymizer(response_seen: threading.Event) -> Callable[[Request], Reply]: + def handle(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/anonymize", request + payload: Final = JSON_OBJECT.validate_json(request.body) + text: Final = payload["text"] + results: Final = payload["analyzer_results"] + assert isinstance(text, str) and isinstance(results, list) and len(results) == 1, payload + if reply in text: + response_seen.set() + return Reply(body=json.dumps({"text": text, "items": [{"entity_type": "PERSON"}]}).encode()) + + return handle + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + no_scope_analyzer_response_seen: Final = threading.Event() + no_scope_anonymizer_response_seen: Final = threading.Event() + scoped_analyzer_response_seen: Final = threading.Event() + scoped_anonymizer_response_seen: Final = threading.Event() + with ( + wire_server(analyzer(no_scope_analyzer_response_seen)) as no_scope_analyzer, + wire_server(anonymizer(no_scope_anonymizer_response_seen)) as no_scope_anonymizer, + wire_server(analyzer(scoped_analyzer_response_seen)) as scoped_analyzer, + wire_server(anonymizer(scoped_anonymizer_response_seen)) as scoped_anonymizer, + ): + no_scope_config: Final = _presidio_configuration( + tmp_path, identity, no_scope_analyzer.url, no_scope_anonymizer.url, None + ) + scoped_config: Final = _presidio_configuration( + tmp_path, identity, scoped_analyzer.url, scoped_anonymizer.url, scope + ) + with owned_proxy_process(gateway, tmp_path, {}, config=no_scope_config, workers=1) as no_scope_owned: + no_scope_proxy: Final = no_scope_owned.gateway + no_scope_guardrails: Final = no_scope_proxy.get("/v2/guardrails/list")["guardrails"] + assert identity in {object_value(item)["guardrail_name"] for item in no_scope_guardrails}, ( + no_scope_guardrails + ) + no_scope_result: Final = _call_client( + "openai_sync", "chat", no_scope_proxy, model, prompt, False, f"{scenario_id}-no-scope" + ) + assert no_scope_result.status == 200 and no_scope_result.text == reply, no_scope_result + no_scope_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(no_scope_upstream) == 1 and prompt in json.dumps(no_scope_upstream[0]["body"]), ( + no_scope_upstream + ) + eventually( + lambda: no_scope_analyzer_response_seen.is_set() and no_scope_anonymizer_response_seen.is_set(), + bool, + seconds=30, + ) + no_scope_analyzer_calls: Final = no_scope_analyzer.drain() + no_scope_anonymizer_calls: Final = no_scope_anonymizer.drain() + no_scope_call_counts: Final = ( + len(no_scope_analyzer_calls), + len(no_scope_anonymizer_calls), + ) + assert no_scope_call_counts[0] == no_scope_call_counts[1] > 0, no_scope_call_counts + with owned_proxy_process(gateway, tmp_path, {}, config=scoped_config, workers=1) as scoped_owned: + scoped_proxy: Final = scoped_owned.gateway + guardrails: Final = scoped_proxy.get("/v2/guardrails/list")["guardrails"] + assert identity in {object_value(item)["guardrail_name"] for item in guardrails}, guardrails + scoped_result: Final = _call_client( + "openai_sync", "chat", scoped_proxy, model, prompt, False, f"{scenario_id}-scoped" + ) + assert (scoped_result.status, _response_body_without_ids(scoped_result.body)) == ( + no_scope_result.status, + _response_body_without_ids(no_scope_result.body), + ), (scoped_result, no_scope_result) + scoped_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(scoped_upstream) == 1 and prompt in json.dumps(scoped_upstream[0]["body"]), ( + scoped_upstream + ) + eventually( + lambda: scoped_analyzer_response_seen.is_set() and scoped_anonymizer_response_seen.is_set(), + bool, + seconds=30, + ) + scoped_analyzer_calls: Final = scoped_analyzer.drain() + scoped_anonymizer_calls: Final = scoped_anonymizer.drain() + scoped_call_counts: Final = ( + len(scoped_analyzer_calls), + len(scoped_anonymizer_calls), + ) + assert scoped_call_counts == no_scope_call_counts, ( + scoped_call_counts, + no_scope_call_counts, + ) + assert tuple(sorted((call.method, call.target, call.body) for call in scoped_analyzer_calls)) == ( + tuple(sorted((call.method, call.target, call.body) for call in no_scope_analyzer_calls)) + ), (scoped_analyzer_calls, no_scope_analyzer_calls) + assert tuple( + sorted((call.method, call.target, call.body) for call in scoped_anonymizer_calls) + ) == tuple(sorted((call.method, call.target, call.body) for call in no_scope_anonymizer_calls)), ( + scoped_anonymizer_calls, + no_scope_anonymizer_calls, + ) + log_text: Final = scoped_owned.log.read_text() + if row_id == "F5" and not _is_base_audit_leg(): + assert "whose logging_only hook scans on its own" in log_text, log_text + if row_id == "F6": + assert "whose logging_only hook scans on its own" not in log_text, log_text + finally: + delete_scenario(upstream_handle) + + +def test_F7_guardrail_ui_settings_classify_directional_scope_support(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _empty_proxy_configuration(tmp_path, f"logging-scope-f7-{uuid.uuid4().hex}") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + response: Final = candidate.get("/guardrails/ui/add_guardrail_settings") + if _is_base_audit_leg(): + assert "providers_without_directional_logging_only_scope" not in response, response + return + unsupported: Final = response.get("providers_without_directional_logging_only_scope") + assert unsupported is not None, response + assert isinstance(unsupported, list), response + assert set(unsupported) == { + "lakera", + "lakera_v2", + "presidio", + "tool_permission", + "cisco_ai_defense", + "xecguard", + "repelloai", + "noma", + "mcp_jwt_signer", + "microsoft_purview", + "agent_365", + "guardrails_ai", + "mcp_security", + "conduct", + "javelin", + "pillar", + "lasso", + "dynamoai", + "pangea", + "aporia", + "aim", + "ibm_guardrails", + "semantic_guard", + "cato_networks", + }, response + assert not {"generic_guardrail_api", "litellm_content_filter", "model_armor"}.intersection(unsupported), response + + +@pytest.mark.parametrize( + ("row_id", "mode", "scope", "expected_status", "expected_directions", "expected_guardrail_mode"), + ( + pytest.param("G1", "pre_call", "input", 400, ("request",), "pre_call", id="G1-yaml-blocking-valid-scope"), + pytest.param("G2", "pre_call", "Input", 400, ("request",), "pre_call", id="G2-yaml-blocking-invalid-literal"), + pytest.param( + "G3", + "logging_only", + "sideways", + 200, + ("request", "response"), + "logging_only", + id="G3-yaml-logging-invalid-literal", + ), + ), +) +def test_yaml_scope_loading_keeps_guardrail_enforcement( + gateway: Gateway, + tmp_path: Path, + row_id: str, + mode: str, + scope: str, + expected_status: int, + expected_directions: tuple[Direction, ...], + expected_guardrail_mode: str, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic yaml blocking marker {identity}" + reply: Final = f"synthetic yaml response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic yaml denial"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, mode=mode) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert baseline.json()["choices"][0]["message"]["content"] == reply, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + guarded: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert guarded.status_code == expected_status, guarded.text + caller_body: Final = JSON_OBJECT.validate_python(guarded.json()) + candidate_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(candidate_upstream) == (0 if expected_status == 400 else 1), candidate_upstream + if expected_status == 200: + assert _response_body_without_ids(caller_body) == _response_body_without_ids( + JSON_OBJECT.validate_python(baseline.json()) + ), (guarded.text, baseline.text) + assert prompt in json.dumps(candidate_upstream[0]["body"]), candidate_upstream + else: + assert "synthetic yaml denial" in guarded.text, guarded.text + if expected_status == 200: + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(expected_directions), + seconds=30, + ) + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(payload) for payload in payloads)) == tuple( + sorted(expected_directions) + ), payloads + assert all(_policy_call_id_matches(payload, call_id) for payload in payloads), payloads + assert all( + payload["texts"] == ([prompt] if _direction(payload) == "request" else [reply]) + for payload in payloads + ), payloads + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple( + (identity, expected_guardrail_mode, "guardrail_intervened") for _ in expected_directions + ), entries + if expected_status == 400: + blocked_row: Final = _spend_row_for_call_id(call_id) + assert _guardrail_entries(blocked_row) == entries, blocked_row + else: + _assert_response_id( + "chat", + str(guarded_rows[0]["request_id"]), + str(caller_body["id"]), + scenario_id, + ) + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope"), + ( + pytest.param("G4", "output", id="G4-database-invalid-combination-before-boot"), + pytest.param("G5", "sideways", id="G5-database-invalid-literal-before-boot"), + ), +) +def test_database_guardrail_load_keeps_pre_call_blocking( + gateway: Gateway, tmp_path: Path, row_id: str, scope: str +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic DB blocked marker {identity}" + reply: Final = f"synthetic DB response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic DB denial"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + _insert_database_guardrail(identity, guardrail.url, scope) + try: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + loaded: Final = read_rows( + 'SELECT guardrail_name, litellm_params FROM "LiteLLM_GuardrailsTable" ' + "WHERE guardrail_name=%s", + (identity,), + ) + assert len(loaded) == 1, loaded + guarded: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert guarded.status_code == 400 and "synthetic DB denial" in guarded.text, guarded.text + assert _drain_upstream(gateway.upstream_url) == () + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(_direction(payload) for payload in payloads) == ("request",), payloads + assert _policy_call_id_matches(payloads[0], call_id), payloads + assert payloads[0]["texts"] == [prompt], payloads + spend_row: Final = _spend_row_for_call_id(call_id) + entries: Final = _guardrail_entries(spend_row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), entries + finally: + _delete_database_guardrail(identity) + finally: + delete_scenario(upstream_handle) + + +def test_G6_database_guardrail_polling_normalizes_invalid_scope_without_reinitializing( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = f"logging-scope-g6-{uuid.uuid4().hex}" + prompt: Final = f"synthetic DB polling marker {identity}" + reply: Final = f"synthetic DB polling response {identity}" + scenario_id: Final = f"phase12-g6-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def denial(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic polling denial"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(denial) as original_policy, wire_server(denial) as updated_policy: + _insert_database_guardrail(identity, original_policy.url, None) + try: + config: Final = _empty_proxy_configuration(tmp_path, identity, reload_seconds=1) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + first_call_id: Final = f"{scenario_id}-before-update" + first: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": first_call_id}, + ) + assert first.status_code == 400 and "synthetic polling denial" in first.text, first.text + first_policy_calls: Final = original_policy.drain() + assert len(first_policy_calls) == 1, first_policy_calls + updated_params: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": updated_policy.url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + "logging_only_scope": "sideways", + } + write_rows( + 'UPDATE "LiteLLM_GuardrailsTable" SET litellm_params=%s::jsonb, updated_at=NOW() ' + "WHERE guardrail_name=%s", + (json.dumps(updated_params), identity), + ) + + def probe_updated_policy() -> tuple[str, int]: + call_id: Final = f"{scenario_id}-poll-{uuid.uuid4().hex}" + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 400 and "synthetic polling denial" in response.text, ( + response.text + ) + return call_id, updated_policy.received.qsize() + + observed_call_id: Final = eventually( + probe_updated_policy, + lambda result: result[1] >= 1, + seconds=30, + )[0] + final_call_id: Final = f"{scenario_id}-after-sync" + final: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": final_call_id}, + ) + assert final.status_code == 400 and "synthetic polling denial" in final.text, final.text + old_payloads: Final = tuple( + JSON_OBJECT.validate_json(call.body) + for call in (*first_policy_calls, *original_policy.drain()) + ) + new_payloads: Final = tuple( + JSON_OBJECT.validate_json(call.body) for call in updated_policy.drain() + ) + assert all(payload["texts"] == [prompt] for payload in (*old_payloads, *new_payloads)), ( + old_payloads, + new_payloads, + ) + all_call_ids: Final = tuple( + str(_policy_call_id(payload)) for payload in (*old_payloads, *new_payloads) + ) + assert len(all_call_ids) == len(set(all_call_ids)), all_call_ids + assert first_call_id in all_call_ids and observed_call_id in all_call_ids, all_call_ids + assert final_call_id in all_call_ids, all_call_ids + assert len(new_payloads) >= 2, new_payloads + for call_id in (first_call_id, observed_call_id, final_call_id): + spend_row: Final = _spend_row_for_call_id(call_id) + entries: Final = _guardrail_entries(spend_row) + assert tuple( + ( + entry["guardrail_name"], + entry["guardrail_mode"], + entry["guardrail_status"], + ) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), (call_id, entries) + assert _drain_upstream(gateway.upstream_url) == () + finally: + _delete_database_guardrail(identity) + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "provider", "mode", "scope", "expected_status", "expected_scope", "expected_message"), + ( + pytest.param( + "H1", + "generic_guardrail_api", + "pre_call", + "input", + 400, + None, + "mode does not include logging_only", + id="H1-post-rejects-scope-outside-logging-only", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + "sideways", + 422, + None, + "logging_only_scope", + id="H2-post-rejects-invalid-string", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + 5, + 422, + None, + "logging_only_scope", + id="H2-post-rejects-invalid-number", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + "", + 422, + None, + "logging_only_scope", + id="H2-post-rejects-invalid-empty-string", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + ["input"], + 422, + None, + "logging_only_scope", + id="H2-post-rejects-invalid-list", + ), + pytest.param( + "H2", + "generic_guardrail_api", + "logging_only", + "x" * 5000, + 422, + None, + "logging_only_scope", + id="H2-post-rejects-oversized-string", + ), + pytest.param( + "H3", + "generic_guardrail_api", + "logging_only", + None, + 200, + None, + "", + id="H3-post-accepts-null-scope", + ), + pytest.param( + "H4", + "presidio", + "logging_only", + "input", + 400, + None, + "whose logging_only hook scans on its own", + id="H4-post-rejects-presidio-input-scope", + ), + pytest.param( + "H4", + "presidio", + "logging_only", + "both", + 200, + "both", + "", + id="H4-post-accepts-presidio-both-scope", + ), + ), +) +def test_management_post_validates_logging_only_scope( + gateway: Gateway, + tmp_path: Path, + row_id: str, + provider: str, + mode: str | list[str], + scope: JsonValue, + expected_status: int, + expected_scope: JsonValue, + expected_message: str, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + api_base: Final = "http://127.0.0.1:9" + config: Final = _empty_proxy_configuration(tmp_path, identity) + try: + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + before_rows: Final = _management_guardrail_rows(identity) + before_list: Final = candidate.get("/v2/guardrails/list")["guardrails"] + expected_leg_status: Final = 200 if _is_base_audit_leg() else expected_status + expected_leg_scope: Final = scope if _is_base_audit_leg() else expected_scope + response: Final = candidate.request( + "POST", + "/guardrails", + _post_guardrail_body(identity, provider, mode, api_base, scope), + ) + assert response.status_code == expected_leg_status, response.text + if expected_message and not _is_base_audit_leg(): + assert expected_message in response.text, response.text + if expected_leg_status == 200: + body: Final = JSON_OBJECT.validate_json(response.content) + params: Final = object_value(body["litellm_params"]) + assert params.get("logging_only_scope") == expected_leg_scope, body + rows: Final = _management_guardrail_rows(identity) + assert ( + len(rows) == 1 + and object_value(rows[0]["litellm_params"]).get("logging_only_scope") == expected_leg_scope + ), rows + else: + if row_id == "H2" and not _is_base_audit_leg(): + details: Final = object_value(JSON_OBJECT.validate_python(response.json())["detail"][0]) + assert details["type"] == "literal_error", details + assert details["loc"] == [ + "body", + "guardrail", + "litellm_params", + "logging_only_scope", + ], details + assert _management_guardrail_rows(identity) == before_rows == () + after_list: Final = candidate.get("/v2/guardrails/list")["guardrails"] + assert after_list == before_list, (before_list, after_list) + assert isinstance(after_list, list), after_list + assert identity not in {object_value(item)["guardrail_name"] for item in after_list}, after_list + finally: + _delete_database_guardrail(identity) + + +def test_H5_management_put_rejection_preserves_database_and_runtime(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h5-{uuid.uuid4().hex}" + prompt: Final = f"synthetic management marker {identity}" + reply: Final = f"synthetic management response {identity}" + scenario_id: Final = f"phase12-h5-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic management denial"}') + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + before_rows: Final = _management_guardrail_rows(identity) + assert len(before_rows) == 1, before_rows + before_info: Final = candidate.get(f"/guardrails/{guardrail_id}/info") + before_list: Final = tuple( + object_value(item) + for item in candidate.get("/v2/guardrails/list")["guardrails"] + if object_value(item)["guardrail_id"] == guardrail_id + ) + before_call_id: Final = f"{scenario_id}-before-put" + before_response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": before_call_id}, + ) + assert before_response.status_code == 400, before_response.text + before_policy: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert len(before_policy) == 1 and _policy_call_id_matches(before_policy[0], before_call_id), ( + before_policy + ) + assert before_policy[0]["texts"] == [prompt], before_policy + update: Final = candidate.request( + "PUT", + f"/guardrails/{guardrail_id}", + _post_guardrail_body( + identity, + "generic_guardrail_api", + "pre_call", + guardrail.url, + "output", + ), + ) + if _is_base_audit_leg(): + assert update.status_code == 200, update.text + updated_rows: Final = _management_guardrail_rows(identity) + assert len(updated_rows) == 1, updated_rows + assert object_value(updated_rows[0]["litellm_params"]).get("logging_only_scope") == "output", ( + updated_rows + ) + assert object_value(updated_rows[0]["litellm_params"])["mode"] == "pre_call", updated_rows + else: + assert update.status_code == 422, update.text + assert "logging_only_scope" in update.text and "logging_only" in update.text, update.text + assert _management_guardrail_rows(identity) == before_rows + after_info: Final = candidate.get(f"/guardrails/{guardrail_id}/info") + assert {key: value for key, value in before_info.items() if key != "updated_at"} == { + key: value for key, value in after_info.items() if key != "updated_at" + } + after_list: Final = tuple( + object_value(item) + for item in candidate.get("/v2/guardrails/list")["guardrails"] + if object_value(item)["guardrail_id"] == guardrail_id + ) + if _is_base_audit_leg(): + assert ( + len(after_list) == 1 + and object_value(after_list[0]["litellm_params"]).get("logging_only_scope") == "output" + ), after_list + else: + assert tuple( + {key: value for key, value in item.items() if key != "updated_at"} for item in after_list + ) == tuple( + {key: value for key, value in item.items() if key != "updated_at"} for item in before_list + ) + after_call_id: Final = f"{scenario_id}-after-put" + after_response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": after_call_id}, + ) + assert after_response.status_code == before_response.status_code, after_response.text + assert after_response.text == before_response.text, (after_response.text, before_response.text) + after_policy: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert len(after_policy) == 1 and _policy_call_id_matches(after_policy[0], after_call_id), after_policy + assert after_policy[0]["texts"] == [prompt], after_policy + assert _drain_upstream(gateway.upstream_url) == () + for call_id in (before_call_id, after_call_id): + entries: Final = _guardrail_entries(_spend_row_for_call_id(call_id)) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), (call_id, entries) + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +def test_H6_management_patch_clears_scope_when_switching_to_blocking_only_mode( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = f"logging-scope-h6-{uuid.uuid4().hex}" + prompt: Final = f"synthetic mode patch marker {identity}" + scenario_id: Final = f"phase12-h6-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, f"response {identity}", False) + ) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic mode patch denial"}') + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": ["pre_call", "logging_only"], + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + "logging_only_scope": "output", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + before: Final = _management_guardrail_rows(identity) + assert len(before) == 1, before + patched: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"mode": ["pre_call"]}}, + ) + assert patched.status_code == 200, patched.text + persisted: Final = _management_guardrail_rows(identity) + assert len(persisted) == 1, persisted + params: Final = object_value(persisted[0]["litellm_params"]) + assert params.get("logging_only_scope") == ("output" if _is_base_audit_leg() else None), persisted + assert params["mode"] == ["pre_call"], persisted + call_id: Final = f"{scenario_id}-after-patch" + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 400 and "synthetic mode patch denial" in response.text, response.text + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(_direction(call) for call in calls) == ("request",), calls + assert _policy_call_id_matches(calls[0], call_id), calls + assert calls[0]["texts"] == [prompt], calls + assert _drain_upstream(gateway.upstream_url) == () + entries: Final = _guardrail_entries(_spend_row_for_call_id(call_id)) + assert tuple( + ( + entry["guardrail_name"], + _guardrail_mode_values(entry["guardrail_mode"]), + entry["guardrail_status"], + ) + for entry in entries + ) == ((identity, ("pre_call",), "guardrail_intervened"),), entries + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +def test_H7_management_patch_rejection_preserves_pre_call_guardrail(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h7-{uuid.uuid4().hex}" + prompt: Final = f"synthetic invalid patch marker {identity}" + scenario_id: Final = f"phase12-h7-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, f"response {identity}", False) + ) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic invalid patch denial"}') + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + before: Final = _management_guardrail_rows(identity) + before_call_id: Final = f"{scenario_id}-before-patch" + response_before: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": before_call_id}, + ) + assert response_before.status_code == 400, response_before.text + policy_before: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + rejected: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": "input"}}, + ) + if _is_base_audit_leg(): + assert rejected.status_code == 200, rejected.text + updated: Final = _management_guardrail_rows(identity) + assert len(updated) == 1, updated + assert object_value(updated[0]["litellm_params"]).get("logging_only_scope") == "input", updated + else: + assert rejected.status_code == 422, rejected.text + assert "logging_only_scope" in rejected.text and "logging_only" in rejected.text, rejected.text + assert _management_guardrail_rows(identity) == before + after_call_id: Final = f"{scenario_id}-after-patch" + response_after: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": after_call_id}, + ) + assert response_after.status_code == response_before.status_code, response_after.text + policy_after: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(_direction(call) for call in policy_before) == ("request",), policy_before + assert tuple(_direction(call) for call in policy_after) == ("request",), policy_after + assert _policy_call_id_matches(policy_before[0], before_call_id), policy_before + assert _policy_call_id_matches(policy_after[0], after_call_id), policy_after + assert policy_before[0]["texts"] == [prompt], policy_before + assert policy_after[0]["texts"] == [prompt], policy_after + assert _drain_upstream(gateway.upstream_url) == () + for call_id in (before_call_id, after_call_id): + entries: Final = _guardrail_entries(_spend_row_for_call_id(call_id)) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), (call_id, entries) + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "stored_scope"), + ( + pytest.param("H8", "output", id="H8-patch-heals-invalid-mode-combination"), + pytest.param("H9", "sideways", id="H9-patch-heals-invalid-stored-literal"), + ), +) +def test_management_patch_default_on_heals_stored_scope( + gateway: Gateway, tmp_path: Path, row_id: str, stored_scope: str +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic healed scope marker {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, f"response {identity}", False) + ) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic healed scope denial"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + _insert_database_guardrail( + identity, + guardrail.url, + stored_scope, + default_on=False, + ) + try: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + before: Final = _management_guardrail_rows(identity) + assert len(before) == 1, before + guardrail_id: Final = string_value(before[0]["guardrail_id"]) + patched: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"default_on": True}}, + ) + assert patched.status_code == 200, patched.text + body: Final = JSON_OBJECT.validate_json(patched.content) + params: Final = object_value(body["litellm_params"]) + assert params["default_on"] is True, body + persisted: Final = _management_guardrail_rows(identity) + assert len(persisted) == 1, persisted + persisted_params: Final = object_value(persisted[0]["litellm_params"]) + assert persisted_params.get("logging_only_scope") == ( + stored_scope if _is_base_audit_leg() else None + ), persisted + assert persisted_params["default_on"] is True, persisted + call_id: Final = f"{scenario_id}-after-patch" + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 400 and "synthetic healed scope denial" in response.text, ( + response.text + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(_direction(call) for call in calls) == ("request",), calls + assert _policy_call_id_matches(calls[0], call_id), calls + assert calls[0]["texts"] == [prompt], calls + assert _drain_upstream(gateway.upstream_url) == () + entries: Final = _guardrail_entries(_spend_row_for_call_id(call_id)) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == ((identity, "pre_call", "guardrail_intervened"),), entries + finally: + _delete_database_guardrail(identity) + finally: + delete_scenario(upstream_handle) + + +def test_H10_management_patch_null_scope_restores_both_directions(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h10-{uuid.uuid4().hex}" + prompt: Final = f"synthetic reset-scope prompt {identity}" + reply: Final = f"synthetic reset-scope response {identity}" + scenario_id: Final = f"phase12-h10-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic reset-scope monitor"}') + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + control: Final = _call_client("openai_sync", "chat", candidate, model, prompt, False, control_call_id) + assert control.status == 200 and control.text == reply, control + assert len(_drain_upstream(gateway.upstream_url)) == 1 + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + "logging_only_scope": "output", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + patched: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": None}}, + ) + assert patched.status_code == 200, patched.text + persisted: Final = _management_guardrail_rows(identity) + assert len(persisted) == 1, persisted + params: Final = object_value(persisted[0]["litellm_params"]) + assert params.get("logging_only_scope") is None, persisted + result: Final = _call_client("openai_sync", "chat", candidate, model, prompt, False, guarded_call_id) + assert (result.status, _response_body_without_ids(result.body)) == ( + control.status, + _response_body_without_ids(control.body), + ), (result, control) + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1, upstream + assert prompt in json.dumps(upstream[0]["body"]), upstream + expected_directions: Final = ("request", "response") + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + policy_calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(call) for call in policy_calls)) == tuple(sorted(expected_directions)), ( + policy_calls + ) + assert all(_policy_call_id_matches(call, guarded_call_id) for call in policy_calls), policy_calls + assert all( + call["texts"] == ([prompt] if _direction(call) == "request" else [reply]) for call in policy_calls + ), policy_calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id("chat", str(guarded_rows[0]["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries + ) == ( + (identity, "logging_only", "guardrail_intervened"), + (identity, "logging_only", "guardrail_intervened"), + ), entries + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +def test_H11_management_patch_same_scope_is_idempotent_and_output_only(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h11-{uuid.uuid4().hex}" + prompt: Final = f"synthetic idempotent prompt {identity}" + reply: Final = f"synthetic idempotent response {identity}" + scenario_id: Final = f"phase12-h11-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic idempotent monitor"}') + + try: + with wire_server(policy) as guardrail: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + control: Final = _call_client("openai_sync", "chat", candidate, model, prompt, False, control_call_id) + assert control.status == 200 and control.text == reply, control + assert len(_drain_upstream(gateway.upstream_url)) == 1 + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "default_on": True, + "api_base": guardrail.url, + "api_key": "synthetic-guardrail-key", + "logging_only_scope": "output", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + before: Final = _management_guardrail_rows(identity) + first: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": "output"}}, + ) + second: Final = candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": "output"}}, + ) + assert first.status_code == 200, first.text + assert second.status_code == 200, second.text + assert _management_guardrail_rows(identity) == before + result: Final = _call_client("openai_sync", "chat", candidate, model, prompt, False, guarded_call_id) + assert (result.status, _response_body_without_ids(result.body)) == ( + control.status, + _response_body_without_ids(control.body), + ), (result, control) + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(call) for call in calls)) == tuple(sorted(expected_directions)), calls + assert all(_policy_call_id_matches(call, guarded_call_id) for call in calls), calls + assert all( + call["texts"] == ([prompt] if _direction(call) == "request" else [reply]) for call in calls + ), calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id("chat", str(guarded_rows[0]["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries + ) == tuple((identity, "logging_only", "guardrail_intervened") for _ in expected_directions), entries + finally: + _delete_database_guardrail(identity) + delete_scenario(upstream_handle) + + +def test_H12_management_reads_expose_typed_logging_only_scope(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h12-{uuid.uuid4().hex}" + try: + with owned_proxy( + gateway, tmp_path, {}, config=_empty_proxy_configuration(tmp_path, identity), workers=1 + ) as candidate: + created: Final = _create_guardrail( + candidate, + identity, + { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "default_on": True, + "api_base": "http://127.0.0.1:9", + "api_key": "synthetic-guardrail-key", + "logging_only_scope": "input", + }, + ) + guardrail_id: Final = string_value(created["guardrail_id"]) + info: Final = candidate.get(f"/guardrails/{guardrail_id}/info") + listed: Final = candidate.get("/v2/guardrails/list")["guardrails"] + assert isinstance(listed, list), listed + matches: Final = tuple( + object_value(item) for item in listed if object_value(item)["guardrail_id"] == guardrail_id + ) + assert len(matches) == 1, listed + assert object_value(info["litellm_params"])["logging_only_scope"] == "input", info + assert object_value(matches[0]["litellm_params"])["logging_only_scope"] == "input", matches + finally: + _delete_database_guardrail(identity) + + +def test_H13_unauthenticated_management_and_chat_requests_do_not_scan(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-h13-{uuid.uuid4().hex}" + prompt: Final = f"synthetic unauthorized prompt {identity}" + scenario_id: Final = f"phase12-h13-{uuid.uuid4().hex}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, f"response {identity}", False) + ) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, "input") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + unauthorized_post: Final = candidate.request( + "POST", + "/guardrails", + _post_guardrail_body( + f"{identity}-unauthorized", + "generic_guardrail_api", + "logging_only", + guardrail.url, + "input", + ), + key="synthetic-invalid-key", + ) + assert unauthorized_post.status_code == 401, unauthorized_post.text + unauthorized_chat: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + key="synthetic-invalid-key", + headers={"x-litellm-call-id": f"{scenario_id}-unauthorized"}, + ) + assert unauthorized_chat.status_code == 401, unauthorized_chat.text + assert _management_guardrail_rows(f"{identity}-unauthorized") == () + assert guardrail.drain() == () + assert _drain_upstream(gateway.upstream_url) == () + finally: + delete_scenario(upstream_handle) diff --git a/tests/integration/observability/test_logging_only_scope_runtime.py b/tests/integration/observability/test_logging_only_scope_runtime.py new file mode 100644 index 00000000000..dd00378401e --- /dev/null +++ b/tests/integration/observability/test_logging_only_scope_runtime.py @@ -0,0 +1,1579 @@ +from __future__ import annotations + +import json +import threading +import uuid +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from datetime import datetime, timezone +from pathlib import Path +from typing import Final + +import pytest +import yaml +from _logging_only_scope_support import ( + BASE_DEFAULT_CACHE_HIT_DIRECTIONS, + BASE_DEFAULT_NORMAL_DIRECTIONS, + BASE_DEFAULT_UPSTREAM_FAILURE_DIRECTIONS, + JSON_OBJECT, + ChaosCall, + ClientKind, + Direction, + Endpoint, + _assert_response_id, + _cache_hit, + _call_cache_client, + _call_client, + _configuration, + _database_guardrail, + _direction, + _directions_for_audit_leg, + _directions_for_scope, + _drain_upstream, + _empty_proxy_configuration, + _guardrail_entries, + _guardrail_mode_status_pairs, + _is_base_audit_leg, + _json_contains_exact_string, + _policy_call_id, + _policy_call_id_matches, + _provider_response, + _response_body_without_ids, + _response_text, + _spend_row_for_call_id, + _spend_row_for_response_id, + _spend_rows, + _spend_rows_for_calls, + _spend_rows_matching_call, + wire_server, +) +from _logging_only_scope_support import ( + _record_audit_properties as _record_audit_properties, +) +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.process import owned_proxy +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request +from pydantic import JsonValue + +from tests.integration.cost_calculation.cost_tracking_case import JsonResponse + + +@pytest.mark.parametrize( + ("row_id", "endpoint", "stream", "client_kind", "scope", "include_scope", "block_directions"), + ( + pytest.param("A1", "chat", False, "openai_sync", "input", True, (), id="A1-chat-input"), + pytest.param("A2", "chat", False, "openai_sync", "output", True, (), id="A2-chat-output"), + pytest.param("A3", "chat", False, "openai_sync", "both", True, (), id="A3-chat-both"), + pytest.param("A4", "chat", False, "httpx", None, False, (), id="A4-chat-missing-scope"), + pytest.param("A5", "chat", False, "httpx", None, True, (), id="A5-chat-null-scope"), + pytest.param("A6", "chat", True, "openai_async", "input", True, (), id="A6-chat-stream-async-input"), + pytest.param("A7", "chat", True, "openai_async", "output", True, (), id="A7-chat-stream-async-output"), + pytest.param("A8", "messages", False, "anthropic_sync", "input", True, (), id="A8-messages-input"), + pytest.param("A9", "messages", False, "anthropic_sync", "output", True, (), id="A9-messages-output"), + pytest.param( + "A10", + "messages", + True, + "anthropic_async", + "output", + True, + (), + id="A10-messages-stream-async-output", + ), + pytest.param( + "A11", + "messages", + True, + "anthropic_async", + "input", + True, + (), + id="A11-messages-stream-async-input", + ), + pytest.param("A12", "responses", False, "openai_async", "input", True, (), id="A12-responses-async-input"), + pytest.param("A13", "responses", False, "openai_async", "output", True, (), id="A13-responses-async-output"), + pytest.param("A14", "responses", True, "openai_sync", "output", True, (), id="A14-responses-stream-output"), + pytest.param("A15", "responses", True, "openai_sync", "input", True, (), id="A15-responses-stream-input"), + pytest.param("B1", "chat", False, "openai_sync", "input", True, ("request",), id="B1-logging-block-input"), + pytest.param("B2", "chat", False, "openai_sync", "output", True, ("response",), id="B2-logging-block-output"), + pytest.param( + "B3", + "chat", + False, + "openai_sync", + "both", + True, + ("request",), + id="B3-logging-block-both", + ), + ), +) +def test_runtime_directional_scope_matches_client_call_and_spend_log( + gateway: Gateway, + tmp_path: Path, + row_id: str, + endpoint: Endpoint, + stream: bool, + client_kind: ClientKind, + scope: str | None, + include_scope: bool, + block_directions: tuple[Direction, ...], +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic request {identity}" + reply: Final = f"synthetic response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + baseline_call_id: Final = f"{scenario_id}-baseline" + guarded_call_id: Final = f"{scenario_id}-guarded" + provider_response: Final = _provider_response(endpoint, scenario_id, reply, stream) + upstream_handle: Final = register_scenario(scenario_id, provider_response) + + def policy(request: Request) -> Reply: + payload: Final = JSON_OBJECT.validate_json(request.body) + direction: Final = _direction(payload) + verdict: Final = ( + {"action": "BLOCKED", "blocked_reason": "synthetic logging-only denial"} + if direction in block_directions + else {"action": "NONE"} + ) + return Reply(body=json.dumps(verdict).encode()) + + try: + with gateway.scenario() as scenario: + api_base: Final = ( + upstream_handle.api_base() if endpoint == "messages" else f"{upstream_handle.api_base()}/v1" + ) + model: Final = scenario.model( + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=api_base, + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, include_scope=include_scope) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = _call_client( + client_kind, endpoint, gateway, model, prompt, stream, baseline_call_id + ) + baseline_upstream: Final = _drain_upstream(gateway.upstream_url) + assert baseline.status == 200, baseline.body + assert baseline.text == reply, baseline + assert len(baseline_upstream) == 1, baseline_upstream + result: Final = _call_client( + client_kind, endpoint, candidate, model, prompt, stream, guarded_call_id + ) + assert (result.status, _response_body_without_ids(result.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (result, baseline) + assert result.text == reply, result + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 1, observed_upstream + assert prompt in json.dumps(observed_upstream[0]["body"]), observed_upstream + base_default_directions: Final = BASE_DEFAULT_NORMAL_DIRECTIONS[(endpoint, stream)] + expected_directions: Final = _directions_for_audit_leg(base_default_directions, scope) + directions_to_collect: Final = _directions_for_audit_leg(base_default_directions, scope) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(directions_to_collect), + seconds=20, + ) + calls: Final = guardrail.drain() + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in calls) + observed_directions: Final = tuple(_direction(payload) for payload in payloads) + assert all(_policy_call_id_matches(payload, guarded_call_id) for payload in payloads), payloads + assert tuple(payload["texts"] for payload in payloads) == tuple( + [prompt] if direction == "request" else [reply] for direction in observed_directions + ), payloads + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id(endpoint, str(guarded_rows[0]["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_rows[0]) + observed_entries: Final = tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) + expected_entries: Final = tuple( + ( + identity, + "logging_only", + "guardrail_intervened" if direction in block_directions else "success", + ) + for direction in expected_directions + ) + assert ( + tuple(sorted(observed_directions)), + tuple(sorted(observed_entries)), + ) == (tuple(sorted(expected_directions)), tuple(sorted(expected_entries))), ( + payloads, + entries, + rows, + ) + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("endpoint", "stream", "client_kind"), + ( + pytest.param("chat", False, "openai_sync", id="A-default-chat-nonstream"), + pytest.param("chat", True, "openai_async", id="A-default-chat-stream"), + pytest.param("messages", False, "anthropic_sync", id="A-default-messages-nonstream"), + pytest.param("messages", True, "anthropic_async", id="A-default-messages-stream"), + pytest.param("responses", False, "openai_async", id="A-default-responses-nonstream"), + pytest.param("responses", True, "openai_sync", id="A-default-responses-stream"), + ), +) +def test_A_unset_scope_matches_measured_endpoint_stream_default( + gateway: Gateway, + tmp_path: Path, + endpoint: Endpoint, + stream: bool, + client_kind: ClientKind, +) -> None: + identity: Final = f"logging-scope-a-default-{endpoint}-{stream}-{uuid.uuid4().hex}" + scenario_id: Final = f"phase12-a-default-{endpoint}-{stream}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic unset-scope request {identity}" + reply: Final = f"synthetic unset-scope response {identity}" + upstream_handle: Final = register_scenario(scenario_id, _provider_response(endpoint, scenario_id, reply, stream)) + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + api_base: Final = ( + upstream_handle.api_base() if endpoint == "messages" else f"{upstream_handle.api_base()}/v1" + ) + model: Final = scenario.model( + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=api_base, + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, None, include_scope=False) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = _call_client( + client_kind, + endpoint, + gateway, + model, + prompt, + stream, + f"{scenario_id}-baseline", + ) + assert baseline.status == 200 and baseline.text == reply, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + result: Final = _call_client( + client_kind, + endpoint, + candidate, + model, + prompt, + stream, + f"{scenario_id}-candidate", + ) + assert (result.status, _response_body_without_ids(result.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (result, baseline) + assert result.text == reply, result + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + expected_directions: Final = BASE_DEFAULT_NORMAL_DIRECTIONS[(endpoint, stream)] + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(payload) for payload in payloads)) == tuple( + sorted(expected_directions) + ), payloads + assert all(_policy_call_id_matches(payload, f"{scenario_id}-candidate") for payload in payloads), ( + payloads + ) + assert tuple(payload["texts"] for payload in payloads) == tuple( + [prompt] if _direction(payload) == "request" else [reply] for payload in payloads + ), payloads + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id( + endpoint, + str(guarded_rows[0]["request_id"]), + result.response_id, + scenario_id, + ) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize("inventory_id", (pytest.param("B4", id="B4-monitor-usage-detail"),)) +def test_logging_only_monitor_counts_only_the_observed_direction( + gateway: Gateway, tmp_path: Path, inventory_id: str +) -> None: + input_identity: Final = f"logging-scope-b4-input-{uuid.uuid4().hex}" + output_identity: Final = f"logging-scope-b4-output-{uuid.uuid4().hex}" + scenario_id: Final = f"phase12-{inventory_id.lower()}-{uuid.uuid4().hex}" + prompt_by_identity: Final = { + input_identity: f"synthetic input {input_identity}", + output_identity: f"synthetic input {output_identity}", + } + provider_reply: Final = f"synthetic response {scenario_id}" + upstream_handle: Final = register_scenario( + scenario_id, _provider_response("chat", scenario_id, provider_reply, False) + ) + + def policy(request: Request) -> Reply: + return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic monitor denial"}).encode()) + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + with wire_server(policy) as input_guardrail, wire_server(policy) as output_guardrail: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [ + { + "guardrail_name": input_identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "logging_only_scope": "input", + "default_on": False, + "api_base": input_guardrail.url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + }, + }, + { + "guardrail_name": output_identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "logging_only", + "logging_only_scope": "output", + "default_on": False, + "api_base": output_guardrail.url, + "api_key": "synthetic-guardrail-key", + "extra_headers": ["x-litellm-call-id"], + }, + }, + ] + config_path: Final = tmp_path / "b4.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=config_path, workers=1) as candidate: + cases: Final = tuple( + ( + identity, + gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt_by_identity[identity]}], + }, + headers={"x-litellm-call-id": f"{scenario_id}-{identity}-baseline"}, + ), + candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt_by_identity[identity]}], + "guardrails": [identity], + }, + headers={"x-litellm-call-id": f"{scenario_id}-{identity}"}, + ), + ) + for identity in (input_identity, output_identity) + ) + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 4, observed_upstream + for identity, baseline_response, guarded_response in cases: + assert baseline_response.status_code == guarded_response.status_code == 200, ( + baseline_response.text, + guarded_response.text, + ) + assert _response_body_without_ids(JSON_OBJECT.validate_python(baseline_response.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(guarded_response.json())) + ), ( + baseline_response.text, + guarded_response.text, + ) + assert baseline_response.json()["choices"][0]["message"]["content"] == provider_reply + assert ( + sum( + prompt_by_identity[identity] in json.dumps(observation["body"]) + for observation in observed_upstream + ) + == 2 + ), observed_upstream + expected_call_ids: Final = tuple( + f"{scenario_id}-{identity}" for identity in (input_identity, output_identity) + ) + eventually( + lambda: input_guardrail.received.qsize(), + lambda count: count >= len(cases), + seconds=20, + ) + eventually( + lambda: output_guardrail.received.qsize(), + lambda count: count >= len(cases), + seconds=20, + ) + policy_payloads: Final = { + input_identity: tuple(JSON_OBJECT.validate_json(call.body) for call in input_guardrail.drain()), + output_identity: tuple( + JSON_OBJECT.validate_json(call.body) for call in output_guardrail.drain() + ), + } + for identity in (input_identity, output_identity): + expected_direction: Final = ( + "request" if _is_base_audit_leg() or identity == input_identity else "response" + ) + assert len(policy_payloads[identity]) == len(expected_call_ids), policy_payloads[identity] + for case_identity in (input_identity, output_identity): + call_id: Final = f"{scenario_id}-{case_identity}" + calls_for_id: Final = tuple( + payload + for payload in policy_payloads[identity] + if _policy_call_id_matches(payload, call_id) + ) + call_summary: Final = tuple( + ( + _direction(payload), + payload.get("litellm_call_id"), + _policy_call_id(payload), + tuple(payload["texts"]), + ) + for payload in calls_for_id + ) + assert tuple(_direction(payload) for payload in calls_for_id) == (expected_direction,), ( + call_id, + call_summary, + ) + expected_text: Final = ( + prompt_by_identity[case_identity] if expected_direction == "request" else provider_reply + ) + assert tuple(tuple(payload["texts"]) for payload in calls_for_id) == ((expected_text,),), ( + call_id, + call_summary, + ) + rows: Final = _spend_rows_for_calls( + (model,), + tuple( + (model, str(response.json()["id"]), f"{scenario_id}-{identity}") + for identity, _baseline_response, response in cases + ), + ) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 2, rows + expected_entries: Final = tuple( + sorted( + ( + (input_identity, "logging_only", "guardrail_intervened"), + (output_identity, "logging_only", "guardrail_intervened"), + ) + ) + ) + for identity, _baseline_response, response in cases: + call_id: Final = f"{scenario_id}-{identity}" + matching_rows: Final = _spend_rows_matching_call(rows, model, call_id) + assert len(matching_rows) == 1, (call_id, rows) + _assert_response_id( + "chat", + str(matching_rows[0]["request_id"]), + str(response.json()["id"]), + scenario_id, + ) + entries: Final = _guardrail_entries(matching_rows[0]) + assert ( + tuple( + sorted( + ( + entry["guardrail_name"], + entry["guardrail_mode"], + entry["guardrail_status"], + ) + for entry in entries + ) + ) + == expected_entries + ), entries + listed: Final = candidate.get("/v2/guardrails/list")["guardrails"] + assert isinstance(listed, list), listed + guardrail_ids: Final = { + object_value(row)["guardrail_name"]: str(object_value(row)["guardrail_id"]) + for row in listed + if object_value(row)["guardrail_name"] in (input_identity, output_identity) + } + assert set(guardrail_ids) == {input_identity, output_identity}, listed + today: Final = datetime.now(timezone.utc).date().isoformat() + for identity in (input_identity, output_identity): + detail: Final = eventually( + lambda identity=identity: candidate.request( + "GET", + f"/guardrails/usage/detail/{guardrail_ids[identity]}", + params={"start_date": today, "end_date": today}, + ).json(), + lambda body: body["requestsEvaluated"] >= 1, + seconds=30, + return_last_on_timeout=True, + ) + assert detail["requestsEvaluated"] == len(cases), detail + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "endpoint", "scope", "expected_direction"), + ( + pytest.param("C1", "chat", "output", "response", id="C1-chat-cache-output"), + pytest.param("C2", "chat", "input", "request", id="C2-chat-cache-input"), + pytest.param("C3", "messages", "output", "response", id="C3-messages-cache-output"), + pytest.param("C4", "responses", "output", "response", id="C4-responses-cache-output"), + ), +) +def test_cache_hit_directional_scope_uses_measured_base_default( + gateway: Gateway, + tmp_path: Path, + row_id: str, + endpoint: Endpoint, + scope: str, + expected_direction: Direction, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + baseline_prompt: Final = f"uncached control {identity}" + cached_prompt: Final = f"repeated cache prompt {identity}" + reply: Final = f"cache response {identity}" + base_default_on_hit: Final = BASE_DEFAULT_CACHE_HIT_DIRECTIONS[endpoint] + expected_miss_directions: Final = _directions_for_audit_leg( + BASE_DEFAULT_NORMAL_DIRECTIONS[(endpoint, False)], scope + ) + expected_hit_directions: Final = _directions_for_audit_leg(base_default_on_hit, scope) + if _is_base_audit_leg(): + assert expected_hit_directions == base_default_on_hit, (base_default_on_hit, expected_hit_directions) + upstream_handle: Final = register_scenario(scenario_id, _provider_response(endpoint, scenario_id, reply, False)) + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + api_base: Final = ( + upstream_handle.api_base() if endpoint == "messages" else f"{upstream_handle.api_base()}/v1" + ) + model: Final = scenario.model( + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=api_base, + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, cache=True) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = _call_cache_client( + endpoint, gateway, model, baseline_prompt, f"{scenario_id}-baseline" + ) + assert baseline.status == 200, baseline.body + assert baseline.text == reply, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + first: Final = _call_cache_client(endpoint, candidate, model, cached_prompt, f"{scenario_id}-first") + assert (first.status, _response_body_without_ids(first.body)) == ( + baseline.status, + _response_body_without_ids(baseline.body), + ), (first, baseline) + assert first.text == reply, first + first_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(first_upstream) == 1, first_upstream + assert cached_prompt in json.dumps(first_upstream[0]["body"]), first_upstream + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(expected_miss_directions), + seconds=20, + ) + miss_payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert all(_policy_call_id_matches(payload, f"{scenario_id}-first") for payload in miss_payloads), ( + miss_payloads + ) + second: Final = _call_cache_client( + endpoint, candidate, model, cached_prompt, f"{scenario_id}-second" + ) + assert (second.status, _response_body_without_ids(second.body)) == ( + first.status, + _response_body_without_ids(first.body), + ), (second, first) + assert _drain_upstream(gateway.upstream_url) == (), "The identical second request must hit Redis" + rows: Final = _spend_rows(model, 3) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(expected_hit_directions), + seconds=20, + ) + hit_payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + for response_id, policy_payloads, expected_directions in ( + (first.response_id, miss_payloads, expected_miss_directions), + (second.response_id, hit_payloads, expected_hit_directions), + ): + assert tuple(sorted(_direction(payload) for payload in policy_payloads)) == tuple( + sorted(expected_directions) + ), (response_id, policy_payloads) + assert all( + payload["texts"] == ([cached_prompt] if _direction(payload) == "request" else [reply]) + for payload in policy_payloads + ), (response_id, policy_payloads) + assert sum(_cache_hit(row["cache_hit"]) for row in rows) == 1, rows + guarded_rows: Final = tuple( + row for row in rows if identity in object_value(row["metadata"]).get("applied_guardrails", []) + ) + assert len(guarded_rows) == 2, rows + miss_rows: Final = tuple(row for row in guarded_rows if not _cache_hit(row["cache_hit"])) + hit_rows: Final = tuple(row for row in guarded_rows if _cache_hit(row["cache_hit"])) + assert len(miss_rows) == len(hit_rows) == 1, rows + _assert_response_id(endpoint, str(miss_rows[0]["request_id"]), first.response_id, scenario_id) + assert str(hit_rows[0]["request_id"]).startswith(f"{second.response_id}_cache_hit"), hit_rows + for row, expected_directions in ( + (miss_rows[0], expected_miss_directions), + (hit_rows[0], expected_hit_directions), + ): + entries: Final = _guardrail_entries(row) + assert len(entries) == len(expected_directions), entries + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope", "modes", "expected_directions", "expected_statuses", "expected_status"), + ( + pytest.param( + "D1", + "output", + ["pre_call", "logging_only"], + ("request",), + ("guardrail_intervened",), + 400, + id="D1-pre-call-blocks-before-upstream", + ), + pytest.param( + "D2", + "output", + ["pre_call", "logging_only"], + ("request", "response"), + ("success", "guardrail_intervened"), + 200, + id="D2-pre-call-and-output-observation", + ), + pytest.param( + "D3", + "input", + ["logging_only", "post_call"], + ("response",), + ("guardrail_intervened",), + 400, + id="D3-post-call-block-remains-enforced", + ), + ), +) +def test_combined_modes_preserve_blocking_and_directional_observation( + gateway: Gateway, + tmp_path: Path, + row_id: str, + scope: str, + modes: list[str], + expected_directions: tuple[Direction, ...], + expected_statuses: tuple[str, ...], + expected_status: int, +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic combined-mode prompt {identity}" + reply: Final = f"synthetic combined-mode response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + payload: Final = JSON_OBJECT.validate_json(request.body) + direction: Final = _direction(payload) + blocked: Final = row_id == "D1" or direction == "response" + verdict: Final = ( + {"action": "BLOCKED", "blocked_reason": f"synthetic denial from {identity}"} + if blocked + else {"action": "NONE"} + ) + return Reply(body=json.dumps(verdict).encode()) + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, mode=modes) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert baseline.json()["choices"][0]["message"]["content"] == reply, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + guarded: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": call_id}, + ) + leg_expected_directions: Final = ( + ("request", "request", "response") + if _is_base_audit_leg() and row_id == "D2" + else expected_directions + ) + leg_expected_statuses: Final = ( + ("success", "success", "guardrail_intervened") + if _is_base_audit_leg() and row_id == "D2" + else expected_statuses + ) + assert guarded.status_code == expected_status, guarded.text + candidate_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(candidate_upstream) == (0 if row_id == "D1" else 1), candidate_upstream + if candidate_upstream: + assert prompt in json.dumps(candidate_upstream[0]["body"]), candidate_upstream + if row_id == "D2": + assert _response_body_without_ids(JSON_OBJECT.validate_python(guarded.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(baseline.json())) + ), (guarded.text, baseline.text) + else: + assert identity in guarded.text, guarded.text + assert f"synthetic denial from {identity}" in guarded.text, guarded.text + expected_policy_call_count: Final = len(leg_expected_directions) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= expected_policy_call_count, + seconds=20, + ) + calls: Final = guardrail.drain() + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in calls) + assert all(_policy_call_id_matches(payload, call_id) for payload in payloads), payloads + observed_directions: Final = tuple(_direction(payload) for payload in payloads) + assert tuple(payload["texts"] for payload in payloads) == tuple( + [prompt] if direction == "request" else [reply] for direction in observed_directions + ), payloads + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + entries: Final = _guardrail_entries(guarded_rows[0]) + observed_modes: Final = _guardrail_mode_status_pairs(entries) + assert all(mode_values == tuple(modes) for mode_values, _ in observed_modes), entries + observed_statuses: Final = tuple(status for _, status in observed_modes) + assert ( + tuple(sorted(observed_directions)), + tuple(sorted(observed_statuses)), + ) == (tuple(sorted(leg_expected_directions)), tuple(sorted(leg_expected_statuses))), ( + payloads, + entries, + rows, + ) + if row_id == "D2": + _assert_response_id( + "chat", + str(guarded_rows[0]["request_id"]), + str(guarded.json()["id"]), + scenario_id, + ) + else: + assert _guardrail_entries(_spend_row_for_call_id(call_id)) == entries, entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("row_id", "scope", "selection"), + ( + pytest.param("E1", "output", "request", id="E1-selected-by-request-body"), + pytest.param("E2", "input", "virtual-key", id="E2-selected-by-key-metadata"), + pytest.param("E3", "output", "unselected", id="E3-no-request-or-key-selection"), + ), +) +def test_logging_only_scope_respects_guardrail_selection_level( + gateway: Gateway, tmp_path: Path, row_id: str, scope: str, selection: str +) -> None: + identity: Final = f"logging-scope-{row_id.lower()}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic selected request {identity}" + reply: Final = f"synthetic selected response {identity}" + scenario_id: Final = f"phase12-{row_id.lower()}-{uuid.uuid4().hex}" + call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + key: Final = ( + scenario.key(metadata={"guardrails": [identity]}) if selection == "virtual-key" else gateway.key + ) + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope, default_on=False) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + request_body: Final[dict[str, JsonValue]] = { + "model": model, + "messages": [{"role": "user", "content": prompt}], + **({"guardrails": [identity]} if selection == "request" else {}), + } + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + key=key, + headers={"x-litellm-call-id": f"{scenario_id}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert baseline.json()["choices"][0]["message"]["content"] == reply, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + guarded: Final = candidate.request( + "POST", + "/v1/chat/completions", + request_body, + key=key, + headers={"x-litellm-call-id": call_id}, + ) + assert guarded.status_code == 200, guarded.text + assert _response_body_without_ids(JSON_OBJECT.validate_python(guarded.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(baseline.json())) + ), (guarded.text, baseline.text) + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 1, observed_upstream + assert prompt in json.dumps(observed_upstream[0]["body"]), observed_upstream + expected_directions: Final = _directions_for_audit_leg(("request", "response"), scope) + directions_to_collect: Final = expected_directions + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(directions_to_collect), + seconds=20, + ) + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + if expected_directions: + assert all(_policy_call_id_matches(payload, call_id) for payload in payloads), payloads + assert tuple(sorted(_direction(payload) for payload in payloads)) == tuple( + sorted(expected_directions) + ), payloads + if expected_directions: + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + row: Final = guarded_rows[0] + _assert_response_id("chat", str(row["request_id"]), str(guarded.json()["id"]), scenario_id) + entries: Final = _guardrail_entries(row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + else: + row: Final = _spend_row_for_response_id(str(guarded.json()["id"])) + assert all(entry["guardrail_name"] != identity for entry in _guardrail_entries(row)), row + finally: + delete_scenario(upstream_handle) + + +def test_X1_missing_null_and_both_scope_have_identical_scans(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-x1-{uuid.uuid4().hex}" + variants: Final = ( + ("both", "both", True), + ("missing", None, False), + ("null", None, True), + ) + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + model: Final = scenario.model() + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": identity}]}, + headers={"x-litellm-call-id": f"{identity}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + baseline_body: Final = JSON_OBJECT.validate_json(baseline.content) + baseline_reply: Final = _response_text("chat", baseline_body) + assert len(_drain_upstream(gateway.upstream_url)) == 1 + for suffix, scope, include_scope in variants: + name: Final = f"{identity}-{suffix}" + call_id: Final = f"{identity}-{suffix}" + with wire_server(policy) as edge: + config: Final = _configuration( + tmp_path, + name, + edge.url, + scope, + include_scope=include_scope, + default_on=False, + ) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": identity}], + "guardrails": [name], + }, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == baseline.status_code, response.text + body: Final = JSON_OBJECT.validate_json(response.content) + assert body["choices"] == baseline_body["choices"], response.text + assert _response_text("chat", body) == baseline_reply, response.text + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and identity in json.dumps(upstream[0]["body"]), upstream + expected_directions: Final = _directions_for_audit_leg(("request", "response"), scope) + eventually( + lambda: edge.received.qsize(), + lambda count, expected_directions=expected_directions: count == len(expected_directions), + seconds=20, + ) + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain()) + assert tuple(sorted(_direction(payload) for payload in payloads)) == tuple( + sorted(expected_directions) + ), payloads + assert all(_policy_call_id_matches(payload, call_id) for payload in payloads), payloads + assert tuple(sorted((_direction(payload), tuple(payload["texts"])) for payload in payloads)) == ( + ("request", (identity,)), + ("response", (baseline_reply,)), + ), payloads + row: Final = _spend_row_for_response_id(str(body["id"])) + entries: Final = _guardrail_entries(row) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((name, "logging_only", "success") for _ in expected_directions), entries + + +def test_X3_five_identical_requests_each_receive_one_response_scan(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-x3-{uuid.uuid4().hex}" + prompt: Final = f"synthetic repeated request {identity}" + + def policy(_request: Request) -> Reply: + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + model: Final = scenario.model() + baseline: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{identity}-baseline"}, + ) + assert baseline.status_code == 200, baseline.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as edge: + config: Final = _configuration(tmp_path, identity, edge.url, "output") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + expected_directions: Final = _directions_for_audit_leg(("request", "response"), "output") + results: Final = tuple( + candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + headers={"x-litellm-call-id": f"{identity}-{index}"}, + ) + for index in range(5) + ) + assert all(result.status_code == baseline.status_code for result in results), results + assert all(result.json()["choices"] == baseline.json()["choices"] for result in results), results + response_ids: Final = tuple(str(result.json()["id"]) for result in results) + assert len(set(response_ids)) == 5, response_ids + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 5, upstream + eventually( + lambda: edge.received.qsize(), + lambda count: count == 5 * len(expected_directions), + seconds=20, + ) + calls: Final = edge.drain() + payloads: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in calls) + for index, result in enumerate(results): + call_id: Final = f"{identity}-{index}" + matching_payloads: Final = tuple( + payload for payload in payloads if _policy_call_id_matches(payload, call_id) + ) + assert tuple(sorted(_direction(payload) for payload in matching_payloads)) == tuple( + sorted(expected_directions) + ), ( + call_id, + matching_payloads, + ) + expected_reply: Final = _response_text("chat", JSON_OBJECT.validate_json(result.content)) + assert all( + payload["texts"] == ([prompt] if _direction(payload) == "request" else [expected_reply]) + for payload in matching_payloads + ), matching_payloads + rows: Final = _spend_rows(model, 6) + for index, result in enumerate(results): + matching_rows: Final = tuple(row for row in rows if row["request_id"] == result.json()["id"]) + assert len(matching_rows) == 1, (result.json()["id"], matching_rows) + _assert_response_id( + "chat", + str(matching_rows[0]["request_id"]), + str(result.json()["id"]), + f"{identity}-{index}", + ) + entries: Final = _guardrail_entries(matching_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "success") for _ in expected_directions), entries + rows: Final = _spend_rows(model, 6) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 5, rows + assert {str(row["request_id"]) for row in guarded_rows} == set(response_ids), guarded_rows + assert all( + tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in _guardrail_entries(row) + ) + == tuple((identity, "logging_only", "success") for _ in expected_directions) + for row in guarded_rows + ), guarded_rows + + +def test_X2_scope_patch_toggles_during_concurrent_requests_keep_one_scan_per_response( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = f"logging-scope-x2-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + scan_started: Final = threading.Event() + release_scans: Final = threading.Event() + + def policy(_request: Request) -> Reply: + scan_started.set() + assert release_scans.wait(timeout=60), identity + return Reply(body=b'{"action":"NONE"}') + + with gateway.scenario() as scenario: + model: Final = scenario.model() + baseline: Final = _call_client( + "openai_sync", "chat", gateway, model, f"synthetic X2 control {marker}", False, f"{marker}-control" + ) + assert baseline.status == 200, baseline + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as edge: + config: Final = _empty_proxy_configuration(tmp_path, identity) + with ( + _database_guardrail(identity, edge.url, "output", mode="logging_only"), + owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate, + ): + guardrails: Final = candidate.get("/v2/guardrails/list")["guardrails"] + guardrail_id: Final = next( + string_value(object_value(guardrail)["guardrail_id"]) + for guardrail in guardrails + if object_value(guardrail)["guardrail_name"] == identity + ) + calls: Final = tuple( + ChaosCall( + index=index, + endpoint="chat", + client_kind="openai_sync", + model=model, + stream=False, + prompt=f"synthetic X2 request {marker}-{index}", + call_id=f"{marker}-x2-{index}", + ) + for index in range(20) + ) + expected_scan_count: Final = len(calls) * (2 if _is_base_audit_leg() else 1) + with ThreadPoolExecutor(max_workers=20) as pool: + futures: Final = tuple( + pool.submit( + _call_client, + call.client_kind, + call.endpoint, + candidate, + call.model, + call.prompt, + call.stream, + call.call_id, + ) + for call in calls + ) + try: + assert eventually(lambda: scan_started.is_set(), bool, seconds=30) + patch_responses: Final = tuple( + candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"logging_only_scope": "input" if index % 2 == 0 else "output"}}, + ) + for index in range(10) + ) + assert all(response.status_code == 200 for response in patch_responses), patch_responses + finally: + release_scans.set() + results: Final = tuple(future.result(timeout=90) for future in futures) + assert all(result.status == baseline.status for result in results), results + assert all(result.text == baseline.text for result in results), results + response_ids: Final = tuple(result.response_id for result in results) + assert len(set(response_ids)) == 20, response_ids + eventually( + lambda: edge.received.qsize(), + lambda count: count == expected_scan_count, + seconds=30, + ) + edge_calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in edge.drain()) + for call in calls: + payloads_for_call: Final = tuple( + payload for payload in edge_calls if _policy_call_id_matches(payload, call.call_id) + ) + directions: Final = tuple(_direction(payload) for payload in payloads_for_call) + if _is_base_audit_leg(): + assert tuple(sorted(directions)) == ("request", "response"), (call, payloads_for_call) + else: + assert len(directions) == 1 and directions[0] in ("request", "response"), ( + call, + payloads_for_call, + ) + assert all( + payload["texts"] == ([call.prompt] if _direction(payload) == "request" else [baseline.text]) + for payload in payloads_for_call + ), payloads_for_call + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 20, upstream + assert ( + tuple( + sum(_json_contains_exact_string(observation["body"], call.prompt) for observation in upstream) + for call in calls + ) + == (1,) * 20 + ), upstream + rows: Final = _spend_rows(model, 21) + for call, result in zip(calls, results): + row: Final = next(row for row in rows if row["request_id"] == result.response_id) + entries: Final = _guardrail_entries(row) + payloads_for_call: Final = tuple( + payload for payload in edge_calls if _policy_call_id_matches(payload, call.call_id) + ) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple( + (identity, "logging_only", "success") + for _ in tuple(_direction(payload) for payload in payloads_for_call) + ), (call.call_id, entries) + + +def test_S1_logging_only_output_scope_fails_open_on_policy_500(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-s1-{uuid.uuid4().hex}" + prompt: Final = f"synthetic policy outage prompt {identity}" + reply: Final = f"synthetic policy outage response {identity}" + scenario_id: Final = f"phase12-s1-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(status=500, body=b'{"error":"synthetic policy outage"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + control: Final = _call_client("openai_sync", "chat", gateway, model, prompt, False, control_call_id) + assert control.status == 200 and control.text == reply, control + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, "output") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = _call_client( + "openai_sync", "chat", candidate, model, prompt, False, guarded_call_id + ) + assert (result.status, _response_body_without_ids(result.body)) == ( + control.status, + _response_body_without_ids(control.body), + ), (result, control) + expected_directions: Final = _directions_for_scope(("request", "response"), "output") + directions_to_collect: Final = _directions_for_audit_leg(("request", "response"), "output") + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + eventually( + lambda: guardrail.received.qsize(), + lambda count: count >= len(directions_to_collect), + seconds=20, + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert all(_policy_call_id_matches(call, guarded_call_id) for call in calls), calls + assert tuple(call["texts"] for call in calls) == tuple( + [prompt] if _direction(call) == "request" else [reply] for call in calls + ), calls + guarded_row: Final = _spend_row_for_response_id(result.response_id) + _assert_response_id("chat", str(guarded_row["request_id"]), result.response_id, scenario_id) + entries: Final = _guardrail_entries(guarded_row) + observed_directions: Final = tuple(_direction(call) for call in calls) + observed_entries: Final = tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) + expected_entries: Final = tuple( + (identity, "logging_only", "guardrail_failed_to_respond") for _ in expected_directions + ) + assert ( + tuple(sorted(observed_directions)), + tuple(sorted(observed_entries)), + ) == (tuple(sorted(expected_directions)), tuple(sorted(expected_entries))), ( + calls, + entries, + guarded_row, + ) + finally: + delete_scenario(upstream_handle) + + +def test_S2_logging_only_output_scope_scans_both_chat_choices(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-s2-{uuid.uuid4().hex}" + prompt: Final = f"synthetic multiple choice prompt {identity}" + first_reply: Final = f"synthetic first choice {identity}" + second_reply: Final = f"synthetic second choice {identity}" + scenario_id: Final = f"phase12-s2-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + provider_response: Final = JsonResponse( + content_type="application/json", + body={ + "id": "chatcmpl-$UNIQUE_ID", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": first_reply}, "finish_reason": "stop"}, + {"index": 1, "message": {"role": "assistant", "content": second_reply}, "finish_reason": "stop"}, + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 10, "total_tokens": 19}, + }, + ) + upstream_handle: Final = register_scenario(scenario_id, provider_response) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic multiple choice monitor"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + control: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}], "n": 2}, + headers={"x-litellm-call-id": control_call_id}, + ) + assert control.status_code == 200, control.text + assert len(control.json()["choices"]) == 2, control.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, "output") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}], "n": 2}, + headers={"x-litellm-call-id": guarded_call_id}, + ) + assert result.status_code == control.status_code, result.text + assert _response_body_without_ids(JSON_OBJECT.validate_python(result.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(control.json())) + ), (result.text, control.text) + upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(upstream) == 1 and prompt in json.dumps(upstream[0]["body"]), upstream + expected_directions: Final = ("response",) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(call) for call in calls)) == tuple(sorted(expected_directions)), ( + calls + ) + assert all(_policy_call_id_matches(call, guarded_call_id) for call in calls), calls + assert tuple(sorted((_direction(call), tuple(call["texts"])) for call in calls)) == tuple( + sorted( + (direction, tuple([prompt] if direction == "request" else [first_reply, second_reply])) + for direction in expected_directions + ) + ), calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id( + "chat", + str(guarded_rows[0]["request_id"]), + str(result.json()["id"]), + scenario_id, + ) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "guardrail_intervened") for _ in expected_directions), entries + finally: + delete_scenario(upstream_handle) + + +@pytest.mark.parametrize( + ("endpoint", "scope"), + ( + pytest.param("chat", "input", id="S3-chat-upstream-401-input"), + pytest.param("chat", "output", id="S3-chat-upstream-401-output"), + pytest.param("messages", "input", id="S3-messages-upstream-401-input"), + pytest.param("messages", "output", id="S3-messages-upstream-401-output"), + pytest.param("responses", "input", id="S3-responses-upstream-401-input"), + pytest.param("responses", "output", id="S3-responses-upstream-401-output"), + ), +) +def test_logging_only_scope_on_upstream_401_uses_measured_base_default( + gateway: Gateway, + tmp_path: Path, + record_property: Callable[[str, object], None], + endpoint: Endpoint, + scope: str, +) -> None: + identity: Final = f"logging-scope-s3-{endpoint}-{scope}-{uuid.uuid4().hex}" + prompt: Final = f"synthetic upstream unauthorized marker {identity}" + scenario_id: Final = f"phase12-s3-{endpoint}-{scope}-{uuid.uuid4().hex}" + base_default_on_failure: Final = BASE_DEFAULT_UPSTREAM_FAILURE_DIRECTIONS[endpoint] + expected_directions: Final = _directions_for_audit_leg(base_default_on_failure, scope) + provider_response: Final = JsonResponse( + content_type="application/json", + body={ + "error": { + "message": f"synthetic upstream unauthorized {identity}", + "type": "invalid_request_error", + "code": "401", + } + }, + status=401, + ) + upstream_handle: Final = register_scenario(scenario_id, provider_response) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"NONE"}') + + try: + with gateway.scenario() as scenario: + api_base: Final = ( + upstream_handle.api_base() if endpoint == "messages" else f"{upstream_handle.api_base()}/v1" + ) + model: Final = scenario.model( + model={ + "chat": "openai/gpt-4o-mini", + "messages": "anthropic/claude-3-7-sonnet-20250219", + "responses": "openai/gpt-4.1-mini", + }[endpoint], + api_base=api_base, + api_key="synthetic-provider-key", + ) + path: Final = { + "chat": "/v1/chat/completions", + "messages": "/v1/messages", + "responses": "/v1/responses", + }[endpoint] + request_body: Final = { + "chat": {"model": model, "messages": [{"role": "user", "content": prompt}]}, + "messages": { + "model": model, + "max_tokens": 1000, + "messages": [{"role": "user", "content": prompt}], + }, + "responses": {"model": model, "input": prompt}, + }[endpoint] + control_call_id: Final = f"{scenario_id}-control" + candidate_call_id: Final = f"{scenario_id}-candidate" + control: Final = gateway.client.request( + "POST", + path, + json=request_body, + headers={ + "Authorization": f"Bearer {gateway.key}", + "x-litellm-call-id": control_call_id, + }, + timeout=60, + ) + control_upstream: Final = _drain_upstream(gateway.upstream_url) + control_upstream_count: Final = len(control_upstream) + record_property("s3_control_upstream_request_count", control_upstream_count) + assert control_upstream_count >= 1, control_upstream + assert control.status_code >= 400, control.text + assert all(prompt in json.dumps(observation["body"]) for observation in control_upstream), control_upstream + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, scope) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = candidate.client.request( + "POST", + path, + json=request_body, + headers={ + "Authorization": f"Bearer {candidate.key}", + "x-litellm-call-id": candidate_call_id, + }, + timeout=60, + ) + upstream: Final = _drain_upstream(gateway.upstream_url) + upstream_count: Final = len(upstream) + record_property("s3_candidate_upstream_request_count", upstream_count) + assert upstream_count >= 1, upstream + assert result.status_code >= 400, result.text + assert result.status_code == control.status_code, result.text + assert _response_body_without_ids(JSON_OBJECT.validate_python(result.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(control.json())) + ), (result.text, control.text) + assert all(prompt in json.dumps(observation["body"]) for observation in upstream), upstream + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert all(_policy_call_id_matches(call, candidate_call_id) for call in calls), calls + assert all(call["texts"] == [prompt] for call in calls if _direction(call) == "request"), calls + spend_rows: Final = _spend_rows(model, 2) + assert all(not _guardrail_entries(row) for row in spend_rows), spend_rows + assert {str(row["request_id"]) for row in spend_rows} >= { + control_call_id, + candidate_call_id, + }, spend_rows + assert tuple(sorted(_direction(call) for call in calls)) == tuple(sorted(expected_directions)), ( + calls, + spend_rows, + ) + finally: + delete_scenario(upstream_handle) + + +def test_S4_logging_only_input_scope_scans_every_multipart_text_part(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = f"logging-scope-s4-{uuid.uuid4().hex}" + first_part: Final = f"synthetic first text part {identity}" + second_part: Final = f"synthetic second text part {identity}" + reply: Final = f"synthetic multipart response {identity}" + scenario_id: Final = f"phase12-s4-{uuid.uuid4().hex}" + control_call_id: Final = f"{scenario_id}-control" + guarded_call_id: Final = f"{scenario_id}-guarded" + upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False)) + + def policy(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=b'{"action":"BLOCKED","blocked_reason":"synthetic multipart monitor"}') + + try: + with gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{upstream_handle.api_base()}/v1", + api_key="synthetic-provider-key", + ) + request_body: Final = { + "model": model, + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": first_part}, + {"type": "text", "text": second_part}, + ], + } + ], + } + control: Final = gateway.request( + "POST", + "/v1/chat/completions", + request_body, + headers={"x-litellm-call-id": control_call_id}, + ) + assert control.status_code == 200, control.text + assert control.json()["choices"][0]["message"]["content"] == reply, control.text + assert len(_drain_upstream(gateway.upstream_url)) == 1 + with wire_server(policy) as guardrail: + config: Final = _configuration(tmp_path, identity, guardrail.url, "input") + with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate: + result: Final = candidate.request( + "POST", + "/v1/chat/completions", + request_body, + headers={"x-litellm-call-id": guarded_call_id}, + ) + assert result.status_code == control.status_code, result.text + assert _response_body_without_ids(JSON_OBJECT.validate_python(result.json())) == ( + _response_body_without_ids(JSON_OBJECT.validate_python(control.json())) + ), (result.text, control.text) + observed_upstream: Final = _drain_upstream(gateway.upstream_url) + assert len(observed_upstream) == 1, observed_upstream + assert first_part in json.dumps(observed_upstream[0]["body"]), observed_upstream + assert second_part in json.dumps(observed_upstream[0]["body"]), observed_upstream + expected_directions: Final = ("request",) + eventually( + lambda: guardrail.received.qsize(), + lambda count: count == len(expected_directions), + seconds=20, + ) + calls: Final = tuple(JSON_OBJECT.validate_json(call.body) for call in guardrail.drain()) + assert tuple(sorted(_direction(call) for call in calls)) == tuple(sorted(expected_directions)), ( + calls + ) + assert all(_policy_call_id_matches(call, guarded_call_id) for call in calls), calls + assert tuple(sorted((_direction(call), tuple(call["texts"])) for call in calls)) == tuple( + sorted( + (direction, tuple([first_part, second_part] if direction == "request" else [reply])) + for direction in expected_directions + ) + ), calls + rows: Final = _spend_rows(model, 2) + guarded_rows: Final = tuple(row for row in rows if _guardrail_entries(row)) + assert len(guarded_rows) == 1, rows + _assert_response_id( + "chat", str(guarded_rows[0]["request_id"]), str(result.json()["id"]), scenario_id + ) + entries: Final = _guardrail_entries(guarded_rows[0]) + assert tuple( + (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) + for entry in entries + ) == tuple((identity, "logging_only", "guardrail_intervened") for _ in expected_directions), entries + finally: + delete_scenario(upstream_handle) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 4339febb0e3..d9d306602f2 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -1,5 +1,6 @@ import json import time +from collections.abc import Callable from datetime import datetime from typing import Dict, List, Optional from unittest.mock import AsyncMock @@ -10,6 +11,7 @@ import yaml from fastapi import HTTPException +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_endpoints import ( CreateGuardrailRequest, @@ -45,10 +47,12 @@ from litellm.proxy.guardrails.guardrail_registry import ( from litellm.types.guardrails import ( ApplyGuardrailRequest, BaseLitellmParams, + GuardrailEventHooks, Guardrail, GuardrailInfoResponse, LitellmParams, ) +from litellm.types.utils import GenericGuardrailAPIInputs # Mock data for testing MOCK_DB_GUARDRAIL = { @@ -88,6 +92,44 @@ MOCK_PATCH_REQUEST = PatchGuardrailRequest( ) +class _PatchScopeSupportedGuardrail(CustomGuardrail): + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: str, + logging_obj: object | None = None, + ) -> GenericGuardrailAPIInputs: + return inputs + + +class _PatchScopeUnsupportedGuardrail(_PatchScopeSupportedGuardrail): + async def async_logging_hook( + self, + kwargs: dict[str, object], + result: object, + call_type: str, + ) -> tuple[dict[str, object], object]: + return kwargs, result + + +def _patch_scope_initializer( + callback_type: type[CustomGuardrail], +) -> Callable[[LitellmParams, Guardrail], CustomGuardrail]: + def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail: + import litellm + + callback = callback_type( + guardrail_name=guardrail["guardrail_name"], + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + litellm.logging_callback_manager.add_litellm_callback(callback) + return callback + + return _initializer + + @pytest.fixture def mock_prisma_client(mocker): """Mock Prisma client for testing""" @@ -127,6 +169,37 @@ def mock_guardrail_registry(mocker): return mock_registry +def _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + callback_type: type[CustomGuardrail], + guardrail_type: str, + litellm_params: dict[str, object], +) -> tuple[InMemoryGuardrailHandler, Guardrail]: + from litellm.proxy.guardrails import guardrail_registry as registry_module + + guardrail: Guardrail = { + "guardrail_id": "patch-scope-test", + "guardrail_name": "Patch scope test", + "litellm_params": {"guardrail": guardrail_type, **litellm_params}, + "guardrail_info": {}, + } + mock_guardrail_registry.get_guardrail_by_id_from_db.return_value = guardrail + mock_guardrail_registry.update_guardrail_in_db.return_value = guardrail + monkeypatch.setitem( + registry_module.guardrail_initializer_registry, + guardrail_type, + _patch_scope_initializer(callback_type), + ) + handler = InMemoryGuardrailHandler() + handler.initialize_guardrail(guardrail=guardrail, source="db") + mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock()) + mocker.patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_guardrail_registry) + mocker.patch("litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", handler) + return handler, guardrail + + @pytest.mark.asyncio async def test_list_guardrails_v2_with_db_and_config(mocker, mock_prisma_client, mock_in_memory_handler): """Test listing guardrails from both DB and config""" @@ -1126,7 +1199,10 @@ async def test_update_guardrail_endpoint( prisma_client=mocker.ANY, ) - mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(guardrail=mocker.ANY) + mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with( + guardrail=mocker.ANY, + reject_invalid_logging_only_scope=True, + ) if scenario == "success_sync_fails_unexpected_error": assert mock_logger is not None @@ -1255,7 +1331,10 @@ async def test_patch_guardrail_endpoint( mock_guardrail_registry.update_guardrail_in_db.assert_called_once() - mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(guardrail=mocker.ANY) + mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with( + guardrail=mocker.ANY, + reject_invalid_logging_only_scope=False, + ) if scenario == "success_sync_fails_unexpected_error": assert mock_logger is not None @@ -1279,6 +1358,192 @@ async def test_patch_guardrail_rejects_mcp_only_on_violation_with_422(mocker, mo mock_guardrail_registry.update_guardrail_in_db.assert_not_called() +@pytest.mark.asyncio +async def test_patch_guardrail_rejects_invalid_logging_only_scope_with_422(mocker, mock_guardrail_registry): + mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock()) + mocker.patch( + "litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", + mock_guardrail_registry, + ) + mock_in_memory_handler = mocker.Mock(spec=InMemoryGuardrailHandler) + mock_in_memory_handler.sync_guardrail_from_db.side_effect = ValueError( + "Guardrail test-db-guardrail: logging_only_scope is set, but mode does not include logging_only" + ) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_in_memory_handler, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(mode="pre_call", logging_only_scope="input")) + + with pytest.raises(HTTPException) as exc_info: + await patch_guardrail("test-guardrail-id", request, user_api_key_dict=MOCK_ADMIN_USER) + + assert exc_info.value.status_code == 422 + assert "update rejected" in str(exc_info.value.detail) + mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with( + guardrail=mocker.ANY, + reject_invalid_logging_only_scope=True, + ) + + +@pytest.mark.asyncio +async def test_patch_guardrail_clears_scope_when_logging_only_mode_is_removed( + mocker, monkeypatch, mock_guardrail_registry +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeSupportedGuardrail, + "patch_scope_supported_test", + { + "mode": ["pre_call", "logging_only"], + "logging_only_scope": "output", + "default_on": True, + }, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(mode=["pre_call"])) + + try: + result = await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert result["guardrail_id"] == stored_guardrail["guardrail_id"] + persisted_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args.kwargs["guardrail"] + assert persisted_guardrail["litellm_params"].logging_only_scope is None + callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]] + assert callback.logging_only_scope is None + assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + +@pytest.mark.asyncio +async def test_patch_guardrail_tolerates_stored_unsupported_scope_on_unrelated_update( + mocker, monkeypatch, mock_guardrail_registry, caplog +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeUnsupportedGuardrail, + "patch_scope_unsupported_test", + {"mode": "logging_only", "logging_only_scope": "output", "default_on": True}, + ) + caplog.clear() + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(default_on=False)) + + try: + result = await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert result["guardrail_id"] == stored_guardrail["guardrail_id"] + callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]] + assert callback.logging_only_scope is None + assert any("Ignoring logging_only_scope" in record.getMessage() for record in caplog.records) + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + +@pytest.mark.asyncio +async def test_patch_guardrail_tolerates_invalid_stored_scope_on_unrelated_update( + mocker, monkeypatch, mock_guardrail_registry +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeUnsupportedGuardrail, + "patch_scope_invalid_literal_test", + {"mode": "logging_only", "logging_only_scope": "sideways", "default_on": True}, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(default_on=False)) + + try: + result = await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert result["guardrail_id"] == stored_guardrail["guardrail_id"] + persisted_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args.kwargs["guardrail"] + assert persisted_guardrail["litellm_params"].logging_only_scope is None + assert persisted_guardrail["litellm_params"].default_on is False + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + +@pytest.mark.asyncio +async def test_patch_guardrail_rejects_explicit_unsupported_scope_and_rolls_back( + mocker, monkeypatch, mock_guardrail_registry +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeUnsupportedGuardrail, + "patch_scope_unsupported_test", + {"mode": "logging_only", "logging_only_scope": "output", "default_on": True}, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(logging_only_scope="output")) + + try: + with pytest.raises(HTTPException) as exc_info: + await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert exc_info.value.status_code == 422 + assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2 + restored_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args_list[-1].kwargs["guardrail"] + assert restored_guardrail["litellm_params"] == stored_guardrail["litellm_params"] + callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]] + assert callback.logging_only_scope is None + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + +@pytest.mark.asyncio +async def test_patch_guardrail_rejected_update_restores_invalid_stored_scope_verbatim( + mocker, monkeypatch, mock_guardrail_registry +): + handler, stored_guardrail = _setup_patch_scope_guardrail( + mocker, + monkeypatch, + mock_guardrail_registry, + _PatchScopeUnsupportedGuardrail, + "patch_scope_invalid_rollback_test", + {"mode": "logging_only", "logging_only_scope": "sideways", "default_on": True}, + ) + request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(logging_only_scope="output")) + + try: + with pytest.raises(HTTPException) as exc_info: + await patch_guardrail( + stored_guardrail["guardrail_id"], + request, + user_api_key_dict=MOCK_ADMIN_USER, + ) + + assert exc_info.value.status_code == 422 + assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2 + restored_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args_list[-1].kwargs["guardrail"] + assert restored_guardrail["litellm_params"] == stored_guardrail["litellm_params"] + callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]] + assert callback.logging_only_scope is None + finally: + handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"]) + + @pytest.mark.parametrize( "scenario,expected_result,expected_exception", [ @@ -2464,6 +2729,13 @@ async def test_ui_settings_map_matches_runtime_supported_event_hooks(): from litellm.proxy.guardrails.guardrail_registry import guardrail_class_registry result = await get_guardrail_ui_settings() + expected_without_directional_scope = { + provider + for provider, guardrail_class in guardrail_class_registry.items() + if not guardrail_class.supports_logging_only_scope() + } + assert set(result.providers_without_directional_logging_only_scope) == expected_without_directional_scope + assert "xecguard" in result.providers_without_directional_logging_only_scope for provider, guardrail_class in guardrail_class_registry.items(): declared = guardrail_class.get_supported_event_hooks() diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 022fe85c779..152fd36e83e 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -1,15 +1,19 @@ from collections.abc import Iterable +from typing import ClassVar, Final from unittest.mock import AsyncMock, MagicMock import pytest +from pydantic import ValidationError from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy.guardrails.guardrail_registry import ( - get_guardrail_initializer_from_hooks, GuardrailRegistry, InMemoryGuardrailHandler, + get_guardrail_initializer_from_hooks, + parse_tolerant_litellm_params, ) -from litellm.types.guardrails import GuardrailEventHooks, Guardrail, LitellmParams +from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams, LoggingOnlyScope, Mode +from litellm.types.utils import GenericGuardrailAPIInputs def test_get_guardrail_initializer_from_hooks(): @@ -472,6 +476,24 @@ def test_unnormalizable_db_params_register_as_changed_without_raising(): assert handler._has_guardrail_params_changed(gid, new) is True +def test_invalid_scope_literal_db_params_compare_equal_after_normalization(): + handler = InMemoryGuardrailHandler() + raw = _db_litellm_params() + gid = "77777777-7777-7777-7777-777777777777" + handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params=LitellmParams(**{**raw, "logging_only_scope": None}), + ) + new = Guardrail( + guardrail_id=gid, + guardrail_name="cf", + litellm_params={**raw, "logging_only_scope": "Input"}, + ) + + assert handler._has_guardrail_params_changed(gid, new) is False + + def _all_callback_lists(): import litellm @@ -932,6 +954,230 @@ class TestScanOnlyToolResultsInitRefusal: ) +class _LoggingOnlyScopeSupportedGuardrail(CustomGuardrail): + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: str, + logging_obj: object | None = None, + ) -> GenericGuardrailAPIInputs: + return inputs + + +class _LoggingOnlyScopeUnsupportedGuardrail(_LoggingOnlyScopeSupportedGuardrail): + async def async_logging_hook( + self, + kwargs: dict[str, object], + result: object, + call_type: str, + ) -> tuple[dict[str, object], object]: + return kwargs, result + + +class _LoggingOnlyScopeNativeGuardrail(_LoggingOnlyScopeSupportedGuardrail): + use_native_lifecycle_hooks: ClassVar[bool] = True + + +def _invalid_scope_content_filter_guardrail() -> Guardrail: + return Guardrail( + guardrail_id="invalid-scope-content-filter-test", + guardrail_name="invalid-scope-content-filter", + litellm_params={ + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "logging_only_scope": "Input", + "blocked_words": [{"keyword": "pineapple", "action": "BLOCK"}], + }, + ) + + +class TestLoggingOnlyScopeValidation: + def _initialize( + self, + mode: str | list[str] | Mode, + scope: LoggingOnlyScope | None, + callback_type: type[CustomGuardrail] = _LoggingOnlyScopeSupportedGuardrail, + reject_invalid_logging_only_scope: bool = False, + assert_registered: bool = False, + ) -> CustomGuardrail: + import litellm + from litellm.proxy.guardrails import guardrail_registry as registry_module + + guardrail_type: Final = "logging_only_scope_test" + created_callbacks: Final[list[CustomGuardrail]] = [] + + def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail: + supported_event_hooks: Final = ( + [GuardrailEventHooks.logging_only] if callback_type.use_native_lifecycle_hooks else None + ) + callback: Final = callback_type( + guardrail_name=guardrail["guardrail_name"], + event_hook=litellm_params.mode, + default_on=True, + supported_event_hooks=supported_event_hooks, + ) + litellm.logging_callback_manager.add_litellm_callback(callback) + created_callbacks.append(callback) + return callback + + registry_module.guardrail_initializer_registry[guardrail_type] = _initializer + lists: Final = _all_callback_lists() + snapshots: Final = [list(callback_list) for callback_list in lists] + try: + handler: Final = InMemoryGuardrailHandler() + result: Final = handler.initialize_guardrail( + guardrail={ + "guardrail_name": "logging-only-scope-guardrail", + "litellm_params": { + "guardrail": guardrail_type, + "mode": mode, + "logging_only_scope": scope, + }, + }, + reject_invalid_logging_only_scope=reject_invalid_logging_only_scope, + ) + assert result is not None + callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]] + assert callback is not None + if assert_registered: + assert callback in lists[0] + return callback + except ValueError: + callback: Final = created_callbacks[0] + assert all(callback not in callback_list for callback_list in lists) + raise + finally: + for callback_list, snapshot in zip(lists, snapshots): + callback_list[:] = snapshot + registry_module.guardrail_initializer_registry.pop(guardrail_type, None) + + def test_scope_without_logging_only_mode_is_ignored_at_load(self) -> None: + callback: Final = self._initialize(mode="pre_call", scope="input", assert_registered=True) + + assert callback.logging_only_scope is None + assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True + + def test_scope_without_logging_only_mode_is_rejected_for_api_writes(self) -> None: + with pytest.raises(ValueError, match="logging_only_scope is set") as exc_info: + self._initialize(mode="pre_call", scope="input", reject_invalid_logging_only_scope=True) + + assert str(exc_info.value) == ( + "Guardrail logging-only-scope-guardrail: logging_only_scope is set, but mode does not include " + "logging_only, so it would never apply. Add logging_only to mode or remove logging_only_scope." + ) + + @pytest.mark.parametrize( + "mode", + ( + "logging_only", + ["pre_call", "logging_only"], + Mode(tags={"audit": "logging_only"}, default="pre_call"), + ), + ) + def test_scope_accepts_logging_only_in_supported_mode_forms(self, mode: str | list[str] | Mode) -> None: + callback: Final = self._initialize(mode=mode, scope="input") + + assert callback.logging_only_scope == "input" + + def test_directional_scope_is_ignored_at_load_when_guardrail_owns_logging_hook(self) -> None: + callback: Final = self._initialize( + mode="logging_only", + scope="input", + callback_type=_LoggingOnlyScopeUnsupportedGuardrail, + assert_registered=True, + ) + + assert callback.logging_only_scope is None + + def test_directional_scope_rejected_for_api_writes_when_guardrail_owns_logging_hook(self) -> None: + with pytest.raises(ValueError, match="logging_only_scope='input' is not supported") as exc_info: + self._initialize( + mode="logging_only", + scope="input", + callback_type=_LoggingOnlyScopeUnsupportedGuardrail, + reject_invalid_logging_only_scope=True, + ) + + assert str(exc_info.value) == ( + "Guardrail logging-only-scope-guardrail: logging_only_scope='input' is not supported by this " + "guardrail, whose logging_only hook scans on its own. Remove logging_only_scope." + ) + + def test_both_scope_accepted_when_guardrail_owns_logging_hook(self) -> None: + callback: Final = self._initialize( + mode="logging_only", + scope="both", + callback_type=_LoggingOnlyScopeUnsupportedGuardrail, + ) + + assert callback.logging_only_scope == "both" + + def test_output_scope_accepted_for_native_lifecycle_guardrail(self) -> None: + callback: Final = self._initialize( + mode="logging_only", + scope="output", + callback_type=_LoggingOnlyScopeNativeGuardrail, + ) + + assert callback.logging_only_scope == "output" + + def test_invalid_scope_fails_litellm_params_validation(self) -> None: + with pytest.raises(ValidationError): + LitellmParams(guardrail="test", mode="logging_only", logging_only_scope="request") + + def test_invalid_scope_literal_keeps_content_filter_registered_and_blocking(self) -> None: + import litellm + from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import ( + ContentFilterGuardrail, + ) + + handler: Final = InMemoryGuardrailHandler() + callback_lists: Final = _all_callback_lists() + callback_snapshots: Final = [list(callback_list) for callback_list in callback_lists] + guardrail: Final = _invalid_scope_content_filter_guardrail() + + try: + result: Final = handler.initialize_guardrail(guardrail=guardrail, source="config") + assert result is not None + callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]] + assert isinstance(callback, ContentFilterGuardrail) + assert callback in litellm.callbacks + assert callback.logging_only_scope is None + assert callback.event_hook == GuardrailEventHooks.pre_call + assert callback._check_blocked_words("pineapple") is not None + finally: + handler.delete_in_memory_guardrail(guardrail["guardrail_id"]) + for callback_list, snapshot in zip(callback_lists, callback_snapshots): + callback_list[:] = snapshot + + def test_invalid_scope_literal_is_rejected_for_strict_initialization_without_callback_leakage(self) -> None: + handler: Final = InMemoryGuardrailHandler() + callback_lists: Final = _all_callback_lists() + callback_snapshots: Final = [list(callback_list) for callback_list in callback_lists] + + with pytest.raises(ValueError, match="logging_only_scope"): + handler.initialize_guardrail( + guardrail=_invalid_scope_content_filter_guardrail(), + source="config", + reject_invalid_logging_only_scope=True, + ) + + assert all(callback_list == snapshot for callback_list, snapshot in zip(callback_lists, callback_snapshots)) + + def test_invalid_scope_literal_does_not_tolerate_other_litellm_params_errors(self) -> None: + with pytest.raises(ValidationError): + parse_tolerant_litellm_params( + { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "logging_only_scope": "Input", + "default_on": "not-a-bool", + }, + "invalid-scope-content-filter", + ) + + @pytest.mark.asyncio async def test_update_guardrail_in_db_raises_when_row_missing(): prisma_client = MagicMock() diff --git a/tests/unit/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py index 4649bddd281..052e3230e0a 100644 --- a/tests/unit/integrations/test_custom_guardrail.py +++ b/tests/unit/integrations/test_custom_guardrail.py @@ -12,7 +12,7 @@ from litellm.integrations.custom_guardrail import ( ) from litellm.litellm_core_utils.litellm_logging import Logging from litellm.proxy._types import CallTypes, UserAPIKeyAuth -from litellm.types.guardrails import GuardrailEventHooks, Mode +from litellm.types.guardrails import GuardrailEventHooks, LoggingOnlyScope, Mode from litellm.types.utils import ( Choices, GenericGuardrailAPIInputs, @@ -2577,6 +2577,32 @@ class TestLoggingOnlyApplyGuardrail: assert "standard_logging_guardrail_information" not in kwargs["litellm_params"]["metadata"] assert kwargs["standard_logging_object"] == {"guardrail_information": None} + @pytest.mark.parametrize( + "scope,expected_calls", + ( + (None, [("request", ["hello there"]), ("response", ["general kenobi"])]), + ("both", [("request", ["hello there"]), ("response", ["general kenobi"])]), + ("input", [("request", ["hello there"])]), + ("output", [("response", ["general kenobi"])]), + ), + ) + @pytest.mark.asyncio + async def test_logging_only_scope_scans_configured_directions( + self, + scope: LoggingOnlyScope | None, + expected_calls: list[tuple[str, list[str]]], + ) -> None: + guardrail: Final = _ApplyOnlyObserver() + guardrail.logging_only_scope = scope + kwargs, response = _logged_call([{"role": "user", "content": "hello there"}]) + + out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert guardrail.calls == expected_calls + entries: Final = out_kwargs["standard_logging_object"]["guardrail_information"] + assert len(entries) == len(expected_calls) + assert [entry["guardrail_mode"] for entry in entries] == ["logging_only"] * len(expected_calls) + @pytest.mark.asyncio async def test_appends_to_pre_call_verdicts_without_duplicating_them(self): guardrail = _ApplyOnlyObserver() @@ -2604,6 +2630,23 @@ class TestLoggingOnlyApplyGuardrail: assert out_kwargs is kwargs assert out_response is response + @pytest.mark.asyncio + async def test_output_scope_scans_response_when_request_copy_fails(self): + import threading + + guardrail: Final = _ApplyOnlyObserver() + guardrail.logging_only_scope = "output" + call: Final = _logged_call([{"role": "user", "content": "hello there", "lock": threading.Lock()}]) + kwargs: Final = call[0] + response: Final = call[1] + + out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert guardrail.calls == [("response", ["general kenobi"])] + entries: Final = out_kwargs["standard_logging_object"]["guardrail_information"] + assert [entry["guardrail_name"] for entry in entries] == ["apply-only-observer"] + assert [entry["guardrail_status"] for entry in entries] == ["success"] + @pytest.mark.asyncio async def test_block_verdict_is_recorded_without_raising(self): guardrail = _ApplyOnlyObserver(block=True) @@ -2611,9 +2654,12 @@ class TestLoggingOnlyApplyGuardrail: out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) - assert guardrail.calls == [("request", ["flagged content"])] + assert guardrail.calls == [ + ("request", ["flagged content"]), + ("response", ["general kenobi"]), + ] entries = out_kwargs["standard_logging_object"]["guardrail_information"] - assert [e["guardrail_status"] for e in entries] == ["guardrail_intervened"] + assert [entry["guardrail_status"] for entry in entries] == ["guardrail_intervened", "guardrail_intervened"] @pytest.mark.asyncio async def test_call_type_without_translation_is_skipped(self): @@ -2932,6 +2978,17 @@ class _NativeLifecycleLoggingGuardrail(CustomGuardrail): return inputs +@pytest.mark.asyncio +async def test_native_lifecycle_guardrail_logging_only_scope_scans_only_input(): + guardrail: Final = _NativeLifecycleLoggingGuardrail() + guardrail.logging_only_scope = "input" + kwargs, response = _logged_call([{"role": "user", "content": "native lifecycle input"}]) + + await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert guardrail.calls == [("request", ["native lifecycle input"])] + + @pytest.mark.asyncio async def test_native_lifecycle_guardrail_logging_only_scans_assembled_response(): """A use_native_lifecycle_hooks guardrail accepts mode logging_only and its @@ -2939,9 +2996,7 @@ async def test_native_lifecycle_guardrail_logging_only_scans_assembled_response( from litellm.types.utils import Choices, Message, ModelResponse guardrail = _NativeLifecycleLoggingGuardrail() - assembled = ModelResponse( - choices=[Choices(message=Message(role="assistant", content="assembled stream text"))] - ) + assembled = ModelResponse(choices=[Choices(message=Message(role="assistant", content="assembled stream text"))]) sentinel_result = object() kwargs = { "model": "gpt-5.4-mini", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailFormField.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailFormField.tsx index 0afd9aab0cf..100d23ad5bb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailFormField.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailFormField.tsx @@ -1,11 +1,17 @@ "use client"; import { CircleHelp } from "lucide-react"; -import React, { useId } from "react"; +import React, { useEffect, useId } from "react"; import { useController, type Control, type ControllerRenderProps, type RegisterOptions } from "react-hook-form"; import { Field, FieldDescription, FieldError, FieldLabel } from "@/components/ui/field"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip"; +import { + getLoggingOnlyScopeOptions, + modeIncludesLoggingOnly, + normalizeLoggingOnlyScopeChoice, + type LoggingOnlyScopeChoice, +} from "./guardrail_info_helpers"; export interface GuardrailCriterion { name: string; @@ -15,6 +21,7 @@ export interface GuardrailCriterion { export interface GuardrailFormValues extends Record { criteria?: GuardrailCriterion[]; + logging_only_scope_choice?: LoggingOnlyScopeChoice; } export type GuardrailFormControl = Control; export type GuardrailFieldRules = Pick, "validate">; @@ -38,6 +45,9 @@ export const asText = (value: unknown): string => { return ""; }; +const isLoggingOnlyScopeChoice = (value: unknown): value is LoggingOnlyScopeChoice => + value === "default" || value === "input" || value === "output" || value === "both"; + export const asStringArray = (value: unknown): string[] => { if (Array.isArray(value)) return value.filter((entry): entry is string => typeof entry === "string"); if (typeof value === "string" && value !== "") return [value]; @@ -123,3 +133,55 @@ export const SkipMessageSelect: React.FC<{ control: GuardrailFieldControlProps } ); }; + +export const LoggingOnlyScopeSelect: React.FC<{ + control: GuardrailFieldControlProps; + directionalScopeSupported: boolean; +}> = ({ control, directionalScopeSupported }) => { + const { id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy } = control; + const items = getLoggingOnlyScopeOptions(directionalScopeSupported); + + useEffect(() => { + const currentChoice = isLoggingOnlyScopeChoice(value) ? value : "default"; + const choice = normalizeLoggingOnlyScopeChoice(currentChoice, directionalScopeSupported); + if (choice !== value) onChange(choice); + }, [value, directionalScopeSupported, onChange]); + + return ( + + ); +}; + +export const LoggingOnlyScopeField: React.FC<{ + control: GuardrailFormControl; + mode: unknown; + directionalScopeSupported: boolean; +}> = ({ control, mode, directionalScopeSupported }) => { + if (!modeIncludesLoggingOnly(mode)) return null; + + return ( + + {(fieldControl) => ( + + )} + + ); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailModeDisplay.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailModeDisplay.tsx new file mode 100644 index 00000000000..e5fa1669a44 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailModeDisplay.tsx @@ -0,0 +1,43 @@ +import React from "react"; +import { Badge } from "@/components/ui/badge"; +import { Card } from "@/components/ui/card"; +import { formatGuardrailMode, formatLoggingOnlyScope, modeIncludesLoggingOnly } from "./guardrail_info_helpers"; + +type GuardrailModeParams = { + mode?: unknown; + default_on?: boolean; + logging_only_scope?: string | null; +}; + +export const GuardrailModeCard: React.FC<{ litellmParams: GuardrailModeParams }> = ({ litellmParams }) => ( + +

Mode

+
+

{formatGuardrailMode(litellmParams.mode) || "-"}

+ + {litellmParams.default_on ? "Default On" : "Default Off"} + +
+ {modeIncludesLoggingOnly(litellmParams.mode) && ( +
+

Logging only scope

+

{formatLoggingOnlyScope(litellmParams.logging_only_scope)}

+
+ )} +
+); + +export const GuardrailModeRows: React.FC<{ litellmParams: GuardrailModeParams }> = ({ litellmParams }) => ( + <> +
+

Mode

+
{formatGuardrailMode(litellmParams.mode) || "-"}
+
+ {modeIncludesLoggingOnly(litellmParams.mode) && ( +
+

Logging only scope

+
{formatLoggingOnlyScope(litellmParams.logging_only_scope)}
+
+ )} + +); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx index aadc9ec0213..230be13ef33 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx @@ -37,6 +37,7 @@ const uiSettings = { supported_entities: [], supported_actions: [], supported_modes: ["pre_call", "post_call"], + providers_without_directional_logging_only_scope: [], pii_entity_categories: [], }; @@ -105,6 +106,67 @@ describe("AddGuardrailForm create payload characterization", () => { expect(payload()).toMatchObject({ litellm_params: { mode: ["pre_call", "post_call"] } }); }); + it("sends the selected output logging-only scope", async () => { + vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({ + ...uiSettings, + supported_modes: ["pre_call", "logging_only"], + }); + const user = userEvent.setup({ delay: null }); + renderForm(); + + await user.type(await screen.findByLabelText("Guardrail Name"), "my-bedrock"); + await pickProvider(user, "Bedrock Guardrail"); + await user.click(screen.getByLabelText("Mode")); + await user.click((await screen.findAllByText("logging_only")).at(-1) as HTMLElement); + await chooseSelectOption(user, await screen.findByLabelText("Logging only scope"), "Output only (response)"); + + await user.type(await screen.findByPlaceholderText("The guardrail id on Bedrock"), "gr-123"); + await user.click(screen.getByRole("button", { name: "Next" })); + await user.click(await screen.findByRole("button", { name: "Create Guardrail" })); + + await waitFor(() => expect(networking.createGuardrailCall).toHaveBeenCalledTimes(1)); + expect(payload()?.litellm_params.logging_only_scope).toBe("output"); + }); + + it("hides directional scope choices for providers that do not support them", async () => { + vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({ + ...uiSettings, + supported_modes: ["pre_call", "logging_only"], + providers_without_directional_logging_only_scope: ["xecguard"], + }); + vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({ + ...providerParams, + xecguard: { ui_friendly_name: "XecGuard" }, + }); + const user = userEvent.setup({ delay: null }); + renderForm(); + + await pickProvider(user, "XecGuard"); + await user.click(screen.getByLabelText("Mode")); + await user.click((await screen.findAllByText("logging_only")).at(-1) as HTMLElement); + await user.click(await screen.findByLabelText("Logging only scope")); + + expect(screen.queryByRole("option", { name: "Input only (request)" })).not.toBeInTheDocument(); + expect(screen.queryByRole("option", { name: "Output only (response)" })).not.toBeInTheDocument(); + expect(screen.getByRole("option", { name: "Both (request and response)" })).toBeInTheDocument(); + }); + + it("hides logging-only scope and omits it from a pre-call payload", async () => { + const user = userEvent.setup({ delay: null }); + renderForm(); + + await user.type(await screen.findByLabelText("Guardrail Name"), "my-bedrock"); + await pickProvider(user, "Bedrock Guardrail"); + expect(screen.queryByLabelText("Logging only scope")).not.toBeInTheDocument(); + + await user.type(await screen.findByPlaceholderText("The guardrail id on Bedrock"), "gr-123"); + await user.click(screen.getByRole("button", { name: "Next" })); + await user.click(await screen.findByRole("button", { name: "Create Guardrail" })); + + await waitFor(() => expect(networking.createGuardrailCall).toHaveBeenCalledTimes(1)); + expect(payload()?.litellm_params).not.toHaveProperty("logging_only_scope"); + }); + it("blocks Next when the user deselects every mode", async () => { const user = userEvent.setup({ delay: null }); renderForm(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index 29df7c8bf3d..1cf17c89d51 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -1,5 +1,5 @@ import React, { useEffect, useMemo, useState } from "react"; -import { useForm, type UseFormReturn } from "react-hook-form"; +import { useForm, useWatch, type UseFormReturn } from "react-hook-form"; import { toast } from "@/lib/toast"; import { createGuardrailCall, @@ -10,18 +10,23 @@ import { import ContentFilterConfiguration from "./content_filter/ContentFilterConfiguration"; import { type CompetitorIntentConfig } from "./content_filter/CompetitorIntentConfiguration"; import { + choiceToLoggingOnlyScope, choiceToSkipSystemForCreate, choiceToSkipToolForCreate, getGuardrailLogo, getGuardrailProviders, getSupportedModesForProvider, guardrail_provider_map, + modeIncludesLoggingOnly, populateGuardrailProviderMap, populateGuardrailProviders, shouldRenderContentFilterConfigSettings, shouldRenderLLMJudgeFields, shouldRenderPIIConfigSettings, + supportsDirectionalLoggingOnlyScope, toModeArray, + type LoggingOnlyScope, + type LoggingOnlyScopeChoice, } from "./guardrail_info_helpers"; import { Logo } from "@/components/molecules/logo/Logo"; import { MultiSelect } from "@/components/shared/MultiSelect"; @@ -49,6 +54,7 @@ import { requiredRule, type GuardrailCriterion, type GuardrailFormValues, + LoggingOnlyScopeField, SkipMessageSelect, } from "./GuardrailFormField"; import GuardrailOptionalParams from "./guardrail_optional_params"; @@ -90,6 +96,7 @@ interface GuardrailSettings { supported_actions: string[]; supported_modes: string[]; supported_modes_by_provider?: Record; + providers_without_directional_logging_only_scope?: string[]; pii_entity_categories: Array<{ category: string; entities: string[]; @@ -160,6 +167,7 @@ type SkipMessageChoice = "inherit" | "yes" | "no"; const INITIAL_VALUES: GuardrailFormValues = { mode: "pre_call", default_on: false, + logging_only_scope_choice: "default", skip_system_message_choice: "inherit", skip_tool_message_choice: "inherit", }; @@ -199,6 +207,7 @@ interface ProviderParamsResponse { const AddGuardrailForm: React.FC = ({ visible, onClose, accessToken, onSuccess, preset }) => { const form = useForm({ defaultValues: INITIAL_VALUES }); + const watchedMode = useWatch({ control: form.control, name: "mode" }); const [loading, setLoading] = useState(false); const [selectedProvider, setSelectedProvider] = useState(null); const [guardrailSettings, setGuardrailSettings] = useState(null); @@ -234,6 +243,7 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a const providerValue = guardrail_provider_map[selectedProvider]; return (providerValue || "").toLowerCase() === "tool_permission"; }, [selectedProvider]); + const directionalScopeSupported = supportsDirectionalLoggingOnlyScope(guardrailSettings, selectedProvider); // Fetch guardrail UI settings + provider params on mount / accessToken change useEffect(() => { @@ -277,6 +287,7 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a guardrail_name: preset.guardrailNameSuggestion, mode: preset.mode, default_on: preset.defaultOn, + logging_only_scope_choice: "default", skip_system_message_choice: "inherit", skip_tool_message_choice: "inherit", }; @@ -439,6 +450,7 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a guardrail_name: string; litellm_params: { guardrail: string; + logging_only_scope?: LoggingOnlyScope | null; [key: string]: unknown; // Allow dynamic properties }; guardrail_info: Record; @@ -462,6 +474,13 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a guardrailData.litellm_params.skip_tool_message_in_guardrail = skipToolForCreate; } + const loggingOnlyScope = choiceToLoggingOnlyScope( + values.logging_only_scope_choice as LoggingOnlyScopeChoice | undefined, + ); + if (modeIncludesLoggingOnly(values.mode) && loggingOnlyScope !== null) { + guardrailData.litellm_params.logging_only_scope = loggingOnlyScope; + } + // For Presidio PII, add the entity and action configurations if (providerKey === "PresidioPII" && selectedEntities.length > 0) { const piiEntitiesConfig: { [key: string]: string } = {}; @@ -796,6 +815,12 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a {(fieldControl) => } + + {/* Use the GuardrailProviderFields component to render provider-specific fields */} {showProviderFields && ( { expect(lastPayload()).toEqual({ litellm_params: { skip_system_message_in_guardrail: true } }); }); + it("shows and updates the logging-only scope", async () => { + const guardrailParams = { + guardrailIdentifier: "gr-abc", + api_key: "sk-old", + mode: "logging_only", + logging_only_scope: "input", + }; + vi.mocked(networking.getGuardrailInfo).mockResolvedValue(guardrail(guardrailParams)); + const user = userEvent.setup({ delay: null }); + renderView(); + + expect(await screen.findAllByText("Input only (request)")).toHaveLength(2); + await openEditor(user); + await chooseSelectOption(user, screen.getByLabelText("Logging only scope"), "Output only (response)"); + await saveChanges(user); + + await waitFor(() => expect(networking.updateGuardrailCall).toHaveBeenCalledTimes(1)); + expect(lastPayload()).toEqual({ litellm_params: { logging_only_scope: "output" } }); + }); + + it("clears the logging-only scope when the edit choice returns to default", async () => { + const guardrailParams = { + guardrailIdentifier: "gr-abc", + api_key: "sk-old", + mode: "logging_only", + logging_only_scope: "input", + }; + vi.mocked(networking.getGuardrailInfo).mockResolvedValue(guardrail(guardrailParams)); + const user = userEvent.setup({ delay: null }); + renderView(); + await openEditor(user); + + await chooseSelectOption(user, screen.getByLabelText("Logging only scope"), "Default (request and response)"); + await saveChanges(user); + + await waitFor(() => expect(networking.updateGuardrailCall).toHaveBeenCalledTimes(1)); + expect(lastPayload()).toEqual({ litellm_params: { logging_only_scope: null } }); + }); + + it("clears a stored directional scope for a provider that does not support it", async () => { + vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({ + ...uiSettings, + providers_without_directional_logging_only_scope: ["bedrock"], + }); + vi.mocked(networking.getGuardrailInfo).mockResolvedValue( + guardrail({ guardrailIdentifier: "gr-abc", mode: "logging_only", logging_only_scope: "output" }), + ); + const user = userEvent.setup({ delay: null }); + renderView(); + await openEditor(user); + await saveChanges(user); + + await waitFor(() => expect(networking.updateGuardrailCall).toHaveBeenCalledTimes(1)); + expect(lastPayload()).toEqual({ litellm_params: { logging_only_scope: null } }); + }); + it("parses the guardrail information textarea into an object", async () => { const user = userEvent.setup({ delay: null }); renderView(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx index c9162d99934..73d9fb8123e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx @@ -28,16 +28,20 @@ import { readRecord, requiredRule, type GuardrailFormValues, + LoggingOnlyScopeField, SkipMessageSelect, } from "./GuardrailFormField"; import ContentFilterManager, { formatContentFilterDataForAPI } from "./content_filter/ContentFilterManager"; import CustomCodeModal, { EditGuardrailData } from "./custom_code/CustomCodeModal"; +import { GuardrailModeCard, GuardrailModeRows } from "./GuardrailModeDisplay"; import { - formatGuardrailMode, + getLoggingOnlyScopeUpdate, getGuardrailLogoAndName, guardrail_provider_map, + loggingOnlyScopeToChoice, skipSystemMessageToChoice, skipToolMessageToChoice, + supportsDirectionalLoggingOnlyScope, type SkipSystemMessageChoice, type SkipToolMessageChoice, } from "./guardrail_info_helpers"; @@ -81,6 +85,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, entities: string[]; }>; supported_modes: string[]; + providers_without_directional_logging_only_scope?: string[]; content_filter_settings?: { prebuilt_patterns: Array<{ name: string; @@ -109,6 +114,8 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, const [toolPermissionConfig, setToolPermissionConfig] = useState(emptyToolPermissionConfig); const [toolPermissionDirty, setToolPermissionDirty] = useState(false); const [customCodeModalVisible, setCustomCodeModalVisible] = useState(false); + const guardrailProvider = guardrailData?.litellm_params?.guardrail ?? null; + const directionalScopeSupported = supportsDirectionalLoggingOnlyScope(guardrailSettings, guardrailProvider); // Content Filter data ref (managed by ContentFilterManager) const contentFilterDataRef = React.useRef<{ @@ -219,6 +226,8 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, if (!guardrailData) return; form.setValue("guardrail_name", guardrailData.guardrail_name); form.setValue("default_on", guardrailData.litellm_params?.default_on); + const storedLoggingOnlyScope = guardrailData.litellm_params?.logging_only_scope; + form.setValue("logging_only_scope_choice", loggingOnlyScopeToChoice(storedLoggingOnlyScope)); form.setValue( "skip_system_message_choice", skipSystemMessageToChoice(guardrailData.litellm_params?.skip_system_message_in_guardrail), @@ -282,7 +291,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, // Prepare update data object - only include changed fields const updateData: any = { - litellm_params: {}, + litellm_params: getLoggingOnlyScopeUpdate(guardrailData.litellm_params, values.logging_only_scope_choice), }; // Only include guardrail_name if it has changed @@ -556,17 +565,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, - -

Mode

-
-

- {formatGuardrailMode(guardrailData.litellm_params?.mode) || "-"} -

- - {guardrailData.litellm_params?.default_on ? "Default On" : "Default Off"} - -
-
+

Created At

@@ -745,6 +744,11 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, > {(fieldControl) => } + {guardrailData.litellm_params?.guardrail === "presidio" && ( <> PII Protection @@ -856,10 +860,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose,

Provider

{displayName}
-
-

Mode

-
{formatGuardrailMode(guardrailData.litellm_params?.mode) || "-"}
-
+

Default On

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx index c5e07fe9624..96ade4b8c1b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx @@ -15,6 +15,14 @@ import { skipToolMessageToChoice, choiceToSkipToolForCreate, formatGuardrailMode, + loggingOnlyScopeToChoice, + choiceToLoggingOnlyScope, + getLoggingOnlyScopeUpdate, + getLoggingOnlyScopeOptions, + formatLoggingOnlyScope, + modeIncludesLoggingOnly, + normalizeLoggingOnlyScopeChoice, + supportsDirectionalLoggingOnlyScope, } from "./guardrail_info_helpers"; describe("guardrail_info_helpers", () => { @@ -27,6 +35,7 @@ describe("guardrail_info_helpers", () => { "PresidioPII", "Bedrock", "Lakera", + "Xecguard", "LitellmContentFilter", "ToolPermission", "BlockCodeExecution", @@ -239,6 +248,100 @@ describe("guardrail_info_helpers", () => { }); }); + describe("logging-only scope helpers", () => { + it("normalizes directional choices only when the provider does not support them", () => { + expect(normalizeLoggingOnlyScopeChoice("input", false)).toBe("default"); + expect(normalizeLoggingOnlyScopeChoice("output", false)).toBe("default"); + expect(normalizeLoggingOnlyScopeChoice("both", false)).toBe("both"); + expect(normalizeLoggingOnlyScopeChoice("default", false)).toBe("default"); + expect(normalizeLoggingOnlyScopeChoice("input", true)).toBe("input"); + expect(normalizeLoggingOnlyScopeChoice("output", true)).toBe("output"); + }); + + it("maps API scope values to choices and back", () => { + expect(loggingOnlyScopeToChoice("input")).toBe("input"); + expect(loggingOnlyScopeToChoice("output")).toBe("output"); + expect(loggingOnlyScopeToChoice("both")).toBe("both"); + expect(loggingOnlyScopeToChoice(undefined)).toBe("default"); + expect(loggingOnlyScopeToChoice(null)).toBe("default"); + expect(loggingOnlyScopeToChoice("invalid")).toBe("default"); + + expect(choiceToLoggingOnlyScope("default")).toBeNull(); + expect(choiceToLoggingOnlyScope(undefined)).toBeNull(); + expect(choiceToLoggingOnlyScope("input")).toBe("input"); + expect(choiceToLoggingOnlyScope("output")).toBe("output"); + expect(choiceToLoggingOnlyScope("both")).toBe("both"); + + expect(getLoggingOnlyScopeUpdate({ logging_only_scope: "input" }, "input")).toEqual({}); + expect(getLoggingOnlyScopeUpdate({ logging_only_scope: "input" }, "output")).toEqual({ + logging_only_scope: "output", + }); + expect(getLoggingOnlyScopeUpdate({ logging_only_scope: "input" }, "default")).toEqual({ + logging_only_scope: null, + }); + }); + + it("formats every scope and falls back to default for missing or unknown values", () => { + expect(formatLoggingOnlyScope("input")).toBe("Input only (request)"); + expect(formatLoggingOnlyScope("output")).toBe("Output only (response)"); + expect(formatLoggingOnlyScope("both")).toBe("Both (request and response)"); + expect(formatLoggingOnlyScope(undefined)).toBe("Default (request and response)"); + expect(formatLoggingOnlyScope(null)).toBe("Default (request and response)"); + expect(formatLoggingOnlyScope("invalid")).toBe("Default (request and response)"); + }); + + it("detects logging_only in string, array, and tagged mode values", () => { + expect(modeIncludesLoggingOnly("logging_only")).toBe(true); + expect(modeIncludesLoggingOnly(["pre_call", "logging_only"])).toBe(true); + expect( + modeIncludesLoggingOnly({ + tags: { "Service-Type: internal-service": "logging_only" }, + default: "pre_call", + }), + ).toBe(true); + expect(modeIncludesLoggingOnly("pre_call")).toBe(false); + }); + + it("filters directional options for unsupported providers and keeps all options otherwise", () => { + expect(getLoggingOnlyScopeOptions(false).map((option) => option.value)).toEqual(["default", "both"]); + expect(getLoggingOnlyScopeOptions(true).map((option) => option.value)).toEqual([ + "default", + "input", + "output", + "both", + ]); + + expect( + supportsDirectionalLoggingOnlyScope( + { providers_without_directional_logging_only_scope: ["xecguard"] }, + "Xecguard", + ), + ).toBe(false); + expect( + supportsDirectionalLoggingOnlyScope( + { providers_without_directional_logging_only_scope: ["xecguard"] }, + "xecguard", + ), + ).toBe(false); + expect( + supportsDirectionalLoggingOnlyScope( + { providers_without_directional_logging_only_scope: ["xecguard"] }, + "Bedrock", + ), + ).toBe(true); + expect( + supportsDirectionalLoggingOnlyScope( + { providers_without_directional_logging_only_scope: ["xecguard"] }, + "unknown-provider", + ), + ).toBe(true); + expect(supportsDirectionalLoggingOnlyScope(null, "Xecguard")).toBe(true); + expect(getLoggingOnlyScopeOptions(supportsDirectionalLoggingOnlyScope(null, "Xecguard"))).toEqual( + getLoggingOnlyScopeOptions(true), + ); + }); + }); + describe("skipSystemMessageToChoice / choiceToSkipSystemForCreate", () => { it("maps API values to form choices and back for create", () => { expect(skipSystemMessageToChoice(undefined)).toBe("inherit"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index 8df7dfb1403..9981979b383 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -114,6 +114,74 @@ export const toModeArray = (raw: unknown): string[] => { return []; }; +export type LoggingOnlyScope = "input" | "output" | "both"; +export type LoggingOnlyScopeChoice = "default" | LoggingOnlyScope; +export type LoggingOnlyScopeOption = { label: string; value: LoggingOnlyScopeChoice }; + +export const normalizeLoggingOnlyScopeChoice = ( + choice: LoggingOnlyScopeChoice, + directionalScopeSupported: boolean, +): LoggingOnlyScopeChoice => + directionalScopeSupported || choice === "default" || choice === "both" ? choice : "default"; + +const LOGGING_ONLY_SCOPE_OPTIONS: LoggingOnlyScopeOption[] = [ + { label: "Default (request and response)", value: "default" }, + { label: "Input only (request)", value: "input" }, + { label: "Output only (response)", value: "output" }, + { label: "Both (request and response)", value: "both" }, +]; + +export const loggingOnlyScopeToChoice = (v: string | null | undefined): LoggingOnlyScopeChoice => + v === "input" || v === "output" || v === "both" ? v : "default"; + +export const choiceToLoggingOnlyScope = (choice: LoggingOnlyScopeChoice | undefined): LoggingOnlyScope | null => + choice === "input" || choice === "output" || choice === "both" ? choice : null; + +export const getLoggingOnlyScopeUpdate = ( + litellmParams: { logging_only_scope?: string | null } | null | undefined, + choice: LoggingOnlyScopeChoice | undefined, +): { logging_only_scope?: LoggingOnlyScope | null } => { + if (choice === undefined || choice === loggingOnlyScopeToChoice(litellmParams?.logging_only_scope)) return {}; + return { logging_only_scope: choiceToLoggingOnlyScope(choice) }; +}; + +export const formatLoggingOnlyScope = (v: string | null | undefined): string => { + if (v === "input") return "Input only (request)"; + if (v === "output") return "Output only (response)"; + if (v === "both") return "Both (request and response)"; + return "Default (request and response)"; +}; + +export const modeIncludesLoggingOnly = (raw: unknown): boolean => { + if (toModeArray(raw).includes("logging_only")) return true; + if (raw === null || typeof raw !== "object") return false; + + const { tags, default: fallback } = raw as { tags?: Record; default?: unknown }; + const taggedModes = + tags && typeof tags === "object" + ? Object.values(tags).some((mode) => toModeArray(mode).includes("logging_only")) + : false; + return toModeArray(fallback).includes("logging_only") || taggedModes; +}; + +export const getLoggingOnlyScopeOptions = (directionalScopeSupported: boolean): LoggingOnlyScopeOption[] => + directionalScopeSupported + ? LOGGING_ONLY_SCOPE_OPTIONS + : LOGGING_ONLY_SCOPE_OPTIONS.filter((option) => option.value === "default" || option.value === "both"); + +export const supportsDirectionalLoggingOnlyScope = ( + settings: { providers_without_directional_logging_only_scope?: string[] } | null, + selectedProvider: string | null, +): boolean => { + const providerKey = selectedProvider + ? ( + guardrail_provider_map[selectedProvider] ?? + Object.values(guardrail_provider_map).find((value) => value.toLowerCase() === selectedProvider.toLowerCase()) + )?.toLowerCase() + : null; + return !providerKey || !settings?.providers_without_directional_logging_only_scope?.includes(providerKey); +}; + export const formatGuardrailMode = (raw: unknown): string => { const flat: string[] = toModeArray(raw); if (flat.length > 0) return flat.join(", "); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 8ffd5c96ab0..229ce097e90 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -25997,6 +25997,11 @@ export interface components { * @description Google Cloud location/region (e.g., us-central1) */ location?: string | null; + /** + * Logging Only Scope + * @description which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking. + */ + logging_only_scope?: ("input" | "output" | "both") | null; /** * Mask Request Content * @description Will mask request content if guardrail makes any changes @@ -32032,6 +32037,27 @@ export interface components { /** Output Text */ output_text: string; }; + /** GuardrailUIAddGuardrailSettings */ + GuardrailUIAddGuardrailSettings: { + /** Content Filter Settings */ + content_filter_settings?: { + [key: string]: unknown; + } | null; + /** Pii Entity Categories */ + pii_entity_categories: components["schemas"]["PiiEntityCategoryMap"][]; + /** Providers Without Directional Logging Only Scope */ + providers_without_directional_logging_only_scope: string[]; + /** Supported Actions */ + supported_actions: string[]; + /** Supported Entities */ + supported_entities: string[]; + /** Supported Modes */ + supported_modes: string[]; + /** Supported Modes By Provider */ + supported_modes_by_provider: { + [key: string]: string[]; + }; + }; /** * HTTPAuthSecurityScheme * @description Defines a security scheme using HTTP authentication. @@ -35456,6 +35482,11 @@ export interface components { * @description Google Cloud location/region (e.g., us-central1) */ location?: string | null; + /** + * Logging Only Scope + * @description which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' (default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking. + */ + logging_only_scope?: ("input" | "output" | "both") | null; /** * Mask * @description Enable content masking using Lasso classifix API @@ -38978,6 +39009,13 @@ export interface components { * @enum {string} */ PiiAction: "BLOCK" | "MASK"; + /** PiiEntityCategoryMap */ + PiiEntityCategoryMap: { + /** Category */ + category: string; + /** Entities */ + entities: string[]; + }; /** * PiiEntityType * @enum {string} @@ -58455,7 +58493,7 @@ export interface operations { [name: string]: unknown; }; content: { - "application/json": unknown; + "application/json": components["schemas"]["GuardrailUIAddGuardrailSettings"]; }; }; };