refactor(tools): lift engine ensure/sub-section boilerplate into _Engine.__init__/_build, contextlib.suppress for swallow-only try/except, collapse yuanbao _err into _YbError, fold pending-dir helpers
This commit is contained in:
@@ -3,18 +3,18 @@ stop phrases, and the TTS self-echo guard. No audio dependencies."""
|
||||
|
||||
import difflib
|
||||
import re
|
||||
from contextlib import suppress
|
||||
from typing import Optional
|
||||
|
||||
|
||||
def _voice_config() -> dict:
|
||||
"""``voice`` section of config.yaml, or ``{}`` when missing, malformed, or the
|
||||
config system can't be imported (broken config mid-install)."""
|
||||
try:
|
||||
with suppress(Exception):
|
||||
from hermes_cli.config import load_config
|
||||
voice_cfg = load_config().get("voice", {})
|
||||
return voice_cfg if isinstance(voice_cfg, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
return {}
|
||||
|
||||
|
||||
# Whisper commonly hallucinates these phrases on silent/near-silent audio
|
||||
@@ -36,7 +36,8 @@ _HALLUCINATION_REPEAT_RE = re.compile(r'^(?:thank you|thanks|bye|you|ok|okay|the
|
||||
def is_whisper_hallucination(transcript: str) -> bool:
|
||||
"""Check if a transcript is a known Whisper hallucination on silence."""
|
||||
cleaned = transcript.strip().lower()
|
||||
return not cleaned or cleaned.rstrip('.!') in WHISPER_HALLUCINATIONS or bool(_HALLUCINATION_REPEAT_RE.match(cleaned))
|
||||
return (not cleaned or cleaned.rstrip('.!') in WHISPER_HALLUCINATIONS
|
||||
or bool(_HALLUCINATION_REPEAT_RE.match(cleaned)))
|
||||
|
||||
|
||||
DEFAULT_VOICE_STOP_PHRASES = ("stop",)
|
||||
@@ -46,15 +47,12 @@ def _load_voice_stop_phrases() -> tuple:
|
||||
"""Configured ``voice.stop_phrases`` (default ``("stop",)``); an empty tuple disables
|
||||
the feature. Malformed config (dict, list of non-strings) falls back to the default
|
||||
rather than crashing the voice loop."""
|
||||
try:
|
||||
with suppress(Exception):
|
||||
raw = _voice_config().get("stop_phrases", DEFAULT_VOICE_STOP_PHRASES)
|
||||
if isinstance(raw, str):
|
||||
raw = [raw]
|
||||
if isinstance(raw, (list, tuple)):
|
||||
return tuple(str(p).strip().lower() for p in raw
|
||||
if isinstance(p, (str, int, float)) and str(p).strip())
|
||||
except Exception:
|
||||
pass
|
||||
return tuple(str(p).strip().lower() for p in raw if isinstance(p, (str, int, float)) and str(p).strip())
|
||||
return DEFAULT_VOICE_STOP_PHRASES
|
||||
|
||||
|
||||
@@ -90,17 +88,12 @@ def _normalize_for_echo_compare(text: str) -> str:
|
||||
|
||||
def is_tts_echo(transcript: str, spoken_text: str,
|
||||
threshold: float = DEFAULT_TTS_ECHO_SIMILARITY_THRESHOLD) -> bool:
|
||||
"""True when *transcript* looks like a self-capture of *spoken_text*.
|
||||
|
||||
Character-level similarity (language-agnostic, no word tokenization): a genuine
|
||||
user interjection is very unlikely to closely match Hermes' own words, so a high
|
||||
ratio signals speaker-bleed (fail-closed guard for the playback-phase listener).
|
||||
The playback-phase capture only spans pre-roll plus time-to-silence, so for long
|
||||
replies the transcript is a short FRAGMENT and the whole-string ratio dilutes toward
|
||||
0; when it misses, a transcript-sized window slides across `spoken_text`. Transcripts
|
||||
shorter than `MIN_FRAGMENT_LENGTH_FOR_ECHO` skip the fallback (a short interjection
|
||||
trivially matches a short window).
|
||||
"""
|
||||
"""True when *transcript* looks like a self-capture of *spoken_text*. Character-level similarity
|
||||
(language-agnostic): a genuine interjection rarely matches Hermes' own words, so a high ratio signals
|
||||
speaker-bleed (fail-closed guard for the playback-phase listener). Playback capture spans only pre-roll
|
||||
plus time-to-silence, so for long replies the transcript is a FRAGMENT and the whole-string ratio dilutes
|
||||
toward 0; when it misses, a transcript-sized window slides across `spoken_text`. Transcripts shorter than
|
||||
`MIN_FRAGMENT_LENGTH_FOR_ECHO` skip the fallback (a short interjection trivially matches a short window)."""
|
||||
a, b = _normalize_for_echo_compare(transcript or ""), _normalize_for_echo_compare(spoken_text or "")
|
||||
if not a or not b:
|
||||
return False
|
||||
|
||||
+75
-131
@@ -16,13 +16,14 @@ import queue
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
|
||||
from tools.wake_word_engines import ( # noqa: F401 (re-exported for callers/tests)
|
||||
_SHERPA_KWS_MODEL_DIR, _SHERPA_KWS_MODEL_URL, _Engine, _OpenWakeWordEngine, _PorcupineEngine,
|
||||
_SherpaKwsEngine, _ensure_sherpa_model, _looks_like_path, _sherpa_model_root,
|
||||
_SherpaKwsEngine, _ensure_sherpa_model, _looks_like_path, _sherpa_model_root, _sub,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -99,19 +100,16 @@ def resolve_inference_framework(cfg: Dict[str, Any]) -> str:
|
||||
coerced to tflite with a one-time warning (a pre-fix pin must not stay deaf)."""
|
||||
global _warned_onnx_coerced
|
||||
|
||||
sub = cfg.get("openwakeword") if isinstance(cfg.get("openwakeword"), dict) else {}
|
||||
framework = str(sub.get("inference_framework") or "").strip().lower()
|
||||
framework = str(_sub(cfg, "openwakeword").get("inference_framework") or "").strip().lower()
|
||||
if not framework:
|
||||
return default_inference_framework()
|
||||
if framework == "onnx" and _is_macos_arm64():
|
||||
if not _warned_onnx_coerced:
|
||||
_warned_onnx_coerced = True
|
||||
logger.warning(
|
||||
"wake: openwakeword.inference_framework='onnx' is set but ONNX's "
|
||||
"embedding model never fires on macOS ARM64 (openWakeWord #336) — "
|
||||
"using tflite instead. Set inference_framework to '' (auto) or "
|
||||
"'tflite' in config.yaml to silence this."
|
||||
)
|
||||
logger.warning("wake: openwakeword.inference_framework='onnx' is set but ONNX's "
|
||||
"embedding model never fires on macOS ARM64 (openWakeWord #336) — "
|
||||
"using tflite instead. Set inference_framework to '' (auto) or "
|
||||
"'tflite' in config.yaml to silence this.")
|
||||
return "tflite"
|
||||
return framework
|
||||
|
||||
@@ -140,11 +138,10 @@ def ensure_tflite_runtime() -> bool:
|
||||
|
||||
def load_wake_word_config() -> Dict[str, Any]:
|
||||
"""Return the ``wake_word`` config section, shape-guarded to a dict."""
|
||||
try:
|
||||
cfg = None
|
||||
with suppress(Exception):
|
||||
from hermes_cli.config import load_config
|
||||
cfg = load_config().get("wake_word")
|
||||
except Exception:
|
||||
cfg = None
|
||||
return cfg if isinstance(cfg, dict) else {}
|
||||
|
||||
|
||||
@@ -169,9 +166,7 @@ def _provider(cfg: Dict[str, Any]) -> str:
|
||||
def _input_device(cfg: Dict[str, Any]) -> int | str | None:
|
||||
"""Configured PortAudio input selector, preserving indices and names."""
|
||||
raw = _get(cfg, "input_device")
|
||||
if raw is None or isinstance(raw, bool):
|
||||
return None
|
||||
return raw if isinstance(raw, int) else (str(raw).strip() or None)
|
||||
return None if isinstance(raw, bool) else raw if raw is None or isinstance(raw, int) else (str(raw).strip() or None)
|
||||
|
||||
|
||||
def _sensitivity(cfg: Dict[str, Any]) -> float:
|
||||
@@ -201,9 +196,7 @@ def resolve_capture_mode(cfg: Optional[Dict[str, Any]] = None, *, prefer_client:
|
||||
raw = str(_get(cfg, "capture") or "auto").strip().lower()
|
||||
if raw in ("client", "remote", "external"):
|
||||
return "client"
|
||||
if raw != "local" and prefer_client and not _local_input_device_ready():
|
||||
return "client"
|
||||
return "local"
|
||||
return "client" if raw != "local" and prefer_client and not _local_input_device_ready() else "local"
|
||||
|
||||
|
||||
def _input_channels(info: Any) -> int:
|
||||
@@ -219,7 +212,8 @@ def _local_input_device_ready() -> bool:
|
||||
if isinstance(devices, dict):
|
||||
return _input_channels(devices) > 0
|
||||
# Also accept a resolvable default input (some hosts list devices oddly).
|
||||
return any(_input_channels(d) > 0 for d in devices) or _input_channels(sd.query_devices(None, "input")) > 0
|
||||
return (any(_input_channels(d) > 0 for d in devices)
|
||||
or _input_channels(sd.query_devices(None, "input")) > 0)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
@@ -228,20 +222,17 @@ def wake_surface_enabled(surface: str, cfg: Optional[Dict[str, Any]] = None) ->
|
||||
"""Should ``surface`` (cli/tui/gui) host the listener? True when enabled and the configured
|
||||
surface is ``auto`` or this one; ``auto`` only makes it eligible — the lock admits one claimant."""
|
||||
cfg = cfg if cfg is not None else load_wake_word_config()
|
||||
if not cfg.get("enabled"):
|
||||
return False
|
||||
want = str(_get(cfg, "surface")).strip().lower() or "auto"
|
||||
return want == "auto" or want == surface.strip().lower()
|
||||
return bool(cfg.get("enabled")) and want in ("auto", surface.strip().lower())
|
||||
|
||||
|
||||
# ── Multi-profile phrase enrollment (open-vocabulary routing) ──
|
||||
|
||||
def _active_profile_name() -> str:
|
||||
try:
|
||||
with suppress(Exception):
|
||||
from hermes_cli.profiles import get_active_profile_name
|
||||
return get_active_profile_name() or "default"
|
||||
except Exception:
|
||||
return "default"
|
||||
return "default"
|
||||
|
||||
|
||||
def enrolled_profile_phrases() -> Dict[str, str]:
|
||||
@@ -249,22 +240,17 @@ def enrolled_profile_phrases() -> Dict[str, str]:
|
||||
``config.yaml`` raw (``load_config()`` targets only the ACTIVE profile). Phrase defaults to
|
||||
``"hey <profile>"``; the sherpa engine listens for all and routes to the match. Unreadable → skipped."""
|
||||
phrases: Dict[str, str] = {}
|
||||
try:
|
||||
with suppress(Exception):
|
||||
from hermes_cli.config import read_user_config_raw
|
||||
from hermes_cli.profiles import get_profile_dir, list_profiles
|
||||
for info in list_profiles():
|
||||
name = getattr(info, "name", None) or str(info)
|
||||
try:
|
||||
with suppress(Exception):
|
||||
wc = read_user_config_raw(Path(get_profile_dir(name)) / "config.yaml").get("wake_word") or {}
|
||||
if not isinstance(wc, dict) or not wc.get("enabled"):
|
||||
continue
|
||||
phrase = str(wc.get("phrase") or f"hey {name}").strip()
|
||||
if phrase:
|
||||
phrases[name] = phrase
|
||||
except Exception:
|
||||
continue
|
||||
except Exception:
|
||||
pass
|
||||
if isinstance(wc, dict) and wc.get("enabled"):
|
||||
phrase = str(wc.get("phrase") or f"hey {name}").strip()
|
||||
if phrase:
|
||||
phrases[name] = phrase
|
||||
return phrases
|
||||
|
||||
|
||||
@@ -277,18 +263,17 @@ def _import_audio():
|
||||
|
||||
|
||||
def _audio_available() -> bool:
|
||||
try:
|
||||
_import_audio()
|
||||
return True
|
||||
except (ImportError, OSError):
|
||||
return False
|
||||
with suppress(ImportError, OSError):
|
||||
return bool(_import_audio())
|
||||
return False
|
||||
|
||||
|
||||
def _describe_input_device(sd, selector: int | str | None) -> Dict[str, Any]:
|
||||
def _describe_input_device(selector: int | str | None, sd=None) -> Dict[str, Any]:
|
||||
"""Resolve a PortAudio selector into JSON-safe diagnostics (``InputStream`` stays the
|
||||
authority on whether the device actually opens)."""
|
||||
authority on whether the device actually opens). Imports sounddevice unless ``sd`` is given."""
|
||||
details: Dict[str, Any] = {"selector": selector}
|
||||
try:
|
||||
sd = sd or _import_audio()[0]
|
||||
info = sd.query_devices(selector, "input")
|
||||
except Exception as e:
|
||||
details["error"] = str(e)
|
||||
@@ -298,23 +283,21 @@ def _describe_input_device(sd, selector: int | str | None) -> Dict[str, Any]:
|
||||
if info.get("name"):
|
||||
details["name"] = str(info["name"])
|
||||
for key, out_key, cast in (("max_input_channels", "max_input_channels", int),
|
||||
("default_samplerate", "default_samplerate", float),
|
||||
("hostapi", "hostapi_index", int)):
|
||||
("default_samplerate", "default_samplerate", float), ("hostapi", "hostapi_index", int)):
|
||||
if isinstance(info.get(key), (int, float)):
|
||||
details[out_key] = cast(info[key])
|
||||
if "hostapi_index" in details:
|
||||
try:
|
||||
with suppress(Exception):
|
||||
hostapi = sd.query_hostapis(details["hostapi_index"])
|
||||
if isinstance(hostapi, dict) and hostapi.get("name"):
|
||||
details["hostapi"] = str(hostapi["name"])
|
||||
except Exception:
|
||||
pass
|
||||
return details
|
||||
|
||||
|
||||
def _device_label(details: Dict[str, Any]) -> str:
|
||||
selector = details.get("selector")
|
||||
label = str(details.get("name") or "").strip() or ("system default" if selector is None else str(selector))
|
||||
name = str(details.get("name") or "").strip()
|
||||
label = name or ("system default" if selector is None else str(selector))
|
||||
hostapi = str(details.get("hostapi") or "").strip()
|
||||
return f"{label} ({hostapi})" if hostapi else label
|
||||
|
||||
@@ -323,10 +306,8 @@ def _capture_sample_rate(details: Dict[str, Any]) -> int:
|
||||
"""Use the selected device's native rate when PortAudio reports one."""
|
||||
rate = details.get("default_samplerate")
|
||||
if isinstance(rate, (int, float)) and not isinstance(rate, bool) and rate > 0:
|
||||
try:
|
||||
with suppress(OverflowError, ValueError):
|
||||
return int(round(rate))
|
||||
except (OverflowError, ValueError):
|
||||
pass
|
||||
return SAMPLE_RATE
|
||||
|
||||
|
||||
@@ -352,11 +333,9 @@ def _resample_audio_frame(np, frame, output_length: int):
|
||||
def silent_audio_hint(details: Dict[str, Any]) -> str:
|
||||
"""Platform-specific remediation for an armed stream delivering silence."""
|
||||
if sys.platform == "darwin":
|
||||
return (
|
||||
"Microphone delivers only silence. Grant the Hermes backend "
|
||||
"microphone access in System Settings > Privacy & Security > "
|
||||
"Microphone, then toggle the wake word."
|
||||
)
|
||||
return ("Microphone delivers only silence. Grant the Hermes backend "
|
||||
"microphone access in System Settings > Privacy & Security > "
|
||||
"Microphone, then toggle the wake word.")
|
||||
fix = ("Set wake_word.input_device to a different PortAudio input device"
|
||||
if sys.platform == "win32" else "Check the selected input device")
|
||||
return (f"Microphone delivers only silence from {_device_label(details)}. "
|
||||
@@ -377,12 +356,11 @@ def _build_engine(cfg: Dict[str, Any]) -> _Engine:
|
||||
def _stt_ready() -> bool:
|
||||
"""Is a speech-to-text provider configured and enabled? (A wake without STT arms the
|
||||
mic but every utterance dies at transcription — same bar as ``check_voice_requirements``.)"""
|
||||
try:
|
||||
with suppress(Exception):
|
||||
from tools.transcription_tools import _get_provider, _load_stt_config, is_stt_enabled
|
||||
stt_config = _load_stt_config()
|
||||
return is_stt_enabled(stt_config) and _get_provider(stt_config) != "none"
|
||||
except Exception:
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
_LAZY_TTS_FEATURES = {"edge": "tts.edge", "elevenlabs": "tts.elevenlabs", "mistral": "tts.mistral"}
|
||||
@@ -421,9 +399,8 @@ def check_wake_word_requirements(cfg: Optional[Dict[str, Any]] = None) -> Dict[s
|
||||
stt_ok, tts_ok = _stt_ready(), _tts_ready()
|
||||
# tflite needs a runtime openWakeWord doesn't declare off Linux; report it as a
|
||||
# remediation instead of arming a detector that can't fire.
|
||||
tflite_ok = True
|
||||
if feature == "wake.openwakeword" and resolve_inference_framework(cfg) == "tflite":
|
||||
tflite_ok = ensure_tflite_runtime() or lazy_deps.is_available("wake.openwakeword.tflite") or lazy_ok
|
||||
tflite_ok = (feature != "wake.openwakeword" or resolve_inference_framework(cfg) != "tflite"
|
||||
or ensure_tflite_runtime() or lazy_deps.is_available("wake.openwakeword.tflite") or lazy_ok)
|
||||
key_ok = provider != "porcupine" or bool((os.getenv("PORCUPINE_ACCESS_KEY") or "").strip())
|
||||
capture_mode = resolve_capture_mode(cfg)
|
||||
missing = " and ".join(n for n, ok in (("speech-to-text", stt_ok), ("text-to-speech", tts_ok)) if not ok)
|
||||
@@ -447,11 +424,9 @@ def check_wake_word_requirements(cfg: Optional[Dict[str, Any]] = None) -> Dict[s
|
||||
else:
|
||||
mic_ok = (deps_ok and audio_ok) or (not deps_ok and lazy_ok)
|
||||
if deps_ok and not audio_ok and not hint:
|
||||
hint = (
|
||||
"No local microphone on this backend. Remote desktop can stream "
|
||||
"the client mic — set wake_word.capture: client or use a desktop "
|
||||
"build with client-capture wake support."
|
||||
)
|
||||
hint = ("No local microphone on this backend. Remote desktop can stream "
|
||||
"the client mic — set wake_word.capture: client or use a desktop "
|
||||
"build with client-capture wake support.")
|
||||
|
||||
return {
|
||||
"available": key_ok and stt_ok and tts_ok and tflite_ok and mic_ok, "provider": provider,
|
||||
@@ -478,35 +453,29 @@ class _Capture:
|
||||
"""One raw block; None when no client frame arrived within 250 ms. Stream errors propagate."""
|
||||
if self.stream is not None:
|
||||
return self.stream.read(self.frame_length)[0]
|
||||
try:
|
||||
with suppress(Exception):
|
||||
return self.queue.get(timeout=0.25)
|
||||
except Exception:
|
||||
return None
|
||||
return None
|
||||
|
||||
def close(self) -> None:
|
||||
try:
|
||||
with suppress(Exception):
|
||||
if self.stream is not None:
|
||||
self.stream.stop()
|
||||
self.stream.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
class WakeWordDetector:
|
||||
"""Background hotword listener; fires ``on_wake()`` when the phrase is heard. The engine is built
|
||||
once and kept across pause/resume — only the stream + reader thread cycle, so mic toggles are cheap."""
|
||||
|
||||
def __init__(self, engine: _Engine, on_wake: Callable[[], None],
|
||||
cooldown: float = _FIRE_COOLDOWN_SECONDS,
|
||||
def __init__(self, engine: _Engine, on_wake: Callable[[], None], cooldown: float = _FIRE_COOLDOWN_SECONDS,
|
||||
on_failure: Optional[Callable[["WakeWordDetector"], None]] = None,
|
||||
input_device: int | str | None = None,
|
||||
external_audio: bool = False):
|
||||
input_device: int | str | None = None, external_audio: bool = False):
|
||||
self.engine, self.on_wake, self.cooldown, self.on_failure = engine, on_wake, cooldown, on_failure
|
||||
self.input_device, self.external_audio = input_device, bool(external_audio)
|
||||
self.input_device_details: Dict[str, Any] = (
|
||||
{"selector": "client", "name": "client capture", "hostapi": "remote"}
|
||||
if self.external_audio else {"selector": input_device}
|
||||
)
|
||||
if self.external_audio else {"selector": input_device})
|
||||
self._thread: Optional[threading.Thread] = None
|
||||
self._stop, self._callback_inflight = threading.Event(), threading.Event()
|
||||
self._lock, self._last_fire = threading.Lock(), 0.0
|
||||
@@ -544,11 +513,9 @@ class WakeWordDetector:
|
||||
try:
|
||||
self._audio_q.put_nowait(chunk)
|
||||
except Exception:
|
||||
try: # full: drop the oldest frame, then retry once
|
||||
with suppress(Exception): # full: drop the oldest frame, then retry once
|
||||
self._audio_q.get_nowait()
|
||||
self._audio_q.put_nowait(chunk)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def start(self) -> None:
|
||||
"""Open the mic (or client feeder) and begin listening. Idempotent."""
|
||||
@@ -599,11 +566,9 @@ class WakeWordDetector:
|
||||
def _open_capture(self, frame_length: int) -> _Capture:
|
||||
"""Open the audio source; raises on any local-mic failure."""
|
||||
if self.external_audio:
|
||||
try: # drain stale frames from a previous arm
|
||||
with suppress(Exception): # drain stale frames from a previous arm
|
||||
while True:
|
||||
self._audio_q.get_nowait()
|
||||
except Exception:
|
||||
pass
|
||||
logger.info("wake word: client-capture mode (frame=%d, rate=%d) — waiting for wake.feed",
|
||||
frame_length, SAMPLE_RATE)
|
||||
return _Capture(queue=self._audio_q, frame_length=frame_length)
|
||||
@@ -613,7 +578,7 @@ class WakeWordDetector:
|
||||
except (ImportError, OSError) as e:
|
||||
logger.error("wake word: audio libraries unavailable: %s", e)
|
||||
raise
|
||||
details = self.input_device_details = _describe_input_device(sd, self.input_device)
|
||||
details = self.input_device_details = _describe_input_device(self.input_device, sd)
|
||||
cap = _Capture(np=np, rate=_capture_sample_rate(details))
|
||||
cap.frame_length = max(1, int(round(frame_length * cap.rate / SAMPLE_RATE)))
|
||||
logger.info("wake word: opening microphone device=%s selector=%r hostapi=%s "
|
||||
@@ -641,13 +606,13 @@ class WakeWordDetector:
|
||||
if self._silent_frames == silent_alert_frames:
|
||||
self.audio_silent = True
|
||||
if frame is not None:
|
||||
logger.warning("wake word: mic delivers only silence (peak<=%d for %ds); %s", _SILENCE_PEAK,
|
||||
_SILENCE_ALERT_SECONDS, silent_audio_hint(self.input_device_details))
|
||||
logger.warning("wake word: mic delivers only silence (peak<=%d for %ds); %s",
|
||||
_SILENCE_PEAK, _SILENCE_ALERT_SECONDS,
|
||||
silent_audio_hint(self.input_device_details))
|
||||
elif self._silent_frames:
|
||||
if self.audio_silent:
|
||||
logger.info("wake word: mic audio detected — stream healthy")
|
||||
self._silent_frames = 0
|
||||
self.audio_silent = False
|
||||
self._silent_frames, self.audio_silent = 0, False
|
||||
|
||||
def _fire(self) -> None:
|
||||
"""Honor the cooldown, then run ``on_wake`` on its own thread (once)."""
|
||||
@@ -671,10 +636,8 @@ class WakeWordDetector:
|
||||
return
|
||||
# Drop buffered audio/feature state so a resume right after a voice turn can't
|
||||
# re-fire on audio captured before the pause (wake → voice → resume → wake loop).
|
||||
try:
|
||||
with suppress(Exception):
|
||||
self.engine.reset()
|
||||
except Exception:
|
||||
pass
|
||||
logger.info("wake word: listening (frame=%d, rate=%d, external=%s)",
|
||||
frame_length, SAMPLE_RATE, self.external_audio)
|
||||
ready.set()
|
||||
@@ -743,7 +706,7 @@ def _acquire_machine_lock(path: Optional[Path] = None):
|
||||
handle = open(lock_path, "a+b")
|
||||
try:
|
||||
_flock(handle, True)
|
||||
except (OSError, BlockingIOError) as e:
|
||||
except OSError as e: # BlockingIOError is an OSError: lock held elsewhere
|
||||
handle.close()
|
||||
raise WakeWordInUse("Wake-word microphone is already owned.") from e
|
||||
return handle
|
||||
@@ -752,12 +715,9 @@ def _acquire_machine_lock(path: Optional[Path] = None):
|
||||
def _release_machine_lock(handle) -> None:
|
||||
if handle is None:
|
||||
return
|
||||
try:
|
||||
with suppress(OSError):
|
||||
_flock(handle, False)
|
||||
except OSError:
|
||||
pass
|
||||
finally:
|
||||
handle.close()
|
||||
handle.close()
|
||||
|
||||
|
||||
def _teardown_locked(close: Callable[[], None]) -> None:
|
||||
@@ -808,24 +768,19 @@ def start_listening(on_wake: Callable[[], None], *, owner: object, config: Optio
|
||||
_detector.start()
|
||||
return _detector
|
||||
except Exception:
|
||||
det = _detector
|
||||
try:
|
||||
_teardown_locked(det.stop if det is not None else lambda: None)
|
||||
except Exception:
|
||||
pass
|
||||
with suppress(Exception):
|
||||
_teardown_locked(_detector.stop if _detector is not None else lambda: None)
|
||||
raise
|
||||
|
||||
|
||||
def _owned_call(owner: object, method: Optional[str] = None) -> bool:
|
||||
"""Under the lock, True iff ``owner`` holds the lease; also invokes ``detector.<method>()`` when given."""
|
||||
def _owned_call(owner: object, action: Optional[Callable[[WakeWordDetector], None]] = None) -> bool:
|
||||
"""Under the lock, True iff ``owner`` holds the lease; also runs ``action(detector)`` when given."""
|
||||
with _detector_lock:
|
||||
det = _owned_detector(owner)
|
||||
if det is None:
|
||||
return False
|
||||
if method == "stop":
|
||||
_teardown_locked(det.stop)
|
||||
elif method:
|
||||
getattr(det, method)()
|
||||
if action is not None:
|
||||
action(det)
|
||||
return True
|
||||
|
||||
|
||||
@@ -835,17 +790,17 @@ def owns_listener(owner: object) -> bool:
|
||||
|
||||
def pause_listening(*, owner: object) -> bool:
|
||||
"""Release the microphone only when ``owner`` holds the lease."""
|
||||
return _owned_call(owner, "pause")
|
||||
return _owned_call(owner, WakeWordDetector.pause)
|
||||
|
||||
|
||||
def resume_listening(*, owner: object) -> bool:
|
||||
"""Re-open the microphone only when ``owner`` holds the lease."""
|
||||
return _owned_call(owner, "resume")
|
||||
return _owned_call(owner, WakeWordDetector.resume)
|
||||
|
||||
|
||||
def stop_listening(*, owner: object) -> bool:
|
||||
"""Fully stop the detector only when ``owner`` holds the lease."""
|
||||
return _owned_call(owner, "stop")
|
||||
return _owned_call(owner, lambda det: _teardown_locked(det.stop))
|
||||
|
||||
|
||||
def _current_detector() -> Optional[WakeWordDetector]:
|
||||
@@ -854,36 +809,26 @@ def _current_detector() -> Optional[WakeWordDetector]:
|
||||
|
||||
|
||||
def is_listening() -> bool:
|
||||
det = _current_detector()
|
||||
return det is not None and det.running
|
||||
return (det := _current_detector()) is not None and det.running
|
||||
|
||||
|
||||
def audio_is_silent() -> bool:
|
||||
"""True when the armed stream opens fine but delivers only silence (dead mic), so
|
||||
detection can never fire; status shows "listening but the microphone appears silent"."""
|
||||
det = _current_detector()
|
||||
return det is not None and det.audio_silent
|
||||
return (det := _current_detector()) is not None and det.audio_silent
|
||||
|
||||
|
||||
def get_input_device_status(cfg: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
"""Return configured/active PortAudio input diagnostics for status UIs."""
|
||||
det = _current_detector()
|
||||
if det is not None:
|
||||
if (det := _current_detector()) is not None:
|
||||
return dict(det.input_device_details)
|
||||
cfg = cfg if cfg is not None else load_wake_word_config()
|
||||
selector = _input_device(cfg)
|
||||
try:
|
||||
sd, _ = _import_audio()
|
||||
except (ImportError, OSError) as e:
|
||||
return {"selector": selector, "error": str(e)}
|
||||
return _describe_input_device(sd, selector)
|
||||
return _describe_input_device(_input_device(cfg if cfg is not None else load_wake_word_config()))
|
||||
|
||||
|
||||
def get_last_match() -> Optional[tuple[str, str]]:
|
||||
"""(matched phrase, profile) of the most recent wake fire when the engine reports
|
||||
per-phrase matches (sherpa multi-profile routing); None otherwise."""
|
||||
det = _current_detector()
|
||||
return None if det is None else getattr(det.engine, "last_match", None)
|
||||
return None if (det := _current_detector()) is None else getattr(det.engine, "last_match", None)
|
||||
|
||||
|
||||
def feed_audio(*, owner: object, pcm_int16) -> bool:
|
||||
@@ -898,8 +843,7 @@ def feed_audio(*, owner: object, pcm_int16) -> bool:
|
||||
|
||||
def detector_frame_info() -> Dict[str, Any]:
|
||||
"""Sample rate + frame length for client capture streamers."""
|
||||
det = _current_detector()
|
||||
if det is None:
|
||||
if (det := _current_detector()) is None:
|
||||
return {"sample_rate": SAMPLE_RATE, "frame_length": 1280}
|
||||
return {"sample_rate": SAMPLE_RATE, "external_audio": bool(det.external_audio),
|
||||
"frame_length": int(getattr(det.engine, "frame_length", 1280) or 1280)}
|
||||
|
||||
+35
-32
@@ -9,6 +9,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
@@ -26,14 +27,24 @@ def _ensure_dep(feature: str) -> None:
|
||||
|
||||
|
||||
class _Engine:
|
||||
"""Minimal hotword-engine contract: feed int16 frames, get a bool."""
|
||||
"""Minimal hotword-engine contract: feed int16 frames, get a bool. Subclasses set ``feature``
|
||||
(lazy_deps name, ensured before ``_build``) and their own ``cfg`` sub-section ``section``."""
|
||||
|
||||
feature: str = ""
|
||||
section: str = ""
|
||||
frame_length: int = 1280 # 80 ms at 16 kHz
|
||||
|
||||
#: (matched phrase, profile name) of the most recent fire. Multi-phrase engines
|
||||
#: (sherpa) set this for profile routing; single-phrase engines leave it None.
|
||||
last_match: Optional[tuple[str, str]] = None
|
||||
|
||||
def __init__(self, cfg: Dict[str, Any]):
|
||||
_ensure_dep(self.feature)
|
||||
self._build(cfg, _sub(cfg, self.section), _ww())
|
||||
|
||||
def _build(self, cfg: Dict[str, Any], sub: Dict[str, Any], ww) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def process(self, frame) -> bool: # frame: 1-D int16 ndarray
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -58,14 +69,13 @@ class _OpenWakeWordEngine(_Engine):
|
||||
``sensitivity`` IS the raw 0..1 threshold (higher = stricter). A real utterance holds the score
|
||||
high across frames while a stray phoneme spikes one, so ``confirmation_frames`` hits are required."""
|
||||
|
||||
feature, section = "wake.openwakeword", "openwakeword"
|
||||
frame_length = 1280 # openWakeWord recommends 80 ms frames.
|
||||
|
||||
def __init__(self, cfg: Dict[str, Any]):
|
||||
_ensure_dep("wake.openwakeword")
|
||||
def _build(self, cfg, sub, ww) -> None:
|
||||
import openwakeword
|
||||
from openwakeword.model import Model
|
||||
ww = _ww()
|
||||
model_ref = str(_sub(cfg, "openwakeword").get("model") or ww._BUNDLED_MODEL_NAME).strip()
|
||||
model_ref = str(sub.get("model") or ww._BUNDLED_MODEL_NAME).strip()
|
||||
framework = self._usable_framework(ww.resolve_inference_framework(cfg))
|
||||
self._threshold = ww._sensitivity(cfg)
|
||||
self._confirm_needed = ww._confirmation_frames(cfg)
|
||||
@@ -114,10 +124,8 @@ class _OpenWakeWordEngine(_Engine):
|
||||
# Clears openWakeWord's rolling feature buffer so stale audio captured before a
|
||||
# pause can't re-fire the moment we resume.
|
||||
self._confirm_streak = 0
|
||||
try:
|
||||
with suppress(Exception):
|
||||
self._model.reset()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def close(self) -> None:
|
||||
self.reset()
|
||||
@@ -161,15 +169,14 @@ class _SherpaKwsEngine(_Engine):
|
||||
"""sherpa-onnx open-vocabulary keyword spotting — any typed phrase, zero training. ``wake_word.phrase``
|
||||
is BPE-tokenized at runtime against the model's vocabulary: DETECTION config, not a cosmetic label."""
|
||||
|
||||
feature, section = "wake.sherpa", "sherpa"
|
||||
frame_length = 1280 # streaming zipformer accepts any chunk; match capture path.
|
||||
|
||||
def __init__(self, cfg: Dict[str, Any]):
|
||||
_ensure_dep("wake.sherpa")
|
||||
def _build(self, cfg, sub, ww) -> None:
|
||||
import sherpa_onnx
|
||||
import tempfile
|
||||
from sherpa_onnx import text2token
|
||||
ww = _ww()
|
||||
model_dir = str(_sub(cfg, "sherpa").get("model_dir") or "").strip()
|
||||
model_dir = str(sub.get("model_dir") or "").strip()
|
||||
d = Path(model_dir) if model_dir else _ensure_sherpa_model()
|
||||
if not (d / "tokens.txt").exists():
|
||||
raise RuntimeError(f"sherpa KWS model not found at {d}")
|
||||
@@ -201,16 +208,16 @@ class _SherpaKwsEngine(_Engine):
|
||||
# matrix tests while 0.25 held zero false fires.
|
||||
threshold = 0.05 + 0.4 * ww._sensitivity(cfg)
|
||||
|
||||
def _model_file(pattern: str) -> str:
|
||||
hits = sorted(d.glob(pattern))
|
||||
def _model_file(part: str) -> str:
|
||||
hits = sorted(d.glob(f"{part}-*[!8].onnx"))
|
||||
if not hits:
|
||||
raise RuntimeError(f"sherpa KWS model file missing: {d}/{pattern}")
|
||||
raise RuntimeError(f"sherpa KWS model file missing: {d}/{part}-*[!8].onnx")
|
||||
return str(hits[0])
|
||||
|
||||
self._spotter = sherpa_onnx.KeywordSpotter(
|
||||
tokens=str(d / "tokens.txt"), encoder=_model_file("encoder-*[!8].onnx"),
|
||||
decoder=_model_file("decoder-*[!8].onnx"), joiner=_model_file("joiner-*[!8].onnx"),
|
||||
keywords_file=self._keywords_file, keywords_threshold=threshold, num_threads=1,
|
||||
tokens=str(d / "tokens.txt"), encoder=_model_file("encoder"), decoder=_model_file("decoder"),
|
||||
joiner=_model_file("joiner"), keywords_file=self._keywords_file, keywords_threshold=threshold,
|
||||
num_threads=1,
|
||||
)
|
||||
self._stream = self._spotter.create_stream()
|
||||
|
||||
@@ -223,38 +230,36 @@ class _SherpaKwsEngine(_Engine):
|
||||
result = self._spotter.get_result(self._stream)
|
||||
if result:
|
||||
fired, display = True, str(result)
|
||||
self.last_match = (display.replace("_", " ").lower(), self._display_to_profile.get(display, ""))
|
||||
self.last_match = (display.replace("_", " ").lower(),
|
||||
self._display_to_profile.get(display, ""))
|
||||
self._spotter.reset_stream(self._stream) # one utterance must not fire repeatedly
|
||||
return fired
|
||||
|
||||
def reset(self) -> None:
|
||||
# Fresh stream drops buffered audio/decoder state (pause → resume must not re-fire).
|
||||
try:
|
||||
with suppress(Exception):
|
||||
self._stream = self._spotter.create_stream()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def close(self) -> None:
|
||||
try:
|
||||
with suppress(OSError):
|
||||
os.unlink(self._keywords_file)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
class _PorcupineEngine(_Engine):
|
||||
"""Picovoice Porcupine — premium, on-device, needs an access key."""
|
||||
|
||||
def __init__(self, cfg: Dict[str, Any]):
|
||||
_ensure_dep("wake.porcupine")
|
||||
feature, section = "wake.porcupine", "porcupine"
|
||||
|
||||
def _build(self, cfg, sub, ww) -> None:
|
||||
import pvporcupine
|
||||
access_key = (os.getenv("PORCUPINE_ACCESS_KEY") or "").strip()
|
||||
if not access_key:
|
||||
raise RuntimeError("Porcupine wake word requires PORCUPINE_ACCESS_KEY "
|
||||
"(get a free key at https://console.picovoice.ai).")
|
||||
keyword = str(_sub(cfg, "porcupine").get("keyword") or "jarvis").strip()
|
||||
keyword = str(sub.get("keyword") or "jarvis").strip()
|
||||
# Porcupine's `sensitivities` runs the OPPOSITE way to our shared knob (higher =
|
||||
# looser); invert so "higher = stricter" holds for every engine.
|
||||
kwargs: Dict[str, Any] = {"access_key": access_key, "sensitivities": [1.0 - _ww()._sensitivity(cfg)]}
|
||||
kwargs: Dict[str, Any] = {"access_key": access_key, "sensitivities": [1.0 - ww._sensitivity(cfg)]}
|
||||
kwargs["keyword_paths" if _looks_like_path(keyword) else "keywords"] = [keyword]
|
||||
self._porcupine = pvporcupine.create(**kwargs)
|
||||
self.frame_length = self._porcupine.frame_length
|
||||
@@ -263,7 +268,5 @@ class _PorcupineEngine(_Engine):
|
||||
return self._porcupine.process(frame) >= 0 # pvporcupine wants a plain sequence of int16
|
||||
|
||||
def close(self) -> None:
|
||||
try:
|
||||
with suppress(Exception):
|
||||
self._porcupine.delete()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -12,6 +12,7 @@ from __future__ import annotations
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
from contextlib import suppress
|
||||
from typing import Dict, List
|
||||
|
||||
from hermes_cli._subprocess_compat import harden_git_argv, noninteractive_git_env
|
||||
@@ -49,13 +50,11 @@ def _untracked_diff(cwd: str, files: List[str]) -> str:
|
||||
"""Render untracked files as new-file diffs via ``git diff --no-index``."""
|
||||
chunks: List[str] = []
|
||||
for rel in files[:_MAX_UNTRACKED_FILES]:
|
||||
try:
|
||||
with suppress(subprocess.TimeoutExpired, OSError):
|
||||
# --no-index exits 1 when files differ — the success path, so the code is ignored.
|
||||
_, out = _run(["diff", "--no-index", "--", os.devnull, rel], cwd)
|
||||
if out.strip():
|
||||
chunks.append(out.rstrip("\n"))
|
||||
except (subprocess.TimeoutExpired, OSError):
|
||||
continue
|
||||
if len(files) > _MAX_UNTRACKED_FILES:
|
||||
chunks.append(f"... ({len(files) - _MAX_UNTRACKED_FILES} more untracked files not shown)")
|
||||
return "\n".join(chunks)
|
||||
|
||||
+20
-37
@@ -18,6 +18,7 @@ import os
|
||||
import re
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
@@ -60,12 +61,8 @@ def _normalize_enabled(value: Any) -> bool:
|
||||
|
||||
# --- Pending store (file-backed) ---
|
||||
|
||||
def _pending_dir(subsystem: str) -> Path:
|
||||
return get_hermes_home() / "pending" / subsystem
|
||||
|
||||
|
||||
def _pending_path(subsystem: str, pending_id: str) -> Path:
|
||||
return _pending_dir(subsystem) / f"{pending_id}.json"
|
||||
return get_hermes_home() / "pending" / subsystem / f"{pending_id}.json"
|
||||
|
||||
|
||||
def _read_record(path: Path) -> Dict[str, Any]:
|
||||
@@ -73,7 +70,7 @@ def _read_record(path: Path) -> Dict[str, Any]:
|
||||
|
||||
|
||||
def _pending_files(subsystem: str) -> list:
|
||||
d = _pending_dir(subsystem)
|
||||
d = _pending_path(subsystem, "").parent
|
||||
return list(d.glob("*.json")) if d.exists() else []
|
||||
|
||||
|
||||
@@ -113,17 +110,16 @@ def list_pending(subsystem: str) -> List[Dict[str, Any]]:
|
||||
|
||||
def get_pending(subsystem: str, pending_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""Return a single pending record by id, or None."""
|
||||
path = _pending_path(subsystem, pending_id)
|
||||
try:
|
||||
with suppress(Exception):
|
||||
path = _pending_path(subsystem, pending_id)
|
||||
return _read_record(path) if path.exists() else None
|
||||
except Exception:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def discard_pending(subsystem: str, pending_id: str) -> bool:
|
||||
"""Delete a pending record. Returns True if it existed."""
|
||||
path = _pending_path(subsystem, pending_id)
|
||||
try:
|
||||
path = _pending_path(subsystem, pending_id)
|
||||
if path.exists():
|
||||
path.unlink()
|
||||
return True
|
||||
@@ -134,10 +130,9 @@ def discard_pending(subsystem: str, pending_id: str) -> bool:
|
||||
|
||||
def pending_count(subsystem: str) -> int:
|
||||
"""Cheap count of pending records (for notification badges)."""
|
||||
try:
|
||||
with suppress(Exception):
|
||||
return len(_pending_files(subsystem))
|
||||
except Exception:
|
||||
return 0
|
||||
return 0
|
||||
|
||||
|
||||
# --- Write origin ---
|
||||
@@ -145,11 +140,10 @@ def pending_count(subsystem: str) -> int:
|
||||
def current_origin() -> str:
|
||||
"""``foreground`` or ``background_review`` — reuses the skill-provenance ContextVar
|
||||
the background review fork sets; foreground turns leave it at the default."""
|
||||
try:
|
||||
with suppress(Exception):
|
||||
from tools.skill_provenance import get_current_write_origin
|
||||
return get_current_write_origin()
|
||||
except Exception:
|
||||
return "foreground"
|
||||
return "foreground"
|
||||
|
||||
|
||||
# --- Gate decision ---
|
||||
@@ -215,11 +209,8 @@ def _prompt_inline_memory_approval(summary: str, detail: str) -> Optional[bool]:
|
||||
|
||||
# --- Skill-specific helpers (gist + diff for the review affordances) ---
|
||||
|
||||
_GIST_TEMPLATES = {
|
||||
"write_file": "write {file_path} in '{name}'",
|
||||
"remove_file": "remove {file_path} from '{name}'",
|
||||
"delete": "delete skill '{name}'",
|
||||
}
|
||||
_GIST_TEMPLATES = {"write_file": "write {file_path} in '{name}'", "remove_file": "remove {file_path} from '{name}'",
|
||||
"delete": "delete skill '{name}'"}
|
||||
|
||||
|
||||
def skill_gist(action: str, name: str, *, content: str = "", file_path: str = "",
|
||||
@@ -229,14 +220,12 @@ def skill_gist(action: str, name: str, *, content: str = "", file_path: str = ""
|
||||
if action in {"create", "edit"} and content:
|
||||
desc = _frontmatter_description(content)
|
||||
size = f"{len(content) // 1024 + 1} KB" if len(content) >= 1024 else f"{len(content)} chars"
|
||||
verb = "create" if action == "create" else "rewrite"
|
||||
return f"{verb} '{name}' — {desc} ({size})" if desc else f"{verb} '{name}' ({size})"
|
||||
return f"{'create' if action == 'create' else 'rewrite'} '{name}'{f' — {desc}' if desc else ''} ({size})"
|
||||
if action == "patch":
|
||||
removed = old_string.count("\n") + 1 if old_string else 0
|
||||
added = new_string.count("\n") + 1 if new_string else 0
|
||||
return f"patch '{name}' {file_path or 'SKILL.md'} (+{added}/-{removed} lines)"
|
||||
template = _GIST_TEMPLATES.get(action, "{action} '{name}'")
|
||||
return template.format(action=action, name=name, file_path=file_path)
|
||||
return _GIST_TEMPLATES.get(action, "{action} '{name}'").format(action=action, name=name, file_path=file_path)
|
||||
|
||||
|
||||
def _frontmatter_description(content: str) -> str:
|
||||
@@ -247,12 +236,11 @@ def _frontmatter_description(content: str) -> str:
|
||||
|
||||
def _find_skill_path(name: str) -> Optional[Path]:
|
||||
"""Directory of an installed skill, or None if unknown / lookup unavailable."""
|
||||
try:
|
||||
with suppress(Exception):
|
||||
from tools.skill_manager_tool import _find_skill
|
||||
found = _find_skill(name)
|
||||
return found["path"] if found else None
|
||||
except Exception:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def skill_pending_diff(record: Dict[str, Any]) -> str:
|
||||
@@ -263,12 +251,9 @@ def skill_pending_diff(record: Dict[str, Any]) -> str:
|
||||
name = payload.get("name", "")
|
||||
if action == "create":
|
||||
return payload.get("content") or ""
|
||||
if action == "remove_file":
|
||||
return f"remove file: {payload.get('file_path')} from skill '{name}'"
|
||||
if action == "delete":
|
||||
return f"delete skill '{name}'"
|
||||
if action not in {"edit", "patch", "write_file"}:
|
||||
return f"({action} on '{name}')"
|
||||
return {"remove_file": f"remove file: {payload.get('file_path')} from skill '{name}'",
|
||||
"delete": f"delete skill '{name}'"}.get(action, f"({action} on '{name}')")
|
||||
|
||||
# patch/write_file target a file inside the skill; edit always targets SKILL.md.
|
||||
target_label, current = "SKILL.md", ""
|
||||
@@ -276,11 +261,9 @@ def skill_pending_diff(record: Dict[str, Any]) -> str:
|
||||
if skill_dir:
|
||||
if action != "edit":
|
||||
target_label = payload.get("file_path") or "SKILL.md"
|
||||
try:
|
||||
with suppress(Exception):
|
||||
p = skill_dir / target_label
|
||||
current = p.read_text(encoding="utf-8") if p.exists() else ""
|
||||
except Exception:
|
||||
current = ""
|
||||
|
||||
if action == "patch":
|
||||
old_s, new_s = payload.get("old_string") or "", payload.get("new_string") or ""
|
||||
|
||||
+38
-56
@@ -10,6 +10,7 @@ from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import logging
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import Tuple
|
||||
|
||||
@@ -36,12 +37,8 @@ class _YbError(Exception):
|
||||
self.payload = {"success": False, "error": msg, **extra}
|
||||
|
||||
|
||||
def _err(msg: str) -> dict:
|
||||
return {"success": False, "error": msg}
|
||||
|
||||
|
||||
def _yb_tool(label: str):
|
||||
"""Handler decorator: ``fn(args)`` → tool_result; ``_YbError`` → its envelope; else logged ``_err``."""
|
||||
"""Handler decorator: ``fn(args)`` → tool_result; ``_YbError`` → its envelope; else logged envelope."""
|
||||
def deco(fn):
|
||||
@functools.wraps(fn)
|
||||
async def handler(args, **kw):
|
||||
@@ -51,18 +48,17 @@ def _yb_tool(label: str):
|
||||
return tool_result(exc.payload)
|
||||
except Exception as exc:
|
||||
logger.exception("[yuanbao_tools] %s error", label)
|
||||
return tool_result(_err(str(exc)))
|
||||
return tool_result(_YbError(str(exc)).payload)
|
||||
return handler
|
||||
return deco
|
||||
|
||||
|
||||
def _get_active_adapter():
|
||||
"""Lazy import to avoid ImportError when gateway.platforms.yuanbao is unavailable."""
|
||||
try:
|
||||
with suppress(ImportError):
|
||||
from gateway.platforms.yuanbao import get_active_adapter
|
||||
return get_active_adapter()
|
||||
except ImportError:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _adapter():
|
||||
@@ -73,11 +69,10 @@ def _adapter():
|
||||
|
||||
|
||||
def _session_env(name: str) -> str:
|
||||
try:
|
||||
with suppress(Exception):
|
||||
from gateway.session_context import get_session_env
|
||||
return get_session_env(name, "")
|
||||
except Exception:
|
||||
return ""
|
||||
return ""
|
||||
|
||||
|
||||
def _nick(m: dict, default: str = "") -> str:
|
||||
@@ -92,10 +87,8 @@ async def _members(adapter, group_code: str) -> list:
|
||||
|
||||
|
||||
async def _resolve_dm_recipient(adapter, group_code: str, name: str) -> Tuple[str, str]:
|
||||
"""Resolve ``name`` to (user_id, nickname) via the group member list.
|
||||
|
||||
>1 partial match raises with ``candidates`` for disambiguation instead of guessing.
|
||||
"""
|
||||
"""Resolve ``name`` to (user_id, nickname) via the group member list; >1 partial match raises
|
||||
with ``candidates`` for disambiguation instead of guessing."""
|
||||
if not group_code:
|
||||
raise _YbError("group_code is required when user_id is not provided")
|
||||
if not name:
|
||||
@@ -119,14 +112,12 @@ async def get_group_info(args) -> dict:
|
||||
"""查询群基本信息(群名、群主、成员数)。"""
|
||||
group_code = args.get("group_code", "")
|
||||
if not group_code:
|
||||
return _err("group_code is required")
|
||||
raise _YbError("group_code is required")
|
||||
gi = await _adapter().query_group_info(group_code)
|
||||
if gi is None:
|
||||
return _err("query_group_info returned None")
|
||||
raise _YbError("query_group_info returned None")
|
||||
return {
|
||||
"success": True,
|
||||
"group_code": group_code,
|
||||
"group_name": gi.get("group_name", ""),
|
||||
"success": True, "group_code": group_code, "group_name": gi.get("group_name", ""),
|
||||
"member_count": gi.get("member_count", 0),
|
||||
"owner": {"user_id": gi.get("owner_id", ""), "nickname": gi.get("owner_nickname", "")},
|
||||
"note": 'The group is called "派 (Pai)" in the app.',
|
||||
@@ -136,20 +127,18 @@ async def get_group_info(args) -> dict:
|
||||
@_yb_tool("query_group_members")
|
||||
async def query_group_members(args) -> dict:
|
||||
"""统一的群成员查询(对齐 TS query_session_members)。
|
||||
|
||||
action: find (按昵称模糊搜索; 无 name 时等同 list_all) / list_bots / list_all (默认).
|
||||
"""
|
||||
action: find (按昵称模糊搜索; 无 name 时等同 list_all) / list_bots / list_all (默认)."""
|
||||
group_code, name = args.get("group_code", ""), args.get("name", "")
|
||||
action = args.get("action", "list_all")
|
||||
if not group_code:
|
||||
return _err("group_code is required")
|
||||
raise _YbError("group_code is required")
|
||||
all_members = [
|
||||
{"user_id": m.get("user_id", ""), "nickname": _nick(m),
|
||||
"role": _USER_TYPE_LABEL.get(m.get("user_type", m.get("role", 0)), "unknown")}
|
||||
for m in await _members(_adapter(), group_code)
|
||||
]
|
||||
if not all_members:
|
||||
return _err("No members found in this group.")
|
||||
raise _YbError("No members found in this group.")
|
||||
|
||||
hint = {"mention_hint": MENTION_HINT} if args.get("mention", False) else {}
|
||||
|
||||
@@ -159,7 +148,7 @@ async def query_group_members(args) -> dict:
|
||||
if action == "list_bots":
|
||||
bots = [m for m in all_members if m["role"] in {"yuanbao_ai", "bot"}]
|
||||
if not bots:
|
||||
return _err("No bots found in this group.")
|
||||
raise _YbError("No bots found in this group.")
|
||||
return _listing(True, f"Found {len(bots)} bot(s).", bots)
|
||||
|
||||
if action == "find" and name:
|
||||
@@ -185,22 +174,21 @@ async def search_sticker(args) -> dict:
|
||||
matches = search_stickers(query or "", limit=safe_limit)
|
||||
return {
|
||||
"success": True, "query": query or "", "count": len(matches),
|
||||
"results": [{k: s.get(k, "") for k in ("sticker_id", "name", "description", "package_id")} for s in matches],
|
||||
"results": [{k: s.get(k, "") for k in ("sticker_id", "name", "description", "package_id")}
|
||||
for s in matches],
|
||||
}
|
||||
|
||||
|
||||
@_yb_tool("send_sticker")
|
||||
async def send_sticker(args) -> dict:
|
||||
"""向 chat_id(缺省取当前会话 HERMES_SESSION_CHAT_ID)发送一张内置贴纸(TIMFaceElem)。
|
||||
|
||||
``sticker``: 名称(如 "六六六")或 sticker_id(如 "278");为空时随机发送。
|
||||
``chat_id``: ``direct:{account_id}`` / ``group:{group_code}`` / 裸 account_id。
|
||||
"""
|
||||
``chat_id``: ``direct:{account_id}`` / ``group:{group_code}`` / 裸 account_id。"""
|
||||
from gateway.platforms.yuanbao_sticker import get_sticker_by_id, get_sticker_by_name, get_random_sticker
|
||||
|
||||
target = (args.get("chat_id", "") or "").strip() or _session_env("HERMES_SESSION_CHAT_ID")
|
||||
if not target:
|
||||
return _err("chat_id is required (no active yuanbao session detected)")
|
||||
raise _YbError("chat_id is required (no active yuanbao session detected)")
|
||||
adapter = _adapter()
|
||||
|
||||
raw = (args.get("sticker", "") or "").strip()
|
||||
@@ -209,13 +197,12 @@ async def send_sticker(args) -> dict:
|
||||
else:
|
||||
sticker_obj = (get_sticker_by_id(raw) if raw.isdigit() else None) or get_sticker_by_name(raw)
|
||||
if sticker_obj is None:
|
||||
return _err(f"Sticker not found: {raw!r}. Use search_sticker first to discover available stickers.")
|
||||
raise _YbError(f"Sticker not found: {raw!r}. Use search_sticker first to discover available stickers.")
|
||||
|
||||
result = await adapter.send_sticker(
|
||||
chat_id=target, sticker_name=sticker_obj.get("name", ""), reply_to=args.get("reply_to", "") or None,
|
||||
)
|
||||
result = await adapter.send_sticker(chat_id=target, sticker_name=sticker_obj.get("name", ""),
|
||||
reply_to=args.get("reply_to", "") or None)
|
||||
if not getattr(result, "success", False):
|
||||
return _err(getattr(result, "error", "send_sticker failed"))
|
||||
raise _YbError(getattr(result, "error", "send_sticker failed"))
|
||||
return {
|
||||
"success": True, "chat_id": target,
|
||||
"sticker": {"sticker_id": sticker_obj.get("sticker_id", ""), "name": sticker_obj.get("name", "")},
|
||||
@@ -226,18 +213,14 @@ async def send_sticker(args) -> dict:
|
||||
|
||||
@_yb_tool("send_dm")
|
||||
async def send_dm(args) -> dict:
|
||||
"""Send a DM to a group member, with optional media.
|
||||
|
||||
group_code defaults to the session's "group:<code>" chat_id. Without ``user_id`` the member
|
||||
list is searched by ``name`` (partial, case-insensitive; >1 match returns candidates).
|
||||
media_files items are {"path", "is_voice"} dicts or (path, is_voice) pairs; ``MEDIA:<path>``
|
||||
tags in the text count too. Partial media failures are reported in ``note``, not as failure.
|
||||
"""
|
||||
"""Send a DM to a group member, with optional media. group_code defaults to the session's
|
||||
"group:<code>" chat_id. Without ``user_id`` the member list is searched by ``name`` (partial,
|
||||
case-insensitive; >1 match returns candidates). media_files items are {"path", "is_voice"} dicts or
|
||||
(path, is_voice) pairs; ``MEDIA:<path>`` tags in the text count too. Partial media failures are
|
||||
reported in ``note``, not as failure."""
|
||||
group_code = args.get("group_code", "")
|
||||
if not group_code:
|
||||
chat_id = _session_env("HERMES_SESSION_CHAT_ID")
|
||||
if chat_id.startswith("group:"):
|
||||
group_code = chat_id.split(":", 1)[1]
|
||||
if not group_code and (chat_id := _session_env("HERMES_SESSION_CHAT_ID")).startswith("group:"):
|
||||
group_code = chat_id.split(":", 1)[1]
|
||||
|
||||
media_files = []
|
||||
for item in args.get("media_files") or []:
|
||||
@@ -248,19 +231,18 @@ async def send_dm(args) -> dict:
|
||||
|
||||
from gateway.platforms.base import BasePlatformAdapter
|
||||
embedded_media, message = BasePlatformAdapter.extract_media(args.get("message", ""))
|
||||
media_files.extend(embedded_media or [])
|
||||
media_files = BasePlatformAdapter.filter_media_delivery_paths(media_files)
|
||||
media_files = BasePlatformAdapter.filter_media_delivery_paths(media_files + list(embedded_media or []))
|
||||
|
||||
if not message and not media_files:
|
||||
return _err("message or media_files is required")
|
||||
raise _YbError("message or media_files is required")
|
||||
adapter = _adapter()
|
||||
|
||||
name, user_id = args.get("name", ""), args.get("user_id", "")
|
||||
resolved_user_id, resolved_nickname = user_id.strip() if user_id else "", name.strip()
|
||||
name = args.get("name", "")
|
||||
resolved_user_id, resolved_nickname = (args.get("user_id", "") or "").strip(), name.strip()
|
||||
if not resolved_user_id:
|
||||
resolved_user_id, resolved_nickname = await _resolve_dm_recipient(adapter, group_code, name)
|
||||
if not resolved_user_id:
|
||||
return _err("Could not resolve user_id")
|
||||
raise _YbError("Could not resolve user_id")
|
||||
|
||||
chat_id = f"direct:{resolved_user_id}"
|
||||
last_result = None
|
||||
@@ -277,9 +259,9 @@ async def send_dm(args) -> dict:
|
||||
errors.append(last_result.error or "media send failed")
|
||||
|
||||
if last_result is None:
|
||||
return _err("No deliverable text or media remained")
|
||||
raise _YbError("No deliverable text or media remained")
|
||||
if errors and not last_result.success:
|
||||
return _err("; ".join(errors))
|
||||
raise _YbError("; ".join(errors))
|
||||
|
||||
note = f'DM sent to "{resolved_nickname}" successfully.'
|
||||
if errors:
|
||||
|
||||
Reference in New Issue
Block a user