diff --git a/libs/partners/openai/langchain_openai/chat_models/base.py b/libs/partners/openai/langchain_openai/chat_models/base.py index cca86677ee..ebf428e58c 100644 --- a/libs/partners/openai/langchain_openai/chat_models/base.py +++ b/libs/partners/openai/langchain_openai/chat_models/base.py @@ -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: diff --git a/libs/partners/openai/tests/unit_tests/chat_models/test_base.py b/libs/partners/openai/tests/unit_tests/chat_models/test_base.py index 3a00101b4c..7bb0b74403 100644 --- a/libs/partners/openai/tests/unit_tests/chat_models/test_base.py +++ b/libs/partners/openai/tests/unit_tests/chat_models/test_base.py @@ -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")