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
3 changes: 1 addition & 2 deletions CHANGES.rst
Original file line number Diff line number Diff line change
Expand Up @@ -472,8 +472,7 @@ Features



- Added :attr:`~aiohttp.ClientResponse.output_size` and
:attr:`~aiohttp.ClientResponse.upload_complete` -- by :user:`Dreamsorcerer`.
- Added ``ClientResponse.output_size`` and ``ClientResponse.upload_complete`` -- by :user:`Dreamsorcerer`.


*Related issues and pull requests on GitHub:*
Expand Down
4 changes: 4 additions & 0 deletions CHANGES/13579.bugfix.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
Fixed a connection being eligible for reuse after its request was cancelled
or failed while waiting for a ``100 Continue`` response or finalizing the
body; the request headers were already sent, so reusing the connection
corrupted the next request on it -- by :user:`Dreamsorcerer`.
2 changes: 2 additions & 0 deletions CHANGES/13579.deprecation.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
Deprecated ``ClientResponse.output_size`` and ``ClientResponse.upload_complete``;
use ``aiohttp.UploadTracker`` instead -- by :user:`Dreamsorcerer`.
1 change: 1 addition & 0 deletions CHANGES/13579.feature.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Added :class:`aiohttp.UploadTracker` for observing a client request's upload progress -- by :user:`Dreamsorcerer`.
2 changes: 1 addition & 1 deletion THREAT_MODEL.md
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@ Key public APIs (non-exhaustive):
| Surface | Entry points |
| --- | --- |
| Server | `aiohttp.web.Application`, `web.RouteTableDef`, `web.run_app`, `web.AppRunner`, `web.WebSocketResponse`, `web.FileResponse` |
| Client | `aiohttp.ClientSession`, `aiohttp.TCPConnector`, `aiohttp.ClientResponse`, `aiohttp.WSMessage`, `aiohttp.BasicAuth` |
| Client | `aiohttp.ClientSession`, `aiohttp.TCPConnector`, `aiohttp.ClientResponse`, `aiohttp.UploadTracker`, `aiohttp.WSMessage`, `aiohttp.BasicAuth` |
| Shared | `aiohttp.MultipartReader`/`MultipartWriter`, `aiohttp.CookieJar`, `aiohttp.TraceConfig`, `aiohttp.resolver.AsyncResolver` |

---
Expand Down
4 changes: 4 additions & 0 deletions aiohttp/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,8 @@
TCPConnector,
TooManyRedirects,
UnixConnector,
UploadAbortedError,
UploadTracker,
WSMessageTypeError,
WSServerHandshakeError,
request,
Expand Down Expand Up @@ -156,6 +158,8 @@
"TCPConnector",
"TooManyRedirects",
"UnixConnector",
"UploadAbortedError",
"UploadTracker",
"NamedPipeConnector",
"WSServerHandshakeError",
"request",
Expand Down
17 changes: 17 additions & 0 deletions aiohttp/abc.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import asyncio
import logging
import socket
from abc import ABC, abstractmethod
Expand Down Expand Up @@ -209,6 +210,22 @@ class AbstractStreamWriter(ABC):
buffer_size: int = 0
output_size: int = 0
length: int | None = 0
# Called with each accepted body chunk's byte count (before any
# transport-level transformation such as compression or chunked
# framing) and the writer's total output_size after the chunk.
# Assigned by the client request machinery for upload progress
# tracking; write()/write_eof() implementations should invoke it
# for every body chunk they accept.
on_body_write: Callable[[int, int], None] | None = None

@property
def transport(self) -> asyncio.WriteTransport | None:
"""The transport this writer writes to, if any.

Used by upload progress tracking to observe the unsent buffer;
writers without one report progress only on completion.
"""
return None

