Support repeated calls to Table.convert()

* Test repeated calls to Table.convert()
* Register Table.convert() functions under their own `lambda_hash` name
* Raise exception on registering identical function names

Refs #525
This commit is contained in:
Martin Carpenter 2023-05-08 23:53:58 +02:00 • committed by GitHub
commit 02f5c4d69d
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 35 additions and 9 deletions

View file

@ -219,6 +219,11 @@ class AlterError(Exception):
pass pass
class FunctionAlreadyRegistered(Exception):
"A function with this name and arity was already registered"
pass
class NoObviousTable(Exception): class NoObviousTable(Exception):
"Could not tell which table this operation refers to" "Could not tell which table this operation refers to"
pass pass
@ -409,7 +414,7 @@ class Database:
fn_name = name or fn.__name__ fn_name = name or fn.__name__
arity = len(inspect.signature(fn).parameters) arity = len(inspect.signature(fn).parameters)
if not replace and (fn_name, arity) in self._registered_functions: if not replace and (fn_name, arity) in self._registered_functions:
return fn raise FunctionAlreadyRegistered(f'Already registered function with name "{fn_name}" and identical arity')
kwargs = {} kwargs = {}
registered = False registered = False
if deterministic: if deterministic:
@ -434,7 +439,7 @@ class Database:
def register_fts4_bm25(self): def register_fts4_bm25(self):
"Register the ``rank_bm25(match_info)`` function used for calculating relevance with SQLite FTS4." "Register the ``rank_bm25(match_info)`` function used for calculating relevance with SQLite FTS4."
self.register_function(rank_bm25, deterministic=True) self.register_function(rank_bm25, deterministic=True, replace=True)
def attach(self, alias: str, filepath: Union[str, pathlib.Path]): def attach(self, alias: str, filepath: Union[str, pathlib.Path]):
""" """
@ -2687,13 +2692,16 @@ class Table(Queryable):
return v return v
return jsonify_if_needed(fn(v)) return jsonify_if_needed(fn(v))
self.db.register_function(convert_value) fn_name = fn.__name__
if fn_name == '<lambda>':
fn_name = f'lambda_{hash(fn)}'
self.db.register_function(convert_value, name=fn_name)
sql = "update [{table}] set {sets}{where};".format( sql = "update [{table}] set {sets}{where};".format(
table=self.name, table=self.name,
sets=", ".join( sets=", ".join(
[ [
"[{output_column}] = convert_value([{column}])".format( "[{output_column}] = {fn_name}([{column}])".format(
output_column=output or column, column=column output_column=output or column, column=column, fn_name=fn_name
) )
for column in columns for column in columns
] ]

View file

@ -147,3 +147,12 @@ def test_convert_multi_exception(fresh_db):
table.insert({"title": "Mixed Case"}) table.insert({"title": "Mixed Case"})
with pytest.raises(BadMultiValues): with pytest.raises(BadMultiValues):
table.convert("title", lambda v: v.upper(), multi=True) table.convert("title", lambda v: v.upper(), multi=True)
def test_convert_repeated(fresh_db):
table = fresh_db["table"]
col = "num"
table.insert({col: 1})
table.convert(col, lambda x: x*2)
table.convert(col, lambda _x: 0)
assert table.get(1) == {col: 0}

View file

@ -3,7 +3,7 @@ import pytest
import sys import sys
from unittest.mock import MagicMock, call from unittest.mock import MagicMock, call
from sqlite_utils.utils import sqlite3 from sqlite_utils.utils import sqlite3
from sqlite_utils.db import FunctionAlreadyRegistered
def test_register_function(fresh_db): def test_register_function(fresh_db):
@fresh_db.register_function @fresh_db.register_function
@ -85,9 +85,10 @@ def test_register_function_replace(fresh_db):
assert "one" == fresh_db.execute("select one()").fetchone()[0] assert "one" == fresh_db.execute("select one()").fetchone()[0]
# This will fail to replace the function: # This will fail to replace the function:
@fresh_db.register_function() with pytest.raises(FunctionAlreadyRegistered):
def one(): # noqa @fresh_db.register_function()
return "two" def one(): # noqa
return "two"
assert "one" == fresh_db.execute("select one()").fetchone()[0] assert "one" == fresh_db.execute("select one()").fetchone()[0]
@ -97,3 +98,11 @@ def test_register_function_replace(fresh_db):
return "two" return "two"
assert "two" == fresh_db.execute("select one()").fetchone()[0] assert "two" == fresh_db.execute("select one()").fetchone()[0]
def test_register_function_duplicate(fresh_db):
def to_lower(s):
return s.lower()
fresh_db.register_function(to_lower)
with pytest.raises(FunctionAlreadyRegistered):
fresh_db.register_function(to_lower)