142 lines
4.8 KiB
Python
142 lines
4.8 KiB
Python
"""Unified selection-time guard registry for model switching surfaces.
|
|
|
|
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 _wrap(kind: str, title: str, warning, model_name: str, provider: Optional[str]):
|
|
"""Lift a raw guard payload into a :class:`SelectionWarning` (None passes through).
|
|
|
|
Duck-typed access: tests (and future guard payloads) may supply objects carrying only
|
|
``.message``.
|
|
"""
|
|
if warning is None:
|
|
return None
|
|
return SelectionWarning(
|
|
kind=kind,
|
|
title=title,
|
|
model=getattr(warning, "model", model_name),
|
|
provider=getattr(warning, "provider", provider or ""),
|
|
message=warning.message,
|
|
)
|
|
|
|
|
|
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
|
|
)
|
|
return _wrap("cost", "Expensive Model Warning", warning, model_name, provider)
|
|
|
|
|
|
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)
|
|
return _wrap("data_policy", "Data-Training Tier Warning", warning, model_name, provider)
|
|
|
|
|
|
# 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.
|
|
|
|
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 not None and (wanted is None or warning.kind in wanted):
|
|
results.append(warning)
|
|
return results
|
|
|
|
|
|
def combined_message(warnings: List[SelectionWarning]) -> str:
|
|
"""Join multiple warnings into one confirm-prompt body.
|
|
|
|
Used by surfaces with a single confirm dialog 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, the single :class:`SelectionWarning` when one fired,
|
|
or a merged ``kind="multiple"`` warning stacking every message — so surfaces rendering one
|
|
confirm dialog from ``warning.message`` can switch without reshaping 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),
|
|
)
|