diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index ae36bac7a..814de9ffa 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -1346,7 +1346,7 @@ class Table(ABC): 2 3 y 3 4 z """ # noqa: E501 - on = [on] if isinstance(on, str) else list(iter(on)) + on = [on] if isinstance(on, str) else list(on) return LanceMergeInsertBuilder(self, on) @@ -5236,7 +5236,7 @@ class AsyncTable: 2 3 y 3 4 z """ # noqa: E501 - on = [on] if isinstance(on, str) else list(iter(on)) + on = [on] if isinstance(on, str) else list(on) return LanceMergeInsertBuilder(self, on) diff --git a/python/python/tests/test_table.py b/python/python/tests/test_table.py index 069527b21..2b2493a8a 100644 --- a/python/python/tests/test_table.py +++ b/python/python/tests/test_table.py @@ -2265,6 +2265,31 @@ def test_update_types(mem_db: DBConnection): assert actual == expected +def test_merge_insert_accepts_column_list(mem_db: DBConnection): + table = mem_db.create_table( + "my_table", + data=pa.table({"a": [1], "b": [2]}), + ) + + builder = table.merge_insert(["a", "b"]) + + assert builder._on == ["a", "b"] + + +@pytest.mark.asyncio +async def test_merge_insert_accepts_column_list_async( + mem_db_async: AsyncConnection, +): + table = await mem_db_async.create_table( + "my_table", + data=pa.table({"a": [1], "b": [2]}), + ) + + builder = table.merge_insert(["a", "b"]) + + assert builder._on == ["a", "b"] + + def test_merge_insert(mem_db: DBConnection): table = mem_db.create_table( "my_table",