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)

Reply via email to