feat(models): unify selection-time guards into one registry across all surfaces
Adds hermes_cli/model_selection_guards.py: a single evaluation point that runs every selection guard (cost + the new data-policy guard) and returns the warnings that fired. All seven model-selection surfaces (CLI picker, cli.py TUI modal, gateway typed /model, dashboard web_server, TUI gateway, Telegram and Discord pickers) now call the registry instead of importing model_cost_guard directly — so the data-training-tier warning from PR #81416 fires everywhere at once, and future guards need zero surface wiring. Guard modules keep their public APIs; existing mock patch points (hermes_cli.model_cost_guard.expensive_model_warning) remain valid.
This commit is contained in:
@@ -9546,9 +9546,9 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
if not getattr(result, "success", False):
|
||||
return True
|
||||
try:
|
||||
from hermes_cli.model_cost_guard import expensive_model_warning
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
|
||||
warning = expensive_model_warning(
|
||||
warning = combined_selection_warning(
|
||||
result.new_model,
|
||||
provider=result.target_provider,
|
||||
base_url=result.base_url or self.base_url or "",
|
||||
@@ -9565,7 +9565,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
("cancel", "Cancel", "Keep the current model."),
|
||||
]
|
||||
raw = self._prompt_text_input_modal(
|
||||
title="!!! Expensive Model Warning !!!",
|
||||
title=f"!!! {warning.title} !!!",
|
||||
detail=warning.message,
|
||||
choices=choices,
|
||||
timeout=120,
|
||||
|
||||
@@ -2433,18 +2433,19 @@ class GatewaySlashCommandsMixin:
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
# Expensive-model confirmation gate (typed /model <name> path).
|
||||
# Selection-guard confirmation gate (typed /model <name> path).
|
||||
# The pickers (Telegram/Discord inline keyboards, TUI, dashboard)
|
||||
# already confirm via their own UI affordances; this covers the
|
||||
# direct text command, which previously bypassed the guard.
|
||||
# expensive_model_warning() may hit models.dev or a /models endpoint
|
||||
# on a cache miss, so run it off the event loop.
|
||||
# Runs the unified registry (cost + data-policy + future 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_cost_guard import expensive_model_warning
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
|
||||
_cost_warning = await asyncio.to_thread(
|
||||
expensive_model_warning,
|
||||
combined_selection_warning,
|
||||
result.new_model,
|
||||
provider=result.target_provider,
|
||||
base_url=result.base_url or current_base_url or "",
|
||||
@@ -2461,7 +2462,7 @@ class GatewaySlashCommandsMixin:
|
||||
f"({current_model or 'unknown'})."
|
||||
)
|
||||
# "once" and "always" both proceed — there is no persistent
|
||||
# opt-out for the cost guard (each expensive switch should be
|
||||
# opt-out for selection guards (each guarded switch should be
|
||||
# an explicit decision).
|
||||
return await _finish_switch()
|
||||
|
||||
@@ -2469,9 +2470,9 @@ class GatewaySlashCommandsMixin:
|
||||
return await self._request_slash_confirm(
|
||||
event=event,
|
||||
command="model",
|
||||
title="Expensive Model Warning",
|
||||
title=_cost_warning.title,
|
||||
message=(
|
||||
f"⚠️ **Expensive Model Warning**\n\n{_cost_warning.message}\n\n"
|
||||
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._"
|
||||
),
|
||||
|
||||
+25
-55
@@ -7461,31 +7461,41 @@ def _reset_config_provider() -> Path:
|
||||
return config_path
|
||||
|
||||
|
||||
def _confirm_expensive_model_selection(
|
||||
def _confirm_selection_guards(
|
||||
model_id: str,
|
||||
*,
|
||||
provider: str = "",
|
||||
base_url: str = "",
|
||||
api_key: str = "",
|
||||
include_kinds: Optional[List[str]] = None,
|
||||
) -> bool:
|
||||
"""Prompt before saving a model whose known pricing exceeds guardrails."""
|
||||
try:
|
||||
from hermes_cli.model_cost_guard import expensive_model_warning
|
||||
"""Prompt before saving a model that trips any selection guard.
|
||||
|
||||
warning = expensive_model_warning(
|
||||
Runs the unified guard registry (cost + data-policy + future guards) via
|
||||
:mod:`hermes_cli.model_selection_guards` and shows one [y/N] confirm with
|
||||
every warning that fired. Returns True to proceed, False to cancel.
|
||||
"""
|
||||
try:
|
||||
from hermes_cli.model_selection_guards import (
|
||||
combined_message,
|
||||
selection_warnings,
|
||||
)
|
||||
|
||||
warnings = selection_warnings(
|
||||
model_id,
|
||||
provider=provider,
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
include_kinds=include_kinds,
|
||||
)
|
||||
except Exception:
|
||||
warning = None
|
||||
if warning is None:
|
||||
warnings = []
|
||||
if not warnings:
|
||||
return True
|
||||
|
||||
print()
|
||||
print("=" * 72)
|
||||
print(warning.message)
|
||||
print(combined_message(warnings))
|
||||
print("=" * 72)
|
||||
try:
|
||||
response = input("Switch anyway? [y/N]: ").strip().lower()
|
||||
@@ -7495,43 +7505,6 @@ def _confirm_expensive_model_selection(
|
||||
return response in {"y", "yes"}
|
||||
|
||||
|
||||
def _confirm_data_policy_selection(
|
||||
model_id: str,
|
||||
*,
|
||||
provider: str = "",
|
||||
base_url: str = "",
|
||||
) -> bool:
|
||||
"""Prompt before saving a model whose tier trains on your prompts/completions.
|
||||
|
||||
Mirrors :func:`_confirm_expensive_model_selection`. Keyed on the model id
|
||||
(e.g. Meta's ``-contributor`` tier), so it is not gated on a known provider.
|
||||
Returns True to proceed, False to cancel.
|
||||
"""
|
||||
try:
|
||||
from hermes_cli.model_data_policy_guard import data_training_warning
|
||||
|
||||
warning = data_training_warning(
|
||||
model_id,
|
||||
provider=provider,
|
||||
base_url=base_url,
|
||||
)
|
||||
except Exception:
|
||||
warning = None
|
||||
if warning is None:
|
||||
return True
|
||||
|
||||
print()
|
||||
print("=" * 72)
|
||||
print(warning.message)
|
||||
print("=" * 72)
|
||||
try:
|
||||
response = input("Use this data-training tier anyway? [y/N]: ").strip().lower()
|
||||
except (KeyboardInterrupt, EOFError):
|
||||
print()
|
||||
return False
|
||||
return response in {"y", "yes"}
|
||||
|
||||
|
||||
def _prompt_model_selection(
|
||||
model_ids: List[str],
|
||||
current_model: str = "",
|
||||
@@ -7564,20 +7537,17 @@ def _prompt_model_selection(
|
||||
def _confirmed_selection(mid: str) -> Optional[str]:
|
||||
if not mid:
|
||||
return None
|
||||
if confirm_provider and not _confirm_expensive_model_selection(
|
||||
# Unified guard registry (hermes_cli.model_selection_guards): the cost
|
||||
# guard only runs when a provider is known (pricing lookups need one);
|
||||
# id-keyed guards like the data-policy guard always run — they must
|
||||
# fire even via a custom endpoint or gateway.
|
||||
_kinds = None if confirm_provider else ["data_policy"]
|
||||
if not _confirm_selection_guards(
|
||||
mid,
|
||||
provider=confirm_provider,
|
||||
base_url=confirm_base_url,
|
||||
api_key=confirm_api_key,
|
||||
):
|
||||
return None
|
||||
# Data-policy guard runs regardless of provider (it keys on the model
|
||||
# id, e.g. a "-contributor" training tier), so it is NOT gated on
|
||||
# confirm_provider like the cost guard above.
|
||||
if not _confirm_data_policy_selection(
|
||||
mid,
|
||||
provider=confirm_provider,
|
||||
base_url=confirm_base_url,
|
||||
include_kinds=_kinds,
|
||||
):
|
||||
return None
|
||||
return mid
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
"""Unified selection-time guard registry for model switching surfaces.
|
||||
|
||||
Hermes has multiple model-selection surfaces (CLI picker, TUI, dashboard,
|
||||
gateway ``/model``, Telegram/Discord pickers, TUI-gateway RPC). Each of them
|
||||
previously imported ``model_cost_guard.expensive_model_warning`` directly, so
|
||||
every new guard class (e.g. the data-training-tier guard) had to be wired into
|
||||
every surface by hand — and inevitably missed some.
|
||||
|
||||
This module is the single evaluation point: ``selection_warnings()`` runs every
|
||||
registered guard and returns the warnings that fired. Surfaces render the
|
||||
result with their own confirm UX (stdin prompt, modal, inline keyboard,
|
||||
``confirm_required`` JSON) — that half stays per-surface; the *evaluation* half
|
||||
lives here. Adding a guard to ``_GUARDS`` makes it appear on every surface at
|
||||
once.
|
||||
|
||||
Guard modules (``model_cost_guard``, ``model_data_policy_guard``) keep their
|
||||
public APIs — existing tests and mock patch points remain valid; this module
|
||||
only aggregates them.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Iterable, List, Optional
|
||||
|
||||
from agent.models_dev import ModelInfo
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SelectionWarning:
|
||||
"""A selection-time warning a surface must confirm before applying."""
|
||||
|
||||
kind: str # "cost" | "data_policy" | future guard kinds
|
||||
title: str
|
||||
model: str
|
||||
provider: str
|
||||
message: str
|
||||
|
||||
|
||||
def _cost_guard(
|
||||
model_name: str,
|
||||
provider: Optional[str],
|
||||
base_url: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model_info: Optional[ModelInfo],
|
||||
) -> Optional[SelectionWarning]:
|
||||
from hermes_cli.model_cost_guard import expensive_model_warning
|
||||
|
||||
warning = expensive_model_warning(
|
||||
model_name,
|
||||
provider=provider,
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
model_info=model_info,
|
||||
)
|
||||
if warning is None:
|
||||
return None
|
||||
# Duck-typed access: tests (and future guard payloads) may supply objects
|
||||
# carrying only ``.message``.
|
||||
return SelectionWarning(
|
||||
kind="cost",
|
||||
title="Expensive Model Warning",
|
||||
model=getattr(warning, "model", model_name),
|
||||
provider=getattr(warning, "provider", provider or ""),
|
||||
message=warning.message,
|
||||
)
|
||||
|
||||
|
||||
def _data_policy_guard(
|
||||
model_name: str,
|
||||
provider: Optional[str],
|
||||
base_url: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model_info: Optional[ModelInfo],
|
||||
) -> Optional[SelectionWarning]:
|
||||
from hermes_cli.model_data_policy_guard import data_training_warning
|
||||
|
||||
warning = data_training_warning(
|
||||
model_name,
|
||||
provider=provider,
|
||||
base_url=base_url,
|
||||
)
|
||||
if warning is None:
|
||||
return None
|
||||
return SelectionWarning(
|
||||
kind="data_policy",
|
||||
title="Data-Training Tier Warning",
|
||||
model=getattr(warning, "model", model_name),
|
||||
provider=getattr(warning, "provider", provider or ""),
|
||||
message=warning.message,
|
||||
)
|
||||
|
||||
|
||||
# Registry, evaluated in order. Add new guard classes here — never at the
|
||||
# individual surfaces.
|
||||
_GUARDS = (
|
||||
_cost_guard,
|
||||
_data_policy_guard,
|
||||
)
|
||||
|
||||
|
||||
def selection_warnings(
|
||||
model_name: str,
|
||||
*,
|
||||
provider: Optional[str] = None,
|
||||
base_url: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
include_kinds: Optional[Iterable[str]] = None,
|
||||
) -> List[SelectionWarning]:
|
||||
"""Run every registered selection guard and return the warnings that fired.
|
||||
|
||||
Returns an empty list in the common case (no guard fired). Callers should
|
||||
run this after model resolution so aliases / provider-specific ids have
|
||||
settled, then surface the messages as a confirm step. ``include_kinds``
|
||||
optionally restricts which guard kinds run (e.g. auth.py's picker only runs
|
||||
the cost guard when a provider is known, but always runs the data-policy
|
||||
guard).
|
||||
|
||||
A misbehaving guard must never break model selection: individual guard
|
||||
exceptions are swallowed.
|
||||
"""
|
||||
wanted = set(include_kinds) if include_kinds is not None else None
|
||||
results: List[SelectionWarning] = []
|
||||
for guard in _GUARDS:
|
||||
try:
|
||||
warning = guard(model_name, provider, base_url, api_key, model_info)
|
||||
except Exception:
|
||||
continue
|
||||
if warning is None:
|
||||
continue
|
||||
if wanted is not None and warning.kind not in wanted:
|
||||
continue
|
||||
results.append(warning)
|
||||
return results
|
||||
|
||||
|
||||
def combined_message(warnings: List[SelectionWarning]) -> str:
|
||||
"""Join multiple warnings into one confirm-prompt body.
|
||||
|
||||
Surfaces that show a single confirm dialog use this when more than one
|
||||
guard fires (rare) — one prompt showing both blocks beats two sequential
|
||||
prompts.
|
||||
"""
|
||||
return "\n\n".join(w.message for w in warnings)
|
||||
|
||||
|
||||
def combined_selection_warning(
|
||||
model_name: str,
|
||||
*,
|
||||
provider: Optional[str] = None,
|
||||
base_url: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
) -> Optional[SelectionWarning]:
|
||||
"""Drop-in replacement for ``expensive_model_warning`` call sites.
|
||||
|
||||
Returns ``None`` when no guard fired, a single :class:`SelectionWarning`
|
||||
when exactly one fired, or a merged warning (``kind="multiple"``) whose
|
||||
``message`` stacks every fired guard. Surfaces that render one confirm
|
||||
dialog with ``warning.message`` can switch to this without reshaping their
|
||||
control flow.
|
||||
"""
|
||||
warnings = selection_warnings(
|
||||
model_name,
|
||||
provider=provider,
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
model_info=model_info,
|
||||
)
|
||||
if not warnings:
|
||||
return None
|
||||
if len(warnings) == 1:
|
||||
return warnings[0]
|
||||
return SelectionWarning(
|
||||
kind="multiple",
|
||||
title="Model Selection Warning",
|
||||
model=warnings[0].model,
|
||||
provider=warnings[0].provider,
|
||||
message=combined_message(warnings),
|
||||
)
|
||||
@@ -6698,12 +6698,12 @@ async def set_model_assignment(body: ModelAssignment, profile: Optional[str] = N
|
||||
# event-loop thread could cross-restore the module globals).
|
||||
if model and not body.confirm_expensive_model:
|
||||
try:
|
||||
from hermes_cli.model_cost_guard import expensive_model_warning
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
|
||||
# Pricing lookup can hit models.dev / a /models endpoint on a
|
||||
# cache miss — keep it off the event loop.
|
||||
warning = await asyncio.to_thread(
|
||||
expensive_model_warning,
|
||||
combined_selection_warning,
|
||||
model,
|
||||
provider=provider,
|
||||
base_url=base_url,
|
||||
|
||||
@@ -9259,12 +9259,12 @@ def _define_discord_view_classes() -> None:
|
||||
|
||||
async def _expensive_warning_for(self, model_id: str):
|
||||
try:
|
||||
from hermes_cli.model_cost_guard import expensive_model_warning
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
|
||||
# Pricing lookup can hit models.dev / a /models endpoint on a
|
||||
# cache miss — keep it off the event loop.
|
||||
return await asyncio.to_thread(
|
||||
expensive_model_warning,
|
||||
combined_selection_warning,
|
||||
model_id,
|
||||
provider=self._selected_provider,
|
||||
)
|
||||
@@ -9363,7 +9363,7 @@ def _define_discord_view_classes() -> None:
|
||||
self._build_expensive_confirm(model_id)
|
||||
await interaction.response.edit_message(
|
||||
embed=discord.Embed(
|
||||
title="⚠ Expensive Model Warning",
|
||||
title=f"⚠ {warning.title}",
|
||||
description=warning.message,
|
||||
color=discord.Color.red(),
|
||||
),
|
||||
|
||||
@@ -6519,12 +6519,12 @@ class TelegramAdapter(BasePlatformAdapter):
|
||||
return
|
||||
|
||||
try:
|
||||
from hermes_cli.model_cost_guard import expensive_model_warning
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
|
||||
# Pricing lookup can hit models.dev / a /models endpoint on a
|
||||
# cache miss — keep it off the event loop.
|
||||
warning = await asyncio.to_thread(
|
||||
expensive_model_warning,
|
||||
combined_selection_warning,
|
||||
model_id,
|
||||
provider=provider_slug,
|
||||
)
|
||||
@@ -6540,12 +6540,12 @@ class TelegramAdapter(BasePlatformAdapter):
|
||||
])
|
||||
await query.edit_message_text(
|
||||
text=self.format_message(
|
||||
f"⚠ *Expensive Model Warning*\n\n{warning.message}"
|
||||
f"⚠ *{warning.title}*\n\n{warning.message}"
|
||||
),
|
||||
parse_mode=ParseMode.MARKDOWN_V2,
|
||||
reply_markup=keyboard,
|
||||
)
|
||||
await query.answer(text="Confirm expensive model")
|
||||
await query.answer(text="Confirm model selection")
|
||||
return
|
||||
|
||||
switch_failed = False
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
"""Tests for the unified model-selection guard registry."""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
from hermes_cli.model_selection_guards import (
|
||||
SelectionWarning,
|
||||
combined_message,
|
||||
combined_selection_warning,
|
||||
selection_warnings,
|
||||
)
|
||||
|
||||
|
||||
def test_no_guard_fires_on_ordinary_model():
|
||||
# No pricing data (no provider), no data-policy rule match.
|
||||
assert selection_warnings("some/ordinary-model") == []
|
||||
assert combined_selection_warning("some/ordinary-model") is None
|
||||
|
||||
|
||||
def test_data_policy_guard_fires_through_registry():
|
||||
warnings = selection_warnings("muse-spark-1.2-contributor", provider="custom")
|
||||
kinds = [w.kind for w in warnings]
|
||||
assert "data_policy" in kinds
|
||||
w = next(w for w in warnings if w.kind == "data_policy")
|
||||
assert "train" in w.message.lower()
|
||||
assert w.title == "Data-Training Tier Warning"
|
||||
|
||||
|
||||
def test_include_kinds_filters_guards():
|
||||
warnings = selection_warnings(
|
||||
"muse-spark-1.2-contributor",
|
||||
provider="custom",
|
||||
include_kinds=["cost"],
|
||||
)
|
||||
assert all(w.kind == "cost" for w in warnings)
|
||||
assert not any(w.kind == "data_policy" for w in warnings)
|
||||
|
||||
|
||||
def test_combined_selection_warning_single():
|
||||
w = combined_selection_warning("muse-spark-1.2-contributor")
|
||||
assert w is not None
|
||||
assert w.kind == "data_policy"
|
||||
|
||||
|
||||
def test_combined_selection_warning_merges_multiple():
|
||||
cost = SelectionWarning(
|
||||
kind="cost",
|
||||
title="Expensive Model Warning",
|
||||
model="m",
|
||||
provider="p",
|
||||
message="COST BLOCK",
|
||||
)
|
||||
policy = SelectionWarning(
|
||||
kind="data_policy",
|
||||
title="Data-Training Tier Warning",
|
||||
model="m",
|
||||
provider="p",
|
||||
message="POLICY BLOCK",
|
||||
)
|
||||
with patch(
|
||||
"hermes_cli.model_selection_guards._GUARDS",
|
||||
(lambda *a: cost, lambda *a: policy),
|
||||
):
|
||||
merged = combined_selection_warning("m")
|
||||
assert merged is not None
|
||||
assert merged.kind == "multiple"
|
||||
assert "COST BLOCK" in merged.message
|
||||
assert "POLICY BLOCK" in merged.message
|
||||
|
||||
|
||||
def test_misbehaving_guard_never_breaks_selection():
|
||||
def _boom(*args):
|
||||
raise RuntimeError("bad guard")
|
||||
|
||||
with patch(
|
||||
"hermes_cli.model_selection_guards._GUARDS",
|
||||
(_boom,),
|
||||
):
|
||||
assert selection_warnings("anything") == []
|
||||
|
||||
|
||||
def test_combined_message_joins_blocks():
|
||||
a = SelectionWarning("cost", "t1", "m", "p", "AAA")
|
||||
b = SelectionWarning("data_policy", "t2", "m", "p", "BBB")
|
||||
assert combined_message([a, b]) == "AAA\n\nBBB"
|
||||
|
||||
|
||||
def test_cost_guard_still_fires_through_registry():
|
||||
# The registry must preserve the existing cost-guard behavior; feed it
|
||||
# explicit model_info so no network lookup is needed.
|
||||
from agent.models_dev import ModelInfo
|
||||
|
||||
info = ModelInfo(
|
||||
id="pricey/model",
|
||||
name="pricey/model",
|
||||
family="",
|
||||
provider_id="test",
|
||||
cost_input=50.0,
|
||||
cost_output=200.0,
|
||||
)
|
||||
warnings = selection_warnings(
|
||||
"pricey/model", provider="test", model_info=info
|
||||
)
|
||||
assert any(w.kind == "cost" for w in warnings)
|
||||
@@ -4650,9 +4650,9 @@ def _apply_model_switch(
|
||||
|
||||
if not confirm_expensive_model:
|
||||
try:
|
||||
from hermes_cli.model_cost_guard import expensive_model_warning
|
||||
from hermes_cli.model_selection_guards import combined_selection_warning
|
||||
|
||||
warning = expensive_model_warning(
|
||||
warning = combined_selection_warning(
|
||||
result.new_model,
|
||||
provider=result.target_provider,
|
||||
base_url=result.base_url or current_base_url,
|
||||
|
||||
Reference in New Issue
Block a user