diff --git a/.codegen.json b/.codegen.json index 1c2e599b8..1153aa33b 100644 --- a/.codegen.json +++ b/.codegen.json @@ -1 +1 @@ -{ "engineHash": "12fa0b7", "specHash": "c5a35a6", "version": "4.17.0" } +{ "engineHash": "fd56d95", "specHash": "c5a35a6", "version": "4.17.0" } diff --git a/box_sdk_gen/networking/multipart_stream.py b/box_sdk_gen/networking/multipart_stream.py index 22c240cd1..78be03017 100644 --- a/box_sdk_gen/networking/multipart_stream.py +++ b/box_sdk_gen/networking/multipart_stream.py @@ -1,4 +1,4 @@ -from io import SEEK_END +from io import SEEK_END, TextIOBase from typing import Iterator, List, Optional, Tuple, Union from urllib3.fields import RequestField @@ -8,7 +8,9 @@ CHUNK_SIZE = 64 * 1024 -MultipartField = Tuple[str, Optional[str], Union[str, ByteStream], Optional[str]] +PartStream = Union[ByteStream, TextIOBase] + +MultipartField = Tuple[str, Optional[str], Union[str, PartStream], Optional[str]] class MultipartStream: @@ -17,13 +19,14 @@ class MultipartStream: so uploads are sent without buffering whole files in memory. Fields are (name, file_name, value, content_type) tuples, where value is - either a string or a binary stream read from its current position. + either a string or a stream read from its current position. Text streams + are encoded as UTF-8. """ def __init__(self, fields: List[MultipartField]): self.boundary = choose_boundary() self.content_type = f'multipart/form-data; boundary={self.boundary}' - self._segments: List[Union[bytes, ByteStream]] = [] + self._segments: List[Union[bytes, PartStream]] = [] for name, file_name, value, content_type in fields: field = RequestField(name=name, data=b'', filename=file_name) field.make_multipart(content_type=content_type) @@ -37,6 +40,8 @@ def __init__(self, fields: List[MultipartField]): self._segments.append(f'--{self.boundary}--\r\n'.encode('utf-8')) self._index = 0 self._offset = 0 + # encoded text stream bytes that didn't fit in the last read + self._pending = b'' # set when a part stream ends before its declared size; retrying won't help self.size_error: Optional[IOError] = None # bytes still to send per stream segment, None when the size is unknown @@ -49,17 +54,31 @@ def __init__(self, fields: List[MultipartField]): self.len = self._compute_length() @staticmethod - def _stream_size(stream: ByteStream) -> Optional[int]: + def _stream_size(stream: PartStream) -> Optional[int]: try: - if not stream.seekable(): + # text stream positions count characters, not encoded bytes + if isinstance(stream, TextIOBase) or not stream.seekable(): return None position = stream.tell() - stream.seek(0, SEEK_END) - end = stream.tell() - stream.seek(position) + # read(0) also catches text-like streams that aren't a TextIOBase + is_text = isinstance(stream.read(0), str) + if stream.tell() != position: + # read(0) shouldn't move a stream, but don't skip data if it does + stream.seek(position) + if is_text: + return None + # like requests and requests-toolbelt, prefer a length the stream reports + if hasattr(stream, '__len__'): + end = len(stream) + elif getattr(stream, 'len', None) is not None: + end = stream.len + else: + stream.seek(0, SEEK_END) + end = stream.tell() + stream.seek(position) # a stream positioned past its end has nothing left to send return max(0, end - position) - except (OSError, AttributeError, TypeError): + except (OSError, AttributeError, TypeError, ValueError): return None def _compute_length(self) -> Optional[int]: @@ -80,7 +99,9 @@ def read(self, size: Optional[int] = -1) -> bytes: chunks = [] while size > 0 and self._index < len(self._segments): segment = self._segments[self._index] - if isinstance(segment, bytes): + if self._pending: + chunk, self._pending = self._pending[:size], self._pending[size:] + elif isinstance(segment, bytes): chunk = segment[self._offset : self._offset + size] self._offset += len(chunk) if self._offset >= len(segment): @@ -103,8 +124,12 @@ def read(self, size: Optional[int] = -1) -> bytes: raise self.size_error self._index += 1 continue + if isinstance(chunk, str): + chunk = chunk.encode('utf-8') if remaining is not None: self._remaining[self._index] = remaining - len(chunk) + # a character can encode to several bytes, so keep what doesn't fit + chunk, self._pending = chunk[:size], chunk[size:] chunks.append(chunk) size -= len(chunk) return b''.join(chunks) diff --git a/test/box_sdk_gen/test/box_network_client.py b/test/box_sdk_gen/test/box_network_client.py index 0af0c65c2..21abdd5b6 100644 --- a/test/box_sdk_gen/test/box_network_client.py +++ b/test/box_sdk_gen/test/box_network_client.py @@ -2,7 +2,7 @@ import json import threading from http.server import BaseHTTPRequestHandler, HTTPServer -from io import BytesIO, RawIOBase, UnsupportedOperation, SEEK_END, SEEK_SET +from io import BytesIO, RawIOBase, StringIO, UnsupportedOperation, SEEK_END, SEEK_SET from unittest import mock from unittest.mock import Mock, patch from requests import Session, Response, RequestException @@ -1513,3 +1513,119 @@ def seek(self, *args): assert len(body) == multipart_stream.len assert b"\r\n\r\n23456789\r\n" in body + + +def test_multipart_upload_text_stream_is_encoded_as_utf8(multipart_server): + url, received, _ = multipart_server + + BoxNetworkClient().fetch(_upload_options(url, StringIO("héllo wörld"))) + + assert len(received) == 1 + headers, body = received[0] + assert headers["Transfer-Encoding"] == "chunked" + assert "Content-Length" not in headers + assert "héllo wörld\r\n".encode("utf-8") in body + + +def test_multipart_stream_encodes_text_file_as_utf8(tmp_path): + path = tmp_path / "file.txt" + path.write_text("héllo wörld\n" * 10000, encoding="utf-8") + + with open(path, "r", encoding="utf-8") as text_file: + multipart_stream = MultipartStream([("file", "file.txt", text_file, None)]) + body = multipart_stream.read() + + assert multipart_stream.len is None + assert body.endswith( + b"\r\n\r\n" + + ("héllo wörld\n" * 10000).encode("utf-8") + + f"\r\n--{multipart_stream.boundary}--\r\n".encode() + ) + + +def test_multipart_stream_uses_length_reported_by_stream(): + class ReportedLength(BytesIO): + len = 10 + + def seek(self, *args): + raise AssertionError("size should come from the reported length") + + stream = ReportedLength(b"0123456789") + stream.read(2) + multipart_stream = MultipartStream([("file", "f", stream, None)]) + + body = multipart_stream.read() + + assert len(body) == multipart_stream.len + assert b"\r\n\r\n23456789\r\n" in body + + +def test_multipart_stream_read_returns_at_most_size_bytes_for_text_stream(): + text = "é" * 100 + "€" * 50 + multipart_stream = MultipartStream([("file", "f", StringIO(text), None)]) + + chunks = list(iter(lambda: multipart_stream.read(3), b"")) + + assert all(len(chunk) <= 3 for chunk in chunks) + assert b"\r\n\r\n" + text.encode("utf-8") + b"\r\n" in b"".join(chunks) + + +def test_multipart_stream_sends_text_like_stream_with_unknown_length(): + class TextLike: + def __init__(self, text): + self._stream = StringIO(text) + + def seekable(self): + return True + + def tell(self): + return self._stream.tell() + + def seek(self, *args): + return self._stream.seek(*args) + + def read(self, size=-1): + return self._stream.read(size) + + text = "héllo wörld " * 100 + multipart_stream = MultipartStream([("file", "f", TextLike(text), None)]) + + body = multipart_stream.read() + + assert multipart_stream.len is None + assert b"\r\n\r\n" + text.encode("utf-8") + b"\r\n" in body + + +def test_multipart_stream_sends_stream_failing_read_zero_with_unknown_length(): + class FailsOnEmptyRead(BytesIO): + def read(self, size=-1): + if size == 0: + raise ValueError("read(0) is not supported") + return super().read(size) + + multipart_stream = MultipartStream( + [("file", "f", FailsOnEmptyRead(b"0123456789"), None)] + ) + + body = multipart_stream.read() + + assert multipart_stream.len is None + assert b"\r\n\r\n0123456789\r\n" in body + + +def test_multipart_stream_restores_position_moved_by_read_zero(): + class MovesOnEmptyRead(BytesIO): + def read(self, size=-1): + if size == 0: + self.seek(self.tell() + 2) + return b"" + return super().read(size) + + multipart_stream = MultipartStream( + [("file", "f", MovesOnEmptyRead(b"0123456789"), None)] + ) + + body = multipart_stream.read() + + assert len(body) == multipart_stream.len + assert b"\r\n\r\n0123456789\r\n" in body diff --git a/test/boxsdk/unit/auth/test_jwt_auth.py b/test/boxsdk/unit/auth/test_jwt_auth.py index e22d8b8b5..720df5776 100644 --- a/test/boxsdk/unit/auth/test_jwt_auth.py +++ b/test/boxsdk/unit/auth/test_jwt_auth.py @@ -208,7 +208,12 @@ def _jwt_auth_init_mocks(**kwargs): backend=default_backend(), ) - yield oauth, assertion, fake_client_id, load_pem_private_key.return_value + yield ( + oauth, + assertion, + fake_client_id, + load_pem_private_key.return_value, + ) if assert_authed: mock_box_session.request.assert_called_once_with(