Skip to content
Draft
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
6 changes: 4 additions & 2 deletions docs/changelog.rst
Original file line number Diff line number Diff line change
Expand Up @@ -329,8 +329,10 @@ v0.54.0 - SQL processing correctness and cleanup
* Spanner adapter modules no longer expose module-level proxy lookup hooks.
* Async migration squash now builds its internal migration runner with a real
migration context, matching the synchronous command path.
* ObStore Arrow streaming no longer resolves cloud ``base_path`` twice for
async streams.
* Arrow batch streaming is now explicitly Parquet-only and reads one row group
at a time across local, fsspec, and obstore backends. Obstore streams through
its seekable reader without buffering the full object, resolves cloud
``base_path`` only once, and closes readers deterministically.
* ``sql.decode()`` now renders a trailing default argument as the ``ELSE``
clause documented for DECODE-style expressions.
* Async drivers can use the statement-cache direct execution path when the
Expand Down
15 changes: 15 additions & 0 deletions docs/reference/storage.rst
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,21 @@ Storage abstraction layer with multiple backend support (local filesystem,
fsspec, obstore), configuration-based registration, and Arrow table
import/export with CSV format support.

Parquet Batch Streaming
=======================

The ``stream_arrow_sync()`` and ``stream_arrow_async()`` backend methods stream
Parquet files in file and row-group order. They accept a keyword-only
``batch_size`` (default ``65_536``) which controls the maximum rows in each
record batch. Each read is restricted to one Parquet row group, so the I/O bound
is one row group rather than one record batch. Choose the Parquet row-group size
when writing files according to the memory bound required while reading them.

These methods intentionally support only ``file_format="parquet"``. Use the
regular Arrow read APIs for CSV, Arrow IPC, JSON, and JSONL payloads. Closing a
sync generator or calling ``aclose()`` on its async iterator closes the active
storage reader.

Pipelines
=========

