Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion nemoguardrails/server/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
8 changes: 7 additions & 1 deletion nemoguardrails/server/schemas/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand Down
44 changes: 43 additions & 1 deletion tests/server/test_guardrail_checks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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