mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #4392 from BerriAI/litellm_gemini_content_policy_errors
fix(vertex_httpx.py): cover gemini content violation (on prompt)
This commit is contained in:
commit
5d570e7c6c
4 changed files with 112 additions and 31 deletions
|
|
@ -562,7 +562,47 @@ class VertexLLM(BaseLLM):
|
|||
status_code=422,
|
||||
)
|
||||
|
||||
## GET MODEL ##
|
||||
model_response.model = model
|
||||
|
||||
## CHECK IF RESPONSE FLAGGED
|
||||
if "promptFeedback" in completion_response:
|
||||
if "blockReason" in completion_response["promptFeedback"]:
|
||||
# If set, the prompt was blocked and no candidates are returned. Rephrase your prompt
|
||||
model_response.choices[0].finish_reason = "content_filter"
|
||||
|
||||
chat_completion_message: ChatCompletionResponseMessage = {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
}
|
||||
|
||||
choice = litellm.Choices(
|
||||
finish_reason="content_filter",
|
||||
index=0,
|
||||
message=chat_completion_message, # type: ignore
|
||||
logprobs=None,
|
||||
enhancements=None,
|
||||
)
|
||||
|
||||
model_response.choices = [choice]
|
||||
|
||||
## GET USAGE ##
|
||||
usage = litellm.Usage(
|
||||
prompt_tokens=completion_response["usageMetadata"][
|
||||
"promptTokenCount"
|
||||
],
|
||||
completion_tokens=completion_response["usageMetadata"].get(
|
||||
"candidatesTokenCount", 0
|
||||
),
|
||||
total_tokens=completion_response["usageMetadata"][
|
||||
"totalTokenCount"
|
||||
],
|
||||
)
|
||||
|
||||
setattr(model_response, "usage", usage)
|
||||
|
||||
return model_response
|
||||
|
||||
if len(completion_response["candidates"]) > 0:
|
||||
content_policy_violations = (
|
||||
VertexGeminiConfig().get_flagged_finish_reasons()
|
||||
|
|
@ -573,26 +613,45 @@ class VertexLLM(BaseLLM):
|
|||
in content_policy_violations.keys()
|
||||
):
|
||||
## CONTENT POLICY VIOLATION ERROR
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message="The response was blocked. Reason={}. Raw Response={}".format(
|
||||
content_policy_violations[
|
||||
completion_response["candidates"][0]["finishReason"]
|
||||
],
|
||||
completion_response,
|
||||
),
|
||||
model_response.choices[0].finish_reason = "content_filter"
|
||||
|
||||
chat_completion_message = {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
}
|
||||
|
||||
choice = litellm.Choices(
|
||||
finish_reason="content_filter",
|
||||
index=0,
|
||||
message=chat_completion_message, # type: ignore
|
||||
logprobs=None,
|
||||
enhancements=None,
|
||||
)
|
||||
|
||||
model_response.choices = [choice]
|
||||
|
||||
## GET USAGE ##
|
||||
usage = litellm.Usage(
|
||||
prompt_tokens=completion_response["usageMetadata"][
|
||||
"promptTokenCount"
|
||||
],
|
||||
completion_tokens=completion_response["usageMetadata"].get(
|
||||
"candidatesTokenCount", 0
|
||||
),
|
||||
total_tokens=completion_response["usageMetadata"][
|
||||
"totalTokenCount"
|
||||
],
|
||||
)
|
||||
|
||||
setattr(model_response, "usage", usage)
|
||||
|
||||
return model_response
|
||||
|
||||
model_response.choices = [] # type: ignore
|
||||
|
||||
## GET MODEL ##
|
||||
model_response.model = model
|
||||
|
||||
try:
|
||||
## GET TEXT ##
|
||||
chat_completion_message: ChatCompletionResponseMessage = {
|
||||
"role": "assistant"
|
||||
}
|
||||
chat_completion_message = {"role": "assistant"}
|
||||
content_str = ""
|
||||
tools: List[ChatCompletionToolCallChunk] = []
|
||||
for idx, candidate in enumerate(completion_response["candidates"]):
|
||||
|
|
@ -632,9 +691,9 @@ class VertexLLM(BaseLLM):
|
|||
## GET USAGE ##
|
||||
usage = litellm.Usage(
|
||||
prompt_tokens=completion_response["usageMetadata"]["promptTokenCount"],
|
||||
completion_tokens=completion_response["usageMetadata"][
|
||||
"candidatesTokenCount"
|
||||
],
|
||||
completion_tokens=completion_response["usageMetadata"].get(
|
||||
"candidatesTokenCount", 0
|
||||
),
|
||||
total_tokens=completion_response["usageMetadata"]["totalTokenCount"],
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,7 @@
|
|||
model_list:
|
||||
- model_name: gemini-1.5-flash-gemini
|
||||
litellm_params:
|
||||
model: gemini/gemini-1.5-flash
|
||||
- litellm_params:
|
||||
api_base: http://0.0.0.0:8080
|
||||
api_key: ''
|
||||
|
|
|
|||
|
|
@ -696,6 +696,18 @@ async def test_gemini_pro_function_calling_httpx(provider, sync_mode):
|
|||
pytest.fail("An unexpected exception occurred - {}".format(str(e)))
|
||||
|
||||
|
||||
def vertex_httpx_mock_reject_prompt_post(*args, **kwargs):
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"Content-Type": "application/json"}
|
||||
mock_response.json.return_value = {
|
||||
"promptFeedback": {"blockReason": "OTHER"},
|
||||
"usageMetadata": {"promptTokenCount": 6285, "totalTokenCount": 6285},
|
||||
}
|
||||
|
||||
return mock_response
|
||||
|
||||
|
||||
# @pytest.mark.skip(reason="exhausted vertex quota. need to refactor to mock the call")
|
||||
def vertex_httpx_mock_post(url, data=None, json=None, headers=None):
|
||||
mock_response = MagicMock()
|
||||
|
|
@ -817,8 +829,11 @@ def vertex_httpx_mock_post(url, data=None, json=None, headers=None):
|
|||
|
||||
|
||||
@pytest.mark.parametrize("provider", ["vertex_ai_beta"]) # "vertex_ai",
|
||||
@pytest.mark.parametrize("content_filter_type", ["prompt", "response"]) # "vertex_ai",
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_pro_json_schema_httpx_content_policy_error(provider):
|
||||
async def test_gemini_pro_json_schema_httpx_content_policy_error(
|
||||
provider, content_filter_type
|
||||
):
|
||||
load_vertex_ai_credentials()
|
||||
litellm.set_verbose = True
|
||||
messages = [
|
||||
|
|
@ -839,16 +854,20 @@ Using this JSON schema:
|
|||
|
||||
client = HTTPHandler()
|
||||
|
||||
with patch.object(client, "post", side_effect=vertex_httpx_mock_post) as mock_call:
|
||||
try:
|
||||
response = completion(
|
||||
model="vertex_ai_beta/gemini-1.5-flash",
|
||||
messages=messages,
|
||||
response_format={"type": "json_object"},
|
||||
client=client,
|
||||
)
|
||||
except litellm.ContentPolicyViolationError as e:
|
||||
pass
|
||||
if content_filter_type == "prompt":
|
||||
_side_effect = vertex_httpx_mock_reject_prompt_post
|
||||
else:
|
||||
_side_effect = vertex_httpx_mock_post
|
||||
|
||||
with patch.object(client, "post", side_effect=_side_effect) as mock_call:
|
||||
response = completion(
|
||||
model="vertex_ai_beta/gemini-1.5-flash",
|
||||
messages=messages,
|
||||
response_format={"type": "json_object"},
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert response.choices[0].finish_reason == "content_filter"
|
||||
|
||||
mock_call.assert_called_once()
|
||||
|
||||
|
|
|
|||
|
|
@ -227,9 +227,9 @@ class PromptFeedback(TypedDict):
|
|||
blockReasonMessage: str
|
||||
|
||||
|
||||
class UsageMetadata(TypedDict):
|
||||
promptTokenCount: int
|
||||
totalTokenCount: int
|
||||
class UsageMetadata(TypedDict, total=False):
|
||||
promptTokenCount: Required[int]
|
||||
totalTokenCount: Required[int]
|
||||
candidatesTokenCount: int
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue