railtracks.guardrails.llm.concrete

  1from __future__ import annotations
  2
  3from typing import (
  4    Any,
  5    Awaitable,
  6    Callable,
  7    cast,
  8)
  9
 10from pydantic import BaseModel
 11
 12from railtracks.events.middleware import (
 13    MiddlewareGuardInputFailureEvent,
 14    MiddlewareGuardInputInvocationEvent,
 15    MiddlewareGuardInputResponseEvent,
 16    MiddlewareGuardOutputFailureEvent,
 17    MiddlewareGuardOutputInvocationEvent,
 18    MiddlewareGuardOutputResponseEvent,
 19)
 20from railtracks.events.send import emit
 21from railtracks.guardrails.llm.llm_guard import BaseLLMGuardrail
 22from railtracks.llm.history import MessageHistory
 23from railtracks.llm.message import AssistantMessage, Message
 24from railtracks.llm.response import Response
 25from railtracks.llm.tools.tool import Tool
 26from railtracks.utils.logging.create import get_rt_logger
 27
 28from ..core.decision import GuardrailDecision
 29from ..core.event import LLMGuardrailEvent, LLMGuardrailPhase
 30
 31logger = get_rt_logger("guardrails")
 32
 33
 34class InputGuard(BaseLLMGuardrail[MessageHistory]):
 35    """Base for guardrails that run on LLM input (e.g. prompt / message history).
 36
 37    Attributes:
 38        phase: Always :attr:`LLMGuardrailPhase.INPUT`.
 39    """
 40
 41    phase = LLMGuardrailPhase.INPUT
 42
 43    async def _middleware_fn(
 44        self,
 45        call: Callable[
 46            [MessageHistory, type[BaseModel] | None, list[Tool] | None],
 47            Awaitable[Response],
 48        ],
 49        message_history: MessageHistory,
 50        schema: type[BaseModel] | None,
 51        tools: list[Tool] | None,
 52    ):
 53        """Run this guard on the message history, then call onward with the result."""
 54        message_history, schema, tools = await self._input_wrapper(
 55            message_history, schema, tools
 56        )
 57        return await call(message_history, schema, tools)
 58
 59    async def _input_wrapper(
 60        self,
 61        message_history: MessageHistory,
 62        schema: type[BaseModel] | None,
 63        tools: list[Tool] | None,
 64    ):
 65        """Build the input event, run this guard, and raise if it blocks."""
 66        node_uuid, run_id = self._node_metadata()
 67        event = LLMGuardrailEvent(
 68            phase=LLMGuardrailPhase.INPUT,
 69            messages=message_history,
 70            node_uuid=node_uuid,
 71            run_id=run_id,
 72        )
 73        input_event = MiddlewareGuardInputInvocationEvent(
 74            message_history=message_history,
 75        )
 76
 77        await emit(input_event)
 78
 79        try:
 80            new_messages, traces, decision = await self.run(
 81                event=event, value=message_history
 82            )
 83        except Exception as e:
 84            failure_event = MiddlewareGuardInputFailureEvent.from_exception(e)
 85            await emit(failure_event)
 86            raise e
 87
 88        result_event = MiddlewareGuardInputResponseEvent(
 89            decision=decision,
 90            message_history=new_messages,
 91        )
 92
 93        await emit(result_event)
 94
 95        self._raise_if_blocked(decision, traces)
 96
 97        return new_messages, schema, tools
 98
 99    def convert(
100        self,
101        value: str | Any | MessageHistory | LLMGuardrailEvent,
102        /,
103    ):
104        """Run this guard without building an :class:`LLMGuardrailEvent` by hand.
105
106        Args:
107            value: A :class:`LLMGuardrailEvent` (passed through), a ``str`` (treated
108                as a single user message), a :class:`~railtracks.llm.message.Message`,
109                or a :class:`~railtracks.llm.history.MessageHistory`.
110
111        Returns:
112            The :class:`GuardrailDecision` from :meth:`__call__`.
113
114        Raises:
115            TypeError: If ``value`` is not a ``str``, ``Message``, ``MessageHistory``,
116                or :class:`LLMGuardrailEvent`.
117        """
118        if isinstance(value, LLMGuardrailEvent):
119            return value
120
121        messages = self._coerce_to_message_history(value)
122        event = LLMGuardrailEvent(
123            phase=LLMGuardrailPhase.INPUT,
124            messages=messages,
125        )
126        return event
127
128    def _extract_transform_value(self, decision: GuardrailDecision) -> Any:
129        """Return the replacement messages from a TRANSFORM decision."""
130        if decision.messages is None:
131            raise ValueError(
132                "Input guardrail returned TRANSFORM without decision.messages."
133            )
134        return decision.messages
135
136    def _sync_event_after_transform(
137        self,
138        event: LLMGuardrailEvent,
139        value: MessageHistory,
140    ) -> LLMGuardrailEvent:
141        """Return a copy of event with its messages replaced."""
142        return event.model_copy(update={"messages": value})
143
144
145def _is_intermediate_tool_call(response: Response) -> bool:
146    """A model round-trip is *intermediate* (not the final reply) when it requests tools.
147
148    Mirrors ``process_message`` in ``llm_helpers`` (tool calls present => "Tool").
149    """
150    return len(response.message.tool_calls) > 0
151
152
153class OutputGuard(BaseLLMGuardrail[Message]):
154    """Base for guardrails that run on LLM output (e.g. model response).
155
156    Inspect ``event.output_message`` for the assistant message produced this turn.
157    ``event.messages`` is conversation context and may not yet include that reply.
158
159    Intermediate tool-call turns pass through untouched, so output rails fire only
160    on the final reply.
161
162    Attributes:
163        phase: Always :attr:`LLMGuardrailPhase.OUTPUT`.
164    """
165
166    phase = LLMGuardrailPhase.OUTPUT
167
168    async def _middleware_fn(
169        self,
170        call: Callable[
171            [MessageHistory, type[BaseModel] | None, list[Tool] | None],
172            Awaitable[Response],
173        ],
174        message_history: MessageHistory,
175        schema: type[BaseModel] | None,
176        tools: list[Tool] | None,
177    ):
178        """Call onward for the response, then run this guard on the final reply.
179
180        Responses that request tools are intermediate steps of the tool-calling
181        loop and pass through unguarded.
182        """
183        result = await call(message_history, schema, tools)
184        if _is_intermediate_tool_call(result):
185            return result
186
187        return await self._output_wrapper(result=result)
188
189    async def _output_wrapper(
190        self,
191        result: Response,
192    ):
193        """Build the output event, run this guard, and rebuild the response if the message changed."""
194        node_uuid, run_id = self._node_metadata()
195        event = LLMGuardrailEvent(
196            phase=LLMGuardrailPhase.OUTPUT,
197            messages=MessageHistory([]),
198            output_message=result.message,
199            node_uuid=node_uuid,
200            run_id=run_id,
201        )
202
203        input_event = MiddlewareGuardOutputInvocationEvent(
204            response=result.message,
205        )
206
207        await emit(input_event)
208        try:
209            new_message, traces, decision = await self.run(
210                event=event, value=result.message
211            )
212        except Exception as e:
213            failure_event = MiddlewareGuardOutputFailureEvent.from_exception(e)
214            await emit(failure_event)
215            raise e
216
217        output_event = MiddlewareGuardOutputResponseEvent(
218            decision=decision,
219            response=new_message,
220        )
221        await emit(output_event)
222
223        self._raise_if_blocked(decision, traces)
224
225        if new_message is result.message:
226            return result
227
228        return Response(message=new_message, message_info=result.message_info)
229
230    def convert(self, output: str | Any | MessageHistory | LLMGuardrailEvent, /):
231        """Run this guard without building an :class:`LLMGuardrailEvent` by hand.
232
233        Args:
234            output: A :class:`LLMGuardrailEvent` (passed through), a ``str`` (becomes
235                the assistant message with empty prior history), a
236                :class:`~railtracks.llm.message.Message`, or a non-empty
237                :class:`~railtracks.llm.history.MessageHistory` (last message is the
238                output under test; earlier entries become ``event.messages``).
239
240        Returns:
241            The :class:`GuardrailDecision` from :meth:`__call__`.
242
243        Raises:
244            ValueError: If ``output`` is an empty :class:`~railtracks.llm.history.MessageHistory`.
245            TypeError: If ``output`` is not a ``str``, ``Message``, ``MessageHistory``,
246                or :class:`LLMGuardrailEvent`.
247        """
248        if isinstance(output, LLMGuardrailEvent):
249            return output
250
251        if isinstance(output, str):
252            output_message = AssistantMessage(output)
253            messages = MessageHistory()
254        elif isinstance(output, Message):
255            output_message = output
256            messages = MessageHistory()
257        elif isinstance(output, MessageHistory):
258            if not output:
259                raise ValueError("Cannot decide with an empty MessageHistory.")
260            output_message = output[-1]
261            messages = MessageHistory(output[:-1])
262        else:
263            raise TypeError(
264                f"Expected str, Message, MessageHistory, or LLMGuardrailEvent, "
265                f"got {type(output).__name__}"
266            )
267
268        event = LLMGuardrailEvent(
269            phase=LLMGuardrailPhase.OUTPUT,
270            messages=messages,
271            output_message=output_message,
272        )
273        return event
274
275    def _extract_transform_value(self, decision: GuardrailDecision) -> Message:
276        """Return the replacement output message from a TRANSFORM decision."""
277        if decision.output_message is None:
278            raise ValueError(
279                "Output guardrail returned TRANSFORM without decision.output_message."
280            )
281        return decision.output_message
282
283    def _sync_event_after_transform(
284        self,
285        event: LLMGuardrailEvent,
286        value: Message,
287    ) -> LLMGuardrailEvent:
288        """Return a copy of event with its output message replaced."""
289        return event.model_copy(update={"output_message": cast(Message, value)})
logger = <RTContextLoggingAdapter RT.guardrails (WARNING)>
 35class InputGuard(BaseLLMGuardrail[MessageHistory]):
 36    """Base for guardrails that run on LLM input (e.g. prompt / message history).
 37
 38    Attributes:
 39        phase: Always :attr:`LLMGuardrailPhase.INPUT`.
 40    """
 41
 42    phase = LLMGuardrailPhase.INPUT
 43
 44    async def _middleware_fn(
 45        self,
 46        call: Callable[
 47            [MessageHistory, type[BaseModel] | None, list[Tool] | None],
 48            Awaitable[Response],
 49        ],
 50        message_history: MessageHistory,
 51        schema: type[BaseModel] | None,
 52        tools: list[Tool] | None,
 53    ):
 54        """Run this guard on the message history, then call onward with the result."""
 55        message_history, schema, tools = await self._input_wrapper(
 56            message_history, schema, tools
 57        )
 58        return await call(message_history, schema, tools)
 59
 60    async def _input_wrapper(
 61        self,
 62        message_history: MessageHistory,
 63        schema: type[BaseModel] | None,
 64        tools: list[Tool] | None,
 65    ):
 66        """Build the input event, run this guard, and raise if it blocks."""
 67        node_uuid, run_id = self._node_metadata()
 68        event = LLMGuardrailEvent(
 69            phase=LLMGuardrailPhase.INPUT,
 70            messages=message_history,
 71            node_uuid=node_uuid,
 72            run_id=run_id,
 73        )
 74        input_event = MiddlewareGuardInputInvocationEvent(
 75            message_history=message_history,
 76        )
 77
 78        await emit(input_event)
 79
 80        try:
 81            new_messages, traces, decision = await self.run(
 82                event=event, value=message_history
 83            )
 84        except Exception as e:
 85            failure_event = MiddlewareGuardInputFailureEvent.from_exception(e)
 86            await emit(failure_event)
 87            raise e
 88
 89        result_event = MiddlewareGuardInputResponseEvent(
 90            decision=decision,
 91            message_history=new_messages,
 92        )
 93
 94        await emit(result_event)
 95
 96        self._raise_if_blocked(decision, traces)
 97
 98        return new_messages, schema, tools
 99