Expand Down
10 changes: 7 additions & 3 deletions sqlspec/protocols.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
and runtime isinstance() checks.
"""

from typing import TYPE_CHECKING, Any, ClassVar, Protocol, overload, runtime_checkable
from typing import TYPE_CHECKING, Any, ClassVar, Literal, Protocol, overload, runtime_checkable

from typing_extensions import Self

Expand Down Expand Up @@ -543,7 +543,9 @@ def write_arrow_sync(self, path: "str | Path", table: "ArrowTable", **kwargs: An
msg = "Arrow writing not implemented"
raise NotImplementedError(msg)

def stream_arrow_sync(self, pattern: str, **kwargs: Any) -> "Iterator[ArrowRecordBatch]":
def stream_arrow_sync(
self, pattern: str, *, file_format: Literal["parquet"] = "parquet", batch_size: int = 65_536, **kwargs: Any
) -> "Iterator[ArrowRecordBatch]":
"""Stream Arrow record batches from matching objects synchronously."""
msg = "Arrow streaming not implemented"
raise NotImplementedError(msg)
Expand Down Expand Up @@ -621,7 +623,9 @@ async def write_arrow_async(self, path: "str | Path", table: "ArrowTable", **kwa
raise NotImplementedError(msg)

# NOTE: Returns AsyncIterator directly; this is intentionally not async def.
def stream_arrow_async(self, pattern: str, **kwargs: Any) -> "AsyncIterator[ArrowRecordBatch]":
def stream_arrow_async(
self, pattern: str, *, file_format: Literal["parquet"] = "parquet", batch_size: int = 65_536, **kwargs: Any
) -> "AsyncIterator[ArrowRecordBatch]":
"""Stream Arrow record batches from matching objects."""
msg = "Async arrow streaming not implemented"
raise NotImplementedError(msg)
Expand Down
34 changes: 34 additions & 0 deletions sqlspec/storage/_arrow_stream.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,34 @@
"""Shared helpers for bounded Parquet batch streaming."""

from pathlib import PurePath
from typing import TYPE_CHECKING, Any

if TYPE_CHECKING:
from collections.abc import Iterator

from sqlspec.typing import ArrowRecordBatch

__all__ = ("iter_parquet_row_groups", "validate_parquet_stream_options")

_NON_PARQUET_SUFFIXES = frozenset({".arrow", ".csv", ".feather", ".ipc", ".json", ".jsonl", ".ndjson"})


def validate_parquet_stream_options(pattern: str, file_format: str, batch_size: int) -> None:
"""Validate a Parquet streaming request before storage is accessed."""
if file_format != "parquet":
msg = f"Arrow batch streaming supports only Parquet files; received file_format={file_format!r}"
raise ValueError(msg)
if batch_size <= 0:
msg = f"batch_size must be greater than zero; received {batch_size}"
raise ValueError(msg)

suffix = PurePath(pattern).suffix.lower()
if suffix in _NON_PARQUET_SUFFIXES:
msg = f"Arrow batch streaming supports only Parquet files; pattern {pattern!r} selects {suffix} files"
raise ValueError(msg)


def iter_parquet_row_groups(parquet_file: Any, *, batch_size: int, **kwargs: Any) -> "Iterator[ArrowRecordBatch]":
"""Yield batches while limiting each PyArrow read to one row group."""
for row_group in range(parquet_file.num_row_groups):
yield from parquet_file.iter_batches(batch_size=batch_size, row_groups=[row_group], **kwargs)
33 changes: 29 additions & 4 deletions sqlspec/storage/backends/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import contextlib
from abc import abstractmethod
from collections.abc import AsyncIterator, Iterator
from typing import TYPE_CHECKING, Any, cast
from typing import TYPE_CHECKING, Any, Literal, cast

from mypy_extensions import mypyc_attr
from typing_extensions import Self
Expand Down Expand Up @@ -58,17 +58,38 @@ def _read_chunk_or_sentinel(file_obj: Any, chunk_size: int) -> Any:
class AsyncArrowBatchIterator:
"""Async iterator wrapper for sync Arrow batch iterators."""

__slots__ = ("_sync_iter",)
__slots__ = ("_closed", "_sync_iter")

def __init__(self, sync_iterator: "Iterator[ArrowRecordBatch]") -> None:
self._sync_iter = sync_iterator
self._closed = False

def __aiter__(self) -> "AsyncArrowBatchIterator":
return self

async def __aenter__(self) -> Self:
return self

async def __aexit__(
self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None"
) -> None:
await self.aclose()

async def aclose(self) -> None:
"""Close the underlying generator and its active storage reader."""
if self._closed:
return
self._closed = True
close = getattr(self._sync_iter, "close", None)
if close is not None:
await asyncio.get_running_loop().run_in_executor(None, close)

def _sync_next(self) -> "ArrowRecordBatch":
if self._closed:
raise _StopAsync()
result = _next_or_sentinel(self._sync_iter)
if result is _EXHAUSTED:
self._closed = True
raise _StopAsync()
return cast("ArrowRecordBatch", result)

Expand Down Expand Up @@ -227,7 +248,9 @@ def write_arrow_sync(self, path: str, table: "ArrowTable", **kwargs: Any) -> Non
raise NotImplementedError

@abstractmethod
def stream_arrow_sync(self, pattern: str, **kwargs: Any) -> "Iterator[ArrowRecordBatch]":
def stream_arrow_sync(
self, pattern: str, *, file_format: Literal["parquet"] = "parquet", batch_size: int = 65_536, **kwargs: Any
) -> "Iterator[ArrowRecordBatch]":
"""Stream Arrow record batches from storage synchronously."""
raise NotImplementedError

Expand Down Expand Up @@ -300,6 +323,8 @@ async def write_arrow_async(self, path: str, table: "ArrowTable", **kwargs: Any)

# NOTE: Returns AsyncIterator directly; keep in sync with ObjectStoreProtocol.
@abstractmethod
def stream_arrow_async(self, pattern: str, **kwargs: Any) -> "AsyncIterator[ArrowRecordBatch]":
def stream_arrow_async(
self, pattern: str, *, file_format: Literal["parquet"] = "parquet", batch_size: int = 65_536, **kwargs: Any
) -> "AsyncIterator[ArrowRecordBatch]":
"""Stream Arrow record batches from storage asynchronously."""
raise NotImplementedError
28 changes: 20 additions & 8 deletions sqlspec/storage/backends/fsspec.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,12 @@
from collections.abc import AsyncIterator, Iterator
from functools import partial
from pathlib import Path
from typing import TYPE_CHECKING, Any, ClassVar, cast, overload
from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast, overload
from urllib.parse import urlparse

from mypy_extensions import mypyc_attr

from sqlspec.storage._arrow_stream import iter_parquet_row_groups, validate_parquet_stream_options
from sqlspec.storage._paths import resolve_storage_path
from sqlspec.storage._utils import _log_storage_event, import_pyarrow_parquet
from sqlspec.storage.backends.base import AsyncArrowBatchIterator, AsyncThreadedBytesIterator
Expand Down Expand Up @@ -396,18 +397,23 @@ def stream_read_sync(self, path: "str | Path", chunk_size: "int | None" = None,
break
yield cast("bytes", chunk)

def stream_arrow_sync(self, pattern: str, **kwargs: Any) -> Iterator["ArrowRecordBatch"]:
def stream_arrow_sync(
self, pattern: str, *, file_format: Literal["parquet"] = "parquet", batch_size: int = 65_536, **kwargs: Any
) -> Iterator["ArrowRecordBatch"]:
"""Stream Arrow record batches from storage synchronously.

