diff --git a/distributed/client.py b/distributed/client.py index 5f14cb4911..125455011a 100644 --- a/distributed/client.py +++ b/distributed/client.py @@ -5217,6 +5217,38 @@ def unregister_scheduler_plugin(self, name: str): """ return self.sync(self.scheduler.unregister_scheduler_plugin, name=name) + def has_scheduler_plugin(self, name: str) -> bool: + """Check if a scheduler plugin is registered. + + Parameters + ---------- + name : str + Name of the plugin to check. + + Returns + ------- + bool + True if a plugin with the given name is registered, False otherwise. + + Examples + -------- + >>> class MyPlugin(SchedulerPlugin): + ... pass + + >>> plugin = MyPlugin() + >>> client.register_plugin(plugin, name="my-plugin") + >>> client.has_scheduler_plugin("my-plugin") + True + >>> client.has_scheduler_plugin("nonexistent") + False + + See Also + -------- + register_scheduler_plugin + unregister_scheduler_plugin + """ + return self.sync(self.scheduler.has_scheduler_plugin, name=name) + def register_worker_callbacks(self, setup=None): """ Registers a setup callback function for all current and future workers. diff --git a/distributed/diagnostics/tests/test_scheduler_plugin.py b/distributed/diagnostics/tests/test_scheduler_plugin.py index 9f312df48e..5f6ae15472 100644 --- a/distributed/diagnostics/tests/test_scheduler_plugin.py +++ b/distributed/diagnostics/tests/test_scheduler_plugin.py @@ -568,3 +568,41 @@ def __init__(self, instance=None): await s.register_scheduler_plugin(plugin=dumps(second), idempotent=False) assert "nonidempotentplugin" in s.plugins assert s.plugins["nonidempotentplugin"].instance == "second" + + +@gen_cluster(client=True) +async def test_has_scheduler_plugin(c, s, a, b): + """Test has_scheduler_plugin method on Client and Scheduler.""" + + class MyPlugin(SchedulerPlugin): + name = "test-plugin" + + # Initially, plugin should not be registered + assert not s.has_scheduler_plugin("test-plugin") + assert not await c.scheduler.has_scheduler_plugin(name="test-plugin") + + # Register the plugin + plugin = MyPlugin() + await c._register_scheduler_plugin( + plugin=plugin, name="test-plugin", idempotent=False + ) + + # Now it should be registered + assert s.has_scheduler_plugin("test-plugin") + assert await c.scheduler.has_scheduler_plugin(name="test-plugin") + + # Check with explicit name + await c._register_scheduler_plugin( + plugin=MyPlugin(), name="another-plugin", idempotent=False + ) + assert s.has_scheduler_plugin("another-plugin") + assert await c.scheduler.has_scheduler_plugin(name="another-plugin") + + # Unregister and verify + s.remove_plugin("test-plugin") + assert not s.has_scheduler_plugin("test-plugin") + assert not await c.scheduler.has_scheduler_plugin(name="test-plugin") + + # Non-existent plugin should return False + assert not s.has_scheduler_plugin("nonexistent") + assert not await c.scheduler.has_scheduler_plugin(name="nonexistent") diff --git a/distributed/scheduler.py b/distributed/scheduler.py index 92f22b807a..9652c8fe66 100644 --- a/distributed/scheduler.py +++ b/distributed/scheduler.py @@ -4185,6 +4185,7 @@ async def post(self) -> None: "get_task_prefix_states": self.get_task_prefix_states, "register_scheduler_plugin": self.register_scheduler_plugin, "unregister_scheduler_plugin": self.unregister_scheduler_plugin, + "has_scheduler_plugin": self.has_scheduler_plugin, "register_worker_plugin": self.register_worker_plugin, "unregister_worker_plugin": self.unregister_worker_plugin, "register_nanny_plugin": self.register_nanny_plugin, @@ -6268,6 +6269,31 @@ async def unregister_scheduler_plugin(self, name: str) -> None: """Unregister a plugin on the scheduler.""" self.remove_plugin(name) + def has_scheduler_plugin(self, name: str) -> bool: + """Check if a scheduler plugin is registered. + + Parameters + ---------- + name : str + Name of the plugin to check. + + Returns + ------- + bool + True if a plugin with the given name is registered, False otherwise. + + Examples + -------- + >>> s.has_scheduler_plugin("my-plugin") # doctest: +SKIP + True + + See Also + -------- + add_plugin + remove_plugin + """ + return name in self.plugins + def worker_send(self, worker: str, msg: dict[str, Any]) -> None: """Send message to worker