Skip to content
Merged
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
36 changes: 28 additions & 8 deletions py/src/braintrust/api/_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@
BraintrustTransportError,
BraintrustTransportRetryExhaustedError,
)
from .policies import RetryMode, RetryPolicy
from .policies import RetryMode, RetryPolicy, is_retryable_request_exception


logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -102,6 +102,9 @@ def __init__(self, base_url: str, adapter: HTTPAdapter | None = None):
self.base_url = base_url
self.token = None
self.adapter = adapter
# An adapter handed to us belongs to the caller. `set_http_adapter` installs
# one instance across every connection, so we must not close it.
self._injected_adapter = adapter

self._reset(total=0)

Expand All @@ -120,6 +123,10 @@ def make_long_lived(self) -> None:
)
self._reset()

def close(self) -> None:
_unmount_adapter(self.session, self._injected_adapter)
self.session.close()

@staticmethod
def sanitize_token(token: str) -> str:
return token.rstrip("\n")
Expand All @@ -131,6 +138,7 @@ def set_token(self, token: str) -> None:

def _set_adapter(self, adapter: HTTPAdapter | None) -> None:
self.adapter = adapter
self._injected_adapter = adapter

def _reset(self, **retry_kwargs: Any) -> None:
self.session = requests.Session()
Expand Down Expand Up @@ -202,6 +210,7 @@ def __init__(
):
custom_transport = session is not None or adapter is not None
self._owns_session = session is None
self._injected_adapter = adapter
self.session = session if session is not None else requests.Session()
if not persist_cookies and self._owns_session:
self.session.cookies.set_policy(_RejectCookiesPolicy())
Expand All @@ -215,6 +224,7 @@ def __init__(

def close(self) -> None:
if self._owns_session:
_unmount_adapter(self.session, self._injected_adapter)
self.session.close()

def __enter__(self) -> "Transport":
Expand Down Expand Up @@ -272,7 +282,7 @@ def request(
**kwargs,
)
except requests.exceptions.RequestException as exc:
if not _is_retryable_request_exception(exc):
if not is_retryable_request_exception(exc):
error = BraintrustTransportError(method=method, url=url, attempts=attempt, retryable=False)
raise error from exc
if attempt >= max_attempts:
Expand Down Expand Up @@ -392,14 +402,24 @@ def _retry_delay(policy: RetryPolicy, attempt: int, retry_after: float | None) -
return min(policy.max_backoff, policy.backoff_factor * (2 ** (attempt - 1)))


def _request_body_is_replayable(data: Any, files: Any) -> bool:
return files is None and (data is None or isinstance(data, (bytes, str)))
def _unmount_adapter(session: requests.Session, adapter: HTTPAdapter | None) -> None:
"""Detach a caller-owned adapter so ``Session.close()`` leaves it open.

``requests.Session.close()`` closes every mounted adapter. A single adapter
installed via ``set_http_adapter`` is mounted on many sessions at once, so
closing one session would otherwise clear the connection pools that the
other sessions are still using.
"""

if adapter is None:
return
for prefix, mounted in list(session.adapters.items()):
if mounted is adapter:
del session.adapters[prefix]

def _is_retryable_request_exception(exc: requests.exceptions.RequestException) -> bool:
return isinstance(exc, (requests.exceptions.ConnectionError, requests.exceptions.Timeout)) and not isinstance(
exc, requests.exceptions.SSLError
)

def _request_body_is_replayable(data: Any, files: Any) -> bool:
return files is None and (data is None or isinstance(data, (bytes, str)))


def _parse_retry_after(value: str | None, wall_time: float) -> float | None:
Expand Down
9 changes: 9 additions & 0 deletions py/src/braintrust/api/policies.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
import enum
from dataclasses import dataclass

import requests


DEFAULT_RETRYABLE_STATUSES = frozenset({408, 429, 500, 502, 503, 504})
DEFAULT_MAX_ATTEMPTS = 4
Expand All @@ -11,6 +13,13 @@
DEFAULT_MAX_BACKOFF = 10.0


def is_retryable_request_exception(exc: requests.exceptions.RequestException) -> bool:
"""Return whether a requests transport failure is safe to retry."""
return isinstance(exc, (requests.exceptions.ConnectionError, requests.exceptions.Timeout)) and not isinstance(
exc, requests.exceptions.SSLError
)


class RetryMode(enum.Enum):
"""The replay safety classification for an API operation."""

Expand Down
62 changes: 57 additions & 5 deletions py/src/braintrust/api/test_transport.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import datetime
import io
from email.utils import format_datetime
from unittest import mock

import pytest
import requests
Expand All @@ -14,7 +15,7 @@
RetryPolicy,
)
from braintrust.api._test_server import scripted_server
from braintrust.api._transport import Transport
from braintrust.api._transport import HTTPConnection, Transport
from braintrust.util import AugmentedHTTPError
from requests.adapters import HTTPAdapter
from urllib3.util.retry import Retry
Expand Down Expand Up @@ -68,13 +69,17 @@ def close(self):
super().close()


def test_transport_closes_owned_session():
def test_transport_closes_owned_session_without_closing_injected_adapter():
adapter = TrackingAdapter()
transport = Transport(adapter=adapter)
session = transport.session

with Transport(adapter=adapter) as transport:
assert transport.session is not None
with mock.patch.object(session, "close", wraps=session.close) as close_spy:
transport.close()

assert adapter.close_count > 0
close_spy.assert_called_once()
# The adapter belongs to the caller and may be mounted on other sessions.
assert adapter.close_count == 0


def test_transport_does_not_close_injected_session():
Expand Down Expand Up @@ -351,3 +356,50 @@ def test_non_retrying_custom_adapter_can_delegate_retries_to_sdk():

assert response.status_code == 200
assert handler.request_count == 2


def test_http_connection_close_does_not_close_shared_adapter():
adapter = TrackingAdapter()
first = HTTPConnection("http://localhost", adapter=adapter)
second = HTTPConnection("http://localhost", adapter=adapter)

first.close()

assert adapter.close_count == 0
assert second.session.get_adapter("http://localhost") is adapter


def test_http_connection_close_closes_self_created_long_lived_adapter():
conn = HTTPConnection("http://localhost")
conn.make_long_lived()
adapter = conn.adapter
assert adapter is not None

with mock.patch.object(adapter, "close", wraps=adapter.close) as close_spy:
conn.close()

# Mounted on both the http:// and https:// prefixes, so closed once per mount.
assert close_spy.call_count > 0


def test_http_connection_close_closes_long_lived_adapter_replaced_by_set_adapter():
adapter = TrackingAdapter()
conn = HTTPConnection("http://localhost")
conn.make_long_lived()
conn._set_adapter(adapter)
conn._reset()

conn.close()

assert adapter.close_count == 0


def test_http_connection_close_does_not_close_adapter_set_after_construction():
adapter = TrackingAdapter()
conn = HTTPConnection("http://localhost")
conn._set_adapter(adapter)
conn._reset()

conn.close()

assert adapter.close_count == 0
Loading