Merge branch 'simp/r2-gw-rest' into simp/integration2
# Conflicts: # tests/gateway/test_48031_model_switch_after_auto_reset.py
This commit is contained in:
+442
-4501
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,485 @@
|
||||
"""Autonomy-loop gateway commands: /goal, /subgoal, /heartbeat, /loop, /refine, /review.
|
||||
|
||||
Split out of ``gateway/slash_commands.py``; bound onto ``GatewayRunner`` through
|
||||
``GatewaySlashCommandsMixin``. Origin internals are imported lazily (``from gateway.slash_commands
|
||||
import ...``) inside the bodies to avoid the import cycle.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from agent.i18n import t
|
||||
from gateway.platforms.base import MessageEvent, MessageType
|
||||
|
||||
# Log-record parity with gateway/run.py and the origin module.
|
||||
logger = logging.getLogger("gateway.run")
|
||||
|
||||
|
||||
class GatewayGoalCommandsMixin:
|
||||
"""Autonomy-loop gateway commands: /goal, /subgoal, /heartbeat, /loop, /refine, /review."""
|
||||
|
||||
async def _handle_goal_command(self, event: "MessageEvent") -> str:
|
||||
"""Handle /goal for gateway platforms.
|
||||
|
||||
Subcommands: status / pause / resume / clear. Setting a new goal queues the goal text as the
|
||||
next turn so the agent starts immediately; the post-turn continuation hook takes over after.
|
||||
"""
|
||||
args = (event.get_command_args() or "").strip()
|
||||
lower = args.lower()
|
||||
|
||||
mgr, session_entry = await self._get_goal_manager_for_event(event)
|
||||
if mgr is None:
|
||||
return t("gateway.goal.unavailable")
|
||||
|
||||
if not args or lower == "status":
|
||||
return mgr.status_line()
|
||||
if lower == "show":
|
||||
return f"{mgr.status_line()}\n{mgr.render_contract()}"
|
||||
if lower == "unwait":
|
||||
return "▶ Wait barrier cleared — goal loop resumes." if mgr.stop_waiting() else "No wait barrier set."
|
||||
if lower in {"clear", "stop", "done"}:
|
||||
had = mgr.has_goal()
|
||||
mgr.clear()
|
||||
self._clear_goal_continuations(event, "clear")
|
||||
return t("gateway.goal_cleared") if had else t("gateway.no_active_goal")
|
||||
if lower == "pause":
|
||||
state = mgr.pause(reason="user-paused")
|
||||
if state is None:
|
||||
return t("gateway.goal.no_goal_set")
|
||||
self._clear_goal_continuations(event, "pause")
|
||||
return t("gateway.goal.paused", goal=state.goal)
|
||||
if lower == "resume":
|
||||
return self._goal_resume(mgr, event)
|
||||
# Verb-prefixed forms take the remainder as their argument.
|
||||
for verb, handler in (("wait ", self._goal_wait), ("gate ", self._goal_gate)):
|
||||
if lower == verb.strip() or lower.startswith(verb):
|
||||
return handler(mgr, args[len(verb) - 1:].strip(), event)
|
||||
return await self._goal_set(mgr, args, lower, event)
|
||||
|
||||
def _clear_goal_continuations(self, event: MessageEvent, verb: str) -> None:
|
||||
try:
|
||||
adapter, _quick_key = self._adapter_and_key_for(event)
|
||||
if adapter and _quick_key:
|
||||
self._clear_goal_pending_continuations(_quick_key, adapter)
|
||||
except Exception as exc:
|
||||
logger.debug("goal %s: pending continuation cleanup failed: %s", verb, exc)
|
||||
|
||||
def _goal_resume(self, mgr, event: MessageEvent) -> str:
|
||||
state = mgr.resume()
|
||||
if state is None:
|
||||
return t("gateway.goal.no_resume")
|
||||
# Resume must restart work, not just flip persisted state: enqueue the canonical
|
||||
# continuation through the adapter FIFO — the same path the post-turn judge uses — so
|
||||
# the next turn fires as soon as this reply is delivered.
|
||||
prompt = mgr.next_continuation_prompt()
|
||||
try:
|
||||
adapter, _quick_key = self._adapter_and_key_for(event)
|
||||
if prompt and adapter and _quick_key:
|
||||
cont_event = MessageEvent(
|
||||
text=prompt,
|
||||
message_type=MessageType.TEXT,
|
||||
source=event.source,
|
||||
message_id=None,
|
||||
channel_prompt=None,
|
||||
)
|
||||
self._enqueue_fifo(_quick_key, cont_event, adapter)
|
||||
except Exception as exc:
|
||||
logger.debug("goal resume: continuation enqueue failed: %s", exc)
|
||||
return t("gateway.goal.resumed", goal=state.goal)
|
||||
|
||||
@staticmethod
|
||||
def _goal_wait(mgr, wait_arg: str, event: MessageEvent) -> str:
|
||||
"""/goal wait <pid> [reason] — park the loop on a background process."""
|
||||
if not wait_arg:
|
||||
return "Usage: /goal wait <pid> [reason]"
|
||||
wtokens = wait_arg.split(None, 1)
|
||||
try:
|
||||
pid = int(wtokens[0])
|
||||
except ValueError:
|
||||
return "/goal wait: <pid> must be an integer process id."
|
||||
reason = wtokens[1].strip() if len(wtokens) > 1 else ""
|
||||
try:
|
||||
mgr.wait_on(pid, reason=reason)
|
||||
except (RuntimeError, ValueError) as exc:
|
||||
return f"/goal wait: {exc}"
|
||||
rtxt = f" ({reason})" if reason else ""
|
||||
return f"⏳ Goal parked on pid {pid}{rtxt}. Loop pauses until it exits."
|
||||
|
||||
def _goal_gate(self, mgr, gate_arg: str, event: MessageEvent) -> str:
|
||||
"""/goal gate [list | add <command> | remove <N> | clear] — deterministic quality gates."""
|
||||
gate_lower = gate_arg.lower()
|
||||
if not gate_arg or gate_lower == "list":
|
||||
return mgr.render_gates()
|
||||
if gate_lower.startswith("add "):
|
||||
# SECURITY: a gate is persisted and later executed with shell=True at every goal turn
|
||||
# boundary (run_gate), with no approval prompt. Letting an allowed but non-admin gateway
|
||||
# sender choose that string is authenticated RCE under the Hermes process account — and
|
||||
# with no admin list configured (the backward-compatible default) every allowed sender
|
||||
# is treated as unrestricted. Gate ONLY this shell-creating operation behind a real,
|
||||
# explicitly-configured admin (the same fail-closed check that guards cross-origin
|
||||
# /resume); list/remove/clear stay open so a non-admin can still recover.
|
||||
if not self._resume_caller_is_admin(event.source):
|
||||
return (
|
||||
"⛔ /goal gate add requires an explicitly configured "
|
||||
"gateway admin (allow_admin_from for DMs, "
|
||||
"group_allow_admin_from for groups)."
|
||||
)
|
||||
try:
|
||||
gate = mgr.add_gate(gate_arg[len("add"):].strip())
|
||||
except (RuntimeError, ValueError) as exc:
|
||||
return f"/goal gate add: {exc}"
|
||||
return (
|
||||
f"⚿ Gate added: $ {gate.command} "
|
||||
f"({gate.max_retries} retries, {gate.timeout_seconds}s timeout). "
|
||||
f"It must pass before the goal can complete."
|
||||
)
|
||||
if gate_lower.startswith("remove ") or gate_lower.startswith("rm "):
|
||||
try:
|
||||
removed = mgr.remove_gate(int(gate_arg.split(None, 1)[1].strip()))
|
||||
except (RuntimeError, ValueError, IndexError) as exc:
|
||||
return f"/goal gate remove: {exc}"
|
||||
return f"✓ Gate removed: $ {removed}"
|
||||
if gate_lower == "clear":
|
||||
try:
|
||||
prev = mgr.clear_gates()
|
||||
except RuntimeError as exc:
|
||||
return f"/goal gate clear: {exc}"
|
||||
return f"✓ Cleared {prev} gate{'s' if prev != 1 else ''}."
|
||||
return "Usage: /goal gate [list | add <command> | remove <N> | clear]"
|
||||
|
||||
async def _goal_set(self, mgr, args: str, lower: str, event: MessageEvent) -> str:
|
||||
"""Set a new goal from free text, inline ``field: value`` contract lines, or ``draft <objective>``."""
|
||||
if lower.startswith("draft"):
|
||||
# Draft a structured completion contract, then set it. The aux LLM call is sync;
|
||||
# run it off the event loop.
|
||||
objective = args[len("draft"):].strip()
|
||||
if not objective:
|
||||
return "Usage: /goal draft <objective in plain language>"
|
||||
try:
|
||||
from hermes_cli.goals import draft_contract
|
||||
|
||||
# _run_in_executor_with_context, not a bare hop: drafting a contract calls the
|
||||
# auxiliary LLM, whose provider/credential resolution reads the profile secret scope
|
||||
# — a contextvar that a default-executor hop drops, leaving it unscoped.
|
||||
contract = await self._run_in_executor_with_context(draft_contract, objective)
|
||||
except Exception as exc:
|
||||
logger.debug("goal draft failed: %s", exc)
|
||||
contract = None
|
||||
args = objective # the goal text is the objective
|
||||
else:
|
||||
# Inline `field: value` lines parse into a completion contract; the remaining prose is
|
||||
# the goal headline. Plain free-form goals (no such lines) behave exactly as before.
|
||||
from hermes_cli.goals import parse_contract
|
||||
|
||||
headline, parsed = parse_contract(args)
|
||||
args = headline or args
|
||||
contract = parsed if not parsed.is_empty() else None
|
||||
|
||||
try:
|
||||
state = mgr.set(args, contract=contract)
|
||||
except ValueError as exc:
|
||||
return t("gateway.goal.invalid", error=str(exc))
|
||||
|
||||
# Queue the goal text as an immediate first turn so the agent starts making progress. The
|
||||
# post-turn hook takes over after.
|
||||
adapter, _quick_key = self._adapter_and_key_for(event)
|
||||
if adapter and _quick_key:
|
||||
try:
|
||||
kickoff_event = MessageEvent(
|
||||
text=state.goal,
|
||||
message_type=MessageType.TEXT,
|
||||
source=event.source,
|
||||
message_id=event.message_id,
|
||||
channel_prompt=event.channel_prompt,
|
||||
)
|
||||
self._enqueue_fifo(_quick_key, kickoff_event, adapter)
|
||||
except Exception as exc:
|
||||
logger.debug("goal kickoff enqueue failed: %s", exc)
|
||||
|
||||
base = t("gateway.goal.set", budget=state.max_turns, goal=state.goal)
|
||||
if state.has_contract():
|
||||
return f"{base}\nCompletion contract:\n{state.contract.render_block()}"
|
||||
if lower.startswith("draft"):
|
||||
# Drafting was requested but the aux model couldn't produce one.
|
||||
return f"{base}\n(Couldn't draft a contract — running as a free-form goal.)"
|
||||
return base
|
||||
|
||||
async def _handle_heartbeat_command(self, event: "MessageEvent") -> str:
|
||||
"""Handle /heartbeat for gateway platforms (mirror of CLI handler).
|
||||
|
||||
Manages the session's one recurring re-entry prompt. The gateway-wide poller injects due
|
||||
heartbeats through the adapter FIFO as ordinary user turns, so alternation and caching hold.
|
||||
"""
|
||||
from hermes_cli.heartbeat import parse_interval, format_interval, MIN_INTERVAL_SECONDS
|
||||
|
||||
args = (event.get_command_args() or "").strip()
|
||||
lower = args.lower()
|
||||
|
||||
mgr, session_entry = await self._get_heartbeat_manager_for_event(event)
|
||||
if mgr is None:
|
||||
return "Heartbeats unavailable (no session)."
|
||||
|
||||
quick_key = self._session_key_for_source(event.source) if event.source else None
|
||||
|
||||
if not args or lower == "status":
|
||||
return mgr.status_line()
|
||||
|
||||
if lower == "pause":
|
||||
state = mgr.pause()
|
||||
return f"⏸ Heartbeat paused: {state.prompt}" if state else "No heartbeat set."
|
||||
|
||||
if lower == "resume":
|
||||
state = mgr.resume()
|
||||
if state is None:
|
||||
return "No heartbeat to resume."
|
||||
if quick_key and event.source is not None:
|
||||
self._register_heartbeat_watch(quick_key, event.source, mgr.session_id)
|
||||
return f"▶ Heartbeat resumed (every {format_interval(state.interval_seconds)}): {state.prompt}"
|
||||
|
||||
if lower in {"clear", "stop", "off"}:
|
||||
had = mgr.clear()
|
||||
if quick_key:
|
||||
self._unregister_heartbeat_watch(quick_key)
|
||||
return "✓ Heartbeat cleared." if had else "No heartbeat set."
|
||||
|
||||
# Set: `/heartbeat every 10m <prompt>` (also accepts `10m <prompt>`).
|
||||
tokens = args.split(None, 2)
|
||||
interval = None
|
||||
prompt = ""
|
||||
if tokens and tokens[0].lower() == "every" and len(tokens) >= 2:
|
||||
interval = parse_interval(f"every {tokens[1]}")
|
||||
prompt = tokens[2] if len(tokens) > 2 else ""
|
||||
elif tokens:
|
||||
interval = parse_interval(tokens[0])
|
||||
prompt = args[len(tokens[0]):].strip() if interval and interval > 0 else ""
|
||||
|
||||
if interval is None:
|
||||
return (
|
||||
"Usage: /heartbeat every <interval> <prompt> (e.g. /heartbeat every 10m Check CI)\n"
|
||||
"Also: /heartbeat status | pause | resume | clear"
|
||||
)
|
||||
if interval < 0:
|
||||
return f"Interval too small — minimum is {MIN_INTERVAL_SECONDS}s."
|
||||
if not prompt.strip():
|
||||
return "Usage: /heartbeat every <interval> <prompt> — the prompt is required."
|
||||
|
||||
try:
|
||||
state = mgr.set(prompt, interval)
|
||||
except ValueError as exc:
|
||||
return f"Invalid heartbeat: {exc}"
|
||||
if quick_key and event.source is not None:
|
||||
self._register_heartbeat_watch(quick_key, event.source, mgr.session_id)
|
||||
return (
|
||||
f"♥ Heartbeat set (every {format_interval(state.interval_seconds)}): {state.prompt}\n"
|
||||
"Fires as a normal turn whenever this session is idle and the interval has "
|
||||
"elapsed. Lives while the gateway runs — use `hermes cron` for durable schedules."
|
||||
)
|
||||
|
||||
def _idle_cached_agent_or_error(self, event: MessageEvent, verb: str):
|
||||
"""``(session_key, cached_agent, None)`` for /refine and /review, or ``(_, _, error_text)``.
|
||||
|
||||
Both need a cached agent from a completed turn and refuse while a run is in flight.
|
||||
"""
|
||||
quick_key = self._session_key_for_source(event.source) if event.source else None
|
||||
if not quick_key:
|
||||
return None, None, f"{verb.capitalize()} unavailable (no session)."
|
||||
if quick_key in self._running_agents:
|
||||
return quick_key, None, f"Agent is running — wait for the turn to finish, then /{verb}."
|
||||
agent = self._cached_agent_for(quick_key)
|
||||
if agent is None:
|
||||
return quick_key, None, f"Nothing to {verb} yet — send a message first."
|
||||
return quick_key, agent, None
|
||||
|
||||
async def _handle_refine_command(self, event: "MessageEvent") -> str:
|
||||
"""Handle /refine — run the memory/skill review fork on demand.
|
||||
|
||||
Runs in a daemon thread against a snapshot of the cached AIAgent's conversation; the live
|
||||
session and prompt cache are untouched. Requires at least one completed turn.
|
||||
"""
|
||||
args = (event.get_command_args() or "").strip()
|
||||
quick_key, agent, error = self._idle_cached_agent_or_error(event, "refine")
|
||||
if error:
|
||||
return error
|
||||
|
||||
snapshot = list(getattr(agent, "_session_messages", None) or [])
|
||||
if not snapshot:
|
||||
return "Nothing to refine yet — the conversation is empty."
|
||||
|
||||
review_skills = "skill_manage" in getattr(agent, "valid_tool_names", set())
|
||||
try:
|
||||
agent._spawn_background_review(
|
||||
messages_snapshot=snapshot,
|
||||
review_memory=True,
|
||||
review_skills=review_skills,
|
||||
focus=args or None,
|
||||
)
|
||||
except Exception as exc:
|
||||
return f"/refine failed to start: {exc}"
|
||||
tail = f" (focus: {args})" if args else ""
|
||||
return (
|
||||
f"⚗ Reviewing this conversation in the background{tail} — "
|
||||
f"any memory/skill updates will be reported when done."
|
||||
)
|
||||
|
||||
async def _handle_review_command(self, event: "MessageEvent") -> str:
|
||||
"""Handle /review — spawn an independent reviewer subagent.
|
||||
|
||||
The approval session-key contextvar is only bound during agent turns, so bind it explicitly
|
||||
here or the completion event carries no gateway route and never re-enters this chat.
|
||||
"""
|
||||
args = (event.get_command_args() or "").strip()
|
||||
quick_key, agent, error = self._idle_cached_agent_or_error(event, "review")
|
||||
if error:
|
||||
return error
|
||||
|
||||
snapshot = list(getattr(agent, "_session_messages", None) or [])
|
||||
|
||||
from tools.approval import (
|
||||
reset_current_session_key,
|
||||
set_current_session_key,
|
||||
)
|
||||
|
||||
def _dispatch():
|
||||
token = set_current_session_key(quick_key)
|
||||
try:
|
||||
from agent.review_engine import start_review
|
||||
|
||||
return start_review(agent, snapshot, args)
|
||||
finally:
|
||||
reset_current_session_key(token)
|
||||
|
||||
try:
|
||||
# _run_in_executor_with_context, not a bare hop: the reviewer
|
||||
# subagent is spawned from the worker and inherits its context,
|
||||
# so a bare hop would run it under the launch home / no secret scope.
|
||||
result = await self._run_in_executor_with_context(_dispatch)
|
||||
except ValueError as exc:
|
||||
return str(exc)
|
||||
except Exception as exc:
|
||||
return f"/review failed to start: {exc}"
|
||||
|
||||
from agent.review_engine import format_dispatch_note
|
||||
|
||||
return format_dispatch_note(result, args)
|
||||
|
||||
async def _handle_subgoal_command(self, event: "MessageEvent") -> str:
|
||||
"""Handle /subgoal for gateway platforms (mirror of CLI handler).
|
||||
|
||||
Subgoals are extra criteria appended to the active goal mid-loop. They modify state read
|
||||
at the next turn boundary, so this is safe to invoke while the agent is running.
|
||||
"""
|
||||
args = (event.get_command_args() or "").strip()
|
||||
mgr, _session_entry = await self._get_goal_manager_for_event(event)
|
||||
if mgr is None:
|
||||
return t("gateway.goal.unavailable")
|
||||
if not mgr.has_goal():
|
||||
return "No active goal. Set one with /goal <text>."
|
||||
|
||||
# No args → list current subgoals.
|
||||
if not args:
|
||||
return f"{mgr.status_line()}\n{mgr.render_subgoals()}"
|
||||
|
||||
tokens = args.split(None, 1)
|
||||
verb = tokens[0].lower()
|
||||
rest = tokens[1].strip() if len(tokens) > 1 else ""
|
||||
|
||||
if verb == "remove":
|
||||
if not rest:
|
||||
return "Usage: /subgoal remove <n>"
|
||||
try:
|
||||
idx = int(rest.split()[0])
|
||||
except ValueError:
|
||||
return "/subgoal remove: <n> must be an integer (1-based index)."
|
||||
try:
|
||||
removed = mgr.remove_subgoal(idx)
|
||||
except (IndexError, RuntimeError) as exc:
|
||||
return f"/subgoal remove: {exc}"
|
||||
return f"✓ Removed subgoal {idx}: {removed}"
|
||||
|
||||
if verb == "clear":
|
||||
try:
|
||||
prev = mgr.clear_subgoals()
|
||||
except RuntimeError as exc:
|
||||
return f"/subgoal clear: {exc}"
|
||||
if prev:
|
||||
return f"✓ Cleared {prev} subgoal{'s' if prev != 1 else ''}."
|
||||
return "No subgoals to clear."
|
||||
|
||||
try:
|
||||
text = mgr.add_subgoal(args)
|
||||
except (ValueError, RuntimeError) as exc:
|
||||
return f"/subgoal: {exc}"
|
||||
idx = len(mgr.state.subgoals) if mgr.state else 0
|
||||
return f"✓ Added subgoal {idx}: {text}"
|
||||
|
||||
async def _get_loop_manager_for_event(self, event: "MessageEvent"):
|
||||
"""Return a LoopManager bound to the session for this gateway event.
|
||||
|
||||
Returns ``(manager, session_entry)``, or ``(None, None)`` when the loops module or session
|
||||
can't be loaded. Mirrors ``_get_goal_manager_for_event``.
|
||||
"""
|
||||
try:
|
||||
from hermes_cli.loops import LoopManager
|
||||
except Exception as exc:
|
||||
logger.debug("loop manager unavailable: %s", exc)
|
||||
return None, None
|
||||
# Warm the SessionDB cache off-loop. A cold cache drops the first
|
||||
# /loop write while the reply claims the loop was set (same class
|
||||
# as the /goal false-ack fix).
|
||||
await self._warm_goals_session_db("loop manager")
|
||||
try:
|
||||
session_entry = await self.async_session_store.get_or_create_session(event.source)
|
||||
except Exception:
|
||||
return None, None
|
||||
sid = getattr(session_entry, "session_id", None) or ""
|
||||
if not sid:
|
||||
return None, None
|
||||
return LoopManager(session_id=sid), session_entry
|
||||
|
||||
async def _handle_loop_command(self, event: "MessageEvent") -> str:
|
||||
"""Handle /loop for gateway platforms — recurring in-session wakeups.
|
||||
|
||||
Mirrors the CLI handler via ``dispatch_loop_command``. New loops capture the event's routing
|
||||
(platform/chat/thread) so the idle loop-wakeup watcher can inject ticks here after a restart.
|
||||
"""
|
||||
try:
|
||||
from hermes_cli.loops import dispatch_loop_command, goal_blocks_loop_tick
|
||||
except Exception as exc:
|
||||
logger.debug("loops module unavailable: %s", exc)
|
||||
return "Loops unavailable."
|
||||
|
||||
mgr, _session_entry = await self._get_loop_manager_for_event(event)
|
||||
if mgr is None:
|
||||
return "Loops unavailable (no active session)."
|
||||
|
||||
route: dict = {}
|
||||
try:
|
||||
src = event.source
|
||||
if src is not None:
|
||||
platform = getattr(src, "platform", "")
|
||||
route = {
|
||||
"platform": platform.value if hasattr(platform, "value") else str(platform or ""),
|
||||
"chat_id": str(getattr(src, "chat_id", "") or ""),
|
||||
"chat_type": str(getattr(src, "chat_type", "") or ""),
|
||||
"thread_id": str(getattr(src, "thread_id", "") or ""),
|
||||
"user_id": str(getattr(src, "user_id", "") or ""),
|
||||
"user_name": str(getattr(src, "user_name", "") or ""),
|
||||
}
|
||||
route = {k: v for k, v in route.items() if v}
|
||||
except Exception:
|
||||
route = {}
|
||||
|
||||
args = (event.get_command_args() or "").strip()
|
||||
result = dispatch_loop_command(mgr, args, route=route)
|
||||
output = result.get("output") or ""
|
||||
if result.get("created"):
|
||||
try:
|
||||
if goal_blocks_loop_tick(mgr.session_id):
|
||||
output += (
|
||||
"\nNote: an active /goal is driving this session — loop "
|
||||
"wakeups defer until the goal finishes, pauses, or parks."
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return output
|
||||
@@ -0,0 +1,981 @@
|
||||
"""Gateway slash commands that switch or tune the model route:
|
||||
/model, /codex-runtime, /reasoning, /fast, /personality.
|
||||
|
||||
Split out of ``gateway/slash_commands.py``; bound onto ``GatewayRunner`` through
|
||||
``GatewaySlashCommandsMixin``. Origin internals are imported lazily (``from gateway.slash_commands
|
||||
import ...``) inside the bodies to avoid the import cycle.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import asyncio
|
||||
from typing import Optional
|
||||
|
||||
from agent.i18n import t
|
||||
from gateway.platforms.base import MessageEvent
|
||||
from hermes_cli.config import atomic_config_write, clear_model_endpoint_credentials
|
||||
from utils import base_url_host_matches
|
||||
|
||||
# Log-record parity with gateway/run.py and the origin module.
|
||||
logger = logging.getLogger("gateway.run")
|
||||
|
||||
|
||||
def _model_switch_skew_guard() -> Optional[str]:
|
||||
"""Refuse a model switch when the gateway is running stale code.
|
||||
|
||||
A long-lived gateway keeps boot-time modules in memory; if the checkout changed underneath it,
|
||||
a first-time lazy import on a new code path can crash on a stale cached dependency. Detect the
|
||||
drift and ask for a restart. Scoped to model switching only (the highest-risk trigger).
|
||||
"""
|
||||
from gateway.code_skew import detect_code_skew
|
||||
|
||||
skew = detect_code_skew()
|
||||
if not skew:
|
||||
return None
|
||||
boot_rev, disk_rev = skew
|
||||
return t(
|
||||
"gateway.model.error_prefix",
|
||||
error=(
|
||||
f"This gateway is running code from {boot_rev} but the checkout on "
|
||||
f"disk is now {disk_rev}. Switching models would risk a stale-module "
|
||||
f"crash — restart the gateway to load the new code: hermes gateway restart"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def _persist_model_switch_to_config(result, config_path) -> None:
|
||||
"""Write-through a resolved /model switch to ``config_path`` (model.default/provider/base_url).
|
||||
|
||||
Write-back round-trip: raw read is correct (merged defaults must not be persisted back to the
|
||||
user's file). A scalar/None ``model:`` is coerced into a dict first — otherwise
|
||||
``cfg.setdefault("model", {})`` returns the existing scalar and the next assignment raises
|
||||
``TypeError``. Named providers re-resolve base_url/api_mode fresh, so leftovers are cleared
|
||||
unconditionally; custom providers have no registry entry to re-derive from, so they need an
|
||||
explicit set-or-clear (a lone ``if base_url:`` leaves stale values).
|
||||
"""
|
||||
from hermes_cli.config import read_user_config_raw, save_config
|
||||
|
||||
cfg = read_user_config_raw(config_path)
|
||||
raw_model = cfg.get("model")
|
||||
if isinstance(raw_model, dict):
|
||||
model_cfg = raw_model
|
||||
elif isinstance(raw_model, str) and raw_model.strip():
|
||||
model_cfg = cfg["model"] = {"default": raw_model.strip()}
|
||||
else:
|
||||
model_cfg = cfg["model"] = {}
|
||||
try:
|
||||
from hermes_cli.route_identity import should_clear_context_pin_async
|
||||
|
||||
if await should_clear_context_pin_async(
|
||||
model_cfg.get("default") or model_cfg.get("model"),
|
||||
result.new_model,
|
||||
model_cfg.get("base_url"),
|
||||
result.base_url,
|
||||
model_cfg.get("provider"),
|
||||
result.target_provider,
|
||||
):
|
||||
model_cfg.pop("context_length", None)
|
||||
except Exception:
|
||||
model_cfg.pop("context_length", None)
|
||||
model_cfg["default"] = result.new_model
|
||||
model_cfg["provider"] = result.target_provider
|
||||
is_custom_target = str(result.target_provider or "").strip().lower() == "custom"
|
||||
if result.base_url:
|
||||
model_cfg["base_url"] = result.base_url
|
||||
elif is_custom_target:
|
||||
model_cfg.pop("base_url", None)
|
||||
if is_custom_target:
|
||||
if result.api_mode:
|
||||
model_cfg["api_mode"] = result.api_mode
|
||||
else:
|
||||
model_cfg.pop("api_mode", None)
|
||||
else:
|
||||
clear_model_endpoint_credentials(model_cfg, clear_base_url=True)
|
||||
save_config(cfg)
|
||||
|
||||
|
||||
def _read_model_command_config(config_path):
|
||||
"""Current (model, provider, base_url, user_providers, custom_providers, excluded) for /model.
|
||||
|
||||
Fail-open: any config read error yields the defaults (``provider="openrouter"``).
|
||||
"""
|
||||
from gateway.run import _load_gateway_config
|
||||
|
||||
current_model, current_provider, current_base_url = "", "openrouter", ""
|
||||
user_provs = custom_provs = None
|
||||
excluded_provs: list = []
|
||||
try:
|
||||
cfg = _load_gateway_config(config_path=config_path)
|
||||
if cfg:
|
||||
model_cfg = cfg.get("model", {})
|
||||
if isinstance(model_cfg, dict):
|
||||
current_model = model_cfg.get("default", "")
|
||||
current_provider = model_cfg.get("provider", current_provider)
|
||||
current_base_url = model_cfg.get("base_url", "")
|
||||
user_provs = cfg.get("providers")
|
||||
try:
|
||||
from hermes_cli.config import get_compatible_custom_providers
|
||||
custom_provs = get_compatible_custom_providers(cfg)
|
||||
except Exception:
|
||||
custom_provs = cfg.get("custom_providers")
|
||||
_excl = cfg.get("model_catalog", {}).get("excluded_providers")
|
||||
if isinstance(_excl, list):
|
||||
excluded_provs = _excl
|
||||
except Exception:
|
||||
pass
|
||||
return current_model, current_provider, current_base_url, user_provs, custom_provs, excluded_provs
|
||||
|
||||
|
||||
def _model_provider_listing_lines(providers) -> list[str]:
|
||||
"""Text-list body for ``/model`` with no args on platforms without a picker."""
|
||||
lines: list[str] = []
|
||||
for p in providers:
|
||||
tag = t("gateway.model.current_tag") if p["is_current"] else ""
|
||||
lines.append(f"**{p['name']}** `--provider {p['slug']}`{tag}:")
|
||||
if p["models"]:
|
||||
model_strs = ", ".join(f"`{m}`" for m in p["models"])
|
||||
extra = t("gateway.model.more_models_suffix", count=p["total_models"] - len(p["models"])) if p["total_models"] > len(p["models"]) else ""
|
||||
lines.append(f" {model_strs}{extra}")
|
||||
elif p.get("api_url"):
|
||||
lines.append(f" `{p['api_url']}`")
|
||||
lines.append("")
|
||||
return lines
|
||||
|
||||
|
||||
class GatewayModelCommandsMixin:
|
||||
"""Model-route slash commands (/model, /codex-runtime, /reasoning, /fast, /personality)."""
|
||||
|
||||
async def _perform_model_switch(
|
||||
self,
|
||||
switch_model,
|
||||
*,
|
||||
raw_input: str,
|
||||
explicit_provider,
|
||||
session_key: str,
|
||||
source,
|
||||
current_model,
|
||||
current_provider,
|
||||
current_base_url,
|
||||
current_api_key,
|
||||
persist_global: bool,
|
||||
user_provs,
|
||||
custom_provs,
|
||||
):
|
||||
"""Resolve a /model switch off-loop. Returns ``(result, None)`` or ``(None, error_text)``."""
|
||||
from gateway.run import _load_gateway_config
|
||||
|
||||
skew_error = _model_switch_skew_guard()
|
||||
if skew_error:
|
||||
return None, skew_error
|
||||
# Offload the switch off the event loop — switch_model() can fall through to a synchronous
|
||||
# models.dev HTTP fetch (requests.get, 15s timeout) on a cold/expired cache, which freezes
|
||||
# the gateway otherwise.
|
||||
result = await asyncio.to_thread(
|
||||
switch_model,
|
||||
raw_input=raw_input,
|
||||
current_provider=current_provider,
|
||||
current_model=current_model,
|
||||
current_base_url=current_base_url,
|
||||
current_api_key=current_api_key,
|
||||
is_global=persist_global,
|
||||
explicit_provider=explicit_provider,
|
||||
user_providers=user_provs,
|
||||
custom_providers=custom_provs,
|
||||
)
|
||||
if not result.success:
|
||||
return None, t("gateway.model.error_prefix", error=result.error_message)
|
||||
try:
|
||||
from hermes_cli.context_switch_guard import enrich_model_switch_warnings_for_gateway
|
||||
|
||||
# Offload: merge_preflight_compression_warning() calls the sync
|
||||
# resolve_display_context_length() provider probe ladder — must not run on the loop.
|
||||
await asyncio.to_thread(
|
||||
enrich_model_switch_warnings_for_gateway,
|
||||
result,
|
||||
self,
|
||||
session_key=session_key,
|
||||
source=source,
|
||||
custom_providers=custom_provs,
|
||||
load_gateway_config=_load_gateway_config,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("preflight-compression switch warning failed: %s", exc)
|
||||
return result, None
|
||||
|
||||
async def _commit_model_switch(
|
||||
self,
|
||||
result,
|
||||
*,
|
||||
session_key: str,
|
||||
source,
|
||||
current_model,
|
||||
current_base_url,
|
||||
current_api_key,
|
||||
custom_provs,
|
||||
persist_global: bool,
|
||||
config_path,
|
||||
one_turn: bool = False,
|
||||
restore_snapshot=None,
|
||||
picker: bool = False,
|
||||
) -> str:
|
||||
"""Apply a resolved switch (cached agent, session, config) and build the confirmation.
|
||||
|
||||
Shared by the typed ``/model <name>`` path and the picker callback (``picker=True``).
|
||||
"""
|
||||
from gateway.run import _load_gateway_config
|
||||
from hermes_cli.model_switch import format_model_for_display, resolve_display_context_length_async
|
||||
|
||||
# If there's a cached agent, update it in-place
|
||||
cached_agent = self._cached_agent_for(session_key)
|
||||
if cached_agent is not None:
|
||||
try:
|
||||
cached_agent.switch_model(
|
||||
new_model=result.new_model,
|
||||
new_provider=result.target_provider,
|
||||
api_key=result.api_key,
|
||||
base_url=result.base_url,
|
||||
api_mode=result.api_mode,
|
||||
capabilities=getattr(result, "runtime_capabilities", None),
|
||||
)
|
||||
except Exception as exc:
|
||||
# In-place swap rolled back to the OLD working model/client and re-raised. Abort the
|
||||
# commit (DB persist, session override, cache eviction, config write) so a failed switch
|
||||
# is a no-op — otherwise the next message rebuilds a broken agent from the override.
|
||||
logger.warning(
|
||||
"%s model switch failed for cached agent: %s", "Picker" if picker else "In-place", exc
|
||||
)
|
||||
return t(
|
||||
"gateway.model.error_prefix",
|
||||
error=f"Model switch to {result.new_model} failed ({exc}); staying on {current_model}.",
|
||||
)
|
||||
|
||||
# Persist the new model to the session DB so the dashboard shows the updated model.
|
||||
_sess_db = getattr(self, "_session_db", None)
|
||||
if _sess_db is not None:
|
||||
try:
|
||||
_sess_entry = await self.async_session_store.get_or_create_session(source)
|
||||
# Typed path: if this session was auto-reset, consume the flag so the next regular
|
||||
# message's cleanup does not wipe the model override just stored below.
|
||||
if not picker and getattr(_sess_entry, "was_auto_reset", False):
|
||||
_sess_entry.was_auto_reset = False
|
||||
await _sess_db.update_session_model(
|
||||
_sess_entry.session_id, result.new_model,
|
||||
provider=result.target_provider,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("Failed to persist model switch to DB: %s", exc)
|
||||
|
||||
# Store a note to prepend to the next user message so the model knows about the switch
|
||||
# (avoids system messages mid-history). Display form strips opaque Palantir RID
|
||||
# prefixes; the override map below keeps the full ID for the wire.
|
||||
if not hasattr(self, "_pending_model_notes"):
|
||||
self._pending_model_notes = {}
|
||||
self._pending_model_notes[session_key] = (
|
||||
f"[Note: model was just switched from {format_model_for_display(current_model)} to "
|
||||
f"{format_model_for_display(result.new_model)} "
|
||||
f"via {result.provider_label or result.target_provider}. "
|
||||
f"{'This override applies to the next turn only. ' if one_turn else ''}"
|
||||
f"Adjust your self-identification accordingly.]"
|
||||
)
|
||||
|
||||
# Store session override so next agent creation uses the new model
|
||||
self._session_model_overrides[session_key] = {
|
||||
"model": result.new_model,
|
||||
"provider": result.target_provider,
|
||||
"api_key": result.api_key,
|
||||
"base_url": result.base_url,
|
||||
"api_mode": result.api_mode,
|
||||
"request_overrides": dict(result.request_overrides or {}),
|
||||
"capabilities": dict(result.runtime_capabilities or {}),
|
||||
}
|
||||
if one_turn:
|
||||
if not hasattr(self, "_pending_one_turn_model_restores"):
|
||||
self._pending_one_turn_model_restores = {}
|
||||
self._pending_one_turn_model_restores[session_key] = (
|
||||
restore_snapshot or {"had_override": False, "override": None}
|
||||
)
|
||||
elif not picker and hasattr(self, "_pending_one_turn_model_restores"):
|
||||
self._pending_one_turn_model_restores.pop(session_key, None)
|
||||
|
||||
# Write-through the non-secret parts (model/provider/base_url) so the override survives a
|
||||
# restart; api_key/api_mode are never persisted (re-resolved on rehydration). /model --once is
|
||||
# EXCLUDED: a one-turn override must not outlive a restart; the pre-once value stays persisted.
|
||||
if not one_turn:
|
||||
try:
|
||||
await self.async_session_store.set_model_override(
|
||||
session_key, self._session_model_overrides[session_key]
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to persist session model override", exc_info=True)
|
||||
|
||||
# Evict cached agent so the next turn creates a fresh agent from the
|
||||
# override rather than relying on cache signature mismatch detection.
|
||||
self._evict_cached_agent(session_key)
|
||||
|
||||
# Persist to config (default) unless --session opted out
|
||||
if persist_global:
|
||||
try:
|
||||
await _persist_model_switch_to_config(result, config_path)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to persist model switch: %s", e)
|
||||
|
||||
# Build confirmation message with full metadata. Display form shortens opaque Palantir
|
||||
# IDs (ri.language-model-service..*) to their trailing slug.
|
||||
provider_label = result.provider_label or result.target_provider
|
||||
lines = [t("gateway.model.switched", model=format_model_for_display(result.new_model))]
|
||||
lines.append(t("gateway.model.provider_label", provider=provider_label))
|
||||
|
||||
# Context: always resolve via the provider-aware chain so Codex OAuth,
|
||||
# Copilot, and Nous-enforced caps win over the raw models.dev entry.
|
||||
mi = result.model_info
|
||||
_sw_config_ctx = None
|
||||
_sw_model_cfg = {}
|
||||
try:
|
||||
_sw_model_cfg = _load_gateway_config().get("model", {})
|
||||
if isinstance(_sw_model_cfg, dict):
|
||||
_sw_raw = _sw_model_cfg.get("context_length")
|
||||
if _sw_raw is not None:
|
||||
_sw_config_ctx = int(_sw_raw)
|
||||
except Exception:
|
||||
pass
|
||||
if not isinstance(_sw_model_cfg, dict):
|
||||
_sw_model_cfg = {}
|
||||
ctx = await resolve_display_context_length_async(
|
||||
result.new_model,
|
||||
result.target_provider,
|
||||
base_url=result.base_url or current_base_url or "",
|
||||
api_key=result.api_key or current_api_key or "",
|
||||
model_info=mi,
|
||||
custom_providers=custom_provs,
|
||||
config_context_length=_sw_config_ctx,
|
||||
configured_model=_sw_model_cfg.get("default") or _sw_model_cfg.get("model"),
|
||||
configured_provider=_sw_model_cfg.get("provider"),
|
||||
configured_base_url=_sw_model_cfg.get("base_url"),
|
||||
)
|
||||
if ctx:
|
||||
lines.append(t("gateway.model.context_label", tokens=f"{ctx:,}"))
|
||||
if mi:
|
||||
if mi.max_output:
|
||||
lines.append(t("gateway.model.max_output_label", tokens=f"{mi.max_output:,}"))
|
||||
lines.append(t("gateway.model.capabilities_label", capabilities=mi.format_capabilities()))
|
||||
|
||||
if not picker:
|
||||
cache_enabled = (
|
||||
(base_url_host_matches(result.base_url or "", "openrouter.ai") and "claude" in result.new_model.lower())
|
||||
or result.api_mode == "anthropic_messages"
|
||||
)
|
||||
if cache_enabled:
|
||||
lines.append(t("gateway.model.prompt_caching_enabled"))
|
||||
|
||||
if result.warning_message:
|
||||
lines.append(t("gateway.model.warning_prefix", warning=result.warning_message))
|
||||
|
||||
if persist_global:
|
||||
lines.append(t("gateway.model.saved_global"))
|
||||
elif one_turn:
|
||||
lines.append(" (next turn only — restores after one response)")
|
||||
else:
|
||||
lines.append(t("gateway.model.session_only_hint"))
|
||||
return "\n".join(lines)
|
||||
|
||||
async def _send_model_picker(self, event: MessageEvent, source, adapter, session_key: str, listing_kwargs: dict, on_model_selected) -> bool:
|
||||
"""Send the interactive /model picker; False when nothing was sent (text fallback).
|
||||
|
||||
*source* is the session-key-normalized source (Telegram topic recovery), so the picker's
|
||||
thread metadata lands where the next turn reads.
|
||||
"""
|
||||
from hermes_cli.model_switch import list_picker_providers
|
||||
|
||||
try:
|
||||
# Offload blocking provider-listing (can fall through to a synchronous urllib HTTP fetch
|
||||
# on a stale cache) off the event loop so the gateway doesn't freeze. See #41289.
|
||||
providers = await asyncio.to_thread(
|
||||
list_picker_providers, max_models=50, include_moa=True, **listing_kwargs
|
||||
)
|
||||
except Exception:
|
||||
providers = []
|
||||
if not providers:
|
||||
return False
|
||||
result = await adapter.send_model_picker(
|
||||
chat_id=source.chat_id,
|
||||
providers=providers,
|
||||
current_model=listing_kwargs["current_model"],
|
||||
current_provider=listing_kwargs["current_provider"],
|
||||
session_key=session_key,
|
||||
on_model_selected=on_model_selected,
|
||||
metadata=self._thread_metadata_for_source(source, self._reply_anchor_for_event(event)),
|
||||
)
|
||||
return bool(result.success)
|
||||
|
||||
async def _handle_model_command(self, event: MessageEvent) -> Optional[str]:
|
||||
"""Handle /model command — switch model."""
|
||||
from gateway.run import _hermes_home
|
||||
from hermes_cli.model_switch import (
|
||||
switch_model as _switch_model, parse_model_switch_args,
|
||||
resolve_persist_behavior,
|
||||
list_authenticated_providers,
|
||||
)
|
||||
from hermes_cli.providers import get_label
|
||||
|
||||
raw_args = event.get_command_args().strip()
|
||||
source = event.source
|
||||
_command_profile_home = None
|
||||
if getattr(getattr(self, "config", None), "multiplex_profiles", False):
|
||||
_command_profile_home = self._resolve_profile_home_for_source(source)
|
||||
|
||||
# Parse --provider, --global, --session, --once, and --refresh flags
|
||||
# via the shared single-owner parser (hermes_cli.model_switch).
|
||||
request = parse_model_switch_args(raw_args)
|
||||
model_input = request.target
|
||||
explicit_provider = request.explicit_provider
|
||||
is_global_flag = request.is_global
|
||||
force_refresh = request.force_refresh
|
||||
is_session = request.is_session
|
||||
one_turn = request.is_once
|
||||
if request.errors:
|
||||
# Gateway decoration: "❌ " prefix over the canonical error copy.
|
||||
return f"❌ {request.error_messages()[0]}"
|
||||
persist_global = resolve_persist_behavior(
|
||||
is_global_flag,
|
||||
is_session,
|
||||
is_once=one_turn,
|
||||
explicit_provider=explicit_provider,
|
||||
)
|
||||
|
||||
# --refresh: bust the disk cache so the picker shows live data.
|
||||
if force_refresh:
|
||||
try:
|
||||
from hermes_cli.models import clear_provider_models_cache
|
||||
clear_provider_models_cache()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Read current model/provider from config
|
||||
config_path = (_command_profile_home or _hermes_home) / "config.yaml"
|
||||
current_model, current_provider, current_base_url, user_provs, custom_provs, excluded_provs = (
|
||||
_read_model_command_config(config_path)
|
||||
)
|
||||
current_api_key = ""
|
||||
|
||||
# Check for session override. Normalize the source the same way a normal message turn does
|
||||
# (Telegram DM topic recovery) before deriving the override key, so the override is stored
|
||||
# under the key the next message turn reads.
|
||||
source = await asyncio.to_thread(self._normalize_source_for_session_key, source)
|
||||
session_key = self._session_key_for_source(source)
|
||||
override = self._session_model_overrides.get(session_key, {})
|
||||
restore_snapshot = (
|
||||
self._snapshot_session_model_override(session_key) if one_turn else None
|
||||
)
|
||||
if override:
|
||||
current_model = override.get("model", current_model)
|
||||
current_provider = override.get("provider", current_provider)
|
||||
current_base_url = override.get("base_url", current_base_url)
|
||||
current_api_key = override.get("api_key", current_api_key)
|
||||
|
||||
async def perform_switch(model_id: str, provider_slug, *, src=source):
|
||||
return await self._perform_model_switch(
|
||||
_switch_model,
|
||||
raw_input=model_id,
|
||||
explicit_provider=provider_slug,
|
||||
session_key=session_key,
|
||||
source=src,
|
||||
current_model=current_model,
|
||||
current_provider=current_provider,
|
||||
current_base_url=current_base_url,
|
||||
current_api_key=current_api_key,
|
||||
persist_global=persist_global,
|
||||
user_provs=user_provs,
|
||||
custom_provs=custom_provs,
|
||||
)
|
||||
|
||||
async def commit_switch(result, *, picker: bool = False, src=source) -> str:
|
||||
"""Apply the resolved switch (agent, session, config) and build the reply."""
|
||||
return await self._commit_model_switch(
|
||||
result,
|
||||
session_key=session_key,
|
||||
source=src,
|
||||
current_model=current_model,
|
||||
current_base_url=current_base_url,
|
||||
current_api_key=current_api_key,
|
||||
custom_provs=custom_provs,
|
||||
persist_global=persist_global,
|
||||
config_path=config_path,
|
||||
one_turn=False if picker else one_turn,
|
||||
restore_snapshot=None if picker else restore_snapshot,
|
||||
picker=picker,
|
||||
)
|
||||
|
||||
async def switch_and_commit(model_id: str, provider_slug, *, picker: bool) -> str:
|
||||
# The picker callback binds the raw event source (pre-normalization), as it always has.
|
||||
src = event.source if picker else source
|
||||
result, error = await perform_switch(model_id, provider_slug, src=src)
|
||||
if error is not None:
|
||||
return error
|
||||
return await commit_switch(result, picker=picker, src=src)
|
||||
|
||||
# No args: show interactive picker (Telegram/Discord) or text list
|
||||
if not model_input and not explicit_provider:
|
||||
listing_kwargs = dict(
|
||||
current_provider=current_provider,
|
||||
current_base_url=current_base_url,
|
||||
current_model=current_model,
|
||||
user_providers=user_provs,
|
||||
custom_providers=custom_provs,
|
||||
excluded_providers=excluded_provs,
|
||||
)
|
||||
# Try interactive picker if the platform supports it
|
||||
adapter = self._adapter_for_source(source)
|
||||
if adapter is not None and getattr(type(adapter), "send_model_picker", None) is not None:
|
||||
async def _on_model_selected(_chat_id: str, model_id: str, provider_slug: str) -> str:
|
||||
"""Perform the model switch and return confirmation text."""
|
||||
if _command_profile_home is None:
|
||||
return await switch_and_commit(model_id, provider_slug, picker=True)
|
||||
from gateway.run import _profile_runtime_scope
|
||||
|
||||
with _profile_runtime_scope(_command_profile_home):
|
||||
return await switch_and_commit(model_id, provider_slug, picker=True)
|
||||
|
||||
if await self._send_model_picker(event, source, adapter, session_key, listing_kwargs, _on_model_selected):
|
||||
return None # Picker sent — adapter handles the response
|
||||
|
||||
# Fallback: text list (for platforms without picker or if picker failed)
|
||||
lines = [t("gateway.model.current_label", model=current_model or "unknown", provider=get_label(current_provider)), ""]
|
||||
try:
|
||||
# Offload blocking provider-listing off the event loop so the
|
||||
# gateway doesn't freeze on a stale-cache HTTP fetch. See #41289.
|
||||
providers = await asyncio.to_thread(list_authenticated_providers, max_models=5, **listing_kwargs)
|
||||
lines.extend(_model_provider_listing_lines(providers))
|
||||
except Exception:
|
||||
pass
|
||||
lines.append(t("gateway.model.usage_switch_model"))
|
||||
lines.append(t("gateway.model.usage_switch_provider"))
|
||||
lines.append(t("gateway.model.usage_persist"))
|
||||
return "\n".join(lines)
|
||||
|
||||
# Perform the switch
|
||||
result, error = await perform_switch(model_input, explicit_provider)
|
||||
if error is not None:
|
||||
return error
|
||||
|
||||
# Selection-guard confirmation for the typed /model <name> path (pickers confirm via their own
|
||||
# UI). Runs the unified registry (cost + data-policy guards); pricing lookups may hit
|
||||
# models.dev or a /models endpoint on a cache miss, so run it off the event loop.
|
||||
_cost_warning = None
|
||||
try:
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
|
||||
_cost_warning = await asyncio.to_thread(
|
||||
combined_selection_warning,
|
||||
result.new_model,
|
||||
provider=result.target_provider,
|
||||
base_url=result.base_url or current_base_url or "",
|
||||
api_key=result.api_key or current_api_key or "",
|
||||
model_info=result.model_info,
|
||||
)
|
||||
except Exception:
|
||||
_cost_warning = None
|
||||
if _cost_warning is not None:
|
||||
async def _on_cost_confirm(choice: str) -> str:
|
||||
if choice == "cancel":
|
||||
return (
|
||||
f"🟡 Model switch cancelled. Current model unchanged "
|
||||
f"({current_model or 'unknown'})."
|
||||
)
|
||||
# "once" and "always" both proceed — there is no persistent
|
||||
# opt-out for selection guards (each guarded switch should be
|
||||
# an explicit decision).
|
||||
return await commit_switch(result)
|
||||
|
||||
_p = self._typed_command_prefix_for(event.source.platform)
|
||||
return await self._request_slash_confirm(
|
||||
event=event,
|
||||
command="model",
|
||||
title=_cost_warning.title,
|
||||
message=(
|
||||
f"⚠️ **{_cost_warning.title}**\n\n{_cost_warning.message}\n\n"
|
||||
f"_Text fallback: reply `{_p}approve` to switch or `{_p}cancel` to keep "
|
||||
"the current model._"
|
||||
),
|
||||
handler=_on_cost_confirm,
|
||||
)
|
||||
|
||||
return await commit_switch(result)
|
||||
|
||||
async def _handle_codex_runtime_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /codex-runtime command in the gateway.
|
||||
|
||||
On change the cached agent is evicted so the next message builds a fresh AIAgent with the
|
||||
new api_mode (avoids prompt-cache invalidation mid-session).
|
||||
"""
|
||||
from hermes_cli import codex_runtime_switch as crs
|
||||
|
||||
raw_args = event.get_command_args().strip() if event else ""
|
||||
new_value, errors = crs.parse_args(raw_args)
|
||||
if errors:
|
||||
return "❌ " + "\n❌ ".join(errors)
|
||||
|
||||
# Load + persist via the same helpers used for /model and /yolo
|
||||
try:
|
||||
from hermes_cli.config import load_config, save_config
|
||||
except Exception as exc:
|
||||
return f"❌ Could not load config: {exc}"
|
||||
cfg = load_config()
|
||||
|
||||
result = crs.apply(
|
||||
cfg,
|
||||
new_value,
|
||||
persist_callback=(save_config if new_value is not None else None),
|
||||
)
|
||||
|
||||
# On a real change, evict the cached agent so the new runtime takes
|
||||
# effect on the next message rather than waiting for cache TTL.
|
||||
if result.success and new_value is not None and result.requires_new_session:
|
||||
try:
|
||||
session_key = self._session_key_for_source(event.source)
|
||||
self._evict_cached_agent(session_key)
|
||||
except Exception:
|
||||
logger.debug("could not evict cached agent after codex-runtime change",
|
||||
exc_info=True)
|
||||
|
||||
prefix = "✓" if result.success else "✗"
|
||||
return f"{prefix} {result.message}"
|
||||
|
||||
async def _handle_personality_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /personality command - list or set a personality.
|
||||
|
||||
All resolution/persistence goes through hermes_cli.personality, the single owner of state.
|
||||
"""
|
||||
from gateway.run import _load_gateway_config
|
||||
from hermes_cli.personality import (
|
||||
active_personality_name,
|
||||
available_personalities,
|
||||
describe_personality,
|
||||
persist_personality,
|
||||
resolve_personality,
|
||||
)
|
||||
|
||||
args = event.get_command_args().strip()
|
||||
|
||||
try:
|
||||
config = _load_gateway_config()
|
||||
except Exception:
|
||||
config = {}
|
||||
personalities = available_personalities(config)
|
||||
|
||||
if not args:
|
||||
current = active_personality_name(config)
|
||||
lines = [t("gateway.personality.header")]
|
||||
lines.append(t("gateway.personality.none_option"))
|
||||
for name, prompt in personalities.items():
|
||||
marker = " ✓" if name == current else ""
|
||||
lines.append(
|
||||
t(
|
||||
"gateway.personality.item",
|
||||
name=f"{name}{marker}",
|
||||
preview=describe_personality(prompt),
|
||||
)
|
||||
)
|
||||
lines.append(t("gateway.personality.usage"))
|
||||
return "\n".join(lines)
|
||||
|
||||
try:
|
||||
name, _new_prompt = resolve_personality(args, config)
|
||||
except ValueError:
|
||||
available = "`none`, " + ", ".join(f"`{n}`" for n in personalities)
|
||||
return t("gateway.personality.unknown", name=args.lower(), available=available)
|
||||
|
||||
# Persist the selection only — hermes_cli.personality never writes agent.system_prompt (user-
|
||||
# owned overlay). persist_personality writes get_hermes_home()/config.yaml (the routed profile
|
||||
# under multiplex) and the next turn re-resolves the prompt from it: no process-global state.
|
||||
if not persist_personality(name):
|
||||
return t("gateway.personality.save_failed", error="config write failed")
|
||||
|
||||
if not name:
|
||||
return t("gateway.personality.cleared")
|
||||
return t("gateway.personality.set_to", name=name)
|
||||
|
||||
def _save_gateway_config_key(self, key_path: str, value) -> bool:
|
||||
"""Save a dot-separated key to config.yaml (shared by /reasoning, /fast
|
||||
and their interactive pickers)."""
|
||||
from gateway.slash_commands import _nested_dict
|
||||
from gateway.run import _gateway_config_home
|
||||
from hermes_cli.config import read_user_config_raw
|
||||
config_path = _gateway_config_home() / "config.yaml"
|
||||
try:
|
||||
# Write-back round-trip: raw read is correct (merged defaults must
|
||||
# not be persisted back to the user's file).
|
||||
user_config = read_user_config_raw(config_path)
|
||||
*parents, leaf = key_path.split(".")
|
||||
_nested_dict(user_config, *parents)[leaf] = value
|
||||
atomic_config_write(config_path, user_config)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error("Failed to save config key %s: %s", key_path, e)
|
||||
return False
|
||||
|
||||
def _apply_reasoning_selection(
|
||||
self,
|
||||
session_key: str,
|
||||
platform_key: str,
|
||||
value: str,
|
||||
persist_global: bool = False,
|
||||
) -> str:
|
||||
"""Apply a /reasoning argument (typed or picked) and return the reply.
|
||||
|
||||
Single path shared by `/reasoning <arg>` and the choice picker so both match the parser.
|
||||
"""
|
||||
from hermes_constants import parse_reasoning_effort
|
||||
|
||||
value = (value or "").strip().lower()
|
||||
|
||||
# Display toggle (per-platform)
|
||||
if value in {"show", "on"}:
|
||||
self._show_reasoning = True
|
||||
self._save_gateway_config_key(
|
||||
f"display.platforms.{platform_key}.show_reasoning", True
|
||||
)
|
||||
return t("gateway.reasoning.display_set_on", platform=platform_key)
|
||||
if value in {"hide", "off"}:
|
||||
self._show_reasoning = False
|
||||
self._save_gateway_config_key(
|
||||
f"display.platforms.{platform_key}.show_reasoning", False
|
||||
)
|
||||
return t("gateway.reasoning.display_set_off", platform=platform_key)
|
||||
|
||||
if value == "reset":
|
||||
if persist_global:
|
||||
return t("gateway.reasoning.reset_global_unsupported")
|
||||
self._set_session_reasoning_override(session_key, None)
|
||||
self._reasoning_config = self._load_reasoning_config()
|
||||
self._evict_cached_agent(session_key)
|
||||
return t("gateway.reasoning.reset_done")
|
||||
|
||||
parsed = parse_reasoning_effort(value)
|
||||
if parsed is None:
|
||||
return t("gateway.reasoning.unknown_arg", arg=value)
|
||||
|
||||
self._reasoning_config = parsed
|
||||
if persist_global:
|
||||
if self._save_gateway_config_key("agent.reasoning_effort", value):
|
||||
self._set_session_reasoning_override(session_key, None)
|
||||
self._evict_cached_agent(session_key)
|
||||
return t("gateway.reasoning.set_global", effort=value)
|
||||
self._set_session_reasoning_override(session_key, parsed)
|
||||
self._evict_cached_agent(session_key)
|
||||
return t("gateway.reasoning.set_global_save_failed", effort=value)
|
||||
|
||||
self._set_session_reasoning_override(session_key, parsed)
|
||||
self._evict_cached_agent(session_key)
|
||||
return t("gateway.reasoning.set_session", effort=value)
|
||||
|
||||
def _reasoning_picker_choices(self, current_effort: str) -> list:
|
||||
"""Build the choice list for the interactive /reasoning picker."""
|
||||
from hermes_constants import VALID_REASONING_EFFORTS
|
||||
|
||||
choices = [{"value": "none", "label": t("gateway.reasoning.choice_none"), "is_current": current_effort == "none"}]
|
||||
choices.extend({"value": level, "label": level, "is_current": level == current_effort} for level in VALID_REASONING_EFFORTS)
|
||||
choices.extend(
|
||||
{"value": v, "label": t(f"gateway.reasoning.choice_{v}"), "is_current": False}
|
||||
for v in ("reset", "show", "hide")
|
||||
)
|
||||
return choices
|
||||
|
||||
async def _try_send_choice_picker(
|
||||
self,
|
||||
event: MessageEvent,
|
||||
session_key: str,
|
||||
title: str,
|
||||
choices: list,
|
||||
on_choice_selected,
|
||||
) -> bool:
|
||||
"""Send an interactive choice picker when the platform supports it.
|
||||
|
||||
Mirrors the `/model` gate: capability is detected on the adapter *type*
|
||||
(``send_choice_picker``); a failed send returns False (text fallback) instead of erroring.
|
||||
"""
|
||||
adapter = self._adapter_for_source(event.source)
|
||||
has_picker = (
|
||||
adapter is not None
|
||||
and getattr(type(adapter), "send_choice_picker", None) is not None
|
||||
)
|
||||
if not has_picker:
|
||||
return False
|
||||
try:
|
||||
metadata = self._reply_metadata(event)
|
||||
result = await adapter.send_choice_picker(
|
||||
chat_id=event.source.chat_id,
|
||||
title=title,
|
||||
choices=choices,
|
||||
session_key=session_key,
|
||||
on_choice_selected=on_choice_selected,
|
||||
metadata=metadata,
|
||||
)
|
||||
return bool(getattr(result, "success", False))
|
||||
except Exception as e:
|
||||
logger.warning("send_choice_picker failed, falling back to text: %s", e)
|
||||
return False
|
||||
|
||||
async def _handle_reasoning_command(self, event: MessageEvent) -> Optional[str]:
|
||||
"""Handle /reasoning command — manage reasoning effort and display toggle."""
|
||||
from gateway.run import _platform_config_key
|
||||
|
||||
raw_args = event.get_command_args().strip()
|
||||
args, persist_global = self._parse_reasoning_command_args(raw_args)
|
||||
# Normalize the source (Telegram DM topic recovery) before deriving
|
||||
# the override key so storage matches the key the next message turn
|
||||
# reads — same fix as /model (#30479).
|
||||
_reasoning_source = await asyncio.to_thread(self._normalize_source_for_session_key, event.source)
|
||||
session_key = self._session_key_for_source(_reasoning_source)
|
||||
self._show_reasoning = self._load_show_reasoning()
|
||||
# Use the session's effective model (session /model override wins over
|
||||
# config default) so per-model reasoning_overrides display correctly.
|
||||
_session_model = str(
|
||||
((getattr(self, "_session_model_overrides", {}) or {}).get(session_key) or {}).get("model") or ""
|
||||
)
|
||||
self._reasoning_config = self._resolve_session_reasoning_config(
|
||||
source=event.source,
|
||||
session_key=session_key,
|
||||
model=_session_model,
|
||||
)
|
||||
|
||||
if not raw_args:
|
||||
# Show current state
|
||||
rc = self._reasoning_config
|
||||
if rc is None:
|
||||
level = t("gateway.reasoning.level_default")
|
||||
current_effort = "medium"
|
||||
elif rc.get("enabled") is False:
|
||||
level = t("gateway.reasoning.level_disabled")
|
||||
current_effort = "none"
|
||||
else:
|
||||
level = rc.get("effort", "medium")
|
||||
current_effort = level
|
||||
display_state = (
|
||||
t("gateway.reasoning.display_on")
|
||||
if self._show_reasoning
|
||||
else t("gateway.reasoning.display_off")
|
||||
)
|
||||
has_session_override = session_key in (getattr(self, "_session_reasoning_overrides", {}) or {})
|
||||
scope = (
|
||||
t("gateway.reasoning.scope_session")
|
||||
if has_session_override
|
||||
else t("gateway.reasoning.scope_global")
|
||||
)
|
||||
|
||||
# Interactive picker on platforms that support it (parity with the
|
||||
# /model picker). Falls through to the text status card otherwise.
|
||||
_picker_platform_key = _platform_config_key(event.source.platform)
|
||||
|
||||
async def _on_reasoning_choice(_chat_id: str, value: str) -> str:
|
||||
return self._apply_reasoning_selection(
|
||||
session_key, _picker_platform_key, value
|
||||
)
|
||||
|
||||
picker_sent = await self._try_send_choice_picker(
|
||||
event,
|
||||
session_key,
|
||||
title=t(
|
||||
"gateway.reasoning.picker_title",
|
||||
level=level,
|
||||
scope=scope,
|
||||
display=display_state,
|
||||
),
|
||||
choices=self._reasoning_picker_choices(current_effort),
|
||||
on_choice_selected=_on_reasoning_choice,
|
||||
)
|
||||
if picker_sent:
|
||||
return None # Picker sent — adapter handles the response
|
||||
|
||||
return t(
|
||||
"gateway.reasoning.status",
|
||||
level=level,
|
||||
scope=scope,
|
||||
display=display_state,
|
||||
)
|
||||
|
||||
# Typed argument path — same applier the picker uses.
|
||||
platform_key = _platform_config_key(event.source.platform)
|
||||
return self._apply_reasoning_selection(
|
||||
session_key, platform_key, args, persist_global=persist_global
|
||||
)
|
||||
|
||||
async def _handle_fast_command(self, event: MessageEvent) -> Optional[str]:
|
||||
"""Handle /fast — mirror the CLI Priority Processing toggle in gateway chats.
|
||||
|
||||
Session-scoped by default; ``--global`` persists agent.service_tier (parity with /model).
|
||||
"""
|
||||
from gateway.run import _load_gateway_config, _resolve_gateway_model
|
||||
from hermes_cli.models import model_supports_fast_mode
|
||||
|
||||
raw_args = event.get_command_args().strip().lower()
|
||||
# Reuse the /reasoning arg parser: strips --global (any position),
|
||||
# normalizes unicode dashes.
|
||||
args, persist_global = self._parse_reasoning_command_args(raw_args)
|
||||
session_key = self._session_key_for_source(event.source)
|
||||
self._service_tier = self._resolve_session_service_tier(
|
||||
session_key=session_key
|
||||
)
|
||||
|
||||
user_config = _load_gateway_config()
|
||||
model = _resolve_gateway_model(user_config)
|
||||
if not model_supports_fast_mode(model):
|
||||
return t("gateway.fast.not_supported")
|
||||
|
||||
def _apply_fast_selection(value: str, persist: bool = False) -> str:
|
||||
"""Apply a /fast argument (typed or picked) and return the reply."""
|
||||
if value in {"fast", "on"}:
|
||||
tier = "priority"
|
||||
saved_value = "fast"
|
||||
label = t("gateway.fast.label_fast")
|
||||
elif value in {"normal", "off"}:
|
||||
tier = None
|
||||
saved_value = "normal"
|
||||
label = t("gateway.fast.label_normal")
|
||||
elif value in {"auto", "cold"}:
|
||||
tier = saved_value = value
|
||||
label = value.upper()
|
||||
else:
|
||||
return t("gateway.fast.unknown_arg", arg=value)
|
||||
self._service_tier = tier
|
||||
if persist:
|
||||
if self._save_gateway_config_key("agent.service_tier", saved_value):
|
||||
# Global write supersedes any session override.
|
||||
self._set_session_service_tier_override(
|
||||
session_key, None, clear=True
|
||||
)
|
||||
self._evict_cached_agent(session_key)
|
||||
return t("gateway.fast.saved", label=label)
|
||||
# Config write failed — fall back to a session override so the
|
||||
# user's choice still applies (mirrors /reasoning --global).
|
||||
self._set_session_service_tier_override(session_key, tier)
|
||||
self._evict_cached_agent(session_key)
|
||||
return t("gateway.fast.session_only", label=label)
|
||||
self._set_session_service_tier_override(session_key, tier)
|
||||
self._evict_cached_agent(session_key)
|
||||
return t("gateway.fast.session_only", label=label)
|
||||
|
||||
if not args or args == "status":
|
||||
is_fast = self._service_tier == "priority"
|
||||
mode = "fast" if is_fast else (self._service_tier or "normal")
|
||||
status = {"fast": t("gateway.fast.status_fast"), "normal": t("gateway.fast.status_normal")}.get(mode, mode)
|
||||
|
||||
async def _on_fast_choice(_chat_id: str, value: str) -> str:
|
||||
return _apply_fast_selection(value, persist=persist_global)
|
||||
|
||||
picker_sent = await self._try_send_choice_picker(
|
||||
event,
|
||||
session_key,
|
||||
title=t("gateway.fast.picker_title", mode=status),
|
||||
choices=[
|
||||
{"value": v, "label": t(f"gateway.fast.choice_{v}"), "is_current": mode == v}
|
||||
for v in ("fast", "normal", "auto", "cold")
|
||||
],
|
||||
on_choice_selected=_on_fast_choice,
|
||||
)
|
||||
if picker_sent:
|
||||
return None # Picker sent — adapter handles the response
|
||||
|
||||
return t("gateway.fast.status", mode=status)
|
||||
|
||||
return _apply_fast_selection(args, persist=persist_global)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,849 @@
|
||||
"""Read-only gateway introspection commands: /status, /context, /usage, /agents, /insights, /topup.
|
||||
|
||||
Split out of ``gateway/slash_commands.py``; bound onto ``GatewayRunner`` through
|
||||
``GatewaySlashCommandsMixin``. Origin internals are imported lazily (``from gateway.slash_commands
|
||||
import ...``) inside the bodies to avoid the import cycle.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import asyncio
|
||||
import hashlib
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from agent.account_usage import fetch_account_usage, render_account_usage_lines
|
||||
from agent.i18n import t
|
||||
from gateway.config import Platform
|
||||
from gateway.platforms.base import MessageEvent
|
||||
|
||||
# Log-record parity with gateway/run.py and the origin module.
|
||||
logger = logging.getLogger("gateway.run")
|
||||
|
||||
|
||||
def _clean_str(value: Any) -> str:
|
||||
"""Strip and return a non-empty string value, or empty string."""
|
||||
return value.strip() if isinstance(value, str) and value.strip() else ""
|
||||
|
||||
|
||||
def _int_value(value: Any) -> int:
|
||||
"""Safely coerce to int."""
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return 0
|
||||
|
||||
|
||||
def _status_model_route(status_agent, persisted_route: dict, session_row: dict, session_entry):
|
||||
"""``(model, provider, context_used, context_total)`` for /status.
|
||||
|
||||
Order: live/cached agent route -> persisted dominant route -> SessionDB row -> gateway config
|
||||
(only loaded when something is still missing).
|
||||
"""
|
||||
from gateway.run import _AGENT_PENDING_SENTINEL, _load_gateway_config, _resolve_gateway_model
|
||||
|
||||
model_name = provider_name = ""
|
||||
route_resolved = False
|
||||
context_used = context_total = 0
|
||||
if status_agent is not None and status_agent is not _AGENT_PENDING_SENTINEL:
|
||||
live_model = _clean_str(getattr(status_agent, "model", ""))
|
||||
live_provider = _clean_str(getattr(status_agent, "provider", ""))
|
||||
if live_model and live_provider:
|
||||
model_name, provider_name, route_resolved = live_model, live_provider, True
|
||||
ctx = getattr(status_agent, "context_compressor", None)
|
||||
if ctx is not None:
|
||||
context_used = _int_value(getattr(ctx, "last_prompt_tokens", 0))
|
||||
context_total = _int_value(getattr(ctx, "context_length", 0))
|
||||
|
||||
persisted_model = _clean_str(persisted_route.get("model"))
|
||||
persisted_provider = _clean_str(persisted_route.get("billing_provider"))
|
||||
if not route_resolved and persisted_model and persisted_provider:
|
||||
model_name, provider_name, route_resolved = persisted_model, persisted_provider, True
|
||||
if not route_resolved:
|
||||
model_name = _clean_str(session_row.get("model"))
|
||||
provider_name = _clean_str(session_row.get("billing_provider"))
|
||||
context_used = context_used or _int_value(getattr(session_entry, "last_prompt_tokens", 0))
|
||||
|
||||
user_config: dict[str, Any] = {}
|
||||
if not model_name or not provider_name or not context_total:
|
||||
try:
|
||||
user_config = _load_gateway_config()
|
||||
except Exception:
|
||||
user_config = {}
|
||||
model_cfg = user_config.get("model", {}) if isinstance(user_config, dict) else {}
|
||||
if not isinstance(model_cfg, dict):
|
||||
model_cfg = {}
|
||||
if not model_name:
|
||||
model_name = _resolve_gateway_model(user_config)
|
||||
if not provider_name:
|
||||
provider_name = _clean_str(model_cfg.get("provider"))
|
||||
if not context_total:
|
||||
configured_context = model_cfg.get("context_length")
|
||||
if isinstance(configured_context, int) and configured_context > 0:
|
||||
context_total = configured_context
|
||||
return model_name, provider_name, context_used, context_total
|
||||
|
||||
|
||||
def _context_compressor_lines(agent, ctx, used: int) -> list[str]:
|
||||
"""/context full view: auto-compression threshold/headroom, compression count + last savings,
|
||||
and cumulative throughput (labelled as throughput, NOT context size)."""
|
||||
lines: list[str] = []
|
||||
threshold = getattr(ctx, "threshold_tokens", 0) or 0
|
||||
threshold_pct = (getattr(ctx, "threshold_percent", 0) or 0) * 100
|
||||
if threshold > 0:
|
||||
if used >= threshold:
|
||||
lines.append(
|
||||
t("gateway.context.over_threshold", threshold=f"{threshold:,}", threshold_pct=f"{threshold_pct:.0f}")
|
||||
)
|
||||
else:
|
||||
lines.append(
|
||||
t(
|
||||
"gateway.context.threshold",
|
||||
threshold=f"{threshold:,}",
|
||||
threshold_pct=f"{threshold_pct:.0f}",
|
||||
to_go=f"{threshold - used:,}",
|
||||
)
|
||||
)
|
||||
compressions = getattr(ctx, "compression_count", 0) or 0
|
||||
lines.append(t("gateway.context.compressions", count=compressions))
|
||||
if compressions:
|
||||
savings = getattr(ctx, "_last_compression_savings_pct", None)
|
||||
if savings is not None:
|
||||
lines.append(t("gateway.context.last_savings", savings=f"{savings:.0f}"))
|
||||
|
||||
def _n(attr):
|
||||
return getattr(agent, attr, 0) or 0
|
||||
|
||||
lines.append("")
|
||||
lines.append(t("gateway.context.totals_header", calls=_n("session_api_calls")))
|
||||
lines.append(
|
||||
t(
|
||||
"gateway.context.totals_line",
|
||||
input=f"{_n('session_input_tokens'):,}",
|
||||
output=f"{_n('session_output_tokens'):,}",
|
||||
reasoning=f"{_n('session_reasoning_tokens'):,}",
|
||||
)
|
||||
)
|
||||
lines.append(t("gateway.context.total_billed", total=f"{_n('session_total_tokens'):,}"))
|
||||
lines.append(t("gateway.context.throughput_note"))
|
||||
return lines
|
||||
|
||||
|
||||
def _agents_delegation_lines(d: dict) -> list[str]:
|
||||
"""/agents rows for one background delegation. Live per-child activity comes from the
|
||||
registry's progress sampler: api calls, current tool, seconds since last activity."""
|
||||
goal = " ".join(str(d.get("goal") or "").split())
|
||||
if len(goal) > 70:
|
||||
goal = goal[:67] + "..."
|
||||
status = d.get("status", "?")
|
||||
row = f"- `{d.get('delegation_id', '?')}` · {status}"
|
||||
if status == "stalling":
|
||||
quiet = d.get("stalled_after_quiet_seconds")
|
||||
if quiet is not None:
|
||||
row += f" · no progress {quiet:.0f}s"
|
||||
elif d.get("seconds_since_progress", 0) >= 60:
|
||||
row += f" · quiet {d['seconds_since_progress']:.0f}s"
|
||||
if goal:
|
||||
row += f" · {goal}"
|
||||
lines = [row]
|
||||
for i, child in enumerate(d.get("children_activity") or []):
|
||||
if not isinstance(child, dict):
|
||||
continue
|
||||
tool = child.get("current_tool")
|
||||
doing = f"`{tool}`" if tool else "between turns"
|
||||
part = f" - child {i + 1}: {child.get('api_calls', '?')} api calls · {doing}"
|
||||
idle = child.get("seconds_since_activity")
|
||||
if idle is not None:
|
||||
part += f" · active {idle:.0f}s ago"
|
||||
lines.append(part)
|
||||
return lines
|
||||
|
||||
|
||||
def _usage_agent_stats_lines(agent) -> list[str]:
|
||||
"""/usage session block for a live agent: rate limits, token breakdown (matches the CLI),
|
||||
context window and compression count."""
|
||||
lines: list[str] = []
|
||||
rl_state = agent.get_rate_limit_state()
|
||||
if rl_state and rl_state.has_data:
|
||||
from agent.rate_limit_tracker import format_rate_limit_compact
|
||||
lines.append(t("gateway.usage.rate_limits", state=format_rate_limit_compact(rl_state)))
|
||||
lines.append("")
|
||||
input_tokens = getattr(agent, "session_input_tokens", 0) or 0
|
||||
output_tokens = getattr(agent, "session_output_tokens", 0) or 0
|
||||
lines.append(t("gateway.usage.header_session"))
|
||||
lines.append(t("gateway.usage.label_model", model=agent.model))
|
||||
lines.append(t("gateway.usage.label_input_tokens", count=f"{input_tokens:,}"))
|
||||
lines.append(t("gateway.usage.label_output_tokens", count=f"{output_tokens:,}"))
|
||||
lines.append(t("gateway.usage.label_total", count=f"{agent.session_total_tokens:,}"))
|
||||
lines.append(t("gateway.usage.label_api_calls", count=agent.session_api_calls))
|
||||
ctx = agent.context_compressor
|
||||
_lpt = ctx.last_prompt_tokens if ctx.last_prompt_tokens > 0 else 0
|
||||
if _lpt:
|
||||
pct = min(100, _lpt / ctx.context_length * 100) if ctx.context_length else 0
|
||||
lines.append(t("gateway.usage.label_context", used=f"{_lpt:,}", total=f"{ctx.context_length:,}", pct=f"{pct:.0f}"))
|
||||
if ctx.compression_count:
|
||||
lines.append(t("gateway.usage.label_compressions", count=ctx.compression_count))
|
||||
return lines
|
||||
|
||||
|
||||
class GatewayStatusCommandsMixin:
|
||||
"""Read-only gateway introspection commands: /status, /context, /usage, /agents, /insights, /topup."""
|
||||
|
||||
async def _handle_status_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /status command."""
|
||||
from gateway.run import _AGENT_PENDING_SENTINEL
|
||||
|
||||
source = event.source
|
||||
session_entry = await self.async_session_store.get_or_create_session(source)
|
||||
|
||||
connected_platforms = [p.value for p in self.adapters]
|
||||
|
||||
# Check if there's an active agent. Keep the sentinel distinct: a
|
||||
# starting/pending run should not be treated as a fully usable agent for
|
||||
# model/context display, but it still occupies the session slot.
|
||||
session_key = session_entry.session_key
|
||||
agent = self._running_agents.get(session_key)
|
||||
is_running = agent is not None and agent is not _AGENT_PENDING_SENTINEL
|
||||
|
||||
# Count pending /queue follow-ups (slot + overflow).
|
||||
adapter = self.adapters.get(source.platform) if source else None
|
||||
queue_depth = self._queue_depth(session_key, adapter=adapter)
|
||||
|
||||
title, session_row, db_total_tokens, persisted_route = await self._status_session_db_facts(
|
||||
session_entry.session_id
|
||||
)
|
||||
# Resolve model/context for cockpit-style status. Prefer the live or cached agent because it
|
||||
# carries the actual runtime route and context compressor; fall back to SessionDB metadata +
|
||||
# last_prompt_tokens so /status stays useful between turns without billing/account calls.
|
||||
status_agent = agent if is_running else self._cached_agent_for(session_key)
|
||||
model_name, provider_name, context_used, context_total = _status_model_route(
|
||||
status_agent, persisted_route, session_row, session_entry
|
||||
)
|
||||
|
||||
model_line = ""
|
||||
if model_name:
|
||||
if provider_name:
|
||||
model_line = t("gateway.status.model_provider", model=model_name, provider=provider_name)
|
||||
else:
|
||||
model_line = t("gateway.status.model", model=model_name)
|
||||
|
||||
context_line = ""
|
||||
if context_total:
|
||||
pct = min(100, round((context_used / context_total) * 100)) if context_total else 0
|
||||
context_line = t(
|
||||
"gateway.status.context",
|
||||
used=f"{context_used:,}",
|
||||
total=f"{context_total:,}",
|
||||
pct=f"{pct}",
|
||||
)
|
||||
elif context_used:
|
||||
context_line = t("gateway.status.context_used", used=f"{context_used:,}")
|
||||
|
||||
lines = [
|
||||
t("gateway.status.header"),
|
||||
"",
|
||||
t("gateway.status.session_id", session_id=session_entry.session_id),
|
||||
]
|
||||
if title:
|
||||
lines.append(t("gateway.status.title", title=title))
|
||||
lines.extend([
|
||||
t("gateway.status.created", timestamp=session_entry.created_at.strftime('%Y-%m-%d %H:%M')),
|
||||
t("gateway.status.last_activity", timestamp=session_entry.updated_at.strftime('%Y-%m-%d %H:%M')),
|
||||
])
|
||||
if model_line:
|
||||
lines.append(model_line)
|
||||
if context_line:
|
||||
lines.append(context_line)
|
||||
lines.extend([
|
||||
t("gateway.status.tokens", tokens=f"{db_total_tokens:,}"),
|
||||
t("gateway.status.agent_running", state=t("gateway.status.state_yes") if is_running else t("gateway.status.state_no")),
|
||||
])
|
||||
if queue_depth:
|
||||
lines.append(t("gateway.status.queued", count=queue_depth))
|
||||
if source.platform == Platform.MATRIX:
|
||||
scope = getattr(self.adapters.get(Platform.MATRIX), "_matrix_session_scope", os.getenv("MATRIX_SESSION_SCOPE", "auto"))
|
||||
thread = source.thread_id or "none"
|
||||
lines.extend([
|
||||
"",
|
||||
t("gateway.status.matrix_scope_header"),
|
||||
t("gateway.status.matrix_scope_room", room=source.chat_name or source.chat_id),
|
||||
t("gateway.status.matrix_scope_room_id", room_id=source.chat_id),
|
||||
t("gateway.status.matrix_scope_thread", thread_id=thread),
|
||||
t("gateway.status.matrix_scope_mode", scope=scope),
|
||||
t(
|
||||
"gateway.status.matrix_scope_key",
|
||||
session_key=self._redact_matrix_session_key(session_key),
|
||||
),
|
||||
])
|
||||
lines.extend([
|
||||
"",
|
||||
t("gateway.status.platforms", platforms=', '.join(connected_platforms)),
|
||||
])
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
async def _status_session_db_facts(self, session_id: str):
|
||||
"""``(title, session_row, db_total_tokens, persisted_route)`` for /status; each fail-open.
|
||||
|
||||
Token totals come from the SQLite session DB rather than the in-memory SessionStore: the
|
||||
agent's per-turn token deltas are persisted into sessions_db (run_agent.py), not into
|
||||
SessionEntry, so session_entry.total_tokens is always 0.
|
||||
"""
|
||||
title = None
|
||||
session_row: dict[str, Any] = {}
|
||||
db_total_tokens = 0
|
||||
persisted_route: dict[str, Any] = {}
|
||||
if not self._session_db:
|
||||
return title, session_row, db_total_tokens, persisted_route
|
||||
try:
|
||||
title = await self._session_db.get_session_title(session_id)
|
||||
except Exception:
|
||||
title = None
|
||||
try:
|
||||
row = await self._session_db.get_session(session_id)
|
||||
if isinstance(row, dict):
|
||||
session_row = row
|
||||
db_total_tokens = sum(
|
||||
_int_value(row.get(k))
|
||||
for k in ("input_tokens", "output_tokens", "cache_read_tokens", "cache_write_tokens", "reasoning_tokens")
|
||||
)
|
||||
except Exception:
|
||||
db_total_tokens = 0
|
||||
try:
|
||||
route = await self._session_db.get_dominant_session_model_route(session_id)
|
||||
if isinstance(route, dict):
|
||||
persisted_route = route
|
||||
except Exception:
|
||||
persisted_route = {}
|
||||
return title, session_row, db_total_tokens, persisted_route
|
||||
|
||||
@staticmethod
|
||||
def _redact_matrix_session_key(session_key: str) -> str:
|
||||
"""Return a stable Matrix session-key fingerprint for shared room status."""
|
||||
text = str(session_key or "")
|
||||
digest = hashlib.sha256(text.encode("utf-8")).hexdigest()[:12]
|
||||
return f"sha256:{digest}"
|
||||
|
||||
async def _handle_context_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /context — the dedicated context-window view.
|
||||
|
||||
/status shows a one-line ``used / total`` summary; this command is the deep view: a usage
|
||||
gauge, auto-compression threshold and headroom, compression count and last savings, and
|
||||
cumulative throughput — the last clearly labelled as throughput, NOT context size.
|
||||
Resolution order: running agent, cached agent, SessionStore/SessionDB metadata, and a
|
||||
transcript estimate only as last resort. ``/context all`` adds per-skill/toolset listings.
|
||||
"""
|
||||
source = event.source
|
||||
session_key = self._session_key_for_source(source)
|
||||
session_entry = await self.async_session_store.get_or_create_session(source)
|
||||
expanded = event.get_command_args().strip().lower() in {"all", "full", "details"}
|
||||
|
||||
# Running agent first (mid-turn), then cached agent (between turns).
|
||||
agent = self._resident_agent_for(session_key)
|
||||
has_agent = bool(agent)
|
||||
|
||||
ctx = getattr(agent, "context_compressor", None) if has_agent else None
|
||||
used, context_length, model_name = await self._resolve_context_figures(
|
||||
agent if has_agent else None, ctx, session_entry, source
|
||||
)
|
||||
|
||||
# Gauge path: real current-context figure
|
||||
if used > 0 and context_length > 0:
|
||||
pct = min(100.0, used / context_length * 100)
|
||||
headroom = max(0, context_length - used)
|
||||
BAR_WIDTH = 24
|
||||
filled = int(round(pct / 100 * BAR_WIDTH))
|
||||
bar = "█" * max(0, filled) + "░" * max(0, BAR_WIDTH - filled)
|
||||
|
||||
lines = [
|
||||
t("gateway.context.header"),
|
||||
"",
|
||||
t("gateway.context.model", model=model_name or "?"),
|
||||
t("gateway.context.window", total=f"{context_length:,}"),
|
||||
t(
|
||||
"gateway.context.in_use",
|
||||
used=f"{used:,}",
|
||||
total=f"{context_length:,}",
|
||||
pct=f"{pct:.0f}",
|
||||
),
|
||||
t("gateway.context.bar", bar=bar),
|
||||
t("gateway.context.headroom", headroom=f"{headroom:,}"),
|
||||
"",
|
||||
]
|
||||
# Full view — compression / throughput need the live agent.
|
||||
if ctx is not None:
|
||||
lines.extend(_context_compressor_lines(agent, ctx, used))
|
||||
else:
|
||||
lines.append(t("gateway.context.detail_after_first"))
|
||||
|
||||
# Per-category estimated breakdown (+ optional expanded listings). Same chars/4 engine
|
||||
# the desktop popover and /usage use; plain text (no glyph grid — monospace isn't
|
||||
# guaranteed on messaging platforms). Fail-open: rendering errors never break /context.
|
||||
if has_agent:
|
||||
breakdown = await asyncio.to_thread(
|
||||
self._context_breakdown_block, agent, source, expanded
|
||||
)
|
||||
if breakdown:
|
||||
lines.append("")
|
||||
lines.extend(breakdown)
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
# Last resort: rough estimate from transcript
|
||||
history = await self.async_session_store.load_transcript(session_entry.session_id)
|
||||
if history:
|
||||
from agent.model_metadata import estimate_messages_tokens_rough
|
||||
|
||||
msgs = [
|
||||
m
|
||||
for m in history
|
||||
if m.get("role") in {"user", "assistant"} and m.get("content")
|
||||
]
|
||||
approx = estimate_messages_tokens_rough(msgs)
|
||||
return "\n".join(
|
||||
[
|
||||
t("gateway.context.header"),
|
||||
"",
|
||||
t(
|
||||
"gateway.context.estimated",
|
||||
count=f"{approx:,}",
|
||||
messages=len(msgs),
|
||||
),
|
||||
t("gateway.context.detail_after_first"),
|
||||
]
|
||||
)
|
||||
return t("gateway.context.no_data")
|
||||
|
||||
async def _resolve_context_figures(self, agent, ctx, session_entry, source):
|
||||
"""``(used, context_length, model_name)`` for /context with cascading fallbacks.
|
||||
|
||||
used : compressor.last_prompt_tokens -> SessionStore.last_prompt_tokens
|
||||
model : agent.model -> SessionDB row model
|
||||
window: compressor.context_length -> effective gateway model route -> model metadata
|
||||
"""
|
||||
used = context_length = 0
|
||||
if ctx is not None:
|
||||
used = getattr(ctx, "last_prompt_tokens", 0) or 0
|
||||
context_length = getattr(ctx, "context_length", 0) or 0
|
||||
model_name = _clean_str(getattr(agent, "model", "")) if agent is not None else ""
|
||||
if not used:
|
||||
used = _int_value(getattr(session_entry, "last_prompt_tokens", 0))
|
||||
if not model_name and self._session_db:
|
||||
try:
|
||||
row = await self._session_db.get_session(session_entry.session_id) or {}
|
||||
if isinstance(row, dict):
|
||||
model_name = _clean_str(row.get("model", ""))
|
||||
except Exception:
|
||||
model_name = ""
|
||||
if not context_length:
|
||||
try:
|
||||
from gateway.run import _profile_runtime_scope, _resolve_gateway_model_context
|
||||
|
||||
def _resolve_nonresident_context():
|
||||
if getattr(getattr(self, "config", None), "multiplex_profiles", False):
|
||||
profile_home = self._resolve_profile_home_for_source(source)
|
||||
with _profile_runtime_scope(profile_home):
|
||||
return _resolve_gateway_model_context(model_name or None)
|
||||
return _resolve_gateway_model_context(model_name or None)
|
||||
|
||||
resolved = await asyncio.to_thread(_resolve_nonresident_context)
|
||||
model_name = model_name or resolved.model
|
||||
context_length = _int_value(resolved.context_length)
|
||||
except Exception:
|
||||
context_length = 0
|
||||
if not context_length and model_name:
|
||||
try:
|
||||
from agent.model_metadata import get_model_context_length
|
||||
|
||||
context_length = _int_value(await asyncio.to_thread(get_model_context_length, model_name))
|
||||
except Exception:
|
||||
context_length = 0
|
||||
return used, context_length, model_name
|
||||
|
||||
async def _handle_agents_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /agents command - list active agents and running tasks."""
|
||||
from gateway.run import _AGENT_PENDING_SENTINEL
|
||||
from tools.process_registry import format_uptime_short, process_registry
|
||||
|
||||
now = time.time()
|
||||
current_session_key = self._session_key_for_source(event.source)
|
||||
|
||||
running_agents: dict = getattr(self, "_running_agents", {}) or {}
|
||||
running_started: dict = getattr(self, "_running_agents_ts", {}) or {}
|
||||
|
||||
agent_rows: list[dict] = []
|
||||
for session_key, agent in running_agents.items():
|
||||
started = float(running_started.get(session_key, now))
|
||||
elapsed = max(0, int(now - started))
|
||||
is_pending = agent is _AGENT_PENDING_SENTINEL
|
||||
agent_rows.append(
|
||||
{
|
||||
"session_key": session_key,
|
||||
"elapsed": elapsed,
|
||||
"state": t("gateway.agents.state_starting") if is_pending else t("gateway.agents.state_running"),
|
||||
"session_id": "" if is_pending else str(getattr(agent, "session_id", "") or ""),
|
||||
"model": "" if is_pending else str(getattr(agent, "model", "") or ""),
|
||||
}
|
||||
)
|
||||
|
||||
agent_rows.sort(key=lambda row: row["elapsed"], reverse=True)
|
||||
|
||||
running_processes: list[dict] = []
|
||||
try:
|
||||
running_processes = [
|
||||
p for p in process_registry.list_sessions()
|
||||
if p.get("status") == "running"
|
||||
]
|
||||
except Exception:
|
||||
running_processes = []
|
||||
|
||||
background_tasks = [
|
||||
t for t in (getattr(self, "_background_tasks", set()) or set())
|
||||
if hasattr(t, "done") and not t.done()
|
||||
]
|
||||
|
||||
lines = [
|
||||
t("gateway.agents.header"),
|
||||
"",
|
||||
t("gateway.agents.active_agents", count=len(agent_rows)),
|
||||
]
|
||||
|
||||
if agent_rows:
|
||||
for idx, row in enumerate(agent_rows[:12], 1):
|
||||
current = t("gateway.agents.this_chat") if row["session_key"] == current_session_key else ""
|
||||
sid = f" · `{row['session_id']}`" if row["session_id"] else ""
|
||||
model = f" · `{row['model']}`" if row["model"] else ""
|
||||
lines.append(
|
||||
f"{idx}. `{row['session_key']}` · {row['state']} · "
|
||||
f"{format_uptime_short(row['elapsed'])}{sid}{model}{current}"
|
||||
)
|
||||
if len(agent_rows) > 12:
|
||||
lines.append(t("gateway.agents.more", count=len(agent_rows) - 12))
|
||||
|
||||
lines.extend(
|
||||
[
|
||||
"",
|
||||
t("gateway.agents.running_processes", count=len(running_processes)),
|
||||
]
|
||||
)
|
||||
if running_processes:
|
||||
for proc in running_processes[:12]:
|
||||
cmd = " ".join(str(proc.get("command", "")).split())
|
||||
if len(cmd) > 90:
|
||||
cmd = cmd[:87] + "..."
|
||||
lines.append(
|
||||
f"- `{proc.get('session_id', '?')}` · "
|
||||
f"{format_uptime_short(int(proc.get('uptime_seconds', 0)))} · `{cmd}`"
|
||||
)
|
||||
if len(running_processes) > 12:
|
||||
lines.append(t("gateway.agents.more", count=len(running_processes) - 12))
|
||||
|
||||
lines.extend(
|
||||
[
|
||||
"",
|
||||
t("gateway.agents.async_jobs", count=len(background_tasks)),
|
||||
]
|
||||
)
|
||||
|
||||
# Background (async) delegations — delegate_task(background=true).
|
||||
try:
|
||||
from tools.async_delegation import list_async_delegations
|
||||
delegations = [
|
||||
d for d in list_async_delegations()
|
||||
if d.get("status") in ("running", "stalling", "finalizing")
|
||||
]
|
||||
except Exception:
|
||||
delegations = []
|
||||
if delegations:
|
||||
lines.extend(["", t("gateway.agents.background_delegations", count=len(delegations))])
|
||||
for d in delegations[:12]:
|
||||
lines.extend(_agents_delegation_lines(d))
|
||||
if len(delegations) > 12:
|
||||
lines.append(t("gateway.agents.more", count=len(delegations) - 12))
|
||||
|
||||
if (
|
||||
not agent_rows
|
||||
and not running_processes
|
||||
and not background_tasks
|
||||
and not delegations
|
||||
):
|
||||
lines.append("")
|
||||
lines.append(t("gateway.agents.none"))
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
async def _handle_topup_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /topup -- show the Nous balance and hand off to the portal.
|
||||
|
||||
Does NOT charge, confirm, or track payment — that happens in the browser; the next /topup
|
||||
shows the new balance. Fetched off the event loop; fail-open.
|
||||
"""
|
||||
from agent.account_usage import build_credits_view
|
||||
|
||||
try:
|
||||
view = await asyncio.to_thread(build_credits_view, markdown=True)
|
||||
except Exception:
|
||||
view = None
|
||||
|
||||
if view is None or not view.logged_in:
|
||||
return t("gateway.credits.not_logged_in")
|
||||
|
||||
lines: list[str] = ["💳 **Nous balance**"]
|
||||
for line in view.balance_lines:
|
||||
if line.lstrip().startswith("📈"):
|
||||
continue # drop the helper's header; we print our own
|
||||
lines.append(line)
|
||||
if view.identity_line:
|
||||
lines.append("")
|
||||
lines.append(view.identity_line)
|
||||
if view.topup_url:
|
||||
lines.append("")
|
||||
lines.append(f"Manage billing on the portal: {view.topup_url}")
|
||||
lines.append("Top up and manage billing in the browser — your balance updates here after.")
|
||||
return "\n".join(lines)
|
||||
|
||||
def _context_breakdown_block(self, agent, source, expanded: bool) -> list[str]:
|
||||
"""Render the /context per-category block (plain text, no grid).
|
||||
|
||||
Estimated (chars/4), same engine as /usage. Runs in a thread; returns [] and never raises.
|
||||
"""
|
||||
try:
|
||||
from agent.context_breakdown import compute_context_details, render_context_breakdown_lines
|
||||
|
||||
payload = self._session_context_breakdown(agent, source)
|
||||
if not (payload.get("categories") or []):
|
||||
return []
|
||||
|
||||
details = None
|
||||
if expanded:
|
||||
try:
|
||||
details = compute_context_details(agent)
|
||||
except Exception:
|
||||
details = {"skills": [], "toolsets": []}
|
||||
|
||||
return render_context_breakdown_lines(payload, details=details, grid=False)
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
def _session_context_breakdown(self, agent, source) -> dict:
|
||||
"""Per-category context estimate (chars/4) for *agent* over the session transcript (sync)."""
|
||||
from agent.context_breakdown import compute_session_context_breakdown
|
||||
|
||||
history: list[dict] = []
|
||||
try:
|
||||
entry = self.session_store.get_or_create_session(source)
|
||||
history = self.session_store.load_transcript(entry.session_id) or []
|
||||
except Exception:
|
||||
history = []
|
||||
return compute_session_context_breakdown(agent, history)
|
||||
|
||||
def _context_breakdown_lines(self, agent, source) -> list[str]:
|
||||
"""Render the per-category context breakdown for /usage.
|
||||
|
||||
Estimated (chars/4). Returns [] and never raises so /usage stays robust.
|
||||
"""
|
||||
try:
|
||||
payload = self._session_context_breakdown(agent, source)
|
||||
categories = payload.get("categories") or []
|
||||
if not categories:
|
||||
return []
|
||||
|
||||
total = payload.get("estimated_total") or 0
|
||||
out = [t("gateway.usage.breakdown_header")]
|
||||
for cat in categories:
|
||||
tokens = int(cat.get("tokens") or 0)
|
||||
if tokens <= 0:
|
||||
continue
|
||||
cat_id = str(cat.get("id") or "")
|
||||
label = t(f"gateway.usage.breakdown_cat_{cat_id}")
|
||||
# Missing key → t() echoes the key back; fall back to the
|
||||
# English label the engine already provides.
|
||||
if label.endswith(f"breakdown_cat_{cat_id}"):
|
||||
label = str(cat.get("label") or cat_id)
|
||||
pct = round(tokens / total * 100) if total else 0
|
||||
out.append(
|
||||
t("gateway.usage.breakdown_line", label=label, count=f"{tokens:,}", pct=pct)
|
||||
)
|
||||
return out if len(out) > 1 else []
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
async def _handle_usage_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /usage command -- show token usage for the current session.
|
||||
|
||||
Checks both _running_agents (mid-turn) and _agent_cache (between turns) so details are
|
||||
available whenever the user asks.
|
||||
"""
|
||||
source = event.source
|
||||
session_key = self._session_key_for_source(source)
|
||||
|
||||
# `/usage reset [--force]` — redeem one banked Codex rate-limit reset
|
||||
# credit. Parsed before the display path so it never mixes with the
|
||||
# stats rendering below.
|
||||
raw_args = event.get_command_args().strip()
|
||||
args = [a.lower() for a in raw_args.split()] if raw_args else []
|
||||
wants_reset = bool(args) and args[0] == "reset"
|
||||
if args and not wants_reset:
|
||||
return t("gateway.usage.unknown_subcommand", args=raw_args)
|
||||
|
||||
# Running agent first (mid-turn), then cached agent (between turns).
|
||||
agent = self._resident_agent_for(session_key)
|
||||
|
||||
# Resolve provider/base_url/api_key for the account-usage fetch. Prefer the live agent; fall
|
||||
# back to persisted billing data on the SessionDB row so `/usage` still returns account info
|
||||
# between turns when no agent is resident.
|
||||
provider = getattr(agent, "provider", None) if agent else None
|
||||
base_url = getattr(agent, "base_url", None) if agent else None
|
||||
api_key = getattr(agent, "api_key", None) if agent else None
|
||||
if not provider and getattr(self, "_session_db", None) is not None:
|
||||
provider, base_url = await self._persisted_billing_route(source)
|
||||
|
||||
if wants_reset:
|
||||
normalized_provider = str(provider or "").strip().lower()
|
||||
if normalized_provider != "openai-codex":
|
||||
return t("gateway.usage.reset_wrong_provider")
|
||||
force = "--force" in args[1:]
|
||||
from agent.account_usage import redeem_codex_reset_credit
|
||||
|
||||
result = await asyncio.to_thread(
|
||||
redeem_codex_reset_credit,
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
force=force,
|
||||
)
|
||||
return result.message
|
||||
|
||||
# Fetch account usage off the event loop so slow provider APIs don't
|
||||
# block the gateway. Failures are non-fatal -- account_lines stays [].
|
||||
account_lines: list[str] = []
|
||||
credits_lines: list[str] = []
|
||||
if provider:
|
||||
try:
|
||||
account_snapshot = await asyncio.to_thread(
|
||||
fetch_account_usage,
|
||||
provider,
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
)
|
||||
except Exception:
|
||||
account_snapshot = None
|
||||
if account_snapshot:
|
||||
account_lines = render_account_usage_lines(account_snapshot, markdown=True)
|
||||
|
||||
# ── Nous credits magnitudes + monthly-grant % gauge ─────────────
|
||||
# Shared with CLI/TUI via nous_credits_lines(); run off the event loop. Gates on "a Nous
|
||||
# account is logged in" — NOT the inference provider, NOT under `if provider:` — so a Nous
|
||||
# user inferring elsewhere still sees a balance. No recovery trigger; fail-open.
|
||||
try:
|
||||
from agent.account_usage import nous_credits_lines
|
||||
|
||||
credits_lines = await asyncio.to_thread(nous_credits_lines, markdown=True)
|
||||
except Exception:
|
||||
credits_lines = [] # fail-open: never break /usage
|
||||
|
||||
def _with_account_blocks(lines: list[str]) -> str:
|
||||
# Each block is preceded by a blank divider only when something precedes it.
|
||||
for block in (account_lines, credits_lines):
|
||||
if block:
|
||||
if lines:
|
||||
lines.append("")
|
||||
lines.extend(block)
|
||||
return "\n".join(lines)
|
||||
|
||||
if agent and hasattr(agent, "session_total_tokens") and agent.session_api_calls > 0:
|
||||
lines = _usage_agent_stats_lines(agent)
|
||||
# Per-category context breakdown (estimated — chars/4 heuristic). Same engine the
|
||||
# desktop popover uses. The system prompt / tools / skills / memory slices read off the
|
||||
# live agent; the conversation slice is estimated from the session transcript.
|
||||
breakdown_lines = await asyncio.to_thread(self._context_breakdown_lines, agent, source)
|
||||
if breakdown_lines:
|
||||
lines.append("")
|
||||
lines.extend(breakdown_lines)
|
||||
return _with_account_blocks(lines)
|
||||
|
||||
# No agent at all -- check session history for a rough count
|
||||
session_entry = await self.async_session_store.get_or_create_session(source)
|
||||
history = await self.async_session_store.load_transcript(session_entry.session_id)
|
||||
if history:
|
||||
from agent.model_metadata import estimate_messages_tokens_rough
|
||||
msgs = [m for m in history if m.get("role") in {"user", "assistant"} and m.get("content")]
|
||||
approx = estimate_messages_tokens_rough(msgs)
|
||||
return _with_account_blocks([
|
||||
t("gateway.usage.header_session_info"),
|
||||
t("gateway.usage.label_messages", count=len(msgs)),
|
||||
t("gateway.usage.label_estimated_context", count=f"{approx:,}"),
|
||||
t("gateway.usage.detailed_after_first"),
|
||||
])
|
||||
if account_lines or credits_lines:
|
||||
return _with_account_blocks([])
|
||||
return t("gateway.usage.no_data")
|
||||
|
||||
async def _persisted_billing_route(self, source):
|
||||
"""``(provider, base_url)`` from the SessionDB row / dominant route when no agent is resident."""
|
||||
try:
|
||||
entry = await self.async_session_store.get_or_create_session(source)
|
||||
persisted = await self._session_db.get_session(entry.session_id) or {}
|
||||
route = await self._session_db.get_dominant_session_model_route(entry.session_id)
|
||||
persisted_route = route if isinstance(route, dict) else {}
|
||||
except Exception:
|
||||
persisted = {}
|
||||
persisted_route = {}
|
||||
if persisted_route.get("billing_provider"):
|
||||
return persisted_route["billing_provider"], persisted_route.get("billing_base_url")
|
||||
return persisted.get("billing_provider"), persisted.get("billing_base_url")
|
||||
|
||||
async def _handle_insights_command(self, event: MessageEvent) -> str:
|
||||
"""Handle /insights command -- show usage insights and analytics."""
|
||||
args = event.get_command_args().strip()
|
||||
|
||||
# Normalize Unicode dashes (Telegram/iOS auto-converts -- to em/en dash)
|
||||
args = re.sub(r'[\u2012\u2013\u2014\u2015](days|source)', r'--\1', args)
|
||||
|
||||
days = 30
|
||||
source = None
|
||||
|
||||
# Parse simple args: /insights 7 or /insights --days 7
|
||||
if args:
|
||||
parts = args.split()
|
||||
i = 0
|
||||
while i < len(parts):
|
||||
if parts[i] == "--days" and i + 1 < len(parts):
|
||||
try:
|
||||
days = int(parts[i + 1])
|
||||
except ValueError:
|
||||
return t("gateway.insights.invalid_days", value=parts[i + 1])
|
||||
i += 2
|
||||
elif parts[i] == "--source" and i + 1 < len(parts):
|
||||
source = parts[i + 1]
|
||||
i += 2
|
||||
elif parts[i].isdigit():
|
||||
days = int(parts[i])
|
||||
i += 1
|
||||
else:
|
||||
i += 1
|
||||
|
||||
try:
|
||||
from hermes_state import get_shared_session_db
|
||||
from agent.insights import InsightsEngine
|
||||
|
||||
def _run_insights():
|
||||
db = get_shared_session_db()
|
||||
try:
|
||||
engine = InsightsEngine(db)
|
||||
report = engine.generate(days=days, source=source)
|
||||
result = engine.format_gateway(report)
|
||||
return result
|
||||
finally:
|
||||
from hermes_state import release_or_close
|
||||
release_or_close(db)
|
||||
|
||||
# Not a bare hop: ``SessionDB()`` resolves ``get_hermes_home()`` at call time, which is
|
||||
# a contextvar set by ``_profile_runtime_scope``; a default-executor hop starts with an
|
||||
# EMPTY context and would read the DEFAULT profile's state.db.
|
||||
return await self._run_in_executor_with_context(_run_insights)
|
||||
except Exception as e:
|
||||
logger.error("Insights command error: %s", e, exc_info=True)
|
||||
return t("gateway.insights.error", error=e)
|
||||
@@ -23,10 +23,8 @@ from __future__ import annotations
|
||||
import ast
|
||||
import inspect
|
||||
|
||||
from gateway import run as gateway_run
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway import slash_commands as gateway_slash
|
||||
from gateway import slash_commands_model as gateway_slash
|
||||
|
||||
|
||||
def _assigns_false(node: ast.AST, attr: str) -> bool:
|
||||
@@ -77,13 +75,13 @@ def test_run_consumes_was_auto_reset_in_cleanup_block():
|
||||
|
||||
|
||||
def test_slash_command_model_path_consumes_was_auto_reset():
|
||||
"""The slash-command model path in gateway/slash_commands.py must consume
|
||||
"""The slash-command model path in gateway/slash_commands_model.py must consume
|
||||
`was_auto_reset` before storing the new model override, so a
|
||||
/model-first-after-auto-reset isn't wiped by the next message's cleanup
|
||||
(#48031)."""
|
||||
src = inspect.getsource(gateway_slash)
|
||||
tree = ast.parse(src)
|
||||
assert _assigns_false(tree, "was_auto_reset"), (
|
||||
"gateway/slash_commands.py model path must set "
|
||||
"gateway/slash_commands_model.py model path must set "
|
||||
"`was_auto_reset = False` before storing the model override (#48031)."
|
||||
)
|
||||
|
||||
@@ -312,6 +312,7 @@ def _slash_host(agent, session_key="tg:123"):
|
||||
return await asyncio.get_running_loop().run_in_executor(None, fn)
|
||||
|
||||
host._run_in_executor_with_context = _run_in_executor_with_context
|
||||
host._cached_agent_for = GatewaySlashCommandsMixin._cached_agent_for.__get__(host)
|
||||
host._compress_codex_app_server_session = (
|
||||
GatewaySlashCommandsMixin._compress_codex_app_server_session.__get__(host)
|
||||
)
|
||||
|
||||
@@ -140,11 +140,11 @@ class TestUsageAccountSection:
|
||||
|
||||
monkeypatch.setattr("gateway.run.asyncio.to_thread", _fake_to_thread)
|
||||
monkeypatch.setattr(
|
||||
"gateway.slash_commands.fetch_account_usage",
|
||||
"gateway.slash_commands_status.fetch_account_usage",
|
||||
lambda provider, base_url=None, api_key=None: object(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"gateway.slash_commands.render_account_usage_lines",
|
||||
"gateway.slash_commands_status.render_account_usage_lines",
|
||||
lambda snapshot, markdown=False: [
|
||||
"📈 **Account limits**",
|
||||
"Provider: openai-codex (Pro)",
|
||||
@@ -186,11 +186,11 @@ class TestUsageAccountSection:
|
||||
|
||||
monkeypatch.setattr("gateway.run.asyncio.to_thread", _fake_to_thread)
|
||||
monkeypatch.setattr(
|
||||
"gateway.slash_commands.fetch_account_usage",
|
||||
"gateway.slash_commands_status.fetch_account_usage",
|
||||
lambda provider, base_url=None, api_key=None: object(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"gateway.slash_commands.render_account_usage_lines",
|
||||
"gateway.slash_commands_status.render_account_usage_lines",
|
||||
lambda snapshot, markdown=False: ["account limits"],
|
||||
)
|
||||
monkeypatch.setattr("agent.account_usage.nous_credits_lines", lambda markdown=False: [])
|
||||
|
||||
@@ -57,13 +57,13 @@ class TestSourceLinesAreClamped:
|
||||
|
||||
def test_gateway_run_clamped(self):
|
||||
# The /usage stats handler was extracted from gateway/run.py into
|
||||
# gateway/slash_commands.py (god-file decomposition Phase 3b).
|
||||
src = self._read_file("gateway/slash_commands.py")
|
||||
# gateway/slash_commands.py and then gateway/slash_commands_status.py.
|
||||
src = self._read_file("gateway/slash_commands_status.py")
|
||||
# Check that the stats handler clamps the context pct with min(100, ...).
|
||||
# Assert the clamp intent, not a specific local name (the occupancy
|
||||
# value is read into a clamped `_lpt` local, #50421).
|
||||
assert "min(100, _lpt / ctx.context_length" in src, (
|
||||
"gateway/slash_commands.py stats pct is not clamped with min(100, ...)"
|
||||
"gateway/slash_commands_status.py stats pct is not clamped with min(100, ...)"
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user