Args:
pattern: The glob pattern to match.
**kwargs: Additional arguments to pass to the glob method.
file_format: Storage format. Only Parquet supports bounded batch streaming.
batch_size: Maximum number of rows in each yielded record batch.
**kwargs: Additional arguments passed to PyArrow batch iteration.

Yields:
Arrow record batches from matching files.
"""
validate_parquet_stream_options(pattern, file_format, batch_size)
pq = import_pyarrow_parquet()
for obj_path in self.glob_sync(pattern, **kwargs):
for obj_path in self.glob_sync(pattern):
file_handle = execute_sync_storage_operation(
partial(self.fs.open, obj_path, mode="rb"),
backend=self.backend_type,
Expand All @@ -421,7 +427,7 @@ def stream_arrow_sync(self, pattern: str, **kwargs: Any) -> Iterator["ArrowRecor
operation="stream_arrow",
path=str(obj_path),
)
yield from parquet_file.iter_batches() # pyright: ignore[reportUnknownMemberType]
yield from iter_parquet_row_groups(parquet_file, batch_size=batch_size, **kwargs)

async def read_bytes_async(self, path: "str | Path", **kwargs: Any) -> bytes:
"""Read bytes from storage asynchronously."""
Expand Down Expand Up @@ -456,17 +462,23 @@ async def stream_read_async(

return AsyncThreadedBytesIterator(file_obj, chunk_size)

def stream_arrow_async(self, pattern: str, **kwargs: Any) -> AsyncIterator["ArrowRecordBatch"]:
def stream_arrow_async(
self, pattern: str, *, file_format: Literal["parquet"] = "parquet", batch_size: int = 65_536, **kwargs: Any
) -> AsyncIterator["ArrowRecordBatch"]:
"""Stream Arrow record batches from storage asynchronously.

Args:
pattern: The glob pattern to match.
**kwargs: Additional arguments to pass to the glob method.
file_format: Storage format. Only Parquet supports bounded batch streaming.
batch_size: Maximum number of rows in each yielded record batch.
**kwargs: Additional arguments passed to PyArrow batch iteration.

