diff --git a/backend/watcher.py b/backend/watcher.py index e7a3c2c..9fdab11 100644 --- a/backend/watcher.py +++ b/backend/watcher.py @@ -5,12 +5,17 @@ import os import threading import time from dataclasses import dataclass +from pathlib import Path from typing import Any +try: + import smbclient # type: ignore[import-not-found] +except ModuleNotFoundError: + smbclient = None from watchdog.events import FileSystemEvent, FileSystemEventHandler, FileSystemMovedEvent from watchdog.observers.polling import PollingObserver -from scanner import resolve_runtime_path, scan_results, path_exists_with_smb +from scanner import resolve_runtime_path, scan_results LOGGER = logging.getLogger("scheduler.watcher") @@ -39,6 +44,49 @@ class _ResultDirEventHandler(FileSystemEventHandler): self._on_new_directory(event.dest_path) +class _SmbPollerThread(threading.Thread): + """Poll a UNC results directory and trigger a scan when entries change.""" + + def __init__(self, unc_path: str, on_change, interval_seconds: float = 5.0): + super().__init__(daemon=True, name=f"smb-poller-{unc_path}") + self._path = unc_path.rstrip("\\") + self._on_change = on_change + self._interval_seconds = max(1.0, float(interval_seconds)) + self._stop_evt = threading.Event() + + def stop(self) -> None: + self._stop_evt.set() + + def run(self) -> None: + prev = self._snapshot() + while not self._stop_evt.wait(self._interval_seconds): + try: + curr = self._snapshot() + added = curr - prev + removed = prev - curr + if added or removed: + LOGGER.info( + "SMB directory change detected (%s): +%d -%d", + self._path, + len(added), + len(removed), + ) + self._on_change("smb-directory-change") + prev = curr + except Exception: + LOGGER.exception("SMB poll error for %s", self._path) + + def _snapshot(self) -> set[str]: + if smbclient is None: + LOGGER.error("smbclient is not installed; SMB polling is unavailable for %s", self._path) + return set() + try: + return {entry.name for entry in smbclient.scandir(self._path) if entry.is_dir()} + except Exception: + LOGGER.exception("SMB scandir failed for %s", self._path) + return set() + + class ResultDirectoryWatcher: """Watch DUT/REF result directories and trigger scans when new result folders appear.""" @@ -47,7 +95,7 @@ class ResultDirectoryWatcher: self._min_scan_interval_seconds = min_scan_interval_seconds self._lock = threading.RLock() self._scan_lock = threading.Lock() - self._observer: PollingObserver | None = None + self._watchers: list[Any] = [] self._current_config: WatchConfig | None = None self._last_scan_monotonic = 0.0 @@ -63,7 +111,7 @@ class ResultDirectoryWatcher: ) with self._lock: - if new_config == self._current_config and self._observer is not None: + if new_config == self._current_config and self._watchers: return self._stop_locked() @@ -73,29 +121,89 @@ class ResultDirectoryWatcher: LOGGER.info("result watcher disabled: DUT/REF result directories are not both set") return - handler = _ResultDirEventHandler(self._on_new_directory) - observer = PollingObserver(timeout=self._poll_interval_seconds) - smb_creds = { "username": new_config.smb_username, "password": new_config.smb_password, "domain": new_config.smb_domain, } - scheduled_count = 0 - for result_dir, label in ((dut_dir, "DUT"), (ref_dir, "REF")): - if not path_exists_with_smb(result_dir, smb_credentials=smb_creds): - LOGGER.warning("%s result directory does not exist yet, skipping watch: %s", label, result_dir) - continue - observer.schedule(handler, result_dir, recursive=False) - scheduled_count += 1 + local_paths: list[Path] = [] + smb_paths: list[str] = [] + seen: set[str] = set() - if scheduled_count == 0: + for result_dir, label in ((dut_dir, "DUT"), (ref_dir, "REF")): + is_watchable, reason = _validate_watch_path(result_dir, smb_creds) + if not is_watchable: + LOGGER.warning( + "%s result directory is not watchable, skipping watch: %s (reason: %s)", + label, + result_dir, + reason, + ) + continue + + if _is_unc_path(result_dir): + unc = _to_unc_path(result_dir) + key = unc.lower() + if key not in seen: + seen.add(key) + smb_paths.append(unc) + else: + p = Path(result_dir) + key = os.path.normcase(str(p)) + if key not in seen: + seen.add(key) + local_paths.append(p) + + watchers: list[Any] = [] + + if local_paths: + handler = _ResultDirEventHandler(self._on_new_directory) + observer = PollingObserver(timeout=self._poll_interval_seconds) + scheduled_count = 0 + scheduled_paths: list[str] = [] + for path in local_paths: + try: + observer.schedule(handler, str(path), recursive=False) + scheduled_count += 1 + scheduled_paths.append(str(path)) + except Exception: + LOGGER.exception("failed to watch local result directory: %s", path) + + if scheduled_count: + try: + observer.start() + watchers.append(observer) + for path in scheduled_paths: + print(f"[watcher] Watching local directory: {path}") + LOGGER.info("watcher successfully watching local path: %s", path) + except Exception: + LOGGER.exception("failed to start local result observer") + try: + observer.stop() + observer.join(timeout=5) + except Exception: + pass + + for unc in smb_paths: + registered, reason = _register_smb_session(unc, smb_creds) + if not registered: + LOGGER.warning("SMB watch skipped for %s (reason: %s)", unc, reason) + continue + poller = _SmbPollerThread( + unc, + on_change=lambda reason: self.scan_now(reason=reason), + interval_seconds=max(2.0, self._poll_interval_seconds), + ) + poller.start() + watchers.append(poller) + LOGGER.info("watcher successfully watching SMB path: %s", unc) + print(f"[watcher] Watching SMB directory: {unc}") + if not watchers: LOGGER.warning("result watcher not started: no valid result directories to watch") return - observer.start() - self._observer = observer + self._watchers = watchers LOGGER.info("result watcher started for DUT=%s REF=%s", dut_dir, ref_dir) self.scan_now(reason="startup") @@ -139,11 +247,19 @@ class ResultDirectoryWatcher: self._stop_locked() def _stop_locked(self) -> None: - observer = self._observer - self._observer = None - if observer is not None: - observer.stop() - observer.join(timeout=5) + watchers = self._watchers + self._watchers = [] + for watcher in watchers: + try: + watcher.stop() + except Exception: + pass + for watcher in watchers: + try: + watcher.join(timeout=5) + except Exception: + pass + if watchers: LOGGER.info("result watcher stopped") @@ -167,3 +283,72 @@ def _normalize_watch_path(path_value: Any) -> str: return str(resolve_runtime_path(cleaned)).strip() +def _is_unc_path(path_value: str) -> bool: + path = str(path_value) + return path.startswith("\\\\") or path.startswith("//") + + +def _to_unc_path(path_value: str) -> str: + path = str(path_value).replace("/", "\\") + if path.startswith("\\\\"): + return path + if path.startswith("//"): + return "\\\\" + path.lstrip("/\\") + return "\\\\" + path.lstrip("/\\") + + +def _validate_watch_path(path_value: str, smb_credentials: dict[str, str]) -> tuple[bool, str]: + path = str(path_value).strip() + if not path: + return False, "path is empty" + + if _is_unc_path(path): + if smbclient is None: + return False, "smbclient is not installed" + unc = _to_unc_path(path) + registered, reason = _register_smb_session(unc, smb_credentials) + if not registered: + return False, reason + try: + # Access one directory entry to validate permissions/connectivity. + next(iter(smbclient.scandir(unc)), None) + except Exception as exc: + return False, f"{type(exc).__name__}: {exc}" + return True, "ok" + + if not os.path.exists(path): + return False, "path does not exist" + if not os.path.isdir(path): + return False, "path is not a directory" + try: + with os.scandir(path): + pass + except Exception as exc: + return False, f"{type(exc).__name__}: {exc}" + return True, "ok" + + +def _register_smb_session(unc_path: str, smb_credentials: dict[str, str]) -> tuple[bool, str]: + if smbclient is None: + return False, "smbclient is not installed" + + server = unc_path[2:].split("\\", 1)[0] + if not server: + return False, "invalid UNC path (missing server)" + + username = str(smb_credentials.get("username") or "").strip() + password = str(smb_credentials.get("password") or "") + domain = str(smb_credentials.get("domain") or "").strip() + if username and domain and "\\" not in username and "@" not in username: + username = f"{domain}\\{username}" + + try: + if username: + smbclient.register_session(server, username=username, password=password) + else: + smbclient.register_session(server) + return True, "ok" + except Exception as exc: + return False, f"SMB session registration failed for {server}: {type(exc).__name__}: {exc}" + +