Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
202 changes: 200 additions & 2 deletions src/openai/_exceptions.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
from __future__ import annotations

import sys
from datetime import timedelta
from typing import TYPE_CHECKING, Any, Optional, cast
from typing_extensions import Literal
from typing_extensions import Literal, override

import httpx2

Expand Down Expand Up @@ -30,9 +32,205 @@
"WebSocketQueueFullError",
]

_RequestTimeoutSnapshot = dict[str, float | None]
_RequestSnapshot = tuple[str, list[tuple[bytes, bytes]], _RequestTimeoutSnapshot | None]
_ResponseExtensionsSnapshot = dict[str, bytes]
_ResponseSnapshot = tuple[
int,
list[tuple[bytes, bytes]],
_RequestSnapshot | None,
timedelta | None,
_ResponseExtensionsSnapshot,
]

_REDACTED_REQUEST_URL = "https://redacted.invalid/"
_REDACTED_HEADER_VALUE = b"<redacted>"
_SAFE_HEADER_VALUES = {
b"content-length",
b"content-type",
b"x-request-id",
}
_TIMEOUT_EXTENSION_KEYS = ("connect", "read", "write", "pool")
_RESPONSE_EXTENSION_KEYS = ("http_version", "reason_phrase")


class OpenAIError(Exception):
pass
@override
def __reduce__(self) -> tuple[object, tuple[object, ...]]:
state = self.__dict__.copy()
Comment thread
sylvesterkaczmarek marked this conversation as resolved.
slot_state = _snapshot_slot_state(self)
Comment thread
sylvesterkaczmarek marked this conversation as resolved.

if isinstance(self, WebSocketConnectionClosedError):
state["unsent_messages"] = []

request_snapshot: _RequestSnapshot | None = None
request = state.get("request")
if _is_http_request(request):
request_snapshot = _snapshot_request(request)
del state["request"]

response_snapshot: _ResponseSnapshot | None = None
response = state.get("response")
if _is_http_response(response):
response_snapshot = _snapshot_response(response)
del state["response"]

return (
_reconstruct_openai_error,
(type(self), self.args, state, slot_state, request_snapshot, response_snapshot),
)


def _is_legacy_httpx_instance(value: object, type_name: str) -> bool:
module = sys.modules.get("httpx")
legacy_type = getattr(module, type_name, None) if module is not None else None
return isinstance(legacy_type, type) and isinstance(value, legacy_type)


def _is_http_request(value: object) -> bool:
return isinstance(value, httpx2.Request) or _is_legacy_httpx_instance(value, "Request")


def _is_http_response(value: object) -> bool:
return isinstance(value, httpx2.Response) or _is_legacy_httpx_instance(value, "Response")


def _snapshot_headers(headers: Any) -> list[tuple[bytes, bytes]]:
return [
(name, value if name.lower() in _SAFE_HEADER_VALUES else _REDACTED_HEADER_VALUE)
for name, value in headers.raw
Comment thread
sylvesterkaczmarek marked this conversation as resolved.
]


def _snapshot_timeout_extension(request: Any) -> _RequestTimeoutSnapshot | None:
extensions = request.extensions
timeout = extensions.get("timeout") if isinstance(extensions, dict) else None
if not isinstance(timeout, dict):
return None

snapshot: _RequestTimeoutSnapshot = {}
for key in _TIMEOUT_EXTENSION_KEYS:
value = timeout.get(key)
if value is None:
if key in timeout:
snapshot[key] = None
elif isinstance(value, (int, float)):
snapshot[key] = float(value)
return snapshot or None


def _snapshot_request(request: Any) -> _RequestSnapshot:
return (
request.method,
_snapshot_headers(request.headers),
_snapshot_timeout_extension(request),
)


def _restore_request(snapshot: _RequestSnapshot) -> httpx2.Request:
method, headers, timeout = snapshot
kwargs: dict[str, Any] = {"headers": headers}
if timeout is not None:
kwargs["extensions"] = {"timeout": timeout}
return httpx2.Request(method, _REDACTED_REQUEST_URL, **kwargs)


def _snapshot_response_extensions(response: Any) -> _ResponseExtensionsSnapshot:
extensions = response.extensions
if not isinstance(extensions, dict):
return {}

snapshot: _ResponseExtensionsSnapshot = {}
for key in _RESPONSE_EXTENSION_KEYS:
value = extensions.get(key)
if isinstance(value, bytes):
snapshot[key] = value
return snapshot


def _snapshot_response(response: Any) -> _ResponseSnapshot:
try:
request_snapshot = _snapshot_request(response.request)
except RuntimeError:
request_snapshot = None

try:
elapsed = response.elapsed
except RuntimeError:
elapsed = None

return (
response.status_code,
_snapshot_headers(response.headers),
request_snapshot,
elapsed,
Comment thread
sylvesterkaczmarek marked this conversation as resolved.
_snapshot_response_extensions(response),
)


def _restore_response(
snapshot: _ResponseSnapshot,
*,
request: httpx2.Request | None,
) -> httpx2.Response:
status_code, headers, response_request_snapshot, elapsed, extensions = snapshot
if request is None and response_request_snapshot is not None:
request = _restore_request(response_request_snapshot)

kwargs: dict[str, Any] = {"headers": headers}
if request is not None:
kwargs["request"] = request
if extensions:
kwargs["extensions"] = extensions

response = httpx2.Response(status_code, **kwargs)
if elapsed is not None:
response.elapsed = elapsed
return response


def _snapshot_slot_state(error: OpenAIError) -> dict[str, Any]:
slot_state: dict[str, Any] = {}
for cls in type(error).__mro__:
slots = cls.__dict__.get("__slots__")
if slots is None:
continue
if isinstance(slots, str):
slots = (slots,)
for slot in slots:
if slot in {"__dict__", "__weakref__"}:
continue
name = slot
if name.startswith("__") and not name.endswith("__"):
name = f"_{cls.__name__.lstrip('_')}{name}"
try:
slot_state[name] = getattr(error, name)
except AttributeError:
pass
return slot_state


def _reconstruct_openai_error(
error_type: type[OpenAIError],
args: tuple[object, ...],
state: dict[str, Any],
slot_state: dict[str, Any],
request_snapshot: _RequestSnapshot | None,
response_snapshot: _ResponseSnapshot | None,
) -> OpenAIError:
error = Exception.__new__(error_type)
Exception.__init__(error, *args)
error.__dict__.update(state)
for name, value in slot_state.items():
setattr(error, name, value)

request = _restore_request(request_snapshot) if request_snapshot is not None else None
if request is not None:
error.__dict__["request"] = request
if response_snapshot is not None:
error.__dict__["response"] = _restore_response(response_snapshot, request=request)

return error


class SubjectTokenProviderError(OpenAIError):
Expand Down
Loading