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
12 changes: 6 additions & 6 deletions src/browserbase/_response.py
Original file line number Diff line number Diff line change
Expand Up @@ -629,12 +629,12 @@ class AsyncResponseContextManager(Generic[_AsyncAPIResponseT]):
when the context manager exits
"""

def __init__(self, api_request: Awaitable[_AsyncAPIResponseT]) -> None:
def __init__(self, api_request: Callable[[], Awaitable[_AsyncAPIResponseT]]) -> None:
self._api_request = api_request
self.__response: _AsyncAPIResponseT | None = None

async def __aenter__(self) -> _AsyncAPIResponseT:
self.__response = await self._api_request
self.__response = await self._api_request()
return self.__response

async def __aexit__(
Expand Down Expand Up @@ -680,9 +680,9 @@ def wrapped(*args: P.args, **kwargs: P.kwargs) -> AsyncResponseContextManager[As

kwargs["extra_headers"] = extra_headers

make_request = func(*args, **kwargs)
make_request = functools.partial(func, *args, **kwargs)

return AsyncResponseContextManager(cast(Awaitable[AsyncAPIResponse[R]], make_request))
return AsyncResponseContextManager(cast(Callable[[], Awaitable[AsyncAPIResponse[R]]], make_request))

return wrapped

Expand Down Expand Up @@ -730,9 +730,9 @@ def wrapped(*args: P.args, **kwargs: P.kwargs) -> AsyncResponseContextManager[_A

kwargs["extra_headers"] = extra_headers

make_request = func(*args, **kwargs)
make_request = functools.partial(func, *args, **kwargs)

return AsyncResponseContextManager(cast(Awaitable[_AsyncAPIResponseT], make_request))
return AsyncResponseContextManager(cast(Callable[[], Awaitable[_AsyncAPIResponseT]], make_request))

return wrapped

Expand Down
52 changes: 52 additions & 0 deletions tests/test_response.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,8 @@
from __future__ import annotations

import gc
import json
import warnings
from typing import Any, List, Union, cast
from typing_extensions import Annotated

Expand All @@ -14,6 +18,7 @@
BinaryAPIResponse,
AsyncBinaryAPIResponse,
extract_response_type,
async_to_streamed_response_wrapper,
)
from browserbase._streaming import Stream
from browserbase._base_client import FinalRequestOptions
Expand All @@ -28,6 +33,53 @@ class ConcreteAPIResponse(APIResponse[List[str]]): ...
class ConcreteAsyncAPIResponse(APIResponse[httpx.Response]): ...


@pytest.mark.asyncio
async def test_async_streaming_response_request_is_lazy(async_client: AsyncBrowserbase) -> None:
calls = 0

async def request(*, extra_headers: dict[str, str] | None = None) -> AsyncAPIResponse[str]:
nonlocal calls
calls += 1
assert extra_headers == {"X-Stainless-Raw-Response": "stream"}
return AsyncAPIResponse(
raw=httpx.Response(200, content=b"response"),
client=async_client,
stream=False,
stream_cls=None,
cast_to=str,
options=FinalRequestOptions.construct(method="get", url="/foo"),
)

wrapped = async_to_streamed_response_wrapper(request)
context_manager = wrapped()

assert calls == 0
async with context_manager as response:
assert calls == 1
assert await response.text() == "response"

assert response.is_closed


@pytest.mark.asyncio
@pytest.mark.parametrize("custom_response", [False, True])
async def test_discarding_async_streaming_response_does_not_warn(
async_client: AsyncBrowserbase, custom_response: bool
) -> None:
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
if custom_response:
context_manager = async_client.sessions.replays.with_streaming_response.retrieve_page(
id="session-id", page_id="page-id"
)
else:
context_manager = async_client.sessions.with_streaming_response.create()
del context_manager
gc.collect()

assert not caught


def test_extract_response_type_direct_classes() -> None:
assert extract_response_type(BaseAPIResponse[str]) == str
assert extract_response_type(APIResponse[str]) == str
Expand Down