mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
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:
@@ -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 };
|
||||
}),
|
||||
};
|
||||
}
|
||||
|
||||
@@ -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", () => {
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -59,8 +59,8 @@ from galaxy.schema.notifications import (
|
||||
NotificationBroadcastUpdateRequest,
|
||||
NotificationCategorySettings,
|
||||
NotificationChannelSettings,
|
||||
NotificationCreatedResponse,
|
||||
NotificationCreateData,
|
||||
NotificationCreatedResponse,
|
||||
NotificationCreateRequest,
|
||||
NotificationRecipients,
|
||||
NotificationResponse,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"])
|
||||
@@ -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}]
|
||||
|
||||
Reference in New Issue
Block a user