100    def convert(
101        self,
102        value: str | Any | MessageHistory | LLMGuardrailEvent,
103        /,
104    ):
105        """Run this guard without building an :class:`LLMGuardrailEvent` by hand.
106
107        Args:
108            value: A :class:`LLMGuardrailEvent` (passed through), a ``str`` (treated
109                as a single user message), a :class:`~railtracks.llm.message.Message`,
110                or a :class:`~railtracks.llm.history.MessageHistory`.
111
112        Returns:
113            The :class:`GuardrailDecision` from :meth:`__call__`.
114
115        Raises:
116            TypeError: If ``value`` is not a ``str``, ``Message``, ``MessageHistory``,
117                or :class:`LLMGuardrailEvent`.
118        """
119        if isinstance(value, LLMGuardrailEvent):
120            return value
121
122        messages = self._coerce_to_message_history(value)
123        event = LLMGuardrailEvent(
124            phase=LLMGuardrailPhase.INPUT,
125            messages=messages,
126        )
127        return event
128
129    def _extract_transform_value(self, decision: GuardrailDecision) -> Any:
130        """Return the replacement messages from a TRANSFORM decision."""
131        if decision.messages is None:
132            raise ValueError(
133                "Input guardrail returned TRANSFORM without decision.messages."
134            )
135        return decision.messages
136
137    def _sync_event_after_transform(
138        self,
139        event: LLMGuardrailEvent,
140        value: MessageHistory,
141    ) -> LLMGuardrailEvent:
142        """Return a copy of event with its messages replaced."""
143        return event.model_copy(update={"messages": value})

