From 486e0cc1cd67e98c22d1f125403144cdedfa3d58 Mon Sep 17 00:00:00 2001 From: Martin Carpenter Date: Thu, 9 Feb 2023 00:59:53 +0100 Subject: [PATCH] Raise exception on registering identical function names --- sqlite_utils/db.py | 9 +++++++-- tests/test_register_function.py | 17 +++++++++++++---- 2 files changed, 20 insertions(+), 6 deletions(-) diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index caf3916..73b0b56 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -219,6 +219,11 @@ class AlterError(Exception): pass +class FunctionAlreadyRegistered(Exception): + "A function with this name and arity was already registered" + pass + + class NoObviousTable(Exception): "Could not tell which table this operation refers to" pass @@ -409,7 +414,7 @@ class Database: fn_name = name or fn.__name__ arity = len(inspect.signature(fn).parameters) 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 = {} registered = False if deterministic: @@ -434,7 +439,7 @@ class Database: def register_fts4_bm25(self): "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]): """ diff --git a/tests/test_register_function.py b/tests/test_register_function.py index 5169a67..05721ea 100644 --- a/tests/test_register_function.py +++ b/tests/test_register_function.py @@ -3,7 +3,7 @@ import pytest import sys from unittest.mock import MagicMock, call from sqlite_utils.utils import sqlite3 - +from sqlite_utils.db import FunctionAlreadyRegistered def test_register_function(fresh_db): @fresh_db.register_function @@ -85,9 +85,10 @@ def test_register_function_replace(fresh_db): assert "one" == fresh_db.execute("select one()").fetchone()[0] # This will fail to replace the function: - @fresh_db.register_function() - def one(): # noqa - return "two" + with pytest.raises(FunctionAlreadyRegistered): + @fresh_db.register_function() + def one(): # noqa + return "two" assert "one" == fresh_db.execute("select one()").fetchone()[0] @@ -97,3 +98,11 @@ def test_register_function_replace(fresh_db): return "two" 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)