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
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
# **Upcoming release**

- #886 Track and close SQLite connections created in worker threads in AutoImport (@mcepl)
- ...

# Release 1.15.0
Expand Down
92 changes: 84 additions & 8 deletions rope/contrib/autoimport/sqlite.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -142,17 +142,21 @@ 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,
memory=memory,
)
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(
Expand Down Expand Up @@ -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:
Expand All @@ -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):
Expand Down Expand Up @@ -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."""
Expand Down
121 changes: 107 additions & 14 deletions ropetest/contrib/autoimport/autoimporttest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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):
Expand Down
4 changes: 3 additions & 1 deletion ropetest/contrib/autoimport/modeltest.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,9 @@

@pytest.fixture
def empty_db():
return sqlite3.connect(":memory:")
conn = sqlite3.connect(":memory:")
yield conn
conn.close()


class TestQuery:
Expand Down
14 changes: 8 additions & 6 deletions ropetest/contrib/autoimporttest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down Expand Up @@ -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):
Expand All @@ -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()

Expand Down
Loading