From 7b3fdf0fcd553ddf25b8d606b7fc34f9fd7979df Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Tue, 22 Jun 2021 11:04:32 -0700 Subject: [PATCH] mypy annotations for rows_from_file(), run mypy in CI Refs #289, #279 --- .github/workflows/test.yml | 3 +++ sqlite_utils/cli.py | 2 +- sqlite_utils/db.py | 8 ++++---- sqlite_utils/utils.py | 31 ++++++++++++++++++------------- 4 files changed, 26 insertions(+), 18 deletions(-) diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index f8a79ae..012db32 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -26,11 +26,14 @@ jobs: - name: Install dependencies run: | pip install -e '.[test]' + pip install mypy - name: Optionally install numpy if: matrix.numpy == 1 run: pip install numpy - name: Run tests run: | pytest + - name: run mypy + run: mypy sqlite_utils - name: Check formatting run: black . --check diff --git a/sqlite_utils/cli.py b/sqlite_utils/cli.py index 0a0c0fe..2de6bf3 100644 --- a/sqlite_utils/cli.py +++ b/sqlite_utils/cli.py @@ -1,6 +1,6 @@ import base64 import click -from click_default_group import DefaultGroup +from click_default_group import DefaultGroup # type: ignore from datetime import datetime import hashlib import pathlib diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 1a88dba..acd1726 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -11,7 +11,7 @@ import json import os import pathlib import re -from sqlite_fts4 import rank_bm25 +from sqlite_fts4 import rank_bm25 # type: ignore import sys import textwrap import uuid @@ -39,14 +39,14 @@ USING\s+(?P\w+) # e.g. USING FTS5 ) try: - import pandas as pd + import pandas as pd # type: ignore except ImportError: pd = None try: - import numpy as np + import numpy as np # type: ignore except ImportError: - np = None + np = None # type: ignore Column = namedtuple( "Column", ("cid", "name", "type", "notnull", "default_value", "is_pk") diff --git a/sqlite_utils/utils.py b/sqlite_utils/utils.py index 4dd5f21..4f3c819 100644 --- a/sqlite_utils/utils.py +++ b/sqlite_utils/utils.py @@ -5,17 +5,18 @@ import enum import io import json import os -from typing import Generator +from typing import cast, BinaryIO, Iterable, Optional, Tuple, Type import click try: - import pysqlite3 as sqlite3 - import pysqlite3.dbapi2 + import pysqlite3 as sqlite3 # type: ignore + import pysqlite3.dbapi2 # type: ignore OperationalError = pysqlite3.dbapi2.OperationalError except ImportError: - import sqlite3 + # https://github.com/python/mypy/issues/1153#issuecomment-253842414 + import sqlite3 # type: ignore OperationalError = sqlite3.OperationalError @@ -143,12 +144,11 @@ class RowsFromFileBadJSON(RowsFromFileError): def rows_from_file( - fp, - format=None, - dialect=None, - encoding=None, - detect_types=False, -) -> Generator[dict, None, None]: + fp: BinaryIO, + format: Optional[Format] = None, + dialect: Optional[Type[csv.Dialect]] = None, + encoding: Optional[str] = None, +) -> Tuple[Iterable[dict], Format]: if format == Format.JSON: decoded = json.load(fp) if isinstance(decoded, dict): @@ -159,8 +159,13 @@ def rows_from_file( elif format == Format.NL: return (json.loads(line) for line in fp if line.strip()), Format.NL elif format == Format.CSV: - decoded_fp = io.TextIOWrapper(fp, encoding=encoding or "utf-8-sig") - return csv.DictReader(decoded_fp, dialect=dialect), Format.CSV + use_encoding: str = encoding or "utf-8-sig" + decoded_fp = io.TextIOWrapper(fp, encoding=use_encoding) + if dialect is not None: + reader = csv.DictReader(decoded_fp, dialect=dialect) + else: + reader = csv.DictReader(decoded_fp) + return reader, Format.CSV elif format == Format.TSV: return ( rows_from_file( @@ -170,7 +175,7 @@ def rows_from_file( ) elif format is None: # Detect the format, then call this recursively - buffered = io.BufferedReader(fp, buffer_size=4096) + buffered = io.BufferedReader(cast(io.RawIOBase, fp), buffer_size=4096) first_bytes = buffered.peek(2048).strip() if first_bytes.startswith(b"[") or first_bytes.startswith(b"{"): # TODO: Detect newline-JSON