diff --git a/.codegen.json b/.codegen.json index b608e90a..b60bd2aa 100644 --- a/.codegen.json +++ b/.codegen.json @@ -1 +1 @@ -{ "engineHash": "daeb1ea", "specHash": "c5a35a6", "version": "4.16.0" } +{ "engineHash": "12fa0b7", "specHash": "c5a35a6", "version": "4.16.0" } diff --git a/box_sdk_gen/networking/__init__.py b/box_sdk_gen/networking/__init__.py index 84da5517..c84d3642 100644 --- a/box_sdk_gen/networking/__init__.py +++ b/box_sdk_gen/networking/__init__.py @@ -16,6 +16,8 @@ from box_sdk_gen.networking.retries import * +from box_sdk_gen.networking.multipart_stream import * + from box_sdk_gen.networking.base_urls import * from box_sdk_gen.networking.version import * diff --git a/box_sdk_gen/networking/box_network_client.py b/box_sdk_gen/networking/box_network_client.py index ee97c9b4..3c12d6e2 100644 --- a/box_sdk_gen/networking/box_network_client.py +++ b/box_sdk_gen/networking/box_network_client.py @@ -1,18 +1,17 @@ import io import time -from collections import OrderedDict from dataclasses import dataclass -from typing import Optional, Dict, Union, Tuple +from typing import Optional, Dict, Union, Tuple, List from sys import version_info as py_version import requests from requests import RequestException, Session, Response from requests.structures import CaseInsensitiveDict -from requests_toolbelt import MultipartEncoder from ..internal.logging import DataSanitizer from .retries import BoxRetryStrategy +from .multipart_stream import MultipartField, MultipartStream from ..networking.fetch_options import FetchOptions from ..networking.fetch_response import FetchResponse from ..box.errors import BoxAPIError, BoxSDKError, RequestInfo, ResponseInfo @@ -40,7 +39,7 @@ class APIRequest: url: str headers: Dict[str, str] params: Dict[str, str] - data: Optional[Union[str, ByteStream, MultipartEncoder]] + data: Optional[Union[str, ByteStream, MultipartStream]] content_type: Optional[str] = None allow_redirects: bool = True timeout: Optional[Union[float, Tuple[Optional[float], Optional[float]]]] = None @@ -160,19 +159,16 @@ def _prepare_request( if options.content_type: if options.content_type == 'multipart/form-data': - fields = OrderedDict() - for part in options.multipart_data: - if part.data: - fields[part.part_name] = sd_to_json(part.data) - else: - fields[part.part_name] = ( - part.file_name or '', - part.file_stream, - part.content_type, - ) - - multipart_stream = MultipartEncoder(fields) + multipart_stream = MultipartStream( + self._prepare_multipart_fields(options) + ) data = multipart_stream + # replace any caller-provided Content-Type, it must carry the boundary + headers = { + name: value + for name, value in headers.items() + if name.lower() != 'content-type' + } headers['Content-Type'] = multipart_stream.content_type else: headers['Content-Type'] = options.content_type @@ -188,6 +184,29 @@ def _prepare_request( timeout=timeout, ) + @staticmethod + def _prepare_multipart_fields( + options: 'FetchOptions', + ) -> List[MultipartField]: + fields = [] + for part in options.multipart_data: + if part.data is not None: + fields.append((part.part_name, None, sd_to_json(part.data), None)) + elif part.file_stream is not None: + fields.append( + ( + part.part_name, + part.file_name or '', + part.file_stream, + part.content_type, + ) + ) + else: + raise BoxSDKError( + message=f'Multipart part "{part.part_name}" has neither data nor file_stream' + ) + return fields + @staticmethod def _get_request_timeout( options: 'FetchOptions', @@ -274,6 +293,13 @@ def _make_request(self, request: APIRequest) -> APIResponse: timeout=timeout, ) except RequestException as request_exc: + if ( + isinstance(request.data, MultipartStream) + and request.data.size_error is not None + ): + raise BoxSDKError( + message=str(request.data.size_error), error=request_exc + ) from request_exc raised_exception = request_exc network_response = None diff --git a/box_sdk_gen/networking/multipart_stream.py b/box_sdk_gen/networking/multipart_stream.py new file mode 100644 index 00000000..22c240cd --- /dev/null +++ b/box_sdk_gen/networking/multipart_stream.py @@ -0,0 +1,116 @@ +from io import SEEK_END +from typing import Iterator, List, Optional, Tuple, Union + +from urllib3.fields import RequestField +from urllib3.filepost import choose_boundary + +from ..internal.utils import ByteStream + +CHUNK_SIZE = 64 * 1024 + +MultipartField = Tuple[str, Optional[str], Union[str, ByteStream], Optional[str]] + + +class MultipartStream: + """ + File-like multipart/form-data body which reads part streams lazily, + 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. + """ + + 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]] = [] + 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) + self._segments.append( + f'--{self.boundary}\r\n{field.render_headers()}'.encode('utf-8') + ) + self._segments.append( + value.encode('utf-8') if isinstance(value, str) else value + ) + self._segments.append(b'\r\n') + self._segments.append(f'--{self.boundary}--\r\n'.encode('utf-8')) + self._index = 0 + self._offset = 0 + # 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 + self._remaining: List[Optional[int]] = [ + None if isinstance(segment, bytes) else self._stream_size(segment) + for segment in self._segments + ] + # requests reads `len` to set Content-Length; None makes it fall back + # to chunked transfer encoding + self.len = self._compute_length() + + @staticmethod + def _stream_size(stream: ByteStream) -> Optional[int]: + try: + if not stream.seekable(): + return None + position = stream.tell() + 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): + return None + + def _compute_length(self) -> Optional[int]: + total = 0 + for segment, remaining in zip(self._segments, self._remaining): + if isinstance(segment, bytes): + total += len(segment) + elif remaining is None: + return None + else: + total += remaining + return total + + def read(self, size: Optional[int] = -1) -> bytes: + if size is None or size < 0: + return b''.join(iter(lambda: self.read(CHUNK_SIZE), b'')) + + chunks = [] + while size > 0 and self._index < len(self._segments): + segment = self._segments[self._index] + if isinstance(segment, bytes): + chunk = segment[self._offset : self._offset + size] + self._offset += len(chunk) + if self._offset >= len(segment): + self._index += 1 + self._offset = 0 + else: + remaining = self._remaining[self._index] + if remaining == 0: + # send exactly the size declared in Content-Length, even if the stream grew + self._index += 1 + continue + chunk = segment.read( + size if remaining is None else min(size, remaining) + ) + if not chunk: + if remaining is not None: + self.size_error = IOError( + f'Multipart stream ended {remaining} bytes before its declared size' + ) + raise self.size_error + self._index += 1 + continue + if remaining is not None: + self._remaining[self._index] = remaining - len(chunk) + chunks.append(chunk) + size -= len(chunk) + return b''.join(chunks) + + def __iter__(self) -> Iterator[bytes]: + return iter(lambda: self.read(CHUNK_SIZE), b'') + + def __repr__(self) -> str: + return f'' diff --git a/test/box_sdk_gen/test/box_network_client.py b/test/box_sdk_gen/test/box_network_client.py index f6bc2828..0af0c65c 100644 --- a/test/box_sdk_gen/test/box_network_client.py +++ b/test/box_sdk_gen/test/box_network_client.py @@ -1,10 +1,12 @@ import pytest import json -from collections import OrderedDict -from io import BytesIO, RawIOBase, UnsupportedOperation, SEEK_SET +import threading +from http.server import BaseHTTPRequestHandler, HTTPServer +from io import BytesIO, RawIOBase, UnsupportedOperation, SEEK_END, SEEK_SET from unittest import mock from unittest.mock import Mock, patch from requests import Session, Response, RequestException +from urllib3.filepost import encode_multipart_formdata from box_sdk_gen import ( NetworkSession, @@ -29,6 +31,7 @@ BoxRetryStrategy, ) from box_sdk_gen.networking.proxy_config import ProxyConfig +from box_sdk_gen.networking.multipart_stream import MultipartStream RETRY_AFTER_HEADER_CASES = [ "retry-after", @@ -370,15 +373,25 @@ def test_prepare_multipart_request(network_client, mock_byte_stream): assert api_request.url == "https://example.com" assert api_request.headers["User-Agent"] == USER_AGENT_HEADER assert api_request.headers["X-Box-UA"] == X_BOX_UA_HEADER - assert api_request.headers["Content-Type"].startswith( - "multipart/form-data; boundary=" + assert api_request.headers["Content-Type"] == ( + f"multipart/form-data; boundary={api_request.data.boundary}" ) assert api_request.params == {} - assert api_request.data.fields == OrderedDict( - [ - ("attributes", '{"name": "file.pdf"}'), - ("file", ("file.pdf", mock_byte_stream, None)), - ] + assert isinstance(api_request.data, MultipartStream) + assert mock_byte_stream.tell() == 0 + + boundary = api_request.data.boundary + assert ( + api_request.data.read() + == ( + f"--{boundary}\r\n" + 'Content-Disposition: form-data; name="attributes"\r\n\r\n' + '{"name": "file.pdf"}\r\n' + f"--{boundary}\r\n" + 'Content-Disposition: form-data; name="file"; filename="file.pdf"\r\n\r\n' + "123\r\n" + f"--{boundary}--\r\n" + ).encode() ) @@ -1258,3 +1271,245 @@ def test_disable_follow_redirects( allow_redirects=False, timeout=(10, 60), ) + + +def test_prepare_multipart_request_drops_custom_content_type_header( + network_client, mock_byte_stream +): + options = FetchOptions( + url="https://example.com", + method="POST", + headers={"content-type": "multipart/form-data"}, + content_type="multipart/form-data", + multipart_data=[ + MultipartItem( + part_name="file", file_stream=mock_byte_stream, file_name="file.pdf" + ), + ], + ) + + api_request = network_client._prepare_request(options=options) + + content_type_headers = [ + value + for name, value in api_request.headers.items() + if name.lower() == "content-type" + ] + assert content_type_headers == [ + f"multipart/form-data; boundary={api_request.data.boundary}" + ] + + +def test_prepare_multipart_request_raises_on_part_without_content(network_client): + options = FetchOptions( + url="https://example.com", + method="POST", + content_type="multipart/form-data", + multipart_data=[MultipartItem(part_name="file", file_name="file.pdf")], + ) + + with pytest.raises(BoxSDKError, match='"file" has neither data nor file_stream'): + network_client._prepare_request(options=options) + + +def _read_request_body(handler): + if handler.headers.get("Transfer-Encoding") == "chunked": + body = b"" + while True: + size = int(handler.rfile.readline().strip(), 16) + chunk = handler.rfile.read(size + 2)[:-2] + if size == 0: + return body + body += chunk + return handler.rfile.read(int(handler.headers["Content-Length"])) + + +@pytest.fixture +def multipart_server(): + received = [] + statuses = [] + + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def do_POST(self): + received.append((dict(self.headers), _read_request_body(self))) + self.send_response(statuses.pop(0) if statuses else 200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", "2") + self.end_headers() + self.wfile.write(b"{}") + + def log_message(self, *args): + pass + + server = HTTPServer(("127.0.0.1", 0), Handler) + threading.Thread(target=server.serve_forever, daemon=True).start() + yield f"http://127.0.0.1:{server.server_port}", received, statuses + server.shutdown() + server.server_close() + + +def _upload_options(url, file_stream, network_session=None): + return FetchOptions( + url=url, + method="POST", + content_type="multipart/form-data", + multipart_data=[ + MultipartItem(part_name="attributes", data={"name": 'fi"le.txt'}), + MultipartItem( + part_name="file", + file_stream=file_stream, + file_name='fi"le.txt', + content_type="text/plain", + ), + ], + network_session=network_session, + ) + + +def _assert_multipart_body(headers, body): + content_type = headers["Content-Type"] + assert content_type.startswith("multipart/form-data; boundary=") + boundary = content_type.split("boundary=")[1] + assert body.endswith(f"--{boundary}--\r\n".encode()) + assert body.index(b'name="attributes"') < body.index(b'name="file"') + assert b'filename="fi%22le.txt"' in body + assert b"Content-Type: text/plain\r\n\r\nfile content\r\n" in body + + +def test_multipart_upload_round_trip_with_retry(multipart_server, network_session_mock): + url, received, statuses = multipart_server + statuses.append(500) + + with patch("time.sleep"): + BoxNetworkClient().fetch( + _upload_options(url, BytesIO(b"file content"), network_session_mock) + ) + + assert len(received) == 2 + for headers, body in received: + assert int(headers["Content-Length"]) == len(body) + assert "Transfer-Encoding" not in headers + _assert_multipart_body(headers, body) + + +def test_multipart_upload_non_seekable_stream_uses_chunked_encoding( + multipart_server, +): + url, received, _ = multipart_server + + BoxNetworkClient().fetch(_upload_options(url, NonSeekableStream(b"file content"))) + + assert len(received) == 1 + headers, body = received[0] + assert headers["Transfer-Encoding"] == "chunked" + assert "Content-Length" not in headers + _assert_multipart_body(headers, body) + + +def test_multipart_stream_matches_urllib3_encoding(): + content = bytes(range(256)) * 1000 + stream = BytesIO(b"skipped" + content) + stream.seek(len(b"skipped")) + multipart_stream = MultipartStream( + [ + ("attributes", None, '{"name": "f\u00e9.bin"}', None), + ("file", "f\u00e9.bin", stream, "application/octet-stream"), + ] + ) + + expected, _ = encode_multipart_formdata( + [ + ("attributes", '{"name": "f\u00e9.bin"}'), + ("file", ("f\u00e9.bin", content, "application/octet-stream")), + ], + boundary=multipart_stream.boundary, + ) + assert multipart_stream.len == len(expected) + assert b"".join(iter(lambda: multipart_stream.read(7), b"")) == expected + + +def test_multipart_stream_reads_file_lazily(): + stream = BytesIO(b"x" * (10 * 1024 * 1024)) + multipart_stream = MultipartStream([("file", "big.bin", stream, None)]) + + assert stream.tell() == 0 + first_chunk = next(iter(multipart_stream)) + assert len(first_chunk) <= 64 * 1024 + assert stream.tell() < 64 * 1024 + + +def test_multipart_stream_length_unknown_for_non_seekable_stream(): + multipart_stream = MultipartStream( + [("file", "file.bin", NonSeekableStream(b"123"), None)] + ) + + assert multipart_stream.len is None + assert multipart_stream.read().endswith( + f"123\r\n--{multipart_stream.boundary}--\r\n".encode() + ) + + +def test_multipart_stream_fails_fast_when_stream_shrinks(): + class Truncated(BytesIO): + def read(self, size=-1): + return super().read(max(0, min(size, 5 - self.tell()))) + + multipart_stream = MultipartStream([("file", "f", Truncated(b"0123456789"), None)]) + + with pytest.raises(IOError, match="ended 5 bytes before its declared size"): + multipart_stream.read() + + +def test_multipart_stream_sends_declared_size_when_stream_grows(): + stream = BytesIO(b"0123456789") + multipart_stream = MultipartStream([("file", "f", stream, None)]) + stream.seek(0, SEEK_END) + stream.write(b"EXTRA") + stream.seek(0) + + body = multipart_stream.read() + + assert len(body) == multipart_stream.len + assert b"EXTRA" not in body + + +def test_multipart_upload_short_stream_fails_without_retry(multipart_server): + url, _, _ = multipart_server + + class Truncated(BytesIO): + def read(self, size=-1): + return super().read(max(0, min(size, 5 - self.tell()))) + + with patch("time.sleep") as sleep: + with pytest.raises(BoxSDKError, match="ended 5 bytes before its declared size"): + BoxNetworkClient().fetch(_upload_options(url, Truncated(b"0123456789"))) + + sleep.assert_not_called() + + +def test_multipart_stream_positioned_past_end_sends_empty_part(): + stream = BytesIO(b"0123456789") + stream.seek(20) + multipart_stream = MultipartStream([("file", "f", stream, None)]) + + body = multipart_stream.read() + + assert len(body) == multipart_stream.len + assert body.count(b"\r\n\r\n\r\n") == 1 + + +def test_multipart_stream_handles_seek_returning_none(): + class LegacySeek(BytesIO): + def seek(self, *args): + super().seek(*args) + + stream = LegacySeek(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 diff --git a/test/box_sdk_gen/test/multipart_uploads.py b/test/box_sdk_gen/test/multipart_uploads.py new file mode 100644 index 00000000..28ff944c --- /dev/null +++ b/test/box_sdk_gen/test/multipart_uploads.py @@ -0,0 +1,64 @@ +from io import BytesIO, RawIOBase + +from box_sdk_gen.client import BoxClient +from box_sdk_gen.internal.utils import ( + generate_byte_buffer, + get_uuid, + read_byte_stream, +) +from box_sdk_gen.managers.uploads import ( + UploadFileAttributes, + UploadFileAttributesParentField, + UploadFileVersionAttributes, +) +from box_sdk_gen.schemas.file_full import FileFull + +from .commons import get_default_client + +client: BoxClient = get_default_client() + + +class NonSeekableStream(RawIOBase): + def __init__(self, content: bytes): + self._content = BytesIO(content) + + def readable(self) -> bool: + return True + + def seekable(self) -> bool: + return False + + def readinto(self, buffer) -> int: + chunk = self._content.read(len(buffer)) + buffer[: len(chunk)] = chunk + return len(chunk) + + +def testUploadFileAndFileVersionFromNonSeekableStream(): + content: bytes = generate_byte_buffer(5 * 1024 * 1024) + new_file_name: str = get_uuid() + uploaded_file: FileFull = client.uploads.upload_file( + UploadFileAttributes( + name=new_file_name, parent=UploadFileAttributesParentField(id='0') + ), + NonSeekableStream(content), + ).entries[0] + try: + assert uploaded_file.name == new_file_name + assert uploaded_file.size == len(content) + assert read_byte_stream(client.downloads.download_file(uploaded_file.id)) == ( + content + ) + + new_content: bytes = generate_byte_buffer(1024 * 1024) + new_file_version: FileFull = client.uploads.upload_file_version( + uploaded_file.id, + UploadFileVersionAttributes(name=get_uuid()), + NonSeekableStream(new_content), + ).entries[0] + assert new_file_version.size == len(new_content) + assert read_byte_stream(client.downloads.download_file(uploaded_file.id)) == ( + new_content + ) + finally: + client.files.delete_file_by_id(uploaded_file.id)