File size: 4,695 Bytes
e201ae4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
"""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