|
40 | 40 | "event_hooks": "event_hooks=... is not supported; use middleware=... instead.", |
41 | 41 | } |
42 | 42 | _TOO_MANY_REDIRECTS_MESSAGE = "Exceeded maximum allowed redirects." |
| 43 | +_NO_AUTH = httpx2.Auth() |
43 | 44 | _BASE_URL_QUERY_MESSAGE = ( |
44 | 45 | "base_url must not contain a query string: httpx2 appends request paths after it, " |
45 | 46 | "producing malformed URLs. Pass the query as params=... instead." |
@@ -133,58 +134,153 @@ def _select_httpx2_options( |
133 | 134 | return forwarded |
134 | 135 |
|
135 | 136 |
|
136 | | -async def _send_following_redirects_async(client: httpx2.AsyncClient, request: httpx2.Request) -> httpx2.Response: |
137 | | - """Send `request` streaming, following redirects hop by hop without reading intermediate bodies.""" |
138 | | - history: list[httpx2.Response] = [] |
139 | | - response = await client.send(request, stream=True, follow_redirects=False) |
140 | | - while client.follow_redirects and response.next_request is not None: |
| 137 | +def _request_auth(client: httpx2.Client | httpx2.AsyncClient, request: httpx2.Request) -> httpx2.Auth: |
| 138 | + """Return the auth httpx2 applies to `request`: the client's, else Basic from URL credentials, else none.""" |
| 139 | + if client.auth is not None: |
| 140 | + return client.auth |
| 141 | + if request.url.username or request.url.password: |
| 142 | + return httpx2.BasicAuth(request.url.username, request.url.password) |
| 143 | + return _NO_AUTH |
| 144 | + |
| 145 | + |
| 146 | +async def _read_capped_and_close_async(streaming: httpx2.Response, cap: int) -> httpx2.Response: |
| 147 | + """Buffer `streaming` under `cap` via `_read_capped_async`, closing it either way.""" |
| 148 | + try: |
| 149 | + return await _read_capped_async(streaming, cap, streaming.request) |
| 150 | + finally: |
| 151 | + await streaming.aclose() |
| 152 | + |
| 153 | + |
| 154 | +async def _send_redirect_hops_async( |
| 155 | + client: httpx2.AsyncClient, |
| 156 | + request: httpx2.Request, |
| 157 | + prior_history: list[httpx2.Response], |
| 158 | +) -> httpx2.Response: |
| 159 | + """Send `request` streaming, following redirects hop by hop without reading intermediate bodies. |
| 160 | +
|
| 161 | + Histories and the `max_redirects` count include `prior_history`, as in httpx2. |
| 162 | + """ |
| 163 | + hops = list(prior_history) |
| 164 | + while True: |
| 165 | + if len(hops) > client.max_redirects: |
| 166 | + raise httpx2.TooManyRedirects(_TOO_MANY_REDIRECTS_MESSAGE, request=request) |
| 167 | + response = await client.send(request, stream=True, follow_redirects=False, auth=_NO_AUTH) |
| 168 | + response.history = list(hops) |
| 169 | + if not client.follow_redirects or response.next_request is None: |
| 170 | + return response |
141 | 171 | await response.aclose() |
142 | | - history.append(response) |
143 | | - if len(history) > client.max_redirects: |
144 | | - raise httpx2.TooManyRedirects(_TOO_MANY_REDIRECTS_MESSAGE, request=response.next_request) |
145 | | - response = await client.send(response.next_request, stream=True, follow_redirects=False, auth=None) |
146 | | - response.history = history |
147 | | - return response |
| 172 | + hops.append(response) |
| 173 | + request = response.next_request |
| 174 | + |
| 175 | + |
| 176 | +async def _send_capped_async(client: httpx2.AsyncClient, request: httpx2.Request, cap: int) -> httpx2.Response: |
| 177 | + """Send `request` streaming, driving the client's auth flow and redirects without reading intermediate bodies. |
| 178 | +
|
| 179 | + An auth that sets `requires_response_body` gets each response buffered under `cap` instead. |
| 180 | + """ |
| 181 | + auth = _request_auth(client, request) |
| 182 | + flow = auth.async_auth_flow(request) |
| 183 | + history: list[httpx2.Response] = [] |
| 184 | + try: |
| 185 | + request = await anext(flow) |
| 186 | + while True: |
| 187 | + response = await _send_redirect_hops_async(client, request, history) |
| 188 | + if auth.requires_response_body: |
| 189 | + response = await _read_capped_and_close_async(response, cap) |
| 190 | + try: |
| 191 | + next_request = await flow.asend(response) |
| 192 | + except StopAsyncIteration: |
| 193 | + return response |
| 194 | + except BaseException: |
| 195 | + await response.aclose() |
| 196 | + raise |
| 197 | + await response.aclose() |
| 198 | + response.history = list(history) |
| 199 | + history.append(response) |
| 200 | + request = next_request |
| 201 | + finally: |
| 202 | + await flow.aclose() |
148 | 203 |
|
149 | 204 |
|
150 | 205 | @contextlib.asynccontextmanager |
151 | | -async def _stream_following_redirects_async( |
| 206 | +async def _stream_capped_async( |
152 | 207 | client: httpx2.AsyncClient, |
153 | 208 | method: str, |
154 | 209 | url: httpx2.URL | str, |
155 | 210 | kwargs: dict[str, typing.Any], |
| 211 | + cap: int, |
156 | 212 | ) -> AsyncIterator[httpx2.Response]: |
157 | | - """Async mirror of `httpx2.AsyncClient.stream` that follows redirects via `_send_following_redirects_async`.""" |
158 | | - response = await _send_following_redirects_async(client, client.build_request(method, url, **kwargs)) |
| 213 | + """Async mirror of `httpx2.AsyncClient.stream` that sends via `_send_capped_async`.""" |
| 214 | + response = await _send_capped_async(client, client.build_request(method, url, **kwargs), cap) |
159 | 215 | try: |
160 | 216 | yield response |
161 | 217 | finally: |
162 | 218 | await response.aclose() |
163 | 219 |
|
164 | 220 |
|
165 | | -def _send_following_redirects(client: httpx2.Client, request: httpx2.Request) -> httpx2.Response: |
166 | | - """Sync mirror of `_send_following_redirects_async`.""" |
167 | | - history: list[httpx2.Response] = [] |
168 | | - response = client.send(request, stream=True, follow_redirects=False) |
169 | | - while client.follow_redirects and response.next_request is not None: |
| 221 | +def _read_capped_and_close(streaming: httpx2.Response, cap: int) -> httpx2.Response: |
| 222 | + """Sync mirror of `_read_capped_and_close_async`.""" |
| 223 | + try: |
| 224 | + return _read_capped(streaming, cap, streaming.request) |
| 225 | + finally: |
| 226 | + streaming.close() |
| 227 | + |
| 228 | + |
| 229 | +def _send_redirect_hops( |
| 230 | + client: httpx2.Client, |
| 231 | + request: httpx2.Request, |
| 232 | + prior_history: list[httpx2.Response], |
| 233 | +) -> httpx2.Response: |
| 234 | + """Sync mirror of `_send_redirect_hops_async`.""" |
| 235 | + hops = list(prior_history) |
| 236 | + while True: |
| 237 | + if len(hops) > client.max_redirects: |
| 238 | + raise httpx2.TooManyRedirects(_TOO_MANY_REDIRECTS_MESSAGE, request=request) |
| 239 | + response = client.send(request, stream=True, follow_redirects=False, auth=_NO_AUTH) |
| 240 | + response.history = list(hops) |
| 241 | + if not client.follow_redirects or response.next_request is None: |
| 242 | + return response |
170 | 243 | response.close() |
171 | | - history.append(response) |
172 | | - if len(history) > client.max_redirects: |
173 | | - raise httpx2.TooManyRedirects(_TOO_MANY_REDIRECTS_MESSAGE, request=response.next_request) |
174 | | - response = client.send(response.next_request, stream=True, follow_redirects=False, auth=None) |
175 | | - response.history = history |
176 | | - return response |
| 244 | + hops.append(response) |
| 245 | + request = response.next_request |
| 246 | + |
| 247 | + |
| 248 | +def _send_capped(client: httpx2.Client, request: httpx2.Request, cap: int) -> httpx2.Response: |
| 249 | + """Sync mirror of `_send_capped_async`.""" |
| 250 | + auth = _request_auth(client, request) |
| 251 | + flow = auth.sync_auth_flow(request) |
| 252 | + history: list[httpx2.Response] = [] |
| 253 | + try: |
| 254 | + request = next(flow) |
| 255 | + while True: |
| 256 | + response = _send_redirect_hops(client, request, history) |
| 257 | + if auth.requires_response_body: |
| 258 | + response = _read_capped_and_close(response, cap) |
| 259 | + try: |
| 260 | + next_request = flow.send(response) |
| 261 | + except StopIteration: |
| 262 | + return response |
| 263 | + except BaseException: |
| 264 | + response.close() |
| 265 | + raise |
| 266 | + response.close() |
| 267 | + response.history = list(history) |
| 268 | + history.append(response) |
| 269 | + request = next_request |
| 270 | + finally: |
| 271 | + flow.close() |
177 | 272 |
|
178 | 273 |
|
179 | 274 | @contextlib.contextmanager |
180 | | -def _stream_following_redirects( |
| 275 | +def _stream_capped( |
181 | 276 | client: httpx2.Client, |
182 | 277 | method: str, |
183 | 278 | url: httpx2.URL | str, |
184 | 279 | kwargs: dict[str, typing.Any], |
| 280 | + cap: int, |
185 | 281 | ) -> Iterator[httpx2.Response]: |
186 | | - """Sync mirror of `_stream_following_redirects_async`.""" |
187 | | - response = _send_following_redirects(client, client.build_request(method, url, **kwargs)) |
| 282 | + """Sync mirror of `_stream_capped_async`.""" |
| 283 | + response = _send_capped(client, client.build_request(method, url, **kwargs), cap) |
188 | 284 | try: |
189 | 285 | yield response |
190 | 286 | finally: |
@@ -284,11 +380,8 @@ async def _terminal(self, request: httpx2.Request) -> httpx2.Response: |
284 | 380 | if cap is None: |
285 | 381 | response = await self._httpx2_client.send(request) |
286 | 382 | else: |
287 | | - streaming = await _send_following_redirects_async(self._httpx2_client, request) |
288 | | - try: |
289 | | - response = await _read_capped_async(streaming, cap, streaming.request) |
290 | | - finally: |
291 | | - await streaming.aclose() |
| 383 | + streaming = await _send_capped_async(self._httpx2_client, request, cap) |
| 384 | + response = await _read_capped_and_close_async(streaming, cap) |
292 | 385 | except RuntimeError as exc: |
293 | 386 | if self._httpx2_client.is_closed: |
294 | 387 | raise TransportError(str(exc)) from exc |
@@ -1140,7 +1233,7 @@ async def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwa |
1140 | 1233 | opened = ( |
1141 | 1234 | self._httpx2_client.stream(method, merged_url, **kwargs) |
1142 | 1235 | if cap is None |
1143 | | - else _stream_following_redirects_async(self._httpx2_client, method, merged_url, kwargs) |
| 1236 | + else _stream_capped_async(self._httpx2_client, method, merged_url, kwargs, cap) |
1144 | 1237 | ) |
1145 | 1238 | async with _httpx2_exception_mapper(), opened as response: |
1146 | 1239 | if HTTPStatus.BAD_REQUEST <= response.status_code < 600: # noqa: PLR2004 — 600 is the synthetic upper bound for 5xx |
@@ -1225,11 +1318,8 @@ def _terminal(self, request: httpx2.Request) -> httpx2.Response: |
1225 | 1318 | if cap is None: |
1226 | 1319 | response = self._httpx2_client.send(request) |
1227 | 1320 | else: |
1228 | | - streaming = _send_following_redirects(self._httpx2_client, request) |
1229 | | - try: |
1230 | | - response = _read_capped(streaming, cap, streaming.request) |
1231 | | - finally: |
1232 | | - streaming.close() |
| 1321 | + streaming = _send_capped(self._httpx2_client, request, cap) |
| 1322 | + response = _read_capped_and_close(streaming, cap) |
1233 | 1323 | except RuntimeError as exc: |
1234 | 1324 | if self._httpx2_client.is_closed: |
1235 | 1325 | raise TransportError(str(exc)) from exc |
@@ -2102,7 +2192,7 @@ def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwargs-fo |
2102 | 2192 | opened = ( |
2103 | 2193 | self._httpx2_client.stream(method, merged_url, **kwargs) |
2104 | 2194 | if cap is None |
2105 | | - else _stream_following_redirects(self._httpx2_client, method, merged_url, kwargs) |
| 2195 | + else _stream_capped(self._httpx2_client, method, merged_url, kwargs, cap) |
2106 | 2196 | ) |
2107 | 2197 | with _httpx2_exception_mapper_sync(), opened as response: |
2108 | 2198 | if HTTPStatus.BAD_REQUEST <= response.status_code < 600: # noqa: PLR2004 — 600 is the synthetic upper bound for 5xx |
|
0 commit comments