Returns:
AsyncIterator yielding Arrow record batches.
"""
return AsyncArrowBatchIterator(self.stream_arrow_sync(pattern, **kwargs))
return AsyncArrowBatchIterator(
self.stream_arrow_sync(pattern, file_format=file_format, batch_size=batch_size, **kwargs)
)

async def read_text_async(self, path: "str | Path", encoding: str = "utf-8", **kwargs: Any) -> str:
"""Read text from storage asynchronously."""
Expand Down
20 changes: 15 additions & 5 deletions sqlspec/storage/backends/local.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,12 +8,13 @@
from collections.abc import AsyncIterator, Iterator
from functools import partial
from pathlib import Path
from typing import TYPE_CHECKING, Any, ClassVar, cast, overload
from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast, overload
from urllib.parse import unquote, urlparse

from mypy_extensions import mypyc_attr

from sqlspec.exceptions import FileNotFoundInStorageError
from sqlspec.storage._arrow_stream import iter_parquet_row_groups, validate_parquet_stream_options
from sqlspec.storage._paths import strip_windows_drive_prefix
from sqlspec.storage._utils import import_pyarrow_parquet
from sqlspec.storage.backends.base import AsyncArrowBatchIterator, AsyncThreadedBytesIterator
Expand Down Expand Up @@ -284,12 +285,15 @@ def write_arrow_sync(self, path: "str | Path", table: "ArrowTable", **kwargs: An
path=str(resolved),
)

def stream_arrow_sync(self, pattern: str, **kwargs: Any) -> Iterator["ArrowRecordBatch"]:
def stream_arrow_sync(
self, pattern: str, *, file_format: Literal["parquet"] = "parquet", batch_size: int = 65_536, **kwargs: Any
) -> Iterator["ArrowRecordBatch"]:
"""Stream Arrow record batches from files matching pattern synchronously.

Yields:
Arrow record batches from matching files.
"""
validate_parquet_stream_options(pattern, file_format, batch_size)
pq = import_pyarrow_parquet()
files = self.glob_sync(pattern)
for file_path in files:
Expand All @@ -301,7 +305,7 @@ def stream_arrow_sync(self, pattern: str, **kwargs: Any) -> Iterator["ArrowRecor
operation="stream_arrow",
path=resolved_str,
)
yield from parquet_file.iter_batches() # pyright: ignore[reportUnknownMemberType]
yield from iter_parquet_row_groups(parquet_file, batch_size=batch_size, **kwargs)

@property
def supports_signing(self) -> bool:
Expand Down Expand Up @@ -409,17 +413,23 @@ async def write_arrow_async(self, path: "str | Path", table: "ArrowTable", **kwa
"""
await async_(self.write_arrow_sync)(path, table, **kwargs)

def stream_arrow_async(self, pattern: str, **kwargs: Any) -> AsyncIterator["ArrowRecordBatch"]:
def stream_arrow_async(
self, pattern: str, *, file_format: Literal["parquet"] = "parquet", batch_size: int = 65_536, **kwargs: Any
) -> AsyncIterator["ArrowRecordBatch"]:
"""Stream Arrow record batches asynchronously.

Args:
pattern: Glob pattern to match files.
file_format: Storage format. Only Parquet supports bounded batch streaming.
batch_size: Maximum number of rows in each yielded record batch.
**kwargs: Additional arguments passed to stream_arrow_sync().

Returns:
AsyncIterator yielding Arrow record batches.
"""
return AsyncArrowBatchIterator(self.stream_arrow_sync(pattern, **kwargs))
return AsyncArrowBatchIterator(
self.stream_arrow_sync(pattern, file_format=file_format, batch_size=batch_size, **kwargs)
)

@overload
async def sign_async(self, paths: str, expires_in: int = 3600, for_upload: bool = False) -> str: ...
Expand Down
Loading
Loading