diff --git a/cli.py b/cli.py index ad793c5dfb..d4c54b2705 100644 --- a/cli.py +++ b/cli.py @@ -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, diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index e49891749b..eb19f61bd0 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -2433,18 +2433,19 @@ class GatewaySlashCommandsMixin: return "\n".join(lines) - # Expensive-model confirmation gate (typed /model path). + # Selection-guard confirmation gate (typed /model 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._" ), diff --git a/hermes_cli/auth.py b/hermes_cli/auth.py index f9c4c0f4bd..3da395831f 100644 --- a/hermes_cli/auth.py +++ b/hermes_cli/auth.py @@ -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 diff --git a/hermes_cli/model_selection_guards.py b/hermes_cli/model_selection_guards.py new file mode 100644 index 0000000000..0d886497cb --- /dev/null +++ b/hermes_cli/model_selection_guards.py @@ -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), + ) diff --git a/hermes_cli/web_server.py b/hermes_cli/web_server.py index 365de042c3..d34e0a16e0 100644 --- a/hermes_cli/web_server.py +++ b/hermes_cli/web_server.py @@ -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, diff --git a/plugins/platforms/discord/adapter.py b/plugins/platforms/discord/adapter.py index c54224895f..b5a1f4ffa8 100644 --- a/plugins/platforms/discord/adapter.py +++ b/plugins/platforms/discord/adapter.py @@ -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(), ), diff --git a/plugins/platforms/telegram/adapter.py b/plugins/platforms/telegram/adapter.py index 093f5b2360..f31cd85e6a 100644 --- a/plugins/platforms/telegram/adapter.py +++ b/plugins/platforms/telegram/adapter.py @@ -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 diff --git a/tests/hermes_cli/test_model_selection_guards.py b/tests/hermes_cli/test_model_selection_guards.py new file mode 100644 index 0000000000..66da614daf --- /dev/null +++ b/tests/hermes_cli/test_model_selection_guards.py @@ -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) diff --git a/tui_gateway/server.py b/tui_gateway/server.py index d5719a7d6a..04b3c956e0 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -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,