Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
# **Upcoming release**

- ...
- #886 Close connections opened by other threads in AutoImport.close() (@cristianchiriac)

# Release 1.15.0

Expand Down
32 changes: 26 additions & 6 deletions rope/contrib/autoimport/sqlite.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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,
)
Expand All @@ -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)
Expand Down Expand Up @@ -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."""
Expand Down
24 changes: 24 additions & 0 deletions ropetest/contrib/autoimport/autoimporttest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading