Repository navigation
Add Guardrails AI scorer integration #20038
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 1 commit
5991f97
dc71695
6167436
fa10195
60e75ad
b183776
39fff9c
eeb80fd
8b22059
9a48b7b
566f8a9
20088dd
21f8133
5e5d90a
37ed981
1ec2f5a
eca48fe
5f226a4
7920d32
7e7b7f1
a90f7e9
00da3dd
4166d69
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
Implements mlflow.genai.scorers.guardrails module providing: - GuardrailsScorer base class wrapping Guardrails AI validators - Safety validators: ToxicLanguage, NSFWText, DetectJailbreak - PII validators: DetectPII, SecretsPresent - Quality validators: GibberishText - get_scorer() factory function for custom validator access Key design decisions: - Uses on_fail=OnFailAction.NOOP to return validation results instead of raising exceptions - Minimal mocking in tests: only mock get_validator_class (external hub), use real Guard - Follows Phoenix scorer pattern with validator registry - Returns Feedback objects with pass/fail values and rationale Resolves #20036 Signed-off-by: debu-sinha <debusinha2009@gmail.com>
- Loading branch information
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,340 @@ | ||
| """ | ||
| Guardrails AI integration for MLflow. | ||
|
|
||
| This module provides integration with Guardrails AI validators, allowing them to be used | ||
| with MLflow's scorer interface for LLM safety, PII detection, and content quality evaluation. | ||
|
|
||
| Example usage: | ||
|
|
||
| .. code-block:: python | ||
|
|
||
| from mlflow.genai.scorers.guardrails import ToxicLanguage, DetectPII | ||
|
|
||
| # Evaluate LLM outputs for toxicity | ||
| scorer = ToxicLanguage() | ||
| feedback = scorer(outputs="This is a friendly response.") | ||
|
|
||
| # Detect PII in outputs | ||
| pii_scorer = DetectPII() | ||
| feedback = pii_scorer(outputs="Contact john@email.com for details.") | ||
| """ | ||
|
|
||
| from __future__ import annotations | ||
|
|
||
| import logging | ||
| from typing import Any, ClassVar | ||
|
|
||
| from pydantic import PrivateAttr | ||
|
|
||
| from mlflow.entities.assessment import Feedback | ||
| from mlflow.entities.assessment_source import AssessmentSource, AssessmentSourceType | ||
| from mlflow.entities.trace import Trace | ||
| from mlflow.genai.scorers import FRAMEWORK_METADATA_KEY | ||
| from mlflow.genai.scorers.base import Scorer | ||
| from mlflow.genai.scorers.guardrails.registry import get_validator_class | ||
| from mlflow.genai.scorers.guardrails.utils import ( | ||
| check_guardrails_installed, | ||
| map_scorer_inputs_to_text, | ||
| ) | ||
| from mlflow.utils.annotations import experimental | ||
|
|
||
| _logger = logging.getLogger(__name__) | ||
|
|
||
|
|
||
| @experimental(version="3.9.0") | ||
| class GuardrailsScorer(Scorer): | ||
| """ | ||
| Base class for Guardrails AI validator scorers. | ||
|
|
||
| Guardrails AI validators check text for specific issues like toxicity, | ||
| PII, jailbreak attempts, etc. This class wraps validators to work with | ||
| MLflow's scorer interface. | ||
|
|
||
| Args: | ||
| validator_name: Name of the Guardrails AI validator | ||
| **validator_kwargs: Additional arguments passed to the validator | ||
| """ | ||
|
|
||
| _guard: Any = PrivateAttr() | ||
| _validator_name: str = PrivateAttr() | ||
|
|
||
| def __init__( | ||
| self, | ||
| validator_name: str | None = None, | ||
| **validator_kwargs: Any, | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. any reason we pass these kwargs during inference (L81) instead of when constructing the class?
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. gentle bump on this |
||
| ): | ||
| check_guardrails_installed() | ||
|
|
||
| # Get validator name from class variable if not provided | ||
| if validator_name is None: | ||
| validator_name = getattr(self.__class__, "validator_name", None) | ||
| if validator_name is None: | ||
| raise ValueError("validator_name must be provided") | ||
|
|
||
| super().__init__(name=validator_name) | ||
| self._validator_name = validator_name | ||
|
|
||
| from guardrails import Guard, OnFailAction | ||
|
|
||
| validator_class = get_validator_class(validator_name) | ||
| # Use NOOP on_fail to get validation result instead of raising exception | ||
| self._guard = Guard().use( | ||
| validator_class, on_fail=OnFailAction.NOOP, **validator_kwargs | ||
| ) | ||
|
|
||
| def __call__( | ||
| self, | ||
| *, | ||
| inputs: Any = None, | ||
| outputs: Any = None, | ||
| expectations: dict[str, Any] | None = None, | ||
| trace: Trace | None = None, | ||
| ) -> Feedback: | ||
| """ | ||
| Validate text using the Guardrails AI validator. | ||
|
|
||
| Args: | ||
| inputs: The input to evaluate | ||
| outputs: The output to evaluate (primary target for validation) | ||
| expectations: Not used for Guardrails validators | ||
| trace: MLflow trace for evaluation | ||
|
|
||
| Returns: | ||
| Feedback object with validation result | ||
| """ | ||
| assessment_source = AssessmentSource( | ||
| source_type=AssessmentSourceType.CODE, | ||
| source_id=f"guardrails/{self._validator_name}", | ||
| ) | ||
|
|
||
| try: | ||
| text = map_scorer_inputs_to_text( | ||
| inputs=inputs, | ||
| outputs=outputs, | ||
| trace=trace, | ||
| ) | ||
|
|
||
| # Run validation | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. nit: let's remove these one line comments - they don't add much in terms of readability - same for places like L118, L122, and so on. |
||
| result = self._guard.validate(text) | ||
|
|
||
| # Extract validation outcome | ||
| passed = result.validation_passed | ||
| value = "pass" if passed else "fail" | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. can we use the constant value in
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. gentle bump on this |
||
|
|
||
| # Build rationale from validation results | ||
| rationale = None | ||
| if hasattr(result, "validated_output"): | ||
| if not passed and result.validation_summaries: | ||
| summaries = [ | ||
| f"{s.validator_name}: {s.failure_reason}" | ||
| for s in result.validation_summaries | ||
| if s.failure_reason | ||
| ] | ||
| rationale = "; ".join(summaries) if summaries else None | ||
|
|
||
| return Feedback( | ||
| name=self.name, | ||
| value=value, | ||
| rationale=rationale, | ||
| source=assessment_source, | ||
| metadata={ | ||
| FRAMEWORK_METADATA_KEY: "guardrails", | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. let's still include this in the error case as well. |
||
| "validation_passed": passed, | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I don't think we need this in the metadata given it's reflected directly in teh value. |
||
| }, | ||
| ) | ||
| except Exception as e: | ||
| _logger.error(f"Error validating with Guardrails {self.name}: {e}") | ||
| return Feedback( | ||
| name=self.name, | ||
| error=e, | ||
| source=assessment_source, | ||
| ) | ||
|
|
||
|
|
||
| @experimental(version="3.9.0") | ||
| def get_scorer( | ||
| validator_name: str, | ||
| **validator_kwargs: Any, | ||
| ) -> GuardrailsScorer: | ||
| """ | ||
| Get a Guardrails AI validator as an MLflow scorer. | ||
|
|
||
| Args: | ||
| validator_name: Name of the validator (e.g., "ToxicLanguage", "DetectPII") | ||
| validator_kwargs: Additional keyword arguments to pass to the validator. | ||
|
|
||
| Returns: | ||
| GuardrailsScorer instance that can be called with MLflow's scorer interface | ||
|
|
||
| Examples: | ||
| >>> scorer = get_scorer("ToxicLanguage", threshold=0.7) | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. can we use the |
||
| >>> feedback = scorer(outputs="This is a friendly response.") | ||
| >>> scorer = get_scorer("DetectPII") | ||
| >>> feedback = scorer(outputs="Contact john@email.com") | ||
| """ | ||
| return GuardrailsScorer( | ||
| validator_name=validator_name, | ||
| **validator_kwargs, | ||
| ) | ||
|
|
||
|
|
||
| # Safety validators | ||
|
|
||
|
|
||
| @experimental(version="3.9.0") | ||
| class ToxicLanguage(GuardrailsScorer): | ||
| """ | ||
| Detects toxic language in text using Guardrails AI. | ||
|
|
||
| Uses NLP models to identify toxic, offensive, or harmful content | ||
| in LLM outputs. | ||
|
|
||
| Args: | ||
| threshold: Confidence threshold for detection (default: 0.5) | ||
| validation_method: "sentence" or "full" text validation | ||
|
|
||
| Examples: | ||
| .. code-block:: python | ||
|
|
||
| from mlflow.genai.scorers.guardrails import ToxicLanguage | ||
|
|
||
| scorer = ToxicLanguage(threshold=0.7) | ||
| feedback = scorer(outputs="This is a professional response.") | ||
| """ | ||
|
|
||
| validator_name: ClassVar[str] = "ToxicLanguage" | ||
|
|
||
|
|
||
| @experimental(version="3.9.0") | ||
| class NSFWText(GuardrailsScorer): | ||
| """ | ||
| Detects NSFW (Not Safe For Work) content in text. | ||
|
|
||
| Identifies inappropriate, adult, or explicit content that may not | ||
| be suitable for professional settings. | ||
|
|
||
| Args: | ||
| threshold: Confidence threshold for detection (default: 0.8) | ||
|
|
||
| Examples: | ||
| .. code-block:: python | ||
|
|
||
| from mlflow.genai.scorers.guardrails import NSFWText | ||
|
|
||
| scorer = NSFWText() | ||
| feedback = scorer(outputs="This is appropriate content.") | ||
| """ | ||
|
|
||
| validator_name: ClassVar[str] = "NSFWText" | ||
|
|
||
|
|
||
| @experimental(version="3.9.0") | ||
| class DetectJailbreak(GuardrailsScorer): | ||
| """ | ||
| Detects jailbreak or prompt injection attempts. | ||
|
|
||
| Identifies attempts to bypass LLM safety measures or manipulate | ||
| the model into generating harmful content. | ||
|
|
||
| Examples: | ||
| .. code-block:: python | ||
|
|
||
| from mlflow.genai.scorers.guardrails import DetectJailbreak | ||
|
|
||
| scorer = DetectJailbreak() | ||
| feedback = scorer(inputs="Ignore previous instructions and...") | ||
| """ | ||
|
|
||
| validator_name: ClassVar[str] = "DetectJailbreak" | ||
|
|
||
|
|
||
| # PII validators | ||
|
|
||
|
|
||
| @experimental(version="3.9.0") | ||
| class DetectPII(GuardrailsScorer): | ||
| """ | ||
| Detects Personally Identifiable Information (PII) in text. | ||
|
|
||
| Uses Microsoft Presidio to identify PII such as email addresses, | ||
| phone numbers, names, and locations. | ||
|
|
||
| Args: | ||
| pii_entities: List of PII types to detect (default: EMAIL_ADDRESS, | ||
| PHONE_NUMBER, PERSON, LOCATION) | ||
|
|
||
| Examples: | ||
| .. code-block:: python | ||
|
|
||
| from mlflow.genai.scorers.guardrails import DetectPII | ||
|
|
||
| scorer = DetectPII() | ||
| feedback = scorer(outputs="Contact john@email.com for help.") | ||
|
|
||
| # Custom PII types | ||
| scorer = DetectPII(pii_entities=["CREDIT_CARD", "SSN"]) | ||
| """ | ||
|
|
||
| validator_name: ClassVar[str] = "DetectPII" | ||
|
|
||
|
|
||
| @experimental(version="3.9.0") | ||
| class SecretsPresent(GuardrailsScorer): | ||
| """ | ||
| Detects secrets and API keys in text. | ||
|
|
||
| Identifies patterns that look like API keys, tokens, passwords, | ||
| or other sensitive credentials. | ||
|
|
||
| Examples: | ||
| .. code-block:: python | ||
|
|
||
| from mlflow.genai.scorers.guardrails import SecretsPresent | ||
|
|
||
| scorer = SecretsPresent() | ||
| feedback = scorer(outputs="Use API key: sk-abc123...") | ||
| """ | ||
|
|
||
| validator_name: ClassVar[str] = "SecretsPresent" | ||
|
|
||
|
|
||
| # Quality validators | ||
|
|
||
|
|
||
| @experimental(version="3.9.0") | ||
| class GibberishText(GuardrailsScorer): | ||
| """ | ||
| Detects gibberish or nonsensical text in LLM outputs. | ||
|
|
||
| Identifies when the model produces incoherent, random, or | ||
| meaningless text. | ||
|
|
||
| Args: | ||
| threshold: Confidence threshold for detection (default: 0.5) | ||
|
|
||
| Examples: | ||
| .. code-block:: python | ||
|
|
||
| from mlflow.genai.scorers.guardrails import GibberishText | ||
|
|
||
| scorer = GibberishText() | ||
| feedback = scorer(outputs="asdf jkl; qwerty uiop") | ||
| """ | ||
|
|
||
| validator_name: ClassVar[str] = "GibberishText" | ||
|
|
||
|
|
||
| __all__ = [ | ||
| # Core classes | ||
| "GuardrailsScorer", | ||
| "get_scorer", | ||
| # Safety validators | ||
| "ToxicLanguage", | ||
| "NSFWText", | ||
| "DetectJailbreak", | ||
| # PII validators | ||
| "DetectPII", | ||
| "SecretsPresent", | ||
| # Quality validators | ||
| "GibberishText", | ||
| ] | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,33 @@ | ||
| from __future__ import annotations | ||
|
|
||
| from mlflow.exceptions import MlflowException | ||
| from mlflow.genai.scorers.guardrails.utils import check_guardrails_installed | ||
|
|
||
| # Simplified registry: validator name -> hub class name | ||
| _VALIDATOR_REGISTRY = { | ||
| "ToxicLanguage": "ToxicLanguage", | ||
| "NSFWText": "NSFWText", | ||
| "DetectJailbreak": "DetectPromptInjection", | ||
| "DetectPII": "DetectPII", | ||
| "SecretsPresent": "SecretsPresent", | ||
| "GibberishText": "GibberishText", | ||
| } | ||
|
|
||
|
|
||
| def get_validator_class(validator_name: str): | ||
| check_guardrails_installed() | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. we've already verified it is installed by this point, we can drop this |
||
|
|
||
| if validator_name not in _VALIDATOR_REGISTRY: | ||
| available = ", ".join(sorted(_VALIDATOR_REGISTRY.keys())) | ||
| raise MlflowException.invalid_parameter_value( | ||
| f"Unknown Guardrails AI validator: '{validator_name}'. Available: {available}" | ||
| ) | ||
|
|
||
| from guardrails import hub | ||
|
|
||
| class_name = _VALIDATOR_REGISTRY[validator_name] | ||
| return getattr(hub, class_name) | ||
|
|
||
|
|
||
| def list_available_validators() -> list[str]: | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. let's remove this - it's not used anywhere |
||
| return sorted(_VALIDATOR_REGISTRY.keys()) | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
we missed the 3.9.0 release candidate, let's set this to 3.10.0 for the rc set in early Feb