Merge branch 'simp/r3-12-w2b' into simp/integration3
This commit is contained in:
@@ -1,9 +1,7 @@
|
||||
"""Gateway slash-command handlers for GatewayRunner.
|
||||
|
||||
Lifted out of ``gateway/run.py`` into a mixin so ``self._handle_*_command`` keeps resolving via the
|
||||
MRO. Cohesive clusters live in the sibling mixins (``slash_commands_model/_session/_status/_goals``);
|
||||
this module keeps the shared helpers plus the one-off commands. run.py helpers are imported lazily.
|
||||
"""
|
||||
"""Gateway slash-command handlers for GatewayRunner: lifted out of ``gateway/run.py`` into a mixin
|
||||
so ``self._handle_*_command`` keeps resolving via the MRO. Cohesive clusters live in the sibling
|
||||
mixins (``slash_commands_model/_session/_status/_goals``); this module keeps the shared helpers plus
|
||||
the one-off commands. run.py helpers are imported lazily."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -28,8 +26,7 @@ from gateway.session import AsyncSessionStore
|
||||
from gateway.slash_commands_goals import GatewayGoalCommandsMixin
|
||||
from gateway.slash_commands_model import ( # noqa: F401 — _model_switch_skew_guard re-exported for tests
|
||||
GatewayModelCommandsMixin,
|
||||
_model_switch_skew_guard,
|
||||
)
|
||||
_model_switch_skew_guard)
|
||||
from gateway.slash_commands_session import GatewaySessionCommandsMixin
|
||||
from gateway.slash_commands_status import GatewayStatusCommandsMixin
|
||||
from hermes_cli.config import atomic_config_write, cfg_get
|
||||
@@ -39,11 +36,9 @@ logger = logging.getLogger("gateway.run")
|
||||
|
||||
|
||||
# /rollback result keys -> i18n line for files the safe restore left alone.
|
||||
_ROLLBACK_SKIP_LINES = (
|
||||
("skipped_user_edits", "gateway.rollback.kept_user_edits"),
|
||||
("skipped_oversize", "gateway.rollback.kept_oversize"),
|
||||
("failed_deletes", "gateway.rollback.failed_deletes"),
|
||||
)
|
||||
_ROLLBACK_SKIP_LINES = (("skipped_user_edits", "gateway.rollback.kept_user_edits"),
|
||||
("skipped_oversize", "gateway.rollback.kept_oversize"),
|
||||
("failed_deletes", "gateway.rollback.failed_deletes"))
|
||||
|
||||
# /busy input modes -> (status-card behavior, set-confirmation behavior).
|
||||
_BUSY_MODE_BEHAVIOR = {
|
||||
@@ -54,37 +49,27 @@ _BUSY_MODE_BEHAVIOR = {
|
||||
}
|
||||
|
||||
# /diff argument -> diff mode (unknown args leave the mode unchanged).
|
||||
_DIFF_MODE_BY_ARG = {
|
||||
"staged": "staged", "--staged": "staged", "cached": "staged", "--cached": "staged",
|
||||
"all": "all", "--all": "all", "head": "all",
|
||||
"session": "session",
|
||||
}
|
||||
_DIFF_MODE_BY_ARG = {**dict.fromkeys(("staged", "--staged", "cached", "--cached"), "staged"),
|
||||
**dict.fromkeys(("all", "--all", "head"), "all"), "session": "session"}
|
||||
|
||||
# /voice subcommand -> stored mode (None = auto-TTS disabled), confirmation i18n key.
|
||||
_VOICE_MODE_BY_ARG = {
|
||||
**dict.fromkeys(("on", "enable"), ("voice_only", "gateway.voice.enabled_voice_only")),
|
||||
**dict.fromkeys(("off", "disable"), ("off", "gateway.voice.disabled_text")),
|
||||
"tts": ("all", "gateway.voice.tts_enabled"),
|
||||
}
|
||||
"tts": ("all", "gateway.voice.tts_enabled")}
|
||||
|
||||
# /footer argument -> new enabled state ("" toggles; anything else is a usage error).
|
||||
_FOOTER_STATE_BY_ARG = {
|
||||
**dict.fromkeys(("on", "enable", "true", "1"), True),
|
||||
**dict.fromkeys(("off", "disable", "false", "0"), False),
|
||||
}
|
||||
_FOOTER_STATE_BY_ARG = {**dict.fromkeys(("on", "enable", "true", "1"), True),
|
||||
**dict.fromkeys(("off", "disable", "false", "0"), False)}
|
||||
|
||||
# /approve modifier tokens -> approval choice (default "once").
|
||||
_APPROVE_CHOICE_BY_ARG = {
|
||||
**dict.fromkeys(("always", "permanent", "permanently"), "always"),
|
||||
**dict.fromkeys(("session", "ses"), "session"),
|
||||
}
|
||||
_APPROVE_CHOICE_BY_ARG = {**dict.fromkeys(("always", "permanent", "permanently"), "always"),
|
||||
**dict.fromkeys(("session", "ses"), "session")}
|
||||
|
||||
_PLATFORM_USAGE = (
|
||||
"Usage: /platform <list|pause|resume> [name]\n"
|
||||
" /platform list — show platform status\n"
|
||||
" /platform pause <name> — stop retrying a failing platform\n"
|
||||
" /platform resume <name> — re-queue a paused platform"
|
||||
)
|
||||
_PLATFORM_USAGE = ("Usage: /platform <list|pause|resume> [name]\n"
|
||||
" /platform list — show platform status\n"
|
||||
" /platform pause <name> — stop retrying a failing platform\n"
|
||||
" /platform resume <name> — re-queue a paused platform")
|
||||
|
||||
_WINDOWS_UPDATE_HELPER = """
|
||||
import os, subprocess, sys
|
||||
@@ -124,30 +109,18 @@ def _restart_notify_payload(event: MessageEvent) -> dict:
|
||||
if source.delivered_via_upstream_relay is True:
|
||||
data["delivered_via_upstream_relay"] = True
|
||||
data.update({k: getattr(source, k) for k in ("user_id", "scope_id") if getattr(source, k)})
|
||||
if source.thread_id:
|
||||
data["thread_id"] = source.thread_id
|
||||
if event.message_id:
|
||||
data["message_id"] = event.message_id
|
||||
return data
|
||||
|
||||
|
||||
def _restart_dedup_payload(event: MessageEvent) -> dict:
|
||||
"""Platform + update_id of the triggering /restart, for redelivery detection."""
|
||||
data = {"platform": event.source.platform.value if event.source.platform else None, "requested_at": time.time()}
|
||||
if event.platform_update_id is not None:
|
||||
data["update_id"] = event.platform_update_id
|
||||
optional = (("thread_id", source.thread_id), ("message_id", event.message_id))
|
||||
data.update({k: v for k, v in optional if v})
|
||||
return data
|
||||
|
||||
|
||||
def _spawn_detached_update(hermes_cmd, output_path, exit_code_path) -> None:
|
||||
"""Spawn ``hermes update --gateway`` detached so it survives the gateway restart it may trigger.
|
||||
|
||||
setsid is portable (works where ``systemd-run --user`` lacks a D-Bus session); ``--gateway``
|
||||
enables file-based IPC so interactive prompts are forwarded; PYTHONUNBUFFERED lets the gateway
|
||||
stream output live. Windows has no setsid: an inline helper runs the updater as a module under
|
||||
stream output live. Windows has no setsid: an inline helper runs the updater as a module under
|
||||
this interpreter (not venv\\Scripts\\hermes.exe — that shim holds its own file open, and the
|
||||
update must replace it), redirects both outputs to one file and writes the exit code.
|
||||
"""
|
||||
update must replace it), redirects both outputs to one file and writes the exit code."""
|
||||
import shutil
|
||||
import subprocess
|
||||
if sys.platform == "win32":
|
||||
@@ -155,8 +128,7 @@ def _spawn_detached_update(hermes_cmd, output_path, exit_code_path) -> None:
|
||||
subprocess.Popen(
|
||||
[sys.executable, "-c", _WINDOWS_UPDATE_HELPER, str(output_path), str(exit_code_path),
|
||||
sys.executable, "-m", "hermes_cli.main", "update", "--gateway"],
|
||||
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, **windows_detach_popen_kwargs(),
|
||||
)
|
||||
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, **windows_detach_popen_kwargs())
|
||||
return
|
||||
hermes_cmd_str = " ".join(shlex.quote(part) for part in hermes_cmd)
|
||||
update_cmd = (
|
||||
@@ -164,8 +136,7 @@ def _spawn_detached_update(hermes_cmd, output_path, exit_code_path) -> None:
|
||||
f" > {shlex.quote(str(output_path))} 2>&1; "
|
||||
# Avoid `status=$?`: `status` is read-only in zsh and this template is reused in
|
||||
# macOS/zsh operator wrappers, so keep it zsh-safe even though bash runs it here.
|
||||
f"rc=$?; printf '%s' \"$rc\" > {shlex.quote(str(exit_code_path))}"
|
||||
)
|
||||
f"rc=$?; printf '%s' \"$rc\" > {shlex.quote(str(exit_code_path))}")
|
||||
# Preferred: setsid creates a new session, fully detached; fallback start_new_session=True
|
||||
# calls os.setsid() in the child.
|
||||
setsid_bin = shutil.which("setsid")
|
||||
@@ -174,12 +145,10 @@ def _spawn_detached_update(hermes_cmd, output_path, exit_code_path) -> None:
|
||||
|
||||
|
||||
def _home_thread_from_source(source) -> Optional[str]:
|
||||
"""The thread id /sethome should persist on the home target, or None.
|
||||
|
||||
Slack thread-per-message keying stamps a top-level message's own id as ``source.thread_id`` (a
|
||||
session key, not a location); persisting it would pin HOME to that ephemeral thread. A thread
|
||||
id equal to the message's own id is synthetic and dropped; a real thread (id = parent's) is kept.
|
||||
"""
|
||||
"""The thread id /sethome should persist on the home target, or None. Slack thread-per-message
|
||||
keying stamps a top-level message's own id as ``source.thread_id`` (a session key, not a
|
||||
location); persisting it would pin HOME to that ephemeral thread. A thread id equal to the
|
||||
message's own id is synthetic and dropped; a real thread (id = parent's) is kept."""
|
||||
thread_id = getattr(source, "thread_id", None)
|
||||
if not thread_id:
|
||||
return None
|
||||
@@ -192,8 +161,7 @@ class GatewaySlashCommandsMixin(
|
||||
GatewayModelCommandsMixin,
|
||||
GatewaySessionCommandsMixin,
|
||||
GatewayStatusCommandsMixin,
|
||||
GatewayGoalCommandsMixin,
|
||||
):
|
||||
GatewayGoalCommandsMixin):
|
||||
"""In-session slash-command handlers for GatewayRunner (plus the helpers the sibling mixins share)."""
|
||||
|
||||
async_session_store: AsyncSessionStore
|
||||
@@ -249,8 +217,8 @@ class GatewaySlashCommandsMixin(
|
||||
if not cp["checkpoints_enabled"]:
|
||||
return None
|
||||
# AIAgent kwargs are ``checkpoint_<field>``; CheckpointManager takes the bare field names.
|
||||
return CheckpointManager(enabled=True, **{k[len("checkpoint_"):]: v for k, v in cp.items()
|
||||
if k.startswith("checkpoint_")})
|
||||
fields = {k[len("checkpoint_"):]: v for k, v in cp.items() if k.startswith("checkpoint_")}
|
||||
return CheckpointManager(enabled=True, **fields)
|
||||
|
||||
def _write_approval_setter(self, section: str, event: MessageEvent):
|
||||
"""``set_mode_fn`` for /memory and /skills: persist ``<section>.write_approval``. Raw read is
|
||||
@@ -282,12 +250,10 @@ class GatewaySlashCommandsMixin(
|
||||
try:
|
||||
await adapter.send(
|
||||
source.chat_id, confirmation_text, reply_to=event.message_id,
|
||||
metadata={"is_approval_prompt": True, "force_proactive_send": True},
|
||||
)
|
||||
metadata={"is_approval_prompt": True, "force_proactive_send": True})
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to send /%s confirmation to %s: %s", verb, source.chat_id, exc, exc_info=True,
|
||||
)
|
||||
logger.warning("Failed to send /%s confirmation to %s: %s", verb, source.chat_id,
|
||||
exc, exc_info=True)
|
||||
return None
|
||||
|
||||
def _typed_command_prefix_for(self, platform) -> str:
|
||||
@@ -308,12 +274,10 @@ class GatewaySlashCommandsMixin(
|
||||
return _gateway_config_home() / "config.yaml", _platform_config_key(event.source.platform)
|
||||
|
||||
async def _handle_profile_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /profile — show the profile serving this source and its home.
|
||||
|
||||
On a multiplexed gateway the process-level profile is the multiplexer's own ("default" in
|
||||
every chat), so with ``multiplex_profiles`` on report ``source.profile`` and resolve home under
|
||||
that profile's runtime scope; when off the stamp is ignored, mirroring ``_run_agent``.
|
||||
"""
|
||||
"""Handle /profile — show the profile serving this source and its home. On a multiplexed
|
||||
gateway the process-level profile is the multiplexer's own ("default" in every chat), so
|
||||
with ``multiplex_profiles`` on report ``source.profile`` and resolve home under that
|
||||
profile's runtime scope; when off the stamp is ignored, mirroring ``_run_agent``."""
|
||||
from hermes_constants import display_hermes_home
|
||||
source = getattr(event, "source", None)
|
||||
profile_name = display = ""
|
||||
@@ -329,10 +293,8 @@ class GatewaySlashCommandsMixin(
|
||||
# Shared executor resolves process-level fallbacks; the multiplexed per-source overrides
|
||||
# (when any) ride in via options.
|
||||
reply = _execute("profile", options={"profile_name": profile_name, "home_display": display})
|
||||
return "\n".join([
|
||||
t("gateway.profile.header", profile=reply.data["profile"]),
|
||||
t("gateway.profile.home", home=reply.data["home"]),
|
||||
])
|
||||
return "\n".join([t("gateway.profile.header", profile=reply.data["profile"]),
|
||||
t("gateway.profile.home", home=reply.data["home"])])
|
||||
|
||||
async def _handle_whoami_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /whoami — platform, DM-vs-group scope, tier and runnable commands (always allowed)."""
|
||||
@@ -424,20 +386,17 @@ class GatewaySlashCommandsMixin(
|
||||
user_id_alt=_field("user_id_alt"),
|
||||
notifier_profile=getattr(self, "_kanban_notifier_profile", None) or self._active_profile_name(),
|
||||
# Subscribing from chat: deliver the passive message and wake the destination agent.
|
||||
delivery_mode="notify+wake", delivery_metadata=delivery_metadata,
|
||||
)
|
||||
delivery_mode="notify+wake", delivery_metadata=delivery_metadata)
|
||||
finally:
|
||||
conn.close()
|
||||
await asyncio.to_thread(_sub)
|
||||
return True
|
||||
|
||||
async def _handle_stop_command(self, event: MessageEvent) -> Union[str, EphemeralReply]:
|
||||
"""Handle /stop command - interrupt a running agent.
|
||||
|
||||
A truly hung agent (blocked thread never checking _interrupt_requested) is caught by the early
|
||||
intercept in _handle_message(); this handler runs via normal dispatch or as a fallback, and
|
||||
force-cleans the session lock in all cases. The session is preserved so the user can continue.
|
||||
"""
|
||||
"""Handle /stop command - interrupt a running agent. A truly hung agent (blocked thread
|
||||
never checking _interrupt_requested) is caught by the early intercept in _handle_message();
|
||||
this handler runs via normal dispatch or as a fallback, and force-cleans the session lock in
|
||||
all cases. The session is preserved so the user can continue."""
|
||||
from gateway.run import _AGENT_PENDING_SENTINEL, _INTERRUPT_REASON_STOP
|
||||
source = event.source
|
||||
session_entry = await self.async_session_store.get_or_create_session(source)
|
||||
@@ -445,8 +404,8 @@ class GatewaySlashCommandsMixin(
|
||||
|
||||
async def _stop(key: str, invalidation_reason: str) -> None:
|
||||
await self._interrupt_and_clear_session(
|
||||
key, source, interrupt_reason=_INTERRUPT_REASON_STOP, invalidation_reason=invalidation_reason,
|
||||
)
|
||||
key, source, interrupt_reason=_INTERRUPT_REASON_STOP,
|
||||
invalidation_reason=invalidation_reason)
|
||||
agent = self._running_agents.get(session_key)
|
||||
if agent is _AGENT_PENDING_SENTINEL: # force-clean the sentinel so the session is unlocked
|
||||
await _stop(session_key, "stop_command_pending")
|
||||
@@ -463,10 +422,8 @@ class GatewaySlashCommandsMixin(
|
||||
if sibling_keys and self._is_user_authorized(source):
|
||||
for sibling_key in sibling_keys:
|
||||
await _stop(sibling_key, "stop_command_thread_sibling")
|
||||
logger.info(
|
||||
"STOP (thread sibling) by %s — interrupted %d run(s) in thread: %s",
|
||||
session_key, len(sibling_keys), ", ".join(sibling_keys),
|
||||
)
|
||||
logger.info("STOP (thread sibling) by %s — interrupted %d run(s) in thread: %s",
|
||||
session_key, len(sibling_keys), ", ".join(sibling_keys))
|
||||
return EphemeralReply(t("gateway.stop.stopped"))
|
||||
|
||||
# No running agent anywhere for this scope. A platform status indicator can still be stuck —
|
||||
@@ -532,12 +489,11 @@ class GatewaySlashCommandsMixin(
|
||||
# update_id) and we see it *again*, it's a redelivery from PTB's graceful-shutdown get_updates
|
||||
# ACK failing on the way out. Ignoring it prevents a loop where every fresh gateway re-restarts.
|
||||
if self._is_stale_restart_redelivery(event):
|
||||
logger.info(
|
||||
"Ignoring redelivered /restart (platform=%s, update_id=%s) — "
|
||||
"already processed by a previous gateway instance.",
|
||||
event.source.platform.value if event.source and event.source.platform else "?",
|
||||
event.platform_update_id,
|
||||
)
|
||||
src = event.source
|
||||
logger.info("Ignoring redelivered /restart (platform=%s, update_id=%s) — "
|
||||
"already processed by a previous gateway instance.",
|
||||
src.platform.value if src and src.platform else "?",
|
||||
event.platform_update_id)
|
||||
return ""
|
||||
if self._restart_requested or self._draining:
|
||||
count = self._running_agent_count()
|
||||
@@ -558,12 +514,20 @@ class GatewaySlashCommandsMixin(
|
||||
self._restart_command_source = event.source
|
||||
return data
|
||||
|
||||
def _dedup_payload() -> dict:
|
||||
# Platform + update_id of the triggering /restart, for redelivery detection.
|
||||
data = {"platform": event.source.platform.value if event.source.platform else None,
|
||||
"requested_at": time.time()}
|
||||
if event.platform_update_id is not None:
|
||||
data["update_id"] = event.platform_update_id
|
||||
return data
|
||||
|
||||
# Save the requester's routing info so the new gateway process can notify them once back.
|
||||
await _write_marker(".restart_notify.json", _notify_payload, "notify file")
|
||||
# Record the triggering platform + update_id in a dedicated dedup marker. Unlike
|
||||
# .restart_notify.json (unlinked once the new gateway sends its notification) this persists
|
||||
# so a delayed Telegram redelivery is still detectable. Overwritten on every /restart.
|
||||
await _write_marker(".restart_last_processed.json", lambda: _restart_dedup_payload(event), "dedup marker")
|
||||
await _write_marker(".restart_last_processed.json", _dedup_payload, "dedup marker")
|
||||
active_agents = self._running_agent_count()
|
||||
# Under a service manager (systemd/launchd) or Docker/Podman, exit 75 so the supervisor /
|
||||
# restart policy restarts us — detached setsid+bash fails there (systemd KillMode=mixed kills
|
||||
@@ -603,20 +567,16 @@ class GatewaySlashCommandsMixin(
|
||||
adapter_for_source = getattr(self, "_adapter_for_source", None)
|
||||
relay_adapter = adapter_for_source(source) if callable(adapter_for_source) else None
|
||||
fronts_platform = getattr(relay_adapter, "fronts_platform", None)
|
||||
if (
|
||||
source.platform in {None, Platform.LOCAL, Platform.RELAY}
|
||||
or not getattr(source, "user_id", None)
|
||||
or not callable(fronts_platform)
|
||||
or not fronts_platform(source.platform)
|
||||
):
|
||||
if (source.platform in {None, Platform.LOCAL, Platform.RELAY}
|
||||
or not getattr(source, "user_id", None)
|
||||
or not callable(fronts_platform) or not fronts_platform(source.platform)):
|
||||
return t("gateway.set_home.save_failed",
|
||||
error="Relay does not authenticate this logical home target")
|
||||
thread_id = _home_thread_from_source(source)
|
||||
home = HomeChannel(
|
||||
platform=source.platform, chat_id=str(chat_id), name=chat_name, thread_id=thread_id,
|
||||
user_id=str(source.user_id) if getattr(source, "user_id", None) else None,
|
||||
scope_id=str(source.scope_id) if getattr(source, "scope_id", None) else None,
|
||||
)
|
||||
scope_id=str(source.scope_id) if getattr(source, "scope_id", None) else None)
|
||||
# config.yaml is canonical because it can persist the authenticated logical-target
|
||||
# provenance required by Relay after a restart.
|
||||
try:
|
||||
@@ -649,9 +609,11 @@ class GatewaySlashCommandsMixin(
|
||||
def _set_mode(mode: str) -> None:
|
||||
self._voice_mode[voice_key] = mode
|
||||
self._save_voice_modes()
|
||||
if adapter and mode == "off":
|
||||
if not adapter:
|
||||
return
|
||||
if mode == "off":
|
||||
self._set_adapter_auto_tts_disabled(adapter, chat_id, disabled=True)
|
||||
elif adapter:
|
||||
else:
|
||||
self._set_adapter_auto_tts_enabled(adapter, chat_id, enabled=True)
|
||||
|
||||
if args in _VOICE_MODE_BY_ARG:
|
||||
@@ -724,10 +686,8 @@ class GatewaySlashCommandsMixin(
|
||||
return msg
|
||||
|
||||
async def _handle_diff_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /diff — show git changes in the working directory.
|
||||
|
||||
Diff body is truncated hard here (chat is not a pager); platform senders clamp further.
|
||||
"""
|
||||
"""Handle /diff — show git changes in the working directory. Diff body is truncated hard
|
||||
here (chat is not a pager); platform senders clamp further."""
|
||||
args = [a.lower() for a in event.get_command_args().strip().split()]
|
||||
stat_only = bool({"--stat", "stat"} & set(args))
|
||||
mode = "working"
|
||||
@@ -796,9 +756,7 @@ class GatewaySlashCommandsMixin(
|
||||
self._track_background_task(self._run_background_task(
|
||||
prompt, event.source, task_id, event_message_id=self._reply_anchor_for_event(event),
|
||||
# Forward image/audio attachments so the background agent can see them.
|
||||
media_urls=list(event.media_urls) if event.media_urls else [],
|
||||
media_types=list(event.media_types) if event.media_types else [],
|
||||
))
|
||||
media_urls=list(event.media_urls or []), media_types=list(event.media_types or [])))
|
||||
return t("gateway.background.started", preview=_preview(prompt), task_id=task_id)
|
||||
|
||||
async def _handle_btw_command(self, event: MessageEvent) -> str:
|
||||
@@ -837,8 +795,7 @@ class GatewaySlashCommandsMixin(
|
||||
try:
|
||||
answer = await asyncio.to_thread(
|
||||
answer_side_question, question, history_snapshot,
|
||||
parent_agent=parent_agent, main_runtime=main_runtime,
|
||||
)
|
||||
parent_agent=parent_agent, main_runtime=main_runtime)
|
||||
reply = t("gateway.btw.answer", preview=preview, answer=answer or "")
|
||||
except Exception as e:
|
||||
logger.warning("/btw side question failed: %s", e)
|
||||
@@ -859,8 +816,7 @@ class GatewaySlashCommandsMixin(
|
||||
# the store persists to the same MEMORY/USER.md and honors the configured char limits).
|
||||
out = handle_pending_subcommand(
|
||||
wa.MEMORY, event.get_command_args().strip().split(), memory_store=load_on_disk_store(),
|
||||
set_mode_fn=self._write_approval_setter("memory", event),
|
||||
)
|
||||
set_mode_fn=self._write_approval_setter("memory", event))
|
||||
return out if out is not None else (
|
||||
"Unknown /memory subcommand. Use: pending, approve <id>, reject <id>, approval <on|off>."
|
||||
)
|
||||
@@ -879,8 +835,7 @@ class GatewaySlashCommandsMixin(
|
||||
"Enable it with /skills approval on, then review staged "
|
||||
"writes here with /skills pending.")
|
||||
out = handle_pending_subcommand(
|
||||
wa.SKILLS, args, set_mode_fn=self._write_approval_setter("skills", event)
|
||||
)
|
||||
wa.SKILLS, args, set_mode_fn=self._write_approval_setter("skills", event))
|
||||
if out is None:
|
||||
return ("Unknown /skills subcommand on this platform. Use: pending, "
|
||||
"approve <id>, reject <id>, diff <id>, approval <on|off>. "
|
||||
@@ -928,9 +883,8 @@ class GatewaySlashCommandsMixin(
|
||||
config_path, platform_key = self._display_config_target(event)
|
||||
try:
|
||||
user_config = _load_gateway_config()
|
||||
gate_enabled = is_truthy_value(
|
||||
cfg_get(user_config, "display", "tool_progress_command"), default=False
|
||||
)
|
||||
gate_enabled = is_truthy_value(cfg_get(user_config, "display", "tool_progress_command"),
|
||||
default=False)
|
||||
except Exception:
|
||||
gate_enabled = False
|
||||
if not gate_enabled:
|
||||
@@ -956,15 +910,11 @@ class GatewaySlashCommandsMixin(
|
||||
mode = self._effective_busy_input_mode(event.source)
|
||||
behavior = _BUSY_MODE_BEHAVIOR.get(mode, _BUSY_MODE_BEHAVIOR["interrupt"])[0]
|
||||
return EphemeralReply(
|
||||
f"**Busy input mode: `{mode}`" + "\n"
|
||||
f"Messages while busy: _{behavior}_" + "\n"
|
||||
f"Change with `/busy queue`, `/busy steer`, or `/busy interrupt`."
|
||||
)
|
||||
|
||||
f"**Busy input mode: `{mode}`\nMessages while busy: _{behavior}_\n"
|
||||
f"Change with `/busy queue`, `/busy steer`, or `/busy interrupt`.")
|
||||
if arg not in _BUSY_MODE_BEHAVIOR:
|
||||
return EphemeralReply(
|
||||
f"Unknown mode `{arg}`. Use `/busy queue`, `/busy steer`, or `/busy interrupt`."
|
||||
)
|
||||
f"Unknown mode `{arg}`. Use `/busy queue`, `/busy steer`, or `/busy interrupt`.")
|
||||
|
||||
# Persist before mutate
|
||||
from cli import save_config_value
|
||||
@@ -984,8 +934,7 @@ class GatewaySlashCommandsMixin(
|
||||
if adapter is not None:
|
||||
adapter._busy_text_mode = self._effective_busy_text_mode(event.source)
|
||||
return EphemeralReply(
|
||||
f"Busy input mode set to **`{arg}`** (saved)." + "\n" f"_{_BUSY_MODE_BEHAVIOR[arg][1]}_"
|
||||
)
|
||||
f"Busy input mode set to **`{arg}`** (saved).\n_{_BUSY_MODE_BEHAVIOR[arg][1]}_")
|
||||
|
||||
async def _handle_footer_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /footer command — toggle the runtime-metadata footer."""
|
||||
@@ -1000,7 +949,6 @@ class GatewaySlashCommandsMixin(
|
||||
arg = parts[1].strip().lower() if len(parts) > 1 else ""
|
||||
except Exception:
|
||||
arg = ""
|
||||
|
||||
try:
|
||||
user_config: dict = _load_gateway_config()
|
||||
except Exception as e:
|
||||
@@ -1010,14 +958,11 @@ class GatewaySlashCommandsMixin(
|
||||
def _state(enabled: bool) -> str:
|
||||
return t("gateway.footer.state_on") if enabled else t("gateway.footer.state_off")
|
||||
if arg in {"status", "?"}:
|
||||
fields = ", ".join(effective.get("fields") or [])
|
||||
return t("gateway.footer.status", state=_state(effective["enabled"]), fields=fields,
|
||||
platform=platform_key)
|
||||
|
||||
return t("gateway.footer.status", state=_state(effective["enabled"]),
|
||||
fields=", ".join(effective.get("fields") or []), platform=platform_key)
|
||||
if arg and arg not in _FOOTER_STATE_BY_ARG:
|
||||
return t("gateway.footer.usage")
|
||||
new_state = _FOOTER_STATE_BY_ARG[arg] if arg else not effective["enabled"]
|
||||
|
||||
try:
|
||||
_nested_dict(user_config, "display", "runtime_footer")["enabled"] = new_state
|
||||
atomic_config_write(config_path, user_config)
|
||||
@@ -1029,8 +974,7 @@ class GatewaySlashCommandsMixin(
|
||||
# Show a preview using current agent state if available.
|
||||
preview = format_runtime_footer(
|
||||
model=_resolve_gateway_model(user_config) or None, context_tokens=0, context_length=None,
|
||||
fields=effective.get("fields") or ["model", "context_pct", "cwd"],
|
||||
)
|
||||
fields=effective.get("fields") or ["model", "context_pct", "cwd"])
|
||||
if preview:
|
||||
example = t("gateway.footer.example_line", preview=preview)
|
||||
return t("gateway.footer.saved", state=_state(new_state), example=example)
|
||||
@@ -1067,8 +1011,7 @@ class GatewaySlashCommandsMixin(
|
||||
return result
|
||||
return await self._request_slash_confirm(
|
||||
event=event, command="reload-mcp", title="/reload-mcp",
|
||||
message=t("gateway.reload_mcp.confirm_prompt"), handler=_on_confirm,
|
||||
)
|
||||
message=t("gateway.reload_mcp.confirm_prompt"), handler=_on_confirm)
|
||||
|
||||
async def _handle_reload_skills_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /reload-skills — rescan skills dir, queue a note for next turn. Skills are invoked at
|
||||
@@ -1093,9 +1036,8 @@ class GatewaySlashCommandsMixin(
|
||||
if inspect.isawaitable(maybe):
|
||||
await maybe
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Adapter %s refresh_skill_group raised: %s", getattr(adapter, "name", adapter), exc,
|
||||
)
|
||||
logger.warning("Adapter %s refresh_skill_group raised: %s",
|
||||
getattr(adapter, "name", adapter), exc)
|
||||
|
||||
lines = [t("gateway.reload_skills.header")]
|
||||
if not added and not removed:
|
||||
@@ -1104,9 +1046,8 @@ class GatewaySlashCommandsMixin(
|
||||
|
||||
def _fmt_line(item: dict) -> str:
|
||||
nm, desc = item.get("name", ""), item.get("description", "")
|
||||
if desc:
|
||||
return t("gateway.reload_skills.item_with_desc", name=nm, desc=desc)
|
||||
return t("gateway.reload_skills.item_no_desc", name=nm)
|
||||
return (t("gateway.reload_skills.item_with_desc", name=nm, desc=desc) if desc
|
||||
else t("gateway.reload_skills.item_no_desc", name=nm))
|
||||
|
||||
# Queue a one-shot note for the next user turn in this session too. Format matches how
|
||||
# the system prompt renders pre-existing skills (`` - name: description``) so the
|
||||
@@ -1114,8 +1055,7 @@ class GatewaySlashCommandsMixin(
|
||||
sections = ["[USER INITIATED SKILLS RELOAD:"]
|
||||
for i18n_key, note_header, items in (
|
||||
("gateway.reload_skills.added_header", "Added Skills:", added),
|
||||
("gateway.reload_skills.removed_header", "Removed Skills:", removed),
|
||||
):
|
||||
("gateway.reload_skills.removed_header", "Removed Skills:", removed)):
|
||||
if items:
|
||||
formatted = [_fmt_line(item) for item in items]
|
||||
lines += [t(i18n_key)] + formatted
|
||||
@@ -1141,13 +1081,9 @@ class GatewaySlashCommandsMixin(
|
||||
return reply.text
|
||||
bundles = reply.data["bundles"]
|
||||
if not bundles:
|
||||
return (
|
||||
"No skill bundles installed.\n"
|
||||
"Create one on the host with:\n"
|
||||
" `hermes bundles create <name> --skill <s1> --skill <s2>`\n"
|
||||
f"Directory: `{reply.data['dir']}`"
|
||||
)
|
||||
|
||||
return ("No skill bundles installed.\nCreate one on the host with:\n"
|
||||
" `hermes bundles create <name> --skill <s1> --skill <s2>`\n"
|
||||
f"Directory: `{reply.data['dir']}`")
|
||||
lines = [f"**Skill Bundles** ({len(bundles)} installed):", ""]
|
||||
for info in bundles:
|
||||
skills = info.get("skills", [])
|
||||
@@ -1171,12 +1107,10 @@ class GatewaySlashCommandsMixin(
|
||||
"""Handle /approve — unblock waiting agent thread(s). They block inside tools/approval.py;
|
||||
signalling the event resumes them so the command executes inline (same flow as the CLI)."""
|
||||
from tools.approval import resolve_gateway_approval
|
||||
session_key, stale = self._blocking_approval_or_stale(
|
||||
event, "gateway.approval_expired", "gateway.approve.no_pending"
|
||||
)
|
||||
session_key, stale = self._blocking_approval_or_stale(event, "gateway.approval_expired",
|
||||
"gateway.approve.no_pending")
|
||||
if stale:
|
||||
return stale
|
||||
|
||||
# Args: "all", "all session", "all always", "session", "always" ("always" beats "session").
|
||||
args = event.get_command_args().strip().lower().split()
|
||||
choices = {_APPROVE_CHOICE_BY_ARG[a] for a in args if a in _APPROVE_CHOICE_BY_ARG}
|
||||
@@ -1192,12 +1126,10 @@ class GatewaySlashCommandsMixin(
|
||||
"""Handle /deny — reject pending dangerous command(s) with a definitive BLOCKED result, as in
|
||||
the CLI. ``/deny`` denies the oldest; ``/deny all`` denies everything."""
|
||||
from tools.approval import resolve_gateway_approval
|
||||
session_key, stale = self._blocking_approval_or_stale(
|
||||
event, "gateway.deny.stale", "gateway.deny.no_pending"
|
||||
)
|
||||
session_key, stale = self._blocking_approval_or_stale(event, "gateway.deny.stale",
|
||||
"gateway.deny.no_pending")
|
||||
if stale:
|
||||
return stale
|
||||
|
||||
# A leading "all" denies every pending command; the rest (or the whole arg string without
|
||||
# "all") is the optional deny reason relayed to the agent, capped to a sane one-liner.
|
||||
raw_args = event.get_command_args().strip()
|
||||
@@ -1207,9 +1139,8 @@ class GatewaySlashCommandsMixin(
|
||||
count = resolve_gateway_approval(session_key, "deny", resolve_all=resolve_all, reason=reason or None)
|
||||
if not count:
|
||||
return t("gateway.deny.no_pending")
|
||||
logger.info(
|
||||
"User denied %d dangerous command(s) via /deny%s", count, " (with reason)" if reason else "",
|
||||
)
|
||||
logger.info("User denied %d dangerous command(s) via /deny%s", count,
|
||||
" (with reason)" if reason else "")
|
||||
key = "gateway.deny.denied" + ("_reason" if reason else "") + ("_plural" if count > 1 else "_singular")
|
||||
confirmation_text = t(key, count=count, reason=reason)
|
||||
return await self._deliver_approval_confirmation(event, confirmation_text, "deny")
|
||||
@@ -1217,13 +1148,11 @@ class GatewaySlashCommandsMixin(
|
||||
async def _handle_debug_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /debug — upload ONLY the summary (system info + log tails), never full logs, to
|
||||
protect privacy; ``hermes debug share`` from the CLI does full uploads."""
|
||||
from hermes_cli.debug import (
|
||||
_GATEWAY_PRIVACY_NOTICE, _best_effort_sweep_expired_pastes, _capture_dump, _schedule_auto_delete,
|
||||
collect_debug_report, upload_to_pastebin,
|
||||
)
|
||||
from hermes_cli.debug import (_GATEWAY_PRIVACY_NOTICE, _best_effort_sweep_expired_pastes,
|
||||
_capture_dump, _schedule_auto_delete, collect_debug_report,
|
||||
upload_to_pastebin)
|
||||
|
||||
# Run blocking I/O (dump capture, log reads, uploads) in a thread.
|
||||
def _collect_and_upload():
|
||||
def _collect_and_upload(): # blocking I/O (dump capture, log reads, uploads) -> thread
|
||||
_best_effort_sweep_expired_pastes()
|
||||
report = collect_debug_report(log_lines=200, dump_text=_capture_dump())
|
||||
try:
|
||||
@@ -1232,11 +1161,10 @@ class GatewaySlashCommandsMixin(
|
||||
return t("gateway.debug.upload_failed", error=exc)
|
||||
_schedule_auto_delete(list(urls.values())) # auto-deletion after 6 hours
|
||||
label_width = max(len(k) for k in urls)
|
||||
return "\n".join([
|
||||
_GATEWAY_PRIVACY_NOTICE, "", t("gateway.debug.header"), "",
|
||||
*(f"`{label:<{label_width}}` {url}" for label, url in urls.items()),
|
||||
"", t("gateway.debug.auto_delete"), t("gateway.debug.full_logs_hint"), t("gateway.debug.share_hint"),
|
||||
])
|
||||
return "\n".join([_GATEWAY_PRIVACY_NOTICE, "", t("gateway.debug.header"), "",
|
||||
*(f"`{label:<{label_width}}` {url}" for label, url in urls.items()),
|
||||
"", t("gateway.debug.auto_delete"), t("gateway.debug.full_logs_hint"),
|
||||
t("gateway.debug.share_hint")])
|
||||
|
||||
# _run_in_executor_with_context, not a bare hop: this collects the profile's logs/config off
|
||||
# ``get_hermes_home()`` and uploads them to a public paste. Losing the contextvar override
|
||||
@@ -1246,10 +1174,9 @@ class GatewaySlashCommandsMixin(
|
||||
async def _handle_update_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /update — spawn ``hermes update`` detached (``setsid``) so it survives the gateway
|
||||
restart it may trigger; marker files let this or the next gateway process notify the user."""
|
||||
from gateway.run import _hermes_home, _resolve_hermes_bin
|
||||
import json
|
||||
from gateway.run import _hermes_home, _resolve_hermes_bin
|
||||
from hermes_cli.config import is_managed, format_managed_message
|
||||
|
||||
# Block non-messaging platforms (API server, webhooks, ACP); plugin platforms with
|
||||
# allow_update_command=True are also allowed.
|
||||
src = event.source
|
||||
@@ -1274,8 +1201,7 @@ class GatewaySlashCommandsMixin(
|
||||
pending = {
|
||||
"platform": src.platform.value, "chat_id": src.chat_id, "chat_type": src.chat_type,
|
||||
"user_id": src.user_id, "session_key": self._session_key_for_source(src),
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
}
|
||||
"timestamp": datetime.now().isoformat()}
|
||||
pending.update({k: v for k, v in (("thread_id", src.thread_id), ("message_id", event.message_id)) if v})
|
||||
_tmp_pending = pending_path.with_suffix(".tmp")
|
||||
_tmp_pending.write_text(json.dumps(pending), encoding="utf-8")
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
"""Gateway slash commands that rotate, switch, fork or rewrite the session transcript:
|
||||
/new, /resume, /sessions, /branch, /title, /save, /undo, /retry, /topic, /compress.
|
||||
|
||||
Split out of ``gateway/slash_commands.py``; bound onto ``GatewayRunner`` through
|
||||
``GatewaySlashCommandsMixin``. Origin internals are imported lazily inside the bodies to avoid
|
||||
the import cycle.
|
||||
"""
|
||||
``GatewaySlashCommandsMixin``. Origin internals are imported lazily inside the bodies to avoid
|
||||
the import cycle."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -30,10 +28,9 @@ _RESET_CLEANUP_TIMEOUT_S = 30.0
|
||||
# chat_type values whose session key is per-user (DM-like), incl. the unknown/blank case.
|
||||
_DM_CHAT_TYPES = {"dm", "direct", "private", ""}
|
||||
|
||||
_BRANCH_COPIED_FIELDS = (
|
||||
"content", "tool_calls", "tool_call_id", "finish_reason", "reasoning", "reasoning_content",
|
||||
"reasoning_details", "codex_reasoning_items", "codex_message_items", "timestamp",
|
||||
)
|
||||
_BRANCH_COPIED_FIELDS = ("content", "tool_calls", "tool_call_id", "finish_reason", "reasoning",
|
||||
"reasoning_content", "reasoning_details", "codex_reasoning_items",
|
||||
"codex_message_items", "timestamp")
|
||||
|
||||
|
||||
def _sattr(obj, name: str) -> str:
|
||||
@@ -43,11 +40,9 @@ def _sattr(obj, name: str) -> str:
|
||||
|
||||
def _manual_compression_reply_lines(summary: dict, compressor, focus_topic) -> list[str]:
|
||||
"""Manual /compress confirmation lines, surfacing summariser/aux-model failures.
|
||||
|
||||
``_last_compress_aborted`` = no usable summary, messages unchanged. Provider exception text is
|
||||
``_last_compress_aborted`` = no usable summary, messages unchanged. Provider exception text is
|
||||
force-redacted at this UI boundary even when global redaction is off; an aux model recovered
|
||||
via main is an info note so the user can fix their config.
|
||||
"""
|
||||
via main is an info note so the user can fix their config."""
|
||||
lines = [f"🗜️ {summary['headline']}"]
|
||||
if focus_topic:
|
||||
lines.append(t("gateway.compress.focus_line", topic=focus_topic))
|
||||
@@ -62,11 +57,8 @@ def _manual_compression_reply_lines(summary: dict, compressor, focus_topic) -> l
|
||||
if getattr(compressor, "_last_compress_aborted", False):
|
||||
lines.append(t("gateway.compress.aborted", error=(summary_err or "unknown error")))
|
||||
elif aux_fail_model:
|
||||
lines.append(t(
|
||||
"gateway.compress.aux_failed",
|
||||
model=aux_fail_model,
|
||||
error=(getattr(compressor, "_last_aux_model_failure_error", None) or "unknown error"),
|
||||
))
|
||||
aux_err = getattr(compressor, "_last_aux_model_failure_error", None) or "unknown error"
|
||||
lines.append(t("gateway.compress.aux_failed", model=aux_fail_model, error=aux_err))
|
||||
return lines
|
||||
|
||||
|
||||
@@ -75,14 +67,10 @@ def _compress_preview_reply(history, partial: bool, keep_last, focus_topic, agg_
|
||||
from agent.model_metadata import estimate_request_tokens_rough
|
||||
from hermes_cli.partial_compress import summarize_compress_preview
|
||||
|
||||
pv_msgs = [
|
||||
{"role": m.get("role"), "content": m.get("content")}
|
||||
for m in history
|
||||
if m.get("role") in {"user", "assistant"} and m.get("content")
|
||||
]
|
||||
report = summarize_compress_preview(
|
||||
pv_msgs, partial, keep_last, focus_topic, estimate_request_tokens_rough(pv_msgs)
|
||||
)
|
||||
pv_msgs = [{"role": m.get("role"), "content": m.get("content")} for m in history
|
||||
if m.get("role") in {"user", "assistant"} and m.get("content")]
|
||||
report = summarize_compress_preview(pv_msgs, partial, keep_last, focus_topic,
|
||||
estimate_request_tokens_rough(pv_msgs))
|
||||
lines = [f"🗜️ {line}" for line in report["lines"]]
|
||||
if agg_note:
|
||||
lines.append(agg_note)
|
||||
@@ -125,58 +113,39 @@ class GatewaySessionCommandsMixin:
|
||||
|
||||
async def _cleanup_old_agent_for_reset(self, session_key: str) -> None:
|
||||
"""Close the old agent's tool resources (sandboxes, browsers, subprocesses) before eviction.
|
||||
|
||||
Blocking work on the event loop (confirm-button click) → offloaded with a bounded timeout.
|
||||
wait_for cancels the await, not the worker thread: a wedged teardown keeps running (or
|
||||
leaks); the reset proceeds either way.
|
||||
"""
|
||||
leaks); the reset proceeds either way."""
|
||||
_old_agent = self._cached_agent_for(session_key)
|
||||
if _old_agent is None:
|
||||
return
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
self._run_in_executor_with_context(self._cleanup_agent_resources, _old_agent),
|
||||
timeout=_RESET_CLEANUP_TIMEOUT_S,
|
||||
)
|
||||
timeout=_RESET_CLEANUP_TIMEOUT_S)
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(
|
||||
"Agent resource cleanup for session %s exceeded %ss during /new reset; proceeding with "
|
||||
"reset (the worker thread is left to finish on its own). (#35994)",
|
||||
session_key, _RESET_CLEANUP_TIMEOUT_S,
|
||||
)
|
||||
session_key, _RESET_CLEANUP_TIMEOUT_S)
|
||||
except Exception as cleanup_exc:
|
||||
logger.warning(
|
||||
"Agent resource cleanup for session %s failed during /new reset: %s (#35994)",
|
||||
session_key, cleanup_exc,
|
||||
)
|
||||
session_key, cleanup_exc)
|
||||
|
||||
async def _fire_session_reset_hooks(
|
||||
self, source: SessionSource, session_key: str, old_sid, new_sid
|
||||
) -> None:
|
||||
async def _fire_session_reset_hooks(self, source: SessionSource, session_key: str, old_sid,
|
||||
new_sid) -> None:
|
||||
"""Session-boundary hooks: plugin finalize (off-loop + bounded — trace exports can block
|
||||
arbitrarily), then session:end and session:reset."""
|
||||
platform_value = source.platform.value if source.platform else ""
|
||||
with contextlib.suppress(Exception):
|
||||
await self._finalize_session_off_loop(
|
||||
session_id=old_sid, platform=platform_value, reason="new_session",
|
||||
old_session_id=old_sid, new_session_id=new_sid,
|
||||
)
|
||||
old_session_id=old_sid, new_session_id=new_sid)
|
||||
hook_payload = {"platform": platform_value, "user_id": source.user_id, "session_key": session_key}
|
||||
await self.hooks.emit("session:end", dict(hook_payload))
|
||||
await self.hooks.emit("session:reset", dict(hook_payload))
|
||||
|
||||
def _invoke_session_reset_lifecycle_hook(self, source: SessionSource, old_sid, new_sid) -> None:
|
||||
"""Plugin on_session_reset hook (new session guaranteed to exist); best-effort."""
|
||||
try:
|
||||
from hermes_cli.lifecycle import invoke_hook as _invoke_hook
|
||||
_invoke_hook(
|
||||
"on_session_reset", session_id=new_sid,
|
||||
platform=source.platform.value if source.platform else "", reason="new_session",
|
||||
old_session_id=old_sid, new_session_id=new_sid,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def _handle_reset_command(self, event: MessageEvent) -> Union[str, EphemeralReply]:
|
||||
"""Handle /new or /reset command."""
|
||||
source = event.source
|
||||
@@ -195,22 +164,16 @@ class GatewaySessionCommandsMixin:
|
||||
self._clear_conversation_scope(session_key, reason="session_reset")
|
||||
# In-flight async delegations end WITH the conversation: once the id rotates their
|
||||
# completions have no live owner. Expire by durable id, routing key as legacy fallback.
|
||||
try:
|
||||
with contextlib.suppress(Exception):
|
||||
from tools.async_delegation import interrupt_for_session
|
||||
interrupt_for_session(
|
||||
session_key=session_key,
|
||||
parent_session_id=str(getattr(old_entry, "session_id", "") or ""),
|
||||
reason="session_reset",
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
interrupt_for_session(session_key=session_key, reason="session_reset",
|
||||
parent_session_id=str(getattr(old_entry, "session_id", "") or ""))
|
||||
_reset_process_scoped_tool_state()
|
||||
|
||||
new_entry = await self.async_session_store.reset_session(session_key)
|
||||
_old_sid = old_entry.session_id if old_entry else None
|
||||
await self._fire_session_reset_hooks(
|
||||
source, session_key, _old_sid, new_entry.session_id if new_entry else None
|
||||
)
|
||||
await self._fire_session_reset_hooks(source, session_key, _old_sid,
|
||||
new_entry.session_id if new_entry else None)
|
||||
# Scoped to the profile serving this source so a multiplexed /new banner reports the
|
||||
# profile's model, not the base config's.
|
||||
try:
|
||||
@@ -233,9 +196,15 @@ class GatewaySessionCommandsMixin:
|
||||
await asyncio.to_thread(self._record_telegram_topic_binding, source, new_entry)
|
||||
except Exception:
|
||||
logger.debug("Failed to rebind Telegram topic after /new", exc_info=True)
|
||||
self._invoke_session_reset_lifecycle_hook(
|
||||
source, _old_sid, new_entry.session_id if new_entry else None
|
||||
)
|
||||
_new_sid = new_entry.session_id if new_entry else None
|
||||
# Plugin on_session_reset hook (new session guaranteed to exist); best-effort.
|
||||
try:
|
||||
from hermes_cli.lifecycle import invoke_hook as _invoke_hook
|
||||
_invoke_hook("on_session_reset", session_id=_new_sid, reason="new_session",
|
||||
platform=source.platform.value if source.platform else "",
|
||||
old_session_id=_old_sid, new_session_id=_new_sid)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
from hermes_cli.tips import get_random_tip
|
||||
_tip_line = t("gateway.reset.tip", tip=get_random_tip())
|
||||
@@ -275,29 +244,21 @@ class GatewaySessionCommandsMixin:
|
||||
return getattr(entry, "origin", None) if entry is not None else None
|
||||
# Test doubles / older stores lack the public lookup; fail closed when nothing resolves.
|
||||
entries = getattr(self.session_store, "_entries", {}) or {}
|
||||
return next(
|
||||
(getattr(e, "origin", None) for e in entries.values() if getattr(e, "session_id", None) == session_id),
|
||||
None,
|
||||
)
|
||||
return next((getattr(e, "origin", None) for e in entries.values()
|
||||
if getattr(e, "session_id", None) == session_id), None)
|
||||
|
||||
@staticmethod
|
||||
def _same_matrix_room(current: SessionSource, origin: Optional[SessionSource]) -> bool:
|
||||
# thread_id is part of the session key, so another thread of the SAME room is a DIFFERENT
|
||||
# session; non-threaded rooms compare "" == "".
|
||||
return (
|
||||
origin is not None
|
||||
and origin.platform == Platform.MATRIX
|
||||
and current.platform == Platform.MATRIX
|
||||
and origin.chat_id == current.chat_id
|
||||
and _sattr(current, "thread_id") == _sattr(origin, "thread_id")
|
||||
)
|
||||
return (origin is not None and origin.platform == Platform.MATRIX
|
||||
and current.platform == Platform.MATRIX and origin.chat_id == current.chat_id
|
||||
and _sattr(current, "thread_id") == _sattr(origin, "thread_id"))
|
||||
|
||||
def _same_origin_chat(self, current: SessionSource, origin: Optional[SessionSource]) -> bool:
|
||||
"""Platform-agnostic counterpart to ``_same_matrix_room``.
|
||||
|
||||
Per-participant sessions must be participant-scoped here too, else a co-member could resume
|
||||
another member's live session (IDOR); only an explicitly shared group/thread shares.
|
||||
"""
|
||||
"""Platform-agnostic counterpart to ``_same_matrix_room``. Per-participant sessions must be
|
||||
participant-scoped here too, else a co-member could resume another member's live session
|
||||
(IDOR); only an explicitly shared group/thread shares."""
|
||||
if origin is None or current is None:
|
||||
return False
|
||||
if origin.platform != current.platform or origin.chat_id != current.chat_id:
|
||||
@@ -327,15 +288,12 @@ class GatewaySessionCommandsMixin:
|
||||
build_session_key's isolation rules so the guards stay in lock-step with the key."""
|
||||
return is_shared_multi_user_session(
|
||||
source, group_sessions_per_user=getattr(self.config, "group_sessions_per_user", True),
|
||||
thread_sessions_per_user=getattr(self.config, "thread_sessions_per_user", False),
|
||||
)
|
||||
thread_sessions_per_user=getattr(self.config, "thread_sessions_per_user", False))
|
||||
|
||||
def _resume_caller_is_admin(self, source: SessionSource) -> bool:
|
||||
"""Whether *source* is an EXPLICITLY-configured admin (cross-origin /resume, /sessions).
|
||||
|
||||
Stricter than ``SlashAccessPolicy.is_admin()``, which is True for every caller when slash
|
||||
gating is DISABLED — the default config would make everyone cross-origin-capable (IDOR).
|
||||
"""
|
||||
gating is DISABLED — the default config would make everyone cross-origin-capable (IDOR)."""
|
||||
try:
|
||||
from gateway.slash_access import policy_for_source
|
||||
policy = policy_for_source(self.config, source)
|
||||
@@ -346,11 +304,9 @@ class GatewaySessionCommandsMixin:
|
||||
|
||||
def _persisted_row_proves_owner(self, source: SessionSource, row: dict) -> bool:
|
||||
"""Whether a persisted (inactive) session *row* provably belongs to *source*'s session key.
|
||||
|
||||
Rows once stored only source + user_id, so the persisted chat/thread origin is compared too
|
||||
and legacy NULL rows fail closed. The table has no user_id_alt column, so an alt-keyed
|
||||
(Signal/Feishu) caller is never proven by user_id alone (CWE-639).
|
||||
"""
|
||||
and legacy NULL rows fail closed. The table has no user_id_alt column, so an alt-keyed
|
||||
(Signal/Feishu) caller is never proven by user_id alone (CWE-639)."""
|
||||
caller_src = source.platform.value if source.platform else None
|
||||
row_src = row.get("source")
|
||||
caller_uid = _sattr(source, "user_id")
|
||||
@@ -380,14 +336,11 @@ class GatewaySessionCommandsMixin:
|
||||
return False
|
||||
return bool(row_uid) and row_uid == caller_uid
|
||||
|
||||
async def _resume_target_allowed(
|
||||
self, source: SessionSource, target_id: str, allow_override: bool = False
|
||||
) -> bool:
|
||||
"""Whether *source* may resume session *target_id* (IDOR guard for every adapter).
|
||||
|
||||
The live origin decides when the target is active; otherwise the DB row must PROVE
|
||||
ownership or fail closed. Admin ``--all`` bypasses.
|
||||
"""
|
||||
async def _resume_target_allowed(self, source: SessionSource, target_id: str,
|
||||
allow_override: bool = False) -> bool:
|
||||
"""Whether *source* may resume session *target_id* (IDOR guard for every adapter). The live
|
||||
origin decides when the target is active; otherwise the DB row must PROVE ownership or fail
|
||||
closed. Admin ``--all`` bypasses."""
|
||||
if allow_override and self._resume_caller_is_admin(source):
|
||||
return True
|
||||
# Only a real SessionSource origin decides; unresolvable/error falls through to DB scoping.
|
||||
@@ -403,9 +356,7 @@ class GatewaySessionCommandsMixin:
|
||||
return False
|
||||
return self._persisted_row_proves_owner(source, row)
|
||||
|
||||
async def _resume_row_visible(
|
||||
self, source: SessionSource, row: dict, allow_all: bool
|
||||
) -> bool:
|
||||
async def _resume_row_visible(self, source: SessionSource, row: dict, allow_all: bool) -> bool:
|
||||
"""Whether a listing *row* belongs to the caller's origin (blocks cross-origin enumeration of
|
||||
ids/previews); Matrix is room-scoped, ``--all`` needs a configured admin everywhere."""
|
||||
if allow_all and self._resume_caller_is_admin(source):
|
||||
@@ -422,19 +373,14 @@ class GatewaySessionCommandsMixin:
|
||||
# The canonical projection skips bookkeeping rows (role=user + display_kind) and pure
|
||||
# handoffs while still recognizing a real ask embedded in a compaction carrier.
|
||||
from agent.context_compressor import (
|
||||
history_before_user_originated_turn,
|
||||
retryable_user_text,
|
||||
split_user_originated_turn,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
history_before_user_originated_turn, retryable_user_text, split_user_originated_turn,
|
||||
user_originated_turn_view)
|
||||
|
||||
source = event.source
|
||||
session_entry = await self.async_session_store.get_or_create_session(source)
|
||||
history = await self.async_session_store.load_transcript(session_entry.session_id)
|
||||
last_user_idx = next(
|
||||
(i for i in range(len(history) - 1, -1, -1) if user_originated_turn_view(history[i]) is not None),
|
||||
None,
|
||||
)
|
||||
last_user_idx = next((i for i in range(len(history) - 1, -1, -1)
|
||||
if user_originated_turn_view(history[i]) is not None), None)
|
||||
if last_user_idx is None:
|
||||
return t("gateway.retry.no_previous")
|
||||
# Resolve text + scaffold-preserving prefix BEFORE any write; messaging retries cannot
|
||||
@@ -452,8 +398,7 @@ class GatewaySessionCommandsMixin:
|
||||
# on the same snapshot so a concurrent newer turn is never removed for stale text.
|
||||
try:
|
||||
rewind_result = await self.async_session_store.rewind_session(
|
||||
session_entry.session_id, 1, require_retryable_composite=True,
|
||||
)
|
||||
session_entry.session_id, 1, require_retryable_composite=True)
|
||||
except ValueError as exc:
|
||||
return f"Cannot retry that message safely: {exc}"
|
||||
if rewind_result is None:
|
||||
@@ -461,15 +406,12 @@ class GatewaySessionCommandsMixin:
|
||||
last_user_msg = rewind_result["target_text"]
|
||||
# active_only preserves the active=0/compacted=1 archive left by in-place compaction.
|
||||
elif not await self.async_session_store.rewrite_transcript(
|
||||
session_entry.session_id, truncated, active_only=True, reject_active_turn_lease=True,
|
||||
):
|
||||
session_entry.session_id, truncated, active_only=True, reject_active_turn_lease=True):
|
||||
return "Retry failed; transcript was not changed."
|
||||
session_entry.last_prompt_tokens = 0 # transcript was truncated
|
||||
retry_event = MessageEvent(
|
||||
return await self._handle_message(MessageEvent(
|
||||
text=last_user_msg, message_type=MessageType.TEXT, source=source,
|
||||
raw_message=event.raw_message, channel_prompt=event.channel_prompt,
|
||||
)
|
||||
return await self._handle_message(retry_event)
|
||||
raw_message=event.raw_message, channel_prompt=event.channel_prompt))
|
||||
|
||||
async def _handle_undo_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /undo [N] — back up N user turns (default 1), soft-deleting the truncated rows and
|
||||
@@ -494,7 +436,8 @@ class GatewaySessionCommandsMixin:
|
||||
logger.debug("undo: cached-agent eviction skipped: %s", e)
|
||||
target_text = result["target_text"]
|
||||
preview = target_text[:200] + "..." if len(target_text) > 200 else target_text
|
||||
return t("gateway.undo.removed", turns=result["turns_undone"], count=result["rewound_count"], preview=preview)
|
||||
return t("gateway.undo.removed", turns=result["turns_undone"],
|
||||
count=result["rewound_count"], preview=preview)
|
||||
|
||||
# --------------------------------------------------------------------- /compress
|
||||
|
||||
@@ -519,8 +462,7 @@ class GatewaySessionCommandsMixin:
|
||||
return (
|
||||
"🗜️ Nothing to compact: this session runs on the Codex app-server runtime, whose "
|
||||
"context lives in a Codex-owned thread that only exists while the agent is active. "
|
||||
"Send a message first, then /compress — or /reset to start fresh."
|
||||
)
|
||||
"Send a message first, then /compress — or /reset to start fresh.")
|
||||
compressor = getattr(agent, "context_compressor", None)
|
||||
count_before = getattr(compressor, "compression_count", 0)
|
||||
try:
|
||||
@@ -530,12 +472,10 @@ class GatewaySessionCommandsMixin:
|
||||
if getattr(compressor, "compression_count", 0) > count_before:
|
||||
return (
|
||||
"🗜️ Codex app-server thread compacted (thread/compact). The transcript mirror is "
|
||||
"unchanged by design — the app-server now carries the compacted context."
|
||||
)
|
||||
"unchanged by design — the app-server now carries the compacted context.")
|
||||
return (
|
||||
"⚠️ Codex app-server compaction did not complete — the thread is unchanged. Check the "
|
||||
"app-server logs, retry /compress, or /reset for a clean session."
|
||||
)
|
||||
"app-server logs, retry /compress, or /reset for a clean session.")
|
||||
|
||||
async def _handle_compress_command_inner(self, event: MessageEvent) -> str:
|
||||
"""Handle /compress -- manually compress conversation context; ``/compress <focus>`` tells
|
||||
@@ -561,25 +501,21 @@ class GatewaySessionCommandsMixin:
|
||||
if _preview:
|
||||
return _compress_preview_reply(history, partial, keep_last, focus_topic, _agg_note)
|
||||
try:
|
||||
return await self._run_manual_compression(
|
||||
source, session_entry, history, partial, keep_last, focus_topic
|
||||
)
|
||||
return await self._run_manual_compression(source, session_entry, history, partial,
|
||||
keep_last, focus_topic)
|
||||
except Exception as e:
|
||||
logger.warning("Manual compress failed: %s", e)
|
||||
return t("gateway.compress.failed", error=e)
|
||||
|
||||
async def _run_manual_compression(
|
||||
self, source, session_entry, history: list, partial: bool, keep_last, focus_topic
|
||||
) -> str:
|
||||
async def _run_manual_compression(self, source, session_entry, history: list, partial: bool,
|
||||
keep_last, focus_topic) -> str:
|
||||
"""Build a temporary agent, compress the transcript, persist, and describe the outcome."""
|
||||
from agent.conversation_compression import finalize_context_engine_compression_notification
|
||||
from agent.manual_compression_feedback import summarize_manual_compression
|
||||
from agent.model_metadata import estimate_request_tokens_rough
|
||||
from gateway.run import _platform_config_key
|
||||
from hermes_cli.partial_compress import (
|
||||
rejoin_compressed_head_and_tail,
|
||||
split_history_for_partial_compress,
|
||||
)
|
||||
from hermes_cli.partial_compress import (rejoin_compressed_head_and_tail,
|
||||
split_history_for_partial_compress)
|
||||
|
||||
session_key = self._session_key_for_source(source)
|
||||
# Platform + stable gateway session key bind this agent (for external context engines) to
|
||||
@@ -622,9 +558,7 @@ class GatewaySessionCommandsMixin:
|
||||
compressed, _ = await self._run_in_executor_with_context(
|
||||
lambda: tmp_agent._compress_context(
|
||||
head, "", approx_tokens=approx_tokens, focus_topic=focus_topic, force=True,
|
||||
defer_context_engine_notification=True,
|
||||
)
|
||||
)
|
||||
defer_context_engine_notification=True))
|
||||
# A held compression lock returns unchanged; say so instead of the misleading no-op text.
|
||||
_lock_skipped = getattr(tmp_agent, "_compression_skipped_due_to_lock", None)
|
||||
if _lock_skipped is True or isinstance(_lock_skipped, str):
|
||||
@@ -635,9 +569,8 @@ class GatewaySessionCommandsMixin:
|
||||
await self._persist_manual_compression(tmp_agent, session_entry, source, compressed)
|
||||
finalize_context_engine_compression_notification(tmp_agent, committed=True)
|
||||
new_tokens = estimate_request_tokens_rough(compressed, system_prompt=_sys_prompt, tools=_tools)
|
||||
summary = summarize_manual_compression(
|
||||
msgs, compressed, approx_tokens, new_tokens, compression_state=compressor,
|
||||
)
|
||||
summary = summarize_manual_compression(msgs, compressed, approx_tokens, new_tokens,
|
||||
compression_state=compressor)
|
||||
finally:
|
||||
finalize_context_engine_compression_notification(tmp_agent, committed=False)
|
||||
self._evict_cached_agent(session_key) # next turn rebuilds the prompt from current files
|
||||
@@ -663,20 +596,17 @@ class GatewaySessionCommandsMixin:
|
||||
logger.warning(
|
||||
"Manual compression could not restore the system prompt for session %s: %s. "
|
||||
"Preserving an empty prompt so the live turn rebuilds it with its configured "
|
||||
"providers.", session_id, exc, exc_info=True,
|
||||
)
|
||||
"providers.", session_id, exc, exc_info=True)
|
||||
|
||||
# compression.checkpoint_required needs the memory provider loaded so _compress_context()
|
||||
# can write the pre-compression checkpoint; otherwise keep the fast path (no provider init).
|
||||
_checkpoint_required = _is_truthy(
|
||||
((_load_cfg() or {}).get("compression") or {}).get("checkpoint_required"),
|
||||
default=False,
|
||||
)
|
||||
tmp_agent = AIAgent(
|
||||
**runtime_kwargs, model=model, max_iterations=4, quiet_mode=True,
|
||||
skip_memory=not _checkpoint_required, enabled_toolsets=["memory"],
|
||||
session_id=session_id, session_db=getattr(self._session_db, "_db", self._session_db),
|
||||
)
|
||||
default=False)
|
||||
tmp_agent = AIAgent(**runtime_kwargs, model=model, max_iterations=4, quiet_mode=True,
|
||||
skip_memory=not _checkpoint_required, enabled_toolsets=["memory"],
|
||||
session_id=session_id,
|
||||
session_db=getattr(self._session_db, "_db", self._session_db))
|
||||
_seed_hygiene_system_prompt(tmp_agent, session_row)
|
||||
# Real platform during construction (context engines bind correctly); afterwards a prompt
|
||||
# rebuilt by compression is stamped as the provider-less fallback, stale for the next turn.
|
||||
@@ -687,29 +617,24 @@ class GatewaySessionCommandsMixin:
|
||||
return tmp_agent
|
||||
|
||||
async def _persist_manual_compression(self, tmp_agent, session_entry, source, compressed) -> None:
|
||||
"""Commit a manual /compress result to the session store.
|
||||
|
||||
Rotation (new continuation id) writes the compressed messages into the NEW session so the
|
||||
original stays searchable; persist BEFORE repointing so a failed write is fatal and old
|
||||
history stays reachable. In-place compaction already archived + inserted rows, and a rewrite
|
||||
would DELETE the archive; an unchanged id without in-place means rotation FAILED.
|
||||
"""
|
||||
"""Commit a manual /compress result to the session store. Rotation (new continuation id)
|
||||
writes the compressed messages into the NEW session so the original stays searchable;
|
||||
persist BEFORE repointing so a failed write is fatal and old history stays reachable.
|
||||
In-place compaction already archived + inserted rows, and a rewrite would DELETE the
|
||||
archive; an unchanged id without in-place means rotation FAILED."""
|
||||
new_session_id = tmp_agent.session_id
|
||||
if new_session_id != session_entry.session_id:
|
||||
if not await self.async_session_store.rewrite_transcript(new_session_id, compressed):
|
||||
raise RuntimeError(
|
||||
f"failed to persist compressed transcript for session {new_session_id}"
|
||||
)
|
||||
f"failed to persist compressed transcript for session {new_session_id}")
|
||||
session_entry.session_id = new_session_id
|
||||
await self.async_session_store._save()
|
||||
await asyncio.to_thread(
|
||||
self._sync_telegram_topic_binding, source, session_entry, reason="compress-command",
|
||||
)
|
||||
await asyncio.to_thread(self._sync_telegram_topic_binding, source, session_entry,
|
||||
reason="compress-command")
|
||||
elif not getattr(tmp_agent, "_last_compaction_in_place", False):
|
||||
logger.warning(
|
||||
"Manual /compress: session rotation did not occur (session_id unchanged) and in-place "
|
||||
"mode is off — preserving original transcript instead of overwriting it (#44794)."
|
||||
)
|
||||
"mode is off — preserving original transcript instead of overwriting it (#44794).")
|
||||
await self.async_session_store.update_session(session_entry.session_key, last_prompt_tokens=0)
|
||||
|
||||
# ------------------------------------------------------------------------ /topic
|
||||
@@ -756,8 +681,7 @@ class GatewaySessionCommandsMixin:
|
||||
await self._session_db.enable_telegram_topic_mode(
|
||||
chat_id=str(source.chat_id), user_id=str(source.user_id), profile_name=profile_name,
|
||||
has_topics_enabled=capabilities.get("has_topics_enabled"),
|
||||
allows_users_to_create_topics=capabilities.get("allows_users_to_create_topics"),
|
||||
)
|
||||
allows_users_to_create_topics=capabilities.get("allows_users_to_create_topics"))
|
||||
except Exception as exc:
|
||||
logger.exception("Failed to enable Telegram topic mode")
|
||||
return t("gateway.topic.enable_failed", error=exc)
|
||||
@@ -768,8 +692,7 @@ class GatewaySessionCommandsMixin:
|
||||
try:
|
||||
binding = await self._session_db.get_telegram_topic_binding(
|
||||
chat_id=str(source.chat_id), thread_id=str(source.thread_id),
|
||||
profile_name=profile_name,
|
||||
)
|
||||
profile_name=profile_name)
|
||||
except Exception:
|
||||
logger.debug("Failed to read Telegram topic binding", exc_info=True)
|
||||
binding = None
|
||||
@@ -780,10 +703,8 @@ class GatewaySessionCommandsMixin:
|
||||
title = await self._session_db.get_session_title(session_id)
|
||||
except Exception:
|
||||
title = None
|
||||
return t(
|
||||
"gateway.topic.bound_status", label=title or t("gateway.topic.untitled_session"),
|
||||
session_id=session_id,
|
||||
)
|
||||
return t("gateway.topic.bound_status", label=title or t("gateway.topic.untitled_session"),
|
||||
session_id=session_id)
|
||||
|
||||
# ------------------------------------------------------------------ /save, /title
|
||||
|
||||
@@ -791,11 +712,7 @@ class GatewaySessionCommandsMixin:
|
||||
"""Handle /save — export the current session and send it as a document."""
|
||||
import tempfile
|
||||
from hermes_cli.session_export import (
|
||||
SAVE_USAGE,
|
||||
default_save_filename,
|
||||
normalize_save_format,
|
||||
render_session_for_save,
|
||||
)
|
||||
SAVE_USAGE, default_save_filename, normalize_save_format, render_session_for_save)
|
||||
|
||||
parts = event.get_command_args().split()
|
||||
redact = bool(parts) and parts[-1].lower() in ("redact", "--redact")
|
||||
@@ -835,10 +752,8 @@ class GatewaySessionCommandsMixin:
|
||||
adapter = self.get_adapter(source.platform)
|
||||
if not adapter:
|
||||
return "Platform adapter not found to send the document."
|
||||
await adapter.send_document(
|
||||
chat_id=source.chat_id, file_path=temp_path, caption=f"Session export: {filename}",
|
||||
file_name=filename,
|
||||
)
|
||||
await adapter.send_document(chat_id=source.chat_id, file_path=temp_path,
|
||||
caption=f"Session export: {filename}", file_name=filename)
|
||||
return "Export complete."
|
||||
except Exception as e:
|
||||
logger.warning("Session /save failed: %s", e)
|
||||
@@ -865,8 +780,7 @@ class GatewaySessionCommandsMixin:
|
||||
session_id=session_id,
|
||||
source=source.platform.value if source.platform else "unknown",
|
||||
user_id=source.user_id, chat_id=source.chat_id, chat_type=source.chat_type,
|
||||
thread_id=source.thread_id,
|
||||
)
|
||||
thread_id=source.thread_id)
|
||||
title_arg = event.get_command_args().strip()
|
||||
if not title_arg:
|
||||
title = await self._session_db.get_session_title(session_id)
|
||||
@@ -899,8 +813,7 @@ class GatewaySessionCommandsMixin:
|
||||
widen = allow_all and self._resume_caller_is_admin(source)
|
||||
sessions = await self._session_db.list_sessions_rich(
|
||||
source=source.platform.value if source.platform else None,
|
||||
session_key=None if widen else session_key, limit=10,
|
||||
)
|
||||
session_key=None if widen else session_key, limit=10)
|
||||
titled = [s for s in sessions if s.get("title")][:10]
|
||||
return [s for s in titled if await self._resume_row_visible(source, s, allow_all)]
|
||||
|
||||
@@ -929,9 +842,8 @@ class GatewaySessionCommandsMixin:
|
||||
logger.debug("Failed to resolve resume continuation for %s: %s", target_id, e)
|
||||
return target_id, name
|
||||
|
||||
async def _resume_access_denied_reply(
|
||||
self, source, target_id: str, name: str, allow_all: bool, allow_cross_room: bool
|
||||
) -> Optional[str]:
|
||||
async def _resume_access_denied_reply(self, source, target_id: str, name: str, allow_all: bool,
|
||||
allow_cross_room: bool) -> Optional[str]:
|
||||
"""IDOR guard: a session id/title is a routing handle, not authority — bind /resume to the
|
||||
caller's own room (Matrix) or platform/user/chat (other adapters)."""
|
||||
if source.platform == Platform.MATRIX:
|
||||
@@ -940,10 +852,8 @@ class GatewaySessionCommandsMixin:
|
||||
return None
|
||||
if target_origin is None:
|
||||
return t("gateway.resume.matrix_blocked_no_origin", name=name)
|
||||
return t(
|
||||
"gateway.resume.matrix_blocked_other_room",
|
||||
room=target_origin.chat_name or target_origin.chat_id, name=name,
|
||||
)
|
||||
return t("gateway.resume.matrix_blocked_other_room", name=name,
|
||||
room=target_origin.chat_name or target_origin.chat_id)
|
||||
if await self._resume_target_allowed(source, target_id, allow_override=(allow_all or allow_cross_room)):
|
||||
return None
|
||||
return t("gateway.resume.blocked_not_owner", name=name)
|
||||
@@ -993,10 +903,8 @@ class GatewaySessionCommandsMixin:
|
||||
msg_count = len([m for m in history if m.get("role") == "user"]) if history else 0
|
||||
if source.platform == Platform.MATRIX and allow_cross_room:
|
||||
msg_part = f" ({msg_count} message{'s' if msg_count != 1 else ''})" if msg_count else ""
|
||||
return t(
|
||||
"gateway.resume.matrix_cross_room_success", title=title,
|
||||
room=source.chat_name or source.chat_id, msg_part=msg_part,
|
||||
)
|
||||
return t("gateway.resume.matrix_cross_room_success", title=title,
|
||||
room=source.chat_name or source.chat_id, msg_part=msg_part)
|
||||
if not msg_count:
|
||||
return t("gateway.resume.resumed_no_count", title=title)
|
||||
if msg_count == 1:
|
||||
@@ -1034,15 +942,10 @@ class GatewaySessionCommandsMixin:
|
||||
if not self._session_db:
|
||||
return self._session_db_unavailable_reply()
|
||||
from hermes_cli.session_listing import (
|
||||
format_gateway_session_listing,
|
||||
parse_session_listing_args,
|
||||
query_session_listing,
|
||||
)
|
||||
|
||||
format_gateway_session_listing, parse_session_listing_args, query_session_listing)
|
||||
try:
|
||||
include_all, include_unnamed, target, search_query = parse_session_listing_args(
|
||||
event.get_command_args().strip()
|
||||
)
|
||||
event.get_command_args().strip())
|
||||
except ValueError as exc:
|
||||
return t("gateway.resume.parse_error", error=exc)
|
||||
if search_query == "":
|
||||
@@ -1059,19 +962,14 @@ class GatewaySessionCommandsMixin:
|
||||
scope_notice = "_Note: `all` (cross-chat listing) requires a configured admin; showing this chat's sessions only._"
|
||||
current_entry = await self.async_session_store.get_or_create_session(source)
|
||||
rows = await asyncio.to_thread(
|
||||
query_session_listing,
|
||||
getattr(self._session_db, "_db", self._session_db),
|
||||
query_session_listing, getattr(self._session_db, "_db", self._session_db),
|
||||
source=source.platform.value if source.platform else None,
|
||||
session_key=None if cross_origin else session_key,
|
||||
current_session_id=current_entry.session_id,
|
||||
include_current_session=True,
|
||||
include_all_sources=cross_origin,
|
||||
include_unnamed=include_unnamed,
|
||||
current_session_id=current_entry.session_id, include_current_session=True,
|
||||
include_all_sources=cross_origin, include_unnamed=include_unnamed,
|
||||
search_query=search_query,
|
||||
# Search filters in SQL: over-fetch so origin-invisible matches don't consume the page.
|
||||
limit=50 if search_query else 10,
|
||||
exclude_sources=["tool"],
|
||||
)
|
||||
limit=50 if search_query else 10, exclude_sources=["tool"])
|
||||
if not cross_origin:
|
||||
rows = [row for row in rows if await self._resume_row_visible(source, row, allow_all=False)]
|
||||
rows = rows[:10]
|
||||
@@ -1079,7 +977,8 @@ class GatewaySessionCommandsMixin:
|
||||
title = f"Sessions matching “{search_query}”"
|
||||
else:
|
||||
title = "Sessions" if include_unnamed else "Named Sessions"
|
||||
return format_gateway_session_listing(rows, include_source=cross_origin, title=title, notice=scope_notice)
|
||||
return format_gateway_session_listing(rows, include_source=cross_origin, title=title,
|
||||
notice=scope_notice)
|
||||
|
||||
# ----------------------------------------------------------------------- /branch
|
||||
|
||||
@@ -1117,15 +1016,10 @@ class GatewaySessionCommandsMixin:
|
||||
source=source.platform.value if source.platform else "gateway",
|
||||
model=(self.config.get("model", {}) or {}).get("default") if isinstance(self.config, dict) else None,
|
||||
model_config={"_branched_from": parent_session_id},
|
||||
parent_session_id=parent_session_id,
|
||||
user_id=source.user_id,
|
||||
session_key=session_key,
|
||||
chat_id=source.chat_id,
|
||||
chat_type=source.chat_type,
|
||||
thread_id=source.thread_id,
|
||||
origin_json=_branch_origin_json,
|
||||
display_name=current_entry.display_name,
|
||||
)
|
||||
parent_session_id=parent_session_id, user_id=source.user_id,
|
||||
session_key=session_key, chat_id=source.chat_id, chat_type=source.chat_type,
|
||||
thread_id=source.thread_id, origin_json=_branch_origin_json,
|
||||
display_name=current_entry.display_name)
|
||||
except Exception as e:
|
||||
logger.error("Failed to create branch session: %s", e)
|
||||
return t("gateway.branch.create_failed", error=e)
|
||||
@@ -1133,8 +1027,7 @@ class GatewaySessionCommandsMixin:
|
||||
# Chunked transactions; best-effort — a failed copy still yields a usable (partial) branch.
|
||||
with contextlib.suppress(Exception):
|
||||
await self._session_db.append_messages_batch(
|
||||
new_session_id, [_branch_row(msg) for msg in history], chunk_rows=500,
|
||||
)
|
||||
new_session_id, [_branch_row(msg) for msg in history], chunk_rows=500)
|
||||
with contextlib.suppress(Exception):
|
||||
await self._session_db.set_session_title(new_session_id, branch_title)
|
||||
new_entry = await self.async_session_store.switch_session(session_key, new_session_id)
|
||||
|
||||
@@ -26,16 +26,13 @@ from gateway.platforms.base import _custom_unit_to_cp
|
||||
from gateway.config import (
|
||||
DEFAULT_STREAMING_EDIT_INTERVAL as _DEFAULT_STREAMING_EDIT_INTERVAL,
|
||||
DEFAULT_STREAMING_BUFFER_THRESHOLD as _DEFAULT_STREAMING_BUFFER_THRESHOLD,
|
||||
DEFAULT_STREAMING_CURSOR as _DEFAULT_STREAMING_CURSOR,
|
||||
)
|
||||
DEFAULT_STREAMING_CURSOR as _DEFAULT_STREAMING_CURSOR)
|
||||
from gateway.response_filters import (
|
||||
is_intentional_silence_response as _is_intentional_silence_response,
|
||||
is_partial_silence_marker as _is_partial_silence_marker,
|
||||
)
|
||||
is_partial_silence_marker as _is_partial_silence_marker)
|
||||
from gateway.stream_consumer_fences import ( # noqa: F401 (re-exported)
|
||||
ensure_closed_code_fences,
|
||||
escape_code_fences_for_display,
|
||||
)
|
||||
escape_code_fences_for_display)
|
||||
from gateway.stream_consumer_transport import StreamTransportMixin
|
||||
from gateway.stream_consumer_fallback import StreamFallbackMixin
|
||||
from gateway.stream_consumer_think import StreamThinkFilterMixin
|
||||
@@ -56,6 +53,7 @@ _FINAL_TEXT = object()
|
||||
_FLUSH = object()
|
||||
_APPROVAL_BOUNDARY = object()
|
||||
_REOPEN_SEED = object()
|
||||
_FUTURE_TYPES = (asyncio.Future, concurrent.futures.Future)
|
||||
|
||||
# Boundary finalize text when nothing has accumulated yet (overridable per boundary).
|
||||
_DEFAULT_BOUNDARY_PLACEHOLDER = "⏸ 等待审批中..."
|
||||
@@ -97,11 +95,7 @@ class _Tick:
|
||||
return not self.got_done and not self.got_segment_break and self.commentary_text is None
|
||||
|
||||
|
||||
class GatewayStreamConsumer(
|
||||
StreamTransportMixin,
|
||||
StreamFallbackMixin,
|
||||
StreamThinkFilterMixin,
|
||||
):
|
||||
class GatewayStreamConsumer(StreamTransportMixin, StreamFallbackMixin, StreamThinkFilterMixin):
|
||||
"""Async consumer that progressively edits a platform message with streamed tokens.
|
||||
Usage: ``agent.stream_delta_callback = consumer.on_delta``; ``create_task(consumer.run())``;
|
||||
after the agent finishes ``consumer.finish()`` then ``await task`` for the final edit."""
|
||||
@@ -123,8 +117,7 @@ class GatewayStreamConsumer(
|
||||
on_new_message: Optional[callable] = None,
|
||||
on_before_finalize: Optional[Callable[[], Any]] = None,
|
||||
initial_reply_to_id: Optional[str] = None,
|
||||
run_still_current: Optional[Callable[[], bool]] = None,
|
||||
):
|
||||
run_still_current: Optional[Callable[[], bool]] = None):
|
||||
self.adapter = adapter
|
||||
self.chat_id = chat_id
|
||||
self.cfg = config or StreamConsumerConfig()
|
||||
@@ -139,9 +132,7 @@ class GatewayStreamConsumer(
|
||||
self._run_still_current = run_still_current or (lambda: True)
|
||||
# Only platforms needing an explicit finalize call (DingTalk AI Cards) force a
|
||||
# redundant final edit; ``is True`` keeps MagicMock adapters out.
|
||||
self._adapter_requires_finalize: bool = (
|
||||
getattr(adapter, "REQUIRES_EDIT_FINALIZE", False) is True
|
||||
)
|
||||
self._adapter_requires_finalize = getattr(adapter, "REQUIRES_EDIT_FINALIZE", False) is True
|
||||
# Telegram bounds edit retries at 5s; a fallback must not wait longer.
|
||||
self._max_fallback_flood_retry_seconds = 5.0
|
||||
|
||||
@@ -320,9 +311,8 @@ class GatewayStreamConsumer(
|
||||
return None
|
||||
if self._delivered_final_text is not None:
|
||||
# A segment break / commentary may have delivered it under another record.
|
||||
return (
|
||||
self._delivered_final_text.strip() == target or self.has_delivered_text(final_text)
|
||||
)
|
||||
return (self._delivered_final_text.strip() == target
|
||||
or self.has_delivered_text(final_text))
|
||||
if self._turn_split_delivery:
|
||||
return False
|
||||
# No recorded payload: judge against the FINAL content, not the flag.
|
||||
@@ -336,11 +326,8 @@ class GatewayStreamConsumer(
|
||||
def has_delivered_text(self, text: str) -> bool:
|
||||
"""Return True if *text* was already delivered as visible chat content."""
|
||||
target = self._clean_for_display(text or "").strip()
|
||||
seen = (
|
||||
self._visible_prefix(),
|
||||
*self._delivered_commentary_texts,
|
||||
*self._delivered_segment_texts,
|
||||
)
|
||||
seen = (self._visible_prefix(), *self._delivered_commentary_texts,
|
||||
*self._delivered_segment_texts)
|
||||
return bool(target) and any(sent.strip() == target for sent in seen)
|
||||
|
||||
def on_segment_break(self) -> None:
|
||||
@@ -348,10 +335,7 @@ class GatewayStreamConsumer(
|
||||
self._queue.put(_NEW_SEGMENT)
|
||||
|
||||
def close_for_approval_prompt(
|
||||
self,
|
||||
placeholder: str | None = None,
|
||||
reason: str = "Approval",
|
||||
reopen: bool = False,
|
||||
self, placeholder: str | None = None, reason: str = "Approval", reopen: bool = False,
|
||||
) -> asyncio.Future:
|
||||
"""Queue an interaction boundary (approval / clarify prompt) from sync context.
|
||||
run() finalizes the current native stream (``placeholder`` when empty), then per
|
||||
@@ -393,11 +377,8 @@ class GatewayStreamConsumer(
|
||||
|
||||
def _reopen_seed_pending(self) -> bool:
|
||||
"""Native stream, reopen requested after a boundary, nothing open yet."""
|
||||
return (
|
||||
self._use_native_streaming
|
||||
and self._awaiting_reopen_after_boundary
|
||||
and not self._native_stream_opened
|
||||
)
|
||||
return (self._use_native_streaming and self._awaiting_reopen_after_boundary
|
||||
and not self._native_stream_opened)
|
||||
|
||||
def request_reopen_seed(self) -> None:
|
||||
"""Thread-safe: request an EAGER native re-seed after a clarify answer. No-op unless
|
||||
@@ -462,20 +443,15 @@ class GatewayStreamConsumer(
|
||||
self._degrade_native_to_buffered_send()
|
||||
self._reset_segment_state()
|
||||
if self._boundary_reopen:
|
||||
logger.info(
|
||||
"[latency] Clarify boundary finalized, awaiting first "
|
||||
"post-answer delta to re-seed (chat=%s, turn=%s)",
|
||||
self.chat_id, self._turn_id,
|
||||
)
|
||||
logger.info("[latency] Clarify boundary finalized, awaiting first "
|
||||
"post-answer delta to re-seed (chat=%s, turn=%s)",
|
||||
self.chat_id, self._turn_id)
|
||||
except Exception as e:
|
||||
logger.warning("%s boundary processing failed: %s", _reason, e)
|
||||
boundary_ok = False
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
if (
|
||||
isinstance(boundary_future, (asyncio.Future, concurrent.futures.Future))
|
||||
and not boundary_future.done()
|
||||
):
|
||||
if isinstance(boundary_future, _FUTURE_TYPES) and not boundary_future.done():
|
||||
boundary_future.set_result(boundary_ok)
|
||||
|
||||
async def _finalize_boundary_stream(self, _reason: str) -> bool:
|
||||
@@ -484,30 +460,23 @@ class GatewayStreamConsumer(
|
||||
finalize_text = self._accumulated or self._boundary_placeholder
|
||||
try:
|
||||
if await self._send_frame(finalize_text, finalize=True):
|
||||
logger.debug(
|
||||
"%s boundary: finalized stream (chat=%s, turn=%s)",
|
||||
_reason, self.chat_id, self._turn_id,
|
||||
)
|
||||
logger.debug("%s boundary: finalized stream (chat=%s, turn=%s)",
|
||||
_reason, self.chat_id, self._turn_id)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning("%s boundary: finalize failed: %s", _reason, e)
|
||||
# Typing bubble may still show partial content; deliver via send().
|
||||
logger.warning(
|
||||
"%s boundary: finalize not confirmed, "
|
||||
"falling back to send() for pre-prompt text (chat=%s)",
|
||||
_reason, self.chat_id,
|
||||
)
|
||||
logger.warning("%s boundary: finalize not confirmed, "
|
||||
"falling back to send() for pre-prompt text (chat=%s)",
|
||||
_reason, self.chat_id)
|
||||
try:
|
||||
send_result = await self.adapter.send(self.chat_id, finalize_text)
|
||||
if getattr(send_result, "success", False):
|
||||
if getattr(await self.adapter.send(self.chat_id, finalize_text), "success", False):
|
||||
return True
|
||||
except Exception as send_err:
|
||||
logger.warning("%s boundary: fallback send also failed: %s", _reason, send_err)
|
||||
logger.error(
|
||||
"%s boundary: both finalize and fallback send failed "
|
||||
"(chat=%s) — pre-prompt text may not have been delivered",
|
||||
_reason, self.chat_id,
|
||||
)
|
||||
logger.error("%s boundary: both finalize and fallback send failed "
|
||||
"(chat=%s) — pre-prompt text may not have been delivered",
|
||||
_reason, self.chat_id)
|
||||
return False
|
||||
|
||||
def on_delta(self, text: str) -> None:
|
||||
@@ -558,13 +527,12 @@ class GatewayStreamConsumer(
|
||||
return
|
||||
|
||||
if self._should_edit(tick) and (
|
||||
self._accumulated
|
||||
or (self._use_native_streaming and self._tool_progress_active)
|
||||
self._accumulated or (self._use_native_streaming and self._tool_progress_active)
|
||||
):
|
||||
# Overflow split. Native streaming bypasses this: the adapter
|
||||
# truncates against the stream protocol's own limit.
|
||||
if not self._use_native_streaming and self._first_send_overflows():
|
||||
if await self._split_first_send(tick) == "return":
|
||||
if await self._split_first_send(tick):
|
||||
return
|
||||
continue
|
||||
await self._seal_overflow_heads()
|
||||
@@ -598,11 +566,8 @@ class GatewayStreamConsumer(
|
||||
def _resolve_length_budget(self) -> "tuple[Callable[[str], int], int]":
|
||||
"""Per-chat length function (relay adapters differ per chat, e.g. utf16) + budget.
|
||||
isinstance gate: MagicMock auto-attributes aren't callables; test doubles use len."""
|
||||
len_fn: "Callable[[str], int]" = (
|
||||
self.adapter.message_len_fn_for_chat(self.chat_id)
|
||||
if isinstance(self.adapter, _BasePlatformAdapter)
|
||||
else len
|
||||
)
|
||||
len_fn = (self.adapter.message_len_fn_for_chat(self.chat_id)
|
||||
if isinstance(self.adapter, _BasePlatformAdapter) else len)
|
||||
return len_fn, max(500, self._raw_message_limit() - len_fn(self.cfg.cursor) - 100)
|
||||
|
||||
async def _start_transports(self) -> None:
|
||||
@@ -611,9 +576,8 @@ class GatewayStreamConsumer(
|
||||
self._use_native_streaming = self._resolve_native_streaming()
|
||||
if self._use_native_streaming:
|
||||
logger.debug("Stream consumer using native-stream transport (chat=%s)", self.chat_id)
|
||||
if await self._try_seed_frame(
|
||||
"Native streaming seed frame raised; disabling native", exc_info=True,
|
||||
):
|
||||
if await self._try_seed_frame("Native streaming seed frame raised; disabling native",
|
||||
exc_info=True):
|
||||
self._native_stream_opened = True
|
||||
self._use_draft_streaming = False
|
||||
return
|
||||
@@ -621,10 +585,8 @@ class GatewayStreamConsumer(
|
||||
self._use_draft_streaming = self._resolve_draft_streaming()
|
||||
if self._use_draft_streaming:
|
||||
self._bump_draft_id()
|
||||
logger.debug(
|
||||
"Stream consumer using native-draft transport (chat=%s draft_id=%s)",
|
||||
self.chat_id, self._draft_id,
|
||||
)
|
||||
logger.debug("Stream consumer using native-draft transport (chat=%s draft_id=%s)",
|
||||
self.chat_id, self._draft_id)
|
||||
|
||||
def _drain_queue(self) -> "_Tick":
|
||||
"""Drain everything queued so far into one tick. Control sentinels stop the drain
|
||||
@@ -636,43 +598,35 @@ class GatewayStreamConsumer(
|
||||
item = self._queue.get_nowait()
|
||||
except queue.Empty:
|
||||
return tick
|
||||
for sentinel, flag in self._QUEUE_SENTINEL_FLAGS:
|
||||
if item is sentinel:
|
||||
setattr(tick, flag, True)
|
||||
return tick
|
||||
handler = None
|
||||
if isinstance(item, tuple):
|
||||
with contextlib.suppress(TypeError): # unhashable head: not one of ours
|
||||
handler = self._QUEUE_TUPLE_HANDLERS.get((item[0], len(item)))
|
||||
if handler is None:
|
||||
self._filter_and_accumulate(item)
|
||||
elif handler(self, tick, item):
|
||||
if item is _DONE:
|
||||
tick.got_done = True
|
||||
return tick
|
||||
|
||||
def _on_final_text(self, tick: "_Tick", item: tuple) -> bool:
|
||||
self._adopt_final_text(item[1])
|
||||
return False
|
||||
|
||||
def _on_approval_boundary(self, tick: "_Tick", item: tuple) -> bool:
|
||||
tick.approval_boundary = (item[1], item[2])
|
||||
return True
|
||||
|
||||
def _on_commentary(self, tick: "_Tick", item: tuple) -> bool:
|
||||
tick.commentary_text = item[1]
|
||||
return True
|
||||
|
||||
def _on_flush(self, tick: "_Tick", item: tuple) -> bool:
|
||||
# Barrier: finalize like a tool boundary, signal at the end of the tick.
|
||||
tick.got_flush = True
|
||||
tick.got_segment_break = True
|
||||
tick.flush_event = item[1]
|
||||
return True
|
||||
|
||||
def _on_tool_progress(self, tick: "_Tick", item: tuple) -> bool:
|
||||
if self._use_native_streaming:
|
||||
self._tool_progress_lines.append(item[1])
|
||||
self._tool_progress_active = True
|
||||
return False # keep draining to batch simultaneous progress lines
|
||||
if item is _NEW_SEGMENT:
|
||||
tick.got_segment_break = True
|
||||
return tick
|
||||
if item is _REOPEN_SEED:
|
||||
tick.got_reopen_seed = True
|
||||
return tick
|
||||
kind = item[0] if isinstance(item, tuple) and item else None
|
||||
if kind is _FINAL_TEXT:
|
||||
self._adopt_final_text(item[1])
|
||||
elif kind is _TOOL_PROGRESS: # keep draining to batch simultaneous lines
|
||||
if self._use_native_streaming:
|
||||
self._tool_progress_lines.append(item[1])
|
||||
self._tool_progress_active = True
|
||||
elif kind is _APPROVAL_BOUNDARY:
|
||||
tick.approval_boundary = (item[1], item[2])
|
||||
return tick
|
||||
elif kind is _COMMENTARY:
|
||||
tick.commentary_text = item[1]
|
||||
return tick
|
||||
elif kind is _FLUSH:
|
||||
# Barrier: finalize like a tool boundary, signal at the end of the tick.
|
||||
tick.got_flush = tick.got_segment_break = True
|
||||
tick.flush_event = item[1]
|
||||
return tick
|
||||
else:
|
||||
self._filter_and_accumulate(item)
|
||||
|
||||
def _adopt_final_text(self, final_raw: str) -> None:
|
||||
"""Adopt the authoritative final (see finish()) as the finalize content — only if this
|
||||
@@ -704,11 +658,8 @@ class GatewayStreamConsumer(
|
||||
self._native_last_pushed_len = 0
|
||||
self._awaiting_reopen_after_boundary = False
|
||||
self._reopen_seeded_eagerly = True
|
||||
logger.info(
|
||||
"[latency] Eager re-seed after clarify answer "
|
||||
"(typing bubble reopened immediately, turn=%s)",
|
||||
self._turn_id,
|
||||
)
|
||||
logger.info("[latency] Eager re-seed after clarify answer "
|
||||
"(typing bubble reopened immediately, turn=%s)", self._turn_id)
|
||||
else:
|
||||
# Degrade to a single buffered send(), like the approval path.
|
||||
self._degrade_native_to_buffered_send()
|
||||
@@ -726,20 +677,17 @@ class GatewayStreamConsumer(
|
||||
elapsed = time.monotonic() - self._last_edit_time
|
||||
# buffer_threshold is a codepoint debounce heuristic, not a
|
||||
# platform-limit check (_len_fn is for overflow).
|
||||
should_edit = bool(
|
||||
(elapsed >= self._current_edit_interval and self._accumulated)
|
||||
or len(self._accumulated) >= self.cfg.buffer_threshold
|
||||
)
|
||||
should_edit = bool((elapsed >= self._current_edit_interval and self._accumulated)
|
||||
or len(self._accumulated) >= self.cfg.buffer_threshold)
|
||||
# Defer mid-stream edits while the buffer could still resolve to a silence
|
||||
# marker ("NO"→"NO_REPLY"); got_done always resolves the buffer.
|
||||
return should_edit and not _is_partial_silence_marker(
|
||||
self._clean_for_display(self._accumulated)
|
||||
)
|
||||
self._clean_for_display(self._accumulated))
|
||||
|
||||
async def _split_first_send(self, tick: "_Tick") -> str:
|
||||
async def _split_first_send(self, tick: "_Tick") -> bool:
|
||||
"""No message to edit yet and the buffer overflows: seal only the head chunks; the
|
||||
tail stays in _accumulated as the active preview later deltas edit in place.
|
||||
Returns "return" (turn finished) or "continue"."""
|
||||
True when the turn finished here (the run loop returns)."""
|
||||
chunks = self._truncate_for_stream(self._accumulated, self._safe_limit, self._len_fn)
|
||||
if len(chunks) <= 1:
|
||||
# Malformed/legacy adapter result must still be splittable.
|
||||
@@ -766,24 +714,23 @@ class GatewayStreamConsumer(
|
||||
self._last_sent_text = ""
|
||||
self._last_edit_time = time.monotonic()
|
||||
if tick.got_done:
|
||||
tail_delivered = not self._accumulated or await self._send_or_edit(
|
||||
self._accumulated, finalize=True,
|
||||
)
|
||||
tail_delivered = (not self._accumulated
|
||||
or await self._send_or_edit(self._accumulated, finalize=True))
|
||||
# ``_already_sent`` may be True from prior state — only heads + tail count.
|
||||
self._final_response_sent = heads_delivered and tail_delivered
|
||||
if self._final_response_sent:
|
||||
self._turn_split_delivery = True
|
||||
self._mark_final_delivered(record=self._accumulated)
|
||||
return "return"
|
||||
return True
|
||||
if tick.got_segment_break:
|
||||
self._fallback_final_send = False
|
||||
self._fallback_prefix = ""
|
||||
if not self._accumulated:
|
||||
return "continue"
|
||||
return False
|
||||
# Early `continue` skips the bottom-of-loop flush signal.
|
||||
if tick.got_flush:
|
||||
self._signal_flush(tick.flush_event)
|
||||
return "continue"
|
||||
return False
|
||||
|
||||
def _overflows(self) -> bool:
|
||||
return self._len_fn(self._accumulated) > self._safe_limit
|
||||
@@ -823,17 +770,14 @@ class GatewayStreamConsumer(
|
||||
|
||||
# A got_done FRESH send via the draft transport already carries finalize=True,
|
||||
# unlike an EDIT, which REQUIRES_EDIT_FINALIZE adapters still need a pass for.
|
||||
tick.draft_final_fresh_send = (
|
||||
tick.got_done and self._use_draft_streaming and self._message_id is None
|
||||
)
|
||||
tick.draft_final_fresh_send = (tick.got_done and self._use_draft_streaming
|
||||
and self._message_id is None)
|
||||
# Segment break finalizes so platforms needing explicit closure (DingTalk AI
|
||||
# Cards) don't leave the segment stuck loading; it closes a preamble, not the
|
||||
# answer.
|
||||
tick.update_visible = await self._send_or_edit(
|
||||
display_text,
|
||||
finalize=(tick.got_done or tick.got_segment_break),
|
||||
is_turn_final=tick.got_done,
|
||||
)
|
||||
display_text, finalize=tick.got_done or tick.got_segment_break,
|
||||
is_turn_final=tick.got_done)
|
||||
self._last_edit_time = time.monotonic()
|
||||
# Lines stay in _tool_progress_lines for the next compose.
|
||||
self._tool_progress_active = False
|
||||
@@ -845,25 +789,15 @@ class GatewayStreamConsumer(
|
||||
if self._reopen_seed_pending() and not self._accumulated:
|
||||
# Lazy reopen, no post-prompt content: nothing is open on screen, so
|
||||
# don't re-seed just to emit a lone "✅".
|
||||
logger.debug(
|
||||
"Clarify reopen boundary with no post-prompt content "
|
||||
"— skipping lone-placeholder finalize (turn=%s)",
|
||||
self._turn_id,
|
||||
)
|
||||
elif (
|
||||
self._reopen_seeded_eagerly
|
||||
and self._native_stream_opened
|
||||
and not self._accumulated
|
||||
and not tick.update_visible
|
||||
):
|
||||
logger.debug("Clarify reopen boundary with no post-prompt content "
|
||||
"— skipping lone-placeholder finalize (turn=%s)", self._turn_id)
|
||||
elif (self._reopen_seeded_eagerly and self._native_stream_opened
|
||||
and not self._accumulated and not tick.update_visible):
|
||||
# Eager seed, no content: the typing bubble IS on screen and would hang
|
||||
# forever — close it with an empty finalize. Delivery flags untouched.
|
||||
await self._close_empty_native_bubble("Eager-seed empty finalize failed: %s")
|
||||
logger.debug(
|
||||
"Eager reopen seed but no post-answer content — "
|
||||
"closed empty typing bubble (turn=%s)",
|
||||
self._turn_id,
|
||||
)
|
||||
logger.debug("Eager reopen seed but no post-answer content — "
|
||||
"closed empty typing bubble (turn=%s)", self._turn_id)
|
||||
elif self._use_native_streaming:
|
||||
# Native streams MUST close with finish=true even when empty (tool-only
|
||||
# turns) — placeholder if needed.
|
||||
@@ -881,11 +815,8 @@ class GatewayStreamConsumer(
|
||||
elif self._final_response_sent:
|
||||
# Fresh-final already delivered; a second finalize would duplicate.
|
||||
self._mark_final_delivered(record=self._accumulated)
|
||||
elif tick.update_visible and (
|
||||
not self._adapter_requires_finalize
|
||||
or self._last_edit_overflowed
|
||||
or tick.draft_final_fresh_send
|
||||
):
|
||||
elif tick.update_visible and (not self._adapter_requires_finalize
|
||||
or self._last_edit_overflowed or tick.draft_final_fresh_send):
|
||||
# The update already delivered the final. A second finalize would re-edit
|
||||
# it (Telegram: editMessageText after sendRichMessage falls back to the
|
||||
# legacy formatter) or overflow-split again, duplicating chunks.
|
||||
@@ -910,9 +841,8 @@ class GatewayStreamConsumer(
|
||||
|
||||
def _cumulative_transport(self) -> bool:
|
||||
"""Stream-is-the-message drafts and WeCom native: one append-only stream per turn."""
|
||||
return (
|
||||
self._stream_is_message() and self._use_draft_streaming
|
||||
) or self._use_native_streaming
|
||||
stream_draft = self._stream_is_message() and self._use_draft_streaming
|
||||
return stream_draft or self._use_native_streaming
|
||||
|
||||
async def _deliver_commentary(self, commentary_text: str) -> None:
|
||||
"""Post commentary as its own message. Cumulative transports keep the stream going —
|
||||
@@ -936,12 +866,8 @@ class GatewayStreamConsumer(
|
||||
return
|
||||
# If the segment-break edit didn't land (flood control / fallback mode),
|
||||
# _accumulated holds unseen pre-boundary text — flush it before the reset.
|
||||
if (
|
||||
self._accumulated
|
||||
and not tick.update_visible
|
||||
and self._message_id
|
||||
and self._message_id != "__no_edit__"
|
||||
):
|
||||
if (self._accumulated and not tick.update_visible and self._message_id
|
||||
and self._message_id != "__no_edit__"):
|
||||
await self._flush_segment_tail_on_edit_failure()
|
||||
self._reset_segment_state(preserve_no_edit=True)
|
||||
|
||||
@@ -953,11 +879,8 @@ class GatewayStreamConsumer(
|
||||
best_effort_ok = False
|
||||
if self._accumulated and self._message_id:
|
||||
with contextlib.suppress(Exception):
|
||||
best_effort_ok = bool(
|
||||
await self._send_or_edit(
|
||||
self._accumulated, finalize=True, is_turn_final=False,
|
||||
)
|
||||
)
|
||||
best_effort_ok = bool(await self._send_or_edit(
|
||||
self._accumulated, finalize=True, is_turn_final=False))
|
||||
elif self._message_id is None:
|
||||
# Draft path keeps _message_id=None; seal in place (else the stream stays
|
||||
# visibly live and the adapter keeps armed interception state).
|
||||
@@ -974,22 +897,6 @@ class GatewayStreamConsumer(
|
||||
if isinstance(item, tuple) and len(item) == 2 and item[0] is _FLUSH:
|
||||
self._signal_flush(item[1])
|
||||
|
||||
# Tuple-shaped queue items keyed on (sentinel, arity); handler returns True
|
||||
# to stop draining. Order-insensitive: each sentinel is a distinct object.
|
||||
_QUEUE_TUPLE_HANDLERS = {
|
||||
(_FINAL_TEXT, 2): _on_final_text,
|
||||
(_APPROVAL_BOUNDARY, 3): _on_approval_boundary,
|
||||
(_COMMENTARY, 2): _on_commentary,
|
||||
(_FLUSH, 2): _on_flush,
|
||||
(_TOOL_PROGRESS, 2): _on_tool_progress,
|
||||
}
|
||||
# Bare control sentinels -> the _Tick flag they set (each ends the drain).
|
||||
_QUEUE_SENTINEL_FLAGS = (
|
||||
(_DONE, "got_done"),
|
||||
(_NEW_SEGMENT, "got_segment_break"),
|
||||
(_REOPEN_SEED, "got_reopen_seed"),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _clean_for_display(text: str) -> str:
|
||||
"""Hide MEDIA:<path> / [[audio_as_voice]] directives; media is delivered post-stream."""
|
||||
|
||||
@@ -17,24 +17,16 @@ logger = logging.getLogger("gateway.stream_consumer")
|
||||
class StreamFallbackMixin:
|
||||
"""Non-streaming delivery paths used once progressive edits fail or the turn ends oddly."""
|
||||
|
||||
async def _send_new_chunk(
|
||||
self,
|
||||
text: str,
|
||||
reply_to_id: Optional[str],
|
||||
*,
|
||||
final: bool = False,
|
||||
) -> Optional[str]:
|
||||
async def _send_new_chunk(self, text: str, reply_to_id: Optional[str], *,
|
||||
final: bool = False) -> Optional[str]:
|
||||
"""Send a new chunk threaded to ``reply_to_id``; returns the new message_id."""
|
||||
text = self._clean_for_display(text)
|
||||
if not text.strip():
|
||||
return reply_to_id
|
||||
try:
|
||||
result = await self.adapter.send(
|
||||
chat_id=self.chat_id,
|
||||
content=text,
|
||||
reply_to=reply_to_id,
|
||||
metadata=self._metadata_for_send(final=final, expect_edits=not final),
|
||||
)
|
||||
chat_id=self.chat_id, content=text, reply_to=reply_to_id,
|
||||
metadata=self._metadata_for_send(final=final, expect_edits=not final))
|
||||
if not (result.success and result.message_id):
|
||||
self._edit_supported = False
|
||||
return reply_to_id
|
||||
@@ -63,19 +55,15 @@ class StreamFallbackMixin:
|
||||
return final_text
|
||||
|
||||
@staticmethod
|
||||
def _split_text_chunks(
|
||||
text: str, limit: int, len_fn: "Callable[[str], int]" = len,
|
||||
) -> list[str]:
|
||||
def _split_text_chunks(text: str, limit: int, len_fn: "Callable[[str], int]" = len,
|
||||
) -> list[str]:
|
||||
"""Split text for fallback sends: newline-preferred, fence-balanced across chunks."""
|
||||
from gateway.platforms.helpers import split_text_fence_aware
|
||||
return split_text_fence_aware(text, limit, len_fn, prefer_paragraphs=False,
|
||||
balance_fences=True)
|
||||
|
||||
return split_text_fence_aware(
|
||||
text, limit, len_fn, prefer_paragraphs=False, balance_fences=True,
|
||||
)
|
||||
|
||||
def _truncate_for_stream(
|
||||
self, text: str, limit: int, len_fn: "Callable[[str], int]",
|
||||
) -> list[str]:
|
||||
def _truncate_for_stream(self, text: str, limit: int, len_fn: "Callable[[str], int]",
|
||||
) -> list[str]:
|
||||
"""Split via the adapter's canonical truncate_message (platform-specific rules);
|
||||
non-base test doubles / legacy adapters keep the two-argument call shape."""
|
||||
truncate = getattr(self.adapter, "truncate_message", None)
|
||||
@@ -111,8 +99,7 @@ class StreamFallbackMixin:
|
||||
sent_any_chunk = False
|
||||
for chunk in chunks:
|
||||
result = await self._send_with_flood_retry(
|
||||
content=chunk, retry_log="Flood control on fallback send, retrying in %.1fs",
|
||||
)
|
||||
content=chunk, retry_log="Flood control on fallback send, retrying in %.1fs")
|
||||
if not result or not result.success:
|
||||
# Partial continuation landed: do NOT set _final_response_sent (the
|
||||
# gateway must still deliver the full answer); _already_sent only
|
||||
@@ -131,15 +118,10 @@ class StreamFallbackMixin:
|
||||
# Best-effort delete of the frozen partial — ONLY when the FULL final was
|
||||
# re-sent. If only the missing tail went out, the partial IS the head of
|
||||
# the answer ("sent only the second half" symptom).
|
||||
if (
|
||||
stale_message_id
|
||||
and stale_message_id != last_message_id
|
||||
and not self._fallback_preserve_partial_messages
|
||||
and continuation == final_text
|
||||
):
|
||||
await self._delete_previews(
|
||||
[stale_message_id], label="Fallback partial", skip_sentinel=False,
|
||||
)
|
||||
if (stale_message_id and stale_message_id != last_message_id
|
||||
and not self._fallback_preserve_partial_messages and continuation == final_text):
|
||||
await self._delete_previews([stale_message_id], label="Fallback partial",
|
||||
skip_sentinel=False)
|
||||
|
||||
self._message_id = last_message_id
|
||||
self._already_sent = True
|
||||
@@ -156,11 +138,8 @@ class StreamFallbackMixin:
|
||||
visible = self._visible_prefix()
|
||||
# Telegram clients can lose (part of) a streamed preview after a failed
|
||||
# final edit, so opt-in adapters commit a fresh final send.
|
||||
if (
|
||||
final_text.strip()
|
||||
and final_text == visible
|
||||
and getattr(self.adapter, "RESEND_FINAL_ON_EMPTY_STREAM_FALLBACK", False) is True
|
||||
):
|
||||
if (final_text.strip() and final_text == visible
|
||||
and getattr(self.adapter, "RESEND_FINAL_ON_EMPTY_STREAM_FALLBACK", False) is True):
|
||||
delivery = await self._send_empty_fallback_final(final_text)
|
||||
if delivery == "delivered":
|
||||
return None
|
||||
@@ -183,12 +162,8 @@ class StreamFallbackMixin:
|
||||
if final_text.strip() and final_text != visible:
|
||||
return final_text
|
||||
# Best-effort strip of a cursor left stuck by the edit failure.
|
||||
if (
|
||||
self._message_id
|
||||
and self._last_sent_text
|
||||
and self.cfg.cursor
|
||||
and self._last_sent_text.endswith(self.cfg.cursor)
|
||||
):
|
||||
if (self._message_id and self._last_sent_text and self.cfg.cursor
|
||||
and self._last_sent_text.endswith(self.cfg.cursor)):
|
||||
clean_text = self._last_sent_text[:-len(self.cfg.cursor)]
|
||||
with contextlib.suppress(Exception):
|
||||
result = await self._edit_message(message_id=self._message_id, content=clean_text)
|
||||
@@ -215,9 +190,8 @@ class StreamFallbackMixin:
|
||||
async def _send_with_flood_retry(self, *, content: str, retry_log: str, reply_to=None):
|
||||
"""adapter.send(final metadata) with ONE bounded flood retry; returns the last
|
||||
SendResult. Exceptions propagate (callers decide whether a raise is "ambiguous")."""
|
||||
kwargs = dict(
|
||||
chat_id=self.chat_id, content=content, metadata=self._metadata_for_send(final=True),
|
||||
)
|
||||
kwargs = dict(chat_id=self.chat_id, content=content,
|
||||
metadata=self._metadata_for_send(final=True))
|
||||
if reply_to is not None:
|
||||
kwargs["reply_to"] = reply_to
|
||||
result = None
|
||||
@@ -240,10 +214,8 @@ class StreamFallbackMixin:
|
||||
stale_ids = self._stale_preview_ids(segment_only=True)
|
||||
try:
|
||||
result = await self._send_with_flood_retry(
|
||||
content=final_text,
|
||||
reply_to=self._initial_reply_to_id,
|
||||
retry_log="Flood control on empty fallback final send; retrying in %.1fs",
|
||||
)
|
||||
content=final_text, reply_to=self._initial_reply_to_id,
|
||||
retry_log="Flood control on empty fallback final send; retrying in %.1fs")
|
||||
except Exception as exc:
|
||||
logger.debug("Empty fallback final send failed: %s", exc)
|
||||
return "ambiguous" if self._send_failure_may_have_delivered(exc) else "failed"
|
||||
@@ -255,9 +227,8 @@ class StreamFallbackMixin:
|
||||
new_message_id = getattr(result, "message_id", None)
|
||||
# Telegram reports delete failure by returning False; the flood window that
|
||||
# broke the finalize can reject this too — one bounded retry.
|
||||
await self._delete_previews(
|
||||
stale_ids, skip=new_message_id, label="Empty fallback", retry_on_false=True,
|
||||
)
|
||||
await self._delete_previews(stale_ids, skip=new_message_id, label="Empty fallback",
|
||||
retry_on_false=True)
|
||||
self._segment_preview_message_ids = set()
|
||||
self._message_id = new_message_id or "__no_edit__"
|
||||
self._already_sent = True
|
||||
@@ -290,10 +261,8 @@ class StreamFallbackMixin:
|
||||
except (TypeError, ValueError):
|
||||
delay = 3.0
|
||||
if delay > self._max_fallback_flood_retry_seconds:
|
||||
logger.debug(
|
||||
"Flood control requests %.1fs; leaving final delivery to the gateway",
|
||||
delay,
|
||||
)
|
||||
logger.debug("Flood control requests %.1fs; leaving final delivery to the gateway",
|
||||
delay)
|
||||
return None
|
||||
return max(0.0, delay)
|
||||
|
||||
@@ -351,11 +320,8 @@ class StreamFallbackMixin:
|
||||
_platform_name = str(_plat or getattr(self.adapter, "name", "")).lower()
|
||||
_needs_reply_anchor = _platform_name in ("buzz", "slack", "mattermost", "feishu")
|
||||
result = await self.adapter.send(
|
||||
chat_id=self.chat_id,
|
||||
content=text,
|
||||
reply_to=self._initial_reply_to_id if _needs_reply_anchor else None,
|
||||
metadata=_md,
|
||||
)
|
||||
chat_id=self.chat_id, content=text,
|
||||
reply_to=self._initial_reply_to_id if _needs_reply_anchor else None, metadata=_md)
|
||||
# Do NOT set _already_sent: commentary is interim, and the flag would
|
||||
# suppress the real final after multiple tool calls.
|
||||
if result.success:
|
||||
|
||||
@@ -25,33 +25,25 @@ class StreamTransportMixin:
|
||||
async def _edit_message(self, *, message_id: str, content: str, finalize: bool = False):
|
||||
"""Edit via the adapter, passing routing metadata when supported."""
|
||||
# Contract: adapters must accept finalize= even when False (test-guarded).
|
||||
kwargs = {
|
||||
"chat_id": self.chat_id,
|
||||
"message_id": message_id,
|
||||
"content": content,
|
||||
"finalize": finalize,
|
||||
}
|
||||
kwargs = dict(chat_id=self.chat_id, message_id=message_id, content=content,
|
||||
finalize=finalize)
|
||||
if self.metadata:
|
||||
try:
|
||||
params = inspect.signature(self.adapter.edit_message).parameters
|
||||
if "metadata" in params or any(
|
||||
param.kind is inspect.Parameter.VAR_KEYWORD for param in params.values()
|
||||
):
|
||||
param.kind is inspect.Parameter.VAR_KEYWORD for param in params.values()):
|
||||
kwargs["metadata"] = self.metadata
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
return await self.adapter.edit_message(**kwargs)
|
||||
|
||||
async def _send_seed_frame(self):
|
||||
"""Open a native stream with an empty seed frame (typing indicator before any token)."""
|
||||
return await self.adapter.send_stream_frame(
|
||||
"", chat_id=self.chat_id, reply_to=self._initial_reply_to_id, turn_id=self._turn_id,
|
||||
)
|
||||
|
||||
async def _try_seed_frame(self, fail_log: str, *, exc_info: bool = False) -> bool:
|
||||
"""_send_seed_frame() as a bool; a raise logs ``fail_log`` at DEBUG (with the error
|
||||
formatted in, or the traceback when ``exc_info``) and reads as False."""
|
||||
return await self._try_frame(self._send_seed_frame(), fail_log, exc_info=exc_info)
|
||||
"""Open a native stream with an empty seed frame (typing indicator before any token) as a
|
||||
bool; a raise logs ``fail_log`` at DEBUG (error formatted in, or the traceback when
|
||||
``exc_info``) and reads as False."""
|
||||
seed = self.adapter.send_stream_frame(
|
||||
"", chat_id=self.chat_id, reply_to=self._initial_reply_to_id, turn_id=self._turn_id)
|
||||
return await self._try_frame(seed, fail_log, exc_info=exc_info)
|
||||
|
||||
@staticmethod
|
||||
async def _try_frame(coro, fail_log: str, *, exc_info: bool = False) -> bool:
|
||||
@@ -68,12 +60,8 @@ class StreamTransportMixin:
|
||||
async def _send_frame(self, text: str, *, finalize: bool):
|
||||
"""One native-stream frame; every frame carries the same chat/reply/turn routing."""
|
||||
return await self.adapter.send_stream_frame(
|
||||
text,
|
||||
finalize=finalize,
|
||||
chat_id=self.chat_id,
|
||||
reply_to=self._initial_reply_to_id,
|
||||
turn_id=self._turn_id,
|
||||
)
|
||||
text, finalize=finalize, chat_id=self.chat_id, reply_to=self._initial_reply_to_id,
|
||||
turn_id=self._turn_id)
|
||||
|
||||
def _close_native_state(self) -> None:
|
||||
"""Mark the native stream closed (next content re-seeds or falls back)."""
|
||||
@@ -103,17 +91,14 @@ class StreamTransportMixin:
|
||||
|
||||
def _stale_preview_ids(self, *, segment_only: bool = False) -> set:
|
||||
"""Preview ids a fresh final replaces; ``segment_only`` spares finalized preambles."""
|
||||
stale_ids = set(
|
||||
self._segment_preview_message_ids if segment_only else self._preview_message_ids
|
||||
)
|
||||
stale_ids = set(self._segment_preview_message_ids if segment_only
|
||||
else self._preview_message_ids)
|
||||
if self._message_id and self._message_id != "__no_edit__":
|
||||
stale_ids.add(str(self._message_id) if segment_only else self._message_id)
|
||||
return stale_ids
|
||||
|
||||
async def _delete_previews(
|
||||
self, stale_ids, *, skip=None, label: str, retry_on_false: bool = False,
|
||||
skip_sentinel: bool = True,
|
||||
) -> None:
|
||||
async def _delete_previews(self, stale_ids, *, skip=None, label: str,
|
||||
retry_on_false: bool = False, skip_sentinel: bool = True) -> None:
|
||||
"""Best-effort delete of stale previews; never the message just sent (``skip``)."""
|
||||
delete_fn = getattr(self.adapter, "delete_message", None)
|
||||
if delete_fn is None:
|
||||
@@ -133,39 +118,31 @@ class StreamTransportMixin:
|
||||
"""cfg.transport "draft"/"auto" → the adapter's supports_draft_streaming probe
|
||||
("draft" logs the downgrade); "edit"/"off" → False."""
|
||||
transport = (self.cfg.transport or "edit").lower()
|
||||
if transport in ("edit", "off"):
|
||||
return False
|
||||
# MagicMock test adapters default to edit.
|
||||
if not isinstance(self.adapter, _BasePlatformAdapter):
|
||||
if transport in ("edit", "off") or not isinstance(self.adapter, _BasePlatformAdapter):
|
||||
return False
|
||||
probe_kwargs = dict(chat_type=self.cfg.chat_type or None, metadata=self.metadata)
|
||||
try:
|
||||
try:
|
||||
# Per-chat probe (relay adapters resolve through the CHAT's
|
||||
# descriptor); older adapters without the kwarg keep the legacy probe.
|
||||
supported = self.adapter.supports_draft_streaming(
|
||||
chat_id=self.chat_id, **probe_kwargs,
|
||||
)
|
||||
supported = self.adapter.supports_draft_streaming(chat_id=self.chat_id,
|
||||
**probe_kwargs)
|
||||
except TypeError:
|
||||
supported = self.adapter.supports_draft_streaming(**probe_kwargs)
|
||||
except Exception:
|
||||
logger.debug("supports_draft_streaming probe raised", exc_info=True)
|
||||
supported = False
|
||||
if not supported and transport == "draft":
|
||||
logger.debug(
|
||||
"Draft streaming requested but unsupported (chat=%s, type=%r) — "
|
||||
"falling back to edit",
|
||||
self.chat_id, self.cfg.chat_type,
|
||||
)
|
||||
logger.debug("Draft streaming requested but unsupported (chat=%s, type=%r) — "
|
||||
"falling back to edit", self.chat_id, self.cfg.chat_type)
|
||||
return bool(supported)
|
||||
|
||||
def _resolve_native_streaming(self) -> bool:
|
||||
"""Native streaming (send_stream_frame for ALL frames): a BasePlatformAdapter with
|
||||
class-level SUPPORTS_NATIVE_STREAMING and a truthy supports_native_streaming probe."""
|
||||
if not (
|
||||
isinstance(self.adapter, _BasePlatformAdapter)
|
||||
and getattr(type(self.adapter), "SUPPORTS_NATIVE_STREAMING", False)
|
||||
):
|
||||
if not (isinstance(self.adapter, _BasePlatformAdapter)
|
||||
and getattr(type(self.adapter), "SUPPORTS_NATIVE_STREAMING", False)):
|
||||
return False
|
||||
probe = getattr(self.adapter, "supports_native_streaming", None)
|
||||
if probe is None:
|
||||
@@ -185,21 +162,16 @@ class StreamTransportMixin:
|
||||
return False
|
||||
try:
|
||||
result = await self.adapter.send_draft(
|
||||
chat_id=self.chat_id,
|
||||
draft_id=self._draft_id,
|
||||
content=text,
|
||||
metadata=self._draft_metadata(),
|
||||
)
|
||||
chat_id=self.chat_id, draft_id=self._draft_id, content=text,
|
||||
metadata=self._draft_metadata())
|
||||
except Exception as e:
|
||||
logger.debug("send_draft raised, disabling draft transport for this run: %s", e)
|
||||
else:
|
||||
if getattr(result, "success", False):
|
||||
self._last_sent_text = text # parity with the edit-based no-op skip
|
||||
return True
|
||||
logger.debug(
|
||||
"send_draft returned success=False, disabling draft transport: %s",
|
||||
getattr(result, "error", "unknown"),
|
||||
)
|
||||
logger.debug("send_draft returned success=False, disabling draft transport: %s",
|
||||
getattr(result, "error", "unknown"))
|
||||
self._draft_failures += 1
|
||||
self._use_draft_streaming = False
|
||||
return False
|
||||
@@ -214,10 +186,8 @@ class StreamTransportMixin:
|
||||
return
|
||||
try:
|
||||
await self.adapter.abandon_open_draft(
|
||||
self.chat_id,
|
||||
self._last_sent_text or self._clean_for_display(self._accumulated),
|
||||
metadata=self._draft_metadata(),
|
||||
)
|
||||
self.chat_id, self._last_sent_text or self._clean_for_display(self._accumulated),
|
||||
metadata=self._draft_metadata())
|
||||
except Exception as e:
|
||||
logger.debug("abandon_open_draft failed (best-effort): %s", e)
|
||||
|
||||
@@ -243,11 +213,8 @@ class StreamTransportMixin:
|
||||
"""Record the primary id plus any continuation ids from an oversized split."""
|
||||
raw = getattr(result, "raw_response", None) or {}
|
||||
raw_ids = raw.get("message_ids") if isinstance(raw, dict) else None
|
||||
for mid in (
|
||||
getattr(result, "message_id", None),
|
||||
*(getattr(result, "continuation_message_ids", None) or ()),
|
||||
*(raw_ids or ()),
|
||||
):
|
||||
for mid in (getattr(result, "message_id", None),
|
||||
*(getattr(result, "continuation_message_ids", None) or ()), *(raw_ids or ())):
|
||||
self._track_preview_id(mid)
|
||||
|
||||
def _adapter_prefers_fresh_final(self, text: str) -> bool:
|
||||
@@ -283,8 +250,7 @@ class StreamTransportMixin:
|
||||
stale_ids = self._stale_preview_ids()
|
||||
try:
|
||||
result = await self.adapter.send(
|
||||
chat_id=self.chat_id, content=text, metadata=self._metadata_for_send(final=True),
|
||||
)
|
||||
chat_id=self.chat_id, content=text, metadata=self._metadata_for_send(final=True))
|
||||
except Exception as e:
|
||||
logger.debug("Fresh-final send failed, falling back to edit: %s", e)
|
||||
return False
|
||||
@@ -312,8 +278,7 @@ class StreamTransportMixin:
|
||||
self._message_created_ts = None
|
||||
|
||||
async def _send_or_edit(
|
||||
self, text: str, *, finalize: bool = False, is_turn_final: bool = True,
|
||||
) -> bool:
|
||||
self, text: str, *, finalize: bool = False, is_turn_final: bool = True) -> bool:
|
||||
"""Send or edit the streaming message; True if delivered. ``finalize`` marks the
|
||||
last edit. Transport order: native frame → draft frame → edit existing → first
|
||||
send; a transport returns None to fall through to the next."""
|
||||
@@ -328,22 +293,15 @@ class StreamTransportMixin:
|
||||
if not visible_stripped:
|
||||
# Native streams MUST still get a finalize frame (placeholder) to close
|
||||
# the thinking bubble, e.g. for a MEDIA-only response.
|
||||
if (
|
||||
finalize and self._use_native_streaming and self._native_stream_opened
|
||||
and await self._try_frame(
|
||||
self._send_frame("✅", finalize=True), "Finalize empty stream failed: %s",
|
||||
)
|
||||
):
|
||||
if (finalize and self._use_native_streaming and self._native_stream_opened
|
||||
and await self._try_frame(self._send_frame("✅", finalize=True),
|
||||
"Finalize empty stream failed: %s")):
|
||||
self._mark_final_delivered()
|
||||
return True # cursor-only / whitespace-only update
|
||||
# Don't open a new message for 1-2 tokens + cursor (rapid tool-calling): if
|
||||
# the cursor-strip edit is then rate-limited, "X ▉" stays forever.
|
||||
if (
|
||||
self._message_id is None
|
||||
and self.cfg.cursor
|
||||
and self.cfg.cursor in text
|
||||
and len(visible_stripped) < self._MIN_NEW_MSG_CHARS
|
||||
):
|
||||
if (self._message_id is None and self.cfg.cursor and self.cfg.cursor in text
|
||||
and len(visible_stripped) < self._MIN_NEW_MSG_CHARS):
|
||||
return True # too short for a standalone message — accumulate more
|
||||
|
||||
# A failed native/draft transport disables itself and falls through so the
|
||||
@@ -353,9 +311,8 @@ class StreamTransportMixin:
|
||||
if ok is not None:
|
||||
return ok
|
||||
if self._use_draft_streaming and self._message_id is None:
|
||||
ok = await self._draft_push(
|
||||
text, pre_fence_text, finalize=finalize, is_turn_final=is_turn_final,
|
||||
)
|
||||
ok = await self._draft_push(text, pre_fence_text, finalize=finalize,
|
||||
is_turn_final=is_turn_final)
|
||||
if ok is not None:
|
||||
return ok
|
||||
self._last_edit_overflowed = False
|
||||
@@ -369,9 +326,8 @@ class StreamTransportMixin:
|
||||
logger.error("Stream send/edit error: %s", e)
|
||||
return False
|
||||
|
||||
async def _native_push(
|
||||
self, text: str, *, finalize: bool, is_turn_final: bool,
|
||||
) -> Optional[bool]:
|
||||
async def _native_push(self, text: str, *, finalize: bool, is_turn_final: bool,
|
||||
) -> Optional[bool]:
|
||||
"""Native streaming: every frame goes through send_stream_frame(); lazy re-seed after
|
||||
a boundary. None when native was disabled (seed/frame failure) → caller falls through."""
|
||||
if not self._native_stream_opened and text:
|
||||
@@ -381,11 +337,8 @@ class StreamTransportMixin:
|
||||
self._native_stream_opened = True
|
||||
self._awaiting_reopen_after_boundary = False
|
||||
# Paired with the boundary-finalize INFO: typing-reappear latency.
|
||||
logger.info(
|
||||
"[latency] Re-opened native stream after boundary "
|
||||
"(turn=%s, waited for first delta)",
|
||||
self._turn_id,
|
||||
)
|
||||
logger.info("[latency] Re-opened native stream after boundary "
|
||||
"(turn=%s, waited for first delta)", self._turn_id)
|
||||
|
||||
# WeCom renders each finalize as a separate bubble: only the turn-final and
|
||||
# boundaries close the stream, not segment breaks.
|
||||
@@ -399,10 +352,8 @@ class StreamTransportMixin:
|
||||
# stream-final-ack-timeout-duplicate.md). A definitive failure rolls it back.
|
||||
if finalize:
|
||||
self._mark_final_delivered(record=text) # recorded: stale frame can't suppress
|
||||
if await self._try_frame(
|
||||
self._send_frame(text, finalize=finalize),
|
||||
"send_stream_frame raised, disabling native streaming: %s",
|
||||
):
|
||||
if await self._try_frame(self._send_frame(text, finalize=finalize),
|
||||
"send_stream_frame raised, disabling native streaming: %s"):
|
||||
self._already_sent = True
|
||||
self._last_sent_text = text
|
||||
self._native_last_pushed_len = len(text)
|
||||
@@ -430,9 +381,8 @@ class StreamTransportMixin:
|
||||
logger.debug("Native fallback: failed to finalize stream: %s", e)
|
||||
return None
|
||||
|
||||
async def _draft_push(
|
||||
self, text: str, pre_fence_text: str, *, finalize: bool, is_turn_final: bool,
|
||||
) -> Optional[bool]:
|
||||
async def _draft_push(self, text: str, pre_fence_text: str, *, finalize: bool,
|
||||
is_turn_final: bool) -> Optional[bool]:
|
||||
"""Draft frame while no message_id exists; None = not applicable / drafts just failed.
|
||||
Skipped when finalizing (the real send clears the draft), EXCEPT stream-is-the-message
|
||||
adapters keep ONE stream per turn: a segment-break finalize must not become a real
|
||||
@@ -455,11 +405,8 @@ class StreamTransportMixin:
|
||||
async def _first_send(self, text: str, *, finalize: bool) -> bool:
|
||||
"""First send, threaded to the user's message (correct topic/thread)."""
|
||||
result = await self.adapter.send(
|
||||
chat_id=self.chat_id,
|
||||
content=text,
|
||||
reply_to=self._initial_reply_to_id,
|
||||
metadata=self._metadata_for_send(final=finalize, expect_edits=not finalize),
|
||||
)
|
||||
chat_id=self.chat_id, content=text, reply_to=self._initial_reply_to_id,
|
||||
metadata=self._metadata_for_send(final=finalize, expect_edits=not finalize))
|
||||
if not result.success:
|
||||
self._edit_supported = False
|
||||
return False
|
||||
@@ -488,29 +435,24 @@ class StreamTransportMixin:
|
||||
# CLASS (MagicMock auto-creates attrs) plus instance __dict__ (test doubles).
|
||||
has_prefers_hook = (
|
||||
hasattr(type(self.adapter), "prefers_fresh_final_streaming")
|
||||
or "prefers_fresh_final_streaming" in getattr(self.adapter, "__dict__", {})
|
||||
)
|
||||
or "prefers_fresh_final_streaming" in getattr(self.adapter, "__dict__", {}))
|
||||
prefers_fresh = self._adapter_prefers_fresh_final(text) # probed every edit (hook contract)
|
||||
if finalize and (
|
||||
prefers_fresh or (not has_prefers_hook and self._should_send_fresh_final())
|
||||
) and await self._try_fresh_final(text, is_turn_final=is_turn_final):
|
||||
return True
|
||||
result = await self._edit_message(
|
||||
message_id=self._message_id, content=text, finalize=finalize,
|
||||
)
|
||||
result = await self._edit_message(message_id=self._message_id, content=text,
|
||||
finalize=finalize)
|
||||
if not result.success:
|
||||
return await self._on_edit_failure(
|
||||
result, text, finalize=finalize, is_turn_final=is_turn_final,
|
||||
)
|
||||
return await self._on_edit_failure(result, text, finalize=finalize,
|
||||
is_turn_final=is_turn_final)
|
||||
self._already_sent = True
|
||||
self._track_preview_ids_from_result(result)
|
||||
# Oversized edit split across continuations: message_id is now the LAST
|
||||
# continuation, which holds only the final chunk — retarget edits and reset
|
||||
# skip-if-same. getattr keeps SimpleNamespace test mocks working.
|
||||
if (
|
||||
(getattr(result, "continuation_message_ids", ()) or ())
|
||||
and result.message_id and result.message_id != self._message_id
|
||||
):
|
||||
if ((getattr(result, "continuation_message_ids", ()) or ())
|
||||
and result.message_id and result.message_id != self._message_id):
|
||||
self._last_edit_overflowed = True
|
||||
self._turn_split_delivery = True
|
||||
self._adopt_message_id(str(result.message_id))
|
||||
@@ -528,18 +470,13 @@ class StreamTransportMixin:
|
||||
self._edit_supported = False
|
||||
self._already_sent = True
|
||||
|
||||
async def _on_edit_failure(
|
||||
self, result, text: str, *, finalize: bool, is_turn_final: bool,
|
||||
) -> bool:
|
||||
async def _on_edit_failure(self, result, text: str, *, finalize: bool, is_turn_final: bool,
|
||||
) -> bool:
|
||||
"""Classify a failed edit: partial overflow, flood backoff, or fallback mode. Always
|
||||
False; the caller's finalize path may still deliver the tail."""
|
||||
turn_final = finalize and is_turn_final
|
||||
if (
|
||||
turn_final
|
||||
and self.cfg.cursor
|
||||
and self._last_sent_text.endswith(self.cfg.cursor)
|
||||
and self._visible_prefix() == text
|
||||
):
|
||||
if (turn_final and self.cfg.cursor and self._last_sent_text.endswith(self.cfg.cursor)
|
||||
and self._visible_prefix() == text):
|
||||
# Cosmetic final edit was rate-limited but the full answer is already on
|
||||
# screen (cursor stuck): mark delivered so the gateway doesn't send it
|
||||
# twice, and record the on-screen payload.
|
||||
@@ -549,9 +486,8 @@ class StreamTransportMixin:
|
||||
if isinstance(raw_response, dict) and raw_response.get("partial_overflow"):
|
||||
# Some overflow chunks landed but not the whole response: preserve the
|
||||
# visible prefix so got_done sends the missing tail.
|
||||
self._message_id = str(
|
||||
raw_response.get("last_message_id") or result.message_id or self._message_id
|
||||
)
|
||||
self._message_id = str(raw_response.get("last_message_id") or result.message_id
|
||||
or self._message_id)
|
||||
delivered_prefix = raw_response.get("delivered_prefix")
|
||||
if isinstance(delivered_prefix, str) and delivered_prefix:
|
||||
self._last_sent_text = delivered_prefix
|
||||
@@ -570,13 +506,10 @@ class StreamTransportMixin:
|
||||
if self._is_flood_error(result):
|
||||
self._flood_strikes += 1
|
||||
self._current_edit_interval = min(self._current_edit_interval * 2, 10.0)
|
||||
logger.debug(
|
||||
"Flood control on edit (strike %d/%d), backoff interval → %.1fs",
|
||||
self._flood_strikes, self._MAX_FLOOD_STRIKES, self._current_edit_interval,
|
||||
)
|
||||
logger.debug("Flood control on edit (strike %d/%d), backoff interval → %.1fs",
|
||||
self._flood_strikes, self._MAX_FLOOD_STRIKES, self._current_edit_interval)
|
||||
immediate_final_fallback = (
|
||||
turn_final and getattr(self.adapter, "FALLBACK_ON_FINAL_EDIT_FLOOD", False) is True
|
||||
)
|
||||
turn_final and getattr(self.adapter, "FALLBACK_ON_FINAL_EDIT_FLOOD", False) is True)
|
||||
if self._flood_strikes < self._MAX_FLOOD_STRIKES and not immediate_final_fallback:
|
||||
self._last_edit_time = time.monotonic() # honor the new interval
|
||||
return False
|
||||
|
||||
Reference in New Issue
Block a user