From b3f813d8df70e2d6d243fe9e5b28991ebba26fe0 Mon Sep 17 00:00:00 2001 From: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 06:23:13 +0000 Subject: [PATCH] fix(python): accept expressions in update filters --- python/python/lancedb/remote/table.py | 11 ++++--- python/python/lancedb/table.py | 42 +++++++++++++++++---------- python/python/tests/test_table.py | 29 ++++++++++++++++++ 3 files changed, 62 insertions(+), 20 deletions(-) diff --git a/python/python/lancedb/remote/table.py b/python/python/lancedb/remote/table.py index acc2f4c9d..0a3ad38d8 100644 --- a/python/python/lancedb/remote/table.py +++ b/python/python/lancedb/remote/table.py @@ -36,6 +36,7 @@ from lancedb._lancedb import ( UpdateResult, ) from lancedb.embeddings.base import EmbeddingFunctionConfig +from lancedb.expr import Expr from lancedb.index import ( FTS, BTree, @@ -855,7 +856,7 @@ class RemoteTable(Table): def update( self, - where: Optional[str] = None, + where: Optional[Union[str, Expr]] = None, values: Optional[dict] = None, *, values_sql: Optional[Dict[str, str]] = None, @@ -866,9 +867,11 @@ class RemoteTable(Table): Parameters ---------- - where: str, optional - The SQL where clause to use when updating rows. For example, 'x = 2' - or 'x IN (1, 2, 3)'. The filter must not be empty, or it will error. + where: str or [Expr][lancedb.expr.Expr], optional + The filter condition. Can be a SQL string or a type-safe + [Expr][lancedb.expr.Expr] built with [col][lancedb.expr.col] and + [lit][lancedb.expr.lit]. The filter must not be empty, or it will + error. values: dict, optional The values to update. The keys are the column names and the values are the values to set. diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index ae36bac7a..2dc8525c3 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -1692,7 +1692,7 @@ class Table(ABC): @abstractmethod def update( self, - where: Optional[str] = None, + where: Optional[Union[str, Expr]] = None, values: Optional[dict] = None, *, values_sql: Optional[Dict[str, str]] = None, @@ -1707,9 +1707,11 @@ class Table(ABC): Parameters ---------- - where: str, optional - The SQL where clause to use when updating rows. For example, 'x = 2' - or 'x IN (1, 2, 3)'. The filter must not be empty, or it will error. + where: str or [Expr][lancedb.expr.Expr], optional + The filter condition. Can be a SQL string or a type-safe + [Expr][lancedb.expr.Expr] built with [col][lancedb.expr.col] and + [lit][lancedb.expr.lit]. The filter must not be empty, or it will + error. values: dict, optional The values to update. The keys are the column names and the values are the values to set. @@ -1727,6 +1729,7 @@ class Table(ABC): Examples -------- >>> import lancedb + >>> from lancedb.expr import col >>> import pandas as pd >>> data = pd.DataFrame({"x": [1, 2, 3], "vector": [[1.0, 2], [3, 4], [5, 6]]}) >>> db = lancedb.connect("./.lancedb") @@ -1736,7 +1739,7 @@ class Table(ABC): 0 1 [1.0, 2.0] 1 2 [3.0, 4.0] 2 3 [5.0, 6.0] - >>> table.update(where="x = 2", values={"vector": [10.0, 10]}) + >>> table.update(where=col("x") == 2, values={"vector": [10.0, 10]}) UpdateResult(rows_updated=1, version=2) >>> table.to_pandas() x vector @@ -3651,7 +3654,7 @@ class LanceTable(Table): def update( self, - where: Optional[str] = None, + where: Optional[Union[str, Expr]] = None, values: Optional[dict] = None, *, values_sql: Optional[Dict[str, str]] = None, @@ -3662,9 +3665,11 @@ class LanceTable(Table): Parameters ---------- - where: str, optional - The SQL where clause to use when updating rows. For example, 'x = 2' - or 'x IN (1, 2, 3)'. The filter must not be empty, or it will error. + where: str or [Expr][lancedb.expr.Expr], optional + The filter condition. Can be a SQL string or a type-safe + [Expr][lancedb.expr.Expr] built with [col][lancedb.expr.col] and + [lit][lancedb.expr.lit]. The filter must not be empty, or it will + error. values: dict, optional The values to update. The keys are the column names and the values are the values to set. @@ -3682,6 +3687,7 @@ class LanceTable(Table): Examples -------- >>> import lancedb + >>> from lancedb.expr import col >>> import pandas as pd >>> data = pd.DataFrame({"x": [1, 2, 3], "vector": [[1.0, 2], [3, 4], [5, 6]]}) >>> db = lancedb.connect("./.lancedb") @@ -3691,7 +3697,7 @@ class LanceTable(Table): 0 1 [1.0, 2.0] 1 2 [3.0, 4.0] 2 3 [5.0, 6.0] - >>> table.update(where="x = 2", values={"vector": [10.0, 10]}) + >>> table.update(where=col("x") == 2, values={"vector": [10.0, 10]}) UpdateResult(rows_updated=1, version=2) >>> table.to_pandas() x vector @@ -5690,7 +5696,7 @@ class AsyncTable: self, updates: Optional[Dict[str, Any]] = None, *, - where: Optional[str] = None, + where: Optional[Union[str, Expr]] = None, updates_sql: Optional[Dict[str, str]] = None, ) -> UpdateResult: """ @@ -5705,9 +5711,11 @@ class AsyncTable: The updates to apply. The keys should be the name of the column to update. The values should be the new values to assign. This is required unless updates_sql is supplied. - where: str, optional - An SQL filter that controls which rows are updated. For example, 'x = 2' - or 'x IN (1, 2, 3)'. Only rows that satisfy this filter will be udpated. + where: str or [Expr][lancedb.expr.Expr], optional + The filter condition. Can be a SQL string or a type-safe + [Expr][lancedb.expr.Expr] built with [col][lancedb.expr.col] and + [lit][lancedb.expr.lit]. Only rows that satisfy this filter will + be updated. updates_sql: dict, optional The updates to apply, expressed as SQL expression strings. The keys should be column names. The values should be SQL expressions. These can be SQL @@ -5725,13 +5733,14 @@ class AsyncTable: -------- >>> import asyncio >>> import lancedb + >>> from lancedb.expr import col >>> import pandas as pd >>> async def demo_update(): ... data = pd.DataFrame({"x": [1, 2], "vector": [[1, 2], [3, 4]]}) ... db = await lancedb.connect_async("./.lancedb") ... table = await db.create_table("my_table", data) ... # x is [1, 2], vector is [[1, 2], [3, 4]] - ... await table.update({"vector": [10, 10]}, where="x = 2") + ... await table.update({"vector": [10, 10]}, where=col("x") == 2) ... # x is [1, 2], vector is [[1, 2], [10, 10]] ... await table.update(updates_sql={"x": "x + 1"}) ... # x is [2, 3], vector is [[1, 2], [10, 10]] @@ -5745,7 +5754,8 @@ class AsyncTable: if updates is not None: updates_sql = {k: value_to_sql(v) for k, v in updates.items()} - return await self._inner.update(updates_sql, where) + predicate = where.to_sql() if isinstance(where, Expr) else where + return await self._inner.update(updates_sql, predicate) async def add_columns( self, transforms: dict[str, str] | pa.field | List[pa.field] | pa.Schema diff --git a/python/python/tests/test_table.py b/python/python/tests/test_table.py index 069527b21..2d47606fd 100644 --- a/python/python/tests/test_table.py +++ b/python/python/tests/test_table.py @@ -309,6 +309,21 @@ async def test_update_async(mem_db_async: AsyncConnection): assert await table.count_rows("id == 10") == 1 +@pytest.mark.asyncio +async def test_update_expr_filter_literals_async(mem_db_async: AsyncConnection): + values = ["5", "4.66e-84", "it's"] + table = await mem_db_async.create_table( + "update_expr_literals", + data=[{"field": value, "result": "original"} for value in values], + ) + + for value in values: + update_res = await table.update({"result": value}, where=col("field") == value) + assert update_res.rows_updated == 1 + + assert (await table.to_arrow())["result"].to_pylist() == values + + def test_create_table(mem_db: DBConnection): schema = pa.schema( { @@ -2196,6 +2211,20 @@ def test_update(mem_db: DBConnection): assert np.allclose(v, np.array([[1.2, 1.9], [1.1, 1.1]])) +def test_update_expr_filter_literals(mem_db: DBConnection): + values = ["5", "4.66e-84", "it's"] + table = mem_db.create_table( + "update_expr_literals", + data=[{"field": value, "result": "original"} for value in values], + ) + + for value in values: + update_res = table.update(where=col("field") == value, values={"result": value}) + assert update_res.rows_updated == 1 + + assert table.to_arrow()["result"].to_pylist() == values + + def test_update_types(mem_db: DBConnection): table = mem_db.create_table( "my_table",