mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
fix(openai): normalize content-policy refusal errors
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
1 parent
884d2d66b4
commit
be83707e5b
2 files changed
+142
-10
No files matched your search
@@ -56,6 +56,7 @@ from langchain_core.exceptions import (
|
||||
ModelAPIError,
|
||||
ModelAuthenticationError,
|
||||
ModelConnectionError,
|
||||
ModelError,
|
||||
ModelInvalidRequestError,
|
||||
ModelNotFoundError,
|
||||
ModelPermissionDeniedError,
|
||||
@@ -635,6 +636,36 @@ def _update_token_usage(
|
||||
return new_usage
|
||||
|
||||
|
||||
class OpenAIRefusalError(ModelError):
|
||||
"""Error raised when OpenAI refuses a request on content-policy grounds."""
|
||||
|
||||
|
||||
class OpenAIBadRequestRefusalError(openai.BadRequestError, OpenAIRefusalError):
|
||||
"""OpenAI bad-request refusal retaining SDK exception compatibility."""
|
||||
|
||||
|
||||
class OpenAIPermissionRefusalError(openai.PermissionDeniedError, OpenAIRefusalError):
|
||||
"""OpenAI permission refusal retaining SDK exception compatibility."""
|
||||
|
||||
|
||||
class OpenAIAPIRefusalError(openai.APIError, OpenAIRefusalError):
|
||||
"""OpenAI streaming refusal retaining SDK exception compatibility."""
|
||||
|
||||
|
||||
def _is_openai_refusal(e: openai.APIError) -> bool:
|
||||
body = e.body
|
||||
if isinstance(body, dict):
|
||||
body = body.get("error", body)
|
||||
if isinstance(body, dict) and any(
|
||||
body.get(key) in ("content_filter", "content_policy_violation")
|
||||
for key in ("type", "code")
|
||||
):
|
||||
return True
|
||||
return type(e) is openai.APIError and (
|
||||
"this content was flagged for possible cybersecurity risk" in e.message.lower()
|
||||
)
|
||||
|
||||
|
||||
class OpenAIContextOverflowError(openai.BadRequestError, ContextOverflowError):
|
||||
"""BadRequestError raised when input exceeds OpenAI's context limit."""
|
||||
|
||||
@@ -678,6 +709,10 @@ class OpenAITimeoutError(openai.APITimeoutError, ModelTimeoutError):
|
||||
|
||||
|
||||
def _handle_openai_bad_request(e: openai.BadRequestError) -> None:
|
||||
if _is_openai_refusal(e):
|
||||
raise OpenAIBadRequestRefusalError(
|
||||
message=e.message, response=e.response, body=e.body
|
||||
) from e
|
||||
if (
|
||||
"context_length_exceeded" in str(e)
|
||||
or "Input tokens exceed the configured limit" in e.message
|
||||
@@ -715,6 +750,15 @@ def _handle_openai_bad_request(e: openai.BadRequestError) -> None:
|
||||
|
||||
|
||||
def _handle_openai_api_error(e: openai.APIError) -> None:
|
||||
if _is_openai_refusal(e):
|
||||
if isinstance(e, openai.PermissionDeniedError):
|
||||
raise OpenAIPermissionRefusalError(
|
||||
message=e.message, response=e.response, body=e.body
|
||||
) from e
|
||||
if type(e) is openai.APIError:
|
||||
raise OpenAIAPIRefusalError(
|
||||
message=e.message, request=e.request, body=e.body
|
||||
) from e
|
||||
error_message = str(e)
|
||||
if "exceeds the context window" in error_message:
|
||||
raise OpenAIAPIContextOverflowError(
|
||||
@@ -4465,16 +4509,6 @@ def _oai_structured_outputs_parser(
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
class OpenAIRefusalError(Exception):
|
||||
"""Error raised when OpenAI Structured Outputs API returns a refusal.
|
||||
|
||||
When using OpenAI's Structured Outputs API with user-generated input, the model
|
||||
may occasionally refuse to fulfill the request for safety reasons.
|
||||
|
||||
See [more on refusals](https://platform.openai.com/docs/guides/structured-outputs/refusals).
|
||||
"""
|
||||
|
||||
|
||||
def _create_usage_metadata(
|
||||
oai_token_usage: dict, service_tier: str | None = None
|
||||
) -> UsageMetadata:
|
||||
|
||||
@@ -5042,6 +5042,104 @@ def test_openai_error_classification(
|
||||
assert exc_info.value.is_retryable is is_retryable
|
||||
|
||||
|
||||
@pytest.mark.parametrize("use_responses_api", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
("status_code", "sdk_error_type"),
|
||||
[(400, openai.BadRequestError), (403, openai.PermissionDeniedError)],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
("body", "is_refusal"),
|
||||
[
|
||||
({"code": "content_filter"}, True),
|
||||
({"type": "content_policy_violation"}, True),
|
||||
({"error": {"type": "invalid_request_error", "code": "content_filter"}}, True),
|
||||
({"code": "invalid_prompt"}, False),
|
||||
(None, False),
|
||||
],
|
||||
)
|
||||
def test_openai_content_policy_error(
|
||||
use_responses_api: bool,
|
||||
status_code: int,
|
||||
sdk_error_type: type[openai.APIStatusError],
|
||||
body: dict[str, object] | None,
|
||||
is_refusal: bool,
|
||||
) -> None:
|
||||
request = httpx2.Request("POST", "https://api.openai.com/v1/responses")
|
||||
response = httpx2.Response(status_code, request=request)
|
||||
sdk_error = sdk_error_type("request rejected", response=response, body=body)
|
||||
model = ChatOpenAI(use_responses_api=use_responses_api)
|
||||
client = model.root_client.responses if use_responses_api else model.client
|
||||
|
||||
with patch.object(client, "with_raw_response") as mock_client:
|
||||
mock_client.create.side_effect = sdk_error
|
||||
with pytest.raises(sdk_error_type) as exc_info:
|
||||
model.invoke("test")
|
||||
|
||||
error = exc_info.value
|
||||
assert isinstance(error, OpenAIRefusalError) is is_refusal
|
||||
assert isinstance(error, ModelError)
|
||||
assert error.is_retryable is False
|
||||
assert error.response is response
|
||||
assert error.body is body
|
||||
assert error.message == sdk_error.message
|
||||
assert error.__cause__ is sdk_error
|
||||
|
||||
|
||||
@pytest.mark.parametrize("use_responses_api", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
("message", "body", "is_refusal"),
|
||||
[
|
||||
("This content was flagged for possible cybersecurity risk.", None, True),
|
||||
("request rejected", {"code": "content_policy_violation"}, True),
|
||||
("unrelated stream failure", None, False),
|
||||
],
|
||||
)
|
||||
async def test_openai_stream_refusal(
|
||||
use_responses_api: bool,
|
||||
message: str,
|
||||
body: dict[str, object] | None,
|
||||
is_refusal: bool,
|
||||
) -> None:
|
||||
request = httpx2.Request("POST", "https://api.openai.com/v1/responses")
|
||||
sdk_error = openai.APIError(message, request=request, body=body)
|
||||
model = ChatOpenAI(use_responses_api=use_responses_api)
|
||||
client = model.root_client.responses if use_responses_api else model.client
|
||||
async_client = (
|
||||
model.root_async_client.responses if use_responses_api else model.async_client
|
||||
)
|
||||
stream = MagicMock()
|
||||
stream.__enter__.return_value = stream
|
||||
stream.__iter__.side_effect = sdk_error
|
||||
async_stream = MagicMock()
|
||||
async_stream.__aenter__.return_value = async_stream
|
||||
async_stream.__aiter__.side_effect = lambda: async_stream
|
||||
async_stream.__anext__.side_effect = sdk_error
|
||||
|
||||
with (
|
||||
patch.object(client, "create", return_value=stream),
|
||||
pytest.raises(openai.APIError) as sync_exc,
|
||||
):
|
||||
list(model.stream("test"))
|
||||
with (
|
||||
patch.object(async_client, "create", AsyncMock(return_value=async_stream)),
|
||||
pytest.raises(openai.APIError) as async_exc,
|
||||
):
|
||||
async for _ in model.astream("test"):
|
||||
pass
|
||||
|
||||
for error in (sync_exc.value, async_exc.value):
|
||||
assert isinstance(error, OpenAIRefusalError) is is_refusal
|
||||
assert error.request is request
|
||||
assert error.body is body
|
||||
assert error.message == message
|
||||
if is_refusal:
|
||||
assert isinstance(error, ModelError)
|
||||
assert error.is_retryable is False
|
||||
assert error.__cause__ is sdk_error
|
||||
else:
|
||||
assert error is sdk_error
|
||||
|
||||
|
||||
def test_openai_transport_error_classification() -> None:
|
||||
"""Timeout and connection failures are classified without a status code."""
|
||||
request = httpx2.Request("POST", "https://api.openai.com/v1/chat/completions")
|
||||
|
||||
Reference in new issue
Block a user