diff --git a/sqlite_utils/utils.py b/sqlite_utils/utils.py index 91fd24ddd..35bdaa9b2 100644 --- a/sqlite_utils/utils.py +++ b/sqlite_utils/utils.py @@ -358,17 +358,18 @@ class Format(enum.Enum): ) return _CloseableIterator(iter(rows), decoded_fp), Format.CSV elif format == Format.TSV: + # The inner call applies the extra-field strategy, so the + # ignore_extras= and extras_key= arguments must be passed to it - + # see https://github.com/simonw/sqlite-utils/issues/892 rows, _ = rows_from_file( - fp, format=Format.CSV, dialect=csv.excel_tab, encoding=encoding - ) - return ( - _extra_key_strategy( - cast(Iterable[dict[str | None, object]], rows), - ignore_extras, - extras_key, - ), - Format.TSV, + fp, + format=Format.CSV, + dialect=csv.excel_tab, + encoding=encoding, + ignore_extras=ignore_extras, + extras_key=extras_key, ) + return rows, Format.TSV elif format is None: # Detect the format, then call this recursively buffered = io.BufferedReader(cast(io.RawIOBase, fp), buffer_size=4096) @@ -393,18 +394,16 @@ class Format(enum.Enum): first_bytes.decode(encoding or "utf-8-sig", "ignore") ) rows, _ = rows_from_file( - buffered, format=Format.CSV, dialect=dialect, encoding=encoding + buffered, + format=Format.CSV, + dialect=dialect, + encoding=encoding, + ignore_extras=ignore_extras, + extras_key=extras_key, ) # Make sure we return the format we detected detected_format = Format.TSV if dialect.delimiter == "\t" else Format.CSV - return ( - _extra_key_strategy( - cast(Iterable[dict[str | None, object]], rows), - ignore_extras, - extras_key, - ), - detected_format, - ) + return rows, detected_format else: raise RowsFromFileError("Bad format") diff --git a/tests/test_rows_from_file.py b/tests/test_rows_from_file.py index a8e7f9d76..9d4a4ee99 100644 --- a/tests/test_rows_from_file.py +++ b/tests/test_rows_from_file.py @@ -57,6 +57,39 @@ def test_rows_from_file_extra_fields_strategies(ignore_extras, extras_key, expec assert list_rows == expected +@pytest.mark.parametrize( + "ignore_extras,extras_key,expected", + ( + (True, None, [{"id": "1", "name": "Cleo"}]), + (False, "_rest", [{"id": "1", "name": "Cleo", "_rest": ["oops"]}]), + # expected of None means expect an error: + (False, False, None), + ), +) +def test_rows_from_file_tsv_extra_fields_strategies( + ignore_extras, extras_key, expected +): + # ignore_extras= and extras_key= must apply to TSV as well as CSV, + # see https://github.com/simonw/sqlite-utils/issues/892 + try: + rows, detected_format = rows_from_file( + BytesIO(b"id\tname\r\n1\tCleo\toops"), + format=Format.TSV, + ignore_extras=ignore_extras, + extras_key=extras_key, + ) + list_rows = list(rows) + except RowError: + if expected is None: + # This is fine, + return + else: + # We did not expect an error + raise + assert detected_format == Format.TSV + assert list_rows == expected + + def test_rows_from_file_error_on_string_io(): with pytest.raises(TypeError) as ex: rows_from_file(StringIO("id,name\r\n1,Cleo")) # type: ignore[arg-type]