Base for guardrails that run on LLM input (e.g. prompt / message history).

Attributes:
  • phase: Always LLMGuardrailPhase.INPUT.
phase = <LLMGuardrailPhase.INPUT: 'llm_input'>
def convert( self, value: Union[str, Any, railtracks.llm.MessageHistory, railtracks.guardrails.LLMGuardrailEvent], /):
100    def convert(
101        self,
102        value: str | Any | MessageHistory | LLMGuardrailEvent,
103        /,
104    ):
105        """Run this guard without building an :class:`LLMGuardrailEvent` by hand.
106
107        Args:
108            value: A :class:`LLMGuardrailEvent` (passed through), a ``str`` (treated
109                as a single user message), a :class:`~railtracks.llm.message.Message`,
110                or a :class:`~railtracks.llm.history.MessageHistory`.
111
112        Returns:
113            The :class:`GuardrailDecision` from :meth:`__call__`.
114
115        Raises:
116            TypeError: If ``value`` is not a ``str``, ``Message``, ``MessageHistory``,
117                or :class:`LLMGuardrailEvent`.
118        """
119        if isinstance(value, LLMGuardrailEvent):
120            return value
121
122        messages = self._coerce_to_message_history(value)
123        event = LLMGuardrailEvent(
124            phase=LLMGuardrailPhase.INPUT,
125            messages=messages,
126        )
127        return event

