"""ml.executor — bounded, timed offload for the synchronous ML fits (analytical tools round). The analytical tools (analyze_forecast / importance / anomaly / cluster / significance) call blocking, CPU-heavy fit routines (scikit-learn, statsmodels). Their compute functions are SYNCHRONOUS by design — like every other `analyze_*` fn (see the STATUS note in `analytics/temporal.py`) — but the invoker runs inside the async slow path. Calling a multi-second fit directly on the event loop would stall every other request in the process. This offloads a fit to a DEDICATED, bounded thread pool with a per-fit timeout — the same pattern the DB executor uses (`src/query/executor/db.py`, F-5): * Dedicated pool — a slow/hung fit occupies an ML-only thread, so its blast radius is other analytics, not the whole service (the shared default `asyncio.to_thread` pool also serves the tabular Parquet loader). * Semaphore — bounds how many fits run at once, so a burst of heavy questions QUEUES instead of oversubscribing the CPU (kept equal to the pool size). * Per-fit timeout — `asyncio.wait_for` stops US waiting. NOTE (the db.py caveat): it cancels the awaiting COROUTINE, not the worker thread, which is not cancellable and runs to completion; the pool bound caps how many such abandoned threads can pile up. sklearn / statsmodels release the GIL inside their C/Fortran kernels, so threads still give real parallelism here. The invoker's never-throw seam (§8.4) catches whatever this raises — including `MLFitTimeoutError` — and returns a `ToolOutput(kind="error")`, so a slow fit degrades that one analysis, never the whole turn. """ from __future__ import annotations import asyncio import functools import os from collections.abc import Callable from concurrent.futures import ThreadPoolExecutor from typing import Any, TypeVar from src.middlewares.logging import get_logger logger = get_logger("ml_executor") _T = TypeVar("_T") def _int_env(name: str, default: int) -> int: """A positive int from the environment, falling back to `default` if unset/bad.""" raw = os.getenv(name) if raw is None: return default try: return max(1, int(raw)) except ValueError: logger.warning("ignoring non-integer env override", var=name, value=raw) return default # ML fits are CPU-bound; more threads than cores just thrash. Small by default, # overridable for a beefier host (or trimmed on a constrained HF Space). _MAX_ML_WORKERS = _int_env("ML_MAX_WORKERS", min(4, os.cpu_count() or 1)) _FIT_TIMEOUT_SECONDS = float(_int_env("ML_FIT_TIMEOUT_SECONDS", 30)) _ML_THREAD_POOL = ThreadPoolExecutor( max_workers=_MAX_ML_WORKERS, thread_name_prefix="mlfit" ) # Equal to the pool size: excess fits wait here rather than piling futures onto a # saturated pool. Constructed without a running loop (fine on 3.10+: asyncio # primitives bind to the loop lazily, on first await). _ML_SEMAPHORE = asyncio.Semaphore(_MAX_ML_WORKERS) class MLFitTimeoutError(TimeoutError): """A single ML fit exceeded its per-fit timeout. Subclasses `TimeoutError` so callers may catch either; the invoker's never-throw seam turns it into a `ToolOutput(kind="error")`. """ async def run_fit( fn: Callable[..., _T], *args: Any, label: str = "", timeout: float | None = None, **kwargs: Any, ) -> _T: """Run a blocking fit ``fn(*args, **kwargs)`` off the event loop — bounded + timed. Args: fn: the synchronous compute/fit callable. *args, **kwargs: forwarded to `fn` unchanged. label: names the fit in logs and the timeout message (pass the tool name). timeout: per-fit budget in seconds; defaults to `_FIT_TIMEOUT_SECONDS` (env `ML_FIT_TIMEOUT_SECONDS`). Returns: Whatever `fn` returns. Raises: MLFitTimeoutError: the fit did not finish within `timeout`. Exception: any error from `fn` is re-raised unchanged (the invoker's never-throw seam wraps it). """ budget = _FIT_TIMEOUT_SECONDS if timeout is None else timeout loop = asyncio.get_running_loop() call = functools.partial(fn, *args, **kwargs) async with _ML_SEMAPHORE: try: return await asyncio.wait_for( loop.run_in_executor(_ML_THREAD_POOL, call), budget ) except TimeoutError: # asyncio.wait_for raises builtin TimeoutError (3.11+) name = label or getattr(fn, "__name__", "fit") logger.warning("ml fit timed out", label=name, timeout=budget) raise MLFitTimeoutError(f"{name} exceeded {budget:g}s") from None