This is an automated email from the ASF dual-hosted git repository.
kojiromike pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/avro.git
The following commit(s) were added to refs/heads/master by this push:
new b2abc70 AVRO-2921: Type Fixes for avro.datafile (#1262)
b2abc70 is described below
commit b2abc70bb85a5b515790d0ada7be212fadd5cce1
Author: Michael A. Smith <[email protected]>
AuthorDate: Fri Jun 18 19:18:46 2021 -0400
AVRO-2921: Type Fixes for avro.datafile (#1262)
...and other small modules
- Type hint and refactor avro.datafile
- Type hint avro.test.mock_tether_parent
- Type hint avro.timezones
Also implements `__slots__` on avro.datafile to reduce the memory footprint
of instances.
---
lang/py/avro/datafile.py | 293 ++++++++++++++++++-----------
lang/py/avro/test/mock_tether_parent.py | 46 ++---
lang/py/avro/test/sample_http_server.py | 24 +--
lang/py/avro/test/test_bench.py | 53 +++---
lang/py/avro/test/test_datafile_interop.py | 35 ++--
lang/py/avro/test/test_tether_task.py | 7 +-
lang/py/avro/timezones.py | 13 +-
lang/py/avro/{timezones.py => utils.py} | 33 +---
8 files changed, 278 insertions(+), 226 deletions(-)
diff --git a/lang/py/avro/datafile.py b/lang/py/avro/datafile.py
index b766854..0e0bd5d 100644
--- a/lang/py/avro/datafile.py
+++ b/lang/py/avro/datafile.py
@@ -17,28 +17,28 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-"""Read/Write Avro File Object Containers."""
+"""
+Read/Write Avro File Object Containers.
+https://avro.apache.org/docs/current/spec.html#Object+Container+Files
+"""
import io
import json
-import os
-import random
-import zlib
+from types import TracebackType
+from typing import BinaryIO, MutableMapping, Optional, Type
import avro.codecs
import avro.errors
import avro.io
import avro.schema
+import avro.utils
-#
-# Constants
-#
VERSION = 1
MAGIC = bytes(b"Obj" + bytearray([VERSION]))
MAGIC_SIZE = len(MAGIC)
SYNC_SIZE = 16
SYNC_INTERVAL = 4000 * SYNC_SIZE # TODO(hammer): make configurable
-META_SCHEMA = avro.schema.parse(
+META_SCHEMA: avro.schema.RecordSchema = avro.schema.parse(
json.dumps(
{
"type": "record",
@@ -59,78 +59,105 @@ VALID_ENCODINGS = ["binary"] # not used yet
CODEC_KEY = "avro.codec"
SCHEMA_KEY = "avro.schema"
-#
-# Write Path
-#
+class _DataFileMetadata:
+ """
+ Mixin for meta properties.
-class _DataFile:
- """Mixin for methods common to both reading and writing."""
+ Files may include arbitrary user-specified metadata.
+ File metadata is written as if defined by the following map schema:
- block_count = 0
- _meta = None
- _sync_marker = None
+ `{"type": "map", "values": "bytes"}`
- def __enter__(self):
- return self
+ All metadata properties that start with "avro." are reserved.
+ The following file metadata properties are currently used:
- def __exit__(self, type, value, traceback):
- # Perform a close if there's no exception
- if type is None:
- self.close()
+ - `avro.schema` contains the schema of objects stored in the file, as JSON
data (required).
+ - `avro.codec`, the name of the compression codec used to compress blocks,
as a string.
+ Implementations are required to support the following codecs: "null" and
"deflate".
+ If codec is absent, it is assumed to be "null". See avro.codecs for
implementation details.
+ """
+
+ __slots__ = ("_meta",)
+
+ _meta: MutableMapping[str, bytes]
- def get_meta(self, key):
+ def get_meta(self, key: str) -> Optional[bytes]:
+ """Get the metadata property at `key`."""
return self.meta.get(key)
- def set_meta(self, key, val):
+ def set_meta(self, key: str, val: bytes) -> None:
+ """Set the metadata property at `key`."""
self.meta[key] = val
- @property
- def sync_marker(self):
- return self._sync_marker
+ def del_meta(self, key: str) -> None:
+ """Unset the metadata property at `key`."""
+ del self.meta[key]
@property
- def meta(self):
- """Read-only dictionary of metadata for this datafile."""
- if self._meta is None:
+ def meta(self) -> MutableMapping[str, bytes]:
+ """Get the dictionary of metadata for this datafile."""
+ if not hasattr(self, "_meta"):
self._meta = {}
return self._meta
@property
- def codec(self):
- """Meta are stored as bytes, but codec is returned as a string."""
- try:
- return self.get_meta(CODEC_KEY).decode()
- except AttributeError:
- return "null"
-
- @codec.setter
- def codec(self, value):
- """Meta are stored as bytes, but codec is set as a string."""
- if value not in VALID_CODECS:
- raise avro.errors.DataFileException(f"Unknown codec: {value!r}")
- self.set_meta(CODEC_KEY, value.encode())
-
- @property
- def schema(self):
- """Meta are stored as bytes, but schema is returned as a string."""
- return self.get_meta(SCHEMA_KEY).decode()
+ def schema(self) -> str:
+ """Get the schema of objects stored in the file from the file's
metadata."""
+ schema_str = self.get_meta(SCHEMA_KEY)
+ if schema_str:
+ return schema_str.decode()
+ raise avro.errors.DataFileException("Missing required schema
metadata.")
@schema.setter
- def schema(self, value):
- """Meta are stored as bytes, but schema is set as a string."""
+ def schema(self, value: str) -> None:
+ """Set the schema of objects stored in the file's metadata."""
self.set_meta(SCHEMA_KEY, value.encode())
+ @property
+ def codec(self) -> str:
+ """Get the file's compression codec algorithm from the file's
metadata."""
+ codec = self.get_meta(CODEC_KEY)
+ return "null" if codec is None else codec.decode()
-class DataFileWriter(_DataFile):
+ @codec.setter
+ def codec(self, value: str) -> None:
+ """Set the file's compression codec algorithm in the file's
metadata."""
+ if value not in VALID_CODECS:
+ raise avro.errors.DataFileException(f"Unknown codec: {value!r}")
+ self.set_meta(CODEC_KEY, value.encode())
- # TODO(hammer): make 'encoder' a metadata property
- def __init__(self, writer, datum_writer, writers_schema=None,
codec=NULL_CODEC):
- """
- If the schema is not present, presume we're appending.
+ @codec.deleter
+ def codec(self) -> None:
+ """Unset the file's compression codec algorithm from the file's
metadata."""
+ self.del_meta(CODEC_KEY)
+
+
+class DataFileWriter(_DataFileMetadata):
+ __slots__ = (
+ "_buffer_encoder",
+ "_buffer_writer",
+ "_datum_writer",
+ "_encoder",
+ "_header_written",
+ "_writer",
+ "block_count",
+ "sync_marker",
+ )
- @param writer: File-like object to write into.
- """
+ _buffer_encoder: avro.io.BinaryEncoder
+ _buffer_writer: io.BytesIO # BinaryIO would have better compatibility,
but we use getvalue right now.
+ _datum_writer: avro.io.DatumWriter
+ _encoder: avro.io.BinaryEncoder
+ _header_written: bool
+ _writer: BinaryIO
+ block_count: int
+ sync_marker: bytes
+
+ def __init__(
+ self, writer: BinaryIO, datum_writer: avro.io.DatumWriter,
writers_schema: Optional[avro.schema.Schema] = None, codec: str = NULL_CODEC
+ ) -> None:
+ """If the schema is not present, presume we're appending."""
self._writer = writer
self._encoder = avro.io.BinaryEncoder(writer)
self._datum_writer = datum_writer
@@ -139,18 +166,13 @@ class DataFileWriter(_DataFile):
self.block_count = 0
self._header_written = False
- if writers_schema is not None:
- self._sync_marker = generate_sixteen_random_bytes()
- self.codec = codec
- self.schema = str(writers_schema)
- self.datum_writer.writers_schema = writers_schema
- else:
+ if writers_schema is None:
# open writer for reading to collect metadata
dfr = DataFileReader(writer, avro.io.DatumReader())
# TODO(hammer): collect arbitrary metadata
# collect metadata
- self._sync_marker = dfr.sync_marker
+ self.sync_marker = dfr.sync_marker
self.codec = dfr.codec
# get schema used to write existing file
@@ -160,33 +182,39 @@ class DataFileWriter(_DataFile):
# seek to the end of the file and prepare for writing
writer.seek(0, 2)
self._header_written = True
+ return
+ self.sync_marker = avro.utils.randbytes(16)
+ self.codec = codec
+ self.schema = str(writers_schema)
+ self.datum_writer.writers_schema = writers_schema
- # read-only properties
- writer = property(lambda self: self._writer)
- encoder = property(lambda self: self._encoder)
- datum_writer = property(lambda self: self._datum_writer)
- buffer_writer = property(lambda self: self._buffer_writer)
- buffer_encoder = property(lambda self: self._buffer_encoder)
+ @property
+ def writer(self) -> BinaryIO:
+ return self._writer
- def _write_header(self):
- header = {"magic": MAGIC, "meta": self.meta, "sync": self.sync_marker}
- self.datum_writer.write_data(META_SCHEMA, header, self.encoder)
- self._header_written = True
+ @property
+ def encoder(self) -> avro.io.BinaryEncoder:
+ return self._encoder
@property
- def codec(self):
- """Meta are stored as bytes, but codec is returned as a string."""
- return self.get_meta(CODEC_KEY).decode()
+ def datum_writer(self) -> avro.io.DatumWriter:
+ return self._datum_writer
- @codec.setter
- def codec(self, value):
- """Meta are stored as bytes, but codec is set as a string."""
- if value not in VALID_CODECS:
- raise avro.errors.DataFileException(f"Unknown codec: {value!r}")
- self.set_meta(CODEC_KEY, value.encode())
+ @property
+ def buffer_writer(self) -> io.BytesIO:
+ return self._buffer_writer
+
+ @property
+ def buffer_encoder(self) -> avro.io.BinaryEncoder:
+ return self._buffer_encoder
+
+ def _write_header(self) -> None:
+ header = {"magic": MAGIC, "meta": self.meta, "sync": self.sync_marker}
+ self.datum_writer.write_data(META_SCHEMA, header, self.encoder)
+ self._header_written = True
# TODO(hammer): make a schema for blocks and use datum_writer
- def _write_block(self):
+ def _write_block(self) -> None:
if not self._header_written:
self._write_header()
@@ -213,7 +241,7 @@ class DataFileWriter(_DataFile):
self.buffer_writer.seek(0)
self.block_count = 0
- def append(self, datum):
+ def append(self, datum: object) -> None:
"""Append a datum to the file."""
self.datum_writer.write(datum, self.buffer_encoder)
self.block_count += 1
@@ -222,7 +250,7 @@ class DataFileWriter(_DataFile):
if self.buffer_writer.tell() >= SYNC_INTERVAL:
self._write_block()
- def sync(self):
+ def sync(self) -> int:
"""
Return the current position as a value that may be passed to
DataFileReader.seek(long). Forces the end of the current block,
@@ -231,20 +259,45 @@ class DataFileWriter(_DataFile):
self._write_block()
return self.writer.tell()
- def flush(self):
+ def flush(self) -> None:
"""Flush the current state of the file, including metadata."""
self._write_block()
self.writer.flush()
- def close(self):
+ def close(self) -> None:
"""Close the file."""
self.flush()
self.writer.close()
+ def __enter__(self) -> "DataFileWriter":
+ return self
+
+ def __exit__(self, type_: Optional[Type[BaseException]], value:
Optional[BaseException], traceback: Optional[TracebackType]) -> None:
+ """Perform a close if there's no exception."""
+ if type_ is None:
+ self.close()
+
-class DataFileReader(_DataFile):
+class DataFileReader(_DataFileMetadata):
"""Read files written by DataFileWriter."""
+ __slots__ = (
+ "_datum_decoder",
+ "_datum_reader",
+ "_file_length",
+ "_raw_decoder",
+ "_reader",
+ "block_count",
+ "sync_marker",
+ )
+ _datum_decoder: Optional[avro.io.BinaryDecoder]
+ _datum_reader: avro.io.DatumReader
+ _file_length: int
+ _raw_decoder: avro.io.BinaryDecoder
+ _reader: BinaryIO
+ block_count: int
+ sync_marker: bytes
+
# TODO(hammer): allow user to specify expected schema?
# TODO(hammer): allow user to specify the encoder
@@ -264,17 +317,30 @@ class DataFileReader(_DataFile):
self.block_count = 0
self.datum_reader.writers_schema = avro.schema.parse(self.schema)
- def __iter__(self):
+ def __iter__(self) -> "DataFileReader":
return self
- # read-only properties
- reader = property(lambda self: self._reader)
- raw_decoder = property(lambda self: self._raw_decoder)
- datum_decoder = property(lambda self: self._datum_decoder)
- datum_reader = property(lambda self: self._datum_reader)
- file_length = property(lambda self: self._file_length)
+ @property
+ def reader(self) -> BinaryIO:
+ return self._reader
+
+ @property
+ def raw_decoder(self) -> avro.io.BinaryDecoder:
+ return self._raw_decoder
+
+ @property
+ def datum_decoder(self) -> Optional[avro.io.BinaryDecoder]:
+ return self._datum_decoder
+
+ @property
+ def datum_reader(self) -> avro.io.DatumReader:
+ return self._datum_reader
- def determine_file_length(self):
+ @property
+ def file_length(self) -> int:
+ return self._file_length
+
+ def determine_file_length(self) -> int:
"""
Get file length and leave file cursor where we found it.
"""
@@ -284,10 +350,10 @@ class DataFileReader(_DataFile):
self.reader.seek(remember_pos)
return file_length
- def is_EOF(self):
+ def is_EOF(self) -> bool:
return self.reader.tell() == self.file_length
- def _read_header(self):
+ def _read_header(self) -> None:
# seek to the beginning of the file to get magic block
self.reader.seek(0, 0)
@@ -302,25 +368,25 @@ class DataFileReader(_DataFile):
self._meta = header["meta"]
# set sync marker
- self._sync_marker = header["sync"]
+ self.sync_marker = header["sync"]
- def _read_block_header(self):
+ def _read_block_header(self) -> None:
self.block_count = self.raw_decoder.read_long()
codec = avro.codecs.get_codec(self.codec)
self._datum_decoder = codec.decompress(self.raw_decoder)
- def _skip_sync(self):
+ def _skip_sync(self) -> bool:
"""
Read the length of the sync marker; if it matches the sync marker,
return True. Otherwise, seek back to where we started and return False.
"""
proposed_sync_marker = self.reader.read(SYNC_SIZE)
- if proposed_sync_marker != self.sync_marker:
- self.reader.seek(-SYNC_SIZE, 1)
- return False
- return True
+ if proposed_sync_marker == self.sync_marker:
+ return True
+ self.reader.seek(-SYNC_SIZE, 1)
+ return False
- def __next__(self):
+ def __next__(self) -> object:
"""Return the next datum in the file."""
while self.block_count == 0:
if self.is_EOF() or (self._skip_sync() and self.is_EOF()):
@@ -331,13 +397,14 @@ class DataFileReader(_DataFile):
self.block_count -= 1
return datum
- def close(self):
+ def close(self) -> None:
"""Close this reader."""
self.reader.close()
+ def __enter__(self) -> "DataFileReader":
+ return self
-def generate_sixteen_random_bytes():
- try:
- return os.urandom(16)
- except NotImplementedError:
- return bytes(random.randrange(256) for i in range(16))
+ def __exit__(self, type_: Optional[Type[BaseException]], value:
Optional[BaseException], traceback: Optional[TracebackType]) -> None:
+ """Perform a close if there's no exception."""
+ if type_ is None:
+ self.close()
diff --git a/lang/py/avro/test/mock_tether_parent.py
b/lang/py/avro/test/mock_tether_parent.py
index a1ec629..9c7c844 100644
--- a/lang/py/avro/test/mock_tether_parent.py
+++ b/lang/py/avro/test/mock_tether_parent.py
@@ -18,8 +18,8 @@
# limitations under the License.
import http.server
-import socket
import sys
+from typing import Mapping
import avro.errors
import avro.ipc
@@ -38,7 +38,7 @@ class MockParentResponder(avro.ipc.Responder):
def __init__(self) -> None:
super().__init__(avro.tether.tether_task.outputProtocol)
- def invoke(self, message, request) -> None:
+ def invoke(self, message: avro.protocol.Message, request: Mapping[str,
str]) -> None:
response = f"MockParentResponder: Received '{message.name}'"
responses = {
"configure": f"{response}': inputPort={request.get('port')}",
@@ -52,7 +52,7 @@ class MockParentResponder(avro.ipc.Responder):
class MockParentHandler(http.server.BaseHTTPRequestHandler):
"""Create a handler for the parent."""
- def do_POST(self):
+ def do_POST(self) -> None:
self.responder = MockParentResponder()
call_request_reader = avro.ipc.FramedReader(self.rfile)
call_request = call_request_reader.read_framed_message()
@@ -64,22 +64,26 @@ class MockParentHandler(http.server.BaseHTTPRequestHandler):
resp_writer.write_framed_message(resp_body)
+def main() -> None:
+ global SERVER_ADDRESS
+
+ if len(sys.argv) != 3 or sys.argv[1].lower() != "start_server":
+ raise avro.errors.UsageError("Usage: mock_tether_parent start_server
port")
+
+ try:
+ port = int(sys.argv[2])
+ except ValueError as e:
+ raise avro.errors.UsageError("Usage: mock_tether_parent start_server
port") from e
+
+ SERVER_ADDRESS = (SERVER_ADDRESS[0], port)
+ print(f"mock_tether_parent: Launching Server on Port: {SERVER_ADDRESS[1]}")
+
+ # flush the output so it shows up in the parent process
+ sys.stdout.flush()
+ parent_server = http.server.HTTPServer(SERVER_ADDRESS, MockParentHandler)
+ parent_server.allow_reuse_address = True
+ parent_server.serve_forever()
+
+
if __name__ == "__main__":
- if len(sys.argv) <= 1:
- raise avro.errors.UsageError("Usage: mock_tether_parent command")
-
- cmd = sys.argv[1].lower()
- if sys.argv[1] == "start_server":
- if len(sys.argv) == 3:
- port = int(sys.argv[2])
- else:
- raise avro.errors.UsageError("Usage: mock_tether_parent
start_server port")
-
- SERVER_ADDRESS = (SERVER_ADDRESS[0], port)
- print(f"mock_tether_parent: Launching Server on Port:
{SERVER_ADDRESS[1]}")
-
- # flush the output so it shows up in the parent process
- sys.stdout.flush()
- parent_server = http.server.HTTPServer(SERVER_ADDRESS,
MockParentHandler)
- parent_server.allow_reuse_address = True
- parent_server.serve_forever()
+ main()
diff --git a/lang/py/avro/test/sample_http_server.py
b/lang/py/avro/test/sample_http_server.py
index e6873ab..670a386 100644
--- a/lang/py/avro/test/sample_http_server.py
+++ b/lang/py/avro/test/sample_http_server.py
@@ -17,16 +17,13 @@
# See the License for the specific language governing permissions and
# limitations under the License.
+import http.server
import json
+from typing import Mapping
import avro.ipc
import avro.protocol
-try:
- import BaseHTTPServer as http_server # type: ignore
-except ImportError:
- import http.server as http_server # type: ignore
-
MAIL_PROTOCOL_JSON = json.dumps(
{
"namespace": "example.proto",
@@ -49,18 +46,19 @@ SERVER_ADDRESS = ("localhost", 9090)
class MailResponder(avro.ipc.Responder):
- def __init__(self):
- avro.ipc.Responder.__init__(self, MAIL_PROTOCOL)
+ def __init__(self) -> None:
+ super().__init__(MAIL_PROTOCOL)
- def invoke(self, message, request):
+ def invoke(self, message: avro.protocol.Message, request: Mapping[str,
Mapping[str, str]]) -> str:
if message.name == "send":
return f"Sent message to {request['message']['to']} from
{request['message']['from']} with body {request['message']['body']}"
if message.name == "replay":
return "replay"
+ raise RuntimeError
-class MailHandler(http_server.BaseHTTPRequestHandler):
- def do_POST(self):
+class MailHandler(http.server.BaseHTTPRequestHandler):
+ def do_POST(self) -> None:
self.responder = MailResponder()
call_request_reader = avro.ipc.FramedReader(self.rfile)
call_request = call_request_reader.read_framed_message()
@@ -72,7 +70,11 @@ class MailHandler(http_server.BaseHTTPRequestHandler):
resp_writer.write_framed_message(resp_body)
-if __name__ == "__main__":
+def main():
mail_server = http_server.HTTPServer(SERVER_ADDRESS, MailHandler)
mail_server.allow_reuse_address = True
mail_server.serve_forever()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/lang/py/avro/test/test_bench.py b/lang/py/avro/test/test_bench.py
index 34c926b..e1d4416 100644
--- a/lang/py/avro/test/test_bench.py
+++ b/lang/py/avro/test/test_bench.py
@@ -26,13 +26,16 @@ import tempfile
import timeit
import unittest
import unittest.mock
+from pathlib import Path
+from typing import List, Mapping, Sequence
import avro.datafile
import avro.io
import avro.schema
+import avro.utils
TYPES = ("A", "CNAME")
-SCHEMA = avro.schema.parse(
+SCHEMA: avro.schema.RecordSchema = avro.schema.parse(
json.dumps(
{
"type": "record",
@@ -52,77 +55,71 @@ MAX_WRITE_SECONDS = 3 if platform.python_implementation()
== "PyPy" else 1
MAX_READ_SECONDS = 3 if platform.python_implementation() == "PyPy" else 1
-try: # pragma: no cover
- randbytes = random.randbytes # type: ignore
-except AttributeError: # pragma: no cover
-
- def randbytes(n):
- """Polyfill for random.randbytes in Python < 3.9"""
- return random.choices(range(256), k=n)
-
-
class TestBench(unittest.TestCase):
- def test_minimum_speed(self):
- with tempfile.NamedTemporaryFile(suffix="avr") as temp:
+ def test_minimum_speed(self) -> None:
+ with tempfile.NamedTemporaryFile(suffix="avr") as temp_:
pass
+ temp = Path(temp_.name)
self.assertLess(
- time_writes(temp.name, NUMBER_OF_TESTS),
+ time_writes(temp, NUMBER_OF_TESTS),
MAX_WRITE_SECONDS,
f"Took longer than {MAX_WRITE_SECONDS} second(s) to write the test
file with {NUMBER_OF_TESTS} values.",
)
self.assertLess(
- time_read(temp.name),
+ time_read(temp),
MAX_READ_SECONDS,
f"Took longer than {MAX_READ_SECONDS} second(s) to read the test
file with {NUMBER_OF_TESTS} values.",
)
-def rand_name():
+def rand_name() -> str:
return "".join(random.sample(string.ascii_lowercase, 15))
-def rand_ip():
- return ".".join(map(str, randbytes(4)))
+def rand_ip() -> str:
+ return ".".join(map(str, avro.utils.randbytes(4)))
-def picks(n):
+def picks(n) -> Sequence[Mapping[str, str]]:
return [{"query": rand_name(), "response": rand_ip(), "type":
random.choice(TYPES)} for _ in range(n)]
-def time_writes(path, number):
- with avro.datafile.DataFileWriter(open(path, "wb"), WRITER, SCHEMA) as dw:
+def time_writes(path: Path, number: int) -> float:
+ with avro.datafile.DataFileWriter(path.open("wb"), WRITER, SCHEMA) as dw:
globals_ = {"dw": dw, "picks": picks(number)}
return timeit.timeit("dw.append(next(p))", number=number,
setup="p=iter(picks)", globals=globals_)
-def time_read(path):
+def time_read(path: Path) -> float:
"""
Time how long it takes to read the file written in the `write` function.
We only do this once, because the size of the file is defined by the
number sent to `write`.
"""
- with avro.datafile.DataFileReader(open(path, "rb"), READER) as dr:
+ with avro.datafile.DataFileReader(path.open("rb"), READER) as dr:
return timeit.timeit("tuple(dr)", number=1, globals={"dr": dr})
-def parse_args(): # pragma: no cover
+def parse_args() -> argparse.Namespace: # pragma: no cover
parser = argparse.ArgumentParser(description="Benchmark writing some
random avro.")
parser.add_argument(
"--number",
"-n",
type=int,
- default=timeit.default_number,
+ default=getattr(timeit, "default_number", 1000000),
help="how many times to run",
)
return parser.parse_args()
-def main(): # pragma: no cover
+def main() -> None: # pragma: no cover
args = parse_args()
- with tempfile.NamedTemporaryFile(suffix=".avr") as temp:
+ with tempfile.NamedTemporaryFile(suffix=".avr") as temp_:
pass
+ temp = Path(temp_.name)
+
print(f"Using file {temp.name}")
- print(f"Writing: {time_writes(temp.name, args.number)}")
- print(f"Reading: {time_read(temp.name)}")
+ print(f"Writing: {time_writes(temp, args.number)}")
+ print(f"Reading: {time_read(temp)}")
if __name__ == "__main__": # pragma: no cover
diff --git a/lang/py/avro/test/test_datafile_interop.py
b/lang/py/avro/test/test_datafile_interop.py
index 382520c..aeb6c21 100644
--- a/lang/py/avro/test/test_datafile_interop.py
+++ b/lang/py/avro/test/test_datafile_interop.py
@@ -19,35 +19,32 @@
import os
import unittest
+from pathlib import Path
+from typing import Optional, cast
import avro
import avro.datafile
import avro.io
-_INTEROP_DATA_DIR = os.path.join(os.path.dirname(avro.__file__), "test",
"interop", "data")
+_INTEROP_DATA_DIR = Path(avro.__file__).parent / "test" / "interop" / "data"
@unittest.skipUnless(os.path.exists(_INTEROP_DATA_DIR), f"{_INTEROP_DATA_DIR}
does not exist")
class TestDataFileInterop(unittest.TestCase):
- def test_interop(self):
+ def test_interop(self) -> None:
"""Test Interop"""
- for f in os.listdir(_INTEROP_DATA_DIR):
- filename = os.path.join(_INTEROP_DATA_DIR, f)
- assert os.stat(filename).st_size > 0
- base_ext = os.path.splitext(os.path.basename(f))[0].split("_", 1)
- if len(base_ext) < 2 or base_ext[1] in avro.datafile.VALID_CODECS:
- print(f"READING {f}\n")
-
- # read data in binary from file
- datum_reader = avro.io.DatumReader()
- with open(filename, "rb") as reader:
- dfr = avro.datafile.DataFileReader(reader, datum_reader)
- i = 0
- for i, datum in enumerate(dfr, 1):
- assert datum is not None
- assert i > 0
- else:
- print(f"SKIPPING {f} due to an unsupported codec\n")
+ datum: Optional[object] = None
+ for filename in _INTEROP_DATA_DIR.iterdir():
+ self.assertGreater(os.stat(filename).st_size, 0)
+ base_ext = filename.stem.split("_", 1)
+ if len(base_ext) < 2 or base_ext[1] not in
avro.codecs.KNOWN_CODECS:
+ print(f"SKIPPING {filename} due to an unsupported codec\n")
+ continue
+ i = None
+ with self.subTest(filename=filename),
avro.datafile.DataFileReader(filename.open("rb"), avro.io.DatumReader()) as dfr:
+ for i, datum in enumerate(cast(avro.datafile.DataFileReader,
dfr), 1):
+ self.assertIsNotNone(datum)
+ self.assertIsNotNone(i)
if __name__ == "__main__":
diff --git a/lang/py/avro/test/test_tether_task.py
b/lang/py/avro/test/test_tether_task.py
index f00251e..5a4e2b2 100644
--- a/lang/py/avro/test/test_tether_task.py
+++ b/lang/py/avro/test/test_tether_task.py
@@ -36,7 +36,7 @@ class TestTetherTask(unittest.TestCase):
TODO: We should validate the the server response by looking at stdout
"""
- def test_tether_task(self):
+ def test_tether_task(self) -> None:
"""
Test that the tether_task is working. We run the mock_tether_parent in
a separate
subprocess
@@ -62,6 +62,8 @@ class TestTetherTask(unittest.TestCase):
# ***************************************************************
# Test the mapper
+ if avro.tether.tether_task.TaskType is None:
+ self.fail()
task.configure(
avro.tether.tether_task.TaskType.MAP,
str(task.inschema),
@@ -89,11 +91,10 @@ class TestTetherTask(unittest.TestCase):
)
# Serialize some data so we can send it to the input function
- datum = {"key": "word", "value": 2}
writer = io.BytesIO()
encoder = avro.io.BinaryEncoder(writer)
datum_writer = avro.io.DatumWriter(task.midschema)
- datum_writer.write(datum, encoder)
+ datum_writer.write({"key": "word", "value": 2}, encoder)
writer.seek(0)
data = writer.read()
diff --git a/lang/py/avro/timezones.py b/lang/py/avro/timezones.py
index 2d7667d..28b114f 100644
--- a/lang/py/avro/timezones.py
+++ b/lang/py/avro/timezones.py
@@ -18,16 +18,17 @@
# limitations under the License.
import datetime
+from typing import Optional
class UTCTzinfo(datetime.tzinfo):
- def utcoffset(self, dt):
+ def utcoffset(self, dt: Optional[datetime.datetime] = None) ->
datetime.timedelta:
return datetime.timedelta(0)
- def tzname(self, dt):
+ def tzname(self, dt: Optional[datetime.datetime] = None) -> str:
return "UTC"
- def dst(self, dt):
+ def dst(self, dt: Optional[datetime.datetime] = None) ->
datetime.timedelta:
return datetime.timedelta(0)
@@ -36,13 +37,13 @@ utc = UTCTzinfo()
# Test Time Zone with fixed offset and no DST
class TSTTzinfo(datetime.tzinfo):
- def utcoffset(self, dt):
+ def utcoffset(self, dt: Optional[datetime.datetime] = None) ->
datetime.timedelta:
return datetime.timedelta(hours=10)
- def tzname(self, dt):
+ def tzname(self, dt: Optional[datetime.datetime] = None) -> str:
return "TST"
- def dst(self, dt):
+ def dst(self, dt: Optional[datetime.datetime] = None) ->
datetime.timedelta:
return datetime.timedelta(0)
diff --git a/lang/py/avro/timezones.py b/lang/py/avro/utils.py
similarity index 60%
copy from lang/py/avro/timezones.py
copy to lang/py/avro/utils.py
index 2d7667d..76cc8c7 100644
--- a/lang/py/avro/timezones.py
+++ b/lang/py/avro/utils.py
@@ -17,33 +17,16 @@
# See the License for the specific language governing permissions and
# limitations under the License.
-import datetime
+"""
+Arbitrary utilities and polyfills.
+"""
+import random
-class UTCTzinfo(datetime.tzinfo):
- def utcoffset(self, dt):
- return datetime.timedelta(0)
- def tzname(self, dt):
- return "UTC"
+def _randbytes(n: int) -> bytes:
+ """Polyfill for random.randbytes in Python < 3.9"""
+ return bytes(random.choices(range(256), k=n))
- def dst(self, dt):
- return datetime.timedelta(0)
-
-utc = UTCTzinfo()
-
-
-# Test Time Zone with fixed offset and no DST
-class TSTTzinfo(datetime.tzinfo):
- def utcoffset(self, dt):
- return datetime.timedelta(hours=10)
-
- def tzname(self, dt):
- return "TST"
-
- def dst(self, dt):
- return datetime.timedelta(0)
-
-
-tst = TSTTzinfo()
+randbytes = getattr(random, "randbytes", _randbytes)