mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-18 12:08:35 +00:00
fix(python): accept expressions in update filters
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user