@abstractmethod
async def write(
Expand Down
206 changes: 116 additions & 90 deletions aiohttp/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,7 @@
ServerTimeoutError,
SocketTimeoutError,
TooManyRedirects,
UploadAbortedError,
WSMessageTypeError,
WSServerHandshakeError,
)
Expand All @@ -77,6 +78,7 @@
Fingerprint,
RequestInfo,
ResponseParams,
UploadTracker,
)
from .client_ws import (
DEFAULT_WS_CLIENT_TIMEOUT,
Expand Down Expand Up @@ -137,12 +139,14 @@
"ServerTimeoutError",
"SocketTimeoutError",
"TooManyRedirects",
"UploadAbortedError",
"WSServerHandshakeError",
# client_reqrep
"ClientRequest",
"ClientResponse",
"Fingerprint",
"RequestInfo",
"UploadTracker",
# connector
"BaseConnector",
"TCPConnector",
Expand Down Expand Up @@ -210,6 +214,7 @@ class _RequestOptions(TypedDict, total=False):
max_field_size: int | None
max_headers: int | None
middlewares: Sequence[ClientMiddlewareType] | None
upload_tracker: UploadTracker | None


class _WSConnectOptions(TypedDict, total=False):
Expand Down Expand Up @@ -250,6 +255,8 @@ class _WSConnectOptions(TypedDict, total=False):
async def _connect_and_send_request(req: ClientRequest) -> ClientResponse:
connector = req._session._connector
assert connector is not None
if (tracker := req._upload_tracker) is not None:
req._upload_gen = tracker._attempt_started()
try:
conn = await connector.connect(req, traces=req._traces, timeout=req._timeout)
except asyncio.TimeoutError as exc:
Expand Down Expand Up @@ -510,116 +517,128 @@ async def _request(
max_field_size: int | None = None,
max_headers: int | None = None,
middlewares: Sequence[ClientMiddlewareType] | None = None,
upload_tracker: UploadTracker | None = None,
) -> ClientResponse:
# NOTE: timeout clamps existing connect and read timeouts. We cannot
# set the default to None because we need to detect if the user wants
# to use the existing timeouts by setting timeout to None.

if self.closed:
raise RuntimeError("Session is closed")
# Bound outside the settle-guaranteeing try block: a rebind error
# must not settle a tracker owned by another request.
if upload_tracker is not None:
upload_tracker._bind()

method = method.upper()
tm: TimeoutHandle | None = None
handle: asyncio.TimerHandle | None = None
# Only traces that saw send_request_start; they must also see a terminal event.
traces: list[Trace] = []
req: ClientRequest | None = None
resp: ClientResponse | None = None
try:
if self.closed:
raise RuntimeError("Session is closed")

if ssl is sentinel:
ssl = self._default_ssl
if not isinstance(ssl, SSL_ALLOWED_TYPES):
raise TypeError(
"ssl should be SSLContext, Fingerprint, or bool, "
f"got {ssl!r} instead."
)
method = method.upper()

if data is not None and json is not None:
raise ValueError(
"data and json parameters can not be used at the same time"
)
elif json is not None:
if self._json_serialize_bytes is not None:
data = payload.JsonBytesPayload(json, dumps=self._json_serialize_bytes)
else:
data = payload.JsonPayload(json, dumps=self._json_serialize)

redirects = 0
history: list[ClientResponse] = []
version = self._version
params = params or {}
if ssl is sentinel:
ssl = self._default_ssl
if not isinstance(ssl, SSL_ALLOWED_TYPES):
raise TypeError(
"ssl should be SSLContext, Fingerprint, or bool, "
f"got {ssl!r} instead."
)

# Merge with default headers and transform to CIMultiDict
headers = self._prepare_headers(headers)
if data is not None and json is not None:
raise ValueError(
"data and json parameters can not be used at the same time"
)
elif json is not None:
if self._json_serialize_bytes is not None:
data = payload.JsonBytesPayload(
json, dumps=self._json_serialize_bytes
)
else:
data = payload.JsonPayload(json, dumps=self._json_serialize)

try:
url = self._build_url(str_or_url)
except ValueError as e:
raise InvalidUrlClientError(str_or_url) from e
redirects = 0
history: list[ClientResponse] = []
version = self._version
params = params or {}

assert self._connector is not None
if url.scheme not in self._connector.allowed_protocol_schema_set:
raise NonHttpUrlClientError(url)
# Merge with default headers and transform to CIMultiDict
headers = self._prepare_headers(headers)

skip_headers: Iterable[istr] | None
if skip_auto_headers is not None:
skip_headers = {
istr(i) for i in skip_auto_headers
} | self._skip_auto_headers
elif self._skip_auto_headers:
skip_headers = self._skip_auto_headers
else:
skip_headers = None

if proxy is None:
proxy = self._default_proxy

resolved_proxy_headers: CIMultiDict[str] | None
if proxy is None:
resolved_proxy_headers = None
else:
resolved_proxy_headers = self._prepare_headers(proxy_headers)
try:
proxy = URL(proxy)
url = self._build_url(str_or_url)
except ValueError as e:
raise InvalidURL(proxy) from e
raise InvalidUrlClientError(str_or_url) from e

assert self._connector is not None
if url.scheme not in self._connector.allowed_protocol_schema_set:
raise NonHttpUrlClientError(url)

skip_headers: Iterable[istr] | None
if skip_auto_headers is not None:
skip_headers = {
istr(i) for i in skip_auto_headers
} | self._skip_auto_headers
elif self._skip_auto_headers:
skip_headers = self._skip_auto_headers
else:
skip_headers = None

if timeout is sentinel or timeout is None:
real_timeout: ClientTimeout = self._timeout
else:
real_timeout = timeout
# timeout is cumulative for all request operations
# (request, redirects, responses, data consuming)
tm = TimeoutHandle(
self._loop, real_timeout.total, ceil_threshold=real_timeout.ceil_threshold
)
handle = tm.start()
if proxy is None:
proxy = self._default_proxy

if read_bufsize is None:
read_bufsize = self._read_bufsize
resolved_proxy_headers: CIMultiDict[str] | None
if proxy is None:
resolved_proxy_headers = None
else:
resolved_proxy_headers = self._prepare_headers(proxy_headers)
try:
proxy = URL(proxy)
except ValueError as e:
raise InvalidURL(proxy) from e

real_timeout = (
self._timeout if timeout is sentinel or timeout is None else timeout
)
# timeout is cumulative for all request operations
# (request, redirects, responses, data consuming)
tm = TimeoutHandle(
self._loop,
real_timeout.total,
ceil_threshold=real_timeout.ceil_threshold,
)
handle = tm.start()

if auto_decompress is None:
auto_decompress = self._auto_decompress
if read_bufsize is None:
read_bufsize = self._read_bufsize

if max_line_size is None:
max_line_size = self._max_line_size
if auto_decompress is None:
auto_decompress = self._auto_decompress

if max_field_size is None:
max_field_size = self._max_field_size
if max_line_size is None:
max_line_size = self._max_line_size

if max_headers is None:
max_headers = self._max_headers
if max_field_size is None:
max_field_size = self._max_field_size

traces = [
Trace(
self,
trace_config,
trace_config.trace_config_ctx(trace_request_ctx=trace_request_ctx),
)
for trace_config in self._trace_configs
]
if max_headers is None:
max_headers = self._max_headers

for trace in traces:
await trace.send_request_start(method, url.update_query(params), headers)
for trace_config in self._trace_configs:
trace = Trace(
self,
trace_config,
trace_config.trace_config_ctx(trace_request_ctx=trace_request_ctx),
)
await trace.send_request_start(
method, url.update_query(params), headers
)
traces.append(trace)

timer = tm.timer()
req: ClientRequest | None = None
resp: ClientResponse | None = None
try:
timer = tm.timer()
with timer:
# https://www.rfc-editor.org/rfc/rfc9112.html#name-retrying-requests
retry_persistent_connection = (
Expand Down Expand Up @@ -727,6 +746,7 @@ async def _request(
traces=traces,
trust_env=self.trust_env,
)
req._upload_tracker = upload_tracker

# Apply middleware (if any) - per-request middleware overrides session middleware
effective_middlewares = (
Expand Down Expand Up @@ -905,14 +925,19 @@ async def _request(
await trace.send_request_end(
method, url.update_query(params), headers, resp
)
if upload_tracker is not None:
upload_tracker._finalize()
return resp

except BaseException as e:
# cleanup timer
tm.close()
if tm is not None:
tm.close()
if handle:
handle.cancel()
handle = None

if upload_tracker is not None:
upload_tracker._finalize()

if resp is not None:
# A failure occurred after the response was received.
Expand All @@ -922,8 +947,9 @@ async def _request(
await req._body.close()

for trace in traces:
# url and headers are bound whenever traces is non-empty.
await trace.send_request_exception(
method, url.update_query(params), headers, e
method, url.update_query(params), headers, e # type: ignore[possibly-undefined, arg-type]
)
raise

Expand Down
5 changes: 5 additions & 0 deletions aiohttp/client_exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
"WSServerHandshakeError",
"ContentTypeError",
"ClientPayloadError",
"UploadAbortedError",
"InvalidURL",
"InvalidUrlClientError",
"RedirectClientError",
Expand Down Expand Up @@ -261,6 +262,10 @@ class ClientPayloadError(ClientError):
"""Response payload error."""


class UploadAbortedError(ClientError):
"""The request body was never fully sent."""


class InvalidURL(ClientError, ValueError):
"""Invalid URL.

Expand Down
Loading
Loading