diff --git a/CHANGELOG.md b/CHANGELOG.md index ff6f27ed8..c3fc78654 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,6 @@ # **Upcoming release** -- ... +- #886 Close connections opened by other threads in AutoImport.close() (@cristianchiriac) # Release 1.15.0 diff --git a/rope/contrib/autoimport/sqlite.py b/rope/contrib/autoimport/sqlite.py index 7ae8fb718..70381d537 100644 --- a/rope/contrib/autoimport/sqlite.py +++ b/rope/contrib/autoimport/sqlite.py @@ -15,9 +15,10 @@ from hashlib import sha256 from itertools import chain from pathlib import Path -from threading import local +from threading import Lock, Thread, current_thread, local from typing import ( TYPE_CHECKING, + Dict, Generator, Iterable, Iterator, @@ -143,6 +144,8 @@ def __init__( DeprecationWarning, ) self.thread_local = local() + self._connections: Dict[Thread, sqlite3.Connection] = {} + self._connections_lock = Lock() self.connection = self.create_database_connection( project=project, memory=memory, @@ -186,10 +189,15 @@ 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, + ) @property def connection(self) -> sqlite3.Connection: @@ -199,7 +207,7 @@ def connection(self) -> sqlite3.Connection: This makes sure AutoImport can be shared across threads. """ if not hasattr(self.thread_local, "connection"): - self.thread_local.connection = self.create_database_connection( + self.connection = self.create_database_connection( project=self.project, memory=self.memory, ) @@ -208,6 +216,14 @@ def connection(self) -> sqlite3.Connection: @connection.setter def connection(self, value: sqlite3.Connection): self.thread_local.connection = value + # Keep track of every thread's connection, so close() can close them + # all. Connections of threads that have finished are closed here. + with self._connections_lock: + for thread, connection in list(self._connections.items()): + if not thread.is_alive(): + connection.close() + del self._connections[thread] + self._connections[current_thread()] = value def _setup_db(self): models.Metadata.create_table(self.connection) @@ -458,9 +474,13 @@ def update_module(self, module: str): self.generate_modules_cache([module]) def close(self): - """Close the autoimport database.""" + """Close the autoimport database, including other threads' connections.""" self.connection.commit() - self.connection.close() + with self._connections_lock: + connections = list(self._connections.values()) + self._connections.clear() + for connection in connections: + connection.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..e9ef3c4ad 100644 --- a/ropetest/contrib/autoimport/autoimporttest.py +++ b/ropetest/contrib/autoimport/autoimporttest.py @@ -123,6 +123,30 @@ def foo(): assert [("from pkg1 import foo", "foo")] == results +def test_close_closes_connections_of_other_threads(project: Project): + autoimport = AutoImport(project, memory=True) + with ThreadPoolExecutor(1) as tp: + worker_connection = tp.submit(lambda: autoimport.connection).result() + + autoimport.close() + + with pytest.raises(sqlite3.ProgrammingError, match="closed database"): + worker_connection.execute("SELECT 1") + + +def test_connections_of_finished_threads_are_closed(project: Project): + with closing(AutoImport(project, memory=True)) as autoimport: + with ThreadPoolExecutor(1) as tp: + worker_connection = tp.submit(lambda: autoimport.connection).result() + + # The next new connection closes the ones left by finished threads + with ThreadPoolExecutor(1) as tp: + tp.submit(lambda: autoimport.connection).result() + + with pytest.raises(sqlite3.ProgrammingError, match="closed database"): + worker_connection.execute("SELECT 1") + + def test_connection(project: Project, project2: Project): ai1 = AutoImport(project) ai2 = AutoImport(project)