Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py: 22%
232 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"""
2Built-in primitives provided to custom code guardrails.
4These functions are injected into the custom code execution environment
5and provide safe, sandboxed functionality for common guardrail operations.
6"""
8import json
9import re
10from collections.abc import Mapping, Sequence
11from typing import Final, Literal
12from urllib.parse import urlparse
14import httpx
15from pydantic import JsonValue
16from typing_extensions import ReadOnly, TypedDict
18from litellm._logging import verbose_proxy_logger
19from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client
20from litellm.types.llms.custom_http import httpxSpecialProvider
22# =============================================================================
23# Result Types - Used by Starlark code to return guardrail decisions
24# =============================================================================
27def allow() -> dict[str, object]:
28 """
29 Allow the request/response to proceed unchanged.
31 Returns:
32 Dict indicating the request should be allowed
33 """
34 return {"action": "allow"}
37def block(reason: str, detection_info: Mapping[str, object] | None = None) -> dict[str, object]:
38 """
39 Block the request/response with a reason.
41 Args:
42 reason: Human-readable reason for blocking
43 detection_info: Optional additional detection metadata
45 Returns:
46 Dict indicating the request should be blocked
47 """
48 result: Final[dict[str, object]] = {"action": "block", "reason": reason}
49 if detection_info:
50 result["detection_info"] = detection_info
51 return result
54class FlagResult(TypedDict):
55 action: ReadOnly[Literal["flag"]]
56 reason: ReadOnly[str]
57 metadata: ReadOnly[Mapping[str, object]]
60def flag(reason: str, metadata: Mapping[str, object] | None = None) -> FlagResult:
61 """
62 Let the request/response proceed unchanged but record a non-blocking violation.
64 Args:
65 reason: Human-readable reason for flagging
66 metadata: Optional structured metadata stored alongside the reason
68 Returns:
69 Dict indicating the request should be flagged but allowed
70 """
71 result: Final[FlagResult] = {
72 "action": "flag",
73 "reason": reason,
74 "metadata": metadata if metadata is not None else {},
75 }
76 return result
79def modify(
80 texts: Sequence[str] | None = None,
81 images: Sequence[object] | None = None,
82 tool_calls: Sequence[object] | None = None,
83) -> dict[str, object]:
84 """
85 Modify the request/response content.
87 Args:
88 texts: Modified text content (if None, keeps original)
89 images: Modified image content (if None, keeps original)
90 tool_calls: Modified tool calls (if None, keeps original)
92 Returns:
93 Dict indicating the content should be modified
94 """
95 result: Final[dict[str, object]] = {"action": "modify"}
96 if texts is not None:
97 result["texts"] = texts
98 if images is not None:
99 result["images"] = images
100 if tool_calls is not None:
101 result["tool_calls"] = tool_calls
102 return result
105# =============================================================================
106# Regex Primitives
107# =============================================================================
110def regex_match(text: str, pattern: str, flags: int = 0) -> bool:
111 """
112 Check if a regex pattern matches anywhere in the text.
114 Args:
115 text: The text to search in
116 pattern: The regex pattern to match
117 flags: Optional regex flags (default: 0)
119 Returns:
120 True if pattern matches, False otherwise
121 """
122 try:
123 return bool(re.search(pattern, text, flags))
124 except re.error as e:
125 verbose_proxy_logger.warning("Starlark regex_match error: %s", e)
126 return False
129def regex_match_all(text: str, pattern: str, flags: int = 0) -> bool:
130 """
131 Check if a regex pattern matches the entire text.
133 Args:
134 text: The text to match
135 pattern: The regex pattern
136 flags: Optional regex flags
138 Returns:
139 True if pattern matches entire text, False otherwise
140 """
141 try:
142 return bool(re.fullmatch(pattern, text, flags))
143 except re.error as e:
144 verbose_proxy_logger.warning("Starlark regex_match_all error: %s", e)
145 return False
148def regex_replace(text: str, pattern: str, replacement: str, flags: int = 0) -> str:
149 """
150 Replace all occurrences of a pattern in text.
152 Args:
153 text: The text to modify
154 pattern: The regex pattern to find
155 replacement: The replacement string
156 flags: Optional regex flags
158 Returns:
159 The text with replacements applied
160 """
161 try:
162 return re.sub(pattern, replacement, text, flags=flags)
163 except re.error as e:
164 verbose_proxy_logger.warning("Starlark regex_replace error: %s", e)
165 return text
168def regex_find_all(text: str, pattern: str, flags: int = 0) -> list[str]:
169 """
170 Find all occurrences of a pattern in text.
172 Args:
173 text: The text to search
174 pattern: The regex pattern to find
175 flags: Optional regex flags
177 Returns:
178 List of all matches
179 """
180 try:
181 return re.findall(pattern, text, flags)
182 except re.error as e:
183 verbose_proxy_logger.warning("Starlark regex_find_all error: %s", e)
184 return []
187# =============================================================================
188# JSON Primitives
189# =============================================================================
192class JsonSchemaNode(TypedDict, total=False):
193 """Subset of JSON Schema keywords understood by the built-in validator."""
195 type: ReadOnly[str]
196 required: ReadOnly[Sequence[str]]
197 properties: ReadOnly[Mapping[str, "JsonSchemaNode"]]
200def json_parse(text: str) -> JsonValue:
201 """
202 Parse a JSON string into a Python object.
204 Args:
205 text: The JSON string to parse
207 Returns:
208 Parsed Python object, or None if parsing fails
209 """
210 try:
211 return json.loads(text)
212 except (json.JSONDecodeError, TypeError) as e:
213 verbose_proxy_logger.debug("Starlark json_parse error: %s", e)
214 return None
217def json_stringify(obj: object) -> str:
218 """
219 Convert a Python object to a JSON string.
221 Args:
222 obj: The object to serialize
224 Returns:
225 JSON string representation
226 """
227 try:
228 return json.dumps(obj)
229 except (TypeError, ValueError) as e:
230 verbose_proxy_logger.warning("Starlark json_stringify error: %s", e)
231 return ""
234def json_schema_valid(obj: JsonValue, schema: JsonSchemaNode) -> bool:
235 """
236 Validate an object against a JSON schema.
238 Args:
239 obj: The object to validate
240 schema: The JSON schema to validate against
242 Returns:
243 True if valid, False otherwise
244 """
245 try:
246 # Try to import jsonschema, fall back to basic validation if not available
247 try:
248 import jsonschema
250 jsonschema.validate(instance=obj, schema=schema)
251 return True
252 except ImportError:
253 # Basic validation without jsonschema library
254 return _basic_json_schema_validate(obj, schema)
255 except Exception as validation_error:
256 # Catch jsonschema.ValidationError and other validation errors
257 if "ValidationError" in type(validation_error).__name__:
258 return False
259 raise
260 except Exception as e:
261 verbose_proxy_logger.warning("Custom code json_schema_valid error: %s", e)
262 return False
265def _basic_json_schema_validate(obj: JsonValue, schema: JsonSchemaNode, max_depth: int = 50) -> bool:
266 """
267 Basic JSON schema validation without external library.
268 Handles: type, required, properties
270 Uses an iterative approach with a stack to avoid recursion limits.
271 max_depth limits nesting to prevent infinite loops from circular schemas.
272 """
273 type_map: Final[Mapping[str, type | tuple[type, ...]]] = {
274 "object": dict,
275 "array": list,
276 "string": str,
277 "number": (int, float),
278 "integer": int,
279 "boolean": bool,
280 "null": type(None),
281 }
283 # Stack of (obj, schema, depth) tuples to process
284 stack: Final[list[tuple[JsonValue, JsonSchemaNode, int]]] = [(obj, schema, 0)]
286 while stack:
287 current_obj, current_schema, depth = stack.pop()
289 # Circuit breaker: stop if we've gone too deep
290 if depth > max_depth:
291 return False
293 # Check type
294 schema_type = current_schema.get("type")
295 if schema_type:
296 expected_type: type | tuple[type, ...] | None = type_map.get(schema_type)
297 if expected_type is not None and not isinstance(current_obj, expected_type):
298 return False
300 # Check required fields and properties for dicts
301 if isinstance(current_obj, dict):
302 required: Sequence[str] = current_schema.get("required", [])
303 for field in required:
304 if field not in current_obj:
305 return False
307 # Queue property validations
308 properties: Mapping[str, JsonSchemaNode] = current_schema.get("properties", {})
309 for prop_name, prop_schema in properties.items():
310 if prop_name in current_obj:
311 stack.append((current_obj[prop_name], prop_schema, depth + 1))
313 return True
316# =============================================================================
317# URL Primitives
318# =============================================================================
321# Common URL pattern for extraction
322_URL_PATTERN: Final = re.compile(r"https?://(?:[-\w.]|(?:%[\da-fA-F]{2}))+[^\s]*", re.IGNORECASE)
325def extract_urls(text: str) -> list[str]:
326 """
327 Extract all URLs from text.
329 Args:
330 text: The text to search for URLs
332 Returns:
333 List of URLs found in the text
334 """
335 return _URL_PATTERN.findall(text)
338def is_valid_url(url: str) -> bool:
339 """
340 Check if a URL is syntactically valid.
342 Args:
343 url: The URL to validate
345 Returns:
346 True if the URL is valid, False otherwise
347 """
348 try:
349 result: Final = urlparse(url)
350 return all([result.scheme, result.netloc])
351 except Exception:
352 return False
355def all_urls_valid(text: str) -> bool:
356 """
357 Check if all URLs in text are valid.
359 Args:
360 text: The text containing URLs
362 Returns:
363 True if all URLs are valid (or no URLs), False otherwise
364 """
365 urls: Final = extract_urls(text)
366 return all(is_valid_url(url) for url in urls)
369def get_url_domain(url: str) -> str | None:
370 """
371 Extract the domain from a URL.
373 Args:
374 url: The URL to parse
376 Returns:
377 The domain, or None if invalid
378 """
379 try:
380 result: Final = urlparse(url)
381 return result.netloc if result.netloc else None
382 except Exception:
383 return None
386# =============================================================================
387# HTTP Request Primitives (Async)
388# =============================================================================
390# Default timeout for HTTP requests (in seconds)
391_HTTP_DEFAULT_TIMEOUT: Final = 30.0
393# Maximum allowed timeout (in seconds)
394_HTTP_MAX_TIMEOUT: Final = 60.0
397class HttpResponseResult(TypedDict):
398 """Outcome of an HTTP primitive call, as handed back to custom code."""
400 status_code: ReadOnly[int]
401 body: ReadOnly[JsonValue]
402 headers: ReadOnly[Mapping[str, str]]
403 success: ReadOnly[bool]
404 error: ReadOnly[str | None]
407def _http_error_response(error: str) -> HttpResponseResult:
408 """Create a standardized error response for HTTP requests."""
409 return {
410 "status_code": 0,
411 "body": None,
412 "headers": {},
413 "success": False,
414 "error": error,
415 }
418def _http_success_response(response: httpx.Response) -> HttpResponseResult:
419 """Create a standardized success response from an httpx Response."""
420 parsed_body: JsonValue
421 try:
422 parsed_body = response.json()
423 except (json.JSONDecodeError, ValueError):
424 parsed_body = response.text
426 return {
427 "status_code": response.status_code,
428 "body": parsed_body,
429 "headers": dict(response.headers),
430 "success": 200 <= response.status_code < 300,
431 "error": None,
432 }
435def _prepare_http_body(
436 body: JsonValue,
437) -> tuple[dict[str, JsonValue] | None, str | None]:
438 """Prepare body arguments for HTTP request - returns (json_body, data_body)."""
439 if body is None:
440 return None, None
441 if isinstance(body, dict):
442 return body, None
443 if isinstance(body, list):
444 return None, json.dumps(body)
445 if isinstance(body, str):
446 return None, body
447 return None, str(body)
450async def http_request(
451 url: str,
452 method: str = "GET",
453 headers: dict[str, str] | None = None,
454 body: JsonValue = None,
455 timeout: float | None = None,
456) -> HttpResponseResult:
457 """
458 Make an async HTTP request to an external service.
460 This function allows custom guardrails to call external APIs for
461 additional validation, content moderation, or data enrichment.
463 Uses LiteLLM's global cached AsyncHTTPHandler for connection pooling
464 and better performance.
466 Args:
467 url: The URL to request
468 method: HTTP method (GET, POST, PUT, DELETE, PATCH). Defaults to GET.
469 headers: Optional dict of HTTP headers
470 body: Optional request body (will be JSON-encoded if dict/list)
471 timeout: Optional timeout in seconds (default: 30, max: 60)
473 Returns:
474 Dict containing:
475 - status_code: HTTP status code
476 - body: Response body (parsed as JSON if possible, otherwise string)
477 - headers: Response headers as dict
478 - success: True if status code is 2xx
479 - error: Error message if request failed, None otherwise
481 Example:
482 # Simple GET request
483 response = await http_request("https://api.example.com/check")
484 if response["success"]:
485 data = response["body"]
487 # POST request with JSON body
488 response = await http_request(
489 "https://api.example.com/moderate",
490 method="POST",
491 headers={"Authorization": "Bearer token"},
492 body={"text": "content to check"}
493 )
494 """
495 # Validate URL
496 if not is_valid_url(url):
497 return _http_error_response(f"Invalid URL: {url}")
499 # Validate and normalize method
500 method = method.upper()
501 allowed_methods: Final = {"GET", "POST", "PUT", "DELETE", "PATCH"}
502 if method not in allowed_methods:
503 return _http_error_response(f"Invalid HTTP method: {method}. Allowed: {', '.join(allowed_methods)}")
505 # Apply timeout limits
506 if timeout is None:
507 timeout = _HTTP_DEFAULT_TIMEOUT
508 else:
509 timeout = min(max(0.1, timeout), _HTTP_MAX_TIMEOUT)
511 # Get the global cached async HTTP client
512 client: Final = get_async_httpx_client(
513 llm_provider=httpxSpecialProvider.GuardrailCallback,
514 params={"timeout": httpx.Timeout(timeout=timeout, connect=5.0)},
515 )
517 try:
518 response: Final = await _execute_http_request(client, method, url, headers, body, timeout)
519 return _http_success_response(response)
521 except httpx.TimeoutException as e:
522 verbose_proxy_logger.warning("Custom code http_request timeout: %s", e)
523 return _http_error_response(f"Request timeout after {timeout}s")
524 except httpx.HTTPStatusError as e:
525 # Return the response even for non-2xx status codes
526 return _http_success_response(e.response)
527 except httpx.RequestError as e:
528 verbose_proxy_logger.warning("Custom code http_request error: %s", e)
529 return _http_error_response(f"Request failed: {e}")
530 except Exception as e:
531 verbose_proxy_logger.warning("Custom code http_request unexpected error: %s", e)
532 return _http_error_response(f"Unexpected error: {e}")
535async def _execute_http_request(
536 client: AsyncHTTPHandler,
537 method: str,
538 url: str,
539 headers: dict[str, str] | None,
540 body: JsonValue,
541 timeout: float,
542) -> httpx.Response:
543 """Execute the HTTP request using the appropriate client method."""
544 json_body, data_body = _prepare_http_body(body)
546 if method == "GET":
547 return await client.get(url=url, headers=headers)
548 elif method == "POST":
549 return await client.post(url=url, headers=headers, json=json_body, data=data_body, timeout=timeout)
550 elif method == "PUT":
551 return await client.put(url=url, headers=headers, json=json_body, data=data_body, timeout=timeout)
552 elif method == "DELETE":
553 return await client.delete(url=url, headers=headers, json=json_body, data=data_body, timeout=timeout)
554 elif method == "PATCH":
555 return await client.patch(url=url, headers=headers, json=json_body, data=data_body, timeout=timeout)
556 else:
557 raise ValueError(f"Unsupported HTTP method: {method}")
560async def http_get(
561 url: str,
562 headers: dict[str, str] | None = None,
563 timeout: float | None = None,
564) -> HttpResponseResult:
565 """
566 Make an async HTTP GET request.
568 Convenience wrapper around http_request for GET requests.
570 Args:
571 url: The URL to request
572 headers: Optional dict of HTTP headers
573 timeout: Optional timeout in seconds
575 Returns:
576 Same as http_request
577 """
578 return await http_request(url=url, method="GET", headers=headers, timeout=timeout)
581async def http_post(
582 url: str,
583 body: JsonValue = None,
584 headers: dict[str, str] | None = None,
585 timeout: float | None = None,
586) -> HttpResponseResult:
587 """
588 Make an async HTTP POST request.
590 Convenience wrapper around http_request for POST requests.
592 Args:
593 url: The URL to request
594 body: Optional request body (will be JSON-encoded if dict/list)
595 headers: Optional dict of HTTP headers
596 timeout: Optional timeout in seconds
598 Returns:
599 Same as http_request
600 """
601 return await http_request(url=url, method="POST", headers=headers, body=body, timeout=timeout)
604# =============================================================================
605# Code Detection Primitives
606# =============================================================================
609# Common code patterns for detection
610_CODE_PATTERNS: Final = {
611 "sql": [
612 r"\b(SELECT|INSERT|UPDATE|DELETE|DROP|CREATE|ALTER|TRUNCATE)\b.*\b(FROM|INTO|TABLE|SET|WHERE)\b",
613 r"\b(SELECT)\s+[\w\*,\s]+\s+FROM\s+\w+",
614 r"\b(INSERT\s+INTO|UPDATE\s+\w+\s+SET|DELETE\s+FROM)\b",
615 ],
616 "python": [
617 r"^\s*(def|class|import|from|if|for|while|try|except|with)\s+",
618 r"^\s*@\w+", # decorators
619 r"\b(print|len|range|str|int|float|list|dict|set)\s*\(",
620 ],
621 "javascript": [
622 r"\b(function|const|let|var|class|import|export)\s+",
623 r"=>", # arrow functions
624 r"\b(console\.(log|error|warn))\s*\(",
625 ],
626 "typescript": [
627 r":\s*(string|number|boolean|any|void|never)\b",
628 r"\b(interface|type|enum)\s+\w+",
629 r"<[A-Z]\w*>", # generics
630 ],
631 "java": [
632 r"\b(public|private|protected)\s+(static\s+)?(class|void|int|String)\b",
633 r"\bSystem\.(out|err)\.print",
634 ],
635 "go": [
636 r"\bfunc\s+\w+\s*\(",
637 r"\b(package|import)\s+",
638 r":=", # short variable declaration
639 ],
640 "rust": [
641 r"\b(fn|let|mut|impl|struct|enum|pub|mod)\s+",
642 r"->", # return type
643 r"\b(println!|format!)\s*\(",
644 ],
645 "shell": [
646 r"^#!.*\b(bash|sh|zsh)\b",
647 r"\b(echo|grep|sed|awk|cat|ls|cd|mkdir|rm)\s+",
648 r"\$\{?\w+\}?", # variable expansion
649 ],
650 "html": [
651 r"<\s*(html|head|body|div|span|p|a|img|script|style)\b[^>]*>",
652 r"</\s*(html|head|body|div|span|p|a|script|style)\s*>",
653 ],
654 "css": [
655 r"\{[^}]*:\s*[^}]+;[^}]*\}",
656 r"@(media|keyframes|import|font-face)\b",
657 ],
658}
661def detect_code(text: str) -> bool:
662 """
663 Check if text contains code of any language.
665 Args:
666 text: The text to check
668 Returns:
669 True if code is detected, False otherwise
670 """
671 return len(detect_code_languages(text)) > 0
674def detect_code_languages(text: str) -> list[str]:
675 """
676 Detect which programming languages are present in text.
678 Args:
679 text: The text to analyze
681 Returns:
682 List of detected language names
683 """
684 detected: Final = []
685 for lang, patterns in _CODE_PATTERNS.items():
686 for pattern in patterns:
687 try:
688 if re.search(pattern, text, re.IGNORECASE | re.MULTILINE):
689 detected.append(lang)
690 break # Only add each language once
691 except re.error:
692 continue
693 return detected
696def contains_code_language(text: str, languages: list[str]) -> bool:
697 """
698 Check if text contains code from specific languages.
700 Args:
701 text: The text to check
702 languages: List of language names to check for
704 Returns:
705 True if any of the specified languages are detected
706 """
707 detected: Final = detect_code_languages(text)
708 return any(lang.lower() in [d.lower() for d in detected] for lang in languages)
711# =============================================================================
712# Text Utility Primitives
713# =============================================================================
716def contains(text: str, substring: str) -> bool:
717 """
718 Check if text contains a substring.
720 Args:
721 text: The text to search in
722 substring: The substring to find
724 Returns:
725 True if substring is found, False otherwise
726 """
727 return substring in text
730def contains_any(text: str, substrings: list[str]) -> bool:
731 """
732 Check if text contains any of the given substrings.
734 Args:
735 text: The text to search in
736 substrings: List of substrings to find
738 Returns:
739 True if any substring is found, False otherwise
740 """
741 return any(s in text for s in substrings)
744def contains_all(text: str, substrings: list[str]) -> bool:
745 """
746 Check if text contains all of the given substrings.
748 Args:
749 text: The text to search in
750 substrings: List of substrings to find
752 Returns:
753 True if all substrings are found, False otherwise
754 """
755 return all(s in text for s in substrings)
758def word_count(text: str) -> int:
759 """
760 Count the number of words in text.
762 Args:
763 text: The text to count words in
765 Returns:
766 Number of words
767 """
768 return len(text.split())
771def char_count(text: str) -> int:
772 """
773 Count the number of characters in text.
775 Args:
776 text: The text to count characters in
778 Returns:
779 Number of characters
780 """
781 return len(text)
784def lower(text: str) -> str:
785 """Convert text to lowercase."""
786 return text.lower()
789def upper(text: str) -> str:
790 """Convert text to uppercase."""
791 return text.upper()
794def trim(text: str) -> str:
795 """Remove leading and trailing whitespace."""
796 return text.strip()
799# =============================================================================
800# Primitives Registry
801# =============================================================================
804def get_custom_code_primitives() -> dict[str, object]:
805 """
806 Get all primitives to inject into the custom code environment.
808 Returns:
809 Dict of function name to function
810 """
811 return {
812 # Result types
813 "allow": allow,
814 "block": block,
815 "flag": flag,
816 "modify": modify,
817 # Regex
818 "regex_match": regex_match,
819 "regex_match_all": regex_match_all,
820 "regex_replace": regex_replace,
821 "regex_find_all": regex_find_all,
822 # JSON
823 "json_parse": json_parse,
824 "json_stringify": json_stringify,
825 "json_schema_valid": json_schema_valid,
826 # URL
827 "extract_urls": extract_urls,
828 "is_valid_url": is_valid_url,
829 "all_urls_valid": all_urls_valid,
830 "get_url_domain": get_url_domain,
831 # HTTP (async)
832 "http_request": http_request,
833 "http_get": http_get,
834 "http_post": http_post,
835 # Code detection
836 "detect_code": detect_code,
837 "detect_code_languages": detect_code_languages,
838 "contains_code_language": contains_code_language,
839 # Text utilities
840 "contains": contains,
841 "contains_any": contains_any,
842 "contains_all": contains_all,
843 "word_count": word_count,
844 "char_count": char_count,
845 "lower": lower,
846 "upper": upper,
847 "trim": trim,
848 # Python builtins (safe subset)
849 "len": len,
850 "str": str,
851 "int": int,
852 "float": float,
853 "bool": bool,
854 "list": list,
855 "dict": dict,
856 "True": True,
857 "False": False,
858 "None": None,
859 }