diff --git a/sqlite_utils/cli.py b/sqlite_utils/cli.py index 86eddfb..173873a 100644 --- a/sqlite_utils/cli.py +++ b/sqlite_utils/cli.py @@ -6,7 +6,7 @@ import hashlib import pathlib import sqlite_utils from sqlite_utils.db import AlterError, BadMultiValues, DescIndex -from sqlite_utils.utils import maximize_csv_field_size_limit +from sqlite_utils.utils import maximize_csv_field_size_limit, _extra_key_strategy_row from sqlite_utils import recipes import textwrap import inspect @@ -797,6 +797,15 @@ _import_options = ( "--encoding", help="Character encoding for input, defaults to utf-8", ), + click.option( + "--ignore-extras", + is_flag=True, + help="If a CSV line has more than the expected number of values, ignore the extras", + ), + click.option( + "--extras-key", + help="If a CSV line has more than the expected number of values put them in a list in this column", + ), ) @@ -885,6 +894,8 @@ def insert_upsert_implementation( sniff, no_headers, encoding, + ignore_extras, + extras_key, batch_size, alter, upsert, @@ -909,6 +920,10 @@ def insert_upsert_implementation( raise click.ClickException("--flatten cannot be used with --csv or --tsv") if encoding and not (csv or tsv): raise click.ClickException("--encoding must be used with --csv or --tsv") + if ignore_extras and extras_key: + raise click.ClickException( + "--ignore-extras and --extras-key cannot be used together" + ) if pk and len(pk) == 1: pk = pk[0] encoding = encoding or "utf-8-sig" @@ -942,7 +957,10 @@ def insert_upsert_implementation( reader = itertools.chain([first_row], reader) else: headers = first_row - docs = (dict(zip(headers, row)) for row in reader) + docs = ( + _extra_key_strategy_row(headers, row, ignore_extras, extras_key) + for row in reader + ) if detect_types: tracker = TypeTracker() docs = tracker.wrap(docs) @@ -1101,6 +1119,8 @@ def insert( sniff, no_headers, encoding, + ignore_extras, + extras_key, batch_size, alter, detect_types, @@ -1176,6 +1196,8 @@ def insert( sniff, no_headers, encoding, + ignore_extras, + extras_key, batch_size, alter=alter, upsert=False, @@ -1214,6 +1236,8 @@ def upsert( sniff, no_headers, encoding, + ignore_extras, + extras_key, alter, not_null, default, @@ -1254,6 +1278,8 @@ def upsert( sniff, no_headers, encoding, + ignore_extras, + extras_key, batch_size, alter=alter, upsert=True, @@ -1297,6 +1323,8 @@ def bulk( sniff, no_headers, encoding, + ignore_extras, + extras_key, load_extension, ): """ @@ -1331,6 +1359,8 @@ def bulk( sniff=sniff, no_headers=no_headers, encoding=encoding, + ignore_extras=ignore_extras, + extras_key=extras_key, batch_size=batch_size, alter=False, upsert=False, diff --git a/sqlite_utils/utils.py b/sqlite_utils/utils.py index d2ccc5f..06dd45d 100644 --- a/sqlite_utils/utils.py +++ b/sqlite_utils/utils.py @@ -9,7 +9,7 @@ import json import os import sys from . import recipes -from typing import Dict, cast, BinaryIO, Iterable, Optional, Tuple, Type +from typing import Any, Dict, cast, BinaryIO, Iterable, Optional, Tuple, Type import click @@ -219,6 +219,25 @@ def _extra_key_strategy( yield row +def _extra_key_strategy_row( + headers: Iterable[str], + row: Iterable[Any], + ignore_extras: Optional[bool] = False, + extras_key: Optional[str] = None, +) -> dict: + # Logic for handling CSV rows with more values than there are headings + if len(headers) <= len(row) or ignore_extras: + return dict(zip(headers, row)) + # There are more row items than headers, what to do with the extras? + if not ignore_extras: + extras = row[len(headers) :] + raise RowError("Row {} contained these extra values: {}".format(row, extras)) + if extras_key: + d = dict(zip(headers, row)) + d[extras_key] = row[len(headers) :] + return d + + def rows_from_file( fp: BinaryIO, format: Optional[Format] = None, diff --git a/tests/test_cli_insert.py b/tests/test_cli_insert.py index 45ad362..e24e654 100644 --- a/tests/test_cli_insert.py +++ b/tests/test_cli_insert.py @@ -222,6 +222,42 @@ def test_insert_csv_tsv(content, options, db_path, tmpdir): assert [{"foo": "1", "bar": "2", "baz": "cat,dog"}] == list(db["data"].rows) +@pytest.mark.parametrize( + "options,expected_rows,expected_error", + [ + (["--csv", "--ignore-extras"], [{"id": "1", "name": "Cleo"}], None), + ( + ["--csv", "--extras-key", "x"], + [{"id": "1", "name": "Cleo", "x": '["extra"]'}], + None, + ), + (["--csv"], None, "Oh no an extra field found"), + ( + ["--csv", "--extras-key", "x", "--ignore-extras"], + None, + "Error: --ignore-extras and --extras-key cannot be used together\n", + ), + ], +) +def test_insert_csv_with_extra_fields( + options, expected_rows, expected_error, db_path, tmpdir +): + db = Database(db_path) + file_path = str(tmpdir / "insert.csv") + open(file_path, "w").write("id,name\n1,Cleo,extra") + result = CliRunner().invoke( + cli.cli, + ["insert", db_path, "data", file_path] + options, + catch_exceptions=False, + ) + if expected_error: + assert result.exit_code == 2 + assert result.output == expected_error + else: + assert result.exit_code == 0 + assert list(db["data"].rows) == expected_rows + + @pytest.mark.parametrize( "options", (