diff --git a/src/openai/_exceptions.py b/src/openai/_exceptions.py index 86f44b0e15..22055e40d8 100644 --- a/src/openai/_exceptions.py +++ b/src/openai/_exceptions.py @@ -2,7 +2,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Optional, cast +from typing import TYPE_CHECKING, Any, Union, Optional, cast from typing_extensions import Literal import httpx @@ -60,7 +60,7 @@ class APIError(OpenAIError): If there was no response associated with this error then it will be `None`. """ - code: Optional[str] = None + code: Optional[Union[str, int]] = None param: Optional[str] = None type: Optional[str] @@ -71,7 +71,8 @@ def __init__(self, message: str, request: httpx.Request, *, body: object | None) self.body = body if is_dict(body): - self.code = cast(Any, construct_type(type_=Optional[str], value=body.get("code"))) + code_value = body.get("code") + self.code = code_value if isinstance(code_value, (str, int)) else None self.param = cast(Any, construct_type(type_=Optional[str], value=body.get("param"))) self.type = cast(Any, construct_type(type_=str, value=body.get("type"))) else: diff --git a/tests/lib/test_exceptions.py b/tests/lib/test_exceptions.py new file mode 100644 index 0000000000..8ec89be5c0 --- /dev/null +++ b/tests/lib/test_exceptions.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +import httpx + +from openai._exceptions import BadRequestError + + +def _make_error(body: dict[str, object]) -> BadRequestError: + request = httpx.Request("POST", "https://api.openai.com/v1/chat/completions") + response = httpx.Response(400, request=request) + return BadRequestError(message="test error", response=response, body=body) + + +def test_integer_error_code_is_preserved() -> None: + # The OpenAI API can return integer error codes (e.g. 1006). Constructing + # the error from such a body must keep the raw int: Pydantic v1 union + # coercion would otherwise stringify it to "1006". + err = _make_error({"message": "test error", "code": 123}) + assert err.code == 123 + assert isinstance(err.code, int) + + +def test_string_error_code_is_preserved() -> None: + # String codes (the common case) must keep working unchanged. + err = _make_error({"message": "test error", "code": "rate_limit_exceeded"}) + assert err.code == "rate_limit_exceeded" + + +def test_missing_or_invalid_error_code_is_none() -> None: + assert _make_error({"message": "test error"}).code is None + assert _make_error({"message": "test error", "code": {"nested": True}}).code is None