Skip to content
43 changes: 33 additions & 10 deletions python/src/iceberg/avro/decoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
import decimal
import struct
from datetime import date, datetime, time
from io import SEEK_CUR

from iceberg.io.base import InputStream
from iceberg.utils.datetime import (
Expand Down Expand Up @@ -56,6 +57,9 @@ def read(self, n: int) -> bytes:
raise ValueError(f"Read {len(read_bytes)} bytes, expected {n} bytes")
return read_bytes

def skip(self, n: int) -> None:
self._input_stream.seek(n, SEEK_CUR)

def read_boolean(self) -> bool:
"""
a boolean is written as a single byte
Expand All @@ -64,11 +68,7 @@ def read_boolean(self) -> bool:
return ord(self.read(1)) == 1

def read_int(self) -> int:
"""int values are written using variable-length, zigzag coding."""
return self.read_long()

def read_long(self) -> int:
"""long values are written using variable-length, zigzag coding."""
"""int/long values are written using variable-length, zigzag coding."""
b = ord(self.read(1))
n = b & 0x7F
shift = 7
Expand Down Expand Up @@ -100,7 +100,7 @@ def read_decimal_from_bytes(self, precision: int, scale: int) -> decimal.Decimal
Decimal bytes are decoded as signed short, int or long depending on the
size of bytes.
"""
size = self.read_long()
size = self.read_int()
return self.read_decimal_from_fixed(precision, scale, size)

def read_decimal_from_fixed(self, _: int, scale: int, size: int) -> decimal.Decimal:
Expand All @@ -116,7 +116,7 @@ def read_bytes(self) -> bytes:
"""
Bytes are encoded as a long followed by that many bytes of data.
"""
return self.read(self.read_long())
return self.read(self.read_int())

def read_utf8(self) -> str:
"""
Expand Down Expand Up @@ -146,14 +146,14 @@ def read_time_micros(self) -> time:
long is decoded as python time object which represents
the number of microseconds after midnight, 00:00:00.000000.
"""
return micros_to_time(self.read_long())
return micros_to_time(self.read_int())

def read_timestamp_micros(self) -> datetime:
"""
long is decoded as python datetime object which represents
the number of microseconds from the unix epoch, 1 January 1970.
"""
return micros_to_timestamp(self.read_long())
return micros_to_timestamp(self.read_int())

def read_timestamptz_micros(self):
"""
Expand All @@ -162,4 +162,27 @@ def read_timestamptz_micros(self):

Adjusted to UTC
"""
return micros_to_timestamptz(self.read_long())
return micros_to_timestamptz(self.read_int())

def skip_null(self) -> None:
pass

def skip_boolean(self) -> None:
self.skip(1)

def skip_int(self) -> None:
b = ord(self.read(1))
while (b & 0x80) != 0:
b = ord(self.read(1))

def skip_float(self) -> None:
self.skip(4)

def skip_double(self) -> None:
self.skip(8)

def skip_bytes(self) -> None:
self.skip(self.read_int())

def skip_utf8(self) -> None:
self.skip_bytes()
15 changes: 11 additions & 4 deletions python/src/iceberg/avro/file.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
from iceberg.avro.codecs import KNOWN_CODECS, Codec
from iceberg.avro.decoder import BinaryDecoder
from iceberg.avro.reader import AvroStruct, ConstructReader, StructReader
from iceberg.avro.resolver import resolve
from iceberg.io.base import InputFile, InputStream
from iceberg.io.memory import MemoryInputStream
from iceberg.schema import Schema, visit
Expand Down Expand Up @@ -107,6 +108,7 @@ def __next__(self) -> AvroStruct:

class AvroFile:
input_file: InputFile
read_schema: Schema | None
input_stream: InputStream
header: AvroFileHeader
schema: Schema
Expand All @@ -116,8 +118,9 @@ class AvroFile:
decoder: BinaryDecoder
block: Block | None = None

def __init__(self, input_file: InputFile) -> None:
def __init__(self, input_file: InputFile, read_schema: Schema | None = None) -> None:
self.input_file = input_file
self.read_schema = read_schema

def __enter__(self):
"""
Expand All @@ -132,7 +135,11 @@ def __enter__(self):
self.header = self._read_header()
self.schema = self.header.get_schema()
self.file_length = len(self.input_file)
self.reader = visit(self.schema, ConstructReader())
if not self.read_schema:
self.reader = visit(self.schema, ConstructReader())
else:
self.reader = resolve(self.schema, self.read_schema)

return self

def __exit__(self, exc_type, exc_val, exc_tb):
Expand All @@ -149,9 +156,9 @@ def _read_block(self) -> int:
raise ValueError(f"Expected sync bytes {self.header.sync!r}, but got {sync_marker!r}")
if self.is_EOF():
raise StopIteration
block_records = self.decoder.read_long()
block_records = self.decoder.read_int()

block_bytes_len = self.decoder.read_long()
block_bytes_len = self.decoder.read_int()
block_bytes = self.decoder.read(block_bytes_len)
if codec := self.header.compression_codec():
block_bytes = codec.decompress(block_bytes)
Expand Down
110 changes: 97 additions & 13 deletions python/src/iceberg/avro/reader.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,66 +76,102 @@ class Reader(Singleton):
def read(self, decoder: BinaryDecoder) -> Any:
...

@abstractmethod
def skip(self, decoder: BinaryDecoder) -> None:
...


class NoneReader(Reader):
def read(self, _: BinaryDecoder) -> None:
return None

def skip(self, decoder: BinaryDecoder) -> None:
return None


class BooleanReader(Reader):
def read(self, decoder: BinaryDecoder) -> bool:
return decoder.read_boolean()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_boolean()


class IntegerReader(Reader):
def read(self, decoder: BinaryDecoder) -> int:
return decoder.read_int()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_int()

class LongReader(Reader):
def read(self, decoder: BinaryDecoder) -> int:
return decoder.read_long()

class LongReader(IntegerReader):
"""Longs and ints are encoded the same way, and there is no long in Python"""
Comment thread
Fokko marked this conversation as resolved.
Outdated


class FloatReader(Reader):
def read(self, decoder: BinaryDecoder) -> float:
return decoder.read_float()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_float()


class DoubleReader(Reader):
def read(self, decoder: BinaryDecoder) -> float:
return decoder.read_double()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_double()


class DateReader(Reader):
def read(self, decoder: BinaryDecoder) -> date:
return decoder.read_date_from_int()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_int()


class TimeReader(Reader):
def read(self, decoder: BinaryDecoder) -> time:
return decoder.read_time_micros()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_int()


class TimestampReader(Reader):
def read(self, decoder: BinaryDecoder) -> datetime:
return decoder.read_timestamp_micros()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_int()


class TimestamptzReader(Reader):
def read(self, decoder: BinaryDecoder) -> datetime:
return decoder.read_timestamptz_micros()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_int()


class StringReader(Reader):
def read(self, decoder: BinaryDecoder) -> str:
return decoder.read_utf8()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_utf8()


class UUIDReader(Reader):
def read(self, decoder: BinaryDecoder) -> UUID:
return UUID(decoder.read_utf8())

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_utf8()


@dataclass(frozen=True)
class FixedReader(Reader):
Expand All @@ -144,11 +180,17 @@ class FixedReader(Reader):
def read(self, decoder: BinaryDecoder) -> bytes:
return decoder.read(self.length)

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip(self.length)


class BinaryReader(Reader):
def read(self, decoder: BinaryDecoder) -> bytes:
return decoder.read_bytes()

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_bytes()


@dataclass(frozen=True)
class DecimalReader(Reader):
Expand All @@ -158,6 +200,9 @@ class DecimalReader(Reader):
def read(self, decoder: BinaryDecoder) -> Decimal:
return decoder.read_decimal_from_bytes(self.precision, self.scale)

def skip(self, decoder: BinaryDecoder) -> None:
decoder.skip_bytes()


@dataclass(frozen=True)
class OptionReader(Reader):
Expand All @@ -177,13 +222,28 @@ def read(self, decoder: BinaryDecoder) -> Any | None:
return self.option.read(decoder)
return None

def skip(self, decoder: BinaryDecoder) -> None:
if decoder.read_int() > 0:
return self.option.skip(decoder)


@dataclass(frozen=True)
class StructReader(Reader):
fields: tuple[Reader, ...] = dataclassfield()
fields: tuple[tuple[int | None, Reader], ...] = dataclassfield()

def read(self, decoder: BinaryDecoder) -> AvroStruct:
return AvroStruct([field.read(decoder) for field in self.fields])
result: list[Any | StructProtocol] = [object] * len(self.fields)
Comment thread
Fokko marked this conversation as resolved.
Outdated
for (pos, field) in self.fields:
if pos is not None:
result[pos] = field.read(decoder)
else:
field.skip(decoder)

return AvroStruct(result)

def skip(self, decoder: BinaryDecoder) -> None:
for _, field in self.fields:
field.skip(decoder)


@dataclass(frozen=True)
Expand All @@ -192,17 +252,28 @@ class ListReader(Reader):

def read(self, decoder: BinaryDecoder) -> list:
read_items = []
block_count = decoder.read_long()
block_count = decoder.read_int()
while block_count != 0:
if block_count < 0:
block_count = -block_count
# We ignore the block size for now
_ = decoder.read_long()
_ = decoder.read_int()
for _ in range(block_count):
read_items.append(self.element.read(decoder))
block_count = decoder.read_long()
block_count = decoder.read_int()
return read_items

def skip(self, decoder: BinaryDecoder) -> None:
block_count = decoder.read_int()
while block_count != 0:
if block_count < 0:
block_count = -block_count
block_size = decoder.read_int()
decoder.skip(block_size)
else:
for _ in range(block_count):
self.element.skip(decoder)
block_count = decoder.read_int()
Comment thread
Fokko marked this conversation as resolved.
Outdated


@dataclass(frozen=True)
class MapReader(Reader):
Expand All @@ -211,26 +282,39 @@ class MapReader(Reader):

def read(self, decoder: BinaryDecoder) -> dict:
read_items = {}
block_count = decoder.read_long()
block_count = decoder.read_int()
while block_count != 0:
if block_count < 0:
block_count = -block_count
# We ignore the block size for now
_ = decoder.read_long()
_ = decoder.read_int()
for _ in range(block_count):
key = self.key.read(decoder)
read_items[key] = self.value.read(decoder)
block_count = decoder.read_long()
block_count = decoder.read_int()

return read_items

def skip(self, decoder: BinaryDecoder) -> None:
block_count = decoder.read_int()
while block_count != 0:
if block_count < 0:
block_count = -block_count
block_size = decoder.read_int()
decoder.skip(block_size)
else:
for _ in range(block_count):
self.key.skip(decoder)
self.value.skip(decoder)
block_count = decoder.read_int()
Comment thread
Fokko marked this conversation as resolved.
Outdated


class ConstructReader(SchemaVisitor[Reader]):
def schema(self, schema: Schema, struct_result: Reader) -> Reader:
return struct_result

def struct(self, struct: StructType, field_results: list[Reader]) -> Reader:
return StructReader(tuple(field_results))
return StructReader(tuple(enumerate(field_results)))

def field(self, field: NestedField, field_result: Reader) -> Reader:
return field_result if field.required else OptionReader(field_result)
Expand Down
Loading