diff --git a/sqlite_utils/cli.py b/sqlite_utils/cli.py index 2f6268f..f15850d 100644 --- a/sqlite_utils/cli.py +++ b/sqlite_utils/cli.py @@ -16,7 +16,7 @@ from sqlite_utils.db import ( NoTable, quote_identifier, ) -from sqlite_utils.plugins import pm, get_plugins +from sqlite_utils.plugins import ensure_plugins_loaded, pm, get_plugins from sqlite_utils.utils import maximize_csv_field_size_limit from sqlite_utils import recipes import textwrap @@ -3412,6 +3412,7 @@ def plugins_list(): click.echo(json.dumps(get_plugins(), indent=2)) +ensure_plugins_loaded() pm.hook.register_commands(cli=cli) cli.add_command(migrate) diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index d447cd2..2005210 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -41,7 +41,7 @@ from typing import ( Tuple, ) import uuid -from sqlite_utils.plugins import pm +from sqlite_utils.plugins import ensure_plugins_loaded, pm try: iterdump = importlib.import_module("sqlite_dump").iterdump @@ -401,6 +401,7 @@ class Database: self._registered_functions: set = set() self.use_counts_table = use_counts_table if execute_plugins: + ensure_plugins_loaded() pm.hook.prepare_connection(conn=self.conn) self.strict = strict diff --git a/sqlite_utils/plugins.py b/sqlite_utils/plugins.py index 457a907..0aff7ff 100644 --- a/sqlite_utils/plugins.py +++ b/sqlite_utils/plugins.py @@ -6,13 +6,19 @@ from . import hookspecs pm: pluggy.PluginManager = pluggy.PluginManager("sqlite_utils") pm.add_hookspecs(hookspecs) +_plugins_loaded = False -if not getattr(sys, "_called_from_test", False): - # Only load plugins if not running tests + +def ensure_plugins_loaded() -> None: + global _plugins_loaded + if _plugins_loaded or getattr(sys, "_called_from_test", False): + return pm.load_setuptools_entrypoints("sqlite_utils") + _plugins_loaded = True def get_plugins() -> List[Dict[str, Union[str, List[str]]]]: + ensure_plugins_loaded() plugins: List[Dict[str, Union[str, List[str]]]] = [] plugin_to_distinfo = dict(pm.list_plugin_distinfo()) for plugin in pm.get_plugins(): diff --git a/tests/test_plugins.py b/tests/test_plugins.py index 1d459c9..c793e32 100644 --- a/tests/test_plugins.py +++ b/tests/test_plugins.py @@ -2,6 +2,7 @@ from click.testing import CliRunner import click import importlib import pytest +import sys from sqlite_utils import cli, Database, hookimpl, plugins @@ -16,6 +17,36 @@ def _supports_pragma_function_list(): db.close() +def test_get_plugins_loads_setuptools_entrypoints_once(monkeypatch): + calls = [] + monkeypatch.delattr(sys, "_called_from_test", raising=False) + monkeypatch.setattr(plugins, "_plugins_loaded", False) + monkeypatch.setattr( + plugins.pm, + "load_setuptools_entrypoints", + lambda group: calls.append(group) or 0, + ) + + plugins.get_plugins() + plugins.get_plugins() + + assert calls == ["sqlite_utils"] + + +def test_get_plugins_does_not_load_setuptools_entrypoints_in_tests(monkeypatch): + calls = [] + monkeypatch.setattr(sys, "_called_from_test", True, raising=False) + monkeypatch.setattr(plugins, "_plugins_loaded", False) + monkeypatch.setattr( + plugins.pm, + "load_setuptools_entrypoints", + lambda group: calls.append(group) or 0, + ) + + assert plugins.get_plugins() == [] + assert calls == [] + + def test_register_commands(): importlib.reload(cli) assert plugins.get_plugins() == []