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
2 changes: 1 addition & 1 deletion .codegen.json
Original file line number Diff line number Diff line change
@@ -1 +1 @@
{ "engineHash": "12fa0b7", "specHash": "c5a35a6", "version": "4.17.0" }
{ "engineHash": "fd56d95", "specHash": "c5a35a6", "version": "4.17.0" }
47 changes: 36 additions & 11 deletions box_sdk_gen/networking/multipart_stream.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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:
Expand All @@ -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)
Expand All @@ -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
Expand All @@ -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]:
Expand All @@ -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):
Expand All @@ -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)
Expand Down
118 changes: 117 additions & 1 deletion test/box_sdk_gen/test/box_network_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
7 changes: 6 additions & 1 deletion test/boxsdk/unit/auth/test_jwt_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading