Skip to content

Commit faebb64

Browse files
committed
feat: stream large Parquet uploads via presigned session API
Replaces the full-file f.read() in upload_parquet() with the presigned upload session API, so only one part_size chunk (not the entire file) is in memory at a time. Falls back to POST /v1/files when the server returns 501 (backends without presigned URL support). Bumps to 0.7.1.
1 parent a075c7c commit faebb64

4 files changed

Lines changed: 166 additions & 6 deletions

File tree

hotdata_framework/client.py

Lines changed: 70 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
from __future__ import annotations
22

33
import functools
4+
import os
45
import time
6+
import urllib3
57
from collections.abc import Iterator
68
from dataclasses import asdict, dataclass
79
from typing import Any, Literal
@@ -20,6 +22,9 @@
2022
from hotdata.models.create_database_request import CreateDatabaseRequest
2123
from hotdata.models.database_default_schema_decl import DatabaseDefaultSchemaDecl
2224
from hotdata.models.database_default_table_decl import DatabaseDefaultTableDecl
25+
from hotdata.models.create_upload_request import CreateUploadRequest
26+
from hotdata.models.finalize_upload_part import FinalizeUploadPart
27+
from hotdata.models.finalize_upload_request import FinalizeUploadRequest
2328
from hotdata.models.load_managed_table_request import LoadManagedTableRequest
2429
from hotdata.models.query_request import QueryRequest
2530
from hotdata.models.query_response import QueryResponse
@@ -306,6 +311,71 @@ def list_managed_tables(
306311
def upload_parquet(self, path: str) -> str:
307312
if not is_parquet_path(path):
308313
raise ValueError(f"Managed table loads require a parquet file (got {path!r})")
314+
file_size = os.path.getsize(path)
315+
try:
316+
session = self.uploads().create_upload_session_handler(
317+
CreateUploadRequest(
318+
declared_size_bytes=file_size,
319+
content_type="application/octet-stream",
320+
)
321+
)
322+
except ApiException as e:
323+
if e.status == 501:
324+
return self._upload_parquet_post(path)
325+
raise RuntimeError(api_error_message(e)) from e
326+
http = urllib3.PoolManager()
327+
parts: list[FinalizeUploadPart] | None = None
328+
try:
329+
if session.mode == "single":
330+
with open(path, "rb") as f:
331+
data = f.read()
332+
resp = http.request(
333+
"PUT",
334+
session.url,
335+
body=data,
336+
headers={"Content-Length": str(file_size), **session.headers},
337+
)
338+
if resp.status not in (200, 201, 204):
339+
raise RuntimeError(f"Storage PUT failed: HTTP {resp.status}")
340+
else:
341+
collected: list[FinalizeUploadPart] = []
342+
with open(path, "rb") as f:
343+
for i, part_url in enumerate(session.part_urls):
344+
chunk = f.read(session.part_size)
345+
resp = http.request(
346+
"PUT",
347+
part_url,
348+
body=chunk,
349+
headers={
350+
"Content-Length": str(len(chunk)),
351+
**session.headers,
352+
},
353+
)
354+
if resp.status not in (200, 201, 204):
355+
raise RuntimeError(
356+
f"Part {i + 1} PUT failed: HTTP {resp.status}"
357+
)
358+
collected.append(
359+
FinalizeUploadPart(
360+
part_number=i + 1,
361+
e_tag=resp.headers["ETag"],
362+
)
363+
)
364+
parts = collected
365+
finally:
366+
http.clear()
367+
try:
368+
finalized = self.uploads().finalize_upload_handler(
369+
upload_id=session.upload_id,
370+
x_upload_finalize_token=session.finalize_token,
371+
finalize_upload_request=FinalizeUploadRequest(parts=parts),
372+
)
373+
except ApiException as e:
374+
raise RuntimeError(api_error_message(e)) from e
375+
return finalized.upload_id
376+
377+
def _upload_parquet_post(self, path: str) -> str:
378+
"""Fallback for storage backends that do not support presigned URLs (501)."""
309379
with open(path, "rb") as f:
310380
data = f.read()
311381
try:

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
44

55
[project]
66
name = "hotdata-framework"
7-
version = "0.7.0"
7+
version = "0.7.1"
88
description = "Python framework for building Hotdata integrations: workspace/session runtime, query execution, and managed databases"
99
readme = "README.md"
1010
requires-python = ">=3.10"

tests/test_databases.py

Lines changed: 94 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,8 @@
11
from __future__ import annotations
22

3+
import io
34
from types import SimpleNamespace
4-
from unittest.mock import mock_open, patch
5+
from unittest.mock import MagicMock, mock_open, patch
56

67
import pytest
78
from hotdata.exceptions import ApiException
@@ -194,16 +195,105 @@ def test_upload_parquet_rejects_non_parquet():
194195
client.upload_parquet("/tmp/data.csv")
195196

196197

197-
def test_upload_parquet_returns_upload_id():
198+
def _mock_open_bytes(data: bytes) -> MagicMock:
199+
"""Return an open() mock whose file handle supports read(n) via BytesIO."""
200+
bio = io.BytesIO(data)
201+
m = MagicMock()
202+
m.__enter__ = lambda s: bio
203+
m.__exit__ = MagicMock(return_value=False)
204+
return MagicMock(return_value=m)
205+
206+
207+
def _session(mode: str, **kw) -> SimpleNamespace:
208+
defaults = dict(
209+
upload_id="upl_sess",
210+
finalize_token="tok",
211+
headers={},
212+
part_size=None,
213+
part_urls=None,
214+
url=None,
215+
)
216+
return SimpleNamespace(mode=mode, **{**defaults, **kw})
217+
218+
219+
def _http_resp(status: int = 200, etag: str = '"abc"') -> SimpleNamespace:
220+
return SimpleNamespace(status=status, headers={"ETag": etag})
221+
222+
223+
def test_upload_parquet_multipart():
224+
client = _client()
225+
data = b"PAR1" + b"\x00" * 6 # 10 bytes -> 2 parts of 5
226+
session = _session("multipart", part_size=5, part_urls=["https://s/1", "https://s/2"])
227+
finalized = SimpleNamespace(upload_id="upl_final")
228+
229+
with (
230+
patch("builtins.open", _mock_open_bytes(data)),
231+
patch("os.path.getsize", return_value=len(data)),
232+
patch.object(client, "uploads") as uploads,
233+
patch("hotdata_framework.client.urllib3.PoolManager") as MockPool,
234+
):
235+
pool = MockPool.return_value
236+
pool.request.return_value = _http_resp()
237+
pool.clear.return_value = None
238+
uploads.return_value.create_upload_session_handler.return_value = session
239+
uploads.return_value.finalize_upload_handler.return_value = finalized
240+
241+
upload_id = client.upload_parquet("/tmp/data.parquet")
242+
243+
assert upload_id == "upl_final"
244+
assert pool.request.call_count == 2
245+
finalize_call = uploads.return_value.finalize_upload_handler.call_args
246+
assert finalize_call.kwargs["upload_id"] == "upl_sess"
247+
assert finalize_call.kwargs["x_upload_finalize_token"] == "tok"
248+
parts = finalize_call.kwargs["finalize_upload_request"].parts
249+
assert len(parts) == 2
250+
assert parts[0].part_number == 1
251+
assert parts[1].part_number == 2
252+
253+
254+
def test_upload_parquet_single_put():
198255
client = _client()
199-
uploaded = SimpleNamespace(id="upl_123")
256+
data = b"PAR1tiny"
257+
session = _session("single", url="https://s/put")
258+
finalized = SimpleNamespace(upload_id="upl_single")
259+
260+
with (
261+
patch("builtins.open", mock_open(read_data=data)),
262+
patch("os.path.getsize", return_value=len(data)),
263+
patch.object(client, "uploads") as uploads,
264+
patch("hotdata_framework.client.urllib3.PoolManager") as MockPool,
265+
):
266+
pool = MockPool.return_value
267+
pool.request.return_value = _http_resp()
268+
pool.clear.return_value = None
269+
uploads.return_value.create_upload_session_handler.return_value = session
270+
uploads.return_value.finalize_upload_handler.return_value = finalized
271+
272+
upload_id = client.upload_parquet("/tmp/data.parquet")
273+
274+
assert upload_id == "upl_single"
275+
pool.request.assert_called_once()
276+
call_args = pool.request.call_args
277+
assert call_args.args[0] == "PUT"
278+
assert call_args.args[1] == "https://s/put"
279+
280+
281+
def test_upload_parquet_fallback_on_501():
282+
client = _client()
283+
uploaded = SimpleNamespace(id="upl_post")
284+
200285
with (
201286
patch("builtins.open", mock_open(read_data=b"PAR1")),
287+
patch("os.path.getsize", return_value=4),
202288
patch.object(client, "uploads") as uploads,
203289
):
290+
err = ApiException(status=501)
291+
uploads.return_value.create_upload_session_handler.side_effect = err
204292
uploads.return_value.upload_file.return_value = uploaded
293+
205294
upload_id = client.upload_parquet("/tmp/data.parquet")
206-
assert upload_id == "upl_123"
295+
296+
assert upload_id == "upl_post"
207297

208298

209299
def test_load_managed_table_with_upload_id():

uv.lock

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)