Files
AI/참고/guardrails-main/guardrails/async_guard.py

551 lines
20 KiB
Python
Raw Normal View History

2026-05-12 19:40:31 +09:00
from builtins import id as object_id
import contextvars
import inspect
from opentelemetry import context as otel_context
from typing import (
Any,
AsyncIterator,
Awaitable,
Callable,
Dict,
Generic,
List,
Optional,
Sequence,
Union,
cast,
)
from guardrails_ai.types import (
ValidationOutcome as IValidationOutcome,
)
from guardrails import Guard
from guardrails.classes import OT, ValidationOutcome
from guardrails.classes.history import Call
from guardrails.classes.history.call_inputs import CallInputs
from guardrails.classes.output_type import OutputTypes
from guardrails.classes.schema.processed_schema import ProcessedSchema
from guardrails.formatters.base_formatter import BaseFormatter
from guardrails.llm_providers import get_async_llm_ask, model_is_supported_server_side
from guardrails.logger import set_scope
from guardrails.run import AsyncRunner, AsyncStreamRunner
from guardrails.stores.context import get_call_kwarg, set_call_kwargs
from guardrails.hub_telemetry.hub_tracing import async_trace
from guardrails.types.pydantic import ModelOrListOfModels
from guardrails.telemetry import trace_async_guard_execution, wrap_with_otel_context
from guardrails.utils.validator_utils import verify_metadata_requirements
from guardrails.validator_base import Validator
class AsyncGuard(Guard, Generic[OT]):
"""The AsyncGuard class.
This class one of the main entry point for using Guardrails. It is
initialized from one of the following class methods:
- `for_rail`
- `for_rail_string`
- `for_pydantic`
- `for_string`
The `__call__`
method functions as a wrapper around LLM APIs. It takes in an Async LLM
API, and optional prompt parameters, and returns the raw output stream from
the LLM and the validated output stream.
"""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
@classmethod
def _for_rail_schema(
cls,
schema: ProcessedSchema,
rail: str,
*,
name: Optional[str] = None,
description: Optional[str] = None,
):
guard = super()._for_rail_schema(
schema,
rail,
name=name,
description=description,
)
if schema.output_type == OutputTypes.STRING:
return cast(AsyncGuard[str], guard)
elif schema.output_type == OutputTypes.LIST:
return cast(AsyncGuard[List], guard)
else:
return cast(AsyncGuard[Dict], guard)
@classmethod
def for_pydantic(
cls,
output_class: ModelOrListOfModels,
*,
messages: Optional[List[Dict]] = None,
reask_messages: Optional[List[Dict]] = None,
name: Optional[str] = None,
description: Optional[str] = None,
output_formatter: Optional[Union[str, BaseFormatter]] = None,
):
guard = super().for_pydantic(
output_class,
messages=messages,
reask_messages=reask_messages,
name=name,
description=description,
output_formatter=output_formatter,
)
if guard._output_type == OutputTypes.LIST:
return cast(AsyncGuard[List], guard)
else:
return cast(AsyncGuard[Dict], guard)
@classmethod
def for_string(
cls,
validators: Sequence[Validator],
*,
string_description: Optional[str] = None,
messages: Optional[List[Dict]] = None,
reask_messages: Optional[List[Dict]] = None,
name: Optional[str] = None,
description: Optional[str] = None,
):
guard = super().for_string(
validators,
string_description=string_description,
messages=messages,
reask_messages=reask_messages,
name=name,
description=description,
)
return cast(AsyncGuard[str], guard)
@classmethod
def from_dict(cls, obj: Optional[Dict[str, Any]]) -> Optional["AsyncGuard"]:
guard = super().from_dict(obj)
return cast(AsyncGuard, guard)
def use(
self,
*validator_spread: Validator,
validators: List[Validator] = [],
on: str = "output",
) -> "AsyncGuard":
guard = super().use(*validator_spread, validators=validators, on=on)
return cast(AsyncGuard, guard)
async def _execute(
self,
*args,
llm_api: Optional[Callable[..., Awaitable[Any]]] = None,
llm_output: Optional[str] = None,
prompt_params: Optional[Dict] = None,
num_reasks: Optional[int] = None,
messages: Optional[List[Dict]] = None,
metadata: Optional[Dict],
full_schema_reask: Optional[bool] = None,
**kwargs,
) -> Union[
ValidationOutcome[OT],
Awaitable[ValidationOutcome[OT]],
AsyncIterator[ValidationOutcome[OT]],
]:
self._fill_validator_map()
self._fill_validators()
metadata = metadata or {}
if not llm_output and llm_api and not (messages):
raise RuntimeError("'messages' must be provided in order to call an LLM!")
# check if validator requirements are fulfilled
missing_keys = verify_metadata_requirements(metadata, self._validators)
if missing_keys:
raise ValueError(
f"Missing required metadata keys: {', '.join(missing_keys)}"
)
async def __exec(
self: AsyncGuard,
*args,
llm_api: Optional[Callable[..., Awaitable[Any]]],
llm_output: Optional[str] = None,
prompt_params: Optional[Dict] = None,
num_reasks: Optional[int] = None,
messages: Optional[List[Dict]] = None,
metadata: Optional[Dict] = None,
full_schema_reask: Optional[bool] = None,
**kwargs,
) -> Union[
ValidationOutcome[OT],
Awaitable[ValidationOutcome[OT]],
AsyncIterator[ValidationOutcome[OT]],
]:
prompt_params = prompt_params or {}
metadata = metadata or {}
if full_schema_reask is None:
full_schema_reask = self._base_model is not None
set_call_kwargs(kwargs)
self._set_num_reasks(num_reasks=num_reasks)
if self._num_reasks is None:
raise RuntimeError(
"`num_reasks` is `None` after calling `configure()`. "
"This should never happen."
)
messages = messages or self._exec_opts.messages
call_inputs = CallInputs(
llmApi=llm_api,
messages=messages,
promptParams=prompt_params,
numReasks=self._num_reasks,
metadata=metadata,
fullSchemaReask=full_schema_reask,
args=list(args),
kwargs=kwargs,
)
if self._use_server and model_is_supported_server_side(
llm_api, *args, **kwargs
):
result = self._call_server(
llm_output=llm_output,
llm_api=llm_api,
num_reasks=self._num_reasks,
prompt_params=prompt_params,
metadata=metadata,
full_schema_reask=full_schema_reask,
messages=messages,
*args,
**kwargs,
)
# If the LLM API is async, return a coroutine
else:
call_log = Call(inputs=call_inputs)
set_scope(str(object_id(call_log)))
self.history.push(call_log)
result = await self._exec(
llm_api=llm_api,
llm_output=llm_output,
prompt_params=prompt_params,
num_reasks=self._num_reasks,
messages=messages,
metadata=metadata,
full_schema_reask=full_schema_reask,
call_log=call_log,
*args,
**kwargs,
)
if inspect.isawaitable(result):
return await result
# TODO: Fix types once async streaming is implemented on server
return result # type: ignore
guard_context = contextvars.Context()
# get the current otel context and wrap the subsequent call
# to preserve otel context if guard call is being called by another
# framework upstream
current_otel_context = otel_context.get_current()
wrapped__exec = wrap_with_otel_context(current_otel_context, __exec)
return await guard_context.run(
wrapped__exec,
self,
llm_api=llm_api,
llm_output=llm_output,
prompt_params=prompt_params,
num_reasks=num_reasks,
messages=messages,
metadata=metadata,
full_schema_reask=full_schema_reask,
*args,
**kwargs,
)
async def _exec(
self,
*args,
llm_api: Optional[Callable[[Any], Awaitable[Any]]],
llm_output: Optional[str] = None,
call_log: Call,
prompt_params: Dict, # Should be defined at this point
num_reasks: int = 0, # Should be defined at this point
metadata: Dict, # Should be defined at this point
full_schema_reask: bool = False, # Should be defined at this point
messages: Optional[List[Dict]],
**kwargs,
) -> Union[
ValidationOutcome[OT],
Awaitable[ValidationOutcome[OT]],
AsyncIterator[ValidationOutcome[OT]],
]:
"""Call the LLM asynchronously and validate the output.
Args:
llm_api: The LLM API to call asynchronously (e.g. openai.Completion.acreate)
prompt_params: The parameters to pass to the prompt.format() method.
num_reasks: The max times to re-ask the LLM for invalid output.
messages: The message history to pass to the LLM.
metadata: Metadata to pass to the validators.
full_schema_reask: When reasking, whether to regenerate the full schema
or just the incorrect values.
Defaults to `True` if a base model is provided,
`False` otherwise.
Returns:
The raw text output from the LLM and the validated output.
"""
api = None
if llm_api is not None or kwargs.get("model") is not None:
api = get_async_llm_ask(llm_api, *args, **kwargs) # type: ignore
if self._output_formatter is not None:
api = self._output_formatter.wrap_async_callable(api) # type: ignore
if kwargs.get("stream", False):
runner = AsyncStreamRunner(
output_type=self._output_type,
output_schema=self.output_schema.model_dump(
exclude_none=True, by_alias=True
),
num_reasks=num_reasks,
validation_map=self._validator_map,
messages=messages,
api=api,
metadata=metadata,
output=llm_output,
base_model=self._base_model,
full_schema_reask=full_schema_reask,
disable_tracer=(
not self._allow_metrics_collection
if isinstance(self._allow_metrics_collection, bool)
else None
),
exec_options=self._exec_opts,
)
# Here we have an async generator
async_generator = runner.async_run(
call_log=call_log, prompt_params=prompt_params
)
return async_generator
else:
runner = AsyncRunner(
output_type=self._output_type,
output_schema=self.output_schema.model_dump(
exclude_none=True, by_alias=True
),
num_reasks=num_reasks,
validation_map=self._validator_map,
messages=messages,
api=api,
metadata=metadata,
output=llm_output,
base_model=self._base_model,
full_schema_reask=full_schema_reask,
disable_tracer=(
not self._allow_metrics_collection
if isinstance(self._allow_metrics_collection, bool)
else None
),
exec_options=self._exec_opts,
)
# Why are we using a different method here instead of just overriding?
call = await runner.async_run(
call_log=call_log, prompt_params=prompt_params
)
return ValidationOutcome[OT].from_guard_history(call)
@async_trace(name="/guard_call", origin="AsyncGuard.__call__")
async def __call__(
self,
llm_api: Optional[Callable[..., Awaitable[Any]]] = None,
*args,
prompt_params: Optional[Dict] = None,
num_reasks: Optional[int] = 1,
messages: Optional[List[Dict]] = None,
metadata: Optional[Dict] = None,
full_schema_reask: Optional[bool] = None,
**kwargs,
) -> Union[
ValidationOutcome[OT],
Awaitable[ValidationOutcome[OT]],
AsyncIterator[ValidationOutcome[OT]],
]:
"""Call the LLM and validate the output. Pass an async LLM API to
return a coroutine.
Args:
llm_api: The LLM API to call
(e.g. openai.completions.create or openai.chat.completions.create)
prompt_params: The parameters to pass to the prompt.format() method.
num_reasks: The max times to re-ask the LLM for invalid output.
messages: The message history to pass to the LLM.
metadata: Metadata to pass to the validators.
full_schema_reask: When reasking, whether to regenerate the full schema
or just the incorrect values.
Defaults to `True` if a base model is provided,
`False` otherwise.
Returns:
The raw text output from the LLM and the validated output.
"""
# Retrieve messages from the provided arguments or default options
messages_from_kwargs = kwargs.pop("messages", None)
messages_from_exec_opts = self._exec_opts.messages
# Determine the final value for messages
messages = messages or messages_from_kwargs or messages_from_exec_opts or []
if messages is not None and not len(messages):
raise RuntimeError(
"You must provide a prompt if messages is empty. "
"Alternatively, you can provide a prompt in the Schema constructor."
)
return await trace_async_guard_execution(
self.name,
self.history,
self._execute,
*args,
llm_api=llm_api,
prompt_params=prompt_params,
num_reasks=num_reasks,
messages=messages,
metadata=metadata,
full_schema_reask=full_schema_reask,
**kwargs,
)
@async_trace(name="/guard_call", origin="AsyncGuard.parse")
async def parse(
self,
llm_output: str,
*args,
metadata: Optional[Dict] = None,
llm_api: Optional[Callable[..., Awaitable[Any]]] = None,
num_reasks: Optional[int] = None,
prompt_params: Optional[Dict] = None,
full_schema_reask: Optional[bool] = None,
**kwargs,
) -> Awaitable[ValidationOutcome[OT]]:
"""Alternate flow to using AsyncGuard where the llm_output is known.
Args:
llm_output: The output being parsed and validated.
metadata: Metadata to pass to the validators.
llm_api: The LLM API to call
(e.g. openai.completions.create or openai.Completion.acreate)
num_reasks: The max times to re-ask the LLM for invalid output.
prompt_params: The parameters to pass to the prompt.format() method.
full_schema_reask: When reasking, whether to regenerate the full schema
or just the incorrect values.
Returns:
The validated response. This is either a string or a dictionary,
determined by the object schema defined in the RAILspec.
"""
final_num_reasks = (
num_reasks
if num_reasks is not None
else self._num_reasks
if self._num_reasks is not None
else 0
if llm_api is None
else 1
)
default_messages = self._exec_opts.messages if llm_api else None
messages = kwargs.pop("messages", default_messages)
return await trace_async_guard_execution( # type: ignore
self.name,
self.history,
self._execute,
*args,
llm_output=llm_output,
llm_api=llm_api,
prompt_params=prompt_params,
num_reasks=final_num_reasks,
messages=messages,
metadata=metadata,
full_schema_reask=full_schema_reask,
**kwargs,
)
async def _stream_server_call(
self, *, payload: Dict[str, Any]
) -> AsyncIterator[ValidationOutcome[OT]]:
# TODO: Once server side supports async streaming, this function will need to
# yield async generators, not generators
if self._api_client:
validation_output: Optional[IValidationOutcome] = None
response = self._api_client.stream_validate(
guard=self, # type: ignore
openai_api_key=get_call_kwarg("api_key"),
**payload,
)
for fragment in response:
validation_output = fragment
if validation_output is None:
yield ValidationOutcome[OT](
call_id="0", # type: ignore
rawLlmOutput=None,
validatedOutput=None,
validationPassed=False,
error="The response from the server was empty!",
)
else:
validated_output = (
cast(OT, validation_output.validated_output)
if validation_output.validated_output
else None
)
yield ValidationOutcome[OT](
call_id=validation_output.call_id, # type: ignore
raw_llm_output=validation_output.raw_llm_output, # type: ignore
validatedOutput=validated_output,
validationPassed=(validation_output.validation_passed is True),
)
# TODO re-enable this once we have a way to get history
# from a multi-node server
# if validation_output:
# guard_history = self._api_client.get_history(
# self.name, validation_output.call_id
# )
# self.history.extend(
# [Call.from_interface(call) for call in guard_history]
# )
else:
raise ValueError("AsyncGuard does not have an api client!")
@async_trace(name="/guard_call", origin="AsyncGuard.validate")
async def validate(
self, llm_output: str, *args, **kwargs
) -> Awaitable[ValidationOutcome[OT]]:
return await self.parse(llm_output=llm_output, *args, **kwargs)
@classmethod
def load(
cls,
name: str,
*,
api_key: Optional[str] = None,
base_url: Optional[str] = None,
history_max_length: Optional[int] = None,
) -> Optional["AsyncGuard"]:
guard = super().load(
name,
api_key=api_key,
base_url=base_url,
history_max_length=history_max_length,
)
if guard:
return cast(AsyncGuard, guard)