diff --git a/nemoguardrails/server/api.py b/nemoguardrails/server/api.py index a9c68c2800..d339d61b64 100644 --- a/nemoguardrails/server/api.py +++ b/nemoguardrails/server/api.py @@ -797,7 +797,7 @@ async def guardrail_check(body: GuardrailCheckRequest, request: Request): if body.guardrails.context: messages.insert(0, {"role": "context", "content": body.guardrails.context}) - result = await llm_rails.check_async(messages=messages) + result = await llm_rails.check_async(messages=messages, rail_types=body.guardrails.rail_types) return GuardrailCheckResponse( status=_map_rail_status(result.status), diff --git a/nemoguardrails/server/schemas/openai.py b/nemoguardrails/server/schemas/openai.py index be96466a85..09f0741027 100644 --- a/nemoguardrails/server/schemas/openai.py +++ b/nemoguardrails/server/schemas/openai.py @@ -21,7 +21,7 @@ from openai.types.chat.chat_completion import ChatCompletion from pydantic import BaseModel, Field, ValidationInfo, field_validator, model_validator -from nemoguardrails.rails.llm.options import GenerationOptions +from nemoguardrails.rails.llm.options import GenerationOptions, RailType class GuardrailsDataOutput(BaseModel): @@ -134,6 +134,12 @@ class GuardrailsDataInput(BaseModel): default=None, description="State object to continue the interaction.", ) + rail_types: Optional[List[RailType]] = Field( + default=None, + description="Rail types to run (checks endpoint only). " + "When omitted, auto-detected from message roles. " + "Valid values: 'input', 'output'.", + ) @model_validator(mode="before") @classmethod diff --git a/tests/server/test_guardrail_checks.py b/tests/server/test_guardrail_checks.py index 6d7d4c8304..dc6eaaad15 100644 --- a/tests/server/test_guardrail_checks.py +++ b/tests/server/test_guardrail_checks.py @@ -20,7 +20,7 @@ pytest.importorskip("openai", reason="openai is required for server tests") from fastapi.testclient import TestClient -from nemoguardrails.rails.llm.options import RailsResult, RailStatus +from nemoguardrails.rails.llm.options import RailsResult, RailStatus, RailType from nemoguardrails.server import api client = TestClient(api.app) @@ -247,3 +247,45 @@ def test_context_prepended_to_messages(): messages = call_args.kwargs.get("messages") or call_args[0][0] assert messages[0]["role"] == "context" assert messages[0]["content"] == {"topic": "science"} + + +@pytest.mark.parametrize( + "rail_types_input, expected", + [ + (["input"], [RailType.INPUT]), + (["output"], [RailType.OUTPUT]), + (["input", "output"], [RailType.INPUT, RailType.OUTPUT]), + (None, None), + ], +) +def test_rail_types_passed_through(rail_types_input, expected): + result = RailsResult(status=RailStatus.PASSED, content="hi") + mock = _mock_rails(result) + + guardrails = {"config_id": "test"} + if rail_types_input is not None: + guardrails["rail_types"] = rail_types_input + + with patch.object(api, "_get_rails", new_callable=AsyncMock, return_value=mock): + resp = _post( + { + "model": "test", + "messages": [{"role": "user", "content": "hi"}], + "guardrails": guardrails, + } + ) + + assert resp.status_code == 200 + mock.check_async.assert_called_once() + assert mock.check_async.call_args.kwargs["rail_types"] == expected + + +def test_rail_types_invalid_value_returns_422(): + resp = _post( + { + "model": "test", + "messages": [{"role": "user", "content": "hi"}], + "guardrails": {"config_id": "test", "rail_types": ["invalid"]}, + } + ) + assert resp.status_code == 422