diff --git a/CHANGELOG.md b/CHANGELOG.md index fef0c21..de5c8ca 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [0.7.1] - 2026-07-15 + +### Changed + +- `upload_parquet()` now uses the presigned upload session API (`POST /v1/uploads`) instead of reading the entire file into memory before uploading. For multipart mode the file is streamed one `part_size` chunk at a time, eliminating the memory spike that caused OOM on large Parquet files. Falls back to `POST /v1/files` when the server returns 501. + ## [0.7.0] - 2026-07-14 ### Added diff --git a/hotdata_framework/client.py b/hotdata_framework/client.py index bbb2822..9a3ef82 100644 --- a/hotdata_framework/client.py +++ b/hotdata_framework/client.py @@ -1,7 +1,9 @@ from __future__ import annotations import functools +import os import time +import urllib3 from collections.abc import Iterator from dataclasses import asdict, dataclass from typing import Any, Literal @@ -20,6 +22,9 @@ from hotdata.models.create_database_request import CreateDatabaseRequest from hotdata.models.database_default_schema_decl import DatabaseDefaultSchemaDecl from hotdata.models.database_default_table_decl import DatabaseDefaultTableDecl +from hotdata.models.create_upload_request import CreateUploadRequest +from hotdata.models.finalize_upload_part import FinalizeUploadPart +from hotdata.models.finalize_upload_request import FinalizeUploadRequest from hotdata.models.load_managed_table_request import LoadManagedTableRequest from hotdata.models.query_request import QueryRequest from hotdata.models.query_response import QueryResponse @@ -306,6 +311,71 @@ def list_managed_tables( def upload_parquet(self, path: str) -> str: if not is_parquet_path(path): raise ValueError(f"Managed table loads require a parquet file (got {path!r})") + file_size = os.path.getsize(path) + try: + session = self.uploads().create_upload_session_handler( + CreateUploadRequest( + declared_size_bytes=file_size, + content_type="application/octet-stream", + ) + ) + except ApiException as e: + if e.status == 501: + return self._upload_parquet_post(path) + raise RuntimeError(api_error_message(e)) from e + http = urllib3.PoolManager() + parts: list[FinalizeUploadPart] | None = None + try: + if session.mode == "single": + with open(path, "rb") as f: + data = f.read() + resp = http.request( + "PUT", + session.url, + body=data, + headers={"Content-Length": str(file_size), **session.headers}, + ) + if resp.status not in (200, 201, 204): + raise RuntimeError(f"Storage PUT failed: HTTP {resp.status}") + else: + collected: list[FinalizeUploadPart] = [] + with open(path, "rb") as f: + for i, part_url in enumerate(session.part_urls): + chunk = f.read(session.part_size) + resp = http.request( + "PUT", + part_url, + body=chunk, + headers={ + "Content-Length": str(len(chunk)), + **session.headers, + }, + ) + if resp.status not in (200, 201, 204): + raise RuntimeError( + f"Part {i + 1} PUT failed: HTTP {resp.status}" + ) + collected.append( + FinalizeUploadPart( + part_number=i + 1, + e_tag=resp.headers["ETag"], + ) + ) + parts = collected + finally: + http.clear() + try: + finalized = self.uploads().finalize_upload_handler( + upload_id=session.upload_id, + x_upload_finalize_token=session.finalize_token, + finalize_upload_request=FinalizeUploadRequest(parts=parts), + ) + except ApiException as e: + raise RuntimeError(api_error_message(e)) from e + return finalized.upload_id + + def _upload_parquet_post(self, path: str) -> str: + """Fallback for storage backends that do not support presigned URLs (501).""" with open(path, "rb") as f: data = f.read() try: diff --git a/pyproject.toml b/pyproject.toml index 5987951..7741379 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "hatchling.build" [project] name = "hotdata-framework" -version = "0.7.0" +version = "0.7.1" description = "Python framework for building Hotdata integrations: workspace/session runtime, query execution, and managed databases" readme = "README.md" requires-python = ">=3.10" diff --git a/tests/test_databases.py b/tests/test_databases.py index accfcd0..2a5cc33 100644 --- a/tests/test_databases.py +++ b/tests/test_databases.py @@ -1,7 +1,8 @@ from __future__ import annotations +import io from types import SimpleNamespace -from unittest.mock import mock_open, patch +from unittest.mock import MagicMock, mock_open, patch import pytest from hotdata.exceptions import ApiException @@ -194,16 +195,105 @@ def test_upload_parquet_rejects_non_parquet(): client.upload_parquet("/tmp/data.csv") -def test_upload_parquet_returns_upload_id(): +def _mock_open_bytes(data: bytes) -> MagicMock: + """Return an open() mock whose file handle supports read(n) via BytesIO.""" + bio = io.BytesIO(data) + m = MagicMock() + m.__enter__ = lambda s: bio + m.__exit__ = MagicMock(return_value=False) + return MagicMock(return_value=m) + + +def _session(mode: str, **kw) -> SimpleNamespace: + defaults = dict( + upload_id="upl_sess", + finalize_token="tok", + headers={}, + part_size=None, + part_urls=None, + url=None, + ) + return SimpleNamespace(mode=mode, **{**defaults, **kw}) + + +def _http_resp(status: int = 200, etag: str = '"abc"') -> SimpleNamespace: + return SimpleNamespace(status=status, headers={"ETag": etag}) + + +def test_upload_parquet_multipart(): + client = _client() + data = b"PAR1" + b"\x00" * 6 # 10 bytes -> 2 parts of 5 + session = _session("multipart", part_size=5, part_urls=["https://s/1", "https://s/2"]) + finalized = SimpleNamespace(upload_id="upl_final") + + with ( + patch("builtins.open", _mock_open_bytes(data)), + patch("os.path.getsize", return_value=len(data)), + patch.object(client, "uploads") as uploads, + patch("hotdata_framework.client.urllib3.PoolManager") as MockPool, + ): + pool = MockPool.return_value + pool.request.return_value = _http_resp() + pool.clear.return_value = None + uploads.return_value.create_upload_session_handler.return_value = session + uploads.return_value.finalize_upload_handler.return_value = finalized + + upload_id = client.upload_parquet("/tmp/data.parquet") + + assert upload_id == "upl_final" + assert pool.request.call_count == 2 + finalize_call = uploads.return_value.finalize_upload_handler.call_args + assert finalize_call.kwargs["upload_id"] == "upl_sess" + assert finalize_call.kwargs["x_upload_finalize_token"] == "tok" + parts = finalize_call.kwargs["finalize_upload_request"].parts + assert len(parts) == 2 + assert parts[0].part_number == 1 + assert parts[1].part_number == 2 + + +def test_upload_parquet_single_put(): client = _client() - uploaded = SimpleNamespace(id="upl_123") + data = b"PAR1tiny" + session = _session("single", url="https://s/put") + finalized = SimpleNamespace(upload_id="upl_single") + + with ( + patch("builtins.open", mock_open(read_data=data)), + patch("os.path.getsize", return_value=len(data)), + patch.object(client, "uploads") as uploads, + patch("hotdata_framework.client.urllib3.PoolManager") as MockPool, + ): + pool = MockPool.return_value + pool.request.return_value = _http_resp() + pool.clear.return_value = None + uploads.return_value.create_upload_session_handler.return_value = session + uploads.return_value.finalize_upload_handler.return_value = finalized + + upload_id = client.upload_parquet("/tmp/data.parquet") + + assert upload_id == "upl_single" + pool.request.assert_called_once() + call_args = pool.request.call_args + assert call_args.args[0] == "PUT" + assert call_args.args[1] == "https://s/put" + + +def test_upload_parquet_fallback_on_501(): + client = _client() + uploaded = SimpleNamespace(id="upl_post") + with ( patch("builtins.open", mock_open(read_data=b"PAR1")), + patch("os.path.getsize", return_value=4), patch.object(client, "uploads") as uploads, ): + err = ApiException(status=501) + uploads.return_value.create_upload_session_handler.side_effect = err uploads.return_value.upload_file.return_value = uploaded + upload_id = client.upload_parquet("/tmp/data.parquet") - assert upload_id == "upl_123" + + assert upload_id == "upl_post" def test_load_managed_table_with_upload_id(): diff --git a/uv.lock b/uv.lock index eaff653..a965abb 100644 --- a/uv.lock +++ b/uv.lock @@ -101,7 +101,7 @@ wheels = [ [[package]] name = "hotdata-framework" -version = "0.7.0" +version = "0.7.1" source = { editable = "." } dependencies = [ { name = "hotdata" },