Add SSE entry-point channel, dispatch observability, declare-queue cache

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.
This commit is contained in:
mvdbeek
2026-04-28 17:18:07 +02:00
parent 5ae948c1df
commit 6e5eada944
37 changed files with 1528 additions and 62 deletions
@@ -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];
@@ -20,14 +20,20 @@ export interface SSEMockState {
onEvent: ((event: MessageEvent) => void) | null;
connect: ReturnType<typeof vi.fn>;
disconnect: ReturnType<typeof vi.fn>;
connected?: Ref<boolean>;
}
/** 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 };
}),
};
}
+14 -1
View File
@@ -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", () => {
+81 -16
View File
@@ -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<EntryPoint[]>([]);
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;
+26
View File
@@ -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``
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
+7 -1
View File
@@ -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)
+2 -1
View File
@@ -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
+1
View File
@@ -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",
+3
View File
@@ -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)
+12 -1
View File
@@ -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"""
@@ -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.
@@ -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
+24 -1
View File
@@ -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:
+1 -1
View File
@@ -59,8 +59,8 @@ from galaxy.schema.notifications import (
NotificationBroadcastUpdateRequest,
NotificationCategorySettings,
NotificationChannelSettings,
NotificationCreatedResponse,
NotificationCreateData,
NotificationCreatedResponse,
NotificationCreateRequest,
NotificationRecipients,
NotificationResponse,
+20 -3
View File
@@ -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(
+68 -7
View File
@@ -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(),
},
)
@@ -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):
+81 -11
View File
@@ -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:
+1
View File
@@ -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]
+1 -1
View File
@@ -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,
@@ -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))
@@ -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,
@@ -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__)
@@ -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,
@@ -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,
@@ -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,
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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,
)
+2 -4
View File
@@ -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:
+158
View File
@@ -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}"
@@ -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")
@@ -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")
@@ -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"])
+226
View File
@@ -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"
@@ -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
@@ -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}]