Files
AI/참고/guardrails-main/tests/integration_tests/test_pydantic.py
2026-05-12 19:40:31 +09:00

228 lines
7.2 KiB
Python

import json
from typing import Dict, List
import pytest
from pydantic import BaseModel
import guardrails as gd
from guardrails.classes.generic.stack import Stack
from guardrails.classes.history.call import Call
from guardrails.classes.llm.llm_response import LLMResponse
from .mock_llm_outputs import pydantic
from .test_assets.pydantic import VALIDATED_RESPONSE_REASK_PROMPT, ListOfPeople
def test_pydantic_with_reask(mocker):
"""Test that the entity extraction works with re-asking."""
mock_invoke_llm = mocker.patch(
"guardrails.llm_providers.LiteLLMCallable._invoke_llm"
)
mock_invoke_llm.side_effect = [
LLMResponse(
output=pydantic.LLM_OUTPUT,
prompt_token_count=123,
response_token_count=1234,
),
LLMResponse(
output=pydantic.LLM_OUTPUT_REASK_1,
prompt_token_count=123,
response_token_count=1234,
),
LLMResponse(
output=pydantic.LLM_OUTPUT_REASK_2,
prompt_token_count=123,
response_token_count=1234,
),
]
guard = gd.Guard.for_pydantic(
ListOfPeople,
messages=[{"role": "user", "content": VALIDATED_RESPONSE_REASK_PROMPT}],
)
final_output = guard(
model="text-davinci-003",
max_tokens=512,
temperature=0.5,
num_reasks=2,
full_schema_reask=False,
)
# Assertions are made on the guard state object.
assert final_output.validation_passed is False
assert final_output.validated_output is None
call = guard.history.first
# Check that the guard state object has the correct number of re-asks.
assert call.iterations.length == 3
# For original prompt and output
assert call.compiled_messages[0]["content"] == pydantic.COMPILED_PROMPT
assert call.iterations.first.raw_output == pydantic.LLM_OUTPUT
assert (
call.iterations.first.validation_response == pydantic.VALIDATED_OUTPUT_REASK_1
)
# For re-asked prompt and output
# Assert through iteration
assert call.iterations.at(1).inputs.messages[1]["content"] == gd.Prompt(
pydantic.COMPILED_PROMPT_REASK_1
)
assert call.iterations.at(1).raw_output == pydantic.LLM_OUTPUT_REASK_1
# Assert through call shortcut properties
assert call.reask_messages.first[1]["content"] == pydantic.COMPILED_PROMPT_REASK_1
assert call.raw_outputs.at(1) == pydantic.LLM_OUTPUT_REASK_1
# We don't track merged validation output anymore
# Each validation_output is instead tracked as it came back from validation
# So this isn't a thing
# assert call.iterations.at(1).validation_response == (
# pydantic.VALIDATED_OUTPUT_REASK_2
# )
# We can, however, merge down to achieve the same thing
intermediate_call_state = Call(
iterations=Stack(call.iterations.first, call.iterations.at(1))
)
intermediate_call_state.inputs.full_schema_reask = False
assert (
intermediate_call_state.validation_response == pydantic.VALIDATED_OUTPUT_REASK_2
)
# For re-asked prompt #2 and output #2
assert call.iterations.last.inputs.messages[1]["content"] == gd.Prompt(
pydantic.COMPILED_PROMPT_REASK_2
)
# Same as above
assert call.reask_messages.last[1]["content"] == pydantic.COMPILED_PROMPT_REASK_2
assert call.raw_outputs.last == pydantic.LLM_OUTPUT_REASK_2
assert call.guarded_output is None
assert call.validation_response == pydantic.VALIDATED_OUTPUT_REASK_3
def test_pydantic_with_full_schema_reask(mocker):
"""Test that the entity extraction works with re-asking."""
mock_invoke_llm = mocker.patch(
"guardrails.llm_providers.LiteLLMCallable._invoke_llm"
)
mock_invoke_llm.side_effect = [
LLMResponse(
output=pydantic.LLM_OUTPUT,
prompt_token_count=123,
response_token_count=1234,
),
LLMResponse(
output=pydantic.LLM_OUTPUT_FULL_REASK_1,
prompt_token_count=123,
response_token_count=1234,
),
LLMResponse(
output=pydantic.LLM_OUTPUT_FULL_REASK_2,
prompt_token_count=123,
response_token_count=1234,
),
]
guard = gd.Guard.for_pydantic(
ListOfPeople,
messages=[
{
"content": VALIDATED_RESPONSE_REASK_PROMPT,
"role": "user",
}
],
)
final_output = guard(
model="gpt-3.5-turbo",
max_tokens=512,
temperature=0.5,
num_reasks=2,
full_schema_reask=True,
)
# Assertions are made on the guard state object.
assert final_output.validation_passed is False
assert final_output.validated_output is None
call = guard.history.first
# Check that the guard state object has the correct number of re-asks.
assert call.iterations.length == 3
# For original prompt and output
assert call.compiled_messages[0]["content"] == pydantic.COMPILED_PROMPT_CHAT
assert call.iterations.first.raw_output == pydantic.LLM_OUTPUT
assert (
call.iterations.first.validation_response == pydantic.VALIDATED_OUTPUT_REASK_1
)
# For re-asked prompt and output
assert (
call.iterations.at(1).inputs.messages[1]["content"]._source
== pydantic.COMPILED_PROMPT_FULL_REASK_1
)
assert (
call.iterations.at(1).inputs.messages[0]["content"]._source
== pydantic.COMPILED_INSTRUCTIONS_CHAT
)
assert call.iterations.at(1).raw_output == pydantic.LLM_OUTPUT_FULL_REASK_1
assert (
call.iterations.at(1).validation_response == pydantic.VALIDATED_OUTPUT_REASK_2
)
# For re-asked prompt #2 and output #2
assert call.iterations.last.inputs.messages[1]["content"] == gd.Prompt(
pydantic.COMPILED_PROMPT_FULL_REASK_2
)
assert call.iterations.last.inputs.messages[0]["content"] == gd.Instructions(
pydantic.COMPILED_INSTRUCTIONS_CHAT
)
assert call.raw_outputs.last == pydantic.LLM_OUTPUT_FULL_REASK_2
assert call.guarded_output is None
assert call.validation_response == pydantic.VALIDATED_OUTPUT_REASK_3
class ContainerModel(BaseModel):
annotated_dict: Dict[str, str] = {}
annotated_dict_in_list: List[Dict[str, str]] = []
annotated_list: List[str] = []
annotated_list_in_dict: Dict[str, List[str]] = {}
class ContainerModel2(BaseModel):
dict_: Dict = {}
dict_in_list: List[Dict] = []
list_: List = []
list_in_dict: Dict[str, List] = {}
@pytest.mark.parametrize(
"model, output",
[
(
ContainerModel,
{
"annotated_dict": {"a": "b"},
"annotated_dict_in_list": [{"a": "b"}],
"annotated_list": ["a"],
"annotated_list_in_dict": {"a": ["b"]},
},
),
(
ContainerModel2,
{
"dict_": {"a": "b"},
"dict_in_list": [{"a": "b"}],
"list_": ["a"],
"list_in_dict": {"a": ["b"]},
},
),
],
)
def test_container_types(model, output):
output_str = json.dumps(output)
guard = gd.Guard.for_pydantic(model)
out = guard.parse(output_str)
assert out.validated_output == output