Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 17 additions & 18 deletions sqlite_utils/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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")

Expand Down
33 changes: 33 additions & 0 deletions tests/test_rows_from_file.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
Loading