From 4a043eb80f601ff060dae67b67d5c755bf78afac Mon Sep 17 00:00:00 2001 From: Sam Hong <20244918+hongsamL@users.noreply.github.com> Date: Mon, 3 Aug 2026 12:45:03 +0200 Subject: [PATCH 1/9] Fix MDS legend overlapping plot on first render --- src/treetracer/callbacks/_helpers.py | 27 ++++++++++++++++++++++++++ src/treetracer/callbacks/treespace.py | 4 ++++ src/treetracer/callbacks/within_run.py | 3 +++ src/treetracer/ui/panels/treespace.py | 7 ++++++- src/treetracer/ui/panels/within_run.py | 6 +++++- 5 files changed, 45 insertions(+), 2 deletions(-) diff --git a/src/treetracer/callbacks/_helpers.py b/src/treetracer/callbacks/_helpers.py index 835909f..274c486 100644 --- a/src/treetracer/callbacks/_helpers.py +++ b/src/treetracer/callbacks/_helpers.py @@ -78,6 +78,33 @@ def file_types_for_filename(filename): return known.get(ext, ("All files (*.*)",)) +def register_autorelayout(graph_id): + """Force one Plotly relayout right after ``graph_id``'s figure paints. + + In pywebview, Plotly's first layout pass can skip legend *automargin*, so a + multi-row horizontal top legend (many files with long names) overlaps the + plot until a manual window resize forces a relayout. This fires that + relayout programmatically — the same thing the drag does — so the legend + lands correctly on first render. ``config.responsive`` doesn't help here + because the container is a stable viewport-height, so nothing triggers a + resize on its own. + """ + from dash import clientside_callback, Input, Output + clientside_callback( + "function(fig){" + " if(fig){setTimeout(function(){" + f" var r=document.getElementById('{graph_id}');" + " var gd=r&&r.querySelector('.js-plotly-plot');" + " if(gd&&window.Plotly){window.Plotly.Plots.resize(gd);}" + " },50);}" + " return window.dash_clientside.no_update;" + "}", + Output(graph_id, "style", allow_duplicate=True), + Input(graph_id, "figure"), + prevent_initial_call=True, + ) + + def _save_file_dialog(default_filename="output.tsv", file_types=None): """Open a native save-file dialog and return the chosen path, or ``None`` if the user cancelled / no dialog could be shown. diff --git a/src/treetracer/callbacks/treespace.py b/src/treetracer/callbacks/treespace.py index 7e46fbf..37ade29 100644 --- a/src/treetracer/callbacks/treespace.py +++ b/src/treetracer/callbacks/treespace.py @@ -197,6 +197,10 @@ def _patch_overlay_bundle(patch, n_traces, panel_data, offsets): def register_treespace_callbacks(): + # Fix the first-paint legend overlap (see _helpers.register_autorelayout). + from ._helpers import register_autorelayout + register_autorelayout("graph") + # Populate the MDS-result selector dropdown @callback( Output("treespace-result-select", "data"), diff --git a/src/treetracer/callbacks/within_run.py b/src/treetracer/callbacks/within_run.py index 2022775..e43b2d5 100644 --- a/src/treetracer/callbacks/within_run.py +++ b/src/treetracer/callbacks/within_run.py @@ -390,6 +390,9 @@ def _make_within_run_figure(df, x, y, z, show_lines=True, def register_within_run_callbacks(): + # Fix the first-paint legend overlap (see _helpers.register_autorelayout). + from ._helpers import register_autorelayout + register_autorelayout("within-run-graph") # ------ selectors ------ @callback( diff --git a/src/treetracer/ui/panels/treespace.py b/src/treetracer/ui/panels/treespace.py index 91305c1..204e04d 100644 --- a/src/treetracer/ui/panels/treespace.py +++ b/src/treetracer/ui/panels/treespace.py @@ -218,7 +218,12 @@ def _add_treespace_panel(): # Hide Plotly's modebar — its toolbar overlaps the legend # when many files are loaded, and the zoom/pan/export # interactions are already exposed via the tab's buttons. - config={"doubleClick": False, "displayModeBar": False}, + # ``responsive`` relayouts once the container reaches its real + # size — without it, the first paint lays out the figure at a + # transitional size and the top legend (wrapped wide by long + # file names) spills over the plot until the user resizes. + config={"doubleClick": False, "displayModeBar": False, + "responsive": True}, ), ], gap="xs"), ], style={"padding": "10px", "position": "relative"}) diff --git a/src/treetracer/ui/panels/within_run.py b/src/treetracer/ui/panels/within_run.py index dac5978..b5395a0 100644 --- a/src/treetracer/ui/panels/within_run.py +++ b/src/treetracer/ui/panels/within_run.py @@ -229,7 +229,11 @@ def _add_within_run_panel(): # Hide Plotly's modebar — its toolbar overlaps the legend # when many files are loaded, and the zoom/pan/export # interactions are already exposed via the tab's buttons. - config={"doubleClick": False, "displayModeBar": False}, + # ``responsive`` relayouts once the container reaches its real + # size, so the top legend lands correctly on first paint + # instead of only after a manual window resize. + config={"doubleClick": False, "displayModeBar": False, + "responsive": True}, ), ], gap="xs"), ], style={"padding": "10px", "position": "relative"}) From 1ace067b211cd09022b6b499aae6874b5f66c1dc Mon Sep 17 00:00:00 2001 From: Sam Hong <20244918+hongsamL@users.noreply.github.com> Date: Mon, 3 Aug 2026 14:00:35 +0200 Subject: [PATCH 2/9] background job manager implementation --- src/test/test_background_jobs.py | 309 +++++++++++++++ src/treetracer/background_jobs.py | 639 ++++++++++++++++++++++++++++++ 2 files changed, 948 insertions(+) create mode 100644 src/test/test_background_jobs.py create mode 100644 src/treetracer/background_jobs.py diff --git a/src/test/test_background_jobs.py b/src/test/test_background_jobs.py new file mode 100644 index 0000000..2aa73d5 --- /dev/null +++ b/src/test/test_background_jobs.py @@ -0,0 +1,309 @@ +"""Correctness tests for sticky background-job lifecycle state.""" + +from __future__ import annotations + +import threading +import time +from concurrent.futures import ThreadPoolExecutor +from dataclasses import replace + +import pytest + +from treetracer.background_jobs import ( + JobBusyError, + JobManager, + JobState, +) + + +def _wait_for_state(manager, ref, expected, timeout=2.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + snapshot = manager.snapshot(ref) + if snapshot is not None and snapshot.state is expected: + return snapshot + time.sleep(0.005) + snapshot = manager.snapshot(ref) + pytest.fail( + f"job did not reach {expected.value}; " + f"last state={None if snapshot is None else snapshot.state.value}" + ) + + +def test_success_is_sticky_until_matching_acknowledgement(): + manager = JobManager() + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "rf", + lambda: {"matrix": "large-result"}, + metadata={"display_name": "RF_001"}, + finalizer=lambda _ref, result: { + "result_ref": "RF_001", + "source": result["matrix"], + }, + ) + terminal = _wait_for_state(manager, ref, JobState.SUCCEEDED) + + assert terminal.progress.fraction == 1.0 + assert terminal.progress.phase == "done" + assert terminal.terminal.payload["result_ref"] == "RF_001" + assert terminal.metadata == {"display_name": "RF_001"} + assert manager.active_ref() == ref + + first = manager.snapshot_for_delivery(ref) + second = manager.snapshot_for_delivery(ref) + assert first.terminal == second.terminal + assert first.delivery_attempt == 1 + assert second.delivery_attempt == 2 + assert manager.snapshot(ref).terminal == terminal.terminal + + assert manager.acknowledge(ref, terminal.terminal.revision) is True + assert manager.snapshot(ref).acknowledged is True + assert manager.active_ref() is None + + +def test_stale_generation_cannot_read_update_or_acknowledge_job(): + manager = JobManager() + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit(executor, "mds", lambda: 1) + terminal = _wait_for_state(manager, ref, JobState.SUCCEEDED) + + stale = replace(ref, generation=ref.generation + 1) + assert manager.snapshot(stale) is None + assert manager.update_progress(stale, 0.5, "wrong job") is False + assert manager.acknowledge(stale, terminal.terminal.revision) is False + assert manager.snapshot(ref).acknowledged is False + + +def test_finalizer_runs_once_while_concurrent_readers_replay_terminal(): + manager = JobManager() + finalizer_entered = threading.Event() + release_finalizer = threading.Event() + calls = 0 + calls_lock = threading.Lock() + + def finalize(_ref, result): + nonlocal calls + with calls_lock: + calls += 1 + finalizer_entered.set() + assert release_finalizer.wait(timeout=2) + return {"result_ref": result} + + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "pseudo_ess", + lambda: "ESS_001", + finalizer=finalize, + ) + assert finalizer_entered.wait(timeout=2) + assert manager.snapshot(ref).state is JobState.FINALIZING + release_finalizer.set() + _wait_for_state(manager, ref, JobState.SUCCEEDED) + + barrier = threading.Barrier(9) + snapshots = [] + snapshots_lock = threading.Lock() + + def read_terminal(): + barrier.wait(timeout=2) + value = manager.snapshot_for_delivery(ref) + with snapshots_lock: + snapshots.append(value) + + threads = [threading.Thread(target=read_terminal) for _ in range(8)] + for thread in threads: + thread.start() + barrier.wait(timeout=2) + for thread in threads: + thread.join(timeout=2) + + assert calls == 1 + assert len(snapshots) == 8 + assert {s.terminal.payload["result_ref"] for s in snapshots} == {"ESS_001"} + assert sorted(s.delivery_attempt for s in snapshots) == list(range(1, 9)) + + +def test_compute_exception_becomes_sticky_failure(): + manager = JobManager() + + def fail(): + raise RuntimeError("worker broke") + + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit(executor, "mds", fail) + terminal = _wait_for_state(manager, ref, JobState.FAILED) + + assert terminal.terminal.payload == { + "message": "worker broke", + "error_type": "RuntimeError", + "stage": "compute", + } + assert manager.snapshot_for_delivery(ref).state is JobState.FAILED + assert manager.snapshot_for_delivery(ref).state is JobState.FAILED + + +def test_configured_worker_cancellation_exception_is_not_an_error(): + class WorkerCancelled(RuntimeError): + pass + + manager = JobManager() + + def cancel(): + raise WorkerCancelled("stopped by user") + + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "consensus", + cancel, + cancel_exceptions=(WorkerCancelled,), + ) + terminal = _wait_for_state(manager, ref, JobState.CANCELLED) + + assert terminal.terminal.payload["message"] == "stopped by user" + assert terminal.terminal.payload["error_type"] == "WorkerCancelled" + + +def test_finalizer_exception_is_terminal_failure_and_not_retried(): + manager = JobManager() + calls = 0 + + def broken_finalizer(_ref, _result): + nonlocal calls + calls += 1 + raise ValueError("could not register result") + + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "rf", + lambda: object(), + finalizer=broken_finalizer, + ) + terminal = _wait_for_state(manager, ref, JobState.FAILED) + + for _ in range(5): + manager.snapshot_for_delivery(ref) + assert calls == 1 + assert terminal.terminal.payload["stage"] == "finalize" + assert terminal.terminal.payload["error_type"] == "ValueError" + + +def test_explicit_cancellation_cannot_be_overwritten_by_late_success(): + manager = JobManager() + task_started = threading.Event() + release_task = threading.Event() + + def work(): + task_started.set() + assert release_task.wait(timeout=2) + return "too late" + + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit(executor, "rf", work) + assert task_started.wait(timeout=2) + assert manager.mark_cancelled(ref, message="reset by user") is True + cancelled = manager.snapshot(ref) + assert cancelled.state is JobState.CANCELLED + release_task.set() + + after_completion = manager.snapshot(ref) + assert after_completion.state is JobState.CANCELLED + assert after_completion.terminal.payload["message"] == "reset by user" + + +def test_single_active_gate_opens_only_after_terminal_acknowledgement(): + manager = JobManager(single_active=True) + task_started = threading.Event() + release_task = threading.Event() + + def work(): + task_started.set() + assert release_task.wait(timeout=2) + + with ThreadPoolExecutor(max_workers=1) as executor: + first = manager.submit(executor, "rf", work) + assert task_started.wait(timeout=2) + with pytest.raises(JobBusyError) as exc_info: + manager.submit(executor, "mds", lambda: None) + assert exc_info.value.active == first + + release_task.set() + terminal = _wait_for_state(manager, first, JobState.SUCCEEDED) + with pytest.raises(JobBusyError): + manager.submit(executor, "mds", lambda: None) + + assert manager.acknowledge(first, terminal.terminal.revision) + second = manager.submit(executor, "mds", lambda: None) + _wait_for_state(manager, second, JobState.SUCCEEDED) + + +def test_progress_is_monotonic_clamped_and_terminal_is_authoritative(): + manager = JobManager() + task_started = threading.Event() + release_task = threading.Event() + + def work(): + task_started.set() + assert release_task.wait(timeout=2) + + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit(executor, "mds", work) + assert task_started.wait(timeout=2) + assert manager.update_progress(ref, 0.7, "eigensolve", "70%") + assert manager.update_progress(ref, 0.4, "stale sidecar", "40%") + assert manager.snapshot(ref).progress.fraction == 0.7 + assert manager.update_progress(ref, 4.0, "finalizing") + assert manager.snapshot(ref).progress.fraction == 1.0 + with pytest.raises(ValueError): + manager.update_progress(ref, float("nan"), "invalid") + + release_task.set() + terminal = _wait_for_state(manager, ref, JobState.SUCCEEDED) + + assert terminal.progress.fraction == 1.0 + assert terminal.progress.phase == "done" + assert manager.update_progress(ref, 0.5, "late sidecar") is False + + +def test_terminal_payload_and_metadata_snapshots_are_defensive_copies(): + manager = JobManager() + metadata = {"nested": {"value": 1}} + payload = {"nested": {"value": 2}} + + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "rf", + lambda: None, + metadata=metadata, + finalizer=lambda _ref, _result: payload, + ) + _wait_for_state(manager, ref, JobState.SUCCEEDED) + + metadata["nested"]["value"] = 100 + payload["nested"]["value"] = 200 + first = manager.snapshot(ref) + first.metadata["nested"]["value"] = 300 + first.terminal.payload["nested"]["value"] = 400 + second = manager.snapshot(ref) + assert second.metadata["nested"]["value"] == 1 + assert second.terminal.payload["nested"]["value"] == 2 + + +def test_acknowledgement_revision_must_match_and_forget_requires_ack(): + manager = JobManager() + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit(executor, "rf", lambda: None) + terminal = _wait_for_state(manager, ref, JobState.SUCCEEDED) + + revision = terminal.terminal.revision + assert manager.forget(ref) is False + assert manager.acknowledge(ref, revision + 1) is False + assert manager.snapshot(ref).acknowledged is False + assert manager.acknowledge(ref, revision) is True + assert manager.forget(ref) is True + assert manager.snapshot(ref) is None diff --git a/src/treetracer/background_jobs.py b/src/treetracer/background_jobs.py new file mode 100644 index 0000000..1bb89ae --- /dev/null +++ b/src/treetracer/background_jobs.py @@ -0,0 +1,639 @@ +"""Thread-safe lifecycle state for background compute jobs. + +The Dash UI polls long-running RF, MDS, Pseudo-ESS, and consensus-tree +computations. A completed job must not become a one-shot event: an HTTP +response can be superseded before the browser applies it. ``JobManager`` +therefore keeps terminal state until the matching browser generation +acknowledges it. + +This module deliberately has no Dash or worker imports. Job-specific code +submits work with a success finalizer that performs its domain side effects +once and returns a small, JSON-friendly terminal payload. Polling reads +snapshots and never consumes them. +""" + +from __future__ import annotations + +import copy +import math +import time +import uuid +from concurrent.futures import CancelledError, Executor, Future +from dataclasses import dataclass +from enum import StrEnum +from threading import RLock +from typing import Any, Callable, Mapping + + +class JobState(StrEnum): + """Lifecycle states exposed to polling and diagnostics.""" + + QUEUED = "queued" + RUNNING = "running" + FINALIZING = "finalizing" + SUCCEEDED = "succeeded" + FAILED = "failed" + CANCELLED = "cancelled" + + +TERMINAL_STATES = frozenset( + {JobState.SUCCEEDED, JobState.FAILED, JobState.CANCELLED} +) + + +class JobBusyError(RuntimeError): + """Raised when a single-active manager already owns a live job.""" + + def __init__(self, active: "JobRef") -> None: + self.active = active + super().__init__( + f"background job {active.job_id} ({active.kind}) is still active" + ) + + +@dataclass(frozen=True, slots=True) +class JobRef: + """Immutable identity passed between the server and a browser store.""" + + job_id: str + generation: int + kind: str + owner_id: str | None = None + + def as_dict(self) -> dict[str, Any]: + return { + "job_id": self.job_id, + "generation": self.generation, + "kind": self.kind, + "owner_id": self.owner_id, + } + + @classmethod + def from_dict(cls, value: Mapping[str, Any]) -> "JobRef": + return cls( + job_id=str(value["job_id"]), + generation=int(value["generation"]), + kind=str(value["kind"]), + owner_id=( + None + if value.get("owner_id") is None + else str(value["owner_id"]) + ), + ) + + +@dataclass(frozen=True, slots=True) +class ProgressSnapshot: + """The latest valid progress observation for a job.""" + + fraction: float + phase: str + label: str | None = None + + def as_dict(self) -> dict[str, Any]: + return { + "fraction": self.fraction, + "phase": self.phase, + "label": self.label, + } + + +@dataclass(frozen=True, slots=True) +class TerminalEvent: + """A stable terminal event retained until browser acknowledgement.""" + + state: JobState + revision: int + payload: Mapping[str, Any] + + def as_dict(self) -> dict[str, Any]: + return { + "state": self.state.value, + "revision": self.revision, + "payload": copy.deepcopy(dict(self.payload)), + } + + +@dataclass(frozen=True, slots=True) +class JobSnapshot: + """Immutable point-in-time view returned to poll callbacks.""" + + ref: JobRef + state: JobState + progress: ProgressSnapshot | None + terminal: TerminalEvent | None + metadata: Mapping[str, Any] + acknowledged: bool + delivery_attempt: int + submitted_at: float + started_at: float | None + finished_at: float | None + + def as_dict(self) -> dict[str, Any]: + """Return a small structure suitable for a ``dcc.Store``.""" + + value = self.ref.as_dict() + value.update( + { + "state": self.state.value, + "progress": ( + None if self.progress is None else self.progress.as_dict() + ), + "terminal": ( + None if self.terminal is None else self.terminal.as_dict() + ), + "metadata": copy.deepcopy(dict(self.metadata)), + "acknowledged": self.acknowledged, + "delivery_attempt": self.delivery_attempt, + "submitted_at": self.submitted_at, + "started_at": self.started_at, + "finished_at": self.finished_at, + } + ) + return value + + +Finalizer = Callable[[JobRef, Any], Mapping[str, Any] | None] + + +@dataclass(slots=True) +class _JobRecord: + ref: JobRef + state: JobState + metadata: dict[str, Any] + submitted_at: float + future: Future[Any] | None = None + progress: ProgressSnapshot | None = None + terminal: TerminalEvent | None = None + acknowledged: bool = False + delivery_attempt: int = 0 + terminal_revision: int = 0 + started_at: float | None = None + finished_at: float | None = None + + +class JobManager: + """Own background-job state and make terminal delivery replayable. + + By default only one unacknowledged job is accepted at a time. This + mirrors TreeTracer's single-thread executor and serialized persistent + worker instead of silently building a queue behind independently enabled + compute buttons. + + The manager does not own the supplied executor. Callers remain + responsible for executor shutdown and for interrupting opaque native work + when a running job is cancelled. + """ + + def __init__( + self, + *, + single_active: bool = True, + max_history: int = 32, + clock: Callable[[], float] = time.monotonic, + id_factory: Callable[[], str] | None = None, + ) -> None: + if max_history < 1: + raise ValueError("max_history must be at least 1") + self._single_active = single_active + self._max_history = max_history + self._clock = clock + self._id_factory = id_factory or (lambda: uuid.uuid4().hex) + self._lock = RLock() + self._records: dict[str, _JobRecord] = {} + self._next_generation = 1 + + def submit( + self, + executor: Executor, + kind: str, + fn: Callable[..., Any], + /, + *args: Any, + owner_id: str | None = None, + metadata: Mapping[str, Any] | None = None, + finalizer: Finalizer | None = None, + cancel_exceptions: tuple[type[BaseException], ...] = (), + **kwargs: Any, + ) -> JobRef: + """Submit work and attach exactly one terminal-state handler. + + ``finalizer`` runs once after successful compute completion. It may + persist the raw result and must return only the small payload needed by + the UI. A finalizer exception becomes a sticky failed terminal event. + + Exceptions raised by ``executor.submit`` are also captured as a failed + terminal event, so a browser that already received the returned job + identity can render the error through the normal polling path. + """ + + kind = str(kind).strip() + if not kind: + raise ValueError("kind must be a non-empty string") + if not isinstance(cancel_exceptions, tuple) or not all( + isinstance(exc_type, type) + and issubclass(exc_type, BaseException) + for exc_type in cancel_exceptions + ): + raise TypeError("cancel_exceptions must be a tuple of exception types") + + with self._lock: + active = self._active_ref_locked() + if self._single_active and active is not None: + raise JobBusyError(active) + + generation = self._next_generation + self._next_generation += 1 + ref = JobRef( + job_id=self._id_factory(), + generation=generation, + kind=kind, + owner_id=owner_id, + ) + if ref.job_id in self._records: + raise ValueError(f"duplicate job id from id_factory: {ref.job_id}") + record = _JobRecord( + ref=ref, + state=JobState.QUEUED, + metadata=copy.deepcopy(dict(metadata or {})), + submitted_at=self._clock(), + ) + self._records[ref.job_id] = record + self._prune_history_locked() + + def run() -> None: + if not self._mark_running(ref): + return + self._run_job( + ref, + fn, + args, + kwargs, + finalizer=finalizer, + cancel_exceptions=cancel_exceptions, + ) + + try: + future = executor.submit(run) + except BaseException as exc: + self._finish_exception(ref, exc, stage="submit") + return ref + + with self._lock: + current = self._matching_record_locked(ref) + if current is not None: + current.future = future + + # This callback must remain tiny. ``Future.add_done_callback`` invokes + # it synchronously when a very fast future finished before registration; + # running domain finalization here would then block ``submit`` itself. + future.add_done_callback( + lambda completed: self._release_future(ref, completed) + ) + return ref + + def snapshot(self, ref: JobRef) -> JobSnapshot | None: + """Read state without consuming or changing terminal delivery.""" + + with self._lock: + record = self._matching_record_locked(ref) + return None if record is None else self._snapshot_locked(record) + + def snapshot_for_delivery(self, ref: JobRef) -> JobSnapshot | None: + """Read state for a poll response and count terminal replays. + + A changing delivery attempt lets a ``dcc.Store`` retrigger browser + reconciliation even though the semantic terminal event and its + revision remain stable. + """ + + with self._lock: + record = self._matching_record_locked(ref) + if record is None: + return None + if record.terminal is not None and not record.acknowledged: + record.delivery_attempt += 1 + return self._snapshot_locked(record) + + def active_ref(self) -> JobRef | None: + """Return the current non-acknowledged job, if any.""" + + with self._lock: + return self._active_ref_locked() + + def update_progress( + self, + ref: JobRef, + fraction: float, + phase: str, + label: str | None = None, + ) -> bool: + """Publish monotonic progress for a matching live job. + + Stale generations and jobs already finalizing/terminal return ``False``. + Fractions are clamped to ``[0, 1]`` and never move backwards. + """ + + fraction = float(fraction) + if not math.isfinite(fraction): + raise ValueError("progress fraction must be finite") + fraction = max(0.0, min(fraction, 1.0)) + phase = str(phase).strip() or "running" + + with self._lock: + record = self._matching_record_locked(ref) + if record is None or record.state not in { + JobState.QUEUED, + JobState.RUNNING, + }: + return False + if record.progress is not None: + fraction = max(record.progress.fraction, fraction) + record.progress = ProgressSnapshot(fraction, phase, label) + return True + + def acknowledge( + self, + ref: JobRef, + terminal_revision: int, + ) -> bool: + """Acknowledge only the exact terminal event applied by the browser.""" + + with self._lock: + record = self._matching_record_locked(ref) + if record is None or record.terminal is None: + return False + if record.terminal.revision != int(terminal_revision): + return False + record.acknowledged = True + self._prune_history_locked() + return True + + def mark_cancelled( + self, + ref: JobRef, + *, + message: str = "Background computation cancelled", + ) -> bool: + """Publish a sticky cancellation for queued/running work. + + The caller must separately interrupt work that ``Future.cancel()`` + cannot stop (for TreeTracer, this means killing the persistent worker). + Cancellation loses to finalization once irreversible domain side + effects have begun. + """ + + with self._lock: + record = self._matching_record_locked(ref) + if record is None or record.state not in { + JobState.QUEUED, + JobState.RUNNING, + }: + return False + future = record.future + self._set_terminal_locked( + record, + JobState.CANCELLED, + {"message": str(message), "stage": "compute"}, + ) + + if future is not None: + future.cancel() + return True + + def forget(self, ref: JobRef) -> bool: + """Remove an acknowledged terminal record from retained history.""" + + with self._lock: + record = self._matching_record_locked(ref) + if ( + record is None + or record.terminal is None + or not record.acknowledged + ): + return False + del self._records[ref.job_id] + return True + + def _mark_running(self, ref: JobRef) -> bool: + with self._lock: + record = self._matching_record_locked(ref) + if record is None or record.state is not JobState.QUEUED: + return False + record.state = JobState.RUNNING + record.started_at = self._clock() + return True + + def _run_job( + self, + ref: JobRef, + fn: Callable[..., Any], + args: tuple[Any, ...], + kwargs: Mapping[str, Any], + *, + finalizer: Finalizer | None, + cancel_exceptions: tuple[type[BaseException], ...], + ) -> None: + try: + result = fn(*args, **kwargs) + except CancelledError as exc: + if self._claim_finalization(ref): + self._finish_cancelled(ref, exc) + return + except BaseException as exc: + if self._claim_finalization(ref): + if cancel_exceptions and isinstance(exc, cancel_exceptions): + self._finish_cancelled(ref, exc) + else: + self._finish_exception(ref, exc, stage="compute") + return + + # Explicit cancellation can win while the opaque compute function is + # running. In that case its late result is intentionally discarded. + if not self._claim_finalization(ref): + return + + try: + payload = {} if finalizer is None else finalizer(ref, result) + if payload is None: + payload = {} + if not isinstance(payload, Mapping): + raise TypeError("job finalizer must return a mapping or None") + payload = copy.deepcopy(dict(payload)) + except BaseException as exc: + self._finish_exception(ref, exc, stage="finalize") + return + + with self._lock: + record = self._matching_record_locked(ref) + if record is None or record.state is not JobState.FINALIZING: + return + self._set_terminal_locked(record, JobState.SUCCEEDED, payload) + + def _claim_finalization(self, ref: JobRef) -> bool: + """Atomically grant one caller permission to finalize a job.""" + + with self._lock: + record = self._matching_record_locked(ref) + if record is None or record.state not in { + JobState.QUEUED, + JobState.RUNNING, + }: + return False + record.state = JobState.FINALIZING + return True + + def _release_future(self, ref: JobRef, future: Future[Any]) -> None: + """Drop the completed Future without doing domain work. + + The wrapped job normally records its own terminal state. The only + additional case handled here is an externally cancelled queued future, + whose wrapper never had a chance to run. + """ + + with self._lock: + record = self._matching_record_locked(ref) + if record is None: + return + if record.future is future: + record.future = None + if future.cancelled() and record.state is JobState.QUEUED: + self._set_terminal_locked( + record, + JobState.CANCELLED, + { + "message": "Background computation cancelled", + "error_type": "CancelledError", + "stage": "compute", + }, + ) + + def _finish_cancelled(self, ref: JobRef, exc: BaseException) -> None: + message = str(exc) or "Background computation cancelled" + with self._lock: + record = self._matching_record_locked(ref) + if record is None or record.state is not JobState.FINALIZING: + return + self._set_terminal_locked( + record, + JobState.CANCELLED, + { + "message": message, + "error_type": type(exc).__name__, + "stage": "compute", + }, + ) + + def _finish_exception( + self, + ref: JobRef, + exc: BaseException, + *, + stage: str, + ) -> None: + with self._lock: + record = self._matching_record_locked(ref) + if record is None: + return + if stage != "submit" and record.state is not JobState.FINALIZING: + return + if stage == "submit" and record.state not in { + JobState.QUEUED, + JobState.RUNNING, + }: + return + self._set_terminal_locked( + record, + JobState.FAILED, + { + "message": str(exc) or type(exc).__name__, + "error_type": type(exc).__name__, + "stage": stage, + }, + ) + + def _set_terminal_locked( + self, + record: _JobRecord, + state: JobState, + payload: Mapping[str, Any], + ) -> None: + if state not in TERMINAL_STATES: + raise ValueError(f"{state!r} is not a terminal state") + record.state = state + record.future = None + record.finished_at = self._clock() + record.terminal_revision += 1 + record.terminal = TerminalEvent( + state=state, + revision=record.terminal_revision, + payload=copy.deepcopy(dict(payload)), + ) + if state is JobState.SUCCEEDED: + record.progress = ProgressSnapshot(1.0, "done", "complete") + else: + last_fraction = ( + 0.0 if record.progress is None else record.progress.fraction + ) + record.progress = ProgressSnapshot(last_fraction, state.value) + + def _matching_record_locked(self, ref: JobRef) -> _JobRecord | None: + record = self._records.get(ref.job_id) + if record is None or record.ref != ref: + return None + return record + + def _active_ref_locked(self) -> JobRef | None: + for record in reversed(tuple(self._records.values())): + if not record.acknowledged: + return record.ref + return None + + def _snapshot_locked(self, record: _JobRecord) -> JobSnapshot: + progress = record.progress + terminal = record.terminal + return JobSnapshot( + ref=record.ref, + state=record.state, + progress=( + None + if progress is None + else ProgressSnapshot( + progress.fraction, + progress.phase, + progress.label, + ) + ), + terminal=( + None + if terminal is None + else TerminalEvent( + terminal.state, + terminal.revision, + copy.deepcopy(dict(terminal.payload)), + ) + ), + metadata=copy.deepcopy(record.metadata), + acknowledged=record.acknowledged, + delivery_attempt=record.delivery_attempt, + submitted_at=record.submitted_at, + started_at=record.started_at, + finished_at=record.finished_at, + ) + + def _prune_history_locked(self) -> None: + excess = len(self._records) - self._max_history + if excess <= 0: + return + removable = [ + job_id + for job_id, record in self._records.items() + if record.acknowledged + ] + for job_id in removable[:excess]: + del self._records[job_id] + + +# Process-local coordinator used by the desktop app. Browser-mode migration +# must provide ``owner_id`` values before supporting multiple independent tabs. +job_manager = JobManager() From df543c3117526a2af0592c6e1ccb0c79ea22feae Mon Sep 17 00:00:00 2001 From: Sam Hong <20244918+hongsamL@users.noreply.github.com> Date: Mon, 3 Aug 2026 14:59:02 +0200 Subject: [PATCH 3/9] RF MDS with job manager --- src/test/test_background_jobs.py | 28 ++ src/test/test_compute_job_lifecycle.py | 250 ++++++++++ src/treetracer/background_jobs.py | 21 + src/treetracer/callbacks/compute.py | 638 ++++++++++++++++++------- src/treetracer/callbacks/sidebar.py | 38 +- src/treetracer/ui/navbar.py | 8 + 6 files changed, 809 insertions(+), 174 deletions(-) create mode 100644 src/test/test_compute_job_lifecycle.py diff --git a/src/test/test_background_jobs.py b/src/test/test_background_jobs.py index 2aa73d5..e90620d 100644 --- a/src/test/test_background_jobs.py +++ b/src/test/test_background_jobs.py @@ -307,3 +307,31 @@ def test_acknowledgement_revision_must_match_and_forget_requires_ack(): assert manager.acknowledge(ref, revision) is True assert manager.forget(ref) is True assert manager.snapshot(ref) is None + + +def test_invalidate_discards_late_result_without_running_finalizer(): + manager = JobManager() + task_started = threading.Event() + release_task = threading.Event() + finalizer_calls = 0 + + def work(): + task_started.set() + assert release_task.wait(timeout=2) + return "stale result" + + def finalize(_ref, _result): + nonlocal finalizer_calls + finalizer_calls += 1 + return {"should_not": "appear"} + + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit(executor, "rf", work, finalizer=finalize) + assert task_started.wait(timeout=2) + assert manager.invalidate(ref) is True + assert manager.snapshot(ref) is None + assert manager.active_ref() is None + release_task.set() + + assert finalizer_calls == 0 + assert manager.snapshot(ref) is None diff --git a/src/test/test_compute_job_lifecycle.py b/src/test/test_compute_job_lifecycle.py new file mode 100644 index 0000000..e38f038 --- /dev/null +++ b/src/test/test_compute_job_lifecycle.py @@ -0,0 +1,250 @@ +"""RF/MDS integration tests for the sticky background-job lifecycle.""" + +from __future__ import annotations + +import threading +import time +from concurrent.futures import ThreadPoolExecutor + +import pytest + +from treetracer.background_jobs import JobManager, JobState +from treetracer.callbacks import compute + + +def _wait_for_terminal(manager, ref, timeout=2.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + snapshot = manager.snapshot(ref) + if snapshot is not None and snapshot.terminal is not None: + return snapshot + time.sleep(0.005) + pytest.fail("background job did not become terminal") + + +def _registered_callback(name): + from dash import _callback + + matches = [] + for callback_data in _callback.GLOBAL_CALLBACK_MAP.values(): + callback_fn = callback_data.get("callback") + if callback_fn is None: + continue + original = getattr(callback_fn, "__wrapped__", callback_fn) + if original.__name__ == name: + matches.append(original) + if not matches: + compute.register_compute_callbacks() + return _registered_callback(name) + return matches[-1] + + +def test_latest_job_ref_uses_generation_and_rejects_invalid_data(): + rf = { + "job_id": "rf-job", + "generation": 4, + "kind": "rf", + "owner_id": None, + } + mds = { + "job_id": "mds-job", + "generation": 7, + "kind": "mds", + "owner_id": None, + } + + assert compute._latest_rf_mds_ref(rf, mds).job_id == "mds-job" + assert compute._latest_rf_mds_ref(rf, None).job_id == "rf-job" + assert compute._latest_rf_mds_ref({"kind": "rf"}, None) is None + assert compute._latest_rf_mds_ref( + {**rf, "kind": "pseudo_ess"}, + None, + ) is None + + +def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): + manager = JobManager(id_factory=lambda: "rf-test-job") + register_calls = [] + expected_index = { + "RF_001": { + "n_trees": 2, + "file_breakdown": {"run.trees": 2}, + "groups_per_file": {"run.trees": ["run"]}, + "is_rooted": True, + } + } + monkeypatch.setattr(compute, "job_manager", manager) + monkeypatch.setattr(compute, "add_log", lambda *_args, **_kwargs: None) + monkeypatch.setattr(compute, "_wlog", lambda *_args, **_kwargs: None) + monkeypatch.setattr( + compute, + "register_distmat", + lambda *args, **kwargs: register_calls.append((args, kwargs)), + ) + monkeypatch.setattr(compute, "get_distmat_index", lambda: expected_index) + + pipeline = { + "rf_name": "RF_001", + "result_names": ["run/tree-1", "run/tree-2"], + "file_breakdown": {"run.trees": 2}, + "groups_per_file": {"run.trees": ["run"]}, + "total_elapsed": 1.25, + "compute_elapsed": 1.0, + "is_rooted": True, + "save_path": "/tmp/rf-test.npy", + } + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "rf", + lambda: pipeline, + finalizer=compute._finalize_rf_job, + ) + terminal = _wait_for_terminal(manager, ref) + + assert terminal.state is JobState.SUCCEEDED + assert len(register_calls) == 1 + assert terminal.terminal.payload["distmat_index"] == expected_index + + poll = _registered_callback("poll_completion") + acknowledge = _registered_callback("acknowledge_terminal_job") + first = poll(10, ref.as_dict(), None, False) + second = poll(11, ref.as_dict(), None, False) + + assert len(first) == 16 + assert first[1] == expected_index + assert first[4] is False + assert first[11] is True + assert first[12]["terminal_revision"] == terminal.terminal.revision + assert first[12]["delivery_attempt"] == 1 + assert second[12]["delivery_attempt"] == 2 + assert len(register_calls) == 1 + assert manager.snapshot(ref).acknowledged is False + + ack_store = acknowledge(second[12]) + assert ack_store["acknowledged"] is True + assert manager.snapshot(ref).acknowledged is True + assert manager.active_ref() is None + + +def test_mds_finalization_stores_full_result_once_and_replays_small_index( + monkeypatch, +): + manager = JobManager(id_factory=lambda: "mds-test-job") + stored = {} + expected_index = { + "RF_002_MDS.tsv": { + "filename": "RF_002_MDS.tsv", + "source_distmat": "RF_002", + "rows": 4, + "dimensions": ["MDS1", "MDS2"], + "groups": ["run-a", "run-b"], + "MIN_TREENUM": 1, + "MAX_TREENUM": 2, + } + } + monkeypatch.setattr(compute, "job_manager", manager) + monkeypatch.setattr(compute, "add_log", lambda *_args, **_kwargs: None) + monkeypatch.setattr( + compute, + "store_mds_result", + lambda key, value: stored.__setitem__(key, value), + ) + monkeypatch.setattr(compute, "get_mds_results_index", lambda: expected_index) + + result = { + "embedding": [ + [0.0, 0.1], + [0.2, 0.3], + [0.4, 0.5], + [0.6, 0.7], + ], + "elapsed": 2.5, + "tree_names": [ + "run-a/tree-1", + "run-a/tree-2", + "run-b/tree-1", + "run-b/tree-2", + ], + "selected_distmat": "RF_002", + "n_components": 2, + } + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "mds", + lambda: result, + finalizer=compute._finalize_mds_job, + ) + terminal = _wait_for_terminal(manager, ref) + + assert list(stored) == ["RF_002_MDS.tsv"] + assert len(stored["RF_002_MDS.tsv"]["data"]) == 4 + assert terminal.terminal.payload["mds_results_index"] == expected_index + assert "embedding" not in terminal.terminal.payload + assert "data" not in terminal.terminal.payload + + manager.snapshot_for_delivery(ref) + manager.snapshot_for_delivery(ref) + assert list(stored) == ["RF_002_MDS.tsv"] + + +def test_reset_waits_for_an_inflight_result_publication(monkeypatch): + manager = JobManager(id_factory=lambda: "reset-test-job") + publication_started = threading.Event() + release_publication = threading.Event() + reset_finished = threading.Event() + register_calls = [] + + def blocking_register(*args, **kwargs): + publication_started.set() + assert release_publication.wait(timeout=2) + register_calls.append((args, kwargs)) + + monkeypatch.setattr(compute, "job_manager", manager) + monkeypatch.setattr(compute, "add_log", lambda *_args, **_kwargs: None) + monkeypatch.setattr(compute, "register_distmat", blocking_register) + monkeypatch.setattr(compute, "get_distmat_index", lambda: {}) + monkeypatch.setattr( + compute.persistent_worker, + "cancel_current_job", + lambda: False, + ) + + pipeline = { + "rf_name": "RF_003", + "result_names": ["run/tree-1", "run/tree-2"], + "file_breakdown": {"run.trees": 2}, + "groups_per_file": {"run.trees": ["run"]}, + "total_elapsed": 1.0, + "compute_elapsed": 0.8, + "is_rooted": True, + "save_path": "/tmp/rf-reset-test.npy", + } + executor = ThreadPoolExecutor(max_workers=1) + ref = manager.submit( + executor, + "rf", + lambda: pipeline, + finalizer=compute._finalize_rf_job, + ) + assert publication_started.wait(timeout=2) + + def run_reset(): + compute.reset() + reset_finished.set() + + reset_thread = threading.Thread(target=run_reset) + reset_thread.start() + deadline = time.monotonic() + 2 + while manager.snapshot(ref) is not None and time.monotonic() < deadline: + time.sleep(0.005) + + assert manager.snapshot(ref) is None + assert reset_finished.wait(timeout=0.05) is False + release_publication.set() + assert reset_finished.wait(timeout=2) + reset_thread.join(timeout=2) + executor.shutdown(wait=True) + + assert len(register_calls) == 1 diff --git a/src/treetracer/background_jobs.py b/src/treetracer/background_jobs.py index 1bb89ae..d921e85 100644 --- a/src/treetracer/background_jobs.py +++ b/src/treetracer/background_jobs.py @@ -415,6 +415,27 @@ def forget(self, ref: JobRef) -> bool: del self._records[ref.job_id] return True + def invalidate(self, ref: JobRef) -> bool: + """Forget a job immediately and reject any late completion. + + This is the reset/clear-data operation, not normal browser delivery. + Removing the record makes every later state transition from the job's + wrapper fail its identity lookup, so a stale result cannot be finalized + into freshly cleared application state. Running native work still has + to be interrupted separately by its owner. + """ + + with self._lock: + record = self._matching_record_locked(ref) + if record is None: + return False + future = record.future + del self._records[ref.job_id] + + if future is not None: + future.cancel() + return True + def _mark_running(self, ref: JobRef) -> bool: with self._lock: record = self._matching_record_locked(ref) diff --git a/src/treetracer/callbacks/compute.py b/src/treetracer/callbacks/compute.py index f9a6a6a..634f83d 100644 --- a/src/treetracer/callbacks/compute.py +++ b/src/treetracer/callbacks/compute.py @@ -1,17 +1,16 @@ from dash import html, callback, Input, Output, State, no_update, ALL, ctx import dash_mantine_components as dmc -import os from concurrent.futures import ThreadPoolExecutor -import numpy as np +from threading import RLock import pandas as pd from ..logger import add_log, notif_id +from ..background_jobs import JobBusyError, JobRef, JobState, job_manager from ..db.tree_service import get_tree_service from ..state import (load_distmat, get_distmat_index, next_distmat_name, get_distmat_path, register_distmat, get_distmat_file_path, - get_distmat_groups_per_file, - store_mds_result, get_mds_results_index, - clear_all_mds_results) + get_distmat_names, store_mds_result, + get_mds_results_index) from ..ui.widgets import computing_banner from ._helpers import _save_file_dialog, extract_group from .._worker_log import log as _wlog @@ -33,10 +32,7 @@ # work runs in the worker subprocess, keeping the pywebview/Dash process # responsive while compute jobs are in flight. _executor = None -_rf_future = None # concurrent.futures.Future for RF job -_rf_meta = {} # metadata needed by poll_completion to save RF result -_mds_future = None # concurrent.futures.Future for between-run MDS job -_mds_meta = {} # metadata needed by poll_completion to build MDS result +_publication_lock = RLock() def _mds_export_filename(source_distmat): @@ -44,6 +40,68 @@ def _mds_export_filename(source_distmat): return f"{stem}_MDS.tsv" +def _job_ref_from_store(data, *, expected_kind=None): + """Parse a browser job reference defensively. + + Browser stores can be empty, stale, or manually modified. Invalid data is + treated as absent; ``JobManager`` performs the authoritative generation + check on every operation. + """ + if not isinstance(data, dict): + return None + try: + ref = JobRef.from_dict(data) + except (KeyError, TypeError, ValueError): + return None + if expected_kind is not None and ref.kind != expected_kind: + return None + if ref.kind not in {"rf", "mds"}: + return None + return ref + + +def _latest_rf_mds_ref(rf_job_data, mds_job_data): + refs = [ + ref + for ref in ( + _job_ref_from_store(rf_job_data, expected_kind="rf"), + _job_ref_from_store(mds_job_data, expected_kind="mds"), + ) + if ref is not None + ] + return max(refs, key=lambda ref: ref.generation, default=None) + + +def _ack_applied_job(applied_data): + """Acknowledge a terminal event known to have reached browser state.""" + ref = _job_ref_from_store(applied_data) + if ref is None: + return False + try: + revision = int(applied_data["terminal_revision"]) + except (KeyError, TypeError, ValueError): + return False + return job_manager.acknowledge(ref, revision) + + +def _terminal_progress_from_store(job_data, expected_kind): + ref = _job_ref_from_store(job_data, expected_kind=expected_kind) + if ref is None: + return None + snapshot = job_manager.snapshot(ref) + if snapshot is None or snapshot.terminal is None: + return None + progress = snapshot.progress + if progress is None: + return None + label = { + JobState.SUCCEEDED: "complete", + JobState.FAILED: "failed", + JobState.CANCELLED: "cancelled", + }[snapshot.state] + return progress.fraction * 100.0, label + + def _get_executor(): global _executor if _executor is None: @@ -139,10 +197,139 @@ def _rf_pipeline(selected_files, save_path, rf_name, is_rooted): result["total_elapsed"] = time.time() - t0 result["is_rooted"] = is_rooted result["progress_path"] = progress_path + result["save_path"] = save_path _wlog(f"[parent] _rf_pipeline: returning result; total_elapsed={result['total_elapsed']:.3f}s") return result +def _job_can_publish(ref): + snapshot = job_manager.snapshot(ref) + return snapshot is not None and snapshot.state is JobState.FINALIZING + + +def _finalize_rf_job(ref, pipeline): + """Persist one RF result and return its small terminal UI payload.""" + with _publication_lock: + if not _job_can_publish(ref): + return {} + return _publish_rf_result(pipeline) + + +def _publish_rf_result(pipeline): + rf_name = pipeline["rf_name"] + result_names = pipeline["result_names"] + file_breakdown = pipeline["file_breakdown"] + groups_per_file = pipeline["groups_per_file"] + elapsed = pipeline["total_elapsed"] + compute_elapsed = pipeline["compute_elapsed"] + is_rooted = pipeline.get("is_rooted", True) + + # The matrix is already on disk; this is the exactly-once publication step. + register_distmat( + rf_name, + result_names, + pipeline["save_path"], + file_breakdown=file_breakdown, + groups_per_file=groups_per_file, + is_rooted=is_rooted, + ) + add_log( + f"Stored RF distance matrix as '{rf_name}' " + f"({len(result_names)}x{len(result_names)})" + ) + add_log( + f"RF pipeline took {elapsed:.2f}s " + f"(rapidtrees compute {compute_elapsed:.2f}s)" + ) + return { + "rf_name": rf_name, + "n_trees": len(result_names), + "elapsed": elapsed, + "compute_elapsed": compute_elapsed, + "distmat_index": get_distmat_index(), + } + + +def _mds_pipeline( + matrix_path, + progress_path, + tree_names, + selected_distmat, + n_components, +): + embedding_list, elapsed = persistent_worker.submit_job( + "compute_mds", + matrix_path=str(matrix_path), + n_components=n_components, + progress_path=progress_path, + ) + return { + "embedding": embedding_list, + "elapsed": elapsed, + "tree_names": [str(name) for name in tree_names], + "selected_distmat": selected_distmat, + "n_components": n_components, + } + + +def _finalize_mds_job(ref, result): + """Persist one MDS result and return its small terminal UI payload.""" + with _publication_lock: + if not _job_can_publish(ref): + return {} + return _publish_mds_result(result) + + +def _publish_mds_result(result): + embedding_list = result["embedding"] + elapsed = result["elapsed"] + tree_names = result["tree_names"] + selected_distmat = result["selected_distmat"] + n_components = result["n_components"] + + mdscols = [f"MDS{i + 1}" for i in range(n_components)] + mds_df = pd.DataFrame(embedding_list, columns=mdscols) + mds_df["tree"] = tree_names + mds_df["group"] = mds_df["tree"].apply(extract_group).astype(str) + group_mapping = { + value: index + for index, value in enumerate(sorted(mds_df["group"].unique())) + } + mds_df["group_col"] = mds_df["group"].map(group_mapping) + mds_df["treenum"] = mds_df.groupby("group").cumcount() + 1 + mds_df["size"] = 6 + mds_filename = _mds_export_filename(selected_distmat) + mds_df["file"] = mds_filename + + metadata = { + "filename": mds_filename, + "source_distmat": selected_distmat, + "rows": len(mds_df), + "dimensions": mdscols, + "groups": mds_df["group"].unique().tolist(), + "MIN_TREENUM": int(mds_df["treenum"].min()), + "MAX_TREENUM": int(mds_df["treenum"].max()), + } + store_mds_result( + mds_filename, + {"metadata": metadata, "data": mds_df.to_dict("records")}, + ) + + n_groups = len(metadata["groups"]) + add_log( + f"PCoA completed in {elapsed:.2f}s: {len(mds_df)} points, " + f"{n_components}D, {n_groups} groups" + ) + return { + "mds_results_index": get_mds_results_index(), + "mds_filename": mds_filename, + "n_points": len(mds_df), + "n_components": n_components, + "n_groups": n_groups, + "elapsed": elapsed, + } + + def _shutdown_executor(): global _executor if _executor is not None: @@ -150,6 +337,25 @@ def _shutdown_executor(): _executor = None +def reset(): + """Invalidate any RF/MDS job before Clear Data wipes its inputs. + + Removing the manager record prevents a late worker result from + re-registering a matrix or MDS result after the application state has been + cleared. The persistent-worker kill interrupts native work when possible. + """ + active = job_manager.active_ref() + if active is None or active.kind not in {"rf", "mds"}: + return + persistent_worker.cancel_current_job() + job_manager.invalidate(active) + # If a finalizer acquired the publication lock just before invalidation, + # let it finish before Clear Data removes published state. If it acquires + # the lock afterward, its identity re-check rejects the invalidated job. + with _publication_lock: + pass + + def register_compute_callbacks(): # Stop button — interrupt whatever compute is currently running by # killing the persistent worker (see @@ -172,7 +378,17 @@ def handle_compute_stop(_stop_clicks, rf_progress_path, mds_progress_path): # n_clicks in the trigger; layout-change fires carry None. if not any(t.get("value") for t in ctx.triggered): return no_update - if persistent_worker.cancel_current_job(): + active = job_manager.active_ref() + managed_cancelled = ( + active is not None + and active.kind in {"rf", "mds"} + and job_manager.mark_cancelled( + active, + message="Computation cancelled by user", + ) + ) + worker_cancelled = persistent_worker.cancel_current_job() + if managed_cancelled or worker_cancelled: # Cancel hard-kills the worker subprocess, so the progress # writer thread's ``finally`` block (which normally unlinks # the sidecar file) never runs. Tidy up here. The file is @@ -293,15 +509,28 @@ def open_export_drawer(_rf_clicks, _mds_clicks): # starts so ``update_rf_progress`` knows which sidecar file to # tail. Early returns leave it at no_update. Output("rf-progress-path", "data", allow_duplicate=True), + Output("rf-job-store", "data"), Input("compute-rf-button", "n_clicks"), State({"type": "compute-tree-checkbox", "index": ALL}, "checked"), State({"type": "compute-tree-checkbox", "index": ALL}, "id"), State("tree-offset-store", "data"), + State("rf-mds-applied-job-store", "data"), prevent_initial_call=True, ) - def handle_compute_rf(n_clicks, checked_list, id_list, stored_summaries): + def handle_compute_rf( + n_clicks, + checked_list, + id_list, + stored_summaries, + applied_job, + ): if not n_clicks or not stored_summaries: - return no_update, no_update, no_update, no_update, no_update + return (no_update,) * 6 + + # The marker is written in the same browser response that rendered the + # prior terminal UI. A fast next click may beat the acknowledgement + # callback, so acknowledge idempotently here before submitting too. + _ack_applied_job(applied_job) # Determine which files are selected selected_files = [ @@ -322,7 +551,7 @@ def handle_compute_rf(n_clicks, checked_list, id_list, stored_summaries): action="show", autoClose=6000, id=notif_id(), - ), no_update, no_update, no_update, no_update + ), no_update, no_update, no_update, no_update, no_update # Collect taxa counts for selected files taxa_counts = {} @@ -346,7 +575,7 @@ def handle_compute_rf(n_clicks, checked_list, id_list, stored_summaries): action="show", autoClose=6000, id=notif_id(), - ), no_update, no_update, no_update, no_update + ), no_update, no_update, no_update, no_update, no_update # --- Rooting consistency check --- # RF over rooted clades and RF over bipartitions are different @@ -371,7 +600,7 @@ def handle_compute_rf(n_clicks, checked_list, id_list, stored_summaries): title="Rooting Mismatch", message=msg, color="red", action="show", autoClose=8000, id=notif_id(), - ), no_update, no_update, no_update, no_update + ), no_update, no_update, no_update, no_update, no_update selected_is_rooted = unique_rootings.pop() # --- Taxa validation passed — submit the pipeline to a worker thread --- @@ -394,14 +623,44 @@ def handle_compute_rf(n_clicks, checked_list, id_list, stored_summaries): rf_name = next_distmat_name() save_path = get_distmat_path(rf_name) - _rf_meta["name"] = rf_name - _rf_meta["save_path"] = save_path - _rf_meta["is_rooted"] = selected_is_rooted - - global _rf_future - _rf_future = _get_executor().submit( - _rf_pipeline, selected_files, save_path, rf_name, selected_is_rooted, - ) + progress_path = save_path + ".progress" + try: + job_ref = job_manager.submit( + _get_executor(), + "rf", + _rf_pipeline, + selected_files, + save_path, + rf_name, + selected_is_rooted, + metadata={ + "display_name": rf_name, + "progress_path": progress_path, + }, + finalizer=_finalize_rf_job, + cancel_exceptions=(persistent_worker.JobCancelled,), + ) + except JobBusyError as exc: + msg = ( + f"Another computation ({exc.active.kind.upper()}) is still " + "finishing. Please wait for it to complete." + ) + add_log(msg, "WARNING") + return ( + dmc.Notification( + title="Computation already running", + message=msg, + color="yellow", + action="show", + autoClose=4000, + id=notif_id(), + ), + no_update, + no_update, + no_update, + no_update, + no_update, + ) computing_indicator = computing_banner( title=f"Computing RF Distances ({rf_name})...", @@ -414,8 +673,14 @@ def handle_compute_rf(n_clicks, checked_list, id_list, stored_summaries): ) # Hand the same path ``_rf_pipeline`` derives down to the # Store so ``update_rf_progress`` reads from the right file. - progress_path = save_path + ".progress" - return no_update, computing_indicator, False, True, progress_path + return ( + no_update, + computing_indicator, + False, + True, + progress_path, + job_ref.as_dict(), + ) # Per-tick reader for the RF progress sidecar file. Runs off the # same ``compute-poll-interval`` as ``poll_completion`` but writes @@ -426,9 +691,13 @@ def handle_compute_rf(n_clicks, checked_list, id_list, stored_summaries): Output("rf-progress-label", "children"), Input("compute-poll-interval", "n_intervals"), State("rf-progress-path", "data"), + State("rf-job-store", "data"), prevent_initial_call=True, ) - def update_rf_progress(_n, progress_path): + def update_rf_progress(_n, progress_path, job_data): + terminal_progress = _terminal_progress_from_store(job_data, "rf") + if terminal_progress is not None: + return terminal_progress if not progress_path: return no_update, no_update import json @@ -439,14 +708,18 @@ def update_rf_progress(_n, progress_path): # File doesn't exist yet, was just deleted, or caught # mid-write — try again next tick. ValueError covers # JSONDecodeError (subclass) too. - return no_update, no_update + terminal_progress = _terminal_progress_from_store(job_data, "rf") + return terminal_progress or (no_update, no_update) val = int(data.get("value", 0)) tot = int(data.get("total", 0)) - frac = float(data.get("fraction", 0.0)) + frac = max(0.0, min(float(data.get("fraction", 0.0)), 1.0)) phase = data.get("phase", "computing") + ref = _job_ref_from_store(job_data, expected_kind="rf") + if ref is not None: + job_manager.update_progress(ref, frac, phase) if tot == 0: return 0, "starting…" - pct = max(0.0, min(100.0, frac * 100.0)) + pct = frac * 100.0 if phase == "finalizing": label = f"{val:,} / {tot:,} pairs — finalizing…" else: @@ -458,9 +731,13 @@ def update_rf_progress(_n, progress_path): Output("mds-progress-label", "children"), Input("compute-poll-interval", "n_intervals"), State("mds-progress-path", "data"), + State("mds-job-store", "data"), prevent_initial_call=True, ) - def update_mds_progress(_n, progress_path): + def update_mds_progress(_n, progress_path, job_data): + terminal_progress = _terminal_progress_from_store(job_data, "mds") + if terminal_progress is not None: + return terminal_progress if not progress_path: return no_update, no_update import json @@ -468,12 +745,16 @@ def update_mds_progress(_n, progress_path): try: data = json.loads(Path(progress_path).read_text()) except (OSError, ValueError): - return no_update, no_update + terminal_progress = _terminal_progress_from_store(job_data, "mds") + return terminal_progress or (no_update, no_update) frac = max(0.0, min(float(data.get("fraction", 0.0)), 1.0)) pct = frac * 100.0 phase = data.get("phase", "computing") label = data.get("label") or phase.replace("_", " ") + ref = _job_ref_from_store(job_data, expected_kind="mds") + if ref is not None: + job_manager.update_progress(ref, frac, phase, label) return pct, f"{label} ({pct:.0f}%)" # ------ RF MATRIX LIST (right column) ------ @@ -568,45 +849,73 @@ def show_distmat_info(selected, distmat_data): Output("compute-poll-interval", "disabled", allow_duplicate=True), Output("compute-mds-button", "disabled", allow_duplicate=True), Output("mds-progress-path", "data", allow_duplicate=True), + Output("mds-job-store", "data"), Input("compute-mds-button", "n_clicks"), State("mds-distmat-select", "value"), + State("rf-mds-applied-job-store", "data"), prevent_initial_call=True, ) - def handle_compute_mds(n_clicks, selected_distmat): + def handle_compute_mds(n_clicks, selected_distmat, applied_job): if not n_clicks or not selected_distmat: - return no_update, no_update, no_update, no_update + return (no_update,) * 5 + + _ack_applied_job(applied_job) try: - tree_names, _ = load_distmat(selected_distmat) + # The worker loads the matrix from its registered path. Pull only + # the small name vector here instead of reading the full n×n array + # into the UI process just to determine dimensions and labels. + tree_names = list(get_distmat_names(selected_distmat)) except KeyError: msg = "Selected distance matrix not available. Please recompute RF distances." add_log(msg, "ERROR") - return dmc.Text(msg, c="red"), no_update, no_update, no_update + return dmc.Text(msg, c="red"), no_update, no_update, no_update, no_update n = len(tree_names) n_components = min(6, n - 1) add_log(f"Computing MDS from {selected_distmat} ({n}x{n}) in background process...") - _mds_meta["tree_names"] = tree_names - _mds_meta["selected_distmat"] = selected_distmat - _mds_meta["n_components"] = n_components - # Route through the persistent worker — same pattern as RF/consensus tree/ # Pseudo-ESS. Per-compute IPC overhead is ~100ms, dwarfed by the # ARPACK eigsh on a 5k×5k matrix; the win is a single unified # background-compute pattern and clean process isolation. - from . import persistent_worker matrix_path = get_distmat_file_path(selected_distmat) progress_path = str(matrix_path) + ".mds.progress" - _mds_meta["progress_path"] = progress_path - global _mds_future - _mds_future = _get_executor().submit( - persistent_worker.submit_job, - "compute_mds", - matrix_path=str(matrix_path), - n_components=n_components, - progress_path=progress_path, - ) + try: + job_ref = job_manager.submit( + _get_executor(), + "mds", + _mds_pipeline, + matrix_path, + progress_path, + tree_names, + selected_distmat, + n_components, + metadata={ + "display_name": _mds_export_filename(selected_distmat), + "progress_path": progress_path, + }, + finalizer=_finalize_mds_job, + cancel_exceptions=(persistent_worker.JobCancelled,), + ) + except JobBusyError as exc: + msg = ( + f"Another computation ({exc.active.kind.upper()}) is still " + "finishing. Please wait for it to complete." + ) + add_log(msg, "WARNING") + return ( + dmc.Alert( + title="Computation already running", + children=dmc.Text(msg, size="sm"), + color="yellow", + variant="light", + ), + no_update, + False, + no_update, + no_update, + ) computing_indicator = computing_banner( title="Computing MDS Embedding...", @@ -617,10 +926,9 @@ def handle_compute_mds(n_clicks, selected_distmat): which="mds", show_progress=True, ) - return computing_indicator, False, True, progress_path + return computing_indicator, False, True, progress_path, job_ref.as_dict() - # ------ POLL + PROCESS: checks futures, processes results in one round trip ------ - # Processing is fast (<20ms) since workers save to disk — no large pickle transfer. + # ------ POLL + RENDER: replay sticky terminal state until browser ack ------ @callback( # RF outputs (5) @@ -635,40 +943,54 @@ def handle_compute_mds(n_clicks, selected_distmat): Output("export-mds-button", "disabled"), Output("plot-config-store", "data", allow_duplicate=True), Output("compute-mds-button", "disabled", allow_duplicate=True), - # Shared outputs (2) + # Shared outputs (3). The applied marker lands in the same browser + # response as the result UI; only that marker authorizes the separate + # acknowledgement callback to consume the retained terminal event. Output("notifications-container", "children", allow_duplicate=True), Output("compute-poll-interval", "disabled", allow_duplicate=True), + Output("rf-mds-applied-job-store", "data"), # Auto-collapse sidebar on RF success (3 outputs mirror the # shell sidebar-toggle callback's outputs). Output("navbar", "style", allow_duplicate=True), Output("sidebar-visible", "data", allow_duplicate=True), Output("appshell", "navbar", allow_duplicate=True), Input("compute-poll-interval", "n_intervals"), + Input("rf-job-store", "data"), + Input("mds-job-store", "data"), State("sidebar-visible", "data"), prevent_initial_call=True, ) - def poll_completion(n_intervals, sidebar_visible): - global _rf_future, _mds_future - - rf_done = _rf_future is not None and _rf_future.done() - mds_done = _mds_future is not None and _mds_future.done() - - # Sample the poll loop sparingly — every 10 ticks (~1 s) we - # log a heartbeat with the futures' state. Useful for telling - # "poll not firing" apart from "poll firing but future never - # done" when diagnosing a hung compute. Once a job IS done - # we always log so we can see the dispatch flowing. - if rf_done or mds_done or (n_intervals or 0) % 10 == 0: - _wlog( - f"[parent] poll_completion tick={n_intervals}: " - f"rf_future={'set' if _rf_future else 'none'}/" - f"{'done' if rf_done else 'pending'}, " - f"mds_future={'set' if _mds_future else 'none'}/" - f"{'done' if mds_done else 'pending'}" - ) - - if not rf_done and not mds_done: - return (no_update,) * 15 + def poll_completion( + n_intervals, + rf_job_data, + mds_job_data, + sidebar_visible, + ): + ref = _latest_rf_mds_ref(rf_job_data, mds_job_data) + if ref is None: + return (no_update,) * 16 + + snapshot = job_manager.snapshot_for_delivery(ref) + if snapshot is None or snapshot.acknowledged: + return (no_update,) * 16 + + if snapshot.terminal is None: + if (n_intervals or 0) % 10 == 0: + _wlog( + f"[parent] poll_completion tick={n_intervals}: " + f"job={ref.job_id}/{ref.kind}/generation-{ref.generation}, " + f"state={snapshot.state.value}" + ) + return (no_update,) * 16 + + terminal = snapshot.terminal + payload = terminal.payload + _wlog( + f"[parent] poll_completion tick={n_intervals}: " + f"delivering job={ref.job_id}/{ref.kind}/generation-{ref.generation}, " + f"state={terminal.state.value}, revision={terminal.revision}, " + f"attempt={snapshot.delivery_attempt}" + ) rf_out = [no_update] * 5 mds_out = [no_update] * 5 @@ -679,14 +1001,8 @@ def poll_completion(n_intervals, sidebar_visible): # last toggle. sidebar_out = [no_update, no_update, no_update] - # --- Process RF --- - if rf_done: - _wlog("[parent] poll_completion: rf_future done; calling .result()") - try: - pipeline = _rf_future.result() - _wlog(f"[parent] poll_completion: .result() returned; keys={sorted(pipeline.keys())}") - except persistent_worker.JobCancelled: - _rf_future = None + if ref.kind == "rf": + if terminal.state is JobState.CANCELLED: add_log("RF computation cancelled by user.", "WARNING") rf_out = [ dmc.Alert( @@ -696,42 +1012,40 @@ def poll_completion(n_intervals, sidebar_visible): ), no_update, no_update, no_update, False, ] - except Exception as e: - msg = f"RF computation failed: {e}" + elif terminal.state is JobState.FAILED: + msg = f"RF computation failed: {payload.get('message', 'Unknown error')}" add_log(msg, "ERROR") - _rf_future = None rf_out = [dmc.Text(msg, c="red"), no_update, no_update, no_update, False] - notif = dmc.Notification(title="RF Computation Error", message=msg, - color="red", action="show", autoClose=6000, - id=notif_id()) + notif = dmc.Notification( + title="RF Computation Error", + message=msg, + color="red", + action="show", + autoClose=6000, + id=f"rf-terminal-{ref.job_id}", + ) else: - _rf_future = None - rf_name = _rf_meta.get("name", "RF") - result_names = pipeline["result_names"] - file_breakdown = pipeline["file_breakdown"] - groups_per_file = pipeline["groups_per_file"] - elapsed = pipeline["total_elapsed"] - compute_elapsed = pipeline["compute_elapsed"] - is_rooted = pipeline.get("is_rooted", - _rf_meta.get("is_rooted", True)) - # Matrix already saved to disk by the worker — just register it - register_distmat(rf_name, result_names, _rf_meta["save_path"], - file_breakdown=file_breakdown, - groups_per_file=groups_per_file, - is_rooted=is_rooted) - add_log(f"Stored RF distance matrix as '{rf_name}' ({len(result_names)}x{len(result_names)})") - add_log(f"RF pipeline took {elapsed:.2f}s (rapidtrees compute {compute_elapsed:.2f}s)") + rf_name = payload["rf_name"] + n_trees = payload["n_trees"] + elapsed = payload["elapsed"] rf_out = [ dmc.Alert(title=f"RF Distance Matrix ({rf_name})", - children=dmc.Text(f"{len(result_names)} x {len(result_names)} trees", size="sm"), + children=dmc.Text(f"{n_trees} x {n_trees} trees", size="sm"), color="green", variant="light"), - get_distmat_index(), + payload["distmat_index"], False, False, False, ] notif = dmc.Notification( title=f"RF Distances Computed ({rf_name})", - message=f"Computed {len(result_names)}x{len(result_names)} RF distance matrix in {elapsed:.2f}s.", - color="green", action="show", autoClose=3000, id=notif_id()) + message=( + f"Computed {n_trees}x{n_trees} RF distance matrix " + f"in {elapsed:.2f}s." + ), + color="green", + action="show", + autoClose=3000, + id=f"rf-terminal-{ref.job_id}", + ) # Auto-collapse the sidebar to give the results area room. if sidebar_visible: sidebar_out = [ @@ -741,12 +1055,8 @@ def poll_completion(n_intervals, sidebar_visible): "collapsed": {"mobile": True}}, ] - # --- Process MDS --- - if mds_done: - try: - embedding_list, elapsed = _mds_future.result() - except persistent_worker.JobCancelled: - _mds_future = None + else: + if terminal.state is JobState.CANCELLED: add_log("MDS computation cancelled by user.", "WARNING") mds_out = [ no_update, @@ -757,66 +1067,60 @@ def poll_completion(n_intervals, sidebar_visible): ), no_update, no_update, False, ] - except Exception as e: - msg = f"MDS computation failed: {e}" + elif terminal.state is JobState.FAILED: + msg = f"MDS computation failed: {payload.get('message', 'Unknown error')}" add_log(msg, "ERROR") - _mds_future = None mds_out = [no_update, dmc.Text(msg, c="red"), no_update, no_update, False] - if notif is no_update: - notif = dmc.Notification(title="MDS Error", message=msg, color="red", - action="show", autoClose=6000, id=notif_id()) + notif = dmc.Notification( + title="MDS Error", + message=msg, + color="red", + action="show", + autoClose=6000, + id=f"mds-terminal-{ref.job_id}", + ) else: - _mds_future = None - tree_names = [str(n) for n in _mds_meta["tree_names"]] - selected_distmat = _mds_meta["selected_distmat"] - n_components = _mds_meta["n_components"] - - mdscols = [f"MDS{i+1}" for i in range(n_components)] - mds_df = pd.DataFrame(embedding_list, columns=mdscols) - mds_df["tree"] = tree_names - mds_df["group"] = mds_df["tree"].apply(extract_group) - mds_df["group"] = mds_df["group"].astype(str) - group_mapping = {val: idx for idx, val in enumerate(sorted(mds_df["group"].unique()))} - mds_df["group_col"] = mds_df["group"].map(group_mapping) - mds_df["treenum"] = mds_df.groupby("group").cumcount() + 1 - mds_df["size"] = 6 - mds_filename = _mds_export_filename(selected_distmat) - mds_df["file"] = mds_filename - - metadata = { - "filename": mds_filename, "source_distmat": selected_distmat, - "rows": len(mds_df), "dimensions": mdscols, - "groups": mds_df["group"].unique().tolist(), - "MIN_TREENUM": int(mds_df["treenum"].min()), - "MAX_TREENUM": int(mds_df["treenum"].max()), - } - mds_entry = {"metadata": metadata, "data": mds_df.to_dict("records")} - - # Store full result server-side, send only metadata through dcc.Store - store_mds_result(mds_filename, mds_entry) - - n_groups = len(mds_df["group"].unique()) - add_log(f"PCoA completed in {elapsed:.2f}s: {len(mds_df)} points, {n_components}D, {n_groups} groups") - + mds_filename = payload["mds_filename"] + n_points = payload["n_points"] + n_components = payload["n_components"] + n_groups = payload["n_groups"] + elapsed = payload["elapsed"] mds_out = [ - get_mds_results_index(), # lightweight metadata only + payload["mds_results_index"], dmc.Alert(title="MDS Embedding Complete", - children=dmc.Text(f"{mds_filename}: {len(mds_df)} points, {n_components}D, {n_groups} groups", size="sm"), + children=dmc.Text(f"{mds_filename}: {n_points} points, {n_components}D, {n_groups} groups", size="sm"), color="green", variant="light"), False, {}, False, ] - if notif is no_update: - notif = dmc.Notification( - title="MDS Computed", - message=f"PCoA: {len(mds_df)} points, {n_components}D in {elapsed:.2f}s.", - color="green", action="show", autoClose=3000, id=notif_id()) - - # Re-enable interval if any jobs are still running - any_running = ((_rf_future is not None and not _rf_future.done()) or - (_mds_future is not None and not _mds_future.done())) - poll_disabled = not any_running - - return (*rf_out, *mds_out, notif, poll_disabled, *sidebar_out) + notif = dmc.Notification( + title="MDS Computed", + message=( + f"PCoA: {n_points} points, {n_components}D " + f"in {elapsed:.2f}s." + ), + color="green", + action="show", + autoClose=3000, + id=f"mds-terminal-{ref.job_id}", + ) + + applied = { + **ref.as_dict(), + "terminal_revision": terminal.revision, + "delivery_attempt": snapshot.delivery_attempt, + } + return (*rf_out, *mds_out, notif, True, applied, *sidebar_out) + + @callback( + Output("rf-mds-job-ack-store", "data"), + Input("rf-mds-applied-job-store", "data"), + prevent_initial_call=True, + ) + def acknowledge_terminal_job(applied_job): + if not isinstance(applied_job, dict): + return no_update + acknowledged = _ack_applied_job(applied_job) + return {**applied_job, "acknowledged": acknowledged} # ------ EXPORT CALLBACKS ------ diff --git a/src/treetracer/callbacks/sidebar.py b/src/treetracer/callbacks/sidebar.py index 76ee510..11abd67 100644 --- a/src/treetracer/callbacks/sidebar.py +++ b/src/treetracer/callbacks/sidebar.py @@ -634,6 +634,15 @@ def remove_file(n_clicks_list, stored_summaries): Output("clade-freq-consensus-tree-select-1", "value", allow_duplicate=True), Output("clade-freq-consensus-tree-select-2", "value", allow_duplicate=True), Output("clade-freq-output-paper", "style", allow_duplicate=True), + # RF/MDS lifecycle state. Clear the browser identities together with + # the server-side job record so no terminal replay can resurrect data. + Output("rf-job-store", "data", allow_duplicate=True), + Output("mds-job-store", "data", allow_duplicate=True), + Output("rf-mds-applied-job-store", "data", allow_duplicate=True), + Output("rf-mds-job-ack-store", "data", allow_duplicate=True), + Output("rf-progress-path", "data", allow_duplicate=True), + Output("mds-progress-path", "data", allow_duplicate=True), + Output("compute-poll-interval", "disabled", allow_duplicate=True), Input("clear-data-button", "n_clicks"), prevent_initial_call=True, ) @@ -643,12 +652,20 @@ def clear_uploads(n_clicks): # Cancel any in-flight subprocess job — descriptors point # into the DB we're about to wipe, and we don't want the # poll callbacks to write results based on stale state. - try: - from . import consensus_tree_compute, pseudo_ess_compute - consensus_tree_compute.reset() - pseudo_ess_compute.reset() - except Exception: - pass + from . import compute, consensus_tree_compute, pseudo_ess_compute + resetters = ( + ("RF/MDS", compute.reset), + ("consensus tree", consensus_tree_compute.reset), + ("Pseudo-ESS", pseudo_ess_compute.reset), + ) + for label, resetter in resetters: + try: + resetter() + except Exception as exc: + add_log( + f"Could not reset {label} computation: {exc}", + "WARNING", + ) # Clear all server-side distance matrices from disk clear_all_distmats() clear_all_mds_results() @@ -720,5 +737,12 @@ def clear_uploads(n_clicks): None, # clade-freq-consensus-tree-select-1.value None, # clade-freq-consensus-tree-select-2.value {"display": "none"}, # clade-freq-output-paper.style + None, # rf-job-store + None, # mds-job-store + None, # rf-mds-applied-job-store + None, # rf-mds-job-ack-store + None, # rf-progress-path + None, # mds-progress-path + True, # compute-poll-interval disabled ) - return (no_update,) * 33 + return (no_update,) * 40 diff --git a/src/treetracer/ui/navbar.py b/src/treetracer/ui/navbar.py index 3f91b7a..5ad45b2 100644 --- a/src/treetracer/ui/navbar.py +++ b/src/treetracer/ui/navbar.py @@ -86,6 +86,14 @@ def add_navbar(): dcc.Store(id="clade-freq-click-store", storage_type="memory"), # Background computation polling dcc.Interval(id="compute-poll-interval", interval=100, disabled=True), + # RF/MDS job identities and two-phase terminal delivery. + # The poll callback writes the applied marker atomically + # with the visible result; only then does the ack callback + # release the server-side sticky terminal event. + dcc.Store(id="rf-job-store", storage_type="memory"), + dcc.Store(id="mds-job-store", storage_type="memory"), + dcc.Store(id="rf-mds-applied-job-store", storage_type="memory"), + dcc.Store(id="rf-mds-job-ack-store", storage_type="memory"), # consensus tree computation polls on its OWN interval. Dash derives an # allow_duplicate output's disambiguation hash from the # callback's Input signature (dash/_utils.py), so sharing From 05395d63d702529da8dcc07707c9f848c01cdd39 Mon Sep 17 00:00:00 2001 From: Sam Hong <20244918+hongsamL@users.noreply.github.com> Date: Mon, 3 Aug 2026 15:29:58 +0200 Subject: [PATCH 4/9] consensus trees and tree ess using job manager --- src/test/test_managed_compute_jobs.py | 275 ++++++++++++ src/treetracer/background_jobs.py | 43 +- src/treetracer/callbacks/compute.py | 45 +- .../callbacks/consensus_tree_compute.py | 399 ++++++++++-------- src/treetracer/callbacks/diagnostics.py | 69 ++- .../callbacks/pseudo_ess_compute.py | 240 +++++------ src/treetracer/callbacks/sidebar.py | 25 +- src/treetracer/callbacks/treespace.py | 44 +- src/treetracer/callbacks/within_run.py | 55 ++- src/treetracer/ui/navbar.py | 25 +- 10 files changed, 781 insertions(+), 439 deletions(-) create mode 100644 src/test/test_managed_compute_jobs.py diff --git a/src/test/test_managed_compute_jobs.py b/src/test/test_managed_compute_jobs.py new file mode 100644 index 0000000..f6e5d47 --- /dev/null +++ b/src/test/test_managed_compute_jobs.py @@ -0,0 +1,275 @@ +"""Lifecycle integration tests for managed Pseudo-ESS and consensus jobs.""" + +from __future__ import annotations + +import time +from concurrent.futures import ThreadPoolExecutor +from functools import partial + +import numpy as np +import pytest +from dash import no_update + +from treetracer.background_jobs import JobManager, JobState +from treetracer.callbacks import ( + compute, + consensus_tree_compute, + pseudo_ess_compute, + sidebar, +) + + +def _wait_for_terminal(manager, ref, timeout=2.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + snapshot = manager.snapshot(ref) + if snapshot is not None and snapshot.terminal is not None: + return snapshot + time.sleep(0.005) + pytest.fail("background job did not become terminal") + + +def _registered_callback(name, register): + from dash import _callback + + def matches(): + found = [] + for callback_data in _callback.GLOBAL_CALLBACK_MAP.values(): + callback_fn = callback_data.get("callback") + if callback_fn is None: + continue + original = getattr(callback_fn, "__wrapped__", callback_fn) + if original.__name__ == name: + found.append(original) + return found + + found = matches() + if not found: + register() + found = matches() + assert found + return found[-1] + + +def test_pseudo_ess_submit_and_terminal_ui_replay_until_ack(monkeypatch): + manager = JobManager(id_factory=lambda: "pseudo-test-job") + worker_calls = [] + + def worker(job_name, **kwargs): + worker_calls.append((job_name, kwargs)) + return { + "results": [ + { + "label": "run-a", + "n_trees": 20, + "burnin_label": "5", + "min": 110.0, + "q2": 210.0, + "max": 310.0, + "n_refs_used": 20, + } + ] + } + + monkeypatch.setattr(pseudo_ess_compute, "job_manager", manager) + monkeypatch.setattr(compute, "job_manager", manager) + monkeypatch.setattr(pseudo_ess_compute, "add_log", lambda *_a, **_k: None) + monkeypatch.setattr(pseudo_ess_compute, "_wlog", lambda *_a, **_k: None) + monkeypatch.setattr(pseudo_ess_compute.persistent_worker, "submit_job", worker) + + with ThreadPoolExecutor(max_workers=1) as executor: + monkeypatch.setattr(pseudo_ess_compute, "_get_executor", lambda: executor) + ref = pseudo_ess_compute.submit_pseudo_ess_job( + distmat_path="/tmp/RF_001.npy", + names=["run-a/tree-1"], + requests=[ + { + "label": "run-a", + "indices": [0, 1, 2, 3], + "burnin_label": "5", + } + ], + n_refs=20, + ) + terminal = _wait_for_terminal(manager, ref) + + assert ref.kind == "pseudo_ess" + assert worker_calls[0][0] == "compute_pseudo_ess" + assert terminal.state is JobState.SUCCEEDED + assert terminal.terminal.payload["n_rows"] == 1 + + poll = _registered_callback( + "poll_pseudo_ess_completion", + pseudo_ess_compute.register_pseudo_ess_compute_callbacks, + ) + first = poll(10, ref.as_dict()) + second = poll(11, ref.as_dict()) + + assert len(first) == 4 + assert first[1] is False + assert first[2] is True + assert first[3]["terminal_revision"] == terminal.terminal.revision + assert first[3]["delivery_attempt"] == 1 + assert second[3]["delivery_attempt"] == 2 + assert manager.snapshot(ref).acknowledged is False + + assert compute._ack_applied_job(second[3]) is True + assert manager.snapshot(ref).acknowledged is True + assert manager.active_ref() is None + + +def test_consensus_finalizer_publishes_once_and_poll_only_replays(monkeypatch): + manager = JobManager(id_factory=lambda: "consensus-test-job") + cache_calls = [] + register_calls = [] + registry = [{"name": "RF_001_Between_consensus_tree_1"}] + + def cache(nexus_bytes): + cache_calls.append(nexus_bytes) + return "tree-uuid" + + def register(**kwargs): + register_calls.append(kwargs) + return {"name": "RF_001_Between_consensus_tree_1"} + + monkeypatch.setattr(consensus_tree_compute, "job_manager", manager) + monkeypatch.setattr(compute, "job_manager", manager) + monkeypatch.setattr(consensus_tree_compute, "add_log", lambda *_a, **_k: None) + monkeypatch.setattr(consensus_tree_compute._state, "cache_consensus_tree", cache) + monkeypatch.setattr( + consensus_tree_compute._state, + "register_consensus_tree", + register, + ) + monkeypatch.setattr( + consensus_tree_compute._state, + "get_consensus_tree_registry", + lambda: registry, + ) + + context = consensus_tree_compute._ConsensusFinalizationContext( + source_distmat="RF_001", + mode="Between", + run=None, + selection=[["run-a", 1], ["run-b", 2]], + tree_names=("run-a/tree-1", "run-b/tree-2"), + coord_by_tree_name={"run-b/tree-2": ("run-b", 2)}, + ) + result = { + "nexus_bytes": b"#NEXUS\n", + "consensus_tree_row": {"name": "run-b/tree-2", "metadata": {}}, + "log_clade_credibility": -2.5, + "counts": np.array([1, 2], dtype=np.int32), + "cols_in_consensus_tree": frozenset({1}), + "missing_taxa": set(), + } + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "consensus", + lambda: result, + metadata={ + "store_target": consensus_tree_compute._TREESPACE_TARGET, + }, + finalizer=partial( + consensus_tree_compute._finalize_consensus_tree_job, + context=context, + ), + ) + terminal = _wait_for_terminal(manager, ref) + + assert terminal.state is JobState.SUCCEEDED + assert cache_calls == [b"#NEXUS\n"] + assert len(register_calls) == 1 + assert register_calls[0]["consensus_tree"] == { + "group": "run-b", + "treenum": 2, + "tree_name": "run-b/tree-2", + } + assert "nexus_bytes" not in terminal.terminal.payload + assert "counts" not in terminal.terminal.payload + + poll = _registered_callback( + "poll_consensus_tree_completion", + consensus_tree_compute.register_consensus_tree_compute_callbacks, + ) + first = poll(10, ref.as_dict()) + second = poll(11, ref.as_dict()) + + assert len(first) == 12 + assert first[0] == { + "uuid": "tree-uuid", + "name": "RF_001_Between_consensus_tree_1", + } + assert first[1] is no_update + assert first[2] == registry + assert first[3:5] == (False, False) + assert first[5] is False + assert first[6] is no_update + assert first[7] == [] + assert first[10] is True + assert first[11]["delivery_attempt"] == 1 + assert second[11]["delivery_attempt"] == 2 + assert len(cache_calls) == 1 + assert len(register_calls) == 1 + + assert compute._ack_applied_job(second[11]) is True + assert manager.snapshot(ref).acknowledged is True + + +def test_consensus_finalization_failure_reenables_origin_button(monkeypatch): + manager = JobManager(id_factory=lambda: "consensus-failed-job") + monkeypatch.setattr(consensus_tree_compute, "job_manager", manager) + monkeypatch.setattr(consensus_tree_compute, "add_log", lambda *_a, **_k: None) + monkeypatch.setattr( + consensus_tree_compute._state, + "cache_consensus_tree", + lambda _value: pytest.fail("invalid result must not be cached"), + ) + + context = consensus_tree_compute._ConsensusFinalizationContext( + source_distmat="RF_001", + mode="Within", + run="run-a", + selection=[["run-a", 1]], + tree_names=("run-a/tree-1",), + coord_by_tree_name={}, + ) + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "consensus", + lambda: {"missing_taxa": {"Taxon C", "Taxon A"}}, + metadata={ + "store_target": consensus_tree_compute._WITHIN_RUN_TARGET, + }, + finalizer=partial( + consensus_tree_compute._finalize_consensus_tree_job, + context=context, + ), + ) + terminal = _wait_for_terminal(manager, ref) + + assert terminal.state is JobState.FAILED + assert terminal.terminal.payload["stage"] == "finalize" + assert "Taxon A, Taxon C" in terminal.terminal.payload["message"] + + poll = _registered_callback( + "poll_consensus_tree_completion", + consensus_tree_compute.register_consensus_tree_compute_callbacks, + ) + output = poll(1, ref.as_dict()) + assert output[3:5] == (False, False) + assert output[5] is no_update + assert output[6] is False + assert output[10] is True + assert output[11]["job_id"] == ref.job_id + assert output[9].id == f"consensus-terminal-{ref.job_id}" + + +def test_clear_data_idle_branch_matches_managed_lifecycle_outputs(): + clear_uploads = _registered_callback( + "clear_uploads", + sidebar.register_sidebar_callbacks, + ) + assert len(clear_uploads(None)) == 45 diff --git a/src/treetracer/background_jobs.py b/src/treetracer/background_jobs.py index d921e85..0042a17 100644 --- a/src/treetracer/background_jobs.py +++ b/src/treetracer/background_jobs.py @@ -152,6 +152,17 @@ def as_dict(self) -> dict[str, Any]: ) return value + def terminal_delivery_marker(self) -> dict[str, Any]: + """Return the browser marker for this terminal delivery attempt.""" + + if self.terminal is None: + raise ValueError("a non-terminal snapshot has no delivery marker") + return { + **self.ref.as_dict(), + "terminal_revision": self.terminal.revision, + "delivery_attempt": self.delivery_attempt, + } + Finalizer = Callable[[JobRef, Any], Mapping[str, Any] | None] @@ -200,6 +211,12 @@ def __init__( self._clock = clock self._id_factory = id_factory or (lambda: uuid.uuid4().hex) self._lock = RLock() + # Domain finalizers run outside ``_lock`` but inside this barrier. + # Reset removes the record first, then crosses the same barrier. This + # guarantees that Clear Data either waits for an already-started + # publication or makes a not-yet-started publication reject the stale + # identity before doing domain work. + self._publication_lock = RLock() self._records: dict[str, _JobRecord] = {} self._next_generation = 1 @@ -434,6 +451,8 @@ def invalidate(self, ref: JobRef) -> bool: if future is not None: future.cancel() + with self._publication_lock: + pass return True def _mark_running(self, ref: JobRef) -> bool: @@ -475,12 +494,24 @@ def _run_job( return try: - payload = {} if finalizer is None else finalizer(ref, result) - if payload is None: - payload = {} - if not isinstance(payload, Mapping): - raise TypeError("job finalizer must return a mapping or None") - payload = copy.deepcopy(dict(payload)) + with self._publication_lock: + # Invalidation may have won after finalization was claimed but + # before this wrapper acquired the publication barrier. + with self._lock: + record = self._matching_record_locked(ref) + if ( + record is None + or record.state is not JobState.FINALIZING + ): + return + payload = {} if finalizer is None else finalizer(ref, result) + if payload is None: + payload = {} + if not isinstance(payload, Mapping): + raise TypeError( + "job finalizer must return a mapping or None" + ) + payload = copy.deepcopy(dict(payload)) except BaseException as exc: self._finish_exception(ref, exc, stage="finalize") return diff --git a/src/treetracer/callbacks/compute.py b/src/treetracer/callbacks/compute.py index 634f83d..93049dd 100644 --- a/src/treetracer/callbacks/compute.py +++ b/src/treetracer/callbacks/compute.py @@ -1,7 +1,6 @@ from dash import html, callback, Input, Output, State, no_update, ALL, ctx import dash_mantine_components as dmc from concurrent.futures import ThreadPoolExecutor -from threading import RLock import pandas as pd from ..logger import add_log, notif_id @@ -32,7 +31,6 @@ # work runs in the worker subprocess, keeping the pywebview/Dash process # responsive while compute jobs are in flight. _executor = None -_publication_lock = RLock() def _mds_export_filename(source_distmat): @@ -55,8 +53,6 @@ def _job_ref_from_store(data, *, expected_kind=None): return None if expected_kind is not None and ref.kind != expected_kind: return None - if ref.kind not in {"rf", "mds"}: - return None return ref @@ -202,17 +198,9 @@ def _rf_pipeline(selected_files, save_path, rf_name, is_rooted): return result -def _job_can_publish(ref): - snapshot = job_manager.snapshot(ref) - return snapshot is not None and snapshot.state is JobState.FINALIZING - - -def _finalize_rf_job(ref, pipeline): +def _finalize_rf_job(_ref, pipeline): """Persist one RF result and return its small terminal UI payload.""" - with _publication_lock: - if not _job_can_publish(ref): - return {} - return _publish_rf_result(pipeline) + return _publish_rf_result(pipeline) def _publish_rf_result(pipeline): @@ -272,12 +260,9 @@ def _mds_pipeline( } -def _finalize_mds_job(ref, result): +def _finalize_mds_job(_ref, result): """Persist one MDS result and return its small terminal UI payload.""" - with _publication_lock: - if not _job_can_publish(ref): - return {} - return _publish_mds_result(result) + return _publish_mds_result(result) def _publish_mds_result(result): @@ -349,11 +334,6 @@ def reset(): return persistent_worker.cancel_current_job() job_manager.invalidate(active) - # If a finalizer acquired the publication lock just before invalidation, - # let it finish before Clear Data removes published state. If it acquires - # the lock afterward, its identity re-check rejects the invalidated job. - with _publication_lock: - pass def register_compute_callbacks(): @@ -381,7 +361,6 @@ def handle_compute_stop(_stop_clicks, rf_progress_path, mds_progress_path): active = job_manager.active_ref() managed_cancelled = ( active is not None - and active.kind in {"rf", "mds"} and job_manager.mark_cancelled( active, message="Computation cancelled by user", @@ -514,7 +493,7 @@ def open_export_drawer(_rf_clicks, _mds_clicks): State({"type": "compute-tree-checkbox", "index": ALL}, "checked"), State({"type": "compute-tree-checkbox", "index": ALL}, "id"), State("tree-offset-store", "data"), - State("rf-mds-applied-job-store", "data"), + State("compute-applied-job-store", "data"), prevent_initial_call=True, ) def handle_compute_rf( @@ -852,7 +831,7 @@ def show_distmat_info(selected, distmat_data): Output("mds-job-store", "data"), Input("compute-mds-button", "n_clicks"), State("mds-distmat-select", "value"), - State("rf-mds-applied-job-store", "data"), + State("compute-applied-job-store", "data"), prevent_initial_call=True, ) def handle_compute_mds(n_clicks, selected_distmat, applied_job): @@ -948,7 +927,7 @@ def handle_compute_mds(n_clicks, selected_distmat, applied_job): # acknowledgement callback to consume the retained terminal event. Output("notifications-container", "children", allow_duplicate=True), Output("compute-poll-interval", "disabled", allow_duplicate=True), - Output("rf-mds-applied-job-store", "data"), + Output("compute-applied-job-store", "data", allow_duplicate=True), # Auto-collapse sidebar on RF success (3 outputs mirror the # shell sidebar-toggle callback's outputs). Output("navbar", "style", allow_duplicate=True), @@ -1104,16 +1083,12 @@ def poll_completion( id=f"mds-terminal-{ref.job_id}", ) - applied = { - **ref.as_dict(), - "terminal_revision": terminal.revision, - "delivery_attempt": snapshot.delivery_attempt, - } + applied = snapshot.terminal_delivery_marker() return (*rf_out, *mds_out, notif, True, applied, *sidebar_out) @callback( - Output("rf-mds-job-ack-store", "data"), - Input("rf-mds-applied-job-store", "data"), + Output("compute-job-ack-store", "data"), + Input("compute-applied-job-store", "data"), prevent_initial_call=True, ) def acknowledge_terminal_job(applied_job): diff --git a/src/treetracer/callbacks/consensus_tree_compute.py b/src/treetracer/callbacks/consensus_tree_compute.py index 19a6c0e..c037eb3 100644 --- a/src/treetracer/callbacks/consensus_tree_compute.py +++ b/src/treetracer/callbacks/consensus_tree_compute.py @@ -1,4 +1,4 @@ -"""Shared consensus tree compute dispatch + polling. +"""Managed consensus-tree dispatch, publication, and terminal polling. Both the Between-run (``treespace``) and Within-run (``within_run``) tabs have a "View consensus tree" button. They used to each call @@ -10,12 +10,10 @@ * Click → enqueue a consensus tree compute job on the persistent worker, show a loading overlay over the active tab, disable the View consensus tree button. -* Wait → ``compute-poll-interval`` ticks at 100 ms; this module's - ``poll_consensus_tree_completion`` runs once per tick. When the future is done, - it caches the NEXUS bytes, calls ``state.register_consensus_tree(...)``, fans - out the result to the right ``*-view-consensus-tree-store`` (which triggers the - per-tab clientside ``window.open(/peartree/)`` callback), and - dismisses the loading overlay. +* Wait → a dedicated interval reads a sticky ``JobManager`` snapshot. The + success finalizer caches the NEXUS bytes and registers the tree exactly once; + polling only renders the retained terminal payload and routes it to the + originating tab. The per-tab callbacks ``view_consensus_tree`` in ``treespace.py`` and ``within_run.py`` shrink to ~30 lines each — they're only responsible @@ -25,42 +23,129 @@ from __future__ import annotations -from concurrent.futures import Future -from typing import Any, Dict, Optional +import copy +from dataclasses import dataclass +from functools import partial +from typing import Any import dash_mantine_components as dmc -from dash import Input, Output, State, callback, no_update +from dash import Input, Output, callback, no_update from .. import state as _state -from ..logger import add_log, notif_id +from ..background_jobs import JobRef, JobState, job_manager +from ..logger import add_log from ..consensus_tree import extract_log_posterior from . import persistent_worker -from .compute import _get_executor +from .compute import _get_executor, _job_ref_from_store -# ── Module state ─────────────────────────────────────────────────────── -# A single in-flight consensus tree job at a time. ``_consensus_tree_meta`` carries the -# tab-specific context the polling callback needs to finalise the -# registry entry and route the result to the right view-consensus-tree-store. +_TREESPACE_TARGET = "treespace-view-consensus-tree-store" +_WITHIN_RUN_TARGET = "within-run-view-consensus-tree-store" +_STORE_TARGETS = frozenset({_TREESPACE_TARGET, _WITHIN_RUN_TARGET}) -_consensus_tree_future: Optional[Future] = None -_consensus_tree_meta: Dict[str, Any] = {} + +@dataclass(frozen=True, slots=True) +class _ConsensusFinalizationContext: + """Parent-only state needed for exactly-once consensus publication.""" + + source_distmat: str + mode: str + run: str | None + selection: list[Any] + tree_names: tuple[str, ...] + coord_by_tree_name: dict[str, tuple[Any, int]] + + +class ConsensusTaxaAlignmentError(ValueError): + """Raised when selected source files cannot share one Translate table.""" def reset() -> None: - """Interrupt any in-flight consensus tree compute. Called by the sidebar's - Clear-Data callback so the persistent worker isn't still - processing a stale request after the DB is wiped. - - ``Future.cancel()`` only drops a not-yet-started future — it can't - stop a job already running in the worker. ``cancel_current_job()`` - kills the worker, which actually interrupts the compute.""" - global _consensus_tree_future, _consensus_tree_meta + """Invalidate an active consensus job before application state clears.""" + active = job_manager.active_ref() + if active is None or active.kind != "consensus": + return persistent_worker.cancel_current_job() - if _consensus_tree_future is not None: - _consensus_tree_future.cancel() - _consensus_tree_future = None - _consensus_tree_meta = {} + job_manager.invalidate(active) + + +def _missing_taxa_message(missing_taxa: Any) -> str: + missing = sorted(str(name) for name in missing_taxa) + sample = ", ".join(missing[:5]) + more = "…" if len(missing) > 5 else "" + return ( + f"Cannot align translate tables: taxa [{sample}{more}] are present " + "in some selected runs but not in the canonical Translate block." + ) + + +def _finalize_consensus_tree_job( + _ref: JobRef, + result: Any, + *, + context: _ConsensusFinalizationContext, +) -> dict[str, Any]: + """Publish one worker result and return a small browser payload.""" + if not isinstance(result, dict): + raise TypeError("consensus-tree worker returned a non-mapping result") + if result.get("missing_taxa"): + raise ConsensusTaxaAlignmentError( + _missing_taxa_message(result["missing_taxa"]) + ) + + nexus_bytes = result.get("nexus_bytes") + consensus_tree_row = result.get("consensus_tree_row") + if not isinstance(nexus_bytes, (bytes, bytearray)): + raise TypeError("consensus-tree worker result is missing NEXUS bytes") + if not isinstance(consensus_tree_row, dict): + raise TypeError("consensus-tree worker result is missing its tree row") + + consensus_tree_name = str(consensus_tree_row["name"]) + coord = context.coord_by_tree_name.get(consensus_tree_name) + if coord is None: + consensus_tree_group, consensus_treenum = None, None + else: + consensus_tree_group, consensus_treenum = coord + + log_clade_cred = result.get("log_clade_credibility") + log_clade_cred = ( + None if log_clade_cred is None else float(log_clade_cred) + ) + consensus_tree_log_posterior = extract_log_posterior(consensus_tree_row) + + # This finalizer runs behind JobManager's publication barrier and can only + # be claimed once. Poll retries never execute these state mutations again. + uid = _state.cache_consensus_tree(bytes(nexus_bytes)) + entry = _state.register_consensus_tree( + source_distmat=context.source_distmat, + mode=context.mode, + run=context.run, + uuid=uid, + consensus_tree={ + "group": consensus_tree_group, + "treenum": consensus_treenum, + "tree_name": consensus_tree_name, + }, + selection=context.selection, + log_clade_credibility=log_clade_cred, + consensus_tree_log_posterior=consensus_tree_log_posterior, + tree_names=context.tree_names, + counts=result.get("counts"), + cols_in_consensus_tree=result.get("cols_in_consensus_tree"), + ) + registered_name = entry["name"] + n_trees = len(context.tree_names) + add_log( + f"Cached consensus tree '{consensus_tree_name}' " + f"(from {n_trees} selected) as {uid}; registered as {registered_name}" + ) + return { + "uuid": uid, + "name": registered_name, + "consensus_tree_name": consensus_tree_name, + "n_trees": n_trees, + "mode": context.mode, + } def submit_consensus_tree_job( @@ -69,10 +154,10 @@ def submit_consensus_tree_job( source_distmat: str, mode: str, selection: list, - run: Optional[str], - consensus_tree_coord_by_tree_name: Dict[str, tuple], + run: str | None, + consensus_tree_coord_by_tree_name: dict[str, tuple], store_target: str, -) -> None: +) -> JobRef: """Enqueue a consensus tree compute job. Called by both tab callbacks. Args: @@ -94,13 +179,15 @@ def submit_consensus_tree_job( ``"within-run-view-consensus-tree-store"`` — tells the polling callback which tab's clientside ``window.open`` to fire. """ - global _consensus_tree_future, _consensus_tree_meta + if not matched_records: + raise ValueError("matched_records must not be empty") + if store_target not in _STORE_TARGETS: + raise ValueError(f"unsupported consensus-tree store target: {store_target}") tree_service = _get_tree_service() db_manager = tree_service.db_manager # ── Build the kwargs the worker needs ────────────────────────────── - canonical_source = matched_records[0]["file_source"] unique_sources = list(dict.fromkeys(r["file_source"] for r in matched_records)) snapshots_path = _state.get_snapshots_path(source_distmat) @@ -120,15 +207,18 @@ def submit_consensus_tree_job( if fs in getattr(db_manager, "_source_preambles", {}) } - _consensus_tree_meta = { - "mode": mode, - "run": run, - "source_distmat": source_distmat, - "selection": selection, - "tree_names": [r["name"] for r in matched_records], - "consensus_tree_coord_by_tree_name": consensus_tree_coord_by_tree_name, - "store_target": store_target, - } + tree_names = tuple(str(r["name"]) for r in matched_records) + context = _ConsensusFinalizationContext( + mode=str(mode), + run=None if run is None else str(run), + source_distmat=str(source_distmat), + selection=copy.deepcopy(selection), + tree_names=tree_names, + coord_by_tree_name={ + str(name): (coord[0], int(coord[1])) + for name, coord in consensus_tree_coord_by_tree_name.items() + }, + ) # Lift the rooting flag off the distmat registry. Defaults to True # for pre-feature distmats; the worker decides based on this whether @@ -140,7 +230,9 @@ def submit_consensus_tree_job( f"({len(matched_records)} trees, source {source_distmat}, " f"{'rooted' if distmat_is_rooted else 'unrooted+midpoint-root'} mode)..." ) - _consensus_tree_future = _get_executor().submit( + ref = job_manager.submit( + _get_executor(), + "consensus", persistent_worker.submit_job, "compute_consensus_tree", matched_records=matched_records, @@ -151,7 +243,19 @@ def submit_consensus_tree_job( source_file_paths=source_file_paths, source_preambles=source_preambles, is_rooted=distmat_is_rooted, + metadata={ + "display_name": f"{mode} consensus tree", + "mode": mode, + "source_distmat": source_distmat, + "store_target": store_target, + }, + finalizer=partial( + _finalize_consensus_tree_job, + context=context, + ), + cancel_exceptions=(persistent_worker.JobCancelled,), ) + return ref def _get_tree_service(): @@ -163,176 +267,103 @@ def _get_tree_service(): def register_consensus_tree_compute_callbacks(): @callback( - # View-consensus-tree-stores: only one fires per completion (based on mode). + # View stores: only the originating tab changes. Output("treespace-view-consensus-tree-store", "data", allow_duplicate=True), Output("within-run-view-consensus-tree-store", "data", allow_duplicate=True), - # Registry — shared by both tabs. Output("consensus-tree-registry-store", "data", allow_duplicate=True), - # Loading overlays — flipped off on both tabs so we don't leave - # a stale overlay on whichever tab the user might have switched - # away from mid-compute. + # Both overlays are dismissed in case the user changed tabs. Output("treespace-loading-overlay", "visible", allow_duplicate=True), Output("within-run-loading-overlay", "visible", allow_duplicate=True), - # View consensus tree buttons — re-enabled on completion. + # Only the originating button is re-enabled. Output("treespace-view-consensus-tree", "disabled", allow_duplicate=True), Output("within-run-view-consensus-tree", "disabled", allow_duplicate=True), - # Selection-ring stores — cleared so the orange "selected" - # marker drops off the MDS view once the compute returns. + # The originating selection is cleared on success. Output("treespace-selected-trees-store", "data", allow_duplicate=True), Output("within-run-selected-trees-store", "data", allow_duplicate=True), - # Notification + the consensus-tree-only poll interval. This callback polls - # ``consensus-tree-poll-interval`` rather than the shared - # ``compute-poll-interval`` so it does NOT share an Input — and - # therefore an allow_duplicate disambiguation hash — with the - # RF/MDS ``poll_completion`` callback. Sharing the input made - # both callbacks emit the identical - # ``compute-poll-interval.disabled`` / ``notifications-container`` - # tokens, which the dash-renderer rejects as duplicates. The hash - # is derived from the Input signature (see dash/_utils.py). Output("notifications-container", "children", allow_duplicate=True), Output("consensus-tree-poll-interval", "disabled", allow_duplicate=True), + Output("compute-applied-job-store", "data", allow_duplicate=True), Input("consensus-tree-poll-interval", "n_intervals"), + Input("consensus-job-store", "data"), prevent_initial_call=True, ) - def poll_consensus_tree_completion(_n): - global _consensus_tree_future - if _consensus_tree_future is None or not _consensus_tree_future.done(): - return (no_update,) * 11 + def poll_consensus_tree_completion(n_intervals, job_data): + ref = _job_ref_from_store(job_data, expected_kind="consensus") + if ref is None: + return (no_update,) * 12 - future = _consensus_tree_future - meta = _consensus_tree_meta - _consensus_tree_future = None # consume the future before any further IO + snapshot = job_manager.snapshot_for_delivery(ref) + if snapshot is None or snapshot.acknowledged: + return (no_update,) * 12 + if snapshot.terminal is None: + return (no_update,) * 12 - mode = meta.get("mode", "Between") - store_target = meta.get("store_target") + terminal = snapshot.terminal + payload = terminal.payload + store_target = snapshot.metadata.get("store_target") - # Idle outputs everywhere except the bits we definitely flip. out_treespace_store = no_update out_within_store = no_update out_registry = no_update - # Dismiss BOTH overlays on completion (cheap, and the user might - # have switched tabs mid-compute). The View buttons, though, are - # scoped to the originating tab in the success path below — - # enabling both here lights up the OTHER tab's View button for a - # consensus tree it can't show (the cross-tab leak bug). out_treespace_overlay = False out_within_overlay = False out_treespace_btn = no_update out_within_btn = no_update out_treespace_sel = no_update out_within_sel = no_update + if store_target == _TREESPACE_TARGET: + out_treespace_btn = False + elif store_target == _WITHIN_RUN_TARGET: + out_within_btn = False - try: - result = future.result() - except persistent_worker.JobCancelled: - # User Stop — dismiss the overlays + re-enable the buttons - # (already set above). The cancel callback showed the - # notification, so don't stack another one here. - add_log("consensus tree computation cancelled by user.", "WARNING") - return (out_treespace_store, out_within_store, out_registry, - out_treespace_overlay, out_within_overlay, - out_treespace_btn, out_within_btn, - out_treespace_sel, out_within_sel, - no_update, True) - except Exception as e: - msg = f"Consensus tree computation failed: {e}" - add_log(msg, "ERROR") + notif = no_update + if terminal.state is JobState.CANCELLED: + if snapshot.delivery_attempt == 1: + add_log("Consensus tree computation cancelled by user.", "WARNING") + elif terminal.state is JobState.FAILED: + message = str(payload.get("message", "Unknown error")) + if snapshot.delivery_attempt == 1: + add_log(f"Consensus tree computation failed: {message}", "ERROR") notif = dmc.Notification( - title="Consensus tree Error", message=str(e), - color="red", action="show", autoClose=6000, id=notif_id(), + title="Consensus tree Error", + message=message, + color="red", + action="show", + autoClose=8000, + id=f"consensus-terminal-{ref.job_id}", ) - return (out_treespace_store, out_within_store, out_registry, - out_treespace_overlay, out_within_overlay, - out_treespace_btn, out_within_btn, - out_treespace_sel, out_within_sel, - notif, True) - - if result.get("missing_taxa"): - missing = sorted(result["missing_taxa"]) - sample = ", ".join(missing[:5]) - more = "…" if len(missing) > 5 else "" + else: + view_payload = {"uuid": payload["uuid"], "name": payload["name"]} + if store_target == _TREESPACE_TARGET: + out_treespace_store = view_payload + out_treespace_sel = [] + elif store_target == _WITHIN_RUN_TARGET: + out_within_store = view_payload + out_within_sel = [] + out_registry = _state.get_consensus_tree_registry() notif = dmc.Notification( - title="Consensus tree Error", + title="Consensus Tree Ready", message=( - f"Cannot align translate tables: taxa [{sample}{more}] " - "are present in some selected runs but not in the " - "canonical Translate block." + f"Consensus tree {payload['name']} computed from " + f"{payload['n_trees']} selected trees." ), - color="red", action="show", autoClose=8000, id=notif_id(), + color="green", + action="show", + autoClose=3000, + id=f"consensus-terminal-{ref.job_id}", ) - return (out_treespace_store, out_within_store, out_registry, - out_treespace_overlay, out_within_overlay, - out_treespace_btn, out_within_btn, - out_treespace_sel, out_within_sel, - notif, True) - - nexus_bytes = result["nexus_bytes"] - consensus_tree_row = result["consensus_tree_row"] - consensus_tree_name = consensus_tree_row["name"] - log_clade_cred = result["log_clade_credibility"] - - uid = _state.cache_consensus_tree(nexus_bytes) - - # consensus tree's (group, treenum) for the green-ring positioning. - coord = meta.get("consensus_tree_coord_by_tree_name", {}).get(consensus_tree_name) - if coord is not None: - consensus_tree_group, consensus_treenum = coord - else: - consensus_tree_group, consensus_treenum = None, None - - entry = _state.register_consensus_tree( - source_distmat=meta["source_distmat"], - mode=mode, - run=meta.get("run"), - uuid=uid, - consensus_tree={ - "group": consensus_tree_group, - "treenum": consensus_treenum, - "tree_name": consensus_tree_name, - }, - selection=meta["selection"], - log_clade_credibility=(None if log_clade_cred is None - else float(log_clade_cred)), - consensus_tree_log_posterior=extract_log_posterior(consensus_tree_row), - tree_names=meta["tree_names"], - counts=result["counts"], - cols_in_consensus_tree=result["cols_in_consensus_tree"], - ) - registered_name = entry["name"] - add_log( - f"Cached consensus tree '{consensus_tree_name}' (from {len(meta['tree_names'])} selected) " - f"as {uid}; registered as {registered_name}" - ) - # The rename modal opens next (see ``forward_compute_to_modal`` - # in ``callbacks/rename_consensus_tree.py``); PearTree only opens once the - # user clicks Save in the modal. Don't promise "opening in - # PearTree" here — the modal title makes the next step obvious. - notif = dmc.Notification( - title="Consensus Tree Ready", - message=( - f"Consensus tree {registered_name} computed from " - f"{len(meta['tree_names'])} selected trees." - ), - color="green", action="show", autoClose=3000, id=notif_id(), - ) - - # Route the {uuid, name} payload — and enable the View button — - # for the originating tab only. The OTHER tab's store and button - # stay untouched (no_update) so each tab governs its own state. - payload = {"uuid": uid, "name": registered_name} - if store_target == "treespace-view-consensus-tree-store": - out_treespace_store = payload - out_treespace_sel = [] - out_treespace_btn = False - else: - out_within_store = payload - out_within_sel = [] - out_within_btn = False - out_registry = _state.get_consensus_tree_registry() - - return (out_treespace_store, out_within_store, out_registry, - out_treespace_overlay, out_within_overlay, - out_treespace_btn, out_within_btn, - out_treespace_sel, out_within_sel, - notif, True) + return ( + out_treespace_store, + out_within_store, + out_registry, + out_treespace_overlay, + out_within_overlay, + out_treespace_btn, + out_within_btn, + out_treespace_sel, + out_within_sel, + notif, + True, + snapshot.terminal_delivery_marker(), + ) diff --git a/src/treetracer/callbacks/diagnostics.py b/src/treetracer/callbacks/diagnostics.py index 7893998..ade1b5b 100644 --- a/src/treetracer/callbacks/diagnostics.py +++ b/src/treetracer/callbacks/diagnostics.py @@ -7,14 +7,13 @@ import numpy as np import pandas as pd +from ..background_jobs import JobBusyError from ..logger import add_log, notif_id from ..db.tree_service import get_tree_service from ..ess.rf_trace import compute_rf_trace_data -from ..ess import compute_pseudo_ess from .. import state from ..theme import get_template from ..ui.widgets import stop_button -from ._helpers import _save_file_dialog def _build_rf_trace_fig(trace_df, ref_group, ref_position, burnin=0): @@ -586,15 +585,25 @@ def toggle_compute_pseudo_ess_button(selected_matrix, checks): Output("pseudo-ess-output", "children", allow_duplicate=True), Output("compute-pseudo-ess-button", "disabled", allow_duplicate=True), Output("compute-poll-interval", "disabled", allow_duplicate=True), + Output("pseudo-ess-job-store", "data"), Input("compute-pseudo-ess-button", "n_clicks"), State("diagnostics-distmat-select", "value"), State("ess-n-refs-input", "value"), State("ess-burnin-input", "value"), State({"type": "ess-run-checkbox", "index": ALL}, "checked"), State({"type": "ess-run-checkbox", "index": ALL}, "id"), + State("compute-applied-job-store", "data"), prevent_initial_call=True, ) - def compute_pseudo_ess_for_runs(n_clicks, selected_matrix, n_refs, burnin, checks, ids): + def compute_pseudo_ess_for_runs( + n_clicks, + selected_matrix, + n_refs, + burnin, + checks, + ids, + applied_job, + ): """Submit a Pseudo-ESS job to the persistent worker. Parent-side: validates input, bins trees per run, applies @@ -608,21 +617,29 @@ def compute_pseudo_ess_for_runs(n_clicks, selected_matrix, n_refs, burnin, check worker's response. """ from . import pseudo_ess_compute + from .compute import _ack_applied_job if not n_clicks or not selected_matrix: - return no_update, no_update, no_update + return (no_update,) * 4 + + # The visible terminal result and this marker arrived in one previous + # browser response. A fast next click can precede the dedicated ack + # callback, so acknowledge it idempotently before requesting a new job. + _ack_applied_job(applied_job) ticked = [i["index"] for i, c in zip(ids, checks) if c] if not ticked: return (dmc.Text("No runs selected.", c="dimmed", size="sm"), - no_update, no_update) + no_update, no_update, no_update) try: - names, _distmat_unused = state.load_distmat(selected_matrix) + # Only labels are needed to form per-run row indices. Avoid loading + # the full n×n matrix into the GUI process; the worker reads it once. + names = list(state.get_distmat_names(selected_matrix)) except KeyError: return (dmc.Text(f"Matrix {selected_matrix!r} is no longer available.", c="red", size="sm"), - no_update, no_update) + no_update, no_update, no_update) # Bucket row indices by group prefix once (matrix-row order # matches MCMC iteration order within each chain). @@ -670,18 +687,36 @@ def compute_pseudo_ess_for_runs(n_clicks, selected_matrix, n_refs, burnin, check return (dmc.Text( "Burn-in leaves fewer than 4 trees per run; nothing to compute.", c="dimmed", size="sm", - ), no_update, no_update) + ), no_update, no_update, no_update) # Hand off to the subprocess. The poll callback in # pseudo_ess_compute.py picks up the result and replaces the # spinner with the result table. - pseudo_ess_compute.submit_pseudo_ess_job( - distmat_path=str(state.get_distmat_file_path(selected_matrix)), - names=list(names), - requests=requests, - n_refs=n_refs_int, - seed=0, - ) + try: + job_ref = pseudo_ess_compute.submit_pseudo_ess_job( + distmat_path=str(state.get_distmat_file_path(selected_matrix)), + names=names, + requests=requests, + n_refs=n_refs_int, + seed=0, + ) + except JobBusyError as exc: + msg = ( + f"Another computation ({exc.active.kind.replace('_', ' ').upper()}) " + "is still finishing. Please wait for it to complete." + ) + add_log(msg, "WARNING") + return ( + dmc.Alert( + title="Computation already running", + children=dmc.Text(msg, size="sm"), + color="yellow", + variant="light", + ), + False, + no_update, + no_update, + ) spinner = dmc.Group([ dmc.Loader(size="sm", type="dots"), @@ -692,5 +727,5 @@ def compute_pseudo_ess_for_runs(n_clicks, selected_matrix, n_refs, burnin, check stop_button("ess"), ], gap="sm") - # spinner, button disabled, poll interval enabled. - return spinner, True, False + # Spinner, button disabled, polling enabled, and immutable job identity. + return spinner, True, False, job_ref.as_dict() diff --git a/src/treetracer/callbacks/pseudo_ess_compute.py b/src/treetracer/callbacks/pseudo_ess_compute.py index 88365a8..e9a7896 100644 --- a/src/treetracer/callbacks/pseudo_ess_compute.py +++ b/src/treetracer/callbacks/pseudo_ess_compute.py @@ -1,78 +1,68 @@ -"""Pseudo-ESS dispatch + polling. +"""Managed Pseudo-ESS dispatch, finalization, and terminal polling. -Same shape as ``consensus_tree_compute.py``: the Diagnostics tab's -"Compute Pseudo-ESS" click handler validates input, slices the -distmat into per-run index lists, and hands the work off to the -persistent worker subprocess via ``persistent_worker.submit_job``. - -The click handler returns immediately with a loading spinner in the -``pseudo-ess-output`` slot. ``poll_pseudo_ess_completion`` listens to -``compute-poll-interval`` and, on the worker's response, builds the -result table parent-side and writes it back. +The Diagnostics tab prepares small per-run index lists and submits one worker +request. ``JobManager`` owns the job identity and terminal state, so polling is +read-only and a completed result remains replayable until the browser confirms +that it applied the matching UI response. """ from __future__ import annotations -from concurrent.futures import Future -from typing import Any, Dict, List, Optional +from typing import Any import dash_mantine_components as dmc import numpy as np -from dash import Input, Output, callback, html, no_update +from dash import Input, Output, callback, no_update -from ..logger import add_log, notif_id +from ..background_jobs import JobRef, JobState, job_manager +from ..logger import add_log from .._worker_log import log as _wlog from . import persistent_worker -from .compute import _get_executor - - -_pseudo_ess_future: Optional[Future] = None -_pseudo_ess_meta: Dict[str, Any] = {} +from .compute import _get_executor, _job_ref_from_store def reset() -> None: - """Interrupt any in-flight Pseudo-ESS compute. Called by sidebar's - Clear-Data callback so the worker isn't still processing against - a distmat that no longer exists. - - ``Future.cancel()`` only drops a not-yet-started future — it can't - stop a job already running in the worker. ``cancel_current_job()`` - kills the worker, which actually interrupts the compute.""" - global _pseudo_ess_future, _pseudo_ess_meta + """Invalidate an active Pseudo-ESS job before application state clears.""" + active = job_manager.active_ref() + if active is None or active.kind != "pseudo_ess": + return persistent_worker.cancel_current_job() - if _pseudo_ess_future is not None: - _pseudo_ess_future.cancel() - _pseudo_ess_future = None - _pseudo_ess_meta = {} + job_manager.invalidate(active) + + +def _finalize_pseudo_ess_job( + _ref: JobRef, + result: Any, +) -> dict[str, Any]: + """Validate the worker response and retain only its small table payload.""" + if not isinstance(result, dict): + raise TypeError("Pseudo-ESS worker returned a non-mapping result") + rows = result.get("results") + if not isinstance(rows, list): + raise TypeError("Pseudo-ESS worker result is missing its results list") + add_log(f"Pseudo-ESS computed for {len(rows)} row(s).") + return {"results": rows, "n_rows": len(rows)} def submit_pseudo_ess_job( *, distmat_path: str, - names: List[str], - requests: List[Dict[str, Any]], + names: list[str], + requests: list[dict[str, Any]], n_refs: int, seed: int = 0, -) -> None: +) -> JobRef: """Enqueue a Pseudo-ESS job covering all ticked runs (+ optional Combined row) in a single subprocess round-trip. Args: distmat_path: path passed straight to the worker. - names: row/col labels (saved on ``_pseudo_ess_meta`` only for - potential future cancellation logging; the worker doesn't - read them). + names: row/column labels forwarded to the worker. requests: list of ``{"label", "indices", "burnin_label"}`` dicts the worker iterates over. n_refs: forwarded to ``compute_pseudo_ess``. seed: forwarded. """ - global _pseudo_ess_future, _pseudo_ess_meta - - _pseudo_ess_meta = { - "distmat_path": distmat_path, - "n_runs": len(requests), - } add_log( f"[Pseudo-ESS] Dispatching to persistent worker " f"({len(requests)} row(s), n_refs={n_refs})..." @@ -81,7 +71,9 @@ def submit_pseudo_ess_job( f"[parent] submit_pseudo_ess_job: {len(requests)} requests, " f"n_refs={n_refs}, distmat_path={distmat_path!r}" ) - _pseudo_ess_future = _get_executor().submit( + ref = job_manager.submit( + _get_executor(), + "pseudo_ess", persistent_worker.submit_job, "compute_pseudo_ess", distmat_path=distmat_path, @@ -89,11 +81,22 @@ def submit_pseudo_ess_job( requests=requests, n_refs=n_refs, seed=seed, + metadata={ + "display_name": "Pseudo-ESS", + "source_distmat_path": distmat_path, + "n_rows": len(requests), + }, + finalizer=_finalize_pseudo_ess_job, + cancel_exceptions=(persistent_worker.JobCancelled,), + ) + _wlog( + "[parent] submit_pseudo_ess_job: managed job created " + f"({ref.job_id}/generation-{ref.generation})" ) - _wlog(f"[parent] submit_pseudo_ess_job: future created (id={id(_pseudo_ess_future):x})") + return ref -def _ess_cell(v: Optional[float]): +def _ess_cell(v: float | None): """One stoplight-coloured ESS table cell. Thresholds mirror Lanfear's rule of thumb.""" if v is None or np.isnan(v): @@ -109,7 +112,7 @@ def _ess_cell(v: Optional[float]): ) -def _build_result_table(results: List[Dict[str, Any]]): +def _build_result_table(results: list[dict[str, Any]]): """Parent-side render of the per-row Pseudo-ESS table. The worker only returns plain dicts; this turns them into Mantine table rows.""" if not results: @@ -153,108 +156,57 @@ def _build_result_table(results: List[Dict[str, Any]]): def register_pseudo_ess_compute_callbacks(): - # Why only 2 outputs (not 4) — see the giant block comment below. @callback( Output("pseudo-ess-output", "children", allow_duplicate=True), Output("compute-pseudo-ess-button", "disabled", allow_duplicate=True), + Output("compute-poll-interval", "disabled", allow_duplicate=True), + Output("compute-applied-job-store", "data", allow_duplicate=True), Input("compute-poll-interval", "n_intervals"), + Input("pseudo-ess-job-store", "data"), prevent_initial_call=True, ) - def poll_pseudo_ess_completion(_n): - # ──────────────────────────────────────────────────────────── - # Output contention: WHY this callback returns ONLY 2 values - # (table + button.disabled) instead of also writing - # ``compute-poll-interval.disabled`` and - # ``notifications-container.children``. - # - # Three poll callbacks share ``Input("compute-poll-interval", - # "n_intervals")``: - # - # * poll_completion (compute.py — RF / MDS) - # * poll_consensus_tree_completion (consensus_tree_compute.py) - # * poll_pseudo_ess_completion (here) - # - # On every 100 ms tick they fire concurrently. When the ESS - # job finishes, this callback returns a real value for - # ``compute-poll-interval.disabled``, while the other two - # return ``no_update`` for everything (no RF/MDS/consensus tree running). - # ``compute-poll-interval.disabled``'s primary writer - # (no ``allow_duplicate``) lives in ``handle_compute_rf`` in - # compute.py. Dash 4.x has a bug — observed empirically on - # the Windows bundle and verified by file-log trace — where - # this secondary-write-during-primary-not-firing pattern - # causes the *entire callback batch* to be silently dropped: - # the table never reaches the GUI, the interval stays - # enabled, the button stays disabled. - # - # Dropping ``compute-poll-interval.disabled`` from THIS - # callback's Outputs removes the conflict. The interval just - # keeps firing forever after ESS finishes — every 100 ms it - # re-enters this callback, sees ``_pseudo_ess_future is None``, - # returns ``(no_update,) * 2``. Cost: a few µs per tick. The - # other two poll callbacks also no-op on idle ticks today. - # - # Dropping ``notifications-container.children`` for the same - # reason — many callbacks fight over it via allow_duplicate. - # The user sees the table appear, which is the only signal - # they need; a green toast on top is redundant. - # ──────────────────────────────────────────────────────────── - global _pseudo_ess_future - future_set = _pseudo_ess_future is not None - future_done = future_set and _pseudo_ess_future.done() - if future_set: - _wlog( - f"[parent] poll_pseudo_ess_completion: " - f"future_set={future_set}, future_done={future_done}, " - f"future_id={id(_pseudo_ess_future):x}" - ) - if not future_set or not future_done: - return no_update, no_update - - future = _pseudo_ess_future - _pseudo_ess_future = None - _wlog("[parent] poll_pseudo_ess_completion: future done, calling .result()") - - try: - result = future.result() - _wlog( - f"[parent] poll_pseudo_ess_completion: .result() returned; " - f"type={type(result).__name__}, " - f"keys={sorted(result.keys()) if isinstance(result, dict) else ''}" + def poll_pseudo_ess_completion(n_intervals, job_data): + ref = _job_ref_from_store(job_data, expected_kind="pseudo_ess") + if ref is None: + return (no_update,) * 4 + + snapshot = job_manager.snapshot_for_delivery(ref) + if snapshot is None or snapshot.acknowledged: + return (no_update,) * 4 + if snapshot.terminal is None: + if (n_intervals or 0) % 10 == 0: + _wlog( + f"[parent] poll_pseudo_ess_completion tick={n_intervals}: " + f"job={ref.job_id}/generation-{ref.generation}, " + f"state={snapshot.state.value}" + ) + return (no_update,) * 4 + + terminal = snapshot.terminal + applied = snapshot.terminal_delivery_marker() + if terminal.state is JobState.CANCELLED: + if snapshot.delivery_attempt == 1: + add_log("Pseudo-ESS computation cancelled by user.", "WARNING") + output = dmc.Alert( + title="Pseudo-ESS computation cancelled", + children=dmc.Text("Stopped before completion.", size="sm"), + color="gray", + variant="light", ) - except persistent_worker.JobCancelled: - _wlog("[parent] poll_pseudo_ess_completion: JobCancelled") - add_log("Pseudo-ESS computation cancelled by user.", "WARNING") - return ( - dmc.Alert( - title="Pseudo-ESS computation cancelled", - children=dmc.Text("Stopped before completion.", size="sm"), - color="gray", variant="light", - ), - False, # re-enable button + elif terminal.state is JobState.FAILED: + msg = ( + "Pseudo-ESS computation failed: " + f"{terminal.payload.get('message', 'Unknown error')}" ) - except Exception as e: - _wlog(f"[parent] poll_pseudo_ess_completion: .result() raised {type(e).__name__}: {e}") - msg = f"Pseudo-ESS computation failed: {e}" - add_log(msg, "ERROR") - return ( - dmc.Text(msg, c="red", size="sm"), - False, # re-enable button - ) - - rows = result.get("results", []) - _wlog(f"[parent] poll_pseudo_ess_completion: building result table for {len(rows)} rows") - add_log(f"Pseudo-ESS computed for {len(rows)} row(s).") - try: - table = _build_result_table(rows) - _wlog("[parent] poll_pseudo_ess_completion: result table built; returning") - except Exception as e: - _wlog( - f"[parent] poll_pseudo_ess_completion: _build_result_table raised " - f"{type(e).__name__}: {e}" - ) - raise - return ( - table, - False, # re-enable button - ) + if snapshot.delivery_attempt == 1: + add_log(msg, "ERROR") + output = dmc.Text(msg, c="red", size="sm") + else: + rows = list(terminal.payload.get("results", [])) + output = _build_result_table(rows) + + # The terminal event remains server-side until the applied marker is + # processed by the shared acknowledgement callback. If this whole Dash + # response is lost, the still-enabled interval requests the same event + # again with a new delivery-attempt marker. + return output, False, True, applied diff --git a/src/treetracer/callbacks/sidebar.py b/src/treetracer/callbacks/sidebar.py index 11abd67..9c74c23 100644 --- a/src/treetracer/callbacks/sidebar.py +++ b/src/treetracer/callbacks/sidebar.py @@ -634,15 +634,21 @@ def remove_file(n_clicks_list, stored_summaries): Output("clade-freq-consensus-tree-select-1", "value", allow_duplicate=True), Output("clade-freq-consensus-tree-select-2", "value", allow_duplicate=True), Output("clade-freq-output-paper", "style", allow_duplicate=True), - # RF/MDS lifecycle state. Clear the browser identities together with - # the server-side job record so no terminal replay can resurrect data. + # Managed-compute lifecycle state. Clear every browser identity together + # with the server-side job record so no terminal replay can resurrect + # data after the reset. Output("rf-job-store", "data", allow_duplicate=True), Output("mds-job-store", "data", allow_duplicate=True), - Output("rf-mds-applied-job-store", "data", allow_duplicate=True), - Output("rf-mds-job-ack-store", "data", allow_duplicate=True), + Output("pseudo-ess-job-store", "data", allow_duplicate=True), + Output("consensus-job-store", "data", allow_duplicate=True), + Output("compute-applied-job-store", "data", allow_duplicate=True), + Output("compute-job-ack-store", "data", allow_duplicate=True), Output("rf-progress-path", "data", allow_duplicate=True), Output("mds-progress-path", "data", allow_duplicate=True), Output("compute-poll-interval", "disabled", allow_duplicate=True), + Output("consensus-tree-poll-interval", "disabled", allow_duplicate=True), + Output("treespace-loading-overlay", "visible", allow_duplicate=True), + Output("within-run-loading-overlay", "visible", allow_duplicate=True), Input("clear-data-button", "n_clicks"), prevent_initial_call=True, ) @@ -739,10 +745,15 @@ def clear_uploads(n_clicks): {"display": "none"}, # clade-freq-output-paper.style None, # rf-job-store None, # mds-job-store - None, # rf-mds-applied-job-store - None, # rf-mds-job-ack-store + None, # pseudo-ess-job-store + None, # consensus-job-store + None, # compute-applied-job-store + None, # compute-job-ack-store None, # rf-progress-path None, # mds-progress-path True, # compute-poll-interval disabled + True, # consensus-tree-poll-interval disabled + False, # treespace-loading-overlay visible + False, # within-run-loading-overlay visible ) - return (no_update,) * 40 + return (no_update,) * 45 diff --git a/src/treetracer/callbacks/treespace.py b/src/treetracer/callbacks/treespace.py index 37ade29..ea1d318 100644 --- a/src/treetracer/callbacks/treespace.py +++ b/src/treetracer/callbacks/treespace.py @@ -3,12 +3,10 @@ from dash import html, callback, clientside_callback, Input, Output, Patch, State, no_update import dash_mantine_components as dmc import plotly.express as px -import plotly.graph_objects as go import pandas as pd from ..logger import add_log from ..state import get_mds_result -from ..theme import get_template from ..plot_utils import ( make_plot_grid, add_trace_multiplot_interleaved, placeholder_fig, retheme_figure, @@ -773,28 +771,34 @@ def export_selected_trees(n_clicks, selected_pairs, plot_config): # poll on a shared allow_duplicate output. Output("consensus-tree-poll-interval", "disabled", allow_duplicate=True), Output("notifications-container", "children", allow_duplicate=True), + Output("consensus-job-store", "data", allow_duplicate=True), Input("treespace-view-consensus-tree", "n_clicks"), State("treespace-selected-trees-store", "data"), State("plot-config-store", "data"), State("treespace-result-select", "value"), State("mds-result-store", "data"), + State("compute-applied-job-store", "data"), prevent_initial_call=True, ) def view_consensus_tree(n_clicks, selected_pairs, plot_config, - selected_key, results): + selected_key, results, applied_job): + from ..background_jobs import JobBusyError from ..logger import notif_id from ..db.tree_service import get_tree_service from . import consensus_tree_compute + from .compute import _ack_applied_job if not n_clicks or not selected_pairs or not plot_config: - return no_update, no_update, no_update, no_update + return (no_update,) * 5 + + _ack_applied_job(applied_job) def _err(msg, autoclose=5000): return (False, False, no_update, dmc.Notification( title="Consensus tree Error", message=msg, color="red", action="show", autoClose=autoclose, id=notif_id(), - )) + ), no_update) results = results or {} if not selected_key or selected_key not in results: @@ -834,22 +838,30 @@ def _err(msg, autoclose=5000): # green ring on the consensus tree's dot. consensus_tree_coord_by_tree_name = { row["tree"]: (row["group"], int(row["treenum"])) - for _, row in combined_df.iterrows() + for _, row in sel_df.iterrows() } - consensus_tree_compute.submit_consensus_tree_job( - matched_records=matched_records, - source_distmat=source_distmat, - mode="Between", - selection=[[g, int(t)] for g, t in selected_pairs], - run=None, - consensus_tree_coord_by_tree_name=consensus_tree_coord_by_tree_name, - store_target="treespace-view-consensus-tree-store", - ) + try: + job_ref = consensus_tree_compute.submit_consensus_tree_job( + matched_records=matched_records, + source_distmat=source_distmat, + mode="Between", + selection=[[g, int(t)] for g, t in selected_pairs], + run=None, + consensus_tree_coord_by_tree_name=( + consensus_tree_coord_by_tree_name + ), + store_target="treespace-view-consensus-tree-store", + ) + except JobBusyError as exc: + return _err( + f"Another computation ({exc.active.kind.replace('_', ' ').upper()}) " + "is still finishing. Please wait for it to complete." + ) # Return: overlay on, button disabled, polling enabled, no # notification yet (notification fires when compute finishes). - return True, True, False, no_update + return True, True, False, no_update, job_ref.as_dict() # The per-tab clientside ``window.open`` that used to live here is # gone — it was a duplicate of the one in within_run.py and the diff --git a/src/treetracer/callbacks/within_run.py b/src/treetracer/callbacks/within_run.py index e43b2d5..facb792 100644 --- a/src/treetracer/callbacks/within_run.py +++ b/src/treetracer/callbacks/within_run.py @@ -787,31 +787,44 @@ def export_selected_trees(n_clicks, selected_treenums, selected_key, selected_ru # consensus tree polling uses its own interval (see navbar.py). Output("consensus-tree-poll-interval", "disabled", allow_duplicate=True), Output("notifications-container", "children", allow_duplicate=True), + Output("consensus-job-store", "data", allow_duplicate=True), Input("within-run-view-consensus-tree", "n_clicks"), State("within-run-selected-trees-store", "data"), State("within-run-result-select", "value"), State("within-run-run-select", "value"), State("mds-result-store", "data"), + State("compute-applied-job-store", "data"), prevent_initial_call=True, ) - def view_consensus_tree(n_clicks, selected_treenums, selected_key, selected_run, results): + def view_consensus_tree( + n_clicks, + selected_treenums, + selected_key, + selected_run, + results, + applied_job, + ): + from ..background_jobs import JobBusyError from ..logger import notif_id from ..db.tree_service import get_tree_service from . import consensus_tree_compute + from .compute import _ack_applied_job if not n_clicks or not selected_treenums: - return no_update, no_update, no_update, no_update + return (no_update,) * 5 + + _ack_applied_job(applied_job) def _err(msg, autoclose=5000): return (False, False, no_update, dmc.Notification( title="Consensus tree Error", message=msg, color="red", action="show", autoClose=autoclose, id=notif_id(), - )) + ), no_update) mds_result = _get_active_result(selected_key, results) if not mds_result or not selected_run: - return no_update, no_update, no_update, no_update + return (no_update,) * 5 source_distmat = (mds_result.get("metadata") or {}).get("source_distmat") if not source_distmat: @@ -819,7 +832,7 @@ def _err(msg, autoclose=5000): df_run, _ = _filter_to_run(mds_result, selected_run) if df_run is None: - return no_update, no_update, no_update, no_update + return (no_update,) * 5 sel_df = df_run[df_run["treenum"].isin(selected_treenums)] tree_names = sel_df["tree"].tolist() if not tree_names: @@ -847,20 +860,30 @@ def _err(msg, autoclose=5000): # treenum. consensus_tree_coord_by_tree_name = { row["tree"]: (selected_run, int(row["treenum"])) - for _, row in df_run.iterrows() + for _, row in sel_df.iterrows() } - consensus_tree_compute.submit_consensus_tree_job( - matched_records=matched_records, - source_distmat=source_distmat, - mode="Within", - selection=[[selected_run, int(t)] for t in selected_treenums], - run=selected_run, - consensus_tree_coord_by_tree_name=consensus_tree_coord_by_tree_name, - store_target="within-run-view-consensus-tree-store", - ) + try: + job_ref = consensus_tree_compute.submit_consensus_tree_job( + matched_records=matched_records, + source_distmat=source_distmat, + mode="Within", + selection=[ + [selected_run, int(t)] for t in selected_treenums + ], + run=selected_run, + consensus_tree_coord_by_tree_name=( + consensus_tree_coord_by_tree_name + ), + store_target="within-run-view-consensus-tree-store", + ) + except JobBusyError as exc: + return _err( + f"Another computation ({exc.active.kind.replace('_', ' ').upper()}) " + "is still finishing. Please wait for it to complete." + ) - return True, True, False, no_update + return True, True, False, no_update, job_ref.as_dict() # The per-tab clientside ``window.open`` that used to live here is # gone; see the parallel note in ``treespace.py``. The diff --git a/src/treetracer/ui/navbar.py b/src/treetracer/ui/navbar.py index 5ad45b2..246fbf4 100644 --- a/src/treetracer/ui/navbar.py +++ b/src/treetracer/ui/navbar.py @@ -86,22 +86,19 @@ def add_navbar(): dcc.Store(id="clade-freq-click-store", storage_type="memory"), # Background computation polling dcc.Interval(id="compute-poll-interval", interval=100, disabled=True), - # RF/MDS job identities and two-phase terminal delivery. - # The poll callback writes the applied marker atomically - # with the visible result; only then does the ack callback - # release the server-side sticky terminal event. + # Per-workflow job identities and shared two-phase terminal + # delivery. The poll that renders a terminal result writes + # an applied marker; only then does the acknowledgement + # callback release the server-side sticky event. dcc.Store(id="rf-job-store", storage_type="memory"), dcc.Store(id="mds-job-store", storage_type="memory"), - dcc.Store(id="rf-mds-applied-job-store", storage_type="memory"), - dcc.Store(id="rf-mds-job-ack-store", storage_type="memory"), - # consensus tree computation polls on its OWN interval. Dash derives an - # allow_duplicate output's disambiguation hash from the - # callback's Input signature (dash/_utils.py), so sharing - # ``compute-poll-interval`` between the RF/MDS poll and the - # consensus tree poll makes them collide on the shared - # ``compute-poll-interval.disabled`` / - # ``notifications-container.children`` outputs. A dedicated - # interval gives ``poll_consensus_tree_completion`` a distinct Input. + dcc.Store(id="pseudo-ess-job-store", storage_type="memory"), + dcc.Store(id="consensus-job-store", storage_type="memory"), + dcc.Store(id="compute-applied-job-store", storage_type="memory"), + dcc.Store(id="compute-job-ack-store", storage_type="memory"), + # Consensus trees keep a dedicated cadence because their + # overlay and button lifecycle can stop independently of + # the shared RF/MDS/Pseudo-ESS interval. dcc.Interval(id="consensus-tree-poll-interval", interval=100, disabled=True), # Path to the RF worker's sidecar progress file # (``.progress``). Set by From 9673878a8a52d9826c80b444f68833c105fef521 Mon Sep 17 00:00:00 2001 From: Sam Hong <20244918+hongsamL@users.noreply.github.com> Date: Mon, 3 Aug 2026 15:57:15 +0200 Subject: [PATCH 5/9] rf-trace and clade comparison using job manager --- src/test/test_managed_analysis_jobs.py | 301 +++++++++ src/test/test_managed_compute_jobs.py | 2 +- src/treetracer/__init__.py | 14 + src/treetracer/callbacks/clade_explore.py | 583 ++++++++++++------ src/treetracer/callbacks/compute.py | 8 +- src/treetracer/callbacks/diagnostics.py | 338 +++++++++- src/treetracer/callbacks/persistent_worker.py | 6 +- src/treetracer/callbacks/sidebar.py | 24 +- .../clade_freq/_subprocess_worker.py | 138 +++++ src/treetracer/ess/_rf_trace_worker.py | 45 ++ src/treetracer/ess/rf_trace.py | 52 +- src/treetracer/state.py | 46 ++ src/treetracer/ui/navbar.py | 6 + 13 files changed, 1298 insertions(+), 265 deletions(-) create mode 100644 src/test/test_managed_analysis_jobs.py create mode 100644 src/treetracer/clade_freq/_subprocess_worker.py create mode 100644 src/treetracer/ess/_rf_trace_worker.py diff --git a/src/test/test_managed_analysis_jobs.py b/src/test/test_managed_analysis_jobs.py new file mode 100644 index 0000000..2bc8f6d --- /dev/null +++ b/src/test/test_managed_analysis_jobs.py @@ -0,0 +1,301 @@ +"""Managed lifecycle and worker tests for Stage 4 analysis jobs.""" + +from __future__ import annotations + +import time +from concurrent.futures import ThreadPoolExecutor +from functools import partial +from types import SimpleNamespace + +import numpy as np +import pandas as pd +import pytest + +from treetracer import state +from treetracer.background_jobs import JobManager, JobState +from treetracer.callbacks import clade_explore, compute, diagnostics +from treetracer.clade_freq._subprocess_worker import ( + compute_clade_frequencies_worker_entry, +) +from treetracer.ess._rf_trace_worker import compute_rf_trace_worker_entry + + +def _wait_for_terminal(manager, ref, timeout=2.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + snapshot = manager.snapshot(ref) + if snapshot is not None and snapshot.terminal is not None: + return snapshot + time.sleep(0.005) + pytest.fail("background job did not become terminal") + + +def _registered_callback(name, register): + from dash import _callback + + def matches(): + found = [] + for callback_data in _callback.GLOBAL_CALLBACK_MAP.values(): + callback_fn = callback_data.get("callback") + if callback_fn is None: + continue + original = getattr(callback_fn, "__wrapped__", callback_fn) + if original.__name__ == name: + found.append(original) + return found + + found = matches() + if not found: + register() + found = matches() + assert found + return found[-1] + + +@pytest.fixture(autouse=True) +def _clean_analysis_results(): + state.clear_all_analysis_results() + yield + state.clear_all_analysis_results() + + +def test_rf_trace_worker_memory_maps_and_returns_only_requested_row(tmp_path): + matrix = np.arange(25, dtype=np.uint16).reshape(5, 5) + path = tmp_path / "matrix.npy" + np.save(path, matrix) + + result = compute_rf_trace_worker_entry( + matrix_path=str(path), + reference_index=3, + ) + + assert result["matrix_size"] == 5 + assert result["distances"] == matrix[3].tolist() + assert "matrix" not in result + + +def test_clade_worker_decodes_only_requested_columns(tmp_path): + path = tmp_path / "snapshot.npz" + presence = np.array( + [ + [1, 1, 0, 0], + [1, 0, 1, 1], + [0, 1, 1, 0], + [1, 1, 0, 1], + ], + dtype=np.uint8, + ) + bits = np.array( + [ + [1, 0, 0], + [0, 1, 1], + [1, 1, 0], + [0, 0, 1], + ], + dtype=np.uint8, + ) + np.savez( + path, + presence=presence, + bipartition_bits=bits, + leaf_names=np.array(["'A'", "B", "C"]), + ) + + result = compute_clade_frequencies_worker_entry( + snapshots_path=str(path), + columns=[1, 3], + counts_1=[1, 1], + counts_2=[2, 1], + n_trees_1=2, + n_trees_2=2, + ) + + assert result["leaf_names"] == ["A", "B", "C"] + assert {row["column_j"] for row in result["rows"]} == {1, 3} + by_column = {row["column_j"]: row for row in result["rows"]} + assert by_column[1]["split_key"] == (1, 2) + assert by_column[1]["freq_1"] == pytest.approx(0.5) + assert by_column[1]["freq_2"] == pytest.approx(1.0) + assert by_column[3]["split_key"] == (2,) + + +def test_rf_trace_terminal_replays_cached_render_until_ack(monkeypatch): + manager = JobManager(id_factory=lambda: "rf-trace-test-job") + db = SimpleNamespace( + _trees=pd.DataFrame( + { + "name": ["run-a/tree-1", "run-a/tree-2", "run-b/tree-1"], + "file_source": ["a.trees", "a.trees", "b.trees"], + } + ), + flush=lambda: None, + ) + fake_figure = SimpleNamespace(to_dict=lambda: {"data": [], "layout": {}}) + monkeypatch.setattr(diagnostics, "job_manager", manager) + monkeypatch.setattr(compute, "job_manager", manager) + monkeypatch.setattr(diagnostics, "add_log", lambda *_a, **_k: None) + monkeypatch.setattr( + diagnostics, + "get_tree_service", + lambda: SimpleNamespace(db_manager=db), + ) + monkeypatch.setattr( + diagnostics, + "_build_rf_trace_fig", + lambda *_a, **_k: fake_figure, + ) + + context = diagnostics._RfTraceFinalizationContext( + selected_matrix="RF_001", + names=("run-a/tree-1", "run-a/tree-2", "run-b/tree-1"), + reference_index=0, + reference_name="run-a/tree-1", + reference_group="run-a", + reference_position="first", + burnin=0, + ) + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "rf_trace", + lambda: { + "distances": [0, 4, 8], + "matrix_size": 3, + "elapsed": 0.01, + }, + finalizer=partial( + diagnostics._finalize_rf_trace_job, + context=context, + ), + ) + terminal = _wait_for_terminal(manager, ref) + + assert terminal.state is JobState.SUCCEEDED + assert terminal.terminal.payload["result_key"] == ref.job_id + assert "distances" not in terminal.terminal.payload + cached = state.get_rf_trace_result(ref.job_id) + assert [row["rf_distance"] for row in cached["records"]] == [4, 8] + + poll = _registered_callback( + "poll_rf_trace_completion", + diagnostics.register_diagnostics_callbacks, + ) + first = poll(1, ref.as_dict()) + second = poll(2, ref.as_dict()) + assert len(first) == 7 + assert len(first[1]) == 2 + assert first[3] is False + assert first[4] is False + assert first[5] is True + assert first[6]["delivery_attempt"] == 1 + assert second[6]["delivery_attempt"] == 2 + + assert compute._ack_applied_job(second[6]) is True + assert manager.snapshot(ref).acknowledged is True + + +def test_stage_four_submit_callbacks_have_matching_idle_output_shapes(): + rf_submit = _registered_callback( + "compute_rf_trace", + diagnostics.register_diagnostics_callbacks, + ) + clade_submit = _registered_callback( + "compute_and_plot_clade_frequencies", + clade_explore.register_clade_explore_callbacks, + ) + + assert len(rf_submit(None, None, None, None, None, None, None)) == 6 + assert len(clade_submit(None, None, None, None, None)) == 8 + + +def test_clade_terminal_keeps_rows_server_side_and_keys_click_resolution( + monkeypatch, +): + manager = JobManager(id_factory=lambda: "clade-test-job") + fake_figure = SimpleNamespace(to_dict=lambda: {"data": [], "layout": {}}) + monkeypatch.setattr(clade_explore, "job_manager", manager) + monkeypatch.setattr(compute, "job_manager", manager) + monkeypatch.setattr(clade_explore, "add_log", lambda *_a, **_k: None) + monkeypatch.setattr( + clade_explore, + "_build_scatter_fig", + lambda *_a, **_k: fake_figure, + ) + + context = clade_explore._CladeComparisonFinalizationContext( + source_distmat="RF_001", + uid_1="uid-a", + uid_2="uid-b", + label_1="Consensus A", + label_2="Consensus B", + consensus_columns_1=frozenset({2}), + consensus_columns_2=frozenset({2, 5}), + min_clade_size=2, + ) + worker_result = { + "rows": [ + { + "split_key": (0, 2), + "column_j": 2, + "freq_1": 0.8, + "freq_2": 0.6, + "clade_size": 2, + }, + { + "split_key": (1,), + "column_j": 5, + "freq_1": 0.0, + "freq_2": 0.4, + "clade_size": 1, + }, + ], + "leaf_names": ["A", "B", "C"], + "elapsed": 0.02, + } + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "clade_compare", + lambda: worker_result, + finalizer=partial( + clade_explore._finalize_clade_comparison_job, + context=context, + ), + ) + terminal = _wait_for_terminal(manager, ref) + + assert terminal.state is JobState.SUCCEEDED + assert "rows" not in terminal.terminal.payload + assert "leaf_names" not in terminal.terminal.payload + resolved = clade_explore._resolved_clade( + {"result_key": ref.job_id, "split_id": 0}, + expected_pair=("uid-a", "uid-b"), + ) + assert resolved == { + "source_distmat": "RF_001", + "column_j": 2, + "tip_names": ("A", "C"), + } + assert clade_explore._resolved_clade( + {"result_key": ref.job_id, "split_id": 0}, + expected_pair=("uid-b", "uid-a"), + ) is None + + poll = _registered_callback( + "poll_clade_frequency_completion", + clade_explore.register_clade_explore_callbacks, + ) + first = poll(1, ref.as_dict()) + second = poll(2, ref.as_dict()) + assert len(first) == 8 + assert len(first[1]) == 2 + assert first[2] == {} + assert first[3] is False + assert first[4] is True + assert first[5] == ref.job_id + assert first[6] is None + assert first[7]["delivery_attempt"] == 1 + assert second[7]["delivery_attempt"] == 2 + + assert compute._ack_applied_job(second[7]) is True + assert manager.snapshot(ref).acknowledged is True diff --git a/src/test/test_managed_compute_jobs.py b/src/test/test_managed_compute_jobs.py index f6e5d47..31c89c0 100644 --- a/src/test/test_managed_compute_jobs.py +++ b/src/test/test_managed_compute_jobs.py @@ -272,4 +272,4 @@ def test_clear_data_idle_branch_matches_managed_lifecycle_outputs(): "clear_uploads", sidebar.register_sidebar_callbacks, ) - assert len(clear_uploads(None)) == 45 + assert len(clear_uploads(None)) == 48 diff --git a/src/treetracer/__init__.py b/src/treetracer/__init__.py index 04fac11..2da1582 100644 --- a/src/treetracer/__init__.py +++ b/src/treetracer/__init__.py @@ -135,6 +135,20 @@ def _recv_exactly(n: int) -> bytes: wlog("calling compute_pseudo_ess_worker_entry") result = compute_pseudo_ess_worker_entry(**kwargs) wlog("compute_pseudo_ess_worker_entry returned") + elif job == "compute_rf_trace": + wlog("importing compute_rf_trace_worker_entry") + from .ess._rf_trace_worker import compute_rf_trace_worker_entry + wlog("calling compute_rf_trace_worker_entry") + result = compute_rf_trace_worker_entry(**kwargs) + wlog("compute_rf_trace_worker_entry returned") + elif job == "compute_clade_frequencies": + wlog("importing compute_clade_frequencies_worker_entry") + from .clade_freq._subprocess_worker import ( + compute_clade_frequencies_worker_entry, + ) + wlog("calling compute_clade_frequencies_worker_entry") + result = compute_clade_frequencies_worker_entry(**kwargs) + wlog("compute_clade_frequencies_worker_entry returned") elif job == "compute_mds": # MDS doesn't need a dedicated worker wrapper — the # ``rf._worker.compute_mds_worker`` function is already diff --git a/src/treetracer/callbacks/clade_explore.py b/src/treetracer/callbacks/clade_explore.py index d3495ac..0df710a 100644 --- a/src/treetracer/callbacks/clade_explore.py +++ b/src/treetracer/callbacks/clade_explore.py @@ -8,34 +8,25 @@ """ import functools +from dataclasses import dataclass +from functools import partial +from typing import Any -from dash import (dcc, html, callback, Input, Output, State, no_update, Patch, - clientside_callback) +from dash import dcc, html, callback, Input, Output, State, no_update, Patch import dash_mantine_components as dmc import plotly.graph_objects as go import numpy as np import pandas as pd from .. import state -from ..clade_freq import compute_clade_frequencies +from ..background_jobs import JobBusyError, JobRef, JobState, job_manager from ..clade_freq.layout import parse_nexus, build_tree_traces, _collect_nodes +from ..logger import add_log from ..theme import get_template, DARK_TEMPLATE from ..plot_utils import retheme_figure - - -# Server-side resolution table for the Clade Frequency scatter → -# tanglegram round-trip. The scatter's ``customdata`` carries an -# integer ``split_id`` (the row index in the DataFrame produced by -# ``compute_clade_frequencies``); this dict maps that id to a -# ``(source_distmat, column_j, split_key)`` triple so the click -# handler can both resolve the clade's tip names AND check consensus tree -# membership via ``column_j in cols_in_consensus_tree`` (the rapidtrees-encoded -# rooted-clade column index in that distmat's snapshot). Rebuilt on -# every Compare click. -# -# Keeping the keys server-side avoids serialising tens of thousands -# of tuples through the browser store on every click. -_split_resolution: dict[int, tuple[str, int, tuple]] = {} +from ..ui.widgets import stop_button +from . import persistent_worker +from .compute import _get_executor, _job_ref_from_store # --------------------------------------------------------------------------- @@ -515,13 +506,12 @@ def clear_clade_freq_caches(): Comparison feature. Called from ``sidebar.clear_uploads`` so a Clear-data click - actually wipes the bipartition→tip-set decode, the parsed-NEXUS - LRU, and the click→split lookup table — they're keyed on consensus tree - uuids that are about to disappear from ``state._consensus_tree_cache``. + actually wipes the parsed-NEXUS and layout LRUs. Managed comparison + payloads and their click-resolution maps live in ``state`` and are cleared + by ``state.clear_all_analysis_results`` in the same reset transaction. """ _get_parsed_consensus_tree.cache_clear() _get_tanglegram_layout.cache_clear() - _split_resolution.clear() def _build_scatter_fig(df_plot, label1, label2): @@ -612,6 +602,138 @@ def _build_scatter_fig(df_plot, label1, label2): return fig +@dataclass(frozen=True, slots=True) +class _CladeComparisonFinalizationContext: + source_distmat: str + uid_1: str + uid_2: str + label_1: str + label_2: str + consensus_columns_1: frozenset[int] + consensus_columns_2: frozenset[int] + min_clade_size: int + + +def _cached_counts_for_columns(entry, columns): + """Return a compact count slice or ``None`` for the worker fallback.""" + counts = entry.get("counts") + n_trees = int(entry.get("n_trees") or 0) + if counts is None or n_trees <= 0: + return None, 0 + values = np.asarray(counts) + if values.ndim != 1: + raise ValueError("cached clade counts must be one-dimensional") + if columns and columns[-1] >= len(values): + raise ValueError("consensus-tree clade column is outside cached counts") + return values[columns].astype(np.int64, copy=False).tolist(), n_trees + + +def _finalize_clade_comparison_job( + ref: JobRef, + result: Any, + *, + context: _CladeComparisonFinalizationContext, +) -> dict[str, Any]: + """Publish compact scatter data and its server-side click resolution.""" + if not isinstance(result, dict): + raise TypeError("clade-frequency worker returned a non-mapping result") + rows = result.get("rows") + leaf_names = result.get("leaf_names") + if not isinstance(rows, list) or not isinstance(leaf_names, list): + raise TypeError("clade-frequency worker returned an invalid payload") + + records = [] + resolution = {} + for row in rows: + if not isinstance(row, dict): + raise TypeError("clade-frequency row must be a mapping") + column_j = int(row["column_j"]) + in_1 = column_j in context.consensus_columns_1 + in_2 = column_j in context.consensus_columns_2 + if not (in_1 or in_2): + continue + if in_1 and in_2: + membership = "both consensus trees" + elif in_1: + membership = f"{context.label_1} only" + else: + membership = f"{context.label_2} only" + + split_key = tuple(int(index) for index in row["split_key"]) + try: + tip_names = tuple(str(leaf_names[index]) for index in split_key) + except IndexError as exc: + raise ValueError( + "clade-frequency split references an unknown leaf index" + ) from exc + + split_id = len(records) + records.append( + { + "split_id": split_id, + "freq_1": float(row["freq_1"]), + "freq_2": float(row["freq_2"]), + "clade_size": int(row["clade_size"]), + "in_consensus_tree_1": in_1, + "in_consensus_tree_2": in_2, + "consensus_tree_membership": membership, + } + ) + resolution[split_id] = { + "source_distmat": context.source_distmat, + "column_j": column_j, + "tip_names": tip_names, + } + + if not records: + raise ValueError( + "None of the selected consensus-tree clades are available to plot" + ) + frame = pd.DataFrame(records) + plotted = frame[frame["clade_size"] >= context.min_clade_size] + figure = _build_scatter_fig(plotted, context.label_1, context.label_2) + state.store_clade_frequency_result( + ref.job_id, + { + "records": records, + "resolution": resolution, + "figure": figure.to_dict(), + "pair": (context.uid_1, context.uid_2), + }, + ) + elapsed = float(result.get("elapsed", 0.0)) + add_log( + f"Compared {len(records)} consensus-tree clades for " + f"{context.source_distmat} in {elapsed:.3f}s." + ) + return { + "result_key": ref.job_id, + "source_distmat": context.source_distmat, + "label_1": context.label_1, + "label_2": context.label_2, + "n_clades": len(records), + "elapsed": elapsed, + } + + +def _resolved_clade(click_data, expected_pair=None): + """Resolve a browser split ID through its immutable managed result key.""" + if not isinstance(click_data, dict): + return None + result_key = click_data.get("result_key") + split_id = click_data.get("split_id") + if result_key is None or split_id is None: + return None + cached = state.get_clade_frequency_result(result_key) + if cached is None: + return None + if expected_pair is not None and tuple(expected_pair) != tuple( + cached.get("pair", ()) + ): + return None + return cached.get("resolution", {}).get(int(split_id)) + + def register_clade_explore_callbacks(): # Monotonic counter that gets stamped on every scatter-plot click # payload (see ``store_scatter_click`` below). Without a unique @@ -816,177 +938,260 @@ def toggle_compare_button(uid1, uid2): # ------ Clade Frequency Comparison: compute and plot ------ - # Instant feedback on click: flip the output paper visible so the - # ``dcc.Loading`` wrapper around ``clade-freq-plot`` (defined in - # ui/panels/clade_freq.py) can render its spinner while the slow - # server compute below runs. Without this the paper stays hidden - # until the compute returns, so the user sees no loader at all. - # Runs clientside (no Python roundtrip) so the spinner appears - # within a frame of the click. - clientside_callback( - """ - function(n_clicks) { - if (!n_clicks) return window.dash_clientside.no_update; - return {}; - } - """, - Output("clade-freq-output-paper", "style", allow_duplicate=True), - Input("clade-freq-compare-button", "n_clicks"), - prevent_initial_call=True, - ) - @callback( - Output("clade-freq-plot", "children"), - Output("clade-freq-data-store", "data"), - # Toggle the output Paper visible only on success; stays - # hidden on any error path or before the first successful - # Compare. Cleared by the sidebar's Clear-data flow. - Output("clade-freq-output-paper", "style"), + Output("clade-freq-plot", "children", allow_duplicate=True), + Output("clade-freq-data-store", "data", allow_duplicate=True), + Output("clade-freq-output-paper", "style", allow_duplicate=True), + Output("clade-freq-compare-button", "disabled", allow_duplicate=True), + Output("compute-poll-interval", "disabled", allow_duplicate=True), + Output("clade-freq-job-store", "data"), + Output("clade-freq-result-key-store", "data", allow_duplicate=True), + Output("clade-freq-click-store", "data", allow_duplicate=True), Input("clade-freq-compare-button", "n_clicks"), State("clade-freq-consensus-tree-select-1", "value"), State("clade-freq-consensus-tree-select-2", "value"), State("clade-freq-min-clade-size", "value"), + State("compute-applied-job-store", "data"), prevent_initial_call=True, ) - def compute_and_plot_clade_frequencies(n_clicks, uid1, uid2, min_clade_size): - """Compute clade frequencies for the two selected consensus tree groups and - render a scatter plot (freq group 1 vs freq group 2). - - Each dot is one bipartition observed in either group. Dot colour - encodes clade_size — the number of tips in the monophyletic - descendant side of the bipartition in the consensus tree(s) that - contain it (so it matches what the tanglegram highlights when - the dot is clicked). Clicking a dot triggers the tanglegram - callback. - """ - if not uid1 or not uid2: - return no_update, no_update, no_update + def compute_and_plot_clade_frequencies( + n_clicks, + uid1, + uid2, + min_clade_size, + applied_job, + ): + """Prepare and submit a selective, process-isolated comparison.""" + from .compute import _ack_applied_job + + if not n_clicks or not uid1 or not uid2: + return (no_update,) * 8 + + _ack_applied_job(applied_job) + + def error(message): + return ( + dmc.Text(message, c="red", size="sm"), + no_update, + {}, + False, + no_update, + no_update, + no_update, + no_update, + ) entry1 = state.get_consensus_tree_registry_entry(uid1) entry2 = state.get_consensus_tree_registry_entry(uid2) if entry1 is None or entry2 is None: - return dmc.Text( + return error( "One or both selected consensus trees are no longer available. " - "Please recompute them.", - c="red", size="sm", - ), no_update, no_update + "Please recompute them." + ) try: - df = compute_clade_frequencies(entry1, entry2) - except (KeyError, FileNotFoundError) as e: - return dmc.Text( - f"Error computing clade frequencies: {e}", - c="red", size="sm", - ), no_update, no_update - - # Annotate every rooted clade with whether it's present in - # consensus tree 1 / consensus tree 2. The dropdowns in ``populate_consensus_tree_selects`` - # restrict choices to the active distmat, so by construction - # ``entry1.source_distmat == entry2.source_distmat`` and both - # ``cols_in_consensus_tree`` lists are in the same column basis as the - # DataFrame's ``column_j`` — a pure int-in-set check. - src1 = entry1["source_distmat"] - cols_in_consensus_tree_1 = entry1.get("cols_in_consensus_tree") or [] - cols_in_consensus_tree_2 = entry2.get("cols_in_consensus_tree") or [] - df["in_consensus_tree_1"] = df["column_j"].isin(set(cols_in_consensus_tree_1)) - df["in_consensus_tree_2"] = df["column_j"].isin(set(cols_in_consensus_tree_2)) - - # Show only clades that are present in at least one of the two - # consensus trees — keeps the tanglegram meaningful when the user clicks. - df = df[df["in_consensus_tree_1"] | df["in_consensus_tree_2"]].reset_index(drop=True) - - # No clade-size re-stamping in rooted mode: each rapidtrees - # column is already a rooted clade with one specific descendant - # set, so the size computed in ``compute_clade_frequencies`` - # (``len(split_key)``) is already the size we want to display. - - if df.empty: - return dmc.Text( - "None of the bipartitions observed in the two groups " - "is a clade of either consensus tree — nothing to plot.", - c="dimmed", size="sm", - ), no_update, no_update - - # Pre-render a human-readable membership label per row for - # the scatter hover. Stored in the DataFrame so the slider - # callback can patch ``customdata`` without rebuilding the - # mapping. - label1_h = entry1["name"] - label2_h = entry2["name"] - membership_labels = [] - for in1, in2 in zip(df["in_consensus_tree_1"], df["in_consensus_tree_2"]): - if in1 and in2: - membership_labels.append("both consensus trees") - elif in1: - membership_labels.append(f"{label1_h} only") - else: - membership_labels.append(f"{label2_h} only") - df["consensus_tree_membership"] = membership_labels - - # Integer row id replaces the old fragile comma-joined string. - # The click-handler + tanglegram callbacks resolve split_id to - # tip names via state.get_canonical_keys at render time. - df["split_id"] = np.arange(len(df), dtype=np.int32) - - # Refresh the click-resolution table: split_id → (distmat, - # column_j, split_key). ``column_j`` is the rapidtrees-encoded - # rooted-clade column index in the snapshot; the click handler - # uses it for the in_1/in_2 check against each consensus tree's - # ``cols_in_consensus_tree``. ``split_key`` is the descendant-set tuple - # of leaf indices, used only to resolve tip names for the - # highlight overlay. - _split_resolution.clear() - for split_id, col_j, key in zip( - df["split_id"].tolist(), - df["column_j"].tolist(), - df["split_key"].tolist(), - ): - _split_resolution[int(split_id)] = (src1, int(col_j), key) - - label1 = entry1["name"] - label2 = entry2["name"] - - # Serialise for the slider callback. split_key is a tuple[int] - # — not JSON-serialisable, so it stays server-side and we only - # ship the integer id through the browser. The in_consensus_tree_1/ - # in_consensus_tree_2 booleans ride along so any future filter or - # colour-coding callback can consume them without re-running - # the clade-membership check. - store_data = df[[ - "split_id", "freq_1", "freq_2", "clade_size", - "in_consensus_tree_1", "in_consensus_tree_2", "consensus_tree_membership", - ]].to_dict("records") + source_1 = str(entry1["source_distmat"]) + source_2 = str(entry2["source_distmat"]) + if source_1 != source_2: + raise ValueError( + "Selected consensus trees belong to different RF matrices." + ) + columns_1 = frozenset( + int(column) + for column in (entry1.get("cols_in_consensus_tree") or []) + ) + columns_2 = frozenset( + int(column) + for column in (entry2.get("cols_in_consensus_tree") or []) + ) + columns = sorted(columns_1 | columns_2) + if not columns: + raise ValueError( + "Selected consensus trees have no cached clade columns." + ) + counts_1, n_trees_1 = _cached_counts_for_columns(entry1, columns) + counts_2, n_trees_2 = _cached_counts_for_columns(entry2, columns) + needs_fallback = counts_1 is None or counts_2 is None + full_names = ( + list(state.get_distmat_names(source_1)) + if needs_fallback + else None + ) + snapshots_path = str(state.get_snapshots_path(source_1)) + except (KeyError, ValueError, FileNotFoundError) as exc: + return error(f"Error preparing clade comparison: {exc}") - min_size = int(min_clade_size or 2) - df_plot = df[df["clade_size"] >= min_size] + try: + min_size = max(1, int(min_clade_size or 2)) + except (TypeError, ValueError): + min_size = 2 + context = _CladeComparisonFinalizationContext( + source_distmat=source_1, + uid_1=str(uid1), + uid_2=str(uid2), + label_1=str(entry1["name"]), + label_2=str(entry2["name"]), + consensus_columns_1=columns_1, + consensus_columns_2=columns_2, + min_clade_size=min_size, + ) + + try: + job_ref = job_manager.submit( + _get_executor(), + "clade_compare", + persistent_worker.submit_job, + "compute_clade_frequencies", + snapshots_path=snapshots_path, + columns=columns, + counts_1=counts_1, + counts_2=counts_2, + n_trees_1=n_trees_1, + n_trees_2=n_trees_2, + tree_names_1=( + list(entry1.get("tree_names") or []) + if counts_1 is None + else None + ), + tree_names_2=( + list(entry2.get("tree_names") or []) + if counts_2 is None + else None + ), + full_distmat_names=full_names, + metadata={ + "display_name": "Clade Frequency Comparison", + "source_distmat": source_1, + }, + finalizer=partial( + _finalize_clade_comparison_job, + context=context, + ), + cancel_exceptions=(persistent_worker.JobCancelled,), + ) + except JobBusyError as exc: + message = ( + f"Another computation ({exc.active.kind.replace('_', ' ').upper()}) " + "is still finishing. Please wait for it to complete." + ) + add_log(message, "WARNING") + return error(message) + + spinner = dmc.Group( + [ + dmc.Loader(size="sm", type="dots"), + dmc.Text( + f"Comparing {len(columns)} consensus-tree clades…", + size="sm", + c="dimmed", + ), + stop_button("clade-compare"), + ], + gap="sm", + ) + return ( + spinner, + None, + {}, + True, + False, + job_ref.as_dict(), + None, + None, + ) + + @callback( + Output("clade-freq-plot", "children", allow_duplicate=True), + Output("clade-freq-data-store", "data", allow_duplicate=True), + Output("clade-freq-output-paper", "style", allow_duplicate=True), + Output("clade-freq-compare-button", "disabled", allow_duplicate=True), + Output("compute-poll-interval", "disabled", allow_duplicate=True), + Output("clade-freq-result-key-store", "data", allow_duplicate=True), + Output("clade-freq-click-store", "data", allow_duplicate=True), + Output("compute-applied-job-store", "data", allow_duplicate=True), + Input("compute-poll-interval", "n_intervals"), + Input("clade-freq-job-store", "data"), + prevent_initial_call=True, + ) + def poll_clade_frequency_completion(_n_intervals, job_data): + ref = _job_ref_from_store(job_data, expected_kind="clade_compare") + if ref is None: + return (no_update,) * 8 + snapshot = job_manager.snapshot_for_delivery(ref) + if ( + snapshot is None + or snapshot.acknowledged + or snapshot.terminal is None + ): + return (no_update,) * 8 + + terminal = snapshot.terminal + payload = terminal.payload + store_data = no_update + result_key = no_update + if terminal.state is JobState.CANCELLED: + if snapshot.delivery_attempt == 1: + add_log("Clade comparison cancelled by user.", "WARNING") + output = dmc.Alert( + title="Clade comparison cancelled", + children=dmc.Text("Stopped before completion.", size="sm"), + color="gray", + variant="light", + ) + elif terminal.state is JobState.FAILED: + message = str(payload.get("message", "Unknown error")) + if snapshot.delivery_attempt == 1: + add_log(f"Clade comparison failed: {message}", "ERROR") + output = dmc.Text( + f"Error computing clade frequencies: {message}", + c="red", + size="sm", + ) + else: + cached = state.get_clade_frequency_result( + payload.get("result_key") + ) + if cached is None: + output = dmc.Text( + "Clade comparison result is no longer available. " + "Please recompute it.", + c="red", + size="sm", + ) + else: + store_data = cached["records"] + result_key = payload["result_key"] + output = dcc.Graph( + id="clade-freq-scatter", + figure=cached["figure"], + config={"displayModeBar": False}, + style={"width": "100%"}, + ) - fig = _build_scatter_fig(df_plot, label1, label2) - # Success path: reveal the output paper. return ( - dcc.Graph( - id="clade-freq-scatter", - figure=fig, - config={"displayModeBar": False}, - style={"width": "100%"}, - ), + output, store_data, {}, + False, + True, + result_key, + None, + snapshot.terminal_delivery_marker(), ) @callback( Output("clade-freq-click-store", "data"), Input("clade-freq-scatter", "clickData"), + State("clade-freq-result-key-store", "data"), prevent_initial_call=True, ) - def store_scatter_click(click_data): + def store_scatter_click(click_data, result_key): """Forward a scatter plot click to the click store. - ``customdata`` is ``[split_id, clade_size]``. The split_id is an - integer row index in the DataFrame produced by the compute - callback; ``draw_tanglegram`` uses it together with the - per-distmat canonical-keys cache to resolve the actual tip - names to highlight. + ``customdata`` is ``[split_id, clade_size, membership]``. The + split ID and immutable result key let ``draw_tanglegram`` resolve the + actual tip names without decoding the full RF snapshot in the UI. The ``_t`` nonce is set to a unique counter each time so ``dcc.Store`` does not deduplicate identical click payloads @@ -995,7 +1200,7 @@ def store_scatter_click(click_data): because the store value matches the previous click. """ nonlocal _click_counter - if not click_data or not click_data.get("points"): + if not result_key or not click_data or not click_data.get("points"): return no_update point = click_data["points"][0] custom = point.get("customdata") @@ -1008,6 +1213,7 @@ def store_scatter_click(click_data): "clade_size": int(custom[1]), "x": float(point["x"]), "y": float(point["y"]), + "result_key": str(result_key), "_t": _click_counter, } except (TypeError, ValueError, IndexError, KeyError): @@ -1068,38 +1274,23 @@ def draw_tanglegram(click_data, uid1, uid2, px_per_tip, current_pair, complement segments + 280 tip markers per tree) is never re-sent. The clicked split is identified by an integer ``split_id``; - the actual tip names are resolved server-side via - ``_split_resolution`` (rebuilt by the Compare callback) and - the per-distmat canonical-keys cache. + the actual tip names are resolved server-side through the immutable + managed comparison result identified in the click payload. """ if not click_data or not uid1 or not uid2: return no_update, no_update, no_update - # ── Resolve the click via _split_resolution ─────────────────────── - # Yields ``(src, column_j, split_key)`` where ``split_key`` is the - # tuple of leaf indices that make up the rooted clade. - split_id = click_data.get("split_id") - if split_id is None: - return no_update, no_update, no_update - resolved = _split_resolution.get(int(split_id)) + resolved = _resolved_clade(click_data, expected_pair=(uid1, uid2)) if resolved is None: - # Click store survived a Compare-button reset and we no - # longer know which split this is. Drop the request - # quietly; the next Compare repopulates _split_resolution. - return no_update, no_update, no_update - src, column_j, split_key = resolved - - try: - canonical = state.get_canonical_keys(src) - except (KeyError, FileNotFoundError): + # The result was reset or evicted; a fresh Compare repopulates it. return no_update, no_update, no_update - leaf_names = canonical["leaf_names"] + column_j = int(resolved["column_j"]) # Single highlight (same on both trees): consensus tree 1 and consensus tree 2 are - # both anchored to ``src`` (the Compare-clade dropdowns + # both anchored to one source matrix (the Compare-clade dropdowns # filter to the active distmat), so a column in the rooted # presence table represents the same descendant set in both. - highlight = {leaf_names[i] for i in split_key} + highlight = set(resolved["tip_names"]) # Containment is an O(1) ``column_j ∈ cols_in_consensus_tree`` check # straight off the registry entries. @@ -1298,8 +1489,8 @@ def toggle_complement_highlights(checked, click_data, uid1, uid2): Patches only the two green overlay traces (9/10); the grey tip base (3/4) is left untouched because the green markers simply draw on top of it. Hiding empties the green x/y/text; showing - resolves the last clicked split from ``_split_resolution`` and - rebuilds them. Their styling was baked in at first render, so + resolves the last clicked split from its managed result and rebuilds + them. Their styling was baked in at first render, so the patch never re-sends marker/line/mode. """ # Need a rendered tanglegram — a prior click plus a live layout @@ -1318,20 +1509,10 @@ def toggle_complement_highlights(checked, click_data, uid1, uid2): patch["data"][idx]["text"] = [] return patch - split_id = click_data.get("split_id") - if split_id is None: - return no_update - resolved = _split_resolution.get(int(split_id)) + resolved = _resolved_clade(click_data, expected_pair=(uid1, uid2)) if resolved is None: return no_update - - src, _column_j, split_key = resolved - try: - canonical = state.get_canonical_keys(src) - except (KeyError, FileNotFoundError): - return no_update - leaf_names = canonical["leaf_names"] - highlight = {leaf_names[i] for i in split_key} + highlight = set(resolved["tip_names"]) _, _, complement_per_tree = _build_mrca_traces(highlight, layout) complement1 = complement_per_tree["tips1"] diff --git a/src/treetracer/callbacks/compute.py b/src/treetracer/callbacks/compute.py index 93049dd..486c4a3 100644 --- a/src/treetracer/callbacks/compute.py +++ b/src/treetracer/callbacks/compute.py @@ -323,14 +323,14 @@ def _shutdown_executor(): def reset(): - """Invalidate any RF/MDS job before Clear Data wipes its inputs. + """Invalidate any managed job before Clear Data wipes its inputs. Removing the manager record prevents a late worker result from - re-registering a matrix or MDS result after the application state has been - cleared. The persistent-worker kill interrupts native work when possible. + publishing into freshly cleared application state. The persistent-worker + kill interrupts native work when possible. """ active = job_manager.active_ref() - if active is None or active.kind not in {"rf", "mds"}: + if active is None: return persistent_worker.cancel_current_job() job_manager.invalidate(active) diff --git a/src/treetracer/callbacks/diagnostics.py b/src/treetracer/callbacks/diagnostics.py index ade1b5b..71adb49 100644 --- a/src/treetracer/callbacks/diagnostics.py +++ b/src/treetracer/callbacks/diagnostics.py @@ -1,3 +1,7 @@ +from dataclasses import dataclass +from functools import partial +from typing import Any + from dash import dcc, html, callback, clientside_callback, Input, Output, State, no_update, ctx, ALL import dash_mantine_components as dmc import plotly.express as px @@ -7,13 +11,15 @@ import numpy as np import pandas as pd -from ..background_jobs import JobBusyError -from ..logger import add_log, notif_id +from ..background_jobs import JobBusyError, JobRef, JobState, job_manager +from ..logger import add_log from ..db.tree_service import get_tree_service -from ..ess.rf_trace import compute_rf_trace_data +from ..ess.rf_trace import find_reference_index from .. import state from ..theme import get_template from ..ui.widgets import stop_button +from . import persistent_worker +from .compute import _get_executor, _job_ref_from_store def _build_rf_trace_fig(trace_df, ref_group, ref_position, burnin=0): @@ -84,6 +90,105 @@ def _build_rf_trace_fig(trace_df, ref_group, ref_position, burnin=0): ) return fig + +@dataclass(frozen=True, slots=True) +class _RfTraceFinalizationContext: + selected_matrix: str + names: tuple[str, ...] + reference_index: int + reference_name: str + reference_group: str + reference_position: str + burnin: int + + +def _finalize_rf_trace_job( + ref: JobRef, + result: Any, + *, + context: _RfTraceFinalizationContext, +) -> dict[str, Any]: + """Build and store one RF trace while retaining a small terminal event.""" + if not isinstance(result, dict) or not isinstance( + result.get("distances"), list + ): + raise TypeError("RF-trace worker returned an invalid distance row") + distances = result["distances"] + if len(distances) != len(context.names): + raise ValueError( + "RF-trace worker row length does not match the matrix name index" + ) + + tree_service = get_tree_service() + tree_service.db_manager.flush() + all_trees = tree_service.db_manager._trees + if len(all_trees) > 0: + name_to_file_source = dict( + zip( + all_trees["name"].astype(str).tolist(), + all_trees["file_source"].astype(str).tolist(), + ) + ) + else: + name_to_file_source = {} + + records = [] + group_counts: dict[str, int] = {} + for index, (tree_name, distance) in enumerate( + zip(context.names, distances) + ): + if index == context.reference_index: + continue + group = ( + tree_name.rsplit("/", 1)[0] + if "/" in tree_name + else tree_name + ) + group_counts[group] = group_counts.get(group, 0) + 1 + records.append( + { + "rf_distance": int(distance), + "group": group, + "name": tree_name, + "file_source": name_to_file_source.get( + tree_name, + "(from distance matrix)", + ), + "treenum": group_counts[group], + } + ) + + if not records: + raise ValueError("No non-reference trees are available for RF Trace") + trace_df = pd.DataFrame(records) + figure = _build_rf_trace_fig( + trace_df, + context.reference_group, + context.reference_position, + burnin=context.burnin, + ) + state.store_rf_trace_result( + ref.job_id, + { + "records": records, + "figure": figure.to_dict(), + }, + ) + elapsed = float(result.get("elapsed", 0.0)) + add_log( + f"RF Trace computed for {len(records)} trees against " + f"{context.reference_name!r} in {elapsed:.3f}s." + ) + return { + "result_key": ref.job_id, + "selected_matrix": context.selected_matrix, + "reference_name": context.reference_name, + "reference_group": context.reference_group, + "reference_position": context.reference_position, + "n_trees": len(records), + "elapsed": elapsed, + } + def register_diagnostics_callbacks(): # ------ Show/hide the RF-dependent diagnostic sections ------ @@ -345,29 +450,44 @@ def populate_groups_from_distmat(selected_matrix, stored_distmats): return group_options, default_group @callback( - Output("rf-trace-plot", "children"), - Output("rf-trace-store", "data"), - Output("notifications-container", "children", allow_duplicate=True), + Output("rf-trace-plot", "children", allow_duplicate=True), + Output("rf-trace-store", "data", allow_duplicate=True), + Output("compute-rf-trace-button", "disabled", allow_duplicate=True), Output("export-rf-trace-button", "disabled", allow_duplicate=True), + Output("compute-poll-interval", "disabled", allow_duplicate=True), + Output("rf-trace-job-store", "data"), Input("compute-rf-trace-button", "n_clicks"), - State("tree-offset-store", "data"), State("rf-reference-group-select", "value"), State("rf-reference-position-select", "value"), State("distmat-store", "data"), State("diagnostics-distmat-select", "value"), State("rf-burnin-input", "value"), + State("compute-applied-job-store", "data"), prevent_initial_call=True, ) - def compute_rf_trace(n_clicks, stored_summaries, ref_group, ref_position, - stored_distmats, selected_matrix, burnin): - """Compute RF distance of every tree to a single shared reference tree using pre-computed distance matrix.""" + def compute_rf_trace( + n_clicks, + ref_group, + ref_position, + stored_distmats, + selected_matrix, + burnin, + applied_job, + ): + """Validate and submit a memory-mapped RF-trace row extraction.""" + from .compute import _ack_applied_job + if not n_clicks: - return no_update, no_update, no_update, no_update + return (no_update,) * 6 + + _ack_applied_job(applied_job) if not ref_group: return ( dmc.Text("Please select a reference group.", c="red"), no_update, + False, + no_update, no_update, no_update, ) @@ -375,45 +495,197 @@ def compute_rf_trace(n_clicks, stored_summaries, ref_group, ref_position, if not stored_distmats: return ( dmc.Text("Please compute RF distances first (Distances tab).", c="red"), - no_update, no_update, no_update, + no_update, False, no_update, no_update, no_update, ) if not selected_matrix or selected_matrix not in stored_distmats: return ( dmc.Text("Please pick an RF matrix at the top of the page.", c="red"), - no_update, no_update, no_update, + no_update, False, no_update, no_update, no_update, ) - result, ref_name = compute_rf_trace_data(selected_matrix, ref_group, ref_position) - - # If result is a string, it's an error message - if isinstance(result, str): - return dmc.Text(result, c="red"), no_update, no_update, no_update - - trace_df = result + try: + names = tuple(str(name) for name in state.get_distmat_names(selected_matrix)) + matrix_path = str(state.get_distmat_file_path(selected_matrix)) + reference_index, reference_name = find_reference_index( + names, + ref_group, + ref_position, + ) + except (KeyError, ValueError) as exc: + return ( + dmc.Text(str(exc), c="red"), + no_update, + False, + no_update, + no_update, + no_update, + ) try: - burnin = int(burnin) if burnin else 0 + burnin_int = max(0, int(burnin)) if burnin else 0 except (ValueError, TypeError): - burnin = 0 - fig = _build_rf_trace_fig(trace_df, ref_group, ref_position, burnin=burnin) - - notification = dmc.Notification( - title="RF Trace Computed", - message=f"Computed RF distances for {len(trace_df)} trees to {ref_position} tree of {ref_group}.", - color="green", - action="show", - autoClose=3000, - id=notif_id(), + burnin_int = 0 + context = _RfTraceFinalizationContext( + selected_matrix=selected_matrix, + names=names, + reference_index=reference_index, + reference_name=reference_name, + reference_group=str(ref_group), + reference_position=str(ref_position), + burnin=burnin_int, ) - store_data = trace_df.to_dict("records") + try: + job_ref = job_manager.submit( + _get_executor(), + "rf_trace", + persistent_worker.submit_job, + "compute_rf_trace", + matrix_path=matrix_path, + reference_index=reference_index, + metadata={ + "display_name": "RF Trace", + "source_distmat": selected_matrix, + }, + finalizer=partial( + _finalize_rf_trace_job, + context=context, + ), + cancel_exceptions=(persistent_worker.JobCancelled,), + ) + except JobBusyError as exc: + msg = ( + f"Another computation ({exc.active.kind.replace('_', ' ').upper()}) " + "is still finishing. Please wait for it to complete." + ) + add_log(msg, "WARNING") + return ( + dmc.Alert( + title="Computation already running", + children=dmc.Text(msg, size="sm"), + color="yellow", + variant="light", + ), + no_update, + False, + no_update, + no_update, + no_update, + ) + + spinner = dmc.Group( + [ + dmc.Loader(size="sm", type="dots"), + dmc.Text( + f"Reading the RF row for {reference_name}…", + size="sm", + c="dimmed", + ), + stop_button("rf-trace"), + ], + gap="sm", + ) + return ( + spinner, + None, + True, + True, + False, + job_ref.as_dict(), + ) + + @callback( + Output("rf-trace-plot", "children", allow_duplicate=True), + Output("rf-trace-store", "data", allow_duplicate=True), + Output("notifications-container", "children", allow_duplicate=True), + Output("export-rf-trace-button", "disabled", allow_duplicate=True), + Output("compute-rf-trace-button", "disabled", allow_duplicate=True), + Output("compute-poll-interval", "disabled", allow_duplicate=True), + Output("compute-applied-job-store", "data", allow_duplicate=True), + Input("compute-poll-interval", "n_intervals"), + Input("rf-trace-job-store", "data"), + prevent_initial_call=True, + ) + def poll_rf_trace_completion(_n_intervals, job_data): + ref = _job_ref_from_store(job_data, expected_kind="rf_trace") + if ref is None: + return (no_update,) * 7 + snapshot = job_manager.snapshot_for_delivery(ref) + if ( + snapshot is None + or snapshot.acknowledged + or snapshot.terminal is None + ): + return (no_update,) * 7 + + terminal = snapshot.terminal + payload = terminal.payload + notification = no_update + store_data = no_update + export_disabled = True + + if terminal.state is JobState.CANCELLED: + if snapshot.delivery_attempt == 1: + add_log("RF Trace computation cancelled by user.", "WARNING") + output = dmc.Alert( + title="RF Trace computation cancelled", + children=dmc.Text("Stopped before completion.", size="sm"), + color="gray", + variant="light", + ) + elif terminal.state is JobState.FAILED: + message = str(payload.get("message", "Unknown error")) + if snapshot.delivery_attempt == 1: + add_log(f"RF Trace computation failed: {message}", "ERROR") + output = dmc.Text( + f"RF Trace computation failed: {message}", + c="red", + ) + notification = dmc.Notification( + title="RF Trace Error", + message=message, + color="red", + action="show", + autoClose=6000, + id=f"rf-trace-terminal-{ref.job_id}", + ) + else: + cached = state.get_rf_trace_result(payload.get("result_key")) + if cached is None: + output = dmc.Text( + "RF Trace result is no longer available. Please recompute it.", + c="red", + ) + else: + store_data = cached["records"] + output = dcc.Graph( + id="rf-trace-graph", + figure=cached["figure"], + config={"displayModeBar": False}, + ) + export_disabled = False + notification = dmc.Notification( + title="RF Trace Computed", + message=( + f"Computed RF distances for {payload['n_trees']} trees " + f"to the {payload['reference_position']} tree of " + f"{payload['reference_group']}." + ), + color="green", + action="show", + autoClose=3000, + id=f"rf-trace-terminal-{ref.job_id}", + ) return ( - dcc.Graph(id="rf-trace-graph", figure=fig, config={"displayModeBar": False}), + output, store_data, notification, + export_disabled, False, + True, + snapshot.terminal_delivery_marker(), ) # Re-render RF trace plot when burnin changes diff --git a/src/treetracer/callbacks/persistent_worker.py b/src/treetracer/callbacks/persistent_worker.py index 93ddb91..f24c83b 100644 --- a/src/treetracer/callbacks/persistent_worker.py +++ b/src/treetracer/callbacks/persistent_worker.py @@ -31,7 +31,7 @@ Wire protocol (parent ↔ worker, framed identically on the socket): request: [4-byte LE length N][N bytes pickle.dumps({ - "job": "compute_rf" | "compute_consensus_tree", + "job": "compute_rf" | "compute_consensus_tree" | ..., "kwargs": {...}, })] response: [4-byte LE length M][M bytes pickle.dumps({ @@ -396,8 +396,8 @@ def submit_job(job_name: str, **kwargs: Any) -> Dict[str, Any]: pywebview event loop on its native thread, etc.) stay responsive. Args: - job_name: ``"compute_rf"`` or ``"compute_consensus_tree"`` — must match a - branch in ``__init__.py:_run_persistent_worker``. + job_name: Managed compute operation name; must match a branch in + ``__init__.py:_run_persistent_worker``. **kwargs: forwarded to the worker function. Returns the worker function's return value (unpickled). Raises diff --git a/src/treetracer/callbacks/sidebar.py b/src/treetracer/callbacks/sidebar.py index 9c74c23..0fccdcd 100644 --- a/src/treetracer/callbacks/sidebar.py +++ b/src/treetracer/callbacks/sidebar.py @@ -4,7 +4,12 @@ from ..logger import add_log, notif_id from ..db.tree_service import get_tree_service -from ..state import clear_all_distmats, clear_all_mds_results, clear_all_consensus_trees +from ..state import ( + clear_all_analysis_results, + clear_all_consensus_trees, + clear_all_distmats, + clear_all_mds_results, +) from ..plot_utils import placeholder_fig from .clade_explore import clear_clade_freq_caches, _tanglegram_placeholder_fig from ._helpers import _open_file_dialog @@ -630,6 +635,7 @@ def remove_file(n_clicks_list, stored_summaries): Output("clade-freq-tanglegram-title", "children", allow_duplicate=True), Output("clade-freq-data-store", "data", allow_duplicate=True), Output("clade-freq-click-store", "data", allow_duplicate=True), + Output("clade-freq-result-key-store", "data", allow_duplicate=True), Output("clade-freq-tanglegram-pair-store", "data", allow_duplicate=True), Output("clade-freq-consensus-tree-select-1", "value", allow_duplicate=True), Output("clade-freq-consensus-tree-select-2", "value", allow_duplicate=True), @@ -641,6 +647,8 @@ def remove_file(n_clicks_list, stored_summaries): Output("mds-job-store", "data", allow_duplicate=True), Output("pseudo-ess-job-store", "data", allow_duplicate=True), Output("consensus-job-store", "data", allow_duplicate=True), + Output("rf-trace-job-store", "data", allow_duplicate=True), + Output("clade-freq-job-store", "data", allow_duplicate=True), Output("compute-applied-job-store", "data", allow_duplicate=True), Output("compute-job-ack-store", "data", allow_duplicate=True), Output("rf-progress-path", "data", allow_duplicate=True), @@ -658,12 +666,8 @@ def clear_uploads(n_clicks): # Cancel any in-flight subprocess job — descriptors point # into the DB we're about to wipe, and we don't want the # poll callbacks to write results based on stale state. - from . import compute, consensus_tree_compute, pseudo_ess_compute - resetters = ( - ("RF/MDS", compute.reset), - ("consensus tree", consensus_tree_compute.reset), - ("Pseudo-ESS", pseudo_ess_compute.reset), - ) + from . import compute + resetters = (("managed", compute.reset),) for label, resetter in resetters: try: resetter() @@ -676,6 +680,7 @@ def clear_uploads(n_clicks): clear_all_distmats() clear_all_mds_results() clear_all_consensus_trees() + clear_all_analysis_results() # Wipe the in-process clade-freq caches (parsed NEXUS # trees, tanglegram layouts, click→split lookup). They're # keyed on consensus tree uuids that no longer exist after the calls @@ -739,6 +744,7 @@ def clear_uploads(n_clicks): None, # clade-freq-tanglegram-title.children None, # clade-freq-data-store None, # clade-freq-click-store + None, # clade-freq-result-key-store None, # clade-freq-tanglegram-pair-store None, # clade-freq-consensus-tree-select-1.value None, # clade-freq-consensus-tree-select-2.value @@ -747,6 +753,8 @@ def clear_uploads(n_clicks): None, # mds-job-store None, # pseudo-ess-job-store None, # consensus-job-store + None, # rf-trace-job-store + None, # clade-freq-job-store None, # compute-applied-job-store None, # compute-job-ack-store None, # rf-progress-path @@ -756,4 +764,4 @@ def clear_uploads(n_clicks): False, # treespace-loading-overlay visible False, # within-run-loading-overlay visible ) - return (no_update,) * 45 + return (no_update,) * 48 diff --git a/src/treetracer/clade_freq/_subprocess_worker.py b/src/treetracer/clade_freq/_subprocess_worker.py new file mode 100644 index 0000000..6f2a6ad --- /dev/null +++ b/src/treetracer/clade_freq/_subprocess_worker.py @@ -0,0 +1,138 @@ +"""Process-isolated clade-frequency comparison worker. + +The UI ultimately displays only clades present in either selected consensus +tree. Decoding every bipartition in the RF snapshot into Python tuples was both +CPU- and memory-heavy, especially on older hardware. This worker decodes only +that small union of column IDs and returns compact numeric rows plus one shared +leaf-name vector. +""" + +from __future__ import annotations + +from typing import Any + + +def _counts_for_columns( + snapshot, + *, + columns, + cached_counts, + cached_n_trees, + tree_names, + full_distmat_names, +): + import numpy as np + + if cached_counts is not None and int(cached_n_trees or 0) > 0: + counts = np.asarray(cached_counts, dtype=np.int64) + if counts.shape != (len(columns),): + raise ValueError( + "cached clade counts do not match the selected column count" + ) + return counts, int(cached_n_trees) + + if not tree_names or not full_distmat_names: + raise ValueError( + "clade counts are unavailable and no tree-name fallback was supplied" + ) + name_to_index = { + str(name): index for index, name in enumerate(full_distmat_names) + } + row_indices = [ + name_to_index[str(name)] + for name in tree_names + if str(name) in name_to_index + ] + if not row_indices: + raise KeyError( + "none of the selected consensus-tree names occur in the RF snapshot" + ) + + presence = snapshot["presence"] + counts = presence[np.ix_(row_indices, columns)].sum(axis=0) + return np.asarray(counts, dtype=np.int64), len(row_indices) + + +def compute_clade_frequencies_worker_entry( + *, + snapshots_path: str, + columns: list[int], + counts_1: list[int] | None, + counts_2: list[int] | None, + n_trees_1: int, + n_trees_2: int, + tree_names_1: list[str] | None = None, + tree_names_2: list[str] | None = None, + full_distmat_names: list[str] | None = None, +) -> dict[str, Any]: + """Compute frequencies and decode only requested snapshot columns.""" + import time + + import numpy as np + + started = time.perf_counter() + columns = sorted({int(column) for column in columns}) + if any(column < 0 for column in columns): + raise ValueError("clade column IDs must be non-negative") + + with np.load(snapshots_path, allow_pickle=False) as snapshot: + if "bipartition_bits" not in snapshot.files: + raise KeyError( + "RF snapshot has no bipartition_bits; recompute the RF matrix" + ) + bits = snapshot["bipartition_bits"] + if columns and columns[-1] >= bits.shape[0]: + raise IndexError( + f"clade column {columns[-1]} is outside a " + f"{bits.shape[0]}-column RF snapshot" + ) + + selected_counts_1, denominator_1 = _counts_for_columns( + snapshot, + columns=columns, + cached_counts=counts_1, + cached_n_trees=n_trees_1, + tree_names=tree_names_1, + full_distmat_names=full_distmat_names, + ) + selected_counts_2, denominator_2 = _counts_for_columns( + snapshot, + columns=columns, + cached_counts=counts_2, + cached_n_trees=n_trees_2, + tree_names=tree_names_2, + full_distmat_names=full_distmat_names, + ) + + selected_bits = bits[columns] + leaf_names = [str(name).strip("'\"") for name in snapshot["leaf_names"]] + + rows = [] + for position, column in enumerate(columns): + split_key = tuple( + int(index) + for index in np.flatnonzero(selected_bits[position]) + ) + freq_1 = float(selected_counts_1[position]) / denominator_1 + freq_2 = float(selected_counts_2[position]) / denominator_2 + rows.append( + { + "split_key": split_key, + "column_j": int(column), + "freq_1": freq_1, + "freq_2": freq_2, + "clade_size": len(split_key), + } + ) + + rows.sort( + key=lambda row: (row["freq_1"] + row["freq_2"]) / 2.0, + reverse=True, + ) + return { + "rows": rows, + "leaf_names": leaf_names, + "n_trees_1": denominator_1, + "n_trees_2": denominator_2, + "elapsed": time.perf_counter() - started, + } diff --git a/src/treetracer/ess/_rf_trace_worker.py b/src/treetracer/ess/_rf_trace_worker.py new file mode 100644 index 0000000..5f82f0d --- /dev/null +++ b/src/treetracer/ess/_rf_trace_worker.py @@ -0,0 +1,45 @@ +"""Process-isolated RF-trace row extraction. + +An RF trace needs one row of an already-computed square distance matrix. The +old callback called ``np.load`` without memory mapping in the GUI process, +materialising the complete O(n²) matrix to read O(n) values. This worker maps +the file read-only and touches only the selected row. +""" + +from __future__ import annotations + +from typing import Any + + +def compute_rf_trace_worker_entry( + *, + matrix_path: str, + reference_index: int, +) -> dict[str, Any]: + """Return one RF-distance row as plain integers for IPC.""" + import time + + import numpy as np + + started = time.perf_counter() + matrix = np.load(matrix_path, mmap_mode="r", allow_pickle=False) + if matrix.ndim != 2 or matrix.shape[0] != matrix.shape[1]: + raise ValueError( + f"RF matrix must be square; received shape {matrix.shape!r}" + ) + + reference_index = int(reference_index) + if not 0 <= reference_index < matrix.shape[0]: + raise IndexError( + f"reference index {reference_index} is outside a " + f"{matrix.shape[0]}-row RF matrix" + ) + + # ``np.asarray`` keeps the mmap-backed view; ``tolist`` performs the only + # materialisation and converts the O(n) row into a compact pickle payload. + distances = np.asarray(matrix[reference_index]).tolist() + return { + "distances": distances, + "matrix_size": int(matrix.shape[0]), + "elapsed": time.perf_counter() - started, + } diff --git a/src/treetracer/ess/rf_trace.py b/src/treetracer/ess/rf_trace.py index 155b562..2e33cc8 100644 --- a/src/treetracer/ess/rf_trace.py +++ b/src/treetracer/ess/rf_trace.py @@ -1,16 +1,32 @@ -"""RF distance trace computation. +"""RF distance trace preparation and compatibility computation. Computes the RF distance from every tree to a user-selected reference tree, using a pre-computed distance matrix. Pure computation — no Dash imports. """ +import numpy as np import pandas as pd from ..logger import add_log -from ..state import load_distmat +from ..state import get_distmat_file_path, get_distmat_names from ..db.tree_service import get_tree_service +def find_reference_index(distmat_names, ref_group, ref_position): + """Return the first/last ``(index, name)`` belonging to a group.""" + matches = [ + (index, name) + for index, name in enumerate(distmat_names) + if (str(name).rsplit("/", 1)[0] if "/" in str(name) else str(name)) + == ref_group + ] + if not matches: + raise ValueError( + f"No trees of group {ref_group!r} in the selected RF matrix." + ) + return matches[0] if ref_position == "first" else matches[-1] + + def compute_rf_trace_data(distmat_key, ref_group, ref_position): """Compute RF distances from every tree to a reference tree. @@ -42,9 +58,15 @@ def compute_rf_trace_data(distmat_key, ref_group, ref_position): (error_message, None) on failure. """ - # Load matrix from disk + # Memory-map the matrix so this compatibility API also touches only the + # selected O(n) row rather than materialising the O(n²) file. try: - distmat_names, distmat_matrix = load_distmat(distmat_key) + distmat_names = list(get_distmat_names(distmat_key)) + distmat_matrix = np.load( + get_distmat_file_path(distmat_key), + mmap_mode="r", + allow_pickle=False, + ) except KeyError: return "Distance matrix not available. Please recompute RF distances.", None @@ -60,21 +82,21 @@ def compute_rf_trace_data(distmat_key, ref_group, ref_position): n.rsplit("/", 1)[0] if "/" in n else n for n in distmat_names ] - ref_trees_in_group = [ - name for name, grp in zip(distmat_names, distmat_groups) - if grp == ref_group - ] - if not ref_trees_in_group: - msg = (f"No trees of group '{ref_group}' in distance matrix " - f"'{distmat_key}'.") + try: + ref_idx, ref_name = find_reference_index( + distmat_names, + ref_group, + ref_position, + ) + except ValueError: + msg = ( + f"No trees of group '{ref_group}' in distance matrix " + f"'{distmat_key}'." + ) add_log(msg, "ERROR") return msg, None - - ref_name = (ref_trees_in_group[0] if ref_position == "first" - else ref_trees_in_group[-1]) add_log(f"Reference tree: '{ref_name}' " f"({ref_position} of '{ref_group}' in '{distmat_key}')") - ref_idx = name_to_idx[ref_name] # Optional per-row ``file_source`` for hover. We pull it from the # DB when available, but a tree missing from the DB (e.g. dropped diff --git a/src/treetracer/state.py b/src/treetracer/state.py index f5418dc..fa5385c 100644 --- a/src/treetracer/state.py +++ b/src/treetracer/state.py @@ -358,6 +358,52 @@ def clear_all_mds_results(): _mds_results.clear() +# --------------------------------------------------------------------------- +# Managed diagnostic/exploration result storage +# --------------------------------------------------------------------------- +# JobManager deliberately retains only small terminal references. These caches +# own the larger render payloads produced by RF Trace and clade comparison so a +# dropped browser response can be replayed without retaining worker results or +# recomputing domain work. They are bounded independently because neither UI +# needs unbounded history. + +_rf_trace_results = {} +_clade_frequency_results = {} +_MAX_ANALYSIS_RESULTS = 16 + + +def _store_bounded_result(cache, key, result): + if len(cache) >= _MAX_ANALYSIS_RESULTS and key not in cache: + del cache[next(iter(cache))] + cache[key] = result + + +def store_rf_trace_result(key, result): + """Store one full RF-trace render payload under a managed-job key.""" + _store_bounded_result(_rf_trace_results, key, result) + + +def get_rf_trace_result(key): + """Return an RF-trace render payload, or ``None`` after eviction/reset.""" + return _rf_trace_results.get(key) + + +def store_clade_frequency_result(key, result): + """Store one clade-comparison render and split-resolution payload.""" + _store_bounded_result(_clade_frequency_results, key, result) + + +def get_clade_frequency_result(key): + """Return a clade-comparison payload, or ``None`` after eviction/reset.""" + return _clade_frequency_results.get(key) + + +def clear_all_analysis_results(): + """Drop managed RF-trace and clade-comparison render payloads.""" + _rf_trace_results.clear() + _clade_frequency_results.clear() + + # --------------------------------------------------------------------------- # In-memory consensus tree cache # --------------------------------------------------------------------------- diff --git a/src/treetracer/ui/navbar.py b/src/treetracer/ui/navbar.py index 246fbf4..f466a47 100644 --- a/src/treetracer/ui/navbar.py +++ b/src/treetracer/ui/navbar.py @@ -94,8 +94,14 @@ def add_navbar(): dcc.Store(id="mds-job-store", storage_type="memory"), dcc.Store(id="pseudo-ess-job-store", storage_type="memory"), dcc.Store(id="consensus-job-store", storage_type="memory"), + dcc.Store(id="rf-trace-job-store", storage_type="memory"), + dcc.Store(id="clade-freq-job-store", storage_type="memory"), dcc.Store(id="compute-applied-job-store", storage_type="memory"), dcc.Store(id="compute-job-ack-store", storage_type="memory"), + # Resolves scatter split IDs through the matching server-side + # managed comparison result; avoids shipping tip sets through + # browser JSON or decoding the full snapshot on click. + dcc.Store(id="clade-freq-result-key-store", storage_type="memory"), # Consensus trees keep a dedicated cadence because their # overlay and button lifecycle can stop independently of # the shared RF/MDS/Pseudo-ESS interval. From 295ea76a5ecc659ce703f4440baa9a50118970c5 Mon Sep 17 00:00:00 2001 From: Sam Hong <20244918+hongsamL@users.noreply.github.com> Date: Mon, 3 Aug 2026 16:42:06 +0200 Subject: [PATCH 6/9] centralize the background job coordination --- src/test/test_app_smoke.py | 52 ++- src/test/test_compute_job_lifecycle.py | 128 +++++-- src/test/test_managed_analysis_jobs.py | 60 ++-- src/test/test_managed_compute_jobs.py | 73 ++-- src/treetracer/background_jobs.py | 12 +- src/treetracer/callbacks/__init__.py | 4 + src/treetracer/callbacks/clade_explore.py | 68 ++-- src/treetracer/callbacks/compute.py | 333 +++++------------- .../callbacks/consensus_tree_compute.py | 77 ++-- src/treetracer/callbacks/diagnostics.py | 139 ++++---- src/treetracer/callbacks/job_reconcile.py | 333 ++++++++++++++++++ .../callbacks/pseudo_ess_compute.py | 72 ++-- src/treetracer/callbacks/sidebar.py | 20 +- src/treetracer/callbacks/treespace.py | 36 +- src/treetracer/callbacks/within_run.py | 23 +- src/treetracer/rf/_subprocess_worker.py | 7 +- src/treetracer/ui/navbar.py | 71 +++- src/treetracer/ui/panels/treespace.py | 2 +- src/treetracer/ui/panels/within_run.py | 2 +- 19 files changed, 928 insertions(+), 584 deletions(-) create mode 100644 src/treetracer/callbacks/job_reconcile.py diff --git a/src/test/test_app_smoke.py b/src/test/test_app_smoke.py index 8233485..1668ec7 100644 --- a/src/test/test_app_smoke.py +++ b/src/test/test_app_smoke.py @@ -10,7 +10,6 @@ from __future__ import annotations import pandas as pd -import pytest def test_app_imports_and_registers_callbacks(): @@ -33,6 +32,57 @@ def test_app_imports_and_registers_callbacks(): register_callbacks(app) +def test_compute_interval_has_one_reconciliation_owner(): + """No feature callback may independently poll or stop the shared timer.""" + from dash import _callback + + owners = set() + for callback_data in _callback.GLOBAL_CALLBACK_MAP.values(): + inputs = callback_data.get("inputs", []) + output = callback_data.get("output") + outputs = output if isinstance(output, list) else [output] + reads_interval = any( + item.get("id") == "compute-poll-interval" for item in inputs + ) + writes_interval = any( + getattr(item, "component_id", None) == "compute-poll-interval" + for item in outputs + ) + if not (reads_interval or writes_interval): + continue + callback_fn = callback_data.get("callback") + callback_fn = getattr(callback_fn, "__wrapped__", callback_fn) + owners.add(getattr(callback_fn, "__name__", "")) + + assert owners == {"reconcile_compute_job"} + + +def test_every_compute_action_reads_the_shared_busy_gate(): + """All entry points must become unavailable while the worker is owned.""" + from dash import _callback + + gated_outputs = set() + for callback_data in _callback.GLOBAL_CALLBACK_MAP.values(): + inputs = callback_data.get("inputs", []) + if not any(item.get("id") == "compute-busy-store" for item in inputs): + continue + output = callback_data.get("output") + outputs = output if isinstance(output, list) else [output] + gated_outputs.update( + getattr(item, "component_id", None) for item in outputs + ) + + assert { + "compute-rf-button", + "compute-mds-button", + "compute-rf-trace-button", + "compute-pseudo-ess-button", + "clade-freq-compare-button", + "treespace-view-consensus-tree", + "within-run-view-consensus-tree", + } <= gated_outputs + + def test_between_run_trailing_overlay_invariant(): """Between-run figure has exactly 8 trailing overlays (4 selection + 4 consensus tree). The patch callbacks address them by fixed diff --git a/src/test/test_compute_job_lifecycle.py b/src/test/test_compute_job_lifecycle.py index e38f038..769803e 100644 --- a/src/test/test_compute_job_lifecycle.py +++ b/src/test/test_compute_job_lifecycle.py @@ -9,7 +9,7 @@ import pytest from treetracer.background_jobs import JobManager, JobState -from treetracer.callbacks import compute +from treetracer.callbacks import compute, job_reconcile def _wait_for_terminal(manager, ref, timeout=2.0): @@ -22,7 +22,7 @@ def _wait_for_terminal(manager, ref, timeout=2.0): pytest.fail("background job did not become terminal") -def _registered_callback(name): +def _registered_callback(name, register=compute.register_compute_callbacks): from dash import _callback matches = [] @@ -34,12 +34,12 @@ def _registered_callback(name): if original.__name__ == name: matches.append(original) if not matches: - compute.register_compute_callbacks() - return _registered_callback(name) + register() + return _registered_callback(name, register) return matches[-1] -def test_latest_job_ref_uses_generation_and_rejects_invalid_data(): +def test_job_ref_parser_rejects_invalid_and_unknown_kinds(): rf = { "job_id": "rf-job", "generation": 4, @@ -53,15 +53,48 @@ def test_latest_job_ref_uses_generation_and_rejects_invalid_data(): "owner_id": None, } - assert compute._latest_rf_mds_ref(rf, mds).job_id == "mds-job" - assert compute._latest_rf_mds_ref(rf, None).job_id == "rf-job" - assert compute._latest_rf_mds_ref({"kind": "rf"}, None) is None - assert compute._latest_rf_mds_ref( - {**rf, "kind": "pseudo_ess"}, - None, + assert job_reconcile.job_ref_from_store(rf).job_id == "rf-job" + assert job_reconcile.job_ref_from_store(mds).job_id == "mds-job" + assert job_reconcile.job_ref_from_store({"kind": "rf"}) is None + assert job_reconcile.job_ref_from_store( + {**rf, "kind": "unknown"}, ) is None +def test_terminal_event_must_match_the_current_browser_generation(): + old_ref = { + "job_id": "old-rf-job", + "generation": 4, + "kind": "rf", + "owner_id": None, + } + current_ref = { + "job_id": "current-rf-job", + "generation": 5, + "kind": "rf", + "owner_id": None, + } + event = { + **old_ref, + "state": "succeeded", + "terminal_revision": 1, + "delivery_attempt": 2, + "payload": {}, + "metadata": {}, + } + + assert job_reconcile.terminal_event_for_job( + event, + current_ref, + expected_kind="rf", + ) is None + assert job_reconcile.terminal_event_for_job( + event, + old_ref, + expected_kind="rf", + ) == event + + def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): manager = JobManager(id_factory=lambda: "rf-test-job") register_calls = [] @@ -74,6 +107,7 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): } } monkeypatch.setattr(compute, "job_manager", manager) + monkeypatch.setattr(job_reconcile, "job_manager", manager) monkeypatch.setattr(compute, "add_log", lambda *_args, **_kwargs: None) monkeypatch.setattr(compute, "_wlog", lambda *_args, **_kwargs: None) monkeypatch.setattr( @@ -106,26 +140,76 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): assert len(register_calls) == 1 assert terminal.terminal.payload["distmat_index"] == expected_index - poll = _registered_callback("poll_completion") - acknowledge = _registered_callback("acknowledge_terminal_job") - first = poll(10, ref.as_dict(), None, False) - second = poll(11, ref.as_dict(), None, False) + reconcile = _registered_callback( + "reconcile_compute_job", + job_reconcile.register_job_reconciliation_callbacks, + ) + render = _registered_callback("render_rf_mds_terminal_event") + acknowledge = _registered_callback( + "acknowledge_terminal_receipt", + job_reconcile.register_job_reconciliation_callbacks, + ) - assert len(first) == 16 + first_reconcile = reconcile( + 10, + ref.as_dict(), + None, + None, + None, + None, + None, + {"busy": False}, + None, + None, + ) + first = render(first_reconcile[1], ref.as_dict(), None, False) + second_reconcile = reconcile( + 11, + ref.as_dict(), + None, + None, + None, + None, + None, + first_reconcile[2], + None, + None, + ) + second = render(second_reconcile[1], ref.as_dict(), None, False) + + assert len(first_reconcile) == 7 + assert first_reconcile[0] is False + assert first_reconcile[2]["busy"] is True + assert first_reconcile[3:5] == (100.0, "complete") + assert len(first) == 12 assert first[1] == expected_index - assert first[4] is False - assert first[11] is True - assert first[12]["terminal_revision"] == terminal.terminal.revision - assert first[12]["delivery_attempt"] == 1 - assert second[12]["delivery_attempt"] == 2 + assert first[2] is False + assert first[8]["terminal_revision"] == terminal.terminal.revision + assert first[8]["delivery_attempt"] == 1 + assert second[8]["delivery_attempt"] == 2 assert len(register_calls) == 1 assert manager.snapshot(ref).acknowledged is False - ack_store = acknowledge(second[12]) + ack_store = acknowledge([second[8], None, None, None, None]) assert ack_store["acknowledged"] is True assert manager.snapshot(ref).acknowledged is True assert manager.active_ref() is None + settled = reconcile( + 12, + ref.as_dict(), + None, + None, + None, + None, + None, + first_reconcile[2], + None, + None, + ) + assert settled[0] is True + assert settled[2] == {"busy": False} + def test_mds_finalization_stores_full_result_once_and_replays_small_index( monkeypatch, diff --git a/src/test/test_managed_analysis_jobs.py b/src/test/test_managed_analysis_jobs.py index 2bc8f6d..faa432c 100644 --- a/src/test/test_managed_analysis_jobs.py +++ b/src/test/test_managed_analysis_jobs.py @@ -13,7 +13,7 @@ from treetracer import state from treetracer.background_jobs import JobManager, JobState -from treetracer.callbacks import clade_explore, compute, diagnostics +from treetracer.callbacks import clade_explore, diagnostics, job_reconcile from treetracer.clade_freq._subprocess_worker import ( compute_clade_frequencies_worker_entry, ) @@ -132,7 +132,6 @@ def test_rf_trace_terminal_replays_cached_render_until_ack(monkeypatch): ) fake_figure = SimpleNamespace(to_dict=lambda: {"data": [], "layout": {}}) monkeypatch.setattr(diagnostics, "job_manager", manager) - monkeypatch.setattr(compute, "job_manager", manager) monkeypatch.setattr(diagnostics, "add_log", lambda *_a, **_k: None) monkeypatch.setattr( diagnostics, @@ -176,21 +175,25 @@ def test_rf_trace_terminal_replays_cached_render_until_ack(monkeypatch): cached = state.get_rf_trace_result(ref.job_id) assert [row["rf_distance"] for row in cached["records"]] == [4, 8] - poll = _registered_callback( - "poll_rf_trace_completion", + render = _registered_callback( + "render_rf_trace_terminal_event", diagnostics.register_diagnostics_callbacks, ) - first = poll(1, ref.as_dict()) - second = poll(2, ref.as_dict()) - assert len(first) == 7 + first_event = job_reconcile._terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + second_event = job_reconcile._terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + first = render(first_event, ref.as_dict()) + second = render(second_event, ref.as_dict()) + assert len(first) == 5 assert len(first[1]) == 2 assert first[3] is False - assert first[4] is False - assert first[5] is True - assert first[6]["delivery_attempt"] == 1 - assert second[6]["delivery_attempt"] == 2 + assert first[4]["delivery_attempt"] == 1 + assert second[4]["delivery_attempt"] == 2 - assert compute._ack_applied_job(second[6]) is True + assert manager.acknowledge(ref, second[4]["terminal_revision"]) is True assert manager.snapshot(ref).acknowledged is True @@ -204,8 +207,8 @@ def test_stage_four_submit_callbacks_have_matching_idle_output_shapes(): clade_explore.register_clade_explore_callbacks, ) - assert len(rf_submit(None, None, None, None, None, None, None)) == 6 - assert len(clade_submit(None, None, None, None, None)) == 8 + assert len(rf_submit(None, None, None, None, None, None)) == 5 + assert len(clade_submit(None, None, None, None)) == 7 def test_clade_terminal_keeps_rows_server_side_and_keys_click_resolution( @@ -214,7 +217,6 @@ def test_clade_terminal_keeps_rows_server_side_and_keys_click_resolution( manager = JobManager(id_factory=lambda: "clade-test-job") fake_figure = SimpleNamespace(to_dict=lambda: {"data": [], "layout": {}}) monkeypatch.setattr(clade_explore, "job_manager", manager) - monkeypatch.setattr(compute, "job_manager", manager) monkeypatch.setattr(clade_explore, "add_log", lambda *_a, **_k: None) monkeypatch.setattr( clade_explore, @@ -281,21 +283,25 @@ def test_clade_terminal_keeps_rows_server_side_and_keys_click_resolution( expected_pair=("uid-b", "uid-a"), ) is None - poll = _registered_callback( - "poll_clade_frequency_completion", + render = _registered_callback( + "render_clade_frequency_terminal_event", clade_explore.register_clade_explore_callbacks, ) - first = poll(1, ref.as_dict()) - second = poll(2, ref.as_dict()) - assert len(first) == 8 + first_event = job_reconcile._terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + second_event = job_reconcile._terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + first = render(first_event, ref.as_dict()) + second = render(second_event, ref.as_dict()) + assert len(first) == 6 assert len(first[1]) == 2 assert first[2] == {} - assert first[3] is False - assert first[4] is True - assert first[5] == ref.job_id - assert first[6] is None - assert first[7]["delivery_attempt"] == 1 - assert second[7]["delivery_attempt"] == 2 + assert first[3] == ref.job_id + assert first[4] is None + assert first[5]["delivery_attempt"] == 1 + assert second[5]["delivery_attempt"] == 2 - assert compute._ack_applied_job(second[7]) is True + assert manager.acknowledge(ref, second[5]["terminal_revision"]) is True assert manager.snapshot(ref).acknowledged is True diff --git a/src/test/test_managed_compute_jobs.py b/src/test/test_managed_compute_jobs.py index 31c89c0..6ea8c28 100644 --- a/src/test/test_managed_compute_jobs.py +++ b/src/test/test_managed_compute_jobs.py @@ -14,6 +14,7 @@ from treetracer.callbacks import ( compute, consensus_tree_compute, + job_reconcile, pseudo_ess_compute, sidebar, ) @@ -98,22 +99,26 @@ def worker(job_name, **kwargs): assert terminal.state is JobState.SUCCEEDED assert terminal.terminal.payload["n_rows"] == 1 - poll = _registered_callback( - "poll_pseudo_ess_completion", + render = _registered_callback( + "render_pseudo_ess_terminal_event", pseudo_ess_compute.register_pseudo_ess_compute_callbacks, ) - first = poll(10, ref.as_dict()) - second = poll(11, ref.as_dict()) - - assert len(first) == 4 - assert first[1] is False - assert first[2] is True - assert first[3]["terminal_revision"] == terminal.terminal.revision - assert first[3]["delivery_attempt"] == 1 - assert second[3]["delivery_attempt"] == 2 + first_event = job_reconcile._terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + second_event = job_reconcile._terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + first = render(first_event, ref.as_dict()) + second = render(second_event, ref.as_dict()) + + assert len(first) == 2 + assert first[1]["terminal_revision"] == terminal.terminal.revision + assert first[1]["delivery_attempt"] == 1 + assert second[1]["delivery_attempt"] == 2 assert manager.snapshot(ref).acknowledged is False - assert compute._ack_applied_job(second[3]) is True + assert manager.acknowledge(ref, second[1]["terminal_revision"]) is True assert manager.snapshot(ref).acknowledged is True assert manager.active_ref() is None @@ -189,14 +194,20 @@ def register(**kwargs): assert "nexus_bytes" not in terminal.terminal.payload assert "counts" not in terminal.terminal.payload - poll = _registered_callback( - "poll_consensus_tree_completion", + render = _registered_callback( + "render_consensus_terminal_event", consensus_tree_compute.register_consensus_tree_compute_callbacks, ) - first = poll(10, ref.as_dict()) - second = poll(11, ref.as_dict()) + first_event = job_reconcile._terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + second_event = job_reconcile._terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + first = render(first_event, ref.as_dict()) + second = render(second_event, ref.as_dict()) - assert len(first) == 12 + assert len(first) == 9 assert first[0] == { "uuid": "tree-uuid", "name": "RF_001_Between_consensus_tree_1", @@ -204,16 +215,14 @@ def register(**kwargs): assert first[1] is no_update assert first[2] == registry assert first[3:5] == (False, False) - assert first[5] is False + assert first[5] == [] assert first[6] is no_update - assert first[7] == [] - assert first[10] is True - assert first[11]["delivery_attempt"] == 1 - assert second[11]["delivery_attempt"] == 2 + assert first[8]["delivery_attempt"] == 1 + assert second[8]["delivery_attempt"] == 2 assert len(cache_calls) == 1 assert len(register_calls) == 1 - assert compute._ack_applied_job(second[11]) is True + assert manager.acknowledge(ref, second[8]["terminal_revision"]) is True assert manager.snapshot(ref).acknowledged is True @@ -254,17 +263,19 @@ def test_consensus_finalization_failure_reenables_origin_button(monkeypatch): assert terminal.terminal.payload["stage"] == "finalize" assert "Taxon A, Taxon C" in terminal.terminal.payload["message"] - poll = _registered_callback( - "poll_consensus_tree_completion", + render = _registered_callback( + "render_consensus_terminal_event", consensus_tree_compute.register_consensus_tree_compute_callbacks, ) - output = poll(1, ref.as_dict()) + event = job_reconcile._terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + output = render(event, ref.as_dict()) assert output[3:5] == (False, False) assert output[5] is no_update - assert output[6] is False - assert output[10] is True - assert output[11]["job_id"] == ref.job_id - assert output[9].id == f"consensus-terminal-{ref.job_id}" + assert output[6] is no_update + assert output[8]["job_id"] == ref.job_id + assert output[7].id == f"consensus-terminal-{ref.job_id}" def test_clear_data_idle_branch_matches_managed_lifecycle_outputs(): @@ -272,4 +283,4 @@ def test_clear_data_idle_branch_matches_managed_lifecycle_outputs(): "clear_uploads", sidebar.register_sidebar_callbacks, ) - assert len(clear_uploads(None)) == 48 + assert len(clear_uploads(None)) == 44 diff --git a/src/treetracer/background_jobs.py b/src/treetracer/background_jobs.py index 0042a17..40d2fb1 100644 --- a/src/treetracer/background_jobs.py +++ b/src/treetracer/background_jobs.py @@ -1,10 +1,10 @@ """Thread-safe lifecycle state for background compute jobs. -The Dash UI polls long-running RF, MDS, Pseudo-ESS, and consensus-tree -computations. A completed job must not become a one-shot event: an HTTP -response can be superseded before the browser applies it. ``JobManager`` -therefore keeps terminal state until the matching browser generation -acknowledges it. +The Dash UI reconciles RF, MDS, Pseudo-ESS, consensus-tree, RF Trace, and +clade-comparison computations through one polling owner. A completed job must +not become a one-shot event: an HTTP response can be superseded before the +browser applies it. ``JobManager`` therefore keeps terminal state until the +matching browser generation acknowledges it. This module deliberately has no Dash or worker imports. Job-specific code submits work with a success finalizer that performs its domain side effects @@ -116,7 +116,7 @@ def as_dict(self) -> dict[str, Any]: @dataclass(frozen=True, slots=True) class JobSnapshot: - """Immutable point-in-time view returned to poll callbacks.""" + """Immutable point-in-time view returned to the reconciler.""" ref: JobRef state: JobState diff --git a/src/treetracer/callbacks/__init__.py b/src/treetracer/callbacks/__init__.py index 61467fd..1633650 100644 --- a/src/treetracer/callbacks/__init__.py +++ b/src/treetracer/callbacks/__init__.py @@ -9,6 +9,7 @@ from .pseudo_ess_compute import register_pseudo_ess_compute_callbacks from .clade_explore import register_clade_explore_callbacks from .rename_consensus_tree import register_rename_consensus_tree_callbacks +from .job_reconcile import register_job_reconciliation_callbacks def register_callbacks(app): @@ -23,3 +24,6 @@ def register_callbacks(app): register_pseudo_ess_compute_callbacks() register_clade_explore_callbacks() register_rename_consensus_tree_callbacks() + # Register last so every feature's submit/presentation adapter is already + # present when the single global compute reconciler is added. + register_job_reconciliation_callbacks() diff --git a/src/treetracer/callbacks/clade_explore.py b/src/treetracer/callbacks/clade_explore.py index 0df710a..fc3ae18 100644 --- a/src/treetracer/callbacks/clade_explore.py +++ b/src/treetracer/callbacks/clade_explore.py @@ -26,7 +26,12 @@ from ..plot_utils import retheme_figure from ..ui.widgets import stop_button from . import persistent_worker -from .compute import _get_executor, _job_ref_from_store +from .compute import _get_executor +from .job_reconcile import ( + is_compute_busy, + terminal_delivery_marker, + terminal_event_for_job, +) # --------------------------------------------------------------------------- @@ -931,10 +936,11 @@ def filter_clade_freq_plot(min_clade_size, store_data): Output("clade-freq-compare-button", "disabled"), Input("clade-freq-consensus-tree-select-1", "value"), Input("clade-freq-consensus-tree-select-2", "value"), + Input("compute-busy-store", "data"), ) - def toggle_compare_button(uid1, uid2): + def toggle_compare_button(uid1, uid2, compute_busy): """Enable the Compare button only when both dropdowns have a selection.""" - return not (uid1 and uid2) + return bool(is_compute_busy(compute_busy) or not (uid1 and uid2)) # ------ Clade Frequency Comparison: compute and plot ------ @@ -943,7 +949,6 @@ def toggle_compare_button(uid1, uid2): Output("clade-freq-data-store", "data", allow_duplicate=True), Output("clade-freq-output-paper", "style", allow_duplicate=True), Output("clade-freq-compare-button", "disabled", allow_duplicate=True), - Output("compute-poll-interval", "disabled", allow_duplicate=True), Output("clade-freq-job-store", "data"), Output("clade-freq-result-key-store", "data", allow_duplicate=True), Output("clade-freq-click-store", "data", allow_duplicate=True), @@ -951,7 +956,6 @@ def toggle_compare_button(uid1, uid2): State("clade-freq-consensus-tree-select-1", "value"), State("clade-freq-consensus-tree-select-2", "value"), State("clade-freq-min-clade-size", "value"), - State("compute-applied-job-store", "data"), prevent_initial_call=True, ) def compute_and_plot_clade_frequencies( @@ -959,15 +963,10 @@ def compute_and_plot_clade_frequencies( uid1, uid2, min_clade_size, - applied_job, ): """Prepare and submit a selective, process-isolated comparison.""" - from .compute import _ack_applied_job - if not n_clicks or not uid1 or not uid2: - return (no_update,) * 8 - - _ack_applied_job(applied_job) + return (no_update,) * 7 def error(message): return ( @@ -978,7 +977,6 @@ def error(message): no_update, no_update, no_update, - no_update, ) entry1 = state.get_consensus_tree_registry_entry(uid1) @@ -1095,7 +1093,6 @@ def error(message): None, {}, True, - False, job_ref.as_dict(), None, None, @@ -1105,33 +1102,32 @@ def error(message): Output("clade-freq-plot", "children", allow_duplicate=True), Output("clade-freq-data-store", "data", allow_duplicate=True), Output("clade-freq-output-paper", "style", allow_duplicate=True), - Output("clade-freq-compare-button", "disabled", allow_duplicate=True), - Output("compute-poll-interval", "disabled", allow_duplicate=True), Output("clade-freq-result-key-store", "data", allow_duplicate=True), Output("clade-freq-click-store", "data", allow_duplicate=True), - Output("compute-applied-job-store", "data", allow_duplicate=True), - Input("compute-poll-interval", "n_intervals"), + Output( + {"type": "compute-terminal-receipt", "kind": "clade-compare"}, + "data", + ), + Input("compute-terminal-event-store", "data"), Input("clade-freq-job-store", "data"), prevent_initial_call=True, ) - def poll_clade_frequency_completion(_n_intervals, job_data): - ref = _job_ref_from_store(job_data, expected_kind="clade_compare") - if ref is None: - return (no_update,) * 8 - snapshot = job_manager.snapshot_for_delivery(ref) - if ( - snapshot is None - or snapshot.acknowledged - or snapshot.terminal is None - ): - return (no_update,) * 8 + def render_clade_frequency_terminal_event(terminal_event, job_data): + event = terminal_event_for_job( + terminal_event, + job_data, + expected_kind="clade_compare", + ) + if event is None: + return (no_update,) * 6 - terminal = snapshot.terminal - payload = terminal.payload + terminal_state = JobState(str(event["state"])) + payload = event["payload"] + first_delivery = int(event["delivery_attempt"]) == 1 store_data = no_update result_key = no_update - if terminal.state is JobState.CANCELLED: - if snapshot.delivery_attempt == 1: + if terminal_state is JobState.CANCELLED: + if first_delivery: add_log("Clade comparison cancelled by user.", "WARNING") output = dmc.Alert( title="Clade comparison cancelled", @@ -1139,9 +1135,9 @@ def poll_clade_frequency_completion(_n_intervals, job_data): color="gray", variant="light", ) - elif terminal.state is JobState.FAILED: + elif terminal_state is JobState.FAILED: message = str(payload.get("message", "Unknown error")) - if snapshot.delivery_attempt == 1: + if first_delivery: add_log(f"Clade comparison failed: {message}", "ERROR") output = dmc.Text( f"Error computing clade frequencies: {message}", @@ -1173,11 +1169,9 @@ def poll_clade_frequency_completion(_n_intervals, job_data): output, store_data, {}, - False, - True, result_key, None, - snapshot.terminal_delivery_marker(), + terminal_delivery_marker(event), ) @callback( diff --git a/src/treetracer/callbacks/compute.py b/src/treetracer/callbacks/compute.py index 486c4a3..2357af7 100644 --- a/src/treetracer/callbacks/compute.py +++ b/src/treetracer/callbacks/compute.py @@ -4,7 +4,7 @@ import pandas as pd from ..logger import add_log, notif_id -from ..background_jobs import JobBusyError, JobRef, JobState, job_manager +from ..background_jobs import JobBusyError, JobState, job_manager from ..db.tree_service import get_tree_service from ..state import (load_distmat, get_distmat_index, next_distmat_name, get_distmat_path, register_distmat, get_distmat_file_path, @@ -14,6 +14,12 @@ from ._helpers import _save_file_dialog, extract_group from .._worker_log import log as _wlog from . import persistent_worker +from .job_reconcile import ( + is_compute_busy, + job_ref_from_store as _job_ref_from_store, + terminal_delivery_marker, + terminal_event_for_job, +) # RF compute happens in a SUBPROCESS, not a thread. See @@ -38,66 +44,6 @@ def _mds_export_filename(source_distmat): return f"{stem}_MDS.tsv" -def _job_ref_from_store(data, *, expected_kind=None): - """Parse a browser job reference defensively. - - Browser stores can be empty, stale, or manually modified. Invalid data is - treated as absent; ``JobManager`` performs the authoritative generation - check on every operation. - """ - if not isinstance(data, dict): - return None - try: - ref = JobRef.from_dict(data) - except (KeyError, TypeError, ValueError): - return None - if expected_kind is not None and ref.kind != expected_kind: - return None - return ref - - -def _latest_rf_mds_ref(rf_job_data, mds_job_data): - refs = [ - ref - for ref in ( - _job_ref_from_store(rf_job_data, expected_kind="rf"), - _job_ref_from_store(mds_job_data, expected_kind="mds"), - ) - if ref is not None - ] - return max(refs, key=lambda ref: ref.generation, default=None) - - -def _ack_applied_job(applied_data): - """Acknowledge a terminal event known to have reached browser state.""" - ref = _job_ref_from_store(applied_data) - if ref is None: - return False - try: - revision = int(applied_data["terminal_revision"]) - except (KeyError, TypeError, ValueError): - return False - return job_manager.acknowledge(ref, revision) - - -def _terminal_progress_from_store(job_data, expected_kind): - ref = _job_ref_from_store(job_data, expected_kind=expected_kind) - if ref is None: - return None - snapshot = job_manager.snapshot(ref) - if snapshot is None or snapshot.terminal is None: - return None - progress = snapshot.progress - if progress is None: - return None - label = { - JobState.SUCCEEDED: "complete", - JobState.FAILED: "failed", - JobState.CANCELLED: "cancelled", - }[snapshot.state] - return progress.fraction * 100.0, label - - def _get_executor(): global _executor if _executor is None: @@ -172,9 +118,8 @@ def _rf_pipeline(selected_files, save_path, rf_name, is_rooted): ) # Sidecar file the worker will keep up-to-date with the rapidtrees - # ProgressCounter state every ~100ms. The Dash poll callback - # (``update_rf_progress``) reads this file to drive the progress bar - # in the computing banner. + # ProgressCounter state every ~100ms. The central 250ms reconciler samples + # this file to drive the progress bar in the computing banner. progress_path = save_path + ".progress" _wlog(f"[parent] _rf_pipeline: about to call persistent_worker.submit_job for {rf_name!r}") @@ -342,8 +287,8 @@ def register_compute_callbacks(): # ``persistent_worker.cancel_current_job``). One callback serves # every banner's Stop button via the {"type": "compute-stop", ...} # pattern-matching id. The kill makes the in-flight job's future - # raise JobCancelled; the per-compute poll callbacks below catch it - # and clear the banner ~one 100ms tick later. + # raise JobCancelled; JobManager records the terminal state and the central + # reconciler delivers it to the matching feature adapter. @callback( Output("notifications-container", "children", allow_duplicate=True), Input({"type": "compute-stop", "which": ALL}, "n_clicks"), @@ -394,8 +339,9 @@ def handle_compute_stop(_stop_clicks, rf_progress_path, mds_progress_path): Output("compute-trees-table", "children"), Output("compute-rf-button", "disabled"), Input("tree-offset-store", "data"), + Input("compute-busy-store", "data"), ) - def render_compute_table(stored_summaries): + def render_compute_table(stored_summaries, compute_busy): if not stored_summaries: return html.Div( dmc.Text("No .trees files loaded yet.", c="dimmed"), @@ -465,7 +411,7 @@ def render_compute_table(stored_summaries): className="tt-compute-table", ) - return table, False + return table, is_compute_busy(compute_busy) @callback( Output("compute-export-drawer", "opened"), @@ -482,18 +428,16 @@ def open_export_drawer(_rf_clicks, _mds_clicks): @callback( Output("notifications-container", "children", allow_duplicate=True), Output("compute-rf-output", "children"), - Output("compute-poll-interval", "disabled"), Output("compute-rf-button", "disabled", allow_duplicate=True), - # Progress-path Store — populated when an RF compute actually - # starts so ``update_rf_progress`` knows which sidecar file to - # tail. Early returns leave it at no_update. + # Progress-path Store — populated when an RF compute actually starts + # so the central reconciler knows which sidecar file to sample. Early + # returns leave it at no_update. Output("rf-progress-path", "data", allow_duplicate=True), Output("rf-job-store", "data"), Input("compute-rf-button", "n_clicks"), State({"type": "compute-tree-checkbox", "index": ALL}, "checked"), State({"type": "compute-tree-checkbox", "index": ALL}, "id"), State("tree-offset-store", "data"), - State("compute-applied-job-store", "data"), prevent_initial_call=True, ) def handle_compute_rf( @@ -501,15 +445,9 @@ def handle_compute_rf( checked_list, id_list, stored_summaries, - applied_job, ): if not n_clicks or not stored_summaries: - return (no_update,) * 6 - - # The marker is written in the same browser response that rendered the - # prior terminal UI. A fast next click may beat the acknowledgement - # callback, so acknowledge idempotently here before submitting too. - _ack_applied_job(applied_job) + return (no_update,) * 5 # Determine which files are selected selected_files = [ @@ -530,7 +468,7 @@ def handle_compute_rf( action="show", autoClose=6000, id=notif_id(), - ), no_update, no_update, no_update, no_update, no_update + ), no_update, no_update, no_update, no_update # Collect taxa counts for selected files taxa_counts = {} @@ -554,7 +492,7 @@ def handle_compute_rf( action="show", autoClose=6000, id=notif_id(), - ), no_update, no_update, no_update, no_update, no_update + ), no_update, no_update, no_update, no_update # --- Rooting consistency check --- # RF over rooted clades and RF over bipartitions are different @@ -579,7 +517,7 @@ def handle_compute_rf( title="Rooting Mismatch", message=msg, color="red", action="show", autoClose=8000, id=notif_id(), - ), no_update, no_update, no_update, no_update, no_update + ), no_update, no_update, no_update, no_update selected_is_rooted = unique_rootings.pop() # --- Taxa validation passed — submit the pipeline to a worker thread --- @@ -638,7 +576,6 @@ def handle_compute_rf( no_update, no_update, no_update, - no_update, ) computing_indicator = computing_banner( @@ -650,92 +587,16 @@ def handle_compute_rf( which="rf", show_progress=True, ) - # Hand the same path ``_rf_pipeline`` derives down to the - # Store so ``update_rf_progress`` reads from the right file. + # Hand the same path ``_rf_pipeline`` derives to the Store so the + # central reconciler reads the right file. return ( no_update, computing_indicator, - False, True, progress_path, job_ref.as_dict(), ) - # Per-tick reader for the RF progress sidecar file. Runs off the - # same ``compute-poll-interval`` as ``poll_completion`` but writes - # to disjoint Outputs (the progress-bar value + label), so it - # doesn't trip Dash 4.x's same-input/same-output duplicate check. - @callback( - Output("rf-progress-bar", "value"), - Output("rf-progress-label", "children"), - Input("compute-poll-interval", "n_intervals"), - State("rf-progress-path", "data"), - State("rf-job-store", "data"), - prevent_initial_call=True, - ) - def update_rf_progress(_n, progress_path, job_data): - terminal_progress = _terminal_progress_from_store(job_data, "rf") - if terminal_progress is not None: - return terminal_progress - if not progress_path: - return no_update, no_update - import json - from pathlib import Path - try: - data = json.loads(Path(progress_path).read_text()) - except (OSError, ValueError): - # File doesn't exist yet, was just deleted, or caught - # mid-write — try again next tick. ValueError covers - # JSONDecodeError (subclass) too. - terminal_progress = _terminal_progress_from_store(job_data, "rf") - return terminal_progress or (no_update, no_update) - val = int(data.get("value", 0)) - tot = int(data.get("total", 0)) - frac = max(0.0, min(float(data.get("fraction", 0.0)), 1.0)) - phase = data.get("phase", "computing") - ref = _job_ref_from_store(job_data, expected_kind="rf") - if ref is not None: - job_manager.update_progress(ref, frac, phase) - if tot == 0: - return 0, "starting…" - pct = frac * 100.0 - if phase == "finalizing": - label = f"{val:,} / {tot:,} pairs — finalizing…" - else: - label = f"{val:,} / {tot:,} pairs ({pct:.1f}%)" - return pct, label - - @callback( - Output("mds-progress-bar", "value"), - Output("mds-progress-label", "children"), - Input("compute-poll-interval", "n_intervals"), - State("mds-progress-path", "data"), - State("mds-job-store", "data"), - prevent_initial_call=True, - ) - def update_mds_progress(_n, progress_path, job_data): - terminal_progress = _terminal_progress_from_store(job_data, "mds") - if terminal_progress is not None: - return terminal_progress - if not progress_path: - return no_update, no_update - import json - from pathlib import Path - try: - data = json.loads(Path(progress_path).read_text()) - except (OSError, ValueError): - terminal_progress = _terminal_progress_from_store(job_data, "mds") - return terminal_progress or (no_update, no_update) - - frac = max(0.0, min(float(data.get("fraction", 0.0)), 1.0)) - pct = frac * 100.0 - phase = data.get("phase", "computing") - label = data.get("label") or phase.replace("_", " ") - ref = _job_ref_from_store(job_data, expected_kind="mds") - if ref is not None: - job_manager.update_progress(ref, frac, phase, label) - return pct, f"{label} ({pct:.0f}%)" - # ------ RF MATRIX LIST (right column) ------ @callback( @@ -789,8 +650,9 @@ def show_rf_matrix_info(selected, distmat_data): Output("mds-distmat-select", "data"), Output("mds-distmat-select", "value"), Input("distmat-store", "data"), + Input("compute-busy-store", "data"), ) - def toggle_mds_button(distmat_data): + def toggle_mds_button(distmat_data, compute_busy): if not distmat_data: return True, dmc.Text("No distance matrix computed yet.", c="dimmed", style={"padding": "20px"}), [], None options = [ @@ -801,7 +663,7 @@ def toggle_mds_button(distmat_data): ] last_key = list(distmat_data.keys())[-1] status = dmc.Text(f"{len(distmat_data)} distance matrix(es) available", c="green") - return False, status, options, last_key + return is_compute_busy(compute_busy), status, options, last_key # Show file breakdown badges when a matrix is selected @callback( @@ -825,20 +687,16 @@ def show_distmat_info(selected, distmat_data): # Start MDS computation in background process @callback( Output("compute-mds-output", "children"), - Output("compute-poll-interval", "disabled", allow_duplicate=True), Output("compute-mds-button", "disabled", allow_duplicate=True), Output("mds-progress-path", "data", allow_duplicate=True), Output("mds-job-store", "data"), Input("compute-mds-button", "n_clicks"), State("mds-distmat-select", "value"), - State("compute-applied-job-store", "data"), prevent_initial_call=True, ) - def handle_compute_mds(n_clicks, selected_distmat, applied_job): + def handle_compute_mds(n_clicks, selected_distmat): if not n_clicks or not selected_distmat: - return (no_update,) * 5 - - _ack_applied_job(applied_job) + return (no_update,) * 4 try: # The worker loads the matrix from its registered path. Pull only @@ -848,7 +706,7 @@ def handle_compute_mds(n_clicks, selected_distmat, applied_job): except KeyError: msg = "Selected distance matrix not available. Please recompute RF distances." add_log(msg, "ERROR") - return dmc.Text(msg, c="red"), no_update, no_update, no_update, no_update + return dmc.Text(msg, c="red"), no_update, no_update, no_update n = len(tree_names) n_components = min(6, n - 1) @@ -890,7 +748,6 @@ def handle_compute_mds(n_clicks, selected_distmat, applied_job): color="yellow", variant="light", ), - no_update, False, no_update, no_update, @@ -905,74 +762,69 @@ def handle_compute_mds(n_clicks, selected_distmat, applied_job): which="mds", show_progress=True, ) - return computing_indicator, False, True, progress_path, job_ref.as_dict() + return computing_indicator, True, progress_path, job_ref.as_dict() - # ------ POLL + RENDER: replay sticky terminal state until browser ack ------ + # ------ PRESENT: render central, generation-checked terminal events ------ @callback( - # RF outputs (5) + # RF result outputs (3). Compute-button availability is owned by the + # shared busy gate and each button's prerequisite callback. Output("compute-rf-output", "children", allow_duplicate=True), Output("distmat-store", "data", allow_duplicate=True), Output("export-rf-button", "disabled", allow_duplicate=True), - Output("compute-rf-trace-button", "disabled", allow_duplicate=True), - Output("compute-rf-button", "disabled", allow_duplicate=True), - # Between-run MDS outputs (5) + # Between-run MDS result outputs (4) Output("mds-result-store", "data"), Output("compute-mds-output", "children", allow_duplicate=True), Output("export-mds-button", "disabled"), Output("plot-config-store", "data", allow_duplicate=True), - Output("compute-mds-button", "disabled", allow_duplicate=True), - # Shared outputs (3). The applied marker lands in the same browser - # response as the result UI; only that marker authorizes the separate - # acknowledgement callback to consume the retained terminal event. + # Shared notification + a dedicated receipt. The receipt lands in the + # same browser response as the terminal UI and is consumed by the one + # acknowledgement sink in ``job_reconcile.py``. Output("notifications-container", "children", allow_duplicate=True), - Output("compute-poll-interval", "disabled", allow_duplicate=True), - Output("compute-applied-job-store", "data", allow_duplicate=True), + Output( + {"type": "compute-terminal-receipt", "kind": "rf-mds"}, + "data", + ), # Auto-collapse sidebar on RF success (3 outputs mirror the # shell sidebar-toggle callback's outputs). Output("navbar", "style", allow_duplicate=True), Output("sidebar-visible", "data", allow_duplicate=True), Output("appshell", "navbar", allow_duplicate=True), - Input("compute-poll-interval", "n_intervals"), + Input("compute-terminal-event-store", "data"), Input("rf-job-store", "data"), Input("mds-job-store", "data"), State("sidebar-visible", "data"), prevent_initial_call=True, ) - def poll_completion( - n_intervals, + def render_rf_mds_terminal_event( + terminal_event, rf_job_data, mds_job_data, sidebar_visible, ): - ref = _latest_rf_mds_ref(rf_job_data, mds_job_data) - if ref is None: - return (no_update,) * 16 - - snapshot = job_manager.snapshot_for_delivery(ref) - if snapshot is None or snapshot.acknowledged: - return (no_update,) * 16 - - if snapshot.terminal is None: - if (n_intervals or 0) % 10 == 0: - _wlog( - f"[parent] poll_completion tick={n_intervals}: " - f"job={ref.job_id}/{ref.kind}/generation-{ref.generation}, " - f"state={snapshot.state.value}" - ) - return (no_update,) * 16 - - terminal = snapshot.terminal - payload = terminal.payload - _wlog( - f"[parent] poll_completion tick={n_intervals}: " - f"delivering job={ref.job_id}/{ref.kind}/generation-{ref.generation}, " - f"state={terminal.state.value}, revision={terminal.revision}, " - f"attempt={snapshot.delivery_attempt}" + event = terminal_event_for_job( + terminal_event, + rf_job_data, + expected_kind="rf", + ) + if event is None: + event = terminal_event_for_job( + terminal_event, + mds_job_data, + expected_kind="mds", ) + if event is None: + return (no_update,) * 12 - rf_out = [no_update] * 5 - mds_out = [no_update] * 5 + ref = _job_ref_from_store(event) + if ref is None: + return (no_update,) * 12 + state = JobState(str(event["state"])) + payload = event["payload"] + first_delivery = int(event["delivery_attempt"]) == 1 + + rf_out = [no_update] * 3 + mds_out = [no_update] * 4 notif = no_update # Sidebar outputs: (navbar.style, sidebar-visible, appshell.navbar). # Only flipped on RF success when the sidebar is currently open; @@ -981,20 +833,23 @@ def poll_completion( sidebar_out = [no_update, no_update, no_update] if ref.kind == "rf": - if terminal.state is JobState.CANCELLED: - add_log("RF computation cancelled by user.", "WARNING") + if state is JobState.CANCELLED: + if first_delivery: + add_log("RF computation cancelled by user.", "WARNING") rf_out = [ dmc.Alert( title="RF computation cancelled", children=dmc.Text("Stopped before completion.", size="sm"), color="gray", variant="light", ), - no_update, no_update, no_update, False, + no_update, + no_update, ] - elif terminal.state is JobState.FAILED: + elif state is JobState.FAILED: msg = f"RF computation failed: {payload.get('message', 'Unknown error')}" - add_log(msg, "ERROR") - rf_out = [dmc.Text(msg, c="red"), no_update, no_update, no_update, False] + if first_delivery: + add_log(msg, "ERROR") + rf_out = [dmc.Text(msg, c="red"), no_update, no_update] notif = dmc.Notification( title="RF Computation Error", message=msg, @@ -1012,7 +867,7 @@ def poll_completion( children=dmc.Text(f"{n_trees} x {n_trees} trees", size="sm"), color="green", variant="light"), payload["distmat_index"], - False, False, False, + False, ] notif = dmc.Notification( title=f"RF Distances Computed ({rf_name})", @@ -1035,8 +890,9 @@ def poll_completion( ] else: - if terminal.state is JobState.CANCELLED: - add_log("MDS computation cancelled by user.", "WARNING") + if state is JobState.CANCELLED: + if first_delivery: + add_log("MDS computation cancelled by user.", "WARNING") mds_out = [ no_update, dmc.Alert( @@ -1044,12 +900,19 @@ def poll_completion( children=dmc.Text("Stopped before completion.", size="sm"), color="gray", variant="light", ), - no_update, no_update, False, + no_update, + no_update, ] - elif terminal.state is JobState.FAILED: + elif state is JobState.FAILED: msg = f"MDS computation failed: {payload.get('message', 'Unknown error')}" - add_log(msg, "ERROR") - mds_out = [no_update, dmc.Text(msg, c="red"), no_update, no_update, False] + if first_delivery: + add_log(msg, "ERROR") + mds_out = [ + no_update, + dmc.Text(msg, c="red"), + no_update, + no_update, + ] notif = dmc.Notification( title="MDS Error", message=msg, @@ -1069,7 +932,8 @@ def poll_completion( dmc.Alert(title="MDS Embedding Complete", children=dmc.Text(f"{mds_filename}: {n_points} points, {n_components}D, {n_groups} groups", size="sm"), color="green", variant="light"), - False, {}, False, + False, + {}, ] notif = dmc.Notification( title="MDS Computed", @@ -1083,19 +947,8 @@ def poll_completion( id=f"mds-terminal-{ref.job_id}", ) - applied = snapshot.terminal_delivery_marker() - return (*rf_out, *mds_out, notif, True, applied, *sidebar_out) - - @callback( - Output("compute-job-ack-store", "data"), - Input("compute-applied-job-store", "data"), - prevent_initial_call=True, - ) - def acknowledge_terminal_job(applied_job): - if not isinstance(applied_job, dict): - return no_update - acknowledged = _ack_applied_job(applied_job) - return {**applied_job, "acknowledged": acknowledged} + receipt = terminal_delivery_marker(event) + return (*rf_out, *mds_out, notif, receipt, *sidebar_out) # ------ EXPORT CALLBACKS ------ diff --git a/src/treetracer/callbacks/consensus_tree_compute.py b/src/treetracer/callbacks/consensus_tree_compute.py index c037eb3..9e1ece5 100644 --- a/src/treetracer/callbacks/consensus_tree_compute.py +++ b/src/treetracer/callbacks/consensus_tree_compute.py @@ -1,4 +1,4 @@ -"""Managed consensus-tree dispatch, publication, and terminal polling. +"""Managed consensus-tree dispatch, publication, and terminal presentation. Both the Between-run (``treespace``) and Within-run (``within_run``) tabs have a "View consensus tree" button. They used to each call @@ -10,10 +10,9 @@ * Click → enqueue a consensus tree compute job on the persistent worker, show a loading overlay over the active tab, disable the View consensus tree button. -* Wait → a dedicated interval reads a sticky ``JobManager`` snapshot. The - success finalizer caches the NEXUS bytes and registers the tree exactly once; - polling only renders the retained terminal payload and routes it to the - originating tab. +* Wait → the shared reconciler publishes a sticky terminal event. The success + finalizer caches the NEXUS bytes and registers the tree exactly once; this + module only renders a matching event and routes it to the originating tab. The per-tab callbacks ``view_consensus_tree`` in ``treespace.py`` and ``within_run.py`` shrink to ~30 lines each — they're only responsible @@ -36,7 +35,11 @@ from ..logger import add_log from ..consensus_tree import extract_log_posterior from . import persistent_worker -from .compute import _get_executor, _job_ref_from_store +from .compute import _get_executor +from .job_reconcile import ( + terminal_delivery_marker, + terminal_event_for_job, +) _TREESPACE_TARGET = "treespace-view-consensus-tree-store" @@ -172,12 +175,12 @@ def submit_consensus_tree_job( re-create the orange selection ring. run: the run/group name for Within mode, else ``None``. consensus_tree_coord_by_tree_name: ``{tree_name: (group, treenum)}`` — - consulted by the polling callback to populate + consulted by the terminal presentation adapter to populate ``consensus_tree.treenum`` (used to put the green ring on the consensus tree's MDS dot). store_target: ``"treespace-view-consensus-tree-store"`` or - ``"within-run-view-consensus-tree-store"`` — tells the polling - callback which tab's clientside ``window.open`` to fire. + ``"within-run-view-consensus-tree-store"`` — tells the terminal + adapter which tab's clientside ``window.open`` to fire. """ if not matched_records: raise ValueError("matched_records must not be empty") @@ -274,55 +277,48 @@ def register_consensus_tree_compute_callbacks(): # Both overlays are dismissed in case the user changed tabs. Output("treespace-loading-overlay", "visible", allow_duplicate=True), Output("within-run-loading-overlay", "visible", allow_duplicate=True), - # Only the originating button is re-enabled. - Output("treespace-view-consensus-tree", "disabled", allow_duplicate=True), - Output("within-run-view-consensus-tree", "disabled", allow_duplicate=True), # The originating selection is cleared on success. Output("treespace-selected-trees-store", "data", allow_duplicate=True), Output("within-run-selected-trees-store", "data", allow_duplicate=True), Output("notifications-container", "children", allow_duplicate=True), - Output("consensus-tree-poll-interval", "disabled", allow_duplicate=True), - Output("compute-applied-job-store", "data", allow_duplicate=True), - Input("consensus-tree-poll-interval", "n_intervals"), + Output( + {"type": "compute-terminal-receipt", "kind": "consensus"}, + "data", + ), + Input("compute-terminal-event-store", "data"), Input("consensus-job-store", "data"), prevent_initial_call=True, ) - def poll_consensus_tree_completion(n_intervals, job_data): - ref = _job_ref_from_store(job_data, expected_kind="consensus") - if ref is None: - return (no_update,) * 12 - - snapshot = job_manager.snapshot_for_delivery(ref) - if snapshot is None or snapshot.acknowledged: - return (no_update,) * 12 - if snapshot.terminal is None: - return (no_update,) * 12 + def render_consensus_terminal_event(terminal_event, job_data): + event = terminal_event_for_job( + terminal_event, + job_data, + expected_kind="consensus", + ) + if event is None: + return (no_update,) * 9 - terminal = snapshot.terminal - payload = terminal.payload - store_target = snapshot.metadata.get("store_target") + ref = JobRef.from_dict(event) + terminal_state = JobState(str(event["state"])) + payload = event["payload"] + store_target = event["metadata"].get("store_target") + first_delivery = int(event["delivery_attempt"]) == 1 out_treespace_store = no_update out_within_store = no_update out_registry = no_update out_treespace_overlay = False out_within_overlay = False - out_treespace_btn = no_update - out_within_btn = no_update out_treespace_sel = no_update out_within_sel = no_update - if store_target == _TREESPACE_TARGET: - out_treespace_btn = False - elif store_target == _WITHIN_RUN_TARGET: - out_within_btn = False notif = no_update - if terminal.state is JobState.CANCELLED: - if snapshot.delivery_attempt == 1: + if terminal_state is JobState.CANCELLED: + if first_delivery: add_log("Consensus tree computation cancelled by user.", "WARNING") - elif terminal.state is JobState.FAILED: + elif terminal_state is JobState.FAILED: message = str(payload.get("message", "Unknown error")) - if snapshot.delivery_attempt == 1: + if first_delivery: add_log(f"Consensus tree computation failed: {message}", "ERROR") notif = dmc.Notification( title="Consensus tree Error", @@ -359,11 +355,8 @@ def poll_consensus_tree_completion(n_intervals, job_data): out_registry, out_treespace_overlay, out_within_overlay, - out_treespace_btn, - out_within_btn, out_treespace_sel, out_within_sel, notif, - True, - snapshot.terminal_delivery_marker(), + terminal_delivery_marker(event), ) diff --git a/src/treetracer/callbacks/diagnostics.py b/src/treetracer/callbacks/diagnostics.py index 71adb49..4ad9da9 100644 --- a/src/treetracer/callbacks/diagnostics.py +++ b/src/treetracer/callbacks/diagnostics.py @@ -19,7 +19,12 @@ from ..theme import get_template from ..ui.widgets import stop_button from . import persistent_worker -from .compute import _get_executor, _job_ref_from_store +from .compute import _get_executor +from .job_reconcile import ( + is_compute_busy, + terminal_delivery_marker, + terminal_event_for_job, +) def _build_rf_trace_fig(trace_df, ref_group, ref_position, burnin=0): @@ -449,12 +454,28 @@ def populate_groups_from_distmat(selected_matrix, stored_distmats): default_group = all_groups[-1] if all_groups else None return group_options, default_group + @callback( + Output("compute-rf-trace-button", "disabled"), + Input("diagnostics-distmat-select", "value"), + Input("rf-reference-group-select", "value"), + Input("compute-busy-store", "data"), + ) + def toggle_compute_rf_trace_button( + selected_matrix, + reference_group, + compute_busy, + ): + return bool( + is_compute_busy(compute_busy) + or not selected_matrix + or not reference_group + ) + @callback( Output("rf-trace-plot", "children", allow_duplicate=True), Output("rf-trace-store", "data", allow_duplicate=True), Output("compute-rf-trace-button", "disabled", allow_duplicate=True), Output("export-rf-trace-button", "disabled", allow_duplicate=True), - Output("compute-poll-interval", "disabled", allow_duplicate=True), Output("rf-trace-job-store", "data"), Input("compute-rf-trace-button", "n_clicks"), State("rf-reference-group-select", "value"), @@ -462,7 +483,6 @@ def populate_groups_from_distmat(selected_matrix, stored_distmats): State("distmat-store", "data"), State("diagnostics-distmat-select", "value"), State("rf-burnin-input", "value"), - State("compute-applied-job-store", "data"), prevent_initial_call=True, ) def compute_rf_trace( @@ -472,15 +492,10 @@ def compute_rf_trace( stored_distmats, selected_matrix, burnin, - applied_job, ): """Validate and submit a memory-mapped RF-trace row extraction.""" - from .compute import _ack_applied_job - if not n_clicks: - return (no_update,) * 6 - - _ack_applied_job(applied_job) + return (no_update,) * 5 if not ref_group: return ( @@ -489,19 +504,18 @@ def compute_rf_trace( False, no_update, no_update, - no_update, ) if not stored_distmats: return ( dmc.Text("Please compute RF distances first (Distances tab).", c="red"), - no_update, False, no_update, no_update, no_update, + no_update, False, no_update, no_update, ) if not selected_matrix or selected_matrix not in stored_distmats: return ( dmc.Text("Please pick an RF matrix at the top of the page.", c="red"), - no_update, False, no_update, no_update, no_update, + no_update, False, no_update, no_update, ) try: @@ -519,7 +533,6 @@ def compute_rf_trace( False, no_update, no_update, - no_update, ) try: @@ -571,7 +584,6 @@ def compute_rf_trace( False, no_update, no_update, - no_update, ) spinner = dmc.Group( @@ -591,7 +603,6 @@ def compute_rf_trace( None, True, True, - False, job_ref.as_dict(), ) @@ -600,33 +611,33 @@ def compute_rf_trace( Output("rf-trace-store", "data", allow_duplicate=True), Output("notifications-container", "children", allow_duplicate=True), Output("export-rf-trace-button", "disabled", allow_duplicate=True), - Output("compute-rf-trace-button", "disabled", allow_duplicate=True), - Output("compute-poll-interval", "disabled", allow_duplicate=True), - Output("compute-applied-job-store", "data", allow_duplicate=True), - Input("compute-poll-interval", "n_intervals"), + Output( + {"type": "compute-terminal-receipt", "kind": "rf-trace"}, + "data", + ), + Input("compute-terminal-event-store", "data"), Input("rf-trace-job-store", "data"), prevent_initial_call=True, ) - def poll_rf_trace_completion(_n_intervals, job_data): - ref = _job_ref_from_store(job_data, expected_kind="rf_trace") - if ref is None: - return (no_update,) * 7 - snapshot = job_manager.snapshot_for_delivery(ref) - if ( - snapshot is None - or snapshot.acknowledged - or snapshot.terminal is None - ): - return (no_update,) * 7 - - terminal = snapshot.terminal - payload = terminal.payload + def render_rf_trace_terminal_event(terminal_event, job_data): + event = terminal_event_for_job( + terminal_event, + job_data, + expected_kind="rf_trace", + ) + if event is None: + return (no_update,) * 5 + + ref = JobRef.from_dict(event) + terminal_state = JobState(str(event["state"])) + payload = event["payload"] + first_delivery = int(event["delivery_attempt"]) == 1 notification = no_update store_data = no_update export_disabled = True - if terminal.state is JobState.CANCELLED: - if snapshot.delivery_attempt == 1: + if terminal_state is JobState.CANCELLED: + if first_delivery: add_log("RF Trace computation cancelled by user.", "WARNING") output = dmc.Alert( title="RF Trace computation cancelled", @@ -634,9 +645,9 @@ def poll_rf_trace_completion(_n_intervals, job_data): color="gray", variant="light", ) - elif terminal.state is JobState.FAILED: + elif terminal_state is JobState.FAILED: message = str(payload.get("message", "Unknown error")) - if snapshot.delivery_attempt == 1: + if first_delivery: add_log(f"RF Trace computation failed: {message}", "ERROR") output = dmc.Text( f"RF Trace computation failed: {message}", @@ -683,9 +694,7 @@ def poll_rf_trace_completion(_n_intervals, job_data): store_data, notification, export_disabled, - False, - True, - snapshot.terminal_delivery_marker(), + terminal_delivery_marker(event), ) # Re-render RF trace plot when burnin changes @@ -836,27 +845,23 @@ def render_ess_runs_table(selected_matrix): Output("compute-pseudo-ess-button", "disabled"), Input("diagnostics-distmat-select", "value"), Input({"type": "ess-run-checkbox", "index": ALL}, "checked"), + Input("compute-busy-store", "data"), ) - def toggle_compute_pseudo_ess_button(selected_matrix, checks): + def toggle_compute_pseudo_ess_button( + selected_matrix, + checks, + compute_busy, + ): # Disabled until a matrix is selected AND at least one run is checked. - if not selected_matrix: + if is_compute_busy(compute_busy) or not selected_matrix: return True if not checks or not any(checks): return True return False @callback( - # pseudo-ess-output.children + compute-pseudo-ess-button.disabled - # are both written from THIS click handler, from - # poll_pseudo_ess_completion (pseudo_ess_compute.py), and from - # the sidebar's Clear-data handler. Making the click handler - # ALSO use allow_duplicate=True means there's no "primary" - # for these Outputs — every writer is equal. This avoids the - # Dash 4.x output-dispatch quirk where a secondary write can - # be dropped if the primary hasn't fired in the same batch. Output("pseudo-ess-output", "children", allow_duplicate=True), Output("compute-pseudo-ess-button", "disabled", allow_duplicate=True), - Output("compute-poll-interval", "disabled", allow_duplicate=True), Output("pseudo-ess-job-store", "data"), Input("compute-pseudo-ess-button", "n_clicks"), State("diagnostics-distmat-select", "value"), @@ -864,7 +869,6 @@ def toggle_compute_pseudo_ess_button(selected_matrix, checks): State("ess-burnin-input", "value"), State({"type": "ess-run-checkbox", "index": ALL}, "checked"), State({"type": "ess-run-checkbox", "index": ALL}, "id"), - State("compute-applied-job-store", "data"), prevent_initial_call=True, ) def compute_pseudo_ess_for_runs( @@ -874,7 +878,6 @@ def compute_pseudo_ess_for_runs( burnin, checks, ids, - applied_job, ): """Submit a Pseudo-ESS job to the persistent worker. @@ -883,26 +886,18 @@ def compute_pseudo_ess_for_runs( worker needs. The actual eigendecomp-heavy ESS compute lives in the worker subprocess — see ``ess._subprocess_worker``. - Returns immediately with a spinner in ``pseudo-ess-output``, - the compute button disabled, and the shared poll interval - enabled so ``poll_pseudo_ess_completion`` will pick up the - worker's response. + Returns immediately with a spinner, a disabled originating button, + and an immutable job identity that wakes the central reconciler. """ from . import pseudo_ess_compute - from .compute import _ack_applied_job if not n_clicks or not selected_matrix: - return (no_update,) * 4 - - # The visible terminal result and this marker arrived in one previous - # browser response. A fast next click can precede the dedicated ack - # callback, so acknowledge it idempotently before requesting a new job. - _ack_applied_job(applied_job) + return (no_update,) * 3 ticked = [i["index"] for i, c in zip(ids, checks) if c] if not ticked: return (dmc.Text("No runs selected.", c="dimmed", size="sm"), - no_update, no_update, no_update) + no_update, no_update) try: # Only labels are needed to form per-run row indices. Avoid loading @@ -911,7 +906,7 @@ def compute_pseudo_ess_for_runs( except KeyError: return (dmc.Text(f"Matrix {selected_matrix!r} is no longer available.", c="red", size="sm"), - no_update, no_update, no_update) + no_update, no_update) # Bucket row indices by group prefix once (matrix-row order # matches MCMC iteration order within each chain). @@ -959,11 +954,11 @@ def compute_pseudo_ess_for_runs( return (dmc.Text( "Burn-in leaves fewer than 4 trees per run; nothing to compute.", c="dimmed", size="sm", - ), no_update, no_update, no_update) + ), no_update, no_update) - # Hand off to the subprocess. The poll callback in - # pseudo_ess_compute.py picks up the result and replaces the - # spinner with the result table. + # Hand off to the subprocess. The central reconciler publishes the + # terminal event; pseudo_ess_compute.py replaces the spinner with the + # matching result table. try: job_ref = pseudo_ess_compute.submit_pseudo_ess_job( distmat_path=str(state.get_distmat_file_path(selected_matrix)), @@ -987,7 +982,6 @@ def compute_pseudo_ess_for_runs( ), False, no_update, - no_update, ) spinner = dmc.Group([ @@ -999,5 +993,4 @@ def compute_pseudo_ess_for_runs( stop_button("ess"), ], gap="sm") - # Spinner, button disabled, polling enabled, and immutable job identity. - return spinner, True, False, job_ref.as_dict() + return spinner, True, job_ref.as_dict() diff --git a/src/treetracer/callbacks/job_reconcile.py b/src/treetracer/callbacks/job_reconcile.py new file mode 100644 index 0000000..a837391 --- /dev/null +++ b/src/treetracer/callbacks/job_reconcile.py @@ -0,0 +1,333 @@ +"""Single-owner reconciliation for every managed background computation. + +The persistent worker and :class:`~treetracer.background_jobs.JobManager` +serialize all heavy work, so the browser needs only one polling owner. This +module is that owner: + +* one 250 ms interval callback observes the active managed job; +* the same callback samples RF/MDS progress sidecars; +* terminal state is copied into one small, replayable browser event; +* job-specific presentation callbacks render that event and write a dedicated + receipt in the same response as their UI; +* one acknowledgement callback consumes those receipts. + +The interval remains enabled until a matching receipt has acknowledged the +sticky terminal event. A lost terminal-event response, presentation response, +or acknowledgement request therefore causes another delivery attempt instead +of a permanently stuck loading state. +""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any, Mapping + +from dash import ALL, Input, Output, State, callback, no_update + +from .._worker_log import log as _wlog +from ..background_jobs import ( + TERMINAL_STATES, + JobRef, + JobSnapshot, + JobState, + job_manager, +) + + +_MANAGED_KINDS = frozenset( + { + "rf", + "mds", + "pseudo_ess", + "consensus", + "rf_trace", + "clade_compare", + } +) + + +def job_ref_from_store(data: Any, *, expected_kind: str | None = None): + """Parse a browser job reference without trusting browser data.""" + if not isinstance(data, Mapping): + return None + try: + ref = JobRef.from_dict(data) + except (KeyError, TypeError, ValueError): + return None + if ref.kind not in _MANAGED_KINDS: + return None + if expected_kind is not None and ref.kind != expected_kind: + return None + return ref + + +def is_compute_busy(data: Any) -> bool: + """Return the shared browser-side compute gate.""" + return isinstance(data, Mapping) and data.get("busy") is True + + +def terminal_event_for_job( + event_data: Any, + job_data: Any, + *, + expected_kind: str, +): + """Return a terminal event only when it matches the current browser job. + + Both values are callback *Inputs*, not States. If a newer job reference + reaches the browser while an older terminal-render request is still in + flight, Dash schedules a newer execution and the generation mismatch below + makes the stale event a no-op. + """ + event_ref = job_ref_from_store(event_data, expected_kind=expected_kind) + job_ref = job_ref_from_store(job_data, expected_kind=expected_kind) + if event_ref is None or event_ref != job_ref: + return None + if not isinstance(event_data, Mapping): + return None + try: + state = JobState(str(event_data["state"])) + revision = int(event_data["terminal_revision"]) + delivery_attempt = int(event_data["delivery_attempt"]) + except (KeyError, TypeError, ValueError): + return None + if state not in TERMINAL_STATES or revision < 1 or delivery_attempt < 1: + return None + if not isinstance(event_data.get("payload"), Mapping): + return None + if not isinstance(event_data.get("metadata"), Mapping): + return None + return event_data + + +def terminal_delivery_marker(event_data: Mapping[str, Any]) -> dict[str, Any]: + """Build the receipt written atomically with feature terminal UI.""" + ref = JobRef.from_dict(event_data) + return { + **ref.as_dict(), + "terminal_revision": int(event_data["terminal_revision"]), + "delivery_attempt": int(event_data["delivery_attempt"]), + } + + +def _terminal_envelope(snapshot: JobSnapshot) -> dict[str, Any]: + """Serialize one small terminal snapshot for feature presentation.""" + terminal = snapshot.terminal + if terminal is None: + raise ValueError("cannot deliver a non-terminal job snapshot") + return { + **snapshot.ref.as_dict(), + "state": terminal.state.value, + "terminal_revision": terminal.revision, + "delivery_attempt": snapshot.delivery_attempt, + "payload": dict(terminal.payload), + "metadata": dict(snapshot.metadata), + } + + +def _busy_payload(ref: JobRef | None) -> dict[str, Any]: + if ref is None: + return {"busy": False} + return {"busy": True, **ref.as_dict()} + + +def _busy_update(ref: JobRef | None, current: Any): + desired = _busy_payload(ref) + return no_update if current == desired else desired + + +def _terminal_progress(snapshot: JobSnapshot) -> tuple[float, str]: + progress = snapshot.progress + fraction = 0.0 if progress is None else progress.fraction + label = { + JobState.SUCCEEDED: "complete", + JobState.FAILED: "failed", + JobState.CANCELLED: "cancelled", + }[snapshot.state] + return fraction * 100.0, label + + +def _read_rf_progress(ref: JobRef, progress_path: Any): + if not progress_path: + return no_update, no_update + try: + data = json.loads(Path(progress_path).read_text(encoding="utf-8")) + except (OSError, ValueError, TypeError): + return no_update, no_update + + value = int(data.get("value", 0)) + total = int(data.get("total", 0)) + fraction = max(0.0, min(float(data.get("fraction", 0.0)), 1.0)) + phase = str(data.get("phase", "computing")) + job_manager.update_progress(ref, fraction, phase) + if total == 0: + return 0, "starting…" + percentage = fraction * 100.0 + if phase == "finalizing": + label = f"{value:,} / {total:,} pairs — finalizing…" + else: + label = f"{value:,} / {total:,} pairs ({percentage:.1f}%)" + return percentage, label + + +def _read_mds_progress(ref: JobRef, progress_path: Any): + if not progress_path: + return no_update, no_update + try: + data = json.loads(Path(progress_path).read_text(encoding="utf-8")) + except (OSError, ValueError, TypeError): + return no_update, no_update + + fraction = max(0.0, min(float(data.get("fraction", 0.0)), 1.0)) + percentage = fraction * 100.0 + phase = str(data.get("phase", "computing")) + label = str(data.get("label") or phase.replace("_", " ")) + job_manager.update_progress(ref, fraction, phase, label) + return percentage, f"{label} ({percentage:.0f}%)" + + +def _matching_active_receipt(receipts: Any): + active = job_manager.active_ref() + if active is None: + return None, None + values = receipts if isinstance(receipts, list) else [receipts] + for value in values: + ref = job_ref_from_store(value) + if ref != active or not isinstance(value, Mapping): + continue + try: + revision = int(value["terminal_revision"]) + except (KeyError, TypeError, ValueError): + continue + return active, {**dict(value), "terminal_revision": revision} + return active, None + + +def register_job_reconciliation_callbacks(): + @callback( + Output("compute-poll-interval", "disabled"), + Output("compute-terminal-event-store", "data"), + Output("compute-busy-store", "data"), + Output("rf-progress-bar", "value"), + Output("rf-progress-label", "children"), + Output("mds-progress-bar", "value"), + Output("mds-progress-label", "children"), + Input("compute-poll-interval", "n_intervals"), + # Per-workflow stores wake this single owner immediately after submit. + Input("rf-job-store", "data"), + Input("mds-job-store", "data"), + Input("pseudo-ess-job-store", "data"), + Input("consensus-job-store", "data"), + Input("rf-trace-job-store", "data"), + Input("clade-freq-job-store", "data"), + State("compute-busy-store", "data"), + State("rf-progress-path", "data"), + State("mds-progress-path", "data"), + prevent_initial_call=True, + ) + def reconcile_compute_job( + n_intervals, + _rf_job_data, + _mds_job_data, + _pseudo_ess_job_data, + _consensus_job_data, + _rf_trace_job_data, + _clade_job_data, + current_busy, + rf_progress_path, + mds_progress_path, + ): + """Own polling, terminal delivery, progress, and the global gate.""" + active = job_manager.active_ref() + busy = _busy_update(active, current_busy) + if active is None: + return ( + True, + no_update, + busy, + no_update, + no_update, + no_update, + no_update, + ) + + snapshot = job_manager.snapshot(active) + if snapshot is None: + return ( + True, + no_update, + _busy_update(None, current_busy), + no_update, + no_update, + no_update, + no_update, + ) + + rf_progress = (no_update, no_update) + mds_progress = (no_update, no_update) + terminal_event = no_update + + if snapshot.terminal is not None: + delivery = job_manager.snapshot_for_delivery(active) + if delivery is None: + return ( + True, + no_update, + _busy_update(None, current_busy), + no_update, + no_update, + no_update, + no_update, + ) + terminal_event = _terminal_envelope(delivery) + if active.kind == "rf": + rf_progress = _terminal_progress(delivery) + elif active.kind == "mds": + mds_progress = _terminal_progress(delivery) + _wlog( + "[parent] reconcile_compute_job: delivering " + f"job={active.job_id}/{active.kind}/generation-" + f"{active.generation}, state={delivery.state.value}, " + f"attempt={delivery.delivery_attempt}" + ) + elif active.kind == "rf": + rf_progress = _read_rf_progress(active, rf_progress_path) + elif active.kind == "mds": + mds_progress = _read_mds_progress(active, mds_progress_path) + elif (n_intervals or 0) % 20 == 0: + _wlog( + "[parent] reconcile_compute_job: " + f"job={active.job_id}/{active.kind}/generation-" + f"{active.generation}, state={snapshot.state.value}" + ) + + # Polling remains enabled through terminal presentation. The receipt + # callback acknowledges server state; the next tick then takes the + # active=None branch and is the sole path that disables this interval. + return ( + False, + terminal_event, + busy, + *rf_progress, + *mds_progress, + ) + + @callback( + Output("compute-job-ack-store", "data"), + Input( + {"type": "compute-terminal-receipt", "kind": ALL}, + "data", + ), + prevent_initial_call=True, + ) + def acknowledge_terminal_receipt(receipts): + """Acknowledge only the receipt for the authoritative active job.""" + active, receipt = _matching_active_receipt(receipts) + if active is None or receipt is None: + return no_update + acknowledged = job_manager.acknowledge( + active, + receipt["terminal_revision"], + ) + return {**receipt, "acknowledged": acknowledged} diff --git a/src/treetracer/callbacks/pseudo_ess_compute.py b/src/treetracer/callbacks/pseudo_ess_compute.py index e9a7896..55a13b4 100644 --- a/src/treetracer/callbacks/pseudo_ess_compute.py +++ b/src/treetracer/callbacks/pseudo_ess_compute.py @@ -1,9 +1,9 @@ -"""Managed Pseudo-ESS dispatch, finalization, and terminal polling. +"""Managed Pseudo-ESS dispatch, finalization, and terminal presentation. The Diagnostics tab prepares small per-run index lists and submits one worker -request. ``JobManager`` owns the job identity and terminal state, so polling is -read-only and a completed result remains replayable until the browser confirms -that it applied the matching UI response. +request. ``JobManager`` owns the job identity and terminal state; the central +reconciler publishes a read-only event, and the completed result remains +replayable until the browser confirms that it applied the matching UI response. """ from __future__ import annotations @@ -18,7 +18,11 @@ from ..logger import add_log from .._worker_log import log as _wlog from . import persistent_worker -from .compute import _get_executor, _job_ref_from_store +from .compute import _get_executor +from .job_reconcile import ( + terminal_delivery_marker, + terminal_event_for_job, +) def reset() -> None: @@ -158,34 +162,28 @@ def _build_result_table(results: list[dict[str, Any]]): def register_pseudo_ess_compute_callbacks(): @callback( Output("pseudo-ess-output", "children", allow_duplicate=True), - Output("compute-pseudo-ess-button", "disabled", allow_duplicate=True), - Output("compute-poll-interval", "disabled", allow_duplicate=True), - Output("compute-applied-job-store", "data", allow_duplicate=True), - Input("compute-poll-interval", "n_intervals"), + Output( + {"type": "compute-terminal-receipt", "kind": "pseudo-ess"}, + "data", + ), + Input("compute-terminal-event-store", "data"), Input("pseudo-ess-job-store", "data"), prevent_initial_call=True, ) - def poll_pseudo_ess_completion(n_intervals, job_data): - ref = _job_ref_from_store(job_data, expected_kind="pseudo_ess") - if ref is None: - return (no_update,) * 4 - - snapshot = job_manager.snapshot_for_delivery(ref) - if snapshot is None or snapshot.acknowledged: - return (no_update,) * 4 - if snapshot.terminal is None: - if (n_intervals or 0) % 10 == 0: - _wlog( - f"[parent] poll_pseudo_ess_completion tick={n_intervals}: " - f"job={ref.job_id}/generation-{ref.generation}, " - f"state={snapshot.state.value}" - ) - return (no_update,) * 4 - - terminal = snapshot.terminal - applied = snapshot.terminal_delivery_marker() - if terminal.state is JobState.CANCELLED: - if snapshot.delivery_attempt == 1: + def render_pseudo_ess_terminal_event(terminal_event, job_data): + event = terminal_event_for_job( + terminal_event, + job_data, + expected_kind="pseudo_ess", + ) + if event is None: + return (no_update,) * 2 + + state = JobState(str(event["state"])) + payload = event["payload"] + first_delivery = int(event["delivery_attempt"]) == 1 + if state is JobState.CANCELLED: + if first_delivery: add_log("Pseudo-ESS computation cancelled by user.", "WARNING") output = dmc.Alert( title="Pseudo-ESS computation cancelled", @@ -193,20 +191,16 @@ def poll_pseudo_ess_completion(n_intervals, job_data): color="gray", variant="light", ) - elif terminal.state is JobState.FAILED: + elif state is JobState.FAILED: msg = ( "Pseudo-ESS computation failed: " - f"{terminal.payload.get('message', 'Unknown error')}" + f"{payload.get('message', 'Unknown error')}" ) - if snapshot.delivery_attempt == 1: + if first_delivery: add_log(msg, "ERROR") output = dmc.Text(msg, c="red", size="sm") else: - rows = list(terminal.payload.get("results", [])) + rows = list(payload.get("results", [])) output = _build_result_table(rows) - # The terminal event remains server-side until the applied marker is - # processed by the shared acknowledgement callback. If this whole Dash - # response is lost, the still-enabled interval requests the same event - # again with a new delivery-attempt marker. - return output, False, True, applied + return output, terminal_delivery_marker(event) diff --git a/src/treetracer/callbacks/sidebar.py b/src/treetracer/callbacks/sidebar.py index 0fccdcd..8e0c5db 100644 --- a/src/treetracer/callbacks/sidebar.py +++ b/src/treetracer/callbacks/sidebar.py @@ -640,21 +640,17 @@ def remove_file(n_clicks_list, stored_summaries): Output("clade-freq-consensus-tree-select-1", "value", allow_duplicate=True), Output("clade-freq-consensus-tree-select-2", "value", allow_duplicate=True), Output("clade-freq-output-paper", "style", allow_duplicate=True), - # Managed-compute lifecycle state. Clear every browser identity together - # with the server-side job record so no terminal replay can resurrect - # data after the reset. + # Clear every browser job identity with the server-side record. Those + # store changes wake the central reconciler, which alone settles its + # interval and busy state after reset. Output("rf-job-store", "data", allow_duplicate=True), Output("mds-job-store", "data", allow_duplicate=True), Output("pseudo-ess-job-store", "data", allow_duplicate=True), Output("consensus-job-store", "data", allow_duplicate=True), Output("rf-trace-job-store", "data", allow_duplicate=True), Output("clade-freq-job-store", "data", allow_duplicate=True), - Output("compute-applied-job-store", "data", allow_duplicate=True), - Output("compute-job-ack-store", "data", allow_duplicate=True), Output("rf-progress-path", "data", allow_duplicate=True), Output("mds-progress-path", "data", allow_duplicate=True), - Output("compute-poll-interval", "disabled", allow_duplicate=True), - Output("consensus-tree-poll-interval", "disabled", allow_duplicate=True), Output("treespace-loading-overlay", "visible", allow_duplicate=True), Output("within-run-loading-overlay", "visible", allow_duplicate=True), Input("clear-data-button", "n_clicks"), @@ -664,8 +660,8 @@ def clear_uploads(n_clicks): if n_clicks: add_log("Data cleared") # Cancel any in-flight subprocess job — descriptors point - # into the DB we're about to wipe, and we don't want the - # poll callbacks to write results based on stale state. + # into the DB we're about to wipe, and we don't want a late + # finalizer to publish results into freshly cleared state. from . import compute resetters = (("managed", compute.reset),) for label, resetter in resetters: @@ -755,13 +751,9 @@ def clear_uploads(n_clicks): None, # consensus-job-store None, # rf-trace-job-store None, # clade-freq-job-store - None, # compute-applied-job-store - None, # compute-job-ack-store None, # rf-progress-path None, # mds-progress-path - True, # compute-poll-interval disabled - True, # consensus-tree-poll-interval disabled False, # treespace-loading-overlay visible False, # within-run-loading-overlay visible ) - return (no_update,) * 48 + return (no_update,) * 44 diff --git a/src/treetracer/callbacks/treespace.py b/src/treetracer/callbacks/treespace.py index ea1d318..83293b0 100644 --- a/src/treetracer/callbacks/treespace.py +++ b/src/treetracer/callbacks/treespace.py @@ -15,6 +15,7 @@ _substitute_newick_labels, _build_canonical_remaps, ) +from .job_reconcile import is_compute_busy # Matches the `tree NAME` token at the start of a NEXUS tree line. @@ -495,15 +496,16 @@ def update_plot_theme(_, current_fig): Output("treespace-export-trees", "disabled"), Output("treespace-view-consensus-tree", "disabled"), Input("treespace-selected-trees-store", "data"), + Input("compute-busy-store", "data"), ) - def update_selection_info(selected): + def update_selection_info(selected, compute_busy): if not selected: return html.Div(), True, True return ( dmc.Badge(f"Selected: {len(selected)} trees", color="red", variant="light", size="sm"), False, - False, + is_compute_busy(compute_busy), ) # ------ selection store change → patch only the last 4 traces ------ @@ -760,16 +762,12 @@ def export_selected_trees(n_clicks, selected_pairs, plot_config): # Validates input, builds the matched-record list + consensus-tree-coord # lookup, then hands off to the persistent worker via # ``consensus_tree_compute.submit_consensus_tree_job``. Completion is handled by - # ``consensus_tree_compute.poll_consensus_tree_completion`` which fans the result back to + # the consensus terminal presentation adapter, which fans the result back to # this tab's view-consensus-tree-store, dismisses the loading overlay, and # re-enables the button. @callback( Output("treespace-loading-overlay", "visible", allow_duplicate=True), Output("treespace-view-consensus-tree", "disabled", allow_duplicate=True), - # consensus tree polling uses its own interval (see navbar.py) so this - # handler and poll_consensus_tree_completion don't collide with the RF/MDS - # poll on a shared allow_duplicate output. - Output("consensus-tree-poll-interval", "disabled", allow_duplicate=True), Output("notifications-container", "children", allow_duplicate=True), Output("consensus-job-store", "data", allow_duplicate=True), Input("treespace-view-consensus-tree", "n_clicks"), @@ -777,24 +775,20 @@ def export_selected_trees(n_clicks, selected_pairs, plot_config): State("plot-config-store", "data"), State("treespace-result-select", "value"), State("mds-result-store", "data"), - State("compute-applied-job-store", "data"), prevent_initial_call=True, ) def view_consensus_tree(n_clicks, selected_pairs, plot_config, - selected_key, results, applied_job): + selected_key, results): from ..background_jobs import JobBusyError from ..logger import notif_id from ..db.tree_service import get_tree_service from . import consensus_tree_compute - from .compute import _ack_applied_job if not n_clicks or not selected_pairs or not plot_config: - return (no_update,) * 5 - - _ack_applied_job(applied_job) + return (no_update,) * 4 def _err(msg, autoclose=5000): - return (False, False, no_update, dmc.Notification( + return (False, False, dmc.Notification( title="Consensus tree Error", message=msg, color="red", action="show", autoClose=autoclose, id=notif_id(), @@ -834,7 +828,7 @@ def _err(msg, autoclose=5000): rec["line_offset"] = int(rec["line_offset"]) rec["line_length"] = int(rec["line_length"]) - # (group, treenum) per tree-name so the poll callback can put the + # (group, treenum) per tree-name so terminal presentation can put the # green ring on the consensus tree's dot. consensus_tree_coord_by_tree_name = { row["tree"]: (row["group"], int(row["treenum"])) @@ -859,18 +853,18 @@ def _err(msg, autoclose=5000): "is still finishing. Please wait for it to complete." ) - # Return: overlay on, button disabled, polling enabled, no - # notification yet (notification fires when compute finishes). - return True, True, False, no_update, job_ref.as_dict() + # The immutable job store wakes the one global reconciler. + return True, True, no_update, job_ref.as_dict() # The per-tab clientside ``window.open`` that used to live here is # gone — it was a duplicate of the one in within_run.py and the # one in consensus_tree_list.py. They all now route through # ``consensus-tree-peartree-open-store`` and the single clientside callback # in ``callbacks/rename_consensus_tree.py`` does the actual ``window.open``. - # The poll callback in ``consensus_tree_compute.py`` still writes the freshly - # registered consensus tree's ``{uuid, name}`` to ``treespace-view-consensus-tree-store`` - # — it's now picked up by ``forward_compute_to_modal`` in + # The terminal adapter in ``consensus_tree_compute.py`` writes the freshly + # registered tree's ``{uuid, name}`` to + # ``treespace-view-consensus-tree-store``. It is picked up by + # ``forward_compute_to_modal`` in # ``rename_consensus_tree.py``, which opens the rename modal (with # ``after='view'``) so the user can confirm or edit the auto-name # before peartree opens on Save. diff --git a/src/treetracer/callbacks/within_run.py b/src/treetracer/callbacks/within_run.py index facb792..7cf82aa 100644 --- a/src/treetracer/callbacks/within_run.py +++ b/src/treetracer/callbacks/within_run.py @@ -30,6 +30,7 @@ from ..theme import get_template from ..plot_utils import retheme_figure +from .job_reconcile import is_compute_busy TREETRACER_BLUE = "#228be6" @@ -701,15 +702,16 @@ def update_plot_theme(_, current_fig): Output("within-run-export-trees", "disabled"), Output("within-run-view-consensus-tree", "disabled"), Input("within-run-selected-trees-store", "data"), + Input("compute-busy-store", "data"), ) - def update_selection_info(selected): + def update_selection_info(selected, compute_busy): if not selected: return html.Div(), True, True return ( dmc.Badge(f"Selected: {len(selected)} trees", color="red", variant="light", size="sm"), False, - False, + is_compute_busy(compute_busy), ) # ------ export selected trees ------ @@ -779,13 +781,11 @@ def export_selected_trees(n_clicks, selected_treenums, selected_key, selected_ru # ------ View consensus tree — thin submit handler ------ # Mirrors the Between-run shape — see callbacks/consensus_tree_compute.py for - # the shared dispatch + polling code, and callbacks/treespace.py + # shared dispatch and terminal presentation, and callbacks/treespace.py # for the parallel implementation. @callback( Output("within-run-loading-overlay", "visible", allow_duplicate=True), Output("within-run-view-consensus-tree", "disabled", allow_duplicate=True), - # consensus tree polling uses its own interval (see navbar.py). - Output("consensus-tree-poll-interval", "disabled", allow_duplicate=True), Output("notifications-container", "children", allow_duplicate=True), Output("consensus-job-store", "data", allow_duplicate=True), Input("within-run-view-consensus-tree", "n_clicks"), @@ -793,7 +793,6 @@ def export_selected_trees(n_clicks, selected_treenums, selected_key, selected_ru State("within-run-result-select", "value"), State("within-run-run-select", "value"), State("mds-result-store", "data"), - State("compute-applied-job-store", "data"), prevent_initial_call=True, ) def view_consensus_tree( @@ -802,21 +801,17 @@ def view_consensus_tree( selected_key, selected_run, results, - applied_job, ): from ..background_jobs import JobBusyError from ..logger import notif_id from ..db.tree_service import get_tree_service from . import consensus_tree_compute - from .compute import _ack_applied_job if not n_clicks or not selected_treenums: - return (no_update,) * 5 - - _ack_applied_job(applied_job) + return (no_update,) * 4 def _err(msg, autoclose=5000): - return (False, False, no_update, dmc.Notification( + return (False, False, dmc.Notification( title="Consensus tree Error", message=msg, color="red", action="show", autoClose=autoclose, id=notif_id(), @@ -824,7 +819,7 @@ def _err(msg, autoclose=5000): mds_result = _get_active_result(selected_key, results) if not mds_result or not selected_run: - return (no_update,) * 5 + return (no_update,) * 4 source_distmat = (mds_result.get("metadata") or {}).get("source_distmat") if not source_distmat: @@ -883,7 +878,7 @@ def _err(msg, autoclose=5000): "is still finishing. Please wait for it to complete." ) - return True, True, False, no_update, job_ref.as_dict() + return True, True, no_update, job_ref.as_dict() # The per-tab clientside ``window.open`` that used to live here is # gone; see the parallel note in ``treespace.py``. The diff --git a/src/treetracer/rf/_subprocess_worker.py b/src/treetracer/rf/_subprocess_worker.py index 9c3fb1c..a6d7799 100644 --- a/src/treetracer/rf/_subprocess_worker.py +++ b/src/treetracer/rf/_subprocess_worker.py @@ -63,7 +63,8 @@ def compute_rf_worker_entry( Defaults to True for backward compatibility with the previous always-rooted behaviour. - Returns a dict consumed by ``poll_completion`` in callbacks/compute.py. + Returns a dict finalized by the managed job wrapper; the central + reconciler later delivers only its small terminal payload to the browser. """ t0 = time.time() @@ -138,8 +139,8 @@ def compute_rf_worker_entry( # When ``progress_path`` is set, share a ``ProgressCounter`` with the # rayon workers via the new rapidtrees 0.6 API, and run a daemon # thread that mirrors the counter state into a small JSON sidecar - # file. The parent's Dash poll callback reads this file every ~100ms - # to drive a progress bar — no IPC changes needed. + # file. The parent's central reconciler samples it every 250ms to drive a + # progress bar — no IPC changes needed. counter = None stop_event: Optional[threading.Event] = None writer: Optional[threading.Thread] = None diff --git a/src/treetracer/ui/navbar.py b/src/treetracer/ui/navbar.py index f466a47..57db801 100644 --- a/src/treetracer/ui/navbar.py +++ b/src/treetracer/ui/navbar.py @@ -84,32 +84,79 @@ def add_navbar(): # scatter-click state, kept server-side-friendly. dcc.Store(id="clade-freq-data-store", storage_type="memory"), dcc.Store(id="clade-freq-click-store", storage_type="memory"), - # Background computation polling - dcc.Interval(id="compute-poll-interval", interval=100, disabled=True), - # Per-workflow job identities and shared two-phase terminal - # delivery. The poll that renders a terminal result writes - # an applied marker; only then does the acknowledgement - # callback release the server-side sticky event. + # One shared reconciliation cadence for every serialized + # background computation. It is enabled only while a + # managed job is awaiting completion or UI receipt. + dcc.Interval( + id="compute-poll-interval", + interval=250, + disabled=True, + ), + # Per-workflow immutable job identities wake the central + # reconciler immediately after a successful submission. dcc.Store(id="rf-job-store", storage_type="memory"), dcc.Store(id="mds-job-store", storage_type="memory"), dcc.Store(id="pseudo-ess-job-store", storage_type="memory"), dcc.Store(id="consensus-job-store", storage_type="memory"), dcc.Store(id="rf-trace-job-store", storage_type="memory"), dcc.Store(id="clade-freq-job-store", storage_type="memory"), - dcc.Store(id="compute-applied-job-store", storage_type="memory"), + # The reconciler publishes one generic terminal envelope. + # Feature adapters render it only when its generation + # matches their current job store, then atomically write a + # dedicated receipt with the terminal UI. + dcc.Store( + id="compute-terminal-event-store", + storage_type="memory", + ), + dcc.Store( + id="compute-busy-store", + storage_type="memory", + data={"busy": False}, + ), dcc.Store(id="compute-job-ack-store", storage_type="memory"), + dcc.Store( + id={ + "type": "compute-terminal-receipt", + "kind": "rf-mds", + }, + storage_type="memory", + ), + dcc.Store( + id={ + "type": "compute-terminal-receipt", + "kind": "pseudo-ess", + }, + storage_type="memory", + ), + dcc.Store( + id={ + "type": "compute-terminal-receipt", + "kind": "consensus", + }, + storage_type="memory", + ), + dcc.Store( + id={ + "type": "compute-terminal-receipt", + "kind": "rf-trace", + }, + storage_type="memory", + ), + dcc.Store( + id={ + "type": "compute-terminal-receipt", + "kind": "clade-compare", + }, + storage_type="memory", + ), # Resolves scatter split IDs through the matching server-side # managed comparison result; avoids shipping tip sets through # browser JSON or decoding the full snapshot on click. dcc.Store(id="clade-freq-result-key-store", storage_type="memory"), - # Consensus trees keep a dedicated cadence because their - # overlay and button lifecycle can stop independently of - # the shared RF/MDS/Pseudo-ESS interval. - dcc.Interval(id="consensus-tree-poll-interval", interval=100, disabled=True), # Path to the RF worker's sidecar progress file # (``.progress``). Set by # ``handle_compute_rf`` when an RF compute starts; - # consumed by ``update_rf_progress`` to drive the + # consumed by the central job reconciler to drive the # progress bar inside the computing banner. dcc.Store(id="rf-progress-path", storage_type="memory"), # Same sidecar-progress pattern for MDS/PCoA. This diff --git a/src/treetracer/ui/panels/treespace.py b/src/treetracer/ui/panels/treespace.py index 204e04d..12ce164 100644 --- a/src/treetracer/ui/panels/treespace.py +++ b/src/treetracer/ui/panels/treespace.py @@ -15,7 +15,7 @@ def _add_treespace_panel(): return html.Div([ # consensus-tree-compute loading overlay. Visible=True flipped on by # ``view_consensus_tree`` (click handler), back to False by - # ``consensus_tree_compute.poll_consensus_tree_completion``. Position relative on the + # consensus terminal presentation adapter. Position relative on the # wrapping Div lets the overlay sit on top. dmc.LoadingOverlay( id="treespace-loading-overlay", diff --git a/src/treetracer/ui/panels/within_run.py b/src/treetracer/ui/panels/within_run.py index b5395a0..13d6bba 100644 --- a/src/treetracer/ui/panels/within_run.py +++ b/src/treetracer/ui/panels/within_run.py @@ -14,7 +14,7 @@ def _add_within_run_panel(): return html.Div([ # consensus-tree-compute loading overlay — see treespace panel for the # full pattern. Toggled by ``view_consensus_tree`` (on) and - # ``consensus_tree_compute.poll_consensus_tree_completion`` (off). + # the consensus terminal presentation adapter (off). dmc.LoadingOverlay( id="within-run-loading-overlay", visible=False, From 116a3b5b29c7e0c7fef191c6b1558758976941a7 Mon Sep 17 00:00:00 2001 From: Sam Hong <20244918+hongsamL@users.noreply.github.com> Date: Mon, 3 Aug 2026 18:30:24 +0200 Subject: [PATCH 7/9] worker protocol and tests --- src/test/README.md | 7 + src/test/test_app_smoke.py | 42 ++ src/test/test_compute_job_lifecycle.py | 12 +- src/test/test_persistent_worker_watchdog.py | 221 +++++++ src/test/test_terminal_delivery_faults.py | 116 ++++ src/test/test_worker_protocol.py | 80 +++ src/treetracer/__init__.py | 100 ++- src/treetracer/callbacks/job_reconcile.py | 98 ++- src/treetracer/callbacks/persistent_worker.py | 613 ++++++++++++++---- src/treetracer/ui/widgets.py | 17 +- src/treetracer/worker_protocol.py | 173 +++++ 11 files changed, 1275 insertions(+), 204 deletions(-) create mode 100644 src/test/test_persistent_worker_watchdog.py create mode 100644 src/test/test_terminal_delivery_faults.py create mode 100644 src/test/test_worker_protocol.py create mode 100644 src/treetracer/worker_protocol.py diff --git a/src/test/README.md b/src/test/README.md index 19200bf..f9ff8b1 100644 --- a/src/test/README.md +++ b/src/test/README.md @@ -24,6 +24,13 @@ bash run_tests.sh -k consensus tree -v | `conftest.py` | Session fixtures: NEXUS parse, rapidtrees presence, DendroPy parse. | | `test.trees` | 100-tree BEAST fixture (5 MB). The CI integration anchor. | | `test_app_smoke.py` | App imports, callback registration, figure-layout invariants. | +| `test_background_jobs.py` | Thread-safe job-state transitions, exactly-once finalization, sticky terminal delivery, acknowledgement, cancellation, and reset. | +| `test_compute_job_lifecycle.py` | RF/MDS lifecycle integration, generation checks, wildcard dynamic progress outputs, and terminal replay. | +| `test_managed_compute_jobs.py` | Pseudo-ESS and consensus managed-job finalization and replay behavior. | +| `test_managed_analysis_jobs.py` | RF Trace and clade-comparison worker/finalizer/cache behavior. | +| `test_worker_protocol.py` | Framed socket messages, EOF/truncation handling, and heartbeat/result ordering. | +| `test_persistent_worker_watchdog.py` | Missing-heartbeat and hard-runtime bounds, protocol validation, configuration, and worker replacement. | +| `test_terminal_delivery_faults.py` | Deterministic terminal-event, terminal-UI/receipt, acknowledgement-loss, and stale-generation scenarios. | | `test_ess.py` | `effective_sample_size` vs AR(1) closed form + arviz cross-check (iid, AR(2), MA(5), heavy-tail, multimodal). | | `test_pseudo_ess.py` | `compute_pseudo_ess` shape, n-cap, rank-norm bound, row-order sensitivity. | | `test_pcoa.py` | `compute_mds` Procrustes-equivalent to scipy on synthetic Euclidean + real RF. | diff --git a/src/test/test_app_smoke.py b/src/test/test_app_smoke.py index 1668ec7..b57d2f9 100644 --- a/src/test/test_app_smoke.py +++ b/src/test/test_app_smoke.py @@ -57,6 +57,48 @@ def test_compute_interval_has_one_reconciliation_owner(): assert owners == {"reconcile_compute_job"} +def test_reconciler_uses_wildcards_for_dynamic_progress_banners(): + """RF and MDS banners never coexist, so concrete Outputs are unsafe. + + Dash rejects the whole reconciler response when a concrete output names + the progress component belonging to the other, currently-unmounted banner. + """ + from dash import _callback + + reconciler_outputs = None + for callback_data in _callback.GLOBAL_CALLBACK_MAP.values(): + callback_fn = callback_data.get("callback") + callback_fn = getattr(callback_fn, "__wrapped__", callback_fn) + if getattr(callback_fn, "__name__", "") != "reconcile_compute_job": + continue + output = callback_data.get("output") + reconciler_outputs = output if isinstance(output, list) else [output] + break + + assert reconciler_outputs is not None + component_ids = [item.component_id for item in reconciler_outputs] + concrete_progress_ids = { + "rf-progress-bar", + "rf-progress-label", + "mds-progress-bar", + "mds-progress-label", + } + assert not any( + isinstance(component_id, str) + and component_id in concrete_progress_ids + for component_id in component_ids + ) + pattern_types = { + component_id.get("type") + for component_id in component_ids + if isinstance(component_id, dict) + } + assert { + "compute-progress-bar", + "compute-progress-label", + } <= pattern_types + + def test_every_compute_action_reads_the_shared_busy_gate(): """All entry points must become unavailable while the worker is owned.""" from dash import _callback diff --git a/src/test/test_compute_job_lifecycle.py b/src/test/test_compute_job_lifecycle.py index 769803e..048a4b6 100644 --- a/src/test/test_compute_job_lifecycle.py +++ b/src/test/test_compute_job_lifecycle.py @@ -149,6 +149,8 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): "acknowledge_terminal_receipt", job_reconcile.register_job_reconciliation_callbacks, ) + rf_bar_ids = [{"type": "compute-progress-bar", "which": "rf"}] + rf_label_ids = [{"type": "compute-progress-label", "which": "rf"}] first_reconcile = reconcile( 10, @@ -161,6 +163,8 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): {"busy": False}, None, None, + rf_bar_ids, + rf_label_ids, ) first = render(first_reconcile[1], ref.as_dict(), None, False) second_reconcile = reconcile( @@ -174,13 +178,15 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): first_reconcile[2], None, None, + rf_bar_ids, + rf_label_ids, ) second = render(second_reconcile[1], ref.as_dict(), None, False) - assert len(first_reconcile) == 7 + assert len(first_reconcile) == 5 assert first_reconcile[0] is False assert first_reconcile[2]["busy"] is True - assert first_reconcile[3:5] == (100.0, "complete") + assert first_reconcile[3:] == ([100.0], ["complete"]) assert len(first) == 12 assert first[1] == expected_index assert first[2] is False @@ -206,6 +212,8 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): first_reconcile[2], None, None, + [], + [], ) assert settled[0] is True assert settled[2] == {"busy": False} diff --git a/src/test/test_persistent_worker_watchdog.py b/src/test/test_persistent_worker_watchdog.py new file mode 100644 index 0000000..bb68903 --- /dev/null +++ b/src/test/test_persistent_worker_watchdog.py @@ -0,0 +1,221 @@ +"""Watchdog behavior without launching the scientific worker process.""" + +from __future__ import annotations + +import socket +import threading +import time + +import pytest + +from treetracer.callbacks import persistent_worker +from treetracer.worker_protocol import ( + WorkerProtocolError, + receive_message, + send_message, +) + + +def test_result_wait_accepts_heartbeats_until_success(): + parent, worker = socket.socketpair() + + def respond(): + send_message(worker, { + "type": "heartbeat", + "job": "compute_rf", + "sequence": 1, + "elapsed_s": 0.0, + }) + send_message(worker, { + "type": "heartbeat", + "job": "compute_rf", + "sequence": 2, + "elapsed_s": 0.01, + }) + send_message(worker, { + "type": "result", + "ok": True, + "result": {"done": True}, + }) + + thread = threading.Thread(target=respond) + thread.start() + try: + result = persistent_worker._receive_worker_result( + parent, + "compute_rf", + persistent_worker.WatchdogPolicy(0.2, 1.0), + started_at=time.monotonic(), + ) + assert result["result"] == {"done": True} + finally: + thread.join(timeout=1) + parent.close() + worker.close() + + +def test_missing_heartbeat_has_a_bounded_wait(): + parent, worker = socket.socketpair() + try: + with pytest.raises( + persistent_worker.WorkerHeartbeatTimeout, + match="no heartbeat or result", + ): + persistent_worker._receive_worker_result( + parent, + "compute_mds", + persistent_worker.WatchdogPolicy(0.03, 1.0), + started_at=time.monotonic(), + ) + finally: + parent.close() + worker.close() + + +def test_maximum_runtime_wins_even_while_heartbeats_continue(): + parent, worker = socket.socketpair() + stop = threading.Event() + + def heartbeat_forever(): + sequence = 0 + while not stop.wait(0.005): + sequence += 1 + try: + send_message(worker, { + "type": "heartbeat", + "job": "compute_pseudo_ess", + "sequence": sequence, + "elapsed_s": sequence * 0.005, + }) + except OSError: + return + + thread = threading.Thread(target=heartbeat_forever) + thread.start() + try: + with pytest.raises( + persistent_worker.WorkerRuntimeExceeded, + match="maximum runtime", + ): + persistent_worker._receive_worker_result( + parent, + "compute_pseudo_ess", + persistent_worker.WatchdogPolicy(0.05, 0.03), + started_at=time.monotonic(), + ) + finally: + stop.set() + parent.close() + worker.close() + thread.join(timeout=1) + + +def test_heartbeat_must_match_job_and_advance_sequence(): + parent, worker = socket.socketpair() + try: + send_message(worker, { + "type": "heartbeat", + "job": "compute_rf_trace", + "sequence": 1, + }) + with pytest.raises(WorkerProtocolError, match="job mismatch"): + persistent_worker._receive_worker_result( + parent, + "compute_rf", + persistent_worker.WatchdogPolicy(0.2, 1.0), + started_at=time.monotonic(), + ) + finally: + parent.close() + worker.close() + + +def test_watchdog_policy_supports_safe_environment_overrides(monkeypatch): + monkeypatch.setenv("TREETRACER_WORKER_HEARTBEAT_INTERVAL_S", "0.01") + monkeypatch.setenv("TREETRACER_WORKER_HEARTBEAT_TIMEOUT_S", "0.02") + monkeypatch.setenv("TREETRACER_WORKER_MAX_RUNTIME_S", "100") + monkeypatch.setenv( + "TREETRACER_WORKER_MAX_RUNTIME_COMPUTE_RF_S", + "50", + ) + + policy = persistent_worker.watchdog_policy("compute_rf") + # Timeout is raised to three heartbeat periods to tolerate jitter. + assert policy.heartbeat_timeout_s == pytest.approx(0.03) + assert policy.max_runtime_s == 50 + + disabled = persistent_worker.watchdog_policy( + "compute_rf", + max_runtime_s=0, + ) + assert disabled.max_runtime_s is None + + +class _FakeProcess: + def __init__(self, pid): + self.pid = pid + self.returncode = None + self.killed = False + + def poll(self): + return -9 if self.killed else None + + def kill(self): + self.killed = True + self.returncode = -9 + + def wait(self, timeout=None): + self.killed = True + self.returncode = -9 + return self.returncode + + +def test_submit_timeout_restarts_worker_before_error_reaches_manager( + monkeypatch, +): + parent, worker = socket.socketpair() + replacement_parent, replacement_worker = socket.socketpair() + old_proc = _FakeProcess(101) + new_proc = _FakeProcess(202) + + monkeypatch.setattr(persistent_worker, "_worker_proc", old_proc) + monkeypatch.setattr(persistent_worker, "_worker_sock", parent) + monkeypatch.setattr(persistent_worker, "_current_job", None) + monkeypatch.setattr(persistent_worker, "_cancelled", False) + monkeypatch.setenv("TREETRACER_WORKER_HEARTBEAT_INTERVAL_S", "0.01") + monkeypatch.setenv("TREETRACER_WORKER_HEARTBEAT_TIMEOUT_S", "0.03") + + def fake_spawn(): + persistent_worker._worker_sock = replacement_parent + return new_proc + + monkeypatch.setattr(persistent_worker, "_spawn_worker", fake_spawn) + + # Consume the request but deliberately send no heartbeat or result. + received = {} + + def consume_request(): + received.update(receive_message(worker)) + + thread = threading.Thread(target=consume_request) + thread.start() + try: + with pytest.raises( + persistent_worker.WorkerHeartbeatTimeout, + match="worker was restarted", + ): + persistent_worker.submit_job( + "compute_rf_trace", + max_runtime_s=1, + matrix_path="ignored.npy", + ) + thread.join(timeout=1) + assert received["job"] == "compute_rf_trace" + assert old_proc.killed is True + assert persistent_worker._worker_proc is new_proc + assert persistent_worker._worker_sock is replacement_parent + finally: + parent.close() + worker.close() + replacement_parent.close() + replacement_worker.close() diff --git a/src/test/test_terminal_delivery_faults.py b/src/test/test_terminal_delivery_faults.py new file mode 100644 index 0000000..91ee0a8 --- /dev/null +++ b/src/test/test_terminal_delivery_faults.py @@ -0,0 +1,116 @@ +"""Protocol-level response-loss scenarios for Stage 6. + +These tests deliberately discard values at each browser-delivery boundary. +They complement the manual real-Dash fault harness in +``tools/stage6_browser_fault_harness.js``. +""" + +from __future__ import annotations + +import time +from concurrent.futures import ThreadPoolExecutor + +import pytest + +from treetracer.background_jobs import JobManager +from treetracer.callbacks.job_reconcile import ( + _terminal_envelope, + terminal_delivery_marker, + terminal_event_for_job, +) + + +def _wait_for_terminal(manager, ref, timeout=1.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + snapshot = manager.snapshot(ref) + if snapshot is not None and snapshot.terminal is not None: + return snapshot + time.sleep(0.002) + pytest.fail("managed job did not become terminal") + + +def test_dropped_event_and_ui_responses_both_replay_before_receipt(): + manager = JobManager(id_factory=lambda: "fault-job") + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit( + executor, + "rf", + lambda: "raw", + finalizer=lambda _ref, _raw: {"result_ref": "RF_001"}, + ) + terminal = _wait_for_terminal(manager, ref) + + # Attempt 1: the coordinator's response is dropped before the generic + # terminal Store changes. No browser adapter runs and there is no receipt. + dropped_event = _terminal_envelope( + manager.snapshot_for_delivery(ref) + ) + assert manager.active_ref() == ref + + # Attempt 2: the generic event lands, but the feature UI response is + # dropped. Building a receipt server-side is not acknowledgement; it must + # reach the browser and trigger the acknowledgement request. + dropped_ui = _terminal_envelope(manager.snapshot_for_delivery(ref)) + assert terminal_event_for_job( + dropped_ui, + ref.as_dict(), + expected_kind="rf", + ) == dropped_ui + _discarded_receipt = terminal_delivery_marker(dropped_ui) + assert manager.snapshot(ref).acknowledged is False + assert manager.active_ref() == ref + + # Attempt 3 lands fully. Semantic terminal state/revision stayed stable; + # only the retry counter advanced. + delivered = _terminal_envelope(manager.snapshot_for_delivery(ref)) + assert delivered["payload"] == dropped_event["payload"] + assert delivered["terminal_revision"] == dropped_event["terminal_revision"] + assert delivered["delivery_attempt"] == 3 + receipt = terminal_delivery_marker(delivered) + assert manager.acknowledge(ref, receipt["terminal_revision"]) + assert manager.active_ref() is None + assert terminal.terminal.revision == receipt["terminal_revision"] + + +def test_dropped_ack_response_is_safe_after_server_acknowledgement(): + manager = JobManager(id_factory=lambda: "ack-fault-job") + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit(executor, "mds", lambda: None) + _wait_for_terminal(manager, ref) + + event = _terminal_envelope(manager.snapshot_for_delivery(ref)) + receipt = terminal_delivery_marker(event) + + # The browser's acknowledgement HTTP response may be lost after the server + # mutation. That is safe: terminal UI and receipt already landed together. + assert manager.acknowledge(ref, receipt["terminal_revision"]) + _dropped_ack_response = True + assert manager.active_ref() is None + assert manager.snapshot_for_delivery(ref).delivery_attempt == 1 + + +def test_delayed_old_event_cannot_render_over_a_new_generation(): + ids = iter(("old-job", "new-job")) + manager = JobManager(id_factory=lambda: next(ids)) + with ThreadPoolExecutor(max_workers=1) as executor: + old_ref = manager.submit(executor, "rf", lambda: None) + old_terminal = _wait_for_terminal(manager, old_ref) + old_event = _terminal_envelope( + manager.snapshot_for_delivery(old_ref) + ) + assert manager.acknowledge( + old_ref, + old_terminal.terminal.revision, + ) + + new_ref = manager.submit(executor, "rf", lambda: None) + _wait_for_terminal(manager, new_ref) + + # This models a delayed Dash response from the preceding computation. + assert terminal_event_for_job( + old_event, + new_ref.as_dict(), + expected_kind="rf", + ) is None + assert manager.snapshot(new_ref).acknowledged is False diff --git a/src/test/test_worker_protocol.py b/src/test/test_worker_protocol.py new file mode 100644 index 0000000..371b271 --- /dev/null +++ b/src/test/test_worker_protocol.py @@ -0,0 +1,80 @@ +"""Deterministic tests for framed worker messages and heartbeats.""" + +from __future__ import annotations + +import socket + +import pytest + +from treetracer.worker_protocol import ( + HeartbeatEmitter, + WorkerConnectionClosed, + WorkerProtocolError, + receive_message, + send_message, +) + + +def test_framed_mapping_round_trip(): + parent, worker = socket.socketpair() + try: + send_message(parent, {"job": "compute_rf", "kwargs": {"n": 3}}) + assert receive_message(worker) == { + "job": "compute_rf", + "kwargs": {"n": 3}, + } + finally: + parent.close() + worker.close() + + +def test_clean_eof_and_truncated_frame_are_distinct(): + parent, worker = socket.socketpair() + parent.close() + with pytest.raises(WorkerConnectionClosed): + receive_message(worker) + worker.close() + + parent, worker = socket.socketpair() + try: + parent.sendall(b"\x08\x00\x00\x00abc") + parent.close() + with pytest.raises(WorkerProtocolError, match="during frame body"): + receive_message(worker) + finally: + worker.close() + + +def test_heartbeat_frames_stop_before_the_result_frame(): + parent, worker = socket.socketpair() + worker.settimeout(1.0) + emitter = HeartbeatEmitter( + parent, + "compute_mds", + interval_s=0.01, + ) + try: + emitter.start() + first = receive_message(worker) + second = receive_message(worker) + assert first["type"] == "heartbeat" + assert first["phase"] == "accepted" + assert second["type"] == "heartbeat" + assert second["sequence"] > first["sequence"] + + emitter.stop() + send_message( + parent, + {"type": "result", "ok": True, "result": [1, 2, 3]}, + lock=emitter.send_lock, + ) + result = receive_message(worker) + assert result == { + "type": "result", + "ok": True, + "result": [1, 2, 3], + } + finally: + emitter.stop() + parent.close() + worker.close() diff --git a/src/treetracer/__init__.py b/src/treetracer/__init__.py index 2da1582..041dab6 100644 --- a/src/treetracer/__init__.py +++ b/src/treetracer/__init__.py @@ -5,10 +5,10 @@ 1. **Worker mode** (env var set) — re-entrant subprocess invocation from ``callbacks/persistent_worker.py``. Skips all GUI imports and - enters a request loop reading length-prefixed pickle frames from - stdin, dispatching to one of the registered worker functions, and - writing the result back to stdout. Lives for the lifetime of the - parent app. + enters a request loop reading length-prefixed pickle frames from a + localhost socket, dispatching to one of the registered worker functions, + and returning heartbeat/result frames on that socket. Lives for the + lifetime of the parent app. The persistent design avoids paying ~1.5 s of Python interpreter boot + import on every Compute RF / View consensus tree click — the boot is @@ -43,12 +43,19 @@ def _run_persistent_worker() -> int: shared worker log (``treetracer._worker_log``) so a stuck worker can be diagnosed post-mortem by reading one file. """ - import pickle + import math import socket - import struct import traceback + from collections.abc import Mapping from ._worker_log import log as wlog, get_log_path + from .worker_protocol import ( + HeartbeatEmitter, + WorkerConnectionClosed, + WorkerProtocolError, + receive_message, + send_message, + ) wlog("entered _run_persistent_worker") # Surface the log path on stderr too. The parent's stderr drainer @@ -85,38 +92,63 @@ def _run_persistent_worker() -> int: return 1 wlog("connected; entering recv loop") - def _recv_exactly(n: int) -> bytes: - data = b"" - while len(data) < n: - chunk = sock.recv(n - len(data)) - if not chunk: - return b"" # EOF — parent closed; loop will return - data += chunk - return data + heartbeat_interval_raw = os.environ.get( + "TREETRACER_WORKER_HEARTBEAT_INTERVAL_S", + "2", + ) + try: + heartbeat_interval_s = float(heartbeat_interval_raw) + if not math.isfinite(heartbeat_interval_s) or heartbeat_interval_s <= 0: + raise ValueError + except ValueError: + heartbeat_interval_s = 2.0 + wlog( + "invalid TREETRACER_WORKER_HEARTBEAT_INTERVAL_S=" + f"{heartbeat_interval_raw!r}; using 2s" + ) while True: - wlog("waiting for next request header (4-byte size)") - size_bytes = _recv_exactly(4) - if not size_bytes: + wlog("waiting for next request frame") + try: + request = receive_message(sock) + except WorkerConnectionClosed: wlog("EOF on socket; clean shutdown") try: sock.close() except OSError: pass return 0 - (size,) = struct.unpack(" bytes: else: wlog(f"FATAL: unknown job {job!r}") raise RuntimeError(f"Unknown job: {job!r}") - response = {"ok": True, "result": result} + response = {"type": "result", "ok": True, "result": result} wlog("job succeeded; serialising response") except BaseException as e: # noqa: BLE001 — defensive: one bad # job must not crash the worker. wlog(f"job raised: {type(e).__name__}: {e}") response = { + "type": "result", "ok": False, "error": f"{type(e).__name__}: {e}", "traceback": traceback.format_exc(), } + finally: + heartbeat.stop() - response_bytes = pickle.dumps(response) - wlog(f"sending response; size={len(response_bytes)}") - # sendall loops internally on short sends — guaranteed to send - # all bytes or raise OSError. The length prefix lets the - # parent know exactly how many bytes to recv. + wlog("sending result response") try: - sock.sendall(struct.pack(" list[Any]: + """Map RF/MDS progress to only the dynamic components in the layout. + + Dash rejects a callback response if it names a concrete Output component + that is not currently mounted. RF and MDS banners are mutually exclusive + dynamic children, so wildcard Outputs plus their matched IDs are required + here. Unknown matches receive ``no_update`` defensively. + """ + updates = [] + for component_id in component_ids or []: + which = ( + component_id.get("which") + if isinstance(component_id, Mapping) + else None + ) + if which == "rf": + updates.append(rf_progress[position]) + elif which == "mds": + updates.append(mds_progress[position]) + else: + updates.append(no_update) + return updates + + +def _dynamic_progress_outputs( + progress_bar_ids: Any, + progress_label_ids: Any, + rf_progress: tuple[Any, Any] = (no_update, no_update), + mds_progress: tuple[Any, Any] = (no_update, no_update), +) -> tuple[list[Any], list[Any]]: + return ( + _progress_component_updates( + progress_bar_ids, + rf_progress, + mds_progress, + position=0, + ), + _progress_component_updates( + progress_label_ids, + rf_progress, + mds_progress, + position=1, + ), + ) + + def _matching_active_receipt(receipts: Any): active = job_manager.active_ref() if active is None: @@ -209,10 +261,8 @@ def register_job_reconciliation_callbacks(): Output("compute-poll-interval", "disabled"), Output("compute-terminal-event-store", "data"), Output("compute-busy-store", "data"), - Output("rf-progress-bar", "value"), - Output("rf-progress-label", "children"), - Output("mds-progress-bar", "value"), - Output("mds-progress-label", "children"), + Output({"type": "compute-progress-bar", "which": ALL}, "value"), + Output({"type": "compute-progress-label", "which": ALL}, "children"), Input("compute-poll-interval", "n_intervals"), # Per-workflow stores wake this single owner immediately after submit. Input("rf-job-store", "data"), @@ -224,6 +274,8 @@ def register_job_reconciliation_callbacks(): State("compute-busy-store", "data"), State("rf-progress-path", "data"), State("mds-progress-path", "data"), + State({"type": "compute-progress-bar", "which": ALL}, "id"), + State({"type": "compute-progress-label", "which": ALL}, "id"), prevent_initial_call=True, ) def reconcile_compute_job( @@ -237,31 +289,35 @@ def reconcile_compute_job( current_busy, rf_progress_path, mds_progress_path, + progress_bar_ids, + progress_label_ids, ): """Own polling, terminal delivery, progress, and the global gate.""" active = job_manager.active_ref() busy = _busy_update(active, current_busy) if active is None: + progress_outputs = _dynamic_progress_outputs( + progress_bar_ids, + progress_label_ids, + ) return ( True, no_update, busy, - no_update, - no_update, - no_update, - no_update, + *progress_outputs, ) snapshot = job_manager.snapshot(active) if snapshot is None: + progress_outputs = _dynamic_progress_outputs( + progress_bar_ids, + progress_label_ids, + ) return ( True, no_update, _busy_update(None, current_busy), - no_update, - no_update, - no_update, - no_update, + *progress_outputs, ) rf_progress = (no_update, no_update) @@ -271,14 +327,15 @@ def reconcile_compute_job( if snapshot.terminal is not None: delivery = job_manager.snapshot_for_delivery(active) if delivery is None: + progress_outputs = _dynamic_progress_outputs( + progress_bar_ids, + progress_label_ids, + ) return ( True, no_update, _busy_update(None, current_busy), - no_update, - no_update, - no_update, - no_update, + *progress_outputs, ) terminal_event = _terminal_envelope(delivery) if active.kind == "rf": @@ -305,12 +362,17 @@ def reconcile_compute_job( # Polling remains enabled through terminal presentation. The receipt # callback acknowledges server state; the next tick then takes the # active=None branch and is the sole path that disables this interval. + progress_outputs = _dynamic_progress_outputs( + progress_bar_ids, + progress_label_ids, + rf_progress, + mds_progress, + ) return ( False, terminal_event, busy, - *rf_progress, - *mds_progress, + *progress_outputs, ) @callback( diff --git a/src/treetracer/callbacks/persistent_worker.py b/src/treetracer/callbacks/persistent_worker.py index f24c83b..dd813a2 100644 --- a/src/treetracer/callbacks/persistent_worker.py +++ b/src/treetracer/callbacks/persistent_worker.py @@ -34,15 +34,12 @@ "job": "compute_rf" | "compute_consensus_tree" | ..., "kwargs": {...}, })] - response: [4-byte LE length M][M bytes pickle.dumps({ - "ok": True, "result": ... - }) | pickle.dumps({ - "ok": False, "error": str, "traceback": str - })] + heartbeat: [frame containing {"type": "heartbeat", ...}] + response: [frame containing {"type": "result", "ok": bool, ...}] -The wire is synchronous: every request gets exactly one response, -in order. ``submit_job`` is serialised via ``_lock`` so concurrent -callbacks don't corrupt the stream. +The wire is synchronous: every request gets zero or more heartbeat frames and +then exactly one result frame. ``submit_job`` is serialised via ``_lock`` so +concurrent callbacks don't corrupt the stream. Bootstrap (parent → worker rendezvous): @@ -64,18 +61,28 @@ ``cancel_current_job`` does exactly that. The kill makes ``submit_job``'s blocking ``recv`` return EOF; it then raises ``JobCancelled``. A fresh worker is spawned immediately so the next compute stays warm. + +Watchdog: + +The worker emits a tiny heartbeat every two seconds from a daemon thread. The +parent treats a prolonged heartbeat gap as an unresponsive worker, kills it, +starts a replacement, and raises a normal compute exception. ``JobManager`` +then delivers that failure through the same sticky terminal protocol as every +other result. A generous per-operation maximum runtime is a final safety net +for a compute function that remains able to heartbeat but never returns. """ from __future__ import annotations import collections +import math import os -import pickle import socket -import struct import subprocess import sys import threading +import time +from dataclasses import dataclass from pathlib import Path from typing import Any, Dict @@ -85,13 +92,21 @@ # parent-side entries are tagged ``[parent]`` so an interleaved # timestamp-sorted view shows the IPC handshake clearly. from .._worker_log import log as _wlog +from ..worker_protocol import ( + WorkerConnectionClosed, + WorkerProtocolError, + receive_message, + send_message, +) _worker_proc: subprocess.Popen | None = None _worker_sock: socket.socket | None = None # connected to current worker _lock = threading.Lock() +_job_state_lock = threading.Lock() -# Cancellation state. ``_current_job`` is the name of the job whose +# Cancellation state, protected by ``_job_state_lock``. ``_current_job`` is +# the name of the job whose # response ``submit_job`` is currently blocked on (``None`` when the # worker is idle); ``_cancelled`` is set by ``cancel_current_job`` so # ``submit_job`` can tell a user Stop apart from a genuine crash. @@ -105,6 +120,131 @@ _stderr_drainer: "_StderrDrainer | None" = None +_DEFAULT_HEARTBEAT_INTERVAL_S = 2.0 +_DEFAULT_HEARTBEAT_TIMEOUT_S = 120.0 +_DEFAULT_MAX_RUNTIME_BY_JOB_S = { + # RF, MDS, and Pseudo-ESS can scale quadratically and legitimately run for + # hours on older machines. The ceilings are intentionally conservative. + "compute_rf": 12 * 60 * 60, + "compute_mds": 6 * 60 * 60, + "compute_pseudo_ess": 6 * 60 * 60, + # These operations work on a selection, row, or selected clade columns. + "compute_consensus_tree": 2 * 60 * 60, + "compute_rf_trace": 60 * 60, + "compute_clade_frequencies": 60 * 60, +} +_DEFAULT_MAX_RUNTIME_S = 6 * 60 * 60 + + +@dataclass(frozen=True, slots=True) +class WatchdogPolicy: + """Resolved health limits for one worker request. + + ``None`` disables the corresponding limit. Environment overrides accept + seconds; setting a value to ``0`` disables that limit deliberately. + """ + + heartbeat_timeout_s: float | None + max_runtime_s: float | None + + +class WorkerUnresponsive(RuntimeError): + """Base class for worker failures detected by the parent watchdog.""" + + +class WorkerHeartbeatTimeout(WorkerUnresponsive): + """No complete heartbeat or result frame arrived within the health limit.""" + + +class WorkerRuntimeExceeded(WorkerUnresponsive): + """A job exceeded its configured maximum runtime while still heartbeating.""" + + +def watchdog_policy( + job_name: str, + *, + max_runtime_s: float | None = None, +) -> WatchdogPolicy: + """Resolve defaults and environment overrides for ``job_name``. + + Override precedence for maximum runtime is: explicit argument, per-job + environment variable, global environment variable, built-in default. + For example, RF can be overridden with + ``TREETRACER_WORKER_MAX_RUNTIME_COMPUTE_RF_S``. + """ + + heartbeat_interval = _heartbeat_interval_s() + heartbeat_timeout = _duration_from_env( + "TREETRACER_WORKER_HEARTBEAT_TIMEOUT_S", + _DEFAULT_HEARTBEAT_TIMEOUT_S, + ) + if ( + heartbeat_timeout is not None + and heartbeat_timeout < heartbeat_interval * 3 + ): + adjusted = heartbeat_interval * 3 + _wlog( + "[parent] heartbeat timeout is shorter than three worker " + f"intervals; using {adjusted:g}s" + ) + heartbeat_timeout = adjusted + if max_runtime_s is None: + per_job_name = ( + "TREETRACER_WORKER_MAX_RUNTIME_" + f"{str(job_name).upper()}_S" + ) + default_runtime = _DEFAULT_MAX_RUNTIME_BY_JOB_S.get( + str(job_name), + _DEFAULT_MAX_RUNTIME_S, + ) + if per_job_name in os.environ: + max_runtime = _duration_from_env(per_job_name, default_runtime) + else: + max_runtime = _duration_from_env( + "TREETRACER_WORKER_MAX_RUNTIME_S", + default_runtime, + ) + else: + max_runtime = _normalise_duration( + max_runtime_s, + name="max_runtime_s", + ) + return WatchdogPolicy(heartbeat_timeout, max_runtime) + + +def _duration_from_env(name: str, default: float) -> float | None: + raw = os.environ.get(name) + if raw is None: + return float(default) + try: + return _normalise_duration(raw, name=name) + except (TypeError, ValueError): + _wlog(f"[parent] ignoring invalid {name}={raw!r}") + return float(default) + + +def _heartbeat_interval_s() -> float: + name = "TREETRACER_WORKER_HEARTBEAT_INTERVAL_S" + raw = os.environ.get(name) + if raw is None: + return _DEFAULT_HEARTBEAT_INTERVAL_S + try: + interval = float(raw) + if not math.isfinite(interval) or interval <= 0: + raise ValueError + return interval + except (TypeError, ValueError): + _wlog(f"[parent] ignoring invalid {name}={raw!r}") + return _DEFAULT_HEARTBEAT_INTERVAL_S + + +def _normalise_duration(value: Any, *, name: str) -> float | None: + seconds = float(value) + if not math.isfinite(seconds) or seconds < 0: + raise ValueError(f"{name} must be a finite non-negative number") + return None if seconds == 0 else seconds + + class _StderrDrainer: """Continuously read the worker's stderr in a background thread so the worker never blocks on a full pipe buffer. @@ -166,7 +306,7 @@ def tail(self) -> bytes: class JobCancelled(RuntimeError): """Raised by ``submit_job`` when the worker was killed via - ``cancel_current_job()`` — lets the polling callbacks render a + ``cancel_current_job()`` — lets the managed terminal adapters render a user-requested Stop as a neutral "cancelled" state rather than a red error.""" @@ -218,6 +358,9 @@ def _spawn_worker() -> subprocess.Popen: env = os.environ.copy() env["TREETRACER_WORKER_MODE"] = "persistent" env["TREETRACER_WORKER_PORT"] = str(port) + env["TREETRACER_WORKER_HEARTBEAT_INTERVAL_S"] = str( + _heartbeat_interval_s() + ) _wlog(f"[parent] _spawn_worker: argv={_worker_argv()!r}") # ── Spawn. Stdio is intentionally DEVNULL for stdin/stdout — @@ -233,32 +376,48 @@ def _spawn_worker() -> subprocess.Popen: ) _wlog(f"[parent] _spawn_worker: Popen returned; worker pid={proc.pid}") + # Drain immediately, including during rendezvous. A noisy import must not + # fill stderr and prevent the worker from ever reaching socket.connect(). + _stderr_drainer = _StderrDrainer(proc.stderr) + _wlog("[parent] _spawn_worker: stderr drainer started") + # ── Wait for the worker to connect back. Generous timeout because # Python startup in the bundle (especially on emulated Windows) can # take several seconds. If the worker dies before connecting we # surface stderr from its drainer so the user knows why. - listener.settimeout(_RENDEZVOUS_TIMEOUT_S) + listener.settimeout(0.25) + deadline = time.monotonic() + _RENDEZVOUS_TIMEOUT_S try: - sock, addr = listener.accept() - except socket.timeout: - # Kill the worker, salvage whatever it printed to stderr. + while True: + try: + sock, addr = listener.accept() + break + except socket.timeout: + if proc.poll() is not None: + raise RuntimeError( + "persistent worker exited before connecting " + f"(exit {proc.returncode})" + ) + if time.monotonic() >= deadline: + raise RuntimeError( + "persistent worker did not connect within " + f"{_RENDEZVOUS_TIMEOUT_S:.0f}s" + ) + except Exception as exc: try: - proc.kill() + if proc.poll() is None: + proc.kill() except OSError: pass - # Drainer might not even exist yet — read stderr directly, - # but bounded so we don't block forever. - stderr_tail = b"" try: - if proc.stderr is not None: - stderr_tail = proc.stderr.read() - except OSError: + proc.wait(timeout=2.0) + except (OSError, subprocess.TimeoutExpired): pass + stderr_tail = _worker_stderr_tail() raise RuntimeError( - f"persistent worker did not connect within " - f"{_RENDEZVOUS_TIMEOUT_S:.0f}s. stderr tail: " + f"{exc}. stderr tail: " f"{stderr_tail.decode('utf-8', errors='replace')[-500:]}" - ) + ) from exc finally: listener.close() @@ -266,11 +425,6 @@ def _spawn_worker() -> subprocess.Popen: _worker_sock = sock _wlog(f"[parent] _spawn_worker: worker connected from {addr}") - # Fresh drainer per worker — the previous drainer (if any) is - # still draining the dead worker's stderr until that pipe EOFs; - # its daemon thread will exit on its own. - _stderr_drainer = _StderrDrainer(proc.stderr) - _wlog("[parent] _spawn_worker: stderr drainer started") return proc @@ -286,11 +440,8 @@ def _worker_stderr_tail() -> bytes: def start() -> None: """Spawn the persistent worker subprocess if it isn't already - running. Returns immediately — the subprocess boots in the - background. The first ``submit_job`` call after this returns will - block on the worker reading from stdin, which is automatic — OS - pipe buffering covers any race between Popen returning and the - worker entering its read loop. + running. The application invokes this on a daemon startup thread; this + function returns after the worker completes its socket rendezvous. """ global _worker_proc with _lock: @@ -336,22 +487,23 @@ def cancel_current_job() -> bool: Safe to call from any thread. It deliberately does **not** acquire ``_lock``: the thread that called ``submit_job`` holds that lock for the whole job, so acquiring it here would block until the job - finished on its own — exactly what a Stop button must avoid. We only - read the ``_worker_proc`` reference (atomic) and signal it, both - thread-safe. + finished on its own — exactly what a Stop button must avoid. A separate, + tiny state lock closes the race between completion and cancellation without + waiting for worker IPC. The kill makes ``submit_job``'s blocking read return EOF; that call then raises ``JobCancelled`` and respawns a fresh worker. """ global _cancelled - proc = _worker_proc # atomic snapshot of the module global - if proc is None or proc.poll() is not None: - return False # no live worker - if _current_job is None: - return False # worker idle — nothing to interrupt - # Order matters: set the flag before the kill so submit_job sees it - # on the EOF the kill is about to cause. - _cancelled = True + with _job_state_lock: + proc = _worker_proc + if proc is None or proc.poll() is not None: + return False # no live worker + if _current_job is None: + return False # worker idle — nothing to interrupt + # Set the flag before the kill so submit_job sees it on the EOF the + # kill is about to cause. + _cancelled = True try: proc.kill() except OSError: @@ -359,6 +511,86 @@ def cancel_current_job() -> bool: return True +def _begin_current_job(job_name: str) -> None: + global _current_job, _cancelled + with _job_state_lock: + _cancelled = False + _current_job = job_name + + +def _finish_current_job() -> bool: + """Close the cancellation window and return whether Stop won it.""" + global _current_job + with _job_state_lock: + was_cancelled = _cancelled + _current_job = None + return was_cancelled + + +def _clear_current_job() -> None: + global _current_job + with _job_state_lock: + _current_job = None + + +def _restart_worker_locked(reason: str) -> bool: + """Retire the current process/socket and warm a replacement. + + ``submit_job`` and ``_ensure_running`` call this while holding ``_lock``. + The old socket is closed before the process is killed so no late frame can + be mistaken for a response from the replacement worker. + """ + global _worker_proc, _worker_sock + + old_proc = _worker_proc + old_sock = _worker_sock + _worker_proc = None + _worker_sock = None + + if old_sock is not None: + try: + old_sock.shutdown(socket.SHUT_RDWR) + except OSError: + pass + try: + old_sock.close() + except OSError: + pass + + if old_proc is not None and old_proc.poll() is None: + try: + old_proc.kill() + except OSError: + pass + try: + old_proc.wait(timeout=2.0) + except (OSError, subprocess.TimeoutExpired): + pass + + _wlog(f"[parent] restarting persistent worker: {reason}") + try: + _worker_proc = _spawn_worker() + except Exception as exc: # noqa: BLE001 — next submission retries startup + _worker_proc = None + _worker_sock = None + _wlog( + "[parent] worker restart failed: " + f"{type(exc).__name__}: {exc}" + ) + return False + _wlog(f"[parent] worker restart complete; pid={_worker_proc.pid}") + return True + + +def _recovery_message(restarted: bool) -> str: + if restarted: + return "The worker was restarted and is ready for another computation." + return ( + "The worker could not be restarted immediately; TreeTracer will retry " + "startup on the next computation." + ) + + def _ensure_running() -> subprocess.Popen: """Return the worker proc, restarting it if it died. Holds ``_lock`` around the check + restart so concurrent submitters can't race.""" @@ -375,117 +607,214 @@ def _ensure_running() -> subprocess.Popen: f"restarting. stderr tail: " f"{stderr_tail.decode('utf-8', errors='replace')[-500:]}\n" ) - # Close the dead worker's socket before spawning a fresh one. - # ``_spawn_worker`` will install a new one alongside the new - # ``_worker_proc``. - if _worker_sock is not None: - try: - _worker_sock.close() - except OSError: - pass - _worker_sock = None - _worker_proc = _spawn_worker() + if not _restart_worker_locked("worker was not running"): + raise RuntimeError("persistent worker could not be restarted") return _worker_proc -def submit_job(job_name: str, **kwargs: Any) -> Dict[str, Any]: - """Send a job to the persistent worker and block on its response. +def _receive_worker_result( + sock: socket.socket, + job_name: str, + policy: WatchdogPolicy, + *, + started_at: float, +) -> dict[str, Any]: + """Consume heartbeat frames until the job's result frame arrives.""" + last_heartbeat_at = started_at + last_sequence = 0 + + while True: + now = time.monotonic() + heartbeat_remaining = ( + None + if policy.heartbeat_timeout_s is None + else policy.heartbeat_timeout_s - (now - last_heartbeat_at) + ) + runtime_remaining = ( + None + if policy.max_runtime_s is None + else policy.max_runtime_s - (now - started_at) + ) + + if runtime_remaining is not None and runtime_remaining <= 0: + raise WorkerRuntimeExceeded( + f"{job_name} exceeded its {policy.max_runtime_s:.0f}s " + "maximum runtime" + ) + if heartbeat_remaining is not None and heartbeat_remaining <= 0: + raise WorkerHeartbeatTimeout( + f"persistent worker sent no heartbeat or result for " + f"{policy.heartbeat_timeout_s:.0f}s during {job_name}" + ) + + waits = [ + value + for value in (heartbeat_remaining, runtime_remaining) + if value is not None + ] + sock.settimeout(min(waits) if waits else None) + try: + message = receive_message(sock) + except socket.timeout as exc: + now = time.monotonic() + if ( + policy.max_runtime_s is not None + and now - started_at >= policy.max_runtime_s + ): + raise WorkerRuntimeExceeded( + f"{job_name} exceeded its {policy.max_runtime_s:.0f}s " + "maximum runtime" + ) from exc + raise WorkerHeartbeatTimeout( + f"persistent worker sent no heartbeat or result for " + f"{policy.heartbeat_timeout_s:.0f}s during {job_name}" + ) from exc + + message_type = message.get("type") + if message_type == "heartbeat": + if message.get("job") != job_name: + raise WorkerProtocolError( + "heartbeat job mismatch: " + f"expected {job_name!r}, got {message.get('job')!r}" + ) + try: + sequence = int(message["sequence"]) + except (KeyError, TypeError, ValueError) as exc: + raise WorkerProtocolError( + "heartbeat is missing a valid sequence" + ) from exc + if sequence < 1: + raise WorkerProtocolError( + f"heartbeat has invalid sequence {sequence}" + ) + if sequence <= last_sequence: + raise WorkerProtocolError( + "heartbeat sequence did not advance: " + f"previous={last_sequence}, current={sequence}" + ) + last_sequence = sequence + last_heartbeat_at = time.monotonic() + if sequence == 1 or sequence % 30 == 0: + _wlog( + "[parent] worker heartbeat: " + f"job={job_name!r}, sequence={sequence}, " + f"elapsed={message.get('elapsed_s', '?')}s" + ) + continue + + # Accept an untyped result for one-version rolling compatibility with + # a worker started just before an application update. + if message_type in (None, "result"): + return message + raise WorkerProtocolError( + f"unexpected worker message type {message_type!r}" + ) - The wait happens inside a ``subprocess`` pipe read, which releases - the GIL — so the parent's other threads (Dash callbacks, the - pywebview event loop on its native thread, etc.) stay responsive. - Args: - job_name: Managed compute operation name; must match a branch in - ``__init__.py:_run_persistent_worker``. - **kwargs: forwarded to the worker function. +def submit_job( + job_name: str, + *, + max_runtime_s: float | None = None, + **kwargs: Any, +) -> Dict[str, Any]: + """Send one job and wait for heartbeats followed by its result. - Returns the worker function's return value (unpickled). Raises - ``JobCancelled`` if the user stopped the job via - ``cancel_current_job``, or ``RuntimeError`` if the worker reported - an error or crashed. + Socket waits release the GIL, so Dash and the desktop event loop remain + responsive. Missing heartbeats, a maximum-runtime breach, a corrupt frame, + or a dead connection retires the worker and warms a replacement before the + exception reaches ``JobManager``. + + ``max_runtime_s`` overrides the environment/default ceiling for this call; + pass ``0`` to disable only the hard runtime ceiling. Heartbeat monitoring + remains independently configurable through + ``TREETRACER_WORKER_HEARTBEAT_TIMEOUT_S``. """ - global _worker_proc, _worker_sock, _current_job, _cancelled - _wlog(f"[parent] submit_job called: job={job_name!r}, kwarg keys={sorted(kwargs.keys())}") + global _worker_sock + + policy = watchdog_policy(job_name, max_runtime_s=max_runtime_s) + _wlog( + f"[parent] submit_job called: job={job_name!r}, " + f"kwarg keys={sorted(kwargs.keys())}, policy={policy}" + ) with _lock: _wlog("[parent] submit_job: _lock acquired") - _cancelled = False proc = _ensure_running() - # _ensure_running guarantees _worker_sock is set alongside - # _worker_proc (both are populated by _spawn_worker). assert _worker_sock is not None sock = _worker_sock - _wlog(f"[parent] submit_job: worker pid={proc.pid}, alive={proc.poll() is None}") + _wlog( + f"[parent] submit_job: worker pid={proc.pid}, " + f"alive={proc.poll() is None}" + ) - request_bytes = pickle.dumps({"job": job_name, "kwargs": kwargs}) - _wlog(f"[parent] submit_job: pickled request; size={len(request_bytes)}") - # ``_current_job`` is the cancellation window: while it's set, - # cancel_current_job() may kill this worker. - _current_job = job_name + _begin_current_job(job_name) + started_at = time.monotonic() try: - _wlog("[parent] submit_job: sending 4-byte size header on socket") - sock.sendall(struct.pack("") - + "\n" + response.get("traceback", "") + str(response.get("error", "")) + + "\n" + + str(response.get("traceback", "")) ) _wlog("[parent] submit_job: returning result to caller") return response["result"] - - -def _recv_exactly(sock: socket.socket, n: int) -> bytes: - """``sock.recv(n)`` can return short reads (e.g. across TCP MSS - boundaries) — keep reading until we have N bytes or hit EOF. - - Returns the bytes received (which may be fewer than ``n`` if the - peer closed the connection mid-frame; callers treat that as EOF). - """ - data = b"" - while len(data) < n: - chunk = sock.recv(n - len(data)) - if not chunk: - return data - data += chunk - return data diff --git a/src/treetracer/ui/widgets.py b/src/treetracer/ui/widgets.py index ccfe842..01b908a 100644 --- a/src/treetracer/ui/widgets.py +++ b/src/treetracer/ui/widgets.py @@ -38,10 +38,9 @@ def computing_banner(title: str, message: str, which: str, When ``show_progress=True`` the banner stacks an additional ``dmc.Progress`` bar and a small status label under the - message/stop row. The bar's id is ``f"{which}-progress-bar"`` and - the label's id is ``f"{which}-progress-label"`` so a separate - polling callback can drive them off a sidecar progress file (see - ``update_rf_progress`` in ``callbacks/compute.py``). + message/stop row. The bar and label use pattern-matching IDs so the + central reconciler can update whichever banner currently exists without + naming the absent RF or MDS banner as a Dash callback output. """ # flex:1 + minWidth:0 lets the (often long) message shrink and # wrap instead of shoving the Stop button past the Alert's right @@ -62,13 +61,17 @@ def computing_banner(title: str, message: str, which: str, [ message_row, dmc.Progress( - id=f"{which}-progress-bar", + id={"type": "compute-progress-bar", "which": which}, value=0, size="md", color="blue", ), - dmc.Text("starting…", size="xs", c="dimmed", - id=f"{which}-progress-label"), + dmc.Text( + "starting…", + size="xs", + c="dimmed", + id={"type": "compute-progress-label", "which": which}, + ), ], gap="xs", ) diff --git a/src/treetracer/worker_protocol.py b/src/treetracer/worker_protocol.py new file mode 100644 index 0000000..f20ea51 --- /dev/null +++ b/src/treetracer/worker_protocol.py @@ -0,0 +1,173 @@ +"""Small framed-message protocol shared by the GUI and compute worker. + +The persistent worker uses one localhost TCP socket for requests, heartbeat +messages, and results. Every message is a length-prefixed pickle frame:: + + [4-byte little-endian payload length][pickled mapping] + +Only TreeTracer processes connected through the one-shot localhost rendezvous +use this protocol. The size limit protects both sides from allocating an +unbounded buffer if a connection is corrupt. +""" + +from __future__ import annotations + +import pickle +import socket +import struct +import threading +import time +from collections.abc import Callable, Mapping +from typing import Any + + +MAX_FRAME_BYTES = 256 * 1024 * 1024 + + +class WorkerConnectionClosed(EOFError): + """The peer closed the socket before another complete frame arrived.""" + + +class WorkerProtocolError(RuntimeError): + """A worker frame is truncated, malformed, or violates the protocol.""" + + +def encode_message(message: Mapping[str, Any]) -> bytes: + """Serialize one mapping with its four-byte length prefix.""" + payload = pickle.dumps(dict(message), protocol=pickle.HIGHEST_PROTOCOL) + if len(payload) > MAX_FRAME_BYTES: + raise WorkerProtocolError( + f"worker frame is {len(payload):,} bytes; " + f"limit is {MAX_FRAME_BYTES:,} bytes" + ) + return struct.pack(" None: + """Send one complete frame, optionally serializing concurrent writers.""" + frame = encode_message(message) + if lock is None: + sock.sendall(frame) + return + with lock: + sock.sendall(frame) + + +def receive_message(sock: socket.socket) -> dict[str, Any]: + """Receive and validate one complete mapping frame.""" + size_bytes = _recv_exactly(sock, 4, field="header") + (size,) = struct.unpack(" MAX_FRAME_BYTES: + raise WorkerProtocolError( + f"worker announced a {size:,}-byte frame; " + f"limit is {MAX_FRAME_BYTES:,} bytes" + ) + payload = _recv_exactly(sock, size, field="body") + try: + message = pickle.loads(payload) + except Exception as exc: + raise WorkerProtocolError( + f"could not decode worker frame: {type(exc).__name__}: {exc}" + ) from exc + if not isinstance(message, Mapping): + raise WorkerProtocolError( + f"worker frame must contain a mapping, got {type(message).__name__}" + ) + return dict(message) + + +def _recv_exactly(sock: socket.socket, size: int, *, field: str) -> bytes: + data = bytearray() + while len(data) < size: + chunk = sock.recv(size - len(data)) + if not chunk: + if not data: + raise WorkerConnectionClosed( + f"worker connection closed before frame {field}" + ) + raise WorkerProtocolError( + f"worker connection closed during frame {field} " + f"({len(data):,}/{size:,} bytes)" + ) + data.extend(chunk) + return bytes(data) + + +class HeartbeatEmitter: + """Send serialized heartbeat frames while one worker job is running. + + The result writer uses :attr:`send_lock` too, so a heartbeat can never be + interleaved with a result frame. ``Event.wait`` makes shutdown immediate + instead of sleeping for the rest of the heartbeat period. + """ + + def __init__( + self, + sock: socket.socket, + job_name: str, + *, + interval_s: float, + log: Callable[[str], None] | None = None, + clock: Callable[[], float] = time.monotonic, + ) -> None: + if interval_s <= 0: + raise ValueError("heartbeat interval must be positive") + self.send_lock = threading.Lock() + self._sock = sock + self._job_name = str(job_name) + self._interval_s = float(interval_s) + self._log = log or (lambda _message: None) + self._clock = clock + self._started_at = self._clock() + self._sequence = 0 + self._stop_event = threading.Event() + self._thread: threading.Thread | None = None + + def start(self) -> None: + """Publish acceptance immediately, then start periodic heartbeats.""" + if self._thread is not None: + raise RuntimeError("heartbeat emitter already started") + self._emit("accepted") + self._thread = threading.Thread( + target=self._run, + name="treetracer-worker-heartbeat", + daemon=True, + ) + self._thread.start() + + def stop(self) -> None: + """Stop before the caller sends the terminal result frame.""" + self._stop_event.set() + thread = self._thread + if thread is not None: + thread.join(timeout=max(1.0, self._interval_s + 0.5)) + + def _run(self) -> None: + while not self._stop_event.wait(self._interval_s): + try: + self._emit("running") + except OSError as exc: + self._log( + "heartbeat send failed: " + f"{type(exc).__name__}: {exc}" + ) + return + + def _emit(self, phase: str) -> None: + self._sequence += 1 + send_message( + self._sock, + { + "type": "heartbeat", + "job": self._job_name, + "sequence": self._sequence, + "phase": phase, + "elapsed_s": max(0.0, self._clock() - self._started_at), + }, + lock=self.send_lock, + ) From 0d50a7d9492137885483ec124737b856612dfbdf Mon Sep 17 00:00:00 2001 From: Sam Hong <20244918+hongsamL@users.noreply.github.com> Date: Mon, 3 Aug 2026 19:38:28 +0200 Subject: [PATCH 8/9] fix polling racing issue with coordinator --- src/test/README.md | 6 +- src/test/test_app_smoke.py | 40 +++++++++ src/test/test_background_jobs.py | 60 ++++++++++--- src/test/test_compute_job_lifecycle.py | 58 ++++++------ src/test/test_managed_analysis_jobs.py | 22 +++-- src/test/test_managed_compute_jobs.py | 24 +++-- src/test/test_terminal_delivery_faults.py | 32 ++++--- src/treetracer/background_jobs.py | 58 ++++++++++-- src/treetracer/callbacks/compute.py | 4 +- src/treetracer/callbacks/job_reconcile.py | 104 ++++++++++++---------- src/treetracer/ui/navbar.py | 1 - 11 files changed, 286 insertions(+), 123 deletions(-) diff --git a/src/test/README.md b/src/test/README.md index f9ff8b1..6e5036a 100644 --- a/src/test/README.md +++ b/src/test/README.md @@ -24,13 +24,13 @@ bash run_tests.sh -k consensus tree -v | `conftest.py` | Session fixtures: NEXUS parse, rapidtrees presence, DendroPy parse. | | `test.trees` | 100-tree BEAST fixture (5 MB). The CI integration anchor. | | `test_app_smoke.py` | App imports, callback registration, figure-layout invariants. | -| `test_background_jobs.py` | Thread-safe job-state transitions, exactly-once finalization, sticky terminal delivery, acknowledgement, cancellation, and reset. | -| `test_compute_job_lifecycle.py` | RF/MDS lifecycle integration, generation checks, wildcard dynamic progress outputs, and terminal replay. | +| `test_background_jobs.py` | Thread-safe job-state transitions, exactly-once finalization, sticky terminal delivery, atomic retry leases, acknowledgement, cancellation, and reset. | +| `test_compute_job_lifecycle.py` | RF/MDS lifecycle integration, generation checks, wildcard dynamic progress outputs, and poll-piggybacked receipt acknowledgement. | | `test_managed_compute_jobs.py` | Pseudo-ESS and consensus managed-job finalization and replay behavior. | | `test_managed_analysis_jobs.py` | RF Trace and clade-comparison worker/finalizer/cache behavior. | | `test_worker_protocol.py` | Framed socket messages, EOF/truncation handling, and heartbeat/result ordering. | | `test_persistent_worker_watchdog.py` | Missing-heartbeat and hard-runtime bounds, protocol validation, configuration, and worker replacement. | -| `test_terminal_delivery_faults.py` | Deterministic terminal-event, terminal-UI/receipt, acknowledgement-loss, and stale-generation scenarios. | +| `test_terminal_delivery_faults.py` | Deterministic terminal-event, terminal-UI/receipt, settling-response-loss, paced-retry, and stale-generation scenarios. | | `test_ess.py` | `effective_sample_size` vs AR(1) closed form + arviz cross-check (iid, AR(2), MA(5), heavy-tail, multimodal). | | `test_pseudo_ess.py` | `compute_pseudo_ess` shape, n-cap, rank-norm bound, row-order sensitivity. | | `test_pcoa.py` | `compute_mds` Procrustes-equivalent to scipy on synthetic Euclidean + real RF. | diff --git a/src/test/test_app_smoke.py b/src/test/test_app_smoke.py index b57d2f9..2f23763 100644 --- a/src/test/test_app_smoke.py +++ b/src/test/test_app_smoke.py @@ -9,6 +9,8 @@ from __future__ import annotations +import json + import pandas as pd @@ -57,6 +59,44 @@ def test_compute_interval_has_one_reconciliation_owner(): assert owners == {"reconcile_compute_job"} +def test_reconciler_piggybacks_receipts_as_state_without_an_ack_callback(): + """Browser receipts settle through the one polling owner. + + Keeping receipts as State avoids a receipt-triggered callback cycle while + removing the independently schedulable acknowledgement request that could + be starved by rapid terminal replays. + """ + from dash import _callback + + reconciler = None + callback_names = set() + for callback_data in _callback.GLOBAL_CALLBACK_MAP.values(): + callback_fn = callback_data.get("callback") + callback_fn = getattr(callback_fn, "__wrapped__", callback_fn) + name = getattr(callback_fn, "__name__", "") + callback_names.add(name) + if name == "reconcile_compute_job": + reconciler = callback_data + + assert reconciler is not None + receipt_states = [] + for state in reconciler.get("state", []): + component_id = state.get("id") + if not isinstance(component_id, str) or not component_id.startswith("{"): + continue + parsed = json.loads(component_id) + if parsed.get("type") == "compute-terminal-receipt": + receipt_states.append((parsed, state.get("property"))) + + assert receipt_states == [ + ( + {"kind": ["ALL"], "type": "compute-terminal-receipt"}, + "data", + ) + ] + assert "acknowledge_terminal_receipt" not in callback_names + + def test_reconciler_uses_wildcards_for_dynamic_progress_banners(): """RF and MDS banners never coexist, so concrete Outputs are unsafe. diff --git a/src/test/test_background_jobs.py b/src/test/test_background_jobs.py index e90620d..c6a9044 100644 --- a/src/test/test_background_jobs.py +++ b/src/test/test_background_jobs.py @@ -31,7 +31,8 @@ def _wait_for_state(manager, ref, expected, timeout=2.0): def test_success_is_sticky_until_matching_acknowledgement(): - manager = JobManager() + now = [10.0] + manager = JobManager(clock=lambda: now[0]) with ThreadPoolExecutor(max_workers=1) as executor: ref = manager.submit( executor, @@ -51,8 +52,10 @@ def test_success_is_sticky_until_matching_acknowledgement(): assert terminal.metadata == {"display_name": "RF_001"} assert manager.active_ref() == ref - first = manager.snapshot_for_delivery(ref) - second = manager.snapshot_for_delivery(ref) + first = manager.claim_terminal_delivery(ref) + assert manager.claim_terminal_delivery(ref) is None + now[0] += 1.0 + second = manager.claim_terminal_delivery(ref) assert first.terminal == second.terminal assert first.delivery_attempt == 1 assert second.delivery_attempt == 2 @@ -76,7 +79,7 @@ def test_stale_generation_cannot_read_update_or_acknowledge_job(): assert manager.snapshot(ref).acknowledged is False -def test_finalizer_runs_once_while_concurrent_readers_replay_terminal(): +def test_finalizer_runs_once_and_delivery_lease_has_one_concurrent_winner(): manager = JobManager() finalizer_entered = threading.Event() release_finalizer = threading.Event() @@ -109,7 +112,7 @@ def finalize(_ref, result): def read_terminal(): barrier.wait(timeout=2) - value = manager.snapshot_for_delivery(ref) + value = manager.claim_terminal_delivery(ref) with snapshots_lock: snapshots.append(value) @@ -122,8 +125,45 @@ def read_terminal(): assert calls == 1 assert len(snapshots) == 8 - assert {s.terminal.payload["result_ref"] for s in snapshots} == {"ESS_001"} - assert sorted(s.delivery_attempt for s in snapshots) == list(range(1, 9)) + deliveries = [snapshot for snapshot in snapshots if snapshot is not None] + assert len(deliveries) == 1 + assert deliveries[0].terminal.payload["result_ref"] == "ESS_001" + assert deliveries[0].delivery_attempt == 1 + assert manager.snapshot(ref).delivery_attempt == 1 + + +def test_terminal_delivery_retries_follow_the_capped_lease_schedule(): + now = [20.0] + manager = JobManager( + clock=lambda: now[0], + terminal_retry_delays=(0.5, 1.0, 2.0), + ) + with ThreadPoolExecutor(max_workers=1) as executor: + ref = manager.submit(executor, "rf", lambda: None) + _wait_for_state(manager, ref, JobState.SUCCEEDED) + + assert manager.claim_terminal_delivery(ref).delivery_attempt == 1 + now[0] += 0.49 + assert manager.claim_terminal_delivery(ref) is None + now[0] += 0.01 + assert manager.claim_terminal_delivery(ref).delivery_attempt == 2 + now[0] += 0.99 + assert manager.claim_terminal_delivery(ref) is None + now[0] += 0.01 + assert manager.claim_terminal_delivery(ref).delivery_attempt == 3 + now[0] += 2.0 + assert manager.claim_terminal_delivery(ref).delivery_attempt == 4 + now[0] += 2.0 + assert manager.claim_terminal_delivery(ref).delivery_attempt == 5 + + +@pytest.mark.parametrize( + "delays", + [(), (0.0,), (float("inf"),), (2.0, 1.0)], +) +def test_terminal_delivery_retry_schedule_must_be_valid(delays): + with pytest.raises(ValueError): + JobManager(terminal_retry_delays=delays) def test_compute_exception_becomes_sticky_failure(): @@ -141,8 +181,8 @@ def fail(): "error_type": "RuntimeError", "stage": "compute", } - assert manager.snapshot_for_delivery(ref).state is JobState.FAILED - assert manager.snapshot_for_delivery(ref).state is JobState.FAILED + assert manager.claim_terminal_delivery(ref).state is JobState.FAILED + assert manager.claim_terminal_delivery(ref) is None def test_configured_worker_cancellation_exception_is_not_an_error(): @@ -186,7 +226,7 @@ def broken_finalizer(_ref, _result): terminal = _wait_for_state(manager, ref, JobState.FAILED) for _ in range(5): - manager.snapshot_for_delivery(ref) + manager.snapshot(ref) assert calls == 1 assert terminal.terminal.payload["stage"] == "finalize" assert terminal.terminal.payload["error_type"] == "ValueError" diff --git a/src/test/test_compute_job_lifecycle.py b/src/test/test_compute_job_lifecycle.py index 048a4b6..53d0a77 100644 --- a/src/test/test_compute_job_lifecycle.py +++ b/src/test/test_compute_job_lifecycle.py @@ -95,7 +95,7 @@ def test_terminal_event_must_match_the_current_browser_generation(): ) == event -def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): +def test_rf_terminal_receipt_is_acknowledged_by_next_reconcile(monkeypatch): manager = JobManager(id_factory=lambda: "rf-test-job") register_calls = [] expected_index = { @@ -145,10 +145,6 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): job_reconcile.register_job_reconciliation_callbacks, ) render = _registered_callback("render_rf_mds_terminal_event") - acknowledge = _registered_callback( - "acknowledge_terminal_receipt", - job_reconcile.register_job_reconciliation_callbacks, - ) rf_bar_ids = [{"type": "compute-progress-bar", "which": "rf"}] rf_label_ids = [{"type": "compute-progress-label", "which": "rf"}] @@ -165,8 +161,22 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): None, rf_bar_ids, rf_label_ids, + [], ) first = render(first_reconcile[1], ref.as_dict(), None, False) + + assert len(first_reconcile) == 5 + assert first_reconcile[0] is False + assert first_reconcile[2]["busy"] is True + assert first_reconcile[3:] == ([100.0], ["complete"]) + assert len(first) == 12 + assert first[1] == expected_index + assert first[2] is False + assert first[8]["terminal_revision"] == terminal.terminal.revision + assert first[8]["delivery_attempt"] == 1 + assert len(register_calls) == 1 + assert manager.snapshot(ref).acknowledged is False + second_reconcile = reconcile( 11, ref.as_dict(), @@ -180,28 +190,17 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): None, rf_bar_ids, rf_label_ids, + [first[8], None, None, None, None], ) - second = render(second_reconcile[1], ref.as_dict(), None, False) - - assert len(first_reconcile) == 5 - assert first_reconcile[0] is False - assert first_reconcile[2]["busy"] is True - assert first_reconcile[3:] == ([100.0], ["complete"]) - assert len(first) == 12 - assert first[1] == expected_index - assert first[2] is False - assert first[8]["terminal_revision"] == terminal.terminal.revision - assert first[8]["delivery_attempt"] == 1 - assert second[8]["delivery_attempt"] == 2 - assert len(register_calls) == 1 - assert manager.snapshot(ref).acknowledged is False - - ack_store = acknowledge([second[8], None, None, None, None]) - assert ack_store["acknowledged"] is True assert manager.snapshot(ref).acknowledged is True assert manager.active_ref() is None + assert second_reconcile[0] is True + assert second_reconcile[2] == {"busy": False} - settled = reconcile( + # If that settling response is lost, the browser still believes it is + # busy and keeps the interval enabled. The next request converges to the + # same idle state without another terminal delivery. + repeated_settle = reconcile( 12, ref.as_dict(), None, @@ -212,11 +211,12 @@ def test_rf_terminal_replays_until_applied_marker_is_acknowledged(monkeypatch): first_reconcile[2], None, None, - [], - [], + rf_bar_ids, + rf_label_ids, + [first[8], None, None, None, None], ) - assert settled[0] is True - assert settled[2] == {"busy": False} + assert repeated_settle[0] is True + assert repeated_settle[2] == {"busy": False} def test_mds_finalization_stores_full_result_once_and_replays_small_index( @@ -276,8 +276,8 @@ def test_mds_finalization_stores_full_result_once_and_replays_small_index( assert "embedding" not in terminal.terminal.payload assert "data" not in terminal.terminal.payload - manager.snapshot_for_delivery(ref) - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) + manager.claim_terminal_delivery(ref) assert list(stored) == ["RF_002_MDS.tsv"] diff --git a/src/test/test_managed_analysis_jobs.py b/src/test/test_managed_analysis_jobs.py index faa432c..9801b75 100644 --- a/src/test/test_managed_analysis_jobs.py +++ b/src/test/test_managed_analysis_jobs.py @@ -120,7 +120,11 @@ def test_clade_worker_decodes_only_requested_columns(tmp_path): def test_rf_trace_terminal_replays_cached_render_until_ack(monkeypatch): - manager = JobManager(id_factory=lambda: "rf-trace-test-job") + now = [30.0] + manager = JobManager( + id_factory=lambda: "rf-trace-test-job", + clock=lambda: now[0], + ) db = SimpleNamespace( _trees=pd.DataFrame( { @@ -180,10 +184,11 @@ def test_rf_trace_terminal_replays_cached_render_until_ack(monkeypatch): diagnostics.register_diagnostics_callbacks, ) first_event = job_reconcile._terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) + now[0] += 1.0 second_event = job_reconcile._terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) first = render(first_event, ref.as_dict()) second = render(second_event, ref.as_dict()) @@ -214,7 +219,11 @@ def test_stage_four_submit_callbacks_have_matching_idle_output_shapes(): def test_clade_terminal_keeps_rows_server_side_and_keys_click_resolution( monkeypatch, ): - manager = JobManager(id_factory=lambda: "clade-test-job") + now = [40.0] + manager = JobManager( + id_factory=lambda: "clade-test-job", + clock=lambda: now[0], + ) fake_figure = SimpleNamespace(to_dict=lambda: {"data": [], "layout": {}}) monkeypatch.setattr(clade_explore, "job_manager", manager) monkeypatch.setattr(clade_explore, "add_log", lambda *_a, **_k: None) @@ -288,10 +297,11 @@ def test_clade_terminal_keeps_rows_server_side_and_keys_click_resolution( clade_explore.register_clade_explore_callbacks, ) first_event = job_reconcile._terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) + now[0] += 1.0 second_event = job_reconcile._terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) first = render(first_event, ref.as_dict()) second = render(second_event, ref.as_dict()) diff --git a/src/test/test_managed_compute_jobs.py b/src/test/test_managed_compute_jobs.py index 6ea8c28..797232c 100644 --- a/src/test/test_managed_compute_jobs.py +++ b/src/test/test_managed_compute_jobs.py @@ -53,7 +53,11 @@ def matches(): def test_pseudo_ess_submit_and_terminal_ui_replay_until_ack(monkeypatch): - manager = JobManager(id_factory=lambda: "pseudo-test-job") + now = [10.0] + manager = JobManager( + id_factory=lambda: "pseudo-test-job", + clock=lambda: now[0], + ) worker_calls = [] def worker(job_name, **kwargs): @@ -104,10 +108,11 @@ def worker(job_name, **kwargs): pseudo_ess_compute.register_pseudo_ess_compute_callbacks, ) first_event = job_reconcile._terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) + now[0] += 1.0 second_event = job_reconcile._terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) first = render(first_event, ref.as_dict()) second = render(second_event, ref.as_dict()) @@ -124,7 +129,11 @@ def worker(job_name, **kwargs): def test_consensus_finalizer_publishes_once_and_poll_only_replays(monkeypatch): - manager = JobManager(id_factory=lambda: "consensus-test-job") + now = [20.0] + manager = JobManager( + id_factory=lambda: "consensus-test-job", + clock=lambda: now[0], + ) cache_calls = [] register_calls = [] registry = [{"name": "RF_001_Between_consensus_tree_1"}] @@ -199,10 +208,11 @@ def register(**kwargs): consensus_tree_compute.register_consensus_tree_compute_callbacks, ) first_event = job_reconcile._terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) + now[0] += 1.0 second_event = job_reconcile._terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) first = render(first_event, ref.as_dict()) second = render(second_event, ref.as_dict()) @@ -268,7 +278,7 @@ def test_consensus_finalization_failure_reenables_origin_button(monkeypatch): consensus_tree_compute.register_consensus_tree_compute_callbacks, ) event = job_reconcile._terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) output = render(event, ref.as_dict()) assert output[3:5] == (False, False) diff --git a/src/test/test_terminal_delivery_faults.py b/src/test/test_terminal_delivery_faults.py index 91ee0a8..b2a7f1d 100644 --- a/src/test/test_terminal_delivery_faults.py +++ b/src/test/test_terminal_delivery_faults.py @@ -31,7 +31,11 @@ def _wait_for_terminal(manager, ref, timeout=1.0): def test_dropped_event_and_ui_responses_both_replay_before_receipt(): - manager = JobManager(id_factory=lambda: "fault-job") + now = [100.0] + manager = JobManager( + id_factory=lambda: "fault-job", + clock=lambda: now[0], + ) with ThreadPoolExecutor(max_workers=1) as executor: ref = manager.submit( executor, @@ -44,14 +48,16 @@ def test_dropped_event_and_ui_responses_both_replay_before_receipt(): # Attempt 1: the coordinator's response is dropped before the generic # terminal Store changes. No browser adapter runs and there is no receipt. dropped_event = _terminal_envelope( - manager.snapshot_for_delivery(ref) + manager.claim_terminal_delivery(ref) ) assert manager.active_ref() == ref + assert manager.claim_terminal_delivery(ref) is None # Attempt 2: the generic event lands, but the feature UI response is # dropped. Building a receipt server-side is not acknowledgement; it must - # reach the browser and trigger the acknowledgement request. - dropped_ui = _terminal_envelope(manager.snapshot_for_delivery(ref)) + # reach browser state and return on a later reconciliation request. + now[0] += 1.0 + dropped_ui = _terminal_envelope(manager.claim_terminal_delivery(ref)) assert terminal_event_for_job( dropped_ui, ref.as_dict(), @@ -63,7 +69,8 @@ def test_dropped_event_and_ui_responses_both_replay_before_receipt(): # Attempt 3 lands fully. Semantic terminal state/revision stayed stable; # only the retry counter advanced. - delivered = _terminal_envelope(manager.snapshot_for_delivery(ref)) + now[0] += 2.0 + delivered = _terminal_envelope(manager.claim_terminal_delivery(ref)) assert delivered["payload"] == dropped_event["payload"] assert delivered["terminal_revision"] == dropped_event["terminal_revision"] assert delivered["delivery_attempt"] == 3 @@ -73,21 +80,22 @@ def test_dropped_event_and_ui_responses_both_replay_before_receipt(): assert terminal.terminal.revision == receipt["terminal_revision"] -def test_dropped_ack_response_is_safe_after_server_acknowledgement(): +def test_dropped_settling_response_is_safe_after_server_acknowledgement(): manager = JobManager(id_factory=lambda: "ack-fault-job") with ThreadPoolExecutor(max_workers=1) as executor: ref = manager.submit(executor, "mds", lambda: None) _wait_for_terminal(manager, ref) - event = _terminal_envelope(manager.snapshot_for_delivery(ref)) + event = _terminal_envelope(manager.claim_terminal_delivery(ref)) receipt = terminal_delivery_marker(event) - # The browser's acknowledgement HTTP response may be lost after the server - # mutation. That is safe: terminal UI and receipt already landed together. + # The coordinator's settling response may be lost after the server + # mutation. That is safe: terminal UI and receipt already landed together, + # and a still-enabled browser interval will ask again. assert manager.acknowledge(ref, receipt["terminal_revision"]) - _dropped_ack_response = True + _dropped_settling_response = True assert manager.active_ref() is None - assert manager.snapshot_for_delivery(ref).delivery_attempt == 1 + assert manager.snapshot(ref).delivery_attempt == 1 def test_delayed_old_event_cannot_render_over_a_new_generation(): @@ -97,7 +105,7 @@ def test_delayed_old_event_cannot_render_over_a_new_generation(): old_ref = manager.submit(executor, "rf", lambda: None) old_terminal = _wait_for_terminal(manager, old_ref) old_event = _terminal_envelope( - manager.snapshot_for_delivery(old_ref) + manager.claim_terminal_delivery(old_ref) ) assert manager.acknowledge( old_ref, diff --git a/src/treetracer/background_jobs.py b/src/treetracer/background_jobs.py index 40d2fb1..c782950 100644 --- a/src/treetracer/background_jobs.py +++ b/src/treetracer/background_jobs.py @@ -178,6 +178,7 @@ class _JobRecord: terminal: TerminalEvent | None = None acknowledged: bool = False delivery_attempt: int = 0 + next_delivery_at: float | None = None terminal_revision: int = 0 started_at: float | None = None finished_at: float | None = None @@ -201,13 +202,28 @@ def __init__( *, single_active: bool = True, max_history: int = 32, + terminal_retry_delays: tuple[float, ...] = (1.0, 2.0, 4.0, 5.0), clock: Callable[[], float] = time.monotonic, id_factory: Callable[[], str] | None = None, ) -> None: if max_history < 1: raise ValueError("max_history must be at least 1") + retry_delays = tuple(float(delay) for delay in terminal_retry_delays) + if not retry_delays or any( + not math.isfinite(delay) or delay <= 0.0 + for delay in retry_delays + ): + raise ValueError( + "terminal_retry_delays must contain positive finite values" + ) + if any( + later < earlier + for earlier, later in zip(retry_delays, retry_delays[1:]) + ): + raise ValueError("terminal_retry_delays must be non-decreasing") self._single_active = single_active self._max_history = max_history + self._terminal_retry_delays = retry_delays self._clock = clock self._id_factory = id_factory or (lambda: uuid.uuid4().hex) self._lock = RLock() @@ -316,20 +332,42 @@ def snapshot(self, ref: JobRef) -> JobSnapshot | None: record = self._matching_record_locked(ref) return None if record is None else self._snapshot_locked(record) - def snapshot_for_delivery(self, ref: JobRef) -> JobSnapshot | None: - """Read state for a poll response and count terminal replays. + def claim_terminal_delivery(self, ref: JobRef) -> JobSnapshot | None: + """Claim a terminal delivery attempt only when its lease is due. + + The first attempt is immediate. Later attempts use the configured, + capped retry schedule. This gives the browser time to apply terminal + UI and return its receipt instead of flooding the callback graph with + a new envelope on every progress-poll tick. - A changing delivery attempt lets a ``dcc.Store`` retrigger browser - reconciliation even though the semantic terminal event and its - revision remain stable. + ``None`` means that the record is missing, non-terminal, already + acknowledged, or still inside the current delivery lease. """ with self._lock: record = self._matching_record_locked(ref) - if record is None: + if ( + record is None + or record.terminal is None + or record.acknowledged + ): + return None + + now = self._clock() + if ( + record.next_delivery_at is not None + and now < record.next_delivery_at + ): return None - if record.terminal is not None and not record.acknowledged: - record.delivery_attempt += 1 + + record.delivery_attempt += 1 + delay_index = min( + record.delivery_attempt - 1, + len(self._terminal_retry_delays) - 1, + ) + record.next_delivery_at = ( + now + self._terminal_retry_delays[delay_index] + ) return self._snapshot_locked(record) def active_ref(self) -> JobRef | None: @@ -383,6 +421,7 @@ def acknowledge( if record.terminal.revision != int(terminal_revision): return False record.acknowledged = True + record.next_delivery_at = None self._prune_history_locked() return True @@ -615,6 +654,9 @@ def _set_terminal_locked( record.state = state record.future = None record.finished_at = self._clock() + record.acknowledged = False + record.delivery_attempt = 0 + record.next_delivery_at = None record.terminal_revision += 1 record.terminal = TerminalEvent( state=state, diff --git a/src/treetracer/callbacks/compute.py b/src/treetracer/callbacks/compute.py index 2357af7..1f7eb06 100644 --- a/src/treetracer/callbacks/compute.py +++ b/src/treetracer/callbacks/compute.py @@ -778,8 +778,8 @@ def handle_compute_mds(n_clicks, selected_distmat): Output("export-mds-button", "disabled"), Output("plot-config-store", "data", allow_duplicate=True), # Shared notification + a dedicated receipt. The receipt lands in the - # same browser response as the terminal UI and is consumed by the one - # acknowledgement sink in ``job_reconcile.py``. + # same browser response as the terminal UI and returns as State on the + # next central reconciliation poll in ``job_reconcile.py``. Output("notifications-container", "children", allow_duplicate=True), Output( {"type": "compute-terminal-receipt", "kind": "rf-mds"}, diff --git a/src/treetracer/callbacks/job_reconcile.py b/src/treetracer/callbacks/job_reconcile.py index 0c56226..574bc86 100644 --- a/src/treetracer/callbacks/job_reconcile.py +++ b/src/treetracer/callbacks/job_reconcile.py @@ -9,12 +9,14 @@ * terminal state is copied into one small, replayable browser event; * job-specific presentation callbacks render that event and write a dedicated receipt in the same response as their UI; -* one acknowledgement callback consumes those receipts. +* the next interval request carries those receipts back as callback state. The interval remains enabled until a matching receipt has acknowledged the -sticky terminal event. A lost terminal-event response, presentation response, -or acknowledgement request therefore causes another delivery attempt instead -of a permanently stuck loading state. +sticky terminal event. Terminal retries use a capped lease schedule rather +than firing on every poll, so a slow browser gets an uncontested window to +apply UI and return its receipt. A lost terminal-event response, presentation +response, or settling response therefore causes a later delivery/settling +attempt instead of a permanently stuck loading state. """ from __future__ import annotations @@ -276,6 +278,10 @@ def register_job_reconciliation_callbacks(): State("mds-progress-path", "data"), State({"type": "compute-progress-bar", "which": ALL}, "id"), State({"type": "compute-progress-label", "which": ALL}, "id"), + State( + {"type": "compute-terminal-receipt", "kind": ALL}, + "data", + ), prevent_initial_call=True, ) def reconcile_compute_job( @@ -291,8 +297,27 @@ def reconcile_compute_job( mds_progress_path, progress_bar_ids, progress_label_ids, + receipts, ): """Own polling, terminal delivery, progress, and the global gate.""" + # A receipt exists in browser state only after its feature callback's + # terminal UI response was applied. Piggyback acknowledgement on this + # already-running poll instead of scheduling another callback in the + # rapidly updating terminal chain. Exact identity/revision matching + # makes stale receipts harmless. + receipt_ref, receipt = _matching_active_receipt(receipts) + if receipt_ref is not None and receipt is not None: + if job_manager.acknowledge( + receipt_ref, + receipt["terminal_revision"], + ): + _wlog( + "[parent] reconcile_compute_job: acknowledged " + f"job={receipt_ref.job_id}/{receipt_ref.kind}/generation-" + f"{receipt_ref.generation}, revision=" + f"{receipt['terminal_revision']}" + ) + active = job_manager.active_ref() busy = _busy_update(active, current_busy) if active is None: @@ -325,29 +350,37 @@ def reconcile_compute_job( terminal_event = no_update if snapshot.terminal is not None: - delivery = job_manager.snapshot_for_delivery(active) + delivery = job_manager.claim_terminal_delivery(active) if delivery is None: - progress_outputs = _dynamic_progress_outputs( - progress_bar_ids, - progress_label_ids, + # ``None`` normally means the previous delivery lease is + # still active. Clear Data can concurrently remove the record, + # so re-check ownership before deciding to keep polling. + current_active = job_manager.active_ref() + if current_active != active: + progress_outputs = _dynamic_progress_outputs( + progress_bar_ids, + progress_label_ids, + ) + return ( + current_active is None, + no_update, + _busy_update(current_active, current_busy), + *progress_outputs, + ) + else: + terminal_event = _terminal_envelope(delivery) + _wlog( + "[parent] reconcile_compute_job: delivering " + f"job={active.job_id}/{active.kind}/generation-" + f"{active.generation}, state={delivery.state.value}, " + f"attempt={delivery.delivery_attempt}" ) - return ( - True, - no_update, - _busy_update(None, current_busy), - *progress_outputs, - ) - terminal_event = _terminal_envelope(delivery) + + terminal_progress = snapshot if delivery is None else delivery if active.kind == "rf": - rf_progress = _terminal_progress(delivery) + rf_progress = _terminal_progress(terminal_progress) elif active.kind == "mds": - mds_progress = _terminal_progress(delivery) - _wlog( - "[parent] reconcile_compute_job: delivering " - f"job={active.job_id}/{active.kind}/generation-" - f"{active.generation}, state={delivery.state.value}, " - f"attempt={delivery.delivery_attempt}" - ) + mds_progress = _terminal_progress(terminal_progress) elif active.kind == "rf": rf_progress = _read_rf_progress(active, rf_progress_path) elif active.kind == "mds": @@ -359,9 +392,9 @@ def reconcile_compute_job( f"{active.generation}, state={snapshot.state.value}" ) - # Polling remains enabled through terminal presentation. The receipt - # callback acknowledges server state; the next tick then takes the - # active=None branch and is the sole path that disables this interval. + # Polling remains enabled through terminal presentation. A later tick + # carries the browser receipt as State, acknowledges it above, and + # takes the active=None branch in that same response. progress_outputs = _dynamic_progress_outputs( progress_bar_ids, progress_label_ids, @@ -374,22 +407,3 @@ def reconcile_compute_job( busy, *progress_outputs, ) - - @callback( - Output("compute-job-ack-store", "data"), - Input( - {"type": "compute-terminal-receipt", "kind": ALL}, - "data", - ), - prevent_initial_call=True, - ) - def acknowledge_terminal_receipt(receipts): - """Acknowledge only the receipt for the authoritative active job.""" - active, receipt = _matching_active_receipt(receipts) - if active is None or receipt is None: - return no_update - acknowledged = job_manager.acknowledge( - active, - receipt["terminal_revision"], - ) - return {**receipt, "acknowledged": acknowledged} diff --git a/src/treetracer/ui/navbar.py b/src/treetracer/ui/navbar.py index 57db801..f3c77f1 100644 --- a/src/treetracer/ui/navbar.py +++ b/src/treetracer/ui/navbar.py @@ -113,7 +113,6 @@ def add_navbar(): storage_type="memory", data={"busy": False}, ), - dcc.Store(id="compute-job-ack-store", storage_type="memory"), dcc.Store( id={ "type": "compute-terminal-receipt", From cc055291ea94e890fc4a43eff763aa93e85fe1a7 Mon Sep 17 00:00:00 2001 From: Sam Hong <20244918+hongsamL@users.noreply.github.com> Date: Mon, 3 Aug 2026 20:03:32 +0200 Subject: [PATCH 9/9] bump version --- pyproject.toml | 4 ++-- src/treetracer/ui/main_body.py | 2 +- uv.lock | 2 +- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 23f8463..c439134 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "treetracer" -version = "0.95rc" +version = "0.96" description = "An app to visualize phylogenetic tree topology convergence" readme = "README.md" authors = [ @@ -75,7 +75,7 @@ package = [ [tool.briefcase] project_name = "TreeTracer" bundle = "community.beast" -version = "0.95rc" +version = "0.96" url = "https://github.com/beast-dev/treetracer" license = "GPL-3.0-or-later" author = "Sam Hong" diff --git a/src/treetracer/ui/main_body.py b/src/treetracer/ui/main_body.py index 5b2a9aa..0fceefc 100644 --- a/src/treetracer/ui/main_body.py +++ b/src/treetracer/ui/main_body.py @@ -21,7 +21,7 @@ def _add_about_modal(): children=[ dmc.Stack( [ - dmc.Title("TreeTracer v0.95rc", order=3), + dmc.Title("TreeTracer v0.96", order=3), dmc.Text( "TreeTracer is a diagnostic tool used to visualize convergence of tree topologies" ), diff --git a/uv.lock b/uv.lock index 35879c3..5ee7eaf 100644 --- a/uv.lock +++ b/uv.lock @@ -1186,7 +1186,7 @@ wheels = [ [[package]] name = "treetracer" -version = "0.95rc0" +version = "0.96" source = { editable = "." } dependencies = [ { name = "dash" },