Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py: 11%
418 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1# +-------------------------------------------------------------+
2#
3# Use Lasso Security Guardrails for your LLM calls
4# https://www.lasso.security/
5#
6# +-------------------------------------------------------------+
8import json
9import os
10import uuid
11from collections.abc import Mapping, Sequence
12from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
14try:
15 import ulid
17 ULID_AVAILABLE = True
18except ImportError:
19 ulid = None
20 ULID_AVAILABLE = False
22try:
23 import httpx
25 HTTPX_AVAILABLE = True
26except ImportError:
27 httpx = None
28 HTTPX_AVAILABLE = False
30from fastapi import HTTPException
32import litellm
33from litellm import DualCache
34from litellm._logging import verbose_proxy_logger
35from litellm.integrations.custom_guardrail import (
36 CustomGuardrail,
37 log_guardrail_information,
38)
39from litellm.integrations.custom_guardrail import dc as global_cache
40from litellm.llms.custom_httpx.http_handler import (
41 get_async_httpx_client,
42 httpxSpecialProvider,
43)
44from litellm.proxy._types import UserAPIKeyAuth
45from litellm.proxy.guardrails._content_utils import (
46 build_inspection_messages,
47 has_non_string_content,
48)
49from litellm.types.guardrails import GuardrailEventHooks
52class LassoResponse(TypedDict):
53 """Type definition for Lasso API response."""
55 violations_detected: bool
56 deputies: dict[str, bool]
57 findings: dict[str, list[dict[str, object]]]
58 messages: list[dict[str, str]] | None
61if TYPE_CHECKING: 61 ↛ 62line 61 didn't jump to line 62 because the condition on line 61 was never true
62 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
65class LassoGuardrailMissingSecrets(Exception):
66 """Exception raised when Lasso API key is missing."""
69class LassoGuardrailAPIError(Exception):
70 """Exception raised when there's an error calling the Lasso API."""
73class LassoGuardrail(CustomGuardrail):
74 """
75 Lasso Security Guardrail integration for LiteLLM.
77 Provides content moderation, PII detection, and policy enforcement
78 through the Lasso Security API.
79 """
81 @classmethod
82 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
83 return [
84 GuardrailEventHooks.pre_call,
85 GuardrailEventHooks.during_call,
86 GuardrailEventHooks.post_call,
87 ]
89 def __init__(
90 self,
91 lasso_api_key: str | None = None,
92 api_key: str | None = None,
93 api_base: str | None = None,
94 user_id: str | None = None,
95 conversation_id: str | None = None,
96 mask: bool | None = False,
97 **kwargs,
98 ):
99 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
100 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
101 self.lasso_api_key = lasso_api_key or api_key or os.environ.get("LASSO_API_KEY")
102 self.user_id = user_id or os.environ.get("LASSO_USER_ID")
103 self.conversation_id = conversation_id or os.environ.get("LASSO_CONVERSATION_ID")
104 self.mask = mask or False
106 if self.lasso_api_key is None:
107 raise LassoGuardrailMissingSecrets(
108 "Couldn't get Lasso api key, either set the `LASSO_API_KEY` in the environment or "
109 "pass it as a parameter to the guardrail in the config file"
110 )
112 self.api_base = api_base or os.getenv("LASSO_API_BASE") or "https://server.lasso.security/gateway/v3"
114 verbose_proxy_logger.debug(
115 "Lasso guardrail initialized: %s, event_hook: %s, mask: %s",
116 kwargs.get("guardrail_name", "unknown"),
117 kwargs.get("event_hook", "unknown"),
118 self.mask,
119 )
121 super().__init__(**kwargs)
123 @staticmethod
124 def _get_field(obj: object, field: str, default: object = None) -> object:
125 """Get a field from either a dict or a Pydantic object."""
126 if isinstance(obj, dict):
127 return obj.get(field, default)
128 return getattr(obj, field, default)
130 @staticmethod
131 def _extract_tool_call_fields(
132 call: object,
133 ) -> tuple[object, object, dict[str, object] | None]:
134 """Extract (call_id, name, parsed_input) from a tool call.
136 Handles both dict-style and Pydantic object-style tool_calls.
137 Parses the JSON arguments string into a dict when possible.
138 """
139 get: Final = LassoGuardrail._get_field
140 call_id: Final = get(call, "id")
141 func: Final = get(call, "function")
142 if not func:
143 return call_id, None, None
144 name: Final = get(func, "name")
145 args_str: Final = get(func, "arguments")
146 input_data: dict[str, object] | None = None
147 if args_str:
148 try:
149 parsed = json.loads(args_str) if isinstance(args_str, (str, bytes, bytearray)) else None
150 except (json.JSONDecodeError, TypeError):
151 parsed = None
152 if isinstance(parsed, dict):
153 input_data = parsed
154 else:
155 # Preserve the raw argument string so Lasso still inspects
156 # callers that smuggle PII/blocked content as malformed JSON
157 # or non-object payloads.
158 input_data = {"arguments": args_str}
159 return call_id, name, input_data
161 def _generate_ulid(self) -> str:
162 """
163 Generate a ULID (Universally Unique Lexicographically Sortable Identifier).
164 Falls back to UUID if ULID library is not available.
165 """
166 if ULID_AVAILABLE and ulid is not None:
167 return str(ulid.ULID())
168 else:
169 verbose_proxy_logger.debug("ULID library not available, using UUID")
170 return str(uuid.uuid4())
172 @log_guardrail_information
173 async def async_pre_call_hook(
174 self,
175 user_api_key_dict: UserAPIKeyAuth,
176 cache: DualCache, # Deprecated, use global_cache instead (kept to align with CustomGuardrail interface)
177 data: dict,
178 call_type: Literal[
179 "completion",
180 "text_completion",
181 "embeddings",
182 "image_generation",
183 "moderation",
184 "audio_transcription",
185 "pass_through_endpoint",
186 "rerank",
187 "mcp_call",
188 "anthropic_messages",
189 ],
190 ) -> Exception | str | dict | None:
191 """
192 Runs before the LLM API call to validate and potentially modify input.
193 Uses 'PROMPT' messageType as this is input to the model.
194 """
195 # Check if this guardrail should run for this request
196 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.pre_call
197 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
198 return data
200 # Get or generate conversation_id and store it in data for post-call consistency
201 # The conversation_id is being stored in the cache so it can be used by the post_call hook
202 self._get_or_generate_conversation_id(data, global_cache)
204 return await self._run_lasso_guardrail(data, global_cache, message_type="PROMPT")
206 @log_guardrail_information
207 async def async_moderation_hook(
208 self,
209 data: dict,
210 user_api_key_dict: UserAPIKeyAuth,
211 call_type: Literal[
212 "completion",
213 "embeddings",
214 "image_generation",
215 "moderation",
216 "audio_transcription",
217 "responses",
218 "mcp_call",
219 "anthropic_messages",
220 ],
221 cache: DualCache,
222 ):
223 """
224 This is used for during_call moderation.
225 Uses 'PROMPT' messageType as this runs concurrently with input processing.
226 """
227 # Check if this guardrail should run for this request
228 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.during_call
229 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
230 return data
232 return await self._run_lasso_guardrail(data, cache, message_type="PROMPT")
234 @log_guardrail_information
235 async def async_post_call_success_hook(
236 self,
237 data: dict,
238 user_api_key_dict: UserAPIKeyAuth,
239 response,
240 ):
241 """
242 Runs after the LLM API call to validate the response.
243 Uses 'COMPLETION' messageType as this is output from the model.
244 """
245 # Check if this guardrail should run for this request
246 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.post_call
247 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
248 return response
250 # Extract messages from the response for validation
251 if isinstance(response, litellm.ModelResponse):
252 response_messages: Final[list[dict[str, object]]] = []
253 for choice in response.choices:
254 if not hasattr(choice, "message"):
255 continue
256 msg = choice.message
257 if msg.content:
258 response_messages.append({"role": "assistant", "content": msg.content})
259 for call in getattr(msg, "tool_calls", None) or []:
260 call_id, name, input_data = self._extract_tool_call_fields(call)
261 if not call_id or not name:
262 continue
263 response_messages.append(
264 {
265 "role": "model",
266 "content": {
267 "type": "tool_use",
268 "id": call_id,
269 "name": name,
270 "input": input_data,
271 },
272 }
273 )
275 if response_messages:
276 # Include litellm_call_id from original data for conversation_id consistency
277 response_data: Final = {
278 "messages": response_messages,
279 "litellm_call_id": data.get("litellm_call_id"),
280 }
282 # Handle masking for post-call
283 if self.mask:
284 headers: Final = self._prepare_headers(response_data, global_cache)
285 payload: Final = self._prepare_payload(response_messages, response_data, global_cache, "COMPLETION")
286 api_url: Final = f"{self.api_base}/classifix"
288 try:
289 lasso_response = await self._call_lasso_api(headers=headers, payload=payload, api_url=api_url)
290 self._process_lasso_response(lasso_response)
292 # Apply masking to the actual response if masked content is available
293 masked_messages: Final = lasso_response.get("messages")
294 if lasso_response.get("violations_detected") and masked_messages:
295 self._apply_masking_to_model_response(response, masked_messages)
296 verbose_proxy_logger.debug("Applied Lasso masking to model response")
297 except Exception as e:
298 if isinstance(e, HTTPException):
299 raise e
300 verbose_proxy_logger.error("Error in post-call Lasso masking: %s", e)
301 raise LassoGuardrailAPIError(f"Failed to apply post-call masking: {e}")
302 else:
303 # Use the same data for conversation_id consistency (no cache access needed)
304 await self._run_lasso_guardrail(response_data, cache=global_cache, message_type="COMPLETION")
305 verbose_proxy_logger.debug("Post-call Lasso validation completed")
306 else:
307 verbose_proxy_logger.warning("No response messages found to validate")
308 else:
309 verbose_proxy_logger.warning("Unexpected response type for post-call hook: %s", type(response))
311 return response
313 def _get_or_generate_conversation_id(self, data: dict, cache: DualCache) -> str:
314 """
315 Get or generate a conversation_id for this request.
317 This method ensures session consistency by using litellm_call_id as a cache key.
318 The same conversation_id is used for both pre-call and post-call hooks within
319 the same request, enabling proper conversation grouping in Lasso UI.
321 Example:
322 >>> guardrail = LassoGuardrail(lasso_api_key="key")
323 >>> data = {"litellm_call_id": "call_123"}
324 >>> conversation_id = guardrail._get_or_generate_conversation_id(data, cache)
325 >>> # Returns consistent ID for same litellm_call_id
327 Args:
328 data: The request data containing litellm_call_id
329 cache: The cache instance for storing conversation_id
331 Returns:
332 str: The conversation_id to use for this request
333 """
334 # Use global conversation_id if set
335 if self.conversation_id:
336 return self.conversation_id
338 # Get the litellm_call_id which is consistent across all hooks for this request
339 litellm_call_id: Final = data.get("litellm_call_id")
341 if not litellm_call_id:
342 # Fallback to generating a new ULID if no litellm_call_id available
343 return self._generate_ulid()
345 # Use litellm_call_id as cache key for conversation_id
346 cache_key: Final = f"lasso_conversation_id:{litellm_call_id}"
348 # Try to get existing conversation_id from cache
349 try:
350 cached_conversation_id: Final = cache.get_cache(cache_key)
351 if cached_conversation_id:
352 return cached_conversation_id
353 except Exception as e:
354 verbose_proxy_logger.warning("Cache retrieval failed: %s", e)
356 # Generate new conversation_id and store in cache
357 generated_id: Final = self._generate_ulid()
359 try:
360 cache.set_cache(cache_key, generated_id, ttl=3600) # Cache for 1 hour
361 except Exception as e:
362 verbose_proxy_logger.warning("Cache storage failed: %s", e)
364 return generated_id
366 async def _run_lasso_guardrail(
367 self,
368 data: dict,
369 cache: DualCache,
370 message_type: Literal["PROMPT", "COMPLETION"] = "PROMPT",
371 ):
372 """
373 Run the Lasso guardrail with the specified message type.
375 This is the core method that handles both classification and masking workflows.
376 It chooses the appropriate API endpoint based on the masking configuration
377 and processes the response according to Lasso's action-based system.
379 Workflow:
380 1. Validate messages are present
381 2. Prepare headers and payload
382 3. Choose API endpoint (classify vs classifix)
383 4. Call Lasso API
384 5. Process response and apply masking if needed
385 6. Handle blocking vs non-blocking violations
387 Args:
388 data: The request data containing messages
389 cache: The cache instance for storing conversation_id (optional for post-call)
390 message_type: Either "PROMPT" for input or "COMPLETION" for output
392 Raises:
393 LassoGuardrailAPIError: If the Lasso API call fails
394 HTTPException: If blocking violations are detected
395 """
396 raw_messages: Final[list[dict[str, object]]] = data.get("messages") or []
397 messages: list[dict[str, Any]] = self._expand_messages_for_classification(raw_messages) if raw_messages else []
398 messages_count: Final = len(messages)
399 if data.get("input") is not None:
400 # Responses-API payloads carry text in data["input"]. Inspect it
401 # alongside any "messages" array — otherwise a caller can attach
402 # benign messages and stash blocked content in input to bypass.
403 messages.extend(build_inspection_messages({"input": data["input"]}))
404 if not messages:
405 return data
407 # Lasso's classifix endpoint returns masked text that we copy back
408 # into ``data["messages"]``. For multimodal/Responses-API input we
409 # would silently strip image/audio parts, so fall back to the
410 # classify endpoint (which still raises on BLOCK actions) and
411 # leave the original payload intact.
412 if self.mask and not has_non_string_content(data):
413 return await self._handle_masking(data, cache, message_type, messages, messages_count)
414 return await self._handle_classification(data, cache, message_type, messages)
416 async def _handle_classification(
417 self,
418 data: dict,
419 cache: DualCache,
420 message_type: Literal["PROMPT", "COMPLETION"],
421 messages: list[dict[str, object]],
422 ) -> dict:
423 """Handle classification without masking."""
424 try:
425 headers: Final = self._prepare_headers(data, cache)
426 payload: Final = self._prepare_payload(messages, data, cache, message_type)
427 response: Final = await self._call_lasso_api(headers=headers, payload=payload)
428 self._process_lasso_response(response)
429 return data
430 except Exception as e:
431 await self._handle_api_error(e, message_type)
432 return data # This line won't be reached due to exception, but satisfies type checker
434 async def _handle_masking(
435 self,
436 data: dict,
437 cache: DualCache,
438 message_type: Literal["PROMPT", "COMPLETION"],
439 messages: list[dict[str, object]],
440 messages_count: int,
441 ) -> dict:
442 """Handle masking with classifix endpoint.
444 ``messages_count`` is the number of inspected items derived from
445 ``data["messages"]``; any items beyond that index came from
446 ``data["input"]`` and must be written back there, not into messages.
447 """
448 try:
449 headers: Final = self._prepare_headers(data, cache)
450 payload: Final = self._prepare_payload(messages, data, cache, message_type)
451 api_url: Final = f"{self.api_base}/classifix"
452 response: Final = await self._call_lasso_api(headers=headers, payload=payload, api_url=api_url)
453 self._process_lasso_response(response)
455 # Apply masking to messages if violations detected and masked messages are available.
456 # Map masked content back onto the original OpenAI-format messages so the
457 # downstream provider receives a compatible payload.
458 masked: Final = response.get("messages")
459 if response.get("violations_detected") and masked:
460 masked_for_messages: Final = masked[:messages_count]
461 masked_for_input: Final = masked[messages_count:]
462 if data.get("messages"):
463 data["messages"] = self._map_masked_messages_back(data["messages"], masked_for_messages)
464 # Also update data["input"] for Responses-API payloads so the
465 # unredacted text doesn't leak through that field.
466 if isinstance(data.get("input"), str):
467 text_parts = [msg["content"] for msg in masked_for_input if isinstance(msg.get("content"), str)]
468 if text_parts:
469 data["input"] = "\n".join(text_parts)
470 self._log_masking_applied(message_type, dict(response))
472 return data
473 except Exception as e:
474 await self._handle_api_error(e, message_type)
475 return data # This line won't be reached due to exception, but satisfies type checker
477 def _map_masked_messages_back(
478 self,
479 original_messages: list[dict[str, Any]],
480 masked_messages: Sequence[Mapping[str, object]],
481 ) -> list[dict[str, object]]:
482 """Map Lasso-format masked messages back onto the original OpenAI-format messages.
484 Lasso receives expanded messages (tool_use / tool_result blocks) and returns them
485 in the same Lasso-internal format with sensitive values replaced. Writing those
486 blocks straight into data["messages"] would corrupt the OpenAI-compatible schema
487 the downstream provider expects. This helper re-applies only the masked content
488 while preserving the original structure.
489 """
490 # Index masked content by type so we can look up by id without caring about order.
491 masked_tool_use: Final[dict[object, dict[str, object]]] = {}
492 masked_tool_result: Final[dict[str, str]] = {}
493 masked_text: Final[list[str]] = []
495 for msg in masked_messages:
496 content = msg.get("content")
497 if isinstance(content, dict):
498 if content.get("type") == "tool_use":
499 call_id = content.get("id")
500 if call_id:
501 masked_tool_use[call_id] = content
502 elif content.get("type") == "tool_result":
503 tool_use_id = content.get("tool_use_id")
504 if tool_use_id:
505 masked_tool_result[tool_use_id] = content.get("content", "")
506 elif isinstance(content, str):
507 masked_text.append(content)
509 # Positional cursor only works if Lasso echoes every text message back.
510 # Skip text remap on count mismatch to avoid writing masked content
511 # onto the wrong original message.
512 original_text_count: Final = sum(
513 1
514 for m in original_messages
515 if m.get("role") != "tool"
516 and ((isinstance(m.get("content"), str) and m.get("content")) or isinstance(m.get("content"), list))
517 )
518 apply_text_cursor: Final = original_text_count == len(masked_text)
519 if not apply_text_cursor and masked_text:
520 verbose_proxy_logger.warning(
521 "Lasso masked-text count mismatch; skipping text remap",
522 extra={
523 "original_text_count": original_text_count,
524 "masked_text_count": len(masked_text),
525 },
526 )
528 result: Final[list[dict[str, object]]] = []
529 text_cursor = 0
531 for orig_msg in original_messages:
532 msg = dict(orig_msg)
533 role = msg.get("role")
534 content = msg.get("content")
536 if role == "tool":
537 tool_call_id = msg.get("tool_call_id")
538 if tool_call_id and tool_call_id in masked_tool_result:
539 msg["content"] = masked_tool_result[tool_call_id]
541 elif isinstance(content, str) and content:
542 if apply_text_cursor and text_cursor < len(masked_text):
543 msg["content"] = masked_text[text_cursor]
544 text_cursor += 1
545 if role == "assistant" and orig_msg.get("tool_calls"):
546 msg["tool_calls"] = self._update_tool_calls_from_masked(orig_msg["tool_calls"], masked_tool_use)
548 elif isinstance(content, list):
549 # Multimodal list content was flattened to a text string before
550 # being sent to Lasso. Replace the list with the masked text
551 # so the cursor stays aligned with subsequent messages.
552 if apply_text_cursor and text_cursor < len(masked_text):
553 msg["content"] = masked_text[text_cursor]
554 text_cursor += 1
555 if role == "assistant" and orig_msg.get("tool_calls"):
556 msg["tool_calls"] = self._update_tool_calls_from_masked(orig_msg["tool_calls"], masked_tool_use)
558 elif role == "assistant" and not content and orig_msg.get("tool_calls"):
559 msg["tool_calls"] = self._update_tool_calls_from_masked(orig_msg["tool_calls"], masked_tool_use)
561 result.append(msg)
563 return result
565 def _update_tool_calls_from_masked(
566 self,
567 tool_calls: list[object],
568 masked_tool_use: Mapping[object, Mapping[str, object]],
569 ) -> list[object]:
570 """Replace tool_call arguments with masked values returned by Lasso."""
571 updated: Final = []
572 for call in tool_calls:
573 call_id = self._get_field(call, "id")
574 if call_id and call_id in masked_tool_use:
575 masked_input = masked_tool_use[call_id].get("input")
576 if masked_input is not None:
577 if isinstance(call, dict):
578 call = dict(call)
579 func_dict = dict(call.get("function", {}))
580 func_dict["arguments"] = json.dumps(masked_input)
581 call["function"] = func_dict
582 else:
583 func_obj = getattr(call, "function", None)
584 if func_obj:
585 func_obj.arguments = json.dumps(masked_input)
586 updated.append(call)
587 return updated
589 async def _handle_api_error(
590 self,
591 error: Exception,
592 message_type: Literal["PROMPT", "COMPLETION"],
593 ) -> None:
594 """Handle API errors with specific error types."""
595 if isinstance(error, HTTPException):
596 raise error
598 # Log error with context
599 verbose_proxy_logger.error(
600 "Error calling Lasso API: %s",
601 error,
602 extra={
603 "guardrail_name": getattr(self, "guardrail_name", "unknown"),
604 "message_type": message_type,
605 "error_type": type(error).__name__,
606 },
607 )
609 # Handle specific error types if httpx is available
610 if HTTPX_AVAILABLE:
611 if isinstance(error, httpx.TimeoutException):
612 raise LassoGuardrailAPIError("Lasso API timeout")
613 elif isinstance(error, httpx.HTTPStatusError):
614 if error.response.status_code == 401:
615 raise LassoGuardrailMissingSecrets("Invalid API key")
616 elif error.response.status_code == 429:
617 raise LassoGuardrailAPIError("Lasso API rate limit exceeded")
618 else:
619 raise LassoGuardrailAPIError(f"API error: {error.response.status_code}")
621 # Generic error handling
622 raise LassoGuardrailAPIError(f"Failed to verify request safety with Lasso API: {error}")
624 def _log_masking_applied(
625 self,
626 message_type: Literal["PROMPT", "COMPLETION"],
627 response: dict[str, Any],
628 ) -> None:
629 """Log masking application with structured context."""
630 conversation_id: Final = getattr(self, "conversation_id", "unknown")
631 verbose_proxy_logger.debug(
632 "Lasso masking applied",
633 extra={
634 "guardrail_name": getattr(self, "guardrail_name", "unknown"),
635 "message_type": message_type,
636 "violations_count": len(response.get("findings", {})),
637 "masked_fields": len(response.get("messages", [])),
638 "conversation_id": conversation_id,
639 },
640 )
642 def _expand_messages_for_classification(self, messages: list[dict[str, Any]]) -> list[dict[str, object]]:
643 """
644 Convert raw OpenAI-format messages to Lasso API format with content blocks.
646 - assistant messages with `tool_calls` → assistant message per tool_use block
647 - role=tool messages → developer role + tool_result block
648 - plain text messages pass through unchanged
649 """
650 expanded: Final[list[dict[str, object]]] = []
651 for msg in messages:
652 role = msg.get("role", "")
653 content = msg.get("content")
655 if role == "tool":
656 tool_call_id = msg.get("tool_call_id")
657 if not tool_call_id:
658 verbose_proxy_logger.warning("Skipping tool message without tool_call_id")
659 continue
660 # Flatten multimodal list content to text so Lasso's
661 # tool_result.content field receives a string.
662 if isinstance(content, list):
663 text_parts = [
664 part["text"]
665 for part in content
666 if isinstance(part, dict) and part.get("type") == "text" and part.get("text")
667 ]
668 tool_result_content = "\n".join(text_parts)
669 else:
670 tool_result_content = content or ""
671 expanded.append(
672 {
673 "role": "developer",
674 "content": {
675 "type": "tool_result",
676 "tool_use_id": tool_call_id,
677 "content": tool_result_content,
678 },
679 }
680 )
681 continue
683 if isinstance(content, list):
684 # Flatten multimodal content arrays to plain text for Lasso.
685 text_parts = [
686 part["text"]
687 for part in content
688 if isinstance(part, dict) and part.get("type") == "text" and part.get("text")
689 ]
690 if text_parts:
691 expanded.append({"role": role, "content": "\n".join(text_parts)})
692 elif content:
693 # Empty string and ``None`` are skipped on purpose: empty
694 # carries no inspectable text and ``None`` is the standard
695 # OpenAI shape for a pure tool-call turn. Dict content
696 # (pre-built tool_use/tool_result blocks from the post-call
697 # path) passes through unchanged.
698 expanded.append({"role": role, "content": content})
700 if role == "assistant":
701 for call in msg.get("tool_calls") or []:
702 call_id, name, input_data = self._extract_tool_call_fields(call)
703 if not call_id or not name:
704 verbose_proxy_logger.warning(
705 "Skipping malformed tool_call",
706 extra={"call_id": call_id, "name": name},
707 )
708 continue
709 expanded.append(
710 {
711 "role": "model",
712 "content": {
713 "type": "tool_use",
714 "id": call_id,
715 "name": name,
716 "input": input_data,
717 },
718 }
719 )
721 return expanded
723 def _prepare_headers(self, data: dict, cache: DualCache) -> dict[str, str]:
724 """Prepare headers for the Lasso API request."""
725 if not self.lasso_api_key:
726 raise LassoGuardrailMissingSecrets(
727 "Couldn't get Lasso api key, either set the `LASSO_API_KEY` in the environment or "
728 "pass it as a parameter to the guardrail in the config file"
729 )
731 headers: Final[dict[str, str]] = {
732 "lasso-api-key": self.lasso_api_key,
733 "Content-Type": "application/json",
734 }
736 # Add optional headers if provided
737 if self.user_id:
738 headers["lasso-user-id"] = self.user_id
740 # Always include conversation_id (generated or provided)
741 conversation_id: Final = self._get_or_generate_conversation_id(data, cache)
743 headers["lasso-conversation-id"] = conversation_id
745 return headers
747 def _prepare_payload(
748 self,
749 messages: list[dict[str, object]],
750 data: dict,
751 cache: DualCache,
752 message_type: Literal["PROMPT", "COMPLETION"] = "PROMPT",
753 ) -> dict[str, object]:
754 """
755 Prepare the payload for the Lasso API request.
757 Args:
758 messages: List of message objects (may contain tool_use/tool_result content blocks)
759 message_type: Type of message - "PROMPT" for input, "COMPLETION" for output
760 data: Request data (used for conversation_id generation and tools extraction)
761 cache: Cache instance for storing conversation_id (optional for post-call)
762 """
763 payload: Final[dict[str, object]] = {
764 "messages": messages,
765 "messageType": message_type,
766 # Drives the "Used By" badge on Lasso Application API Keys: every call from this
767 # integration is attributed as "litellm" on the keys list.
768 "source": {"type": "litellm"},
769 }
771 # Add optional parameters if available
772 if self.user_id:
773 payload["userId"] = self.user_id
775 # Always include sessionId (conversation_id - generated or provided)
776 conversation_id: Final = self._get_or_generate_conversation_id(data, cache)
777 payload["sessionId"] = conversation_id
779 # Map OpenAI ChatCompletionToolParam array → ToolDefinition array
780 tools_data: Final[list[dict[str, object]]] = data.get("tools") or []
781 if tools_data:
782 get: Final = self._get_field
783 tool_definitions: Final = []
784 for tool in tools_data:
785 func = get(tool, "function")
786 if not func:
787 continue
788 name = get(func, "name")
789 if not name:
790 continue
791 td: dict[str, object] = {"name": name}
792 description = get(func, "description")
793 if description:
794 td["description"] = description
795 parameters = get(func, "parameters")
796 if parameters:
797 td["parameters"] = parameters
798 tool_definitions.append(td)
799 if tool_definitions:
800 payload["tools"] = tool_definitions
802 return payload
804 async def _call_lasso_api(
805 self,
806 headers: dict[str, str],
807 payload: dict[str, object],
808 api_url: str | None = None,
809 ) -> LassoResponse:
810 """Call the Lasso API and return the response."""
811 url: Final = api_url or f"{self.api_base}/classify"
812 verbose_proxy_logger.debug("Calling Lasso API with messageType: %s", payload.get("messageType"))
813 response: Final = await self.async_handler.post(
814 url=url,
815 headers=headers,
816 json=payload,
817 timeout=10.0,
818 )
819 response.raise_for_status()
820 return response.json()
822 def _process_lasso_response(self, response: LassoResponse) -> None:
823 """
824 Process the Lasso API response and handle violations according to action types.
826 This method implements the action-based blocking logic:
827 - BLOCK: Raises HTTPException to stop request/response
828 - AUTO_MASKING: Logs warning and continues (masking applied elsewhere)
829 - WARN: Logs warning and continues
831 Example Response:
832 {
833 "violations_detected": true,
834 "findings": {
835 "jailbreak": [{
836 "action": "BLOCK",
837 "severity": "HIGH"
838 }]
839 }
840 }
842 Args:
843 response: The response dictionary from Lasso API
845 Raises:
846 HTTPException: If any finding has "action": "BLOCK"
847 """
848 if response and response.get("violations_detected") is True:
849 violated_deputies: Final = self._parse_violated_deputies(response)
850 verbose_proxy_logger.warning("Lasso guardrail detected violations: %s", violated_deputies)
852 # Check if any findings have "BLOCK" action
853 blocking_violations: Final = self._check_for_blocking_actions(response)
855 if blocking_violations:
856 # Block the request/response for findings with "BLOCK" action
857 raise HTTPException(
858 status_code=400,
859 detail={
860 "error": "Violated Lasso guardrail policy",
861 "detection_message": f"Blocking violations detected: {', '.join(blocking_violations)}",
862 "lasso_response": response,
863 },
864 )
865 else:
866 # Continue with warning for non-blocking violations (e.g., AUTO_MASKING)
867 verbose_proxy_logger.info(
868 "Non-blocking Lasso violations detected, continuing with warning: %s", violated_deputies
869 )
871 def _check_for_blocking_actions(self, response: LassoResponse) -> list[str]:
872 """
873 Check findings for actions that should block the request/response.
875 Examines the findings section of the Lasso response to identify which
876 deputies have violations with "BLOCK" action. This enables granular
877 control where some violations (like PII) can be masked while others
878 (like jailbreaks) are blocked entirely.
880 Args:
881 response: The response dictionary from Lasso API
883 Returns:
884 List[str]: Names of deputies with blocking violations
886 Example:
887 >>> response = {
888 ... "findings": {
889 ... "jailbreak": [{"action": "BLOCK"}],
890 ... "pattern-detection": [{"action": "AUTO_MASKING"}]
891 ... }
892 ... }
893 >>> guardrail._check_for_blocking_actions(response)
894 ['jailbreak']
895 """
896 blocking_violations: Final = []
897 findings: Final = response.get("findings", {})
899 for deputy_name, deputy_findings in findings.items():
900 if isinstance(deputy_findings, list):
901 for finding in deputy_findings:
902 if isinstance(finding, dict) and finding.get("action") == "BLOCK":
903 if deputy_name not in blocking_violations:
904 blocking_violations.append(deputy_name)
905 break # No need to check other findings for this deputy
907 return blocking_violations
909 def _parse_violated_deputies(self, response: LassoResponse) -> list[str]:
910 """Parse the response to extract violated deputies."""
911 violated_deputies: Final = []
912 if "deputies" in response:
913 for deputy, is_violated in response["deputies"].items():
914 if is_violated:
915 violated_deputies.append(deputy)
916 return violated_deputies
918 def _apply_masking_to_model_response(
919 self,
920 model_response: litellm.ModelResponse,
921 masked_messages: Sequence[Mapping[str, object]],
922 ) -> None:
923 """Apply masking to the actual model response when mask=True and masked content is available."""
924 # Index masked tool_use blocks by id for O(1) lookup.
925 masked_tool_use: Final[dict[object, dict[str, object]]] = {}
926 masked_text: Final[list[str]] = []
927 for masked_msg in masked_messages:
928 content = masked_msg.get("content")
929 if isinstance(content, dict) and content.get("type") == "tool_use":
930 call_id = content.get("id")
931 if call_id:
932 masked_tool_use[call_id] = content
933 elif isinstance(content, str):
934 masked_text.append(content)
936 # Count text-bearing choices to verify 1:1 mapping with masked texts.
937 original_text_count = sum(1 for c in model_response.choices if hasattr(c, "message") and c.message.content)
938 apply_text: Final = original_text_count == len(masked_text)
939 if not apply_text and masked_text:
940 verbose_proxy_logger.warning(
941 "Lasso masked-text count mismatch in model response; skipping text remap",
942 extra={
943 "original_text_count": original_text_count,
944 "masked_text_count": len(masked_text),
945 },
946 )
948 text_cursor = 0
949 for choice in model_response.choices:
950 if not hasattr(choice, "message"):
951 continue
952 msg = choice.message
954 if msg.content and apply_text and text_cursor < len(masked_text):
955 msg.content = masked_text[text_cursor]
956 text_cursor += 1
957 verbose_proxy_logger.debug("Applied masked text content to choice %s", text_cursor)
959 for call in getattr(msg, "tool_calls", None) or []:
960 call_id = self._get_field(call, "id")
961 if call_id and call_id in masked_tool_use:
962 masked_input = masked_tool_use[call_id].get("input")
963 if masked_input is not None:
964 if isinstance(call, dict):
965 func = call.get("function", {})
966 if isinstance(func, dict):
967 func["arguments"] = json.dumps(masked_input)
968 else:
969 func = getattr(call, "function", None)
970 if func:
971 func.arguments = json.dumps(masked_input)
972 verbose_proxy_logger.debug("Applied masked tool_call arguments for call_id=%s", call_id)
974 @staticmethod
975 def get_config_model() -> type["GuardrailConfigModel"] | None:
976 from litellm.types.proxy.guardrails.guardrail_hooks.lasso import (
977 LassoGuardrailConfigModel,
978 )
980 return LassoGuardrailConfigModel