Run this guard without building an LLMGuardrailEvent by hand.

Arguments:
Returns:

The GuardrailDecision from __call__().

Raises:
  • TypeError: If value is not a str, Message, MessageHistory, or LLMGuardrailEvent.
154class OutputGuard(BaseLLMGuardrail[Message]):
155    """Base for guardrails that run on LLM output (e.g. model response).
156
157    Inspect ``event.output_message`` for the assistant message produced this turn.
158    ``event.messages`` is conversation context and may not yet include that reply.
159
160    Intermediate tool-call turns pass through untouched, so output rails fire only
161    on the final reply.
162
163    Attributes:
164        phase: Always :attr:`LLMGuardrailPhase.OUTPUT`.
165    """
166
167    phase = LLMGuardrailPhase.OUTPUT
168
169    async def _middleware_fn(
170        self,
171        call: Callable[
172            [MessageHistory, type[BaseModel] | None, list[Tool] | None],
173            Awaitable[Response],
174        ],
175        message_history: MessageHistory,
176        schema: type[BaseModel] | None,
177        tools: list[Tool] | None,
178    ):
179        """Call onward for the response, then run this guard on the final reply.
180
181        Responses that request tools are intermediate steps of the tool-calling
182        loop and pass through unguarded.
183        """
184        result = await call(message_history, schema, tools)
185        if _is_intermediate_tool_call(result):
186            return result
187
188        return await self._output_wrapper(result=result)
189
190    async def _output_wrapper(
191        self,
192        result: Response,
193    ):
194        """Build the output event, run this guard, and rebuild the response if the message changed."""
195        node_uuid, run_id = self._node_metadata()
196        event = LLMGuardrailEvent(
197            phase=LLMGuardrailPhase.OUTPUT,
198            messages=MessageHistory([]),
199            output_message=result.message,
200            node_uuid=node_uuid,
201            run_id=run_id,
202        )
203
204        input_event = MiddlewareGuardOutputInvocationEvent(
205            response=result.message,
206        )
207
208        await emit(input_event)
209        try:
210            new_message, traces, decision = await self.run(
211                event=event, value=result.message
212            )
213        except Exception as e:
214            failure_event = MiddlewareGuardOutputFailureEvent.from_exception(e)
215            await emit(failure_event)
216            raise e
217
218        output_event = MiddlewareGuardOutputResponseEvent(
219            decision=decision,
220            response=new_message,
221        )
222        await emit(output_event)
223
224        self._raise_if_blocked(decision, traces)
225
226        if new_message is result.message:
227            return result
228
229        return Response(message=new_message, message_info=result.message_info)
230
231    def convert(self, output: str | Any | MessageHistory | LLMGuardrailEvent, /):
232        """Run this guard without building an :class:`LLMGuardrailEvent` by hand.
233
234        Args:
235            output: A :class:`LLMGuardrailEvent` (passed through), a ``str`` (becomes
236                the assistant message with empty prior history), a
237                :class:`~railtracks.llm.message.Message`, or a non-empty
238                :class:`~railtracks.llm.history.MessageHistory` (last message is the
239                output under test; earlier entries become ``event.messages``).
240
241        Returns:
242            The :class:`GuardrailDecision` from :meth:`__call__`.
243
244        Raises:
245            ValueError: If ``output`` is an empty :class:`~railtracks.llm.history.MessageHistory`.
246            TypeError: If ``output`` is not a ``str``, ``Message``, ``MessageHistory``,
247                or :class:`LLMGuardrailEvent`.
248        """
249        if isinstance(output, LLMGuardrailEvent):
250            return output
251
252        if isinstance(output, str):
253            output_message = AssistantMessage(output)
254            messages = MessageHistory()
255        elif isinstance(output, Message):
256            output_message = output
257            messages = MessageHistory()
258        elif isinstance(output, MessageHistory):
259            if not output:
260                raise ValueError("Cannot decide with an empty MessageHistory.")
261            output_message = output[-1]
262            messages = MessageHistory(output[:-1])
263        else:
264            raise TypeError(
265                f"Expected str, Message, MessageHistory, or LLMGuardrailEvent, "
266                f"got {type(output).__name__}"
267            )
268
269        event = LLMGuardrailEvent(
270            phase=LLMGuardrailPhase.OUTPUT,
271            messages=messages,
272            output_message=output_message,
273        )
274        return event
275
276    def _extract_transform_value(self, decision: GuardrailDecision) -> Message:
277        """Return the replacement output message from a TRANSFORM decision."""
278        if decision.output_message is None:
279            raise ValueError(
280                "Output guardrail returned TRANSFORM without decision.output_message."
281            )
282        return decision.output_message
283
284    def _sync_event_after_transform(
285        self,
286        event: LLMGuardrailEvent,
287        value: Message,
288    ) -> LLMGuardrailEvent:
289        """Return a copy of event with its output message replaced."""
290        return event.model_copy(update={"output_message": cast(Message, value)})

