fix(openai): normalize content-policy refusal errors

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
Mason Daughertyandopen-swe[bot] committed 2026-10-02 00:12:08 +00:00
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")