Better support check for deterministic=True, closes #425

Bug first discovered in #421
This commit is contained in:
Simon Willison 2022-04-13 15:31:37 -07:00
commit 6f3ae864f1
2 changed files with 39 additions and 28 deletions

View file

@ -392,13 +392,18 @@ class Database:
if not replace and (name, arity) in self._registered_functions:
return fn
kwargs = {}
if (
deterministic
and sys.version_info >= (3, 8)
and self.sqlite_version >= (3, 8, 3)
):
kwargs["deterministic"] = True
self.conn.create_function(name, arity, fn, **kwargs)
registered = False
if deterministic:
# Try this, but fall back if sqlite3.NotSupportedError
try:
self.conn.create_function(
name, arity, fn, **dict(kwargs, deterministic=True)
)
registered = True
except sqlite3.NotSupportedError:
pass
if not registered:
self.conn.create_function(name, arity, fn, **kwargs)
self._registered_functions.add((name, arity))
return fn

View file

@ -1,7 +1,8 @@
# flake8: noqa
import pytest
import sys
from unittest.mock import MagicMock
from unittest.mock import MagicMock, call
from sqlite_utils.utils import sqlite3
def test_register_function(fresh_db):
@ -31,36 +32,41 @@ def test_register_function_deterministic(fresh_db):
assert result == "bob"
@pytest.mark.skipif(
sys.version_info < (3, 8), reason="deterministic=True was added in Python 3.8"
)
@pytest.mark.parametrize(
"fake_sqlite_version,should_use_deterministic",
(
("3.36.0", True),
("3.8.3", True),
("3.8.2", False),
),
)
def test_register_function_deterministic_registered(
fresh_db, fake_sqlite_version, should_use_deterministic
):
def test_register_function_deterministic_tries_again_if_exception_raised(fresh_db):
fresh_db.conn = MagicMock()
fresh_db.conn.create_function = MagicMock()
fresh_db.conn.execute().fetchall.return_value = [(fake_sqlite_version,)]
@fresh_db.register_function(deterministic=True)
def to_lower_2(s):
return s.lower()
expected_kwargs = {}
if should_use_deterministic:
expected_kwargs = dict(deterministic=True)
fresh_db.conn.create_function.assert_called_with(
"to_lower_2", 1, to_lower_2, **expected_kwargs
"to_lower_2", 1, to_lower_2, deterministic=True
)
first = True
def side_effect(*args, **kwargs):
# Raise exception only first time this is called
nonlocal first
if first:
first = False
raise sqlite3.NotSupportedError()
# But if sqlite3.NotSupportedError is raised, it tries again
fresh_db.conn.create_function.reset_mock()
fresh_db.conn.create_function.side_effect = side_effect
@fresh_db.register_function(deterministic=True)
def to_lower_3(s):
return s.lower()
# Should have been called once with deterministic=True and once without
assert fresh_db.conn.create_function.call_args_list == [
call("to_lower_3", 1, to_lower_3, deterministic=True),
call("to_lower_3", 1, to_lower_3),
]
def test_register_function_replace(fresh_db):
@fresh_db.register_function()