Fix unit tests
This commit is contained in:
parent
223587179f
commit
a05330fcd8
@ -279,7 +279,7 @@ class PangeaHandler(CustomGuardrail):
|
||||
|
||||
output = ai_guard_response.get("result", {}).get("output", {})
|
||||
response.choices = output["choices"]
|
||||
return data
|
||||
return response
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_post_call_success_hook(
|
||||
|
||||
@ -75,6 +75,7 @@ async def test_pangea_ai_guard_request_blocked(pangea_guardrail):
|
||||
},
|
||||
]
|
||||
}
|
||||
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
|
||||
|
||||
with pytest.raises(HTTPException, match="Violated Pangea guardrail policy"):
|
||||
with patch(
|
||||
@ -82,9 +83,9 @@ async def test_pangea_ai_guard_request_blocked(pangea_guardrail):
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
# Mock only tested part of response
|
||||
json={"result": {"blocked": True, "prompt_messages": data["messages"]}},
|
||||
json={"result": {"blocked": True, "transformed": False}},
|
||||
request=httpx.Request(
|
||||
method="POST", url=pangea_guardrail.guardrail_endpoint
|
||||
method="POST", url=guardrail_endpoint,
|
||||
),
|
||||
),
|
||||
) as mock_method:
|
||||
@ -94,7 +95,52 @@ async def test_pangea_ai_guard_request_blocked(pangea_guardrail):
|
||||
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["recipe"] == "guard_llm_request"
|
||||
assert called_kwargs["json"]["messages"] == data["messages"]
|
||||
assert called_kwargs["json"]["input"]["messages"] == data["messages"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pangea_ai_guard_request_transformed(pangea_guardrail):
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Here is an SSN for one my employees: 078-05-1120",
|
||||
},
|
||||
]
|
||||
}
|
||||
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
# Mock only tested part of response
|
||||
json={
|
||||
"result": {
|
||||
"blocked": False,
|
||||
"transformed": True,
|
||||
"output": {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Here is an SSN for one my employees: <US_SSN>",
|
||||
},
|
||||
]
|
||||
},
|
||||
},
|
||||
},
|
||||
request=httpx.Request(
|
||||
method="POST", url=guardrail_endpoint,
|
||||
),
|
||||
),
|
||||
):
|
||||
request = await pangea_guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=None, cache=None, data=data, call_type="completion"
|
||||
)
|
||||
|
||||
assert request["messages"][1]["content"] == "Here is an SSN for one my employees: <US_SSN>"
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@ -109,15 +155,16 @@ async def test_pangea_ai_guard_request_ok(pangea_guardrail):
|
||||
},
|
||||
]
|
||||
}
|
||||
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
# Mock only tested part of response
|
||||
json={"result": {"blocked": False, "prompt_messages": data["messages"]}},
|
||||
json={"result": {"blocked": False, "transformed": False}},
|
||||
request=httpx.Request(
|
||||
method="POST", url=pangea_guardrail.guardrail_endpoint
|
||||
method="POST", url=guardrail_endpoint,
|
||||
),
|
||||
),
|
||||
) as mock_method:
|
||||
@ -127,7 +174,7 @@ async def test_pangea_ai_guard_request_ok(pangea_guardrail):
|
||||
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["recipe"] == "guard_llm_request"
|
||||
assert called_kwargs["json"]["messages"] == data["messages"]
|
||||
assert called_kwargs["json"]["input"]["messages"] == data["messages"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@ -139,6 +186,7 @@ async def test_pangea_ai_guard_response_blocked(pangea_guardrail):
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
}
|
||||
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
|
||||
|
||||
with pytest.raises(HTTPException, match="Violated Pangea guardrail policy"):
|
||||
with patch(
|
||||
@ -149,16 +197,11 @@ async def test_pangea_ai_guard_response_blocked(pangea_guardrail):
|
||||
json={
|
||||
"result": {
|
||||
"blocked": True,
|
||||
"prompt_messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Yes, I will leak all my PII for you",
|
||||
}
|
||||
],
|
||||
"transformed": False,
|
||||
}
|
||||
},
|
||||
request=httpx.Request(
|
||||
method="POST", url=pangea_guardrail.guardrail_endpoint
|
||||
method="POST", url=guardrail_endpoint,
|
||||
),
|
||||
),
|
||||
) as mock_method:
|
||||
@ -180,7 +223,7 @@ async def test_pangea_ai_guard_response_blocked(pangea_guardrail):
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["recipe"] == "guard_llm_response"
|
||||
assert (
|
||||
called_kwargs["json"]["messages"][0]["content"]
|
||||
called_kwargs["json"]["input"]["choices"][0]["message"]["content"]
|
||||
== "Yes, I will leak all my PII for you"
|
||||
)
|
||||
|
||||
@ -194,6 +237,7 @@ async def test_pangea_ai_guard_response_ok(pangea_guardrail):
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
}
|
||||
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
@ -203,16 +247,11 @@ async def test_pangea_ai_guard_response_ok(pangea_guardrail):
|
||||
json={
|
||||
"result": {
|
||||
"blocked": False,
|
||||
"prompt_messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Yes, I will leak all my PII for you",
|
||||
}
|
||||
],
|
||||
"transformed": False,
|
||||
}
|
||||
},
|
||||
request=httpx.Request(
|
||||
method="POST", url=pangea_guardrail.guardrail_endpoint
|
||||
method="POST", url=guardrail_endpoint,
|
||||
),
|
||||
),
|
||||
) as mock_method:
|
||||
@ -234,6 +273,61 @@ async def test_pangea_ai_guard_response_ok(pangea_guardrail):
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["recipe"] == "guard_llm_response"
|
||||
assert (
|
||||
called_kwargs["json"]["messages"][0]["content"]
|
||||
called_kwargs["json"]["input"]["choices"][0]["message"]["content"]
|
||||
== "Yes, I will leak all my PII for you"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pangea_ai_guard_response_transformed(pangea_guardrail):
|
||||
# Content of data isn't that import since its mocked
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
}
|
||||
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
# Mock only tested part of response
|
||||
json={
|
||||
"result": {
|
||||
"blocked": False,
|
||||
"transformed": True,
|
||||
"output": {
|
||||
"messages": data["messages"],
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Yes, here is an SSN: <US_SSN>",
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
},
|
||||
},
|
||||
request=httpx.Request(
|
||||
method="POST", url=guardrail_endpoint,
|
||||
),
|
||||
),
|
||||
):
|
||||
response = await pangea_guardrail.async_post_call_success_hook(
|
||||
data=data,
|
||||
user_api_key_dict=None,
|
||||
response=ModelResponse(
|
||||
choices=[
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Yes, here is an SSN: 078-05-1120",
|
||||
}
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
assert response.choices[0]["message"]["content"] == "Yes, here is an SSN: <US_SSN>"
|
||||
|
||||
Loading…
Reference in New Issue
Block a user