Base for guardrails that run on LLM output (e.g. model response).

Inspect event.output_message for the assistant message produced this turn. event.messages is conversation context and may not yet include that reply.

Intermediate tool-call turns pass through untouched, so output rails fire only on the final reply.

Attributes:
  • phase: Always LLMGuardrailPhase.OUTPUT.
phase = <LLMGuardrailPhase.OUTPUT: 'llm_output'>
def convert( self, output: Union[str, Any, railtracks.llm.MessageHistory, railtracks.guardrails.LLMGuardrailEvent], /):
231    def convert(self, output: str | Any | MessageHistory | LLMGuardrailEvent, /):
232        """Run this guard without building an :class:`LLMGuardrailEvent` by hand.
233
234        Args:
235            output: A :class:`LLMGuardrailEvent` (passed through), a ``str`` (becomes
236                the assistant message with empty prior history), a
237                :class:`~railtracks.llm.message.Message`, or a non-empty
238                :class:`~railtracks.llm.history.MessageHistory` (last message is the
239                output under test; earlier entries become ``event.messages``).
240
241        Returns:
242            The :class:`GuardrailDecision` from :meth:`__call__`.
243
244        Raises:
245            ValueError: If ``output`` is an empty :class:`~railtracks.llm.history.MessageHistory`.
246            TypeError: If ``output`` is not a ``str``, ``Message``, ``MessageHistory``,
247                or :class:`LLMGuardrailEvent`.
248        """
249        if isinstance(output, LLMGuardrailEvent):
250            return output
251
252        if isinstance(output, str):
253            output_message = AssistantMessage(output)
254            messages = MessageHistory()
255        elif isinstance(output, Message):
256            output_message = output
257            messages = MessageHistory()
258        elif isinstance(output, MessageHistory):
259            if not output:
260                raise ValueError("Cannot decide with an empty MessageHistory.")
261            output_message = output[-1]
262            messages = MessageHistory(output[:-1])
263        else:
264            raise TypeError(
265                f"Expected str, Message, MessageHistory, or LLMGuardrailEvent, "
266                f"got {type(output).__name__}"
267            )
268
269        event = LLMGuardrailEvent(
270            phase=LLMGuardrailPhase.OUTPUT,
271            messages=messages,
272            output_message=output_message,
273        )
274        return event

Run this guard without building an LLMGuardrailEvent by hand.

Arguments:
Returns:

The GuardrailDecision from __call__().

Raises: