fix: 补全 vendor 快照(model_tools/utils/toolsets/gateway/cron/hermes_time)+ python-dotenv 依赖
- vendor_hermes.sh: 补全遗漏的单文件模块和目录 - pyproject.toml: 添加 python-dotenv 隐性依赖 - 全量 import 测试通过(216 .py, 180K 行) - firecrawl/fal_client 为可选工具,缺失时优雅跳过
This commit is contained in:
@@ -1 +0,0 @@
|
|||||||
16f9d020
|
|
||||||
@@ -1,8 +1,5 @@
|
|||||||
# MindOS CLI Vendor Snapshot
|
# MindOS CLI Vendor Snapshot
|
||||||
# 此目录包含从 Hermes 主仓打包的快照副本
|
source: hermes
|
||||||
# 运行时只 import 此副本,不依赖用户环境中的 hermes
|
commit: 16f9d020
|
||||||
|
|
||||||
source: mindOSv2/hermes
|
|
||||||
commit: f8a855a
|
|
||||||
snapshot_date: 2026-04-29
|
snapshot_date: 2026-04-29
|
||||||
snapshot_by: vendor_hermes.sh
|
snapshot_by: vendor_hermes.sh
|
||||||
|
|||||||
@@ -0,0 +1,42 @@
|
|||||||
|
"""
|
||||||
|
Cron job scheduling system for Hermes Agent.
|
||||||
|
|
||||||
|
This module provides scheduled task execution, allowing the agent to:
|
||||||
|
- Run automated tasks on schedules (cron expressions, intervals, one-shot)
|
||||||
|
- Self-schedule reminders and follow-up tasks
|
||||||
|
- Execute tasks in isolated sessions (no prior context)
|
||||||
|
|
||||||
|
Cron jobs are executed automatically by the gateway daemon:
|
||||||
|
hermes gateway install # Install as a user service
|
||||||
|
sudo hermes gateway install --system # Linux servers: boot-time system service
|
||||||
|
hermes gateway # Or run in foreground
|
||||||
|
|
||||||
|
The gateway ticks the scheduler every 60 seconds. A file lock prevents
|
||||||
|
duplicate execution if multiple processes overlap.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from cron.jobs import (
|
||||||
|
create_job,
|
||||||
|
get_job,
|
||||||
|
list_jobs,
|
||||||
|
remove_job,
|
||||||
|
update_job,
|
||||||
|
pause_job,
|
||||||
|
resume_job,
|
||||||
|
trigger_job,
|
||||||
|
JOBS_FILE,
|
||||||
|
)
|
||||||
|
from cron.scheduler import tick
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"create_job",
|
||||||
|
"get_job",
|
||||||
|
"list_jobs",
|
||||||
|
"remove_job",
|
||||||
|
"update_job",
|
||||||
|
"pause_job",
|
||||||
|
"resume_job",
|
||||||
|
"trigger_job",
|
||||||
|
"tick",
|
||||||
|
"JOBS_FILE",
|
||||||
|
]
|
||||||
@@ -0,0 +1,762 @@
|
|||||||
|
"""
|
||||||
|
Cron job storage and management.
|
||||||
|
|
||||||
|
Jobs are stored in ~/.hermes/cron/jobs.json
|
||||||
|
Output is saved to ~/.hermes/cron/output/{job_id}/{timestamp}.md
|
||||||
|
"""
|
||||||
|
|
||||||
|
import copy
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import tempfile
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime, timedelta
|
||||||
|
from pathlib import Path
|
||||||
|
from hermes_constants import get_hermes_home
|
||||||
|
from typing import Optional, Dict, List, Any
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
from hermes_time import now as _hermes_now
|
||||||
|
|
||||||
|
try:
|
||||||
|
from croniter import croniter
|
||||||
|
HAS_CRONITER = True
|
||||||
|
except ImportError:
|
||||||
|
HAS_CRONITER = False
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Configuration
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
HERMES_DIR = get_hermes_home().resolve()
|
||||||
|
CRON_DIR = HERMES_DIR / "cron"
|
||||||
|
JOBS_FILE = CRON_DIR / "jobs.json"
|
||||||
|
OUTPUT_DIR = CRON_DIR / "output"
|
||||||
|
ONESHOT_GRACE_SECONDS = 120
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_skill_list(skill: Optional[str] = None, skills: Optional[Any] = None) -> List[str]:
|
||||||
|
"""Normalize legacy/single-skill and multi-skill inputs into a unique ordered list."""
|
||||||
|
if skills is None:
|
||||||
|
raw_items = [skill] if skill else []
|
||||||
|
elif isinstance(skills, str):
|
||||||
|
raw_items = [skills]
|
||||||
|
else:
|
||||||
|
raw_items = list(skills)
|
||||||
|
|
||||||
|
normalized: List[str] = []
|
||||||
|
for item in raw_items:
|
||||||
|
text = str(item or "").strip()
|
||||||
|
if text and text not in normalized:
|
||||||
|
normalized.append(text)
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_skill_fields(job: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Return a job dict with canonical `skills` and legacy `skill` fields aligned."""
|
||||||
|
normalized = dict(job)
|
||||||
|
skills = _normalize_skill_list(normalized.get("skill"), normalized.get("skills"))
|
||||||
|
normalized["skills"] = skills
|
||||||
|
normalized["skill"] = skills[0] if skills else None
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
|
||||||
|
def _secure_dir(path: Path):
|
||||||
|
"""Set directory to owner-only access (0700). No-op on Windows."""
|
||||||
|
try:
|
||||||
|
os.chmod(path, 0o700)
|
||||||
|
except (OSError, NotImplementedError):
|
||||||
|
pass # Windows or other platforms where chmod is not supported
|
||||||
|
|
||||||
|
|
||||||
|
def _secure_file(path: Path):
|
||||||
|
"""Set file to owner-only read/write (0600). No-op on Windows."""
|
||||||
|
try:
|
||||||
|
if path.exists():
|
||||||
|
os.chmod(path, 0o600)
|
||||||
|
except (OSError, NotImplementedError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def ensure_dirs():
|
||||||
|
"""Ensure cron directories exist with secure permissions."""
|
||||||
|
CRON_DIR.mkdir(parents=True, exist_ok=True)
|
||||||
|
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||||
|
_secure_dir(CRON_DIR)
|
||||||
|
_secure_dir(OUTPUT_DIR)
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Schedule Parsing
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
def parse_duration(s: str) -> int:
|
||||||
|
"""
|
||||||
|
Parse duration string into minutes.
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
"30m" → 30
|
||||||
|
"2h" → 120
|
||||||
|
"1d" → 1440
|
||||||
|
"""
|
||||||
|
s = s.strip().lower()
|
||||||
|
match = re.match(r'^(\d+)\s*(m|min|mins|minute|minutes|h|hr|hrs|hour|hours|d|day|days)$', s)
|
||||||
|
if not match:
|
||||||
|
raise ValueError(f"Invalid duration: '{s}'. Use format like '30m', '2h', or '1d'")
|
||||||
|
|
||||||
|
value = int(match.group(1))
|
||||||
|
unit = match.group(2)[0] # First char: m, h, or d
|
||||||
|
|
||||||
|
multipliers = {'m': 1, 'h': 60, 'd': 1440}
|
||||||
|
return value * multipliers[unit]
|
||||||
|
|
||||||
|
|
||||||
|
def parse_schedule(schedule: str) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
Parse schedule string into structured format.
|
||||||
|
|
||||||
|
Returns dict with:
|
||||||
|
- kind: "once" | "interval" | "cron"
|
||||||
|
- For "once": "run_at" (ISO timestamp)
|
||||||
|
- For "interval": "minutes" (int)
|
||||||
|
- For "cron": "expr" (cron expression)
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
"30m" → once in 30 minutes
|
||||||
|
"2h" → once in 2 hours
|
||||||
|
"every 30m" → recurring every 30 minutes
|
||||||
|
"every 2h" → recurring every 2 hours
|
||||||
|
"0 9 * * *" → cron expression
|
||||||
|
"2026-02-03T14:00" → once at timestamp
|
||||||
|
"""
|
||||||
|
schedule = schedule.strip()
|
||||||
|
original = schedule
|
||||||
|
schedule_lower = schedule.lower()
|
||||||
|
|
||||||
|
# "every X" pattern → recurring interval
|
||||||
|
if schedule_lower.startswith("every "):
|
||||||
|
duration_str = schedule[6:].strip()
|
||||||
|
minutes = parse_duration(duration_str)
|
||||||
|
return {
|
||||||
|
"kind": "interval",
|
||||||
|
"minutes": minutes,
|
||||||
|
"display": f"every {minutes}m"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Check for cron expression (5 or 6 space-separated fields)
|
||||||
|
# Cron fields: minute hour day month weekday [year]
|
||||||
|
parts = schedule.split()
|
||||||
|
if len(parts) >= 5 and all(
|
||||||
|
re.match(r'^[\d\*\-,/]+$', p) for p in parts[:5]
|
||||||
|
):
|
||||||
|
if not HAS_CRONITER:
|
||||||
|
raise ValueError("Cron expressions require 'croniter' package. Install with: pip install croniter")
|
||||||
|
# Validate cron expression
|
||||||
|
try:
|
||||||
|
croniter(schedule)
|
||||||
|
except Exception as e:
|
||||||
|
raise ValueError(f"Invalid cron expression '{schedule}': {e}")
|
||||||
|
return {
|
||||||
|
"kind": "cron",
|
||||||
|
"expr": schedule,
|
||||||
|
"display": schedule
|
||||||
|
}
|
||||||
|
|
||||||
|
# ISO timestamp (contains T or looks like date)
|
||||||
|
if 'T' in schedule or re.match(r'^\d{4}-\d{2}-\d{2}', schedule):
|
||||||
|
try:
|
||||||
|
# Parse and validate
|
||||||
|
dt = datetime.fromisoformat(schedule.replace('Z', '+00:00'))
|
||||||
|
# Make naive timestamps timezone-aware at parse time so the stored
|
||||||
|
# value doesn't depend on the system timezone matching at check time.
|
||||||
|
if dt.tzinfo is None:
|
||||||
|
dt = dt.astimezone() # Interpret as local timezone
|
||||||
|
return {
|
||||||
|
"kind": "once",
|
||||||
|
"run_at": dt.isoformat(),
|
||||||
|
"display": f"once at {dt.strftime('%Y-%m-%d %H:%M')}"
|
||||||
|
}
|
||||||
|
except ValueError as e:
|
||||||
|
raise ValueError(f"Invalid timestamp '{schedule}': {e}")
|
||||||
|
|
||||||
|
# Duration like "30m", "2h", "1d" → one-shot from now
|
||||||
|
try:
|
||||||
|
minutes = parse_duration(schedule)
|
||||||
|
run_at = _hermes_now() + timedelta(minutes=minutes)
|
||||||
|
return {
|
||||||
|
"kind": "once",
|
||||||
|
"run_at": run_at.isoformat(),
|
||||||
|
"display": f"once in {original}"
|
||||||
|
}
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
raise ValueError(
|
||||||
|
f"Invalid schedule '{original}'. Use:\n"
|
||||||
|
f" - Duration: '30m', '2h', '1d' (one-shot)\n"
|
||||||
|
f" - Interval: 'every 30m', 'every 2h' (recurring)\n"
|
||||||
|
f" - Cron: '0 9 * * *' (cron expression)\n"
|
||||||
|
f" - Timestamp: '2026-02-03T14:00:00' (one-shot at time)"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _ensure_aware(dt: datetime) -> datetime:
|
||||||
|
"""Return a timezone-aware datetime in Hermes configured timezone.
|
||||||
|
|
||||||
|
Backward compatibility:
|
||||||
|
- Older stored timestamps may be naive.
|
||||||
|
- Naive values are interpreted as *system-local wall time* (the timezone
|
||||||
|
`datetime.now()` used when they were created), then converted to the
|
||||||
|
configured Hermes timezone.
|
||||||
|
|
||||||
|
This preserves relative ordering for legacy naive timestamps across
|
||||||
|
timezone changes and avoids false not-due results.
|
||||||
|
"""
|
||||||
|
target_tz = _hermes_now().tzinfo
|
||||||
|
if dt.tzinfo is None:
|
||||||
|
local_tz = datetime.now().astimezone().tzinfo
|
||||||
|
return dt.replace(tzinfo=local_tz).astimezone(target_tz)
|
||||||
|
return dt.astimezone(target_tz)
|
||||||
|
|
||||||
|
|
||||||
|
def _recoverable_oneshot_run_at(
|
||||||
|
schedule: Dict[str, Any],
|
||||||
|
now: datetime,
|
||||||
|
*,
|
||||||
|
last_run_at: Optional[str] = None,
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Return a one-shot run time if it is still eligible to fire.
|
||||||
|
|
||||||
|
One-shot jobs get a small grace window so jobs created a few seconds after
|
||||||
|
their requested minute still run on the next tick. Once a one-shot has
|
||||||
|
already run, it is never eligible again.
|
||||||
|
"""
|
||||||
|
if schedule.get("kind") != "once":
|
||||||
|
return None
|
||||||
|
if last_run_at:
|
||||||
|
return None
|
||||||
|
|
||||||
|
run_at = schedule.get("run_at")
|
||||||
|
if not run_at:
|
||||||
|
return None
|
||||||
|
|
||||||
|
run_at_dt = _ensure_aware(datetime.fromisoformat(run_at))
|
||||||
|
if run_at_dt >= now - timedelta(seconds=ONESHOT_GRACE_SECONDS):
|
||||||
|
return run_at
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _compute_grace_seconds(schedule: dict) -> int:
|
||||||
|
"""Compute how late a job can be and still catch up instead of fast-forwarding.
|
||||||
|
|
||||||
|
Uses half the schedule period, clamped between 120 seconds and 2 hours.
|
||||||
|
This ensures daily jobs can catch up if missed by up to 2 hours,
|
||||||
|
while frequent jobs (every 5-10 min) still fast-forward quickly.
|
||||||
|
"""
|
||||||
|
MIN_GRACE = 120
|
||||||
|
MAX_GRACE = 7200 # 2 hours
|
||||||
|
|
||||||
|
kind = schedule.get("kind")
|
||||||
|
|
||||||
|
if kind == "interval":
|
||||||
|
period_seconds = schedule.get("minutes", 1) * 60
|
||||||
|
grace = period_seconds // 2
|
||||||
|
return max(MIN_GRACE, min(grace, MAX_GRACE))
|
||||||
|
|
||||||
|
if kind == "cron" and HAS_CRONITER:
|
||||||
|
try:
|
||||||
|
now = _hermes_now()
|
||||||
|
cron = croniter(schedule["expr"], now)
|
||||||
|
first = cron.get_next(datetime)
|
||||||
|
second = cron.get_next(datetime)
|
||||||
|
period_seconds = int((second - first).total_seconds())
|
||||||
|
grace = period_seconds // 2
|
||||||
|
return max(MIN_GRACE, min(grace, MAX_GRACE))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return MIN_GRACE
|
||||||
|
|
||||||
|
|
||||||
|
def compute_next_run(schedule: Dict[str, Any], last_run_at: Optional[str] = None) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
Compute the next run time for a schedule.
|
||||||
|
|
||||||
|
Returns ISO timestamp string, or None if no more runs.
|
||||||
|
"""
|
||||||
|
now = _hermes_now()
|
||||||
|
|
||||||
|
if schedule["kind"] == "once":
|
||||||
|
return _recoverable_oneshot_run_at(schedule, now, last_run_at=last_run_at)
|
||||||
|
|
||||||
|
elif schedule["kind"] == "interval":
|
||||||
|
minutes = schedule["minutes"]
|
||||||
|
if last_run_at:
|
||||||
|
# Next run is last_run + interval
|
||||||
|
last = _ensure_aware(datetime.fromisoformat(last_run_at))
|
||||||
|
next_run = last + timedelta(minutes=minutes)
|
||||||
|
else:
|
||||||
|
# First run is now + interval
|
||||||
|
next_run = now + timedelta(minutes=minutes)
|
||||||
|
return next_run.isoformat()
|
||||||
|
|
||||||
|
elif schedule["kind"] == "cron":
|
||||||
|
if not HAS_CRONITER:
|
||||||
|
return None
|
||||||
|
cron = croniter(schedule["expr"], now)
|
||||||
|
next_run = cron.get_next(datetime)
|
||||||
|
return next_run.isoformat()
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Job CRUD Operations
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
def load_jobs() -> List[Dict[str, Any]]:
|
||||||
|
"""Load all jobs from storage."""
|
||||||
|
ensure_dirs()
|
||||||
|
if not JOBS_FILE.exists():
|
||||||
|
return []
|
||||||
|
|
||||||
|
try:
|
||||||
|
with open(JOBS_FILE, 'r', encoding='utf-8') as f:
|
||||||
|
data = json.load(f)
|
||||||
|
return data.get("jobs", [])
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
# Retry with strict=False to handle bare control chars in string values
|
||||||
|
try:
|
||||||
|
with open(JOBS_FILE, 'r', encoding='utf-8') as f:
|
||||||
|
data = json.loads(f.read(), strict=False)
|
||||||
|
jobs = data.get("jobs", [])
|
||||||
|
if jobs:
|
||||||
|
# Auto-repair: rewrite with proper escaping
|
||||||
|
save_jobs(jobs)
|
||||||
|
logger.warning("Auto-repaired jobs.json (had invalid control characters)")
|
||||||
|
return jobs
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to auto-repair jobs.json: %s", e)
|
||||||
|
raise RuntimeError(f"Cron database corrupted and unrepairable: {e}") from e
|
||||||
|
except IOError as e:
|
||||||
|
logger.error("IOError reading jobs.json: %s", e)
|
||||||
|
raise RuntimeError(f"Failed to read cron database: {e}") from e
|
||||||
|
|
||||||
|
|
||||||
|
def save_jobs(jobs: List[Dict[str, Any]]):
|
||||||
|
"""Save all jobs to storage."""
|
||||||
|
ensure_dirs()
|
||||||
|
fd, tmp_path = tempfile.mkstemp(dir=str(JOBS_FILE.parent), suffix='.tmp', prefix='.jobs_')
|
||||||
|
try:
|
||||||
|
with os.fdopen(fd, 'w', encoding='utf-8') as f:
|
||||||
|
json.dump({"jobs": jobs, "updated_at": _hermes_now().isoformat()}, f, indent=2)
|
||||||
|
f.flush()
|
||||||
|
os.fsync(f.fileno())
|
||||||
|
os.replace(tmp_path, JOBS_FILE)
|
||||||
|
_secure_file(JOBS_FILE)
|
||||||
|
except BaseException:
|
||||||
|
try:
|
||||||
|
os.unlink(tmp_path)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
def create_job(
|
||||||
|
prompt: str,
|
||||||
|
schedule: str,
|
||||||
|
name: Optional[str] = None,
|
||||||
|
repeat: Optional[int] = None,
|
||||||
|
deliver: Optional[str] = None,
|
||||||
|
origin: Optional[Dict[str, Any]] = None,
|
||||||
|
skill: Optional[str] = None,
|
||||||
|
skills: Optional[List[str]] = None,
|
||||||
|
model: Optional[str] = None,
|
||||||
|
provider: Optional[str] = None,
|
||||||
|
base_url: Optional[str] = None,
|
||||||
|
script: Optional[str] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
Create a new cron job.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt: The prompt to run (must be self-contained, or a task instruction when skill is set)
|
||||||
|
schedule: Schedule string (see parse_schedule)
|
||||||
|
name: Optional friendly name
|
||||||
|
repeat: How many times to run (None = forever, 1 = once)
|
||||||
|
deliver: Where to deliver output ("origin", "local", "telegram", etc.)
|
||||||
|
origin: Source info where job was created (for "origin" delivery)
|
||||||
|
skill: Optional legacy single skill name to load before running the prompt
|
||||||
|
skills: Optional ordered list of skills to load before running the prompt
|
||||||
|
model: Optional per-job model override
|
||||||
|
provider: Optional per-job provider override
|
||||||
|
base_url: Optional per-job base URL override
|
||||||
|
script: Optional path to a Python script whose stdout is injected into the
|
||||||
|
prompt each run. The script runs before the agent turn, and its output
|
||||||
|
is prepended as context. Useful for data collection / change detection.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The created job dict
|
||||||
|
"""
|
||||||
|
parsed_schedule = parse_schedule(schedule)
|
||||||
|
|
||||||
|
# Normalize repeat: treat 0 or negative values as None (infinite)
|
||||||
|
if repeat is not None and repeat <= 0:
|
||||||
|
repeat = None
|
||||||
|
|
||||||
|
# Auto-set repeat=1 for one-shot schedules if not specified
|
||||||
|
if parsed_schedule["kind"] == "once" and repeat is None:
|
||||||
|
repeat = 1
|
||||||
|
|
||||||
|
# Default delivery to origin if available, otherwise local
|
||||||
|
if deliver is None:
|
||||||
|
deliver = "origin" if origin else "local"
|
||||||
|
|
||||||
|
job_id = uuid.uuid4().hex[:12]
|
||||||
|
now = _hermes_now().isoformat()
|
||||||
|
|
||||||
|
normalized_skills = _normalize_skill_list(skill, skills)
|
||||||
|
normalized_model = str(model).strip() if isinstance(model, str) else None
|
||||||
|
normalized_provider = str(provider).strip() if isinstance(provider, str) else None
|
||||||
|
normalized_base_url = str(base_url).strip().rstrip("/") if isinstance(base_url, str) else None
|
||||||
|
normalized_model = normalized_model or None
|
||||||
|
normalized_provider = normalized_provider or None
|
||||||
|
normalized_base_url = normalized_base_url or None
|
||||||
|
normalized_script = str(script).strip() if isinstance(script, str) else None
|
||||||
|
normalized_script = normalized_script or None
|
||||||
|
|
||||||
|
label_source = (prompt or (normalized_skills[0] if normalized_skills else None)) or "cron job"
|
||||||
|
job = {
|
||||||
|
"id": job_id,
|
||||||
|
"name": name or label_source[:50].strip(),
|
||||||
|
"prompt": prompt,
|
||||||
|
"skills": normalized_skills,
|
||||||
|
"skill": normalized_skills[0] if normalized_skills else None,
|
||||||
|
"model": normalized_model,
|
||||||
|
"provider": normalized_provider,
|
||||||
|
"base_url": normalized_base_url,
|
||||||
|
"script": normalized_script,
|
||||||
|
"schedule": parsed_schedule,
|
||||||
|
"schedule_display": parsed_schedule.get("display", schedule),
|
||||||
|
"repeat": {
|
||||||
|
"times": repeat, # None = forever
|
||||||
|
"completed": 0
|
||||||
|
},
|
||||||
|
"enabled": True,
|
||||||
|
"state": "scheduled",
|
||||||
|
"paused_at": None,
|
||||||
|
"paused_reason": None,
|
||||||
|
"created_at": now,
|
||||||
|
"next_run_at": compute_next_run(parsed_schedule),
|
||||||
|
"last_run_at": None,
|
||||||
|
"last_status": None,
|
||||||
|
"last_error": None,
|
||||||
|
"last_delivery_error": None,
|
||||||
|
# Delivery configuration
|
||||||
|
"deliver": deliver,
|
||||||
|
"origin": origin, # Tracks where job was created for "origin" delivery
|
||||||
|
}
|
||||||
|
|
||||||
|
jobs = load_jobs()
|
||||||
|
jobs.append(job)
|
||||||
|
save_jobs(jobs)
|
||||||
|
|
||||||
|
return job
|
||||||
|
|
||||||
|
|
||||||
|
def get_job(job_id: str) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Get a job by ID."""
|
||||||
|
jobs = load_jobs()
|
||||||
|
for job in jobs:
|
||||||
|
if job["id"] == job_id:
|
||||||
|
return _apply_skill_fields(job)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def list_jobs(include_disabled: bool = False) -> List[Dict[str, Any]]:
|
||||||
|
"""List all jobs, optionally including disabled ones."""
|
||||||
|
jobs = [_apply_skill_fields(j) for j in load_jobs()]
|
||||||
|
if not include_disabled:
|
||||||
|
jobs = [j for j in jobs if j.get("enabled", True)]
|
||||||
|
return jobs
|
||||||
|
|
||||||
|
|
||||||
|
def update_job(job_id: str, updates: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Update a job by ID, refreshing derived schedule fields when needed."""
|
||||||
|
jobs = load_jobs()
|
||||||
|
for i, job in enumerate(jobs):
|
||||||
|
if job["id"] != job_id:
|
||||||
|
continue
|
||||||
|
|
||||||
|
updated = _apply_skill_fields({**job, **updates})
|
||||||
|
schedule_changed = "schedule" in updates
|
||||||
|
|
||||||
|
if "skills" in updates or "skill" in updates:
|
||||||
|
normalized_skills = _normalize_skill_list(updated.get("skill"), updated.get("skills"))
|
||||||
|
updated["skills"] = normalized_skills
|
||||||
|
updated["skill"] = normalized_skills[0] if normalized_skills else None
|
||||||
|
|
||||||
|
if schedule_changed:
|
||||||
|
updated_schedule = updated["schedule"]
|
||||||
|
updated["schedule_display"] = updates.get(
|
||||||
|
"schedule_display",
|
||||||
|
updated_schedule.get("display", updated.get("schedule_display")),
|
||||||
|
)
|
||||||
|
if updated.get("state") != "paused":
|
||||||
|
updated["next_run_at"] = compute_next_run(updated_schedule)
|
||||||
|
|
||||||
|
if updated.get("enabled", True) and updated.get("state") != "paused" and not updated.get("next_run_at"):
|
||||||
|
updated["next_run_at"] = compute_next_run(updated["schedule"])
|
||||||
|
|
||||||
|
jobs[i] = updated
|
||||||
|
save_jobs(jobs)
|
||||||
|
return _apply_skill_fields(jobs[i])
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def pause_job(job_id: str, reason: Optional[str] = None) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Pause a job without deleting it."""
|
||||||
|
return update_job(
|
||||||
|
job_id,
|
||||||
|
{
|
||||||
|
"enabled": False,
|
||||||
|
"state": "paused",
|
||||||
|
"paused_at": _hermes_now().isoformat(),
|
||||||
|
"paused_reason": reason,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def resume_job(job_id: str) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Resume a paused job and compute the next future run from now."""
|
||||||
|
job = get_job(job_id)
|
||||||
|
if not job:
|
||||||
|
return None
|
||||||
|
|
||||||
|
next_run_at = compute_next_run(job["schedule"])
|
||||||
|
return update_job(
|
||||||
|
job_id,
|
||||||
|
{
|
||||||
|
"enabled": True,
|
||||||
|
"state": "scheduled",
|
||||||
|
"paused_at": None,
|
||||||
|
"paused_reason": None,
|
||||||
|
"next_run_at": next_run_at,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def trigger_job(job_id: str) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Schedule a job to run on the next scheduler tick."""
|
||||||
|
job = get_job(job_id)
|
||||||
|
if not job:
|
||||||
|
return None
|
||||||
|
return update_job(
|
||||||
|
job_id,
|
||||||
|
{
|
||||||
|
"enabled": True,
|
||||||
|
"state": "scheduled",
|
||||||
|
"paused_at": None,
|
||||||
|
"paused_reason": None,
|
||||||
|
"next_run_at": _hermes_now().isoformat(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def remove_job(job_id: str) -> bool:
|
||||||
|
"""Remove a job by ID."""
|
||||||
|
jobs = load_jobs()
|
||||||
|
original_len = len(jobs)
|
||||||
|
jobs = [j for j in jobs if j["id"] != job_id]
|
||||||
|
if len(jobs) < original_len:
|
||||||
|
save_jobs(jobs)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def mark_job_run(job_id: str, success: bool, error: Optional[str] = None,
|
||||||
|
delivery_error: Optional[str] = None):
|
||||||
|
"""
|
||||||
|
Mark a job as having been run.
|
||||||
|
|
||||||
|
Updates last_run_at, last_status, increments completed count,
|
||||||
|
computes next_run_at, and auto-deletes if repeat limit reached.
|
||||||
|
|
||||||
|
``delivery_error`` is tracked separately from the agent error — a job
|
||||||
|
can succeed (agent produced output) but fail delivery (platform down).
|
||||||
|
"""
|
||||||
|
jobs = load_jobs()
|
||||||
|
for i, job in enumerate(jobs):
|
||||||
|
if job["id"] == job_id:
|
||||||
|
now = _hermes_now().isoformat()
|
||||||
|
job["last_run_at"] = now
|
||||||
|
job["last_status"] = "ok" if success else "error"
|
||||||
|
job["last_error"] = error if not success else None
|
||||||
|
# Track delivery failures separately — cleared on successful delivery
|
||||||
|
job["last_delivery_error"] = delivery_error
|
||||||
|
|
||||||
|
# Increment completed count
|
||||||
|
if job.get("repeat"):
|
||||||
|
job["repeat"]["completed"] = job["repeat"].get("completed", 0) + 1
|
||||||
|
|
||||||
|
# Check if we've hit the repeat limit
|
||||||
|
times = job["repeat"].get("times")
|
||||||
|
completed = job["repeat"]["completed"]
|
||||||
|
if times is not None and times > 0 and completed >= times:
|
||||||
|
# Remove the job (limit reached)
|
||||||
|
jobs.pop(i)
|
||||||
|
save_jobs(jobs)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Compute next run
|
||||||
|
job["next_run_at"] = compute_next_run(job["schedule"], now)
|
||||||
|
|
||||||
|
# If no next run (one-shot completed), disable
|
||||||
|
if job["next_run_at"] is None:
|
||||||
|
job["enabled"] = False
|
||||||
|
job["state"] = "completed"
|
||||||
|
elif job.get("state") != "paused":
|
||||||
|
job["state"] = "scheduled"
|
||||||
|
|
||||||
|
save_jobs(jobs)
|
||||||
|
return
|
||||||
|
|
||||||
|
logger.warning("mark_job_run: job_id %s not found, skipping save", job_id)
|
||||||
|
|
||||||
|
|
||||||
|
def advance_next_run(job_id: str) -> bool:
|
||||||
|
"""Preemptively advance next_run_at for a recurring job before execution.
|
||||||
|
|
||||||
|
Call this BEFORE run_job() so that if the process crashes mid-execution,
|
||||||
|
the job won't re-fire on the next gateway restart. This converts the
|
||||||
|
scheduler from at-least-once to at-most-once for recurring jobs — missing
|
||||||
|
one run is far better than firing dozens of times in a crash loop.
|
||||||
|
|
||||||
|
One-shot jobs are left unchanged so they can still retry on restart.
|
||||||
|
|
||||||
|
Returns True if next_run_at was advanced, False otherwise.
|
||||||
|
"""
|
||||||
|
jobs = load_jobs()
|
||||||
|
for job in jobs:
|
||||||
|
if job["id"] == job_id:
|
||||||
|
kind = job.get("schedule", {}).get("kind")
|
||||||
|
if kind not in ("cron", "interval"):
|
||||||
|
return False
|
||||||
|
now = _hermes_now().isoformat()
|
||||||
|
new_next = compute_next_run(job["schedule"], now)
|
||||||
|
if new_next and new_next != job.get("next_run_at"):
|
||||||
|
job["next_run_at"] = new_next
|
||||||
|
save_jobs(jobs)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def get_due_jobs() -> List[Dict[str, Any]]:
|
||||||
|
"""Get all jobs that are due to run now.
|
||||||
|
|
||||||
|
For recurring jobs (cron/interval), if the scheduled time is stale
|
||||||
|
(more than one period in the past, e.g. because the gateway was down),
|
||||||
|
the job is fast-forwarded to the next future run instead of firing
|
||||||
|
immediately. This prevents a burst of missed jobs on gateway restart.
|
||||||
|
"""
|
||||||
|
now = _hermes_now()
|
||||||
|
raw_jobs = load_jobs()
|
||||||
|
jobs = [_apply_skill_fields(j) for j in copy.deepcopy(raw_jobs)]
|
||||||
|
due = []
|
||||||
|
needs_save = False
|
||||||
|
|
||||||
|
for job in jobs:
|
||||||
|
if not job.get("enabled", True):
|
||||||
|
continue
|
||||||
|
|
||||||
|
next_run = job.get("next_run_at")
|
||||||
|
if not next_run:
|
||||||
|
recovered_next = _recoverable_oneshot_run_at(
|
||||||
|
job.get("schedule", {}),
|
||||||
|
now,
|
||||||
|
last_run_at=job.get("last_run_at"),
|
||||||
|
)
|
||||||
|
if not recovered_next:
|
||||||
|
continue
|
||||||
|
|
||||||
|
job["next_run_at"] = recovered_next
|
||||||
|
next_run = recovered_next
|
||||||
|
logger.info(
|
||||||
|
"Job '%s' had no next_run_at; recovering one-shot run at %s",
|
||||||
|
job.get("name", job["id"]),
|
||||||
|
recovered_next,
|
||||||
|
)
|
||||||
|
for rj in raw_jobs:
|
||||||
|
if rj["id"] == job["id"]:
|
||||||
|
rj["next_run_at"] = recovered_next
|
||||||
|
needs_save = True
|
||||||
|
break
|
||||||
|
|
||||||
|
next_run_dt = _ensure_aware(datetime.fromisoformat(next_run))
|
||||||
|
if next_run_dt <= now:
|
||||||
|
schedule = job.get("schedule", {})
|
||||||
|
kind = schedule.get("kind")
|
||||||
|
|
||||||
|
# For recurring jobs, check if the scheduled time is stale
|
||||||
|
# (gateway was down and missed the window). Fast-forward to
|
||||||
|
# the next future occurrence instead of firing a stale run.
|
||||||
|
grace = _compute_grace_seconds(schedule)
|
||||||
|
if kind in ("cron", "interval") and (now - next_run_dt).total_seconds() > grace:
|
||||||
|
# Job is past its catch-up grace window — this is a stale missed run.
|
||||||
|
# Grace scales with schedule period: daily=2h, hourly=30m, 10min=5m.
|
||||||
|
new_next = compute_next_run(schedule, now.isoformat())
|
||||||
|
if new_next:
|
||||||
|
logger.info(
|
||||||
|
"Job '%s' missed its scheduled time (%s, grace=%ds). "
|
||||||
|
"Fast-forwarding to next run: %s",
|
||||||
|
job.get("name", job["id"]),
|
||||||
|
next_run,
|
||||||
|
grace,
|
||||||
|
new_next,
|
||||||
|
)
|
||||||
|
# Update the job in storage
|
||||||
|
for rj in raw_jobs:
|
||||||
|
if rj["id"] == job["id"]:
|
||||||
|
rj["next_run_at"] = new_next
|
||||||
|
needs_save = True
|
||||||
|
break
|
||||||
|
continue # Skip this run
|
||||||
|
|
||||||
|
due.append(job)
|
||||||
|
|
||||||
|
if needs_save:
|
||||||
|
save_jobs(raw_jobs)
|
||||||
|
|
||||||
|
return due
|
||||||
|
|
||||||
|
|
||||||
|
def save_job_output(job_id: str, output: str):
|
||||||
|
"""Save job output to file."""
|
||||||
|
ensure_dirs()
|
||||||
|
job_output_dir = OUTPUT_DIR / job_id
|
||||||
|
job_output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
_secure_dir(job_output_dir)
|
||||||
|
|
||||||
|
timestamp = _hermes_now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||||
|
output_file = job_output_dir / f"{timestamp}.md"
|
||||||
|
|
||||||
|
fd, tmp_path = tempfile.mkstemp(dir=str(job_output_dir), suffix='.tmp', prefix='.output_')
|
||||||
|
try:
|
||||||
|
with os.fdopen(fd, 'w', encoding='utf-8') as f:
|
||||||
|
f.write(output)
|
||||||
|
f.flush()
|
||||||
|
os.fsync(f.fileno())
|
||||||
|
os.replace(tmp_path, output_file)
|
||||||
|
_secure_file(output_file)
|
||||||
|
except BaseException:
|
||||||
|
try:
|
||||||
|
os.unlink(tmp_path)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
raise
|
||||||
|
|
||||||
|
return output_file
|
||||||
@@ -0,0 +1,992 @@
|
|||||||
|
"""
|
||||||
|
Cron job scheduler - executes due jobs.
|
||||||
|
|
||||||
|
Provides tick() which checks for due jobs and runs them. The gateway
|
||||||
|
calls this every 60 seconds from a background thread.
|
||||||
|
|
||||||
|
Uses a file-based lock (~/.hermes/cron/.tick.lock) so only one tick
|
||||||
|
runs at a time if multiple processes overlap.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import concurrent.futures
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import subprocess
|
||||||
|
import sys
|
||||||
|
|
||||||
|
# fcntl is Unix-only; on Windows use msvcrt for file locking
|
||||||
|
try:
|
||||||
|
import fcntl
|
||||||
|
except ImportError:
|
||||||
|
fcntl = None
|
||||||
|
try:
|
||||||
|
import msvcrt
|
||||||
|
except ImportError:
|
||||||
|
msvcrt = None
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
# Add parent directory to path for imports BEFORE repo-level imports.
|
||||||
|
# Without this, standalone invocations (e.g. after `hermes update` reloads
|
||||||
|
# the module) fail with ModuleNotFoundError for hermes_time et al.
|
||||||
|
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||||
|
|
||||||
|
from hermes_constants import get_hermes_home
|
||||||
|
from hermes_cli.config import load_config
|
||||||
|
from hermes_time import now as _hermes_now
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Valid delivery platforms — used to validate user-supplied platform names
|
||||||
|
# in cron delivery targets, preventing env var enumeration via crafted names.
|
||||||
|
_KNOWN_DELIVERY_PLATFORMS = frozenset({
|
||||||
|
"telegram", "discord", "slack", "whatsapp", "signal",
|
||||||
|
"matrix", "mattermost", "homeassistant", "dingtalk", "feishu",
|
||||||
|
"wecom", "wecom_callback", "weixin", "sms", "email", "webhook", "bluebubbles",
|
||||||
|
"qqbot",
|
||||||
|
})
|
||||||
|
|
||||||
|
from cron.jobs import get_due_jobs, mark_job_run, save_job_output, advance_next_run
|
||||||
|
|
||||||
|
# Sentinel: when a cron agent has nothing new to report, it can start its
|
||||||
|
# response with this marker to suppress delivery. Output is still saved
|
||||||
|
# locally for audit.
|
||||||
|
SILENT_MARKER = "[SILENT]"
|
||||||
|
|
||||||
|
# Resolve Hermes home directory (respects HERMES_HOME override)
|
||||||
|
_hermes_home = get_hermes_home()
|
||||||
|
|
||||||
|
# File-based lock prevents concurrent ticks from gateway + daemon + systemd timer
|
||||||
|
_LOCK_DIR = _hermes_home / "cron"
|
||||||
|
_LOCK_FILE = _LOCK_DIR / ".tick.lock"
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_origin(job: dict) -> Optional[dict]:
|
||||||
|
"""Extract origin info from a job, preserving any extra routing metadata."""
|
||||||
|
origin = job.get("origin")
|
||||||
|
if not origin:
|
||||||
|
return None
|
||||||
|
platform = origin.get("platform")
|
||||||
|
chat_id = origin.get("chat_id")
|
||||||
|
if platform and chat_id:
|
||||||
|
return origin
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_delivery_target(job: dict) -> Optional[dict]:
|
||||||
|
"""Resolve the concrete auto-delivery target for a cron job, if any."""
|
||||||
|
deliver = job.get("deliver", "local")
|
||||||
|
origin = _resolve_origin(job)
|
||||||
|
|
||||||
|
if deliver == "local":
|
||||||
|
return None
|
||||||
|
|
||||||
|
if deliver == "origin":
|
||||||
|
if origin:
|
||||||
|
return {
|
||||||
|
"platform": origin["platform"],
|
||||||
|
"chat_id": str(origin["chat_id"]),
|
||||||
|
"thread_id": origin.get("thread_id"),
|
||||||
|
}
|
||||||
|
# Origin missing (e.g. job created via API/script) — try each
|
||||||
|
# platform's home channel as a fallback instead of silently dropping.
|
||||||
|
for platform_name in ("matrix", "telegram", "discord", "slack", "bluebubbles"):
|
||||||
|
chat_id = os.getenv(f"{platform_name.upper()}_HOME_CHANNEL", "")
|
||||||
|
if chat_id:
|
||||||
|
logger.info(
|
||||||
|
"Job '%s' has deliver=origin but no origin; falling back to %s home channel",
|
||||||
|
job.get("name", job.get("id", "?")),
|
||||||
|
platform_name,
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"platform": platform_name,
|
||||||
|
"chat_id": chat_id,
|
||||||
|
"thread_id": None,
|
||||||
|
}
|
||||||
|
return None
|
||||||
|
|
||||||
|
if ":" in deliver:
|
||||||
|
platform_name, rest = deliver.split(":", 1)
|
||||||
|
platform_key = platform_name.lower()
|
||||||
|
|
||||||
|
from tools.send_message_tool import _parse_target_ref
|
||||||
|
|
||||||
|
parsed_chat_id, parsed_thread_id, is_explicit = _parse_target_ref(platform_key, rest)
|
||||||
|
if is_explicit:
|
||||||
|
chat_id, thread_id = parsed_chat_id, parsed_thread_id
|
||||||
|
else:
|
||||||
|
chat_id, thread_id = rest, None
|
||||||
|
|
||||||
|
# Resolve human-friendly labels like "Alice (dm)" to real IDs.
|
||||||
|
try:
|
||||||
|
from gateway.channel_directory import resolve_channel_name
|
||||||
|
resolved = resolve_channel_name(platform_key, chat_id)
|
||||||
|
if resolved:
|
||||||
|
parsed_chat_id, parsed_thread_id, resolved_is_explicit = _parse_target_ref(platform_key, resolved)
|
||||||
|
if resolved_is_explicit:
|
||||||
|
chat_id, thread_id = parsed_chat_id, parsed_thread_id
|
||||||
|
else:
|
||||||
|
chat_id = resolved
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return {
|
||||||
|
"platform": platform_name,
|
||||||
|
"chat_id": chat_id,
|
||||||
|
"thread_id": thread_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
platform_name = deliver
|
||||||
|
if origin and origin.get("platform") == platform_name:
|
||||||
|
return {
|
||||||
|
"platform": platform_name,
|
||||||
|
"chat_id": str(origin["chat_id"]),
|
||||||
|
"thread_id": origin.get("thread_id"),
|
||||||
|
}
|
||||||
|
|
||||||
|
if platform_name.lower() not in _KNOWN_DELIVERY_PLATFORMS:
|
||||||
|
return None
|
||||||
|
chat_id = os.getenv(f"{platform_name.upper()}_HOME_CHANNEL", "")
|
||||||
|
if not chat_id:
|
||||||
|
return None
|
||||||
|
|
||||||
|
return {
|
||||||
|
"platform": platform_name,
|
||||||
|
"chat_id": chat_id,
|
||||||
|
"thread_id": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# Media extension sets — keep in sync with gateway/platforms/base.py:_process_message_background
|
||||||
|
_AUDIO_EXTS = frozenset({'.ogg', '.opus', '.mp3', '.wav', '.m4a'})
|
||||||
|
_VIDEO_EXTS = frozenset({'.mp4', '.mov', '.avi', '.mkv', '.webm', '.3gp'})
|
||||||
|
_IMAGE_EXTS = frozenset({'.jpg', '.jpeg', '.png', '.webp', '.gif'})
|
||||||
|
|
||||||
|
|
||||||
|
def _send_media_via_adapter(adapter, chat_id: str, media_files: list, metadata: dict | None, loop, job: dict) -> None:
|
||||||
|
"""Send extracted MEDIA files as native platform attachments via a live adapter.
|
||||||
|
|
||||||
|
Routes each file to the appropriate adapter method (send_voice, send_image_file,
|
||||||
|
send_video, send_document) based on file extension — mirroring the routing logic
|
||||||
|
in ``BasePlatformAdapter._process_message_background``.
|
||||||
|
"""
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
for media_path, _is_voice in media_files:
|
||||||
|
try:
|
||||||
|
ext = Path(media_path).suffix.lower()
|
||||||
|
if ext in _AUDIO_EXTS:
|
||||||
|
coro = adapter.send_voice(chat_id=chat_id, audio_path=media_path, metadata=metadata)
|
||||||
|
elif ext in _VIDEO_EXTS:
|
||||||
|
coro = adapter.send_video(chat_id=chat_id, video_path=media_path, metadata=metadata)
|
||||||
|
elif ext in _IMAGE_EXTS:
|
||||||
|
coro = adapter.send_image_file(chat_id=chat_id, image_path=media_path, metadata=metadata)
|
||||||
|
else:
|
||||||
|
coro = adapter.send_document(chat_id=chat_id, file_path=media_path, metadata=metadata)
|
||||||
|
|
||||||
|
future = asyncio.run_coroutine_threadsafe(coro, loop)
|
||||||
|
result = future.result(timeout=30)
|
||||||
|
if result and not getattr(result, "success", True):
|
||||||
|
logger.warning(
|
||||||
|
"Job '%s': media send failed for %s: %s",
|
||||||
|
job.get("id", "?"), media_path, getattr(result, "error", "unknown"),
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Job '%s': failed to send media %s: %s", job.get("id", "?"), media_path, e)
|
||||||
|
|
||||||
|
|
||||||
|
def _deliver_result(job: dict, content: str, adapters=None, loop=None) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
Deliver job output to the configured target (origin chat, specific platform, etc.).
|
||||||
|
|
||||||
|
When ``adapters`` and ``loop`` are provided (gateway is running), tries to
|
||||||
|
use the live adapter first — this supports E2EE rooms (e.g. Matrix) where
|
||||||
|
the standalone HTTP path cannot encrypt. Falls back to standalone send if
|
||||||
|
the adapter path fails or is unavailable.
|
||||||
|
|
||||||
|
Returns None on success, or an error string on failure.
|
||||||
|
"""
|
||||||
|
target = _resolve_delivery_target(job)
|
||||||
|
if not target:
|
||||||
|
if job.get("deliver", "local") != "local":
|
||||||
|
msg = f"no delivery target resolved for deliver={job.get('deliver', 'local')}"
|
||||||
|
logger.warning("Job '%s': %s", job["id"], msg)
|
||||||
|
return msg
|
||||||
|
return None # local-only jobs don't deliver — not a failure
|
||||||
|
|
||||||
|
platform_name = target["platform"]
|
||||||
|
chat_id = target["chat_id"]
|
||||||
|
thread_id = target.get("thread_id")
|
||||||
|
|
||||||
|
# Diagnostic: log thread_id for topic-aware delivery debugging
|
||||||
|
origin = job.get("origin") or {}
|
||||||
|
origin_thread = origin.get("thread_id")
|
||||||
|
if origin_thread and not thread_id:
|
||||||
|
logger.warning(
|
||||||
|
"Job '%s': origin has thread_id=%s but delivery target lost it "
|
||||||
|
"(deliver=%s, target=%s)",
|
||||||
|
job["id"], origin_thread, job.get("deliver", "local"), target,
|
||||||
|
)
|
||||||
|
elif thread_id:
|
||||||
|
logger.debug(
|
||||||
|
"Job '%s': delivering to %s:%s thread_id=%s",
|
||||||
|
job["id"], platform_name, chat_id, thread_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
from tools.send_message_tool import _send_to_platform
|
||||||
|
from gateway.config import load_gateway_config, Platform
|
||||||
|
|
||||||
|
platform_map = {
|
||||||
|
"telegram": Platform.TELEGRAM,
|
||||||
|
"discord": Platform.DISCORD,
|
||||||
|
"slack": Platform.SLACK,
|
||||||
|
"whatsapp": Platform.WHATSAPP,
|
||||||
|
"signal": Platform.SIGNAL,
|
||||||
|
"matrix": Platform.MATRIX,
|
||||||
|
"mattermost": Platform.MATTERMOST,
|
||||||
|
"homeassistant": Platform.HOMEASSISTANT,
|
||||||
|
"dingtalk": Platform.DINGTALK,
|
||||||
|
"feishu": Platform.FEISHU,
|
||||||
|
"wecom": Platform.WECOM,
|
||||||
|
"wecom_callback": Platform.WECOM_CALLBACK,
|
||||||
|
"weixin": Platform.WEIXIN,
|
||||||
|
"email": Platform.EMAIL,
|
||||||
|
"sms": Platform.SMS,
|
||||||
|
"bluebubbles": Platform.BLUEBUBBLES,
|
||||||
|
"qqbot": Platform.QQBOT,
|
||||||
|
}
|
||||||
|
platform = platform_map.get(platform_name.lower())
|
||||||
|
if not platform:
|
||||||
|
msg = f"unknown platform '{platform_name}'"
|
||||||
|
logger.warning("Job '%s': %s", job["id"], msg)
|
||||||
|
return msg
|
||||||
|
|
||||||
|
try:
|
||||||
|
config = load_gateway_config()
|
||||||
|
except Exception as e:
|
||||||
|
msg = f"failed to load gateway config: {e}"
|
||||||
|
logger.error("Job '%s': %s", job["id"], msg)
|
||||||
|
return msg
|
||||||
|
|
||||||
|
pconfig = config.platforms.get(platform)
|
||||||
|
if not pconfig or not pconfig.enabled:
|
||||||
|
msg = f"platform '{platform_name}' not configured/enabled"
|
||||||
|
logger.warning("Job '%s': %s", job["id"], msg)
|
||||||
|
return msg
|
||||||
|
|
||||||
|
# Optionally wrap the content with a header/footer so the user knows this
|
||||||
|
# is a cron delivery. Wrapping is on by default; set cron.wrap_response: false
|
||||||
|
# in config.yaml for clean output.
|
||||||
|
wrap_response = True
|
||||||
|
try:
|
||||||
|
user_cfg = load_config()
|
||||||
|
wrap_response = user_cfg.get("cron", {}).get("wrap_response", True)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if wrap_response:
|
||||||
|
task_name = job.get("name", job["id"])
|
||||||
|
delivery_content = (
|
||||||
|
f"Cronjob Response: {task_name}\n"
|
||||||
|
f"-------------\n\n"
|
||||||
|
f"{content}\n\n"
|
||||||
|
f"Note: The agent cannot see this message, and therefore cannot respond to it."
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
delivery_content = content
|
||||||
|
|
||||||
|
# Extract MEDIA: tags so attachments are forwarded as files, not raw text
|
||||||
|
from gateway.platforms.base import BasePlatformAdapter
|
||||||
|
media_files, cleaned_delivery_content = BasePlatformAdapter.extract_media(delivery_content)
|
||||||
|
|
||||||
|
# Prefer the live adapter when the gateway is running — this supports E2EE
|
||||||
|
# rooms (e.g. Matrix) where the standalone HTTP path cannot encrypt.
|
||||||
|
runtime_adapter = (adapters or {}).get(platform)
|
||||||
|
if runtime_adapter is not None and loop is not None and getattr(loop, "is_running", lambda: False)():
|
||||||
|
send_metadata = {"thread_id": thread_id} if thread_id else None
|
||||||
|
try:
|
||||||
|
# Send cleaned text (MEDIA tags stripped) — not the raw content
|
||||||
|
text_to_send = cleaned_delivery_content.strip()
|
||||||
|
adapter_ok = True
|
||||||
|
if text_to_send:
|
||||||
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
|
runtime_adapter.send(chat_id, text_to_send, metadata=send_metadata),
|
||||||
|
loop,
|
||||||
|
)
|
||||||
|
send_result = future.result(timeout=60)
|
||||||
|
if send_result and not getattr(send_result, "success", True):
|
||||||
|
err = getattr(send_result, "error", "unknown")
|
||||||
|
logger.warning(
|
||||||
|
"Job '%s': live adapter send to %s:%s failed (%s), falling back to standalone",
|
||||||
|
job["id"], platform_name, chat_id, err,
|
||||||
|
)
|
||||||
|
adapter_ok = False # fall through to standalone path
|
||||||
|
|
||||||
|
# Send extracted media files as native attachments via the live adapter
|
||||||
|
if adapter_ok and media_files:
|
||||||
|
_send_media_via_adapter(runtime_adapter, chat_id, media_files, send_metadata, loop, job)
|
||||||
|
|
||||||
|
if adapter_ok:
|
||||||
|
logger.info("Job '%s': delivered to %s:%s via live adapter", job["id"], platform_name, chat_id)
|
||||||
|
return None
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"Job '%s': live adapter delivery to %s:%s failed (%s), falling back to standalone",
|
||||||
|
job["id"], platform_name, chat_id, e,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Standalone path: run the async send in a fresh event loop (safe from any thread)
|
||||||
|
coro = _send_to_platform(platform, pconfig, chat_id, cleaned_delivery_content, thread_id=thread_id, media_files=media_files)
|
||||||
|
try:
|
||||||
|
result = asyncio.run(coro)
|
||||||
|
except RuntimeError:
|
||||||
|
# asyncio.run() checks for a running loop before awaiting the coroutine;
|
||||||
|
# when it raises, the original coro was never started — close it to
|
||||||
|
# prevent "coroutine was never awaited" RuntimeWarning, then retry in a
|
||||||
|
# fresh thread that has no running loop.
|
||||||
|
coro.close()
|
||||||
|
import concurrent.futures
|
||||||
|
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||||
|
future = pool.submit(asyncio.run, _send_to_platform(platform, pconfig, chat_id, cleaned_delivery_content, thread_id=thread_id, media_files=media_files))
|
||||||
|
result = future.result(timeout=30)
|
||||||
|
except Exception as e:
|
||||||
|
msg = f"delivery to {platform_name}:{chat_id} failed: {e}"
|
||||||
|
logger.error("Job '%s': %s", job["id"], msg)
|
||||||
|
return msg
|
||||||
|
|
||||||
|
if result and result.get("error"):
|
||||||
|
msg = f"delivery error: {result['error']}"
|
||||||
|
logger.error("Job '%s': %s", job["id"], msg)
|
||||||
|
return msg
|
||||||
|
|
||||||
|
logger.info("Job '%s': delivered to %s:%s", job["id"], platform_name, chat_id)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
_DEFAULT_SCRIPT_TIMEOUT = 120 # seconds
|
||||||
|
# Backward-compatible module override used by tests and emergency monkeypatches.
|
||||||
|
_SCRIPT_TIMEOUT = _DEFAULT_SCRIPT_TIMEOUT
|
||||||
|
|
||||||
|
|
||||||
|
def _get_script_timeout() -> int:
|
||||||
|
"""Resolve cron pre-run script timeout from module/env/config with a safe default."""
|
||||||
|
if _SCRIPT_TIMEOUT != _DEFAULT_SCRIPT_TIMEOUT:
|
||||||
|
try:
|
||||||
|
timeout = int(float(_SCRIPT_TIMEOUT))
|
||||||
|
if timeout > 0:
|
||||||
|
return timeout
|
||||||
|
except Exception:
|
||||||
|
logger.warning("Invalid patched _SCRIPT_TIMEOUT=%r; using env/config/default", _SCRIPT_TIMEOUT)
|
||||||
|
|
||||||
|
env_value = os.getenv("HERMES_CRON_SCRIPT_TIMEOUT", "").strip()
|
||||||
|
if env_value:
|
||||||
|
try:
|
||||||
|
timeout = int(float(env_value))
|
||||||
|
if timeout > 0:
|
||||||
|
return timeout
|
||||||
|
except Exception:
|
||||||
|
logger.warning("Invalid HERMES_CRON_SCRIPT_TIMEOUT=%r; using config/default", env_value)
|
||||||
|
|
||||||
|
try:
|
||||||
|
cfg = load_config() or {}
|
||||||
|
cron_cfg = cfg.get("cron", {}) if isinstance(cfg, dict) else {}
|
||||||
|
configured = cron_cfg.get("script_timeout_seconds")
|
||||||
|
if configured is not None:
|
||||||
|
timeout = int(float(configured))
|
||||||
|
if timeout > 0:
|
||||||
|
return timeout
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug("Failed to load cron script timeout from config: %s", exc)
|
||||||
|
|
||||||
|
return _DEFAULT_SCRIPT_TIMEOUT
|
||||||
|
|
||||||
|
|
||||||
|
def _run_job_script(script_path: str) -> tuple[bool, str]:
|
||||||
|
"""Execute a cron job's data-collection script and capture its output.
|
||||||
|
|
||||||
|
Scripts must reside within HERMES_HOME/scripts/. Both relative and
|
||||||
|
absolute paths are resolved and validated against this directory to
|
||||||
|
prevent arbitrary script execution via path traversal or absolute
|
||||||
|
path injection.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
script_path: Path to a Python script. Relative paths are resolved
|
||||||
|
against HERMES_HOME/scripts/. Absolute and ~-prefixed paths
|
||||||
|
are also validated to ensure they stay within the scripts dir.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(success, output) — on failure *output* contains the error message so the
|
||||||
|
LLM can report the problem to the user.
|
||||||
|
"""
|
||||||
|
from hermes_constants import get_hermes_home
|
||||||
|
|
||||||
|
scripts_dir = get_hermes_home() / "scripts"
|
||||||
|
scripts_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
scripts_dir_resolved = scripts_dir.resolve()
|
||||||
|
|
||||||
|
raw = Path(script_path).expanduser()
|
||||||
|
if raw.is_absolute():
|
||||||
|
path = raw.resolve()
|
||||||
|
else:
|
||||||
|
path = (scripts_dir / raw).resolve()
|
||||||
|
|
||||||
|
# Guard against path traversal, absolute path injection, and symlink
|
||||||
|
# escape — scripts MUST reside within HERMES_HOME/scripts/.
|
||||||
|
try:
|
||||||
|
path.relative_to(scripts_dir_resolved)
|
||||||
|
except ValueError:
|
||||||
|
return False, (
|
||||||
|
f"Blocked: script path resolves outside the scripts directory "
|
||||||
|
f"({scripts_dir_resolved}): {script_path!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if not path.exists():
|
||||||
|
return False, f"Script not found: {path}"
|
||||||
|
if not path.is_file():
|
||||||
|
return False, f"Script path is not a file: {path}"
|
||||||
|
|
||||||
|
script_timeout = _get_script_timeout()
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
[sys.executable, str(path)],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=script_timeout,
|
||||||
|
cwd=str(path.parent),
|
||||||
|
)
|
||||||
|
stdout = (result.stdout or "").strip()
|
||||||
|
stderr = (result.stderr or "").strip()
|
||||||
|
|
||||||
|
# Redact secrets from both stdout and stderr before any return path.
|
||||||
|
try:
|
||||||
|
from agent.redact import redact_sensitive_text
|
||||||
|
stdout = redact_sensitive_text(stdout)
|
||||||
|
stderr = redact_sensitive_text(stderr)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if result.returncode != 0:
|
||||||
|
parts = [f"Script exited with code {result.returncode}"]
|
||||||
|
if stderr:
|
||||||
|
parts.append(f"stderr:\n{stderr}")
|
||||||
|
if stdout:
|
||||||
|
parts.append(f"stdout:\n{stdout}")
|
||||||
|
return False, "\n".join(parts)
|
||||||
|
|
||||||
|
return True, stdout
|
||||||
|
|
||||||
|
except subprocess.TimeoutExpired:
|
||||||
|
return False, f"Script timed out after {script_timeout}s: {path}"
|
||||||
|
except Exception as exc:
|
||||||
|
return False, f"Script execution failed: {exc}"
|
||||||
|
|
||||||
|
|
||||||
|
def _build_job_prompt(job: dict) -> str:
|
||||||
|
"""Build the effective prompt for a cron job, optionally loading one or more skills first."""
|
||||||
|
prompt = job.get("prompt", "")
|
||||||
|
skills = job.get("skills")
|
||||||
|
|
||||||
|
# Run data-collection script if configured, inject output as context.
|
||||||
|
script_path = job.get("script")
|
||||||
|
if script_path:
|
||||||
|
success, script_output = _run_job_script(script_path)
|
||||||
|
if success:
|
||||||
|
if script_output:
|
||||||
|
prompt = (
|
||||||
|
"## Script Output\n"
|
||||||
|
"The following data was collected by a pre-run script. "
|
||||||
|
"Use it as context for your analysis.\n\n"
|
||||||
|
f"```\n{script_output}\n```\n\n"
|
||||||
|
f"{prompt}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
prompt = (
|
||||||
|
"[Script ran successfully but produced no output.]\n\n"
|
||||||
|
f"{prompt}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
prompt = (
|
||||||
|
"## Script Error\n"
|
||||||
|
"The data-collection script failed. Report this to the user.\n\n"
|
||||||
|
f"```\n{script_output}\n```\n\n"
|
||||||
|
f"{prompt}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Always prepend cron execution guidance so the agent knows how
|
||||||
|
# delivery works and can suppress delivery when appropriate.
|
||||||
|
cron_hint = (
|
||||||
|
"[SYSTEM: You are running as a scheduled cron job. "
|
||||||
|
"DELIVERY: Your final response will be automatically delivered "
|
||||||
|
"to the user — do NOT use send_message or try to deliver "
|
||||||
|
"the output yourself. Just produce your report/output as your "
|
||||||
|
"final response and the system handles the rest. "
|
||||||
|
"SILENT: If there is genuinely nothing new to report, respond "
|
||||||
|
"with exactly \"[SILENT]\" (nothing else) to suppress delivery. "
|
||||||
|
"Never combine [SILENT] with content — either report your "
|
||||||
|
"findings normally, or say [SILENT] and nothing more.]\n\n"
|
||||||
|
)
|
||||||
|
prompt = cron_hint + prompt
|
||||||
|
if skills is None:
|
||||||
|
legacy = job.get("skill")
|
||||||
|
skills = [legacy] if legacy else []
|
||||||
|
|
||||||
|
skill_names = [str(name).strip() for name in skills if str(name).strip()]
|
||||||
|
if not skill_names:
|
||||||
|
return prompt
|
||||||
|
|
||||||
|
from tools.skills_tool import skill_view
|
||||||
|
|
||||||
|
parts = []
|
||||||
|
skipped: list[str] = []
|
||||||
|
for skill_name in skill_names:
|
||||||
|
loaded = json.loads(skill_view(skill_name))
|
||||||
|
if not loaded.get("success"):
|
||||||
|
error = loaded.get("error") or f"Failed to load skill '{skill_name}'"
|
||||||
|
logger.warning("Cron job '%s': skill not found, skipping — %s", job.get("name", job.get("id")), error)
|
||||||
|
skipped.append(skill_name)
|
||||||
|
continue
|
||||||
|
|
||||||
|
content = str(loaded.get("content") or "").strip()
|
||||||
|
if parts:
|
||||||
|
parts.append("")
|
||||||
|
parts.extend(
|
||||||
|
[
|
||||||
|
f'[SYSTEM: The user has invoked the "{skill_name}" skill, indicating they want you to follow its instructions. The full skill content is loaded below.]',
|
||||||
|
"",
|
||||||
|
content,
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
if skipped:
|
||||||
|
notice = (
|
||||||
|
f"[SYSTEM: The following skill(s) were listed for this job but could not be found "
|
||||||
|
f"and were skipped: {', '.join(skipped)}. "
|
||||||
|
f"Start your response with a brief notice so the user is aware, e.g.: "
|
||||||
|
f"'⚠️ Skill(s) not found and skipped: {', '.join(skipped)}']"
|
||||||
|
)
|
||||||
|
parts.insert(0, notice)
|
||||||
|
|
||||||
|
if prompt:
|
||||||
|
parts.extend(["", f"The user has provided the following instruction alongside the skill invocation: {prompt}"])
|
||||||
|
return "\n".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
def run_job(job: dict) -> tuple[bool, str, str, Optional[str]]:
|
||||||
|
"""
|
||||||
|
Execute a single cron job.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (success, full_output_doc, final_response, error_message)
|
||||||
|
"""
|
||||||
|
from run_agent import AIAgent
|
||||||
|
|
||||||
|
# Initialize SQLite session store so cron job messages are persisted
|
||||||
|
# and discoverable via session_search (same pattern as gateway/run.py).
|
||||||
|
_session_db = None
|
||||||
|
try:
|
||||||
|
from hermes_state import SessionDB
|
||||||
|
_session_db = SessionDB()
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Job '%s': SQLite session store not available: %s", job.get("id", "?"), e)
|
||||||
|
|
||||||
|
job_id = job["id"]
|
||||||
|
job_name = job["name"]
|
||||||
|
prompt = _build_job_prompt(job)
|
||||||
|
origin = _resolve_origin(job)
|
||||||
|
_cron_session_id = f"cron_{job_id}_{_hermes_now().strftime('%Y%m%d_%H%M%S')}"
|
||||||
|
|
||||||
|
logger.info("Running job '%s' (ID: %s)", job_name, job_id)
|
||||||
|
logger.info("Prompt: %s", prompt[:100])
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Inject origin context so the agent's send_message tool knows the chat.
|
||||||
|
# Must be INSIDE the try block so the finally cleanup always runs.
|
||||||
|
if origin:
|
||||||
|
os.environ["HERMES_SESSION_PLATFORM"] = origin["platform"]
|
||||||
|
os.environ["HERMES_SESSION_CHAT_ID"] = str(origin["chat_id"])
|
||||||
|
if origin.get("chat_name"):
|
||||||
|
os.environ["HERMES_SESSION_CHAT_NAME"] = origin["chat_name"]
|
||||||
|
# Re-read .env and config.yaml fresh every run so provider/key
|
||||||
|
# changes take effect without a gateway restart.
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
try:
|
||||||
|
load_dotenv(str(_hermes_home / ".env"), override=True, encoding="utf-8")
|
||||||
|
except UnicodeDecodeError:
|
||||||
|
load_dotenv(str(_hermes_home / ".env"), override=True, encoding="latin-1")
|
||||||
|
|
||||||
|
delivery_target = _resolve_delivery_target(job)
|
||||||
|
if delivery_target:
|
||||||
|
os.environ["HERMES_CRON_AUTO_DELIVER_PLATFORM"] = delivery_target["platform"]
|
||||||
|
os.environ["HERMES_CRON_AUTO_DELIVER_CHAT_ID"] = str(delivery_target["chat_id"])
|
||||||
|
if delivery_target.get("thread_id") is not None:
|
||||||
|
os.environ["HERMES_CRON_AUTO_DELIVER_THREAD_ID"] = str(delivery_target["thread_id"])
|
||||||
|
|
||||||
|
model = job.get("model") or os.getenv("HERMES_MODEL") or ""
|
||||||
|
|
||||||
|
# Load config.yaml for model, reasoning, prefill, toolsets, provider routing
|
||||||
|
_cfg = {}
|
||||||
|
try:
|
||||||
|
import yaml
|
||||||
|
_cfg_path = str(_hermes_home / "config.yaml")
|
||||||
|
if os.path.exists(_cfg_path):
|
||||||
|
with open(_cfg_path) as _f:
|
||||||
|
_cfg = yaml.safe_load(_f) or {}
|
||||||
|
_model_cfg = _cfg.get("model", {})
|
||||||
|
if not job.get("model"):
|
||||||
|
if isinstance(_model_cfg, str):
|
||||||
|
model = _model_cfg
|
||||||
|
elif isinstance(_model_cfg, dict):
|
||||||
|
model = _model_cfg.get("default", model)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Job '%s': failed to load config.yaml, using defaults: %s", job_id, e)
|
||||||
|
|
||||||
|
# Apply IPv4 preference if configured.
|
||||||
|
try:
|
||||||
|
from hermes_constants import apply_ipv4_preference
|
||||||
|
_net_cfg = _cfg.get("network", {})
|
||||||
|
if isinstance(_net_cfg, dict) and _net_cfg.get("force_ipv4"):
|
||||||
|
apply_ipv4_preference(force=True)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Reasoning config from config.yaml
|
||||||
|
from hermes_constants import parse_reasoning_effort
|
||||||
|
effort = str(_cfg.get("agent", {}).get("reasoning_effort", "")).strip()
|
||||||
|
reasoning_config = parse_reasoning_effort(effort)
|
||||||
|
|
||||||
|
# Prefill messages from env or config.yaml
|
||||||
|
prefill_messages = None
|
||||||
|
prefill_file = os.getenv("HERMES_PREFILL_MESSAGES_FILE", "") or _cfg.get("prefill_messages_file", "")
|
||||||
|
if prefill_file:
|
||||||
|
import json as _json
|
||||||
|
pfpath = Path(prefill_file).expanduser()
|
||||||
|
if not pfpath.is_absolute():
|
||||||
|
pfpath = _hermes_home / pfpath
|
||||||
|
if pfpath.exists():
|
||||||
|
try:
|
||||||
|
with open(pfpath, "r", encoding="utf-8") as _pf:
|
||||||
|
prefill_messages = _json.load(_pf)
|
||||||
|
if not isinstance(prefill_messages, list):
|
||||||
|
prefill_messages = None
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Job '%s': failed to parse prefill messages file '%s': %s", job_id, pfpath, e)
|
||||||
|
prefill_messages = None
|
||||||
|
|
||||||
|
# Max iterations
|
||||||
|
max_iterations = _cfg.get("agent", {}).get("max_turns") or _cfg.get("max_turns") or 90
|
||||||
|
|
||||||
|
# Provider routing
|
||||||
|
pr = _cfg.get("provider_routing", {})
|
||||||
|
smart_routing = _cfg.get("smart_model_routing", {}) or {}
|
||||||
|
|
||||||
|
from hermes_cli.runtime_provider import (
|
||||||
|
resolve_runtime_provider,
|
||||||
|
format_runtime_provider_error,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
runtime_kwargs = {
|
||||||
|
"requested": job.get("provider") or os.getenv("HERMES_INFERENCE_PROVIDER"),
|
||||||
|
}
|
||||||
|
if job.get("base_url"):
|
||||||
|
runtime_kwargs["explicit_base_url"] = job.get("base_url")
|
||||||
|
runtime = resolve_runtime_provider(**runtime_kwargs)
|
||||||
|
except Exception as exc:
|
||||||
|
message = format_runtime_provider_error(exc)
|
||||||
|
raise RuntimeError(message) from exc
|
||||||
|
|
||||||
|
from agent.smart_model_routing import resolve_turn_route
|
||||||
|
turn_route = resolve_turn_route(
|
||||||
|
prompt,
|
||||||
|
smart_routing,
|
||||||
|
{
|
||||||
|
"model": model,
|
||||||
|
"api_key": runtime.get("api_key"),
|
||||||
|
"base_url": runtime.get("base_url"),
|
||||||
|
"provider": runtime.get("provider"),
|
||||||
|
"api_mode": runtime.get("api_mode"),
|
||||||
|
"command": runtime.get("command"),
|
||||||
|
"args": list(runtime.get("args") or []),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
fallback_model = _cfg.get("fallback_providers") or _cfg.get("fallback_model") or None
|
||||||
|
credential_pool = None
|
||||||
|
runtime_provider = str(turn_route["runtime"].get("provider") or "").strip().lower()
|
||||||
|
if runtime_provider:
|
||||||
|
try:
|
||||||
|
from agent.credential_pool import load_pool
|
||||||
|
pool = load_pool(runtime_provider)
|
||||||
|
if pool.has_credentials():
|
||||||
|
credential_pool = pool
|
||||||
|
logger.info(
|
||||||
|
"Job '%s': loaded credential pool for provider %s with %d entries",
|
||||||
|
job_id,
|
||||||
|
runtime_provider,
|
||||||
|
len(pool.entries()),
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Job '%s': failed to load credential pool for %s: %s", job_id, runtime_provider, e)
|
||||||
|
|
||||||
|
agent = AIAgent(
|
||||||
|
model=turn_route["model"],
|
||||||
|
api_key=turn_route["runtime"].get("api_key"),
|
||||||
|
base_url=turn_route["runtime"].get("base_url"),
|
||||||
|
provider=turn_route["runtime"].get("provider"),
|
||||||
|
api_mode=turn_route["runtime"].get("api_mode"),
|
||||||
|
acp_command=turn_route["runtime"].get("command"),
|
||||||
|
acp_args=turn_route["runtime"].get("args"),
|
||||||
|
max_iterations=max_iterations,
|
||||||
|
reasoning_config=reasoning_config,
|
||||||
|
prefill_messages=prefill_messages,
|
||||||
|
fallback_model=fallback_model,
|
||||||
|
credential_pool=credential_pool,
|
||||||
|
providers_allowed=pr.get("only"),
|
||||||
|
providers_ignored=pr.get("ignore"),
|
||||||
|
providers_order=pr.get("order"),
|
||||||
|
provider_sort=pr.get("sort"),
|
||||||
|
disabled_toolsets=["cronjob", "messaging", "clarify"],
|
||||||
|
quiet_mode=True,
|
||||||
|
skip_context_files=True, # Don't inject SOUL.md/AGENTS.md from scheduler cwd
|
||||||
|
skip_memory=True, # Cron system prompts would corrupt user representations
|
||||||
|
platform="cron",
|
||||||
|
session_id=_cron_session_id,
|
||||||
|
session_db=_session_db,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Run the agent with an *inactivity*-based timeout: the job can run
|
||||||
|
# for hours if it's actively calling tools / receiving stream tokens,
|
||||||
|
# but a hung API call or stuck tool with no activity for the configured
|
||||||
|
# duration is caught and killed. Default 600s (10 min inactivity);
|
||||||
|
# override via HERMES_CRON_TIMEOUT env var. 0 = unlimited.
|
||||||
|
#
|
||||||
|
# Uses the agent's built-in activity tracker (updated by
|
||||||
|
# _touch_activity() on every tool call, API call, and stream delta).
|
||||||
|
_cron_timeout = float(os.getenv("HERMES_CRON_TIMEOUT", 600))
|
||||||
|
_cron_inactivity_limit = _cron_timeout if _cron_timeout > 0 else None
|
||||||
|
_POLL_INTERVAL = 5.0
|
||||||
|
_cron_pool = concurrent.futures.ThreadPoolExecutor(max_workers=1)
|
||||||
|
_cron_future = _cron_pool.submit(agent.run_conversation, prompt)
|
||||||
|
_inactivity_timeout = False
|
||||||
|
try:
|
||||||
|
if _cron_inactivity_limit is None:
|
||||||
|
# Unlimited — just wait for the result.
|
||||||
|
result = _cron_future.result()
|
||||||
|
else:
|
||||||
|
result = None
|
||||||
|
while True:
|
||||||
|
done, _ = concurrent.futures.wait(
|
||||||
|
{_cron_future}, timeout=_POLL_INTERVAL,
|
||||||
|
)
|
||||||
|
if done:
|
||||||
|
result = _cron_future.result()
|
||||||
|
break
|
||||||
|
# Agent still running — check inactivity.
|
||||||
|
_idle_secs = 0.0
|
||||||
|
if hasattr(agent, "get_activity_summary"):
|
||||||
|
try:
|
||||||
|
_act = agent.get_activity_summary()
|
||||||
|
_idle_secs = _act.get("seconds_since_activity", 0.0)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
if _idle_secs >= _cron_inactivity_limit:
|
||||||
|
_inactivity_timeout = True
|
||||||
|
break
|
||||||
|
except Exception:
|
||||||
|
_cron_pool.shutdown(wait=False, cancel_futures=True)
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
_cron_pool.shutdown(wait=False, cancel_futures=True)
|
||||||
|
|
||||||
|
if _inactivity_timeout:
|
||||||
|
# Build diagnostic summary from the agent's activity tracker.
|
||||||
|
_activity = {}
|
||||||
|
if hasattr(agent, "get_activity_summary"):
|
||||||
|
try:
|
||||||
|
_activity = agent.get_activity_summary()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
_last_desc = _activity.get("last_activity_desc", "unknown")
|
||||||
|
_secs_ago = _activity.get("seconds_since_activity", 0)
|
||||||
|
_cur_tool = _activity.get("current_tool")
|
||||||
|
_iter_n = _activity.get("api_call_count", 0)
|
||||||
|
_iter_max = _activity.get("max_iterations", 0)
|
||||||
|
|
||||||
|
logger.error(
|
||||||
|
"Job '%s' idle for %.0fs (inactivity limit %.0fs) "
|
||||||
|
"| last_activity=%s | iteration=%s/%s | tool=%s",
|
||||||
|
job_name, _secs_ago, _cron_inactivity_limit,
|
||||||
|
_last_desc, _iter_n, _iter_max,
|
||||||
|
_cur_tool or "none",
|
||||||
|
)
|
||||||
|
if hasattr(agent, "interrupt"):
|
||||||
|
agent.interrupt("Cron job timed out (inactivity)")
|
||||||
|
raise TimeoutError(
|
||||||
|
f"Cron job '{job_name}' idle for "
|
||||||
|
f"{int(_secs_ago)}s (limit {int(_cron_inactivity_limit)}s) "
|
||||||
|
f"— last activity: {_last_desc}"
|
||||||
|
)
|
||||||
|
|
||||||
|
final_response = result.get("final_response", "") or ""
|
||||||
|
# Use a separate variable for log display; keep final_response clean
|
||||||
|
# for delivery logic (empty response = no delivery).
|
||||||
|
logged_response = final_response if final_response else "(No response generated)"
|
||||||
|
|
||||||
|
output = f"""# Cron Job: {job_name}
|
||||||
|
|
||||||
|
**Job ID:** {job_id}
|
||||||
|
**Run Time:** {_hermes_now().strftime('%Y-%m-%d %H:%M:%S')}
|
||||||
|
**Schedule:** {job.get('schedule_display', 'N/A')}
|
||||||
|
|
||||||
|
## Prompt
|
||||||
|
|
||||||
|
{prompt}
|
||||||
|
|
||||||
|
## Response
|
||||||
|
|
||||||
|
{logged_response}
|
||||||
|
"""
|
||||||
|
|
||||||
|
logger.info("Job '%s' completed successfully", job_name)
|
||||||
|
return True, output, final_response, None
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
error_msg = f"{type(e).__name__}: {str(e)}"
|
||||||
|
logger.exception("Job '%s' failed: %s", job_name, error_msg)
|
||||||
|
|
||||||
|
output = f"""# Cron Job: {job_name} (FAILED)
|
||||||
|
|
||||||
|
**Job ID:** {job_id}
|
||||||
|
**Run Time:** {_hermes_now().strftime('%Y-%m-%d %H:%M:%S')}
|
||||||
|
**Schedule:** {job.get('schedule_display', 'N/A')}
|
||||||
|
|
||||||
|
## Prompt
|
||||||
|
|
||||||
|
{prompt}
|
||||||
|
|
||||||
|
## Error
|
||||||
|
|
||||||
|
```
|
||||||
|
{error_msg}
|
||||||
|
```
|
||||||
|
"""
|
||||||
|
return False, output, "", error_msg
|
||||||
|
|
||||||
|
finally:
|
||||||
|
# Clean up injected env vars so they don't leak to other jobs
|
||||||
|
for key in (
|
||||||
|
"HERMES_SESSION_PLATFORM",
|
||||||
|
"HERMES_SESSION_CHAT_ID",
|
||||||
|
"HERMES_SESSION_CHAT_NAME",
|
||||||
|
"HERMES_CRON_AUTO_DELIVER_PLATFORM",
|
||||||
|
"HERMES_CRON_AUTO_DELIVER_CHAT_ID",
|
||||||
|
"HERMES_CRON_AUTO_DELIVER_THREAD_ID",
|
||||||
|
):
|
||||||
|
os.environ.pop(key, None)
|
||||||
|
if _session_db:
|
||||||
|
try:
|
||||||
|
_session_db.end_session(_cron_session_id, "cron_complete")
|
||||||
|
except (Exception, KeyboardInterrupt) as e:
|
||||||
|
logger.debug("Job '%s': failed to end session: %s", job_id, e)
|
||||||
|
try:
|
||||||
|
_session_db.close()
|
||||||
|
except (Exception, KeyboardInterrupt) as e:
|
||||||
|
logger.debug("Job '%s': failed to close SQLite session store: %s", job_id, e)
|
||||||
|
|
||||||
|
|
||||||
|
def tick(verbose: bool = True, adapters=None, loop=None) -> int:
|
||||||
|
"""
|
||||||
|
Check and run all due jobs.
|
||||||
|
|
||||||
|
Uses a file lock so only one tick runs at a time, even if the gateway's
|
||||||
|
in-process ticker and a standalone daemon or manual tick overlap.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
verbose: Whether to print status messages
|
||||||
|
adapters: Optional dict mapping Platform → live adapter (from gateway)
|
||||||
|
loop: Optional asyncio event loop (from gateway) for live adapter sends
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Number of jobs executed (0 if another tick is already running)
|
||||||
|
"""
|
||||||
|
_LOCK_DIR.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
# Cross-platform file locking: fcntl on Unix, msvcrt on Windows
|
||||||
|
lock_fd = None
|
||||||
|
try:
|
||||||
|
lock_fd = open(_LOCK_FILE, "w")
|
||||||
|
if fcntl:
|
||||||
|
fcntl.flock(lock_fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||||
|
elif msvcrt:
|
||||||
|
msvcrt.locking(lock_fd.fileno(), msvcrt.LK_NBLCK, 1)
|
||||||
|
except (OSError, IOError):
|
||||||
|
logger.debug("Tick skipped — another instance holds the lock")
|
||||||
|
if lock_fd is not None:
|
||||||
|
lock_fd.close()
|
||||||
|
return 0
|
||||||
|
|
||||||
|
try:
|
||||||
|
due_jobs = get_due_jobs()
|
||||||
|
|
||||||
|
if verbose and not due_jobs:
|
||||||
|
logger.info("%s - No jobs due", _hermes_now().strftime('%H:%M:%S'))
|
||||||
|
return 0
|
||||||
|
|
||||||
|
if verbose:
|
||||||
|
logger.info("%s - %s job(s) due", _hermes_now().strftime('%H:%M:%S'), len(due_jobs))
|
||||||
|
|
||||||
|
executed = 0
|
||||||
|
for job in due_jobs:
|
||||||
|
try:
|
||||||
|
# For recurring jobs (cron/interval), advance next_run_at to the
|
||||||
|
# next future occurrence BEFORE execution. This way, if the
|
||||||
|
# process crashes mid-run, the job won't re-fire on restart.
|
||||||
|
# One-shot jobs are left alone so they can retry on restart.
|
||||||
|
advance_next_run(job["id"])
|
||||||
|
|
||||||
|
success, output, final_response, error = run_job(job)
|
||||||
|
|
||||||
|
output_file = save_job_output(job["id"], output)
|
||||||
|
if verbose:
|
||||||
|
logger.info("Output saved to: %s", output_file)
|
||||||
|
|
||||||
|
# Deliver the final response to the origin/target chat.
|
||||||
|
# If the agent responded with [SILENT], skip delivery (but
|
||||||
|
# output is already saved above). Failed jobs always deliver.
|
||||||
|
deliver_content = final_response if success else f"⚠️ Cron job '{job.get('name', job['id'])}' failed:\n{error}"
|
||||||
|
should_deliver = bool(deliver_content)
|
||||||
|
if should_deliver and success and SILENT_MARKER in deliver_content.strip().upper():
|
||||||
|
logger.info("Job '%s': agent returned %s — skipping delivery", job["id"], SILENT_MARKER)
|
||||||
|
should_deliver = False
|
||||||
|
|
||||||
|
delivery_error = None
|
||||||
|
if should_deliver:
|
||||||
|
try:
|
||||||
|
delivery_error = _deliver_result(job, deliver_content, adapters=adapters, loop=loop)
|
||||||
|
except Exception as de:
|
||||||
|
delivery_error = str(de)
|
||||||
|
logger.error("Delivery failed for job %s: %s", job["id"], de)
|
||||||
|
|
||||||
|
mark_job_run(job["id"], success, error, delivery_error=delivery_error)
|
||||||
|
executed += 1
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Error processing job %s: %s", job['id'], e)
|
||||||
|
mark_job_run(job["id"], False, str(e))
|
||||||
|
|
||||||
|
return executed
|
||||||
|
finally:
|
||||||
|
if fcntl:
|
||||||
|
fcntl.flock(lock_fd, fcntl.LOCK_UN)
|
||||||
|
elif msvcrt:
|
||||||
|
try:
|
||||||
|
msvcrt.locking(lock_fd.fileno(), msvcrt.LK_UNLCK, 1)
|
||||||
|
except (OSError, IOError):
|
||||||
|
pass
|
||||||
|
lock_fd.close()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
tick(verbose=True)
|
||||||
@@ -0,0 +1,35 @@
|
|||||||
|
"""
|
||||||
|
Hermes Gateway - Multi-platform messaging integration.
|
||||||
|
|
||||||
|
This module provides a unified gateway for connecting the Hermes agent
|
||||||
|
to various messaging platforms (Telegram, Discord, WhatsApp) with:
|
||||||
|
- Session management (persistent conversations with reset policies)
|
||||||
|
- Dynamic context injection (agent knows where messages come from)
|
||||||
|
- Delivery routing (cron job outputs to appropriate channels)
|
||||||
|
- Platform-specific toolsets (different capabilities per platform)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from .config import GatewayConfig, PlatformConfig, HomeChannel, load_gateway_config
|
||||||
|
from .session import (
|
||||||
|
SessionContext,
|
||||||
|
SessionStore,
|
||||||
|
SessionResetPolicy,
|
||||||
|
build_session_context_prompt,
|
||||||
|
)
|
||||||
|
from .delivery import DeliveryRouter, DeliveryTarget
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Config
|
||||||
|
"GatewayConfig",
|
||||||
|
"PlatformConfig",
|
||||||
|
"HomeChannel",
|
||||||
|
"load_gateway_config",
|
||||||
|
# Session
|
||||||
|
"SessionContext",
|
||||||
|
"SessionStore",
|
||||||
|
"SessionResetPolicy",
|
||||||
|
"build_session_context_prompt",
|
||||||
|
# Delivery
|
||||||
|
"DeliveryRouter",
|
||||||
|
"DeliveryTarget",
|
||||||
|
]
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Built-in gateway hooks that are always registered."""
|
||||||
@@ -0,0 +1,85 @@
|
|||||||
|
"""Built-in boot-md hook — run ~/.hermes/BOOT.md on gateway startup.
|
||||||
|
|
||||||
|
This hook is always registered. It silently skips if no BOOT.md exists.
|
||||||
|
To activate, create ``~/.hermes/BOOT.md`` with instructions for the
|
||||||
|
agent to execute on every gateway restart.
|
||||||
|
|
||||||
|
Example BOOT.md::
|
||||||
|
|
||||||
|
# Startup Checklist
|
||||||
|
|
||||||
|
1. Check if any cron jobs failed overnight
|
||||||
|
2. Send a status update to Discord #general
|
||||||
|
3. If there are errors in /opt/app/deploy.log, summarize them
|
||||||
|
|
||||||
|
The agent runs in a background thread so it doesn't block gateway
|
||||||
|
startup. If nothing needs attention, it replies with [SILENT] to
|
||||||
|
suppress delivery.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import threading
|
||||||
|
|
||||||
|
logger = logging.getLogger("hooks.boot-md")
|
||||||
|
|
||||||
|
from hermes_constants import get_hermes_home
|
||||||
|
HERMES_HOME = get_hermes_home()
|
||||||
|
BOOT_FILE = HERMES_HOME / "BOOT.md"
|
||||||
|
|
||||||
|
|
||||||
|
def _build_boot_prompt(content: str) -> str:
|
||||||
|
"""Wrap BOOT.md content in a system-level instruction."""
|
||||||
|
return (
|
||||||
|
"You are running a startup boot checklist. Follow the BOOT.md "
|
||||||
|
"instructions below exactly.\n\n"
|
||||||
|
"---\n"
|
||||||
|
f"{content}\n"
|
||||||
|
"---\n\n"
|
||||||
|
"Execute each instruction. If you need to send a message to a "
|
||||||
|
"platform, use the send_message tool.\n"
|
||||||
|
"If nothing needs attention and there is nothing to report, "
|
||||||
|
"reply with ONLY: [SILENT]"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _run_boot_agent(content: str) -> None:
|
||||||
|
"""Spawn a one-shot agent session to execute the boot instructions."""
|
||||||
|
try:
|
||||||
|
from run_agent import AIAgent
|
||||||
|
|
||||||
|
prompt = _build_boot_prompt(content)
|
||||||
|
agent = AIAgent(
|
||||||
|
quiet_mode=True,
|
||||||
|
skip_context_files=True,
|
||||||
|
skip_memory=True,
|
||||||
|
max_iterations=20,
|
||||||
|
)
|
||||||
|
result = agent.run_conversation(prompt)
|
||||||
|
response = result.get("final_response", "")
|
||||||
|
if response and "[SILENT]" not in response:
|
||||||
|
logger.info("boot-md completed: %s", response[:200])
|
||||||
|
else:
|
||||||
|
logger.info("boot-md completed (nothing to report)")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("boot-md agent failed: %s", e)
|
||||||
|
|
||||||
|
|
||||||
|
async def handle(event_type: str, context: dict) -> None:
|
||||||
|
"""Gateway startup handler — run BOOT.md if it exists."""
|
||||||
|
if not BOOT_FILE.exists():
|
||||||
|
return
|
||||||
|
|
||||||
|
content = BOOT_FILE.read_text(encoding="utf-8").strip()
|
||||||
|
if not content:
|
||||||
|
return
|
||||||
|
|
||||||
|
logger.info("Running BOOT.md (%d chars)", len(content))
|
||||||
|
|
||||||
|
# Run in a background thread so we don't block gateway startup.
|
||||||
|
thread = threading.Thread(
|
||||||
|
target=_run_boot_agent,
|
||||||
|
args=(content,),
|
||||||
|
name="boot-md",
|
||||||
|
daemon=True,
|
||||||
|
)
|
||||||
|
thread.start()
|
||||||
@@ -0,0 +1,276 @@
|
|||||||
|
"""
|
||||||
|
Channel directory -- cached map of reachable channels/contacts per platform.
|
||||||
|
|
||||||
|
Built on gateway startup, refreshed periodically (every 5 min), and saved to
|
||||||
|
~/.hermes/channel_directory.json. The send_message tool reads this file for
|
||||||
|
action="list" and for resolving human-friendly channel names to numeric IDs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
from hermes_cli.config import get_hermes_home
|
||||||
|
from utils import atomic_json_write
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
DIRECTORY_PATH = get_hermes_home() / "channel_directory.json"
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_channel_query(value: str) -> str:
|
||||||
|
return value.lstrip("#").strip().lower()
|
||||||
|
|
||||||
|
|
||||||
|
def _channel_target_name(platform_name: str, channel: Dict[str, Any]) -> str:
|
||||||
|
"""Return the human-facing target label shown to users for a channel entry."""
|
||||||
|
name = channel["name"]
|
||||||
|
if platform_name == "discord" and channel.get("guild"):
|
||||||
|
return f"#{name}"
|
||||||
|
if platform_name != "discord" and channel.get("type"):
|
||||||
|
return f"{name} ({channel['type']})"
|
||||||
|
return name
|
||||||
|
|
||||||
|
|
||||||
|
def _session_entry_id(origin: Dict[str, Any]) -> Optional[str]:
|
||||||
|
chat_id = origin.get("chat_id")
|
||||||
|
if not chat_id:
|
||||||
|
return None
|
||||||
|
thread_id = origin.get("thread_id")
|
||||||
|
if thread_id:
|
||||||
|
return f"{chat_id}:{thread_id}"
|
||||||
|
return str(chat_id)
|
||||||
|
|
||||||
|
|
||||||
|
def _session_entry_name(origin: Dict[str, Any]) -> str:
|
||||||
|
base_name = origin.get("chat_name") or origin.get("user_name") or str(origin.get("chat_id"))
|
||||||
|
thread_id = origin.get("thread_id")
|
||||||
|
if not thread_id:
|
||||||
|
return base_name
|
||||||
|
|
||||||
|
topic_label = origin.get("chat_topic") or f"topic {thread_id}"
|
||||||
|
return f"{base_name} / {topic_label}"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Build / refresh
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def build_channel_directory(adapters: Dict[Any, Any]) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
Build a channel directory from connected platform adapters and session data.
|
||||||
|
|
||||||
|
Returns the directory dict and writes it to DIRECTORY_PATH.
|
||||||
|
"""
|
||||||
|
from gateway.config import Platform
|
||||||
|
|
||||||
|
platforms: Dict[str, List[Dict[str, str]]] = {}
|
||||||
|
|
||||||
|
for platform, adapter in adapters.items():
|
||||||
|
try:
|
||||||
|
if platform == Platform.DISCORD:
|
||||||
|
platforms["discord"] = _build_discord(adapter)
|
||||||
|
elif platform == Platform.SLACK:
|
||||||
|
platforms["slack"] = _build_slack(adapter)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Channel directory: failed to build %s: %s", platform.value, e)
|
||||||
|
|
||||||
|
# Platforms that don't support direct channel enumeration get session-based
|
||||||
|
# discovery automatically. Skip infrastructure entries that aren't messaging
|
||||||
|
# platforms — everything else falls through to _build_from_sessions().
|
||||||
|
_SKIP_SESSION_DISCOVERY = frozenset({"local", "api_server", "webhook"})
|
||||||
|
for plat in Platform:
|
||||||
|
plat_name = plat.value
|
||||||
|
if plat_name in _SKIP_SESSION_DISCOVERY or plat_name in platforms:
|
||||||
|
continue
|
||||||
|
platforms[plat_name] = _build_from_sessions(plat_name)
|
||||||
|
|
||||||
|
directory = {
|
||||||
|
"updated_at": datetime.now().isoformat(),
|
||||||
|
"platforms": platforms,
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
atomic_json_write(DIRECTORY_PATH, directory)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Channel directory: failed to write: %s", e)
|
||||||
|
|
||||||
|
return directory
|
||||||
|
|
||||||
|
|
||||||
|
def _build_discord(adapter) -> List[Dict[str, str]]:
|
||||||
|
"""Enumerate all text channels the Discord bot can see."""
|
||||||
|
channels = []
|
||||||
|
client = getattr(adapter, "_client", None)
|
||||||
|
if not client:
|
||||||
|
return channels
|
||||||
|
|
||||||
|
try:
|
||||||
|
import discord as _discord # noqa: F401 — SDK presence check
|
||||||
|
except ImportError:
|
||||||
|
return channels
|
||||||
|
|
||||||
|
for guild in client.guilds:
|
||||||
|
for ch in guild.text_channels:
|
||||||
|
channels.append({
|
||||||
|
"id": str(ch.id),
|
||||||
|
"name": ch.name,
|
||||||
|
"guild": guild.name,
|
||||||
|
"type": "channel",
|
||||||
|
})
|
||||||
|
# Also include DM-capable users we've interacted with is not
|
||||||
|
# feasible via guild enumeration; those come from sessions.
|
||||||
|
|
||||||
|
# Merge any DMs from session history
|
||||||
|
channels.extend(_build_from_sessions("discord"))
|
||||||
|
return channels
|
||||||
|
|
||||||
|
|
||||||
|
def _build_slack(adapter) -> List[Dict[str, str]]:
|
||||||
|
"""List Slack channels the bot has joined."""
|
||||||
|
# Slack adapter may expose a web client
|
||||||
|
client = getattr(adapter, "_app", None) or getattr(adapter, "_client", None)
|
||||||
|
if not client:
|
||||||
|
return _build_from_sessions("slack")
|
||||||
|
|
||||||
|
try:
|
||||||
|
from tools.send_message_tool import _send_slack # noqa: F401
|
||||||
|
# Use the Slack Web API directly if available
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Fallback to session data
|
||||||
|
return _build_from_sessions("slack")
|
||||||
|
|
||||||
|
|
||||||
|
def _build_from_sessions(platform_name: str) -> List[Dict[str, str]]:
|
||||||
|
"""Pull known channels/contacts from sessions.json origin data."""
|
||||||
|
sessions_path = get_hermes_home() / "sessions" / "sessions.json"
|
||||||
|
if not sessions_path.exists():
|
||||||
|
return []
|
||||||
|
|
||||||
|
entries = []
|
||||||
|
try:
|
||||||
|
with open(sessions_path, encoding="utf-8") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
|
||||||
|
seen_ids = set()
|
||||||
|
for _key, session in data.items():
|
||||||
|
origin = session.get("origin") or {}
|
||||||
|
if origin.get("platform") != platform_name:
|
||||||
|
continue
|
||||||
|
entry_id = _session_entry_id(origin)
|
||||||
|
if not entry_id or entry_id in seen_ids:
|
||||||
|
continue
|
||||||
|
seen_ids.add(entry_id)
|
||||||
|
entries.append({
|
||||||
|
"id": entry_id,
|
||||||
|
"name": _session_entry_name(origin),
|
||||||
|
"type": session.get("chat_type", "dm"),
|
||||||
|
"thread_id": origin.get("thread_id"),
|
||||||
|
})
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Channel directory: failed to read sessions for %s: %s", platform_name, e)
|
||||||
|
|
||||||
|
return entries
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Read / resolve
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def load_directory() -> Dict[str, Any]:
|
||||||
|
"""Load the cached channel directory from disk."""
|
||||||
|
if not DIRECTORY_PATH.exists():
|
||||||
|
return {"updated_at": None, "platforms": {}}
|
||||||
|
try:
|
||||||
|
with open(DIRECTORY_PATH, encoding="utf-8") as f:
|
||||||
|
return json.load(f)
|
||||||
|
except Exception:
|
||||||
|
return {"updated_at": None, "platforms": {}}
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_channel_name(platform_name: str, name: str) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
Resolve a human-friendly channel name to a numeric ID.
|
||||||
|
|
||||||
|
Matching strategy (case-insensitive, first match wins):
|
||||||
|
- Discord: "bot-home", "#bot-home", "GuildName/bot-home"
|
||||||
|
- Telegram: display name or group name
|
||||||
|
- Slack: "engineering", "#engineering"
|
||||||
|
"""
|
||||||
|
directory = load_directory()
|
||||||
|
channels = directory.get("platforms", {}).get(platform_name, [])
|
||||||
|
if not channels:
|
||||||
|
return None
|
||||||
|
|
||||||
|
query = _normalize_channel_query(name)
|
||||||
|
|
||||||
|
# 1. Exact name match, including the display labels shown by send_message(action="list")
|
||||||
|
for ch in channels:
|
||||||
|
if _normalize_channel_query(ch["name"]) == query:
|
||||||
|
return ch["id"]
|
||||||
|
if _normalize_channel_query(_channel_target_name(platform_name, ch)) == query:
|
||||||
|
return ch["id"]
|
||||||
|
|
||||||
|
# 2. Guild-qualified match for Discord ("GuildName/channel")
|
||||||
|
if "/" in query:
|
||||||
|
guild_part, ch_part = query.rsplit("/", 1)
|
||||||
|
for ch in channels:
|
||||||
|
guild = ch.get("guild", "").strip().lower()
|
||||||
|
if guild == guild_part and _normalize_channel_query(ch["name"]) == ch_part:
|
||||||
|
return ch["id"]
|
||||||
|
|
||||||
|
# 3. Partial prefix match (only if unambiguous)
|
||||||
|
matches = [ch for ch in channels if _normalize_channel_query(ch["name"]).startswith(query)]
|
||||||
|
if len(matches) == 1:
|
||||||
|
return matches[0]["id"]
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def format_directory_for_display() -> str:
|
||||||
|
"""Format the channel directory as a human-readable list for the model."""
|
||||||
|
directory = load_directory()
|
||||||
|
platforms = directory.get("platforms", {})
|
||||||
|
|
||||||
|
if not any(platforms.values()):
|
||||||
|
return "No messaging platforms connected or no channels discovered yet."
|
||||||
|
|
||||||
|
lines = ["Available messaging targets:\n"]
|
||||||
|
|
||||||
|
for plat_name, channels in sorted(platforms.items()):
|
||||||
|
if not channels:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Group Discord channels by guild
|
||||||
|
if plat_name == "discord":
|
||||||
|
guilds: Dict[str, List] = {}
|
||||||
|
dms: List = []
|
||||||
|
for ch in channels:
|
||||||
|
guild = ch.get("guild")
|
||||||
|
if guild:
|
||||||
|
guilds.setdefault(guild, []).append(ch)
|
||||||
|
else:
|
||||||
|
dms.append(ch)
|
||||||
|
|
||||||
|
for guild_name, guild_channels in sorted(guilds.items()):
|
||||||
|
lines.append(f"Discord ({guild_name}):")
|
||||||
|
for ch in sorted(guild_channels, key=lambda c: c["name"]):
|
||||||
|
lines.append(f" discord:{_channel_target_name(plat_name, ch)}")
|
||||||
|
if dms:
|
||||||
|
lines.append("Discord (DMs):")
|
||||||
|
for ch in dms:
|
||||||
|
lines.append(f" discord:{_channel_target_name(plat_name, ch)}")
|
||||||
|
lines.append("")
|
||||||
|
else:
|
||||||
|
lines.append(f"{plat_name.title()}:")
|
||||||
|
for ch in channels:
|
||||||
|
lines.append(f" {plat_name}:{_channel_target_name(plat_name, ch)}")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
lines.append('Use these as the "target" parameter when sending.')
|
||||||
|
lines.append('Bare platform name (e.g. "telegram") sends to home channel.')
|
||||||
|
|
||||||
|
return "\n".join(lines)
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,256 @@
|
|||||||
|
"""
|
||||||
|
Delivery routing for cron job outputs and agent responses.
|
||||||
|
|
||||||
|
Routes messages to the appropriate destination based on:
|
||||||
|
- Explicit targets (e.g., "telegram:123456789")
|
||||||
|
- Platform home channels (e.g., "telegram" → home channel)
|
||||||
|
- Origin (back to where the job was created)
|
||||||
|
- Local (always saved to files)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
from datetime import datetime
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Dict, List, Optional, Any
|
||||||
|
|
||||||
|
from hermes_cli.config import get_hermes_home
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
MAX_PLATFORM_OUTPUT = 4000
|
||||||
|
TRUNCATED_VISIBLE = 3800
|
||||||
|
|
||||||
|
from .config import Platform, GatewayConfig
|
||||||
|
from .session import SessionSource
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DeliveryTarget:
|
||||||
|
"""
|
||||||
|
A single delivery target.
|
||||||
|
|
||||||
|
Represents where a message should be sent:
|
||||||
|
- "origin" → back to source
|
||||||
|
- "local" → save to local files
|
||||||
|
- "telegram" → Telegram home channel
|
||||||
|
- "telegram:123456" → specific Telegram chat
|
||||||
|
"""
|
||||||
|
platform: Platform
|
||||||
|
chat_id: Optional[str] = None # None means use home channel
|
||||||
|
thread_id: Optional[str] = None
|
||||||
|
is_origin: bool = False
|
||||||
|
is_explicit: bool = False # True if chat_id was explicitly specified
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def parse(cls, target: str, origin: Optional[SessionSource] = None) -> "DeliveryTarget":
|
||||||
|
"""
|
||||||
|
Parse a delivery target string.
|
||||||
|
|
||||||
|
Formats:
|
||||||
|
- "origin" → back to source
|
||||||
|
- "local" → local files only
|
||||||
|
- "telegram" → Telegram home channel
|
||||||
|
- "telegram:123456" → specific Telegram chat
|
||||||
|
"""
|
||||||
|
target = target.strip().lower()
|
||||||
|
|
||||||
|
if target == "origin":
|
||||||
|
if origin:
|
||||||
|
return cls(
|
||||||
|
platform=origin.platform,
|
||||||
|
chat_id=origin.chat_id,
|
||||||
|
thread_id=origin.thread_id,
|
||||||
|
is_origin=True,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Fallback to local if no origin
|
||||||
|
return cls(platform=Platform.LOCAL, is_origin=True)
|
||||||
|
|
||||||
|
if target == "local":
|
||||||
|
return cls(platform=Platform.LOCAL)
|
||||||
|
|
||||||
|
# Check for platform:chat_id or platform:chat_id:thread_id format
|
||||||
|
if ":" in target:
|
||||||
|
parts = target.split(":", 2)
|
||||||
|
platform_str = parts[0]
|
||||||
|
chat_id = parts[1] if len(parts) > 1 else None
|
||||||
|
thread_id = parts[2] if len(parts) > 2 else None
|
||||||
|
try:
|
||||||
|
platform = Platform(platform_str)
|
||||||
|
return cls(platform=platform, chat_id=chat_id, thread_id=thread_id, is_explicit=True)
|
||||||
|
except ValueError:
|
||||||
|
# Unknown platform, treat as local
|
||||||
|
return cls(platform=Platform.LOCAL)
|
||||||
|
|
||||||
|
# Just a platform name (use home channel)
|
||||||
|
try:
|
||||||
|
platform = Platform(target)
|
||||||
|
return cls(platform=platform)
|
||||||
|
except ValueError:
|
||||||
|
# Unknown platform, treat as local
|
||||||
|
return cls(platform=Platform.LOCAL)
|
||||||
|
|
||||||
|
def to_string(self) -> str:
|
||||||
|
"""Convert back to string format."""
|
||||||
|
if self.is_origin:
|
||||||
|
return "origin"
|
||||||
|
if self.platform == Platform.LOCAL:
|
||||||
|
return "local"
|
||||||
|
if self.chat_id and self.thread_id:
|
||||||
|
return f"{self.platform.value}:{self.chat_id}:{self.thread_id}"
|
||||||
|
if self.chat_id:
|
||||||
|
return f"{self.platform.value}:{self.chat_id}"
|
||||||
|
return self.platform.value
|
||||||
|
|
||||||
|
|
||||||
|
class DeliveryRouter:
|
||||||
|
"""
|
||||||
|
Routes messages to appropriate destinations.
|
||||||
|
|
||||||
|
Handles the logic of resolving delivery targets and dispatching
|
||||||
|
messages to the right platform adapters.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, config: GatewayConfig, adapters: Dict[Platform, Any] = None):
|
||||||
|
"""
|
||||||
|
Initialize the delivery router.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: Gateway configuration
|
||||||
|
adapters: Dict mapping platforms to their adapter instances
|
||||||
|
"""
|
||||||
|
self.config = config
|
||||||
|
self.adapters = adapters or {}
|
||||||
|
self.output_dir = get_hermes_home() / "cron" / "output"
|
||||||
|
|
||||||
|
async def deliver(
|
||||||
|
self,
|
||||||
|
content: str,
|
||||||
|
targets: List[DeliveryTarget],
|
||||||
|
job_id: Optional[str] = None,
|
||||||
|
job_name: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
Deliver content to all specified targets.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
content: The message/output to deliver
|
||||||
|
targets: List of delivery targets
|
||||||
|
job_id: Optional job ID (for cron jobs)
|
||||||
|
job_name: Optional job name
|
||||||
|
metadata: Additional metadata to include
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with delivery results per target
|
||||||
|
"""
|
||||||
|
results = {}
|
||||||
|
|
||||||
|
for target in targets:
|
||||||
|
try:
|
||||||
|
if target.platform == Platform.LOCAL:
|
||||||
|
result = self._deliver_local(content, job_id, job_name, metadata)
|
||||||
|
else:
|
||||||
|
result = await self._deliver_to_platform(target, content, metadata)
|
||||||
|
|
||||||
|
results[target.to_string()] = {
|
||||||
|
"success": True,
|
||||||
|
"result": result
|
||||||
|
}
|
||||||
|
except Exception as e:
|
||||||
|
results[target.to_string()] = {
|
||||||
|
"success": False,
|
||||||
|
"error": str(e)
|
||||||
|
}
|
||||||
|
|
||||||
|
return results
|
||||||
|
|
||||||
|
def _deliver_local(
|
||||||
|
self,
|
||||||
|
content: str,
|
||||||
|
job_id: Optional[str],
|
||||||
|
job_name: Optional[str],
|
||||||
|
metadata: Optional[Dict[str, Any]]
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Save content to local files."""
|
||||||
|
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||||
|
|
||||||
|
if job_id:
|
||||||
|
output_path = self.output_dir / job_id / f"{timestamp}.md"
|
||||||
|
else:
|
||||||
|
output_path = self.output_dir / "misc" / f"{timestamp}.md"
|
||||||
|
|
||||||
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
# Build the output document
|
||||||
|
lines = []
|
||||||
|
if job_name:
|
||||||
|
lines.append(f"# {job_name}")
|
||||||
|
else:
|
||||||
|
lines.append("# Delivery Output")
|
||||||
|
|
||||||
|
lines.append("")
|
||||||
|
lines.append(f"**Timestamp:** {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
|
||||||
|
|
||||||
|
if job_id:
|
||||||
|
lines.append(f"**Job ID:** {job_id}")
|
||||||
|
|
||||||
|
if metadata:
|
||||||
|
for key, value in metadata.items():
|
||||||
|
lines.append(f"**{key}:** {value}")
|
||||||
|
|
||||||
|
lines.append("")
|
||||||
|
lines.append("---")
|
||||||
|
lines.append("")
|
||||||
|
lines.append(content)
|
||||||
|
|
||||||
|
output_path.write_text("\n".join(lines))
|
||||||
|
|
||||||
|
return {
|
||||||
|
"path": str(output_path),
|
||||||
|
"timestamp": timestamp
|
||||||
|
}
|
||||||
|
|
||||||
|
def _save_full_output(self, content: str, job_id: str) -> Path:
|
||||||
|
"""Save full cron output to disk and return the file path."""
|
||||||
|
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||||
|
out_dir = get_hermes_home() / "cron" / "output"
|
||||||
|
out_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
path = out_dir / f"{job_id}_{timestamp}.txt"
|
||||||
|
path.write_text(content)
|
||||||
|
return path
|
||||||
|
|
||||||
|
async def _deliver_to_platform(
|
||||||
|
self,
|
||||||
|
target: DeliveryTarget,
|
||||||
|
content: str,
|
||||||
|
metadata: Optional[Dict[str, Any]]
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Deliver content to a messaging platform."""
|
||||||
|
adapter = self.adapters.get(target.platform)
|
||||||
|
|
||||||
|
if not adapter:
|
||||||
|
raise ValueError(f"No adapter configured for {target.platform.value}")
|
||||||
|
|
||||||
|
if not target.chat_id:
|
||||||
|
raise ValueError(f"No chat ID for {target.platform.value} delivery")
|
||||||
|
|
||||||
|
# Guard: truncate oversized cron output to stay within platform limits
|
||||||
|
if len(content) > MAX_PLATFORM_OUTPUT:
|
||||||
|
job_id = (metadata or {}).get("job_id", "unknown")
|
||||||
|
saved_path = self._save_full_output(content, job_id)
|
||||||
|
logger.info("Cron output truncated (%d chars) — full output: %s", len(content), saved_path)
|
||||||
|
content = (
|
||||||
|
content[:TRUNCATED_VISIBLE]
|
||||||
|
+ f"\n\n... [truncated, full output saved to {saved_path}]"
|
||||||
|
)
|
||||||
|
|
||||||
|
send_metadata = dict(metadata or {})
|
||||||
|
if target.thread_id and "thread_id" not in send_metadata:
|
||||||
|
send_metadata["thread_id"] = target.thread_id
|
||||||
|
return await adapter.send(target.chat_id, content, metadata=send_metadata or None)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,187 @@
|
|||||||
|
"""Per-platform display/verbosity configuration resolver.
|
||||||
|
|
||||||
|
Provides ``resolve_display_setting()`` — the single entry-point for reading
|
||||||
|
display settings with platform-specific overrides and sensible defaults.
|
||||||
|
|
||||||
|
Resolution order (first non-None wins):
|
||||||
|
1. ``display.platforms.<platform>.<key>`` — explicit per-platform user override
|
||||||
|
2. ``display.<key>`` — global user setting
|
||||||
|
3. ``_PLATFORM_DEFAULTS[<platform>][<key>]`` — built-in sensible default
|
||||||
|
4. ``_GLOBAL_DEFAULTS[<key>]`` — built-in global default
|
||||||
|
|
||||||
|
Backward compatibility: ``display.tool_progress_overrides`` is still read as a
|
||||||
|
fallback for ``tool_progress`` when no ``display.platforms`` entry exists. A
|
||||||
|
config migration (version bump) automatically moves the old format into the new
|
||||||
|
``display.platforms`` structure.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Overrideable display settings and their global defaults
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# These are the settings that can be configured per-platform.
|
||||||
|
# Other display settings (compact, personality, skin, etc.) are CLI-only
|
||||||
|
# and don't participate in per-platform resolution.
|
||||||
|
|
||||||
|
_GLOBAL_DEFAULTS: dict[str, Any] = {
|
||||||
|
"tool_progress": "all",
|
||||||
|
"show_reasoning": False,
|
||||||
|
"tool_preview_length": 0,
|
||||||
|
"streaming": None, # None = follow top-level streaming config
|
||||||
|
}
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Sensible per-platform defaults — tiered by platform capability
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Tier 1 (high): Supports message editing, typically personal/team use
|
||||||
|
# Tier 2 (medium): Supports editing but often workspace/customer-facing
|
||||||
|
# Tier 3 (low): No edit support — each progress msg is permanent
|
||||||
|
# Tier 4 (minimal): Batch/non-interactive delivery
|
||||||
|
|
||||||
|
_TIER_HIGH = {
|
||||||
|
"tool_progress": "all",
|
||||||
|
"show_reasoning": False,
|
||||||
|
"tool_preview_length": 40,
|
||||||
|
"streaming": None, # follow global
|
||||||
|
}
|
||||||
|
|
||||||
|
_TIER_MEDIUM = {
|
||||||
|
"tool_progress": "new",
|
||||||
|
"show_reasoning": False,
|
||||||
|
"tool_preview_length": 40,
|
||||||
|
"streaming": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
_TIER_LOW = {
|
||||||
|
"tool_progress": "off",
|
||||||
|
"show_reasoning": False,
|
||||||
|
"tool_preview_length": 40,
|
||||||
|
"streaming": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
_TIER_MINIMAL = {
|
||||||
|
"tool_progress": "off",
|
||||||
|
"show_reasoning": False,
|
||||||
|
"tool_preview_length": 0,
|
||||||
|
"streaming": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
_PLATFORM_DEFAULTS: dict[str, dict[str, Any]] = {
|
||||||
|
# Tier 1 — full edit support, personal/team use
|
||||||
|
"telegram": _TIER_HIGH,
|
||||||
|
"discord": _TIER_HIGH,
|
||||||
|
|
||||||
|
# Tier 2 — edit support, often customer/workspace channels
|
||||||
|
"slack": _TIER_MEDIUM,
|
||||||
|
"mattermost": _TIER_MEDIUM,
|
||||||
|
"matrix": _TIER_MEDIUM,
|
||||||
|
"feishu": _TIER_MEDIUM,
|
||||||
|
|
||||||
|
# Tier 3 — no edit support, progress messages are permanent
|
||||||
|
"signal": _TIER_LOW,
|
||||||
|
"whatsapp": _TIER_MEDIUM, # Baileys bridge supports /edit
|
||||||
|
"bluebubbles": _TIER_LOW,
|
||||||
|
"weixin": _TIER_LOW,
|
||||||
|
"wecom": _TIER_LOW,
|
||||||
|
"wecom_callback": _TIER_LOW,
|
||||||
|
"dingtalk": _TIER_LOW,
|
||||||
|
|
||||||
|
# Tier 4 — batch or non-interactive delivery
|
||||||
|
"email": _TIER_MINIMAL,
|
||||||
|
"sms": _TIER_MINIMAL,
|
||||||
|
"webhook": _TIER_MINIMAL,
|
||||||
|
"homeassistant": _TIER_MINIMAL,
|
||||||
|
"api_server": {**_TIER_HIGH, "tool_preview_length": 0},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Canonical set of per-platform overrideable keys (for validation).
|
||||||
|
OVERRIDEABLE_KEYS = frozenset(_GLOBAL_DEFAULTS.keys())
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_display_setting(
|
||||||
|
user_config: dict,
|
||||||
|
platform_key: str,
|
||||||
|
setting: str,
|
||||||
|
fallback: Any = None,
|
||||||
|
) -> Any:
|
||||||
|
"""Resolve a display setting with per-platform override support.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
user_config : dict
|
||||||
|
The full parsed config.yaml dict.
|
||||||
|
platform_key : str
|
||||||
|
Platform config key (e.g. ``"telegram"``, ``"slack"``). Use
|
||||||
|
``_platform_config_key(source.platform)`` from gateway/run.py.
|
||||||
|
setting : str
|
||||||
|
Display setting name (e.g. ``"tool_progress"``, ``"show_reasoning"``).
|
||||||
|
fallback : Any
|
||||||
|
Fallback value when the setting isn't found anywhere.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
The resolved value, or *fallback* if nothing is configured.
|
||||||
|
"""
|
||||||
|
display_cfg = user_config.get("display") or {}
|
||||||
|
|
||||||
|
# 1. Explicit per-platform override (display.platforms.<platform>.<key>)
|
||||||
|
platforms = display_cfg.get("platforms") or {}
|
||||||
|
plat_overrides = platforms.get(platform_key)
|
||||||
|
if isinstance(plat_overrides, dict):
|
||||||
|
val = plat_overrides.get(setting)
|
||||||
|
if val is not None:
|
||||||
|
return _normalise(setting, val)
|
||||||
|
|
||||||
|
# 1b. Backward compat: display.tool_progress_overrides.<platform>
|
||||||
|
if setting == "tool_progress":
|
||||||
|
legacy = display_cfg.get("tool_progress_overrides")
|
||||||
|
if isinstance(legacy, dict):
|
||||||
|
val = legacy.get(platform_key)
|
||||||
|
if val is not None:
|
||||||
|
return _normalise(setting, val)
|
||||||
|
|
||||||
|
# 2. Global user setting (display.<key>)
|
||||||
|
val = display_cfg.get(setting)
|
||||||
|
if val is not None:
|
||||||
|
return _normalise(setting, val)
|
||||||
|
|
||||||
|
# 3. Built-in platform default
|
||||||
|
plat_defaults = _PLATFORM_DEFAULTS.get(platform_key)
|
||||||
|
if plat_defaults:
|
||||||
|
val = plat_defaults.get(setting)
|
||||||
|
if val is not None:
|
||||||
|
return val
|
||||||
|
|
||||||
|
# 4. Built-in global default
|
||||||
|
val = _GLOBAL_DEFAULTS.get(setting)
|
||||||
|
if val is not None:
|
||||||
|
return val
|
||||||
|
|
||||||
|
return fallback
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _normalise(setting: str, value: Any) -> Any:
|
||||||
|
"""Normalise YAML quirks (bare ``off`` → False in YAML 1.1)."""
|
||||||
|
if setting == "tool_progress":
|
||||||
|
if value is False:
|
||||||
|
return "off"
|
||||||
|
if value is True:
|
||||||
|
return "all"
|
||||||
|
return str(value).lower()
|
||||||
|
if setting in ("show_reasoning", "streaming"):
|
||||||
|
if isinstance(value, str):
|
||||||
|
return value.lower() in ("true", "1", "yes", "on")
|
||||||
|
return bool(value)
|
||||||
|
if setting == "tool_preview_length":
|
||||||
|
try:
|
||||||
|
return int(value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return 0
|
||||||
|
return value
|
||||||
@@ -0,0 +1,170 @@
|
|||||||
|
"""
|
||||||
|
Event Hook System
|
||||||
|
|
||||||
|
A lightweight event-driven system that fires handlers at key lifecycle points.
|
||||||
|
Hooks are discovered from ~/.hermes/hooks/ directories, each containing:
|
||||||
|
- HOOK.yaml (metadata: name, description, events list)
|
||||||
|
- handler.py (Python handler with async def handle(event_type, context))
|
||||||
|
|
||||||
|
Events:
|
||||||
|
- gateway:startup -- Gateway process starts
|
||||||
|
- session:start -- New session created (first message of a new session)
|
||||||
|
- session:end -- Session ends (user ran /new or /reset)
|
||||||
|
- session:reset -- Session reset completed (new session entry created)
|
||||||
|
- agent:start -- Agent begins processing a message
|
||||||
|
- agent:step -- Each turn in the tool-calling loop
|
||||||
|
- agent:end -- Agent finishes processing
|
||||||
|
- command:* -- Any slash command executed (wildcard match)
|
||||||
|
|
||||||
|
Errors in hooks are caught and logged but never block the main pipeline.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import importlib.util
|
||||||
|
from typing import Any, Callable, Dict, List, Optional
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
from hermes_cli.config import get_hermes_home
|
||||||
|
|
||||||
|
|
||||||
|
HOOKS_DIR = get_hermes_home() / "hooks"
|
||||||
|
|
||||||
|
|
||||||
|
class HookRegistry:
|
||||||
|
"""
|
||||||
|
Discovers, loads, and fires event hooks.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
registry = HookRegistry()
|
||||||
|
registry.discover_and_load()
|
||||||
|
await registry.emit("agent:start", {"platform": "telegram", ...})
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
# event_type -> [handler_fn, ...]
|
||||||
|
self._handlers: Dict[str, List[Callable]] = {}
|
||||||
|
self._loaded_hooks: List[dict] = [] # metadata for listing
|
||||||
|
|
||||||
|
@property
|
||||||
|
def loaded_hooks(self) -> List[dict]:
|
||||||
|
"""Return metadata about all loaded hooks."""
|
||||||
|
return list(self._loaded_hooks)
|
||||||
|
|
||||||
|
def _register_builtin_hooks(self) -> None:
|
||||||
|
"""Register built-in hooks that are always active."""
|
||||||
|
try:
|
||||||
|
from gateway.builtin_hooks.boot_md import handle as boot_md_handle
|
||||||
|
|
||||||
|
self._handlers.setdefault("gateway:startup", []).append(boot_md_handle)
|
||||||
|
self._loaded_hooks.append({
|
||||||
|
"name": "boot-md",
|
||||||
|
"description": "Run ~/.hermes/BOOT.md on gateway startup",
|
||||||
|
"events": ["gateway:startup"],
|
||||||
|
"path": "(builtin)",
|
||||||
|
})
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[hooks] Could not load built-in boot-md hook: {e}", flush=True)
|
||||||
|
|
||||||
|
def discover_and_load(self) -> None:
|
||||||
|
"""
|
||||||
|
Scan the hooks directory for hook directories and load their handlers.
|
||||||
|
|
||||||
|
Also registers built-in hooks that are always active.
|
||||||
|
|
||||||
|
Each hook directory must contain:
|
||||||
|
- HOOK.yaml with at least 'name' and 'events' keys
|
||||||
|
- handler.py with a top-level 'handle' function (sync or async)
|
||||||
|
"""
|
||||||
|
self._register_builtin_hooks()
|
||||||
|
|
||||||
|
if not HOOKS_DIR.exists():
|
||||||
|
return
|
||||||
|
|
||||||
|
for hook_dir in sorted(HOOKS_DIR.iterdir()):
|
||||||
|
if not hook_dir.is_dir():
|
||||||
|
continue
|
||||||
|
|
||||||
|
manifest_path = hook_dir / "HOOK.yaml"
|
||||||
|
handler_path = hook_dir / "handler.py"
|
||||||
|
|
||||||
|
if not manifest_path.exists() or not handler_path.exists():
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
manifest = yaml.safe_load(manifest_path.read_text(encoding="utf-8"))
|
||||||
|
if not manifest or not isinstance(manifest, dict):
|
||||||
|
print(f"[hooks] Skipping {hook_dir.name}: invalid HOOK.yaml", flush=True)
|
||||||
|
continue
|
||||||
|
|
||||||
|
hook_name = manifest.get("name", hook_dir.name)
|
||||||
|
events = manifest.get("events", [])
|
||||||
|
if not events:
|
||||||
|
print(f"[hooks] Skipping {hook_name}: no events declared", flush=True)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Dynamically load the handler module
|
||||||
|
spec = importlib.util.spec_from_file_location(
|
||||||
|
f"hermes_hook_{hook_name}", handler_path
|
||||||
|
)
|
||||||
|
if spec is None or spec.loader is None:
|
||||||
|
print(f"[hooks] Skipping {hook_name}: could not load handler.py", flush=True)
|
||||||
|
continue
|
||||||
|
|
||||||
|
module = importlib.util.module_from_spec(spec)
|
||||||
|
spec.loader.exec_module(module)
|
||||||
|
|
||||||
|
handle_fn = getattr(module, "handle", None)
|
||||||
|
if handle_fn is None:
|
||||||
|
print(f"[hooks] Skipping {hook_name}: no 'handle' function found", flush=True)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Register the handler for each declared event
|
||||||
|
for event in events:
|
||||||
|
self._handlers.setdefault(event, []).append(handle_fn)
|
||||||
|
|
||||||
|
self._loaded_hooks.append({
|
||||||
|
"name": hook_name,
|
||||||
|
"description": manifest.get("description", ""),
|
||||||
|
"events": events,
|
||||||
|
"path": str(hook_dir),
|
||||||
|
})
|
||||||
|
|
||||||
|
print(f"[hooks] Loaded hook '{hook_name}' for events: {events}", flush=True)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[hooks] Error loading hook {hook_dir.name}: {e}", flush=True)
|
||||||
|
|
||||||
|
async def emit(self, event_type: str, context: Optional[Dict[str, Any]] = None) -> None:
|
||||||
|
"""
|
||||||
|
Fire all handlers registered for an event.
|
||||||
|
|
||||||
|
Supports wildcard matching: handlers registered for "command:*" will
|
||||||
|
fire for any "command:..." event. Handlers registered for a base type
|
||||||
|
like "agent" won't fire for "agent:start" -- only exact matches and
|
||||||
|
explicit wildcards.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
event_type: The event identifier (e.g. "agent:start").
|
||||||
|
context: Optional dict with event-specific data.
|
||||||
|
"""
|
||||||
|
if context is None:
|
||||||
|
context = {}
|
||||||
|
|
||||||
|
# Collect handlers: exact match + wildcard match
|
||||||
|
handlers = list(self._handlers.get(event_type, []))
|
||||||
|
|
||||||
|
# Check for wildcard patterns (e.g., "command:*" matches "command:reset")
|
||||||
|
if ":" in event_type:
|
||||||
|
base = event_type.split(":")[0]
|
||||||
|
wildcard_key = f"{base}:*"
|
||||||
|
handlers.extend(self._handlers.get(wildcard_key, []))
|
||||||
|
|
||||||
|
for fn in handlers:
|
||||||
|
try:
|
||||||
|
result = fn(event_type, context)
|
||||||
|
# Support both sync and async handlers
|
||||||
|
if asyncio.iscoroutine(result):
|
||||||
|
await result
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[hooks] Error in handler for '{event_type}': {e}", flush=True)
|
||||||
@@ -0,0 +1,132 @@
|
|||||||
|
"""
|
||||||
|
Session mirroring for cross-platform message delivery.
|
||||||
|
|
||||||
|
When a message is sent to a platform (via send_message or cron delivery),
|
||||||
|
this module appends a "delivery-mirror" record to the target session's
|
||||||
|
transcript so the receiving-side agent has context about what was sent.
|
||||||
|
|
||||||
|
Standalone -- works from CLI, cron, and gateway contexts without needing
|
||||||
|
the full SessionStore machinery.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from hermes_cli.config import get_hermes_home
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_SESSIONS_DIR = get_hermes_home() / "sessions"
|
||||||
|
_SESSIONS_INDEX = _SESSIONS_DIR / "sessions.json"
|
||||||
|
|
||||||
|
|
||||||
|
def mirror_to_session(
|
||||||
|
platform: str,
|
||||||
|
chat_id: str,
|
||||||
|
message_text: str,
|
||||||
|
source_label: str = "cli",
|
||||||
|
thread_id: Optional[str] = None,
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
Append a delivery-mirror message to the target session's transcript.
|
||||||
|
|
||||||
|
Finds the gateway session that matches the given platform + chat_id,
|
||||||
|
then writes a mirror entry to both the JSONL transcript and SQLite DB.
|
||||||
|
|
||||||
|
Returns True if mirrored successfully, False if no matching session or error.
|
||||||
|
All errors are caught -- this is never fatal.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
session_id = _find_session_id(platform, str(chat_id), thread_id=thread_id)
|
||||||
|
if not session_id:
|
||||||
|
logger.debug("Mirror: no session found for %s:%s:%s", platform, chat_id, thread_id)
|
||||||
|
return False
|
||||||
|
|
||||||
|
mirror_msg = {
|
||||||
|
"role": "assistant",
|
||||||
|
"content": message_text,
|
||||||
|
"timestamp": datetime.now().isoformat(),
|
||||||
|
"mirror": True,
|
||||||
|
"mirror_source": source_label,
|
||||||
|
}
|
||||||
|
|
||||||
|
_append_to_jsonl(session_id, mirror_msg)
|
||||||
|
_append_to_sqlite(session_id, mirror_msg)
|
||||||
|
|
||||||
|
logger.debug("Mirror: wrote to session %s (from %s)", session_id, source_label)
|
||||||
|
return True
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Mirror failed for %s:%s:%s: %s", platform, chat_id, thread_id, e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _find_session_id(platform: str, chat_id: str, thread_id: Optional[str] = None) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
Find the active session_id for a platform + chat_id pair.
|
||||||
|
|
||||||
|
Scans sessions.json entries and matches where origin.chat_id == chat_id
|
||||||
|
on the right platform. DM session keys don't embed the chat_id
|
||||||
|
(e.g. "agent:main:telegram:dm"), so we check the origin dict.
|
||||||
|
"""
|
||||||
|
if not _SESSIONS_INDEX.exists():
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
with open(_SESSIONS_INDEX, encoding="utf-8") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
|
||||||
|
platform_lower = platform.lower()
|
||||||
|
best_match = None
|
||||||
|
best_updated = ""
|
||||||
|
|
||||||
|
for _key, entry in data.items():
|
||||||
|
origin = entry.get("origin") or {}
|
||||||
|
entry_platform = (origin.get("platform") or entry.get("platform", "")).lower()
|
||||||
|
|
||||||
|
if entry_platform != platform_lower:
|
||||||
|
continue
|
||||||
|
|
||||||
|
origin_chat_id = str(origin.get("chat_id", ""))
|
||||||
|
if origin_chat_id == str(chat_id):
|
||||||
|
origin_thread_id = origin.get("thread_id")
|
||||||
|
if thread_id is not None and str(origin_thread_id or "") != str(thread_id):
|
||||||
|
continue
|
||||||
|
updated = entry.get("updated_at", "")
|
||||||
|
if updated > best_updated:
|
||||||
|
best_updated = updated
|
||||||
|
best_match = entry.get("session_id")
|
||||||
|
|
||||||
|
return best_match
|
||||||
|
|
||||||
|
|
||||||
|
def _append_to_jsonl(session_id: str, message: dict) -> None:
|
||||||
|
"""Append a message to the JSONL transcript file."""
|
||||||
|
transcript_path = _SESSIONS_DIR / f"{session_id}.jsonl"
|
||||||
|
try:
|
||||||
|
with open(transcript_path, "a", encoding="utf-8") as f:
|
||||||
|
f.write(json.dumps(message, ensure_ascii=False) + "\n")
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Mirror JSONL write failed: %s", e)
|
||||||
|
|
||||||
|
|
||||||
|
def _append_to_sqlite(session_id: str, message: dict) -> None:
|
||||||
|
"""Append a message to the SQLite session database."""
|
||||||
|
db = None
|
||||||
|
try:
|
||||||
|
from hermes_state import SessionDB
|
||||||
|
db = SessionDB()
|
||||||
|
db.append_message(
|
||||||
|
session_id=session_id,
|
||||||
|
role=message.get("role", "assistant"),
|
||||||
|
content=message.get("content"),
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Mirror SQLite write failed: %s", e)
|
||||||
|
finally:
|
||||||
|
if db is not None:
|
||||||
|
db.close()
|
||||||
@@ -0,0 +1,309 @@
|
|||||||
|
"""
|
||||||
|
DM Pairing System
|
||||||
|
|
||||||
|
Code-based approval flow for authorizing new users on messaging platforms.
|
||||||
|
Instead of static allowlists with user IDs, unknown users receive a one-time
|
||||||
|
pairing code that the bot owner approves via the CLI.
|
||||||
|
|
||||||
|
Security features (based on OWASP + NIST SP 800-63-4 guidance):
|
||||||
|
- 8-char codes from 32-char unambiguous alphabet (no 0/O/1/I)
|
||||||
|
- Cryptographic randomness via secrets.choice()
|
||||||
|
- 1-hour code expiry
|
||||||
|
- Max 3 pending codes per platform
|
||||||
|
- Rate limiting: 1 request per user per 10 minutes
|
||||||
|
- Lockout after 5 failed approval attempts (1 hour)
|
||||||
|
- File permissions: chmod 0600 on all data files
|
||||||
|
- Codes are never logged to stdout
|
||||||
|
|
||||||
|
Storage: ~/.hermes/pairing/
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import secrets
|
||||||
|
import tempfile
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from hermes_constants import get_hermes_dir
|
||||||
|
|
||||||
|
|
||||||
|
# Unambiguous alphabet -- excludes 0/O, 1/I to prevent confusion
|
||||||
|
ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
|
||||||
|
CODE_LENGTH = 8
|
||||||
|
|
||||||
|
# Timing constants
|
||||||
|
CODE_TTL_SECONDS = 3600 # Codes expire after 1 hour
|
||||||
|
RATE_LIMIT_SECONDS = 600 # 1 request per user per 10 minutes
|
||||||
|
LOCKOUT_SECONDS = 3600 # Lockout duration after too many failures
|
||||||
|
|
||||||
|
# Limits
|
||||||
|
MAX_PENDING_PER_PLATFORM = 3 # Max pending codes per platform
|
||||||
|
MAX_FAILED_ATTEMPTS = 5 # Failed approvals before lockout
|
||||||
|
|
||||||
|
PAIRING_DIR = get_hermes_dir("platforms/pairing", "pairing")
|
||||||
|
|
||||||
|
|
||||||
|
def _secure_write(path: Path, data: str) -> None:
|
||||||
|
"""Write data to file with restrictive permissions (owner read/write only).
|
||||||
|
|
||||||
|
Uses a temp-file + atomic rename so readers always see either the old
|
||||||
|
complete file or the new one — never a partial write.
|
||||||
|
"""
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
fd, tmp_path = tempfile.mkstemp(dir=str(path.parent), suffix=".tmp")
|
||||||
|
try:
|
||||||
|
with os.fdopen(fd, "w", encoding="utf-8") as f:
|
||||||
|
f.write(data)
|
||||||
|
f.flush()
|
||||||
|
os.fsync(f.fileno())
|
||||||
|
os.replace(tmp_path, str(path))
|
||||||
|
try:
|
||||||
|
os.chmod(path, 0o600)
|
||||||
|
except OSError:
|
||||||
|
pass # Windows doesn't support chmod the same way
|
||||||
|
except BaseException:
|
||||||
|
try:
|
||||||
|
os.unlink(tmp_path)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
class PairingStore:
|
||||||
|
"""
|
||||||
|
Manages pairing codes and approved user lists.
|
||||||
|
|
||||||
|
Data files per platform:
|
||||||
|
- {platform}-pending.json : pending pairing requests
|
||||||
|
- {platform}-approved.json : approved (paired) users
|
||||||
|
- _rate_limits.json : rate limit tracking
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
PAIRING_DIR.mkdir(parents=True, exist_ok=True)
|
||||||
|
# Protects all read-modify-write cycles. The gateway runs multiple
|
||||||
|
# platform adapters concurrently in threads sharing one PairingStore.
|
||||||
|
self._lock = threading.RLock()
|
||||||
|
|
||||||
|
def _pending_path(self, platform: str) -> Path:
|
||||||
|
return PAIRING_DIR / f"{platform}-pending.json"
|
||||||
|
|
||||||
|
def _approved_path(self, platform: str) -> Path:
|
||||||
|
return PAIRING_DIR / f"{platform}-approved.json"
|
||||||
|
|
||||||
|
def _rate_limit_path(self) -> Path:
|
||||||
|
return PAIRING_DIR / "_rate_limits.json"
|
||||||
|
|
||||||
|
def _load_json(self, path: Path) -> dict:
|
||||||
|
if path.exists():
|
||||||
|
try:
|
||||||
|
return json.loads(path.read_text(encoding="utf-8"))
|
||||||
|
except (json.JSONDecodeError, OSError):
|
||||||
|
return {}
|
||||||
|
return {}
|
||||||
|
|
||||||
|
def _save_json(self, path: Path, data: dict) -> None:
|
||||||
|
_secure_write(path, json.dumps(data, indent=2, ensure_ascii=False))
|
||||||
|
|
||||||
|
# ----- Approved users -----
|
||||||
|
|
||||||
|
def is_approved(self, platform: str, user_id: str) -> bool:
|
||||||
|
"""Check if a user is approved (paired) on a platform."""
|
||||||
|
approved = self._load_json(self._approved_path(platform))
|
||||||
|
return user_id in approved
|
||||||
|
|
||||||
|
def list_approved(self, platform: str = None) -> list:
|
||||||
|
"""List approved users, optionally filtered by platform."""
|
||||||
|
results = []
|
||||||
|
platforms = [platform] if platform else self._all_platforms("approved")
|
||||||
|
for p in platforms:
|
||||||
|
approved = self._load_json(self._approved_path(p))
|
||||||
|
for uid, info in approved.items():
|
||||||
|
results.append({"platform": p, "user_id": uid, **info})
|
||||||
|
return results
|
||||||
|
|
||||||
|
def _approve_user(self, platform: str, user_id: str, user_name: str = "") -> None:
|
||||||
|
"""Add a user to the approved list. Must be called under self._lock."""
|
||||||
|
approved = self._load_json(self._approved_path(platform))
|
||||||
|
approved[user_id] = {
|
||||||
|
"user_name": user_name,
|
||||||
|
"approved_at": time.time(),
|
||||||
|
}
|
||||||
|
self._save_json(self._approved_path(platform), approved)
|
||||||
|
|
||||||
|
def revoke(self, platform: str, user_id: str) -> bool:
|
||||||
|
"""Remove a user from the approved list. Returns True if found."""
|
||||||
|
path = self._approved_path(platform)
|
||||||
|
with self._lock:
|
||||||
|
approved = self._load_json(path)
|
||||||
|
if user_id in approved:
|
||||||
|
del approved[user_id]
|
||||||
|
self._save_json(path, approved)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
# ----- Pending codes -----
|
||||||
|
|
||||||
|
def generate_code(
|
||||||
|
self, platform: str, user_id: str, user_name: str = ""
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
Generate a pairing code for a new user.
|
||||||
|
|
||||||
|
Returns the code string, or None if:
|
||||||
|
- User is rate-limited (too recent request)
|
||||||
|
- Max pending codes reached for this platform
|
||||||
|
- User/platform is in lockout due to failed attempts
|
||||||
|
"""
|
||||||
|
with self._lock:
|
||||||
|
self._cleanup_expired(platform)
|
||||||
|
|
||||||
|
# Check lockout
|
||||||
|
if self._is_locked_out(platform):
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Check rate limit for this specific user
|
||||||
|
if self._is_rate_limited(platform, user_id):
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Check max pending
|
||||||
|
pending = self._load_json(self._pending_path(platform))
|
||||||
|
if len(pending) >= MAX_PENDING_PER_PLATFORM:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Generate cryptographically random code
|
||||||
|
code = "".join(secrets.choice(ALPHABET) for _ in range(CODE_LENGTH))
|
||||||
|
|
||||||
|
# Store pending request
|
||||||
|
pending[code] = {
|
||||||
|
"user_id": user_id,
|
||||||
|
"user_name": user_name,
|
||||||
|
"created_at": time.time(),
|
||||||
|
}
|
||||||
|
self._save_json(self._pending_path(platform), pending)
|
||||||
|
|
||||||
|
# Record rate limit
|
||||||
|
self._record_rate_limit(platform, user_id)
|
||||||
|
|
||||||
|
return code
|
||||||
|
|
||||||
|
def approve_code(self, platform: str, code: str) -> Optional[dict]:
|
||||||
|
"""
|
||||||
|
Approve a pairing code. Adds the user to the approved list.
|
||||||
|
|
||||||
|
Returns {user_id, user_name} on success, None if code is invalid/expired.
|
||||||
|
"""
|
||||||
|
with self._lock:
|
||||||
|
self._cleanup_expired(platform)
|
||||||
|
code = code.upper().strip()
|
||||||
|
|
||||||
|
pending = self._load_json(self._pending_path(platform))
|
||||||
|
if code not in pending:
|
||||||
|
self._record_failed_attempt(platform)
|
||||||
|
return None
|
||||||
|
|
||||||
|
entry = pending.pop(code)
|
||||||
|
self._save_json(self._pending_path(platform), pending)
|
||||||
|
|
||||||
|
# Add to approved list
|
||||||
|
self._approve_user(platform, entry["user_id"], entry.get("user_name", ""))
|
||||||
|
|
||||||
|
return {
|
||||||
|
"user_id": entry["user_id"],
|
||||||
|
"user_name": entry.get("user_name", ""),
|
||||||
|
}
|
||||||
|
|
||||||
|
def list_pending(self, platform: str = None) -> list:
|
||||||
|
"""List pending pairing requests, optionally filtered by platform."""
|
||||||
|
results = []
|
||||||
|
platforms = [platform] if platform else self._all_platforms("pending")
|
||||||
|
for p in platforms:
|
||||||
|
self._cleanup_expired(p)
|
||||||
|
pending = self._load_json(self._pending_path(p))
|
||||||
|
for code, info in pending.items():
|
||||||
|
age_min = int((time.time() - info["created_at"]) / 60)
|
||||||
|
results.append({
|
||||||
|
"platform": p,
|
||||||
|
"code": code,
|
||||||
|
"user_id": info["user_id"],
|
||||||
|
"user_name": info.get("user_name", ""),
|
||||||
|
"age_minutes": age_min,
|
||||||
|
})
|
||||||
|
return results
|
||||||
|
|
||||||
|
def clear_pending(self, platform: str = None) -> int:
|
||||||
|
"""Clear all pending requests. Returns count removed."""
|
||||||
|
with self._lock:
|
||||||
|
count = 0
|
||||||
|
platforms = [platform] if platform else self._all_platforms("pending")
|
||||||
|
for p in platforms:
|
||||||
|
pending = self._load_json(self._pending_path(p))
|
||||||
|
count += len(pending)
|
||||||
|
self._save_json(self._pending_path(p), {})
|
||||||
|
return count
|
||||||
|
|
||||||
|
# ----- Rate limiting and lockout -----
|
||||||
|
|
||||||
|
def _is_rate_limited(self, platform: str, user_id: str) -> bool:
|
||||||
|
"""Check if a user has requested a code too recently."""
|
||||||
|
limits = self._load_json(self._rate_limit_path())
|
||||||
|
key = f"{platform}:{user_id}"
|
||||||
|
last_request = limits.get(key, 0)
|
||||||
|
return (time.time() - last_request) < RATE_LIMIT_SECONDS
|
||||||
|
|
||||||
|
def _record_rate_limit(self, platform: str, user_id: str) -> None:
|
||||||
|
"""Record the time of a pairing request for rate limiting."""
|
||||||
|
limits = self._load_json(self._rate_limit_path())
|
||||||
|
key = f"{platform}:{user_id}"
|
||||||
|
limits[key] = time.time()
|
||||||
|
self._save_json(self._rate_limit_path(), limits)
|
||||||
|
|
||||||
|
def _is_locked_out(self, platform: str) -> bool:
|
||||||
|
"""Check if a platform is in lockout due to failed approval attempts."""
|
||||||
|
limits = self._load_json(self._rate_limit_path())
|
||||||
|
lockout_key = f"_lockout:{platform}"
|
||||||
|
lockout_until = limits.get(lockout_key, 0)
|
||||||
|
return time.time() < lockout_until
|
||||||
|
|
||||||
|
def _record_failed_attempt(self, platform: str) -> None:
|
||||||
|
"""Record a failed approval attempt. Triggers lockout after MAX_FAILED_ATTEMPTS."""
|
||||||
|
limits = self._load_json(self._rate_limit_path())
|
||||||
|
fail_key = f"_failures:{platform}"
|
||||||
|
fails = limits.get(fail_key, 0) + 1
|
||||||
|
limits[fail_key] = fails
|
||||||
|
if fails >= MAX_FAILED_ATTEMPTS:
|
||||||
|
lockout_key = f"_lockout:{platform}"
|
||||||
|
limits[lockout_key] = time.time() + LOCKOUT_SECONDS
|
||||||
|
limits[fail_key] = 0 # Reset counter
|
||||||
|
print(f"[pairing] Platform {platform} locked out for {LOCKOUT_SECONDS}s "
|
||||||
|
f"after {MAX_FAILED_ATTEMPTS} failed attempts", flush=True)
|
||||||
|
self._save_json(self._rate_limit_path(), limits)
|
||||||
|
|
||||||
|
# ----- Cleanup -----
|
||||||
|
|
||||||
|
def _cleanup_expired(self, platform: str) -> None:
|
||||||
|
"""Remove expired pending codes."""
|
||||||
|
path = self._pending_path(platform)
|
||||||
|
pending = self._load_json(path)
|
||||||
|
now = time.time()
|
||||||
|
expired = [
|
||||||
|
code for code, info in pending.items()
|
||||||
|
if (now - info["created_at"]) > CODE_TTL_SECONDS
|
||||||
|
]
|
||||||
|
if expired:
|
||||||
|
for code in expired:
|
||||||
|
del pending[code]
|
||||||
|
self._save_json(path, pending)
|
||||||
|
|
||||||
|
def _all_platforms(self, suffix: str) -> list:
|
||||||
|
"""List all platforms that have data files of a given suffix."""
|
||||||
|
platforms = []
|
||||||
|
for f in PAIRING_DIR.iterdir():
|
||||||
|
if f.name.endswith(f"-{suffix}.json"):
|
||||||
|
platform = f.name.replace(f"-{suffix}.json", "")
|
||||||
|
if not platform.startswith("_"):
|
||||||
|
platforms.append(platform)
|
||||||
|
return platforms
|
||||||
@@ -0,0 +1,313 @@
|
|||||||
|
# Adding a New Messaging Platform
|
||||||
|
|
||||||
|
Checklist for integrating a new messaging platform into the Hermes gateway.
|
||||||
|
Use this as a reference when building a new adapter — every item here is a
|
||||||
|
real integration point that exists in the codebase. Missing any of them will
|
||||||
|
cause broken functionality, missing features, or inconsistent behavior.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. Core Adapter (`gateway/platforms/<platform>.py`)
|
||||||
|
|
||||||
|
The adapter is a subclass of `BasePlatformAdapter` from `gateway/platforms/base.py`.
|
||||||
|
|
||||||
|
### Required methods
|
||||||
|
|
||||||
|
| Method | Purpose |
|
||||||
|
|--------|---------|
|
||||||
|
| `__init__(self, config)` | Parse config, init state. Call `super().__init__(config, Platform.YOUR_PLATFORM)` |
|
||||||
|
| `connect() -> bool` | Connect to the platform, start listeners. Return True on success |
|
||||||
|
| `disconnect()` | Stop listeners, close connections, cancel tasks |
|
||||||
|
| `send(chat_id, text, ...) -> SendResult` | Send a text message |
|
||||||
|
| `send_typing(chat_id)` | Send typing indicator |
|
||||||
|
| `send_image(chat_id, image_url, caption) -> SendResult` | Send an image |
|
||||||
|
| `get_chat_info(chat_id) -> dict` | Return `{name, type, chat_id}` for a chat |
|
||||||
|
|
||||||
|
### Optional methods (have default stubs in base)
|
||||||
|
|
||||||
|
| Method | Purpose |
|
||||||
|
|--------|---------|
|
||||||
|
| `send_document(chat_id, path, caption)` | Send a file attachment |
|
||||||
|
| `send_voice(chat_id, path)` | Send a voice message |
|
||||||
|
| `send_video(chat_id, path, caption)` | Send a video |
|
||||||
|
| `send_animation(chat_id, path, caption)` | Send a GIF/animation |
|
||||||
|
| `send_image_file(chat_id, path, caption)` | Send image from local file |
|
||||||
|
|
||||||
|
### Required function
|
||||||
|
|
||||||
|
```python
|
||||||
|
def check_<platform>_requirements() -> bool:
|
||||||
|
"""Check if this platform's dependencies are available."""
|
||||||
|
```
|
||||||
|
|
||||||
|
### Key patterns to follow
|
||||||
|
|
||||||
|
- Use `self.build_source(...)` to construct `SessionSource` objects
|
||||||
|
- Call `self.handle_message(event)` to dispatch inbound messages to the gateway
|
||||||
|
- Use `MessageEvent`, `MessageType`, `SendResult` from base
|
||||||
|
- Use `cache_image_from_bytes`, `cache_audio_from_bytes`, `cache_document_from_bytes` for attachments
|
||||||
|
- Filter self-messages (prevent reply loops)
|
||||||
|
- Filter sync/echo messages if the platform has them
|
||||||
|
- Redact sensitive identifiers (phone numbers, tokens) in all log output
|
||||||
|
- Implement reconnection with exponential backoff + jitter for streaming connections
|
||||||
|
- Set `MAX_MESSAGE_LENGTH` if the platform has message size limits
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. Platform Enum (`gateway/config.py`)
|
||||||
|
|
||||||
|
Add the platform to the `Platform` enum:
|
||||||
|
|
||||||
|
```python
|
||||||
|
class Platform(Enum):
|
||||||
|
...
|
||||||
|
YOUR_PLATFORM = "your_platform"
|
||||||
|
```
|
||||||
|
|
||||||
|
Add env var loading in `_apply_env_overrides()`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# Your Platform
|
||||||
|
your_token = os.getenv("YOUR_PLATFORM_TOKEN")
|
||||||
|
if your_token:
|
||||||
|
if Platform.YOUR_PLATFORM not in config.platforms:
|
||||||
|
config.platforms[Platform.YOUR_PLATFORM] = PlatformConfig()
|
||||||
|
config.platforms[Platform.YOUR_PLATFORM].enabled = True
|
||||||
|
config.platforms[Platform.YOUR_PLATFORM].token = your_token
|
||||||
|
```
|
||||||
|
|
||||||
|
Update `get_connected_platforms()` if your platform doesn't use token/api_key
|
||||||
|
(e.g., WhatsApp uses `enabled` flag, Signal uses `extra` dict).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. Adapter Factory (`gateway/run.py`)
|
||||||
|
|
||||||
|
Add to `_create_adapter()`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
elif platform == Platform.YOUR_PLATFORM:
|
||||||
|
from gateway.platforms.your_platform import YourAdapter, check_your_requirements
|
||||||
|
if not check_your_requirements():
|
||||||
|
logger.warning("Your Platform: dependencies not met")
|
||||||
|
return None
|
||||||
|
return YourAdapter(config)
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. Authorization Maps (`gateway/run.py`)
|
||||||
|
|
||||||
|
Add to BOTH dicts in `_is_user_authorized()`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
platform_env_map = {
|
||||||
|
...
|
||||||
|
Platform.YOUR_PLATFORM: "YOUR_PLATFORM_ALLOWED_USERS",
|
||||||
|
}
|
||||||
|
platform_allow_all_map = {
|
||||||
|
...
|
||||||
|
Platform.YOUR_PLATFORM: "YOUR_PLATFORM_ALLOW_ALL_USERS",
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. Session Source (`gateway/session.py`)
|
||||||
|
|
||||||
|
If your platform needs extra identity fields (e.g., Signal's UUID alongside
|
||||||
|
phone number), add them to the `SessionSource` dataclass with `Optional` defaults,
|
||||||
|
and update `to_dict()`, `from_dict()`, and `build_source()` in base.py.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. System Prompt Hints (`agent/prompt_builder.py`)
|
||||||
|
|
||||||
|
Add a `PLATFORM_HINTS` entry so the agent knows what platform it's on:
|
||||||
|
|
||||||
|
```python
|
||||||
|
PLATFORM_HINTS = {
|
||||||
|
...
|
||||||
|
"your_platform": (
|
||||||
|
"You are on Your Platform. "
|
||||||
|
"Describe formatting capabilities, media support, etc."
|
||||||
|
),
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Without this, the agent won't know it's on your platform and may use
|
||||||
|
inappropriate formatting (e.g., markdown on platforms that don't render it).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. Toolset (`toolsets.py`)
|
||||||
|
|
||||||
|
Add a named toolset for your platform:
|
||||||
|
|
||||||
|
```python
|
||||||
|
"hermes-your-platform": {
|
||||||
|
"description": "Your Platform bot toolset",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
```
|
||||||
|
|
||||||
|
And add it to the `hermes-gateway` composite:
|
||||||
|
|
||||||
|
```python
|
||||||
|
"hermes-gateway": {
|
||||||
|
"includes": [..., "hermes-your-platform"]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 8. Cron Delivery (`cron/scheduler.py`)
|
||||||
|
|
||||||
|
Add to `platform_map` in `_deliver_result()`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
platform_map = {
|
||||||
|
...
|
||||||
|
"your_platform": Platform.YOUR_PLATFORM,
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Without this, `cronjob(action="create", deliver="your_platform", ...)` silently fails.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 9. Send Message Tool (`tools/send_message_tool.py`)
|
||||||
|
|
||||||
|
Add to `platform_map` in `send_message_tool()`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
platform_map = {
|
||||||
|
...
|
||||||
|
"your_platform": Platform.YOUR_PLATFORM,
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Add routing in `_send_to_platform()`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
elif platform == Platform.YOUR_PLATFORM:
|
||||||
|
return await _send_your_platform(pconfig, chat_id, message)
|
||||||
|
```
|
||||||
|
|
||||||
|
Implement `_send_your_platform()` — a standalone async function that sends
|
||||||
|
a single message without requiring the full adapter (for use by cron jobs
|
||||||
|
and the send_message tool outside the gateway process).
|
||||||
|
|
||||||
|
Update the tool schema `target` description to include your platform example.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 10. Cronjob Tool Schema (`tools/cronjob_tools.py`)
|
||||||
|
|
||||||
|
Update the `deliver` parameter description and docstring to mention your
|
||||||
|
platform as a delivery option.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 11. Channel Directory (`gateway/channel_directory.py`)
|
||||||
|
|
||||||
|
If your platform can't enumerate chats (most can't), add it to the
|
||||||
|
session-based discovery list:
|
||||||
|
|
||||||
|
```python
|
||||||
|
for plat_name in ("telegram", "whatsapp", "signal", "your_platform"):
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 12. Status Display (`hermes_cli/status.py`)
|
||||||
|
|
||||||
|
Add to the `platforms` dict in the Messaging Platforms section:
|
||||||
|
|
||||||
|
```python
|
||||||
|
platforms = {
|
||||||
|
...
|
||||||
|
"Your Platform": ("YOUR_PLATFORM_TOKEN", "YOUR_PLATFORM_HOME_CHANNEL"),
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 13. Gateway Setup Wizard (`hermes_cli/gateway.py`)
|
||||||
|
|
||||||
|
Add to the `_PLATFORMS` list:
|
||||||
|
|
||||||
|
```python
|
||||||
|
{
|
||||||
|
"key": "your_platform",
|
||||||
|
"label": "Your Platform",
|
||||||
|
"emoji": "📱",
|
||||||
|
"token_var": "YOUR_PLATFORM_TOKEN",
|
||||||
|
"setup_instructions": [...],
|
||||||
|
"vars": [...],
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
If your platform needs custom setup logic (connectivity testing, QR codes,
|
||||||
|
policy choices), add a `_setup_your_platform()` function and route to it
|
||||||
|
in the platform selection switch.
|
||||||
|
|
||||||
|
Update `_platform_status()` if your platform's "configured" check differs
|
||||||
|
from the standard `bool(get_env_value(token_var))`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 14. Phone/ID Redaction (`agent/redact.py`)
|
||||||
|
|
||||||
|
If your platform uses sensitive identifiers (phone numbers, etc.), add a
|
||||||
|
regex pattern and redaction function to `agent/redact.py`. This ensures
|
||||||
|
identifiers are masked in ALL log output, not just your adapter's logs.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 15. Documentation
|
||||||
|
|
||||||
|
| File | What to update |
|
||||||
|
|------|---------------|
|
||||||
|
| `README.md` | Platform list in feature table + documentation table |
|
||||||
|
| `AGENTS.md` | Gateway description + env var config section |
|
||||||
|
| `website/docs/user-guide/messaging/<platform>.md` | **NEW** — Full setup guide (see existing platform docs for template) |
|
||||||
|
| `website/docs/user-guide/messaging/index.md` | Architecture diagram, toolset table, security examples, Next Steps links |
|
||||||
|
| `website/docs/reference/environment-variables.md` | All env vars for the platform |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 16. Tests (`tests/gateway/test_<platform>.py`)
|
||||||
|
|
||||||
|
Recommended test coverage:
|
||||||
|
|
||||||
|
- Platform enum exists with correct value
|
||||||
|
- Config loading from env vars via `_apply_env_overrides`
|
||||||
|
- Adapter init (config parsing, allowlist handling, default values)
|
||||||
|
- Helper functions (redaction, parsing, file type detection)
|
||||||
|
- Session source round-trip (to_dict → from_dict)
|
||||||
|
- Authorization integration (platform in allowlist maps)
|
||||||
|
- Send message tool routing (platform in platform_map)
|
||||||
|
|
||||||
|
Optional but valuable:
|
||||||
|
- Async tests for message handling flow (mock the platform API)
|
||||||
|
- SSE/WebSocket reconnection logic
|
||||||
|
- Attachment processing
|
||||||
|
- Group message filtering
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Quick Verification
|
||||||
|
|
||||||
|
After implementing everything, verify with:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# All tests pass
|
||||||
|
python -m pytest tests/ -q
|
||||||
|
|
||||||
|
# Grep for your platform name to find any missed integration points
|
||||||
|
grep -r "telegram\|discord\|whatsapp\|slack" gateway/ tools/ agent/ cron/ hermes_cli/ toolsets.py \
|
||||||
|
--include="*.py" -l | sort -u
|
||||||
|
# Check each file in the output — if it mentions other platforms but not yours, you missed it
|
||||||
|
```
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
"""
|
||||||
|
Platform adapters for messaging integrations.
|
||||||
|
|
||||||
|
Each adapter handles:
|
||||||
|
- Receiving messages from a platform
|
||||||
|
- Sending messages/responses back
|
||||||
|
- Platform-specific authentication
|
||||||
|
- Message formatting and media handling
|
||||||
|
"""
|
||||||
|
|
||||||
|
from .base import BasePlatformAdapter, MessageEvent, SendResult
|
||||||
|
from .qqbot import QQAdapter
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"BasePlatformAdapter",
|
||||||
|
"MessageEvent",
|
||||||
|
"SendResult",
|
||||||
|
"QQAdapter",
|
||||||
|
]
|
||||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,897 @@
|
|||||||
|
"""BlueBubbles iMessage platform adapter.
|
||||||
|
|
||||||
|
Uses the local BlueBubbles macOS server for outbound REST sends and inbound
|
||||||
|
webhooks. Supports text messaging, media attachments (images, voice, video,
|
||||||
|
documents), tapback reactions, typing indicators, and read receipts.
|
||||||
|
|
||||||
|
Architecture based on PR #5869 (benjaminsehl) with inbound attachment
|
||||||
|
downloading from PR #4588 (YuhangLin).
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
from urllib.parse import quote
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
from gateway.platforms.base import (
|
||||||
|
BasePlatformAdapter,
|
||||||
|
MessageEvent,
|
||||||
|
MessageType,
|
||||||
|
SendResult,
|
||||||
|
cache_image_from_bytes,
|
||||||
|
cache_audio_from_bytes,
|
||||||
|
cache_document_from_bytes,
|
||||||
|
)
|
||||||
|
from gateway.platforms.helpers import strip_markdown
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Constants
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
DEFAULT_WEBHOOK_HOST = "127.0.0.1"
|
||||||
|
DEFAULT_WEBHOOK_PORT = 8645
|
||||||
|
DEFAULT_WEBHOOK_PATH = "/bluebubbles-webhook"
|
||||||
|
MAX_TEXT_LENGTH = 4000
|
||||||
|
|
||||||
|
# Tapback reaction codes (BlueBubbles associatedMessageType values)
|
||||||
|
_TAPBACK_ADDED = {
|
||||||
|
2000: "love", 2001: "like", 2002: "dislike",
|
||||||
|
2003: "laugh", 2004: "emphasize", 2005: "question",
|
||||||
|
}
|
||||||
|
_TAPBACK_REMOVED = {
|
||||||
|
3000: "love", 3001: "like", 3002: "dislike",
|
||||||
|
3003: "laugh", 3004: "emphasize", 3005: "question",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Webhook event types that carry user messages
|
||||||
|
_MESSAGE_EVENTS = {"new-message", "message", "updated-message"}
|
||||||
|
|
||||||
|
# Log redaction patterns
|
||||||
|
_PHONE_RE = re.compile(r"\+?\d{7,15}")
|
||||||
|
_EMAIL_RE = re.compile(r"[\w.+-]+@[\w-]+\.[\w.]+")
|
||||||
|
|
||||||
|
|
||||||
|
def _redact(text: str) -> str:
|
||||||
|
"""Redact phone numbers and emails from log output."""
|
||||||
|
text = _PHONE_RE.sub("[REDACTED]", text)
|
||||||
|
text = _EMAIL_RE.sub("[REDACTED]", text)
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def check_bluebubbles_requirements() -> bool:
|
||||||
|
try:
|
||||||
|
import aiohttp # noqa: F401
|
||||||
|
import httpx as _httpx # noqa: F401
|
||||||
|
except ImportError:
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_server_url(raw: str) -> str:
|
||||||
|
value = (raw or "").strip()
|
||||||
|
if not value:
|
||||||
|
return ""
|
||||||
|
if not re.match(r"^https?://", value, flags=re.I):
|
||||||
|
value = f"http://{value}"
|
||||||
|
return value.rstrip("/")
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Adapter
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class BlueBubblesAdapter(BasePlatformAdapter):
|
||||||
|
platform = Platform.BLUEBUBBLES
|
||||||
|
MAX_MESSAGE_LENGTH = MAX_TEXT_LENGTH
|
||||||
|
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.BLUEBUBBLES)
|
||||||
|
extra = config.extra or {}
|
||||||
|
self.server_url = _normalize_server_url(
|
||||||
|
extra.get("server_url") or os.getenv("BLUEBUBBLES_SERVER_URL", "")
|
||||||
|
)
|
||||||
|
self.password = extra.get("password") or os.getenv("BLUEBUBBLES_PASSWORD", "")
|
||||||
|
self.webhook_host = (
|
||||||
|
extra.get("webhook_host")
|
||||||
|
or os.getenv("BLUEBUBBLES_WEBHOOK_HOST", DEFAULT_WEBHOOK_HOST)
|
||||||
|
)
|
||||||
|
self.webhook_port = int(
|
||||||
|
extra.get("webhook_port")
|
||||||
|
or os.getenv("BLUEBUBBLES_WEBHOOK_PORT", str(DEFAULT_WEBHOOK_PORT))
|
||||||
|
)
|
||||||
|
self.webhook_path = (
|
||||||
|
extra.get("webhook_path")
|
||||||
|
or os.getenv("BLUEBUBBLES_WEBHOOK_PATH", DEFAULT_WEBHOOK_PATH)
|
||||||
|
)
|
||||||
|
if not str(self.webhook_path).startswith("/"):
|
||||||
|
self.webhook_path = f"/{self.webhook_path}"
|
||||||
|
self.send_read_receipts = bool(extra.get("send_read_receipts", True))
|
||||||
|
self.client: Optional[httpx.AsyncClient] = None
|
||||||
|
self._runner = None
|
||||||
|
self._private_api_enabled: Optional[bool] = None
|
||||||
|
self._helper_connected: bool = False
|
||||||
|
self._guid_cache: Dict[str, str] = {}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# API helpers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _api_url(self, path: str) -> str:
|
||||||
|
sep = "&" if "?" in path else "?"
|
||||||
|
return f"{self.server_url}{path}{sep}password={quote(self.password, safe='')}"
|
||||||
|
|
||||||
|
async def _api_get(self, path: str) -> Dict[str, Any]:
|
||||||
|
assert self.client is not None
|
||||||
|
res = await self.client.get(self._api_url(path))
|
||||||
|
res.raise_for_status()
|
||||||
|
return res.json()
|
||||||
|
|
||||||
|
async def _api_post(self, path: str, payload: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
assert self.client is not None
|
||||||
|
res = await self.client.post(self._api_url(path), json=payload)
|
||||||
|
res.raise_for_status()
|
||||||
|
return res.json()
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Lifecycle
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
if not self.server_url or not self.password:
|
||||||
|
logger.error(
|
||||||
|
"[bluebubbles] BLUEBUBBLES_SERVER_URL and BLUEBUBBLES_PASSWORD are required"
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
self.client = httpx.AsyncClient(timeout=30.0)
|
||||||
|
try:
|
||||||
|
await self._api_get("/api/v1/ping")
|
||||||
|
info = await self._api_get("/api/v1/server/info")
|
||||||
|
server_data = (info or {}).get("data", {})
|
||||||
|
self._private_api_enabled = bool(server_data.get("private_api"))
|
||||||
|
self._helper_connected = bool(server_data.get("helper_connected"))
|
||||||
|
logger.info(
|
||||||
|
"[bluebubbles] connected to %s (private_api=%s, helper=%s)",
|
||||||
|
self.server_url,
|
||||||
|
self._private_api_enabled,
|
||||||
|
self._helper_connected,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(
|
||||||
|
"[bluebubbles] cannot reach server at %s: %s", self.server_url, exc
|
||||||
|
)
|
||||||
|
if self.client:
|
||||||
|
await self.client.aclose()
|
||||||
|
self.client = None
|
||||||
|
return False
|
||||||
|
|
||||||
|
app = web.Application()
|
||||||
|
app.router.add_get("/health", lambda _: web.Response(text="ok"))
|
||||||
|
app.router.add_post(self.webhook_path, self._handle_webhook)
|
||||||
|
self._runner = web.AppRunner(app)
|
||||||
|
await self._runner.setup()
|
||||||
|
site = web.TCPSite(self._runner, self.webhook_host, self.webhook_port)
|
||||||
|
await site.start()
|
||||||
|
self._mark_connected()
|
||||||
|
logger.info(
|
||||||
|
"[bluebubbles] webhook listening on http://%s:%s%s",
|
||||||
|
self.webhook_host,
|
||||||
|
self.webhook_port,
|
||||||
|
self.webhook_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Register webhook with BlueBubbles server
|
||||||
|
# This is required for the server to know where to send events
|
||||||
|
await self._register_webhook()
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
# Unregister webhook before cleaning up
|
||||||
|
await self._unregister_webhook()
|
||||||
|
|
||||||
|
if self.client:
|
||||||
|
await self.client.aclose()
|
||||||
|
self.client = None
|
||||||
|
if self._runner:
|
||||||
|
await self._runner.cleanup()
|
||||||
|
self._runner = None
|
||||||
|
self._mark_disconnected()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def _webhook_url(self) -> str:
|
||||||
|
"""Compute the external webhook URL for BlueBubbles registration."""
|
||||||
|
host = self.webhook_host
|
||||||
|
if host in ("0.0.0.0", "127.0.0.1", "localhost", "::"):
|
||||||
|
host = "localhost"
|
||||||
|
return f"http://{host}:{self.webhook_port}{self.webhook_path}"
|
||||||
|
|
||||||
|
async def _find_registered_webhooks(self, url: str) -> list:
|
||||||
|
"""Return list of BB webhook entries matching *url*."""
|
||||||
|
try:
|
||||||
|
res = await self._api_get("/api/v1/webhook")
|
||||||
|
data = res.get("data")
|
||||||
|
if isinstance(data, list):
|
||||||
|
return [wh for wh in data if wh.get("url") == url]
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return []
|
||||||
|
|
||||||
|
async def _register_webhook(self) -> bool:
|
||||||
|
"""Register this webhook URL with the BlueBubbles server.
|
||||||
|
|
||||||
|
BlueBubbles requires webhooks to be registered via API before
|
||||||
|
it will send events. Checks for an existing registration first
|
||||||
|
to avoid duplicates (e.g. after a crash without clean shutdown).
|
||||||
|
"""
|
||||||
|
if not self.client:
|
||||||
|
return False
|
||||||
|
|
||||||
|
webhook_url = self._webhook_url
|
||||||
|
|
||||||
|
# Crash resilience — reuse an existing registration if present
|
||||||
|
existing = await self._find_registered_webhooks(webhook_url)
|
||||||
|
if existing:
|
||||||
|
logger.info(
|
||||||
|
"[bluebubbles] webhook already registered: %s", webhook_url
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"url": webhook_url,
|
||||||
|
"events": ["new-message", "updated-message", "message"],
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
res = await self._api_post("/api/v1/webhook", payload)
|
||||||
|
status = res.get("status", 0)
|
||||||
|
if 200 <= status < 300:
|
||||||
|
logger.info(
|
||||||
|
"[bluebubbles] webhook registered with server: %s",
|
||||||
|
webhook_url,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"[bluebubbles] webhook registration returned status %s: %s",
|
||||||
|
status,
|
||||||
|
res.get("message"),
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"[bluebubbles] failed to register webhook with server: %s",
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _unregister_webhook(self) -> bool:
|
||||||
|
"""Unregister this webhook URL from the BlueBubbles server.
|
||||||
|
|
||||||
|
Removes *all* matching registrations to clean up any duplicates
|
||||||
|
left by prior crashes.
|
||||||
|
"""
|
||||||
|
if not self.client:
|
||||||
|
return False
|
||||||
|
|
||||||
|
webhook_url = self._webhook_url
|
||||||
|
removed = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
for wh in await self._find_registered_webhooks(webhook_url):
|
||||||
|
wh_id = wh.get("id")
|
||||||
|
if wh_id:
|
||||||
|
res = await self.client.delete(
|
||||||
|
self._api_url(f"/api/v1/webhook/{wh_id}")
|
||||||
|
)
|
||||||
|
res.raise_for_status()
|
||||||
|
removed = True
|
||||||
|
if removed:
|
||||||
|
logger.info(
|
||||||
|
"[bluebubbles] webhook unregistered: %s", webhook_url
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug(
|
||||||
|
"[bluebubbles] failed to unregister webhook (non-critical): %s",
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return removed
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Chat GUID resolution
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _resolve_chat_guid(self, target: str) -> Optional[str]:
|
||||||
|
"""Resolve an email/phone to a BlueBubbles chat GUID.
|
||||||
|
|
||||||
|
If *target* already contains a semicolon (raw GUID format like
|
||||||
|
``iMessage;-;user@example.com``), it is returned as-is. Otherwise
|
||||||
|
the adapter queries the BlueBubbles chat list and matches on
|
||||||
|
``chatIdentifier`` or participant address.
|
||||||
|
"""
|
||||||
|
target = (target or "").strip()
|
||||||
|
if not target:
|
||||||
|
return None
|
||||||
|
# Already a raw GUID
|
||||||
|
if ";" in target:
|
||||||
|
return target
|
||||||
|
if target in self._guid_cache:
|
||||||
|
return self._guid_cache[target]
|
||||||
|
try:
|
||||||
|
payload = await self._api_post(
|
||||||
|
"/api/v1/chat/query",
|
||||||
|
{"limit": 100, "offset": 0, "with": ["participants"]},
|
||||||
|
)
|
||||||
|
for chat in payload.get("data", []) or []:
|
||||||
|
guid = chat.get("guid") or chat.get("chatGuid")
|
||||||
|
identifier = chat.get("chatIdentifier") or chat.get("identifier")
|
||||||
|
if identifier == target:
|
||||||
|
if guid:
|
||||||
|
self._guid_cache[target] = guid
|
||||||
|
return guid
|
||||||
|
for part in chat.get("participants", []) or []:
|
||||||
|
if (part.get("address") or "").strip() == target and guid:
|
||||||
|
self._guid_cache[target] = guid
|
||||||
|
return guid
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def _create_chat_for_handle(
|
||||||
|
self, address: str, message: str
|
||||||
|
) -> SendResult:
|
||||||
|
"""Create a new chat by sending the first message to *address*."""
|
||||||
|
payload = {
|
||||||
|
"addresses": [address],
|
||||||
|
"message": message,
|
||||||
|
"tempGuid": f"temp-{datetime.utcnow().timestamp()}",
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
res = await self._api_post("/api/v1/chat/new", payload)
|
||||||
|
data = res.get("data") or {}
|
||||||
|
msg_id = data.get("guid") or data.get("messageGuid") or "ok"
|
||||||
|
return SendResult(success=True, message_id=str(msg_id), raw_response=res)
|
||||||
|
except Exception as exc:
|
||||||
|
return SendResult(success=False, error=str(exc))
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Text sending
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
text = strip_markdown(content or "")
|
||||||
|
if not text:
|
||||||
|
return SendResult(success=False, error="BlueBubbles send requires text")
|
||||||
|
chunks = self.truncate_message(text, max_length=self.MAX_MESSAGE_LENGTH)
|
||||||
|
last = SendResult(success=True)
|
||||||
|
for chunk in chunks:
|
||||||
|
guid = await self._resolve_chat_guid(chat_id)
|
||||||
|
if not guid:
|
||||||
|
# If the target looks like an address, try creating a new chat
|
||||||
|
if self._private_api_enabled and (
|
||||||
|
"@" in chat_id or re.match(r"^\+\d+", chat_id)
|
||||||
|
):
|
||||||
|
return await self._create_chat_for_handle(chat_id, chunk)
|
||||||
|
return SendResult(
|
||||||
|
success=False,
|
||||||
|
error=f"BlueBubbles chat not found for target: {chat_id}",
|
||||||
|
)
|
||||||
|
payload: Dict[str, Any] = {
|
||||||
|
"chatGuid": guid,
|
||||||
|
"tempGuid": f"temp-{datetime.utcnow().timestamp()}",
|
||||||
|
"message": chunk,
|
||||||
|
}
|
||||||
|
if reply_to and self._private_api_enabled and self._helper_connected:
|
||||||
|
payload["method"] = "private-api"
|
||||||
|
payload["selectedMessageGuid"] = reply_to
|
||||||
|
payload["partIndex"] = 0
|
||||||
|
try:
|
||||||
|
res = await self._api_post("/api/v1/message/text", payload)
|
||||||
|
data = res.get("data") or {}
|
||||||
|
msg_id = data.get("guid") or data.get("messageGuid") or "ok"
|
||||||
|
last = SendResult(
|
||||||
|
success=True, message_id=str(msg_id), raw_response=res
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
return SendResult(success=False, error=str(exc))
|
||||||
|
return last
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Media sending (outbound)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _send_attachment(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
file_path: str,
|
||||||
|
filename: Optional[str] = None,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
is_audio_message: bool = False,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a file attachment via BlueBubbles multipart upload."""
|
||||||
|
if not self.client:
|
||||||
|
return SendResult(success=False, error="Not connected")
|
||||||
|
if not os.path.isfile(file_path):
|
||||||
|
return SendResult(success=False, error=f"File not found: {file_path}")
|
||||||
|
|
||||||
|
guid = await self._resolve_chat_guid(chat_id)
|
||||||
|
if not guid:
|
||||||
|
return SendResult(success=False, error=f"Chat not found: {chat_id}")
|
||||||
|
|
||||||
|
fname = filename or os.path.basename(file_path)
|
||||||
|
try:
|
||||||
|
with open(file_path, "rb") as f:
|
||||||
|
files = {"attachment": (fname, f, "application/octet-stream")}
|
||||||
|
data: Dict[str, str] = {
|
||||||
|
"chatGuid": guid,
|
||||||
|
"name": fname,
|
||||||
|
"tempGuid": uuid.uuid4().hex,
|
||||||
|
}
|
||||||
|
if is_audio_message:
|
||||||
|
data["isAudioMessage"] = "true"
|
||||||
|
res = await self.client.post(
|
||||||
|
self._api_url("/api/v1/message/attachment"),
|
||||||
|
files=files,
|
||||||
|
data=data,
|
||||||
|
timeout=120,
|
||||||
|
)
|
||||||
|
res.raise_for_status()
|
||||||
|
result = res.json()
|
||||||
|
|
||||||
|
if caption:
|
||||||
|
await self.send(chat_id, caption)
|
||||||
|
|
||||||
|
if result.get("status") == 200:
|
||||||
|
rdata = result.get("data") or {}
|
||||||
|
msg_id = rdata.get("guid") if isinstance(rdata, dict) else None
|
||||||
|
return SendResult(
|
||||||
|
success=True, message_id=msg_id, raw_response=result
|
||||||
|
)
|
||||||
|
return SendResult(
|
||||||
|
success=False,
|
||||||
|
error=result.get("message", "Attachment upload failed"),
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
async def send_image(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
image_url: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
try:
|
||||||
|
from gateway.platforms.base import cache_image_from_url
|
||||||
|
|
||||||
|
local_path = await cache_image_from_url(image_url)
|
||||||
|
return await self._send_attachment(chat_id, local_path, caption=caption)
|
||||||
|
except Exception:
|
||||||
|
return await super().send_image(chat_id, image_url, caption, reply_to)
|
||||||
|
|
||||||
|
async def send_image_file(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
image_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
return await self._send_attachment(chat_id, image_path, caption=caption)
|
||||||
|
|
||||||
|
async def send_voice(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
audio_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
return await self._send_attachment(
|
||||||
|
chat_id, audio_path, caption=caption, is_audio_message=True
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_video(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
video_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
return await self._send_attachment(chat_id, video_path, caption=caption)
|
||||||
|
|
||||||
|
async def send_document(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
file_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
file_name: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
return await self._send_attachment(
|
||||||
|
chat_id, file_path, filename=file_name, caption=caption
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_animation(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
animation_url: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
return await self.send_image(
|
||||||
|
chat_id, animation_url, caption, reply_to, metadata
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Typing indicators
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def send_typing(self, chat_id: str, metadata=None) -> None:
|
||||||
|
if not self._private_api_enabled or not self._helper_connected or not self.client:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
guid = await self._resolve_chat_guid(chat_id)
|
||||||
|
if guid:
|
||||||
|
encoded = quote(guid, safe="")
|
||||||
|
await self.client.post(
|
||||||
|
self._api_url(f"/api/v1/chat/{encoded}/typing"), timeout=5
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop_typing(self, chat_id: str) -> None:
|
||||||
|
if not self._private_api_enabled or not self._helper_connected or not self.client:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
guid = await self._resolve_chat_guid(chat_id)
|
||||||
|
if guid:
|
||||||
|
encoded = quote(guid, safe="")
|
||||||
|
await self.client.delete(
|
||||||
|
self._api_url(f"/api/v1/chat/{encoded}/typing"), timeout=5
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Read receipts
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def mark_read(self, chat_id: str) -> bool:
|
||||||
|
if not self._private_api_enabled or not self._helper_connected or not self.client:
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
guid = await self._resolve_chat_guid(chat_id)
|
||||||
|
if guid:
|
||||||
|
encoded = quote(guid, safe="")
|
||||||
|
await self.client.post(
|
||||||
|
self._api_url(f"/api/v1/chat/{encoded}/read"), timeout=5
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return False
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Tapback reactions
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Chat info
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
is_group = ";+;" in (chat_id or "")
|
||||||
|
info: Dict[str, Any] = {
|
||||||
|
"name": chat_id,
|
||||||
|
"type": "group" if is_group else "dm",
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
guid = await self._resolve_chat_guid(chat_id)
|
||||||
|
if guid:
|
||||||
|
encoded = quote(guid, safe="")
|
||||||
|
res = await self._api_get(
|
||||||
|
f"/api/v1/chat/{encoded}?with=participants"
|
||||||
|
)
|
||||||
|
data = (res or {}).get("data", {})
|
||||||
|
display_name = (
|
||||||
|
data.get("displayName")
|
||||||
|
or data.get("chatIdentifier")
|
||||||
|
or chat_id
|
||||||
|
)
|
||||||
|
participants = []
|
||||||
|
for p in data.get("participants", []) or []:
|
||||||
|
addr = (p.get("address") or "").strip()
|
||||||
|
if addr:
|
||||||
|
participants.append(addr)
|
||||||
|
info["name"] = display_name
|
||||||
|
if participants:
|
||||||
|
info["participants"] = participants
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return info
|
||||||
|
|
||||||
|
def format_message(self, content: str) -> str:
|
||||||
|
return strip_markdown(content)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Inbound attachment downloading (from #4588)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _download_attachment(
|
||||||
|
self, att_guid: str, att_meta: Dict[str, Any]
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Download an attachment from BlueBubbles and cache it locally.
|
||||||
|
|
||||||
|
Returns the local file path on success, None on failure.
|
||||||
|
"""
|
||||||
|
if not self.client:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
encoded = quote(att_guid, safe="")
|
||||||
|
resp = await self.client.get(
|
||||||
|
self._api_url(f"/api/v1/attachment/{encoded}/download"),
|
||||||
|
timeout=60,
|
||||||
|
follow_redirects=True,
|
||||||
|
)
|
||||||
|
resp.raise_for_status()
|
||||||
|
data = resp.content
|
||||||
|
|
||||||
|
mime = (att_meta.get("mimeType") or "").lower()
|
||||||
|
transfer_name = att_meta.get("transferName", "")
|
||||||
|
|
||||||
|
if mime.startswith("image/"):
|
||||||
|
ext_map = {
|
||||||
|
"image/jpeg": ".jpg",
|
||||||
|
"image/png": ".png",
|
||||||
|
"image/gif": ".gif",
|
||||||
|
"image/webp": ".webp",
|
||||||
|
"image/heic": ".jpg",
|
||||||
|
"image/heif": ".jpg",
|
||||||
|
"image/tiff": ".jpg",
|
||||||
|
}
|
||||||
|
ext = ext_map.get(mime, ".jpg")
|
||||||
|
return cache_image_from_bytes(data, ext)
|
||||||
|
|
||||||
|
if mime.startswith("audio/"):
|
||||||
|
ext_map = {
|
||||||
|
"audio/mp3": ".mp3",
|
||||||
|
"audio/mpeg": ".mp3",
|
||||||
|
"audio/ogg": ".ogg",
|
||||||
|
"audio/wav": ".wav",
|
||||||
|
"audio/x-caf": ".mp3",
|
||||||
|
"audio/mp4": ".m4a",
|
||||||
|
"audio/aac": ".m4a",
|
||||||
|
}
|
||||||
|
ext = ext_map.get(mime, ".mp3")
|
||||||
|
return cache_audio_from_bytes(data, ext)
|
||||||
|
|
||||||
|
# Videos, documents, and everything else
|
||||||
|
filename = transfer_name or f"file_{uuid.uuid4().hex[:8]}"
|
||||||
|
return cache_document_from_bytes(data, filename)
|
||||||
|
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"[bluebubbles] failed to download attachment %s: %s",
|
||||||
|
_redact(att_guid),
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Webhook handling
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _extract_payload_record(
|
||||||
|
self, payload: Dict[str, Any]
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
data = payload.get("data")
|
||||||
|
if isinstance(data, dict):
|
||||||
|
return data
|
||||||
|
if isinstance(data, list):
|
||||||
|
for item in data:
|
||||||
|
if isinstance(item, dict):
|
||||||
|
return item
|
||||||
|
if isinstance(payload.get("message"), dict):
|
||||||
|
return payload.get("message")
|
||||||
|
return payload if isinstance(payload, dict) else None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _value(*candidates: Any) -> Optional[str]:
|
||||||
|
for candidate in candidates:
|
||||||
|
if isinstance(candidate, str) and candidate.strip():
|
||||||
|
return candidate.strip()
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def _handle_webhook(self, request):
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
token = (
|
||||||
|
request.query.get("password")
|
||||||
|
or request.query.get("guid")
|
||||||
|
or request.headers.get("x-password")
|
||||||
|
or request.headers.get("x-guid")
|
||||||
|
or request.headers.get("x-bluebubbles-guid")
|
||||||
|
)
|
||||||
|
if token != self.password:
|
||||||
|
return web.json_response({"error": "unauthorized"}, status=401)
|
||||||
|
try:
|
||||||
|
raw = await request.read()
|
||||||
|
body = raw.decode("utf-8", errors="replace")
|
||||||
|
try:
|
||||||
|
payload = json.loads(body)
|
||||||
|
except Exception:
|
||||||
|
from urllib.parse import parse_qs
|
||||||
|
|
||||||
|
form = parse_qs(body)
|
||||||
|
payload_str = (
|
||||||
|
form.get("payload")
|
||||||
|
or form.get("data")
|
||||||
|
or form.get("message")
|
||||||
|
or [""]
|
||||||
|
)[0]
|
||||||
|
payload = json.loads(payload_str) if payload_str else {}
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("[bluebubbles] webhook parse error: %s", exc)
|
||||||
|
return web.json_response({"error": "invalid payload"}, status=400)
|
||||||
|
|
||||||
|
event_type = self._value(payload.get("type"), payload.get("event")) or ""
|
||||||
|
# Only process message events; silently acknowledge everything else
|
||||||
|
if event_type and event_type not in _MESSAGE_EVENTS:
|
||||||
|
return web.Response(text="ok")
|
||||||
|
|
||||||
|
record = self._extract_payload_record(payload) or {}
|
||||||
|
is_from_me = bool(
|
||||||
|
record.get("isFromMe")
|
||||||
|
or record.get("fromMe")
|
||||||
|
or record.get("is_from_me")
|
||||||
|
)
|
||||||
|
if is_from_me:
|
||||||
|
return web.Response(text="ok")
|
||||||
|
|
||||||
|
# Skip tapback reactions delivered as messages
|
||||||
|
assoc_type = record.get("associatedMessageType")
|
||||||
|
if isinstance(assoc_type, int) and assoc_type in {
|
||||||
|
**_TAPBACK_ADDED,
|
||||||
|
**_TAPBACK_REMOVED,
|
||||||
|
}:
|
||||||
|
return web.Response(text="ok")
|
||||||
|
|
||||||
|
text = (
|
||||||
|
self._value(
|
||||||
|
record.get("text"), record.get("message"), record.get("body")
|
||||||
|
)
|
||||||
|
or ""
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- Inbound attachment handling ---
|
||||||
|
attachments = record.get("attachments") or []
|
||||||
|
media_urls: List[str] = []
|
||||||
|
media_types: List[str] = []
|
||||||
|
msg_type = MessageType.TEXT
|
||||||
|
|
||||||
|
for att in attachments:
|
||||||
|
att_guid = att.get("guid", "")
|
||||||
|
if not att_guid:
|
||||||
|
continue
|
||||||
|
cached = await self._download_attachment(att_guid, att)
|
||||||
|
if cached:
|
||||||
|
mime = (att.get("mimeType") or "").lower()
|
||||||
|
media_urls.append(cached)
|
||||||
|
media_types.append(mime)
|
||||||
|
if mime.startswith("image/"):
|
||||||
|
msg_type = MessageType.PHOTO
|
||||||
|
elif mime.startswith("audio/") or (att.get("uti") or "").endswith(
|
||||||
|
"caf"
|
||||||
|
):
|
||||||
|
msg_type = MessageType.VOICE
|
||||||
|
elif mime.startswith("video/"):
|
||||||
|
msg_type = MessageType.VIDEO
|
||||||
|
else:
|
||||||
|
msg_type = MessageType.DOCUMENT
|
||||||
|
|
||||||
|
# With multiple attachments, prefer PHOTO if any images present
|
||||||
|
if len(media_urls) > 1:
|
||||||
|
mime_prefixes = {(m or "").split("/")[0] for m in media_types}
|
||||||
|
if "image" in mime_prefixes:
|
||||||
|
msg_type = MessageType.PHOTO
|
||||||
|
|
||||||
|
if not text and media_urls:
|
||||||
|
text = "(attachment)"
|
||||||
|
# --- End attachment handling ---
|
||||||
|
|
||||||
|
chat_guid = self._value(
|
||||||
|
record.get("chatGuid"),
|
||||||
|
payload.get("chatGuid"),
|
||||||
|
record.get("chat_guid"),
|
||||||
|
payload.get("chat_guid"),
|
||||||
|
payload.get("guid"),
|
||||||
|
)
|
||||||
|
chat_identifier = self._value(
|
||||||
|
record.get("chatIdentifier"),
|
||||||
|
record.get("identifier"),
|
||||||
|
payload.get("chatIdentifier"),
|
||||||
|
payload.get("identifier"),
|
||||||
|
)
|
||||||
|
sender = (
|
||||||
|
self._value(
|
||||||
|
record.get("handle", {}).get("address")
|
||||||
|
if isinstance(record.get("handle"), dict)
|
||||||
|
else None,
|
||||||
|
record.get("sender"),
|
||||||
|
record.get("from"),
|
||||||
|
record.get("address"),
|
||||||
|
)
|
||||||
|
or chat_identifier
|
||||||
|
or chat_guid
|
||||||
|
)
|
||||||
|
if not (chat_guid or chat_identifier) and sender:
|
||||||
|
chat_identifier = sender
|
||||||
|
if not sender or not (chat_guid or chat_identifier) or not text:
|
||||||
|
return web.json_response({"error": "missing message fields"}, status=400)
|
||||||
|
|
||||||
|
session_chat_id = chat_guid or chat_identifier
|
||||||
|
is_group = bool(record.get("isGroup")) or (";+;" in (chat_guid or ""))
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id=session_chat_id,
|
||||||
|
chat_name=chat_identifier or sender,
|
||||||
|
chat_type="group" if is_group else "dm",
|
||||||
|
user_id=sender,
|
||||||
|
user_name=sender,
|
||||||
|
chat_id_alt=chat_identifier,
|
||||||
|
)
|
||||||
|
event = MessageEvent(
|
||||||
|
text=text,
|
||||||
|
message_type=msg_type,
|
||||||
|
source=source,
|
||||||
|
raw_message=payload,
|
||||||
|
message_id=self._value(
|
||||||
|
record.get("guid"),
|
||||||
|
record.get("messageGuid"),
|
||||||
|
record.get("id"),
|
||||||
|
),
|
||||||
|
reply_to_message_id=self._value(
|
||||||
|
record.get("threadOriginatorGuid"),
|
||||||
|
record.get("associatedMessageGuid"),
|
||||||
|
),
|
||||||
|
media_urls=media_urls,
|
||||||
|
media_types=media_types,
|
||||||
|
)
|
||||||
|
task = asyncio.create_task(self.handle_message(event))
|
||||||
|
self._background_tasks.add(task)
|
||||||
|
task.add_done_callback(self._background_tasks.discard)
|
||||||
|
|
||||||
|
# Fire-and-forget read receipt
|
||||||
|
if self.send_read_receipts and session_chat_id:
|
||||||
|
asyncio.create_task(self.mark_read(session_chat_id))
|
||||||
|
|
||||||
|
return web.Response(text="ok")
|
||||||
|
|
||||||
@@ -0,0 +1,359 @@
|
|||||||
|
"""
|
||||||
|
MIND OS 3.0 — 实时录音 WS 管线 (Phase B)
|
||||||
|
|
||||||
|
从 V2 asr-proxy.cjs 移植到 Python asyncio + aiohttp。
|
||||||
|
|
||||||
|
端点:GET /mindos-next/ws/record?token={JWT}&chatId={chatId}
|
||||||
|
|
||||||
|
协议(与 V2 客户端兼容):
|
||||||
|
客户端 → 服务端: binary PCM16 帧 / {"type":"stop"}
|
||||||
|
服务端 → 客户端: {"type":"proxy_connected","chatId":"..."}
|
||||||
|
{"type":"partial","text":"..."}
|
||||||
|
{"type":"final","text":"..."}
|
||||||
|
{"type":"speech_started"/"speech_stopped"}
|
||||||
|
{"type":"asr_reconnecting","attempt":N}
|
||||||
|
{"type":"error","message":"..."}
|
||||||
|
|
||||||
|
DashScope Realtime API(V5 协议,与 V2 asr-proxy.cjs 完全相同):
|
||||||
|
wss://dashscope.aliyuncs.com/api-ws/v1/realtime?model=qwen3-asr-flash-realtime
|
||||||
|
Auth: Authorization: bearer {DASHSCOPE_API_KEY}
|
||||||
|
发送: {"type":"input_audio_buffer.append","audio":"<base64>"}
|
||||||
|
停止: {"type":"input_audio_buffer.commit"} + {"type":"session.finish"}
|
||||||
|
接收: session.created / conversation.item.input_audio_transcription.* /
|
||||||
|
input_audio_buffer.speech_* / session.finished / error
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import base64
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
from datetime import datetime
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
logger = logging.getLogger("dashscope_realtime")
|
||||||
|
|
||||||
|
DASHSCOPE_ASR_URL = "wss://dashscope.aliyuncs.com/api-ws/v1/realtime"
|
||||||
|
DASHSCOPE_ASR_MODEL = os.getenv("DASHSCOPE_REALTIME_MODEL", "qwen3-asr-flash-realtime")
|
||||||
|
MAX_RECONNECTS = 3
|
||||||
|
RECONNECT_BASE_MS = 1000
|
||||||
|
|
||||||
|
# ★ 架构修复:这里直接存放 MindOSSSEServer 实例的引用。
|
||||||
|
# mindos_sse.py 在 start() 里导入本模块时调用 register_server(self),
|
||||||
|
# 始终使用同一个对象,彻底绕开 sys.modules 命名空间问题。
|
||||||
|
_sse_server = None
|
||||||
|
|
||||||
|
def register_server(server) -> None:
|
||||||
|
global _sse_server
|
||||||
|
_sse_server = server
|
||||||
|
logger.info("[Realtime] SSE server 已注入: %s", type(server).__name__)
|
||||||
|
|
||||||
|
|
||||||
|
# ─── 公开入口(由 mindos_sse.py 注册为路由) ─────────────────
|
||||||
|
|
||||||
|
async def handleWsRecord(request: web.Request) -> web.WebSocketResponse:
|
||||||
|
"""
|
||||||
|
GET /mindos-next/ws/record?token={JWT}&chatId={chatId}
|
||||||
|
aiohttp WS handler。
|
||||||
|
"""
|
||||||
|
from gateway.platforms.mindos_sse import _verifyTokenAsync, _CORS_HEADERS # type: ignore
|
||||||
|
|
||||||
|
# 1. Auth via ?token= query param
|
||||||
|
token = request.rel_url.query.get("token", "")
|
||||||
|
chatId = request.rel_url.query.get("chatId", "")
|
||||||
|
if not token:
|
||||||
|
raise web.HTTPForbidden()
|
||||||
|
|
||||||
|
user = await _verifyTokenAsync(token)
|
||||||
|
if not user:
|
||||||
|
raise web.HTTPUnauthorized()
|
||||||
|
userId = user.get("sub") or user.get("userId", "")
|
||||||
|
|
||||||
|
if not chatId:
|
||||||
|
chatId = f"chat_{int(time.time() * 1000)}"
|
||||||
|
|
||||||
|
meetingId = request.rel_url.query.get("meetingId") or f"rec_{int(time.time() * 1000)}"
|
||||||
|
source = request.rel_url.query.get("source", "mic") # "mic" | "system"
|
||||||
|
|
||||||
|
logger.info("[Realtime] 客户端连接 userId=%s chatId=%s meeting=%s source=%s", userId, chatId, meetingId, source)
|
||||||
|
|
||||||
|
# 2. 升级到 WS
|
||||||
|
clientWs = web.WebSocketResponse(heartbeat=30)
|
||||||
|
await clientWs.prepare(request)
|
||||||
|
|
||||||
|
# 3. 初始化 MD 文件(wiki/{userId}/raw/)
|
||||||
|
wiki_root = os.getenv("MINDOS_WIKI_DIR", os.path.expanduser("~/.hermes/wiki"))
|
||||||
|
wiki_dir = Path(wiki_root) / userId / "raw"
|
||||||
|
wiki_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
now = datetime.now()
|
||||||
|
suffix = "system" if source == "system" else "mic"
|
||||||
|
md_filename = f"{now.strftime('%Y-%m-%d_%H%M')}_{suffix}_recording.md"
|
||||||
|
md_file = wiki_dir / md_filename
|
||||||
|
rel_path = f"raw/{md_filename}"
|
||||||
|
|
||||||
|
# 写文件头
|
||||||
|
source_label = "系统拾音" if source == "system" else "麦克风"
|
||||||
|
md_file.write_text(
|
||||||
|
f"# 录音转写 {now.strftime('%Y-%m-%d %H:%M')}({source_label})\n\n> 实时录制于 MindOS\n\n",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
total_chars = 0
|
||||||
|
sentences = [] # 收集完整句子
|
||||||
|
reconnects = 0
|
||||||
|
ds_ready = asyncio.Event()
|
||||||
|
stop_event = asyncio.Event()
|
||||||
|
audio_queue: asyncio.Queue[bytes | None] = asyncio.Queue()
|
||||||
|
|
||||||
|
# ─── DashScope WS 协程 ──────────────────────────────────
|
||||||
|
|
||||||
|
async def run_dashscope():
|
||||||
|
nonlocal reconnects, total_chars
|
||||||
|
|
||||||
|
api_key = os.getenv("DASHSCOPE_API_KEY", "")
|
||||||
|
if not api_key:
|
||||||
|
await _send_client(clientWs, {"type": "error", "message": "ASR 未配置 DASHSCOPE_API_KEY"})
|
||||||
|
return
|
||||||
|
|
||||||
|
import aiohttp as _aio
|
||||||
|
url = f"{DASHSCOPE_ASR_URL}?model={DASHSCOPE_ASR_MODEL}"
|
||||||
|
headers = {"Authorization": f"bearer {api_key}"}
|
||||||
|
|
||||||
|
while reconnects <= MAX_RECONNECTS and not stop_event.is_set():
|
||||||
|
try:
|
||||||
|
async with _aio.ClientSession() as session:
|
||||||
|
async with session.ws_connect(url, headers=headers) as dsWs:
|
||||||
|
ds_ready.clear() # 连接成功,但等 session.created 后再就绪
|
||||||
|
reconnects = 0 # 连上就重置
|
||||||
|
logger.info("[Realtime] DashScope 已连接 meeting=%s", meetingId)
|
||||||
|
|
||||||
|
async def send_audio_loop():
|
||||||
|
"""从队列取音频帧发给 DashScope"""
|
||||||
|
while True:
|
||||||
|
chunk = await audio_queue.get()
|
||||||
|
if chunk is None:
|
||||||
|
# stop 信号
|
||||||
|
try:
|
||||||
|
await dsWs.send_str(json.dumps(
|
||||||
|
{"type": "input_audio_buffer.commit"}))
|
||||||
|
await dsWs.send_str(json.dumps(
|
||||||
|
{"type": "session.finish"}))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return
|
||||||
|
if dsWs.closed:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
b64 = base64.b64encode(chunk).decode()
|
||||||
|
await dsWs.send_str(json.dumps(
|
||||||
|
{"type": "input_audio_buffer.append", "audio": b64}))
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[Realtime] send audio err: %s", e)
|
||||||
|
|
||||||
|
# 并发:发音频 + 收结果
|
||||||
|
send_task = asyncio.create_task(send_audio_loop())
|
||||||
|
|
||||||
|
async for msg in dsWs:
|
||||||
|
if msg.type != _aio.WSMsgType.TEXT:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
evt = json.loads(msg.data)
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
|
||||||
|
etype = evt.get("type", "")
|
||||||
|
|
||||||
|
if etype == "session.created":
|
||||||
|
# V2 验证:DashScope V5 Realtime 不需要 session.update
|
||||||
|
# 直接就绪,开始接收音频帧
|
||||||
|
ds_ready.set()
|
||||||
|
await _send_client(clientWs, {
|
||||||
|
"type": "proxy_connected", "chatId": chatId})
|
||||||
|
|
||||||
|
elif etype == "conversation.item.input_audio_transcription.completed":
|
||||||
|
text = (evt.get("transcript") or "").strip()
|
||||||
|
if text:
|
||||||
|
sentences.append(text)
|
||||||
|
total_chars += len(text)
|
||||||
|
elapsed = int(time.time() * 1000 - start_ms) // 1000
|
||||||
|
mm, ss = divmod(elapsed, 60)
|
||||||
|
# 实时写 MD
|
||||||
|
with md_file.open("a", encoding="utf-8") as f:
|
||||||
|
f.write(f"[{mm:02d}:{ss:02d}] {text}\n")
|
||||||
|
await _send_client(clientWs, {"type": "final", "text": text})
|
||||||
|
|
||||||
|
elif etype == "conversation.item.input_audio_transcription.text":
|
||||||
|
text = (evt.get("transcript") or "").strip()
|
||||||
|
if text:
|
||||||
|
await _send_client(clientWs, {"type": "partial", "text": text})
|
||||||
|
|
||||||
|
elif etype in (
|
||||||
|
"input_audio_buffer.speech_started",
|
||||||
|
"input_audio_buffer.speech_stopped",
|
||||||
|
):
|
||||||
|
short = "speech_started" if "started" in etype else "speech_stopped"
|
||||||
|
await _send_client(clientWs, {"type": short})
|
||||||
|
|
||||||
|
elif etype == "session.finished":
|
||||||
|
logger.info("[Realtime] session.finished meeting=%s", meetingId)
|
||||||
|
send_task.cancel()
|
||||||
|
stop_event.set()
|
||||||
|
return
|
||||||
|
|
||||||
|
elif etype == "error":
|
||||||
|
err_msg = evt.get("error", {}).get("message") or str(evt)
|
||||||
|
logger.error("[Realtime] DashScope error: %s", err_msg)
|
||||||
|
await _send_client(clientWs, {"type": "error", "message": err_msg})
|
||||||
|
|
||||||
|
send_task.cancel()
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[Realtime] DashScope 连接异常: %s", e)
|
||||||
|
|
||||||
|
if stop_event.is_set():
|
||||||
|
return
|
||||||
|
|
||||||
|
# 重连
|
||||||
|
reconnects += 1
|
||||||
|
if reconnects > MAX_RECONNECTS:
|
||||||
|
logger.error("[Realtime] 重连次数耗尽 meeting=%s", meetingId)
|
||||||
|
break
|
||||||
|
|
||||||
|
delay = RECONNECT_BASE_MS * (2 ** (reconnects - 1)) / 1000
|
||||||
|
logger.warning("[Realtime] %ss 后重连(%d/%d) meeting=%s",
|
||||||
|
delay, reconnects, MAX_RECONNECTS, meetingId)
|
||||||
|
await _send_client(clientWs, {
|
||||||
|
"type": "asr_reconnecting",
|
||||||
|
"attempt": reconnects,
|
||||||
|
"maxAttempts": MAX_RECONNECTS,
|
||||||
|
})
|
||||||
|
ds_ready.clear()
|
||||||
|
await asyncio.sleep(delay)
|
||||||
|
|
||||||
|
# ─── 主循环:接收客户端消息 ──────────────────────────────
|
||||||
|
|
||||||
|
start_ms = int(time.time() * 1000)
|
||||||
|
ds_task = asyncio.create_task(run_dashscope())
|
||||||
|
|
||||||
|
import aiohttp as _aio2
|
||||||
|
async for msg in clientWs:
|
||||||
|
if msg.type == _aio2.WSMsgType.BINARY:
|
||||||
|
# PCM16 帧 → 入队
|
||||||
|
await ds_ready.wait() # 等 DashScope 握手完成再转发
|
||||||
|
await audio_queue.put(msg.data)
|
||||||
|
|
||||||
|
elif msg.type == _aio2.WSMsgType.TEXT:
|
||||||
|
try:
|
||||||
|
cmd = json.loads(msg.data)
|
||||||
|
if cmd.get("type") == "stop":
|
||||||
|
logger.info("[Realtime] 收到 stop 命令 meeting=%s", meetingId)
|
||||||
|
await audio_queue.put(None) # 通知 send_audio_loop 发 commit+finish
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
elif msg.type in (_aio2.WSMsgType.ERROR, _aio2.WSMsgType.CLOSE):
|
||||||
|
break
|
||||||
|
|
||||||
|
# 客户端断开:确保停止
|
||||||
|
stop_event.set()
|
||||||
|
await audio_queue.put(None)
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(ds_task, timeout=10)
|
||||||
|
except (asyncio.TimeoutError, Exception):
|
||||||
|
pass
|
||||||
|
|
||||||
|
# ─── Finalize:写 MD 结尾 + 写 DB + 推 SSE ──────────────
|
||||||
|
|
||||||
|
# P0-4: 计算实时录音总时长
|
||||||
|
duration_seconds = (time.time() * 1000 - start_ms) / 1000.0
|
||||||
|
|
||||||
|
await _finalize(
|
||||||
|
md_file=md_file,
|
||||||
|
rel_path=rel_path,
|
||||||
|
sentences=sentences,
|
||||||
|
total_chars=total_chars,
|
||||||
|
chatId=chatId,
|
||||||
|
meetingId=meetingId,
|
||||||
|
userId=userId,
|
||||||
|
duration_seconds=duration_seconds,
|
||||||
|
source=source,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info("[Realtime] 会话结束 meeting=%s chars=%d", meetingId, total_chars)
|
||||||
|
return clientWs
|
||||||
|
|
||||||
|
|
||||||
|
# ─── 内部工具函数 ─────────────────────────────────────────────
|
||||||
|
|
||||||
|
async def _send_client(ws: web.WebSocketResponse, data: dict) -> None:
|
||||||
|
"""安全发送 JSON 到客户端"""
|
||||||
|
if ws.closed:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
await ws.send_str(json.dumps(data, ensure_ascii=False))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
async def _finalize(
|
||||||
|
md_file: Path,
|
||||||
|
rel_path: str,
|
||||||
|
sentences: list[str],
|
||||||
|
total_chars: int,
|
||||||
|
chatId: str,
|
||||||
|
meetingId: str,
|
||||||
|
userId: str,
|
||||||
|
duration_seconds: float = 0.0,
|
||||||
|
source: str = "mic",
|
||||||
|
) -> None:
|
||||||
|
"""写 MD 结尾、持久化到 DB、推送 SSE md:appended。
|
||||||
|
|
||||||
|
使用 voice2md_atoms 共享原子层(④ db_persist + ⑤ sse_push)。
|
||||||
|
"""
|
||||||
|
from voice2md_atoms import ( # type: ignore
|
||||||
|
persist_audio_result, push_md_appended, deduct_asr_credits,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 写结尾标记
|
||||||
|
try:
|
||||||
|
with md_file.open("a", encoding="utf-8") as f:
|
||||||
|
f.write("\n---(录音结束)\n")
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[Realtime] 写 MD 结尾失败: %s", e)
|
||||||
|
|
||||||
|
# 读完整内容(用于内嵌推送)
|
||||||
|
md_content = ""
|
||||||
|
try:
|
||||||
|
md_content = md_file.read_text(encoding="utf-8")
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# ④ db_persist(chars=0 也写 DB,确保刷新后可恢复)
|
||||||
|
persist_audio_result(
|
||||||
|
chat_id=chatId, user_id=userId,
|
||||||
|
file_name=md_file.name, md_path=rel_path,
|
||||||
|
chars=total_chars, md_content=md_content,
|
||||||
|
oss_read_url="", # 实时录音无 OSS URL
|
||||||
|
)
|
||||||
|
|
||||||
|
# 积分扣减(实时录音:2 credits/秒)
|
||||||
|
asr_credits = max(1, int(duration_seconds * 2))
|
||||||
|
deduct_asr_credits(
|
||||||
|
user_id=userId, chat_id=chatId, credits=asr_credits,
|
||||||
|
tx_type="asr_realtime", model="qwen3-asr-flash-realtime",
|
||||||
|
seconds=duration_seconds,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ⑤ sse_push
|
||||||
|
# ★ 使用模块级 _sse_server(由 mindos_sse.start() 直接注入)
|
||||||
|
push_md_appended(
|
||||||
|
sse_server=_sse_server, user_id=userId, chat_id=chatId,
|
||||||
|
file=rel_path, chars=total_chars, md_content=md_content,
|
||||||
|
message=f"✅ 录音转写完成({total_chars} 字)",
|
||||||
|
source=source,
|
||||||
|
)
|
||||||
|
|
||||||
@@ -0,0 +1,425 @@
|
|||||||
|
"""
|
||||||
|
⚠️ DEPRECATED(文档解析部分)— 请勿复制或修改以下函数 ⚠️
|
||||||
|
|
||||||
|
以下函数已被 infra.pipelines.anyfile2md 统一管线取代:
|
||||||
|
- _is_vlm_needed() → infra.atoms.text_sniffer
|
||||||
|
- _extract_pdf_vlm() → infra.atoms.page_rasterizer + vlm_ocr + oss_presign
|
||||||
|
- _extract_pdf_to_md() → infra.atoms.text_extractor
|
||||||
|
- _extract_docx() → infra.atoms.text_extractor
|
||||||
|
- _extract_txt() → infra.atoms.text_extractor
|
||||||
|
- _simulate_extract() → 不再需要
|
||||||
|
|
||||||
|
✅ _extract_audio_asr() 是 deepview 专有的音频 ASR 管线,不在废弃范围。
|
||||||
|
|
||||||
|
迁移方式参考 xinzong_materials.py(2026-04-20 已完成迁移):
|
||||||
|
from infra.pipelines.anyfile2md import parseLocal
|
||||||
|
result = await parseLocal(str(raw_path), original_filename)
|
||||||
|
|
||||||
|
新的统一管线位于:
|
||||||
|
mindOSv2/hermes-overlay/infra/pipelines/anyfile2md.py
|
||||||
|
mindOSv2/hermes-overlay/infra/atoms/ (6 个原子操作)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import hashlib
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
import logging
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
from pathlib import Path
|
||||||
|
from hermes_constants import get_hermes_home
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
async def process_uploaded_material_oss(
|
||||||
|
oss_key: str,
|
||||||
|
original_filename: str,
|
||||||
|
context_id: str,
|
||||||
|
push_event_fn,
|
||||||
|
user_id: str,
|
||||||
|
org_id: str = "org_001"
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
1. Downloads native PDF from OSS
|
||||||
|
2. Runs deduplication
|
||||||
|
3. Sniffs doc to decide between VLM or pure text
|
||||||
|
4. Executes the relevant pipeline to produce .md
|
||||||
|
"""
|
||||||
|
from hermes_constants import get_hermes_home
|
||||||
|
# 1. Fetch from OSS
|
||||||
|
oss_ak = os.getenv("ALIYUN_ACCESS_KEY_ID", "")
|
||||||
|
oss_sk = os.getenv("ALIYUN_ACCESS_KEY_SECRET", "")
|
||||||
|
oss_bucket = os.getenv("OSS_BUCKET", "meetings-dev")
|
||||||
|
oss_endpoint = os.getenv("OSS_ENDPOINT", "oss-cn-beijing.aliyuncs.com")
|
||||||
|
|
||||||
|
if not all([oss_ak, oss_sk]):
|
||||||
|
logger.error("[DeepviewMaterials] Missing OSS keys in backend .env. Skipping ingestion.")
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
import oss2
|
||||||
|
auth = oss2.Auth(oss_ak, oss_sk)
|
||||||
|
bucket = oss2.Bucket(auth, oss_endpoint, oss_bucket)
|
||||||
|
|
||||||
|
logger.info(f"[DeepviewMaterials] Over internal network fetching {oss_key}...")
|
||||||
|
result = bucket.get_object(oss_key)
|
||||||
|
file_bytes = result.read()
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"[DeepviewMaterials] OSS Download failed: {e}")
|
||||||
|
return
|
||||||
|
|
||||||
|
# 2. Hashing
|
||||||
|
file_hash = hashlib.md5(file_bytes).hexdigest()
|
||||||
|
|
||||||
|
# 3. 解析 Context 路由(深维专属双上下文 FS-as-Database 管线)
|
||||||
|
storageDir = os.getenv("DEEPVIEW_STORAGE_DIR", os.path.expanduser("~/Downloads/Coding/医生助理智能体/backend/storage"))
|
||||||
|
storage_root = Path(storageDir)
|
||||||
|
user_id_safe = user_id if user_id else "unknown"
|
||||||
|
org_id_safe = org_id if org_id else "org_001"
|
||||||
|
|
||||||
|
user_dir = storage_root / "users" / user_id_safe
|
||||||
|
org_dir = storage_root / "orgs" / org_id_safe
|
||||||
|
platform_dir = storage_root / "platform"
|
||||||
|
|
||||||
|
ext = os.path.splitext(original_filename)[1].lower()
|
||||||
|
|
||||||
|
# 过滤恶意文件
|
||||||
|
if ext == '.url' or ext == '.exe' or ext == '.sh':
|
||||||
|
logger.warning(f"Rejected unsafe file: {original_filename}")
|
||||||
|
return
|
||||||
|
|
||||||
|
# 定义路由目标
|
||||||
|
if context_id.startswith("recording:"):
|
||||||
|
recordingId = context_id.split(":", 1)[1]
|
||||||
|
parts = recordingId.split("/")
|
||||||
|
if len(parts) > 1:
|
||||||
|
# 旧格式(已归档):recording:clientId/asrId
|
||||||
|
clientId = parts[0]
|
||||||
|
asrId = parts[-1]
|
||||||
|
base_dir = user_dir / "clients" / clientId / "history"
|
||||||
|
raw_path = base_dir / f"{asrId}_raw{ext}"
|
||||||
|
md_path = base_dir / f"{asrId}.md"
|
||||||
|
else:
|
||||||
|
# 新格式(Inbox):recording:asrId
|
||||||
|
asrId = parts[0]
|
||||||
|
base_dir = user_dir / "inbox" / asrId
|
||||||
|
raw_path = base_dir / f"asr_raw{ext}"
|
||||||
|
md_path = base_dir / "asr.md"
|
||||||
|
elif context_id.startswith("client:"):
|
||||||
|
clientId = context_id.split(":", 1)[1]
|
||||||
|
base_dir = user_dir / "clients" / clientId
|
||||||
|
# 客户全景上下文:碎片化资料存档(可能是合同、病历单据等)
|
||||||
|
raw_path = base_dir / f"{file_hash}_raw{ext}"
|
||||||
|
md_path = base_dir / f"{file_hash}.md"
|
||||||
|
elif context_id.startswith("wiki:") or context_id == "deepview":
|
||||||
|
# 存入本机构域的 wiki 库
|
||||||
|
base_dir = org_dir / "wiki"
|
||||||
|
raw_path = base_dir / f"{file_hash}_raw{ext}"
|
||||||
|
md_path = base_dir / f"{file_hash}.md"
|
||||||
|
elif context_id.startswith("platform:"):
|
||||||
|
# 预留平台运维口
|
||||||
|
base_dir = platform_dir / "wiki"
|
||||||
|
raw_path = base_dir / f"{file_hash}_raw{ext}"
|
||||||
|
md_path = base_dir / f"{file_hash}.md"
|
||||||
|
elif context_id.startswith("doctor:"):
|
||||||
|
base_dir = user_dir
|
||||||
|
raw_path = base_dir / f"doctor_profile_raw{ext}"
|
||||||
|
md_path = base_dir / "doctor_profile.md"
|
||||||
|
else:
|
||||||
|
# Fallback 容错隔离区
|
||||||
|
base_dir = user_dir / "misc" / context_id.replace(":", "_")
|
||||||
|
raw_path = base_dir / f"{file_hash}_raw{ext}"
|
||||||
|
md_path = base_dir / f"{file_hash}.md"
|
||||||
|
|
||||||
|
base_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
# 幂等性检查:文件是否已经存在
|
||||||
|
if md_path.exists():
|
||||||
|
logger.info(f"[DeepviewMaterials] File {original_filename} already processed for {context_id}.")
|
||||||
|
push_event_fn(user_id, "material:done", {
|
||||||
|
"projectId": context_id,
|
||||||
|
"filename": original_filename,
|
||||||
|
"fileId": file_hash
|
||||||
|
})
|
||||||
|
return
|
||||||
|
|
||||||
|
# 存储原始二进制文件(用于溯源和重试)
|
||||||
|
with open(raw_path, "wb") as f:
|
||||||
|
f.write(file_bytes)
|
||||||
|
|
||||||
|
# Process
|
||||||
|
loop = asyncio.get_event_loop()
|
||||||
|
try:
|
||||||
|
if ext == ".pdf":
|
||||||
|
# Sniffer
|
||||||
|
if _is_vlm_needed(str(raw_path)):
|
||||||
|
logger.info("[DeepviewMaterials] Sniffer: Routing to OSS+VLM Pipeline.")
|
||||||
|
await loop.run_in_executor(None, _extract_pdf_vlm, str(raw_path), str(md_path), bucket, file_hash, original_filename)
|
||||||
|
else:
|
||||||
|
logger.info("[DeepviewMaterials] Sniffer: Routing to Pymupdf4llm text Pipeline.")
|
||||||
|
await loop.run_in_executor(None, _extract_pdf_to_md, str(raw_path), str(md_path), original_filename)
|
||||||
|
elif ext in [".docx", ".doc"]:
|
||||||
|
await loop.run_in_executor(None, _extract_docx, str(raw_path), str(md_path), original_filename)
|
||||||
|
elif ext in [".txt", ".md", ".csv"]:
|
||||||
|
await loop.run_in_executor(None, _extract_txt, str(raw_path), str(md_path), original_filename)
|
||||||
|
elif ext in [".m4a", ".mp3", ".wav", ".webm", ".ogg"]:
|
||||||
|
await loop.run_in_executor(None, _extract_audio_asr, str(raw_path), str(md_path), original_filename, bucket, oss_key)
|
||||||
|
else:
|
||||||
|
await loop.run_in_executor(None, _simulate_extract, str(md_path), original_filename)
|
||||||
|
|
||||||
|
# ★ 持久化到 deepview_materials 表(DB 唯一真相)
|
||||||
|
try:
|
||||||
|
from hermes_state import SessionDB
|
||||||
|
db = SessionDB()
|
||||||
|
def _do(conn):
|
||||||
|
conn.execute(
|
||||||
|
"INSERT OR IGNORE INTO deepview_materials (id, filename, context_id, user_id, source, created_at) "
|
||||||
|
"VALUES (?, ?, ?, ?, 'upload', ?)",
|
||||||
|
(file_hash, original_filename, context_id, user_id, time.time())
|
||||||
|
)
|
||||||
|
db._execute_write(_do)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"[DeepviewMaterials] Failed to persist to DB: {e}")
|
||||||
|
|
||||||
|
# Emit SSE success
|
||||||
|
logger.info(f"[DeepviewMaterials] Successfully ingested {original_filename}")
|
||||||
|
push_event_fn(user_id, "material:done", {
|
||||||
|
"projectId": context_id,
|
||||||
|
"filename": original_filename,
|
||||||
|
"fileId": file_hash
|
||||||
|
})
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"[DeepviewMaterials] Failed to extract {original_filename}: {e}", exc_info=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_vlm_needed(pdf_path: str) -> bool:
|
||||||
|
try:
|
||||||
|
import fitz
|
||||||
|
doc = fitz.open(pdf_path)
|
||||||
|
if len(doc) == 0: return False
|
||||||
|
|
||||||
|
check_pages = min(2, len(doc))
|
||||||
|
total_text_len = 0
|
||||||
|
vlm_vote = 0
|
||||||
|
|
||||||
|
for i in range(check_pages):
|
||||||
|
page = doc[i]
|
||||||
|
rect = page.rect
|
||||||
|
width, height = rect.width, rect.height
|
||||||
|
if width > height and (width / height) > 1.2:
|
||||||
|
vlm_vote += 1
|
||||||
|
text = page.get_text()
|
||||||
|
total_text_len += len(text.strip())
|
||||||
|
|
||||||
|
doc.close()
|
||||||
|
|
||||||
|
# If any page is landscape => PPT => VLM
|
||||||
|
if vlm_vote > 0:
|
||||||
|
return True
|
||||||
|
|
||||||
|
# If extremely low text density => Scanned => VLM
|
||||||
|
if check_pages > 0 and (total_text_len / check_pages) < 100:
|
||||||
|
return True
|
||||||
|
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Sniffer error: {e}")
|
||||||
|
return False # Fallback to pymupdf text
|
||||||
|
|
||||||
|
def _extract_pdf_vlm(pdf_path: str, md_path: str, bucket, file_hash: str, original_filename: str):
|
||||||
|
import fitz
|
||||||
|
import tempfile
|
||||||
|
from openai import OpenAI
|
||||||
|
|
||||||
|
# 统一收拢于服务端的 LiteLLM 代理管线(彻底废弃单独的 DASHSCOPE 授权)
|
||||||
|
litellm_key = os.getenv("GEMINI_API_KEY")
|
||||||
|
litellm_base = os.getenv("GEMINI_BASE_URL", "http://127.0.0.1:4000/v1")
|
||||||
|
vlm_model = os.getenv("DEEPVIEW_MODEL", "gemini3.1pro-vertex")
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||||
|
doc = fitz.open(pdf_path)
|
||||||
|
zoom = 200 / 72
|
||||||
|
mat = fitz.Matrix(zoom, zoom)
|
||||||
|
|
||||||
|
urls = []
|
||||||
|
for i, page in enumerate(doc):
|
||||||
|
pix = page.get_pixmap(matrix=mat)
|
||||||
|
img_path = os.path.join(tmp_dir, f"page_{i + 1:03d}.png")
|
||||||
|
pix.save(img_path)
|
||||||
|
pix = None
|
||||||
|
|
||||||
|
oss_key = f"deepview-assets/{file_hash}/slides/page_{i+1:03d}.png"
|
||||||
|
bucket.put_object_from_file(oss_key, img_path)
|
||||||
|
url = bucket.sign_url('GET', oss_key, 3600*24)
|
||||||
|
urls.append(url)
|
||||||
|
|
||||||
|
doc.close()
|
||||||
|
|
||||||
|
if not litellm_key or not litellm_base:
|
||||||
|
with open(md_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(f"# Document: VLM Extraction Skipped\n\n缺少LiteLLM网关配置 (GEMINI_API_KEY/GEMINI_BASE_URL)。已将 {len(urls)} 页图片传至 OSS。")
|
||||||
|
return
|
||||||
|
|
||||||
|
client = OpenAI(base_url=litellm_base, api_key=litellm_key, timeout=120)
|
||||||
|
|
||||||
|
markdown_blocks = []
|
||||||
|
for i, u in enumerate(urls):
|
||||||
|
prompt = "详细分析这张页面图片。如果是PPT请提炼核心观点、标题和要素。如果是扫描件请保留每一段具体文字。用纯粹的Markdown格式输出,不要使用```markdown包裹。"
|
||||||
|
try:
|
||||||
|
resp = client.chat.completions.create(
|
||||||
|
model=vlm_model,
|
||||||
|
messages=[{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "image_url", "image_url": {"url": u}},
|
||||||
|
{"type": "text", "text": prompt}
|
||||||
|
]
|
||||||
|
}],
|
||||||
|
max_tokens=2048,
|
||||||
|
temperature=0.1
|
||||||
|
)
|
||||||
|
txt = resp.choices[0].message.content.strip()
|
||||||
|
if txt.startswith("```markdown"):
|
||||||
|
txt = txt[11:]
|
||||||
|
txt = txt.strip("`\n ")
|
||||||
|
markdown_blocks.append(f"## 第 {i+1} 页\n\n{txt}\n")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"LiteLLM VLM error page {i+1}: {e}")
|
||||||
|
markdown_blocks.append(f"## 第 {i+1} 页\n*(视觉提取超时或失败)*\n")
|
||||||
|
|
||||||
|
with open(md_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(f"# Document: {original_filename}\n\n")
|
||||||
|
f.write("\n---\n".join(markdown_blocks))
|
||||||
|
|
||||||
|
def _extract_pdf_to_md(pdf_path: str, md_path: str, original_filename: str):
|
||||||
|
try:
|
||||||
|
import pymupdf4llm
|
||||||
|
md_text = pymupdf4llm.to_markdown(pdf_path)
|
||||||
|
with open(md_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(f"# Document: {original_filename}\n\n")
|
||||||
|
f.write(md_text)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"pymupdf4llm error: {e}")
|
||||||
|
_simulate_extract(md_path, os.path.basename(pdf_path))
|
||||||
|
|
||||||
|
def _extract_docx(docx_path: str, md_path: str, original_filename: str):
|
||||||
|
try:
|
||||||
|
import docx
|
||||||
|
doc = docx.Document(docx_path)
|
||||||
|
text = "\n\n".join([p.text for p in doc.paragraphs if p.text.strip()])
|
||||||
|
with open(md_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(f"# Document: {original_filename}\n\n")
|
||||||
|
f.write(text)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"docx error: {e}")
|
||||||
|
_simulate_extract(md_path, os.path.basename(docx_path))
|
||||||
|
|
||||||
|
def _extract_txt(txt_path: str, md_path: str, original_filename: str):
|
||||||
|
try:
|
||||||
|
with open(txt_path, "r", encoding="utf-8") as f:
|
||||||
|
text = f.read()
|
||||||
|
with open(md_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(f"# Document: {original_filename}\n\n")
|
||||||
|
f.write(text)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"txt error: {e}")
|
||||||
|
_simulate_extract(md_path, os.path.basename(txt_path))
|
||||||
|
|
||||||
|
def _extract_audio_asr(audio_path: str, md_path: str, original_filename: str, bucket, oss_key: str):
|
||||||
|
import requests
|
||||||
|
import time
|
||||||
|
try:
|
||||||
|
dashscope_key = os.getenv("DASHSCOPE_API_KEY", "")
|
||||||
|
if not dashscope_key:
|
||||||
|
logger.error("[DeepviewMaterials] Missing DASHSCOPE_API_KEY for ASR.")
|
||||||
|
_simulate_extract(md_path, f"{original_filename} (ASR Failed: No DASHSCOPE_API_KEY)")
|
||||||
|
return
|
||||||
|
|
||||||
|
logger.info(f"[DeepviewMaterials] Submitting Long Audio ASR Task for {original_filename}")
|
||||||
|
|
||||||
|
# 1. 签名前端传上来的 OSS URL (给阿里大模型长音频异步读取用)
|
||||||
|
audio_url = bucket.sign_url('GET', oss_key, 3600 * 24)
|
||||||
|
|
||||||
|
# 2. 提交异步听写任务
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {dashscope_key}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
"X-DashScope-Async": "enable"
|
||||||
|
}
|
||||||
|
payload = {
|
||||||
|
"model": "paraformer-v2", # qwen-audio 长语音引擎(唯一支持说话人分离的离线异步服务)
|
||||||
|
"input": {"file_urls": [audio_url]},
|
||||||
|
"parameters": {
|
||||||
|
"diarization_enabled": True # 开启说话人角色分离 (音色解析)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
resp = requests.post("https://dashscope.aliyuncs.com/api/v1/services/audio/asr/transcription", json=payload, headers=headers)
|
||||||
|
if resp.status_code != 200:
|
||||||
|
logger.error(f"[DeepviewMaterials] ASR Submit Failed: {resp.text}")
|
||||||
|
_simulate_extract(md_path, f"{original_filename} (ASR Submit Failed)")
|
||||||
|
return
|
||||||
|
|
||||||
|
task_id = resp.json()["output"]["task_id"]
|
||||||
|
logger.info(f"[DeepviewMaterials] ASR Task Submitted: {task_id}. Polling...")
|
||||||
|
|
||||||
|
# 3. 轮询结果
|
||||||
|
polling_url = f"https://dashscope.aliyuncs.com/api/v1/tasks/{task_id}"
|
||||||
|
while True:
|
||||||
|
status_resp = requests.get(polling_url, headers=headers)
|
||||||
|
if status_resp.status_code != 200:
|
||||||
|
logger.error(f"[DeepviewMaterials] ASR Polling HTTP Error: {status_resp.text}")
|
||||||
|
break
|
||||||
|
|
||||||
|
data = status_resp.json()
|
||||||
|
status = data["output"]["task_status"]
|
||||||
|
|
||||||
|
if status == "SUCCEEDED":
|
||||||
|
result_url = data["output"]["results"][0]["transcription_url"]
|
||||||
|
result_resp = requests.get(result_url)
|
||||||
|
transcripts = result_resp.json().get("transcripts", [])
|
||||||
|
|
||||||
|
if not transcripts:
|
||||||
|
with open(md_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(f"# 🎵 面诊录音: {original_filename}\n\n*(未能提取出任何语音)*")
|
||||||
|
return
|
||||||
|
|
||||||
|
sentences = transcripts[0].get("sentences", [])
|
||||||
|
|
||||||
|
# 4. 格式化落盘:音色换行隔离
|
||||||
|
with open(md_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(f"# 🎵 面诊录音: {original_filename}\n\n")
|
||||||
|
last_speaker = None
|
||||||
|
for s in sentences:
|
||||||
|
spk = s.get("speaker_id", "Unknown")
|
||||||
|
# Paraformer 通常返回 spk_0, spk_1 或者根据音色聚类
|
||||||
|
spk_label = f"**说话人 {spk}**" if spk != "Unknown" else "**未知说话人**"
|
||||||
|
text = s.get("text", "")
|
||||||
|
|
||||||
|
if spk != last_speaker:
|
||||||
|
f.write(f"\n\n{spk_label}: {text}")
|
||||||
|
last_speaker = spk
|
||||||
|
else:
|
||||||
|
f.write(f" {text}")
|
||||||
|
|
||||||
|
logger.info(f"[DeepviewMaterials] ASR Diarization successfully completed for {original_filename}")
|
||||||
|
return
|
||||||
|
|
||||||
|
elif status == "FAILED":
|
||||||
|
logger.error(f"[DeepviewMaterials] ASR Task Failed Internally: {data}")
|
||||||
|
break
|
||||||
|
|
||||||
|
time.sleep(3) # 轮询间隔
|
||||||
|
|
||||||
|
# 兜底
|
||||||
|
_simulate_extract(md_path, f"{original_filename} (ASR Polling Failed or Timeout)")
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"[DeepviewMaterials] audio error: {e}", exc_info=True)
|
||||||
|
_simulate_extract(md_path, os.path.basename(audio_path))
|
||||||
|
|
||||||
|
def _simulate_extract(md_path: str, original_filename: str):
|
||||||
|
with open(md_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(f"# Document: {original_filename}\n\nNotice: Extracted via simple fallback parser.")
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,333 @@
|
|||||||
|
"""
|
||||||
|
DingTalk platform adapter using Stream Mode.
|
||||||
|
|
||||||
|
Uses dingtalk-stream SDK for real-time message reception without webhooks.
|
||||||
|
Responses are sent via DingTalk's session webhook (markdown format).
|
||||||
|
|
||||||
|
Requires:
|
||||||
|
pip install dingtalk-stream httpx
|
||||||
|
DINGTALK_CLIENT_ID and DINGTALK_CLIENT_SECRET env vars
|
||||||
|
|
||||||
|
Configuration in config.yaml:
|
||||||
|
platforms:
|
||||||
|
dingtalk:
|
||||||
|
enabled: true
|
||||||
|
extra:
|
||||||
|
client_id: "your-app-key" # or DINGTALK_CLIENT_ID env var
|
||||||
|
client_secret: "your-secret" # or DINGTALK_CLIENT_SECRET env var
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
|
try:
|
||||||
|
import dingtalk_stream
|
||||||
|
from dingtalk_stream import ChatbotHandler, ChatbotMessage
|
||||||
|
DINGTALK_STREAM_AVAILABLE = True
|
||||||
|
except ImportError:
|
||||||
|
DINGTALK_STREAM_AVAILABLE = False
|
||||||
|
dingtalk_stream = None # type: ignore[assignment]
|
||||||
|
|
||||||
|
try:
|
||||||
|
import httpx
|
||||||
|
HTTPX_AVAILABLE = True
|
||||||
|
except ImportError:
|
||||||
|
HTTPX_AVAILABLE = False
|
||||||
|
httpx = None # type: ignore[assignment]
|
||||||
|
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
from gateway.platforms.helpers import MessageDeduplicator
|
||||||
|
from gateway.platforms.base import (
|
||||||
|
BasePlatformAdapter,
|
||||||
|
MessageEvent,
|
||||||
|
MessageType,
|
||||||
|
SendResult,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
MAX_MESSAGE_LENGTH = 20000
|
||||||
|
RECONNECT_BACKOFF = [2, 5, 10, 30, 60]
|
||||||
|
_SESSION_WEBHOOKS_MAX = 500
|
||||||
|
_DINGTALK_WEBHOOK_RE = re.compile(r'^https://api\.dingtalk\.com/')
|
||||||
|
|
||||||
|
|
||||||
|
def check_dingtalk_requirements() -> bool:
|
||||||
|
"""Check if DingTalk dependencies are available and configured."""
|
||||||
|
if not DINGTALK_STREAM_AVAILABLE or not HTTPX_AVAILABLE:
|
||||||
|
return False
|
||||||
|
if not os.getenv("DINGTALK_CLIENT_ID") or not os.getenv("DINGTALK_CLIENT_SECRET"):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
class DingTalkAdapter(BasePlatformAdapter):
|
||||||
|
"""DingTalk chatbot adapter using Stream Mode.
|
||||||
|
|
||||||
|
The dingtalk-stream SDK maintains a long-lived WebSocket connection.
|
||||||
|
Incoming messages arrive via a ChatbotHandler callback. Replies are
|
||||||
|
sent via the incoming message's session_webhook URL using httpx.
|
||||||
|
"""
|
||||||
|
|
||||||
|
MAX_MESSAGE_LENGTH = MAX_MESSAGE_LENGTH
|
||||||
|
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.DINGTALK)
|
||||||
|
|
||||||
|
extra = config.extra or {}
|
||||||
|
self._client_id: str = extra.get("client_id") or os.getenv("DINGTALK_CLIENT_ID", "")
|
||||||
|
self._client_secret: str = extra.get("client_secret") or os.getenv("DINGTALK_CLIENT_SECRET", "")
|
||||||
|
|
||||||
|
self._stream_client: Any = None
|
||||||
|
self._stream_task: Optional[asyncio.Task] = None
|
||||||
|
self._http_client: Optional["httpx.AsyncClient"] = None
|
||||||
|
|
||||||
|
# Message deduplication
|
||||||
|
self._dedup = MessageDeduplicator(max_size=1000)
|
||||||
|
# Map chat_id -> session_webhook for reply routing
|
||||||
|
self._session_webhooks: Dict[str, str] = {}
|
||||||
|
|
||||||
|
# -- Connection lifecycle -----------------------------------------------
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
"""Connect to DingTalk via Stream Mode."""
|
||||||
|
if not DINGTALK_STREAM_AVAILABLE:
|
||||||
|
logger.warning("[%s] dingtalk-stream not installed. Run: pip install dingtalk-stream", self.name)
|
||||||
|
return False
|
||||||
|
if not HTTPX_AVAILABLE:
|
||||||
|
logger.warning("[%s] httpx not installed. Run: pip install httpx", self.name)
|
||||||
|
return False
|
||||||
|
if not self._client_id or not self._client_secret:
|
||||||
|
logger.warning("[%s] DINGTALK_CLIENT_ID and DINGTALK_CLIENT_SECRET required", self.name)
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
self._http_client = httpx.AsyncClient(timeout=30.0)
|
||||||
|
|
||||||
|
credential = dingtalk_stream.Credential(self._client_id, self._client_secret)
|
||||||
|
self._stream_client = dingtalk_stream.DingTalkStreamClient(credential)
|
||||||
|
|
||||||
|
# Capture the current event loop for cross-thread dispatch
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
handler = _IncomingHandler(self, loop)
|
||||||
|
self._stream_client.register_callback_handler(
|
||||||
|
dingtalk_stream.ChatbotMessage.TOPIC, handler
|
||||||
|
)
|
||||||
|
|
||||||
|
self._stream_task = asyncio.create_task(self._run_stream())
|
||||||
|
self._mark_connected()
|
||||||
|
logger.info("[%s] Connected via Stream Mode", self.name)
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[%s] Failed to connect: %s", self.name, e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _run_stream(self) -> None:
|
||||||
|
"""Run the blocking stream client with auto-reconnection."""
|
||||||
|
backoff_idx = 0
|
||||||
|
while self._running:
|
||||||
|
try:
|
||||||
|
logger.debug("[%s] Starting stream client...", self.name)
|
||||||
|
await asyncio.to_thread(self._stream_client.start)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
return
|
||||||
|
except Exception as e:
|
||||||
|
if not self._running:
|
||||||
|
return
|
||||||
|
logger.warning("[%s] Stream client error: %s", self.name, e)
|
||||||
|
|
||||||
|
if not self._running:
|
||||||
|
return
|
||||||
|
|
||||||
|
delay = RECONNECT_BACKOFF[min(backoff_idx, len(RECONNECT_BACKOFF) - 1)]
|
||||||
|
logger.info("[%s] Reconnecting in %ds...", self.name, delay)
|
||||||
|
await asyncio.sleep(delay)
|
||||||
|
backoff_idx += 1
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
"""Disconnect from DingTalk."""
|
||||||
|
self._running = False
|
||||||
|
self._mark_disconnected()
|
||||||
|
|
||||||
|
if self._stream_task:
|
||||||
|
self._stream_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._stream_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
self._stream_task = None
|
||||||
|
|
||||||
|
if self._http_client:
|
||||||
|
await self._http_client.aclose()
|
||||||
|
self._http_client = None
|
||||||
|
|
||||||
|
self._stream_client = None
|
||||||
|
self._session_webhooks.clear()
|
||||||
|
self._dedup.clear()
|
||||||
|
logger.info("[%s] Disconnected", self.name)
|
||||||
|
|
||||||
|
# -- Inbound message processing -----------------------------------------
|
||||||
|
|
||||||
|
async def _on_message(self, message: "ChatbotMessage") -> None:
|
||||||
|
"""Process an incoming DingTalk chatbot message."""
|
||||||
|
msg_id = getattr(message, "message_id", None) or uuid.uuid4().hex
|
||||||
|
if self._dedup.is_duplicate(msg_id):
|
||||||
|
logger.debug("[%s] Duplicate message %s, skipping", self.name, msg_id)
|
||||||
|
return
|
||||||
|
|
||||||
|
text = self._extract_text(message)
|
||||||
|
if not text:
|
||||||
|
logger.debug("[%s] Empty message, skipping", self.name)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Chat context
|
||||||
|
conversation_id = getattr(message, "conversation_id", "") or ""
|
||||||
|
conversation_type = getattr(message, "conversation_type", "1")
|
||||||
|
is_group = str(conversation_type) == "2"
|
||||||
|
sender_id = getattr(message, "sender_id", "") or ""
|
||||||
|
sender_nick = getattr(message, "sender_nick", "") or sender_id
|
||||||
|
sender_staff_id = getattr(message, "sender_staff_id", "") or ""
|
||||||
|
|
||||||
|
chat_id = conversation_id or sender_id
|
||||||
|
chat_type = "group" if is_group else "dm"
|
||||||
|
|
||||||
|
# Store session webhook for reply routing (validate origin to prevent SSRF)
|
||||||
|
session_webhook = getattr(message, "session_webhook", None) or ""
|
||||||
|
if session_webhook and chat_id and _DINGTALK_WEBHOOK_RE.match(session_webhook):
|
||||||
|
if len(self._session_webhooks) >= _SESSION_WEBHOOKS_MAX:
|
||||||
|
# Evict oldest entry to cap memory growth
|
||||||
|
try:
|
||||||
|
self._session_webhooks.pop(next(iter(self._session_webhooks)))
|
||||||
|
except StopIteration:
|
||||||
|
pass
|
||||||
|
self._session_webhooks[chat_id] = session_webhook
|
||||||
|
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id=chat_id,
|
||||||
|
chat_name=getattr(message, "conversation_title", None),
|
||||||
|
chat_type=chat_type,
|
||||||
|
user_id=sender_id,
|
||||||
|
user_name=sender_nick,
|
||||||
|
user_id_alt=sender_staff_id if sender_staff_id else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Parse timestamp
|
||||||
|
create_at = getattr(message, "create_at", None)
|
||||||
|
try:
|
||||||
|
timestamp = datetime.fromtimestamp(int(create_at) / 1000, tz=timezone.utc) if create_at else datetime.now(tz=timezone.utc)
|
||||||
|
except (ValueError, OSError, TypeError):
|
||||||
|
timestamp = datetime.now(tz=timezone.utc)
|
||||||
|
|
||||||
|
event = MessageEvent(
|
||||||
|
text=text,
|
||||||
|
message_type=MessageType.TEXT,
|
||||||
|
source=source,
|
||||||
|
message_id=msg_id,
|
||||||
|
raw_message=message,
|
||||||
|
timestamp=timestamp,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.debug("[%s] Message from %s in %s: %s",
|
||||||
|
self.name, sender_nick, chat_id[:20] if chat_id else "?", text[:50])
|
||||||
|
await self.handle_message(event)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_text(message: "ChatbotMessage") -> str:
|
||||||
|
"""Extract plain text from a DingTalk chatbot message."""
|
||||||
|
text = getattr(message, "text", None) or ""
|
||||||
|
if isinstance(text, dict):
|
||||||
|
content = text.get("content", "").strip()
|
||||||
|
else:
|
||||||
|
content = str(text).strip()
|
||||||
|
|
||||||
|
# Fall back to rich text if present
|
||||||
|
if not content:
|
||||||
|
rich_text = getattr(message, "rich_text", None)
|
||||||
|
if rich_text and isinstance(rich_text, list):
|
||||||
|
parts = [item["text"] for item in rich_text
|
||||||
|
if isinstance(item, dict) and item.get("text")]
|
||||||
|
content = " ".join(parts).strip()
|
||||||
|
return content
|
||||||
|
|
||||||
|
# -- Outbound messaging -------------------------------------------------
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a markdown reply via DingTalk session webhook."""
|
||||||
|
metadata = metadata or {}
|
||||||
|
|
||||||
|
session_webhook = metadata.get("session_webhook") or self._session_webhooks.get(chat_id)
|
||||||
|
if not session_webhook:
|
||||||
|
return SendResult(success=False,
|
||||||
|
error="No session_webhook available. Reply must follow an incoming message.")
|
||||||
|
|
||||||
|
if not self._http_client:
|
||||||
|
return SendResult(success=False, error="HTTP client not initialized")
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"msgtype": "markdown",
|
||||||
|
"markdown": {"title": "Hermes", "text": content[:self.MAX_MESSAGE_LENGTH]},
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
resp = await self._http_client.post(session_webhook, json=payload, timeout=15.0)
|
||||||
|
if resp.status_code < 300:
|
||||||
|
return SendResult(success=True, message_id=uuid.uuid4().hex[:12])
|
||||||
|
body = resp.text
|
||||||
|
logger.warning("[%s] Send failed HTTP %d: %s", self.name, resp.status_code, body[:200])
|
||||||
|
return SendResult(success=False, error=f"HTTP {resp.status_code}: {body[:200]}")
|
||||||
|
except httpx.TimeoutException:
|
||||||
|
return SendResult(success=False, error="Timeout sending message to DingTalk")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[%s] Send error: %s", self.name, e)
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
async def send_typing(self, chat_id: str, metadata=None) -> None:
|
||||||
|
"""DingTalk does not support typing indicators."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
"""Return basic info about a DingTalk conversation."""
|
||||||
|
return {"name": chat_id, "type": "group" if "group" in chat_id.lower() else "dm"}
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Internal stream handler
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class _IncomingHandler(ChatbotHandler if DINGTALK_STREAM_AVAILABLE else object):
|
||||||
|
"""dingtalk-stream ChatbotHandler that forwards messages to the adapter."""
|
||||||
|
|
||||||
|
def __init__(self, adapter: DingTalkAdapter, loop: asyncio.AbstractEventLoop):
|
||||||
|
if DINGTALK_STREAM_AVAILABLE:
|
||||||
|
super().__init__()
|
||||||
|
self._adapter = adapter
|
||||||
|
self._loop = loop
|
||||||
|
|
||||||
|
def process(self, message: "ChatbotMessage"):
|
||||||
|
"""Called by dingtalk-stream in its thread when a message arrives.
|
||||||
|
|
||||||
|
Schedules the async handler on the main event loop.
|
||||||
|
"""
|
||||||
|
loop = self._loop
|
||||||
|
if loop is None or loop.is_closed():
|
||||||
|
logger.error("[DingTalk] Event loop unavailable, cannot dispatch message")
|
||||||
|
return dingtalk_stream.AckMessage.STATUS_OK, "OK"
|
||||||
|
|
||||||
|
future = asyncio.run_coroutine_threadsafe(self._adapter._on_message(message), loop)
|
||||||
|
try:
|
||||||
|
future.result(timeout=60)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("[DingTalk] Error processing incoming message")
|
||||||
|
|
||||||
|
return dingtalk_stream.AckMessage.STATUS_OK, "OK"
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,335 @@
|
|||||||
|
"""
|
||||||
|
⚠️ DEPRECATED — 请勿修改、复制或引用此文件 ⚠️
|
||||||
|
|
||||||
|
本文件已被 infra.pipelines.anyfile2md 统一管线取代。
|
||||||
|
mindos_sse.py 已完成迁移(2026-04-20),不再调用此模块。
|
||||||
|
|
||||||
|
新的统一管线位于:
|
||||||
|
mindOSv2/hermes-overlay/infra/pipelines/anyfile2md.py
|
||||||
|
mindOSv2/hermes-overlay/infra/atoms/ (6 个原子操作)
|
||||||
|
|
||||||
|
原始文件保留仅供参考,下次部署时将删除。
|
||||||
|
|
||||||
|
──────────────────────────────────────────────────────
|
||||||
|
以下为原始 docstring(仅供考古):
|
||||||
|
|
||||||
|
doc_parser.py — AnyFile2MD 通用文件解析模块(已废弃)
|
||||||
|
|
||||||
|
架构约束(SPEC_anyfile2md_doc_parser.md v1.1):
|
||||||
|
❌ 不依赖 LLMAgno(V2 遗留)——前代隔离铁律
|
||||||
|
❌ 不调用 LLM(VL 降级除外),不写 DB,不持久化任何状态
|
||||||
|
✅ VL 调用直接走 LiteLLM Gateway(OpenAI SDK)
|
||||||
|
✅ 封闭函数(铁律 1):只做格式分发 + 文本提取 + 返回 Markdown
|
||||||
|
|
||||||
|
公开接口:
|
||||||
|
parse_to_markdown(read_url, filename, ...) → str
|
||||||
|
is_supported_doc(filename) → bool
|
||||||
|
SUPPORTED_DOC_EXTS → set
|
||||||
|
|
||||||
|
依赖:
|
||||||
|
python-docx (≈V2 mammoth —— DOCX 文本提取)
|
||||||
|
pdfplumber (≈V2 pdf-parse —— PDF 文字提取)
|
||||||
|
PyMuPDF (fitz) (≈V2 pdf-to-img —— PDF 逐页光栅化,VL 降级用)
|
||||||
|
openai (VL 调用,hermes 已有)
|
||||||
|
httpx (OSS 下载/上传,hermes 已有)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# ─── 格式分类(移植自 V2 normalize-manager.cjs L14-30) ─────
|
||||||
|
|
||||||
|
AUDIO_EXTS = {
|
||||||
|
".mp3", ".mp4", ".m4a", ".wav", ".webm",
|
||||||
|
".ogg", ".aac", ".flac", ".opus", ".amr",
|
||||||
|
}
|
||||||
|
IMAGE_EXTS = {
|
||||||
|
".png", ".jpg", ".jpeg", ".bmp", ".gif",
|
||||||
|
".webp", ".tiff", ".tif",
|
||||||
|
}
|
||||||
|
DOCX_EXTS = {".docx", ".doc"}
|
||||||
|
PDF_EXTS = {".pdf"}
|
||||||
|
TEXT_EXTS = {".txt", ".md", ".markdown", ".rst", ".log"}
|
||||||
|
|
||||||
|
# 所有支持的非音频格式(前端 accept 属性用)
|
||||||
|
SUPPORTED_DOC_EXTS = IMAGE_EXTS | DOCX_EXTS | PDF_EXTS | TEXT_EXTS
|
||||||
|
|
||||||
|
|
||||||
|
def is_supported_doc(filename: str) -> bool:
|
||||||
|
"""判断文件是否为支持的文档/图片格式。"""
|
||||||
|
ext = Path(filename).suffix.lower()
|
||||||
|
return ext in SUPPORTED_DOC_EXTS
|
||||||
|
|
||||||
|
|
||||||
|
# ─── VL 提示词(从 LLMAgno/server.py 复制并固化) ────────────
|
||||||
|
|
||||||
|
VL_EXTRACT_PROMPT = """请仔细阅读这张图片中的所有内容,并以结构化的纯文本形式输出:
|
||||||
|
|
||||||
|
1. 完整转录图片中的所有文字内容,保留原始排版结构(标题、段落、列表等)
|
||||||
|
2. 如果图片包含表格,用 Markdown 表格格式输出
|
||||||
|
3. 如果图片包含图表(柱状图、折线图、饼图等),用文字描述图表的数据和趋势
|
||||||
|
4. 如果图片包含流程图或示意图,用文字描述其结构和逻辑关系
|
||||||
|
5. 忽略页眉页脚、页码、水印等装饰性元素
|
||||||
|
|
||||||
|
直接输出内容,不要加任何前缀说明或总结。"""
|
||||||
|
|
||||||
|
|
||||||
|
# ─── 主函数 ─────────────────────────────────────────────────
|
||||||
|
|
||||||
|
async def parse_to_markdown(
|
||||||
|
read_url: str,
|
||||||
|
filename: str,
|
||||||
|
*,
|
||||||
|
max_pages: int = 30,
|
||||||
|
vl_model: str = "qwen-vl",
|
||||||
|
) -> dict:
|
||||||
|
"""
|
||||||
|
下载 OSS 文件 → 按扩展名分流解析 → 返回解析结果。
|
||||||
|
|
||||||
|
封闭函数(铁律 1):不写 DB,不持久化状态。
|
||||||
|
VL 调用直接走 LiteLLM Gateway(前代隔离铁律)。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
read_url: OSS 可读 URL(由 generate_oss_presign 返回)
|
||||||
|
filename: 原始文件名(用于判断扩展名)
|
||||||
|
max_pages: PDF 最大处理页数(安全守卫,V2 验证值=30)
|
||||||
|
vl_model: VL 模型别名(LiteLLM 标准别名)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
{
|
||||||
|
"text": str, # Markdown 格式的文本内容
|
||||||
|
"provider": str, # 使用的解析方式
|
||||||
|
"vl_pages": int, # VL 识别的页数(用于积分计算)
|
||||||
|
}
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: 不支持的文件格式
|
||||||
|
RuntimeError: 解析失败
|
||||||
|
"""
|
||||||
|
ext = Path(filename).suffix.lower()
|
||||||
|
|
||||||
|
if ext in AUDIO_EXTS:
|
||||||
|
raise ValueError(f"音频文件请走 /api/audio/transcribe 管线: {ext}")
|
||||||
|
|
||||||
|
if ext not in SUPPORTED_DOC_EXTS:
|
||||||
|
raise ValueError(f"不支持的文件格式: {ext}")
|
||||||
|
|
||||||
|
# 1. 下载到临时文件
|
||||||
|
tmp_path = await _download_to_temp(read_url, filename)
|
||||||
|
|
||||||
|
try:
|
||||||
|
if ext in TEXT_EXTS:
|
||||||
|
text = _read_text_file(tmp_path)
|
||||||
|
return {"text": text, "provider": "text_read", "vl_pages": 0}
|
||||||
|
|
||||||
|
if ext in IMAGE_EXTS:
|
||||||
|
text = await _vl_extract_image(read_url, filename, vl_model)
|
||||||
|
return {"text": text, "provider": "vl_image", "vl_pages": 1}
|
||||||
|
|
||||||
|
if ext in DOCX_EXTS:
|
||||||
|
text = _extract_docx(tmp_path)
|
||||||
|
return {"text": text, "provider": "python_docx", "vl_pages": 0}
|
||||||
|
|
||||||
|
if ext in PDF_EXTS:
|
||||||
|
# PDF 双策略降级(移植自 V2 normalize-manager L260-280)
|
||||||
|
text = _extract_pdf_text(tmp_path)
|
||||||
|
if len(text.strip()) >= 50:
|
||||||
|
return {"text": text, "provider": "pdfplumber", "vl_pages": 0}
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"[DocParser] PDF 文字过少 (%d字),降级到 VL 逐页识别",
|
||||||
|
len(text.strip()),
|
||||||
|
)
|
||||||
|
text, page_count = await _vl_extract_pdf_pages(
|
||||||
|
tmp_path, filename, max_pages, vl_model,
|
||||||
|
)
|
||||||
|
return {"text": text, "provider": "vl_pdf_pages", "vl_pages": page_count}
|
||||||
|
|
||||||
|
# 兜底:尝试直读
|
||||||
|
text = _read_text_file(tmp_path)
|
||||||
|
return {"text": text, "provider": "text_fallback", "vl_pages": 0}
|
||||||
|
|
||||||
|
finally:
|
||||||
|
_cleanup_temp(tmp_path)
|
||||||
|
|
||||||
|
|
||||||
|
# ─── 各格式解析函数 ──────────────────────────────────────────
|
||||||
|
|
||||||
|
def _extract_docx(file_path: str) -> str:
|
||||||
|
"""DOCX → Markdown 文本。等价于 V2 mammoth.extractRawText()。"""
|
||||||
|
import docx # python-docx
|
||||||
|
|
||||||
|
doc = docx.Document(file_path)
|
||||||
|
paragraphs = [p.text for p in doc.paragraphs if p.text.strip()]
|
||||||
|
return "\n\n".join(paragraphs)
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_pdf_text(file_path: str) -> str:
|
||||||
|
"""PDF 文字提取。等价于 V2 pdf-parse。"""
|
||||||
|
import pdfplumber
|
||||||
|
|
||||||
|
texts = []
|
||||||
|
with pdfplumber.open(file_path) as pdf:
|
||||||
|
for page in pdf.pages:
|
||||||
|
text = page.extract_text()
|
||||||
|
if text:
|
||||||
|
texts.append(text)
|
||||||
|
return "\n\n".join(texts)
|
||||||
|
|
||||||
|
|
||||||
|
async def _vl_extract_image(
|
||||||
|
image_url: str,
|
||||||
|
filename: str,
|
||||||
|
model: str = "qwen-vl",
|
||||||
|
) -> str:
|
||||||
|
"""单张图片 VL 识别。直接走 LiteLLM Gateway(前代隔离铁律)。"""
|
||||||
|
import openai
|
||||||
|
|
||||||
|
client = openai.AsyncOpenAI(
|
||||||
|
base_url=os.getenv("LITELLM_BASE_URL", "http://127.0.0.1:4000/v1"),
|
||||||
|
api_key=os.getenv("LITELLM_API_KEY", ""),
|
||||||
|
)
|
||||||
|
|
||||||
|
start = time.time()
|
||||||
|
response = await client.chat.completions.create(
|
||||||
|
model=model,
|
||||||
|
messages=[{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "image_url", "image_url": {"url": image_url}},
|
||||||
|
{"type": "text", "text": VL_EXTRACT_PROMPT},
|
||||||
|
],
|
||||||
|
}],
|
||||||
|
max_tokens=4096,
|
||||||
|
temperature=0.1,
|
||||||
|
)
|
||||||
|
|
||||||
|
text = response.choices[0].message.content or ""
|
||||||
|
elapsed = int((time.time() - start) * 1000)
|
||||||
|
logger.info(
|
||||||
|
"[DocParser] VL ✅ model=%s, text_len=%d, elapsed=%dms, file=%s",
|
||||||
|
model, len(text), elapsed, filename,
|
||||||
|
)
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
async def _vl_extract_pdf_pages(
|
||||||
|
file_path: str,
|
||||||
|
filename: str,
|
||||||
|
max_pages: int,
|
||||||
|
model: str,
|
||||||
|
) -> tuple[str, int]:
|
||||||
|
"""扫描型 PDF 逐页 VL 识别。移植自 V2 extractPdfViaVL。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(markdown_text, page_count) — page_count 用于积分计算
|
||||||
|
"""
|
||||||
|
import fitz # PyMuPDF
|
||||||
|
|
||||||
|
doc = fitz.open(file_path)
|
||||||
|
total_pages = len(doc)
|
||||||
|
pages = []
|
||||||
|
processed_count = 0
|
||||||
|
|
||||||
|
for i, page in enumerate(doc):
|
||||||
|
if i >= max_pages:
|
||||||
|
logger.warning(
|
||||||
|
"[DocParser] PDF 页数超限 (%d/%d),截断", max_pages, total_pages,
|
||||||
|
)
|
||||||
|
break
|
||||||
|
|
||||||
|
# 光栅化 → PNG 字节
|
||||||
|
pix = page.get_pixmap(dpi=200)
|
||||||
|
img_bytes = pix.tobytes("png")
|
||||||
|
|
||||||
|
# 每页独立 try-except(V2 验证的韧性策略)
|
||||||
|
try:
|
||||||
|
img_url = await _upload_temp_image(
|
||||||
|
img_bytes, f"{filename}_p{i}.png",
|
||||||
|
)
|
||||||
|
text = await _vl_extract_image(img_url, f"{filename}_p{i}", model)
|
||||||
|
if text.strip():
|
||||||
|
pages.append(f"## 第 {i + 1} 页\n\n{text}")
|
||||||
|
processed_count += 1
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"[DocParser] PDF 第%d页 VL 失败(跳过): %s", i + 1, e,
|
||||||
|
)
|
||||||
|
|
||||||
|
doc.close()
|
||||||
|
return "\n\n".join(pages), processed_count
|
||||||
|
|
||||||
|
|
||||||
|
# ─── 辅助函数 ────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _read_text_file(file_path: str) -> str:
|
||||||
|
"""直读文本文件,含二进制安全检查(移植自 V2 normalize-manager L287-299)。"""
|
||||||
|
with open(file_path, "rb") as f:
|
||||||
|
raw = f.read()
|
||||||
|
|
||||||
|
# null 字节检测:防止二进制文件误读
|
||||||
|
if b"\x00" in raw[:1024]:
|
||||||
|
raise ValueError("文件包含二进制内容,无法作为文本处理")
|
||||||
|
|
||||||
|
return raw.decode("utf-8", errors="replace")
|
||||||
|
|
||||||
|
|
||||||
|
async def _download_to_temp(read_url: str, filename: str) -> str:
|
||||||
|
"""从 OSS URL 下载到临时文件,返回临时路径。"""
|
||||||
|
suffix = Path(filename).suffix or ".bin"
|
||||||
|
|
||||||
|
async with httpx.AsyncClient(timeout=120) as client:
|
||||||
|
resp = await client.get(read_url)
|
||||||
|
resp.raise_for_status()
|
||||||
|
|
||||||
|
tmp = tempfile.NamedTemporaryFile(delete=False, suffix=suffix)
|
||||||
|
tmp.write(resp.content)
|
||||||
|
tmp.close()
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"[DocParser] 下载完成: %s → %s (%d bytes)",
|
||||||
|
filename, tmp.name, len(resp.content),
|
||||||
|
)
|
||||||
|
return tmp.name
|
||||||
|
|
||||||
|
|
||||||
|
async def _upload_temp_image(img_bytes: bytes, filename: str) -> str:
|
||||||
|
"""上传临时图片到 OSS,返回可读 URL(VL 逐页降级用)。
|
||||||
|
|
||||||
|
复用 flash_asr.generate_oss_presign + httpx PUT。
|
||||||
|
"""
|
||||||
|
from flash_asr import generate_oss_presign # type: ignore
|
||||||
|
|
||||||
|
ext = Path(filename).suffix or ".png"
|
||||||
|
# 使用独立的 OSS 前缀,与音频/文档区分
|
||||||
|
presign = generate_oss_presign(
|
||||||
|
user_id="vl_temp", ext=ext,
|
||||||
|
prefix="mindos-next/vl-temp",
|
||||||
|
)
|
||||||
|
|
||||||
|
async with httpx.AsyncClient(timeout=60) as client:
|
||||||
|
resp = await client.put(
|
||||||
|
presign["upload_url"],
|
||||||
|
content=img_bytes,
|
||||||
|
headers={"Content-Type": presign["content_type"]},
|
||||||
|
)
|
||||||
|
resp.raise_for_status()
|
||||||
|
|
||||||
|
logger.debug("[DocParser] 临时图片已上传: %s", presign["oss_key"])
|
||||||
|
return presign["read_url"]
|
||||||
|
|
||||||
|
|
||||||
|
def _cleanup_temp(tmp_path: str) -> None:
|
||||||
|
"""删除临时文件,fire-and-forget。"""
|
||||||
|
try:
|
||||||
|
os.unlink(tmp_path)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
@@ -0,0 +1,625 @@
|
|||||||
|
"""
|
||||||
|
Email platform adapter for the Hermes gateway.
|
||||||
|
|
||||||
|
Allows users to interact with Hermes by sending emails.
|
||||||
|
Uses IMAP to receive and SMTP to send messages.
|
||||||
|
|
||||||
|
Environment variables:
|
||||||
|
EMAIL_IMAP_HOST — IMAP server host (e.g., imap.gmail.com)
|
||||||
|
EMAIL_IMAP_PORT — IMAP server port (default: 993)
|
||||||
|
EMAIL_SMTP_HOST — SMTP server host (e.g., smtp.gmail.com)
|
||||||
|
EMAIL_SMTP_PORT — SMTP server port (default: 587)
|
||||||
|
EMAIL_ADDRESS — Email address for the agent
|
||||||
|
EMAIL_PASSWORD — Email password or app-specific password
|
||||||
|
EMAIL_POLL_INTERVAL — Seconds between mailbox checks (default: 15)
|
||||||
|
EMAIL_ALLOWED_USERS — Comma-separated list of allowed sender addresses
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import email as email_lib
|
||||||
|
import imaplib
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
import smtplib
|
||||||
|
import ssl
|
||||||
|
import uuid
|
||||||
|
from email.header import decode_header
|
||||||
|
from email.mime.multipart import MIMEMultipart
|
||||||
|
from email.mime.text import MIMEText
|
||||||
|
from email.mime.base import MIMEBase
|
||||||
|
from email import encoders
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
from gateway.platforms.base import (
|
||||||
|
BasePlatformAdapter,
|
||||||
|
MessageEvent,
|
||||||
|
MessageType,
|
||||||
|
SendResult,
|
||||||
|
cache_document_from_bytes,
|
||||||
|
cache_image_from_bytes,
|
||||||
|
)
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
# Automated sender patterns — emails from these are silently ignored
|
||||||
|
_NOREPLY_PATTERNS = (
|
||||||
|
"noreply", "no-reply", "no_reply", "donotreply", "do-not-reply",
|
||||||
|
"mailer-daemon", "postmaster", "bounce", "notifications@",
|
||||||
|
"automated@", "auto-confirm", "auto-reply", "automailer",
|
||||||
|
)
|
||||||
|
|
||||||
|
# RFC headers that indicate bulk/automated mail
|
||||||
|
_AUTOMATED_HEADERS = {
|
||||||
|
"Auto-Submitted": lambda v: v.lower() != "no",
|
||||||
|
"Precedence": lambda v: v.lower() in ("bulk", "list", "junk"),
|
||||||
|
"X-Auto-Response-Suppress": lambda v: bool(v),
|
||||||
|
"List-Unsubscribe": lambda v: bool(v),
|
||||||
|
}
|
||||||
|
|
||||||
|
# Gmail-safe max length per email body
|
||||||
|
MAX_MESSAGE_LENGTH = 50_000
|
||||||
|
|
||||||
|
# Supported image extensions for inline detection
|
||||||
|
_IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".gif", ".webp"}
|
||||||
|
|
||||||
|
def _is_automated_sender(address: str, headers: dict) -> bool:
|
||||||
|
"""Return True if this email is from an automated/noreply source."""
|
||||||
|
addr = address.lower()
|
||||||
|
if any(pattern in addr for pattern in _NOREPLY_PATTERNS):
|
||||||
|
return True
|
||||||
|
for header, check in _AUTOMATED_HEADERS.items():
|
||||||
|
value = headers.get(header, "")
|
||||||
|
if value and check(value):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
def check_email_requirements() -> bool:
|
||||||
|
"""Check if email platform dependencies are available."""
|
||||||
|
addr = os.getenv("EMAIL_ADDRESS")
|
||||||
|
pwd = os.getenv("EMAIL_PASSWORD")
|
||||||
|
imap = os.getenv("EMAIL_IMAP_HOST")
|
||||||
|
smtp = os.getenv("EMAIL_SMTP_HOST")
|
||||||
|
if not all([addr, pwd, imap, smtp]):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def _decode_header_value(raw: str) -> str:
|
||||||
|
"""Decode an RFC 2047 encoded email header into a plain string."""
|
||||||
|
parts = decode_header(raw)
|
||||||
|
decoded = []
|
||||||
|
for part, charset in parts:
|
||||||
|
if isinstance(part, bytes):
|
||||||
|
decoded.append(part.decode(charset or "utf-8", errors="replace"))
|
||||||
|
else:
|
||||||
|
decoded.append(part)
|
||||||
|
return " ".join(decoded)
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_text_body(msg: email_lib.message.Message) -> str:
|
||||||
|
"""Extract the plain-text body from a potentially multipart email."""
|
||||||
|
if msg.is_multipart():
|
||||||
|
for part in msg.walk():
|
||||||
|
content_type = part.get_content_type()
|
||||||
|
disposition = str(part.get("Content-Disposition", ""))
|
||||||
|
# Skip attachments
|
||||||
|
if "attachment" in disposition:
|
||||||
|
continue
|
||||||
|
if content_type == "text/plain":
|
||||||
|
payload = part.get_payload(decode=True)
|
||||||
|
if payload:
|
||||||
|
charset = part.get_content_charset() or "utf-8"
|
||||||
|
return payload.decode(charset, errors="replace")
|
||||||
|
# Fallback: try text/html and strip tags
|
||||||
|
for part in msg.walk():
|
||||||
|
content_type = part.get_content_type()
|
||||||
|
disposition = str(part.get("Content-Disposition", ""))
|
||||||
|
if "attachment" in disposition:
|
||||||
|
continue
|
||||||
|
if content_type == "text/html":
|
||||||
|
payload = part.get_payload(decode=True)
|
||||||
|
if payload:
|
||||||
|
charset = part.get_content_charset() or "utf-8"
|
||||||
|
html = payload.decode(charset, errors="replace")
|
||||||
|
return _strip_html(html)
|
||||||
|
return ""
|
||||||
|
else:
|
||||||
|
payload = msg.get_payload(decode=True)
|
||||||
|
if payload:
|
||||||
|
charset = msg.get_content_charset() or "utf-8"
|
||||||
|
text = payload.decode(charset, errors="replace")
|
||||||
|
if msg.get_content_type() == "text/html":
|
||||||
|
return _strip_html(text)
|
||||||
|
return text
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _strip_html(html: str) -> str:
|
||||||
|
"""Naive HTML tag stripper for fallback text extraction."""
|
||||||
|
text = re.sub(r"<br\s*/?>", "\n", html, flags=re.IGNORECASE)
|
||||||
|
text = re.sub(r"<p[^>]*>", "\n", text, flags=re.IGNORECASE)
|
||||||
|
text = re.sub(r"</p>", "\n", text, flags=re.IGNORECASE)
|
||||||
|
text = re.sub(r"<[^>]+>", "", text)
|
||||||
|
text = re.sub(r" ", " ", text)
|
||||||
|
text = re.sub(r"&", "&", text)
|
||||||
|
text = re.sub(r"<", "<", text)
|
||||||
|
text = re.sub(r">", ">", text)
|
||||||
|
text = re.sub(r"\n{3,}", "\n\n", text)
|
||||||
|
return text.strip()
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_email_address(raw: str) -> str:
|
||||||
|
"""Extract bare email address from 'Name <addr>' format."""
|
||||||
|
match = re.search(r"<([^>]+)>", raw)
|
||||||
|
if match:
|
||||||
|
return match.group(1).strip().lower()
|
||||||
|
return raw.strip().lower()
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_attachments(
|
||||||
|
msg: email_lib.message.Message,
|
||||||
|
skip_attachments: bool = False,
|
||||||
|
) -> List[Dict[str, Any]]:
|
||||||
|
"""Extract attachment metadata and cache files locally.
|
||||||
|
|
||||||
|
When *skip_attachments* is True, all attachment/inline parts are ignored
|
||||||
|
(useful for malware protection or bandwidth savings).
|
||||||
|
"""
|
||||||
|
attachments = []
|
||||||
|
if not msg.is_multipart():
|
||||||
|
return attachments
|
||||||
|
|
||||||
|
for part in msg.walk():
|
||||||
|
disposition = str(part.get("Content-Disposition", ""))
|
||||||
|
if skip_attachments and ("attachment" in disposition or "inline" in disposition):
|
||||||
|
continue
|
||||||
|
if "attachment" not in disposition and "inline" not in disposition:
|
||||||
|
continue
|
||||||
|
# Skip text/plain and text/html body parts
|
||||||
|
content_type = part.get_content_type()
|
||||||
|
if content_type in ("text/plain", "text/html") and "attachment" not in disposition:
|
||||||
|
continue
|
||||||
|
|
||||||
|
filename = part.get_filename()
|
||||||
|
if filename:
|
||||||
|
filename = _decode_header_value(filename)
|
||||||
|
else:
|
||||||
|
ext = part.get_content_subtype() or "bin"
|
||||||
|
filename = f"attachment.{ext}"
|
||||||
|
|
||||||
|
payload = part.get_payload(decode=True)
|
||||||
|
if not payload:
|
||||||
|
continue
|
||||||
|
|
||||||
|
ext = Path(filename).suffix.lower()
|
||||||
|
if ext in _IMAGE_EXTS:
|
||||||
|
try:
|
||||||
|
cached_path = cache_image_from_bytes(payload, ext)
|
||||||
|
except ValueError:
|
||||||
|
logger.debug("Skipping non-image attachment %s (invalid magic bytes)", filename)
|
||||||
|
continue
|
||||||
|
attachments.append({
|
||||||
|
"path": cached_path,
|
||||||
|
"filename": filename,
|
||||||
|
"type": "image",
|
||||||
|
"media_type": content_type,
|
||||||
|
})
|
||||||
|
else:
|
||||||
|
cached_path = cache_document_from_bytes(payload, filename)
|
||||||
|
attachments.append({
|
||||||
|
"path": cached_path,
|
||||||
|
"filename": filename,
|
||||||
|
"type": "document",
|
||||||
|
"media_type": content_type,
|
||||||
|
})
|
||||||
|
|
||||||
|
return attachments
|
||||||
|
|
||||||
|
|
||||||
|
class EmailAdapter(BasePlatformAdapter):
|
||||||
|
"""Email gateway adapter using IMAP (receive) and SMTP (send)."""
|
||||||
|
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.EMAIL)
|
||||||
|
|
||||||
|
self._address = os.getenv("EMAIL_ADDRESS", "")
|
||||||
|
self._password = os.getenv("EMAIL_PASSWORD", "")
|
||||||
|
self._imap_host = os.getenv("EMAIL_IMAP_HOST", "")
|
||||||
|
self._imap_port = int(os.getenv("EMAIL_IMAP_PORT", "993"))
|
||||||
|
self._smtp_host = os.getenv("EMAIL_SMTP_HOST", "")
|
||||||
|
self._smtp_port = int(os.getenv("EMAIL_SMTP_PORT", "587"))
|
||||||
|
self._poll_interval = int(os.getenv("EMAIL_POLL_INTERVAL", "15"))
|
||||||
|
|
||||||
|
# Skip attachments — configured via config.yaml:
|
||||||
|
# platforms:
|
||||||
|
# email:
|
||||||
|
# skip_attachments: true
|
||||||
|
extra = config.extra or {}
|
||||||
|
self._skip_attachments = extra.get("skip_attachments", False)
|
||||||
|
|
||||||
|
# Track message IDs we've already processed to avoid duplicates
|
||||||
|
self._seen_uids: set = set()
|
||||||
|
self._seen_uids_max: int = 2000 # cap to prevent unbounded memory growth
|
||||||
|
self._poll_task: Optional[asyncio.Task] = None
|
||||||
|
|
||||||
|
# Map chat_id (sender email) -> last subject + message-id for threading
|
||||||
|
self._thread_context: Dict[str, Dict[str, str]] = {}
|
||||||
|
|
||||||
|
logger.info("[Email] Adapter initialized for %s", self._address)
|
||||||
|
|
||||||
|
def _trim_seen_uids(self) -> None:
|
||||||
|
"""Keep only the most recent UIDs to prevent unbounded memory growth.
|
||||||
|
|
||||||
|
IMAP UIDs are monotonically increasing integers. When the set grows
|
||||||
|
beyond the cap, we keep only the highest half — old UIDs are safe to
|
||||||
|
drop because new messages always have higher UIDs and IMAP's UNSEEN
|
||||||
|
flag prevents re-delivery regardless.
|
||||||
|
"""
|
||||||
|
if len(self._seen_uids) <= self._seen_uids_max:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
# UIDs are bytes like b'1234' — sort numerically and keep top half
|
||||||
|
sorted_uids = sorted(self._seen_uids, key=lambda u: int(u))
|
||||||
|
keep = self._seen_uids_max // 2
|
||||||
|
self._seen_uids = set(sorted_uids[-keep:])
|
||||||
|
logger.debug("[Email] Trimmed seen UIDs to %d entries", len(self._seen_uids))
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
# Fallback: just clear old entries if sort fails
|
||||||
|
self._seen_uids = set(list(self._seen_uids)[-self._seen_uids_max // 2:])
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
"""Connect to the IMAP server and start polling for new messages."""
|
||||||
|
try:
|
||||||
|
# Test IMAP connection
|
||||||
|
imap = imaplib.IMAP4_SSL(self._imap_host, self._imap_port, timeout=30)
|
||||||
|
imap.login(self._address, self._password)
|
||||||
|
# Mark all existing messages as seen so we only process new ones
|
||||||
|
imap.select("INBOX")
|
||||||
|
status, data = imap.uid("search", None, "ALL")
|
||||||
|
if status == "OK" and data and data[0]:
|
||||||
|
for uid in data[0].split():
|
||||||
|
self._seen_uids.add(uid)
|
||||||
|
# Keep only the most recent UIDs to prevent unbounded growth
|
||||||
|
self._trim_seen_uids()
|
||||||
|
imap.logout()
|
||||||
|
logger.info("[Email] IMAP connection test passed. %d existing messages skipped.", len(self._seen_uids))
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[Email] IMAP connection failed: %s", e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Test SMTP connection
|
||||||
|
smtp = smtplib.SMTP(self._smtp_host, self._smtp_port, timeout=30)
|
||||||
|
smtp.starttls(context=ssl.create_default_context())
|
||||||
|
smtp.login(self._address, self._password)
|
||||||
|
smtp.quit()
|
||||||
|
logger.info("[Email] SMTP connection test passed.")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[Email] SMTP connection failed: %s", e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
self._running = True
|
||||||
|
self._poll_task = asyncio.create_task(self._poll_loop())
|
||||||
|
print(f"[Email] Connected as {self._address}")
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
"""Stop polling and disconnect."""
|
||||||
|
self._running = False
|
||||||
|
if self._poll_task:
|
||||||
|
self._poll_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._poll_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
self._poll_task = None
|
||||||
|
logger.info("[Email] Disconnected.")
|
||||||
|
|
||||||
|
async def _poll_loop(self) -> None:
|
||||||
|
"""Poll IMAP for new messages at regular intervals."""
|
||||||
|
while self._running:
|
||||||
|
try:
|
||||||
|
await self._check_inbox()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
break
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[Email] Poll error: %s", e)
|
||||||
|
await asyncio.sleep(self._poll_interval)
|
||||||
|
|
||||||
|
async def _check_inbox(self) -> None:
|
||||||
|
"""Check INBOX for unseen messages and dispatch them."""
|
||||||
|
# Run IMAP operations in a thread to avoid blocking the event loop
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
messages = await loop.run_in_executor(None, self._fetch_new_messages)
|
||||||
|
for msg_data in messages:
|
||||||
|
await self._dispatch_message(msg_data)
|
||||||
|
|
||||||
|
def _fetch_new_messages(self) -> List[Dict[str, Any]]:
|
||||||
|
"""Fetch new (unseen) messages from IMAP. Runs in executor thread."""
|
||||||
|
results = []
|
||||||
|
try:
|
||||||
|
imap = imaplib.IMAP4_SSL(self._imap_host, self._imap_port, timeout=30)
|
||||||
|
try:
|
||||||
|
imap.login(self._address, self._password)
|
||||||
|
imap.select("INBOX")
|
||||||
|
|
||||||
|
status, data = imap.uid("search", None, "UNSEEN")
|
||||||
|
if status != "OK" or not data or not data[0]:
|
||||||
|
return results
|
||||||
|
|
||||||
|
for uid in data[0].split():
|
||||||
|
if uid in self._seen_uids:
|
||||||
|
continue
|
||||||
|
self._seen_uids.add(uid)
|
||||||
|
# Trim periodically to prevent unbounded memory growth
|
||||||
|
if len(self._seen_uids) > self._seen_uids_max:
|
||||||
|
self._trim_seen_uids()
|
||||||
|
|
||||||
|
status, msg_data = imap.uid("fetch", uid, "(RFC822)")
|
||||||
|
if status != "OK":
|
||||||
|
continue
|
||||||
|
|
||||||
|
raw_email = msg_data[0][1]
|
||||||
|
msg = email_lib.message_from_bytes(raw_email)
|
||||||
|
|
||||||
|
sender_raw = msg.get("From", "")
|
||||||
|
sender_addr = _extract_email_address(sender_raw)
|
||||||
|
sender_name = _decode_header_value(sender_raw)
|
||||||
|
# Remove email from name if present
|
||||||
|
if "<" in sender_name:
|
||||||
|
sender_name = sender_name.split("<")[0].strip().strip('"')
|
||||||
|
|
||||||
|
subject = _decode_header_value(msg.get("Subject", "(no subject)"))
|
||||||
|
message_id = msg.get("Message-ID", "")
|
||||||
|
in_reply_to = msg.get("In-Reply-To", "")
|
||||||
|
# Skip automated/noreply senders before any processing
|
||||||
|
msg_headers = dict(msg.items())
|
||||||
|
if _is_automated_sender(sender_addr, msg_headers):
|
||||||
|
logger.debug("[Email] Skipping automated sender: %s", sender_addr)
|
||||||
|
continue
|
||||||
|
body = _extract_text_body(msg)
|
||||||
|
attachments = _extract_attachments(msg, skip_attachments=self._skip_attachments)
|
||||||
|
|
||||||
|
results.append({
|
||||||
|
"uid": uid,
|
||||||
|
"sender_addr": sender_addr,
|
||||||
|
"sender_name": sender_name,
|
||||||
|
"subject": subject,
|
||||||
|
"message_id": message_id,
|
||||||
|
"in_reply_to": in_reply_to,
|
||||||
|
"body": body,
|
||||||
|
"attachments": attachments,
|
||||||
|
"date": msg.get("Date", ""),
|
||||||
|
})
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
imap.logout()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[Email] IMAP fetch error: %s", e)
|
||||||
|
return results
|
||||||
|
|
||||||
|
async def _dispatch_message(self, msg_data: Dict[str, Any]) -> None:
|
||||||
|
"""Convert a fetched email into a MessageEvent and dispatch it."""
|
||||||
|
sender_addr = msg_data["sender_addr"]
|
||||||
|
|
||||||
|
# Skip self-messages
|
||||||
|
if sender_addr == self._address.lower():
|
||||||
|
return
|
||||||
|
|
||||||
|
# Never reply to automated senders
|
||||||
|
if _is_automated_sender(sender_addr, {}):
|
||||||
|
logger.debug("[Email] Dropping automated sender at dispatch: %s", sender_addr)
|
||||||
|
return
|
||||||
|
|
||||||
|
subject = msg_data["subject"]
|
||||||
|
body = msg_data["body"].strip()
|
||||||
|
attachments = msg_data["attachments"]
|
||||||
|
|
||||||
|
# Build message text: include subject as context
|
||||||
|
text = body
|
||||||
|
if subject and not subject.startswith("Re:"):
|
||||||
|
text = f"[Subject: {subject}]\n\n{body}"
|
||||||
|
|
||||||
|
# Determine message type and media
|
||||||
|
media_urls = []
|
||||||
|
media_types = []
|
||||||
|
msg_type = MessageType.TEXT
|
||||||
|
|
||||||
|
for att in attachments:
|
||||||
|
media_urls.append(att["path"])
|
||||||
|
media_types.append(att["media_type"])
|
||||||
|
if att["type"] == "image":
|
||||||
|
msg_type = MessageType.PHOTO
|
||||||
|
|
||||||
|
# Store thread context for reply threading
|
||||||
|
self._thread_context[sender_addr] = {
|
||||||
|
"subject": subject,
|
||||||
|
"message_id": msg_data["message_id"],
|
||||||
|
}
|
||||||
|
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id=sender_addr,
|
||||||
|
chat_name=msg_data["sender_name"] or sender_addr,
|
||||||
|
chat_type="dm",
|
||||||
|
user_id=sender_addr,
|
||||||
|
user_name=msg_data["sender_name"] or sender_addr,
|
||||||
|
)
|
||||||
|
|
||||||
|
event = MessageEvent(
|
||||||
|
text=text or "(empty email)",
|
||||||
|
message_type=msg_type,
|
||||||
|
source=source,
|
||||||
|
message_id=msg_data["message_id"],
|
||||||
|
media_urls=media_urls,
|
||||||
|
media_types=media_types,
|
||||||
|
reply_to_message_id=msg_data["in_reply_to"] or None,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info("[Email] New message from %s: %s", sender_addr, subject)
|
||||||
|
await self.handle_message(event)
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send an email reply to the given address."""
|
||||||
|
try:
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
message_id = await loop.run_in_executor(
|
||||||
|
None, self._send_email, chat_id, content, reply_to
|
||||||
|
)
|
||||||
|
return SendResult(success=True, message_id=message_id)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[Email] Send failed to %s: %s", chat_id, e)
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
def _send_email(
|
||||||
|
self,
|
||||||
|
to_addr: str,
|
||||||
|
body: str,
|
||||||
|
reply_to_msg_id: Optional[str] = None,
|
||||||
|
) -> str:
|
||||||
|
"""Send an email via SMTP. Runs in executor thread."""
|
||||||
|
msg = MIMEMultipart()
|
||||||
|
msg["From"] = self._address
|
||||||
|
msg["To"] = to_addr
|
||||||
|
|
||||||
|
# Thread context for reply
|
||||||
|
ctx = self._thread_context.get(to_addr, {})
|
||||||
|
subject = ctx.get("subject", "Hermes Agent")
|
||||||
|
if not subject.startswith("Re:"):
|
||||||
|
subject = f"Re: {subject}"
|
||||||
|
msg["Subject"] = subject
|
||||||
|
|
||||||
|
# Threading headers
|
||||||
|
original_msg_id = reply_to_msg_id or ctx.get("message_id")
|
||||||
|
if original_msg_id:
|
||||||
|
msg["In-Reply-To"] = original_msg_id
|
||||||
|
msg["References"] = original_msg_id
|
||||||
|
|
||||||
|
msg_id = f"<hermes-{uuid.uuid4().hex[:12]}@{self._address.split('@')[1]}>"
|
||||||
|
msg["Message-ID"] = msg_id
|
||||||
|
|
||||||
|
msg.attach(MIMEText(body, "plain", "utf-8"))
|
||||||
|
|
||||||
|
smtp = smtplib.SMTP(self._smtp_host, self._smtp_port, timeout=30)
|
||||||
|
try:
|
||||||
|
smtp.starttls(context=ssl.create_default_context())
|
||||||
|
smtp.login(self._address, self._password)
|
||||||
|
smtp.send_message(msg)
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
smtp.quit()
|
||||||
|
except Exception:
|
||||||
|
smtp.close()
|
||||||
|
|
||||||
|
logger.info("[Email] Sent reply to %s (subject: %s)", to_addr, subject)
|
||||||
|
return msg_id
|
||||||
|
|
||||||
|
async def send_typing(self, chat_id: str, metadata: Optional[Dict[str, Any]] = None) -> None:
|
||||||
|
"""Email has no typing indicator — no-op."""
|
||||||
|
|
||||||
|
async def send_image(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
image_url: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send an image URL as part of an email body."""
|
||||||
|
text = caption or ""
|
||||||
|
text += f"\n\nImage: {image_url}"
|
||||||
|
return await self.send(chat_id, text.strip(), reply_to)
|
||||||
|
|
||||||
|
async def send_document(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
file_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
file_name: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a file as an email attachment."""
|
||||||
|
try:
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
message_id = await loop.run_in_executor(
|
||||||
|
None,
|
||||||
|
self._send_email_with_attachment,
|
||||||
|
chat_id,
|
||||||
|
caption or "",
|
||||||
|
file_path,
|
||||||
|
file_name,
|
||||||
|
)
|
||||||
|
return SendResult(success=True, message_id=message_id)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[Email] Send document failed: %s", e)
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
def _send_email_with_attachment(
|
||||||
|
self,
|
||||||
|
to_addr: str,
|
||||||
|
body: str,
|
||||||
|
file_path: str,
|
||||||
|
file_name: Optional[str] = None,
|
||||||
|
) -> str:
|
||||||
|
"""Send an email with a file attachment via SMTP."""
|
||||||
|
msg = MIMEMultipart()
|
||||||
|
msg["From"] = self._address
|
||||||
|
msg["To"] = to_addr
|
||||||
|
|
||||||
|
ctx = self._thread_context.get(to_addr, {})
|
||||||
|
subject = ctx.get("subject", "Hermes Agent")
|
||||||
|
if not subject.startswith("Re:"):
|
||||||
|
subject = f"Re: {subject}"
|
||||||
|
msg["Subject"] = subject
|
||||||
|
|
||||||
|
original_msg_id = ctx.get("message_id")
|
||||||
|
if original_msg_id:
|
||||||
|
msg["In-Reply-To"] = original_msg_id
|
||||||
|
msg["References"] = original_msg_id
|
||||||
|
|
||||||
|
msg_id = f"<hermes-{uuid.uuid4().hex[:12]}@{self._address.split('@')[1]}>"
|
||||||
|
msg["Message-ID"] = msg_id
|
||||||
|
|
||||||
|
if body:
|
||||||
|
msg.attach(MIMEText(body, "plain", "utf-8"))
|
||||||
|
|
||||||
|
# Attach file
|
||||||
|
p = Path(file_path)
|
||||||
|
fname = file_name or p.name
|
||||||
|
with open(p, "rb") as f:
|
||||||
|
part = MIMEBase("application", "octet-stream")
|
||||||
|
part.set_payload(f.read())
|
||||||
|
encoders.encode_base64(part)
|
||||||
|
part.add_header("Content-Disposition", f"attachment; filename={fname}")
|
||||||
|
msg.attach(part)
|
||||||
|
|
||||||
|
smtp = smtplib.SMTP(self._smtp_host, self._smtp_port, timeout=30)
|
||||||
|
try:
|
||||||
|
smtp.starttls(context=ssl.create_default_context())
|
||||||
|
smtp.login(self._address, self._password)
|
||||||
|
smtp.send_message(msg)
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
smtp.quit()
|
||||||
|
except Exception:
|
||||||
|
smtp.close()
|
||||||
|
|
||||||
|
return msg_id
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
"""Return basic info about the email chat."""
|
||||||
|
ctx = self._thread_context.get(chat_id, {})
|
||||||
|
return {
|
||||||
|
"name": chat_id,
|
||||||
|
"type": "dm",
|
||||||
|
"chat_id": chat_id,
|
||||||
|
"subject": ctx.get("subject", ""),
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,193 @@
|
|||||||
|
"""
|
||||||
|
flash_asr.py — AnyFile2MD 音频转写模块 v2.0
|
||||||
|
|
||||||
|
架构修正(2026-04-15):
|
||||||
|
❌ 旧方案:base64 multipart → DashScope 同步 API(10MB 限制)
|
||||||
|
✅ 新方案:OSS 签名 URL → DashScope paraformer-v2 异步任务(无大小限制)
|
||||||
|
|
||||||
|
职责(封闭函数,铁律 1):
|
||||||
|
transcribe_from_oss_url(url) → AsyncIterator[{"text", "begin_ms", "end_ms"}]
|
||||||
|
generate_oss_presign(user_id, ext) → {"upload_url", "oss_key", "read_url"}
|
||||||
|
is_supported_audio(filename) → bool
|
||||||
|
|
||||||
|
不调用 LLM,不写 DB,不持久化任何状态。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import hmac
|
||||||
|
import hashlib
|
||||||
|
import base64
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
import logging
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import AsyncIterator
|
||||||
|
from urllib.parse import quote
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# ─── 支持格式 ────────────────────────────────────────────────
|
||||||
|
SUPPORTED_EXTENSIONS = {
|
||||||
|
".mp3", ".mp4", ".m4a", ".wav", ".webm",
|
||||||
|
".ogg", ".aac", ".flac", ".opus", ".amr",
|
||||||
|
}
|
||||||
|
|
||||||
|
def is_supported_audio(filename: str) -> bool:
|
||||||
|
ext = os.path.splitext(filename)[1].lower()
|
||||||
|
return ext in SUPPORTED_EXTENSIONS
|
||||||
|
|
||||||
|
|
||||||
|
# ─── OSS 配置 ─────────────────────────────────────────────────
|
||||||
|
def _oss_config() -> dict:
|
||||||
|
return {
|
||||||
|
"access_key_id": os.getenv("OSS_ACCESS_KEY_ID", ""),
|
||||||
|
"access_key_secret": os.getenv("OSS_ACCESS_KEY_SECRET", ""),
|
||||||
|
"bucket": os.getenv("OSS_BUCKET", "meetings-dev"),
|
||||||
|
"endpoint": os.getenv("OSS_ENDPOINT", "oss-cn-guangzhou.aliyuncs.com"),
|
||||||
|
"prefix": os.getenv("OSS_AUDIO_PREFIX", "mindos-next/audio"),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def generate_oss_presign(user_id: str, ext: str, expire_seconds: int = 3600,
|
||||||
|
prefix: str | None = None) -> dict:
|
||||||
|
"""
|
||||||
|
生成 OSS 预签名 URL(使用 oss2 SDK,V2 验证过的可靠方式)。
|
||||||
|
返回:{ upload_url, oss_key, read_url, content_type }
|
||||||
|
- upload_url :前端直传用(HTTP PUT,必须带 Content-Type: application/octet-stream)
|
||||||
|
- read_url :交给 DashScope paraformer-v2 下载用
|
||||||
|
- content_type:前端 PUT 时必须传此 Content-Type
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prefix: OSS 路径前缀。默认 None 时使用 OSS_AUDIO_PREFIX 环境变量。
|
||||||
|
文档管线传 "mindos-next/docs",音频管线不传(使用默认值)。
|
||||||
|
"""
|
||||||
|
import oss2 # type: ignore
|
||||||
|
|
||||||
|
cfg = _oss_config()
|
||||||
|
if not cfg["access_key_id"] or not cfg["access_key_secret"]:
|
||||||
|
raise RuntimeError("[FlashASR] OSS_ACCESS_KEY_ID / OSS_ACCESS_KEY_SECRET 未配置")
|
||||||
|
|
||||||
|
actual_prefix = prefix or cfg["prefix"]
|
||||||
|
ts = int(time.time())
|
||||||
|
oss_key = f"{actual_prefix}/{user_id}/{ts}_{user_id[:8]}{ext}"
|
||||||
|
|
||||||
|
auth = oss2.Auth(cfg["access_key_id"], cfg["access_key_secret"])
|
||||||
|
bucket = oss2.Bucket(auth, f"https://{cfg['endpoint']}", cfg["bucket"])
|
||||||
|
|
||||||
|
# Content-Type 必须写入签名,前端 PUT 时须传相同值
|
||||||
|
content_type = "application/octet-stream"
|
||||||
|
upload_url = bucket.sign_url(
|
||||||
|
"PUT", oss_key, expire_seconds,
|
||||||
|
slash_safe=True,
|
||||||
|
headers={"Content-Type": content_type},
|
||||||
|
)
|
||||||
|
read_url = bucket.sign_url("GET", oss_key, expire_seconds, slash_safe=True)
|
||||||
|
|
||||||
|
host = f"{cfg['bucket']}.{cfg['endpoint']}"
|
||||||
|
logger.info("[FlashASR] presign ok oss_key=%s", oss_key)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"upload_url": upload_url,
|
||||||
|
"read_url": read_url,
|
||||||
|
"oss_key": oss_key,
|
||||||
|
"host": host,
|
||||||
|
"content_type": content_type,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ─── DashScope paraformer-v2 异步任务 ─────────────────────────
|
||||||
|
DASHSCOPE_API_KEY = lambda: os.getenv("DASHSCOPE_API_KEY", "")
|
||||||
|
DASHSCOPE_HOST = "https://dashscope.aliyuncs.com"
|
||||||
|
TRANSCRIPTION_URL = f"{DASHSCOPE_HOST}/api/v1/services/audio/asr/transcription"
|
||||||
|
|
||||||
|
_HEADERS = lambda: {
|
||||||
|
"Authorization": f"Bearer {DASHSCOPE_API_KEY()}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
"X-DashScope-Async": "enable",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _submit_task(file_url: str) -> str:
|
||||||
|
"""提交 paraformer-v2 异步转写任务,返回 task_id。"""
|
||||||
|
body = {
|
||||||
|
"model": "paraformer-v2",
|
||||||
|
"input": {"file_urls": [file_url]},
|
||||||
|
"parameters": {
|
||||||
|
"language_hints": ["zh", "en"],
|
||||||
|
"timestamp_alignment": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
async with httpx.AsyncClient(timeout=30) as client:
|
||||||
|
resp = await client.post(TRANSCRIPTION_URL, json=body, headers=_HEADERS())
|
||||||
|
resp.raise_for_status()
|
||||||
|
data = resp.json()
|
||||||
|
|
||||||
|
task_id = data.get("output", {}).get("task_id")
|
||||||
|
if not task_id:
|
||||||
|
raise RuntimeError(f"[FlashASR] 提交任务失败: {data}")
|
||||||
|
logger.info("[FlashASR] 任务已提交 task_id=%s", task_id)
|
||||||
|
return task_id
|
||||||
|
|
||||||
|
|
||||||
|
async def _poll_task(task_id: str, poll_interval: float = 5.0, timeout: float = 600.0) -> dict:
|
||||||
|
"""轮询任务状态,返回完成后的 output。"""
|
||||||
|
url = f"{DASHSCOPE_HOST}/api/v1/tasks/{task_id}"
|
||||||
|
deadline = time.monotonic() + timeout
|
||||||
|
|
||||||
|
async with httpx.AsyncClient(timeout=30) as client:
|
||||||
|
while time.monotonic() < deadline:
|
||||||
|
resp = await client.get(url, headers=_HEADERS())
|
||||||
|
resp.raise_for_status()
|
||||||
|
data = resp.json()
|
||||||
|
status = data.get("output", {}).get("task_status", "")
|
||||||
|
|
||||||
|
if status == "SUCCEEDED":
|
||||||
|
logger.info("[FlashASR] 任务完成 task_id=%s", task_id)
|
||||||
|
return data["output"]
|
||||||
|
elif status in ("FAILED", "CANCELED"):
|
||||||
|
raise RuntimeError(f"[FlashASR] 任务失败 task_id={task_id} status={status}: {data}")
|
||||||
|
|
||||||
|
logger.debug("[FlashASR] 轮询中 task_id=%s status=%s", task_id, status)
|
||||||
|
await asyncio.sleep(poll_interval)
|
||||||
|
|
||||||
|
raise RuntimeError(f"[FlashASR] 任务超时 task_id={task_id} (>{timeout}s)")
|
||||||
|
|
||||||
|
|
||||||
|
async def transcribe_from_oss_url(read_url: str) -> AsyncIterator[dict]:
|
||||||
|
"""
|
||||||
|
调用 DashScope paraformer-v2 异步转写。
|
||||||
|
yield {"text": str, "begin_ms": int, "end_ms": int}
|
||||||
|
"""
|
||||||
|
if not DASHSCOPE_API_KEY():
|
||||||
|
raise RuntimeError("[FlashASR] DASHSCOPE_API_KEY 未配置")
|
||||||
|
|
||||||
|
task_id = await _submit_task(read_url)
|
||||||
|
output = await _poll_task(task_id)
|
||||||
|
|
||||||
|
# 下载转写结果 JSON
|
||||||
|
result_url = output.get("results", [{}])[0].get("transcription_url", "")
|
||||||
|
if not result_url:
|
||||||
|
raise RuntimeError(f"[FlashASR] 未找到 transcription_url: {output}")
|
||||||
|
|
||||||
|
async with httpx.AsyncClient(timeout=60) as client:
|
||||||
|
resp = await client.get(result_url)
|
||||||
|
resp.raise_for_status()
|
||||||
|
transcript_data = resp.json()
|
||||||
|
|
||||||
|
# 解析句子列表
|
||||||
|
sentences = transcript_data.get("transcripts", [{}])[0].get("sentences", [])
|
||||||
|
if not sentences:
|
||||||
|
# fallback:整段文本
|
||||||
|
text = transcript_data.get("transcripts", [{}])[0].get("text", "")
|
||||||
|
if text:
|
||||||
|
yield {"text": text, "begin_ms": 0, "end_ms": 0}
|
||||||
|
return
|
||||||
|
|
||||||
|
for s in sentences:
|
||||||
|
yield {
|
||||||
|
"text": s.get("text", "").strip(),
|
||||||
|
"begin_ms": int(s.get("begin_time", 0)),
|
||||||
|
"end_ms": int(s.get("end_time", 0)),
|
||||||
|
}
|
||||||
@@ -0,0 +1,261 @@
|
|||||||
|
"""Shared helper classes for gateway platform adapters.
|
||||||
|
|
||||||
|
Extracts common patterns that were duplicated across 5-7 adapters:
|
||||||
|
message deduplication, text batch aggregation, markdown stripping,
|
||||||
|
and thread participation tracking.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, Dict, Optional
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from gateway.platforms.base import BasePlatformAdapter, MessageEvent
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Message Deduplication ────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
class MessageDeduplicator:
|
||||||
|
"""TTL-based message deduplication cache.
|
||||||
|
|
||||||
|
Replaces the identical ``_seen_messages`` / ``_is_duplicate()`` pattern
|
||||||
|
previously duplicated in discord, slack, dingtalk, wecom, weixin,
|
||||||
|
mattermost, and feishu adapters.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
self._dedup = MessageDeduplicator()
|
||||||
|
|
||||||
|
# In message handler:
|
||||||
|
if self._dedup.is_duplicate(msg_id):
|
||||||
|
return
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, max_size: int = 2000, ttl_seconds: float = 300):
|
||||||
|
self._seen: Dict[str, float] = {}
|
||||||
|
self._max_size = max_size
|
||||||
|
self._ttl = ttl_seconds
|
||||||
|
|
||||||
|
def is_duplicate(self, msg_id: str) -> bool:
|
||||||
|
"""Return True if *msg_id* was already seen within the TTL window."""
|
||||||
|
if not msg_id:
|
||||||
|
return False
|
||||||
|
now = time.time()
|
||||||
|
if msg_id in self._seen:
|
||||||
|
return True
|
||||||
|
self._seen[msg_id] = now
|
||||||
|
if len(self._seen) > self._max_size:
|
||||||
|
cutoff = now - self._ttl
|
||||||
|
self._seen = {k: v for k, v in self._seen.items() if v > cutoff}
|
||||||
|
return False
|
||||||
|
|
||||||
|
def clear(self):
|
||||||
|
"""Clear all tracked messages."""
|
||||||
|
self._seen.clear()
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Text Batch Aggregation ──────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
class TextBatchAggregator:
|
||||||
|
"""Aggregates rapid-fire text events into single messages.
|
||||||
|
|
||||||
|
Replaces the ``_enqueue_text_event`` / ``_flush_text_batch`` pattern
|
||||||
|
previously duplicated in telegram, discord, matrix, wecom, and feishu.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
self._text_batcher = TextBatchAggregator(
|
||||||
|
handler=self._message_handler,
|
||||||
|
batch_delay=0.6,
|
||||||
|
split_threshold=1900,
|
||||||
|
)
|
||||||
|
|
||||||
|
# In message dispatch:
|
||||||
|
if msg_type == MessageType.TEXT and self._text_batcher.is_enabled():
|
||||||
|
self._text_batcher.enqueue(event, session_key)
|
||||||
|
return
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
handler,
|
||||||
|
*,
|
||||||
|
batch_delay: float = 0.6,
|
||||||
|
split_delay: float = 2.0,
|
||||||
|
split_threshold: int = 4000,
|
||||||
|
):
|
||||||
|
self._handler = handler
|
||||||
|
self._batch_delay = batch_delay
|
||||||
|
self._split_delay = split_delay
|
||||||
|
self._split_threshold = split_threshold
|
||||||
|
self._pending: Dict[str, "MessageEvent"] = {}
|
||||||
|
self._pending_tasks: Dict[str, asyncio.Task] = {}
|
||||||
|
|
||||||
|
def is_enabled(self) -> bool:
|
||||||
|
"""Return True if batching is active (delay > 0)."""
|
||||||
|
return self._batch_delay > 0
|
||||||
|
|
||||||
|
def enqueue(self, event: "MessageEvent", key: str) -> None:
|
||||||
|
"""Add *event* to the pending batch for *key*."""
|
||||||
|
chunk_len = len(event.text or "")
|
||||||
|
existing = self._pending.get(key)
|
||||||
|
if not existing:
|
||||||
|
event._last_chunk_len = chunk_len # type: ignore[attr-defined]
|
||||||
|
self._pending[key] = event
|
||||||
|
else:
|
||||||
|
existing.text = f"{existing.text}\n{event.text}"
|
||||||
|
existing._last_chunk_len = chunk_len # type: ignore[attr-defined]
|
||||||
|
|
||||||
|
# Cancel prior flush timer, start a new one
|
||||||
|
prior = self._pending_tasks.get(key)
|
||||||
|
if prior and not prior.done():
|
||||||
|
prior.cancel()
|
||||||
|
self._pending_tasks[key] = asyncio.create_task(self._flush(key))
|
||||||
|
|
||||||
|
async def _flush(self, key: str) -> None:
|
||||||
|
"""Wait then dispatch the batched event for *key*."""
|
||||||
|
current_task = self._pending_tasks.get(key)
|
||||||
|
pending = self._pending.get(key)
|
||||||
|
last_len = getattr(pending, "_last_chunk_len", 0) if pending else 0
|
||||||
|
|
||||||
|
# Use longer delay when the last chunk looks like a split message
|
||||||
|
delay = self._split_delay if last_len >= self._split_threshold else self._batch_delay
|
||||||
|
await asyncio.sleep(delay)
|
||||||
|
|
||||||
|
event = self._pending.pop(key, None)
|
||||||
|
if event:
|
||||||
|
try:
|
||||||
|
await self._handler(event)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("[TextBatchAggregator] Error dispatching batched event for %s", key)
|
||||||
|
|
||||||
|
if self._pending_tasks.get(key) is current_task:
|
||||||
|
self._pending_tasks.pop(key, None)
|
||||||
|
|
||||||
|
def cancel_all(self) -> None:
|
||||||
|
"""Cancel all pending flush tasks."""
|
||||||
|
for task in self._pending_tasks.values():
|
||||||
|
if not task.done():
|
||||||
|
task.cancel()
|
||||||
|
self._pending_tasks.clear()
|
||||||
|
self._pending.clear()
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Markdown Stripping ──────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
# Pre-compiled regexes for performance
|
||||||
|
_RE_BOLD = re.compile(r"\*\*(.+?)\*\*", re.DOTALL)
|
||||||
|
_RE_ITALIC_STAR = re.compile(r"\*(.+?)\*", re.DOTALL)
|
||||||
|
_RE_BOLD_UNDER = re.compile(r"__(.+?)__", re.DOTALL)
|
||||||
|
_RE_ITALIC_UNDER = re.compile(r"_(.+?)_", re.DOTALL)
|
||||||
|
_RE_CODE_BLOCK = re.compile(r"```[a-zA-Z0-9_+-]*\n?")
|
||||||
|
_RE_INLINE_CODE = re.compile(r"`(.+?)`")
|
||||||
|
_RE_HEADING = re.compile(r"^#{1,6}\s+", re.MULTILINE)
|
||||||
|
_RE_LINK = re.compile(r"\[([^\]]+)\]\([^\)]+\)")
|
||||||
|
_RE_MULTI_NEWLINE = re.compile(r"\n{3,}")
|
||||||
|
|
||||||
|
|
||||||
|
def strip_markdown(text: str) -> str:
|
||||||
|
"""Strip markdown formatting for plain-text platforms (SMS, iMessage, etc.).
|
||||||
|
|
||||||
|
Replaces the identical ``_strip_markdown()`` functions previously
|
||||||
|
duplicated in sms.py, bluebubbles.py, and feishu.py.
|
||||||
|
"""
|
||||||
|
text = _RE_BOLD.sub(r"\1", text)
|
||||||
|
text = _RE_ITALIC_STAR.sub(r"\1", text)
|
||||||
|
text = _RE_BOLD_UNDER.sub(r"\1", text)
|
||||||
|
text = _RE_ITALIC_UNDER.sub(r"\1", text)
|
||||||
|
text = _RE_CODE_BLOCK.sub("", text)
|
||||||
|
text = _RE_INLINE_CODE.sub(r"\1", text)
|
||||||
|
text = _RE_HEADING.sub("", text)
|
||||||
|
text = _RE_LINK.sub(r"\1", text)
|
||||||
|
text = _RE_MULTI_NEWLINE.sub("\n\n", text)
|
||||||
|
return text.strip()
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Thread Participation Tracking ───────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
class ThreadParticipationTracker:
|
||||||
|
"""Persistent tracking of threads the bot has participated in.
|
||||||
|
|
||||||
|
Replaces the identical ``_load/_save_participated_threads`` +
|
||||||
|
``_mark_thread_participated`` pattern previously duplicated in
|
||||||
|
discord.py and matrix.py.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
self._threads = ThreadParticipationTracker("discord")
|
||||||
|
|
||||||
|
# Check membership:
|
||||||
|
if thread_id in self._threads:
|
||||||
|
...
|
||||||
|
|
||||||
|
# Mark participation:
|
||||||
|
self._threads.mark(thread_id)
|
||||||
|
"""
|
||||||
|
|
||||||
|
_MAX_TRACKED = 500
|
||||||
|
|
||||||
|
def __init__(self, platform_name: str, max_tracked: int = 500):
|
||||||
|
self._platform = platform_name
|
||||||
|
self._max_tracked = max_tracked
|
||||||
|
self._threads: set = self._load()
|
||||||
|
|
||||||
|
def _state_path(self) -> Path:
|
||||||
|
from hermes_constants import get_hermes_home
|
||||||
|
return get_hermes_home() / f"{self._platform}_threads.json"
|
||||||
|
|
||||||
|
def _load(self) -> set:
|
||||||
|
path = self._state_path()
|
||||||
|
if path.exists():
|
||||||
|
try:
|
||||||
|
return set(json.loads(path.read_text(encoding="utf-8")))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return set()
|
||||||
|
|
||||||
|
def _save(self) -> None:
|
||||||
|
path = self._state_path()
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
thread_list = list(self._threads)
|
||||||
|
if len(thread_list) > self._max_tracked:
|
||||||
|
thread_list = thread_list[-self._max_tracked:]
|
||||||
|
self._threads = set(thread_list)
|
||||||
|
path.write_text(json.dumps(thread_list), encoding="utf-8")
|
||||||
|
|
||||||
|
def mark(self, thread_id: str) -> None:
|
||||||
|
"""Mark *thread_id* as participated and persist."""
|
||||||
|
if thread_id not in self._threads:
|
||||||
|
self._threads.add(thread_id)
|
||||||
|
self._save()
|
||||||
|
|
||||||
|
def __contains__(self, thread_id: str) -> bool:
|
||||||
|
return thread_id in self._threads
|
||||||
|
|
||||||
|
def clear(self) -> None:
|
||||||
|
self._threads.clear()
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Phone Number Redaction ──────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def redact_phone(phone: str) -> str:
|
||||||
|
"""Redact a phone number for logging, preserving country code and last 4.
|
||||||
|
|
||||||
|
Replaces the identical ``_redact_phone()`` functions in signal.py,
|
||||||
|
sms.py, and bluebubbles.py.
|
||||||
|
"""
|
||||||
|
if not phone:
|
||||||
|
return "<none>"
|
||||||
|
if len(phone) <= 8:
|
||||||
|
return phone[:2] + "****" + phone[-2:] if len(phone) > 4 else "****"
|
||||||
|
return phone[:4] + "****" + phone[-4:]
|
||||||
@@ -0,0 +1,449 @@
|
|||||||
|
"""
|
||||||
|
Home Assistant platform adapter.
|
||||||
|
|
||||||
|
Connects to the HA WebSocket API for real-time event monitoring.
|
||||||
|
State-change events are converted to MessageEvent objects and forwarded
|
||||||
|
to the agent for processing. Outbound messages are delivered as HA
|
||||||
|
persistent notifications.
|
||||||
|
|
||||||
|
Requires:
|
||||||
|
- aiohttp (already in messaging extras)
|
||||||
|
- HASS_TOKEN env var (Long-Lived Access Token)
|
||||||
|
- HASS_URL env var (default: http://homeassistant.local:8123)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Any, Dict, Optional, Set
|
||||||
|
|
||||||
|
try:
|
||||||
|
import aiohttp
|
||||||
|
AIOHTTP_AVAILABLE = True
|
||||||
|
except ImportError:
|
||||||
|
AIOHTTP_AVAILABLE = False
|
||||||
|
aiohttp = None # type: ignore[assignment]
|
||||||
|
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
from gateway.platforms.base import (
|
||||||
|
BasePlatformAdapter,
|
||||||
|
MessageEvent,
|
||||||
|
MessageType,
|
||||||
|
SendResult,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def check_ha_requirements() -> bool:
|
||||||
|
"""Check if Home Assistant dependencies are available and configured."""
|
||||||
|
if not AIOHTTP_AVAILABLE:
|
||||||
|
return False
|
||||||
|
if not os.getenv("HASS_TOKEN"):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
class HomeAssistantAdapter(BasePlatformAdapter):
|
||||||
|
"""
|
||||||
|
Home Assistant WebSocket adapter.
|
||||||
|
|
||||||
|
Subscribes to ``state_changed`` events and forwards them as
|
||||||
|
MessageEvent objects. Supports domain/entity filtering and
|
||||||
|
per-entity cooldowns to avoid event floods.
|
||||||
|
"""
|
||||||
|
|
||||||
|
MAX_MESSAGE_LENGTH = 4096
|
||||||
|
|
||||||
|
# Reconnection backoff schedule (seconds)
|
||||||
|
_BACKOFF_STEPS = [5, 10, 30, 60]
|
||||||
|
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.HOMEASSISTANT)
|
||||||
|
|
||||||
|
# Connection state
|
||||||
|
self._session: Optional["aiohttp.ClientSession"] = None
|
||||||
|
self._ws: Optional["aiohttp.ClientWebSocketResponse"] = None
|
||||||
|
self._rest_session: Optional["aiohttp.ClientSession"] = None
|
||||||
|
self._listen_task: Optional[asyncio.Task] = None
|
||||||
|
self._msg_id: int = 0
|
||||||
|
|
||||||
|
# Configuration from extra
|
||||||
|
extra = config.extra or {}
|
||||||
|
token = config.token or os.getenv("HASS_TOKEN", "")
|
||||||
|
url = extra.get("url") or os.getenv("HASS_URL", "http://homeassistant.local:8123")
|
||||||
|
self._hass_url: str = url.rstrip("/")
|
||||||
|
self._hass_token: str = token
|
||||||
|
|
||||||
|
# Event filtering
|
||||||
|
self._watch_domains: Set[str] = set(extra.get("watch_domains", []))
|
||||||
|
self._watch_entities: Set[str] = set(extra.get("watch_entities", []))
|
||||||
|
self._ignore_entities: Set[str] = set(extra.get("ignore_entities", []))
|
||||||
|
self._watch_all: bool = bool(extra.get("watch_all", False))
|
||||||
|
self._cooldown_seconds: int = int(extra.get("cooldown_seconds", 30))
|
||||||
|
|
||||||
|
# Cooldown tracking: entity_id -> last_event_timestamp
|
||||||
|
self._last_event_time: Dict[str, float] = {}
|
||||||
|
|
||||||
|
def _next_id(self) -> int:
|
||||||
|
"""Return the next WebSocket message ID."""
|
||||||
|
self._msg_id += 1
|
||||||
|
return self._msg_id
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Connection lifecycle
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
"""Connect to HA WebSocket API and subscribe to events."""
|
||||||
|
if not AIOHTTP_AVAILABLE:
|
||||||
|
logger.warning("[%s] aiohttp not installed. Run: pip install aiohttp", self.name)
|
||||||
|
return False
|
||||||
|
|
||||||
|
if not self._hass_token:
|
||||||
|
logger.warning("[%s] No HASS_TOKEN configured", self.name)
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
success = await self._ws_connect()
|
||||||
|
if not success:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Dedicated REST session for send() calls
|
||||||
|
self._rest_session = aiohttp.ClientSession(
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Warn if no event filters are configured
|
||||||
|
if not self._watch_domains and not self._watch_entities and not self._watch_all:
|
||||||
|
logger.warning(
|
||||||
|
"[%s] No watch_domains, watch_entities, or watch_all configured. "
|
||||||
|
"All state_changed events will be dropped. Configure filters in "
|
||||||
|
"your HA platform config to receive events.",
|
||||||
|
self.name,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Start background listener
|
||||||
|
self._listen_task = asyncio.create_task(self._listen_loop())
|
||||||
|
self._running = True
|
||||||
|
logger.info("[%s] Connected to %s", self.name, self._hass_url)
|
||||||
|
return True
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[%s] Failed to connect: %s", self.name, e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _ws_connect(self) -> bool:
|
||||||
|
"""Establish WebSocket connection and authenticate."""
|
||||||
|
ws_url = self._hass_url.replace("http://", "ws://").replace("https://", "wss://")
|
||||||
|
ws_url = f"{ws_url}/api/websocket"
|
||||||
|
|
||||||
|
self._session = aiohttp.ClientSession(
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30)
|
||||||
|
)
|
||||||
|
self._ws = await self._session.ws_connect(ws_url, heartbeat=30, timeout=30)
|
||||||
|
|
||||||
|
# Step 1: Receive auth_required
|
||||||
|
msg = await self._ws.receive_json()
|
||||||
|
if msg.get("type") != "auth_required":
|
||||||
|
logger.error("Expected auth_required, got: %s", msg.get("type"))
|
||||||
|
await self._cleanup_ws()
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Step 2: Send auth
|
||||||
|
await self._ws.send_json({
|
||||||
|
"type": "auth",
|
||||||
|
"access_token": self._hass_token,
|
||||||
|
})
|
||||||
|
|
||||||
|
# Step 3: Wait for auth_ok
|
||||||
|
msg = await self._ws.receive_json()
|
||||||
|
if msg.get("type") != "auth_ok":
|
||||||
|
logger.error("Auth failed: %s", msg)
|
||||||
|
await self._cleanup_ws()
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Step 4: Subscribe to state_changed events
|
||||||
|
sub_id = self._next_id()
|
||||||
|
await self._ws.send_json({
|
||||||
|
"id": sub_id,
|
||||||
|
"type": "subscribe_events",
|
||||||
|
"event_type": "state_changed",
|
||||||
|
})
|
||||||
|
|
||||||
|
# Verify subscription acknowledgement
|
||||||
|
msg = await self._ws.receive_json()
|
||||||
|
if not msg.get("success"):
|
||||||
|
logger.error("Failed to subscribe to events: %s", msg)
|
||||||
|
await self._cleanup_ws()
|
||||||
|
return False
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def _cleanup_ws(self) -> None:
|
||||||
|
"""Close WebSocket and session."""
|
||||||
|
if self._ws and not self._ws.closed:
|
||||||
|
await self._ws.close()
|
||||||
|
self._ws = None
|
||||||
|
if self._session and not self._session.closed:
|
||||||
|
await self._session.close()
|
||||||
|
self._session = None
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
"""Disconnect from Home Assistant."""
|
||||||
|
self._running = False
|
||||||
|
if self._listen_task:
|
||||||
|
self._listen_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._listen_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
self._listen_task = None
|
||||||
|
|
||||||
|
await self._cleanup_ws()
|
||||||
|
if self._rest_session and not self._rest_session.closed:
|
||||||
|
await self._rest_session.close()
|
||||||
|
self._rest_session = None
|
||||||
|
logger.info("[%s] Disconnected", self.name)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Event listener
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _listen_loop(self) -> None:
|
||||||
|
"""Main event loop with automatic reconnection."""
|
||||||
|
backoff_idx = 0
|
||||||
|
|
||||||
|
while self._running:
|
||||||
|
try:
|
||||||
|
await self._read_events()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
return
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[%s] WebSocket error: %s", self.name, e)
|
||||||
|
|
||||||
|
if not self._running:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Reconnect with backoff
|
||||||
|
delay = self._BACKOFF_STEPS[min(backoff_idx, len(self._BACKOFF_STEPS) - 1)]
|
||||||
|
logger.info("[%s] Reconnecting in %ds...", self.name, delay)
|
||||||
|
await asyncio.sleep(delay)
|
||||||
|
backoff_idx += 1
|
||||||
|
|
||||||
|
try:
|
||||||
|
await self._cleanup_ws()
|
||||||
|
success = await self._ws_connect()
|
||||||
|
if success:
|
||||||
|
backoff_idx = 0 # Reset on successful reconnect
|
||||||
|
logger.info("[%s] Reconnected", self.name)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[%s] Reconnection failed: %s", self.name, e)
|
||||||
|
|
||||||
|
async def _read_events(self) -> None:
|
||||||
|
"""Read events from WebSocket until disconnected."""
|
||||||
|
if self._ws is None or self._ws.closed:
|
||||||
|
return
|
||||||
|
async for ws_msg in self._ws:
|
||||||
|
if ws_msg.type == aiohttp.WSMsgType.TEXT:
|
||||||
|
try:
|
||||||
|
data = json.loads(ws_msg.data)
|
||||||
|
if data.get("type") == "event":
|
||||||
|
await self._handle_ha_event(data.get("event", {}))
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
logger.debug("Invalid JSON from HA WS: %s", ws_msg.data[:200])
|
||||||
|
elif ws_msg.type in (aiohttp.WSMsgType.CLOSED, aiohttp.WSMsgType.ERROR):
|
||||||
|
break
|
||||||
|
|
||||||
|
async def _handle_ha_event(self, event: Dict[str, Any]) -> None:
|
||||||
|
"""Process a state_changed event from Home Assistant."""
|
||||||
|
event_data = event.get("data", {})
|
||||||
|
entity_id: str = event_data.get("entity_id", "")
|
||||||
|
|
||||||
|
if not entity_id:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Apply ignore filter
|
||||||
|
if entity_id in self._ignore_entities:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Apply domain/entity watch filters (closed by default — require
|
||||||
|
# explicit watch_domains, watch_entities, or watch_all to forward)
|
||||||
|
domain = entity_id.split(".")[0] if "." in entity_id else ""
|
||||||
|
if self._watch_domains or self._watch_entities:
|
||||||
|
domain_match = domain in self._watch_domains if self._watch_domains else False
|
||||||
|
entity_match = entity_id in self._watch_entities if self._watch_entities else False
|
||||||
|
if not domain_match and not entity_match:
|
||||||
|
return
|
||||||
|
elif not self._watch_all:
|
||||||
|
# No filters configured and watch_all is off — drop the event
|
||||||
|
return
|
||||||
|
|
||||||
|
# Apply cooldown
|
||||||
|
now = time.time()
|
||||||
|
last = self._last_event_time.get(entity_id, 0)
|
||||||
|
if (now - last) < self._cooldown_seconds:
|
||||||
|
return
|
||||||
|
self._last_event_time[entity_id] = now
|
||||||
|
|
||||||
|
# Build human-readable message
|
||||||
|
old_state = event_data.get("old_state", {})
|
||||||
|
new_state = event_data.get("new_state", {})
|
||||||
|
message = self._format_state_change(entity_id, old_state, new_state)
|
||||||
|
|
||||||
|
if not message:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Build MessageEvent and forward to handler
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id="ha_events",
|
||||||
|
chat_name="Home Assistant Events",
|
||||||
|
chat_type="channel",
|
||||||
|
user_id="homeassistant",
|
||||||
|
user_name="Home Assistant",
|
||||||
|
)
|
||||||
|
|
||||||
|
msg_event = MessageEvent(
|
||||||
|
text=message,
|
||||||
|
message_type=MessageType.TEXT,
|
||||||
|
source=source,
|
||||||
|
message_id=f"ha_{entity_id}_{int(now)}",
|
||||||
|
timestamp=datetime.now(),
|
||||||
|
)
|
||||||
|
|
||||||
|
await self.handle_message(msg_event)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _format_state_change(
|
||||||
|
entity_id: str,
|
||||||
|
old_state: Dict[str, Any],
|
||||||
|
new_state: Dict[str, Any],
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Convert a state_changed event into a human-readable description."""
|
||||||
|
if not new_state:
|
||||||
|
return None
|
||||||
|
|
||||||
|
old_val = old_state.get("state", "unknown") if old_state else "unknown"
|
||||||
|
new_val = new_state.get("state", "unknown")
|
||||||
|
|
||||||
|
# Skip if state didn't actually change
|
||||||
|
if old_val == new_val:
|
||||||
|
return None
|
||||||
|
|
||||||
|
friendly_name = new_state.get("attributes", {}).get("friendly_name", entity_id)
|
||||||
|
domain = entity_id.split(".")[0] if "." in entity_id else ""
|
||||||
|
|
||||||
|
# Domain-specific formatting
|
||||||
|
if domain == "climate":
|
||||||
|
attrs = new_state.get("attributes", {})
|
||||||
|
temp = attrs.get("current_temperature", "?")
|
||||||
|
target = attrs.get("temperature", "?")
|
||||||
|
return (
|
||||||
|
f"[Home Assistant] {friendly_name}: HVAC mode changed from "
|
||||||
|
f"'{old_val}' to '{new_val}' (current: {temp}, target: {target})"
|
||||||
|
)
|
||||||
|
|
||||||
|
if domain == "sensor":
|
||||||
|
unit = new_state.get("attributes", {}).get("unit_of_measurement", "")
|
||||||
|
return (
|
||||||
|
f"[Home Assistant] {friendly_name}: changed from "
|
||||||
|
f"{old_val}{unit} to {new_val}{unit}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if domain == "binary_sensor":
|
||||||
|
return (
|
||||||
|
f"[Home Assistant] {friendly_name}: "
|
||||||
|
f"{'triggered' if new_val == 'on' else 'cleared'} "
|
||||||
|
f"(was {'triggered' if old_val == 'on' else 'cleared'})"
|
||||||
|
)
|
||||||
|
|
||||||
|
if domain in ("light", "switch", "fan"):
|
||||||
|
return (
|
||||||
|
f"[Home Assistant] {friendly_name}: turned "
|
||||||
|
f"{'on' if new_val == 'on' else 'off'}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if domain == "alarm_control_panel":
|
||||||
|
return (
|
||||||
|
f"[Home Assistant] {friendly_name}: alarm state changed from "
|
||||||
|
f"'{old_val}' to '{new_val}'"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Generic fallback
|
||||||
|
return (
|
||||||
|
f"[Home Assistant] {friendly_name} ({entity_id}): "
|
||||||
|
f"changed from '{old_val}' to '{new_val}'"
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Outbound messaging
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a notification via HA REST API (persistent_notification.create).
|
||||||
|
|
||||||
|
Uses the REST API instead of WebSocket to avoid a race condition
|
||||||
|
with the event listener loop that reads from the same WS connection.
|
||||||
|
"""
|
||||||
|
url = f"{self._hass_url}/api/services/persistent_notification/create"
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {self._hass_token}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
}
|
||||||
|
payload = {
|
||||||
|
"title": "Hermes Agent",
|
||||||
|
"message": content[:self.MAX_MESSAGE_LENGTH],
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
if self._rest_session:
|
||||||
|
async with self._rest_session.post(
|
||||||
|
url,
|
||||||
|
headers=headers,
|
||||||
|
json=payload,
|
||||||
|
timeout=aiohttp.ClientTimeout(total=10),
|
||||||
|
) as resp:
|
||||||
|
if resp.status < 300:
|
||||||
|
return SendResult(success=True, message_id=uuid.uuid4().hex[:12])
|
||||||
|
else:
|
||||||
|
body = await resp.text()
|
||||||
|
return SendResult(success=False, error=f"HTTP {resp.status}: {body}")
|
||||||
|
else:
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
async with session.post(
|
||||||
|
url,
|
||||||
|
headers=headers,
|
||||||
|
json=payload,
|
||||||
|
timeout=aiohttp.ClientTimeout(total=10),
|
||||||
|
) as resp:
|
||||||
|
if resp.status < 300:
|
||||||
|
return SendResult(success=True, message_id=uuid.uuid4().hex[:12])
|
||||||
|
else:
|
||||||
|
body = await resp.text()
|
||||||
|
return SendResult(success=False, error=f"HTTP {resp.status}: {body}")
|
||||||
|
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
return SendResult(success=False, error="Timeout sending notification to HA")
|
||||||
|
except Exception as e:
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
async def send_typing(self, chat_id: str, metadata=None) -> None:
|
||||||
|
"""No typing indicator for Home Assistant."""
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
"""Return basic info about the HA event channel."""
|
||||||
|
return {
|
||||||
|
"name": "Home Assistant Events",
|
||||||
|
"type": "channel",
|
||||||
|
"url": self._hass_url,
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,733 @@
|
|||||||
|
"""Mattermost gateway adapter.
|
||||||
|
|
||||||
|
Connects to a self-hosted (or cloud) Mattermost instance via its REST API
|
||||||
|
(v4) and WebSocket for real-time events. No external Mattermost library
|
||||||
|
required — uses aiohttp which is already a Hermes dependency.
|
||||||
|
|
||||||
|
Environment variables:
|
||||||
|
MATTERMOST_URL Server URL (e.g. https://mm.example.com)
|
||||||
|
MATTERMOST_TOKEN Bot token or personal-access token
|
||||||
|
MATTERMOST_ALLOWED_USERS Comma-separated user IDs
|
||||||
|
MATTERMOST_HOME_CHANNEL Channel ID for cron/notification delivery
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
from gateway.platforms.helpers import MessageDeduplicator
|
||||||
|
from gateway.platforms.base import (
|
||||||
|
BasePlatformAdapter,
|
||||||
|
MessageEvent,
|
||||||
|
MessageType,
|
||||||
|
SendResult,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Mattermost post size limit (server default is 16383, but 4000 is the
|
||||||
|
# practical limit for readable messages — matching OpenClaw's choice).
|
||||||
|
MAX_POST_LENGTH = 4000
|
||||||
|
|
||||||
|
# Channel type codes returned by the Mattermost API.
|
||||||
|
_CHANNEL_TYPE_MAP = {
|
||||||
|
"D": "dm",
|
||||||
|
"G": "group",
|
||||||
|
"P": "group", # private channel → treat as group
|
||||||
|
"O": "channel",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Reconnect parameters (exponential backoff).
|
||||||
|
_RECONNECT_BASE_DELAY = 2.0
|
||||||
|
_RECONNECT_MAX_DELAY = 60.0
|
||||||
|
_RECONNECT_JITTER = 0.2
|
||||||
|
|
||||||
|
|
||||||
|
def check_mattermost_requirements() -> bool:
|
||||||
|
"""Return True if the Mattermost adapter can be used."""
|
||||||
|
token = os.getenv("MATTERMOST_TOKEN", "")
|
||||||
|
url = os.getenv("MATTERMOST_URL", "")
|
||||||
|
if not token:
|
||||||
|
logger.debug("Mattermost: MATTERMOST_TOKEN not set")
|
||||||
|
return False
|
||||||
|
if not url:
|
||||||
|
logger.warning("Mattermost: MATTERMOST_URL not set")
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
import aiohttp # noqa: F401
|
||||||
|
return True
|
||||||
|
except ImportError:
|
||||||
|
logger.warning("Mattermost: aiohttp not installed")
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class MattermostAdapter(BasePlatformAdapter):
|
||||||
|
"""Gateway adapter for Mattermost (self-hosted or cloud)."""
|
||||||
|
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.MATTERMOST)
|
||||||
|
|
||||||
|
self._base_url: str = (
|
||||||
|
config.extra.get("url", "")
|
||||||
|
or os.getenv("MATTERMOST_URL", "")
|
||||||
|
).rstrip("/")
|
||||||
|
self._token: str = config.token or os.getenv("MATTERMOST_TOKEN", "")
|
||||||
|
|
||||||
|
self._bot_user_id: str = ""
|
||||||
|
self._bot_username: str = ""
|
||||||
|
|
||||||
|
# aiohttp session + websocket handle
|
||||||
|
self._session: Any = None # aiohttp.ClientSession
|
||||||
|
self._ws: Any = None # aiohttp.ClientWebSocketResponse
|
||||||
|
self._ws_task: Optional[asyncio.Task] = None
|
||||||
|
self._reconnect_task: Optional[asyncio.Task] = None
|
||||||
|
self._closing = False
|
||||||
|
|
||||||
|
# Reply mode: "thread" to nest replies, "off" for flat messages.
|
||||||
|
self._reply_mode: str = (
|
||||||
|
config.extra.get("reply_mode", "")
|
||||||
|
or os.getenv("MATTERMOST_REPLY_MODE", "off")
|
||||||
|
).lower()
|
||||||
|
|
||||||
|
# Dedup cache (prevent reprocessing)
|
||||||
|
self._dedup = MessageDeduplicator()
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# HTTP helpers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _headers(self) -> Dict[str, str]:
|
||||||
|
return {
|
||||||
|
"Authorization": f"Bearer {self._token}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
}
|
||||||
|
|
||||||
|
async def _api_get(self, path: str) -> Dict[str, Any]:
|
||||||
|
"""GET /api/v4/{path}."""
|
||||||
|
import aiohttp
|
||||||
|
url = f"{self._base_url}/api/v4/{path.lstrip('/')}"
|
||||||
|
try:
|
||||||
|
async with self._session.get(url, headers=self._headers(), timeout=aiohttp.ClientTimeout(total=30)) as resp:
|
||||||
|
if resp.status >= 400:
|
||||||
|
body = await resp.text()
|
||||||
|
logger.error("MM API GET %s → %s: %s", path, resp.status, body[:200])
|
||||||
|
return {}
|
||||||
|
return await resp.json()
|
||||||
|
except aiohttp.ClientError as exc:
|
||||||
|
logger.error("MM API GET %s network error: %s", path, exc)
|
||||||
|
return {}
|
||||||
|
|
||||||
|
async def _api_post(
|
||||||
|
self, path: str, payload: Dict[str, Any]
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""POST /api/v4/{path} with JSON body."""
|
||||||
|
import aiohttp
|
||||||
|
url = f"{self._base_url}/api/v4/{path.lstrip('/')}"
|
||||||
|
try:
|
||||||
|
async with self._session.post(
|
||||||
|
url, headers=self._headers(), json=payload,
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30)
|
||||||
|
) as resp:
|
||||||
|
if resp.status >= 400:
|
||||||
|
body = await resp.text()
|
||||||
|
logger.error("MM API POST %s → %s: %s", path, resp.status, body[:200])
|
||||||
|
return {}
|
||||||
|
return await resp.json()
|
||||||
|
except aiohttp.ClientError as exc:
|
||||||
|
logger.error("MM API POST %s network error: %s", path, exc)
|
||||||
|
return {}
|
||||||
|
|
||||||
|
async def _api_put(
|
||||||
|
self, path: str, payload: Dict[str, Any]
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""PUT /api/v4/{path} with JSON body."""
|
||||||
|
import aiohttp
|
||||||
|
url = f"{self._base_url}/api/v4/{path.lstrip('/')}"
|
||||||
|
try:
|
||||||
|
async with self._session.put(
|
||||||
|
url, headers=self._headers(), json=payload
|
||||||
|
) as resp:
|
||||||
|
if resp.status >= 400:
|
||||||
|
body = await resp.text()
|
||||||
|
logger.error("MM API PUT %s → %s: %s", path, resp.status, body[:200])
|
||||||
|
return {}
|
||||||
|
return await resp.json()
|
||||||
|
except aiohttp.ClientError as exc:
|
||||||
|
logger.error("MM API PUT %s network error: %s", path, exc)
|
||||||
|
return {}
|
||||||
|
|
||||||
|
async def _upload_file(
|
||||||
|
self, channel_id: str, file_data: bytes, filename: str, content_type: str = "application/octet-stream"
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Upload a file and return its file ID, or None on failure."""
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
url = f"{self._base_url}/api/v4/files"
|
||||||
|
form = aiohttp.FormData()
|
||||||
|
form.add_field("channel_id", channel_id)
|
||||||
|
form.add_field(
|
||||||
|
"files",
|
||||||
|
file_data,
|
||||||
|
filename=filename,
|
||||||
|
content_type=content_type,
|
||||||
|
)
|
||||||
|
headers = {"Authorization": f"Bearer {self._token}"}
|
||||||
|
async with self._session.post(url, headers=headers, data=form, timeout=aiohttp.ClientTimeout(total=60)) as resp:
|
||||||
|
if resp.status >= 400:
|
||||||
|
body = await resp.text()
|
||||||
|
logger.error("MM file upload → %s: %s", resp.status, body[:200])
|
||||||
|
return None
|
||||||
|
data = await resp.json()
|
||||||
|
infos = data.get("file_infos", [])
|
||||||
|
return infos[0]["id"] if infos else None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Required overrides
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
"""Connect to Mattermost and start the WebSocket listener."""
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
if not self._base_url or not self._token:
|
||||||
|
logger.error("Mattermost: URL or token not configured")
|
||||||
|
return False
|
||||||
|
|
||||||
|
self._session = aiohttp.ClientSession(
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30)
|
||||||
|
)
|
||||||
|
self._closing = False
|
||||||
|
|
||||||
|
# Verify credentials and fetch bot identity.
|
||||||
|
me = await self._api_get("users/me")
|
||||||
|
if not me or "id" not in me:
|
||||||
|
logger.error("Mattermost: failed to authenticate — check MATTERMOST_TOKEN and MATTERMOST_URL")
|
||||||
|
await self._session.close()
|
||||||
|
return False
|
||||||
|
|
||||||
|
self._bot_user_id = me["id"]
|
||||||
|
self._bot_username = me.get("username", "")
|
||||||
|
logger.info(
|
||||||
|
"Mattermost: authenticated as @%s (%s) on %s",
|
||||||
|
self._bot_username,
|
||||||
|
self._bot_user_id,
|
||||||
|
self._base_url,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Start WebSocket in background.
|
||||||
|
self._ws_task = asyncio.create_task(self._ws_loop())
|
||||||
|
self._mark_connected()
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
"""Disconnect from Mattermost."""
|
||||||
|
self._closing = True
|
||||||
|
|
||||||
|
if self._ws_task and not self._ws_task.done():
|
||||||
|
self._ws_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._ws_task
|
||||||
|
except (asyncio.CancelledError, Exception):
|
||||||
|
pass
|
||||||
|
|
||||||
|
if self._reconnect_task and not self._reconnect_task.done():
|
||||||
|
self._reconnect_task.cancel()
|
||||||
|
|
||||||
|
if self._ws:
|
||||||
|
await self._ws.close()
|
||||||
|
self._ws = None
|
||||||
|
|
||||||
|
if self._session and not self._session.closed:
|
||||||
|
await self._session.close()
|
||||||
|
|
||||||
|
logger.info("Mattermost: disconnected")
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a message (or multiple chunks) to a channel."""
|
||||||
|
if not content:
|
||||||
|
return SendResult(success=True)
|
||||||
|
|
||||||
|
formatted = self.format_message(content)
|
||||||
|
chunks = self.truncate_message(formatted, MAX_POST_LENGTH)
|
||||||
|
|
||||||
|
last_id = None
|
||||||
|
for chunk in chunks:
|
||||||
|
payload: Dict[str, Any] = {
|
||||||
|
"channel_id": chat_id,
|
||||||
|
"message": chunk,
|
||||||
|
}
|
||||||
|
# Thread support: reply_to is the root post ID.
|
||||||
|
if reply_to and self._reply_mode == "thread":
|
||||||
|
payload["root_id"] = reply_to
|
||||||
|
|
||||||
|
data = await self._api_post("posts", payload)
|
||||||
|
if not data or "id" not in data:
|
||||||
|
return SendResult(success=False, error="Failed to create post")
|
||||||
|
last_id = data["id"]
|
||||||
|
|
||||||
|
return SendResult(success=True, message_id=last_id)
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
"""Return channel name and type."""
|
||||||
|
data = await self._api_get(f"channels/{chat_id}")
|
||||||
|
if not data:
|
||||||
|
return {"name": chat_id, "type": "channel"}
|
||||||
|
|
||||||
|
ch_type = _CHANNEL_TYPE_MAP.get(data.get("type", "O"), "channel")
|
||||||
|
display_name = data.get("display_name") or data.get("name") or chat_id
|
||||||
|
return {"name": display_name, "type": ch_type}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Optional overrides
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def send_typing(
|
||||||
|
self, chat_id: str, metadata: Optional[Dict[str, Any]] = None
|
||||||
|
) -> None:
|
||||||
|
"""Send a typing indicator."""
|
||||||
|
await self._api_post(
|
||||||
|
f"users/{self._bot_user_id}/typing",
|
||||||
|
{"channel_id": chat_id},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def edit_message(
|
||||||
|
self, chat_id: str, message_id: str, content: str
|
||||||
|
) -> SendResult:
|
||||||
|
"""Edit an existing post."""
|
||||||
|
formatted = self.format_message(content)
|
||||||
|
data = await self._api_put(
|
||||||
|
f"posts/{message_id}/patch",
|
||||||
|
{"message": formatted},
|
||||||
|
)
|
||||||
|
if not data or "id" not in data:
|
||||||
|
return SendResult(success=False, error="Failed to edit post")
|
||||||
|
return SendResult(success=True, message_id=data["id"])
|
||||||
|
|
||||||
|
async def send_image(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
image_url: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Download an image and upload it as a file attachment."""
|
||||||
|
return await self._send_url_as_file(
|
||||||
|
chat_id, image_url, caption, reply_to, "image"
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_image_file(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
image_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Upload a local image file."""
|
||||||
|
return await self._send_local_file(
|
||||||
|
chat_id, image_path, caption, reply_to
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_document(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
file_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
file_name: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Upload a local file as a document."""
|
||||||
|
return await self._send_local_file(
|
||||||
|
chat_id, file_path, caption, reply_to, file_name
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_voice(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
audio_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Upload an audio file."""
|
||||||
|
return await self._send_local_file(
|
||||||
|
chat_id, audio_path, caption, reply_to
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_video(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
video_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Upload a video file."""
|
||||||
|
return await self._send_local_file(
|
||||||
|
chat_id, video_path, caption, reply_to
|
||||||
|
)
|
||||||
|
|
||||||
|
def format_message(self, content: str) -> str:
|
||||||
|
"""Mattermost uses standard Markdown — mostly pass through.
|
||||||
|
|
||||||
|
Strip image markdown into plain links (files are uploaded separately).
|
||||||
|
"""
|
||||||
|
# Convert  to just the URL — Mattermost renders
|
||||||
|
# image URLs as inline previews automatically.
|
||||||
|
content = re.sub(r"!\[([^\]]*)\]\(([^)]+)\)", r"\2", content)
|
||||||
|
return content
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# File helpers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _send_url_as_file(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
url: str,
|
||||||
|
caption: Optional[str],
|
||||||
|
reply_to: Optional[str],
|
||||||
|
kind: str = "file",
|
||||||
|
) -> SendResult:
|
||||||
|
"""Download a URL and upload it as a file attachment."""
|
||||||
|
from tools.url_safety import is_safe_url
|
||||||
|
if not is_safe_url(url):
|
||||||
|
logger.warning("Mattermost: blocked unsafe URL (SSRF protection)")
|
||||||
|
return await self.send(chat_id, f"{caption or ''}\n{url}".strip(), reply_to)
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
last_exc = None
|
||||||
|
file_data = None
|
||||||
|
ct = "application/octet-stream"
|
||||||
|
fname = url.rsplit("/", 1)[-1].split("?")[0] or f"{kind}.png"
|
||||||
|
|
||||||
|
for attempt in range(3):
|
||||||
|
try:
|
||||||
|
async with self._session.get(url, timeout=aiohttp.ClientTimeout(total=30)) as resp:
|
||||||
|
if resp.status >= 500 or resp.status == 429:
|
||||||
|
if attempt < 2:
|
||||||
|
logger.debug("Mattermost download retry %d/2 for %s (status %d)",
|
||||||
|
attempt + 1, url[:80], resp.status)
|
||||||
|
await asyncio.sleep(1.5 * (attempt + 1))
|
||||||
|
continue
|
||||||
|
if resp.status >= 400:
|
||||||
|
return await self.send(chat_id, f"{caption or ''}\n{url}".strip(), reply_to)
|
||||||
|
file_data = await resp.read()
|
||||||
|
ct = resp.content_type or "application/octet-stream"
|
||||||
|
break
|
||||||
|
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
||||||
|
if attempt < 2:
|
||||||
|
await asyncio.sleep(1.5 * (attempt + 1))
|
||||||
|
continue
|
||||||
|
logger.warning("Mattermost: failed to download %s after %d attempts: %s", url, attempt + 1, exc)
|
||||||
|
return await self.send(chat_id, f"{caption or ''}\n{url}".strip(), reply_to)
|
||||||
|
|
||||||
|
if file_data is None:
|
||||||
|
logger.warning("Mattermost: download returned no data for %s", url)
|
||||||
|
return await self.send(chat_id, f"{caption or ''}\n{url}".strip(), reply_to)
|
||||||
|
|
||||||
|
file_id = await self._upload_file(chat_id, file_data, fname, ct)
|
||||||
|
if not file_id:
|
||||||
|
return await self.send(chat_id, f"{caption or ''}\n{url}".strip(), reply_to)
|
||||||
|
|
||||||
|
payload: Dict[str, Any] = {
|
||||||
|
"channel_id": chat_id,
|
||||||
|
"message": caption or "",
|
||||||
|
"file_ids": [file_id],
|
||||||
|
}
|
||||||
|
if reply_to and self._reply_mode == "thread":
|
||||||
|
payload["root_id"] = reply_to
|
||||||
|
|
||||||
|
data = await self._api_post("posts", payload)
|
||||||
|
if not data or "id" not in data:
|
||||||
|
return SendResult(success=False, error="Failed to post with file")
|
||||||
|
return SendResult(success=True, message_id=data["id"])
|
||||||
|
|
||||||
|
async def _send_local_file(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
file_path: str,
|
||||||
|
caption: Optional[str],
|
||||||
|
reply_to: Optional[str],
|
||||||
|
file_name: Optional[str] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Upload a local file and attach it to a post."""
|
||||||
|
import mimetypes
|
||||||
|
|
||||||
|
p = Path(file_path)
|
||||||
|
if not p.exists():
|
||||||
|
return await self.send(
|
||||||
|
chat_id, f"{caption or ''}\n(file not found: {file_path})", reply_to
|
||||||
|
)
|
||||||
|
|
||||||
|
fname = file_name or p.name
|
||||||
|
ct = mimetypes.guess_type(fname)[0] or "application/octet-stream"
|
||||||
|
file_data = p.read_bytes()
|
||||||
|
|
||||||
|
file_id = await self._upload_file(chat_id, file_data, fname, ct)
|
||||||
|
if not file_id:
|
||||||
|
return SendResult(success=False, error="File upload failed")
|
||||||
|
|
||||||
|
payload: Dict[str, Any] = {
|
||||||
|
"channel_id": chat_id,
|
||||||
|
"message": caption or "",
|
||||||
|
"file_ids": [file_id],
|
||||||
|
}
|
||||||
|
if reply_to and self._reply_mode == "thread":
|
||||||
|
payload["root_id"] = reply_to
|
||||||
|
|
||||||
|
data = await self._api_post("posts", payload)
|
||||||
|
if not data or "id" not in data:
|
||||||
|
return SendResult(success=False, error="Failed to post with file")
|
||||||
|
return SendResult(success=True, message_id=data["id"])
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# WebSocket
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _ws_loop(self) -> None:
|
||||||
|
"""Connect to the WebSocket and listen for events, reconnecting on failure."""
|
||||||
|
delay = _RECONNECT_BASE_DELAY
|
||||||
|
while not self._closing:
|
||||||
|
try:
|
||||||
|
await self._ws_connect_and_listen()
|
||||||
|
# Clean disconnect — reset delay.
|
||||||
|
delay = _RECONNECT_BASE_DELAY
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
return
|
||||||
|
except Exception as exc:
|
||||||
|
if self._closing:
|
||||||
|
return
|
||||||
|
# Detect permanent auth/permission failures that will never
|
||||||
|
# succeed on retry — stop reconnecting instead of looping forever.
|
||||||
|
import aiohttp
|
||||||
|
err_str = str(exc).lower()
|
||||||
|
if isinstance(exc, aiohttp.WSServerHandshakeError) and exc.status in (401, 403):
|
||||||
|
logger.error("Mattermost WS auth failed (HTTP %d) — stopping reconnect", exc.status)
|
||||||
|
return
|
||||||
|
if "401" in err_str or "403" in err_str or "unauthorized" in err_str:
|
||||||
|
logger.error("Mattermost WS permanent error: %s — stopping reconnect", exc)
|
||||||
|
return
|
||||||
|
logger.warning("Mattermost WS error: %s — reconnecting in %.0fs", exc, delay)
|
||||||
|
|
||||||
|
if self._closing:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Exponential backoff with jitter.
|
||||||
|
import random
|
||||||
|
jitter = delay * _RECONNECT_JITTER * random.random()
|
||||||
|
await asyncio.sleep(delay + jitter)
|
||||||
|
delay = min(delay * 2, _RECONNECT_MAX_DELAY)
|
||||||
|
|
||||||
|
async def _ws_connect_and_listen(self) -> None:
|
||||||
|
"""Single WebSocket session: connect, authenticate, process events."""
|
||||||
|
# Build WS URL: https:// → wss://, http:// → ws://
|
||||||
|
ws_url = re.sub(r"^http", "ws", self._base_url) + "/api/v4/websocket"
|
||||||
|
logger.info("Mattermost: connecting to %s", ws_url)
|
||||||
|
|
||||||
|
self._ws = await self._session.ws_connect(ws_url, heartbeat=30.0)
|
||||||
|
|
||||||
|
# Authenticate via the WebSocket.
|
||||||
|
auth_msg = {
|
||||||
|
"seq": 1,
|
||||||
|
"action": "authentication_challenge",
|
||||||
|
"data": {"token": self._token},
|
||||||
|
}
|
||||||
|
await self._ws.send_json(auth_msg)
|
||||||
|
logger.info("Mattermost: WebSocket connected and authenticated")
|
||||||
|
|
||||||
|
async for raw_msg in self._ws:
|
||||||
|
if self._closing:
|
||||||
|
return
|
||||||
|
|
||||||
|
if raw_msg.type in (
|
||||||
|
raw_msg.type.TEXT,
|
||||||
|
raw_msg.type.BINARY,
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
event = json.loads(raw_msg.data)
|
||||||
|
except (json.JSONDecodeError, TypeError):
|
||||||
|
continue
|
||||||
|
await self._handle_ws_event(event)
|
||||||
|
elif raw_msg.type in (
|
||||||
|
raw_msg.type.ERROR,
|
||||||
|
raw_msg.type.CLOSE,
|
||||||
|
raw_msg.type.CLOSING,
|
||||||
|
raw_msg.type.CLOSED,
|
||||||
|
):
|
||||||
|
logger.info("Mattermost: WebSocket closed (%s)", raw_msg.type)
|
||||||
|
break
|
||||||
|
|
||||||
|
async def _handle_ws_event(self, event: Dict[str, Any]) -> None:
|
||||||
|
"""Process a single WebSocket event."""
|
||||||
|
event_type = event.get("event")
|
||||||
|
if event_type != "posted":
|
||||||
|
return
|
||||||
|
|
||||||
|
data = event.get("data", {})
|
||||||
|
raw_post_str = data.get("post")
|
||||||
|
if not raw_post_str:
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
post = json.loads(raw_post_str)
|
||||||
|
except (json.JSONDecodeError, TypeError):
|
||||||
|
return
|
||||||
|
|
||||||
|
# Ignore own messages.
|
||||||
|
if post.get("user_id") == self._bot_user_id:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Ignore system posts.
|
||||||
|
if post.get("type"):
|
||||||
|
return
|
||||||
|
|
||||||
|
post_id = post.get("id", "")
|
||||||
|
|
||||||
|
# Dedup.
|
||||||
|
if self._dedup.is_duplicate(post_id):
|
||||||
|
return
|
||||||
|
|
||||||
|
# Build message event.
|
||||||
|
channel_id = post.get("channel_id", "")
|
||||||
|
channel_type_raw = data.get("channel_type", "O")
|
||||||
|
chat_type = _CHANNEL_TYPE_MAP.get(channel_type_raw, "channel")
|
||||||
|
|
||||||
|
# For DMs, user_id is sufficient. For channels, check for @mention.
|
||||||
|
message_text = post.get("message", "")
|
||||||
|
|
||||||
|
# Mention-gating for non-DM channels.
|
||||||
|
# Config (env vars):
|
||||||
|
# MATTERMOST_REQUIRE_MENTION: Require @mention in channels (default: true)
|
||||||
|
# MATTERMOST_FREE_RESPONSE_CHANNELS: Channel IDs where bot responds without mention
|
||||||
|
if channel_type_raw != "D":
|
||||||
|
require_mention = os.getenv(
|
||||||
|
"MATTERMOST_REQUIRE_MENTION", "true"
|
||||||
|
).lower() not in ("false", "0", "no")
|
||||||
|
|
||||||
|
free_channels_raw = os.getenv("MATTERMOST_FREE_RESPONSE_CHANNELS", "")
|
||||||
|
free_channels = {ch.strip() for ch in free_channels_raw.split(",") if ch.strip()}
|
||||||
|
is_free_channel = channel_id in free_channels
|
||||||
|
|
||||||
|
mention_patterns = [
|
||||||
|
f"@{self._bot_username}",
|
||||||
|
f"@{self._bot_user_id}",
|
||||||
|
]
|
||||||
|
has_mention = any(
|
||||||
|
pattern.lower() in message_text.lower()
|
||||||
|
for pattern in mention_patterns
|
||||||
|
)
|
||||||
|
|
||||||
|
if require_mention and not is_free_channel and not has_mention:
|
||||||
|
logger.debug(
|
||||||
|
"Mattermost: skipping non-DM message without @mention (channel=%s)",
|
||||||
|
channel_id,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Strip @mention from the message text so the agent sees clean input.
|
||||||
|
if has_mention:
|
||||||
|
for pattern in mention_patterns:
|
||||||
|
message_text = re.sub(
|
||||||
|
re.escape(pattern), "", message_text, flags=re.IGNORECASE
|
||||||
|
).strip()
|
||||||
|
|
||||||
|
# Resolve sender info.
|
||||||
|
sender_id = post.get("user_id", "")
|
||||||
|
sender_name = data.get("sender_name", "").lstrip("@") or sender_id
|
||||||
|
|
||||||
|
# Thread support: if the post is in a thread, use root_id.
|
||||||
|
thread_id = post.get("root_id") or None
|
||||||
|
|
||||||
|
# Determine message type.
|
||||||
|
file_ids = post.get("file_ids") or []
|
||||||
|
msg_type = MessageType.TEXT
|
||||||
|
if message_text.startswith("/"):
|
||||||
|
msg_type = MessageType.COMMAND
|
||||||
|
|
||||||
|
# Download file attachments immediately (URLs require auth headers
|
||||||
|
# that downstream tools won't have).
|
||||||
|
media_urls: List[str] = []
|
||||||
|
media_types: List[str] = []
|
||||||
|
for fid in file_ids:
|
||||||
|
try:
|
||||||
|
file_info = await self._api_get(f"files/{fid}/info")
|
||||||
|
fname = file_info.get("name", f"file_{fid}")
|
||||||
|
ext = Path(fname).suffix or ""
|
||||||
|
mime = file_info.get("mime_type", "application/octet-stream")
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
dl_url = f"{self._base_url}/api/v4/files/{fid}"
|
||||||
|
async with self._session.get(
|
||||||
|
dl_url,
|
||||||
|
headers={"Authorization": f"Bearer {self._token}"},
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30),
|
||||||
|
) as resp:
|
||||||
|
if resp.status < 400:
|
||||||
|
file_data = await resp.read()
|
||||||
|
from gateway.platforms.base import cache_image_from_bytes, cache_document_from_bytes
|
||||||
|
if mime.startswith("image/"):
|
||||||
|
local_path = cache_image_from_bytes(file_data, ext or ".png")
|
||||||
|
media_urls.append(local_path)
|
||||||
|
media_types.append(mime)
|
||||||
|
elif mime.startswith("audio/"):
|
||||||
|
from gateway.platforms.base import cache_audio_from_bytes
|
||||||
|
local_path = cache_audio_from_bytes(file_data, ext or ".ogg")
|
||||||
|
media_urls.append(local_path)
|
||||||
|
media_types.append(mime)
|
||||||
|
else:
|
||||||
|
local_path = cache_document_from_bytes(file_data, fname)
|
||||||
|
media_urls.append(local_path)
|
||||||
|
media_types.append(mime)
|
||||||
|
else:
|
||||||
|
logger.warning("Mattermost: failed to download file %s: HTTP %s", fid, resp.status)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Mattermost: error downloading file %s: %s", fid, exc)
|
||||||
|
|
||||||
|
# Set message type based on downloaded media types.
|
||||||
|
if media_types and msg_type == MessageType.TEXT:
|
||||||
|
if any(m.startswith("image/") for m in media_types):
|
||||||
|
msg_type = MessageType.PHOTO
|
||||||
|
elif any(m.startswith("audio/") for m in media_types):
|
||||||
|
msg_type = MessageType.VOICE
|
||||||
|
elif media_types:
|
||||||
|
msg_type = MessageType.DOCUMENT
|
||||||
|
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id=channel_id,
|
||||||
|
chat_type=chat_type,
|
||||||
|
user_id=sender_id,
|
||||||
|
user_name=sender_name,
|
||||||
|
thread_id=thread_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
msg_event = MessageEvent(
|
||||||
|
text=message_text,
|
||||||
|
message_type=msg_type,
|
||||||
|
source=source,
|
||||||
|
raw_message=post,
|
||||||
|
message_id=post_id,
|
||||||
|
media_urls=media_urls if media_urls else None,
|
||||||
|
media_types=media_types if media_types else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
await self.handle_message(msg_event)
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,113 @@
|
|||||||
|
"""
|
||||||
|
MindOS NEXT — mdConverter:transcript 片段 → Markdown 文件追加
|
||||||
|
|
||||||
|
封闭函数(铁律 1):
|
||||||
|
输入:transcript 文本片段(str)+ 可选毫秒偏移(int)
|
||||||
|
输出:追加到 wiki/meetings/*.md
|
||||||
|
副作用:仅磁盘文件写入;不调用 LLM,不写 DB。
|
||||||
|
|
||||||
|
MD-First(铁律 6):
|
||||||
|
音频/文件的唯一持久化形式就是这个 .md 文件。
|
||||||
|
无 DB 记录,无 UUID,文件名即身份。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
import threading
|
||||||
|
from datetime import datetime
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
|
||||||
|
_MIME_TO_EXT = {
|
||||||
|
"audio/mpeg": "mp3", "audio/mp4": "m4a", "audio/mp3": "mp3",
|
||||||
|
"audio/wav": "wav", "audio/webm": "webm", "audio/ogg": "ogg",
|
||||||
|
"audio/aac": "aac", "audio/flac": "flac",
|
||||||
|
}
|
||||||
|
|
||||||
|
_SUPPORTED_EXTS = {".mp3", ".mp4", ".m4a", ".wav", ".webm", ".ogg", ".aac", ".flac"}
|
||||||
|
|
||||||
|
|
||||||
|
def is_supported_audio(filename: str) -> bool:
|
||||||
|
return Path(filename).suffix.lower() in _SUPPORTED_EXTS
|
||||||
|
|
||||||
|
|
||||||
|
class MdConverter:
|
||||||
|
"""transcript 片段追加为 Markdown 文件(L0 原文层)。"""
|
||||||
|
|
||||||
|
def __init__(self, wiki_dir: str):
|
||||||
|
self.meetings_dir = Path(wiki_dir) / "raw"
|
||||||
|
self.meetings_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
def new_file(
|
||||||
|
self,
|
||||||
|
title: str = "未命名",
|
||||||
|
source_filename: str = "",
|
||||||
|
duration_hint: str = "",
|
||||||
|
) -> Path:
|
||||||
|
"""
|
||||||
|
创建新 MD 文件,写入 YAML-style header,返回文件路径。
|
||||||
|
文件名:YYYY-MM-DD_HHmm_<title>.md
|
||||||
|
"""
|
||||||
|
now = datetime.now()
|
||||||
|
safe_title = title.replace("/", "-").replace("\\", "-")[:20]
|
||||||
|
filename = f"{now.strftime('%Y-%m-%d_%H%M')}_{safe_title}.md"
|
||||||
|
file_path = self.meetings_dir / filename
|
||||||
|
|
||||||
|
source_line = f"来源:`{source_filename}` | " if source_filename else ""
|
||||||
|
duration_line = f"时长:{duration_hint} | " if duration_hint else ""
|
||||||
|
|
||||||
|
header = (
|
||||||
|
f"# {now.strftime('%Y-%m-%d %H:%M')} {title}\n\n"
|
||||||
|
f"> {source_line}{duration_line}"
|
||||||
|
f"转写模型:qwen3-asr-flash | 入库时间:{now.strftime('%Y-%m-%d %H:%M')}\n\n"
|
||||||
|
f"---\n\n"
|
||||||
|
)
|
||||||
|
with self._lock:
|
||||||
|
file_path.write_text(header, encoding="utf-8")
|
||||||
|
|
||||||
|
return file_path
|
||||||
|
|
||||||
|
def append_segment(
|
||||||
|
self,
|
||||||
|
file_path: Path,
|
||||||
|
text: str,
|
||||||
|
offset_ms: int = 0,
|
||||||
|
) -> int:
|
||||||
|
"""
|
||||||
|
追加一个 transcript 片段。
|
||||||
|
格式:[MM:SS] 文字内容
|
||||||
|
返回追加的字符数。
|
||||||
|
"""
|
||||||
|
text = text.strip()
|
||||||
|
if not text:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
if offset_ms > 0:
|
||||||
|
total_s = offset_ms // 1000
|
||||||
|
mm, ss = divmod(total_s, 60)
|
||||||
|
timestamp = f"[{mm:02d}:{ss:02d}] "
|
||||||
|
else:
|
||||||
|
timestamp = ""
|
||||||
|
|
||||||
|
line = f"{timestamp}{text}\n\n"
|
||||||
|
with self._lock:
|
||||||
|
with file_path.open("a", encoding="utf-8") as f:
|
||||||
|
f.write(line)
|
||||||
|
|
||||||
|
return len(text)
|
||||||
|
|
||||||
|
def finalize(self, file_path: Path, char_count: int = 0) -> None:
|
||||||
|
"""写入结束标记。"""
|
||||||
|
summary = f"\n---\n\n> ✅ 转写完成,共 {char_count} 字\n"
|
||||||
|
with self._lock:
|
||||||
|
with file_path.open("a", encoding="utf-8") as f:
|
||||||
|
f.write(summary)
|
||||||
|
|
||||||
|
def relative_path(self, file_path: Path) -> str:
|
||||||
|
"""返回相对于 meetings_dir 父目录的路径,用于日志/SSE。"""
|
||||||
|
try:
|
||||||
|
return str(file_path.relative_to(self.meetings_dir.parent))
|
||||||
|
except ValueError:
|
||||||
|
return str(file_path)
|
||||||
@@ -0,0 +1,249 @@
|
|||||||
|
"""
|
||||||
|
Mind CLI Bridge — Cloud 端 WebSocket Tunnel 管理器。
|
||||||
|
|
||||||
|
管理所有 Mind CLI 用户的 Tunnel 连接:
|
||||||
|
- 接收 CLI 的 WebSocket 握手 + JWT 认证
|
||||||
|
- 能力审计(审批通过的工具白名单)
|
||||||
|
- 工具调用派发(Cloud LLM → Tunnel → CLI → 结果)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
logger = logging.getLogger("mindcli_bridge")
|
||||||
|
|
||||||
|
# ── 初版白名单(所有内置工具默认通过) ────────────────────
|
||||||
|
DEFAULT_APPROVED_TOOLS = [
|
||||||
|
"terminal",
|
||||||
|
"file_read",
|
||||||
|
"file_write",
|
||||||
|
"file_ops",
|
||||||
|
"grep",
|
||||||
|
"code_execution",
|
||||||
|
]
|
||||||
|
|
||||||
|
# 工具调用超时
|
||||||
|
TOOL_CALL_TIMEOUT = 120 # 秒
|
||||||
|
# 心跳检测:60s 无消息判定离线
|
||||||
|
HEARTBEAT_TIMEOUT = 60
|
||||||
|
|
||||||
|
|
||||||
|
class MindCLIBridge:
|
||||||
|
"""
|
||||||
|
管理所有 CLI 用户的 Tunnel 连接。
|
||||||
|
|
||||||
|
每个 userId 最多维持一个活跃 WebSocket 连接。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
# userId → WebSocketResponse
|
||||||
|
self._connections: dict[str, web.WebSocketResponse] = {}
|
||||||
|
# userId → 能力报告
|
||||||
|
self._capabilities: dict[str, dict] = {}
|
||||||
|
# userId → 最后心跳时间
|
||||||
|
self._last_heartbeat: dict[str, float] = {}
|
||||||
|
# 待处理的 tool_call 请求
|
||||||
|
self._pending: dict[str, asyncio.Future] = {}
|
||||||
|
|
||||||
|
async def handle_tunnel_connect(self, request: web.Request) -> web.WebSocketResponse:
|
||||||
|
"""
|
||||||
|
WebSocket 握手入口。
|
||||||
|
|
||||||
|
路由:GET /mindos-next/ws/cli-tunnel
|
||||||
|
认证:Header Authorization: Bearer <JWT>
|
||||||
|
"""
|
||||||
|
ws = web.WebSocketResponse(heartbeat=30, autoping=True)
|
||||||
|
await ws.prepare(request)
|
||||||
|
|
||||||
|
# ── JWT 认证 ──
|
||||||
|
token = _extract_token(request)
|
||||||
|
if not token:
|
||||||
|
await ws.send_json({"type": "error", "message": "Missing Authorization header"})
|
||||||
|
await ws.close()
|
||||||
|
return ws
|
||||||
|
|
||||||
|
# 使用 MindPass 认证策略
|
||||||
|
user = await self._verify_token(token, request)
|
||||||
|
if not user:
|
||||||
|
await ws.send_json({"type": "error", "message": "Invalid or expired token"})
|
||||||
|
await ws.close()
|
||||||
|
return ws
|
||||||
|
|
||||||
|
user_id = user["userId"]
|
||||||
|
|
||||||
|
# 关闭旧连接
|
||||||
|
old_ws = self._connections.get(user_id)
|
||||||
|
if old_ws and not old_ws.closed:
|
||||||
|
logger.info("[Bridge] 关闭 %s 的旧连接", user_id)
|
||||||
|
await old_ws.close()
|
||||||
|
|
||||||
|
self._connections[user_id] = ws
|
||||||
|
self._last_heartbeat[user_id] = time.time()
|
||||||
|
logger.info("[Bridge] CLI 连接: userId=%s", user_id)
|
||||||
|
|
||||||
|
# ── 发送连接确认 ──
|
||||||
|
await ws.send_json({
|
||||||
|
"type": "connected",
|
||||||
|
"userId": user_id,
|
||||||
|
})
|
||||||
|
|
||||||
|
# ── 等待能力报告 ──
|
||||||
|
try:
|
||||||
|
msg = await asyncio.wait_for(ws.receive_json(), timeout=10)
|
||||||
|
if msg.get("type") == "capability_report":
|
||||||
|
self._capabilities[user_id] = msg
|
||||||
|
logger.info(
|
||||||
|
"[Bridge] 能力报告: userId=%s, tools=%d, mcp=%d",
|
||||||
|
user_id,
|
||||||
|
len(msg.get("tools", [])),
|
||||||
|
len(msg.get("mcp_servers", [])),
|
||||||
|
)
|
||||||
|
|
||||||
|
# 审批工具
|
||||||
|
approved = self._audit_capabilities(msg)
|
||||||
|
await ws.send_json({
|
||||||
|
"type": "approved_tools",
|
||||||
|
"tools": approved,
|
||||||
|
})
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.warning("[Bridge] %s 未在 10s 内发送能力报告", user_id)
|
||||||
|
|
||||||
|
# ── 消息循环 ──
|
||||||
|
try:
|
||||||
|
async for msg in ws:
|
||||||
|
if msg.type == web.WSMsgType.TEXT:
|
||||||
|
data = json.loads(msg.data)
|
||||||
|
await self._handle_message(user_id, data)
|
||||||
|
elif msg.type in (web.WSMsgType.ERROR, web.WSMsgType.CLOSE):
|
||||||
|
break
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[Bridge] %s 连接异常: %s", user_id, e)
|
||||||
|
finally:
|
||||||
|
self._cleanup(user_id)
|
||||||
|
|
||||||
|
return ws
|
||||||
|
|
||||||
|
def is_connected(self, user_id: str) -> bool:
|
||||||
|
"""检查用户的 CLI 是否在线。"""
|
||||||
|
ws = self._connections.get(user_id)
|
||||||
|
return ws is not None and not ws.closed
|
||||||
|
|
||||||
|
def get_capabilities(self, user_id: str) -> dict | None:
|
||||||
|
"""获取用户 CLI 的能力报告。"""
|
||||||
|
return self._capabilities.get(user_id)
|
||||||
|
|
||||||
|
async def dispatch_tool_call(
|
||||||
|
self, user_id: str, tool_name: str, args: dict
|
||||||
|
) -> dict:
|
||||||
|
"""
|
||||||
|
派发工具调用:Cloud LLM → Tunnel → CLI → 结果。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
user_id: 用户 ID
|
||||||
|
tool_name: 工具名(如 "terminal")
|
||||||
|
args: 工具参数
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
{"output": "...", "exit_code": 0} 或 {"error": "..."}
|
||||||
|
"""
|
||||||
|
ws = self._connections.get(user_id)
|
||||||
|
if not ws or ws.closed:
|
||||||
|
return {"error": "CLI not connected"}
|
||||||
|
|
||||||
|
call_id = str(uuid.uuid4())[:8]
|
||||||
|
|
||||||
|
# 创建 Future 等待 CLI 响应
|
||||||
|
future: asyncio.Future = asyncio.get_event_loop().create_future()
|
||||||
|
self._pending[call_id] = future
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 发送 JSON-RPC tool_call
|
||||||
|
await ws.send_json({
|
||||||
|
"jsonrpc": "2.0",
|
||||||
|
"method": "tool_call",
|
||||||
|
"params": {"tool": tool_name, "args": args},
|
||||||
|
"id": call_id,
|
||||||
|
})
|
||||||
|
logger.info("[Bridge] → %s: %s (id=%s)", user_id, tool_name, call_id)
|
||||||
|
|
||||||
|
# 等待结果
|
||||||
|
result = await asyncio.wait_for(future, timeout=TOOL_CALL_TIMEOUT)
|
||||||
|
logger.info("[Bridge] ← %s: %s 完成 (id=%s)", user_id, tool_name, call_id)
|
||||||
|
return result
|
||||||
|
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.warning("[Bridge] %s: %s 超时 (id=%s)", user_id, tool_name, call_id)
|
||||||
|
return {"error": f"Tool call timed out after {TOOL_CALL_TIMEOUT}s"}
|
||||||
|
finally:
|
||||||
|
self._pending.pop(call_id, None)
|
||||||
|
|
||||||
|
# ── 内部方法 ──────────────────────────────────────
|
||||||
|
|
||||||
|
async def _verify_token(self, token: str, request: web.Request) -> dict | None:
|
||||||
|
"""JWT 验证(复用 MindPass 认证策略)。"""
|
||||||
|
try:
|
||||||
|
from platforms.sse_base.auth_strategies import MindPassAuth
|
||||||
|
auth = MindPassAuth()
|
||||||
|
return await auth.verify(token)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[Bridge] Token 验证失败: %s", e)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _audit_capabilities(self, report: dict) -> list[str]:
|
||||||
|
"""
|
||||||
|
审计 CLI 上报的能力,返回审批通过的工具列表。
|
||||||
|
|
||||||
|
初版:白名单 = DEFAULT_APPROVED_TOOLS 与上报工具的交集。
|
||||||
|
后续:从 DB 读取 per-user 配置。
|
||||||
|
"""
|
||||||
|
reported_tools = {t["name"] for t in report.get("tools", [])}
|
||||||
|
approved = [t for t in DEFAULT_APPROVED_TOOLS if t in reported_tools]
|
||||||
|
return approved
|
||||||
|
|
||||||
|
async def _handle_message(self, user_id: str, data: dict) -> None:
|
||||||
|
"""处理 CLI 发来的消息。"""
|
||||||
|
self._last_heartbeat[user_id] = time.time()
|
||||||
|
|
||||||
|
# JSON-RPC 响应(工具调用结果)
|
||||||
|
if data.get("jsonrpc") == "2.0" and "id" in data:
|
||||||
|
call_id = data["id"]
|
||||||
|
future = self._pending.get(call_id)
|
||||||
|
if future and not future.done():
|
||||||
|
result = data.get("result", {"error": "Empty result"})
|
||||||
|
future.set_result(result)
|
||||||
|
return
|
||||||
|
|
||||||
|
# 心跳 pong
|
||||||
|
if data.get("type") == "pong":
|
||||||
|
return
|
||||||
|
|
||||||
|
logger.debug("[Bridge] %s 未知消息: %s", user_id, data.get("type"))
|
||||||
|
|
||||||
|
def _cleanup(self, user_id: str) -> None:
|
||||||
|
"""连接关闭时清理。"""
|
||||||
|
self._connections.pop(user_id, None)
|
||||||
|
self._capabilities.pop(user_id, None)
|
||||||
|
self._last_heartbeat.pop(user_id, None)
|
||||||
|
|
||||||
|
# 取消所有待处理的 Future
|
||||||
|
to_cancel = [k for k, v in self._pending.items() if not v.done()]
|
||||||
|
for k in to_cancel:
|
||||||
|
self._pending[k].set_exception(ConnectionError("CLI disconnected"))
|
||||||
|
del self._pending[k]
|
||||||
|
|
||||||
|
logger.info("[Bridge] CLI 断开: userId=%s", user_id)
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_token(request: web.Request) -> str | None:
|
||||||
|
"""从 Authorization header 或 query 参数提取 token。"""
|
||||||
|
auth_header = request.headers.get("Authorization", "")
|
||||||
|
if auth_header.startswith("Bearer "):
|
||||||
|
return auth_header[7:]
|
||||||
|
# 备选:query 参数(WebSocket 某些客户端不支持 header)
|
||||||
|
return request.query.get("token")
|
||||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,825 @@
|
|||||||
|
"""Signal messenger platform adapter.
|
||||||
|
|
||||||
|
Connects to a signal-cli daemon running in HTTP mode.
|
||||||
|
Inbound messages arrive via SSE (Server-Sent Events) streaming.
|
||||||
|
Outbound messages and actions use JSON-RPC 2.0 over HTTP.
|
||||||
|
|
||||||
|
Based on PR #268 by ibhagwan, rebuilt with bug fixes.
|
||||||
|
|
||||||
|
Requires:
|
||||||
|
- signal-cli installed and running: signal-cli daemon --http 127.0.0.1:8080
|
||||||
|
- SIGNAL_HTTP_URL and SIGNAL_ACCOUNT environment variables set
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import base64
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import random
|
||||||
|
import time
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List, Optional, Any
|
||||||
|
from urllib.parse import quote, unquote
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
from gateway.platforms.base import (
|
||||||
|
BasePlatformAdapter,
|
||||||
|
MessageEvent,
|
||||||
|
MessageType,
|
||||||
|
SendResult,
|
||||||
|
cache_image_from_bytes,
|
||||||
|
cache_audio_from_bytes,
|
||||||
|
cache_document_from_bytes,
|
||||||
|
cache_image_from_url,
|
||||||
|
)
|
||||||
|
from gateway.platforms.helpers import redact_phone
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Constants
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
SIGNAL_MAX_ATTACHMENT_SIZE = 100 * 1024 * 1024 # 100 MB
|
||||||
|
MAX_MESSAGE_LENGTH = 8000 # Signal message size limit
|
||||||
|
TYPING_INTERVAL = 8.0 # seconds between typing indicator refreshes
|
||||||
|
SSE_RETRY_DELAY_INITIAL = 2.0
|
||||||
|
SSE_RETRY_DELAY_MAX = 60.0
|
||||||
|
HEALTH_CHECK_INTERVAL = 30.0 # seconds between health checks
|
||||||
|
HEALTH_CHECK_STALE_THRESHOLD = 120.0 # seconds without SSE activity before concern
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_comma_list(value: str) -> List[str]:
|
||||||
|
"""Split a comma-separated string into a list, stripping whitespace."""
|
||||||
|
return [v.strip() for v in value.split(",") if v.strip()]
|
||||||
|
|
||||||
|
|
||||||
|
def _guess_extension(data: bytes) -> str:
|
||||||
|
"""Guess file extension from magic bytes."""
|
||||||
|
if data[:4] == b"\x89PNG":
|
||||||
|
return ".png"
|
||||||
|
if data[:2] == b"\xff\xd8":
|
||||||
|
return ".jpg"
|
||||||
|
if data[:4] == b"GIF8":
|
||||||
|
return ".gif"
|
||||||
|
if len(data) >= 12 and data[:4] == b"RIFF" and data[8:12] == b"WEBP":
|
||||||
|
return ".webp"
|
||||||
|
if data[:4] == b"%PDF":
|
||||||
|
return ".pdf"
|
||||||
|
if len(data) >= 8 and data[4:8] == b"ftyp":
|
||||||
|
return ".mp4"
|
||||||
|
if data[:4] == b"OggS":
|
||||||
|
return ".ogg"
|
||||||
|
if len(data) >= 2 and data[0] == 0xFF and (data[1] & 0xE0) == 0xE0:
|
||||||
|
return ".mp3"
|
||||||
|
if data[:2] == b"PK":
|
||||||
|
return ".zip"
|
||||||
|
return ".bin"
|
||||||
|
|
||||||
|
|
||||||
|
def _is_image_ext(ext: str) -> bool:
|
||||||
|
return ext.lower() in (".jpg", ".jpeg", ".png", ".gif", ".webp")
|
||||||
|
|
||||||
|
|
||||||
|
def _is_audio_ext(ext: str) -> bool:
|
||||||
|
return ext.lower() in (".mp3", ".wav", ".ogg", ".m4a", ".aac")
|
||||||
|
|
||||||
|
|
||||||
|
_EXT_TO_MIME = {
|
||||||
|
".jpg": "image/jpeg", ".jpeg": "image/jpeg", ".png": "image/png",
|
||||||
|
".gif": "image/gif", ".webp": "image/webp",
|
||||||
|
".ogg": "audio/ogg", ".mp3": "audio/mpeg", ".wav": "audio/wav",
|
||||||
|
".m4a": "audio/mp4", ".aac": "audio/aac",
|
||||||
|
".mp4": "video/mp4", ".pdf": "application/pdf", ".zip": "application/zip",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _ext_to_mime(ext: str) -> str:
|
||||||
|
"""Map file extension to MIME type."""
|
||||||
|
return _EXT_TO_MIME.get(ext.lower(), "application/octet-stream")
|
||||||
|
|
||||||
|
|
||||||
|
def _render_mentions(text: str, mentions: list) -> str:
|
||||||
|
"""Replace Signal mention placeholders (\\uFFFC) with readable @identifiers.
|
||||||
|
|
||||||
|
Signal encodes @mentions as the Unicode object replacement character
|
||||||
|
with out-of-band metadata containing the mentioned user's UUID/number.
|
||||||
|
"""
|
||||||
|
if not mentions or "\uFFFC" not in text:
|
||||||
|
return text
|
||||||
|
# Sort mentions by start position (reverse) to replace from end to start
|
||||||
|
# so indices don't shift as we replace
|
||||||
|
sorted_mentions = sorted(mentions, key=lambda m: m.get("start", 0), reverse=True)
|
||||||
|
for mention in sorted_mentions:
|
||||||
|
start = mention.get("start", 0)
|
||||||
|
length = mention.get("length", 1)
|
||||||
|
# Use the mention's number or UUID as the replacement
|
||||||
|
identifier = mention.get("number") or mention.get("uuid") or "user"
|
||||||
|
replacement = f"@{identifier}"
|
||||||
|
text = text[:start] + replacement + text[start + length:]
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def check_signal_requirements() -> bool:
|
||||||
|
"""Check if Signal is configured (has URL and account)."""
|
||||||
|
return bool(os.getenv("SIGNAL_HTTP_URL") and os.getenv("SIGNAL_ACCOUNT"))
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Signal Adapter
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class SignalAdapter(BasePlatformAdapter):
|
||||||
|
"""Signal messenger adapter using signal-cli HTTP daemon."""
|
||||||
|
|
||||||
|
platform = Platform.SIGNAL
|
||||||
|
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.SIGNAL)
|
||||||
|
|
||||||
|
extra = config.extra or {}
|
||||||
|
self.http_url = extra.get("http_url", "http://127.0.0.1:8080").rstrip("/")
|
||||||
|
self.account = extra.get("account", "")
|
||||||
|
self.ignore_stories = extra.get("ignore_stories", True)
|
||||||
|
|
||||||
|
# Parse allowlists — group policy is derived from presence of group allowlist
|
||||||
|
group_allowed_str = os.getenv("SIGNAL_GROUP_ALLOWED_USERS", "")
|
||||||
|
self.group_allow_from = set(_parse_comma_list(group_allowed_str))
|
||||||
|
|
||||||
|
# HTTP client
|
||||||
|
self.client: Optional[httpx.AsyncClient] = None
|
||||||
|
|
||||||
|
# Background tasks
|
||||||
|
self._sse_task: Optional[asyncio.Task] = None
|
||||||
|
self._health_monitor_task: Optional[asyncio.Task] = None
|
||||||
|
self._typing_tasks: Dict[str, asyncio.Task] = {}
|
||||||
|
self._running = False
|
||||||
|
self._last_sse_activity = 0.0
|
||||||
|
self._sse_response: Optional[httpx.Response] = None
|
||||||
|
|
||||||
|
# Normalize account for self-message filtering
|
||||||
|
self._account_normalized = self.account.strip()
|
||||||
|
|
||||||
|
# Track recently sent message timestamps to prevent echo-back loops
|
||||||
|
# in Note to Self / self-chat mode (mirrors WhatsApp recentlySentIds)
|
||||||
|
self._recent_sent_timestamps: set = set()
|
||||||
|
self._max_recent_timestamps = 50
|
||||||
|
|
||||||
|
logger.info("Signal adapter initialized: url=%s account=%s groups=%s",
|
||||||
|
self.http_url, redact_phone(self.account),
|
||||||
|
"enabled" if self.group_allow_from else "disabled")
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Lifecycle
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
"""Connect to signal-cli daemon and start SSE listener."""
|
||||||
|
if not self.http_url or not self.account:
|
||||||
|
logger.error("Signal: SIGNAL_HTTP_URL and SIGNAL_ACCOUNT are required")
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Acquire scoped lock to prevent duplicate Signal listeners for the same phone
|
||||||
|
try:
|
||||||
|
if not self._acquire_platform_lock('signal-phone', self.account, 'Signal account'):
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Signal: Could not acquire phone lock (non-fatal): %s", e)
|
||||||
|
|
||||||
|
self.client = httpx.AsyncClient(timeout=30.0)
|
||||||
|
|
||||||
|
# Health check — verify signal-cli daemon is reachable
|
||||||
|
try:
|
||||||
|
resp = await self.client.get(f"{self.http_url}/api/v1/check", timeout=10.0)
|
||||||
|
if resp.status_code != 200:
|
||||||
|
logger.error("Signal: health check failed (status %d)", resp.status_code)
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Signal: cannot reach signal-cli at %s: %s", self.http_url, e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
self._running = True
|
||||||
|
self._last_sse_activity = time.time()
|
||||||
|
self._sse_task = asyncio.create_task(self._sse_listener())
|
||||||
|
self._health_monitor_task = asyncio.create_task(self._health_monitor())
|
||||||
|
|
||||||
|
logger.info("Signal: connected to %s", self.http_url)
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
"""Stop SSE listener and clean up."""
|
||||||
|
self._running = False
|
||||||
|
|
||||||
|
if self._sse_task:
|
||||||
|
self._sse_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._sse_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if self._health_monitor_task:
|
||||||
|
self._health_monitor_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._health_monitor_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Cancel all typing tasks
|
||||||
|
for task in self._typing_tasks.values():
|
||||||
|
task.cancel()
|
||||||
|
self._typing_tasks.clear()
|
||||||
|
|
||||||
|
if self.client:
|
||||||
|
await self.client.aclose()
|
||||||
|
self.client = None
|
||||||
|
|
||||||
|
self._release_platform_lock()
|
||||||
|
|
||||||
|
logger.info("Signal: disconnected")
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# SSE Streaming (inbound messages)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _sse_listener(self) -> None:
|
||||||
|
"""Listen for SSE events from signal-cli daemon."""
|
||||||
|
url = f"{self.http_url}/api/v1/events?account={quote(self.account, safe='')}"
|
||||||
|
backoff = SSE_RETRY_DELAY_INITIAL
|
||||||
|
|
||||||
|
while self._running:
|
||||||
|
try:
|
||||||
|
logger.debug("Signal SSE: connecting to %s", url)
|
||||||
|
async with self.client.stream(
|
||||||
|
"GET", url,
|
||||||
|
headers={"Accept": "text/event-stream"},
|
||||||
|
timeout=None,
|
||||||
|
) as response:
|
||||||
|
self._sse_response = response
|
||||||
|
backoff = SSE_RETRY_DELAY_INITIAL # Reset on successful connection
|
||||||
|
self._last_sse_activity = time.time()
|
||||||
|
logger.info("Signal SSE: connected")
|
||||||
|
|
||||||
|
buffer = ""
|
||||||
|
async for chunk in response.aiter_text():
|
||||||
|
if not self._running:
|
||||||
|
break
|
||||||
|
buffer += chunk
|
||||||
|
while "\n" in buffer:
|
||||||
|
line, buffer = buffer.split("\n", 1)
|
||||||
|
line = line.strip()
|
||||||
|
if not line:
|
||||||
|
continue
|
||||||
|
# SSE keepalive comments (":") prove the connection
|
||||||
|
# is alive — update activity so the health monitor
|
||||||
|
# doesn't report false idle warnings.
|
||||||
|
if line.startswith(":"):
|
||||||
|
self._last_sse_activity = time.time()
|
||||||
|
continue
|
||||||
|
# Parse SSE data lines
|
||||||
|
if line.startswith("data:"):
|
||||||
|
data_str = line[5:].strip()
|
||||||
|
if not data_str:
|
||||||
|
continue
|
||||||
|
self._last_sse_activity = time.time()
|
||||||
|
try:
|
||||||
|
data = json.loads(data_str)
|
||||||
|
await self._handle_envelope(data)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
logger.debug("Signal SSE: invalid JSON: %s", data_str[:100])
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Signal SSE: error handling event")
|
||||||
|
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
break
|
||||||
|
except httpx.HTTPError as e:
|
||||||
|
if self._running:
|
||||||
|
logger.warning("Signal SSE: HTTP error: %s (reconnecting in %.0fs)", e, backoff)
|
||||||
|
except Exception as e:
|
||||||
|
if self._running:
|
||||||
|
logger.warning("Signal SSE: error: %s (reconnecting in %.0fs)", e, backoff)
|
||||||
|
|
||||||
|
if self._running:
|
||||||
|
# Add 20% jitter to prevent thundering herd on reconnection
|
||||||
|
jitter = backoff * 0.2 * random.random()
|
||||||
|
await asyncio.sleep(backoff + jitter)
|
||||||
|
backoff = min(backoff * 2, SSE_RETRY_DELAY_MAX)
|
||||||
|
|
||||||
|
self._sse_response = None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Health Monitor
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _health_monitor(self) -> None:
|
||||||
|
"""Monitor SSE connection health and force reconnect if stale."""
|
||||||
|
while self._running:
|
||||||
|
await asyncio.sleep(HEALTH_CHECK_INTERVAL)
|
||||||
|
if not self._running:
|
||||||
|
break
|
||||||
|
|
||||||
|
elapsed = time.time() - self._last_sse_activity
|
||||||
|
if elapsed > HEALTH_CHECK_STALE_THRESHOLD:
|
||||||
|
logger.warning("Signal: SSE idle for %.0fs, checking daemon health", elapsed)
|
||||||
|
try:
|
||||||
|
resp = await self.client.get(
|
||||||
|
f"{self.http_url}/api/v1/check", timeout=10.0
|
||||||
|
)
|
||||||
|
if resp.status_code == 200:
|
||||||
|
# Daemon is alive but SSE is idle — update activity to
|
||||||
|
# avoid repeated warnings (connection may just be quiet)
|
||||||
|
self._last_sse_activity = time.time()
|
||||||
|
logger.debug("Signal: daemon healthy, SSE idle")
|
||||||
|
else:
|
||||||
|
logger.warning("Signal: health check failed (%d), forcing reconnect", resp.status_code)
|
||||||
|
self._force_reconnect()
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Signal: health check error: %s, forcing reconnect", e)
|
||||||
|
self._force_reconnect()
|
||||||
|
|
||||||
|
def _force_reconnect(self) -> None:
|
||||||
|
"""Force SSE reconnection by closing the current response."""
|
||||||
|
if self._sse_response and not self._sse_response.is_stream_consumed:
|
||||||
|
try:
|
||||||
|
task = asyncio.create_task(self._sse_response.aclose())
|
||||||
|
self._background_tasks.add(task)
|
||||||
|
task.add_done_callback(self._background_tasks.discard)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
self._sse_response = None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Message Handling
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _handle_envelope(self, envelope: dict) -> None:
|
||||||
|
"""Process an incoming signal-cli envelope."""
|
||||||
|
# Unwrap nested envelope if present
|
||||||
|
envelope_data = envelope.get("envelope", envelope)
|
||||||
|
|
||||||
|
# Handle syncMessage: extract "Note to Self" messages (sent to own account)
|
||||||
|
# while still filtering other sync events (read receipts, typing, etc.)
|
||||||
|
is_note_to_self = False
|
||||||
|
if "syncMessage" in envelope_data:
|
||||||
|
sync_msg = envelope_data.get("syncMessage")
|
||||||
|
if sync_msg and isinstance(sync_msg, dict):
|
||||||
|
sent_msg = sync_msg.get("sentMessage")
|
||||||
|
if sent_msg and isinstance(sent_msg, dict):
|
||||||
|
dest = sent_msg.get("destinationNumber") or sent_msg.get("destination")
|
||||||
|
sent_ts = sent_msg.get("timestamp")
|
||||||
|
if dest == self._account_normalized:
|
||||||
|
# Check if this is an echo of our own outbound reply
|
||||||
|
if sent_ts and sent_ts in self._recent_sent_timestamps:
|
||||||
|
self._recent_sent_timestamps.discard(sent_ts)
|
||||||
|
return
|
||||||
|
# Genuine user Note to Self — promote to dataMessage
|
||||||
|
is_note_to_self = True
|
||||||
|
envelope_data = {**envelope_data, "dataMessage": sent_msg}
|
||||||
|
if not is_note_to_self:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Extract sender info
|
||||||
|
sender = (
|
||||||
|
envelope_data.get("sourceNumber")
|
||||||
|
or envelope_data.get("sourceUuid")
|
||||||
|
or envelope_data.get("source")
|
||||||
|
)
|
||||||
|
sender_name = envelope_data.get("sourceName", "")
|
||||||
|
sender_uuid = envelope_data.get("sourceUuid", "")
|
||||||
|
|
||||||
|
if not sender:
|
||||||
|
logger.debug("Signal: ignoring envelope with no sender")
|
||||||
|
return
|
||||||
|
|
||||||
|
# Self-message filtering — prevent reply loops (but allow Note to Self)
|
||||||
|
if self._account_normalized and sender == self._account_normalized and not is_note_to_self:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Filter stories
|
||||||
|
if self.ignore_stories and envelope_data.get("storyMessage"):
|
||||||
|
return
|
||||||
|
|
||||||
|
# Get data message — also check editMessage (edited messages contain
|
||||||
|
# their updated dataMessage inside editMessage.dataMessage)
|
||||||
|
data_message = (
|
||||||
|
envelope_data.get("dataMessage")
|
||||||
|
or (envelope_data.get("editMessage") or {}).get("dataMessage")
|
||||||
|
)
|
||||||
|
if not data_message:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Check for group message
|
||||||
|
group_info = data_message.get("groupInfo")
|
||||||
|
group_id = group_info.get("groupId") if group_info else None
|
||||||
|
is_group = bool(group_id)
|
||||||
|
|
||||||
|
# Group message filtering — derived from SIGNAL_GROUP_ALLOWED_USERS:
|
||||||
|
# - No env var set → groups disabled (default safe behavior)
|
||||||
|
# - Env var set with group IDs → only those groups allowed
|
||||||
|
# - Env var set with "*" → all groups allowed
|
||||||
|
# DM auth is fully handled by run.py (_is_user_authorized)
|
||||||
|
if is_group:
|
||||||
|
if not self.group_allow_from:
|
||||||
|
logger.debug("Signal: ignoring group message (no SIGNAL_GROUP_ALLOWED_USERS)")
|
||||||
|
return
|
||||||
|
if "*" not in self.group_allow_from and group_id not in self.group_allow_from:
|
||||||
|
logger.debug("Signal: group %s not in allowlist", group_id[:8] if group_id else "?")
|
||||||
|
return
|
||||||
|
|
||||||
|
# Build chat info
|
||||||
|
chat_id = sender if not is_group else f"group:{group_id}"
|
||||||
|
chat_type = "group" if is_group else "dm"
|
||||||
|
|
||||||
|
# Extract text and render mentions
|
||||||
|
text = data_message.get("message", "")
|
||||||
|
mentions = data_message.get("mentions", [])
|
||||||
|
if text and mentions:
|
||||||
|
text = _render_mentions(text, mentions)
|
||||||
|
|
||||||
|
# Process attachments
|
||||||
|
attachments_data = data_message.get("attachments", [])
|
||||||
|
media_urls = []
|
||||||
|
media_types = []
|
||||||
|
|
||||||
|
if attachments_data and not getattr(self, "ignore_attachments", False):
|
||||||
|
for att in attachments_data:
|
||||||
|
att_id = att.get("id")
|
||||||
|
att_size = att.get("size", 0)
|
||||||
|
if not att_id:
|
||||||
|
continue
|
||||||
|
if att_size > SIGNAL_MAX_ATTACHMENT_SIZE:
|
||||||
|
logger.warning("Signal: attachment too large (%d bytes), skipping", att_size)
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
cached_path, ext = await self._fetch_attachment(att_id)
|
||||||
|
if cached_path:
|
||||||
|
# Use contentType from Signal if available, else map from extension
|
||||||
|
content_type = att.get("contentType") or _ext_to_mime(ext)
|
||||||
|
media_urls.append(cached_path)
|
||||||
|
media_types.append(content_type)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Signal: failed to fetch attachment %s", att_id)
|
||||||
|
|
||||||
|
# Build session source
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id=chat_id,
|
||||||
|
chat_name=group_info.get("groupName") if group_info else sender_name,
|
||||||
|
chat_type=chat_type,
|
||||||
|
user_id=sender,
|
||||||
|
user_name=sender_name or sender,
|
||||||
|
user_id_alt=sender_uuid if sender_uuid else None,
|
||||||
|
chat_id_alt=group_id if is_group else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Determine message type from media
|
||||||
|
msg_type = MessageType.TEXT
|
||||||
|
if media_types:
|
||||||
|
if any(mt.startswith("audio/") for mt in media_types):
|
||||||
|
msg_type = MessageType.VOICE
|
||||||
|
elif any(mt.startswith("image/") for mt in media_types):
|
||||||
|
msg_type = MessageType.PHOTO
|
||||||
|
|
||||||
|
# Parse timestamp from envelope data (milliseconds since epoch)
|
||||||
|
ts_ms = envelope_data.get("timestamp", 0)
|
||||||
|
if ts_ms:
|
||||||
|
try:
|
||||||
|
timestamp = datetime.fromtimestamp(ts_ms / 1000, tz=timezone.utc)
|
||||||
|
except (ValueError, OSError):
|
||||||
|
timestamp = datetime.now(tz=timezone.utc)
|
||||||
|
else:
|
||||||
|
timestamp = datetime.now(tz=timezone.utc)
|
||||||
|
|
||||||
|
# Build and dispatch event
|
||||||
|
event = MessageEvent(
|
||||||
|
source=source,
|
||||||
|
text=text or "",
|
||||||
|
message_type=msg_type,
|
||||||
|
media_urls=media_urls,
|
||||||
|
media_types=media_types,
|
||||||
|
timestamp=timestamp,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.debug("Signal: message from %s in %s: %s",
|
||||||
|
redact_phone(sender), chat_id[:20], (text or "")[:50])
|
||||||
|
|
||||||
|
await self.handle_message(event)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Attachment Handling
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _fetch_attachment(self, attachment_id: str) -> tuple:
|
||||||
|
"""Fetch an attachment via JSON-RPC and cache it. Returns (path, ext)."""
|
||||||
|
result = await self._rpc("getAttachment", {
|
||||||
|
"account": self.account,
|
||||||
|
"id": attachment_id,
|
||||||
|
})
|
||||||
|
|
||||||
|
if not result:
|
||||||
|
return None, ""
|
||||||
|
|
||||||
|
# Handle dict response (signal-cli returns {"data": "base64..."})
|
||||||
|
if isinstance(result, dict):
|
||||||
|
result = result.get("data")
|
||||||
|
if not result:
|
||||||
|
logger.warning("Signal: attachment response missing 'data' key")
|
||||||
|
return None, ""
|
||||||
|
|
||||||
|
# Result is base64-encoded file content
|
||||||
|
raw_data = base64.b64decode(result)
|
||||||
|
ext = _guess_extension(raw_data)
|
||||||
|
|
||||||
|
if _is_image_ext(ext):
|
||||||
|
path = cache_image_from_bytes(raw_data, ext)
|
||||||
|
elif _is_audio_ext(ext):
|
||||||
|
path = cache_audio_from_bytes(raw_data, ext)
|
||||||
|
else:
|
||||||
|
path = cache_document_from_bytes(raw_data, ext)
|
||||||
|
|
||||||
|
return path, ext
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# JSON-RPC Communication
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _rpc(self, method: str, params: dict, rpc_id: str = None) -> Any:
|
||||||
|
"""Send a JSON-RPC 2.0 request to signal-cli daemon."""
|
||||||
|
if not self.client:
|
||||||
|
logger.warning("Signal: RPC called but client not connected")
|
||||||
|
return None
|
||||||
|
|
||||||
|
if rpc_id is None:
|
||||||
|
rpc_id = f"{method}_{int(time.time() * 1000)}"
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"jsonrpc": "2.0",
|
||||||
|
"method": method,
|
||||||
|
"params": params,
|
||||||
|
"id": rpc_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
resp = await self.client.post(
|
||||||
|
f"{self.http_url}/api/v1/rpc",
|
||||||
|
json=payload,
|
||||||
|
timeout=30.0,
|
||||||
|
)
|
||||||
|
resp.raise_for_status()
|
||||||
|
data = resp.json()
|
||||||
|
|
||||||
|
if "error" in data:
|
||||||
|
logger.warning("Signal RPC error (%s): %s", method, data["error"])
|
||||||
|
return None
|
||||||
|
|
||||||
|
return data.get("result")
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Signal RPC %s failed: %s", method, e)
|
||||||
|
return None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Sending
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a text message."""
|
||||||
|
await self._stop_typing_indicator(chat_id)
|
||||||
|
|
||||||
|
params: Dict[str, Any] = {
|
||||||
|
"account": self.account,
|
||||||
|
"message": content,
|
||||||
|
}
|
||||||
|
|
||||||
|
if chat_id.startswith("group:"):
|
||||||
|
params["groupId"] = chat_id[6:]
|
||||||
|
else:
|
||||||
|
params["recipient"] = [chat_id]
|
||||||
|
|
||||||
|
result = await self._rpc("send", params)
|
||||||
|
|
||||||
|
if result is not None:
|
||||||
|
self._track_sent_timestamp(result)
|
||||||
|
# Use the timestamp from the RPC result as a pseudo message_id.
|
||||||
|
# Signal doesn't have real message IDs, but the stream consumer
|
||||||
|
# needs a truthy value to follow its edit→fallback path correctly.
|
||||||
|
_msg_id = str(result.get("timestamp", "")) if isinstance(result, dict) else None
|
||||||
|
return SendResult(success=True, message_id=_msg_id or None)
|
||||||
|
return SendResult(success=False, error="RPC send failed")
|
||||||
|
|
||||||
|
def _track_sent_timestamp(self, rpc_result) -> None:
|
||||||
|
"""Record outbound message timestamp for echo-back filtering."""
|
||||||
|
ts = rpc_result.get("timestamp") if isinstance(rpc_result, dict) else None
|
||||||
|
if ts:
|
||||||
|
self._recent_sent_timestamps.add(ts)
|
||||||
|
if len(self._recent_sent_timestamps) > self._max_recent_timestamps:
|
||||||
|
self._recent_sent_timestamps.pop()
|
||||||
|
|
||||||
|
async def send_typing(self, chat_id: str, metadata=None) -> None:
|
||||||
|
"""Send a typing indicator."""
|
||||||
|
params: Dict[str, Any] = {
|
||||||
|
"account": self.account,
|
||||||
|
}
|
||||||
|
|
||||||
|
if chat_id.startswith("group:"):
|
||||||
|
params["groupId"] = chat_id[6:]
|
||||||
|
else:
|
||||||
|
params["recipient"] = [chat_id]
|
||||||
|
|
||||||
|
await self._rpc("sendTyping", params, rpc_id="typing")
|
||||||
|
|
||||||
|
async def send_image(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
image_url: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send an image. Supports http(s):// and file:// URLs."""
|
||||||
|
await self._stop_typing_indicator(chat_id)
|
||||||
|
|
||||||
|
# Resolve image to local path
|
||||||
|
if image_url.startswith("file://"):
|
||||||
|
file_path = unquote(image_url[7:])
|
||||||
|
else:
|
||||||
|
# Download remote image to cache
|
||||||
|
try:
|
||||||
|
file_path = await cache_image_from_url(image_url)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Signal: failed to download image: %s", e)
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
if not file_path or not Path(file_path).exists():
|
||||||
|
return SendResult(success=False, error="Image file not found")
|
||||||
|
|
||||||
|
# Validate size
|
||||||
|
file_size = Path(file_path).stat().st_size
|
||||||
|
if file_size > SIGNAL_MAX_ATTACHMENT_SIZE:
|
||||||
|
return SendResult(success=False, error=f"Image too large ({file_size} bytes)")
|
||||||
|
|
||||||
|
params: Dict[str, Any] = {
|
||||||
|
"account": self.account,
|
||||||
|
"message": caption or "",
|
||||||
|
"attachments": [file_path],
|
||||||
|
}
|
||||||
|
|
||||||
|
if chat_id.startswith("group:"):
|
||||||
|
params["groupId"] = chat_id[6:]
|
||||||
|
else:
|
||||||
|
params["recipient"] = [chat_id]
|
||||||
|
|
||||||
|
result = await self._rpc("send", params)
|
||||||
|
if result is not None:
|
||||||
|
self._track_sent_timestamp(result)
|
||||||
|
return SendResult(success=True)
|
||||||
|
return SendResult(success=False, error="RPC send with attachment failed")
|
||||||
|
|
||||||
|
async def _send_attachment(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
file_path: str,
|
||||||
|
media_label: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send any file as a Signal attachment via RPC.
|
||||||
|
|
||||||
|
Shared implementation for send_document, send_image_file, send_voice,
|
||||||
|
and send_video — avoids duplicating the validation/routing/RPC logic.
|
||||||
|
"""
|
||||||
|
await self._stop_typing_indicator(chat_id)
|
||||||
|
|
||||||
|
try:
|
||||||
|
file_size = Path(file_path).stat().st_size
|
||||||
|
except FileNotFoundError:
|
||||||
|
return SendResult(success=False, error=f"{media_label} file not found: {file_path}")
|
||||||
|
|
||||||
|
if file_size > SIGNAL_MAX_ATTACHMENT_SIZE:
|
||||||
|
return SendResult(success=False, error=f"{media_label} too large ({file_size} bytes)")
|
||||||
|
|
||||||
|
params: Dict[str, Any] = {
|
||||||
|
"account": self.account,
|
||||||
|
"message": caption or "",
|
||||||
|
"attachments": [file_path],
|
||||||
|
}
|
||||||
|
|
||||||
|
if chat_id.startswith("group:"):
|
||||||
|
params["groupId"] = chat_id[6:]
|
||||||
|
else:
|
||||||
|
params["recipient"] = [chat_id]
|
||||||
|
|
||||||
|
result = await self._rpc("send", params)
|
||||||
|
if result is not None:
|
||||||
|
self._track_sent_timestamp(result)
|
||||||
|
return SendResult(success=True)
|
||||||
|
return SendResult(success=False, error=f"RPC send {media_label.lower()} failed")
|
||||||
|
|
||||||
|
async def send_document(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
file_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
filename: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a document/file attachment."""
|
||||||
|
return await self._send_attachment(chat_id, file_path, "File", caption)
|
||||||
|
|
||||||
|
async def send_image_file(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
image_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a local image file as a native Signal attachment.
|
||||||
|
|
||||||
|
Called by the gateway media delivery flow when MEDIA: tags containing
|
||||||
|
image paths are extracted from agent responses.
|
||||||
|
"""
|
||||||
|
return await self._send_attachment(chat_id, image_path, "Image", caption)
|
||||||
|
|
||||||
|
async def send_voice(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
audio_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send an audio file as a Signal attachment.
|
||||||
|
|
||||||
|
Signal does not distinguish voice messages from file attachments at
|
||||||
|
the API level, so this routes through the same RPC send path.
|
||||||
|
"""
|
||||||
|
return await self._send_attachment(chat_id, audio_path, "Audio", caption)
|
||||||
|
|
||||||
|
async def send_video(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
video_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a video file as a Signal attachment."""
|
||||||
|
return await self._send_attachment(chat_id, video_path, "Video", caption)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Typing Indicators
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _stop_typing_indicator(self, chat_id: str) -> None:
|
||||||
|
"""Stop a typing indicator loop for a chat."""
|
||||||
|
task = self._typing_tasks.pop(chat_id, None)
|
||||||
|
if task:
|
||||||
|
task.cancel()
|
||||||
|
try:
|
||||||
|
await task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def stop_typing(self, chat_id: str) -> None:
|
||||||
|
"""Public interface for stopping typing — called by base adapter's
|
||||||
|
_keep_typing finally block to clean up platform-level typing tasks."""
|
||||||
|
await self._stop_typing_indicator(chat_id)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Chat Info
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
"""Get information about a chat/contact."""
|
||||||
|
if chat_id.startswith("group:"):
|
||||||
|
return {
|
||||||
|
"name": chat_id,
|
||||||
|
"type": "group",
|
||||||
|
"chat_id": chat_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Try to resolve contact name
|
||||||
|
result = await self._rpc("getContact", {
|
||||||
|
"account": self.account,
|
||||||
|
"contactAddress": chat_id,
|
||||||
|
})
|
||||||
|
|
||||||
|
name = chat_id
|
||||||
|
if result and isinstance(result, dict):
|
||||||
|
name = result.get("name") or result.get("profileName") or chat_id
|
||||||
|
|
||||||
|
return {
|
||||||
|
"name": name,
|
||||||
|
"type": "dm",
|
||||||
|
"chat_id": chat_id,
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,373 @@
|
|||||||
|
"""SMS (Twilio) platform adapter.
|
||||||
|
|
||||||
|
Connects to the Twilio REST API for outbound SMS and runs an aiohttp
|
||||||
|
webhook server to receive inbound messages.
|
||||||
|
|
||||||
|
Shares credentials with the optional telephony skill — same env vars:
|
||||||
|
- TWILIO_ACCOUNT_SID
|
||||||
|
- TWILIO_AUTH_TOKEN
|
||||||
|
- TWILIO_PHONE_NUMBER (E.164 from-number, e.g. +15551234567)
|
||||||
|
|
||||||
|
Gateway-specific env vars:
|
||||||
|
- SMS_WEBHOOK_PORT (default 8080)
|
||||||
|
- SMS_WEBHOOK_HOST (default 0.0.0.0)
|
||||||
|
- SMS_WEBHOOK_URL (public URL for Twilio signature validation — required)
|
||||||
|
- SMS_INSECURE_NO_SIGNATURE (true to disable signature validation — dev only)
|
||||||
|
- SMS_ALLOWED_USERS (comma-separated E.164 phone numbers)
|
||||||
|
- SMS_ALLOW_ALL_USERS (true/false)
|
||||||
|
- SMS_HOME_CHANNEL (phone number for cron delivery)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import base64
|
||||||
|
import hashlib
|
||||||
|
import hmac
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import urllib.parse
|
||||||
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
from gateway.platforms.base import (
|
||||||
|
BasePlatformAdapter,
|
||||||
|
MessageEvent,
|
||||||
|
MessageType,
|
||||||
|
SendResult,
|
||||||
|
)
|
||||||
|
from gateway.platforms.helpers import redact_phone, strip_markdown
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
TWILIO_API_BASE = "https://api.twilio.com/2010-04-01/Accounts"
|
||||||
|
MAX_SMS_LENGTH = 1600 # ~10 SMS segments
|
||||||
|
DEFAULT_WEBHOOK_PORT = 8080
|
||||||
|
DEFAULT_WEBHOOK_HOST = "0.0.0.0"
|
||||||
|
|
||||||
|
|
||||||
|
def check_sms_requirements() -> bool:
|
||||||
|
"""Check if SMS adapter dependencies are available."""
|
||||||
|
try:
|
||||||
|
import aiohttp # noqa: F401
|
||||||
|
except ImportError:
|
||||||
|
return False
|
||||||
|
return bool(os.getenv("TWILIO_ACCOUNT_SID") and os.getenv("TWILIO_AUTH_TOKEN"))
|
||||||
|
|
||||||
|
|
||||||
|
class SmsAdapter(BasePlatformAdapter):
|
||||||
|
"""
|
||||||
|
Twilio SMS <-> Hermes gateway adapter.
|
||||||
|
|
||||||
|
Each inbound phone number gets its own Hermes session (multi-tenant).
|
||||||
|
Replies are always sent from the configured TWILIO_PHONE_NUMBER.
|
||||||
|
"""
|
||||||
|
|
||||||
|
MAX_MESSAGE_LENGTH = MAX_SMS_LENGTH
|
||||||
|
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.SMS)
|
||||||
|
self._account_sid: str = os.environ["TWILIO_ACCOUNT_SID"]
|
||||||
|
self._auth_token: str = os.environ["TWILIO_AUTH_TOKEN"]
|
||||||
|
self._from_number: str = os.getenv("TWILIO_PHONE_NUMBER", "")
|
||||||
|
self._webhook_port: int = int(
|
||||||
|
os.getenv("SMS_WEBHOOK_PORT", str(DEFAULT_WEBHOOK_PORT))
|
||||||
|
)
|
||||||
|
self._webhook_host: str = os.getenv("SMS_WEBHOOK_HOST", DEFAULT_WEBHOOK_HOST)
|
||||||
|
self._webhook_url: str = os.getenv("SMS_WEBHOOK_URL", "").strip()
|
||||||
|
self._runner = None
|
||||||
|
self._http_session: Optional["aiohttp.ClientSession"] = None
|
||||||
|
|
||||||
|
def _basic_auth_header(self) -> str:
|
||||||
|
"""Build HTTP Basic auth header value for Twilio."""
|
||||||
|
creds = f"{self._account_sid}:{self._auth_token}"
|
||||||
|
encoded = base64.b64encode(creds.encode("ascii")).decode("ascii")
|
||||||
|
return f"Basic {encoded}"
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Required abstract methods
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
import aiohttp
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
if not self._from_number:
|
||||||
|
logger.error("[sms] TWILIO_PHONE_NUMBER not set — cannot send replies")
|
||||||
|
return False
|
||||||
|
|
||||||
|
insecure_no_sig = os.getenv("SMS_INSECURE_NO_SIGNATURE", "").lower() == "true"
|
||||||
|
|
||||||
|
if not self._webhook_url and not insecure_no_sig:
|
||||||
|
logger.error(
|
||||||
|
"[sms] Refusing to start: SMS_WEBHOOK_URL is required for Twilio "
|
||||||
|
"signature validation. Set it to the public URL configured in your "
|
||||||
|
"Twilio console (e.g. https://example.com/webhooks/twilio). "
|
||||||
|
"For local development without validation, set "
|
||||||
|
"SMS_INSECURE_NO_SIGNATURE=true (NOT recommended for production).",
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
if insecure_no_sig and not self._webhook_url:
|
||||||
|
logger.warning(
|
||||||
|
"[sms] SMS_INSECURE_NO_SIGNATURE=true — Twilio signature validation "
|
||||||
|
"is DISABLED. Any client that can reach port %d can inject messages. "
|
||||||
|
"Do NOT use this in production.",
|
||||||
|
self._webhook_port,
|
||||||
|
)
|
||||||
|
|
||||||
|
app = web.Application()
|
||||||
|
app.router.add_post("/webhooks/twilio", self._handle_webhook)
|
||||||
|
app.router.add_get("/health", lambda _: web.Response(text="ok"))
|
||||||
|
|
||||||
|
self._runner = web.AppRunner(app)
|
||||||
|
await self._runner.setup()
|
||||||
|
site = web.TCPSite(self._runner, self._webhook_host, self._webhook_port)
|
||||||
|
await site.start()
|
||||||
|
self._http_session = aiohttp.ClientSession(
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30),
|
||||||
|
)
|
||||||
|
self._running = True
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"[sms] Twilio webhook server listening on %s:%d, from: %s",
|
||||||
|
self._webhook_host,
|
||||||
|
self._webhook_port,
|
||||||
|
redact_phone(self._from_number),
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
if self._http_session:
|
||||||
|
await self._http_session.close()
|
||||||
|
self._http_session = None
|
||||||
|
if self._runner:
|
||||||
|
await self._runner.cleanup()
|
||||||
|
self._runner = None
|
||||||
|
self._running = False
|
||||||
|
logger.info("[sms] Disconnected")
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
formatted = self.format_message(content)
|
||||||
|
chunks = self.truncate_message(formatted)
|
||||||
|
last_result = SendResult(success=True)
|
||||||
|
|
||||||
|
url = f"{TWILIO_API_BASE}/{self._account_sid}/Messages.json"
|
||||||
|
headers = {
|
||||||
|
"Authorization": self._basic_auth_header(),
|
||||||
|
}
|
||||||
|
|
||||||
|
session = self._http_session or aiohttp.ClientSession(
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
for chunk in chunks:
|
||||||
|
form_data = aiohttp.FormData()
|
||||||
|
form_data.add_field("From", self._from_number)
|
||||||
|
form_data.add_field("To", chat_id)
|
||||||
|
form_data.add_field("Body", chunk)
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with session.post(url, data=form_data, headers=headers) as resp:
|
||||||
|
body = await resp.json()
|
||||||
|
if resp.status >= 400:
|
||||||
|
error_msg = body.get("message", str(body))
|
||||||
|
logger.error(
|
||||||
|
"[sms] send failed to %s: %s %s",
|
||||||
|
redact_phone(chat_id),
|
||||||
|
resp.status,
|
||||||
|
error_msg,
|
||||||
|
)
|
||||||
|
return SendResult(
|
||||||
|
success=False,
|
||||||
|
error=f"Twilio {resp.status}: {error_msg}",
|
||||||
|
)
|
||||||
|
msg_sid = body.get("sid", "")
|
||||||
|
last_result = SendResult(success=True, message_id=msg_sid)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[sms] send error to %s: %s", redact_phone(chat_id), e)
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
finally:
|
||||||
|
# Close session only if we created a fallback (no persistent session)
|
||||||
|
if not self._http_session and session:
|
||||||
|
await session.close()
|
||||||
|
|
||||||
|
return last_result
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
return {"name": chat_id, "type": "dm"}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# SMS-specific formatting
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def format_message(self, content: str) -> str:
|
||||||
|
"""Strip markdown — SMS renders it as literal characters."""
|
||||||
|
return strip_markdown(content)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Twilio signature validation
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _validate_twilio_signature(
|
||||||
|
self, url: str, post_params: dict, signature: str,
|
||||||
|
) -> bool:
|
||||||
|
"""Validate ``X-Twilio-Signature`` header (HMAC-SHA1, base64).
|
||||||
|
|
||||||
|
Tries both with and without the default port for the URL scheme,
|
||||||
|
since Twilio may sign with either variant.
|
||||||
|
|
||||||
|
Algorithm: https://www.twilio.com/docs/usage/security#validating-requests
|
||||||
|
"""
|
||||||
|
if self._check_signature(url, post_params, signature):
|
||||||
|
return True
|
||||||
|
|
||||||
|
variant = self._port_variant_url(url)
|
||||||
|
if variant and self._check_signature(variant, post_params, signature):
|
||||||
|
return True
|
||||||
|
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _check_signature(
|
||||||
|
self, url: str, post_params: dict, signature: str,
|
||||||
|
) -> bool:
|
||||||
|
"""Compute and compare a single Twilio signature."""
|
||||||
|
data_to_sign = url
|
||||||
|
for key in sorted(post_params.keys()):
|
||||||
|
data_to_sign += key + post_params[key]
|
||||||
|
mac = hmac.new(
|
||||||
|
self._auth_token.encode("utf-8"),
|
||||||
|
data_to_sign.encode("utf-8"),
|
||||||
|
hashlib.sha1,
|
||||||
|
)
|
||||||
|
computed = base64.b64encode(mac.digest()).decode("utf-8")
|
||||||
|
return hmac.compare_digest(computed, signature)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _port_variant_url(url: str) -> str | None:
|
||||||
|
"""Return the URL with the default port toggled, or None.
|
||||||
|
|
||||||
|
Only toggles default ports (443 for https, 80 for http).
|
||||||
|
Non-standard ports are never modified.
|
||||||
|
"""
|
||||||
|
parsed = urllib.parse.urlparse(url)
|
||||||
|
default_ports = {"https": 443, "http": 80}
|
||||||
|
default_port = default_ports.get(parsed.scheme)
|
||||||
|
if default_port is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if parsed.port == default_port:
|
||||||
|
# Has explicit default port → strip it
|
||||||
|
return urllib.parse.urlunparse(
|
||||||
|
(parsed.scheme, parsed.hostname, parsed.path,
|
||||||
|
parsed.params, parsed.query, parsed.fragment)
|
||||||
|
)
|
||||||
|
elif parsed.port is None:
|
||||||
|
# No port → add default
|
||||||
|
netloc = f"{parsed.hostname}:{default_port}"
|
||||||
|
return urllib.parse.urlunparse(
|
||||||
|
(parsed.scheme, netloc, parsed.path,
|
||||||
|
parsed.params, parsed.query, parsed.fragment)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Non-standard port — no variant
|
||||||
|
return None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Twilio webhook handler
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _handle_webhook(self, request) -> "aiohttp.web.Response":
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
try:
|
||||||
|
raw = await request.read()
|
||||||
|
# Twilio sends form-encoded data, not JSON
|
||||||
|
form = urllib.parse.parse_qs(raw.decode("utf-8"), keep_blank_values=True)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[sms] webhook parse error: %s", e)
|
||||||
|
return web.Response(
|
||||||
|
text='<?xml version="1.0" encoding="UTF-8"?><Response></Response>',
|
||||||
|
content_type="application/xml",
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Validate Twilio request signature when SMS_WEBHOOK_URL is configured
|
||||||
|
if self._webhook_url:
|
||||||
|
twilio_sig = request.headers.get("X-Twilio-Signature", "")
|
||||||
|
if not twilio_sig:
|
||||||
|
logger.warning("[sms] Rejected: missing X-Twilio-Signature header")
|
||||||
|
return web.Response(
|
||||||
|
text='<?xml version="1.0" encoding="UTF-8"?><Response></Response>',
|
||||||
|
content_type="application/xml",
|
||||||
|
status=403,
|
||||||
|
)
|
||||||
|
flat_params = {k: v[0] for k, v in form.items() if v}
|
||||||
|
if not self._validate_twilio_signature(
|
||||||
|
self._webhook_url, flat_params, twilio_sig
|
||||||
|
):
|
||||||
|
logger.warning("[sms] Rejected: invalid Twilio signature")
|
||||||
|
return web.Response(
|
||||||
|
text='<?xml version="1.0" encoding="UTF-8"?><Response></Response>',
|
||||||
|
content_type="application/xml",
|
||||||
|
status=403,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Extract fields (parse_qs returns lists)
|
||||||
|
from_number = (form.get("From", [""]))[0].strip()
|
||||||
|
to_number = (form.get("To", [""]))[0].strip()
|
||||||
|
text = (form.get("Body", [""]))[0].strip()
|
||||||
|
message_sid = (form.get("MessageSid", [""]))[0].strip()
|
||||||
|
|
||||||
|
if not from_number or not text:
|
||||||
|
return web.Response(
|
||||||
|
text='<?xml version="1.0" encoding="UTF-8"?><Response></Response>',
|
||||||
|
content_type="application/xml",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Ignore messages from our own number (echo prevention)
|
||||||
|
if from_number == self._from_number:
|
||||||
|
logger.debug("[sms] ignoring echo from own number %s", redact_phone(from_number))
|
||||||
|
return web.Response(
|
||||||
|
text='<?xml version="1.0" encoding="UTF-8"?><Response></Response>',
|
||||||
|
content_type="application/xml",
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"[sms] inbound from %s -> %s: %s",
|
||||||
|
redact_phone(from_number),
|
||||||
|
redact_phone(to_number),
|
||||||
|
text[:80],
|
||||||
|
)
|
||||||
|
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id=from_number,
|
||||||
|
chat_name=from_number,
|
||||||
|
chat_type="dm",
|
||||||
|
user_id=from_number,
|
||||||
|
user_name=from_number,
|
||||||
|
)
|
||||||
|
event = MessageEvent(
|
||||||
|
text=text,
|
||||||
|
message_type=MessageType.TEXT,
|
||||||
|
source=source,
|
||||||
|
raw_message=form,
|
||||||
|
message_id=message_sid,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Non-blocking: Twilio expects a fast response
|
||||||
|
task = asyncio.create_task(self.handle_message(event))
|
||||||
|
self._background_tasks.add(task)
|
||||||
|
task.add_done_callback(self._background_tasks.discard)
|
||||||
|
|
||||||
|
# Return empty TwiML — we send replies via the REST API, not inline TwiML
|
||||||
|
return web.Response(
|
||||||
|
text='<?xml version="1.0" encoding="UTF-8"?><Response></Response>',
|
||||||
|
content_type="application/xml",
|
||||||
|
)
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,246 @@
|
|||||||
|
"""Telegram-specific network helpers.
|
||||||
|
|
||||||
|
Provides a hostname-preserving fallback transport for networks where
|
||||||
|
api.telegram.org resolves to an endpoint that is unreachable from the current
|
||||||
|
host. The transport keeps the logical request host and TLS SNI as
|
||||||
|
api.telegram.org while retrying the TCP connection against one or more fallback
|
||||||
|
IPv4 addresses.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import ipaddress
|
||||||
|
import logging
|
||||||
|
import socket
|
||||||
|
from typing import Iterable, Optional
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_TELEGRAM_API_HOST = "api.telegram.org"
|
||||||
|
|
||||||
|
# DNS-over-HTTPS providers used to discover Telegram API IPs that may differ
|
||||||
|
# from the (potentially unreachable) IP returned by the local system resolver.
|
||||||
|
_DOH_TIMEOUT = 4.0 # seconds — bounded so connect() isn't noticeably delayed
|
||||||
|
|
||||||
|
_DOH_PROVIDERS: list[dict] = [
|
||||||
|
{
|
||||||
|
"url": "https://dns.google/resolve",
|
||||||
|
"params": {"name": _TELEGRAM_API_HOST, "type": "A"},
|
||||||
|
"headers": {},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"url": "https://cloudflare-dns.com/dns-query",
|
||||||
|
"params": {"name": _TELEGRAM_API_HOST, "type": "A"},
|
||||||
|
"headers": {"Accept": "application/dns-json"},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
# Last-resort IPs when DoH is also blocked. These are stable Telegram Bot API
|
||||||
|
# endpoints in the 149.154.160.0/20 block (same seed used by OpenClaw).
|
||||||
|
_SEED_FALLBACK_IPS: list[str] = ["149.154.167.220"]
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_proxy_url() -> str | None:
|
||||||
|
# Delegate to shared implementation (env vars + macOS system proxy detection)
|
||||||
|
from gateway.platforms.base import resolve_proxy_url
|
||||||
|
return resolve_proxy_url()
|
||||||
|
|
||||||
|
|
||||||
|
class TelegramFallbackTransport(httpx.AsyncBaseTransport):
|
||||||
|
"""Retry Telegram Bot API requests via fallback IPs while preserving TLS/SNI.
|
||||||
|
|
||||||
|
Requests continue to target https://api.telegram.org/... logically, but on
|
||||||
|
connect failures the underlying TCP connection is retried against a known
|
||||||
|
reachable IP. This is effectively the programmatic equivalent of
|
||||||
|
``curl --resolve api.telegram.org:443:<ip>``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, fallback_ips: Iterable[str], **transport_kwargs):
|
||||||
|
self._fallback_ips = [ip for ip in dict.fromkeys(_normalize_fallback_ips(fallback_ips))]
|
||||||
|
proxy_url = _resolve_proxy_url()
|
||||||
|
if proxy_url and "proxy" not in transport_kwargs:
|
||||||
|
transport_kwargs["proxy"] = proxy_url
|
||||||
|
self._primary = httpx.AsyncHTTPTransport(**transport_kwargs)
|
||||||
|
self._fallbacks = {
|
||||||
|
ip: httpx.AsyncHTTPTransport(**transport_kwargs) for ip in self._fallback_ips
|
||||||
|
}
|
||||||
|
self._sticky_ip: Optional[str] = None
|
||||||
|
self._sticky_lock = asyncio.Lock()
|
||||||
|
|
||||||
|
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||||
|
if request.url.host != _TELEGRAM_API_HOST or not self._fallback_ips:
|
||||||
|
return await self._primary.handle_async_request(request)
|
||||||
|
|
||||||
|
sticky_ip = self._sticky_ip
|
||||||
|
attempt_order: list[Optional[str]] = [sticky_ip] if sticky_ip else [None]
|
||||||
|
for ip in self._fallback_ips:
|
||||||
|
if ip != sticky_ip:
|
||||||
|
attempt_order.append(ip)
|
||||||
|
|
||||||
|
last_error: Exception | None = None
|
||||||
|
for ip in attempt_order:
|
||||||
|
candidate = request if ip is None else _rewrite_request_for_ip(request, ip)
|
||||||
|
transport = self._primary if ip is None else self._fallbacks[ip]
|
||||||
|
try:
|
||||||
|
response = await transport.handle_async_request(candidate)
|
||||||
|
if ip is not None and self._sticky_ip != ip:
|
||||||
|
async with self._sticky_lock:
|
||||||
|
if self._sticky_ip != ip:
|
||||||
|
self._sticky_ip = ip
|
||||||
|
logger.warning(
|
||||||
|
"[Telegram] Primary api.telegram.org path unreachable; using sticky fallback IP %s",
|
||||||
|
ip,
|
||||||
|
)
|
||||||
|
return response
|
||||||
|
except Exception as exc:
|
||||||
|
last_error = exc
|
||||||
|
if not _is_retryable_connect_error(exc):
|
||||||
|
raise
|
||||||
|
if ip is None:
|
||||||
|
logger.warning(
|
||||||
|
"[Telegram] Primary api.telegram.org connection failed (%s); trying fallback IPs %s",
|
||||||
|
exc,
|
||||||
|
", ".join(self._fallback_ips),
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
logger.warning("[Telegram] Fallback IP %s failed: %s", ip, exc)
|
||||||
|
continue
|
||||||
|
|
||||||
|
if last_error is None:
|
||||||
|
raise RuntimeError("All Telegram fallback IPs exhausted but no error was recorded")
|
||||||
|
raise last_error
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
await self._primary.aclose()
|
||||||
|
for transport in self._fallbacks.values():
|
||||||
|
await transport.aclose()
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_fallback_ips(values: Iterable[str]) -> list[str]:
|
||||||
|
normalized: list[str] = []
|
||||||
|
for value in values:
|
||||||
|
raw = str(value).strip()
|
||||||
|
if not raw:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
addr = ipaddress.ip_address(raw)
|
||||||
|
except ValueError:
|
||||||
|
logger.warning("Ignoring invalid Telegram fallback IP: %r", raw)
|
||||||
|
continue
|
||||||
|
if addr.version != 4:
|
||||||
|
logger.warning("Ignoring non-IPv4 Telegram fallback IP: %s", raw)
|
||||||
|
continue
|
||||||
|
if addr.is_private or addr.is_loopback or addr.is_link_local or addr.is_unspecified:
|
||||||
|
logger.warning("Ignoring private/internal Telegram fallback IP: %s", raw)
|
||||||
|
continue
|
||||||
|
normalized.append(str(addr))
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
|
||||||
|
def parse_fallback_ip_env(value: str | None) -> list[str]:
|
||||||
|
if not value:
|
||||||
|
return []
|
||||||
|
parts = [part.strip() for part in value.split(",")]
|
||||||
|
return _normalize_fallback_ips(parts)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_system_dns() -> set[str]:
|
||||||
|
"""Return the IPv4 addresses that the OS resolver gives for api.telegram.org."""
|
||||||
|
try:
|
||||||
|
results = socket.getaddrinfo(_TELEGRAM_API_HOST, 443, socket.AF_INET)
|
||||||
|
return {addr[4][0] for addr in results}
|
||||||
|
except Exception:
|
||||||
|
return set()
|
||||||
|
|
||||||
|
|
||||||
|
async def _query_doh_provider(
|
||||||
|
client: httpx.AsyncClient, provider: dict
|
||||||
|
) -> list[str]:
|
||||||
|
"""Query one DoH provider and return A-record IPs."""
|
||||||
|
try:
|
||||||
|
resp = await client.get(
|
||||||
|
provider["url"], params=provider["params"], headers=provider["headers"]
|
||||||
|
)
|
||||||
|
resp.raise_for_status()
|
||||||
|
data = resp.json()
|
||||||
|
ips: list[str] = []
|
||||||
|
for answer in data.get("Answer", []):
|
||||||
|
if answer.get("type") != 1: # A record
|
||||||
|
continue
|
||||||
|
raw = answer.get("data", "").strip()
|
||||||
|
try:
|
||||||
|
ipaddress.ip_address(raw)
|
||||||
|
ips.append(raw)
|
||||||
|
except ValueError:
|
||||||
|
continue
|
||||||
|
return ips
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug("DoH query to %s failed: %s", provider["url"], exc)
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
async def discover_fallback_ips() -> list[str]:
|
||||||
|
"""Auto-discover Telegram API IPs via DNS-over-HTTPS.
|
||||||
|
|
||||||
|
Resolves api.telegram.org through Google and Cloudflare DoH, collects all
|
||||||
|
unique IPs, and excludes the system-DNS-resolved IP (which is presumably
|
||||||
|
unreachable on this network). Falls back to a hardcoded seed list when DoH
|
||||||
|
is also unavailable.
|
||||||
|
"""
|
||||||
|
async with httpx.AsyncClient(timeout=httpx.Timeout(_DOH_TIMEOUT)) as client:
|
||||||
|
doh_tasks = [_query_doh_provider(client, p) for p in _DOH_PROVIDERS]
|
||||||
|
system_dns_task = asyncio.to_thread(_resolve_system_dns)
|
||||||
|
results = await asyncio.gather(system_dns_task, *doh_tasks, return_exceptions=True)
|
||||||
|
|
||||||
|
# results[0] = system DNS IPs (set), results[1:] = DoH IP lists
|
||||||
|
system_ips: set[str] = results[0] if isinstance(results[0], set) else set()
|
||||||
|
|
||||||
|
doh_ips: list[str] = []
|
||||||
|
for r in results[1:]:
|
||||||
|
if isinstance(r, list):
|
||||||
|
doh_ips.extend(r)
|
||||||
|
|
||||||
|
# Deduplicate preserving order, exclude system-DNS IPs
|
||||||
|
seen: set[str] = set()
|
||||||
|
candidates: list[str] = []
|
||||||
|
for ip in doh_ips:
|
||||||
|
if ip not in seen and ip not in system_ips:
|
||||||
|
seen.add(ip)
|
||||||
|
candidates.append(ip)
|
||||||
|
|
||||||
|
# Validate through existing normalization
|
||||||
|
validated = _normalize_fallback_ips(candidates)
|
||||||
|
|
||||||
|
if validated:
|
||||||
|
logger.debug("Discovered Telegram fallback IPs via DoH: %s", ", ".join(validated))
|
||||||
|
return validated
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"DoH discovery yielded no new IPs (system DNS: %s); using seed fallback IPs %s",
|
||||||
|
", ".join(system_ips) or "unknown",
|
||||||
|
", ".join(_SEED_FALLBACK_IPS),
|
||||||
|
)
|
||||||
|
return list(_SEED_FALLBACK_IPS)
|
||||||
|
|
||||||
|
|
||||||
|
def _rewrite_request_for_ip(request: httpx.Request, ip: str) -> httpx.Request:
|
||||||
|
original_host = request.url.host or _TELEGRAM_API_HOST
|
||||||
|
url = request.url.copy_with(host=ip)
|
||||||
|
headers = request.headers.copy()
|
||||||
|
headers["host"] = original_host
|
||||||
|
extensions = dict(request.extensions)
|
||||||
|
extensions["sni_hostname"] = original_host
|
||||||
|
return httpx.Request(
|
||||||
|
method=request.method,
|
||||||
|
url=url,
|
||||||
|
headers=headers,
|
||||||
|
stream=request.stream,
|
||||||
|
extensions=extensions,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_retryable_connect_error(exc: Exception) -> bool:
|
||||||
|
return isinstance(exc, (httpx.ConnectTimeout, httpx.ConnectError))
|
||||||
@@ -0,0 +1,136 @@
|
|||||||
|
"""
|
||||||
|
voice2md_atoms.py — Voice2MD 共享原子层
|
||||||
|
|
||||||
|
Spec 来源: docs/SPEC_voice2md_atomization_design.md §九
|
||||||
|
提取自: mindos_sse.py._runOssTranscribePipeline (管线A)
|
||||||
|
dashscope_realtime._finalize (管线B)
|
||||||
|
|
||||||
|
共享原子:
|
||||||
|
④ persist_audio_result() — 统一 DB 持久化
|
||||||
|
⑤ push_md_appended() — 统一 SSE 推送
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger("voice2md_atoms")
|
||||||
|
|
||||||
|
|
||||||
|
# ─── ④ db_persist: 统一音频结果持久化 ─────────────────────────
|
||||||
|
|
||||||
|
def persist_audio_result(
|
||||||
|
chat_id: str,
|
||||||
|
user_id: str,
|
||||||
|
file_name: str,
|
||||||
|
md_path: str,
|
||||||
|
chars: int,
|
||||||
|
md_content: str,
|
||||||
|
oss_read_url: str = "",
|
||||||
|
) -> None:
|
||||||
|
"""将音频转写结果写入 SessionDB。
|
||||||
|
|
||||||
|
管线 A (Import) 和管线 B (Live) 共用此函数。
|
||||||
|
payload 结构统一为 {type:"audio", fileName, mdPath, chars, ossReadUrl, mdContent}。
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
不抛出异常。失败仅 logger.warning。
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from hermes_state import SessionDB # type: ignore
|
||||||
|
db = SessionDB()
|
||||||
|
db.create_session(
|
||||||
|
session_id=chat_id, source="mindos",
|
||||||
|
user_id=user_id, model="audio",
|
||||||
|
)
|
||||||
|
db.append_message(
|
||||||
|
session_id=chat_id,
|
||||||
|
role="assistant",
|
||||||
|
content=json.dumps({
|
||||||
|
"type": "audio",
|
||||||
|
"fileName": file_name,
|
||||||
|
"mdPath": md_path,
|
||||||
|
"chars": chars,
|
||||||
|
"ossReadUrl": oss_read_url,
|
||||||
|
"mdContent": md_content,
|
||||||
|
}, ensure_ascii=False),
|
||||||
|
)
|
||||||
|
logger.info("[voice2md] DB persist ok chatId=%s chars=%d", chat_id, chars)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[voice2md] DB persist failed (non-fatal): %s", e)
|
||||||
|
|
||||||
|
|
||||||
|
# ─── ⑤ sse_push: 统一 SSE md:appended 推送 ──────────────────
|
||||||
|
|
||||||
|
def push_md_appended(
|
||||||
|
sse_server,
|
||||||
|
user_id: str,
|
||||||
|
chat_id: str,
|
||||||
|
file: str,
|
||||||
|
chars: int,
|
||||||
|
md_content: str,
|
||||||
|
message: str = "",
|
||||||
|
source: str = "mic",
|
||||||
|
) -> None:
|
||||||
|
"""推送 SSE md:appended 事件到前端。
|
||||||
|
|
||||||
|
管线 A 和管线 B 共用此函数。
|
||||||
|
chars=0 时跳过推送(无内容不扰前端)。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
sse_server: MindOSSSEServer 实例(具有 _pushEvent 方法)。
|
||||||
|
管线 A 传 self,管线 B 传 _sse_server 模块级变量。
|
||||||
|
source: 音频来源标识("mic" 或 "system"),前端据此区分气泡类型。
|
||||||
|
"""
|
||||||
|
if chars == 0:
|
||||||
|
logger.info("[voice2md] 无转写内容,跳过 SSE 推送 chatId=%s", chat_id)
|
||||||
|
return
|
||||||
|
|
||||||
|
if not sse_server:
|
||||||
|
logger.warning("[voice2md] SSE push 跳过: sse_server 未注入")
|
||||||
|
return
|
||||||
|
|
||||||
|
if not message:
|
||||||
|
message = f"✅ 转写完成({chars} 字)"
|
||||||
|
|
||||||
|
try:
|
||||||
|
sse_server._pushEvent(user_id, "md:appended", {
|
||||||
|
"chatId": chat_id,
|
||||||
|
"file": file,
|
||||||
|
"chars": chars,
|
||||||
|
"mdContent": md_content,
|
||||||
|
"message": message,
|
||||||
|
"source": source,
|
||||||
|
})
|
||||||
|
logger.info("[voice2md] SSE md:appended pushed userId=%s chars=%d source=%s", user_id, chars, source)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[voice2md] SSE push failed: %s", e)
|
||||||
|
|
||||||
|
|
||||||
|
# ─── ASR 积分扣减(管线 A/B 各自计费方式不同,不统一) ─────────
|
||||||
|
|
||||||
|
def deduct_asr_credits(
|
||||||
|
user_id: str,
|
||||||
|
chat_id: str,
|
||||||
|
credits: int,
|
||||||
|
tx_type: str,
|
||||||
|
model: str,
|
||||||
|
seconds: float,
|
||||||
|
) -> None:
|
||||||
|
"""ASR 积分扣减。管线 A 和 B 调用时传不同的 tx_type 和 model。"""
|
||||||
|
try:
|
||||||
|
from hermes_state import SessionDB # type: ignore
|
||||||
|
db = SessionDB()
|
||||||
|
db.deduct_credits(
|
||||||
|
user_id=user_id,
|
||||||
|
credits=credits,
|
||||||
|
tx_type=tx_type,
|
||||||
|
session_id=chat_id,
|
||||||
|
model=model,
|
||||||
|
raw_metric=json.dumps({"seconds": round(seconds, 1)}),
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"[ASR Cost] userId=%s model=%s type=%s seconds=%.1f credits=%d chatId=%s",
|
||||||
|
user_id, model, tx_type, seconds, credits, chat_id,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[voice2md] credit deduction failed: %s", e)
|
||||||
@@ -0,0 +1,99 @@
|
|||||||
|
"""
|
||||||
|
voice_import.py — Voice2MD 管线 A:离线导入(原子组合器)
|
||||||
|
|
||||||
|
Spec 来源: docs/SPEC_voice2md_atomization_design.md §二
|
||||||
|
提取自: mindos_sse.py._runOssTranscribePipeline() 82 行 → 本文件 ~30 行编排
|
||||||
|
|
||||||
|
原子组合:
|
||||||
|
⑦ asr_batch (flash_asr.transcribe_from_oss_url)
|
||||||
|
⑧ md_writer (md_converter.MdConverter)
|
||||||
|
④ db_persist (voice2md_atoms.persist_audio_result)
|
||||||
|
⑤ sse_push (voice2md_atoms.push_md_appended)
|
||||||
|
+ 积分扣减 (voice2md_atoms.deduct_asr_credits)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
|
logger = logging.getLogger("voice_import")
|
||||||
|
|
||||||
|
|
||||||
|
async def run(
|
||||||
|
*,
|
||||||
|
sse_server,
|
||||||
|
user_id: str,
|
||||||
|
chat_id: str,
|
||||||
|
read_url: str,
|
||||||
|
oss_key: str,
|
||||||
|
title: str,
|
||||||
|
) -> None:
|
||||||
|
"""离线音频文件 → Markdown。管线只做编排,不做逻辑。
|
||||||
|
|
||||||
|
调用方: mindos_sse.py 的音频转写路由,通过 asyncio.create_task(voice_import.run(...))
|
||||||
|
"""
|
||||||
|
from flash_asr import transcribe_from_oss_url # type: ignore ⑦ asr_batch
|
||||||
|
from md_converter import MdConverter # type: ignore ⑧ md_writer
|
||||||
|
from voice2md_atoms import ( # type: ignore ④⑤ 共享终点
|
||||||
|
persist_audio_result, push_md_appended, deduct_asr_credits,
|
||||||
|
)
|
||||||
|
|
||||||
|
wiki_root = os.getenv("MINDOS_WIKI_DIR", os.path.expanduser("~/.hermes/wiki"))
|
||||||
|
wiki_dir = os.path.join(wiki_root, user_id)
|
||||||
|
converter = MdConverter(wiki_dir)
|
||||||
|
|
||||||
|
original_filename = oss_key.split("/")[-1]
|
||||||
|
|
||||||
|
# 进度推送:开始
|
||||||
|
sse_server._pushEvent(user_id, "md:progress", {
|
||||||
|
"chatId": chat_id, "stage": "asr_start", "filename": original_filename,
|
||||||
|
})
|
||||||
|
|
||||||
|
# ⑧ md_writer: 创建文件
|
||||||
|
md_file = converter.new_file(title=title, source_filename=original_filename)
|
||||||
|
total_chars = 0
|
||||||
|
last_offset_ms = 0
|
||||||
|
|
||||||
|
# ⑦ asr_batch: paraformer-v2 异步转写
|
||||||
|
async for seg in transcribe_from_oss_url(read_url):
|
||||||
|
if seg.get("text"):
|
||||||
|
offset_ms = seg.get("begin_ms", 0)
|
||||||
|
if offset_ms > last_offset_ms:
|
||||||
|
last_offset_ms = offset_ms
|
||||||
|
total_chars += converter.append_segment(
|
||||||
|
md_file, seg["text"], offset_ms=offset_ms,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ⑧ md_writer: 写结束标记
|
||||||
|
converter.finalize(md_file, char_count=total_chars)
|
||||||
|
rel_path = converter.relative_path(md_file)
|
||||||
|
|
||||||
|
# 积分扣减
|
||||||
|
audio_seconds = last_offset_ms / 1000.0
|
||||||
|
asr_credits = max(1, int(audio_seconds))
|
||||||
|
deduct_asr_credits(
|
||||||
|
user_id=user_id, chat_id=chat_id, credits=asr_credits,
|
||||||
|
tx_type="asr_file", model="paraformer-v2", seconds=audio_seconds,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 读取完整 MD 内容
|
||||||
|
md_content = ""
|
||||||
|
try:
|
||||||
|
md_content = md_file.read_text(encoding="utf-8")
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# ④ db_persist
|
||||||
|
persist_audio_result(
|
||||||
|
chat_id=chat_id, user_id=user_id,
|
||||||
|
file_name=title, md_path=rel_path,
|
||||||
|
chars=total_chars, md_content=md_content,
|
||||||
|
oss_read_url=read_url,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ⑤ sse_push
|
||||||
|
push_md_appended(
|
||||||
|
sse_server=sse_server, user_id=user_id, chat_id=chat_id,
|
||||||
|
file=rel_path, chars=total_chars, md_content=md_content,
|
||||||
|
message=f"✅ 会议记录已入库:{md_file.name}({total_chars} 字)",
|
||||||
|
)
|
||||||
@@ -0,0 +1,672 @@
|
|||||||
|
"""Generic webhook platform adapter.
|
||||||
|
|
||||||
|
Runs an aiohttp HTTP server that receives webhook POSTs from external
|
||||||
|
services (GitHub, GitLab, JIRA, Stripe, etc.), validates HMAC signatures,
|
||||||
|
transforms payloads into agent prompts, and routes responses back to the
|
||||||
|
source or to another configured platform.
|
||||||
|
|
||||||
|
Configuration lives in config.yaml under platforms.webhook.extra.routes.
|
||||||
|
Each route defines:
|
||||||
|
- events: which event types to accept (header-based filtering)
|
||||||
|
- secret: HMAC secret for signature validation (REQUIRED)
|
||||||
|
- prompt: template string formatted with the webhook payload
|
||||||
|
- skills: optional list of skills to load for the agent
|
||||||
|
- deliver: where to send the response (github_comment, telegram, etc.)
|
||||||
|
- deliver_extra: additional delivery config (repo, pr_number, chat_id)
|
||||||
|
|
||||||
|
Security:
|
||||||
|
- HMAC secret is required per route (validated at startup)
|
||||||
|
- Rate limiting per route (fixed-window, configurable)
|
||||||
|
- Idempotency cache prevents duplicate agent runs on webhook retries
|
||||||
|
- Body size limits checked before reading payload
|
||||||
|
- Set secret to "INSECURE_NO_AUTH" to skip validation (testing only)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import hashlib
|
||||||
|
import hmac
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
import subprocess
|
||||||
|
import time
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
try:
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
AIOHTTP_AVAILABLE = True
|
||||||
|
except ImportError:
|
||||||
|
AIOHTTP_AVAILABLE = False
|
||||||
|
web = None # type: ignore[assignment]
|
||||||
|
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
from gateway.platforms.base import (
|
||||||
|
BasePlatformAdapter,
|
||||||
|
MessageEvent,
|
||||||
|
MessageType,
|
||||||
|
SendResult,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
DEFAULT_HOST = "0.0.0.0"
|
||||||
|
DEFAULT_PORT = 8644
|
||||||
|
_INSECURE_NO_AUTH = "INSECURE_NO_AUTH"
|
||||||
|
_DYNAMIC_ROUTES_FILENAME = "webhook_subscriptions.json"
|
||||||
|
|
||||||
|
|
||||||
|
def check_webhook_requirements() -> bool:
|
||||||
|
"""Check if webhook adapter dependencies are available."""
|
||||||
|
return AIOHTTP_AVAILABLE
|
||||||
|
|
||||||
|
|
||||||
|
class WebhookAdapter(BasePlatformAdapter):
|
||||||
|
"""Generic webhook receiver that triggers agent runs from HTTP POSTs."""
|
||||||
|
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.WEBHOOK)
|
||||||
|
self._host: str = config.extra.get("host", DEFAULT_HOST)
|
||||||
|
self._port: int = int(config.extra.get("port", DEFAULT_PORT))
|
||||||
|
self._global_secret: str = config.extra.get("secret", "")
|
||||||
|
self._static_routes: Dict[str, dict] = config.extra.get("routes", {})
|
||||||
|
self._dynamic_routes: Dict[str, dict] = {}
|
||||||
|
self._dynamic_routes_mtime: float = 0.0
|
||||||
|
self._routes: Dict[str, dict] = dict(self._static_routes)
|
||||||
|
self._runner = None
|
||||||
|
|
||||||
|
# Delivery info keyed by session chat_id.
|
||||||
|
#
|
||||||
|
# Read by every send() invocation for the chat_id (status messages
|
||||||
|
# AND the final response). Cleaned up via TTL on each POST so the
|
||||||
|
# dict stays bounded — see _prune_delivery_info(). Do NOT pop on
|
||||||
|
# send(), or interim status messages (e.g. fallback notifications,
|
||||||
|
# context-pressure warnings) will consume the entry before the
|
||||||
|
# final response arrives, causing the response to silently fall
|
||||||
|
# back to the "log" deliver type.
|
||||||
|
self._delivery_info: Dict[str, dict] = {}
|
||||||
|
self._delivery_info_created: Dict[str, float] = {}
|
||||||
|
|
||||||
|
# Reference to gateway runner for cross-platform delivery (set externally)
|
||||||
|
self.gateway_runner = None
|
||||||
|
|
||||||
|
# Idempotency: TTL cache of recently processed delivery IDs.
|
||||||
|
# Prevents duplicate agent runs when webhook providers retry.
|
||||||
|
self._seen_deliveries: Dict[str, float] = {}
|
||||||
|
self._idempotency_ttl: int = 3600 # 1 hour
|
||||||
|
|
||||||
|
# Rate limiting: per-route timestamps in a fixed window.
|
||||||
|
self._rate_counts: Dict[str, List[float]] = {}
|
||||||
|
self._rate_limit: int = int(config.extra.get("rate_limit", 30)) # per minute
|
||||||
|
|
||||||
|
# Body size limit (auth-before-body pattern)
|
||||||
|
self._max_body_bytes: int = int(
|
||||||
|
config.extra.get("max_body_bytes", 1_048_576)
|
||||||
|
) # 1MB
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Lifecycle
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
# Load agent-created subscriptions before validating
|
||||||
|
self._reload_dynamic_routes()
|
||||||
|
|
||||||
|
# Validate routes at startup — secret is required per route
|
||||||
|
for name, route in self._routes.items():
|
||||||
|
secret = route.get("secret", self._global_secret)
|
||||||
|
if not secret:
|
||||||
|
raise ValueError(
|
||||||
|
f"[webhook] Route '{name}' has no HMAC secret. "
|
||||||
|
f"Set 'secret' on the route or globally. "
|
||||||
|
f"For testing without auth, set secret to '{_INSECURE_NO_AUTH}'."
|
||||||
|
)
|
||||||
|
|
||||||
|
app = web.Application()
|
||||||
|
app.router.add_get("/health", self._handle_health)
|
||||||
|
app.router.add_post("/webhooks/{route_name}", self._handle_webhook)
|
||||||
|
|
||||||
|
# Port conflict detection — fail fast if port is already in use
|
||||||
|
import socket as _socket
|
||||||
|
try:
|
||||||
|
with _socket.socket(_socket.AF_INET, _socket.SOCK_STREAM) as _s:
|
||||||
|
_s.settimeout(1)
|
||||||
|
_s.connect(('127.0.0.1', self._port))
|
||||||
|
logger.error('[webhook] Port %d already in use. Set a different port in config.yaml: platforms.webhook.port', self._port)
|
||||||
|
return False
|
||||||
|
except (ConnectionRefusedError, OSError):
|
||||||
|
pass # port is free
|
||||||
|
|
||||||
|
self._runner = web.AppRunner(app)
|
||||||
|
await self._runner.setup()
|
||||||
|
site = web.TCPSite(self._runner, self._host, self._port)
|
||||||
|
await site.start()
|
||||||
|
self._mark_connected()
|
||||||
|
|
||||||
|
route_names = ", ".join(self._routes.keys()) or "(none configured)"
|
||||||
|
logger.info(
|
||||||
|
"[webhook] Listening on %s:%d — routes: %s",
|
||||||
|
self._host,
|
||||||
|
self._port,
|
||||||
|
route_names,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
if self._runner:
|
||||||
|
await self._runner.cleanup()
|
||||||
|
self._runner = None
|
||||||
|
self._mark_disconnected()
|
||||||
|
logger.info("[webhook] Disconnected")
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Deliver the agent's response to the configured destination.
|
||||||
|
|
||||||
|
chat_id is ``webhook:{route}:{delivery_id}``. The delivery info
|
||||||
|
stored during webhook receipt is read with ``.get()`` (not popped)
|
||||||
|
so that interim status messages emitted before the final response
|
||||||
|
— fallback-model notifications, context-pressure warnings, etc. —
|
||||||
|
do not consume the entry and silently downgrade the final response
|
||||||
|
to the ``log`` deliver type. TTL cleanup happens on POST.
|
||||||
|
"""
|
||||||
|
delivery = self._delivery_info.get(chat_id, {})
|
||||||
|
deliver_type = delivery.get("deliver", "log")
|
||||||
|
|
||||||
|
if deliver_type == "log":
|
||||||
|
logger.info("[webhook] Response for %s: %s", chat_id, content[:200])
|
||||||
|
return SendResult(success=True)
|
||||||
|
|
||||||
|
if deliver_type == "github_comment":
|
||||||
|
return await self._deliver_github_comment(content, delivery)
|
||||||
|
|
||||||
|
# Cross-platform delivery — any platform with a gateway adapter
|
||||||
|
if self.gateway_runner and deliver_type in (
|
||||||
|
"telegram",
|
||||||
|
"discord",
|
||||||
|
"slack",
|
||||||
|
"signal",
|
||||||
|
"sms",
|
||||||
|
"whatsapp",
|
||||||
|
"matrix",
|
||||||
|
"mattermost",
|
||||||
|
"homeassistant",
|
||||||
|
"email",
|
||||||
|
"dingtalk",
|
||||||
|
"feishu",
|
||||||
|
"wecom",
|
||||||
|
"wecom_callback",
|
||||||
|
"weixin",
|
||||||
|
"bluebubbles",
|
||||||
|
"qqbot",
|
||||||
|
):
|
||||||
|
return await self._deliver_cross_platform(
|
||||||
|
deliver_type, content, delivery
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.warning("[webhook] Unknown deliver type: %s", deliver_type)
|
||||||
|
return SendResult(
|
||||||
|
success=False, error=f"Unknown deliver type: {deliver_type}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _prune_delivery_info(self, now: float) -> None:
|
||||||
|
"""Drop delivery_info entries older than the idempotency TTL.
|
||||||
|
|
||||||
|
Mirrors the cleanup pattern used for ``_seen_deliveries``. Called
|
||||||
|
on each POST so the dict size is bounded by ``rate_limit * TTL``
|
||||||
|
even if many webhooks fire and never receive a final response.
|
||||||
|
"""
|
||||||
|
cutoff = now - self._idempotency_ttl
|
||||||
|
stale = [
|
||||||
|
k
|
||||||
|
for k, t in self._delivery_info_created.items()
|
||||||
|
if t < cutoff
|
||||||
|
]
|
||||||
|
for k in stale:
|
||||||
|
self._delivery_info.pop(k, None)
|
||||||
|
self._delivery_info_created.pop(k, None)
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
return {"name": chat_id, "type": "webhook"}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# HTTP handlers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _handle_health(self, request: "web.Request") -> "web.Response":
|
||||||
|
"""GET /health — simple health check."""
|
||||||
|
return web.json_response({"status": "ok", "platform": "webhook"})
|
||||||
|
|
||||||
|
def _reload_dynamic_routes(self) -> None:
|
||||||
|
"""Reload agent-created subscriptions from disk if the file changed."""
|
||||||
|
from hermes_constants import get_hermes_home
|
||||||
|
hermes_home = get_hermes_home()
|
||||||
|
subs_path = hermes_home / _DYNAMIC_ROUTES_FILENAME
|
||||||
|
if not subs_path.exists():
|
||||||
|
if self._dynamic_routes:
|
||||||
|
self._dynamic_routes = {}
|
||||||
|
self._routes = dict(self._static_routes)
|
||||||
|
logger.debug("[webhook] Dynamic subscriptions file removed, cleared dynamic routes")
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
mtime = subs_path.stat().st_mtime
|
||||||
|
if mtime <= self._dynamic_routes_mtime:
|
||||||
|
return # No change
|
||||||
|
data = json.loads(subs_path.read_text(encoding="utf-8"))
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
return
|
||||||
|
# Merge: static routes take precedence over dynamic ones
|
||||||
|
self._dynamic_routes = {
|
||||||
|
k: v for k, v in data.items()
|
||||||
|
if k not in self._static_routes
|
||||||
|
}
|
||||||
|
self._routes = {**self._dynamic_routes, **self._static_routes}
|
||||||
|
self._dynamic_routes_mtime = mtime
|
||||||
|
logger.info(
|
||||||
|
"[webhook] Reloaded %d dynamic route(s): %s",
|
||||||
|
len(self._dynamic_routes),
|
||||||
|
", ".join(self._dynamic_routes.keys()) or "(none)",
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[webhook] Failed to reload dynamic routes: %s", e)
|
||||||
|
|
||||||
|
async def _handle_webhook(self, request: "web.Request") -> "web.Response":
|
||||||
|
"""POST /webhooks/{route_name} — receive and process a webhook event."""
|
||||||
|
# Hot-reload dynamic subscriptions on each request (mtime-gated, cheap)
|
||||||
|
self._reload_dynamic_routes()
|
||||||
|
|
||||||
|
route_name = request.match_info.get("route_name", "")
|
||||||
|
route_config = self._routes.get(route_name)
|
||||||
|
|
||||||
|
if not route_config:
|
||||||
|
return web.json_response(
|
||||||
|
{"error": f"Unknown route: {route_name}"}, status=404
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── Auth-before-body ─────────────────────────────────────
|
||||||
|
# Check Content-Length before reading the full payload.
|
||||||
|
content_length = request.content_length or 0
|
||||||
|
if content_length > self._max_body_bytes:
|
||||||
|
return web.json_response(
|
||||||
|
{"error": "Payload too large"}, status=413
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── Rate limiting ────────────────────────────────────────
|
||||||
|
now = time.time()
|
||||||
|
window = self._rate_counts.setdefault(route_name, [])
|
||||||
|
window[:] = [t for t in window if now - t < 60]
|
||||||
|
if len(window) >= self._rate_limit:
|
||||||
|
return web.json_response(
|
||||||
|
{"error": "Rate limit exceeded"}, status=429
|
||||||
|
)
|
||||||
|
window.append(now)
|
||||||
|
|
||||||
|
# Read body
|
||||||
|
try:
|
||||||
|
raw_body = await request.read()
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[webhook] Failed to read body: %s", e)
|
||||||
|
return web.json_response({"error": "Bad request"}, status=400)
|
||||||
|
|
||||||
|
# Validate HMAC signature (skip for INSECURE_NO_AUTH testing mode)
|
||||||
|
secret = route_config.get("secret", self._global_secret)
|
||||||
|
if secret and secret != _INSECURE_NO_AUTH:
|
||||||
|
if not self._validate_signature(request, raw_body, secret):
|
||||||
|
logger.warning(
|
||||||
|
"[webhook] Invalid signature for route %s", route_name
|
||||||
|
)
|
||||||
|
return web.json_response(
|
||||||
|
{"error": "Invalid signature"}, status=401
|
||||||
|
)
|
||||||
|
|
||||||
|
# Parse payload
|
||||||
|
try:
|
||||||
|
payload = json.loads(raw_body)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
# Try form-encoded as fallback
|
||||||
|
try:
|
||||||
|
import urllib.parse
|
||||||
|
|
||||||
|
payload = dict(
|
||||||
|
urllib.parse.parse_qsl(raw_body.decode("utf-8"))
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
return web.json_response(
|
||||||
|
{"error": "Cannot parse body"}, status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check event type filter
|
||||||
|
event_type = (
|
||||||
|
request.headers.get("X-GitHub-Event", "")
|
||||||
|
or request.headers.get("X-GitLab-Event", "")
|
||||||
|
or payload.get("event_type", "")
|
||||||
|
or "unknown"
|
||||||
|
)
|
||||||
|
allowed_events = route_config.get("events", [])
|
||||||
|
if allowed_events and event_type not in allowed_events:
|
||||||
|
logger.debug(
|
||||||
|
"[webhook] Ignoring event %s for route %s (allowed: %s)",
|
||||||
|
event_type,
|
||||||
|
route_name,
|
||||||
|
allowed_events,
|
||||||
|
)
|
||||||
|
return web.json_response(
|
||||||
|
{"status": "ignored", "event": event_type}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Format prompt from template
|
||||||
|
prompt_template = route_config.get("prompt", "")
|
||||||
|
prompt = self._render_prompt(
|
||||||
|
prompt_template, payload, event_type, route_name
|
||||||
|
)
|
||||||
|
|
||||||
|
# Inject skill content if configured.
|
||||||
|
# We call build_skill_invocation_message() directly rather than
|
||||||
|
# using /skill-name slash commands — the gateway's command parser
|
||||||
|
# would intercept those and break the flow.
|
||||||
|
skills = route_config.get("skills", [])
|
||||||
|
if skills:
|
||||||
|
try:
|
||||||
|
from agent.skill_commands import (
|
||||||
|
build_skill_invocation_message,
|
||||||
|
get_skill_commands,
|
||||||
|
)
|
||||||
|
|
||||||
|
skill_cmds = get_skill_commands()
|
||||||
|
for skill_name in skills:
|
||||||
|
cmd_key = f"/{skill_name}"
|
||||||
|
if cmd_key in skill_cmds:
|
||||||
|
skill_content = build_skill_invocation_message(
|
||||||
|
cmd_key, user_instruction=prompt
|
||||||
|
)
|
||||||
|
if skill_content:
|
||||||
|
prompt = skill_content
|
||||||
|
break # Load the first matching skill
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"[webhook] Skill '%s' not found", skill_name
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[webhook] Skill loading failed: %s", e)
|
||||||
|
|
||||||
|
# Build a unique delivery ID
|
||||||
|
delivery_id = request.headers.get(
|
||||||
|
"X-GitHub-Delivery",
|
||||||
|
request.headers.get("X-Request-ID", str(int(time.time() * 1000))),
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── Idempotency ─────────────────────────────────────────
|
||||||
|
# Skip duplicate deliveries (webhook retries).
|
||||||
|
now = time.time()
|
||||||
|
# Prune expired entries
|
||||||
|
self._seen_deliveries = {
|
||||||
|
k: v
|
||||||
|
for k, v in self._seen_deliveries.items()
|
||||||
|
if now - v < self._idempotency_ttl
|
||||||
|
}
|
||||||
|
if delivery_id in self._seen_deliveries:
|
||||||
|
logger.info(
|
||||||
|
"[webhook] Skipping duplicate delivery %s", delivery_id
|
||||||
|
)
|
||||||
|
return web.json_response(
|
||||||
|
{"status": "duplicate", "delivery_id": delivery_id},
|
||||||
|
status=200,
|
||||||
|
)
|
||||||
|
self._seen_deliveries[delivery_id] = now
|
||||||
|
|
||||||
|
# Use delivery_id in session key so concurrent webhooks on the
|
||||||
|
# same route get independent agent runs (not queued/interrupted).
|
||||||
|
session_chat_id = f"webhook:{route_name}:{delivery_id}"
|
||||||
|
|
||||||
|
# Store delivery info for send(). Read by every send() invocation
|
||||||
|
# for this chat_id (interim status messages and the final response),
|
||||||
|
# so we do NOT pop on send. TTL-based cleanup keeps the dict bounded.
|
||||||
|
deliver_config = {
|
||||||
|
"deliver": route_config.get("deliver", "log"),
|
||||||
|
"deliver_extra": self._render_delivery_extra(
|
||||||
|
route_config.get("deliver_extra", {}), payload
|
||||||
|
),
|
||||||
|
"payload": payload,
|
||||||
|
}
|
||||||
|
self._delivery_info[session_chat_id] = deliver_config
|
||||||
|
self._delivery_info_created[session_chat_id] = now
|
||||||
|
self._prune_delivery_info(now)
|
||||||
|
|
||||||
|
# Build source and event
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id=session_chat_id,
|
||||||
|
chat_name=f"webhook/{route_name}",
|
||||||
|
chat_type="webhook",
|
||||||
|
user_id=f"webhook:{route_name}",
|
||||||
|
user_name=route_name,
|
||||||
|
)
|
||||||
|
event = MessageEvent(
|
||||||
|
text=prompt,
|
||||||
|
message_type=MessageType.TEXT,
|
||||||
|
source=source,
|
||||||
|
raw_message=payload,
|
||||||
|
message_id=delivery_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"[webhook] %s event=%s route=%s prompt_len=%d delivery=%s",
|
||||||
|
request.method,
|
||||||
|
event_type,
|
||||||
|
route_name,
|
||||||
|
len(prompt),
|
||||||
|
delivery_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Non-blocking — return 202 Accepted immediately
|
||||||
|
task = asyncio.create_task(self.handle_message(event))
|
||||||
|
self._background_tasks.add(task)
|
||||||
|
task.add_done_callback(self._background_tasks.discard)
|
||||||
|
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"status": "accepted",
|
||||||
|
"route": route_name,
|
||||||
|
"event": event_type,
|
||||||
|
"delivery_id": delivery_id,
|
||||||
|
},
|
||||||
|
status=202,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Signature validation
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _validate_signature(
|
||||||
|
self, request: "web.Request", body: bytes, secret: str
|
||||||
|
) -> bool:
|
||||||
|
"""Validate webhook signature (GitHub, GitLab, generic HMAC-SHA256)."""
|
||||||
|
# GitHub: X-Hub-Signature-256 = sha256=<hex>
|
||||||
|
gh_sig = request.headers.get("X-Hub-Signature-256", "")
|
||||||
|
if gh_sig:
|
||||||
|
expected = "sha256=" + hmac.new(
|
||||||
|
secret.encode(), body, hashlib.sha256
|
||||||
|
).hexdigest()
|
||||||
|
return hmac.compare_digest(gh_sig, expected)
|
||||||
|
|
||||||
|
# GitLab: X-Gitlab-Token = <plain secret>
|
||||||
|
gl_token = request.headers.get("X-Gitlab-Token", "")
|
||||||
|
if gl_token:
|
||||||
|
return hmac.compare_digest(gl_token, secret)
|
||||||
|
|
||||||
|
# Generic: X-Webhook-Signature = <hex HMAC-SHA256>
|
||||||
|
generic_sig = request.headers.get("X-Webhook-Signature", "")
|
||||||
|
if generic_sig:
|
||||||
|
expected = hmac.new(
|
||||||
|
secret.encode(), body, hashlib.sha256
|
||||||
|
).hexdigest()
|
||||||
|
return hmac.compare_digest(generic_sig, expected)
|
||||||
|
|
||||||
|
# No recognised signature header but secret is configured → reject
|
||||||
|
logger.debug(
|
||||||
|
"[webhook] Secret configured but no signature header found"
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Prompt rendering
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _render_prompt(
|
||||||
|
self,
|
||||||
|
template: str,
|
||||||
|
payload: dict,
|
||||||
|
event_type: str,
|
||||||
|
route_name: str,
|
||||||
|
) -> str:
|
||||||
|
"""Render a prompt template with the webhook payload.
|
||||||
|
|
||||||
|
Supports dot-notation access into nested dicts:
|
||||||
|
``{pull_request.title}`` → ``payload["pull_request"]["title"]``
|
||||||
|
|
||||||
|
Special token ``{__raw__}`` dumps the entire payload as indented
|
||||||
|
JSON (truncated to 4000 chars). Useful for monitoring alerts or
|
||||||
|
any webhook where the agent needs to see the full payload.
|
||||||
|
"""
|
||||||
|
if not template:
|
||||||
|
truncated = json.dumps(payload, indent=2)[:4000]
|
||||||
|
return (
|
||||||
|
f"Webhook event '{event_type}' on route "
|
||||||
|
f"'{route_name}':\n\n```json\n{truncated}\n```"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _resolve(match: re.Match) -> str:
|
||||||
|
key = match.group(1)
|
||||||
|
# Special token: dump the entire payload as JSON
|
||||||
|
if key == "__raw__":
|
||||||
|
return json.dumps(payload, indent=2)[:4000]
|
||||||
|
value: Any = payload
|
||||||
|
for part in key.split("."):
|
||||||
|
if isinstance(value, dict):
|
||||||
|
value = value.get(part, f"{{{key}}}")
|
||||||
|
else:
|
||||||
|
return f"{{{key}}}"
|
||||||
|
if isinstance(value, (dict, list)):
|
||||||
|
return json.dumps(value, indent=2)[:2000]
|
||||||
|
return str(value)
|
||||||
|
|
||||||
|
return re.sub(r"\{([a-zA-Z0-9_.]+)\}", _resolve, template)
|
||||||
|
|
||||||
|
def _render_delivery_extra(
|
||||||
|
self, extra: dict, payload: dict
|
||||||
|
) -> dict:
|
||||||
|
"""Render delivery_extra template values with payload data."""
|
||||||
|
rendered: Dict[str, Any] = {}
|
||||||
|
for key, value in extra.items():
|
||||||
|
if isinstance(value, str):
|
||||||
|
rendered[key] = self._render_prompt(value, payload, "", "")
|
||||||
|
else:
|
||||||
|
rendered[key] = value
|
||||||
|
return rendered
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Response delivery
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _deliver_github_comment(
|
||||||
|
self, content: str, delivery: dict
|
||||||
|
) -> SendResult:
|
||||||
|
"""Post agent response as a GitHub PR/issue comment via ``gh`` CLI."""
|
||||||
|
extra = delivery.get("deliver_extra", {})
|
||||||
|
repo = extra.get("repo", "")
|
||||||
|
pr_number = extra.get("pr_number", "")
|
||||||
|
|
||||||
|
if not repo or not pr_number:
|
||||||
|
logger.error(
|
||||||
|
"[webhook] github_comment delivery missing repo or pr_number"
|
||||||
|
)
|
||||||
|
return SendResult(
|
||||||
|
success=False, error="Missing repo or pr_number"
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
[
|
||||||
|
"gh",
|
||||||
|
"pr",
|
||||||
|
"comment",
|
||||||
|
str(pr_number),
|
||||||
|
"--repo",
|
||||||
|
repo,
|
||||||
|
"--body",
|
||||||
|
content,
|
||||||
|
],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
if result.returncode == 0:
|
||||||
|
logger.info(
|
||||||
|
"[webhook] Posted comment on %s#%s", repo, pr_number
|
||||||
|
)
|
||||||
|
return SendResult(success=True)
|
||||||
|
else:
|
||||||
|
logger.error(
|
||||||
|
"[webhook] gh pr comment failed: %s", result.stderr
|
||||||
|
)
|
||||||
|
return SendResult(success=False, error=result.stderr)
|
||||||
|
except FileNotFoundError:
|
||||||
|
logger.error(
|
||||||
|
"[webhook] 'gh' CLI not found — install GitHub CLI for "
|
||||||
|
"github_comment delivery"
|
||||||
|
)
|
||||||
|
return SendResult(
|
||||||
|
success=False, error="gh CLI not installed"
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[webhook] github_comment delivery error: %s", e)
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
async def _deliver_cross_platform(
|
||||||
|
self, platform_name: str, content: str, delivery: dict
|
||||||
|
) -> SendResult:
|
||||||
|
"""Route response to another platform (telegram, discord, etc.)."""
|
||||||
|
if not self.gateway_runner:
|
||||||
|
return SendResult(
|
||||||
|
success=False,
|
||||||
|
error="No gateway runner for cross-platform delivery",
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
target_platform = Platform(platform_name)
|
||||||
|
except ValueError:
|
||||||
|
return SendResult(
|
||||||
|
success=False, error=f"Unknown platform: {platform_name}"
|
||||||
|
)
|
||||||
|
|
||||||
|
adapter = self.gateway_runner.adapters.get(target_platform)
|
||||||
|
if not adapter:
|
||||||
|
return SendResult(
|
||||||
|
success=False,
|
||||||
|
error=f"Platform {platform_name} not connected",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Use home channel if no specific chat_id in deliver_extra
|
||||||
|
extra = delivery.get("deliver_extra", {})
|
||||||
|
chat_id = extra.get("chat_id", "")
|
||||||
|
if not chat_id:
|
||||||
|
home = self.gateway_runner.config.get_home_channel(target_platform)
|
||||||
|
if home:
|
||||||
|
chat_id = home.chat_id
|
||||||
|
else:
|
||||||
|
return SendResult(
|
||||||
|
success=False,
|
||||||
|
error=f"No chat_id or home channel for {platform_name}",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Pass thread_id from deliver_extra so Telegram forum topics work
|
||||||
|
metadata = None
|
||||||
|
thread_id = extra.get("message_thread_id") or extra.get("thread_id")
|
||||||
|
if thread_id:
|
||||||
|
metadata = {"thread_id": thread_id}
|
||||||
|
|
||||||
|
return await adapter.send(chat_id, content, metadata=metadata)
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,387 @@
|
|||||||
|
"""WeCom callback-mode adapter for self-built enterprise applications.
|
||||||
|
|
||||||
|
Unlike the bot/websocket adapter in ``wecom.py``, this handles the standard
|
||||||
|
WeCom callback flow: WeCom POSTs encrypted XML to an HTTP endpoint, the
|
||||||
|
adapter decrypts it, queues the message for the agent, and immediately
|
||||||
|
acknowledges. The agent's reply is delivered later via the proactive
|
||||||
|
``message/send`` API using an access-token.
|
||||||
|
|
||||||
|
Supports multiple self-built apps under one gateway instance, scoped by
|
||||||
|
``corp_id:user_id`` to avoid cross-corp collisions.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import socket as _socket
|
||||||
|
import time
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
from xml.etree import ElementTree as ET
|
||||||
|
|
||||||
|
try:
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
AIOHTTP_AVAILABLE = True
|
||||||
|
except ImportError:
|
||||||
|
web = None # type: ignore[assignment]
|
||||||
|
AIOHTTP_AVAILABLE = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
HTTPX_AVAILABLE = True
|
||||||
|
except ImportError:
|
||||||
|
httpx = None # type: ignore[assignment]
|
||||||
|
HTTPX_AVAILABLE = False
|
||||||
|
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
from gateway.platforms.base import BasePlatformAdapter, MessageEvent, MessageType, SendResult
|
||||||
|
from gateway.platforms.wecom_crypto import WXBizMsgCrypt, WeComCryptoError
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
DEFAULT_HOST = "0.0.0.0"
|
||||||
|
DEFAULT_PORT = 8645
|
||||||
|
DEFAULT_PATH = "/wecom/callback"
|
||||||
|
ACCESS_TOKEN_TTL_SECONDS = 7200
|
||||||
|
MESSAGE_DEDUP_TTL_SECONDS = 300
|
||||||
|
|
||||||
|
|
||||||
|
def check_wecom_callback_requirements() -> bool:
|
||||||
|
return AIOHTTP_AVAILABLE and HTTPX_AVAILABLE
|
||||||
|
|
||||||
|
|
||||||
|
class WecomCallbackAdapter(BasePlatformAdapter):
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.WECOM_CALLBACK)
|
||||||
|
extra = config.extra or {}
|
||||||
|
self._host = str(extra.get("host") or DEFAULT_HOST)
|
||||||
|
self._port = int(extra.get("port") or DEFAULT_PORT)
|
||||||
|
self._path = str(extra.get("path") or DEFAULT_PATH)
|
||||||
|
self._apps: List[Dict[str, Any]] = self._normalize_apps(extra)
|
||||||
|
self._runner: Optional[web.AppRunner] = None
|
||||||
|
self._site: Optional[web.TCPSite] = None
|
||||||
|
self._app: Optional[web.Application] = None
|
||||||
|
self._http_client: Optional[httpx.AsyncClient] = None
|
||||||
|
self._message_queue: asyncio.Queue[MessageEvent] = asyncio.Queue()
|
||||||
|
self._poll_task: Optional[asyncio.Task] = None
|
||||||
|
self._seen_messages: Dict[str, float] = {}
|
||||||
|
self._user_app_map: Dict[str, str] = {}
|
||||||
|
self._access_tokens: Dict[str, Dict[str, Any]] = {}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# App normalisation
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _user_app_key(corp_id: str, user_id: str) -> str:
|
||||||
|
return f"{corp_id}:{user_id}" if corp_id else user_id
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_apps(extra: Dict[str, Any]) -> List[Dict[str, Any]]:
|
||||||
|
apps = extra.get("apps")
|
||||||
|
if isinstance(apps, list) and apps:
|
||||||
|
return [dict(app) for app in apps if isinstance(app, dict)]
|
||||||
|
if extra.get("corp_id"):
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"name": extra.get("name") or "default",
|
||||||
|
"corp_id": extra.get("corp_id", ""),
|
||||||
|
"corp_secret": extra.get("corp_secret", ""),
|
||||||
|
"agent_id": str(extra.get("agent_id", "")),
|
||||||
|
"token": extra.get("token", ""),
|
||||||
|
"encoding_aes_key": extra.get("encoding_aes_key", ""),
|
||||||
|
}
|
||||||
|
]
|
||||||
|
return []
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Lifecycle
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
if not self._apps:
|
||||||
|
logger.warning("[WecomCallback] No callback apps configured")
|
||||||
|
return False
|
||||||
|
if not check_wecom_callback_requirements():
|
||||||
|
logger.warning("[WecomCallback] aiohttp/httpx not installed")
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Quick port-in-use check.
|
||||||
|
try:
|
||||||
|
with _socket.socket(_socket.AF_INET, _socket.SOCK_STREAM) as sock:
|
||||||
|
sock.settimeout(1)
|
||||||
|
sock.connect(("127.0.0.1", self._port))
|
||||||
|
logger.error("[WecomCallback] Port %d already in use", self._port)
|
||||||
|
return False
|
||||||
|
except (ConnectionRefusedError, OSError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
try:
|
||||||
|
self._http_client = httpx.AsyncClient(timeout=20.0)
|
||||||
|
self._app = web.Application()
|
||||||
|
self._app.router.add_get("/health", self._handle_health)
|
||||||
|
self._app.router.add_get(self._path, self._handle_verify)
|
||||||
|
self._app.router.add_post(self._path, self._handle_callback)
|
||||||
|
self._runner = web.AppRunner(self._app)
|
||||||
|
await self._runner.setup()
|
||||||
|
self._site = web.TCPSite(self._runner, self._host, self._port)
|
||||||
|
await self._site.start()
|
||||||
|
self._poll_task = asyncio.create_task(self._poll_loop())
|
||||||
|
self._mark_connected()
|
||||||
|
logger.info(
|
||||||
|
"[WecomCallback] HTTP server listening on %s:%s%s",
|
||||||
|
self._host, self._port, self._path,
|
||||||
|
)
|
||||||
|
for app in self._apps:
|
||||||
|
try:
|
||||||
|
await self._refresh_access_token(app)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"[WecomCallback] Initial token refresh failed for app '%s': %s",
|
||||||
|
app.get("name", "default"), exc,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
except Exception:
|
||||||
|
await self._cleanup()
|
||||||
|
logger.exception("[WecomCallback] Failed to start")
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
self._running = False
|
||||||
|
if self._poll_task:
|
||||||
|
self._poll_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._poll_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
self._poll_task = None
|
||||||
|
await self._cleanup()
|
||||||
|
self._mark_disconnected()
|
||||||
|
logger.info("[WecomCallback] Disconnected")
|
||||||
|
|
||||||
|
async def _cleanup(self) -> None:
|
||||||
|
self._site = None
|
||||||
|
if self._runner:
|
||||||
|
await self._runner.cleanup()
|
||||||
|
self._runner = None
|
||||||
|
self._app = None
|
||||||
|
if self._http_client:
|
||||||
|
await self._http_client.aclose()
|
||||||
|
self._http_client = None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Outbound: proactive send via access-token API
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
app = self._resolve_app_for_chat(chat_id)
|
||||||
|
touser = chat_id.split(":", 1)[1] if ":" in chat_id else chat_id
|
||||||
|
try:
|
||||||
|
token = await self._get_access_token(app)
|
||||||
|
payload = {
|
||||||
|
"touser": touser,
|
||||||
|
"msgtype": "text",
|
||||||
|
"agentid": int(str(app.get("agent_id") or 0)),
|
||||||
|
"text": {"content": content[:2048]},
|
||||||
|
"safe": 0,
|
||||||
|
}
|
||||||
|
resp = await self._http_client.post(
|
||||||
|
f"https://qyapi.weixin.qq.com/cgi-bin/message/send?access_token={token}",
|
||||||
|
json=payload,
|
||||||
|
)
|
||||||
|
data = resp.json()
|
||||||
|
if data.get("errcode") != 0:
|
||||||
|
return SendResult(success=False, error=str(data))
|
||||||
|
return SendResult(
|
||||||
|
success=True,
|
||||||
|
message_id=str(data.get("msgid", "")),
|
||||||
|
raw_response=data,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
return SendResult(success=False, error=str(exc))
|
||||||
|
|
||||||
|
def _resolve_app_for_chat(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
"""Pick the app associated with *chat_id*, falling back sensibly."""
|
||||||
|
app_name = self._user_app_map.get(chat_id)
|
||||||
|
if not app_name and ":" not in chat_id:
|
||||||
|
# Legacy bare user_id — try to find a unique match.
|
||||||
|
matching = [k for k in self._user_app_map if k.endswith(f":{chat_id}")]
|
||||||
|
if len(matching) == 1:
|
||||||
|
app_name = self._user_app_map.get(matching[0])
|
||||||
|
app = self._get_app_by_name(app_name) if app_name else None
|
||||||
|
return app or self._apps[0]
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
return {"name": chat_id, "type": "dm"}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Inbound: HTTP callback handlers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _handle_health(self, request: web.Request) -> web.Response:
|
||||||
|
return web.json_response({"status": "ok", "platform": "wecom_callback"})
|
||||||
|
|
||||||
|
async def _handle_verify(self, request: web.Request) -> web.Response:
|
||||||
|
"""GET endpoint — WeCom URL verification handshake."""
|
||||||
|
msg_signature = request.query.get("msg_signature", "")
|
||||||
|
timestamp = request.query.get("timestamp", "")
|
||||||
|
nonce = request.query.get("nonce", "")
|
||||||
|
echostr = request.query.get("echostr", "")
|
||||||
|
for app in self._apps:
|
||||||
|
try:
|
||||||
|
crypt = self._crypt_for_app(app)
|
||||||
|
plain = crypt.verify_url(msg_signature, timestamp, nonce, echostr)
|
||||||
|
return web.Response(text=plain, content_type="text/plain")
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
return web.Response(status=403, text="signature verification failed")
|
||||||
|
|
||||||
|
async def _handle_callback(self, request: web.Request) -> web.Response:
|
||||||
|
"""POST endpoint — receive an encrypted message callback."""
|
||||||
|
msg_signature = request.query.get("msg_signature", "")
|
||||||
|
timestamp = request.query.get("timestamp", "")
|
||||||
|
nonce = request.query.get("nonce", "")
|
||||||
|
body = await request.text()
|
||||||
|
|
||||||
|
for app in self._apps:
|
||||||
|
try:
|
||||||
|
decrypted = self._decrypt_request(
|
||||||
|
app, body, msg_signature, timestamp, nonce,
|
||||||
|
)
|
||||||
|
event = self._build_event(app, decrypted)
|
||||||
|
if event is not None:
|
||||||
|
# Record which app this user belongs to.
|
||||||
|
if event.source and event.source.user_id:
|
||||||
|
map_key = self._user_app_key(
|
||||||
|
str(app.get("corp_id") or ""), event.source.user_id,
|
||||||
|
)
|
||||||
|
self._user_app_map[map_key] = app["name"]
|
||||||
|
await self._message_queue.put(event)
|
||||||
|
# Immediately acknowledge — the agent's reply will arrive
|
||||||
|
# later via the proactive message/send API.
|
||||||
|
return web.Response(text="success", content_type="text/plain")
|
||||||
|
except WeComCryptoError:
|
||||||
|
continue
|
||||||
|
except Exception:
|
||||||
|
logger.exception("[WecomCallback] Error handling message")
|
||||||
|
break
|
||||||
|
return web.Response(status=400, text="invalid callback payload")
|
||||||
|
|
||||||
|
async def _poll_loop(self) -> None:
|
||||||
|
"""Drain the message queue and dispatch to the gateway runner."""
|
||||||
|
while True:
|
||||||
|
event = await self._message_queue.get()
|
||||||
|
try:
|
||||||
|
task = asyncio.create_task(self.handle_message(event))
|
||||||
|
self._background_tasks.add(task)
|
||||||
|
task.add_done_callback(self._background_tasks.discard)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("[WecomCallback] Failed to enqueue event")
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# XML / crypto helpers
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _decrypt_request(
|
||||||
|
self, app: Dict[str, Any], body: str,
|
||||||
|
msg_signature: str, timestamp: str, nonce: str,
|
||||||
|
) -> str:
|
||||||
|
root = ET.fromstring(body)
|
||||||
|
encrypt = root.findtext("Encrypt", default="")
|
||||||
|
crypt = self._crypt_for_app(app)
|
||||||
|
return crypt.decrypt(msg_signature, timestamp, nonce, encrypt).decode("utf-8")
|
||||||
|
|
||||||
|
def _build_event(self, app: Dict[str, Any], xml_text: str) -> Optional[MessageEvent]:
|
||||||
|
root = ET.fromstring(xml_text)
|
||||||
|
msg_type = (root.findtext("MsgType") or "").lower()
|
||||||
|
# Silently acknowledge lifecycle events.
|
||||||
|
if msg_type == "event":
|
||||||
|
event_name = (root.findtext("Event") or "").lower()
|
||||||
|
if event_name in {"enter_agent", "subscribe"}:
|
||||||
|
return None
|
||||||
|
if msg_type not in {"text", "event"}:
|
||||||
|
return None
|
||||||
|
|
||||||
|
user_id = root.findtext("FromUserName", default="")
|
||||||
|
corp_id = root.findtext("ToUserName", default=app.get("corp_id", ""))
|
||||||
|
scoped_chat_id = self._user_app_key(corp_id, user_id)
|
||||||
|
content = root.findtext("Content", default="").strip()
|
||||||
|
if not content and msg_type == "event":
|
||||||
|
content = "/start"
|
||||||
|
msg_id = (
|
||||||
|
root.findtext("MsgId")
|
||||||
|
or f"{user_id}:{root.findtext('CreateTime', default='0')}"
|
||||||
|
)
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id=scoped_chat_id,
|
||||||
|
chat_name=user_id,
|
||||||
|
chat_type="dm",
|
||||||
|
user_id=user_id,
|
||||||
|
user_name=user_id,
|
||||||
|
)
|
||||||
|
return MessageEvent(
|
||||||
|
text=content,
|
||||||
|
message_type=MessageType.TEXT,
|
||||||
|
source=source,
|
||||||
|
raw_message=xml_text,
|
||||||
|
message_id=msg_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _crypt_for_app(self, app: Dict[str, Any]) -> WXBizMsgCrypt:
|
||||||
|
return WXBizMsgCrypt(
|
||||||
|
token=str(app.get("token") or ""),
|
||||||
|
encoding_aes_key=str(app.get("encoding_aes_key") or ""),
|
||||||
|
receive_id=str(app.get("corp_id") or ""),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _get_app_by_name(self, name: Optional[str]) -> Optional[Dict[str, Any]]:
|
||||||
|
if not name:
|
||||||
|
return None
|
||||||
|
for app in self._apps:
|
||||||
|
if app.get("name") == name:
|
||||||
|
return app
|
||||||
|
return None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Access-token management
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def _get_access_token(self, app: Dict[str, Any]) -> str:
|
||||||
|
cached = self._access_tokens.get(app["name"])
|
||||||
|
now = time.time()
|
||||||
|
if cached and cached.get("expires_at", 0) > now + 60:
|
||||||
|
return cached["token"]
|
||||||
|
return await self._refresh_access_token(app)
|
||||||
|
|
||||||
|
async def _refresh_access_token(self, app: Dict[str, Any]) -> str:
|
||||||
|
resp = await self._http_client.get(
|
||||||
|
"https://qyapi.weixin.qq.com/cgi-bin/gettoken",
|
||||||
|
params={
|
||||||
|
"corpid": app.get("corp_id"),
|
||||||
|
"corpsecret": app.get("corp_secret"),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
data = resp.json()
|
||||||
|
if data.get("errcode") != 0:
|
||||||
|
raise RuntimeError(f"WeCom token refresh failed: {data}")
|
||||||
|
token = data["access_token"]
|
||||||
|
expires_in = int(data.get("expires_in", ACCESS_TOKEN_TTL_SECONDS))
|
||||||
|
self._access_tokens[app["name"]] = {
|
||||||
|
"token": token,
|
||||||
|
"expires_at": time.time() + expires_in,
|
||||||
|
}
|
||||||
|
logger.info(
|
||||||
|
"[WecomCallback] Token refreshed for app '%s' (corp=%s), expires in %ss",
|
||||||
|
app.get("name", "default"),
|
||||||
|
app.get("corp_id", ""),
|
||||||
|
expires_in,
|
||||||
|
)
|
||||||
|
return token
|
||||||
@@ -0,0 +1,142 @@
|
|||||||
|
"""WeCom BizMsgCrypt-compatible AES-CBC encryption for callback mode.
|
||||||
|
|
||||||
|
Implements the same wire format as Tencent's official ``WXBizMsgCrypt``
|
||||||
|
SDK so that WeCom can verify, encrypt, and decrypt callback payloads.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import base64
|
||||||
|
import hashlib
|
||||||
|
import os
|
||||||
|
import secrets
|
||||||
|
import socket
|
||||||
|
import struct
|
||||||
|
from typing import Optional
|
||||||
|
from xml.etree import ElementTree as ET
|
||||||
|
|
||||||
|
from cryptography.hazmat.backends import default_backend
|
||||||
|
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||||
|
|
||||||
|
|
||||||
|
class WeComCryptoError(Exception):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class SignatureError(WeComCryptoError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class DecryptError(WeComCryptoError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class EncryptError(WeComCryptoError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class PKCS7Encoder:
|
||||||
|
block_size = 32
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def encode(cls, text: bytes) -> bytes:
|
||||||
|
amount_to_pad = cls.block_size - (len(text) % cls.block_size)
|
||||||
|
if amount_to_pad == 0:
|
||||||
|
amount_to_pad = cls.block_size
|
||||||
|
pad = bytes([amount_to_pad]) * amount_to_pad
|
||||||
|
return text + pad
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def decode(cls, decrypted: bytes) -> bytes:
|
||||||
|
if not decrypted:
|
||||||
|
raise DecryptError("empty decrypted payload")
|
||||||
|
pad = decrypted[-1]
|
||||||
|
if pad < 1 or pad > cls.block_size:
|
||||||
|
raise DecryptError("invalid PKCS7 padding")
|
||||||
|
if decrypted[-pad:] != bytes([pad]) * pad:
|
||||||
|
raise DecryptError("malformed PKCS7 padding")
|
||||||
|
return decrypted[:-pad]
|
||||||
|
|
||||||
|
|
||||||
|
def _sha1_signature(token: str, timestamp: str, nonce: str, encrypt: str) -> str:
|
||||||
|
parts = sorted([token, timestamp, nonce, encrypt])
|
||||||
|
return hashlib.sha1("".join(parts).encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
class WXBizMsgCrypt:
|
||||||
|
"""Minimal WeCom callback crypto helper compatible with BizMsgCrypt semantics."""
|
||||||
|
|
||||||
|
def __init__(self, token: str, encoding_aes_key: str, receive_id: str):
|
||||||
|
if not token:
|
||||||
|
raise ValueError("token is required")
|
||||||
|
if not encoding_aes_key:
|
||||||
|
raise ValueError("encoding_aes_key is required")
|
||||||
|
if len(encoding_aes_key) != 43:
|
||||||
|
raise ValueError("encoding_aes_key must be 43 chars")
|
||||||
|
if not receive_id:
|
||||||
|
raise ValueError("receive_id is required")
|
||||||
|
|
||||||
|
self.token = token
|
||||||
|
self.receive_id = receive_id
|
||||||
|
self.key = base64.b64decode(encoding_aes_key + "=")
|
||||||
|
self.iv = self.key[:16]
|
||||||
|
|
||||||
|
def verify_url(self, msg_signature: str, timestamp: str, nonce: str, echostr: str) -> str:
|
||||||
|
plain = self.decrypt(msg_signature, timestamp, nonce, echostr)
|
||||||
|
return plain.decode("utf-8")
|
||||||
|
|
||||||
|
def decrypt(self, msg_signature: str, timestamp: str, nonce: str, encrypt: str) -> bytes:
|
||||||
|
expected = _sha1_signature(self.token, timestamp, nonce, encrypt)
|
||||||
|
if expected != msg_signature:
|
||||||
|
raise SignatureError("signature mismatch")
|
||||||
|
try:
|
||||||
|
cipher_text = base64.b64decode(encrypt)
|
||||||
|
except Exception as exc:
|
||||||
|
raise DecryptError(f"invalid base64 payload: {exc}") from exc
|
||||||
|
try:
|
||||||
|
cipher = Cipher(algorithms.AES(self.key), modes.CBC(self.iv), backend=default_backend())
|
||||||
|
decryptor = cipher.decryptor()
|
||||||
|
padded = decryptor.update(cipher_text) + decryptor.finalize()
|
||||||
|
plain = PKCS7Encoder.decode(padded)
|
||||||
|
content = plain[16:] # skip 16-byte random prefix
|
||||||
|
xml_length = socket.ntohl(struct.unpack("I", content[:4])[0])
|
||||||
|
xml_content = content[4:4 + xml_length]
|
||||||
|
receive_id = content[4 + xml_length:].decode("utf-8")
|
||||||
|
except WeComCryptoError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
raise DecryptError(f"decrypt failed: {exc}") from exc
|
||||||
|
|
||||||
|
if receive_id != self.receive_id:
|
||||||
|
raise DecryptError("receive_id mismatch")
|
||||||
|
return xml_content
|
||||||
|
|
||||||
|
def encrypt(self, plaintext: str, nonce: Optional[str] = None, timestamp: Optional[str] = None) -> str:
|
||||||
|
nonce = nonce or self._random_nonce()
|
||||||
|
timestamp = timestamp or str(int(__import__("time").time()))
|
||||||
|
encrypt = self._encrypt_bytes(plaintext.encode("utf-8"))
|
||||||
|
signature = _sha1_signature(self.token, timestamp, nonce, encrypt)
|
||||||
|
root = ET.Element("xml")
|
||||||
|
ET.SubElement(root, "Encrypt").text = encrypt
|
||||||
|
ET.SubElement(root, "MsgSignature").text = signature
|
||||||
|
ET.SubElement(root, "TimeStamp").text = timestamp
|
||||||
|
ET.SubElement(root, "Nonce").text = nonce
|
||||||
|
return ET.tostring(root, encoding="unicode")
|
||||||
|
|
||||||
|
def _encrypt_bytes(self, raw: bytes) -> str:
|
||||||
|
try:
|
||||||
|
random_prefix = os.urandom(16)
|
||||||
|
msg_len = struct.pack("I", socket.htonl(len(raw)))
|
||||||
|
payload = random_prefix + msg_len + raw + self.receive_id.encode("utf-8")
|
||||||
|
padded = PKCS7Encoder.encode(payload)
|
||||||
|
cipher = Cipher(algorithms.AES(self.key), modes.CBC(self.iv), backend=default_backend())
|
||||||
|
encryptor = cipher.encryptor()
|
||||||
|
encrypted = encryptor.update(padded) + encryptor.finalize()
|
||||||
|
return base64.b64encode(encrypted).decode("utf-8")
|
||||||
|
except Exception as exc:
|
||||||
|
raise EncryptError(f"encrypt failed: {exc}") from exc
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _random_nonce(length: int = 10) -> str:
|
||||||
|
alphabet = "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"
|
||||||
|
return "".join(secrets.choice(alphabet) for _ in range(length))
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,989 @@
|
|||||||
|
"""
|
||||||
|
WhatsApp platform adapter.
|
||||||
|
|
||||||
|
WhatsApp integration is more complex than Telegram/Discord because:
|
||||||
|
- No official bot API for personal accounts
|
||||||
|
- Business API requires Meta Business verification
|
||||||
|
- Most solutions use web-based automation
|
||||||
|
|
||||||
|
This adapter supports multiple backends:
|
||||||
|
1. WhatsApp Business API (requires Meta verification)
|
||||||
|
2. whatsapp-web.js (via Node.js subprocess) - for personal accounts
|
||||||
|
3. Baileys (via Node.js subprocess) - alternative for personal accounts
|
||||||
|
|
||||||
|
For simplicity, we'll implement a generic interface that can work
|
||||||
|
with different backends via a bridge pattern.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import platform
|
||||||
|
import re
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
_IS_WINDOWS = platform.system() == "Windows"
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, Optional, Any
|
||||||
|
|
||||||
|
from hermes_constants import get_hermes_dir
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _kill_port_process(port: int) -> None:
|
||||||
|
"""Kill any process listening on the given TCP port."""
|
||||||
|
try:
|
||||||
|
if _IS_WINDOWS:
|
||||||
|
# Use netstat to find the PID bound to this port, then taskkill
|
||||||
|
result = subprocess.run(
|
||||||
|
["netstat", "-ano", "-p", "TCP"],
|
||||||
|
capture_output=True, text=True, timeout=5,
|
||||||
|
)
|
||||||
|
for line in result.stdout.splitlines():
|
||||||
|
parts = line.split()
|
||||||
|
if len(parts) >= 5 and parts[3] == "LISTENING":
|
||||||
|
local_addr = parts[1]
|
||||||
|
if local_addr.endswith(f":{port}"):
|
||||||
|
try:
|
||||||
|
subprocess.run(
|
||||||
|
["taskkill", "/PID", parts[4], "/F"],
|
||||||
|
capture_output=True, timeout=5,
|
||||||
|
)
|
||||||
|
except subprocess.SubprocessError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
result = subprocess.run(
|
||||||
|
["fuser", f"{port}/tcp"],
|
||||||
|
capture_output=True, timeout=5,
|
||||||
|
)
|
||||||
|
if result.returncode == 0:
|
||||||
|
subprocess.run(
|
||||||
|
["fuser", "-k", f"{port}/tcp"],
|
||||||
|
capture_output=True, timeout=5,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
import sys
|
||||||
|
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
|
||||||
|
|
||||||
|
from gateway.config import Platform, PlatformConfig
|
||||||
|
from gateway.platforms.base import (
|
||||||
|
BasePlatformAdapter,
|
||||||
|
MessageEvent,
|
||||||
|
MessageType,
|
||||||
|
SendResult,
|
||||||
|
SUPPORTED_DOCUMENT_TYPES,
|
||||||
|
cache_image_from_url,
|
||||||
|
cache_audio_from_url,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def check_whatsapp_requirements() -> bool:
|
||||||
|
"""
|
||||||
|
Check if WhatsApp dependencies are available.
|
||||||
|
|
||||||
|
WhatsApp requires a Node.js bridge for most implementations.
|
||||||
|
"""
|
||||||
|
# Check for Node.js
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
["node", "--version"],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=5
|
||||||
|
)
|
||||||
|
return result.returncode == 0
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class WhatsAppAdapter(BasePlatformAdapter):
|
||||||
|
"""
|
||||||
|
WhatsApp adapter.
|
||||||
|
|
||||||
|
This implementation uses a simple HTTP bridge pattern where:
|
||||||
|
1. A Node.js process runs the WhatsApp Web client
|
||||||
|
2. Messages are forwarded via HTTP/IPC to this Python adapter
|
||||||
|
3. Responses are sent back through the bridge
|
||||||
|
|
||||||
|
The actual Node.js bridge implementation can vary:
|
||||||
|
- whatsapp-web.js based
|
||||||
|
- Baileys based
|
||||||
|
- Business API based
|
||||||
|
|
||||||
|
Configuration:
|
||||||
|
- bridge_script: Path to the Node.js bridge script
|
||||||
|
- bridge_port: Port for HTTP communication (default: 3000)
|
||||||
|
- session_path: Path to store WhatsApp session data
|
||||||
|
"""
|
||||||
|
|
||||||
|
# WhatsApp message limits — practical UX limit, not protocol max.
|
||||||
|
# WhatsApp allows ~65K but long messages are unreadable on mobile.
|
||||||
|
MAX_MESSAGE_LENGTH = 4096
|
||||||
|
|
||||||
|
# Default bridge location relative to the hermes-agent install
|
||||||
|
_DEFAULT_BRIDGE_DIR = Path(__file__).resolve().parents[2] / "scripts" / "whatsapp-bridge"
|
||||||
|
|
||||||
|
def __init__(self, config: PlatformConfig):
|
||||||
|
super().__init__(config, Platform.WHATSAPP)
|
||||||
|
self._bridge_process: Optional[subprocess.Popen] = None
|
||||||
|
self._bridge_port: int = config.extra.get("bridge_port", 3000)
|
||||||
|
self._bridge_script: Optional[str] = config.extra.get(
|
||||||
|
"bridge_script",
|
||||||
|
str(self._DEFAULT_BRIDGE_DIR / "bridge.js"),
|
||||||
|
)
|
||||||
|
self._session_path: Path = Path(config.extra.get(
|
||||||
|
"session_path",
|
||||||
|
get_hermes_dir("platforms/whatsapp/session", "whatsapp/session")
|
||||||
|
))
|
||||||
|
self._reply_prefix: Optional[str] = config.extra.get("reply_prefix")
|
||||||
|
self._mention_patterns = self._compile_mention_patterns()
|
||||||
|
self._message_queue: asyncio.Queue = asyncio.Queue()
|
||||||
|
self._bridge_log_fh = None
|
||||||
|
self._bridge_log: Optional[Path] = None
|
||||||
|
self._poll_task: Optional[asyncio.Task] = None
|
||||||
|
self._http_session: Optional["aiohttp.ClientSession"] = None
|
||||||
|
|
||||||
|
def _whatsapp_require_mention(self) -> bool:
|
||||||
|
configured = self.config.extra.get("require_mention")
|
||||||
|
if configured is not None:
|
||||||
|
if isinstance(configured, str):
|
||||||
|
return configured.lower() in ("true", "1", "yes", "on")
|
||||||
|
return bool(configured)
|
||||||
|
return os.getenv("WHATSAPP_REQUIRE_MENTION", "false").lower() in ("true", "1", "yes", "on")
|
||||||
|
|
||||||
|
def _whatsapp_free_response_chats(self) -> set[str]:
|
||||||
|
raw = self.config.extra.get("free_response_chats")
|
||||||
|
if raw is None:
|
||||||
|
raw = os.getenv("WHATSAPP_FREE_RESPONSE_CHATS", "")
|
||||||
|
if isinstance(raw, list):
|
||||||
|
return {str(part).strip() for part in raw if str(part).strip()}
|
||||||
|
return {part.strip() for part in str(raw).split(",") if part.strip()}
|
||||||
|
|
||||||
|
def _compile_mention_patterns(self):
|
||||||
|
patterns = self.config.extra.get("mention_patterns")
|
||||||
|
if patterns is None:
|
||||||
|
raw = os.getenv("WHATSAPP_MENTION_PATTERNS", "").strip()
|
||||||
|
if raw:
|
||||||
|
try:
|
||||||
|
patterns = json.loads(raw)
|
||||||
|
except Exception:
|
||||||
|
patterns = [part.strip() for part in raw.splitlines() if part.strip()]
|
||||||
|
if not patterns:
|
||||||
|
patterns = [part.strip() for part in raw.split(",") if part.strip()]
|
||||||
|
if patterns is None:
|
||||||
|
return []
|
||||||
|
if isinstance(patterns, str):
|
||||||
|
patterns = [patterns]
|
||||||
|
if not isinstance(patterns, list):
|
||||||
|
logger.warning("[%s] whatsapp mention_patterns must be a list or string; got %s", self.name, type(patterns).__name__)
|
||||||
|
return []
|
||||||
|
|
||||||
|
compiled = []
|
||||||
|
for pattern in patterns:
|
||||||
|
if not isinstance(pattern, str) or not pattern.strip():
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
compiled.append(re.compile(pattern, re.IGNORECASE))
|
||||||
|
except re.error as exc:
|
||||||
|
logger.warning("[%s] Invalid WhatsApp mention pattern %r: %s", self.name, pattern, exc)
|
||||||
|
if compiled:
|
||||||
|
logger.info("[%s] Loaded %d WhatsApp mention pattern(s)", self.name, len(compiled))
|
||||||
|
return compiled
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_whatsapp_id(value: Optional[str]) -> str:
|
||||||
|
if not value:
|
||||||
|
return ""
|
||||||
|
normalized = str(value).strip()
|
||||||
|
if ":" in normalized and "@" in normalized:
|
||||||
|
normalized = normalized.replace(":", "@", 1)
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
def _bot_ids_from_message(self, data: Dict[str, Any]) -> set[str]:
|
||||||
|
bot_ids = set()
|
||||||
|
for candidate in data.get("botIds") or []:
|
||||||
|
normalized = self._normalize_whatsapp_id(candidate)
|
||||||
|
if normalized:
|
||||||
|
bot_ids.add(normalized)
|
||||||
|
return bot_ids
|
||||||
|
|
||||||
|
def _message_is_reply_to_bot(self, data: Dict[str, Any]) -> bool:
|
||||||
|
quoted_participant = self._normalize_whatsapp_id(data.get("quotedParticipant"))
|
||||||
|
if not quoted_participant:
|
||||||
|
return False
|
||||||
|
return quoted_participant in self._bot_ids_from_message(data)
|
||||||
|
|
||||||
|
def _message_mentions_bot(self, data: Dict[str, Any]) -> bool:
|
||||||
|
bot_ids = self._bot_ids_from_message(data)
|
||||||
|
if not bot_ids:
|
||||||
|
return False
|
||||||
|
mentioned_ids = {
|
||||||
|
nid
|
||||||
|
for candidate in (data.get("mentionedIds") or [])
|
||||||
|
if (nid := self._normalize_whatsapp_id(candidate))
|
||||||
|
}
|
||||||
|
if mentioned_ids & bot_ids:
|
||||||
|
return True
|
||||||
|
|
||||||
|
body = str(data.get("body") or "")
|
||||||
|
lower_body = body.lower()
|
||||||
|
for bot_id in bot_ids:
|
||||||
|
bare_id = bot_id.split("@", 1)[0].lower()
|
||||||
|
if bare_id and (f"@{bare_id}" in lower_body or bare_id in lower_body):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _message_matches_mention_patterns(self, data: Dict[str, Any]) -> bool:
|
||||||
|
if not self._mention_patterns:
|
||||||
|
return False
|
||||||
|
body = str(data.get("body") or "")
|
||||||
|
return any(pattern.search(body) for pattern in self._mention_patterns)
|
||||||
|
|
||||||
|
def _clean_bot_mention_text(self, text: str, data: Dict[str, Any]) -> str:
|
||||||
|
if not text:
|
||||||
|
return text
|
||||||
|
bot_ids = self._bot_ids_from_message(data)
|
||||||
|
cleaned = text
|
||||||
|
for bot_id in bot_ids:
|
||||||
|
bare_id = bot_id.split("@", 1)[0]
|
||||||
|
if bare_id:
|
||||||
|
cleaned = re.sub(rf"@{re.escape(bare_id)}\b[,:\-]*\s*", "", cleaned)
|
||||||
|
return cleaned.strip() or text
|
||||||
|
|
||||||
|
def _should_process_message(self, data: Dict[str, Any]) -> bool:
|
||||||
|
if not data.get("isGroup"):
|
||||||
|
return True
|
||||||
|
chat_id = str(data.get("chatId") or "")
|
||||||
|
if chat_id in self._whatsapp_free_response_chats():
|
||||||
|
return True
|
||||||
|
if not self._whatsapp_require_mention():
|
||||||
|
return True
|
||||||
|
body = str(data.get("body") or "").strip()
|
||||||
|
if body.startswith("/"):
|
||||||
|
return True
|
||||||
|
if self._message_is_reply_to_bot(data):
|
||||||
|
return True
|
||||||
|
if self._message_mentions_bot(data):
|
||||||
|
return True
|
||||||
|
return self._message_matches_mention_patterns(data)
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
"""
|
||||||
|
Start the WhatsApp bridge.
|
||||||
|
|
||||||
|
This launches the Node.js bridge process and waits for it to be ready.
|
||||||
|
"""
|
||||||
|
if not check_whatsapp_requirements():
|
||||||
|
logger.warning("[%s] Node.js not found. WhatsApp requires Node.js.", self.name)
|
||||||
|
return False
|
||||||
|
|
||||||
|
bridge_path = Path(self._bridge_script)
|
||||||
|
if not bridge_path.exists():
|
||||||
|
logger.warning("[%s] Bridge script not found: %s", self.name, bridge_path)
|
||||||
|
return False
|
||||||
|
|
||||||
|
logger.info("[%s] Bridge found at %s", self.name, bridge_path)
|
||||||
|
|
||||||
|
# Acquire scoped lock to prevent duplicate sessions
|
||||||
|
try:
|
||||||
|
if not self._acquire_platform_lock('whatsapp-session', str(self._session_path), 'WhatsApp session'):
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("[%s] Could not acquire session lock (non-fatal): %s", self.name, e)
|
||||||
|
|
||||||
|
# Auto-install npm dependencies if node_modules doesn't exist
|
||||||
|
bridge_dir = bridge_path.parent
|
||||||
|
if not (bridge_dir / "node_modules").exists():
|
||||||
|
print(f"[{self.name}] Installing WhatsApp bridge dependencies...")
|
||||||
|
try:
|
||||||
|
install_result = subprocess.run(
|
||||||
|
["npm", "install", "--silent"],
|
||||||
|
cwd=str(bridge_dir),
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=60,
|
||||||
|
)
|
||||||
|
if install_result.returncode != 0:
|
||||||
|
print(f"[{self.name}] npm install failed: {install_result.stderr}")
|
||||||
|
return False
|
||||||
|
print(f"[{self.name}] Dependencies installed")
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[{self.name}] Failed to install dependencies: {e}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Ensure session directory exists
|
||||||
|
self._session_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
# Check if bridge is already running and connected
|
||||||
|
import aiohttp
|
||||||
|
import asyncio
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
async with session.get(
|
||||||
|
f"http://127.0.0.1:{self._bridge_port}/health",
|
||||||
|
timeout=aiohttp.ClientTimeout(total=2)
|
||||||
|
) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
data = await resp.json()
|
||||||
|
bridge_status = data.get("status", "unknown")
|
||||||
|
if bridge_status == "connected":
|
||||||
|
print(f"[{self.name}] Using existing bridge (status: {bridge_status})")
|
||||||
|
self._mark_connected()
|
||||||
|
self._bridge_process = None # Not managed by us
|
||||||
|
self._http_session = aiohttp.ClientSession()
|
||||||
|
self._poll_task = asyncio.create_task(self._poll_messages())
|
||||||
|
return True
|
||||||
|
else:
|
||||||
|
print(f"[{self.name}] Bridge found but not connected (status: {bridge_status}), restarting")
|
||||||
|
except Exception:
|
||||||
|
pass # Bridge not running, start a new one
|
||||||
|
|
||||||
|
# Kill any orphaned bridge from a previous gateway run
|
||||||
|
_kill_port_process(self._bridge_port)
|
||||||
|
await asyncio.sleep(1)
|
||||||
|
|
||||||
|
# Start the bridge process in its own process group.
|
||||||
|
# Route output to a log file so QR codes, errors, and reconnection
|
||||||
|
# messages are preserved for troubleshooting.
|
||||||
|
whatsapp_mode = os.getenv("WHATSAPP_MODE", "self-chat")
|
||||||
|
self._bridge_log = self._session_path.parent / "bridge.log"
|
||||||
|
bridge_log_fh = open(self._bridge_log, "a")
|
||||||
|
self._bridge_log_fh = bridge_log_fh
|
||||||
|
|
||||||
|
# Build bridge subprocess environment.
|
||||||
|
# Pass WHATSAPP_REPLY_PREFIX from config.yaml so the Node bridge
|
||||||
|
# can use it without the user needing to set a separate env var.
|
||||||
|
bridge_env = os.environ.copy()
|
||||||
|
if self._reply_prefix is not None:
|
||||||
|
bridge_env["WHATSAPP_REPLY_PREFIX"] = self._reply_prefix
|
||||||
|
|
||||||
|
self._bridge_process = subprocess.Popen(
|
||||||
|
[
|
||||||
|
"node",
|
||||||
|
str(bridge_path),
|
||||||
|
"--port", str(self._bridge_port),
|
||||||
|
"--session", str(self._session_path),
|
||||||
|
"--mode", whatsapp_mode,
|
||||||
|
],
|
||||||
|
stdout=bridge_log_fh,
|
||||||
|
stderr=bridge_log_fh,
|
||||||
|
preexec_fn=None if _IS_WINDOWS else os.setsid,
|
||||||
|
env=bridge_env,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Wait for the bridge to connect to WhatsApp.
|
||||||
|
# Phase 1: wait for the HTTP server to come up (up to 15s).
|
||||||
|
# Phase 2: wait for WhatsApp status: connected (up to 15s more).
|
||||||
|
import aiohttp
|
||||||
|
http_ready = False
|
||||||
|
data = {}
|
||||||
|
for attempt in range(15):
|
||||||
|
await asyncio.sleep(1)
|
||||||
|
if self._bridge_process.poll() is not None:
|
||||||
|
print(f"[{self.name}] Bridge process died (exit code {self._bridge_process.returncode})")
|
||||||
|
print(f"[{self.name}] Check log: {self._bridge_log}")
|
||||||
|
self._close_bridge_log()
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
async with session.get(
|
||||||
|
f"http://127.0.0.1:{self._bridge_port}/health",
|
||||||
|
timeout=aiohttp.ClientTimeout(total=2)
|
||||||
|
) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
http_ready = True
|
||||||
|
data = await resp.json()
|
||||||
|
if data.get("status") == "connected":
|
||||||
|
print(f"[{self.name}] Bridge ready (status: connected)")
|
||||||
|
break
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if not http_ready:
|
||||||
|
print(f"[{self.name}] Bridge HTTP server did not start in 15s")
|
||||||
|
print(f"[{self.name}] Check log: {self._bridge_log}")
|
||||||
|
self._close_bridge_log()
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Phase 2: HTTP is up but WhatsApp may still be connecting.
|
||||||
|
# Give it more time to authenticate with saved credentials.
|
||||||
|
if data.get("status") != "connected":
|
||||||
|
print(f"[{self.name}] Bridge HTTP ready, waiting for WhatsApp connection...")
|
||||||
|
for attempt in range(15):
|
||||||
|
await asyncio.sleep(1)
|
||||||
|
if self._bridge_process.poll() is not None:
|
||||||
|
print(f"[{self.name}] Bridge process died during connection")
|
||||||
|
print(f"[{self.name}] Check log: {self._bridge_log}")
|
||||||
|
self._close_bridge_log()
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
async with session.get(
|
||||||
|
f"http://127.0.0.1:{self._bridge_port}/health",
|
||||||
|
timeout=aiohttp.ClientTimeout(total=2)
|
||||||
|
) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
data = await resp.json()
|
||||||
|
if data.get("status") == "connected":
|
||||||
|
print(f"[{self.name}] Bridge ready (status: connected)")
|
||||||
|
break
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
# Still not connected — warn but proceed (bridge may
|
||||||
|
# auto-reconnect later, e.g. after a code 515 restart).
|
||||||
|
print(f"[{self.name}] ⚠ WhatsApp not connected after 30s")
|
||||||
|
print(f"[{self.name}] Bridge log: {self._bridge_log}")
|
||||||
|
print(f"[{self.name}] If session expired, re-pair: hermes whatsapp")
|
||||||
|
|
||||||
|
# Create a persistent HTTP session for all bridge communication
|
||||||
|
self._http_session = aiohttp.ClientSession()
|
||||||
|
|
||||||
|
# Start message polling task
|
||||||
|
self._poll_task = asyncio.create_task(self._poll_messages())
|
||||||
|
|
||||||
|
self._mark_connected()
|
||||||
|
print(f"[{self.name}] Bridge started on port {self._bridge_port}")
|
||||||
|
return True
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
self._release_platform_lock()
|
||||||
|
logger.error("[%s] Failed to start bridge: %s", self.name, e, exc_info=True)
|
||||||
|
self._close_bridge_log()
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _close_bridge_log(self) -> None:
|
||||||
|
"""Close the bridge log file handle if open."""
|
||||||
|
if self._bridge_log_fh:
|
||||||
|
try:
|
||||||
|
self._bridge_log_fh.close()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
self._bridge_log_fh = None
|
||||||
|
|
||||||
|
async def _check_managed_bridge_exit(self) -> Optional[str]:
|
||||||
|
"""Return a fatal error message if the managed bridge child exited."""
|
||||||
|
if self._bridge_process is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
returncode = self._bridge_process.poll()
|
||||||
|
if returncode is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
message = f"WhatsApp bridge process exited unexpectedly (code {returncode})."
|
||||||
|
if not self.has_fatal_error:
|
||||||
|
logger.error("[%s] %s", self.name, message)
|
||||||
|
self._set_fatal_error("whatsapp_bridge_exited", message, retryable=True)
|
||||||
|
self._close_bridge_log()
|
||||||
|
await self._notify_fatal_error()
|
||||||
|
return self.fatal_error_message or message
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
"""Stop the WhatsApp bridge and clean up any orphaned processes."""
|
||||||
|
if self._bridge_process:
|
||||||
|
try:
|
||||||
|
# Kill the entire process group so child node processes die too
|
||||||
|
import signal
|
||||||
|
try:
|
||||||
|
if _IS_WINDOWS:
|
||||||
|
self._bridge_process.terminate()
|
||||||
|
else:
|
||||||
|
os.killpg(os.getpgid(self._bridge_process.pid), signal.SIGTERM)
|
||||||
|
except (ProcessLookupError, PermissionError):
|
||||||
|
self._bridge_process.terminate()
|
||||||
|
await asyncio.sleep(1)
|
||||||
|
if self._bridge_process.poll() is None:
|
||||||
|
try:
|
||||||
|
if _IS_WINDOWS:
|
||||||
|
self._bridge_process.kill()
|
||||||
|
else:
|
||||||
|
os.killpg(os.getpgid(self._bridge_process.pid), signal.SIGKILL)
|
||||||
|
except (ProcessLookupError, PermissionError):
|
||||||
|
self._bridge_process.kill()
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[{self.name}] Error stopping bridge: {e}")
|
||||||
|
else:
|
||||||
|
# Bridge was not started by us, don't kill it
|
||||||
|
print(f"[{self.name}] Disconnecting (external bridge left running)")
|
||||||
|
|
||||||
|
# Cancel the poll task explicitly
|
||||||
|
if self._poll_task and not self._poll_task.done():
|
||||||
|
self._poll_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._poll_task
|
||||||
|
except (asyncio.CancelledError, Exception):
|
||||||
|
pass
|
||||||
|
self._poll_task = None
|
||||||
|
|
||||||
|
# Close the persistent HTTP session
|
||||||
|
if self._http_session and not self._http_session.closed:
|
||||||
|
await self._http_session.close()
|
||||||
|
self._http_session = None
|
||||||
|
|
||||||
|
self._release_platform_lock()
|
||||||
|
|
||||||
|
self._mark_disconnected()
|
||||||
|
self._bridge_process = None
|
||||||
|
self._close_bridge_log()
|
||||||
|
print(f"[{self.name}] Disconnected")
|
||||||
|
|
||||||
|
def format_message(self, content: str) -> str:
|
||||||
|
"""Convert standard markdown to WhatsApp-compatible formatting.
|
||||||
|
|
||||||
|
WhatsApp supports: *bold*, _italic_, ~strikethrough~, ```code```,
|
||||||
|
and monospaced `inline`. Standard markdown uses different syntax
|
||||||
|
for bold/italic/strikethrough, so we convert here.
|
||||||
|
|
||||||
|
Code blocks (``` fenced) and inline code (`) are protected from
|
||||||
|
conversion via placeholder substitution.
|
||||||
|
"""
|
||||||
|
if not content:
|
||||||
|
return content
|
||||||
|
|
||||||
|
# --- 1. Protect fenced code blocks from formatting changes ---
|
||||||
|
_FENCE_PH = "\x00FENCE"
|
||||||
|
fences: list[str] = []
|
||||||
|
|
||||||
|
def _save_fence(m: re.Match) -> str:
|
||||||
|
fences.append(m.group(0))
|
||||||
|
return f"{_FENCE_PH}{len(fences) - 1}\x00"
|
||||||
|
|
||||||
|
result = re.sub(r"```[\s\S]*?```", _save_fence, content)
|
||||||
|
|
||||||
|
# --- 2. Protect inline code ---
|
||||||
|
_CODE_PH = "\x00CODE"
|
||||||
|
codes: list[str] = []
|
||||||
|
|
||||||
|
def _save_code(m: re.Match) -> str:
|
||||||
|
codes.append(m.group(0))
|
||||||
|
return f"{_CODE_PH}{len(codes) - 1}\x00"
|
||||||
|
|
||||||
|
result = re.sub(r"`[^`\n]+`", _save_code, result)
|
||||||
|
|
||||||
|
# --- 3. Convert markdown formatting to WhatsApp syntax ---
|
||||||
|
# Bold: **text** or __text__ → *text*
|
||||||
|
result = re.sub(r"\*\*(.+?)\*\*", r"*\1*", result)
|
||||||
|
result = re.sub(r"__(.+?)__", r"*\1*", result)
|
||||||
|
# Strikethrough: ~~text~~ → ~text~
|
||||||
|
result = re.sub(r"~~(.+?)~~", r"~\1~", result)
|
||||||
|
# Italic: *text* is already WhatsApp italic — leave as-is
|
||||||
|
# _text_ is already WhatsApp italic — leave as-is
|
||||||
|
|
||||||
|
# --- 4. Convert markdown headers to bold text ---
|
||||||
|
# # Header → *Header*
|
||||||
|
result = re.sub(r"^#{1,6}\s+(.+)$", r"*\1*", result, flags=re.MULTILINE)
|
||||||
|
|
||||||
|
# --- 5. Convert markdown links: [text](url) → text (url) ---
|
||||||
|
result = re.sub(r"\[([^\]]+)\]\(([^)]+)\)", r"\1 (\2)", result)
|
||||||
|
|
||||||
|
# --- 6. Restore protected sections ---
|
||||||
|
for i, fence in enumerate(fences):
|
||||||
|
result = result.replace(f"{_FENCE_PH}{i}\x00", fence)
|
||||||
|
for i, code in enumerate(codes):
|
||||||
|
result = result.replace(f"{_CODE_PH}{i}\x00", code)
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
async def send(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
content: str,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
metadata: Optional[Dict[str, Any]] = None
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a message via the WhatsApp bridge.
|
||||||
|
|
||||||
|
Formats markdown for WhatsApp, splits long messages into chunks
|
||||||
|
that preserve code block boundaries, and sends each chunk sequentially.
|
||||||
|
"""
|
||||||
|
if not self._running or not self._http_session:
|
||||||
|
return SendResult(success=False, error="Not connected")
|
||||||
|
bridge_exit = await self._check_managed_bridge_exit()
|
||||||
|
if bridge_exit:
|
||||||
|
return SendResult(success=False, error=bridge_exit)
|
||||||
|
|
||||||
|
if not content or not content.strip():
|
||||||
|
return SendResult(success=True, message_id=None)
|
||||||
|
|
||||||
|
try:
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
# Format and chunk the message
|
||||||
|
formatted = self.format_message(content)
|
||||||
|
chunks = self.truncate_message(formatted, self.MAX_MESSAGE_LENGTH)
|
||||||
|
|
||||||
|
last_message_id = None
|
||||||
|
for chunk in chunks:
|
||||||
|
payload: Dict[str, Any] = {
|
||||||
|
"chatId": chat_id,
|
||||||
|
"message": chunk,
|
||||||
|
}
|
||||||
|
if reply_to and last_message_id is None:
|
||||||
|
# Only reply-to on the first chunk
|
||||||
|
payload["replyTo"] = reply_to
|
||||||
|
|
||||||
|
async with self._http_session.post(
|
||||||
|
f"http://127.0.0.1:{self._bridge_port}/send",
|
||||||
|
json=payload,
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30)
|
||||||
|
) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
data = await resp.json()
|
||||||
|
last_message_id = data.get("messageId")
|
||||||
|
else:
|
||||||
|
error = await resp.text()
|
||||||
|
return SendResult(success=False, error=error)
|
||||||
|
|
||||||
|
# Small delay between chunks to avoid rate limiting
|
||||||
|
if len(chunks) > 1:
|
||||||
|
await asyncio.sleep(0.3)
|
||||||
|
|
||||||
|
return SendResult(
|
||||||
|
success=True,
|
||||||
|
message_id=last_message_id,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
async def edit_message(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
message_id: str,
|
||||||
|
content: str,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Edit a previously sent message via the WhatsApp bridge."""
|
||||||
|
if not self._running or not self._http_session:
|
||||||
|
return SendResult(success=False, error="Not connected")
|
||||||
|
bridge_exit = await self._check_managed_bridge_exit()
|
||||||
|
if bridge_exit:
|
||||||
|
return SendResult(success=False, error=bridge_exit)
|
||||||
|
try:
|
||||||
|
import aiohttp
|
||||||
|
async with self._http_session.post(
|
||||||
|
f"http://127.0.0.1:{self._bridge_port}/edit",
|
||||||
|
json={
|
||||||
|
"chatId": chat_id,
|
||||||
|
"messageId": message_id,
|
||||||
|
"message": content,
|
||||||
|
},
|
||||||
|
timeout=aiohttp.ClientTimeout(total=15)
|
||||||
|
) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
return SendResult(success=True, message_id=message_id)
|
||||||
|
else:
|
||||||
|
error = await resp.text()
|
||||||
|
return SendResult(success=False, error=error)
|
||||||
|
except Exception as e:
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
async def _send_media_to_bridge(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
file_path: str,
|
||||||
|
media_type: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
file_name: Optional[str] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send any media file via bridge /send-media endpoint."""
|
||||||
|
if not self._running or not self._http_session:
|
||||||
|
return SendResult(success=False, error="Not connected")
|
||||||
|
bridge_exit = await self._check_managed_bridge_exit()
|
||||||
|
if bridge_exit:
|
||||||
|
return SendResult(success=False, error=bridge_exit)
|
||||||
|
try:
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
if not os.path.exists(file_path):
|
||||||
|
return SendResult(success=False, error=f"File not found: {file_path}")
|
||||||
|
|
||||||
|
payload: Dict[str, Any] = {
|
||||||
|
"chatId": chat_id,
|
||||||
|
"filePath": file_path,
|
||||||
|
"mediaType": media_type,
|
||||||
|
}
|
||||||
|
if caption:
|
||||||
|
payload["caption"] = caption
|
||||||
|
if file_name:
|
||||||
|
payload["fileName"] = file_name
|
||||||
|
|
||||||
|
async with self._http_session.post(
|
||||||
|
f"http://127.0.0.1:{self._bridge_port}/send-media",
|
||||||
|
json=payload,
|
||||||
|
timeout=aiohttp.ClientTimeout(total=120),
|
||||||
|
) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
data = await resp.json()
|
||||||
|
return SendResult(
|
||||||
|
success=True,
|
||||||
|
message_id=data.get("messageId"),
|
||||||
|
raw_response=data,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
error = await resp.text()
|
||||||
|
return SendResult(success=False, error=error)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
return SendResult(success=False, error=str(e))
|
||||||
|
|
||||||
|
async def send_image(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
image_url: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Download image URL to cache, send natively via bridge."""
|
||||||
|
try:
|
||||||
|
local_path = await cache_image_from_url(image_url)
|
||||||
|
return await self._send_media_to_bridge(chat_id, local_path, "image", caption)
|
||||||
|
except Exception:
|
||||||
|
return await super().send_image(chat_id, image_url, caption, reply_to)
|
||||||
|
|
||||||
|
async def send_image_file(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
image_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a local image file natively via bridge."""
|
||||||
|
return await self._send_media_to_bridge(chat_id, image_path, "image", caption)
|
||||||
|
|
||||||
|
async def send_video(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
video_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a video natively via bridge — plays inline in WhatsApp."""
|
||||||
|
return await self._send_media_to_bridge(chat_id, video_path, "video", caption)
|
||||||
|
|
||||||
|
async def send_document(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
file_path: str,
|
||||||
|
caption: Optional[str] = None,
|
||||||
|
file_name: Optional[str] = None,
|
||||||
|
reply_to: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> SendResult:
|
||||||
|
"""Send a document/file as a downloadable attachment via bridge."""
|
||||||
|
return await self._send_media_to_bridge(
|
||||||
|
chat_id, file_path, "document", caption,
|
||||||
|
file_name or os.path.basename(file_path),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_typing(self, chat_id: str, metadata=None) -> None:
|
||||||
|
"""Send typing indicator via bridge."""
|
||||||
|
if not self._running or not self._http_session:
|
||||||
|
return
|
||||||
|
if await self._check_managed_bridge_exit():
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
await self._http_session.post(
|
||||||
|
f"http://127.0.0.1:{self._bridge_port}/typing",
|
||||||
|
json={"chatId": chat_id},
|
||||||
|
timeout=aiohttp.ClientTimeout(total=5)
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass # Ignore typing indicator failures
|
||||||
|
|
||||||
|
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||||
|
"""Get information about a WhatsApp chat."""
|
||||||
|
if not self._running or not self._http_session:
|
||||||
|
return {"name": "Unknown", "type": "dm"}
|
||||||
|
if await self._check_managed_bridge_exit():
|
||||||
|
return {"name": chat_id, "type": "dm"}
|
||||||
|
|
||||||
|
try:
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
async with self._http_session.get(
|
||||||
|
f"http://127.0.0.1:{self._bridge_port}/chat/{chat_id}",
|
||||||
|
timeout=aiohttp.ClientTimeout(total=10)
|
||||||
|
) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
data = await resp.json()
|
||||||
|
return {
|
||||||
|
"name": data.get("name", chat_id),
|
||||||
|
"type": "group" if data.get("isGroup") else "dm",
|
||||||
|
"participants": data.get("participants", []),
|
||||||
|
}
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Could not get WhatsApp chat info for %s: %s", chat_id, e)
|
||||||
|
|
||||||
|
return {"name": chat_id, "type": "dm"}
|
||||||
|
|
||||||
|
async def _poll_messages(self) -> None:
|
||||||
|
"""Poll the bridge for incoming messages."""
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
while self._running:
|
||||||
|
if not self._http_session:
|
||||||
|
break
|
||||||
|
bridge_exit = await self._check_managed_bridge_exit()
|
||||||
|
if bridge_exit:
|
||||||
|
print(f"[{self.name}] {bridge_exit}")
|
||||||
|
break
|
||||||
|
try:
|
||||||
|
async with self._http_session.get(
|
||||||
|
f"http://127.0.0.1:{self._bridge_port}/messages",
|
||||||
|
timeout=aiohttp.ClientTimeout(total=30)
|
||||||
|
) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
messages = await resp.json()
|
||||||
|
for msg_data in messages:
|
||||||
|
event = await self._build_message_event(msg_data)
|
||||||
|
if event:
|
||||||
|
await self.handle_message(event)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
break
|
||||||
|
except Exception as e:
|
||||||
|
bridge_exit = await self._check_managed_bridge_exit()
|
||||||
|
if bridge_exit:
|
||||||
|
print(f"[{self.name}] {bridge_exit}")
|
||||||
|
break
|
||||||
|
print(f"[{self.name}] Poll error: {e}")
|
||||||
|
await asyncio.sleep(5)
|
||||||
|
|
||||||
|
await asyncio.sleep(1) # Poll interval
|
||||||
|
|
||||||
|
async def _build_message_event(self, data: Dict[str, Any]) -> Optional[MessageEvent]:
|
||||||
|
"""Build a MessageEvent from bridge message data, downloading images to cache."""
|
||||||
|
try:
|
||||||
|
if not self._should_process_message(data):
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Determine message type
|
||||||
|
msg_type = MessageType.TEXT
|
||||||
|
if data.get("hasMedia"):
|
||||||
|
media_type = data.get("mediaType", "")
|
||||||
|
if "image" in media_type:
|
||||||
|
msg_type = MessageType.PHOTO
|
||||||
|
elif "video" in media_type:
|
||||||
|
msg_type = MessageType.VIDEO
|
||||||
|
elif "audio" in media_type or "ptt" in media_type: # ptt = voice note
|
||||||
|
msg_type = MessageType.VOICE
|
||||||
|
else:
|
||||||
|
msg_type = MessageType.DOCUMENT
|
||||||
|
|
||||||
|
# Determine chat type
|
||||||
|
is_group = data.get("isGroup", False)
|
||||||
|
chat_type = "group" if is_group else "dm"
|
||||||
|
|
||||||
|
# Build source
|
||||||
|
source = self.build_source(
|
||||||
|
chat_id=data.get("chatId", ""),
|
||||||
|
chat_name=data.get("chatName"),
|
||||||
|
chat_type=chat_type,
|
||||||
|
user_id=data.get("senderId"),
|
||||||
|
user_name=data.get("senderName"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Download media URLs to the local cache so agent tools
|
||||||
|
# can access them reliably regardless of URL expiration.
|
||||||
|
raw_urls = data.get("mediaUrls", [])
|
||||||
|
cached_urls = []
|
||||||
|
media_types = []
|
||||||
|
for url in raw_urls:
|
||||||
|
if msg_type == MessageType.PHOTO and url.startswith(("http://", "https://")):
|
||||||
|
try:
|
||||||
|
cached_path = await cache_image_from_url(url, ext=".jpg")
|
||||||
|
cached_urls.append(cached_path)
|
||||||
|
media_types.append("image/jpeg")
|
||||||
|
print(f"[{self.name}] Cached user image: {cached_path}", flush=True)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[{self.name}] Failed to cache image: {e}", flush=True)
|
||||||
|
cached_urls.append(url)
|
||||||
|
media_types.append("image/jpeg")
|
||||||
|
elif msg_type == MessageType.PHOTO and os.path.isabs(url):
|
||||||
|
# Local file path — bridge already downloaded the image
|
||||||
|
cached_urls.append(url)
|
||||||
|
media_types.append("image/jpeg")
|
||||||
|
print(f"[{self.name}] Using bridge-cached image: {url}", flush=True)
|
||||||
|
elif msg_type == MessageType.VOICE and url.startswith(("http://", "https://")):
|
||||||
|
try:
|
||||||
|
cached_path = await cache_audio_from_url(url, ext=".ogg")
|
||||||
|
cached_urls.append(cached_path)
|
||||||
|
media_types.append("audio/ogg")
|
||||||
|
print(f"[{self.name}] Cached user voice: {cached_path}", flush=True)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[{self.name}] Failed to cache voice: {e}", flush=True)
|
||||||
|
cached_urls.append(url)
|
||||||
|
media_types.append("audio/ogg")
|
||||||
|
elif msg_type == MessageType.VOICE and os.path.isabs(url):
|
||||||
|
# Local file path — bridge already downloaded the audio
|
||||||
|
cached_urls.append(url)
|
||||||
|
media_types.append("audio/ogg")
|
||||||
|
print(f"[{self.name}] Using bridge-cached audio: {url}", flush=True)
|
||||||
|
elif msg_type == MessageType.DOCUMENT and os.path.isabs(url):
|
||||||
|
# Local file path — bridge already downloaded the document
|
||||||
|
cached_urls.append(url)
|
||||||
|
ext = Path(url).suffix.lower()
|
||||||
|
mime = SUPPORTED_DOCUMENT_TYPES.get(ext, "application/octet-stream")
|
||||||
|
media_types.append(mime)
|
||||||
|
print(f"[{self.name}] Using bridge-cached document: {url}", flush=True)
|
||||||
|
elif msg_type == MessageType.VIDEO and os.path.isabs(url):
|
||||||
|
cached_urls.append(url)
|
||||||
|
media_types.append("video/mp4")
|
||||||
|
print(f"[{self.name}] Using bridge-cached video: {url}", flush=True)
|
||||||
|
else:
|
||||||
|
cached_urls.append(url)
|
||||||
|
media_types.append("unknown")
|
||||||
|
|
||||||
|
# For text-readable documents, inject file content directly into
|
||||||
|
# the message text so the agent can read it inline.
|
||||||
|
# Cap at 100KB to match Telegram/Discord/Slack behaviour.
|
||||||
|
body = data.get("body", "")
|
||||||
|
if data.get("isGroup"):
|
||||||
|
body = self._clean_bot_mention_text(body, data)
|
||||||
|
MAX_TEXT_INJECT_BYTES = 100 * 1024
|
||||||
|
if msg_type == MessageType.DOCUMENT and cached_urls:
|
||||||
|
for doc_path in cached_urls:
|
||||||
|
ext = Path(doc_path).suffix.lower()
|
||||||
|
if ext in (".txt", ".md", ".csv", ".json", ".xml", ".yaml", ".yml", ".log", ".py", ".js", ".ts", ".html", ".css"):
|
||||||
|
try:
|
||||||
|
file_size = Path(doc_path).stat().st_size
|
||||||
|
if file_size > MAX_TEXT_INJECT_BYTES:
|
||||||
|
print(f"[{self.name}] Skipping text injection for {doc_path} ({file_size} bytes > {MAX_TEXT_INJECT_BYTES})", flush=True)
|
||||||
|
continue
|
||||||
|
content = Path(doc_path).read_text(errors="replace")
|
||||||
|
fname = Path(doc_path).name
|
||||||
|
# Remove the doc_<hex>_ prefix for display
|
||||||
|
display_name = fname
|
||||||
|
if "_" in fname:
|
||||||
|
parts = fname.split("_", 2)
|
||||||
|
if len(parts) >= 3:
|
||||||
|
display_name = parts[2]
|
||||||
|
injection = f"[Content of {display_name}]:\n{content}"
|
||||||
|
if body:
|
||||||
|
body = f"{injection}\n\n{body}"
|
||||||
|
else:
|
||||||
|
body = injection
|
||||||
|
print(f"[{self.name}] Injected text content from: {doc_path}", flush=True)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[{self.name}] Failed to read document text: {e}", flush=True)
|
||||||
|
|
||||||
|
return MessageEvent(
|
||||||
|
text=body,
|
||||||
|
message_type=msg_type,
|
||||||
|
source=source,
|
||||||
|
raw_message=data,
|
||||||
|
message_id=data.get("messageId"),
|
||||||
|
media_urls=cached_urls,
|
||||||
|
media_types=media_types,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[{self.name}] Error building event: {e}")
|
||||||
|
return None
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
"""Shared gateway restart constants and parsing helpers."""
|
||||||
|
|
||||||
|
from hermes_cli.config import DEFAULT_CONFIG
|
||||||
|
|
||||||
|
# EX_TEMPFAIL from sysexits.h — used to ask the service manager to restart
|
||||||
|
# the gateway after a graceful drain/reload path completes.
|
||||||
|
GATEWAY_SERVICE_RESTART_EXIT_CODE = 75
|
||||||
|
|
||||||
|
DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT = float(
|
||||||
|
DEFAULT_CONFIG["agent"]["restart_drain_timeout"]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def parse_restart_drain_timeout(raw: object) -> float:
|
||||||
|
"""Parse a configured drain timeout, falling back to the shared default."""
|
||||||
|
try:
|
||||||
|
value = float(raw) if str(raw or "").strip() else DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT
|
||||||
|
return max(0.0, value)
|
||||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,146 @@
|
|||||||
|
"""
|
||||||
|
Session-scoped context variables for the Hermes gateway.
|
||||||
|
|
||||||
|
Replaces the previous ``os.environ``-based session state
|
||||||
|
(``HERMES_SESSION_PLATFORM``, ``HERMES_SESSION_CHAT_ID``, etc.) with
|
||||||
|
Python's ``contextvars.ContextVar``.
|
||||||
|
|
||||||
|
**Why this matters**
|
||||||
|
|
||||||
|
The gateway processes messages concurrently via ``asyncio``. When two
|
||||||
|
messages arrive at the same time the old code did:
|
||||||
|
|
||||||
|
os.environ["HERMES_SESSION_THREAD_ID"] = str(context.source.thread_id)
|
||||||
|
|
||||||
|
Because ``os.environ`` is *process-global*, Message A's value was
|
||||||
|
silently overwritten by Message B before Message A's agent finished
|
||||||
|
running. Background-task notifications and tool calls therefore routed
|
||||||
|
to the wrong thread.
|
||||||
|
|
||||||
|
``contextvars.ContextVar`` values are *task-local*: each ``asyncio``
|
||||||
|
task (and any ``run_in_executor`` thread it spawns) gets its own copy,
|
||||||
|
so concurrent messages never interfere.
|
||||||
|
|
||||||
|
**Backward compatibility**
|
||||||
|
|
||||||
|
The public helper ``get_session_env(name, default="")`` mirrors the old
|
||||||
|
``os.getenv("HERMES_SESSION_*", ...)`` calls. Existing tool code only
|
||||||
|
needs to replace the import + call site:
|
||||||
|
|
||||||
|
# before
|
||||||
|
import os
|
||||||
|
platform = os.getenv("HERMES_SESSION_PLATFORM", "")
|
||||||
|
|
||||||
|
# after
|
||||||
|
from gateway.session_context import get_session_env
|
||||||
|
platform = get_session_env("HERMES_SESSION_PLATFORM", "")
|
||||||
|
"""
|
||||||
|
|
||||||
|
from contextvars import ContextVar
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Per-task session variables
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
_SESSION_PLATFORM: ContextVar[str] = ContextVar("HERMES_SESSION_PLATFORM", default="")
|
||||||
|
_SESSION_CHAT_ID: ContextVar[str] = ContextVar("HERMES_SESSION_CHAT_ID", default="")
|
||||||
|
_SESSION_CHAT_NAME: ContextVar[str] = ContextVar("HERMES_SESSION_CHAT_NAME", default="")
|
||||||
|
_SESSION_THREAD_ID: ContextVar[str] = ContextVar("HERMES_SESSION_THREAD_ID", default="")
|
||||||
|
_SESSION_USER_ID: ContextVar[str] = ContextVar("HERMES_SESSION_USER_ID", default="")
|
||||||
|
_SESSION_USER_NAME: ContextVar[str] = ContextVar("HERMES_SESSION_USER_NAME", default="")
|
||||||
|
_SESSION_KEY: ContextVar[str] = ContextVar("HERMES_SESSION_KEY", default="")
|
||||||
|
_SESSION_SKILLS_DIRS: ContextVar[str] = ContextVar("HERMES_SESSION_SKILLS_DIRS", default="")
|
||||||
|
|
||||||
|
_VAR_MAP = {
|
||||||
|
"HERMES_SESSION_PLATFORM": _SESSION_PLATFORM,
|
||||||
|
"HERMES_SESSION_CHAT_ID": _SESSION_CHAT_ID,
|
||||||
|
"HERMES_SESSION_CHAT_NAME": _SESSION_CHAT_NAME,
|
||||||
|
"HERMES_SESSION_THREAD_ID": _SESSION_THREAD_ID,
|
||||||
|
"HERMES_SESSION_USER_ID": _SESSION_USER_ID,
|
||||||
|
"HERMES_SESSION_USER_NAME": _SESSION_USER_NAME,
|
||||||
|
"HERMES_SESSION_KEY": _SESSION_KEY,
|
||||||
|
"HERMES_SESSION_SKILLS_DIRS": _SESSION_SKILLS_DIRS,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def set_session_env(name: str, value: str):
|
||||||
|
"""Set a session context variable by its ``HERMES_SESSION_*`` name.
|
||||||
|
|
||||||
|
Returns the reset token (pass to ``var.reset(token)`` to restore).
|
||||||
|
If the variable name is unknown, sets ``os.environ`` as fallback.
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
|
||||||
|
var = _VAR_MAP.get(name)
|
||||||
|
if var is not None:
|
||||||
|
return var.set(value)
|
||||||
|
# Fallback: 未知变量名写入 os.environ(CLI 兼容)
|
||||||
|
os.environ[name] = value
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def set_session_vars(
|
||||||
|
platform: str = "",
|
||||||
|
chat_id: str = "",
|
||||||
|
chat_name: str = "",
|
||||||
|
thread_id: str = "",
|
||||||
|
user_id: str = "",
|
||||||
|
user_name: str = "",
|
||||||
|
session_key: str = "",
|
||||||
|
) -> list:
|
||||||
|
"""Set all session context variables and return reset tokens.
|
||||||
|
|
||||||
|
Call ``clear_session_vars(tokens)`` in a ``finally`` block to restore
|
||||||
|
the previous values when the handler exits.
|
||||||
|
|
||||||
|
Returns a list of ``Token`` objects (one per variable) that can be
|
||||||
|
passed to ``clear_session_vars``.
|
||||||
|
"""
|
||||||
|
tokens = [
|
||||||
|
_SESSION_PLATFORM.set(platform),
|
||||||
|
_SESSION_CHAT_ID.set(chat_id),
|
||||||
|
_SESSION_CHAT_NAME.set(chat_name),
|
||||||
|
_SESSION_THREAD_ID.set(thread_id),
|
||||||
|
_SESSION_USER_ID.set(user_id),
|
||||||
|
_SESSION_USER_NAME.set(user_name),
|
||||||
|
_SESSION_KEY.set(session_key),
|
||||||
|
]
|
||||||
|
return tokens
|
||||||
|
|
||||||
|
|
||||||
|
def clear_session_vars(tokens: list) -> None:
|
||||||
|
"""Restore session context variables to their pre-handler values."""
|
||||||
|
if not tokens:
|
||||||
|
return
|
||||||
|
vars_in_order = [
|
||||||
|
_SESSION_PLATFORM,
|
||||||
|
_SESSION_CHAT_ID,
|
||||||
|
_SESSION_CHAT_NAME,
|
||||||
|
_SESSION_THREAD_ID,
|
||||||
|
_SESSION_USER_ID,
|
||||||
|
_SESSION_USER_NAME,
|
||||||
|
_SESSION_KEY,
|
||||||
|
]
|
||||||
|
for var, token in zip(vars_in_order, tokens):
|
||||||
|
var.reset(token)
|
||||||
|
|
||||||
|
|
||||||
|
def get_session_env(name: str, default: str = "") -> str:
|
||||||
|
"""Read a session context variable by its legacy ``HERMES_SESSION_*`` name.
|
||||||
|
|
||||||
|
Drop-in replacement for ``os.getenv("HERMES_SESSION_*", default)``.
|
||||||
|
|
||||||
|
Resolution order:
|
||||||
|
1. Context variable (set by the gateway for concurrency-safe access)
|
||||||
|
2. ``os.environ`` (used by CLI, cron scheduler, and tests)
|
||||||
|
3. *default*
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
|
||||||
|
var = _VAR_MAP.get(name)
|
||||||
|
if var is not None:
|
||||||
|
value = var.get()
|
||||||
|
if value:
|
||||||
|
return value
|
||||||
|
# Fall back to os.environ for CLI, cron, and test compatibility
|
||||||
|
return os.getenv(name, default)
|
||||||
@@ -0,0 +1,439 @@
|
|||||||
|
"""
|
||||||
|
Gateway runtime status helpers.
|
||||||
|
|
||||||
|
Provides PID-file based detection of whether the gateway daemon is running,
|
||||||
|
used by send_message's check_fn to gate availability in the CLI.
|
||||||
|
|
||||||
|
The PID file lives at ``{HERMES_HOME}/gateway.pid``. HERMES_HOME defaults to
|
||||||
|
``~/.hermes`` but can be overridden via the environment variable. This means
|
||||||
|
separate HERMES_HOME directories naturally get separate PID files — a property
|
||||||
|
that will be useful when we add named profiles (multiple agents running
|
||||||
|
concurrently under distinct configurations).
|
||||||
|
"""
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import signal
|
||||||
|
import subprocess
|
||||||
|
import sys
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from pathlib import Path
|
||||||
|
from hermes_constants import get_hermes_home
|
||||||
|
from typing import Any, Optional
|
||||||
|
|
||||||
|
_GATEWAY_KIND = "hermes-gateway"
|
||||||
|
_RUNTIME_STATUS_FILE = "gateway_state.json"
|
||||||
|
_LOCKS_DIRNAME = "gateway-locks"
|
||||||
|
_IS_WINDOWS = sys.platform == "win32"
|
||||||
|
_UNSET = object()
|
||||||
|
|
||||||
|
|
||||||
|
def _get_pid_path() -> Path:
|
||||||
|
"""Return the path to the gateway PID file, respecting HERMES_HOME."""
|
||||||
|
home = get_hermes_home()
|
||||||
|
return home / "gateway.pid"
|
||||||
|
|
||||||
|
|
||||||
|
def _get_runtime_status_path() -> Path:
|
||||||
|
"""Return the persisted runtime health/status file path."""
|
||||||
|
return _get_pid_path().with_name(_RUNTIME_STATUS_FILE)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_lock_dir() -> Path:
|
||||||
|
"""Return the machine-local directory for token-scoped gateway locks."""
|
||||||
|
override = os.getenv("HERMES_GATEWAY_LOCK_DIR")
|
||||||
|
if override:
|
||||||
|
return Path(override)
|
||||||
|
state_home = Path(os.getenv("XDG_STATE_HOME", Path.home() / ".local" / "state"))
|
||||||
|
return state_home / "hermes" / _LOCKS_DIRNAME
|
||||||
|
|
||||||
|
|
||||||
|
def _utc_now_iso() -> str:
|
||||||
|
return datetime.now(timezone.utc).isoformat()
|
||||||
|
|
||||||
|
|
||||||
|
def terminate_pid(pid: int, *, force: bool = False) -> None:
|
||||||
|
"""Terminate a PID with platform-appropriate force semantics.
|
||||||
|
|
||||||
|
POSIX uses SIGTERM/SIGKILL. Windows uses taskkill /T /F for true force-kill
|
||||||
|
because os.kill(..., SIGTERM) is not equivalent to a tree-killing hard stop.
|
||||||
|
"""
|
||||||
|
if force and _IS_WINDOWS:
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
["taskkill", "/PID", str(pid), "/T", "/F"],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=10,
|
||||||
|
)
|
||||||
|
except FileNotFoundError:
|
||||||
|
os.kill(pid, signal.SIGTERM)
|
||||||
|
return
|
||||||
|
|
||||||
|
if result.returncode != 0:
|
||||||
|
details = (result.stderr or result.stdout or "").strip()
|
||||||
|
raise OSError(details or f"taskkill failed for PID {pid}")
|
||||||
|
return
|
||||||
|
|
||||||
|
sig = signal.SIGTERM if not force else getattr(signal, "SIGKILL", signal.SIGTERM)
|
||||||
|
os.kill(pid, sig)
|
||||||
|
|
||||||
|
|
||||||
|
def _scope_hash(identity: str) -> str:
|
||||||
|
return hashlib.sha256(identity.encode("utf-8")).hexdigest()[:16]
|
||||||
|
|
||||||
|
|
||||||
|
def _get_scope_lock_path(scope: str, identity: str) -> Path:
|
||||||
|
return _get_lock_dir() / f"{scope}-{_scope_hash(identity)}.lock"
|
||||||
|
|
||||||
|
|
||||||
|
def _get_process_start_time(pid: int) -> Optional[int]:
|
||||||
|
"""Return the kernel start time for a process when available."""
|
||||||
|
stat_path = Path(f"/proc/{pid}/stat")
|
||||||
|
try:
|
||||||
|
# Field 22 in /proc/<pid>/stat is process start time (clock ticks).
|
||||||
|
return int(stat_path.read_text().split()[21])
|
||||||
|
except (FileNotFoundError, IndexError, PermissionError, ValueError, OSError):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _read_process_cmdline(pid: int) -> Optional[str]:
|
||||||
|
"""Return the process command line as a space-separated string."""
|
||||||
|
cmdline_path = Path(f"/proc/{pid}/cmdline")
|
||||||
|
try:
|
||||||
|
raw = cmdline_path.read_bytes()
|
||||||
|
except (FileNotFoundError, PermissionError, OSError):
|
||||||
|
return None
|
||||||
|
|
||||||
|
if not raw:
|
||||||
|
return None
|
||||||
|
return raw.replace(b"\x00", b" ").decode("utf-8", errors="ignore").strip()
|
||||||
|
|
||||||
|
|
||||||
|
def _looks_like_gateway_process(pid: int) -> bool:
|
||||||
|
"""Return True when the live PID still looks like the Hermes gateway."""
|
||||||
|
cmdline = _read_process_cmdline(pid)
|
||||||
|
if not cmdline:
|
||||||
|
return False
|
||||||
|
|
||||||
|
patterns = (
|
||||||
|
"hermes_cli.main gateway",
|
||||||
|
"hermes_cli/main.py gateway",
|
||||||
|
"hermes gateway",
|
||||||
|
"gateway/run.py",
|
||||||
|
)
|
||||||
|
return any(pattern in cmdline for pattern in patterns)
|
||||||
|
|
||||||
|
|
||||||
|
def _record_looks_like_gateway(record: dict[str, Any]) -> bool:
|
||||||
|
"""Validate gateway identity from PID-file metadata when cmdline is unavailable."""
|
||||||
|
if record.get("kind") != _GATEWAY_KIND:
|
||||||
|
return False
|
||||||
|
|
||||||
|
argv = record.get("argv")
|
||||||
|
if not isinstance(argv, list) or not argv:
|
||||||
|
return False
|
||||||
|
|
||||||
|
cmdline = " ".join(str(part) for part in argv)
|
||||||
|
patterns = (
|
||||||
|
"hermes_cli.main gateway",
|
||||||
|
"hermes_cli/main.py gateway",
|
||||||
|
"hermes gateway",
|
||||||
|
"gateway/run.py",
|
||||||
|
)
|
||||||
|
return any(pattern in cmdline for pattern in patterns)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_pid_record() -> dict:
|
||||||
|
return {
|
||||||
|
"pid": os.getpid(),
|
||||||
|
"kind": _GATEWAY_KIND,
|
||||||
|
"argv": list(sys.argv),
|
||||||
|
"start_time": _get_process_start_time(os.getpid()),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _build_runtime_status_record() -> dict[str, Any]:
|
||||||
|
payload = _build_pid_record()
|
||||||
|
payload.update({
|
||||||
|
"gateway_state": "starting",
|
||||||
|
"exit_reason": None,
|
||||||
|
"restart_requested": False,
|
||||||
|
"active_agents": 0,
|
||||||
|
"platforms": {},
|
||||||
|
"updated_at": _utc_now_iso(),
|
||||||
|
})
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
def _read_json_file(path: Path) -> Optional[dict[str, Any]]:
|
||||||
|
if not path.exists():
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
raw = path.read_text().strip()
|
||||||
|
except OSError:
|
||||||
|
return None
|
||||||
|
if not raw:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
payload = json.loads(raw)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return None
|
||||||
|
return payload if isinstance(payload, dict) else None
|
||||||
|
|
||||||
|
|
||||||
|
def _write_json_file(path: Path, payload: dict[str, Any]) -> None:
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
path.write_text(json.dumps(payload))
|
||||||
|
|
||||||
|
|
||||||
|
def _read_pid_record() -> Optional[dict]:
|
||||||
|
pid_path = _get_pid_path()
|
||||||
|
if not pid_path.exists():
|
||||||
|
return None
|
||||||
|
|
||||||
|
raw = pid_path.read_text().strip()
|
||||||
|
if not raw:
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
payload = json.loads(raw)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
try:
|
||||||
|
return {"pid": int(raw)}
|
||||||
|
except ValueError:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if isinstance(payload, int):
|
||||||
|
return {"pid": payload}
|
||||||
|
if isinstance(payload, dict):
|
||||||
|
return payload
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def write_pid_file() -> None:
|
||||||
|
"""Write the current process PID and metadata to the gateway PID file."""
|
||||||
|
_write_json_file(_get_pid_path(), _build_pid_record())
|
||||||
|
|
||||||
|
|
||||||
|
def write_runtime_status(
|
||||||
|
*,
|
||||||
|
gateway_state: Any = _UNSET,
|
||||||
|
exit_reason: Any = _UNSET,
|
||||||
|
restart_requested: Any = _UNSET,
|
||||||
|
active_agents: Any = _UNSET,
|
||||||
|
platform: Any = _UNSET,
|
||||||
|
platform_state: Any = _UNSET,
|
||||||
|
error_code: Any = _UNSET,
|
||||||
|
error_message: Any = _UNSET,
|
||||||
|
) -> None:
|
||||||
|
"""Persist gateway runtime health information for diagnostics/status."""
|
||||||
|
path = _get_runtime_status_path()
|
||||||
|
payload = _read_json_file(path) or _build_runtime_status_record()
|
||||||
|
payload.setdefault("platforms", {})
|
||||||
|
payload.setdefault("kind", _GATEWAY_KIND)
|
||||||
|
payload["pid"] = os.getpid()
|
||||||
|
payload["start_time"] = _get_process_start_time(os.getpid())
|
||||||
|
payload["updated_at"] = _utc_now_iso()
|
||||||
|
|
||||||
|
if gateway_state is not _UNSET:
|
||||||
|
payload["gateway_state"] = gateway_state
|
||||||
|
if exit_reason is not _UNSET:
|
||||||
|
payload["exit_reason"] = exit_reason
|
||||||
|
if restart_requested is not _UNSET:
|
||||||
|
payload["restart_requested"] = bool(restart_requested)
|
||||||
|
if active_agents is not _UNSET:
|
||||||
|
payload["active_agents"] = max(0, int(active_agents))
|
||||||
|
|
||||||
|
if platform is not _UNSET:
|
||||||
|
platform_payload = payload["platforms"].get(platform, {})
|
||||||
|
if platform_state is not _UNSET:
|
||||||
|
platform_payload["state"] = platform_state
|
||||||
|
if error_code is not _UNSET:
|
||||||
|
platform_payload["error_code"] = error_code
|
||||||
|
if error_message is not _UNSET:
|
||||||
|
platform_payload["error_message"] = error_message
|
||||||
|
platform_payload["updated_at"] = _utc_now_iso()
|
||||||
|
payload["platforms"][platform] = platform_payload
|
||||||
|
|
||||||
|
_write_json_file(path, payload)
|
||||||
|
|
||||||
|
|
||||||
|
def read_runtime_status() -> Optional[dict[str, Any]]:
|
||||||
|
"""Read the persisted gateway runtime health/status information."""
|
||||||
|
return _read_json_file(_get_runtime_status_path())
|
||||||
|
|
||||||
|
|
||||||
|
def remove_pid_file() -> None:
|
||||||
|
"""Remove the gateway PID file if it exists."""
|
||||||
|
try:
|
||||||
|
_get_pid_path().unlink(missing_ok=True)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def acquire_scoped_lock(scope: str, identity: str, metadata: Optional[dict[str, Any]] = None) -> tuple[bool, Optional[dict[str, Any]]]:
|
||||||
|
"""Acquire a machine-local lock keyed by scope + identity.
|
||||||
|
|
||||||
|
Used to prevent multiple local gateways from using the same external identity
|
||||||
|
at once (e.g. the same Telegram bot token across different HERMES_HOME dirs).
|
||||||
|
"""
|
||||||
|
lock_path = _get_scope_lock_path(scope, identity)
|
||||||
|
lock_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
record = {
|
||||||
|
**_build_pid_record(),
|
||||||
|
"scope": scope,
|
||||||
|
"identity_hash": _scope_hash(identity),
|
||||||
|
"metadata": metadata or {},
|
||||||
|
"updated_at": _utc_now_iso(),
|
||||||
|
}
|
||||||
|
|
||||||
|
existing = _read_json_file(lock_path)
|
||||||
|
if existing is None and lock_path.exists():
|
||||||
|
# Lock file exists but is empty or contains invalid JSON — treat as
|
||||||
|
# stale. This happens when a previous process was killed between
|
||||||
|
# O_CREAT|O_EXCL and the subsequent json.dump() (e.g. DNS failure
|
||||||
|
# during rapid Slack reconnect retries).
|
||||||
|
try:
|
||||||
|
lock_path.unlink(missing_ok=True)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
if existing:
|
||||||
|
try:
|
||||||
|
existing_pid = int(existing["pid"])
|
||||||
|
except (KeyError, TypeError, ValueError):
|
||||||
|
existing_pid = None
|
||||||
|
|
||||||
|
if existing_pid == os.getpid() and existing.get("start_time") == record.get("start_time"):
|
||||||
|
_write_json_file(lock_path, record)
|
||||||
|
return True, existing
|
||||||
|
|
||||||
|
stale = existing_pid is None
|
||||||
|
if not stale:
|
||||||
|
try:
|
||||||
|
os.kill(existing_pid, 0)
|
||||||
|
except (ProcessLookupError, PermissionError):
|
||||||
|
stale = True
|
||||||
|
else:
|
||||||
|
current_start = _get_process_start_time(existing_pid)
|
||||||
|
if (
|
||||||
|
existing.get("start_time") is not None
|
||||||
|
and current_start is not None
|
||||||
|
and current_start != existing.get("start_time")
|
||||||
|
):
|
||||||
|
stale = True
|
||||||
|
# Check if process is stopped (Ctrl+Z / SIGTSTP) — stopped
|
||||||
|
# processes still respond to os.kill(pid, 0) but are not
|
||||||
|
# actually running. Treat them as stale so --replace works.
|
||||||
|
if not stale:
|
||||||
|
try:
|
||||||
|
_proc_status = Path(f"/proc/{existing_pid}/status")
|
||||||
|
if _proc_status.exists():
|
||||||
|
for _line in _proc_status.read_text().splitlines():
|
||||||
|
if _line.startswith("State:"):
|
||||||
|
_state = _line.split()[1]
|
||||||
|
if _state in ("T", "t"): # stopped or tracing stop
|
||||||
|
stale = True
|
||||||
|
break
|
||||||
|
except (OSError, PermissionError):
|
||||||
|
pass
|
||||||
|
if stale:
|
||||||
|
try:
|
||||||
|
lock_path.unlink(missing_ok=True)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
return False, existing
|
||||||
|
|
||||||
|
try:
|
||||||
|
fd = os.open(lock_path, os.O_CREAT | os.O_EXCL | os.O_WRONLY)
|
||||||
|
except FileExistsError:
|
||||||
|
return False, _read_json_file(lock_path)
|
||||||
|
try:
|
||||||
|
with os.fdopen(fd, "w", encoding="utf-8") as handle:
|
||||||
|
json.dump(record, handle)
|
||||||
|
except Exception:
|
||||||
|
try:
|
||||||
|
lock_path.unlink(missing_ok=True)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
raise
|
||||||
|
return True, None
|
||||||
|
|
||||||
|
|
||||||
|
def release_scoped_lock(scope: str, identity: str) -> None:
|
||||||
|
"""Release a previously-acquired scope lock when owned by this process."""
|
||||||
|
lock_path = _get_scope_lock_path(scope, identity)
|
||||||
|
existing = _read_json_file(lock_path)
|
||||||
|
if not existing:
|
||||||
|
return
|
||||||
|
if existing.get("pid") != os.getpid():
|
||||||
|
return
|
||||||
|
if existing.get("start_time") != _get_process_start_time(os.getpid()):
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
lock_path.unlink(missing_ok=True)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def release_all_scoped_locks() -> int:
|
||||||
|
"""Remove all scoped lock files in the lock directory.
|
||||||
|
|
||||||
|
Called during --replace to clean up stale locks left by stopped/killed
|
||||||
|
gateway processes that did not release their locks gracefully.
|
||||||
|
Returns the number of lock files removed.
|
||||||
|
"""
|
||||||
|
lock_dir = _get_lock_dir()
|
||||||
|
removed = 0
|
||||||
|
if lock_dir.exists():
|
||||||
|
for lock_file in lock_dir.glob("*.lock"):
|
||||||
|
try:
|
||||||
|
lock_file.unlink(missing_ok=True)
|
||||||
|
removed += 1
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
return removed
|
||||||
|
|
||||||
|
|
||||||
|
def get_running_pid() -> Optional[int]:
|
||||||
|
"""Return the PID of a running gateway instance, or ``None``.
|
||||||
|
|
||||||
|
Checks the PID file and verifies the process is actually alive.
|
||||||
|
Cleans up stale PID files automatically.
|
||||||
|
"""
|
||||||
|
record = _read_pid_record()
|
||||||
|
if not record:
|
||||||
|
remove_pid_file()
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
pid = int(record["pid"])
|
||||||
|
except (KeyError, TypeError, ValueError):
|
||||||
|
remove_pid_file()
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
os.kill(pid, 0) # signal 0 = existence check, no actual signal sent
|
||||||
|
except (ProcessLookupError, PermissionError):
|
||||||
|
remove_pid_file()
|
||||||
|
return None
|
||||||
|
|
||||||
|
recorded_start = record.get("start_time")
|
||||||
|
current_start = _get_process_start_time(pid)
|
||||||
|
if recorded_start is not None and current_start is not None and current_start != recorded_start:
|
||||||
|
remove_pid_file()
|
||||||
|
return None
|
||||||
|
|
||||||
|
if not _looks_like_gateway_process(pid):
|
||||||
|
if not _record_looks_like_gateway(record):
|
||||||
|
remove_pid_file()
|
||||||
|
return None
|
||||||
|
|
||||||
|
return pid
|
||||||
|
|
||||||
|
|
||||||
|
def is_gateway_running() -> bool:
|
||||||
|
"""Check if the gateway daemon is currently running."""
|
||||||
|
return get_running_pid() is not None
|
||||||
@@ -0,0 +1,111 @@
|
|||||||
|
"""
|
||||||
|
Sticker description cache for Telegram.
|
||||||
|
|
||||||
|
When users send stickers, we describe them via the vision tool and cache
|
||||||
|
the descriptions keyed by file_unique_id so we don't re-analyze the same
|
||||||
|
sticker image on every send. Descriptions are concise (1-2 sentences).
|
||||||
|
|
||||||
|
Cache location: ~/.hermes/sticker_cache.json
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from hermes_cli.config import get_hermes_home
|
||||||
|
|
||||||
|
|
||||||
|
CACHE_PATH = get_hermes_home() / "sticker_cache.json"
|
||||||
|
|
||||||
|
# Vision prompt for describing stickers -- kept concise to save tokens
|
||||||
|
STICKER_VISION_PROMPT = (
|
||||||
|
"Describe this sticker in 1-2 sentences. Focus on what it depicts -- "
|
||||||
|
"character, action, emotion. Be concise and objective."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _load_cache() -> dict:
|
||||||
|
"""Load the sticker cache from disk."""
|
||||||
|
if CACHE_PATH.exists():
|
||||||
|
try:
|
||||||
|
return json.loads(CACHE_PATH.read_text(encoding="utf-8"))
|
||||||
|
except (json.JSONDecodeError, OSError):
|
||||||
|
return {}
|
||||||
|
return {}
|
||||||
|
|
||||||
|
|
||||||
|
def _save_cache(cache: dict) -> None:
|
||||||
|
"""Save the sticker cache to disk."""
|
||||||
|
CACHE_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
CACHE_PATH.write_text(
|
||||||
|
json.dumps(cache, indent=2, ensure_ascii=False),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_cached_description(file_unique_id: str) -> Optional[dict]:
|
||||||
|
"""
|
||||||
|
Look up a cached sticker description.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
dict with keys {description, emoji, set_name, cached_at} or None.
|
||||||
|
"""
|
||||||
|
cache = _load_cache()
|
||||||
|
return cache.get(file_unique_id)
|
||||||
|
|
||||||
|
|
||||||
|
def cache_sticker_description(
|
||||||
|
file_unique_id: str,
|
||||||
|
description: str,
|
||||||
|
emoji: str = "",
|
||||||
|
set_name: str = "",
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Store a sticker description in the cache.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_unique_id: Telegram's stable sticker identifier.
|
||||||
|
description: Vision-generated description text.
|
||||||
|
emoji: Associated emoji (e.g. "😀").
|
||||||
|
set_name: Sticker set name if available.
|
||||||
|
"""
|
||||||
|
cache = _load_cache()
|
||||||
|
cache[file_unique_id] = {
|
||||||
|
"description": description,
|
||||||
|
"emoji": emoji,
|
||||||
|
"set_name": set_name,
|
||||||
|
"cached_at": time.time(),
|
||||||
|
}
|
||||||
|
_save_cache(cache)
|
||||||
|
|
||||||
|
|
||||||
|
def build_sticker_injection(
|
||||||
|
description: str,
|
||||||
|
emoji: str = "",
|
||||||
|
set_name: str = "",
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
Build the warm-style injection text for a sticker description.
|
||||||
|
|
||||||
|
Returns a string like:
|
||||||
|
[The user sent a sticker 😀 from "MyPack"~ It shows: "A cat waving" (=^.w.^=)]
|
||||||
|
"""
|
||||||
|
context = ""
|
||||||
|
if set_name and emoji:
|
||||||
|
context = f" {emoji} from \"{set_name}\""
|
||||||
|
elif emoji:
|
||||||
|
context = f" {emoji}"
|
||||||
|
|
||||||
|
return f"[The user sent a sticker{context}~ It shows: \"{description}\" (=^.w.^=)]"
|
||||||
|
|
||||||
|
|
||||||
|
def build_animated_sticker_injection(emoji: str = "") -> str:
|
||||||
|
"""
|
||||||
|
Build injection text for animated/video stickers we can't analyze.
|
||||||
|
"""
|
||||||
|
if emoji:
|
||||||
|
return (
|
||||||
|
f"[The user sent an animated sticker {emoji}~ "
|
||||||
|
f"I can't see animated ones yet, but the emoji suggests: {emoji}]"
|
||||||
|
)
|
||||||
|
return "[The user sent an animated sticker~ I can't see animated ones yet]"
|
||||||
@@ -0,0 +1,744 @@
|
|||||||
|
"""Gateway streaming consumer — bridges sync agent callbacks to async platform delivery.
|
||||||
|
|
||||||
|
The agent fires stream_delta_callback(text) synchronously from its worker thread.
|
||||||
|
GatewayStreamConsumer:
|
||||||
|
1. Receives deltas via on_delta() (thread-safe, sync)
|
||||||
|
2. Queues them to an asyncio task via queue.Queue
|
||||||
|
3. The async run() task buffers, rate-limits, and progressively edits
|
||||||
|
a single message on the target platform
|
||||||
|
|
||||||
|
Design: Uses the edit transport (send initial message, then editMessageText).
|
||||||
|
This is universally supported across Telegram, Discord, and Slack.
|
||||||
|
|
||||||
|
Credit: jobless0x (#774, #1312), OutThisLife (#798), clicksingh (#697).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import queue
|
||||||
|
import re
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Optional
|
||||||
|
|
||||||
|
logger = logging.getLogger("gateway.stream_consumer")
|
||||||
|
|
||||||
|
# Sentinel to signal the stream is complete
|
||||||
|
_DONE = object()
|
||||||
|
|
||||||
|
# Sentinel to signal a tool boundary — finalize current message and start a
|
||||||
|
# new one so that subsequent text appears below tool progress messages.
|
||||||
|
_NEW_SEGMENT = object()
|
||||||
|
|
||||||
|
# Queue marker for a completed assistant commentary message emitted between
|
||||||
|
# API/tool iterations (for example: "I'll inspect the repo first.").
|
||||||
|
_COMMENTARY = object()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class StreamConsumerConfig:
|
||||||
|
"""Runtime config for a single stream consumer instance."""
|
||||||
|
edit_interval: float = 1.0
|
||||||
|
buffer_threshold: int = 40
|
||||||
|
cursor: str = " ▉"
|
||||||
|
|
||||||
|
|
||||||
|
class GatewayStreamConsumer:
|
||||||
|
"""Async consumer that progressively edits a platform message with streamed tokens.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
consumer = GatewayStreamConsumer(adapter, chat_id, config, metadata=metadata)
|
||||||
|
# Pass consumer.on_delta as stream_delta_callback to AIAgent
|
||||||
|
agent = AIAgent(..., stream_delta_callback=consumer.on_delta)
|
||||||
|
# Start the consumer as an asyncio task
|
||||||
|
task = asyncio.create_task(consumer.run())
|
||||||
|
# ... run agent in thread pool ...
|
||||||
|
consumer.finish() # signal completion
|
||||||
|
await task # wait for final edit
|
||||||
|
"""
|
||||||
|
|
||||||
|
# After this many consecutive flood-control failures, permanently disable
|
||||||
|
# progressive edits for the remainder of the stream.
|
||||||
|
_MAX_FLOOD_STRIKES = 3
|
||||||
|
|
||||||
|
# Reasoning/thinking tags that models emit inline in content.
|
||||||
|
# Must stay in sync with cli.py _OPEN_TAGS/_CLOSE_TAGS and
|
||||||
|
# run_agent.py _strip_think_blocks() tag variants.
|
||||||
|
_OPEN_THINK_TAGS = (
|
||||||
|
"<REASONING_SCRATCHPAD>", "<think>", "<reasoning>",
|
||||||
|
"<THINKING>", "<thinking>", "<thought>",
|
||||||
|
)
|
||||||
|
_CLOSE_THINK_TAGS = (
|
||||||
|
"</REASONING_SCRATCHPAD>", "</think>", "</reasoning>",
|
||||||
|
"</THINKING>", "</thinking>", "</thought>",
|
||||||
|
)
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
adapter: Any,
|
||||||
|
chat_id: str,
|
||||||
|
config: Optional[StreamConsumerConfig] = None,
|
||||||
|
metadata: Optional[dict] = None,
|
||||||
|
):
|
||||||
|
self.adapter = adapter
|
||||||
|
self.chat_id = chat_id
|
||||||
|
self.cfg = config or StreamConsumerConfig()
|
||||||
|
self.metadata = metadata
|
||||||
|
self._queue: queue.Queue = queue.Queue()
|
||||||
|
self._accumulated = ""
|
||||||
|
self._message_id: Optional[str] = None
|
||||||
|
self._already_sent = False
|
||||||
|
self._edit_supported = True # Disabled when progressive edits are no longer usable
|
||||||
|
self._last_edit_time = 0.0
|
||||||
|
self._last_sent_text = "" # Track last-sent text to skip redundant edits
|
||||||
|
self._fallback_final_send = False
|
||||||
|
self._fallback_prefix = ""
|
||||||
|
self._flood_strikes = 0 # Consecutive flood-control edit failures
|
||||||
|
self._current_edit_interval = self.cfg.edit_interval # Adaptive backoff
|
||||||
|
self._final_response_sent = False
|
||||||
|
|
||||||
|
# Think-block filter state (mirrors CLI's _stream_delta tag suppression)
|
||||||
|
self._in_think_block = False
|
||||||
|
self._think_buffer = ""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def already_sent(self) -> bool:
|
||||||
|
"""True if at least one message was sent or edited during the run."""
|
||||||
|
return self._already_sent
|
||||||
|
|
||||||
|
@property
|
||||||
|
def final_response_sent(self) -> bool:
|
||||||
|
"""True when the stream consumer delivered the final assistant reply."""
|
||||||
|
return self._final_response_sent
|
||||||
|
|
||||||
|
def on_segment_break(self) -> None:
|
||||||
|
"""Finalize the current stream segment and start a fresh message."""
|
||||||
|
self._queue.put(_NEW_SEGMENT)
|
||||||
|
|
||||||
|
def on_commentary(self, text: str) -> None:
|
||||||
|
"""Queue a completed interim assistant commentary message."""
|
||||||
|
if text:
|
||||||
|
self._queue.put((_COMMENTARY, text))
|
||||||
|
|
||||||
|
def _reset_segment_state(self, *, preserve_no_edit: bool = False) -> None:
|
||||||
|
if preserve_no_edit and self._message_id == "__no_edit__":
|
||||||
|
return
|
||||||
|
self._message_id = None
|
||||||
|
self._accumulated = ""
|
||||||
|
self._last_sent_text = ""
|
||||||
|
self._fallback_final_send = False
|
||||||
|
self._fallback_prefix = ""
|
||||||
|
|
||||||
|
def on_delta(self, text: str) -> None:
|
||||||
|
"""Thread-safe callback — called from the agent's worker thread.
|
||||||
|
|
||||||
|
When *text* is ``None``, signals a tool boundary: the current message
|
||||||
|
is finalized and subsequent text will be sent as a new message so it
|
||||||
|
appears below any tool-progress messages the gateway sent in between.
|
||||||
|
"""
|
||||||
|
if text:
|
||||||
|
self._queue.put(text)
|
||||||
|
elif text is None:
|
||||||
|
self.on_segment_break()
|
||||||
|
|
||||||
|
def finish(self) -> None:
|
||||||
|
"""Signal that the stream is complete."""
|
||||||
|
self._queue.put(_DONE)
|
||||||
|
|
||||||
|
# ── Think-block filtering ────────────────────────────────────────
|
||||||
|
# Models like MiniMax emit inline <think>...</think> blocks in their
|
||||||
|
# content. The CLI's _stream_delta suppresses these via a state
|
||||||
|
# machine; we do the same here so gateway users never see raw
|
||||||
|
# reasoning tags. The agent also strips them from the final
|
||||||
|
# response (run_agent.py _strip_think_blocks), but the stream
|
||||||
|
# consumer sends intermediate edits before that stripping happens.
|
||||||
|
|
||||||
|
def _filter_and_accumulate(self, text: str) -> None:
|
||||||
|
"""Add a text delta to the accumulated buffer, suppressing think blocks.
|
||||||
|
|
||||||
|
Uses a state machine that tracks whether we are inside a
|
||||||
|
reasoning/thinking block. Text inside such blocks is silently
|
||||||
|
discarded. Partial tags at buffer boundaries are held back in
|
||||||
|
``_think_buffer`` until enough characters arrive to decide.
|
||||||
|
"""
|
||||||
|
buf = self._think_buffer + text
|
||||||
|
self._think_buffer = ""
|
||||||
|
|
||||||
|
while buf:
|
||||||
|
if self._in_think_block:
|
||||||
|
# Look for the earliest closing tag
|
||||||
|
best_idx = -1
|
||||||
|
best_len = 0
|
||||||
|
for tag in self._CLOSE_THINK_TAGS:
|
||||||
|
idx = buf.find(tag)
|
||||||
|
if idx != -1 and (best_idx == -1 or idx < best_idx):
|
||||||
|
best_idx = idx
|
||||||
|
best_len = len(tag)
|
||||||
|
|
||||||
|
if best_len:
|
||||||
|
# Found closing tag — discard block, process remainder
|
||||||
|
self._in_think_block = False
|
||||||
|
buf = buf[best_idx + best_len:]
|
||||||
|
else:
|
||||||
|
# No closing tag yet — hold tail that could be a
|
||||||
|
# partial closing tag prefix, discard the rest.
|
||||||
|
max_tag = max(len(t) for t in self._CLOSE_THINK_TAGS)
|
||||||
|
self._think_buffer = buf[-max_tag:] if len(buf) > max_tag else buf
|
||||||
|
return
|
||||||
|
else:
|
||||||
|
# Look for earliest opening tag at a block boundary
|
||||||
|
# (start of text / preceded by newline + optional whitespace).
|
||||||
|
# This prevents false positives when models *mention* tags
|
||||||
|
# in prose (e.g. "the <think> tag is used for…").
|
||||||
|
best_idx = -1
|
||||||
|
best_len = 0
|
||||||
|
for tag in self._OPEN_THINK_TAGS:
|
||||||
|
search_start = 0
|
||||||
|
while True:
|
||||||
|
idx = buf.find(tag, search_start)
|
||||||
|
if idx == -1:
|
||||||
|
break
|
||||||
|
# Block-boundary check (mirrors cli.py logic)
|
||||||
|
if idx == 0:
|
||||||
|
is_boundary = (
|
||||||
|
not self._accumulated
|
||||||
|
or self._accumulated.endswith("\n")
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
preceding = buf[:idx]
|
||||||
|
last_nl = preceding.rfind("\n")
|
||||||
|
if last_nl == -1:
|
||||||
|
is_boundary = (
|
||||||
|
(not self._accumulated
|
||||||
|
or self._accumulated.endswith("\n"))
|
||||||
|
and preceding.strip() == ""
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
is_boundary = preceding[last_nl + 1:].strip() == ""
|
||||||
|
|
||||||
|
if is_boundary and (best_idx == -1 or idx < best_idx):
|
||||||
|
best_idx = idx
|
||||||
|
best_len = len(tag)
|
||||||
|
break # first boundary hit for this tag is enough
|
||||||
|
search_start = idx + 1
|
||||||
|
|
||||||
|
if best_len:
|
||||||
|
# Emit text before the tag, enter think block
|
||||||
|
self._accumulated += buf[:best_idx]
|
||||||
|
self._in_think_block = True
|
||||||
|
buf = buf[best_idx + best_len:]
|
||||||
|
else:
|
||||||
|
# No opening tag — check for a partial tag at the tail
|
||||||
|
held_back = 0
|
||||||
|
for tag in self._OPEN_THINK_TAGS:
|
||||||
|
for i in range(1, len(tag)):
|
||||||
|
if buf.endswith(tag[:i]) and i > held_back:
|
||||||
|
held_back = i
|
||||||
|
if held_back:
|
||||||
|
self._accumulated += buf[:-held_back]
|
||||||
|
self._think_buffer = buf[-held_back:]
|
||||||
|
else:
|
||||||
|
self._accumulated += buf
|
||||||
|
return
|
||||||
|
|
||||||
|
def _flush_think_buffer(self) -> None:
|
||||||
|
"""Flush any held-back partial-tag buffer into accumulated text.
|
||||||
|
|
||||||
|
Called when the stream ends (got_done) so that partial text that
|
||||||
|
was held back waiting for a possible opening tag is not lost.
|
||||||
|
"""
|
||||||
|
if self._think_buffer and not self._in_think_block:
|
||||||
|
self._accumulated += self._think_buffer
|
||||||
|
self._think_buffer = ""
|
||||||
|
|
||||||
|
async def run(self) -> None:
|
||||||
|
"""Async task that drains the queue and edits the platform message."""
|
||||||
|
# Platform message length limit — leave room for cursor + formatting
|
||||||
|
_raw_limit = getattr(self.adapter, "MAX_MESSAGE_LENGTH", 4096)
|
||||||
|
_safe_limit = max(500, _raw_limit - len(self.cfg.cursor) - 100)
|
||||||
|
|
||||||
|
try:
|
||||||
|
while True:
|
||||||
|
# Drain all available items from the queue
|
||||||
|
got_done = False
|
||||||
|
got_segment_break = False
|
||||||
|
commentary_text = None
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
item = self._queue.get_nowait()
|
||||||
|
if item is _DONE:
|
||||||
|
got_done = True
|
||||||
|
break
|
||||||
|
if item is _NEW_SEGMENT:
|
||||||
|
got_segment_break = True
|
||||||
|
break
|
||||||
|
if isinstance(item, tuple) and len(item) == 2 and item[0] is _COMMENTARY:
|
||||||
|
commentary_text = item[1]
|
||||||
|
break
|
||||||
|
self._filter_and_accumulate(item)
|
||||||
|
except queue.Empty:
|
||||||
|
break
|
||||||
|
|
||||||
|
# Flush any held-back partial-tag buffer on stream end
|
||||||
|
# so trailing text that was waiting for a potential open
|
||||||
|
# tag is not lost.
|
||||||
|
if got_done:
|
||||||
|
self._flush_think_buffer()
|
||||||
|
|
||||||
|
# Decide whether to flush an edit
|
||||||
|
now = time.monotonic()
|
||||||
|
elapsed = now - self._last_edit_time
|
||||||
|
should_edit = (
|
||||||
|
got_done
|
||||||
|
or got_segment_break
|
||||||
|
or commentary_text is not None
|
||||||
|
or (elapsed >= self._current_edit_interval
|
||||||
|
and self._accumulated)
|
||||||
|
or len(self._accumulated) >= self.cfg.buffer_threshold
|
||||||
|
)
|
||||||
|
|
||||||
|
current_update_visible = False
|
||||||
|
if should_edit and self._accumulated:
|
||||||
|
# Split overflow: if accumulated text exceeds the platform
|
||||||
|
# limit, split into properly sized chunks.
|
||||||
|
if (
|
||||||
|
len(self._accumulated) > _safe_limit
|
||||||
|
and self._message_id is None
|
||||||
|
):
|
||||||
|
# No existing message to edit (first message or after a
|
||||||
|
# segment break). Use truncate_message — the same
|
||||||
|
# helper the non-streaming path uses — to split with
|
||||||
|
# proper word/code-fence boundaries and chunk
|
||||||
|
# indicators like "(1/2)".
|
||||||
|
chunks = self.adapter.truncate_message(
|
||||||
|
self._accumulated, _safe_limit
|
||||||
|
)
|
||||||
|
for chunk in chunks:
|
||||||
|
await self._send_new_chunk(chunk, self._message_id)
|
||||||
|
self._accumulated = ""
|
||||||
|
self._last_sent_text = ""
|
||||||
|
self._last_edit_time = time.monotonic()
|
||||||
|
if got_done:
|
||||||
|
self._final_response_sent = self._already_sent
|
||||||
|
return
|
||||||
|
if got_segment_break:
|
||||||
|
self._message_id = None
|
||||||
|
self._fallback_final_send = False
|
||||||
|
self._fallback_prefix = ""
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Existing message: edit it with the first chunk, then
|
||||||
|
# start a new message for the overflow remainder.
|
||||||
|
while (
|
||||||
|
len(self._accumulated) > _safe_limit
|
||||||
|
and self._message_id is not None
|
||||||
|
and self._edit_supported
|
||||||
|
):
|
||||||
|
split_at = self._accumulated.rfind("\n", 0, _safe_limit)
|
||||||
|
if split_at < _safe_limit // 2:
|
||||||
|
split_at = _safe_limit
|
||||||
|
chunk = self._accumulated[:split_at]
|
||||||
|
ok = await self._send_or_edit(chunk)
|
||||||
|
if self._fallback_final_send or not ok:
|
||||||
|
# Edit failed (or backed off due to flood control)
|
||||||
|
# while attempting to split an oversized message.
|
||||||
|
# Keep the full accumulated text intact so the
|
||||||
|
# fallback final-send path can deliver the remaining
|
||||||
|
# continuation without dropping content.
|
||||||
|
break
|
||||||
|
self._accumulated = self._accumulated[split_at:].lstrip("\n")
|
||||||
|
self._message_id = None
|
||||||
|
self._last_sent_text = ""
|
||||||
|
|
||||||
|
display_text = self._accumulated
|
||||||
|
if not got_done and not got_segment_break and commentary_text is None:
|
||||||
|
display_text += self.cfg.cursor
|
||||||
|
|
||||||
|
current_update_visible = await self._send_or_edit(display_text)
|
||||||
|
self._last_edit_time = time.monotonic()
|
||||||
|
|
||||||
|
if got_done:
|
||||||
|
# Final edit without cursor. If progressive editing failed
|
||||||
|
# mid-stream, send a single continuation/fallback message
|
||||||
|
# here instead of letting the base gateway path send the
|
||||||
|
# full response again.
|
||||||
|
if self._accumulated:
|
||||||
|
if self._fallback_final_send:
|
||||||
|
await self._send_fallback_final(self._accumulated)
|
||||||
|
elif current_update_visible:
|
||||||
|
self._final_response_sent = True
|
||||||
|
elif self._message_id:
|
||||||
|
self._final_response_sent = await self._send_or_edit(self._accumulated)
|
||||||
|
elif not self._already_sent:
|
||||||
|
self._final_response_sent = await self._send_or_edit(self._accumulated)
|
||||||
|
return
|
||||||
|
|
||||||
|
if commentary_text is not None:
|
||||||
|
self._reset_segment_state()
|
||||||
|
await self._send_commentary(commentary_text)
|
||||||
|
self._last_edit_time = time.monotonic()
|
||||||
|
self._reset_segment_state()
|
||||||
|
|
||||||
|
# Tool boundary: reset message state so the next text chunk
|
||||||
|
# creates a fresh message below any tool-progress messages.
|
||||||
|
#
|
||||||
|
# Exception: when _message_id is "__no_edit__" the platform
|
||||||
|
# never returned a real message ID (e.g. Signal, webhook with
|
||||||
|
# github_comment delivery). Resetting to None would re-enter
|
||||||
|
# the "first send" path on every tool boundary and post one
|
||||||
|
# platform message per tool call — that is what caused 155
|
||||||
|
# comments under a single PR. Instead, preserve the sentinel
|
||||||
|
# so the full continuation is delivered once via
|
||||||
|
# _send_fallback_final.
|
||||||
|
# (When editing fails mid-stream due to flood control the id is
|
||||||
|
# a real string like "msg_1", not "__no_edit__", so that case
|
||||||
|
# still resets and creates a fresh segment as intended.)
|
||||||
|
if got_segment_break:
|
||||||
|
self._reset_segment_state(preserve_no_edit=True)
|
||||||
|
|
||||||
|
await asyncio.sleep(0.05) # Small yield to not busy-loop
|
||||||
|
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
# Best-effort final edit on cancellation
|
||||||
|
if self._accumulated and self._message_id:
|
||||||
|
try:
|
||||||
|
await self._send_or_edit(self._accumulated)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
# If we delivered any content before being cancelled, mark the
|
||||||
|
# final response as sent so the gateway's already_sent check
|
||||||
|
# doesn't trigger a duplicate message. The 5-second
|
||||||
|
# stream_task timeout (gateway/run.py) can cancel us while
|
||||||
|
# waiting on a slow Telegram API call — without this flag the
|
||||||
|
# gateway falls through to the normal send path.
|
||||||
|
if self._already_sent:
|
||||||
|
self._final_response_sent = True
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Stream consumer error: %s", e)
|
||||||
|
|
||||||
|
# Pattern to strip MEDIA:<path> tags (including optional surrounding quotes).
|
||||||
|
# Matches the simple cleanup regex used by the non-streaming path in
|
||||||
|
# gateway/platforms/base.py for post-processing.
|
||||||
|
_MEDIA_RE = re.compile(r'''[`"']?MEDIA:\s*\S+[`"']?''')
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _clean_for_display(text: str) -> str:
|
||||||
|
"""Strip MEDIA: directives and internal markers from text before display.
|
||||||
|
|
||||||
|
The streaming path delivers raw text chunks that may include
|
||||||
|
``MEDIA:<path>`` tags and ``[[audio_as_voice]]`` directives meant for
|
||||||
|
the platform adapter's post-processing. The actual media files are
|
||||||
|
delivered separately via ``_deliver_media_from_response()`` after the
|
||||||
|
stream finishes — we just need to hide the raw directives from the
|
||||||
|
user.
|
||||||
|
"""
|
||||||
|
if "MEDIA:" not in text and "[[audio_as_voice]]" not in text:
|
||||||
|
return text
|
||||||
|
cleaned = text.replace("[[audio_as_voice]]", "")
|
||||||
|
cleaned = GatewayStreamConsumer._MEDIA_RE.sub("", cleaned)
|
||||||
|
# Collapse excessive blank lines left behind by removed tags
|
||||||
|
cleaned = re.sub(r'\n{3,}', '\n\n', cleaned)
|
||||||
|
# Strip trailing whitespace/newlines but preserve leading content
|
||||||
|
return cleaned.rstrip()
|
||||||
|
|
||||||
|
async def _send_new_chunk(self, text: str, reply_to_id: Optional[str]) -> Optional[str]:
|
||||||
|
"""Send a new message chunk, optionally threaded to a previous message.
|
||||||
|
|
||||||
|
Returns the message_id so callers can thread subsequent chunks.
|
||||||
|
"""
|
||||||
|
text = self._clean_for_display(text)
|
||||||
|
if not text.strip():
|
||||||
|
return reply_to_id
|
||||||
|
try:
|
||||||
|
meta = dict(self.metadata) if self.metadata else {}
|
||||||
|
result = await self.adapter.send(
|
||||||
|
chat_id=self.chat_id,
|
||||||
|
content=text,
|
||||||
|
reply_to=reply_to_id,
|
||||||
|
metadata=meta,
|
||||||
|
)
|
||||||
|
if result.success and result.message_id:
|
||||||
|
self._message_id = str(result.message_id)
|
||||||
|
self._already_sent = True
|
||||||
|
self._last_sent_text = text
|
||||||
|
return str(result.message_id)
|
||||||
|
else:
|
||||||
|
self._edit_supported = False
|
||||||
|
return reply_to_id
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Stream send chunk error: %s", e)
|
||||||
|
return reply_to_id
|
||||||
|
|
||||||
|
def _visible_prefix(self) -> str:
|
||||||
|
"""Return the visible text already shown in the streamed message."""
|
||||||
|
prefix = self._last_sent_text or ""
|
||||||
|
if self.cfg.cursor and prefix.endswith(self.cfg.cursor):
|
||||||
|
prefix = prefix[:-len(self.cfg.cursor)]
|
||||||
|
return self._clean_for_display(prefix)
|
||||||
|
|
||||||
|
def _continuation_text(self, final_text: str) -> str:
|
||||||
|
"""Return only the part of final_text the user has not already seen."""
|
||||||
|
prefix = self._fallback_prefix or self._visible_prefix()
|
||||||
|
if prefix and final_text.startswith(prefix):
|
||||||
|
return final_text[len(prefix):].lstrip()
|
||||||
|
return final_text
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _split_text_chunks(text: str, limit: int) -> list[str]:
|
||||||
|
"""Split text into reasonably sized chunks for fallback sends."""
|
||||||
|
if len(text) <= limit:
|
||||||
|
return [text]
|
||||||
|
chunks: list[str] = []
|
||||||
|
remaining = text
|
||||||
|
while len(remaining) > limit:
|
||||||
|
split_at = remaining.rfind("\n", 0, limit)
|
||||||
|
if split_at < limit // 2:
|
||||||
|
split_at = limit
|
||||||
|
chunks.append(remaining[:split_at])
|
||||||
|
remaining = remaining[split_at:].lstrip("\n")
|
||||||
|
if remaining:
|
||||||
|
chunks.append(remaining)
|
||||||
|
return chunks
|
||||||
|
|
||||||
|
async def _send_fallback_final(self, text: str) -> None:
|
||||||
|
"""Send the final continuation after streaming edits stop working.
|
||||||
|
|
||||||
|
Retries each chunk once on flood-control failures with a short delay.
|
||||||
|
"""
|
||||||
|
final_text = self._clean_for_display(text)
|
||||||
|
continuation = self._continuation_text(final_text)
|
||||||
|
self._fallback_final_send = False
|
||||||
|
if not continuation.strip():
|
||||||
|
# Nothing new to send — the visible partial already matches final text.
|
||||||
|
self._already_sent = True
|
||||||
|
self._final_response_sent = True
|
||||||
|
return
|
||||||
|
|
||||||
|
raw_limit = getattr(self.adapter, "MAX_MESSAGE_LENGTH", 4096)
|
||||||
|
safe_limit = max(500, raw_limit - 100)
|
||||||
|
chunks = self._split_text_chunks(continuation, safe_limit)
|
||||||
|
|
||||||
|
last_message_id: Optional[str] = None
|
||||||
|
last_successful_chunk = ""
|
||||||
|
sent_any_chunk = False
|
||||||
|
for chunk in chunks:
|
||||||
|
# Try sending with one retry on flood-control errors.
|
||||||
|
result = None
|
||||||
|
for attempt in range(2):
|
||||||
|
result = await self.adapter.send(
|
||||||
|
chat_id=self.chat_id,
|
||||||
|
content=chunk,
|
||||||
|
metadata=self.metadata,
|
||||||
|
)
|
||||||
|
if result.success:
|
||||||
|
break
|
||||||
|
if attempt == 0 and self._is_flood_error(result):
|
||||||
|
logger.debug(
|
||||||
|
"Flood control on fallback send, retrying in 3s"
|
||||||
|
)
|
||||||
|
await asyncio.sleep(3.0)
|
||||||
|
else:
|
||||||
|
break # non-flood error or second attempt failed
|
||||||
|
|
||||||
|
if not result or not result.success:
|
||||||
|
if sent_any_chunk:
|
||||||
|
# Some continuation text already reached the user. Suppress
|
||||||
|
# the base gateway final-send path so we don't resend the
|
||||||
|
# full response and create another duplicate.
|
||||||
|
self._already_sent = True
|
||||||
|
self._final_response_sent = True
|
||||||
|
self._message_id = last_message_id
|
||||||
|
self._last_sent_text = last_successful_chunk
|
||||||
|
self._fallback_prefix = ""
|
||||||
|
return
|
||||||
|
# No fallback chunk reached the user — allow the normal gateway
|
||||||
|
# final-send path to try one more time.
|
||||||
|
self._already_sent = False
|
||||||
|
self._message_id = None
|
||||||
|
self._last_sent_text = ""
|
||||||
|
self._fallback_prefix = ""
|
||||||
|
return
|
||||||
|
sent_any_chunk = True
|
||||||
|
last_successful_chunk = chunk
|
||||||
|
last_message_id = result.message_id or last_message_id
|
||||||
|
|
||||||
|
self._message_id = last_message_id
|
||||||
|
self._already_sent = True
|
||||||
|
self._final_response_sent = True
|
||||||
|
self._last_sent_text = chunks[-1]
|
||||||
|
self._fallback_prefix = ""
|
||||||
|
|
||||||
|
def _is_flood_error(self, result) -> bool:
|
||||||
|
"""Check if a SendResult failure is due to flood control / rate limiting."""
|
||||||
|
err = getattr(result, "error", "") or ""
|
||||||
|
err_lower = err.lower()
|
||||||
|
return "flood" in err_lower or "retry after" in err_lower or "rate" in err_lower
|
||||||
|
|
||||||
|
async def _try_strip_cursor(self) -> None:
|
||||||
|
"""Best-effort edit to remove the cursor from the last visible message.
|
||||||
|
|
||||||
|
Called when entering fallback mode so the user doesn't see a stuck
|
||||||
|
cursor (▉) in the partial message.
|
||||||
|
"""
|
||||||
|
if not self._message_id or self._message_id == "__no_edit__":
|
||||||
|
return
|
||||||
|
prefix = self._visible_prefix()
|
||||||
|
if not prefix or not prefix.strip():
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
await self.adapter.edit_message(
|
||||||
|
chat_id=self.chat_id,
|
||||||
|
message_id=self._message_id,
|
||||||
|
content=prefix,
|
||||||
|
)
|
||||||
|
self._last_sent_text = prefix
|
||||||
|
except Exception:
|
||||||
|
pass # best-effort — don't let this block the fallback path
|
||||||
|
|
||||||
|
async def _send_commentary(self, text: str) -> bool:
|
||||||
|
"""Send a completed interim assistant commentary message."""
|
||||||
|
text = self._clean_for_display(text)
|
||||||
|
if not text.strip():
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
result = await self.adapter.send(
|
||||||
|
chat_id=self.chat_id,
|
||||||
|
content=text,
|
||||||
|
metadata=self.metadata,
|
||||||
|
)
|
||||||
|
if result.success:
|
||||||
|
self._already_sent = True
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Commentary send error: %s", e)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _send_or_edit(self, text: str) -> bool:
|
||||||
|
"""Send or edit the streaming message.
|
||||||
|
|
||||||
|
Returns True if the text was successfully delivered (sent or edited),
|
||||||
|
False otherwise. Callers like the overflow split loop use this to
|
||||||
|
decide whether to advance past the delivered chunk.
|
||||||
|
"""
|
||||||
|
# Strip MEDIA: directives so they don't appear as visible text.
|
||||||
|
# Media files are delivered as native attachments after the stream
|
||||||
|
# finishes (via _deliver_media_from_response in gateway/run.py).
|
||||||
|
text = self._clean_for_display(text)
|
||||||
|
# A bare streaming cursor is not meaningful user-visible content and
|
||||||
|
# can render as a stray tofu/white-box message on some clients.
|
||||||
|
visible_without_cursor = text
|
||||||
|
if self.cfg.cursor:
|
||||||
|
visible_without_cursor = visible_without_cursor.replace(self.cfg.cursor, "")
|
||||||
|
_visible_stripped = visible_without_cursor.strip()
|
||||||
|
if not _visible_stripped:
|
||||||
|
return True # cursor-only / whitespace-only update
|
||||||
|
if not text.strip():
|
||||||
|
return True # nothing to send is "success"
|
||||||
|
# Guard: do not create a brand-new standalone message when the only
|
||||||
|
# visible content is a handful of characters alongside the streaming
|
||||||
|
# cursor. During rapid tool-calling the model often emits 1-2 tokens
|
||||||
|
# before switching to tool calls; the resulting "X ▉" message risks
|
||||||
|
# leaving the cursor permanently visible if the follow-up edit (to
|
||||||
|
# strip the cursor on segment break) is rate-limited by the platform.
|
||||||
|
# This was reported on Telegram, Matrix, and other clients where the
|
||||||
|
# ▉ block character renders as a visible white box ("tofu").
|
||||||
|
# Existing messages (edits) are unaffected — only first sends gated.
|
||||||
|
_MIN_NEW_MSG_CHARS = 4
|
||||||
|
if (self._message_id is None
|
||||||
|
and self.cfg.cursor
|
||||||
|
and self.cfg.cursor in text
|
||||||
|
and len(_visible_stripped) < _MIN_NEW_MSG_CHARS):
|
||||||
|
return True # too short for a standalone message — accumulate more
|
||||||
|
try:
|
||||||
|
if self._message_id is not None:
|
||||||
|
if self._edit_supported:
|
||||||
|
# Skip if text is identical to what we last sent
|
||||||
|
if text == self._last_sent_text:
|
||||||
|
return True
|
||||||
|
# Edit existing message
|
||||||
|
result = await self.adapter.edit_message(
|
||||||
|
chat_id=self.chat_id,
|
||||||
|
message_id=self._message_id,
|
||||||
|
content=text,
|
||||||
|
)
|
||||||
|
if result.success:
|
||||||
|
self._already_sent = True
|
||||||
|
self._last_sent_text = text
|
||||||
|
# Successful edit — reset flood strike counter
|
||||||
|
self._flood_strikes = 0
|
||||||
|
return True
|
||||||
|
else:
|
||||||
|
# Edit failed. If this looks like flood control / rate
|
||||||
|
# limiting, use adaptive backoff: double the edit interval
|
||||||
|
# and retry on the next cycle. Only permanently disable
|
||||||
|
# edits after _MAX_FLOOD_STRIKES consecutive failures.
|
||||||
|
if self._is_flood_error(result):
|
||||||
|
self._flood_strikes += 1
|
||||||
|
self._current_edit_interval = min(
|
||||||
|
self._current_edit_interval * 2, 10.0,
|
||||||
|
)
|
||||||
|
logger.debug(
|
||||||
|
"Flood control on edit (strike %d/%d), "
|
||||||
|
"backoff interval → %.1fs",
|
||||||
|
self._flood_strikes,
|
||||||
|
self._MAX_FLOOD_STRIKES,
|
||||||
|
self._current_edit_interval,
|
||||||
|
)
|
||||||
|
if self._flood_strikes < self._MAX_FLOOD_STRIKES:
|
||||||
|
# Don't disable edits yet — just slow down.
|
||||||
|
# Update _last_edit_time so the next edit
|
||||||
|
# respects the new interval.
|
||||||
|
self._last_edit_time = time.monotonic()
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Non-flood error OR flood strikes exhausted: enter
|
||||||
|
# fallback mode — send only the missing tail once the
|
||||||
|
# final response is available.
|
||||||
|
logger.debug(
|
||||||
|
"Edit failed (strikes=%d), entering fallback mode",
|
||||||
|
self._flood_strikes,
|
||||||
|
)
|
||||||
|
self._fallback_prefix = self._visible_prefix()
|
||||||
|
self._fallback_final_send = True
|
||||||
|
self._edit_supported = False
|
||||||
|
self._already_sent = True
|
||||||
|
# Best-effort: strip the cursor from the last visible
|
||||||
|
# message so the user doesn't see a stuck ▉.
|
||||||
|
await self._try_strip_cursor()
|
||||||
|
return False
|
||||||
|
else:
|
||||||
|
# Editing not supported — skip intermediate updates.
|
||||||
|
# The final response will be sent by the fallback path.
|
||||||
|
return False
|
||||||
|
else:
|
||||||
|
# First message — send new
|
||||||
|
result = await self.adapter.send(
|
||||||
|
chat_id=self.chat_id,
|
||||||
|
content=text,
|
||||||
|
metadata=self.metadata,
|
||||||
|
)
|
||||||
|
if result.success:
|
||||||
|
if result.message_id:
|
||||||
|
self._message_id = result.message_id
|
||||||
|
else:
|
||||||
|
self._edit_supported = False
|
||||||
|
self._already_sent = True
|
||||||
|
self._last_sent_text = text
|
||||||
|
if not result.message_id:
|
||||||
|
self._fallback_prefix = self._visible_prefix()
|
||||||
|
self._fallback_final_send = True
|
||||||
|
# Sentinel prevents re-entering the first-send path on
|
||||||
|
# every delta/tool boundary when platforms accept a
|
||||||
|
# message but do not return an editable message id.
|
||||||
|
self._message_id = "__no_edit__"
|
||||||
|
return True
|
||||||
|
else:
|
||||||
|
# Initial send failed — disable streaming for this session
|
||||||
|
self._edit_supported = False
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Stream send/edit error: %s", e)
|
||||||
|
return False
|
||||||
+150
-38
@@ -31,7 +31,7 @@ T = TypeVar("T")
|
|||||||
|
|
||||||
DEFAULT_DB_PATH = get_hermes_home() / "state.db"
|
DEFAULT_DB_PATH = get_hermes_home() / "state.db"
|
||||||
|
|
||||||
SCHEMA_VERSION = 8
|
SCHEMA_VERSION = 9
|
||||||
|
|
||||||
SCHEMA_SQL = """
|
SCHEMA_SQL = """
|
||||||
CREATE TABLE IF NOT EXISTS schema_version (
|
CREATE TABLE IF NOT EXISTS schema_version (
|
||||||
@@ -81,7 +81,8 @@ CREATE TABLE IF NOT EXISTS messages (
|
|||||||
finish_reason TEXT,
|
finish_reason TEXT,
|
||||||
reasoning TEXT,
|
reasoning TEXT,
|
||||||
reasoning_details TEXT,
|
reasoning_details TEXT,
|
||||||
codex_reasoning_items TEXT
|
codex_reasoning_items TEXT,
|
||||||
|
archived INTEGER DEFAULT 0
|
||||||
);
|
);
|
||||||
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_sessions_source ON sessions(source);
|
CREATE INDEX IF NOT EXISTS idx_sessions_source ON sessions(source);
|
||||||
@@ -366,6 +367,13 @@ class SessionDB:
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
cursor.execute("UPDATE schema_version SET version = 8")
|
cursor.execute("UPDATE schema_version SET version = 8")
|
||||||
|
if current_version < 9:
|
||||||
|
# v9: messages 表增加 archived 列(Context Compaction 消息归档)
|
||||||
|
try:
|
||||||
|
cursor.execute("ALTER TABLE messages ADD COLUMN archived INTEGER DEFAULT 0")
|
||||||
|
except sqlite3.OperationalError:
|
||||||
|
pass # Column already exists
|
||||||
|
cursor.execute("UPDATE schema_version SET version = 9")
|
||||||
|
|
||||||
# Unique title index — always ensure it exists (safe to run after migrations
|
# Unique title index — always ensure it exists (safe to run after migrations
|
||||||
# since the title column is guaranteed to exist at this point)
|
# since the title column is guaranteed to exist at this point)
|
||||||
@@ -920,19 +928,125 @@ class SessionDB:
|
|||||||
result.append(msg)
|
result.append(msg)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
def get_messages_as_conversation(self, session_id: str) -> List[Dict[str, Any]]:
|
def archive_messages(self, session_id: str) -> int:
|
||||||
|
"""将指定 session 的所有活跃消息标记为 archived。
|
||||||
|
|
||||||
|
用于 Context Compaction:原始消息保留在 DB 中供前端历史查看,
|
||||||
|
但不再被 get_messages_as_conversation() 返回给 LLM。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
归档的消息数量。
|
||||||
"""
|
"""
|
||||||
Load messages in the OpenAI conversation format (role + content dicts).
|
def _do(conn):
|
||||||
Used by the gateway to restore conversation history.
|
cursor = conn.execute(
|
||||||
"""
|
"UPDATE messages SET archived = 1 "
|
||||||
with self._lock:
|
"WHERE session_id = ? AND archived = 0",
|
||||||
cursor = self._conn.execute(
|
|
||||||
"SELECT role, content, tool_call_id, tool_calls, tool_name, "
|
|
||||||
"reasoning, reasoning_details, codex_reasoning_items "
|
|
||||||
"FROM messages WHERE session_id = ? ORDER BY timestamp, id",
|
|
||||||
(session_id,),
|
(session_id,),
|
||||||
)
|
)
|
||||||
rows = cursor.fetchall()
|
return cursor.rowcount
|
||||||
|
return self._execute_write(_do)
|
||||||
|
|
||||||
|
def get_messages_as_conversation(self, session_id: str) -> List[Dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
Load active (non-archived) messages in the OpenAI conversation format.
|
||||||
|
Used by the gateway to restore conversation history for LLM context.
|
||||||
|
Archived messages (from Context Compaction) are excluded.
|
||||||
|
"""
|
||||||
|
with self._lock:
|
||||||
|
rows = self._conn.execute(
|
||||||
|
"SELECT role, content, tool_call_id, tool_calls, tool_name, "
|
||||||
|
"reasoning, reasoning_details, codex_reasoning_items "
|
||||||
|
"FROM messages WHERE session_id = ? AND archived = 0 ORDER BY timestamp, id",
|
||||||
|
(session_id,),
|
||||||
|
).fetchall()
|
||||||
|
return self._rows_to_conversation(rows)
|
||||||
|
|
||||||
|
# ── Compaction 标记前缀(与 context_compressor.py 保持一致) ──
|
||||||
|
_COMPACTION_MARKER = "[CONTEXT COMPACTION"
|
||||||
|
|
||||||
|
def get_all_messages_as_conversation(self, session_id: str) -> List[Dict[str, Any]]:
|
||||||
|
"""Load ALL messages (including archived) for frontend history display.
|
||||||
|
|
||||||
|
智能合并策略:
|
||||||
|
1. 如果没有归档消息 → 直接返回活跃消息(未发生过 compaction)。
|
||||||
|
2. 如果有归档消息 → 以归档消息为"基础历史",然后从活跃消息中
|
||||||
|
**去除 compaction summary + tail 重复副本**,只追加真正的新消息。
|
||||||
|
|
||||||
|
这确保前端看到完整、无重复的对话历史,同时 LLM 侧的
|
||||||
|
get_messages_as_conversation() 仍然只返回活跃(精简)上下文。
|
||||||
|
"""
|
||||||
|
with self._lock:
|
||||||
|
archived_rows = self._conn.execute(
|
||||||
|
"SELECT role, content, tool_call_id, tool_calls, tool_name, "
|
||||||
|
"reasoning, reasoning_details, codex_reasoning_items "
|
||||||
|
"FROM messages WHERE session_id = ? AND archived = 1 "
|
||||||
|
"ORDER BY timestamp, id",
|
||||||
|
(session_id,),
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
active_rows = self._conn.execute(
|
||||||
|
"SELECT role, content, tool_call_id, tool_calls, tool_name, "
|
||||||
|
"reasoning, reasoning_details, codex_reasoning_items "
|
||||||
|
"FROM messages WHERE session_id = ? AND archived = 0 "
|
||||||
|
"ORDER BY timestamp, id",
|
||||||
|
(session_id,),
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
# 无归档 → 从未压缩过,直接返回全部活跃消息
|
||||||
|
if not archived_rows:
|
||||||
|
return self._rows_to_conversation(active_rows)
|
||||||
|
|
||||||
|
archived_msgs = self._rows_to_conversation(archived_rows)
|
||||||
|
|
||||||
|
# 过滤归档中的 compaction summary(多次压缩可能产生多条)
|
||||||
|
archived_msgs = [
|
||||||
|
m for m in archived_msgs
|
||||||
|
if not (m.get("content") or "").lstrip().startswith(self._COMPACTION_MARKER)
|
||||||
|
]
|
||||||
|
|
||||||
|
# 构建归档消息的签名集(role + content 前 300 字),用于去重
|
||||||
|
# tail 消息的 timestamp 在 flush 时被重写,可能出现在 archived
|
||||||
|
# 的任意位置,因此必须对全量 archived 构建签名
|
||||||
|
all_sigs: set = set()
|
||||||
|
for msg in archived_msgs:
|
||||||
|
sig = self._msg_dedup_sig(msg)
|
||||||
|
all_sigs.add(sig)
|
||||||
|
|
||||||
|
# 从活跃消息中筛选真正的新消息
|
||||||
|
active_msgs = self._rows_to_conversation(active_rows)
|
||||||
|
new_msgs = []
|
||||||
|
for msg in active_msgs:
|
||||||
|
content = msg.get("content") or ""
|
||||||
|
# 跳过 compaction summary(LLM 内部参考,不应展示给用户)
|
||||||
|
if content.lstrip().startswith(self._COMPACTION_MARKER):
|
||||||
|
continue
|
||||||
|
# 跳过 tail 重复副本
|
||||||
|
sig = self._msg_dedup_sig(msg)
|
||||||
|
if sig in all_sigs:
|
||||||
|
continue
|
||||||
|
new_msgs.append(msg)
|
||||||
|
|
||||||
|
return archived_msgs + new_msgs
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _msg_dedup_sig(msg: dict) -> tuple:
|
||||||
|
"""生成消息去重签名。
|
||||||
|
|
||||||
|
签名组成:(role, content_prefix, tool_call_id, tc_fingerprint)
|
||||||
|
- tool_call_id: tool 角色消息的关联 ID
|
||||||
|
- tc_fingerprint: assistant 消息携带的 tool_calls 首个 call ID
|
||||||
|
(防止空 content 但不同 tool_calls 的 assistant 消息误判重复)
|
||||||
|
"""
|
||||||
|
content = (msg.get("content") or "")[:300]
|
||||||
|
tc_fp = ""
|
||||||
|
tool_calls = msg.get("tool_calls")
|
||||||
|
if tool_calls and isinstance(tool_calls, list) and len(tool_calls) > 0:
|
||||||
|
first_tc = tool_calls[0] if isinstance(tool_calls[0], dict) else {}
|
||||||
|
tc_fp = first_tc.get("id", "") or first_tc.get("call_id", "")
|
||||||
|
return (msg.get("role", ""), content, msg.get("tool_call_id") or "", tc_fp)
|
||||||
|
|
||||||
|
def _rows_to_conversation(self, rows) -> List[Dict[str, Any]]:
|
||||||
|
"""将 DB rows 转换为 OpenAI conversation 格式(共享解析逻辑)。"""
|
||||||
messages = []
|
messages = []
|
||||||
for row in rows:
|
for row in rows:
|
||||||
msg = {"role": row["role"], "content": row["content"]}
|
msg = {"role": row["role"], "content": row["content"]}
|
||||||
@@ -944,11 +1058,8 @@ class SessionDB:
|
|||||||
try:
|
try:
|
||||||
msg["tool_calls"] = json.loads(row["tool_calls"])
|
msg["tool_calls"] = json.loads(row["tool_calls"])
|
||||||
except (json.JSONDecodeError, TypeError):
|
except (json.JSONDecodeError, TypeError):
|
||||||
logger.warning("Failed to deserialize tool_calls in conversation replay, falling back to []")
|
logger.warning("Failed to deserialize tool_calls, falling back to []")
|
||||||
msg["tool_calls"] = []
|
msg["tool_calls"] = []
|
||||||
# Restore reasoning fields on assistant messages so providers
|
|
||||||
# that replay reasoning (OpenRouter, OpenAI, Nous) receive
|
|
||||||
# coherent multi-turn reasoning context.
|
|
||||||
if row["role"] == "assistant":
|
if row["role"] == "assistant":
|
||||||
if row["reasoning"]:
|
if row["reasoning"]:
|
||||||
msg["reasoning"] = row["reasoning"]
|
msg["reasoning"] = row["reasoning"]
|
||||||
@@ -956,13 +1067,11 @@ class SessionDB:
|
|||||||
try:
|
try:
|
||||||
msg["reasoning_details"] = json.loads(row["reasoning_details"])
|
msg["reasoning_details"] = json.loads(row["reasoning_details"])
|
||||||
except (json.JSONDecodeError, TypeError):
|
except (json.JSONDecodeError, TypeError):
|
||||||
logger.warning("Failed to deserialize reasoning_details, falling back to None")
|
|
||||||
msg["reasoning_details"] = None
|
msg["reasoning_details"] = None
|
||||||
if row["codex_reasoning_items"]:
|
if row["codex_reasoning_items"]:
|
||||||
try:
|
try:
|
||||||
msg["codex_reasoning_items"] = json.loads(row["codex_reasoning_items"])
|
msg["codex_reasoning_items"] = json.loads(row["codex_reasoning_items"])
|
||||||
except (json.JSONDecodeError, TypeError):
|
except (json.JSONDecodeError, TypeError):
|
||||||
logger.warning("Failed to deserialize codex_reasoning_items, falling back to None")
|
|
||||||
msg["codex_reasoning_items"] = None
|
msg["codex_reasoning_items"] = None
|
||||||
messages.append(msg)
|
messages.append(msg)
|
||||||
return messages
|
return messages
|
||||||
@@ -1162,13 +1271,24 @@ class SessionDB:
|
|||||||
cursor = self._conn.execute("SELECT COUNT(*) FROM sessions")
|
cursor = self._conn.execute("SELECT COUNT(*) FROM sessions")
|
||||||
return cursor.fetchone()[0]
|
return cursor.fetchone()[0]
|
||||||
|
|
||||||
def message_count(self, session_id: str = None) -> int:
|
def message_count(self, session_id: str = None, active_only: bool = False) -> int:
|
||||||
"""Count messages, optionally for a specific session."""
|
"""Count messages, optionally for a specific session.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session_id: If provided, count only messages for this session.
|
||||||
|
active_only: If True, exclude archived messages (from Context Compaction).
|
||||||
|
"""
|
||||||
with self._lock:
|
with self._lock:
|
||||||
if session_id:
|
if session_id:
|
||||||
cursor = self._conn.execute(
|
if active_only:
|
||||||
"SELECT COUNT(*) FROM messages WHERE session_id = ?", (session_id,)
|
cursor = self._conn.execute(
|
||||||
)
|
"SELECT COUNT(*) FROM messages WHERE session_id = ? AND archived = 0",
|
||||||
|
(session_id,),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
cursor = self._conn.execute(
|
||||||
|
"SELECT COUNT(*) FROM messages WHERE session_id = ?", (session_id,)
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
cursor = self._conn.execute("SELECT COUNT(*) FROM messages")
|
cursor = self._conn.execute("SELECT COUNT(*) FROM messages")
|
||||||
return cursor.fetchone()[0]
|
return cursor.fetchone()[0]
|
||||||
@@ -1305,21 +1425,13 @@ class SessionDB:
|
|||||||
s.started_at
|
s.started_at
|
||||||
) AS last_active,
|
) AS last_active,
|
||||||
(
|
(
|
||||||
COALESCE(
|
SELECT COUNT(*) FROM messages m3
|
||||||
(SELECT SUM(
|
WHERE m3.session_id = s.id
|
||||||
(LENGTH(m3.tool_calls) - LENGTH(REPLACE(m3.tool_calls, '"write_file"', ''))) / LENGTH('"write_file"')
|
AND (
|
||||||
) FROM messages m3
|
m3.tool_calls LIKE '%"write_file"%'
|
||||||
WHERE m3.session_id = s.id
|
OR m3.tool_calls LIKE '%"patch"%'
|
||||||
AND m3.tool_calls LIKE '%"write_file"%'),
|
OR (m3.role = 'assistant' AND m3.content LIKE '%"type"%audio"%')
|
||||||
0)
|
)
|
||||||
+
|
|
||||||
(SELECT COUNT(*) FROM messages m4
|
|
||||||
WHERE m4.session_id = s.id
|
|
||||||
AND m4.tool_calls LIKE '%"patch"%')
|
|
||||||
+
|
|
||||||
(SELECT COUNT(*) FROM messages m5
|
|
||||||
WHERE m5.session_id = s.id
|
|
||||||
AND m5.role = 'assistant' AND m5.content LIKE '%"type"%audio"%')
|
|
||||||
) AS work_product_count
|
) AS work_product_count
|
||||||
FROM sessions s
|
FROM sessions s
|
||||||
WHERE s.user_id = ?
|
WHERE s.user_id = ?
|
||||||
|
|||||||
@@ -0,0 +1,104 @@
|
|||||||
|
"""
|
||||||
|
Timezone-aware clock for Hermes.
|
||||||
|
|
||||||
|
Provides a single ``now()`` helper that returns a timezone-aware datetime
|
||||||
|
based on the user's configured IANA timezone (e.g. ``Asia/Kolkata``).
|
||||||
|
|
||||||
|
Resolution order:
|
||||||
|
1. ``HERMES_TIMEZONE`` environment variable
|
||||||
|
2. ``timezone`` key in ``~/.hermes/config.yaml``
|
||||||
|
3. Falls back to the server's local time (``datetime.now().astimezone()``)
|
||||||
|
|
||||||
|
Invalid timezone values log a warning and fall back safely — Hermes never
|
||||||
|
crashes due to a bad timezone string.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from datetime import datetime
|
||||||
|
from hermes_constants import get_config_path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
try:
|
||||||
|
from zoneinfo import ZoneInfo
|
||||||
|
except ImportError:
|
||||||
|
# Python 3.8 fallback (shouldn't be needed — Hermes requires 3.9+)
|
||||||
|
from backports.zoneinfo import ZoneInfo # type: ignore[no-redef]
|
||||||
|
|
||||||
|
# Cached state — resolved once, reused on every call.
|
||||||
|
# Call reset_cache() to force re-resolution (e.g. after config changes).
|
||||||
|
_cached_tz: Optional[ZoneInfo] = None
|
||||||
|
_cached_tz_name: Optional[str] = None
|
||||||
|
_cache_resolved: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_timezone_name() -> str:
|
||||||
|
"""Read the configured IANA timezone string (or empty string).
|
||||||
|
|
||||||
|
This does file I/O when falling through to config.yaml, so callers
|
||||||
|
should cache the result rather than calling on every ``now()``.
|
||||||
|
"""
|
||||||
|
# 1. Environment variable (highest priority — set by Supervisor, etc.)
|
||||||
|
tz_env = os.getenv("HERMES_TIMEZONE", "").strip()
|
||||||
|
if tz_env:
|
||||||
|
return tz_env
|
||||||
|
|
||||||
|
# 2. config.yaml ``timezone`` key
|
||||||
|
try:
|
||||||
|
import yaml
|
||||||
|
config_path = get_config_path()
|
||||||
|
if config_path.exists():
|
||||||
|
with open(config_path) as f:
|
||||||
|
cfg = yaml.safe_load(f) or {}
|
||||||
|
tz_cfg = cfg.get("timezone", "")
|
||||||
|
if isinstance(tz_cfg, str) and tz_cfg.strip():
|
||||||
|
return tz_cfg.strip()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _get_zoneinfo(name: str) -> Optional[ZoneInfo]:
|
||||||
|
"""Validate and return a ZoneInfo, or None if invalid."""
|
||||||
|
if not name:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
return ZoneInfo(name)
|
||||||
|
except (KeyError, Exception) as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Invalid timezone '%s': %s. Falling back to server local time.",
|
||||||
|
name, exc,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def get_timezone() -> Optional[ZoneInfo]:
|
||||||
|
"""Return the user's configured ZoneInfo, or None (meaning server-local).
|
||||||
|
|
||||||
|
Resolved once and cached. Call ``reset_cache()`` after config changes.
|
||||||
|
"""
|
||||||
|
global _cached_tz, _cached_tz_name, _cache_resolved
|
||||||
|
if not _cache_resolved:
|
||||||
|
_cached_tz_name = _resolve_timezone_name()
|
||||||
|
_cached_tz = _get_zoneinfo(_cached_tz_name)
|
||||||
|
_cache_resolved = True
|
||||||
|
return _cached_tz
|
||||||
|
|
||||||
|
|
||||||
|
def now() -> datetime:
|
||||||
|
"""
|
||||||
|
Return the current time as a timezone-aware datetime.
|
||||||
|
|
||||||
|
If a valid timezone is configured, returns wall-clock time in that zone.
|
||||||
|
Otherwise returns the server's local time (via ``astimezone()``).
|
||||||
|
"""
|
||||||
|
tz = get_timezone()
|
||||||
|
if tz is not None:
|
||||||
|
return datetime.now(tz)
|
||||||
|
# No timezone configured — use server-local (still tz-aware)
|
||||||
|
return datetime.now().astimezone()
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,601 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""
|
||||||
|
Model Tools Module
|
||||||
|
|
||||||
|
Thin orchestration layer over the tool registry. Each tool file in tools/
|
||||||
|
self-registers its schema, handler, and metadata via tools.registry.register().
|
||||||
|
This module triggers discovery (by importing all tool modules), then provides
|
||||||
|
the public API that run_agent.py, cli.py, batch_runner.py, and the RL
|
||||||
|
environments consume.
|
||||||
|
|
||||||
|
Public API (signatures preserved from the original 2,400-line version):
|
||||||
|
get_tool_definitions(enabled_toolsets, disabled_toolsets, quiet_mode) -> list
|
||||||
|
handle_function_call(function_name, function_args, task_id, user_task) -> str
|
||||||
|
TOOL_TO_TOOLSET_MAP: dict (for batch_runner.py)
|
||||||
|
TOOLSET_REQUIREMENTS: dict (for cli.py, doctor.py)
|
||||||
|
get_all_tool_names() -> list
|
||||||
|
get_toolset_for_tool(name) -> str
|
||||||
|
get_available_toolsets() -> dict
|
||||||
|
check_toolset_requirements() -> dict
|
||||||
|
check_tool_availability(quiet) -> tuple
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import threading
|
||||||
|
from typing import Dict, Any, List, Optional, Tuple
|
||||||
|
|
||||||
|
from tools.registry import registry
|
||||||
|
from toolsets import resolve_toolset, validate_toolset
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Async Bridging (single source of truth -- used by registry.dispatch too)
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
_tool_loop = None # persistent loop for the main (CLI) thread
|
||||||
|
_tool_loop_lock = threading.Lock()
|
||||||
|
_worker_thread_local = threading.local() # per-worker-thread persistent loops
|
||||||
|
|
||||||
|
|
||||||
|
def _get_tool_loop():
|
||||||
|
"""Return a long-lived event loop for running async tool handlers.
|
||||||
|
|
||||||
|
Using a persistent loop (instead of asyncio.run() which creates and
|
||||||
|
*closes* a fresh loop every time) prevents "Event loop is closed"
|
||||||
|
errors that occur when cached httpx/AsyncOpenAI clients attempt to
|
||||||
|
close their transport on a dead loop during garbage collection.
|
||||||
|
"""
|
||||||
|
global _tool_loop
|
||||||
|
with _tool_loop_lock:
|
||||||
|
if _tool_loop is None or _tool_loop.is_closed():
|
||||||
|
_tool_loop = asyncio.new_event_loop()
|
||||||
|
return _tool_loop
|
||||||
|
|
||||||
|
|
||||||
|
def _get_worker_loop():
|
||||||
|
"""Return a persistent event loop for the current worker thread.
|
||||||
|
|
||||||
|
Each worker thread (e.g., delegate_task's ThreadPoolExecutor threads)
|
||||||
|
gets its own long-lived loop stored in thread-local storage. This
|
||||||
|
prevents the "Event loop is closed" errors that occurred when
|
||||||
|
asyncio.run() was used per-call: asyncio.run() creates a loop, runs
|
||||||
|
the coroutine, then *closes* the loop — but cached httpx/AsyncOpenAI
|
||||||
|
clients remain bound to that now-dead loop and raise RuntimeError
|
||||||
|
during garbage collection or subsequent use.
|
||||||
|
|
||||||
|
By keeping the loop alive for the thread's lifetime, cached clients
|
||||||
|
stay valid and their cleanup runs on a live loop.
|
||||||
|
"""
|
||||||
|
loop = getattr(_worker_thread_local, 'loop', None)
|
||||||
|
if loop is None or loop.is_closed():
|
||||||
|
loop = asyncio.new_event_loop()
|
||||||
|
asyncio.set_event_loop(loop)
|
||||||
|
_worker_thread_local.loop = loop
|
||||||
|
return loop
|
||||||
|
|
||||||
|
|
||||||
|
def _run_async(coro):
|
||||||
|
"""Run an async coroutine from a sync context.
|
||||||
|
|
||||||
|
If the current thread already has a running event loop (e.g., inside
|
||||||
|
the gateway's async stack or Atropos's event loop), we spin up a
|
||||||
|
disposable thread so asyncio.run() can create its own loop without
|
||||||
|
conflicting.
|
||||||
|
|
||||||
|
For the common CLI path (no running loop), we use a persistent event
|
||||||
|
loop so that cached async clients (httpx / AsyncOpenAI) remain bound
|
||||||
|
to a live loop and don't trigger "Event loop is closed" on GC.
|
||||||
|
|
||||||
|
When called from a worker thread (parallel tool execution), we use a
|
||||||
|
per-thread persistent loop to avoid both contention with the main
|
||||||
|
thread's shared loop AND the "Event loop is closed" errors caused by
|
||||||
|
asyncio.run()'s create-and-destroy lifecycle.
|
||||||
|
|
||||||
|
This is the single source of truth for sync->async bridging in tool
|
||||||
|
handlers. The RL paths (agent_loop.py, tool_context.py) also provide
|
||||||
|
outer thread-pool wrapping as defense-in-depth, but each handler is
|
||||||
|
self-protecting via this function.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
except RuntimeError:
|
||||||
|
loop = None
|
||||||
|
|
||||||
|
if loop and loop.is_running():
|
||||||
|
# Inside an async context (gateway, RL env) — run in a fresh thread.
|
||||||
|
import concurrent.futures
|
||||||
|
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||||
|
future = pool.submit(asyncio.run, coro)
|
||||||
|
return future.result(timeout=300)
|
||||||
|
|
||||||
|
# If we're on a worker thread (e.g., parallel tool execution in
|
||||||
|
# delegate_task), use a per-thread persistent loop. This avoids
|
||||||
|
# contention with the main thread's shared loop while keeping cached
|
||||||
|
# httpx/AsyncOpenAI clients bound to a live loop for the thread's
|
||||||
|
# lifetime — preventing "Event loop is closed" on GC cleanup.
|
||||||
|
if threading.current_thread() is not threading.main_thread():
|
||||||
|
worker_loop = _get_worker_loop()
|
||||||
|
return worker_loop.run_until_complete(coro)
|
||||||
|
|
||||||
|
tool_loop = _get_tool_loop()
|
||||||
|
return tool_loop.run_until_complete(coro)
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Tool Discovery (importing each module triggers its registry.register calls)
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
def _discover_tools():
|
||||||
|
"""Import all tool modules to trigger their registry.register() calls.
|
||||||
|
|
||||||
|
Wrapped in a function so import errors in optional tools (e.g., fal_client
|
||||||
|
not installed) don't prevent the rest from loading.
|
||||||
|
"""
|
||||||
|
_modules = [
|
||||||
|
"tools.web_tools",
|
||||||
|
"tools.terminal_tool",
|
||||||
|
"tools.file_tools",
|
||||||
|
"tools.vision_tools",
|
||||||
|
"tools.mixture_of_agents_tool",
|
||||||
|
"tools.image_generation_tool",
|
||||||
|
"tools.skills_tool",
|
||||||
|
"tools.skill_manager_tool",
|
||||||
|
"tools.browser_tool",
|
||||||
|
"tools.cronjob_tools",
|
||||||
|
"tools.rl_training_tool",
|
||||||
|
"tools.tts_tool",
|
||||||
|
"tools.todo_tool",
|
||||||
|
"tools.memory_tool",
|
||||||
|
"tools.session_search_tool",
|
||||||
|
"tools.clarify_tool",
|
||||||
|
"tools.code_execution_tool",
|
||||||
|
"tools.delegate_tool",
|
||||||
|
"tools.process_registry",
|
||||||
|
"tools.send_message_tool",
|
||||||
|
# "tools.honcho_tools", # Removed — Honcho is now a memory provider plugin
|
||||||
|
"tools.homeassistant_tool",
|
||||||
|
"tools.cli_tunnel_tool",
|
||||||
|
]
|
||||||
|
import importlib
|
||||||
|
for mod_name in _modules:
|
||||||
|
try:
|
||||||
|
importlib.import_module(mod_name)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Could not import tool module %s: %s", mod_name, e)
|
||||||
|
|
||||||
|
|
||||||
|
_discover_tools()
|
||||||
|
|
||||||
|
# MCP tool discovery (external MCP servers from config)
|
||||||
|
try:
|
||||||
|
from tools.mcp_tool import discover_mcp_tools
|
||||||
|
discover_mcp_tools()
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("MCP tool discovery failed: %s", e)
|
||||||
|
|
||||||
|
# Plugin tool discovery (user/project/pip plugins)
|
||||||
|
try:
|
||||||
|
from hermes_cli.plugins import discover_plugins
|
||||||
|
discover_plugins()
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Plugin discovery failed: %s", e)
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Backward-compat constants (built once after discovery)
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
TOOL_TO_TOOLSET_MAP: Dict[str, str] = registry.get_tool_to_toolset_map()
|
||||||
|
|
||||||
|
TOOLSET_REQUIREMENTS: Dict[str, dict] = registry.get_toolset_requirements()
|
||||||
|
|
||||||
|
# Resolved tool names from the last get_tool_definitions() call.
|
||||||
|
# Used by code_execution_tool to know which tools are available in this session.
|
||||||
|
_last_resolved_tool_names: List[str] = []
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Legacy toolset name mapping (old _tools-suffixed names -> tool name lists)
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
_LEGACY_TOOLSET_MAP = {
|
||||||
|
"web_tools": ["web_search", "web_extract"],
|
||||||
|
"terminal_tools": ["terminal"],
|
||||||
|
"vision_tools": ["vision_analyze"],
|
||||||
|
"moa_tools": ["mixture_of_agents"],
|
||||||
|
"image_tools": ["image_generate"],
|
||||||
|
"skills_tools": ["skills_list", "skill_view", "skill_manage"],
|
||||||
|
"browser_tools": [
|
||||||
|
"browser_navigate", "browser_snapshot", "browser_click",
|
||||||
|
"browser_type", "browser_scroll", "browser_back",
|
||||||
|
"browser_press", "browser_get_images",
|
||||||
|
"browser_vision", "browser_console"
|
||||||
|
],
|
||||||
|
"cronjob_tools": ["cronjob"],
|
||||||
|
"rl_tools": [
|
||||||
|
"rl_list_environments", "rl_select_environment",
|
||||||
|
"rl_get_current_config", "rl_edit_config",
|
||||||
|
"rl_start_training", "rl_check_status",
|
||||||
|
"rl_stop_training", "rl_get_results",
|
||||||
|
"rl_list_runs", "rl_test_inference"
|
||||||
|
],
|
||||||
|
"file_tools": ["read_file", "write_file", "patch", "search_files"],
|
||||||
|
"tts_tools": ["text_to_speech"],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# get_tool_definitions (the main schema provider)
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
def get_tool_definitions(
|
||||||
|
enabled_toolsets: List[str] = None,
|
||||||
|
disabled_toolsets: List[str] = None,
|
||||||
|
quiet_mode: bool = False,
|
||||||
|
) -> List[Dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
Get tool definitions for model API calls with toolset-based filtering.
|
||||||
|
|
||||||
|
All tools must be part of a toolset to be accessible.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
enabled_toolsets: Only include tools from these toolsets.
|
||||||
|
disabled_toolsets: Exclude tools from these toolsets (if enabled_toolsets is None).
|
||||||
|
quiet_mode: Suppress status prints.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Filtered list of OpenAI-format tool definitions.
|
||||||
|
"""
|
||||||
|
# Determine which tool names the caller wants
|
||||||
|
tools_to_include: set = set()
|
||||||
|
|
||||||
|
if enabled_toolsets is not None:
|
||||||
|
for toolset_name in enabled_toolsets:
|
||||||
|
if validate_toolset(toolset_name):
|
||||||
|
resolved = resolve_toolset(toolset_name)
|
||||||
|
tools_to_include.update(resolved)
|
||||||
|
if not quiet_mode:
|
||||||
|
print(f"✅ Enabled toolset '{toolset_name}': {', '.join(resolved) if resolved else 'no tools'}")
|
||||||
|
elif toolset_name in _LEGACY_TOOLSET_MAP:
|
||||||
|
legacy_tools = _LEGACY_TOOLSET_MAP[toolset_name]
|
||||||
|
tools_to_include.update(legacy_tools)
|
||||||
|
if not quiet_mode:
|
||||||
|
print(f"✅ Enabled legacy toolset '{toolset_name}': {', '.join(legacy_tools)}")
|
||||||
|
else:
|
||||||
|
if not quiet_mode:
|
||||||
|
print(f"⚠️ Unknown toolset: {toolset_name}")
|
||||||
|
|
||||||
|
elif disabled_toolsets:
|
||||||
|
from toolsets import get_all_toolsets
|
||||||
|
for ts_name in get_all_toolsets():
|
||||||
|
tools_to_include.update(resolve_toolset(ts_name))
|
||||||
|
|
||||||
|
for toolset_name in disabled_toolsets:
|
||||||
|
if validate_toolset(toolset_name):
|
||||||
|
resolved = resolve_toolset(toolset_name)
|
||||||
|
tools_to_include.difference_update(resolved)
|
||||||
|
if not quiet_mode:
|
||||||
|
print(f"🚫 Disabled toolset '{toolset_name}': {', '.join(resolved) if resolved else 'no tools'}")
|
||||||
|
elif toolset_name in _LEGACY_TOOLSET_MAP:
|
||||||
|
legacy_tools = _LEGACY_TOOLSET_MAP[toolset_name]
|
||||||
|
tools_to_include.difference_update(legacy_tools)
|
||||||
|
if not quiet_mode:
|
||||||
|
print(f"🚫 Disabled legacy toolset '{toolset_name}': {', '.join(legacy_tools)}")
|
||||||
|
else:
|
||||||
|
if not quiet_mode:
|
||||||
|
print(f"⚠️ Unknown toolset: {toolset_name}")
|
||||||
|
else:
|
||||||
|
from toolsets import get_all_toolsets
|
||||||
|
for ts_name in get_all_toolsets():
|
||||||
|
tools_to_include.update(resolve_toolset(ts_name))
|
||||||
|
|
||||||
|
# Plugin-registered tools are now resolved through the normal toolset
|
||||||
|
# path — validate_toolset() / resolve_toolset() / get_all_toolsets()
|
||||||
|
# all check the tool registry for plugin-provided toolsets. No bypass
|
||||||
|
# needed; plugins respect enabled_toolsets / disabled_toolsets like any
|
||||||
|
# other toolset.
|
||||||
|
|
||||||
|
# Ask the registry for schemas (only returns tools whose check_fn passes)
|
||||||
|
filtered_tools = registry.get_definitions(tools_to_include, quiet=quiet_mode)
|
||||||
|
|
||||||
|
# The set of tool names that actually passed check_fn filtering.
|
||||||
|
# Use this (not tools_to_include) for any downstream schema that references
|
||||||
|
# other tools by name — otherwise the model sees tools mentioned in
|
||||||
|
# descriptions that don't actually exist, and hallucinates calls to them.
|
||||||
|
available_tool_names = {t["function"]["name"] for t in filtered_tools}
|
||||||
|
|
||||||
|
# Rebuild execute_code schema to only list sandbox tools that are actually
|
||||||
|
# available. Without this, the model sees "web_search is available in
|
||||||
|
# execute_code" even when the API key isn't configured or the toolset is
|
||||||
|
# disabled (#560-discord).
|
||||||
|
if "execute_code" in available_tool_names:
|
||||||
|
from tools.code_execution_tool import SANDBOX_ALLOWED_TOOLS, build_execute_code_schema
|
||||||
|
sandbox_enabled = SANDBOX_ALLOWED_TOOLS & available_tool_names
|
||||||
|
dynamic_schema = build_execute_code_schema(sandbox_enabled)
|
||||||
|
for i, td in enumerate(filtered_tools):
|
||||||
|
if td.get("function", {}).get("name") == "execute_code":
|
||||||
|
filtered_tools[i] = {"type": "function", "function": dynamic_schema}
|
||||||
|
break
|
||||||
|
|
||||||
|
# Strip web tool cross-references from browser_navigate description when
|
||||||
|
# web_search / web_extract are not available. The static schema says
|
||||||
|
# "prefer web_search or web_extract" which causes the model to hallucinate
|
||||||
|
# those tools when they're missing.
|
||||||
|
if "browser_navigate" in available_tool_names:
|
||||||
|
web_tools_available = {"web_search", "web_extract"} & available_tool_names
|
||||||
|
if not web_tools_available:
|
||||||
|
for i, td in enumerate(filtered_tools):
|
||||||
|
if td.get("function", {}).get("name") == "browser_navigate":
|
||||||
|
desc = td["function"].get("description", "")
|
||||||
|
desc = desc.replace(
|
||||||
|
" For simple information retrieval, prefer web_search or web_extract (faster, cheaper).",
|
||||||
|
"",
|
||||||
|
)
|
||||||
|
filtered_tools[i] = {
|
||||||
|
"type": "function",
|
||||||
|
"function": {**td["function"], "description": desc},
|
||||||
|
}
|
||||||
|
break
|
||||||
|
|
||||||
|
if not quiet_mode:
|
||||||
|
if filtered_tools:
|
||||||
|
tool_names = [t["function"]["name"] for t in filtered_tools]
|
||||||
|
print(f"🛠️ Final tool selection ({len(filtered_tools)} tools): {', '.join(tool_names)}")
|
||||||
|
else:
|
||||||
|
print("🛠️ No tools selected (all filtered out or unavailable)")
|
||||||
|
|
||||||
|
global _last_resolved_tool_names
|
||||||
|
_last_resolved_tool_names = [t["function"]["name"] for t in filtered_tools]
|
||||||
|
|
||||||
|
return filtered_tools
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# handle_function_call (the main dispatcher)
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
# Tools whose execution is intercepted by the agent loop (run_agent.py)
|
||||||
|
# because they need agent-level state (TodoStore, MemoryStore, etc.).
|
||||||
|
# The registry still holds their schemas; dispatch just returns a stub error
|
||||||
|
# so if something slips through, the LLM sees a sensible message.
|
||||||
|
_AGENT_LOOP_TOOLS = {"todo", "memory", "session_search", "delegate_task"}
|
||||||
|
_READ_SEARCH_TOOLS = {"read_file", "search_files"}
|
||||||
|
|
||||||
|
|
||||||
|
# =========================================================================
|
||||||
|
# Tool argument type coercion
|
||||||
|
# =========================================================================
|
||||||
|
|
||||||
|
def coerce_tool_args(tool_name: str, args: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Coerce tool call arguments to match their JSON Schema types.
|
||||||
|
|
||||||
|
LLMs frequently return numbers as strings (``"42"`` instead of ``42``)
|
||||||
|
and booleans as strings (``"true"`` instead of ``true``). This compares
|
||||||
|
each argument value against the tool's registered JSON Schema and attempts
|
||||||
|
safe coercion when the value is a string but the schema expects a different
|
||||||
|
type. Original values are preserved when coercion fails.
|
||||||
|
|
||||||
|
Handles ``"type": "integer"``, ``"type": "number"``, ``"type": "boolean"``,
|
||||||
|
and union types (``"type": ["integer", "string"]``).
|
||||||
|
"""
|
||||||
|
if not args or not isinstance(args, dict):
|
||||||
|
return args
|
||||||
|
|
||||||
|
schema = registry.get_schema(tool_name)
|
||||||
|
if not schema:
|
||||||
|
return args
|
||||||
|
|
||||||
|
properties = (schema.get("parameters") or {}).get("properties")
|
||||||
|
if not properties:
|
||||||
|
return args
|
||||||
|
|
||||||
|
for key, value in args.items():
|
||||||
|
if not isinstance(value, str):
|
||||||
|
continue
|
||||||
|
prop_schema = properties.get(key)
|
||||||
|
if not prop_schema:
|
||||||
|
continue
|
||||||
|
expected = prop_schema.get("type")
|
||||||
|
if not expected:
|
||||||
|
continue
|
||||||
|
coerced = _coerce_value(value, expected)
|
||||||
|
if coerced is not value:
|
||||||
|
args[key] = coerced
|
||||||
|
|
||||||
|
return args
|
||||||
|
|
||||||
|
|
||||||
|
def _coerce_value(value: str, expected_type):
|
||||||
|
"""Attempt to coerce a string *value* to *expected_type*.
|
||||||
|
|
||||||
|
Returns the original string when coercion is not applicable or fails.
|
||||||
|
"""
|
||||||
|
if isinstance(expected_type, list):
|
||||||
|
# Union type — try each in order, return first successful coercion
|
||||||
|
for t in expected_type:
|
||||||
|
result = _coerce_value(value, t)
|
||||||
|
if result is not value:
|
||||||
|
return result
|
||||||
|
return value
|
||||||
|
|
||||||
|
if expected_type in ("integer", "number"):
|
||||||
|
return _coerce_number(value, integer_only=(expected_type == "integer"))
|
||||||
|
if expected_type == "boolean":
|
||||||
|
return _coerce_boolean(value)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _coerce_number(value: str, integer_only: bool = False):
|
||||||
|
"""Try to parse *value* as a number. Returns original string on failure."""
|
||||||
|
try:
|
||||||
|
f = float(value)
|
||||||
|
except (ValueError, OverflowError):
|
||||||
|
return value
|
||||||
|
# Guard against inf/nan before int() conversion
|
||||||
|
if f != f or f == float("inf") or f == float("-inf"):
|
||||||
|
return f
|
||||||
|
# If it looks like an integer (no fractional part), return int
|
||||||
|
if f == int(f):
|
||||||
|
return int(f)
|
||||||
|
if integer_only:
|
||||||
|
# Schema wants an integer but value has decimals — keep as string
|
||||||
|
return value
|
||||||
|
return f
|
||||||
|
|
||||||
|
|
||||||
|
def _coerce_boolean(value: str):
|
||||||
|
"""Try to parse *value* as a boolean. Returns original string on failure."""
|
||||||
|
low = value.strip().lower()
|
||||||
|
if low == "true":
|
||||||
|
return True
|
||||||
|
if low == "false":
|
||||||
|
return False
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def handle_function_call(
|
||||||
|
function_name: str,
|
||||||
|
function_args: Dict[str, Any],
|
||||||
|
task_id: Optional[str] = None,
|
||||||
|
tool_call_id: Optional[str] = None,
|
||||||
|
session_id: Optional[str] = None,
|
||||||
|
user_task: Optional[str] = None,
|
||||||
|
enabled_tools: Optional[List[str]] = None,
|
||||||
|
skip_pre_tool_call_hook: bool = False,
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
Main function call dispatcher that routes calls to the tool registry.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
function_name: Name of the function to call.
|
||||||
|
function_args: Arguments for the function.
|
||||||
|
task_id: Unique identifier for terminal/browser session isolation.
|
||||||
|
user_task: The user's original task (for browser_snapshot context).
|
||||||
|
enabled_tools: Tool names enabled for this session. When provided,
|
||||||
|
execute_code uses this list to determine which sandbox
|
||||||
|
tools to generate. Falls back to the process-global
|
||||||
|
``_last_resolved_tool_names`` for backward compat.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Function result as a JSON string.
|
||||||
|
"""
|
||||||
|
# Coerce string arguments to their schema-declared types (e.g. "42"→42)
|
||||||
|
function_args = coerce_tool_args(function_name, function_args)
|
||||||
|
|
||||||
|
try:
|
||||||
|
if function_name in _AGENT_LOOP_TOOLS:
|
||||||
|
return json.dumps({"error": f"{function_name} must be handled by the agent loop"})
|
||||||
|
|
||||||
|
# Check plugin hooks for a block directive (unless caller already
|
||||||
|
# checked — e.g. run_agent._invoke_tool passes skip=True to
|
||||||
|
# avoid double-firing the hook).
|
||||||
|
if not skip_pre_tool_call_hook:
|
||||||
|
block_message: Optional[str] = None
|
||||||
|
try:
|
||||||
|
from hermes_cli.plugins import get_pre_tool_call_block_message
|
||||||
|
block_message = get_pre_tool_call_block_message(
|
||||||
|
function_name,
|
||||||
|
function_args,
|
||||||
|
task_id=task_id or "",
|
||||||
|
session_id=session_id or "",
|
||||||
|
tool_call_id=tool_call_id or "",
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if block_message is not None:
|
||||||
|
return json.dumps({"error": block_message}, ensure_ascii=False)
|
||||||
|
else:
|
||||||
|
# Still fire the hook for observers — just don't check for blocking
|
||||||
|
# (the caller already did that).
|
||||||
|
try:
|
||||||
|
from hermes_cli.plugins import invoke_hook
|
||||||
|
invoke_hook(
|
||||||
|
"pre_tool_call",
|
||||||
|
tool_name=function_name,
|
||||||
|
args=function_args,
|
||||||
|
task_id=task_id or "",
|
||||||
|
session_id=session_id or "",
|
||||||
|
tool_call_id=tool_call_id or "",
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Notify the read-loop tracker when a non-read/search tool runs,
|
||||||
|
# so the *consecutive* counter resets (reads after other work are fine).
|
||||||
|
if function_name not in _READ_SEARCH_TOOLS:
|
||||||
|
try:
|
||||||
|
from tools.file_tools import notify_other_tool_call
|
||||||
|
notify_other_tool_call(task_id or "default")
|
||||||
|
except Exception:
|
||||||
|
pass # file_tools may not be loaded yet
|
||||||
|
|
||||||
|
if function_name == "execute_code":
|
||||||
|
# Prefer the caller-provided list so subagents can't overwrite
|
||||||
|
# the parent's tool set via the process-global.
|
||||||
|
sandbox_enabled = enabled_tools if enabled_tools is not None else _last_resolved_tool_names
|
||||||
|
result = registry.dispatch(
|
||||||
|
function_name, function_args,
|
||||||
|
task_id=task_id,
|
||||||
|
enabled_tools=sandbox_enabled,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
result = registry.dispatch(
|
||||||
|
function_name, function_args,
|
||||||
|
task_id=task_id,
|
||||||
|
user_task=user_task,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
from hermes_cli.plugins import invoke_hook
|
||||||
|
invoke_hook(
|
||||||
|
"post_tool_call",
|
||||||
|
tool_name=function_name,
|
||||||
|
args=function_args,
|
||||||
|
result=result,
|
||||||
|
task_id=task_id or "",
|
||||||
|
session_id=session_id or "",
|
||||||
|
tool_call_id=tool_call_id or "",
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
error_msg = f"Error executing {function_name}: {str(e)}"
|
||||||
|
logger.error(error_msg)
|
||||||
|
return json.dumps({"error": error_msg}, ensure_ascii=False)
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Backward-compat wrapper functions
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
def get_all_tool_names() -> List[str]:
|
||||||
|
"""Return all registered tool names."""
|
||||||
|
return registry.get_all_tool_names()
|
||||||
|
|
||||||
|
|
||||||
|
def get_toolset_for_tool(tool_name: str) -> Optional[str]:
|
||||||
|
"""Return the toolset a tool belongs to."""
|
||||||
|
return registry.get_toolset_for_tool(tool_name)
|
||||||
|
|
||||||
|
|
||||||
|
def get_available_toolsets() -> Dict[str, dict]:
|
||||||
|
"""Return toolset availability info for UI display."""
|
||||||
|
return registry.get_available_toolsets()
|
||||||
|
|
||||||
|
|
||||||
|
def check_toolset_requirements() -> Dict[str, bool]:
|
||||||
|
"""Return {toolset: available_bool} for every registered toolset."""
|
||||||
|
return registry.check_toolset_requirements()
|
||||||
|
|
||||||
|
|
||||||
|
def check_tool_availability(quiet: bool = False) -> Tuple[List[str], List[dict]]:
|
||||||
|
"""Return (available_toolsets, unavailable_info)."""
|
||||||
|
return registry.check_tool_availability(quiet=quiet)
|
||||||
@@ -6804,17 +6804,17 @@ class AIAgent:
|
|||||||
# 这样前端的 chatId 始终有效,history API 能找到压缩后的消息。
|
# 这样前端的 chatId 始终有效,history API 能找到压缩后的消息。
|
||||||
old_title = self._session_db.get_session_title(self.session_id)
|
old_title = self._session_db.get_session_title(self.session_id)
|
||||||
|
|
||||||
# 清除该 session 的所有旧消息(压缩摘要已在内存 messages[] 中)
|
# 归档旧消息(保留完整历史,但 LLM 不再看到它们)
|
||||||
try:
|
try:
|
||||||
self._session_db._conn.execute(
|
archived_count = self._session_db.archive_messages(self.session_id)
|
||||||
"DELETE FROM messages WHERE session_id = ?",
|
logger.info(
|
||||||
(self.session_id,),
|
"[compress] Archived %d messages for session %s",
|
||||||
|
archived_count, self.session_id,
|
||||||
)
|
)
|
||||||
self._session_db._conn.commit()
|
except Exception as e:
|
||||||
except Exception:
|
logger.warning("Message archival failed (non-fatal): %s", e)
|
||||||
pass # non-fatal: worst case 有重复旧消息
|
|
||||||
|
|
||||||
# 重置 flush cursor — 压缩后的消息从头写入
|
# 重置 flush cursor — 压缩后的消息从头写入(archived 消息不会被重复读取)
|
||||||
self._last_flushed_db_idx = 0
|
self._last_flushed_db_idx = 0
|
||||||
|
|
||||||
# 更新日志文件路径(仅 JSON 日志,不影响 DB)
|
# 更新日志文件路径(仅 JSON 日志,不影响 DB)
|
||||||
|
|||||||
@@ -0,0 +1,324 @@
|
|||||||
|
"""
|
||||||
|
CLI Tunnel 工具 — 动态注册 Mind CLI 本地工具到 Hermes 编排器。
|
||||||
|
|
||||||
|
设计遵循 SPEC 铁律 B(能力报告义务):
|
||||||
|
CLI 连接时上报能力 → Cloud 审批 → 动态注册到 ToolRegistry
|
||||||
|
CLI 断开 → 从 ToolRegistry 注销
|
||||||
|
|
||||||
|
注册模式参照 MCP 动态发现(mcp_tool._register_server_tools),
|
||||||
|
用户隔离参照飞书连接器(feishu_tool.handler(args.user_id))。
|
||||||
|
|
||||||
|
与内置工具通过 `cli_` 前缀区分:
|
||||||
|
cli_terminal = 在用户本地电脑执行命令
|
||||||
|
terminal = 在云端服务器执行命令
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from tools.registry import registry, tool_error, tool_result
|
||||||
|
|
||||||
|
logger = logging.getLogger("tools.cli_tunnel")
|
||||||
|
|
||||||
|
# ── 模块级 bridge 引用(由 mindos_sse.py 注入) ────────────
|
||||||
|
_bridge = None
|
||||||
|
|
||||||
|
|
||||||
|
def set_bridge(bridge) -> None:
|
||||||
|
"""注入 MindCLIBridge 实例(进程生命周期内调用一次)。"""
|
||||||
|
global _bridge
|
||||||
|
_bridge = bridge
|
||||||
|
|
||||||
|
|
||||||
|
# ── 工具 Schema 模板 ─────────────────────────────────────
|
||||||
|
# 每个 CLI 内置工具对应一个 schema 模板,用于注册到 LLM function calling。
|
||||||
|
# 只有 capability_report 中出现的工具才会被注册。
|
||||||
|
|
||||||
|
_TOOL_SCHEMAS: dict[str, dict] = {
|
||||||
|
"terminal": {
|
||||||
|
"name": "cli_terminal",
|
||||||
|
"description": (
|
||||||
|
"在用户的本地电脑上执行终端命令。"
|
||||||
|
"用于需要访问用户本地文件系统、开发环境或系统工具的场景。"
|
||||||
|
"注意:这是在用户个人电脑上执行,不是在服务器上。"
|
||||||
|
),
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"command": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "要执行的 shell 命令",
|
||||||
|
},
|
||||||
|
"cwd": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "工作目录(可选,默认用户 home)",
|
||||||
|
},
|
||||||
|
"timeout": {
|
||||||
|
"type": "integer",
|
||||||
|
"description": "超时秒数(可选,默认 30)",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["command"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"file_read": {
|
||||||
|
"name": "cli_file_read",
|
||||||
|
"description": (
|
||||||
|
"读取用户本地电脑上的文件内容。"
|
||||||
|
"用于查看用户电脑上的配置文件、代码、文档等。"
|
||||||
|
),
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"path": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "文件绝对路径",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["path"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"file_write": {
|
||||||
|
"name": "cli_file_write",
|
||||||
|
"description": (
|
||||||
|
"写入文件到用户本地电脑。"
|
||||||
|
"用于创建或更新用户电脑上的文件。"
|
||||||
|
),
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"path": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "文件绝对路径",
|
||||||
|
},
|
||||||
|
"content": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "文件内容",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["path", "content"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"file_ops": {
|
||||||
|
"name": "cli_file_ops",
|
||||||
|
"description": (
|
||||||
|
"在用户本地电脑上执行文件操作(复制/移动/删除/列表)。"
|
||||||
|
),
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"operation": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "操作类型",
|
||||||
|
"enum": ["copy", "move", "delete", "list"],
|
||||||
|
},
|
||||||
|
"source": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "源路径",
|
||||||
|
},
|
||||||
|
"destination": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "目标路径(copy/move 时必填)",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["operation", "source"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"grep": {
|
||||||
|
"name": "cli_grep",
|
||||||
|
"description": (
|
||||||
|
"在用户本地电脑上搜索文件内容。"
|
||||||
|
"用于在用户项目中查找代码、配置等。"
|
||||||
|
),
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"pattern": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "搜索模式",
|
||||||
|
},
|
||||||
|
"path": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "搜索路径(可选,默认当前目录)",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["pattern"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"code_execution": {
|
||||||
|
"name": "cli_code_execution",
|
||||||
|
"description": (
|
||||||
|
"在用户本地电脑上执行代码片段(Python)。"
|
||||||
|
"用于需要在用户本地环境中运行脚本的场景。"
|
||||||
|
),
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"code": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "要执行的代码",
|
||||||
|
},
|
||||||
|
"language": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "编程语言(当前仅支持 python)",
|
||||||
|
"enum": ["python"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["code"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ── 动态注册 / 注销 ──────────────────────────────────────
|
||||||
|
|
||||||
|
def register_cli_tools(
|
||||||
|
user_id: str,
|
||||||
|
capabilities: dict,
|
||||||
|
) -> list[str]:
|
||||||
|
"""
|
||||||
|
CLI 连接时调用——根据能力报告动态注册工具到 Hermes ToolRegistry。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
user_id: 用户 ID
|
||||||
|
capabilities: CLI 上报的 capability_report
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
已注册的工具名列表(带 cli_ 前缀)
|
||||||
|
"""
|
||||||
|
from toolsets import create_custom_toolset, TOOLSETS
|
||||||
|
|
||||||
|
if not _bridge:
|
||||||
|
logger.warning("[CliTunnel] bridge 未初始化,跳过工具注册")
|
||||||
|
return []
|
||||||
|
|
||||||
|
reported_tools = {t["name"] for t in capabilities.get("tools", [])}
|
||||||
|
registered: list[str] = []
|
||||||
|
|
||||||
|
for local_name, schema in _TOOL_SCHEMAS.items():
|
||||||
|
if local_name not in reported_tools:
|
||||||
|
continue
|
||||||
|
|
||||||
|
prefixed_name = schema["name"] # e.g. "cli_terminal"
|
||||||
|
|
||||||
|
# 避免与非 CLI 内置工具冲突
|
||||||
|
existing_toolset = registry.get_toolset_for_tool(prefixed_name)
|
||||||
|
if existing_toolset and existing_toolset != "connectors":
|
||||||
|
logger.warning(
|
||||||
|
"[CliTunnel] 工具 '%s' 与 toolset '%s' 冲突,跳过",
|
||||||
|
prefixed_name, existing_toolset,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
registry.register(
|
||||||
|
name=prefixed_name,
|
||||||
|
toolset="connectors",
|
||||||
|
schema=schema,
|
||||||
|
handler=_make_cli_handler(local_name, user_id),
|
||||||
|
check_fn=_make_check_fn(user_id),
|
||||||
|
is_async=False,
|
||||||
|
description=schema["description"],
|
||||||
|
emoji="💻",
|
||||||
|
)
|
||||||
|
registered.append(prefixed_name)
|
||||||
|
|
||||||
|
# 注入到 hermes-* umbrella toolsets 以确保通过 enabled_toolsets 过滤
|
||||||
|
if registered:
|
||||||
|
for ts_name, ts in TOOLSETS.items():
|
||||||
|
if ts_name.startswith("hermes-"):
|
||||||
|
for tool_name in registered:
|
||||||
|
if tool_name not in ts["tools"]:
|
||||||
|
ts["tools"].append(tool_name)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"[CliTunnel] 为用户 %s 注册 %d 个本地工具: %s",
|
||||||
|
user_id, len(registered), ", ".join(registered),
|
||||||
|
)
|
||||||
|
return registered
|
||||||
|
|
||||||
|
|
||||||
|
def deregister_cli_tools(tools: list[str]) -> None:
|
||||||
|
"""
|
||||||
|
CLI 断开时调用——从 ToolRegistry 注销工具。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
tools: 之前 register_cli_tools 返回的工具名列表
|
||||||
|
"""
|
||||||
|
from toolsets import TOOLSETS
|
||||||
|
|
||||||
|
for name in tools:
|
||||||
|
registry.deregister(name)
|
||||||
|
# 从 hermes-* umbrella toolsets 中移除
|
||||||
|
for ts_name, ts in TOOLSETS.items():
|
||||||
|
if ts_name.startswith("hermes-"):
|
||||||
|
try:
|
||||||
|
ts["tools"].remove(name)
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if tools:
|
||||||
|
logger.info("[CliTunnel] 注销 %d 个本地工具: %s", len(tools), ", ".join(tools))
|
||||||
|
|
||||||
|
|
||||||
|
# ── Handler 工厂 ─────────────────────────────────────────
|
||||||
|
|
||||||
|
def _make_cli_handler(local_tool_name: str, user_id: str):
|
||||||
|
"""
|
||||||
|
返回闭包 handler,走 Tunnel 派发工具调用。
|
||||||
|
|
||||||
|
签名:handler(args, **kwargs) -> str(符合 registry.dispatch 要求)
|
||||||
|
"""
|
||||||
|
def _handler(args: dict, **kwargs) -> str:
|
||||||
|
if not _bridge:
|
||||||
|
return tool_error("CLI Bridge 未初始化")
|
||||||
|
|
||||||
|
if not _bridge.is_connected(user_id):
|
||||||
|
return tool_error(
|
||||||
|
f"本地 CLI 未连接。请确保 Mind CLI 正在运行并已连接隧道。"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 同步调用异步 dispatch(复用 model_tools._run_async 模式)
|
||||||
|
try:
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
except RuntimeError:
|
||||||
|
loop = None
|
||||||
|
|
||||||
|
coro = _bridge.dispatch_tool_call(user_id, local_tool_name, args)
|
||||||
|
|
||||||
|
try:
|
||||||
|
if loop and loop.is_running():
|
||||||
|
# 在 async 上下文中(SSE gateway)→ 开线程
|
||||||
|
import concurrent.futures
|
||||||
|
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||||
|
future = pool.submit(asyncio.run, coro)
|
||||||
|
result = future.result(timeout=130)
|
||||||
|
else:
|
||||||
|
result = asyncio.run(coro)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("[CliTunnel] 工具 %s 调用失败: %s", local_tool_name, e)
|
||||||
|
return tool_error(f"本地工具调用失败: {e}")
|
||||||
|
|
||||||
|
# result 是 dict(从 CLI 返回的 JSON-RPC result)
|
||||||
|
if isinstance(result, dict) and "error" in result:
|
||||||
|
return tool_error(result["error"])
|
||||||
|
|
||||||
|
return json.dumps(
|
||||||
|
{"ok": True, "tool": local_tool_name, "result": result},
|
||||||
|
ensure_ascii=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _handler
|
||||||
|
|
||||||
|
|
||||||
|
def _make_check_fn(user_id: str):
|
||||||
|
"""
|
||||||
|
返回 check_fn 闭包——registry 每次 get_definitions 时调用。
|
||||||
|
|
||||||
|
只有当 bridge 存在且该用户 CLI 在线时才返回 True,
|
||||||
|
离线时工具自动从 LLM 视野消失。
|
||||||
|
"""
|
||||||
|
def _check() -> bool:
|
||||||
|
return bool(_bridge and _bridge.is_connected(user_id))
|
||||||
|
return _check
|
||||||
@@ -0,0 +1,364 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""
|
||||||
|
Toolset Distributions Module
|
||||||
|
|
||||||
|
This module defines distributions of toolsets for data generation runs.
|
||||||
|
Each distribution specifies which toolsets should be used and their probability
|
||||||
|
of being selected for any given prompt during the batch processing.
|
||||||
|
|
||||||
|
A distribution is a dictionary mapping toolset names to their selection probability (%).
|
||||||
|
Probabilities should sum to 100, but the system will normalize if they don't.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
from toolset_distributions import get_distribution, list_distributions
|
||||||
|
|
||||||
|
# Get a specific distribution
|
||||||
|
dist = get_distribution("image_gen")
|
||||||
|
|
||||||
|
# List all available distributions
|
||||||
|
all_dists = list_distributions()
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
import random
|
||||||
|
from toolsets import validate_toolset
|
||||||
|
|
||||||
|
|
||||||
|
# Distribution definitions
|
||||||
|
# Each key is a distribution name, and the value is a dict of toolset_name: probability_percentage
|
||||||
|
DISTRIBUTIONS = {
|
||||||
|
# Default: All tools available 100% of the time
|
||||||
|
"default": {
|
||||||
|
"description": "All available tools, all the time",
|
||||||
|
"toolsets": {
|
||||||
|
"web": 100,
|
||||||
|
"vision": 100,
|
||||||
|
"image_gen": 100,
|
||||||
|
"terminal": 100,
|
||||||
|
"file": 100,
|
||||||
|
"moa": 100,
|
||||||
|
"browser": 100
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Image generation focused distribution
|
||||||
|
"image_gen": {
|
||||||
|
"description": "Heavy focus on image generation with vision and web support",
|
||||||
|
"toolsets": {
|
||||||
|
"image_gen": 90, # 80% chance of image generation tools
|
||||||
|
"vision": 90, # 60% chance of vision tools
|
||||||
|
"web": 55, # 40% chance of web tools
|
||||||
|
"terminal": 45,
|
||||||
|
"moa": 10 # 20% chance of reasoning tools
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Research-focused distribution
|
||||||
|
"research": {
|
||||||
|
"description": "Web research with vision analysis and reasoning",
|
||||||
|
"toolsets": {
|
||||||
|
"web": 90, # 90% chance of web tools
|
||||||
|
"browser": 70, # 70% chance of browser tools for deep research
|
||||||
|
"vision": 50, # 50% chance of vision tools
|
||||||
|
"moa": 40, # 40% chance of reasoning tools
|
||||||
|
"terminal": 10 # 10% chance of terminal tools
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Scientific problem solving focused distribution
|
||||||
|
"science": {
|
||||||
|
"description": "Scientific research with web, terminal, file, and browser capabilities",
|
||||||
|
"toolsets": {
|
||||||
|
"web": 94, # 94% chance of web tools
|
||||||
|
"terminal": 94, # 94% chance of terminal tools
|
||||||
|
"file": 94, # 94% chance of file tools
|
||||||
|
"vision": 65, # 65% chance of vision tools
|
||||||
|
"browser": 50, # 50% chance of browser for accessing papers/databases
|
||||||
|
"image_gen": 15, # 15% chance of image generation tools
|
||||||
|
"moa": 10 # 10% chance of reasoning tools
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Development-focused distribution
|
||||||
|
"development": {
|
||||||
|
"description": "Terminal, file tools, and reasoning with occasional web lookup",
|
||||||
|
"toolsets": {
|
||||||
|
"terminal": 80, # 80% chance of terminal tools
|
||||||
|
"file": 80, # 80% chance of file tools (read, write, patch, search)
|
||||||
|
"moa": 60, # 60% chance of reasoning tools
|
||||||
|
"web": 30, # 30% chance of web tools
|
||||||
|
"vision": 10 # 10% chance of vision tools
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Safe mode (no terminal)
|
||||||
|
"safe": {
|
||||||
|
"description": "All tools except terminal for safety",
|
||||||
|
"toolsets": {
|
||||||
|
"web": 80,
|
||||||
|
"browser": 70, # Browser is safe (no local filesystem access)
|
||||||
|
"vision": 60,
|
||||||
|
"image_gen": 60,
|
||||||
|
"moa": 50
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Balanced distribution
|
||||||
|
"balanced": {
|
||||||
|
"description": "Equal probability of all toolsets",
|
||||||
|
"toolsets": {
|
||||||
|
"web": 50,
|
||||||
|
"vision": 50,
|
||||||
|
"image_gen": 50,
|
||||||
|
"terminal": 50,
|
||||||
|
"file": 50,
|
||||||
|
"moa": 50,
|
||||||
|
"browser": 50
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Minimal (web only)
|
||||||
|
"minimal": {
|
||||||
|
"description": "Only web tools for basic research",
|
||||||
|
"toolsets": {
|
||||||
|
"web": 100
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Terminal only
|
||||||
|
"terminal_only": {
|
||||||
|
"description": "Terminal and file tools for code execution tasks",
|
||||||
|
"toolsets": {
|
||||||
|
"terminal": 100,
|
||||||
|
"file": 100
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Terminal + web (common for coding tasks that need docs)
|
||||||
|
"terminal_web": {
|
||||||
|
"description": "Terminal and file tools with web search for documentation lookup",
|
||||||
|
"toolsets": {
|
||||||
|
"terminal": 100,
|
||||||
|
"file": 100,
|
||||||
|
"web": 100
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Creative (vision + image generation)
|
||||||
|
"creative": {
|
||||||
|
"description": "Image generation and vision analysis focus",
|
||||||
|
"toolsets": {
|
||||||
|
"image_gen": 90,
|
||||||
|
"vision": 90,
|
||||||
|
"web": 30
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Reasoning heavy
|
||||||
|
"reasoning": {
|
||||||
|
"description": "Heavy mixture of agents usage with minimal other tools",
|
||||||
|
"toolsets": {
|
||||||
|
"moa": 90,
|
||||||
|
"web": 30,
|
||||||
|
"terminal": 20
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Browser-based web interaction
|
||||||
|
"browser_use": {
|
||||||
|
"description": "Full browser-based web interaction with search, vision, and page control",
|
||||||
|
"toolsets": {
|
||||||
|
"browser": 100, # All browser tools always available
|
||||||
|
"web": 80, # Web search for finding URLs and quick lookups
|
||||||
|
"vision": 70 # Vision analysis for images found on pages
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Browser only (no other tools)
|
||||||
|
"browser_only": {
|
||||||
|
"description": "Only browser automation tools for pure web interaction tasks",
|
||||||
|
"toolsets": {
|
||||||
|
"browser": 100
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Browser-focused tasks distribution (for browser-use-tasks.jsonl)
|
||||||
|
"browser_tasks": {
|
||||||
|
"description": "Browser-focused distribution (browser toolset includes web_search for finding URLs since Google blocks direct browser searches)",
|
||||||
|
"toolsets": {
|
||||||
|
"browser": 97, # 97% - browser tools (includes web_search) almost always available
|
||||||
|
"vision": 12, # 12% - vision analysis occasionally
|
||||||
|
"terminal": 15 # 15% - terminal occasionally for local operations
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Terminal-focused tasks distribution (for nous-terminal-tasks.jsonl)
|
||||||
|
"terminal_tasks": {
|
||||||
|
"description": "Terminal-focused distribution with high terminal/file availability, occasional other tools",
|
||||||
|
"toolsets": {
|
||||||
|
"terminal": 97, # 97% - terminal almost always available
|
||||||
|
"file": 97, # 97% - file tools almost always available
|
||||||
|
"web": 97, # 15% - web search/scrape for documentation
|
||||||
|
"browser": 75, # 10% - browser occasionally for web interaction
|
||||||
|
"vision": 50, # 8% - vision analysis rarely
|
||||||
|
"image_gen": 10 # 3% - image generation very rarely
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
# Mixed browser+terminal tasks distribution (for mixed-browser-terminal-tasks.jsonl)
|
||||||
|
"mixed_tasks": {
|
||||||
|
"description": "Mixed distribution with high browser, terminal, and file availability for complex tasks",
|
||||||
|
"toolsets": {
|
||||||
|
"browser": 92, # 92% - browser tools highly available
|
||||||
|
"terminal": 92, # 92% - terminal highly available
|
||||||
|
"file": 92, # 92% - file tools highly available
|
||||||
|
"web": 35, # 35% - web search/scrape fairly common
|
||||||
|
"vision": 15, # 15% - vision analysis occasionally
|
||||||
|
"image_gen": 15 # 15% - image generation occasionally
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def get_distribution(name: str) -> Optional[Dict[str, any]]:
|
||||||
|
"""
|
||||||
|
Get a toolset distribution by name.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name (str): Name of the distribution
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict: Distribution definition with description and toolsets
|
||||||
|
None: If distribution not found
|
||||||
|
"""
|
||||||
|
return DISTRIBUTIONS.get(name)
|
||||||
|
|
||||||
|
|
||||||
|
def list_distributions() -> Dict[str, Dict]:
|
||||||
|
"""
|
||||||
|
List all available distributions.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict: All distribution definitions
|
||||||
|
"""
|
||||||
|
return DISTRIBUTIONS.copy()
|
||||||
|
|
||||||
|
|
||||||
|
def sample_toolsets_from_distribution(distribution_name: str) -> List[str]:
|
||||||
|
"""
|
||||||
|
Sample toolsets based on a distribution's probabilities.
|
||||||
|
|
||||||
|
Each toolset in the distribution has a % chance of being included.
|
||||||
|
This allows multiple toolsets to be active simultaneously.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
distribution_name (str): Name of the distribution to sample from
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List[str]: List of sampled toolset names
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If distribution name is not found
|
||||||
|
"""
|
||||||
|
dist = get_distribution(distribution_name)
|
||||||
|
if not dist:
|
||||||
|
raise ValueError(f"Unknown distribution: {distribution_name}")
|
||||||
|
|
||||||
|
# Sample each toolset independently based on its probability
|
||||||
|
selected_toolsets = []
|
||||||
|
|
||||||
|
for toolset_name, probability in dist["toolsets"].items():
|
||||||
|
# Validate toolset exists
|
||||||
|
if not validate_toolset(toolset_name):
|
||||||
|
print(f"⚠️ Warning: Toolset '{toolset_name}' in distribution '{distribution_name}' is not valid")
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Roll the dice - if random value is less than probability, include this toolset
|
||||||
|
if random.random() * 100 < probability:
|
||||||
|
selected_toolsets.append(toolset_name)
|
||||||
|
|
||||||
|
# If no toolsets were selected (can happen with low probabilities),
|
||||||
|
# ensure at least one toolset is selected by picking the highest probability one
|
||||||
|
if not selected_toolsets and dist["toolsets"]:
|
||||||
|
# Find toolset with highest probability
|
||||||
|
highest_prob_toolset = max(dist["toolsets"].items(), key=lambda x: x[1])[0]
|
||||||
|
if validate_toolset(highest_prob_toolset):
|
||||||
|
selected_toolsets.append(highest_prob_toolset)
|
||||||
|
|
||||||
|
return selected_toolsets
|
||||||
|
|
||||||
|
|
||||||
|
def validate_distribution(distribution_name: str) -> bool:
|
||||||
|
"""
|
||||||
|
Check if a distribution name is valid.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
distribution_name (str): Distribution name to validate
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: True if valid, False otherwise
|
||||||
|
"""
|
||||||
|
return distribution_name in DISTRIBUTIONS
|
||||||
|
|
||||||
|
|
||||||
|
def print_distribution_info(distribution_name: str) -> None:
|
||||||
|
"""
|
||||||
|
Print detailed information about a distribution.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
distribution_name (str): Distribution name
|
||||||
|
"""
|
||||||
|
dist = get_distribution(distribution_name)
|
||||||
|
if not dist:
|
||||||
|
print(f"❌ Unknown distribution: {distribution_name}")
|
||||||
|
return
|
||||||
|
|
||||||
|
print(f"\n📊 Distribution: {distribution_name}")
|
||||||
|
print(f" Description: {dist['description']}")
|
||||||
|
print(" Toolsets:")
|
||||||
|
for toolset, prob in sorted(dist["toolsets"].items(), key=lambda x: x[1], reverse=True):
|
||||||
|
print(f" • {toolset:15} : {prob:3}% chance")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
"""
|
||||||
|
Demo and testing of the distributions system
|
||||||
|
"""
|
||||||
|
print("📊 Toolset Distributions Demo")
|
||||||
|
print("=" * 60)
|
||||||
|
|
||||||
|
# List all distributions
|
||||||
|
print("\n📋 Available Distributions:")
|
||||||
|
print("-" * 40)
|
||||||
|
for name, dist in list_distributions().items():
|
||||||
|
print(f"\n {name}:")
|
||||||
|
print(f" {dist['description']}")
|
||||||
|
toolset_list = ", ".join([f"{ts}({p}%)" for ts, p in dist["toolsets"].items()])
|
||||||
|
print(f" Toolsets: {toolset_list}")
|
||||||
|
|
||||||
|
# Demo sampling
|
||||||
|
print("\n\n🎲 Sampling Examples:")
|
||||||
|
print("-" * 40)
|
||||||
|
|
||||||
|
test_distributions = ["image_gen", "research", "balanced", "default"]
|
||||||
|
|
||||||
|
for dist_name in test_distributions:
|
||||||
|
print(f"\n{dist_name}:")
|
||||||
|
# Sample 5 times to show variability
|
||||||
|
samples = []
|
||||||
|
for _ in range(5):
|
||||||
|
sampled = sample_toolsets_from_distribution(dist_name)
|
||||||
|
samples.append(sorted(sampled))
|
||||||
|
|
||||||
|
print(f" Sample 1: {samples[0]}")
|
||||||
|
print(f" Sample 2: {samples[1]}")
|
||||||
|
print(f" Sample 3: {samples[2]}")
|
||||||
|
print(f" Sample 4: {samples[3]}")
|
||||||
|
print(f" Sample 5: {samples[4]}")
|
||||||
|
|
||||||
|
# Show detailed info
|
||||||
|
print("\n\n📊 Detailed Distribution Info:")
|
||||||
|
print("-" * 40)
|
||||||
|
print_distribution_info("image_gen")
|
||||||
|
print_distribution_info("research")
|
||||||
|
|
||||||
@@ -0,0 +1,661 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""
|
||||||
|
Toolsets Module
|
||||||
|
|
||||||
|
This module provides a flexible system for defining and managing tool aliases/toolsets.
|
||||||
|
Toolsets allow you to group tools together for specific scenarios and can be composed
|
||||||
|
from individual tools or other toolsets.
|
||||||
|
|
||||||
|
Features:
|
||||||
|
- Define custom toolsets with specific tools
|
||||||
|
- Compose toolsets from other toolsets
|
||||||
|
- Built-in common toolsets for typical use cases
|
||||||
|
- Easy extension for new toolsets
|
||||||
|
- Support for dynamic toolset resolution
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
from toolsets import get_toolset, resolve_toolset, get_all_toolsets
|
||||||
|
|
||||||
|
# Get tools for a specific toolset
|
||||||
|
tools = get_toolset("research")
|
||||||
|
|
||||||
|
# Resolve a toolset to get all tool names (including from composed toolsets)
|
||||||
|
all_tools = resolve_toolset("full_stack")
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import List, Dict, Any, Set, Optional
|
||||||
|
|
||||||
|
|
||||||
|
# Shared tool list for CLI and all messaging platform toolsets.
|
||||||
|
# Edit this once to update all platforms simultaneously.
|
||||||
|
_HERMES_CORE_TOOLS = [
|
||||||
|
# Web
|
||||||
|
"web_search", "web_extract",
|
||||||
|
# Terminal + process management
|
||||||
|
"terminal", "process",
|
||||||
|
# File manipulation
|
||||||
|
"read_file", "write_file", "patch", "search_files",
|
||||||
|
# Vision + image generation
|
||||||
|
"vision_analyze", "image_generate",
|
||||||
|
# Skills
|
||||||
|
"skills_list", "skill_view", "skill_manage",
|
||||||
|
# Browser automation
|
||||||
|
"browser_navigate", "browser_snapshot", "browser_click",
|
||||||
|
"browser_type", "browser_scroll", "browser_back",
|
||||||
|
"browser_press", "browser_get_images",
|
||||||
|
"browser_vision", "browser_console",
|
||||||
|
# Text-to-speech
|
||||||
|
"text_to_speech",
|
||||||
|
# Planning & memory
|
||||||
|
"todo", "memory",
|
||||||
|
# Session history search
|
||||||
|
"session_search",
|
||||||
|
# Clarifying questions
|
||||||
|
"clarify",
|
||||||
|
# Code execution + delegation
|
||||||
|
"execute_code", "delegate_task",
|
||||||
|
# Cronjob management
|
||||||
|
"cronjob",
|
||||||
|
# Cross-platform messaging (gated on gateway running via check_fn)
|
||||||
|
"send_message",
|
||||||
|
# Home Assistant smart home control (gated on HASS_TOKEN via check_fn)
|
||||||
|
"ha_list_entities", "ha_get_state", "ha_list_services", "ha_call_service",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
# Core toolset definitions
|
||||||
|
# These can include individual tools or reference other toolsets
|
||||||
|
TOOLSETS = {
|
||||||
|
# Basic toolsets - individual tool categories
|
||||||
|
"web": {
|
||||||
|
"description": "Web research and content extraction tools",
|
||||||
|
"tools": ["web_search", "web_extract"],
|
||||||
|
"includes": [] # No other toolsets included
|
||||||
|
},
|
||||||
|
|
||||||
|
"search": {
|
||||||
|
"description": "Web search only (no content extraction/scraping)",
|
||||||
|
"tools": ["web_search"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"vision": {
|
||||||
|
"description": "Image analysis and vision tools",
|
||||||
|
"tools": ["vision_analyze"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"image_gen": {
|
||||||
|
"description": "Creative generation tools (images)",
|
||||||
|
"tools": ["image_generate"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"terminal": {
|
||||||
|
"description": "Terminal/command execution and process management tools",
|
||||||
|
"tools": ["terminal", "process"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"moa": {
|
||||||
|
"description": "Advanced reasoning and problem-solving tools",
|
||||||
|
"tools": ["mixture_of_agents"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"skills": {
|
||||||
|
"description": "Access, create, edit, and manage skill documents with specialized instructions and knowledge",
|
||||||
|
"tools": ["skills_list", "skill_view", "skill_manage"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"browser": {
|
||||||
|
"description": "Browser automation for web interaction (navigate, click, type, scroll, iframes, hold-click) with web search for finding URLs",
|
||||||
|
"tools": [
|
||||||
|
"browser_navigate", "browser_snapshot", "browser_click",
|
||||||
|
"browser_type", "browser_scroll", "browser_back",
|
||||||
|
"browser_press", "browser_get_images",
|
||||||
|
"browser_vision", "browser_console", "web_search"
|
||||||
|
],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"cronjob": {
|
||||||
|
"description": "Cronjob management tool - create, list, update, pause, resume, remove, and trigger scheduled tasks",
|
||||||
|
"tools": ["cronjob"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"messaging": {
|
||||||
|
"description": "Cross-platform messaging: send messages to Telegram, Discord, Slack, SMS, etc.",
|
||||||
|
"tools": ["send_message"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"rl": {
|
||||||
|
"description": "RL training tools for running reinforcement learning on Tinker-Atropos",
|
||||||
|
"tools": [
|
||||||
|
"rl_list_environments", "rl_select_environment",
|
||||||
|
"rl_get_current_config", "rl_edit_config",
|
||||||
|
"rl_start_training", "rl_check_status",
|
||||||
|
"rl_stop_training", "rl_get_results",
|
||||||
|
"rl_list_runs", "rl_test_inference"
|
||||||
|
],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"file": {
|
||||||
|
"description": "File manipulation tools: read, write, patch (with fuzzy matching), and search (content + files)",
|
||||||
|
"tools": ["read_file", "write_file", "patch", "search_files"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"tts": {
|
||||||
|
"description": "Text-to-speech: convert text to audio with Edge TTS (free), ElevenLabs, or OpenAI",
|
||||||
|
"tools": ["text_to_speech"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"todo": {
|
||||||
|
"description": "Task planning and tracking for multi-step work",
|
||||||
|
"tools": ["todo"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"memory": {
|
||||||
|
"description": "Persistent memory across sessions (personal notes + user profile)",
|
||||||
|
"tools": ["memory"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"session_search": {
|
||||||
|
"description": "Search and recall past conversations with summarization",
|
||||||
|
"tools": ["session_search"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"clarify": {
|
||||||
|
"description": "Ask the user clarifying questions (multiple-choice or open-ended)",
|
||||||
|
"tools": ["clarify"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"code_execution": {
|
||||||
|
"description": "Run Python scripts that call tools programmatically (reduces LLM round trips)",
|
||||||
|
"tools": ["execute_code"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"delegation": {
|
||||||
|
"description": "Spawn subagents with isolated context for complex subtasks",
|
||||||
|
"tools": ["delegate_task"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
# "honcho" toolset removed — Honcho is now a memory provider plugin.
|
||||||
|
# Tools are injected via MemoryManager, not the toolset system.
|
||||||
|
|
||||||
|
"homeassistant": {
|
||||||
|
"description": "Home Assistant smart home control and monitoring",
|
||||||
|
"tools": ["ha_list_entities", "ha_get_state", "ha_list_services", "ha_call_service"],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
|
||||||
|
# Scenario-specific toolsets
|
||||||
|
|
||||||
|
"debugging": {
|
||||||
|
"description": "Debugging and troubleshooting toolkit",
|
||||||
|
"tools": ["terminal", "process"],
|
||||||
|
"includes": ["web", "file"] # For searching error messages and solutions, and file operations
|
||||||
|
},
|
||||||
|
|
||||||
|
"safe": {
|
||||||
|
"description": "Safe toolkit without terminal access",
|
||||||
|
"tools": [],
|
||||||
|
"includes": ["web", "vision", "image_gen"]
|
||||||
|
},
|
||||||
|
|
||||||
|
# ==========================================================================
|
||||||
|
# Full Hermes toolsets (CLI + messaging platforms)
|
||||||
|
#
|
||||||
|
# All platforms share the same core tools (including send_message,
|
||||||
|
# which is gated on gateway running via its check_fn).
|
||||||
|
# ==========================================================================
|
||||||
|
|
||||||
|
"hermes-acp": {
|
||||||
|
"description": "Editor integration (VS Code, Zed, JetBrains) — coding-focused tools without messaging, audio, or clarify UI",
|
||||||
|
"tools": [
|
||||||
|
"web_search", "web_extract",
|
||||||
|
"terminal", "process",
|
||||||
|
"read_file", "write_file", "patch", "search_files",
|
||||||
|
"vision_analyze",
|
||||||
|
"skills_list", "skill_view", "skill_manage",
|
||||||
|
"browser_navigate", "browser_snapshot", "browser_click",
|
||||||
|
"browser_type", "browser_scroll", "browser_back",
|
||||||
|
"browser_press", "browser_get_images",
|
||||||
|
"browser_vision", "browser_console",
|
||||||
|
"todo", "memory",
|
||||||
|
"session_search",
|
||||||
|
"execute_code", "delegate_task",
|
||||||
|
],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-api-server": {
|
||||||
|
"description": "OpenAI-compatible API server — full agent tools accessible via HTTP (no interactive UI tools like clarify or send_message)",
|
||||||
|
"tools": [
|
||||||
|
# Web
|
||||||
|
"web_search", "web_extract",
|
||||||
|
# Terminal + process management
|
||||||
|
"terminal", "process",
|
||||||
|
# File manipulation
|
||||||
|
"read_file", "write_file", "patch", "search_files",
|
||||||
|
# Vision + image generation
|
||||||
|
"vision_analyze", "image_generate",
|
||||||
|
# Skills
|
||||||
|
"skills_list", "skill_view", "skill_manage",
|
||||||
|
# Browser automation
|
||||||
|
"browser_navigate", "browser_snapshot", "browser_click",
|
||||||
|
"browser_type", "browser_scroll", "browser_back",
|
||||||
|
"browser_press", "browser_get_images",
|
||||||
|
"browser_vision", "browser_console",
|
||||||
|
# Planning & memory
|
||||||
|
"todo", "memory",
|
||||||
|
# Session history search
|
||||||
|
"session_search",
|
||||||
|
# Code execution + delegation
|
||||||
|
"execute_code", "delegate_task",
|
||||||
|
# Cronjob management
|
||||||
|
"cronjob",
|
||||||
|
# Home Assistant smart home control (gated on HASS_TOKEN via check_fn)
|
||||||
|
"ha_list_entities", "ha_get_state", "ha_list_services", "ha_call_service",
|
||||||
|
|
||||||
|
],
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-cli": {
|
||||||
|
"description": "Full interactive CLI toolset - all default tools plus cronjob management",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-telegram": {
|
||||||
|
"description": "Telegram bot toolset - full access for personal use (terminal has safety checks)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-discord": {
|
||||||
|
"description": "Discord bot toolset - full access (terminal has safety checks via dangerous command approval)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-whatsapp": {
|
||||||
|
"description": "WhatsApp bot toolset - similar to Telegram (personal messaging, more trusted)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-slack": {
|
||||||
|
"description": "Slack bot toolset - full access for workspace use (terminal has safety checks)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-signal": {
|
||||||
|
"description": "Signal bot toolset - encrypted messaging platform (full access)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-bluebubbles": {
|
||||||
|
"description": "BlueBubbles iMessage bot toolset - Apple iMessage via local BlueBubbles server",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-homeassistant": {
|
||||||
|
"description": "Home Assistant bot toolset - smart home event monitoring and control",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-email": {
|
||||||
|
"description": "Email bot toolset - interact with Hermes via email (IMAP/SMTP)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-mattermost": {
|
||||||
|
"description": "Mattermost bot toolset - self-hosted team messaging (full access)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-matrix": {
|
||||||
|
"description": "Matrix bot toolset - decentralized encrypted messaging (full access)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-dingtalk": {
|
||||||
|
"description": "DingTalk bot toolset - enterprise messaging platform (full access)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-feishu": {
|
||||||
|
"description": "Feishu/Lark bot toolset - enterprise messaging via Feishu/Lark (full access)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-weixin": {
|
||||||
|
"description": "Weixin bot toolset - personal WeChat messaging via iLink (full access)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-qqbot": {
|
||||||
|
"description": "QQBot toolset - QQ messaging via Official Bot API v2 (full access)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-wecom": {
|
||||||
|
"description": "WeCom bot toolset - enterprise WeChat messaging (full access)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-wecom-callback": {
|
||||||
|
"description": "WeCom callback toolset - enterprise self-built app messaging (full access)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-sms": {
|
||||||
|
"description": "SMS bot toolset - interact with Hermes via SMS (Twilio)",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-webhook": {
|
||||||
|
"description": "Webhook toolset - receive and process external webhook events",
|
||||||
|
"tools": _HERMES_CORE_TOOLS,
|
||||||
|
"includes": []
|
||||||
|
},
|
||||||
|
|
||||||
|
"hermes-gateway": {
|
||||||
|
"description": "Gateway toolset - union of all messaging platform tools",
|
||||||
|
"tools": [],
|
||||||
|
"includes": ["hermes-telegram", "hermes-discord", "hermes-whatsapp", "hermes-slack", "hermes-signal", "hermes-bluebubbles", "hermes-homeassistant", "hermes-email", "hermes-sms", "hermes-mattermost", "hermes-matrix", "hermes-dingtalk", "hermes-feishu", "hermes-wecom", "hermes-wecom-callback", "hermes-weixin", "hermes-qqbot", "hermes-webhook"]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
def get_toolset(name: str) -> Optional[Dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
Get a toolset definition by name.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name (str): Name of the toolset
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict: Toolset definition with description, tools, and includes
|
||||||
|
None: If toolset not found
|
||||||
|
"""
|
||||||
|
# Return toolset definition
|
||||||
|
return TOOLSETS.get(name)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_toolset(name: str, visited: Set[str] = None) -> List[str]:
|
||||||
|
"""
|
||||||
|
Recursively resolve a toolset to get all tool names.
|
||||||
|
|
||||||
|
This function handles toolset composition by recursively resolving
|
||||||
|
included toolsets and combining all tools.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name (str): Name of the toolset to resolve
|
||||||
|
visited (Set[str]): Set of already visited toolsets (for cycle detection)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List[str]: List of all tool names in the toolset
|
||||||
|
"""
|
||||||
|
if visited is None:
|
||||||
|
visited = set()
|
||||||
|
|
||||||
|
# Special aliases that represent all tools across every toolset
|
||||||
|
# This ensures future toolsets are automatically included without changes.
|
||||||
|
if name in {"all", "*"}:
|
||||||
|
all_tools: Set[str] = set()
|
||||||
|
for toolset_name in get_toolset_names():
|
||||||
|
# Use a fresh visited set per branch to avoid cross-branch contamination
|
||||||
|
resolved = resolve_toolset(toolset_name, visited.copy())
|
||||||
|
all_tools.update(resolved)
|
||||||
|
return list(all_tools)
|
||||||
|
|
||||||
|
# Check for cycles / already-resolved (diamond deps).
|
||||||
|
# Silently return [] — either this is a diamond (not a bug, tools already
|
||||||
|
# collected via another path) or a genuine cycle (safe to skip).
|
||||||
|
if name in visited:
|
||||||
|
return []
|
||||||
|
|
||||||
|
visited.add(name)
|
||||||
|
|
||||||
|
# Get toolset definition
|
||||||
|
toolset = TOOLSETS.get(name)
|
||||||
|
if not toolset:
|
||||||
|
# Fall back to tool registry for plugin-provided toolsets
|
||||||
|
if name in _get_plugin_toolset_names():
|
||||||
|
try:
|
||||||
|
from tools.registry import registry
|
||||||
|
return registry.get_tool_names_for_toolset(name)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return []
|
||||||
|
|
||||||
|
# Collect direct tools
|
||||||
|
tools = set(toolset.get("tools", []))
|
||||||
|
|
||||||
|
# Recursively resolve included toolsets, sharing the visited set across
|
||||||
|
# sibling includes so diamond dependencies are only resolved once and
|
||||||
|
# cycle warnings don't fire multiple times for the same cycle.
|
||||||
|
for included_name in toolset.get("includes", []):
|
||||||
|
included_tools = resolve_toolset(included_name, visited)
|
||||||
|
tools.update(included_tools)
|
||||||
|
|
||||||
|
return list(tools)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_multiple_toolsets(toolset_names: List[str]) -> List[str]:
|
||||||
|
"""
|
||||||
|
Resolve multiple toolsets and combine their tools.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
toolset_names (List[str]): List of toolset names to resolve
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List[str]: Combined list of all tool names (deduplicated)
|
||||||
|
"""
|
||||||
|
all_tools = set()
|
||||||
|
|
||||||
|
for name in toolset_names:
|
||||||
|
tools = resolve_toolset(name)
|
||||||
|
all_tools.update(tools)
|
||||||
|
|
||||||
|
return list(all_tools)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_plugin_toolset_names() -> Set[str]:
|
||||||
|
"""Return toolset names registered by plugins (from the tool registry).
|
||||||
|
|
||||||
|
These are toolsets that exist in the registry but not in the static
|
||||||
|
``TOOLSETS`` dict — i.e. they were added by plugins at load time.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from tools.registry import registry
|
||||||
|
return {
|
||||||
|
toolset_name
|
||||||
|
for toolset_name in registry.get_registered_toolset_names()
|
||||||
|
if toolset_name not in TOOLSETS
|
||||||
|
}
|
||||||
|
except Exception:
|
||||||
|
return set()
|
||||||
|
|
||||||
|
|
||||||
|
def get_all_toolsets() -> Dict[str, Dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
Get all available toolsets with their definitions.
|
||||||
|
|
||||||
|
Includes both statically-defined toolsets and plugin-registered ones.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict: All toolset definitions
|
||||||
|
"""
|
||||||
|
result = TOOLSETS.copy()
|
||||||
|
# Add plugin-provided toolsets (synthetic entries)
|
||||||
|
for ts_name in _get_plugin_toolset_names():
|
||||||
|
if ts_name not in result:
|
||||||
|
try:
|
||||||
|
from tools.registry import registry
|
||||||
|
tools = registry.get_tool_names_for_toolset(ts_name)
|
||||||
|
result[ts_name] = {
|
||||||
|
"description": f"Plugin toolset: {ts_name}",
|
||||||
|
"tools": tools,
|
||||||
|
}
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def get_toolset_names() -> List[str]:
|
||||||
|
"""
|
||||||
|
Get names of all available toolsets (excluding aliases).
|
||||||
|
|
||||||
|
Includes plugin-registered toolset names.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List[str]: List of toolset names
|
||||||
|
"""
|
||||||
|
names = set(TOOLSETS.keys())
|
||||||
|
names |= _get_plugin_toolset_names()
|
||||||
|
return sorted(names)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
def validate_toolset(name: str) -> bool:
|
||||||
|
"""
|
||||||
|
Check if a toolset name is valid.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name (str): Toolset name to validate
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: True if valid, False otherwise
|
||||||
|
"""
|
||||||
|
# Accept special alias names for convenience
|
||||||
|
if name in {"all", "*"}:
|
||||||
|
return True
|
||||||
|
if name in TOOLSETS:
|
||||||
|
return True
|
||||||
|
# Check tool registry for plugin-provided toolsets
|
||||||
|
return name in _get_plugin_toolset_names()
|
||||||
|
|
||||||
|
|
||||||
|
def create_custom_toolset(
|
||||||
|
name: str,
|
||||||
|
description: str,
|
||||||
|
tools: List[str] = None,
|
||||||
|
includes: List[str] = None
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Create a custom toolset at runtime.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name (str): Name for the new toolset
|
||||||
|
description (str): Description of the toolset
|
||||||
|
tools (List[str]): Direct tools to include
|
||||||
|
includes (List[str]): Other toolsets to include
|
||||||
|
"""
|
||||||
|
TOOLSETS[name] = {
|
||||||
|
"description": description,
|
||||||
|
"tools": tools or [],
|
||||||
|
"includes": includes or []
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
def get_toolset_info(name: str) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
Get detailed information about a toolset including resolved tools.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name (str): Toolset name
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict: Detailed toolset information
|
||||||
|
"""
|
||||||
|
toolset = get_toolset(name)
|
||||||
|
if not toolset:
|
||||||
|
return None
|
||||||
|
|
||||||
|
resolved_tools = resolve_toolset(name)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"name": name,
|
||||||
|
"description": toolset["description"],
|
||||||
|
"direct_tools": toolset["tools"],
|
||||||
|
"includes": toolset["includes"],
|
||||||
|
"resolved_tools": resolved_tools,
|
||||||
|
"tool_count": len(resolved_tools),
|
||||||
|
"is_composite": bool(toolset["includes"])
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
print("Toolsets System Demo")
|
||||||
|
print("=" * 60)
|
||||||
|
|
||||||
|
print("\nAvailable Toolsets:")
|
||||||
|
print("-" * 40)
|
||||||
|
for name, toolset in get_all_toolsets().items():
|
||||||
|
info = get_toolset_info(name)
|
||||||
|
composite = "[composite]" if info["is_composite"] else "[leaf]"
|
||||||
|
print(f" {composite} {name:20} - {toolset['description']}")
|
||||||
|
print(f" Tools: {len(info['resolved_tools'])} total")
|
||||||
|
|
||||||
|
print("\nToolset Resolution Examples:")
|
||||||
|
print("-" * 40)
|
||||||
|
for name in ["web", "terminal", "safe", "debugging"]:
|
||||||
|
tools = resolve_toolset(name)
|
||||||
|
print(f"\n {name}:")
|
||||||
|
print(f" Resolved to {len(tools)} tools: {', '.join(sorted(tools))}")
|
||||||
|
|
||||||
|
print("\nMultiple Toolset Resolution:")
|
||||||
|
print("-" * 40)
|
||||||
|
combined = resolve_multiple_toolsets(["web", "vision", "terminal"])
|
||||||
|
print(" Combining ['web', 'vision', 'terminal']:")
|
||||||
|
print(f" Result: {', '.join(sorted(combined))}")
|
||||||
|
|
||||||
|
print("\nCustom Toolset Creation:")
|
||||||
|
print("-" * 40)
|
||||||
|
create_custom_toolset(
|
||||||
|
name="my_custom",
|
||||||
|
description="My custom toolset for specific tasks",
|
||||||
|
tools=["web_search"],
|
||||||
|
includes=["terminal", "vision"]
|
||||||
|
)
|
||||||
|
custom_info = get_toolset_info("my_custom")
|
||||||
|
print(" Created 'my_custom' toolset:")
|
||||||
|
print(f" Description: {custom_info['description']}")
|
||||||
|
print(f" Resolved tools: {', '.join(custom_info['resolved_tools'])}")
|
||||||
@@ -0,0 +1,164 @@
|
|||||||
|
"""Shared utility functions for hermes-agent."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Union
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
TRUTHY_STRINGS = frozenset({"1", "true", "yes", "on"})
|
||||||
|
|
||||||
|
|
||||||
|
def is_truthy_value(value: Any, default: bool = False) -> bool:
|
||||||
|
"""Coerce bool-ish values using the project's shared truthy string set."""
|
||||||
|
if value is None:
|
||||||
|
return default
|
||||||
|
if isinstance(value, bool):
|
||||||
|
return value
|
||||||
|
if isinstance(value, str):
|
||||||
|
return value.strip().lower() in TRUTHY_STRINGS
|
||||||
|
return bool(value)
|
||||||
|
|
||||||
|
|
||||||
|
def env_var_enabled(name: str, default: str = "") -> bool:
|
||||||
|
"""Return True when an environment variable is set to a truthy value."""
|
||||||
|
return is_truthy_value(os.getenv(name, default), default=False)
|
||||||
|
|
||||||
|
|
||||||
|
def atomic_json_write(
|
||||||
|
path: Union[str, Path],
|
||||||
|
data: Any,
|
||||||
|
*,
|
||||||
|
indent: int = 2,
|
||||||
|
**dump_kwargs: Any,
|
||||||
|
) -> None:
|
||||||
|
"""Write JSON data to a file atomically.
|
||||||
|
|
||||||
|
Uses temp file + fsync + os.replace to ensure the target file is never
|
||||||
|
left in a partially-written state. If the process crashes mid-write,
|
||||||
|
the previous version of the file remains intact.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
path: Target file path (will be created or overwritten).
|
||||||
|
data: JSON-serializable data to write.
|
||||||
|
indent: JSON indentation (default 2).
|
||||||
|
**dump_kwargs: Additional keyword args forwarded to json.dump(), such
|
||||||
|
as default=str for non-native types.
|
||||||
|
"""
|
||||||
|
path = Path(path)
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
fd, tmp_path = tempfile.mkstemp(
|
||||||
|
dir=str(path.parent),
|
||||||
|
prefix=f".{path.stem}_",
|
||||||
|
suffix=".tmp",
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
with os.fdopen(fd, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(
|
||||||
|
data,
|
||||||
|
f,
|
||||||
|
indent=indent,
|
||||||
|
ensure_ascii=False,
|
||||||
|
**dump_kwargs,
|
||||||
|
)
|
||||||
|
f.flush()
|
||||||
|
os.fsync(f.fileno())
|
||||||
|
os.replace(tmp_path, path)
|
||||||
|
except BaseException:
|
||||||
|
# Intentionally catch BaseException so temp-file cleanup still runs for
|
||||||
|
# KeyboardInterrupt/SystemExit before re-raising the original signal.
|
||||||
|
try:
|
||||||
|
os.unlink(tmp_path)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
def atomic_yaml_write(
|
||||||
|
path: Union[str, Path],
|
||||||
|
data: Any,
|
||||||
|
*,
|
||||||
|
default_flow_style: bool = False,
|
||||||
|
sort_keys: bool = False,
|
||||||
|
extra_content: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Write YAML data to a file atomically.
|
||||||
|
|
||||||
|
Uses temp file + fsync + os.replace to ensure the target file is never
|
||||||
|
left in a partially-written state. If the process crashes mid-write,
|
||||||
|
the previous version of the file remains intact.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
path: Target file path (will be created or overwritten).
|
||||||
|
data: YAML-serializable data to write.
|
||||||
|
default_flow_style: YAML flow style (default False).
|
||||||
|
sort_keys: Whether to sort dict keys (default False).
|
||||||
|
extra_content: Optional string to append after the YAML dump
|
||||||
|
(e.g. commented-out sections for user reference).
|
||||||
|
"""
|
||||||
|
path = Path(path)
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
fd, tmp_path = tempfile.mkstemp(
|
||||||
|
dir=str(path.parent),
|
||||||
|
prefix=f".{path.stem}_",
|
||||||
|
suffix=".tmp",
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
with os.fdopen(fd, "w", encoding="utf-8") as f:
|
||||||
|
yaml.dump(data, f, default_flow_style=default_flow_style, sort_keys=sort_keys)
|
||||||
|
if extra_content:
|
||||||
|
f.write(extra_content)
|
||||||
|
f.flush()
|
||||||
|
os.fsync(f.fileno())
|
||||||
|
os.replace(tmp_path, path)
|
||||||
|
except BaseException:
|
||||||
|
# Match atomic_json_write: cleanup must also happen for process-level
|
||||||
|
# interruptions before we re-raise them.
|
||||||
|
try:
|
||||||
|
os.unlink(tmp_path)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
# ─── JSON Helpers ─────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def safe_json_loads(text: str, default: Any = None) -> Any:
|
||||||
|
"""Parse JSON, returning *default* on any parse error.
|
||||||
|
|
||||||
|
Replaces the ``try: json.loads(x) except (JSONDecodeError, TypeError)``
|
||||||
|
pattern duplicated across display.py, anthropic_adapter.py,
|
||||||
|
auxiliary_client.py, and others.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
return json.loads(text)
|
||||||
|
except (json.JSONDecodeError, TypeError, ValueError):
|
||||||
|
return default
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Environment Variable Helpers ─────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def env_int(key: str, default: int = 0) -> int:
|
||||||
|
"""Read an environment variable as an integer, with fallback."""
|
||||||
|
raw = os.getenv(key, "").strip()
|
||||||
|
if not raw:
|
||||||
|
return default
|
||||||
|
try:
|
||||||
|
return int(raw)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
return default
|
||||||
|
|
||||||
|
|
||||||
|
def env_bool(key: str, default: bool = False) -> bool:
|
||||||
|
"""Read an environment variable as a boolean."""
|
||||||
|
return is_truthy_value(os.getenv(key, ""), default=default)
|
||||||
@@ -28,6 +28,8 @@ dependencies = [
|
|||||||
# Mind CLI 专属
|
# Mind CLI 专属
|
||||||
"websockets>=12.0",
|
"websockets>=12.0",
|
||||||
"click>=8.0",
|
"click>=8.0",
|
||||||
|
# Hermes _vendor 隐性依赖
|
||||||
|
"python-dotenv>=1.0",
|
||||||
]
|
]
|
||||||
|
|
||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
|
|||||||
+27
-12
@@ -40,24 +40,39 @@ echo "📦 Vendoring hermes@$COMMIT → $VENDOR_DIR"
|
|||||||
rm -rf "$VENDOR_DIR"
|
rm -rf "$VENDOR_DIR"
|
||||||
mkdir -p "$VENDOR_DIR"
|
mkdir -p "$VENDOR_DIR"
|
||||||
|
|
||||||
# 核心模块
|
# 核心模块(单文件)
|
||||||
cp "$HERMES_SRC/cli.py" "$VENDOR_DIR/"
|
for pyfile in cli.py run_agent.py mcp_serve.py hermes_state.py hermes_constants.py \
|
||||||
cp "$HERMES_SRC/run_agent.py" "$VENDOR_DIR/"
|
batch_runner.py model_tools.py utils.py toolsets.py toolset_distributions.py \
|
||||||
cp "$HERMES_SRC/mcp_serve.py" "$VENDOR_DIR/"
|
hermes_time.py; do
|
||||||
cp "$HERMES_SRC/hermes_state.py" "$VENDOR_DIR/"
|
if [ -f "$HERMES_SRC/$pyfile" ]; then
|
||||||
cp "$HERMES_SRC/hermes_constants.py" "$VENDOR_DIR/"
|
cp "$HERMES_SRC/$pyfile" "$VENDOR_DIR/"
|
||||||
cp "$HERMES_SRC/batch_runner.py" "$VENDOR_DIR/" 2>/dev/null || true
|
echo " ✓ $pyfile"
|
||||||
|
else
|
||||||
|
echo " ⚠ $pyfile 不存在,跳过"
|
||||||
|
fi
|
||||||
|
done
|
||||||
|
|
||||||
# 子模块
|
# 子模块(目录)
|
||||||
cp -r "$HERMES_SRC/hermes_cli" "$VENDOR_DIR/"
|
for subdir in hermes_cli tools agent model_tools utils toolsets toolset_distributions gateway cron; do
|
||||||
cp -r "$HERMES_SRC/tools" "$VENDOR_DIR/"
|
if [ -d "$HERMES_SRC/$subdir" ]; then
|
||||||
cp -r "$HERMES_SRC/agent" "$VENDOR_DIR/"
|
cp -r "$HERMES_SRC/$subdir" "$VENDOR_DIR/"
|
||||||
|
echo " ✓ $subdir/"
|
||||||
|
else
|
||||||
|
echo " ⚠ $subdir/ 不存在,跳过"
|
||||||
|
fi
|
||||||
|
done
|
||||||
|
|
||||||
# __init__.py
|
# __init__.py
|
||||||
touch "$VENDOR_DIR/__init__.py"
|
touch "$VENDOR_DIR/__init__.py"
|
||||||
|
|
||||||
# 版本锁定标记
|
# 版本锁定标记
|
||||||
echo "$COMMIT" > "$VENDOR_DIR/HERMES_COMMIT"
|
cat > "$VENDOR_DIR/VENDOR_COMMIT" << MARKER
|
||||||
|
# MindOS CLI Vendor Snapshot
|
||||||
|
source: hermes
|
||||||
|
commit: $COMMIT
|
||||||
|
snapshot_date: $(date +%Y-%m-%d)
|
||||||
|
snapshot_by: vendor_hermes.sh
|
||||||
|
MARKER
|
||||||
|
|
||||||
# 统计
|
# 统计
|
||||||
FILE_COUNT=$(find "$VENDOR_DIR" -name "*.py" | wc -l | tr -d ' ')
|
FILE_COUNT=$(find "$VENDOR_DIR" -name "*.py" | wc -l | tr -d ' ')
|
||||||
|
|||||||
Reference in New Issue
Block a user