diff --git a/CHANGELOG.md b/CHANGELOG.md index ff6f27ed8..bb5851055 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,6 @@ # **Upcoming release** +- #886 Track and close SQLite connections created in worker threads in AutoImport (@mcepl) - ... # Release 1.15.0 diff --git a/rope/contrib/autoimport/sqlite.py b/rope/contrib/autoimport/sqlite.py index 7ae8fb718..524f29f45 100644 --- a/rope/contrib/autoimport/sqlite.py +++ b/rope/contrib/autoimport/sqlite.py @@ -15,7 +15,7 @@ from hashlib import sha256 from itertools import chain from pathlib import Path -from threading import local +from threading import Lock, local from typing import ( TYPE_CHECKING, Generator, @@ -142,6 +142,10 @@ def __init__( "`AutoImport(memory=True)` explicitly.", DeprecationWarning, ) + self._closed = False + self._connections: Set[sqlite3.Connection] = set() + self._connections_lock = Lock() + self._observer: Optional[resourceobserver.ResourceObserver] = None self.thread_local = local() self.connection = self.create_database_connection( project=project, @@ -149,10 +153,10 @@ def __init__( ) self._setup_db() if observe: - observer = resourceobserver.ResourceObserver( + self._observer = resourceobserver.ResourceObserver( changed=self._changed, moved=self._moved, removed=self._removed ) - project.add_observer(observer) + project.add_observer(self._observer) @classmethod def create_database_connection( @@ -186,10 +190,25 @@ def calculate_project_hash(data: str) -> str: else: project_hash = calculate_project_hash(project.ropefolder.real_path) return sqlite3.connect( - f"file:rope-{project_hash}:?mode=memory&cache=shared", uri=True + f"file:rope-{project_hash}:?mode=memory&cache=shared", + uri=True, + check_same_thread=False, ) else: - return sqlite3.connect(project.ropefolder.pathlib / "autoimport.db") + return sqlite3.connect( + project.ropefolder.pathlib / "autoimport.db", + check_same_thread=False, + ) + + def _register_connection(self, conn: sqlite3.Connection) -> None: + with self._connections_lock: + if self._closed: + with contextlib.suppress( + sqlite3.ProgrammingError, sqlite3.OperationalError + ): + conn.close() + raise exceptions.RopeError("AutoImport instance has been closed") + self._connections.add(conn) @property def connection(self) -> sqlite3.Connection: @@ -198,15 +217,26 @@ def connection(self) -> sqlite3.Connection: This makes sure AutoImport can be shared across threads. """ + if self._closed: + raise exceptions.RopeError("AutoImport instance has been closed") if not hasattr(self.thread_local, "connection"): - self.thread_local.connection = self.create_database_connection( + conn = self.create_database_connection( project=self.project, memory=self.memory, ) + self._register_connection(conn) + self.thread_local.connection = conn return self.thread_local.connection @connection.setter def connection(self, value: sqlite3.Connection): + if self._closed: + raise exceptions.RopeError("AutoImport instance has been closed") + old_conn = getattr(self.thread_local, "connection", None) + if old_conn is not None and old_conn is not value: + with self._connections_lock: + self._connections.discard(old_conn) + self._register_connection(value) self.thread_local.connection = value def _setup_db(self): @@ -457,10 +487,56 @@ def update_module(self, module: str): self._del_if_exist(module) self.generate_modules_cache([module]) + def close_thread_connection(self): + """Close the SQLite connection for the current thread.""" + conn = getattr(self.thread_local, "connection", None) + if conn is not None: + with self._connections_lock: + self._connections.discard(conn) + with contextlib.suppress(AttributeError): + del self.thread_local.connection + with contextlib.suppress( + sqlite3.ProgrammingError, sqlite3.OperationalError + ): + conn.commit() + with contextlib.suppress( + sqlite3.ProgrammingError, sqlite3.OperationalError + ): + conn.close() + def close(self): """Close the autoimport database.""" - self.connection.commit() - self.connection.close() + with self._connections_lock: + if self._closed: + return + self._closed = True + connections = list(self._connections) + self._connections.clear() + + if self._observer is not None: + with contextlib.suppress(Exception): + self.project.remove_observer(self._observer) + self._observer = None + + for conn in connections: + with contextlib.suppress( + sqlite3.ProgrammingError, sqlite3.OperationalError + ): + conn.commit() + with contextlib.suppress( + sqlite3.ProgrammingError, sqlite3.OperationalError + ): + conn.close() + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + self.close() + + def __del__(self): + with contextlib.suppress(Exception): + self.close() def get_name_locations(self, name): """Return a list of ``(resource, lineno)`` tuples.""" diff --git a/ropetest/contrib/autoimport/autoimporttest.py b/ropetest/contrib/autoimport/autoimporttest.py index 072ab1ed8..b813af3a0 100644 --- a/ropetest/contrib/autoimport/autoimporttest.py +++ b/ropetest/contrib/autoimport/autoimporttest.py @@ -6,6 +6,7 @@ import pytest +from rope.base import exceptions from rope.base.project import Project from rope.base.resources import File, Folder from rope.contrib.autoimport import models @@ -46,16 +47,16 @@ def test_autoimport_connection_parameter_with_in_memory( project: Project, autoimport: AutoImport, ): - connection = AutoImport.create_database_connection(memory=True) - assert is_in_memory_database(connection) + with closing(AutoImport.create_database_connection(memory=True)) as connection: + assert is_in_memory_database(connection) def test_autoimport_connection_parameter_with_project( project: Project, autoimport: AutoImport, ): - connection = AutoImport.create_database_connection(project=project) - assert not is_in_memory_database(connection) + with closing(AutoImport.create_database_connection(project=project)) as connection: + assert not is_in_memory_database(connection) def test_autoimport_create_database_connection_conflicting_parameter( @@ -102,25 +103,117 @@ def foo(): def test_multithreading( - autoimport: AutoImport, project: Project, pkg1: Folder, mod1: File, ): mod1_init = pkg1.get_child("__init__.py") - mod1_init.write(dedent("""\ + mod1_init.write( + dedent("""\ def foo(): pass - """)) - mod1.write(dedent("""\ + """) + ) + mod1.write( + dedent("""\ foo - """)) - autoimport = AutoImport(project, memory=False) - autoimport.generate_cache([mod1_init]) + """) + ) + with closing(AutoImport(project, memory=False)) as autoimport: + autoimport.generate_cache([mod1_init]) + + with ThreadPoolExecutor(1) as tp: + results = tp.submit(autoimport.search, "foo", True).result() + assert [("from pkg1 import foo", "foo")] == results + + +def test_multithread_connections_closed_on_close(project: Project): + with AutoImport(project, memory=True) as ai: + main_conn = ai.connection + worker_conns = [] + + def worker(): + conn = ai.connection + worker_conns.append(conn) + return list(ai.search("foo")) + + with ThreadPoolExecutor(3) as tp: + futures = [tp.submit(worker) for _ in range(3)] + for f in futures: + f.result() + + all_conns = {main_conn} | set(worker_conns) + assert len(all_conns) > 1 + assert all_conns.issubset(ai._connections) + + assert ai._closed + assert len(ai._connections) == 0 + for conn in all_conns: + with pytest.raises(sqlite3.ProgrammingError, match="Cannot operate on a closed database"): + conn.execute("SELECT 1") + + with pytest.raises(exceptions.RopeError, match="AutoImport instance has been closed"): + _ = ai.connection + + +def test_close_thread_connection(project: Project): + with AutoImport(project, memory=True) as ai: + worker_conn = None + + def worker(): + nonlocal worker_conn + worker_conn = ai.connection + assert worker_conn in ai._connections + ai.close_thread_connection() + assert worker_conn not in ai._connections + with pytest.raises(sqlite3.ProgrammingError, match="Cannot operate on a closed database"): + worker_conn.execute("SELECT 1") + new_conn = ai.connection + assert new_conn is not worker_conn + assert new_conn in ai._connections + + with ThreadPoolExecutor(1) as tp: + tp.submit(worker).result() + + +def test_close_idempotent(project: Project): + ai = AutoImport(project, memory=True) + conn = ai.connection + ai.close() + assert ai._closed + ai.close() + with pytest.raises(sqlite3.ProgrammingError, match="Cannot operate on a closed database"): + conn.execute("SELECT 1") - tp = ThreadPoolExecutor(1) - results = tp.submit(autoimport.search, "foo", True).result() - assert [("from pkg1 import foo", "foo")] == results + +def test_register_connection_after_close(project: Project): + ai = AutoImport(project, memory=True) + ai.close() + conn = AutoImport.create_database_connection(memory=True) + with pytest.raises(exceptions.RopeError, match="AutoImport instance has been closed"): + ai._register_connection(conn) + assert conn not in ai._connections + with pytest.raises(sqlite3.ProgrammingError, match="Cannot operate on a closed database"): + conn.execute("SELECT 1") + + +def test_connection_setter_after_close(project: Project): + ai = AutoImport(project, memory=True) + ai.close() + with closing(AutoImport.create_database_connection(memory=True)) as conn: + with pytest.raises(exceptions.RopeError, match="AutoImport instance has been closed"): + ai.connection = conn + + +def test_connection_setter_replaces_existing(project: Project): + with AutoImport(project, memory=True) as ai: + old_conn = ai.connection + assert old_conn in ai._connections + with closing(AutoImport.create_database_connection(memory=True)) as conn: + ai.connection = conn + assert ai.connection is conn + assert conn in ai._connections + assert old_conn not in ai._connections def test_connection(project: Project, project2: Project): diff --git a/ropetest/contrib/autoimport/modeltest.py b/ropetest/contrib/autoimport/modeltest.py index 0044434cc..8eef06d35 100644 --- a/ropetest/contrib/autoimport/modeltest.py +++ b/ropetest/contrib/autoimport/modeltest.py @@ -8,7 +8,9 @@ @pytest.fixture def empty_db(): - return sqlite3.connect(":memory:") + conn = sqlite3.connect(":memory:") + yield conn + conn.close() class TestQuery: diff --git a/ropetest/contrib/autoimporttest.py b/ropetest/contrib/autoimporttest.py index d42ec14f7..ebf21af56 100644 --- a/ropetest/contrib/autoimporttest.py +++ b/ropetest/contrib/autoimporttest.py @@ -15,6 +15,7 @@ def setUp(self): self.importer = autoimport.AutoImport(self.project, observe=False) def tearDown(self): + self.importer.close() testutils.remove_project(self.project) super().tearDown() @@ -178,12 +179,12 @@ def test_skipping_directories_not_accessible_because_of_permission_error(self): def test_search_submodule(project, external_fixturepkg): - importer = autoimport.AutoImport(project, observe=False) - importer.update_module("external_fixturepkg") - import_statement = ("from external_fixturepkg import mod1", "mod1") - assert import_statement in importer.search("mod1", exact_match=True) - assert import_statement in importer.search("mo") - assert import_statement in importer.search("mod1") + with autoimport.AutoImport(project, observe=False) as importer: + importer.update_module("external_fixturepkg") + import_statement = ("from external_fixturepkg import mod1", "mod1") + assert import_statement in importer.search("mod1", exact_match=True) + assert import_statement in importer.search("mo") + assert import_statement in importer.search("mod1") class AutoImportObservingTest(unittest.TestCase): @@ -196,6 +197,7 @@ def setUp(self): self.importer = autoimport.AutoImport(self.project, observe=True) def tearDown(self): + self.importer.close() testutils.remove_project(self.project) super().tearDown()