3 Commits

Author SHA1 Message Date
m4 7b0ebaf1a4 Bump version to 0.1.20 2026-07-13 09:24:13 +08:00
m4 0827f21a7a Add release script and commit workflow docs 2026-07-13 09:23:16 +08:00
m4 c2743251e9 Initial commit of EvoScientist framework
Self-evolving AI scientist framework built on LangGraph/LangChain with
CLI/TUI core, FastAPI gateway, and Next.js frontend.

Co-Authored-By: Claude Opus 4 <noreply@anthropic.com>
2026-07-13 08:07:45 +08:00
422 changed files with 45051 additions and 81380 deletions
-65
View File
@@ -1,65 +0,0 @@
# VCS
.git
.gitignore
.gitattributes
# CI / project meta (image doesn't need these)
.github/
docs/
CONTRIBUTING.md
README.zh-CN.md
LICENSE
# Editor / tooling state
.vscode/
.idea/
.cursor/
.codex/
.claude/
.agents/
.cursorrules
.ruff_cache/
# Python build artifacts and caches
__pycache__/
*.py[cod]
*.egg-info/
*.egg
build/
dist/
.venv/
venv/
.pytest_cache/
.ruff_cache/
.coverage
# Tests aren't needed at runtime
tests/
# Notebooks
*.ipynb
.ipynb_checkpoints/
# Local runtime data (must never leak into the image)
.env
.env.*
!.env.example
runs/
workspace/
skills/
memory/
memories/
media/
conversation_history/
.deno_cache/
.langgraph_api/
large_tool_results/
*.log
botpy.log
# Docker outputs themselves
Dockerfile.*
docker-compose*.override.yml
# OS
.DS_Store
+22 -27
View File
@@ -1,32 +1,27 @@
# EvoScientist — cp .env.example .env && fill in your keys
# EvoScientist CLI environment variables
# The preferred configuration flow is `evosci onboard`, which writes
# ~/.evoscientist/config/settings.yaml. Environment variables can override it.
# LLM provider (pick at least one)
ANTHROPIC_API_KEY= # console.anthropic.com
OPENAI_API_KEY= # platform.openai.com
GOOGLE_API_KEY= # aistudio.google.com/api-keys
NVIDIA_API_KEY= # build.nvidia.com
# Optional application directories
# EVOSCIENTIST_HOME=~/.evoscientist
# EVOSCIENTIST_DATA_ROOT=~/.evoscientist/data
# Direct providers (optional)
MINIMAX_API_KEY= # platform.minimaxi.com (China, default) or platform.minimax.io (Global)
MINIMAX_BASE_URL= # https://api.minimaxi.com/anthropic (default) or https://api.minimax.io/anthropic
ZHIPU_API_KEY= # open.bigmodel.cn (智谱)
VOLCENGINE_API_KEY= # volcengine.com (火山引擎)
DASHSCOPE_API_KEY= # dashscope.aliyuncs.com (阿里云)
MOONSHOT_API_KEY= # platform.moonshot.cn (月之暗面)
KIMI_API_KEY= # kimi.com/code (Kimi 代码计划)
# Logging
EVOSCIENTIST_LOG_LEVEL=INFO
# EVOSCIENTIST_LOG_DIR=~/.evoscientist/data/logs
EVOSCIENTIST_LOG_RETENTION_DAYS=30
# Aggregator platforms (optional)
SILICONFLOW_API_KEY= # siliconflow.cn
OPENROUTER_API_KEY= # openrouter.ai
# Optional PostgreSQL checkpoint storage. Without this value the CLI uses an
# in-memory checkpointer for the current process.
# EVOSCIENTIST_SESSION_DB_URL=postgresql://user:password@localhost:5432/evoscientist
# Custom endpoints (optional)
CUSTOM_OPENAI_API_KEY= # OpenAI-compatible endpoint
CUSTOM_OPENAI_BASE_URL=
CUSTOM_ANTHROPIC_API_KEY= # Anthropic-compatible endpoint
CUSTOM_ANTHROPIC_BASE_URL=
# Model provider examples. Provider-specific settings can also be configured
# through `evosci onboard`.
# OPENAI_API_KEY=sk-...
# OPENAI_BASE_URL=https://api.openai.com/v1
# ANTHROPIC_API_KEY=sk-ant-...
# GOOGLE_API_KEY=...
# TAVILY_API_KEY=tvly-...
# Local models (optional)
OLLAMA_BASE_URL= # http://localhost:11434 (default)
# Web search (optional)
TAVILY_API_KEY= # app.tavily.com
# Optional default model
# DEFAULT_MODEL=openai/gpt-5.4
Binary file not shown.

Before

Width:  |  Height:  |  Size: 234 KiB

+1 -1
View File
@@ -5,5 +5,5 @@
<rect x="54" y="5" width="62" height="24" rx="6" fill="#1565c0"/>
<text x="85" y="22" text-anchor="middle"
font-family="Inter, -apple-system, system-ui, sans-serif"
font-size="13" font-weight="700" fill="#ffffff">v0.2.2</text>
font-size="13" font-weight="700" fill="#ffffff">v0.0.7</text>
</svg>

Before

Width:  |  Height:  |  Size: 555 B

After

Width:  |  Height:  |  Size: 555 B

+1 -1
View File
@@ -5,5 +5,5 @@
<rect x="54" y="5" width="62" height="24" rx="6" fill="#2563eb"/>
<text x="85" y="22" text-anchor="middle"
font-family="Inter, -apple-system, system-ui, sans-serif"
font-size="13" font-weight="700" fill="#ffffff">v0.2.2</text>
font-size="13" font-weight="700" fill="#ffffff">v0.0.7</text>
</svg>

Before

Width:  |  Height:  |  Size: 555 B

After

Width:  |  Height:  |  Size: 555 B

Binary file not shown.

After

Width:  |  Height:  |  Size: 654 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 213 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 287 KiB

After

Width:  |  Height:  |  Size: 428 KiB

-20
View File
@@ -1,20 +0,0 @@
version: 2
updates:
# Base images in Dockerfile (BASE_IMAGE / NODE_IMAGE ARG defaults).
# Dependabot reads `FROM`, `COPY --from=`, and ARG-bound base refs, and
# bumps both the @sha256 digest and the trailing # vX.Y.Z comment.
- package-ecosystem: "docker"
directory: "/"
schedule:
interval: "weekly"
open-pull-requests-limit: 5
commit-message:
prefix: "chore(docker)"
labels:
- "dependencies"
- "docker"
# Single PR per cadence rather than one per image — keeps reviewer load low
# and lets us validate trixie/uv/node bumps as a coherent set.
groups:
base-images:
patterns: ["*"]
-67
View File
@@ -1,67 +0,0 @@
name: Docker
on:
push:
branches: ["main"]
tags: ["v*"]
pull_request:
paths:
- "Dockerfile"
- ".dockerignore"
- "pyproject.toml"
- "uv.lock"
- "EvoScientist/**"
- ".github/workflows/docker.yml"
workflow_dispatch:
concurrency:
group: docker-${{ github.ref }}
cancel-in-progress: true
env:
REGISTRY: ghcr.io
IMAGE_NAME: ${{ github.repository }}
jobs:
build:
runs-on: ubuntu-latest
timeout-minutes: 45
permissions:
contents: read
packages: write
steps:
- uses: actions/checkout@93cb6efe18208431cddfb8368fd83d5badbf9bfd # v5.0.1
- uses: docker/setup-qemu-action@c7c53464625b32c7a7e944ae62b3e17d2b600130 # v3.7.0
- uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12.0
- name: Log in to ${{ env.REGISTRY }}
if: github.event_name != 'pull_request'
uses: docker/login-action@c94ce9fb468520275223c153574b00df6fe4bcc9 # v3.7.0
with:
registry: ${{ env.REGISTRY }}
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Extract image metadata
id: meta
uses: docker/metadata-action@c299e40c65443455700f0fdfc63efafe5b349051 # v5.10.0
with:
images: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}
tags: |
type=ref,event=branch
type=ref,event=pr
type=semver,pattern={{version}}
type=semver,pattern={{major}}.{{minor}}
type=raw,value=latest,enable={{is_default_branch}}
- name: Build and push
uses: docker/build-push-action@10e90e3645eae34f1e60eeb005ba3a3d33f178e8 # v6.19.2
with:
context: .
platforms: linux/amd64,linux/arm64
push: ${{ github.event_name != 'pull_request' }}
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
cache-from: type=gha
cache-to: type=gha,mode=max
+1 -8
View File
@@ -7,18 +7,11 @@ on:
jobs:
pytest:
runs-on: ubuntu-latest
timeout-minutes: 15
# ``fail-fast: false`` so a single failing (os, python-version) cell
# doesn't cancel the rest of the matrix. Useful while the Windows
# leg is being brought up — we want to see all four cell results
# in one CI run instead of playing whack-a-mole one failure at a
# time. See #207.
strategy:
fail-fast: false
matrix:
os: [ubuntu-latest, windows-latest]
python-version: ["3.11", "3.12"]
runs-on: ${{ matrix.os }}
steps:
- uses: actions/checkout@v5
- uses: astral-sh/setup-uv@v6
+20 -3
View File
@@ -9,13 +9,13 @@ dist/
build/
*.egg
*.pytest_cache/
.benchmarks/
.coverage
.ipynb_checkpoints/
# Environment
.env
.env.*
.env_*
!.env.example
.venv/
venv/
@@ -36,9 +36,7 @@ bridge/package-lock.json
.langgraph_api/
workspace/
skills/
!EvoScientist/skills/
memory/
!EvoScientist/memory/
media/
conversation_history/
.deno_cache/
@@ -48,3 +46,22 @@ conversation_history/
*meals/
botpy.log
large_tool_results/
# Docker runtime data
docker/data/
# Project-level config data
.data/
# Sensitive / credentials (never commit)
postgresql:*
_s3_backup/
# Local scratch / tooling data
.superpowers/
.test-home/
tmp/
# Root-level debug scratch scripts (proper tests live in tests/)
/test_*.py
/research_lookup_temp.py
+2 -2
View File
@@ -1,10 +1,10 @@
repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
# Ruff version.
rev: v0.15.17
rev: v0.9.9
hooks:
# Run the linter.
- id: ruff-check
- id: ruff
args: [ --fix ]
# Run the formatter.
- id: ruff-format
+1 -1
View File
@@ -69,7 +69,7 @@ EvoScientist is a multi-agent AI system for automated scientific experimentation
| Framework | [DeepAgents](https://github.com/langchain-ai/deepagents) + [LangChain](https://python.langchain.com/) + [LangGraph](https://langchain-ai.github.io/langgraph/) |
| Default model | `claude-sonnet-4-6` (Anthropic) |
| Tests | ~890 across 36 files, no API keys needed |
| Config file | `~/.config/evoscientist/config.yaml` |
| Config file | `.data/.config/settings.yaml` (project root) |
### Sub-Agents (defined in `EvoScientist/subagent.yaml`)
-75
View File
@@ -1,75 +0,0 @@
# syntax=docker/dockerfile:1.7
ARG BASE_IMAGE=ghcr.io/astral-sh/uv:python3.11-trixie-slim@sha256:7936cc6625ca04cafa6ecc3c2881ddfe90a747c55c74480cd4ac6ffad6a5af1e
ARG NODE_IMAGE=node:24-trixie-slim@sha256:735dd688da64d22ebd9dd374b3e7e5a874635668fd2a6ec20ca1f99264294086
FROM ${NODE_IMAGE} AS nodejs
# ---------- Builder ----------
FROM ${BASE_IMAGE} AS builder
ENV UV_COMPILE_BYTECODE=1 \
UV_LINK_MODE=copy \
UV_PYTHON_DOWNLOADS=never \
UV_PROJECT_ENVIRONMENT=/opt/venv
WORKDIR /src
COPY pyproject.toml uv.lock README.md ./
RUN --mount=type=cache,target=/root/.cache/uv \
uv sync --frozen --no-install-project --no-dev \
--extra all-channels
COPY EvoScientist ./EvoScientist
RUN --mount=type=cache,target=/root/.cache/uv \
uv sync --frozen --no-dev --no-editable \
--extra all-channels
# ---------- Runtime ----------
FROM ${BASE_IMAGE} AS runtime
RUN apt-get update \
&& apt-get install -y --no-install-recommends \
git \
ca-certificates \
tini \
curl \
&& rm -rf /var/lib/apt/lists/*
COPY --from=nodejs /usr/local/bin/node /usr/local/bin/node
COPY --from=nodejs /usr/local/lib/node_modules /usr/local/lib/node_modules
RUN ln -sf /usr/local/lib/node_modules/npm/bin/npm-cli.js /usr/local/bin/npm \
&& ln -sf /usr/local/lib/node_modules/npm/bin/npx-cli.js /usr/local/bin/npx
ARG UID=1000
ARG GID=1000
RUN groupadd --gid ${GID} evosci \
&& useradd --uid ${UID} --gid ${GID} --create-home --shell /bin/bash evosci
COPY --from=builder /opt/venv /opt/venv
ENV PATH="/opt/venv/bin:/home/evosci/.evoscientist/.local/bin:${PATH}" \
PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1 \
EVOSCIENTIST_WORKSPACE_DIR=/workspace \
EVOSCIENTIST_DATA_DIR=/home/evosci/.evoscientist \
XDG_CONFIG_HOME=/home/evosci/.evoscientist/.config \
UV_TOOL_DIR=/home/evosci/.evoscientist/.local/share/uv/tools \
UV_TOOL_BIN_DIR=/home/evosci/.evoscientist/.local/bin
RUN mkdir -p /workspace \
/home/evosci/.evoscientist/.config/evoscientist \
/home/evosci/.evoscientist/.local/bin \
/home/evosci/.evoscientist/.local/share/uv/tools \
&& chown -R ${UID}:${GID} /workspace /home/evosci
USER evosci
WORKDIR /workspace
LABEL org.opencontainers.image.title="EvoScientist" \
org.opencontainers.image.description="EvoScientist agent with core + all-channels dependencies pre-installed." \
org.opencontainers.image.source="https://github.com/EvoScientist/EvoScientist" \
org.opencontainers.image.documentation="https://github.com/EvoScientist/EvoScientist#-docker" \
org.opencontainers.image.licenses="Apache-2.0"
ENTRYPOINT ["tini", "--", "evosci"]
File diff suppressed because it is too large Load Diff
+9 -3
View File
@@ -9,13 +9,14 @@ from __future__ import annotations
from importlib import import_module
from ._version import __version__
_EXPORTS: dict[str, tuple[str, str]] = {
# Agent graph (lazy to avoid expensive initialization at import time)
"EvoScientist_agent": (".EvoScientist", "EvoScientist_agent"),
"create_cli_agent": (".EvoScientist", "create_cli_agent"),
# Backends
"CustomSandboxBackend": (".backends", "CustomSandboxBackend"),
"MemoryFilesystemBackend": (".backends", "MemoryFilesystemBackend"),
"ReadOnlyFilesystemBackend": (".backends", "ReadOnlyFilesystemBackend"),
# Configuration
"EvoScientistConfig": (".config", "EvoScientistConfig"),
@@ -25,19 +26,24 @@ _EXPORTS: dict[str, tuple[str, str]] = {
"get_config_path": (".config", "get_config_path"),
# LLM
"get_chat_model": (".llm", "get_chat_model"),
"MODELS": (".llm", "MODELS"),
"list_models": (".llm", "list_models"),
"DEFAULT_MODEL": (".llm", "DEFAULT_MODEL"),
# Prompts
"get_system_prompt": (".prompts", "get_system_prompt"),
"RESEARCHER_INSTRUCTIONS": (".prompts", "RESEARCHER_INSTRUCTIONS"),
# Tools
"tavily_search": (".tools", "tavily_search"),
"web_search": (".tools", "web_search"),
"web_extract": (".tools", "web_extract"),
"web_crawl": (".tools", "web_crawl"),
"think_tool": (".tools", "think_tool"),
# Sessions
"get_checkpointer": (".sessions", "get_checkpointer"),
"generate_thread_id": (".sessions", "generate_thread_id"),
"list_threads": (".sessions", "list_threads"),
"delete_thread": (".sessions", "delete_thread"),
"get_storage_stats": (".sessions", "get_storage_stats"),
"get_aggregated_storage_stats": (".sessions", "get_aggregated_storage_stats"),
"list_all_session_db_paths": (".sessions", "list_all_session_db_paths"),
}
+1
View File
@@ -0,0 +1 @@
__version__ = "0.1.20"
-50
View File
@@ -1,50 +0,0 @@
"""Windows asyncio event-loop policy compatibility.
On Windows, MCP stdio servers are launched as subprocesses by the MCP SDK's
stdio transport, which uses ``anyio.open_process`` → ``asyncio`` async
subprocess support. ``asyncio``'s *Selector* event loop does **not** implement
async subprocess creation, so on a Selector loop the stdio transport falls back
to a synchronous ``subprocess.Popen`` inside an ``async`` function. Under
``langgraph dev`` (which enables ``blockbuster`` by default to police blocking
I/O) that synchronous call is flagged as a ``BlockingError`` — see issue #283.
The *Proactor* loop supports async subprocesses natively, so the fallback never
happens and ``blockbuster`` allows the (now genuinely async) spawn.
Windows has defaulted to ``WindowsProactorEventLoopPolicy`` since Python 3.8, so
this is normally a no-op. We set it explicitly anyway as a safeguard: a
dependency, IDE, or notebook host may have installed a Selector policy earlier
in the process, and the ``langgraph dev`` subprocess in particular runs code we
don't fully control. Calling this at each process entrypoint — **before any
event loop is created** — guarantees the MCP subprocess path stays async.
This must run at import/startup time, ahead of the first ``asyncio.run`` /
``new_event_loop`` call; once a loop exists, swapping the policy does not change
the already-running loop.
"""
from __future__ import annotations
import sys
def ensure_proactor_event_loop_policy() -> bool:
"""Install ``WindowsProactorEventLoopPolicy`` on Windows if needed.
Returns ``True`` if a Proactor policy is in effect afterwards (always
``False`` off Windows, where the concept doesn't apply). Safe and idempotent
to call multiple times; a no-op on non-Windows platforms.
"""
if sys.platform != "win32":
return False
import asyncio
proactor_policy = getattr(asyncio, "WindowsProactorEventLoopPolicy", None)
if proactor_policy is None: # pragma: no cover - non-Windows / stripped build
return False
current = asyncio.get_event_loop_policy()
if not isinstance(current, proactor_policy):
asyncio.set_event_loop_policy(proactor_policy())
return True
+636 -856
View File
File diff suppressed because it is too large Load Diff
-334
View File
@@ -1,334 +0,0 @@
"""Background OS-process execution for the sandbox.
A *process* here is a single detached OS process launched via ``run_in_background``
(distinct from an async sub-agent *task* and a future cron *schedule* — the word
"job" is intentionally never used).
The registry is **module-global (process-level)**: processes survive ``/new`` and
``/resume`` within the same CLI process, but are not persisted across a CLI restart.
The live ``Popen`` handle is held so ``poll()`` / ``returncode`` stay authoritative
(no PID-reuse risk).
Command validation and cwd resolution happen at the tool layer
(``middleware/background.py``); this module is the pure execution + tracking mechanism
and is safe to unit-test on its own. A future scheduler (cron) would reuse ``launch``.
"""
from __future__ import annotations
import logging
import os
import signal
import subprocess
import threading
import time
import uuid
from collections.abc import Callable
from dataclasses import dataclass, field
from datetime import UTC, datetime
from pathlib import Path
import psutil
logger = logging.getLogger(__name__)
_BG_DIRNAME = ".bg_processes"
_KILL_GRACE_SECONDS = 2.0
@dataclass
class BgProcess:
"""A tracked background OS process."""
process_id: str
name: str
command: str
popen: subprocess.Popen
pid: int
log_path: Path
started_at: str # ISO-8601 UTC (record/display)
started_ts: float # epoch seconds (elapsed computation)
origin_thread_id: str | None = None # CLI thread/session that launched it
returncode: int | None = None
finished_at: str | None = None
finished_ts: float | None = None # epoch at exit; freezes elapsed once done
stopped: bool = False # set by stop(); suppresses the completion notification
# epoch each thread last checked this process (status/list); keyed by thread_id
# so a check from one session can't dedup another session's completion ping.
last_checked_by_thread: dict[str | None, float] = field(default_factory=dict)
_PROCESSES: dict[str, BgProcess] = {}
_LOCK = threading.Lock()
def _now_iso() -> str:
return datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ")
def _record_exit(proc: BgProcess) -> None:
"""Record terminal state on first observed exit. Caller MUST hold ``_LOCK``.
``finished_ts`` is set when the exit is first observed. The per-process daemon
watcher (:func:`_watch`) calls this right after ``popen.wait()`` returns, so in
practice ``finished_ts`` ≈ the real exit time. Calls from ``status`` / ``list_all`` /
``stop`` are a fallback for the brief window before the watcher runs.
"""
rc = proc.popen.poll()
if rc is not None and proc.returncode is None:
proc.returncode = rc
proc.finished_at = _now_iso()
proc.finished_ts = time.time()
def _elapsed(proc: BgProcess) -> int:
"""Seconds the process has run — frozen at first-observed exit once it has exited."""
end = proc.finished_ts if proc.finished_ts is not None else time.time()
return int(end - proc.started_ts)
def was_observed_done(process_id: str, origin_thread_id: str | None = None) -> bool:
"""True if ``origin_thread_id`` already saw this process's completion itself.
i.e. the process has exited AND was checked (``status``/``list_all``) from that thread
at or after it finished. Used to dedup the completion notification (routed to the
launching thread), so a check from a *different* session can't suppress it.
"""
with _LOCK:
proc = _PROCESSES.get(process_id)
if proc is None or proc.finished_ts is None:
return False
seen_ts = proc.last_checked_by_thread.get(origin_thread_id)
return seen_ts is not None and seen_ts >= proc.finished_ts
def _read_tail(log_path: Path, tail_bytes: int) -> str:
# Seek from the end so a huge log isn't fully read into memory on each status check.
try:
with log_path.open("rb") as f:
f.seek(0, os.SEEK_END)
size = f.tell()
if size == 0:
return "(no output yet)"
if size > tail_bytes:
f.seek(-tail_bytes, os.SEEK_END)
return "...(truncated)...\n" + f.read().decode("utf-8", "replace")
f.seek(0)
data = f.read()
except OSError:
return "(no output captured yet)"
return data.decode("utf-8", "replace")
def _watch(proc: BgProcess, on_exit: Callable[[BgProcess], None] | None) -> None:
"""Block until ``proc`` exits, record the exit promptly, then fire ``on_exit``.
Running in a daemon thread, ``popen.wait()`` lets us record ``finished_ts`` at (very
close to) the real exit time — fixing the observation-time inflation — and gives a
hook the CLI layer wires to a completion notification, without ``background.py``
importing the notifier (kept decoupled via the callback).
"""
try:
proc.popen.wait()
except Exception:
pass
with _LOCK:
_record_exit(proc)
if on_exit is not None:
try:
on_exit(proc)
except Exception:
logger.warning("background on_exit callback failed", exc_info=True)
def launch(
command: str,
cwd: str,
name: str | None = None,
*,
origin_thread_id: str | None = None,
on_exit: Callable[[BgProcess], None] | None = None,
) -> str:
"""Launch ``command`` detached in ``cwd``; return a short ``process_id``.
The command is run via ``shell=True`` with output redirected to a per-process log
file under ``<cwd>/.bg_processes/`` and ``start_new_session=True`` so the child is a
process-group leader (survives this call's return and can be killed as a group).
The caller is responsible for validating ``command`` first.
``origin_thread_id`` records the launching CLI session so ``list_all`` can scope to it.
``on_exit`` (optional) is called with the ``BgProcess`` from a daemon watcher thread
once the process exits — used by the CLI layer to emit a completion notification.
"""
process_id = uuid.uuid4().hex[:8]
log_dir = Path(cwd) / _BG_DIRNAME
log_dir.mkdir(parents=True, exist_ok=True)
log_path = log_dir / f"{process_id}.log"
log_file = open(log_path, "w")
try:
popen = subprocess.Popen(
command,
shell=True,
cwd=cwd,
stdout=log_file,
stderr=subprocess.STDOUT,
stdin=subprocess.DEVNULL,
start_new_session=True,
)
finally:
# The child inherited its own dup of the fd during spawn; the parent's copy
# is no longer needed (and must be closed so the pipe/file isn't held open).
log_file.close()
proc = BgProcess(
process_id=process_id,
name=name or command[:40],
command=command,
popen=popen,
pid=popen.pid,
log_path=log_path,
started_at=_now_iso(),
started_ts=time.time(),
origin_thread_id=origin_thread_id,
)
with _LOCK:
_PROCESSES[process_id] = proc
# Daemon watcher: records the precise exit time and fires on_exit when done.
threading.Thread(target=_watch, args=(proc, on_exit), daemon=True).start()
return process_id
def status(
process_id: str, *, thread_id: str | None = None, tail_bytes: int = 16_000
) -> str:
"""Return a human-readable status + recent output tail for ``process_id``."""
with _LOCK:
proc = _PROCESSES.get(process_id)
if proc is None:
return (
f"No such background process: {process_id!r}. "
"Use list_processes to see tracked processes."
)
_record_exit(proc)
proc.last_checked_by_thread[thread_id] = time.time() # this thread observed it
running = proc.returncode is None
elapsed = _elapsed(proc)
name, pid, command, returncode, log_path = (
proc.name,
proc.pid,
proc.command,
proc.returncode,
proc.log_path,
)
if running:
head = f"Process {process_id} (name={name!r}) RUNNING — {elapsed}s elapsed, pid {pid}."
else:
head = f"Process {process_id} (name={name!r}) EXITED code {returncode} after ~{elapsed}s."
tail = _read_tail(log_path, tail_bytes) # file IO outside the lock
return (
f"{head}\nCommand: {command}\n--- output (last {tail_bytes} bytes) ---\n{tail}"
)
def _kill_process_tree(popen: subprocess.Popen, *, forceful: bool) -> None:
"""Kill the process group/tree in a cross-platform way.
On POSIX ``start_new_session=True`` makes the child a process-group
leader; ``os.killpg`` terminates the entire group (shell + any
grandchildren). On Windows ``TerminateProcess`` (used by
``Popen.terminate()`` / ``Popen.kill()``) only kills the direct
child — it does *not* cascade to grandchildren. We use ``psutil``
to walk the process tree and signal every descendant.
"""
if os.name == "nt":
try:
proc = psutil.Process(popen.pid)
targets = [proc, *proc.children(recursive=True)]
except (psutil.NoSuchProcess, psutil.AccessDenied):
return
for p in targets:
try:
if forceful:
p.kill()
else:
p.terminate()
except (psutil.NoSuchProcess, psutil.AccessDenied):
pass
else:
sig = signal.SIGKILL if forceful else signal.SIGTERM
try:
os.killpg(os.getpgid(popen.pid), sig)
except ProcessLookupError:
pass
def stop(process_id: str) -> str:
"""Terminate ``process_id`` and its process group (SIGTERM, then SIGKILL)."""
with _LOCK:
proc = _PROCESSES.get(process_id)
if proc is None:
return f"No such background process: {process_id!r}."
if proc.popen.poll() is not None:
_record_exit(proc)
return f"Process {process_id} already finished (code {proc.returncode})."
# Mark as user-stopped so the watcher's on_exit suppresses the completion
# notification (the user already knows — no need to ping them).
proc.stopped = True
# The watcher's popen.wait() reaps without the lock, so a tiny PID-reuse race
# remains (getpgid on a recycled pid). On POSIX ProcessLookupError covers the
# common case; on Windows ``Popen.terminate()`` is a no-op on a dead handle
# so we poll after the call instead.
_kill_process_tree(proc.popen, forceful=False)
if proc.popen.poll() is not None:
_record_exit(proc)
return f"Process {process_id} is no longer running."
deadline = time.time() + _KILL_GRACE_SECONDS
while time.time() < deadline:
with _LOCK:
if proc.popen.poll() is not None:
_record_exit(proc)
break
time.sleep(0.1)
else:
with _LOCK:
if proc.popen.poll() is None:
_kill_process_tree(proc.popen, forceful=True)
_record_exit(proc)
with _LOCK:
_record_exit(proc)
name = proc.name
return f"Stopped background process {process_id} (name={name!r})."
def list_all(thread_id: str | None = None, *, include_all: bool = False) -> str:
"""List tracked background processes with live statuses.
Scoped to the launching session (``thread_id``) unless ``include_all`` is set.
"""
with _LOCK:
all_procs = list(_PROCESSES.values())
procs = (
all_procs
if include_all
else [p for p in all_procs if p.origin_thread_id == thread_id]
)
if not procs:
if all_procs and not include_all:
return (
"No background processes in this session "
f"({len(all_procs)} in other sessions — pass all_threads=True to see them)."
)
return "No background processes tracked."
lines = []
now = time.time()
for p in procs:
_record_exit(p)
p.last_checked_by_thread[thread_id] = now # this thread observed it
state = "RUNNING" if p.returncode is None else f"exited({p.returncode})"
lines.append(
f" {p.process_id} {state:12} {_elapsed(p)}s name={p.name!r}"
)
return f"{len(procs)} background process(es):\n" + "\n".join(lines)
+1 -1
View File
@@ -2,7 +2,7 @@
EvoScientist provides unified integration with 10 messaging platforms. This document covers the architecture overview, message processing pipeline, capability matrix, security model, deployment guides, and troubleshooting.
Configuration file: `~/.config/evoscientist/config.yaml` (or use environment variables with the `EVOSCIENTIST_` prefix).
Configuration file: `~/.config/ai4scientist/settings.yaml` (or use environment variables with the `EVOSCIENTIST_` prefix).
## Table of Contents
+59 -2
View File
@@ -2,20 +2,23 @@
Channels push messages to the inbound queue; the agent (or any consumer)
reads from inbound, processes, and pushes responses to the outbound queue.
``ChannelManager._dispatch_outbound`` routes outbound messages to the
correct channel by looking up its registered :class:`Channel` instance.
A background dispatcher routes outbound messages to the correct channel
via subscriber callbacks.
Deduplication is handled at the Channel level (single dedup point).
"""
import asyncio
import logging
from collections.abc import Awaitable, Callable
from ..debug import TraceMixin, debug_trace_enabled
from .events import InboundMessage, OutboundMessage
logger = logging.getLogger(__name__)
OutboundCallback = Callable[[OutboundMessage], Awaitable[None]]
class MessageBus(TraceMixin):
"""Async message bus that decouples chat channels from the agent core."""
@@ -25,6 +28,8 @@ class MessageBus(TraceMixin):
def __init__(self):
self.inbound: asyncio.Queue[InboundMessage] = asyncio.Queue(maxsize=5000)
self.outbound: asyncio.Queue[OutboundMessage] = asyncio.Queue(maxsize=5000)
self._outbound_subscribers: dict[str, list[OutboundCallback]] = {}
self._running = False
self._debug_trace = debug_trace_enabled()
self._trace_logger = logger
@@ -48,6 +53,58 @@ class MessageBus(TraceMixin):
"""Consume the next outbound message (blocks until available)."""
return await self.outbound.get()
# ── subscriber routing ──
def subscribe_outbound(
self,
channel: str,
callback: OutboundCallback,
) -> None:
"""Register a callback for outbound messages targeting *channel*."""
if channel not in self._outbound_subscribers:
self._outbound_subscribers[channel] = []
self._outbound_subscribers[channel].append(callback)
async def dispatch_outbound(self) -> None:
"""Route outbound messages to subscribed channels.
Run as a background task — loops until :meth:`stop` is called.
"""
self._running = True
while self._running:
try:
msg = await asyncio.wait_for(
self.outbound.get(),
timeout=1.0,
)
except TimeoutError:
continue
subscribers = self._outbound_subscribers.get(msg.channel, [])
if not subscribers:
self._trace_event(
"bus_dispatch_drop",
target_channel=msg.channel,
reason="no_subscriber",
chat_id=msg.chat_id,
)
logger.warning(f"No subscriber for channel: {msg.channel}")
continue
for callback in subscribers:
try:
await callback(msg)
except Exception as e:
self._trace_event(
"bus_dispatch_error",
target_channel=msg.channel,
chat_id=msg.chat_id,
error_type=type(e).__name__,
)
logger.error(f"Error dispatching to {msg.channel}: {e}")
def stop(self) -> None:
"""Stop the dispatcher loop."""
self._running = False
@property
def inbound_size(self) -> int:
return self.inbound.qsize()
-1
View File
@@ -160,7 +160,6 @@ QQ = ChannelCapabilities(
format_type="plain",
max_text_length=4096,
typing=False, # no typing API for QQ bots
inline_buttons=True, # markdown + keyboard payload (C2C only)
media_send=True,
media_receive=True,
voice=False, # qq-botpy does not expose voice as a distinct message type
+96 -135
View File
@@ -11,12 +11,13 @@ from __future__ import annotations
import asyncio
import logging
import time
import uuid
from collections import OrderedDict
from collections.abc import AsyncIterator, Callable
from dataclasses import dataclass
from typing import Any, TypeVar
from ..gateway import GraphGateway, GraphRunInput, GraphTarget, RunRequest
from .base import Channel
from .bus import MessageBus
from .bus.events import InboundMessage, OutboundMessage
@@ -118,7 +119,7 @@ def _should_auto_approve(action_requests: list[dict]) -> bool:
return True
try:
from ..config.settings import HITL_SHELL_TOOLS, load_config
from ..config.settings import load_config
cfg = load_config()
except Exception:
@@ -134,10 +135,14 @@ def _should_auto_approve(action_requests: list[dict]) -> bool:
)
for req in action_requests:
name = req.get("name", "")
if name not in HITL_SHELL_TOOLS:
name = (
req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
)
if name != "execute":
continue
args = req.get("args", {})
args = (
req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
)
command = args.get("command", "") if isinstance(args, dict) else ""
cmd = command.strip()
if not any(cmd.startswith(prefix) for prefix in shell_allow_list):
@@ -145,18 +150,16 @@ def _should_auto_approve(action_requests: list[dict]) -> bool:
return True
def _format_approval_prompt(
action_requests: list[dict], *, with_buttons: bool = False
) -> str:
"""Format an approval prompt as a text message for channel users.
When *with_buttons* is True, the trailing "Reply: 1=Approve..."
instruction is dropped — the buttons replace the textual cue.
"""
def _format_approval_prompt(action_requests: list[dict]) -> str:
"""Format an approval prompt as a text message for channel users."""
lines = ["\u26a0\ufe0f Approval Required\n"]
for i, req in enumerate(action_requests, 1):
name = req.get("name", "")
args = req.get("args", {})
name = (
req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
)
args = (
req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
)
if isinstance(args, dict):
command = args.get("command", args.get("path", ""))
else:
@@ -165,10 +168,9 @@ def _format_approval_prompt(
lines.append(f" {i}. {name}: {command}")
else:
lines.append(f" {i}. {name}")
if not with_buttons:
lines.append("")
lines.append("Reply: 1=Approve, 2=Reject, 3=Approve all")
lines.append("(Auto-reject in 2 min if no reply)")
lines.append("")
lines.append("Reply: 1=Approve, 2=Reject, 3=Approve all")
lines.append("(Auto-reject in 2 min if no reply)")
return "\n".join(lines)
@@ -187,25 +189,6 @@ def _parse_approval_reply(text: str) -> str | None:
return None
def _approval_prompt_metadata(
base_metadata: dict | None, *, with_buttons: bool
) -> dict:
"""Outbound metadata for the HITL approval prompt.
When *with_buttons* is True, attaches Approve/Reject/Auto buttons whose
values match ``_parse_approval_reply`` so a click flows through the same
path as a typed ``"1"``/``"2"``/``"3"`` reply.
"""
metadata = dict(base_metadata or {})
if with_buttons:
metadata["buttons"] = [
{"text": "Approve", "value": "1", "type": "primary"},
{"text": "Reject", "value": "2", "type": "danger"},
{"text": "Approve all", "value": "3"},
]
return metadata
@dataclass
class _PendingInterrupt:
"""Stored state for a pending HITL interrupt awaiting channel user reply."""
@@ -234,11 +217,9 @@ class InboundConsumer:
manager:
The ChannelManager (used to look up channel instances).
agent:
The local agent object used by local graph gateway targets.
The agent object (must support ``stream_agent_events``).
thread_id:
Default thread ID for agent conversations.
graph_gateway:
Gateway used for thread creation and graph streaming.
send_thinking:
Whether to forward thinking messages to the channel.
on_message_received:
@@ -269,7 +250,6 @@ class InboundConsumer:
agent: Any,
thread_id: str,
*,
graph_gateway: GraphGateway,
send_thinking: bool = False,
on_message_received: Callable[[InboundMessage], None] | None = None,
on_streaming_event: Callable[[dict], None] | None = None,
@@ -283,7 +263,6 @@ class InboundConsumer:
self.manager = manager
self.agent = agent
self.thread_id = thread_id
self.graph_gateway = graph_gateway
self.send_thinking = send_thinking
self._on_message_received = on_message_received
self._on_streaming_event = on_streaming_event
@@ -317,7 +296,7 @@ class InboundConsumer:
# ask_user: pending reply per session_key
self._pending_ask_user_replies: dict[str, _PendingAskUserReply] = {}
async def _get_thread_id(self, sender_id: str) -> str:
def _get_thread_id(self, sender_id: str) -> str:
"""Get or create a thread ID for the given sender.
Uses LRU ordering: recently accessed senders are moved to the
@@ -333,9 +312,7 @@ class InboundConsumer:
if self.thread_id:
self._sessions[sender_id] = f"{self.thread_id}:{sender_id}"
else:
self._sessions[sender_id] = await self.graph_gateway.create_thread(
GraphTarget(local_graph=self.agent)
)
self._sessions[sender_id] = str(uuid.uuid4())
return self._sessions[sender_id]
def _get_channel(self, channel_name: str) -> Channel | None:
@@ -429,7 +406,7 @@ class InboundConsumer:
pass
channel = self._get_channel(msg.channel)
thread_id = await self._get_thread_id(msg.sender_id)
thread_id = self._get_thread_id(msg.sender_id)
session_key = msg.session_key # "channel:chat_id"
# Lazily create per-chat lock; evict stale locks when too many
@@ -472,16 +449,15 @@ class InboundConsumer:
session_key: str,
) -> None:
"""Stream agent events with HITL interrupt handling."""
from langgraph.types import Command
from ..stream.events import stream_agent_events
stream_input: GraphRunInput = msg.content
stream_input: Any = msg.content
_t0 = time.monotonic()
try:
if channel:
await channel.start_typing(msg.chat_id)
_last_sent_thinking: str | None = None
for _hitl_round in range(_MAX_HITL_ROUNDS):
final_content = ""
thinking_buffer: list[str] = []
@@ -490,38 +466,14 @@ class InboundConsumer:
thinking_sent = False
interrupt_data: dict | None = None
async def _flush_thinking_buffer(
buffer: list[str] = thinking_buffer,
) -> bool:
"""Send the current thinking buffer, dedup by content."""
nonlocal thinking_sent, _last_sent_thinking
if not channel or thinking_sent or not buffer:
return False
full_thinking = "".join(buffer).rstrip()
buffer.clear()
if not full_thinking or full_thinking == _last_sent_thinking:
return False
await channel.send_thinking_message(
msg.sender_id,
full_thinking,
msg.metadata,
)
thinking_sent = True
_last_sent_thinking = full_thinking
return True
async for event in _timeout_aiter(
self.graph_gateway.stream_events(
RunRequest(
message=stream_input,
thread_id=thread_id,
media=msg.media or None
if isinstance(stream_input, str)
else None,
target=GraphTarget(local_graph=self.agent),
)
stream_agent_events(
self.agent,
stream_input,
thread_id,
media=msg.media or None
if isinstance(stream_input, str)
else None,
),
self._inference_timeout,
):
@@ -542,7 +494,16 @@ class InboundConsumer:
if event.get("name") == "write_todos" and not todo_sent:
todos = event.get("args", {}).get("todos", [])
if todos and channel:
await _flush_thinking_buffer()
if thinking_buffer and not thinking_sent:
full_thinking = "".join(thinking_buffer)
if full_thinking:
await channel.send_thinking_message(
msg.sender_id,
full_thinking,
msg.metadata,
)
thinking_sent = True
thinking_buffer.clear()
await channel.send_todo_message(
msg.sender_id,
_format_todo_list(todos),
@@ -553,11 +514,14 @@ class InboundConsumer:
elif event_type == "text":
final_content += event.get("content", "")
elif event_type == "progress":
# Internal agent planning — ignore for channel delivery.
# This text should NOT appear in the user-facing response.
pass
elif event_type == "subagent_text":
sa_name = event.get("subagent", "unknown")
instance_id = event.get("instance_id")
if not instance_id:
continue
instance_id = event.get("instance_id") or sa_name
if instance_id not in subagent_text_buffers:
subagent_text_buffers[instance_id] = (sa_name, [])
subagent_text_buffers[instance_id][1].append(
@@ -576,7 +540,14 @@ class InboundConsumer:
break # exit async for to handle ask_user
# Flush thinking
await _flush_thinking_buffer()
if thinking_buffer and not thinking_sent and channel:
full_thinking = "".join(thinking_buffer)
if full_thinking:
await channel.send_thinking_message(
msg.sender_id,
full_thinking,
msg.metadata,
)
# No interrupt — normal completion
if interrupt_data is None:
@@ -591,6 +562,13 @@ class InboundConsumer:
)
await self.bus.publish_outbound(outbound)
self._metrics.total_successes += 1
elapsed = time.monotonic() - _t0
logger.info(
"stream completed: %d chars, %.2fs, session=%s",
len(outbound.content),
elapsed,
session_key,
)
if self._on_message_sent:
try:
self._on_message_sent(outbound)
@@ -605,6 +583,7 @@ class InboundConsumer:
interrupt_data,
session_key,
)
from langgraph.types import Command # type: ignore[import-untyped]
stream_input = Command(resume=result)
continue
@@ -615,6 +594,8 @@ class InboundConsumer:
# Session auto-approve (user previously chose "Approve all")
if session_key in self._auto_approve_sessions:
from langgraph.types import Command # type: ignore[import-untyped]
stream_input = Command(
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
)
@@ -622,27 +603,21 @@ class InboundConsumer:
# Config auto-approve (auto_approve, non-execute, allow_list)
if _should_auto_approve(action_reqs):
from langgraph.types import Command # type: ignore[import-untyped]
stream_input = Command(
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
)
continue
# Needs user approval — send prompt to channel
has_buttons = (
channel is not None and channel.capabilities.inline_buttons
)
prompt_text = _format_approval_prompt(
action_reqs, with_buttons=has_buttons
)
approval_metadata = _approval_prompt_metadata(
msg.metadata, with_buttons=has_buttons
)
prompt_text = _format_approval_prompt(action_reqs)
await self.bus.publish_outbound(
OutboundMessage(
channel=msg.channel,
chat_id=msg.chat_id,
content=prompt_text,
metadata=approval_metadata,
metadata=msg.metadata,
)
)
@@ -654,61 +629,35 @@ class InboundConsumer:
)
self._pending_interrupts[session_key] = pending
timed_out = False
try:
await asyncio.wait_for(
pending.event.wait(),
timeout=_HITL_APPROVAL_TIMEOUT,
)
except TimeoutError:
timed_out = True
# Auto-approve on timeout
pending.decision = "approve"
finally:
# Unregister BEFORE any further await so a late reply can't flip
# the decision back to approve during the notification round-trip.
self._pending_interrupts.pop(session_key, None)
if timed_out:
# Reject on timeout (fail-closed; matches cli/channel.py). Decision
# is a local constant, not pending.decision, so it can't be
# overwritten by a late reply after we unregistered above.
decision = "reject"
decision = pending.decision or "approve"
if decision == "reject":
await self.bus.publish_outbound(
OutboundMessage(
channel=msg.channel,
chat_id=msg.chat_id,
content="⏰ Approval timed out. Action rejected.",
content="Tool execution rejected.",
metadata=msg.metadata,
)
)
else:
decision = pending.decision or "reject"
# Visible confirmation so the click/reply registers (QQ has no
# message recall API for C2C). Only fires when the user
# actually responded — silent on timeout to avoid claiming
# the user approved when they just walked away.
if pending.event.is_set():
feedback_text = {
"approve": "\u2705 已批准",
"auto": "\u2705 已批准(后续自动通过)",
"reject": "\u274c 已拒绝",
}.get(decision)
if feedback_text:
await self.bus.publish_outbound(
OutboundMessage(
channel=msg.channel,
chat_id=msg.chat_id,
content=feedback_text,
metadata=msg.metadata,
)
)
if decision == "reject":
return
if decision == "auto":
self._auto_approve_sessions.add(session_key)
from langgraph.types import Command # type: ignore[import-untyped]
stream_input = Command(
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
)
@@ -716,9 +665,14 @@ class InboundConsumer:
except TimeoutError:
self._metrics.total_timeouts += 1
elapsed = time.monotonic() - _t0
logger.error(
f"Inference timeout ({self._inference_timeout}s idle) "
f"for {msg.sender_id} in {session_key}"
"Inference timeout (%ds idle) for %s in %s, elapsed=%.2fs, response=%d chars",
self._inference_timeout,
msg.sender_id,
session_key,
elapsed,
len(final_content),
)
await self.bus.publish_outbound(
OutboundMessage(
@@ -731,7 +685,14 @@ class InboundConsumer:
except Exception as e:
self._metrics.total_failures += 1
logger.error(f"Agent error: {e}")
elapsed = time.monotonic() - _t0
logger.error(
"Agent error: %s | elapsed=%.2fs, response=%d chars, session=%s",
e,
elapsed,
len(final_content),
session_key,
)
await self.bus.publish_outbound(
OutboundMessage(
channel=msg.channel,
+2 -5
View File
@@ -19,15 +19,11 @@ Examples:
import argparse
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from .channel import DingTalkChannel, DingTalkConfig
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger(__name__)
@@ -72,6 +68,7 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
config = DingTalkConfig(
+2 -5
View File
@@ -19,15 +19,11 @@ Examples:
import argparse
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from .channel import DiscordChannel, DiscordConfig
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger(__name__)
@@ -73,6 +69,7 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
config = DiscordConfig(
+2 -5
View File
@@ -19,15 +19,11 @@ Examples:
import argparse
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from .channel import EmailChannel, EmailConfig
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger(__name__)
@@ -99,6 +95,7 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
config = EmailConfig(
+1 -2
View File
@@ -1,8 +1,7 @@
from ..channel_manager import _parse_csv, register_channel
from .channel import FeishuChannel, FeishuConfig
from .onboard import qr_register
__all__ = ["FeishuChannel", "FeishuConfig", "qr_register"]
__all__ = ["FeishuChannel", "FeishuConfig"]
def create_from_config(config) -> FeishuChannel:
-25
View File
@@ -376,31 +376,6 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
.build()
)
# Silently absorb events we don't have a handler for. Feishu auto-
# subscribes a PersonalAgent app to many event types (reactions,
# read receipts, recalls, member changes…) that EvoScientist doesn't
# care about. Without this wrapper, ``_do_without_validation``
# raises ``EventException("processor not found, type: ...")``,
# which lark-oapi's WS client (ws/client.py) catches and turns into
# an HTTP 500 reply on the WebSocket frame — Feishu then marks the
# event as failed and retries it. This is especially noisy because
# our own ``_send_ack_reaction`` triggers ``im.message.reaction.
# created_v1`` on every inbound message, causing a feedback loop.
from lark_oapi.core.exception import EventException
_original_dispatch = handler._do_without_validation
def _silent_dispatch(payload: bytes):
try:
return _original_dispatch(payload)
except EventException as exc:
if "processor not found" in str(exc):
logger.debug("Feishu: ignored unsubscribed event (%s)", exc)
return None
raise
handler._do_without_validation = _silent_dispatch
ws_client = lark.ws.Client(
self.config.app_id,
self.config.app_secret,
-361
View File
@@ -1,361 +0,0 @@
"""Feishu / Lark scan-to-create (QR code onboard) flow.
Drives the Feishu open-platform device-code flow at
``accounts.feishu.cn/oauth/v1/app/registration`` (and the Lark equivalent
at ``accounts.larksuite.com``). The user scans a terminal QR code with
Feishu / Lark mobile, the platform provisions a ``PersonalAgent``-archetype
bot application with the required IM permissions pre-attached, and the
poll endpoint returns ``client_id`` / ``client_secret`` — enough to fully
configure :class:`FeishuChannel`.
Domain auto-switches from ``feishu`` to ``lark`` if the poll response's
``user_info.tenant_brand`` reports a Lark tenant.
The HTTP shape mirrors RFC 8628 (OAuth Device Authorization Grant) with
a vendor-specific ``action`` form field selecting init / begin / poll.
Style follows :mod:`EvoScientist.channels.qq.onboard` — httpx, plain
``print`` for progress, and an optional ``qrcode`` dependency for ASCII
rendering.
"""
from __future__ import annotations
import logging
import time
from typing import Any
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
_ACCOUNTS_URLS: dict[str, str] = {
"feishu": "https://accounts.feishu.cn",
"lark": "https://accounts.larksuite.com",
}
_OPEN_URLS: dict[str, str] = {
"feishu": "https://open.feishu.cn",
"lark": "https://open.larksuite.com",
}
_REGISTRATION_PATH = "/oauth/v1/app/registration"
_REQUEST_TIMEOUT_S = 10.0
_DEFAULT_POLL_INTERVAL_S = 5
_DEFAULT_EXPIRE_S = 600
def _accounts_base_url(domain: str) -> str:
return _ACCOUNTS_URLS.get(domain, _ACCOUNTS_URLS["feishu"])
def _open_base_url(domain: str) -> str:
return _OPEN_URLS.get(domain, _OPEN_URLS["feishu"])
# ---------------------------------------------------------------------------
# QR rendering
# ---------------------------------------------------------------------------
try:
import qrcode as _qrcode_mod
except (ImportError, TypeError):
_qrcode_mod = None # type: ignore[assignment]
def _render_qr(url: str) -> bool:
"""Render *url* as an ASCII QR in the terminal. Returns True on success."""
if _qrcode_mod is None:
return False
try:
qr = _qrcode_mod.QRCode(
error_correction=_qrcode_mod.constants.ERROR_CORRECT_M,
border=2,
)
qr.add_data(url)
qr.make(fit=True)
qr.print_ascii(invert=True)
return True
except Exception:
return False
# ---------------------------------------------------------------------------
# Registration HTTP
# ---------------------------------------------------------------------------
def _post_registration(base_url: str, body: dict[str, str]) -> dict:
"""POST form-encoded *body* to the registration endpoint.
The endpoint replies with JSON even on 4xx responses (``authorization_pending``
comes back as HTTP 400 with a parseable body), so we always read the body
and only fall back to raising if the bytes are missing or not JSON.
"""
import httpx
url = f"{base_url}{_REGISTRATION_PATH}"
headers = {"Content-Type": "application/x-www-form-urlencoded"}
with httpx.Client(timeout=_REQUEST_TIMEOUT_S, follow_redirects=True) as client:
resp = client.post(url, data=body, headers=headers)
# Don't raise_for_status — 4xx may still carry a usable JSON body.
try:
return resp.json()
except ValueError:
resp.raise_for_status() # re-raise underlying HTTP error
raise # pragma: no cover — raise_for_status already raised
def _init_registration(domain: str) -> None:
"""Probe the registration environment. Raises if client_secret auth is unavailable."""
res = _post_registration(_accounts_base_url(domain), {"action": "init"})
methods = res.get("supported_auth_methods") or []
if "client_secret" not in methods:
raise RuntimeError(
f"Feishu / Lark registration environment does not support "
f"client_secret auth (got: {methods})"
)
def _begin_registration(domain: str) -> dict:
"""Start the device-code flow.
Returns a dict with ``device_code``, ``qr_url``, ``user_code``,
``interval``, and ``expire_in``.
"""
res = _post_registration(
_accounts_base_url(domain),
{
"action": "begin",
"archetype": "PersonalAgent",
"auth_method": "client_secret",
"request_user_info": "open_id",
},
)
device_code = res.get("device_code")
if not device_code:
raise RuntimeError(
f"Feishu / Lark registration did not return a device_code: {res}"
)
qr_url = res.get("verification_uri_complete") or ""
sep = "&" if "?" in qr_url else "?"
qr_url = f"{qr_url}{sep}from=evoscientist&tp=evoscientist"
return {
"device_code": device_code,
"qr_url": qr_url,
"user_code": res.get("user_code", ""),
"interval": int(res.get("interval") or _DEFAULT_POLL_INTERVAL_S),
"expire_in": int(res.get("expire_in") or _DEFAULT_EXPIRE_S),
}
def _poll_registration(
*,
device_code: str,
interval: int,
expire_in: int,
domain: str,
) -> dict | None:
"""Poll until the user scans, or the device_code expires / is denied.
Auto-switches the polling domain to ``lark`` if the server reports
``user_info.tenant_brand == "lark"`` — the credentials only resolve
against the matching open-platform host.
Returns a dict with ``app_id``, ``app_secret``, ``domain``, ``open_id``
on success, or ``None`` on timeout / explicit denial.
"""
deadline = time.monotonic() + expire_in
current_domain = domain
domain_switched = False
poll_count = 0
while time.monotonic() < deadline:
try:
res = _post_registration(
_accounts_base_url(current_domain),
{
"action": "poll",
"device_code": device_code,
"tp": "ob_app",
},
)
except Exception as exc:
logger.debug("[Feishu onboard] poll request error: %s", exc)
time.sleep(interval)
continue
poll_count += 1
if poll_count == 1:
print(" Waiting for scan…", end="", flush=True)
elif poll_count % 6 == 0:
print(".", end="", flush=True)
# Domain auto-detection — the server may still return creds in
# this same poll, so we fall through rather than restarting.
user_info = res.get("user_info") or {}
if (
user_info.get("tenant_brand") == "lark"
and not domain_switched
and current_domain != "lark"
):
current_domain = "lark"
domain_switched = True
if res.get("client_id") and res.get("client_secret"):
print() # newline after the dots
return {
"app_id": res["client_id"],
"app_secret": res["client_secret"],
"domain": current_domain,
"open_id": user_info.get("open_id"),
}
error = res.get("error", "")
if error in {"access_denied", "expired_token"}:
print()
logger.warning("[Feishu onboard] Registration %s", error)
return None
# authorization_pending / slow_down / unknown — keep polling
time.sleep(interval)
print()
logger.warning("[Feishu onboard] Poll timed out after %ds", expire_in)
return None
# ---------------------------------------------------------------------------
# Bot probe (best-effort, uses tenant_access_token + /bot/v3/info)
# ---------------------------------------------------------------------------
def _probe_bot(app_id: str, app_secret: str, domain: str) -> dict | None:
"""Fetch bot name / bot_open_id via the open-platform REST API.
Best-effort: failures return ``None`` and the caller proceeds without
a friendly bot name. Uses raw HTTP so we don't require ``lark-oapi``
to be installed at onboard time (it's only needed for WebSocket mode).
"""
import httpx
base = _open_base_url(domain)
token_url = f"{base}/open-apis/auth/v3/tenant_access_token/internal"
info_url = f"{base}/open-apis/bot/v3/info"
try:
with httpx.Client(timeout=_REQUEST_TIMEOUT_S, follow_redirects=True) as client:
tok_resp = client.post(
token_url,
json={"app_id": app_id, "app_secret": app_secret},
)
tok_data = tok_resp.json()
if tok_data.get("code") != 0:
logger.debug("[Feishu onboard] token fetch failed: %s", tok_data)
return None
token = tok_data.get("tenant_access_token")
if not token:
return None
info_resp = client.get(
info_url,
headers={"Authorization": f"Bearer {token}"},
)
info_data = info_resp.json()
except Exception as exc:
logger.debug("[Feishu onboard] bot probe failed: %s", exc)
return None
if info_data.get("code") != 0:
return None
bot = info_data.get("bot") or info_data.get("data", {}).get("bot") or {}
return {
"bot_name": bot.get("app_name") or bot.get("bot_name"),
"bot_open_id": bot.get("open_id"),
}
# ---------------------------------------------------------------------------
# Public entry-point
# ---------------------------------------------------------------------------
def qr_register(
*,
initial_domain: str = "feishu",
timeout_seconds: int = 600,
) -> dict[str, Any] | None:
"""Run the Feishu / Lark scan-to-create QR registration flow.
Args:
initial_domain: ``"feishu"`` (default, mainland) or ``"lark"`` (overseas).
Auto-switches mid-flow if the scanning user is on the other tenant.
timeout_seconds: Wall-clock budget for the whole flow.
Returns on success::
{
"app_id": str,
"app_secret": str,
"domain": "feishu" | "lark",
"open_id": str | None,
"bot_name": str | None,
"bot_open_id": str | None,
}
Returns ``None`` on expected failures (network, denial, timeout).
"""
try:
return _qr_register_inner(
initial_domain=initial_domain,
timeout_seconds=timeout_seconds,
)
except Exception as exc:
logger.warning("[Feishu onboard] Registration failed: %s", exc)
return None
def _qr_register_inner(
*,
initial_domain: str,
timeout_seconds: int,
) -> dict[str, Any] | None:
print(" Connecting to Feishu / Lark…", end="", flush=True)
_init_registration(initial_domain)
begin = _begin_registration(initial_domain)
print(" done.")
print()
qr_url = begin["qr_url"]
if _render_qr(qr_url):
print(
f"\n Scan the QR code above with Feishu / Lark on your phone,\n"
f" or open this URL directly:\n {qr_url}"
)
else:
print(f" Open this URL in Feishu / Lark on your phone:\n\n {qr_url}\n")
print(
" Tip: pip install qrcode to display a scannable QR code here next time"
)
print()
result = _poll_registration(
device_code=begin["device_code"],
interval=begin["interval"],
expire_in=min(begin["expire_in"], timeout_seconds),
domain=initial_domain,
)
if not result:
return None
bot_info = _probe_bot(result["app_id"], result["app_secret"], result["domain"])
if bot_info:
result["bot_name"] = bot_info.get("bot_name")
result["bot_open_id"] = bot_info.get("bot_open_id")
else:
result["bot_name"] = None
result["bot_open_id"] = None
return result
+2 -5
View File
@@ -20,15 +20,11 @@ Examples:
import argparse
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from .channel import FeishuChannel, FeishuConfig
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger(__name__)
@@ -96,6 +92,7 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
config = FeishuConfig(
+2
View File
@@ -19,6 +19,7 @@ Examples:
import argparse
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from . import IMessageChannel, IMessageConfig
@@ -67,6 +68,7 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
config = IMessageConfig(
+5 -10
View File
@@ -75,13 +75,11 @@ class DedupCache:
max_size: int = _DEDUP_MAX,
trim_to: int = _DEDUP_TRIM,
ttl_seconds: float = _DEDUP_TTL,
clock: Callable[[], float] | None = None,
) -> None:
self._seen: OrderedDict[str, float] = OrderedDict()
self._max = max_size
self._trim = trim_to
self._ttl = ttl_seconds
self._clock = clock or time.monotonic
# ── public API ──────────────────────────────────────────────────
@@ -95,16 +93,15 @@ class DedupCache:
if not msg_id:
return False
now = self._clock()
self._prune(now)
self._prune()
if msg_id in self._seen:
# LRU: refresh position and timestamp
self._seen.move_to_end(msg_id)
self._seen[msg_id] = now
self._seen[msg_id] = time.monotonic()
return True
self._seen[msg_id] = now
self._seen[msg_id] = time.monotonic()
if len(self._seen) > self._max:
while len(self._seen) > self._trim:
self._seen.popitem(last=False)
@@ -121,9 +118,9 @@ class DedupCache:
# ── internal ────────────────────────────────────────────────────
def _prune(self, now: float | None = None) -> None:
def _prune(self) -> None:
"""Remove entries older than *ttl_seconds*."""
cutoff = (self._clock() if now is None else now) - self._ttl
cutoff = time.monotonic() - self._ttl
# OrderedDict is insertion-ordered; oldest entries are first.
while self._seen:
_key, ts = next(iter(self._seen.items()))
@@ -431,13 +428,11 @@ class DedupMiddleware(InboundMiddleware):
max_size: int = 1000,
trim_to: int = 500,
ttl_seconds: float = 3600.0,
clock: Callable[[], float] | None = None,
) -> None:
self._cache = DedupCache(
max_size=max_size,
trim_to=trim_to,
ttl_seconds=ttl_seconds,
clock=clock,
)
async def process_inbound(
+1 -2
View File
@@ -10,9 +10,8 @@ Usage in config:
from ..channel_manager import _parse_csv, register_channel
from .channel import QQChannel, QQConfig
from .onboard import qr_register
__all__ = ["QQChannel", "QQConfig", "qr_register"]
__all__ = ["QQChannel", "QQConfig"]
def create_from_config(config) -> QQChannel:
+15 -211
View File
@@ -26,55 +26,6 @@ except ImportError:
GroupMessage = None
# ── Inline keyboard (button) helpers ─────────────────────────────────
def _normalize_button(btn: dict) -> tuple[str, str] | None:
"""Return ``(label, value)`` for a button, or ``None`` if no label."""
label = (btn.get("text") or "").strip()
if not label:
return None
raw = btn.get("value")
return label, str(raw) if raw is not None else label
def _build_qq_keyboard(buttons: list[dict]) -> dict | None:
"""Build a QQ Bot keyboard payload (one button per row).
Render style: 1 = primary (blue), 0 = secondary (grey) — QQ has no danger.
``action.permission`` is required by the schema; ``type=2`` is harmless for
C2C (the click always comes from the DM peer). Returns ``None`` if no
button has a usable label.
"""
rows: list[dict] = []
for idx, btn in enumerate(buttons):
norm = _normalize_button(btn)
if norm is None:
continue
label, value = norm
style = 1 if btn.get("type") == "primary" else 0
rows.append(
{
"buttons": [
{
"id": btn.get("id") or f"btn_{idx}",
"render_data": {
"label": label,
"visited_label": label,
"style": style,
},
"action": {
"type": 1, # callback (server pushes interaction event)
"permission": {"type": 2},
"data": value,
},
}
]
}
)
return {"content": {"rows": rows}} if rows else None
@dataclass
class QQConfig(BaseChannelConfig):
app_id: str = ""
@@ -84,11 +35,7 @@ class QQConfig(BaseChannelConfig):
def _make_bot_class(channel: "QQChannel") -> "type[botpy.Client]":
"""Create a botpy Client subclass bound to the given channel."""
intents = botpy.Intents(
public_messages=True,
direct_message=True,
interaction=True, # button clicks → on_interaction_create
)
intents = botpy.Intents(public_messages=True, direct_message=True)
class _Bot(botpy.Client):
def __init__(self):
@@ -103,9 +50,6 @@ def _make_bot_class(channel: "QQChannel") -> "type[botpy.Client]":
async def on_group_at_message_create(self, message: "GroupMessage"):
await channel._on_msg(message, "group")
async def on_interaction_create(self, interaction):
await channel._on_interaction(interaction)
return _Bot
@@ -230,76 +174,6 @@ class QQChannel(Channel):
except Exception as e:
logger.error(f"Error handling QQ message: {e}")
async def _on_interaction(self, interaction) -> None:
"""Handle ``on_interaction_create`` (button click).
Surfaces the click as an :class:`InboundMessage` whose ``content`` is
the button's ``data`` verbatim — so a "1"/"approve"/… click flows
through ``_parse_approval_reply`` exactly like a typed reply.
The click runs through inbound middleware (Dedup suppresses QQ
retries) but is published directly to the bus so the per-sender
debounce buffer doesn't merge the click value with subsequent text.
Group-scope clicks are ignored (DM-only by design).
"""
# ACK first — QQ requires a response within ~5s or the button UI
# shows "expired". Code 0 just means "received"; downstream still
# decides the actual approval/rejection.
interaction_id = getattr(interaction, "id", "") or ""
if interaction_id and self._client:
try:
await self._client.api.on_interaction_result(interaction_id, 0)
except Exception as ack_exc:
logger.debug("QQ interaction ack failed: %s", ack_exc)
try:
user_openid = getattr(interaction, "user_openid", "") or ""
if not user_openid:
logger.debug("QQ interaction ignored (no user_openid; not C2C)")
return
resolved = getattr(getattr(interaction, "data", None), "resolved", None)
button_data = getattr(resolved, "button_data", "") or ""
button_id = getattr(resolved, "button_id", "") or ""
triggering_msg_id = getattr(resolved, "message_id", "") or ""
# QQ may serialize non-str values; coerce. Fall back to button id
# when no data — same path as a typed reply via _parse_approval_reply.
button_value = str(button_data) if button_data != "" else ""
text = button_value or button_id
# Stable id so DedupMiddleware suppresses any QQ retry callbacks.
message_id = (
f"{triggering_msg_id}:action:{interaction_id}"
if interaction_id
else f"qq_action:{datetime.now().timestamp()}"
)
raw = RawIncoming(
sender_id=user_openid,
chat_id=user_openid, # C2C: chat_id == user_openid
text=text,
timestamp=datetime.now(),
message_id=message_id,
metadata={
"chat_id": user_openid,
"msg_type": "c2c",
"event_id": triggering_msg_id,
"backend": "qq",
"button_click": True,
"button_id": button_id,
"button_value": button_value,
},
is_group=False,
was_mentioned=True,
)
inbound = await self._build_inbound_async(raw)
if inbound is not None and self._bus:
await self._bus.publish_inbound(inbound)
except Exception:
logger.exception("QQ interaction handler error")
# ── Send ──────────────────────────────────────────────────────
def _next_msg_seq(self, msg_id: str) -> int:
@@ -321,81 +195,17 @@ class QQChannel(Channel):
msg_type = (metadata or {}).get("msg_type", "c2c")
msg_id = (metadata or {}).get("event_id", "")
seq = self._next_msg_seq(msg_id)
# Inline keyboard is C2C-only here — group keyboards have stricter
# permission semantics and are out of scope for now.
buttons = (metadata or {}).get("buttons") if msg_type == "c2c" else None
keyboard = _build_qq_keyboard(buttons) if buttons else None
try:
await self._post_markdown_message(
chat_id, raw_text, msg_type, msg_id, seq, keyboard=keyboard
)
await self._post_markdown_message(chat_id, raw_text, msg_type, msg_id, seq)
return
except Exception as exc:
if not self._should_fallback_to_plain_text(exc):
logger.error(
"QQ markdown send failed with non-fallbackable error "
"(chat_id=%s, msg_id=%s, seq=%s): %r",
chat_id,
msg_id,
seq,
exc,
)
raise
self._record_markdown_fallback(chat_id, raw_text, exc)
logger.warning(
"QQ markdown send failed, falling back to plain text "
"(chat_id=%s, msg_id=%s, seq=%s): %r",
chat_id,
msg_id,
seq,
exc,
)
logger.debug("QQ markdown send failed, falling back to plain text: %s", exc)
# QQ may have already consumed `seq` server-side even on failure.
# Reusing it for the plain retry triggers "duplicate msg_seq", so
# always advance to a fresh seq before the fallback send.
fallback_seq = self._next_msg_seq(msg_id)
plain_text = self._plain_formatter.format(raw_text)
# Plain-text fallback can't carry a keyboard. Append `value=label`
# pairs so the user can still type "1"/"approve"/… instead of
# tapping (`_parse_approval_reply` accepts the same values).
if buttons:
pairs = []
for btn in buttons:
norm = _normalize_button(btn)
if norm is not None:
label, value = norm
pairs.append(f"{value}={label}")
if pairs:
plain_text = f"{plain_text}\n\nReply: {', '.join(pairs)}"
try:
await self._post_plain_message(
chat_id, plain_text, msg_type, msg_id, fallback_seq
)
except Exception as plain_exc:
logger.error(
"QQ plain fallback also failed (chat_id=%s, msg_id=%s, seq=%s): %r",
chat_id,
msg_id,
fallback_seq,
plain_exc,
)
raise
# QQ server-side error codes / fragments that indicate the markdown
# request itself is invalid (template not configured, format rejected,
# content audit, etc.). Seeing any of these means we should retry with
# plain text rather than re-raise.
_QQ_MARKDOWN_ERROR_MARKERS: ClassVar[tuple[str, ...]] = (
"304014", # markdown template not configured
"304003", # invalid markdown params
"40034059", # generic send message failed (often markdown-related)
"模板", # CN: template (standard form)
"模版", # CN: template (variant form)
"审核", # CN: audit
)
await self._post_plain_message(chat_id, plain_text, msg_type, msg_id, seq)
def _should_fallback_to_plain_text(self, exc: Exception) -> bool:
"""Return True only for markdown compatibility/validation failures."""
@@ -404,20 +214,17 @@ class QQChannel(Channel):
msg = str(exc).lower()
compatibility_tokens = ("unsupported", "unexpected", "unknown", "invalid")
if "unexpected keyword argument" in msg:
return True
if "markdown" in msg and any(token in msg for token in compatibility_tokens):
return True
if "msg_type" in msg and any(token in msg for token in compatibility_tokens):
return True
# QQ-specific server error codes returned by qq-botpy as strings.
raw = str(exc)
for marker in self._QQ_MARKDOWN_ERROR_MARKERS:
if marker in raw or marker.lower() in msg:
return True
return False
return (
"unexpected keyword argument" in msg
or (
"markdown" in msg
and any(token in msg for token in compatibility_tokens)
)
or (
"msg_type" in msg
and any(token in msg for token in compatibility_tokens)
)
)
def _record_markdown_fallback(
self,
@@ -447,7 +254,6 @@ class QQChannel(Channel):
msg_type: str,
msg_id: str,
seq: int,
keyboard: dict | None = None,
) -> None:
payload = {
"msg_type": 2,
@@ -455,8 +261,6 @@ class QQChannel(Channel):
"msg_id": msg_id,
"msg_seq": seq,
}
if keyboard is not None:
payload["keyboard"] = keyboard
if msg_type == "group":
await self._client.api.post_group_message(
group_openid=chat_id,
-49
View File
@@ -1,49 +0,0 @@
"""AES-256-GCM utilities for QQ Bot scan-to-configure credential decryption.
Ported from hermes-agent/gateway/platforms/qqbot/crypto.py — the q.qq.com
``create_bind_task`` / ``poll_bind_result`` flow uses AES-256-GCM to keep
the bot's *client_secret* off the wire in plaintext.
"""
from __future__ import annotations
import base64
import os
def generate_bind_key() -> str:
"""Generate a 256-bit random AES key, base64-encoded.
The key is sent to ``create_bind_task`` so the server can encrypt
the bot's *client_secret* before returning it. Only this client
holds the key, so the secret never travels in plaintext.
"""
return base64.b64encode(os.urandom(32)).decode()
def decrypt_secret(encrypted_base64: str, key_base64: str) -> str:
"""Decrypt a base64-encoded AES-256-GCM ciphertext.
Ciphertext layout (after base64-decoding)::
IV (12 bytes) ‖ ciphertext (N bytes) ‖ AuthTag (16 bytes)
Args:
encrypted_base64: The ``bot_encrypt_secret`` value returned by
``poll_bind_result``.
key_base64: The base64 AES key produced by :func:`generate_bind_key`.
Returns:
The decrypted *client_secret* as a UTF-8 string.
"""
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
key = base64.b64decode(key_base64)
raw = base64.b64decode(encrypted_base64)
iv = raw[:12]
ciphertext_with_tag = raw[12:] # AESGCM expects ciphertext + tag concatenated
aesgcm = AESGCM(key)
plaintext = aesgcm.decrypt(iv, ciphertext_with_tag, None)
return plaintext.decode("utf-8")
-300
View File
@@ -1,300 +0,0 @@
"""QQ Bot scan-to-configure (QR code onboard) flow.
Ported from hermes-agent/gateway/platforms/qqbot/onboard.py.
Calls the ``q.qq.com`` ``create_bind_task`` / ``poll_bind_result`` APIs to
generate a QR code URL and poll for scan completion. On success the caller
receives the bot's *app_id*, *client_secret* (decrypted locally), and the
scanner's *user_openid* — enough to fully configure the QQ channel.
The bot must already be registered at https://q.qq.com — scanning binds
the QQ user (developer / admin) to the existing application; it does not
create a new one.
Reference: https://bot.q.qq.com/wiki/develop/api-v2/
"""
from __future__ import annotations
import logging
import os
import platform
import sys
import time
from enum import IntEnum
from urllib.parse import quote
from .crypto import decrypt_secret, generate_bind_key
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Endpoints / timing
# ---------------------------------------------------------------------------
# The portal domain is configurable for corporate proxies / sandbox routing.
PORTAL_HOST = os.getenv("QQ_PORTAL_HOST", "q.qq.com")
ONBOARD_CREATE_PATH = "/lite/create_bind_task"
ONBOARD_POLL_PATH = "/lite/poll_bind_result"
QR_URL_TEMPLATE = (
"https://q.qq.com/qqbot/openclaw/connect.html"
"?task_id={task_id}&_wv=2&source=evoscientist"
)
ONBOARD_API_TIMEOUT = 10.0
ONBOARD_POLL_INTERVAL = 2.0
_MAX_REFRESHES = 3
# ---------------------------------------------------------------------------
# Bind status
# ---------------------------------------------------------------------------
class BindStatus(IntEnum):
"""Status codes returned by ``poll_bind_result``."""
NONE = 0
PENDING = 1
COMPLETED = 2
EXPIRED = 3
# ---------------------------------------------------------------------------
# HTTP headers
# ---------------------------------------------------------------------------
def _get_evoscientist_version() -> str:
try:
from importlib.metadata import version
return version("evoscientist")
except Exception:
return "dev"
def _build_user_agent() -> str:
py_version = (
f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}"
)
os_name = platform.system().lower()
return (
f"EvoScientistQQ/1.0.0 (Python/{py_version}; {os_name}; "
f"EvoScientist/{_get_evoscientist_version()})"
)
def _api_headers() -> dict[str, str]:
"""Standard HTTP headers for q.qq.com onboard API requests.
``q.qq.com`` requires ``Accept: application/json`` — without it,
the server returns a JavaScript anti-bot challenge page.
"""
return {
"Content-Type": "application/json",
"Accept": "application/json",
"User-Agent": _build_user_agent(),
}
# ---------------------------------------------------------------------------
# QR rendering
# ---------------------------------------------------------------------------
try:
import qrcode as _qrcode_mod
except (ImportError, TypeError):
_qrcode_mod = None # type: ignore[assignment]
def _render_qr(url: str) -> bool:
"""Render a QR code to the terminal. Returns True on success."""
if _qrcode_mod is None:
return False
try:
qr = _qrcode_mod.QRCode(
error_correction=_qrcode_mod.constants.ERROR_CORRECT_M,
border=2,
)
qr.add_data(url)
qr.make(fit=True)
qr.print_ascii(invert=True)
return True
except Exception:
return False
# ---------------------------------------------------------------------------
# HTTP helpers
# ---------------------------------------------------------------------------
def _create_bind_task(timeout: float = ONBOARD_API_TIMEOUT) -> tuple[str, str]:
"""Create a bind task and return *(task_id, aes_key_base64)*.
Raises:
RuntimeError: if the API returns a non-zero ``retcode``.
"""
import httpx
url = f"https://{PORTAL_HOST}{ONBOARD_CREATE_PATH}"
key = generate_bind_key()
with httpx.Client(timeout=timeout, follow_redirects=True) as client:
resp = client.post(url, json={"key": key}, headers=_api_headers())
resp.raise_for_status()
data = resp.json()
if data.get("retcode") != 0:
raise RuntimeError(data.get("msg", "create_bind_task failed"))
task_id = data.get("data", {}).get("task_id")
if not task_id:
raise RuntimeError("create_bind_task: missing task_id in response")
logger.debug("create_bind_task ok: task_id=%s", task_id)
return task_id, key
def _poll_bind_result(
task_id: str,
timeout: float = ONBOARD_API_TIMEOUT,
) -> tuple[BindStatus, str, str, str]:
"""Poll the bind result for *task_id*.
Returns:
``(status, bot_appid, bot_encrypt_secret, user_openid)``.
Raises:
RuntimeError: if the API returns a non-zero ``retcode``.
"""
import httpx
url = f"https://{PORTAL_HOST}{ONBOARD_POLL_PATH}"
with httpx.Client(timeout=timeout, follow_redirects=True) as client:
resp = client.post(url, json={"task_id": task_id}, headers=_api_headers())
resp.raise_for_status()
data = resp.json()
if data.get("retcode") != 0:
raise RuntimeError(data.get("msg", "poll_bind_result failed"))
d = data.get("data", {})
return (
BindStatus(d.get("status", 0)),
str(d.get("bot_appid", "")),
d.get("bot_encrypt_secret", ""),
d.get("user_openid", ""),
)
def build_connect_url(task_id: str) -> str:
"""Build the QR-code target URL for a given *task_id*."""
return QR_URL_TEMPLATE.format(task_id=quote(task_id))
# ---------------------------------------------------------------------------
# Public entry-point
# ---------------------------------------------------------------------------
def qr_register(timeout_seconds: int = 600) -> dict | None:
"""Run the QQ Bot scan-to-configure QR registration flow.
Handles create → display → poll → decrypt in one call. The QR
auto-refreshes up to ``_MAX_REFRESHES`` times if the user takes
too long to scan.
Args:
timeout_seconds: Total wall-clock budget across all refreshes.
Returns:
``{"app_id": ..., "client_secret": ..., "user_openid": ...}`` on
success, or ``None`` on failure / expiry / cancellation.
"""
deadline = time.monotonic() + timeout_seconds
for refresh_count in range(_MAX_REFRESHES + 1):
# ── Create bind task ──
try:
task_id, aes_key = _create_bind_task()
except Exception as exc:
logger.warning("[QQ onboard] Failed to create bind task: %s", exc)
return None
url = build_connect_url(task_id)
# ── Display QR code + URL ──
print()
if _render_qr(url):
print(f" Scan the QR code above, or open this URL on your phone:\n {url}")
else:
print(f" Open this URL in QQ on your phone:\n {url}")
print(" Tip: pip install qrcode to display a scannable QR code here")
print()
# ── Poll loop ──
consecutive_errors = 0
while time.monotonic() < deadline:
try:
status, app_id, encrypted_secret, user_openid = _poll_bind_result(
task_id
)
except Exception as exc:
consecutive_errors += 1
logger.warning(
"[QQ onboard] poll_bind_result failed (%d consecutive): %s",
consecutive_errors,
exc,
)
if consecutive_errors >= 5:
print(
"\n Repeated polling failures — aborting."
" See logs for details."
)
return None
time.sleep(ONBOARD_POLL_INTERVAL)
continue
consecutive_errors = 0
if status == BindStatus.COMPLETED:
try:
client_secret = decrypt_secret(encrypted_secret, aes_key)
except Exception as exc:
logger.warning("[QQ onboard] decrypt_secret failed: %s", exc)
return None
print()
print(f" QR scan complete! (App ID: {app_id})")
if user_openid:
print(f" Scanner's OpenID: {user_openid}")
return {
"app_id": app_id,
"client_secret": client_secret,
"user_openid": user_openid,
}
if status == BindStatus.EXPIRED:
if refresh_count >= _MAX_REFRESHES:
logger.warning(
"[QQ onboard] QR code expired %d times — giving up",
_MAX_REFRESHES,
)
return None
print(
f"\n QR code expired, refreshing... "
f"({refresh_count + 1}/{_MAX_REFRESHES})"
)
break # next outer iteration creates a new task
time.sleep(ONBOARD_POLL_INTERVAL)
else:
# deadline reached without completing
logger.warning("[QQ onboard] Poll timed out after %ds", timeout_seconds)
return None
return None
+2 -5
View File
@@ -19,15 +19,11 @@ Examples:
import argparse
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from .channel import QQChannel, QQConfig
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger(__name__)
@@ -68,6 +64,7 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
config = QQConfig(
+2 -5
View File
@@ -19,15 +19,11 @@ Examples:
import argparse
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from .channel import SignalChannel, SignalConfig
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger(__name__)
@@ -78,6 +74,7 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
config = SignalConfig(
+2 -5
View File
@@ -19,15 +19,11 @@ Examples:
import argparse
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from .channel import SlackChannel, SlackConfig
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger(__name__)
@@ -78,6 +74,7 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
config = SlackConfig(
-3
View File
@@ -108,10 +108,8 @@ async def _async_main(
if use_agent:
logger.info("Loading EvoScientist agent...")
from ..EvoScientist import create_cli_agent
from ..gateway import create_runtime_gateways
agent = create_cli_agent()
runtime_gateways = create_runtime_gateways()
logger.info("Agent loaded")
consumer = InboundConsumer(
@@ -119,7 +117,6 @@ async def _async_main(
manager=manager,
agent=agent,
thread_id="",
graph_gateway=runtime_gateways.graph_gateway,
send_thinking=send_thinking,
)
manager.register_health_provider("consumer", lambda: consumer.metrics)
+2 -5
View File
@@ -19,15 +19,11 @@ Examples:
import argparse
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from .channel import TelegramChannel, TelegramConfig
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger(__name__)
@@ -63,6 +59,7 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
config = TelegramConfig(
+21 -59
View File
@@ -5,16 +5,13 @@ Supports multiple WeChat backends:
— Most stable, pure HTTP, no third-party dependencies
- **wechatmp**: 微信公众号 (WeChat Official Account) via official API
— Pure HTTP webhook, suitable for public-facing bots
- **personal**: 个人微信 via Tencent's iLink Bot API
— Long-poll + AES-128-ECB CDN media protocol; QR-code login required.
Adapted from hermes-agent.
Backends 1+2 share the HTTP-webhook ``WeChatChannel``; backend 3 uses the
long-poll ``WeixinPersonalChannel``.
Both backends use httpx (already a core dependency) and receive messages
via HTTP webhook, send replies via REST API.
Usage in config:
channel_enabled = "wechat"
wechat_backend = "wecom" # or "wechatmp" or "personal"
wechat_backend = "wecom" # or "wechatmp"
# WeCom settings
wechat_wecom_corp_id = "..."
@@ -30,54 +27,19 @@ Usage in config:
wechat_mp_token = "..."
wechat_mp_encoding_aes_key = "..."
wechat_webhook_port = 9001
# OR: Personal WeChat (iLink Bot)
# First run `python -m EvoScientist.channels.wechat.serve --qr-login`
# to obtain an account_id + token via QR-code scan.
wechat_personal_account_id = "..."
wechat_personal_token = "..." # optional if persisted on disk
wechat_personal_dm_policy = "open" # open | allowlist
wechat_personal_group_policy = "disabled"
"""
from ..channel_manager import _parse_csv, register_channel
from .channel import WeChatChannel, WeChatMPConfig, WeComConfig
from .personal import WeixinPersonalChannel, WeixinPersonalConfig, qr_login
__all__ = [
"WeChatChannel",
"WeChatMPConfig",
"WeComConfig",
"WeixinPersonalChannel",
"WeixinPersonalConfig",
"qr_login",
]
__all__ = ["WeChatChannel", "WeChatMPConfig", "WeComConfig"]
def create_from_config(config):
"""Factory dispatched on ``config.wechat_backend``."""
backend = (getattr(config, "wechat_backend", "") or "wecom").lower()
allowed = _parse_csv(getattr(config, "wechat_allowed_senders", ""))
proxy = getattr(config, "wechat_proxy", "") or None
port = int(getattr(config, "wechat_webhook_port", 9001) or 9001)
if backend == "personal":
group_allowed = _parse_csv(getattr(config, "wechat_personal_group_allowed", ""))
cfg = WeixinPersonalConfig(
account_id=getattr(config, "wechat_personal_account_id", ""),
token=getattr(config, "wechat_personal_token", ""),
base_url=getattr(config, "wechat_personal_base_url", "")
or "https://ilinkai.weixin.qq.com",
cdn_base_url=getattr(config, "wechat_personal_cdn_base_url", "")
or "https://novac2c.cdn.weixin.qq.com/c2c",
dm_policy=getattr(config, "wechat_personal_dm_policy", "open") or "open",
group_policy=getattr(config, "wechat_personal_group_policy", "disabled")
or "disabled",
group_allowed_senders=group_allowed,
allowed_senders=allowed,
proxy=proxy,
)
return WeixinPersonalChannel(cfg)
def create_from_config(config) -> WeChatChannel:
backend = config.wechat_backend or "wecom"
allowed = _parse_csv(config.wechat_allowed_senders)
proxy = config.wechat_proxy or None
port = int(config.wechat_webhook_port or 9001)
if backend == "wechatmp":
mp_config = WeChatMPConfig(
@@ -90,18 +52,18 @@ def create_from_config(config):
proxy=proxy,
)
return WeChatChannel(mp_config, backend="wechatmp")
wecom_config = WeComConfig(
corp_id=config.wechat_wecom_corp_id,
agent_id=config.wechat_wecom_agent_id,
secret=config.wechat_wecom_secret,
token=config.wechat_wecom_token,
encoding_aes_key=config.wechat_wecom_encoding_aes_key,
webhook_port=port,
allowed_senders=allowed,
proxy=proxy,
)
return WeChatChannel(wecom_config, backend="wecom")
else:
wecom_config = WeComConfig(
corp_id=config.wechat_wecom_corp_id,
agent_id=config.wechat_wecom_agent_id,
secret=config.wechat_wecom_secret,
token=config.wechat_wecom_token,
encoding_aes_key=config.wechat_wecom_encoding_aes_key,
webhook_port=port,
allowed_senders=allowed,
proxy=proxy,
)
return WeChatChannel(wecom_config, backend="wecom")
register_channel("wechat", create_from_config)
-47
View File
@@ -83,53 +83,6 @@ def _aes_encrypt(key: bytes, iv: bytes, plaintext: bytes) -> bytes:
) from None
def aes128_ecb_decrypt(ciphertext: bytes, key: bytes) -> bytes:
"""AES-128-ECB decryption with PKCS#7 unpadding.
Used by the personal-WeChat (iLink) backend for CDN-encrypted media
payloads. Block size is 16; the WeChat CDN protocol pads with PKCS#7.
"""
if _HAS_PYCRYPTO:
cipher = AES.new(key, AES.MODE_ECB)
padded = cipher.decrypt(ciphertext)
else:
try:
import pyaes
decrypter = pyaes.Decrypter(pyaes.AESModeOfOperationECB(key))
padded = decrypter.feed(ciphertext)
padded += decrypter.feed()
except ImportError:
raise ImportError(
"WeChat CDN media decryption requires pycryptodome or pyaes. "
"Install with: pip install pycryptodome"
) from None
if not padded:
return padded
pad_len = padded[-1]
if 1 <= pad_len <= 16 and padded.endswith(bytes([pad_len]) * pad_len):
return padded[:-pad_len]
return padded
def parse_ilink_aes_key(aes_key_b64: str) -> bytes:
"""Parse the iLink CDN AES key.
iLink encodes the 16-byte key in two formats:
- direct base64 of 16 bytes
- base64 of a 32-char ASCII hex string (which decodes to 16 raw bytes)
"""
decoded = base64.b64decode(aes_key_b64)
if len(decoded) == 16:
return decoded
if len(decoded) == 32:
text = decoded.decode("ascii", errors="ignore")
if text and all(ch in "0123456789abcdefABCDEF" for ch in text):
return bytes.fromhex(text)
raise ValueError(f"unexpected aes_key format ({len(decoded)} decoded bytes)")
class WeChatCrypto:
"""Handles WeChat/WeCom message encryption and decryption.
File diff suppressed because it is too large Load Diff
-37
View File
@@ -70,40 +70,3 @@ async def validate_wechat_mp(
return False, f"Error ({data.get('errcode')}): {data.get('errmsg')}"
except Exception as e:
return False, f"Error: {e}"
async def validate_wechat_personal(
account_id: str,
token: str = "",
) -> tuple[bool, str]:
"""Validate that a personal WeChat (iLink) account has been logged in.
Personal-WeChat credentials are obtained via QR-code scan and persisted
on disk; there is no offline credential format the user can paste in.
This probe simply checks that:
- the account_id is set, and
- either *token* is supplied inline, or a saved-account file exists for
*account_id* under ``DATA_DIR/wechat_personal/accounts/``.
Online liveness is not checked because the iLink long-poll endpoint is
not designed for cheap probes.
"""
if not account_id:
return False, (
"account_id is required. Run "
"`python -m EvoScientist.channels.wechat.serve --qr-login` first."
)
if token:
return True, f"Personal WeChat account {account_id[:8]}… token provided"
from .personal import load_account
persisted = load_account(account_id)
if not persisted or not persisted.get("token"):
return False, (
f"No saved credentials for account_id={account_id[:8]}…. "
"Run `python -m EvoScientist.channels.wechat.serve --qr-login`."
)
return True, f"Personal WeChat account {account_id[:8]}… loaded from disk"
+6 -81
View File
@@ -20,13 +20,6 @@ Usage:
--token TOKEN \\
--aes-key AES_KEY
# Personal WeChat (个人微信 via iLink Bot)
# First, log in via QR scan to obtain credentials:
python -m EvoScientist.channels.wechat.serve --qr-login
# Then run with the saved account_id:
python -m EvoScientist.channels.wechat.serve \\
--backend personal --account-id <id>
Options:
--port PORT Webhook listen port (default: 9001)
--allow USER_ID Allowed sender (repeatable)
@@ -35,24 +28,13 @@ Options:
"""
import argparse
import asyncio
import logging
from ...logging_config import configure_logging_from_settings
from ..bus import MessageBus
from ..standalone import run_standalone
from .channel import WeChatChannel, WeChatMPConfig, WeComConfig
from .personal import (
WeixinPersonalChannel,
WeixinPersonalConfig,
load_account,
qr_login,
)
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger(__name__)
@@ -64,15 +46,10 @@ def parse_args():
)
parser.add_argument(
"--backend",
choices=["wecom", "wechatmp", "personal"],
choices=["wecom", "wechatmp"],
default="wecom",
help="WeChat backend type (default: wecom)",
)
parser.add_argument(
"--qr-login",
action="store_true",
help="Run interactive QR-code login for personal WeChat and exit",
)
parser.add_argument("--port", type=int, default=9001, help="Webhook port")
parser.add_argument(
"--allow",
@@ -108,31 +85,6 @@ def parse_args():
mp.add_argument("--app-id", default="", help="MP App ID")
mp.add_argument("--app-secret", default="", help="MP App Secret")
# Personal-WeChat settings
personal = parser.add_argument_group("Personal WeChat (iLink Bot)")
personal.add_argument(
"--account-id",
default="",
help="iLink account_id (obtained via --qr-login)",
)
personal.add_argument(
"--bot-token",
default="",
help="iLink bearer token; if omitted, loaded from disk via account-id",
)
personal.add_argument(
"--dm-policy",
choices=["open", "allowlist"],
default="open",
help="Direct-message policy (default: open)",
)
personal.add_argument(
"--group-policy",
choices=["open", "allowlist", "disabled"],
default="disabled",
help="Group-message policy (default: disabled — iLink rarely delivers)",
)
# Shared settings
parser.add_argument("--token", default="", help="Callback verification token")
parser.add_argument("--aes-key", default="", help="EncodingAESKey")
@@ -143,14 +95,8 @@ def parse_args():
def main():
"""Entry point."""
configure_logging_from_settings(default_level=logging.INFO)
args = parse_args()
if args.qr_login:
result = asyncio.run(qr_login())
if not result:
raise SystemExit(1)
return
allowed = set(args.allowed_senders) if args.allowed_senders else None
allowed_channels = set(args.allowed_channels) if args.allowed_channels else None
proxy = args.proxy or None
@@ -167,8 +113,7 @@ def main():
allowed_channels=allowed_channels,
proxy=proxy,
)
channel = WeChatChannel(config, backend=args.backend)
elif args.backend == "wechatmp":
else:
config = WeChatMPConfig(
app_id=args.app_id,
app_secret=args.app_secret,
@@ -179,31 +124,11 @@ def main():
allowed_channels=allowed_channels,
proxy=proxy,
)
channel = WeChatChannel(config, backend=args.backend)
else: # personal
token = args.bot_token
if not token and args.account_id:
persisted = load_account(args.account_id)
if persisted:
token = persisted.get("token", "")
if not args.account_id or not token:
raise SystemExit(
"Personal WeChat requires --account-id (and a saved token, "
"obtained via --qr-login)."
)
personal_config = WeixinPersonalConfig(
account_id=args.account_id,
token=token,
allowed_senders=allowed,
allowed_channels=allowed_channels,
dm_policy=args.dm_policy,
group_policy=args.group_policy,
proxy=proxy,
)
channel = WeixinPersonalChannel(personal_config)
send_thinking = args.thinking and args.agent
bus = MessageBus()
channel = WeChatChannel(config, backend=args.backend)
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
+26 -66
View File
@@ -1,67 +1,27 @@
"""EvoScientist CLI package.
"""EvoScientist CLI package."""
Most re-exports are served lazily through ``__getattr__`` so that a bare
``import EvoScientist.cli`` only costs what ``main()`` actually needs. That
keeps ``evosci --help`` fast — the heavy chat-model/TUI/langgraph imports
only pay their cost when someone actually touches those names.
"""
from __future__ import annotations
from .. import deploy as _deploy_pkg # noqa: F401 — registers `deploy` @app.command
# Backward-compat re-exports (tests import these from EvoScientist.cli)
from ..stream.state import ( # noqa: F401
StreamState,
SubAgentState,
_build_todo_stats,
_parse_todo_items,
)
from . import commands # noqa: F401 — registers @app.command decorators
from ._app import app
from ._constants import WELCOME_SLOGANS # noqa: F401
from .agent import _deduplicate_run_name # noqa: F401
from .channel import _channels_is_running, _channels_stop # noqa: F401
__all__ = [
"DEFAULT_UI_BACKEND",
"SUPPORTED_UI_BACKENDS",
"WELCOME_SLOGANS",
"StreamState",
"SubAgentState",
"_build_todo_stats",
"_channels_is_running",
"_channels_stop",
"_deduplicate_run_name",
"_parse_todo_items",
"app",
"get_backend",
"main",
"normalize_ui_backend",
"resolve_ui_backend",
"run_streaming",
]
# Map attribute name -> (relative-module, attribute-in-module).
# Paths starting with ".." reach out of this package.
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
"StreamState": ("..stream.state", "StreamState"),
"SubAgentState": ("..stream.state", "SubAgentState"),
"_build_todo_stats": ("..stream.state", "_build_todo_stats"),
"_parse_todo_items": ("..stream.state", "_parse_todo_items"),
"WELCOME_SLOGANS": ("._constants", "WELCOME_SLOGANS"),
"_deduplicate_run_name": (".agent", "_deduplicate_run_name"),
"_channels_is_running": (".channel", "_channels_is_running"),
"_channels_stop": (".channel", "_channels_stop"),
"DEFAULT_UI_BACKEND": (".tui_runtime", "DEFAULT_UI_BACKEND"),
"SUPPORTED_UI_BACKENDS": (".tui_runtime", "SUPPORTED_UI_BACKENDS"),
"get_backend": (".tui_runtime", "get_backend"),
"normalize_ui_backend": (".tui_runtime", "normalize_ui_backend"),
"resolve_ui_backend": (".tui_runtime", "resolve_ui_backend"),
"run_streaming": (".tui_runtime", "run_streaming"),
}
def __getattr__(name: str):
target = _LAZY_EXPORTS.get(name)
if target is None:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
from importlib import import_module
module_path, attr = target
module = import_module(module_path, package=__name__)
value = getattr(module, attr)
globals()[name] = value
return value
# UI runtime re-exports (merged from former tui/ package)
from .tui_runtime import ( # noqa: F401
DEFAULT_UI_BACKEND,
SUPPORTED_UI_BACKENDS,
get_backend,
normalize_ui_backend,
resolve_ui_backend,
run_streaming,
)
def main():
@@ -69,12 +29,6 @@ def main():
import os
import warnings
# Keep MCP stdio subprocess spawning async on Windows (see #283). Must run
# before any event loop is created, hence at the very top of the entrypoint.
from .._winloop import ensure_proactor_event_loop_policy
ensure_proactor_event_loop_policy()
warnings.filterwarnings("ignore", message=".*not known to support tools.*")
warnings.filterwarnings(
"ignore", message=".*type is unknown and inference may fail.*"
@@ -87,3 +41,9 @@ def main():
_log_level = os.environ.get("EVOSCIENTIST_LOG_LEVEL", "") or config.log_level
_configure_logging()
app()
def admin_main():
"""evo-admin CLI entry point — admin management commands."""
from ._app import admin_app
admin_app()
-200
View File
@@ -1,200 +0,0 @@
"""Background MCP/agent load lifecycle shared by CLI and TUI surfaces.
Holds no references to Rich, prompt_toolkit, or Textual — UI-specific
rendering and thread-hopping plug in via callbacks.
"""
from __future__ import annotations
import asyncio
import logging
from collections.abc import Callable
from typing import Any, Generic, TypeVar
_logger = logging.getLogger(__name__)
ProgressEvent = str # "start" | "success" | "error"
ProgressState = str # "pending" | "ok" | "error"
AgentT = TypeVar("AgentT")
ProgressCallback = Callable[[ProgressEvent, str, str], None]
SuccessCallback = Callable[[AgentT], None]
FailureCallback = Callable[[BaseException], None]
class MCPProgressTracker:
"""Per-server MCP load progress state.
Reads and writes are GIL-atomic but iteration must go through
:meth:`snapshot` — events can fire from a worker thread while the
main thread renders.
"""
__slots__ = ("progress",)
def __init__(self) -> None:
self.progress: dict[str, tuple[ProgressState, str]] = {}
def prime(self) -> None:
"""Seed a ``pending`` entry for every configured server.
Keeps the UI's "N / M" denominator stable from the first render.
"""
try:
from ..mcp import load_mcp_config
cfg = load_mcp_config() or {}
self.progress = dict.fromkeys(cfg, ("pending", ""))
except Exception:
self.progress = {}
def record(
self, event: ProgressEvent, server: str, detail: str
) -> ProgressState | None:
"""Apply an event and return the new state, or ``None`` if unknown."""
if event == "start":
self.progress.setdefault(server, ("pending", ""))
return "pending"
if event == "success":
self.progress[server] = ("ok", detail)
return "ok"
if event == "error":
self.progress[server] = ("error", detail)
return "error"
return None
def snapshot(self) -> list[tuple[ProgressState, str]]:
return list(self.progress.values())
def totals(self) -> tuple[int, int]:
"""``(done, total)`` — done excludes ``pending``."""
snap = self.snapshot()
total = len(snap)
done = sum(1 for state, _ in snap if state != "pending")
return done, total
class BackgroundAgentLoader(Generic[AgentT]):
"""Owns the background ``_load_agent`` task and its generation token.
Each :meth:`start` bumps an internal id; callbacks from a superseded
load (the old worker thread keeps running after cancel, since
``asyncio.to_thread`` can't preempt arbitrary Python code) compare
against it and drop silently.
``on_progress`` fires on the **worker thread**; UI callers hop
threads inside it if needed. ``on_success`` / ``on_failure`` fire
on the event loop when the task completes.
"""
def __init__(
self,
loader_fn: Callable[..., AgentT],
*,
on_progress: ProgressCallback | None = None,
on_success: SuccessCallback | None = None,
on_failure: FailureCallback | None = None,
) -> None:
self._loader_fn = loader_fn
self._on_progress = on_progress
self._on_success = on_success
self._on_failure = on_failure
self.agent: AgentT | None = None
self._task: asyncio.Task[AgentT] | None = None
self._load_id: int = 0
@property
def task(self) -> asyncio.Task[AgentT] | None:
return self._task
@property
def is_pending(self) -> bool:
return self.agent is None and self._task is not None and not self._task.done()
@property
def needs_restart(self) -> bool:
"""True when no load is in flight and no agent is ready.
Callers that want auto-retry behavior (e.g. TUI on the next
user send after a failure) check this before :meth:`start`.
"""
return self.agent is None and (self._task is None or self._task.done())
def start(self, **loader_kwargs: Any) -> None:
prev = self._task
if prev is not None and not prev.done():
prev.cancel()
self._load_id += 1
load_id = self._load_id
self.agent = None
def _gated_progress(event: str, server: str, detail: str) -> None:
if load_id != self._load_id:
return
if self._on_progress is None:
return
try:
self._on_progress(event, server, detail)
except Exception:
_logger.debug("MCP progress callback raised", exc_info=True)
self._task = asyncio.create_task(
asyncio.to_thread(
self._loader_fn,
on_mcp_progress=_gated_progress,
**loader_kwargs,
)
)
self._task.add_done_callback(lambda task, lid=load_id: self._on_done(task, lid))
def adopt(self, agent: AgentT) -> None:
"""Install an externally-built agent and supersede any in-flight load.
Used by ``/model`` (and any other caller that constructs a
replacement agent directly): bumps the generation token so a
late-arriving background load can't clobber ``self.agent`` via
the done-callback, cancels the in-flight wrapper, and seats the
new agent immediately.
"""
prev = self._task
if prev is not None and not prev.done():
prev.cancel()
self._load_id += 1
self._task = None
self.agent = agent
async def await_ready(self) -> AgentT:
"""Return the loaded agent; re-raises on load failure.
Idempotent. State transitions (setting ``self.agent``, calling
``on_success`` / ``on_failure``) are handled exclusively by
:meth:`_on_done`, which fires before this ``await`` resumes
(asyncio guarantees done-callbacks run in registration order).
"""
if self.agent is not None:
return self.agent
if self._task is None:
raise RuntimeError(
"BackgroundAgentLoader.await_ready called before start()"
)
await self._task
if self.agent is None:
raise RuntimeError("BackgroundAgentLoader completed without an agent")
return self.agent
def _on_done(self, task: asyncio.Task[AgentT], load_id: int) -> None:
if load_id != self._load_id:
return
if task.cancelled():
return
try:
self.agent = task.result()
except Exception as exc:
# Keep ``_task`` set so a later ``await_ready`` re-raises the
# real exception instead of the "before start()" sentinel.
self.agent = None
if self._on_failure is not None:
self._on_failure(exc)
return
if self._on_success is not None:
self._on_success(self.agent)
+3 -15
View File
@@ -52,18 +52,6 @@ app.add_typer(mcp_app, name="mcp")
channel_app = typer.Typer(help="Channel management commands")
app.add_typer(channel_app, name="channel")
# Sessions subcommand group — diagnostic tools for the LangGraph checkpoint DB
sessions_app = typer.Typer(
help="Inspect and manage the sessions DB (~/.evoscientist/sessions.db)",
invoke_without_command=True,
)
app.add_typer(sessions_app, name="sessions")
# Configure subcommand group — re-run a single onboarding section.
configure_app = typer.Typer(
help=(
"Re-run one onboarding section without going through the full wizard.\n"
"Example: EvoSci configure provider"
),
)
app.add_typer(configure_app, name="configure")
# Admin subcommand group
admin_app = typer.Typer(help="Admin management commands")
app.add_typer(admin_app, name="admin")
+2 -16
View File
@@ -2,21 +2,7 @@
from datetime import UTC, datetime
def _agent_name() -> str:
# Deferred import: ``sessions`` pulls in langgraph/aiosqlite (~300 ms)
# and is only needed when ``build_metadata`` is actually called.
from ..sessions import AGENT_NAME
return AGENT_NAME
# Dangerous-mode warning banner — shared by Rich CLI, Textual TUI, and serve so
# the wording never drifts. Label is rendered white-on-red, message in red.
DANGEROUS_BANNER_LABEL = "DANGEROUS MODE"
DANGEROUS_BANNER_MESSAGE = (
"Real-filesystem access • the agent can read/write/delete anywhere."
)
from ..sessions import AGENT_NAME
WELCOME_SLOGANS = [
"Ready for vibe research? What do you want cooking?",
@@ -48,7 +34,7 @@ LOGO_GRADIENT = ["#1a237e", "#1565c0", "#1e88e5", "#42a5f5", "#64b5f6", "#90caf9
def build_metadata(workspace_dir: str | None, model: str | None) -> dict:
"""Build metadata dict for LangGraph checkpoint persistence."""
return {
"agent_name": _agent_name(),
"agent_name": AGENT_NAME,
"updated_at": datetime.now(UTC).isoformat(),
"workspace_dir": workspace_dir or "",
"model": model or "",
+3 -23
View File
@@ -3,13 +3,9 @@
import os
from datetime import datetime
from pathlib import Path
from typing import TYPE_CHECKING
from ..paths import new_run_dir
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
def _shorten_path(path: str) -> str:
"""Shorten absolute path to relative path from current directory."""
@@ -62,34 +58,18 @@ def _create_session_workspace(name: str | None = None) -> str:
return workspace_dir
def _load_agent(
workspace_dir: str | None = None,
checkpointer=None,
config=None,
chat_model=None,
*,
on_mcp_progress=None,
) -> "CompiledStateGraph":
def _load_agent(workspace_dir: str | None = None, checkpointer=None, config=None):
"""Load the CLI agent with optional persistent checkpointer.
Args:
workspace_dir: Optional per-session workspace directory.
checkpointer: Optional LangGraph checkpointer (e.g. ``AsyncSqliteSaver``).
checkpointer: Optional LangGraph checkpointer.
Falls back to ``InMemorySaver`` when ``None``.
config: Optional pre-loaded ``EvoScientistConfig``. Forwarded to
``create_cli_agent`` to avoid double config loading.
chat_model: Optional pre-built chat model. Forwarded to
``create_cli_agent``; combined with an explicit ``config`` it
selects the pure (no module-global write) build path.
on_mcp_progress: Optional per-server MCP progress callback.
Signature ``(event, server_name, detail) -> None``.
"""
from ..EvoScientist import create_cli_agent
return create_cli_agent(
workspace_dir=workspace_dir,
checkpointer=checkpointer,
config=config,
chat_model=chat_model,
on_mcp_progress=on_mcp_progress,
workspace_dir=workspace_dir, checkpointer=checkpointer, config=config
)
-595
View File
@@ -1,595 +0,0 @@
"""Async sub-agent auto-notification.
When a sub-agent on langgraph dev reaches a terminal state, a watcher coroutine
pushes a lightweight notification onto a thread-safe queue. The CLI loop drains
the queue, dedups against deepagents' async_tasks state, batches survivors,
and injects a synthetic user message that triggers one LLM turn.
"""
from __future__ import annotations
import asyncio
import json
import logging
import queue
import threading
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Final, TypeAlias, TypedDict
if TYPE_CHECKING:
from ..gateway import GraphGateway, GraphTarget
TERMINAL_STATUSES: Final = frozenset({"success", "error", "timeout", "interrupted"})
"""Aligned with langgraph_sdk.schema.RunStatus terminal values.
Cancel operations transition runs into ``interrupted`` (not ``cancelled``).
"""
# How many times the watcher will re-join the SSE stream when it closes
# cleanly but ``runs.get`` reports the run is still alive (typical cause:
# HTTP keep-alive timeout on long static periods). Bounded to prevent an
# unbounded loop if the server permanently misreports status.
_MAX_RECONNECT_ATTEMPTS: Final = 10
class AsyncTaskState(TypedDict, total=False):
status: str
last_checked_at: str
last_updated_at: str
AsyncTasksState: TypeAlias = dict[str, AsyncTaskState]
@dataclass(frozen=True)
class AsyncTaskNotification:
"""A completed-async-task signal pushed by a watcher."""
task_id: str
agent_name: str
status: str # one of TERMINAL_STATUSES
received_at: str # ISO-8601 UTC timestamp
prompt: str = "" # original task description sent to the sub-agent
kind: str = "agent" # "agent" (sub-agent) | "bg-process" (background shell)
# The CLI/main-agent thread_id under which the watcher was spawned. Used
# to route the notification back to the originating CLI session so a
# /new between launch and completion does not inject the synthetic
# message into an unrelated thread (where ``check_async_task`` cannot
# find the task_id). ``None`` means "unrouted" — the notification
# drains for any current_thread_id (back-compat for direct callers).
origin_cli_thread_id: str | None = None
# Per-thread routing: notifications with ``origin_cli_thread_id`` land in
# the matching sub-queue. Notifications without one go to ``_unrouted_queue``
# and drain regardless of current thread (back-compat for legacy callers
# and direct-put test paths).
_notifications_by_thread: dict[str, queue.Queue[AsyncTaskNotification]] = {}
_notifications_lock = threading.Lock()
_unrouted_queue: queue.Queue[AsyncTaskNotification] = queue.Queue()
# Public alias for the unrouted bucket — preserved so legacy tests and any
# external direct callers that did ``_notification_queue.put(...)`` keep
# working unchanged. New code should call ``_enqueue`` instead.
_notification_queue = _unrouted_queue
# Track active watcher tasks/futures for clean shutdown.
# dict[handle, origin_cli_thread_id] so the consumer's batching grace loop
# can filter for watchers tied to the current CLI thread (or unrouted)
# without being delayed by sibling-thread watchers.
_active_watchers: dict[object, str | None] = {}
# Map thread_id (sub-agent thread) → current watcher handle (supports
# replacement on update_async_task).
_watcher_by_thread: dict[str, asyncio.Task[None]] = {}
def _has_relevant_active_watchers(current_thread_id: str | None) -> bool:
"""Are there any in-flight watchers whose notifications would drain on
a ``consume_notifications`` call for ``current_thread_id``?
A watcher is relevant if its ``origin_cli_thread_id`` matches the
current CLI thread or is ``None`` (unrouted bucket drains for any
consumer). Sibling-thread watchers are ignored.
"""
if current_thread_id is None:
return bool(_active_watchers)
return any(
origin == current_thread_id or origin is None
for origin in _active_watchers.values()
)
logger = logging.getLogger(__name__)
def _enqueue(notification: AsyncTaskNotification) -> None:
"""Route a notification to its origin-thread queue or the unrouted bucket."""
tid = notification.origin_cli_thread_id
if not tid:
_unrouted_queue.put(notification)
return
with _notifications_lock:
q = _notifications_by_thread.get(tid)
if q is None:
q = queue.Queue()
_notifications_by_thread[tid] = q
q.put(notification)
def has_pending_notifications(current_thread_id: str | None = None) -> bool:
"""Cheap predicate for poller idle paths — true iff there's anything to consume.
If ``current_thread_id`` is given, only the matching thread queue and
the unrouted bucket count. With no argument, only the unrouted bucket
counts (legacy behavior).
"""
if not _unrouted_queue.empty():
return True
if current_thread_id is None:
return False
with _notifications_lock:
q = _notifications_by_thread.get(current_thread_id)
return q is not None and not q.empty()
def pending_thread_ids() -> set[str]:
"""Return the set of thread_ids with pending routed notifications."""
with _notifications_lock:
return {tid for tid, q in _notifications_by_thread.items() if not q.empty()}
async def read_async_tasks_from_gateway(
gateway: GraphGateway,
target: GraphTarget,
thread_id: str,
) -> AsyncTasksState:
"""Read async_tasks state through the active graph gateway."""
try:
values = await gateway.get_state_values(target, thread_id)
except Exception:
return {}
return values.get("async_tasks", {})
async def watch_run_and_notify(
client,
thread_id: str,
run_id: str,
agent_name: str,
prompt: str = "",
origin_cli_thread_id: str | None = None,
) -> None:
"""Subscribe to a run's event stream; enqueue notification when it terminates.
Status detection strategy (priority order):
1. **In-band ``event="error"`` SSE part** — authoritative error signal
from langgraph dev, no race against server-side state writeback.
2. **Server-side state via ``runs.get``** — invoked after the stream
closes (cleanly or with exception) to verify the run is actually
done. Required because SSE long-poll can close on HTTP keep-alive
timeout while the run is still running, which would otherwise be
misread as ``"success"`` (observed in production with long-running
literature search tasks under concurrency).
3. **Re-join loop** — if ``runs.get`` reports ``pending`` / ``running``,
the run is alive but we lost the stream; re-join up to
``_MAX_RECONNECT_ATTEMPTS`` times before giving up.
The previous implementation trusted clean stream exits as success
without any verification, which produced false-positive notifications
when SSE keep-alive timeouts closed the stream early.
Race-safety note: ``runs.get`` returning ``"error"`` immediately after
a clean stream close can be a transient state for an actually-successful
run (server hasn't finalized the writeback). We trust the absence of
in-band error event over a stale ``runs.get="error"`` — see the
``status == "error" and not saw_error_event`` branch below.
"""
for attempt in range(_MAX_RECONNECT_ATTEMPTS + 1):
stream_failed = False
saw_error_event = False
try:
async for chunk in client.runs.join_stream(
thread_id=thread_id, run_id=run_id, stream_mode="values"
):
ev = getattr(chunk, "event", None)
data = getattr(chunk, "data", None)
if ev == "error":
saw_error_event = True
logger.info(
"Watcher saw error event for task %s: %r", thread_id, data
)
except Exception:
stream_failed = True
logger.warning(
"Watcher stream failed for task %s", thread_id, exc_info=True
)
if saw_error_event:
status = "error"
break
# Verify with server before deciding the run is done — clean stream
# close does NOT guarantee terminal state.
try:
run = await client.runs.get(thread_id=thread_id, run_id=run_id)
raw = run.get("status", "")
except Exception:
# Cannot verify terminal state. Defaulting to "success" here would
# reintroduce the false-positive class this watcher exists to
# prevent (clean stream + transient runs.get failure → unverified
# success). Retry within the reconnect budget; on exhaustion drop
# the notification rather than guess.
if attempt >= _MAX_RECONNECT_ATTEMPTS:
logger.warning(
"Watcher runs.get failed for task %s after %d reconnects; "
"unable to verify terminal state, skipping notification",
thread_id,
_MAX_RECONNECT_ATTEMPTS,
exc_info=True,
)
return
logger.warning(
"Watcher runs.get failed for task %s; retrying after backoff "
"(attempt %d)",
thread_id,
attempt + 1,
exc_info=True,
)
await asyncio.sleep(min(0.25 * (attempt + 1), 2.0))
continue
if raw not in TERMINAL_STATUSES:
# Non-terminal status — includes the documented ``pending`` /
# ``running`` values AND any future / unknown status the SDK may
# introduce. Stream closed early but run is not done; re-join
# unless we've exhausted attempts. Treating unknown statuses as
# non-terminal is the safe default — better to retry once more
# than to enqueue a false-positive on an unrecognized state.
if attempt >= _MAX_RECONNECT_ATTEMPTS:
logger.warning(
"Watcher gave up on task %s after %d reconnects "
"(server still reports %r); skipping notification",
thread_id,
_MAX_RECONNECT_ATTEMPTS,
raw,
)
return
logger.info(
"Watcher SSE closed for task %s but run reports %r; "
"re-joining (attempt %d)",
thread_id,
raw,
attempt + 1,
)
continue
if raw == "error":
# Race-safe interpretation: no in-band error event → trust the
# absence over the server-side ``error`` (likely transient
# writeback state for a successful run). Stream-failure path
# is the one case where we DO trust ``error`` — the stream
# blowing up usually means something genuinely went wrong.
status = "error" if stream_failed else "success"
break
# success / timeout / interrupted — trust authoritative terminal status.
status = raw
break
else:
# Loop exhausted without a break — should be unreachable because the
# re-join branch returns explicitly when attempts are exhausted, but
# guard against future refactors.
return
notification = AsyncTaskNotification(
task_id=thread_id,
agent_name=agent_name,
status=status,
received_at=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"),
prompt=prompt,
origin_cli_thread_id=origin_cli_thread_id,
)
_enqueue(notification)
logger.info(
"Enqueued async notification: task=%s agent=%s status=%s origin_thread=%s",
thread_id,
agent_name,
status,
origin_cli_thread_id or "<unrouted>",
)
def spawn_watcher(
client,
thread_id: str,
run_id: str,
agent_name: str,
prompt: str = "",
origin_cli_thread_id: str | None = None,
) -> asyncio.Task[None]:
"""Spawn a watcher on the caller's asyncio loop.
Replacement semantics support ``update_async_task`` which creates a new
run_id on the same thread_id — we want the new watcher to take over
without the old (now obsolete) watcher firing a stale notification.
Cancellation propagates ``CancelledError`` (a BaseException), which the
watcher's ``except Exception:`` does NOT catch — so ``_enqueue(...)``
never executes for the cancelled watcher (no stale notification).
``origin_cli_thread_id`` tags the resulting notification so the consumer
only injects it back into the originating CLI session.
Caller must already be in a running asyncio event loop. Serve mode's
ephemeral per-turn loop kills watchers spawned during a turn — that
limitation is tracked separately.
"""
old_task = _watcher_by_thread.get(thread_id)
if old_task is not None and not old_task.done():
old_task.cancel()
task = asyncio.create_task(
watch_run_and_notify(
client,
thread_id,
run_id,
agent_name,
prompt,
origin_cli_thread_id=origin_cli_thread_id,
)
)
_watcher_by_thread[thread_id] = task
_active_watchers[task] = origin_cli_thread_id
def _cleanup(t: asyncio.Task[None]) -> None:
_active_watchers.pop(t, None)
# Only remove if THIS task is still the registered one — could
# have been replaced by a newer spawn_watcher call already.
if _watcher_by_thread.get(thread_id) is t:
del _watcher_by_thread[thread_id]
task.add_done_callback(_cleanup)
return task
def _drain_one_queue(q: queue.Queue) -> list[AsyncTaskNotification]:
items: list[AsyncTaskNotification] = []
while True:
try:
items.append(q.get_nowait())
except queue.Empty:
return items
def drain_notifications(
current_thread_id: str | None = None,
) -> list[AsyncTaskNotification]:
"""Pull pending notifications off the queue (non-blocking).
With ``current_thread_id``: drains the matching per-thread queue plus
the unrouted bucket. Without it: drains EVERY queue (legacy behavior;
used by tests and diagnostics).
"""
if current_thread_id is None:
items: list[AsyncTaskNotification] = _drain_one_queue(_unrouted_queue)
with _notifications_lock:
queues = list(_notifications_by_thread.values())
for q in queues:
items.extend(_drain_one_queue(q))
return items
items = _drain_one_queue(_unrouted_queue)
with _notifications_lock:
q = _notifications_by_thread.get(current_thread_id)
if q is not None:
items.extend(_drain_one_queue(q))
return items
def dedup_notifications(
notifs: list[AsyncTaskNotification],
async_tasks: AsyncTasksState | None,
) -> list[AsyncTaskNotification]:
"""Filter notifications the agent has already 'seen' via prior check.
Logic: skip a notification if `async_tasks[task_id]` exists with a TERMINAL
status and `last_checked_at >= last_updated_at` (timestamps are ISO-8601
so lexicographic comparison is correct). Also skip if `last_checked_at`
is empty (brand-new task where agent hasn't checked yet).
"""
from .. import background # cli -> core import; lazy to avoid import-order issues
async_tasks = async_tasks or {}
survivors: list[AsyncTaskNotification] = []
for n in notifs:
if n.kind == "bg-process":
# Background process: skip if the launching session already inspected it
# after it finished (check_process / list_processes) — mirrors the task
# dedup below. Per-thread: another session's check doesn't suppress this.
if background.was_observed_done(n.task_id, n.origin_cli_thread_id):
logger.debug("Dedup: skipping shell notification for %s", n.task_id)
continue
survivors.append(n)
continue
task = async_tasks.get(n.task_id)
if (
task
and task.get("status") in TERMINAL_STATUSES
and task.get("last_checked_at", "") >= task.get("last_updated_at", "")
and task.get("last_checked_at", "") != ""
):
logger.debug(
"Dedup: skipping notification for already-checked task %s", n.task_id
)
continue
survivors.append(n)
return survivors
def _render_notification_group(
notifs: list[AsyncTaskNotification], title: str, label: str
) -> list[tuple[str, str]]:
"""Render one group of notifications inside a titled open-right frame.
Open-right compact frame; bottom matches the top's width:
╭── ✦ Agent Teams ✦ ────
✔ writing Task: ... success
╰─────────────────────────
"""
top_divider = "╭──" + title + "────" # 4 dashes on the right (2x of left)
bottom_divider = "╰" + "─" * (len(top_divider) - 1)
lines: list[tuple[str, str]] = [(top_divider, "dim")]
for n in notifs:
# `writing-agent` → `writing`.
name = n.agent_name.removesuffix("-agent")
if n.status == "success":
icon, color = "✔", "#e67e22" # carrot orange (CSS hex; Rich+Textual)
elif n.status == "error":
icon, color = "✗", "red"
else: # cancelled, timeout, interrupted
icon, color = "⚠", "yellow"
# Collapse newlines, truncate prompt/command preview to 60 chars.
prompt_preview = (n.prompt or "").replace("\n", " ").strip()
if len(prompt_preview) > 60:
prompt_preview = prompt_preview[:60] + "…"
if prompt_preview:
text = f" {icon} {name:18s} {label}: {prompt_preview} {n.status}"
else:
# Fallback: short task_id when no prompt is available
short_tid = (
f"{n.task_id[:8]}…{n.task_id[-4:]}"
if len(n.task_id) > 12
else n.task_id
)
text = f" {icon} {name:18s} ({short_tid}) {n.status}"
lines.append((text, color))
lines.append((bottom_divider, "dim"))
return lines
def format_notification_lines(
notifs: list[AsyncTaskNotification],
) -> list[tuple[str, str]]:
"""Render notifications as compact tool-result-style lines for screen display.
Async sub-agents and background processes get SEPARATE titled frames so a shell
background process is never mislabeled as an "Agent Team". Returns (text, rich_style)
tuples. The LLM still receives the full ``format_batch_message`` text; this is purely
the visual representation for the human operator.
"""
if not notifs:
return []
tasks = [n for n in notifs if n.kind == "agent"]
shell = [n for n in notifs if n.kind == "bg-process"]
unknown = [n for n in notifs if n.kind not in {"agent", "bg-process"}]
lines: list[tuple[str, str]] = []
if tasks:
lines += _render_notification_group(tasks, " ✦ Agent Teams ✦ ", "Task")
if shell:
lines += _render_notification_group(shell, " ✦ Background ✦ ", "Cmd")
if unknown:
# Fallback so a future kind is never silently dropped from the display.
lines += _render_notification_group(unknown, " ✦ Updates ✦ ", "Task")
return lines
def format_batch_message(notifs: list[AsyncTaskNotification]) -> str:
"""Compose the synthetic user message that wakes the supervisor.
Each task is rendered as a compact JSON object (one per line) so the LLM
can reliably parse agent name, status, and task_id without ambiguity.
``ensure_ascii=False`` lets non-ASCII agent names pass through unchanged.
Visual decoration lives in ``format_notification_lines``.
"""
if not notifs:
return ""
lines = ["[Async tasks update]"]
for n in notifs:
lines.append(
json.dumps(
{
"agent": n.agent_name,
"kind": n.kind,
"status": n.status,
"task_id": n.task_id,
},
ensure_ascii=False,
)
)
# bg-process is inspected with check_process; sub-agents with check_async_task.
hints: list[str] = []
if any(n.kind == "agent" for n in notifs):
hints.append("check_async_task (sub-agents)")
if any(n.kind == "bg-process" for n in notifs):
hints.append("check_process (background processes)")
# Fallback when a batch has only unrecognized kinds (hints empty).
hint_text = " or ".join(hints) if hints else "the appropriate status tool"
lines.append(
f"(Signal only — fetch full result via {hint_text} if relevant to "
"the current step, else acknowledge & continue.)"
)
return "\n".join(lines)
# Brief grace window after the last drain: catch one final burst of arrivals
NOTIFICATION_BATCH_GRACE_SECONDS = 0.3
# Max time we'll wait for in-flight watchers to settle before triggering the
# agent turn — bounds latency for long-running tasks while still batching
# co-completing ones.
NOTIFICATION_ACTIVE_WATCHER_WAIT_SECONDS = 3.0
async def consume_notifications(
run_message: Callable[[str, list[AsyncTaskNotification]], Awaitable[None]],
read_async_tasks_state: Callable[[], Awaitable[AsyncTasksState]],
current_thread_id: str | None = None,
) -> None:
"""Drain queue, dedup, batch, and inject as a synthetic user message.
Args:
run_message: async callable receiving (llm_text, notifs_list).
``llm_text`` is the full structured message for the LLM
(from ``format_batch_message``). ``notifs_list`` is the
survivors list so callers can render per-task visual lines
without re-parsing the text.
read_async_tasks_state: async callable returning current ``async_tasks``
from the agent's state for dedup.
current_thread_id: the active CLI thread id. When given, only
notifications whose ``origin_cli_thread_id`` matches (or that
were enqueued unrouted) are drained — notifications belonging
to other threads stay queued and naturally drain on the next
poller tick after the user ``/resume``s back into them. When
omitted (legacy callers / tests), every queue drains.
"""
notifs = drain_notifications(current_thread_id)
if not notifs:
return
# Adaptive grace: if other watchers tied to THIS thread (or unrouted) are
# still in flight, wait briefly for them to settle so co-completing tasks
# batch into a single agent turn. Sibling-thread watchers don't count —
# their notifications wouldn't drain on this tick anyway.
loop = asyncio.get_running_loop()
deadline = loop.time() + NOTIFICATION_ACTIVE_WATCHER_WAIT_SECONDS
while _has_relevant_active_watchers(current_thread_id) and loop.time() < deadline:
await asyncio.sleep(0.2)
notifs.extend(drain_notifications(current_thread_id))
# Final brief grace to catch arrivals enqueued just before this tick
await asyncio.sleep(NOTIFICATION_BATCH_GRACE_SECONDS)
notifs.extend(drain_notifications(current_thread_id))
try:
async_tasks = await read_async_tasks_state()
except Exception:
logger.warning("Failed to read async_tasks state for dedup", exc_info=True)
async_tasks = {}
survivors = dedup_notifications(notifs, async_tasks)
if not survivors:
logger.info(
"All %d notifications deduped (already known to agent)", len(notifs)
)
return
text = format_batch_message(survivors)
await run_message(text, survivors)
+147 -587
View File
@@ -9,26 +9,20 @@ enqueues a ``ChannelMessage`` on a thread-safe ``queue.Queue`` and waits
for the main thread to set a response via ``_set_channel_response()``.
"""
from __future__ import annotations
import asyncio
import logging
import queue
import threading
import time
import uuid
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from typing import Any
from rich.panel import Panel
from rich.table import Table
from rich.text import Text
from ..commands.base import ChannelRuntime
from ..stream.console import console
if TYPE_CHECKING:
from ..gateway import GraphGateway
from ..stream.display import console
_channel_logger = logging.getLogger(__name__)
@@ -65,10 +59,6 @@ _response_lock = threading.Lock()
_RESPONSE_TIMEOUT = 600.0
_LATE_RESPONSE_TIMEOUT = 86400.0
_LATE_RESPONSE_NOTICE = "Still working on it. I'll send the result when it's ready."
_channel_request_lock = threading.Lock()
_channel_requests: dict[str, dict[str, str]] = {}
_session_requests: dict[str, list[str]] = {}
_cancelled_channel_messages: set[str] = set()
def _enqueue_channel_message(msg: ChannelMessage) -> asyncio.Future[str]:
@@ -81,7 +71,6 @@ def _enqueue_channel_message(msg: ChannelMessage) -> asyncio.Future[str]:
"loop": loop,
"response": None,
}
_register_channel_request(msg)
_message_queue.put(msg)
return future
@@ -117,336 +106,6 @@ def _pop_channel_response(msg_id: str, *, cancel_pending: bool = False) -> str |
return slot["response"]
def _channel_session_key(channel_type: str, chat_id: str) -> str:
return f"{channel_type}:{chat_id}"
def _channel_message_session_key(msg: ChannelMessage) -> str:
return _channel_session_key(msg.channel_type, msg.chat_id)
def _channel_message_cancel_scope(msg: ChannelMessage) -> str:
return f"channel:{msg.channel_type}:{msg.chat_id}:{msg.msg_id}"
def _register_channel_request(msg: ChannelMessage) -> None:
"""Track a queued channel request so `/stop` can find it later."""
session_key = _channel_message_session_key(msg)
with _channel_request_lock:
_channel_requests[msg.msg_id] = {
"session_key": session_key,
"cancel_scope": _channel_message_cancel_scope(msg),
"state": "queued",
}
_session_requests.setdefault(session_key, []).append(msg.msg_id)
def _claim_channel_request(msg: ChannelMessage) -> bool:
"""Mark a queued request active. Returns False if it was cancelled first."""
with _channel_request_lock:
slot = _channel_requests.get(msg.msg_id)
if slot is None or msg.msg_id in _cancelled_channel_messages:
return False
slot["state"] = "active"
return True
def _claim_or_complete_channel_request(msg: ChannelMessage) -> bool:
"""Claim a request, or clean it up if `/stop` cancelled it while queued."""
if _claim_channel_request(msg):
return True
_complete_channel_request(msg.msg_id)
return False
def _channel_request_state(msg_id: str) -> str | None:
with _channel_request_lock:
slot = _channel_requests.get(msg_id)
return slot.get("state") if slot is not None else None
def _complete_channel_request(
msg_id: str,
*,
discard_cancel_scope: bool = True,
) -> None:
"""Forget a request once its waiter is resolved or cancelled."""
with _channel_request_lock:
slot = _channel_requests.pop(msg_id, None)
_cancelled_channel_messages.discard(msg_id)
if slot is not None:
request_ids = _session_requests.get(slot["session_key"])
if request_ids:
try:
request_ids.remove(msg_id)
except ValueError:
pass
if not request_ids:
_session_requests.pop(slot["session_key"], None)
if slot is not None and discard_cancel_scope:
from ..stream.display import discard_stream_cancel
discard_stream_cancel(slot["cancel_scope"])
def _cancel_channel_session(channel_type: str, chat_id: str) -> tuple[int, int]:
"""Cancel queued and active work for one channel chat session."""
session_key = _channel_session_key(channel_type, chat_id)
with _channel_request_lock:
request_ids: list[str] = []
cancelled_ids: list[str] = []
active_scopes: list[str] = []
with _response_lock:
for msg_id in tuple(_session_requests.get(session_key, ())):
request_slot = _channel_requests.get(msg_id)
if request_slot is None:
continue
response_slot = _pending_responses.get(msg_id)
response_resolved = False
if response_slot is not None:
future = response_slot["future"]
# Once a response is already resolved, leave the slot alone
# so the bus waiter can still publish it instead of falling
# back to "No response".
response_resolved = (
response_slot.get("response") is not None or future.done()
)
if not response_resolved:
request_ids.append(msg_id)
should_cancel = False
if response_slot is None:
should_cancel = request_slot.get("state") == "active"
else:
should_cancel = not response_resolved
if should_cancel:
cancelled_ids.append(msg_id)
if request_slot.get("state") == "active" and should_cancel:
active_scopes.append(request_slot["cancel_scope"])
_cancelled_channel_messages.update(cancelled_ids)
for msg_id in request_ids:
_pop_channel_response(msg_id, cancel_pending=True)
if active_scopes:
from ..stream.display import request_stream_cancel
for cancel_scope in active_scopes:
request_stream_cancel(cancel_scope)
return len(request_ids), len(active_scopes)
# ---------------------------------------------------------------------------
# Slash command dispatch for channel messages
# ---------------------------------------------------------------------------
# Shared by all three UI surfaces that accept inbound channel messages:
# Rich CLI (``cli/interactive.py::_process_channel_message``), Textual
# TUI (``cli/tui_interactive.py``'s channel handler), and headless
# serve (``cli/commands.py::_serve_process_message``). They all route
# ``/foo`` text through ``cmd_manager`` instead of feeding it to the
# LLM as a plain prompt.
async def dispatch_channel_slash_command(
msg: ChannelMessage,
*,
agent: Any,
thread_id: str,
workspace_dir: str | None,
checkpointer: Any,
append_system: Callable[[str, str], None],
graph_gateway: GraphGateway,
start_new_session_cb: Callable[[], Awaitable[None]] | None = None,
handle_session_resume_cb: Callable[..., Awaitable[None]] | None = None,
await_agent_ready: Callable[[], Awaitable[Any]] | None = None,
on_cmd_completed: Callable[..., Awaitable[None]] | None = None,
channel_runtime: ChannelRuntime | None = None,
) -> bool:
"""Dispatch a slash command from a channel message.
Returns True if the helper handled the message (successfully or with
an error) — the caller must then return without streaming anything
to the agent. Returns False for non-slash content or unresolved
slash commands, so the caller can fall through to the agent
streaming path (matches TUI behavior).
Parameters
----------
msg:
The inbound ``ChannelMessage`` to inspect.
agent:
Default agent handle for the ``CommandContext``. Commands that
do not need the agent use this value directly.
thread_id, workspace_dir, checkpointer:
Populate ``CommandContext``.
append_system:
``(text, style)`` callback for local CLI/TUI log output. Used
by ``ChannelCommandUI`` to surface system breadcrumbs and by
this helper to print the "Executed command from ..." line.
start_new_session_cb, handle_session_resume_cb:
Optional lifecycle callbacks forwarded to ``ChannelCommandUI``.
Headless serve passes ``None`` — ``/new`` and ``/resume`` degrade
gracefully via the default ``ChannelCommandUI`` messages.
graph_gateway:
Graph gateway forwarded to slash commands and channel resume-history
rendering.
await_agent_ready:
Optional async resolver that blocks until the background agent
load finishes. Called only when ``cmd.needs_agent(args)`` is
True. Headless serve passes ``None`` because the agent is
loaded up-front before the bus starts.
on_cmd_completed:
Optional ``async (ctx, original_agent, cmd) -> None`` callback
fired only after ``cmd_manager.execute`` returns True. The
``original_agent`` argument is the agent handle command execution
started against: ``agent_for_ctx`` after any ``await_agent_ready``
resolution, or the dispatcher's input agent when no resolver is
supplied. Callers can compare ``ctx.agent`` with
``original_agent`` to detect command-driven swaps. Used by Rich
CLI to (a) adopt an agent swap (``/model``) back into the
running session and (b) refresh the status snapshot for
commands that mutate session-level state (``/new``,
``/compact``) — mirrors the REPL dispatch at
``cli/interactive.py:1002-1030``. Headless serve passes
``None`` since it cannot hot-swap its polling-loop agent.
"""
if not msg.content.strip().startswith("/"):
return False
try:
return await _dispatch_channel_slash_impl(
msg,
agent=agent,
thread_id=thread_id,
workspace_dir=workspace_dir,
checkpointer=checkpointer,
append_system=append_system,
start_new_session_cb=start_new_session_cb,
handle_session_resume_cb=handle_session_resume_cb,
await_agent_ready=await_agent_ready,
on_cmd_completed=on_cmd_completed,
channel_runtime=channel_runtime,
graph_gateway=graph_gateway,
)
except Exception as exc:
# Last-ditch safety: any uncaught exception from inside the
# dispatch pipeline (lazy import failure, ChannelCommandUI
# construction, terminal I/O from ``append_system``, bus
# publish races, ...) must not take down the caller's polling
# loop — a crashed serve / dead channel queue task is worse
# than one failed command.
_channel_logger.exception(
"Unexpected slash dispatch failure for %s (msg=%s)",
msg.channel_type,
msg.msg_id,
)
try:
_set_channel_response(msg.msg_id, f"Command error: {exc}")
except Exception: # pragma: no cover — defensive
pass
# Return True so the caller treats the message as handled and
# does not fall through to the agent streaming path.
return True
async def _dispatch_channel_slash_impl(
msg: ChannelMessage,
*,
agent: Any,
thread_id: str,
workspace_dir: str | None,
checkpointer: Any,
append_system: Callable[[str, str], None],
graph_gateway: GraphGateway,
start_new_session_cb: Callable[[], Awaitable[None]] | None,
handle_session_resume_cb: Callable[..., Awaitable[None]] | None,
await_agent_ready: Callable[[], Awaitable[Any]] | None,
on_cmd_completed: Callable[..., Awaitable[None]] | None,
channel_runtime: ChannelRuntime | None,
) -> bool:
"""Inner body of ``dispatch_channel_slash_command``.
Split from the public wrapper so the wrapper can guard with a
top-level try/except without visually obscuring the main flow.
"""
# Lazy imports: avoid coupling the channel module to ``commands`` at
# import time (tui_interactive.py does the same).
from ..commands.base import CommandContext
from ..commands.channel_ui import ChannelCommandUI
from ..commands.manager import manager as cmd_manager
parsed = cmd_manager.resolve(msg.content)
if parsed is None:
# Unknown slash command — let the agent handle it (matches TUI).
return False
cmd, cmd_args = parsed
agent_for_ctx = agent
if cmd.needs_agent(cmd_args) and await_agent_ready is not None:
try:
agent_for_ctx = await await_agent_ready()
except Exception as exc:
_set_channel_response(msg.msg_id, f"Command error: {exc}")
return True
ui = ChannelCommandUI(
msg,
append_system_callback=append_system,
start_new_session_callback=start_new_session_cb,
handle_session_resume_callback=handle_session_resume_cb,
graph_gateway=graph_gateway,
)
ctx = CommandContext(
agent=agent_for_ctx,
thread_id=thread_id,
ui=ui,
workspace_dir=workspace_dir,
checkpointer=checkpointer,
channel_runtime=channel_runtime,
graph_gateway=graph_gateway,
)
try:
cmd_executed = await cmd_manager.execute(msg.content, ctx)
except Exception as exc:
_channel_logger.debug(f"Channel command error: {exc}", exc_info=True)
_set_channel_response(msg.msg_id, f"Command error: {exc}")
return True # must return — do NOT fall through to the agent
if cmd_executed:
if ctx.command_error is not None:
details = ctx.command_error or "(no details)"
_set_channel_response(msg.msg_id, f"Command error: {details}")
return True
if on_cmd_completed is not None:
try:
# Command output already flushed by ``cmd_manager.execute``
# via ``ctx.ui.flush()`` — the hook does internal state
# sync (agent adoption, status snapshot refresh) only,
# so swallowing its errors keeps the user-visible reply
# intact even if the sync path is broken.
await on_cmd_completed(ctx, agent_for_ctx, cmd)
except Exception as exc:
_channel_logger.debug(
f"Channel command post-exec callback error: {exc}",
exc_info=True,
)
append_system(
f"[{msg.channel_type}: Executed command from {msg.sender}]",
"dim",
)
_set_channel_response(msg.msg_id, f"Command executed: {msg.content}")
return True
# ``cmd_manager.execute`` returned False (empty / unparseable input).
# Fall through to the agent streaming path.
return False
# ---------------------------------------------------------------------------
# HITL approval intercept: bus thread ⇄ main CLI thread
# ---------------------------------------------------------------------------
@@ -462,143 +121,6 @@ _HITL_APPROVAL_TIMEOUT = 120.0 # seconds to wait for HITL approval reply
_ASK_USER_TIMEOUT = (
300.0 # seconds to wait for ask_user reply (longer for thinking time)
)
_STOP_COMMANDS = frozenset(("/stop", "/cancel"))
# ---------------------------------------------------------------------------
# Per-thread channel-origin registry
# ---------------------------------------------------------------------------
# When a channel-originated message starts an agent turn, the turn's
# thread_id is remembered against its (channel_type, chat_id, metadata).
# Later, when an async sub-agent notification fires a synthetic agent turn
# for that same thread_id, the notifier path pushes the synthesized final
# response back to the same chat — otherwise the follow-up would only render
# locally and the channel user would never see it. v1 forwards only the
# final response (no mid-turn thinking/todo/media).
@dataclass(frozen=True)
class _ChannelOrigin:
"""Channel destination remembered for a thread, for notifier push-back."""
channel_type: str
chat_id: str
sender: str
metadata: dict | None = None
_thread_channel_origins: dict[str, _ChannelOrigin] = {}
_thread_channel_origins_lock = threading.Lock()
def remember_channel_origin(thread_id: str | None, msg: ChannelMessage) -> None:
"""Record that ``thread_id`` is currently bound to ``msg``'s channel chat.
Called on entry to each channel-triggered agent turn (Rich CLI / TUI /
serve). The latest channel turn for a given thread wins — re-registering
is intentional, since the user can keep talking on the same thread from
the same channel and we always want the most recent metadata.
"""
if not thread_id:
return
with _thread_channel_origins_lock:
_thread_channel_origins[thread_id] = _ChannelOrigin(
channel_type=msg.channel_type,
chat_id=msg.chat_id,
sender=msg.sender,
metadata=dict(msg.metadata) if msg.metadata else None,
)
def get_channel_origin(thread_id: str | None) -> _ChannelOrigin | None:
"""Return the channel origin remembered for ``thread_id``, or ``None``."""
if not thread_id:
return None
with _thread_channel_origins_lock:
return _thread_channel_origins.get(thread_id)
def forget_channel_origin(thread_id: str | None) -> None:
"""Drop the registry entry for ``thread_id`` (e.g. on ``/new`` rotation)."""
if not thread_id:
return
with _thread_channel_origins_lock:
_thread_channel_origins.pop(thread_id, None)
def publish_to_channel_origin(thread_id: str | None, content: str) -> bool:
"""Schedule pushing ``content`` to the channel remembered for ``thread_id``.
Fire-and-forget: returns ``True`` iff a publish coroutine was scheduled
on the bus loop; returns ``False`` if no origin is registered, the bus
isn't running, ``content`` is empty/whitespace, or scheduling itself
fails. The publish runs asynchronously — failures inside the coroutine
are logged via a done-callback so callers (which are often on event
loops that must not block) don't pay any latency.
"""
from ..channels.bus.events import OutboundMessage
if not content or not content.strip():
return False
origin = get_channel_origin(thread_id)
if origin is None:
return False
loop = _bus_loop
manager = _manager
if loop is None or manager is None:
return False
bus = getattr(manager, "bus", None)
if bus is None:
return False
async def _publish_and_record() -> None:
await bus.publish_outbound(
OutboundMessage(
channel=origin.channel_type,
chat_id=origin.chat_id,
content=content,
metadata=origin.metadata or {},
)
)
# Mirror the normal channel-reply path, which records a "sent"
# message after a successful publish so per-channel stats stay
# accurate for forwarded notifications too.
manager.record_message(origin.channel_type, "sent")
try:
future = asyncio.run_coroutine_threadsafe(_publish_and_record(), loop)
except Exception as exc:
_channel_logger.warning(
"Async notification publish to %s:%s failed to schedule: %s",
origin.channel_type,
origin.chat_id,
exc,
)
return False
def _on_publish_done(fut) -> None:
"""Log any exception raised by the fire-and-forget publish coroutine."""
# A cancelled future raises CancelledError from .exception() rather
# than returning it (e.g. bus loop torn down mid-publish); treat that
# as a benign shutdown, not a failure to log.
if fut.cancelled():
return
exc = fut.exception()
if exc is not None:
_channel_logger.warning(
"Async notification publish to %s:%s failed: %s",
origin.channel_type,
origin.chat_id,
exc,
)
future.add_done_callback(_on_publish_done)
return True
def _is_stop_command(content: str | None) -> bool:
"""Whether incoming content is a stop/cancel slash command."""
return (content or "").strip().lower() in _STOP_COMMANDS
def _register_hitl_wait(channel_type: str, chat_id: str) -> threading.Event:
@@ -632,7 +154,7 @@ def _try_set_hitl_reply(channel_type: str, chat_id: str, content: str) -> bool:
def channel_ask_user_prompt(
ask_user_data: dict,
msg: ChannelMessage | None = None,
msg: "ChannelMessage | None" = None,
) -> dict:
"""Format ask_user questions and collect answers from a channel user.
@@ -664,7 +186,7 @@ def channel_ask_user_prompt(
channel=msg.channel_type,
chat_id=msg.chat_id,
content=content,
metadata=msg.metadata or {},
metadata=msg.metadata,
)
),
bus_loop,
@@ -721,8 +243,6 @@ def channel_ask_user_prompt(
return {"status": "cancelled"}
raw = reply_text.strip()
if _is_stop_command(raw):
return {"status": "cancelled"}
if raw.lower() == "cancel":
return {"status": "cancelled"}
@@ -740,8 +260,6 @@ def channel_ask_user_prompt(
if not replied or not other_text:
_send("\u23f0 Response timed out.")
return {"status": "cancelled"}
if _is_stop_command(other_text):
return {"status": "cancelled"}
if other_text.strip().lower() == "cancel":
return {"status": "cancelled"}
answers.append(other_text.strip())
@@ -761,7 +279,7 @@ def channel_ask_user_prompt(
def channel_hitl_prompt(
action_requests: list,
msg: ChannelMessage,
msg: "ChannelMessage",
) -> list[dict] | None:
"""Send HITL approval prompt to channel user and wait for reply.
@@ -772,7 +290,6 @@ def channel_hitl_prompt(
"""
from ..channels.bus.events import OutboundMessage
from ..channels.consumer import (
_approval_prompt_metadata,
_format_approval_prompt,
_parse_approval_reply,
)
@@ -787,17 +304,7 @@ def channel_hitl_prompt(
_channel_logger.debug("HITL: no bus_loop or bus_ref, rejecting")
return None
# Look up the channel instance so we can attach buttons when the channel
# supports `inline_buttons` (Feishu cards, QQ keyboards, …).
channel_obj = (
_manager.get_channel(msg.channel_type) if _manager is not None else None
)
has_buttons = channel_obj is not None and channel_obj.capabilities.inline_buttons
approval_metadata = _approval_prompt_metadata(
msg.metadata, with_buttons=has_buttons
)
def _send(content: str, *, metadata: dict | None = None) -> bool:
def _send(content: str) -> bool:
"""Send a message to the channel user. Returns True on success."""
try:
asyncio.run_coroutine_threadsafe(
@@ -806,9 +313,7 @@ def channel_hitl_prompt(
channel=msg.channel_type,
chat_id=msg.chat_id,
content=content,
metadata=metadata
if metadata is not None
else msg.metadata or {},
metadata=msg.metadata,
)
),
bus_loop,
@@ -819,8 +324,8 @@ def channel_hitl_prompt(
return False
# 1. Send approval prompt
prompt_text = _format_approval_prompt(action_requests, with_buttons=has_buttons)
if not _send(prompt_text, metadata=approval_metadata):
prompt_text = _format_approval_prompt(action_requests)
if not _send(prompt_text):
return None
# 2. Wait for channel user's reply
@@ -832,24 +337,16 @@ def channel_hitl_prompt(
_send("\u23f0 Approval timed out. Action rejected.")
return None
if _is_stop_command(reply_text):
# `/stop` already got its own immediate ack from the bus fast-path.
# Treat it as a pure cancel signal here so we don't send a second,
# contradictory "Unrecognized reply" message.
return None
# 3. Parse decision
decision = _parse_approval_reply(reply_text)
if decision == "auto":
_hitl_auto_approve.add(session_key)
_send("\u2705 已批准(后续自动通过)")
return [{"type": "approve"} for _ in action_requests]
if decision == "approve":
_send("\u2705 已批准")
return [{"type": "approve"} for _ in action_requests]
feedback = (
"\u274c 已拒绝"
"Action rejected."
if decision == "reject"
else "Unrecognized reply. Action rejected."
)
@@ -864,6 +361,8 @@ def channel_hitl_prompt(
_manager: Any | None = None # ChannelManager
_bus_loop: asyncio.AbstractEventLoop | None = None
_bus_thread: threading.Thread | None = None
_cli_agent: Any = None # shared agent reference (same as CLI)
_cli_thread_id: str | None = None # shared thread_id (same conversation)
def _channels_is_running(channel_type: str | None = None) -> bool:
@@ -881,18 +380,9 @@ def _channels_running_list() -> list[str]:
return _manager.running_channels() if _manager else []
def _channels_stop(
channel_type: str | None = None,
*,
runtime: ChannelRuntime | None = None,
) -> None:
"""Stop channel(s) and clean up module-level state.
``runtime`` is the ``ChannelRuntime`` whose binding should be
cleared once the channels are gone — the caller owns it (commands
keep a reference via ``ctx.channel_runtime``).
"""
global _manager, _bus_loop, _bus_thread
def _channels_stop(channel_type: str | None = None) -> None:
"""Stop channel(s) and clean up module-level state."""
global _manager, _bus_loop, _bus_thread, _cli_agent, _cli_thread_id
if channel_type is None:
# Stop everything
@@ -905,13 +395,15 @@ def _channels_stop(
future.result(timeout=10)
except Exception as e:
_channel_logger.debug(f"Error stopping channels: {e}")
if _manager:
_manager.bus.stop()
if _bus_thread:
_bus_thread.join(timeout=5)
_manager = None
_bus_loop = None
_bus_thread = None
if runtime is not None:
runtime.clear()
_cli_agent = None
_cli_thread_id = None
return
# Stop a specific channel
@@ -925,8 +417,9 @@ def _channels_stop(
except Exception as e:
_channel_logger.debug(f"Error removing channel {channel_type}: {e}")
if _manager and not _manager.running_channels() and runtime is not None:
runtime.clear()
if _manager and not _manager.running_channels():
_cli_agent = None
_cli_thread_id = None
def _start_channels_bus_mode(
@@ -1040,20 +533,6 @@ async def _bus_inbound_consumer(bus, manager) -> None:
except asyncio.CancelledError:
break
# /stop should preempt HITL interception so cancel works while
# waiting for approvals/questions. If a HITL wait is pending,
# still release it so the blocking prompt can unwind immediately.
if _is_stop_command(msg.content):
if _try_set_hitl_reply(msg.channel, msg.chat_id, msg.content):
_channel_logger.info(
f"[bus] stop request released HITL wait for "
f"{msg.channel}:{msg.chat_id}"
)
_task = asyncio.create_task(_handle_bus_message(bus, manager, msg))
_tasks.add(_task)
_task.add_done_callback(_tasks.discard)
continue
# Check if this message is a HITL approval reply
if _try_set_hitl_reply(msg.channel, msg.chat_id, msg.content):
_channel_logger.info(
@@ -1082,37 +561,6 @@ async def _handle_bus_message(bus, manager, msg) -> None:
)
manager.record_message(msg.channel, "received")
# Fast-path: /stop intercept. Handle on the bus task itself so we
# don't deadlock behind the main-thread stream we're trying to
# interrupt. No typing indicator, no queue entry.
if _is_stop_command(msg.content):
cancelled_count, active_count = _cancel_channel_session(
msg.channel, msg.chat_id
)
try:
await bus.publish_outbound(
OutboundMessage(
channel=msg.channel,
chat_id=msg.chat_id,
content="Stopped.",
reply_to=msg.message_id or None,
metadata=msg.metadata,
)
)
manager.record_message(msg.channel, "sent")
except Exception as e:
_channel_logger.error(f"[bus] /stop ack send error: {e}")
else:
if cancelled_count or active_count:
_channel_logger.info(
"[bus] /stop cancelled %d request(s) (%d active) for %s:%s",
cancelled_count,
active_count,
msg.channel,
msg.chat_id,
)
return
channel = manager.get_channel(msg.channel)
typing_active = False
if channel:
@@ -1182,8 +630,6 @@ async def _handle_bus_message(bus, manager, msg) -> None:
f"for {cm.msg_id}"
)
_pop_channel_response(cm.msg_id, cancel_pending=True)
if _channel_request_state(cm.msg_id) != "active":
_complete_channel_request(cm.msg_id)
return
response = _pop_channel_response(cm.msg_id) or "No response"
@@ -1199,8 +645,6 @@ async def _handle_bus_message(bus, manager, msg) -> None:
manager.record_message(msg.channel, "sent")
except asyncio.CancelledError:
_pop_channel_response(cm.msg_id, cancel_pending=True)
if _channel_request_state(cm.msg_id) != "active":
_complete_channel_request(cm.msg_id)
raise
except Exception as e:
_channel_logger.error(f"[bus] Outbound error: {e}")
@@ -1238,13 +682,131 @@ def _print_channel_panel(channels: list[tuple[str, bool, str]]) -> None:
console.print()
def _cmd_channel(
args: str,
agent: Any,
thread_id: str,
*,
send_thinking: bool | None = None,
) -> None:
"""Start a channel in background using bus mode.
Usage:
/channel [telegram|discord|imessage] -- start channel (default from config)
/channel status -- show current channel status
/channel stop -- stop running channel
"""
global _cli_agent, _cli_thread_id
from ..config import load_config
app_config = load_config()
channel_type = args.strip().lower() if args and args.strip() else ""
if channel_type == "status":
running = _channels_running_list()
if running and _manager:
detailed = _manager.get_detailed_status()
table = Table(title="Channel Status", show_header=True, expand=False)
table.add_column("Channel", style="cyan")
table.add_column("Status")
table.add_column("Uptime", style="dim")
table.add_column("Rx", justify="right")
table.add_column("Tx", justify="right")
for ch_name in running:
info = detailed.get(ch_name, {})
secs = info.get("uptime_seconds", 0)
mins, s = divmod(int(secs), 60)
hours, mins = divmod(mins, 60)
uptime = f"{hours}h{mins:02d}m" if hours else f"{mins}m{s:02d}s"
rx = str(info.get("received", 0))
tx = str(info.get("sent", 0))
table.add_row(ch_name, "[green]running[/green]", uptime, rx, tx)
console.print(table)
console.print()
else:
console.print("[dim]No channel running[/dim]\n")
return
if not channel_type:
channel_type = app_config.channel_enabled
if not channel_type:
console.print("[yellow]No channel configured.[/yellow]")
console.print(
"[dim]Run[/dim] evosci onboard [dim]or specify:[/dim] /channel telegram\n"
)
return
requested = [t.strip() for t in channel_type.split(",") if t.strip()]
if _channels_is_running():
running = _channels_running_list()
results: list[tuple[str, bool, str]] = []
for ct in requested:
if ct in running:
results.append((ct, True, "already running"))
else:
try:
_add_channel_to_running_bus(
ct,
app_config,
send_thinking=send_thinking,
)
results.append((ct, True, "connected (bus)"))
except Exception as e:
results.append((ct, False, str(e)))
_print_channel_panel(results)
return
_cli_agent = agent
_cli_thread_id = thread_id
# Override channel_enabled for this invocation
original = app_config.channel_enabled
app_config.channel_enabled = channel_type
try:
_start_channels_bus_mode(
app_config,
agent,
thread_id,
send_thinking=send_thinking,
)
results = [(ct, True, "connected (bus)") for ct in requested]
except Exception as e:
results = [(ct, False, str(e)) for ct in requested]
finally:
app_config.channel_enabled = original
_print_channel_panel(results)
def _cmd_channel_stop(channel_type: str | None = None) -> None:
"""Stop background channel(s).
Args:
channel_type: Specific channel to stop, or None to stop all.
"""
if not _channels_is_running():
console.print("[dim]No channel running[/dim]\n")
return
if channel_type:
if not _channels_is_running(channel_type):
console.print(f"[dim]{channel_type} is not running[/dim]\n")
return
_channels_stop(channel_type)
console.print(f"[dim]{channel_type} stopped[/dim]\n")
else:
running = _channels_running_list()
_channels_stop()
console.print(f"[dim]{', '.join(running)} stopped[/dim]\n")
def _auto_start_channel(
agent: Any,
thread_id: str,
config,
*,
send_thinking: bool | None = None,
runtime: ChannelRuntime | None = None,
) -> None:
"""Start channels automatically from config (bus mode).
@@ -1252,23 +814,21 @@ def _auto_start_channel(
agent: Compiled agent graph.
thread_id: Current thread ID.
config: EvoScientistConfig with channel settings.
runtime: Caller-owned ``ChannelRuntime`` to bind so commands
running over the channels can swap the agent later. ``None``
is accepted for callers that don't yet pass one.
"""
global _cli_agent, _cli_thread_id
if not config.channel_enabled:
return
_cli_agent = agent
_cli_thread_id = thread_id
_start_channels_bus_mode(
config,
agent,
thread_id,
send_thinking=send_thinking,
)
# Bind only after startup succeeds; a failure above must not leave
# a stale runtime binding pointing at channels that never started.
if runtime is not None:
runtime.bind(agent, thread_id)
types = [t.strip() for t in config.channel_enabled.split(",") if t.strip()]
results = [(ct, True, "connected (bus)") for ct in types]
_print_channel_panel(results)
+16 -45
View File
@@ -21,7 +21,7 @@ import os
import pathlib
import subprocess
import sys
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from textual.app import App
@@ -29,15 +29,6 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
_PREVIEW_MAX = 40
_pyperclip_notify_shown = False
def _is_remote_session() -> bool:
"""Return True when running over SSH without a local display."""
if os.environ.get("SSH_CLIENT") or os.environ.get("SSH_CONNECTION"):
if not os.environ.get("DISPLAY") and not os.environ.get("WAYLAND_DISPLAY"):
return True
return False
# ── Platform-native clipboard read ────────────────────────────────
@@ -141,6 +132,8 @@ def copy_selection_to_clipboard(app: App) -> None:
if not hasattr(widget, "text_selection") or not widget.text_selection:
continue
selection = widget.text_selection
if selection.end is None:
continue
try:
result = widget.get_selection(selection)
except (AttributeError, TypeError, ValueError, IndexError) as exc:
@@ -161,59 +154,37 @@ def copy_selection_to_clipboard(app: App) -> None:
combined = "\n".join(selected_texts)
# Build method list: (fn, reliable) — reliable means we *know* the text
# reached the system clipboard (e.g. pyperclip). OSC 52 / Textual write
# to the terminal and succeed even when the terminal silently ignores the
# sequence (PuTTY, older terminals).
copy_methods: list[tuple[Any, bool]] = [
(app.copy_to_clipboard, False),
]
# Try methods in priority order
copy_methods = [app.copy_to_clipboard]
try:
import pyperclip
copy_methods.insert(0, (pyperclip.copy, True))
copy_methods.insert(0, pyperclip.copy)
except ImportError:
global _pyperclip_notify_shown
if not _pyperclip_notify_shown:
_pyperclip_notify_shown = True
app.notify(
'Failed to import "pyperclip", text copying might not work.',
severity="information",
timeout=3,
)
pass
copy_methods.append((_copy_osc52, False))
copy_methods.append(_copy_osc52)
remote = _is_remote_session()
for fn, reliable in copy_methods:
for fn in copy_methods:
try:
fn(combined)
except (OSError, RuntimeError, TypeError) as exc:
logger.debug(
"Clipboard method %s failed: %s", getattr(fn, "__name__", repr(fn)), exc
)
continue
if reliable or not remote:
app.notify(
f'"{_shorten(selected_texts)}" copied',
severity="information",
timeout=2,
markup=False,
)
else:
# OSC 52 over SSH — may be silently ignored (e.g. Windows Terminal, PuTTY)
app.notify(
"Copied text - if paste fails, use Shift+mouse-select for native copy",
severity="information",
timeout=3,
except (OSError, RuntimeError, TypeError) as exc:
logger.debug(
"Clipboard method %s failed: %s", getattr(fn, "__name__", repr(fn)), exc
)
return
continue
else:
return
app.notify(
"Copy failed — use Shift+mouse-select for native terminal copy",
"Failed to copy — no clipboard method available",
severity="warning",
timeout=3,
)
+304 -1225
View File
File diff suppressed because it is too large Load Diff
+42 -168
View File
@@ -19,35 +19,16 @@ from pathlib import Path
_PATH_CHARS = r"A-Za-z0-9._~/\\:-"
FILE_MENTION_PATTERN = re.compile(
r"@(?:"
r'"(?P<dquoted>[^"\n]+)"'
r"|'(?P<squoted>[^'\n]+)'"
r"|(?P<bare>(?:\\.|[" + _PATH_CHARS + r"])+)"
r")"
)
FILE_MENTION_PATTERN = re.compile(r"@(?P<path>(?:\\.|[" + _PATH_CHARS + r"])+)")
"""Matches ``@path/to/file`` in user input.
Three forms are supported, in priority order:
1. ``@"path with spaces.pdf"`` — explicit double-quoted path
2. ``@'path with spaces.pdf'`` — explicit single-quoted path
3. ``@bare/path`` — backslash-escaped spaces (``@my\\\\ folder/file``) work;
raw unescaped spaces are handled via greedy expansion in
:func:`parse_file_mentions`.
Bare ``@`` with no path characters is not matched (uses ``+`` not ``*``).
Escaped spaces (``@my\\\\ folder/file``) are supported. Bare ``@`` with no
path characters is not matched (uses ``+`` not ``*``).
"""
_EMAIL_PREFIX = re.compile(r"[a-zA-Z0-9._%+-]$")
"""If the character immediately before ``@`` matches this, it's an email address."""
# Hard cap on tokens consumed during greedy expansion across whitespace.
_GREEDY_MAX_TOKENS = 20
# Trailing punctuation stripped before checking if a greedy candidate exists.
_GREEDY_TRAIL_PUNCT = ",;:!?)]}>"
# Files larger than this are referenced by path only (not embedded inline).
_MAX_EMBED_BYTES = 256 * 1024 # 256 KB
@@ -212,75 +193,6 @@ def _read_file(path: Path) -> str:
return f"\n### {path.name}\nPath: `{path}`\n```\n{content}\n```"
def _resolve_path(raw: str, cwd: Path) -> Path | None:
"""Resolve *raw* to an existing file path, or ``None``.
Honors backslash-escaped spaces and ``~`` expansion. Returns ``None``
when the path does not exist, is not a regular file, or raises
``OSError``/``RuntimeError`` during resolution.
"""
clean = raw.replace("\\ ", " ")
try:
p = Path(clean).expanduser()
if not p.is_absolute():
p = cwd / p
resolved = p.resolve()
except (OSError, RuntimeError):
return None
if resolved.is_file():
return resolved
return None
def _greedy_extend(
text: str,
raw: str,
match_end: int,
cwd: Path,
) -> tuple[str, Path, int] | None:
"""Try to extend *raw* across whitespace until the path resolves.
Walks the text after *match_end*, capped at the next newline or the
start of another ``@`` mention. Tries the longest plausible suffix
first and shrinks one token at a time, stripping trailing punctuation
that is unlikely to be part of a filename.
Returns ``(extended_raw, resolved_file, new_end_pos)`` on success,
else ``None``.
"""
rest = text[match_end:]
# Hard boundaries that should never be crossed.
boundary = len(rest)
nl = rest.find("\n")
if nl >= 0:
boundary = nl
next_at = re.search(r"\s@", rest[:boundary])
if next_at:
boundary = next_at.start()
region = rest[:boundary]
if not region or not region[0].isspace():
return None
tokens = list(re.finditer(r"\S+", region))
if not tokens:
return None
# Try longest-first so we prefer the most specific match.
for i in range(min(len(tokens), _GREEDY_MAX_TOKENS), 0, -1):
end = tokens[i - 1].end()
suffix = region[:end].rstrip(_GREEDY_TRAIL_PUNCT)
if not suffix:
continue
candidate = raw + suffix
resolved = _resolve_path(candidate, cwd)
if resolved is not None:
return candidate, resolved, match_end + len(suffix)
return None
def parse_file_mentions(
text: str,
cwd: Path | None = None,
@@ -306,54 +218,40 @@ def parse_file_mentions(
files: list[Path] = []
warnings: list[str] = []
seen: set[Path] = set()
# finditer would normally re-scan from each match's end, but greedy
# expansion can consume bytes past that point. Track a manual cursor
# and skip matches that start before it.
cursor = 0
for match in FILE_MENTION_PATTERN.finditer(text):
if match.start() < cursor:
continue
# Skip email addresses — character immediately before @ is alphanumeric
before = text[: match.start()]
if before and _EMAIL_PREFIX.search(before):
continue
dquoted = match.group("dquoted")
squoted = match.group("squoted")
bare = match.group("bare")
quoted_raw = dquoted if dquoted is not None else squoted
raw = quoted_raw if quoted_raw is not None else bare
is_quoted = quoted_raw is not None
raw = match.group("path")
clean = raw.replace("\\ ", " ")
resolved = _resolve_path(raw, cwd)
end_pos = match.end()
if resolved is None and not is_quoted:
extended = _greedy_extend(text, raw, match.end(), cwd)
if extended is not None:
raw, resolved, end_pos = extended
cursor = end_pos
if resolved is None:
warnings.append(f"@file not found: {raw}")
continue
# Deduplicate: skip paths already seen in this message.
if resolved in seen:
continue
seen.add(resolved)
files.append(resolved)
# Warn when the file lives outside the workspace root — it may
# contain sensitive content (e.g. @~/.ssh/id_rsa).
# Checked after dedup so a repeated mention only warns once.
try:
resolved.relative_to(workspace_root)
except ValueError:
warnings.append(
f"@{raw} is outside the workspace "
f"({workspace_root}) — embedding may expose sensitive files"
)
p = Path(clean).expanduser()
if not p.is_absolute():
p = cwd / p
resolved = p.resolve()
if not resolved.exists() or not resolved.is_file():
warnings.append(f"@file not found: {raw}")
continue
# Deduplicate: skip paths already seen in this message.
if resolved in seen:
continue
seen.add(resolved)
files.append(resolved)
# Warn when the file lives outside the workspace root — it may
# contain sensitive content (e.g. @~/.ssh/id_rsa).
# Checked after dedup so a repeated mention only warns once.
try:
resolved.relative_to(workspace_root)
except ValueError:
warnings.append(
f"@{raw} is outside the workspace "
f"({workspace_root}) — embedding may expose sensitive files"
)
except (OSError, RuntimeError) as exc:
warnings.append(f"invalid @file path {raw!r}: {exc}")
return files, warnings
@@ -404,13 +302,6 @@ def _type_hint(rel_path: str) -> str:
return suffix or "file"
def _format_mention(rel_path: str) -> str:
"""Render *rel_path* as an ``@`` mention, quoting if it contains spaces."""
if " " in rel_path:
return f'@"{rel_path}"'
return f"@{rel_path}"
def complete_file_mention(
text: str,
workspace_dir: str | None = None,
@@ -429,22 +320,13 @@ def complete_file_mention(
List of ``(completion_string, type_hint)`` tuples, e.g.
``[("@results/v2.json", "json"), ("@README.md", "md")]``.
Directories have a trailing ``/`` and type hint ``"dir"``.
Paths containing spaces are returned in double-quoted form,
e.g. ``@"my docs/file.pdf"``.
"""
# Find the last @token. Allow whitespace inside a quoted partial so
# completion keeps working as the user types ``@"PRE`` → ``@"PREPING_ B``.
quoted_match = re.search(r'@"([^"\n]*)$', text)
if quoted_match:
partial = quoted_match.group(1)
quoted = True
else:
match = re.search(r"@([^\s\"']*)$", text)
if not match:
return []
partial = match.group(1).replace("\\ ", " ")
quoted = False
# Find the last @token
match = re.search(r"@([^\s]*)$", text)
if not match:
return []
partial = match.group(1).replace("\\ ", " ")
base_str = workspace_dir or str(Path.cwd())
base = Path(base_str)
@@ -462,10 +344,10 @@ def complete_file_mention(
rel = entry.relative_to(base)
suffix = "/" if entry.is_dir() else ""
candidates_raw.append(rel.as_posix() + suffix)
except (OSError, ValueError):
except OSError:
return []
return [
(_format_mention(r), "dir" if r.endswith("/") else _type_hint(r))
(f"@{r}", "dir" if r.endswith("/") else _type_hint(r))
for r in candidates_raw[:10]
]
@@ -481,24 +363,16 @@ def complete_file_mention(
except OSError:
pass
combined = all_files + dir_candidates
# Determine query: if partial has a slash, search within that subtree
if "/" in partial:
# Search within the subtree for the given directory prefix
combined = all_files + dir_candidates
# Filter candidates to those starting with the directory prefix
dir_prefix = partial.rsplit("/", 1)[0] + "/"
file_query = partial.rsplit("/", 1)[1]
subtree = [c for c in combined if c.startswith(dir_prefix)]
results = _fuzzy_search(file_query, subtree)
else:
# Depth 1 only: top-level files and directories
top_files = [f for f in all_files if "/" not in f]
results = _fuzzy_search(partial, top_files + dir_candidates)
results = _fuzzy_search(partial, combined)
if quoted:
# User opened a quoted mention — close it for them.
return [
(f'@"{r}"', "dir" if r.endswith("/") else _type_hint(r)) for r in results
]
return [
(_format_mention(r), "dir" if r.endswith("/") else _type_hint(r))
for r in results
]
return [(f"@{r}", "dir" if r.endswith("/") else _type_hint(r)) for r in results]
+1 -1
View File
@@ -1,7 +1,7 @@
"""History-based auto-suggest for Textual TUI Input widget.
Reads prompt_toolkit FileHistory format so Rich CLI and TUI share the same
history file at ~/.evoscientist/history.
history file at ~/.config/ai4scientist/history.
"""
from __future__ import annotations
File diff suppressed because it is too large Load Diff
+15 -4
View File
@@ -9,6 +9,7 @@ from __future__ import annotations
from collections import Counter
import questionary
from prompt_toolkit.styles import Style as PtStyle
from questionary import Choice
from ..mcp.registry import (
@@ -20,8 +21,18 @@ from ..mcp.registry import (
install_mcp_server,
install_mcp_servers,
)
from ..stream.console import console
from .widgets.thread_selector import PICKER_STYLE
from ..stream.display import console
_PICKER_STYLE = PtStyle.from_dict(
{
"questionmark": "#888888",
"question": "",
"pointer": "bold",
"highlighted": "bold",
"text": "#888888",
"answer": "bold",
}
)
_INSTALLED_INDICATOR = ("fg:#4caf50", "\u2713 ")
@@ -46,7 +57,7 @@ def _checkbox_ask(choices, message: str, **kwargs):
return questionary.checkbox(
message,
choices=choices,
style=PICKER_STYLE,
style=_PICKER_STYLE,
qmark="\u276f",
**kwargs,
).ask()
@@ -89,7 +100,7 @@ def _browse_and_select(
selected_tag = questionary.select(
"Filter by tag:",
choices=tag_choices,
style=PICKER_STYLE,
style=_PICKER_STYLE,
qmark="\u276f",
).ask()
+128 -2
View File
@@ -1,10 +1,10 @@
"""Shared UI helpers for MCP server display and operations (used by the Typer `mcp` commands)."""
"""MCP server display, operations, and /mcp slash-command dispatcher."""
from typing import Any
from rich.table import Table
from ..stream.console import console
from ..stream.display import console
def _mcp_list_servers() -> None:
@@ -181,3 +181,129 @@ def _show_mcp_config(name: str = "", *, show_blank_line: bool = True) -> str:
if show_blank_line:
console.print()
return "ok"
def _cmd_mcp_add(args_str: str) -> None:
"""Handle ``/mcp add ...``."""
import shlex
from ..mcp import parse_mcp_add_args
if not args_str.strip():
console.print("[bold]Usage:[/bold] /mcp add <name> <command-or-url> [args...]")
console.print()
console.print(
"[dim]Transport is auto-detected: URLs \u2192 http, commands \u2192 stdio[/dim]"
)
console.print()
console.print("[bold]Examples:[/bold]")
console.print(
" /mcp add sequential-thinking npx -y @modelcontextprotocol/server-sequential-thinking"
)
console.print(" /mcp add docs-langchain https://docs.langchain.com/mcp")
console.print(
" /mcp add my-sse http://localhost:9090/sse --transport sse --expose-to research-agent"
)
console.print()
console.print("[dim]Options:[/dim]")
console.print(" --transport T Transport type (default: auto-detect)")
console.print(
" --tools t1,t2 Tool allowlist (supports wildcards: *_exa, read_*)"
)
console.print(" --expose-to a1,a2 Target agents (default: main)")
console.print(" --header Key:Value HTTP header (repeatable)")
console.print(" --env KEY=VALUE Env var for stdio (repeatable)")
console.print(
" --env-ref KEY Env var as runtime ${KEY} reference (repeatable)"
)
console.print()
return
try:
tokens = shlex.split(args_str)
kwargs = parse_mcp_add_args(tokens)
_mcp_add_server_from_kwargs(kwargs, show_reload_hint=True)
except ValueError as exc:
console.print(f"[red]{exc}[/red]")
console.print()
def _cmd_mcp_edit(args_str: str) -> None:
"""Handle ``/mcp edit <name> --field value ...``."""
import shlex
from ..mcp import parse_mcp_edit_args
if not args_str.strip():
console.print("[bold]Usage:[/bold] /mcp edit <name> --<field> <value> ...")
console.print()
console.print(
"[dim]Fields:[/dim] --transport, --command, --url, --args, --tools, --expose-to, --header, --env"
)
console.print(
"[dim]Use[/dim] --tools none [dim]or[/dim] --expose-to none [dim]to clear a field.[/dim]"
)
console.print()
console.print("[bold]Examples:[/bold]")
console.print(" /mcp edit filesystem --expose-to main,code-agent")
console.print(" /mcp edit filesystem --tools read_file,write_file")
console.print(" /mcp edit my-api --url http://new-host:8080/mcp")
console.print(" /mcp edit my-api --tools none")
console.print()
return
try:
tokens = shlex.split(args_str)
name, fields = parse_mcp_edit_args(tokens)
_mcp_edit_server_fields(name, fields, show_reload_hint=True)
except ValueError as exc:
console.print(f"[red]{exc}[/red]")
console.print()
def _cmd_mcp_remove(name: str) -> None:
"""Handle ``/mcp remove <name>``."""
_mcp_remove_server(name, show_reload_hint=True)
console.print()
def _cmd_mcp_config(name: str) -> None:
"""Handle ``/mcp config [name]``."""
_show_mcp_config(name, show_blank_line=True)
def _cmd_mcp(args: str) -> None:
"""Dispatch ``/mcp`` subcommands."""
args = args.strip()
if not args:
_mcp_list_servers()
return
parts = args.split(maxsplit=1)
subcmd = parts[0].lower()
subargs = parts[1] if len(parts) > 1 else ""
if subcmd == "list":
_mcp_list_servers()
elif subcmd == "add":
_cmd_mcp_add(subargs)
elif subcmd == "edit":
_cmd_mcp_edit(subargs)
elif subcmd == "remove":
_cmd_mcp_remove(subargs)
elif subcmd == "config":
_cmd_mcp_config(subargs)
elif subcmd == "install":
from .mcp_install_cmd import _cmd_install_mcp
_cmd_install_mcp(subargs)
else:
console.print("[bold]MCP commands:[/bold]")
console.print(" /mcp List configured servers")
console.print(" /mcp list List configured servers")
console.print(" /mcp config Show detailed server config")
console.print(" /mcp add ... Add a server")
console.print(" /mcp edit ... Edit an existing server")
console.print(" /mcp remove ... Remove a server")
console.print(" /mcp install ... Browse and install servers")
console.print()
-21
View File
@@ -1,21 +0,0 @@
"""Helper for printing the session-exit Goodbye message and resume hint."""
from __future__ import annotations
from rich.console import Console
from rich.markup import escape
def print_resume_hint(
thread_id: str | None,
console: Console | None = None,
) -> None:
"""Print ``Goodbye!`` and, when available, a resume hint for *thread_id*."""
out = console or Console()
out.print("[dim]Goodbye![/dim]")
if thread_id:
from ..sessions import short_thread_id
out.print()
out.print("[dim]Resume this session with:[/dim]")
out.print(f"[cyan]EvoSci --resume {escape(short_thread_id(thread_id))}[/cyan]")
-256
View File
@@ -1,256 +0,0 @@
"""CommandUI Protocol adapter for the Rich CLI surface.
Lifecycle methods (``request_quit``, ``force_quit``, ``clear_chat``,
``start_new_session``, ``handle_session_resume``, ``update_status_after_compact``)
are callback-driven: when their corresponding ``on_*`` constructor kwarg
is ``None``, the method is a silent no-op, mirroring
``ChannelCommandUI``'s fallback pattern. Callers that need a specific
side-effect (REPL quit flag flip, status-bar refresh, …) wire the
callback at construction time; non-interactive surfaces (tests,
alternate REPLs) can leave callbacks unset without crashing.
``wait_for_*`` methods return ``None`` on cancel / fallback and are
always safe to ``await``.
"""
from __future__ import annotations
from collections.abc import Awaitable, Callable
from typing import Any
from rich.console import Console
from rich.table import Table
from ..commands.base import CommandUI
class RichCLICommandUI(CommandUI):
"""CommandUI implementation that prints to a Rich ``Console``.
Commands that affect CLI-closure state (session lifecycle, exit flag,
status-bar snapshot) go through optional callbacks wired by the REPL.
This mirrors ``ChannelCommandUI``'s injection pattern and keeps
``interactive.py``'s ``state`` dict as the single source of truth.
"""
def __init__(
self,
console: Console,
*,
on_request_quit: Callable[[], None] | None = None,
on_force_quit: Callable[[], None] | None = None,
on_clear_chat: Callable[[], None] | None = None,
on_status_after_compact: Callable[[int], None] | None = None,
on_start_new_session: Callable[[], Awaitable[None]] | None = None,
on_handle_session_resume: (
Callable[[str, str | None], Awaitable[None]] | None
) = None,
) -> None:
self.console = console
self._on_request_quit = on_request_quit
self._on_force_quit = on_force_quit
self._on_clear_chat = on_clear_chat
self._on_status_after_compact = on_status_after_compact
self._on_start_new_session = on_start_new_session
self._on_handle_session_resume = on_handle_session_resume
# Bound ``console.status(...)`` context manager used by
# /compact's start/stop indicator pair.
self._compact_status_ctx: Any = None
# ── Core I/O ─────────────────────────────────────────────
@property
def supports_interactive(self) -> bool:
return True
def append_system(self, text: str, style: str = "dim") -> None:
self.console.print(text, style=style)
def mount_renderable(self, renderable: Any) -> None:
self.console.print(renderable)
async def flush(self) -> None:
# Rich console flushes synchronously; nothing to await.
return
# ── /model interactive picker fallback ──────────────────
async def wait_for_model_pick(
self,
entries: list[tuple[str, str, str]],
current_model: str | None,
current_provider: str | None,
) -> tuple[str, str] | None:
"""Print the model table and return ``None``; user re-runs with
``/model <name>`` since the CLI has no interactive picker."""
table = Table(
title="Available Models",
show_header=True,
header_style="bold cyan",
)
table.add_column("Name", style="bold")
table.add_column("Provider", style="dim")
for name, _mid, prov in entries:
marker = " *" if name == current_model and prov == current_provider else ""
table.add_row(f"{name}{marker}", prov)
self.console.print(table)
self.console.print(
"[dim]Usage: /model <name> [provider] [--save] — "
"provider is optional, auto-detected from model name[/dim]"
)
return None
def update_status_after_model_change(
self, new_model: str, new_provider: str | None = None
) -> None:
"""No-op; the CLI REPL refreshes status itself after detecting an
``ctx.agent`` change post-``cmd_manager.execute``."""
return
# ── Interactive pickers ────────────────────────────────
async def wait_for_thread_pick(
self, threads: list[dict], current_thread: str, title: str
) -> str | None:
"""Interactive workspace-grouped thread picker using ``questionary``.
Ported from the pre-migration ``_cmd_resume`` implementation.
Returns the selected ``thread_id`` string, or ``None`` on cancel.
Callers (``ResumeCommand``/``DeleteCommand``) pre-check for
empty thread lists before invoking this method.
"""
import questionary # type: ignore[import-untyped]
from prompt_toolkit.layout.dimension import ( # type: ignore[import-untyped]
Dimension,
)
from questionary.prompts.common import ( # type: ignore[import-untyped]
InquirerControl,
)
from ..sessions import _format_relative_time
from .widgets.thread_selector import PICKER_STYLE, _build_items
choices: list[Any] = []
for item in _build_items(threads):
if item["type"] == "header":
choices.append(questionary.Separator(f"── \U0001f4c2 {item['label']}"))
elif item["type"] == "subheader":
choices.append(questionary.Separator(f" {item['label']}"))
else:
t = item["thread"]
tid = t["thread_id"]
preview = t.get("preview", "") or ""
msgs = t.get("message_count", 0)
model = t.get("model", "") or ""
when = _format_relative_time(t.get("updated_at"))
indent = " " if item.get("indented") else " "
marker = " *" if tid == current_thread else ""
parts = [f"{indent}{tid}{marker}"]
if preview:
parts.append(preview[:40] + "…" if len(preview) > 40 else preview)
parts.append(f"({msgs} msgs)")
if model:
parts.append(model)
if when:
parts.append(when)
label = " ".join(parts)
choices.append(questionary.Choice(title=label, value=tid))
prompt = questionary.select(title, choices=choices, style=PICKER_STYLE)
# Limit visible list to 10 rows with scrolling. Touches
# questionary/prompt-toolkit private internals so guard against
# library-shape changes — picker stays functional at default
# height even if the cap fails.
try:
for window in prompt.application.layout.find_all_windows():
if isinstance(window.content, InquirerControl):
window.height = Dimension(max=10)
break
except Exception:
pass
# ``ask_async`` (questionary >= 2.0.1) avoids blocking the
# asyncio event loop while the user interacts with the picker.
return await prompt.ask_async()
# ── Lifecycle callbacks ───────────────────────────────
def clear_chat(self) -> None:
if self._on_clear_chat is not None:
self._on_clear_chat()
else:
self.console.clear()
def request_quit(self) -> None:
if self._on_request_quit is not None:
self._on_request_quit()
def force_quit(self) -> None:
if self._on_force_quit is not None:
self._on_force_quit()
async def start_new_session(self) -> None:
if self._on_start_new_session is not None:
await self._on_start_new_session()
async def handle_session_resume(
self, thread_id: str, workspace_dir: str | None = None
) -> None:
if self._on_handle_session_resume is not None:
await self._on_handle_session_resume(thread_id, workspace_dir)
# /compact indicator pair — duck-typed by ``CompactCommand`` via
# ``getattr``, not declared on the ``CommandUI`` Protocol.
def start_compacting_indicator(self) -> None:
# Idempotent: close any lingering context before starting a new
# one so a double-call (e.g. two overlapping /compact attempts
# via the message queue) can't leak a Rich Live handle.
if self._compact_status_ctx is not None:
try:
self._compact_status_ctx.__exit__(None, None, None)
except Exception:
pass
self._compact_status_ctx = None
status = self.console.status("[cyan]Compacting conversation...[/cyan]")
status.__enter__()
self._compact_status_ctx = status
def stop_compacting_indicator(self) -> None:
ctx = self._compact_status_ctx
self._compact_status_ctx = None
if ctx is not None:
try:
ctx.__exit__(None, None, None)
except Exception:
pass
def update_status_after_compact(self, input_tokens: int) -> None:
if self._on_status_after_compact is not None:
self._on_status_after_compact(input_tokens)
# ── Skill / MCP browse (delegated to worker threads) ──
async def wait_for_skill_browse(
self, index: list[dict], installed_names: set[str], pre_filter_tag: str
) -> list[str] | None:
"""Delegate to the extracted questionary picker on a worker
thread — questionary blocks the event loop so the call must
not happen on the main asyncio thread."""
import asyncio
from .skills_cmd import _pick_skills_interactive
return await asyncio.to_thread(
_pick_skills_interactive, index, installed_names, pre_filter_tag
)
async def wait_for_mcp_browse(
self, servers: list, installed_names: set[str], pre_filter_tag: str
) -> list | None:
"""Delegate to the MCP browse picker on a worker thread."""
import asyncio
from .mcp_install_cmd import _browse_and_select
return await asyncio.to_thread(
_browse_and_select, servers, installed_names, pre_filter_tag
)
+224 -30
View File
@@ -1,30 +1,171 @@
"""Shared UI helpers for skill-management commands (picker used by /evoskills)."""
"""Slash commands for skill management: /skills, /install-skill, /uninstall-skill, /evoskills."""
from ..stream.console import console
from pathlib import Path
from ..stream.display import console
from .agent import _shorten_path
def _pick_skills_interactive(
index: list[dict],
installed_names: set[str],
pre_filter_tag: str,
) -> list[str] | None:
"""Interactive questionary picker for EvoSkills browse.
def _cmd_list_skills() -> None:
"""List all available skills (workspace, global, and built-in)."""
from ..paths import GLOBAL_SKILLS_DIR, USER_SKILLS_DIR
from ..tools.skills_manager import list_skills
Two-phase picker:
1. tag filter — ``questionary.select`` (skipped if ``pre_filter_tag``)
2. multi-select — ``questionary.checkbox`` with installed items disabled
skills = list_skills(include_system=True)
Returns:
list of ``install_source`` strings selected by the user,
``None`` if the user cancelled at either phase, or
``[]`` if nothing was selectable / all-installed in the filter.
if not skills:
console.print("[dim]No skills available.[/dim]")
console.print("[dim]Install with:[/dim] /install-skill <path-or-url>")
console.print(
f"[dim]Global skills:[/dim] [cyan]{_shorten_path(str(GLOBAL_SKILLS_DIR))}[/cyan]"
)
console.print()
return
workspace_skills = [s for s in skills if s.source == "workspace"]
global_skills = [s for s in skills if s.source == "global"]
builtin_skills = [s for s in skills if s.source == "builtin"]
sections = [
("Workspace Skills", workspace_skills, "green"),
("Global Skills", global_skills, "cyan"),
("Built-in Skills", builtin_skills, "blue"),
]
printed = False
for title, group, color in sections:
if not group:
continue
if printed:
console.print()
console.print(f"[bold]{title}[/bold] ({len(group)}):")
for skill in group:
tags_str = f" [dim]({', '.join(skill.tags)})[/dim]" if skill.tags else ""
console.print(
f" [{color}]{skill.name}[/{color}] - {skill.description}{tags_str}"
)
printed = True
console.print(
f"\n[dim]Global skills:[/dim] [cyan]{_shorten_path(str(GLOBAL_SKILLS_DIR))}[/cyan]"
)
console.print(
f"[dim]Workspace skills:[/dim] [green]{_shorten_path(str(USER_SKILLS_DIR))}[/green]"
)
console.print()
def _cmd_install_skill(args: str) -> None:
"""Install a skill from local path or GitHub URL.
By default, installs to the global skills directory (~/.config/ai4scientist/skills/).
Append --local to install to the current workspace instead.
Usage: /install-skill <path-or-url> [--local]
"""
from ..paths import GLOBAL_SKILLS_DIR, USER_SKILLS_DIR
from ..tools.skills_manager import install_skill
# Parse --local flag out of the args string
local = "--local" in args.split()
source = args.replace("--local", "").strip()
if not source:
console.print("[red]Usage:[/red] /install-skill <path-or-url> [--local]")
console.print("[dim]Examples:[/dim]")
console.print(" /install-skill ./my-skill")
console.print(
" /install-skill https://github.com/user/repo/tree/main/skill-name"
)
console.print(" /install-skill user/repo@skill-name")
console.print(
" /install-skill ./my-skill --local [dim](workspace only)[/dim]"
)
console.print()
return
dest_label = (
f"[cyan]{_shorten_path(str(USER_SKILLS_DIR))}[/cyan] [dim](workspace)[/dim]"
if local
else f"[cyan]{_shorten_path(str(GLOBAL_SKILLS_DIR))}[/cyan] [dim](global)[/dim]"
)
console.print(f"[dim]Installing skill from:[/dim] {source}")
console.print(f"[dim]Destination:[/dim] {dest_label}")
result = install_skill(source, global_install=not local)
if result.get("batch"):
# Batch install — multiple skills
for item in result.get("installed", []):
console.print(f"[green]Installed:[/green] {item['name']}")
console.print(
f" [dim]Description:[/dim] {item.get('description', '(none)')}"
)
console.print(
f" [dim]Path:[/dim] [cyan]{_shorten_path(item['path'])}[/cyan]"
)
for item in result.get("failed", []):
console.print(f"[red]Failed:[/red] {item['name']} — {item['error']}")
installed_count = len(result.get("installed", []))
if installed_count:
console.print(f"\n[green]{installed_count} skill(s) installed.[/green]")
console.print("[dim]Reload with /new to apply.[/dim]")
elif result["success"]:
console.print(f"[green]Installed:[/green] {result['name']}")
console.print(f"[dim]Description:[/dim] {result.get('description', '(none)')}")
console.print(f"[dim]Path:[/dim] [cyan]{_shorten_path(result['path'])}[/cyan]")
console.print()
console.print("[dim]Reload with /new to apply.[/dim]")
else:
console.print(f"[red]Failed:[/red] {result['error']}")
console.print()
def _cmd_uninstall_skill(name: str) -> None:
"""Uninstall a user-installed skill."""
from ..tools.skills_manager import uninstall_skill
if not name:
console.print("[red]Usage:[/red] /uninstall-skill <skill-name>")
console.print("[dim]Use /skills to see installed skills.[/dim]")
console.print()
return
result = uninstall_skill(name)
if result["success"]:
console.print(f"[green]Uninstalled:[/green] {name}")
console.print("[dim]Reload with /new to apply.[/dim]")
else:
console.print(f"[red]Failed:[/red] {result['error']}")
console.print()
def _cmd_install_skills(args: str = "") -> None:
"""Browse and install skills from the EvoSkills repository.
Args:
args: Optional tag name to pre-filter (e.g. "core").
"""
from collections import Counter
import questionary
from prompt_toolkit.styles import Style as PtStyle
from questionary import Choice
from .widgets.thread_selector import PICKER_STYLE
from ..paths import GLOBAL_SKILLS_DIR, USER_SKILLS_DIR
from ..tools.skills_manager import fetch_remote_skill_index, install_skill
_PICKER_STYLE = PtStyle.from_dict(
{
"questionmark": "#888888",
"question": "",
"pointer": "bold",
"highlighted": "bold",
"text": "#888888",
"answer": "bold",
}
)
# Installed-item indicator style for disabled checkbox choices.
_INSTALLED_INDICATOR = ("fg:#4caf50", "✓ ")
@@ -49,24 +190,46 @@ def _pick_skills_interactive(
return questionary.checkbox(
message,
choices=choices,
style=PICKER_STYLE,
style=_PICKER_STYLE,
qmark="❯",
**kwargs,
).ask()
finally:
InquirerControl._get_choice_tokens = original
pre_filter_tag = (pre_filter_tag or "").strip().lower()
# Step 1: Fetch remote index
console.print("[dim]Fetching skill index...[/dim]")
try:
index = fetch_remote_skill_index()
except Exception as e:
console.print(f"[red]Failed to fetch skill index: {e}[/red]")
console.print(
"[dim]Try installing directly: /install-skill EvoScientist/EvoSkills@skills[/dim]"
)
console.print()
return
# Phase 1: tag filter (skip if pre-filtered via args)
if not index:
console.print("[yellow]No skills found in the repository.[/yellow]")
console.print()
return
# Detect already-installed skills (both global and workspace tiers)
installed_names: set[str] = set()
for skills_dir in (Path(GLOBAL_SKILLS_DIR), Path(USER_SKILLS_DIR)):
if skills_dir.exists():
installed_names.update(e.name for e in skills_dir.iterdir() if e.is_dir())
pre_filter_tag = args.strip().lower() if args else ""
# Step 2: Tag filter (skip if pre-filtered via args)
if pre_filter_tag:
filtered = [
s for s in index if pre_filter_tag in [t.lower() for t in s.get("tags", [])]
]
if not filtered:
console.print(
f"[yellow]No skills found with tag: {pre_filter_tag}[/yellow]"
)
console.print(f"[yellow]No skills found with tag: {args.strip()}[/yellow]")
# Show available tags
tag_counter: Counter[str] = Counter()
for s in index:
for t in s.get("tags", []):
@@ -75,8 +238,10 @@ def _pick_skills_interactive(
sorted_tags = sorted(tag_counter.items(), key=lambda x: (-x[1], x[0]))
tags_str = ", ".join(f"{tag} ({count})" for tag, count in sorted_tags)
console.print(f"[dim]Available tags: {tags_str}[/dim]")
return []
console.print()
return
else:
# Build tag choices for interactive picker
tag_counter = Counter()
for s in index:
for t in s.get("tags", []):
@@ -90,12 +255,13 @@ def _pick_skills_interactive(
selected_tag = questionary.select(
"Filter by tag:",
choices=tag_choices,
style=PICKER_STYLE,
style=_PICKER_STYLE,
qmark="❯",
).ask()
if selected_tag is None:
return None
console.print()
return
if selected_tag == "__all__":
filtered = index
@@ -106,12 +272,14 @@ def _pick_skills_interactive(
if selected_tag in [t.lower() for t in s.get("tags", [])]
]
# Phase 2: skill selection checkbox
if all(s["name"] in installed_names for s in filtered):
# Step 3: Skill selection checkbox
all_installed = all(s["name"] in installed_names for s in filtered)
if all_installed:
console.print(
"[green]All skills in this category are already installed.[/green]"
)
return []
console.print()
return
choices = []
for s in filtered:
@@ -137,5 +305,31 @@ def _pick_skills_interactive(
selected = _checkbox_ask(choices, "Select skills to install:")
if selected is None:
return None
return list(selected)
console.print()
return
if not selected:
console.print("[dim]No skills selected.[/dim]")
console.print()
return
# Step 4: Install selected skills (default: global)
installed_count = 0
for source in selected:
result = install_skill(source, global_install=True)
if result.get("batch"):
for item in result.get("installed", []):
console.print(f"[green]Installed:[/green] {item['name']}")
installed_count += 1
for item in result.get("failed", []):
console.print(f"[red]Failed:[/red] {item['name']} — {item['error']}")
elif result.get("success"):
console.print(f"[green]Installed:[/green] {result['name']}")
installed_count += 1
else:
console.print(f"[red]Failed:[/red] {result.get('error', 'unknown')}")
if installed_count:
console.print(f"\n[green]{installed_count} skill(s) installed.[/green]")
console.print("[dim]Reload with /new to apply.[/dim]")
console.print()
+13 -110
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
from dataclasses import dataclass, replace
from datetime import datetime
from typing import TYPE_CHECKING, Any
from typing import Any
from langchain_core.messages import AIMessage, HumanMessage
from langchain_core.messages.utils import count_tokens_approximately
@@ -13,15 +13,7 @@ from ..llm.context_window import (
DEFAULT_CONTEXT_WINDOW_FALLBACK,
resolve_context_window,
)
from ..memory.worker_activity import (
MemoryWorkerStatusSnapshot,
ObservationLinkerStatusSnapshot,
memory_worker_status,
observation_linker_status,
)
if TYPE_CHECKING:
from ..gateway import GraphGateway
from ..sessions import get_thread_messages
_FALLBACK_CONTEXT_WINDOW = DEFAULT_CONTEXT_WINDOW_FALLBACK
STATUS_BAR_BG = "#171a20"
@@ -34,11 +26,6 @@ STATUS_BAD = "#d08c61"
STATUS_CRITICAL = "#d86f6f"
STATUS_HINT_IDLE = "#8b9bb0"
STATUS_HINT_BUSY = "#f0c36a"
STATUS_HINT_WRITING = "#7eb8e0"
# Braille spinner frames used by the CLI bottom toolbar and TUI status bar
# to animate the "Loading MCP tools" indicator.
SPINNER_FRAMES = "\u280b\u2819\u2839\u2838\u283c\u2834\u2826\u2827\u2807\u280f"
@dataclass(slots=True)
@@ -99,15 +86,15 @@ def format_token_count_compact(value: int) -> str:
"""Format large token counts into a compact human-readable form."""
abs_value = abs(int(value))
if abs_value >= 1_000_000:
num = float(value) / 1_000_000
num = value / 1_000_000
suffix = "M"
elif abs_value >= 1_000:
num = float(value) / 1_000
num = value / 1_000
suffix = "K"
else:
return str(value)
if num == int(num):
if num.is_integer():
return f"{int(num)}{suffix}"
return f"{num:.1f}{suffix}"
@@ -168,10 +155,15 @@ def trim_status_text(text: str, max_width: int) -> str:
if max_width <= ellipsis_width:
return ellipsis[:max_width]
try:
from prompt_toolkit.utils import get_cwidth
except Exception:
get_cwidth = None
out: list[str] = []
width = 0
for ch in text:
ch_width = _display_width(ch)
ch_width = get_cwidth(ch) if get_cwidth else len(ch)
if width + ch_width + ellipsis_width > max_width:
break
out.append(ch)
@@ -179,94 +171,13 @@ def trim_status_text(text: str, max_width: int) -> str:
return "".join(out).rstrip() + ellipsis
def get_memory_worker_status() -> MemoryWorkerStatusSnapshot | None:
"""Read completed EvoMemory save counts without making rendering fail."""
try:
return memory_worker_status()
except Exception:
return None
def get_observation_linker_status() -> ObservationLinkerStatusSnapshot | None:
"""Read active observation-linker status without making rendering fail."""
try:
return observation_linker_status()
except Exception:
return None
def _plural(count: int, singular: str, plural: str | None = None) -> str:
word = singular if count == 1 else (plural or f"{singular}s")
return f"{count} {word}"
def _memory_activity_label(
*,
worker_status: MemoryWorkerStatusSnapshot | None,
linker_status: ObservationLinkerStatusSnapshot | None,
) -> str:
parts: list[str] = []
if worker_status is not None and worker_status.is_running:
parts.append("🧠")
if linker_status is not None and linker_status.is_running:
parts.append("🔗")
saved: list[str] = []
if worker_status is not None:
if worker_status.profile_updates:
saved.append(_plural(worker_status.profile_updates, "profile edit"))
if worker_status.observations_recorded:
saved.append(_plural(worker_status.observations_recorded, "observation"))
if saved:
parts.append(f"Saved {', '.join(saved)}")
if linker_status is not None and linker_status.relations_linked:
parts.append(
f"Created {_plural(linker_status.relations_linked, 'memory link')}"
)
return " ".join(parts)
def _append_memory_indicator(
frags: list[tuple[str, str]],
*,
worker_status: MemoryWorkerStatusSnapshot | None,
linker_status: ObservationLinkerStatusSnapshot | None,
width: int,
) -> None:
if worker_status is None and linker_status is None:
return
label = _memory_activity_label(
worker_status=worker_status,
linker_status=linker_status,
)
if not label:
return
tail: list[tuple[str, str]] = []
if frags and frags[-1] == ("class:status-bar", " "):
tail.append(frags.pop())
separator = " │ " if width >= 76 else " · "
frags.extend(
[
("class:status-bar-dim", separator),
("class:status-bar-warn", label),
]
)
frags.extend(tail)
def build_status_fragments(
snapshot: SessionStatusSnapshot,
started_at: datetime,
width: int,
) -> list[tuple[str, str]]:
"""Build prompt_toolkit formatted-text fragments for the status bar."""
now = datetime.now()
duration_label = format_duration_compact(started_at, now=now)
duration_label = format_duration_compact(started_at)
percent = snapshot.context_percent
percent_label = f"{percent}%"
if width < 52:
@@ -304,13 +215,6 @@ def build_status_fragments(
("class:status-bar", " "),
]
_append_memory_indicator(
frags,
worker_status=get_memory_worker_status(),
linker_status=get_observation_linker_status(),
width=width,
)
total_width = sum(_display_width(text) for _, text in frags)
if total_width > width:
plain_text = "".join(text for _, text in frags)
@@ -441,12 +345,11 @@ async def build_session_status_snapshot(
model_name: str | None = None,
model_obj: Any | None = None,
pending_user_text: str | None = None,
graph_gateway: GraphGateway,
) -> SessionStatusSnapshot:
"""Count current thread context and return a display snapshot."""
resolved_name = _resolve_model_name(model_name, model_obj)
window = _resolve_context_window(model_obj)
messages = list(await graph_gateway.get_thread_messages(thread_id))
messages = list(await get_thread_messages(thread_id))
pending = (pending_user_text or "").strip()
if pending:
-7
View File
@@ -6,7 +6,6 @@ from collections.abc import Callable
from dataclasses import dataclass
from typing import Any, Protocol
from ..gateway import GraphGateway
from ..stream.display import _run_streaming
@@ -31,8 +30,6 @@ class StreamingTUIBackend(Protocol):
metadata: dict | None = None,
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
cancel_scope: str | None = None,
gateway: GraphGateway,
) -> str:
"""Run streaming and return final response text."""
@@ -59,8 +56,6 @@ class RichStreamingBackend:
metadata: dict | None = None,
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
cancel_scope: str | None = None,
gateway: GraphGateway,
) -> str:
return _run_streaming(
agent=agent,
@@ -76,6 +71,4 @@ class RichStreamingBackend:
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
gateway=gateway,
)
File diff suppressed because it is too large Load Diff
+2 -13
View File
@@ -5,16 +5,11 @@ from __future__ import annotations
from collections.abc import Callable
from typing import Any
from ..gateway import GraphGateway
from ..stream.console import console
from ..stream.display import console
from .tui_backends import RichStreamingBackend, StreamingTUIBackend
DEFAULT_UI_BACKEND = "cli"
# "webui" launches the browser front-end instead of an in-terminal UI; it is
# intercepted earlier (cli/commands.py:_main_callback) and never reaches the
# streaming backends, but is listed here so normalize/resolve preserve it
# rather than falling back to "cli".
SUPPORTED_UI_BACKENDS = ("cli", "tui", "webui")
SUPPORTED_UI_BACKENDS = ("cli", "tui")
_LEGACY_BACKEND_MAP = {"textual": "tui", "rich": "cli"}
@@ -79,8 +74,6 @@ def run_streaming(
metadata: dict | None = None,
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
cancel_scope: str | None = None,
gateway: GraphGateway,
) -> str:
"""Run streaming with the selected backend."""
backend = get_backend(ui_backend, warn_fallback=True)
@@ -99,8 +92,6 @@ def run_streaming(
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
gateway=gateway,
)
except RuntimeError:
requested = normalize_ui_backend(ui_backend)
@@ -122,7 +113,5 @@ def run_streaming(
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
gateway=gateway,
)
raise
-2
View File
@@ -6,7 +6,6 @@ from .assistant_message import AssistantMessage
from .compact_summary_widget import CompactSummaryWidget
from .compacting_widget import CompactingWidget
from .loading_widget import LoadingWidget
from .mcp_loader_widget import MCPLoaderWidget
from .subagent_widget import SubAgentWidget
from .summarization_widget import SummarizationWidget
from .system_message import SystemMessage
@@ -24,7 +23,6 @@ __all__ = [
"CompactSummaryWidget",
"CompactingWidget",
"LoadingWidget",
"MCPLoaderWidget",
"SubAgentWidget",
"SummarizationWidget",
"SystemMessage",
+15 -3
View File
@@ -105,7 +105,11 @@ class ApprovalWidget(Widget):
self._option_widgets = []
count = len(self._action_requests)
if count == 1:
name = self._action_requests[0].get("name", "")
name = (
self._action_requests[0].get("name", "")
if isinstance(self._action_requests[0], dict)
else getattr(self._action_requests[0], "name", "")
)
title = f">>> {name} Requires Approval <<<"
else:
title = f">>> {count} Tool Calls Require Approval <<<"
@@ -113,8 +117,16 @@ class ApprovalWidget(Widget):
# Show each action request as a compact line
for req in self._action_requests:
name = req.get("name", "")
args = req.get("args", {})
name = (
req.get("name", "")
if isinstance(req, dict)
else getattr(req, "name", "")
)
args = (
req.get("args", {})
if isinstance(req, dict)
else getattr(req, "args", {})
)
if isinstance(args, dict):
command = args.get("command", args.get("path", ""))
else:
+4 -10
View File
@@ -5,7 +5,6 @@ from __future__ import annotations
from textual.containers import Vertical
from textual.widgets import Markdown
from ...stream.display import _fix_markdown_heading_spacing
from .timestamp_mixin import TimestampClickMixin
@@ -40,11 +39,8 @@ class AssistantMessage(TimestampClickMixin, Vertical):
yield Markdown("")
def on_mount(self) -> None:
"""Render ``initial_content`` once the widget enters the DOM."""
if self._content:
self.query_one(Markdown).update(
_fix_markdown_heading_spacing(self._content)
)
self.query_one(Markdown).update(self._content)
async def append_content(self, text: str) -> None:
"""Append text and schedule a debounced Markdown re-render."""
@@ -54,14 +50,12 @@ class AssistantMessage(TimestampClickMixin, Vertical):
self.set_timer(0.1, self._flush_markdown)
def _flush_markdown(self) -> None:
"""Flush accumulated content to the Markdown widget on a display copy."""
"""Flush accumulated content to the Markdown widget."""
self._flush_pending = False
self.query_one(Markdown).update(_fix_markdown_heading_spacing(self._content))
self.query_one(Markdown).update(self._content)
async def stop_stream(self) -> None:
"""Finalize the stream — ensure final content is rendered."""
self._flush_pending = False
if self._content:
self.query_one(Markdown).update(
_fix_markdown_heading_spacing(self._content)
)
self.query_one(Markdown).update(self._content)
@@ -1,191 +0,0 @@
"""Live-updating widget that shows per-server MCP load progress.
Mounted above the chat input while MCP tools are being fetched in the
background. Re-renders on a 100 ms tick so the spinner animates and the
per-server states transition smoothly from pending → ok/error.
When the load finishes:
- All-success runs auto-dismiss after a short grace period so the chat
area isn't permanently crowded.
- Failures stick around longer so the user has time to read the error
detail, then auto-dismiss — otherwise the widget pins itself above
the input forever.
"""
from __future__ import annotations
import time
from rich.text import Text
from textual.widgets import Static
from ..status_bar import SPINNER_FRAMES
_DIM = "#7c8594"
_STRONG = "#e5e7eb"
_GOOD = "#5fcf8b"
_WARN = "#d7b45a"
_BAD = "#d86f6f"
# How long to wait after an all-success load before auto-dismissing.
_AUTO_DISMISS_SECONDS = 2.5
# Longer grace on failure so the user has time to read error detail.
_AUTO_DISMISS_ON_ERROR_SECONDS = 12.0
class MCPLoaderWidget(Static):
"""Shows a header line + one line per MCP server with its live status."""
DEFAULT_CSS = """
MCPLoaderWidget {
height: auto;
padding: 0 1;
margin: 0 0 1 0;
}
"""
TICK_SECONDS = 0.1
def __init__(self, servers: list[str]) -> None:
# server_name -> (state, detail); state ∈ {"pending","ok","error"}.
self._progress: dict[str, tuple[str, str]] = dict.fromkeys(
servers, ("pending", "")
)
self._frame = 0
self._tick_handle = None
self._finished = False
self._dismissed = False
self._auto_dismiss_at: float | None = None
# Seed with real content so Textual can measure us before the
# first tick; ``self.update()`` during ``__init__`` is unsafe
# (widget isn't attached yet), but we can pass the renderable
# straight into ``Static.__init__``.
super().__init__(self._build_renderable())
def on_mount(self) -> None:
self._tick_handle = self.set_interval(self.TICK_SECONDS, self._tick)
def on_unmount(self) -> None:
if self._tick_handle is not None:
self._tick_handle.stop()
self._tick_handle = None
# ── Public API ───────────────────────────────────────────────────
@property
def dismissed(self) -> bool:
"""Whether the widget has already removed itself from the DOM."""
return self._dismissed
def update_server(self, name: str, state: str, detail: str = "") -> None:
"""Record a progress event for one server and re-render."""
if self._dismissed or state not in ("pending", "ok", "error"):
return
# First-time-seen servers (e.g., ones missing from the initial
# prime set because the config file changed mid-load) just get
# appended — order stays stable for already-known entries.
self._progress[name] = (state, detail)
self._refresh_content()
def mark_finished(self) -> None:
"""Call once the background load task resolves (success or error).
If nothing ever progressed past ``pending``, the load was served
from cache (no events emitted) — drop the widget immediately
instead of flashing a misleading "0/N loaded" header.
Otherwise schedule an auto-dismiss: short on full success so the
chat area isn't cluttered, longer on failure so the user has
time to read the error detail before it goes away.
"""
if self._finished:
return
self._finished = True
progressed = any(state != "pending" for state, _ in self._progress.values())
if not progressed:
self._dismiss()
return
has_errors = any(state == "error" for state, _ in self._progress.values())
delay = _AUTO_DISMISS_ON_ERROR_SECONDS if has_errors else _AUTO_DISMISS_SECONDS
self._auto_dismiss_at = time.monotonic() + delay
self._refresh_content()
# ── Internal ─────────────────────────────────────────────────────
def _tick(self) -> None:
self._frame = (self._frame + 1) % len(SPINNER_FRAMES)
if (
self._auto_dismiss_at is not None
and time.monotonic() >= self._auto_dismiss_at
):
self._auto_dismiss_at = None
self._dismiss()
return
if not self._finished:
self._refresh_content()
def _dismiss(self) -> None:
"""Stop the tick timer and detach from the DOM.
Sets :attr:`dismissed` so the app can clear its widget reference
and late progress events become no-ops.
"""
if self._dismissed:
return
self._dismissed = True
if self._tick_handle is not None:
self._tick_handle.stop()
self._tick_handle = None
# Fire-and-forget remove — nothing awaits us.
self.remove()
def _build_renderable(self) -> Text:
spinner = SPINNER_FRAMES[self._frame]
pending = sum(1 for state, _ in self._progress.values() if state == "pending")
total = len(self._progress)
done = total - pending
header = Text()
if self._finished:
errors = sum(1 for state, _ in self._progress.values() if state == "error")
if errors:
header.append("✗ MCP ", style=f"{_BAD} bold")
header.append(
f"{done - errors}/{total} loaded, {errors} failed",
style=_STRONG,
)
else:
header.append("✓ MCP ", style=f"{_GOOD} bold")
header.append(f"{done}/{total} servers loaded", style=_STRONG)
else:
header.append(f"{spinner} ", style=f"{_WARN} bold")
header.append("Loading MCP tools ", style=_STRONG)
header.append(f"{done}/{total}", style=_DIM)
lines: list[Text] = [header]
for name, (state, detail) in self._progress.items():
line = Text(" ")
if state == "pending":
line.append(f"{spinner} ", style=_WARN)
line.append(name, style=_DIM)
elif state == "ok":
line.append("✓ ", style=_GOOD)
line.append(name, style=_STRONG)
if detail:
line.append(f" {detail} tools", style=_DIM)
else: # error
line.append("✗ ", style=_BAD)
line.append(name, style=_STRONG)
if detail:
summary = detail if len(detail) <= 80 else detail[:77] + "…"
line.append(f" {summary}", style=_BAD)
lines.append(line)
return Text("\n").join(lines)
def _refresh_content(self) -> None:
# NB: don't name this ``_render`` — that shadows Textual's internal
# ``Widget._render`` which must return a ``Visual``. Silently
# breaking that contract triggers ``'NoneType' object has no
# attribute 'get_height'`` during layout.
self.update(self._build_renderable())
-390
View File
@@ -1,390 +0,0 @@
"""Inline model picker widget for /model command in TUI.
Keyboard-driven widget mounted directly into the chat container.
Models are grouped by provider with a search/filter input.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, ClassVar, Literal
from rich.text import Text
from textual.binding import Binding, BindingType
from textual.containers import Container
from textual.message import Message
from textual.widget import Widget
from textual.widgets import Input, Static
if TYPE_CHECKING:
from textual import events
from textual.app import ComposeResult
# Sentinel ``model_id`` used for the "Custom Ollama model..." pseudo-row.
# Selecting this row switches the widget into free-text input mode instead
# of posting ``Picked`` — the user types a model name, Enter confirms.
_CUSTOM_OLLAMA_ID = "__custom_ollama__"
def _build_items(
entries: list[tuple[str, str, str]],
current_model: str | None = None,
current_provider: str | None = None,
filter_text: str = "",
) -> list[dict]:
"""Build the flat item list rendered by ModelPickerWidget.
Returns a list of::
{"type": "header", "label": str}
{"type": "model", "name": str, "model_id": str, "provider": str, "current": bool}
"""
# Apply filter. The Custom Ollama sentinel is the user's escape hatch
# when no local models match; it must remain visible regardless of filter.
if filter_text:
ft = filter_text.lower()
entries = [
(n, mid, p)
for n, mid, p in entries
if mid == _CUSTOM_OLLAMA_ID or ft in n.lower() or ft in p.lower()
]
# Group by provider preserving order. Deduplicate the Custom Ollama
# sentinel defensively — if callers somehow pass two sentinel rows
# (state reuse, stale merges), collapse them into one to avoid
# rendering duplicate "Custom Ollama model..." rows in the picker.
groups: dict[str, list[tuple[str, str, str]]] = {}
seen_sentinel = False
for name, model_id, provider in entries:
if model_id == _CUSTOM_OLLAMA_ID:
if seen_sentinel:
continue
seen_sentinel = True
if provider not in groups:
groups[provider] = []
groups[provider].append((name, model_id, provider))
items: list[dict] = []
for provider, models in groups.items():
items.append({"type": "header", "label": provider})
for name, model_id, prov in models:
is_current = name == current_model and prov == current_provider
items.append(
{
"type": "model",
"name": name,
"model_id": model_id,
"provider": prov,
"current": is_current,
}
)
return items
class ModelPickerWidget(Widget):
"""Inline model picker -- mounts in chat, keyboard-driven.
Posts ``Picked(name, provider)`` on Enter, ``Cancelled()`` on Esc.
Type to filter models.
"""
can_focus = True
# Required so the Custom Ollama ``Input`` child can hold focus when the
# user is typing a model name.
can_focus_children = True
DEFAULT_CSS = """
ModelPickerWidget {
height: auto;
max-height: 30;
margin: 1 0;
padding: 0 1;
background: $surface;
border: solid $primary;
}
ModelPickerWidget .picker-custom-input {
height: 3;
margin: 1 0 0 0;
}
ModelPickerWidget .picker-title {
height: 1;
text-style: bold;
color: $primary;
}
ModelPickerWidget .picker-filter {
height: 1;
padding: 0 1;
color: $text;
}
ModelPickerWidget .picker-rows {
height: auto;
max-height: 22;
overflow-y: auto;
}
ModelPickerWidget .picker-header {
height: 1;
padding: 0 1;
margin-top: 1;
}
ModelPickerWidget .picker-row {
height: 1;
padding: 0 1;
}
ModelPickerWidget .picker-row-selected {
background: $primary;
text-style: bold;
}
ModelPickerWidget .picker-help {
height: 1;
color: $text-muted;
text-style: italic;
}
"""
BINDINGS: ClassVar[list[BindingType]] = [
Binding("up", "move_up", "Up", show=False),
Binding("down", "move_down", "Down", show=False),
Binding("enter", "select", "Select", show=False),
Binding("escape", "cancel", "Cancel", show=False),
Binding("backspace", "backspace", "Backspace", show=False),
]
class Picked(Message):
def __init__(self, name: str, provider: str) -> None:
super().__init__()
self.name = name
self.provider = provider
class Cancelled(Message):
"""Posted when user cancels selection."""
def __init__(
self,
entries: list[tuple[str, str, str]],
*,
current_model: str | None = None,
current_provider: str | None = None,
title: str = ">>> Select model <<<",
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
self._entries = entries
self._current_model = current_model
self._current_provider = current_provider
self._title = title
self._filter_text = ""
self._items = _build_items(
entries,
current_model=current_model,
current_provider=current_provider,
)
self._selected = self._first_model_index()
self._row_widgets: list[Static] = []
self._filter_widget: Static | None = None
# "list" = arrow-key selection over models; "input" = free-text entry
# for Custom Ollama model name. Transitions: selecting the sentinel
# row enters input mode; Esc or Up arrow inside input mode returns to
# list mode without closing the picker.
self._mode: Literal["list", "input"] = "list"
self._custom_input: Input | None = None
def _first_model_index(self) -> int:
for i, item in enumerate(self._items):
if item["type"] == "model":
return i
return 0
def _move(self, direction: int) -> None:
if not self._items:
return
i = (self._selected + direction) % len(self._items)
steps = 0
while self._items[i]["type"] != "model" and steps < len(self._items):
i = (i + direction) % len(self._items)
steps += 1
if self._items[i]["type"] == "model":
self._selected = i
self._update_rows()
def _rebuild(self) -> None:
"""Rebuild items from filter and re-render."""
self._items = _build_items(
self._entries,
current_model=self._current_model,
current_provider=self._current_provider,
filter_text=self._filter_text,
)
self._selected = self._first_model_index()
# Re-mount rows
rows_container = self.query_one(".picker-rows", Container)
for w in list(rows_container.children):
w.remove()
self._row_widgets.clear()
for item in self._items:
css = "picker-header" if item["type"] == "header" else "picker-row"
widget = Static("", classes=css)
self._row_widgets.append(widget)
rows_container.mount(widget)
self._update_rows()
self._update_filter()
def compose(self) -> ComposeResult:
yield Static(self._title, classes="picker-title")
self._filter_widget = Static("", classes="picker-filter")
yield self._filter_widget
with Container(classes="picker-rows"):
for item in self._items:
css = "picker-header" if item["type"] == "header" else "picker-row"
widget = Static("", classes=css)
self._row_widgets.append(widget)
yield widget
# Hidden until the user selects "Custom Ollama model..." \u2014 then shown
# and focused for free-text entry of an Ollama model name.
self._custom_input = Input(
placeholder="Type Ollama model name (e.g. llama3.3)...",
classes="picker-custom-input",
)
self._custom_input.display = False
yield self._custom_input
yield Static(
"\u2191/\u2193 navigate \u00b7 Enter select \u00b7 Type to filter \u00b7 Esc cancel",
classes="picker-help",
)
def on_mount(self) -> None:
self._update_rows()
self._update_filter()
self.call_later(self.focus)
def _update_filter(self) -> None:
if self._filter_widget is not None:
if self._filter_text:
t = Text()
t.append(" Filter: ", style="dim")
t.append(self._filter_text, style="bold")
t.append("\u2588", style="blink")
self._filter_widget.update(t)
else:
self._filter_widget.update(
Text(" Type to filter...", style="dim italic")
)
def _update_rows(self) -> None:
for i, (item, widget) in enumerate(
zip(self._items, self._row_widgets, strict=False)
):
widget.remove_class("picker-row-selected")
if item["type"] == "header":
t = Text()
t.append("\u2500\u2500 ", style="bold cyan")
t.append(item["label"], style="bold cyan")
widget.update(t)
else:
is_selected = i == self._selected
t = Text()
cursor = "\u25b8 " if is_selected else " "
t.append(cursor, style="bold cyan" if is_selected else "dim")
t.append(item["name"], style="bold" if is_selected else "")
if item["current"]:
t.append(" *", style="bold green")
t.append(f" ({item['provider']})", style="dim italic")
widget.update(t)
if is_selected:
widget.add_class("picker-row-selected")
widget.scroll_visible()
def on_key(self, event: events.Key) -> None:
# In input mode, the Input child owns printable keys + backspace.
if self._mode == "input":
return
# Let bindings handle special keys
if event.key in ("up", "down", "enter", "escape", "backspace"):
return
# Printable characters -> filter
if event.character and event.character.isprintable():
self._filter_text += event.character
self._rebuild()
event.prevent_default()
def action_backspace(self) -> None:
if self._mode == "input":
# Input widget handles its own backspace.
return
if self._filter_text:
self._filter_text = self._filter_text[:-1]
self._rebuild()
def action_move_up(self) -> None:
if self._mode == "input":
# Up from the Input field escapes back to list selection.
self._exit_input_mode()
return
self._move(-1)
def action_move_down(self) -> None:
if self._mode == "input":
# Down in input mode is ambiguous; absorb rather than toggle.
return
self._move(1)
def action_select(self) -> None:
if self._mode == "input":
self._submit_custom_input()
return
if not self._items or self._selected >= len(self._items):
self.post_message(self.Cancelled())
return
item = self._items[self._selected]
if item["type"] != "model":
self.post_message(self.Cancelled())
return
if item["provider"] == "ollama" and item["model_id"] == _CUSTOM_OLLAMA_ID:
self._enter_input_mode()
return
self.post_message(self.Picked(item["name"], item["provider"]))
def action_cancel(self) -> None:
if self._mode == "input":
# Esc returns to list selection; does NOT close the picker.
self._exit_input_mode()
return
self.post_message(self.Cancelled())
def on_blur(self, event: events.Blur) -> None:
# When the Input child has focus we must NOT steal it back.
if self._mode == "input":
return
self.call_after_refresh(self.focus)
def on_input_submitted(self, event: Input.Submitted) -> None:
"""Safety net: Enter fired inside the Input widget rather than
bubbling to ``action_select``. Route to the same submit path."""
if event.input is self._custom_input:
event.stop()
self._submit_custom_input()
def _enter_input_mode(self) -> None:
"""Show the Custom Ollama Input and move focus into it."""
self._mode = "input"
if self._custom_input is not None:
self._custom_input.display = True
# Carry any filter text over as a nice touch — user may have
# started typing a model name thinking it would filter.
self._custom_input.value = self._filter_text
self._custom_input.focus()
def _exit_input_mode(self) -> None:
"""Hide the Input and return focus to the list."""
self._mode = "list"
if self._custom_input is not None:
self._custom_input.display = False
self._custom_input.value = ""
self.focus()
def _submit_custom_input(self) -> None:
"""Confirm the typed Ollama model name. Empty input is a no-op —
user can Esc out or keep typing."""
typed = (self._custom_input.value if self._custom_input else "").strip()
if not typed:
return
self.post_message(self.Picked(typed, "ollama"))
+3 -4
View File
@@ -22,7 +22,7 @@ class SubAgentWidget(Vertical):
┌ ▶ Cooking with research-agent — Search literature ─┐
│ ✓ 8 completed │
│ ● tavily_search query="LLM attention" │
│ ● web_search query="LLM attention" │
│ ✓ 3 results │
└─────────────────────────────────────────────────────┘
@@ -208,11 +208,10 @@ class SubAgentWidget(Vertical):
widget.set_success(content)
else:
widget.set_error(content)
# Move from running to completed (dedup guards against repeat
# deliveries of the same tool result inflating the collapse summary).
# Move from running to completed
if matched_key and matched_key in self._running_ids:
self._running_ids.remove(matched_key)
if matched_key and matched_key not in self._completed_ids:
if matched_key:
self._completed_ids.append(matched_key)
self._update_visibility()
+2 -19
View File
@@ -18,7 +18,6 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Any, ClassVar
from prompt_toolkit.styles import Style as PtStyle # type: ignore[import-untyped]
from rich.text import Text
from textual.binding import Binding, BindingType
from textual.containers import Container
@@ -31,22 +30,6 @@ if TYPE_CHECKING:
from textual.app import ComposeResult
# Style for questionary pickers used by the Rich CLI ``/resume`` and
# ``/delete`` interactive selectors. Matches the slash-completion menu's
# visual language: gray (#888888) for non-selected, bold for selected,
# no background changes.
PICKER_STYLE = PtStyle.from_dict(
{
"questionmark": "#888888",
"question": "",
"pointer": "bold",
"highlighted": "bold",
"text": "#888888",
"answer": "bold",
}
)
# ---------------------------------------------------------------------------
# Path helpers
# ---------------------------------------------------------------------------
@@ -202,9 +185,9 @@ def build_row_text(
indented: bool = False,
) -> Text:
"""Thread row. *indented* adds extra leading space for L2-grouped rows."""
from ...sessions import _format_relative_time, short_thread_id
from ...sessions import _format_relative_time
tid = short_thread_id(thread["thread_id"])
tid = thread["thread_id"]
preview = thread.get("preview", "") or ""
msgs = thread.get("message_count", 0)
model = thread.get("model", "") or ""
+1 -9
View File
@@ -16,12 +16,7 @@ class UsageWidget(Static):
}
"""
def __init__(
self,
input_tokens: int,
output_tokens: int,
elapsed: str | None = None,
) -> None:
def __init__(self, input_tokens: int, output_tokens: int) -> None:
stats = Text(justify="right")
stats.append("[", style="dim italic")
stats.append("Usage: ", style="dim italic")
@@ -29,8 +24,5 @@ class UsageWidget(Static):
stats.append(" in · ", style="dim italic")
stats.append(f"{output_tokens:,}", style="green italic")
stats.append(" out", style="dim italic")
if elapsed:
stats.append(" · ", style="dim italic")
stats.append(f"Elapsed: {elapsed}", style="dim italic")
stats.append("]", style="dim italic")
super().__init__(stats)
@@ -1,39 +0,0 @@
"""Transient widget shown while a /resume restarts the langgraph dev subprocess.
Mirrors ``CompactingWidget`` — a timer-backed status line that ticks elapsed
seconds so the user has live feedback during the up-to-60s langgraph dev
workspace sync (subprocess stop + restart so deployed sub-agents see the
resumed thread's workspace).
"""
from __future__ import annotations
from .timed_status_widget import TimedStatusWidget
class WorkspaceSyncWidget(TimedStatusWidget):
"""Timer-backed status line for an in-progress workspace sync."""
DEFAULT_CSS = """
WorkspaceSyncWidget {
height: auto;
color: #94a3b8;
padding: 0 0;
margin: 0 0 1 0;
}
"""
def __init__(self) -> None:
super().__init__()
def _refresh_display(self) -> None:
self.update(
f"Syncing async sub-agent server to resumed workspace... "
f"({self.elapsed_seconds}s)"
)
async def cleanup(self) -> None:
"""Stop timer and remove from DOM."""
self._stop_timer()
if self.is_mounted:
await self.remove()
+1 -2
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from . import implementation
from .base import Argument, Command, CommandContext, CommandUI, SubCommand
from .base import Argument, Command, CommandContext, CommandUI
from .channel_ui import ChannelCommandUI
from .manager import CommandManager, manager
@@ -12,7 +12,6 @@ __all__ = [
"CommandContext",
"CommandManager",
"CommandUI",
"SubCommand",
"implementation",
"manager",
]
-177
View File
@@ -1,177 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass
from enum import StrEnum
class CompletionKind(StrEnum):
"""Discriminator for the kind of completion result."""
COMMANDS = "commands"
SUBCOMMANDS = "subcommands"
EMPTY = "empty"
_CATEGORY_ORDER = ["Session", "Skills", "MCP", "Channels", "Model", "General"]
@dataclass(frozen=True)
class CompletionCandidate:
"""A single completion suggestion with its replacement range."""
text: str
description: str
replace_start: int
replace_end: int
category: str = ""
@dataclass(frozen=True)
class CompletionResult:
"""The result of parsing a slash command input for completions."""
kind: CompletionKind
candidates: list[CompletionCandidate]
def compute_completions(text: str, cursor_pos: int) -> CompletionResult:
"""Parse *text* up to *cursor_pos* and return completion candidates.
This is the shared engine used by both the Rich CLI
(``SlashCommandCompleter``) and the TUI (``on_text_area_changed``).
Both thin adapters only need to translate the returned candidates
into their respective render/apply primitives.
"""
from .manager import manager as cmd_manager
before = text[:cursor_pos]
if not before.startswith("/"):
return CompletionResult(CompletionKind.EMPTY, [])
parts = before.split()
if not parts:
return CompletionResult(CompletionKind.EMPTY, [])
cmd_name = parts[0].lower()
has_trailing_space = before.endswith(" ")
# --- Top-level command completion ---
if len(parts) == 1:
prefix = before.lower().rstrip()
# Match commands by canonical name AND aliases
by_cat: dict[str, list[tuple[str, str]]] = {}
for cmd in cmd_manager.get_all_commands():
all_names = [cmd.name.lower()] + [
a.lower() if a.startswith("/") else f"/{a.lower()}" for a in cmd.alias
]
if any(n.startswith(prefix) for n in all_names):
by_cat.setdefault(cmd.category, []).append((cmd.name, cmd.description))
# Whether the typed prefix resolves to an exact command/alias
exact_cmd = cmd_manager.get_command(prefix)
if exact_cmd and not has_trailing_space:
return CompletionResult(CompletionKind.EMPTY, [])
if exact_cmd and has_trailing_space:
completions = exact_cmd.get_completions([""])
if completions:
insert_pos = len(before)
return CompletionResult(
CompletionKind.SUBCOMMANDS,
[
CompletionCandidate(
text=name,
description=desc,
replace_start=insert_pos,
replace_end=insert_pos,
)
for name, desc in completions
],
)
return CompletionResult(CompletionKind.EMPTY, [])
all_matches = [v for vs in by_cat.values() for v in vs]
if not all_matches:
return CompletionResult(CompletionKind.EMPTY, [])
# Build candidates ordered by category
candidates: list[CompletionCandidate] = []
for cat in _CATEGORY_ORDER:
for cmd_text, desc in by_cat.get(cat, []):
candidates.append(
CompletionCandidate(
text=cmd_text,
description=desc,
replace_start=0,
replace_end=len(before),
category=cat,
)
)
for cat, items in by_cat.items():
if cat not in _CATEGORY_ORDER:
for cmd_text, desc in items:
candidates.append(
CompletionCandidate(
text=cmd_text,
description=desc,
replace_start=0,
replace_end=len(before),
category=cat,
)
)
return CompletionResult(CompletionKind.COMMANDS, candidates)
# --- Subcommand / argument completion (len(parts) >= 2) ---
cmd = cmd_manager.get_command(cmd_name)
if cmd is None:
return CompletionResult(CompletionKind.EMPTY, [])
# Delegate to Command.get_completions for all depths
tokens = parts[1:]
if has_trailing_space:
tokens.append("")
completions = cmd.get_completions(tokens)
if not completions:
return CompletionResult(CompletionKind.EMPTY, [])
# Compute replacement range.
if tokens[-1]:
# User is typing a partial — replace it
sub_start = before.rfind(tokens[-1])
if sub_start < 0:
sub_start = len(before)
replace_end = len(before)
elif len(tokens) >= 2 and tokens[-2]:
# Trailing space after a token. Check if the previous token is
# a known subcommand name — if so, the completion is for the
# NEXT argument (insert at cursor). If not, the completions
# refine the partial (replace it).
prev = tokens[-2]
is_known_sub = any(sc.name == prev for sc in cmd.subcommands)
if not is_known_sub:
sub_start = before.rfind(prev)
if sub_start < 0:
sub_start = len(before)
else:
sub_start = len(before)
replace_end = len(before)
else:
sub_start = len(before)
replace_end = len(before)
return CompletionResult(
CompletionKind.SUBCOMMANDS,
[
CompletionCandidate(
text=name,
description=desc,
replace_start=sub_start,
replace_end=replace_end,
)
for name, desc in completions
],
)
+3 -87
View File
@@ -1,11 +1,8 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, ClassVar, Protocol, runtime_checkable
if TYPE_CHECKING:
from ..gateway import GraphGateway
from dataclasses import dataclass
from typing import Any, ClassVar, Protocol, runtime_checkable
@dataclass
@@ -18,15 +15,6 @@ class Argument:
required: bool = True
@dataclass
class SubCommand:
"""A subcommand of a parent slash command."""
name: str
description: str
arguments: list[Argument] = field(default_factory=list)
@runtime_checkable
class CommandUI(Protocol):
"""Protocol for UI operations that commands can perform."""
@@ -47,38 +35,16 @@ class CommandUI(Protocol):
async def wait_for_mcp_browse(
self, servers: list, installed_names: set[str], pre_filter_tag: str
) -> list | None: ...
async def wait_for_model_pick(
self,
entries: list[tuple[str, str, str]],
current_model: str | None,
current_provider: str | None,
) -> tuple[str, str] | None: ...
def clear_chat(self) -> None: ...
def request_quit(self) -> None: ...
def force_quit(self) -> None: ...
async def start_new_session(self) -> None: ...
def start_new_session(self) -> None: ...
async def handle_session_resume(
self, thread_id: str, workspace_dir: str | None = None
) -> None: ...
async def flush(self) -> None: ...
@dataclass
class ChannelRuntime:
"""Mutable handle to the agent + thread bound to running channels."""
agent: Any = None
thread_id: str | None = None
def bind(self, agent: Any, thread_id: str) -> None:
self.agent = agent
self.thread_id = thread_id
def clear(self) -> None:
self.agent = None
self.thread_id = None
@dataclass
class CommandContext:
"""Context passed to commands during execution."""
@@ -89,9 +55,6 @@ class CommandContext:
workspace_dir: str | None = None
checkpointer: Any = None
config: Any = None
channel_runtime: ChannelRuntime | None = None
graph_gateway: GraphGateway | None = None
command_error: str | None = None
# Real LLM input token count from last usage_metadata (includes system
# prompt + tool schemas). Used by /compact for accurate display.
input_tokens_hint: int | None = None
@@ -104,53 +67,6 @@ class Command(ABC):
alias: ClassVar[list[str]] = []
description: str
arguments: ClassVar[list[Argument]] = []
category: ClassVar[str] = "General"
subcommands: ClassVar[list[SubCommand]] = []
# When False, callers may dispatch this command without waiting for
# the background agent load to finish — important so recovery
# commands like ``/mcp add`` can run even when the MCP load is
# failing and ``_await_agent_ready`` would hang.
requires_agent: ClassVar[bool] = False
def needs_agent(self, args: list[str]) -> bool:
"""Whether this specific invocation needs the agent.
Default returns :attr:`requires_agent`. Override when a command
has a mix of agent-using and agent-free subcommands (e.g.
``/channel start`` vs ``/channel status``).
"""
return self.requires_agent
def get_completions(self, tokens: list[str]) -> list[tuple[str, str]]:
"""Return completions for args typed after the command name.
Default walks :attr:`subcommands` for the first positional token
only. Override for deeper levels (e.g. server names, thread IDs).
"""
if not self.subcommands:
return []
if len(tokens) <= 1:
prefix = tokens[0].lower() if tokens else ""
matches = [
(sc.name, sc.description)
for sc in self.subcommands
if sc.name.startswith(prefix)
]
# Exact match — subcommand already complete, hide popup
if len(matches) == 1 and matches[0][0] == prefix:
return []
return matches
# partial + trailing space: /mcp a → still show "add"
if len(tokens) == 2 and tokens[1] == "":
prefix = tokens[0].lower()
if any(sc.name == prefix for sc in self.subcommands):
return []
return [
(sc.name, sc.description)
for sc in self.subcommands
if sc.name.startswith(prefix)
]
return []
@abstractmethod
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
+7 -99
View File
@@ -1,23 +1,14 @@
from __future__ import annotations
import asyncio
import logging
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any
from typing import Any
from .base import CommandUI
if TYPE_CHECKING:
from ..gateway import GraphGateway
_logger = logging.getLogger(__name__)
class ChannelCommandUI(CommandUI):
"""CommandUI implementation for messaging channels with output buffering."""
_TEXT_CHUNK_LIMIT = 3500
@property
def supports_interactive(self) -> bool:
return False
@@ -25,67 +16,24 @@ class ChannelCommandUI(CommandUI):
def __init__(
self,
channel_msg: Any,
*,
graph_gateway: GraphGateway,
append_system_callback: Any = None,
start_new_session_callback: Callable[[], Awaitable[None]] | None = None,
start_new_session_callback: Any = None,
handle_session_resume_callback: Any = None,
):
self.msg = channel_msg
self.append_system_callback = append_system_callback
self.start_new_session_callback = start_new_session_callback
self.handle_session_resume_callback = handle_session_resume_callback
self.graph_gateway = graph_gateway
self._system_buffer: list[str] = []
def _queue_system(
self,
text: str,
style: str = "dim",
*,
mirror_local: bool = True,
) -> None:
if mirror_local and self.append_system_callback:
def append_system(self, text: str, style: str = "dim") -> None:
if self.append_system_callback:
self.append_system_callback(text, style)
# Buffer the text for grouped delivery to the channel
# We ignore style for grouping but keep it for individual lines if needed
self._system_buffer.append(text)
def append_system(self, text: str, style: str = "dim") -> None:
self._queue_system(text, style)
@staticmethod
def _extract_message_text(message: Any) -> str:
content = getattr(message, "content", "") or ""
if isinstance(content, list):
parts = [
block.get("text", "")
for block in content
if isinstance(block, dict) and block.get("type") == "text"
]
content = " ".join(parts) if parts else ""
return str(content).strip()
async def _send_text_chunks(self, text: str, *, mirror_local: bool = True) -> None:
"""Flush long plain-text payloads in channel-safe chunks."""
text = (text or "").strip()
if not text:
return
pending = text
while pending:
chunk = pending[: self._TEXT_CHUNK_LIMIT]
if len(pending) > self._TEXT_CHUNK_LIMIT:
split_at = chunk.rfind("\n")
if split_at > 0:
chunk = chunk[:split_at]
chunk = chunk.rstrip()
if not chunk:
chunk = pending[: self._TEXT_CHUNK_LIMIT]
self._queue_system(chunk, mirror_local=mirror_local)
await self.flush()
pending = pending[len(chunk) :].lstrip("\n")
async def flush(self) -> None:
"""Send all buffered system messages as a single grouped message."""
if not self._system_buffer:
@@ -181,9 +129,9 @@ class ChannelCommandUI(CommandUI):
def force_quit(self) -> None:
self.request_quit()
async def start_new_session(self) -> None:
def start_new_session(self) -> None:
if self.start_new_session_callback:
await self.start_new_session_callback()
self.start_new_session_callback()
else:
self.append_system(
"New session requested. Please restart the channel link or use /new if supported."
@@ -192,45 +140,5 @@ class ChannelCommandUI(CommandUI):
async def handle_session_resume(
self, thread_id: str, workspace_dir: str | None = None
) -> None:
mirror_local = self.handle_session_resume_callback is None
if self.handle_session_resume_callback:
await self.handle_session_resume_callback(thread_id, workspace_dir)
lines = [f"Resumed session: {thread_id}"]
try:
messages = await self.graph_gateway.get_thread_messages(thread_id)
except Exception as exc:
_logger.exception(
"Failed to load saved history for resumed thread %s",
thread_id,
)
lines.append(f"(history unavailable: {exc})")
await self._send_text_chunks("\n".join(lines), mirror_local=mirror_local)
return
display = [m for m in messages if getattr(m, "type", None) in ("human", "ai")]
if not display:
if messages:
lines.append("No displayable messages in this session.")
else:
lines.append("No saved messages in this session.")
await self._send_text_chunks("\n".join(lines), mirror_local=mirror_local)
return
HISTORY_WINDOW = 10
if len(display) > HISTORY_WINDOW:
display = display[-HISTORY_WINDOW:]
lines.append(f"Conversation history (last {HISTORY_WINDOW} messages):")
else:
lines.append("Conversation history:")
for message in display:
text = self._extract_message_text(message)
if not text:
continue
if getattr(message, "type", None) == "human":
lines.append(f"User: {text}")
else:
lines.append(f"EvoScientist: {text}")
await self._send_text_chunks("\n".join(lines), mirror_local=mirror_local)
@@ -1,25 +1,5 @@
from __future__ import annotations
from . import (
autoskills,
channel,
general,
mcp,
model,
model_fallback,
schedule,
session,
skills,
)
from . import channel, general, mcp, session, skills
__all__ = [
"autoskills",
"channel",
"general",
"mcp",
"model",
"model_fallback",
"schedule",
"session",
"skills",
]
__all__ = ["channel", "general", "mcp", "session", "skills"]
@@ -1,435 +0,0 @@
from __future__ import annotations
import asyncio
from enum import Enum
from typing import ClassVar
from rich.table import Table
from ..base import Command, CommandContext, SubCommand
from ..manager import manager
AUTOSKILLS_COMMAND = "/autoskills"
_PROPOSAL_STATUSES = {
"review": "pending",
"approved": "approved",
"rejected": "rejected",
}
class AutoSkillsCommand(Command):
"""Manage EvoMemory AutoSkills proposals."""
name = AUTOSKILLS_COMMAND
alias: ClassVar[list[str]] = ["/skills-review"]
description = "Review EvoMemory autoskill proposals"
subcommands: ClassVar[list[SubCommand]] = [
SubCommand("status", "Show AutoSkills config and proposals for review"),
SubCommand("help", "Show AutoSkills command examples"),
SubCommand("list", "List autoskill proposals, optionally filtered by status"),
SubCommand("review", "Review autoskill proposals awaiting a decision"),
SubCommand("approve", "Approve an autoskill proposal by id"),
SubCommand("reject", "Reject an autoskill proposal by id"),
SubCommand("run", "Run AutoSkills once now"),
SubCommand("on", "Enable periodic AutoSkills"),
SubCommand("off", "Disable periodic AutoSkills"),
SubCommand("mode", "Set review or auto approval mode"),
SubCommand("cadence", "Set nightly, weekly, or monthly cadence"),
SubCommand("time", "Set local run time as HH:MM"),
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
sub = args[0].lower() if args else "help"
rest = args[1:]
if sub in {"help", "-h", "--help", "?"}:
self._show_help(ctx)
elif sub in {"status", "show"}:
await self._status(ctx)
elif sub in {"list", "ls", "proposals"}:
await self._list_command(ctx, rest)
elif sub == "review":
await self._list(ctx, status="pending")
elif sub in {"approve", "accept"}:
await self._approve(ctx, self._first_arg(rest))
elif sub in {"reject", "deny", "decline"}:
await self._reject(ctx, self._first_arg(rest))
elif sub in {"run", "now"}:
await self._run(ctx)
elif sub in {"on", "enable"}:
await self._set_config(ctx, "memory_skill_synthesis_enabled", "true")
elif sub in {"off", "disable"}:
await self._set_config(ctx, "memory_skill_synthesis_enabled", "false")
elif sub == "mode":
await self._set_config(
ctx,
"memory_skill_synthesis_mode",
self._first_arg(rest),
)
elif sub in {"auto", "automatic"}:
await self._set_config(ctx, "memory_skill_synthesis_mode", "auto")
elif sub == "manual":
await self._set_config(ctx, "memory_skill_synthesis_mode", "review")
elif sub == "cadence":
await self._set_config(
ctx,
"memory_skill_synthesis_cadence",
self._first_arg(rest),
)
elif sub in {"nightly", "weekly", "monthly"}:
await self._set_config(ctx, "memory_skill_synthesis_cadence", sub)
elif sub == "time":
await self._set_config(
ctx,
"memory_skill_synthesis_time",
self._first_arg(rest),
)
else:
self._show_help(ctx, prefix=f"Unknown AutoSkills command: {sub}")
async def _status(self, ctx: CommandContext) -> None:
from ... import paths
from ...config import get_effective_config
from ...memory.autoskills.proposals import list_skill_proposals
from ...memory.autoskills.schedule import alist_autoskill_schedules
cfg = get_effective_config()
workspace_dir = self._workspace_dir(ctx)
pending = list_skill_proposals(
paths.MEMORIES_DIR,
status="pending",
workspace_dir=workspace_dir,
)
ctx.ui.append_system(
(
"AutoSkills: "
f"{'on' if cfg.memory_skill_synthesis_enabled else 'off'} | "
f"mode={cfg.memory_skill_synthesis_mode.value} | "
f"cadence={cfg.memory_skill_synthesis_cadence.value} | "
f"time={cfg.memory_skill_synthesis_time}"
),
style="dim",
)
ctx.ui.append_system(
f"AutoSkill proposal(s) ready for review: {len(pending)}",
style="yellow" if pending else "dim",
)
if pending:
ctx.ui.append_system(
(
f"Next: {AUTOSKILLS_COMMAND} review, then "
f"{AUTOSKILLS_COMMAND} approve <id> or "
f"{AUTOSKILLS_COMMAND} reject <id>."
),
style="dim",
)
elif cfg.memory_skill_synthesis_enabled:
ctx.ui.append_system(
f"Next: {AUTOSKILLS_COMMAND} run to search now, or "
f"{AUTOSKILLS_COMMAND} help for commands.",
style="dim",
)
else:
ctx.ui.append_system(
f"Next: {AUTOSKILLS_COMMAND} run to search once, or "
f"{AUTOSKILLS_COMMAND} on to enable scheduled runs.",
style="dim",
)
if cfg.memory_skill_synthesis_enabled:
try:
rows = await alist_autoskill_schedules(cfg, limit=1)
except Exception:
rows = []
if rows:
ctx.ui.append_system(
f"Background schedule id: {str(rows[0].get('cron_id', ''))[:8]}",
style="dim",
)
async def _list_command(self, ctx: CommandContext, args: list[str]) -> None:
if not args or args[0].lower() == "all":
await self._list(ctx)
return
status = _PROPOSAL_STATUSES.get(args[0].lower())
if status is None:
ctx.ui.append_system(
f"Usage: {AUTOSKILLS_COMMAND} list [review|approved|rejected|all]",
style="yellow",
)
return
await self._list(ctx, status=status)
async def _list(self, ctx: CommandContext, *, status: str | None = None) -> None:
from ... import paths
from ...memory.autoskills.proposals import list_skill_proposals
workspace_dir = self._workspace_dir(ctx)
proposals = list_skill_proposals(
paths.MEMORIES_DIR,
status=status,
workspace_dir=workspace_dir,
)
if not proposals:
if status:
label = self._status_label(status)
ctx.ui.append_system(
f"No autoskill proposals {label}.",
style="dim",
)
else:
ctx.ui.append_system("No autoskill proposals.", style="dim")
return
title = "EvoMemory AutoSkill Proposals"
if status:
title = (
f"EvoMemory AutoSkill Proposals {self._status_label(status).title()}"
)
table = Table(title=title, show_header=True)
table.add_column("ID", style="cyan")
table.add_column("Action", style="magenta")
table.add_column("AutoSkill", style="green")
table.add_column("Status", style="yellow")
table.add_column("Observations", justify="right")
table.add_column("Description", style="dim")
for proposal in proposals:
table.add_row(
proposal.proposal_id,
proposal.operation,
proposal.skill_name,
proposal.status,
str(len(proposal.source_observation_ids)),
proposal.description,
)
ctx.ui.mount_renderable(table)
ctx.ui.append_system(
f"Use {AUTOSKILLS_COMMAND} approve <id> or "
f"{AUTOSKILLS_COMMAND} reject <id>.",
style="dim",
)
async def _approve(self, ctx: CommandContext, proposal_id: str | None) -> None:
from ... import paths
from ...memory.autoskills.proposals import approve_skill_proposal
if not proposal_id:
ctx.ui.append_system(
f"Usage: {AUTOSKILLS_COMMAND} approve <id>",
style="yellow",
)
ctx.ui.append_system(
f"Run {AUTOSKILLS_COMMAND} review to copy a proposal ID.",
style="dim",
)
return
workspace_dir = self._workspace_dir(ctx)
result = await asyncio.to_thread(
approve_skill_proposal,
paths.MEMORIES_DIR,
proposal_id,
workspace_dir=workspace_dir,
)
if result.get("approved"):
verb = "Updated" if result.get("operation") == "update" else "Approved"
ctx.ui.append_system(
f"{verb} autoskill: {result['skill_name']} ({result['path']})",
style="green",
)
ctx.ui.append_system(
"Reload with /new to apply the new skill.", style="dim"
)
else:
ctx.ui.append_system(f"Approval failed: {result.get('error')}", style="red")
async def _reject(self, ctx: CommandContext, proposal_id: str | None) -> None:
from ... import paths
from ...memory.autoskills.proposals import reject_skill_proposal
if not proposal_id:
ctx.ui.append_system(
f"Usage: {AUTOSKILLS_COMMAND} reject <id>",
style="yellow",
)
ctx.ui.append_system(
f"Run {AUTOSKILLS_COMMAND} review to copy a proposal ID.",
style="dim",
)
return
workspace_dir = self._workspace_dir(ctx)
result = await asyncio.to_thread(
reject_skill_proposal,
paths.MEMORIES_DIR,
proposal_id,
workspace_dir=workspace_dir,
)
if result.get("rejected"):
ctx.ui.append_system(
f"Rejected proposal: {result['proposal_id']}",
style="green",
)
else:
ctx.ui.append_system(f"Reject failed: {result.get('error')}", style="red")
async def _run(self, ctx: CommandContext) -> None:
from ...config import get_effective_config
from ...memory.autoskills.schedule import arun_autoskill_now
workspace_dir = self._workspace_dir(ctx)
try:
result = await arun_autoskill_now(
get_effective_config(),
workspace_dir=workspace_dir,
)
except Exception as exc:
ctx.ui.append_system(f"Failed to start AutoSkills: {exc}", style="red")
return
ctx.ui.append_system(
f"Started AutoSkills run {result['run_id']}.",
style="green",
)
async def _set_config(
self,
ctx: CommandContext,
key: str,
value: str | None,
) -> None:
from ...config import get_effective_config, set_config_value
from ...memory.autoskills.schedule import reconcile_autoskill_schedule
workspace_dir = self._workspace_dir(ctx)
if not value:
cfg = get_effective_config()
current = self._display_value(getattr(cfg, key))
ctx.ui.append_system(
f"Current {self._config_label(key)}: {current}",
style="dim",
)
ctx.ui.append_system(
f"Usage: {self._config_usage(key)}",
style="yellow",
)
return
if not await asyncio.to_thread(set_config_value, key, value):
valid = self._config_values(key)
suffix = f" Valid values: {valid}." if valid else ""
ctx.ui.append_system(
f"Invalid value for {self._config_label(key)}: {value}.{suffix}",
style="red",
)
return
cfg = get_effective_config()
if ctx.config is not None and hasattr(ctx.config, key):
setattr(ctx.config, key, getattr(cfg, key))
await asyncio.to_thread(
reconcile_autoskill_schedule,
cfg,
workspace_dir=workspace_dir,
)
ctx.ui.append_system(
f"Updated {self._config_label(key)} = {self._display_value(getattr(cfg, key))}",
style="green",
)
@staticmethod
def _config_label(key: str) -> str:
labels = {
"memory_skill_synthesis_enabled": "AutoSkills",
"memory_skill_synthesis_mode": "AutoSkills mode",
"memory_skill_synthesis_cadence": "AutoSkills cadence",
"memory_skill_synthesis_time": "AutoSkills time",
}
return labels.get(key, key)
@staticmethod
def _display_value(value: object) -> object:
return getattr(value, "value", value)
@staticmethod
def _enum_values(enum_type: type[Enum], *, separator: str = ", ") -> str:
return separator.join(str(member.value) for member in enum_type)
@classmethod
def _config_usage(cls, key: str) -> str:
from ...config import MemorySkillSynthesisCadence, MemorySkillSynthesisMode
if key == "memory_skill_synthesis_mode":
values = cls._enum_values(MemorySkillSynthesisMode, separator="|")
return f"{AUTOSKILLS_COMMAND} mode {values}"
if key == "memory_skill_synthesis_cadence":
values = cls._enum_values(MemorySkillSynthesisCadence, separator="|")
return f"{AUTOSKILLS_COMMAND} cadence {values}"
if key == "memory_skill_synthesis_time":
return f"{AUTOSKILLS_COMMAND} time HH:MM"
return f"{AUTOSKILLS_COMMAND} <value>"
@classmethod
def _config_values(cls, key: str) -> str | None:
from ...config import MemorySkillSynthesisCadence, MemorySkillSynthesisMode
if key == "memory_skill_synthesis_mode":
return cls._enum_values(MemorySkillSynthesisMode)
if key == "memory_skill_synthesis_cadence":
return cls._enum_values(MemorySkillSynthesisCadence)
if key == "memory_skill_synthesis_time":
return "24-hour local time, for example 03:00"
return None
@staticmethod
def _status_label(status: str) -> str:
if status == "pending":
return "ready for review"
return status
@staticmethod
def _show_help(ctx: CommandContext, *, prefix: str | None = None) -> None:
if prefix:
ctx.ui.append_system(prefix, style="yellow")
ctx.ui.append_system(
(
f"Usage: {AUTOSKILLS_COMMAND} "
"[status|review|approve|reject|run|on|off|mode|cadence|time]"
),
style="bold",
)
table = Table(title="AutoSkills Commands", show_header=True)
table.add_column("Command", style="cyan")
table.add_column("Use when", style="dim")
rows = [
(AUTOSKILLS_COMMAND, "Show this command reference"),
(f"{AUTOSKILLS_COMMAND} status", "Show config and the next useful action"),
(f"{AUTOSKILLS_COMMAND} review", "Review proposals waiting for a decision"),
(f"{AUTOSKILLS_COMMAND} approve <id>", "Install a reviewed autoskill"),
(f"{AUTOSKILLS_COMMAND} reject <id>", "Dismiss a reviewed proposal"),
(f"{AUTOSKILLS_COMMAND} run", "Start a one-off background autoskill run"),
(f"{AUTOSKILLS_COMMAND} on|off", "Enable or disable scheduled runs"),
(f"{AUTOSKILLS_COMMAND} auto|manual", "Switch approval behavior"),
(
f"{AUTOSKILLS_COMMAND} nightly|weekly|monthly",
"Set the built-in schedule cadence",
),
(f"{AUTOSKILLS_COMMAND} time 03:00", "Set the local schedule time"),
(
f"{AUTOSKILLS_COMMAND} list [status]",
"List all proposals or filter by review, approved, or rejected",
),
]
for command, description in rows:
table.add_row(command, description)
ctx.ui.mount_renderable(table)
ctx.ui.append_system(
"Aliases: /skills-review, ls, proposals, accept, deny, enable, disable, now.",
style="dim",
)
@staticmethod
def _workspace_dir(ctx: CommandContext) -> str:
from ... import paths
return str(ctx.workspace_dir or paths.WORKSPACE_ROOT)
@staticmethod
def _first_arg(args: list[str]) -> str | None:
return args[0] if args else None
manager.register(AutoSkillsCommand())
@@ -1,12 +1,10 @@
from __future__ import annotations
from typing import ClassVar
from rich.panel import Panel
from rich.table import Table
from rich.text import Text
from ..base import Command, CommandContext, SubCommand
from ..base import Command, CommandContext
from ..manager import manager
@@ -15,27 +13,6 @@ class ChannelCommand(Command):
name = "/channel"
description = "Configure messaging channels"
category = "Channels"
subcommands: ClassVar[list[SubCommand]] = [
SubCommand("status", "Show channel status"),
SubCommand("stop", "Stop running channels"),
SubCommand("telegram", "Start Telegram channel"),
SubCommand("discord", "Start Discord channel"),
SubCommand("slack", "Start Slack channel"),
SubCommand("feishu", "Start Feishu channel"),
SubCommand("dingtalk", "Start DingTalk channel"),
SubCommand("wechat", "Start WeChat channel"),
SubCommand("email", "Start Email channel"),
SubCommand("imessage", "Start iMessage channel"),
]
def needs_agent(self, args: list[str]) -> bool:
# ``status`` and ``stop`` are introspection / teardown; they
# must work even when the agent load is still in flight or has
# failed. Only start/add flows feed ``ctx.agent`` into
# ``_start_channels_bus_mode``.
subcmd = args[0].lower() if args else ""
return subcmd not in {"status", "stop"}
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
import EvoScientist.cli.channel as _ch_mod
@@ -86,10 +63,10 @@ class ChannelCommand(Command):
ctx.ui.append_system("No channels are running.", style="dim")
else:
if target:
_channels_stop(target, runtime=ctx.channel_runtime)
_channels_stop(target)
ctx.ui.append_system(f"Channel '{target}' stopped.", style="green")
else:
_channels_stop(runtime=ctx.channel_runtime)
_channels_stop()
ctx.ui.append_system("All channels stopped.", style="green")
return
@@ -113,11 +90,6 @@ class ChannelCommand(Command):
ctx.ui.append_system(f"Adding channel(s): {', '.join(requested)}...")
from ...cli.channel import _add_channel_to_running_bus
# Bind the runtime up-front so partial-success states (one
# channel attached, next one raises) still leave the bus
# observing the latest agent/thread refs.
if ctx.channel_runtime is not None:
ctx.channel_runtime.bind(ctx.agent, ctx.thread_id)
try:
for ct in requested:
_add_channel_to_running_bus(ct, config, send_thinking=send_thinking)
@@ -151,8 +123,6 @@ class ChannelCommand(Command):
ctx.thread_id,
send_thinking=send_thinking,
)
if ctx.channel_runtime is not None:
ctx.channel_runtime.bind(ctx.agent, ctx.thread_id)
# Show status panel
if _ch_mod._manager:
@@ -29,9 +29,6 @@ class HelpCommand(Command):
if cmd.alias:
desc += f" (aliases: {', '.join(cmd.alias)})"
help_text.append(f"{desc}\n", style="dim")
if cmd.subcommands:
names = ", ".join(sc.name for sc in cmd.subcommands)
help_text.append(f" subcommands: {names}\n", style="dim italic")
ctx.ui.mount_renderable(help_text)
@@ -52,13 +49,17 @@ class CurrentCommand(Command):
f"Workspace: {_shorten_path(ctx.workspace_dir)}",
style="dim",
)
memory_path = paths.MEMORIES_DIR
memory_path = paths.MEMORY_DIR
if memory_path:
from ...cli.agent import _shorten_path
ctx.ui.append_system(
f"Memory dir: {_shorten_path(str(memory_path))}", style="dim"
)
# How to determine UI type here?
# Maybe ctx.ui has a name or we pass it in ctx.
# For now, let's keep it simple.
ctx.ui.append_system("UI: auto", style="dim")
# Register commands
+18 -53
View File
@@ -1,10 +1,8 @@
from __future__ import annotations
from typing import ClassVar
from rich.table import Table
from ..base import Command, CommandContext, SubCommand
from ..base import Command, CommandContext
from ..manager import manager
@@ -13,46 +11,8 @@ class MCPCommand(Command):
name = "/mcp"
description = "Manage MCP servers"
category = "MCP"
subcommands: ClassVar[list[SubCommand]] = [
SubCommand("list", "List configured MCP servers"),
SubCommand("config", "Show server configuration details"),
SubCommand("add", "Add a new MCP server"),
SubCommand("edit", "Edit an MCP server configuration"),
SubCommand("remove", "Remove an MCP server"),
SubCommand("install", "Browse and install MCP servers"),
]
_server_names_cache: list[str] | None = None
def _get_server_names(self) -> list[str]:
if self._server_names_cache is None:
try:
from ...mcp import load_mcp_config
self._server_names_cache = list(load_mcp_config().keys())
except Exception:
return []
return self._server_names_cache
def _invalidate_server_cache(self) -> None:
self._server_names_cache = None
def get_completions(self, tokens: list[str]) -> list[tuple[str, str]]:
if len(tokens) <= 1:
return super().get_completions(tokens)
subcmd = tokens[0].lower()
if subcmd in ("config", "remove", "edit") and len(tokens) == 2:
prefix = tokens[1].lower()
return [
(name, "")
for name in self._get_server_names()
if name.lower().startswith(prefix)
]
return super().get_completions(tokens)
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
"""Dispatch to the appropriate MCP subcommand."""
if not args or args[0] == "list":
await self._mcp_list(ctx)
return
@@ -64,26 +24,35 @@ class MCPCommand(Command):
await self._mcp_config(ctx, subargs[0] if subargs else "")
elif subcmd == "add":
await self._mcp_add(ctx, subargs)
self._invalidate_server_cache()
elif subcmd == "edit":
await self._mcp_edit(ctx, subargs)
self._invalidate_server_cache()
elif subcmd == "remove":
await self._mcp_remove(ctx, subargs[0] if subargs else "")
self._invalidate_server_cache()
elif subcmd == "install":
from .mcp_install import InstallMCPCommand
await InstallMCPCommand().execute(ctx, subargs)
else:
ctx.ui.append_system("MCP commands:", style="bold")
for sub in self.subcommands:
ctx.ui.append_system(
f" /mcp {sub.name:<12} {sub.description}", style="dim"
)
ctx.ui.append_system(
" /mcp List configured servers", style="dim"
)
ctx.ui.append_system(
" /mcp list List configured servers", style="dim"
)
ctx.ui.append_system(
" /mcp config Show detailed server config", style="dim"
)
ctx.ui.append_system(" /mcp add ... Add a server", style="dim")
ctx.ui.append_system(
" /mcp edit ... Edit an existing server", style="dim"
)
ctx.ui.append_system(" /mcp remove ... Remove a server", style="dim")
ctx.ui.append_system(
" /mcp install ... Browse and install servers", style="dim"
)
async def _mcp_list(self, ctx: CommandContext) -> None:
"""Display a table of all configured MCP servers."""
from ...mcp import load_mcp_config
from ...mcp.client import USER_MCP_CONFIG
@@ -116,7 +85,6 @@ class MCPCommand(Command):
ctx.ui.append_system(f"Config file: {USER_MCP_CONFIG}", style="dim")
async def _mcp_config(self, ctx: CommandContext, name: str) -> None:
"""Show detailed configuration for one or all MCP servers."""
from ...mcp import load_mcp_config
from ...mcp.client import USER_MCP_CONFIG
@@ -165,7 +133,6 @@ class MCPCommand(Command):
ctx.ui.append_system(f"Config file: {USER_MCP_CONFIG}", style="dim")
async def _mcp_add(self, ctx: CommandContext, tokens: list[str]) -> None:
"""Add a new MCP server from parsed arguments."""
from ...mcp import add_mcp_server, parse_mcp_add_args
if not tokens:
@@ -186,7 +153,6 @@ class MCPCommand(Command):
ctx.ui.append_system(f"Error: {exc}", style="red")
async def _mcp_edit(self, ctx: CommandContext, tokens: list[str]) -> None:
"""Edit fields of an existing MCP server configuration."""
from ...mcp import edit_mcp_server, parse_mcp_edit_args
if not tokens:
@@ -204,7 +170,6 @@ class MCPCommand(Command):
ctx.ui.append_system(f"Error: {exc}", style="red")
async def _mcp_remove(self, ctx: CommandContext, name: str) -> None:
"""Remove an MCP server by name."""
from ...mcp import remove_mcp_server
if not name:
@@ -10,7 +10,6 @@ class InstallMCPCommand(Command):
name = "/install-mcp"
description = "Browse and install MCP servers"
category = "MCP"
arguments: ClassVar[list[Argument]] = [
Argument(
name="source",
@@ -1,204 +0,0 @@
from __future__ import annotations
from typing import ClassVar
from ..base import Argument, Command, CommandContext
from ..manager import manager
def extract_model_and_provider(args: list[str]) -> tuple[str, str]:
"""Parse model name and provider from argument list.
Args:
args: Non-empty argument list (model_name [provider]).
Returns:
``(model_name, provider)`` tuple.
Raises:
ValueError: If the model is not in the registry. Skipped when
``provider_override == "ollama"``, since Ollama models are
locally-installed and never appear in ``MODELS``.
"""
from ...llm.models import MODELS
model_name = args[0]
provider_override = args[1] if len(args) > 1 else None
# Ollama models are locally-installed — not in the registry. Pass the name
# through verbatim; get_chat_model's "Assume full model ID" fallback
# (models.py) accepts them.
if provider_override == "ollama":
return model_name, "ollama"
if model_name not in MODELS:
raise ValueError(f"Unknown model '{model_name}'")
if provider_override:
provider = provider_override
else:
_, provider = MODELS[model_name]
return model_name, provider
class ModelCommand(Command):
"""Switch the LLM model for the current session."""
name = "/model"
description = "Switch model (--save to persist)"
category = "Model"
# ``--save`` is parsed manually in ``execute`` via ``"--save" in args``;
# ``type=bool`` below is declarative metadata, not enforced by the manager.
arguments: ClassVar[list[Argument]] = [
Argument(
name="model_name",
type=str,
description="Model short name (e.g. claude-sonnet-4-6). Opens picker if omitted.",
required=False,
),
Argument(
name="--save",
type=bool,
description="Save the choice to config file",
required=False,
),
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...EvoScientist import _ensure_config
from ...llm.models import list_model_picker_entries
cfg = _ensure_config()
current_model = cfg.model
current_provider = cfg.provider
# Parse --save flag
save = "--save" in args
args = [a for a in args if a != "--save"]
if args:
try:
model_name, provider = extract_model_and_provider(args)
except ValueError:
ctx.ui.append_system(
f"Unknown model '{args[0]}'. Use /model to browse available models.",
style="red",
)
return
await self._apply_model(ctx, model_name, provider, save=save)
return
# Interactive picker
if not ctx.ui.supports_interactive:
ctx.ui.append_system(
"Usage: /model <name> [provider] [--save]",
style="yellow",
)
return
entries = await list_model_picker_entries(
getattr(cfg, "ollama_base_url", None),
include_custom_ollama=True,
)
result = await ctx.ui.wait_for_model_pick(
entries,
current_model=current_model,
current_provider=current_provider,
)
if result is None:
return
name, provider = result
# Defense-in-depth: the widget should have replaced the sentinel with
# the user-typed name. If it didn't, treat as cancel rather than try
# to switch to a literal "__custom_ollama__" model.
if provider == "ollama" and name in (
"Custom Ollama model...",
"__custom_ollama__",
):
return
await self._apply_model(ctx, name, provider, save=save)
async def _apply_model(
self,
ctx: CommandContext,
model_name: str,
provider: str,
*,
save: bool = False,
) -> None:
import copy
from ...cli.agent import _load_agent
from ...EvoScientist import (
_build_chat_model,
_ensure_config,
set_active_config,
set_chat_model_instance,
)
cfg = _ensure_config()
# Build a temporary config + its chat model and verify the agent can be
# built before committing anything. ``create_cli_agent(config=...,
# chat_model=...)`` is pure (issue #183) — it writes none of the cached
# config/model module globals — so a failure below leaves the session
# on the original model with no snapshot/restore needed.
temp_cfg = copy.copy(cfg)
temp_cfg.model = model_name
temp_cfg.provider = provider
try:
new_chat_model = _build_chat_model(temp_cfg)
new_agent = _load_agent(
workspace_dir=ctx.workspace_dir,
checkpointer=ctx.checkpointer,
config=temp_cfg,
chat_model=new_chat_model,
)
except Exception as e:
ctx.ui.append_system(f"Failed to switch model: {e}", style="red")
return
# Agent built with no global mutation — commit the switch atomically.
# These are pure assignments and cannot fail, so the session can never
# be left half-switched. Apply the switch to the LIVE ``cfg`` in place
# (the active config object) instead of rebinding ``_config`` to the
# fresh ``temp_cfg`` — callers that hold the active config by reference
# (e.g. serve's ``agent_holder["config"]`` and its workspace-changing
# ``/resume`` reload) must observe the new model/provider. The verify
# build above used the ``temp_cfg`` copy, so a failed build never reaches
# here and the live ``cfg`` stays untouched (failure still no-ops).
cfg.model = model_name
cfg.provider = provider
set_active_config(cfg)
set_chat_model_instance(new_chat_model, (model_name, provider))
ctx.agent = new_agent
# Persist to config file if --save was given
if save:
from ...config.settings import set_config_value
set_config_value("model", model_name)
set_config_value("provider", provider)
# Propagate to the channel runtime if channels are running so the
# bus picks up the new agent on the next inbound message.
if ctx.channel_runtime is not None and ctx.channel_runtime.agent is not None:
ctx.channel_runtime.agent = new_agent
# Update status bar if available
update_model_fn = getattr(ctx.ui, "update_status_after_model_change", None)
if callable(update_model_fn):
update_model_fn(model_name, provider)
saved_note = " (saved to config)" if save else ""
ctx.ui.append_system(
f"Switched to {model_name} ({provider}){saved_note}", style="green"
)
manager.register(ModelCommand())
@@ -1,304 +0,0 @@
"""Slash command for managing the model fallback chain.
Provides ``/model-fallback`` (alias ``/fallback``) with subcommands to
add, remove, list, clear, save, and display help for fallback models.
"""
from __future__ import annotations
from typing import ClassVar
from ..base import Argument, Command, CommandContext, SubCommand
from ..manager import manager
class ModelFallbackCommand(Command):
"""Manage the model fallback chain."""
name = "/model-fallback"
alias: ClassVar[list[str]] = ["/fallback"]
description = "Manage fallback models (add/remove/list/clear)"
category = "Model"
arguments: ClassVar[list[Argument]] = [
Argument(
name="action",
type=str,
description="add|remove|list|clear|save|help",
required=False,
),
]
subcommands: ClassVar[list[SubCommand]] = [
SubCommand("list", "Display the current fallback chain"),
SubCommand("add", "Append a model to the fallback chain"),
SubCommand("remove", "Remove a model by position"),
SubCommand("clear", "Remove all fallback entries"),
SubCommand("save", "Persist the chain to config"),
SubCommand("help", "Show subcommand reference"),
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...llm.models import MODELS
from ...middleware.model_fallback import (
add_fallback,
clear_fallbacks,
get_fallback_chain,
remove_fallback_at,
serialize_fallback_chain,
)
save = "--save" in args
args = [a for a in args if a != "--save"]
if not args:
await self._show_list(ctx, get_fallback_chain())
return
action = args[0].lower()
if action == "list":
await self._show_list(ctx, get_fallback_chain())
elif action == "add":
if len(args) >= 2:
model_name = args[1]
provider = args[2] if len(args) > 2 else None
if provider is None:
if model_name in MODELS:
_, provider = MODELS[model_name]
else:
ctx.ui.append_system(
f"Unknown model '{model_name}'. Specify provider explicitly: "
f"/model-fallback add {model_name} <provider>",
style="red",
)
return
else:
picked = await self._pick_model(ctx)
if picked is None:
return
model_name, provider = picked
if add_fallback(model_name, provider):
ctx.ui.append_system(
f"Added {model_name} ({provider}) to fallback chain", style="green"
)
else:
ctx.ui.append_system(
f"{model_name} ({provider}) is already in the fallback chain",
style="yellow",
)
return
if save:
self._save_to_config(serialize_fallback_chain())
elif action == "remove":
chain = get_fallback_chain()
if not chain:
ctx.ui.append_system("Fallback chain is empty", style="yellow")
return
if len(args) >= 2:
arg = args[1]
try:
idx = int(arg) - 1
except ValueError:
ctx.ui.append_system(
f"Expected a position number (1-{len(chain)}), got '{arg}'. "
"Use /model-fallback list to see positions.",
style="red",
)
return
removed = remove_fallback_at(idx)
if removed is None:
ctx.ui.append_system(
f"Invalid position {arg}. "
f"Use a number between 1 and {len(chain)}.",
style="red",
)
return
model_name, provider = removed
else:
picked = await self._pick_fallback_to_remove(ctx, chain)
if picked is None:
return
model_name, provider = picked
live_chain = get_fallback_chain()
try:
idx = live_chain.index((model_name, provider))
except ValueError:
ctx.ui.append_system(
f"{model_name} ({provider}) is no longer in the fallback chain",
style="yellow",
)
return
remove_fallback_at(idx)
ctx.ui.append_system(
f"Removed {model_name} ({provider}) from fallback chain",
style="green",
)
if save:
self._save_to_config(serialize_fallback_chain())
elif action == "clear":
clear_fallbacks()
ctx.ui.append_system("Cleared all fallback models", style="green")
if save:
self._save_to_config("")
elif action == "save":
self._save_to_config(serialize_fallback_chain())
ctx.ui.append_system("Fallback chain saved to config", style="green")
elif action == "help":
self._show_help(ctx)
else:
self._show_help(ctx)
async def _pick_model(self, ctx: CommandContext) -> tuple[str, str] | None:
"""Open the interactive model picker to select a fallback model.
Falls back to a usage hint when the UI does not support interactive
widgets (CLI mode without a model argument).
Args:
ctx: Current command context.
Returns:
``(model_name, provider)`` tuple, or ``None`` if cancelled.
"""
if not ctx.ui.supports_interactive:
ctx.ui.append_system(
"Usage: /model-fallback add <model> [provider]", style="yellow"
)
return None
from ...EvoScientist import _ensure_config
from ...llm.models import list_model_picker_entries
cfg = _ensure_config()
entries = await list_model_picker_entries(
getattr(cfg, "ollama_base_url", None),
include_custom_ollama=True,
)
result = await ctx.ui.wait_for_model_pick(
entries,
current_model=cfg.model,
current_provider=cfg.provider,
)
if result is None:
return None
name, provider = result
if provider == "ollama" and name in (
"Custom Ollama model...",
"__custom_ollama__",
):
return None
return name, provider
async def _pick_fallback_to_remove(
self, ctx: CommandContext, chain: list[tuple[str, str]]
) -> tuple[str, str] | None:
"""Open the model picker populated with the current fallback chain.
Falls back to a usage hint in CLI mode.
Args:
ctx: Current command context.
chain: The current fallback chain to choose from.
Returns:
``(model_name, provider)`` tuple, or ``None`` if cancelled.
"""
if not ctx.ui.supports_interactive:
ctx.ui.append_system(
"Usage: /model-fallback remove <position> "
"(use /model-fallback list to see positions)",
style="yellow",
)
return None
entries = [(m, m, p) for m, p in chain]
result = await ctx.ui.wait_for_model_pick(
entries, current_model=None, current_provider=None
)
if result is None:
return None
return result
def _show_help(self, ctx: CommandContext) -> None:
"""Render the subcommand reference table.
Args:
ctx: Current command context.
"""
from rich.text import Text
text = Text("/model-fallback subcommands:\n", style="bold")
for cmd, desc in (
(
"add [model] [provider]",
"Add a fallback model (opens picker if omitted)",
),
(
"remove [position]",
"Remove a fallback by position (opens picker in TUI)",
),
("list", "Show the current fallback chain"),
("clear", "Remove all fallback models"),
("save", "Save current fallback chain to config file"),
("help", "Show this help message"),
):
text.append(f" {cmd:<26}", style="cyan")
text.append(f"{desc}\n", style="dim")
text.append(
"\nAdd --save to add/remove/clear to persist the change immediately.\n",
style="dim",
)
ctx.ui.mount_renderable(text)
async def _show_list(
self, ctx: CommandContext, chain: list[tuple[str, str]]
) -> None:
"""Display the current fallback chain as a numbered list.
Args:
ctx: Current command context.
chain: The fallback chain to display.
"""
if not chain:
ctx.ui.append_system("No fallback models configured", style="dim")
ctx.ui.append_system(
"Use /model-fallback add <model> [provider] to add one",
style="dim",
)
return
from rich.text import Text
text = Text("Fallback chain:\n", style="bold")
for idx, (model, provider) in enumerate(chain, 1):
text.append(f" {idx}. ", style="dim")
text.append(model, style="cyan")
text.append(f" ({provider})\n", style="dim")
ctx.ui.mount_renderable(text)
def _save_to_config(self, value: str) -> None:
"""Persist the fallback chain string to the config file.
Args:
value: Serialized chain (``"model:provider,..."``).
"""
from ...config.settings import set_config_value
set_config_value("model_fallbacks", value)
manager.register(ModelFallbackCommand())
@@ -1,234 +0,0 @@
from __future__ import annotations
import asyncio
import re
from typing import ClassVar
from rich.table import Table
from ..base import Command, CommandContext, SubCommand
from ..manager import manager
class ScheduleCommand(Command):
"""Manage scheduled (cron) tasks."""
name = "/schedule"
description = "Manage scheduled (cron) tasks"
subcommands: ClassVar[list[SubCommand]] = [
SubCommand("add", 'Add: /schedule add <m h dom mon dow> "<prompt>"'),
SubCommand("list", "List scheduled tasks"),
SubCommand("remove", "Remove a schedule by id"),
SubCommand("run", "Run a schedule's prompt once now (test)"),
SubCommand("pause", "Disable a schedule by id"),
SubCommand("resume", "Enable a schedule by id"),
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
"""Dispatch to the appropriate /schedule subcommand."""
cfg = getattr(ctx, "config", None)
if cfg is not None and not getattr(cfg, "enable_scheduler", True):
ctx.ui.append_system(
"Scheduled tasks are disabled (`enable_scheduler` is off).",
style="yellow",
)
return
from ...cron import schedule as crons
# Cron SDK calls are sync HTTP; offload to a thread so backend latency
# can never freeze the interactive event loop.
if not await asyncio.to_thread(crons.is_available):
ctx.ui.append_system(
"Scheduler unavailable: the langgraph dev backend is not running.",
style="yellow",
)
return
if not args or args[0].lower() == "list":
await self._list(ctx, crons)
return
sub = args[0].lower()
rest = args[1:]
if sub == "add":
await self._add(ctx, crons, rest)
elif sub == "remove":
await self._remove(ctx, crons, rest[0] if rest else "")
elif sub == "run":
await self._run(ctx, crons, rest[0] if rest else "")
elif sub in ("pause", "resume"):
await self._set_enabled(
ctx, crons, rest[0] if rest else "", sub == "resume"
)
else:
ctx.ui.append_system("Schedule commands:", style="bold")
for s in self.subcommands:
ctx.ui.append_system(
f" /schedule {s.name:<8} {s.description}", style="dim"
)
async def _add(self, ctx: CommandContext, crons, rest: list[str]) -> None:
# Cron may arrive as 5 separate tokens (unquoted) or 1 token (shlex-quoted).
# Split on any whitespace so extra spaces don't break detection; the
# backend rejects genuinely malformed expressions.
if rest and len(rest[0].split()) == 5: # quoted 5-field cron
schedule, prompt_tokens = " ".join(rest[0].split()), rest[1:]
elif len(rest) >= 5: # 5 separate cron fields
schedule, prompt_tokens = " ".join(rest[:5]), rest[5:]
else:
ctx.ui.append_system(
'Usage: /schedule add "<m h dom mon dow>" "<prompt>"', style="yellow"
)
return
prompt = " ".join(prompt_tokens).strip().strip('"').strip("'")
if not prompt:
ctx.ui.append_system("A task prompt is required.", style="yellow")
return
# B3: strip unsafe chars; keep only alphanumerics + hyphens (kebab-case).
raw = prompt[:48].lower()
name = re.sub(r"[^a-z0-9]+", "-", raw).strip("-")[:32] or "task"
try:
rec = await asyncio.to_thread(
crons.create_schedule, name=name, schedule=schedule, prompt=prompt
)
except Exception as exc:
ctx.ui.append_system(f"Error: {exc}", style="red")
return
ctx.ui.append_system(
f"Scheduled '{name}' [{schedule}] — id {rec.get('cron_id')}. "
"Runs unattended in the background.",
style="green",
)
async def _list(self, ctx: CommandContext, crons) -> None:
# B1: guard SDK call — backend may die after the is_available() check.
try:
rows = await asyncio.to_thread(crons.list_schedules)
except Exception as exc:
ctx.ui.append_system(f"Error: {exc}", style="red")
return
if not rows:
ctx.ui.append_system(
"No scheduled tasks. Add one: /schedule add ...", style="dim"
)
return
table = Table(title="Scheduled Tasks", show_header=True)
table.add_column("ID", style="cyan")
table.add_column("Name", style="magenta")
table.add_column("Schedule", style="green")
table.add_column("Enabled", style="yellow")
table.add_column("Next run (UTC)", style="white")
for r in rows:
meta = r.get("metadata") or {}
table.add_row(
str(r.get("cron_id", ""))[:8],
str(meta.get("name", "")),
str(r.get("schedule", "")),
"yes" if r.get("enabled", True) else "no",
str(r.get("next_run_date", "")),
)
ctx.ui.mount_renderable(table)
_AMBIGUOUS = object() # B2: sentinel returned when multiple crons match a prefix
_BACKEND_ERROR = object() # sentinel returned when list_schedules() raises
async def _resolve(self, crons, prefix: str):
"""Return the unique matching record, _AMBIGUOUS if >1 match, _BACKEND_ERROR on error, or None."""
# B1: guard SDK call — backend may die after is_available() check.
try:
all_rows = await asyncio.to_thread(crons.list_schedules)
except Exception as exc:
# Store the exception text so _resolve_or_report can surface it.
self._last_backend_exc = exc
return self._BACKEND_ERROR
# B2: collect ALL matches; ambiguous prefix → sentinel so callers can warn.
matches = [r for r in all_rows if str(r.get("cron_id", "")).startswith(prefix)]
if len(matches) > 1:
return self._AMBIGUOUS
return matches[0] if matches else None
async def _resolve_or_report(self, ctx: CommandContext, crons, prefix: str):
"""Resolve prefix → record, emit UI error on ambiguity/miss/error, return None on failure."""
match = await self._resolve(crons, prefix)
if match is self._BACKEND_ERROR:
exc = getattr(self, "_last_backend_exc", None)
ctx.ui.append_system(
f"Error: scheduler backend unavailable ({exc})",
style="red",
)
return None
if match is self._AMBIGUOUS:
ctx.ui.append_system(
f"Multiple schedules match '{prefix}' — use a longer id.",
style="yellow",
)
return None
if not match:
ctx.ui.append_system(f"No schedule matching {prefix}.", style="yellow")
return None
return match
async def _remove(self, ctx: CommandContext, crons, prefix: str) -> None:
if not prefix:
ctx.ui.append_system("Usage: /schedule remove <id>", style="yellow")
return
match = await self._resolve_or_report(ctx, crons, prefix)
if match is None:
return
cron_id = str(match.get("cron_id", ""))
try:
await asyncio.to_thread(crons.delete_schedule, cron_id)
except Exception as exc:
ctx.ui.append_system(f"Error: {exc}", style="red")
return
ctx.ui.append_system(f"Removed schedule {cron_id}.", style="green")
async def _run(self, ctx: CommandContext, crons, prefix: str) -> None:
if not prefix:
ctx.ui.append_system("Usage: /schedule run <id>", style="yellow")
return
match = await self._resolve_or_report(ctx, crons, prefix)
if match is None:
return
prompt = (match.get("metadata") or {}).get("prompt", "")
if not str(prompt).strip():
ctx.ui.append_system(
f"Schedule {prefix} has no stored prompt — cannot run it.",
style="yellow",
)
return
try:
rec = await asyncio.to_thread(crons.run_now, prompt)
except Exception as exc:
ctx.ui.append_system(f"Error: {exc}", style="red")
return
# Don't promise a location; the task's own prompt decides where output goes.
ctx.ui.append_system(
f"Fired schedule {prefix} once now (run {rec.get('run_id')}). "
"Any output goes wherever the task's instruction specifies.",
style="green",
)
async def _set_enabled(
self, ctx: CommandContext, crons, prefix: str, enabled: bool
) -> None:
if not prefix:
ctx.ui.append_system("Usage: /schedule pause|resume <id>", style="yellow")
return
match = await self._resolve_or_report(ctx, crons, prefix)
if match is None:
return
cron_id = str(match.get("cron_id", ""))
try:
await asyncio.to_thread(crons.set_enabled, cron_id, enabled)
except Exception as exc:
ctx.ui.append_system(f"Error: {exc}", style="red")
return
ctx.ui.append_system(
f"{'Resumed' if enabled else 'Paused'} schedule {cron_id}.", style="green"
)
# Register schedule command
manager.register(ScheduleCommand())
+46 -51
View File
@@ -5,24 +5,15 @@ from typing import ClassVar
from rich.table import Table
from ...gateway import GraphGateway, GraphTarget
from ..base import Argument, Command, CommandContext
from ..manager import manager
def _graph_gateway(ctx: CommandContext) -> GraphGateway:
if ctx.graph_gateway is None:
raise RuntimeError("Session commands require a graph_gateway")
return ctx.graph_gateway
class CompactCommand(Command):
"""Compact conversation to free context."""
name = "/compact"
description = "Compact conversation to free context"
requires_agent = True
category = "Session"
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...cli.commands import (
@@ -44,12 +35,8 @@ class CompactCommand(Command):
try:
result = await compact_conversation(
graph_gateway=_graph_gateway(ctx),
agent=ctx.agent,
thread_id=ctx.thread_id,
target=GraphTarget(
local_graph=ctx.agent,
workspace_dir=ctx.workspace_dir,
),
input_tokens_hint=ctx.input_tokens_hint,
)
finally:
@@ -83,13 +70,11 @@ class ThreadsCommand(Command):
name = "/threads"
description = "List recent sessions"
category = "Session"
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...sessions import _format_relative_time, short_thread_id
from ...sessions import _format_relative_time, list_threads
gateway = _graph_gateway(ctx)
threads = await gateway.list_threads(
threads = await list_threads(
limit=0,
include_message_count=True,
include_preview=True,
@@ -117,7 +102,7 @@ class ThreadsCommand(Command):
marker = " *" if thread_id_value == ctx.thread_id else ""
row = [
f"{short_thread_id(thread_id_value)}{marker}",
f"{thread_id_value}{marker}",
thread.get("preview", "") or "",
str(thread.get("message_count", 0)),
]
@@ -127,11 +112,6 @@ class ThreadsCommand(Command):
table.add_row(*row)
ctx.ui.mount_renderable(table)
if not is_channel:
ctx.ui.append_system(
" /resume to continue a session "
"/delete <id> to remove /new to start fresh",
)
class ResumeCommand(Command):
@@ -139,7 +119,6 @@ class ResumeCommand(Command):
name = "/resume"
description = "Resume a previous session"
category = "Session"
arguments: ClassVar[list[Argument]] = [
Argument(
name="thread_id",
@@ -150,10 +129,14 @@ class ResumeCommand(Command):
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
gateway = _graph_gateway(ctx)
from ...sessions import (
get_thread_metadata,
list_threads,
)
arg = args[0] if args else ""
if not arg:
threads = await gateway.list_threads(
threads = await list_threads(
limit=0,
include_message_count=True,
include_preview=True,
@@ -177,7 +160,7 @@ class ResumeCommand(Command):
if not resolved:
return
metadata = await gateway.get_thread_metadata(resolved)
metadata = await get_thread_metadata(resolved)
restored_workspace = (metadata or {}).get("workspace_dir", "")
if restored_workspace:
ctx.workspace_dir = restored_workspace
@@ -189,16 +172,21 @@ class ResumeCommand(Command):
await ctx.ui.handle_session_resume(resolved, restored_workspace)
async def _resolve_thread_id(self, prefix: str, ctx: CommandContext) -> str | None:
resolution = await _graph_gateway(ctx).resolve_thread(prefix)
if resolution.thread_id:
return resolution.thread_id
from ...sessions import find_similar_threads, thread_exists
if resolution.matches:
if await thread_exists(prefix):
return prefix
similar = await find_similar_threads(prefix)
if len(similar) == 1:
return similar[0]
if len(similar) > 1:
ctx.ui.append_system(
f"Ambiguous thread ID '{prefix}'. Use a longer prefix.",
style="yellow",
)
for thread in resolution.matches:
for thread in similar:
ctx.ui.append_system(f" - {thread}", style="dim")
return None
@@ -211,10 +199,9 @@ class NewCommand(Command):
name = "/new"
description = "Start a new session"
category = "Session"
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
await ctx.ui.start_new_session()
ctx.ui.start_new_session()
class ClearCommand(Command):
@@ -222,7 +209,6 @@ class ClearCommand(Command):
name = "/clear"
description = "Clear chat history"
category = "Session"
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
ctx.ui.clear_chat()
@@ -233,7 +219,6 @@ class DeleteCommand(Command):
name = "/delete"
description = "Delete a saved session"
category = "Session"
arguments: ClassVar[list[Argument]] = [
Argument(
name="thread_id",
@@ -244,10 +229,16 @@ class DeleteCommand(Command):
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
gateway = _graph_gateway(ctx)
from ...sessions import (
delete_thread,
find_similar_threads,
list_threads,
thread_exists,
)
arg = args[0] if args else ""
if not arg:
threads = await gateway.list_threads(
threads = await list_threads(
limit=0,
include_message_count=True,
include_preview=True,
@@ -267,17 +258,22 @@ class DeleteCommand(Command):
arg = selected
# Resolve thread_id
resolution = await gateway.resolve_thread(arg)
if resolution.matches:
ctx.ui.append_system(
f"Ambiguous thread ID '{arg}'. Use a longer prefix.",
style="yellow",
)
for thread in resolution.matches:
ctx.ui.append_system(f" - {thread}", style="dim")
return
resolved = None
if await thread_exists(arg):
resolved = arg
else:
similar = await find_similar_threads(arg)
if len(similar) == 1:
resolved = similar[0]
elif len(similar) > 1:
ctx.ui.append_system(
f"Ambiguous thread ID '{arg}'. Use a longer prefix.",
style="yellow",
)
for thread in similar:
ctx.ui.append_system(f" - {thread}", style="dim")
return
resolved = resolution.thread_id
if not resolved:
ctx.ui.append_system(f"Session '{arg}' not found.", style="red")
return
@@ -289,7 +285,7 @@ class DeleteCommand(Command):
)
return
deleted = await gateway.delete_thread(resolved)
deleted = await delete_thread(resolved)
if deleted:
ctx.ui.append_system(f"Deleted session {resolved}.", style="green")
else:
@@ -302,7 +298,6 @@ class ExitCommand(Command):
name = "/exit"
alias: ClassVar[list[str]] = ["/quit", "/q"]
description = "Quit EvoScientist"
category = "Session"
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
ctx.ui.force_quit()
+1 -12
View File
@@ -13,7 +13,6 @@ class SkillsCommand(Command):
name = "/skills"
description = "List installed skills"
category = "Skills"
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...cli.agent import _shorten_path
@@ -66,7 +65,6 @@ class InstallSkill(Command):
name: ClassVar[str] = "/install-skill"
description: ClassVar[str] = "Add a skill from path or GitHub"
category: ClassVar[str] = "Skills"
arguments: ClassVar[list[Argument]] = [
Argument(
name="source",
@@ -144,7 +142,6 @@ class InstallSkills(Command):
description: ClassVar[str] = (
"Browse and install EvoSkills (optional: /evoskills <tag>)"
)
category: ClassVar[str] = "Skills"
arguments: ClassVar[list[Argument]] = [
Argument(
name="tag", type=str, description="Tag to filter skills by", required=False
@@ -216,18 +213,11 @@ class InstallSkills(Command):
pre_filter_tag=tag,
)
# ``None`` means user cancelled (Esc / Ctrl-C). An empty list means
# the picker handled a "nothing to do" state (all-installed / no
# tag matches) and already printed its own specific message; the
# outer layer should stay silent rather than claim a cancel.
if selected_sources is None:
if not selected_sources:
if not is_channel:
ctx.ui.append_system("Browse cancelled.", style="dim")
return
if not selected_sources:
return
# Install selected skills
installed_count = 0
for source in selected_sources:
@@ -258,7 +248,6 @@ class UninstallSkill(Command):
name: ClassVar[str] = "/uninstall-skill"
description: ClassVar[str] = "Remove an installed skill"
category: ClassVar[str] = "Skills"
arguments: ClassVar[list[Argument]] = [
Argument(
name="name",
+1 -35
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
import logging
import shlex
from .base import Command, CommandContext, SubCommand
from .base import Command, CommandContext
_logger = logging.getLogger(__name__)
@@ -27,27 +27,6 @@ class CommandManager:
"""Lookup a command by name."""
return self._commands.get(name.lower())
def resolve(self, command_str: str) -> tuple[Command, list[str]] | None:
"""Return ``(command, args)`` for the dispatch of ``command_str``.
Uses the same parsing as :meth:`execute` so callers can inspect
metadata (e.g. call :meth:`Command.needs_agent`) without
re-implementing ``shlex`` quirks.
"""
command_str = command_str.strip()
if not command_str:
return None
try:
parts = shlex.split(command_str)
except ValueError:
parts = command_str.split()
if not parts:
return None
cmd = self.get_command(parts[0])
if cmd is None:
return None
return cmd, parts[1:]
def list_commands(self) -> list[tuple[str, str]]:
"""List all registered command names and descriptions."""
seen = set()
@@ -58,17 +37,6 @@ class CommandManager:
seen.add(cmd)
return results
def get_subcommands(self, command_name: str) -> list[SubCommand]:
"""Return subcommands declared by *command_name*, or empty list."""
cmd = self.get_command(command_name)
if cmd is None:
return []
return cmd.subcommands
def list_subcommands(self, command_name: str) -> list[tuple[str, str]]:
"""Return ``(name, description)`` pairs for completion rendering."""
return [(sc.name, sc.description) for sc in self.get_subcommands(command_name)]
def get_all_commands(self) -> list[Command]:
"""Return all registered command instances."""
seen = set()
@@ -104,14 +72,12 @@ class CommandManager:
if not cmd:
return False
ctx.command_error = None
try:
await cmd.execute(ctx, args)
await ctx.ui.flush()
return True
except Exception as e:
_logger.exception(f"Error executing command {cmd_name}: {e}")
ctx.command_error = str(e)
ctx.ui.append_system(f"Error executing {cmd_name}: {e}", style="red")
await ctx.ui.flush()
return True
-10
View File
@@ -10,11 +10,6 @@ The onboard module is loaded lazily because it pulls in heavy dependencies
from .settings import (
EvoScientistConfig,
MemoryControls,
MemoryObservationTarget,
MemoryObservationWriter,
MemorySkillSynthesisCadence,
MemorySkillSynthesisMode,
apply_config_to_env,
get_config_dir,
get_config_path,
@@ -29,11 +24,6 @@ from .settings import (
__all__ = [
"EvoScientistConfig",
"MemoryControls",
"MemoryObservationTarget",
"MemoryObservationWriter",
"MemorySkillSynthesisCadence",
"MemorySkillSynthesisMode",
"apply_config_to_env",
# settings
"get_config_dir",
+316
View File
@@ -0,0 +1,316 @@
"""Configuration helpers for dedicated image-generation models.
Image generation models are service/tool models, not chat models. Keeping
them in a separate config section prevents image-only models such as
``gpt-image-2`` from being offered in the normal chat model selector.
"""
from __future__ import annotations
import os
import tempfile
from typing import Any
import yaml
from pydantic import BaseModel, Field, field_validator
IMAGE_GENERATION_SECTION = "image_generation"
DEFAULT_IMAGE_GENERATION_TIMEOUT_SECONDS = 120.0
IMAGE_GENERATION_USAGE_NOTES = [
"Normal chat model lists filter out image-only models such as gpt-image-* and dall-e-*.",
"If an image-only model is accidentally saved in LLM settings, it is moved into image_generation.",
"If the frontend sends an image-only model as the chat model, the backend returns IMAGE_MODEL_NOT_CHAT_MODEL.",
]
_IMAGE_MODEL_PREFIXES = (
"gpt-image",
"chatgpt-image",
"dall-e",
"dalle",
)
_IMAGE_MODEL_MARKERS = (
"wanx",
"seedream",
)
class ImageModelEntry(BaseModel):
"""A single dedicated image-generation model."""
id: str = Field(..., description="Model ID sent to the image API")
name: str = Field("", description="Display name or short alias")
provider: str = Field("openai-compatible", description="Provider label")
api_key: str = Field("", description="API key or ${ENV_VAR} reference")
base_url: str = Field("", description="API base URL or ${ENV_VAR} reference")
supports_generation: bool = True
supports_edit: bool = True
default_size: str = "1024x1024"
default_quality: str = "auto"
params: dict[str, Any] = Field(default_factory=dict)
@field_validator("id")
@classmethod
def id_not_empty(cls, value: str) -> str:
value = value.strip()
if not value:
raise ValueError("image model id must not be empty")
return value
def resolved_api_key(self) -> str:
return _resolve_env_ref(self.api_key)
def resolved_base_url(self) -> str:
return _resolve_env_ref(self.base_url)
def display_name(self) -> str:
return self.name or self.id
class ImageGenerationSettings(BaseModel):
"""Dedicated image-generation model settings."""
default_model: str = ""
timeout_seconds: float = Field(
DEFAULT_IMAGE_GENERATION_TIMEOUT_SECONDS,
description="Per-request timeout for image provider HTTP calls",
)
models: list[ImageModelEntry] = Field(default_factory=list)
@field_validator("timeout_seconds")
@classmethod
def timeout_must_be_positive(cls, value: float) -> float:
if value <= 0:
raise ValueError("image generation timeout_seconds must be greater than 0")
return value
def is_image_generation_model(model_ref: str | None) -> bool:
"""Return True when a model ID is known to be image-generation only."""
if not model_ref:
return False
value = str(model_ref).strip().lower()
if "/" in value:
value = value.rsplit("/", 1)[1]
return value.startswith(_IMAGE_MODEL_PREFIXES) or any(
marker in value for marker in _IMAGE_MODEL_MARKERS
)
def load_image_generation_settings(
*,
include_legacy: bool = True,
resolve_env: bool = True,
) -> ImageGenerationSettings:
"""Load dedicated image-generation settings with legacy fallback."""
raw = _load_settings_yaml()
legacy_present = _legacy_config_present(raw)
section = raw.get(IMAGE_GENERATION_SECTION)
if isinstance(section, dict):
source = _resolve_nested_env(section) if resolve_env else section
settings = ImageGenerationSettings.model_validate(source)
else:
settings = ImageGenerationSettings()
legacy = _legacy_model_entry(raw)
if not settings.models:
if include_legacy or legacy_present:
settings.models = [legacy]
elif include_legacy and _legacy_env_overrides_present():
_merge_entry(settings, legacy)
if not settings.default_model:
if settings.models:
settings.default_model = settings.models[0].id
elif include_legacy or legacy_present:
settings.default_model = legacy.id
return settings
def save_image_generation_settings(settings: ImageGenerationSettings) -> ImageGenerationSettings:
"""Save image-generation settings into ``settings.yaml``."""
validated = ImageGenerationSettings.model_validate(settings.model_dump(mode="python"))
raw = _load_settings_yaml()
raw[IMAGE_GENERATION_SECTION] = validated.model_dump(mode="python")
_atomic_write_settings_yaml(raw)
return validated
def merge_image_generation_models(
entries: list[ImageModelEntry],
*,
default_model: str | None = None,
) -> ImageGenerationSettings:
"""Merge image model entries into the dedicated image model list."""
raw = _load_settings_yaml()
settings = load_image_generation_settings(include_legacy=_legacy_config_present(raw))
for entry in entries:
_merge_entry(settings, entry)
if default_model:
settings.default_model = default_model
elif entries and not settings.default_model:
settings.default_model = entries[0].id
return save_image_generation_settings(settings)
def resolve_image_generation_model(model_ref: str | None = None) -> ImageModelEntry:
"""Resolve a model ID/name to a configured image model entry."""
settings = load_image_generation_settings()
requested = (model_ref or settings.default_model or "").strip()
if not requested and settings.models:
requested = settings.models[0].id
for entry in settings.models:
if requested in {entry.id, entry.name}:
return _with_legacy_fallbacks(entry)
if requested:
legacy = _legacy_model_entry(_load_settings_yaml())
if requested == legacy.id:
return legacy
raise ValueError(f"Unknown image generation model: {requested}")
raise ValueError("No image generation model configured")
def list_image_generation_models(*, include_sensitive: bool = False) -> dict[str, Any]:
"""Return image-generation model settings for API/tool display."""
settings = load_image_generation_settings()
models = []
for entry in settings.models:
item = entry.model_dump(mode="python")
item["name"] = entry.display_name()
if include_sensitive:
item["api_key"] = entry.resolved_api_key()
item["base_url"] = entry.resolved_base_url()
else:
item["api_key"] = _mask_secret(entry.resolved_api_key())
item["base_url"] = entry.resolved_base_url()
models.append(item)
return {
"default_model": settings.default_model,
"timeout_seconds": settings.timeout_seconds,
"models": models,
"usage_notes": IMAGE_GENERATION_USAGE_NOTES,
}
def _merge_entry(settings: ImageGenerationSettings, entry: ImageModelEntry) -> None:
for idx, existing in enumerate(settings.models):
if existing.id == entry.id:
data = existing.model_dump(mode="python")
update = entry.model_dump(mode="python")
for key, value in update.items():
if value not in ("", None, {}, []):
data[key] = value
settings.models[idx] = ImageModelEntry.model_validate(data)
return
settings.models.append(entry)
def _with_legacy_fallbacks(entry: ImageModelEntry) -> ImageModelEntry:
legacy = _legacy_model_entry(_load_settings_yaml())
data = entry.model_dump(mode="python")
if not data.get("api_key") or _is_masked_secret(str(data.get("api_key") or "")):
data["api_key"] = legacy.api_key
if not data.get("base_url"):
data["base_url"] = legacy.base_url
return ImageModelEntry.model_validate(data)
def _legacy_env_overrides_present() -> bool:
return any(os.environ.get(key) for key in ("IMAGE_GEN_MODEL", "IMAGE_GEN_API_KEY", "IMAGE_GEN_BASE_URL"))
def _legacy_config_present(raw: dict[str, Any]) -> bool:
return _legacy_env_overrides_present() or any(
str(raw.get(key) or "").strip()
for key in ("image_gen_model", "image_gen_api_key", "image_gen_base_url")
)
def _legacy_model_entry(raw: dict[str, Any]) -> ImageModelEntry:
model = (
os.environ.get("IMAGE_GEN_MODEL")
or str(raw.get("image_gen_model") or "").strip()
or "dall-e-3"
)
api_key = (
os.environ.get("IMAGE_GEN_API_KEY")
or str(raw.get("image_gen_api_key") or "").strip()
or os.environ.get("OPENAI_API_KEY", "")
)
base_url = (
os.environ.get("IMAGE_GEN_BASE_URL")
or str(raw.get("image_gen_base_url") or "").strip()
or os.environ.get("OPENAI_BASE_URL", "")
)
return ImageModelEntry(
id=model,
name=model,
api_key=api_key,
base_url=base_url,
provider="openai-compatible",
)
def _load_settings_yaml() -> dict[str, Any]:
from EvoScientist.config.settings import get_config_path
path = get_config_path()
if not path.exists():
return {}
try:
with path.open(encoding="utf-8") as fh:
data = yaml.safe_load(fh) or {}
except Exception:
return {}
return data if isinstance(data, dict) else {}
def _atomic_write_settings_yaml(data: dict[str, Any]) -> None:
from EvoScientist.config.settings import get_config_path
path = get_config_path()
path.parent.mkdir(parents=True, exist_ok=True)
fd, tmp_path = tempfile.mkstemp(
dir=str(path.parent),
prefix=".settings_",
suffix=".yaml.tmp",
text=True,
)
try:
with os.fdopen(fd, "w", encoding="utf-8") as fh:
yaml.safe_dump(data, fh, default_flow_style=False, sort_keys=False)
fh.flush()
os.fsync(fh.fileno())
os.replace(tmp_path, path)
finally:
if os.path.exists(tmp_path):
os.unlink(tmp_path)
def _resolve_nested_env(value: Any) -> Any:
if isinstance(value, dict):
return {k: _resolve_nested_env(v) for k, v in value.items()}
if isinstance(value, list):
return [_resolve_nested_env(v) for v in value]
if isinstance(value, str):
return _resolve_env_ref(value)
return value
def _resolve_env_ref(value: str) -> str:
if isinstance(value, str) and value.startswith("${") and value.endswith("}"):
return os.environ.get(value[2:-1], "")
return value or ""
def _mask_secret(value: str) -> str:
if not value:
return ""
if len(value) <= 8:
return "***"
return f"{value[:3]}...{value[-4:]}"
def _is_masked_secret(value: str) -> bool:
return bool(value and (value == "***" or value == "********" or "..." in value))
+947
View File
@@ -0,0 +1,947 @@
"""Structured LLM provider/model configuration.
This module defines the YAML-driven configuration for providers, models,
and their parameters. It supports:
- Multiple providers with api_key, base_url, and protocol
- Multiple models per provider with alias, capabilities, and params
- Three-level parameter inheritance: defaults ← model ← runtime
- Environment variable references in api_key fields
"""
from __future__ import annotations
import logging
import os
import random
import re
import tempfile
import threading
import time
from pathlib import Path
from typing import Any, Literal
logger = logging.getLogger(__name__)
import yaml
from pydantic import BaseModel, Field, field_validator
# ─── Environment variable reference pattern ──────────────────────────────
_ENV_REF_RE = re.compile(r"^\$\{(\w+)\}$")
_FLAT_LLM_CONFIG_KEYS = {
"provider_routes",
"provider",
"model",
"reasoning_effort",
"anthropic_api_key",
"anthropic_base_url",
"openai_api_key",
"nvidia_api_key",
"google_api_key",
"minimax_api_key",
"siliconflow_api_key",
"openrouter_api_key",
"deepseek_api_key",
"zhipu_api_key",
"volcengine_api_key",
"dashscope_api_key",
"moonshot_api_key",
"kimi_api_key",
"custom_openai_api_key",
"custom_openai_base_url",
"custom_anthropic_api_key",
"custom_anthropic_base_url",
"ollama_base_url",
}
def _resolve_env_ref(value: str) -> str:
"""Resolve ${ENV_VAR} references in string values.
If the value matches the pattern ${ENV_VAR}, return the environment
variable's value. Otherwise return the string as-is.
"""
m = _ENV_REF_RE.match(value.strip())
if m:
return os.environ.get(m.group(1), "")
return value
# ─── Weighted round-robin load balancer ───────────────────────────────────
# Per-provider counter for round-robin position
_round_robin_counters: dict[str, int] = {}
# ─── Endpoint call statistics ──────────────────────────────────────────
class EndpointStats:
"""Thread-safe in-memory endpoint call and token usage statistics.
Tracks per-endpoint:
- Call counts (from resolve_model)
- Token usage (input_tokens, output_tokens)
- Last call timestamp
- Per-model call distribution
Stats are kept in memory. Call ``snapshot()`` for a point-in-time
copy, ``reset()`` to clear counters, or ``summary()`` for a
human-readable report.
"""
def __init__(self) -> None:
self._lock = threading.Lock()
# key = (provider_name, endpoint_name)
# value = dict with keys: calls, last_call_ts, models, input_tokens, output_tokens
self._stats: dict[tuple[str, str], dict[str, Any]] = {}
# Track last endpoint selected per model for token attribution
# key = model_id, value = (provider, endpoint_name)
self._last_endpoint: dict[str, tuple[str, str]] = {}
# Track most recently used endpoint (for fallback when model_id is empty)
self._last_recorded: tuple[str, str] | None = None
# Stack of pending endpoint attributions for sequential matching.
# Each call to record() pushes, each record_tokens_for_model() pops.
# This ensures tokens are attributed to the correct endpoint in tool-call loops
# where resolve_model() is called multiple times before usage_metadata arrives.
self._attribution_stack: list[tuple[str, str]] = []
def _ensure_entry(self, key: tuple[str, str]) -> dict[str, Any]:
"""Get or create a stats entry for the given key."""
return self._stats.setdefault(key, {
"calls": 0,
"last_call_ts": 0.0,
"models": {},
"input_tokens": 0,
"output_tokens": 0,
})
def record(self, provider: str, endpoint: str, model_id: str) -> None:
"""Record one endpoint selection (called from resolve_model)."""
with self._lock:
entry = self._ensure_entry((provider, endpoint))
entry["calls"] += 1
entry["last_call_ts"] = time.time()
entry["models"][model_id] = entry["models"].get(model_id, 0) + 1
# Remember last endpoint for this model (for token attribution)
self._last_endpoint[model_id] = (provider, endpoint)
# Track most recently used endpoint overall
self._last_recorded = (provider, endpoint)
# Push to attribution stack for sequential token matching
self._attribution_stack.append((provider, endpoint))
def record_tokens(
self,
provider: str,
endpoint: str,
input_tokens: int = 0,
output_tokens: int = 0,
) -> None:
"""Record token usage for an endpoint.
Can be called after an API response is received to accumulate
token counts. The endpoint entry is created if it doesn't exist
(e.g. for single-endpoint providers that aren't tracked by
``record()``).
"""
if not input_tokens and not output_tokens:
return
with self._lock:
entry = self._ensure_entry((provider, endpoint))
entry["input_tokens"] += input_tokens
entry["output_tokens"] += output_tokens
def record_tokens_for_model(
self,
model_id: str,
input_tokens: int = 0,
output_tokens: int = 0,
) -> None:
"""Record token usage, auto-routing to the correct endpoint.
Resolution order:
1. Attribution stack — pops the oldest pending endpoint (FIFO matching
with record() calls, handles tool-call loops correctly).
2. ``_last_endpoint[model_id]`` — direct model→endpoint mapping.
3. Scan ``_stats`` for any endpoint that has this model registered.
4. ``_last_recorded`` — fallback for empty model_id.
5. ``("unknown", "unknown")`` — last resort (debug level).
"""
if not input_tokens and not output_tokens:
return
with self._lock:
provider, endpoint = None, None
# Strategy 1: Pop from attribution stack (matches record() calls in order)
if self._attribution_stack:
provider, endpoint = self._attribution_stack.pop(0)
# Strategy 2: Direct lookup by model_id
if provider is None and model_id:
p, e = self._last_endpoint.get(model_id, (None, None))
if p is not None:
provider, endpoint = p, e
# Strategy 3: Scan _stats for this model
if provider is None and model_id:
for (p, e), data in self._stats.items():
if model_id in data.get("models", {}):
provider, endpoint = p, e
self._last_endpoint[model_id] = (p, e)
break
# Strategy 4: Use most recently recorded endpoint
if provider is None and self._last_recorded:
provider, endpoint = self._last_recorded
# Strategy 5: Unknown — no resolve_model() was called beforehand
if provider is None:
# When model_id is empty and all state is empty, this is a
# known race condition (TOCTOU between events.py guard and
# this method's lock). Silently skip — no useful attribution
# is possible and logging it just creates noise.
if not model_id and not self._last_recorded:
return
provider, endpoint = "unknown", "unknown"
# Common when LLM calls bypass resolve_model() (e.g. LangChain
# internal bindings, sub-agents). Debug level to avoid log spam.
logger.debug(
"EndpointStats: model %r not found, stack=%d _last_endpoint=%s "
"_last_recorded=%s",
model_id,
len(self._attribution_stack),
list(self._last_endpoint.keys()),
self._last_recorded,
)
entry = self._ensure_entry((provider, endpoint))
entry["input_tokens"] += input_tokens
entry["output_tokens"] += output_tokens
def snapshot(self) -> dict[str, Any]:
"""Return a deep copy of current stats."""
with self._lock:
import copy
return copy.deepcopy(self._stats)
async def restore_today_from_db(self) -> int:
"""Restore today's endpoint stats from endpoint_usage_daily.
Called once at gateway startup so that the in-memory EndpointStats
reflects today's accumulated usage after a process restart.
Returns the number of rows restored.
"""
try:
from EvoScientist.runtime_integrations import current_date, get_app_connection
db = await get_app_connection()
today = current_date()
rows = await db.execute_fetchall(
"""
SELECT provider, endpoint, model, calls,
input_tokens, output_tokens
FROM endpoint_usage_daily
WHERE date = $1
""",
(today,),
)
with self._lock:
for r in rows:
key = (r["provider"], r["endpoint"])
entry = self._ensure_entry(key)
entry["calls"] += r["calls"]
entry["input_tokens"] += r["input_tokens"]
entry["output_tokens"] += r["output_tokens"]
if r["model"]:
entry["models"][r["model"]] = (
entry["models"].get(r["model"], 0) + r["calls"]
)
self._last_endpoint[r["model"]] = key
restored = len(rows)
if restored:
logger.info(
"EndpointStats: restored %d endpoint entries from DB for today",
restored,
)
return restored
except Exception:
logger.warning("EndpointStats: failed to restore from DB", exc_info=True)
return 0
def reset(self) -> None:
"""Clear all counters."""
with self._lock:
self._stats.clear()
@staticmethod
def _fmt_tokens(n: int) -> str:
"""Format token count compactly."""
if n >= 1_000_000:
return f"{n / 1_000_000:.1f}M"
if n >= 1_000:
return f"{n / 1_000:.1f}K"
return str(n)
def summary(self) -> str:
"""Human-readable summary for CLI / logging."""
snap = self.snapshot()
if not snap:
return "No endpoint calls recorded."
total_calls = sum(d["calls"] for d in snap.values())
total_input = sum(d["input_tokens"] for d in snap.values())
total_output = sum(d["output_tokens"] for d in snap.values())
if total_calls == 0:
return "No endpoint calls recorded."
lines: list[str] = []
for (provider, endpoint), data in sorted(snap.items()):
calls = data["calls"]
pct = calls / total_calls * 100
elapsed = time.time() - data["last_call_ts"] if data["last_call_ts"] else 0
if elapsed < 60:
ago = f"{elapsed:.0f}s ago"
elif elapsed < 3600:
ago = f"{elapsed / 60:.0f}m ago"
else:
ago = f"{elapsed / 3600:.1f}h ago"
bar = "█" * int(pct / 5) + "░" * (20 - int(pct / 5))
inp = self._fmt_tokens(data["input_tokens"])
out = self._fmt_tokens(data["output_tokens"])
models = ", ".join(f"{m}({c})" for m, c in sorted(data["models"].items()))
token_str = ""
if data["input_tokens"] or data["output_tokens"]:
token_str = f" in={inp} out={out}"
lines.append(
f" {provider}/{endpoint} {bar} {calls} ({pct:.0f}%) last {ago}{token_str}\n"
f" models: {models}"
)
header = f"Endpoint usage — {total_calls} calls"
if total_input or total_output:
header = (
f"Endpoint usage — {total_calls} calls, "
f"{self._fmt_tokens(total_input)} in / {self._fmt_tokens(total_output)} out tokens"
)
return header + ":\n" + "\n".join(lines)
# Global singleton
_endpoint_stats = EndpointStats()
def get_endpoint_stats() -> EndpointStats:
"""Return the global endpoint call statistics instance."""
return _endpoint_stats
def _select_endpoint(
endpoints: list,
provider_name: str,
preferred_name: str = "",
) -> tuple:
"""Select an endpoint using weighted round-robin or explicit pinning.
Args:
endpoints: List of EndpointConfig objects with valid credentials.
provider_name: Provider name for round-robin counter.
preferred_name: If set, try to find an endpoint with this name first.
Returns:
The selected EndpointConfig.
"""
if not endpoints:
raise ValueError(f"No available endpoints for provider '{provider_name}'")
# 1. If model pins to a specific endpoint, find it
if preferred_name:
for ep in endpoints:
if ep.name == preferred_name and ep.resolved_api_key():
return ep
# 2. Filter to endpoints with valid credentials
available = [ep for ep in endpoints if ep.resolved_api_key()]
if not available:
raise ValueError(f"No endpoints with valid credentials for provider '{provider_name}'")
if len(available) == 1:
return available[0]
# 3. Weighted round-robin selection
total_weight = sum(ep.weight for ep in available)
if total_weight <= 0:
return available[0]
key = provider_name
pos = _round_robin_counters.get(key, 0)
_round_robin_counters[key] = (pos + 1) % total_weight
# Walk through endpoints by cumulative weight
cumulative = 0
for ep in available:
cumulative += ep.weight
if pos < cumulative:
return ep
return available[-1]
# ═══════════════════════════════════════════════════════════════════════════
# Pydantic models for the new structured config
# ═══════════════════════════════════════════════════════════════════════════
class EndpointConfig(BaseModel):
"""A single API endpoint within a provider — its own api_key, base_url, and optional weights."""
name: str = Field("", description="Endpoint name for reference (e.g. 'coding', 'general')")
api_key: str = Field("", description="API key (literal or ${ENV_VAR} reference)")
base_url: str = Field("", description="API base URL override (empty = provider default)")
weight: int = Field(1, description="Load-balancing weight (higher = more traffic)")
extra_body: dict[str, Any] = Field(
default_factory=dict,
description="Extra JSON body fields for requests through this endpoint",
)
default_headers: dict[str, str] = Field(
default_factory=dict,
description="Custom HTTP headers for requests through this endpoint",
)
params: dict[str, Any] = Field(
default_factory=dict,
description="Endpoint-level model/client parameter overrides",
)
def resolved_api_key(self) -> str:
"""Return the API key with environment variable references resolved."""
return _resolve_env_ref(self.api_key) if self.api_key else ""
def resolved_base_url(self) -> str:
"""Return the base URL with environment variable references resolved."""
return _resolve_env_ref(self.base_url) if self.base_url else ""
class AccessConfig(BaseModel):
"""Plan/role access constraints for providers and models.
Empty lists mean unrestricted. Gateway treats the admin role as allowed.
"""
allowed_plans: list[str] = Field(default_factory=list, description="Allowed subscription plans")
allowed_roles: list[str] = Field(default_factory=list, description="Allowed user roles")
class ModelEntry(BaseModel):
"""A single model definition within a provider."""
id: str = Field(..., description="Full model ID sent to the API")
alias: str = Field("", description="Short alias for easy reference")
tier: str = Field("", description="Billing tier controlled by the server config")
currency: str = Field("CNY", description="Settlement currency for model pricing")
max_tokens: int = Field(4096, description="Maximum output tokens")
supports_vision: bool = Field(False, description="Whether the model supports image inputs")
supports_reasoning: bool = Field(False, description="Whether the model supports reasoning/thinking")
endpoint: str = Field("", description="Preferred endpoint name (empty = auto-select via load balancing)")
params: dict[str, Any] = Field(
default_factory=dict,
description="Model-level parameter overrides (temperature, thinking, reasoning, etc.)",
)
pricing: dict[str, float | str | None] = Field(
default_factory=dict,
description="Optional billing price: input_per_million, output_per_million, cached_input_per_million.",
)
permission_mode: Literal["inherit", "custom"] = Field(
"inherit",
description="Model access mode: inherit provider permissions or use model-level access.",
)
access: AccessConfig = Field(
default_factory=AccessConfig,
description="Model-level access constraints used when permission_mode is custom.",
)
@field_validator("id")
@classmethod
def id_not_empty(cls, v: str) -> str:
if not v.strip():
raise ValueError("model id must not be empty")
return v.strip()
class ProviderConfig(BaseModel):
"""Configuration for a single LLM provider.
Supports two modes:
1. Single endpoint (legacy): set api_key / base_url directly at provider level.
2. Multi-endpoint (new): define an 'endpoints' list, each with its own
api_key, base_url, weight. Models can pin to a specific endpoint or
let the system auto-select via weighted round-robin for load balancing.
"""
api_key: str = Field("", description="API key (literal or ${ENV_VAR} reference)")
base_url: str = Field("", description="API base URL override (empty = provider default)")
protocol: Literal["openai", "anthropic", "google-genai", "ollama", "openrouter"] = Field(
"openai",
description="LLM protocol to use for this provider",
)
extra_body: dict[str, Any] = Field(
default_factory=dict,
description="Extra JSON body fields for all requests to this provider",
)
default_headers: dict[str, str] = Field(
default_factory=dict,
description="Custom HTTP headers for all requests to this provider",
)
params: dict[str, Any] = Field(
default_factory=dict,
description="Provider-level model/client parameter defaults",
)
access: AccessConfig = Field(
default_factory=AccessConfig,
description="Provider-level access constraints.",
)
models: list[ModelEntry] = Field(
default_factory=list,
description="Models available under this provider",
)
endpoints: list[EndpointConfig] = Field(
default_factory=list,
description="Multiple API endpoints for load balancing (overrides top-level api_key/base_url when set)",
)
def resolved_api_key(self) -> str:
"""Return the API key with environment variable references resolved.
For multi-endpoint providers, returns the first available resolved key.
"""
# Multi-endpoint mode: return first non-empty key
if self.endpoints:
for ep in self.endpoints:
key = ep.resolved_api_key()
if key:
return key
return ""
# Single-endpoint mode (legacy)
return _resolve_env_ref(self.api_key) if self.api_key else ""
def get_resolved_endpoints(self) -> list[EndpointConfig]:
"""Return only endpoints that have valid credentials."""
return [ep for ep in self.endpoints if ep.resolved_api_key() or (not ep.api_key and self.resolved_api_key())]
def has_credentials(self) -> bool:
"""Check if this provider has any usable credentials."""
if self.endpoints:
return any(ep.resolved_api_key() for ep in self.endpoints)
return bool(self.resolved_api_key())
class ModelDefaults(BaseModel):
"""Global default parameters for all models."""
temperature: float | None = Field(None, description="Sampling temperature")
max_tokens: int = Field(4096, description="Default max output tokens")
reasoning_effort: str | None = Field("high", description="Reasoning effort level (low/medium/high/xhigh)")
stream_usage: bool = Field(True, description="Whether to enable streaming token usage stats")
class StructuredConfig(BaseModel):
"""Top-level structured configuration for EvoScientist LLM settings."""
default_model: str = Field("", description="Default model ID or alias")
providers: dict[str, ProviderConfig] = Field(
default_factory=dict,
description="Provider definitions keyed by name",
)
model_defaults: ModelDefaults = Field(
default_factory=ModelDefaults,
description="Global default model parameters",
)
# ═══════════════════════════════════════════════════════════════════════════
# Config loading with layered discovery
# ═══════════════════════════════════════════════════════════════════════════
def _get_global_settings_path() -> Path:
"""Get global settings.yaml path via get_config_path()."""
from EvoScientist.config.settings import get_config_path
return get_config_path()
def _load_yaml_file(path: Path) -> dict:
"""Load a YAML file, returning empty dict on failure."""
if not path.is_file():
return {}
try:
with open(path) as f:
data = yaml.safe_load(f)
return data if isinstance(data, dict) else {}
except Exception:
return {}
def _detect_new_config_format(data: dict) -> bool:
"""Detect whether a YAML dict uses the new structured format.
The new format is identified by the presence of a 'providers' top-level key
that is a dict (not a flat string value).
"""
providers = data.get("providers")
return isinstance(providers, dict) and len(providers) > 0
def _deep_merge(base: dict, override: dict) -> dict:
"""Deep merge override into base. Override values take precedence."""
result = base.copy()
for key, value in override.items():
if key in result and isinstance(result[key], dict) and isinstance(value, dict):
result[key] = _deep_merge(result[key], value)
elif key in result and isinstance(result[key], list) and isinstance(value, list):
# For lists (like models), override completely replaces
result[key] = value
else:
result[key] = value
return result
def load_structured_config(
cli_overrides: dict[str, Any] | None = None,
) -> StructuredConfig:
"""Load structured config from ``settings.yaml`` (single source of truth).
Returns code defaults if the file doesn't exist or is invalid.
Args:
cli_overrides: Optional CLI argument overrides.
Returns:
StructuredConfig instance.
"""
settings_path = _get_global_settings_path()
settings_data = _load_yaml_file(settings_path)
config = StructuredConfig()
if _detect_new_config_format(settings_data):
config = StructuredConfig(**settings_data)
_remove_image_generation_models(config)
# Apply CLI overrides
if cli_overrides:
if "default_model" in cli_overrides and cli_overrides["default_model"]:
config.default_model = cli_overrides["default_model"]
return config
def _remove_image_generation_models(config: StructuredConfig) -> list[tuple[str, ProviderConfig, ModelEntry]]:
"""Remove image-only models from chat LLM providers and return them."""
from EvoScientist.config.image_models import is_image_generation_model
removed: list[tuple[str, ProviderConfig, ModelEntry]] = []
for prov_name, provider in config.providers.items():
chat_models: list[ModelEntry] = []
for model in provider.models:
if is_image_generation_model(model.id) or is_image_generation_model(model.alias):
removed.append((prov_name, provider, model))
else:
chat_models.append(model)
provider.models = chat_models
if config.default_model and is_image_generation_model(config.default_model):
config.default_model = ""
return removed
def move_image_models_to_dedicated_config(config: StructuredConfig) -> StructuredConfig:
"""Move image-only model entries out of chat LLM config."""
removed = _remove_image_generation_models(config)
if not removed:
return config
from EvoScientist.config.image_models import ImageModelEntry, merge_image_generation_models
entries: list[ImageModelEntry] = []
for prov_name, provider, model in removed:
api_key = provider.api_key
base_url = provider.base_url
if model.endpoint:
for endpoint in provider.endpoints:
if endpoint.name == model.endpoint:
api_key = endpoint.api_key or api_key
base_url = endpoint.base_url or base_url
break
entries.append(
ImageModelEntry(
id=model.id,
name=model.alias or model.id,
provider=prov_name,
api_key=api_key,
base_url=base_url,
supports_generation=True,
supports_edit=True,
default_size=str(model.params.get("size") or "1024x1024"),
default_quality=str(model.params.get("quality") or "auto"),
params=model.params,
)
)
merge_image_generation_models(entries, default_model=entries[0].id if entries else None)
return config
def save_structured_config(config: StructuredConfig | dict[str, Any]) -> StructuredConfig:
"""Validate and atomically persist the global structured config."""
validated = config if isinstance(config, StructuredConfig) else StructuredConfig.model_validate(config)
validated = move_image_models_to_dedicated_config(validated)
path = _get_global_settings_path()
path.parent.mkdir(parents=True, exist_ok=True)
existing = _load_yaml_file(path)
output = existing if isinstance(existing, dict) else {}
for key in _FLAT_LLM_CONFIG_KEYS:
output.pop(key, None)
output.update(validated.model_dump(mode="python"))
yaml_text = yaml.safe_dump(
output,
allow_unicode=False,
default_flow_style=False,
sort_keys=False,
)
fd, tmp_path = tempfile.mkstemp(
dir=str(path.parent),
prefix=".llm_config_",
suffix=".yaml.tmp",
text=True,
)
try:
with os.fdopen(fd, "w", encoding="utf-8") as f:
f.write(yaml_text)
f.flush()
os.fsync(f.fileno())
os.replace(tmp_path, path)
finally:
if os.path.exists(tmp_path):
os.unlink(tmp_path)
return validated
# ═══════════════════════════════════════════════════════════════════════════
# Model registry: lookup and resolution
# ═══════════════════════════════════════════════════════════════════════════
class ResolvedModel(BaseModel):
"""Fully resolved model information ready for get_chat_model()."""
provider_name: str = Field(..., description="Provider name (e.g. 'anthropic', 'deepseek')")
model_id: str = Field(..., description="Full model ID for the API call")
protocol: str = Field(..., description="Protocol to use (openai/anthropic/google-genai/ollama)")
api_key: str = Field("", description="Resolved API key")
base_url: str = Field("", description="Resolved base URL")
endpoint_name: str = Field("", description="Name of the selected endpoint (empty for single-endpoint providers)")
params: dict[str, Any] = Field(default_factory=dict, description="Merged parameters")
supports_vision: bool = False
supports_reasoning: bool = False
max_tokens: int = 4096
def _build_alias_index(config: StructuredConfig) -> dict[str, tuple[str, str]]:
"""Build alias → (provider_name, model_id) index.
Also indexes model IDs directly so both alias and full ID can be looked up.
"""
index: dict[str, tuple[str, str]] = {}
for prov_name, prov in config.providers.items():
for model in prov.models:
# Index by alias (if set)
if model.alias:
index[model.alias] = (prov_name, model.id)
# Index by full model ID
index[model.id] = (prov_name, model.id)
return index
def resolve_model(
model_ref: str,
config: StructuredConfig | None = None,
runtime_params: dict[str, Any] | None = None,
) -> ResolvedModel:
"""Resolve a model reference to a fully configured ResolvedModel.
Args:
model_ref: Model ID, alias, or provider-prefixed name (e.g. "anthropic/claude-sonnet-4-6").
config: StructuredConfig to use (loaded automatically if None).
runtime_params: Additional runtime parameter overrides.
Returns:
ResolvedModel with all fields populated.
Raises:
ValueError: If the model cannot be resolved.
"""
if config is None:
config = load_structured_config()
runtime_params = runtime_params or {}
# Handle provider-prefixed references: "anthropic/claude-sonnet-4-6"
explicit_provider = None
if "/" in model_ref:
parts = model_ref.split("/", 1)
explicit_provider = parts[0]
model_ref = parts[1]
# Try alias/ID lookup from structured configuration.
alias_index = _build_alias_index(config)
if model_ref in alias_index:
prov_name, model_id = alias_index[model_ref]
if explicit_provider and prov_name != explicit_provider:
prov_name = explicit_provider
model_id = model_ref
elif explicit_provider:
# Not in registry but user specified provider — use as-is
prov_name = explicit_provider
model_id = model_ref
else:
raise ValueError(
f"Model '{model_ref}' is not declared in structured LLM config. "
"Add it under providers.*.models or call it as 'provider/model'."
)
# Get provider config
provider = config.providers.get(prov_name)
if not provider:
raise ValueError(
f"Provider '{prov_name}' is not declared in structured LLM config."
)
# Find the model entry (if registered)
model_entry = None
for m in provider.models:
if m.id == model_id or m.alias == model_ref:
model_entry = m
break
# Three-level parameter merge: defaults ← model ← runtime
params: dict[str, Any] = {}
defaults = config.model_defaults
# Level 1: global defaults
if defaults.temperature is not None:
params["temperature"] = defaults.temperature
params["max_tokens"] = defaults.max_tokens
params["stream_usage"] = defaults.stream_usage
# Level 2: provider-level params
params.update(provider.params)
# Level 3: model-level params
if model_entry:
params.update(model_entry.params)
params["max_tokens"] = model_entry.max_tokens
# Override defaults with model-level values
if "temperature" not in model_entry.params and defaults.temperature is not None:
params["temperature"] = defaults.temperature
# ── Resolve credentials from provider or endpoint ────────────
resolved_api_key = provider.resolved_api_key()
resolved_base_url = _resolve_env_ref(provider.base_url) if provider.base_url else ""
resolved_extra_body = dict(provider.extra_body) if provider.extra_body else {}
resolved_headers = dict(provider.default_headers) if provider.default_headers else {}
selected_endpoint_name = ""
if provider.endpoints:
# Multi-endpoint mode: select endpoint via load balancing
preferred_ep = model_entry.endpoint if model_entry else ""
available_endpoints = [ep for ep in provider.endpoints if ep.resolved_api_key()]
if available_endpoints:
selected = _select_endpoint(available_endpoints, prov_name, preferred_ep)
resolved_api_key = selected.resolved_api_key()
resolved_base_url = selected.resolved_base_url() or resolved_base_url
selected_endpoint_name = selected.name
# Merge endpoint-level extra_body and headers (provider-level first, endpoint overrides)
if selected.extra_body:
resolved_extra_body.update(selected.extra_body)
if selected.default_headers:
resolved_headers.update(selected.default_headers)
if selected.params:
params.update(selected.params)
# Record endpoint call statistics
_endpoint_stats.record(prov_name, selected.name or "default", model_id)
logger.debug(
"endpoint selected: provider=%s endpoint=%s model=%s",
prov_name, selected.name or "default", model_id,
)
else:
# Single-endpoint provider — still record for token attribution
# so that usage_stats events carry correct endpoint info
_endpoint_stats.record(prov_name, "default", model_id)
# Level 4: runtime params
params.update(runtime_params)
# Store merged extra_body/headers in params for downstream consumption
if resolved_extra_body:
params["_extra_body"] = resolved_extra_body
if resolved_headers:
params["_default_headers"] = resolved_headers
return ResolvedModel(
provider_name=prov_name,
model_id=model_id,
protocol=provider.protocol,
api_key=resolved_api_key,
base_url=resolved_base_url,
endpoint_name=selected_endpoint_name,
params=params,
supports_vision=model_entry.supports_vision if model_entry else False,
supports_reasoning=model_entry.supports_reasoning if model_entry else False,
max_tokens=model_entry.max_tokens if model_entry else defaults.max_tokens,
)
def get_default_model(config: StructuredConfig | None = None) -> str:
"""Get the default model reference from config."""
if config is None:
config = load_structured_config()
return config.default_model
def list_available_models(config: StructuredConfig | None = None) -> list[dict[str, Any]]:
"""List all available models from configured providers.
Returns a list of dicts with keys: id, alias, provider, protocol,
max_tokens, supports_vision, supports_reasoning.
"""
if config is None:
config = load_structured_config()
models = []
for prov_name, provider in config.providers.items():
# Only include models from providers that have credentials
# (either via top-level api_key or via at least one endpoint)
has_creds = provider.has_credentials() or (
prov_name == "ollama" and bool(provider.base_url)
)
if not has_creds:
continue
for model in provider.models:
models.append({
"id": model.id,
"alias": model.alias,
"provider": prov_name,
"protocol": provider.protocol,
"max_tokens": model.max_tokens,
"supports_vision": model.supports_vision,
"supports_reasoning": model.supports_reasoning,
"provider_access": provider.access.model_dump(mode="python"),
"permission_mode": model.permission_mode,
"model_access": model.access.model_dump(mode="python"),
})
return models
File diff suppressed because it is too large Load Diff
-34
View File
@@ -1,34 +0,0 @@
"""Onboarding package.
The wizard's only package-level public entry point is :func:`run_onboard`.
Everything else lives in submodules — import directly from them:
- :mod:`EvoScientist.config.onboard.wizard` — orchestrator, ``run_onboard``,
``STEPS``, ``render_progress``
- :mod:`EvoScientist.config.onboard.steps` — per-step functions
- :mod:`EvoScientist.config.onboard.channels` — channel selection + setup
- :mod:`EvoScientist.config.onboard.helpers` — API-key prompt, ccproxy,
npx/node, LaTeX, iMessage helpers
- :mod:`EvoScientist.config.onboard.style` — Rich styles + ``_checkbox_ask``
- :mod:`EvoScientist.config.onboard.validators` — input validators
- :mod:`EvoScientist.config.onboard.prompter` — ``NonInteractivePrompter``
(CLI-answer container) + ``select_navigation_active`` / ``GoBack`` for
keyboard nav
- :mod:`EvoScientist.config.onboard.constants` — canonical valid-value sets
This module used to re-export every symbol from every submodule for
backward compat during the initial refactor; those re-exports have since
been removed to keep the public surface narrow. New code should always
import from the submodule that owns the symbol; test code should use
``patch("EvoScientist.config.onboard.<submodule>.<name>")`` paths.
"""
from __future__ import annotations
# Sole package-level public entry. ``EvoScientist.config`` re-exports this
# (via lazy ``__getattr__``) so ``from EvoScientist.config import
# run_onboard`` keeps working — that import path is used by the CLI and is
# the only documented external API.
from .wizard import run_onboard
__all__ = ["run_onboard"]
-958
View File
@@ -1,958 +0,0 @@
"""Channel selection + per-channel configuration.
`_step_channels` is the big one — over 700 lines that walk the user through
selecting which messaging channels to enable and collecting credentials for
each.
"""
from __future__ import annotations
import questionary
from questionary import Choice
from ..settings import EvoScientistConfig
from .helpers import (
_setup_imessage,
)
from .style import (
QMARK,
WIZARD_STYLE,
console,
)
def _step_channels(config: EvoScientistConfig) -> dict[str, object]:
"""Step: Select channels to enable on startup.
Presents a multi-select list of supported channels.
For each selected channel, prompts for required credentials
and validates them via the channel's probe function.
Args:
config: Current configuration.
Returns:
Dict mapping config field names to their new values.
Empty dict when the user skips or selects nothing.
"""
# Currently enabled channels
_currently_enabled = {
t.strip()
for t in (getattr(config, "channel_enabled", "") or "").split(",")
if t.strip()
}
# Legacy iMessage compat
if (
getattr(config, "imessage_enabled", False)
and "imessage" not in _currently_enabled
):
_currently_enabled.add("imessage")
# Direct pip packages for each channel extra. Used to install the
# exact dependency without requiring the evoscientist package itself
# to be resolvable on PyPI (e.g. editable / dev installs).
_CHANNEL_PIP_DEPS: dict[str, list[str]] = {
"telegram": ["python-telegram-bot>=21.0"],
"discord": ["discord.py>=2.3"],
"slack": ["slack-sdk>=3.27", "aiohttp>=3.9"],
"feishu": ["aiohttp>=3.9", "qrcode>=7.4"],
"dingtalk": ["aiohttp>=3.9"],
"wechat": [
"pycryptodome>=3.20",
"aiohttp>=3.9",
"qrcode>=7.4",
"certifi>=2024.0",
],
"qq": ["qq-botpy>=1.0", "cryptography>=41.0", "qrcode>=7.4"],
}
# Channel definitions:
# (value, display_name, required_fields, import_check, pip_extra)
# required_fields entries are (field_name, prompt_label, is_secret).
# ``is_secret=True`` triggers a password prompt (no echo, no default echo)
# so bot tokens / OAuth secrets / IMAP+SMTP passwords don't leak into
# terminal scrollback, screen recordings, or support sessions.
_CHANNELS = [
(
"telegram",
"Telegram",
[("telegram_bot_token", "Bot token (from @BotFather)", True)],
"telegram",
"telegram",
),
(
"discord",
"Discord",
[("discord_bot_token", "Bot token", True)],
"discord",
"discord",
),
(
"slack",
"Slack",
[
("slack_bot_token", "Bot token (xoxb-...)", True),
("slack_app_token", "App token for Socket Mode (xapp-...)", True),
],
"slack_sdk",
"slack",
),
(
"feishu",
"Feishu",
[
("feishu_app_id", "App ID", False),
("feishu_app_secret", "App Secret", True),
],
"aiohttp",
"feishu",
),
(
"dingtalk",
"DingTalk",
[
("dingtalk_client_id", "Client ID (AppKey)", False),
("dingtalk_client_secret", "Client Secret (AppSecret)", True),
],
"aiohttp",
"dingtalk",
),
(
"wechat",
"WeChat",
[], # backend-specific fields prompted in the wechat branch below
("aiohttp", "qrcode", "Crypto", "certifi"),
"wechat",
),
(
"email",
"Email",
[
("email_imap_host", "IMAP host", False),
("email_imap_username", "IMAP username", False),
("email_imap_password", "IMAP password", True),
("email_smtp_host", "SMTP host", False),
("email_smtp_username", "SMTP username", False),
("email_smtp_password", "SMTP password", True),
("email_from_address", "From address", False),
],
None,
None,
),
(
"qq",
"QQ",
[
("qq_app_id", "App ID", False),
("qq_app_secret", "App Secret", True),
],
"botpy",
"qq",
),
(
"signal",
"Signal",
[("signal_phone_number", "Phone number (E.164)", False)],
None,
None,
),
("imessage", "iMessage", [], None, None), # handled via _setup_imessage()
]
choices = [
Choice(
title=display,
value=value,
checked=value in _currently_enabled,
)
for value, display, *_ in _CHANNELS
]
selected = questionary.checkbox(
"Select channels to enable (Space to toggle, Enter to confirm):",
choices=choices,
style=WIZARD_STYLE,
qmark=QMARK,
).ask()
if selected is None:
raise KeyboardInterrupt()
updates: dict[str, object] = {}
if not selected:
updates["channel_enabled"] = ""
updates["imessage_enabled"] = False
return updates
from ...mcp.registry import install_library, pip_install_hint
# Build a lookup for channel definitions
_ch_lookup = {
v: (v, d, fields, imp, extra) for v, d, fields, imp, extra in _CHANNELS
}
enabled_channels: list[str] = []
for ch_name in selected:
_, display, required_fields, import_check, pip_extra = _ch_lookup[ch_name]
console.print(f"\n [bold cyan]── {display} ──[/bold cyan]")
# Check pip dependency before proceeding
if import_check:
_required_imports: tuple[str, ...] = (
(import_check,)
if isinstance(import_check, str)
else tuple(import_check)
)
_pkg_ready = False
try:
for _module_name in _required_imports:
__import__(_module_name)
_pkg_ready = True
except ImportError:
console.print(" [yellow]✗ Required package not installed.[/yellow]")
# Determine packages to install
_pip_pkgs = _CHANNEL_PIP_DEPS.get(pip_extra, []) if pip_extra else []
_pkg_display = (
" ".join(f'"{p}"' for p in _pip_pkgs)
if _pip_pkgs
else f'"evoscientist[{pip_extra}]"'
)
install_now = questionary.confirm(
f"Install {_pkg_display} now?",
default=True,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if install_now is None:
raise KeyboardInterrupt() from None
if install_now:
console.print(f" [dim]Installing {_pkg_display}...[/dim]")
if _pip_pkgs:
_ok = all(install_library(p) for p in _pip_pkgs)
else:
_ok = install_library(f"evoscientist[{pip_extra}]")
if _ok:
# Verify the imports actually work now
try:
for _module_name in _required_imports:
__import__(_module_name)
console.print(" [green]✓ Installed successfully.[/green]")
_pkg_ready = True
except ImportError:
console.print(
" [red]✗ Package installed but import failed.[/red]"
)
console.print(
" [dim]Try restarting and running:[/dim] evosci channel setup"
)
else:
console.print(" [red]✗ Installation failed.[/red]")
console.print(
f" [dim]Run manually:[/dim] {pip_install_hint()} {_pkg_display}"
)
if not _pkg_ready:
# Previously-enabled channels are silently dropped from
# ``channel_enabled`` if we just ``continue`` — warn.
if ch_name in _currently_enabled:
console.print(
f" [bold yellow]⚠ {display} will be DISABLED[/bold yellow]"
" [dim](dependency missing — re-run after install)[/dim]"
)
else:
console.print(
f" [dim]Skipping {display} — dependency not installed.[/dim]"
)
continue
# Special handling for iMessage
if ch_name == "imessage":
ready = _setup_imessage()
if not ready:
console.print()
enable_anyway = questionary.confirm(
"Enable iMessage anyway? (will try to connect on startup)",
default=False,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if enable_anyway is None:
raise KeyboardInterrupt()
if not enable_anyway:
continue
# Allowed senders
senders = questionary.text(
"Allowed senders (comma-separated, empty = all):",
default=getattr(config, "imessage_allowed_senders", ""),
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if senders is None:
raise KeyboardInterrupt()
updates["imessage_enabled"] = True
updates["imessage_allowed_senders"] = senders.strip()
enabled_channels.append("imessage")
continue
# QQ: offer scan-to-configure before falling back to manual entry.
# The bot must already exist at q.qq.com — scanning binds the
# developer's QQ account to it and returns app_id + client_secret.
_qq_scanned = False
_feishu_scanned = False
if ch_name == "qq":
scan_choices = [
Choice(
title="Scan QR code (recommended — auto-fill App ID & Secret)",
value="scan",
),
Choice(title="Enter App ID and Secret manually", value="manual"),
]
scan_choice = questionary.select(
"Configure QQ Bot:",
choices=scan_choices,
default="scan",
style=WIZARD_STYLE,
qmark=f" {QMARK}",
use_indicator=True,
).ask()
if scan_choice is None:
raise KeyboardInterrupt()
if scan_choice == "scan":
# Preflight: AES-GCM decryption needs `cryptography`.
# `qrcode` is a soft dep — onboard.py degrades to URL-only display.
try:
import cryptography # noqa: F401
except ImportError:
console.print(
' [yellow]✗ QR scan requires "cryptography".[/yellow]'
)
install_now = questionary.confirm(
'Install "cryptography" now?',
default=True,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if install_now is None:
raise KeyboardInterrupt() from None
if install_now and install_library("cryptography>=41.0"):
console.print(" [green]✓ Installed cryptography.[/green]")
else:
console.print(
" [yellow]⚠ Falling back to manual entry.[/yellow]"
)
scan_choice = "manual"
if scan_choice == "scan":
from ...channels.qq.onboard import qr_register
console.print(
" [dim]Make sure the bot is registered at"
" https://q.qq.com first — scanning binds an"
" existing app, it does not create one.[/dim]"
)
try:
creds = qr_register()
except Exception as exc:
console.print(f" [red]✗ Scan failed: {exc}[/red]")
creds = None
if creds:
updates["qq_app_id"] = creds["app_id"]
updates["qq_app_secret"] = creds["client_secret"]
console.print(
f" [green]✓ Bound QQ Bot (App ID: {creds['app_id']})[/green]"
)
_qq_scanned = True
else:
console.print(
" [yellow]⚠ Scan did not complete — falling"
" back to manual entry.[/yellow]"
)
# Feishu: offer scan-to-create before falling back to manual entry.
# Unlike QQ, this provisions a brand-new PersonalAgent app with the
# required IM permissions attached, then returns app_id + app_secret.
if ch_name == "feishu":
scan_choices = [
Choice(
title="Scan QR code (recommended — auto-create app, fill App ID & Secret)",
value="scan",
),
Choice(title="Enter App ID and Secret manually", value="manual"),
]
scan_choice = questionary.select(
"Configure Feishu / Lark:",
choices=scan_choices,
default="scan",
style=WIZARD_STYLE,
qmark=f" {QMARK}",
use_indicator=True,
).ask()
if scan_choice is None:
raise KeyboardInterrupt()
if scan_choice == "scan":
# `qrcode` is the only soft dep needed — onboard prints the URL
# if it's missing, but the UX is much worse, so offer to install.
try:
import qrcode # noqa: F401
except ImportError:
console.print(
' [yellow]✗ QR scan looks best with "qrcode".[/yellow]'
)
install_now = questionary.confirm(
'Install "qrcode" now?',
default=True,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if install_now is None:
raise KeyboardInterrupt() from None
if install_now and install_library("qrcode>=7.4"):
console.print(" [green]✓ Installed qrcode.[/green]")
else:
console.print(
" [yellow]⚠ Falling back to manual entry.[/yellow]"
)
scan_choice = "manual"
if scan_choice == "scan":
# Region selection — accounts.feishu.cn vs accounts.larksuite.com.
# The poll endpoint auto-switches if the scanning user is on the
# other tenant, so this is just a starting hint.
region_choices = [
Choice(title="Feishu (飞书, mainland China)", value="feishu"),
Choice(title="Lark (overseas)", value="lark"),
]
region = questionary.select(
"Region:",
choices=region_choices,
default="feishu",
style=WIZARD_STYLE,
qmark=f" {QMARK}",
use_indicator=True,
).ask()
if region is None:
raise KeyboardInterrupt()
if scan_choice == "scan":
from ...channels.feishu.onboard import qr_register
console.print(
" [dim]A QR code will be printed below — open Feishu or"
" Lark on your phone and scan it. The platform will"
" auto-create a bot app with IM permissions and return"
" the credentials here.[/dim]"
)
try:
creds = qr_register(initial_domain=region)
except Exception as exc:
console.print(f" [red]✗ Scan failed: {exc}[/red]")
creds = None
if creds:
updates["feishu_app_id"] = creds["app_id"]
updates["feishu_app_secret"] = creds["app_secret"]
# Sync open-platform domain to the resolved region
updates["feishu_domain"] = (
"https://open.larksuite.com"
if creds.get("domain") == "lark"
else "https://open.feishu.cn"
)
bot_name = creds.get("bot_name")
if bot_name:
console.print(
f' [green]✓ Bound Feishu bot "{bot_name}"'
f" (App ID: {creds['app_id']})[/green]"
)
else:
console.print(
f" [green]✓ Bound Feishu app"
f" (App ID: {creds['app_id']})[/green]"
)
_feishu_scanned = True
else:
console.print(
" [yellow]⚠ Scan did not complete — falling"
" back to manual entry.[/yellow]"
)
# WeChat: pick backend (wecom / wechatmp / personal), then prompt
# backend-specific fields. Personal-WeChat has no static credentials —
# we offer an interactive QR-scan that obtains and persists them.
if ch_name == "wechat":
backend_choices = [
Choice(
title="WeCom (企业微信应用) — most stable, official API",
value="wecom",
),
Choice(
title="Official Account (微信公众号) — public-facing bots",
value="wechatmp",
),
Choice(
title="Personal WeChat (个人微信, iLink) — QR-code scan login",
value="personal",
),
]
wechat_backend = questionary.select(
"WeChat backend:",
choices=backend_choices,
default=getattr(config, "wechat_backend", "") or "wecom",
style=WIZARD_STYLE,
qmark=f" {QMARK}",
use_indicator=True,
).ask()
if wechat_backend is None:
raise KeyboardInterrupt()
updates["wechat_backend"] = wechat_backend
# Both WeCom and WeChat MP need the same non-empty-required
# treatment as the generic required_fields loop below — newly
# enabling either with blank credentials would leave the channel
# half-configured and only fail at first message.
wechat_newly_enabled = "wechat" not in _currently_enabled
wechat_fields_for_backend: list[tuple[str, str, bool]] = []
if wechat_backend == "wecom":
wechat_fields_for_backend = [
("wechat_wecom_corp_id", "WeCom Corp ID", False),
("wechat_wecom_agent_id", "WeCom Agent ID", False),
("wechat_wecom_secret", "WeCom Secret", True),
]
elif wechat_backend == "wechatmp":
wechat_fields_for_backend = [
("wechat_mp_app_id", "Official Account App ID", False),
("wechat_mp_app_secret", "Official Account App Secret", True),
]
if wechat_backend in ("wecom", "wechatmp"):
for field_name, prompt_label, is_secret in wechat_fields_for_backend:
current = getattr(config, field_name, "")
while True:
if is_secret:
masked_hint = (
f" (current: ***{current[-4:]})" if current else ""
)
value = questionary.password(
f"{prompt_label}{masked_hint}:",
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
else:
value = questionary.text(
f"{prompt_label}:",
default=current,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if value is None:
raise KeyboardInterrupt()
value = value.strip()
if not value and current:
break # keep existing
if not value and wechat_newly_enabled:
console.print(
f" [yellow]{prompt_label} is required to "
"enable WeChat. Press Ctrl+C to cancel.[/yellow]"
)
continue
updates[field_name] = value
break
elif wechat_backend == "personal":
personal_choices = [
Choice(
title="Scan QR code now (recommended — login to a personal WeChat account)",
value="scan",
),
Choice(
title="I already have an account_id — enter it manually",
value="manual",
),
]
personal_choice = questionary.select(
"Personal WeChat login:",
choices=personal_choices,
default="scan",
style=WIZARD_STYLE,
qmark=f" {QMARK}",
use_indicator=True,
).ask()
if personal_choice is None:
raise KeyboardInterrupt()
if personal_choice == "scan":
from ...channels.wechat.personal import _account_dir as _wp_dir
_accounts_path = _wp_dir()
console.print(
" [dim]A QR code will be printed below — open WeChat on"
" your phone and scan it. The session token is saved"
f" to {_accounts_path}.[/dim]"
)
try:
import asyncio
from ...channels.wechat.personal import qr_login
creds = asyncio.run(qr_login())
except Exception as exc:
console.print(f" [red]✗ Scan failed: {exc}[/red]")
creds = None
if creds:
updates["wechat_personal_account_id"] = creds["account_id"]
# Token is persisted on disk by qr_login(); the channel
# reads it from the per-account store at runtime, so we
# intentionally do NOT copy it into the main config here
# (avoids stale duplicates and broader secret exposure).
console.print(
f" [green]✓ Logged in (account_id: "
f"{creds['account_id'][:12]}…)[/green]"
)
else:
console.print(
" [yellow]⚠ QR login did not complete — falling"
" back to manual entry.[/yellow]"
)
personal_choice = "manual"
if personal_choice == "manual":
current_id = getattr(config, "wechat_personal_account_id", "")
wechat_newly_enabled = "wechat" not in _currently_enabled
while True:
account_id = questionary.text(
"iLink account_id (from a previous --qr-login run):",
default=current_id,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if account_id is None:
raise KeyboardInterrupt()
account_id = account_id.strip()
if not account_id and current_id:
break # keep existing
if not account_id and wechat_newly_enabled:
console.print(
" [yellow]account_id is required to enable "
"WeChat Personal. Press Ctrl+C to cancel.[/yellow]"
)
continue
updates["wechat_personal_account_id"] = account_id
break
# Prompt for required fields. Secret fields use ``questionary.password``
# so the entered value (and the existing one shown as a hint) are
# never echoed to the terminal — see _CHANNELS docstring above.
if not _qq_scanned and not _feishu_scanned:
# "Newly enabled" = this channel wasn't in the user's prior
# ``channel_enabled`` list. Required fields with no existing
# value must be non-empty for newly enabled channels — saving
# blanks leaves the channel half-configured and only surfaces
# the problem on the first message.
newly_enabled = ch_name not in _currently_enabled
for field_name, prompt_label, is_secret in required_fields:
current = getattr(config, field_name, "")
while True:
if is_secret:
masked_hint = (
f" (current: ***{current[-4:]})" if current else ""
)
value = questionary.password(
f"{prompt_label}{masked_hint}:",
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
else:
value = questionary.text(
f"{prompt_label}:",
default=current,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if value is None:
raise KeyboardInterrupt()
value = value.strip()
# Empty input + existing value → keep existing (this is
# the "re-run wizard, no change to this field" path).
if not value and current:
break
# Empty input + newly enabling channel → not OK; the
# channel would be enabled with broken creds. Re-prompt.
if not value and newly_enabled:
console.print(
f" [yellow]{prompt_label} is required to enable "
f"{display}. Press Ctrl+C to cancel instead.[/yellow]"
)
continue
# Empty input + previously enabled but never set
# (unlikely, but tolerate) → still allow blank-through
# so the user isn't blocked re-running configure later.
updates[field_name] = value
break
# Feishu: subscription mode + optional fields
if ch_name == "feishu":
mode_choices = [
Choice(
title="Webhook (requires public IP / port forwarding)",
value="webhook",
),
Choice(
title="WebSocket long connection (no public IP needed)",
value="websocket",
),
]
sub_mode = questionary.select(
"Subscription mode:",
choices=mode_choices,
default="webhook",
style=WIZARD_STYLE,
qmark=f" {QMARK}",
use_indicator=True,
).ask()
if sub_mode is None:
raise KeyboardInterrupt()
updates["feishu_subscription_mode"] = sub_mode
if sub_mode == "websocket":
# WebSocket mode needs lark-oapi SDK
try:
__import__("lark_oapi")
except ImportError:
console.print(
' [yellow]✗ WebSocket mode requires "lark-oapi".[/yellow]'
)
install_sdk = questionary.confirm(
'Install "lark-oapi>=1.4.0" now?',
default=True,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if install_sdk is None:
raise KeyboardInterrupt() from None
if install_sdk:
console.print(' [dim]Installing "lark-oapi"...[/dim]')
if install_library("lark-oapi>=1.4.0"):
console.print(" [green]✓ Installed successfully.[/green]")
else:
console.print(" [red]✗ Installation failed.[/red]")
console.print(
f" [dim]Run manually:[/dim] {pip_install_hint()} "
'"lark-oapi>=1.4.0"'
)
else:
# Webhook mode: prompt optional verification/encryption fields.
# Both are credentials — use password() so they don't echo.
console.print(
" [dim]The following fields are optional"
" (press Enter to skip):[/dim]"
)
for field_name, prompt_label in [
("feishu_verification_token", "Verification Token (optional)"),
("feishu_encrypt_key", "Encrypt Key (optional)"),
]:
current = getattr(config, field_name, "")
masked_hint = f" (current: ***{current[-4:]})" if current else ""
value = questionary.password(
f"{prompt_label}{masked_hint}:",
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if value is None:
raise KeyboardInterrupt()
value = value.strip()
if not value and current:
# Keep existing value when user just presses Enter.
continue
updates[field_name] = value
# Allowed senders (common for all channels)
senders_field = f"{ch_name}_allowed_senders"
if hasattr(config, senders_field):
senders = questionary.text(
"Allowed senders (comma-separated, empty = all):",
default=getattr(config, senders_field, ""),
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if senders is None:
raise KeyboardInterrupt()
updates[senders_field] = senders.strip()
# Probe validation
_probe_channel(ch_name, config, updates)
enabled_channels.append(ch_name)
updates["channel_enabled"] = ",".join(enabled_channels)
# Keep legacy field in sync
updates["imessage_enabled"] = "imessage" in enabled_channels
# --- Common prompt: send thinking (shown when any channel is enabled) ---
if enabled_channels:
console.print("\n [bold cyan]── Channel Settings ──[/bold cyan]")
thinking_choices = [
Choice(title="On (forward model reasoning)", value=True),
Choice(title="Off (only send final responses)", value=False),
]
send_thinking = questionary.select(
"Send thinking panel in channel?",
choices=thinking_choices,
default=config.channel_send_thinking,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
use_indicator=True,
).ask()
if send_thinking is None:
raise KeyboardInterrupt()
updates["channel_send_thinking"] = send_thinking
return updates
def _probe_channel(
ch_name: str,
config: EvoScientistConfig,
updates: dict[str, object],
) -> None:
"""Run the probe for a channel type and print the result.
Non-fatal: prints a warning on failure but does not prevent enabling.
"""
import asyncio
def _val(key: str, fallback: str = "") -> str:
"""Get a value from updates first, then config, then fallback."""
if key in updates:
return str(updates[key])
return str(getattr(config, key, fallback))
console.print(" [dim]Validating credentials...[/dim]")
async def _run() -> tuple[bool, str]:
if ch_name == "telegram":
from ...channels.telegram.probe import validate_telegram_token
return await validate_telegram_token(
_val("telegram_bot_token"),
_val("telegram_proxy") or None,
)
elif ch_name == "discord":
from ...channels.discord.probe import validate_discord_token
return await validate_discord_token(
_val("discord_bot_token"),
_val("discord_proxy") or None,
)
elif ch_name == "slack":
from ...channels.slack.probe import validate_slack_tokens
return await validate_slack_tokens(
_val("slack_bot_token"),
_val("slack_app_token") or None,
_val("slack_proxy") or None,
)
elif ch_name == "wechat":
backend = _val("wechat_backend", "wecom")
if backend == "wechatmp":
from ...channels.wechat.probe import validate_wechat_mp
return await validate_wechat_mp(
_val("wechat_mp_app_id"),
_val("wechat_mp_app_secret"),
_val("wechat_proxy") or None,
)
elif backend == "personal":
from ...channels.wechat.probe import validate_wechat_personal
return await validate_wechat_personal(
_val("wechat_personal_account_id"),
_val("wechat_personal_token"),
)
else:
from ...channels.wechat.probe import validate_wecom
return await validate_wecom(
_val("wechat_wecom_corp_id"),
_val("wechat_wecom_secret"),
_val("wechat_proxy") or None,
)
elif ch_name == "feishu":
from ...channels.feishu.probe import validate_feishu_credentials
return await validate_feishu_credentials(
_val("feishu_app_id"),
_val("feishu_app_secret"),
_val("feishu_domain", "https://open.feishu.cn"),
)
elif ch_name == "dingtalk":
from ...channels.dingtalk.probe import validate_dingtalk
return await validate_dingtalk(
_val("dingtalk_client_id"),
_val("dingtalk_client_secret"),
_val("dingtalk_proxy") or None,
)
elif ch_name == "email":
from ...channels.email.probe import validate_email_imap
return await validate_email_imap(
_val("email_imap_host"),
int(_val("email_imap_port", "993")),
_val("email_imap_username"),
_val("email_imap_password"),
_val("email_imap_use_ssl", "True").lower() not in ("false", "0", "no"),
)
elif ch_name == "qq":
from ...channels.qq.probe import validate_qq
return await validate_qq(
_val("qq_app_id"),
_val("qq_app_secret"),
)
elif ch_name == "signal":
from ...channels.signal.probe import validate_signal
return await validate_signal(
_val("signal_phone_number"),
_val("signal_cli_path", "signal-cli"),
int(_val("signal_rpc_port", "7583")),
)
else:
return True, "No probe available"
try:
try:
loop = asyncio.get_event_loop()
if loop.is_running():
import nest_asyncio # type: ignore[import-untyped]
nest_asyncio.apply()
except RuntimeError:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
ok, detail = loop.run_until_complete(_run())
if ok:
console.print(f" [green]✓ {detail}[/green]")
else:
console.print(f" [yellow]⚠ {detail}[/yellow]")
console.print(
" [dim]Channel will still be enabled — check credentials later.[/dim]"
)
except Exception as e:
console.print(f" [yellow]⚠ Could not validate: {e}[/yellow]")
console.print(
" [dim]Channel will still be enabled — check credentials later.[/dim]"
)
# =============================================================================
# Progress Rendering (for tests and potential future use)
# =============================================================================

Some files were not shown because too many files have changed in this diff Show More