Compare commits
92 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| befdb17e0b | |||
| ddf337fc40 | |||
| 2e5db60dc5 | |||
| 4c338ed914 | |||
| 470cf75722 | |||
| d4b53bfb08 | |||
| ea99ce9f7e | |||
| 57bae2e6ae | |||
| 418abca4f4 | |||
| 0410b40f57 | |||
| a486b85851 | |||
| be8c23861d | |||
| 9bd7a37a77 | |||
| b36c19a22a | |||
| dae67e9299 | |||
| dedd53bad6 | |||
| a45563ea7f | |||
| 7dbb68d807 | |||
| 382ed305cc | |||
| 324011ab3f | |||
| 812e05d29d | |||
| 561e161123 | |||
| ee14afe3a3 | |||
| 5d893c1dc6 | |||
| c683f6e739 | |||
| 8d68dfa652 | |||
| a95f5d4c45 | |||
| 99075345ab | |||
| 8383f343ae | |||
| 3d6cc959e3 | |||
| ce0bec1f91 | |||
| 3683cbfc13 | |||
| c7be701abf | |||
| 386c8130ea | |||
| 8d7d95a20d | |||
| 836ff5c939 | |||
| 8764f2012f | |||
| bcee009917 | |||
| 47c8da3e5b | |||
| 1b324906fd | |||
| b500ccc311 | |||
| 1d02a909d3 | |||
| 71d5649414 | |||
| da15b70535 | |||
| 932c934485 | |||
| 1f3e8f57a3 | |||
| fb329e4aaa | |||
| 3f45ebd6fa | |||
| a1be0d1437 | |||
| e086f76da7 | |||
| 12adc62868 | |||
| 0c21a01f6f | |||
| b40b6f784d | |||
| aff63ccd65 | |||
| ab1a6b0062 | |||
| bd2464423a | |||
| 3cda9894c7 | |||
| 3a1dbf0a0f | |||
| b5b01d50c2 | |||
| 33979e5371 | |||
| 3c5cc831c0 | |||
| a9a57cd828 | |||
| c7d93fb574 | |||
| d97ee2b917 | |||
| 2a17d14e33 | |||
| d249e320bd | |||
| f81a8b086e | |||
| 4ddf7ebe52 | |||
| 562ce0eb83 | |||
| a6a8e19dcc | |||
| 8cac50e2ef | |||
| 10c032450e | |||
| 8b1451cdda | |||
| ac58caab7b | |||
| fe70599d3d | |||
| 1863f0730c | |||
| 172c8409d4 | |||
| cb4e127d8c | |||
| 5bd3234973 | |||
| 71e480d463 | |||
| f802a49535 | |||
| e9857dde22 | |||
| 1abfc8236d | |||
| 042da63d54 | |||
| 06a9511bdd | |||
| 584b9d24ac | |||
| 89bfb548b8 | |||
| e8399b7c94 | |||
| 0f709cff8b | |||
| 05dfffbc73 | |||
| 01845f4311 | |||
| db1abce8d8 |
@@ -18,6 +18,9 @@ KIMI_API_KEY= # kimi.com/code (Kimi 代码计划)
|
|||||||
# Aggregator platforms (optional)
|
# Aggregator platforms (optional)
|
||||||
SILICONFLOW_API_KEY= # siliconflow.cn
|
SILICONFLOW_API_KEY= # siliconflow.cn
|
||||||
OPENROUTER_API_KEY= # openrouter.ai
|
OPENROUTER_API_KEY= # openrouter.ai
|
||||||
|
REQUESTY_API_KEY= # requesty.ai
|
||||||
|
ATLASCLOUD_API_KEY= # atlascloud.ai
|
||||||
|
NOVITA_API_KEY= # novita.ai
|
||||||
|
|
||||||
# Custom endpoints (optional)
|
# Custom endpoints (optional)
|
||||||
CUSTOM_OPENAI_API_KEY= # OpenAI-compatible endpoint
|
CUSTOM_OPENAI_API_KEY= # OpenAI-compatible endpoint
|
||||||
|
|||||||
@@ -5,5 +5,5 @@
|
|||||||
<rect x="54" y="5" width="62" height="24" rx="6" fill="#1565c0"/>
|
<rect x="54" y="5" width="62" height="24" rx="6" fill="#1565c0"/>
|
||||||
<text x="85" y="22" text-anchor="middle"
|
<text x="85" y="22" text-anchor="middle"
|
||||||
font-family="Inter, -apple-system, system-ui, sans-serif"
|
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.3.0</text>
|
||||||
</svg>
|
</svg>
|
||||||
|
Before Width: | Height: | Size: 555 B After Width: | Height: | Size: 555 B |
@@ -5,5 +5,5 @@
|
|||||||
<rect x="54" y="5" width="62" height="24" rx="6" fill="#2563eb"/>
|
<rect x="54" y="5" width="62" height="24" rx="6" fill="#2563eb"/>
|
||||||
<text x="85" y="22" text-anchor="middle"
|
<text x="85" y="22" text-anchor="middle"
|
||||||
font-family="Inter, -apple-system, system-ui, sans-serif"
|
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.3.0</text>
|
||||||
</svg>
|
</svg>
|
||||||
|
Before Width: | Height: | Size: 555 B After Width: | Height: | Size: 555 B |
Binary file not shown.
|
Before Width: | Height: | Size: 287 KiB After Width: | Height: | Size: 221 KiB |
@@ -11,7 +11,9 @@ jobs:
|
|||||||
timeout-minutes: 10
|
timeout-minutes: 10
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v5
|
- uses: actions/checkout@v5
|
||||||
- uses: astral-sh/setup-uv@v6
|
# Full tag required — setup-uv dropped major/minor tags in v8.0.0, so
|
||||||
|
# `@v9` does not resolve. See the note in lint.yml.
|
||||||
|
- uses: astral-sh/setup-uv@v9.0.0
|
||||||
with:
|
with:
|
||||||
python-version: "3.11"
|
python-version: "3.11"
|
||||||
cache-dependency-glob: "**/pyproject.toml"
|
cache-dependency-glob: "**/pyproject.toml"
|
||||||
|
|||||||
@@ -3,7 +3,10 @@ name: Docker
|
|||||||
on:
|
on:
|
||||||
push:
|
push:
|
||||||
branches: ["main"]
|
branches: ["main"]
|
||||||
tags: ["v*"]
|
# Version images build when a GitHub Release is published — the same event
|
||||||
|
# that triggers the PyPI upload (publish.yml), so the two channels stay in sync.
|
||||||
|
release:
|
||||||
|
types: [published]
|
||||||
pull_request:
|
pull_request:
|
||||||
paths:
|
paths:
|
||||||
- "Dockerfile"
|
- "Dockerfile"
|
||||||
@@ -32,6 +35,19 @@ jobs:
|
|||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@93cb6efe18208431cddfb8368fd83d5badbf9bfd # v5.0.1
|
- uses: actions/checkout@93cb6efe18208431cddfb8368fd83d5badbf9bfd # v5.0.1
|
||||||
|
|
||||||
|
# Same guard as publish.yml — a release whose tag mismatches pyproject
|
||||||
|
# must not publish versioned images either.
|
||||||
|
- name: Guard — package version must match the release tag
|
||||||
|
if: github.event_name == 'release'
|
||||||
|
run: |
|
||||||
|
VERSION=$(grep -m1 '^version = ' pyproject.toml | sed -E 's/^version = "(.*)"/\1/')
|
||||||
|
TAG="${GITHUB_REF_NAME#v}"
|
||||||
|
echo "pyproject version: $VERSION | release tag: $TAG"
|
||||||
|
if [ "$VERSION" != "$TAG" ]; then
|
||||||
|
echo "::error::pyproject version ($VERSION) does not match release tag ($TAG); refusing to publish images."
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
- uses: docker/setup-qemu-action@c7c53464625b32c7a7e944ae62b3e17d2b600130 # v3.7.0
|
- uses: docker/setup-qemu-action@c7c53464625b32c7a7e944ae62b3e17d2b600130 # v3.7.0
|
||||||
- uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12.0
|
- uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12.0
|
||||||
|
|
||||||
|
|||||||
@@ -11,7 +11,11 @@ jobs:
|
|||||||
timeout-minutes: 5
|
timeout-minutes: 5
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v5
|
- uses: actions/checkout@v5
|
||||||
- uses: astral-sh/setup-uv@v6
|
# Pinned to a full tag on purpose: setup-uv stopped publishing major /
|
||||||
|
# minor tags in v8.0.0 (supply-chain hardening), so `@v9` does not exist
|
||||||
|
# and would fail the job. Releases are immutable, so the tag is as safe
|
||||||
|
# as a SHA. v7 is where the action moved off the deprecated node20.
|
||||||
|
- uses: astral-sh/setup-uv@v9.0.0
|
||||||
with:
|
with:
|
||||||
python-version: "3.11"
|
python-version: "3.11"
|
||||||
cache-dependency-glob: "**/pyproject.toml"
|
cache-dependency-glob: "**/pyproject.toml"
|
||||||
|
|||||||
@@ -0,0 +1,67 @@
|
|||||||
|
name: Publish to PyPI
|
||||||
|
|
||||||
|
# Publishes EvoScientist to PyPI via OpenID Connect (trusted publishing) — no
|
||||||
|
# API token or password involved. Fires when a GitHub Release is published
|
||||||
|
# (i.e. after `gh release create vX.Y.Z`), builds the sdist + wheel, guards
|
||||||
|
# that the package version matches the release tag, then uploads with a
|
||||||
|
# short-lived OIDC token.
|
||||||
|
#
|
||||||
|
# One-time PyPI setup (Manage project -> Publishing -> Add a new publisher):
|
||||||
|
# Owner: EvoScientist
|
||||||
|
# Repository: EvoScientist
|
||||||
|
# Workflow name: publish.yml
|
||||||
|
# Environment name: pypi
|
||||||
|
on:
|
||||||
|
release:
|
||||||
|
types: [published]
|
||||||
|
# Manual re-run escape hatch — dispatch it from the release tag; the version
|
||||||
|
# guard rejects any non-tag ref.
|
||||||
|
workflow_dispatch:
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
publish:
|
||||||
|
name: Build and publish to PyPI
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
timeout-minutes: 15
|
||||||
|
environment:
|
||||||
|
name: pypi
|
||||||
|
url: https://pypi.org/project/EvoScientist/
|
||||||
|
permissions:
|
||||||
|
id-token: write # required to mint the OIDC token PyPI verifies
|
||||||
|
steps:
|
||||||
|
# Actions in this job are pinned to full commit SHAs (like docker.yml):
|
||||||
|
# it holds id-token: write and PyPI publishing authority.
|
||||||
|
- uses: actions/checkout@93cb6efe18208431cddfb8368fd83d5badbf9bfd # v5.0.1
|
||||||
|
with:
|
||||||
|
persist-credentials: false
|
||||||
|
- uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0
|
||||||
|
with:
|
||||||
|
python-version: "3.11"
|
||||||
|
# No shared cache in the publishing job — build from clean sources only.
|
||||||
|
enable-cache: false
|
||||||
|
|
||||||
|
- name: Build sdist + wheel
|
||||||
|
run: uv build
|
||||||
|
|
||||||
|
- name: Guard — package version must match the release tag
|
||||||
|
run: |
|
||||||
|
if [ "${GITHUB_REF_TYPE}" != "tag" ]; then
|
||||||
|
echo "::error::This workflow must run from a version tag (got ${GITHUB_REF_TYPE} '${GITHUB_REF_NAME}'); dispatch it from the release tag."
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
VERSION=$(grep -m1 '^version = ' pyproject.toml | sed -E 's/^version = "(.*)"/\1/')
|
||||||
|
TAG="${GITHUB_REF_NAME#v}"
|
||||||
|
echo "pyproject version: $VERSION | release tag: $TAG"
|
||||||
|
if [ "$VERSION" != "$TAG" ]; then
|
||||||
|
echo "::error::pyproject version ($VERSION) does not match release tag ($TAG); refusing to publish."
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
- name: Twine metadata check
|
||||||
|
run: uvx twine check dist/*
|
||||||
|
|
||||||
|
- name: Publish to PyPI (trusted publishing)
|
||||||
|
uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2
|
||||||
@@ -21,11 +21,13 @@ jobs:
|
|||||||
runs-on: ${{ matrix.os }}
|
runs-on: ${{ matrix.os }}
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v5
|
- uses: actions/checkout@v5
|
||||||
- uses: astral-sh/setup-uv@v6
|
# Full tag required — setup-uv dropped major/minor tags in v8.0.0, so
|
||||||
|
# `@v9` does not resolve. See the note in lint.yml.
|
||||||
|
- uses: astral-sh/setup-uv@v9.0.0
|
||||||
with:
|
with:
|
||||||
python-version: ${{ matrix.python-version }}
|
python-version: ${{ matrix.python-version }}
|
||||||
cache-dependency-glob: "**/pyproject.toml"
|
cache-dependency-glob: "**/pyproject.toml"
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: uv sync --dev
|
run: uv sync --dev --extra all-channels
|
||||||
- name: Run pytest
|
- name: Run pytest
|
||||||
run: uv run pytest -v --timeout=30
|
run: uv run pytest -v --timeout=30
|
||||||
|
|||||||
+531
-62
@@ -23,7 +23,11 @@ from collections.abc import Sequence
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from langchain.agents.middleware import AgentMiddleware, HumanInTheLoopMiddleware
|
from langchain.agents.middleware import (
|
||||||
|
AgentMiddleware,
|
||||||
|
HumanInTheLoopMiddleware,
|
||||||
|
TodoListMiddleware,
|
||||||
|
)
|
||||||
|
|
||||||
from . import paths as _paths_mod
|
from . import paths as _paths_mod
|
||||||
from .config import (
|
from .config import (
|
||||||
@@ -42,6 +46,9 @@ logging.getLogger("deepagents.middleware.skills").setLevel(logging.ERROR)
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from langgraph.graph.state import CompiledStateGraph
|
from langgraph.graph.state import CompiledStateGraph
|
||||||
|
|
||||||
|
from .middleware.events import MiddlewareEventSink
|
||||||
|
from .runtime import AsyncRuntime
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
# Constants
|
# Constants
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -50,6 +57,15 @@ SUBAGENTS_CONFIG = Path(__file__).parent / "subagents"
|
|||||||
SKILLS_DIR = str(Path(__file__).parent / "skills")
|
SKILLS_DIR = str(Path(__file__).parent / "skills")
|
||||||
DEFAULT_SKILL_SOURCES = ("/skills/",)
|
DEFAULT_SKILL_SOURCES = ("/skills/",)
|
||||||
|
|
||||||
|
# Tools requiring human approval on attended agents (deepagents 0.7.0 ships a
|
||||||
|
# recursive `delete` FS tool that would otherwise bypass the execute blocklist).
|
||||||
|
HITL_INTERRUPT_ON: dict[str, bool] = {
|
||||||
|
"execute": True,
|
||||||
|
"run_in_background": True,
|
||||||
|
"schedule_task": True,
|
||||||
|
"delete": True,
|
||||||
|
}
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
# Lazy state — initialized on first use, not at import time
|
# Lazy state — initialized on first use, not at import time
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -110,6 +126,23 @@ def _apply_env_from_config(cfg) -> None:
|
|||||||
apply_config_to_env(cfg)
|
apply_config_to_env(cfg)
|
||||||
|
|
||||||
|
|
||||||
|
def _deepagents_provides_todo_list() -> bool:
|
||||||
|
"""True when deepagents' own default chain already includes the todo list.
|
||||||
|
|
||||||
|
deepagents 0.7.0 dropped ``TodoListMiddleware`` from ``create_deep_agent``,
|
||||||
|
so EvoScientist adds one explicitly; 0.6.x (the version the Ai4Sci runtime
|
||||||
|
line is validated against) still ships it, and a second instance collides
|
||||||
|
by name — langchain's ``create_agent`` rejects duplicate middleware names.
|
||||||
|
"""
|
||||||
|
from importlib.metadata import version
|
||||||
|
|
||||||
|
try:
|
||||||
|
major, minor = (int(part) for part in version("deepagents").split(".")[:2])
|
||||||
|
except Exception: # pragma: no cover - unparseable metadata
|
||||||
|
return False
|
||||||
|
return (major, minor) < (0, 7)
|
||||||
|
|
||||||
|
|
||||||
def _ensure_config(config=None):
|
def _ensure_config(config=None):
|
||||||
"""Return cached config. If *config* is passed, cache and use it."""
|
"""Return cached config. If *config* is passed, cache and use it."""
|
||||||
if config is not None:
|
if config is not None:
|
||||||
@@ -245,7 +278,11 @@ def _load_mcp_config_once() -> tuple[str, dict]:
|
|||||||
return sig, cfg
|
return sig, cfg
|
||||||
|
|
||||||
|
|
||||||
def _load_mcp_tools_cached(on_progress=None) -> dict[str, list]:
|
def _load_mcp_tools_cached(
|
||||||
|
on_progress=None,
|
||||||
|
*,
|
||||||
|
runtime: "AsyncRuntime | None" = None,
|
||||||
|
) -> dict[str, list]:
|
||||||
"""Load MCP tools with config-aware caching.
|
"""Load MCP tools with config-aware caching.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -266,7 +303,11 @@ def _load_mcp_tools_cached(on_progress=None) -> dict[str, list]:
|
|||||||
if _MCP_TOOLS_CACHE_KEY == cfg_key and _MCP_TOOLS_CACHE_VALUE is not None:
|
if _MCP_TOOLS_CACHE_KEY == cfg_key and _MCP_TOOLS_CACHE_VALUE is not None:
|
||||||
return {k: list(v) for k, v in _MCP_TOOLS_CACHE_VALUE.items()}
|
return {k: list(v) for k, v in _MCP_TOOLS_CACHE_VALUE.items()}
|
||||||
|
|
||||||
loaded = load_mcp_tools(config=cfg, on_progress=on_progress)
|
loaded = load_mcp_tools(
|
||||||
|
config=cfg,
|
||||||
|
on_progress=on_progress,
|
||||||
|
runtime=runtime,
|
||||||
|
)
|
||||||
_MCP_TOOLS_CACHE_KEY = cfg_key
|
_MCP_TOOLS_CACHE_KEY = cfg_key
|
||||||
_MCP_TOOLS_CACHE_VALUE = {k: list(v) for k, v in loaded.items()}
|
_MCP_TOOLS_CACHE_VALUE = {k: list(v) for k, v in loaded.items()}
|
||||||
return {k: list(v) for k, v in loaded.items()}
|
return {k: list(v) for k, v in loaded.items()}
|
||||||
@@ -316,6 +357,7 @@ def _inject_subagent_middleware(
|
|||||||
RecoverableToolEffectMiddleware,
|
RecoverableToolEffectMiddleware,
|
||||||
RepetitiveToolCallGuardMiddleware,
|
RepetitiveToolCallGuardMiddleware,
|
||||||
ToolErrorHandlerMiddleware,
|
ToolErrorHandlerMiddleware,
|
||||||
|
ToolHistoryRepairMiddleware,
|
||||||
ToolProtocolGuardMiddleware,
|
ToolProtocolGuardMiddleware,
|
||||||
create_context_editing_middleware,
|
create_context_editing_middleware,
|
||||||
create_memory_lifecycle_middleware,
|
create_memory_lifecycle_middleware,
|
||||||
@@ -366,12 +408,15 @@ def _inject_subagent_middleware(
|
|||||||
max_consecutive_errors=max_consecutive_tool_errors,
|
max_consecutive_errors=max_consecutive_tool_errors,
|
||||||
),
|
),
|
||||||
ToolProtocolGuardMiddleware(),
|
ToolProtocolGuardMiddleware(),
|
||||||
|
# Sync subagents replay their own history to strict providers too.
|
||||||
|
ToolHistoryRepairMiddleware(),
|
||||||
# Subagents share the main agent's model: use the threaded
|
# Subagents share the main agent's model: use the threaded
|
||||||
# ``chat_model`` on the pure path, else defer to the factory's
|
# ``chat_model`` on the pure path, else defer to the factory's
|
||||||
# ``_ensure_chat_model()`` fallback (when ``chat_model=None``).
|
# ``_ensure_chat_model()`` fallback (when ``chat_model=None``).
|
||||||
create_context_editing_middleware(chat_model),
|
create_context_editing_middleware(chat_model),
|
||||||
create_runtime_context_middleware(),
|
create_runtime_context_middleware(),
|
||||||
ToolErrorHandlerMiddleware(),
|
ToolErrorHandlerMiddleware(),
|
||||||
|
TodoListMiddleware(),
|
||||||
ContextOverflowMapperMiddleware(),
|
ContextOverflowMapperMiddleware(),
|
||||||
]
|
]
|
||||||
if memory_controls.memory_enabled:
|
if memory_controls.memory_enabled:
|
||||||
@@ -443,8 +488,46 @@ def _apply_budgeted_skill_context(kwargs: dict, backend) -> dict:
|
|||||||
return updated
|
return updated
|
||||||
|
|
||||||
|
|
||||||
|
def _fold_expert_subagents(subs: list[dict], tool_registry: dict) -> None:
|
||||||
|
"""Append expert-skill sub-agent specs to ``subs``, guarding names.
|
||||||
|
|
||||||
|
Each installed expert skill becomes an in-process sub-agent entry so
|
||||||
|
the main agent's ``task`` tool (and the QuickJS ``task()`` global) can
|
||||||
|
dispatch to it in-turn by name. The same experts independently get a
|
||||||
|
background reach via ``build_expert_async_subagent_specs``; the two
|
||||||
|
reaches land on separate tool schemas, so sharing the name is safe.
|
||||||
|
|
||||||
|
Skips (with a warning) any expert whose ``name`` collides with a
|
||||||
|
subagent already in ``subs`` or with ``general-purpose``. The reserved
|
||||||
|
name matters because ``_ensure_general_purpose_subagent`` runs right
|
||||||
|
after this and early-returns when it sees the slot occupied — an expert
|
||||||
|
named ``general-purpose`` would silently take the slot and deepagents'
|
||||||
|
default subagent prompt would never reach the agent.
|
||||||
|
"""
|
||||||
|
from deepagents.middleware.subagents import GENERAL_PURPOSE_SUBAGENT
|
||||||
|
|
||||||
|
from .subagents.expert_container import build_expert_subagent_specs
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
taken = {s.get("name") for s in subs} | {GENERAL_PURPOSE_SUBAGENT["name"]}
|
||||||
|
for spec in build_expert_subagent_specs(tool_registry=tool_registry):
|
||||||
|
name = spec["name"]
|
||||||
|
if name in taken:
|
||||||
|
logger.warning(
|
||||||
|
"Expert skill %r collides with an existing sub-agent name; skipping.",
|
||||||
|
name,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
taken.add(name)
|
||||||
|
subs.append(spec)
|
||||||
|
|
||||||
|
|
||||||
def _maybe_swap_async_subagents(
|
def _maybe_swap_async_subagents(
|
||||||
subs: list, middleware: list | None = None, *, cfg=None
|
subs: list,
|
||||||
|
middleware: list | None = None,
|
||||||
|
*,
|
||||||
|
tool_registry: dict | None = None,
|
||||||
|
cfg=None,
|
||||||
) -> list:
|
) -> list:
|
||||||
"""Replace ``_async``-flagged sub-agents with ``AsyncSubAgent`` specs when enabled.
|
"""Replace ``_async``-flagged sub-agents with ``AsyncSubAgent`` specs when enabled.
|
||||||
|
|
||||||
@@ -460,17 +543,24 @@ def _maybe_swap_async_subagents(
|
|||||||
Adding a new async sub-agent requires no change here — flip
|
Adding a new async sub-agent requires no change here — flip
|
||||||
``async: true`` in its yaml and create the matching deployment graph.
|
``async: true`` in its yaml and create the matching deployment graph.
|
||||||
|
|
||||||
All return paths strip the internal ``_async`` field from sub-agent dicts
|
YAML tool names stay in the internal ``_tool_names`` field until this
|
||||||
before handoff, since deepagents may schema-validate the kwarg.
|
decision point. In-process specs resolve them against ``tool_registry``;
|
||||||
|
swapped remote specs discard them because their graph factory resolves
|
||||||
|
tools in its own process. All return paths strip internal fields before
|
||||||
|
handoff, since deepagents may schema-validate the kwargs.
|
||||||
|
|
||||||
When async subagents are actually swapped in and ``middleware`` is provided,
|
When async subagents are actually swapped in and ``middleware`` is provided,
|
||||||
appends ``AsyncWatcherMiddleware`` so launches spawn an
|
appends ``AsyncWatcherMiddleware`` so launches spawn an
|
||||||
``async_notifier`` watcher.
|
``async_notifier`` watcher.
|
||||||
"""
|
"""
|
||||||
|
from .utils import resolve_subagent_tools
|
||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
cfg = cfg if cfg is not None else _ensure_config()
|
||||||
|
tool_registry = tool_registry or {}
|
||||||
if not getattr(cfg, "enable_async_subagents", False):
|
if not getattr(cfg, "enable_async_subagents", False):
|
||||||
# Async fully disabled — strip the internal flag before handoff.
|
# Async fully disabled: every spec will run in-process.
|
||||||
for s in subs:
|
for s in subs:
|
||||||
|
resolve_subagent_tools(s, tool_registry)
|
||||||
s.pop("_async", None)
|
s.pop("_async", None)
|
||||||
return subs
|
return subs
|
||||||
|
|
||||||
@@ -484,9 +574,9 @@ def _maybe_swap_async_subagents(
|
|||||||
"enable_async_subagents=true but langgraph dev is not reachable; "
|
"enable_async_subagents=true but langgraph dev is not reachable; "
|
||||||
"falling back to in-process sync delegation for all sub-agents."
|
"falling back to in-process sync delegation for all sub-agents."
|
||||||
)
|
)
|
||||||
# Strip the internal ``_async`` flag (carried from ``load_subagents``)
|
# Every spec falls back to in-process execution.
|
||||||
# before sub-agents reach deepagents — it's never a deepagents key.
|
|
||||||
for s in subs:
|
for s in subs:
|
||||||
|
resolve_subagent_tools(s, tool_registry)
|
||||||
s.pop("_async", None)
|
s.pop("_async", None)
|
||||||
return subs
|
return subs
|
||||||
|
|
||||||
@@ -498,14 +588,18 @@ def _maybe_swap_async_subagents(
|
|||||||
|
|
||||||
if not async_specs:
|
if not async_specs:
|
||||||
for s in subs:
|
for s in subs:
|
||||||
|
resolve_subagent_tools(s, tool_registry)
|
||||||
s.pop("_async", None)
|
s.pop("_async", None)
|
||||||
return subs
|
return subs
|
||||||
|
|
||||||
from deepagents import AsyncSubAgent
|
from deepagents import AsyncSubAgent
|
||||||
|
|
||||||
from .langgraph_dev.sdk import configured_langgraph_dev_url
|
from .langgraph_dev.sdk import langgraph_dev_url
|
||||||
|
|
||||||
runtime_url = configured_langgraph_dev_url()
|
# Self-dispatch target. Resolved through ``langgraph_dev_url`` so it tracks
|
||||||
|
# both ``langgraph_dev_port`` and ``langgraph_dev_host`` — a wildcard bind
|
||||||
|
# maps back to loopback, a pinned interface is honored verbatim.
|
||||||
|
dev_url = langgraph_dev_url(cfg)
|
||||||
out = []
|
out = []
|
||||||
agent_specs: dict[str, AsyncSubAgent] = {}
|
agent_specs: dict[str, AsyncSubAgent] = {}
|
||||||
# MCP tools routed to async sub-agents (via ``expose_to: <name>`` in
|
# MCP tools routed to async sub-agents (via ``expose_to: <name>`` in
|
||||||
@@ -520,19 +614,22 @@ def _maybe_swap_async_subagents(
|
|||||||
name=name,
|
name=name,
|
||||||
description=async_specs[name],
|
description=async_specs[name],
|
||||||
graph_id=name,
|
graph_id=name,
|
||||||
url=runtime_url,
|
url=dev_url,
|
||||||
)
|
)
|
||||||
agent_specs[name] = spec
|
agent_specs[name] = spec
|
||||||
out.append(spec)
|
out.append(spec)
|
||||||
else:
|
else:
|
||||||
# Strip the internal flag before handoff to deepagents.
|
resolve_subagent_tools(s, tool_registry)
|
||||||
s.pop("_async", None)
|
s.pop("_async", None)
|
||||||
out.append(s)
|
out.append(s)
|
||||||
|
|
||||||
if agent_specs and middleware is not None:
|
if agent_specs and middleware is not None:
|
||||||
|
from .cli import async_notifier
|
||||||
from .middleware.async_watcher import AsyncWatcherMiddleware
|
from .middleware.async_watcher import AsyncWatcherMiddleware
|
||||||
|
|
||||||
middleware.append(AsyncWatcherMiddleware(agent_specs))
|
# Composition root wires the concrete notifier port into the middleware;
|
||||||
|
# the middleware itself never imports the CLI layer.
|
||||||
|
middleware.append(AsyncWatcherMiddleware(agent_specs, notifier=async_notifier))
|
||||||
|
|
||||||
# Forward the CLI's live (model, provider) into deepagents'
|
# Forward the CLI's live (model, provider) into deepagents'
|
||||||
# start/update_async_task tool calls so the deployed graph can
|
# start/update_async_task tool calls so the deployed graph can
|
||||||
@@ -546,6 +643,138 @@ def _maybe_swap_async_subagents(
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def _route_async_specs_through_evo_middleware(
|
||||||
|
subs: list, base_middleware: list, *, cfg=None
|
||||||
|
) -> list:
|
||||||
|
"""Move ``AsyncSubAgent`` specs from ``subs`` into ``EvoAsyncSubAgentMiddleware``.
|
||||||
|
|
||||||
|
Deepagents' ``create_deep_agent`` auto-composes the vanilla
|
||||||
|
``AsyncSubAgentMiddleware`` when it sees ``graph_id``-carrying entries
|
||||||
|
in ``subagents=``. We need our payload-aware subclass to handle those
|
||||||
|
(see ``EvoScientist/middleware/expert_async_subagent.py`` for the
|
||||||
|
upstream-workaround rationale). To prevent the auto-composition and
|
||||||
|
route all async dispatch through our subclass, we strip AsyncSubAgent
|
||||||
|
specs from ``subs`` here and hand them to our middleware.
|
||||||
|
|
||||||
|
Also folds in ``AsyncSubAgent`` specs for installed expert skills —
|
||||||
|
all pointing at the shared ``expert-container-async`` graph, marked
|
||||||
|
``is_expert=True`` so the middleware requires a payload with
|
||||||
|
``skill_name``.
|
||||||
|
|
||||||
|
The completion watcher (``AsyncWatcherMiddleware``) is found or created
|
||||||
|
before the middleware so the middleware's resolve-on-miss start tool can
|
||||||
|
hold the watcher's agent dict by reference — see the wiring block below.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``subs`` with ``graph_id``-carrying entries removed. Safe to pass
|
||||||
|
as ``create_deep_agent(subagents=...)`` — the async-auto-compose
|
||||||
|
branch is skipped for empty async lists.
|
||||||
|
"""
|
||||||
|
from .middleware.async_watcher import AsyncWatcherMiddleware
|
||||||
|
from .middleware.expert_async_subagent import EvoAsyncSubAgentMiddleware
|
||||||
|
from .subagents.expert_container_async import build_expert_async_subagent_specs
|
||||||
|
|
||||||
|
cfg = cfg if cfg is not None else _ensure_config()
|
||||||
|
|
||||||
|
async_specs = [s for s in subs if "graph_id" in s]
|
||||||
|
sync_subs = [s for s in subs if "graph_id" not in s]
|
||||||
|
expert_specs = build_expert_async_subagent_specs(cfg=cfg)
|
||||||
|
async_specs.extend(expert_specs)
|
||||||
|
|
||||||
|
# Find or create the completion watcher BEFORE constructing
|
||||||
|
# ``EvoAsyncSubAgentMiddleware``: the middleware's resolve-on-miss start
|
||||||
|
# tool must hold the watcher's agent dict by reference, so an expert
|
||||||
|
# discovered mid-session lands in the dispatch table and the watcher in
|
||||||
|
# one step. Without the watcher update, dispatch succeeds but the
|
||||||
|
# watcher's ``get_async(agent_name)`` raises KeyError inside its
|
||||||
|
# ``try/except`` and the completion notification silently never fires.
|
||||||
|
watcher_agents: dict | None = None
|
||||||
|
watcher = next(
|
||||||
|
(m for m in base_middleware if isinstance(m, AsyncWatcherMiddleware)),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if watcher is None:
|
||||||
|
# No YAML async subagents were registered, so ``_maybe_swap`` did
|
||||||
|
# not install the watcher. Install it now so experts still get
|
||||||
|
# completion notifications. Appended before the middleware's
|
||||||
|
# index-0 insert below, which yields the same final order as
|
||||||
|
# append-after-insert: ``[EvoAsync..., ..., watcher]``.
|
||||||
|
if expert_specs:
|
||||||
|
from .cli import async_notifier
|
||||||
|
|
||||||
|
watcher = AsyncWatcherMiddleware(
|
||||||
|
{s["name"]: s for s in expert_specs},
|
||||||
|
notifier=async_notifier,
|
||||||
|
)
|
||||||
|
base_middleware.append(watcher)
|
||||||
|
elif expert_specs:
|
||||||
|
# Extend AsyncWatcherMiddleware's client cache with expert specs so
|
||||||
|
# start_async_task launches for experts spawn a completion watcher —
|
||||||
|
# otherwise the watcher's ``get_async(agent_name)`` KeyErrors on the
|
||||||
|
# expert name, no notification is enqueued, and the main agent never
|
||||||
|
# learns the task finished. ``_maybe_swap_async_subagents`` above only
|
||||||
|
# populates the watcher with YAML-defined async subagents
|
||||||
|
# (writing-agent, data-analysis-agent, scheduler); this hook folds in
|
||||||
|
# the experts too.
|
||||||
|
#
|
||||||
|
# The mutation reaches through two layers of private state:
|
||||||
|
# ``AsyncWatcherMiddleware._clients`` (our own) and
|
||||||
|
# ``_ClientCache._agents`` (upstream deepagents). If upstream ever
|
||||||
|
# renames ``_agents`` or wraps it in an immutable snapshot, the
|
||||||
|
# ``.update(...)`` below silently lands on nothing — expert
|
||||||
|
# completion nudges then stop firing without a diagnostic surface.
|
||||||
|
# Convert that silent-drop into a grep-able error line and leave
|
||||||
|
# ``watcher_agents`` unset; expert dispatches still work (without
|
||||||
|
# completion notifications and without mid-session resolution into
|
||||||
|
# the watcher) until upstream drift is fixed.
|
||||||
|
if not hasattr(watcher._clients, "_agents"):
|
||||||
|
logging.getLogger(__name__).error(
|
||||||
|
"AsyncWatcherMiddleware._clients has no `_agents` slot — "
|
||||||
|
"deepagents internal renamed; expert completion "
|
||||||
|
"notifications will not fire until the extension hook is "
|
||||||
|
"updated to the new attribute name."
|
||||||
|
)
|
||||||
|
watcher = None
|
||||||
|
else:
|
||||||
|
watcher._clients._agents.update({s["name"]: s for s in expert_specs})
|
||||||
|
# The second ``hasattr`` is not redundant with the one in the ``elif``
|
||||||
|
# above: that check only runs when ``expert_specs`` is non-empty. When a
|
||||||
|
# pre-existing watcher has an empty expert set, this is the only guard
|
||||||
|
# standing between an upstream rename of ``_ClientCache._agents`` and an
|
||||||
|
# AttributeError that would kill agent construction — without it, the
|
||||||
|
# drift degrades to "no completion nudges" instead of crashing.
|
||||||
|
if watcher is not None and hasattr(watcher._clients, "_agents"):
|
||||||
|
watcher_agents = watcher._clients._agents
|
||||||
|
|
||||||
|
if async_specs:
|
||||||
|
# ``_maybe_swap_async_subagents`` installs the model-passthrough patch
|
||||||
|
# only when the yaml-async spec list is non-empty. An expert-only setup
|
||||||
|
# (no ``writing-agent`` / ``data-analysis-agent`` / ``scheduler`` in
|
||||||
|
# yaml) would otherwise miss the patch entirely, so we install it here
|
||||||
|
# too. Idempotent — the shared ``_model_passthrough_patched`` flag
|
||||||
|
# guards against double-patching.
|
||||||
|
from .llm.patches import _patch_deepagents_model_passthrough
|
||||||
|
|
||||||
|
_patch_deepagents_model_passthrough()
|
||||||
|
|
||||||
|
# Prepend rather than append so the ``## Async subagents`` prompt
|
||||||
|
# section stays in the stable prefix. Appending pushes it past the
|
||||||
|
# volatile memory tail, invalidating the cached prefix on every
|
||||||
|
# memory change.
|
||||||
|
base_middleware.insert(
|
||||||
|
0,
|
||||||
|
EvoAsyncSubAgentMiddleware(
|
||||||
|
async_subagents=async_specs,
|
||||||
|
watcher_agents=watcher_agents,
|
||||||
|
# The construction cfg, so resolve-on-miss specs the same
|
||||||
|
# langgraph_dev_port the construction-time specs used instead
|
||||||
|
# of re-reading config from disk at dispatch time.
|
||||||
|
cfg=cfg,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return sync_subs
|
||||||
|
|
||||||
|
|
||||||
def _build_base_kwargs(
|
def _build_base_kwargs(
|
||||||
base_backend, base_middleware, *, cfg=None, chat_model=None, workspace_dir=None
|
base_backend, base_middleware, *, cfg=None, chat_model=None, workspace_dir=None
|
||||||
):
|
):
|
||||||
@@ -555,19 +784,29 @@ def _build_base_kwargs(
|
|||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
cfg = cfg if cfg is not None else _ensure_config()
|
||||||
tool_registry = {"think_tool": think_tool}
|
tool_registry = {"think_tool": think_tool}
|
||||||
|
base_tools = [think_tool, skill_manager]
|
||||||
if os.environ.get("TAVILY_API_KEY"):
|
if os.environ.get("TAVILY_API_KEY"):
|
||||||
tool_registry["tavily_search"] = tavily_search
|
tool_registry["tavily_search"] = tavily_search
|
||||||
base_tools = [think_tool, skill_manager]
|
base_tools.append(tavily_search)
|
||||||
|
|
||||||
subs = load_subagents(
|
subs = load_subagents(
|
||||||
SUBAGENTS_CONFIG,
|
SUBAGENTS_CONFIG,
|
||||||
tool_registry=tool_registry,
|
|
||||||
)
|
)
|
||||||
|
_fold_expert_subagents(subs, tool_registry)
|
||||||
_ensure_general_purpose_subagent(subs)
|
_ensure_general_purpose_subagent(subs)
|
||||||
_inject_subagent_middleware(
|
_inject_subagent_middleware(
|
||||||
subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model
|
subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model
|
||||||
)
|
)
|
||||||
subs = _maybe_swap_async_subagents(subs, base_middleware, cfg=cfg)
|
subs = _maybe_swap_async_subagents(
|
||||||
|
subs,
|
||||||
|
base_middleware,
|
||||||
|
tool_registry=tool_registry,
|
||||||
|
cfg=cfg,
|
||||||
|
)
|
||||||
|
# Route AsyncSubAgent specs (both standard and expert) through
|
||||||
|
# EvoAsyncSubAgentMiddleware so the payload-aware start_async_task tool
|
||||||
|
# replaces upstream's non-parameterisable one.
|
||||||
|
subs = _route_async_specs_through_evo_middleware(subs, base_middleware, cfg=cfg)
|
||||||
return {
|
return {
|
||||||
"name": "EvoScientist",
|
"name": "EvoScientist",
|
||||||
"model": chat_model if chat_model is not None else _ensure_chat_model(),
|
"model": chat_model if chat_model is not None else _ensure_chat_model(),
|
||||||
@@ -588,6 +827,7 @@ def load_mcp_and_build_kwargs(
|
|||||||
cfg=None,
|
cfg=None,
|
||||||
chat_model=None,
|
chat_model=None,
|
||||||
workspace_dir=None,
|
workspace_dir=None,
|
||||||
|
runtime: "AsyncRuntime | None" = None,
|
||||||
):
|
):
|
||||||
"""Load MCP tools (cached by config) and build agent kwargs.
|
"""Load MCP tools (cached by config) and build agent kwargs.
|
||||||
|
|
||||||
@@ -606,7 +846,10 @@ def load_mcp_and_build_kwargs(
|
|||||||
from .utils import load_subagents
|
from .utils import load_subagents
|
||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
cfg = cfg if cfg is not None else _ensure_config()
|
||||||
mcp_by_agent = _load_mcp_tools_cached(on_progress=on_mcp_progress)
|
mcp_by_agent = _load_mcp_tools_cached(
|
||||||
|
on_progress=on_mcp_progress,
|
||||||
|
runtime=runtime,
|
||||||
|
)
|
||||||
if not mcp_by_agent:
|
if not mcp_by_agent:
|
||||||
return _build_base_kwargs(
|
return _build_base_kwargs(
|
||||||
base_backend,
|
base_backend,
|
||||||
@@ -617,36 +860,105 @@ def load_mcp_and_build_kwargs(
|
|||||||
)
|
)
|
||||||
|
|
||||||
tool_registry = {"think_tool": think_tool}
|
tool_registry = {"think_tool": think_tool}
|
||||||
|
base_tools = [think_tool, skill_manager]
|
||||||
if os.environ.get("TAVILY_API_KEY"):
|
if os.environ.get("TAVILY_API_KEY"):
|
||||||
tool_registry["tavily_search"] = tavily_search
|
tool_registry["tavily_search"] = tavily_search
|
||||||
base_tools = [think_tool, skill_manager]
|
base_tools.append(tavily_search)
|
||||||
|
|
||||||
# Fresh tool registry — start from base tools + MCP tools
|
# DeepAgents installs these outside ``base_tools`` through middleware.
|
||||||
|
# MCP tools must never shadow them inside any one agent namespace.
|
||||||
|
middleware_tool_names = {
|
||||||
|
"ls",
|
||||||
|
"read_file",
|
||||||
|
"write_file",
|
||||||
|
"edit_file",
|
||||||
|
"glob",
|
||||||
|
"grep",
|
||||||
|
"execute",
|
||||||
|
"write_todos",
|
||||||
|
"task",
|
||||||
|
"start_async_task",
|
||||||
|
"check_async_task",
|
||||||
|
"update_async_task",
|
||||||
|
"cancel_async_task",
|
||||||
|
"list_async_tasks",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Fresh tool registry — start from built-ins, then add one representative
|
||||||
|
# MCP implementation for YAML name resolution. A tool may be exposed to
|
||||||
|
# several agents, but no agent may contain duplicate names and no MCP tool
|
||||||
|
# may override a built-in implementation.
|
||||||
registry = dict(tool_registry)
|
registry = dict(tool_registry)
|
||||||
for tools in mcp_by_agent.values():
|
builtin_names = {
|
||||||
|
*(str(getattr(tool, "name", "")) for tool in base_tools),
|
||||||
|
*middleware_tool_names,
|
||||||
|
}
|
||||||
|
for agent_name, tools in mcp_by_agent.items():
|
||||||
|
seen_for_agent: set[str] = set()
|
||||||
for t in tools:
|
for t in tools:
|
||||||
registry[t.name] = t
|
tool_name = str(t.name)
|
||||||
|
if tool_name in builtin_names or tool_name in seen_for_agent:
|
||||||
|
from .llm.contracts import EvoRuntimeError
|
||||||
|
|
||||||
|
raise EvoRuntimeError(
|
||||||
|
"TOOL_REGISTRY_CONFLICT",
|
||||||
|
details=(
|
||||||
|
{"agent_name": str(agent_name), "tool_name": tool_name},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
seen_for_agent.add(tool_name)
|
||||||
|
registry.setdefault(tool_name, t)
|
||||||
|
|
||||||
mcp_main = mcp_by_agent.pop("main", [])
|
mcp_main = mcp_by_agent.pop("main", [])
|
||||||
|
|
||||||
subs = load_subagents(
|
subs = load_subagents(
|
||||||
SUBAGENTS_CONFIG,
|
SUBAGENTS_CONFIG,
|
||||||
tool_registry=registry,
|
|
||||||
)
|
)
|
||||||
|
_fold_expert_subagents(subs, registry)
|
||||||
|
|
||||||
_ensure_general_purpose_subagent(subs)
|
_ensure_general_purpose_subagent(subs)
|
||||||
_inject_subagent_middleware(
|
_inject_subagent_middleware(
|
||||||
subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model
|
subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model
|
||||||
)
|
)
|
||||||
|
|
||||||
# Inject MCP tools into subagents by name
|
# Inject MCP tools into subagents by name. YAML-resolved tools already
|
||||||
|
# belong to that agent namespace, so a second tool with the same name is a
|
||||||
|
# configuration conflict rather than an item to append silently.
|
||||||
for sa in subs:
|
for sa in subs:
|
||||||
if sa_tools := mcp_by_agent.get(sa["name"], []):
|
if sa_tools := mcp_by_agent.get(sa["name"], []):
|
||||||
sa.setdefault("tools", []).extend(sa_tools)
|
target_tools = sa.setdefault("tools", [])
|
||||||
|
existing_names = {
|
||||||
|
str(getattr(tool, "name", tool)) for tool in target_tools
|
||||||
|
}
|
||||||
|
for tool in sa_tools:
|
||||||
|
tool_name = str(tool.name)
|
||||||
|
if tool_name in existing_names:
|
||||||
|
from .llm.contracts import EvoRuntimeError
|
||||||
|
|
||||||
|
raise EvoRuntimeError(
|
||||||
|
"TOOL_REGISTRY_CONFLICT",
|
||||||
|
details=(
|
||||||
|
{
|
||||||
|
"agent_name": str(sa["name"]),
|
||||||
|
"tool_name": tool_name,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
existing_names.add(tool_name)
|
||||||
|
target_tools.append(tool)
|
||||||
|
|
||||||
# Swap selected sub-agents to AsyncSubAgent (must happen AFTER MCP injection
|
# Swap selected sub-agents to AsyncSubAgent (must happen AFTER MCP injection
|
||||||
# since async sub-agents are remote graphs that load their own tools).
|
# since async sub-agents are remote graphs that load their own tools).
|
||||||
subs = _maybe_swap_async_subagents(subs, base_middleware, cfg=cfg)
|
subs = _maybe_swap_async_subagents(
|
||||||
|
subs,
|
||||||
|
base_middleware,
|
||||||
|
tool_registry=registry,
|
||||||
|
cfg=cfg,
|
||||||
|
)
|
||||||
|
# Mirror the base path: route AsyncSubAgent specs through
|
||||||
|
# EvoAsyncSubAgentMiddleware so the payload-aware start_async_task tool
|
||||||
|
# is the one composed into the main agent.
|
||||||
|
subs = _route_async_specs_through_evo_middleware(subs, base_middleware, cfg=cfg)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"name": "EvoScientist",
|
"name": "EvoScientist",
|
||||||
@@ -665,8 +977,19 @@ def load_mcp_and_build_kwargs(
|
|||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
def _get_legacy_backend():
|
def _get_legacy_backend(
|
||||||
"""Build the deployment-root backend used outside Web full deploy."""
|
*, guard_dangerous: bool | None = None, refuse_delete: bool = False
|
||||||
|
):
|
||||||
|
"""Build the deployment-root backend used outside Web full deploy.
|
||||||
|
|
||||||
|
``guard_dangerous`` — when ``None`` (default) follows ``cfg.auto_approve``;
|
||||||
|
the two research async sub-agent graphs (``writing-agent`` /
|
||||||
|
``data-analysis-agent``) pass ``True`` because their remote thread has no
|
||||||
|
approval path at all (see ``subagents/_factory._GUARDED_ASYNC_SUBAGENTS``).
|
||||||
|
``refuse_delete`` — the same two async graphs pass ``True`` so the recursive
|
||||||
|
``delete`` FS tool is refused and relayed to the orchestrator for approval,
|
||||||
|
rather than deleting unattended.
|
||||||
|
"""
|
||||||
from deepagents.backends import CompositeBackend
|
from deepagents.backends import CompositeBackend
|
||||||
|
|
||||||
from .backends import (
|
from .backends import (
|
||||||
@@ -676,6 +999,8 @@ def _get_legacy_backend():
|
|||||||
)
|
)
|
||||||
|
|
||||||
cfg = _ensure_config()
|
cfg = _ensure_config()
|
||||||
|
if guard_dangerous is None:
|
||||||
|
guard_dangerous = cfg.auto_approve
|
||||||
workspace_dir = str(_paths_mod.WORKSPACE_ROOT)
|
workspace_dir = str(_paths_mod.WORKSPACE_ROOT)
|
||||||
set_active_workspace(workspace_dir)
|
set_active_workspace(workspace_dir)
|
||||||
memory_dir = str(_paths_mod.MEMORIES_DIR)
|
memory_dir = str(_paths_mod.MEMORIES_DIR)
|
||||||
@@ -689,6 +1014,8 @@ def _get_legacy_backend():
|
|||||||
virtual_mode=True,
|
virtual_mode=True,
|
||||||
timeout=cfg.sandbox_execute_timeout,
|
timeout=cfg.sandbox_execute_timeout,
|
||||||
dangerous=cfg.dangerous_mode,
|
dangerous=cfg.dangerous_mode,
|
||||||
|
guard_dangerous=guard_dangerous,
|
||||||
|
refuse_delete=refuse_delete,
|
||||||
)
|
)
|
||||||
sk_backend = MergedSkillsBackend(
|
sk_backend = MergedSkillsBackend(
|
||||||
primary_dir=user_skills_dir,
|
primary_dir=user_skills_dir,
|
||||||
@@ -708,14 +1035,20 @@ def _get_legacy_backend():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _get_default_backend():
|
def _get_default_backend(
|
||||||
|
*, guard_dangerous: bool | None = None, refuse_delete: bool = False
|
||||||
|
):
|
||||||
"""Use Origin's conversation-scoped backend for Web full deploy."""
|
"""Use Origin's conversation-scoped backend for Web full deploy."""
|
||||||
if os.environ.get("EVOSCIENTIST_DEPLOY_MODE", "").lower() != "full":
|
if os.environ.get("EVOSCIENTIST_DEPLOY_MODE", "").lower() != "full":
|
||||||
return _get_legacy_backend()
|
return _get_legacy_backend(
|
||||||
from .workspace_scope import create_workspace_backend_factory
|
guard_dangerous=guard_dangerous, refuse_delete=refuse_delete
|
||||||
|
)
|
||||||
|
from .workspace_scope import create_deferred_scoped_backend
|
||||||
|
|
||||||
cfg = _ensure_config()
|
cfg = _ensure_config()
|
||||||
return create_workspace_backend_factory(
|
# deepagents 0.7 removed backend factories, so hand the middleware an
|
||||||
|
# instance that resolves this run's scope from the runnable config.
|
||||||
|
return create_deferred_scoped_backend(
|
||||||
_get_legacy_backend,
|
_get_legacy_backend,
|
||||||
dangerous=cfg.dangerous_mode,
|
dangerous=cfg.dangerous_mode,
|
||||||
allow_unscoped_legacy=False,
|
allow_unscoped_legacy=False,
|
||||||
@@ -739,6 +1072,7 @@ def _get_default_middleware(
|
|||||||
enable_scheduler: bool | None = None,
|
enable_scheduler: bool | None = None,
|
||||||
enable_memory_workers: bool | None = None,
|
enable_memory_workers: bool | None = None,
|
||||||
install_subagent_guard: bool = False,
|
install_subagent_guard: bool = False,
|
||||||
|
events: "MiddlewareEventSink | None" = None,
|
||||||
):
|
):
|
||||||
"""Build the default middleware list.
|
"""Build the default middleware list.
|
||||||
|
|
||||||
@@ -758,6 +1092,11 @@ def _get_default_middleware(
|
|||||||
(avoids writing module globals on the pure path).
|
(avoids writing module globals on the pure path).
|
||||||
memory_source_agent: Attribution name for profile/observation writes.
|
memory_source_agent: Attribution name for profile/observation writes.
|
||||||
Async sub-agent factories pass their deployed agent name here.
|
Async sub-agent factories pass their deployed agent name here.
|
||||||
|
events: Frontend/session-supplied event sink. Middleware report
|
||||||
|
tool-selection events and model-fallback notices to it.
|
||||||
|
Defaults to the current stream run's sink for main agents; async
|
||||||
|
sub-agent stacks are always forced to ``NoOpSink`` (they must not
|
||||||
|
drive the main-agent widgets).
|
||||||
"""
|
"""
|
||||||
from .middleware import (
|
from .middleware import (
|
||||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||||
@@ -770,7 +1109,9 @@ def _get_default_middleware(
|
|||||||
RecoverableToolEffectMiddleware,
|
RecoverableToolEffectMiddleware,
|
||||||
RepetitiveToolCallGuardMiddleware,
|
RepetitiveToolCallGuardMiddleware,
|
||||||
ToolErrorHandlerMiddleware,
|
ToolErrorHandlerMiddleware,
|
||||||
|
ToolHistoryRepairMiddleware,
|
||||||
ToolProtocolGuardMiddleware,
|
ToolProtocolGuardMiddleware,
|
||||||
|
create_active_team_middleware,
|
||||||
create_code_interpreter_middleware,
|
create_code_interpreter_middleware,
|
||||||
create_context_editing_middleware,
|
create_context_editing_middleware,
|
||||||
create_memory_lifecycle_middleware,
|
create_memory_lifecycle_middleware,
|
||||||
@@ -781,6 +1122,13 @@ def _get_default_middleware(
|
|||||||
default_memory_scheduler,
|
default_memory_scheduler,
|
||||||
load_fallback_chain,
|
load_fallback_chain,
|
||||||
)
|
)
|
||||||
|
from .middleware.events import NO_OP_SINK, RunScopedEventSink
|
||||||
|
|
||||||
|
# Subagent stacks never drive the main-agent frontend widgets; force the
|
||||||
|
# no-op sink there regardless of what the caller passed. Main stacks built
|
||||||
|
# without an explicit frontend/session sink report into the active stream
|
||||||
|
# run's sink, preserving selector suppression for headless local runs.
|
||||||
|
events = NO_OP_SINK if for_async_subagent else (events or RunScopedEventSink())
|
||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
cfg = cfg if cfg is not None else _ensure_config()
|
||||||
repetitive_tool_call_threshold = getattr(
|
repetitive_tool_call_threshold = getattr(
|
||||||
@@ -821,6 +1169,8 @@ def _get_default_middleware(
|
|||||||
MemoryObservationTarget.AGENT
|
MemoryObservationTarget.AGENT
|
||||||
),
|
),
|
||||||
"memory_scheduler": memory_scheduler,
|
"memory_scheduler": memory_scheduler,
|
||||||
|
# First-contact intro: main agent only, and never in unattended runs.
|
||||||
|
"enable_profile_bootstrap": not for_async_subagent and not bool(cfg.auto_mode),
|
||||||
}
|
}
|
||||||
if memory_max_inline_profile_chars is not None:
|
if memory_max_inline_profile_chars is not None:
|
||||||
memory_kwargs["max_inline_profile_chars"] = memory_max_inline_profile_chars
|
memory_kwargs["max_inline_profile_chars"] = memory_max_inline_profile_chars
|
||||||
@@ -853,7 +1203,10 @@ def _get_default_middleware(
|
|||||||
else {}
|
else {}
|
||||||
),
|
),
|
||||||
model=resolved_selector_model,
|
model=resolved_selector_model,
|
||||||
track_stream_selection=not for_async_subagent,
|
# A frontend sink enables the streaming selection lifecycle; subagent /
|
||||||
|
# headless stacks stay silent (upstream replaced track_stream_selection
|
||||||
|
# with the events sink).
|
||||||
|
events=events,
|
||||||
)
|
)
|
||||||
mw = [
|
mw = [
|
||||||
# Outermost — catches provider-SDK exceptions from the model
|
# Outermost — catches provider-SDK exceptions from the model
|
||||||
@@ -863,14 +1216,24 @@ def _get_default_middleware(
|
|||||||
ErrorNormalizationMiddleware(),
|
ErrorNormalizationMiddleware(),
|
||||||
RecoverableMeteringMiddleware(),
|
RecoverableMeteringMiddleware(),
|
||||||
RecoverableToolEffectMiddleware(),
|
RecoverableToolEffectMiddleware(),
|
||||||
|
ToolHistoryRepairMiddleware(),
|
||||||
create_context_editing_middleware(model),
|
create_context_editing_middleware(model),
|
||||||
*([ModelFallbackMiddleware()] if enable_legacy_model_fallback else []),
|
*(
|
||||||
|
[ModelFallbackMiddleware(events=events)]
|
||||||
|
if enable_legacy_model_fallback
|
||||||
|
else []
|
||||||
|
),
|
||||||
RepetitiveToolCallGuardMiddleware(
|
RepetitiveToolCallGuardMiddleware(
|
||||||
threshold=repetitive_tool_call_threshold,
|
threshold=repetitive_tool_call_threshold,
|
||||||
max_consecutive_errors=max_consecutive_tool_errors,
|
max_consecutive_errors=max_consecutive_tool_errors,
|
||||||
),
|
),
|
||||||
ContextOverflowMapperMiddleware(),
|
ContextOverflowMapperMiddleware(),
|
||||||
ToolErrorHandlerMiddleware(),
|
ToolErrorHandlerMiddleware(),
|
||||||
|
# deepagents 0.7.0 dropped TodoListMiddleware from its defaults;
|
||||||
|
# EXPERIMENT_WORKFLOW planning and the todo UI pipeline require it.
|
||||||
|
# deepagents 0.6.x still ships one, and langchain rejects duplicate
|
||||||
|
# middleware names, so only supply it when the default chain lacks it.
|
||||||
|
*(() if _deepagents_provides_todo_list() else (TodoListMiddleware(),)),
|
||||||
*selector_middlewares,
|
*selector_middlewares,
|
||||||
ToolProtocolGuardMiddleware(),
|
ToolProtocolGuardMiddleware(),
|
||||||
# Interpreter prompt must land before runtime/memory context, so this
|
# Interpreter prompt must land before runtime/memory context, so this
|
||||||
@@ -908,13 +1271,32 @@ def _get_default_middleware(
|
|||||||
|
|
||||||
mw.insert(0, AskUserMiddleware())
|
mw.insert(0, AskUserMiddleware())
|
||||||
|
|
||||||
|
# Expert prompt for the main agent — injects the ## Experts concept every
|
||||||
|
# turn (plus the invited-expert list when experts are invited). Inserted
|
||||||
|
# AFTER AskUser so it sits ahead of AskUser in the stack and runs first,
|
||||||
|
# landing its block right after ## Skills System (experts mirror skills).
|
||||||
|
# Main agent only: a running expert graph must not inject the expert prompt
|
||||||
|
# into its own baked-in persona.
|
||||||
|
if not for_async_subagent:
|
||||||
|
mw.insert(0, create_active_team_middleware())
|
||||||
|
|
||||||
# Background-process tools (run_in_background / check_process / stop_process /
|
# Background-process tools (run_in_background / check_process / stop_process /
|
||||||
# list_processes) — main agent only. Async sub-agents run on langgraph-dev and
|
# list_processes) — main agent only. Async sub-agents run on langgraph-dev and
|
||||||
# must not spawn local OS processes.
|
# must not spawn local OS processes.
|
||||||
if not for_async_subagent and enable_background_execution:
|
if not for_async_subagent and enable_background_execution:
|
||||||
|
from .cli import async_notifier
|
||||||
from .middleware.background import BackgroundExecutionMiddleware
|
from .middleware.background import BackgroundExecutionMiddleware
|
||||||
|
|
||||||
mw.append(BackgroundExecutionMiddleware())
|
# Inject the notifier port + the assembly-time dangerous-mode policy
|
||||||
|
# (agents rebuild on config change, so the captured flag never staler
|
||||||
|
# than the agent it lives on).
|
||||||
|
mw.append(
|
||||||
|
BackgroundExecutionMiddleware(
|
||||||
|
async_notifier,
|
||||||
|
dangerous=cfg.dangerous_mode,
|
||||||
|
guard_dangerous=cfg.auto_approve,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
if install_subagent_guard:
|
if install_subagent_guard:
|
||||||
mw.append(DisableSubagentToolMiddleware())
|
mw.append(DisableSubagentToolMiddleware())
|
||||||
@@ -922,6 +1304,19 @@ def _get_default_middleware(
|
|||||||
return mw
|
return mw
|
||||||
|
|
||||||
|
|
||||||
|
def _build_hitl_interrupt_on(*, auto_approve: bool) -> dict[str, bool] | None:
|
||||||
|
"""Return :data:`HITL_INTERRUPT_ON` for ``create_deep_agent``, or ``None``
|
||||||
|
when the user opted out (``auto_approve`` / ``auto_mode`` /
|
||||||
|
``dangerous_mode``) so nothing is armed and unattended runs never pause.
|
||||||
|
Passing it to ``create_deep_agent`` (not ``HumanInTheLoopMiddleware``) lets
|
||||||
|
declarative sub-agents inherit it while ``AsyncSubAgent`` specs do not — so
|
||||||
|
async agents can't hang on an approval nobody can deliver.
|
||||||
|
"""
|
||||||
|
if auto_approve:
|
||||||
|
return None
|
||||||
|
return dict(HITL_INTERRUPT_ON)
|
||||||
|
|
||||||
|
|
||||||
def _get_default_agent():
|
def _get_default_agent():
|
||||||
"""Build the default agent (no checkpointer) on first access.
|
"""Build the default agent (no checkpointer) on first access.
|
||||||
|
|
||||||
@@ -972,20 +1367,18 @@ def _get_default_agent():
|
|||||||
else _get_default_middleware()
|
else _get_default_middleware()
|
||||||
)
|
)
|
||||||
|
|
||||||
# HITL on main agent only (mirrors create_cli_agent). Use middleware,
|
# HITL on main agent only (mirrors create_cli_agent). Arm it through the
|
||||||
# not interrupt_on= kwarg — the kwarg propagates to every subagent and
|
# review middleware and NEVER through `create_deep_agent(interrupt_on=)`:
|
||||||
# breaks parallel execute calls (multi-pending-interrupt LangGraph
|
# that kwarg makes deepagents append its own plain
|
||||||
# error). See PR #202.
|
# HumanInTheLoopMiddleware, and the second, auto-blind layer interrupts
|
||||||
|
# even on a run the gateway verified as auto — silently disabling
|
||||||
|
# automatic approval. It also propagates to every subagent and breaks
|
||||||
|
# parallel execute calls (multi-pending-interrupt LangGraph error).
|
||||||
|
# See PR #202. `HITL_INTERRUPT_ON` is the single source of the tool set.
|
||||||
from .middleware import DynamicReviewMiddleware
|
from .middleware import DynamicReviewMiddleware
|
||||||
|
|
||||||
mw.append(
|
mw.append(
|
||||||
DynamicReviewMiddleware(
|
DynamicReviewMiddleware(interrupt_on=dict(HITL_INTERRUPT_ON))
|
||||||
interrupt_on={
|
|
||||||
"execute": True,
|
|
||||||
"run_in_background": True,
|
|
||||||
"schedule_task": True,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if web_full:
|
if web_full:
|
||||||
@@ -1019,6 +1412,8 @@ def _get_default_agent():
|
|||||||
)
|
)
|
||||||
kwargs = _apply_budgeted_skill_context(kwargs, be)
|
kwargs = _apply_budgeted_skill_context(kwargs, be)
|
||||||
|
|
||||||
|
# No `interrupt_on=` here: HITL is armed above, as a single layer, by the
|
||||||
|
# auto-aware review middleware (see the comment there).
|
||||||
_EvoScientist_agent = create_deep_agent(
|
_EvoScientist_agent = create_deep_agent(
|
||||||
**kwargs,
|
**kwargs,
|
||||||
).with_config({"recursion_limit": cfg.recursion_limit})
|
).with_config({"recursion_limit": cfg.recursion_limit})
|
||||||
@@ -1043,6 +1438,30 @@ def __getattr__(name: str):
|
|||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
def _create_run_summarization_middleware(model, backend, summarizer):
|
||||||
|
"""Resolve context thresholds against the consuming model, not the summarizer."""
|
||||||
|
from deepagents.middleware.summarization import (
|
||||||
|
SummarizationMiddleware, compute_summarization_defaults,
|
||||||
|
)
|
||||||
|
|
||||||
|
defaults = compute_summarization_defaults(model)
|
||||||
|
window = (model.profile or {}).get("max_input_tokens")
|
||||||
|
|
||||||
|
def absolute(value):
|
||||||
|
if isinstance(value, tuple) and value[0] == "fraction":
|
||||||
|
if not isinstance(window, int) or window <= 0:
|
||||||
|
raise ValueError("fraction threshold requires a main-model context window")
|
||||||
|
return ("tokens", int(window * value[1]))
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return {key: absolute(item) for key, item in value.items()}
|
||||||
|
return value
|
||||||
|
|
||||||
|
return SummarizationMiddleware(
|
||||||
|
model=summarizer, backend=backend,
|
||||||
|
trim_tokens_to_summarize=None, **absolute(defaults),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def create_cli_agent(
|
def create_cli_agent(
|
||||||
workspace_dir: str | None = None,
|
workspace_dir: str | None = None,
|
||||||
checkpointer=None,
|
checkpointer=None,
|
||||||
@@ -1060,6 +1479,8 @@ def create_cli_agent(
|
|||||||
main_agent_route_middleware: AgentMiddleware | None = None,
|
main_agent_route_middleware: AgentMiddleware | None = None,
|
||||||
execution_profile=None,
|
execution_profile=None,
|
||||||
agent_model_set=None,
|
agent_model_set=None,
|
||||||
|
events: "MiddlewareEventSink | None" = None,
|
||||||
|
runtime: "AsyncRuntime | None" = None,
|
||||||
) -> "CompiledStateGraph":
|
) -> "CompiledStateGraph":
|
||||||
"""Create agent with checkpointer for CLI multi-turn support.
|
"""Create agent with checkpointer for CLI multi-turn support.
|
||||||
|
|
||||||
@@ -1102,6 +1523,8 @@ def create_cli_agent(
|
|||||||
after ConfigurableModelMiddleware and before tool selection. When
|
after ConfigurableModelMiddleware and before tool selection. When
|
||||||
provided, EvoScientist's legacy model fallback is disabled for the
|
provided, EvoScientist's legacy model fallback is disabled for the
|
||||||
top-level agent so the host is the only fallback authority.
|
top-level agent so the host is the only fallback authority.
|
||||||
|
runtime: Optional application-scoped runtime for synchronous MCP tool
|
||||||
|
discovery. Direct callers get a scoped runtime when omitted.
|
||||||
"""
|
"""
|
||||||
import os as _os
|
import os as _os
|
||||||
|
|
||||||
@@ -1119,9 +1542,15 @@ def create_cli_agent(
|
|||||||
# locals and write no module globals. Otherwise keep the legacy
|
# locals and write no module globals. Otherwise keep the legacy
|
||||||
# global-writing behavior — callers that pass config= only (CLI startup,
|
# global-writing behavior — callers that pass config= only (CLI startup,
|
||||||
# langgraph dev) rely on it to seat the active config/model.
|
# langgraph dev) rely on it to seat the active config/model.
|
||||||
|
is_web = getattr(execution_profile, "name", "") in {"web_v1", "web_v3"}
|
||||||
|
if is_web and any(value is None for value in (
|
||||||
|
config, chat_model, workspace_dir, memory_dir, workspace_backend, checkpointer
|
||||||
|
)):
|
||||||
|
raise ValueError("Web execution requires explicit config/model/workspace/memory/backend/checkpointer")
|
||||||
if config is not None and chat_model is not None:
|
if config is not None and chat_model is not None:
|
||||||
cfg = config
|
cfg = config
|
||||||
_apply_env_from_config(cfg)
|
if not is_web:
|
||||||
|
_apply_env_from_config(cfg)
|
||||||
else:
|
else:
|
||||||
cfg = _ensure_config(config)
|
cfg = _ensure_config(config)
|
||||||
chat_model = None
|
chat_model = None
|
||||||
@@ -1167,7 +1596,8 @@ def create_cli_agent(
|
|||||||
|
|
||||||
# Always construct fresh backends from current paths (avoids stale
|
# Always construct fresh backends from current paths (avoids stale
|
||||||
# module-level backend when workspace root changed at runtime).
|
# module-level backend when workspace root changed at runtime).
|
||||||
set_active_workspace(workspace_dir)
|
if not is_web:
|
||||||
|
set_active_workspace(workspace_dir)
|
||||||
ws_backend = workspace_backend
|
ws_backend = workspace_backend
|
||||||
if ws_backend is None:
|
if ws_backend is None:
|
||||||
ws_backend = CustomSandboxBackend(
|
ws_backend = CustomSandboxBackend(
|
||||||
@@ -1175,6 +1605,7 @@ def create_cli_agent(
|
|||||||
virtual_mode=True,
|
virtual_mode=True,
|
||||||
timeout=cfg.sandbox_execute_timeout,
|
timeout=cfg.sandbox_execute_timeout,
|
||||||
dangerous=cfg.dangerous_mode,
|
dangerous=cfg.dangerous_mode,
|
||||||
|
guard_dangerous=cfg.auto_approve,
|
||||||
)
|
)
|
||||||
sk_backend = MergedSkillsBackend(
|
sk_backend = MergedSkillsBackend(
|
||||||
primary_dir=_usr_skills_dir,
|
primary_dir=_usr_skills_dir,
|
||||||
@@ -1187,6 +1618,7 @@ def create_cli_agent(
|
|||||||
)
|
)
|
||||||
be = CompositeBackend(
|
be = CompositeBackend(
|
||||||
default=ws_backend,
|
default=ws_backend,
|
||||||
|
artifacts_root="/workspace" if is_web else "/",
|
||||||
routes={
|
routes={
|
||||||
"/skills/": sk_backend,
|
"/skills/": sk_backend,
|
||||||
"/memories/": mem_backend,
|
"/memories/": mem_backend,
|
||||||
@@ -1225,6 +1657,7 @@ def create_cli_agent(
|
|||||||
bool(profile.memory_workers) if profile is not None else None
|
bool(profile.memory_workers) if profile is not None else None
|
||||||
),
|
),
|
||||||
install_subagent_guard=(profile is not None and not profile.subagents),
|
install_subagent_guard=(profile is not None and not profile.subagents),
|
||||||
|
events=events,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
from .middleware import ProviderContextMediaMiddleware
|
from .middleware import ProviderContextMediaMiddleware
|
||||||
@@ -1242,7 +1675,9 @@ def create_cli_agent(
|
|||||||
)
|
)
|
||||||
mw.insert(
|
mw.insert(
|
||||||
(error_index + 1) if error_index is not None else 0,
|
(error_index + 1) if error_index is not None else 0,
|
||||||
ProviderContextMediaMiddleware(be),
|
ProviderContextMediaMiddleware(
|
||||||
|
be, media_prefix="/workspace/artifacts/model-output" if is_web else "/artifacts/model-output"
|
||||||
|
),
|
||||||
)
|
)
|
||||||
if main_agent_route_middleware is not None:
|
if main_agent_route_middleware is not None:
|
||||||
configurable_index = next(
|
configurable_index = next(
|
||||||
@@ -1260,19 +1695,20 @@ def create_cli_agent(
|
|||||||
if main_agent_outer_middlewares:
|
if main_agent_outer_middlewares:
|
||||||
mw = [*main_agent_outer_middlewares, *mw]
|
mw = [*main_agent_outer_middlewares, *mw]
|
||||||
|
|
||||||
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
|
# HITL on main agent only. Arm it as a middleware and NEVER through
|
||||||
# would propagate it to every subagent, breaking parallel execute calls
|
# `create_deep_agent(interrupt_on=)`: that kwarg makes deepagents append its
|
||||||
# (multi-pending-interrupt LangGraph error).
|
# own plain HumanInTheLoopMiddleware, and that second, auto-blind layer
|
||||||
if not cfg.auto_approve:
|
# interrupts even on a run the gateway verified as auto — silently disabling
|
||||||
mw.append(
|
# automatic approval. It also propagates to every subagent, breaking parallel
|
||||||
HumanInTheLoopMiddleware(
|
# execute calls (multi-pending-interrupt LangGraph error). Web runs defer to
|
||||||
interrupt_on={
|
# the gateway's per-thread review mode, so they always arm the auto-aware
|
||||||
"execute": True,
|
# middleware; `HITL_INTERRUPT_ON` is the single source of the tool set.
|
||||||
"run_in_background": True,
|
if is_web:
|
||||||
"schedule_task": True,
|
from .middleware.dynamic_review import DynamicReviewMiddleware
|
||||||
}
|
|
||||||
)
|
mw.append(DynamicReviewMiddleware(interrupt_on=dict(HITL_INTERRUPT_ON)))
|
||||||
)
|
elif _build_hitl_interrupt_on(auto_approve=cfg.auto_approve) is not None:
|
||||||
|
mw.append(HumanInTheLoopMiddleware(interrupt_on=dict(HITL_INTERRUPT_ON)))
|
||||||
|
|
||||||
# Re-load MCP tools from current config (picks up /mcp add changes)
|
# Re-load MCP tools from current config (picks up /mcp add changes)
|
||||||
kwargs = load_mcp_and_build_kwargs(
|
kwargs = load_mcp_and_build_kwargs(
|
||||||
@@ -1282,11 +1718,44 @@ def create_cli_agent(
|
|||||||
cfg=cfg,
|
cfg=cfg,
|
||||||
chat_model=chat_model,
|
chat_model=chat_model,
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
|
runtime=runtime,
|
||||||
)
|
)
|
||||||
if not enable_subagents:
|
if not enable_subagents:
|
||||||
kwargs = {**kwargs, "subagents": []}
|
kwargs = {**kwargs, "subagents": []}
|
||||||
|
if is_web:
|
||||||
|
kwargs = {**kwargs, "tools": [
|
||||||
|
tool for tool in kwargs.get("tools", [])
|
||||||
|
if getattr(tool, "name", "") != "skill_manager"
|
||||||
|
]}
|
||||||
kwargs = _apply_budgeted_skill_context(kwargs, be)
|
kwargs = _apply_budgeted_skill_context(kwargs, be)
|
||||||
|
|
||||||
|
if agent_model_set is not None:
|
||||||
|
from types import FunctionType
|
||||||
|
def create_run_summarizer(model, backend):
|
||||||
|
return _create_run_summarization_middleware(
|
||||||
|
model, backend, agent_model_set.deepagents_summarizer,
|
||||||
|
)
|
||||||
|
|
||||||
|
# deepagents currently has no per-call summarizer factory parameter.
|
||||||
|
# Scope its existing assembly function to this run, never patch module
|
||||||
|
# globals or register a process-wide profile containing tenant models.
|
||||||
|
if (not isinstance(create_deep_agent, FunctionType)
|
||||||
|
or "create_summarization_middleware" not in create_deep_agent.__code__.co_names):
|
||||||
|
raise RuntimeError("Unsupported deepagents assembly; revalidate summarizer adapter")
|
||||||
|
create_deep_agent = FunctionType(
|
||||||
|
create_deep_agent.__code__,
|
||||||
|
{**create_deep_agent.__globals__,
|
||||||
|
"create_summarization_middleware": create_run_summarizer},
|
||||||
|
create_deep_agent.__name__,
|
||||||
|
create_deep_agent.__defaults__,
|
||||||
|
create_deep_agent.__closure__,
|
||||||
|
)
|
||||||
|
from deepagents import create_deep_agent as original_factory
|
||||||
|
create_deep_agent.__kwdefaults__ = original_factory.__kwdefaults__
|
||||||
|
|
||||||
|
# No `interrupt_on=` here: HITL is armed above, as a single middleware layer
|
||||||
|
# (see the comment there). Passing the kwarg would add a second, auto-blind
|
||||||
|
# layer that defeats verified auto approval.
|
||||||
return create_deep_agent(
|
return create_deep_agent(
|
||||||
**kwargs,
|
**kwargs,
|
||||||
checkpointer=checkpointer,
|
checkpointer=checkpointer,
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from importlib import import_module
|
from importlib import import_module
|
||||||
|
|
||||||
__version__ = "0.2.2"
|
from ._version import __version__
|
||||||
|
|
||||||
_EXPORTS: dict[str, tuple[str, str]] = {
|
_EXPORTS: dict[str, tuple[str, str]] = {
|
||||||
# Agent graph (lazy to avoid expensive initialization at import time)
|
# Agent graph (lazy to avoid expensive initialization at import time)
|
||||||
@@ -73,4 +73,5 @@ def __dir__() -> list[str]:
|
|||||||
return sorted(set(globals()) | set(_EXPORTS))
|
return sorted(set(globals()) | set(_EXPORTS))
|
||||||
|
|
||||||
|
|
||||||
__all__ = list(_EXPORTS)
|
__all__ = ["__version__"]
|
||||||
|
__all__.extend(_EXPORTS)
|
||||||
|
|||||||
@@ -0,0 +1,3 @@
|
|||||||
|
"""Package version shared by builds and runtime."""
|
||||||
|
|
||||||
|
__version__ = "0.3.0"
|
||||||
+558
-7
@@ -4,9 +4,16 @@ import os
|
|||||||
import posixpath
|
import posixpath
|
||||||
import re
|
import re
|
||||||
import shlex
|
import shlex
|
||||||
|
import signal
|
||||||
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import StrEnum
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from deepagents.backends import FilesystemBackend, LocalShellBackend
|
from deepagents.backends import FilesystemBackend, LocalShellBackend
|
||||||
from deepagents.backends.protocol import (
|
from deepagents.backends.protocol import (
|
||||||
@@ -22,7 +29,30 @@ from deepagents.backends.protocol import (
|
|||||||
)
|
)
|
||||||
from filelock import FileLock
|
from filelock import FileLock
|
||||||
|
|
||||||
|
try: # deepagents>=0.7 adds file deletion; the Ai4Sci runtime line holds 0.6.x
|
||||||
|
from deepagents.backends.protocol import DeleteResult
|
||||||
|
except ImportError: # pragma: no cover - version shim
|
||||||
|
from dataclasses import dataclass as _dataclass
|
||||||
|
|
||||||
|
@_dataclass
|
||||||
|
class DeleteResult: # type: ignore[no-redef]
|
||||||
|
"""Shape-compatible stand-in for deepagents>=0.7's ``DeleteResult``.
|
||||||
|
|
||||||
|
Upstream v0.3.0 implements ``delete`` refusals on several backends.
|
||||||
|
deepagents 0.6.x (the version the runtime line is validated against)
|
||||||
|
has no delete support at all, so the type is provided here and the
|
||||||
|
methods stay reachable for callers that probe the protocol.
|
||||||
|
"""
|
||||||
|
|
||||||
|
error: str | None = None
|
||||||
|
path: str | None = None
|
||||||
|
|
||||||
|
|
||||||
from . import paths
|
from . import paths
|
||||||
|
from .cancellation import current_cancel_event
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from langgraph.types import Command
|
||||||
|
|
||||||
# Reproduced here to dodge a circular import from .EvoScientist (the canonical
|
# Reproduced here to dodge a circular import from .EvoScientist (the canonical
|
||||||
# SKILLS_DIR constant).
|
# SKILLS_DIR constant).
|
||||||
@@ -69,6 +99,87 @@ BLOCKED_COMMANDS = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
_active_shell_processes_lock = threading.RLock()
|
||||||
|
_active_shell_processes: dict[threading.Event, set[subprocess.Popen[str]]] = {}
|
||||||
|
_PROCESS_DRAIN_GRACE_SECONDS = 1.0
|
||||||
|
|
||||||
|
|
||||||
|
def _terminate_process_tree(process: subprocess.Popen[str]) -> None:
|
||||||
|
"""Force-stop a shell and its descendants without waiting for reaping."""
|
||||||
|
# A completed Popen has already reaped its PID, which the OS may reuse.
|
||||||
|
# Inspect the recorded state rather than calling poll(): an exited but
|
||||||
|
# unreaped shell can still have live descendants in its process group.
|
||||||
|
if process.returncode is not None:
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
if os.name == "nt":
|
||||||
|
# CREATE_NEW_PROCESS_GROUP alone does not make terminate() recursive.
|
||||||
|
# taskkill is the native way to stop the complete descendant tree.
|
||||||
|
subprocess.run(
|
||||||
|
["taskkill", "/PID", str(process.pid), "/T", "/F"],
|
||||||
|
check=False,
|
||||||
|
capture_output=True,
|
||||||
|
timeout=5,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
os.killpg(process.pid, signal.SIGKILL)
|
||||||
|
except (OSError, subprocess.SubprocessError):
|
||||||
|
try:
|
||||||
|
process.kill()
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def _stop_collecting_process_output(process: subprocess.Popen[str]) -> None:
|
||||||
|
"""Close inherited pipes and reap *process* without blocking the caller."""
|
||||||
|
for pipe in (process.stdout, process.stderr):
|
||||||
|
if pipe is not None:
|
||||||
|
try:
|
||||||
|
pipe.close()
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if process.poll() is None:
|
||||||
|
threading.Thread(target=process.wait, daemon=True).start()
|
||||||
|
|
||||||
|
|
||||||
|
def cancel_active_shell_processes(event: threading.Event) -> None:
|
||||||
|
"""Terminate every active shell command associated with *event*."""
|
||||||
|
with _active_shell_processes_lock:
|
||||||
|
processes = tuple(_active_shell_processes.get(event, ()))
|
||||||
|
for process in processes:
|
||||||
|
_terminate_process_tree(process)
|
||||||
|
|
||||||
|
|
||||||
|
def _register_shell_process(
|
||||||
|
event: threading.Event | None,
|
||||||
|
process: subprocess.Popen[str],
|
||||||
|
) -> None:
|
||||||
|
if event is None:
|
||||||
|
return
|
||||||
|
with _active_shell_processes_lock:
|
||||||
|
_active_shell_processes.setdefault(event, set()).add(process)
|
||||||
|
cancel_now = event.is_set()
|
||||||
|
if cancel_now:
|
||||||
|
_terminate_process_tree(process)
|
||||||
|
|
||||||
|
|
||||||
|
def _unregister_shell_process(
|
||||||
|
event: threading.Event | None,
|
||||||
|
process: subprocess.Popen[str],
|
||||||
|
) -> None:
|
||||||
|
if event is None:
|
||||||
|
return
|
||||||
|
with _active_shell_processes_lock:
|
||||||
|
processes = _active_shell_processes.get(event)
|
||||||
|
if processes is None:
|
||||||
|
return
|
||||||
|
processes.discard(process)
|
||||||
|
if not processes:
|
||||||
|
_active_shell_processes.pop(event, None)
|
||||||
|
|
||||||
|
|
||||||
def _shell_token_spans(command: str) -> list[dict[str, object]]:
|
def _shell_token_spans(command: str) -> list[dict[str, object]]:
|
||||||
"""Tokenize enough shell syntax to find quoted SSH remote commands.
|
"""Tokenize enough shell syntax to find quoted SSH remote commands.
|
||||||
|
|
||||||
@@ -88,7 +199,11 @@ def _shell_token_spans(command: str) -> list[dict[str, object]]:
|
|||||||
if ch in "`();|&":
|
if ch in "`();|&":
|
||||||
return ch
|
return ch
|
||||||
if ch in "<>":
|
if ch in "<>":
|
||||||
if index + 1 < n and command[index + 1] == ch:
|
# `>>`/`<<` and `>|` (force-clobber redirect) are single redirection
|
||||||
|
# operators, NOT a pipe — the trailing `|` must not read as a boundary.
|
||||||
|
if index + 1 < n and (
|
||||||
|
command[index + 1] == ch or (ch == ">" and command[index + 1] == "|")
|
||||||
|
):
|
||||||
return command[index : index + 2]
|
return command[index : index + 2]
|
||||||
return ch
|
return ch
|
||||||
if ch.isdigit():
|
if ch.isdigit():
|
||||||
@@ -97,7 +212,9 @@ def _shell_token_spans(command: str) -> list[dict[str, object]]:
|
|||||||
j += 1
|
j += 1
|
||||||
if j < n and command[j] in "<>":
|
if j < n and command[j] in "<>":
|
||||||
end = j + 1
|
end = j + 1
|
||||||
if end < n and command[end] in ("&", command[j]):
|
# `2>&1`, `2>>`, and `2>|` (fd force-clobber) are single
|
||||||
|
# redirection operators — the trailing `|` is not a pipe.
|
||||||
|
if end < n and command[end] in ("&", "|", command[j]):
|
||||||
end += 1
|
end += 1
|
||||||
return command[index:end]
|
return command[index:end]
|
||||||
return None
|
return None
|
||||||
@@ -159,6 +276,175 @@ def _shell_token_spans(command: str) -> list[dict[str, object]]:
|
|||||||
return tokens
|
return tokens
|
||||||
|
|
||||||
|
|
||||||
|
# Commands that are dangerous as the RIGHT-HAND SIDE of a pipe (they consume
|
||||||
|
# piped data as code or ship it off-box). Everything else piping is normal.
|
||||||
|
_PIPE_NETWORKING_RHS = frozenset(
|
||||||
|
{
|
||||||
|
"nc",
|
||||||
|
"ncat",
|
||||||
|
"netcat",
|
||||||
|
"ssh",
|
||||||
|
"curl",
|
||||||
|
"wget",
|
||||||
|
"telnet",
|
||||||
|
"socat",
|
||||||
|
"scp",
|
||||||
|
"sftp",
|
||||||
|
"rsync",
|
||||||
|
"ftp",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
_PIPE_INTERPRETER_RHS = frozenset(
|
||||||
|
{
|
||||||
|
"sh",
|
||||||
|
"bash",
|
||||||
|
"zsh",
|
||||||
|
"dash",
|
||||||
|
"ash",
|
||||||
|
"ksh",
|
||||||
|
"fish",
|
||||||
|
"python",
|
||||||
|
"python2",
|
||||||
|
"python3",
|
||||||
|
"node",
|
||||||
|
"bun",
|
||||||
|
"deno",
|
||||||
|
"ruby",
|
||||||
|
"perl",
|
||||||
|
"php",
|
||||||
|
"lua",
|
||||||
|
"iex",
|
||||||
|
"elixir",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
_PIPE_DANGEROUS_RHS = _PIPE_INTERPRETER_RHS | _PIPE_NETWORKING_RHS
|
||||||
|
|
||||||
|
|
||||||
|
def check_dangerous_command(command: str) -> str | None:
|
||||||
|
"""Return a reason if *command* pipes output into an interpreter or a
|
||||||
|
network tool, else ``None``.
|
||||||
|
|
||||||
|
Deliberately narrow: this guards indirect prompt injection (the agent
|
||||||
|
ingests untrusted web content and could be induced to run
|
||||||
|
``curl … | bash``). Everyday research shell — pipes into ``grep``/``head``,
|
||||||
|
redirects, ``python -c``, ``..``/``~`` paths — is NOT flagged here.
|
||||||
|
Workspace confinement stays in :func:`validate_command`.
|
||||||
|
|
||||||
|
Only the token immediately after the pipe is inspected, so wrapper
|
||||||
|
commands like ``env bash``, ``xargs bash``, or ``timeout 5 bash`` are
|
||||||
|
not detected — this is a known limitation, not a bug to fix here.
|
||||||
|
"""
|
||||||
|
after_pipe = False
|
||||||
|
for token in _shell_token_spans(command):
|
||||||
|
if token.get("type") == "op":
|
||||||
|
value = token.get("value")
|
||||||
|
if value == "|":
|
||||||
|
after_pipe = True
|
||||||
|
elif value == "&" and after_pipe:
|
||||||
|
# `|&` (pipe stdout+stderr) tokenizes as `|` then `&`;
|
||||||
|
# keep the pipe context open across the `&`.
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
after_pipe = False
|
||||||
|
continue
|
||||||
|
if after_pipe:
|
||||||
|
base = str(token.get("value", "")).split("/")[-1]
|
||||||
|
# strip trailing version digits: python3.11 -> python, lua5.4 -> lua
|
||||||
|
normalized = re.sub(r"[0-9.]+$", "", base) or base
|
||||||
|
if base in _PIPE_DANGEROUS_RHS or normalized in _PIPE_DANGEROUS_RHS:
|
||||||
|
kind = (
|
||||||
|
"networking tool"
|
||||||
|
if base in _PIPE_NETWORKING_RHS
|
||||||
|
or normalized in _PIPE_NETWORKING_RHS
|
||||||
|
else "interpreter"
|
||||||
|
)
|
||||||
|
return f"pipes output into {kind} '{base}'"
|
||||||
|
after_pipe = False
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
class ActionDecision(StrEnum):
|
||||||
|
"""Outcome of the shell-action policy."""
|
||||||
|
|
||||||
|
APPROVE = "approve"
|
||||||
|
REJECT = "reject"
|
||||||
|
PROMPT = "prompt"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ActionVerdict:
|
||||||
|
"""A decision plus the reason to show the user or feed back to the agent."""
|
||||||
|
|
||||||
|
decision: ActionDecision
|
||||||
|
reason: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_action_decision(
|
||||||
|
command: str,
|
||||||
|
*,
|
||||||
|
auto_approve: bool = False,
|
||||||
|
dangerous_mode: bool = False,
|
||||||
|
allow_list: list[str] | None = None,
|
||||||
|
) -> ActionVerdict:
|
||||||
|
"""Single source of truth for approve / reject / prompt.
|
||||||
|
|
||||||
|
Precedence:
|
||||||
|
1. ``dangerous_mode`` — the user asked for full power; run everything.
|
||||||
|
2. dangerous detection — pipe into interpreter/network.
|
||||||
|
3. ``auto_approve`` — opt-out means *never prompt*: approve, or reject
|
||||||
|
a dangerous command with a reason the agent can act on.
|
||||||
|
4. ``allow_list`` — case-sensitive match on a whole command or a
|
||||||
|
command-plus-space prefix; blank entries are ignored. Every segment of
|
||||||
|
a chain (``a; b``, ``a && b``, ``a | b``) must match, so an allow-listed
|
||||||
|
prefix cannot carry a non-listed command in behind it.
|
||||||
|
"""
|
||||||
|
if dangerous_mode:
|
||||||
|
return ActionVerdict(ActionDecision.APPROVE)
|
||||||
|
|
||||||
|
reason = check_dangerous_command(command)
|
||||||
|
|
||||||
|
if auto_approve:
|
||||||
|
if reason:
|
||||||
|
return ActionVerdict(ActionDecision.REJECT, reason)
|
||||||
|
return ActionVerdict(ActionDecision.APPROVE)
|
||||||
|
|
||||||
|
if reason:
|
||||||
|
return ActionVerdict(ActionDecision.PROMPT, reason)
|
||||||
|
|
||||||
|
if allow_list:
|
||||||
|
prefixes = [p.strip() for p in allow_list if p.strip()]
|
||||||
|
# Match on a token boundary so allow-listing `ls` does not also approve
|
||||||
|
# `lsof` (case-sensitive, like the shell). Require EVERY segment of a
|
||||||
|
# chain to match, so `ls; rm -rf x` cannot ride in on an allow-listed
|
||||||
|
# `ls`. ``None`` means an unparseable construct (substitution/newline) —
|
||||||
|
# decline rather than risk approving a hidden command.
|
||||||
|
segments = _split_command_segments(command)
|
||||||
|
if segments is not None:
|
||||||
|
segments = segments or [command.strip()]
|
||||||
|
if prefixes and all(
|
||||||
|
any(seg == p or seg.startswith(p + " ") for p in prefixes)
|
||||||
|
for seg in segments
|
||||||
|
):
|
||||||
|
return ActionVerdict(ActionDecision.APPROVE)
|
||||||
|
|
||||||
|
return ActionVerdict(ActionDecision.PROMPT)
|
||||||
|
|
||||||
|
|
||||||
|
def build_hitl_resume(interrupt_id: str, decisions: list[dict]) -> "Command":
|
||||||
|
"""Build a HITL resume Command keyed by interrupt_id.
|
||||||
|
|
||||||
|
Keying by id (not the flat ``{"decisions": …}``) is REQUIRED whenever the
|
||||||
|
graph has more than one pending interrupt — parallel sub-agents that each
|
||||||
|
call ``execute`` do exactly that, and a flat resume raises
|
||||||
|
``RuntimeError: When there are multiple pending interrupts …``. Resuming a
|
||||||
|
single id resolves that interrupt and re-parks the rest (they re-emit on the
|
||||||
|
next stream), so callers drain them one at a time. Safe for N=1 too.
|
||||||
|
"""
|
||||||
|
from langgraph.types import Command
|
||||||
|
|
||||||
|
return Command(resume={interrupt_id: {"decisions": decisions}})
|
||||||
|
|
||||||
|
|
||||||
_SSH_OPTIONS_WITH_VALUE = {
|
_SSH_OPTIONS_WITH_VALUE = {
|
||||||
"-B",
|
"-B",
|
||||||
"-b",
|
"-b",
|
||||||
@@ -401,6 +687,40 @@ def _split_shell_commands(command: str) -> list[str]:
|
|||||||
return base_commands
|
return base_commands
|
||||||
|
|
||||||
|
|
||||||
|
def _split_command_segments(command: str) -> list[str] | None:
|
||||||
|
"""Split a compound command into raw segment strings on command boundaries.
|
||||||
|
|
||||||
|
Quote-aware (via ``_shell_token_spans``). Boundaries are ``;`` ``&&`` ``||``
|
||||||
|
``|`` ``&`` and grouping; redirections are not boundaries. Lets the allow-list
|
||||||
|
clear a chain only when *every* segment is allow-listed, not just the leading
|
||||||
|
one (``ls; rm -rf x`` must not ride in on an allow-listed ``ls``).
|
||||||
|
|
||||||
|
Returns ``None`` when the command contains a construct this small tokenizer
|
||||||
|
cannot safely reason about — command substitution (``$(...)`` or backticks,
|
||||||
|
which run a hidden command even inside double quotes) or a newline separator —
|
||||||
|
so the caller declines to allow-list it rather than approve a hidden command.
|
||||||
|
Deliberately a substring over-approximation: a literal/quoted ``$(``, backtick,
|
||||||
|
or newline also declines (a safe extra prompt, never a bypass). Quote/escape
|
||||||
|
awareness is intentionally not attempted — that fragility caused the original
|
||||||
|
chaining gap.
|
||||||
|
"""
|
||||||
|
if "$(" in command or "`" in command or "\n" in command or "\r" in command:
|
||||||
|
return None
|
||||||
|
boundaries = {"&&", "||", ";", "|", "&", "(", ")"}
|
||||||
|
segments: list[str] = []
|
||||||
|
seg_start = 0
|
||||||
|
for token in _shell_token_spans(command):
|
||||||
|
if token.get("type") == "op" and token.get("value") in boundaries:
|
||||||
|
seg = command[seg_start : int(token["start"])].strip()
|
||||||
|
if seg:
|
||||||
|
segments.append(seg)
|
||||||
|
seg_start = int(token["end"])
|
||||||
|
tail = command[seg_start:].strip()
|
||||||
|
if tail:
|
||||||
|
segments.append(tail)
|
||||||
|
return segments
|
||||||
|
|
||||||
|
|
||||||
def _has_traversal_component(command: str) -> bool:
|
def _has_traversal_component(command: str) -> bool:
|
||||||
"""Check if command contains '..' as a path component (not substring)."""
|
"""Check if command contains '..' as a path component (not substring)."""
|
||||||
from pathlib import PurePosixPath
|
from pathlib import PurePosixPath
|
||||||
@@ -819,6 +1139,11 @@ class ReadOnlyFilesystemBackend(FilesystemBackend):
|
|||||||
for file_path, _ in files
|
for file_path, _ in files
|
||||||
]
|
]
|
||||||
|
|
||||||
|
def delete(self, file_path: str) -> DeleteResult:
|
||||||
|
return DeleteResult(
|
||||||
|
error="This directory is read-only. Delete operations are not permitted here."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class MemoryFilesystemBackend(FilesystemBackend):
|
class MemoryFilesystemBackend(FilesystemBackend):
|
||||||
"""Filesystem backend for memory files with structured-write enforcement.
|
"""Filesystem backend for memory files with structured-write enforcement.
|
||||||
@@ -835,6 +1160,10 @@ class MemoryFilesystemBackend(FilesystemBackend):
|
|||||||
"Raw edits under /memories are limited to existing "
|
"Raw edits under /memories are limited to existing "
|
||||||
"/memories/profile/... files. Use memory tools for observations."
|
"/memories/profile/... files. Use memory tools for observations."
|
||||||
)
|
)
|
||||||
|
_RAW_DELETE_ERROR = (
|
||||||
|
"Deletes under /memories are blocked. Manage memory files through "
|
||||||
|
"memory tools instead."
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -871,6 +1200,9 @@ class MemoryFilesystemBackend(FilesystemBackend):
|
|||||||
for file_path, _ in files
|
for file_path, _ in files
|
||||||
]
|
]
|
||||||
|
|
||||||
|
def delete(self, file_path: str) -> DeleteResult:
|
||||||
|
return DeleteResult(error=self._RAW_DELETE_ERROR)
|
||||||
|
|
||||||
|
|
||||||
def build_memory_agent_backend(
|
def build_memory_agent_backend(
|
||||||
*,
|
*,
|
||||||
@@ -1082,7 +1414,12 @@ class MergedSkillsBackend(BackendProtocol):
|
|||||||
|
|
||||||
|
|
||||||
def prepare_sandbox_command(
|
def prepare_sandbox_command(
|
||||||
command: str, cwd: str | Path, *, virtual_mode: bool = True, dangerous: bool = False
|
command: str,
|
||||||
|
cwd: str | Path,
|
||||||
|
*,
|
||||||
|
virtual_mode: bool = True,
|
||||||
|
dangerous: bool = False,
|
||||||
|
guard_dangerous: bool = False,
|
||||||
) -> tuple[str, str | None]:
|
) -> tuple[str, str | None]:
|
||||||
"""Normalize workspace paths in ``command`` and validate it for the sandbox.
|
"""Normalize workspace paths in ``command`` and validate it for the sandbox.
|
||||||
|
|
||||||
@@ -1092,7 +1429,16 @@ def prepare_sandbox_command(
|
|||||||
|
|
||||||
Returns ``(prepared_command, error)``: ``error`` is a message string when the command
|
Returns ``(prepared_command, error)``: ``error`` is a message string when the command
|
||||||
is rejected (the caller must NOT run it), otherwise ``None``.
|
is rejected (the caller must NOT run it), otherwise ``None``.
|
||||||
|
|
||||||
|
``guard_dangerous`` (see :func:`check_dangerous_command`) does not see inside an SSH
|
||||||
|
remote payload: a dangerous pipe *inside* a quoted ``ssh host '...'`` argument is not
|
||||||
|
detected, because the quoted payload is a single opaque token. Piping *into* ``ssh``
|
||||||
|
itself (e.g. ``cat secret | ssh host x``) is detected — the check runs on the original,
|
||||||
|
unmasked command so the SSH-masking done below (which also replaces the literal ``ssh``
|
||||||
|
token) does not blind it.
|
||||||
"""
|
"""
|
||||||
|
original_command = command
|
||||||
|
|
||||||
ssh_error = _validate_ssh_remote_command_format(command)
|
ssh_error = _validate_ssh_remote_command_format(command)
|
||||||
if ssh_error:
|
if ssh_error:
|
||||||
return command, ssh_error
|
return command, ssh_error
|
||||||
@@ -1128,6 +1474,19 @@ def prepare_sandbox_command(
|
|||||||
)
|
)
|
||||||
if error:
|
if error:
|
||||||
return command, error
|
return command, error
|
||||||
|
|
||||||
|
# No interactive approval is reachable here (unattended main agent, or an
|
||||||
|
# async sub-agent on a remote thread), so refuse the narrow dangerous set
|
||||||
|
# with a reason the agent can act on rather than running it blind.
|
||||||
|
if guard_dangerous and not dangerous:
|
||||||
|
dangerous_reason = check_dangerous_command(original_command)
|
||||||
|
if dangerous_reason:
|
||||||
|
return _restore_spans(command, ssh_replacements), (
|
||||||
|
f"Command blocked: {dangerous_reason}. "
|
||||||
|
f"Rewrite it to avoid that, or request approval from the user "
|
||||||
|
f"(the orchestrator can re-issue it after approval)."
|
||||||
|
)
|
||||||
|
|
||||||
return _restore_spans(command, ssh_replacements), None
|
return _restore_spans(command, ssh_replacements), None
|
||||||
|
|
||||||
|
|
||||||
@@ -1153,6 +1512,8 @@ class CustomSandboxBackend(LocalShellBackend):
|
|||||||
env: dict[str, str] | None = None,
|
env: dict[str, str] | None = None,
|
||||||
inherit_env: bool = True,
|
inherit_env: bool = True,
|
||||||
dangerous: bool = False,
|
dangerous: bool = False,
|
||||||
|
guard_dangerous: bool = False,
|
||||||
|
refuse_delete: bool = False,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Initialize custom sandbox backend.
|
Initialize custom sandbox backend.
|
||||||
@@ -1168,8 +1529,20 @@ class CustomSandboxBackend(LocalShellBackend):
|
|||||||
paths anywhere on disk (no workspace confinement). Forces
|
paths anywhere on disk (no workspace confinement). Forces
|
||||||
``virtual_mode=False`` and relaxes path validation while keeping
|
``virtual_mode=False`` and relaxes path validation while keeping
|
||||||
the privileged-command blocklist. Defaults to False.
|
the privileged-command blocklist. Defaults to False.
|
||||||
|
guard_dangerous: Refuse the narrow dangerous-command set (see
|
||||||
|
:func:`check_dangerous_command`) outright, for contexts where
|
||||||
|
no interactive approval is reachable (unattended auto-approve
|
||||||
|
runs, async sub-agents). Bypassed when ``dangerous=True``.
|
||||||
|
Defaults to False.
|
||||||
|
refuse_delete: Refuse the recursive ``delete`` FS tool outright,
|
||||||
|
relaying an approval request to the orchestrator. Used for async
|
||||||
|
research sub-agents (writing / data-analysis) that have no
|
||||||
|
interactive approval path. Bypassed when ``dangerous=True``.
|
||||||
|
Defaults to False.
|
||||||
"""
|
"""
|
||||||
self._dangerous = dangerous
|
self._dangerous = dangerous
|
||||||
|
self._guard_dangerous = guard_dangerous
|
||||||
|
self._refuse_delete = refuse_delete
|
||||||
if dangerous:
|
if dangerous:
|
||||||
# Real paths require the legacy (non-virtual) resolution path so the
|
# Real paths require the legacy (non-virtual) resolution path so the
|
||||||
# parent backend returns absolute paths as-is.
|
# parent backend returns absolute paths as-is.
|
||||||
@@ -1239,6 +1612,22 @@ class CustomSandboxBackend(LocalShellBackend):
|
|||||||
|
|
||||||
return super()._resolve_path(key)
|
return super()._resolve_path(key)
|
||||||
|
|
||||||
|
_DELETE_APPROVAL_ERROR = (
|
||||||
|
"Delete blocked: needs approval. Report it to the orchestrator, which "
|
||||||
|
"can re-issue it after approval."
|
||||||
|
)
|
||||||
|
|
||||||
|
def delete(self, file_path: str) -> DeleteResult:
|
||||||
|
"""Refuse ``delete`` for guarded async sub-agents (no approval path).
|
||||||
|
|
||||||
|
No ``adelete`` override is needed: the inherited ``BackendProtocol.adelete``
|
||||||
|
runs ``asyncio.to_thread(self.delete, ...)``, so async sub-agents reach
|
||||||
|
this refusal too.
|
||||||
|
"""
|
||||||
|
if self._refuse_delete and not self._dangerous:
|
||||||
|
return DeleteResult(error=self._DELETE_APPROVAL_ERROR)
|
||||||
|
return super().delete(file_path)
|
||||||
|
|
||||||
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
|
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
|
||||||
"""
|
"""
|
||||||
Execute shell command in sandbox environment.
|
Execute shell command in sandbox environment.
|
||||||
@@ -1248,16 +1637,172 @@ class CustomSandboxBackend(LocalShellBackend):
|
|||||||
- Access to paths outside workspace
|
- Access to paths outside workspace
|
||||||
- Dangerous system commands
|
- Dangerous system commands
|
||||||
|
|
||||||
Then delegates to LocalShellBackend.execute() for actual execution.
|
The validated command is handed to the owned process runner so
|
||||||
|
cancelling an agent turn can terminate the complete process tree.
|
||||||
"""
|
"""
|
||||||
|
# Preserve LocalShellBackend's public validation contract. This
|
||||||
|
# override cannot delegate execution to the base implementation because
|
||||||
|
# it must retain the Popen handle for cancellation, so validate before
|
||||||
|
# command preparation and process launch instead.
|
||||||
|
if not command or not isinstance(command, str):
|
||||||
|
return ExecuteResponse(
|
||||||
|
output="Error: Command must be a non-empty string.",
|
||||||
|
exit_code=1,
|
||||||
|
truncated=False,
|
||||||
|
)
|
||||||
|
|
||||||
command, error = prepare_sandbox_command(
|
command, error = prepare_sandbox_command(
|
||||||
command, self.cwd, virtual_mode=self.virtual_mode, dangerous=self._dangerous
|
command,
|
||||||
|
self.cwd,
|
||||||
|
virtual_mode=self.virtual_mode,
|
||||||
|
dangerous=self._dangerous,
|
||||||
|
guard_dangerous=self._guard_dangerous,
|
||||||
)
|
)
|
||||||
if error:
|
if error:
|
||||||
return ExecuteResponse(output=error, exit_code=1, truncated=False)
|
return ExecuteResponse(output=error, exit_code=1, truncated=False)
|
||||||
|
|
||||||
# Delegate to parent for subprocess execution
|
return self._execute_prepared_command(command, timeout=timeout)
|
||||||
response = super().execute(command, timeout=timeout)
|
|
||||||
|
def _execute_prepared_command(
|
||||||
|
self,
|
||||||
|
command: str,
|
||||||
|
*,
|
||||||
|
timeout: int | None = None,
|
||||||
|
) -> ExecuteResponse:
|
||||||
|
"""Execute an already validated command in an owned process group."""
|
||||||
|
|
||||||
|
effective_timeout = timeout if timeout is not None else self._default_timeout
|
||||||
|
if effective_timeout <= 0:
|
||||||
|
msg = f"timeout must be positive, got {effective_timeout}"
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
cancel_event = current_cancel_event()
|
||||||
|
if cancel_event is not None and cancel_event.is_set():
|
||||||
|
return ExecuteResponse(
|
||||||
|
output="Command cancelled before execution.",
|
||||||
|
exit_code=130,
|
||||||
|
truncated=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
process: subprocess.Popen[str] | None = None
|
||||||
|
termination_reason: str | None = None
|
||||||
|
output_abandoned = False
|
||||||
|
try:
|
||||||
|
process_options: dict[str, object] = {}
|
||||||
|
if os.name == "nt":
|
||||||
|
process_options["creationflags"] = subprocess.CREATE_NEW_PROCESS_GROUP
|
||||||
|
else:
|
||||||
|
process_options["start_new_session"] = True
|
||||||
|
|
||||||
|
process = subprocess.Popen(
|
||||||
|
command,
|
||||||
|
shell=True,
|
||||||
|
stdout=subprocess.PIPE,
|
||||||
|
stderr=subprocess.PIPE,
|
||||||
|
stdin=subprocess.DEVNULL,
|
||||||
|
text=True,
|
||||||
|
env=self._env,
|
||||||
|
cwd=str(self.cwd),
|
||||||
|
**process_options,
|
||||||
|
)
|
||||||
|
_register_shell_process(cancel_event, process)
|
||||||
|
deadline = time.monotonic() + effective_timeout
|
||||||
|
drain_deadline: float | None = None
|
||||||
|
|
||||||
|
while True:
|
||||||
|
now = time.monotonic()
|
||||||
|
if (
|
||||||
|
termination_reason is None
|
||||||
|
and cancel_event is not None
|
||||||
|
and cancel_event.is_set()
|
||||||
|
):
|
||||||
|
termination_reason = "cancelled"
|
||||||
|
_terminate_process_tree(process)
|
||||||
|
drain_deadline = now + _PROCESS_DRAIN_GRACE_SECONDS
|
||||||
|
elif termination_reason is None and now >= deadline:
|
||||||
|
termination_reason = "timed_out"
|
||||||
|
_terminate_process_tree(process)
|
||||||
|
drain_deadline = now + _PROCESS_DRAIN_GRACE_SECONDS
|
||||||
|
|
||||||
|
if drain_deadline is not None and now >= drain_deadline:
|
||||||
|
_stop_collecting_process_output(process)
|
||||||
|
stdout = stderr = ""
|
||||||
|
output_abandoned = True
|
||||||
|
break
|
||||||
|
|
||||||
|
communicate_deadline = (
|
||||||
|
drain_deadline if drain_deadline is not None else deadline
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
stdout, stderr = process.communicate(
|
||||||
|
timeout=max(
|
||||||
|
0.01,
|
||||||
|
min(0.1, communicate_deadline - time.monotonic()),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
break
|
||||||
|
except subprocess.TimeoutExpired:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if termination_reason == "timed_out":
|
||||||
|
if timeout is not None:
|
||||||
|
timeout_output = (
|
||||||
|
"Error: Command timed out after "
|
||||||
|
f"{effective_timeout} seconds (custom timeout). The command "
|
||||||
|
"may be stuck or require more time."
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
timeout_output = (
|
||||||
|
f"Error: Command timed out after {effective_timeout} seconds. "
|
||||||
|
"For long-running commands, re-run using the timeout parameter."
|
||||||
|
)
|
||||||
|
response = ExecuteResponse(
|
||||||
|
output=timeout_output,
|
||||||
|
exit_code=124,
|
||||||
|
truncated=output_abandoned,
|
||||||
|
)
|
||||||
|
elif termination_reason == "cancelled" or (
|
||||||
|
cancel_event is not None and cancel_event.is_set()
|
||||||
|
):
|
||||||
|
response = ExecuteResponse(
|
||||||
|
output="Command cancelled.",
|
||||||
|
exit_code=130,
|
||||||
|
truncated=output_abandoned,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
output_parts = []
|
||||||
|
if stdout:
|
||||||
|
output_parts.append(stdout)
|
||||||
|
if stderr:
|
||||||
|
stderr_lines = stderr.strip().split("\n")
|
||||||
|
output_parts.extend(f"[stderr] {line}" for line in stderr_lines)
|
||||||
|
output = "\n".join(output_parts) if output_parts else "<no output>"
|
||||||
|
|
||||||
|
truncated = False
|
||||||
|
if len(output) > self._max_output_bytes:
|
||||||
|
output = output[: self._max_output_bytes]
|
||||||
|
output += (
|
||||||
|
f"\n\n... Output truncated at {self._max_output_bytes} bytes."
|
||||||
|
)
|
||||||
|
truncated = True
|
||||||
|
if process.returncode != 0:
|
||||||
|
output = f"{output.rstrip()}\n\nExit code: {process.returncode}"
|
||||||
|
response = ExecuteResponse(
|
||||||
|
output=output,
|
||||||
|
exit_code=process.returncode,
|
||||||
|
truncated=truncated,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
if process is not None:
|
||||||
|
_terminate_process_tree(process)
|
||||||
|
response = ExecuteResponse(
|
||||||
|
output=f"Error executing command ({type(exc).__name__}): {exc}",
|
||||||
|
exit_code=1,
|
||||||
|
truncated=False,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
if process is not None:
|
||||||
|
_unregister_shell_process(cancel_event, process)
|
||||||
|
|
||||||
# Enhance timeout errors with actionable recovery guidance
|
# Enhance timeout errors with actionable recovery guidance
|
||||||
if response.exit_code == 124:
|
if response.exit_code == 124:
|
||||||
@@ -1325,6 +1870,12 @@ class AutoskillProposalSandboxBackend(CustomSandboxBackend):
|
|||||||
for file_path, _ in files
|
for file_path, _ in files
|
||||||
]
|
]
|
||||||
|
|
||||||
|
def delete(self, file_path: str) -> DeleteResult:
|
||||||
|
return DeleteResult(
|
||||||
|
error="Deletes are blocked for AutoSkills. Manage proposal files "
|
||||||
|
"under /autoskill-proposals/ instead."
|
||||||
|
)
|
||||||
|
|
||||||
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
|
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
|
||||||
return super().execute(
|
return super().execute(
|
||||||
self._rewrite_autoskill_mount(command),
|
self._rewrite_autoskill_mount(command),
|
||||||
|
|||||||
@@ -0,0 +1,27 @@
|
|||||||
|
"""Cancellation context shared by streaming frontends and blocking tools."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextvars
|
||||||
|
import threading
|
||||||
|
from collections.abc import Iterator
|
||||||
|
from contextlib import contextmanager
|
||||||
|
|
||||||
|
_current_cancel_event: contextvars.ContextVar[threading.Event | None] = (
|
||||||
|
contextvars.ContextVar("evoscientist_cancel_event", default=None)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def bind_cancel_event(event: threading.Event) -> Iterator[None]:
|
||||||
|
"""Make a stream's cancellation event visible to nested sync tool calls."""
|
||||||
|
token = _current_cancel_event.set(event)
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
_current_cancel_event.reset(token)
|
||||||
|
|
||||||
|
|
||||||
|
def current_cancel_event() -> threading.Event | None:
|
||||||
|
"""Return the cancellation event bound to the current agent run, if any."""
|
||||||
|
return _current_cancel_event.get()
|
||||||
+187
-66
@@ -7,20 +7,24 @@ This module defines the Channel interface that all messaging channels
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import re
|
import re
|
||||||
|
import threading
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||||
from collections.abc import Callable as CallableABC
|
from collections.abc import Callable as CallableABC
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import datetime
|
from datetime import UTC, datetime
|
||||||
|
from email.utils import parsedate_to_datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from ..paths import MEDIA_DIR
|
from ..paths import MEDIA_DIR
|
||||||
|
from ..runtime import AsyncRuntime
|
||||||
from .bus.events import InboundMessage, OutboundMessage
|
from .bus.events import InboundMessage, OutboundMessage
|
||||||
from .capabilities import ChannelCapabilities
|
from .capabilities import ChannelCapabilities
|
||||||
from .debug import TraceMixin, debug_trace_enabled
|
from .debug import TraceMixin, debug_trace_enabled
|
||||||
from .formatter import UnifiedFormatter
|
from .formatter import UnifiedFormatter
|
||||||
|
from .interaction import is_slash_command
|
||||||
from .plugin import ChannelMeta, ChannelPlugin
|
from .plugin import ChannelMeta, ChannelPlugin
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_logger = logging.getLogger(__name__)
|
||||||
@@ -298,6 +302,8 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
|
|||||||
maxsize=queue_maxsize
|
maxsize=queue_maxsize
|
||||||
)
|
)
|
||||||
self._running = False
|
self._running = False
|
||||||
|
self._startup_event = threading.Event()
|
||||||
|
self._startup_error: str | None = None
|
||||||
|
|
||||||
# Global tracing can be enabled via shared config/env even when
|
# Global tracing can be enabled via shared config/env even when
|
||||||
# individual channel factories have not been updated yet.
|
# individual channel factories have not been updated yet.
|
||||||
@@ -741,68 +747,144 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
|
|||||||
|
|
||||||
# ── Send retry abstraction ──────────────────────────────────────
|
# ── Send retry abstraction ──────────────────────────────────────
|
||||||
|
|
||||||
_non_retryable_patterns: tuple[str, ...] = ()
|
# HTTP status codes that should never be retried. Listed explicitly
|
||||||
|
# rather than as a 4xx range: 408 and 425 are retryable by definition and
|
||||||
|
# 429 is handled by the rate-limit path.
|
||||||
|
_non_retryable_status_codes: tuple[int, ...] = (400, 401, 403, 404)
|
||||||
|
|
||||||
|
# Structured SDK error codes that should never be retried (e.g. Slack invalid_auth)
|
||||||
|
# Channel-specific message patterns (e.g. Feishu 10003, DingTalk 40014) are handled
|
||||||
|
# via _non_retryable_patterns in respective channel subclasses.
|
||||||
|
_non_retryable_error_codes: tuple[str, ...] = (
|
||||||
|
"invalid_auth",
|
||||||
|
"invalid_token",
|
||||||
|
"expired_token",
|
||||||
|
"token_expired",
|
||||||
|
"token_revoked",
|
||||||
|
"account_inactive",
|
||||||
|
"not_authed",
|
||||||
|
"no_permission",
|
||||||
|
"missing_scope",
|
||||||
|
)
|
||||||
|
|
||||||
|
_non_retryable_patterns: tuple[str, ...] = (
|
||||||
|
"unauthorized",
|
||||||
|
"forbidden",
|
||||||
|
"permission denied",
|
||||||
|
"invalid token",
|
||||||
|
"invalid api key",
|
||||||
|
"authentication failed",
|
||||||
|
)
|
||||||
_rate_limit_patterns: tuple[str, ...] = ("429", "ratelimit")
|
_rate_limit_patterns: tuple[str, ...] = ("429", "ratelimit")
|
||||||
_rate_limit_delay: float = 1.0
|
_rate_limit_delay: float = 1.0
|
||||||
|
|
||||||
def _extract_retry_after(self, exc: Exception) -> float | None:
|
def _extract_retry_after(self, exc: Exception) -> float | None:
|
||||||
"""Extract retry-wait seconds from an exception.
|
"""Extract retry-wait seconds from an exception.
|
||||||
|
|
||||||
Returns ``None`` to signal that the error is **not retryable**.
|
Returns a retry delay in seconds, or ``None`` when the error is
|
||||||
|
explicitly non-retryable.
|
||||||
|
|
||||||
Pipeline:
|
Pipeline:
|
||||||
1. SDK-provided ``retry_after`` attribute (Telegram / Slack SDKs).
|
1. Non-retryable detection → ``None``. Evaluates HTTP status codes
|
||||||
2. HTTP ``Retry-After`` header via :meth:`_parse_retry_after_header`.
|
(e.g. 401, 403), structured SDK error codes (e.g. Slack
|
||||||
3. Non-retryable pattern match → ``None``.
|
``"invalid_auth"``), and message pattern matching
|
||||||
4. Rate-limit pattern match → ``_rate_limit_delay``.
|
(e.g. ``"unauthorized"``, ``"forbidden"``).
|
||||||
5. Default ``1.0`` s (generic transient-error retry).
|
2. Server-supplied delay via :meth:`_extract_retry_delay`
|
||||||
|
(httpx ``Retry-After``; channels override for their SDK).
|
||||||
|
3. Rate-limit pattern match → ``_rate_limit_delay``.
|
||||||
|
4. Default ``1.0`` s for generic transient errors.
|
||||||
|
|
||||||
Channels can customize behavior declaratively via class attributes
|
Channels can customize behavior declaratively via class attributes
|
||||||
``_non_retryable_patterns``, ``_rate_limit_patterns``, and
|
``_non_retryable_patterns``, ``_rate_limit_patterns``,
|
||||||
``_rate_limit_delay``, or override this method entirely.
|
``_non_retryable_status_codes``, ``_non_retryable_error_codes``,
|
||||||
|
and ``_rate_limit_delay``, or override this method entirely.
|
||||||
"""
|
"""
|
||||||
# 1. SDK retry_after attribute
|
# 1. Non-retryable detection: evaluate status codes, structured SDK
|
||||||
retry = getattr(exc, "retry_after", None)
|
# error codes, and message patterns independently.
|
||||||
if retry is not None:
|
status_code = self._extract_status_code(exc)
|
||||||
return float(retry)
|
if status_code is not None and status_code in self._non_retryable_status_codes:
|
||||||
|
return None
|
||||||
|
|
||||||
# 2. HTTP Retry-After header
|
sdk_error = self._extract_sdk_error_code(exc)
|
||||||
header_val = self._parse_retry_after_header(exc)
|
if sdk_error is not None and sdk_error in self._non_retryable_error_codes:
|
||||||
if header_val is not None:
|
return None
|
||||||
return header_val
|
|
||||||
|
|
||||||
msg = str(exc).lower()
|
msg = str(exc).lower()
|
||||||
|
|
||||||
# 3. Non-retryable patterns
|
|
||||||
if self._non_retryable_patterns and any(
|
if self._non_retryable_patterns and any(
|
||||||
p in msg for p in self._non_retryable_patterns
|
p in msg for p in self._non_retryable_patterns
|
||||||
):
|
):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# 4. Rate-limit patterns
|
# 2. Server-supplied delay
|
||||||
|
delay = self._extract_retry_delay(exc)
|
||||||
|
if delay is not None:
|
||||||
|
return delay
|
||||||
|
|
||||||
|
# 3. Rate-limit patterns
|
||||||
if self._rate_limit_patterns and any(
|
if self._rate_limit_patterns and any(
|
||||||
p in msg for p in self._rate_limit_patterns
|
p in msg for p in self._rate_limit_patterns
|
||||||
):
|
):
|
||||||
return self._rate_limit_delay
|
return self._rate_limit_delay
|
||||||
|
|
||||||
# 5. Default
|
# 4. Default: transient error, retry with the standard delay
|
||||||
return 1.0
|
return 1.0
|
||||||
|
|
||||||
def _parse_retry_after_header(self, exc: Exception) -> float | None:
|
def _extract_status_code(self, exc: Exception) -> int | None:
|
||||||
"""Try to extract a ``Retry-After`` value from an HTTP response."""
|
"""Extract HTTP status from an httpx error.
|
||||||
resp = getattr(exc, "response", None)
|
|
||||||
if resp is None:
|
Channels with other SDKs (e.g. ``SlackChannel``, ``DiscordChannel``)
|
||||||
return None
|
override this method.
|
||||||
headers = getattr(resp, "headers", None)
|
"""
|
||||||
if not headers:
|
import httpx
|
||||||
return None
|
|
||||||
raw = headers.get("Retry-After") or headers.get("retry-after")
|
if isinstance(exc, httpx.HTTPStatusError):
|
||||||
if raw is None:
|
return exc.response.status_code
|
||||||
return None
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _extract_sdk_error_code(self, exc: Exception) -> str | None:
|
||||||
|
"""Extract structured SDK error code string from an exception.
|
||||||
|
|
||||||
|
Plain HTTP carries no structured error code by default (returns ``None``).
|
||||||
|
Subclasses with specialized SDKs (e.g. ``SlackChannel``, ``DiscordChannel``)
|
||||||
|
override this method.
|
||||||
|
"""
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _extract_retry_delay(self, exc: Exception) -> float | None:
|
||||||
|
"""Retry delay the server asked for, in seconds, or ``None``.
|
||||||
|
|
||||||
|
Base implementation reads the ``Retry-After`` header of an httpx
|
||||||
|
error. Channels whose SDK reports the delay differently
|
||||||
|
(``SlackChannel``, ``TelegramChannel``, ``DiscordChannel``) override.
|
||||||
|
"""
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
if isinstance(exc, httpx.HTTPStatusError):
|
||||||
|
raw = exc.response.headers.get("retry-after")
|
||||||
|
if raw is not None:
|
||||||
|
return self._parse_retry_after(raw)
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _parse_retry_after(raw: str) -> float | None:
|
||||||
|
"""Convert a ``Retry-After`` header value to seconds.
|
||||||
|
|
||||||
|
RFC 9110 allows either delay-seconds or an HTTP-date; a date is
|
||||||
|
returned as the non-negative number of seconds until it. Unparseable
|
||||||
|
values yield ``None`` so the caller can fall back to its own delay.
|
||||||
|
"""
|
||||||
try:
|
try:
|
||||||
return float(raw)
|
return float(raw)
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
when = parsedate_to_datetime(raw)
|
||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
return None
|
return None
|
||||||
|
if when.tzinfo is None:
|
||||||
|
when = when.replace(tzinfo=UTC)
|
||||||
|
return max(0.0, (when - datetime.now(UTC)).total_seconds())
|
||||||
|
|
||||||
async def _send_with_retry(
|
async def _send_with_retry(
|
||||||
self,
|
self,
|
||||||
@@ -923,34 +1005,27 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
|
|||||||
return None
|
return None
|
||||||
return self._raw_to_inbound(current)
|
return self._raw_to_inbound(current)
|
||||||
|
|
||||||
def _build_inbound(self, raw: RawIncoming) -> InboundMessage | None:
|
def _build_inbound(
|
||||||
|
self,
|
||||||
|
raw: RawIncoming,
|
||||||
|
*,
|
||||||
|
runtime: AsyncRuntime | None = None,
|
||||||
|
) -> InboundMessage | None:
|
||||||
"""Run *raw* through inbound middlewares and convert to InboundMessage.
|
"""Run *raw* through inbound middlewares and convert to InboundMessage.
|
||||||
|
|
||||||
Synchronous wrapper around :meth:`_build_inbound_async`. When an
|
Compatibility wrapper for synchronous integrations. Internal channel
|
||||||
event loop is already running, the coroutine is scheduled on that
|
implementations should await :meth:`_build_inbound_async` on their
|
||||||
loop via :func:`asyncio.run_coroutine_threadsafe` to avoid
|
transport loop. A caller may provide its application runtime to reuse
|
||||||
thread-safety issues with middleware state (DedupCache,
|
that owner; otherwise a runtime is scoped to this call.
|
||||||
GroupHistoryBuffer, etc.).
|
|
||||||
|
This method deliberately rejects callers already running an event
|
||||||
|
loop. Blocking such a loop while scheduling the coroutine back onto it
|
||||||
|
deadlocks; async callers must await :meth:`_build_inbound_async`.
|
||||||
"""
|
"""
|
||||||
import asyncio
|
if runtime is None:
|
||||||
|
with AsyncRuntime(thread_name="evosci-channel-adapter-runtime") as owned:
|
||||||
try:
|
return self._build_inbound(raw, runtime=owned)
|
||||||
loop = asyncio.get_running_loop()
|
return runtime.run_sync(lambda: self._build_inbound_async(raw))
|
||||||
except RuntimeError:
|
|
||||||
loop = None
|
|
||||||
|
|
||||||
if loop is not None and loop.is_running():
|
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
|
||||||
self._build_inbound_async(raw),
|
|
||||||
loop,
|
|
||||||
)
|
|
||||||
return future.result()
|
|
||||||
else:
|
|
||||||
new_loop = asyncio.new_event_loop()
|
|
||||||
try:
|
|
||||||
return new_loop.run_until_complete(self._build_inbound_async(raw))
|
|
||||||
finally:
|
|
||||||
new_loop.close()
|
|
||||||
|
|
||||||
def _raw_to_inbound(self, raw: RawIncoming) -> InboundMessage | None:
|
def _raw_to_inbound(self, raw: RawIncoming) -> InboundMessage | None:
|
||||||
"""Convert a RawIncoming to InboundMessage (pure transformation, no filtering).
|
"""Convert a RawIncoming to InboundMessage (pure transformation, no filtering).
|
||||||
@@ -1054,6 +1129,43 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
|
|||||||
"""Buffer *msg* with debounce, then publish to bus."""
|
"""Buffer *msg* with debounce, then publish to bus."""
|
||||||
sender = msg.sender_id
|
sender = msg.sender_id
|
||||||
|
|
||||||
|
if self._on_activity:
|
||||||
|
try:
|
||||||
|
self._on_activity(sender, "received")
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Slash commands are control messages, not prompt fragments. Flush any
|
||||||
|
# prompt already waiting for this sender, then publish the command as
|
||||||
|
# its own message so either arrival order cannot newline-merge them.
|
||||||
|
if is_slash_command(msg.content) and self._bus:
|
||||||
|
# A flush removes itself from this mapping before awaiting the bus
|
||||||
|
# publish. Therefore a task still present here has not detached
|
||||||
|
# its buffered payload yet and is safe to cancel; an in-flight,
|
||||||
|
# backpressured publish is deliberately left alone.
|
||||||
|
debounce_task = self._debounce_tasks.pop(sender, None)
|
||||||
|
if debounce_task is not None:
|
||||||
|
debounce_task.cancel()
|
||||||
|
try:
|
||||||
|
await debounce_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
# Awaiting a cancelled child normally raises here with no
|
||||||
|
# cancellation pending on this task. If our caller also
|
||||||
|
# cancelled queue_message(), preserve that outer signal.
|
||||||
|
current = asyncio.current_task()
|
||||||
|
if current is not None and current.cancelling() > 0:
|
||||||
|
raise
|
||||||
|
try:
|
||||||
|
await self._process_buffered_messages(sender)
|
||||||
|
except Exception:
|
||||||
|
_logger.error(
|
||||||
|
f"{self.name} buffered-prompt flush failed for {sender}; "
|
||||||
|
"publishing the command anyway",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
await self._bus.publish_inbound(msg)
|
||||||
|
return
|
||||||
|
|
||||||
if sender not in self._message_buffers:
|
if sender not in self._message_buffers:
|
||||||
self._message_buffers[sender] = []
|
self._message_buffers[sender] = []
|
||||||
self._message_metadata[sender] = msg.metadata
|
self._message_metadata[sender] = msg.metadata
|
||||||
@@ -1066,12 +1178,6 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
|
|||||||
if msg.media:
|
if msg.media:
|
||||||
self._message_media[sender].extend(msg.media)
|
self._message_media[sender].extend(msg.media)
|
||||||
|
|
||||||
if self._on_activity:
|
|
||||||
try:
|
|
||||||
self._on_activity(sender, "received")
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
if sender in self._debounce_tasks:
|
if sender in self._debounce_tasks:
|
||||||
self._debounce_tasks[sender].cancel()
|
self._debounce_tasks[sender].cancel()
|
||||||
|
|
||||||
@@ -1086,8 +1192,10 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
|
|||||||
await asyncio.sleep(_w)
|
await asyncio.sleep(_w)
|
||||||
try:
|
try:
|
||||||
await self._process_buffered_messages(_s)
|
await self._process_buffered_messages(_s)
|
||||||
except Exception as e:
|
except Exception:
|
||||||
_logger.error(f"{self.name} debounce flush error for {_s}: {e}")
|
_logger.error(
|
||||||
|
f"{self.name} debounce flush error for {_s}", exc_info=True
|
||||||
|
)
|
||||||
|
|
||||||
self._debounce_tasks[sender] = asyncio.create_task(debounce_callback())
|
self._debounce_tasks[sender] = asyncio.create_task(debounce_callback())
|
||||||
|
|
||||||
@@ -1167,16 +1275,25 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
|
|||||||
"""Run the channel with auto-reconnect (exponential backoff)."""
|
"""Run the channel with auto-reconnect (exponential backoff)."""
|
||||||
backoff = 1.0
|
backoff = 1.0
|
||||||
max_backoff = 60.0
|
max_backoff = 60.0
|
||||||
|
self._startup_event.clear()
|
||||||
|
self._startup_error = None
|
||||||
self._running = True
|
self._running = True
|
||||||
while self._running:
|
while self._running:
|
||||||
try:
|
try:
|
||||||
await self.start()
|
await self.start()
|
||||||
|
self._startup_error = None
|
||||||
|
self._startup_event.set()
|
||||||
backoff = 1.0
|
backoff = 1.0
|
||||||
async for msg in self.receive():
|
async for msg in self.receive():
|
||||||
await self.queue_message(msg)
|
await self.queue_message(msg)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
|
if not self._startup_event.is_set():
|
||||||
|
self._startup_error = "startup cancelled"
|
||||||
|
self._startup_event.set()
|
||||||
break
|
break
|
||||||
except ChannelError as e:
|
except ChannelError as e:
|
||||||
|
self._startup_error = str(e)
|
||||||
|
self._startup_event.set()
|
||||||
self._trace_event(
|
self._trace_event(
|
||||||
"channel_fatal_error",
|
"channel_fatal_error",
|
||||||
error_type=type(e).__name__,
|
error_type=type(e).__name__,
|
||||||
@@ -1204,6 +1321,10 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
|
|||||||
await asyncio.sleep(backoff)
|
await asyncio.sleep(backoff)
|
||||||
backoff = min(backoff * 2, max_backoff)
|
backoff = min(backoff * 2, max_backoff)
|
||||||
|
|
||||||
|
if not self._startup_event.is_set():
|
||||||
|
self._startup_error = "channel stopped before startup completed"
|
||||||
|
self._startup_event.set()
|
||||||
|
|
||||||
# ── Channel allow-list check ─────────────────────────────────────
|
# ── Channel allow-list check ─────────────────────────────────────
|
||||||
|
|
||||||
def is_channel_allowed(self, channel_id: str) -> bool:
|
def is_channel_allowed(self, channel_id: str) -> bool:
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ class OutboundMessage:
|
|||||||
reply_to: str | None = None
|
reply_to: str | None = None
|
||||||
media: list[str] = field(default_factory=list)
|
media: list[str] = field(default_factory=list)
|
||||||
metadata: dict[str, Any] = field(default_factory=dict)
|
metadata: dict[str, Any] = field(default_factory=dict)
|
||||||
|
failure_notice: str | None = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def recipient(self) -> str:
|
def recipient(self) -> str:
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ import logging
|
|||||||
import pkgutil
|
import pkgutil
|
||||||
import time
|
import time
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field, replace
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -29,6 +29,11 @@ from .plugin import ChannelPlugin
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Best-effort failure notices must never wedge the dispatcher on a hung send.
|
||||||
|
_FAILURE_NOTICE_TIMEOUT = 15.0
|
||||||
|
|
||||||
|
CHANNEL_STARTUP_PENDING_DETAIL = "starting (bus)"
|
||||||
|
|
||||||
|
|
||||||
# ═════════════════════════════════════════════════════════════════════
|
# ═════════════════════════════════════════════════════════════════════
|
||||||
# Account management (formerly account.py)
|
# Account management (formerly account.py)
|
||||||
@@ -741,6 +746,12 @@ class ChannelManager:
|
|||||||
delivery_failed = True
|
delivery_failed = True
|
||||||
if not delivery_failed and (msg.content or msg.media):
|
if not delivery_failed and (msg.content or msg.media):
|
||||||
drained += 1
|
drained += 1
|
||||||
|
elif delivery_failed:
|
||||||
|
await self._send_failure_notice(
|
||||||
|
channel,
|
||||||
|
msg,
|
||||||
|
timeout=max(1.0, deadline - time.monotonic()),
|
||||||
|
)
|
||||||
dropped = self.bus.outbound.qsize()
|
dropped = self.bus.outbound.qsize()
|
||||||
if drained or dropped:
|
if drained or dropped:
|
||||||
logger.info(f"Outbound drain: {drained} sent, {dropped} dropped")
|
logger.info(f"Outbound drain: {drained} sent, {dropped} dropped")
|
||||||
@@ -841,6 +852,50 @@ class ChannelManager:
|
|||||||
|
|
||||||
# ── outbound routing ──
|
# ── outbound routing ──
|
||||||
|
|
||||||
|
def _record_outbound_failure(self, channel_name: str, error: str) -> None:
|
||||||
|
health = self._health.get(channel_name)
|
||||||
|
if health is None:
|
||||||
|
return
|
||||||
|
health.consecutive_failures += 1
|
||||||
|
health.total_failures += 1
|
||||||
|
health.last_failure_time = time.time()
|
||||||
|
health.last_failure_error = error
|
||||||
|
|
||||||
|
async def _send_failure_notice(
|
||||||
|
self,
|
||||||
|
channel: Channel,
|
||||||
|
msg: OutboundMessage,
|
||||||
|
*,
|
||||||
|
timeout: float | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Best-effort short notice when the real payload could not be sent."""
|
||||||
|
if not msg.failure_notice:
|
||||||
|
return
|
||||||
|
fallback = replace(
|
||||||
|
msg,
|
||||||
|
content=msg.failure_notice,
|
||||||
|
media=[],
|
||||||
|
failure_notice=None,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
coro = channel.send(fallback)
|
||||||
|
if timeout is not None:
|
||||||
|
coro = asyncio.wait_for(coro, timeout=timeout)
|
||||||
|
fallback_ok = await coro
|
||||||
|
except Exception as fallback_error:
|
||||||
|
logger.error(
|
||||||
|
"Error sending delivery failure notice to %s: %s",
|
||||||
|
msg.channel,
|
||||||
|
fallback_error,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
if not fallback_ok:
|
||||||
|
logger.error(
|
||||||
|
"Error sending delivery failure notice to %s: "
|
||||||
|
"send() returned False",
|
||||||
|
msg.channel,
|
||||||
|
)
|
||||||
|
|
||||||
async def _dispatch_outbound(self) -> None:
|
async def _dispatch_outbound(self) -> None:
|
||||||
"""Route outbound messages from the bus to the correct channel."""
|
"""Route outbound messages from the bus to the correct channel."""
|
||||||
logger.info("Outbound dispatcher started")
|
logger.info("Outbound dispatcher started")
|
||||||
@@ -870,13 +925,20 @@ class ChannelManager:
|
|||||||
msg = processed
|
msg = processed
|
||||||
|
|
||||||
delivery_failed = False
|
delivery_failed = False
|
||||||
|
failure_error = "one or more outbound deliveries failed"
|
||||||
if msg.content:
|
if msg.content:
|
||||||
text_ok = await channel.send(msg)
|
try:
|
||||||
if not text_ok:
|
text_ok = await channel.send(msg)
|
||||||
logger.error(
|
except Exception as e:
|
||||||
f"Error sending to {msg.channel}: send() returned False"
|
logger.error(f"Error sending to {msg.channel}", exc_info=True)
|
||||||
)
|
failure_error = str(e)
|
||||||
delivery_failed = True
|
delivery_failed = True
|
||||||
|
else:
|
||||||
|
if not text_ok:
|
||||||
|
logger.error(
|
||||||
|
f"Error sending to {msg.channel}: send() returned False"
|
||||||
|
)
|
||||||
|
delivery_failed = True
|
||||||
|
|
||||||
for media_path in msg.media:
|
for media_path in msg.media:
|
||||||
try:
|
try:
|
||||||
@@ -892,11 +954,18 @@ class ChannelManager:
|
|||||||
)
|
)
|
||||||
delivery_failed = True
|
delivery_failed = True
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error sending media to {msg.channel}: {e}")
|
logger.error(
|
||||||
|
f"Error sending media to {msg.channel}", exc_info=True
|
||||||
|
)
|
||||||
|
failure_error = str(e)
|
||||||
delivery_failed = True
|
delivery_failed = True
|
||||||
|
|
||||||
if delivery_failed:
|
if delivery_failed:
|
||||||
raise RuntimeError("one or more outbound deliveries failed")
|
await self._send_failure_notice(
|
||||||
|
channel, msg, timeout=_FAILURE_NOTICE_TIMEOUT
|
||||||
|
)
|
||||||
|
self._record_outbound_failure(msg.channel, failure_error)
|
||||||
|
continue
|
||||||
|
|
||||||
# Success
|
# Success
|
||||||
health = self._health.get(msg.channel)
|
health = self._health.get(msg.channel)
|
||||||
@@ -904,13 +973,12 @@ class ChannelManager:
|
|||||||
health.consecutive_failures = 0
|
health.consecutive_failures = 0
|
||||||
health.total_successes += 1
|
health.total_successes += 1
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error sending to {msg.channel}: {e}")
|
# Unexpected internal error (pipeline, bookkeeping) — the
|
||||||
health = self._health.get(msg.channel)
|
# transport paths above handle their own failures.
|
||||||
if health is not None:
|
logger.error(
|
||||||
health.consecutive_failures += 1
|
f"Outbound dispatch error for {msg.channel}", exc_info=True
|
||||||
health.total_failures += 1
|
)
|
||||||
health.last_failure_time = time.monotonic()
|
self._record_outbound_failure(msg.channel, str(e))
|
||||||
health.last_failure_error = str(e)
|
|
||||||
|
|
||||||
# ── per-account lifecycle ──
|
# ── per-account lifecycle ──
|
||||||
|
|
||||||
@@ -978,6 +1046,31 @@ class ChannelManager:
|
|||||||
"""Return names of currently running channels."""
|
"""Return names of currently running channels."""
|
||||||
return [name for name, ch in self._channels.items() if ch._running]
|
return [name for name, ch in self._channels.items() if ch._running]
|
||||||
|
|
||||||
|
def startup_results(self, *, timeout: float = 0.0) -> list[tuple[str, bool, str]]:
|
||||||
|
"""Return each channel's initial connection result.
|
||||||
|
|
||||||
|
The optional timeout is shared across all channels, which start
|
||||||
|
concurrently. Channels still connecting when it expires are reported
|
||||||
|
as starting rather than connected.
|
||||||
|
"""
|
||||||
|
deadline = time.monotonic() + max(timeout, 0.0)
|
||||||
|
for channel in self._channels.values():
|
||||||
|
remaining = deadline - time.monotonic()
|
||||||
|
if remaining > 0 and not channel._startup_event.is_set():
|
||||||
|
channel._startup_event.wait(remaining)
|
||||||
|
|
||||||
|
results: list[tuple[str, bool, str]] = []
|
||||||
|
for name, channel in self._channels.items():
|
||||||
|
if not channel._startup_event.is_set():
|
||||||
|
results.append((name, False, CHANNEL_STARTUP_PENDING_DETAIL))
|
||||||
|
elif channel._startup_error:
|
||||||
|
results.append((name, False, f"failed: {channel._startup_error}"))
|
||||||
|
elif channel._running:
|
||||||
|
results.append((name, True, "connected (bus)"))
|
||||||
|
else:
|
||||||
|
results.append((name, False, "stopped during startup"))
|
||||||
|
return results
|
||||||
|
|
||||||
def get_stats(self) -> dict:
|
def get_stats(self) -> dict:
|
||||||
"""Return summary stats for all channels."""
|
"""Return summary stats for all channels."""
|
||||||
return {
|
return {
|
||||||
|
|||||||
+138
-371
@@ -20,6 +20,17 @@ from ..gateway import GraphGateway, GraphRunInput, GraphTarget, RunRequest
|
|||||||
from .base import Channel
|
from .base import Channel
|
||||||
from .bus import MessageBus
|
from .bus import MessageBus
|
||||||
from .bus.events import InboundMessage, OutboundMessage
|
from .bus.events import InboundMessage, OutboundMessage
|
||||||
|
from .capabilities import ChannelCapabilities
|
||||||
|
from .interaction import (
|
||||||
|
ASK_USER_TIMEOUT,
|
||||||
|
HITL_APPROVAL_TIMEOUT,
|
||||||
|
REJECTED_FEEDBACK,
|
||||||
|
ApprovalPolicy,
|
||||||
|
InteractionIO,
|
||||||
|
PendingReplyRegistry,
|
||||||
|
resolve_approval,
|
||||||
|
resolve_ask_user,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -28,10 +39,6 @@ T = TypeVar("T")
|
|||||||
_MAX_CHAT_LOCKS = 10_000
|
_MAX_CHAT_LOCKS = 10_000
|
||||||
_MAX_SESSIONS = 10_000
|
_MAX_SESSIONS = 10_000
|
||||||
_MAX_HITL_ROUNDS = 50
|
_MAX_HITL_ROUNDS = 50
|
||||||
_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)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -108,120 +115,56 @@ def _join_subagent_text(buffers: dict[str, tuple[str, list[str]]]) -> str:
|
|||||||
return "\n\n".join(sections)
|
return "\n\n".join(sections)
|
||||||
|
|
||||||
|
|
||||||
def _should_auto_approve(action_requests: list[dict]) -> bool:
|
class _ConsumerIO(InteractionIO):
|
||||||
"""Check if all action requests can be auto-approved via config.
|
""":class:`InteractionIO` over the consumer's bus + reply registry.
|
||||||
|
|
||||||
Returns True if no manual approval is needed (config auto_approve,
|
Publishes prompts through ``bus.publish_outbound`` and blocks for
|
||||||
non-execute tools, or shell_allow_list match).
|
replies on the consumer's shared :class:`PendingReplyRegistry` — both
|
||||||
|
on the consumer's own event loop, so the engine runs natively async
|
||||||
|
here with no thread hand-off.
|
||||||
"""
|
"""
|
||||||
if not action_requests:
|
|
||||||
|
def __init__(
|
||||||
|
self, consumer: InboundConsumer, msg: InboundMessage, session_key: str
|
||||||
|
) -> None:
|
||||||
|
self._consumer = consumer
|
||||||
|
self._msg = msg
|
||||||
|
self._session_key = session_key
|
||||||
|
self._last_reply_message: InboundMessage | None = None
|
||||||
|
channel = consumer._get_channel(msg.channel)
|
||||||
|
self.capabilities = (
|
||||||
|
channel.capabilities if channel is not None else ChannelCapabilities()
|
||||||
|
)
|
||||||
|
self.base_metadata = msg.metadata
|
||||||
|
|
||||||
|
async def send(self, content: str, *, metadata: dict | None = None) -> bool:
|
||||||
|
await self._consumer.bus.publish_outbound(
|
||||||
|
OutboundMessage(
|
||||||
|
channel=self._msg.channel,
|
||||||
|
chat_id=self._msg.chat_id,
|
||||||
|
content=content,
|
||||||
|
metadata=metadata if metadata is not None else self._msg.metadata,
|
||||||
|
)
|
||||||
|
)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
try:
|
async def wait_reply(self, *, timeout: float) -> str | None:
|
||||||
from ..config.settings import HITL_SHELL_TOOLS, load_config
|
reply = await self._consumer._reply_registry.wait_event(
|
||||||
|
self._session_key, timeout
|
||||||
|
)
|
||||||
|
if reply is None:
|
||||||
|
self._last_reply_message = None
|
||||||
|
return None
|
||||||
|
self._last_reply_message = (
|
||||||
|
reply.context if isinstance(reply.context, InboundMessage) else None
|
||||||
|
)
|
||||||
|
return reply.content
|
||||||
|
|
||||||
cfg = load_config()
|
def take_reply_context(self) -> InboundMessage | None:
|
||||||
except Exception:
|
"""Consume the last inbound reply context captured by ``wait_reply``."""
|
||||||
return False # fail-closed
|
msg = self._last_reply_message
|
||||||
|
self._last_reply_message = None
|
||||||
if cfg.auto_approve:
|
return msg
|
||||||
return True
|
|
||||||
|
|
||||||
shell_allow_list = (
|
|
||||||
[s.strip() for s in cfg.shell_allow_list.split(",") if s.strip()]
|
|
||||||
if cfg.shell_allow_list
|
|
||||||
else []
|
|
||||||
)
|
|
||||||
|
|
||||||
for req in action_requests:
|
|
||||||
name = req.get("name", "")
|
|
||||||
if name not in HITL_SHELL_TOOLS:
|
|
||||||
continue
|
|
||||||
args = req.get("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):
|
|
||||||
return False
|
|
||||||
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.
|
|
||||||
"""
|
|
||||||
lines = ["\u26a0\ufe0f Approval Required\n"]
|
|
||||||
for i, req in enumerate(action_requests, 1):
|
|
||||||
name = req.get("name", "")
|
|
||||||
args = req.get("args", {})
|
|
||||||
if isinstance(args, dict):
|
|
||||||
command = args.get("command", args.get("path", ""))
|
|
||||||
else:
|
|
||||||
command = ""
|
|
||||||
if command:
|
|
||||||
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)")
|
|
||||||
return "\n".join(lines)
|
|
||||||
|
|
||||||
|
|
||||||
def _parse_approval_reply(text: str) -> str | None:
|
|
||||||
"""Parse a channel user's reply as an approval decision.
|
|
||||||
|
|
||||||
Returns "approve", "reject", "auto", or None if not recognized.
|
|
||||||
"""
|
|
||||||
t = text.strip().lower()
|
|
||||||
if t in ("1", "y", "yes", "approve", "ok"):
|
|
||||||
return "approve"
|
|
||||||
if t in ("2", "n", "no", "reject"):
|
|
||||||
return "reject"
|
|
||||||
if t in ("3", "a", "auto", "approve all"):
|
|
||||||
return "auto"
|
|
||||||
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."""
|
|
||||||
|
|
||||||
thread_id: str
|
|
||||||
action_requests: list
|
|
||||||
event: asyncio.Event # set when user replies
|
|
||||||
decision: str | None = None # "approve", "reject", "auto"
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class _PendingAskUserReply:
|
|
||||||
"""Stored state for a pending ask_user question awaiting channel user reply."""
|
|
||||||
|
|
||||||
event: asyncio.Event # set when user replies
|
|
||||||
reply: str | None = None # raw reply text
|
|
||||||
|
|
||||||
|
|
||||||
class InboundConsumer:
|
class InboundConsumer:
|
||||||
@@ -310,12 +253,12 @@ class InboundConsumer:
|
|||||||
# Metrics
|
# Metrics
|
||||||
self._metrics = ConsumerMetrics()
|
self._metrics = ConsumerMetrics()
|
||||||
|
|
||||||
# HITL: pending interrupts per session_key, and auto-approve sessions
|
# Interaction engine state: one reply registry (routes the next
|
||||||
self._pending_interrupts: dict[str, _PendingInterrupt] = {}
|
# message from a chat into a waiting prompt) and one approval
|
||||||
self._auto_approve_sessions: set[str] = set()
|
# policy (config rule + session "Approve all" grants), shared by
|
||||||
|
# the ask_user and HITL flows via ``channels.interaction``.
|
||||||
# ask_user: pending reply per session_key
|
self._reply_registry = PendingReplyRegistry()
|
||||||
self._pending_ask_user_replies: dict[str, _PendingAskUserReply] = {}
|
self._approval_policy = ApprovalPolicy()
|
||||||
|
|
||||||
async def _get_thread_id(self, sender_id: str) -> str:
|
async def _get_thread_id(self, sender_id: str) -> str:
|
||||||
"""Get or create a thread ID for the given sender.
|
"""Get or create a thread ID for the given sender.
|
||||||
@@ -428,8 +371,6 @@ class InboundConsumer:
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
channel = self._get_channel(msg.channel)
|
|
||||||
thread_id = await self._get_thread_id(msg.sender_id)
|
|
||||||
session_key = msg.session_key # "channel:chat_id"
|
session_key = msg.session_key # "channel:chat_id"
|
||||||
|
|
||||||
# Lazily create per-chat lock; evict stale locks when too many
|
# Lazily create per-chat lock; evict stale locks when too many
|
||||||
@@ -440,29 +381,39 @@ class InboundConsumer:
|
|||||||
|
|
||||||
self._metrics.total_processed += 1
|
self._metrics.total_processed += 1
|
||||||
|
|
||||||
# ask_user: check if this message is a reply to a pending question.
|
# Reply interception: if a prompt (ask_user question or HITL
|
||||||
# Must be checked BEFORE HITL approval — any text is a valid answer.
|
# approval) is waiting on this chat, hand it this message instead
|
||||||
if session_key in self._pending_ask_user_replies:
|
# of starting a fresh agent turn. The engine parses it (stop /
|
||||||
pending_ask = self._pending_ask_user_replies[session_key]
|
# cancel / choice / approval grammar), so the registry only routes
|
||||||
pending_ask.reply = msg.content
|
# text plus the original inbound context — one path for both flows.
|
||||||
pending_ask.event.set()
|
if self._reply_registry.try_resolve(session_key, msg.content, context=msg):
|
||||||
return # consumed as ask_user answer
|
return
|
||||||
|
|
||||||
# HITL: check if this message is a reply to a pending approval
|
# Resolved only for real agent turns — a consumed prompt reply must
|
||||||
if session_key in self._pending_interrupts:
|
# not create a graph thread or touch the sender-session LRU.
|
||||||
pending = self._pending_interrupts[session_key]
|
channel = self._get_channel(msg.channel)
|
||||||
decision = _parse_approval_reply(msg.content)
|
thread_id = await self._get_thread_id(msg.sender_id)
|
||||||
if decision is not None:
|
|
||||||
pending.decision = decision
|
|
||||||
pending.event.set()
|
|
||||||
return # don't process as a new agent message
|
|
||||||
# Unrecognized reply — treat as new message, cancel pending
|
|
||||||
pending.decision = "reject"
|
|
||||||
pending.event.set()
|
|
||||||
del self._pending_interrupts[session_key]
|
|
||||||
|
|
||||||
async with self._chat_locks[session_key]:
|
async with self._chat_locks[session_key]:
|
||||||
await self._stream_with_hitl(msg, channel, thread_id, session_key)
|
refeed = await self._stream_with_hitl(msg, channel, thread_id, session_key)
|
||||||
|
|
||||||
|
# An unrecognized reply to a pending approval rejects the action and
|
||||||
|
# then becomes a new agent turn. The lock was released above, so the
|
||||||
|
# previous turn has fully unwound before the refeed turn acquires it.
|
||||||
|
# Loops in case the refeed turn hits another approval that is again
|
||||||
|
# answered with unparseable text.
|
||||||
|
while refeed is not None:
|
||||||
|
channel = self._get_channel(refeed.channel)
|
||||||
|
thread_id = await self._get_thread_id(refeed.sender_id)
|
||||||
|
session_key = refeed.session_key
|
||||||
|
if session_key not in self._chat_locks:
|
||||||
|
self._chat_locks[session_key] = asyncio.Lock()
|
||||||
|
if len(self._chat_locks) > _MAX_CHAT_LOCKS:
|
||||||
|
self._evict_chat_locks()
|
||||||
|
async with self._chat_locks[session_key]:
|
||||||
|
refeed = await self._stream_with_hitl(
|
||||||
|
refeed, channel, thread_id, session_key
|
||||||
|
)
|
||||||
|
|
||||||
async def _stream_with_hitl(
|
async def _stream_with_hitl(
|
||||||
self,
|
self,
|
||||||
@@ -470,8 +421,13 @@ class InboundConsumer:
|
|||||||
channel: Channel | None,
|
channel: Channel | None,
|
||||||
thread_id: str,
|
thread_id: str,
|
||||||
session_key: str,
|
session_key: str,
|
||||||
) -> None:
|
) -> InboundMessage | None:
|
||||||
"""Stream agent events with HITL interrupt handling."""
|
"""Stream agent events with HITL interrupt handling.
|
||||||
|
|
||||||
|
Returns ``None`` normally. When a pending approval is answered
|
||||||
|
with unrecognized text, returns the intercepted inbound reply so the
|
||||||
|
caller can refeed it as a new agent turn after this one unwinds.
|
||||||
|
"""
|
||||||
from langgraph.types import Command
|
from langgraph.types import Command
|
||||||
|
|
||||||
stream_input: GraphRunInput = msg.content
|
stream_input: GraphRunInput = msg.content
|
||||||
@@ -609,108 +565,42 @@ class InboundConsumer:
|
|||||||
stream_input = Command(resume=result)
|
stream_input = Command(resume=result)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# HITL: resolve the interrupt
|
# HITL: resolve the interrupt through the shared engine.
|
||||||
|
# ``resolve_approval`` handles session/config auto-approve,
|
||||||
|
# the approval prompt (with capability-driven buttons), the
|
||||||
|
# reply wait, parsing (incl. /stop), and feedback strings.
|
||||||
action_reqs = interrupt_data.get("action_requests", [])
|
action_reqs = interrupt_data.get("action_requests", [])
|
||||||
n = len(action_reqs) or 1
|
io = _ConsumerIO(self, msg, session_key)
|
||||||
|
outcome = await resolve_approval(
|
||||||
# Session auto-approve (user previously chose "Approve all")
|
action_reqs,
|
||||||
if session_key in self._auto_approve_sessions:
|
io,
|
||||||
stream_input = Command(
|
self._approval_policy,
|
||||||
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
|
session_key,
|
||||||
)
|
timeout=HITL_APPROVAL_TIMEOUT,
|
||||||
continue
|
|
||||||
|
|
||||||
# Config auto-approve (auto_approve, non-execute, allow_list)
|
|
||||||
if _should_auto_approve(action_reqs):
|
|
||||||
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(
|
if outcome.unrecognized_reply is not None:
|
||||||
action_reqs, with_buttons=has_buttons
|
# Serve-mode policy: an unrecognized reply rejects the
|
||||||
)
|
# pending action, confirms with reject feedback, and is
|
||||||
approval_metadata = _approval_prompt_metadata(
|
# then processed as a new agent turn. The refeed is
|
||||||
msg.metadata, with_buttons=has_buttons
|
# returned to ``_handle_message`` so chat-lock ordering
|
||||||
)
|
# stays serialized.
|
||||||
await self.bus.publish_outbound(
|
await io.send(REJECTED_FEEDBACK)
|
||||||
OutboundMessage(
|
# In this flow, the final wait_reply call is exactly the
|
||||||
channel=msg.channel,
|
# unrecognized approval reply. ask_user does not read this.
|
||||||
chat_id=msg.chat_id,
|
refeed_msg = io.take_reply_context()
|
||||||
content=prompt_text,
|
if refeed_msg is None:
|
||||||
metadata=approval_metadata,
|
logger.warning(
|
||||||
)
|
"Unrecognized approval reply had no inbound context; "
|
||||||
)
|
"dropping refeed"
|
||||||
|
|
||||||
# Wait for user reply
|
|
||||||
pending = _PendingInterrupt(
|
|
||||||
thread_id=thread_id,
|
|
||||||
action_requests=action_reqs,
|
|
||||||
event=asyncio.Event(),
|
|
||||||
)
|
|
||||||
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
|
|
||||||
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"
|
|
||||||
await self.bus.publish_outbound(
|
|
||||||
OutboundMessage(
|
|
||||||
channel=msg.channel,
|
|
||||||
chat_id=msg.chat_id,
|
|
||||||
content="⏰ Approval timed out. Action rejected.",
|
|
||||||
metadata=msg.metadata,
|
|
||||||
)
|
)
|
||||||
)
|
return refeed_msg
|
||||||
else:
|
if outcome.decisions is None:
|
||||||
decision = pending.decision or "reject"
|
return None # reject / timeout / stop — end the turn
|
||||||
|
|
||||||
# Visible confirmation so the click/reply registers (QQ has no
|
from ..backends import build_hitl_resume
|
||||||
# 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":
|
stream_input = build_hitl_resume(
|
||||||
return
|
interrupt_data.get("interrupt_id"), outcome.decisions
|
||||||
|
|
||||||
if decision == "auto":
|
|
||||||
self._auto_approve_sessions.add(session_key)
|
|
||||||
|
|
||||||
stream_input = Command(
|
|
||||||
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
|
|
||||||
)
|
)
|
||||||
# continue to next HITL round
|
# continue to next HITL round
|
||||||
|
|
||||||
@@ -773,150 +663,27 @@ class InboundConsumer:
|
|||||||
|
|
||||||
# ── ask_user helpers ──
|
# ── ask_user helpers ──
|
||||||
|
|
||||||
async def _wait_for_ask_user_reply(
|
|
||||||
self,
|
|
||||||
session_key: str,
|
|
||||||
timeout: float,
|
|
||||||
) -> str | None:
|
|
||||||
"""Register a pending ask_user slot and wait for the user to reply.
|
|
||||||
|
|
||||||
Returns the raw reply text, or ``None`` on timeout.
|
|
||||||
"""
|
|
||||||
pending = _PendingAskUserReply(event=asyncio.Event())
|
|
||||||
self._pending_ask_user_replies[session_key] = pending
|
|
||||||
try:
|
|
||||||
await asyncio.wait_for(pending.event.wait(), timeout=timeout)
|
|
||||||
except TimeoutError:
|
|
||||||
pass
|
|
||||||
finally:
|
|
||||||
self._pending_ask_user_replies.pop(session_key, None)
|
|
||||||
return pending.reply
|
|
||||||
|
|
||||||
async def _resolve_ask_user(
|
async def _resolve_ask_user(
|
||||||
self,
|
self,
|
||||||
msg: InboundMessage,
|
msg: InboundMessage,
|
||||||
event_data: dict,
|
event_data: dict,
|
||||||
session_key: str,
|
session_key: str,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""Handle an ask_user interrupt: send questions to channel, collect answers.
|
"""Handle an ask_user interrupt via the shared engine.
|
||||||
|
|
||||||
Mirrors the logic of ``cli.channel.channel_ask_user_prompt`` but runs
|
Delegates the whole question/answer flow (prompt formatting, choice
|
||||||
fully async inside the consumer event loop.
|
+ "Other" grammar, ``/stop`` handling) to
|
||||||
|
:func:`channels.interaction.resolve_ask_user` over a
|
||||||
|
:class:`_ConsumerIO` adapter, so serve mode and the CLI bridge
|
||||||
|
cannot drift.
|
||||||
|
|
||||||
Returns a dict suitable for ``Command(resume=...)``:
|
Returns a dict suitable for ``Command(resume=...)``:
|
||||||
``{"answers": [...], "status": "answered"}`` or
|
``{"answers": [...], "status": "answered"}`` or
|
||||||
``{"status": "cancelled"}``.
|
``{"status": "cancelled"}``.
|
||||||
"""
|
"""
|
||||||
questions = event_data.get("questions", [])
|
questions = event_data.get("questions", [])
|
||||||
if not questions:
|
io = _ConsumerIO(self, msg, session_key)
|
||||||
return {"answers": [], "status": "answered"}
|
return await resolve_ask_user(questions, io, timeout=ASK_USER_TIMEOUT)
|
||||||
|
|
||||||
total = len(questions)
|
|
||||||
answers: list[str] = []
|
|
||||||
|
|
||||||
for i, q in enumerate(questions):
|
|
||||||
q_text = q.get("question", "")
|
|
||||||
q_type = q.get("type", "text")
|
|
||||||
required = q.get("required", True)
|
|
||||||
|
|
||||||
# -- Format question header --
|
|
||||||
if total == 1:
|
|
||||||
header = "\u2753 Quick check-in from EvoScientist\n"
|
|
||||||
else:
|
|
||||||
header = f"\u2753 Question {i + 1}/{total}\n"
|
|
||||||
|
|
||||||
lines: list[str] = [header, f"{i + 1}. {q_text}"]
|
|
||||||
if not required:
|
|
||||||
lines[-1] += " (optional)"
|
|
||||||
|
|
||||||
if q_type == "multiple_choice":
|
|
||||||
choices = q.get("choices", [])
|
|
||||||
for j, choice in enumerate(choices):
|
|
||||||
label = choice.get("value", str(choice))
|
|
||||||
letter = chr(ord("A") + j)
|
|
||||||
lines.append(f" {letter}. {label}")
|
|
||||||
other_letter = chr(ord("A") + len(choices))
|
|
||||||
lines.append(f" {other_letter}. Other")
|
|
||||||
letters = "/".join(chr(ord("A") + k) for k in range(len(choices) + 1))
|
|
||||||
lines.append(f"\nReply with a letter ({letters}), or 'cancel'.")
|
|
||||||
else:
|
|
||||||
skip_hint = " Leave empty to skip." if not required else ""
|
|
||||||
lines.append(f"\nReply with your answer, or 'cancel'.{skip_hint}")
|
|
||||||
|
|
||||||
# -- Send question --
|
|
||||||
await self.bus.publish_outbound(
|
|
||||||
OutboundMessage(
|
|
||||||
channel=msg.channel,
|
|
||||||
chat_id=msg.chat_id,
|
|
||||||
content="\n".join(lines),
|
|
||||||
metadata=msg.metadata,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
# -- Wait for user reply --
|
|
||||||
reply = await self._wait_for_ask_user_reply(
|
|
||||||
session_key,
|
|
||||||
_ASK_USER_TIMEOUT,
|
|
||||||
)
|
|
||||||
|
|
||||||
if not reply:
|
|
||||||
await self.bus.publish_outbound(
|
|
||||||
OutboundMessage(
|
|
||||||
channel=msg.channel,
|
|
||||||
chat_id=msg.chat_id,
|
|
||||||
content="\u23f0 Response timed out.",
|
|
||||||
metadata=msg.metadata,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
return {"status": "cancelled"}
|
|
||||||
|
|
||||||
raw = reply.strip()
|
|
||||||
if raw.lower() == "cancel":
|
|
||||||
return {"status": "cancelled"}
|
|
||||||
|
|
||||||
# -- Parse answer --
|
|
||||||
if q_type == "multiple_choice":
|
|
||||||
choices = q.get("choices", [])
|
|
||||||
other_letter = chr(ord("A") + len(choices))
|
|
||||||
if len(raw) == 1 and raw.upper() == other_letter:
|
|
||||||
# "Other" selected — ask for free-form input
|
|
||||||
await self.bus.publish_outbound(
|
|
||||||
OutboundMessage(
|
|
||||||
channel=msg.channel,
|
|
||||||
chat_id=msg.chat_id,
|
|
||||||
content="Please type your answer:",
|
|
||||||
metadata=msg.metadata,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
other_reply = await self._wait_for_ask_user_reply(
|
|
||||||
session_key,
|
|
||||||
_ASK_USER_TIMEOUT,
|
|
||||||
)
|
|
||||||
if not other_reply:
|
|
||||||
await self.bus.publish_outbound(
|
|
||||||
OutboundMessage(
|
|
||||||
channel=msg.channel,
|
|
||||||
chat_id=msg.chat_id,
|
|
||||||
content="\u23f0 Response timed out.",
|
|
||||||
metadata=msg.metadata,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
return {"status": "cancelled"}
|
|
||||||
if other_reply.strip().lower() == "cancel":
|
|
||||||
return {"status": "cancelled"}
|
|
||||||
answers.append(other_reply.strip())
|
|
||||||
elif len(raw) == 1 and raw.upper().isalpha():
|
|
||||||
idx = ord(raw.upper()) - ord("A")
|
|
||||||
if 0 <= idx < len(choices):
|
|
||||||
answers.append(choices[idx].get("value", raw))
|
|
||||||
else:
|
|
||||||
answers.append(raw)
|
|
||||||
else:
|
|
||||||
answers.append(raw)
|
|
||||||
else:
|
|
||||||
answers.append(raw)
|
|
||||||
|
|
||||||
return {"answers": answers, "status": "answered"}
|
|
||||||
|
|
||||||
# ── internal ──
|
# ── internal ──
|
||||||
|
|
||||||
|
|||||||
@@ -35,7 +35,11 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
|||||||
capabilities = DINGTALK_CAPS
|
capabilities = DINGTALK_CAPS
|
||||||
name = "dingtalk"
|
name = "dingtalk"
|
||||||
_ready_attrs = ("_http_client", "_access_token")
|
_ready_attrs = ("_http_client", "_access_token")
|
||||||
_non_retryable_patterns = ("invalidauthentication", "forbidden", "40014")
|
_non_retryable_patterns = (
|
||||||
|
*Channel._non_retryable_patterns,
|
||||||
|
"invalidauthentication",
|
||||||
|
"40014",
|
||||||
|
)
|
||||||
_mention_pattern = r"@\S+\s*"
|
_mention_pattern = r"@\S+\s*"
|
||||||
_mention_strip_count = 1
|
_mention_strip_count = 1
|
||||||
|
|
||||||
|
|||||||
@@ -208,6 +208,25 @@ class DiscordChannel(Channel):
|
|||||||
return str(self._client.user.id)
|
return str(self._client.user.id)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
# ── Retry error code extraction (override base) ─────────────────
|
||||||
|
|
||||||
|
def _extract_status_code(self, exc: Exception) -> int | None:
|
||||||
|
"""Extract HTTP status code from discord.HTTPException or fallback to base."""
|
||||||
|
import discord
|
||||||
|
|
||||||
|
if isinstance(exc, discord.HTTPException):
|
||||||
|
return exc.status
|
||||||
|
return super()._extract_status_code(exc)
|
||||||
|
|
||||||
|
def _extract_retry_delay(self, exc: Exception) -> float | None:
|
||||||
|
"""Honor ``discord.RateLimited``, raised when a 429 exceeds
|
||||||
|
``max_ratelimit_timeout`` and discord.py stops retrying internally."""
|
||||||
|
import discord
|
||||||
|
|
||||||
|
if isinstance(exc, discord.RateLimited):
|
||||||
|
return exc.retry_after
|
||||||
|
return super()._extract_retry_delay(exc)
|
||||||
|
|
||||||
# ── Inbound ─────────────────────────────────────────────────────
|
# ── Inbound ─────────────────────────────────────────────────────
|
||||||
|
|
||||||
async def _on_message(self, message) -> None:
|
async def _on_message(self, message) -> None:
|
||||||
|
|||||||
@@ -73,7 +73,12 @@ class EmailChannel(Channel, PollingMixin):
|
|||||||
name = "email"
|
name = "email"
|
||||||
|
|
||||||
capabilities = EMAIL_CAPS
|
capabilities = EMAIL_CAPS
|
||||||
_non_retryable_patterns = ("auth", "login", "credential")
|
_non_retryable_patterns = (
|
||||||
|
*Channel._non_retryable_patterns,
|
||||||
|
"auth",
|
||||||
|
"login",
|
||||||
|
"credential",
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(self, config: EmailConfig):
|
def __init__(self, config: EmailConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ async def validate_email_imap(
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_running_loop()
|
||||||
|
|
||||||
def _check():
|
def _check():
|
||||||
try:
|
try:
|
||||||
@@ -62,7 +62,7 @@ async def validate_email_smtp(
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_running_loop()
|
||||||
|
|
||||||
def _check():
|
def _check():
|
||||||
server = None
|
server = None
|
||||||
|
|||||||
@@ -257,6 +257,7 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
|||||||
name = "feishu"
|
name = "feishu"
|
||||||
_ready_attrs = ("_http_client", "_access_token")
|
_ready_attrs = ("_http_client", "_access_token")
|
||||||
_non_retryable_patterns = (
|
_non_retryable_patterns = (
|
||||||
|
*Channel._non_retryable_patterns,
|
||||||
"app_access_token is empty", # invalid credentials
|
"app_access_token is empty", # invalid credentials
|
||||||
"10003", # invalid app_id
|
"10003", # invalid app_id
|
||||||
"10014", # invalid app_secret
|
"10014", # invalid app_secret
|
||||||
@@ -828,8 +829,18 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
|||||||
except Exception:
|
except Exception:
|
||||||
return web.Response(status=400)
|
return web.Response(status=400)
|
||||||
|
|
||||||
# ── Decrypt if encrypt_key is configured ──
|
# When encryption is configured the inbound POST MUST carry an
|
||||||
if self.config.encrypt_key and "encrypt" in body:
|
# ``encrypt`` field. A plaintext body used to skip decryption and
|
||||||
|
# reach the agent directly, defeating the encryption setup (issue
|
||||||
|
# #392). Treat a missing ``encrypt`` field on an
|
||||||
|
# encryption-configured channel as an authentication failure.
|
||||||
|
if self.config.encrypt_key:
|
||||||
|
if not isinstance(body, dict) or "encrypt" not in body:
|
||||||
|
logger.warning(
|
||||||
|
"Feishu event rejected: encrypt_key is configured but the "
|
||||||
|
"body has no 'encrypt' field (possible signature bypass)"
|
||||||
|
)
|
||||||
|
return web.Response(status=403)
|
||||||
try:
|
try:
|
||||||
body = self._decrypt_event(body["encrypt"])
|
body = self._decrypt_event(body["encrypt"])
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|||||||
@@ -150,7 +150,7 @@ class ImsgRpcClient:
|
|||||||
"params": params or {},
|
"params": params or {},
|
||||||
}
|
}
|
||||||
|
|
||||||
future: asyncio.Future = asyncio.get_event_loop().create_future()
|
future: asyncio.Future = asyncio.get_running_loop().create_future()
|
||||||
self._pending[request_id] = future
|
self._pending[request_id] = future
|
||||||
|
|
||||||
line = json.dumps(payload) + "\n"
|
line = json.dumps(payload) + "\n"
|
||||||
|
|||||||
@@ -0,0 +1,563 @@
|
|||||||
|
"""Transport-agnostic HITL and ask_user interaction engine.
|
||||||
|
|
||||||
|
The module defines the channel-side protocol shared by serve mode and the
|
||||||
|
CLI/TUI bridge: prompt formatting, reply grammar, stop handling, approval
|
||||||
|
policy, pending-reply routing, and the async engine coroutines for approval
|
||||||
|
and ask_user flows. Drivers provide transport-specific IO through
|
||||||
|
:class:`InteractionIO`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING, Protocol
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from .capabilities import ChannelCapabilities
|
||||||
|
|
||||||
|
# ── timeout constants ──────────────────────────────────────────────────
|
||||||
|
# Per-flow defaults. HITL approval is short (a yes/no gate); ask_user is
|
||||||
|
# longer because the human may need thinking time.
|
||||||
|
HITL_APPROVAL_TIMEOUT = 120.0 # seconds to wait for a HITL approval reply
|
||||||
|
ASK_USER_TIMEOUT = 300.0 # seconds to wait for an ask_user reply
|
||||||
|
|
||||||
|
# ── stop-command grammar ──────────────────────────────────────────
|
||||||
|
# Checked before reply parsing in *both* flows so a `/stop` mid-prompt
|
||||||
|
# always cancels instead of being captured as a literal answer.
|
||||||
|
_STOP_COMMANDS = frozenset(("/stop", "/cancel"))
|
||||||
|
|
||||||
|
# ── feedback strings ─────────────────────────────────────────
|
||||||
|
# Visible confirmations so a click/reply registers on channels without a
|
||||||
|
# message-recall API (e.g. QQ C2C).
|
||||||
|
APPROVED_FEEDBACK = "✅ Approved"
|
||||||
|
APPROVED_AUTO_FEEDBACK = "✅ Approved (auto-approving future actions)"
|
||||||
|
REJECTED_FEEDBACK = "❌ Rejected"
|
||||||
|
UNRECOGNIZED_FEEDBACK = "Unrecognized reply. Action rejected."
|
||||||
|
APPROVAL_TIMEOUT_FEEDBACK = "⏰ Approval timed out. Action rejected."
|
||||||
|
ASK_USER_TIMEOUT_FEEDBACK = "⏰ Response timed out."
|
||||||
|
OTHER_PROMPT = "Please type your answer:"
|
||||||
|
|
||||||
|
# ── stop / cancel helpers ──────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def is_slash_command(text: str | None) -> bool:
|
||||||
|
"""Whether inbound content is a slash command (a control message, not a
|
||||||
|
prompt fragment)."""
|
||||||
|
return (text or "").lstrip().startswith("/")
|
||||||
|
|
||||||
|
|
||||||
|
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 is_cancel_reply(content: str | None) -> bool:
|
||||||
|
"""Whether a reply is the literal ``cancel`` sentinel (case-insensitive)."""
|
||||||
|
return (content or "").strip().lower() == "cancel"
|
||||||
|
|
||||||
|
|
||||||
|
# ── approval reply grammar ─────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def parse_approval_reply(text: str) -> str | None:
|
||||||
|
"""Parse a channel user's reply as an approval decision.
|
||||||
|
|
||||||
|
Returns "approve", "reject", "auto", or None if not recognized.
|
||||||
|
"""
|
||||||
|
t = text.strip().lower()
|
||||||
|
if t in ("1", "y", "yes", "approve", "ok"):
|
||||||
|
return "approve"
|
||||||
|
if t in ("2", "n", "no", "reject"):
|
||||||
|
return "reject"
|
||||||
|
if t in ("3", "a", "auto", "approve all"):
|
||||||
|
return "auto"
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def approve_decisions(action_requests: list) -> list[dict]:
|
||||||
|
"""Build the ``decisions`` payload that approves every action request.
|
||||||
|
|
||||||
|
Length matches ``action_requests`` (with a floor of 1, matching the
|
||||||
|
consumer's historical ``len(...) or 1`` so an empty request list still
|
||||||
|
yields a single approve — the shape ``Command(resume=...)`` expects).
|
||||||
|
"""
|
||||||
|
n = len(action_requests) or 1
|
||||||
|
return [{"type": "approve"} for _ in range(n)]
|
||||||
|
|
||||||
|
|
||||||
|
# ── approval prompt formatting ─────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
lines = ["⚠️ Approval Required\n"]
|
||||||
|
for i, req in enumerate(action_requests, 1):
|
||||||
|
name = req.get("name", "")
|
||||||
|
args = req.get("args", {})
|
||||||
|
if isinstance(args, dict):
|
||||||
|
# deepagents 0.7.0's `delete` tool uses `file_path`, not
|
||||||
|
# `command`/`path` — without this fallback the prompt shows
|
||||||
|
# only "delete" with no target.
|
||||||
|
command = args.get("command", args.get("path", args.get("file_path", "")))
|
||||||
|
else:
|
||||||
|
command = ""
|
||||||
|
if command:
|
||||||
|
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)")
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
# ── ask_user question formatting & answer grammar ──────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def _choice_value(choice: object, fallback: str = "") -> str:
|
||||||
|
"""Normalize one ask_user choice to its display/answer string.
|
||||||
|
|
||||||
|
Choices arrive from model-produced tool args; the schema says dicts with
|
||||||
|
a ``value`` key, but nothing enforces that at runtime, so plain strings
|
||||||
|
(or anything else) must not crash the prompt.
|
||||||
|
"""
|
||||||
|
if isinstance(choice, dict):
|
||||||
|
return str(choice.get("value", fallback or choice))
|
||||||
|
return str(choice)
|
||||||
|
|
||||||
|
|
||||||
|
def format_question_prompt(question: dict, index: int, total: int) -> str:
|
||||||
|
"""Format one ask_user *question* as a channel message.
|
||||||
|
|
||||||
|
*index* is 0-based; *total* is the number of questions in the batch.
|
||||||
|
"""
|
||||||
|
q_text = question.get("question", "")
|
||||||
|
q_type = question.get("type", "text")
|
||||||
|
required = question.get("required", True)
|
||||||
|
|
||||||
|
if total == 1:
|
||||||
|
header = "❓ Quick check-in from EvoScientist\n"
|
||||||
|
else:
|
||||||
|
header = f"❓ Question {index + 1}/{total}\n"
|
||||||
|
|
||||||
|
lines: list[str] = [header, f"{index + 1}. {q_text}"]
|
||||||
|
if not required:
|
||||||
|
lines[-1] += " (optional)"
|
||||||
|
|
||||||
|
if q_type == "multiple_choice":
|
||||||
|
choices = question.get("choices", [])
|
||||||
|
for j, choice in enumerate(choices):
|
||||||
|
label = _choice_value(choice)
|
||||||
|
letter = chr(ord("A") + j)
|
||||||
|
lines.append(f" {letter}. {label}")
|
||||||
|
other_letter = chr(ord("A") + len(choices))
|
||||||
|
lines.append(f" {other_letter}. Other")
|
||||||
|
letters = "/".join(chr(ord("A") + k) for k in range(len(choices) + 1))
|
||||||
|
lines.append(f"\nReply with a letter ({letters}), or 'cancel'.")
|
||||||
|
else:
|
||||||
|
skip_hint = " Leave empty to skip." if not required else ""
|
||||||
|
lines.append(f"\nReply with your answer, or 'cancel'.{skip_hint}")
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def parse_choice_answer(raw: str, choices: list) -> tuple[str, str | None]:
|
||||||
|
"""Classify a multiple-choice reply.
|
||||||
|
|
||||||
|
Returns ``(kind, value)``:
|
||||||
|
|
||||||
|
* ``("other", None)`` — the "Other" letter was chosen; the caller must
|
||||||
|
run the free-form sub-flow (send :data:`OTHER_PROMPT`, wait again).
|
||||||
|
* ``("answer", value)`` — a resolved answer string (the chosen
|
||||||
|
choice's ``value``, or the raw text when it isn't a valid letter).
|
||||||
|
"""
|
||||||
|
other_letter = chr(ord("A") + len(choices))
|
||||||
|
if len(raw) == 1 and raw.upper() == other_letter:
|
||||||
|
return ("other", None)
|
||||||
|
if len(raw) == 1 and raw.upper().isalpha():
|
||||||
|
idx = ord(raw.upper()) - ord("A")
|
||||||
|
if 0 <= idx < len(choices):
|
||||||
|
return ("answer", _choice_value(choices[idx], raw))
|
||||||
|
return ("answer", raw)
|
||||||
|
return ("answer", raw)
|
||||||
|
|
||||||
|
|
||||||
|
# ── approval policy ────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def config_auto_approve(action_requests: list[dict]) -> bool:
|
||||||
|
"""Whether config rules alone clear every action request.
|
||||||
|
|
||||||
|
Returns True if no manual approval is needed via config: the global
|
||||||
|
``auto_approve`` flag, non-execute tools, or every shell command
|
||||||
|
resolving to :attr:`~EvoScientist.backends.ActionDecision.APPROVE` via
|
||||||
|
:func:`~EvoScientist.backends.resolve_action_decision` (token-boundary
|
||||||
|
``shell_allow_list`` match, dangerous commands never auto-cleared).
|
||||||
|
Fail-closed on config load errors.
|
||||||
|
"""
|
||||||
|
if not action_requests:
|
||||||
|
return True
|
||||||
|
|
||||||
|
try:
|
||||||
|
from ..backends import ActionDecision, resolve_action_decision
|
||||||
|
from ..config.settings import (
|
||||||
|
HITL_ALWAYS_PROMPT_TOOLS,
|
||||||
|
HITL_SHELL_TOOLS,
|
||||||
|
load_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
cfg = load_config()
|
||||||
|
except Exception:
|
||||||
|
return False # fail-closed
|
||||||
|
|
||||||
|
shell_allow_list = (
|
||||||
|
[s.strip() for s in cfg.shell_allow_list.split(",") if s.strip()]
|
||||||
|
if cfg.shell_allow_list
|
||||||
|
else []
|
||||||
|
)
|
||||||
|
|
||||||
|
for req in action_requests:
|
||||||
|
if not isinstance(req, dict):
|
||||||
|
return False # malformed request — never auto-clear
|
||||||
|
name = req.get("name", "")
|
||||||
|
if name in HITL_ALWAYS_PROMPT_TOOLS:
|
||||||
|
return False
|
||||||
|
if name not in HITL_SHELL_TOOLS:
|
||||||
|
continue
|
||||||
|
args = req.get("args", {})
|
||||||
|
command = args.get("command", "") if isinstance(args, dict) else ""
|
||||||
|
verdict = resolve_action_decision(
|
||||||
|
command,
|
||||||
|
auto_approve=cfg.auto_approve,
|
||||||
|
dangerous_mode=cfg.dangerous_mode,
|
||||||
|
allow_list=shell_allow_list,
|
||||||
|
)
|
||||||
|
if verdict.decision is not ActionDecision.APPROVE:
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
class ApprovalPolicy:
|
||||||
|
"""Auto-approve policy backed by config rules and session grants.
|
||||||
|
|
||||||
|
One instance is owned per process. The consumer keeps one on its event
|
||||||
|
loop; the CLI bridge keeps one on the bus loop.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._granted_sessions: set[str] = set()
|
||||||
|
|
||||||
|
def is_session_granted(self, session_key: str) -> bool:
|
||||||
|
"""Whether the user previously chose "Approve all" for this session."""
|
||||||
|
return session_key in self._granted_sessions
|
||||||
|
|
||||||
|
def grant_session(self, session_key: str) -> None:
|
||||||
|
"""Record an "Approve all" grant for this session."""
|
||||||
|
self._granted_sessions.add(session_key)
|
||||||
|
|
||||||
|
def clear_sessions(self) -> None:
|
||||||
|
"""Forget all session grants (test hygiene / session reset)."""
|
||||||
|
self._granted_sessions.clear()
|
||||||
|
|
||||||
|
def auto_decision(
|
||||||
|
self, session_key: str, action_requests: list[dict]
|
||||||
|
) -> list[dict] | None:
|
||||||
|
"""Return an approve-all ``decisions`` list if this can auto-resolve.
|
||||||
|
|
||||||
|
Auto-resolves when the session was granted "Approve all" or when
|
||||||
|
config rules clear every request; otherwise returns ``None`` and
|
||||||
|
the caller must prompt the user.
|
||||||
|
"""
|
||||||
|
if self.is_session_granted(session_key) or config_auto_approve(action_requests):
|
||||||
|
return approve_decisions(action_requests)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
# ── transport adapter + reply registry ─────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
class InteractionIO(Protocol):
|
||||||
|
"""One conversation partner on one channel chat.
|
||||||
|
|
||||||
|
A transport adapter: the engine coroutines below drive a human
|
||||||
|
interaction entirely through this interface, so the same protocol
|
||||||
|
logic runs over the consumer's async loop and over the CLI bus loop.
|
||||||
|
|
||||||
|
Attributes
|
||||||
|
----------
|
||||||
|
capabilities:
|
||||||
|
The channel's :class:`ChannelCapabilities` — the engine reads
|
||||||
|
``inline_buttons`` to decide whether to attach approval buttons.
|
||||||
|
base_metadata:
|
||||||
|
The default outbound metadata for this chat (echoed back on each
|
||||||
|
send unless the engine supplies richer metadata, e.g. buttons).
|
||||||
|
"""
|
||||||
|
|
||||||
|
capabilities: ChannelCapabilities
|
||||||
|
base_metadata: dict | None
|
||||||
|
|
||||||
|
async def send(self, content: str, *, metadata: dict | None = None) -> bool:
|
||||||
|
"""Send *content* to the user; return True on success."""
|
||||||
|
...
|
||||||
|
|
||||||
|
async def wait_reply(self, *, timeout: float) -> str | None:
|
||||||
|
"""Wait for the user's next reply; return None on timeout."""
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class PendingReply:
|
||||||
|
"""A pending prompt reply plus optional transport-specific context."""
|
||||||
|
|
||||||
|
content: str
|
||||||
|
context: object | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class PendingReplyRegistry:
|
||||||
|
"""Route "the next message from this chat" into a waiting coroutine.
|
||||||
|
|
||||||
|
One instance per process (the consumer owns one on its loop; the CLI
|
||||||
|
bridge owns one on the bus loop). ``register`` / ``wait`` are used by
|
||||||
|
an :class:`InteractionIO` adapter to block for a reply; the inbound
|
||||||
|
interception point calls ``try_resolve`` to hand a message to that
|
||||||
|
waiter instead of enqueuing it as a fresh turn.
|
||||||
|
|
||||||
|
Asyncio-based: register/resolve must happen on the same event loop.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._pending: dict[str, asyncio.Future[PendingReply]] = {}
|
||||||
|
|
||||||
|
def register(self, session_key: str) -> asyncio.Future[PendingReply]:
|
||||||
|
"""Create and store a future awaiting the next reply for *session_key*."""
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
# A stale waiter for the same chat should never linger; cancel it
|
||||||
|
# so its coroutine unwinds instead of hanging until timeout.
|
||||||
|
stale = self._pending.get(session_key)
|
||||||
|
if stale is not None and not stale.done():
|
||||||
|
stale.cancel()
|
||||||
|
fut: asyncio.Future[PendingReply] = loop.create_future()
|
||||||
|
self._pending[session_key] = fut
|
||||||
|
return fut
|
||||||
|
|
||||||
|
def try_resolve(
|
||||||
|
self,
|
||||||
|
session_key: str,
|
||||||
|
content: str,
|
||||||
|
*,
|
||||||
|
context: object | None = None,
|
||||||
|
) -> bool:
|
||||||
|
"""Deliver *content* to a pending waiter. Returns True if consumed."""
|
||||||
|
fut = self._pending.get(session_key)
|
||||||
|
if fut is not None and not fut.done():
|
||||||
|
fut.set_result(PendingReply(content=content, context=context))
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
def discard(self, session_key: str) -> None:
|
||||||
|
"""Drop any pending waiter for *session_key* (idempotent)."""
|
||||||
|
self._pending.pop(session_key, None)
|
||||||
|
|
||||||
|
async def wait(self, session_key: str, timeout: float) -> str | None:
|
||||||
|
"""Register, await a reply for *timeout* seconds, then clean up.
|
||||||
|
|
||||||
|
Returns the reply text, or ``None`` on timeout / cancellation.
|
||||||
|
"""
|
||||||
|
reply = await self.wait_event(session_key, timeout)
|
||||||
|
return reply.content if reply is not None else None
|
||||||
|
|
||||||
|
async def wait_event(self, session_key: str, timeout: float) -> PendingReply | None:
|
||||||
|
"""Register, await a reply event, then clean up.
|
||||||
|
|
||||||
|
Returns the full reply envelope, or ``None`` on timeout / registry
|
||||||
|
cancellation. Cancellation of the task awaiting this method propagates.
|
||||||
|
"""
|
||||||
|
fut = self.register(session_key)
|
||||||
|
try:
|
||||||
|
done, _pending = await asyncio.wait({fut}, timeout=timeout)
|
||||||
|
if not done:
|
||||||
|
fut.cancel()
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
return fut.result()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
return None
|
||||||
|
finally:
|
||||||
|
# Identity-safe: only drop *our* slot, never a newer waiter that
|
||||||
|
# re-registered on the same chat while we were unwinding.
|
||||||
|
if self._pending.get(session_key) is fut:
|
||||||
|
self._pending.pop(session_key, None)
|
||||||
|
|
||||||
|
def clear(self) -> None:
|
||||||
|
"""Cancel and forget every pending waiter (shutdown / test hygiene)."""
|
||||||
|
for fut in self._pending.values():
|
||||||
|
if not fut.done():
|
||||||
|
fut.cancel()
|
||||||
|
self._pending.clear()
|
||||||
|
|
||||||
|
def __contains__(self, session_key: str) -> bool:
|
||||||
|
return session_key in self._pending
|
||||||
|
|
||||||
|
|
||||||
|
# ── engine coroutines ──────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
async def resolve_ask_user(
|
||||||
|
questions: list[dict], io: InteractionIO, *, timeout: float = ASK_USER_TIMEOUT
|
||||||
|
) -> dict:
|
||||||
|
"""Drive an ask_user interrupt to a resume payload.
|
||||||
|
|
||||||
|
Sends each question in turn, collects answers, and handles the choice
|
||||||
|
grammar (letters + the "Other" free-form sub-flow). ``/stop`` and
|
||||||
|
``cancel`` are checked *before* parsing every reply.
|
||||||
|
|
||||||
|
Returns a dict suitable for ``Command(resume=...)``:
|
||||||
|
``{"answers": [...], "status": "answered"}`` or ``{"status": "cancelled"}``.
|
||||||
|
"""
|
||||||
|
if not questions:
|
||||||
|
return {"answers": [], "status": "answered"}
|
||||||
|
|
||||||
|
total = len(questions)
|
||||||
|
answers: list[str] = []
|
||||||
|
|
||||||
|
for i, q in enumerate(questions):
|
||||||
|
if not await io.send(format_question_prompt(q, i, total)):
|
||||||
|
return {"status": "cancelled"}
|
||||||
|
|
||||||
|
reply = await io.wait_reply(timeout=timeout)
|
||||||
|
if reply is None:
|
||||||
|
await io.send(ASK_USER_TIMEOUT_FEEDBACK)
|
||||||
|
return {"status": "cancelled"}
|
||||||
|
|
||||||
|
raw = reply.strip()
|
||||||
|
required = q.get("required", True) is not False
|
||||||
|
if raw == "":
|
||||||
|
if required:
|
||||||
|
return {"status": "cancelled"}
|
||||||
|
answers.append("")
|
||||||
|
continue
|
||||||
|
|
||||||
|
if is_stop_command(raw) or is_cancel_reply(raw):
|
||||||
|
return {"status": "cancelled"}
|
||||||
|
|
||||||
|
if q.get("type", "text") == "multiple_choice":
|
||||||
|
choices = q.get("choices", [])
|
||||||
|
kind, value = parse_choice_answer(raw, choices)
|
||||||
|
if kind == "other":
|
||||||
|
if not await io.send(OTHER_PROMPT):
|
||||||
|
return {"status": "cancelled"}
|
||||||
|
other = await io.wait_reply(timeout=timeout)
|
||||||
|
if other is None:
|
||||||
|
await io.send(ASK_USER_TIMEOUT_FEEDBACK)
|
||||||
|
return {"status": "cancelled"}
|
||||||
|
other_raw = other.strip()
|
||||||
|
if other_raw == "":
|
||||||
|
if required:
|
||||||
|
return {"status": "cancelled"}
|
||||||
|
answers.append("")
|
||||||
|
continue
|
||||||
|
if is_stop_command(other_raw) or is_cancel_reply(other_raw):
|
||||||
|
return {"status": "cancelled"}
|
||||||
|
answers.append(other_raw)
|
||||||
|
else:
|
||||||
|
answers.append(value)
|
||||||
|
else:
|
||||||
|
answers.append(raw)
|
||||||
|
|
||||||
|
return {"answers": answers, "status": "answered"}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ApprovalOutcome:
|
||||||
|
"""Result of :func:`resolve_approval`.
|
||||||
|
|
||||||
|
``decisions`` is the approve-all payload on approve/auto, or ``None``
|
||||||
|
when the action was declined (reject / timeout / stop / unrecognized).
|
||||||
|
|
||||||
|
``unrecognized_reply`` carries the raw reply text when parsing failed.
|
||||||
|
The engine centralizes *parsing* but does not decide the transport
|
||||||
|
policy for unparseable text — the drivers do: the consumer rejects the
|
||||||
|
pending action and refeeds the text as a new agent turn (a channel user
|
||||||
|
who ignores the prompt and types a fresh instruction must not lose it);
|
||||||
|
the CLI bridge sends :data:`UNRECOGNIZED_FEEDBACK` and declines.
|
||||||
|
"""
|
||||||
|
|
||||||
|
decisions: list[dict] | None = None
|
||||||
|
unrecognized_reply: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
async def resolve_approval(
|
||||||
|
action_requests: list,
|
||||||
|
io: InteractionIO,
|
||||||
|
policy: ApprovalPolicy,
|
||||||
|
session_key: str,
|
||||||
|
*,
|
||||||
|
timeout: float = HITL_APPROVAL_TIMEOUT,
|
||||||
|
) -> ApprovalOutcome:
|
||||||
|
"""Drive a HITL approval interrupt to an :class:`ApprovalOutcome`.
|
||||||
|
|
||||||
|
Auto-resolves via *policy* (session grant or config rule) without
|
||||||
|
prompting. Otherwise sends the approval prompt (with capability-driven
|
||||||
|
buttons), waits for a reply, and parses it. ``/stop`` cancels silently
|
||||||
|
(it already got its own ack from the transport's stop fast-path). An
|
||||||
|
unrecognized reply declines *without feedback* and hands the raw text
|
||||||
|
back to the driver via ``unrecognized_reply`` (see
|
||||||
|
:class:`ApprovalOutcome` for the per-driver policy).
|
||||||
|
"""
|
||||||
|
auto = policy.auto_decision(session_key, action_requests)
|
||||||
|
if auto is not None:
|
||||||
|
return ApprovalOutcome(decisions=auto)
|
||||||
|
|
||||||
|
has_buttons = bool(io.capabilities.inline_buttons)
|
||||||
|
prompt = format_approval_prompt(action_requests, with_buttons=has_buttons)
|
||||||
|
metadata = approval_prompt_metadata(io.base_metadata, with_buttons=has_buttons)
|
||||||
|
if not await io.send(prompt, metadata=metadata):
|
||||||
|
return ApprovalOutcome()
|
||||||
|
|
||||||
|
reply = await io.wait_reply(timeout=timeout)
|
||||||
|
if reply is None:
|
||||||
|
await io.send(APPROVAL_TIMEOUT_FEEDBACK)
|
||||||
|
return ApprovalOutcome()
|
||||||
|
|
||||||
|
if is_stop_command(reply):
|
||||||
|
return ApprovalOutcome()
|
||||||
|
|
||||||
|
decision = parse_approval_reply(reply)
|
||||||
|
if decision == "auto":
|
||||||
|
policy.grant_session(session_key)
|
||||||
|
await io.send(APPROVED_AUTO_FEEDBACK)
|
||||||
|
return ApprovalOutcome(decisions=approve_decisions(action_requests))
|
||||||
|
if decision == "approve":
|
||||||
|
await io.send(APPROVED_FEEDBACK)
|
||||||
|
return ApprovalOutcome(decisions=approve_decisions(action_requests))
|
||||||
|
if decision == "reject":
|
||||||
|
await io.send(REJECTED_FEEDBACK)
|
||||||
|
return ApprovalOutcome()
|
||||||
|
|
||||||
|
# Unrecognized — decline and report the raw text; the driver chooses
|
||||||
|
# the feedback / refeed policy.
|
||||||
|
return ApprovalOutcome(unrecognized_reply=reply)
|
||||||
@@ -23,6 +23,7 @@ from typing import Any
|
|||||||
from .base import RawIncoming
|
from .base import RawIncoming
|
||||||
from .bus.events import InboundMessage, OutboundMessage
|
from .bus.events import InboundMessage, OutboundMessage
|
||||||
from .debug import emit_debug_event_if
|
from .debug import emit_debug_event_if
|
||||||
|
from .interaction import is_slash_command
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -811,8 +812,21 @@ class MentionGatingMiddleware(InboundMiddleware):
|
|||||||
policy=self.require_mention,
|
policy=self.require_mention,
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
# Strip mentions from group messages
|
# A slash command's platform target belongs only to its first token;
|
||||||
if raw.is_group and self._strip_fn:
|
# preserve mentions in its arguments. Ordinary group messages may
|
||||||
|
# still carry a bot mention elsewhere and use the full-message strip.
|
||||||
|
if self._strip_fn and is_slash_command(raw.text):
|
||||||
|
text = raw.text
|
||||||
|
token_start = len(text) - len(text.lstrip())
|
||||||
|
token_end = token_start
|
||||||
|
while token_end < len(text) and not text[token_end].isspace():
|
||||||
|
token_end += 1
|
||||||
|
stripped_token = self._strip_fn(text[token_start:token_end])
|
||||||
|
raw = dataclasses.replace(
|
||||||
|
raw,
|
||||||
|
text=text[:token_start] + stripped_token + text[token_end:],
|
||||||
|
)
|
||||||
|
elif self._strip_fn and raw.is_group:
|
||||||
raw = dataclasses.replace(raw, text=self._strip_fn(raw.text))
|
raw = dataclasses.replace(raw, text=self._strip_fn(raw.text))
|
||||||
return raw
|
return raw
|
||||||
|
|
||||||
@@ -928,6 +942,12 @@ class GroupHistoryMiddleware(InboundMiddleware):
|
|||||||
# Don't drop here — let MentionGatingMiddleware handle that
|
# Don't drop here — let MentionGatingMiddleware handle that
|
||||||
return raw
|
return raw
|
||||||
|
|
||||||
|
# Slash commands must remain the leading content so channel command
|
||||||
|
# dispatchers can recognize them. Keep buffered chatter for the next
|
||||||
|
# normal mentioned message instead of injecting it ahead of a command.
|
||||||
|
if is_slash_command(raw.text):
|
||||||
|
return raw
|
||||||
|
|
||||||
# Mentioned: inject history context
|
# Mentioned: inject history context
|
||||||
history_context = self._buffer.format_context(raw.chat_id)
|
history_context = self._buffer.format_context(raw.chat_id)
|
||||||
if history_context:
|
if history_context:
|
||||||
|
|||||||
@@ -116,7 +116,7 @@ class QQChannel(Channel):
|
|||||||
|
|
||||||
capabilities = QQ_CAPS
|
capabilities = QQ_CAPS
|
||||||
_ready_attrs = ("_client", "_running")
|
_ready_attrs = ("_client", "_running")
|
||||||
_non_retryable_patterns = ()
|
_non_retryable_patterns = Channel._non_retryable_patterns
|
||||||
_mention_pattern = r"@\S+\s*"
|
_mention_pattern = r"@\S+\s*"
|
||||||
_mention_strip_count = 1
|
_mention_strip_count = 1
|
||||||
_markdown_fallback_exc_types: ClassVar[tuple[type[Exception], ...]] = (
|
_markdown_fallback_exc_types: ClassVar[tuple[type[Exception], ...]] = (
|
||||||
@@ -235,7 +235,7 @@ class QQChannel(Channel):
|
|||||||
|
|
||||||
Surfaces the click as an :class:`InboundMessage` whose ``content`` is
|
Surfaces the click as an :class:`InboundMessage` whose ``content`` is
|
||||||
the button's ``data`` verbatim — so a "1"/"approve"/… click flows
|
the button's ``data`` verbatim — so a "1"/"approve"/… click flows
|
||||||
through ``_parse_approval_reply`` exactly like a typed reply.
|
through ``parse_approval_reply`` exactly like a typed reply.
|
||||||
|
|
||||||
The click runs through inbound middleware (Dedup suppresses QQ
|
The click runs through inbound middleware (Dedup suppresses QQ
|
||||||
retries) but is published directly to the bus so the per-sender
|
retries) but is published directly to the bus so the per-sender
|
||||||
@@ -264,7 +264,7 @@ class QQChannel(Channel):
|
|||||||
triggering_msg_id = getattr(resolved, "message_id", "") or ""
|
triggering_msg_id = getattr(resolved, "message_id", "") or ""
|
||||||
|
|
||||||
# QQ may serialize non-str values; coerce. Fall back to button id
|
# 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.
|
# when no data — same path as a typed reply via parse_approval_reply.
|
||||||
button_value = str(button_data) if button_data != "" else ""
|
button_value = str(button_data) if button_data != "" else ""
|
||||||
text = button_value or button_id
|
text = button_value or button_id
|
||||||
|
|
||||||
@@ -360,7 +360,7 @@ class QQChannel(Channel):
|
|||||||
plain_text = self._plain_formatter.format(raw_text)
|
plain_text = self._plain_formatter.format(raw_text)
|
||||||
# Plain-text fallback can't carry a keyboard. Append `value=label`
|
# Plain-text fallback can't carry a keyboard. Append `value=label`
|
||||||
# pairs so the user can still type "1"/"approve"/… instead of
|
# pairs so the user can still type "1"/"approve"/… instead of
|
||||||
# tapping (`_parse_approval_reply` accepts the same values).
|
# tapping (`parse_approval_reply` accepts the same values).
|
||||||
if buttons:
|
if buttons:
|
||||||
pairs = []
|
pairs = []
|
||||||
for btn in buttons:
|
for btn in buttons:
|
||||||
|
|||||||
@@ -32,7 +32,11 @@ class SignalChannel(Channel):
|
|||||||
name = "signal"
|
name = "signal"
|
||||||
|
|
||||||
capabilities = SIGNAL_CAPS
|
capabilities = SIGNAL_CAPS
|
||||||
_non_retryable_patterns = ("unregistered", "auth")
|
_non_retryable_patterns = (
|
||||||
|
*Channel._non_retryable_patterns,
|
||||||
|
"unregistered",
|
||||||
|
"auth",
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(self, config: SignalConfig):
|
def __init__(self, config: SignalConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ async def validate_signal(
|
|||||||
return False, "phone_number is required"
|
return False, "phone_number is required"
|
||||||
|
|
||||||
# Check signal-cli binary
|
# Check signal-cli binary
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_running_loop()
|
||||||
|
|
||||||
def _check():
|
def _check():
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -19,6 +19,13 @@ class SlackConfig(BaseChannelConfig):
|
|||||||
text_chunk_limit: int = 4096
|
text_chunk_limit: int = 4096
|
||||||
|
|
||||||
|
|
||||||
|
def _slack_response_types() -> tuple[type, ...]:
|
||||||
|
from slack_sdk.web.async_slack_response import AsyncSlackResponse
|
||||||
|
from slack_sdk.web.slack_response import SlackResponse
|
||||||
|
|
||||||
|
return (SlackResponse, AsyncSlackResponse)
|
||||||
|
|
||||||
|
|
||||||
class SlackChannel(Channel):
|
class SlackChannel(Channel):
|
||||||
"""Slack channel using slack-sdk Socket Mode."""
|
"""Slack channel using slack-sdk Socket Mode."""
|
||||||
|
|
||||||
@@ -195,6 +202,45 @@ class SlackChannel(Channel):
|
|||||||
def _get_bot_identifier(self) -> str | None:
|
def _get_bot_identifier(self) -> str | None:
|
||||||
return getattr(self, "_bot_user_id", None)
|
return getattr(self, "_bot_user_id", None)
|
||||||
|
|
||||||
|
# ── Retry error code extraction (override base) ─────────────────
|
||||||
|
|
||||||
|
def _extract_status_code(self, exc: Exception) -> int | None:
|
||||||
|
"""Extract HTTP status code from SlackApiError or fallback to base."""
|
||||||
|
from slack_sdk.errors import SlackApiError
|
||||||
|
|
||||||
|
if isinstance(exc, SlackApiError) and isinstance(
|
||||||
|
exc.response, _slack_response_types()
|
||||||
|
):
|
||||||
|
return exc.response.status_code
|
||||||
|
return super()._extract_status_code(exc)
|
||||||
|
|
||||||
|
def _extract_sdk_error_code(self, exc: Exception) -> str | None:
|
||||||
|
"""Extract structured error code string from SlackApiError."""
|
||||||
|
from slack_sdk.errors import SlackApiError
|
||||||
|
|
||||||
|
if isinstance(exc, SlackApiError) and isinstance(
|
||||||
|
exc.response, _slack_response_types()
|
||||||
|
):
|
||||||
|
error = exc.response.get("error")
|
||||||
|
return error.lower() if isinstance(error, str) else None
|
||||||
|
return super()._extract_sdk_error_code(exc)
|
||||||
|
|
||||||
|
def _extract_retry_delay(self, exc: Exception) -> float | None:
|
||||||
|
"""Read Slack's ``Retry-After`` header from a SlackApiError.
|
||||||
|
|
||||||
|
``SlackResponse.headers`` is a plain ``dict`` whose key casing depends
|
||||||
|
on the HTTP client, so match the key case-insensitively (the same
|
||||||
|
approach slack_sdk's own ``RateLimitErrorRetryHandler`` takes).
|
||||||
|
"""
|
||||||
|
from slack_sdk.errors import SlackApiError
|
||||||
|
|
||||||
|
if isinstance(exc, SlackApiError):
|
||||||
|
for key, raw in exc.response.headers.items():
|
||||||
|
if key.lower() == "retry-after":
|
||||||
|
return self._parse_retry_after(raw)
|
||||||
|
return None
|
||||||
|
return super()._extract_retry_delay(exc)
|
||||||
|
|
||||||
# ── ACK Reactions ───────────────────────────────────────────────
|
# ── ACK Reactions ───────────────────────────────────────────────
|
||||||
|
|
||||||
async def _send_ack_reaction(
|
async def _send_ack_reaction(
|
||||||
|
|||||||
@@ -26,6 +26,13 @@ from .debug import emit_debug_event
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
async def _create_standalone_agent():
|
||||||
|
"""Construct the synchronous agent without blocking the channel loop."""
|
||||||
|
from ..EvoScientist import create_cli_agent
|
||||||
|
|
||||||
|
return await asyncio.to_thread(create_cli_agent)
|
||||||
|
|
||||||
|
|
||||||
def _channel_trace_enabled(channel: Channel) -> bool:
|
def _channel_trace_enabled(channel: Channel) -> bool:
|
||||||
"""Check if debug tracing is enabled on the channel."""
|
"""Check if debug tracing is enabled on the channel."""
|
||||||
try:
|
try:
|
||||||
@@ -107,10 +114,12 @@ async def _async_main(
|
|||||||
consumer: InboundConsumer | None = None
|
consumer: InboundConsumer | None = None
|
||||||
if use_agent:
|
if use_agent:
|
||||||
logger.info("Loading EvoScientist agent...")
|
logger.info("Loading EvoScientist agent...")
|
||||||
from ..EvoScientist import create_cli_agent
|
|
||||||
from ..gateway import create_runtime_gateways
|
from ..gateway import create_runtime_gateways
|
||||||
|
|
||||||
agent = create_cli_agent()
|
# Agent construction performs synchronous MCP discovery through the
|
||||||
|
# owned-runtime bridge. Keep it off this already-running channel loop
|
||||||
|
# (and avoid blocking channel health/startup work while it loads).
|
||||||
|
agent = await _create_standalone_agent()
|
||||||
runtime_gateways = create_runtime_gateways()
|
runtime_gateways = create_runtime_gateways()
|
||||||
logger.info("Agent loaded")
|
logger.info("Agent loaded")
|
||||||
|
|
||||||
@@ -151,7 +160,7 @@ async def _async_main(
|
|||||||
await channel.stop()
|
await channel.stop()
|
||||||
await manager.stop_health()
|
await manager.stop_health()
|
||||||
|
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_running_loop()
|
||||||
for sig in (signal.SIGINT, signal.SIGTERM):
|
for sig in (signal.SIGINT, signal.SIGTERM):
|
||||||
loop.add_signal_handler(
|
loop.add_signal_handler(
|
||||||
sig,
|
sig,
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime
|
from datetime import datetime, timedelta
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import ClassVar
|
from typing import ClassVar
|
||||||
|
|
||||||
@@ -34,7 +34,11 @@ class TelegramChannel(Channel):
|
|||||||
capabilities = TELEGRAM_CAPS
|
capabilities = TELEGRAM_CAPS
|
||||||
_typing_interval: float = 4.0
|
_typing_interval: float = 4.0
|
||||||
_ready_attrs = ("_app",)
|
_ready_attrs = ("_app",)
|
||||||
_non_retryable_patterns = ("parse", "can't parse")
|
_non_retryable_patterns = (
|
||||||
|
*Channel._non_retryable_patterns,
|
||||||
|
"parse",
|
||||||
|
"can't parse",
|
||||||
|
)
|
||||||
_mention_pattern = r"(?i)@{bot_id}\s*"
|
_mention_pattern = r"(?i)@{bot_id}\s*"
|
||||||
|
|
||||||
def __init__(self, config: TelegramConfig):
|
def __init__(self, config: TelegramConfig):
|
||||||
@@ -80,9 +84,7 @@ class TelegramChannel(Channel):
|
|||||||
| filters.LOCATION
|
| filters.LOCATION
|
||||||
)
|
)
|
||||||
|
|
||||||
self._app.add_handler(
|
self._app.add_handler(MessageHandler(media_filter, self._on_message))
|
||||||
MessageHandler(media_filter & ~filters.COMMAND, self._on_message)
|
|
||||||
)
|
|
||||||
|
|
||||||
await self._app.initialize()
|
await self._app.initialize()
|
||||||
# Cache bot username for @mention detection in groups
|
# Cache bot username for @mention detection in groups
|
||||||
@@ -94,12 +96,17 @@ class TelegramChannel(Channel):
|
|||||||
logger.info("Telegram channel started (polling)")
|
logger.info("Telegram channel started (polling)")
|
||||||
|
|
||||||
async def _cleanup(self) -> None:
|
async def _cleanup(self) -> None:
|
||||||
if self._app:
|
app = self._app
|
||||||
if self._app.updater and self._app.updater.running:
|
self._app = None
|
||||||
await self._app.updater.stop()
|
if app is None:
|
||||||
await self._app.stop()
|
return
|
||||||
await self._app.shutdown()
|
|
||||||
logger.info("Telegram channel stopped")
|
if app.updater and app.updater.running:
|
||||||
|
await app.updater.stop()
|
||||||
|
if app.running:
|
||||||
|
await app.stop()
|
||||||
|
await app.shutdown()
|
||||||
|
logger.info("Telegram channel stopped")
|
||||||
|
|
||||||
# ── Typing indicator (override base) ────────────────────────────
|
# ── Typing indicator (override base) ────────────────────────────
|
||||||
|
|
||||||
@@ -111,6 +118,21 @@ class TelegramChannel(Channel):
|
|||||||
action="typing",
|
action="typing",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ── Retry delay extraction (override base) ─────────────────────
|
||||||
|
|
||||||
|
def _extract_retry_delay(self, exc: Exception) -> float | None:
|
||||||
|
"""Honor Telegram flood control (``telegram.error.RetryAfter``).
|
||||||
|
|
||||||
|
``retry_after`` is an ``int`` by default and a ``timedelta`` when the
|
||||||
|
``PTB_TIMEDELTA`` opt-in is enabled.
|
||||||
|
"""
|
||||||
|
from telegram.error import RetryAfter
|
||||||
|
|
||||||
|
if isinstance(exc, RetryAfter):
|
||||||
|
ra = exc.retry_after
|
||||||
|
return ra.total_seconds() if isinstance(ra, timedelta) else float(ra)
|
||||||
|
return super()._extract_retry_delay(exc)
|
||||||
|
|
||||||
# ── Send (template method overrides) ──────────────────────────
|
# ── Send (template method overrides) ──────────────────────────
|
||||||
|
|
||||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||||
@@ -161,6 +183,21 @@ class TelegramChannel(Channel):
|
|||||||
def _get_bot_identifier(self) -> str | None:
|
def _get_bot_identifier(self) -> str | None:
|
||||||
return self._bot_username or None
|
return self._bot_username or None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _command_target(text: str) -> str | None:
|
||||||
|
"""Return a Telegram command's target username.
|
||||||
|
|
||||||
|
An empty string represents a bare command; ``None`` means the message
|
||||||
|
is not command-shaped.
|
||||||
|
"""
|
||||||
|
parts = text.lstrip().split(None, 1)
|
||||||
|
if not parts or not parts[0].startswith("/"):
|
||||||
|
return None
|
||||||
|
command_token = parts[0][1:]
|
||||||
|
if "@" not in command_token:
|
||||||
|
return ""
|
||||||
|
return command_token.rsplit("@", 1)[1].lower()
|
||||||
|
|
||||||
async def _send_ack_reaction(
|
async def _send_ack_reaction(
|
||||||
self, chat_id: str, message_id: str, emoji: str = "👀"
|
self, chat_id: str, message_id: str, emoji: str = "👀"
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -202,10 +239,19 @@ class TelegramChannel(Channel):
|
|||||||
|
|
||||||
# Detect group and mention status for centralized gating
|
# Detect group and mention status for centralized gating
|
||||||
is_group = message.chat.type in ("group", "supergroup")
|
is_group = message.chat.type in ("group", "supergroup")
|
||||||
was_mentioned = True # DM default
|
was_mentioned = not is_group
|
||||||
if is_group and self._bot_username:
|
if is_group:
|
||||||
text_check = (message.text or message.caption or "").lower()
|
text_check = (message.text or message.caption or "").lower()
|
||||||
was_mentioned = f"@{self._bot_username}" in text_check
|
command_target = self._command_target(text_check)
|
||||||
|
if command_target is not None:
|
||||||
|
# A bare command that Telegram delivered to this bot is
|
||||||
|
# actionable. Commands explicitly addressed to another bot
|
||||||
|
# must remain ignored.
|
||||||
|
was_mentioned = not command_target or (
|
||||||
|
bool(self._bot_username) and command_target == self._bot_username
|
||||||
|
)
|
||||||
|
elif self._bot_username:
|
||||||
|
was_mentioned = f"@{self._bot_username}" in text_check
|
||||||
|
|
||||||
content_parts: list[str] = []
|
content_parts: list[str] = []
|
||||||
media_paths: list[str] = []
|
media_paths: list[str] = []
|
||||||
|
|||||||
@@ -337,9 +337,21 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
|||||||
logger.info(f"WeChat callback POST received, body length={len(body)}")
|
logger.info(f"WeChat callback POST received, body length={len(body)}")
|
||||||
xml_data = parse_xml(body)
|
xml_data = parse_xml(body)
|
||||||
|
|
||||||
# If encrypted, decrypt first
|
# If encryption is configured, the inbound POST MUST carry an
|
||||||
|
# <Encrypt> element and a matching msg_signature. An unsigned body
|
||||||
|
# used to fall through to _safe_process_message and reach the agent
|
||||||
|
# regardless of credentials, which made the encryption setup
|
||||||
|
# ineffective (issue #392). Treat a missing <Encrypt> on an
|
||||||
|
# encryption-configured channel as an authentication failure.
|
||||||
encrypt = xml_data.get("Encrypt", "")
|
encrypt = xml_data.get("Encrypt", "")
|
||||||
if encrypt and self._crypto:
|
if self._crypto:
|
||||||
|
if not encrypt:
|
||||||
|
logger.warning(
|
||||||
|
"WeChat POST rejected: encryption is configured but the "
|
||||||
|
"body has no <Encrypt> element (possible signature bypass)"
|
||||||
|
)
|
||||||
|
return web.Response(status=403)
|
||||||
|
|
||||||
signature = request.query.get("msg_signature", "")
|
signature = request.query.get("msg_signature", "")
|
||||||
timestamp = request.query.get("timestamp", "")
|
timestamp = request.query.get("timestamp", "")
|
||||||
nonce = request.query.get("nonce", "")
|
nonce = request.query.get("nonce", "")
|
||||||
|
|||||||
@@ -79,6 +79,11 @@ def main():
|
|||||||
warnings.filterwarnings(
|
warnings.filterwarnings(
|
||||||
"ignore", message=".*type is unknown and inference may fail.*"
|
"ignore", message=".*type is unknown and inference may fail.*"
|
||||||
)
|
)
|
||||||
|
# v3 streaming is a deliberate choice (#268), so its beta notice is noise.
|
||||||
|
# Matched by message, not category, to keep other beta warnings visible.
|
||||||
|
warnings.filterwarnings(
|
||||||
|
"ignore", message=".*v3 streaming protocol on Pregel is experimental.*"
|
||||||
|
)
|
||||||
from ..config import load_config
|
from ..config import load_config
|
||||||
from .commands import _configure_logging
|
from .commands import _configure_logging
|
||||||
|
|
||||||
|
|||||||
@@ -59,6 +59,12 @@ sessions_app = typer.Typer(
|
|||||||
)
|
)
|
||||||
app.add_typer(sessions_app, name="sessions")
|
app.add_typer(sessions_app, name="sessions")
|
||||||
|
|
||||||
|
# Background langgraph dev server management — the explicit counterpart to
|
||||||
|
# langgraph_dev_keepalive: a server that outlives its CLI needs a first-class
|
||||||
|
# way to inspect and stop it.
|
||||||
|
server_app = typer.Typer(help="Manage the background langgraph dev server")
|
||||||
|
app.add_typer(server_app, name="server")
|
||||||
|
|
||||||
# Configure subcommand group — re-run a single onboarding section.
|
# Configure subcommand group — re-run a single onboarding section.
|
||||||
configure_app = typer.Typer(
|
configure_app = typer.Typer(
|
||||||
help=(
|
help=(
|
||||||
|
|||||||
@@ -10,6 +10,8 @@ from ..paths import new_run_dir
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from langgraph.graph.state import CompiledStateGraph
|
from langgraph.graph.state import CompiledStateGraph
|
||||||
|
|
||||||
|
from ..runtime import AsyncRuntime
|
||||||
|
|
||||||
|
|
||||||
def _shorten_path(path: str) -> str:
|
def _shorten_path(path: str) -> str:
|
||||||
"""Shorten absolute path to relative path from current directory."""
|
"""Shorten absolute path to relative path from current directory."""
|
||||||
@@ -69,6 +71,8 @@ def _load_agent(
|
|||||||
chat_model=None,
|
chat_model=None,
|
||||||
*,
|
*,
|
||||||
on_mcp_progress=None,
|
on_mcp_progress=None,
|
||||||
|
events=None,
|
||||||
|
runtime: "AsyncRuntime | None" = None,
|
||||||
) -> "CompiledStateGraph":
|
) -> "CompiledStateGraph":
|
||||||
"""Load the CLI agent with optional persistent checkpointer.
|
"""Load the CLI agent with optional persistent checkpointer.
|
||||||
|
|
||||||
@@ -83,6 +87,7 @@ def _load_agent(
|
|||||||
selects the pure (no module-global write) build path.
|
selects the pure (no module-global write) build path.
|
||||||
on_mcp_progress: Optional per-server MCP progress callback.
|
on_mcp_progress: Optional per-server MCP progress callback.
|
||||||
Signature ``(event, server_name, detail) -> None``.
|
Signature ``(event, server_name, detail) -> None``.
|
||||||
|
runtime: Optional application-scoped runtime used for MCP discovery.
|
||||||
"""
|
"""
|
||||||
from ..EvoScientist import create_cli_agent
|
from ..EvoScientist import create_cli_agent
|
||||||
|
|
||||||
@@ -92,4 +97,6 @@ def _load_agent(
|
|||||||
config=config,
|
config=config,
|
||||||
chat_model=chat_model,
|
chat_model=chat_model,
|
||||||
on_mcp_progress=on_mcp_progress,
|
on_mcp_progress=on_mcp_progress,
|
||||||
|
events=events,
|
||||||
|
runtime=runtime,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -117,6 +117,71 @@ def _enqueue(notification: AsyncTaskNotification) -> None:
|
|||||||
q.put(notification)
|
q.put(notification)
|
||||||
|
|
||||||
|
|
||||||
|
def enqueue_task_notification(notification: AsyncTaskNotification) -> None:
|
||||||
|
"""Public :class:`~EvoScientist.middleware.notifier.NotifierPort` entry point.
|
||||||
|
|
||||||
|
Route a completed-task notification onto the consumer queue. Thin wrapper
|
||||||
|
over :func:`_enqueue` so middleware can enqueue without reaching into the
|
||||||
|
module's private symbols.
|
||||||
|
"""
|
||||||
|
_enqueue(notification)
|
||||||
|
|
||||||
|
|
||||||
|
def enqueue_bg_process_notification(
|
||||||
|
*,
|
||||||
|
task_id: str,
|
||||||
|
agent_name: str,
|
||||||
|
status: str,
|
||||||
|
prompt: str = "",
|
||||||
|
origin_cli_thread_id: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Build and enqueue a background-process completion notification.
|
||||||
|
|
||||||
|
:class:`~EvoScientist.middleware.notifier.NotifierPort` entry point used by
|
||||||
|
the background middleware so it never constructs the CLI-owned
|
||||||
|
:class:`AsyncTaskNotification` itself — the ``kind="bg-process"`` tag and the
|
||||||
|
UTC ``received_at`` timestamp are filled in here.
|
||||||
|
"""
|
||||||
|
_enqueue(
|
||||||
|
AsyncTaskNotification(
|
||||||
|
task_id=task_id,
|
||||||
|
agent_name=agent_name,
|
||||||
|
status=status,
|
||||||
|
received_at=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"),
|
||||||
|
prompt=prompt,
|
||||||
|
kind="bg-process",
|
||||||
|
origin_cli_thread_id=origin_cli_thread_id,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def pre_cancel_watcher(task_id: str) -> None:
|
||||||
|
"""Cancel a stale watcher for ``task_id`` before a new run replaces it.
|
||||||
|
|
||||||
|
``update_async_task`` starts a new run on the same ``thread_id`` with
|
||||||
|
``multitask_strategy="interrupt"``, which closes the old run's stream
|
||||||
|
cleanly. Without pre-cancellation the old watcher would observe that clean
|
||||||
|
exit and enqueue a stale "success" notification before the new spawn can
|
||||||
|
replace it. Cancellation propagates ``CancelledError`` (a ``BaseException``)
|
||||||
|
which the watcher's ``except Exception:`` does not catch, so ``_enqueue``
|
||||||
|
never runs for the cancelled watcher.
|
||||||
|
|
||||||
|
No-op when there is no live watcher; swallows any error (a failed
|
||||||
|
pre-cancel only risks one stale notification, never a crashed tool call).
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
old = _watcher_by_thread.get(task_id)
|
||||||
|
if old is not None and not old.done():
|
||||||
|
old.cancel()
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Pre-cancel of stale watcher for task %s failed; a stale success "
|
||||||
|
"notification may be enqueued",
|
||||||
|
task_id,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def has_pending_notifications(current_thread_id: str | None = None) -> bool:
|
def has_pending_notifications(current_thread_id: str | None = None) -> bool:
|
||||||
"""Cheap predicate for poller idle paths — true iff there's anything to consume.
|
"""Cheap predicate for poller idle paths — true iff there's anything to consume.
|
||||||
|
|
||||||
|
|||||||
+277
-270
@@ -12,6 +12,7 @@ for the main thread to set a response via ``_set_channel_response()``.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import concurrent.futures
|
||||||
import logging
|
import logging
|
||||||
import queue
|
import queue
|
||||||
import threading
|
import threading
|
||||||
@@ -24,11 +25,25 @@ from typing import TYPE_CHECKING, Any
|
|||||||
from rich.panel import Panel
|
from rich.panel import Panel
|
||||||
from rich.text import Text
|
from rich.text import Text
|
||||||
|
|
||||||
|
from ..channels.capabilities import ChannelCapabilities
|
||||||
|
from ..channels.interaction import (
|
||||||
|
ASK_USER_TIMEOUT,
|
||||||
|
HITL_APPROVAL_TIMEOUT,
|
||||||
|
UNRECOGNIZED_FEEDBACK,
|
||||||
|
ApprovalPolicy,
|
||||||
|
InteractionIO,
|
||||||
|
PendingReplyRegistry,
|
||||||
|
is_slash_command,
|
||||||
|
is_stop_command,
|
||||||
|
resolve_approval,
|
||||||
|
resolve_ask_user,
|
||||||
|
)
|
||||||
from ..commands.base import ChannelRuntime
|
from ..commands.base import ChannelRuntime
|
||||||
from ..stream.console import console
|
from ..stream.console import console
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from ..gateway import GraphGateway
|
from ..gateway import GraphGateway
|
||||||
|
from ..runtime import AsyncRuntime
|
||||||
|
|
||||||
_channel_logger = logging.getLogger(__name__)
|
_channel_logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -60,6 +75,9 @@ _message_queue: queue.Queue[ChannelMessage] = queue.Queue()
|
|||||||
# Pending responses:
|
# Pending responses:
|
||||||
# main → bus (msg_id → {"future": Future[str], "loop": loop, "response": str|None})
|
# main → bus (msg_id → {"future": Future[str], "loop": loop, "response": str|None})
|
||||||
_pending_responses: dict[str, dict] = {}
|
_pending_responses: dict[str, dict] = {}
|
||||||
|
# Sentinel response: the command's output already reached the channel via the
|
||||||
|
# command UI, so the bus consumer must not deliver a second message.
|
||||||
|
COMMAND_OUTPUT_ALREADY_SENT = "__evosci-command-output-already-sent__"
|
||||||
_response_lock = threading.Lock()
|
_response_lock = threading.Lock()
|
||||||
|
|
||||||
_RESPONSE_TIMEOUT = 600.0
|
_RESPONSE_TIMEOUT = 600.0
|
||||||
@@ -264,6 +282,7 @@ async def dispatch_channel_slash_command(
|
|||||||
await_agent_ready: Callable[[], Awaitable[Any]] | None = None,
|
await_agent_ready: Callable[[], Awaitable[Any]] | None = None,
|
||||||
on_cmd_completed: Callable[..., Awaitable[None]] | None = None,
|
on_cmd_completed: Callable[..., Awaitable[None]] | None = None,
|
||||||
channel_runtime: ChannelRuntime | None = None,
|
channel_runtime: ChannelRuntime | None = None,
|
||||||
|
async_runtime: AsyncRuntime | None = None,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Dispatch a slash command from a channel message.
|
"""Dispatch a slash command from a channel message.
|
||||||
|
|
||||||
@@ -313,7 +332,7 @@ async def dispatch_channel_slash_command(
|
|||||||
``cli/interactive.py:1002-1030``. Headless serve passes
|
``cli/interactive.py:1002-1030``. Headless serve passes
|
||||||
``None`` since it cannot hot-swap its polling-loop agent.
|
``None`` since it cannot hot-swap its polling-loop agent.
|
||||||
"""
|
"""
|
||||||
if not msg.content.strip().startswith("/"):
|
if not is_slash_command(msg.content):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -330,6 +349,7 @@ async def dispatch_channel_slash_command(
|
|||||||
on_cmd_completed=on_cmd_completed,
|
on_cmd_completed=on_cmd_completed,
|
||||||
channel_runtime=channel_runtime,
|
channel_runtime=channel_runtime,
|
||||||
graph_gateway=graph_gateway,
|
graph_gateway=graph_gateway,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
# Last-ditch safety: any uncaught exception from inside the
|
# Last-ditch safety: any uncaught exception from inside the
|
||||||
@@ -366,6 +386,7 @@ async def _dispatch_channel_slash_impl(
|
|||||||
await_agent_ready: Callable[[], Awaitable[Any]] | None,
|
await_agent_ready: Callable[[], Awaitable[Any]] | None,
|
||||||
on_cmd_completed: Callable[..., Awaitable[None]] | None,
|
on_cmd_completed: Callable[..., Awaitable[None]] | None,
|
||||||
channel_runtime: ChannelRuntime | None,
|
channel_runtime: ChannelRuntime | None,
|
||||||
|
async_runtime: AsyncRuntime | None,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Inner body of ``dispatch_channel_slash_command``.
|
"""Inner body of ``dispatch_channel_slash_command``.
|
||||||
|
|
||||||
@@ -378,10 +399,17 @@ async def _dispatch_channel_slash_impl(
|
|||||||
from ..commands.channel_ui import ChannelCommandUI
|
from ..commands.channel_ui import ChannelCommandUI
|
||||||
from ..commands.manager import manager as cmd_manager
|
from ..commands.manager import manager as cmd_manager
|
||||||
|
|
||||||
|
# The wrapper only forwards slash-prefixed content, so an unresolved
|
||||||
|
# parse is always an unknown command — answer instead of feeding a typo
|
||||||
|
# to the agent.
|
||||||
parsed = cmd_manager.resolve(msg.content)
|
parsed = cmd_manager.resolve(msg.content)
|
||||||
if parsed is None:
|
if parsed is None:
|
||||||
# Unknown slash command — let the agent handle it (matches TUI).
|
bad_cmd = msg.content.split(None, 1)[0]
|
||||||
return False
|
_set_channel_response(
|
||||||
|
msg.msg_id,
|
||||||
|
f"Unknown command: {bad_cmd}\nType /help to see available commands.",
|
||||||
|
)
|
||||||
|
return True
|
||||||
cmd, cmd_args = parsed
|
cmd, cmd_args = parsed
|
||||||
|
|
||||||
agent_for_ctx = agent
|
agent_for_ctx = agent
|
||||||
@@ -407,6 +435,7 @@ async def _dispatch_channel_slash_impl(
|
|||||||
checkpointer=checkpointer,
|
checkpointer=checkpointer,
|
||||||
channel_runtime=channel_runtime,
|
channel_runtime=channel_runtime,
|
||||||
graph_gateway=graph_gateway,
|
graph_gateway=graph_gateway,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -418,8 +447,11 @@ async def _dispatch_channel_slash_impl(
|
|||||||
|
|
||||||
if cmd_executed:
|
if cmd_executed:
|
||||||
if ctx.command_error is not None:
|
if ctx.command_error is not None:
|
||||||
details = ctx.command_error or "(no details)"
|
if ui.sent_to_channel:
|
||||||
_set_channel_response(msg.msg_id, f"Command error: {details}")
|
_set_channel_response(msg.msg_id, COMMAND_OUTPUT_ALREADY_SENT)
|
||||||
|
else:
|
||||||
|
details = ctx.command_error or "(no details)"
|
||||||
|
_set_channel_response(msg.msg_id, f"Command error: {details}")
|
||||||
return True
|
return True
|
||||||
|
|
||||||
if on_cmd_completed is not None:
|
if on_cmd_completed is not None:
|
||||||
@@ -439,7 +471,12 @@ async def _dispatch_channel_slash_impl(
|
|||||||
f"[{msg.channel_type}: Executed command from {msg.sender}]",
|
f"[{msg.channel_type}: Executed command from {msg.sender}]",
|
||||||
"dim",
|
"dim",
|
||||||
)
|
)
|
||||||
_set_channel_response(msg.msg_id, f"Command executed: {msg.content}")
|
if ui.sent_to_channel:
|
||||||
|
# The user already saw the command's own output — a second
|
||||||
|
# "Command executed" message is just noise.
|
||||||
|
_set_channel_response(msg.msg_id, COMMAND_OUTPUT_ALREADY_SENT)
|
||||||
|
else:
|
||||||
|
_set_channel_response(msg.msg_id, f"Command executed: {msg.content}")
|
||||||
return True
|
return True
|
||||||
|
|
||||||
# ``cmd_manager.execute`` returned False (empty / unparseable input).
|
# ``cmd_manager.execute`` returned False (empty / unparseable input).
|
||||||
@@ -448,21 +485,99 @@ async def _dispatch_channel_slash_impl(
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# HITL approval intercept: bus thread ⇄ main CLI thread
|
# HITL / ask_user interaction bridge: bus loop ⇄ main CLI thread
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# When the main thread needs HITL approval from a channel user, it registers
|
# The interaction protocol itself (prompt formatting, reply grammar,
|
||||||
# a pending HITL wait for (channel, chat_id). The bus consumer checks this
|
# feedback, auto-approve policy) lives in ``channels.interaction``. Here we
|
||||||
# BEFORE normal enqueue, so the next reply from that user is intercepted.
|
# only bridge it: the whole engine coroutine runs on the bus loop via
|
||||||
|
# ``run_coroutine_threadsafe`` while the calling (main / TUI) thread blocks
|
||||||
|
# on the resulting future. Replies are routed by a single asyncio-based
|
||||||
|
# ``PendingReplyRegistry`` fed from the inbound interception point — the bus
|
||||||
|
# consumer checks it BEFORE normal enqueue, so the next reply from that chat
|
||||||
|
# is intercepted.
|
||||||
|
|
||||||
_pending_hitl: dict[str, dict] = {} # "channel:chat_id" -> {event, reply}
|
|
||||||
_hitl_lock = threading.Lock()
|
|
||||||
_hitl_auto_approve: set[str] = set() # "channel:chat_id" keys with auto-approve
|
|
||||||
|
|
||||||
_HITL_APPROVAL_TIMEOUT = 120.0 # seconds to wait for HITL approval reply
|
# Extra head-room on the outer ``.result()`` wait so the engine's own
|
||||||
_ASK_USER_TIMEOUT = (
|
# per-flow timeout always fires first and returns a clean cancelled/None
|
||||||
300.0 # seconds to wait for ask_user reply (longer for thinking time)
|
# instead of the bridge tearing the coroutine down mid-flight.
|
||||||
)
|
_ENGINE_RESULT_SLACK = 30.0
|
||||||
_STOP_COMMANDS = frozenset(("/stop", "/cancel"))
|
_ENGINE_CANCEL_SETTLE_TIMEOUT = 1.0
|
||||||
|
# Send timeout inside the bridge IO adapter (kept per-flow-independent, as
|
||||||
|
# the standalone consumer has no send timeout).
|
||||||
|
_BRIDGE_SEND_TIMEOUT = 15.0
|
||||||
|
_ASK_USER_WAITS_PER_QUESTION = 2
|
||||||
|
_ASK_USER_SENDS_PER_QUESTION = 3
|
||||||
|
_HITL_SENDS_PER_APPROVAL = 2
|
||||||
|
|
||||||
|
# One reply registry + one approval policy for the whole bridge process,
|
||||||
|
# both living on the bus loop (replacing the old ``_pending_hitl`` /
|
||||||
|
# ``_hitl_lock`` / ``_hitl_auto_approve`` module globals).
|
||||||
|
_reply_registry = PendingReplyRegistry()
|
||||||
|
_approval_policy = ApprovalPolicy()
|
||||||
|
|
||||||
|
|
||||||
|
class _BridgeIO(InteractionIO):
|
||||||
|
""":class:`InteractionIO` for the CLI bridge, running on the bus loop.
|
||||||
|
|
||||||
|
``send`` publishes outbound (bounded by :data:`_BRIDGE_SEND_TIMEOUT`);
|
||||||
|
``wait_reply`` blocks on the shared :data:`_reply_registry`. Both run on
|
||||||
|
the bus loop because the engine coroutine is scheduled there via
|
||||||
|
``run_coroutine_threadsafe`` — no per-message thread hop.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
bus: Any,
|
||||||
|
msg: ChannelMessage,
|
||||||
|
capabilities: ChannelCapabilities,
|
||||||
|
session_key: str,
|
||||||
|
) -> None:
|
||||||
|
self._bus = bus
|
||||||
|
self._msg = msg
|
||||||
|
self.capabilities = capabilities
|
||||||
|
self.base_metadata = msg.metadata
|
||||||
|
self._session_key = session_key
|
||||||
|
|
||||||
|
async def send(self, content: str, *, metadata: dict | None = None) -> bool:
|
||||||
|
from ..channels.bus.events import OutboundMessage
|
||||||
|
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(
|
||||||
|
self._bus.publish_outbound(
|
||||||
|
OutboundMessage(
|
||||||
|
channel=self._msg.channel_type,
|
||||||
|
chat_id=self._msg.chat_id,
|
||||||
|
content=content,
|
||||||
|
metadata=metadata
|
||||||
|
if metadata is not None
|
||||||
|
else self._msg.metadata or {},
|
||||||
|
)
|
||||||
|
),
|
||||||
|
timeout=_BRIDGE_SEND_TIMEOUT,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
except Exception as exc:
|
||||||
|
_channel_logger.debug("bridge send failed: %s", exc)
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def wait_reply(self, *, timeout: float) -> str | None:
|
||||||
|
return await _reply_registry.wait(self._session_key, timeout)
|
||||||
|
|
||||||
|
|
||||||
|
def _ask_user_result_timeout(question_count: int) -> float:
|
||||||
|
per_question = (
|
||||||
|
ASK_USER_TIMEOUT * _ASK_USER_WAITS_PER_QUESTION
|
||||||
|
+ _BRIDGE_SEND_TIMEOUT * _ASK_USER_SENDS_PER_QUESTION
|
||||||
|
)
|
||||||
|
return per_question * question_count + _ENGINE_RESULT_SLACK
|
||||||
|
|
||||||
|
|
||||||
|
def _hitl_result_timeout() -> float:
|
||||||
|
return (
|
||||||
|
HITL_APPROVAL_TIMEOUT
|
||||||
|
+ _BRIDGE_SEND_TIMEOUT * _HITL_SENDS_PER_APPROVAL
|
||||||
|
+ _ENGINE_RESULT_SLACK
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -596,265 +711,132 @@ def publish_to_channel_origin(thread_id: str | None, content: str) -> bool:
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
def _is_stop_command(content: str | None) -> bool:
|
def _run_engine_on_bus(coro, *, result_timeout: float, on_error):
|
||||||
"""Whether incoming content is a stop/cancel slash command."""
|
"""Run *coro* (an engine coroutine) on the bus loop and block for it.
|
||||||
return (content or "").strip().lower() in _STOP_COMMANDS
|
|
||||||
|
|
||||||
|
Schedules the coroutine on ``_bus_loop`` via ``run_coroutine_threadsafe``
|
||||||
|
and waits up to *result_timeout* seconds for it (the outer bound is the
|
||||||
|
engine's own per-flow timeout plus slack, so the engine's timeout fires
|
||||||
|
first). Returns *on_error* (a zero-arg factory) on any failure.
|
||||||
|
"""
|
||||||
|
bus_loop = _bus_loop
|
||||||
|
if bus_loop is None:
|
||||||
|
coro.close()
|
||||||
|
return on_error()
|
||||||
|
try:
|
||||||
|
fut = asyncio.run_coroutine_threadsafe(coro, bus_loop)
|
||||||
|
except Exception as exc:
|
||||||
|
coro.close()
|
||||||
|
_channel_logger.debug("interaction engine bridge failed: %s", exc)
|
||||||
|
return on_error()
|
||||||
|
|
||||||
def _register_hitl_wait(channel_type: str, chat_id: str) -> threading.Event:
|
try:
|
||||||
"""Register a pending HITL wait. Returns a threading.Event to block on."""
|
return fut.result(timeout=result_timeout)
|
||||||
key = f"{channel_type}:{chat_id}"
|
except concurrent.futures.TimeoutError as exc:
|
||||||
event = threading.Event()
|
fut.cancel()
|
||||||
with _hitl_lock:
|
try:
|
||||||
_pending_hitl[key] = {"event": event, "reply": None}
|
asyncio.run_coroutine_threadsafe(asyncio.sleep(0), bus_loop).result(
|
||||||
return event
|
timeout=_ENGINE_CANCEL_SETTLE_TIMEOUT
|
||||||
|
)
|
||||||
|
except concurrent.futures.TimeoutError:
|
||||||
def _pop_hitl_reply(channel_type: str, chat_id: str) -> str | None:
|
_channel_logger.debug("interaction engine cancellation did not settle")
|
||||||
"""Pop and return the HITL reply (or None if not set)."""
|
except Exception as settle_exc:
|
||||||
key = f"{channel_type}:{chat_id}"
|
_channel_logger.debug(
|
||||||
with _hitl_lock:
|
"interaction engine failed while settling cancellation: %s",
|
||||||
slot = _pending_hitl.pop(key, None)
|
settle_exc,
|
||||||
return slot["reply"] if slot else None
|
)
|
||||||
|
_channel_logger.debug("interaction engine bridge timed out: %s", exc)
|
||||||
|
return on_error()
|
||||||
def _try_set_hitl_reply(channel_type: str, chat_id: str, content: str) -> bool:
|
except Exception as exc:
|
||||||
"""Try to intercept a message as a HITL reply. Returns True if consumed."""
|
_channel_logger.debug("interaction engine bridge failed: %s", exc)
|
||||||
key = f"{channel_type}:{chat_id}"
|
return on_error()
|
||||||
with _hitl_lock:
|
|
||||||
slot = _pending_hitl.get(key)
|
|
||||||
if slot:
|
|
||||||
slot["reply"] = content
|
|
||||||
slot["event"].set()
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
def channel_ask_user_prompt(
|
def channel_ask_user_prompt(
|
||||||
ask_user_data: dict,
|
ask_user_data: dict,
|
||||||
msg: ChannelMessage | None = None,
|
msg: ChannelMessage | None = None,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""Format ask_user questions and collect answers from a channel user.
|
"""Collect answers to ask_user questions from a channel user.
|
||||||
|
|
||||||
If *msg* is provided, sends questions via the bus and waits for a reply.
|
Thin bridge: runs :func:`channels.interaction.resolve_ask_user` on the
|
||||||
Otherwise falls back to returning a cancelled result.
|
bus loop over a :class:`_BridgeIO` and blocks for the result. Signature
|
||||||
|
and return shape are unchanged (callers in ``interactive.py`` /
|
||||||
|
``commands.py`` / ``tui_interactive.py`` are untouched).
|
||||||
|
|
||||||
Returns:
|
Returns ``{"answers": [...], "status": "answered"}`` or
|
||||||
``{"answers": [...], "status": "answered"}`` or
|
``{"status": "cancelled"}``.
|
||||||
``{"status": "cancelled"}``.
|
|
||||||
"""
|
"""
|
||||||
from ..channels.bus.events import OutboundMessage
|
|
||||||
|
|
||||||
questions = ask_user_data.get("questions", [])
|
questions = ask_user_data.get("questions", [])
|
||||||
if not questions:
|
if not questions:
|
||||||
return {"answers": [], "status": "answered"}
|
return {"answers": [], "status": "answered"}
|
||||||
|
if msg is None or not msg.bus_ref or _bus_loop is None:
|
||||||
if msg is None or not msg.bus_ref:
|
|
||||||
return {"status": "cancelled"}
|
return {"status": "cancelled"}
|
||||||
|
|
||||||
bus_loop = _bus_loop
|
# ask_user never uses buttons; a plain capability set suffices.
|
||||||
if not bus_loop:
|
io = _BridgeIO(
|
||||||
return {"status": "cancelled"}
|
msg.bus_ref, msg, ChannelCapabilities(), _channel_message_session_key(msg)
|
||||||
|
)
|
||||||
def _send(content: str) -> bool:
|
return _run_engine_on_bus(
|
||||||
try:
|
resolve_ask_user(questions, io, timeout=ASK_USER_TIMEOUT),
|
||||||
asyncio.run_coroutine_threadsafe(
|
result_timeout=_ask_user_result_timeout(len(questions)),
|
||||||
msg.bus_ref.publish_outbound(
|
on_error=lambda: {"status": "cancelled"},
|
||||||
OutboundMessage(
|
)
|
||||||
channel=msg.channel_type,
|
|
||||||
chat_id=msg.chat_id,
|
|
||||||
content=content,
|
|
||||||
metadata=msg.metadata or {},
|
|
||||||
)
|
|
||||||
),
|
|
||||||
bus_loop,
|
|
||||||
).result(timeout=15)
|
|
||||||
return True
|
|
||||||
except Exception as exc:
|
|
||||||
_channel_logger.debug("ask_user send failed: %s", exc)
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Ask one question at a time (consistent with Rich CLI / TUI)
|
|
||||||
total = len(questions)
|
|
||||||
answers: list[str] = []
|
|
||||||
|
|
||||||
for i, q in enumerate(questions):
|
|
||||||
q_text = q.get("question", "")
|
|
||||||
q_type = q.get("type", "text")
|
|
||||||
required = q.get("required", True)
|
|
||||||
|
|
||||||
# Format single question
|
|
||||||
if total == 1:
|
|
||||||
header = "\u2753 Quick check-in from EvoScientist\n"
|
|
||||||
else:
|
|
||||||
header = f"\u2753 Question {i + 1}/{total}\n"
|
|
||||||
|
|
||||||
lines = [header, f"{i + 1}. {q_text}"]
|
|
||||||
if not required:
|
|
||||||
lines[-1] += " (optional)"
|
|
||||||
|
|
||||||
if q_type == "multiple_choice":
|
|
||||||
choices = q.get("choices", [])
|
|
||||||
for j, choice in enumerate(choices):
|
|
||||||
label = choice.get("value", str(choice))
|
|
||||||
letter = chr(ord("A") + j)
|
|
||||||
lines.append(f" {letter}. {label}")
|
|
||||||
other_letter = chr(ord("A") + len(choices))
|
|
||||||
lines.append(f" {other_letter}. Other")
|
|
||||||
lines.append(
|
|
||||||
f"\nReply with a letter ({'/'.join(chr(ord('A') + k) for k in range(len(choices) + 1))}), or 'cancel'."
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
skip_hint = " Leave empty to skip." if not required else ""
|
|
||||||
lines.append(f"\nReply with your answer, or 'cancel'.{skip_hint}")
|
|
||||||
|
|
||||||
if not _send("\n".join(lines)):
|
|
||||||
return {"status": "cancelled"}
|
|
||||||
|
|
||||||
# Wait for reply
|
|
||||||
hitl_event = _register_hitl_wait(msg.channel_type, msg.chat_id)
|
|
||||||
replied = hitl_event.wait(timeout=_ASK_USER_TIMEOUT)
|
|
||||||
reply_text = _pop_hitl_reply(msg.channel_type, msg.chat_id)
|
|
||||||
|
|
||||||
if not replied or not reply_text:
|
|
||||||
_send("\u23f0 Response timed out.")
|
|
||||||
return {"status": "cancelled"}
|
|
||||||
|
|
||||||
raw = reply_text.strip()
|
|
||||||
if _is_stop_command(raw):
|
|
||||||
return {"status": "cancelled"}
|
|
||||||
if raw.lower() == "cancel":
|
|
||||||
return {"status": "cancelled"}
|
|
||||||
|
|
||||||
# Parse answer
|
|
||||||
if q_type == "multiple_choice":
|
|
||||||
choices = q.get("choices", [])
|
|
||||||
other_letter = chr(ord("A") + len(choices))
|
|
||||||
if len(raw) == 1 and raw.upper() == other_letter:
|
|
||||||
# Other selected — ask for free-form input
|
|
||||||
if not _send("Please type your answer:"):
|
|
||||||
return {"status": "cancelled"}
|
|
||||||
hitl_event = _register_hitl_wait(msg.channel_type, msg.chat_id)
|
|
||||||
replied = hitl_event.wait(timeout=_ASK_USER_TIMEOUT)
|
|
||||||
other_text = _pop_hitl_reply(msg.channel_type, msg.chat_id)
|
|
||||||
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())
|
|
||||||
elif len(raw) == 1 and raw.upper().isalpha():
|
|
||||||
idx = ord(raw.upper()) - ord("A")
|
|
||||||
if 0 <= idx < len(choices):
|
|
||||||
answers.append(choices[idx].get("value", raw))
|
|
||||||
else:
|
|
||||||
answers.append(raw)
|
|
||||||
else:
|
|
||||||
answers.append(raw)
|
|
||||||
else:
|
|
||||||
answers.append(raw)
|
|
||||||
|
|
||||||
return {"answers": answers, "status": "answered"}
|
|
||||||
|
|
||||||
|
|
||||||
def channel_hitl_prompt(
|
def channel_hitl_prompt(
|
||||||
action_requests: list,
|
action_requests: list,
|
||||||
msg: ChannelMessage,
|
msg: ChannelMessage,
|
||||||
) -> list[dict] | None:
|
) -> list[dict] | None:
|
||||||
"""Send HITL approval prompt to channel user and wait for reply.
|
"""Resolve a HITL approval prompt with a channel user.
|
||||||
|
|
||||||
Blocking function — uses threading.Event.wait(). Safe to call from a
|
Thin bridge: runs :func:`channels.interaction.resolve_approval` on the
|
||||||
background thread (CLI channel processing or asyncio.to_thread in TUI).
|
bus loop over a :class:`_BridgeIO` and blocks for the result. Signature
|
||||||
|
and return shape are unchanged (callers are untouched). Safe to call
|
||||||
|
from a background thread (CLI channel processing / TUI ``to_thread``).
|
||||||
|
|
||||||
Returns approval decisions list on approve/auto, or None on reject/timeout.
|
Returns the approval decisions list on approve/auto, or None on
|
||||||
|
reject / unrecognized / timeout / stop.
|
||||||
"""
|
"""
|
||||||
from ..channels.bus.events import OutboundMessage
|
session_key = _channel_message_session_key(msg)
|
||||||
from ..channels.consumer import (
|
decisions = _approval_policy.auto_decision(session_key, action_requests)
|
||||||
_approval_prompt_metadata,
|
if decisions is not None:
|
||||||
_format_approval_prompt,
|
return decisions
|
||||||
_parse_approval_reply,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Check session auto-approve (set by a previous "3" reply)
|
if not (_bus_loop and msg.bus_ref):
|
||||||
session_key = f"{msg.channel_type}:{msg.chat_id}"
|
|
||||||
if session_key in _hitl_auto_approve:
|
|
||||||
return [{"type": "approve"} for _ in action_requests]
|
|
||||||
|
|
||||||
bus_loop = _bus_loop
|
|
||||||
if not (bus_loop and msg.bus_ref):
|
|
||||||
_channel_logger.debug("HITL: no bus_loop or bus_ref, rejecting")
|
_channel_logger.debug("HITL: no bus_loop or bus_ref, rejecting")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# Look up the channel instance so we can attach buttons when the channel
|
# Look up the channel instance so the engine can attach buttons when the
|
||||||
# supports `inline_buttons` (Feishu cards, QQ keyboards, …).
|
# channel supports `inline_buttons` (Feishu cards, QQ keyboards, …).
|
||||||
channel_obj = (
|
channel_obj = (
|
||||||
_manager.get_channel(msg.channel_type) if _manager is not None else None
|
_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
|
capabilities = (
|
||||||
approval_metadata = _approval_prompt_metadata(
|
channel_obj.capabilities if channel_obj is not None else ChannelCapabilities()
|
||||||
msg.metadata, with_buttons=has_buttons
|
|
||||||
)
|
)
|
||||||
|
io = _BridgeIO(msg.bus_ref, msg, capabilities, session_key)
|
||||||
|
|
||||||
def _send(content: str, *, metadata: dict | None = None) -> bool:
|
async def _hitl_flow() -> list[dict] | None:
|
||||||
"""Send a message to the channel user. Returns True on success."""
|
outcome = await resolve_approval(
|
||||||
try:
|
action_requests,
|
||||||
asyncio.run_coroutine_threadsafe(
|
io,
|
||||||
msg.bus_ref.publish_outbound(
|
_approval_policy,
|
||||||
OutboundMessage(
|
session_key,
|
||||||
channel=msg.channel_type,
|
timeout=HITL_APPROVAL_TIMEOUT,
|
||||||
chat_id=msg.chat_id,
|
)
|
||||||
content=content,
|
if outcome.unrecognized_reply is not None:
|
||||||
metadata=metadata
|
# CLI-bridge policy: an unparseable reply declines with the
|
||||||
if metadata is not None
|
# explicit notice. Only the serve-mode consumer refeeds the
|
||||||
else msg.metadata or {},
|
# text as a new turn.
|
||||||
)
|
await io.send(UNRECOGNIZED_FEEDBACK)
|
||||||
),
|
return None
|
||||||
bus_loop,
|
return outcome.decisions
|
||||||
).result(timeout=15)
|
|
||||||
return True
|
|
||||||
except Exception as exc:
|
|
||||||
_channel_logger.debug("HITL send failed: %s", exc)
|
|
||||||
return False
|
|
||||||
|
|
||||||
# 1. Send approval prompt
|
return _run_engine_on_bus(
|
||||||
prompt_text = _format_approval_prompt(action_requests, with_buttons=has_buttons)
|
_hitl_flow(),
|
||||||
if not _send(prompt_text, metadata=approval_metadata):
|
result_timeout=_hitl_result_timeout(),
|
||||||
return None
|
on_error=lambda: None,
|
||||||
|
|
||||||
# 2. Wait for channel user's reply
|
|
||||||
hitl_event = _register_hitl_wait(msg.channel_type, msg.chat_id)
|
|
||||||
replied = hitl_event.wait(timeout=_HITL_APPROVAL_TIMEOUT)
|
|
||||||
reply_text = _pop_hitl_reply(msg.channel_type, msg.chat_id)
|
|
||||||
|
|
||||||
if not replied or not reply_text:
|
|
||||||
_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 已拒绝"
|
|
||||||
if decision == "reject"
|
|
||||||
else "Unrecognized reply. Action rejected."
|
|
||||||
)
|
)
|
||||||
_send(feedback)
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -866,6 +848,11 @@ _bus_loop: asyncio.AbstractEventLoop | None = None
|
|||||||
_bus_thread: threading.Thread | None = None
|
_bus_thread: threading.Thread | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_channel_startup_results() -> list[tuple[str, bool, str]]:
|
||||||
|
"""Return the current channel startup snapshot without waiting."""
|
||||||
|
return _manager.startup_results() if _manager is not None else []
|
||||||
|
|
||||||
|
|
||||||
def _channels_is_running(channel_type: str | None = None) -> bool:
|
def _channels_is_running(channel_type: str | None = None) -> bool:
|
||||||
"""Check whether channels are running."""
|
"""Check whether channels are running."""
|
||||||
if _manager is None:
|
if _manager is None:
|
||||||
@@ -896,7 +883,7 @@ def _channels_stop(
|
|||||||
|
|
||||||
if channel_type is None:
|
if channel_type is None:
|
||||||
# Stop everything
|
# Stop everything
|
||||||
if _bus_loop and _manager:
|
if _bus_loop and _manager and not _bus_loop.is_closed():
|
||||||
try:
|
try:
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
_manager.stop_all(),
|
_manager.stop_all(),
|
||||||
@@ -935,7 +922,7 @@ def _start_channels_bus_mode(
|
|||||||
thread_id: str,
|
thread_id: str,
|
||||||
*,
|
*,
|
||||||
send_thinking: bool | None = None,
|
send_thinking: bool | None = None,
|
||||||
) -> None:
|
) -> list[tuple[str, bool, str]]:
|
||||||
"""Start all channels in bus mode with MessageBus + ChannelManager.
|
"""Start all channels in bus mode with MessageBus + ChannelManager.
|
||||||
|
|
||||||
Creates a single event loop in a daemon thread running the bus,
|
Creates a single event loop in a daemon thread running the bus,
|
||||||
@@ -968,6 +955,10 @@ def _start_channels_bus_mode(
|
|||||||
try:
|
try:
|
||||||
await mgr.start_all()
|
await mgr.start_all()
|
||||||
finally:
|
finally:
|
||||||
|
# ``start_all`` returns when all channel tasks terminate. This
|
||||||
|
# includes immediate fatal startup failures, so tear down the
|
||||||
|
# dispatcher and health server before closing the bus loop.
|
||||||
|
await mgr.stop_all()
|
||||||
consumer.cancel()
|
consumer.cancel()
|
||||||
try:
|
try:
|
||||||
await consumer
|
await consumer
|
||||||
@@ -994,6 +985,8 @@ def _start_channels_bus_mode(
|
|||||||
break
|
break
|
||||||
time.sleep(0.1)
|
time.sleep(0.1)
|
||||||
|
|
||||||
|
return mgr.startup_results(timeout=2.0)
|
||||||
|
|
||||||
|
|
||||||
def _add_channel_to_running_bus(
|
def _add_channel_to_running_bus(
|
||||||
channel_type: str,
|
channel_type: str,
|
||||||
@@ -1040,13 +1033,16 @@ async def _bus_inbound_consumer(bus, manager) -> None:
|
|||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
break
|
break
|
||||||
|
|
||||||
# /stop should preempt HITL interception so cancel works while
|
session_key = _channel_session_key(msg.channel, msg.chat_id)
|
||||||
# waiting for approvals/questions. If a HITL wait is pending,
|
|
||||||
# still release it so the blocking prompt can unwind immediately.
|
# /stop should preempt interaction interception so cancel works
|
||||||
if _is_stop_command(msg.content):
|
# while waiting for approvals/questions. If a prompt wait is
|
||||||
if _try_set_hitl_reply(msg.channel, msg.chat_id, msg.content):
|
# pending, still deliver /stop into it so the blocking engine
|
||||||
|
# unwinds immediately (it treats /stop as a clean cancel).
|
||||||
|
if is_stop_command(msg.content):
|
||||||
|
if _reply_registry.try_resolve(session_key, msg.content):
|
||||||
_channel_logger.info(
|
_channel_logger.info(
|
||||||
f"[bus] stop request released HITL wait for "
|
f"[bus] stop request released interaction wait for "
|
||||||
f"{msg.channel}:{msg.chat_id}"
|
f"{msg.channel}:{msg.chat_id}"
|
||||||
)
|
)
|
||||||
_task = asyncio.create_task(_handle_bus_message(bus, manager, msg))
|
_task = asyncio.create_task(_handle_bus_message(bus, manager, msg))
|
||||||
@@ -1054,10 +1050,12 @@ async def _bus_inbound_consumer(bus, manager) -> None:
|
|||||||
_task.add_done_callback(_tasks.discard)
|
_task.add_done_callback(_tasks.discard)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Check if this message is a HITL approval reply
|
# Reply interception sits ahead of normal enqueue — if a prompt
|
||||||
if _try_set_hitl_reply(msg.channel, msg.chat_id, msg.content):
|
# is waiting on this chat, the next message is its reply and
|
||||||
|
# must NOT be enqueued as a fresh agent turn.
|
||||||
|
if _reply_registry.try_resolve(session_key, msg.content):
|
||||||
_channel_logger.info(
|
_channel_logger.info(
|
||||||
f"[bus] HITL reply from {msg.channel}:{msg.sender_id}: "
|
f"[bus] interaction reply from {msg.channel}:{msg.sender_id}: "
|
||||||
f"{msg.content[:60]}"
|
f"{msg.content[:60]}"
|
||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
@@ -1085,7 +1083,7 @@ async def _handle_bus_message(bus, manager, msg) -> None:
|
|||||||
# Fast-path: /stop intercept. Handle on the bus task itself so we
|
# Fast-path: /stop intercept. Handle on the bus task itself so we
|
||||||
# don't deadlock behind the main-thread stream we're trying to
|
# don't deadlock behind the main-thread stream we're trying to
|
||||||
# interrupt. No typing indicator, no queue entry.
|
# interrupt. No typing indicator, no queue entry.
|
||||||
if _is_stop_command(msg.content):
|
if is_stop_command(msg.content):
|
||||||
cancelled_count, active_count = _cancel_channel_session(
|
cancelled_count, active_count = _cancel_channel_session(
|
||||||
msg.channel, msg.chat_id
|
msg.channel, msg.chat_id
|
||||||
)
|
)
|
||||||
@@ -1187,16 +1185,21 @@ async def _handle_bus_message(bus, manager, msg) -> None:
|
|||||||
return
|
return
|
||||||
|
|
||||||
response = _pop_channel_response(cm.msg_id) or "No response"
|
response = _pop_channel_response(cm.msg_id) or "No response"
|
||||||
await bus.publish_outbound(
|
if response != COMMAND_OUTPUT_ALREADY_SENT:
|
||||||
OutboundMessage(
|
await bus.publish_outbound(
|
||||||
channel=msg.channel,
|
OutboundMessage(
|
||||||
chat_id=msg.chat_id,
|
channel=msg.channel,
|
||||||
content=response,
|
chat_id=msg.chat_id,
|
||||||
reply_to=msg.message_id or None,
|
content=response,
|
||||||
metadata=msg.metadata,
|
reply_to=msg.message_id or None,
|
||||||
|
metadata=msg.metadata,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
)
|
manager.record_message(msg.channel, "sent")
|
||||||
manager.record_message(msg.channel, "sent")
|
else:
|
||||||
|
# The command UI published its own response before returning the
|
||||||
|
# sentinel, so account for that delivery without sending an ack.
|
||||||
|
manager.record_message(msg.channel, "sent")
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
_pop_channel_response(cm.msg_id, cancel_pending=True)
|
_pop_channel_response(cm.msg_id, cancel_pending=True)
|
||||||
if _channel_request_state(cm.msg_id) != "active":
|
if _channel_request_state(cm.msg_id) != "active":
|
||||||
@@ -1245,7 +1248,7 @@ def _auto_start_channel(
|
|||||||
*,
|
*,
|
||||||
send_thinking: bool | None = None,
|
send_thinking: bool | None = None,
|
||||||
runtime: ChannelRuntime | None = None,
|
runtime: ChannelRuntime | None = None,
|
||||||
) -> None:
|
) -> list[tuple[str, bool, str]]:
|
||||||
"""Start channels automatically from config (bus mode).
|
"""Start channels automatically from config (bus mode).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -1257,18 +1260,22 @@ def _auto_start_channel(
|
|||||||
is accepted for callers that don't yet pass one.
|
is accepted for callers that don't yet pass one.
|
||||||
"""
|
"""
|
||||||
if not config.channel_enabled:
|
if not config.channel_enabled:
|
||||||
return
|
return []
|
||||||
|
|
||||||
_start_channels_bus_mode(
|
results = _start_channels_bus_mode(
|
||||||
config,
|
config,
|
||||||
agent,
|
agent,
|
||||||
thread_id,
|
thread_id,
|
||||||
send_thinking=send_thinking,
|
send_thinking=send_thinking,
|
||||||
)
|
)
|
||||||
# Bind only after startup succeeds; a failure above must not leave
|
# A channel that is still starting may connect later and needs the runtime
|
||||||
# a stale runtime binding pointing at channels that never started.
|
# binding. Immediate failures must not leave a stale binding behind.
|
||||||
if runtime is not None:
|
from ..channels.channel_manager import CHANNEL_STARTUP_PENDING_DETAIL
|
||||||
|
|
||||||
|
has_active_channel = any(
|
||||||
|
ok or detail == CHANNEL_STARTUP_PENDING_DETAIL for _, ok, detail in results
|
||||||
|
)
|
||||||
|
if runtime is not None and has_active_channel:
|
||||||
runtime.bind(agent, thread_id)
|
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)
|
_print_channel_panel(results)
|
||||||
|
return results
|
||||||
|
|||||||
@@ -0,0 +1,112 @@
|
|||||||
|
"""Non-blocking bridge for streaming callbacks sent through a channel loop."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import concurrent.futures
|
||||||
|
import logging
|
||||||
|
import threading
|
||||||
|
from collections.abc import Coroutine
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
|
class PendingChannelSends:
|
||||||
|
"""Schedule channel I/O without blocking the caller's event loop.
|
||||||
|
|
||||||
|
Streaming callbacks run on the owned async runtime, while channel clients
|
||||||
|
belong to the channel bus loop. Submissions therefore only enqueue work;
|
||||||
|
the frontend settles the returned futures after streaming has unwound.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
loop: asyncio.AbstractEventLoop | None,
|
||||||
|
logger: logging.Logger,
|
||||||
|
) -> None:
|
||||||
|
self._loop = loop
|
||||||
|
self._logger = logger
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
self._pending: list[tuple[concurrent.futures.Future[Any], str, int]] = []
|
||||||
|
self._tail: concurrent.futures.Future[Any] | None = None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _close(coro: Coroutine[Any, Any, Any]) -> None:
|
||||||
|
coro.close()
|
||||||
|
|
||||||
|
async def _run_after(
|
||||||
|
self,
|
||||||
|
predecessor: concurrent.futures.Future[Any] | None,
|
||||||
|
coro: Coroutine[Any, Any, Any],
|
||||||
|
) -> Any:
|
||||||
|
if predecessor is not None:
|
||||||
|
try:
|
||||||
|
await asyncio.shield(asyncio.wrap_future(predecessor))
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
task = asyncio.current_task()
|
||||||
|
if task is not None and task.cancelling():
|
||||||
|
self._close(coro)
|
||||||
|
raise
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return await coro
|
||||||
|
|
||||||
|
def submit(
|
||||||
|
self,
|
||||||
|
coro: Coroutine[Any, Any, Any],
|
||||||
|
label: str,
|
||||||
|
timeout: int = 15,
|
||||||
|
) -> None:
|
||||||
|
"""Schedule one send and return immediately."""
|
||||||
|
if self._loop is None:
|
||||||
|
self._close(coro)
|
||||||
|
return
|
||||||
|
with self._lock:
|
||||||
|
ordered_coro = self._run_after(self._tail, coro)
|
||||||
|
try:
|
||||||
|
future = asyncio.run_coroutine_threadsafe(ordered_coro, self._loop)
|
||||||
|
except Exception as exc:
|
||||||
|
self._close(ordered_coro)
|
||||||
|
self._close(coro)
|
||||||
|
self._logger.debug("%s send failed: %s", label, exc)
|
||||||
|
return
|
||||||
|
self._tail = future
|
||||||
|
self._pending.append((future, label, timeout))
|
||||||
|
|
||||||
|
def _take_pending(
|
||||||
|
self,
|
||||||
|
) -> list[tuple[concurrent.futures.Future[Any], str, int]]:
|
||||||
|
with self._lock:
|
||||||
|
pending = self._pending
|
||||||
|
self._pending = []
|
||||||
|
return pending
|
||||||
|
|
||||||
|
def settle(self) -> None:
|
||||||
|
"""Wait for scheduled sends from a synchronous frontend thread."""
|
||||||
|
for future, label, timeout in self._take_pending():
|
||||||
|
try:
|
||||||
|
future.result(timeout=timeout)
|
||||||
|
except Exception as exc:
|
||||||
|
future.cancel()
|
||||||
|
self._logger.debug("%s send failed: %s", label, exc)
|
||||||
|
|
||||||
|
async def settle_async(self) -> None:
|
||||||
|
"""Wait for scheduled sends without blocking the frontend loop."""
|
||||||
|
pending = self._take_pending()
|
||||||
|
try:
|
||||||
|
for future, label, timeout in pending:
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(asyncio.wrap_future(future), timeout=timeout)
|
||||||
|
except TimeoutError as exc:
|
||||||
|
future.cancel()
|
||||||
|
self._logger.debug("%s send failed: %s", label, exc)
|
||||||
|
except asyncio.CancelledError as exc:
|
||||||
|
task = asyncio.current_task()
|
||||||
|
if task is not None and task.cancelling():
|
||||||
|
raise
|
||||||
|
self._logger.debug("%s send failed: %s", label, exc)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.debug("%s send failed: %s", label, exc)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
for future, _label, _timeout in pending:
|
||||||
|
future.cancel()
|
||||||
|
raise
|
||||||
+232
-107
@@ -1,6 +1,5 @@
|
|||||||
"""Typer command registrations — onboard, config, mcp, main callback."""
|
"""Typer command registrations — onboard, config, mcp, main callback."""
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import queue
|
import queue
|
||||||
@@ -12,22 +11,30 @@ from importlib.metadata import version as _pkg_version
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Annotated, Any, cast
|
from typing import TYPE_CHECKING, Annotated, Any, cast
|
||||||
|
|
||||||
|
import click
|
||||||
import typer
|
import typer
|
||||||
from rich.markup import escape
|
from rich.markup import escape
|
||||||
from rich.table import Table
|
from rich.table import Table
|
||||||
|
|
||||||
from ..commands.base import ChannelRuntime, Command, CommandContext
|
from ..commands.base import (
|
||||||
|
ChannelRuntime,
|
||||||
|
Command,
|
||||||
|
CommandContext,
|
||||||
|
active_teams_configurable_extra,
|
||||||
|
)
|
||||||
from ..gateway import (
|
from ..gateway import (
|
||||||
GraphGateway,
|
GraphGateway,
|
||||||
GraphTarget,
|
GraphTarget,
|
||||||
RunRequest,
|
RunRequest,
|
||||||
RuntimeGateways,
|
|
||||||
create_runtime_gateways,
|
|
||||||
)
|
)
|
||||||
from ..llm.context_window import DEFAULT_CONTEXT_WINDOW_FALLBACK, resolve_context_window
|
from ..llm.context_window import DEFAULT_CONTEXT_WINDOW_FALLBACK, resolve_context_window
|
||||||
from ..paths import ensure_dirs, set_active_workspace, set_workspace_root
|
from ..paths import ensure_dirs, set_active_workspace, set_workspace_root
|
||||||
|
from ..runtime import AsyncRuntime
|
||||||
from ..stream.console import console
|
from ..stream.console import console
|
||||||
from . import async_notifier
|
from . import (
|
||||||
|
async_notifier,
|
||||||
|
server_cmd, # noqa: F401 — registers `EvoSci server` commands
|
||||||
|
)
|
||||||
from ._app import app, channel_app, config_app, configure_app, mcp_app, sessions_app
|
from ._app import app, channel_app, config_app, configure_app, mcp_app, sessions_app
|
||||||
from ._constants import build_metadata
|
from ._constants import build_metadata
|
||||||
from .agent import (
|
from .agent import (
|
||||||
@@ -53,6 +60,7 @@ from .channel import (
|
|||||||
publish_to_channel_origin,
|
publish_to_channel_origin,
|
||||||
remember_channel_origin,
|
remember_channel_origin,
|
||||||
)
|
)
|
||||||
|
from .channel_sends import PendingChannelSends
|
||||||
from .mcp_ui import (
|
from .mcp_ui import (
|
||||||
_mcp_add_server_from_kwargs,
|
_mcp_add_server_from_kwargs,
|
||||||
_mcp_edit_server_fields,
|
_mcp_edit_server_fields,
|
||||||
@@ -65,6 +73,36 @@ if TYPE_CHECKING:
|
|||||||
from langgraph.graph.state import CompiledStateGraph
|
from langgraph.graph.state import CompiledStateGraph
|
||||||
|
|
||||||
from ..config import EvoScientistConfig
|
from ..config import EvoScientistConfig
|
||||||
|
from ..gateway import RuntimeGateways
|
||||||
|
|
||||||
|
|
||||||
|
_ASYNC_RUNTIME_META_KEY = "evoscientist.async_runtime"
|
||||||
|
|
||||||
|
|
||||||
|
def _close_cli_async_runtime(runtime: AsyncRuntime) -> None:
|
||||||
|
"""Close the owned runtime or surface a controlled CLI shutdown failure."""
|
||||||
|
try:
|
||||||
|
runtime.close()
|
||||||
|
except TimeoutError as exc:
|
||||||
|
click.echo(
|
||||||
|
f"Error: Async runtime shutdown did not complete: {exc}",
|
||||||
|
err=True,
|
||||||
|
)
|
||||||
|
raise click.exceptions.Exit(1) from None
|
||||||
|
|
||||||
|
|
||||||
|
def _get_cli_async_runtime(ctx: typer.Context) -> AsyncRuntime:
|
||||||
|
"""Return the application-scoped runtime owned by this CLI invocation."""
|
||||||
|
root = ctx.find_root()
|
||||||
|
runtime = root.meta.get(_ASYNC_RUNTIME_META_KEY)
|
||||||
|
if runtime is None:
|
||||||
|
runtime = AsyncRuntime()
|
||||||
|
root.meta[_ASYNC_RUNTIME_META_KEY] = runtime
|
||||||
|
root.call_on_close(lambda: _close_cli_async_runtime(runtime))
|
||||||
|
if not isinstance(runtime, AsyncRuntime): # pragma: no cover - defensive
|
||||||
|
raise RuntimeError("CLI async runtime context is invalid")
|
||||||
|
return runtime
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
# Onboard command
|
# Onboard command
|
||||||
@@ -73,6 +111,7 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
@app.command()
|
@app.command()
|
||||||
def onboard(
|
def onboard(
|
||||||
|
ctx: typer.Context,
|
||||||
skip_validation: bool = typer.Option(
|
skip_validation: bool = typer.Option(
|
||||||
False, "--skip-validation", help="Skip API key validation during setup"
|
False, "--skip-validation", help="Skip API key validation during setup"
|
||||||
),
|
),
|
||||||
@@ -201,7 +240,11 @@ def onboard(
|
|||||||
strict=non_interactive,
|
strict=non_interactive,
|
||||||
)
|
)
|
||||||
|
|
||||||
_run_onboard_cli(skip_validation=skip_validation, prompter=prompter)
|
_run_onboard_cli(
|
||||||
|
skip_validation=skip_validation,
|
||||||
|
prompter=prompter,
|
||||||
|
runtime=_get_cli_async_runtime(ctx),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -243,11 +286,21 @@ def _run_onboard_cli(**kwargs: Any) -> None:
|
|||||||
raise typer.Exit(code=1) from exc
|
raise typer.Exit(code=1) from exc
|
||||||
|
|
||||||
|
|
||||||
def _configure_section(section: str, skip_validation: bool = False) -> None:
|
def _configure_section(
|
||||||
|
section: str,
|
||||||
|
skip_validation: bool = False,
|
||||||
|
*,
|
||||||
|
runtime: AsyncRuntime | None = None,
|
||||||
|
) -> None:
|
||||||
"""Run a single onboarding section, reusing the wizard's step logic."""
|
"""Run a single onboarding section, reusing the wizard's step logic."""
|
||||||
|
kwargs: dict[str, Any] = {
|
||||||
|
"skip_validation": skip_validation,
|
||||||
|
"only_sections": {section},
|
||||||
|
}
|
||||||
|
if runtime is not None:
|
||||||
|
kwargs["runtime"] = runtime
|
||||||
_run_onboard_cli(
|
_run_onboard_cli(
|
||||||
skip_validation=skip_validation,
|
**kwargs,
|
||||||
only_sections={section},
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -325,9 +378,9 @@ def configure_latex():
|
|||||||
|
|
||||||
|
|
||||||
@configure_app.command("channels")
|
@configure_app.command("channels")
|
||||||
def configure_channels():
|
def configure_channels(ctx: typer.Context):
|
||||||
"""Re-run channels selection and per-channel configuration."""
|
"""Re-run channels selection and per-channel configuration."""
|
||||||
_configure_section("channels")
|
_configure_section("channels", runtime=_get_cli_async_runtime(ctx))
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -336,24 +389,17 @@ def configure_channels():
|
|||||||
|
|
||||||
|
|
||||||
@channel_app.command("setup")
|
@channel_app.command("setup")
|
||||||
def channel_setup():
|
def channel_setup(ctx: typer.Context):
|
||||||
"""Interactive channel configuration wizard.
|
"""Interactive channel configuration wizard.
|
||||||
|
|
||||||
Guides you through selecting and configuring messaging channels
|
Guides you through selecting and configuring messaging channels
|
||||||
(Telegram, Discord, or iMessage).
|
(Telegram, Discord, or iMessage).
|
||||||
"""
|
"""
|
||||||
import asyncio
|
|
||||||
|
|
||||||
try:
|
|
||||||
asyncio.get_event_loop()
|
|
||||||
except RuntimeError:
|
|
||||||
asyncio.set_event_loop(asyncio.new_event_loop())
|
|
||||||
|
|
||||||
from ..config import load_config, save_config
|
from ..config import load_config, save_config
|
||||||
from ..config.onboard.channels import _step_channels
|
from ..config.onboard.channels import _step_channels
|
||||||
|
|
||||||
config = load_config()
|
config = load_config()
|
||||||
updates = _step_channels(config)
|
updates = _step_channels(config, runtime=_get_cli_async_runtime(ctx))
|
||||||
if updates:
|
if updates:
|
||||||
for key, value in updates.items():
|
for key, value in updates.items():
|
||||||
setattr(config, key, value)
|
setattr(config, key, value)
|
||||||
@@ -464,7 +510,13 @@ def _ensure_async_subagent_server(config: Any, *, workspace_dir: str) -> None:
|
|||||||
state would route async sub-agent calls to a process pinned to /A
|
state would route async sub-agent calls to a process pinned to /A
|
||||||
while the main agent runs in /B.
|
while the main agent runs in /B.
|
||||||
"""
|
"""
|
||||||
from ..langgraph_dev.manager import WorkspaceMismatchError, ensure_langgraph_dev
|
from ..langgraph_dev.manager import (
|
||||||
|
_DEFAULT_HOST,
|
||||||
|
WorkspaceMismatchError,
|
||||||
|
_is_loopback_host,
|
||||||
|
ensure_langgraph_dev,
|
||||||
|
is_async_subagents_available,
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
with console.status(
|
with console.status(
|
||||||
@@ -477,6 +529,32 @@ def _ensure_async_subagent_server(config: Any, *, workspace_dir: str) -> None:
|
|||||||
console.print(f"[red]{exc}[/red]")
|
console.print(f"[red]{exc}[/red]")
|
||||||
raise typer.Exit(1) from exc
|
raise typer.Exit(1) from exc
|
||||||
|
|
||||||
|
from ..langgraph_dev import manager as _lg_manager
|
||||||
|
|
||||||
|
if _lg_manager.CONFIG_DRIFT_SINCE_LAUNCH:
|
||||||
|
console.print(
|
||||||
|
"[yellow]⚠ Config changed since the background agent server was "
|
||||||
|
"launched — async sub-agents still use the old settings. Apply "
|
||||||
|
"them with [bold]EvoSci server stop[/bold], then restart "
|
||||||
|
"EvoSci.[/yellow]"
|
||||||
|
)
|
||||||
|
|
||||||
|
# The backend is shared by every UI mode, so the exposure warning lives
|
||||||
|
# here, not just in deploy/WebUI. Gated on the server being up: warning
|
||||||
|
# about a bind that never happened would be worse than saying nothing.
|
||||||
|
bind_host = str(getattr(config, "langgraph_dev_host", _DEFAULT_HOST) or "").strip()
|
||||||
|
if (
|
||||||
|
bind_host
|
||||||
|
and not _is_loopback_host(bind_host)
|
||||||
|
and is_async_subagents_available()
|
||||||
|
):
|
||||||
|
console.print(
|
||||||
|
"[bold white on red] ⚠ PUBLIC BIND [/bold white on red] "
|
||||||
|
f"[bold red]Agent server listening on {bind_host} — no auth, and "
|
||||||
|
f"the agent can run shell. Use --host 127.0.0.1 on untrusted "
|
||||||
|
f"networks.[/bold red]"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _reconcile_autoskill_schedule(config: Any, *, workspace_dir: str) -> None:
|
def _reconcile_autoskill_schedule(config: Any, *, workspace_dir: str) -> None:
|
||||||
"""Best-effort reconciliation for EvoMemory's hidden AutoSkills cron."""
|
"""Best-effort reconciliation for EvoMemory's hidden AutoSkills cron."""
|
||||||
@@ -672,9 +750,6 @@ async def compact_conversation(
|
|||||||
Returns a structured ``CompactResult``.
|
Returns a structured ``CompactResult``.
|
||||||
"""
|
"""
|
||||||
from langchain_core.messages.utils import count_tokens_approximately
|
from langchain_core.messages.utils import count_tokens_approximately
|
||||||
from langchain_core.runnables import RunnableConfig
|
|
||||||
|
|
||||||
config: RunnableConfig = {"configurable": {"thread_id": thread_id}}
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
state_values = await graph_gateway.get_state_values(target, thread_id)
|
state_values = await graph_gateway.get_state_values(target, thread_id)
|
||||||
@@ -778,22 +853,18 @@ async def compact_conversation(
|
|||||||
# Generate summary (LLM call)
|
# Generate summary (LLM call)
|
||||||
summary = await middleware._acreate_summary(to_summarize)
|
summary = await middleware._acreate_summary(to_summarize)
|
||||||
|
|
||||||
# Inject thread_id into LangGraph contextvar so _get_thread_id() finds it
|
# Reuse the persisted _summarization_session_id (or generate one) so
|
||||||
# (compact runs outside a runnable context, so get_config() would fail
|
# history keeps appending to a single file; re-persisted below.
|
||||||
# and the middleware would generate a random "session_xxx" filename instead
|
session_id = middleware._get_session_id(state_values)
|
||||||
# of reusing the real thread_id).
|
|
||||||
from langgraph.config import var_child_runnable_config
|
|
||||||
|
|
||||||
_token = var_child_runnable_config.set(config)
|
|
||||||
|
|
||||||
# Offload old messages to backend
|
# Offload old messages to backend
|
||||||
file_path: str | None = None
|
file_path: str | None = None
|
||||||
try:
|
try:
|
||||||
file_path = await middleware._aoffload_to_backend(backend, to_summarize)
|
file_path = await middleware._aoffload_to_backend(
|
||||||
|
backend, to_summarize, session_id
|
||||||
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass # non-fatal — proceed without offloaded history
|
pass # non-fatal — proceed without offloaded history
|
||||||
finally:
|
|
||||||
var_child_runnable_config.reset(_token)
|
|
||||||
|
|
||||||
from langchain_core.messages import HumanMessage
|
from langchain_core.messages import HumanMessage
|
||||||
|
|
||||||
@@ -839,7 +910,7 @@ async def compact_conversation(
|
|||||||
await graph_gateway.update_state_values(
|
await graph_gateway.update_state_values(
|
||||||
target,
|
target,
|
||||||
thread_id,
|
thread_id,
|
||||||
{"_summarization_event": new_event},
|
{"_summarization_event": new_event, "_summarization_session_id": session_id},
|
||||||
)
|
)
|
||||||
|
|
||||||
return CompactResult(
|
return CompactResult(
|
||||||
@@ -874,7 +945,8 @@ class ServeRuntimeState:
|
|||||||
thread_id: str
|
thread_id: str
|
||||||
workspace_dir: str | None
|
workspace_dir: str | None
|
||||||
config: "EvoScientistConfig | None"
|
config: "EvoScientistConfig | None"
|
||||||
runtime_gateways: RuntimeGateways
|
runtime_gateways: "RuntimeGateways"
|
||||||
|
async_runtime: AsyncRuntime
|
||||||
resume_warning_thread_id: str | None = None
|
resume_warning_thread_id: str | None = None
|
||||||
|
|
||||||
def set_agent(
|
def set_agent(
|
||||||
@@ -969,6 +1041,7 @@ async def _apply_serve_resume_state(
|
|||||||
_load_agent,
|
_load_agent,
|
||||||
workspace_dir=new_workspace,
|
workspace_dir=new_workspace,
|
||||||
config=effective_config,
|
config=effective_config,
|
||||||
|
runtime=runtime_state.async_runtime,
|
||||||
)
|
)
|
||||||
await _sync_background_agent_server_workspace(
|
await _sync_background_agent_server_workspace(
|
||||||
effective_config,
|
effective_config,
|
||||||
@@ -1119,8 +1192,6 @@ def _serve_process_message(
|
|||||||
via the ``on_cmd_completed`` hook because the command mutates
|
via the ``on_cmd_completed`` hook because the command mutates
|
||||||
``ctx.thread_id`` / ``ctx.workspace_dir`` directly.
|
``ctx.thread_id`` / ``ctx.workspace_dir`` directly.
|
||||||
"""
|
"""
|
||||||
import asyncio
|
|
||||||
|
|
||||||
from .channel import _bus_loop
|
from .channel import _bus_loop
|
||||||
from .tui_runtime import run_streaming
|
from .tui_runtime import run_streaming
|
||||||
|
|
||||||
@@ -1139,14 +1210,10 @@ def _serve_process_message(
|
|||||||
|
|
||||||
# -- channel callback helpers (same pattern as interactive.py) --
|
# -- channel callback helpers (same pattern as interactive.py) --
|
||||||
|
|
||||||
|
pending_channel_sends = PendingChannelSends(_bus_loop, _serve_logger)
|
||||||
|
|
||||||
def _send_to_channel(coro, label: str, timeout: int = 15) -> None:
|
def _send_to_channel(coro, label: str, timeout: int = 15) -> None:
|
||||||
loop = _bus_loop
|
pending_channel_sends.submit(coro, label, timeout)
|
||||||
if not loop:
|
|
||||||
return
|
|
||||||
try:
|
|
||||||
asyncio.run_coroutine_threadsafe(coro, loop).result(timeout=timeout)
|
|
||||||
except Exception as e:
|
|
||||||
_serve_logger.debug(f"{label} send failed: {e}")
|
|
||||||
|
|
||||||
def _send_thinking(thinking: str) -> None:
|
def _send_thinking(thinking: str) -> None:
|
||||||
ch = msg.channel_ref
|
ch = msg.channel_ref
|
||||||
@@ -1196,31 +1263,15 @@ def _serve_process_message(
|
|||||||
# commands like ``/evoskills`` actually execute in serve mode instead
|
# commands like ``/evoskills`` actually execute in serve mode instead
|
||||||
# of being fed to the LLM as a plain prompt. ``await_agent_ready`` is
|
# of being fed to the LLM as a plain prompt. ``await_agent_ready`` is
|
||||||
# None because the agent is always loaded before the serve loop polls.
|
# None because the agent is always loaded before the serve loop polls.
|
||||||
# Uses a dedicated event loop (not ``asyncio.run``) so SIGINT handling
|
# Slash commands run on the application-owned runtime. The main thread
|
||||||
# installed by ``serve()`` remains authoritative — ``asyncio.run``
|
# remains the signal owner while command coroutines share one stable loop.
|
||||||
# swaps ``signal.set_wakeup_fd`` and can leave it dangling on edge
|
|
||||||
# cases, which breaks Ctrl+C between messages.
|
|
||||||
# ``set_event_loop`` is needed because some downstream commands
|
|
||||||
# (e.g. ``/install-mcp``) call ``asyncio.get_event_loop()``, which
|
|
||||||
# raises ``RuntimeError`` on Python 3.12+ when the thread has no
|
|
||||||
# current loop set. The prior loop (often ``None``) is restored in
|
|
||||||
# the ``finally`` below so subsequent messages start from a clean
|
|
||||||
# slate. Loop creation lives inside the try so an exception between
|
|
||||||
# creation and ``set_event_loop`` still closes the loop.
|
|
||||||
try:
|
try:
|
||||||
_prev_loop: asyncio.AbstractEventLoop | None
|
|
||||||
try:
|
|
||||||
_prev_loop = asyncio.get_event_loop_policy().get_event_loop()
|
|
||||||
except RuntimeError:
|
|
||||||
_prev_loop = None
|
|
||||||
_slash_loop: asyncio.AbstractEventLoop | None = None
|
|
||||||
_slash_handled = False
|
_slash_handled = False
|
||||||
_slash_error: Exception | None = None
|
_slash_error: Exception | None = None
|
||||||
try:
|
try:
|
||||||
_slash_loop = asyncio.new_event_loop()
|
async_runtime = runtime_state.async_runtime
|
||||||
asyncio.set_event_loop(_slash_loop)
|
_slash_handled = async_runtime.run_sync(
|
||||||
_slash_handled = _slash_loop.run_until_complete(
|
lambda: dispatch_channel_slash_command(
|
||||||
dispatch_channel_slash_command(
|
|
||||||
msg,
|
msg,
|
||||||
agent=runtime_state.agent,
|
agent=runtime_state.agent,
|
||||||
thread_id=runtime_state.thread_id,
|
thread_id=runtime_state.thread_id,
|
||||||
@@ -1245,15 +1296,12 @@ def _serve_process_message(
|
|||||||
),
|
),
|
||||||
channel_runtime=channel_runtime,
|
channel_runtime=channel_runtime,
|
||||||
graph_gateway=runtime_gateways.graph_gateway,
|
graph_gateway=runtime_gateways.graph_gateway,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
_slash_error = exc
|
_slash_error = exc
|
||||||
_serve_logger.exception("Slash dispatch failed for %s", msg.channel_type)
|
_serve_logger.exception("Slash dispatch failed for %s", msg.channel_type)
|
||||||
finally:
|
|
||||||
if _slash_loop is not None:
|
|
||||||
_slash_loop.close()
|
|
||||||
asyncio.set_event_loop(_prev_loop)
|
|
||||||
|
|
||||||
if _slash_error is not None:
|
if _slash_error is not None:
|
||||||
_set_channel_response(msg.msg_id, f"Command error: {_slash_error}")
|
_set_channel_response(msg.msg_id, f"Command error: {_slash_error}")
|
||||||
@@ -1280,6 +1328,7 @@ def _serve_process_message(
|
|||||||
show_thinking=show_thinking,
|
show_thinking=show_thinking,
|
||||||
interactive=True,
|
interactive=True,
|
||||||
metadata=meta,
|
metadata=meta,
|
||||||
|
configurable_extra=active_teams_configurable_extra(channel_runtime),
|
||||||
on_thinking=_send_thinking,
|
on_thinking=_send_thinking,
|
||||||
on_todo=_send_todo,
|
on_todo=_send_todo,
|
||||||
on_file_write=_send_media,
|
on_file_write=_send_media,
|
||||||
@@ -1287,11 +1336,13 @@ def _serve_process_message(
|
|||||||
ask_user_prompt_fn=_ask_user_prompt,
|
ask_user_prompt_fn=_ask_user_prompt,
|
||||||
cancel_scope=_channel_message_cancel_scope(msg),
|
cancel_scope=_channel_message_cancel_scope(msg),
|
||||||
gateway=runtime_gateways.graph_gateway,
|
gateway=runtime_gateways.graph_gateway,
|
||||||
|
runtime=runtime_state.async_runtime,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
response = f"Error: {e}"
|
response = f"Error: {e}"
|
||||||
console.print(f"[red]Serve error: {e}[/red]")
|
console.print(f"[red]Serve error: {e}[/red]")
|
||||||
|
|
||||||
|
pending_channel_sends.settle()
|
||||||
_set_channel_response(msg.msg_id, response)
|
_set_channel_response(msg.msg_id, response)
|
||||||
console.print(f"[dim][{msg.channel_type}] Replied to {msg.sender}[/dim]")
|
console.print(f"[dim][{msg.channel_type}] Replied to {msg.sender}[/dim]")
|
||||||
finally:
|
finally:
|
||||||
@@ -1309,6 +1360,7 @@ def _serve_drain_notifications(
|
|||||||
model: str | None,
|
model: str | None,
|
||||||
workspace_dir: str,
|
workspace_dir: str,
|
||||||
show_thinking: bool,
|
show_thinking: bool,
|
||||||
|
channel_runtime: ChannelRuntime | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Drain the async-task notification queue in headless serve mode.
|
"""Drain the async-task notification queue in headless serve mode.
|
||||||
|
|
||||||
@@ -1340,7 +1392,9 @@ def _serve_drain_notifications(
|
|||||||
show_thinking=show_thinking,
|
show_thinking=show_thinking,
|
||||||
interactive=True,
|
interactive=True,
|
||||||
metadata=meta,
|
metadata=meta,
|
||||||
|
configurable_extra=active_teams_configurable_extra(channel_runtime),
|
||||||
gateway=runtime_state.runtime_gateways.graph_gateway,
|
gateway=runtime_state.runtime_gateways.graph_gateway,
|
||||||
|
runtime=runtime_state.async_runtime,
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
_serve_logger.warning("Notification agent turn failed: %s", exc)
|
_serve_logger.warning("Notification agent turn failed: %s", exc)
|
||||||
@@ -1378,25 +1432,28 @@ def _serve_drain_notifications(
|
|||||||
current_thread_id=runtime_state.thread_id,
|
current_thread_id=runtime_state.thread_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
_notif_loop: _aio.AbstractEventLoop | None = None
|
|
||||||
try:
|
try:
|
||||||
_notif_loop = _aio.new_event_loop()
|
runtime_state.async_runtime.run_sync(_consume)
|
||||||
_notif_loop.run_until_complete(_consume())
|
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
_serve_logger.warning("Notification drain failed: %s", exc)
|
_serve_logger.warning("Notification drain failed: %s", exc)
|
||||||
finally:
|
|
||||||
if _notif_loop is not None:
|
|
||||||
_notif_loop.close()
|
|
||||||
|
|
||||||
|
|
||||||
@app.command()
|
@app.command()
|
||||||
def serve(
|
def serve(
|
||||||
|
ctx: typer.Context,
|
||||||
no_thinking: bool = typer.Option(
|
no_thinking: bool = typer.Option(
|
||||||
False, "--no-thinking", help="Disable thinking relay to channels"
|
False, "--no-thinking", help="Disable thinking relay to channels"
|
||||||
),
|
),
|
||||||
workdir: str | None = typer.Option(
|
workdir: str | None = typer.Option(
|
||||||
None, "--workdir", help="Override workspace directory"
|
None, "--workdir", help="Override workspace directory"
|
||||||
),
|
),
|
||||||
|
host: str | None = typer.Option(
|
||||||
|
None,
|
||||||
|
"--host",
|
||||||
|
help="Interface to bind the langgraph dev backend to (default: "
|
||||||
|
"langgraph_dev_host = 127.0.0.1). Pass 0.0.0.0 to reach it from "
|
||||||
|
"another machine — the backend has no auth.",
|
||||||
|
),
|
||||||
auto_approve: bool = typer.Option(
|
auto_approve: bool = typer.Option(
|
||||||
False,
|
False,
|
||||||
"--auto-approve",
|
"--auto-approve",
|
||||||
@@ -1431,6 +1488,9 @@ def serve(
|
|||||||
from ..config import apply_config_to_env, get_effective_config
|
from ..config import apply_config_to_env, get_effective_config
|
||||||
|
|
||||||
cli_overrides = {}
|
cli_overrides = {}
|
||||||
|
# serve starts no front-end, so only the backend bind applies here.
|
||||||
|
if host is not None and host.strip():
|
||||||
|
cli_overrides["langgraph_dev_host"] = host.strip()
|
||||||
if auto_approve:
|
if auto_approve:
|
||||||
cli_overrides["auto_approve"] = True
|
cli_overrides["auto_approve"] = True
|
||||||
if auto_mode:
|
if auto_mode:
|
||||||
@@ -1445,6 +1505,7 @@ def serve(
|
|||||||
cli_overrides["log_level"] = "DEBUG"
|
cli_overrides["log_level"] = "DEBUG"
|
||||||
cli_overrides["channel_debug_tracing"] = True
|
cli_overrides["channel_debug_tracing"] = True
|
||||||
config = get_effective_config(cli_overrides)
|
config = get_effective_config(cli_overrides)
|
||||||
|
async_runtime = _get_cli_async_runtime(ctx)
|
||||||
if debug:
|
if debug:
|
||||||
os.environ["EVOSCIENTIST_LOG_LEVEL"] = "DEBUG"
|
os.environ["EVOSCIENTIST_LOG_LEVEL"] = "DEBUG"
|
||||||
os.environ["EVOSCIENTIST_CHANNEL_DEBUG_TRACING"] = "true"
|
os.environ["EVOSCIENTIST_CHANNEL_DEBUG_TRACING"] = "true"
|
||||||
@@ -1495,11 +1556,15 @@ def serve(
|
|||||||
f"[bold red]{DANGEROUS_BANNER_MESSAGE}[/bold red]"
|
f"[bold red]{DANGEROUS_BANNER_MESSAGE}[/bold red]"
|
||||||
)
|
)
|
||||||
console.print("[dim]Loading agent...[/dim]")
|
console.print("[dim]Loading agent...[/dim]")
|
||||||
agent = _load_agent(workspace_dir=ws, config=config)
|
agent = _load_agent(workspace_dir=ws, config=config, runtime=async_runtime)
|
||||||
|
|
||||||
|
from ..gateway import create_runtime_gateways
|
||||||
|
|
||||||
runtime_gateways = create_runtime_gateways()
|
runtime_gateways = create_runtime_gateways()
|
||||||
tid = asyncio.run(
|
tid = async_runtime.run_sync(
|
||||||
runtime_gateways.graph_gateway.create_thread(GraphTarget(workspace_dir=ws))
|
lambda: runtime_gateways.graph_gateway.create_thread(
|
||||||
|
GraphTarget(workspace_dir=ws)
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
# Mutable runtime shared with _serve_process_message so channel slash
|
# Mutable runtime shared with _serve_process_message so channel slash
|
||||||
@@ -1511,6 +1576,7 @@ def serve(
|
|||||||
workspace_dir=ws,
|
workspace_dir=ws,
|
||||||
config=config,
|
config=config,
|
||||||
runtime_gateways=runtime_gateways,
|
runtime_gateways=runtime_gateways,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
|
|
||||||
channel_runtime = ChannelRuntime(agent=agent, thread_id=tid)
|
channel_runtime = ChannelRuntime(agent=agent, thread_id=tid)
|
||||||
@@ -1551,9 +1617,22 @@ def serve(
|
|||||||
import threading
|
import threading
|
||||||
|
|
||||||
shutdown_event = threading.Event()
|
shutdown_event = threading.Event()
|
||||||
|
no_active_cancel_scope = object()
|
||||||
|
active_cancel_scope: str | object | None = no_active_cancel_scope
|
||||||
|
|
||||||
def _handle_shutdown(signum: int, _frame: Any) -> None:
|
def _handle_shutdown(signum: int, _frame: Any) -> None:
|
||||||
shutdown_event.set()
|
shutdown_event.set()
|
||||||
|
# Cancelling the owned asyncio task is not enough when it is awaiting a
|
||||||
|
# blocking execute call: the executor thread and its isolated process
|
||||||
|
# group keep running until the matching stream event is set. Request
|
||||||
|
# scope cancellation before KeyboardInterrupt unwinds message cleanup
|
||||||
|
# (which discards that scope). SIGTERM also needs this to unblock the
|
||||||
|
# synchronous serve call so the poll loop can observe shutdown_event.
|
||||||
|
scope = active_cancel_scope
|
||||||
|
if scope is not no_active_cancel_scope:
|
||||||
|
from ..stream.display import request_stream_cancel
|
||||||
|
|
||||||
|
request_stream_cancel(cast(str | None, scope))
|
||||||
# Fall back to Python's default SIGINT behavior (raises
|
# Fall back to Python's default SIGINT behavior (raises
|
||||||
# KeyboardInterrupt) so blocking I/O inside ``run_streaming``
|
# KeyboardInterrupt) so blocking I/O inside ``run_streaming``
|
||||||
# is still interrupted. For SIGTERM there's no default that
|
# is still interrupted. For SIGTERM there's no default that
|
||||||
@@ -1573,6 +1652,7 @@ def serve(
|
|||||||
if shutdown_event.is_set():
|
if shutdown_event.is_set():
|
||||||
break
|
break
|
||||||
if msg is not None:
|
if msg is not None:
|
||||||
|
active_cancel_scope = _channel_message_cancel_scope(msg)
|
||||||
try:
|
try:
|
||||||
_serve_process_message(
|
_serve_process_message(
|
||||||
msg,
|
msg,
|
||||||
@@ -1588,15 +1668,23 @@ def serve(
|
|||||||
except KeyboardInterrupt:
|
except KeyboardInterrupt:
|
||||||
shutdown_event.set()
|
shutdown_event.set()
|
||||||
break
|
break
|
||||||
|
finally:
|
||||||
|
active_cancel_scope = no_active_cancel_scope
|
||||||
|
|
||||||
# Poll notification queue when idle (no channel message was pending).
|
# Poll notification queue when idle (no channel message was pending).
|
||||||
if async_notifier.has_pending_notifications(runtime_state.thread_id):
|
if async_notifier.has_pending_notifications(runtime_state.thread_id):
|
||||||
_serve_drain_notifications(
|
# Notification turns use the default stream cancellation scope.
|
||||||
runtime_state=runtime_state,
|
active_cancel_scope = None
|
||||||
model=config.model,
|
try:
|
||||||
workspace_dir=ws,
|
_serve_drain_notifications(
|
||||||
show_thinking=effective_channel_thinking,
|
runtime_state=runtime_state,
|
||||||
)
|
model=config.model,
|
||||||
|
workspace_dir=ws,
|
||||||
|
show_thinking=effective_channel_thinking,
|
||||||
|
channel_runtime=channel_runtime,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
active_cancel_scope = no_active_cancel_scope
|
||||||
except KeyboardInterrupt:
|
except KeyboardInterrupt:
|
||||||
shutdown_event.set()
|
shutdown_event.set()
|
||||||
finally:
|
finally:
|
||||||
@@ -1954,20 +2042,16 @@ def sessions_callback(ctx: typer.Context):
|
|||||||
so the bare command is informative rather than silent.
|
so the bare command is informative rather than silent.
|
||||||
"""
|
"""
|
||||||
if ctx.invoked_subcommand is None:
|
if ctx.invoked_subcommand is None:
|
||||||
sessions_stats()
|
sessions_stats(ctx)
|
||||||
|
|
||||||
|
|
||||||
@sessions_app.command("stats")
|
@sessions_app.command("stats")
|
||||||
def sessions_stats():
|
def sessions_stats(ctx: typer.Context):
|
||||||
"""Show DB size, thread count, total checkpoints, top heaviest threads."""
|
"""Show DB size, thread count, total checkpoints, top heaviest threads."""
|
||||||
import asyncio
|
|
||||||
|
|
||||||
from ..sessions import db_stats
|
from ..sessions import db_stats
|
||||||
|
|
||||||
try:
|
runtime = _get_cli_async_runtime(ctx)
|
||||||
stats = asyncio.get_event_loop().run_until_complete(db_stats())
|
stats = runtime.run_sync(db_stats)
|
||||||
except RuntimeError:
|
|
||||||
stats = asyncio.new_event_loop().run_until_complete(db_stats())
|
|
||||||
|
|
||||||
table = Table(title="EvoScientist sessions DB", show_header=True)
|
table = Table(title="EvoScientist sessions DB", show_header=True)
|
||||||
table.add_column("Metric", style="cyan")
|
table.add_column("Metric", style="cyan")
|
||||||
@@ -2094,6 +2178,15 @@ def _main_callback(
|
|||||||
"--ui",
|
"--ui",
|
||||||
help="UI backend: tui (default), cli, or webui.",
|
help="UI backend: tui (default), cli, or webui.",
|
||||||
),
|
),
|
||||||
|
host: str | None = typer.Option(
|
||||||
|
None,
|
||||||
|
"--host",
|
||||||
|
help="Interface to bind servers to (default: 127.0.0.1 for both). "
|
||||||
|
"Sets langgraph_dev_host — the backend shared by every UI mode — and "
|
||||||
|
"webui_host (WebUI mode only). Applies to the default entry; the "
|
||||||
|
"serve and deploy subcommands take their own --host. Pass 0.0.0.0 to "
|
||||||
|
"reach both from another machine (the backend has no auth).",
|
||||||
|
),
|
||||||
output_format: str | None = typer.Option(
|
output_format: str | None = typer.Option(
|
||||||
None,
|
None,
|
||||||
"--output-format",
|
"--output-format",
|
||||||
@@ -2108,6 +2201,8 @@ def _main_callback(
|
|||||||
if ctx.invoked_subcommand is not None:
|
if ctx.invoked_subcommand is not None:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
async_runtime = _get_cli_async_runtime(ctx)
|
||||||
|
|
||||||
# Load and apply configuration
|
# Load and apply configuration
|
||||||
from ..config import apply_config_to_env, get_effective_config
|
from ..config import apply_config_to_env, get_effective_config
|
||||||
|
|
||||||
@@ -2152,6 +2247,11 @@ def _main_callback(
|
|||||||
cli_overrides["show_thinking"] = False
|
cli_overrides["show_thinking"] = False
|
||||||
if ui:
|
if ui:
|
||||||
cli_overrides["ui_backend"] = ui
|
cli_overrides["ui_backend"] = ui
|
||||||
|
if host is not None and host.strip():
|
||||||
|
# One flag drives both servers; the backend applies in EVERY UI mode
|
||||||
|
# (auto-started for tui/cli/serve too), webui_host only in WebUI mode.
|
||||||
|
cli_overrides["webui_host"] = host.strip()
|
||||||
|
cli_overrides["langgraph_dev_host"] = host.strip()
|
||||||
if auto_approve:
|
if auto_approve:
|
||||||
cli_overrides["auto_approve"] = True
|
cli_overrides["auto_approve"] = True
|
||||||
if effective_auto_mode:
|
if effective_auto_mode:
|
||||||
@@ -2319,6 +2419,7 @@ def _main_callback(
|
|||||||
# Single-shot mode: wrap in persistent checkpointer
|
# Single-shot mode: wrap in persistent checkpointer
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
|
from ..gateway import create_runtime_gateways
|
||||||
from ..sessions import get_checkpointer
|
from ..sessions import get_checkpointer
|
||||||
from ..stream.json_sink import stream_json
|
from ..stream.json_sink import stream_json
|
||||||
from .interactive import _wait_for_memory_workers_before_exit, cmd_run
|
from .interactive import _wait_for_memory_workers_before_exit, cmd_run
|
||||||
@@ -2350,10 +2451,12 @@ def _main_callback(
|
|||||||
else:
|
else:
|
||||||
tid = await graph_gateway.create_thread()
|
tid = await graph_gateway.create_thread()
|
||||||
console.print("[dim]Loading agent...[/dim]")
|
console.print("[dim]Loading agent...[/dim]")
|
||||||
agent = _load_agent(
|
agent = await asyncio.to_thread(
|
||||||
|
_load_agent,
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
checkpointer=checkpointer,
|
checkpointer=checkpointer,
|
||||||
config=config,
|
config=config,
|
||||||
|
runtime=async_runtime,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
if effective_output_format == "stream-json":
|
if effective_output_format == "stream-json":
|
||||||
@@ -2382,26 +2485,47 @@ def _main_callback(
|
|||||||
# matching the text path (cmd_run does this itself).
|
# matching the text path (cmd_run does this itself).
|
||||||
_wait_for_memory_workers_before_exit()
|
_wait_for_memory_workers_before_exit()
|
||||||
else:
|
else:
|
||||||
cmd_run(
|
stream_worker = asyncio.create_task(
|
||||||
agent,
|
asyncio.to_thread(
|
||||||
prompt,
|
cmd_run,
|
||||||
thread_id=tid,
|
agent,
|
||||||
show_thinking=show_thinking,
|
prompt,
|
||||||
workspace_dir=workspace_dir,
|
thread_id=tid,
|
||||||
model=config.model,
|
show_thinking=show_thinking,
|
||||||
ui_backend=config.ui_backend,
|
workspace_dir=workspace_dir,
|
||||||
runtime_gateways=runtime_gateways,
|
model=config.model,
|
||||||
|
ui_backend=config.ui_backend,
|
||||||
|
runtime_gateways=runtime_gateways,
|
||||||
|
async_runtime=async_runtime,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
try:
|
||||||
|
await asyncio.shield(stream_worker)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
from ..stream.display import request_stream_cancel
|
||||||
|
from .tui_runtime import settle_cancelled_worker
|
||||||
|
|
||||||
|
await settle_cancelled_worker(
|
||||||
|
stream_worker,
|
||||||
|
on_cancel=request_stream_cancel,
|
||||||
|
)
|
||||||
|
raise
|
||||||
finally:
|
finally:
|
||||||
|
# Model failures can bypass middleware ``after_agent``
|
||||||
|
# hooks. Close any remaining QuickJS workers while this
|
||||||
|
# event loop is still available; their synchronous GC
|
||||||
|
# fallback can deadlock during interpreter shutdown.
|
||||||
|
from ..middleware.code_interpreter import (
|
||||||
|
aclose_code_interpreters,
|
||||||
|
)
|
||||||
|
|
||||||
|
await aclose_code_interpreters()
|
||||||
try:
|
try:
|
||||||
print_resume_hint(tid, console=console)
|
print_resume_hint(tid, console=console)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
import nest_asyncio
|
async_runtime.run_sync(_single_shot)
|
||||||
|
|
||||||
nest_asyncio.apply()
|
|
||||||
asyncio.get_event_loop().run_until_complete(_single_shot())
|
|
||||||
else:
|
else:
|
||||||
from .interactive import cmd_interactive
|
from .interactive import cmd_interactive
|
||||||
|
|
||||||
@@ -2418,6 +2542,7 @@ def _main_callback(
|
|||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
ui_backend=config.ui_backend,
|
ui_backend=config.ui_backend,
|
||||||
config=config,
|
config=config,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+160
-36
@@ -4,8 +4,10 @@ import asyncio
|
|||||||
import logging
|
import logging
|
||||||
import queue
|
import queue
|
||||||
import random
|
import random
|
||||||
|
import signal
|
||||||
import sys
|
import sys
|
||||||
from collections.abc import Callable
|
import threading
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
@@ -62,6 +64,7 @@ from .channel import (
|
|||||||
_set_channel_response,
|
_set_channel_response,
|
||||||
dispatch_channel_slash_command,
|
dispatch_channel_slash_command,
|
||||||
)
|
)
|
||||||
|
from .channel_sends import PendingChannelSends
|
||||||
from .file_mentions import complete_file_mention, resolve_file_mentions
|
from .file_mentions import complete_file_mention, resolve_file_mentions
|
||||||
from .rich_command_ui import RichCLICommandUI
|
from .rich_command_ui import RichCLICommandUI
|
||||||
from .status_bar import (
|
from .status_bar import (
|
||||||
@@ -83,7 +86,12 @@ from .status_bar import (
|
|||||||
make_usage_status_snapshot,
|
make_usage_status_snapshot,
|
||||||
)
|
)
|
||||||
from .tui_interactive import run_textual_interactive
|
from .tui_interactive import run_textual_interactive
|
||||||
from .tui_runtime import resolve_ui_backend, run_streaming
|
from .tui_runtime import (
|
||||||
|
StreamCancellationTimeout,
|
||||||
|
resolve_ui_backend,
|
||||||
|
run_streaming,
|
||||||
|
run_streaming_async,
|
||||||
|
)
|
||||||
|
|
||||||
_MEMORY_WORKER_SHUTDOWN_WAIT_SECONDS = 120.0
|
_MEMORY_WORKER_SHUTDOWN_WAIT_SECONDS = 120.0
|
||||||
_MEMORY_WORKER_SHUTDOWN_POLL_SECONDS = 0.5
|
_MEMORY_WORKER_SHUTDOWN_POLL_SECONDS = 0.5
|
||||||
@@ -97,6 +105,8 @@ _background_tasks: set[asyncio.Task] = set()
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from langgraph.graph.state import CompiledStateGraph
|
from langgraph.graph.state import CompiledStateGraph
|
||||||
|
|
||||||
|
from ..runtime import AsyncRuntime
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
class _StartupSession:
|
class _StartupSession:
|
||||||
@@ -107,6 +117,15 @@ class _StartupSession:
|
|||||||
resumed: bool
|
resumed: bool
|
||||||
|
|
||||||
|
|
||||||
|
async def _run_serialized_turn(
|
||||||
|
turn_lock: asyncio.Lock,
|
||||||
|
operation: Callable[[], Awaitable[Any]],
|
||||||
|
) -> Any:
|
||||||
|
"""Run one session turn without overlapping another frontend source."""
|
||||||
|
async with turn_lock:
|
||||||
|
return await operation()
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
# Banner
|
# Banner
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -328,6 +347,47 @@ async def _resolve_startup_session(
|
|||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
async def _run_rich_cli_streaming_turn(**kwargs: Any) -> str:
|
||||||
|
"""Run one Rich CLI turn with a fresh, turn-local SIGINT policy.
|
||||||
|
|
||||||
|
``asyncio.run`` installs a SIGINT handler whose interrupt count lasts for
|
||||||
|
the lifetime of the runner. The Rich CLI intentionally recovers after a
|
||||||
|
cancelled turn, so relying on that handler makes Ctrl+C on a later turn
|
||||||
|
look like the runner's second interrupt and raises ``KeyboardInterrupt``.
|
||||||
|
|
||||||
|
While a model turn is active, route the first Ctrl+C to a child task
|
||||||
|
instead. Restoring the runner's handler after every turn keeps Ctrl+C at
|
||||||
|
the prompt unchanged and resets the force-quit boundary for the next turn.
|
||||||
|
A second Ctrl+C before the current turn settles remains a force quit.
|
||||||
|
"""
|
||||||
|
stream_task = asyncio.create_task(
|
||||||
|
run_streaming_async(**kwargs, recover_on_cancel=True)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Interactive CLI execution belongs on the main thread, but retaining the
|
||||||
|
# ordinary await makes this helper safe in embedded/test environments where
|
||||||
|
# Python does not permit installing process signal handlers.
|
||||||
|
if threading.current_thread() is not threading.main_thread():
|
||||||
|
return await stream_task
|
||||||
|
|
||||||
|
previous_sigint = signal.getsignal(signal.SIGINT)
|
||||||
|
interrupted = False
|
||||||
|
|
||||||
|
def _cancel_turn(signum: int, frame: Any) -> None:
|
||||||
|
nonlocal interrupted
|
||||||
|
if interrupted or stream_task.done():
|
||||||
|
signal.default_int_handler(signum, frame)
|
||||||
|
return
|
||||||
|
interrupted = True
|
||||||
|
stream_task.cancel()
|
||||||
|
|
||||||
|
signal.signal(signal.SIGINT, _cancel_turn)
|
||||||
|
try:
|
||||||
|
return await stream_task
|
||||||
|
finally:
|
||||||
|
signal.signal(signal.SIGINT, previous_sigint)
|
||||||
|
|
||||||
|
|
||||||
def cmd_interactive(
|
def cmd_interactive(
|
||||||
show_thinking: bool = True,
|
show_thinking: bool = True,
|
||||||
channel_send_thinking: bool = True,
|
channel_send_thinking: bool = True,
|
||||||
@@ -340,6 +400,7 @@ def cmd_interactive(
|
|||||||
thread_id: str | None = None,
|
thread_id: str | None = None,
|
||||||
ui_backend: str = "cli",
|
ui_backend: str = "cli",
|
||||||
config=None,
|
config=None,
|
||||||
|
async_runtime: "AsyncRuntime | None" = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Interactive conversation mode with streaming output.
|
"""Interactive conversation mode with streaming output.
|
||||||
|
|
||||||
@@ -358,15 +419,15 @@ def cmd_interactive(
|
|||||||
thread_id: Optional thread ID to resume a previous session
|
thread_id: Optional thread ID to resume a previous session
|
||||||
ui_backend: UI backend ('cli' or 'tui')
|
ui_backend: UI backend ('cli' or 'tui')
|
||||||
"""
|
"""
|
||||||
import nest_asyncio
|
|
||||||
|
|
||||||
nest_asyncio.apply()
|
|
||||||
|
|
||||||
resolved_ui_backend = resolve_ui_backend(ui_backend, warn_fallback=True)
|
resolved_ui_backend = resolve_ui_backend(ui_backend, warn_fallback=True)
|
||||||
if resolved_ui_backend == "tui":
|
if resolved_ui_backend == "tui":
|
||||||
from functools import partial
|
from functools import partial
|
||||||
|
|
||||||
load_agent = partial(_load_agent, config=config)
|
load_agent = partial(
|
||||||
|
_load_agent,
|
||||||
|
config=config,
|
||||||
|
runtime=async_runtime,
|
||||||
|
)
|
||||||
run_textual_interactive(
|
run_textual_interactive(
|
||||||
show_thinking=show_thinking,
|
show_thinking=show_thinking,
|
||||||
channel_send_thinking=channel_send_thinking,
|
channel_send_thinking=channel_send_thinking,
|
||||||
@@ -380,6 +441,7 @@ def cmd_interactive(
|
|||||||
load_agent=load_agent,
|
load_agent=load_agent,
|
||||||
create_session_workspace=_create_session_workspace,
|
create_session_workspace=_create_session_workspace,
|
||||||
config=config,
|
config=config,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -419,7 +481,7 @@ def cmd_interactive(
|
|||||||
width = console.size.width
|
width = console.size.width
|
||||||
console.print(Text("\u2500" * width, style="dim"))
|
console.print(Text("\u2500" * width, style="dim"))
|
||||||
|
|
||||||
from ..commands.base import ChannelRuntime
|
from ..commands.base import ChannelRuntime, active_teams_configurable_extra
|
||||||
|
|
||||||
channel_runtime = ChannelRuntime()
|
channel_runtime = ChannelRuntime()
|
||||||
|
|
||||||
@@ -448,7 +510,17 @@ def cmd_interactive(
|
|||||||
on_progress=_on_mcp_progress,
|
on_progress=_on_mcp_progress,
|
||||||
)
|
)
|
||||||
|
|
||||||
runtime_gateways = create_runtime_gateways()
|
# One frontend event sink for the whole session — injected into the agent's
|
||||||
|
# middleware (write side) and the local gateway's streaming path (read side)
|
||||||
|
# so both share one owner. It survives agent rebuilds (/model, /new, MCP
|
||||||
|
# reload) because the session, not the agent, holds it.
|
||||||
|
from ..stream.sink import SessionEventSink
|
||||||
|
|
||||||
|
event_sink = SessionEventSink(
|
||||||
|
fallback_display=lambda text, style: console.print(text, style=style)
|
||||||
|
)
|
||||||
|
|
||||||
|
runtime_gateways = create_runtime_gateways(events=event_sink)
|
||||||
graph_gateway = runtime_gateways.graph_gateway
|
graph_gateway = runtime_gateways.graph_gateway
|
||||||
requested_thread_id = thread_id
|
requested_thread_id = thread_id
|
||||||
|
|
||||||
@@ -486,6 +558,8 @@ def cmd_interactive(
|
|||||||
workspace_dir=state["workspace_dir"],
|
workspace_dir=state["workspace_dir"],
|
||||||
checkpointer=checkpointer,
|
checkpointer=checkpointer,
|
||||||
config=config,
|
config=config,
|
||||||
|
events=event_sink,
|
||||||
|
runtime=async_runtime,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _await_agent_ready() -> "CompiledStateGraph":
|
async def _await_agent_ready() -> "CompiledStateGraph":
|
||||||
@@ -857,6 +931,8 @@ def cmd_interactive(
|
|||||||
|
|
||||||
# ---- Channel queue processing (bus → main thread) ----
|
# ---- Channel queue processing (bus → main thread) ----
|
||||||
|
|
||||||
|
turn_lock = asyncio.Lock()
|
||||||
|
|
||||||
async def _process_channel_message(msg: ChannelMessage) -> None:
|
async def _process_channel_message(msg: ChannelMessage) -> None:
|
||||||
"""Process a single channel message with real-time streaming.
|
"""Process a single channel message with real-time streaming.
|
||||||
|
|
||||||
@@ -894,17 +970,12 @@ def cmd_interactive(
|
|||||||
console.print(rx)
|
console.print(rx)
|
||||||
_print_separator()
|
_print_separator()
|
||||||
|
|
||||||
|
pending_channel_sends = PendingChannelSends(
|
||||||
|
_ch_mod._bus_loop, _channel_logger
|
||||||
|
)
|
||||||
|
|
||||||
def _send_to_channel(coro, label: str, timeout: int = 15) -> None:
|
def _send_to_channel(coro, label: str, timeout: int = 15) -> None:
|
||||||
"""Schedule an async channel send on the bus loop."""
|
pending_channel_sends.submit(coro, label, timeout)
|
||||||
loop = _ch_mod._bus_loop
|
|
||||||
if not loop:
|
|
||||||
return
|
|
||||||
try:
|
|
||||||
asyncio.run_coroutine_threadsafe(coro, loop).result(
|
|
||||||
timeout=timeout
|
|
||||||
)
|
|
||||||
except Exception as e:
|
|
||||||
_channel_logger.debug(f"{label} send failed: {e}")
|
|
||||||
|
|
||||||
def _send_thinking_to_channel(thinking: str) -> None:
|
def _send_thinking_to_channel(thinking: str) -> None:
|
||||||
ch = msg.channel_ref
|
ch = msg.channel_ref
|
||||||
@@ -1018,6 +1089,7 @@ def cmd_interactive(
|
|||||||
on_cmd_completed=_on_channel_cmd_completed,
|
on_cmd_completed=_on_channel_cmd_completed,
|
||||||
channel_runtime=channel_runtime,
|
channel_runtime=channel_runtime,
|
||||||
graph_gateway=runtime_gateways.graph_gateway,
|
graph_gateway=runtime_gateways.graph_gateway,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
if _slash_handled:
|
if _slash_handled:
|
||||||
# A channel-issued /new or /resume rotates the thread
|
# A channel-issued /new or /resume rotates the thread
|
||||||
@@ -1036,7 +1108,7 @@ def cmd_interactive(
|
|||||||
await _refresh_status_snapshot(
|
await _refresh_status_snapshot(
|
||||||
msg.content, reset_streaming_text=True
|
msg.content, reset_streaming_text=True
|
||||||
)
|
)
|
||||||
response = run_streaming(
|
response = await run_streaming_async(
|
||||||
ui_backend=state["ui_backend"],
|
ui_backend=state["ui_backend"],
|
||||||
agent=ready_agent,
|
agent=ready_agent,
|
||||||
message=msg.content,
|
message=msg.content,
|
||||||
@@ -1044,6 +1116,9 @@ def cmd_interactive(
|
|||||||
show_thinking=show_thinking,
|
show_thinking=show_thinking,
|
||||||
interactive=True,
|
interactive=True,
|
||||||
metadata=meta,
|
metadata=meta,
|
||||||
|
configurable_extra=active_teams_configurable_extra(
|
||||||
|
channel_runtime
|
||||||
|
),
|
||||||
on_thinking=_send_thinking_to_channel,
|
on_thinking=_send_thinking_to_channel,
|
||||||
on_todo=_send_todo_to_channel,
|
on_todo=_send_todo_to_channel,
|
||||||
on_file_write=_send_media_to_channel,
|
on_file_write=_send_media_to_channel,
|
||||||
@@ -1053,11 +1128,13 @@ def cmd_interactive(
|
|||||||
status_footer_builder=_stream_status_footer,
|
status_footer_builder=_stream_status_footer,
|
||||||
cancel_scope=_ch_mod._channel_message_cancel_scope(msg),
|
cancel_scope=_ch_mod._channel_message_cancel_scope(msg),
|
||||||
gateway=runtime_gateways.graph_gateway,
|
gateway=runtime_gateways.graph_gateway,
|
||||||
|
runtime=async_runtime,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
response = f"Error: {e}"
|
response = f"Error: {e}"
|
||||||
console.print(f"[red]Channel error: {e}[/red]")
|
console.print(f"[red]Channel error: {e}[/red]")
|
||||||
|
|
||||||
|
await pending_channel_sends.settle_async()
|
||||||
_set_channel_response(msg.msg_id, response)
|
_set_channel_response(msg.msg_id, response)
|
||||||
await _refresh_status_snapshot(reset_streaming_text=True)
|
await _refresh_status_snapshot(reset_streaming_text=True)
|
||||||
|
|
||||||
@@ -1094,7 +1171,7 @@ def cmd_interactive(
|
|||||||
meta = build_metadata(state["workspace_dir"], model)
|
meta = build_metadata(state["workspace_dir"], model)
|
||||||
await _refresh_status_snapshot(text, reset_streaming_text=True)
|
await _refresh_status_snapshot(text, reset_streaming_text=True)
|
||||||
ready_agent = await _await_agent_ready()
|
ready_agent = await _await_agent_ready()
|
||||||
response = run_streaming(
|
response = await run_streaming_async(
|
||||||
ui_backend=state["ui_backend"],
|
ui_backend=state["ui_backend"],
|
||||||
agent=ready_agent,
|
agent=ready_agent,
|
||||||
message=text,
|
message=text,
|
||||||
@@ -1107,9 +1184,11 @@ def cmd_interactive(
|
|||||||
show_thinking=show_thinking,
|
show_thinking=show_thinking,
|
||||||
interactive=True,
|
interactive=True,
|
||||||
metadata=meta,
|
metadata=meta,
|
||||||
|
configurable_extra=active_teams_configurable_extra(channel_runtime),
|
||||||
on_stream_event=_handle_stream_status_event,
|
on_stream_event=_handle_stream_status_event,
|
||||||
status_footer_builder=_stream_status_footer,
|
status_footer_builder=_stream_status_footer,
|
||||||
gateway=runtime_gateways.graph_gateway,
|
gateway=runtime_gateways.graph_gateway,
|
||||||
|
runtime=async_runtime,
|
||||||
)
|
)
|
||||||
_notif_tid = target_thread_id or state["thread_id"]
|
_notif_tid = target_thread_id or state["thread_id"]
|
||||||
if _ch_mod.publish_to_channel_origin(_notif_tid, response):
|
if _ch_mod.publish_to_channel_origin(_notif_tid, response):
|
||||||
@@ -1165,7 +1244,10 @@ def cmd_interactive(
|
|||||||
except queue.Empty:
|
except queue.Empty:
|
||||||
msg = None
|
msg = None
|
||||||
if msg is not None:
|
if msg is not None:
|
||||||
await _process_channel_message(msg)
|
await _run_serialized_turn(
|
||||||
|
turn_lock,
|
||||||
|
lambda _msg=msg: _process_channel_message(_msg),
|
||||||
|
)
|
||||||
continue # check queues again immediately
|
continue # check queues again immediately
|
||||||
|
|
||||||
# Notification path (only when no channel message was pending).
|
# Notification path (only when no channel message was pending).
|
||||||
@@ -1182,8 +1264,13 @@ def cmd_interactive(
|
|||||||
try:
|
try:
|
||||||
await async_notifier.consume_notifications(
|
await async_notifier.consume_notifications(
|
||||||
run_message=lambda text, notifs, _tid=current_tid: (
|
run_message=lambda text, notifs, _tid=current_tid: (
|
||||||
_inject_notification_message(
|
_run_serialized_turn(
|
||||||
text, notifs, target_thread_id=_tid
|
turn_lock,
|
||||||
|
lambda: _inject_notification_message(
|
||||||
|
text,
|
||||||
|
notifs,
|
||||||
|
target_thread_id=_tid,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
read_async_tasks_state=read_async_tasks_state,
|
read_async_tasks_state=read_async_tasks_state,
|
||||||
@@ -1318,6 +1405,7 @@ def cmd_interactive(
|
|||||||
input_tokens_hint=state.get("status_last_input_tokens"),
|
input_tokens_hint=state.get("status_last_input_tokens"),
|
||||||
channel_runtime=channel_runtime,
|
channel_runtime=channel_runtime,
|
||||||
graph_gateway=runtime_gateways.graph_gateway,
|
graph_gateway=runtime_gateways.graph_gateway,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
await cmd_manager.execute(user_input, ctx)
|
await cmd_manager.execute(user_input, ctx)
|
||||||
|
|
||||||
@@ -1400,17 +1488,26 @@ def cmd_interactive(
|
|||||||
await _refresh_status_snapshot(
|
await _refresh_status_snapshot(
|
||||||
message_to_send, reset_streaming_text=True
|
message_to_send, reset_streaming_text=True
|
||||||
)
|
)
|
||||||
run_streaming(
|
await _run_serialized_turn(
|
||||||
ui_backend=state["ui_backend"],
|
turn_lock,
|
||||||
agent=ready_agent,
|
lambda _agent=ready_agent, _message=message_to_send, _thread_id=state["thread_id"], _meta=meta: (
|
||||||
message=message_to_send,
|
_run_rich_cli_streaming_turn(
|
||||||
thread_id=state["thread_id"],
|
ui_backend=state["ui_backend"],
|
||||||
show_thinking=show_thinking,
|
agent=_agent,
|
||||||
interactive=True,
|
message=_message,
|
||||||
metadata=meta,
|
thread_id=_thread_id,
|
||||||
on_stream_event=_handle_stream_status_event,
|
show_thinking=show_thinking,
|
||||||
status_footer_builder=_stream_status_footer,
|
interactive=True,
|
||||||
gateway=runtime_gateways.graph_gateway,
|
metadata=_meta,
|
||||||
|
configurable_extra=active_teams_configurable_extra(
|
||||||
|
channel_runtime
|
||||||
|
),
|
||||||
|
on_stream_event=_handle_stream_status_event,
|
||||||
|
status_footer_builder=_stream_status_footer,
|
||||||
|
gateway=runtime_gateways.graph_gateway,
|
||||||
|
runtime=async_runtime,
|
||||||
|
)
|
||||||
|
),
|
||||||
)
|
)
|
||||||
await _refresh_status_snapshot(reset_streaming_text=True)
|
await _refresh_status_snapshot(reset_streaming_text=True)
|
||||||
console.print()
|
console.print()
|
||||||
@@ -1425,6 +1522,14 @@ def cmd_interactive(
|
|||||||
console.print()
|
console.print()
|
||||||
state["running"] = False
|
state["running"] = False
|
||||||
break
|
break
|
||||||
|
except StreamCancellationTimeout as e:
|
||||||
|
console.print(f"[red]{escape(str(e))}[/red]")
|
||||||
|
console.print(
|
||||||
|
"[dim]Exiting because the active turn could not be "
|
||||||
|
"stopped safely.[/dim]"
|
||||||
|
)
|
||||||
|
state["running"] = False
|
||||||
|
break
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
error_msg = str(e)
|
error_msg = str(e)
|
||||||
if (
|
if (
|
||||||
@@ -1445,6 +1550,17 @@ def cmd_interactive(
|
|||||||
await queue_task
|
await queue_task
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
pass
|
pass
|
||||||
|
try:
|
||||||
|
from ..middleware.code_interpreter import (
|
||||||
|
aclose_code_interpreters,
|
||||||
|
)
|
||||||
|
|
||||||
|
await aclose_code_interpreters()
|
||||||
|
except Exception:
|
||||||
|
_channel_logger.debug(
|
||||||
|
"code interpreter cleanup failed",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
# Best-effort: guard so a DB lookup failure here can't
|
# Best-effort: guard so a DB lookup failure here can't
|
||||||
# shadow the original exception exiting _async_main_loop.
|
# shadow the original exception exiting _async_main_loop.
|
||||||
current_tid = state.get("thread_id")
|
current_tid = state.get("thread_id")
|
||||||
@@ -1482,6 +1598,7 @@ def cmd_run(
|
|||||||
ui_backend: str = "cli",
|
ui_backend: str = "cli",
|
||||||
*,
|
*,
|
||||||
runtime_gateways: RuntimeGateways,
|
runtime_gateways: RuntimeGateways,
|
||||||
|
async_runtime: "AsyncRuntime | None" = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Single-shot execution with streaming display.
|
"""Single-shot execution with streaming display.
|
||||||
|
|
||||||
@@ -1515,6 +1632,7 @@ def cmd_run(
|
|||||||
interactive=False,
|
interactive=False,
|
||||||
metadata=meta,
|
metadata=meta,
|
||||||
gateway=runtime_gateways.graph_gateway,
|
gateway=runtime_gateways.graph_gateway,
|
||||||
|
runtime=async_runtime,
|
||||||
)
|
)
|
||||||
_wait_for_memory_workers_before_exit()
|
_wait_for_memory_workers_before_exit()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -1527,7 +1645,13 @@ def cmd_run(
|
|||||||
raise typer.Exit(1) from e
|
raise typer.Exit(1) from e
|
||||||
else:
|
else:
|
||||||
console.print(f"[red]Error: {e}[/red]")
|
console.print(f"[red]Error: {e}[/red]")
|
||||||
raise
|
# This is the process boundary for single-shot text mode. Letting
|
||||||
|
# provider exceptions escape makes Typer/Rich render the complete
|
||||||
|
# async exception chain after we already printed a concise error;
|
||||||
|
# large OpenAI/httpx chains can keep the CLI busy well after the
|
||||||
|
# resume hint is shown. Convert the failure to Click's controlled
|
||||||
|
# exit signal while preserving the cause for programmatic callers.
|
||||||
|
raise typer.Exit(1) from e
|
||||||
|
|
||||||
|
|
||||||
def _wait_for_memory_workers_before_exit(
|
def _wait_for_memory_workers_before_exit(
|
||||||
|
|||||||
@@ -0,0 +1,99 @@
|
|||||||
|
"""``EvoSci server`` — inspect and stop the background langgraph dev server.
|
||||||
|
|
||||||
|
The explicit counterpart to ``langgraph_dev_keepalive``: an opt-in server
|
||||||
|
that outlives its CLI needs an equally explicit way to see and stop it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import sys
|
||||||
|
|
||||||
|
from ..stream.console import console
|
||||||
|
from ._app import server_app
|
||||||
|
|
||||||
|
|
||||||
|
@server_app.command("status")
|
||||||
|
def server_status() -> None:
|
||||||
|
"""Show the background langgraph dev server's state."""
|
||||||
|
from ..config import get_effective_config
|
||||||
|
from ..langgraph_dev.manager import (
|
||||||
|
_DEFAULT_HOST,
|
||||||
|
_DEFAULT_PORT,
|
||||||
|
_pid_serves_port,
|
||||||
|
_read_workspace_sidecar,
|
||||||
|
is_langgraph_dev_running,
|
||||||
|
)
|
||||||
|
|
||||||
|
config = get_effective_config()
|
||||||
|
port = int(getattr(config, "langgraph_dev_port", _DEFAULT_PORT))
|
||||||
|
host = (
|
||||||
|
str(getattr(config, "langgraph_dev_host", _DEFAULT_HOST) or _DEFAULT_HOST)
|
||||||
|
).strip() or _DEFAULT_HOST
|
||||||
|
running = is_langgraph_dev_running(port=port, host=host)
|
||||||
|
sidecar = _read_workspace_sidecar()
|
||||||
|
|
||||||
|
if not running and sidecar is None:
|
||||||
|
console.print("[dim]No background langgraph dev server is running.[/dim]")
|
||||||
|
return
|
||||||
|
state = "[green]running[/green]" if running else "[red]not responding[/red]"
|
||||||
|
console.print(f"[bold]langgraph dev[/bold] on port {port}: {state}")
|
||||||
|
if sidecar is not None:
|
||||||
|
console.print(f" workspace: {sidecar.get('workspace')}")
|
||||||
|
pid = sidecar.get("pid")
|
||||||
|
if _pid_serves_port(pid, port):
|
||||||
|
console.print(f" pid: {pid}")
|
||||||
|
else:
|
||||||
|
console.print(
|
||||||
|
f" pid: {pid} [yellow](stale record — this pid does "
|
||||||
|
f"not serve port {port})[/yellow]"
|
||||||
|
)
|
||||||
|
elif running:
|
||||||
|
console.print(
|
||||||
|
" [yellow]no sidecar — externally managed or pre-keepalive server[/yellow]"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@server_app.command("stop")
|
||||||
|
def server_stop() -> None:
|
||||||
|
"""Stop the background langgraph dev server started by EvoSci."""
|
||||||
|
from ..config import get_effective_config
|
||||||
|
from ..langgraph_dev.manager import (
|
||||||
|
_DEFAULT_HOST,
|
||||||
|
_DEFAULT_PORT,
|
||||||
|
is_langgraph_dev_running,
|
||||||
|
stop_recorded_server,
|
||||||
|
)
|
||||||
|
|
||||||
|
pid = stop_recorded_server()
|
||||||
|
if pid is not None:
|
||||||
|
console.print(f"[green]✓[/green] Stopped langgraph dev (pid {pid}).")
|
||||||
|
return
|
||||||
|
config = get_effective_config()
|
||||||
|
port = int(getattr(config, "langgraph_dev_port", _DEFAULT_PORT))
|
||||||
|
host = (
|
||||||
|
str(getattr(config, "langgraph_dev_host", _DEFAULT_HOST) or _DEFAULT_HOST)
|
||||||
|
).strip() or _DEFAULT_HOST
|
||||||
|
if is_langgraph_dev_running(port=port, host=host):
|
||||||
|
# A server without ownership records (crashed session, deleted state
|
||||||
|
# files) can't be verified as ours — refuse to guess, hand the user
|
||||||
|
# the manual path instead of a silent no-op.
|
||||||
|
console.print(
|
||||||
|
f"[yellow]⚠ A langgraph dev is still serving port {port}, but "
|
||||||
|
f"EvoSci has no ownership record for it, so it was not "
|
||||||
|
f"touched.[/yellow]"
|
||||||
|
)
|
||||||
|
if sys.platform == "win32":
|
||||||
|
manual = (
|
||||||
|
f'powershell "Get-NetTCPConnection -LocalPort {port} | '
|
||||||
|
f'Select-Object -ExpandProperty OwningProcess | Stop-Process"'
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
manual = f"kill $(lsof -ti :{port})"
|
||||||
|
console.print(
|
||||||
|
f"[dim]If it is yours, stop it manually: [bold]{manual}[/bold][/dim]"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
console.print(
|
||||||
|
"[dim]No EvoSci-owned langgraph dev server to stop "
|
||||||
|
"(stale state, if any, was cleaned up).[/dim]"
|
||||||
|
)
|
||||||
@@ -4,11 +4,14 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Protocol
|
from typing import TYPE_CHECKING, Any, Protocol
|
||||||
|
|
||||||
from ..gateway import GraphGateway
|
from ..gateway import GraphGateway
|
||||||
from ..stream.display import _run_streaming
|
from ..stream.display import _run_streaming
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from ..runtime import AsyncRuntime
|
||||||
|
|
||||||
|
|
||||||
class StreamingTUIBackend(Protocol):
|
class StreamingTUIBackend(Protocol):
|
||||||
"""Protocol for TUI backends that can render agent streaming output."""
|
"""Protocol for TUI backends that can render agent streaming output."""
|
||||||
@@ -29,10 +32,12 @@ class StreamingTUIBackend(Protocol):
|
|||||||
on_stream_event: Callable[[str, Any], Any] | None = None,
|
on_stream_event: Callable[[str, Any], Any] | None = None,
|
||||||
status_footer_builder: Callable[[], Any] | None = None,
|
status_footer_builder: Callable[[], Any] | None = None,
|
||||||
metadata: dict | None = None,
|
metadata: dict | None = None,
|
||||||
|
configurable_extra: dict[str, Any] | None = None,
|
||||||
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
||||||
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
||||||
cancel_scope: str | None = None,
|
cancel_scope: str | None = None,
|
||||||
gateway: GraphGateway,
|
gateway: GraphGateway,
|
||||||
|
runtime: AsyncRuntime | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Run streaming and return final response text."""
|
"""Run streaming and return final response text."""
|
||||||
|
|
||||||
@@ -57,10 +62,12 @@ class RichStreamingBackend:
|
|||||||
on_stream_event: Callable[[str, Any], Any] | None = None,
|
on_stream_event: Callable[[str, Any], Any] | None = None,
|
||||||
status_footer_builder: Callable[[], Any] | None = None,
|
status_footer_builder: Callable[[], Any] | None = None,
|
||||||
metadata: dict | None = None,
|
metadata: dict | None = None,
|
||||||
|
configurable_extra: dict[str, Any] | None = None,
|
||||||
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
||||||
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
||||||
cancel_scope: str | None = None,
|
cancel_scope: str | None = None,
|
||||||
gateway: GraphGateway,
|
gateway: GraphGateway,
|
||||||
|
runtime: AsyncRuntime | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
return _run_streaming(
|
return _run_streaming(
|
||||||
agent=agent,
|
agent=agent,
|
||||||
@@ -74,8 +81,10 @@ class RichStreamingBackend:
|
|||||||
on_stream_event=on_stream_event,
|
on_stream_event=on_stream_event,
|
||||||
status_footer_builder=status_footer_builder,
|
status_footer_builder=status_footer_builder,
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
|
configurable_extra=configurable_extra,
|
||||||
hitl_prompt_fn=hitl_prompt_fn,
|
hitl_prompt_fn=hitl_prompt_fn,
|
||||||
ask_user_prompt_fn=ask_user_prompt_fn,
|
ask_user_prompt_fn=ask_user_prompt_fn,
|
||||||
cancel_scope=cancel_scope,
|
cancel_scope=cancel_scope,
|
||||||
gateway=gateway,
|
gateway=gateway,
|
||||||
|
runtime=runtime,
|
||||||
)
|
)
|
||||||
|
|||||||
+429
-104
@@ -11,6 +11,7 @@ import logging
|
|||||||
import queue
|
import queue
|
||||||
import random
|
import random
|
||||||
import sys
|
import sys
|
||||||
|
import threading
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
@@ -53,11 +54,11 @@ from .channel import (
|
|||||||
ChannelMessage,
|
ChannelMessage,
|
||||||
_auto_start_channel,
|
_auto_start_channel,
|
||||||
_channels_is_running,
|
_channels_is_running,
|
||||||
_channels_running_list,
|
|
||||||
_channels_stop,
|
_channels_stop,
|
||||||
_message_queue,
|
_message_queue,
|
||||||
_set_channel_response,
|
_set_channel_response,
|
||||||
dispatch_channel_slash_command,
|
dispatch_channel_slash_command,
|
||||||
|
get_channel_startup_results,
|
||||||
)
|
)
|
||||||
from .file_mentions import complete_file_mention, resolve_file_mentions
|
from .file_mentions import complete_file_mention, resolve_file_mentions
|
||||||
from .history_suggester import HistorySuggester
|
from .history_suggester import HistorySuggester
|
||||||
@@ -77,6 +78,9 @@ from .status_bar import (
|
|||||||
make_usage_status_snapshot,
|
make_usage_status_snapshot,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from ..runtime import AsyncRuntime
|
||||||
|
|
||||||
_channel_logger = logging.getLogger(__name__)
|
_channel_logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -92,6 +96,45 @@ def _shorten_path(path: str) -> str:
|
|||||||
return _sp(path)
|
return _sp(path)
|
||||||
|
|
||||||
|
|
||||||
|
async def _auto_start_channel_in_worker(
|
||||||
|
agent: Any,
|
||||||
|
thread_id: str,
|
||||||
|
config: Any,
|
||||||
|
*,
|
||||||
|
send_thinking: bool,
|
||||||
|
runtime: Any,
|
||||||
|
stop_requested: threading.Event,
|
||||||
|
) -> list[tuple[str, bool, str]]:
|
||||||
|
"""Run blocking channel startup without occupying the TUI event loop."""
|
||||||
|
|
||||||
|
def _start() -> list[tuple[str, bool, str]]:
|
||||||
|
try:
|
||||||
|
return _auto_start_channel(
|
||||||
|
agent,
|
||||||
|
thread_id,
|
||||||
|
config,
|
||||||
|
send_thinking=send_thinking,
|
||||||
|
runtime=runtime,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
if stop_requested.is_set():
|
||||||
|
_channels_stop(runtime=runtime)
|
||||||
|
|
||||||
|
worker = asyncio.create_task(asyncio.to_thread(_start))
|
||||||
|
try:
|
||||||
|
return await asyncio.shield(worker)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
stop_requested.set()
|
||||||
|
try:
|
||||||
|
await worker
|
||||||
|
except Exception:
|
||||||
|
_channel_logger.debug(
|
||||||
|
"Channel startup worker failed during cancellation",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
def _build_welcome_banner(
|
def _build_welcome_banner(
|
||||||
*,
|
*,
|
||||||
thread_id: str,
|
thread_id: str,
|
||||||
@@ -220,6 +263,9 @@ async def _sync_tui_command_completion(
|
|||||||
cmd: Command,
|
cmd: Command,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Adopt successful command-side state changes back into the TUI app."""
|
"""Adopt successful command-side state changes back into the TUI app."""
|
||||||
|
if app._exiting:
|
||||||
|
return
|
||||||
|
|
||||||
agent_swapped = ctx.agent is not None and ctx.agent is not original_agent
|
agent_swapped = ctx.agent is not None and ctx.agent is not original_agent
|
||||||
if agent_swapped:
|
if agent_swapped:
|
||||||
from ..EvoScientist import _ensure_config
|
from ..EvoScientist import _ensure_config
|
||||||
@@ -273,6 +319,165 @@ def _stopped_response_after_narration(
|
|||||||
return display_current, display_stopped, full_stopped
|
return display_current, display_stopped, full_stopped
|
||||||
|
|
||||||
|
|
||||||
|
# (kind, payload, item_index): kind is "header"/"sep"/"item"; payload is the
|
||||||
|
# category name for headers or the candidate for items; item_index is the
|
||||||
|
# candidate's position in the source list (-1 for non-item rows).
|
||||||
|
_CompletionRow = tuple[str, Any, int]
|
||||||
|
|
||||||
|
|
||||||
|
def _build_completion_rows(items: list[Any]) -> list[_CompletionRow]:
|
||||||
|
"""Flatten completion candidates into render rows with category headers."""
|
||||||
|
rows: list[_CompletionRow] = []
|
||||||
|
last_cat = ""
|
||||||
|
for i, candidate in enumerate(items):
|
||||||
|
cat = getattr(candidate, "category", "")
|
||||||
|
if cat and cat != last_cat:
|
||||||
|
if last_cat:
|
||||||
|
rows.append(("sep", "", -1))
|
||||||
|
rows.append(("header", cat, -1))
|
||||||
|
last_cat = cat
|
||||||
|
rows.append(("item", candidate, i))
|
||||||
|
return rows
|
||||||
|
|
||||||
|
|
||||||
|
def _window_completion_rows(
|
||||||
|
rows: list[_CompletionRow],
|
||||||
|
selected: int,
|
||||||
|
max_rows: int,
|
||||||
|
) -> tuple[list[_CompletionRow], int, int]:
|
||||||
|
"""Slice *rows* to a window of at most *max_rows* total display lines.
|
||||||
|
|
||||||
|
The window always contains the selected item (top of the list when
|
||||||
|
nothing is selected) and reserves one line per overflow indicator.
|
||||||
|
Returns ``(visible_rows, hidden_items_above, hidden_items_below)``.
|
||||||
|
"""
|
||||||
|
max_rows = max(max_rows, 5)
|
||||||
|
if len(rows) <= max_rows:
|
||||||
|
return list(rows), 0, 0
|
||||||
|
|
||||||
|
sel_row = 0
|
||||||
|
if selected >= 0:
|
||||||
|
for r, (kind, _payload, idx) in enumerate(rows):
|
||||||
|
if kind == "item" and idx == selected:
|
||||||
|
sel_row = r
|
||||||
|
break
|
||||||
|
|
||||||
|
# Center the selection; centering keeps it clear of the indicator
|
||||||
|
# lines that replace the window's edge rows when content is clipped.
|
||||||
|
start = min(max(sel_row - max_rows // 2, 0), len(rows) - max_rows)
|
||||||
|
end = start + max_rows
|
||||||
|
content_start = start + (1 if start > 0 else 0)
|
||||||
|
content_end = end - (1 if end < len(rows) else 0)
|
||||||
|
|
||||||
|
above = sum(1 for kind, _p, _i in rows[:content_start] if kind == "item")
|
||||||
|
below = sum(1 for kind, _p, _i in rows[content_end:] if kind == "item")
|
||||||
|
return rows[content_start:content_end], above, below
|
||||||
|
|
||||||
|
|
||||||
|
def _render_completion_text(items: list[Any], selected: int, max_rows: int) -> Text:
|
||||||
|
"""Render the completion popup content bounded to *max_rows* lines."""
|
||||||
|
rows = _build_completion_rows(items)
|
||||||
|
visible, above, below = _window_completion_rows(rows, selected, max_rows)
|
||||||
|
|
||||||
|
# Blank separator lines are cosmetic — drop them at the window edges.
|
||||||
|
while visible and visible[0][0] == "sep":
|
||||||
|
visible = visible[1:]
|
||||||
|
while visible and visible[-1][0] == "sep":
|
||||||
|
visible = visible[:-1]
|
||||||
|
|
||||||
|
lines: list[Text] = []
|
||||||
|
if above:
|
||||||
|
lines.append(Text(f" ↑ {above} more", style="dim italic"))
|
||||||
|
for kind, payload, idx in visible:
|
||||||
|
if kind == "sep":
|
||||||
|
lines.append(Text())
|
||||||
|
elif kind == "header":
|
||||||
|
lines.append(Text(f" {payload}", style="bold #6b7280"))
|
||||||
|
elif idx == selected:
|
||||||
|
lines.append(
|
||||||
|
Text.assemble(
|
||||||
|
(" ▸ ", "bold"),
|
||||||
|
(f"{payload.text:<28}", "bold"),
|
||||||
|
(payload.description, "bold"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
lines.append(
|
||||||
|
Text.assemble(
|
||||||
|
(" ", "#888888"),
|
||||||
|
(f"{payload.text:<28}", "#888888"),
|
||||||
|
(payload.description, "#888888"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if below:
|
||||||
|
lines.append(Text(f" ↓ {below} more", style="dim italic"))
|
||||||
|
return Text("\n").join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
# Hard cap on popup lines so the popup never dwarfs the chat area
|
||||||
|
# (mainstream CLI behavior); matches the pre-#354 max-height.
|
||||||
|
_COMPLETION_MAX_VISIBLE_ROWS = 15
|
||||||
|
# Rows kept free for the input row, status bar and a slice of chat. On
|
||||||
|
# terminals shorter than ~17 rows the 5-row floor wins over this
|
||||||
|
# reservation — a smaller popup would be unusable.
|
||||||
|
_COMPLETION_RESERVED_ROWS = 12
|
||||||
|
|
||||||
|
|
||||||
|
def _completion_row_budget(height: int) -> int:
|
||||||
|
"""Popup line budget for a terminal of *height* rows."""
|
||||||
|
if height <= 0:
|
||||||
|
return _COMPLETION_MAX_VISIBLE_ROWS
|
||||||
|
return max(5, min(height - _COMPLETION_RESERVED_ROWS, _COMPLETION_MAX_VISIBLE_ROWS))
|
||||||
|
|
||||||
|
|
||||||
|
# Textual converts rich Text to Content and drops rich no_wrap/overflow
|
||||||
|
# attributes, so line cropping must be enforced here in CSS.
|
||||||
|
_COMPLETIONS_CSS = """
|
||||||
|
#completions {
|
||||||
|
display: none;
|
||||||
|
height: auto;
|
||||||
|
background: #1e1f26;
|
||||||
|
padding: 0 1;
|
||||||
|
border-bottom: solid #0284c7;
|
||||||
|
text-wrap: nowrap;
|
||||||
|
text-overflow: ellipsis;
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_chat_scroll(container: Any) -> None:
|
||||||
|
"""Repair the chat scroll state after the popup resized the viewport.
|
||||||
|
|
||||||
|
Textual's compositor recomputes ``scroll_y`` for anchored containers
|
||||||
|
bypassing the validator, so when the popup hides and the content fits
|
||||||
|
again, ``scroll_y`` can go negative — the scrollbar then renders as if
|
||||||
|
scrolled to the bottom while the content sits at the top (issue #301
|
||||||
|
family). Runs after refresh so sizes are current.
|
||||||
|
"""
|
||||||
|
# force=True: with the content fitting, the scrollbar is hidden and
|
||||||
|
# allow_vertical_scroll is False — an unforced scroll_home would
|
||||||
|
# silently no-op and leave the negative scroll_y in place.
|
||||||
|
if container.is_anchored:
|
||||||
|
if container.max_scroll_y <= 0:
|
||||||
|
container.anchor(False)
|
||||||
|
container.scroll_home(animate=False, immediate=True, force=True)
|
||||||
|
elif container.scroll_y < 0:
|
||||||
|
container.scroll_home(animate=False, immediate=True, force=True)
|
||||||
|
# Resync the scrollbar thumb: watch_scroll_y skips the update while
|
||||||
|
# the scrollbar is hidden (or when the compositor wrote scroll_y via
|
||||||
|
# set_reactive), so a stale position survives until the scrollbar
|
||||||
|
# reappears — rendering as "scrolled to bottom" at the top.
|
||||||
|
scrollbar = getattr(container, "vertical_scrollbar", None)
|
||||||
|
if scrollbar is not None and scrollbar.position != container.scroll_y:
|
||||||
|
scrollbar.position = container.scroll_y
|
||||||
|
|
||||||
|
|
||||||
|
def _session_auto_approve_decisions(action_requests: list) -> list[dict]:
|
||||||
|
"""TUI session "approve all": an explicit human opt-in, so blanket-approve
|
||||||
|
everything (dangerous set included), matching the Rich CLI and channel."""
|
||||||
|
return [{"type": "approve"} for _ in action_requests]
|
||||||
|
|
||||||
|
|
||||||
def run_textual_interactive(
|
def run_textual_interactive(
|
||||||
*,
|
*,
|
||||||
show_thinking: bool,
|
show_thinking: bool,
|
||||||
@@ -287,6 +492,7 @@ def run_textual_interactive(
|
|||||||
load_agent: Callable[..., Any],
|
load_agent: Callable[..., Any],
|
||||||
create_session_workspace: Callable[[str | None], str],
|
create_session_workspace: Callable[[str | None], str],
|
||||||
config: Any | None = None,
|
config: Any | None = None,
|
||||||
|
async_runtime: AsyncRuntime | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Run full-screen Textual interactive chat loop."""
|
"""Run full-screen Textual interactive chat loop."""
|
||||||
if config is None:
|
if config is None:
|
||||||
@@ -294,7 +500,15 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
config = get_effective_config()
|
config = get_effective_config()
|
||||||
|
|
||||||
runtime_gateways = create_runtime_gateways()
|
# One frontend event sink for the whole TUI session — injected into the
|
||||||
|
# agent's middleware (write side) and the local gateway's streaming path
|
||||||
|
# (read side). The fallback-notice display is bound to the App's
|
||||||
|
# _append_system once the App exists (on_mount); tool-selection needs no
|
||||||
|
# display hook (its widget is mounted from the stream event).
|
||||||
|
from ..stream.sink import SessionEventSink
|
||||||
|
|
||||||
|
event_sink = SessionEventSink()
|
||||||
|
runtime_gateways = create_runtime_gateways(events=event_sink)
|
||||||
graph_gateway = runtime_gateways.graph_gateway
|
graph_gateway = runtime_gateways.graph_gateway
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -311,6 +525,7 @@ def run_textual_interactive(
|
|||||||
CompactingWidget,
|
CompactingWidget,
|
||||||
LoadingWidget,
|
LoadingWidget,
|
||||||
MCPLoaderWidget,
|
MCPLoaderWidget,
|
||||||
|
PanelWidget,
|
||||||
SubAgentWidget,
|
SubAgentWidget,
|
||||||
SummarizationWidget,
|
SummarizationWidget,
|
||||||
SystemMessage,
|
SystemMessage,
|
||||||
@@ -333,7 +548,8 @@ def run_textual_interactive(
|
|||||||
def supports_interactive(self) -> bool:
|
def supports_interactive(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
CSS = """
|
CSS = (
|
||||||
|
"""
|
||||||
Screen {
|
Screen {
|
||||||
layout: vertical;
|
layout: vertical;
|
||||||
background: #16161a;
|
background: #16161a;
|
||||||
@@ -385,14 +601,9 @@ def run_textual_interactive(
|
|||||||
padding: 0 2;
|
padding: 0 2;
|
||||||
color: #9ca3af;
|
color: #9ca3af;
|
||||||
}
|
}
|
||||||
#completions {
|
"""
|
||||||
display: none;
|
+ _COMPLETIONS_CSS
|
||||||
height: auto;
|
+ """
|
||||||
max-height: 15;
|
|
||||||
background: #1e1f26;
|
|
||||||
padding: 0 1;
|
|
||||||
border-bottom: solid #0284c7;
|
|
||||||
}
|
|
||||||
#status {
|
#status {
|
||||||
height: 1;
|
height: 1;
|
||||||
min-height: 1;
|
min-height: 1;
|
||||||
@@ -401,6 +612,7 @@ def run_textual_interactive(
|
|||||||
padding: 0 1;
|
padding: 0 1;
|
||||||
}
|
}
|
||||||
"""
|
"""
|
||||||
|
)
|
||||||
BINDINGS: ClassVar[list[Binding]] = [
|
BINDINGS: ClassVar[list[Binding]] = [
|
||||||
Binding("ctrl+c", "request_quit", "Quit", show=False, priority=True),
|
Binding("ctrl+c", "request_quit", "Quit", show=False, priority=True),
|
||||||
Binding("ctrl+v", "paste_clipboard", "Paste", show=False),
|
Binding("ctrl+v", "paste_clipboard", "Paste", show=False),
|
||||||
@@ -438,7 +650,8 @@ def run_textual_interactive(
|
|||||||
self._resumed = resumed
|
self._resumed = resumed
|
||||||
self._resume_warning = resume_warning
|
self._resume_warning = resume_warning
|
||||||
self._channel_timer: Any = None
|
self._channel_timer: Any = None
|
||||||
self._started_channel_types: list[str] = []
|
self._channel_start_results: list[tuple[str, bool, str]] = []
|
||||||
|
self._channel_start_stop = threading.Event()
|
||||||
self._busy = False
|
self._busy = False
|
||||||
self._notification_consuming: bool = (
|
self._notification_consuming: bool = (
|
||||||
False # prevent overlapping consume coroutines
|
False # prevent overlapping consume coroutines
|
||||||
@@ -449,6 +662,7 @@ def run_textual_interactive(
|
|||||||
] = [] # queued messages to send after current turn
|
] = [] # queued messages to send after current turn
|
||||||
self._comp_items: list = []
|
self._comp_items: list = []
|
||||||
self._comp_index: int = -1
|
self._comp_index: int = -1
|
||||||
|
self._comp_last_height: int = 0
|
||||||
self._comp_base: str = ""
|
self._comp_base: str = ""
|
||||||
self._hitl_auto_approve: bool = False
|
self._hitl_auto_approve: bool = False
|
||||||
self._approval_future: asyncio.Future | None = None
|
self._approval_future: asyncio.Future | None = None
|
||||||
@@ -465,6 +679,7 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
self._channel_runtime = ChannelRuntime()
|
self._channel_runtime = ChannelRuntime()
|
||||||
self._quit_pending: bool = False
|
self._quit_pending: bool = False
|
||||||
|
self._exiting: bool = False
|
||||||
self._current_model: str | None = model
|
self._current_model: str | None = model
|
||||||
self._current_provider: str | None = provider
|
self._current_provider: str | None = provider
|
||||||
self._status_started_at = datetime.now()
|
self._status_started_at = datetime.now()
|
||||||
@@ -516,6 +731,7 @@ def run_textual_interactive(
|
|||||||
self._agent_loader.start(
|
self._agent_loader.start(
|
||||||
workspace_dir=workspace,
|
workspace_dir=workspace,
|
||||||
checkpointer=self._checkpointer,
|
checkpointer=self._checkpointer,
|
||||||
|
events=self._runtime_gateways.graph_gateway.events,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _mount_mcp_loader_widget(self) -> None:
|
def _mount_mcp_loader_widget(self) -> None:
|
||||||
@@ -794,11 +1010,13 @@ def run_textual_interactive(
|
|||||||
yield Static("", id="status")
|
yield Static("", id="status")
|
||||||
|
|
||||||
def on_mount(self) -> None:
|
def on_mount(self) -> None:
|
||||||
# Register fallback middleware UI callback so messages appear
|
# Bind the session sink's fallback-notice display so model-fallback
|
||||||
# as SystemMessage widgets in the chat container.
|
# messages appear as SystemMessage widgets in the chat container.
|
||||||
from ..middleware.model_fallback import set_ui_emit
|
# ``event_sink`` is the concrete SessionEventSink created by the
|
||||||
|
# enclosing factory — the same instance the gateway carries.
|
||||||
set_ui_emit(lambda text, style: self._append_system(text, style))
|
event_sink.set_fallback_display(
|
||||||
|
lambda text, style: self._append_system(text, style)
|
||||||
|
)
|
||||||
|
|
||||||
self._render_welcome()
|
self._render_welcome()
|
||||||
self._render_status()
|
self._render_status()
|
||||||
@@ -850,12 +1068,23 @@ def run_textual_interactive(
|
|||||||
exc_info=True,
|
exc_info=True,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
self._start_channels()
|
await self._start_channels()
|
||||||
|
|
||||||
ch_task = asyncio.create_task(_deferred_start_channels())
|
ch_task = asyncio.create_task(_deferred_start_channels())
|
||||||
self._background_tasks.add(ch_task)
|
self._background_tasks.add(ch_task)
|
||||||
ch_task.add_done_callback(self._background_tasks.discard)
|
ch_task.add_done_callback(self._background_tasks.discard)
|
||||||
|
|
||||||
|
def on_resize(self, event: Any) -> None:
|
||||||
|
"""Re-window the completion popup for the new terminal height."""
|
||||||
|
try:
|
||||||
|
comp_widget = self.query_one("#completions", Static)
|
||||||
|
except Exception:
|
||||||
|
return
|
||||||
|
if comp_widget.display and self._comp_items:
|
||||||
|
# Deferred: this handler can run before the base App
|
||||||
|
# handler updates self.size with the new dimensions.
|
||||||
|
self.call_after_refresh(self._render_completions)
|
||||||
|
|
||||||
# ── Update check ──────────────────────────────────────
|
# ── Update check ──────────────────────────────────────
|
||||||
|
|
||||||
async def _check_for_updates(self) -> None:
|
async def _check_for_updates(self) -> None:
|
||||||
@@ -877,28 +1106,46 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
# ── Channel integration ────────────────────────────────
|
# ── Channel integration ────────────────────────────────
|
||||||
|
|
||||||
def _start_channels(self) -> None:
|
async def _start_channels(self) -> None:
|
||||||
"""Auto-start channels if enabled in config."""
|
"""Auto-start channels if enabled in config."""
|
||||||
try:
|
try:
|
||||||
from ..config import load_config
|
from ..config import load_config
|
||||||
|
|
||||||
cfg = load_config()
|
cfg = await asyncio.to_thread(load_config)
|
||||||
if cfg and cfg.channel_enabled and not _channels_is_running():
|
if cfg and cfg.channel_enabled and not _channels_is_running():
|
||||||
_auto_start_channel(
|
results = await _auto_start_channel_in_worker(
|
||||||
self._agent_loader.agent,
|
self._agent_loader.agent,
|
||||||
self._conversation_tid,
|
self._conversation_tid,
|
||||||
cfg,
|
cfg,
|
||||||
send_thinking=self._channel_send_thinking,
|
send_thinking=self._channel_send_thinking,
|
||||||
runtime=self._channel_runtime,
|
runtime=self._channel_runtime,
|
||||||
|
stop_requested=self._channel_start_stop,
|
||||||
)
|
)
|
||||||
types = [
|
if self._exiting:
|
||||||
t.strip() for t in cfg.channel_enabled.split(",") if t.strip()
|
return
|
||||||
]
|
current_agent = self._agent_loader.agent
|
||||||
self._started_channel_types = types
|
if current_agent is not None and _channels_is_running():
|
||||||
|
self._channel_runtime.bind(
|
||||||
|
current_agent,
|
||||||
|
self._conversation_tid,
|
||||||
|
)
|
||||||
|
self._channel_start_results = results
|
||||||
self._render_welcome()
|
self._render_welcome()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
self._channel_start_stop.set()
|
||||||
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
_channel_logger.debug(f"Channel auto-start failed: {e}")
|
_channel_logger.debug(f"Channel auto-start failed: {e}")
|
||||||
self._channel_timer = self.set_interval(0.1, self._poll_channel_queue)
|
finally:
|
||||||
|
if (
|
||||||
|
not self._exiting
|
||||||
|
and not self._channel_start_stop.is_set()
|
||||||
|
and self._channel_timer is None
|
||||||
|
):
|
||||||
|
self._channel_timer = self.set_interval(
|
||||||
|
0.1,
|
||||||
|
self._poll_channel_queue,
|
||||||
|
)
|
||||||
|
|
||||||
def _poll_channel_queue(self) -> None:
|
def _poll_channel_queue(self) -> None:
|
||||||
"""Poll the channel + notification queues (every 100ms)."""
|
"""Poll the channel + notification queues (every 100ms)."""
|
||||||
@@ -1075,7 +1322,9 @@ def run_textual_interactive(
|
|||||||
go negative and pushes the welcome banner out of view (issue #301).
|
go negative and pushes the welcome banner out of view (issue #301).
|
||||||
"""
|
"""
|
||||||
container.anchor(False)
|
container.anchor(False)
|
||||||
container.scroll_home(animate=False, immediate=True)
|
# force=True: with content fitting, the scrollbar is hidden and
|
||||||
|
# allow_vertical_scroll is False — unforced scroll_home no-ops.
|
||||||
|
container.scroll_home(animate=False, immediate=True, force=True)
|
||||||
|
|
||||||
def _append_system(self, text: str, style: str = "dim") -> None:
|
def _append_system(self, text: str, style: str = "dim") -> None:
|
||||||
"""Mount a SystemMessage widget into #chat."""
|
"""Mount a SystemMessage widget into #chat."""
|
||||||
@@ -1125,7 +1374,7 @@ def run_textual_interactive(
|
|||||||
Returns the ``ApprovalWidget.Decided`` message, or ``None`` on
|
Returns the ``ApprovalWidget.Decided`` message, or ``None`` on
|
||||||
timeout / cancellation.
|
timeout / cancellation.
|
||||||
"""
|
"""
|
||||||
self._approval_future = asyncio.get_event_loop().create_future()
|
self._approval_future = asyncio.get_running_loop().create_future()
|
||||||
try:
|
try:
|
||||||
return await asyncio.wait_for(self._approval_future, timeout=300)
|
return await asyncio.wait_for(self._approval_future, timeout=300)
|
||||||
except (TimeoutError, asyncio.CancelledError):
|
except (TimeoutError, asyncio.CancelledError):
|
||||||
@@ -1171,7 +1420,7 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
Returns the selected thread_id, or ``None`` on cancel/timeout.
|
Returns the selected thread_id, or ``None`` on cancel/timeout.
|
||||||
"""
|
"""
|
||||||
self._picker_future = asyncio.get_event_loop().create_future()
|
self._picker_future = asyncio.get_running_loop().create_future()
|
||||||
try:
|
try:
|
||||||
return await asyncio.wait_for(self._picker_future, timeout=120)
|
return await asyncio.wait_for(self._picker_future, timeout=120)
|
||||||
except (TimeoutError, asyncio.CancelledError):
|
except (TimeoutError, asyncio.CancelledError):
|
||||||
@@ -1199,7 +1448,7 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
Returns list of install sources, or None on cancel/timeout.
|
Returns list of install sources, or None on cancel/timeout.
|
||||||
"""
|
"""
|
||||||
self._browser_future = asyncio.get_event_loop().create_future()
|
self._browser_future = asyncio.get_running_loop().create_future()
|
||||||
try:
|
try:
|
||||||
return await asyncio.wait_for(self._browser_future, timeout=300)
|
return await asyncio.wait_for(self._browser_future, timeout=300)
|
||||||
except (TimeoutError, asyncio.CancelledError):
|
except (TimeoutError, asyncio.CancelledError):
|
||||||
@@ -1226,7 +1475,7 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
async def _wait_for_mcp_browse(self, browser_widget) -> list | None:
|
async def _wait_for_mcp_browse(self, browser_widget) -> list | None:
|
||||||
"""Wait for user to complete MCP server browsing."""
|
"""Wait for user to complete MCP server browsing."""
|
||||||
self._mcp_browser_future = asyncio.get_event_loop().create_future()
|
self._mcp_browser_future = asyncio.get_running_loop().create_future()
|
||||||
try:
|
try:
|
||||||
return await asyncio.wait_for(self._mcp_browser_future, timeout=300)
|
return await asyncio.wait_for(self._mcp_browser_future, timeout=300)
|
||||||
except (TimeoutError, asyncio.CancelledError):
|
except (TimeoutError, asyncio.CancelledError):
|
||||||
@@ -1254,7 +1503,7 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
Returns ``(name, provider)`` or ``None`` on cancel/timeout.
|
Returns ``(name, provider)`` or ``None`` on cancel/timeout.
|
||||||
"""
|
"""
|
||||||
self._model_picker_future = asyncio.get_event_loop().create_future()
|
self._model_picker_future = asyncio.get_running_loop().create_future()
|
||||||
try:
|
try:
|
||||||
return await asyncio.wait_for(self._model_picker_future, timeout=120)
|
return await asyncio.wait_for(self._model_picker_future, timeout=120)
|
||||||
except (TimeoutError, asyncio.CancelledError):
|
except (TimeoutError, asyncio.CancelledError):
|
||||||
@@ -1316,6 +1565,7 @@ def run_textual_interactive(
|
|||||||
"""
|
"""
|
||||||
from ..stream.display import (
|
from ..stream.display import (
|
||||||
is_stream_cancel_requested,
|
is_stream_cancel_requested,
|
||||||
|
iter_with_stream_cancel,
|
||||||
)
|
)
|
||||||
|
|
||||||
container = self.query_one("#chat", VerticalScroll)
|
container = self.query_one("#chat", VerticalScroll)
|
||||||
@@ -1341,6 +1591,7 @@ def run_textual_interactive(
|
|||||||
todo_w: TodoWidget | None = None
|
todo_w: TodoWidget | None = None
|
||||||
tool_widgets: dict[str, ToolCallWidget] = {}
|
tool_widgets: dict[str, ToolCallWidget] = {}
|
||||||
subagent_widgets: dict[str, SubAgentWidget] = {}
|
subagent_widgets: dict[str, SubAgentWidget] = {}
|
||||||
|
panel_widgets: dict[str, PanelWidget] = {}
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class _ResponseDisplayState:
|
class _ResponseDisplayState:
|
||||||
@@ -1542,16 +1793,26 @@ def run_textual_interactive(
|
|||||||
summarization_w = None
|
summarization_w = None
|
||||||
try:
|
try:
|
||||||
_anchor_engaged = False
|
_anchor_engaged = False
|
||||||
async for event in graph_gateway.stream_events(
|
_active_teams = list(self._channel_runtime.active_teams)
|
||||||
RunRequest(
|
_configurable_extra = (
|
||||||
message=_stream_input,
|
{"active_teams": _active_teams} if _active_teams else None
|
||||||
thread_id=thread_id_override or self._conversation_tid,
|
)
|
||||||
metadata=metadata,
|
async for event in iter_with_stream_cancel(
|
||||||
target=GraphTarget(
|
graph_gateway.stream_events(
|
||||||
local_graph=agent,
|
RunRequest(
|
||||||
workspace_dir=self._workspace_dir,
|
message=_stream_input,
|
||||||
),
|
thread_id=(
|
||||||
)
|
thread_id_override or self._conversation_tid
|
||||||
|
),
|
||||||
|
metadata=metadata,
|
||||||
|
target=GraphTarget(
|
||||||
|
local_graph=agent,
|
||||||
|
workspace_dir=self._workspace_dir,
|
||||||
|
),
|
||||||
|
configurable_extra=_configurable_extra,
|
||||||
|
)
|
||||||
|
),
|
||||||
|
cancel_scope,
|
||||||
):
|
):
|
||||||
if is_stream_cancel_requested(cancel_scope):
|
if is_stream_cancel_requested(cancel_scope):
|
||||||
response = await _mark_cancelled_response()
|
response = await _mark_cancelled_response()
|
||||||
@@ -1846,6 +2107,40 @@ def run_textual_interactive(
|
|||||||
if sa_w is not None:
|
if sa_w is not None:
|
||||||
sa_w.finalize()
|
sa_w.finalize()
|
||||||
|
|
||||||
|
elif event_type == "panel_dispatch_start":
|
||||||
|
eval_id = event.get("eval_id", "") or "_unbatched"
|
||||||
|
panel_w = panel_widgets.get(eval_id)
|
||||||
|
if panel_w is None:
|
||||||
|
panel_w = PanelWidget(eval_id)
|
||||||
|
# Register before awaiting mount: a cancel
|
||||||
|
# during the await would otherwise orphan a
|
||||||
|
# ticking panel outside the cleanup loop.
|
||||||
|
panel_widgets[eval_id] = panel_w
|
||||||
|
await container.mount(panel_w)
|
||||||
|
await panel_w.start_dispatch(
|
||||||
|
event["id"],
|
||||||
|
event.get("subagent_type", ""),
|
||||||
|
event.get("label", "") or event.get("description", ""),
|
||||||
|
)
|
||||||
|
|
||||||
|
elif event_type == "panel_dispatch_complete":
|
||||||
|
eval_id = event.get("eval_id", "") or "_unbatched"
|
||||||
|
panel_w = panel_widgets.get(eval_id)
|
||||||
|
if panel_w is not None:
|
||||||
|
panel_w.complete_dispatch(
|
||||||
|
event["id"], int(event.get("duration_ms", 0))
|
||||||
|
)
|
||||||
|
|
||||||
|
elif event_type == "panel_dispatch_error":
|
||||||
|
eval_id = event.get("eval_id", "") or "_unbatched"
|
||||||
|
panel_w = panel_widgets.get(eval_id)
|
||||||
|
if panel_w is not None:
|
||||||
|
panel_w.fail_dispatch(
|
||||||
|
event["id"],
|
||||||
|
int(event.get("duration_ms", 0)),
|
||||||
|
event.get("error", ""),
|
||||||
|
)
|
||||||
|
|
||||||
elif event_type == "ask_user":
|
elif event_type == "ask_user":
|
||||||
questions = event.get("questions", [])
|
questions = event.get("questions", [])
|
||||||
if questions:
|
if questions:
|
||||||
@@ -1888,20 +2183,15 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
elif event_type == "interrupt":
|
elif event_type == "interrupt":
|
||||||
action_reqs = event.get("action_requests", [])
|
action_reqs = event.get("action_requests", [])
|
||||||
n = len(action_reqs) or 1
|
interrupt_id = event.get("interrupt_id")
|
||||||
|
|
||||||
# HITL: check session auto-approve first
|
# HITL: session "approve all" blanket-approves.
|
||||||
if self._hitl_auto_approve:
|
if self._hitl_auto_approve:
|
||||||
from langgraph.types import (
|
from ..backends import build_hitl_resume
|
||||||
Command, # type: ignore[import-untyped]
|
|
||||||
)
|
|
||||||
|
|
||||||
_stream_input = Command(
|
decisions = _session_auto_approve_decisions(action_reqs)
|
||||||
resume={
|
_stream_input = build_hitl_resume(
|
||||||
"decisions": [
|
interrupt_id, decisions
|
||||||
{"type": "approve"} for _ in range(n)
|
|
||||||
]
|
|
||||||
}
|
|
||||||
)
|
)
|
||||||
_hitl_resuming = True
|
_hitl_resuming = True
|
||||||
break # re-enter outer HITL loop
|
break # re-enter outer HITL loop
|
||||||
@@ -1921,12 +2211,10 @@ def run_textual_interactive(
|
|||||||
response = await _mark_cancelled_response()
|
response = await _mark_cancelled_response()
|
||||||
break
|
break
|
||||||
if decisions is not None:
|
if decisions is not None:
|
||||||
from langgraph.types import (
|
from ..backends import build_hitl_resume
|
||||||
Command, # type: ignore[import-untyped]
|
|
||||||
)
|
|
||||||
|
|
||||||
_stream_input = Command(
|
_stream_input = build_hitl_resume(
|
||||||
resume={"decisions": decisions}
|
interrupt_id, decisions
|
||||||
)
|
)
|
||||||
_hitl_resuming = True
|
_hitl_resuming = True
|
||||||
break # re-enter outer HITL loop
|
break # re-enter outer HITL loop
|
||||||
@@ -1955,12 +2243,10 @@ def run_textual_interactive(
|
|||||||
if decided_event and decided_event.decisions is not None:
|
if decided_event and decided_event.decisions is not None:
|
||||||
if decided_event.auto_approve_session:
|
if decided_event.auto_approve_session:
|
||||||
self._hitl_auto_approve = True
|
self._hitl_auto_approve = True
|
||||||
from langgraph.types import (
|
from ..backends import build_hitl_resume
|
||||||
Command, # type: ignore[import-untyped]
|
|
||||||
)
|
|
||||||
|
|
||||||
_stream_input = Command(
|
_stream_input = build_hitl_resume(
|
||||||
resume={"decisions": decided_event.decisions}
|
interrupt_id, decided_event.decisions
|
||||||
)
|
)
|
||||||
_hitl_resuming = True
|
_hitl_resuming = True
|
||||||
break # re-enter outer HITL loop with resume
|
break # re-enter outer HITL loop with resume
|
||||||
@@ -2075,6 +2361,13 @@ def run_textual_interactive(
|
|||||||
sa_w.finalize()
|
sa_w.finalize()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
# Finalize any still-running panel dispatches so their
|
||||||
|
# per-row spinner timers stop instead of ticking forever.
|
||||||
|
for panel_w in panel_widgets.values():
|
||||||
|
try:
|
||||||
|
panel_w.finalize_running()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
# Finalize thinking widget
|
# Finalize thinking widget
|
||||||
if thinking_w is not None and thinking_w._is_active:
|
if thinking_w is not None and thinking_w._is_active:
|
||||||
try:
|
try:
|
||||||
@@ -2142,6 +2435,11 @@ def run_textual_interactive(
|
|||||||
cancelled = False
|
cancelled = False
|
||||||
response = ""
|
response = ""
|
||||||
try:
|
try:
|
||||||
|
# Foreground turns share the legacy default scope. Reset it at
|
||||||
|
# the turn boundary; scoped channel stop requests remain armed.
|
||||||
|
from ..stream.display import clear_stream_cancel
|
||||||
|
|
||||||
|
clear_stream_cancel()
|
||||||
self._busy = True
|
self._busy = True
|
||||||
self._turn_started_at = datetime.now()
|
self._turn_started_at = datetime.now()
|
||||||
self._status_phase = ResearchPhase.THINKING
|
self._status_phase = ResearchPhase.THINKING
|
||||||
@@ -2176,6 +2474,17 @@ def run_textual_interactive(
|
|||||||
)
|
)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
cancelled = True
|
cancelled = True
|
||||||
|
try:
|
||||||
|
from ..middleware.code_interpreter import (
|
||||||
|
aclose_code_interpreters,
|
||||||
|
)
|
||||||
|
|
||||||
|
await aclose_code_interpreters()
|
||||||
|
except Exception:
|
||||||
|
_channel_logger.debug(
|
||||||
|
"code interpreter cleanup after cancellation failed",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
self._append_system("\nInterrupted by user", style="dim italic #ffe082")
|
self._append_system("\nInterrupted by user", style="dim italic #ffe082")
|
||||||
finally:
|
finally:
|
||||||
self._busy = False
|
self._busy = False
|
||||||
@@ -2309,6 +2618,7 @@ def run_textual_interactive(
|
|||||||
on_cmd_completed=self._on_channel_cmd_completed,
|
on_cmd_completed=self._on_channel_cmd_completed,
|
||||||
channel_runtime=self._channel_runtime,
|
channel_runtime=self._channel_runtime,
|
||||||
graph_gateway=self._runtime_gateways.graph_gateway,
|
graph_gateway=self._runtime_gateways.graph_gateway,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
if _slash_handled:
|
if _slash_handled:
|
||||||
# A channel-issued /new or /resume rotates the thread in
|
# A channel-issued /new or /resume rotates the thread in
|
||||||
@@ -2784,30 +3094,36 @@ def run_textual_interactive(
|
|||||||
def _hide_completions(self) -> None:
|
def _hide_completions(self) -> None:
|
||||||
self._comp_items = []
|
self._comp_items = []
|
||||||
self._comp_index = -1
|
self._comp_index = -1
|
||||||
self.query_one("#completions", Static).display = False
|
self._comp_last_height = 0
|
||||||
|
comp_widget = self.query_one("#completions", Static)
|
||||||
|
was_visible = comp_widget.display
|
||||||
|
comp_widget.display = False
|
||||||
|
# Called on every ordinary input change — only a popup that was
|
||||||
|
# actually visible changed the chat viewport.
|
||||||
|
if was_visible:
|
||||||
|
self.call_after_refresh(self._normalize_chat_after_popup)
|
||||||
|
|
||||||
|
def _completion_max_rows(self) -> int:
|
||||||
|
return _completion_row_budget(int(getattr(self.size, "height", 0) or 0))
|
||||||
|
|
||||||
def _render_completions(self) -> None:
|
def _render_completions(self) -> None:
|
||||||
comp_text = Text()
|
comp_text = _render_completion_text(
|
||||||
last_cat = ""
|
self._comp_items, self._comp_index, self._completion_max_rows()
|
||||||
for i, candidate in enumerate(self._comp_items):
|
)
|
||||||
cmd, desc = candidate.text, candidate.description
|
|
||||||
cat = getattr(candidate, "category", "")
|
|
||||||
if cat and cat != last_cat:
|
|
||||||
if last_cat:
|
|
||||||
comp_text.append("\n")
|
|
||||||
comp_text.append(f" {cat}\n", style="bold #6b7280")
|
|
||||||
last_cat = cat
|
|
||||||
if i == self._comp_index:
|
|
||||||
comp_text.append(" \u25b8 ", style="bold")
|
|
||||||
comp_text.append(f"{cmd:<28}", style="bold")
|
|
||||||
comp_text.append(desc, style="bold")
|
|
||||||
else:
|
|
||||||
comp_text.append(" ", style="#888888")
|
|
||||||
comp_text.append(f"{cmd:<28}", style="#888888")
|
|
||||||
comp_text.append(desc, style="#888888")
|
|
||||||
if i < len(self._comp_items) - 1:
|
|
||||||
comp_text.append("\n")
|
|
||||||
self.query_one("#completions", Static).update(comp_text)
|
self.query_one("#completions", Static).update(comp_text)
|
||||||
|
# Selection-only navigation keeps the height — skip the (cheap
|
||||||
|
# but per-keystroke) normalize unless the viewport can change.
|
||||||
|
n_lines = len(comp_text.plain.splitlines()) if comp_text.plain else 0
|
||||||
|
if n_lines != self._comp_last_height:
|
||||||
|
self._comp_last_height = n_lines
|
||||||
|
self.call_after_refresh(self._normalize_chat_after_popup)
|
||||||
|
|
||||||
|
def _normalize_chat_after_popup(self) -> None:
|
||||||
|
try:
|
||||||
|
container = self.query_one("#chat", VerticalScroll)
|
||||||
|
except Exception:
|
||||||
|
return
|
||||||
|
_normalize_chat_scroll(container)
|
||||||
|
|
||||||
# ── Slash commands ─────────────────────────────────────
|
# ── Slash commands ─────────────────────────────────────
|
||||||
|
|
||||||
@@ -2847,6 +3163,7 @@ def run_textual_interactive(
|
|||||||
input_tokens_hint=self._status_last_input_tokens,
|
input_tokens_hint=self._status_last_input_tokens,
|
||||||
channel_runtime=self._channel_runtime,
|
channel_runtime=self._channel_runtime,
|
||||||
graph_gateway=self._runtime_gateways.graph_gateway,
|
graph_gateway=self._runtime_gateways.graph_gateway,
|
||||||
|
async_runtime=async_runtime,
|
||||||
)
|
)
|
||||||
|
|
||||||
if await cmd_manager.execute(command, ctx):
|
if await cmd_manager.execute(command, ctx):
|
||||||
@@ -2864,8 +3181,9 @@ def run_textual_interactive(
|
|||||||
self._render_status()
|
self._render_status()
|
||||||
finally:
|
finally:
|
||||||
self._busy = False
|
self._busy = False
|
||||||
prompt_widget.disabled = False
|
if not self._exiting:
|
||||||
prompt_widget.focus()
|
prompt_widget.disabled = False
|
||||||
|
prompt_widget.focus()
|
||||||
|
|
||||||
async def _render_history(self, thread_id_value: str) -> None:
|
async def _render_history(self, thread_id_value: str) -> None:
|
||||||
"""Render conversation history from a saved thread.
|
"""Render conversation history from a saved thread.
|
||||||
@@ -2970,13 +3288,13 @@ def run_textual_interactive(
|
|||||||
|
|
||||||
def _do_exit(self) -> None:
|
def _do_exit(self) -> None:
|
||||||
"""Clean up channels, unregister callbacks, and exit."""
|
"""Clean up channels, unregister callbacks, and exit."""
|
||||||
from ..middleware.model_fallback import set_ui_emit
|
self._exiting = True
|
||||||
|
self._channel_start_stop.set()
|
||||||
set_ui_emit(None)
|
event_sink.set_fallback_display(None)
|
||||||
if self._channel_timer is not None:
|
if self._channel_timer is not None:
|
||||||
self._channel_timer.stop()
|
self._channel_timer.stop()
|
||||||
self._channel_timer = None
|
self._channel_timer = None
|
||||||
self._started_channel_types.clear()
|
self._channel_start_results.clear()
|
||||||
if _channels_is_running():
|
if _channels_is_running():
|
||||||
try:
|
try:
|
||||||
_channels_stop(runtime=self._channel_runtime)
|
_channels_stop(runtime=self._channel_runtime)
|
||||||
@@ -2992,6 +3310,9 @@ def run_textual_interactive(
|
|||||||
self._queued_messages.clear()
|
self._queued_messages.clear()
|
||||||
self._render_queue_indicator()
|
self._render_queue_indicator()
|
||||||
if self._run_task is not None and not self._run_task.done():
|
if self._run_task is not None and not self._run_task.done():
|
||||||
|
from ..stream.display import request_stream_cancel
|
||||||
|
|
||||||
|
request_stream_cancel()
|
||||||
self._run_task.cancel()
|
self._run_task.cancel()
|
||||||
else:
|
else:
|
||||||
# Edge case: busy but no task — force reset
|
# Edge case: busy but no task — force reset
|
||||||
@@ -3105,11 +3426,11 @@ def run_textual_interactive(
|
|||||||
def _render_welcome(self) -> None:
|
def _render_welcome(self) -> None:
|
||||||
channels_info: list[tuple[str, bool, str]] | None = None
|
channels_info: list[tuple[str, bool, str]] | None = None
|
||||||
try:
|
try:
|
||||||
running = _channels_running_list()
|
current = get_channel_startup_results()
|
||||||
started = self._started_channel_types
|
if current:
|
||||||
if running or started:
|
self._channel_start_results = current
|
||||||
all_types = list(dict.fromkeys(running + started))
|
if self._channel_start_results:
|
||||||
channels_info = [(ct, True, "connected (bus)") for ct in all_types]
|
channels_info = self._channel_start_results
|
||||||
else:
|
else:
|
||||||
from ..config import load_config
|
from ..config import load_config
|
||||||
|
|
||||||
@@ -3359,6 +3680,18 @@ def run_textual_interactive(
|
|||||||
finally:
|
finally:
|
||||||
from .resume_hint import print_resume_hint
|
from .resume_hint import print_resume_hint
|
||||||
|
|
||||||
|
try:
|
||||||
|
from ..middleware.code_interpreter import (
|
||||||
|
aclose_code_interpreters,
|
||||||
|
)
|
||||||
|
|
||||||
|
await aclose_code_interpreters()
|
||||||
|
except Exception:
|
||||||
|
_channel_logger.debug(
|
||||||
|
"code interpreter cleanup failed",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
|
||||||
# Best-effort resume hint — guarded so failures here (e.g.
|
# Best-effort resume hint — guarded so failures here (e.g.
|
||||||
# DB teardown race during abnormal shutdown) cannot shadow
|
# DB teardown race during abnormal shutdown) cannot shadow
|
||||||
# the original run_async traceback.
|
# the original run_async traceback.
|
||||||
@@ -3378,12 +3711,4 @@ def run_textual_interactive(
|
|||||||
except Exception:
|
except Exception:
|
||||||
_channel_logger.debug("print_resume_hint failed", exc_info=True)
|
_channel_logger.debug("print_resume_hint failed", exc_info=True)
|
||||||
|
|
||||||
import nest_asyncio # type: ignore[import-untyped]
|
asyncio.run(_amain())
|
||||||
|
|
||||||
nest_asyncio.apply()
|
|
||||||
try:
|
|
||||||
loop = asyncio.get_event_loop()
|
|
||||||
except RuntimeError:
|
|
||||||
loop = asyncio.new_event_loop()
|
|
||||||
asyncio.set_event_loop(loop)
|
|
||||||
loop.run_until_complete(_amain())
|
|
||||||
|
|||||||
@@ -2,14 +2,20 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from typing import Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from ..gateway import GraphGateway
|
from ..gateway import GraphGateway
|
||||||
|
from ..runtime import AsyncRuntimeError
|
||||||
from ..stream.console import console
|
from ..stream.console import console
|
||||||
from .tui_backends import RichStreamingBackend, StreamingTUIBackend
|
from .tui_backends import RichStreamingBackend, StreamingTUIBackend
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from ..runtime import AsyncRuntime
|
||||||
|
|
||||||
DEFAULT_UI_BACKEND = "cli"
|
DEFAULT_UI_BACKEND = "cli"
|
||||||
|
STREAM_CANCEL_SETTLE_TIMEOUT = 5.0
|
||||||
# "webui" launches the browser front-end instead of an in-terminal UI; it is
|
# "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
|
# intercepted earlier (cli/commands.py:_main_callback) and never reaches the
|
||||||
# streaming backends, but is listed here so normalize/resolve preserve it
|
# streaming backends, but is listed here so normalize/resolve preserve it
|
||||||
@@ -18,6 +24,41 @@ SUPPORTED_UI_BACKENDS = ("cli", "tui", "webui")
|
|||||||
_LEGACY_BACKEND_MAP = {"textual": "tui", "rich": "cli"}
|
_LEGACY_BACKEND_MAP = {"textual": "tui", "rich": "cli"}
|
||||||
|
|
||||||
|
|
||||||
|
class StreamCancellationTimeout(RuntimeError):
|
||||||
|
"""A blocking renderer did not settle after its turn was cancelled."""
|
||||||
|
|
||||||
|
|
||||||
|
def _consume_late_worker_result(worker: asyncio.Task[Any]) -> None:
|
||||||
|
"""Retrieve a detached worker result so eventual failure is not unhandled."""
|
||||||
|
try:
|
||||||
|
worker.exception()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
async def settle_cancelled_worker(
|
||||||
|
worker: asyncio.Task[Any],
|
||||||
|
*,
|
||||||
|
on_cancel: Callable[[], Any],
|
||||||
|
) -> Any:
|
||||||
|
"""Request cooperative cancellation and wait a bounded time for settlement."""
|
||||||
|
on_cancel()
|
||||||
|
done, _ = await asyncio.wait(
|
||||||
|
{worker},
|
||||||
|
timeout=STREAM_CANCEL_SETTLE_TIMEOUT,
|
||||||
|
)
|
||||||
|
if not done:
|
||||||
|
worker.add_done_callback(_consume_late_worker_result)
|
||||||
|
raise StreamCancellationTimeout(
|
||||||
|
"The active turn did not stop within "
|
||||||
|
f"{STREAM_CANCEL_SETTLE_TIMEOUT:g} seconds after cancellation."
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
return worker.result()
|
||||||
|
except Exception:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
def normalize_ui_backend(value: str | None) -> str:
|
def normalize_ui_backend(value: str | None) -> str:
|
||||||
"""Normalize user-provided backend name with a safe default."""
|
"""Normalize user-provided backend name with a safe default."""
|
||||||
if not value:
|
if not value:
|
||||||
@@ -77,10 +118,12 @@ def run_streaming(
|
|||||||
on_stream_event: Callable[[str, Any], Any] | None = None,
|
on_stream_event: Callable[[str, Any], Any] | None = None,
|
||||||
status_footer_builder: Callable[[], Any] | None = None,
|
status_footer_builder: Callable[[], Any] | None = None,
|
||||||
metadata: dict | None = None,
|
metadata: dict | None = None,
|
||||||
|
configurable_extra: dict[str, Any] | None = None,
|
||||||
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
||||||
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
||||||
cancel_scope: str | None = None,
|
cancel_scope: str | None = None,
|
||||||
gateway: GraphGateway,
|
gateway: GraphGateway,
|
||||||
|
runtime: AsyncRuntime | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Run streaming with the selected backend."""
|
"""Run streaming with the selected backend."""
|
||||||
backend = get_backend(ui_backend, warn_fallback=True)
|
backend = get_backend(ui_backend, warn_fallback=True)
|
||||||
@@ -97,11 +140,15 @@ def run_streaming(
|
|||||||
on_stream_event=on_stream_event,
|
on_stream_event=on_stream_event,
|
||||||
status_footer_builder=status_footer_builder,
|
status_footer_builder=status_footer_builder,
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
|
configurable_extra=configurable_extra,
|
||||||
hitl_prompt_fn=hitl_prompt_fn,
|
hitl_prompt_fn=hitl_prompt_fn,
|
||||||
ask_user_prompt_fn=ask_user_prompt_fn,
|
ask_user_prompt_fn=ask_user_prompt_fn,
|
||||||
cancel_scope=cancel_scope,
|
cancel_scope=cancel_scope,
|
||||||
gateway=gateway,
|
gateway=gateway,
|
||||||
|
runtime=runtime,
|
||||||
)
|
)
|
||||||
|
except AsyncRuntimeError:
|
||||||
|
raise
|
||||||
except RuntimeError:
|
except RuntimeError:
|
||||||
requested = normalize_ui_backend(ui_backend)
|
requested = normalize_ui_backend(ui_backend)
|
||||||
if requested == "tui":
|
if requested == "tui":
|
||||||
@@ -120,9 +167,45 @@ def run_streaming(
|
|||||||
on_stream_event=on_stream_event,
|
on_stream_event=on_stream_event,
|
||||||
status_footer_builder=status_footer_builder,
|
status_footer_builder=status_footer_builder,
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
|
configurable_extra=configurable_extra,
|
||||||
hitl_prompt_fn=hitl_prompt_fn,
|
hitl_prompt_fn=hitl_prompt_fn,
|
||||||
ask_user_prompt_fn=ask_user_prompt_fn,
|
ask_user_prompt_fn=ask_user_prompt_fn,
|
||||||
cancel_scope=cancel_scope,
|
cancel_scope=cancel_scope,
|
||||||
gateway=gateway,
|
gateway=gateway,
|
||||||
|
runtime=runtime,
|
||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
async def run_streaming_async(
|
||||||
|
*,
|
||||||
|
recover_on_cancel: bool = False,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> str:
|
||||||
|
"""Run the synchronous Rich renderer without blocking a frontend loop.
|
||||||
|
|
||||||
|
Cancellation requests the matching stream scope and gives the worker a
|
||||||
|
bounded interval to unwind. Foreground interactive turns may opt into
|
||||||
|
recovering the frontend task after cleanup so Ctrl+C returns to the prompt.
|
||||||
|
"""
|
||||||
|
from ..stream.display import request_stream_cancel
|
||||||
|
|
||||||
|
worker = asyncio.create_task(asyncio.to_thread(run_streaming, **kwargs))
|
||||||
|
try:
|
||||||
|
return await asyncio.shield(worker)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
try:
|
||||||
|
response = await settle_cancelled_worker(
|
||||||
|
worker,
|
||||||
|
on_cancel=lambda: request_stream_cancel(kwargs.get("cancel_scope")),
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
from ..middleware.code_interpreter import aclose_code_interpreters
|
||||||
|
|
||||||
|
await aclose_code_interpreters()
|
||||||
|
if recover_on_cancel:
|
||||||
|
current = asyncio.current_task()
|
||||||
|
if current is not None and current.uncancel() > 0:
|
||||||
|
raise
|
||||||
|
return response
|
||||||
|
raise
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from .compact_summary_widget import CompactSummaryWidget
|
|||||||
from .compacting_widget import CompactingWidget
|
from .compacting_widget import CompactingWidget
|
||||||
from .loading_widget import LoadingWidget
|
from .loading_widget import LoadingWidget
|
||||||
from .mcp_loader_widget import MCPLoaderWidget
|
from .mcp_loader_widget import MCPLoaderWidget
|
||||||
|
from .panel_widget import PanelWidget
|
||||||
from .subagent_widget import SubAgentWidget
|
from .subagent_widget import SubAgentWidget
|
||||||
from .summarization_widget import SummarizationWidget
|
from .summarization_widget import SummarizationWidget
|
||||||
from .system_message import SystemMessage
|
from .system_message import SystemMessage
|
||||||
@@ -25,6 +26,7 @@ __all__ = [
|
|||||||
"CompactingWidget",
|
"CompactingWidget",
|
||||||
"LoadingWidget",
|
"LoadingWidget",
|
||||||
"MCPLoaderWidget",
|
"MCPLoaderWidget",
|
||||||
|
"PanelWidget",
|
||||||
"SubAgentWidget",
|
"SubAgentWidget",
|
||||||
"SummarizationWidget",
|
"SummarizationWidget",
|
||||||
"SystemMessage",
|
"SystemMessage",
|
||||||
|
|||||||
@@ -101,6 +101,15 @@ class ApprovalWidget(Widget):
|
|||||||
self._selected = 0
|
self._selected = 0
|
||||||
self._option_widgets: list[Static] = []
|
self._option_widgets: list[Static] = []
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_command(args: dict) -> str:
|
||||||
|
"""Pull the display-worthy target out of a tool's args dict.
|
||||||
|
|
||||||
|
Checks `command`/`path` first, then deepagents 0.7.0's `delete`
|
||||||
|
tool key `file_path` — without it, `delete` shows no target.
|
||||||
|
"""
|
||||||
|
return args.get("command", args.get("path", args.get("file_path", "")))
|
||||||
|
|
||||||
def compose(self) -> ComposeResult:
|
def compose(self) -> ComposeResult:
|
||||||
self._option_widgets = []
|
self._option_widgets = []
|
||||||
count = len(self._action_requests)
|
count = len(self._action_requests)
|
||||||
@@ -115,10 +124,7 @@ class ApprovalWidget(Widget):
|
|||||||
for req in self._action_requests:
|
for req in self._action_requests:
|
||||||
name = req.get("name", "")
|
name = req.get("name", "")
|
||||||
args = req.get("args", {})
|
args = req.get("args", {})
|
||||||
if isinstance(args, dict):
|
command = self._extract_command(args) if isinstance(args, dict) else ""
|
||||||
command = args.get("command", args.get("path", ""))
|
|
||||||
else:
|
|
||||||
command = ""
|
|
||||||
if command:
|
if command:
|
||||||
cmd_str = str(command)
|
cmd_str = str(command)
|
||||||
if len(cmd_str) > _COMMAND_TRUNCATE_LENGTH:
|
if len(cmd_str) > _COMMAND_TRUNCATE_LENGTH:
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Inline MCP server browser widget for /install-mcp in TUI.
|
"""Inline MCP server browser widget for /install-mcp in TUI.
|
||||||
|
|
||||||
Two-phase keyboard-driven widget (mirrors SkillBrowserWidget):
|
Two-phase keyboard-driven widget built on the shared picker engine
|
||||||
|
(``picker_base.TagCheckboxBrowserBase``):
|
||||||
Phase 1 — tag picker (arrow keys + Enter to select, or Esc for all)
|
Phase 1 — tag picker (arrow keys + Enter to select, or Esc for all)
|
||||||
Phase 2 — server checkbox (arrow keys to navigate, Space to toggle, Enter to confirm)
|
Phase 2 — server checkbox (arrow keys to navigate, Space to toggle, Enter to confirm)
|
||||||
|
|
||||||
@@ -12,73 +13,20 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import TYPE_CHECKING, Any, ClassVar
|
from typing import TYPE_CHECKING, Any, ClassVar
|
||||||
|
|
||||||
from rich.text import Text
|
|
||||||
from textual.binding import Binding, BindingType
|
|
||||||
from textual.containers import Container
|
|
||||||
from textual.message import Message
|
from textual.message import Message
|
||||||
from textual.widget import Widget
|
|
||||||
from textual.widgets import Static
|
from .picker_base import TagCheckboxBrowserBase
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from textual import events
|
|
||||||
from textual.app import ComposeResult
|
|
||||||
|
|
||||||
from ...mcp.registry import MCPServerEntry
|
from ...mcp.registry import MCPServerEntry
|
||||||
|
|
||||||
|
|
||||||
class MCPBrowserWidget(Widget):
|
class MCPBrowserWidget(TagCheckboxBrowserBase):
|
||||||
"""Inline MCP server browser — mounts in chat, keyboard-driven.
|
"""Inline MCP server browser — mounts in chat, keyboard-driven."""
|
||||||
|
|
||||||
Phase 1: Tag picker (select a tag filter or "All").
|
_INSTALLED_SUFFIX: ClassVar[str] = " (configured)"
|
||||||
Phase 2: Server checkbox (toggle servers, confirm to install).
|
_PHASE2_TITLE: ClassVar[str] = "Select MCP servers to install"
|
||||||
"""
|
_PHASE2_CONFIRM_LABEL: ClassVar[str] = "install"
|
||||||
|
|
||||||
can_focus = True
|
|
||||||
can_focus_children = False
|
|
||||||
|
|
||||||
DEFAULT_CSS = """
|
|
||||||
MCPBrowserWidget {
|
|
||||||
height: auto;
|
|
||||||
max-height: 30;
|
|
||||||
margin: 1 0;
|
|
||||||
padding: 0 1;
|
|
||||||
background: $surface;
|
|
||||||
border: solid $primary;
|
|
||||||
}
|
|
||||||
MCPBrowserWidget .browser-title {
|
|
||||||
height: 1;
|
|
||||||
text-style: bold;
|
|
||||||
color: $primary;
|
|
||||||
}
|
|
||||||
MCPBrowserWidget .browser-rows {
|
|
||||||
height: auto;
|
|
||||||
max-height: 20;
|
|
||||||
overflow-y: auto;
|
|
||||||
}
|
|
||||||
MCPBrowserWidget .browser-row {
|
|
||||||
height: 1;
|
|
||||||
padding: 0 1;
|
|
||||||
}
|
|
||||||
MCPBrowserWidget .browser-row-selected {
|
|
||||||
background: $primary;
|
|
||||||
text-style: bold;
|
|
||||||
}
|
|
||||||
MCPBrowserWidget .browser-help {
|
|
||||||
height: 1;
|
|
||||||
color: $text-muted;
|
|
||||||
text-style: italic;
|
|
||||||
}
|
|
||||||
"""
|
|
||||||
|
|
||||||
BINDINGS: ClassVar[list[BindingType]] = [
|
|
||||||
Binding("up", "move_up", "Up", show=False),
|
|
||||||
Binding("k", "move_up", "Up", show=False),
|
|
||||||
Binding("down", "move_down", "Down", show=False),
|
|
||||||
Binding("j", "move_down", "Down", show=False),
|
|
||||||
Binding("enter", "confirm", "Confirm", show=False),
|
|
||||||
Binding("space", "toggle", "Toggle", show=False),
|
|
||||||
Binding("escape", "cancel", "Cancel", show=False),
|
|
||||||
]
|
|
||||||
|
|
||||||
class Confirmed(Message):
|
class Confirmed(Message):
|
||||||
"""Posted when user confirms server selection."""
|
"""Posted when user confirms server selection."""
|
||||||
@@ -90,242 +38,17 @@ class MCPBrowserWidget(Widget):
|
|||||||
class Cancelled(Message):
|
class Cancelled(Message):
|
||||||
"""Posted when user cancels."""
|
"""Posted when user cancels."""
|
||||||
|
|
||||||
def __init__(
|
def _item_name(self, item: Any) -> str:
|
||||||
self,
|
return item.name
|
||||||
servers: list[MCPServerEntry],
|
|
||||||
installed_names: set[str],
|
|
||||||
*,
|
|
||||||
pre_filter_tag: str = "",
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> None:
|
|
||||||
super().__init__(**kwargs)
|
|
||||||
self._servers = servers
|
|
||||||
self._installed_names = installed_names
|
|
||||||
self._pre_filter_tag = pre_filter_tag.lower()
|
|
||||||
self._selected = 0
|
|
||||||
self._row_widgets: list[Static] = []
|
|
||||||
self._title_widget: Static | None = None
|
|
||||||
self._help_widget: Static | None = None
|
|
||||||
|
|
||||||
# Phase 1: tag picker
|
def _item_tags(self, item: Any) -> list[str]:
|
||||||
# Phase 2: server checkbox
|
return item.tags
|
||||||
self._phase: int = 1
|
|
||||||
self._tag_items: list[tuple[str, int]] = []
|
|
||||||
self._server_items: list[MCPServerEntry] = []
|
|
||||||
self._checked: set[int] = set()
|
|
||||||
|
|
||||||
# Build tag list
|
def _item_desc(self, item: Any) -> str:
|
||||||
from collections import Counter
|
return item.description or item.label
|
||||||
|
|
||||||
tag_counter: Counter[str] = Counter()
|
def _post_confirmed(self, items: list[Any]) -> None:
|
||||||
for s in self._servers:
|
self.post_message(self.Confirmed(items))
|
||||||
for t in s.tags:
|
|
||||||
tag_counter[t.lower()] += 1
|
|
||||||
sorted_tags = sorted(tag_counter.items(), key=lambda x: (-x[1], x[0]))
|
|
||||||
self._tag_items = [("all", len(self._servers)), *sorted_tags]
|
|
||||||
|
|
||||||
# If pre-filtered, skip to phase 2
|
def _post_cancelled(self) -> None:
|
||||||
if self._pre_filter_tag:
|
self.post_message(self.Cancelled())
|
||||||
self._server_items = [
|
|
||||||
s
|
|
||||||
for s in self._servers
|
|
||||||
if self._pre_filter_tag in [t.lower() for t in s.tags]
|
|
||||||
]
|
|
||||||
if self._server_items:
|
|
||||||
self._phase = 2
|
|
||||||
else:
|
|
||||||
self._pre_filter_tag = ""
|
|
||||||
|
|
||||||
def compose(self) -> ComposeResult:
|
|
||||||
self._title_widget = Static("", classes="browser-title")
|
|
||||||
yield self._title_widget
|
|
||||||
with Container(classes="browser-rows"):
|
|
||||||
max_rows = max(len(self._tag_items), len(self._servers))
|
|
||||||
for _ in range(max_rows):
|
|
||||||
widget = Static("", classes="browser-row")
|
|
||||||
self._row_widgets.append(widget)
|
|
||||||
yield widget
|
|
||||||
self._help_widget = Static("", classes="browser-help")
|
|
||||||
yield self._help_widget
|
|
||||||
|
|
||||||
def on_mount(self) -> None:
|
|
||||||
self.call_after_refresh(self._update_display)
|
|
||||||
self.call_later(self.focus)
|
|
||||||
|
|
||||||
def _update_display(self) -> None:
|
|
||||||
if self._phase == 1:
|
|
||||||
self._render_tag_picker()
|
|
||||||
else:
|
|
||||||
self._render_server_checkbox()
|
|
||||||
|
|
||||||
def _render_tag_picker(self) -> None:
|
|
||||||
if self._title_widget:
|
|
||||||
self._title_widget.update("Filter by tag:")
|
|
||||||
if self._help_widget:
|
|
||||||
self._help_widget.update(
|
|
||||||
"\u2191/\u2193 navigate \u00b7 Enter select \u00b7 Esc cancel"
|
|
||||||
)
|
|
||||||
|
|
||||||
for i, widget in enumerate(self._row_widgets):
|
|
||||||
if i < len(self._tag_items):
|
|
||||||
tag, count = self._tag_items[i]
|
|
||||||
is_selected = i == self._selected
|
|
||||||
text = Text()
|
|
||||||
cursor = "\u25b8 " if is_selected else " "
|
|
||||||
text.append(cursor, style="bold cyan" if is_selected else "dim")
|
|
||||||
label = f"{tag} ({count})"
|
|
||||||
text.append(label, style="bold" if is_selected else "")
|
|
||||||
widget.update(text)
|
|
||||||
widget.display = True
|
|
||||||
widget.remove_class("browser-row-selected")
|
|
||||||
if is_selected:
|
|
||||||
widget.add_class("browser-row-selected")
|
|
||||||
widget.scroll_visible()
|
|
||||||
else:
|
|
||||||
widget.update("")
|
|
||||||
widget.display = False
|
|
||||||
|
|
||||||
def _row_content_width(self) -> int:
|
|
||||||
try:
|
|
||||||
w = self.size.width
|
|
||||||
if w > 0:
|
|
||||||
return w - 6
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
try:
|
|
||||||
return self.app.size.width - 10
|
|
||||||
except Exception:
|
|
||||||
return 100
|
|
||||||
|
|
||||||
def _truncate(self, desc: str, name: str, *, suffix: str = "") -> str:
|
|
||||||
overhead = 2 + 2 + len(name) + 3 + len(suffix)
|
|
||||||
max_len = max(20, self._row_content_width() - overhead)
|
|
||||||
if len(desc) <= max_len:
|
|
||||||
return desc
|
|
||||||
return desc[: max_len - 1] + "\u2026"
|
|
||||||
|
|
||||||
def _render_server_checkbox(self) -> None:
|
|
||||||
n_checked = len(
|
|
||||||
[
|
|
||||||
i
|
|
||||||
for i in self._checked
|
|
||||||
if self._server_items[i].name not in self._installed_names
|
|
||||||
]
|
|
||||||
)
|
|
||||||
if self._title_widget:
|
|
||||||
self._title_widget.update(
|
|
||||||
f"Select MCP servers to install ({n_checked} selected):"
|
|
||||||
)
|
|
||||||
if self._help_widget:
|
|
||||||
self._help_widget.update(
|
|
||||||
"\u2191/\u2193 navigate \u00b7 Space toggle \u00b7 Enter install \u00b7 Esc cancel"
|
|
||||||
)
|
|
||||||
|
|
||||||
for i, widget in enumerate(self._row_widgets):
|
|
||||||
if i < len(self._server_items):
|
|
||||||
entry = self._server_items[i]
|
|
||||||
is_selected = i == self._selected
|
|
||||||
is_installed = entry.name in self._installed_names
|
|
||||||
is_checked = i in self._checked
|
|
||||||
|
|
||||||
text = Text()
|
|
||||||
cursor = "\u25b8 " if is_selected else " "
|
|
||||||
text.append(cursor, style="bold cyan" if is_selected else "dim")
|
|
||||||
|
|
||||||
desc = entry.description or entry.label
|
|
||||||
|
|
||||||
if is_installed:
|
|
||||||
suffix = " (configured)"
|
|
||||||
desc = self._truncate(desc, entry.name, suffix=suffix)
|
|
||||||
text.append("\u2713 ", style="green")
|
|
||||||
text.append(entry.name, style="green dim")
|
|
||||||
text.append(f" \u2014 {desc}", style="dim")
|
|
||||||
text.append(suffix, style="dim italic")
|
|
||||||
elif is_checked:
|
|
||||||
desc = self._truncate(desc, entry.name)
|
|
||||||
text.append("\u25cf ", style="green bold")
|
|
||||||
text.append(entry.name, style="bold")
|
|
||||||
text.append(f" \u2014 {desc}", style="")
|
|
||||||
else:
|
|
||||||
desc = self._truncate(desc, entry.name)
|
|
||||||
text.append("\u25cb ", style="dim")
|
|
||||||
text.append(entry.name, style="bold" if is_selected else "")
|
|
||||||
text.append(f" \u2014 {desc}", style="dim")
|
|
||||||
|
|
||||||
widget.update(text)
|
|
||||||
widget.display = True
|
|
||||||
widget.remove_class("browser-row-selected")
|
|
||||||
if is_selected:
|
|
||||||
widget.add_class("browser-row-selected")
|
|
||||||
widget.scroll_visible()
|
|
||||||
else:
|
|
||||||
widget.update("")
|
|
||||||
widget.display = False
|
|
||||||
|
|
||||||
def _current_items_count(self) -> int:
|
|
||||||
if self._phase == 1:
|
|
||||||
return len(self._tag_items)
|
|
||||||
return len(self._server_items)
|
|
||||||
|
|
||||||
def action_move_up(self) -> None:
|
|
||||||
n = self._current_items_count()
|
|
||||||
if not n:
|
|
||||||
return
|
|
||||||
self._selected = (self._selected - 1) % n
|
|
||||||
self._update_display()
|
|
||||||
|
|
||||||
def action_move_down(self) -> None:
|
|
||||||
n = self._current_items_count()
|
|
||||||
if not n:
|
|
||||||
return
|
|
||||||
self._selected = (self._selected + 1) % n
|
|
||||||
self._update_display()
|
|
||||||
|
|
||||||
def action_toggle(self) -> None:
|
|
||||||
if self._phase != 2:
|
|
||||||
return
|
|
||||||
if not self._server_items:
|
|
||||||
return
|
|
||||||
entry = self._server_items[self._selected]
|
|
||||||
if entry.name in self._installed_names:
|
|
||||||
return
|
|
||||||
if self._selected in self._checked:
|
|
||||||
self._checked.discard(self._selected)
|
|
||||||
else:
|
|
||||||
self._checked.add(self._selected)
|
|
||||||
self._update_display()
|
|
||||||
|
|
||||||
def action_confirm(self) -> None:
|
|
||||||
if self._phase == 1:
|
|
||||||
if not self._tag_items:
|
|
||||||
return
|
|
||||||
tag, _ = self._tag_items[self._selected]
|
|
||||||
if tag == "all":
|
|
||||||
self._server_items = list(self._servers)
|
|
||||||
else:
|
|
||||||
self._server_items = [
|
|
||||||
s for s in self._servers if tag in [t.lower() for t in s.tags]
|
|
||||||
]
|
|
||||||
self._phase = 2
|
|
||||||
self._selected = 0
|
|
||||||
self._checked = set()
|
|
||||||
self._update_display()
|
|
||||||
else:
|
|
||||||
entries = [
|
|
||||||
self._server_items[i]
|
|
||||||
for i in sorted(self._checked)
|
|
||||||
if self._server_items[i].name not in self._installed_names
|
|
||||||
]
|
|
||||||
self.post_message(self.Confirmed(entries))
|
|
||||||
|
|
||||||
def action_cancel(self) -> None:
|
|
||||||
if self._phase == 2 and not self._pre_filter_tag:
|
|
||||||
self._phase = 1
|
|
||||||
self._selected = 0
|
|
||||||
self._checked = set()
|
|
||||||
self._update_display()
|
|
||||||
else:
|
|
||||||
self.post_message(self.Cancelled())
|
|
||||||
|
|
||||||
def on_blur(self, event: events.Blur) -> None:
|
|
||||||
self.call_after_refresh(self.focus)
|
|
||||||
|
|||||||
@@ -12,9 +12,10 @@ from rich.text import Text
|
|||||||
from textual.binding import Binding, BindingType
|
from textual.binding import Binding, BindingType
|
||||||
from textual.containers import Container
|
from textual.containers import Container
|
||||||
from textual.message import Message
|
from textual.message import Message
|
||||||
from textual.widget import Widget
|
|
||||||
from textual.widgets import Input, Static
|
from textual.widgets import Input, Static
|
||||||
|
|
||||||
|
from .picker_base import PickerWidgetBase, first_selectable_index, move_selection
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from textual import events
|
from textual import events
|
||||||
from textual.app import ComposeResult
|
from textual.app import ComposeResult
|
||||||
@@ -81,14 +82,13 @@ def _build_items(
|
|||||||
return items
|
return items
|
||||||
|
|
||||||
|
|
||||||
class ModelPickerWidget(Widget):
|
class ModelPickerWidget(PickerWidgetBase):
|
||||||
"""Inline model picker -- mounts in chat, keyboard-driven.
|
"""Inline model picker -- mounts in chat, keyboard-driven.
|
||||||
|
|
||||||
Posts ``Picked(name, provider)`` on Enter, ``Cancelled()`` on Esc.
|
Posts ``Picked(name, provider)`` on Enter, ``Cancelled()`` on Esc.
|
||||||
Type to filter models.
|
Type to filter models.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
can_focus = True
|
|
||||||
# Required so the Custom Ollama ``Input`` child can hold focus when the
|
# Required so the Custom Ollama ``Input`` child can hold focus when the
|
||||||
# user is typing a model name.
|
# user is typing a model name.
|
||||||
can_focus_children = True
|
can_focus_children = True
|
||||||
@@ -188,22 +188,19 @@ class ModelPickerWidget(Widget):
|
|||||||
self._mode: Literal["list", "input"] = "list"
|
self._mode: Literal["list", "input"] = "list"
|
||||||
self._custom_input: Input | None = None
|
self._custom_input: Input | None = None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_model(item: dict) -> bool:
|
||||||
|
return item["type"] == "model"
|
||||||
|
|
||||||
def _first_model_index(self) -> int:
|
def _first_model_index(self) -> int:
|
||||||
for i, item in enumerate(self._items):
|
return first_selectable_index(self._items, self._is_model)
|
||||||
if item["type"] == "model":
|
|
||||||
return i
|
|
||||||
return 0
|
|
||||||
|
|
||||||
def _move(self, direction: int) -> None:
|
def _move(self, direction: int) -> None:
|
||||||
if not self._items:
|
if not self._items:
|
||||||
return
|
return
|
||||||
i = (self._selected + direction) % len(self._items)
|
new = move_selection(self._items, self._selected, direction, self._is_model)
|
||||||
steps = 0
|
if self._is_model(self._items[new]):
|
||||||
while self._items[i]["type"] != "model" and steps < len(self._items):
|
self._selected = new
|
||||||
i = (i + direction) % len(self._items)
|
|
||||||
steps += 1
|
|
||||||
if self._items[i]["type"] == "model":
|
|
||||||
self._selected = i
|
|
||||||
self._update_rows()
|
self._update_rows()
|
||||||
|
|
||||||
def _rebuild(self) -> None:
|
def _rebuild(self) -> None:
|
||||||
@@ -251,10 +248,9 @@ class ModelPickerWidget(Widget):
|
|||||||
classes="picker-help",
|
classes="picker-help",
|
||||||
)
|
)
|
||||||
|
|
||||||
def on_mount(self) -> None:
|
def _refresh_view(self) -> None:
|
||||||
self._update_rows()
|
self._update_rows()
|
||||||
self._update_filter()
|
self._update_filter()
|
||||||
self.call_later(self.focus)
|
|
||||||
|
|
||||||
def _update_filter(self) -> None:
|
def _update_filter(self) -> None:
|
||||||
if self._filter_widget is not None:
|
if self._filter_widget is not None:
|
||||||
@@ -273,8 +269,8 @@ class ModelPickerWidget(Widget):
|
|||||||
for i, (item, widget) in enumerate(
|
for i, (item, widget) in enumerate(
|
||||||
zip(self._items, self._row_widgets, strict=False)
|
zip(self._items, self._row_widgets, strict=False)
|
||||||
):
|
):
|
||||||
widget.remove_class("picker-row-selected")
|
|
||||||
if item["type"] == "header":
|
if item["type"] == "header":
|
||||||
|
widget.remove_class("picker-row-selected")
|
||||||
t = Text()
|
t = Text()
|
||||||
t.append("\u2500\u2500 ", style="bold cyan")
|
t.append("\u2500\u2500 ", style="bold cyan")
|
||||||
t.append(item["label"], style="bold cyan")
|
t.append(item["label"], style="bold cyan")
|
||||||
@@ -289,9 +285,7 @@ class ModelPickerWidget(Widget):
|
|||||||
t.append(" *", style="bold green")
|
t.append(" *", style="bold green")
|
||||||
t.append(f" ({item['provider']})", style="dim italic")
|
t.append(f" ({item['provider']})", style="dim italic")
|
||||||
widget.update(t)
|
widget.update(t)
|
||||||
if is_selected:
|
self.apply_row_highlight(widget, is_selected)
|
||||||
widget.add_class("picker-row-selected")
|
|
||||||
widget.scroll_visible()
|
|
||||||
|
|
||||||
def on_key(self, event: events.Key) -> None:
|
def on_key(self, event: events.Key) -> None:
|
||||||
# In input mode, the Input child owns printable keys + backspace.
|
# In input mode, the Input child owns printable keys + backspace.
|
||||||
@@ -350,11 +344,9 @@ class ModelPickerWidget(Widget):
|
|||||||
return
|
return
|
||||||
self.post_message(self.Cancelled())
|
self.post_message(self.Cancelled())
|
||||||
|
|
||||||
def on_blur(self, event: events.Blur) -> None:
|
def _should_refocus_on_blur(self) -> bool:
|
||||||
# When the Input child has focus we must NOT steal it back.
|
# When the Input child has focus we must NOT steal it back.
|
||||||
if self._mode == "input":
|
return self._mode != "input"
|
||||||
return
|
|
||||||
self.call_after_refresh(self.focus)
|
|
||||||
|
|
||||||
def on_input_submitted(self, event: Input.Submitted) -> None:
|
def on_input_submitted(self, event: Input.Submitted) -> None:
|
||||||
"""Safety net: Enter fired inside the Input widget rather than
|
"""Safety net: Enter fired inside the Input widget rather than
|
||||||
|
|||||||
@@ -0,0 +1,257 @@
|
|||||||
|
"""Panel widget — in-eval ``task()`` fan-out live view.
|
||||||
|
|
||||||
|
Groups all sub-agent dispatches from a single ``code_interpreter`` eval into
|
||||||
|
one bordered container, one row per dispatch. Each row shows the expert /
|
||||||
|
subagent type, the label, a running elapsed timer, and a status dot that
|
||||||
|
flips to ``ok``/``err`` on completion.
|
||||||
|
|
||||||
|
Sourced from the ``custom`` stream events emitted by ``langchain_quickjs``
|
||||||
|
(see ``.stream.emitter.panel_dispatch_start`` etc.). Keyed by ``eval_id``
|
||||||
|
so parallel dispatches from the same eval appear stacked; distinct evals
|
||||||
|
get distinct panels.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import time
|
||||||
|
|
||||||
|
from rich.text import Text
|
||||||
|
from textual.containers import Vertical
|
||||||
|
from textual.widget import Widget
|
||||||
|
from textual.widgets import Static
|
||||||
|
|
||||||
|
from ..status_bar import SPINNER_FRAMES
|
||||||
|
|
||||||
|
_ROW_LABEL_MAX_CHARS = 48
|
||||||
|
|
||||||
|
|
||||||
|
class _DispatchRow(Widget):
|
||||||
|
"""One row inside a PanelWidget — a single ``task()`` dispatch.
|
||||||
|
|
||||||
|
Subclasses ``Widget`` directly and overrides ``render()`` to build the
|
||||||
|
row's ``Text`` on demand. Earlier revisions subclassed ``Static`` (visual
|
||||||
|
stayed ``None`` past the first paint) and ``Vertical`` with an inner
|
||||||
|
``Static`` (container-layout race); rendering from ``render()`` is the
|
||||||
|
standard pattern for single-line widgets and dodges both issues by
|
||||||
|
letting Textual manage the visual lifecycle itself.
|
||||||
|
"""
|
||||||
|
|
||||||
|
DEFAULT_CSS = """
|
||||||
|
_DispatchRow {
|
||||||
|
height: 1;
|
||||||
|
width: 100%;
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, subagent_type: str, label: str) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self._subagent_type = subagent_type
|
||||||
|
self._label = label
|
||||||
|
self._started_at = time.monotonic()
|
||||||
|
self._status: str = "running" # "running" | "ok" | "err"
|
||||||
|
self._duration_ms: int | None = None
|
||||||
|
self._error: str = ""
|
||||||
|
self._frame = 0
|
||||||
|
|
||||||
|
def render(self) -> Text:
|
||||||
|
line = Text()
|
||||||
|
if self._status == "running":
|
||||||
|
line.append(f" {SPINNER_FRAMES[self._frame]} ", style="cyan")
|
||||||
|
elif self._status == "ok":
|
||||||
|
line.append(" \u2713 ", style="green")
|
||||||
|
else:
|
||||||
|
line.append(" \u2717 ", style="red")
|
||||||
|
line.append(f"{self._subagent_type} ", style="bold")
|
||||||
|
if self._label:
|
||||||
|
trimmed = self._label
|
||||||
|
if len(trimmed) > _ROW_LABEL_MAX_CHARS:
|
||||||
|
trimmed = trimmed[: _ROW_LABEL_MAX_CHARS - 1] + "\u2026"
|
||||||
|
line.append(f"\u2014 {trimmed} ", style="dim")
|
||||||
|
line.append(self._elapsed_display(), style="dim")
|
||||||
|
if self._status == "err" and self._error:
|
||||||
|
err = self._error.split("\n", 1)[0]
|
||||||
|
if len(err) > 60:
|
||||||
|
err = err[:59] + "\u2026"
|
||||||
|
line.append(f" {err}", style="red")
|
||||||
|
return line
|
||||||
|
|
||||||
|
def tick(self) -> None:
|
||||||
|
if self._status == "running":
|
||||||
|
self._frame = (self._frame + 1) % len(SPINNER_FRAMES)
|
||||||
|
self.refresh()
|
||||||
|
|
||||||
|
def complete(self, duration_ms: int) -> None:
|
||||||
|
self._status = "ok"
|
||||||
|
self._duration_ms = duration_ms
|
||||||
|
self.refresh()
|
||||||
|
|
||||||
|
def fail(self, duration_ms: int, error: str) -> None:
|
||||||
|
self._status = "err"
|
||||||
|
self._duration_ms = duration_ms
|
||||||
|
self._error = error
|
||||||
|
self.refresh()
|
||||||
|
|
||||||
|
def _elapsed_display(self) -> str:
|
||||||
|
if self._duration_ms is not None:
|
||||||
|
secs = self._duration_ms / 1000.0
|
||||||
|
else:
|
||||||
|
secs = time.monotonic() - self._started_at
|
||||||
|
return f"{secs:5.1f}s"
|
||||||
|
|
||||||
|
|
||||||
|
class PanelWidget(Vertical):
|
||||||
|
"""Container for one eval's fan-out — bordered box, one row per dispatch."""
|
||||||
|
|
||||||
|
DEFAULT_CSS = """
|
||||||
|
PanelWidget {
|
||||||
|
height: auto;
|
||||||
|
margin: 0 0;
|
||||||
|
}
|
||||||
|
PanelWidget .panel-header {
|
||||||
|
height: auto;
|
||||||
|
color: #22d3ee;
|
||||||
|
}
|
||||||
|
PanelWidget .panel-rows {
|
||||||
|
height: auto;
|
||||||
|
padding: 0 0 0 2;
|
||||||
|
}
|
||||||
|
PanelWidget .panel-footer {
|
||||||
|
height: auto;
|
||||||
|
color: #22d3ee;
|
||||||
|
}
|
||||||
|
PanelWidget.--completed .panel-header {
|
||||||
|
color: #4ade80;
|
||||||
|
}
|
||||||
|
PanelWidget.--completed .panel-footer {
|
||||||
|
color: #4ade80;
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, eval_id: str) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self._eval_id = eval_id
|
||||||
|
self._rows: dict[str, _DispatchRow] = {}
|
||||||
|
self._timer_handle = None
|
||||||
|
self._is_active = True
|
||||||
|
|
||||||
|
@property
|
||||||
|
def eval_id(self) -> str:
|
||||||
|
return self._eval_id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def dispatch_count(self) -> int:
|
||||||
|
return len(self._rows)
|
||||||
|
|
||||||
|
def compose(self):
|
||||||
|
yield Static("", classes="panel-header")
|
||||||
|
yield Vertical(classes="panel-rows")
|
||||||
|
yield Static("", classes="panel-footer")
|
||||||
|
|
||||||
|
def on_mount(self) -> None:
|
||||||
|
self._timer_handle = self.set_interval(0.1, self._tick)
|
||||||
|
self._render_header()
|
||||||
|
self._render_footer()
|
||||||
|
|
||||||
|
def _tick(self) -> None:
|
||||||
|
for row in self._rows.values():
|
||||||
|
row.tick()
|
||||||
|
self._render_header()
|
||||||
|
|
||||||
|
async def start_dispatch(
|
||||||
|
self, dispatch_id: str, subagent_type: str, label: str
|
||||||
|
) -> None:
|
||||||
|
if dispatch_id in self._rows:
|
||||||
|
return
|
||||||
|
# Re-arm if the panel already finalized: a Promise.allSettled retry
|
||||||
|
# of a failed subset dispatches under the same eval_id, so a new
|
||||||
|
# start after _maybe_finalize() has stopped the timer must undo
|
||||||
|
# the three effects of finalize (latch, class, timer) or the new
|
||||||
|
# row's spinner/elapsed stay frozen and future completions never
|
||||||
|
# refresh the header.
|
||||||
|
if not self._is_active:
|
||||||
|
self._is_active = True
|
||||||
|
self.remove_class("--completed")
|
||||||
|
self._timer_handle = self.set_interval(0.1, self._tick)
|
||||||
|
self._render_footer()
|
||||||
|
row = _DispatchRow(subagent_type, label)
|
||||||
|
rows_container = self.query_one(".panel-rows", Vertical)
|
||||||
|
await rows_container.mount(row)
|
||||||
|
self._rows[dispatch_id] = row
|
||||||
|
self._render_header()
|
||||||
|
|
||||||
|
def complete_dispatch(self, dispatch_id: str, duration_ms: int) -> None:
|
||||||
|
row = self._rows.get(dispatch_id)
|
||||||
|
if row is not None:
|
||||||
|
row.complete(duration_ms)
|
||||||
|
self._maybe_finalize()
|
||||||
|
|
||||||
|
def fail_dispatch(self, dispatch_id: str, duration_ms: int, error: str) -> None:
|
||||||
|
row = self._rows.get(dispatch_id)
|
||||||
|
if row is not None:
|
||||||
|
row.fail(duration_ms, error)
|
||||||
|
self._maybe_finalize()
|
||||||
|
|
||||||
|
def finalize_running(self, reason: str = "interrupted") -> None:
|
||||||
|
"""Fail all still-running rows with their measured elapsed time.
|
||||||
|
|
||||||
|
Called from the TUI turn-cleanup ``finally`` so the 100ms interval
|
||||||
|
timer stops when a turn is cancelled mid-dispatch — without this,
|
||||||
|
no reference remains to stop the timer once the outer scope exits.
|
||||||
|
"""
|
||||||
|
if not self._is_active:
|
||||||
|
return
|
||||||
|
now = time.monotonic()
|
||||||
|
for dispatch_id, row in list(self._rows.items()):
|
||||||
|
if row._status == "running":
|
||||||
|
elapsed_ms = int((now - row._started_at) * 1000)
|
||||||
|
self.fail_dispatch(dispatch_id, elapsed_ms, reason)
|
||||||
|
# Cover the zero-running-rows case: cancel-during-mount can leave
|
||||||
|
# the timer armed with either no rows registered (first dispatch)
|
||||||
|
# or every registered row already terminal (allSettled retry).
|
||||||
|
# The loop skips both, so call _maybe_finalize unconditionally —
|
||||||
|
# it is a no-op once _is_active has flipped.
|
||||||
|
self._maybe_finalize()
|
||||||
|
|
||||||
|
def _maybe_finalize(self) -> None:
|
||||||
|
if not self._is_active:
|
||||||
|
return
|
||||||
|
if all(row._status != "running" for row in self._rows.values()):
|
||||||
|
self._is_active = False
|
||||||
|
if self._timer_handle is not None:
|
||||||
|
self._timer_handle.stop()
|
||||||
|
self._timer_handle = None
|
||||||
|
self.add_class("--completed")
|
||||||
|
self._render_header()
|
||||||
|
self._render_footer()
|
||||||
|
|
||||||
|
def _summary_counts(self) -> tuple[int, int, int]:
|
||||||
|
running = ok = err = 0
|
||||||
|
for row in self._rows.values():
|
||||||
|
if row._status == "running":
|
||||||
|
running += 1
|
||||||
|
elif row._status == "ok":
|
||||||
|
ok += 1
|
||||||
|
else:
|
||||||
|
err += 1
|
||||||
|
return running, ok, err
|
||||||
|
|
||||||
|
def _render_header(self) -> None:
|
||||||
|
header = self.query_one(".panel-header", Static)
|
||||||
|
running, ok, err = self._summary_counts()
|
||||||
|
line = Text()
|
||||||
|
if self._is_active:
|
||||||
|
line.append("\u250c \u25b6 Expert panel ", style="bold cyan")
|
||||||
|
line.append(
|
||||||
|
f"({running} running, {ok} done, {err} failed)", style="dim cyan"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
line.append("\u2713 Expert panel ", style="bold green")
|
||||||
|
line.append(f"({ok} done, {err} failed)", style="dim green")
|
||||||
|
header.update(line)
|
||||||
|
|
||||||
|
def _render_footer(self) -> None:
|
||||||
|
footer = self.query_one(".panel-footer", Static)
|
||||||
|
if self._is_active:
|
||||||
|
footer.update(Text("\u2514 running...", style="dim cyan"))
|
||||||
|
else:
|
||||||
|
footer.update(Text(""))
|
||||||
@@ -0,0 +1,427 @@
|
|||||||
|
"""Shared engine for the TUI's keyboard-driven picker/browser widgets.
|
||||||
|
|
||||||
|
Every inline picker (model picker, thread picker, skill/MCP browsers)
|
||||||
|
follows the same pattern: a flat item list where some rows are
|
||||||
|
selectable, a wrapping highlight cursor, Enter/Esc terminal messages,
|
||||||
|
and focus trapped inside the widget until a decision is made. This
|
||||||
|
module owns that machinery so the widgets only provide their data
|
||||||
|
model and row rendering.
|
||||||
|
|
||||||
|
Subclassing contract: Textual dispatches same-named handlers at EVERY
|
||||||
|
level of the MRO, so subclasses must NOT define ``on_mount``/``on_blur``
|
||||||
|
— they implement the ``_refresh_view()`` hook (and override
|
||||||
|
``_should_refocus_on_blur()`` if focus may legitimately leave, e.g. a
|
||||||
|
child ``Input``). Message classes (``Picked``/``Confirmed``/
|
||||||
|
``Cancelled``) stay defined in each widget: their handler names
|
||||||
|
(``on_<widget>_<message>``) derive from the defining class.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Callable
|
||||||
|
from typing import TYPE_CHECKING, Any, ClassVar
|
||||||
|
|
||||||
|
from rich.text import Text
|
||||||
|
from textual.binding import Binding, BindingType
|
||||||
|
from textual.containers import Container
|
||||||
|
from textual.widget import Widget
|
||||||
|
from textual.widgets import Static
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from textual import events
|
||||||
|
from textual.app import ComposeResult
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Pure selection helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def first_selectable_index(
|
||||||
|
items: list[Any], is_selectable: Callable[[Any], bool]
|
||||||
|
) -> int:
|
||||||
|
"""Index of the first selectable item, or 0 when none qualifies."""
|
||||||
|
for i, item in enumerate(items):
|
||||||
|
if is_selectable(item):
|
||||||
|
return i
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def move_selection(
|
||||||
|
items: list[Any],
|
||||||
|
current: int,
|
||||||
|
direction: int,
|
||||||
|
is_selectable: Callable[[Any], bool],
|
||||||
|
) -> int:
|
||||||
|
"""Next selectable index from *current*, wrapping around the list.
|
||||||
|
|
||||||
|
Non-selectable rows (headers, separators) are skipped; when no
|
||||||
|
selectable row exists the *current* index is returned unchanged.
|
||||||
|
"""
|
||||||
|
if not items:
|
||||||
|
return current
|
||||||
|
i = (current + direction) % len(items)
|
||||||
|
steps = 0
|
||||||
|
while not is_selectable(items[i]) and steps < len(items):
|
||||||
|
i = (i + direction) % len(items)
|
||||||
|
steps += 1
|
||||||
|
return i if is_selectable(items[i]) else current
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Widget base
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class PickerWidgetBase(Widget):
|
||||||
|
"""Focus-trapped inline picker: mount-focus, blur-refocus, row
|
||||||
|
highlight bookkeeping and description truncation."""
|
||||||
|
|
||||||
|
can_focus = True
|
||||||
|
can_focus_children = False
|
||||||
|
|
||||||
|
def _refresh_view(self) -> None:
|
||||||
|
"""Render the current state into the row widgets."""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def _should_refocus_on_blur(self) -> bool:
|
||||||
|
"""Whether blur should snap focus back (focus trap)."""
|
||||||
|
return True
|
||||||
|
|
||||||
|
def on_mount(self) -> None:
|
||||||
|
# Deferred so self.size is populated for width-aware rendering.
|
||||||
|
self.call_after_refresh(self._refresh_view)
|
||||||
|
self.call_later(self.focus)
|
||||||
|
|
||||||
|
def on_blur(self, event: events.Blur) -> None:
|
||||||
|
if self._should_refocus_on_blur():
|
||||||
|
self.call_after_refresh(self.focus)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def apply_row_highlight(
|
||||||
|
widget: Static, selected: bool, css_class: str = "picker-row-selected"
|
||||||
|
) -> None:
|
||||||
|
"""Toggle the selected-row CSS class and keep the row in view."""
|
||||||
|
widget.remove_class(css_class)
|
||||||
|
if selected:
|
||||||
|
widget.add_class(css_class)
|
||||||
|
widget.scroll_visible()
|
||||||
|
|
||||||
|
def _row_content_width(self) -> int:
|
||||||
|
"""Usable character width for a row's text content (accounts for
|
||||||
|
widget border/padding; falls back to terminal width pre-layout)."""
|
||||||
|
try:
|
||||||
|
w = self.size.width
|
||||||
|
if w > 0:
|
||||||
|
# border (2) + widget padding (2) + row padding (2)
|
||||||
|
return w - 6
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
return self.app.size.width - 10
|
||||||
|
except Exception:
|
||||||
|
return 100
|
||||||
|
|
||||||
|
def _truncate(self, desc: str, name: str, *, suffix: str = "") -> str:
|
||||||
|
"""Truncate a description to fit the row, adding ellipsis."""
|
||||||
|
# cursor(2) + indicator(2) + name + " — "(3) + suffix
|
||||||
|
overhead = 2 + 2 + len(name) + 3 + len(suffix)
|
||||||
|
max_len = max(20, self._row_content_width() - overhead)
|
||||||
|
if len(desc) <= max_len:
|
||||||
|
return desc
|
||||||
|
return desc[: max_len - 1] + "…"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Two-phase tag-filter → checkbox browser
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class TagCheckboxBrowserBase(PickerWidgetBase):
|
||||||
|
"""Two-phase multi-select browser shared by the skill and MCP browsers.
|
||||||
|
|
||||||
|
Phase 1 — tag picker (Enter selects a tag filter, "all" included).
|
||||||
|
Phase 2 — checkbox list (Space toggles, Enter confirms, Esc goes back
|
||||||
|
to phase 1 unless the widget was constructed pre-filtered).
|
||||||
|
|
||||||
|
Subclasses provide the data adapters (``_item_name`` / ``_item_tags``
|
||||||
|
/ ``_item_desc``), the phase-2 texts, and ``_post_confirmed()``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
DEFAULT_CSS = """
|
||||||
|
TagCheckboxBrowserBase {
|
||||||
|
height: auto;
|
||||||
|
max-height: 30;
|
||||||
|
margin: 1 0;
|
||||||
|
padding: 0 1;
|
||||||
|
background: $surface;
|
||||||
|
border: solid $primary;
|
||||||
|
}
|
||||||
|
TagCheckboxBrowserBase .browser-title {
|
||||||
|
height: 1;
|
||||||
|
text-style: bold;
|
||||||
|
color: $primary;
|
||||||
|
}
|
||||||
|
TagCheckboxBrowserBase .browser-rows {
|
||||||
|
height: auto;
|
||||||
|
max-height: 20;
|
||||||
|
overflow-y: auto;
|
||||||
|
}
|
||||||
|
TagCheckboxBrowserBase .browser-row {
|
||||||
|
height: 1;
|
||||||
|
padding: 0 1;
|
||||||
|
}
|
||||||
|
TagCheckboxBrowserBase .browser-row-selected {
|
||||||
|
background: $primary;
|
||||||
|
text-style: bold;
|
||||||
|
}
|
||||||
|
TagCheckboxBrowserBase .browser-help {
|
||||||
|
height: 1;
|
||||||
|
color: $text-muted;
|
||||||
|
text-style: italic;
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
|
||||||
|
BINDINGS: ClassVar[list[BindingType]] = [
|
||||||
|
Binding("up", "move_up", "Up", show=False),
|
||||||
|
Binding("k", "move_up", "Up", show=False),
|
||||||
|
Binding("down", "move_down", "Down", show=False),
|
||||||
|
Binding("j", "move_down", "Down", show=False),
|
||||||
|
Binding("enter", "confirm", "Confirm", show=False),
|
||||||
|
Binding("space", "toggle", "Toggle", show=False),
|
||||||
|
Binding("escape", "cancel", "Cancel", show=False),
|
||||||
|
]
|
||||||
|
|
||||||
|
# -- subclass adapters --------------------------------------------
|
||||||
|
|
||||||
|
_INSTALLED_SUFFIX: ClassVar[str] = " (installed)"
|
||||||
|
_PHASE2_TITLE: ClassVar[str] = "Select items to install"
|
||||||
|
_PHASE2_CONFIRM_LABEL: ClassVar[str] = "install"
|
||||||
|
|
||||||
|
def _item_name(self, item: Any) -> str:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def _item_tags(self, item: Any) -> list[str]:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def _item_desc(self, item: Any) -> str:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def _post_confirmed(self, items: list[Any]) -> None:
|
||||||
|
"""Post the widget-specific ``Confirmed`` message."""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def _post_cancelled(self) -> None:
|
||||||
|
"""Post the widget-specific ``Cancelled`` message."""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
# -- state ---------------------------------------------------------
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
items: list[Any],
|
||||||
|
installed_names: set[str],
|
||||||
|
*,
|
||||||
|
pre_filter_tag: str = "",
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> None:
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
self._all_items = items
|
||||||
|
self._installed_names = installed_names
|
||||||
|
self._pre_filter_tag = pre_filter_tag.lower()
|
||||||
|
self._selected = 0
|
||||||
|
self._row_widgets: list[Static] = []
|
||||||
|
self._title_widget: Static | None = None
|
||||||
|
self._help_widget: Static | None = None
|
||||||
|
|
||||||
|
self._phase: int = 1
|
||||||
|
self._filtered_items: list[Any] = []
|
||||||
|
self._checked: set[int] = set()
|
||||||
|
|
||||||
|
# Build tag list (sorted by count desc, then alphabetically)
|
||||||
|
from collections import Counter
|
||||||
|
|
||||||
|
tag_counter: Counter[str] = Counter()
|
||||||
|
for item in self._all_items:
|
||||||
|
for t in self._item_tags(item):
|
||||||
|
tag_counter[t.lower()] += 1
|
||||||
|
sorted_tags = sorted(tag_counter.items(), key=lambda x: (-x[1], x[0]))
|
||||||
|
self._tag_items: list[tuple[str, int]] = [
|
||||||
|
("all", len(self._all_items)),
|
||||||
|
*sorted_tags,
|
||||||
|
]
|
||||||
|
|
||||||
|
# If pre-filtered, skip to phase 2
|
||||||
|
if self._pre_filter_tag:
|
||||||
|
self._filtered_items = self._items_with_tag(self._pre_filter_tag)
|
||||||
|
if self._filtered_items:
|
||||||
|
self._phase = 2
|
||||||
|
else:
|
||||||
|
self._pre_filter_tag = ""
|
||||||
|
|
||||||
|
def _items_with_tag(self, tag: str) -> list[Any]:
|
||||||
|
if tag == "all":
|
||||||
|
return list(self._all_items)
|
||||||
|
return [
|
||||||
|
item
|
||||||
|
for item in self._all_items
|
||||||
|
if tag in [t.lower() for t in self._item_tags(item)]
|
||||||
|
]
|
||||||
|
|
||||||
|
# -- layout ---------------------------------------------------------
|
||||||
|
|
||||||
|
def compose(self) -> ComposeResult:
|
||||||
|
self._title_widget = Static("", classes="browser-title")
|
||||||
|
yield self._title_widget
|
||||||
|
with Container(classes="browser-rows"):
|
||||||
|
max_rows = max(len(self._tag_items), len(self._all_items))
|
||||||
|
for _ in range(max_rows):
|
||||||
|
widget = Static("", classes="browser-row")
|
||||||
|
self._row_widgets.append(widget)
|
||||||
|
yield widget
|
||||||
|
self._help_widget = Static("", classes="browser-help")
|
||||||
|
yield self._help_widget
|
||||||
|
|
||||||
|
# -- rendering ------------------------------------------------------
|
||||||
|
|
||||||
|
def _refresh_view(self) -> None:
|
||||||
|
if self._phase == 1:
|
||||||
|
self._render_tag_picker()
|
||||||
|
else:
|
||||||
|
self._render_checkbox_list()
|
||||||
|
|
||||||
|
def _render_tag_picker(self) -> None:
|
||||||
|
if self._title_widget:
|
||||||
|
self._title_widget.update("Filter by tag:")
|
||||||
|
if self._help_widget:
|
||||||
|
self._help_widget.update("↑/↓ navigate · Enter select · Esc cancel")
|
||||||
|
|
||||||
|
for i, widget in enumerate(self._row_widgets):
|
||||||
|
if i < len(self._tag_items):
|
||||||
|
tag, count = self._tag_items[i]
|
||||||
|
is_selected = i == self._selected
|
||||||
|
text = Text()
|
||||||
|
cursor = "▸ " if is_selected else " "
|
||||||
|
text.append(cursor, style="bold cyan" if is_selected else "dim")
|
||||||
|
text.append(f"{tag} ({count})", style="bold" if is_selected else "")
|
||||||
|
widget.update(text)
|
||||||
|
widget.display = True
|
||||||
|
self.apply_row_highlight(widget, is_selected, "browser-row-selected")
|
||||||
|
else:
|
||||||
|
widget.update("")
|
||||||
|
widget.display = False
|
||||||
|
|
||||||
|
def _render_checkbox_list(self) -> None:
|
||||||
|
n_checked = len(
|
||||||
|
[
|
||||||
|
i
|
||||||
|
for i in self._checked
|
||||||
|
if self._item_name(self._filtered_items[i]) not in self._installed_names
|
||||||
|
]
|
||||||
|
)
|
||||||
|
if self._title_widget:
|
||||||
|
self._title_widget.update(f"{self._PHASE2_TITLE} ({n_checked} selected):")
|
||||||
|
if self._help_widget:
|
||||||
|
self._help_widget.update(
|
||||||
|
"↑/↓ navigate · Space toggle · "
|
||||||
|
f"Enter {self._PHASE2_CONFIRM_LABEL} · Esc cancel"
|
||||||
|
)
|
||||||
|
|
||||||
|
for i, widget in enumerate(self._row_widgets):
|
||||||
|
if i < len(self._filtered_items):
|
||||||
|
item = self._filtered_items[i]
|
||||||
|
name = self._item_name(item)
|
||||||
|
is_selected = i == self._selected
|
||||||
|
is_installed = name in self._installed_names
|
||||||
|
is_checked = i in self._checked
|
||||||
|
|
||||||
|
text = Text()
|
||||||
|
cursor = "▸ " if is_selected else " "
|
||||||
|
text.append(cursor, style="bold cyan" if is_selected else "dim")
|
||||||
|
|
||||||
|
if is_installed:
|
||||||
|
suffix = self._INSTALLED_SUFFIX
|
||||||
|
desc = self._truncate(self._item_desc(item), name, suffix=suffix)
|
||||||
|
text.append("✓ ", style="green")
|
||||||
|
text.append(name, style="green dim")
|
||||||
|
text.append(f" — {desc}", style="dim")
|
||||||
|
text.append(suffix, style="dim italic")
|
||||||
|
elif is_checked:
|
||||||
|
desc = self._truncate(self._item_desc(item), name)
|
||||||
|
text.append("● ", style="green bold")
|
||||||
|
text.append(name, style="bold")
|
||||||
|
text.append(f" — {desc}", style="")
|
||||||
|
else:
|
||||||
|
desc = self._truncate(self._item_desc(item), name)
|
||||||
|
text.append("○ ", style="dim")
|
||||||
|
text.append(name, style="bold" if is_selected else "")
|
||||||
|
text.append(f" — {desc}", style="dim")
|
||||||
|
|
||||||
|
widget.update(text)
|
||||||
|
widget.display = True
|
||||||
|
self.apply_row_highlight(widget, is_selected, "browser-row-selected")
|
||||||
|
else:
|
||||||
|
widget.update("")
|
||||||
|
widget.display = False
|
||||||
|
|
||||||
|
# -- actions ----------------------------------------------------------
|
||||||
|
|
||||||
|
def _current_items_count(self) -> int:
|
||||||
|
if self._phase == 1:
|
||||||
|
return len(self._tag_items)
|
||||||
|
return len(self._filtered_items)
|
||||||
|
|
||||||
|
def action_move_up(self) -> None:
|
||||||
|
n = self._current_items_count()
|
||||||
|
if not n:
|
||||||
|
return
|
||||||
|
self._selected = (self._selected - 1) % n
|
||||||
|
self._refresh_view()
|
||||||
|
|
||||||
|
def action_move_down(self) -> None:
|
||||||
|
n = self._current_items_count()
|
||||||
|
if not n:
|
||||||
|
return
|
||||||
|
self._selected = (self._selected + 1) % n
|
||||||
|
self._refresh_view()
|
||||||
|
|
||||||
|
def action_toggle(self) -> None:
|
||||||
|
"""Toggle checkbox selection (phase 2 only)."""
|
||||||
|
if self._phase != 2 or not self._filtered_items:
|
||||||
|
return
|
||||||
|
if self._item_name(self._filtered_items[self._selected]) in (
|
||||||
|
self._installed_names
|
||||||
|
):
|
||||||
|
return # Can't toggle already-installed items
|
||||||
|
if self._selected in self._checked:
|
||||||
|
self._checked.discard(self._selected)
|
||||||
|
else:
|
||||||
|
self._checked.add(self._selected)
|
||||||
|
self._refresh_view()
|
||||||
|
|
||||||
|
def action_confirm(self) -> None:
|
||||||
|
if self._phase == 1:
|
||||||
|
if not self._tag_items:
|
||||||
|
return
|
||||||
|
tag, _ = self._tag_items[self._selected]
|
||||||
|
self._filtered_items = self._items_with_tag(tag)
|
||||||
|
self._phase = 2
|
||||||
|
self._selected = 0
|
||||||
|
self._checked = set()
|
||||||
|
self._refresh_view()
|
||||||
|
else:
|
||||||
|
items = [
|
||||||
|
self._filtered_items[i]
|
||||||
|
for i in sorted(self._checked)
|
||||||
|
if self._item_name(self._filtered_items[i]) not in self._installed_names
|
||||||
|
]
|
||||||
|
self._post_confirmed(items)
|
||||||
|
|
||||||
|
def action_cancel(self) -> None:
|
||||||
|
if self._phase == 2 and not self._pre_filter_tag:
|
||||||
|
# Go back to tag picker
|
||||||
|
self._phase = 1
|
||||||
|
self._selected = 0
|
||||||
|
self._checked = set()
|
||||||
|
self._refresh_view()
|
||||||
|
else:
|
||||||
|
self._post_cancelled()
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Inline skill browser widget for /evoskills in TUI.
|
"""Inline skill browser widget for /evoskills in TUI.
|
||||||
|
|
||||||
Two-phase keyboard-driven widget:
|
Two-phase keyboard-driven widget built on the shared picker engine
|
||||||
|
(``picker_base.TagCheckboxBrowserBase``):
|
||||||
Phase 1 — tag picker (arrow keys + Enter to select, or Esc for all)
|
Phase 1 — tag picker (arrow keys + Enter to select, or Esc for all)
|
||||||
Phase 2 — skill checkbox (arrow keys to navigate, Space to toggle, Enter to confirm)
|
Phase 2 — skill checkbox (arrow keys to navigate, Space to toggle, Enter to confirm)
|
||||||
|
|
||||||
@@ -10,73 +11,19 @@ or ``SkillBrowserWidget.Cancelled`` on Esc.
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Any, ClassVar
|
from typing import Any, ClassVar
|
||||||
|
|
||||||
from rich.text import Text
|
|
||||||
from textual.binding import Binding, BindingType
|
|
||||||
from textual.containers import Container
|
|
||||||
from textual.message import Message
|
from textual.message import Message
|
||||||
from textual.widget import Widget
|
|
||||||
from textual.widgets import Static
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
from .picker_base import TagCheckboxBrowserBase
|
||||||
from textual import events
|
|
||||||
from textual.app import ComposeResult
|
|
||||||
|
|
||||||
|
|
||||||
class SkillBrowserWidget(Widget):
|
class SkillBrowserWidget(TagCheckboxBrowserBase):
|
||||||
"""Inline skill browser — mounts in chat, keyboard-driven.
|
"""Inline skill browser — mounts in chat, keyboard-driven."""
|
||||||
|
|
||||||
Phase 1: Tag picker (select a tag filter or "All").
|
_INSTALLED_SUFFIX: ClassVar[str] = " (installed)"
|
||||||
Phase 2: Skill checkbox (toggle skills, confirm to install).
|
_PHASE2_TITLE: ClassVar[str] = "Select skills to install"
|
||||||
"""
|
_PHASE2_CONFIRM_LABEL: ClassVar[str] = "install"
|
||||||
|
|
||||||
can_focus = True
|
|
||||||
can_focus_children = False
|
|
||||||
|
|
||||||
DEFAULT_CSS = """
|
|
||||||
SkillBrowserWidget {
|
|
||||||
height: auto;
|
|
||||||
max-height: 30;
|
|
||||||
margin: 1 0;
|
|
||||||
padding: 0 1;
|
|
||||||
background: $surface;
|
|
||||||
border: solid $primary;
|
|
||||||
}
|
|
||||||
SkillBrowserWidget .browser-title {
|
|
||||||
height: 1;
|
|
||||||
text-style: bold;
|
|
||||||
color: $primary;
|
|
||||||
}
|
|
||||||
SkillBrowserWidget .browser-rows {
|
|
||||||
height: auto;
|
|
||||||
max-height: 20;
|
|
||||||
overflow-y: auto;
|
|
||||||
}
|
|
||||||
SkillBrowserWidget .browser-row {
|
|
||||||
height: 1;
|
|
||||||
padding: 0 1;
|
|
||||||
}
|
|
||||||
SkillBrowserWidget .browser-row-selected {
|
|
||||||
background: $primary;
|
|
||||||
text-style: bold;
|
|
||||||
}
|
|
||||||
SkillBrowserWidget .browser-help {
|
|
||||||
height: 1;
|
|
||||||
color: $text-muted;
|
|
||||||
text-style: italic;
|
|
||||||
}
|
|
||||||
"""
|
|
||||||
|
|
||||||
BINDINGS: ClassVar[list[BindingType]] = [
|
|
||||||
Binding("up", "move_up", "Up", show=False),
|
|
||||||
Binding("k", "move_up", "Up", show=False),
|
|
||||||
Binding("down", "move_down", "Down", show=False),
|
|
||||||
Binding("j", "move_down", "Down", show=False),
|
|
||||||
Binding("enter", "confirm", "Confirm", show=False),
|
|
||||||
Binding("space", "toggle", "Toggle", show=False),
|
|
||||||
Binding("escape", "cancel", "Cancel", show=False),
|
|
||||||
]
|
|
||||||
|
|
||||||
class Confirmed(Message):
|
class Confirmed(Message):
|
||||||
"""Posted when user confirms skill selection."""
|
"""Posted when user confirms skill selection."""
|
||||||
@@ -88,261 +35,17 @@ class SkillBrowserWidget(Widget):
|
|||||||
class Cancelled(Message):
|
class Cancelled(Message):
|
||||||
"""Posted when user cancels."""
|
"""Posted when user cancels."""
|
||||||
|
|
||||||
def __init__(
|
def _item_name(self, item: Any) -> str:
|
||||||
self,
|
return item["name"]
|
||||||
index: list[dict],
|
|
||||||
installed_names: set[str],
|
|
||||||
*,
|
|
||||||
pre_filter_tag: str = "",
|
|
||||||
**kwargs: Any,
|
|
||||||
) -> None:
|
|
||||||
super().__init__(**kwargs)
|
|
||||||
self._index = index
|
|
||||||
self._installed_names = installed_names
|
|
||||||
self._pre_filter_tag = pre_filter_tag.lower()
|
|
||||||
self._selected = 0
|
|
||||||
self._row_widgets: list[Static] = []
|
|
||||||
self._title_widget: Static | None = None
|
|
||||||
self._help_widget: Static | None = None
|
|
||||||
|
|
||||||
# Phase 1: tag picker
|
def _item_tags(self, item: Any) -> list[str]:
|
||||||
# Phase 2: skill checkbox
|
return item.get("tags", [])
|
||||||
self._phase: int = 1
|
|
||||||
self._tag_items: list[tuple[str, int]] = [] # (tag, count)
|
|
||||||
self._skill_items: list[dict] = [] # filtered skills
|
|
||||||
self._checked: set[int] = set() # indices of checked skills
|
|
||||||
|
|
||||||
# Build tag list (sorted by count desc, then alphabetically)
|
def _item_desc(self, item: Any) -> str:
|
||||||
from collections import Counter
|
return item["description"]
|
||||||
|
|
||||||
tag_counter: Counter[str] = Counter()
|
def _post_confirmed(self, items: list[Any]) -> None:
|
||||||
for s in self._index:
|
self.post_message(self.Confirmed([s["install_source"] for s in items]))
|
||||||
for t in s.get("tags", []):
|
|
||||||
tag_counter[t.lower()] += 1
|
|
||||||
sorted_tags = sorted(tag_counter.items(), key=lambda x: (-x[1], x[0]))
|
|
||||||
self._tag_items = [("all", len(self._index)), *sorted_tags]
|
|
||||||
|
|
||||||
# If pre-filtered, skip to phase 2
|
def _post_cancelled(self) -> None:
|
||||||
if self._pre_filter_tag:
|
self.post_message(self.Cancelled())
|
||||||
self._skill_items = [
|
|
||||||
s
|
|
||||||
for s in self._index
|
|
||||||
if self._pre_filter_tag in [t.lower() for t in s.get("tags", [])]
|
|
||||||
]
|
|
||||||
if self._skill_items:
|
|
||||||
self._phase = 2
|
|
||||||
else:
|
|
||||||
# No matches — show tag picker anyway
|
|
||||||
self._pre_filter_tag = ""
|
|
||||||
|
|
||||||
def compose(self) -> ComposeResult:
|
|
||||||
self._title_widget = Static("", classes="browser-title")
|
|
||||||
yield self._title_widget
|
|
||||||
with Container(classes="browser-rows"):
|
|
||||||
# Pre-allocate enough rows for the larger of tag list or skill list
|
|
||||||
max_rows = max(len(self._tag_items), len(self._index))
|
|
||||||
for _ in range(max_rows):
|
|
||||||
widget = Static("", classes="browser-row")
|
|
||||||
self._row_widgets.append(widget)
|
|
||||||
yield widget
|
|
||||||
self._help_widget = Static("", classes="browser-help")
|
|
||||||
yield self._help_widget
|
|
||||||
|
|
||||||
def on_mount(self) -> None:
|
|
||||||
# Defer rendering until after layout so self.size is populated
|
|
||||||
self.call_after_refresh(self._update_display)
|
|
||||||
self.call_later(self.focus)
|
|
||||||
|
|
||||||
def _update_display(self) -> None:
|
|
||||||
if self._phase == 1:
|
|
||||||
self._render_tag_picker()
|
|
||||||
else:
|
|
||||||
self._render_skill_checkbox()
|
|
||||||
|
|
||||||
def _render_tag_picker(self) -> None:
|
|
||||||
if self._title_widget:
|
|
||||||
self._title_widget.update("Filter by tag:")
|
|
||||||
if self._help_widget:
|
|
||||||
self._help_widget.update("↑/↓ navigate · Enter select · Esc cancel")
|
|
||||||
|
|
||||||
for i, widget in enumerate(self._row_widgets):
|
|
||||||
if i < len(self._tag_items):
|
|
||||||
tag, count = self._tag_items[i]
|
|
||||||
is_selected = i == self._selected
|
|
||||||
text = Text()
|
|
||||||
cursor = "▸ " if is_selected else " "
|
|
||||||
text.append(cursor, style="bold cyan" if is_selected else "dim")
|
|
||||||
label = f"{tag} ({count})"
|
|
||||||
text.append(label, style="bold" if is_selected else "")
|
|
||||||
widget.update(text)
|
|
||||||
widget.display = True
|
|
||||||
widget.remove_class("browser-row-selected")
|
|
||||||
if is_selected:
|
|
||||||
widget.add_class("browser-row-selected")
|
|
||||||
widget.scroll_visible()
|
|
||||||
else:
|
|
||||||
widget.update("")
|
|
||||||
widget.display = False
|
|
||||||
|
|
||||||
def _row_content_width(self) -> int:
|
|
||||||
"""Get the usable character width for a row's text content.
|
|
||||||
|
|
||||||
Accounts for widget border, widget padding, and row padding.
|
|
||||||
Falls back to terminal width if the widget hasn't been laid out yet.
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
w = self.size.width
|
|
||||||
if w > 0:
|
|
||||||
# border (2) + widget padding-left/right (2) + row padding-left/right (2)
|
|
||||||
return w - 6
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
# Fallback: use terminal width minus reasonable chrome
|
|
||||||
try:
|
|
||||||
return self.app.size.width - 10
|
|
||||||
except Exception:
|
|
||||||
return 100
|
|
||||||
|
|
||||||
def _truncate(self, desc: str, name: str, *, suffix: str = "") -> str:
|
|
||||||
"""Truncate a description to fit the row, adding ellipsis if needed."""
|
|
||||||
# cursor(2) + indicator(2) + name + " — "(3) + suffix
|
|
||||||
overhead = 2 + 2 + len(name) + 3 + len(suffix)
|
|
||||||
max_len = max(20, self._row_content_width() - overhead)
|
|
||||||
if len(desc) <= max_len:
|
|
||||||
return desc
|
|
||||||
return desc[: max_len - 1] + "…"
|
|
||||||
|
|
||||||
def _render_skill_checkbox(self) -> None:
|
|
||||||
n_checked = len(
|
|
||||||
[
|
|
||||||
i
|
|
||||||
for i in self._checked
|
|
||||||
if self._skill_items[i]["name"] not in self._installed_names
|
|
||||||
]
|
|
||||||
)
|
|
||||||
if self._title_widget:
|
|
||||||
self._title_widget.update(
|
|
||||||
f"Select skills to install ({n_checked} selected):"
|
|
||||||
)
|
|
||||||
if self._help_widget:
|
|
||||||
self._help_widget.update(
|
|
||||||
"↑/↓ navigate · Space toggle · Enter install · Esc cancel"
|
|
||||||
)
|
|
||||||
|
|
||||||
for i, widget in enumerate(self._row_widgets):
|
|
||||||
if i < len(self._skill_items):
|
|
||||||
skill = self._skill_items[i]
|
|
||||||
is_selected = i == self._selected
|
|
||||||
is_installed = skill["name"] in self._installed_names
|
|
||||||
is_checked = i in self._checked
|
|
||||||
|
|
||||||
text = Text()
|
|
||||||
cursor = "▸ " if is_selected else " "
|
|
||||||
text.append(cursor, style="bold cyan" if is_selected else "dim")
|
|
||||||
|
|
||||||
if is_installed:
|
|
||||||
suffix = " (installed)"
|
|
||||||
desc = self._truncate(
|
|
||||||
desc=skill["description"],
|
|
||||||
name=skill["name"],
|
|
||||||
suffix=suffix,
|
|
||||||
)
|
|
||||||
text.append("✓ ", style="green")
|
|
||||||
text.append(skill["name"], style="green dim")
|
|
||||||
text.append(f" — {desc}", style="dim")
|
|
||||||
text.append(suffix, style="dim italic")
|
|
||||||
elif is_checked:
|
|
||||||
desc = self._truncate(skill["description"], skill["name"])
|
|
||||||
text.append("● ", style="green bold")
|
|
||||||
text.append(skill["name"], style="bold")
|
|
||||||
text.append(f" — {desc}", style="")
|
|
||||||
else:
|
|
||||||
desc = self._truncate(skill["description"], skill["name"])
|
|
||||||
text.append("○ ", style="dim")
|
|
||||||
text.append(skill["name"], style="bold" if is_selected else "")
|
|
||||||
text.append(f" — {desc}", style="dim")
|
|
||||||
|
|
||||||
widget.update(text)
|
|
||||||
widget.display = True
|
|
||||||
widget.remove_class("browser-row-selected")
|
|
||||||
if is_selected:
|
|
||||||
widget.add_class("browser-row-selected")
|
|
||||||
widget.scroll_visible()
|
|
||||||
else:
|
|
||||||
widget.update("")
|
|
||||||
widget.display = False
|
|
||||||
|
|
||||||
def _current_items_count(self) -> int:
|
|
||||||
if self._phase == 1:
|
|
||||||
return len(self._tag_items)
|
|
||||||
return len(self._skill_items)
|
|
||||||
|
|
||||||
def action_move_up(self) -> None:
|
|
||||||
n = self._current_items_count()
|
|
||||||
if not n:
|
|
||||||
return
|
|
||||||
self._selected = (self._selected - 1) % n
|
|
||||||
self._update_display()
|
|
||||||
|
|
||||||
def action_move_down(self) -> None:
|
|
||||||
n = self._current_items_count()
|
|
||||||
if not n:
|
|
||||||
return
|
|
||||||
self._selected = (self._selected + 1) % n
|
|
||||||
self._update_display()
|
|
||||||
|
|
||||||
def action_toggle(self) -> None:
|
|
||||||
"""Toggle skill selection (phase 2 only)."""
|
|
||||||
if self._phase != 2:
|
|
||||||
return
|
|
||||||
if not self._skill_items:
|
|
||||||
return
|
|
||||||
skill = self._skill_items[self._selected]
|
|
||||||
if skill["name"] in self._installed_names:
|
|
||||||
return # Can't toggle installed skills
|
|
||||||
if self._selected in self._checked:
|
|
||||||
self._checked.discard(self._selected)
|
|
||||||
else:
|
|
||||||
self._checked.add(self._selected)
|
|
||||||
self._update_display()
|
|
||||||
|
|
||||||
def action_confirm(self) -> None:
|
|
||||||
if self._phase == 1:
|
|
||||||
# Transition to phase 2
|
|
||||||
if not self._tag_items:
|
|
||||||
return
|
|
||||||
tag, _ = self._tag_items[self._selected]
|
|
||||||
if tag == "all":
|
|
||||||
self._skill_items = list(self._index)
|
|
||||||
else:
|
|
||||||
self._skill_items = [
|
|
||||||
s
|
|
||||||
for s in self._index
|
|
||||||
if tag in [t.lower() for t in s.get("tags", [])]
|
|
||||||
]
|
|
||||||
self._phase = 2
|
|
||||||
self._selected = 0
|
|
||||||
self._checked = set()
|
|
||||||
self._update_display()
|
|
||||||
else:
|
|
||||||
# Confirm selection
|
|
||||||
sources = [
|
|
||||||
self._skill_items[i]["install_source"]
|
|
||||||
for i in sorted(self._checked)
|
|
||||||
if self._skill_items[i]["name"] not in self._installed_names
|
|
||||||
]
|
|
||||||
self.post_message(self.Confirmed(sources))
|
|
||||||
|
|
||||||
def action_cancel(self) -> None:
|
|
||||||
if self._phase == 2 and not self._pre_filter_tag:
|
|
||||||
# Go back to tag picker
|
|
||||||
self._phase = 1
|
|
||||||
self._selected = 0
|
|
||||||
self._checked = set()
|
|
||||||
self._update_display()
|
|
||||||
else:
|
|
||||||
self.post_message(self.Cancelled())
|
|
||||||
|
|
||||||
def on_blur(self, event: events.Blur) -> None:
|
|
||||||
"""Re-focus to keep focus trapped until decision is made."""
|
|
||||||
self.call_after_refresh(self.focus)
|
|
||||||
|
|||||||
@@ -23,11 +23,11 @@ from rich.text import Text
|
|||||||
from textual.binding import Binding, BindingType
|
from textual.binding import Binding, BindingType
|
||||||
from textual.containers import Container
|
from textual.containers import Container
|
||||||
from textual.message import Message
|
from textual.message import Message
|
||||||
from textual.widget import Widget
|
|
||||||
from textual.widgets import Static
|
from textual.widgets import Static
|
||||||
|
|
||||||
|
from .picker_base import PickerWidgetBase, first_selectable_index, move_selection
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from textual import events
|
|
||||||
from textual.app import ComposeResult
|
from textual.app import ComposeResult
|
||||||
|
|
||||||
|
|
||||||
@@ -237,16 +237,13 @@ def build_row_text(
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
class ThreadPickerWidget(Widget):
|
class ThreadPickerWidget(PickerWidgetBase):
|
||||||
"""Inline thread picker — mounts in chat, keyboard-driven.
|
"""Inline thread picker — mounts in chat, keyboard-driven.
|
||||||
|
|
||||||
Posts ``Picked(thread_id)`` on Enter, ``Cancelled()`` on Esc.
|
Posts ``Picked(thread_id)`` on Enter, ``Cancelled()`` on Esc.
|
||||||
Threads are displayed in a two-level workspace hierarchy.
|
Threads are displayed in a two-level workspace hierarchy.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
can_focus = True
|
|
||||||
can_focus_children = False
|
|
||||||
|
|
||||||
DEFAULT_CSS = """
|
DEFAULT_CSS = """
|
||||||
ThreadPickerWidget {
|
ThreadPickerWidget {
|
||||||
height: auto;
|
height: auto;
|
||||||
@@ -323,22 +320,19 @@ class ThreadPickerWidget(Widget):
|
|||||||
self._selected = self._first_thread_index()
|
self._selected = self._first_thread_index()
|
||||||
self._row_widgets: list[Static] = []
|
self._row_widgets: list[Static] = []
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_thread(item: dict) -> bool:
|
||||||
|
return item["type"] == "thread"
|
||||||
|
|
||||||
def _first_thread_index(self) -> int:
|
def _first_thread_index(self) -> int:
|
||||||
for i, item in enumerate(self._items):
|
return first_selectable_index(self._items, self._is_thread)
|
||||||
if item["type"] == "thread":
|
|
||||||
return i
|
|
||||||
return 0
|
|
||||||
|
|
||||||
def _move(self, direction: int) -> None:
|
def _move(self, direction: int) -> None:
|
||||||
if not self._items:
|
if not self._items:
|
||||||
return
|
return
|
||||||
i = (self._selected + direction) % len(self._items)
|
new = move_selection(self._items, self._selected, direction, self._is_thread)
|
||||||
steps = 0
|
if self._is_thread(self._items[new]):
|
||||||
while self._items[i]["type"] != "thread" and steps < len(self._items):
|
self._selected = new
|
||||||
i = (i + direction) % len(self._items)
|
|
||||||
steps += 1
|
|
||||||
if self._items[i]["type"] == "thread":
|
|
||||||
self._selected = i
|
|
||||||
self._update_rows()
|
self._update_rows()
|
||||||
|
|
||||||
def compose(self) -> ComposeResult:
|
def compose(self) -> ComposeResult:
|
||||||
@@ -358,18 +352,18 @@ class ThreadPickerWidget(Widget):
|
|||||||
classes="picker-help",
|
classes="picker-help",
|
||||||
)
|
)
|
||||||
|
|
||||||
def on_mount(self) -> None:
|
def _refresh_view(self) -> None:
|
||||||
self._update_rows()
|
self._update_rows()
|
||||||
self.call_later(self.focus)
|
|
||||||
|
|
||||||
def _update_rows(self) -> None:
|
def _update_rows(self) -> None:
|
||||||
for i, (item, widget) in enumerate(
|
for i, (item, widget) in enumerate(
|
||||||
zip(self._items, self._row_widgets, strict=False)
|
zip(self._items, self._row_widgets, strict=False)
|
||||||
):
|
):
|
||||||
widget.remove_class("picker-row-selected")
|
|
||||||
if item["type"] == "header":
|
if item["type"] == "header":
|
||||||
|
widget.remove_class("picker-row-selected")
|
||||||
widget.update(build_header_text(item["label"]))
|
widget.update(build_header_text(item["label"]))
|
||||||
elif item["type"] == "subheader":
|
elif item["type"] == "subheader":
|
||||||
|
widget.remove_class("picker-row-selected")
|
||||||
widget.update(build_subheader_text(item["label"]))
|
widget.update(build_subheader_text(item["label"]))
|
||||||
else:
|
else:
|
||||||
thread = item["thread"]
|
thread = item["thread"]
|
||||||
@@ -381,9 +375,7 @@ class ThreadPickerWidget(Widget):
|
|||||||
indented=item.get("indented", False),
|
indented=item.get("indented", False),
|
||||||
)
|
)
|
||||||
widget.update(text)
|
widget.update(text)
|
||||||
if is_selected:
|
self.apply_row_highlight(widget, is_selected)
|
||||||
widget.add_class("picker-row-selected")
|
|
||||||
widget.scroll_visible()
|
|
||||||
|
|
||||||
def action_move_up(self) -> None:
|
def action_move_up(self) -> None:
|
||||||
self._move(-1)
|
self._move(-1)
|
||||||
@@ -403,6 +395,3 @@ class ThreadPickerWidget(Widget):
|
|||||||
|
|
||||||
def action_cancel(self) -> None:
|
def action_cancel(self) -> None:
|
||||||
self.post_message(self.Cancelled())
|
self.post_message(self.Cancelled())
|
||||||
|
|
||||||
def on_blur(self, event: events.Blur) -> None:
|
|
||||||
self.call_after_refresh(self.focus)
|
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any, ClassVar, Protocol, runtime_checkable
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from ..gateway import GraphGateway
|
from ..gateway import GraphGateway
|
||||||
|
from ..runtime import AsyncRuntime
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -65,20 +66,47 @@ class CommandUI(Protocol):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ChannelRuntime:
|
class ChannelRuntime:
|
||||||
"""Mutable handle to the agent + thread bound to running channels."""
|
"""Mutable handle to the agent + thread bound to running channels.
|
||||||
|
|
||||||
|
Also holds session-scoped bindings mutated by slash commands — the
|
||||||
|
``active_teams`` list backs the ``/expert`` command, feeding into
|
||||||
|
``RunRequest.configurable_extra`` at stream call time.
|
||||||
|
"""
|
||||||
|
|
||||||
agent: Any = None
|
agent: Any = None
|
||||||
thread_id: str | None = None
|
thread_id: str | None = None
|
||||||
|
active_teams: list[str] = field(default_factory=list)
|
||||||
|
|
||||||
def bind(self, agent: Any, thread_id: str) -> None:
|
def bind(self, agent: Any, thread_id: str) -> None:
|
||||||
self.agent = agent
|
self.agent = agent
|
||||||
self.thread_id = thread_id
|
self.thread_id = thread_id
|
||||||
|
|
||||||
def clear(self) -> None:
|
def clear(self) -> None:
|
||||||
|
# ``active_teams`` is session-scoped and reset explicitly by ``/new``
|
||||||
|
# (session.py) and ``/expert clear`` — not tied to channel lifecycle.
|
||||||
|
# Clearing here on channel shutdown would silently dismiss the user's
|
||||||
|
# invited experts, which they never asked for.
|
||||||
self.agent = None
|
self.agent = None
|
||||||
self.thread_id = None
|
self.thread_id = None
|
||||||
|
|
||||||
|
|
||||||
|
def active_teams_configurable_extra(
|
||||||
|
runtime: ChannelRuntime | None,
|
||||||
|
) -> dict[str, Any] | None:
|
||||||
|
"""Build ``RunRequest.configurable_extra`` from a channel runtime.
|
||||||
|
|
||||||
|
Returns ``{"active_teams": [...]}`` when the runtime has invited
|
||||||
|
experts, or ``None`` when there is no runtime or no active invites —
|
||||||
|
lets stream call sites forward the field unconditionally without
|
||||||
|
each duplicating the "read runtime slot, build dict, drop when
|
||||||
|
empty" three-liner.
|
||||||
|
"""
|
||||||
|
if runtime is None:
|
||||||
|
return None
|
||||||
|
invited = list(runtime.active_teams)
|
||||||
|
return {"active_teams": invited} if invited else None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class CommandContext:
|
class CommandContext:
|
||||||
"""Context passed to commands during execution."""
|
"""Context passed to commands during execution."""
|
||||||
@@ -91,6 +119,7 @@ class CommandContext:
|
|||||||
config: Any = None
|
config: Any = None
|
||||||
channel_runtime: ChannelRuntime | None = None
|
channel_runtime: ChannelRuntime | None = None
|
||||||
graph_gateway: GraphGateway | None = None
|
graph_gateway: GraphGateway | None = None
|
||||||
|
async_runtime: AsyncRuntime | None = None
|
||||||
command_error: str | None = None
|
command_error: str | None = None
|
||||||
# Real LLM input token count from last usage_metadata (includes system
|
# Real LLM input token count from last usage_metadata (includes system
|
||||||
# prompt + tool schemas). Used by /compact for accurate display.
|
# prompt + tool schemas). Used by /compact for accurate display.
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_COMMAND_OUTPUT_FAILURE_NOTICE = "Command output could not be delivered."
|
||||||
|
|
||||||
|
|
||||||
class ChannelCommandUI(CommandUI):
|
class ChannelCommandUI(CommandUI):
|
||||||
"""CommandUI implementation for messaging channels with output buffering."""
|
"""CommandUI implementation for messaging channels with output buffering."""
|
||||||
@@ -37,6 +39,10 @@ class ChannelCommandUI(CommandUI):
|
|||||||
self.handle_session_resume_callback = handle_session_resume_callback
|
self.handle_session_resume_callback = handle_session_resume_callback
|
||||||
self.graph_gateway = graph_gateway
|
self.graph_gateway = graph_gateway
|
||||||
self._system_buffer: list[str] = []
|
self._system_buffer: list[str] = []
|
||||||
|
# Whether any output was delivered (or scheduled for delivery) to the
|
||||||
|
# channel. The slash dispatcher consults this to decide between a
|
||||||
|
# bare completion ack and staying silent.
|
||||||
|
self.sent_to_channel: bool = False
|
||||||
|
|
||||||
def _queue_system(
|
def _queue_system(
|
||||||
self,
|
self,
|
||||||
@@ -107,6 +113,7 @@ class ChannelCommandUI(CommandUI):
|
|||||||
content=grouped_text,
|
content=grouped_text,
|
||||||
reply_to=self.msg.message_id,
|
reply_to=self.msg.message_id,
|
||||||
metadata=self.msg.metadata,
|
metadata=self.msg.metadata,
|
||||||
|
failure_notice=_COMMAND_OUTPUT_FAILURE_NOTICE,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.msg.bus_ref:
|
if self.msg.bus_ref:
|
||||||
@@ -114,6 +121,7 @@ class ChannelCommandUI(CommandUI):
|
|||||||
else:
|
else:
|
||||||
coro = self.msg.channel_ref.send(outbound)
|
coro = self.msg.channel_ref.send(outbound)
|
||||||
|
|
||||||
|
self.sent_to_channel = True
|
||||||
asyncio.run_coroutine_threadsafe(coro, loop)
|
asyncio.run_coroutine_threadsafe(coro, loop)
|
||||||
|
|
||||||
def mount_renderable(self, renderable: Any) -> None:
|
def mount_renderable(self, renderable: Any) -> None:
|
||||||
@@ -147,6 +155,7 @@ class ChannelCommandUI(CommandUI):
|
|||||||
content=f"```\n{text}\n```",
|
content=f"```\n{text}\n```",
|
||||||
reply_to=self.msg.message_id,
|
reply_to=self.msg.message_id,
|
||||||
metadata=self.msg.metadata,
|
metadata=self.msg.metadata,
|
||||||
|
failure_notice=_COMMAND_OUTPUT_FAILURE_NOTICE,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.msg.bus_ref:
|
if self.msg.bus_ref:
|
||||||
@@ -154,6 +163,7 @@ class ChannelCommandUI(CommandUI):
|
|||||||
else:
|
else:
|
||||||
coro = self.msg.channel_ref.send(outbound)
|
coro = self.msg.channel_ref.send(outbound)
|
||||||
|
|
||||||
|
self.sent_to_channel = True
|
||||||
asyncio.run_coroutine_threadsafe(coro, loop)
|
asyncio.run_coroutine_threadsafe(coro, loop)
|
||||||
|
|
||||||
async def wait_for_thread_pick(
|
async def wait_for_thread_pick(
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
|||||||
from . import (
|
from . import (
|
||||||
autoskills,
|
autoskills,
|
||||||
channel,
|
channel,
|
||||||
|
experts,
|
||||||
general,
|
general,
|
||||||
mcp,
|
mcp,
|
||||||
model,
|
model,
|
||||||
@@ -15,6 +16,7 @@ from . import (
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
"autoskills",
|
"autoskills",
|
||||||
"channel",
|
"channel",
|
||||||
|
"experts",
|
||||||
"general",
|
"general",
|
||||||
"mcp",
|
"mcp",
|
||||||
"model",
|
"model",
|
||||||
|
|||||||
@@ -0,0 +1,257 @@
|
|||||||
|
"""Slash commands for TUI expert-skill selection.
|
||||||
|
|
||||||
|
``/experts`` — list installed expert skills.
|
||||||
|
``/expert <name>`` — toggle an expert into the current session's
|
||||||
|
``active_teams`` list; the next turn's ``configurable.active_teams`` picks
|
||||||
|
this up and ``ActiveTeamMiddleware`` biases the main-agent's delegation
|
||||||
|
toward the invited expert(s).
|
||||||
|
``/expert clear`` — reset the list.
|
||||||
|
|
||||||
|
User-facing verbs match the WebUI gallery: **invite** to add an expert,
|
||||||
|
**dismiss** to remove one. Internal state field stays ``active_teams``
|
||||||
|
for wire compatibility.
|
||||||
|
|
||||||
|
Backing store is ``ChannelRuntime.active_teams`` (see
|
||||||
|
``EvoScientist/commands/base.py``). WebUI users get the same effect via
|
||||||
|
its gallery + langgraph-sdk ``config.configurable``; these commands are
|
||||||
|
the TUI-side equivalent.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING, ClassVar
|
||||||
|
|
||||||
|
from rich.table import Table
|
||||||
|
|
||||||
|
from ..base import Argument, Command, CommandContext, SubCommand
|
||||||
|
from ..manager import manager
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from ...tools.skills_manager import SkillInfo
|
||||||
|
|
||||||
|
_dispatchable_experts_cache: list[SkillInfo] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def invalidate_experts_cache() -> None:
|
||||||
|
"""Reset the /expert dispatchable-experts cache.
|
||||||
|
|
||||||
|
Called after ``install_skill`` / ``uninstall_skill`` mutations so a
|
||||||
|
freshly installed expert shows up in the /expert popup on the next
|
||||||
|
keystroke.
|
||||||
|
"""
|
||||||
|
global _dispatchable_experts_cache
|
||||||
|
_dispatchable_experts_cache = None
|
||||||
|
|
||||||
|
|
||||||
|
def _subscribe_cache_invalidation() -> None:
|
||||||
|
"""Register with ``skills_manager`` so every install/uninstall path
|
||||||
|
(slash commands, agent ``skill_manager`` @tool, onboarding) busts
|
||||||
|
the /expert popup — no caller has to remember.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from ...tools.skills_manager import register_skills_changed_callback
|
||||||
|
|
||||||
|
register_skills_changed_callback(invalidate_experts_cache)
|
||||||
|
except Exception:
|
||||||
|
# ``skills_manager`` not importable in some early-init contexts;
|
||||||
|
# cache staleness is a UX inconvenience, not a correctness bug.
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
_subscribe_cache_invalidation()
|
||||||
|
|
||||||
|
|
||||||
|
def _dispatchable_experts() -> list[SkillInfo]:
|
||||||
|
"""Cached list of experts that /expert can safely invite.
|
||||||
|
|
||||||
|
Filters ``list_expert_skills`` down to those that pass the same
|
||||||
|
empty-body + name-collision guards ``build_expert_subagent_specs``
|
||||||
|
and ``_fold_expert_subagents`` apply at agent-construction time, so
|
||||||
|
the /expert popup and invite-accept path only ever surface names
|
||||||
|
that will actually reach ``ActiveTeamMiddleware``'s cue.
|
||||||
|
"""
|
||||||
|
global _dispatchable_experts_cache
|
||||||
|
if _dispatchable_experts_cache is None:
|
||||||
|
try:
|
||||||
|
from ...subagents.expert_container import list_dispatchable_experts
|
||||||
|
|
||||||
|
_dispatchable_experts_cache = list_dispatchable_experts(include_system=True)
|
||||||
|
except Exception:
|
||||||
|
return []
|
||||||
|
return _dispatchable_experts_cache
|
||||||
|
|
||||||
|
|
||||||
|
class ExpertsCommand(Command):
|
||||||
|
"""List installed expert skills."""
|
||||||
|
|
||||||
|
name: ClassVar[str] = "/experts"
|
||||||
|
description: ClassVar[str] = "List installed expert skills"
|
||||||
|
category: ClassVar[str] = "Experts"
|
||||||
|
|
||||||
|
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||||
|
from ...tools.skills_manager import list_expert_skills
|
||||||
|
|
||||||
|
experts = list_expert_skills(include_system=True)
|
||||||
|
active = _current_active_teams(ctx)
|
||||||
|
|
||||||
|
if not experts:
|
||||||
|
ctx.ui.append_system("No expert skills installed.", style="dim")
|
||||||
|
ctx.ui.append_system(
|
||||||
|
"Install with: /install-skill <path-or-url>", style="dim"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
table = Table(title=f"Expert Skills ({len(experts)})", show_header=True)
|
||||||
|
table.add_column("Name", style="cyan")
|
||||||
|
table.add_column("Role", style="dim")
|
||||||
|
table.add_column("Active", style="green")
|
||||||
|
for skill in experts:
|
||||||
|
marker = "*" if skill.name in active else ""
|
||||||
|
table.add_row(
|
||||||
|
skill.name,
|
||||||
|
skill.role or skill.description,
|
||||||
|
marker,
|
||||||
|
)
|
||||||
|
ctx.ui.mount_renderable(table)
|
||||||
|
|
||||||
|
if active:
|
||||||
|
ctx.ui.append_system(
|
||||||
|
f"Active: {', '.join(active)}. Toggle with `/expert <name>`, "
|
||||||
|
"clear with `/expert clear`.",
|
||||||
|
style="dim",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
ctx.ui.append_system(
|
||||||
|
"No experts invited. `/expert <name>` to invite one.",
|
||||||
|
style="dim",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ExpertCommand(Command):
|
||||||
|
"""Invite, dismiss, or clear expert skills for the current thread."""
|
||||||
|
|
||||||
|
name: ClassVar[str] = "/expert"
|
||||||
|
description: ClassVar[str] = "Invite or dismiss an expert skill"
|
||||||
|
category: ClassVar[str] = "Experts"
|
||||||
|
arguments: ClassVar[list[Argument]] = [
|
||||||
|
Argument(
|
||||||
|
name="name_or_clear",
|
||||||
|
type=str,
|
||||||
|
description="Expert skill name to toggle, or 'clear' to reset",
|
||||||
|
required=True,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
subcommands: ClassVar[list[SubCommand]] = [
|
||||||
|
SubCommand("clear", "Dismiss all invited experts"),
|
||||||
|
]
|
||||||
|
|
||||||
|
def _get_expert_candidates(self) -> list[tuple[str, str]]:
|
||||||
|
return [(s.name, s.role or s.description) for s in _dispatchable_experts()]
|
||||||
|
|
||||||
|
def get_completions(self, tokens: list[str]) -> list[tuple[str, str]]:
|
||||||
|
"""Complete expert names + the ``clear`` subcommand."""
|
||||||
|
# /expert takes a single positional arg; anything past it (including a
|
||||||
|
# trailing space that turns tokens into ["name", ""]) has nothing to offer.
|
||||||
|
if len(tokens) > 1:
|
||||||
|
return []
|
||||||
|
prefix = tokens[0].lower() if tokens else ""
|
||||||
|
candidates = [
|
||||||
|
*self._get_expert_candidates(),
|
||||||
|
("clear", "Dismiss all invited experts"),
|
||||||
|
]
|
||||||
|
matches = [
|
||||||
|
(name, desc) for name, desc in candidates if name.lower().startswith(prefix)
|
||||||
|
]
|
||||||
|
# Exact match — argument already complete, hide the popup.
|
||||||
|
if len(matches) == 1 and matches[0][0].lower() == prefix:
|
||||||
|
return []
|
||||||
|
return matches
|
||||||
|
|
||||||
|
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||||
|
runtime = ctx.channel_runtime
|
||||||
|
if runtime is None:
|
||||||
|
ctx.ui.append_system(
|
||||||
|
"/expert requires a session runtime; not available in this context.",
|
||||||
|
style="yellow",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if not args:
|
||||||
|
ctx.ui.append_system(
|
||||||
|
"Usage: /expert <name> toggle an expert into the invited list",
|
||||||
|
style="yellow",
|
||||||
|
)
|
||||||
|
ctx.ui.append_system(
|
||||||
|
" /expert clear dismiss all invited experts",
|
||||||
|
style="dim",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
target = args[0].strip()
|
||||||
|
if target.lower() == "clear":
|
||||||
|
if not runtime.active_teams:
|
||||||
|
ctx.ui.append_system("No experts invited.", style="dim")
|
||||||
|
return
|
||||||
|
dismissed = list(runtime.active_teams)
|
||||||
|
runtime.active_teams = []
|
||||||
|
ctx.ui.append_system(
|
||||||
|
f"Dismissed experts: {', '.join(dismissed)}", style="dim"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Completion matches case-insensitively; honour the same here by
|
||||||
|
# resolving a case-variant to the on-disk name before membership.
|
||||||
|
by_lower = {s.name.lower(): s.name for s in _dispatchable_experts()}
|
||||||
|
canonical = by_lower.get(target.lower())
|
||||||
|
if canonical is None:
|
||||||
|
from ...tools.skills_manager import list_expert_skills
|
||||||
|
|
||||||
|
installed = {
|
||||||
|
s.name.lower() for s in list_expert_skills(include_system=True)
|
||||||
|
}
|
||||||
|
if target.lower() not in installed:
|
||||||
|
ctx.ui.append_system(
|
||||||
|
f"No expert skill named '{target}'. `/experts` lists "
|
||||||
|
"installed ones.",
|
||||||
|
style="red",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
ctx.ui.append_system(
|
||||||
|
f"Expert '{target}' can't be dispatched (empty actor "
|
||||||
|
"definition or name collision with a built-in sub-agent).",
|
||||||
|
style="red",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if canonical in runtime.active_teams:
|
||||||
|
runtime.active_teams = [n for n in runtime.active_teams if n != canonical]
|
||||||
|
ctx.ui.append_system(f"Dismissed expert: {canonical}", style="dim")
|
||||||
|
else:
|
||||||
|
runtime.active_teams = [*runtime.active_teams, canonical]
|
||||||
|
ctx.ui.append_system(f"Invited expert: {canonical}", style="green")
|
||||||
|
# An expert installed mid-session: the background reach
|
||||||
|
# (``start_async_task``) resolves it on first dispatch, but the
|
||||||
|
# in-turn ``task`` reach is frozen into the running agent, so it
|
||||||
|
# needs a rebuilt agent. An expert installed before this session
|
||||||
|
# started is already inside that frozen set — its in-turn reach
|
||||||
|
# works without a rebuild — so the hint scopes the /new boundary
|
||||||
|
# to newly installed experts instead of stating it
|
||||||
|
# unconditionally.
|
||||||
|
ctx.ui.append_system(
|
||||||
|
"Newly installed experts: background dispatch is available "
|
||||||
|
"immediately; in-turn task dispatch needs /new.",
|
||||||
|
style="dim",
|
||||||
|
)
|
||||||
|
if runtime.active_teams:
|
||||||
|
ctx.ui.append_system(
|
||||||
|
f"Active: {', '.join(runtime.active_teams)}", style="dim"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _current_active_teams(ctx: CommandContext) -> list[str]:
|
||||||
|
runtime = ctx.channel_runtime
|
||||||
|
return list(runtime.active_teams) if runtime is not None else []
|
||||||
|
|
||||||
|
|
||||||
|
manager.register(ExpertsCommand())
|
||||||
|
manager.register(ExpertCommand())
|
||||||
@@ -37,7 +37,7 @@ class InstallMCPCommand(Command):
|
|||||||
try:
|
try:
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
servers = await asyncio.get_event_loop().run_in_executor(
|
servers = await asyncio.get_running_loop().run_in_executor(
|
||||||
None, fetch_marketplace_index
|
None, fetch_marketplace_index
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -130,6 +130,7 @@ class ModelCommand(Command):
|
|||||||
*,
|
*,
|
||||||
save: bool = False,
|
save: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
import asyncio
|
||||||
import copy
|
import copy
|
||||||
|
|
||||||
from ...cli.agent import _load_agent
|
from ...cli.agent import _load_agent
|
||||||
@@ -139,6 +140,7 @@ class ModelCommand(Command):
|
|||||||
set_active_config,
|
set_active_config,
|
||||||
set_chat_model_instance,
|
set_chat_model_instance,
|
||||||
)
|
)
|
||||||
|
from ...runtime import AsyncRuntime
|
||||||
|
|
||||||
cfg = _ensure_config()
|
cfg = _ensure_config()
|
||||||
|
|
||||||
@@ -151,13 +153,26 @@ class ModelCommand(Command):
|
|||||||
temp_cfg.model = model_name
|
temp_cfg.model = model_name
|
||||||
temp_cfg.provider = provider
|
temp_cfg.provider = provider
|
||||||
|
|
||||||
|
# Re-thread the session's frontend event sink so the rebuilt agent's
|
||||||
|
# middleware keeps driving the tool-selection widget / fallback notices
|
||||||
|
# after a /model switch (the sink lives on the gateway, not the agent).
|
||||||
|
events = ctx.graph_gateway.events
|
||||||
|
|
||||||
try:
|
try:
|
||||||
new_chat_model = _build_chat_model(temp_cfg)
|
new_chat_model = _build_chat_model(temp_cfg)
|
||||||
new_agent = _load_agent(
|
load_kwargs = {
|
||||||
workspace_dir=ctx.workspace_dir,
|
"workspace_dir": ctx.workspace_dir,
|
||||||
checkpointer=ctx.checkpointer,
|
"checkpointer": ctx.checkpointer,
|
||||||
config=temp_cfg,
|
"config": temp_cfg,
|
||||||
chat_model=new_chat_model,
|
"chat_model": new_chat_model,
|
||||||
|
"events": events,
|
||||||
|
}
|
||||||
|
async_runtime = getattr(ctx, "async_runtime", None)
|
||||||
|
if isinstance(async_runtime, AsyncRuntime):
|
||||||
|
load_kwargs["runtime"] = async_runtime
|
||||||
|
new_agent = await asyncio.to_thread(
|
||||||
|
_load_agent,
|
||||||
|
**load_kwargs,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
ctx.ui.append_system(f"Failed to switch model: {e}", style="red")
|
ctx.ui.append_system(f"Failed to switch model: {e}", style="red")
|
||||||
|
|||||||
@@ -10,13 +10,21 @@ from ..base import Command, CommandContext, SubCommand
|
|||||||
from ..manager import manager
|
from ..manager import manager
|
||||||
|
|
||||||
|
|
||||||
|
def _clean(text: str) -> str:
|
||||||
|
"""Trim a shlex-joined argument and drop a stray wrapping quote pair."""
|
||||||
|
return text.strip().strip('"').strip("'")
|
||||||
|
|
||||||
|
|
||||||
class ScheduleCommand(Command):
|
class ScheduleCommand(Command):
|
||||||
"""Manage scheduled (cron) tasks."""
|
"""Manage scheduled (cron) tasks."""
|
||||||
|
|
||||||
name = "/schedule"
|
name = "/schedule"
|
||||||
description = "Manage scheduled (cron) tasks"
|
description = "Manage scheduled (cron) tasks"
|
||||||
subcommands: ClassVar[list[SubCommand]] = [
|
subcommands: ClassVar[list[SubCommand]] = [
|
||||||
SubCommand("add", 'Add: /schedule add <m h dom mon dow> "<prompt>"'),
|
SubCommand(
|
||||||
|
"add",
|
||||||
|
'Add: /schedule add <m h dom mon dow> "<prompt>" [--rubric "<checklist>"]',
|
||||||
|
),
|
||||||
SubCommand("list", "List scheduled tasks"),
|
SubCommand("list", "List scheduled tasks"),
|
||||||
SubCommand("remove", "Remove a schedule by id"),
|
SubCommand("remove", "Remove a schedule by id"),
|
||||||
SubCommand("run", "Run a schedule's prompt once now (test)"),
|
SubCommand("run", "Run a schedule's prompt once now (test)"),
|
||||||
@@ -78,10 +86,19 @@ class ScheduleCommand(Command):
|
|||||||
schedule, prompt_tokens = " ".join(rest[:5]), rest[5:]
|
schedule, prompt_tokens = " ".join(rest[:5]), rest[5:]
|
||||||
else:
|
else:
|
||||||
ctx.ui.append_system(
|
ctx.ui.append_system(
|
||||||
'Usage: /schedule add "<m h dom mon dow>" "<prompt>"', style="yellow"
|
'Usage: /schedule add "<m h dom mon dow>" "<prompt>" '
|
||||||
|
'[--rubric "<checklist>"]',
|
||||||
|
style="yellow",
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
prompt = " ".join(prompt_tokens).strip().strip('"').strip("'")
|
# Optional trailing acceptance checklist; everything after --rubric is it.
|
||||||
|
rubric = None
|
||||||
|
if "--rubric" in prompt_tokens:
|
||||||
|
# Last occurrence wins so an unquoted prompt may mention the flag.
|
||||||
|
split_at = len(prompt_tokens) - 1 - prompt_tokens[::-1].index("--rubric")
|
||||||
|
rubric = _clean(" ".join(prompt_tokens[split_at + 1 :])) or None
|
||||||
|
prompt_tokens = prompt_tokens[:split_at]
|
||||||
|
prompt = _clean(" ".join(prompt_tokens))
|
||||||
if not prompt:
|
if not prompt:
|
||||||
ctx.ui.append_system("A task prompt is required.", style="yellow")
|
ctx.ui.append_system("A task prompt is required.", style="yellow")
|
||||||
return
|
return
|
||||||
@@ -90,7 +107,11 @@ class ScheduleCommand(Command):
|
|||||||
name = re.sub(r"[^a-z0-9]+", "-", raw).strip("-")[:32] or "task"
|
name = re.sub(r"[^a-z0-9]+", "-", raw).strip("-")[:32] or "task"
|
||||||
try:
|
try:
|
||||||
rec = await asyncio.to_thread(
|
rec = await asyncio.to_thread(
|
||||||
crons.create_schedule, name=name, schedule=schedule, prompt=prompt
|
crons.create_schedule,
|
||||||
|
name=name,
|
||||||
|
schedule=schedule,
|
||||||
|
prompt=prompt,
|
||||||
|
rubric=rubric,
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
ctx.ui.append_system(f"Error: {exc}", style="red")
|
ctx.ui.append_system(f"Error: {exc}", style="red")
|
||||||
@@ -119,6 +140,7 @@ class ScheduleCommand(Command):
|
|||||||
table.add_column("Schedule", style="green")
|
table.add_column("Schedule", style="green")
|
||||||
table.add_column("Enabled", style="yellow")
|
table.add_column("Enabled", style="yellow")
|
||||||
table.add_column("Next run (UTC)", style="white")
|
table.add_column("Next run (UTC)", style="white")
|
||||||
|
table.add_column("Rubric", style="blue")
|
||||||
for r in rows:
|
for r in rows:
|
||||||
meta = r.get("metadata") or {}
|
meta = r.get("metadata") or {}
|
||||||
table.add_row(
|
table.add_row(
|
||||||
@@ -127,6 +149,7 @@ class ScheduleCommand(Command):
|
|||||||
str(r.get("schedule", "")),
|
str(r.get("schedule", "")),
|
||||||
"yes" if r.get("enabled", True) else "no",
|
"yes" if r.get("enabled", True) else "no",
|
||||||
str(r.get("next_run_date", "")),
|
str(r.get("next_run_date", "")),
|
||||||
|
"yes" if meta.get("rubric") else "",
|
||||||
)
|
)
|
||||||
ctx.ui.mount_renderable(table)
|
ctx.ui.mount_renderable(table)
|
||||||
|
|
||||||
@@ -191,7 +214,8 @@ class ScheduleCommand(Command):
|
|||||||
match = await self._resolve_or_report(ctx, crons, prefix)
|
match = await self._resolve_or_report(ctx, crons, prefix)
|
||||||
if match is None:
|
if match is None:
|
||||||
return
|
return
|
||||||
prompt = (match.get("metadata") or {}).get("prompt", "")
|
meta = match.get("metadata") or {}
|
||||||
|
prompt = meta.get("prompt", "")
|
||||||
if not str(prompt).strip():
|
if not str(prompt).strip():
|
||||||
ctx.ui.append_system(
|
ctx.ui.append_system(
|
||||||
f"Schedule {prefix} has no stored prompt — cannot run it.",
|
f"Schedule {prefix} has no stored prompt — cannot run it.",
|
||||||
@@ -199,7 +223,9 @@ class ScheduleCommand(Command):
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
rec = await asyncio.to_thread(crons.run_now, prompt)
|
rec = await asyncio.to_thread(
|
||||||
|
crons.run_now, prompt, rubric=meta.get("rubric") or None
|
||||||
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
ctx.ui.append_system(f"Error: {exc}", style="red")
|
ctx.ui.append_system(f"Error: {exc}", style="red")
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -182,8 +182,21 @@ class ResumeCommand(Command):
|
|||||||
if restored_workspace:
|
if restored_workspace:
|
||||||
ctx.workspace_dir = restored_workspace
|
ctx.workspace_dir = restored_workspace
|
||||||
|
|
||||||
|
switched_thread = resolved != ctx.thread_id
|
||||||
ctx.thread_id = resolved
|
ctx.thread_id = resolved
|
||||||
|
|
||||||
|
# Invitations are session-scoped (see ChannelRuntime.active_teams);
|
||||||
|
# resuming a different thread is a session switch, so release them —
|
||||||
|
# uniform with /new. Resuming the current thread keeps them.
|
||||||
|
runtime = ctx.channel_runtime
|
||||||
|
if switched_thread and runtime is not None and runtime.active_teams:
|
||||||
|
dismissed = list(runtime.active_teams)
|
||||||
|
runtime.active_teams = []
|
||||||
|
ctx.ui.append_system(
|
||||||
|
f"Dismissed experts on session switch: {', '.join(dismissed)}",
|
||||||
|
style="dim",
|
||||||
|
)
|
||||||
|
|
||||||
# Signal session change to UI
|
# Signal session change to UI
|
||||||
if hasattr(ctx.ui, "handle_session_resume"):
|
if hasattr(ctx.ui, "handle_session_resume"):
|
||||||
await ctx.ui.handle_session_resume(resolved, restored_workspace)
|
await ctx.ui.handle_session_resume(resolved, restored_workspace)
|
||||||
@@ -214,7 +227,23 @@ class NewCommand(Command):
|
|||||||
category = "Session"
|
category = "Session"
|
||||||
|
|
||||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||||
|
# ``/new`` means fresh state — release any invited experts. Uniform
|
||||||
|
# with the explicit ``/expert clear`` path; avoids
|
||||||
|
# the "why is idea-brainstorm still active in my new thread?"
|
||||||
|
# surprise. Users who want to reuse an invite in the next thread can
|
||||||
|
# re-invite explicitly. Cleared only after the new session actually
|
||||||
|
# exists, so a failed start leaves the current session intact.
|
||||||
|
runtime = ctx.channel_runtime
|
||||||
|
dismissed: list[str] = []
|
||||||
|
if runtime is not None and runtime.active_teams:
|
||||||
|
dismissed = list(runtime.active_teams)
|
||||||
await ctx.ui.start_new_session()
|
await ctx.ui.start_new_session()
|
||||||
|
if dismissed:
|
||||||
|
runtime.active_teams = []
|
||||||
|
ctx.ui.append_system(
|
||||||
|
f"Dismissed experts on new session: {', '.join(dismissed)}",
|
||||||
|
style="dim",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ClearCommand(Command):
|
class ClearCommand(Command):
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from __future__ import annotations
|
|||||||
import questionary
|
import questionary
|
||||||
from questionary import Choice
|
from questionary import Choice
|
||||||
|
|
||||||
|
from ...runtime import AsyncRuntime
|
||||||
from ..settings import EvoScientistConfig
|
from ..settings import EvoScientistConfig
|
||||||
from .helpers import (
|
from .helpers import (
|
||||||
_setup_imessage,
|
_setup_imessage,
|
||||||
@@ -21,7 +22,11 @@ from .style import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _step_channels(config: EvoScientistConfig) -> dict[str, object]:
|
def _step_channels(
|
||||||
|
config: EvoScientistConfig,
|
||||||
|
*,
|
||||||
|
runtime: AsyncRuntime | None = None,
|
||||||
|
) -> dict[str, object]:
|
||||||
"""Step: Select channels to enable on startup.
|
"""Step: Select channels to enable on startup.
|
||||||
|
|
||||||
Presents a multi-select list of supported channels.
|
Presents a multi-select list of supported channels.
|
||||||
@@ -35,6 +40,12 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]:
|
|||||||
Dict mapping config field names to their new values.
|
Dict mapping config field names to their new values.
|
||||||
Empty dict when the user skips or selects nothing.
|
Empty dict when the user skips or selects nothing.
|
||||||
"""
|
"""
|
||||||
|
# Direct/programmatic callers still get a single owned runtime for the
|
||||||
|
# whole step. CLI callers pass their application-scoped runtime instead.
|
||||||
|
if runtime is None:
|
||||||
|
with AsyncRuntime(thread_name="evosci-onboard-runtime") as owned_runtime:
|
||||||
|
return _step_channels(config, runtime=owned_runtime)
|
||||||
|
|
||||||
# Currently enabled channels
|
# Currently enabled channels
|
||||||
_currently_enabled = {
|
_currently_enabled = {
|
||||||
t.strip()
|
t.strip()
|
||||||
@@ -592,11 +603,9 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]:
|
|||||||
f" to {_accounts_path}.[/dim]"
|
f" to {_accounts_path}.[/dim]"
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
import asyncio
|
|
||||||
|
|
||||||
from ...channels.wechat.personal import qr_login
|
from ...channels.wechat.personal import qr_login
|
||||||
|
|
||||||
creds = asyncio.run(qr_login())
|
creds = runtime.run_sync(qr_login)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
console.print(f" [red]✗ Scan failed: {exc}[/red]")
|
console.print(f" [red]✗ Scan failed: {exc}[/red]")
|
||||||
creds = None
|
creds = None
|
||||||
@@ -783,7 +792,7 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]:
|
|||||||
updates[senders_field] = senders.strip()
|
updates[senders_field] = senders.strip()
|
||||||
|
|
||||||
# Probe validation
|
# Probe validation
|
||||||
_probe_channel(ch_name, config, updates)
|
_probe_channel(ch_name, config, updates, runtime=runtime)
|
||||||
|
|
||||||
enabled_channels.append(ch_name)
|
enabled_channels.append(ch_name)
|
||||||
|
|
||||||
@@ -820,12 +829,13 @@ def _probe_channel(
|
|||||||
ch_name: str,
|
ch_name: str,
|
||||||
config: EvoScientistConfig,
|
config: EvoScientistConfig,
|
||||||
updates: dict[str, object],
|
updates: dict[str, object],
|
||||||
|
*,
|
||||||
|
runtime: AsyncRuntime,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Run the probe for a channel type and print the result.
|
"""Run the probe for a channel type and print the result.
|
||||||
|
|
||||||
Non-fatal: prints a warning on failure but does not prevent enabling.
|
Non-fatal: prints a warning on failure but does not prevent enabling.
|
||||||
"""
|
"""
|
||||||
import asyncio
|
|
||||||
|
|
||||||
def _val(key: str, fallback: str = "") -> str:
|
def _val(key: str, fallback: str = "") -> str:
|
||||||
"""Get a value from updates first, then config, then fallback."""
|
"""Get a value from updates first, then config, then fallback."""
|
||||||
@@ -928,17 +938,7 @@ def _probe_channel(
|
|||||||
return True, "No probe available"
|
return True, "No probe available"
|
||||||
|
|
||||||
try:
|
try:
|
||||||
try:
|
ok, detail = runtime.run_sync(_run)
|
||||||
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:
|
if ok:
|
||||||
console.print(f" [green]✓ {detail}[/green]")
|
console.print(f" [green]✓ {detail}[/green]")
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ VALID_PROVIDERS: frozenset[str] = frozenset(
|
|||||||
"zhipu",
|
"zhipu",
|
||||||
"zhipu-code",
|
"zhipu-code",
|
||||||
"volcengine",
|
"volcengine",
|
||||||
|
"volcengine-code",
|
||||||
"dashscope",
|
"dashscope",
|
||||||
"dashscope-code",
|
"dashscope-code",
|
||||||
"deepseek",
|
"deepseek",
|
||||||
@@ -30,6 +31,9 @@ VALID_PROVIDERS: frozenset[str] = frozenset(
|
|||||||
"nvidia",
|
"nvidia",
|
||||||
"siliconflow",
|
"siliconflow",
|
||||||
"openrouter",
|
"openrouter",
|
||||||
|
"atlascloud",
|
||||||
|
"requesty",
|
||||||
|
"novita",
|
||||||
"custom-openai",
|
"custom-openai",
|
||||||
"custom-anthropic",
|
"custom-anthropic",
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from ..settings import EvoScientistConfig
|
|||||||
from .style import QMARK, WIZARD_STYLE, console
|
from .style import QMARK, WIZARD_STYLE, console
|
||||||
from .validators import (
|
from .validators import (
|
||||||
validate_anthropic_key,
|
validate_anthropic_key,
|
||||||
|
validate_atlascloud_key,
|
||||||
validate_dashscope_code_key,
|
validate_dashscope_code_key,
|
||||||
validate_dashscope_key,
|
validate_dashscope_key,
|
||||||
validate_deepseek_key,
|
validate_deepseek_key,
|
||||||
@@ -22,9 +23,11 @@ from .validators import (
|
|||||||
validate_kimi_key,
|
validate_kimi_key,
|
||||||
validate_minimax_key,
|
validate_minimax_key,
|
||||||
validate_moonshot_key,
|
validate_moonshot_key,
|
||||||
|
validate_novita_key,
|
||||||
validate_nvidia_key,
|
validate_nvidia_key,
|
||||||
validate_openai_key,
|
validate_openai_key,
|
||||||
validate_openrouter_key,
|
validate_openrouter_key,
|
||||||
|
validate_requesty_key,
|
||||||
validate_siliconflow_key,
|
validate_siliconflow_key,
|
||||||
validate_volcengine_key,
|
validate_volcengine_key,
|
||||||
validate_zhipu_key,
|
validate_zhipu_key,
|
||||||
@@ -70,6 +73,21 @@ def _provider_key_info(config: EvoScientistConfig, provider: str):
|
|||||||
config.openrouter_api_key or os.environ.get("OPENROUTER_API_KEY", ""),
|
config.openrouter_api_key or os.environ.get("OPENROUTER_API_KEY", ""),
|
||||||
validate_openrouter_key,
|
validate_openrouter_key,
|
||||||
),
|
),
|
||||||
|
"atlascloud": (
|
||||||
|
"Atlas Cloud",
|
||||||
|
config.atlascloud_api_key or os.environ.get("ATLASCLOUD_API_KEY", ""),
|
||||||
|
validate_atlascloud_key,
|
||||||
|
),
|
||||||
|
"requesty": (
|
||||||
|
"Requesty",
|
||||||
|
config.requesty_api_key or os.environ.get("REQUESTY_API_KEY", ""),
|
||||||
|
validate_requesty_key,
|
||||||
|
),
|
||||||
|
"novita": (
|
||||||
|
"Novita",
|
||||||
|
config.novita_api_key or os.environ.get("NOVITA_API_KEY", ""),
|
||||||
|
validate_novita_key,
|
||||||
|
),
|
||||||
"deepseek": (
|
"deepseek": (
|
||||||
"DeepSeek",
|
"DeepSeek",
|
||||||
config.deepseek_api_key or os.environ.get("DEEPSEEK_API_KEY", ""),
|
config.deepseek_api_key or os.environ.get("DEEPSEEK_API_KEY", ""),
|
||||||
@@ -90,6 +108,11 @@ def _provider_key_info(config: EvoScientistConfig, provider: str):
|
|||||||
config.volcengine_api_key or os.environ.get("VOLCENGINE_API_KEY", ""),
|
config.volcengine_api_key or os.environ.get("VOLCENGINE_API_KEY", ""),
|
||||||
validate_volcengine_key,
|
validate_volcengine_key,
|
||||||
),
|
),
|
||||||
|
"volcengine-code": (
|
||||||
|
"Volcengine Coding Plan",
|
||||||
|
config.volcengine_api_key or os.environ.get("VOLCENGINE_API_KEY", ""),
|
||||||
|
validate_volcengine_key,
|
||||||
|
),
|
||||||
"dashscope": (
|
"dashscope": (
|
||||||
"DashScope",
|
"DashScope",
|
||||||
config.dashscope_api_key or os.environ.get("DASHSCOPE_API_KEY", ""),
|
config.dashscope_api_key or os.environ.get("DASHSCOPE_API_KEY", ""),
|
||||||
|
|||||||
@@ -156,8 +156,14 @@ def _step_langgraph_dev_port(config: EvoScientistConfig) -> int:
|
|||||||
f"EvoSci config set langgraph_dev_port <other-port>[/yellow]"
|
f"EvoSci config set langgraph_dev_port <other-port>[/yellow]"
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
# Render the address the configured bind actually produces rather than
|
||||||
|
# a hard-coded loopback URL — the two diverge once langgraph_dev_host
|
||||||
|
# is pinned to a specific interface.
|
||||||
|
from ...langgraph_dev.manager import _base_url
|
||||||
|
|
||||||
|
host = getattr(config, "langgraph_dev_host", "")
|
||||||
console.print(
|
console.print(
|
||||||
f" [green]✓ EvoScientist will run on http://127.0.0.1:{port}[/green]"
|
f" [green]✓ EvoScientist will run on {_base_url(port, host)}[/green]"
|
||||||
)
|
)
|
||||||
return port
|
return port
|
||||||
|
|
||||||
@@ -220,7 +226,15 @@ def _step_webui_port(config: EvoScientistConfig) -> int:
|
|||||||
raise KeyboardInterrupt()
|
raise KeyboardInterrupt()
|
||||||
|
|
||||||
port = int(raw) if raw else current_port
|
port = int(raw) if raw else current_port
|
||||||
console.print(f" [green]✓ WebUI will open at http://localhost:{port}[/green]")
|
# Same reasoning as the langgraph-dev step: render the configured bind, not
|
||||||
|
# a hard-coded localhost. A wildcard bind still shows loopback here — that
|
||||||
|
# is the address this machine's own browser opens.
|
||||||
|
from ...langgraph_dev.manager import _format_hostport
|
||||||
|
|
||||||
|
host = getattr(config, "webui_host", "")
|
||||||
|
console.print(
|
||||||
|
f" [green]✓ WebUI will open at http://{_format_hostport(host, port)}[/green]"
|
||||||
|
)
|
||||||
console.print(
|
console.print(
|
||||||
" [yellow]⚠️ Note: the WebUI won't show your CLI/TUI chat history "
|
" [yellow]⚠️ Note: the WebUI won't show your CLI/TUI chat history "
|
||||||
"yet.[/yellow]"
|
"yet.[/yellow]"
|
||||||
@@ -264,6 +278,10 @@ def _step_provider(
|
|||||||
title="Volcengine (火山引擎 — Doubao models)",
|
title="Volcengine (火山引擎 — Doubao models)",
|
||||||
value="volcengine",
|
value="volcengine",
|
||||||
),
|
),
|
||||||
|
Choice(
|
||||||
|
title="Volcengine Coding Plan (火山引擎代码计划 — coding models)",
|
||||||
|
value="volcengine-code",
|
||||||
|
),
|
||||||
Choice(
|
Choice(
|
||||||
title="DashScope (阿里云 — Qwen models)",
|
title="DashScope (阿里云 — Qwen models)",
|
||||||
value="dashscope",
|
value="dashscope",
|
||||||
@@ -296,6 +314,18 @@ def _step_provider(
|
|||||||
title="OpenRouter (aggregator — Grok, Gemini, Qwen, etc.)",
|
title="OpenRouter (aggregator — Grok, Gemini, Qwen, etc.)",
|
||||||
value="openrouter",
|
value="openrouter",
|
||||||
),
|
),
|
||||||
|
Choice(
|
||||||
|
title="Atlas Cloud (aggregator — DeepSeek, Qwen, etc.)",
|
||||||
|
value="atlascloud",
|
||||||
|
),
|
||||||
|
Choice(
|
||||||
|
title="Requesty (aggregator — OpenAI, Anthropic, Gemini, xAI, etc.)",
|
||||||
|
value="requesty",
|
||||||
|
),
|
||||||
|
Choice(
|
||||||
|
title="Novita (aggregator — DeepSeek, Qwen, GLM, etc.)",
|
||||||
|
value="novita",
|
||||||
|
),
|
||||||
Choice(
|
Choice(
|
||||||
title="OpenAI-compatible (third-party OpenAI endpoint)",
|
title="OpenAI-compatible (third-party OpenAI endpoint)",
|
||||||
value="custom-openai",
|
value="custom-openai",
|
||||||
@@ -981,6 +1011,11 @@ _RECOMMENDED_SKILLS = [
|
|||||||
"label": "HuggingFace Skills (dataset creation, model training & evaluation, third party by HuggingFace)",
|
"label": "HuggingFace Skills (dataset creation, model training & evaluation, third party by HuggingFace)",
|
||||||
"source": "huggingface/skills@skills",
|
"source": "huggingface/skills@skills",
|
||||||
},
|
},
|
||||||
|
# ── Third-party (NVIDIA BioNeMo) ──
|
||||||
|
{
|
||||||
|
"label": "BioNeMo Skills (31 protein folding, docking, generative chemistry & genomics skills, third party by NVIDIA)",
|
||||||
|
"source": "NVIDIA-BioNeMo/bionemo-agent-toolkit@plugins/bionemo-agent-toolkit/skills",
|
||||||
|
},
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -321,6 +321,159 @@ def validate_openrouter_key(api_key: str) -> tuple[bool, str]:
|
|||||||
return False, f"Error: {e}"
|
return False, f"Error: {e}"
|
||||||
|
|
||||||
|
|
||||||
|
def validate_atlascloud_key(api_key: str) -> tuple[bool, str]:
|
||||||
|
"""Validate an Atlas Cloud key with a nonexistent sentinel model.
|
||||||
|
|
||||||
|
The probe deliberately targets a nonexistent sentinel model. A 404 means
|
||||||
|
authentication passed and model resolution failed; 200 also confirms
|
||||||
|
authentication if the sentinel unexpectedly resolves. A 401/403 means the
|
||||||
|
key was rejected. Other statuses remain inconclusive until verified.
|
||||||
|
"""
|
||||||
|
if not api_key:
|
||||||
|
return True, "Skipped (no key provided)"
|
||||||
|
|
||||||
|
try:
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
resp = httpx.post(
|
||||||
|
"https://api.atlascloud.ai/v1/chat/completions",
|
||||||
|
headers={
|
||||||
|
"Authorization": f"Bearer {api_key}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
},
|
||||||
|
json={
|
||||||
|
"model": "atlascloud/auth-preflight",
|
||||||
|
"messages": [{"role": "user", "content": "ping"}],
|
||||||
|
"max_tokens": 1,
|
||||||
|
},
|
||||||
|
timeout=10,
|
||||||
|
)
|
||||||
|
if resp.status_code in (200, 404):
|
||||||
|
return True, "Valid"
|
||||||
|
# Atlas checks account balance before model resolution: a valid key
|
||||||
|
# on an uncredited account gets 402 from the sentinel probe.
|
||||||
|
if resp.status_code == 402:
|
||||||
|
return True, "Valid (insufficient balance — top up to use)"
|
||||||
|
if resp.status_code in (401, 403):
|
||||||
|
return False, "Invalid API key"
|
||||||
|
return False, f"Validation inconclusive (HTTP {resp.status_code})"
|
||||||
|
except Exception as e:
|
||||||
|
classified = _classify_validation_error(e)
|
||||||
|
if classified is not None:
|
||||||
|
return classified
|
||||||
|
return False, f"Error: {e}"
|
||||||
|
|
||||||
|
|
||||||
|
def validate_requesty_key(api_key: str) -> tuple[bool, str]:
|
||||||
|
"""Validate a Requesty API key against the router's auth layer.
|
||||||
|
|
||||||
|
Unlike OpenRouter, Requesty's ``/v1/models`` endpoint returns HTTP 200
|
||||||
|
(the public model catalog) even for a missing or invalid key, so it
|
||||||
|
cannot be used to check a key. We instead issue a minimal
|
||||||
|
``/v1/chat/completions`` request, but deliberately target a nonexistent
|
||||||
|
sentinel model: the router checks auth *before* resolving the model, so
|
||||||
|
the response distinguishes the two failures without depending on any
|
||||||
|
real model staying available upstream.
|
||||||
|
|
||||||
|
- valid key → 404 ("Model and/or policy not supported"), i.e. auth passed
|
||||||
|
(or 200 in the unlikely event the sentinel ever resolves);
|
||||||
|
- invalid/missing key → 401/403 ("Invalid authorization token");
|
||||||
|
- 429 (rate-limit) / 5xx (router incident) leave validity unknown, so a
|
||||||
|
transient outage doesn't reject a good key.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (is_valid, message).
|
||||||
|
"""
|
||||||
|
if not api_key:
|
||||||
|
return True, "Skipped (no key provided)"
|
||||||
|
|
||||||
|
try:
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
resp = httpx.post(
|
||||||
|
"https://router.requesty.ai/v1/chat/completions",
|
||||||
|
headers={
|
||||||
|
"Authorization": f"Bearer {api_key}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
},
|
||||||
|
json={
|
||||||
|
# Deliberately nonexistent sentinel: auth is resolved before
|
||||||
|
# the model, so a valid key gets a 404 (model-not-found)
|
||||||
|
# rather than depending on a specific model being available.
|
||||||
|
"model": "requesty/auth-preflight",
|
||||||
|
"messages": [{"role": "user", "content": "ping"}],
|
||||||
|
"max_tokens": 1,
|
||||||
|
},
|
||||||
|
timeout=10,
|
||||||
|
)
|
||||||
|
# 200 (accepted) or 404 (auth passed, model not found) → key is good.
|
||||||
|
if resp.status_code in (200, 404):
|
||||||
|
return True, "Valid"
|
||||||
|
# Only 401/403 mean the key is actually rejected. 429 (rate-limit)
|
||||||
|
# and 5xx (router incident) leave the key validity unknown — surface
|
||||||
|
# the real status so the user doesn't go re-roll a good key during
|
||||||
|
# an outage.
|
||||||
|
if resp.status_code in (401, 403):
|
||||||
|
return False, "Invalid API key"
|
||||||
|
return False, f"Validation inconclusive (HTTP {resp.status_code})"
|
||||||
|
except Exception as e:
|
||||||
|
classified = _classify_validation_error(e)
|
||||||
|
if classified is not None:
|
||||||
|
return classified
|
||||||
|
return False, f"Error: {e}"
|
||||||
|
|
||||||
|
|
||||||
|
def validate_novita_key(api_key: str) -> tuple[bool, str]:
|
||||||
|
"""Validate a Novita API key against the router's auth layer.
|
||||||
|
|
||||||
|
Like Requesty and Atlas Cloud, Novita's ``/v1/models`` endpoint returns
|
||||||
|
HTTP 200 (the public model catalog) even for a missing or invalid key, so
|
||||||
|
it cannot be used to check a key (verified against the live endpoint). We
|
||||||
|
instead issue a minimal ``/v1/chat/completions`` request with a
|
||||||
|
deliberately nonexistent sentinel model: auth is resolved before the
|
||||||
|
model, so a valid key doesn't depend on any real model staying available
|
||||||
|
upstream.
|
||||||
|
|
||||||
|
- invalid/missing key → 401/403 (confirmed against the live endpoint);
|
||||||
|
- valid key → 200 or 404 (model-not-found, auth passed), mirroring the
|
||||||
|
Requesty/Atlas Cloud sentinel pattern;
|
||||||
|
- 429 (rate-limit) / 5xx (service incident) leave validity unknown, so a
|
||||||
|
transient outage doesn't reject a good key.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (is_valid, message).
|
||||||
|
"""
|
||||||
|
if not api_key:
|
||||||
|
return True, "Skipped (no key provided)"
|
||||||
|
|
||||||
|
try:
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
resp = httpx.post(
|
||||||
|
"https://api.novita.ai/openai/v1/chat/completions",
|
||||||
|
headers={
|
||||||
|
"Authorization": f"Bearer {api_key}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
},
|
||||||
|
json={
|
||||||
|
"model": "novita/auth-preflight",
|
||||||
|
"messages": [{"role": "user", "content": "ping"}],
|
||||||
|
"max_tokens": 1,
|
||||||
|
},
|
||||||
|
timeout=10,
|
||||||
|
)
|
||||||
|
if resp.status_code in (200, 404):
|
||||||
|
return True, "Valid"
|
||||||
|
if resp.status_code in (401, 403):
|
||||||
|
return False, "Invalid API key"
|
||||||
|
return False, f"Validation inconclusive (HTTP {resp.status_code})"
|
||||||
|
except Exception as e:
|
||||||
|
classified = _classify_validation_error(e)
|
||||||
|
if classified is not None:
|
||||||
|
return classified
|
||||||
|
return False, f"Error: {e}"
|
||||||
|
|
||||||
|
|
||||||
def validate_deepseek_key(api_key: str) -> tuple[bool, str]:
|
def validate_deepseek_key(api_key: str) -> tuple[bool, str]:
|
||||||
"""Validate a DeepSeek API key by making a test request.
|
"""Validate a DeepSeek API key by making a test request.
|
||||||
|
|
||||||
@@ -373,6 +526,9 @@ def validate_zhipu_key(api_key: str) -> tuple[bool, str]:
|
|||||||
def validate_volcengine_key(api_key: str) -> tuple[bool, str]:
|
def validate_volcengine_key(api_key: str) -> tuple[bool, str]:
|
||||||
"""Validate a Volcengine API key by making a test request.
|
"""Validate a Volcengine API key by making a test request.
|
||||||
|
|
||||||
|
Uses the general endpoint for validation; volcengine and volcengine-code
|
||||||
|
share the same API key and only differ in their runtime base URL.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple of (is_valid, message).
|
Tuple of (is_valid, message).
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import questionary
|
|||||||
from rich.panel import Panel
|
from rich.panel import Panel
|
||||||
from rich.text import Text
|
from rich.text import Text
|
||||||
|
|
||||||
|
from ...runtime import AsyncRuntime
|
||||||
from ..settings import (
|
from ..settings import (
|
||||||
EvoScientistConfig,
|
EvoScientistConfig,
|
||||||
get_config_path,
|
get_config_path,
|
||||||
@@ -117,10 +118,14 @@ _PROVIDER_KEY_ATTR = {
|
|||||||
"google-genai": "google_api_key",
|
"google-genai": "google_api_key",
|
||||||
"siliconflow": "siliconflow_api_key",
|
"siliconflow": "siliconflow_api_key",
|
||||||
"openrouter": "openrouter_api_key",
|
"openrouter": "openrouter_api_key",
|
||||||
|
"atlascloud": "atlascloud_api_key",
|
||||||
|
"requesty": "requesty_api_key",
|
||||||
|
"novita": "novita_api_key",
|
||||||
"deepseek": "deepseek_api_key",
|
"deepseek": "deepseek_api_key",
|
||||||
"zhipu": "zhipu_api_key",
|
"zhipu": "zhipu_api_key",
|
||||||
"zhipu-code": "zhipu_api_key",
|
"zhipu-code": "zhipu_api_key",
|
||||||
"volcengine": "volcengine_api_key",
|
"volcengine": "volcengine_api_key",
|
||||||
|
"volcengine-code": "volcengine_api_key",
|
||||||
"dashscope": "dashscope_api_key",
|
"dashscope": "dashscope_api_key",
|
||||||
"dashscope-code": "dashscope_api_key",
|
"dashscope-code": "dashscope_api_key",
|
||||||
"moonshot": "moonshot_api_key",
|
"moonshot": "moonshot_api_key",
|
||||||
@@ -475,6 +480,7 @@ def run_onboard(
|
|||||||
skip_validation: bool = False,
|
skip_validation: bool = False,
|
||||||
prompter=None,
|
prompter=None,
|
||||||
only_sections: set[str] | frozenset[str] | None = None,
|
only_sections: set[str] | frozenset[str] | None = None,
|
||||||
|
runtime: AsyncRuntime | None = None,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Run the interactive onboarding wizard.
|
"""Run the interactive onboarding wizard.
|
||||||
|
|
||||||
@@ -487,6 +493,9 @@ def run_onboard(
|
|||||||
only_sections: If given, restrict the wizard to exactly these section
|
only_sections: If given, restrict the wizard to exactly these section
|
||||||
ids — the Keep/Modify/Reset prompt is skipped. Used by ``EvoSci
|
ids — the Keep/Modify/Reset prompt is skipped. Used by ``EvoSci
|
||||||
configure <section>`` to re-run a single phase.
|
configure <section>`` to re-run a single phase.
|
||||||
|
runtime: Optional application-scoped async runtime used by channel
|
||||||
|
login and credential probes. Direct callers may omit it; the
|
||||||
|
channel step then owns a runtime for the duration of that step.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
True if configuration was saved, False if cancelled.
|
True if configuration was saved, False if cancelled.
|
||||||
@@ -883,7 +892,7 @@ def run_onboard(
|
|||||||
_step_tinytex()
|
_step_tinytex()
|
||||||
|
|
||||||
if "channels" in sections_to_run:
|
if "channels" in sections_to_run:
|
||||||
for key, value in _step_channels(config).items():
|
for key, value in _step_channels(config, runtime=runtime).items():
|
||||||
setattr(config, key, value)
|
setattr(config, key, value)
|
||||||
_autosave(config)
|
_autosave(config)
|
||||||
|
|
||||||
|
|||||||
+148
-15
@@ -1,8 +1,9 @@
|
|||||||
"""Configuration management for EvoScientist.
|
"""Configuration management for EvoScientist.
|
||||||
|
|
||||||
Handles loading, saving, and merging configuration from multiple sources
|
Handles loading, saving, and merging configuration from multiple sources.
|
||||||
with the following priority (highest to lowest):
|
See :func:`get_effective_config` for the authoritative priority chain —
|
||||||
CLI arguments > Environment variables > Config file > Defaults
|
``EVOSCIENTIST_*`` shell values and third-party keys are treated
|
||||||
|
asymmetrically with respect to workspace ``.env`` handling.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -16,14 +17,18 @@ from pathlib import Path
|
|||||||
from typing import Any, Literal, get_type_hints
|
from typing import Any, Literal, get_type_hints
|
||||||
|
|
||||||
import yaml
|
import yaml
|
||||||
from dotenv import find_dotenv, load_dotenv
|
from dotenv import dotenv_values, find_dotenv
|
||||||
|
|
||||||
# Tools that run shell commands and need manual HITL approval (subject to
|
# Tools that run shell commands and need manual HITL approval (subject to
|
||||||
# shell_allow_list). Single source of truth for every interrupt consumer
|
# shell_allow_list). Single source of truth for every interrupt consumer
|
||||||
# (stream/display.py, channels/consumer.py) — keep aligned with the agent's
|
# (stream/display.py, channels/interaction.py) — keep aligned with the agent's
|
||||||
# `interrupt_on` set in EvoScientist.py.
|
# `interrupt_on` set in EvoScientist.py.
|
||||||
HITL_SHELL_TOOLS = ("execute", "run_in_background")
|
HITL_SHELL_TOOLS = ("execute", "run_in_background")
|
||||||
|
|
||||||
|
# Armed non-shell destructive tools must always prompt — no allow-list carve-outs
|
||||||
|
# (their args carry paths, not commands). Keep aligned with HITL_INTERRUPT_ON.
|
||||||
|
HITL_ALWAYS_PROMPT_TOOLS = ("delete", "schedule_task")
|
||||||
|
|
||||||
|
|
||||||
class MemoryObservationTarget(StrEnum):
|
class MemoryObservationTarget(StrEnum):
|
||||||
"""Runtime locations that can receive `record_observation`."""
|
"""Runtime locations that can receive `record_observation`."""
|
||||||
@@ -178,6 +183,9 @@ class EvoScientistConfig:
|
|||||||
minimax_base_url: str = ""
|
minimax_base_url: str = ""
|
||||||
siliconflow_api_key: str = ""
|
siliconflow_api_key: str = ""
|
||||||
openrouter_api_key: str = ""
|
openrouter_api_key: str = ""
|
||||||
|
atlascloud_api_key: str = ""
|
||||||
|
requesty_api_key: str = ""
|
||||||
|
novita_api_key: str = ""
|
||||||
deepseek_api_key: str = ""
|
deepseek_api_key: str = ""
|
||||||
zhipu_api_key: str = ""
|
zhipu_api_key: str = ""
|
||||||
volcengine_api_key: str = ""
|
volcengine_api_key: str = ""
|
||||||
@@ -218,11 +226,24 @@ class EvoScientistConfig:
|
|||||||
# the Ai4Sci-Web Gateway's recoverable runtime URL.
|
# the Ai4Sci-Web Gateway's recoverable runtime URL.
|
||||||
langgraph_dev_port: int = 3076
|
langgraph_dev_port: int = 3076
|
||||||
|
|
||||||
|
# Network interface the langgraph dev subprocess binds to. Loopback by
|
||||||
|
# default — this is the unauthenticated agent API (the agent can run
|
||||||
|
# shell), so "0.0.0.0" is opt-in and every launcher prints a PUBLIC BIND
|
||||||
|
# banner while exposed. Internal callers *connect* via manager._probe_host,
|
||||||
|
# so widening never redirects their traffic off-box.
|
||||||
|
langgraph_dev_host: str = "127.0.0.1"
|
||||||
|
|
||||||
# Port for the WebUI front-end (Next.js server from @evoscientist/webui),
|
# Port for the WebUI front-end (Next.js server from @evoscientist/webui),
|
||||||
# used only when ui_backend == "webui". The backend keeps
|
# used only when ui_backend == "webui". The backend keeps
|
||||||
# its own port (langgraph_dev_port); this is just the browser server.
|
# its own port (langgraph_dev_port); this is just the browser server.
|
||||||
webui_port: int = 4716
|
webui_port: int = 4716
|
||||||
|
|
||||||
|
# Network interface the WebUI front-end binds to. Loopback by default,
|
||||||
|
# matching langgraph_dev_host: this server is not a passive app shell —
|
||||||
|
# its API reads, writes and uploads workspace files and installs skills,
|
||||||
|
# all unauthenticated. Set "0.0.0.0" (with langgraph_dev_host) for LAN.
|
||||||
|
webui_host: str = "127.0.0.1"
|
||||||
|
|
||||||
# --- Scheduled tasks (cron) ---
|
# --- Scheduled tasks (cron) ---
|
||||||
# Master switch for scheduled tasks (/schedule, NL tools, scheduler context). Defaults
|
# Master switch for scheduled tasks (/schedule, NL tools, scheduler context). Defaults
|
||||||
# True so the feature is available out-of-the-box; set False to disable.
|
# True so the feature is available out-of-the-box; set False to disable.
|
||||||
@@ -248,6 +269,15 @@ class EvoScientistConfig:
|
|||||||
# slowdown.
|
# slowdown.
|
||||||
langgraph_dev_jobs_per_worker: int = 10
|
langgraph_dev_jobs_per_worker: int = 10
|
||||||
|
|
||||||
|
# Keep the auto-started langgraph dev subprocess running after the CLI
|
||||||
|
# exits. The next `EvoSci` start in the same workspace reuses it instantly
|
||||||
|
# instead of paying the cold boot (~15s). Starting in a DIFFERENT workspace
|
||||||
|
# raises WorkspaceMismatchError with the leftover server's pid — stop it
|
||||||
|
# manually (the server is pinned to one workspace per process). Known
|
||||||
|
# limitation: changing langgraph_dev_port/host while a keepalive server
|
||||||
|
# runs orphans its records — run `EvoSci server stop` before switching.
|
||||||
|
langgraph_dev_keepalive: bool = False
|
||||||
|
|
||||||
# Max LangGraph super-steps (LLM call / tool call / sub-agent delegation
|
# Max LangGraph super-steps (LLM call / tool call / sub-agent delegation
|
||||||
# each count as 1) before raising GraphRecursionError. Resets on every
|
# each count as 1) before raising GraphRecursionError. Resets on every
|
||||||
# ``agent.invoke()`` — i.e., this is per-turn, NOT per-conversation. For
|
# ``agent.invoke()`` — i.e., this is per-turn, NOT per-conversation. For
|
||||||
@@ -294,6 +324,14 @@ class EvoScientistConfig:
|
|||||||
DEFAULT_MEMORY_SKILL_SYNTHESIS_CADENCE
|
DEFAULT_MEMORY_SKILL_SYNTHESIS_CADENCE
|
||||||
)
|
)
|
||||||
memory_skill_synthesis_time: str = DEFAULT_MEMORY_SKILL_SYNTHESIS_TIME
|
memory_skill_synthesis_time: str = DEFAULT_MEMORY_SKILL_SYNTHESIS_TIME
|
||||||
|
# Max number of parsed observation files kept in the process-wide parse
|
||||||
|
# cache. Each entry holds one parsed document keyed on the file path; at
|
||||||
|
# the end of a call the LRU trims down to max(cap, entries touched by
|
||||||
|
# the call), so an active store larger than the cap temporarily exceeds
|
||||||
|
# it instead of thrashing. 2048 is generous for the single-workspace
|
||||||
|
# deploy model; raise for a long-running server that cycles through many
|
||||||
|
# large workspaces.
|
||||||
|
memory_observation_cache_max_files: int = 2048
|
||||||
|
|
||||||
# Workspace Settings
|
# Workspace Settings
|
||||||
default_mode: Literal["daemon", "run"] = "daemon"
|
default_mode: Literal["daemon", "run"] = "daemon"
|
||||||
@@ -310,7 +348,9 @@ class EvoScientistConfig:
|
|||||||
openrouter_anthropic_prompt_cache: bool = True
|
openrouter_anthropic_prompt_cache: bool = True
|
||||||
# OpenRouter app attribution (issue #339). Sent only for the openrouter
|
# OpenRouter app attribution (issue #339). Sent only for the openrouter
|
||||||
# provider; identifies EvoScientist in OpenRouter's app rankings/analytics.
|
# provider; identifies EvoScientist in OpenRouter's app rankings/analytics.
|
||||||
# Override (e.g. a private fork) via these fields or their env vars.
|
# Override (e.g. a private fork) via these fields or their env vars. A custom
|
||||||
|
# title only takes effect together with a custom referer: OpenRouter keys app
|
||||||
|
# pages by referer, so a lone title would rename the shared EvoScientist page.
|
||||||
# Defaults live in the module constants above (also imported by llm/models.py).
|
# Defaults live in the module constants above (also imported by llm/models.py).
|
||||||
openrouter_http_referer: str = OPENROUTER_DEFAULT_HTTP_REFERER
|
openrouter_http_referer: str = OPENROUTER_DEFAULT_HTTP_REFERER
|
||||||
openrouter_app_title: str = OPENROUTER_DEFAULT_APP_TITLE
|
openrouter_app_title: str = OPENROUTER_DEFAULT_APP_TITLE
|
||||||
@@ -461,6 +501,9 @@ class EvoScientistConfig:
|
|||||||
# DM access control policy
|
# DM access control policy
|
||||||
dm_policy: str = "allowlist"
|
dm_policy: str = "allowlist"
|
||||||
|
|
||||||
|
# OpenAI API mode - "" = auto, "true" = force Responses, "false" = force Completions
|
||||||
|
use_responses_api: str = ""
|
||||||
|
|
||||||
# ccproxy
|
# ccproxy
|
||||||
ccproxy_port: int = 8000
|
ccproxy_port: int = 8000
|
||||||
|
|
||||||
@@ -492,14 +535,36 @@ class EvoScientistConfig:
|
|||||||
)
|
)
|
||||||
self.sandbox_execute_timeout = 300
|
self.sandbox_execute_timeout = 300
|
||||||
|
|
||||||
# Dangerous mode implies auto_approve regardless of source (CLI, env,
|
# A non-positive cache cap would evict every file entry immediately,
|
||||||
# config file). Mirrors how auto_mode implies auto_approve — done here so
|
# defeating the cache entirely.
|
||||||
# the coupling holds even when dangerous_mode is set via `config set`.
|
cap = self.memory_observation_cache_max_files
|
||||||
if self.dangerous_mode:
|
if not isinstance(cap, int) or isinstance(cap, bool) or cap < 1:
|
||||||
|
logging.getLogger(__name__).warning(
|
||||||
|
"Invalid memory_observation_cache_max_files %r; falling back to 2048.",
|
||||||
|
cap,
|
||||||
|
)
|
||||||
|
self.memory_observation_cache_max_files = 2048
|
||||||
|
|
||||||
|
# auto_mode and dangerous_mode both imply auto_approve regardless of
|
||||||
|
# source (CLI, env, config file, direct construction) — done here so the
|
||||||
|
# "unattended → zero prompts" contract holds even when either is set via
|
||||||
|
# `config set` or a config file rather than a CLI flag.
|
||||||
|
if self.auto_mode or self.dangerous_mode:
|
||||||
self.auto_approve = True
|
self.auto_approve = True
|
||||||
|
|
||||||
_normalize_str_enum_fields(self)
|
_normalize_str_enum_fields(self)
|
||||||
|
|
||||||
|
# Bind hosts reach socket.bind() / the langgraph CLI verbatim, where a
|
||||||
|
# stray-whitespace or empty value surfaces as an opaque gaierror at
|
||||||
|
# startup. Normalize to the field's own default instead.
|
||||||
|
for _host_field, _host_default in (
|
||||||
|
("langgraph_dev_host", "127.0.0.1"),
|
||||||
|
("webui_host", "127.0.0.1"),
|
||||||
|
):
|
||||||
|
_host = getattr(self, _host_field, _host_default)
|
||||||
|
_host = _host.strip() if isinstance(_host, str) else ""
|
||||||
|
setattr(self, _host_field, _host or _host_default)
|
||||||
|
|
||||||
synthesis_time = _normalize_hhmm(self.memory_skill_synthesis_time)
|
synthesis_time = _normalize_hhmm(self.memory_skill_synthesis_time)
|
||||||
if synthesis_time is None:
|
if synthesis_time is None:
|
||||||
logging.getLogger(__name__).warning(
|
logging.getLogger(__name__).warning(
|
||||||
@@ -783,6 +848,9 @@ _ENV_MAPPINGS = {
|
|||||||
"minimax_base_url": "MINIMAX_BASE_URL",
|
"minimax_base_url": "MINIMAX_BASE_URL",
|
||||||
"siliconflow_api_key": "SILICONFLOW_API_KEY",
|
"siliconflow_api_key": "SILICONFLOW_API_KEY",
|
||||||
"openrouter_api_key": "OPENROUTER_API_KEY",
|
"openrouter_api_key": "OPENROUTER_API_KEY",
|
||||||
|
"atlascloud_api_key": "ATLASCLOUD_API_KEY",
|
||||||
|
"requesty_api_key": "REQUESTY_API_KEY",
|
||||||
|
"novita_api_key": "NOVITA_API_KEY",
|
||||||
"deepseek_api_key": "DEEPSEEK_API_KEY",
|
"deepseek_api_key": "DEEPSEEK_API_KEY",
|
||||||
"zhipu_api_key": "ZHIPU_API_KEY",
|
"zhipu_api_key": "ZHIPU_API_KEY",
|
||||||
"volcengine_api_key": "VOLCENGINE_API_KEY",
|
"volcengine_api_key": "VOLCENGINE_API_KEY",
|
||||||
@@ -812,10 +880,13 @@ _ENV_MAPPINGS = {
|
|||||||
"dangerous_mode": "EVOSCIENTIST_DANGEROUS_MODE",
|
"dangerous_mode": "EVOSCIENTIST_DANGEROUS_MODE",
|
||||||
"channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING",
|
"channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING",
|
||||||
"ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT",
|
"ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT",
|
||||||
|
"use_responses_api": "EVOSCIENTIST_USE_RESPONSES_API",
|
||||||
"checkpoint_keep_per_thread": "EVOSCIENTIST_CHECKPOINT_KEEP_PER_THREAD",
|
"checkpoint_keep_per_thread": "EVOSCIENTIST_CHECKPOINT_KEEP_PER_THREAD",
|
||||||
"enable_async_subagents": "EVOSCIENTIST_ENABLE_ASYNC_SUBAGENTS",
|
"enable_async_subagents": "EVOSCIENTIST_ENABLE_ASYNC_SUBAGENTS",
|
||||||
"langgraph_dev_port": "EVOSCIENTIST_LANGGRAPH_DEV_PORT",
|
"langgraph_dev_port": "EVOSCIENTIST_LANGGRAPH_DEV_PORT",
|
||||||
|
"langgraph_dev_host": "EVOSCIENTIST_LANGGRAPH_DEV_HOST",
|
||||||
"webui_port": "EVOSCIENTIST_WEBUI_PORT",
|
"webui_port": "EVOSCIENTIST_WEBUI_PORT",
|
||||||
|
"webui_host": "EVOSCIENTIST_WEBUI_HOST",
|
||||||
"enable_scheduler": "EVOSCIENTIST_ENABLE_SCHEDULER",
|
"enable_scheduler": "EVOSCIENTIST_ENABLE_SCHEDULER",
|
||||||
"scheduler_default_timezone": "EVOSCIENTIST_SCHEDULER_DEFAULT_TIMEZONE",
|
"scheduler_default_timezone": "EVOSCIENTIST_SCHEDULER_DEFAULT_TIMEZONE",
|
||||||
"code_interpreter_timeout": "EVOSCIENTIST_CODE_INTERPRETER_TIMEOUT",
|
"code_interpreter_timeout": "EVOSCIENTIST_CODE_INTERPRETER_TIMEOUT",
|
||||||
@@ -823,6 +894,7 @@ _ENV_MAPPINGS = {
|
|||||||
"sandbox_execute_timeout": "EVOSCIENTIST_SANDBOX_EXECUTE_TIMEOUT",
|
"sandbox_execute_timeout": "EVOSCIENTIST_SANDBOX_EXECUTE_TIMEOUT",
|
||||||
"langgraph_dev_file_persistence": "EVOSCIENTIST_LANGGRAPH_DEV_FILE_PERSISTENCE",
|
"langgraph_dev_file_persistence": "EVOSCIENTIST_LANGGRAPH_DEV_FILE_PERSISTENCE",
|
||||||
"langgraph_dev_jobs_per_worker": "EVOSCIENTIST_LANGGRAPH_DEV_JOBS_PER_WORKER",
|
"langgraph_dev_jobs_per_worker": "EVOSCIENTIST_LANGGRAPH_DEV_JOBS_PER_WORKER",
|
||||||
|
"langgraph_dev_keepalive": "EVOSCIENTIST_LANGGRAPH_DEV_KEEPALIVE",
|
||||||
"recursion_limit": "EVOSCIENTIST_RECURSION_LIMIT",
|
"recursion_limit": "EVOSCIENTIST_RECURSION_LIMIT",
|
||||||
"repetitive_tool_call_threshold": (
|
"repetitive_tool_call_threshold": (
|
||||||
"EVOSCIENTIST_REPETITIVE_TOOL_CALL_THRESHOLD"
|
"EVOSCIENTIST_REPETITIVE_TOOL_CALL_THRESHOLD"
|
||||||
@@ -836,6 +908,7 @@ _ENV_MAPPINGS = {
|
|||||||
"memory_skill_synthesis_mode": "EVOSCIENTIST_MEMORY_SKILL_SYNTHESIS_MODE",
|
"memory_skill_synthesis_mode": "EVOSCIENTIST_MEMORY_SKILL_SYNTHESIS_MODE",
|
||||||
"memory_skill_synthesis_cadence": "EVOSCIENTIST_MEMORY_SKILL_SYNTHESIS_CADENCE",
|
"memory_skill_synthesis_cadence": "EVOSCIENTIST_MEMORY_SKILL_SYNTHESIS_CADENCE",
|
||||||
"memory_skill_synthesis_time": "EVOSCIENTIST_MEMORY_SKILL_SYNTHESIS_TIME",
|
"memory_skill_synthesis_time": "EVOSCIENTIST_MEMORY_SKILL_SYNTHESIS_TIME",
|
||||||
|
"memory_observation_cache_max_files": "EVOSCIENTIST_MAX_CACHED_FILES",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -845,10 +918,33 @@ def get_effective_config(
|
|||||||
"""Get effective configuration by merging all sources.
|
"""Get effective configuration by merging all sources.
|
||||||
|
|
||||||
Priority (highest to lowest):
|
Priority (highest to lowest):
|
||||||
1. CLI arguments (cli_overrides)
|
1. CLI arguments (``cli_overrides``)
|
||||||
2. Environment variables
|
2. Parent-process environment variables for any ``EVOSCIENTIST_*`` key
|
||||||
3. Config file
|
3. ``.env`` file at (or above) the current working directory
|
||||||
4. Defaults
|
4. Parent-process environment variables for everything else
|
||||||
|
(third-party API keys / base URLs, plus arbitrary unmapped keys)
|
||||||
|
5. Config file (``~/.config/evoscientist/config.yaml``)
|
||||||
|
6. Dataclass defaults
|
||||||
|
|
||||||
|
Rows 2 and 4 differ because ``.env`` values need different treatment
|
||||||
|
for our own namespaced config knobs vs third-party credentials.
|
||||||
|
Third-party keys (``ANTHROPIC_API_KEY``, ``OPENAI_API_KEY``, ...)
|
||||||
|
follow the industry convention that ``.env`` is the per-project
|
||||||
|
credential store; extending shell-wins to them would silently flip
|
||||||
|
a workspace key back to a global ``.bashrc`` key. Our own
|
||||||
|
``EVOSCIENTIST_*`` keys are the opposite: an explicit CLI/parent-
|
||||||
|
process value (e.g. the bind port that ``EvoSci deploy --port X``
|
||||||
|
hands to the langgraph dev subprocess) must not be shadowed by a
|
||||||
|
workspace ``.env``. We implement this by reading ``.env`` into a
|
||||||
|
dict via ``dotenv_values`` (no ``os.environ`` mutation), then
|
||||||
|
writing third-party keys unconditionally and ``EVOSCIENTIST_*`` keys
|
||||||
|
only when the shell doesn't already have a non-empty value.
|
||||||
|
|
||||||
|
Tradeoff: ``OPENAI_API_KEY=xxx evoscientist ...`` inline overrides
|
||||||
|
still lose to a workspace ``.env`` containing ``OPENAI_API_KEY``,
|
||||||
|
because the merge writes third-party keys from ``.env``
|
||||||
|
unconditionally. Users who need to override a ``.env``-defined
|
||||||
|
credential inline must edit or unset the ``.env`` entry.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
cli_overrides: Dictionary of CLI argument overrides.
|
cli_overrides: Dictionary of CLI argument overrides.
|
||||||
@@ -856,7 +952,34 @@ def get_effective_config(
|
|||||||
Returns:
|
Returns:
|
||||||
EvoScientistConfig with merged values.
|
EvoScientistConfig with merged values.
|
||||||
"""
|
"""
|
||||||
load_dotenv(find_dotenv(usecwd=True), override=True)
|
# Merge workspace ``.env`` into ``os.environ`` without going through
|
||||||
|
# ``load_dotenv``. The previous snapshot → ``load_dotenv`` → restore
|
||||||
|
# sequence was a read-modify-write on ``os.environ`` that could race with
|
||||||
|
# concurrent ``get_effective_config`` calls in the langgraph dev subprocess
|
||||||
|
# (per-request threads in ``langgraph_dev/http.py``, ``sessions.py``
|
||||||
|
# checkpoint writes, memory workers): one thread's mid-flight ``.env``
|
||||||
|
# value could be re-captured by another as "parent env" and then restored
|
||||||
|
# last, promoting the ``.env`` value into the snapshot permanently.
|
||||||
|
#
|
||||||
|
# ``dotenv_values`` returns a dict without touching ``os.environ``, so the
|
||||||
|
# merge below is a pure write sequence and idempotent under interleaving.
|
||||||
|
# Third-party keys keep ``.env``-wins (industry convention).
|
||||||
|
# ``EVOSCIENTIST_*`` keys are our own namespaced config knobs where
|
||||||
|
# CLI/parent-process intent should stay authoritative — write from ``.env``
|
||||||
|
# only when the shell doesn't already have a non-empty value. Treating an
|
||||||
|
# empty shell value as "unset" matches the ``if env_value:`` truthy check
|
||||||
|
# in the ``_ENV_MAPPINGS`` loop below; without this, an empty parent export
|
||||||
|
# would silently regress vs main by falling through to file/defaults.
|
||||||
|
dotenv_path = find_dotenv(usecwd=True)
|
||||||
|
dotenv_map = dotenv_values(dotenv_path) if dotenv_path else {}
|
||||||
|
for env_key, env_value in dotenv_map.items():
|
||||||
|
if env_value is None:
|
||||||
|
continue # bare ``FOO`` without ``=`` — nothing to write
|
||||||
|
if env_key.startswith("EVOSCIENTIST_"):
|
||||||
|
if not os.environ.get(env_key):
|
||||||
|
os.environ[env_key] = env_value
|
||||||
|
else:
|
||||||
|
os.environ[env_key] = env_value
|
||||||
|
|
||||||
# Start with file config (includes defaults for missing values)
|
# Start with file config (includes defaults for missing values)
|
||||||
config = load_config()
|
config = load_config()
|
||||||
@@ -913,6 +1036,12 @@ def apply_config_to_env(config: EvoScientistConfig) -> None:
|
|||||||
os.environ["SILICONFLOW_API_KEY"] = config.siliconflow_api_key
|
os.environ["SILICONFLOW_API_KEY"] = config.siliconflow_api_key
|
||||||
if config.openrouter_api_key and not os.environ.get("OPENROUTER_API_KEY"):
|
if config.openrouter_api_key and not os.environ.get("OPENROUTER_API_KEY"):
|
||||||
os.environ["OPENROUTER_API_KEY"] = config.openrouter_api_key
|
os.environ["OPENROUTER_API_KEY"] = config.openrouter_api_key
|
||||||
|
if config.atlascloud_api_key and not os.environ.get("ATLASCLOUD_API_KEY"):
|
||||||
|
os.environ["ATLASCLOUD_API_KEY"] = config.atlascloud_api_key
|
||||||
|
if config.requesty_api_key and not os.environ.get("REQUESTY_API_KEY"):
|
||||||
|
os.environ["REQUESTY_API_KEY"] = config.requesty_api_key
|
||||||
|
if config.novita_api_key and not os.environ.get("NOVITA_API_KEY"):
|
||||||
|
os.environ["NOVITA_API_KEY"] = config.novita_api_key
|
||||||
if config.deepseek_api_key and not os.environ.get("DEEPSEEK_API_KEY"):
|
if config.deepseek_api_key and not os.environ.get("DEEPSEEK_API_KEY"):
|
||||||
os.environ["DEEPSEEK_API_KEY"] = config.deepseek_api_key
|
os.environ["DEEPSEEK_API_KEY"] = config.deepseek_api_key
|
||||||
if config.zhipu_api_key and not os.environ.get("ZHIPU_API_KEY"):
|
if config.zhipu_api_key and not os.environ.get("ZHIPU_API_KEY"):
|
||||||
@@ -973,3 +1102,7 @@ def apply_config_to_env(config: EvoScientistConfig) -> None:
|
|||||||
os.environ["EVOSCIENTIST_DANGEROUS_MODE"] = "true"
|
os.environ["EVOSCIENTIST_DANGEROUS_MODE"] = "true"
|
||||||
else:
|
else:
|
||||||
os.environ.pop("EVOSCIENTIST_DANGEROUS_MODE", None)
|
os.environ.pop("EVOSCIENTIST_DANGEROUS_MODE", None)
|
||||||
|
if config.use_responses_api and not os.environ.get(
|
||||||
|
"EVOSCIENTIST_USE_RESPONSES_API"
|
||||||
|
):
|
||||||
|
os.environ["EVOSCIENTIST_USE_RESPONSES_API"] = config.use_responses_api
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ multiple clients at one hand-started server they will share the same cron store.
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from langgraph_sdk.schema import Cron, Run
|
from langgraph_sdk.schema import Cron, Run
|
||||||
@@ -28,6 +28,33 @@ SCHEDULER_GRAPH_ID = "scheduler"
|
|||||||
SCHEDULED_RUN_KIND = "scheduled_task"
|
SCHEDULED_RUN_KIND = "scheduled_task"
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_rubric(rubric: str | None) -> str | None:
|
||||||
|
text = (rubric or "").strip()
|
||||||
|
return text or None
|
||||||
|
|
||||||
|
|
||||||
|
def _scheduled_input(prompt: str, rubric: str | None) -> dict[str, Any]:
|
||||||
|
"""Run input for the scheduler graph; ``rubric`` rides along only when set.
|
||||||
|
|
||||||
|
The key is read by ``RubricMiddleware`` mounted on the scheduler graph — an
|
||||||
|
absent key means no grading pass at all, so unset stays byte-identical to
|
||||||
|
the pre-rubric payload.
|
||||||
|
"""
|
||||||
|
payload: dict[str, Any] = messages_input(prompt)
|
||||||
|
if rubric:
|
||||||
|
payload["rubric"] = rubric
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
def _scheduled_metadata(
|
||||||
|
*, name: str, prompt: str, rubric: str | None
|
||||||
|
) -> dict[str, str]:
|
||||||
|
metadata = {"run_kind": SCHEDULED_RUN_KIND, "name": name, "prompt": prompt}
|
||||||
|
if rubric:
|
||||||
|
metadata["rubric"] = rubric
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
|
||||||
def _scheduler_url() -> str:
|
def _scheduler_url() -> str:
|
||||||
return configured_langgraph_dev_url()
|
return configured_langgraph_dev_url()
|
||||||
|
|
||||||
@@ -48,16 +75,26 @@ def is_available() -> bool:
|
|||||||
|
|
||||||
|
|
||||||
def create_schedule(
|
def create_schedule(
|
||||||
*, name: str, schedule: str, prompt: str, timezone: str | None = None
|
*,
|
||||||
|
name: str,
|
||||||
|
schedule: str,
|
||||||
|
prompt: str,
|
||||||
|
timezone: str | None = None,
|
||||||
|
rubric: str | None = None,
|
||||||
) -> Cron:
|
) -> Cron:
|
||||||
"""Create a recurring scheduled task on the scheduler graph."""
|
"""Create a recurring scheduled task on the scheduler graph.
|
||||||
|
|
||||||
|
``rubric`` is an optional acceptance checklist graded after each run; blank
|
||||||
|
means the run is never graded.
|
||||||
|
"""
|
||||||
|
rubric = _normalize_rubric(rubric)
|
||||||
# Crons are stored in the langgraph-dev process's .langgraph_api store, not
|
# Crons are stored in the langgraph-dev process's .langgraph_api store, not
|
||||||
# tagged by workspace. Isolation is process-level (see module docstring).
|
# tagged by workspace. Isolation is process-level (see module docstring).
|
||||||
return _client().crons.create(
|
return _client().crons.create(
|
||||||
assistant_id=SCHEDULER_GRAPH_ID,
|
assistant_id=SCHEDULER_GRAPH_ID,
|
||||||
schedule=schedule,
|
schedule=schedule,
|
||||||
input=messages_input(prompt),
|
input=_scheduled_input(prompt, rubric),
|
||||||
metadata={"run_kind": SCHEDULED_RUN_KIND, "name": name, "prompt": prompt},
|
metadata=_scheduled_metadata(name=name, prompt=prompt, rubric=rubric),
|
||||||
timezone=timezone or _default_timezone(),
|
timezone=timezone or _default_timezone(),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -87,20 +124,17 @@ def set_enabled(cron_id: str, enabled: bool) -> Cron:
|
|||||||
return _client().crons.update(cron_id, enabled=enabled)
|
return _client().crons.update(cron_id, enabled=enabled)
|
||||||
|
|
||||||
|
|
||||||
def run_now(prompt: str) -> Run:
|
def run_now(prompt: str, *, rubric: str | None = None) -> Run:
|
||||||
"""Fire a one-off scheduler run immediately (for ``/schedule run``).
|
"""Fire a one-off scheduler run immediately (for ``/schedule run``).
|
||||||
|
|
||||||
Output goes wherever the task's prompt specifies; there is no push notification.
|
Output goes wherever the task's prompt specifies; there is no push notification.
|
||||||
"""
|
"""
|
||||||
|
rubric = _normalize_rubric(rubric)
|
||||||
client = _client()
|
client = _client()
|
||||||
thread = client.threads.create(graph_id=SCHEDULER_GRAPH_ID)
|
thread = client.threads.create(graph_id=SCHEDULER_GRAPH_ID)
|
||||||
return client.runs.create(
|
return client.runs.create(
|
||||||
thread_id=str(thread["thread_id"]),
|
thread_id=str(thread["thread_id"]),
|
||||||
assistant_id=SCHEDULER_GRAPH_ID,
|
assistant_id=SCHEDULER_GRAPH_ID,
|
||||||
input=messages_input(prompt),
|
input=_scheduled_input(prompt, rubric),
|
||||||
metadata={
|
metadata=_scheduled_metadata(name="manual-run", prompt=prompt, rubric=rubric),
|
||||||
"run_kind": SCHEDULED_RUN_KIND,
|
|
||||||
"name": "manual-run",
|
|
||||||
"prompt": prompt,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -46,6 +46,13 @@ def deploy(
|
|||||||
"--port",
|
"--port",
|
||||||
help="Port for langgraph dev (default: config.langgraph_dev_port = 3076)",
|
help="Port for langgraph dev (default: config.langgraph_dev_port = 3076)",
|
||||||
),
|
),
|
||||||
|
host: str | None = typer.Option(
|
||||||
|
None,
|
||||||
|
"--host",
|
||||||
|
help="Interface to bind (default: config.langgraph_dev_host = "
|
||||||
|
"127.0.0.1, i.e. this machine only — pass 0.0.0.0 to reach it from "
|
||||||
|
"other machines, but note the server has no auth)",
|
||||||
|
),
|
||||||
tunnel: bool = typer.Option(
|
tunnel: bool = typer.Option(
|
||||||
False,
|
False,
|
||||||
"--tunnel",
|
"--tunnel",
|
||||||
@@ -66,9 +73,15 @@ def deploy(
|
|||||||
"""
|
"""
|
||||||
from ..config import apply_config_to_env, get_effective_config
|
from ..config import apply_config_to_env, get_effective_config
|
||||||
from ..langgraph_dev.manager import (
|
from ..langgraph_dev.manager import (
|
||||||
|
_DEFAULT_HOST,
|
||||||
_DEFAULT_PORT,
|
_DEFAULT_PORT,
|
||||||
RUNTIME,
|
RUNTIME,
|
||||||
|
_base_url,
|
||||||
|
_is_loopback_host,
|
||||||
_is_port_occupied,
|
_is_port_occupied,
|
||||||
|
_pid_serves_port,
|
||||||
|
_read_workspace_sidecar,
|
||||||
|
_server_config_fingerprint,
|
||||||
is_langgraph_dev_running,
|
is_langgraph_dev_running,
|
||||||
read_tunnel_url,
|
read_tunnel_url,
|
||||||
start_langgraph_dev,
|
start_langgraph_dev,
|
||||||
@@ -114,20 +127,46 @@ def deploy(
|
|||||||
)
|
)
|
||||||
raise typer.Exit(1)
|
raise typer.Exit(1)
|
||||||
|
|
||||||
|
# A blank ``--host`` means "not passed" (matching serve), so it can never
|
||||||
|
# discard the configured bind. Both branches strip: whitespace reaching
|
||||||
|
# socket.bind() surfaces as an opaque gaierror, and duck-typed configs
|
||||||
|
# handed to this function never ran ``__post_init__`` normalization.
|
||||||
|
cli_host = host.strip() if host is not None else ""
|
||||||
|
effective_host = (
|
||||||
|
cli_host
|
||||||
|
or str(getattr(config, "langgraph_dev_host", _DEFAULT_HOST) or "").strip()
|
||||||
|
or _DEFAULT_HOST
|
||||||
|
)
|
||||||
|
|
||||||
# 4. Pre-flight port check — refuse to start if a non-EvoSci process is
|
# 4. Pre-flight port check — refuse to start if a non-EvoSci process is
|
||||||
# holding the port. If an existing EvoSci langgraph dev is already up,
|
# holding the port. If an existing EvoSci langgraph dev is already up,
|
||||||
# also refuse (deploy is the "primary server" — running multiple on the
|
# also refuse (deploy is the "primary server" — running multiple on the
|
||||||
# same port is a configuration error).
|
# same port is a configuration error).
|
||||||
if _is_port_occupied(effective_port):
|
if _is_port_occupied(effective_port, effective_host):
|
||||||
if is_langgraph_dev_running(port=effective_port):
|
if is_langgraph_dev_running(port=effective_port, host=effective_host):
|
||||||
console.print(
|
console.print(
|
||||||
f"[red]Port {effective_port} is already serving a langgraph dev "
|
f"[red]Port {effective_port} is already serving a langgraph dev "
|
||||||
f"instance.[/red]"
|
f"instance.[/red]"
|
||||||
)
|
)
|
||||||
console.print(
|
sidecar = _read_workspace_sidecar()
|
||||||
"[dim]Stop the existing EvoSci/serve session first, or use "
|
if sidecar is not None and _pid_serves_port(
|
||||||
"[bold]--port[/bold] to deploy on a different port.[/dim]"
|
sidecar.get("pid"), effective_port
|
||||||
)
|
):
|
||||||
|
# Surface what we know about the occupant — with keepalive it
|
||||||
|
# may be an ownerless leftover rather than a live session.
|
||||||
|
# Only when the recorded pid verifiably serves THIS port, so a
|
||||||
|
# stale or other-port record is never blamed.
|
||||||
|
console.print(
|
||||||
|
f"[dim]It serves workspace {sidecar.get('workspace')} "
|
||||||
|
f"(pid {sidecar.get('pid')}). Stop it with "
|
||||||
|
f"[bold]EvoSci server stop[/bold], or use "
|
||||||
|
f"[bold]--port[/bold] to deploy on a different port.[/dim]"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
console.print(
|
||||||
|
"[dim]Stop the existing EvoSci/serve session first, or use "
|
||||||
|
"[bold]--port[/bold] to deploy on a different port.[/dim]"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
console.print(
|
console.print(
|
||||||
f"[red]Port {effective_port} is occupied by another process.[/red]"
|
f"[red]Port {effective_port} is occupied by another process.[/red]"
|
||||||
@@ -144,6 +183,7 @@ def deploy(
|
|||||||
Panel(
|
Panel(
|
||||||
Text.from_markup(
|
Text.from_markup(
|
||||||
f"[bold]Workspace:[/bold] {_shorten(ws)}\n"
|
f"[bold]Workspace:[/bold] {_shorten(ws)}\n"
|
||||||
|
f"[bold]Host:[/bold] {effective_host}\n"
|
||||||
f"[bold]Port:[/bold] {effective_port}\n"
|
f"[bold]Port:[/bold] {effective_port}\n"
|
||||||
f"[bold]Auth:[/bold] {_auth_label}"
|
f"[bold]Auth:[/bold] {_auth_label}"
|
||||||
),
|
),
|
||||||
@@ -162,6 +202,13 @@ def deploy(
|
|||||||
f"[bold red]{DANGEROUS_BANNER_MESSAGE}[/bold red]"
|
f"[bold red]{DANGEROUS_BANNER_MESSAGE}[/bold red]"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if not _is_loopback_host(effective_host):
|
||||||
|
console.print(
|
||||||
|
"[bold white on red] ⚠ PUBLIC BIND [/bold white on red] "
|
||||||
|
f"[bold red]Listening on {effective_host} — no auth, and the agent "
|
||||||
|
f"can run shell. Trusted networks only.[/bold red]"
|
||||||
|
)
|
||||||
|
|
||||||
if tunnel:
|
if tunnel:
|
||||||
console.print(
|
console.print(
|
||||||
"[bold white on red] ⚠ PUBLIC TUNNEL [/bold white on red] "
|
"[bold white on red] ⚠ PUBLIC TUNNEL [/bold white on red] "
|
||||||
@@ -197,10 +244,12 @@ def deploy(
|
|||||||
proc = start_langgraph_dev(
|
proc = start_langgraph_dev(
|
||||||
workspace_dir=Path(ws),
|
workspace_dir=Path(ws),
|
||||||
port=effective_port,
|
port=effective_port,
|
||||||
|
host=effective_host,
|
||||||
file_persistence=file_persistence,
|
file_persistence=file_persistence,
|
||||||
jobs_per_worker=jobs_per_worker,
|
jobs_per_worker=jobs_per_worker,
|
||||||
deploy_mode=True,
|
deploy_mode=True,
|
||||||
tunnel=tunnel,
|
tunnel=tunnel,
|
||||||
|
config_fingerprint=_server_config_fingerprint(config),
|
||||||
)
|
)
|
||||||
atexit.register(stop_langgraph_dev, proc)
|
atexit.register(stop_langgraph_dev, proc)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -236,7 +285,7 @@ def deploy(
|
|||||||
Panel(
|
Panel(
|
||||||
Text.from_markup(
|
Text.from_markup(
|
||||||
f"[bold]Endpoint:[/bold] "
|
f"[bold]Endpoint:[/bold] "
|
||||||
f"http://localhost:{effective_port}\n"
|
f"{_base_url(effective_port, effective_host)}\n"
|
||||||
f"{public_line}"
|
f"{public_line}"
|
||||||
f"[bold]Assistant ID:[/bold] EvoScientist\n"
|
f"[bold]Assistant ID:[/bold] EvoScientist\n"
|
||||||
f"[bold]Connect via:[/bold] any LangChain SDK / "
|
f"[bold]Connect via:[/bold] any LangChain SDK / "
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ from ..stream.console import console
|
|||||||
# Front-end npm package + spec. ``@latest`` → always the newest published UI.
|
# Front-end npm package + spec. ``@latest`` → always the newest published UI.
|
||||||
_WEBUI_PACKAGE = "@evoscientist/webui@latest"
|
_WEBUI_PACKAGE = "@evoscientist/webui@latest"
|
||||||
_DEFAULT_WEBUI_PORT = 4716
|
_DEFAULT_WEBUI_PORT = 4716
|
||||||
|
_DEFAULT_WEBUI_HOST = "127.0.0.1"
|
||||||
|
|
||||||
|
|
||||||
def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
||||||
@@ -58,10 +59,15 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
|||||||
"""
|
"""
|
||||||
from ..config import apply_config_to_env
|
from ..config import apply_config_to_env
|
||||||
from ..langgraph_dev.manager import (
|
from ..langgraph_dev.manager import (
|
||||||
|
_DEFAULT_HOST,
|
||||||
_DEFAULT_PORT,
|
_DEFAULT_PORT,
|
||||||
RUNTIME,
|
RUNTIME,
|
||||||
|
_base_url,
|
||||||
|
_format_hostport,
|
||||||
|
_is_loopback_host,
|
||||||
_is_port_occupied,
|
_is_port_occupied,
|
||||||
_read_workspace_sidecar,
|
_read_workspace_sidecar,
|
||||||
|
_server_config_fingerprint,
|
||||||
is_langgraph_dev_running,
|
is_langgraph_dev_running,
|
||||||
start_langgraph_dev,
|
start_langgraph_dev,
|
||||||
stop_langgraph_dev,
|
stop_langgraph_dev,
|
||||||
@@ -84,6 +90,14 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
|||||||
# webui_port = the local Next.js server the browser actually opens.
|
# webui_port = the local Next.js server the browser actually opens.
|
||||||
backend_port = int(getattr(config, "langgraph_dev_port", _DEFAULT_PORT))
|
backend_port = int(getattr(config, "langgraph_dev_port", _DEFAULT_PORT))
|
||||||
webui_port = int(getattr(config, "webui_port", _DEFAULT_WEBUI_PORT))
|
webui_port = int(getattr(config, "webui_port", _DEFAULT_WEBUI_PORT))
|
||||||
|
# ...and their bind interfaces, both loopback by default — the front-end
|
||||||
|
# carries workspace/skill APIs of its own (see config.webui_host).
|
||||||
|
backend_host = (
|
||||||
|
str(getattr(config, "langgraph_dev_host", _DEFAULT_HOST) or _DEFAULT_HOST)
|
||||||
|
).strip() or _DEFAULT_HOST
|
||||||
|
webui_host = (
|
||||||
|
str(getattr(config, "webui_host", _DEFAULT_WEBUI_HOST) or _DEFAULT_WEBUI_HOST)
|
||||||
|
).strip() or _DEFAULT_WEBUI_HOST
|
||||||
for label, p in (("langgraph dev", backend_port), ("WebUI", webui_port)):
|
for label, p in (("langgraph dev", backend_port), ("WebUI", webui_port)):
|
||||||
if not (1 <= p <= 65535):
|
if not (1 <= p <= 65535):
|
||||||
console.print(
|
console.print(
|
||||||
@@ -128,8 +142,8 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
|||||||
# else start a fresh deploy-mode one (full MCP + async). Refuse a foreign
|
# else start a fresh deploy-mode one (full MCP + async). Refuse a foreign
|
||||||
# occupant — that's a configuration error, not something to silently share.
|
# occupant — that's a configuration error, not something to silently share.
|
||||||
started_proc = None
|
started_proc = None
|
||||||
if _is_port_occupied(backend_port):
|
if _is_port_occupied(backend_port, backend_host):
|
||||||
if is_langgraph_dev_running(port=backend_port):
|
if is_langgraph_dev_running(port=backend_port, host=backend_host):
|
||||||
# Reuse an existing EvoSci server only when it serves THIS workspace
|
# Reuse an existing EvoSci server only when it serves THIS workspace
|
||||||
# — mirror the sidecar guard in ensure_langgraph_dev so WebUI started
|
# — mirror the sidecar guard in ensure_langgraph_dev so WebUI started
|
||||||
# from workspace B never silently binds to a server pinned to
|
# from workspace B never silently binds to a server pinned to
|
||||||
@@ -150,6 +164,31 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
|||||||
f"[/dim]"
|
f"[/dim]"
|
||||||
)
|
)
|
||||||
raise typer.Exit(1)
|
raise typer.Exit(1)
|
||||||
|
if sidecar is not None and sidecar.get("deploy_mode") is False:
|
||||||
|
# A stripped (CLI-started) server has no MCP and no async
|
||||||
|
# sub-agents — silently reusing it would degrade the WebUI
|
||||||
|
# with no visible cause. Refuse; never auto-kill.
|
||||||
|
console.print(
|
||||||
|
f"[red]Port {backend_port} is serving a stripped "
|
||||||
|
f"(CLI-mode) langgraph dev — the WebUI needs the full "
|
||||||
|
f"deploy-mode server (MCP + async sub-agents).[/red]"
|
||||||
|
)
|
||||||
|
console.print(
|
||||||
|
"[dim]Stop it with [bold]EvoSci server stop[/bold], then "
|
||||||
|
"re-run [bold]EvoSci[/bold].[/dim]"
|
||||||
|
)
|
||||||
|
raise typer.Exit(1)
|
||||||
|
if sidecar is not None:
|
||||||
|
recorded_fp = sidecar.get("config_fingerprint")
|
||||||
|
if isinstance(
|
||||||
|
recorded_fp, str
|
||||||
|
) and recorded_fp != _server_config_fingerprint(config):
|
||||||
|
console.print(
|
||||||
|
"[yellow]⚠ Config changed since this server was "
|
||||||
|
"launched — it still serves the old settings. Apply "
|
||||||
|
"them with [bold]EvoSci server stop[/bold], then "
|
||||||
|
"re-run EvoSci.[/yellow]"
|
||||||
|
)
|
||||||
console.print(
|
console.print(
|
||||||
f"[green]✓[/green] Reusing langgraph dev already serving "
|
f"[green]✓[/green] Reusing langgraph dev already serving "
|
||||||
f"port {backend_port}"
|
f"port {backend_port}"
|
||||||
@@ -174,17 +213,28 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
|||||||
started_proc = start_langgraph_dev(
|
started_proc = start_langgraph_dev(
|
||||||
workspace_dir=Path(ws),
|
workspace_dir=Path(ws),
|
||||||
port=backend_port,
|
port=backend_port,
|
||||||
|
host=backend_host,
|
||||||
file_persistence=file_persistence,
|
file_persistence=file_persistence,
|
||||||
jobs_per_worker=jobs_per_worker,
|
jobs_per_worker=jobs_per_worker,
|
||||||
deploy_mode=True,
|
deploy_mode=True,
|
||||||
|
config_fingerprint=_server_config_fingerprint(config),
|
||||||
)
|
)
|
||||||
atexit.register(stop_langgraph_dev, started_proc)
|
if getattr(config, "langgraph_dev_keepalive", False):
|
||||||
|
# Keepalive: the deploy-mode backend outlives this session so
|
||||||
|
# the next same-workspace launch reuses it instantly. The npx
|
||||||
|
# front-end below still stops on exit as usual.
|
||||||
|
console.print(
|
||||||
|
"[dim]keepalive: backend server stays up after exit — "
|
||||||
|
"stop it with [bold]EvoSci server stop[/bold].[/dim]"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
atexit.register(stop_langgraph_dev, started_proc)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
console.print(f"[red]langgraph dev startup failed:[/red] {exc}")
|
console.print(f"[red]langgraph dev startup failed:[/red] {exc}")
|
||||||
raise typer.Exit(1) from exc
|
raise typer.Exit(1) from exc
|
||||||
console.print("[green]✓[/green] langgraph dev ready")
|
console.print("[green]✓[/green] langgraph dev ready")
|
||||||
|
|
||||||
if _is_port_occupied(webui_port):
|
if _is_port_occupied(webui_port, webui_host):
|
||||||
console.print(
|
console.print(
|
||||||
f"[yellow]⚠ Port {webui_port} is already in use; the WebUI server "
|
f"[yellow]⚠ Port {webui_port} is already in use; the WebUI server "
|
||||||
f"may fail to start. Change it with "
|
f"may fail to start. Change it with "
|
||||||
@@ -197,20 +247,37 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
|||||||
# inherited so it all shows in THIS terminal. EVOSCIENTIST_LANGGRAPH_DEV_PORT
|
# inherited so it all shows in THIS terminal. EVOSCIENTIST_LANGGRAPH_DEV_PORT
|
||||||
# lets the UI's config prefill point at our backend automatically. Secrets
|
# lets the UI's config prefill point at our backend automatically. Secrets
|
||||||
# are scrubbed — the browser UI never needs LLM provider API keys.
|
# are scrubbed — the browser UI never needs LLM provider API keys.
|
||||||
|
#
|
||||||
|
# HOSTNAME is the front-end's only bind knob: the package has no --host
|
||||||
|
# flag; its launcher forwards `HOSTNAME || "127.0.0.1"` to the Next server.
|
||||||
webui_env = _scrubbed_env(
|
webui_env = _scrubbed_env(
|
||||||
{
|
{
|
||||||
"EVOSCIENTIST_LANGGRAPH_DEV_PORT": str(backend_port),
|
"EVOSCIENTIST_LANGGRAPH_DEV_PORT": str(backend_port),
|
||||||
"PORT": str(webui_port),
|
"PORT": str(webui_port),
|
||||||
|
"HOSTNAME": webui_host,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
# The UI reaches the backend from the BROWSER; when only the front-end is
|
||||||
|
# exposed, remote pages load but every request fails — say so.
|
||||||
|
remote_backend_hint = ""
|
||||||
|
if not _is_loopback_host(webui_host) and _is_loopback_host(backend_host):
|
||||||
|
remote_backend_hint = (
|
||||||
|
f"\n[yellow]Note:[/yellow] the UI connects to the backend from the "
|
||||||
|
f"browser. Remote visitors cannot reach a loopback backend — run "
|
||||||
|
f"[bold]EvoSci config set langgraph_dev_host 0.0.0.0[/bold] and "
|
||||||
|
f"point the UI at [bold]http://<this-machine-ip>:{backend_port}"
|
||||||
|
f"[/bold].\n"
|
||||||
|
)
|
||||||
console.print(
|
console.print(
|
||||||
Panel(
|
Panel(
|
||||||
Text.from_markup(
|
Text.from_markup(
|
||||||
f"[bold]Backend:[/bold] http://localhost:{backend_port} "
|
f"[bold]Backend:[/bold] {_base_url(backend_port, backend_host)} "
|
||||||
f"[dim](langgraph dev — Assistant: EvoScientist)[/dim]\n"
|
f"[dim](langgraph dev — Assistant: EvoScientist)[/dim]\n"
|
||||||
f"[bold]WebUI:[/bold] http://localhost:{webui_port} "
|
f"[bold]WebUI:[/bold] "
|
||||||
|
f"http://{_format_hostport(webui_host, webui_port)} "
|
||||||
f"[dim](opens in your browser)[/dim]\n"
|
f"[dim](opens in your browser)[/dim]\n"
|
||||||
f"[bold]Logs:[/bold] {_shorten(str(RUNTIME.log_file))}\n\n"
|
f"[bold]Logs:[/bold] {_shorten(str(RUNTIME.log_file))}\n"
|
||||||
|
f"{remote_backend_hint}\n"
|
||||||
f"[dim]Fetching {_WEBUI_PACKAGE} via npx (first run may take a "
|
f"[dim]Fetching {_WEBUI_PACKAGE} via npx (first run may take a "
|
||||||
f"moment)… Press Ctrl+C to stop.[/dim]"
|
f"moment)… Press Ctrl+C to stop.[/dim]"
|
||||||
),
|
),
|
||||||
@@ -218,6 +285,19 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
|
|||||||
border_style="green",
|
border_style="green",
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
if not _is_loopback_host(backend_host):
|
||||||
|
console.print(
|
||||||
|
"[bold white on red] ⚠ PUBLIC BIND [/bold white on red] "
|
||||||
|
f"[bold red]Backend listening on {backend_host} — no auth, and the "
|
||||||
|
f"agent can run shell. Trusted networks only.[/bold red]"
|
||||||
|
)
|
||||||
|
if not _is_loopback_host(webui_host):
|
||||||
|
console.print(
|
||||||
|
"[bold white on red] ⚠ PUBLIC BIND [/bold white on red] "
|
||||||
|
f"[bold red]WebUI listening on {webui_host} — its API reads, writes "
|
||||||
|
f"and uploads workspace files and installs skills, with no auth. "
|
||||||
|
f"Trusted networks only.[/bold red]"
|
||||||
|
)
|
||||||
|
|
||||||
popen_kwargs: dict[str, Any] = {"env": webui_env}
|
popen_kwargs: dict[str, Any] = {"env": webui_env}
|
||||||
if os.name == "posix":
|
if os.name == "posix":
|
||||||
|
|||||||
@@ -0,0 +1,466 @@
|
|||||||
|
"""Bounded document extraction and non-text file policy for workspaces."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
import subprocess
|
||||||
|
import sys
|
||||||
|
import tempfile
|
||||||
|
import zipfile
|
||||||
|
from pathlib import Path
|
||||||
|
from xml.etree import ElementTree as ET
|
||||||
|
|
||||||
|
MAX_DOCUMENT_BYTES = 50 * 1024 * 1024
|
||||||
|
MAX_DOCUMENT_RESULT_CHARS = 50_000
|
||||||
|
MAX_CONVERTED_DOCUMENT_BYTES = 10 * 1024 * 1024
|
||||||
|
DOCUMENT_CONVERSION_TIMEOUT_SECONDS = 60
|
||||||
|
MAX_IMAGE_BYTES = 25 * 1024 * 1024
|
||||||
|
MAX_IMAGE_EDGE = 2048
|
||||||
|
MAX_IMAGE_PIXELS = 40_000_000
|
||||||
|
MAX_OOXML_MEMBERS = 10_000
|
||||||
|
MAX_OOXML_MEMBER_BYTES = 50 * 1024 * 1024
|
||||||
|
MAX_OOXML_EXPANDED_BYTES = 200 * 1024 * 1024
|
||||||
|
MAX_OOXML_COMPRESSION_RATIO = 100
|
||||||
|
|
||||||
|
IMAGE_EXTENSIONS = frozenset(
|
||||||
|
{".bmp", ".gif", ".ico", ".jpeg", ".jpg", ".png", ".tif", ".tiff", ".webp"}
|
||||||
|
)
|
||||||
|
DOCUMENT_EXTENSIONS = frozenset(
|
||||||
|
{
|
||||||
|
".doc",
|
||||||
|
".docm",
|
||||||
|
".docx",
|
||||||
|
".epub",
|
||||||
|
".odp",
|
||||||
|
".ods",
|
||||||
|
".odt",
|
||||||
|
".pdf",
|
||||||
|
".pot",
|
||||||
|
".pps",
|
||||||
|
".ppsm",
|
||||||
|
".ppsx",
|
||||||
|
".ppt",
|
||||||
|
".pptm",
|
||||||
|
".pptx",
|
||||||
|
".rtf",
|
||||||
|
".xls",
|
||||||
|
".xlsb",
|
||||||
|
".xlsm",
|
||||||
|
".xlsx",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
ARCHIVE_EXTENSIONS = frozenset(
|
||||||
|
{".7z", ".bz2", ".gz", ".rar", ".tar", ".tgz", ".xz", ".zip"}
|
||||||
|
)
|
||||||
|
DATABASE_EXTENSIONS = frozenset({".db", ".sqlite", ".sqlite3"})
|
||||||
|
EXECUTABLE_EXTENSIONS = frozenset(
|
||||||
|
{".app", ".deb", ".dll", ".dylib", ".elf", ".exe", ".msi", ".rpm", ".so"}
|
||||||
|
)
|
||||||
|
DATASET_EXTENSIONS = frozenset(
|
||||||
|
{".arrow", ".feather", ".h5", ".hdf5", ".npy", ".npz", ".parquet"}
|
||||||
|
)
|
||||||
|
MEDIA_EXTENSIONS = frozenset(
|
||||||
|
{".aac", ".avi", ".flac", ".m4a", ".mkv", ".mov", ".mp3", ".mp4", ".ogg", ".wav", ".webm"}
|
||||||
|
)
|
||||||
|
|
||||||
|
_W = "http://schemas.openxmlformats.org/wordprocessingml/2006/main"
|
||||||
|
_A = "http://schemas.openxmlformats.org/drawingml/2006/main"
|
||||||
|
_S = "http://schemas.openxmlformats.org/spreadsheetml/2006/main"
|
||||||
|
_R = "http://schemas.openxmlformats.org/officeDocument/2006/relationships"
|
||||||
|
|
||||||
|
|
||||||
|
class DocumentExtractionError(RuntimeError):
|
||||||
|
"""A supported document could not be converted to bounded text."""
|
||||||
|
|
||||||
|
|
||||||
|
def classify_file(path: str, head: bytes = b"") -> str:
|
||||||
|
"""Classify a workspace file into one policy category."""
|
||||||
|
|
||||||
|
extension = Path(path).suffix.lower()
|
||||||
|
if extension in IMAGE_EXTENSIONS:
|
||||||
|
return "image"
|
||||||
|
if extension in DOCUMENT_EXTENSIONS:
|
||||||
|
return "document"
|
||||||
|
if extension in ARCHIVE_EXTENSIONS:
|
||||||
|
return "archive"
|
||||||
|
if extension in DATABASE_EXTENSIONS or head.startswith(b"SQLite format 3\x00"):
|
||||||
|
return "database"
|
||||||
|
if extension in EXECUTABLE_EXTENSIONS or head.startswith((b"MZ", b"\x7fELF")):
|
||||||
|
return "executable"
|
||||||
|
if extension in DATASET_EXTENSIONS:
|
||||||
|
return "dataset"
|
||||||
|
if extension in MEDIA_EXTENSIONS:
|
||||||
|
return "media"
|
||||||
|
return "unknown"
|
||||||
|
|
||||||
|
|
||||||
|
def binary_processing_guidance(path: str, kind: str, size_bytes: int) -> str:
|
||||||
|
"""Return bounded, actionable JSON-like guidance for Agent-side programming."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
|
||||||
|
common = {
|
||||||
|
"code": "BINARY_PROCESSING_REQUIRED"
|
||||||
|
if kind != "binary"
|
||||||
|
else "UNSUPPORTED_BINARY_FILE",
|
||||||
|
"path": path,
|
||||||
|
"kind": kind,
|
||||||
|
"size_bytes": size_bytes,
|
||||||
|
}
|
||||||
|
if kind == "archive":
|
||||||
|
common.update(
|
||||||
|
action=(
|
||||||
|
"Use execute with Python to list and validate archive members before "
|
||||||
|
"selective extraction; never use extractall."
|
||||||
|
),
|
||||||
|
constraints={
|
||||||
|
"list_before_extract": True,
|
||||||
|
"max_members": 2000,
|
||||||
|
"max_total_uncompressed_bytes": 500 * 1024 * 1024,
|
||||||
|
"max_member_bytes": 100 * 1024 * 1024,
|
||||||
|
"max_compression_ratio": 100,
|
||||||
|
"reject_absolute_or_parent_paths": True,
|
||||||
|
"do_not_execute_members": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
elif kind == "database":
|
||||||
|
common.update(
|
||||||
|
action=(
|
||||||
|
"Use execute with Python sqlite3 in read-only mode: "
|
||||||
|
"file:<path>?mode=ro&immutable=1; set PRAGMA query_only=ON; "
|
||||||
|
"inspect schema, run bounded SELECT queries with LIMIT, and write "
|
||||||
|
"large results under artifacts/."
|
||||||
|
),
|
||||||
|
constraints={
|
||||||
|
"read_only": True,
|
||||||
|
"mode": "mode=ro",
|
||||||
|
"query_only": True,
|
||||||
|
"max_rows": 1000,
|
||||||
|
"forbid_attach_database": True,
|
||||||
|
"forbid_load_extension": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
elif kind == "executable":
|
||||||
|
common.update(
|
||||||
|
action=(
|
||||||
|
"Use execute only for bounded static metadata inspection (hash, file "
|
||||||
|
"headers, signature, imports, strings); this file must not be executed."
|
||||||
|
),
|
||||||
|
constraints={"must_not_be_executed": True, "static_analysis_only": True},
|
||||||
|
)
|
||||||
|
elif kind == "dataset":
|
||||||
|
common.update(
|
||||||
|
action=(
|
||||||
|
"Use execute with the appropriate library to inspect schema, dimensions, "
|
||||||
|
"statistics, and a bounded sample; do not serialize the whole dataset."
|
||||||
|
)
|
||||||
|
)
|
||||||
|
elif kind == "media":
|
||||||
|
common.update(
|
||||||
|
action=(
|
||||||
|
"Use execute with ffprobe/ffmpeg or an available transcription workflow "
|
||||||
|
"to inspect metadata and selected ranges; do not inline the complete file."
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
common.update(
|
||||||
|
kind="binary",
|
||||||
|
action=(
|
||||||
|
"This is an unsupported binary file. Use execute only for bounded static "
|
||||||
|
"inspection; do not execute it or inline its bytes."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return json.dumps(common, ensure_ascii=False, sort_keys=True)
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_image_bytes(data: bytes, path: str) -> bytes:
|
||||||
|
"""Validate and downsample an image before it becomes a model media block."""
|
||||||
|
|
||||||
|
import io
|
||||||
|
|
||||||
|
if len(data) > MAX_IMAGE_BYTES:
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"IMAGE_TOO_LARGE: {len(data)} bytes exceeds {MAX_IMAGE_BYTES}"
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
image = Image.open(io.BytesIO(data))
|
||||||
|
width, height = image.size
|
||||||
|
if width * height > MAX_IMAGE_PIXELS:
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"IMAGE_PIXEL_BUDGET_EXCEEDED: {width}x{height} exceeds "
|
||||||
|
f"{MAX_IMAGE_PIXELS} pixels"
|
||||||
|
)
|
||||||
|
image.load()
|
||||||
|
except DocumentExtractionError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"IMAGE_PROCESSING_FAILED: {path}: {type(exc).__name__}: {exc}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
frame_count = int(getattr(image, "n_frames", 1) or 1)
|
||||||
|
if frame_count > 1:
|
||||||
|
image.seek(0)
|
||||||
|
work = image.convert("RGBA" if image.mode in {"RGBA", "LA"} else "RGB")
|
||||||
|
else:
|
||||||
|
work = image.copy()
|
||||||
|
|
||||||
|
if max(work.size) <= MAX_IMAGE_EDGE and frame_count == 1:
|
||||||
|
return data
|
||||||
|
|
||||||
|
work.thumbnail((MAX_IMAGE_EDGE, MAX_IMAGE_EDGE), Image.Resampling.LANCZOS)
|
||||||
|
has_alpha = work.mode in {"RGBA", "LA"} or (
|
||||||
|
work.mode == "P" and "transparency" in work.info
|
||||||
|
)
|
||||||
|
output = io.BytesIO()
|
||||||
|
if has_alpha:
|
||||||
|
if work.mode == "P":
|
||||||
|
work = work.convert("RGBA")
|
||||||
|
work.save(output, "PNG", optimize=True)
|
||||||
|
else:
|
||||||
|
if work.mode != "RGB":
|
||||||
|
work = work.convert("RGB")
|
||||||
|
work.save(output, "JPEG", quality=85, optimize=True)
|
||||||
|
return output.getvalue()
|
||||||
|
|
||||||
|
|
||||||
|
def extract_document_bytes(data: bytes, path: str) -> str:
|
||||||
|
"""Extract readable text from a supported document without exposing bytes."""
|
||||||
|
|
||||||
|
if len(data) > MAX_DOCUMENT_BYTES:
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"DOCUMENT_TOO_LARGE: {len(data)} bytes exceeds {MAX_DOCUMENT_BYTES}"
|
||||||
|
)
|
||||||
|
extension = Path(path).suffix.lower()
|
||||||
|
try:
|
||||||
|
if extension == ".docx":
|
||||||
|
return _extract_docx(data)
|
||||||
|
if extension == ".pptx":
|
||||||
|
return _extract_pptx(data)
|
||||||
|
if extension == ".xlsx":
|
||||||
|
return _extract_xlsx(data)
|
||||||
|
return _extract_anydoc(data, extension)
|
||||||
|
except DocumentExtractionError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"DOCUMENT_EXTRACTION_FAILED: {type(exc).__name__}: {exc}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
|
||||||
|
def paginate_document_text(text: str, *, offset: int, limit: int) -> str:
|
||||||
|
"""Apply line and character budgets to extracted document text."""
|
||||||
|
|
||||||
|
lines = text.splitlines(keepends=True)
|
||||||
|
if not lines:
|
||||||
|
return "(document contains no extractable text)"
|
||||||
|
if offset >= len(lines):
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"Line offset {offset} exceeds extracted document length ({len(lines)} lines)"
|
||||||
|
)
|
||||||
|
selected = "".join(lines[offset : offset + limit])
|
||||||
|
if len(selected) <= MAX_DOCUMENT_RESULT_CHARS:
|
||||||
|
return selected
|
||||||
|
trimmed = selected[:MAX_DOCUMENT_RESULT_CHARS]
|
||||||
|
boundary = trimmed.rfind("\n")
|
||||||
|
if boundary > 0:
|
||||||
|
trimmed = trimmed[: boundary + 1]
|
||||||
|
consumed = max(1, len(trimmed.splitlines()))
|
||||||
|
return (
|
||||||
|
trimmed
|
||||||
|
+ f"\n[DOCUMENT_OUTPUT_TRUNCATED: use offset={offset + consumed} to continue; "
|
||||||
|
+ f"single-read limit is {MAX_DOCUMENT_RESULT_CHARS} characters]\n"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _validated_ooxml_archive(data: bytes) -> zipfile.ZipFile:
|
||||||
|
try:
|
||||||
|
archive = zipfile.ZipFile(_bytes_path(data))
|
||||||
|
members = archive.infolist()
|
||||||
|
except zipfile.BadZipFile as exc:
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
"DOCUMENT_EXTRACTION_FAILED: invalid OOXML container"
|
||||||
|
) from exc
|
||||||
|
expanded = 0
|
||||||
|
names: set[str] = set()
|
||||||
|
if len(members) > MAX_OOXML_MEMBERS:
|
||||||
|
archive.close()
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"DOCUMENT_RESOURCE_LIMIT: OOXML has {len(members)} members; "
|
||||||
|
f"limit is {MAX_OOXML_MEMBERS}"
|
||||||
|
)
|
||||||
|
for member in members:
|
||||||
|
normalized = member.filename.replace("\\", "/")
|
||||||
|
parts = tuple(part for part in normalized.split("/") if part)
|
||||||
|
if (
|
||||||
|
normalized.startswith("/")
|
||||||
|
or ".." in parts
|
||||||
|
or member.filename in names
|
||||||
|
or bool(member.flag_bits & 0x1)
|
||||||
|
):
|
||||||
|
archive.close()
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
"DOCUMENT_RESOURCE_LIMIT: OOXML contains an unsafe, duplicate, "
|
||||||
|
"or encrypted member"
|
||||||
|
)
|
||||||
|
names.add(member.filename)
|
||||||
|
expanded += member.file_size
|
||||||
|
ratio = member.file_size / max(member.compress_size, 1)
|
||||||
|
if (
|
||||||
|
member.file_size > MAX_OOXML_MEMBER_BYTES
|
||||||
|
or expanded > MAX_OOXML_EXPANDED_BYTES
|
||||||
|
or ratio > MAX_OOXML_COMPRESSION_RATIO
|
||||||
|
):
|
||||||
|
archive.close()
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
"DOCUMENT_RESOURCE_LIMIT: OOXML member expansion exceeds safety limits"
|
||||||
|
)
|
||||||
|
return archive
|
||||||
|
|
||||||
|
|
||||||
|
def _zip_xml(data: bytes, member: str) -> ET.Element:
|
||||||
|
try:
|
||||||
|
with _validated_ooxml_archive(data) as archive:
|
||||||
|
raw = archive.read(member)
|
||||||
|
except KeyError as exc:
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"DOCUMENT_EXTRACTION_FAILED: missing {member}"
|
||||||
|
) from exc
|
||||||
|
return ET.fromstring(raw)
|
||||||
|
|
||||||
|
|
||||||
|
def _bytes_path(data: bytes):
|
||||||
|
import io
|
||||||
|
|
||||||
|
return io.BytesIO(data)
|
||||||
|
|
||||||
|
|
||||||
|
def _ooxml_part_number(name: str) -> int:
|
||||||
|
stem = Path(name).stem
|
||||||
|
digits = "".join(character for character in stem if character.isdigit())
|
||||||
|
return int(digits) if digits else 0
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_docx(data: bytes) -> str:
|
||||||
|
root = _zip_xml(data, "word/document.xml")
|
||||||
|
paragraphs: list[str] = []
|
||||||
|
for paragraph in root.iter(f"{{{_W}}}p"):
|
||||||
|
text = "".join(node.text or "" for node in paragraph.iter(f"{{{_W}}}t"))
|
||||||
|
if text:
|
||||||
|
paragraphs.append(text)
|
||||||
|
if not paragraphs:
|
||||||
|
raise DocumentExtractionError("DOCUMENT_EXTRACTION_FAILED: DOCX has no text")
|
||||||
|
return "\n".join(paragraphs) + "\n"
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_pptx(data: bytes) -> str:
|
||||||
|
try:
|
||||||
|
with _validated_ooxml_archive(data) as archive:
|
||||||
|
names = sorted(
|
||||||
|
name
|
||||||
|
for name in archive.namelist()
|
||||||
|
if name.startswith("ppt/slides/slide") and name.endswith(".xml")
|
||||||
|
)
|
||||||
|
slides: list[str] = []
|
||||||
|
for index, name in enumerate(sorted(names, key=_ooxml_part_number), 1):
|
||||||
|
root = ET.fromstring(archive.read(name))
|
||||||
|
texts = [node.text or "" for node in root.iter(f"{{{_A}}}t")]
|
||||||
|
slides.append(f"## Slide {index}\n" + "\n".join(t for t in texts if t))
|
||||||
|
except zipfile.BadZipFile as exc:
|
||||||
|
raise DocumentExtractionError("DOCUMENT_EXTRACTION_FAILED: invalid PPTX") from exc
|
||||||
|
if not slides:
|
||||||
|
raise DocumentExtractionError("DOCUMENT_EXTRACTION_FAILED: PPTX has no slides")
|
||||||
|
return "\n\n".join(slides) + "\n"
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_xlsx(data: bytes) -> str:
|
||||||
|
try:
|
||||||
|
with _validated_ooxml_archive(data) as archive:
|
||||||
|
shared: list[str] = []
|
||||||
|
if "xl/sharedStrings.xml" in archive.namelist():
|
||||||
|
root = ET.fromstring(archive.read("xl/sharedStrings.xml"))
|
||||||
|
shared = [
|
||||||
|
"".join(node.text or "" for node in item.iter(f"{{{_S}}}t"))
|
||||||
|
for item in root.iter(f"{{{_S}}}si")
|
||||||
|
]
|
||||||
|
sheets = sorted(
|
||||||
|
name
|
||||||
|
for name in archive.namelist()
|
||||||
|
if name.startswith("xl/worksheets/sheet") and name.endswith(".xml")
|
||||||
|
)
|
||||||
|
output: list[str] = []
|
||||||
|
for index, name in enumerate(sheets, 1):
|
||||||
|
root = ET.fromstring(archive.read(name))
|
||||||
|
output.append(f"## Sheet {index}")
|
||||||
|
for row in root.iter(f"{{{_S}}}row"):
|
||||||
|
values: list[str] = []
|
||||||
|
for cell in row.iter(f"{{{_S}}}c"):
|
||||||
|
value_node = cell.find(f"{{{_S}}}v")
|
||||||
|
value = value_node.text if value_node is not None else ""
|
||||||
|
if cell.get("t") == "s" and value and value.isdigit():
|
||||||
|
shared_index = int(value)
|
||||||
|
value = shared[shared_index] if shared_index < len(shared) else value
|
||||||
|
values.append(value or "")
|
||||||
|
output.append("\t".join(values))
|
||||||
|
except zipfile.BadZipFile as exc:
|
||||||
|
raise DocumentExtractionError("DOCUMENT_EXTRACTION_FAILED: invalid XLSX") from exc
|
||||||
|
if len(output) <= 1:
|
||||||
|
raise DocumentExtractionError("DOCUMENT_EXTRACTION_FAILED: XLSX has no sheets")
|
||||||
|
return "\n".join(output) + "\n"
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_anydoc(data: bytes, extension: str) -> str:
|
||||||
|
source = ""
|
||||||
|
output = ""
|
||||||
|
try:
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=extension, delete=False) as handle:
|
||||||
|
handle.write(data)
|
||||||
|
source = handle.name
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".md", delete=False) as handle:
|
||||||
|
output = handle.name
|
||||||
|
script = (
|
||||||
|
"import pathlib,sys; import anydoc; "
|
||||||
|
"text=anydoc.to_markdown(sys.argv[1]); "
|
||||||
|
"pathlib.Path(sys.argv[2]).write_text(text, encoding='utf-8')"
|
||||||
|
)
|
||||||
|
subprocess.run(
|
||||||
|
[sys.executable, "-c", script, source, output],
|
||||||
|
check=True,
|
||||||
|
capture_output=True,
|
||||||
|
timeout=DOCUMENT_CONVERSION_TIMEOUT_SECONDS,
|
||||||
|
)
|
||||||
|
output_path = Path(output)
|
||||||
|
if output_path.stat().st_size > MAX_CONVERTED_DOCUMENT_BYTES:
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
"DOCUMENT_RESOURCE_LIMIT: converted document exceeds output budget"
|
||||||
|
)
|
||||||
|
text = output_path.read_text(encoding="utf-8")
|
||||||
|
except subprocess.TimeoutExpired as exc:
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"DOCUMENT_CONVERSION_TIMEOUT: exceeded {DOCUMENT_CONVERSION_TIMEOUT_SECONDS}s"
|
||||||
|
) from exc
|
||||||
|
except subprocess.CalledProcessError as exc:
|
||||||
|
detail = exc.stderr.decode("utf-8", errors="replace")[-1000:]
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"DOCUMENT_EXTRACTION_FAILED: converter exited {exc.returncode}: {detail}"
|
||||||
|
) from exc
|
||||||
|
except Exception as exc:
|
||||||
|
if isinstance(exc, DocumentExtractionError):
|
||||||
|
raise
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
f"DOCUMENT_EXTRACTION_FAILED: {type(exc).__name__}: {exc}"
|
||||||
|
) from exc
|
||||||
|
finally:
|
||||||
|
for temporary in (source, output):
|
||||||
|
if temporary:
|
||||||
|
try:
|
||||||
|
os.unlink(temporary)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
if not isinstance(text, str) or not text.strip():
|
||||||
|
raise DocumentExtractionError(
|
||||||
|
"DOCUMENT_EXTRACTION_FAILED: document contains no extractable text"
|
||||||
|
)
|
||||||
|
return text.rstrip("\n") + "\n"
|
||||||
@@ -4,19 +4,16 @@ The gateway package is the migration seam between UI surfaces and graph
|
|||||||
execution. CLI, TUI, channels, and future frontends should depend on this
|
execution. CLI, TUI, channels, and future frontends should depend on this
|
||||||
package for thread/run operations instead of reaching directly into
|
package for thread/run operations instead of reaching directly into
|
||||||
``sessions.py``, ``stream.events``, or the LangGraph SDK.
|
``sessions.py``, ``stream.events``, or the LangGraph SDK.
|
||||||
|
|
||||||
|
Backend implementations are attached lazily via :mod:`lazy_loader` (SPEC-1 /
|
||||||
|
PEP 562): importing the shared :mod:`.types` protocols must not cascade into
|
||||||
|
``sessions``/langgraph/langgraph_sdk, which every CLI invocation would pay.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from . import background_runs
|
from typing import TYPE_CHECKING
|
||||||
from .local import LocalGraphGateway, LocalThreadStore
|
|
||||||
from .runtime import (
|
import lazy_loader as _lazy
|
||||||
RuntimeGatewayBackend,
|
|
||||||
RuntimeGateways,
|
|
||||||
create_runtime_gateways,
|
|
||||||
)
|
|
||||||
from .server import (
|
|
||||||
LangGraphServerGateway,
|
|
||||||
LangGraphServerThreadStore,
|
|
||||||
)
|
|
||||||
from .types import (
|
from .types import (
|
||||||
DEFAULT_GRAPH_ID,
|
DEFAULT_GRAPH_ID,
|
||||||
GraphEvent,
|
GraphEvent,
|
||||||
@@ -29,6 +26,44 @@ from .types import (
|
|||||||
ThreadStore,
|
ThreadStore,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
# Static counterparts of the lazy attach below — type checkers don't
|
||||||
|
# infer names served through __getattr__.
|
||||||
|
from . import background_runs
|
||||||
|
from .local import LocalGraphGateway, LocalThreadStore
|
||||||
|
from .runtime import (
|
||||||
|
RuntimeGatewayBackend,
|
||||||
|
RuntimeGateways,
|
||||||
|
create_runtime_gateways,
|
||||||
|
)
|
||||||
|
from .server import (
|
||||||
|
LangGraphServerGateway,
|
||||||
|
LangGraphServerThreadStore,
|
||||||
|
)
|
||||||
|
|
||||||
|
__getattr__, _attach_dir, _ = _lazy.attach(
|
||||||
|
__name__,
|
||||||
|
submodules=["background_runs"],
|
||||||
|
submod_attrs={
|
||||||
|
"local": ["LocalGraphGateway", "LocalThreadStore"],
|
||||||
|
"runtime": [
|
||||||
|
"RuntimeGatewayBackend",
|
||||||
|
"RuntimeGateways",
|
||||||
|
"create_runtime_gateways",
|
||||||
|
],
|
||||||
|
"server": [
|
||||||
|
"LangGraphServerGateway",
|
||||||
|
"LangGraphServerThreadStore",
|
||||||
|
],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def __dir__() -> list[str]:
|
||||||
|
# attach() only knows the lazy names; include the eager type exports too.
|
||||||
|
return sorted(set(_attach_dir()) | set(__all__))
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"DEFAULT_GRAPH_ID",
|
"DEFAULT_GRAPH_ID",
|
||||||
"GraphEvent",
|
"GraphEvent",
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ import asyncio
|
|||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
from collections.abc import Callable, Mapping
|
from collections.abc import Callable, Mapping, Sequence
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, Protocol, TypedDict
|
from typing import TYPE_CHECKING, Protocol, TypedDict
|
||||||
|
|
||||||
@@ -75,6 +75,12 @@ class _SyncRunsClient(Protocol):
|
|||||||
|
|
||||||
def get(self, thread_id: str, run_id: str) -> Run: ...
|
def get(self, thread_id: str, run_id: str) -> Run: ...
|
||||||
|
|
||||||
|
def list(
|
||||||
|
self, thread_id: str, *, limit: int, offset: int, status: str
|
||||||
|
) -> list[Run]: ...
|
||||||
|
|
||||||
|
def cancel_many(self, *, thread_id: str, run_ids: Sequence[str]) -> object: ...
|
||||||
|
|
||||||
|
|
||||||
class SyncLangGraphClient(Protocol):
|
class SyncLangGraphClient(Protocol):
|
||||||
"""Sync subset of the LangGraph SDK used by background runs."""
|
"""Sync subset of the LangGraph SDK used by background runs."""
|
||||||
@@ -107,6 +113,14 @@ class _AsyncRunsClient(Protocol):
|
|||||||
|
|
||||||
async def get(self, thread_id: str, run_id: str) -> Run: ...
|
async def get(self, thread_id: str, run_id: str) -> Run: ...
|
||||||
|
|
||||||
|
async def list(
|
||||||
|
self, thread_id: str, *, limit: int, offset: int, status: str
|
||||||
|
) -> list[Run]: ...
|
||||||
|
|
||||||
|
async def cancel_many(
|
||||||
|
self, *, thread_id: str, run_ids: Sequence[str]
|
||||||
|
) -> object: ...
|
||||||
|
|
||||||
|
|
||||||
class AsyncLangGraphClient(Protocol):
|
class AsyncLangGraphClient(Protocol):
|
||||||
"""Async subset of the LangGraph SDK used by background runs."""
|
"""Async subset of the LangGraph SDK used by background runs."""
|
||||||
@@ -247,12 +261,97 @@ async def _aget_run_status(
|
|||||||
return run["status"]
|
return run["status"]
|
||||||
|
|
||||||
|
|
||||||
|
# Page size for enumerating a thread's runs before deletion. The SDK's
|
||||||
|
# ``runs.list`` defaults to limit=10, which would silently skip runs on
|
||||||
|
# threads with a longer history.
|
||||||
|
_RUN_CANCEL_PAGE_SIZE = 100
|
||||||
|
|
||||||
|
# Statuses worth cancelling; listed server-side so terminal history is
|
||||||
|
# never paged through.
|
||||||
|
_CANCELABLE_RUN_STATUSES = ("pending", "running")
|
||||||
|
|
||||||
|
|
||||||
|
def _cancel_thread_runs(
|
||||||
|
client: SyncLangGraphClient,
|
||||||
|
thread_id: str,
|
||||||
|
*,
|
||||||
|
name: str,
|
||||||
|
) -> None:
|
||||||
|
"""Best-effort interrupt of the thread's pending/running runs.
|
||||||
|
|
||||||
|
The server's ``threads.delete`` cascade-removes queued runs from the
|
||||||
|
registry, but it does not interrupt a run that is already executing —
|
||||||
|
cancelling first sends the interrupt control message so in-flight work
|
||||||
|
actually stops (issue #358). It also protects cleanup paths that
|
||||||
|
mutate the registry without going through ``threads.delete``. The bulk
|
||||||
|
cancel is skipped when nothing is cancellable (the server 404s on an
|
||||||
|
empty cancel set), which keeps the common terminal-only path to two
|
||||||
|
cheap filtered GETs.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
run_ids: list[str] = []
|
||||||
|
for status in _CANCELABLE_RUN_STATUSES:
|
||||||
|
offset = 0
|
||||||
|
while True:
|
||||||
|
page = client.runs.list(
|
||||||
|
thread_id,
|
||||||
|
limit=_RUN_CANCEL_PAGE_SIZE,
|
||||||
|
offset=offset,
|
||||||
|
status=status,
|
||||||
|
)
|
||||||
|
run_ids.extend(run["run_id"] for run in page)
|
||||||
|
if len(page) < _RUN_CANCEL_PAGE_SIZE:
|
||||||
|
break
|
||||||
|
offset += _RUN_CANCEL_PAGE_SIZE
|
||||||
|
if run_ids:
|
||||||
|
client.runs.cancel_many(
|
||||||
|
thread_id=thread_id, run_ids=list(dict.fromkeys(run_ids))
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to cancel %s runs on thread %s", name, thread_id, exc_info=True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _acancel_thread_runs(
|
||||||
|
client: AsyncLangGraphClient,
|
||||||
|
thread_id: str,
|
||||||
|
*,
|
||||||
|
name: str,
|
||||||
|
) -> None:
|
||||||
|
"""Async variant of :func:`_cancel_thread_runs`."""
|
||||||
|
try:
|
||||||
|
run_ids: list[str] = []
|
||||||
|
for status in _CANCELABLE_RUN_STATUSES:
|
||||||
|
offset = 0
|
||||||
|
while True:
|
||||||
|
page = await client.runs.list(
|
||||||
|
thread_id,
|
||||||
|
limit=_RUN_CANCEL_PAGE_SIZE,
|
||||||
|
offset=offset,
|
||||||
|
status=status,
|
||||||
|
)
|
||||||
|
run_ids.extend(run["run_id"] for run in page)
|
||||||
|
if len(page) < _RUN_CANCEL_PAGE_SIZE:
|
||||||
|
break
|
||||||
|
offset += _RUN_CANCEL_PAGE_SIZE
|
||||||
|
if run_ids:
|
||||||
|
await client.runs.cancel_many(
|
||||||
|
thread_id=thread_id, run_ids=list(dict.fromkeys(run_ids))
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to cancel %s runs on thread %s", name, thread_id, exc_info=True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _delete_thread(
|
def _delete_thread(
|
||||||
client: SyncLangGraphClient,
|
client: SyncLangGraphClient,
|
||||||
thread_id: str,
|
thread_id: str,
|
||||||
*,
|
*,
|
||||||
name: str,
|
name: str,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
_cancel_thread_runs(client, thread_id, name=name)
|
||||||
try:
|
try:
|
||||||
client.threads.delete(thread_id)
|
client.threads.delete(thread_id)
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -265,6 +364,7 @@ async def _adelete_thread(
|
|||||||
*,
|
*,
|
||||||
name: str,
|
name: str,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
await _acancel_thread_runs(client, thread_id, name=name)
|
||||||
try:
|
try:
|
||||||
await client.threads.delete(thread_id)
|
await client.threads.delete(thread_id)
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|||||||
@@ -19,6 +19,8 @@ from .types import (
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from langgraph.graph.state import CompiledStateGraph
|
from langgraph.graph.state import CompiledStateGraph
|
||||||
|
|
||||||
|
from ..middleware.events import SessionEvents
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
class LocalThreadStore:
|
class LocalThreadStore:
|
||||||
@@ -61,9 +63,16 @@ class LocalThreadStore:
|
|||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
class LocalGraphGateway:
|
class LocalGraphGateway:
|
||||||
"""Gateway backed by the current in-process graph and session helpers."""
|
"""Gateway backed by the current in-process graph and session helpers.
|
||||||
|
|
||||||
|
``events`` is the frontend/session event sink for this runtime — normally
|
||||||
|
the same instance injected into the agent's middleware. If it is ``None``,
|
||||||
|
``stream_agent_events`` creates a per-run session sink and binds it for
|
||||||
|
default main-agent middleware via ``RunScopedEventSink``.
|
||||||
|
"""
|
||||||
|
|
||||||
thread_store: ThreadStore = field(default_factory=LocalThreadStore)
|
thread_store: ThreadStore = field(default_factory=LocalThreadStore)
|
||||||
|
events: SessionEvents | None = None
|
||||||
|
|
||||||
async def create_thread(
|
async def create_thread(
|
||||||
self,
|
self,
|
||||||
@@ -155,6 +164,8 @@ class LocalGraphGateway:
|
|||||||
request.thread_id,
|
request.thread_id,
|
||||||
metadata=request.metadata,
|
metadata=request.metadata,
|
||||||
media=request.media,
|
media=request.media,
|
||||||
|
events=self.events,
|
||||||
|
configurable_extra=request.configurable_extra,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
async for event in inner:
|
async for event in inner:
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Literal
|
from typing import TYPE_CHECKING, Literal
|
||||||
|
|
||||||
from langgraph_sdk import get_client
|
from langgraph_sdk import get_client
|
||||||
from langgraph_sdk.client import LangGraphClient
|
from langgraph_sdk.client import LangGraphClient
|
||||||
@@ -14,6 +14,9 @@ from .server import (
|
|||||||
)
|
)
|
||||||
from .types import GraphGateway, ThreadStore
|
from .types import GraphGateway, ThreadStore
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from ..middleware.events import SessionEvents
|
||||||
|
|
||||||
RuntimeGatewayBackend = Literal["local", "langgraph_server"]
|
RuntimeGatewayBackend = Literal["local", "langgraph_server"]
|
||||||
|
|
||||||
|
|
||||||
@@ -32,8 +35,14 @@ def create_runtime_gateways(
|
|||||||
graph_id: str = DEFAULT_GRAPH_ID,
|
graph_id: str = DEFAULT_GRAPH_ID,
|
||||||
headers: dict[str, str] | None = None,
|
headers: dict[str, str] | None = None,
|
||||||
langgraph_client: LangGraphClient | None = None,
|
langgraph_client: LangGraphClient | None = None,
|
||||||
|
events: SessionEvents | None = None,
|
||||||
) -> RuntimeGateways:
|
) -> RuntimeGateways:
|
||||||
"""Create gateway handles for CLI/TUI/serve execution."""
|
"""Create gateway handles for CLI/TUI/serve execution.
|
||||||
|
|
||||||
|
``events`` is the frontend event sink; it is attached to the local gateway
|
||||||
|
so the streaming path shares the same sink instance the frontend injects
|
||||||
|
into the agent's middleware. Server backends ignore it (headless).
|
||||||
|
"""
|
||||||
if backend == "langgraph_server":
|
if backend == "langgraph_server":
|
||||||
if base_url is None and langgraph_client is None:
|
if base_url is None and langgraph_client is None:
|
||||||
raise ValueError("base_url is required for langgraph_server gateways")
|
raise ValueError("base_url is required for langgraph_server gateways")
|
||||||
@@ -59,5 +68,5 @@ def create_runtime_gateways(
|
|||||||
|
|
||||||
return RuntimeGateways(
|
return RuntimeGateways(
|
||||||
thread_store=local_thread_store,
|
thread_store=local_thread_store,
|
||||||
graph_gateway=LocalGraphGateway(thread_store=local_thread_store),
|
graph_gateway=LocalGraphGateway(thread_store=local_thread_store, events=events),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -7,7 +7,10 @@ import uuid
|
|||||||
from collections.abc import AsyncIterator, Mapping
|
from collections.abc import AsyncIterator, Mapping
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import UTC, datetime
|
from datetime import UTC, datetime
|
||||||
from typing import Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from ..middleware.events import SessionEvents
|
||||||
|
|
||||||
from langchain_core.messages import BaseMessage, convert_to_messages, messages_from_dict
|
from langchain_core.messages import BaseMessage, convert_to_messages, messages_from_dict
|
||||||
from langgraph.types import Command
|
from langgraph.types import Command
|
||||||
@@ -25,6 +28,7 @@ from ..stream.events import (
|
|||||||
)
|
)
|
||||||
from ..stream.summarization import _find_summarization_event_payload
|
from ..stream.summarization import _find_summarization_event_payload
|
||||||
from ..stream.v3_payloads import _as_raw_map, _event_namespace
|
from ..stream.v3_payloads import _as_raw_map, _event_namespace
|
||||||
|
from .background_runs import _acancel_thread_runs
|
||||||
from .types import (
|
from .types import (
|
||||||
DEFAULT_GRAPH_ID,
|
DEFAULT_GRAPH_ID,
|
||||||
GraphEvent,
|
GraphEvent,
|
||||||
@@ -321,6 +325,10 @@ class LangGraphServerThreadStore(ThreadStore):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
async def delete_thread(self, thread_id: str) -> bool:
|
async def delete_thread(self, thread_id: str) -> bool:
|
||||||
|
# Interrupt live runs first: the server's cascade delete clears
|
||||||
|
# queued runs from the registry but does not stop a run that is
|
||||||
|
# already executing (issue #358).
|
||||||
|
await _acancel_thread_runs(self.client, thread_id, name="thread delete")
|
||||||
try:
|
try:
|
||||||
await self.client.threads.delete(thread_id)
|
await self.client.threads.delete(thread_id)
|
||||||
except NotFoundError:
|
except NotFoundError:
|
||||||
@@ -448,6 +456,7 @@ class LangGraphServerGateway:
|
|||||||
thread_store: LangGraphServerThreadStore
|
thread_store: LangGraphServerThreadStore
|
||||||
graph_id: str = DEFAULT_GRAPH_ID
|
graph_id: str = DEFAULT_GRAPH_ID
|
||||||
interrupt_wait_seconds: float = 5.0
|
interrupt_wait_seconds: float = 5.0
|
||||||
|
events: SessionEvents | None = None
|
||||||
|
|
||||||
def _target_graph_id(self, target: GraphTarget | None = None) -> str:
|
def _target_graph_id(self, target: GraphTarget | None = None) -> str:
|
||||||
return target.graph_id if target is not None else self.graph_id
|
return target.graph_id if target is not None else self.graph_id
|
||||||
|
|||||||
@@ -6,13 +6,15 @@ from collections.abc import AsyncIterator
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, Any, Protocol, TypeAlias
|
from typing import TYPE_CHECKING, Any, Protocol, TypeAlias
|
||||||
|
|
||||||
from langgraph.types import Command
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from langgraph.graph.state import CompiledStateGraph
|
from langgraph.graph.state import CompiledStateGraph
|
||||||
|
from langgraph.types import Command
|
||||||
|
|
||||||
|
from ..middleware.events import SessionEvents
|
||||||
|
|
||||||
GraphEvent: TypeAlias = dict[str, Any]
|
GraphEvent: TypeAlias = dict[str, Any]
|
||||||
GraphRunInput: TypeAlias = str | Command
|
# String alias keeps this module langgraph-free at import time (~950 modules).
|
||||||
|
GraphRunInput: TypeAlias = "str | Command"
|
||||||
GraphStateValues: TypeAlias = dict[str, Any]
|
GraphStateValues: TypeAlias = dict[str, Any]
|
||||||
DEFAULT_GRAPH_ID = "EvoScientist"
|
DEFAULT_GRAPH_ID = "EvoScientist"
|
||||||
|
|
||||||
@@ -39,6 +41,13 @@ class RunRequest:
|
|||||||
metadata: dict[str, Any] | None = None
|
metadata: dict[str, Any] | None = None
|
||||||
media: list[str] | None = None
|
media: list[str] | None = None
|
||||||
target: GraphTarget | None = None
|
target: GraphTarget | None = None
|
||||||
|
configurable_extra: dict[str, Any] | None = None
|
||||||
|
"""Extra keys to merge into the LangGraph ``configurable`` dict alongside
|
||||||
|
``thread_id`` — e.g. ``{"active_teams": [...]}`` from the TUI
|
||||||
|
``/expert`` command. WebUI callers achieve the same effect via
|
||||||
|
``langgraph_sdk``'s ``config.configurable`` on their own; this field is
|
||||||
|
the local-gateway equivalent so CLI / TUI / headless serve can bias
|
||||||
|
the run identically."""
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
@@ -94,6 +103,8 @@ class ThreadStore(Protocol):
|
|||||||
class GraphGateway(Protocol):
|
class GraphGateway(Protocol):
|
||||||
"""One authority for graph runs and thread lifecycle operations."""
|
"""One authority for graph runs and thread lifecycle operations."""
|
||||||
|
|
||||||
|
events: SessionEvents | None
|
||||||
|
|
||||||
async def create_thread(
|
async def create_thread(
|
||||||
self,
|
self,
|
||||||
target: GraphTarget | None = None,
|
target: GraphTarget | None = None,
|
||||||
|
|||||||
@@ -0,0 +1,17 @@
|
|||||||
|
"""Authentication headers for LangGraph-to-Gateway internal calls."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
|
||||||
|
|
||||||
|
def internal_service_token() -> str:
|
||||||
|
return (
|
||||||
|
os.environ.get("EVOSCIENTIST_BACKEND_SERVICE_TOKEN", "").strip()
|
||||||
|
or os.environ.get("AI4SCI_EVO_RUNTIME_GRANT_SECRET", "").strip()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def internal_service_headers() -> dict[str, str]:
|
||||||
|
token = internal_service_token()
|
||||||
|
return {"X-Ai4Sci-Service-Token": token} if token else {}
|
||||||
@@ -29,10 +29,19 @@ from EvoScientist.memory.agents import (
|
|||||||
)
|
)
|
||||||
from EvoScientist.memory.types import MemorySourceType
|
from EvoScientist.memory.types import MemorySourceType
|
||||||
from EvoScientist.subagents._factory import build_async_subagent_graph
|
from EvoScientist.subagents._factory import build_async_subagent_graph
|
||||||
|
from EvoScientist.subagents.expert_container_async import (
|
||||||
|
build_expert_container_async_graph,
|
||||||
|
)
|
||||||
|
|
||||||
writing_agent = build_async_subagent_graph("writing-agent")
|
writing_agent = build_async_subagent_graph("writing-agent")
|
||||||
data_analysis_agent = build_async_subagent_graph("data-analysis-agent")
|
data_analysis_agent = build_async_subagent_graph("data-analysis-agent")
|
||||||
scheduler = build_async_subagent_graph("scheduler")
|
scheduler = build_async_subagent_graph("scheduler")
|
||||||
|
# Generic async container for expert-skill dispatch. One graph, parameterised
|
||||||
|
# per invocation by the ``skill_name`` payload the main agent passes through
|
||||||
|
# ``EvoAsyncSubAgentMiddleware.start_async_task``. Any installed expert skill
|
||||||
|
# dispatches through this graph; the loader middleware resolves the skill
|
||||||
|
# body at model-call time.
|
||||||
|
expert_container_async = build_expert_container_async_graph()
|
||||||
evomemory_subagent_worker = build_memory_worker_graph(MemorySourceType.SUBAGENT)
|
evomemory_subagent_worker = build_memory_worker_graph(MemorySourceType.SUBAGENT)
|
||||||
evomemory_turn_worker = build_memory_worker_graph(MemorySourceType.TURN)
|
evomemory_turn_worker = build_memory_worker_graph(MemorySourceType.TURN)
|
||||||
evomemory_observation_linker = build_observation_linker_graph()
|
evomemory_observation_linker = build_observation_linker_graph()
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ import asyncio
|
|||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import secrets
|
import secrets
|
||||||
from typing import Any
|
from typing import Any, cast
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from starlette.applications import Starlette
|
from starlette.applications import Starlette
|
||||||
@@ -104,6 +104,11 @@ async def recoverable_run_capabilities(_request: Request) -> JSONResponse:
|
|||||||
{
|
{
|
||||||
"version": 1,
|
"version": 1,
|
||||||
"deterministic_run_id": True,
|
"deterministic_run_id": True,
|
||||||
|
# The dev adapter uses process-local Runs storage. Checkpoint
|
||||||
|
# durability does not make execution identity restart-safe.
|
||||||
|
"durable_run_identity": False,
|
||||||
|
"worker_exit_confirmation": True,
|
||||||
|
"run_not_found_proves_absence": False,
|
||||||
"stream_resumable": True,
|
"stream_resumable": True,
|
||||||
"durability_sync": True,
|
"durability_sync": True,
|
||||||
"multitask_enqueue": True,
|
"multitask_enqueue": True,
|
||||||
@@ -270,16 +275,149 @@ async def bind_workspace_run(request: Request) -> JSONResponse:
|
|||||||
return JSONResponse(_run_payload(run))
|
return JSONResponse(_run_payload(run))
|
||||||
|
|
||||||
|
|
||||||
|
def _interrupt_ids(value: Any) -> set[str]:
|
||||||
|
"""从 __interrupt__ 写入值里取出中断标识(Interrupt 对象或字典都兼容)。"""
|
||||||
|
|
||||||
|
items = value if isinstance(value, (list, tuple)) else [value]
|
||||||
|
found: set[str] = set()
|
||||||
|
for item in items:
|
||||||
|
ident = getattr(item, "id", None)
|
||||||
|
if ident is None and isinstance(item, dict):
|
||||||
|
ident = item.get("id") or item.get("interrupt_id")
|
||||||
|
if ident:
|
||||||
|
found.add(str(ident))
|
||||||
|
return found
|
||||||
|
|
||||||
|
|
||||||
|
def _anchor_has_interrupt(checkpoint: Any, expected: str) -> bool:
|
||||||
|
"""锚点处是否仍承载待审批中断(并核对中断标识)。
|
||||||
|
|
||||||
|
这是"按锚点恢复"的准入判据:只要承载该中断的检查点写入还在 PG,暂停就仍然
|
||||||
|
有效 —— 与运行时进程是否重启过、距暂停多久都无关。
|
||||||
|
"""
|
||||||
|
|
||||||
|
found: set[str] = set()
|
||||||
|
for write in getattr(checkpoint, "pending_writes", None) or ():
|
||||||
|
try:
|
||||||
|
if len(write) < 3 or str(write[1]) != "__interrupt__":
|
||||||
|
continue
|
||||||
|
found |= _interrupt_ids(write[2])
|
||||||
|
except TypeError:
|
||||||
|
continue
|
||||||
|
if not found:
|
||||||
|
return False
|
||||||
|
# 网关未提供标识(或回退值 default)时只做存在性判定。
|
||||||
|
if not expected or expected == "default":
|
||||||
|
return True
|
||||||
|
return expected in found
|
||||||
|
|
||||||
|
|
||||||
|
async def _compatible_checkpoint(conn, thread_id: str, assistant_id: str,
|
||||||
|
config: dict, anchor: dict | None = None
|
||||||
|
) -> tuple[bool, bool]:
|
||||||
|
"""Read through the API-owned saver and graph factory, never execute here."""
|
||||||
|
from langgraph_api._checkpointer import get_checkpointer
|
||||||
|
from langgraph_api.graph import get_graph, graph_exists
|
||||||
|
from langgraph_api.store import get_store
|
||||||
|
|
||||||
|
saver = await get_checkpointer(conn=conn)
|
||||||
|
read_config = {**config, "configurable": {
|
||||||
|
**config.get("configurable", {}), "thread_id": thread_id,
|
||||||
|
"checkpoint_ns": str((anchor or {}).get("checkpoint_ns") or ""),
|
||||||
|
}}
|
||||||
|
if anchor and str(anchor.get("checkpoint_id") or ""):
|
||||||
|
# 按锚点恢复:调用方(网关)指定了承载该中断的祖先检查点,暂停时就已落库。
|
||||||
|
# 只有该锚点处确实还有待审批写入才放行 —— 不依赖运行时当前头部。
|
||||||
|
read_config["configurable"]["checkpoint_id"] = str(anchor["checkpoint_id"])
|
||||||
|
else:
|
||||||
|
# Admission always checks the current head, never a caller-selected ancestor.
|
||||||
|
read_config["configurable"].pop("checkpoint_id", None)
|
||||||
|
checkpoint = await saver.aget_tuple(read_config)
|
||||||
|
if checkpoint is None:
|
||||||
|
return False, False
|
||||||
|
graph_id = assistant_id
|
||||||
|
if not graph_exists(graph_id):
|
||||||
|
from langgraph_runtime.ops import Assistants
|
||||||
|
from langgraph_api.utils import fetchone
|
||||||
|
assistant = await fetchone(await Assistants.get(conn, UUID(assistant_id)))
|
||||||
|
graph_id = assistant["graph_id"]
|
||||||
|
if checkpoint.metadata.get("graph_id", graph_id) != graph_id:
|
||||||
|
raise ValueError("checkpoint graph mismatch")
|
||||||
|
if anchor and str(anchor.get("checkpoint_id") or ""):
|
||||||
|
return True, _anchor_has_interrupt(
|
||||||
|
checkpoint, str(anchor.get("interrupt_id") or "")
|
||||||
|
)
|
||||||
|
# get_graph enters coroutine/async-context-manager factories and binds the
|
||||||
|
# same API saver used by the worker. aget_state also validates delta seeds.
|
||||||
|
async with get_graph(graph_id, read_config, checkpointer=saver,
|
||||||
|
store=await get_store(), access_context="threads.read") as graph:
|
||||||
|
state = await graph.aget_state(read_config, subgraphs=True)
|
||||||
|
pending = bool(getattr(state, "interrupts", ())) or any(
|
||||||
|
task.interrupts for task in state.tasks
|
||||||
|
)
|
||||||
|
return True, pending
|
||||||
|
|
||||||
|
|
||||||
|
async def _has_legacy_history(conn, thread_id: str) -> bool:
|
||||||
|
"""Resolve Web history from PostgreSQL and the Runtime registry only."""
|
||||||
|
dsn = os.getenv("EVOSCIENTIST_WEB_CHECKPOINT_DSN", "")
|
||||||
|
if dsn:
|
||||||
|
from psycopg import AsyncConnection
|
||||||
|
async with await AsyncConnection.connect(
|
||||||
|
dsn, autocommit=True, connect_timeout=5,
|
||||||
|
options="-c default_transaction_read_only=on -c search_path=public",
|
||||||
|
) as pg:
|
||||||
|
async with pg.cursor() as cursor:
|
||||||
|
await cursor.execute("SELECT EXISTS(SELECT 1 FROM threads WHERE id=%s::uuid)",
|
||||||
|
(thread_id,))
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
if row and row[0]:
|
||||||
|
return True
|
||||||
|
from langgraph_api.utils import fetchone
|
||||||
|
from langgraph_runtime.ops import Threads
|
||||||
|
try:
|
||||||
|
await fetchone(await Threads.get(conn, UUID(thread_id)))
|
||||||
|
except Exception as exc:
|
||||||
|
if getattr(exc, "status_code", None) != 404:
|
||||||
|
raise
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
async def _history_admission(conn, thread_id, assistant_id, config, operation, history,
|
||||||
|
anchor: dict | None = None):
|
||||||
|
exists, pending = await _compatible_checkpoint(
|
||||||
|
conn, thread_id, assistant_id, config, anchor
|
||||||
|
)
|
||||||
|
if operation == "resume":
|
||||||
|
return "resume" if exists and pending else "CHECKPOINT_RESUME_UNAVAILABLE"
|
||||||
|
if pending:
|
||||||
|
return "THREAD_AWAITING_INPUT"
|
||||||
|
if exists:
|
||||||
|
return "append"
|
||||||
|
if history is not None:
|
||||||
|
return "initialize"
|
||||||
|
return "HISTORY_REQUIRED" if await _has_legacy_history(conn, thread_id) else "new"
|
||||||
|
|
||||||
|
|
||||||
async def create_recoverable_run(request: Request) -> JSONResponse:
|
async def create_recoverable_run(request: Request) -> JSONResponse:
|
||||||
"""Create a LangGraph Run with a caller-owned deterministic UUID.
|
"""Create a LangGraph Run with a caller-owned deterministic UUID.
|
||||||
|
|
||||||
LangGraph's public create endpoint always generates its own UUID. This
|
LangGraph's public create endpoint always generates its own UUID. This
|
||||||
adapter performs lookup and insertion while holding the process-wide run
|
adapter performs lookup and insertion while holding the process-wide run
|
||||||
creation lock and passes the durable request UUID to ``create_valid_run``.
|
creation lock and passes the durable request UUID to ``create_valid_run``.
|
||||||
Retrying after a lost HTTP response therefore cannot create another Run.
|
This prevents duplicate creation only while this process retains the Run.
|
||||||
|
A lost response followed by a restart MUST NOT be retried as a fresh create:
|
||||||
|
the dev backend forgets Runs, and a 404 is not evidence of non-execution.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if request.headers.get("x-auth-scheme") != "langsmith":
|
from EvoScientist.internal_service import internal_service_token
|
||||||
|
|
||||||
|
token = internal_service_token()
|
||||||
|
if not token:
|
||||||
|
return JSONResponse({"code": "WORKSPACE_SERVICE_UNAVAILABLE"}, status_code=503)
|
||||||
|
header = request.headers.get("authorization", "")
|
||||||
|
if not header.startswith("Bearer ") or not secrets.compare_digest(header[7:], token):
|
||||||
return JSONResponse({"code": "UNAUTHORIZED"}, status_code=401)
|
return JSONResponse({"code": "UNAUTHORIZED"}, status_code=401)
|
||||||
value = await request.json()
|
value = await request.json()
|
||||||
if not isinstance(value, dict):
|
if not isinstance(value, dict):
|
||||||
@@ -303,17 +441,30 @@ async def create_recoverable_run(request: Request) -> JSONResponse:
|
|||||||
if operation == "resume":
|
if operation == "resume":
|
||||||
if (
|
if (
|
||||||
value.get("input") is not None
|
value.get("input") is not None
|
||||||
|
or value.get("history") is not None
|
||||||
or not isinstance(command, dict)
|
or not isinstance(command, dict)
|
||||||
or set(command) != {"resume"}
|
or set(command) != {"resume"}
|
||||||
):
|
):
|
||||||
return JSONResponse({"code": "INVALID_RESUME_REQUEST"}, status_code=400)
|
return JSONResponse({"code": "INVALID_RESUME_REQUEST"}, status_code=400)
|
||||||
elif command is not None:
|
elif command is not None:
|
||||||
return JSONResponse({"code": "INVALID_START_REQUEST"}, status_code=400)
|
return JSONResponse({"code": "INVALID_START_REQUEST"}, status_code=400)
|
||||||
|
# 按锚点恢复(可选,仅 resume):网关在暂停时把"承载该中断的检查点"落库,
|
||||||
|
# 继续时回传。这样续接只依赖该检查点仍在 PG,而不依赖运行时当前头部。
|
||||||
|
anchor = value.get("anchor") if operation == "resume" else None
|
||||||
|
if not isinstance(anchor, dict) or not str(anchor.get("checkpoint_id") or ""):
|
||||||
|
anchor = None
|
||||||
|
history = value.get("history")
|
||||||
|
# History validation is performed under the create lock, after idempotent
|
||||||
|
# lookup. Only that branch can attest this request did not create a Run.
|
||||||
|
|
||||||
from langgraph_api.models.run import Runs, create_valid_run
|
from langgraph_api.models.run import Runs, create_valid_run
|
||||||
from langgraph_api.utils import fetchone
|
from langgraph_api.utils import fetchone
|
||||||
from langgraph_runtime.database import connect
|
from langgraph_runtime.database import connect
|
||||||
|
|
||||||
|
from EvoScientist.langgraph_dev import worker_exit
|
||||||
|
|
||||||
|
worker_exit.install()
|
||||||
|
|
||||||
payload = {
|
payload = {
|
||||||
"assistant_id": assistant_id,
|
"assistant_id": assistant_id,
|
||||||
"input": value.get("input"),
|
"input": value.get("input"),
|
||||||
@@ -331,6 +482,19 @@ async def create_recoverable_run(request: Request) -> JSONResponse:
|
|||||||
"run_request_id": run_request_id,
|
"run_request_id": run_request_id,
|
||||||
"request_hash": request_hash,
|
"request_hash": request_hash,
|
||||||
}
|
}
|
||||||
|
# The admission read and the worker must address the same checkpoint.
|
||||||
|
# Without an anchor that is the current head (forking from an ancestor is
|
||||||
|
# not part of the recoverable-run contract); with an anchor it is the
|
||||||
|
# ancestor recorded at pause time — which is exactly what makes the
|
||||||
|
# continuation independent of restarts and elapsed time.
|
||||||
|
configurable = dict(payload["config"].get("configurable", {}))
|
||||||
|
for key in ("checkpoint_id", "checkpoint_map"):
|
||||||
|
configurable.pop(key, None)
|
||||||
|
configurable.update(thread_id=thread_id, checkpoint_ns="")
|
||||||
|
if anchor is not None:
|
||||||
|
configurable["checkpoint_id"] = str(anchor["checkpoint_id"])
|
||||||
|
configurable["checkpoint_ns"] = str(anchor.get("checkpoint_ns") or "")
|
||||||
|
payload["config"] = {**payload["config"], "configurable": configurable}
|
||||||
async with _recoverable_run_lock:
|
async with _recoverable_run_lock:
|
||||||
async with connect() as conn:
|
async with connect() as conn:
|
||||||
existing_iter = await Runs.get(conn, run_id, thread_id=UUID(thread_id))
|
existing_iter = await Runs.get(conn, run_id, thread_id=UUID(thread_id))
|
||||||
@@ -353,6 +517,51 @@ async def create_recoverable_run(request: Request) -> JSONResponse:
|
|||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
{"run_id": str(existing["run_id"]), "status": existing["status"], "created": False}
|
{"run_id": str(existing["run_id"]), "status": existing["status"], "created": False}
|
||||||
)
|
)
|
||||||
|
try:
|
||||||
|
admission = await _history_admission(
|
||||||
|
conn, thread_id, assistant_id, payload["config"], operation, history,
|
||||||
|
anchor,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
return JSONResponse({
|
||||||
|
"code": "CHECKPOINT_UNAVAILABLE", "create_disposition": "not_created",
|
||||||
|
"run_id": str(run_id), "request_hash": request_hash,
|
||||||
|
}, status_code=503)
|
||||||
|
if admission not in {"new", "initialize", "append", "resume"}:
|
||||||
|
return JSONResponse({
|
||||||
|
"code": admission, "create_disposition": "not_created",
|
||||||
|
"run_id": str(run_id), "request_hash": request_hash,
|
||||||
|
}, status_code=409)
|
||||||
|
if history is not None:
|
||||||
|
from EvoScientist.llm.history_rebuild import committed_history_input
|
||||||
|
|
||||||
|
try:
|
||||||
|
if not isinstance(history, dict):
|
||||||
|
raise ValueError("history must be an object")
|
||||||
|
import hashlib
|
||||||
|
|
||||||
|
history_hash = hashlib.sha256(json.dumps(
|
||||||
|
history, ensure_ascii=False, sort_keys=True, separators=(",", ":"),
|
||||||
|
).encode()).hexdigest()
|
||||||
|
if value.get("history_hash") != history_hash:
|
||||||
|
raise ValueError("history hash mismatch")
|
||||||
|
payload["input"] = committed_history_input(
|
||||||
|
history, payload["input"], thread_id=thread_id, run_id=str(run_id),
|
||||||
|
checkpoint_exists=admission == "append",
|
||||||
|
)
|
||||||
|
if admission == "initialize":
|
||||||
|
payload["metadata"].update(
|
||||||
|
history_hash=history_hash,
|
||||||
|
history_revision=history["conversation_revision"],
|
||||||
|
history_schema=history["schema"],
|
||||||
|
)
|
||||||
|
except (ValueError, TypeError, KeyError, AttributeError):
|
||||||
|
return JSONResponse({
|
||||||
|
"code": "INVALID_HISTORY_REQUEST",
|
||||||
|
"create_disposition": "not_created",
|
||||||
|
"run_id": str(run_id), "request_hash": request_hash,
|
||||||
|
}, status_code=400)
|
||||||
|
worker_exit.reserve(thread_id, str(run_id), request_hash)
|
||||||
created = await create_valid_run(
|
created = await create_valid_run(
|
||||||
conn,
|
conn,
|
||||||
thread_id,
|
thread_id,
|
||||||
@@ -366,9 +575,110 @@ async def create_recoverable_run(request: Request) -> JSONResponse:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def cancel_recoverable_run(request: Request) -> JSONResponse:
|
||||||
|
from EvoScientist.internal_service import internal_service_token
|
||||||
|
|
||||||
|
token = internal_service_token()
|
||||||
|
if not token:
|
||||||
|
return JSONResponse({"code": "WORKSPACE_SERVICE_UNAVAILABLE"}, status_code=503)
|
||||||
|
header = request.headers.get("authorization", "")
|
||||||
|
if not header.startswith("Bearer ") or not secrets.compare_digest(header[7:], token):
|
||||||
|
return JSONResponse({"code": "UNAUTHORIZED"}, status_code=401)
|
||||||
|
from EvoScientist.langgraph_dev import worker_exit
|
||||||
|
from langgraph_api.models.run import Runs
|
||||||
|
from langgraph_runtime.database import connect
|
||||||
|
|
||||||
|
thread_id = UUID(request.path_params["thread_id"])
|
||||||
|
run_id = UUID(request.path_params["run_id"])
|
||||||
|
# This service principal controls only pairs admitted by authenticated create.
|
||||||
|
if not worker_exit.is_reserved(str(thread_id), str(run_id)):
|
||||||
|
return JSONResponse({"code": "RUN_NOT_AUTHORIZED"}, status_code=404)
|
||||||
|
# Close worker admission before notifying the original runtime control queue.
|
||||||
|
initial = worker_exit.cancel_and_inspect(str(thread_id), str(run_id))
|
||||||
|
if initial.get("execution_exited") is not True:
|
||||||
|
async with connect() as conn:
|
||||||
|
try:
|
||||||
|
await Runs.cancel(conn, [run_id], thread_id=thread_id, action="interrupt")
|
||||||
|
except Exception as exc:
|
||||||
|
if getattr(exc, "status_code", None) not in {404, 409}:
|
||||||
|
raise
|
||||||
|
receipt = await worker_exit.wait_for_exit(str(thread_id), str(run_id))
|
||||||
|
if receipt.get("execution_exited") is True:
|
||||||
|
async with connect() as conn:
|
||||||
|
try:
|
||||||
|
await Runs.delete(cast(Any, conn), run_id, thread_id=thread_id)
|
||||||
|
except Exception as exc:
|
||||||
|
if getattr(exc, "status_code", None) != 404:
|
||||||
|
raise
|
||||||
|
receipt = {**receipt, "checkpoint_cleanup": "completed"}
|
||||||
|
return JSONResponse(receipt)
|
||||||
|
|
||||||
|
|
||||||
|
async def get_teams(_request: Request) -> JSONResponse:
|
||||||
|
"""Return installed expert skills as ``{teams: [...]}`` for the WebUI gallery.
|
||||||
|
|
||||||
|
A "team" in the WebUI vocabulary is an installed expert skill — a skill
|
||||||
|
directory carrying a sibling ``EXPERT.md`` (or, on the deprecated path,
|
||||||
|
``type: expert`` SKILL.md frontmatter). The response is a curated,
|
||||||
|
gallery-safe projection: name + description, plus optional ``byline`` /
|
||||||
|
``capability_tags`` / ``avatar_hint`` when the skill populates them.
|
||||||
|
|
||||||
|
Cards for experts on the current contract carry name + description only:
|
||||||
|
the decoration fields were actor metadata in SKILL.md frontmatter, which
|
||||||
|
that contract removes rather than relocates (``EXPERT.md`` has no
|
||||||
|
frontmatter to hold them). The omit-when-unpopulated projection below is
|
||||||
|
what makes those cards degrade rather than break; restoring richer cards
|
||||||
|
means sourcing decoration from index metadata, not re-adding frontmatter
|
||||||
|
fields.
|
||||||
|
|
||||||
|
Backend implementation details (SKILL.md body / system prompt, role
|
||||||
|
line, tool list, source tier, filesystem path,
|
||||||
|
tags) are intentionally NOT projected. The gallery only needs
|
||||||
|
identity + descriptor fields to render the card; anything richer
|
||||||
|
belongs in a dedicated info endpoint.
|
||||||
|
|
||||||
|
Sourced from ``list_expert_skills(include_system=True)`` so
|
||||||
|
first-party experts shipped as builtin skills surface alongside
|
||||||
|
workspace/global installs.
|
||||||
|
|
||||||
|
Offloaded to a thread because the skill loader does synchronous
|
||||||
|
filesystem walking + yaml parsing, which langgraph-dev's
|
||||||
|
``blockbuster`` middleware refuses on the async event loop.
|
||||||
|
|
||||||
|
Response shape (each entry): ``{name, description, byline?,
|
||||||
|
capability_tags?, avatar_hint?}`` — the WebUI gallery consumes these.
|
||||||
|
"""
|
||||||
|
from EvoScientist.tools.skills_manager import list_expert_skills
|
||||||
|
|
||||||
|
experts = await asyncio.to_thread(list_expert_skills, True)
|
||||||
|
teams = []
|
||||||
|
for info in experts:
|
||||||
|
entry = {
|
||||||
|
"name": info.name,
|
||||||
|
"description": info.description,
|
||||||
|
}
|
||||||
|
# Optional gallery fields — omit when unpopulated so the WebUI
|
||||||
|
# card degrades gracefully (SkillInfo defaults `byline` /
|
||||||
|
# `avatar_hint` to "" and `capability_tags` to [], which we
|
||||||
|
# treat as "not declared").
|
||||||
|
if info.byline:
|
||||||
|
entry["byline"] = info.byline
|
||||||
|
if info.capability_tags:
|
||||||
|
entry["capability_tags"] = list(info.capability_tags)
|
||||||
|
if info.avatar_hint:
|
||||||
|
entry["avatar_hint"] = info.avatar_hint
|
||||||
|
teams.append(entry)
|
||||||
|
return JSONResponse({"teams": teams})
|
||||||
|
|
||||||
|
|
||||||
app = Starlette(
|
app = Starlette(
|
||||||
routes=[
|
routes=[
|
||||||
Route("/api/models", get_models, methods=["GET"]),
|
Route("/api/models", get_models, methods=["GET"]),
|
||||||
|
Route(
|
||||||
|
"/api/ai4sci/recoverable-runs/{thread_id}/{run_id}/cancel",
|
||||||
|
cancel_recoverable_run,
|
||||||
|
methods=["POST"],
|
||||||
|
),
|
||||||
Route(
|
Route(
|
||||||
"/api/ai4sci/recoverable-runs/capabilities",
|
"/api/ai4sci/recoverable-runs/capabilities",
|
||||||
recoverable_run_capabilities,
|
recoverable_run_capabilities,
|
||||||
@@ -384,5 +694,6 @@ app = Starlette(
|
|||||||
Route("/internal/workspace-scopes/by-thread/{thread_id}", delete_workspace_scope, methods=["DELETE"]),
|
Route("/internal/workspace-scopes/by-thread/{thread_id}", delete_workspace_scope, methods=["DELETE"]),
|
||||||
Route("/internal/workspace-scopes/{scope_id}/runs/reserve", reserve_workspace_run, methods=["POST"]),
|
Route("/internal/workspace-scopes/{scope_id}/runs/reserve", reserve_workspace_run, methods=["POST"]),
|
||||||
Route("/internal/workspace-scopes/{scope_id}/runs/{run_request_id}", bind_workspace_run, methods=["PATCH"]),
|
Route("/internal/workspace-scopes/{scope_id}/runs/{run_request_id}", bind_workspace_run, methods=["PATCH"]),
|
||||||
|
Route("/api/teams", get_teams, methods=["GET"]),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -5,6 +5,7 @@
|
|||||||
"writing-agent": "EvoScientist.langgraph_dev.graphs:writing_agent",
|
"writing-agent": "EvoScientist.langgraph_dev.graphs:writing_agent",
|
||||||
"data-analysis-agent": "EvoScientist.langgraph_dev.graphs:data_analysis_agent",
|
"data-analysis-agent": "EvoScientist.langgraph_dev.graphs:data_analysis_agent",
|
||||||
"scheduler": "EvoScientist.langgraph_dev.graphs:scheduler",
|
"scheduler": "EvoScientist.langgraph_dev.graphs:scheduler",
|
||||||
|
"expert-container-async": "EvoScientist.langgraph_dev.graphs:expert_container_async",
|
||||||
"evomemory-subagent-worker": "EvoScientist.langgraph_dev.graphs:evomemory_subagent_worker",
|
"evomemory-subagent-worker": "EvoScientist.langgraph_dev.graphs:evomemory_subagent_worker",
|
||||||
"evomemory-turn-worker": "EvoScientist.langgraph_dev.graphs:evomemory_turn_worker",
|
"evomemory-turn-worker": "EvoScientist.langgraph_dev.graphs:evomemory_turn_worker",
|
||||||
"evomemory-observation-linker": "EvoScientist.langgraph_dev.graphs:evomemory_observation_linker",
|
"evomemory-observation-linker": "EvoScientist.langgraph_dev.graphs:evomemory_observation_linker",
|
||||||
|
|||||||
@@ -0,0 +1,19 @@
|
|||||||
|
{
|
||||||
|
"dependencies": ["."],
|
||||||
|
"graphs": {
|
||||||
|
"EvoScientist": "EvoScientist.langgraph_dev.main_graph:EvoScientist_agent",
|
||||||
|
"writing-agent": "EvoScientist.langgraph_dev.graphs:writing_agent",
|
||||||
|
"data-analysis-agent": "EvoScientist.langgraph_dev.graphs:data_analysis_agent",
|
||||||
|
"scheduler": "EvoScientist.langgraph_dev.graphs:scheduler",
|
||||||
|
"evomemory-subagent-worker": "EvoScientist.langgraph_dev.graphs:evomemory_subagent_worker",
|
||||||
|
"evomemory-turn-worker": "EvoScientist.langgraph_dev.graphs:evomemory_turn_worker",
|
||||||
|
"evomemory-observation-linker": "EvoScientist.langgraph_dev.graphs:evomemory_observation_linker",
|
||||||
|
"evomemory-autoskills": "EvoScientist.langgraph_dev.graphs:evomemory_autoskills"
|
||||||
|
},
|
||||||
|
"checkpointer": {
|
||||||
|
"backend": "custom",
|
||||||
|
"path": "EvoScientist.web_checkpointer.create_web_checkpointer"
|
||||||
|
},
|
||||||
|
"config": {"recursion_limit": 5000},
|
||||||
|
"http": {"app": "EvoScientist.langgraph_dev.http:app"}
|
||||||
|
}
|
||||||
@@ -12,6 +12,7 @@ Mirrors the lifecycle pattern used by ``ccproxy_manager.py``.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import atexit
|
import atexit
|
||||||
|
import hashlib
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
@@ -21,6 +22,7 @@ import subprocess
|
|||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
from dataclasses import fields as dataclass_fields
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
@@ -110,15 +112,56 @@ def needs_langgraph_dev(config: EvoScientistConfig) -> bool:
|
|||||||
_LOCK = threading.RLock()
|
_LOCK = threading.RLock()
|
||||||
|
|
||||||
|
|
||||||
|
# Set by ``ensure_langgraph_dev`` when it reuses a keepalive server whose
|
||||||
|
# recorded launch-time config fingerprint differs from the current effective
|
||||||
|
# config. The CLI reads it after startup to surface a "restart to apply"
|
||||||
|
# hint — the server itself is never restarted automatically.
|
||||||
|
CONFIG_DRIFT_SINCE_LAUNCH = False
|
||||||
|
|
||||||
|
|
||||||
# Default port shared with the Ai4Sci-Web recoverable runtime.
|
# Default port shared with the Ai4Sci-Web recoverable runtime.
|
||||||
# Overridable per-call via ``start_langgraph_dev(port=...)`` /
|
# Overridable per-call via ``start_langgraph_dev(port=...)`` /
|
||||||
# ``ensure_langgraph_dev`` (which reads ``config.langgraph_dev_port``) and the
|
# ``ensure_langgraph_dev`` (which reads ``config.langgraph_dev_port``) and the
|
||||||
# corresponding url= field on AsyncSubAgent specs.
|
# corresponding url= field on AsyncSubAgent specs.
|
||||||
_DEFAULT_PORT = 3076
|
_DEFAULT_PORT = 3076
|
||||||
|
|
||||||
|
# Default bind interface — loopback, matching ``config.langgraph_dev_host``.
|
||||||
|
# SECURITY: this is the unauthenticated agent API; launchers print a PUBLIC
|
||||||
|
# BIND banner while it is exposed.
|
||||||
|
_DEFAULT_HOST = "127.0.0.1"
|
||||||
|
|
||||||
def _base_url(port: int = _DEFAULT_PORT) -> str:
|
# Wildcard bind addresses: the server listens on every interface, but you
|
||||||
return f"http://localhost:{port}"
|
# cannot meaningfully *connect* to them (0.0.0.0 is routed to loopback on
|
||||||
|
# Linux and outright rejected on Windows), so clients target loopback instead.
|
||||||
|
_WILDCARD_HOSTS = frozenset({"0.0.0.0", "::", ""})
|
||||||
|
|
||||||
|
|
||||||
|
def _probe_host(host: str = _DEFAULT_HOST) -> str:
|
||||||
|
"""Map a bind address to one a client can actually connect to.
|
||||||
|
|
||||||
|
A wildcard bind includes loopback, so clients use ``127.0.0.1``; a
|
||||||
|
specific interface is returned as-is — loopback would not reach it.
|
||||||
|
"""
|
||||||
|
return "127.0.0.1" if host in _WILDCARD_HOSTS else host
|
||||||
|
|
||||||
|
|
||||||
|
def _is_loopback_host(host: str) -> bool:
|
||||||
|
"""Return True if binding ``host`` keeps the server unreachable off-box.
|
||||||
|
|
||||||
|
Drives the PUBLIC BIND warning, so it is conservative: anything not
|
||||||
|
provably loopback counts as exposed.
|
||||||
|
"""
|
||||||
|
return host.strip().lower() in {"127.0.0.1", "::1", "localhost"}
|
||||||
|
|
||||||
|
|
||||||
|
def _format_hostport(host: str, port: int) -> str:
|
||||||
|
"""Render ``host:port`` for a URL, bracketing IPv6 literals per RFC 3986."""
|
||||||
|
probe = _probe_host(host)
|
||||||
|
return f"[{probe}]:{port}" if ":" in probe else f"{probe}:{port}"
|
||||||
|
|
||||||
|
|
||||||
|
def _base_url(port: int = _DEFAULT_PORT, host: str = _DEFAULT_HOST) -> str:
|
||||||
|
return f"http://{_format_hostport(host, port)}"
|
||||||
|
|
||||||
|
|
||||||
# Default rollover threshold for ``RUNTIME.log_file`` — once the log
|
# Default rollover threshold for ``RUNTIME.log_file`` — once the log
|
||||||
@@ -177,9 +220,17 @@ class WorkspaceMismatchError(RuntimeError):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
def _write_workspace_sidecar(workspace_dir: Path, pid: int) -> None:
|
def _write_workspace_sidecar(
|
||||||
|
workspace_dir: Path,
|
||||||
|
pid: int,
|
||||||
|
config_fingerprint: str | None = None,
|
||||||
|
deploy_mode: bool | None = None,
|
||||||
|
) -> None:
|
||||||
"""Record the workspace + pid of the langgraph dev we just started.
|
"""Record the workspace + pid of the langgraph dev we just started.
|
||||||
|
|
||||||
|
``config_fingerprint`` (optional) captures the launch-time config subset
|
||||||
|
the server consumed; keepalive reuse compares it to detect drift.
|
||||||
|
|
||||||
Atomic write via temp-file + ``os.replace``: without this, a concurrent
|
Atomic write via temp-file + ``os.replace``: without this, a concurrent
|
||||||
reader could observe a partially-written file, fail JSON parse, and
|
reader could observe a partially-written file, fail JSON parse, and
|
||||||
silently downgrade to the "no sidecar" fallback path — which skips the
|
silently downgrade to the "no sidecar" fallback path — which skips the
|
||||||
@@ -194,9 +245,12 @@ def _write_workspace_sidecar(workspace_dir: Path, pid: int) -> None:
|
|||||||
try:
|
try:
|
||||||
RUNTIME.pid_dir.mkdir(parents=True, exist_ok=True)
|
RUNTIME.pid_dir.mkdir(parents=True, exist_ok=True)
|
||||||
tmp = RUNTIME.workspace_sidecar.with_suffix(".json.tmp")
|
tmp = RUNTIME.workspace_sidecar.with_suffix(".json.tmp")
|
||||||
tmp.write_text(
|
payload: dict = {"workspace": str(workspace_dir), "pid": pid}
|
||||||
json.dumps({"workspace": str(workspace_dir), "pid": pid}), encoding="utf-8"
|
if config_fingerprint is not None:
|
||||||
)
|
payload["config_fingerprint"] = config_fingerprint
|
||||||
|
if deploy_mode is not None:
|
||||||
|
payload["deploy_mode"] = deploy_mode
|
||||||
|
tmp.write_text(json.dumps(payload), encoding="utf-8")
|
||||||
os.replace(tmp, RUNTIME.workspace_sidecar)
|
os.replace(tmp, RUNTIME.workspace_sidecar)
|
||||||
except OSError as exc:
|
except OSError as exc:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -331,32 +385,37 @@ def is_langgraph_dev_running(
|
|||||||
base_url: str | None = None,
|
base_url: str | None = None,
|
||||||
*,
|
*,
|
||||||
port: int = _DEFAULT_PORT,
|
port: int = _DEFAULT_PORT,
|
||||||
|
host: str = _DEFAULT_HOST,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Check whether a langgraph dev API is already serving at ``base_url``.
|
"""Check whether a langgraph dev API is already serving at ``base_url``.
|
||||||
|
|
||||||
``base_url`` overrides ``port`` when given.
|
``base_url`` overrides ``port``/``host`` when given.
|
||||||
"""
|
"""
|
||||||
url = base_url or _base_url(port)
|
url = base_url or _base_url(port, host)
|
||||||
try:
|
try:
|
||||||
return httpx.get(f"{url}/ok", timeout=1.0).status_code == 200
|
return httpx.get(f"{url}/ok", timeout=1.0).status_code == 200
|
||||||
except (httpx.TransportError, OSError):
|
except (httpx.TransportError, OSError):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def _is_port_occupied(port: int) -> bool:
|
def _is_port_occupied(port: int, host: str = _DEFAULT_HOST) -> bool:
|
||||||
"""Return True if anything is listening on ``port`` (TCP, IPv4)."""
|
"""Return True if anything is listening on ``host:port`` (TCP)."""
|
||||||
import socket as _socket
|
import socket as _socket
|
||||||
|
|
||||||
s = _socket.socket(_socket.AF_INET, _socket.SOCK_STREAM)
|
probe = _probe_host(host)
|
||||||
|
family = _socket.AF_INET6 if ":" in probe else _socket.AF_INET
|
||||||
|
s = _socket.socket(family, _socket.SOCK_STREAM)
|
||||||
try:
|
try:
|
||||||
s.settimeout(0.5)
|
s.settimeout(0.5)
|
||||||
# connect_ex returns 0 on success (something accepted), nonzero otherwise
|
# connect_ex returns 0 on success (something accepted), nonzero otherwise
|
||||||
return s.connect_ex(("127.0.0.1", port)) == 0
|
return s.connect_ex((probe, port)) == 0
|
||||||
finally:
|
finally:
|
||||||
s.close()
|
s.close()
|
||||||
|
|
||||||
|
|
||||||
def _wait_for_port_release(port: int, timeout: float = 10.0) -> bool:
|
def _wait_for_port_release(
|
||||||
|
port: int, timeout: float = 10.0, host: str = _DEFAULT_HOST
|
||||||
|
) -> bool:
|
||||||
"""Poll until ``port`` is released or ``timeout`` elapses.
|
"""Poll until ``port`` is released or ``timeout`` elapses.
|
||||||
|
|
||||||
Used after ``stop_langgraph_dev`` / ``_kill_owned_stale_process`` to
|
Used after ``stop_langgraph_dev`` / ``_kill_owned_stale_process`` to
|
||||||
@@ -364,13 +423,13 @@ def _wait_for_port_release(port: int, timeout: float = 10.0) -> bool:
|
|||||||
True if the port is free, False on timeout.
|
True if the port is free, False on timeout.
|
||||||
"""
|
"""
|
||||||
deadline = time.monotonic() + timeout
|
deadline = time.monotonic() + timeout
|
||||||
while _is_port_occupied(port) and time.monotonic() < deadline:
|
while _is_port_occupied(port, host) and time.monotonic() < deadline:
|
||||||
time.sleep(0.5)
|
time.sleep(0.5)
|
||||||
return not _is_port_occupied(port)
|
return not _is_port_occupied(port, host)
|
||||||
|
|
||||||
|
|
||||||
def _can_bind_port(port: int) -> bool:
|
def _can_bind_port(port: int, host: str = _DEFAULT_HOST) -> bool:
|
||||||
"""Return True if a fresh ``bind()`` to ``port`` succeeds right now.
|
"""Return True if a fresh ``bind()`` to ``host:port`` succeeds right now.
|
||||||
|
|
||||||
More reliable than ``_is_port_occupied`` when the previous listener has
|
More reliable than ``_is_port_occupied`` when the previous listener has
|
||||||
just exited: ``connect_ex`` can already report "free" while ``bind()``
|
just exited: ``connect_ex`` can already report "free" while ``bind()``
|
||||||
@@ -378,12 +437,17 @@ def _can_bind_port(port: int) -> bool:
|
|||||||
(TIME_WAIT for accepted connections, SO_REUSEADDR rules, etc.). This
|
(TIME_WAIT for accepted connections, SO_REUSEADDR rules, etc.). This
|
||||||
actually attempts the bind that langgraph dev would attempt, then
|
actually attempts the bind that langgraph dev would attempt, then
|
||||||
closes immediately.
|
closes immediately.
|
||||||
|
|
||||||
|
Binds the *literal* ``host`` — not ``_probe_host(host)`` — because this
|
||||||
|
must replicate the server's own bind: a loopback probe can succeed while
|
||||||
|
the real wildcard bind still fails on another interface's conflict.
|
||||||
"""
|
"""
|
||||||
import socket as _socket
|
import socket as _socket
|
||||||
|
|
||||||
s = _socket.socket(_socket.AF_INET, _socket.SOCK_STREAM)
|
family = _socket.AF_INET6 if ":" in host else _socket.AF_INET
|
||||||
|
s = _socket.socket(family, _socket.SOCK_STREAM)
|
||||||
try:
|
try:
|
||||||
s.bind(("127.0.0.1", port))
|
s.bind((host, port))
|
||||||
return True
|
return True
|
||||||
except OSError:
|
except OSError:
|
||||||
return False
|
return False
|
||||||
@@ -394,7 +458,9 @@ def _can_bind_port(port: int) -> bool:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
def _wait_for_port_bindable(port: int, timeout: float = 60.0) -> bool:
|
def _wait_for_port_bindable(
|
||||||
|
port: int, timeout: float = 60.0, host: str = _DEFAULT_HOST
|
||||||
|
) -> bool:
|
||||||
"""Poll until a real ``bind()`` to ``port`` can succeed, or timeout.
|
"""Poll until a real ``bind()`` to ``port`` can succeed, or timeout.
|
||||||
|
|
||||||
Use this immediately before ``subprocess.Popen("langgraph dev")`` —
|
Use this immediately before ``subprocess.Popen("langgraph dev")`` —
|
||||||
@@ -408,7 +474,7 @@ def _wait_for_port_bindable(port: int, timeout: float = 60.0) -> bool:
|
|||||||
"""
|
"""
|
||||||
deadline = time.monotonic() + timeout
|
deadline = time.monotonic() + timeout
|
||||||
while time.monotonic() < deadline:
|
while time.monotonic() < deadline:
|
||||||
if _can_bind_port(port):
|
if _can_bind_port(port, host):
|
||||||
return True
|
return True
|
||||||
time.sleep(0.5)
|
time.sleep(0.5)
|
||||||
return False
|
return False
|
||||||
@@ -523,6 +589,186 @@ def _kill_owned_stale_process(port: int) -> bool:
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
# Config fields that provably never reach the langgraph dev subprocess:
|
||||||
|
# the channel stack + STT run in the CLI process, display/workspace/frontend
|
||||||
|
# knobs shape the CLI itself, and keepalive is a lifecycle flag. Everything
|
||||||
|
# NOT listed here counts toward the drift fingerprint, so a newly added
|
||||||
|
# config field defaults to "affects the server" — the failure mode is a
|
||||||
|
# spurious restart hint, never silent staleness.
|
||||||
|
# Packaged sub-agent specs — consumed at graph build; module constant so
|
||||||
|
# tests can redirect it.
|
||||||
|
_SUBAGENTS_DIR = Path(__file__).resolve().parent.parent / "subagents"
|
||||||
|
|
||||||
|
_FINGERPRINT_EXCLUDED_PREFIXES = (
|
||||||
|
"channel_",
|
||||||
|
"imessage_",
|
||||||
|
"telegram_",
|
||||||
|
"discord_",
|
||||||
|
"slack_",
|
||||||
|
"feishu_",
|
||||||
|
"wechat_",
|
||||||
|
"dingtalk_",
|
||||||
|
"email_",
|
||||||
|
"qq_",
|
||||||
|
"signal_",
|
||||||
|
"stt_",
|
||||||
|
)
|
||||||
|
_FINGERPRINT_EXCLUDED_FIELDS = frozenset(
|
||||||
|
{
|
||||||
|
"require_mention",
|
||||||
|
"text_chunk_limit",
|
||||||
|
"allowed_channels",
|
||||||
|
"dm_policy",
|
||||||
|
"shared_webhook_port",
|
||||||
|
"show_thinking",
|
||||||
|
"ui_backend",
|
||||||
|
"log_level",
|
||||||
|
"default_mode",
|
||||||
|
"default_workdir",
|
||||||
|
"webui_port",
|
||||||
|
"webui_host",
|
||||||
|
"langgraph_dev_keepalive",
|
||||||
|
"shell_allow_list",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _server_config_fingerprint(config: EvoScientistConfig) -> str:
|
||||||
|
"""Hash of everything the langgraph dev subprocess consumes at launch.
|
||||||
|
|
||||||
|
Deployed graphs read config once at import (``subagents/_factory.py``,
|
||||||
|
``EvoScientist.py``), so a keepalive server keeps serving those values
|
||||||
|
until restarted. Iterates the full ``EvoScientistConfig`` field list
|
||||||
|
minus the explicit exclusion set above — a new config field counts
|
||||||
|
toward drift by default — and folds in ``mcp.yaml`` plus the packaged
|
||||||
|
``subagents/*.yaml``, which are consumed at graph build too. Secrets
|
||||||
|
only feed a truncated one-way digest; nothing recoverable is stored.
|
||||||
|
getattr with defaults: deploy/WebUI (and their tests) routinely hand
|
||||||
|
this module duck-typed config objects missing dataclass fields.
|
||||||
|
"""
|
||||||
|
parts = []
|
||||||
|
for field in dataclass_fields(EvoScientistConfig):
|
||||||
|
name = field.name
|
||||||
|
if name in _FINGERPRINT_EXCLUDED_FIELDS or name.startswith(
|
||||||
|
_FINGERPRINT_EXCLUDED_PREFIXES
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
parts.append((name, str(getattr(config, name, None))))
|
||||||
|
digest = hashlib.sha256(repr(parts).encode("utf-8"))
|
||||||
|
try:
|
||||||
|
from EvoScientist.config.settings import get_config_dir
|
||||||
|
|
||||||
|
mcp_yaml = get_config_dir() / "mcp.yaml"
|
||||||
|
if mcp_yaml.exists():
|
||||||
|
digest.update(mcp_yaml.read_bytes())
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
for yaml_path in sorted(_SUBAGENTS_DIR.glob("*.yaml")):
|
||||||
|
digest.update(yaml_path.name.encode("utf-8"))
|
||||||
|
digest.update(yaml_path.read_bytes())
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
return digest.hexdigest()[:16]
|
||||||
|
|
||||||
|
|
||||||
|
def stop_recorded_server() -> int | None:
|
||||||
|
"""Explicitly stop the langgraph dev recorded in our PID file.
|
||||||
|
|
||||||
|
Backs the user-facing ``EvoSci server stop`` command — the deliberate
|
||||||
|
counterpart to ``langgraph_dev_keepalive``: an opt-in server that
|
||||||
|
outlives its CLI needs a first-class way to stop it. Ownership = our
|
||||||
|
PID file + a live process whose cmdline still contains ``langgraph``
|
||||||
|
(same loose anti-PID-recycling match as ``_kill_owned_stale_process``,
|
||||||
|
with PID-file ownership as the primary guard). Holds the cross-process
|
||||||
|
file lock so a concurrent start can't have its fresh PID/sidecar records
|
||||||
|
wiped by this stop's cleanup. Kills the whole process tree, then removes
|
||||||
|
the PID file + sidecar. Returns the stopped pid, or ``None`` when nothing
|
||||||
|
was stopped (stale/corrupt files, if any, are still cleaned up).
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
with FileLock(str(RUNTIME.lock_file), timeout=_FILE_LOCK_TIMEOUT):
|
||||||
|
return _stop_recorded_server_locked()
|
||||||
|
except FileLockTimeout:
|
||||||
|
logger.warning(
|
||||||
|
"Timed out waiting for the langgraph dev lock — another EvoSci "
|
||||||
|
"process is mid lifecycle change; not stopping anything."
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _stop_recorded_server_locked() -> int | None:
|
||||||
|
with _LOCK:
|
||||||
|
if _PROCESS is not None and _PROCESS.poll() is None:
|
||||||
|
pid = _PROCESS.pid
|
||||||
|
stop_langgraph_dev()
|
||||||
|
return pid
|
||||||
|
if not RUNTIME.pid_file.exists():
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
owned_pid = int(RUNTIME.pid_file.read_text(encoding="utf-8").strip())
|
||||||
|
except ValueError:
|
||||||
|
stop_langgraph_dev() # corrupt PID file — clean it up as promised
|
||||||
|
return None
|
||||||
|
except OSError:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
proc = psutil.Process(owned_pid)
|
||||||
|
cmdline = proc.cmdline()
|
||||||
|
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||||
|
stop_langgraph_dev() # dead/inaccessible — clean the stale files
|
||||||
|
return None
|
||||||
|
if not any("langgraph" in arg for arg in cmdline):
|
||||||
|
stop_langgraph_dev() # pid recycled by a foreign process — files only
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
children = proc.children(recursive=True)
|
||||||
|
for child in children:
|
||||||
|
try:
|
||||||
|
child.terminate()
|
||||||
|
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||||
|
pass
|
||||||
|
proc.terminate()
|
||||||
|
try:
|
||||||
|
proc.wait(timeout=5)
|
||||||
|
except psutil.TimeoutExpired:
|
||||||
|
for child in children:
|
||||||
|
try:
|
||||||
|
child.kill()
|
||||||
|
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||||
|
pass
|
||||||
|
proc.kill()
|
||||||
|
# The parent exiting promptly doesn't prove its workers did — sweep
|
||||||
|
# the pre-kill snapshot for survivors.
|
||||||
|
for child in children:
|
||||||
|
try:
|
||||||
|
if child.is_running():
|
||||||
|
child.kill()
|
||||||
|
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||||
|
pass
|
||||||
|
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||||
|
pass
|
||||||
|
stop_langgraph_dev()
|
||||||
|
return owned_pid
|
||||||
|
|
||||||
|
|
||||||
|
def _pid_serves_port(pid: object, port: int) -> bool:
|
||||||
|
"""Best-effort check that ``pid`` is a langgraph dev serving ``port``.
|
||||||
|
|
||||||
|
Used to attribute an occupied port to the sidecar's recorded server
|
||||||
|
before printing its details — avoids blaming a stale record. Relies on
|
||||||
|
``--port`` always being in ``start_langgraph_dev``'s argv, not on
|
||||||
|
port→PID mapping (root-only on macOS via psutil).
|
||||||
|
"""
|
||||||
|
if not isinstance(pid, int) or isinstance(pid, bool) or pid <= 0:
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
cmdline = psutil.Process(pid).cmdline()
|
||||||
|
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||||
|
return False
|
||||||
|
return any("langgraph" in arg for arg in cmdline) and str(port) in cmdline
|
||||||
|
|
||||||
|
|
||||||
def _packaged_langgraph_config() -> Path:
|
def _packaged_langgraph_config() -> Path:
|
||||||
"""Return path to the package-shipped ``langgraph.json``.
|
"""Return path to the package-shipped ``langgraph.json``.
|
||||||
|
|
||||||
@@ -544,10 +790,12 @@ def start_langgraph_dev(
|
|||||||
workspace_dir: Path | None = None,
|
workspace_dir: Path | None = None,
|
||||||
*,
|
*,
|
||||||
port: int = _DEFAULT_PORT,
|
port: int = _DEFAULT_PORT,
|
||||||
|
host: str = _DEFAULT_HOST,
|
||||||
file_persistence: bool = True,
|
file_persistence: bool = True,
|
||||||
jobs_per_worker: int = 10,
|
jobs_per_worker: int = 10,
|
||||||
deploy_mode: bool = False,
|
deploy_mode: bool = False,
|
||||||
tunnel: bool = False,
|
tunnel: bool = False,
|
||||||
|
config_fingerprint: str | None = None,
|
||||||
) -> subprocess.Popen:
|
) -> subprocess.Popen:
|
||||||
"""Start langgraph dev as a background subprocess.
|
"""Start langgraph dev as a background subprocess.
|
||||||
|
|
||||||
@@ -557,6 +805,9 @@ def start_langgraph_dev(
|
|||||||
(``CustomSandboxBackend`` derives its workspace root from cwd via
|
(``CustomSandboxBackend`` derives its workspace root from cwd via
|
||||||
``paths.WORKSPACE_ROOT``). Defaults to ``Path.cwd()``.
|
``paths.WORKSPACE_ROOT``). Defaults to ``Path.cwd()``.
|
||||||
port: TCP port to bind. Defaults to 3076.
|
port: TCP port to bind. Defaults to 3076.
|
||||||
|
host: Network interface to bind. Defaults to loopback. SECURITY:
|
||||||
|
widening this exposes an unauthenticated API whose agent can run
|
||||||
|
shell commands — only pass ``0.0.0.0`` on trusted networks.
|
||||||
file_persistence: When True (default), langgraph dev writes its full
|
file_persistence: When True (default), langgraph dev writes its full
|
||||||
``.langgraph_api/`` cache so async-task / Store / scheduler state
|
``.langgraph_api/`` cache so async-task / Store / scheduler state
|
||||||
survives subprocess restarts. Set False to suppress periodic
|
survives subprocess restarts. Set False to suppress periodic
|
||||||
@@ -609,7 +860,9 @@ def start_langgraph_dev(
|
|||||||
# only verifies PID-file ownership, so absence of a match conflates "stale
|
# only verifies PID-file ownership, so absence of a match conflates "stale
|
||||||
# TIME_WAIT" with "foreign process". Falling through to the bind poll
|
# TIME_WAIT" with "foreign process". Falling through to the bind poll
|
||||||
# disambiguates by behavior — TIME_WAIT clears, foreign listeners don't.
|
# disambiguates by behavior — TIME_WAIT clears, foreign listeners don't.
|
||||||
if not is_langgraph_dev_running(port=port) and _is_port_occupied(port):
|
if not is_langgraph_dev_running(port=port, host=host) and _is_port_occupied(
|
||||||
|
port, host
|
||||||
|
):
|
||||||
if _kill_owned_stale_process(port):
|
if _kill_owned_stale_process(port):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Cleaned up stale langgraph dev (pid from %s) on port %d",
|
"Cleaned up stale langgraph dev (pid from %s) on port %d",
|
||||||
@@ -620,7 +873,7 @@ def start_langgraph_dev(
|
|||||||
# several seconds before fully releasing it. Poll until the port
|
# several seconds before fully releasing it. Poll until the port
|
||||||
# is genuinely free so the upcoming bind() doesn't race a
|
# is genuinely free so the upcoming bind() doesn't race a
|
||||||
# half-released socket and crash with "Port already in use".
|
# half-released socket and crash with "Port already in use".
|
||||||
_wait_for_port_release(port)
|
_wait_for_port_release(port, host=host)
|
||||||
else:
|
else:
|
||||||
# No owned stale PID — could be foreign or kernel-only TIME_WAIT
|
# No owned stale PID — could be foreign or kernel-only TIME_WAIT
|
||||||
# from a previous subprocess. Defer to the bind poll below.
|
# from a previous subprocess. Defer to the bind poll below.
|
||||||
@@ -638,9 +891,9 @@ def start_langgraph_dev(
|
|||||||
# "Port already in use" even though our pre-checks passed. By probing
|
# "Port already in use" even though our pre-checks passed. By probing
|
||||||
# the same operation langgraph dev will do, we either wait it out or
|
# the same operation langgraph dev will do, we either wait it out or
|
||||||
# fail clearly with an actionable message. 60s covers macOS TIME_WAIT.
|
# fail clearly with an actionable message. 60s covers macOS TIME_WAIT.
|
||||||
if not _wait_for_port_bindable(port):
|
if not _wait_for_port_bindable(port, host=host):
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"Port {port} cannot be bound after waiting 60s (kernel TIME_WAIT "
|
f"{host}:{port} cannot be bound after waiting 60s (kernel TIME_WAIT "
|
||||||
f"or another process holds it). Free the port with `lsof -ti:{port}`, "
|
f"or another process holds it). Free the port with `lsof -ti:{port}`, "
|
||||||
f"or change ports with: `EvoSci config set langgraph_dev_port <other-port>`"
|
f"or change ports with: `EvoSci config set langgraph_dev_port <other-port>`"
|
||||||
)
|
)
|
||||||
@@ -713,6 +966,25 @@ def start_langgraph_dev(
|
|||||||
sub_env.pop("EVOSCIENTIST_DEPLOY_MODE", None)
|
sub_env.pop("EVOSCIENTIST_DEPLOY_MODE", None)
|
||||||
sub_env["EVOSCIENTIST_DEPLOY_MODE"] = "full" if deploy_mode else "stripped"
|
sub_env["EVOSCIENTIST_DEPLOY_MODE"] = "full" if deploy_mode else "stripped"
|
||||||
|
|
||||||
|
# Propagate the effective bind port into the subprocess's config resolution
|
||||||
|
# via the standard ``EVOSCIENTIST_LANGGRAPH_DEV_PORT`` override (see
|
||||||
|
# ``EvoScientist/config/settings.py``). Without this, ``EvoSci deploy
|
||||||
|
# --port X`` binds to X but the deployed main agent still reads
|
||||||
|
# ``cfg.langgraph_dev_port`` from disk and dispatches self-loop async
|
||||||
|
# tasks (start_async_task → http://localhost:{cfg.port}) to whatever
|
||||||
|
# the config file says — which desyncs from the bind port whenever
|
||||||
|
# ``--port`` differs from the persisted ``langgraph_dev_port``, and
|
||||||
|
# every async subagent launch fails with "All connection attempts failed".
|
||||||
|
# ``get_effective_config`` treats ``EVOSCIENTIST_*`` shell values as
|
||||||
|
# authoritative over any workspace ``.env`` (see its docstring), so a
|
||||||
|
# ``.env`` in the subprocess cwd cannot shadow the caller-resolved port.
|
||||||
|
sub_env["EVOSCIENTIST_LANGGRAPH_DEV_PORT"] = str(port)
|
||||||
|
# Same reasoning for the bind interface: the deployed agent resolves its
|
||||||
|
# self-dispatch URL from ``cfg.langgraph_dev_host``, so a host resolved by
|
||||||
|
# this caller (``EvoSci deploy --host X``) must reach the subprocess too,
|
||||||
|
# or async sub-agent launches would target whatever the config file says.
|
||||||
|
sub_env["EVOSCIENTIST_LANGGRAPH_DEV_HOST"] = host
|
||||||
|
|
||||||
try:
|
try:
|
||||||
logger.info("Starting langgraph dev with CLI: %s", exe)
|
logger.info("Starting langgraph dev with CLI: %s", exe)
|
||||||
proc = subprocess.Popen(
|
proc = subprocess.Popen(
|
||||||
@@ -721,6 +993,8 @@ def start_langgraph_dev(
|
|||||||
"dev",
|
"dev",
|
||||||
"--config",
|
"--config",
|
||||||
str(config_file),
|
str(config_file),
|
||||||
|
"--host",
|
||||||
|
host,
|
||||||
"--port",
|
"--port",
|
||||||
str(port),
|
str(port),
|
||||||
"--n-jobs-per-worker",
|
"--n-jobs-per-worker",
|
||||||
@@ -743,7 +1017,12 @@ def start_langgraph_dev(
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
RUNTIME.pid_file.write_text(str(proc.pid), encoding="utf-8")
|
RUNTIME.pid_file.write_text(str(proc.pid), encoding="utf-8")
|
||||||
_write_workspace_sidecar(workspace_dir=workspace_dir, pid=proc.pid)
|
_write_workspace_sidecar(
|
||||||
|
workspace_dir=workspace_dir,
|
||||||
|
pid=proc.pid,
|
||||||
|
config_fingerprint=config_fingerprint,
|
||||||
|
deploy_mode=deploy_mode,
|
||||||
|
)
|
||||||
global _PROCESS_WORKSPACE
|
global _PROCESS_WORKSPACE
|
||||||
_PROCESS = proc
|
_PROCESS = proc
|
||||||
_PROCESS_WORKSPACE = workspace_dir
|
_PROCESS_WORKSPACE = workspace_dir
|
||||||
@@ -773,9 +1052,9 @@ def start_langgraph_dev(
|
|||||||
f"langgraph dev exited immediately with code {proc.returncode}.\n"
|
f"langgraph dev exited immediately with code {proc.returncode}.\n"
|
||||||
f"Log tail:\n{tail}"
|
f"Log tail:\n{tail}"
|
||||||
)
|
)
|
||||||
if is_langgraph_dev_running(port=port):
|
if is_langgraph_dev_running(port=port, host=host):
|
||||||
logger.info(
|
logger.info(
|
||||||
"langgraph dev started on %s (pid=%d)", _base_url(port), proc.pid
|
"langgraph dev started on %s (pid=%d)", _base_url(port, host), proc.pid
|
||||||
)
|
)
|
||||||
return proc
|
return proc
|
||||||
time.sleep(0.5)
|
time.sleep(0.5)
|
||||||
@@ -920,7 +1199,8 @@ def ensure_langgraph_dev(
|
|||||||
still chat with sync sub-agents; only async sub-agent calls and EvoMemory
|
still chat with sync sub-agents; only async sub-agent calls and EvoMemory
|
||||||
background workers will fail.
|
background workers will fail.
|
||||||
"""
|
"""
|
||||||
global _ASYNC_SUBAGENTS_AVAILABLE
|
global _ASYNC_SUBAGENTS_AVAILABLE, CONFIG_DRIFT_SINCE_LAUNCH
|
||||||
|
CONFIG_DRIFT_SINCE_LAUNCH = False
|
||||||
|
|
||||||
if not needs_langgraph_dev(config):
|
if not needs_langgraph_dev(config):
|
||||||
_ASYNC_SUBAGENTS_AVAILABLE = False
|
_ASYNC_SUBAGENTS_AVAILABLE = False
|
||||||
@@ -959,8 +1239,10 @@ def _ensure_langgraph_dev_locked(
|
|||||||
workspace_dir: Path | str | None,
|
workspace_dir: Path | str | None,
|
||||||
) -> subprocess.Popen | None:
|
) -> subprocess.Popen | None:
|
||||||
"""Locked critical section of ``ensure_langgraph_dev`` — must hold ``_LOCK``."""
|
"""Locked critical section of ``ensure_langgraph_dev`` — must hold ``_LOCK``."""
|
||||||
global _ASYNC_SUBAGENTS_AVAILABLE
|
global _ASYNC_SUBAGENTS_AVAILABLE, CONFIG_DRIFT_SINCE_LAUNCH
|
||||||
|
config_fp = _server_config_fingerprint(config)
|
||||||
port = int(getattr(config, "langgraph_dev_port", _DEFAULT_PORT))
|
port = int(getattr(config, "langgraph_dev_port", _DEFAULT_PORT))
|
||||||
|
host = str(getattr(config, "langgraph_dev_host", _DEFAULT_HOST) or _DEFAULT_HOST)
|
||||||
file_persistence = bool(getattr(config, "langgraph_dev_file_persistence", True))
|
file_persistence = bool(getattr(config, "langgraph_dev_file_persistence", True))
|
||||||
jobs_per_worker = int(getattr(config, "langgraph_dev_jobs_per_worker", 10))
|
jobs_per_worker = int(getattr(config, "langgraph_dev_jobs_per_worker", 10))
|
||||||
|
|
||||||
@@ -993,10 +1275,10 @@ def _ensure_langgraph_dev_locked(
|
|||||||
# and abort with a hard "non-langgraph process" error — turning a
|
# and abort with a hard "non-langgraph process" error — turning a
|
||||||
# clean owned restart into a permanent async-disable. Wait inline for
|
# clean owned restart into a permanent async-disable. Wait inline for
|
||||||
# the kernel to release the port before continuing.
|
# the kernel to release the port before continuing.
|
||||||
_wait_for_port_release(port)
|
_wait_for_port_release(port, host=host)
|
||||||
_ASYNC_SUBAGENTS_AVAILABLE = False # cleared until restart succeeds
|
_ASYNC_SUBAGENTS_AVAILABLE = False # cleared until restart succeeds
|
||||||
|
|
||||||
if is_langgraph_dev_running(port=port):
|
if is_langgraph_dev_running(port=port, host=host):
|
||||||
# If WE own the running process AND it's still alive, workspace was
|
# If WE own the running process AND it's still alive, workspace was
|
||||||
# already verified above via _PROCESS_WORKSPACE comparison. Otherwise
|
# already verified above via _PROCESS_WORKSPACE comparison. Otherwise
|
||||||
# — we never owned it (EvoSci deploy in another terminal, or a
|
# — we never owned it (EvoSci deploy in another terminal, or a
|
||||||
@@ -1012,17 +1294,38 @@ def _ensure_langgraph_dev_locked(
|
|||||||
if sidecar is not None:
|
if sidecar is not None:
|
||||||
recorded = Path(sidecar["workspace"]).resolve()
|
recorded = Path(sidecar["workspace"]).resolve()
|
||||||
if recorded != ws_path.resolve():
|
if recorded != ws_path.resolve():
|
||||||
|
hint = ""
|
||||||
|
if getattr(config, "langgraph_dev_keepalive", False):
|
||||||
|
# Only under keepalive can the server be an ownerless
|
||||||
|
# leftover; without the flag the mismatch means a live
|
||||||
|
# session, where a stop suggestion would be misleading.
|
||||||
|
# Point at `EvoSci server stop` (not a raw kill): it
|
||||||
|
# verifies ownership and cleans the PID/sidecar files,
|
||||||
|
# so no stale records are left behind.
|
||||||
|
hint = (
|
||||||
|
" If it is a leftover keepalive server, stop it"
|
||||||
|
" with: EvoSci server stop."
|
||||||
|
)
|
||||||
raise WorkspaceMismatchError(
|
raise WorkspaceMismatchError(
|
||||||
f"An EvoSci langgraph dev is already running on "
|
f"An EvoSci langgraph dev is already running on "
|
||||||
f"{_base_url(port)} for workspace {recorded}, but the "
|
f"{_base_url(port, host)} for workspace {recorded}, but the "
|
||||||
f"current process requested workspace {ws_path}. "
|
f"current process requested workspace {ws_path}. "
|
||||||
f"Stop the other EvoSci session (deploy / TUI / serve) "
|
f"Stop the other EvoSci session (deploy / TUI / serve) "
|
||||||
f"or rerun with --workdir {recorded}."
|
f"or rerun with --workdir {recorded}." + hint
|
||||||
|
)
|
||||||
|
recorded_fp = sidecar.get("config_fingerprint")
|
||||||
|
if isinstance(recorded_fp, str) and recorded_fp != config_fp:
|
||||||
|
CONFIG_DRIFT_SINCE_LAUNCH = True
|
||||||
|
logger.warning(
|
||||||
|
"Config changed since the running langgraph dev was "
|
||||||
|
"launched — async sub-agents still use the old "
|
||||||
|
"settings until the server is restarted "
|
||||||
|
"(EvoSci server stop)."
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"Reusing externally-managed langgraph dev on %s; sidecar "
|
"Reusing externally-managed langgraph dev on %s; sidecar "
|
||||||
"confirms matching workspace %s.",
|
"confirms matching workspace %s.",
|
||||||
_base_url(port),
|
_base_url(port, host),
|
||||||
recorded,
|
recorded,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -1034,11 +1337,13 @@ def _ensure_langgraph_dev_locked(
|
|||||||
"workspace sidecar, cannot verify it matches the requested "
|
"workspace sidecar, cannot verify it matches the requested "
|
||||||
"%s. Async sub-agents may operate on a different workspace's "
|
"%s. Async sub-agents may operate on a different workspace's "
|
||||||
"files.",
|
"files.",
|
||||||
_base_url(port),
|
_base_url(port, host),
|
||||||
ws_path,
|
ws_path,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.info("langgraph dev already running on %s, reusing", _base_url(port))
|
logger.info(
|
||||||
|
"langgraph dev already running on %s, reusing", _base_url(port, host)
|
||||||
|
)
|
||||||
_ASYNC_SUBAGENTS_AVAILABLE = True
|
_ASYNC_SUBAGENTS_AVAILABLE = True
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -1046,8 +1351,10 @@ def _ensure_langgraph_dev_locked(
|
|||||||
proc = start_langgraph_dev(
|
proc = start_langgraph_dev(
|
||||||
workspace_dir=ws_path,
|
workspace_dir=ws_path,
|
||||||
port=port,
|
port=port,
|
||||||
|
host=host,
|
||||||
file_persistence=file_persistence,
|
file_persistence=file_persistence,
|
||||||
jobs_per_worker=jobs_per_worker,
|
jobs_per_worker=jobs_per_worker,
|
||||||
|
config_fingerprint=config_fp,
|
||||||
)
|
)
|
||||||
except (FileNotFoundError, RuntimeError) as exc:
|
except (FileNotFoundError, RuntimeError) as exc:
|
||||||
# Startup failed — keep async subagents disabled so the main agent
|
# Startup failed — keep async subagents disabled so the main agent
|
||||||
@@ -1064,5 +1371,10 @@ def _ensure_langgraph_dev_locked(
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
_ASYNC_SUBAGENTS_AVAILABLE = True
|
_ASYNC_SUBAGENTS_AVAILABLE = True
|
||||||
atexit.register(stop_langgraph_dev, proc)
|
if getattr(config, "langgraph_dev_keepalive", False):
|
||||||
|
# Keepalive: leave the server (plus PID file + sidecar) behind on CLI
|
||||||
|
# exit so the next start in this workspace reuses it instantly.
|
||||||
|
logger.info("langgraph_dev_keepalive enabled — server will outlive this CLI.")
|
||||||
|
else:
|
||||||
|
atexit.register(stop_langgraph_dev, proc)
|
||||||
return proc
|
return proc
|
||||||
|
|||||||
@@ -6,20 +6,47 @@ import os
|
|||||||
from collections.abc import Mapping
|
from collections.abc import Mapping
|
||||||
|
|
||||||
DEFAULT_LANGGRAPH_DEV_PORT = 3076
|
DEFAULT_LANGGRAPH_DEV_PORT = 3076
|
||||||
|
# Mirrors ``config.langgraph_dev_host`` / ``manager._DEFAULT_HOST``. The value
|
||||||
|
# only matters as a stand-in for the *bind* host — ``_format_hostport`` runs it
|
||||||
|
# through ``_probe_host``, so both this and "0.0.0.0" yield the same client URL.
|
||||||
|
DEFAULT_LANGGRAPH_DEV_HOST = "127.0.0.1"
|
||||||
LANGGRAPH_DEV_AUTH_HEADERS = {"x-auth-scheme": "langsmith"}
|
LANGGRAPH_DEV_AUTH_HEADERS = {"x-auth-scheme": "langsmith"}
|
||||||
|
|
||||||
|
|
||||||
def langgraph_dev_url(config: object | None = None, *, port: int | None = None) -> str:
|
def langgraph_dev_url(
|
||||||
"""Return the local langgraph-dev base URL for a config or explicit port."""
|
config: object | None = None,
|
||||||
|
*,
|
||||||
|
port: int | None = None,
|
||||||
|
host: str | None = None,
|
||||||
|
) -> str:
|
||||||
|
"""Return the local langgraph-dev base URL for a config or explicit port/host.
|
||||||
|
|
||||||
|
An explicit ``LANGGRAPH_SERVER_URL`` (container / prod deploy) wins when no
|
||||||
|
port or host override is supplied. Otherwise the configured bind interface is
|
||||||
|
mapped through ``manager._probe_host``: a wildcard bind (``0.0.0.0``) still
|
||||||
|
resolves to loopback here, while a specific interface is honored so
|
||||||
|
self-dispatch keeps working when the server is pinned to one address.
|
||||||
|
"""
|
||||||
runtime_url = os.environ.get("LANGGRAPH_SERVER_URL", "").strip().rstrip("/")
|
runtime_url = os.environ.get("LANGGRAPH_SERVER_URL", "").strip().rstrip("/")
|
||||||
if port is None and runtime_url:
|
if port is None and host is None and runtime_url:
|
||||||
return runtime_url
|
return runtime_url
|
||||||
|
|
||||||
|
from .manager import _format_hostport
|
||||||
|
|
||||||
selected_port = (
|
selected_port = (
|
||||||
int(port)
|
int(port)
|
||||||
if port is not None
|
if port is not None
|
||||||
else int(getattr(config, "langgraph_dev_port", DEFAULT_LANGGRAPH_DEV_PORT))
|
else int(getattr(config, "langgraph_dev_port", DEFAULT_LANGGRAPH_DEV_PORT))
|
||||||
)
|
)
|
||||||
return f"http://localhost:{selected_port}"
|
selected_host = (
|
||||||
|
host
|
||||||
|
if host is not None
|
||||||
|
else str(
|
||||||
|
getattr(config, "langgraph_dev_host", DEFAULT_LANGGRAPH_DEV_HOST)
|
||||||
|
or DEFAULT_LANGGRAPH_DEV_HOST
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return f"http://{_format_hostport(selected_host, selected_port)}"
|
||||||
|
|
||||||
|
|
||||||
def configured_langgraph_dev_url() -> str:
|
def configured_langgraph_dev_url() -> str:
|
||||||
|
|||||||
@@ -0,0 +1,236 @@
|
|||||||
|
"""Process-local exit receipts for the existing LangGraph worker.
|
||||||
|
|
||||||
|
No receipt survives a restart. Missing receipts never prove execution absent.
|
||||||
|
The admission gate also covers a queued worker that has not started yet.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from concurrent.futures import Future
|
||||||
|
from contextvars import ContextVar
|
||||||
|
from functools import wraps
|
||||||
|
from threading import RLock
|
||||||
|
from typing import Any, cast
|
||||||
|
|
||||||
|
_lock = RLock()
|
||||||
|
_runs: dict[tuple[str, str], dict] = {}
|
||||||
|
_installed = False
|
||||||
|
_owned: ContextVar[bool] = ContextVar("recoverable_worker_owned", default=False)
|
||||||
|
_evidence: ContextVar[dict | None] = ContextVar("worker_exit_evidence", default=None)
|
||||||
|
|
||||||
|
|
||||||
|
class _ExitBoundary:
|
||||||
|
def __init__(self, context, *, run=False):
|
||||||
|
self.context = context
|
||||||
|
self.run = run
|
||||||
|
|
||||||
|
async def __aenter__(self):
|
||||||
|
return await self.context.__aenter__()
|
||||||
|
|
||||||
|
async def __aexit__(self, typ, value, tb):
|
||||||
|
evidence = _evidence.get()
|
||||||
|
try:
|
||||||
|
result = await self.context.__aexit__(typ, value, tb)
|
||||||
|
except BaseException as exc:
|
||||||
|
# An asynccontextmanager propagating its body exception is not a
|
||||||
|
# cleanup failure. Replacement exceptions are fail-closed.
|
||||||
|
if evidence is not None and exc is not value:
|
||||||
|
evidence["cleanup_failed"] = True
|
||||||
|
if evidence is not None and self.run and exc is value:
|
||||||
|
evidence["run_exited"] = True
|
||||||
|
raise
|
||||||
|
else:
|
||||||
|
if evidence is not None and self.run:
|
||||||
|
evidence["run_exited"] = True
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
async def _drain(task):
|
||||||
|
while not task.done():
|
||||||
|
try:
|
||||||
|
await asyncio.shield(task)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
continue
|
||||||
|
except BaseException:
|
||||||
|
break
|
||||||
|
if not task.cancelled():
|
||||||
|
task.exception()
|
||||||
|
|
||||||
|
|
||||||
|
async def _await_remote_future(remote: Future):
|
||||||
|
"""Do not let caller cancellation detach a coroutine on a worker thread."""
|
||||||
|
async def wait():
|
||||||
|
return await asyncio.wrap_future(remote)
|
||||||
|
|
||||||
|
task = asyncio.create_task(wait())
|
||||||
|
try:
|
||||||
|
return await asyncio.shield(task)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
await _drain(task)
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
async def _persistent_cancellation_listener(original, queue, run_id, thread_id, done):
|
||||||
|
"""The in-memory listener's idle timeout is a poll, not end-of-listening."""
|
||||||
|
while not done.is_set():
|
||||||
|
await original(queue, run_id, thread_id, done)
|
||||||
|
|
||||||
|
|
||||||
|
def install() -> None:
|
||||||
|
global _installed
|
||||||
|
from langgraph_api import worker, stream
|
||||||
|
from langgraph_runtime_inmem import ops as inmem_ops
|
||||||
|
from langgraph.pregel._loop import AsyncPregelLoop
|
||||||
|
from quickjs_rs.threading import ThreadWorker
|
||||||
|
|
||||||
|
with _lock:
|
||||||
|
if _installed:
|
||||||
|
return
|
||||||
|
original = worker.worker
|
||||||
|
original_to_thread = asyncio.to_thread
|
||||||
|
original_enter = worker.Runs.enter
|
||||||
|
original_closing = stream.aclosing
|
||||||
|
original_stack = stream.AsyncExitStack
|
||||||
|
original_loop_exit = AsyncPregelLoop.__aexit__
|
||||||
|
original_cancellation_listener = inmem_ops.listen_for_cancellation
|
||||||
|
|
||||||
|
|
||||||
|
class ObservedStack(original_stack):
|
||||||
|
async def __aexit__(self, typ, value, tb):
|
||||||
|
return await _ExitBoundary(super()).__aexit__(typ, value, tb)
|
||||||
|
|
||||||
|
async def loop_exit(self, exc_type, exc_value, traceback):
|
||||||
|
evidence = _evidence.get()
|
||||||
|
if evidence is None:
|
||||||
|
return await original_loop_exit(self, exc_type, exc_value, traceback)
|
||||||
|
try:
|
||||||
|
return await original_loop_exit(self, exc_type, exc_value, traceback)
|
||||||
|
except asyncio.CancelledError as exc:
|
||||||
|
# Installed Pregel exposes its outstanding exit task in args.
|
||||||
|
tasks = [arg for arg in exc.args if isinstance(arg, asyncio.Task)]
|
||||||
|
for task in tasks:
|
||||||
|
await _drain(task)
|
||||||
|
if evidence is not None and (
|
||||||
|
not tasks or any(task.cancelled() or task.exception() is not None for task in tasks)
|
||||||
|
):
|
||||||
|
evidence["cleanup_failed"] = True
|
||||||
|
raise
|
||||||
|
except BaseException:
|
||||||
|
if evidence is not None:
|
||||||
|
evidence["cleanup_failed"] = True
|
||||||
|
raise
|
||||||
|
|
||||||
|
cast(Any, worker.Runs).enter = staticmethod(
|
||||||
|
lambda *args, **kw: _ExitBoundary(cast(Any, original_enter)(*args, **kw), run=True)
|
||||||
|
)
|
||||||
|
stream.aclosing = lambda iterator: _ExitBoundary(original_closing(iterator))
|
||||||
|
stream.AsyncExitStack = ObservedStack
|
||||||
|
cast(Any, AsyncPregelLoop).__aexit__ = loop_exit
|
||||||
|
|
||||||
|
async def persistent_cancellation_listener(queue, run_id, thread_id, done):
|
||||||
|
return await _persistent_cancellation_listener(
|
||||||
|
original_cancellation_listener, queue, run_id, thread_id, done
|
||||||
|
)
|
||||||
|
|
||||||
|
def quickjs_run_async(self, coro):
|
||||||
|
self._ensure_started()
|
||||||
|
remote = asyncio.run_coroutine_threadsafe(coro, self._loop)
|
||||||
|
return asyncio.create_task(_await_remote_future(remote))
|
||||||
|
|
||||||
|
inmem_ops.listen_for_cancellation = persistent_cancellation_listener
|
||||||
|
cast(Any, ThreadWorker).run_async = quickjs_run_async
|
||||||
|
|
||||||
|
@wraps(original_to_thread)
|
||||||
|
async def owned_to_thread(func, /, *args, **kwargs):
|
||||||
|
if not _owned.get():
|
||||||
|
return await original_to_thread(func, *args, **kwargs)
|
||||||
|
task = asyncio.create_task(original_to_thread(func, *args, **kwargs))
|
||||||
|
try:
|
||||||
|
return await asyncio.shield(task)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
await _drain(task)
|
||||||
|
raise
|
||||||
|
|
||||||
|
@wraps(original)
|
||||||
|
async def observed(run, attempt, main_loop, **kwargs):
|
||||||
|
key = (str(run["thread_id"]), str(run["run_id"]))
|
||||||
|
identity = (attempt, object())
|
||||||
|
with _lock:
|
||||||
|
entry = _runs.get(key)
|
||||||
|
if entry is not None:
|
||||||
|
if entry["cancel_requested"]:
|
||||||
|
# Cancel closed admission before this queued attempt ran.
|
||||||
|
return None
|
||||||
|
entry["active"] += 1
|
||||||
|
entry["exited"] = False
|
||||||
|
entry["attempts"][identity] = False
|
||||||
|
result = None
|
||||||
|
token = _owned.set(entry is not None)
|
||||||
|
evidence = {"run_exited": False, "cleanup_failed": False}
|
||||||
|
evidence_token = _evidence.set(evidence if entry is not None else None)
|
||||||
|
try:
|
||||||
|
result = await original(run, attempt, main_loop, **kwargs)
|
||||||
|
return result
|
||||||
|
finally:
|
||||||
|
_owned.reset(token)
|
||||||
|
_evidence.reset(evidence_token)
|
||||||
|
if entry is not None:
|
||||||
|
with _lock:
|
||||||
|
entry["active"] -= 1
|
||||||
|
# Only this invocation can discharge its own evidence.
|
||||||
|
# Business error/timeout/retry is independent of cleanup.
|
||||||
|
clean = evidence["run_exited"] and not evidence["cleanup_failed"]
|
||||||
|
entry["attempts"][identity] = clean
|
||||||
|
entry["uncertain"] = not all(entry["attempts"].values())
|
||||||
|
if clean:
|
||||||
|
entry["exited"] = True
|
||||||
|
entry["status"] = result["status"] if result else "retry"
|
||||||
|
|
||||||
|
worker.worker = observed
|
||||||
|
asyncio.to_thread = owned_to_thread
|
||||||
|
_installed = True
|
||||||
|
|
||||||
|
|
||||||
|
def reserve(thread_id: str, run_id: str, request_hash: str) -> None:
|
||||||
|
with _lock:
|
||||||
|
_runs.setdefault((thread_id, run_id), {
|
||||||
|
"request_hash": request_hash, "active": 0,
|
||||||
|
"cancel_requested": False, "exited": False,
|
||||||
|
"uncertain": False, "status": "pending", "attempts": {},
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
def is_reserved(thread_id: str, run_id: str) -> bool:
|
||||||
|
with _lock:
|
||||||
|
return (thread_id, run_id) in _runs
|
||||||
|
|
||||||
|
|
||||||
|
def cancel_and_inspect(thread_id: str, run_id: str) -> dict:
|
||||||
|
with _lock:
|
||||||
|
entry = _runs.get((thread_id, run_id))
|
||||||
|
if entry is None:
|
||||||
|
return {"run_id": run_id, "thread_id": thread_id, "execution_exited": False}
|
||||||
|
entry["cancel_requested"] = True
|
||||||
|
# No active worker plus closed admission is safe even for pending work.
|
||||||
|
# A worker that escaped abnormally keeps uncertainty latched.
|
||||||
|
stopped = entry["active"] == 0 and not entry["uncertain"]
|
||||||
|
return {
|
||||||
|
"run_id": run_id, "thread_id": thread_id,
|
||||||
|
"request_hash": entry["request_hash"],
|
||||||
|
"execution_exited": stopped,
|
||||||
|
"exit_kind": "worker_returned" if entry["exited"] else "admission_closed",
|
||||||
|
"status": entry["status"] if entry["exited"] else "interrupted",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def wait_for_exit(thread_id: str, run_id: str, timeout: float = 3.0) -> dict:
|
||||||
|
deadline = asyncio.get_running_loop().time() + timeout
|
||||||
|
while True:
|
||||||
|
receipt = cancel_and_inspect(thread_id, run_id)
|
||||||
|
if receipt.get("execution_exited") is True:
|
||||||
|
return receipt
|
||||||
|
remaining = deadline - asyncio.get_running_loop().time()
|
||||||
|
if remaining <= 0:
|
||||||
|
return receipt
|
||||||
|
await asyncio.sleep(min(0.02, remaining))
|
||||||
@@ -18,6 +18,7 @@ __getattr__, __dir__, __all__ = _lazy.attach(
|
|||||||
"context_window",
|
"context_window",
|
||||||
"models",
|
"models",
|
||||||
"patches",
|
"patches",
|
||||||
|
"registry",
|
||||||
"contracts",
|
"contracts",
|
||||||
"config_admin",
|
"config_admin",
|
||||||
"configuration",
|
"configuration",
|
||||||
@@ -35,9 +36,12 @@ __getattr__, __dir__, __all__ = _lazy.attach(
|
|||||||
"resolve_context_window",
|
"resolve_context_window",
|
||||||
],
|
],
|
||||||
"models": [
|
"models": [
|
||||||
|
"get_chat_model",
|
||||||
|
],
|
||||||
|
# Registry data resolves without the langchain/provider-SDK stack.
|
||||||
|
"registry": [
|
||||||
"DEFAULT_MODEL",
|
"DEFAULT_MODEL",
|
||||||
"MODELS",
|
"MODELS",
|
||||||
"get_chat_model",
|
|
||||||
"get_model_info",
|
"get_model_info",
|
||||||
"get_models_for_provider",
|
"get_models_for_provider",
|
||||||
"list_models",
|
"list_models",
|
||||||
|
|||||||
@@ -398,6 +398,8 @@ class AdapterRegistration:
|
|||||||
)
|
)
|
||||||
if reasoning not in {None, "off"}:
|
if reasoning not in {None, "off"}:
|
||||||
result["reasoning"] = {"effort": reasoning}
|
result["reasoning"] = {"effort": reasoning}
|
||||||
|
if self.adapter_id in {"openai", "xai"}:
|
||||||
|
result["reasoning"]["summary"] = "auto"
|
||||||
elif self.adapter_id == "dashscope" and thinking:
|
elif self.adapter_id == "dashscope" and thinking:
|
||||||
result["reasoning"] = {"effort": "medium"}
|
result["reasoning"] = {"effort": "medium"}
|
||||||
elif self.adapter_id == "google-gemini":
|
elif self.adapter_id == "google-gemini":
|
||||||
|
|||||||
@@ -17,11 +17,15 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = {
|
|||||||
# Qwen 3.6 open-source variants — exceptions to the ``qwen3.6`` family.
|
# Qwen 3.6 open-source variants — exceptions to the ``qwen3.6`` family.
|
||||||
"qwen3.6-27b": 262_000,
|
"qwen3.6-27b": 262_000,
|
||||||
"qwen3.6-35b-a3b": 262_000,
|
"qwen3.6-35b-a3b": 262_000,
|
||||||
|
# Qwen 3.8 closed-source tiers — Max flagship and Flash (1M).
|
||||||
|
"qwen3.8-max": 1_000_000,
|
||||||
|
"qwen3.8-flash": 1_000_000,
|
||||||
# Qwen 3.7 closed-source tiers — Max flagship and Plus (1M).
|
# Qwen 3.7 closed-source tiers — Max flagship and Plus (1M).
|
||||||
"qwen3.7-max": 1_000_000,
|
"qwen3.7-max": 1_000_000,
|
||||||
"qwen3.7-plus": 1_000_000,
|
"qwen3.7-plus": 1_000_000,
|
||||||
# xAI Grok — per-model windows (build-0.1: 256K, 4.5: 500K).
|
# xAI Grok — per-model windows (build-0.1: 256K, 4.5/4.6: 500K).
|
||||||
"grok-build-0.1": 256_000,
|
"grok-build-0.1": 256_000,
|
||||||
|
"grok-4.6": 500_000,
|
||||||
"grok-4.5": 500_000,
|
"grok-4.5": 500_000,
|
||||||
# Claude Haiku 4.5 — exception to the ``claude-`` family (200K, not 1M).
|
# Claude Haiku 4.5 — exception to the ``claude-`` family (200K, not 1M).
|
||||||
"claude-haiku-4-5": 200_000,
|
"claude-haiku-4-5": 200_000,
|
||||||
@@ -29,10 +33,15 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = {
|
|||||||
# Covers OpenRouter ``minimax/minimax-m3`` (via split('/')[-1]) and direct
|
# Covers OpenRouter ``minimax/minimax-m3`` (via split('/')[-1]) and direct
|
||||||
# ``MiniMax-M3`` (via lowercased exact match).
|
# ``MiniMax-M3`` (via lowercased exact match).
|
||||||
"minimax-m3": 1_000_000,
|
"minimax-m3": 1_000_000,
|
||||||
# Zhipu GLM-5.2 — 1M context, an exception to the ``glm-5`` family (203K).
|
# Zhipu GLM-5.3/5.2 — 1M context, exceptions to the ``glm-5`` family (203K).
|
||||||
# Matches OpenRouter ``z-ai/glm-5.2`` via split('/')[-1].
|
# Matches OpenRouter ``z-ai/glm-5.x`` via split('/')[-1].
|
||||||
|
"glm-5.3": 1_000_000,
|
||||||
|
"glm-5.3-flash": 1_000_000,
|
||||||
"glm-5.2": 1_000_000,
|
"glm-5.2": 1_000_000,
|
||||||
# Tencent Hunyuan HY3 — 262K context (OpenRouter ``tencent/hy3``).
|
# Volcengine Coding Plan's OpenAI-compatible alias for GLM-5.2.
|
||||||
|
"glm-5-2": 1_000_000,
|
||||||
|
# Tencent Hunyuan — HY4 preview 1M, HY3 262K (OpenRouter ``tencent/hy*``).
|
||||||
|
"hy4-preview": 1_048_576,
|
||||||
"hy3": 262_000,
|
"hy3": 262_000,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -42,12 +51,17 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = {
|
|||||||
_KNOWN_MODEL_FAMILIES: list[tuple[str, int]] = [
|
_KNOWN_MODEL_FAMILIES: list[tuple[str, int]] = [
|
||||||
# All Claude — 1M via the ``context-1m-2025-08-07`` beta header.
|
# All Claude — 1M via the ``context-1m-2025-08-07`` beta header.
|
||||||
("claude-", 1_000_000),
|
("claude-", 1_000_000),
|
||||||
|
# OpenAI GPT-6 family — astra, astra-pro, future variants
|
||||||
|
("gpt-6", 1_050_000),
|
||||||
# OpenAI GPT-5.6 family — sol, terra, luna variants
|
# OpenAI GPT-5.6 family — sol, terra, luna variants
|
||||||
("gpt-5.6", 1_050_000),
|
("gpt-5.6", 1_050_000),
|
||||||
# OpenAI GPT-5.5 family — base, pro, future variants
|
# OpenAI GPT-5.5 family — base, pro, future variants
|
||||||
("gpt-5.5", 1_050_000),
|
("gpt-5.5", 1_050_000),
|
||||||
# Google Gemini 3.x family — flash, flash-lite, pro (1.05M). Excludes 2.5.
|
# Google Gemini 3.x family — flash, flash-lite, pro (1.05M). Excludes 2.5.
|
||||||
("gemini-3", 1_050_000),
|
("gemini-3", 1_050_000),
|
||||||
|
# Moonshot Kimi K3 — 1M context; covers bare ``kimi-k3`` (native Moonshot),
|
||||||
|
# OpenRouter ``moonshotai/kimi-k3``, and dated slugs like ``kimi-k3-20260715``.
|
||||||
|
("kimi-k3", 1_048_576),
|
||||||
# Moonshot Kimi K2 family — k2.5, k2.6, k2-thinking, k2-thinking-turbo
|
# Moonshot Kimi K2 family — k2.5, k2.6, k2-thinking, k2-thinking-turbo
|
||||||
("kimi-k2", 262_000),
|
("kimi-k2", 262_000),
|
||||||
# Zhipu GLM-5 family — base, 5.1, 5-turbo, 5v-turbo, etc.
|
# Zhipu GLM-5 family — base, 5.1, 5-turbo, 5v-turbo, etc.
|
||||||
@@ -56,6 +70,8 @@ _KNOWN_MODEL_FAMILIES: list[tuple[str, int]] = [
|
|||||||
("deepseek-v4", 1_050_000),
|
("deepseek-v4", 1_050_000),
|
||||||
# Xiaomi MiMo v2.5 family — base, pro, future variants
|
# Xiaomi MiMo v2.5 family — base, pro, future variants
|
||||||
("mimo-v2.5", 1_050_000),
|
("mimo-v2.5", 1_050_000),
|
||||||
|
# Meta Muse Spark family — 1.1/1.2/1.3 (OpenRouter ``meta/muse-spark-*``, 1M).
|
||||||
|
("muse-spark", 1_048_576),
|
||||||
# Qwen 3.6 closed-source family — flash, plus, max-preview, etc.
|
# Qwen 3.6 closed-source family — flash, plus, max-preview, etc.
|
||||||
# Open-source ``-<size>b`` variants are 262K — listed in the dict above.
|
# Open-source ``-<size>b`` variants are 262K — listed in the dict above.
|
||||||
("qwen3.6", 1_000_000),
|
("qwen3.6", 1_000_000),
|
||||||
|
|||||||
@@ -55,6 +55,11 @@ class EvoRuntimeError(RuntimeError):
|
|||||||
self.code = code
|
self.code = code
|
||||||
self.details = tuple(dict(item) for item in details)
|
self.details = tuple(dict(item) for item in details)
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
# LangGraph persists task failures using repr(exc). Keep that snapshot
|
||||||
|
# machine-readable without serializing provider messages or details.
|
||||||
|
return f"{type(self).__name__}(code={self.code!r})"
|
||||||
|
|
||||||
|
|
||||||
def now_ms() -> int:
|
def now_ms() -> int:
|
||||||
return time.time_ns() // 1_000_000
|
return time.time_ns() // 1_000_000
|
||||||
@@ -63,6 +68,14 @@ def now_ms() -> int:
|
|||||||
def _unsigned(value: Any) -> dict[str, Any]:
|
def _unsigned(value: Any) -> dict[str, Any]:
|
||||||
payload = asdict(value)
|
payload = asdict(value)
|
||||||
payload.pop("signature", None)
|
payload.pop("signature", None)
|
||||||
|
if payload.get("predecessor_owner_epoch") == 0:
|
||||||
|
payload.pop("predecessor_owner_epoch")
|
||||||
|
for name in ("continuation_pending_hash", "continuation_decision_hash"):
|
||||||
|
if payload.get(name) == "":
|
||||||
|
payload.pop(name)
|
||||||
|
for name in ("execution_id", "predecessor_execution_id", "predecessor_checkpoint_id"):
|
||||||
|
if payload.get(name) == "":
|
||||||
|
payload.pop(name)
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
|
|
||||||
@@ -90,6 +103,11 @@ class RoutePreparationGrant:
|
|||||||
key_id: str
|
key_id: str
|
||||||
signature: str
|
signature: str
|
||||||
contract_version: int = CONTRACT_VERSION
|
contract_version: int = CONTRACT_VERSION
|
||||||
|
predecessor_execution_id: str = ""
|
||||||
|
predecessor_checkpoint_id: str = ""
|
||||||
|
predecessor_owner_epoch: int = 0
|
||||||
|
continuation_pending_hash: str = ""
|
||||||
|
continuation_decision_hash: str = ""
|
||||||
|
|
||||||
def unsigned_payload(self) -> dict[str, Any]:
|
def unsigned_payload(self) -> dict[str, Any]:
|
||||||
return _unsigned(self)
|
return _unsigned(self)
|
||||||
@@ -190,6 +208,7 @@ class PreparedRunQuote:
|
|||||||
key_id: str
|
key_id: str
|
||||||
signature: str
|
signature: str
|
||||||
contract_version: int = CONTRACT_VERSION
|
contract_version: int = CONTRACT_VERSION
|
||||||
|
execution_id: str = ""
|
||||||
|
|
||||||
def unsigned_payload(self) -> dict[str, Any]:
|
def unsigned_payload(self) -> dict[str, Any]:
|
||||||
return _unsigned(self)
|
return _unsigned(self)
|
||||||
@@ -617,11 +636,7 @@ class HmacGrantAuthority:
|
|||||||
payload = {**defaults, **kwargs}
|
payload = {**defaults, **kwargs}
|
||||||
if "roles" in payload:
|
if "roles" in payload:
|
||||||
payload["roles"] = tuple(sorted({str(role) for role in payload["roles"]}))
|
payload["roles"] = tuple(sorted({str(role) for role in payload["roles"]}))
|
||||||
unsigned = {
|
unsigned = _unsigned(contract_class(signature="", **payload))
|
||||||
key_name: value
|
|
||||||
for key_name, value in payload.items()
|
|
||||||
if key_name != "signature"
|
|
||||||
}
|
|
||||||
signature = sign_contract(contract_class.__name__, unsigned, key)
|
signature = sign_contract(contract_class.__name__, unsigned, key)
|
||||||
return contract_class(signature=signature, **payload)
|
return contract_class(signature=signature, **payload)
|
||||||
|
|
||||||
@@ -761,6 +776,22 @@ class AgentExecutionProfile:
|
|||||||
return cls.web_v3()
|
return cls.web_v3()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ModelFactoryResult:
|
||||||
|
"""Explicit transfer of exclusively owned closeable resources to one run.
|
||||||
|
|
||||||
|
Plain custom factory return values are borrowed. Do not list shared clients.
|
||||||
|
The runtime closes only these exact objects, never their reachable children.
|
||||||
|
Built-in factories register newly allocated transports during construction.
|
||||||
|
Cleanup tries every resource even after failures, at most three times per
|
||||||
|
resource. Failed resources and their exception history remain owned; no
|
||||||
|
successful terminal event is emitted until all resources close.
|
||||||
|
"""
|
||||||
|
|
||||||
|
model: Any
|
||||||
|
owned_clients: tuple[Any, ...] = ()
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
class AgentModelSet:
|
class AgentModelSet:
|
||||||
main_agent: Any
|
main_agent: Any
|
||||||
@@ -859,7 +890,23 @@ class EvoWebRun(Protocol):
|
|||||||
self, after_sequence: int | None = None
|
self, after_sequence: int | None = None
|
||||||
) -> AsyncIterator[EvoRuntimeEvent]: ...
|
) -> AsyncIterator[EvoRuntimeEvent]: ...
|
||||||
|
|
||||||
async def cancel(self, reason: str) -> str: ...
|
async def cancel(self, reason: str, *, owner_epoch: int | None = None,
|
||||||
|
boot_id: str | None = None) -> str: ...
|
||||||
|
|
||||||
|
async def wait_stopped(self, timeout: float | None = 2.0) -> str:
|
||||||
|
"""Return a terminal outcome only after owned resource cleanup.
|
||||||
|
|
||||||
|
A bounded timeout returns ``unknown`` without cancelling execution or
|
||||||
|
cleanup. Cancelling the waiter also leaves cleanup owned by the run.
|
||||||
|
Host checkpointers and workspace backends are borrowed, never closed.
|
||||||
|
``awaiting_input`` ends this execution, not the checkpoint workflow;
|
||||||
|
a separate run may resume it. Its run_terminal payload includes the
|
||||||
|
final checkpoint_thread_id, checkpoint_id, checkpoint_ns and JSON
|
||||||
|
pending_interrupts (id/value records). Consumers must not interpret
|
||||||
|
every non-failed outcome as completed. Final checkpoint read failure
|
||||||
|
yields failed with FINAL_CHECKPOINT_READ_FAILED, never completed.
|
||||||
|
"""
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
def event_payload(value: Any) -> dict[str, Any]:
|
def event_payload(value: Any) -> dict[str, Any]:
|
||||||
|
|||||||
@@ -0,0 +1,80 @@
|
|||||||
|
"""DeepSeek chat model integration."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from collections.abc import Mapping
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from langchain_core.language_models import LanguageModelInput
|
||||||
|
from langchain_core.messages import AIMessage, BaseMessage
|
||||||
|
from langchain_deepseek import ChatDeepSeek
|
||||||
|
|
||||||
|
from .openai_compat import OpenAICompatContentMixin
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
DEEPSEEK_THINKING_DISABLED = {"type": "disabled"}
|
||||||
|
|
||||||
|
|
||||||
|
def is_deepseek_thinking_disabled(
|
||||||
|
extra_body: Mapping[str, object] | None,
|
||||||
|
) -> bool:
|
||||||
|
"""Return whether a request body explicitly disables DeepSeek thinking."""
|
||||||
|
if not extra_body:
|
||||||
|
return False
|
||||||
|
thinking = extra_body.get("thinking")
|
||||||
|
return isinstance(thinking, Mapping) and thinking.get("type") == "disabled"
|
||||||
|
|
||||||
|
|
||||||
|
def _inject_reasoning_content(
|
||||||
|
messages: list[BaseMessage],
|
||||||
|
payload: dict[str, object],
|
||||||
|
) -> dict[str, object]:
|
||||||
|
"""Copy captured DeepSeek reasoning into serialized assistant messages."""
|
||||||
|
reasoning = [
|
||||||
|
message.additional_kwargs.get("reasoning_content")
|
||||||
|
for message in messages
|
||||||
|
if isinstance(message, AIMessage)
|
||||||
|
]
|
||||||
|
serialized = payload.get("messages")
|
||||||
|
if not isinstance(serialized, list):
|
||||||
|
return payload
|
||||||
|
|
||||||
|
ai_index = 0
|
||||||
|
for message in serialized:
|
||||||
|
if not isinstance(message, dict) or message.get("role") != "assistant":
|
||||||
|
continue
|
||||||
|
value = reasoning[ai_index] if ai_index < len(reasoning) else None
|
||||||
|
if value:
|
||||||
|
message["reasoning_content"] = value
|
||||||
|
elif "reasoning_content" not in message:
|
||||||
|
message["reasoning_content"] = ""
|
||||||
|
ai_index += 1
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
class EvoChatDeepSeek(OpenAICompatContentMixin, ChatDeepSeek):
|
||||||
|
"""ChatDeepSeek with EvoScientist's media and history compatibility."""
|
||||||
|
|
||||||
|
def _get_request_payload(
|
||||||
|
self,
|
||||||
|
input_: LanguageModelInput,
|
||||||
|
*,
|
||||||
|
stop: list[str] | None = None,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
payload = super()._get_request_payload(input_, stop=stop, **kwargs)
|
||||||
|
if is_deepseek_thinking_disabled(self.extra_body):
|
||||||
|
return payload
|
||||||
|
|
||||||
|
try:
|
||||||
|
messages = self._convert_input(input_).to_messages()
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"DeepSeek reasoning passback: input conversion failed",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
return payload
|
||||||
|
|
||||||
|
return _inject_reasoning_content(messages, payload)
|
||||||
@@ -353,6 +353,7 @@ def _redact_api_keys(message: str) -> str:
|
|||||||
_HOST_TO_PROVIDER: dict[str, str] = {
|
_HOST_TO_PROVIDER: dict[str, str] = {
|
||||||
"api.openai.com": "openai",
|
"api.openai.com": "openai",
|
||||||
"api.anthropic.com": "anthropic",
|
"api.anthropic.com": "anthropic",
|
||||||
|
"api.atlascloud.ai": "atlascloud",
|
||||||
"api.deepseek.com": "deepseek",
|
"api.deepseek.com": "deepseek",
|
||||||
"api.moonshot.cn": "moonshot",
|
"api.moonshot.cn": "moonshot",
|
||||||
"api.siliconflow.cn": "siliconflow",
|
"api.siliconflow.cn": "siliconflow",
|
||||||
@@ -363,6 +364,7 @@ _HOST_TO_PROVIDER: dict[str, str] = {
|
|||||||
"api.minimaxi.com": "minimax",
|
"api.minimaxi.com": "minimax",
|
||||||
"api.kimi.com": "kimi", # kimi-coding shares this host
|
"api.kimi.com": "kimi", # kimi-coding shares this host
|
||||||
"openrouter.ai": "openrouter",
|
"openrouter.ai": "openrouter",
|
||||||
|
"api.novita.ai": "novita",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -377,16 +379,22 @@ def _provider_from_model(model: Any) -> str | None:
|
|||||||
(``ErrorNormalizationMiddleware``) then passes the exception
|
(``ErrorNormalizationMiddleware``) then passes the exception
|
||||||
through unchanged.
|
through unchanged.
|
||||||
"""
|
"""
|
||||||
cls_module = type(model).__module__ or ""
|
cls_modules = {cls.__module__ for cls in type(model).__mro__}
|
||||||
if cls_module.startswith("langchain_openrouter"):
|
|
||||||
|
def _uses_sdk(module_prefix: str) -> bool:
|
||||||
|
return any(module.startswith(module_prefix) for module in cls_modules)
|
||||||
|
|
||||||
|
if _uses_sdk("langchain_openrouter"):
|
||||||
return "openrouter"
|
return "openrouter"
|
||||||
if cls_module.startswith("langchain_google_genai"):
|
if _uses_sdk("langchain_google_genai"):
|
||||||
return "google_genai"
|
return "google_genai"
|
||||||
if cls_module.startswith("langchain_openai"):
|
if _uses_sdk("langchain_deepseek"):
|
||||||
|
return "deepseek"
|
||||||
|
if _uses_sdk("langchain_openai"):
|
||||||
return _lookup_host_or_compat(
|
return _lookup_host_or_compat(
|
||||||
getattr(model, "openai_api_base", None), module_tag="openai"
|
getattr(model, "openai_api_base", None), module_tag="openai"
|
||||||
)
|
)
|
||||||
if cls_module.startswith("langchain_anthropic"):
|
if _uses_sdk("langchain_anthropic"):
|
||||||
return _lookup_host_or_compat(
|
return _lookup_host_or_compat(
|
||||||
getattr(model, "anthropic_api_url", None), module_tag="anthropic"
|
getattr(model, "anthropic_api_url", None), module_tag="anthropic"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,17 +2,22 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
from collections.abc import Mapping, Sequence
|
from collections.abc import Mapping, Sequence
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from langchain_core.language_models.chat_models import BaseChatModel
|
from langchain_core.language_models.chat_models import BaseChatModel
|
||||||
from langchain_core.messages import BaseMessage, messages_from_dict, messages_to_dict
|
from langchain_core.messages import BaseMessage, messages_from_dict, messages_to_dict
|
||||||
from langchain_core.outputs import ChatGeneration, ChatResult
|
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult
|
||||||
from langchain_core.tools import BaseTool
|
from langchain_core.tools import BaseTool
|
||||||
from langchain_core.utils.function_calling import convert_to_openai_tool
|
from langchain_core.utils.function_calling import convert_to_openai_tool
|
||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
|
|
||||||
|
from EvoScientist.internal_service import internal_service_headers
|
||||||
|
|
||||||
|
from .contracts import EvoRuntimeError
|
||||||
|
|
||||||
|
|
||||||
class GatewayProxyChatModel(BaseChatModel):
|
class GatewayProxyChatModel(BaseChatModel):
|
||||||
gateway_url: str
|
gateway_url: str
|
||||||
@@ -44,6 +49,26 @@ class GatewayProxyChatModel(BaseChatModel):
|
|||||||
update={"bound_tools": serialized, "bound_tool_choice": tool_choice}
|
update={"bound_tools": serialized, "bound_tool_choice": tool_choice}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _attempt_id(self, run_manager: Any) -> str:
|
||||||
|
# The metering callback publishes the per-call tracing run_id into the
|
||||||
|
# shared configurable dict (see RecoverableMeteringCallback). The legacy
|
||||||
|
# BaseChatModel.astream path does NOT forward run_manager to _astream, so
|
||||||
|
# run_manager.run_id is unavailable here; the configurable channel is the
|
||||||
|
# only reliable source of the id the gateway registered under.
|
||||||
|
try:
|
||||||
|
from langgraph.config import get_config
|
||||||
|
|
||||||
|
config = get_config()
|
||||||
|
except Exception:
|
||||||
|
config = None
|
||||||
|
if isinstance(config, dict):
|
||||||
|
configurable = config.get("configurable")
|
||||||
|
if isinstance(configurable, dict):
|
||||||
|
current = configurable.get("ai4sci_attempt_id")
|
||||||
|
if current:
|
||||||
|
return str(current)
|
||||||
|
return str(getattr(run_manager, "run_id", None) or self.run_id)
|
||||||
|
|
||||||
def _generate(self, *args: Any, **kwargs: Any) -> ChatResult:
|
def _generate(self, *args: Any, **kwargs: Any) -> ChatResult:
|
||||||
del args, kwargs
|
del args, kwargs
|
||||||
raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH")
|
raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH")
|
||||||
@@ -56,8 +81,10 @@ class GatewayProxyChatModel(BaseChatModel):
|
|||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> ChatResult:
|
) -> ChatResult:
|
||||||
del stop, kwargs
|
del stop, kwargs
|
||||||
attempt_id = str(getattr(run_manager, "run_id", None) or self.run_id)
|
attempt_id = self._attempt_id(run_manager)
|
||||||
async with httpx.AsyncClient(timeout=httpx.Timeout(660.0, connect=5.0)) as client:
|
async with httpx.AsyncClient(
|
||||||
|
timeout=httpx.Timeout(660.0, connect=5.0)
|
||||||
|
) as client:
|
||||||
response = await client.post(
|
response = await client.post(
|
||||||
f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/invoke",
|
f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/invoke",
|
||||||
json={
|
json={
|
||||||
@@ -68,6 +95,7 @@ class GatewayProxyChatModel(BaseChatModel):
|
|||||||
"tools": self.bound_tools,
|
"tools": self.bound_tools,
|
||||||
"tool_choice": self.bound_tool_choice,
|
"tool_choice": self.bound_tool_choice,
|
||||||
},
|
},
|
||||||
|
headers=internal_service_headers(),
|
||||||
)
|
)
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
value = response.json()
|
value = response.json()
|
||||||
@@ -76,6 +104,101 @@ class GatewayProxyChatModel(BaseChatModel):
|
|||||||
raise RuntimeError("AI4SCI_MODEL_PROXY_RESPONSE_INVALID")
|
raise RuntimeError("AI4SCI_MODEL_PROXY_RESPONSE_INVALID")
|
||||||
return ChatResult(generations=[ChatGeneration(message=parsed[0])])
|
return ChatResult(generations=[ChatGeneration(message=parsed[0])])
|
||||||
|
|
||||||
|
async def _astream(
|
||||||
|
self,
|
||||||
|
messages: list[BaseMessage],
|
||||||
|
stop: list[str] | None = None,
|
||||||
|
run_manager: Any = None,
|
||||||
|
**kwargs: Any,
|
||||||
|
):
|
||||||
|
del stop, kwargs
|
||||||
|
attempt_id = self._attempt_id(run_manager)
|
||||||
|
payload = {
|
||||||
|
"run_id": self.run_id,
|
||||||
|
"attempt_id": attempt_id,
|
||||||
|
"envelope_signature": self.envelope_signature,
|
||||||
|
"messages": messages_to_dict(messages),
|
||||||
|
"tools": self.bound_tools,
|
||||||
|
"tool_choice": self.bound_tool_choice,
|
||||||
|
"stream": True,
|
||||||
|
}
|
||||||
|
async with httpx.AsyncClient(
|
||||||
|
timeout=httpx.Timeout(660.0, connect=5.0)
|
||||||
|
) as client:
|
||||||
|
async with client.stream(
|
||||||
|
"POST",
|
||||||
|
f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/stream",
|
||||||
|
json=payload,
|
||||||
|
headers=internal_service_headers(),
|
||||||
|
) as response:
|
||||||
|
response.raise_for_status()
|
||||||
|
saw_done = False
|
||||||
|
async for line in response.aiter_lines():
|
||||||
|
if not line.startswith("data:"):
|
||||||
|
continue
|
||||||
|
data = line[len("data:") :].strip()
|
||||||
|
if data == "[DONE]":
|
||||||
|
saw_done = True
|
||||||
|
break
|
||||||
|
chunk = json.loads(data)
|
||||||
|
if chunk.get("type") == "error":
|
||||||
|
code = str(chunk.get("code") or "MODEL_PROVIDER_ERROR")
|
||||||
|
message = str(chunk.get("message") or code)
|
||||||
|
details = {
|
||||||
|
key: value
|
||||||
|
for key, value in {
|
||||||
|
"http_status": chunk.get("status"),
|
||||||
|
"retryable": chunk.get("retryable"),
|
||||||
|
}.items()
|
||||||
|
if isinstance(value, int | bool)
|
||||||
|
}
|
||||||
|
raise EvoRuntimeError(
|
||||||
|
code,
|
||||||
|
message,
|
||||||
|
details=(details,) if details else (),
|
||||||
|
)
|
||||||
|
message = _chunk_to_message(chunk)
|
||||||
|
yield ChatGenerationChunk(
|
||||||
|
message=message,
|
||||||
|
generation_info=chunk.get("generation_info"),
|
||||||
|
)
|
||||||
|
if not saw_done:
|
||||||
|
raise RuntimeError("AI4SCI_MODEL_STREAM_INCOMPLETE")
|
||||||
|
|
||||||
|
|
||||||
|
def _chunk_to_message(chunk: dict[str, Any]) -> BaseMessage:
|
||||||
|
delta = chunk.get("delta") or {}
|
||||||
|
message_dict = delta.get("message")
|
||||||
|
if message_dict is None:
|
||||||
|
raise RuntimeError("AI4SCI_MODEL_STREAM_DELTA_INVALID")
|
||||||
|
parsed = messages_from_dict([message_dict])
|
||||||
|
if len(parsed) != 1:
|
||||||
|
raise RuntimeError("AI4SCI_MODEL_STREAM_DELTA_INVALID")
|
||||||
|
message = parsed[0]
|
||||||
|
summary = delta.get("reasoning_summary")
|
||||||
|
if isinstance(summary, str) and summary:
|
||||||
|
content = list(message.content) if isinstance(message.content, list) else []
|
||||||
|
existing_summary = "".join(
|
||||||
|
str(part.get("text") or "")
|
||||||
|
for block in content
|
||||||
|
if isinstance(block, Mapping) and block.get("type") == "reasoning"
|
||||||
|
for part in (block.get("summary") or [])
|
||||||
|
if isinstance(part, Mapping) and part.get("type") == "summary_text"
|
||||||
|
)
|
||||||
|
if summary.startswith(existing_summary):
|
||||||
|
summary_delta = summary[len(existing_summary):]
|
||||||
|
elif existing_summary.startswith(summary):
|
||||||
|
summary_delta = ""
|
||||||
|
else:
|
||||||
|
summary_delta = summary
|
||||||
|
if summary_delta:
|
||||||
|
content.append({
|
||||||
|
"type": "reasoning",
|
||||||
|
"summary": [{"type": "summary_text", "text": summary_delta}],
|
||||||
|
})
|
||||||
|
message.content = content
|
||||||
|
return message
|
||||||
|
|
||||||
|
|
||||||
def proxy_from_config(
|
def proxy_from_config(
|
||||||
value: Mapping[str, Any], *, provider_id: str = "", model_id: str = ""
|
value: Mapping[str, Any], *, provider_id: str = "", model_id: str = ""
|
||||||
|
|||||||
@@ -3,7 +3,9 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
import logging
|
||||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||||
|
from types import SimpleNamespace
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from langchain_core.language_models.chat_models import BaseChatModel
|
from langchain_core.language_models.chat_models import BaseChatModel
|
||||||
@@ -20,6 +22,60 @@ from langchain_core.tools import BaseTool
|
|||||||
from langchain_core.utils.function_calling import convert_to_openai_tool
|
from langchain_core.utils.function_calling import convert_to_openai_tool
|
||||||
from pydantic import Field, SecretStr
|
from pydantic import Field, SecretStr
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _cleanup_outcome(
|
||||||
|
primary: BaseException | None, errors: list[BaseException]
|
||||||
|
) -> None:
|
||||||
|
if primary is not None or not errors:
|
||||||
|
return
|
||||||
|
for error in errors:
|
||||||
|
if not isinstance(error, Exception):
|
||||||
|
raise error
|
||||||
|
raise RuntimeError("MODEL_PROVIDER_CLEANUP_ERROR") from None
|
||||||
|
|
||||||
|
|
||||||
|
def _record_cleanup_error(resource: str) -> None:
|
||||||
|
# Never log exception messages, reprs or tracebacks containing provider secrets.
|
||||||
|
try:
|
||||||
|
logger.warning("MODEL_PROVIDER_CLEANUP_ERROR resource=%s", resource)
|
||||||
|
except Exception:
|
||||||
|
# A broken logging handler must not replace the provider exception either.
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def _close_owned(client: Any, stream: Any, primary: BaseException | None) -> None:
|
||||||
|
errors: list[BaseException] = []
|
||||||
|
for resource, target in (("stream", stream), ("client", client)):
|
||||||
|
if target is None:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
target.close()
|
||||||
|
except BaseException as error:
|
||||||
|
errors.append(error)
|
||||||
|
_record_cleanup_error(resource)
|
||||||
|
_cleanup_outcome(primary, errors)
|
||||||
|
|
||||||
|
|
||||||
|
async def _aclose_owned(
|
||||||
|
client: Any, stream: Any, primary: BaseException | None
|
||||||
|
) -> None:
|
||||||
|
errors: list[BaseException] = []
|
||||||
|
for resource in ("stream", "async_client", "client"):
|
||||||
|
try:
|
||||||
|
if resource == "stream":
|
||||||
|
if stream is not None:
|
||||||
|
await stream.close()
|
||||||
|
elif resource == "async_client":
|
||||||
|
await client.aio.aclose()
|
||||||
|
else:
|
||||||
|
client.close()
|
||||||
|
except BaseException as error:
|
||||||
|
errors.append(error)
|
||||||
|
_record_cleanup_error(resource)
|
||||||
|
_cleanup_outcome(primary, errors)
|
||||||
|
|
||||||
|
|
||||||
class GeminiInteractionsChatModel(BaseChatModel):
|
class GeminiInteractionsChatModel(BaseChatModel):
|
||||||
"""Minimal native bridge that preserves signed Provider content blocks."""
|
"""Minimal native bridge that preserves signed Provider content blocks."""
|
||||||
@@ -96,7 +152,7 @@ class GeminiInteractionsChatModel(BaseChatModel):
|
|||||||
"generation_config": generation_config,
|
"generation_config": generation_config,
|
||||||
"tools": list(self.bound_tools),
|
"tools": list(self.bound_tools),
|
||||||
"store": False,
|
"store": False,
|
||||||
"stream": False,
|
"stream": True,
|
||||||
}
|
}
|
||||||
|
|
||||||
def _generate(
|
def _generate(
|
||||||
@@ -108,8 +164,20 @@ class GeminiInteractionsChatModel(BaseChatModel):
|
|||||||
request = self._request(messages)
|
request = self._request(messages)
|
||||||
if stop:
|
if stop:
|
||||||
request["generation_config"]["stop_sequences"] = stop
|
request["generation_config"]["stop_sequences"] = stop
|
||||||
response = self._client().interactions.create(**request)
|
client = self._client()
|
||||||
return _chat_result(response)
|
stream = None
|
||||||
|
state = _InteractionAccumulator()
|
||||||
|
primary = None
|
||||||
|
try:
|
||||||
|
stream = client.interactions.create(**request)
|
||||||
|
for event in stream:
|
||||||
|
state.accept(_dump(event))
|
||||||
|
return state.result()
|
||||||
|
except BaseException as error:
|
||||||
|
primary = error
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
_close_owned(client, stream, primary)
|
||||||
|
|
||||||
async def _agenerate(
|
async def _agenerate(
|
||||||
self,
|
self,
|
||||||
@@ -120,8 +188,20 @@ class GeminiInteractionsChatModel(BaseChatModel):
|
|||||||
request = self._request(messages)
|
request = self._request(messages)
|
||||||
if stop:
|
if stop:
|
||||||
request["generation_config"]["stop_sequences"] = stop
|
request["generation_config"]["stop_sequences"] = stop
|
||||||
response = await self._client().aio.interactions.create(**request)
|
client = self._client()
|
||||||
return _chat_result(response)
|
stream = None
|
||||||
|
state = _InteractionAccumulator()
|
||||||
|
primary = None
|
||||||
|
try:
|
||||||
|
stream = await client.aio.interactions.create(**request)
|
||||||
|
async for event in stream:
|
||||||
|
state.accept(_dump(event))
|
||||||
|
return state.result()
|
||||||
|
except BaseException as error:
|
||||||
|
primary = error
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
await _aclose_owned(client, stream, primary)
|
||||||
|
|
||||||
async def _astream(
|
async def _astream(
|
||||||
self,
|
self,
|
||||||
@@ -133,8 +213,22 @@ class GeminiInteractionsChatModel(BaseChatModel):
|
|||||||
request["stream"] = True
|
request["stream"] = True
|
||||||
if stop:
|
if stop:
|
||||||
request["generation_config"]["stop_sequences"] = stop
|
request["generation_config"]["stop_sequences"] = stop
|
||||||
stream = await self._client().aio.interactions.create(**request)
|
client = self._client()
|
||||||
|
stream = None
|
||||||
|
primary = None
|
||||||
|
try:
|
||||||
|
stream = await client.aio.interactions.create(**request)
|
||||||
|
async for chunk in self._astream_events(stream):
|
||||||
|
yield chunk
|
||||||
|
except BaseException as error:
|
||||||
|
primary = error
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
await _aclose_owned(client, stream, primary)
|
||||||
|
|
||||||
|
async def _astream_events(self, stream: Any) -> AsyncIterator[ChatGenerationChunk]:
|
||||||
blocks: dict[int, dict[str, Any]] = {}
|
blocks: dict[int, dict[str, Any]] = {}
|
||||||
|
completed = False
|
||||||
async for event in stream:
|
async for event in stream:
|
||||||
payload = _dump(event)
|
payload = _dump(event)
|
||||||
event_type = payload.get("event_type")
|
event_type = payload.get("event_type")
|
||||||
@@ -175,6 +269,7 @@ class GeminiInteractionsChatModel(BaseChatModel):
|
|||||||
if event_type == "error":
|
if event_type == "error":
|
||||||
raise RuntimeError("MODEL_PROVIDER_PROTOCOL_ERROR")
|
raise RuntimeError("MODEL_PROVIDER_PROTOCOL_ERROR")
|
||||||
if event_type == "interaction.complete":
|
if event_type == "interaction.complete":
|
||||||
|
completed = True
|
||||||
interaction = payload.get("interaction") or {}
|
interaction = payload.get("interaction") or {}
|
||||||
ordered_blocks = [blocks[index] for index in sorted(blocks)]
|
ordered_blocks = [blocks[index] for index in sorted(blocks)]
|
||||||
yield ChatGenerationChunk(
|
yield ChatGenerationChunk(
|
||||||
@@ -188,9 +283,7 @@ class GeminiInteractionsChatModel(BaseChatModel):
|
|||||||
provider_request_id=interaction.get("id"),
|
provider_request_id=interaction.get("id"),
|
||||||
),
|
),
|
||||||
response_metadata={
|
response_metadata={
|
||||||
"model_name": str(
|
"model_name": _interaction_model_name(interaction),
|
||||||
(interaction.get("model") or {}).get("id") or ""
|
|
||||||
),
|
|
||||||
"finish_reason": str(
|
"finish_reason": str(
|
||||||
interaction.get("status") or "unknown"
|
interaction.get("status") or "unknown"
|
||||||
),
|
),
|
||||||
@@ -198,6 +291,50 @@ class GeminiInteractionsChatModel(BaseChatModel):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if not completed:
|
||||||
|
raise RuntimeError("MODEL_PROVIDER_PROTOCOL_ERROR")
|
||||||
|
|
||||||
|
|
||||||
|
class _InteractionAccumulator:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.blocks: dict[int, dict[str, Any]] = {}
|
||||||
|
self.interaction: dict[str, Any] | None = None
|
||||||
|
|
||||||
|
def accept(self, payload: dict[str, Any]) -> None:
|
||||||
|
kind = payload.get("event_type")
|
||||||
|
if kind == "content.start":
|
||||||
|
self.blocks[int(payload["index"])] = dict(payload.get("content") or {})
|
||||||
|
elif kind == "content.delta":
|
||||||
|
_merge_stream_delta(
|
||||||
|
self.blocks.setdefault(int(payload["index"]), {}),
|
||||||
|
dict(payload.get("delta") or {}),
|
||||||
|
)
|
||||||
|
elif kind == "error":
|
||||||
|
raise RuntimeError("MODEL_PROVIDER_PROTOCOL_ERROR")
|
||||||
|
elif kind == "interaction.complete":
|
||||||
|
self.interaction = payload.get("interaction") or {}
|
||||||
|
|
||||||
|
def result(self) -> ChatResult:
|
||||||
|
if self.interaction is None:
|
||||||
|
raise RuntimeError("MODEL_PROVIDER_PROTOCOL_ERROR")
|
||||||
|
interaction = self.interaction
|
||||||
|
return _chat_result(
|
||||||
|
SimpleNamespace(
|
||||||
|
outputs=[self.blocks[index] for index in sorted(self.blocks)],
|
||||||
|
usage=interaction.get("usage"),
|
||||||
|
id=interaction.get("id"),
|
||||||
|
status=interaction.get("status", "unknown"),
|
||||||
|
model=SimpleNamespace(id=_interaction_model_name(interaction)),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _interaction_model_name(interaction: Mapping[str, Any]) -> str:
|
||||||
|
model = interaction.get("model")
|
||||||
|
if isinstance(model, Mapping):
|
||||||
|
return str(model.get("id") or "")
|
||||||
|
return str(model or "")
|
||||||
|
|
||||||
|
|
||||||
def _message_content(message: BaseMessage) -> list[dict[str, Any]]:
|
def _message_content(message: BaseMessage) -> list[dict[str, Any]]:
|
||||||
if isinstance(message, AIMessage):
|
if isinstance(message, AIMessage):
|
||||||
|
|||||||
@@ -0,0 +1,112 @@
|
|||||||
|
"""Local durable initialization gate. All graph access must use this gate.
|
||||||
|
|
||||||
|
The file lock covers graph writes as well as owner changes; a row CAS alone
|
||||||
|
cannot fence an old writer while it is suspended in a graph await.
|
||||||
|
"""
|
||||||
|
import asyncio
|
||||||
|
import fcntl
|
||||||
|
import sqlite3
|
||||||
|
from contextlib import asynccontextmanager, closing
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from langgraph.graph import START, END
|
||||||
|
|
||||||
|
|
||||||
|
class SqliteInitializationStore:
|
||||||
|
def __init__(self, path):
|
||||||
|
if str(path) == ":memory:":
|
||||||
|
raise ValueError("initialization requires a durable file")
|
||||||
|
self.path = str(Path(path).resolve())
|
||||||
|
with closing(self.connect()) as db, db:
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS history_initialization (
|
||||||
|
key TEXT PRIMARY KEY, attempt TEXT NOT NULL, digest TEXT NOT NULL,
|
||||||
|
owner TEXT NOT NULL, fence INTEGER NOT NULL, status TEXT NOT NULL)""")
|
||||||
|
|
||||||
|
def connect(self):
|
||||||
|
db = sqlite3.connect(self.path)
|
||||||
|
db.row_factory = sqlite3.Row
|
||||||
|
return db
|
||||||
|
|
||||||
|
def get(self, key):
|
||||||
|
with closing(self.connect()) as db:
|
||||||
|
row = db.execute("SELECT * FROM history_initialization WHERE key=?", (key,)).fetchone()
|
||||||
|
return dict(row) if row else None
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def locked(self, key):
|
||||||
|
with open(self.path + "." + key.split(":")[-1] + ".lock", "a") as lock:
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
fcntl.flock(lock, fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||||
|
break
|
||||||
|
except BlockingIOError:
|
||||||
|
await asyncio.sleep(0.01)
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
fcntl.flock(lock, fcntl.LOCK_UN)
|
||||||
|
|
||||||
|
def claim(self, key, attempt, digest, owner, fence):
|
||||||
|
if not isinstance(attempt, str) or not attempt.strip() or not isinstance(owner, str) or not owner.strip():
|
||||||
|
raise ValueError("attempt and owner required")
|
||||||
|
if type(fence) is not int or fence < 1:
|
||||||
|
raise ValueError("positive owner fence required")
|
||||||
|
with closing(self.connect()) as db, db:
|
||||||
|
db.execute("BEGIN IMMEDIATE")
|
||||||
|
row = db.execute("SELECT * FROM history_initialization WHERE key=?", (key,)).fetchone()
|
||||||
|
if row:
|
||||||
|
if row["attempt"] != attempt or row["digest"] != digest:
|
||||||
|
raise ValueError("history attempt/content conflict")
|
||||||
|
if fence < row["fence"] or (fence == row["fence"] and owner != row["owner"]):
|
||||||
|
raise ValueError("stale history owner fence")
|
||||||
|
if row["status"] not in ("INITIALIZING", "READY"):
|
||||||
|
raise ValueError("history checkpoint already consumed")
|
||||||
|
db.execute("UPDATE history_initialization SET owner=?, fence=? WHERE key=?", (owner, fence, key))
|
||||||
|
else:
|
||||||
|
db.execute("INSERT INTO history_initialization VALUES (?,?,?,?,?,?)",
|
||||||
|
(key, attempt, digest, owner, fence, "INITIALIZING"))
|
||||||
|
|
||||||
|
def status(self, key, attempt, owner, fence, before, after):
|
||||||
|
with closing(self.connect()) as db, db:
|
||||||
|
changed = db.execute("""UPDATE history_initialization SET status=?
|
||||||
|
WHERE key=? AND attempt=? AND owner=? AND fence=? AND status=?""",
|
||||||
|
(after, key, attempt, owner, fence, before)).rowcount
|
||||||
|
if changed != 1:
|
||||||
|
raise ValueError("history requires READY and current owner fence")
|
||||||
|
|
||||||
|
|
||||||
|
async def initialize(graph, key, messages, digest, store, attempt, owner, fence):
|
||||||
|
config = {"configurable": {"thread_id": key}}
|
||||||
|
marker = {"history_attempt": attempt, "history_digest": digest}
|
||||||
|
async with store.locked(key):
|
||||||
|
current = await graph.aget_state(config)
|
||||||
|
if store.get(key) is None and current.created_at is not None:
|
||||||
|
raise ValueError("unowned history checkpoint already exists")
|
||||||
|
store.claim(key, attempt, digest, owner, fence)
|
||||||
|
if current.created_at is None:
|
||||||
|
await graph.aupdate_state(dict(config, metadata=dict(marker, history_stage="START")),
|
||||||
|
{"messages": list(messages)}, as_node=START)
|
||||||
|
current = await graph.aget_state(config)
|
||||||
|
metadata = current.metadata or {}
|
||||||
|
if any(metadata.get(k) != v for k, v in marker.items()):
|
||||||
|
raise ValueError("unowned history checkpoint")
|
||||||
|
if current.values.get("messages", []) != list(messages):
|
||||||
|
raise ValueError("history checkpoint content changed")
|
||||||
|
stage = metadata.get("history_stage")
|
||||||
|
if stage == "START":
|
||||||
|
await graph.aupdate_state(dict(config, metadata=dict(marker, history_stage="END")),
|
||||||
|
None, as_node=END)
|
||||||
|
current = await graph.aget_state(config)
|
||||||
|
elif stage != "END":
|
||||||
|
raise ValueError("unknown history initialization stage")
|
||||||
|
if current.next or current.tasks:
|
||||||
|
raise ValueError("history checkpoint is not READY")
|
||||||
|
store.status(key, attempt, owner, fence, store.get(key)["status"], "READY")
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
async def invoke(graph, key, store, attempt, owner, fence, input):
|
||||||
|
async with store.locked(key):
|
||||||
|
# Consume before execution: uncertain invocation must not be replayed.
|
||||||
|
store.status(key, attempt, owner, fence, "READY", "CONSUMED")
|
||||||
|
return await graph.ainvoke(input, {"configurable": {"thread_id": key}})
|
||||||
@@ -0,0 +1,244 @@
|
|||||||
|
"""Committed history to a fresh turn checkpoint, never an old stack resume.
|
||||||
|
|
||||||
|
The host supplies authorized, revision-consistent records and owns the turn
|
||||||
|
fence across creation and invocation. This module does not read or modify PG.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import copy
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Sequence
|
||||||
|
|
||||||
|
from langchain_core.messages import BaseMessage, HumanMessage, ToolMessage
|
||||||
|
from langchain_core.runnables import RunnableConfig
|
||||||
|
from langgraph.graph import END, START
|
||||||
|
|
||||||
|
from .patches import _sanitize_openai_tool_history, _validate_openai_tool_history
|
||||||
|
from .history_initialization import SqliteInitializationStore, initialize, invoke
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class HistoryScope:
|
||||||
|
tenant_id: str
|
||||||
|
thread_id: str
|
||||||
|
turn_id: str
|
||||||
|
workspace_id: str
|
||||||
|
graph_version: str
|
||||||
|
tool_version: str
|
||||||
|
history_revision: int
|
||||||
|
|
||||||
|
def validate(self) -> None:
|
||||||
|
for value in (self.tenant_id, self.thread_id, self.turn_id,
|
||||||
|
self.workspace_id, self.graph_version, self.tool_version):
|
||||||
|
if not isinstance(value, str) or not value.strip():
|
||||||
|
raise ValueError("history scope fields must be nonempty strings")
|
||||||
|
if type(self.history_revision) is not int or self.history_revision < 0:
|
||||||
|
raise ValueError("invalid history revision")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class HistoryRecord:
|
||||||
|
message_id: str
|
||||||
|
revision: int
|
||||||
|
message: BaseMessage
|
||||||
|
partial: bool = False
|
||||||
|
file_refs: tuple[str, ...] = ()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class NormalizedHistory:
|
||||||
|
scope: HistoryScope
|
||||||
|
records: tuple[HistoryRecord, ...]
|
||||||
|
messages: tuple[BaseMessage, ...]
|
||||||
|
|
||||||
|
|
||||||
|
def _check_scope(source: HistoryScope, expected: HistoryScope) -> None:
|
||||||
|
source.validate()
|
||||||
|
expected.validate()
|
||||||
|
if source != expected:
|
||||||
|
raise ValueError("history scope/version/revision mismatch")
|
||||||
|
|
||||||
|
|
||||||
|
def _text(content: Any) -> str:
|
||||||
|
# Only text is replayed. Media bytes stay outside model history; authorized
|
||||||
|
# file references are supplied separately by the host's record adapter.
|
||||||
|
if isinstance(content, str):
|
||||||
|
return "[inline media omitted]" if "base64," in content or "data:" in content else content
|
||||||
|
if isinstance(content, list):
|
||||||
|
return "\n".join(_text(b.get("text", "")) for b in content
|
||||||
|
if isinstance(b, dict) and b.get("type") == "text")
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_history(records: Sequence[HistoryRecord], *, source: HistoryScope,
|
||||||
|
expected: HistoryScope) -> NormalizedHistory:
|
||||||
|
"""Pure normalization of already-authorized stable-ID committed records.
|
||||||
|
|
||||||
|
Complete parallel tool pairs survive; unmatched calls become sourced facts,
|
||||||
|
never invented ToolMessages. Raw messages and provider metadata are untouched.
|
||||||
|
"""
|
||||||
|
_check_scope(source, expected)
|
||||||
|
seen: set[str] = set()
|
||||||
|
messages = []
|
||||||
|
notes = []
|
||||||
|
for record in records:
|
||||||
|
if (not isinstance(record.message_id, str) or not record.message_id.strip()
|
||||||
|
or record.message_id in seen or type(record.revision) is not int
|
||||||
|
or record.revision < 0 or not isinstance(record.message, BaseMessage)):
|
||||||
|
raise ValueError("invalid or duplicate history record")
|
||||||
|
seen.add(record.message_id)
|
||||||
|
message = copy.deepcopy(record.message)
|
||||||
|
message.id = record.message_id
|
||||||
|
messages.append(message)
|
||||||
|
calls = {}
|
||||||
|
results = set()
|
||||||
|
for message in messages:
|
||||||
|
for call in getattr(message, "tool_calls", []):
|
||||||
|
call_id = call.get("id")
|
||||||
|
if not call_id or call_id in calls:
|
||||||
|
raise ValueError("missing or duplicate tool call ID")
|
||||||
|
calls[call_id] = call.get("name")
|
||||||
|
if isinstance(message, ToolMessage):
|
||||||
|
call_id = message.tool_call_id
|
||||||
|
if not call_id or call_id not in calls or call_id in results:
|
||||||
|
raise ValueError("unassociated or duplicate tool result")
|
||||||
|
if message.name is not None and message.name != calls[call_id]:
|
||||||
|
raise ValueError("tool result name conflicts with call ID")
|
||||||
|
results.add(call_id)
|
||||||
|
repaired = _sanitize_openai_tool_history(messages)
|
||||||
|
retained = {m.id: m for m in repaired}
|
||||||
|
for record, original in zip(records, messages):
|
||||||
|
repaired_message = retained.get(record.message_id)
|
||||||
|
kept = {c.get("id") for c in getattr(repaired_message, "tool_calls", [])}
|
||||||
|
missing = [c.get("id") or "unidentified" for c in getattr(original, "tool_calls", [])
|
||||||
|
if c.get("id") not in kept]
|
||||||
|
provenance = f"source={record.message_id}@{record.revision}"
|
||||||
|
if missing or getattr(original, "invalid_tool_calls", []):
|
||||||
|
notes.append(f"[{provenance}] incomplete tool calls (result unavailable; execution unknown): {missing}")
|
||||||
|
if record.partial:
|
||||||
|
notes.append(f"[{provenance}] partial result, not a completed answer")
|
||||||
|
for ref in record.file_refs:
|
||||||
|
if not isinstance(ref, str) or not ref.strip() or "data:" in ref or "base64" in ref:
|
||||||
|
raise ValueError("invalid file reference")
|
||||||
|
notes.append(f"[{provenance}] file reference: {ref}")
|
||||||
|
for message in repaired:
|
||||||
|
message.content = _text(message.content)
|
||||||
|
message.additional_kwargs = {}
|
||||||
|
message.response_metadata = {}
|
||||||
|
if hasattr(message, "artifact"):
|
||||||
|
message.artifact = None
|
||||||
|
# Tool arguments are historical context, not executable input. Reject
|
||||||
|
# embedded binary payload rather than stringify it into provider text.
|
||||||
|
if "base64" in json.dumps(getattr(message, "tool_calls", [])):
|
||||||
|
raise ValueError("inline binary tool arguments are not replayable")
|
||||||
|
if notes:
|
||||||
|
repaired.append(HumanMessage(content="Historical context notes:\n" + "\n".join(notes),
|
||||||
|
id="history-notes:" + str(source.history_revision)))
|
||||||
|
_validate_openai_tool_history(repaired)
|
||||||
|
return NormalizedHistory(source, tuple(copy.deepcopy(records)), tuple(repaired))
|
||||||
|
|
||||||
|
|
||||||
|
def committed_history_input(history: dict, current: dict, *, thread_id: str,
|
||||||
|
run_id: str, checkpoint_exists: bool = False) -> dict:
|
||||||
|
"""Project PG display records into the existing worker's message channel.
|
||||||
|
|
||||||
|
Display tool items are sourced facts, not provider tool protocol. The
|
||||||
|
reducer sentinel is only used when the host attests no compatible checkpoint
|
||||||
|
exists. Later turns append without resetting the summarizer's cut indexes.
|
||||||
|
No second invocation or checkpoint mutation is performed by the HTTP host.
|
||||||
|
"""
|
||||||
|
from langchain_core.messages import AIMessage, RemoveMessage, convert_to_messages
|
||||||
|
from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
||||||
|
|
||||||
|
if (history.get("schema") != "ai4sci.committed-history.v1"
|
||||||
|
or history.get("thread_id") != thread_id
|
||||||
|
or not isinstance(history.get("records"), list)
|
||||||
|
or not isinstance(current, dict)
|
||||||
|
or not isinstance(current.get("messages"), list)):
|
||||||
|
raise ValueError("invalid committed history input")
|
||||||
|
scope = HistoryScope(
|
||||||
|
tenant_id=history["user_uid"], thread_id=thread_id, turn_id=run_id,
|
||||||
|
workspace_id=thread_id, graph_version="worker-input-v1",
|
||||||
|
tool_version="display-facts-v2", history_revision=history["conversation_revision"],
|
||||||
|
)
|
||||||
|
records = []
|
||||||
|
for row in sorted(history["records"], key=lambda r: (r["message_index"], r["message_id"])):
|
||||||
|
if row["message_id"] == history.get("excluded_message_id"):
|
||||||
|
continue
|
||||||
|
payload = row["payload"]
|
||||||
|
role = row["role"]
|
||||||
|
if role not in {"user", "assistant"}:
|
||||||
|
raise ValueError("unsupported committed history role")
|
||||||
|
text = []
|
||||||
|
refs = []
|
||||||
|
if role == "user":
|
||||||
|
text.append(_text(payload.get("content", "")))
|
||||||
|
for attachment in payload.get("attached_files") or []:
|
||||||
|
if isinstance(attachment, dict):
|
||||||
|
ref = attachment.get("virtual_path") or attachment.get("path")
|
||||||
|
if isinstance(ref, str) and ref.startswith("/workspace/"):
|
||||||
|
refs.append(ref)
|
||||||
|
else:
|
||||||
|
for item in sorted(payload.get("items") or [], key=lambda i: i["item_sequence"]):
|
||||||
|
if item.get("status") == "superseded":
|
||||||
|
continue
|
||||||
|
kind = item.get("type")
|
||||||
|
source = f"[{row['message_id']}@{row['revision']}:{item['item_id']}]"
|
||||||
|
if kind == "message":
|
||||||
|
text.extend(part["text"] for part in item.get("content", [])
|
||||||
|
if part.get("type") in {"output_text", "refusal"})
|
||||||
|
elif kind in {"tool_call", "tool_output", "agent_status", "summarization"}:
|
||||||
|
fields = {key: item[key] for key in
|
||||||
|
("name", "input", "output", "result", "summary", "status") if key in item}
|
||||||
|
text.append(f"{source} historical {kind}: " + json.dumps(fields, ensure_ascii=False))
|
||||||
|
elif kind == "artifact" and str(item.get("virtual_path", "")).startswith("/workspace/"):
|
||||||
|
refs.append(item["virtual_path"])
|
||||||
|
message = (HumanMessage if role == "user" else AIMessage)(content="\n".join(text))
|
||||||
|
records.append(HistoryRecord(row["message_id"], row["revision"], message,
|
||||||
|
partial=bool(payload.get("incomplete")), file_refs=tuple(refs)))
|
||||||
|
normalized = normalize_history(records, source=scope, expected=scope)
|
||||||
|
incoming = convert_to_messages(copy.deepcopy(current["messages"]))
|
||||||
|
for index, message in enumerate(incoming):
|
||||||
|
# Stable IDs ensure replay cannot duplicate the current turn either.
|
||||||
|
message.id = f"current:{run_id}:{index}"
|
||||||
|
if checkpoint_exists:
|
||||||
|
return {**current, "messages": [m.model_dump(mode="json") for m in incoming]}
|
||||||
|
messages = [RemoveMessage(id=REMOVE_ALL_MESSAGES), *normalized.messages, *incoming]
|
||||||
|
# Old summary cut indexes refer to the replaced checkpoint message list.
|
||||||
|
# The existing summarizer/budget middleware recomputes them for this input.
|
||||||
|
return {**current, "messages": [m.model_dump(mode="json") for m in messages],
|
||||||
|
"_summarization_event": None}
|
||||||
|
|
||||||
|
|
||||||
|
def history_key(scope: HistoryScope) -> str:
|
||||||
|
scope.validate()
|
||||||
|
identity = json.dumps(list(vars(scope).values()), separators=(",", ":"))
|
||||||
|
return "history-v1:" + hashlib.sha256(identity.encode()).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
async def invoke_history_checkpoint(graph, *, scope, store, attempt, owner, fence, input):
|
||||||
|
return await invoke(graph, history_key(scope), store, attempt, owner, fence, input)
|
||||||
|
|
||||||
|
|
||||||
|
async def create_history_checkpoint(graph: Any, history: NormalizedHistory, *,
|
||||||
|
expected: HistoryScope, store=None, attempt=None,
|
||||||
|
owner=None, fence=None) -> RunnableConfig:
|
||||||
|
"""Seed through a durable gate; dispatch only via invoke_history_checkpoint.
|
||||||
|
|
||||||
|
A matching attempt/content may recover initialization, never execution.
|
||||||
|
The host owns authorization and issues monotonically increasing fences.
|
||||||
|
"""
|
||||||
|
_check_scope(history.scope, expected)
|
||||||
|
normalized = normalize_history(history.records, source=history.scope, expected=expected)
|
||||||
|
key = history_key(expected)
|
||||||
|
if store is not None:
|
||||||
|
payload = [{"id": r.message_id, "revision": r.revision,
|
||||||
|
"message": r.message.model_dump(mode="json"),
|
||||||
|
"partial": r.partial, "files": r.file_refs} for r in history.records]
|
||||||
|
digest = hashlib.sha256(json.dumps(payload, sort_keys=True,
|
||||||
|
separators=(",", ":")).encode()).hexdigest()
|
||||||
|
return await initialize(graph, key, normalized.messages, digest,
|
||||||
|
store, attempt, owner, fence)
|
||||||
|
raise ValueError("durable initialization store is required")
|
||||||
@@ -0,0 +1,370 @@
|
|||||||
|
"""Opt-in host identity prototype. Records never prove resource quiescence.
|
||||||
|
|
||||||
|
Trusted local host API, not an authenticated remote control endpoint. A future
|
||||||
|
PG host store can implement this protocol without introducing a dispatcher.
|
||||||
|
Only the runtime resource owner may attest cleanup and release its claim.
|
||||||
|
"""
|
||||||
|
from contextlib import closing, contextmanager
|
||||||
|
from pathlib import Path
|
||||||
|
import sqlite3
|
||||||
|
import json
|
||||||
|
import hashlib
|
||||||
|
import os
|
||||||
|
from typing import Protocol
|
||||||
|
|
||||||
|
from .contracts import EvoRuntimeError, canonical_json_v1
|
||||||
|
|
||||||
|
|
||||||
|
class HostExecutionRegistry(Protocol):
|
||||||
|
def claim_checkpoint(self, execution_id: str, *, store_id: str, checkpoint_thread_id: str, subject_id: str) -> None: ...
|
||||||
|
def release_unbound_checkpoint(self, execution_id: str) -> None: ...
|
||||||
|
def continuation(self, execution_id: str) -> dict: ...
|
||||||
|
def prepare_terminal(self, execution_id: str, *, event: dict) -> None: ...
|
||||||
|
def terminal_intent(self, execution_id: str) -> dict | None: ...
|
||||||
|
def confirm_terminal(self, execution_id: str, *, digest: str) -> None: ...
|
||||||
|
def finish(self, execution_id: str, *, outcome: str, checkpoint_id: str = "") -> None: ...
|
||||||
|
def bind(self, *, execution_id: str, grant_id: str, digest: str,
|
||||||
|
thread_id: str, turn_id: str, predecessor_execution_id: str = "",
|
||||||
|
predecessor_checkpoint_id: str = "", predecessor_owner_epoch: int = 0,
|
||||||
|
continuation_pending_hash: str = "", continuation_decision_hash: str = "") -> None: ...
|
||||||
|
def lookup_grant(self, grant_id: str, digest: str) -> dict | None: ...
|
||||||
|
def inspect(self, execution_id: str) -> dict: ...
|
||||||
|
def transfer_control(self, execution_id: str, *, expected_epoch: int, new_epoch: int) -> int: ...
|
||||||
|
def require_control(self, execution_id: str, *, owner_epoch: int) -> None: ...
|
||||||
|
def inspect_control(self, execution_id: str, *, owner_epoch: int | None, boot_id: str | None) -> dict: ...
|
||||||
|
def accept_cancel(self, execution_id: str, *, owner_epoch: int | None, boot_id: str | None, reason: str) -> int: ...
|
||||||
|
|
||||||
|
|
||||||
|
class SQLiteHostRegistry:
|
||||||
|
def __init__(self, path: str | Path, *, host_id: str, boot_id: str):
|
||||||
|
self.path, self.host_id, self.boot_id = str(path), host_id, boot_id
|
||||||
|
directory = Path(path).parent
|
||||||
|
directory.mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||||
|
# Exclusive creation gives new user-content stores safe defaults without
|
||||||
|
# changing permissions on an existing host's directory or database.
|
||||||
|
try:
|
||||||
|
fd = os.open(self.path, os.O_CREAT | os.O_EXCL | os.O_WRONLY, 0o600)
|
||||||
|
except FileExistsError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
os.close(fd)
|
||||||
|
with self._transaction() as db:
|
||||||
|
db.execute("BEGIN IMMEDIATE")
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS executions (
|
||||||
|
execution_id TEXT PRIMARY KEY, grant_id TEXT UNIQUE NOT NULL,
|
||||||
|
digest TEXT NOT NULL, thread_id TEXT NOT NULL,
|
||||||
|
turn_id TEXT NOT NULL, host_id TEXT NOT NULL, boot_id TEXT NOT NULL,
|
||||||
|
owner_epoch INTEGER NOT NULL DEFAULT 1)""")
|
||||||
|
schema = db.execute("SELECT sql FROM sqlite_master WHERE name='executions'").fetchone()[0]
|
||||||
|
if "UNIQUE(thread_id, turn_id)" in schema:
|
||||||
|
db.execute("ALTER TABLE executions RENAME TO legacy_executions")
|
||||||
|
db.execute("""CREATE TABLE executions (
|
||||||
|
execution_id TEXT PRIMARY KEY, grant_id TEXT UNIQUE NOT NULL,
|
||||||
|
digest TEXT NOT NULL, thread_id TEXT NOT NULL, turn_id TEXT NOT NULL,
|
||||||
|
host_id TEXT NOT NULL, boot_id TEXT NOT NULL,
|
||||||
|
owner_epoch INTEGER NOT NULL DEFAULT 1)""")
|
||||||
|
db.execute("INSERT INTO executions SELECT * FROM legacy_executions")
|
||||||
|
db.execute("DROP TABLE legacy_executions")
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS cancel_intents (
|
||||||
|
intent_id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
|
execution_id TEXT NOT NULL, owner_epoch INTEGER NOT NULL,
|
||||||
|
boot_id TEXT NOT NULL, reason TEXT NOT NULL)""")
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS terminal_evidence (
|
||||||
|
execution_id TEXT PRIMARY KEY, outcome TEXT NOT NULL)""")
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS terminal_intents (
|
||||||
|
execution_id TEXT PRIMARY KEY, event_json TEXT NOT NULL,
|
||||||
|
digest TEXT NOT NULL, phase TEXT NOT NULL,
|
||||||
|
cleanup_boot_id TEXT NOT NULL)""")
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS pending_continuations (
|
||||||
|
execution_id TEXT PRIMARY KEY, checkpoint_id TEXT NOT NULL,
|
||||||
|
consumed_by TEXT UNIQUE)""")
|
||||||
|
columns = {row[1] for row in db.execute("PRAGMA table_info(pending_continuations)")}
|
||||||
|
for name in ("pending_hash", "decision_hash"):
|
||||||
|
if name not in columns:
|
||||||
|
db.execute(f"ALTER TABLE pending_continuations ADD COLUMN {name} TEXT NOT NULL DEFAULT ''")
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS active_claims (
|
||||||
|
execution_id TEXT PRIMARY KEY, thread_id TEXT NOT NULL,
|
||||||
|
turn_id TEXT NOT NULL, UNIQUE(thread_id, turn_id))""")
|
||||||
|
db.execute("""INSERT OR IGNORE INTO active_claims
|
||||||
|
SELECT execution_id, thread_id, turn_id FROM executions
|
||||||
|
WHERE execution_id NOT IN (SELECT execution_id FROM terminal_evidence)""")
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS checkpoint_scopes (
|
||||||
|
store_id TEXT NOT NULL, checkpoint_thread_id TEXT NOT NULL,
|
||||||
|
checkpoint_ns TEXT NOT NULL, subject_id TEXT NOT NULL,
|
||||||
|
PRIMARY KEY(store_id, checkpoint_thread_id, checkpoint_ns))""")
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS checkpoint_writers (
|
||||||
|
execution_id TEXT PRIMARY KEY, store_id TEXT NOT NULL,
|
||||||
|
checkpoint_thread_id TEXT NOT NULL, checkpoint_ns TEXT NOT NULL,
|
||||||
|
host_id TEXT NOT NULL, boot_id TEXT NOT NULL,
|
||||||
|
UNIQUE(store_id, checkpoint_thread_id, checkpoint_ns))""")
|
||||||
|
db.execute("""CREATE TABLE IF NOT EXISTS execution_checkpoint_scopes (
|
||||||
|
execution_id TEXT PRIMARY KEY, store_id TEXT NOT NULL,
|
||||||
|
checkpoint_thread_id TEXT NOT NULL, checkpoint_ns TEXT NOT NULL)""")
|
||||||
|
|
||||||
|
def claim_checkpoint(self, execution_id: str, *, store_id: str,
|
||||||
|
checkpoint_thread_id: str, subject_id: str) -> None:
|
||||||
|
"""Root Graph writer, including child namespaces; no expiry/takeover."""
|
||||||
|
scope = (store_id, checkpoint_thread_id, "")
|
||||||
|
if not all((execution_id, store_id, checkpoint_thread_id, subject_id)):
|
||||||
|
raise EvoRuntimeError("CHECKPOINT_SCOPE_INVALID")
|
||||||
|
with self._transaction() as db:
|
||||||
|
db.execute("BEGIN IMMEDIATE")
|
||||||
|
owner = db.execute("SELECT subject_id FROM checkpoint_scopes WHERE store_id=? AND checkpoint_thread_id=? AND checkpoint_ns=?", scope).fetchone()
|
||||||
|
if owner and owner[0] != subject_id:
|
||||||
|
raise EvoRuntimeError("CHECKPOINT_SUBJECT_MISMATCH")
|
||||||
|
if db.execute("SELECT 1 FROM checkpoint_writers WHERE store_id=? AND checkpoint_thread_id=? AND checkpoint_ns=?", scope).fetchone():
|
||||||
|
raise EvoRuntimeError("CHECKPOINT_WRITER_BUSY")
|
||||||
|
if db.execute("SELECT 1 FROM active_claims WHERE execution_id NOT IN (SELECT execution_id FROM execution_checkpoint_scopes)").fetchone():
|
||||||
|
raise EvoRuntimeError("CHECKPOINT_LEGACY_WRITER_UNKNOWN")
|
||||||
|
db.execute("INSERT OR IGNORE INTO checkpoint_scopes VALUES (?, ?, ?, ?)", (*scope, subject_id))
|
||||||
|
db.execute("INSERT INTO checkpoint_writers VALUES (?, ?, ?, ?, ?, ?)",
|
||||||
|
(execution_id, *scope, self.host_id, self.boot_id))
|
||||||
|
|
||||||
|
def release_unbound_checkpoint(self, execution_id: str) -> None:
|
||||||
|
"""Only preparation failure, before bind/construction/writes began."""
|
||||||
|
with self._transaction() as db:
|
||||||
|
db.execute("BEGIN IMMEDIATE")
|
||||||
|
if db.execute("SELECT 1 FROM executions WHERE execution_id=?", (execution_id,)).fetchone():
|
||||||
|
raise EvoRuntimeError("CHECKPOINT_WRITER_ALREADY_BOUND")
|
||||||
|
db.execute("DELETE FROM checkpoint_writers WHERE execution_id=? AND host_id=? AND boot_id=?",
|
||||||
|
(execution_id, self.host_id, self.boot_id))
|
||||||
|
|
||||||
|
def _connect(self):
|
||||||
|
db = sqlite3.connect(self.path, timeout=2)
|
||||||
|
db.row_factory = sqlite3.Row
|
||||||
|
db.execute("PRAGMA synchronous=FULL")
|
||||||
|
return db
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _transaction(self):
|
||||||
|
try:
|
||||||
|
with closing(self._connect()) as db, db:
|
||||||
|
yield db
|
||||||
|
except sqlite3.OperationalError as exc:
|
||||||
|
code = getattr(exc, "sqlite_errorcode", 0) & 255
|
||||||
|
if code in {sqlite3.SQLITE_BUSY, sqlite3.SQLITE_LOCKED}:
|
||||||
|
raise EvoRuntimeError("HOST_REGISTRY_BUSY") from exc
|
||||||
|
raise
|
||||||
|
|
||||||
|
def bind(self, *, execution_id: str, grant_id: str, digest: str,
|
||||||
|
thread_id: str, turn_id: str, predecessor_execution_id: str = "",
|
||||||
|
predecessor_checkpoint_id: str = "", predecessor_owner_epoch: int = 0,
|
||||||
|
continuation_pending_hash: str = "", continuation_decision_hash: str = "") -> None:
|
||||||
|
with self._transaction() as db:
|
||||||
|
db.execute("BEGIN IMMEDIATE")
|
||||||
|
if db.execute("SELECT 1 FROM executions WHERE execution_id=? OR grant_id=?",
|
||||||
|
(execution_id, grant_id)).fetchone():
|
||||||
|
raise EvoRuntimeError("EXECUTION_IDENTITY_CONFLICT")
|
||||||
|
writer = db.execute("SELECT * FROM checkpoint_writers WHERE execution_id=?", (execution_id,)).fetchone()
|
||||||
|
if writer and (writer["host_id"] != self.host_id or writer["boot_id"] != self.boot_id):
|
||||||
|
raise EvoRuntimeError("EXECUTION_BOOT_MISMATCH")
|
||||||
|
if writer and predecessor_execution_id:
|
||||||
|
previous = db.execute("SELECT * FROM execution_checkpoint_scopes WHERE execution_id=?", (predecessor_execution_id,)).fetchone()
|
||||||
|
if previous is None or any(previous[k] != writer[k] for k in ("store_id", "checkpoint_thread_id", "checkpoint_ns")):
|
||||||
|
raise EvoRuntimeError("CONTINUATION_CHECKPOINT_SCOPE_MISMATCH")
|
||||||
|
if db.execute("SELECT 1 FROM active_claims WHERE thread_id=? AND turn_id=?",
|
||||||
|
(thread_id, turn_id)).fetchone():
|
||||||
|
raise EvoRuntimeError("TURN_EXECUTION_UNKNOWN")
|
||||||
|
prior = db.execute("SELECT 1 FROM executions WHERE thread_id=? AND turn_id=?",
|
||||||
|
(thread_id, turn_id)).fetchone()
|
||||||
|
if prior or predecessor_execution_id or predecessor_checkpoint_id:
|
||||||
|
if not predecessor_execution_id:
|
||||||
|
raise EvoRuntimeError("CONTINUATION_REQUIRED")
|
||||||
|
pending = db.execute("""SELECT p.*, e.owner_epoch, e.host_id FROM pending_continuations p
|
||||||
|
JOIN executions e USING(execution_id)
|
||||||
|
WHERE p.execution_id=? AND e.thread_id=? AND e.turn_id=?""",
|
||||||
|
(predecessor_execution_id, thread_id, turn_id)).fetchone()
|
||||||
|
if (pending is not None and pending["consumed_by"] is not None
|
||||||
|
and pending["checkpoint_id"] == predecessor_checkpoint_id):
|
||||||
|
failed = db.execute(
|
||||||
|
"SELECT 1 FROM terminal_evidence WHERE execution_id=? AND outcome='failed'",
|
||||||
|
(pending["consumed_by"],),
|
||||||
|
).fetchone()
|
||||||
|
if failed:
|
||||||
|
# Never unconsume a decision on failure. A new grant alone
|
||||||
|
# is insufficient; recovery needs fresh pending authority.
|
||||||
|
raise EvoRuntimeError("CONTINUATION_CONSUMED_FAILURE_REQUIRES_REAUTHORIZATION")
|
||||||
|
if (pending is None or not predecessor_checkpoint_id
|
||||||
|
or pending["checkpoint_id"] != predecessor_checkpoint_id
|
||||||
|
or pending["consumed_by"] is not None):
|
||||||
|
raise EvoRuntimeError("CONTINUATION_INVALID")
|
||||||
|
if (pending["host_id"] != self.host_id or not pending["pending_hash"]
|
||||||
|
or pending["owner_epoch"] != predecessor_owner_epoch
|
||||||
|
or isinstance(predecessor_owner_epoch, bool)
|
||||||
|
or pending["pending_hash"] != continuation_pending_hash
|
||||||
|
or len(continuation_decision_hash) != 64):
|
||||||
|
raise EvoRuntimeError("CONTINUATION_AUTHORIZATION_INVALID")
|
||||||
|
if db.execute("SELECT 1 FROM cancel_intents WHERE execution_id=?",
|
||||||
|
(predecessor_execution_id,)).fetchone():
|
||||||
|
raise EvoRuntimeError("CONTINUATION_CANCELLED")
|
||||||
|
db.execute("UPDATE pending_continuations SET consumed_by=?, decision_hash=? WHERE execution_id=?",
|
||||||
|
(execution_id, continuation_decision_hash, predecessor_execution_id))
|
||||||
|
db.execute("INSERT INTO executions VALUES (?, ?, ?, ?, ?, ?, ?, 1)",
|
||||||
|
(execution_id, grant_id, digest, thread_id, turn_id,
|
||||||
|
self.host_id, self.boot_id))
|
||||||
|
db.execute("INSERT INTO active_claims VALUES (?, ?, ?)",
|
||||||
|
(execution_id, thread_id, turn_id))
|
||||||
|
if writer:
|
||||||
|
db.execute("INSERT INTO execution_checkpoint_scopes VALUES (?, ?, ?, ?)",
|
||||||
|
(execution_id, writer["store_id"], writer["checkpoint_thread_id"], writer["checkpoint_ns"]))
|
||||||
|
|
||||||
|
def finish(self, execution_id: str, *, outcome: str, checkpoint_id: str = "") -> None:
|
||||||
|
"""Trusted resource-owner attestation, not transferable control authority."""
|
||||||
|
self._finish(execution_id, outcome=outcome, checkpoint_id=checkpoint_id)
|
||||||
|
|
||||||
|
def prepare_terminal(self, execution_id: str, *, event: dict) -> None:
|
||||||
|
"""Original resource owner only, AFTER all owned cleanup returns."""
|
||||||
|
body = canonical_json_v1(event).decode()
|
||||||
|
digest = hashlib.sha256(body.encode()).hexdigest()
|
||||||
|
if (event.get("run_id") != execution_id or event.get("kind") != "run"
|
||||||
|
or event.get("payload", {}).get("kind") != "run_terminal"):
|
||||||
|
raise EvoRuntimeError("EXECUTION_TERMINAL_CONFLICT")
|
||||||
|
with self._transaction() as db:
|
||||||
|
db.execute("BEGIN IMMEDIATE")
|
||||||
|
row = db.execute("SELECT * FROM executions WHERE execution_id=?", (execution_id,)).fetchone()
|
||||||
|
if row is None or row['host_id'] != self.host_id or row['boot_id'] != self.boot_id:
|
||||||
|
raise EvoRuntimeError("EXECUTION_BOOT_MISMATCH")
|
||||||
|
prior = db.execute("SELECT digest FROM terminal_intents WHERE execution_id=?", (execution_id,)).fetchone()
|
||||||
|
if prior and prior['digest'] != digest:
|
||||||
|
raise EvoRuntimeError("EXECUTION_TERMINAL_CONFLICT")
|
||||||
|
db.execute("INSERT OR IGNORE INTO terminal_intents VALUES (?, ?, ?, 'prepared', ?)",
|
||||||
|
(execution_id, body, digest, self.boot_id))
|
||||||
|
|
||||||
|
def terminal_intent(self, execution_id: str) -> dict | None:
|
||||||
|
with closing(self._connect()) as db:
|
||||||
|
row = db.execute("""SELECT i.* FROM terminal_intents i JOIN executions e USING(execution_id)
|
||||||
|
WHERE execution_id=? AND e.host_id=?""", (execution_id, self.host_id)).fetchone()
|
||||||
|
if row is None:
|
||||||
|
return None
|
||||||
|
return {**dict(row), 'event': json.loads(row['event_json']), 'cleanup_confirmed': True}
|
||||||
|
|
||||||
|
def confirm_terminal(self, execution_id: str, *, digest: str) -> None:
|
||||||
|
with self._transaction() as db:
|
||||||
|
row = db.execute("""SELECT i.* FROM terminal_intents i JOIN executions e USING(execution_id)
|
||||||
|
WHERE execution_id=? AND e.host_id=?""", (execution_id, self.host_id)).fetchone()
|
||||||
|
if row is None or row['digest'] != digest:
|
||||||
|
raise EvoRuntimeError("EXECUTION_TERMINAL_CONFLICT")
|
||||||
|
db.execute("UPDATE terminal_intents SET phase='sink_confirmed' WHERE execution_id=? AND phase='prepared'",
|
||||||
|
(execution_id,))
|
||||||
|
|
||||||
|
def _finish(self, execution_id: str, *, outcome: str, checkpoint_id: str) -> None:
|
||||||
|
if outcome not in {"completed", "cancelled", "failed", "awaiting_input"}:
|
||||||
|
raise EvoRuntimeError("EXECUTION_OUTCOME_INVALID")
|
||||||
|
with self._transaction() as db:
|
||||||
|
db.execute("BEGIN IMMEDIATE")
|
||||||
|
row = db.execute("SELECT * FROM executions WHERE execution_id=?",
|
||||||
|
(execution_id,)).fetchone()
|
||||||
|
intent = db.execute("SELECT * FROM terminal_intents WHERE execution_id=?", (execution_id,)).fetchone()
|
||||||
|
if row is None or row["host_id"] != self.host_id or (row["boot_id"] != self.boot_id and
|
||||||
|
(intent is None or intent['phase'] not in {'sink_confirmed', 'registry_finished'})):
|
||||||
|
raise EvoRuntimeError("EXECUTION_BOOT_MISMATCH")
|
||||||
|
if intent is not None:
|
||||||
|
payload = json.loads(intent['event_json'])['payload']
|
||||||
|
if (intent['phase'] not in {'sink_confirmed', 'registry_finished'}
|
||||||
|
or payload['outcome'] != outcome
|
||||||
|
or str(payload.get('checkpoint_id') or '') != checkpoint_id):
|
||||||
|
raise EvoRuntimeError("EXECUTION_TERMINAL_CONFLICT")
|
||||||
|
prior = db.execute("SELECT outcome FROM terminal_evidence WHERE execution_id=?",
|
||||||
|
(execution_id,)).fetchone()
|
||||||
|
if prior is not None and prior["outcome"] != outcome:
|
||||||
|
raise EvoRuntimeError("EXECUTION_TERMINAL_CONFLICT")
|
||||||
|
db.execute("INSERT OR IGNORE INTO terminal_evidence VALUES (?, ?)",
|
||||||
|
(execution_id, outcome))
|
||||||
|
if outcome == "awaiting_input" and checkpoint_id:
|
||||||
|
pending = db.execute("SELECT checkpoint_id FROM pending_continuations WHERE execution_id=?",
|
||||||
|
(execution_id,)).fetchone()
|
||||||
|
if pending is not None and pending["checkpoint_id"] != checkpoint_id:
|
||||||
|
raise EvoRuntimeError("EXECUTION_TERMINAL_CONFLICT")
|
||||||
|
identity = {}
|
||||||
|
if intent is not None:
|
||||||
|
payload = json.loads(intent['event_json'])['payload']
|
||||||
|
identity = {k: payload.get(k) for k in (
|
||||||
|
"checkpoint_thread_id", "checkpoint_id", "checkpoint_ns", "pending_interrupts")}
|
||||||
|
pending_hash = hashlib.sha256(canonical_json_v1(identity)).hexdigest() if identity else ""
|
||||||
|
db.execute("INSERT OR IGNORE INTO pending_continuations "
|
||||||
|
"(execution_id, checkpoint_id, consumed_by, pending_hash, decision_hash) VALUES (?, ?, NULL, ?, '')",
|
||||||
|
(execution_id, checkpoint_id, pending_hash))
|
||||||
|
db.execute("DELETE FROM active_claims WHERE execution_id=?", (execution_id,))
|
||||||
|
db.execute("DELETE FROM checkpoint_writers WHERE execution_id=?", (execution_id,))
|
||||||
|
db.execute("UPDATE terminal_intents SET phase='registry_finished' WHERE execution_id=?", (execution_id,))
|
||||||
|
|
||||||
|
def continuation(self, execution_id: str) -> dict:
|
||||||
|
with closing(self._connect()) as db:
|
||||||
|
row = db.execute("""SELECT p.*, e.owner_epoch FROM pending_continuations p
|
||||||
|
JOIN executions e USING(execution_id) WHERE execution_id=? AND e.host_id=?""",
|
||||||
|
(execution_id, self.host_id)).fetchone()
|
||||||
|
if row is None:
|
||||||
|
raise EvoRuntimeError("CONTINUATION_INVALID")
|
||||||
|
return dict(row)
|
||||||
|
|
||||||
|
def lookup_grant(self, grant_id: str, digest: str) -> dict | None:
|
||||||
|
with closing(self._connect()) as db:
|
||||||
|
row = db.execute("SELECT * FROM executions WHERE grant_id=?", (grant_id,)).fetchone()
|
||||||
|
if row is None:
|
||||||
|
return None
|
||||||
|
if row["digest"] != digest:
|
||||||
|
raise EvoRuntimeError("CONTRACT_REPLAYED")
|
||||||
|
return dict(row)
|
||||||
|
|
||||||
|
def transfer_control(self, execution_id: str, *, expected_epoch: int, new_epoch: int) -> int:
|
||||||
|
if new_epoch <= expected_epoch:
|
||||||
|
raise EvoRuntimeError("OWNER_EPOCH_STALE")
|
||||||
|
with self._transaction() as db:
|
||||||
|
changed = db.execute(
|
||||||
|
"UPDATE executions SET owner_epoch=? WHERE execution_id=? AND owner_epoch=? AND host_id=? AND boot_id=?",
|
||||||
|
(new_epoch, execution_id, expected_epoch, self.host_id, self.boot_id),
|
||||||
|
).rowcount
|
||||||
|
if changed != 1:
|
||||||
|
raise EvoRuntimeError("OWNER_EPOCH_STALE")
|
||||||
|
return new_epoch
|
||||||
|
|
||||||
|
def require_control(self, execution_id: str, *, owner_epoch: int) -> None:
|
||||||
|
with closing(self._connect()) as db:
|
||||||
|
row = db.execute("SELECT owner_epoch FROM executions WHERE execution_id=?",
|
||||||
|
(execution_id,)).fetchone()
|
||||||
|
if row is None or row["owner_epoch"] != owner_epoch:
|
||||||
|
raise EvoRuntimeError("OWNER_EPOCH_STALE")
|
||||||
|
|
||||||
|
def inspect(self, execution_id: str) -> dict:
|
||||||
|
with closing(self._connect()) as db:
|
||||||
|
row = db.execute("SELECT * FROM executions WHERE execution_id=?", (execution_id,)).fetchone()
|
||||||
|
terminal = db.execute("SELECT outcome FROM terminal_evidence WHERE execution_id=?",
|
||||||
|
(execution_id,)).fetchone()
|
||||||
|
cancel = db.execute("SELECT 1 FROM cancel_intents WHERE execution_id=? LIMIT 1",
|
||||||
|
(execution_id,)).fetchone()
|
||||||
|
return {**(dict(row) if row else {"execution_id": execution_id}),
|
||||||
|
"cancel_requested": cancel is not None,
|
||||||
|
"recovery_action": "inspect_only",
|
||||||
|
"status": terminal["outcome"] if terminal else "unknown",
|
||||||
|
"resources_confirmed_exited": terminal is not None,
|
||||||
|
"source": "host_binding_only" if row else "no_host_binding"}
|
||||||
|
|
||||||
|
def _control_row(self, db, execution_id, owner_epoch, boot_id):
|
||||||
|
if owner_epoch is None:
|
||||||
|
raise EvoRuntimeError("OWNER_EPOCH_REQUIRED")
|
||||||
|
row = db.execute("SELECT * FROM executions WHERE execution_id=?", (execution_id,)).fetchone()
|
||||||
|
if row is None:
|
||||||
|
raise EvoRuntimeError("EXECUTION_UNKNOWN")
|
||||||
|
if row["host_id"] != self.host_id or row["boot_id"] != self.boot_id or boot_id != self.boot_id:
|
||||||
|
raise EvoRuntimeError("EXECUTION_BOOT_MISMATCH")
|
||||||
|
if isinstance(owner_epoch, bool) or row["owner_epoch"] != owner_epoch:
|
||||||
|
raise EvoRuntimeError("OWNER_EPOCH_STALE")
|
||||||
|
return row
|
||||||
|
|
||||||
|
def inspect_control(self, execution_id: str, *, owner_epoch: int | None, boot_id: str | None) -> dict:
|
||||||
|
with closing(self._connect()) as db:
|
||||||
|
row = self._control_row(db, execution_id, owner_epoch, boot_id)
|
||||||
|
return self.inspect(execution_id)
|
||||||
|
|
||||||
|
def accept_cancel(self, execution_id: str, *, owner_epoch: int | None, boot_id: str | None, reason: str) -> int:
|
||||||
|
# COMMIT is the linearization point shared with transfer's conditional UPDATE.
|
||||||
|
# Accepted commands survive a later transfer; no write lock crosses an await.
|
||||||
|
with self._transaction() as db:
|
||||||
|
db.execute("BEGIN IMMEDIATE")
|
||||||
|
self._control_row(db, execution_id, owner_epoch, boot_id)
|
||||||
|
cursor = db.execute(
|
||||||
|
"INSERT INTO cancel_intents (execution_id, owner_epoch, boot_id, reason) VALUES (?, ?, ?, ?)",
|
||||||
|
(execution_id, owner_epoch, boot_id, reason),
|
||||||
|
)
|
||||||
|
assert cursor.lastrowid is not None
|
||||||
|
return cursor.lastrowid
|
||||||
@@ -83,8 +83,9 @@ def compile_invocation_plan(
|
|||||||
"""Validate and freeze adapter output before constructing a provider SDK."""
|
"""Validate and freeze adapter output before constructing a provider SDK."""
|
||||||
|
|
||||||
params = dict(sdk_params)
|
params = dict(sdk_params)
|
||||||
streaming = purpose == "main_agent"
|
streaming = True
|
||||||
params["streaming"] = streaming
|
params["streaming"] = streaming
|
||||||
|
params["disable_streaming"] = False
|
||||||
if runtime_provider == "openai" and streaming:
|
if runtime_provider == "openai" and streaming:
|
||||||
params["stream_usage"] = True
|
params["stream_usage"] = True
|
||||||
token_fields = tuple(
|
token_fields = tuple(
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ import stat
|
|||||||
import tempfile
|
import tempfile
|
||||||
import time
|
import time
|
||||||
from collections.abc import Mapping, Sequence
|
from collections.abc import Mapping, Sequence
|
||||||
from dataclasses import asdict, dataclass, field
|
from dataclasses import asdict, dataclass, field, replace
|
||||||
from decimal import Decimal, InvalidOperation
|
from decimal import Decimal, InvalidOperation
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -477,12 +477,49 @@ class EvoModelConfig:
|
|||||||
pools = _parse_endpoint_pools(raw.get("endpoint_pools"), providers)
|
pools = _parse_endpoint_pools(raw.get("endpoint_pools"), providers)
|
||||||
health = _parse_health(raw.get("route_health"))
|
health = _parse_health(raw.get("route_health"))
|
||||||
selectors = _parse_selectors(raw.get("route_selectors"), providers, pools)
|
selectors = _parse_selectors(raw.get("route_selectors"), providers, pools)
|
||||||
|
# Schema2's historical non_streaming Web tool routes and native routes
|
||||||
|
# are administrator declarations, not probes or HTTP streaming policy.
|
||||||
|
declared_tool_models = {
|
||||||
|
(selector.provider, selector.model)
|
||||||
|
for selector in selectors.values()
|
||||||
|
if selector.tool_call_transport in {"native", "non_streaming"}
|
||||||
|
}
|
||||||
|
providers = {
|
||||||
|
provider_id: replace(
|
||||||
|
provider,
|
||||||
|
models={
|
||||||
|
model_id: replace(
|
||||||
|
model,
|
||||||
|
capabilities={
|
||||||
|
**model.capabilities,
|
||||||
|
"tools": (provider_id, model_id) in declared_tool_models,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
for model_id, model in provider.models.items()
|
||||||
|
},
|
||||||
|
)
|
||||||
|
for provider_id, provider in providers.items()
|
||||||
|
}
|
||||||
main_routes, title_selector_id = _parse_purpose_routes(
|
main_routes, title_selector_id = _parse_purpose_routes(
|
||||||
raw.get("purpose_routes"), selectors
|
raw.get("purpose_routes"), selectors
|
||||||
)
|
)
|
||||||
limits = _parse_call_limits(raw.get("purpose_call_limits"))
|
limits = _parse_call_limits(raw.get("purpose_call_limits"))
|
||||||
web_runtime = _parse_web_runtime(raw.get("web_runtime"))
|
web_runtime = _parse_web_runtime(raw.get("web_runtime"))
|
||||||
fallbacks = _parse_fallbacks(raw.get("tool_protocol_fallbacks"), selectors)
|
fallbacks = _parse_fallbacks(raw.get("tool_protocol_fallbacks"), selectors)
|
||||||
|
execution_selectors = set(main_routes.selectable.values()) | {
|
||||||
|
title_selector_id
|
||||||
|
}
|
||||||
|
for primary, candidates in fallbacks.items():
|
||||||
|
execution_selectors.add(primary)
|
||||||
|
execution_selectors.update(candidates)
|
||||||
|
for selector_id in sorted(execution_selectors):
|
||||||
|
selector = selectors[selector_id]
|
||||||
|
if (selector.provider, selector.model) not in declared_tool_models:
|
||||||
|
raise EvoRuntimeError(
|
||||||
|
"LLM_ROUTE_CONFIGURATION_REQUIRED",
|
||||||
|
"schema2 tool capability migration requires an explicit "
|
||||||
|
f"native or historical non_streaming declaration for {selector_id}",
|
||||||
|
)
|
||||||
config = cls(
|
config = cls(
|
||||||
config_revision=revision,
|
config_revision=revision,
|
||||||
config_identity_key_id=identity_key_id,
|
config_identity_key_id=identity_key_id,
|
||||||
@@ -1385,7 +1422,7 @@ def _parse_v3_config(
|
|||||||
allowed=_REASONING_EFFORTS - {"disabled"},
|
allowed=_REASONING_EFFORTS - {"disabled"},
|
||||||
)
|
)
|
||||||
default_reasoning_effort = str(
|
default_reasoning_effort = str(
|
||||||
reasoning_policy.get("default_effort") or "high"
|
reasoning_policy.get("default_effort") or "medium"
|
||||||
)
|
)
|
||||||
if default_reasoning_effort not in allowed_reasoning_efforts:
|
if default_reasoning_effort not in allowed_reasoning_efforts:
|
||||||
raise EvoRuntimeError(
|
raise EvoRuntimeError(
|
||||||
@@ -1973,6 +2010,7 @@ def _parse_providers(value: Any) -> Mapping[str, ProviderConfig]:
|
|||||||
_string_set(access.get("allowed_plans"), "allowed_plans"),
|
_string_set(access.get("allowed_plans"), "allowed_plans"),
|
||||||
_string_set(access.get("allowed_roles"), "allowed_roles"),
|
_string_set(access.get("allowed_roles"), "allowed_roles"),
|
||||||
quote,
|
quote,
|
||||||
|
capabilities={"tools": False, "text": True},
|
||||||
)
|
)
|
||||||
result[key] = ProviderConfig(
|
result[key] = ProviderConfig(
|
||||||
key,
|
key,
|
||||||
@@ -2137,10 +2175,12 @@ def _parse_selectors(
|
|||||||
transport = _text(
|
transport = _text(
|
||||||
item.get("tool_call_transport"), "selector.tool_call_transport"
|
item.get("tool_call_transport"), "selector.tool_call_transport"
|
||||||
)
|
)
|
||||||
if transport != "non_streaming":
|
# Legacy values remain readable for signed route evidence, but do not
|
||||||
|
# control HTTP streaming. Tool support and HTTP mode are independent.
|
||||||
|
if transport not in {"native", "streaming", "non_streaming"}:
|
||||||
raise EvoRuntimeError(
|
raise EvoRuntimeError(
|
||||||
"LLM_ROUTE_CONFIGURATION_REQUIRED",
|
"LLM_ROUTE_CONFIGURATION_REQUIRED",
|
||||||
"Web tool routes must be non_streaming",
|
"Web tool routes require native tool support",
|
||||||
)
|
)
|
||||||
_validate_params(
|
_validate_params(
|
||||||
provider.params, name=f"providers.{provider_id}.params", allowed=allowed
|
provider.params, name=f"providers.{provider_id}.params", allowed=allowed
|
||||||
@@ -3180,14 +3220,16 @@ def convert_v2_to_v3_draft(
|
|||||||
"api_mode": selector.api_mode
|
"api_mode": selector.api_mode
|
||||||
if selector
|
if selector
|
||||||
else "chat_completions",
|
else "chat_completions",
|
||||||
"tool_call_transport": "native",
|
"tool_call_transport": "native"
|
||||||
|
if model.capabilities.get("tools", False)
|
||||||
|
else "disabled",
|
||||||
},
|
},
|
||||||
"capabilities": {
|
"capabilities": {
|
||||||
"text": True,
|
"text": True,
|
||||||
"vision": model.supports_vision,
|
"vision": model.supports_vision,
|
||||||
"video": False,
|
"video": False,
|
||||||
"documents": False,
|
"documents": False,
|
||||||
"tools": True,
|
"tools": bool(model.capabilities.get("tools", False)),
|
||||||
"structured_output": False,
|
"structured_output": False,
|
||||||
"thinking": model.supports_reasoning,
|
"thinking": model.supports_reasoning,
|
||||||
},
|
},
|
||||||
|
|||||||
+259
-318
@@ -1,10 +1,11 @@
|
|||||||
"""LLM model configuration based on LangChain init_chat_model.
|
"""LLM model configuration based on LangChain init_chat_model.
|
||||||
|
|
||||||
This module provides a unified interface for creating chat model instances
|
This module provides a unified interface for creating chat model instances
|
||||||
with support for multiple providers (Anthropic, OpenAI, Google GenAI, MiniMax
|
with support for multiple providers (Anthropic, OpenAI, Google GenAI, Atlas
|
||||||
(Anthropic-compatible), NVIDIA, SiliconFlow, OpenRouter, ZhipuAI, Volcengine,
|
Cloud, MiniMax (Anthropic-compatible), NVIDIA, SiliconFlow, OpenRouter, Requesty,
|
||||||
DashScope, DashScope-Code, DeepSeek, Ollama, and custom OpenAI/Anthropic-compatible
|
Novita, ZhipuAI, Volcengine, DashScope, DashScope-Code, DeepSeek, Ollama, and
|
||||||
endpoints) and convenient short names for common models.
|
custom OpenAI/Anthropic-compatible endpoints) and convenient short names for
|
||||||
|
common models.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -15,6 +16,7 @@ import subprocess
|
|||||||
import warnings
|
import warnings
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
from langchain.chat_models import init_chat_model
|
from langchain.chat_models import init_chat_model
|
||||||
|
|
||||||
@@ -24,26 +26,31 @@ from ..config.settings import (
|
|||||||
OPENROUTER_DEFAULT_HTTP_REFERER,
|
OPENROUTER_DEFAULT_HTTP_REFERER,
|
||||||
)
|
)
|
||||||
from .context_window import apply_known_context_window
|
from .context_window import apply_known_context_window
|
||||||
|
from .deepseek import EvoChatDeepSeek
|
||||||
from .patches import (
|
from .patches import (
|
||||||
_is_ccproxy_codex,
|
_is_ccproxy_codex,
|
||||||
|
_patch_anthropic_strip_foreign_reasoning,
|
||||||
|
_patch_anthropic_structured_output,
|
||||||
_patch_ccproxy_system_to_developer,
|
_patch_ccproxy_system_to_developer,
|
||||||
_patch_deepseek_reasoning_passback,
|
|
||||||
_patch_openai_compat_content,
|
_patch_openai_compat_content,
|
||||||
_patch_openrouter_strip_responses_reasoning,
|
_patch_openrouter_strip_responses_reasoning,
|
||||||
|
_patch_openrouter_structured_output,
|
||||||
|
)
|
||||||
|
from .registry import (
|
||||||
|
_ANTHROPIC_ROUTED_PROVIDERS,
|
||||||
|
_MODEL_ENTRIES,
|
||||||
|
_OPENAI_ROUTED_PROVIDERS,
|
||||||
|
_OPENROUTER_JSON_SCHEMA_STRUCTURED_OUTPUT_MODELS, # noqa: F401 — re-exported
|
||||||
|
_THINKING_CAPABLE_PROVIDERS,
|
||||||
|
DEFAULT_MODEL,
|
||||||
|
MODELS,
|
||||||
|
_is_mandatory_thinking_kimi,
|
||||||
|
get_model_info, # noqa: F401 — re-exported for existing import sites
|
||||||
|
get_models_for_provider, # noqa: F401 — re-exported for existing import sites
|
||||||
|
list_model_picker_entries, # noqa: F401 — re-exported for existing import sites
|
||||||
|
list_models, # noqa: F401 — re-exported for existing import sites
|
||||||
|
list_models_by_provider, # noqa: F401 — re-exported for existing import sites
|
||||||
)
|
)
|
||||||
|
|
||||||
_MINIMAX_ANTHROPIC_BASE_URL = "https://api.minimaxi.com/anthropic"
|
|
||||||
_SILICONFLOW_BASE_URL = "https://api.siliconflow.cn/v1"
|
|
||||||
|
|
||||||
_ZHIPU_BASE_URL = "https://open.bigmodel.cn/api/paas/v4"
|
|
||||||
_ZHIPU_CODE_BASE_URL = "https://open.bigmodel.cn/api/coding/paas/v4"
|
|
||||||
_VOLCENGINE_BASE_URL = "https://ark.cn-beijing.volces.com/api/v3"
|
|
||||||
_DASHSCOPE_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
|
||||||
_DASHSCOPE_CODE_BASE_URL = "https://coding.dashscope.aliyuncs.com/v1"
|
|
||||||
|
|
||||||
_DEEPSEEK_BASE_URL = "https://api.deepseek.com"
|
|
||||||
_MOONSHOT_BASE_URL = "https://api.moonshot.cn/v1"
|
|
||||||
_KIMI_CODING_BASE_URL = "https://api.kimi.com/coding/"
|
|
||||||
|
|
||||||
# Minimum Codex CLI version advertised when no explicit override is set. Newer
|
# Minimum Codex CLI version advertised when no explicit override is set. Newer
|
||||||
# installed versions are advertised automatically.
|
# installed versions are advertised automatically.
|
||||||
@@ -84,33 +91,79 @@ def _resolve_codex_client_version() -> str:
|
|||||||
return _CODEX_CLIENT_VERSION_FALLBACK
|
return _CODEX_CLIENT_VERSION_FALLBACK
|
||||||
|
|
||||||
|
|
||||||
# Providers routed through the OpenAI provider with a custom base_url.
|
|
||||||
# Maps provider name → (base_url or None, env var for API key).
|
|
||||||
_OPENAI_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = {
|
|
||||||
"deepseek": (_DEEPSEEK_BASE_URL, "DEEPSEEK_API_KEY"),
|
|
||||||
"moonshot": (_MOONSHOT_BASE_URL, "MOONSHOT_API_KEY"),
|
|
||||||
"siliconflow": (_SILICONFLOW_BASE_URL, "SILICONFLOW_API_KEY"),
|
|
||||||
"zhipu": (_ZHIPU_BASE_URL, "ZHIPU_API_KEY"),
|
|
||||||
"zhipu-code": (_ZHIPU_CODE_BASE_URL, "ZHIPU_API_KEY"),
|
|
||||||
"volcengine": (_VOLCENGINE_BASE_URL, "VOLCENGINE_API_KEY"),
|
|
||||||
"dashscope": (_DASHSCOPE_BASE_URL, "DASHSCOPE_API_KEY"),
|
|
||||||
"dashscope-code": (_DASHSCOPE_CODE_BASE_URL, "DASHSCOPE_API_KEY"),
|
|
||||||
"custom-openai": (
|
|
||||||
None,
|
|
||||||
"CUSTOM_OPENAI_API_KEY",
|
|
||||||
), # base_url from CUSTOM_OPENAI_BASE_URL env
|
|
||||||
}
|
|
||||||
|
|
||||||
# Providers routed through the Anthropic provider with a custom base_url.
|
|
||||||
# Maps provider name → (base_url or None, env var for API key).
|
|
||||||
_ANTHROPIC_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = {
|
|
||||||
"minimax": (_MINIMAX_ANTHROPIC_BASE_URL, "MINIMAX_API_KEY"),
|
|
||||||
"kimi-coding": (_KIMI_CODING_BASE_URL, "KIMI_API_KEY"),
|
|
||||||
"custom-anthropic": (None, "CUSTOM_ANTHROPIC_API_KEY"),
|
|
||||||
}
|
|
||||||
|
|
||||||
# Anthropic-routed providers that support extended thinking.
|
def _resolve_reasoning_effort(default: str) -> str:
|
||||||
_THINKING_CAPABLE_PROVIDERS: set[str] = {"minimax"}
|
"""Return the configured reasoning effort or a provider-specific default."""
|
||||||
|
return os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or default
|
||||||
|
|
||||||
|
|
||||||
|
# Qwen 3.8 Max canonical levels and documented OpenAI alias mappings:
|
||||||
|
# https://docs.qwencloud.com/api-reference/chat/openai-chat#reasoning-effort
|
||||||
|
_DASHSCOPE_QWEN38_REASONING_EFFORTS = frozenset(
|
||||||
|
{"none", "minimal", "low", "medium", "high", "xhigh", "max"}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_dashscope_reasoning_effort(
|
||||||
|
provider: str,
|
||||||
|
model_id: str,
|
||||||
|
effort: str,
|
||||||
|
) -> None:
|
||||||
|
"""Reject reasoning levels unsupported by DashScope Qwen 3.8 Max."""
|
||||||
|
if effort not in _DASHSCOPE_QWEN38_REASONING_EFFORTS:
|
||||||
|
choices = ", ".join(sorted(_DASHSCOPE_QWEN38_REASONING_EFFORTS))
|
||||||
|
raise ValueError(
|
||||||
|
f"Unsupported EVOSCIENTIST_REASONING_EFFORT={effort!r} for "
|
||||||
|
f"{provider} model {model_id!r}. Supported values: {choices}."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_openai_compat_reasoning_config(
|
||||||
|
provider: str,
|
||||||
|
model_id: str,
|
||||||
|
kwargs: dict[str, Any],
|
||||||
|
) -> None:
|
||||||
|
"""Apply reasoning controls supported by OpenAI-compatible providers.
|
||||||
|
|
||||||
|
Routed providers deliberately skip the native-OpenAI branch in
|
||||||
|
:func:`_apply_auto_config`, because most compatible endpoints reject
|
||||||
|
OpenAI-only ``reasoning`` payloads. A small subset does support the
|
||||||
|
standard ``reasoning_effort`` field, though:
|
||||||
|
|
||||||
|
* DashScope Qwen 3.8 Max supports ``low`` / ``medium`` / ``xhigh`` and
|
||||||
|
maps the OpenAI aliases (including ``none``). Its server default is
|
||||||
|
extremely large, so use the standard ``medium`` level unless the user
|
||||||
|
selected another level.
|
||||||
|
* ``custom-openai`` is user-owned. Forward an *explicit* setting only;
|
||||||
|
with no setting, preserve compatibility with endpoints that reject the
|
||||||
|
field (including many non-reasoning OpenAI-compatible APIs).
|
||||||
|
|
||||||
|
Explicit caller kwargs always win.
|
||||||
|
"""
|
||||||
|
configured = os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip()
|
||||||
|
short_model_id = model_id.rsplit("/", 1)[-1]
|
||||||
|
|
||||||
|
if provider == "dashscope" and short_model_id.startswith("qwen3.8-max"):
|
||||||
|
if "reasoning_effort" not in kwargs:
|
||||||
|
effort = configured or "medium"
|
||||||
|
_validate_dashscope_reasoning_effort(provider, model_id, effort)
|
||||||
|
kwargs["reasoning_effort"] = effort
|
||||||
|
return
|
||||||
|
|
||||||
|
if provider == "custom-openai" and configured:
|
||||||
|
kwargs.setdefault("reasoning_effort", configured)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_deepseek_endpoint(base_url: str | None) -> bool:
|
||||||
|
"""Return whether an OpenAI-compatible endpoint is DeepSeek's API."""
|
||||||
|
if not base_url:
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
return urlparse(base_url).hostname == "api.deepseek.com"
|
||||||
|
except ValueError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
_TRUTHY_ENV_VALUES = {"1", "true", "yes", "on"}
|
_TRUTHY_ENV_VALUES = {"1", "true", "yes", "on"}
|
||||||
_FALSEY_ENV_VALUES = {"0", "false", "no", "off"}
|
_FALSEY_ENV_VALUES = {"0", "false", "no", "off"}
|
||||||
@@ -128,192 +181,6 @@ _OPENROUTER_MAX_CATEGORIES_PER_REQUEST = 2
|
|||||||
# LangChain move them into model_kwargs and can later leak them into SDK calls.
|
# LangChain move them into model_kwargs and can later leak them into SDK calls.
|
||||||
_UNSUPPORTED_CHAT_MODEL_KWARGS = frozenset({"sanitize_openai_sdk_headers"})
|
_UNSUPPORTED_CHAT_MODEL_KWARGS = frozenset({"sanitize_openai_sdk_headers"})
|
||||||
|
|
||||||
# Model registry: list of (short_name, model_id, provider)
|
|
||||||
# Allows same short_name across different providers.
|
|
||||||
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
|
|
||||||
# Custom Anthropic (third-party Claude-compatible endpoints, current-gen defaults)
|
|
||||||
# Listed BEFORE native anthropic so MODELS dict defaults to native provider
|
|
||||||
("claude-sonnet-4-6", "claude-sonnet-4-6", "custom-anthropic"),
|
|
||||||
("claude-haiku-4-5", "claude-haiku-4-5", "custom-anthropic"),
|
|
||||||
# Custom OpenAI (third-party OpenAI-compatible endpoints, 3 defaults)
|
|
||||||
# Listed BEFORE native openai so MODELS dict defaults to native provider
|
|
||||||
("gpt-5.5-pro", "gpt-5.5-pro", "custom-openai"),
|
|
||||||
("gpt-5.5", "gpt-5.5", "custom-openai"),
|
|
||||||
("gpt-5.4", "gpt-5.4", "custom-openai"),
|
|
||||||
("gpt-5.3-codex", "gpt-5.3-codex", "custom-openai"),
|
|
||||||
("gpt-5-mini", "gpt-5-mini", "custom-openai"),
|
|
||||||
# Anthropic (current generation)
|
|
||||||
("claude-fable-5", "claude-fable-5", "anthropic"),
|
|
||||||
("claude-opus-4-8", "claude-opus-4-8", "anthropic"),
|
|
||||||
("claude-sonnet-5", "claude-sonnet-5", "anthropic"),
|
|
||||||
("claude-sonnet-4-6", "claude-sonnet-4-6", "anthropic"),
|
|
||||||
("claude-haiku-4-5", "claude-haiku-4-5", "anthropic"),
|
|
||||||
# OpenAI
|
|
||||||
("gpt-5.6-sol", "gpt-5.6-sol", "openai"),
|
|
||||||
("gpt-5.6-terra", "gpt-5.6-terra", "openai"),
|
|
||||||
("gpt-5.6-luna", "gpt-5.6-luna", "openai"),
|
|
||||||
("gpt-5.5-pro", "gpt-5.5-pro", "openai"),
|
|
||||||
("gpt-5.5", "gpt-5.5", "openai"),
|
|
||||||
("gpt-5.4", "gpt-5.4", "openai"),
|
|
||||||
("gpt-5.4-mini", "gpt-5.4-mini", "openai"),
|
|
||||||
("gpt-5.4-nano", "gpt-5.4-nano", "openai"),
|
|
||||||
("gpt-5.3-codex", "gpt-5.3-codex", "openai"),
|
|
||||||
("gpt-5.2-codex", "gpt-5.2-codex", "openai"),
|
|
||||||
("gpt-5.2", "gpt-5.2", "openai"),
|
|
||||||
("gpt-5.1", "gpt-5.1", "openai"),
|
|
||||||
("gpt-5", "gpt-5", "openai"),
|
|
||||||
("gpt-5-mini", "gpt-5-mini", "openai"),
|
|
||||||
("gpt-5-nano", "gpt-5-nano", "openai"),
|
|
||||||
# Google GenAI
|
|
||||||
("gemini-3.5-flash", "gemini-3.5-flash", "google-genai"),
|
|
||||||
("gemini-3.1-pro", "gemini-3.1-pro-preview", "google-genai"),
|
|
||||||
(
|
|
||||||
"gemini-3.1-pro-customtools",
|
|
||||||
"gemini-3.1-pro-preview-customtools",
|
|
||||||
"google-genai",
|
|
||||||
),
|
|
||||||
("gemini-3.1-flash-lite", "gemini-3.1-flash-lite-preview", "google-genai"),
|
|
||||||
("gemini-3-flash", "gemini-3-flash-preview", "google-genai"),
|
|
||||||
("gemini-2.5-flash", "gemini-2.5-flash", "google-genai"),
|
|
||||||
("gemini-2.5-flash-lite", "gemini-2.5-flash-lite", "google-genai"),
|
|
||||||
("gemini-2.5-pro", "gemini-2.5-pro", "google-genai"),
|
|
||||||
# MiniMax (direct API — Anthropic-compatible; default: api.minimaxi.com, global: api.minimax.io)
|
|
||||||
("minimax-m3", "MiniMax-M3", "minimax"),
|
|
||||||
("minimax-m2.7", "MiniMax-M2.7", "minimax"),
|
|
||||||
("minimax-m2.7-highspeed", "MiniMax-M2.7-highspeed", "minimax"),
|
|
||||||
("minimax-m2.5", "MiniMax-M2.5", "minimax"),
|
|
||||||
("minimax-m2.5-highspeed", "MiniMax-M2.5-highspeed", "minimax"),
|
|
||||||
# NVIDIA
|
|
||||||
("nemotron-super", "nvidia/nemotron-3-super-120b-a12b", "nvidia"),
|
|
||||||
("nemotron-nano", "nvidia/nemotron-3-nano-30b-a3b", "nvidia"),
|
|
||||||
("glm-5.2", "z-ai/glm-5.2", "nvidia"),
|
|
||||||
("glm4.7", "z-ai/glm4.7", "nvidia"),
|
|
||||||
("deepseek-v3.2", "deepseek-ai/deepseek-v3.2", "nvidia"),
|
|
||||||
("deepseek-v3.1", "deepseek-ai/deepseek-v3.1-terminus", "nvidia"),
|
|
||||||
("kimi-k2.5", "moonshotai/kimi-k2.5", "nvidia"),
|
|
||||||
("kimi-k2-thinking", "moonshotai/kimi-k2-thinking", "nvidia"),
|
|
||||||
("minimax-m2.5", "minimaxai/minimax-m2.5", "nvidia"),
|
|
||||||
("minimax-m2.1", "minimaxai/minimax-m2.1", "nvidia"),
|
|
||||||
("qwen3.5-397b", "qwen/qwen3.5-397b-a17b", "nvidia"),
|
|
||||||
("step-3.5-flash", "stepfun-ai/step-3.5-flash", "nvidia"),
|
|
||||||
# SiliconFlow
|
|
||||||
("minimax-m2.5", "Pro/MiniMaxAI/MiniMax-M2.5", "siliconflow"),
|
|
||||||
("glm-5.2", "Pro/zai-org/GLM-5.2", "siliconflow"),
|
|
||||||
("glm-5", "Pro/zai-org/GLM-5", "siliconflow"),
|
|
||||||
("kimi-k2.5", "Pro/moonshotai/Kimi-K2.5", "siliconflow"),
|
|
||||||
("glm-4.7", "Pro/zai-org/GLM-4.7", "siliconflow"),
|
|
||||||
# OpenRouter
|
|
||||||
("claude-fable-5", "anthropic/claude-fable-5", "openrouter"),
|
|
||||||
("claude-opus-4.8", "anthropic/claude-opus-4.8", "openrouter"),
|
|
||||||
("claude-opus-4.8-fast", "anthropic/claude-opus-4.8-fast", "openrouter"),
|
|
||||||
("claude-sonnet-5", "anthropic/claude-sonnet-5", "openrouter"),
|
|
||||||
("claude-sonnet-4.6", "anthropic/claude-sonnet-4.6", "openrouter"),
|
|
||||||
("gpt-5.6-sol", "openai/gpt-5.6-sol", "openrouter"),
|
|
||||||
("gpt-5.6-terra", "openai/gpt-5.6-terra", "openrouter"),
|
|
||||||
("gpt-5.6-luna", "openai/gpt-5.6-luna", "openrouter"),
|
|
||||||
("gpt-5.5-pro", "openai/gpt-5.5-pro", "openrouter"),
|
|
||||||
("gpt-5.5", "openai/gpt-5.5", "openrouter"),
|
|
||||||
("gpt-5.4", "openai/gpt-5.4", "openrouter"),
|
|
||||||
("gpt-5.3-codex", "openai/gpt-5.3-codex", "openrouter"),
|
|
||||||
("gemini-3.5-flash", "google/gemini-3.5-flash", "openrouter"),
|
|
||||||
("gemini-3.1-pro", "google/gemini-3.1-pro-preview", "openrouter"),
|
|
||||||
("gemini-3-flash", "google/gemini-3-flash-preview", "openrouter"),
|
|
||||||
("kimi-k2.6", "moonshotai/kimi-k2.6", "openrouter"),
|
|
||||||
("glm-5.2", "z-ai/glm-5.2", "openrouter"),
|
|
||||||
("glm-5v-turbo", "z-ai/glm-5v-turbo", "openrouter"),
|
|
||||||
("minimax-m3", "minimax/minimax-m3", "openrouter"),
|
|
||||||
("mimo-v2.5-pro", "xiaomi/mimo-v2.5-pro", "openrouter"),
|
|
||||||
("mimo-v2.5", "xiaomi/mimo-v2.5", "openrouter"),
|
|
||||||
("grok-build-0.1", "x-ai/grok-build-0.1", "openrouter"),
|
|
||||||
("grok-4.5", "x-ai/grok-4.5", "openrouter"),
|
|
||||||
("hy3", "tencent/hy3", "openrouter"),
|
|
||||||
("qwen3.7-max", "qwen/qwen3.7-max", "openrouter"),
|
|
||||||
("qwen3.7-plus", "qwen/qwen3.7-plus", "openrouter"),
|
|
||||||
("qwen3.6-flash", "qwen/qwen3.6-flash", "openrouter"),
|
|
||||||
("qwen3.5-122b", "qwen/qwen3.5-122b-a10b", "openrouter"),
|
|
||||||
("deepseek-v4-pro", "deepseek/deepseek-v4-pro", "openrouter"),
|
|
||||||
("deepseek-v4-flash", "deepseek/deepseek-v4-flash", "openrouter"),
|
|
||||||
# Zhipu CodePlan (智谱代码计划 — coding-only endpoint)
|
|
||||||
("glm-5.2", "glm-5.2", "zhipu-code"),
|
|
||||||
("glm-5.1", "glm-5.1", "zhipu-code"),
|
|
||||||
("glm-5", "glm-5", "zhipu-code"),
|
|
||||||
("glm-5-turbo", "glm-5-turbo", "zhipu-code"),
|
|
||||||
("glm-5v-turbo", "glm-5v-turbo", "zhipu-code"),
|
|
||||||
("glm-4.7", "glm-4.7", "zhipu-code"),
|
|
||||||
# Zhipu (智谱 — general endpoint, default for simple lookups)
|
|
||||||
("glm-5.2", "glm-5.2", "zhipu"),
|
|
||||||
("glm-5.1", "glm-5.1", "zhipu"),
|
|
||||||
("glm-5", "glm-5", "zhipu"),
|
|
||||||
("glm-5-turbo", "glm-5-turbo", "zhipu"),
|
|
||||||
("glm-5v-turbo", "glm-5v-turbo", "zhipu"),
|
|
||||||
("glm-4.7", "glm-4.7", "zhipu"),
|
|
||||||
# Volcengine (火山引擎 — Doubao models)
|
|
||||||
("doubao-seed-2.0-pro", "doubao-seed-2-0-pro-260215", "volcengine"),
|
|
||||||
("doubao-seed-2.0-lite", "doubao-seed-2-0-lite-260215", "volcengine"),
|
|
||||||
("doubao-seed-2.0-mini", "doubao-seed-2-0-mini-260215", "volcengine"),
|
|
||||||
("doubao-seed-2.0-code", "doubao-seed-2-0-code-preview-260215", "volcengine"),
|
|
||||||
("doubao-seed-1.6", "doubao-seed-1.6", "volcengine"),
|
|
||||||
("doubao-1.5-pro", "doubao-1.5-pro-256k", "volcengine"),
|
|
||||||
("doubao-1.5-thinking-pro", "doubao-1.5-thinking-pro", "volcengine"),
|
|
||||||
# DashScope Coding Plan (阿里云代码计划 — subscription sk-sp-* endpoint)
|
|
||||||
("qwen3.7-max", "qwen3.7-max", "dashscope-code"),
|
|
||||||
("qwen3.7-plus", "qwen3.7-plus", "dashscope-code"),
|
|
||||||
("qwen3.6-max", "qwen3.6-max-preview", "dashscope-code"),
|
|
||||||
("qwen3.6-plus", "qwen3.6-plus", "dashscope-code"),
|
|
||||||
("qwen3.6-flash", "qwen3.6-flash", "dashscope-code"),
|
|
||||||
("qwen3-coder", "qwen3-coder-plus", "dashscope-code"),
|
|
||||||
("qwen3-coder-next", "qwen3-coder-next", "dashscope-code"),
|
|
||||||
("qwen3-max", "qwen3-max", "dashscope-code"),
|
|
||||||
("qwen3.5-plus", "qwen3.5-plus", "dashscope-code"),
|
|
||||||
# DashScope (阿里云 — Qwen models, default for simple lookups)
|
|
||||||
("qwen3.7-max", "qwen3.7-max", "dashscope"),
|
|
||||||
("qwen3.7-plus", "qwen3.7-plus", "dashscope"),
|
|
||||||
("qwen3.6-max", "qwen3.6-max-preview", "dashscope"),
|
|
||||||
("qwen3.6-plus", "qwen3.6-plus", "dashscope"),
|
|
||||||
("qwen3.6-flash", "qwen3.6-flash", "dashscope"),
|
|
||||||
("qwen3-coder", "qwen3-coder-plus", "dashscope"),
|
|
||||||
("qwen3-235b", "qwen3-235b-a22b", "dashscope"),
|
|
||||||
("qwen-max", "qwen-max", "dashscope"),
|
|
||||||
("qwq-plus", "qwq-plus", "dashscope"),
|
|
||||||
# DeepSeek
|
|
||||||
("deepseek-v4-pro", "deepseek-v4-pro", "deepseek"),
|
|
||||||
("deepseek-v4-flash", "deepseek-v4-flash", "deepseek"),
|
|
||||||
# Legacy aliases (deprecated 2026-07-24; route to v4-flash thinking/non-thinking)
|
|
||||||
("deepseek-r1", "deepseek-reasoner", "deepseek"),
|
|
||||||
("deepseek-v3", "deepseek-chat", "deepseek"),
|
|
||||||
# Moonshot (OpenAI-compatible)
|
|
||||||
("kimi-k2.6", "kimi-k2.6", "moonshot"),
|
|
||||||
("kimi-k2.5", "kimi-k2.5", "moonshot"),
|
|
||||||
("kimi-k2-thinking", "kimi-k2-thinking", "moonshot"),
|
|
||||||
("kimi-k2-thinking-turbo", "kimi-k2-thinking-turbo", "moonshot"),
|
|
||||||
("moonshot-v1-auto", "moonshot-v1-auto", "moonshot"),
|
|
||||||
("moonshot-v1-128k", "moonshot-v1-128k", "moonshot"),
|
|
||||||
("moonshot-v1-32k", "moonshot-v1-32k", "moonshot"),
|
|
||||||
("moonshot-v1-8k", "moonshot-v1-8k", "moonshot"),
|
|
||||||
# Kimi Coding Plan (Anthropic-compatible)
|
|
||||||
("kimi-for-coding", "kimi-for-coding", "kimi-coding"),
|
|
||||||
]
|
|
||||||
|
|
||||||
# Public dict for simple lookups (last entry wins for duplicate names).
|
|
||||||
# Use get_models_for_provider() for provider-aware lookups.
|
|
||||||
MODELS: dict[str, tuple[str, str]] = {
|
|
||||||
name: (model_id, provider) for name, model_id, provider in _MODEL_ENTRIES
|
|
||||||
}
|
|
||||||
|
|
||||||
DEFAULT_MODEL = "claude-sonnet-4-6"
|
|
||||||
|
|
||||||
|
|
||||||
def get_models_for_provider(provider: str) -> list[tuple[str, str]]:
|
|
||||||
"""Get all models for a specific provider.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
provider: Provider name (e.g., 'anthropic', 'openrouter').
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of (short_name, model_id) tuples for the provider.
|
|
||||||
"""
|
|
||||||
return [(name, model_id) for name, model_id, p in _MODEL_ENTRIES if p == provider]
|
|
||||||
|
|
||||||
|
|
||||||
def _env_flag_enabled(name: str) -> bool:
|
def _env_flag_enabled(name: str) -> bool:
|
||||||
return os.environ.get(name, "").strip().lower() in _TRUTHY_ENV_VALUES
|
return os.environ.get(name, "").strip().lower() in _TRUTHY_ENV_VALUES
|
||||||
@@ -333,11 +200,36 @@ def _drop_unsupported_chat_model_kwargs(kwargs: dict[str, Any]) -> None:
|
|||||||
model_kwargs.pop(key, None)
|
model_kwargs.pop(key, None)
|
||||||
|
|
||||||
|
|
||||||
def _supports_openrouter_anthropic_prompt_cache(provider: str, model_id: str) -> bool:
|
_IMPLICIT_CACHE_PROVIDERS = frozenset({"zhipu", "zhipu-code", "siliconflow", "nvidia"})
|
||||||
"""Return whether EvoScientist should declare OpenRouter Claude caching."""
|
|
||||||
return provider == "openrouter" and model_id.startswith(
|
# OpenAI-compatible routers that forward an Anthropic-style ``cache_control``
|
||||||
|
# declaration through to Claude models addressed as ``anthropic/...``.
|
||||||
|
_EXPLICIT_CACHE_PROVIDERS = frozenset({"openrouter", "requesty"})
|
||||||
|
|
||||||
|
|
||||||
|
def _cache_strategy(provider: str, model_id: str) -> str:
|
||||||
|
"""Return the provider's prompt-cache mechanism.
|
||||||
|
|
||||||
|
``explicit`` — needs Anthropic-style ``cache_control`` markers (Claude routes
|
||||||
|
on the OpenAI-compatible routers).
|
||||||
|
``implicit`` — provider prefixes-cache automatically; no markers, but the
|
||||||
|
prompt prefix must stay byte-stable for hits (see memory injection order).
|
||||||
|
``none`` — no cache model to declare.
|
||||||
|
"""
|
||||||
|
if provider in _EXPLICIT_CACHE_PROVIDERS and model_id.startswith(
|
||||||
("anthropic/", "~anthropic/")
|
("anthropic/", "~anthropic/")
|
||||||
)
|
):
|
||||||
|
return "explicit"
|
||||||
|
if provider in _IMPLICIT_CACHE_PROVIDERS:
|
||||||
|
return "implicit"
|
||||||
|
return "none"
|
||||||
|
|
||||||
|
|
||||||
|
def _supports_openrouter_anthropic_prompt_cache(
|
||||||
|
provider: str | None, model_id: str
|
||||||
|
) -> bool:
|
||||||
|
"""Return whether EvoScientist should declare Claude caching for a router."""
|
||||||
|
return _cache_strategy(provider or "", model_id) == "explicit"
|
||||||
|
|
||||||
|
|
||||||
def _has_cache_control_override(kwargs: dict[str, Any]) -> bool:
|
def _has_cache_control_override(kwargs: dict[str, Any]) -> bool:
|
||||||
@@ -360,16 +252,25 @@ def _has_cache_control_override(kwargs: dict[str, Any]) -> bool:
|
|||||||
|
|
||||||
|
|
||||||
def _apply_openrouter_anthropic_prompt_cache(
|
def _apply_openrouter_anthropic_prompt_cache(
|
||||||
provider: str,
|
provider: str | None,
|
||||||
model_id: str,
|
model_id: str,
|
||||||
kwargs: dict[str, Any],
|
kwargs: dict[str, Any],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Declare OpenRouter Claude prompt caching unless explicitly disabled.
|
"""Declare router Claude prompt caching unless explicitly disabled.
|
||||||
|
|
||||||
OpenRouter already handles implicit caching for most providers, but Claude
|
OpenRouter and Requesty both handle implicit caching for most providers,
|
||||||
prompt caching needs Anthropic-style cache-control declaration.
|
but Claude prompt caching needs an Anthropic-style cache-control
|
||||||
|
declaration. Each router honours its own opt-out env flag
|
||||||
|
(``EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE`` /
|
||||||
|
``EVOSCIENTIST_REQUESTY_ANTHROPIC_PROMPT_CACHE``).
|
||||||
"""
|
"""
|
||||||
if _env_flag_disabled("EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"):
|
if provider is None:
|
||||||
|
return
|
||||||
|
disable_flag = {
|
||||||
|
"openrouter": "EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE",
|
||||||
|
"requesty": "EVOSCIENTIST_REQUESTY_ANTHROPIC_PROMPT_CACHE",
|
||||||
|
}.get(provider)
|
||||||
|
if disable_flag is not None and _env_flag_disabled(disable_flag):
|
||||||
return
|
return
|
||||||
if not _supports_openrouter_anthropic_prompt_cache(provider, model_id):
|
if not _supports_openrouter_anthropic_prompt_cache(provider, model_id):
|
||||||
return
|
return
|
||||||
@@ -378,6 +279,22 @@ def _apply_openrouter_anthropic_prompt_cache(
|
|||||||
kwargs.setdefault("model_kwargs", {})["cache_control"] = {"type": "ephemeral"}
|
kwargs.setdefault("model_kwargs", {})["cache_control"] = {"type": "ephemeral"}
|
||||||
|
|
||||||
|
|
||||||
|
def _enable_openrouter_429_retry(chat_model: Any) -> None:
|
||||||
|
"""Add 429 to the OpenRouter SDK's retryable status codes (default ["5XX"]).
|
||||||
|
|
||||||
|
Upstream rate limits ("temporarily rate-limited upstream", whose
|
||||||
|
Retry-After the SDK backoff already honors) otherwise fail the run outright.
|
||||||
|
"""
|
||||||
|
sdk_config = getattr(getattr(chat_model, "client", None), "sdk_configuration", None)
|
||||||
|
retry_config: Any = getattr(sdk_config, "retry_config", None)
|
||||||
|
# Skip the UNSET sentinel (max_retries=0) and explicit caller overrides.
|
||||||
|
if not hasattr(retry_config, "status_codes_override"):
|
||||||
|
return
|
||||||
|
if retry_config.status_codes_override:
|
||||||
|
return
|
||||||
|
retry_config.status_codes_override = ["429", "5XX"]
|
||||||
|
|
||||||
|
|
||||||
def _apply_auto_config(
|
def _apply_auto_config(
|
||||||
provider: str,
|
provider: str,
|
||||||
model_id: str,
|
model_id: str,
|
||||||
@@ -410,8 +327,15 @@ def _apply_auto_config(
|
|||||||
else:
|
else:
|
||||||
_is_proxy = False
|
_is_proxy = False
|
||||||
if _is_proxy or (is_third_party and not _supports_thinking):
|
if _is_proxy or (is_third_party and not _supports_thinking):
|
||||||
pass
|
# Mandatory-thinking Kimi models (K3 / Kimi For Coding) must declare
|
||||||
elif "fable" in model_id or model_id.endswith(("4-6", "4-7", "4-8")):
|
# thinking so with_structured_output avoids forced tool_choice (400).
|
||||||
|
# max_tokens must exceed budget_tokens (default resolves to 4096).
|
||||||
|
if is_third_party and _is_mandatory_thinking_kimi(model_id):
|
||||||
|
kwargs["thinking"] = {"type": "enabled", "budget_tokens": 10000}
|
||||||
|
kwargs.setdefault("max_tokens", 16000)
|
||||||
|
elif "fable" in model_id or model_id.endswith(
|
||||||
|
("opus-5", "sonnet-5", "4-6", "4-7", "4-8")
|
||||||
|
):
|
||||||
kwargs["thinking"] = {"type": "adaptive", "display": "summarized"}
|
kwargs["thinking"] = {"type": "adaptive", "display": "summarized"}
|
||||||
kwargs.setdefault("effort", "max")
|
kwargs.setdefault("effort", "max")
|
||||||
else:
|
else:
|
||||||
@@ -472,6 +396,9 @@ def get_chat_model(
|
|||||||
>>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID
|
>>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID
|
||||||
"""
|
"""
|
||||||
model = model or DEFAULT_MODEL
|
model = model or DEFAULT_MODEL
|
||||||
|
# Captured before any auto-configuration: an explicit caller reasoning block
|
||||||
|
# must survive an explicit `use_responses_api=False`, auto-injected one must not.
|
||||||
|
_caller_supplied_reasoning = "reasoning" in kwargs
|
||||||
|
|
||||||
# Look up short name in registry (provider-aware)
|
# Look up short name in registry (provider-aware)
|
||||||
model_id = None
|
model_id = None
|
||||||
@@ -550,12 +477,17 @@ def get_chat_model(
|
|||||||
if api_key:
|
if api_key:
|
||||||
kwargs.setdefault("api_key", api_key)
|
kwargs.setdefault("api_key", api_key)
|
||||||
|
|
||||||
|
elif provider == "deepseek":
|
||||||
|
api_key = os.environ.get("DEEPSEEK_API_KEY", "")
|
||||||
|
if api_key:
|
||||||
|
kwargs["api_key"] = api_key
|
||||||
|
|
||||||
# OpenAI-routed providers → route through OpenAI provider with base_url
|
# OpenAI-routed providers → route through OpenAI provider with base_url
|
||||||
elif provider in _OPENAI_ROUTED_PROVIDERS:
|
elif provider in _OPENAI_ROUTED_PROVIDERS:
|
||||||
_original_provider = provider
|
_original_provider = provider
|
||||||
base_url_default, api_key_env = _OPENAI_ROUTED_PROVIDERS[provider]
|
base_url_default, api_key_env = _OPENAI_ROUTED_PROVIDERS[provider]
|
||||||
if provider == "custom-openai":
|
if provider == "custom-openai":
|
||||||
base_url = os.environ.get("CUSTOM_OPENAI_BASE_URL", "")
|
base_url = str(kwargs.get("base_url") or os.environ.get("CUSTOM_OPENAI_BASE_URL", ""))
|
||||||
if not base_url:
|
if not base_url:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"CUSTOM_OPENAI_BASE_URL environment variable is required when using "
|
"CUSTOM_OPENAI_BASE_URL environment variable is required when using "
|
||||||
@@ -570,14 +502,17 @@ def get_chat_model(
|
|||||||
api_key = os.environ.get(api_key_env, "")
|
api_key = os.environ.get(api_key_env, "")
|
||||||
if api_key:
|
if api_key:
|
||||||
kwargs.setdefault("api_key", api_key)
|
kwargs.setdefault("api_key", api_key)
|
||||||
|
_apply_openai_compat_reasoning_config(provider, model_id, kwargs)
|
||||||
# SiliconFlow: disable thinking — LangChain drops reasoning_content
|
# SiliconFlow: disable thinking — LangChain drops reasoning_content
|
||||||
# from history, causing error 20015 on multi-turn requests.
|
# from history, causing error 20015 on multi-turn requests.
|
||||||
if provider == "siliconflow":
|
if provider == "siliconflow":
|
||||||
kwargs.setdefault("extra_body", {})["enable_thinking"] = False
|
kwargs.setdefault("extra_body", {})["enable_thinking"] = False
|
||||||
# Moonshot: disable thinking for all models to prevent LangChain from dropping
|
# Moonshot: disable thinking for pre-K3 models to prevent LangChain from
|
||||||
# reasoning_content, which causes multi-turn conversation errors (error 20015).
|
# dropping reasoning_content, which causes multi-turn conversation errors
|
||||||
# Even native thinking models like kimi-k2-thinking operate in non-thinking mode.
|
# (error 20015). Even native thinking models like kimi-k2-thinking operate
|
||||||
if provider == "moonshot":
|
# in non-thinking mode. kimi-k3+ is exempt: always-thinking, and
|
||||||
|
# Moonshot's K3 guide forbids the K2.x `thinking` parameter for it.
|
||||||
|
if provider == "moonshot" and not model_id.startswith("kimi-k3"):
|
||||||
kwargs.setdefault("extra_body", {})["thinking"] = {"type": "disabled"}
|
kwargs.setdefault("extra_body", {})["thinking"] = {"type": "disabled"}
|
||||||
provider = "openai"
|
provider = "openai"
|
||||||
|
|
||||||
@@ -593,7 +528,13 @@ def get_chat_model(
|
|||||||
# passback (OpenRouter's `/responses` beta is stateless, store=false —
|
# passback (OpenRouter's `/responses` beta is stateless, store=false —
|
||||||
# "Item with id 'rs_...' not found"); the patch strips them on passback,
|
# "Item with id 'rs_...' not found"); the patch strips them on passback,
|
||||||
# so enabling `summary` is safe. See langchain-ai/langchain#37777.
|
# so enabling `summary` is safe. See langchain-ai/langchain#37777.
|
||||||
kwargs.setdefault("reasoning", {"effort": "high", "summary": "auto"})
|
# Note: mandatory-reasoning endpoints (kimi-k3, grok-4.5, …) reject
|
||||||
|
# effort "none" with HTTP 400 — that error is surfaced to the user
|
||||||
|
# as-is; pick a real effort (low+) for those models.
|
||||||
|
# Ai4Sci: an invocation parameter is owned by the compiled plan, so the
|
||||||
|
# deployment environment must not alter it; medium is the fixed default.
|
||||||
|
effort = "medium"
|
||||||
|
kwargs.setdefault("reasoning", {"effort": effort, "summary": "auto"})
|
||||||
# App attribution (issue #339): identify EvoScientist to OpenRouter so
|
# App attribution (issue #339): identify EvoScientist to OpenRouter so
|
||||||
# usage is credited to the project (app rankings, model app tabs,
|
# usage is credited to the project (app rankings, model app tabs,
|
||||||
# analytics) rather than langchain-openrouter's LangChain-branded
|
# analytics) rather than langchain-openrouter's LangChain-branded
|
||||||
@@ -611,6 +552,11 @@ def get_chat_model(
|
|||||||
os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE", "").strip()
|
os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE", "").strip()
|
||||||
or OPENROUTER_DEFAULT_APP_TITLE,
|
or OPENROUTER_DEFAULT_APP_TITLE,
|
||||||
)
|
)
|
||||||
|
# OpenRouter keys app pages by HTTP-Referer and X-Title only renames that
|
||||||
|
# page, so a custom title on the default referer would rename the shared
|
||||||
|
# EvoScientist page for everyone. Honor it only with a custom referer.
|
||||||
|
if kwargs["app_url"] == OPENROUTER_DEFAULT_HTTP_REFERER:
|
||||||
|
kwargs["app_title"] = OPENROUTER_DEFAULT_APP_TITLE
|
||||||
# app_categories must be a list[str] (langchain-openrouter joins it into
|
# app_categories must be a list[str] (langchain-openrouter joins it into
|
||||||
# the X-OpenRouter-Categories header); split the comma-separated config
|
# the X-OpenRouter-Categories header); split the comma-separated config
|
||||||
# value and drop blanks so a stray comma/space can't emit an empty one.
|
# value and drop blanks so a stray comma/space can't emit an empty one.
|
||||||
@@ -641,6 +587,7 @@ def get_chat_model(
|
|||||||
if _app_categories:
|
if _app_categories:
|
||||||
kwargs.setdefault("app_categories", _app_categories)
|
kwargs.setdefault("app_categories", _app_categories)
|
||||||
_patch_openrouter_strip_responses_reasoning()
|
_patch_openrouter_strip_responses_reasoning()
|
||||||
|
_patch_openrouter_structured_output()
|
||||||
|
|
||||||
# Anthropic-routed providers → route through Anthropic provider with base_url
|
# Anthropic-routed providers → route through Anthropic provider with base_url
|
||||||
elif provider in _ANTHROPIC_ROUTED_PROVIDERS:
|
elif provider in _ANTHROPIC_ROUTED_PROVIDERS:
|
||||||
@@ -676,25 +623,79 @@ def get_chat_model(
|
|||||||
|
|
||||||
_drop_unsupported_chat_model_kwargs(kwargs)
|
_drop_unsupported_chat_model_kwargs(kwargs)
|
||||||
_apply_auto_config(provider, model_id, _is_third_party, kwargs, _original_provider)
|
_apply_auto_config(provider, model_id, _is_third_party, kwargs, _original_provider)
|
||||||
_apply_openrouter_anthropic_prompt_cache(provider, model_id, kwargs)
|
# OpenAI-routed routers (e.g. Requesty) reassign ``provider`` to "openai"
|
||||||
|
# above, so use the original provider name to detect router-level caching.
|
||||||
|
_cache_provider = _original_provider or provider
|
||||||
|
_apply_openrouter_anthropic_prompt_cache(_cache_provider, model_id, kwargs)
|
||||||
|
|
||||||
|
_uses_native_deepseek = provider == "deepseek" or (
|
||||||
|
provider == "openai"
|
||||||
|
and _original_provider == "custom-openai"
|
||||||
|
and _is_deepseek_endpoint(kwargs.get("base_url"))
|
||||||
|
)
|
||||||
|
|
||||||
|
# User-level override for the OpenAI Responses API vs Chat Completions.
|
||||||
|
# When "false", force Chat Completions and drop reasoning (which triggers
|
||||||
|
# the Responses API path in langchain-openai). Only applies to OpenAI.
|
||||||
|
if _uses_native_deepseek:
|
||||||
|
if kwargs.get("use_responses_api") is True:
|
||||||
|
raise ValueError(
|
||||||
|
"DeepSeek does not support the OpenAI Responses API. "
|
||||||
|
"Remove use_responses_api=True."
|
||||||
|
)
|
||||||
|
kwargs.pop("use_responses_api", None)
|
||||||
|
elif provider == "openai":
|
||||||
|
if "use_responses_api" in kwargs:
|
||||||
|
# An explicit per-call plan always wins over the deployment env
|
||||||
|
# (Ai4Sci compiles the invocation plan; env is only a default).
|
||||||
|
if kwargs["use_responses_api"] is False and not _caller_supplied_reasoning:
|
||||||
|
# Chat Completions cannot carry the Responses-only reasoning block.
|
||||||
|
kwargs.pop("reasoning", None)
|
||||||
|
else:
|
||||||
|
_responses_api_setting = (
|
||||||
|
os.environ.get("EVOSCIENTIST_USE_RESPONSES_API", "").strip().lower()
|
||||||
|
)
|
||||||
|
if _responses_api_setting == "false":
|
||||||
|
kwargs["use_responses_api"] = False
|
||||||
|
kwargs.pop("reasoning", None)
|
||||||
|
elif _responses_api_setting == "true":
|
||||||
|
kwargs["use_responses_api"] = True
|
||||||
|
|
||||||
|
if _is_openai_proxy and kwargs.get("use_responses_api") is True:
|
||||||
|
reasoning = kwargs.setdefault("reasoning", {})
|
||||||
|
if isinstance(reasoning, dict):
|
||||||
|
reasoning = dict(reasoning)
|
||||||
|
reasoning.setdefault("context", "all_turns")
|
||||||
|
kwargs["reasoning"] = reasoning
|
||||||
|
|
||||||
|
# Ai4Sci: an ambient ANTHROPIC_AUTH_TOKEN would silently override the
|
||||||
|
# explicit api_key resolved for this provider, so hide it for this call.
|
||||||
anthropic_auth_token = None
|
anthropic_auth_token = None
|
||||||
if provider == "anthropic" and kwargs.get("api_key"):
|
if provider == "anthropic" and kwargs.get("api_key"):
|
||||||
anthropic_auth_token = os.environ.pop("ANTHROPIC_AUTH_TOKEN", None)
|
anthropic_auth_token = os.environ.pop("ANTHROPIC_AUTH_TOKEN", None)
|
||||||
try:
|
try:
|
||||||
chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs)
|
if _uses_native_deepseek:
|
||||||
|
chat_model = EvoChatDeepSeek(model=model_id, **kwargs)
|
||||||
|
else:
|
||||||
|
chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs)
|
||||||
finally:
|
finally:
|
||||||
if anthropic_auth_token is not None:
|
if anthropic_auth_token is not None:
|
||||||
os.environ["ANTHROPIC_AUTH_TOKEN"] = anthropic_auth_token
|
os.environ["ANTHROPIC_AUTH_TOKEN"] = anthropic_auth_token
|
||||||
|
|
||||||
# Flatten list content to strings for strict OpenAI-compatible providers
|
# Flatten list content to strings for strict OpenAI-compatible providers
|
||||||
# (DeepSeek, SiliconFlow, OpenRouter, custom-openai, etc.) and
|
# (SiliconFlow, OpenRouter, custom-openai, etc.) and
|
||||||
# native OpenAI through a proxy, to avoid "sequence expected string" errors.
|
# native OpenAI through a proxy, to avoid "sequence expected string" errors.
|
||||||
# Moonshot and Kimi Coding support standard format, no patch needed.
|
# Moonshot and Kimi Coding support standard format, no patch needed.
|
||||||
|
# Mandatory-thinking Kimi models on Anthropic-routed endpoints are exempt:
|
||||||
|
# flatten drops thinking blocks, which Kimi requires on tool-call turns.
|
||||||
_no_patch_providers = {"moonshot", "kimi-coding"}
|
_no_patch_providers = {"moonshot", "kimi-coding"}
|
||||||
if (
|
if (
|
||||||
_is_third_party or _is_openai_proxy
|
(_is_third_party or _is_openai_proxy)
|
||||||
) and _original_provider not in _no_patch_providers:
|
and _original_provider not in _no_patch_providers
|
||||||
|
and not _uses_native_deepseek
|
||||||
|
and not (provider == "anthropic" and _is_mandatory_thinking_kimi(model_id))
|
||||||
|
and _original_provider not in _ANTHROPIC_ROUTED_PROVIDERS
|
||||||
|
):
|
||||||
# Anthropic-routed providers accept media in tool results natively;
|
# Anthropic-routed providers accept media in tool results natively;
|
||||||
# only OpenAI-compatible providers need tool-media hoisting.
|
# only OpenAI-compatible providers need tool-media hoisting.
|
||||||
_hoist = _original_provider not in _ANTHROPIC_ROUTED_PROVIDERS
|
_hoist = _original_provider not in _ANTHROPIC_ROUTED_PROVIDERS
|
||||||
@@ -712,77 +713,17 @@ def get_chat_model(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# DeepSeek thinking mode requires reasoning_content passback in multi-turn
|
|
||||||
# + tool_use scenarios.
|
|
||||||
if _original_provider == "deepseek":
|
|
||||||
_patch_deepseek_reasoning_passback(chat_model)
|
|
||||||
|
|
||||||
if _is_openai_proxy:
|
if _is_openai_proxy:
|
||||||
_patch_ccproxy_system_to_developer(chat_model)
|
_patch_ccproxy_system_to_developer(chat_model)
|
||||||
|
|
||||||
|
if provider == "openrouter":
|
||||||
|
_enable_openrouter_429_retry(chat_model)
|
||||||
|
|
||||||
|
if provider == "anthropic":
|
||||||
|
_patch_anthropic_strip_foreign_reasoning()
|
||||||
|
_patch_anthropic_structured_output()
|
||||||
|
|
||||||
apply_known_context_window(chat_model)
|
apply_known_context_window(chat_model)
|
||||||
|
|
||||||
return chat_model
|
return chat_model
|
||||||
|
|
||||||
|
|
||||||
def list_models() -> list[str]:
|
|
||||||
"""List all available model short names.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of unique model short names that can be passed to get_chat_model().
|
|
||||||
"""
|
|
||||||
seen = set()
|
|
||||||
result = []
|
|
||||||
for name, _, _ in _MODEL_ENTRIES:
|
|
||||||
if name not in seen:
|
|
||||||
seen.add(name)
|
|
||||||
result.append(name)
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
def list_models_by_provider() -> list[tuple[str, str, str]]:
|
|
||||||
"""List all unique (short_name, model_id, provider) entries.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
De-duplicated list of model entries preserving registry order.
|
|
||||||
"""
|
|
||||||
seen: set[tuple[str, str]] = set()
|
|
||||||
result: list[tuple[str, str, str]] = []
|
|
||||||
for name, model_id, provider in _MODEL_ENTRIES:
|
|
||||||
key = (name, provider)
|
|
||||||
if key not in seen:
|
|
||||||
seen.add(key)
|
|
||||||
result.append((name, model_id, provider))
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
async def list_model_picker_entries(
|
|
||||||
ollama_base_url: str | None,
|
|
||||||
*,
|
|
||||||
include_custom_ollama: bool,
|
|
||||||
) -> list[tuple[str, str, str]]:
|
|
||||||
"""Return model picker entries, optionally including local Ollama models."""
|
|
||||||
entries = list_models_by_provider()
|
|
||||||
if ollama_base_url:
|
|
||||||
from .ollama_discovery import discover_ollama_models
|
|
||||||
|
|
||||||
for detected_name in await discover_ollama_models(
|
|
||||||
ollama_base_url,
|
|
||||||
timeout=1.5,
|
|
||||||
):
|
|
||||||
entries.append((detected_name, detected_name, "ollama"))
|
|
||||||
if include_custom_ollama:
|
|
||||||
entries.append(("Custom Ollama model...", "__custom_ollama__", "ollama"))
|
|
||||||
return entries
|
|
||||||
|
|
||||||
|
|
||||||
def get_model_info(model: str) -> tuple[str, str] | None:
|
|
||||||
"""Get the (model_id, provider) tuple for a short name.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
model: Short model name.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (model_id, provider) or None if not found.
|
|
||||||
"""
|
|
||||||
return MODELS.get(model)
|
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user