diff --git a/tests/test_plaud_download.py b/tests/test_plaud_download.py new file mode 100644 index 000000000..c5b9413b3 --- /dev/null +++ b/tests/test_plaud_download.py @@ -0,0 +1,108 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +import logging +from unittest.mock import Mock + +import pytest +import requests + + +class _Response: + status_code = 200 + headers = {"Content-Length": "0"} + text = "" + + def __init__(self, chunks=None, exc=None): + self._chunks = chunks or [] + self._exc = exc + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc, tb): + return False + + def iter_content(self, chunk_size): + if self._exc: + raise self._exc + yield from self._chunks + + +@pytest.mark.timeout(5) +def test_download_to_file_returns_false_on_read_timeout(tmp_path, caplog): + from think.importers.plaud import download_to_file + + session = Mock() + session.get.return_value = _Response(exc=requests.exceptions.ReadTimeout("stalled")) + dest_path = tmp_path / "recording.opus" + + caplog.set_level(logging.WARNING) + + assert download_to_file(session, "https://example.test/file", dest_path) is False + assert not dest_path.exists() + assert "Plaud download for recording.opus failed: stalled" in caplog.text + + +def test_download_to_file_calls_progress_cb_throttled(tmp_path, monkeypatch): + from think.importers import plaud + + session = Mock() + session.get.return_value = _Response(chunks=[b"x"] * 12) + progress_cb = Mock() + ticks = iter(range(13)) + monkeypatch.setattr(plaud.time, "monotonic", lambda: next(ticks)) + + dest_path = tmp_path / "recording.opus" + + assert ( + plaud.download_to_file( + session, + "https://example.test/file", + dest_path, + progress_cb=progress_cb, + ) + is True + ) + assert dest_path.read_bytes() == b"x" * 12 + assert progress_cb.call_count == 2 + + +def test_sync_inactivity_timer_trips_when_progress_stops(tmp_path, monkeypatch, caplog): + from think.importers import plaud + + files = [ + { + "id": "file1", + "filename": "One", + "fullname": "one.opus", + "filesize": 10, + "start_time": 1737000000000, + "duration": 60000, + }, + { + "id": "file2", + "filename": "Two", + "fullname": "two.opus", + "filesize": 10, + "start_time": 1737000300000, + "duration": 60000, + }, + ] + + monkeypatch.setenv("PLAUD_ACCESS_TOKEN", "test-token") + monkeypatch.setattr(plaud, "SYNC_INACTIVITY_TIMEOUT", 1) + monkeypatch.setattr(plaud, "list_files", lambda _session, _token: files) + monkeypatch.setattr( + plaud, "get_temp_url", lambda *_args: "https://example.test/file" + ) + monkeypatch.setattr(plaud, "download_to_file", lambda *_args, **_kwargs: False) + ticks = iter([0.0, 0.5, 2.0]) + monkeypatch.setattr(plaud.time, "monotonic", lambda: next(ticks)) + + caplog.set_level(logging.WARNING) + + result = plaud.PlaudBackend().sync(tmp_path, dry_run=False) + + assert any("Sync stalled" in error for error in result["errors"]) + assert "Sync stalled" in caplog.text diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index 524f6b88e..40558bd58 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -158,6 +158,116 @@ class TestLoadConfig: assert load_config() == {} + def test_max_runtime_valid_string_round_trips(self, journal_path): + # D-E/D-F: assert the accepted Plaud cap via test-local config, + # leaving the synthetic fixture schedule minimal. + _write_config( + journal_path, + { + "sync:plaud": { + "cmd": ["sol", "import", "--sync", "plaud"], + "every": "hourly", + "max_runtime": "30m", + }, + }, + ) + from think.scheduler import load_config + + entries = load_config() + assert entries["sync:plaud"]["max_runtime"] == 1800 + + def test_max_runtime_valid_int_round_trips(self, journal_path): + _write_config( + journal_path, + { + "sync:plaud": { + "cmd": ["sol", "import", "--sync", "plaud"], + "every": "hourly", + "max_runtime": 1800, + }, + }, + ) + from think.scheduler import load_config + + entries = load_config() + assert entries["sync:plaud"]["max_runtime"] == 1800 + + def test_max_runtime_invalid_negative_logged_and_dropped( + self, journal_path, caplog + ): + _write_config( + journal_path, + { + "sync:plaud": { + "cmd": ["sol", "import", "--sync", "plaud"], + "every": "hourly", + "max_runtime": -5, + }, + }, + ) + from think.scheduler import load_config + + entries = load_config() + assert "max_runtime" not in entries["sync:plaud"] + assert "Schedule 'sync:plaud': invalid max_runtime -5" in caplog.text + + def test_max_runtime_invalid_garbage_logged_and_dropped(self, journal_path, caplog): + _write_config( + journal_path, + { + "sync:plaud": { + "cmd": ["sol", "import", "--sync", "plaud"], + "every": "hourly", + "max_runtime": "garbage", + }, + }, + ) + from think.scheduler import load_config + + entries = load_config() + assert "max_runtime" not in entries["sync:plaud"] + assert "Schedule 'sync:plaud': invalid max_runtime 'garbage'" in caplog.text + + def test_max_runtime_invalid_type_logged_and_dropped(self, journal_path, caplog): + _write_config( + journal_path, + { + "sync:plaud": { + "cmd": ["sol", "import", "--sync", "plaud"], + "every": "hourly", + "max_runtime": [1, 2], + }, + }, + ) + from think.scheduler import load_config + + entries = load_config() + assert "max_runtime" not in entries["sync:plaud"] + assert "Schedule 'sync:plaud': invalid max_runtime [1, 2]" in caplog.text + + def test_collect_runtime_caps_returns_only_capped_entries(self, journal_path): + _write_config( + journal_path, + { + "sync:plaud": { + "cmd": ["sol", "import", "--sync", "plaud"], + "every": "hourly", + "max_runtime": "30m", + }, + "heartbeat": { + "cmd": ["sol", "heartbeat"], + "every": "daily", + }, + }, + ) + import think.scheduler as mod + + mod.init(Mock()) + + assert mod.collect_runtime_caps() == [ + (["sol", "import", "--sync", "plaud"], 1800) + ] + # --------------------------------------------------------------------------- # load_state / save_state diff --git a/tests/test_supervisor.py b/tests/test_supervisor.py index 841777343..77760fe5a 100644 --- a/tests/test_supervisor.py +++ b/tests/test_supervisor.py @@ -4,6 +4,7 @@ import importlib import io import os +import signal import subprocess import sys from unittest.mock import MagicMock @@ -551,6 +552,89 @@ def test_stale_queue_detected_on_submit(monkeypatch): assert len(mod._task_queue._queues["indexer"]) == 1 +class _TaskProcessStub: + def __init__(self): + self.send_signal = MagicMock() + self.kill = MagicMock() + + +class _TaskManagedStub: + def __init__(self, *, cmd, start_time=100.0): + self.cmd = cmd + self.start_time = start_time + self.process = _TaskProcessStub() + + +def test_taskqueue_set_cap_records_cap(): + mod = importlib.import_module("think.supervisor") + queue = mod.TaskQueue(on_queue_change=None) + + queue.set_cap("import", 1800) + + assert queue._caps["import"] == 1800 + + +def test_enforce_deadlines_sends_sigterm_when_elapsed_exceeds_cap(caplog): + mod = importlib.import_module("think.supervisor") + queue = mod.TaskQueue(on_queue_change=None) + managed = _TaskManagedStub( + cmd=["sol", "import", "--sync", "plaud", "--save"], + start_time=100.0, + ) + queue._active["ref-1"] = managed + queue.set_cap("import", 50) + + caplog.set_level("WARNING") + queue.enforce_deadlines(200.0) + + managed.process.send_signal.assert_called_once_with(signal.SIGTERM) + assert queue._terminations["ref-1"] == 200.0 + assert ( + "Task import (cmd=sol import --sync plaud --save, ref=ref-1) exceeded " + "max_runtime of 50s (elapsed=100s); sending SIGTERM" + ) in caplog.text + + +def test_enforce_deadlines_escalates_to_sigkill_after_15s(caplog): + mod = importlib.import_module("think.supervisor") + queue = mod.TaskQueue(on_queue_change=None) + managed = _TaskManagedStub(cmd=["sol", "import"], start_time=100.0) + queue._active["ref-1"] = managed + queue._terminations["ref-1"] = 200.0 + + caplog.set_level("WARNING") + queue.enforce_deadlines(216.0) + + managed.process.kill.assert_called_once_with() + assert queue._terminations["ref-1"] == 0.0 + assert ( + "Task import (ref=ref-1) did not exit 15s after SIGTERM; sending SIGKILL" + ) in caplog.text + + +def test_enforce_deadlines_clears_termination_state_when_ref_exits(): + mod = importlib.import_module("think.supervisor") + queue = mod.TaskQueue(on_queue_change=None) + queue._terminations["ref-1"] = 200.0 + + queue.enforce_deadlines(216.0) + + assert "ref-1" not in queue._terminations + + +def test_enforce_deadlines_noop_when_no_cap(): + mod = importlib.import_module("think.supervisor") + queue = mod.TaskQueue(on_queue_change=None) + managed = _TaskManagedStub(cmd=["sol", "import"], start_time=100.0) + queue._active["ref-1"] = managed + + queue.enforce_deadlines(10_000.0) + + managed.process.send_signal.assert_not_called() + managed.process.kill.assert_not_called() + assert queue._terminations == {} + + def test_supervisor_singleton_lock_acquired(tmp_path, monkeypatch): mod = importlib.reload(importlib.import_module("think.supervisor")) diff --git a/tests/test_utils.py b/tests/test_utils.py new file mode 100644 index 000000000..6c8bf7279 --- /dev/null +++ b/tests/test_utils.py @@ -0,0 +1,36 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +import pytest + +from think.utils import parse_duration_seconds + + +@pytest.mark.parametrize( + ("spec", "expected"), + [ + (30, 30), + ("45s", 45), + ("30m", 1800), + ("1h", 3600), + ], +) +def test_parse_duration_seconds_valid(spec, expected): + assert parse_duration_seconds(spec) == expected + + +@pytest.mark.parametrize( + "spec", + [ + 0, + -5, + "garbage", + "5x", + "30 m", + None, + [], + ], +) +def test_parse_duration_seconds_invalid(spec): + with pytest.raises(ValueError, match="invalid duration"): + parse_duration_seconds(spec) diff --git a/think/importers/plaud.py b/think/importers/plaud.py index d7afc77a2..0148cd5d1 100644 --- a/think/importers/plaud.py +++ b/think/importers/plaud.py @@ -13,7 +13,7 @@ import subprocess import tempfile import time from pathlib import Path -from typing import Any, Dict, List, Optional +from typing import Any, Callable, Dict, List, Optional import requests from requests.adapters import HTTPAdapter @@ -25,6 +25,7 @@ API_BASE = "https://api.plaud.ai" # Skip recordings shorter than this (milliseconds) MIN_DURATION_MS = 30_000 +SYNC_INACTIVITY_TIMEOUT = 3600 def make_session() -> requests.Session: @@ -163,39 +164,59 @@ def sanitize_filename(filename: str) -> str: def download_to_file( - session: requests.Session, url: str, dest_path: pathlib.Path + session: requests.Session, + url: str, + dest_path: pathlib.Path, + progress_cb: Callable[[], None] | None = None, ) -> bool: """Stream-download URL to dest_path atomically.""" dest_path.parent.mkdir(parents=True, exist_ok=True) - with session.get(url, stream=True, timeout=60) as r: - if r.status_code != 200: - logger.warning( - "[%s] Download error %s: %s", - dest_path.stem, - r.status_code, - r.text[:200], - ) - return False - total = int(r.headers.get("Content-Length", "0")) or None - # Write to a temp file then atomically move - with tempfile.NamedTemporaryFile( - dir=str(dest_path.parent), delete=False - ) as tmp: - tmp_path = pathlib.Path(tmp.name) - try: - downloaded = 0 - for chunk in r.iter_content(chunk_size=1024 * 256): - if not chunk: - continue - tmp.write(chunk) - downloaded += len(chunk) - tmp.flush() - os.fsync(tmp.fileno()) - except Exception as e: - tmp.close() - tmp_path.unlink(missing_ok=True) - logger.warning("[%s] Error while writing file: %s", dest_path.stem, e) + tmp_path: pathlib.Path | None = None + try: + # Design §4-§5: fail stalled streaming reads and refresh outer progress. + with session.get(url, stream=True, timeout=(30, 45)) as r: + if r.status_code != 200: + logger.warning( + "[%s] Download error %s: %s", + dest_path.stem, + r.status_code, + r.text[:200], + ) return False + total = int(r.headers.get("Content-Length", "0")) or None + # Write to a temp file then atomically move + with tempfile.NamedTemporaryFile( + dir=str(dest_path.parent), delete=False + ) as tmp: + tmp_path = pathlib.Path(tmp.name) + try: + downloaded = 0 + last_refresh = time.monotonic() + for chunk in r.iter_content(chunk_size=1024 * 256): + if not chunk: + continue + tmp.write(chunk) + downloaded += len(chunk) + now = time.monotonic() + if progress_cb is not None and now - last_refresh >= 5: + progress_cb() + last_refresh = now + tmp.flush() + os.fsync(tmp.fileno()) + except requests.exceptions.RequestException: + raise + except (IOError, OSError) as e: + tmp.close() + tmp_path.unlink(missing_ok=True) + logger.warning( + "[%s] Error while writing file: %s", dest_path.stem, e + ) + return False + except requests.exceptions.RequestException as exc: + if tmp_path is not None: + tmp_path.unlink(missing_ok=True) + logger.warning("Plaud download for %s failed: %s", dest_path.name, exc) + return False tmp_path.replace(dest_path) size_info = f" ({total} bytes)" if total else "" @@ -433,9 +454,13 @@ class PlaudBackend: # within this window. The inner import process has its own # 600s inactivity timeout for stall detection, so this is a # generous outer safety net. - sync_timeout = 3600 + sync_timeout = SYNC_INACTIVITY_TIMEOUT last_completed = time.monotonic() + def _refresh() -> None: + nonlocal last_completed + last_completed = time.monotonic() + for idx, (file_id, info) in enumerate(to_process, 1): # Check inactivity timeout (time since last file completed) idle_duration = time.monotonic() - last_completed @@ -496,7 +521,9 @@ class PlaudBackend: errors.append(msg) continue - if not download_to_file(session, temp_url, dest_path): + if not download_to_file( + session, temp_url, dest_path, progress_cb=_refresh + ): msg = f"{filename}: download failed" logger.warning(" FAILED — %s", msg) errors.append(msg) diff --git a/think/scheduler.py b/think/scheduler.py index ab95fa4c0..afa8b1efe 100644 --- a/think/scheduler.py +++ b/think/scheduler.py @@ -22,7 +22,13 @@ from datetime import datetime, timedelta from pathlib import Path from typing import Any -from think.utils import get_journal, now_ms, require_solstone, setup_cli +from think.utils import ( + get_journal, + now_ms, + parse_duration_seconds, + require_solstone, + setup_cli, +) logger = logging.getLogger(__name__) @@ -125,7 +131,21 @@ def load_config() -> dict[str, dict[str, Any]]: if not entry.get("enabled", True): continue - entries[name] = {"cmd": cmd, "every": every} + validated = {"cmd": cmd, "every": every} + max_runtime = entry.get("max_runtime") + if max_runtime is not None: + # D-C / design §2: preserve caps for TaskQueue registration, + # not as extra supervisor.request payload fields. + try: + validated["max_runtime"] = parse_duration_seconds(max_runtime) + except ValueError: + logger.warning( + "Schedule '%s': invalid max_runtime %r, dropping cap", + name, + max_runtime, + ) + + entries[name] = validated return entries @@ -305,6 +325,16 @@ def init(callosum: Any) -> None: logger.info("Scheduler initialized (no schedules configured)") +def collect_runtime_caps() -> list[tuple[list[str], int]]: + """Return configured task runtime caps from loaded schedule entries.""" + caps: list[tuple[list[str], int]] = [] + for entry in _entries.values(): + max_runtime = entry.get("max_runtime") + if max_runtime is not None: + caps.append((list(entry["cmd"]), max_runtime)) + return caps + + def register_defaults() -> None: """Ensure built-in default schedules exist in the config file. diff --git a/think/supervisor.py b/think/supervisor.py index e29f3ae40..a74c5e304 100644 --- a/think/supervisor.py +++ b/think/supervisor.py @@ -197,6 +197,8 @@ class TaskQueue: ] = {} # command_name -> {"ref": str, "thread": Thread} self._queues: dict[str, list] = {} # command_name -> list of {refs, cmd} dicts self._active: dict[str, RunnerManagedProcess] = {} # ref -> process + self._caps: dict[str, int] = {} + self._terminations: dict[str, float] = {} self._pending: list[dict] = [] self._ready = ready self._lock = threading.Lock() @@ -319,6 +321,63 @@ class TaskQueue: return None + def set_cap(self, cmd_name: str, seconds: int) -> None: + """Set a max runtime cap in seconds for a queued command name.""" + with self._lock: + self._caps[cmd_name] = seconds + + def enforce_deadlines(self, now: float) -> None: + """Enforce configured task runtime caps without blocking the supervisor tick.""" + with self._lock: + active_refs = set(self._active) + for ref in list(self._terminations): + if ref not in active_refs: + self._terminations.pop(ref, None) + + for ref, managed in list(self._active.items()): + cmd_name = self.get_command_name(managed.cmd) + termination_started = self._terminations.get(ref) + if termination_started is not None: + if termination_started <= 0: + continue + if now - termination_started >= 15: + logging.warning( + "Task %s (ref=%s) did not exit 15s after SIGTERM; " + "sending SIGKILL", + cmd_name, + ref, + ) + try: + managed.process.kill() + except (ProcessLookupError, OSError): + pass + self._terminations[ref] = 0.0 + continue + + cap = self._caps.get(cmd_name) + if not cap: + continue + + elapsed = now - managed.start_time + if elapsed <= cap: + continue + + elapsed_seconds = int(elapsed) + logging.warning( + "Task %s (cmd=%s, ref=%s) exceeded max_runtime of %ds " + "(elapsed=%ds); sending SIGTERM", + cmd_name, + " ".join(managed.cmd), + ref, + cap, + elapsed_seconds, + ) + try: + managed.process.send_signal(signal.SIGTERM) + self._terminations[ref] = now + except (ProcessLookupError, OSError): + pass + def set_ready(self) -> None: """Allow buffered tasks to start dispatching through the normal queue path.""" with self._lock: @@ -365,7 +424,8 @@ class TaskQueue: managed = RunnerManagedProcess.spawn( cmd, ref=primary_ref, callosum=callosum, day=day ) - self._active[primary_ref] = managed + with self._lock: + self._active[primary_ref] = managed callosum.emit( "supervisor", @@ -415,7 +475,8 @@ class TaskQueue: managed.cleanup() except Exception: logging.exception(f"Task {cmd_name} ({primary_ref}): cleanup failed") - self._active.pop(primary_ref, None) + with self._lock: + self._active.pop(primary_ref, None) try: callosum.stop() except Exception: @@ -1383,6 +1444,9 @@ async def supervise( logging.error(f"Failed to kill {service}: {e}") # Don't delete here - let handle_runner_exits clean up + if _task_queue: + _task_queue.enforce_deadlines(time.time()) + # Check for runner exits first (immediate alert) if procs: await handle_runner_exits(procs) @@ -1660,6 +1724,15 @@ def main() -> None: if schedule_enabled and _supervisor_callosum: scheduler.init(_supervisor_callosum) scheduler.register_defaults() + if _task_queue: + for cmd, seconds in scheduler.collect_runtime_caps(): + cmd_name = TaskQueue.get_command_name(cmd) + _task_queue.set_cap(cmd_name, seconds) + logging.info( + "Registered max_runtime cap for %s: %ss", + cmd_name, + seconds, + ) routines.init(_supervisor_callosum) if _task_queue: diff --git a/think/utils.py b/think/utils.py index 1bd0d6858..fdc1376ae 100644 --- a/think/utils.py +++ b/think/utils.py @@ -153,6 +153,24 @@ def get_journal() -> str: return path +def parse_duration_seconds(spec) -> int: + # D-D: shared parser for scheduler max_runtime values. + if isinstance(spec, int) and not isinstance(spec, bool): + if spec > 0: + return spec + raise ValueError(f"invalid duration: {spec!r}") + + if isinstance(spec, str): + match = re.fullmatch(r"(\d+)([smh])", spec) + if match: + amount = int(match.group(1)) + if amount <= 0: + raise ValueError(f"invalid duration: {spec!r}") + return amount * {"s": 1, "m": 60, "h": 3600}[match.group(2)] + + raise ValueError(f"invalid duration: {spec!r}") + + def resolve_journal_path(journal: str | Path, rel: str) -> Path: """Resolve a chronicle-free journal-relative path to its on-disk location.""" if not rel: