From 6e5eada9441285f309aa4b71e0ba9ca10a2c81d7 Mon Sep 17 00:00:00 2001 From: mvdbeek Date: Mon, 20 Apr 2026 14:35:21 +0200 Subject: [PATCH] Add SSE entry-point channel, dispatch observability, declare-queue cache MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three related additions to the SSE notification pipeline, bundled here because they share ``SSEEventDispatcher._send`` as their modification point: 1. Interactive-tool entry-point SSE channel - ``entry_point_update`` dispatcher method + queue-worker handler. - ``InteractiveToolManager.configure_entry_points`` dispatches a wake-up event after the DB commit; the client refetches ``/api/entry_points`` on receipt (no payload). - Frontend: new SSE event type, store subscription with XOR polling fallback via ``enable_sse_entry_point_updates`` config flag. - Integration + selenium tests. 2. Queue and SSE observability metrics - Counters, timers, and periodic gauges for SSE dispatch, control-queue task execution, control-queue depth (via kombu passive declare), active SSE connections, dropped events, and active WorkerProcess rows. Flow through the existing ``galaxy_statsd_client`` — no new infra. Gauges are scheduled via Celery beat at ``queue_metrics_interval`` seconds (default 15). All instrumentation no-ops when statsd isn't configured. - Sub-emitter failures log once at WARNING and bump a ``galaxy.queue_metrics.error`` counter tagged by emitter so broken emitters are visible in metrics without log spam. 3. Active-worker control-queue cache - 30 s TTL cache on ``all_control_queues_for_declare`` with RLock stampede protection. At 1000+ events/s this eliminates ~30 DB round-trips/s per webapp for data that only changes on 60 s heartbeat cadence. Empty results are not cached — would otherwise silently drop every SSE event during the startup window. --- client/src/composables/useNotificationSSE.ts | 1 + client/src/stores/_testing/sseStoreSupport.ts | 8 +- client/src/stores/entryPointStore.test.js | 15 +- client/src/stores/entryPointStore.ts | 97 +++++-- doc/source/admin/galaxy_options.rst | 26 ++ lib/galaxy/app/__init__.py | 8 +- lib/galaxy/app_unittest_utils/galaxy_mock.py | 3 +- lib/galaxy/authnz/psa_authnz.py | 1 + lib/galaxy/celery/__init__.py | 3 + lib/galaxy/celery/tasks.py | 13 +- lib/galaxy/config/sample/galaxy.yml.sample | 12 + lib/galaxy/config/schemas/config_schema.yml | 18 ++ lib/galaxy/managers/interactivetool.py | 25 +- lib/galaxy/managers/notification.py | 2 +- lib/galaxy/managers/sse.py | 23 +- lib/galaxy/managers/sse_dispatch.py | 75 +++++- lib/galaxy/model/unittest_utils/data_app.py | 6 + lib/galaxy/queue_worker/__init__.py | 92 ++++++- lib/galaxy/structured_app/__init__.py | 1 + lib/galaxy/webapps/galaxy/api/tool_data.py | 2 +- lib/galaxy/webapps/galaxy/metrics/__init__.py | 0 .../webapps/galaxy/metrics/queue_metrics.py | 147 +++++++++++ lib/galaxy/webapps/galaxy/services/base.py | 1 - .../webapps/galaxy/services/datasets.py | 6 +- .../webapps/galaxy/services/histories.py | 2 +- .../galaxy/services/history_contents.py | 2 +- .../webapps/galaxy/services/invocations.py | 2 +- lib/galaxy/webapps/galaxy/services/jobs.py | 2 +- lib/galaxy/webapps/galaxy/services/pages.py | 2 +- lib/galaxy/webapps/galaxy/services/users.py | 6 +- test/integration/test_entry_point_sse.py | 158 ++++++++++++ .../test_entry_point_sse.py | 107 ++++++++ .../test_notification_sse.py | 4 +- test/unit/app/managers/test_queue_metrics.py | 243 ++++++++++++++++++ test/unit/app/managers/test_sse_dispatch.py | 226 ++++++++++++++++ .../app/managers/test_sse_dispatch_cache.py | 140 ++++++++++ .../app/queue_worker/test_queue_worker.py | 111 ++++++++ 37 files changed, 1528 insertions(+), 62 deletions(-) create mode 100644 lib/galaxy/webapps/galaxy/metrics/__init__.py create mode 100644 lib/galaxy/webapps/galaxy/metrics/queue_metrics.py create mode 100644 test/integration/test_entry_point_sse.py create mode 100644 test/integration_selenium/test_entry_point_sse.py create mode 100644 test/unit/app/managers/test_queue_metrics.py create mode 100644 test/unit/app/managers/test_sse_dispatch.py create mode 100644 test/unit/app/managers/test_sse_dispatch_cache.py diff --git a/client/src/composables/useNotificationSSE.ts b/client/src/composables/useNotificationSSE.ts index b0d34d8e8a5..944af5fb8d6 100644 --- a/client/src/composables/useNotificationSSE.ts +++ b/client/src/composables/useNotificationSSE.ts @@ -10,6 +10,7 @@ export const SSE_EVENT_TYPES = [ "broadcast_update", "notification_status", "history_update", + "entry_point_update", ] as const; export type SSEEventType = (typeof SSE_EVENT_TYPES)[number]; diff --git a/client/src/stores/_testing/sseStoreSupport.ts b/client/src/stores/_testing/sseStoreSupport.ts index 4613430bb25..b3e5d97fcb8 100644 --- a/client/src/stores/_testing/sseStoreSupport.ts +++ b/client/src/stores/_testing/sseStoreSupport.ts @@ -20,14 +20,20 @@ export interface SSEMockState { onEvent: ((event: MessageEvent) => void) | null; connect: ReturnType; disconnect: ReturnType; + connected?: Ref; } /** Build the factory used with ``vi.mock("@/composables/useNotificationSSE", ...)``. */ export function sseMockFactory(state: SSEMockState) { + // Lazily initialize ``connected`` so existing callers that don't pass it + // still get a working ref. + if (!state.connected) { + state.connected = ref(false); + } return { useSSE: vi.fn((onEvent: (event: MessageEvent) => void) => { state.onEvent = onEvent; - return { connect: state.connect, disconnect: state.disconnect }; + return { connect: state.connect, disconnect: state.disconnect, connected: state.connected }; }), }; } diff --git a/client/src/stores/entryPointStore.test.js b/client/src/stores/entryPointStore.test.js index d0c6b69b457..fdae2799d30 100644 --- a/client/src/stores/entryPointStore.test.js +++ b/client/src/stores/entryPointStore.test.js @@ -1,12 +1,25 @@ import flushPromises from "flush-promises"; import { createPinia, setActivePinia } from "pinia"; -import { beforeEach, describe, expect, it } from "vitest"; +import { beforeEach, describe, expect, it, vi } from "vitest"; import { HttpResponse, useServerMock } from "@/api/client/__mocks__"; import testInteractiveToolsResponse from "../components/InteractiveTools/testData/testInteractiveToolsResponse"; +import { sseMockFactory } from "./_testing/sseStoreSupport"; import { useEntryPointStore } from "./entryPointStore"; +// ``vi.mock`` is hoisted above module-level declarations, so the capture-state +// has to be built via ``vi.hoisted`` to be visible to the factory. Prevents +// these tests from opening a real EventSource against ``/api/events/stream`` +// when ``useEntryPointStore()`` is invoked. +const sseState = vi.hoisted(() => ({ + onEvent: null, + connect: vi.fn(), + disconnect: vi.fn(), + connected: null, +})); +vi.mock("@/composables/useNotificationSSE", () => sseMockFactory(sseState)); + const { server, http } = useServerMock(); describe("stores/EntryPointStore", () => { diff --git a/client/src/stores/entryPointStore.ts b/client/src/stores/entryPointStore.ts index 805124a713d..68193406967 100644 --- a/client/src/stores/entryPointStore.ts +++ b/client/src/stores/entryPointStore.ts @@ -1,10 +1,12 @@ import axios from "axios"; import isEqual from "lodash.isequal"; import { defineStore } from "pinia"; -import { computed, ref } from "vue"; +import { computed, ref, watch } from "vue"; import { useResourceWatcher } from "@/composables/resourceWatcher"; +import { useSSE } from "@/composables/useNotificationSSE"; import { getAppRoot } from "@/onload/loadConfig"; +import { useConfigStore } from "@/stores/configurationStore"; import { rethrowSimple } from "@/utils/simple-error"; const ACTIVE_POLLING_INTERVAL = 10000; @@ -23,23 +25,8 @@ interface EntryPoint { } export const useEntryPointStore = defineStore("entryPointStore", () => { - const { startWatchingResource: startWatchingEntryPoints, stopWatchingResource: stopWatchingEntryPoints } = - useResourceWatcher(fetchEntryPoints, { - shortPollingInterval: ACTIVE_POLLING_INTERVAL, - enableBackgroundPolling: false, // No need to poll in the background - }); - const entryPoints = ref([]); - const entryPointsForJob = computed(() => { - return (jobId: string) => entryPoints.value.filter((entryPoint) => entryPoint["job_id"] === jobId); - }); - - const entryPointsForHda = computed(() => { - return (hdaId: string) => - entryPoints.value.filter((entryPoint) => entryPoint["output_datasets_ids"].includes(hdaId)); - }); - async function fetchEntryPoints() { const url = `${getAppRoot()}api/entry_points`; const params = { running: true }; @@ -51,6 +38,79 @@ export const useEntryPointStore = defineStore("entryPointStore", () => { } } + // SSE-driven path: on each entry_point_update signal, refetch the canonical + // list from REST. The event carries no data — it's a pure wake-up. + function handleEntryPointSSEEvent(_event: MessageEvent) { + fetchEntryPoints().catch((err) => console.error("Error refreshing entry points from SSE push:", err)); + } + const { + connect: sseConnect, + disconnect: sseDisconnect, + connected: sseConnected, + } = useSSE(handleEntryPointSSEEvent, ["entry_point_update"]); + + let watchingInitialized = false; + let stopWatchingEntryPointsResource: (() => void) | null = null; + + // Callers opt in via ``startWatchingEntryPoints()`` (App.vue gates this on + // ``interactivetools_enable``). We then pick SSE or polling based on the + // server flag — mutually exclusive, mirroring historyStore / notificationsStore. + // ``useConfigStore`` is resolved lazily here so tests that only exercise + // the data methods don't need a ``/api/configuration`` handler registered. + function startWatchingEntryPoints() { + if (watchingInitialized) { + return; + } + watchingInitialized = true; + const configStore = useConfigStore(); + + const decide = () => { + if (configStore.config?.enable_sse_entry_point_updates) { + // Baseline fetch + SSE. Reconnect-refetch closes the "user + // navigated away and missed events" window. + fetchEntryPoints().catch((err) => console.warn("Initial entry-point load failed", err)); + sseConnect(); + watch(sseConnected, (isConnected, wasConnected) => { + if (isConnected && !wasConnected) { + fetchEntryPoints().catch((err) => + console.error("Error refreshing entry points on SSE reconnect:", err), + ); + } + }); + } else { + const { startWatchingResource, stopWatchingResource } = useResourceWatcher(fetchEntryPoints, { + shortPollingInterval: ACTIVE_POLLING_INTERVAL, + enableBackgroundPolling: false, + }); + stopWatchingEntryPointsResource = stopWatchingResource; + startWatchingResource(); + } + }; + + if (configStore.isLoaded) { + decide(); + } else { + const stop = watch( + () => configStore.isLoaded, + (loaded) => { + if (loaded) { + stop(); + decide(); + } + }, + ); + } + } + + const entryPointsForJob = computed(() => { + return (jobId: string) => entryPoints.value.filter((entryPoint) => entryPoint["job_id"] === jobId); + }); + + const entryPointsForHda = computed(() => { + return (hdaId: string) => + entryPoints.value.filter((entryPoint) => entryPoint["output_datasets_ids"].includes(hdaId)); + }); + function updateEntryPoints(data: EntryPoint[]) { let hasChanged = entryPoints.value.length !== data.length ? true : false; if (entryPoints.value.length === 0) { @@ -76,6 +136,11 @@ export const useEntryPointStore = defineStore("entryPointStore", () => { return { ...original, ...updated }; } + function stopWatchingEntryPoints() { + sseDisconnect(); + stopWatchingEntryPointsResource?.(); + } + function removeEntryPoint(toolId: string) { const index = entryPoints.value.findIndex((ep) => { return ep.id === toolId ? true : false; diff --git a/doc/source/admin/galaxy_options.rst b/doc/source/admin/galaxy_options.rst index da52bf02d82..e87d6ac0ef4 100644 --- a/doc/source/admin/galaxy_options.rst +++ b/doc/source/admin/galaxy_options.rst @@ -3408,6 +3408,18 @@ :Type: bool +~~~~~~~~~~~~~~~~~~~~~~~~~~ +``queue_metrics_interval`` +~~~~~~~~~~~~~~~~~~~~~~~~~~ + +:Description: + How often (in seconds) the Celery beat task emits queue-depth, + SSE-connection, and WorkerProcess gauges. Only active when + statsd_host is set. Set to 0 to disable. +:Default: ``15`` +:Type: int + + ~~~~~~~~~~~~~~~~~~~~~~ ``library_import_dir`` ~~~~~~~~~~~~~~~~~~~~~~ @@ -5818,6 +5830,20 @@ :Type: bool +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +``enable_sse_entry_point_updates`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +:Description: + Enables real-time interactive-tool entry-point update + notifications via Server-Sent Events. When enabled, the client + subscribes to entry_point_update SSE events and refetches the + entry-point list on each event, replacing the 10-second polling + loop. When disabled, polling remains the source of updates. +:Default: ``false`` +:Type: bool + + ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ ``history_audit_monitor_poll_interval`` ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ diff --git a/lib/galaxy/app/__init__.py b/lib/galaxy/app/__init__.py index 991944962cf..28ed51e4c3a 100644 --- a/lib/galaxy/app/__init__.py +++ b/lib/galaxy/app/__init__.py @@ -695,6 +695,7 @@ class GalaxyManagerApplication(MinimalManagerApp, MinimalGalaxyApplication): SSEEventDispatcher( queue_worker=getattr(self, "queue_worker", None), application_stack=self.application_stack, + statsd_client=self.execution_timer_factory.galaxy_statsd_client, ), ) self.notification_manager = self._register_singleton(NotificationManager) @@ -858,7 +859,12 @@ class UniverseApplication(StructuredApp, GalaxyManagerApplication, InstallationT # SSE connection manager for real-time notification push. # Consumed via ``depends(SSEConnectionManager)`` / ``app[SSEConnectionManager]``, # so no module-level attribute is needed — keep the container wiring only. - self._register_singleton(SSEConnectionManager) + self._register_singleton( + SSEConnectionManager, + SSEConnectionManager( + statsd_client=self.execution_timer_factory.galaxy_statsd_client, + ), + ) # AI agent registry and service agent_registry = build_agent_registry(self.config) diff --git a/lib/galaxy/app_unittest_utils/galaxy_mock.py b/lib/galaxy/app_unittest_utils/galaxy_mock.py index 82a0b788665..9f4fe40139e 100644 --- a/lib/galaxy/app_unittest_utils/galaxy_mock.py +++ b/lib/galaxy/app_unittest_utils/galaxy_mock.py @@ -119,6 +119,7 @@ class MockApp(di.Container, GalaxyDataTestApp): history_manager: HistoryManager job_metrics: JobMetrics vault: Optional[Vault] = None + execution_timer_factory: Any stop: bool is_webapp: bool = True @@ -159,7 +160,7 @@ class MockApp(di.Container, GalaxyDataTestApp): self.application_stack = ApplicationStack() self.auth_manager = AuthManager(self.config) self.user_manager = UserManager(cast(BasicSharedApp, self)) - self.execution_timer_factory = Bunch(get_timer=StructuredExecutionTimer) + self.execution_timer_factory = Bunch(get_timer=StructuredExecutionTimer, galaxy_statsd_client=None) self.interactivetool_manager = Bunch(create_interactivetool=lambda *args, **kwargs: None) self.is_job_handler = False self.biotools_metadata_source = None diff --git a/lib/galaxy/authnz/psa_authnz.py b/lib/galaxy/authnz/psa_authnz.py index 75247adc4ce..022272e3d9d 100644 --- a/lib/galaxy/authnz/psa_authnz.py +++ b/lib/galaxy/authnz/psa_authnz.py @@ -798,6 +798,7 @@ def _send_oidc_profile_update_notification(trans, user, updates: list[str]) -> N NotificationVariant, PersonalNotificationCategory, ) + labels: dict[str, str] = { "email": "email address", "username": "public name", diff --git a/lib/galaxy/celery/__init__.py b/lib/galaxy/celery/__init__.py index 67917cca658..dce38961f3c 100644 --- a/lib/galaxy/celery/__init__.py +++ b/lib/galaxy/celery/__init__.py @@ -252,6 +252,9 @@ def setup_periodic_tasks(config, celery_app): schedule_task("prune_history_audit_table", config.history_audit_table_prune_interval) schedule_task("cleanup_short_term_storage", config.short_term_storage_cleanup_interval) + if config.statsd_host: + schedule_task("emit_queue_metrics_task", config.queue_metrics_interval) + if config.enable_notification_system: schedule_task("cleanup_expired_notifications", config.expired_notifications_cleanup_interval) schedule_task("dispatch_pending_notifications", config.dispatch_notifications_interval) diff --git a/lib/galaxy/celery/tasks.py b/lib/galaxy/celery/tasks.py index b630fdce26a..eef0c951bbb 100644 --- a/lib/galaxy/celery/tasks.py +++ b/lib/galaxy/celery/tasks.py @@ -76,7 +76,10 @@ from galaxy.security.vault import ( Vault, ) from galaxy.short_term_storage import ShortTermStorageMonitor -from galaxy.structured_app import MinimalManagerApp +from galaxy.structured_app import ( + MinimalManagerApp, + StructuredApp, +) from galaxy.tools import create_tool_from_representation from galaxy.tools.data_fetch import do_fetch from galaxy.util import galaxy_directory @@ -628,6 +631,14 @@ def dispatch_pending_notifications(notification_manager: NotificationManager): log.info(f"Successfully dispatched {count} notifications.") +@galaxy_task(action="emit queue and SSE observability metrics") +def emit_queue_metrics_task(app: StructuredApp): + """Sample control-queue depth, SSE connection count, and worker rows → statsd.""" + from galaxy.webapps.galaxy.metrics.queue_metrics import emit_queue_metrics + + emit_queue_metrics(app) + + @galaxy_task(action="clean up job working directories") def cleanup_jwds(sa_session: galaxy_scoped_session, object_store: BaseObjectStore, config: GalaxyAppConfiguration): """Cleanup job working directories for failed jobs that are older than X days""" diff --git a/lib/galaxy/config/sample/galaxy.yml.sample b/lib/galaxy/config/sample/galaxy.yml.sample index 34ae86e4b1b..b4fe9d9556e 100644 --- a/lib/galaxy/config/sample/galaxy.yml.sample +++ b/lib/galaxy/config/sample/galaxy.yml.sample @@ -1962,6 +1962,11 @@ galaxy: # really. Do not set this in production environments. #statsd_mock_calls: false + # How often (in seconds) the Celery beat task emits queue-depth, + # SSE-connection, and WorkerProcess gauges. Only active when + # statsd_host is set. Set to 0 to disable. + #queue_metrics_interval: 15 + # Add an option to the library upload form which allows administrators # to upload a directory of files. #library_import_dir: null @@ -3136,6 +3141,13 @@ galaxy: # browsers, replacing aggressive 3-second polling. #enable_sse_history_updates: false + # Enables real-time interactive-tool entry-point update notifications + # via Server-Sent Events. When enabled, the client subscribes to + # entry_point_update SSE events and refetches the entry-point list on + # each event, replacing the 10-second polling loop. When disabled, + # polling remains the source of updates. + #enable_sse_entry_point_updates: false + # The interval in seconds between history audit table polls when using # the polling fallback (SQLite or when PostgreSQL LISTEN/NOTIFY is # unavailable). Only used when enable_sse_history_updates is true. diff --git a/lib/galaxy/config/schemas/config_schema.yml b/lib/galaxy/config/schemas/config_schema.yml index aba88b87cee..cbba5927ac7 100644 --- a/lib/galaxy/config/schemas/config_schema.yml +++ b/lib/galaxy/config/schemas/config_schema.yml @@ -2512,6 +2512,14 @@ mapping: Mock out statsd client calls - only used by testing infrastructure really. Do not set this in production environments. + queue_metrics_interval: + type: int + default: 15 + required: false + desc: | + How often (in seconds) the Celery beat task emits queue-depth, SSE-connection, + and WorkerProcess gauges. Only active when statsd_host is set. Set to 0 to disable. + library_import_dir: type: str required: false @@ -4303,6 +4311,16 @@ mapping: LISTEN/NOTIFY or audit table polling as a fallback for SQLite) and pushes update signals to connected browsers, replacing aggressive 3-second polling. + enable_sse_entry_point_updates: + type: bool + default: false + required: false + desc: | + Enables real-time interactive-tool entry-point update notifications via + Server-Sent Events. When enabled, the client subscribes to entry_point_update + SSE events and refetches the entry-point list on each event, replacing the + 10-second polling loop. When disabled, polling remains the source of updates. + history_audit_monitor_poll_interval: type: int default: 2 diff --git a/lib/galaxy/managers/interactivetool.py b/lib/galaxy/managers/interactivetool.py index 009d8f8dc36..21b872f146e 100644 --- a/lib/galaxy/managers/interactivetool.py +++ b/lib/galaxy/managers/interactivetool.py @@ -6,6 +6,7 @@ from collections.abc import ( ) from typing import ( Any, + Optional, TYPE_CHECKING, Union, ) @@ -28,6 +29,7 @@ from sqlalchemy import ( ) from galaxy import exceptions +from galaxy.managers.sse_dispatch import SSEEventDispatcher from galaxy.model import ( InteractiveToolEntryPoint, Job, @@ -147,7 +149,11 @@ class InteractiveToolManager: Manager for dealing with InteractiveTools """ - def __init__(self, app: "MinimalManagerApp") -> None: + def __init__( + self, + app: "MinimalManagerApp", + dispatcher: Optional[SSEEventDispatcher] = None, + ) -> None: self.app = app self.security = app.security self.sa_session = app.model.context @@ -157,6 +163,12 @@ class InteractiveToolManager: app.config.interactivetoolsproxy_map or app.config.interactivetools_map, self.encoder.encode_id, ) + # Lagom can't auto-inject ``SSEEventDispatcher`` here because the + # ``app: "MinimalManagerApp"`` hint is only a forward reference + # (TYPE_CHECKING import), so ``get_type_hints`` on this signature + # fails. Resolve through the container explicitly — ``resolve_or_none`` + # returns ``None`` for mocks/test apps that never registered one. + self.dispatcher = dispatcher if dispatcher is not None else app.resolve_or_none(SSEEventDispatcher) def create_entry_points( self, job: Job, tool: "Tool", entry_points=Union[Iterable[dict[str, Any]], None], flush: bool = True @@ -198,6 +210,17 @@ class InteractiveToolManager: configured.append(ep) if configured: self.sa_session.commit() + # Fan out an SSE push so the user's browser can refresh the entry + # point list immediately instead of waiting for the 10 s poll. + # Anonymous jobs fall back to polling — ``push_to_user`` keys on + # user_id, and anonymous clients sit in the broadcast-only set. + if self.dispatcher is not None and job.user_id is not None: + try: + self.dispatcher.entry_point_update(user_id=job.user_id) + except Exception: + # The DB commit is authoritative; the SSE event is best + # effort. Never let a dispatch failure poison the caller. + log.exception("Failed to dispatch entry_point_update SSE event for job %s", job.id) return dict(not_configured=not_configured, configured=configured) def save_entry_point(self, entry_point: InteractiveToolEntryPoint) -> None: diff --git a/lib/galaxy/managers/notification.py b/lib/galaxy/managers/notification.py index 67666034253..18de7e795a5 100644 --- a/lib/galaxy/managers/notification.py +++ b/lib/galaxy/managers/notification.py @@ -59,8 +59,8 @@ from galaxy.schema.notifications import ( NotificationBroadcastUpdateRequest, NotificationCategorySettings, NotificationChannelSettings, - NotificationCreatedResponse, NotificationCreateData, + NotificationCreatedResponse, NotificationCreateRequest, NotificationRecipients, NotificationResponse, diff --git a/lib/galaxy/managers/sse.py b/lib/galaxy/managers/sse.py index 2c83a549e08..66b54bb002d 100644 --- a/lib/galaxy/managers/sse.py +++ b/lib/galaxy/managers/sse.py @@ -20,6 +20,7 @@ from typing import ( ) from galaxy.model.orm.now import now +from galaxy.web.statsd_client import VanillaGalaxyStatsdClient log = logging.getLogger(__name__) @@ -78,10 +79,11 @@ class SSEConnectionManager: (typically the Kombu daemon thread via control task handlers). """ - def __init__(self) -> None: + def __init__(self, statsd_client: Optional[VanillaGalaxyStatsdClient] = None) -> None: self._connections: dict[int, set[asyncio.Queue]] = defaultdict(set) self._broadcast_connections: set[asyncio.Queue] = set() self._loop: Optional[asyncio.AbstractEventLoop] = None + self._statsd_client = statsd_client def _ensure_loop(self) -> None: """Capture the running asyncio event loop. Must be called from async context.""" @@ -148,13 +150,14 @@ class SSEConnectionManager: # Event loop is closed or shutting down pass - @staticmethod - def _do_put(queue: asyncio.Queue, event: SSEEvent) -> None: + def _do_put(self, queue: asyncio.Queue, event: SSEEvent) -> None: """Runs ON the event loop thread. Safe to touch asyncio.Queue here.""" try: queue.put_nowait(event) except asyncio.QueueFull: log.warning("SSE queue full, dropping event: %s", event.event) + if self._statsd_client is not None: + self._statsd_client.incr("galaxy.sse.connections.dropped") @property def connected_user_ids(self) -> set[int]: @@ -164,6 +167,20 @@ class SSEConnectionManager: def total_connections(self) -> int: return len(self._broadcast_connections) + @property + def total_broadcast_connections(self) -> int: + """Number of all active SSE connections (includes anonymous). + + Every connection is added to ``_broadcast_connections``; this is the + most accurate "SSE clients currently connected" gauge. + """ + return len(self._broadcast_connections) + + @property + def total_per_user_connections(self) -> int: + """Number of active SSE connections bound to a specific user_id.""" + return sum(len(queues) for queues in self._connections.values()) + # -- High-level streaming helper -- async def stream( diff --git a/lib/galaxy/managers/sse_dispatch.py b/lib/galaxy/managers/sse_dispatch.py index 8db0cb6e63a..927cb1ab0d8 100644 --- a/lib/galaxy/managers/sse_dispatch.py +++ b/lib/galaxy/managers/sse_dispatch.py @@ -9,17 +9,24 @@ inline imports in the hot path. """ import logging +import threading +import time +from collections.abc import Callable from typing import ( Any, Optional, ) +from cachetools import TTLCache +from kombu import Queue + from galaxy.managers.sse import make_event_id from galaxy.queue_worker import ( ControlTask, GalaxyQueueWorker, ) from galaxy.queues import all_control_queues_for_declare +from galaxy.web.statsd_client import VanillaGalaxyStatsdClient from galaxy.web_stack import ApplicationStack log = logging.getLogger(__name__) @@ -32,31 +39,70 @@ class SSEEventDispatcher: without a full ``app``. ``queue_worker`` is ``Optional`` because unit-test mock apps and Galaxy configurations without AMQP don't construct one — the dispatcher silently no-ops in that case. + + ``statsd_client`` is optional — if ``None`` (statsd not configured), all + instrumentation becomes a cheap attribute-lookup no-op. """ + # TTL for the active-worker declare-queue cache. The WorkerProcess heartbeat + # writes every 60 s and ``all_control_queues_for_declare`` filters on a 120 s + # window, so a 30 s cache cannot produce a result that wasn't also valid in + # the non-cached call. Surfaced as a class constant so tests can monkey-patch. + _DECLARE_QUEUES_TTL_SECONDS = 30 + def __init__( self, queue_worker: Optional[GalaxyQueueWorker], application_stack: ApplicationStack, + statsd_client: Optional[VanillaGalaxyStatsdClient] = None, + clock: Callable[[], float] = time.monotonic, ) -> None: self._queue_worker = queue_worker self._application_stack = application_stack + self._statsd_client = statsd_client + self._clock = clock + self._declare_queues_cache: TTLCache = TTLCache(maxsize=1, ttl=self._DECLARE_QUEUES_TTL_SECONDS, timer=clock) + self._declare_queues_lock = threading.RLock() + + def _get_declare_queues(self) -> list[Queue]: + # Empty results (startup before DatabaseHeartbeat registers this process, + # or a transient DB error swallowed by ``all_control_queues_for_declare``) + # must not be pinned for the full TTL — they'd silently drop every SSE + # event until the next expiry. Only cache non-empty results. + with self._declare_queues_lock: + try: + return self._declare_queues_cache["webapp"] + except KeyError: + queues = all_control_queues_for_declare(self._application_stack, webapp_only=True) + if queues: + self._declare_queues_cache["webapp"] = queues + return queues def _send(self, task: str, kwargs: dict[str, Any]) -> None: if self._queue_worker is None: # AMQP not configured at all (e.g. unit-test mock app). Skip silently. log.debug("SSE dispatch skipped: no queue_worker configured (task=%s)", task) + if self._statsd_client is not None: + self._statsd_client.incr("galaxy.sse.dispatch.skipped_no_qw") return + if self._statsd_client is not None: + self._statsd_client.incr("galaxy.sse.dispatch.count", tags={"task": task}) # Only fan out to webapp processes — job handlers and workflow schedulers # don't have browser SSE connections to push to. - declare_queues = all_control_queues_for_declare(self._application_stack, webapp_only=True) + declare_queues = self._get_declare_queues() control_task = ControlTask(self._queue_worker) - control_task.send_task( - payload={"task": task, "kwargs": kwargs}, - routing_key="control.*", - expiration=10, - declare_queues=declare_queues, - ) + start_time = time.perf_counter() if self._statsd_client is not None else 0.0 + try: + control_task.send_task( + payload={"task": task, "kwargs": kwargs}, + routing_key="control.*", + expiration=10, + declare_queues=declare_queues, + ) + finally: + if self._statsd_client is not None: + dt_ms = int((time.perf_counter() - start_time) * 1000) + self._statsd_client.timing("galaxy.sse.dispatch.latency_ms", dt_ms, tags={"task": task}) def notify_users(self, user_ids: list[int], payload: str, event_id: Optional[str] = None) -> None: self._send( @@ -85,3 +131,18 @@ class SSEEventDispatcher: "event_id": event_id or make_event_id(), }, ) + + def entry_point_update(self, user_id: int, event_id: Optional[str] = None) -> None: + """Fan out a wake-up ``entry_point_update`` event for one user. + + The client always refetches the canonical entry-point list on receipt, + so no IDs are sent — keeping the payload small and the dispatch path + free of per-event encoding work. + """ + self._send( + "entry_point_update", + { + "user_id": user_id, + "event_id": event_id or make_event_id(), + }, + ) diff --git a/lib/galaxy/model/unittest_utils/data_app.py b/lib/galaxy/model/unittest_utils/data_app.py index 0d424cca362..2eb7e78bf3e 100644 --- a/lib/galaxy/model/unittest_utils/data_app.py +++ b/lib/galaxy/model/unittest_utils/data_app.py @@ -9,6 +9,7 @@ galaxy-data. import os import shutil import tempfile +from types import SimpleNamespace from typing import Optional from galaxy import ( @@ -105,6 +106,11 @@ class GalaxyDataTestApp: model.setup_global_object_store_for_models(self.object_store) self.security_agent = self.model.security_agent self.tag_handler = GalaxyTagHandler(self.model.session) + # statsd/observability plumbing — stubbed out so paths that read + # ``app.execution_timer_factory.galaxy_statsd_client`` (e.g. the + # queue-worker instrumentation) degrade to a no-op if ever exercised + # against this mock instead of raising ``AttributeError``. + self.execution_timer_factory = SimpleNamespace(galaxy_statsd_client=None) self.init_datatypes() def init_datatypes(self): diff --git a/lib/galaxy/queue_worker/__init__.py b/lib/galaxy/queue_worker/__init__.py index 04824d8a849..a01fc27efc9 100644 --- a/lib/galaxy/queue_worker/__init__.py +++ b/lib/galaxy/queue_worker/__init__.py @@ -13,8 +13,10 @@ import threading import time from inspect import ismodule from typing import ( + cast, Optional, TYPE_CHECKING, + TypedDict, ) from kombu import ( @@ -48,6 +50,39 @@ if TYPE_CHECKING: ) +class NotifyUsersPayload(TypedDict, total=False): + """Wire contract for the ``notify_users`` control-task kwargs.""" + + user_ids: list[int] + payload: str + event_id: Optional[str] + + +class NotifyBroadcastPayload(TypedDict, total=False): + """Wire contract for the ``notify_broadcast`` control-task kwargs.""" + + payload: str + event_id: Optional[str] + + +class HistoryUpdatePayload(TypedDict, total=False): + """Wire contract for the ``history_update`` control-task kwargs. + + ``user_updates`` maps stringified user IDs to lists of (unencoded) history IDs. + Stringified because AMQP JSON serialization coerces dict keys to strings. + """ + + user_updates: dict[str, list[int]] + event_id: Optional[str] + + +class EntryPointUpdatePayload(TypedDict, total=False): + """Wire contract for the ``entry_point_update`` control-task kwargs.""" + + user_id: int + event_id: Optional[str] + + def send_local_control_task( app: "StructuredApp", task: str, @@ -361,21 +396,26 @@ def admin_job_lock(app, **kwargs): def notify_users(app: "MinimalManagerApp", **kwargs) -> None: """Push SSE events to connected users on this worker process.""" + payload = cast(NotifyUsersPayload, kwargs) sse_manager = app[SSEConnectionManager] - user_ids = kwargs.get("user_ids", []) - payload = kwargs.get("payload", "{}") - event_id = kwargs.get("event_id") - event = SSEEvent(event="notification_update", data=payload, id=event_id) - for user_id in user_ids: + event = SSEEvent( + event="notification_update", + data=payload.get("payload", "{}"), + id=payload.get("event_id"), + ) + for user_id in payload.get("user_ids", []): sse_manager.push_to_user(user_id, event) def notify_broadcast(app: "MinimalManagerApp", **kwargs) -> None: """Push SSE broadcast events to all connected clients on this worker process.""" + payload = cast(NotifyBroadcastPayload, kwargs) sse_manager = app[SSEConnectionManager] - payload = kwargs.get("payload", "{}") - event_id = kwargs.get("event_id") - event = SSEEvent(event="broadcast_update", data=payload, id=event_id) + event = SSEEvent( + event="broadcast_update", + data=payload.get("payload", "{}"), + id=payload.get("event_id"), + ) sse_manager.push_broadcast(event) @@ -385,11 +425,11 @@ def history_update(app: "MinimalManagerApp", **kwargs) -> None: Encodes integer history IDs here (not in the monitor) so the manager layer stays free of presentation/security concerns. """ + payload = cast(HistoryUpdatePayload, kwargs) sse_manager = app[SSEConnectionManager] - user_updates = kwargs.get("user_updates", {}) - event_id = kwargs.get("event_id") + event_id = payload.get("event_id") encode = app.security.encode_id - for user_id_str, history_ids in user_updates.items(): + for user_id_str, history_ids in payload.get("user_updates", {}).items(): user_id = int(user_id_str) encoded_ids = [encode(hid) for hid in history_ids] data = json.dumps({"history_ids": encoded_ids}) @@ -397,6 +437,20 @@ def history_update(app: "MinimalManagerApp", **kwargs) -> None: sse_manager.push_to_user(user_id, event) +def entry_point_update(app: "MinimalManagerApp", **kwargs) -> None: + """Push a wake-up SSE event to a single connected user. + + The payload is empty by design: the client refetches ``/api/entry_points`` + (the canonical source) on receipt, so there's nothing to narrow or merge. + Dropping the IDs from the payload also avoids per-event ``encode_id`` work + at 1000+ events/s. + """ + payload = cast(EntryPointUpdatePayload, kwargs) + sse_manager = app[SSEConnectionManager] + event = SSEEvent(event="entry_point_update", data="{}", id=payload.get("event_id")) + sse_manager.push_to_user(int(payload["user_id"]), event) + + control_message_to_task = { "create_panel_section": create_panel_section, "reload_tool": reload_tool, @@ -415,6 +469,7 @@ control_message_to_task = { "notify_users": notify_users, "notify_broadcast": notify_broadcast, "history_update": history_update, + "entry_point_update": entry_point_update, } @@ -504,8 +559,14 @@ class GalaxyQueueWorker(ConsumerProducerMixin, threading.Thread): def process_task(self, body, message): result = "NO_RESULT" + task_name = body.get("task") + statsd_client = self.app.execution_timer_factory.galaxy_statsd_client + if statsd_client is not None and task_name is not None: + statsd_client.incr("galaxy.control_queue.task.count", tags={"task": task_name}) if body["task"] in self.task_mapping: if body.get("noop", None) != self.app.config.server_name: + outcome = "ok" + handler_start = time.perf_counter() if statsd_client is not None else 0.0 try: f = self.task_mapping[body["task"]] if message.headers.get("epoch", math.inf) > self.epoch: @@ -526,7 +587,16 @@ class GalaxyQueueWorker(ConsumerProducerMixin, threading.Thread): result = "NO_OP" except Exception: # this shouldn't ever throw an exception, but... + outcome = "error" log.exception("Error running control task type: %s", body["task"]) + finally: + if statsd_client is not None and task_name is not None: + dt_ms = int((time.perf_counter() - handler_start) * 1000) + statsd_client.timing( + "galaxy.control_queue.task.latency_ms", + dt_ms, + tags={"task": task_name, "outcome": outcome}, + ) else: result = "NO_OP" else: diff --git a/lib/galaxy/structured_app/__init__.py b/lib/galaxy/structured_app/__init__.py index af78cf03d31..fc342739f5f 100644 --- a/lib/galaxy/structured_app/__init__.py +++ b/lib/galaxy/structured_app/__init__.py @@ -176,6 +176,7 @@ class StructuredApp(MinimalManagerApp): vault: Vault webhooks_registry: WebhooksRegistry queue_worker: Any # 'galaxy.queue_worker.GalaxyQueueWorker' + execution_timer_factory: Any # 'galaxy.app.ExecutionTimerFactory' data_provider_registry: Any # 'galaxy.visualization.data_providers.registry.DataProviderRegistry' tool_cache: "ToolCache" tool_shed_repository_cache: Optional[ToolShedRepositoryCache] diff --git a/lib/galaxy/webapps/galaxy/api/tool_data.py b/lib/galaxy/webapps/galaxy/api/tool_data.py index c7ae44f45ae..3d1f1a9abb8 100644 --- a/lib/galaxy/webapps/galaxy/api/tool_data.py +++ b/lib/galaxy/webapps/galaxy/api/tool_data.py @@ -11,6 +11,7 @@ from pydantic import ( Field, ) +from galaxy.celery.helpers import async_task_summary from galaxy.celery.tasks import import_data_bundle from galaxy.managers.context import ProvidesUserContext from galaxy.managers.tool_data import ToolDataManager @@ -25,7 +26,6 @@ from galaxy.tool_util.data._schema import ( ToolDataItem, ) from galaxy.webapps.base.api import GalaxyFileResponse -from galaxy.webapps.galaxy.services.base import async_task_summary from . import ( depends, DependsOnTrans, diff --git a/lib/galaxy/webapps/galaxy/metrics/__init__.py b/lib/galaxy/webapps/galaxy/metrics/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/lib/galaxy/webapps/galaxy/metrics/queue_metrics.py b/lib/galaxy/webapps/galaxy/metrics/queue_metrics.py new file mode 100644 index 00000000000..c45dcfd7110 --- /dev/null +++ b/lib/galaxy/webapps/galaxy/metrics/queue_metrics.py @@ -0,0 +1,147 @@ +"""Periodic gauge emitter for control-queue depth and SSE-connection counts. + +Scheduled by Celery beat (see ``galaxy.celery.__init__.setup_periodic_tasks``) +at a fixed cadence. Opens a short-lived kombu connection, iterates the control +queues returned by ``all_control_queues_for_declare`` and samples each queue's +message-count via a passive declare. Also samples in-memory connection counts +from ``SSEConnectionManager`` and the active-``WorkerProcess`` count from the +database. + +All instrumentation no-ops when ``app.execution_timer_factory.galaxy_statsd_client`` +is ``None`` — i.e. statsd isn't configured. +""" + +import datetime +import logging +from collections import defaultdict +from typing import TYPE_CHECKING + +from lagom.exceptions import UnresolvableType +from sqlalchemy import ( + func, + select, +) + +from galaxy.managers.sse import SSEConnectionManager +from galaxy.model import WorkerProcess +from galaxy.model.orm.now import now +from galaxy.queues import ( + all_control_queues_for_declare, + DEFAULT_ACTIVE_PROCESS_WINDOW_SECONDS, +) + +if TYPE_CHECKING: + from galaxy.structured_app import StructuredApp + from galaxy.web.statsd_client import VanillaGalaxyStatsdClient + +log = logging.getLogger(__name__) + + +def emit_sse_connection_gauges( + statsd_client: "VanillaGalaxyStatsdClient", + sse_manager: SSEConnectionManager, +) -> None: + """Emit ``galaxy.sse.connections.active`` gauges by connection kind.""" + statsd_client.timing( + "galaxy.sse.connections.active", + sse_manager.total_broadcast_connections, + tags={"kind": "broadcast"}, + ) + statsd_client.timing( + "galaxy.sse.connections.active", + sse_manager.total_per_user_connections, + tags={"kind": "per_user"}, + ) + + +def emit_control_queue_depth( + statsd_client: "VanillaGalaxyStatsdClient", + app: "StructuredApp", +) -> None: + """Emit ``galaxy.control_queue.depth`` per active webapp/handler queue. + + A per-queue passive declare can fail on transports that don't implement it + (e.g. the sqlalchemy kombu transport) or for queues that don't yet exist on + the broker. Those are expected and quiet — logged at DEBUG, no metric, move + on. Errors at the broker-connection layer propagate up so the caller can + surface them. + """ + connection = app.amqp_internal_connection_obj + if connection is None: + return + queues = all_control_queues_for_declare(app.application_stack) + if not queues: + return + with connection.clone() as conn: + channel = conn.channel() + try: + for queue in queues: + try: + declared = queue.bind(channel).queue_declare(passive=True) + except Exception: + log.debug( + "queue_metrics: passive declare failed for %s", + queue.name, + exc_info=True, + ) + continue + statsd_client.timing( + "galaxy.control_queue.depth", + declared.message_count, + tags={"queue_name": queue.name}, + ) + finally: + channel.close() + + +def emit_worker_process_gauge( + statsd_client: "VanillaGalaxyStatsdClient", + app: "StructuredApp", +) -> None: + """Emit ``galaxy.worker_process.active`` gauge grouped by ``app_type``.""" + cutoff = now() - datetime.timedelta(seconds=DEFAULT_ACTIVE_PROCESS_WINDOW_SECONDS) + stmt = ( + select(WorkerProcess.app_type, func.count(WorkerProcess.id)) + .where(WorkerProcess.update_time > cutoff) + .group_by(WorkerProcess.app_type) + ) + counts: dict[str, int] = defaultdict(int) + with app.model.new_session() as session: + for app_type, count in session.execute(stmt): + counts[app_type or "unknown"] = int(count) + for app_type, count in counts.items(): + statsd_client.timing( + "galaxy.worker_process.active", + count, + tags={"app_type": app_type}, + ) + + +def _run(name: str, statsd_client: "VanillaGalaxyStatsdClient", fn) -> None: + """Run a sub-emitter, isolating its failures. + + A broken sub-emitter logs once at WARNING and increments + ``galaxy.queue_metrics.error`` (tagged by emitter name) so the failure is + observable in metrics without the Celery-beat wrapper logging on every + tick. The other sub-emitters continue to run on this tick. + """ + try: + fn() + except Exception: + log.warning("queue_metrics: %s emitter failed", name, exc_info=True) + statsd_client.incr("galaxy.queue_metrics.error", tags={"emitter": name}) + + +def emit_queue_metrics(app: "StructuredApp") -> None: + """Periodic entry-point — no-ops when statsd isn't configured.""" + statsd_client = app.execution_timer_factory.galaxy_statsd_client + if statsd_client is None: + return + try: + sse_manager = app[SSEConnectionManager] + except UnresolvableType: + sse_manager = None + if sse_manager is not None: + _run("sse_connections", statsd_client, lambda: emit_sse_connection_gauges(statsd_client, sse_manager)) + _run("control_queue_depth", statsd_client, lambda: emit_control_queue_depth(statsd_client, app)) + _run("worker_process", statsd_client, lambda: emit_worker_process_gauge(statsd_client, app)) diff --git a/lib/galaxy/webapps/galaxy/services/base.py b/lib/galaxy/webapps/galaxy/services/base.py index d485d277d5b..e8692496618 100644 --- a/lib/galaxy/webapps/galaxy/services/base.py +++ b/lib/galaxy/webapps/galaxy/services/base.py @@ -7,7 +7,6 @@ from typing import ( Optional, ) -from galaxy.celery.helpers import async_task_summary as async_task_summary # re-export for existing callers from galaxy.exceptions import ( AuthenticationRequired, ConfigDoesNotAllowException, diff --git a/lib/galaxy/webapps/galaxy/services/datasets.py b/lib/galaxy/webapps/galaxy/services/datasets.py index aae67a7ec74..d3f7f8ee8b3 100644 --- a/lib/galaxy/webapps/galaxy/services/datasets.py +++ b/lib/galaxy/webapps/galaxy/services/datasets.py @@ -24,6 +24,7 @@ from galaxy import ( util, web, ) +from galaxy.celery.helpers import async_task_summary from galaxy.celery.tasks import compute_dataset_hash from galaxy.datatypes.binary import Binary from galaxy.datatypes.dataproviders.exceptions import NoProviderAvailable @@ -88,10 +89,7 @@ from galaxy.visualization.data_providers.genome import ( ) from galaxy.visualization.data_providers.registry import DataProviderRegistry from galaxy.webapps.base.controller import UsesVisualizationMixin -from galaxy.webapps.galaxy.services.base import ( - async_task_summary, - ServiceBase, -) +from galaxy.webapps.galaxy.services.base import ServiceBase log = logging.getLogger(__name__) diff --git a/lib/galaxy/webapps/galaxy/services/histories.py b/lib/galaxy/webapps/galaxy/services/histories.py index 25e8982da5d..097fbc1c96f 100644 --- a/lib/galaxy/webapps/galaxy/services/histories.py +++ b/lib/galaxy/webapps/galaxy/services/histories.py @@ -24,6 +24,7 @@ from galaxy import ( exceptions as glx_exceptions, model, ) +from galaxy.celery.helpers import async_task_summary from galaxy.celery.tasks import ( import_model_store, prepare_history_download, @@ -87,7 +88,6 @@ from galaxy.security.idencoding import IdEncodingHelper from galaxy.short_term_storage import ShortTermStorageAllocator from galaxy.util import restore_text from galaxy.webapps.galaxy.services.base import ( - async_task_summary, ConsumesModelStores, model_store_storage_target, ServesExportStores, diff --git a/lib/galaxy/webapps/galaxy/services/history_contents.py b/lib/galaxy/webapps/galaxy/services/history_contents.py index 581d1dc772e..53926007bd3 100644 --- a/lib/galaxy/webapps/galaxy/services/history_contents.py +++ b/lib/galaxy/webapps/galaxy/services/history_contents.py @@ -21,6 +21,7 @@ from typing_extensions import ( ) from galaxy import exceptions +from galaxy.celery.helpers import async_task_summary from galaxy.celery.tasks import ( change_datatype, materialize as materialize_task, @@ -118,7 +119,6 @@ from galaxy.security.idencoding import IdEncodingHelper from galaxy.short_term_storage import ShortTermStorageAllocator from galaxy.util.zipstream import ZipstreamWrapper from galaxy.webapps.galaxy.services.base import ( - async_task_summary, ConsumesModelStores, ensure_celery_tasks_enabled, model_store_storage_target, diff --git a/lib/galaxy/webapps/galaxy/services/invocations.py b/lib/galaxy/webapps/galaxy/services/invocations.py index 29654d885a4..dc5cc27a577 100644 --- a/lib/galaxy/webapps/galaxy/services/invocations.py +++ b/lib/galaxy/webapps/galaxy/services/invocations.py @@ -6,6 +6,7 @@ from typing import ( from pydantic import Field +from galaxy.celery.helpers import async_task_summary from galaxy.celery.tasks import ( prepare_invocation_download, write_invocation_to, @@ -61,7 +62,6 @@ from galaxy.schema.tasks import ( from galaxy.security.idencoding import IdEncodingHelper from galaxy.short_term_storage import ShortTermStorageAllocator from galaxy.webapps.galaxy.services.base import ( - async_task_summary, ConsumesModelStores, ensure_celery_tasks_enabled, model_store_storage_target, diff --git a/lib/galaxy/webapps/galaxy/services/jobs.py b/lib/galaxy/webapps/galaxy/services/jobs.py index fac6ea9ecf3..daadbe37b13 100644 --- a/lib/galaxy/webapps/galaxy/services/jobs.py +++ b/lib/galaxy/webapps/galaxy/services/jobs.py @@ -16,6 +16,7 @@ from galaxy import ( exceptions, model, ) +from galaxy.celery.helpers import async_task_summary from galaxy.celery.tasks import queue_jobs from galaxy.managers import hdas from galaxy.managers.base import security_check @@ -60,7 +61,6 @@ from galaxy.tool_util.parameters import ( ToolParameterBundleModel, ) from galaxy.webapps.galaxy.services.base import ( - async_task_summary, ServiceBase, ) from .tools import validate_tool_for_running diff --git a/lib/galaxy/webapps/galaxy/services/pages.py b/lib/galaxy/webapps/galaxy/services/pages.py index d77560efd27..e95e6dd959d 100644 --- a/lib/galaxy/webapps/galaxy/services/pages.py +++ b/lib/galaxy/webapps/galaxy/services/pages.py @@ -4,6 +4,7 @@ from typing import ( ) from galaxy import exceptions +from galaxy.celery.helpers import async_task_summary from galaxy.celery.tasks import prepare_pdf_download from galaxy.managers import base from galaxy.managers.markdown_util import ( @@ -34,7 +35,6 @@ from galaxy.security.idencoding import IdEncodingHelper from galaxy.short_term_storage import ShortTermStorageAllocator from galaxy.webapps.galaxy.api.common import PageIdPathParam from galaxy.webapps.galaxy.services.base import ( - async_task_summary, ensure_celery_tasks_enabled, ServiceBase, ) diff --git a/lib/galaxy/webapps/galaxy/services/users.py b/lib/galaxy/webapps/galaxy/services/users.py index 39aa35c03a4..dd8fa7650c1 100644 --- a/lib/galaxy/webapps/galaxy/services/users.py +++ b/lib/galaxy/webapps/galaxy/services/users.py @@ -9,6 +9,7 @@ from galaxy import ( exceptions as glx_exceptions, util, ) +from galaxy.celery.helpers import async_task_summary from galaxy.managers import api_keys from galaxy.managers.context import ( ProvidesHistoryContext, @@ -34,10 +35,7 @@ from galaxy.schema.schema import ( UserModel, ) from galaxy.security.idencoding import IdEncodingHelper -from galaxy.webapps.galaxy.services.base import ( - async_task_summary, - ServiceBase, -) +from galaxy.webapps.galaxy.services.base import ServiceBase from galaxy.webapps.galaxy.services.roles import role_to_model if TYPE_CHECKING: diff --git a/test/integration/test_entry_point_sse.py b/test/integration/test_entry_point_sse.py new file mode 100644 index 00000000000..ead6e24b10e --- /dev/null +++ b/test/integration/test_entry_point_sse.py @@ -0,0 +1,158 @@ +"""Integration tests for SSE-based interactive-tool entry-point update notifications. + +Mirrors ``test_history_sse.py``. Rather than spin up a real containerized +interactive tool (which would require docker), these tests exercise the +dispatch path by creating a ``Job`` and ``InteractiveToolEntryPoint`` rows +directly via the live app's SQLAlchemy session and invoking +``InteractiveToolManager.configure_entry_points`` with a stub ``ports_dict``. +This is exactly the moment the event fires in production — the job runner's +port-routing hook is the only upstream caller, and the SSE dispatch happens +after the DB commit regardless of how the ports were obtained. +""" + +import time +from urllib.parse import urljoin +from uuid import uuid4 + +from galaxy.model import ( + InteractiveToolEntryPoint, + Job, +) +from galaxy_test.base.populators import DatasetPopulator +from galaxy_test.base.sse import SSELineListener +from galaxy_test.driver.integration_util import IntegrationTestCase + + +def _make_ports_dict(tool_port: int) -> dict: + """Stub the runner's port-routing payload — one tool_port, fake host/proto.""" + return { + str(tool_port): { + "host": "host.invalid", + "port": 12345, + "protocol": "http", + } + } + + +class TestEntryPointSSEIntegration(IntegrationTestCase): + dataset_populator: DatasetPopulator + framework_tool_and_types = True + + @classmethod + def handle_galaxy_config_kwds(cls, config): + super().handle_galaxy_config_kwds(config) + config["enable_celery_tasks"] = False + + def setUp(self): + super().setUp() + self.dataset_populator = DatasetPopulator(self.galaxy_interactor) + + def _events_stream_url(self) -> str: + return urljoin(self.url, "api/events/stream") + + def _user_id_for_api_key(self, api_key: str) -> int: + """Return the integer ``User.id`` for the user owning ``api_key``.""" + # The ``/api/users/current`` endpoint returns the encoded id; decode + # via the app's security helper so we get the raw int the job row + # needs. + import requests + + response = requests.get(urljoin(self.url, "api/users/current"), params={"key": api_key}) + response.raise_for_status() + encoded_id = response.json()["id"] + return self._app.security.decode_id(encoded_id) + + def _create_it_job_with_entry_point(self, user_id: int, tool_port: int = 8888) -> tuple[int, int]: + """Create a minimal Job + unconfigured InteractiveToolEntryPoint row pair. + + Returns ``(job_id, entry_point_id)``. The session used is the live + app's; the rows are real and survive the call. + """ + sa_session = self._app.model.context + job = Job() + job.user_id = user_id + job.tool_id = "interactivetool_simple" + job.state = Job.states.RUNNING + sa_session.add(job) + sa_session.flush() + ep = InteractiveToolEntryPoint( + job=job, + tool_port=tool_port, + entry_url="/", + name="test entry point", + label="test", + requires_domain=True, + requires_path_in_url=False, + requires_path_in_header_named=None, + ) + sa_session.add(ep) + sa_session.commit() + return job.id, ep.id + + def test_entry_point_update_event_fires_on_configure(self): + """configure_entry_points should fire an ``entry_point_update`` wake-up event.""" + api_key = self.galaxy_interactor.api_key + assert api_key is not None + user_id = self._user_id_for_api_key(api_key) + job_id, _ = self._create_it_job_with_entry_point(user_id) + + listener = SSELineListener(self._events_stream_url(), api_key) + listener.start() + try: + sa_session = self._app.model.context + job = sa_session.get(Job, job_id) + assert job is not None + self._app.interactivetool_manager.configure_entry_points(job, _make_ports_dict(8888)) + + entry_point_events = listener.wait_for_event("entry_point_update") + finally: + listener.stop() + + # The event carries no payload: the client refetches ``/api/entry_points`` + # (the canonical source) on receipt, so the event just needs to arrive. + assert len(entry_point_events) >= 1, f"Expected entry_point_update wake-up, got: {entry_point_events}" + assert entry_point_events[0]["event"] == "entry_point_update" + + def test_entry_point_update_is_scoped_to_owning_user(self): + """User A must not see entry_point_update events for user B's jobs. + + The event has no payload to cross-check with, so we assert on event + count: user A's stream should receive exactly one event for its own + ``configure_entry_points`` call and none for user B's. + """ + user_b = self._setup_user(f"{uuid4()}@galaxy.test") + _, user_b_api_key = self._setup_user_get_key(user_b["email"]) + user_b_id = self._user_id_for_api_key(user_b_api_key) + + user_a_api_key = self.galaxy_interactor.api_key + assert user_a_api_key is not None + user_a_id = self._user_id_for_api_key(user_a_api_key) + + job_a_id, _ = self._create_it_job_with_entry_point(user_a_id, tool_port=7001) + job_b_id, _ = self._create_it_job_with_entry_point(user_b_id, tool_port=7002) + + listener = SSELineListener(self._events_stream_url(), user_a_api_key) + listener.start() + try: + sa_session = self._app.model.context + job_b = sa_session.get(Job, job_b_id) + assert job_b is not None + # User B's job — user A must NOT see this. + self._app.interactivetool_manager.configure_entry_points(job_b, _make_ports_dict(7002)) + # Give the broker a moment so a leaked event (if any) would arrive + # before we fire user A's event — the assertion would then catch + # more than one event on user A's stream. + time.sleep(0.5) + + job_a = sa_session.get(Job, job_a_id) + assert job_a is not None + # User A's own job — this is what A's stream must observe. + self._app.interactivetool_manager.configure_entry_points(job_a, _make_ports_dict(7001)) + + entry_point_events = listener.wait_for_event("entry_point_update") + finally: + listener.stop() + + assert ( + len(entry_point_events) == 1 + ), f"User A expected exactly one entry_point_update (own job); saw {len(entry_point_events)}: {entry_point_events}" diff --git a/test/integration_selenium/test_entry_point_sse.py b/test/integration_selenium/test_entry_point_sse.py new file mode 100644 index 00000000000..e661cb0a79f --- /dev/null +++ b/test/integration_selenium/test_entry_point_sse.py @@ -0,0 +1,107 @@ +"""Playwright E2E test for the interactive-tool entry-point SSE pipeline. + +Verifies that when an interactive-tool entry point transitions to ``configured`` +server-side (the runner's port-routing hook calls ``configure_entry_points``), +a logged-in user's browser receives the ``entry_point_update`` SSE event and +the entry-point store refetches, without the 10 s polling interval. + +This test stubs the server-side runner callback: it creates the Job and entry +point directly and calls ``InteractiveToolManager.configure_entry_points`` on +the live app. That invocation is the exact dispatch site in production. +""" + +from uuid import uuid4 + +from galaxy.util.wait import wait_on +from galaxy_test.selenium.framework import ( + managed_history, + selenium_test, +) +from .framework import SeleniumIntegrationTestCase + +SSE_CONNECT_TIMEOUT_SECONDS = 15 +SSE_EVENT_TIMEOUT_SECONDS = 15 + + +class TestEntryPointSSESeleniumIntegration(SeleniumIntegrationTestCase): + ensure_registered = True + + @classmethod + def handle_galaxy_config_kwds(cls, config): + super().handle_galaxy_config_kwds(config) + config["enable_celery_tasks"] = False + + def _wait_for_sse_connected(self) -> None: + """Block until the frontend confirms the SSE pipeline is live.""" + wait_on( + lambda: True if self.driver.execute_script("return window.__galaxy_sse_connected === true") else None, + "window.__galaxy_sse_connected === true", + timeout=SSE_CONNECT_TIMEOUT_SECONDS, + ) + + def _last_sse_event_ts(self) -> int: + return self.driver.execute_script("return window.__galaxy_sse_last_event_ts || 0") or 0 + + def _wait_for_sse_event_after(self, baseline_ts: int) -> None: + wait_on( + lambda: True if self._last_sse_event_ts() > baseline_ts else None, + "window.__galaxy_sse_last_event_ts advanced past baseline", + timeout=SSE_EVENT_TIMEOUT_SECONDS, + ) + + def _create_it_job_with_entry_point(self, tool_port: int = 8888) -> tuple[int, int]: + from galaxy.model import ( + InteractiveToolEntryPoint, + Job, + ) + + user_info = self._get("users/current").json() + user_id = self._app.security.decode_id(user_info["id"]) + sa_session = self._app.model.context + job = Job() + job.user_id = user_id + job.tool_id = "interactivetool_simple" + job.state = Job.states.RUNNING + sa_session.add(job) + sa_session.flush() + ep = InteractiveToolEntryPoint( + job=job, + tool_port=tool_port, + entry_url="/", + name=f"selenium entry {uuid4()}", + label="selenium", + requires_domain=True, + requires_path_in_url=False, + requires_path_in_header_named=None, + ) + sa_session.add(ep) + sa_session.commit() + return job.id, ep.id + + @selenium_test + @managed_history + def test_entry_point_update_pushed_via_sse(self): + """configure_entry_points should trigger an SSE push the client observes.""" + # Navigate home so the entry-point store is mounted and subscribed. + self.home() + self._wait_for_sse_connected() + baseline_ts = self._last_sse_event_ts() + self.screenshot("entry_point_sse_before") + + job_id, _ep_id = self._create_it_job_with_entry_point(tool_port=8888) + + # Stub the runner hook: call configure_entry_points on the live app. + from galaxy.model import Job + + sa_session = self._app.model.context + job = sa_session.get(Job, job_id) + assert job is not None + self._app.interactivetool_manager.configure_entry_points( + job, + {"8888": {"host": "host.invalid", "port": 12345, "protocol": "http"}}, + ) + + # Prove the update arrived via SSE (not polling): the composable's + # event-timestamp hook only advances when useSSE's listener fires. + self._wait_for_sse_event_after(baseline_ts) + self.screenshot("entry_point_sse_after") diff --git a/test/integration_selenium/test_notification_sse.py b/test/integration_selenium/test_notification_sse.py index 715898bbd2e..cae3b6e7b6b 100644 --- a/test/integration_selenium/test_notification_sse.py +++ b/test/integration_selenium/test_notification_sse.py @@ -129,7 +129,5 @@ class TestNotificationSSESeleniumIntegration(SeleniumIntegrationTestCase): self._wait_for_sse_event_after(baseline_ts) # The indicator dot should appear on the bell (within the #activity-notifications element) - self.wait_for_selector_visible( - "#activity-notifications .indicator", timeout=SSE_EVENT_TIMEOUT_SECONDS * 1000 - ) + self.wait_for_selector_visible("#activity-notifications .indicator", timeout=SSE_EVENT_TIMEOUT_SECONDS * 1000) self.screenshot("notification_bell_indicator") diff --git a/test/unit/app/managers/test_queue_metrics.py b/test/unit/app/managers/test_queue_metrics.py new file mode 100644 index 00000000000..fddaa576924 --- /dev/null +++ b/test/unit/app/managers/test_queue_metrics.py @@ -0,0 +1,243 @@ +"""Unit tests for :mod:`galaxy.webapps.galaxy.metrics.queue_metrics`. + +The DI container (``app``), the ``SSEConnectionManager``, and the kombu +connection are small hand-built fakes — no broker or database is required. +Assertions are on the recorded state of the statsd client (counters, timings) +rather than on mock call-lists so a regression that stops emitting a gauge +fails the test for the right reason. + +The failure-isolation test drives real sub-emitters into their error paths by +handing them genuinely broken collaborators (a connection whose ``clone()`` +raises, a model whose ``new_session()`` raises). That way the test exercises +the real ``_run`` wrapper rather than asserting a monkey-patched side_effect. +""" + +from dataclasses import ( + dataclass, + field, +) +from types import SimpleNamespace +from typing import ( + cast, + Optional, +) +from unittest.mock import MagicMock + +import pytest + +from galaxy.managers.sse import SSEConnectionManager +from galaxy.structured_app import StructuredApp +from galaxy.web.statsd_client import VanillaGalaxyStatsdClient +from galaxy.webapps.galaxy.metrics import queue_metrics + +# --------------------------------------------------------------------------- +# Fakes +# --------------------------------------------------------------------------- + + +@dataclass +class FakeStatsdClient: + """In-memory statsd recorder — see docstring in ``test_sse_dispatch.py``.""" + + counters: dict[tuple[str, tuple[tuple[str, str], ...]], int] = field(default_factory=dict) + timings: list[tuple[str, float, tuple[tuple[str, str], ...]]] = field(default_factory=list) + + def incr(self, metric: str, tags: Optional[dict[str, str]] = None) -> None: + key = (metric, tuple(sorted((tags or {}).items()))) + self.counters[key] = self.counters.get(key, 0) + 1 + + def timing(self, metric: str, value: float, tags: Optional[dict[str, str]] = None) -> None: + self.timings.append((metric, value, tuple(sorted((tags or {}).items())))) + + def counter(self, metric: str, tags: Optional[dict[str, str]] = None) -> int: + return self.counters.get((metric, tuple(sorted((tags or {}).items()))), 0) + + def timings_for(self, metric: str) -> list[tuple[float, dict[str, str]]]: + return [(v, dict(t)) for m, v, t in self.timings if m == metric] + + +class _ContainerApp: + """Tiny stand-in for ``StructuredApp`` + the Lagom container. + + Supports ``app[ClassName]`` lookup for ``SSEConnectionManager`` and + arbitrary attribute access. + """ + + def __init__(self, **attrs): + self._container: dict[type, object] = {} + for k, v in attrs.items(): + setattr(self, k, v) + + def register(self, cls, instance): + self._container[cls] = instance + + def __getitem__(self, cls): + return self._container[cls] + + +def _fake_sse_manager(broadcast: int = 3, per_user: int = 5): + m = MagicMock(spec=SSEConnectionManager) + m.total_broadcast_connections = broadcast + m.total_per_user_connections = per_user + return m + + +def _make_fake_queue(name: str, count: int): + """Fake kombu Queue exposing ``.bind(channel).queue_declare(passive=True).message_count``.""" + declared = SimpleNamespace(message_count=count) + bound = SimpleNamespace(queue_declare=lambda passive: declared) + return SimpleNamespace(name=name, bind=lambda channel: bound) + + +def _make_fake_connection(channel=None): + """Fake kombu Connection whose ``.clone()`` is usable as a context manager.""" + channel = channel or MagicMock() + conn_cm = MagicMock() + conn_cm.channel.return_value = channel + cm = MagicMock() + cm.__enter__ = lambda self: conn_cm + cm.__exit__ = lambda self, *a: False + connection = MagicMock() + connection.clone.return_value = cm + return connection + + +# --------------------------------------------------------------------------- +# Sub-emitter tests — assert on recorded state +# --------------------------------------------------------------------------- + + +def test_emit_sse_connection_gauges_emits_both_kinds(): + statsd = FakeStatsdClient() + queue_metrics.emit_sse_connection_gauges( + cast(VanillaGalaxyStatsdClient, statsd), _fake_sse_manager(broadcast=4, per_user=7) + ) + + assert statsd.timings_for("galaxy.sse.connections.active") == [ + (4, {"kind": "broadcast"}), + (7, {"kind": "per_user"}), + ] + + +def test_emit_control_queue_depth_emits_per_queue(monkeypatch): + statsd = FakeStatsdClient() + fake_queues = [ + _make_fake_queue("control.main@h", 3), + _make_fake_queue("control.main.1@h", 0), + ] + monkeypatch.setattr( + queue_metrics, + "all_control_queues_for_declare", + lambda application_stack: fake_queues, + ) + + app = _ContainerApp( + amqp_internal_connection_obj=_make_fake_connection(), + application_stack=MagicMock(), + ) + queue_metrics.emit_control_queue_depth(cast(VanillaGalaxyStatsdClient, statsd), cast(StructuredApp, app)) + + assert statsd.timings_for("galaxy.control_queue.depth") == [ + (3, {"queue_name": "control.main@h"}), + (0, {"queue_name": "control.main.1@h"}), + ] + + +def test_emit_control_queue_depth_skips_failed_passive_declare(monkeypatch): + """One bad queue → we skip it and keep going for the rest.""" + statsd = FakeStatsdClient() + + def bad_declare(passive): + raise RuntimeError("queue does not exist yet") + + good_queue = _make_fake_queue("control.good@h", 9) + bad_queue = SimpleNamespace( + name="control.bad@h", + bind=lambda channel: SimpleNamespace(queue_declare=bad_declare), + ) + monkeypatch.setattr(queue_metrics, "all_control_queues_for_declare", lambda stack: [good_queue, bad_queue]) + + app = _ContainerApp( + amqp_internal_connection_obj=_make_fake_connection(), + application_stack=MagicMock(), + ) + queue_metrics.emit_control_queue_depth(cast(VanillaGalaxyStatsdClient, statsd), cast(StructuredApp, app)) + + assert statsd.timings_for("galaxy.control_queue.depth") == [ + (9, {"queue_name": "control.good@h"}), + ] + + +def test_emit_control_queue_depth_no_broker_connection_is_noop(): + statsd = FakeStatsdClient() + app = _ContainerApp(amqp_internal_connection_obj=None, application_stack=MagicMock()) + queue_metrics.emit_control_queue_depth(cast(VanillaGalaxyStatsdClient, statsd), cast(StructuredApp, app)) + assert statsd.timings == [] + + +# --------------------------------------------------------------------------- +# emit_queue_metrics — aggregate entry-point +# --------------------------------------------------------------------------- + + +def test_emit_queue_metrics_is_silent_when_statsd_is_none(): + """No statsd client → every sub-call is skipped, no DB or broker access.""" + app = _ContainerApp( + execution_timer_factory=SimpleNamespace(galaxy_statsd_client=None), + ) + # Would raise AttributeError if the short-circuit didn't fire before any + # real collaborator was touched. + queue_metrics.emit_queue_metrics(cast(StructuredApp, app)) + + +def test_emit_queue_metrics_isolates_real_subemitter_failures(monkeypatch): + """When real sub-emitters raise, ``_run`` contains the failure and logs an error counter. + + We drive the failures through the actual sub-emitter bodies — not monkey- + patched side_effects — by handing in a connection whose ``.clone()`` raises + (for ``emit_control_queue_depth``) and a model whose ``.new_session()`` + raises (for ``emit_worker_process_gauge``). The SSE gauge is given healthy + collaborators and should still land. + """ + statsd = FakeStatsdClient() + sse_manager = _fake_sse_manager(broadcast=2, per_user=1) + + broken_connection = MagicMock() + broken_connection.clone.side_effect = RuntimeError("broker is gone") + + broken_model = MagicMock() + broken_model.new_session.side_effect = RuntimeError("db is gone") + + # Ensure the broker path reaches .clone() rather than short-circuiting on + # an empty queue list. + monkeypatch.setattr( + queue_metrics, + "all_control_queues_for_declare", + lambda stack: [_make_fake_queue("control.main@h", 0)], + ) + + app = _ContainerApp( + execution_timer_factory=SimpleNamespace(galaxy_statsd_client=statsd), + amqp_internal_connection_obj=broken_connection, + application_stack=MagicMock(), + model=broken_model, + ) + app.register(SSEConnectionManager, sse_manager) + + # Must not raise — the SSE gauge still lands. + queue_metrics.emit_queue_metrics(cast(StructuredApp, app)) + + # SSE gauge landed despite the other two failing. + sse_timings = statsd.timings_for("galaxy.sse.connections.active") + assert (2, {"kind": "broadcast"}) in sse_timings + assert (1, {"kind": "per_user"}) in sse_timings + + # Each failing sub-emitter bumped its error counter tagged by name. + assert statsd.counter("galaxy.queue_metrics.error", {"emitter": "control_queue_depth"}) == 1 + assert statsd.counter("galaxy.queue_metrics.error", {"emitter": "worker_process"}) == 1 + # The healthy SSE sub-emitter did not. + assert statsd.counter("galaxy.queue_metrics.error", {"emitter": "sse_connections"}) == 0 + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/test/unit/app/managers/test_sse_dispatch.py b/test/unit/app/managers/test_sse_dispatch.py new file mode 100644 index 00000000000..08da83632ed --- /dev/null +++ b/test/unit/app/managers/test_sse_dispatch.py @@ -0,0 +1,226 @@ +"""Unit tests for :class:`galaxy.managers.sse_dispatch.SSEEventDispatcher` observability. + +Focus is on the statsd instrumentation contract AND the effect of dispatch: +counters/timers fire on the happy path and the ``_queue_worker is None`` +early-return, payloads reach the broker with the expected task+kwargs, and the +dispatcher is a silent no-op when ``statsd_client`` is ``None``. + +These tests use lightweight fakes (``FakeStatsdClient``, ``FakeControlTask``) +that record state we can assert against — rather than ``MagicMock`` call-lists — +so a regression that silently drops dispatch or stops recording metrics fails +the test for the right reason. +""" + +from dataclasses import ( + dataclass, + field, +) +from typing import ( + Any, + Optional, +) +from unittest.mock import MagicMock + +import pytest + +from galaxy.managers.sse_dispatch import SSEEventDispatcher + + +@dataclass +class FakeStatsdClient: + """In-memory stand-in for ``VanillaGalaxyStatsdClient``. + + Records ``incr`` and ``timing`` calls as plain data so tests assert on + observable state (counter totals, recorded timings) instead of mock + call-lists. + """ + + counters: dict[tuple[str, tuple[tuple[str, str], ...]], int] = field(default_factory=dict) + timings: list[tuple[str, float, tuple[tuple[str, str], ...]]] = field(default_factory=list) + + def incr(self, metric: str, tags: Optional[dict[str, str]] = None) -> None: + key = (metric, tuple(sorted((tags or {}).items()))) + self.counters[key] = self.counters.get(key, 0) + 1 + + def timing(self, metric: str, value: float, tags: Optional[dict[str, str]] = None) -> None: + self.timings.append((metric, value, tuple(sorted((tags or {}).items())))) + + def counter(self, metric: str, tags: Optional[dict[str, str]] = None) -> int: + return self.counters.get((metric, tuple(sorted((tags or {}).items()))), 0) + + +@dataclass +class RecordedTask: + payload: dict[str, Any] + routing_key: str + expiration: Optional[int] + declare_queues: Any + + +class FakeControlTask: + """Stand-in for ``ControlTask`` that records dispatches instead of touching AMQP.""" + + instances: list["FakeControlTask"] = [] + + def __init__(self, queue_worker) -> None: + self.queue_worker = queue_worker + self.sent: list[RecordedTask] = [] + FakeControlTask.instances.append(self) + + def send_task( + self, + payload: dict[str, Any], + routing_key: str, + expiration: Optional[int] = None, + declare_queues: Any = None, + **_: Any, + ) -> None: + self.sent.append( + RecordedTask( + payload=payload, + routing_key=routing_key, + expiration=expiration, + declare_queues=declare_queues, + ) + ) + + +class BoomControlTask: + """``ControlTask`` fake whose ``send_task`` always raises — exercises the finally block.""" + + def __init__(self, queue_worker) -> None: + self.queue_worker = queue_worker + + def send_task(self, **kwargs) -> None: + raise RuntimeError("broker down") + + +@pytest.fixture +def application_stack(): + return MagicMock(name="ApplicationStack") + + +@pytest.fixture +def queue_worker(): + return MagicMock(name="GalaxyQueueWorker") + + +@pytest.fixture +def statsd() -> FakeStatsdClient: + return FakeStatsdClient() + + +@pytest.fixture(autouse=True) +def _reset_fake_control_task_instances(): + FakeControlTask.instances.clear() + yield + FakeControlTask.instances.clear() + + +@pytest.fixture +def fake_control_task(monkeypatch) -> type[FakeControlTask]: + monkeypatch.setattr("galaxy.managers.sse_dispatch.ControlTask", FakeControlTask) + monkeypatch.setattr( + "galaxy.managers.sse_dispatch.all_control_queues_for_declare", + lambda *args, **kwargs: [], + ) + return FakeControlTask + + +def test_dispatcher_no_op_when_queue_worker_is_none_and_no_statsd(application_stack): + """No statsd client set and no queue_worker → silent no-op, no AttributeError.""" + dispatcher = SSEEventDispatcher( + queue_worker=None, + application_stack=application_stack, + statsd_client=None, + ) + # Must not raise, must not attempt to declare queues. + dispatcher.notify_users([1, 2], "hello") + dispatcher.notify_broadcast("hi") + dispatcher.history_update({"1": [42]}) + + +def test_dispatcher_records_skipped_counter_when_queue_worker_is_none(application_stack, statsd): + """Two dispatches with no queue_worker → two skipped_no_qw increments, no timings.""" + dispatcher = SSEEventDispatcher( + queue_worker=None, + application_stack=application_stack, + statsd_client=statsd, + ) + dispatcher.notify_users([1], "hello") + dispatcher.notify_broadcast("world") + + assert statsd.counter("galaxy.sse.dispatch.skipped_no_qw") == 2 + # No latency timing — we never got as far as the broker call. + assert statsd.timings == [] + + +def test_dispatcher_enqueues_payload_and_records_metrics_on_send( + application_stack, queue_worker, statsd, fake_control_task +): + """Happy path: payload reaches the broker AND counter+timer are recorded.""" + dispatcher = SSEEventDispatcher( + queue_worker=queue_worker, + application_stack=application_stack, + statsd_client=statsd, + ) + dispatcher.notify_users([1, 2], "hello") + + # Exactly one ControlTask constructed, one send_task recorded with the right + # payload — asserts the dispatch *effect*, not just that a mock was called. + assert len(fake_control_task.instances) == 1 + sent = fake_control_task.instances[0].sent + assert len(sent) == 1 + assert sent[0].payload["task"] == "notify_users" + assert sent[0].payload["kwargs"]["user_ids"] == [1, 2] + assert sent[0].payload["kwargs"]["payload"] == "hello" + assert "event_id" in sent[0].payload["kwargs"] + assert sent[0].routing_key == "control.*" + assert sent[0].expiration == 10 + + # Counter + timer both recorded with matching task tag. + assert statsd.counter("galaxy.sse.dispatch.count", {"task": "notify_users"}) == 1 + assert len(statsd.timings) == 1 + metric, _value, tags = statsd.timings[0] + assert metric == "galaxy.sse.dispatch.latency_ms" + assert dict(tags) == {"task": "notify_users"} + + +def test_dispatcher_timer_still_fires_on_send_exception(monkeypatch, application_stack, queue_worker, statsd): + """Timer lives in ``finally`` — broker errors don't mask the latency metric.""" + monkeypatch.setattr( + "galaxy.managers.sse_dispatch.all_control_queues_for_declare", + lambda *args, **kwargs: [], + ) + monkeypatch.setattr("galaxy.managers.sse_dispatch.ControlTask", BoomControlTask) + + dispatcher = SSEEventDispatcher( + queue_worker=queue_worker, + application_stack=application_stack, + statsd_client=statsd, + ) + with pytest.raises(RuntimeError): + dispatcher.history_update({"7": [1]}) + + assert statsd.counter("galaxy.sse.dispatch.count", {"task": "history_update"}) == 1 + assert len(statsd.timings) == 1 + metric, _value, tags = statsd.timings[0] + assert metric == "galaxy.sse.dispatch.latency_ms" + assert dict(tags) == {"task": "history_update"} + + +def test_dispatcher_no_statsd_means_no_instrumentation(application_stack, queue_worker, fake_control_task): + """When ``statsd_client`` is ``None`` instrumentation is bypassed entirely. + + The dispatch still happens — we assert via the ControlTask fake — but there + is nothing to observe on the (absent) statsd side. + """ + dispatcher = SSEEventDispatcher( + queue_worker=queue_worker, + application_stack=application_stack, + statsd_client=None, + ) + dispatcher.notify_broadcast("hi") + assert len(fake_control_task.instances) == 1 + assert len(fake_control_task.instances[0].sent) == 1 + assert fake_control_task.instances[0].sent[0].payload["task"] == "notify_broadcast" diff --git a/test/unit/app/managers/test_sse_dispatch_cache.py b/test/unit/app/managers/test_sse_dispatch_cache.py new file mode 100644 index 00000000000..e7d1a7e9ad4 --- /dev/null +++ b/test/unit/app/managers/test_sse_dispatch_cache.py @@ -0,0 +1,140 @@ +"""Tests for the TTL cache on ``SSEEventDispatcher._get_declare_queues``. + +The cache exists because ``_send`` is on a hot path (1000+ events/s at target +load) and without it each dispatch fires a ``WorkerProcess`` DB query. The +underlying data only changes on a 60 s heartbeat cadence, so a 30 s TTL is safe. +""" + +from concurrent.futures import ( + as_completed, + ThreadPoolExecutor, +) +from typing import Any +from unittest.mock import MagicMock + +import pytest + +from galaxy.managers import sse_dispatch +from galaxy.managers.sse_dispatch import SSEEventDispatcher + + +@pytest.fixture +def fake_declare(monkeypatch): + """Replace ``all_control_queues_for_declare`` with a call-counting fake. + + Returns a non-empty list by default so the cache stores a value — empty + results are intentionally not cached (see ``test_empty_result_not_cached``). + """ + calls: dict[str, Any] = {"count": 0, "returns": [MagicMock(name="queue")]} + + def _fake(application_stack, webapp_only=False): + calls["count"] += 1 + # Sanity: the dispatcher must always ask for webapp-only queues. + assert webapp_only is True + return calls["returns"] + + monkeypatch.setattr(sse_dispatch, "all_control_queues_for_declare", _fake) + return calls + + +class FakeClock: + """Controllable time source passed to ``SSEEventDispatcher``. + + Lets tests advance the dispatcher's TTL cache deterministically without + reaching into ``cachetools`` internals. + """ + + def __init__(self, start: float = 0.0) -> None: + self.now = start + + def __call__(self) -> float: + return self.now + + def advance(self, seconds: float) -> None: + self.now += seconds + + +@pytest.fixture +def clock() -> FakeClock: + return FakeClock() + + +@pytest.fixture +def dispatcher(monkeypatch, clock): + """Build a dispatcher with a stub queue_worker and stub application_stack. + + ``ControlTask`` is swapped for a no-op so ``_send`` doesn't try to open a + real AMQP connection. + """ + queue_worker = MagicMock(name="queue_worker") + application_stack = MagicMock(name="application_stack") + + fake_control_task = MagicMock() + fake_control_task.return_value.send_task = MagicMock() + monkeypatch.setattr(sse_dispatch, "ControlTask", fake_control_task) + + return SSEEventDispatcher(queue_worker=queue_worker, application_stack=application_stack, clock=clock) + + +def test_declare_queues_cached_within_ttl(dispatcher, fake_declare): + """Repeated dispatches inside the TTL window only hit the DB once.""" + for _ in range(10): + dispatcher.notify_broadcast("payload") + assert fake_declare["count"] == 1 + + +def test_declare_queues_refetched_after_ttl(dispatcher, fake_declare, clock): + """Once the TTL expires, the next call refetches exactly once.""" + dispatcher.notify_broadcast("payload") + assert fake_declare["count"] == 1 + + # Advance the injected clock past the TTL so the cache sees the entry as + # expired on the next read. + clock.advance(dispatcher._DECLARE_QUEUES_TTL_SECONDS + 1) + + dispatcher.notify_broadcast("payload") + assert fake_declare["count"] == 2 + + # Further dispatches at the advanced time reuse the newly populated entry. + dispatcher.notify_broadcast("payload") + assert fake_declare["count"] == 2 + + +def test_empty_result_not_cached(dispatcher, fake_declare): + """An empty list must not be pinned in the cache for the full TTL. + + Empty results arise during startup (before ``DatabaseHeartbeat`` writes the + row) and on swallowed DB errors. Caching them would silently drop every SSE + event until the next TTL expiry. + """ + fake_declare["returns"] = [] + dispatcher.notify_broadcast("payload") + dispatcher.notify_broadcast("payload") + assert fake_declare["count"] == 2 + + # Once the upstream starts returning a non-empty result, caching resumes. + fake_declare["returns"] = [MagicMock(name="queue")] + dispatcher.notify_broadcast("payload") + dispatcher.notify_broadcast("payload") + assert fake_declare["count"] == 3 + + +def test_declare_queues_thread_safe_single_query_under_load(dispatcher, fake_declare): + """Concurrent ``_send`` from many threads still only triggers one DB query. + + With stampede protection (RLock around the miss) all 500 dispatches should + collapse to a single ``all_control_queues_for_declare`` call inside one TTL + window. The assertion is exact (== 1), not loose, because the lock + serializes the miss. + """ + iterations = 500 + + def work(): + dispatcher.notify_broadcast("payload") + + with ThreadPoolExecutor(max_workers=16) as pool: + futures = [pool.submit(work) for _ in range(iterations)] + for future in as_completed(futures): + future.result() + + assert fake_declare["count"] == 1 diff --git a/test/unit/app/queue_worker/test_queue_worker.py b/test/unit/app/queue_worker/test_queue_worker.py index 09bc8a32fa5..857bdeb2cac 100644 --- a/test/unit/app/queue_worker/test_queue_worker.py +++ b/test/unit/app/queue_worker/test_queue_worker.py @@ -1,6 +1,13 @@ import datetime import time +from dataclasses import ( + dataclass, + field, +) from math import inf +from types import SimpleNamespace +from typing import Optional +from unittest.mock import MagicMock import pytest @@ -14,6 +21,27 @@ from galaxy.queues import connection_from_config from galaxy.web_stack import application_stack_instance +@dataclass +class FakeStatsdClient: + """In-memory statsd recorder — see docstring in ``test_sse_dispatch.py``.""" + + counters: dict[tuple[str, tuple[tuple[str, str], ...]], int] = field(default_factory=dict) + timings: list[tuple[str, float, tuple[tuple[str, str], ...]]] = field(default_factory=list) + + def incr(self, metric: str, tags: Optional[dict[str, str]] = None) -> None: + key = (metric, tuple(sorted((tags or {}).items()))) + self.counters[key] = self.counters.get(key, 0) + 1 + + def timing(self, metric: str, value: float, tags: Optional[dict[str, str]] = None) -> None: + self.timings.append((metric, value, tuple(sorted((tags or {}).items())))) + + def counter(self, metric: str, tags: Optional[dict[str, str]] = None) -> int: + return self.counters.get((metric, tuple(sorted((tags or {}).items()))), 0) + + def timings_for(self, metric: str) -> list[tuple[float, dict[str, str]]]: + return [(v, dict(t)) for m, v, t in self.timings if m == metric] + + def bar(app, **kwargs): app.some_var = "bar" app.tasks_executed.append("echo") @@ -119,3 +147,86 @@ def wait_for_var(obj, var, value, tries=10, sleep=0.25): tries -= 1 time.sleep(sleep) assert getattr(obj, var) == value + + +# --------------------------------------------------------------------------- +# process_task observability tests +# +# These don't need a real broker or DB — we just instantiate GalaxyQueueWorker +# via ``__new__`` and drive ``process_task`` directly with fake ``body`` / +# ``message`` objects. The statsd client is pulled through +# ``app.execution_timer_factory.galaxy_statsd_client``. +# --------------------------------------------------------------------------- + + +def _make_fake_worker(task_fn, statsd_client): + worker = GalaxyQueueWorker.__new__(GalaxyQueueWorker) + worker.app = SimpleNamespace( + config=SimpleNamespace(server_name="test.server"), + execution_timer_factory=SimpleNamespace(galaxy_statsd_client=statsd_client), + ) + worker.task_mapping = {"echo": task_fn} + worker.epoch = 0 + # ``producer`` is a read-only property on the mixin; we avoid the publisher + # path entirely by leaving ``reply_to`` out of the fake message properties. + return worker + + +def _fake_message(): + message = MagicMock() + message.headers = {"epoch": inf} # always greater than worker.epoch + message.properties = {} # no reply_to — skip publisher path + return message + + +def test_process_task_emits_counter_and_ok_timer(): + statsd = FakeStatsdClient() + calls: list[dict] = [] + + def handler(app, **kwargs): + calls.append(kwargs) + return "done" + + worker = _make_fake_worker(handler, statsd) + worker.process_task({"task": "echo", "kwargs": {"x": 1}}, _fake_message()) + + # Handler actually ran — if the worker silently dropped the task, this list + # would be empty and the test would fail loudly. + assert calls == [{"x": 1}] + assert statsd.counter("galaxy.control_queue.task.count", {"task": "echo"}) == 1 + assert len(statsd.timings_for("galaxy.control_queue.task.latency_ms")) == 1 + _value, tags = statsd.timings_for("galaxy.control_queue.task.latency_ms")[0] + assert tags == {"task": "echo", "outcome": "ok"} + + +def test_process_task_emits_error_timer_on_handler_exception(): + statsd = FakeStatsdClient() + invocations: list[bool] = [] + + def boom(app, **kwargs): + invocations.append(True) + raise RuntimeError("handler failed") + + worker = _make_fake_worker(boom, statsd) + # process_task swallows handler exceptions (logged, not raised). + worker.process_task({"task": "echo", "kwargs": {}}, _fake_message()) + + # Handler was actually invoked before raising. + assert invocations == [True] + assert statsd.counter("galaxy.control_queue.task.count", {"task": "echo"}) == 1 + assert len(statsd.timings_for("galaxy.control_queue.task.latency_ms")) == 1 + _value, tags = statsd.timings_for("galaxy.control_queue.task.latency_ms")[0] + assert tags == {"task": "echo", "outcome": "error"} + + +def test_process_task_no_statsd_is_silent_no_op(): + calls: list[dict] = [] + + def handler(app, **kwargs): + calls.append(kwargs) + return "done" + + worker = _make_fake_worker(handler, statsd_client=None) + worker.process_task({"task": "echo", "kwargs": {"x": 2}}, _fake_message()) + # Handler ran — the test guards both "no exception" AND "task not silently dropped". + assert calls == [{"x": 2}]