From 72490fa6cda887d2ae89c95591e5a60287b96fc7 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 16:39:35 -0700 Subject: [PATCH] refactor(model_setup_flows): extract custom-endpoint, Azure Foundry and Bedrock wizards into concern modules hermes_cli/model_setup_flows_custom.py (_model_flow_custom, _model_flow_named_custom + helpers), model_setup_flows_azure.py (_model_flow_azure_foundry + Entra preflight / picker), model_setup_flows_bedrock.py (BEDROCK_GEO_PREFIXES, routability predicates, both Bedrock flows). Bodies moved verbatim (AST slices); every name re-exported from model_setup_flows (noqa: F401) so hermes_cli.main and test imports keep resolving. Origin: 1975 -> 1159 lines. Flow corpus (396 cases): 0 diffs; 34 test files / 364 tests green. --- hermes_cli/model_setup_flows.py | 876 +----------------------- hermes_cli/model_setup_flows_azure.py | 255 +++++++ hermes_cli/model_setup_flows_bedrock.py | 224 ++++++ hermes_cli/model_setup_flows_common.py | 2 + hermes_cli/model_setup_flows_custom.py | 409 +++++++++++ 5 files changed, 920 insertions(+), 846 deletions(-) create mode 100644 hermes_cli/model_setup_flows_azure.py create mode 100644 hermes_cli/model_setup_flows_bedrock.py create mode 100644 hermes_cli/model_setup_flows_custom.py diff --git a/hermes_cli/model_setup_flows.py b/hermes_cli/model_setup_flows.py index 8893ab793f..0868d7e7f7 100644 --- a/hermes_cli/model_setup_flows.py +++ b/hermes_cli/model_setup_flows.py @@ -15,13 +15,10 @@ from __future__ import annotations import argparse import os -import subprocess -import urllib.parse -from hermes_cli.cli_output import line_input from hermes_cli.config import clear_model_endpoint_credentials -from hermes_cli.providers import custom_provider_slug from hermes_cli.model_setup_flows_common import ( # noqa: F401 + _HTTP, _activate_provider_model, _ask, _begin_model_config, @@ -43,43 +40,35 @@ from hermes_cli.model_setup_flows_common import ( # noqa: F401 _say, _show_curated, ) - -_HTTP = ("http://", "https://") - - -# AWS cross-region inference profile prefixes. A geo-prefixed profile only routes -# from endpoints in its own geography (us.* from eu-central-2 is rejected by AWS -# regardless of credentials); global.* routes from everywhere. -BEDROCK_GEO_PREFIXES = ("us.", "eu.", "ap.", "apac.", "jp.", "ca.", "sa.", "me.", "af.") - -# region-name prefixes -> inference-profile geo prefix -_REGION_GEO = (("us.", ("us-", "us_gov")), ("eu.", ("eu-",)), ("ap.", ("ap-",)), ("ca.", ("ca-",)), - ("sa.", ("sa-",)), ("me.", ("me-",)), ("af.", ("af-",))) - - -def bedrock_region_geo_prefix(region_name: str) -> str: - """Map an AWS region name to its inference-profile geo prefix ('' = unknown).""" - r = (region_name or "").lower() - return next((geo for geo, prefixes in _REGION_GEO if r.startswith(prefixes)), "") - - -def bedrock_model_routable_from_region(model_id: str, region_name: str) -> bool: - """True when *model_id* can be invoked from *region_name*'s endpoint. - - Bare foundation-model ids and ``global.*`` profiles route from anywhere; - geo-prefixed profiles only from their own geography. Unknown regions hide nothing. - """ - mid = (model_id or "").lower() - matched_geo = next((p for p in BEDROCK_GEO_PREFIXES if mid.startswith(p)), None) - if matched_geo is None or mid.startswith("global."): - return True - geo = bedrock_region_geo_prefix(region_name) - if not geo: - return True - if geo == "ap.": - # Asia-Pacific regions can carry ap./apac./jp. profile spellings. - return matched_geo in ("ap.", "apac.", "jp.") - return matched_geo == geo +from hermes_cli.model_setup_flows_custom import ( # noqa: F401 + _parse_context_length, + _probe_custom_endpoint, + _pick_detected_model, + _model_flow_custom, + _configured_model_ids, + _discover_named_custom_models, + _pick_named_custom_model, + _model_flow_named_custom, +) +from hermes_cli.model_setup_flows_azure import ( # noqa: F401 + _azure_mode_label, + _azure_entra_preflight, + _azure_pick_model, + _model_flow_azure_foundry, +) +from hermes_cli.model_setup_flows_bedrock import ( # noqa: F401 + BEDROCK_GEO_PREFIXES, + _REGION_GEO, + bedrock_region_geo_prefix, + bedrock_model_routable_from_region, + _model_flow_bedrock_api_key, + _BEDROCK_EXCLUDE_PREFIXES, + _BEDROCK_EXCLUDE_SUBSTRINGS, + _BEDROCK_PROFILE_PREFIXES, + _BEDROCK_RECOMMENDED_BASES, + _bedrock_text_model_ids, + _model_flow_bedrock, +) def _model_flow_openrouter(config, current_model=""): @@ -519,641 +508,6 @@ def _model_flow_minimax_oauth(config, current_model="", args=None): _activate_provider_model(selected, "minimax-oauth", creds["base_url"], f"\u2713 Using MiniMax model: {selected}", no_change=None) -def _parse_context_length(text: str): - """``128k`` / ``128,000`` -> int; None when blank, non-positive, or unparsable (warns).""" - if not text: - return None - try: - value = int(text.replace(",", "").replace("k", "000").replace("K", "000")) - except ValueError: - print(f"Invalid context length: {text} — will auto-detect.") - return None - return value if value > 0 else None - - -def _probe_custom_endpoint(effective_key: str, effective_url: str) -> tuple[dict, str]: - """Verify a custom endpoint via ``probe_api_models`` and report; returns - ``(probe, effective_url)`` where the URL may be the working fallback base.""" - from hermes_cli.models import probe_api_models - - probe = probe_api_models(effective_key, effective_url) - if probe.get("used_fallback") and probe.get("resolved_base_url"): - print(f"Warning: endpoint verification worked at {probe['resolved_base_url']}/models, " - f"not the exact URL you entered. Saving the working base URL instead.") - effective_url = probe["resolved_base_url"] - elif probe.get("models") is not None: - print(f"Verified endpoint via {probe.get('probed_url')} ({len(probe.get('models') or [])} model(s) visible)") - else: - print(f"Warning: could not verify this endpoint via {probe.get('probed_url')}. Hermes will still save it.") - suggested = probe.get("suggested_base_url") - if suggested and suggested.endswith("/v1"): - print(f" If this server expects /v1 in the path, try base URL: {suggested}") - elif suggested: - print(f" If /v1 should not be in the base URL, try: {suggested}") - return probe, effective_url - - -def _pick_detected_model(detected_models: list) -> str: - """Model-name step of the custom flow: confirm a single detection, number-pick from - several, or type one. Raises KeyboardInterrupt/EOFError like the prompts it wraps.""" - manual = "Model name (e.g. gpt-4, llama-3-70b): " - if len(detected_models) == 1: - print(f" Detected model: {detected_models[0]}") - if input(" Use this model? [Y/n]: ").strip().lower() in {"", "y", "yes"}: - return detected_models[0] - return line_input(manual).strip() - if len(detected_models) > 1: - print(" Available models:") - for i, m in enumerate(detected_models, 1): - print(f" {i}. {m}") - pick = input(f" Select model [1-{len(detected_models)}] or type name: ").strip() - if pick.isdigit() and 1 <= int(pick) <= len(detected_models): - return detected_models[int(pick) - 1] - return pick - return line_input(manual).strip() - - -def _model_flow_custom(config): - """Custom endpoint: collect URL, API key, and model name. - - Also saves the endpoint to ``custom_providers`` in config.yaml so it appears - in the provider menu on subsequent runs. - """ - from hermes_cli.main import _auto_provider_name, _prompt_custom_api_mode_selection, _save_custom_provider - from hermes_cli.auth import _save_model_choice, deactivate_provider - from hermes_cli.config import custom_endpoint_key_env, get_env_value, save_env_value - from hermes_cli.secret_prompt import masked_secret_prompt - - current_url = get_env_value("OPENAI_BASE_URL") or "" - current_key = get_env_value("OPENAI_API_KEY") or "" - - print("Custom OpenAI-compatible endpoint configuration:") - if current_url: - print(f" Current URL: {current_url}") - if current_key: - print(f" Current key: {current_key[:8]}...") - print() - - try: - base_url = line_input(f"API base URL [{current_url or 'e.g. https://api.example.com/v1'}]: ").strip() - api_key = masked_secret_prompt(f"API key [{current_key[:8] + '...' if current_key else 'optional'}]: ").strip() - except (KeyboardInterrupt, EOFError): - print("\nCancelled.") - return - - if not base_url and not current_url: - print("No URL provided. Cancelled.") - return - effective_url = base_url or current_url - if not effective_url.startswith(_HTTP): - print(f"Invalid URL: {effective_url} (must start with http:// or https://)") - return - effective_key = api_key or current_key - - # Most local servers (Ollama, vLLM, llama.cpp) need /v1 for OpenAI-compatible - # chat completions — offer to append it when the URL looks local without it. - _url_lower = effective_url.rstrip("/").lower() - _looks_local = any(h in _url_lower for h in ("localhost", "127.0.0.1", "0.0.0.0", ":11434", ":8080", ":5000")) - if _looks_local and not _url_lower.endswith("/v1"): - _say("", " Hint: Did you mean to add /v1 at the end?", - " Most local model servers (Ollama, vLLM, llama.cpp) require it.", f" e.g. {effective_url.rstrip('/')}/v1") - if _ask(" Add /v1? [Y/n]: ", raw=True, cancel_msg=None, on_cancel="n").lower() in {"", "y", "yes"}: - effective_url = effective_url.rstrip("/") + "/v1" - print(f" Updated URL: {effective_url}") - print() - - probe, effective_url = _probe_custom_endpoint(effective_key, effective_url) - - # Ask for the API mode explicitly so codex-compatible custom providers don't - # silently fall back to chat_completions. - current_model_cfg = config.get("model") - current_api_mode = str(current_model_cfg.get("api_mode") or "").strip() if isinstance(current_model_cfg, dict) else "" - api_mode = _prompt_custom_api_mode_selection(effective_url, current_api_mode=current_api_mode) - print(f" API mode: {api_mode}" if api_mode else " API mode: auto-detect") - - # Select model — use probe results when available, fall back to manual input - try: - model_name = _pick_detected_model(probe.get("models") or []) - context_length_str = line_input("Context length in tokens [leave blank for auto-detect]: ").strip() - # Display name — shown in the provider menu on future runs - default_name = _auto_provider_name(effective_url) - display_name = line_input(f"Display name [{default_name}]: ").strip() or default_name - except (KeyboardInterrupt, EOFError): - print("\nCancelled.") - return - context_length = _parse_context_length(context_length_str) - - # The key goes to .env and config.yaml only references it. Keyed on host:port - # so two servers on one machine keep separate credentials. - custom_key_env = "" - if effective_key: - _parsed = urllib.parse.urlparse(effective_url) - _identity = _parsed.hostname or "" - if _parsed.port: - _identity = f"{_identity}_{_parsed.port}" - custom_key_env = custom_endpoint_key_env(_identity) - save_env_value(custom_key_env, effective_key) - print(f" API key saved to .env as {custom_key_env}") - - def _apply_endpoint(model: dict) -> None: - model["provider"] = "custom" - model["base_url"] = effective_url - if custom_key_env: - model["api_key"] = f"${{{custom_key_env}}}" - if api_mode: - model["api_mode"] = api_mode - else: - model.pop("api_mode", None) - - if model_name: - _save_model_choice(model_name) - cfg, model = _load_config_model_section() - _apply_endpoint(model) - _commit_model_config(cfg) - # Sync the caller's config dict so the setup wizard's final save_config(config) - # doesn't overwrite model.provider/base_url with its stale values. - config["model"] = dict(model) - print(f"Default model set to: {model_name} (via {effective_url})") - else: - if base_url or api_key: - deactivate_provider() - # Even without a model name, persist the endpoint on the caller's config dict. - _caller_model = config.get("model") - if not isinstance(_caller_model, dict): - _caller_model = {"default": _caller_model} if _caller_model else {} - _apply_endpoint(_caller_model) - config["model"] = _caller_model - print("Endpoint saved. Use `/model` in chat or `hermes model` to set a model.") - - # Auto-save to custom_providers so it appears in the menu next time - _save_custom_provider(effective_url, effective_key, model_name or "", context_length=context_length, - name=display_name, api_mode=api_mode, key_env=custom_key_env) - _prune_replaced_custom_model_config_credentials(effective_url, provider_name=display_name) - - -def _azure_mode_label(mode: str) -> str: - return "OpenAI-style" if mode == "chat_completions" else "Anthropic-style" - - -def _azure_entra_preflight(current_entra: dict): - """Entra ID credential preflight for the Azure flow. Returns - ``(token_provider, entra_overrides)``; ``None`` when the user cancelled; - ``False`` when the adapter is missing (caller falls back to API-key auth).""" - try: - from agent.azure_identity_adapter import ( - EntraIdentityConfig, SCOPE_AI_AZURE_DEFAULT, build_token_provider, describe_active_credential, - has_azure_identity_installed, - ) - except ImportError as exc: - _say("", f"⚠ Could not import azure-identity adapter: {exc}", " Falling back to API key auth.") - return False - - print() - if not has_azure_identity_installed(): - _say("◐ The 'azure-identity' package is not installed yet.", - " Hermes will install it now (the preflight below triggers the lazy-install). " - "To skip lazy installs, run: pip install azure-identity") - - # Only the optional scope override is persisted; identity selection (tenant, - # user-assigned MI, workload identity, SP) stays in AZURE_* SDK env vars. - entra_overrides: dict = {} - _persisted_scope_override = str(current_entra.get("scope") or "").strip() - entra_scope = _persisted_scope_override or SCOPE_AI_AZURE_DEFAULT - if _persisted_scope_override: - entra_overrides["scope"] = _persisted_scope_override - - _say("", "◐ Probing Microsoft Entra ID credential chain (up to 10s)...") - _config = EntraIdentityConfig(scope=entra_scope) - info = describe_active_credential(config=_config, timeout_seconds=10.0) - if info.get("ok"): - env_sources = info.get("env_sources") or [] - tag = ", ".join(env_sources) if env_sources else "default chain" - print(f"✓ Entra ID token acquired ({tag}, scope={entra_scope})") - else: - err = info.get("error") or "credential chain exhausted" - hint = info.get("hint") or ( - "Run `az login`, attach a managed identity to this VM, or set AZURE_TENANT_ID/AZURE_CLIENT_ID/AZURE_CLIENT_SECRET." - ) - _say(f"⚠ {err}", f" Hint: {hint}") - ans = _ask("Save Entra config anyway and validate later? [Y/n]: ", raw=True) - if ans is None: - return None - if ans.lower() not in ("", "y", "yes"): - print("Cancelled.") - return None - - # Best-effort token provider for the detection probe; on failure the probe - # falls back to manual entry. - try: - token_provider = build_token_provider(config=_config) - except Exception as exc: - print(f"⚠ Could not build token provider for probing: {exc}") - token_provider = None - return token_provider, entra_overrides - - -def _azure_pick_model(discovered_models: list, current_model: str): - """Model/deployment step of the Azure flow; None when cancelled.""" - if not discovered_models: - model_name = _ask(f"Model / deployment name [{current_model or 'e.g. gpt-5.4, claude-sonnet-4-6'}]: ") - return None if model_name is None else (model_name or current_model) - print("Available models on this endpoint:") - for i, mid in enumerate(discovered_models[:30], start=1): - print(f" {i:>2}. {mid}") - if len(discovered_models) > 30: - print(f" ... and {len(discovered_models) - 30} more (type name manually if not shown)") - print() - pick = _ask(f"Pick by number, or type a deployment name [{current_model or discovered_models[0]}]: ", raw=True) - if pick is None: - return None - if not pick: - return current_model or discovered_models[0] - if pick.isdigit() and 1 <= int(pick) <= min(len(discovered_models), 30): - return discovered_models[int(pick) - 1] - return pick - - -def _model_flow_azure_foundry(config, current_model=""): - """Azure Foundry provider: configure endpoint, auth mode, API mode, and model. - - Two transports (OpenAI-style ``/v1/chat/completions``, Anthropic-style - ``/v1/messages``) and two auth modes: **API key** (``AZURE_FOUNDRY_API_KEY``) or - **Microsoft Entra ID** (keyless RBAC via ``azure-identity``; the same ``Azure AI - User`` role covers both transports — OpenAI SDK takes a callable ``api_key``, - Anthropic gets a bearer-injecting ``httpx.Client`` from - :func:`agent.azure_identity_adapter.build_bearer_http_client`). - - Detection order: ``/anthropic`` URL suffix → Anthropic; ``GET /models`` - success → OpenAI-style + model picker; Anthropic Messages probe; manual entry. - Context length resolves via :func:`agent.model_metadata.get_model_context_length`. - """ - from hermes_cli.config import get_env_value, save_env_value - from hermes_cli import azure_detect - - # ── Load current Azure Foundry configuration ───────────────────── - model_cfg = config.get("model", {}) - current_base_url = current_api_mode = "" - current_auth_mode, current_entra = "api_key", {} - if isinstance(model_cfg, dict) and model_cfg.get("provider") == "azure-foundry": - current_base_url = str(model_cfg.get("base_url", "") or "") - current_api_mode = str(model_cfg.get("api_mode", "") or "") - current_auth_mode = str(model_cfg.get("auth_mode") or "api_key").strip().lower() or "api_key" - _cur_entra = model_cfg.get("entra") or {} - current_entra = _cur_entra if isinstance(_cur_entra, dict) else {} - current_api_key = get_env_value("AZURE_FOUNDRY_API_KEY") or "" - - _say("", "Azure Foundry Configuration", "=" * 50, "", - "Azure Foundry can host models with either OpenAI-style or", - "Anthropic-style API endpoints. Hermes will probe your", - "endpoint to auto-detect the transport and the deployed", - "models when possible.", "") - if current_base_url: - print(f" Current endpoint: {current_base_url}") - if current_api_mode: - print(f" Current API mode: {_azure_mode_label(current_api_mode)}") - if current_auth_mode == "entra_id": - print(" Current auth mode: Microsoft Entra ID (keyless)") - elif current_api_key: - print(f" Current auth mode: API key ({current_api_key[:8]}...)") - print() - - # ── Step 1: endpoint URL ───────────────────────────────────────── - _placeholder = current_base_url or ( - "e.g. https://.openai.azure.com/openai/v1 or https://.services.ai.azure.com/anthropic" - ) - base_url = _ask(f"API endpoint URL [{_placeholder}]: ") - if base_url is None: - return - effective_url = (base_url or current_base_url).rstrip("/") - if not effective_url: - print("No endpoint URL provided. Cancelled.") - return - if not effective_url.startswith(_HTTP): - print(f"Invalid URL: {effective_url} (must start with http:// or https://)") - return - - # ── Step 2: authentication mode ────────────────────────────────── - _say("", "Authentication:", " 1. API key (AZURE_FOUNDRY_API_KEY in .env)", - " 2. Microsoft Entra ID (managed identity / workload identity / az login)", - " Recommended by Microsoft. Works for both OpenAI-style and Anthropic-style endpoints.", - " Requires the 'Azure AI User' role on the Foundry resource.") - _auth_default = "2" if current_auth_mode == "entra_id" else "1" - auth_choice = _ask(f"Authentication mode [1/2] ({_auth_default}): ", raw=True) - if auth_choice is None: - return - use_entra = (auth_choice or _auth_default) == "2" - - # ── Step 3: credentials (key OR Entra preflight) ───────────────── - effective_key: str = "" - entra_overrides: dict = {} - token_provider = None # callable when entra - if use_entra: - preflight = _azure_entra_preflight(current_entra) - if preflight is None: - return - if preflight is False: - use_entra = False - else: - token_provider, entra_overrides = preflight - if not use_entra: - print() - api_key = _ask(f"API key [{current_api_key[:8] + '...' if current_api_key else 'required'}]: ", secret=True) - if api_key is None: - return - effective_key = api_key or current_api_key - if not effective_key: - print("No API key provided. Cancelled.") - return - - # ── Step 4: auto-detect transport + models ─────────────────────── - _say("", "◐ Probing endpoint to auto-detect transport and models...") - detection = azure_detect.detect(effective_url, api_key=effective_key, token_provider=token_provider) - discovered_models: list[str] = list(detection.models) - api_mode: str = detection.api_mode or "" - if api_mode: - print(f"✓ Detected API transport: {_azure_mode_label(api_mode)}") - if detection.reason: - print(f" ({detection.reason})") - if discovered_models: - print(f"✓ Found {len(discovered_models)} deployed model(s) on this endpoint") - else: - _say(f"⚠ Auto-detection incomplete: {detection.reason}", "", - "Select the API format your Azure Foundry endpoint uses:", - " 1. OpenAI-style (POST /v1/chat/completions)", - " For: GPT models, Llama, Mistral, and most open models", - " 2. Anthropic-style (POST /v1/messages)", - " For: Claude models deployed via Anthropic API format") - default_choice = "2" if current_api_mode == "anthropic_messages" else "1" - mode_choice = _ask(f"API format [1/2] ({default_choice}): ", raw=True) - if mode_choice is None: - return - api_mode = "anthropic_messages" if (mode_choice or default_choice) == "2" else "chat_completions" - - # ── Step 5: model name ─────────────────────────────────────────── - print() - effective_model = _azure_pick_model(discovered_models, current_model) - if effective_model is None: - return - if not effective_model: - print("No model name provided. Cancelled.") - return - - # ── Step 6: context-length lookup ──────────────────────────────── - ctx_len = azure_detect.lookup_context_length(effective_model, effective_url, api_key=effective_key, token_provider=token_provider) - - # ── Step 7: persist ────────────────────────────────────────────── - if not use_entra: - save_env_value("AZURE_FOUNDRY_API_KEY", effective_key) - cfg, model = _load_config_model_section() - model["provider"] = "azure-foundry" - model["base_url"] = effective_url - model["api_mode"] = api_mode - model["default"] = effective_model - model["auth_mode"] = "entra_id" if use_entra else "api_key" - clear_model_endpoint_credentials(model, clear_api_mode=False) - # Persist only a non-default Entra scope so config.yaml stays tidy. - clean_entra = {k: v for k in ("scope",) if (v := entra_overrides.get(k))} - if use_entra and clean_entra: - model["entra"] = clean_entra - else: - model.pop("entra", None) - if ctx_len: - model["context_length"] = ctx_len - _commit_model_config(cfg) - config["model"] = dict(model) - - # Clear conflicting env vars so auxiliary clients don't pick up a stale - # OpenAI base URL / key. - for var in ("OPENAI_BASE_URL", "OPENAI_API_KEY"): - if get_env_value(var): - save_env_value(var, "") - - _say("", "✓ Azure Foundry configured:", f" Endpoint: {effective_url}", - f" API mode: {_azure_mode_label(api_mode)}", - f" Auth: {'Microsoft Entra ID (keyless)' if use_entra else 'API key'}", - f" Model: {effective_model}", - f" Context length: {ctx_len:,} tokens" if ctx_len else " Context length: not auto-detected (will fall back at runtime)", - "") - - -def _configured_model_ids(cfg_models) -> list[str]: - """Model ids from a ``custom_providers[].models`` mapping or list (marker keys skipped).""" - if isinstance(cfg_models, dict): - markers = {"__explicit_model_allowlist__", "__discovered_model_catalog__"} - return [str(m) for m in cfg_models if m not in markers and str(m).strip()] - out: list[str] = [] - if isinstance(cfg_models, list): - for entry in cfg_models: - if isinstance(entry, dict): - model_id = str(entry.get("id") or entry.get("model") or "").strip() - else: - model_id = str(entry).strip() if isinstance(entry, str) else "" - if model_id: - out.append(model_id) - return out - - -def _discover_named_custom_models(provider_info: dict, api_key: str, configured_models: list, explicit_catalog: bool): - """Live catalog probe for a named custom endpoint (native ``/api/tags`` for Ollama). - Returns ``(models, native_catalog_empty)``; persists the live catalog as a side effect.""" - from hermes_cli.config import normalize_extra_headers - from hermes_cli.models import ( - fetch_api_models, fetch_ollama_local_models, _get_ollama_native_headers, _normalize_openai_base_url, - should_use_ollama_native_catalog, - ) - - name, base_url = provider_info["name"], provider_info["base_url"] - api_mode = provider_info.get("api_mode", "") - provider_key = (provider_info.get("provider_key") or "").strip() - print("Fetching available models...") - fetch_kwargs = {"timeout": 8.0} - if api_mode: - fetch_kwargs["api_mode"] = api_mode - native_catalog_provider = "ollama" if provider_key.lower() == "ollama" or name.strip().lower() == "ollama" else "custom" - extra_headers = normalize_extra_headers(provider_info.get("extra_headers")) or {} - candidate_headers = _get_ollama_native_headers(base_url, api_key=api_key) - for key in tuple(candidate_headers): - if any(key.lower() == existing.lower() for existing in extra_headers): - del candidate_headers[key] - candidate_headers.update(extra_headers) - caller_has_authorization = any(key.lower() == "authorization" for key in extra_headers) - if api_key and not caller_has_authorization: - for key in tuple(candidate_headers): - if key.lower() == "authorization": - del candidate_headers[key] - candidate_headers["Authorization"] = f"Bearer {api_key}" - use_native = should_use_ollama_native_catalog(native_catalog_provider, base_url, headers=candidate_headers or None) - native_headers_arg = candidate_headers or None if use_native else (extra_headers or None) - native_catalog_empty = False - if use_native: - if explicit_catalog and configured_models: - live_models = configured_models - else: - live_models = fetch_ollama_local_models(base_url, timeout=8.0, headers=native_headers_arg) - native_catalog_empty = live_models == [] - if live_models is None: - live_models = fetch_api_models(api_key, _normalize_openai_base_url(base_url), headers=native_headers_arg, **fetch_kwargs) - native_catalog_empty = False - else: - live_models = fetch_api_models(api_key, base_url, headers=native_headers_arg, **fetch_kwargs) - models = configured_models if explicit_catalog else [] if native_catalog_empty else (live_models or configured_models) - # Persist the live catalog to the custom_providers entry so no-probe surfaces - # (dashboard, desktop, ACP) show the full list; mirrors model_switch.py's - # _save_discovered_models_to_config. A failed save is non-fatal. - if live_models: - try: - from hermes_cli.model_switch import _save_discovered_models_to_config - - _save_discovered_models_to_config(base_url, live_models, api_mode=api_mode, headers=extra_headers or None) - except Exception: - pass - return models, native_catalog_empty - - -def _pick_named_custom_model(name: str, models: list, saved_model: str): - """Searchable radiolist over *models* (numbered prompt without curses); None = cancelled.""" - default_idx = models.index(saved_model) if saved_model and saved_model in models else 0 - print(f"Found {len(models)} model(s):\n") - try: - from hermes_cli.curses_ui import curses_radiolist - - menu_items = [f"{m} (current)" if m == saved_model else m for m in models] + ["Cancel"] - idx = curses_radiolist(f"Select model from {name}:", menu_items, selected=default_idx, cancel_returns=-1, searchable=True) - print() - except (ImportError, NotImplementedError, OSError, subprocess.SubprocessError): - for i, m in enumerate(models, 1): - print(f" {i}. {m}{' (current)' if m == saved_model else ''}") - _say(f" {len(models) + 1}. Cancel", "") - try: - val = input(f"Choice [1-{len(models) + 1}]: ").strip() - if not val: - print("Cancelled.") - return None - idx = int(val) - 1 - except (ValueError, KeyboardInterrupt, EOFError): - print("\nCancelled.") - return None - if idx < 0 or idx >= len(models): - print("Cancelled.") - return None - return models[idx] - - -def _model_flow_named_custom(config, provider_info): - """Handle a named custom provider from config.yaml custom_providers list. - - Probes the endpoint's model catalog (native ``/api/tags`` for endpoints - conservatively identified as Ollama); a previously saved model is pre-selected - and used as the fallback when probing fails. - """ - from hermes_cli.main import _custom_provider_api_key_config_value, _custom_provider_base_url_config_value, _save_custom_provider - from hermes_cli.auth import _save_model_choice - from hermes_cli.config import load_config, save_config - from hermes_cli.model_switch import _entry_models_discovered, _models_config_is_allowlist - - name = provider_info["name"] - base_url = provider_info["base_url"] - api_mode = provider_info.get("api_mode", "") - api_key = provider_info.get("api_key", "") - key_env = provider_info.get("key_env", "") - saved_model = provider_info.get("model", "") - provider_key = (provider_info.get("provider_key") or "").strip() - - # Resolve key from env var if api_key not set directly - if not api_key and key_env: - api_key = os.environ.get(key_env, "") - config_api_key = _custom_provider_api_key_config_value(provider_info, api_key) - - # ``discover_models: false`` (default True) uses the configured ``models:`` list - # verbatim and skips the live probe, so operators can restrict the picker to the - # subset their plan serves. Same semantics as the slash-command picker. - discover = provider_info.get("discover_models", True) - if isinstance(discover, str): - discover = discover.lower() not in {"false", "no", "0"} - cfg_models = provider_info.get("models", {}) - explicit_catalog = _models_config_is_allowlist(cfg_models, _entry_models_discovered(provider_info)) - configured_models = _configured_model_ids(cfg_models) - - print(f" Provider: {name}") - print(f" URL: {base_url}") - if saved_model: - print(f" Current: {saved_model}") - print() - - native_catalog_empty = False - if not discover: - # Never probe. The active model is a usable sole choice, not a catalog. - models = configured_models or ([saved_model] if saved_model else []) - print(f"Using configured models (discover_models: false): {len(models)}") - else: - models, native_catalog_empty = _discover_named_custom_models(provider_info, api_key, configured_models, explicit_catalog) - - if models: - model_name = _pick_named_custom_model(name, models, saved_model) - if model_name is None: - return - elif saved_model and not native_catalog_empty: - print("Could not fetch models from endpoint.") - model_name = _ask(f"Model name [{saved_model}]: ") - if model_name is None: - return - model_name = model_name or saved_model - else: - print("Could not fetch models from endpoint. Enter model name manually.") - model_name = _ask("Model name: ") - if model_name is None: - return - if not model_name: - print("No model specified. Cancelled.") - return - - # Activate and save the model to the custom_providers entry - _save_model_choice(model_name) - cfg, model = _load_config_model_section() - if provider_key: - model["provider"] = custom_provider_slug(name, provider_key) - model.pop("base_url", None) - model.pop("api_key", None) - else: - model["provider"] = "custom" - model["base_url"] = _custom_provider_base_url_config_value(provider_info, base_url) - if config_api_key: - model["api_key"] = config_api_key - # Apply api_mode from custom_providers entry, or clear stale value - if api_mode: - model["api_mode"] = api_mode - else: - model.pop("api_mode", None) # let runtime auto-detect from URL - _commit_model_config(cfg) - - # Persist the selected model back to whichever schema owns this endpoint. - if provider_key: - cfg = load_config() - providers_cfg = cfg.get("providers") - provider_entry = providers_cfg.get(provider_key) if isinstance(providers_cfg, dict) else None - if isinstance(provider_entry, dict): - provider_entry["default_model"] = model_name - # Only persist an inline api_key when the user originally had one - # (literal or ``${VAR}``). Entries relying on ``key_env`` must not get - # a synthesized api_key — the runtime resolves key_env directly and - # writing it would downgrade credential hygiene. - had_inline_api_key = bool( - str(provider_info.get("api_key_ref", "") or "").strip() or str(provider_info.get("api_key", "") or "").strip() - ) - if had_inline_api_key and config_api_key and not str(provider_entry.get("api_key", "") or "").strip(): - provider_entry["api_key"] = config_api_key - if key_env and not str(provider_entry.get("key_env", "") or "").strip(): - provider_entry["key_env"] = key_env - cfg["providers"] = providers_cfg - save_config(cfg) - else: - # Save model name to the custom_providers entry for next time - _save_custom_provider(base_url, config_api_key, model_name, api_mode=api_mode) - - print(f"\n✅ Model set to: {model_name}") - print(f" Provider: {name} ({base_url})") - - def _copilot_model_list(live_ids) -> list: """Live GitHub Copilot ids, or the curated fallback with a warning.""" from hermes_cli.models import _PROVIDER_MODELS @@ -1442,176 +796,6 @@ def _model_flow_stepfun(config, current_model=""): config["model"] = dict(model) -def _model_flow_bedrock_api_key(config, region, current_model=""): - """Bedrock API Key mode — uses the OpenAI-compatible bedrock-mantle endpoint. - - For developers without an AWS account who received a Bedrock API Key from - their AWS admin. Works like any OpenAI-compatible endpoint. - """ - from hermes_cli.auth import _resolve_api_key_provider_secret, ProviderConfig - from hermes_cli.config import save_env_value - from hermes_cli.models import _PROVIDER_MODELS - - mantle_base_url = f"https://bedrock-mantle.{region}.api.aws/v1" - - # Check env var and credential pool (keys added via `hermes auth`) - bedrock_pconfig = ProviderConfig(id="bedrock", name="Bedrock", auth_type="api_key", api_key_env_vars=("AWS_BEARER_TOKEN_BEDROCK",)) - existing_key, existing_source = _resolve_api_key_provider_secret("bedrock", bedrock_pconfig) - if existing_key: - from hermes_cli.env_loader import format_secret_source_suffix - - source_suffix = format_secret_source_suffix(existing_source or "AWS_BEARER_TOKEN_BEDROCK") - print(f" Bedrock API Key: {existing_key[:12]}... ✓{source_suffix}") - else: - _say(f" Endpoint: {mantle_base_url}", "") - api_key = _ask(" Bedrock API Key: ", secret=True, cancel_msg="") - if api_key is None: - return - if not api_key: - print(" Cancelled.") - return - save_env_value("AWS_BEARER_TOKEN_BEDROCK", api_key) - existing_key = api_key - print(" ✓ API key saved.") - print() - - # Static list — mantle doesn't need boto3 for discovery - model_list = _PROVIDER_MODELS.get("bedrock", []) - print(f" Showing {len(model_list)} curated models") - selected = _pick_model_or_prompt( - model_list, " Model ID: ", current_model=current_model, confirm_provider="custom", - confirm_base_url=mantle_base_url, confirm_api_key=existing_key, - ) - - def _finish(cfg, _model): - # The bearer token rides on a named provider entry: a bare ``provider: custom`` - # cannot carry a credential for this host because OPENAI_API_KEY is gated to - # openai.com, so requests would go out as "no-key-required". - providers = _ensure_dict_section(cfg, "providers") - mantle_entry = providers.get("bedrock-mantle") - if not isinstance(mantle_entry, dict): - mantle_entry = {} - mantle_entry["base_url"] = mantle_base_url - mantle_entry["key_env"] = "AWS_BEARER_TOKEN_BEDROCK" - providers["bedrock-mantle"] = mantle_entry - # Also save region in bedrock config for reference - _ensure_dict_section(cfg, "bedrock")["region"] = region - - # Saved as a custom provider pointing to bedrock-mantle (no inline endpoint fields). - if _finish_model(selected, "custom:bedrock-mantle", f" Default model set to: {selected} (via Bedrock API Key, {region})", - no_change=" No change.", drop_base_url=True, drop_api_mode=True, finish=_finish) is not None: - print(f" Endpoint: {mantle_base_url}") - - -_BEDROCK_EXCLUDE_PREFIXES = ("stability.", "cohere.embed", "twelvelabs.", "us.stability.", "us.cohere.embed", - "us.twelvelabs.", "global.cohere.embed", "global.twelvelabs.") -_BEDROCK_EXCLUDE_SUBSTRINGS = ("safeguard", "voxtral", "palmyra-vision") -_BEDROCK_PROFILE_PREFIXES = BEDROCK_GEO_PREFIXES + ("global.",) -# Recommended models, matched geo-agnostically so an EU (eu.*) or APAC (apac.*) -# picker pins its own region's profile rather than a us.* one. -_BEDROCK_RECOMMENDED_BASES = ( - "anthropic.claude-sonnet-4-6", "anthropic.claude-opus-4-6", "anthropic.claude-haiku-4-5", "amazon.nova-pro", - "amazon.nova-lite", "amazon.nova-micro", "deepseek.v3", "meta.llama4-maverick", "meta.llama4-scout", -) - - -def _bedrock_text_model_ids(live_models: list, region: str) -> list[str]: - """Filter live Bedrock models to routable text models, dedupe bare ids against their - inference profiles, and order: recommended (in-region profile before global.*), - then other global.* profiles, then the rest.""" - def _base_id(mid: str) -> str: - _pp = next((p for p in _BEDROCK_PROFILE_PREFIXES if mid.startswith(p)), None) - return mid[len(_pp):] if _pp else mid - - filtered = [ - m for m in live_models - if not any(m["id"].startswith(p) for p in _BEDROCK_EXCLUDE_PREFIXES) - and not any(s in m["id"].lower() for s in _BEDROCK_EXCLUDE_SUBSTRINGS) - and bedrock_model_routable_from_region(m["id"], region) - ] - # Deduplicate: prefer inference profiles (geo-prefixed or global.*) over bare foundation model IDs. - profile_base_ids = {_base_id(m["id"]) for m in filtered if m["id"].startswith(_BEDROCK_PROFILE_PREFIXES)} - deduped = [m for m in filtered if m["id"].startswith(_BEDROCK_PROFILE_PREFIXES) or m["id"] not in profile_base_ids] - - def _sort_key(m): - mid = m["id"] - base = _base_id(mid) - for i, rec in enumerate(_BEDROCK_RECOMMENDED_BASES): - if base.startswith(rec): - # In-region geo profile beats global.* for the same model - return (0, i, 0 if not mid.startswith("global.") else 1, mid) - if mid.startswith("global."): - return (1, 0, 0, mid) - return (2, 0, 0, mid) - - deduped.sort(key=_sort_key) - return [m["id"] for m in deduped] - - -def _model_flow_bedrock(config, current_model=""): - """AWS Bedrock provider: verify credentials, pick region, discover models. - - Uses the native Converse API via boto3 — not the OpenAI-compatible endpoint. - Auth is the AWS SDK default credential chain (env vars, profile, instance - role), so no API key prompt is needed. - """ - from hermes_cli.models import _PROVIDER_MODELS - - # 1. Check for AWS credentials - try: - from agent.bedrock_adapter import has_aws_credentials, resolve_aws_auth_env_var, resolve_bedrock_region, discover_bedrock_models - except ImportError: - _say(" ✗ boto3 is not installed. Install it with:", " pip install boto3", "") - return - - if not has_aws_credentials(): - _say(" ⚠ No AWS credentials detected via environment variables.", - " Bedrock will use boto3's default credential chain (IMDS, SSO, etc.)", "") - auth_var = resolve_aws_auth_env_var() - print(f" AWS credentials: {auth_var} ✓" if auth_var else " AWS credentials: boto3 default chain (instance role / SSO)") - print() - - # 2. Region selection - current_region = resolve_bedrock_region() - region_input = _ask(f" AWS Region [{current_region}]: ", cancel_msg="") - if region_input is None: - return - region = region_input or current_region - - # 2b. Authentication mode - _say(" Choose authentication method:", "", " 1. IAM credential chain (recommended)", - " Works with EC2 instance roles, SSO, env vars, aws configure", " 2. Bedrock API Key", - " Enter your Bedrock API Key directly — also supports", - " team scenarios where an admin distributes keys", "") - auth_choice = _ask(" Choice [1]: ", raw=True, cancel_msg="") - if auth_choice is None: - return - if auth_choice == "2": - _model_flow_bedrock_api_key(config, region, current_model) - return - - # 3. Model discovery — try live API first, fall back to static list - print(f" Discovering models in {region}...") - live_models = discover_bedrock_models(region) - if live_models: - model_list = _bedrock_text_model_ids(live_models, region) - print(f" Found {len(model_list)} text model(s) (filtered from {len(live_models)} total)") - else: - model_list = _PROVIDER_MODELS.get("bedrock", []) - if not model_list: - print(" No models found. Check IAM permissions for bedrock:ListFoundationModels.") - return - print(f" Using {len(model_list)} curated models (live discovery unavailable)") - - # 4. Model selection - runtime_url = f"https://bedrock-runtime.{region}.amazonaws.com" - selected = _pick_model_or_prompt(model_list, " Model ID: ", current_model=current_model, confirm_provider="bedrock", confirm_base_url=runtime_url) - # api_mode is dropped: bedrock_converse is auto-detected. - _finish_model(selected, "bedrock", f" Default model set to: {selected} (via AWS Bedrock, {region})", no_change=" No change.", - base_url=runtime_url, drop_api_mode=True, - finish=lambda cfg, _m: _ensure_dict_section(cfg, "bedrock").__setitem__("region", region)) - - def _model_flow_vertex(config, current_model=""): """Google Vertex AI provider: Gemini via the OpenAI-compatible endpoint. diff --git a/hermes_cli/model_setup_flows_azure.py b/hermes_cli/model_setup_flows_azure.py new file mode 100644 index 0000000000..7e138c2a0a --- /dev/null +++ b/hermes_cli/model_setup_flows_azure.py @@ -0,0 +1,255 @@ +"""Azure Foundry wizard (OpenAI-style or Anthropic-style transport, API-key or Entra ID auth). + +Imports of hermes_cli.config / azure_detect stay lazy (tests patch them at call time). +Prompt strings and config write order are behavior. +""" + +from __future__ import annotations + +from hermes_cli.config import clear_model_endpoint_credentials +from hermes_cli.model_setup_flows_common import _HTTP, _ask, _commit_model_config, _load_config_model_section, _say + + +def _azure_mode_label(mode: str) -> str: + return "OpenAI-style" if mode == "chat_completions" else "Anthropic-style" + + +def _azure_entra_preflight(current_entra: dict): + """Entra ID credential preflight for the Azure flow. Returns + ``(token_provider, entra_overrides)``; ``None`` when the user cancelled; + ``False`` when the adapter is missing (caller falls back to API-key auth).""" + try: + from agent.azure_identity_adapter import ( + EntraIdentityConfig, SCOPE_AI_AZURE_DEFAULT, build_token_provider, describe_active_credential, + has_azure_identity_installed, + ) + except ImportError as exc: + _say("", f"⚠ Could not import azure-identity adapter: {exc}", " Falling back to API key auth.") + return False + + print() + if not has_azure_identity_installed(): + _say("◐ The 'azure-identity' package is not installed yet.", + " Hermes will install it now (the preflight below triggers the lazy-install). " + "To skip lazy installs, run: pip install azure-identity") + + # Only the optional scope override is persisted; identity selection (tenant, + # user-assigned MI, workload identity, SP) stays in AZURE_* SDK env vars. + entra_overrides: dict = {} + _persisted_scope_override = str(current_entra.get("scope") or "").strip() + entra_scope = _persisted_scope_override or SCOPE_AI_AZURE_DEFAULT + if _persisted_scope_override: + entra_overrides["scope"] = _persisted_scope_override + + _say("", "◐ Probing Microsoft Entra ID credential chain (up to 10s)...") + _config = EntraIdentityConfig(scope=entra_scope) + info = describe_active_credential(config=_config, timeout_seconds=10.0) + if info.get("ok"): + env_sources = info.get("env_sources") or [] + tag = ", ".join(env_sources) if env_sources else "default chain" + print(f"✓ Entra ID token acquired ({tag}, scope={entra_scope})") + else: + err = info.get("error") or "credential chain exhausted" + hint = info.get("hint") or ( + "Run `az login`, attach a managed identity to this VM, or set AZURE_TENANT_ID/AZURE_CLIENT_ID/AZURE_CLIENT_SECRET." + ) + _say(f"⚠ {err}", f" Hint: {hint}") + ans = _ask("Save Entra config anyway and validate later? [Y/n]: ", raw=True) + if ans is None: + return None + if ans.lower() not in ("", "y", "yes"): + print("Cancelled.") + return None + + # Best-effort token provider for the detection probe; on failure the probe + # falls back to manual entry. + try: + token_provider = build_token_provider(config=_config) + except Exception as exc: + print(f"⚠ Could not build token provider for probing: {exc}") + token_provider = None + return token_provider, entra_overrides + + +def _azure_pick_model(discovered_models: list, current_model: str): + """Model/deployment step of the Azure flow; None when cancelled.""" + if not discovered_models: + model_name = _ask(f"Model / deployment name [{current_model or 'e.g. gpt-5.4, claude-sonnet-4-6'}]: ") + return None if model_name is None else (model_name or current_model) + print("Available models on this endpoint:") + for i, mid in enumerate(discovered_models[:30], start=1): + print(f" {i:>2}. {mid}") + if len(discovered_models) > 30: + print(f" ... and {len(discovered_models) - 30} more (type name manually if not shown)") + print() + pick = _ask(f"Pick by number, or type a deployment name [{current_model or discovered_models[0]}]: ", raw=True) + if pick is None: + return None + if not pick: + return current_model or discovered_models[0] + if pick.isdigit() and 1 <= int(pick) <= min(len(discovered_models), 30): + return discovered_models[int(pick) - 1] + return pick + + +def _model_flow_azure_foundry(config, current_model=""): + """Azure Foundry provider: configure endpoint, auth mode, API mode, and model. + + Two transports (OpenAI-style ``/v1/chat/completions``, Anthropic-style + ``/v1/messages``) and two auth modes: **API key** (``AZURE_FOUNDRY_API_KEY``) or + **Microsoft Entra ID** (keyless RBAC via ``azure-identity``; the same ``Azure AI + User`` role covers both transports — OpenAI SDK takes a callable ``api_key``, + Anthropic gets a bearer-injecting ``httpx.Client`` from + :func:`agent.azure_identity_adapter.build_bearer_http_client`). + + Detection order: ``/anthropic`` URL suffix → Anthropic; ``GET /models`` + success → OpenAI-style + model picker; Anthropic Messages probe; manual entry. + Context length resolves via :func:`agent.model_metadata.get_model_context_length`. + """ + from hermes_cli.config import get_env_value, save_env_value + from hermes_cli import azure_detect + + # ── Load current Azure Foundry configuration ───────────────────── + model_cfg = config.get("model", {}) + current_base_url = current_api_mode = "" + current_auth_mode, current_entra = "api_key", {} + if isinstance(model_cfg, dict) and model_cfg.get("provider") == "azure-foundry": + current_base_url = str(model_cfg.get("base_url", "") or "") + current_api_mode = str(model_cfg.get("api_mode", "") or "") + current_auth_mode = str(model_cfg.get("auth_mode") or "api_key").strip().lower() or "api_key" + _cur_entra = model_cfg.get("entra") or {} + current_entra = _cur_entra if isinstance(_cur_entra, dict) else {} + current_api_key = get_env_value("AZURE_FOUNDRY_API_KEY") or "" + + _say("", "Azure Foundry Configuration", "=" * 50, "", + "Azure Foundry can host models with either OpenAI-style or", + "Anthropic-style API endpoints. Hermes will probe your", + "endpoint to auto-detect the transport and the deployed", + "models when possible.", "") + if current_base_url: + print(f" Current endpoint: {current_base_url}") + if current_api_mode: + print(f" Current API mode: {_azure_mode_label(current_api_mode)}") + if current_auth_mode == "entra_id": + print(" Current auth mode: Microsoft Entra ID (keyless)") + elif current_api_key: + print(f" Current auth mode: API key ({current_api_key[:8]}...)") + print() + + # ── Step 1: endpoint URL ───────────────────────────────────────── + _placeholder = current_base_url or ( + "e.g. https://.openai.azure.com/openai/v1 or https://.services.ai.azure.com/anthropic" + ) + base_url = _ask(f"API endpoint URL [{_placeholder}]: ") + if base_url is None: + return + effective_url = (base_url or current_base_url).rstrip("/") + if not effective_url: + print("No endpoint URL provided. Cancelled.") + return + if not effective_url.startswith(_HTTP): + print(f"Invalid URL: {effective_url} (must start with http:// or https://)") + return + + # ── Step 2: authentication mode ────────────────────────────────── + _say("", "Authentication:", " 1. API key (AZURE_FOUNDRY_API_KEY in .env)", + " 2. Microsoft Entra ID (managed identity / workload identity / az login)", + " Recommended by Microsoft. Works for both OpenAI-style and Anthropic-style endpoints.", + " Requires the 'Azure AI User' role on the Foundry resource.") + _auth_default = "2" if current_auth_mode == "entra_id" else "1" + auth_choice = _ask(f"Authentication mode [1/2] ({_auth_default}): ", raw=True) + if auth_choice is None: + return + use_entra = (auth_choice or _auth_default) == "2" + + # ── Step 3: credentials (key OR Entra preflight) ───────────────── + effective_key: str = "" + entra_overrides: dict = {} + token_provider = None # callable when entra + if use_entra: + preflight = _azure_entra_preflight(current_entra) + if preflight is None: + return + if preflight is False: + use_entra = False + else: + token_provider, entra_overrides = preflight + if not use_entra: + print() + api_key = _ask(f"API key [{current_api_key[:8] + '...' if current_api_key else 'required'}]: ", secret=True) + if api_key is None: + return + effective_key = api_key or current_api_key + if not effective_key: + print("No API key provided. Cancelled.") + return + + # ── Step 4: auto-detect transport + models ─────────────────────── + _say("", "◐ Probing endpoint to auto-detect transport and models...") + detection = azure_detect.detect(effective_url, api_key=effective_key, token_provider=token_provider) + discovered_models: list[str] = list(detection.models) + api_mode: str = detection.api_mode or "" + if api_mode: + print(f"✓ Detected API transport: {_azure_mode_label(api_mode)}") + if detection.reason: + print(f" ({detection.reason})") + if discovered_models: + print(f"✓ Found {len(discovered_models)} deployed model(s) on this endpoint") + else: + _say(f"⚠ Auto-detection incomplete: {detection.reason}", "", + "Select the API format your Azure Foundry endpoint uses:", + " 1. OpenAI-style (POST /v1/chat/completions)", + " For: GPT models, Llama, Mistral, and most open models", + " 2. Anthropic-style (POST /v1/messages)", + " For: Claude models deployed via Anthropic API format") + default_choice = "2" if current_api_mode == "anthropic_messages" else "1" + mode_choice = _ask(f"API format [1/2] ({default_choice}): ", raw=True) + if mode_choice is None: + return + api_mode = "anthropic_messages" if (mode_choice or default_choice) == "2" else "chat_completions" + + # ── Step 5: model name ─────────────────────────────────────────── + print() + effective_model = _azure_pick_model(discovered_models, current_model) + if effective_model is None: + return + if not effective_model: + print("No model name provided. Cancelled.") + return + + # ── Step 6: context-length lookup ──────────────────────────────── + ctx_len = azure_detect.lookup_context_length(effective_model, effective_url, api_key=effective_key, token_provider=token_provider) + + # ── Step 7: persist ────────────────────────────────────────────── + if not use_entra: + save_env_value("AZURE_FOUNDRY_API_KEY", effective_key) + cfg, model = _load_config_model_section() + model["provider"] = "azure-foundry" + model["base_url"] = effective_url + model["api_mode"] = api_mode + model["default"] = effective_model + model["auth_mode"] = "entra_id" if use_entra else "api_key" + clear_model_endpoint_credentials(model, clear_api_mode=False) + # Persist only a non-default Entra scope so config.yaml stays tidy. + clean_entra = {k: v for k in ("scope",) if (v := entra_overrides.get(k))} + if use_entra and clean_entra: + model["entra"] = clean_entra + else: + model.pop("entra", None) + if ctx_len: + model["context_length"] = ctx_len + _commit_model_config(cfg) + config["model"] = dict(model) + + # Clear conflicting env vars so auxiliary clients don't pick up a stale + # OpenAI base URL / key. + for var in ("OPENAI_BASE_URL", "OPENAI_API_KEY"): + if get_env_value(var): + save_env_value(var, "") + + _say("", "✓ Azure Foundry configured:", f" Endpoint: {effective_url}", + f" API mode: {_azure_mode_label(api_mode)}", + f" Auth: {'Microsoft Entra ID (keyless)' if use_entra else 'API key'}", + f" Model: {effective_model}", + f" Context length: {ctx_len:,} tokens" if ctx_len else " Context length: not auto-detected (will fall back at runtime)", + "") diff --git a/hermes_cli/model_setup_flows_bedrock.py b/hermes_cli/model_setup_flows_bedrock.py new file mode 100644 index 0000000000..201aaf9095 --- /dev/null +++ b/hermes_cli/model_setup_flows_bedrock.py @@ -0,0 +1,224 @@ +"""AWS Bedrock wizards: native Converse API (IAM chain, region-scoped model discovery) +and the Bedrock API Key mode on the OpenAI-compatible bedrock-mantle endpoint. + +Imports of hermes_cli.auth / config / models stay lazy (tests patch them at call time). +Prompt strings and config write order are behavior. +""" + +from __future__ import annotations + +from hermes_cli.model_setup_flows_common import ( + _ask, _ensure_dict_section, _finish_model, _pick_model_or_prompt, _say, +) + + +# AWS cross-region inference profile prefixes. A geo-prefixed profile only routes +# from endpoints in its own geography (us.* from eu-central-2 is rejected by AWS +# regardless of credentials); global.* routes from everywhere. +BEDROCK_GEO_PREFIXES = ("us.", "eu.", "ap.", "apac.", "jp.", "ca.", "sa.", "me.", "af.") + + +# region-name prefixes -> inference-profile geo prefix +_REGION_GEO = (("us.", ("us-", "us_gov")), ("eu.", ("eu-",)), ("ap.", ("ap-",)), ("ca.", ("ca-",)), + ("sa.", ("sa-",)), ("me.", ("me-",)), ("af.", ("af-",))) + + +def bedrock_region_geo_prefix(region_name: str) -> str: + """Map an AWS region name to its inference-profile geo prefix ('' = unknown).""" + r = (region_name or "").lower() + return next((geo for geo, prefixes in _REGION_GEO if r.startswith(prefixes)), "") + + +def bedrock_model_routable_from_region(model_id: str, region_name: str) -> bool: + """True when *model_id* can be invoked from *region_name*'s endpoint. + + Bare foundation-model ids and ``global.*`` profiles route from anywhere; + geo-prefixed profiles only from their own geography. Unknown regions hide nothing. + """ + mid = (model_id or "").lower() + matched_geo = next((p for p in BEDROCK_GEO_PREFIXES if mid.startswith(p)), None) + if matched_geo is None or mid.startswith("global."): + return True + geo = bedrock_region_geo_prefix(region_name) + if not geo: + return True + if geo == "ap.": + # Asia-Pacific regions can carry ap./apac./jp. profile spellings. + return matched_geo in ("ap.", "apac.", "jp.") + return matched_geo == geo + + +def _model_flow_bedrock_api_key(config, region, current_model=""): + """Bedrock API Key mode — uses the OpenAI-compatible bedrock-mantle endpoint. + + For developers without an AWS account who received a Bedrock API Key from + their AWS admin. Works like any OpenAI-compatible endpoint. + """ + from hermes_cli.auth import _resolve_api_key_provider_secret, ProviderConfig + from hermes_cli.config import save_env_value + from hermes_cli.models import _PROVIDER_MODELS + + mantle_base_url = f"https://bedrock-mantle.{region}.api.aws/v1" + + # Check env var and credential pool (keys added via `hermes auth`) + bedrock_pconfig = ProviderConfig(id="bedrock", name="Bedrock", auth_type="api_key", api_key_env_vars=("AWS_BEARER_TOKEN_BEDROCK",)) + existing_key, existing_source = _resolve_api_key_provider_secret("bedrock", bedrock_pconfig) + if existing_key: + from hermes_cli.env_loader import format_secret_source_suffix + + source_suffix = format_secret_source_suffix(existing_source or "AWS_BEARER_TOKEN_BEDROCK") + print(f" Bedrock API Key: {existing_key[:12]}... ✓{source_suffix}") + else: + _say(f" Endpoint: {mantle_base_url}", "") + api_key = _ask(" Bedrock API Key: ", secret=True, cancel_msg="") + if api_key is None: + return + if not api_key: + print(" Cancelled.") + return + save_env_value("AWS_BEARER_TOKEN_BEDROCK", api_key) + existing_key = api_key + print(" ✓ API key saved.") + print() + + # Static list — mantle doesn't need boto3 for discovery + model_list = _PROVIDER_MODELS.get("bedrock", []) + print(f" Showing {len(model_list)} curated models") + selected = _pick_model_or_prompt( + model_list, " Model ID: ", current_model=current_model, confirm_provider="custom", + confirm_base_url=mantle_base_url, confirm_api_key=existing_key, + ) + + def _finish(cfg, _model): + # The bearer token rides on a named provider entry: a bare ``provider: custom`` + # cannot carry a credential for this host because OPENAI_API_KEY is gated to + # openai.com, so requests would go out as "no-key-required". + providers = _ensure_dict_section(cfg, "providers") + mantle_entry = providers.get("bedrock-mantle") + if not isinstance(mantle_entry, dict): + mantle_entry = {} + mantle_entry["base_url"] = mantle_base_url + mantle_entry["key_env"] = "AWS_BEARER_TOKEN_BEDROCK" + providers["bedrock-mantle"] = mantle_entry + # Also save region in bedrock config for reference + _ensure_dict_section(cfg, "bedrock")["region"] = region + + # Saved as a custom provider pointing to bedrock-mantle (no inline endpoint fields). + if _finish_model(selected, "custom:bedrock-mantle", f" Default model set to: {selected} (via Bedrock API Key, {region})", + no_change=" No change.", drop_base_url=True, drop_api_mode=True, finish=_finish) is not None: + print(f" Endpoint: {mantle_base_url}") + + +_BEDROCK_EXCLUDE_PREFIXES = ("stability.", "cohere.embed", "twelvelabs.", "us.stability.", "us.cohere.embed", + "us.twelvelabs.", "global.cohere.embed", "global.twelvelabs.") + + +_BEDROCK_EXCLUDE_SUBSTRINGS = ("safeguard", "voxtral", "palmyra-vision") + + +_BEDROCK_PROFILE_PREFIXES = BEDROCK_GEO_PREFIXES + ("global.",) + + +# Recommended models, matched geo-agnostically so an EU (eu.*) or APAC (apac.*) +# picker pins its own region's profile rather than a us.* one. +_BEDROCK_RECOMMENDED_BASES = ( + "anthropic.claude-sonnet-4-6", "anthropic.claude-opus-4-6", "anthropic.claude-haiku-4-5", "amazon.nova-pro", + "amazon.nova-lite", "amazon.nova-micro", "deepseek.v3", "meta.llama4-maverick", "meta.llama4-scout", +) + + +def _bedrock_text_model_ids(live_models: list, region: str) -> list[str]: + """Filter live Bedrock models to routable text models, dedupe bare ids against their + inference profiles, and order: recommended (in-region profile before global.*), + then other global.* profiles, then the rest.""" + def _base_id(mid: str) -> str: + _pp = next((p for p in _BEDROCK_PROFILE_PREFIXES if mid.startswith(p)), None) + return mid[len(_pp):] if _pp else mid + + filtered = [ + m for m in live_models + if not any(m["id"].startswith(p) for p in _BEDROCK_EXCLUDE_PREFIXES) + and not any(s in m["id"].lower() for s in _BEDROCK_EXCLUDE_SUBSTRINGS) + and bedrock_model_routable_from_region(m["id"], region) + ] + # Deduplicate: prefer inference profiles (geo-prefixed or global.*) over bare foundation model IDs. + profile_base_ids = {_base_id(m["id"]) for m in filtered if m["id"].startswith(_BEDROCK_PROFILE_PREFIXES)} + deduped = [m for m in filtered if m["id"].startswith(_BEDROCK_PROFILE_PREFIXES) or m["id"] not in profile_base_ids] + + def _sort_key(m): + mid = m["id"] + base = _base_id(mid) + for i, rec in enumerate(_BEDROCK_RECOMMENDED_BASES): + if base.startswith(rec): + # In-region geo profile beats global.* for the same model + return (0, i, 0 if not mid.startswith("global.") else 1, mid) + if mid.startswith("global."): + return (1, 0, 0, mid) + return (2, 0, 0, mid) + + deduped.sort(key=_sort_key) + return [m["id"] for m in deduped] + + +def _model_flow_bedrock(config, current_model=""): + """AWS Bedrock provider: verify credentials, pick region, discover models. + + Uses the native Converse API via boto3 — not the OpenAI-compatible endpoint. + Auth is the AWS SDK default credential chain (env vars, profile, instance + role), so no API key prompt is needed. + """ + from hermes_cli.models import _PROVIDER_MODELS + + # 1. Check for AWS credentials + try: + from agent.bedrock_adapter import has_aws_credentials, resolve_aws_auth_env_var, resolve_bedrock_region, discover_bedrock_models + except ImportError: + _say(" ✗ boto3 is not installed. Install it with:", " pip install boto3", "") + return + + if not has_aws_credentials(): + _say(" ⚠ No AWS credentials detected via environment variables.", + " Bedrock will use boto3's default credential chain (IMDS, SSO, etc.)", "") + auth_var = resolve_aws_auth_env_var() + print(f" AWS credentials: {auth_var} ✓" if auth_var else " AWS credentials: boto3 default chain (instance role / SSO)") + print() + + # 2. Region selection + current_region = resolve_bedrock_region() + region_input = _ask(f" AWS Region [{current_region}]: ", cancel_msg="") + if region_input is None: + return + region = region_input or current_region + + # 2b. Authentication mode + _say(" Choose authentication method:", "", " 1. IAM credential chain (recommended)", + " Works with EC2 instance roles, SSO, env vars, aws configure", " 2. Bedrock API Key", + " Enter your Bedrock API Key directly — also supports", + " team scenarios where an admin distributes keys", "") + auth_choice = _ask(" Choice [1]: ", raw=True, cancel_msg="") + if auth_choice is None: + return + if auth_choice == "2": + _model_flow_bedrock_api_key(config, region, current_model) + return + + # 3. Model discovery — try live API first, fall back to static list + print(f" Discovering models in {region}...") + live_models = discover_bedrock_models(region) + if live_models: + model_list = _bedrock_text_model_ids(live_models, region) + print(f" Found {len(model_list)} text model(s) (filtered from {len(live_models)} total)") + else: + model_list = _PROVIDER_MODELS.get("bedrock", []) + if not model_list: + print(" No models found. Check IAM permissions for bedrock:ListFoundationModels.") + return + print(f" Using {len(model_list)} curated models (live discovery unavailable)") + + # 4. Model selection + runtime_url = f"https://bedrock-runtime.{region}.amazonaws.com" + selected = _pick_model_or_prompt(model_list, " Model ID: ", current_model=current_model, confirm_provider="bedrock", confirm_base_url=runtime_url) + # api_mode is dropped: bedrock_converse is auto-detected. + _finish_model(selected, "bedrock", f" Default model set to: {selected} (via AWS Bedrock, {region})", no_change=" No change.", + base_url=runtime_url, drop_api_mode=True, + finish=lambda cfg, _m: _ensure_dict_section(cfg, "bedrock").__setitem__("region", region)) diff --git a/hermes_cli/model_setup_flows_common.py b/hermes_cli/model_setup_flows_common.py index 6b5c94753a..7accc4bcdd 100644 --- a/hermes_cli/model_setup_flows_common.py +++ b/hermes_cli/model_setup_flows_common.py @@ -14,6 +14,8 @@ from __future__ import annotations from hermes_cli.cli_output import line_input from hermes_cli.config import clear_model_endpoint_credentials +_HTTP = ("http://", "https://") + def _say(*lines: str) -> None: """``print`` each line (``""`` = blank line); one call per banner block.""" diff --git a/hermes_cli/model_setup_flows_custom.py b/hermes_cli/model_setup_flows_custom.py new file mode 100644 index 0000000000..b92c117a25 --- /dev/null +++ b/hermes_cli/model_setup_flows_custom.py @@ -0,0 +1,409 @@ +"""Custom OpenAI-compatible endpoint wizards: the ad-hoc ``custom`` flow and the +``custom_providers`` / ``providers.`` named-endpoint flow. + +Imports of hermes_cli.main / auth / config / models stay lazy (main.py import cycle; +tests patch them at call time). Prompt strings and config write order are behavior. +""" + +from __future__ import annotations + +import os +import subprocess +import urllib.parse + +from hermes_cli.cli_output import line_input +from hermes_cli.providers import custom_provider_slug +from hermes_cli.model_setup_flows_common import ( + _HTTP, _ask, _commit_model_config, _load_config_model_section, + _prune_replaced_custom_model_config_credentials, _say, +) + + +def _parse_context_length(text: str): + """``128k`` / ``128,000`` -> int; None when blank, non-positive, or unparsable (warns).""" + if not text: + return None + try: + value = int(text.replace(",", "").replace("k", "000").replace("K", "000")) + except ValueError: + print(f"Invalid context length: {text} — will auto-detect.") + return None + return value if value > 0 else None + + +def _probe_custom_endpoint(effective_key: str, effective_url: str) -> tuple[dict, str]: + """Verify a custom endpoint via ``probe_api_models`` and report; returns + ``(probe, effective_url)`` where the URL may be the working fallback base.""" + from hermes_cli.models import probe_api_models + + probe = probe_api_models(effective_key, effective_url) + if probe.get("used_fallback") and probe.get("resolved_base_url"): + print(f"Warning: endpoint verification worked at {probe['resolved_base_url']}/models, " + f"not the exact URL you entered. Saving the working base URL instead.") + effective_url = probe["resolved_base_url"] + elif probe.get("models") is not None: + print(f"Verified endpoint via {probe.get('probed_url')} ({len(probe.get('models') or [])} model(s) visible)") + else: + print(f"Warning: could not verify this endpoint via {probe.get('probed_url')}. Hermes will still save it.") + suggested = probe.get("suggested_base_url") + if suggested and suggested.endswith("/v1"): + print(f" If this server expects /v1 in the path, try base URL: {suggested}") + elif suggested: + print(f" If /v1 should not be in the base URL, try: {suggested}") + return probe, effective_url + + +def _pick_detected_model(detected_models: list) -> str: + """Model-name step of the custom flow: confirm a single detection, number-pick from + several, or type one. Raises KeyboardInterrupt/EOFError like the prompts it wraps.""" + manual = "Model name (e.g. gpt-4, llama-3-70b): " + if len(detected_models) == 1: + print(f" Detected model: {detected_models[0]}") + if input(" Use this model? [Y/n]: ").strip().lower() in {"", "y", "yes"}: + return detected_models[0] + return line_input(manual).strip() + if len(detected_models) > 1: + print(" Available models:") + for i, m in enumerate(detected_models, 1): + print(f" {i}. {m}") + pick = input(f" Select model [1-{len(detected_models)}] or type name: ").strip() + if pick.isdigit() and 1 <= int(pick) <= len(detected_models): + return detected_models[int(pick) - 1] + return pick + return line_input(manual).strip() + + +def _model_flow_custom(config): + """Custom endpoint: collect URL, API key, and model name. + + Also saves the endpoint to ``custom_providers`` in config.yaml so it appears + in the provider menu on subsequent runs. + """ + from hermes_cli.main import _auto_provider_name, _prompt_custom_api_mode_selection, _save_custom_provider + from hermes_cli.auth import _save_model_choice, deactivate_provider + from hermes_cli.config import custom_endpoint_key_env, get_env_value, save_env_value + from hermes_cli.secret_prompt import masked_secret_prompt + + current_url = get_env_value("OPENAI_BASE_URL") or "" + current_key = get_env_value("OPENAI_API_KEY") or "" + + print("Custom OpenAI-compatible endpoint configuration:") + if current_url: + print(f" Current URL: {current_url}") + if current_key: + print(f" Current key: {current_key[:8]}...") + print() + + try: + base_url = line_input(f"API base URL [{current_url or 'e.g. https://api.example.com/v1'}]: ").strip() + api_key = masked_secret_prompt(f"API key [{current_key[:8] + '...' if current_key else 'optional'}]: ").strip() + except (KeyboardInterrupt, EOFError): + print("\nCancelled.") + return + + if not base_url and not current_url: + print("No URL provided. Cancelled.") + return + effective_url = base_url or current_url + if not effective_url.startswith(_HTTP): + print(f"Invalid URL: {effective_url} (must start with http:// or https://)") + return + effective_key = api_key or current_key + + # Most local servers (Ollama, vLLM, llama.cpp) need /v1 for OpenAI-compatible + # chat completions — offer to append it when the URL looks local without it. + _url_lower = effective_url.rstrip("/").lower() + _looks_local = any(h in _url_lower for h in ("localhost", "127.0.0.1", "0.0.0.0", ":11434", ":8080", ":5000")) + if _looks_local and not _url_lower.endswith("/v1"): + _say("", " Hint: Did you mean to add /v1 at the end?", + " Most local model servers (Ollama, vLLM, llama.cpp) require it.", f" e.g. {effective_url.rstrip('/')}/v1") + if _ask(" Add /v1? [Y/n]: ", raw=True, cancel_msg=None, on_cancel="n").lower() in {"", "y", "yes"}: + effective_url = effective_url.rstrip("/") + "/v1" + print(f" Updated URL: {effective_url}") + print() + + probe, effective_url = _probe_custom_endpoint(effective_key, effective_url) + + # Ask for the API mode explicitly so codex-compatible custom providers don't + # silently fall back to chat_completions. + current_model_cfg = config.get("model") + current_api_mode = str(current_model_cfg.get("api_mode") or "").strip() if isinstance(current_model_cfg, dict) else "" + api_mode = _prompt_custom_api_mode_selection(effective_url, current_api_mode=current_api_mode) + print(f" API mode: {api_mode}" if api_mode else " API mode: auto-detect") + + # Select model — use probe results when available, fall back to manual input + try: + model_name = _pick_detected_model(probe.get("models") or []) + context_length_str = line_input("Context length in tokens [leave blank for auto-detect]: ").strip() + # Display name — shown in the provider menu on future runs + default_name = _auto_provider_name(effective_url) + display_name = line_input(f"Display name [{default_name}]: ").strip() or default_name + except (KeyboardInterrupt, EOFError): + print("\nCancelled.") + return + context_length = _parse_context_length(context_length_str) + + # The key goes to .env and config.yaml only references it. Keyed on host:port + # so two servers on one machine keep separate credentials. + custom_key_env = "" + if effective_key: + _parsed = urllib.parse.urlparse(effective_url) + _identity = _parsed.hostname or "" + if _parsed.port: + _identity = f"{_identity}_{_parsed.port}" + custom_key_env = custom_endpoint_key_env(_identity) + save_env_value(custom_key_env, effective_key) + print(f" API key saved to .env as {custom_key_env}") + + def _apply_endpoint(model: dict) -> None: + model["provider"] = "custom" + model["base_url"] = effective_url + if custom_key_env: + model["api_key"] = f"${{{custom_key_env}}}" + if api_mode: + model["api_mode"] = api_mode + else: + model.pop("api_mode", None) + + if model_name: + _save_model_choice(model_name) + cfg, model = _load_config_model_section() + _apply_endpoint(model) + _commit_model_config(cfg) + # Sync the caller's config dict so the setup wizard's final save_config(config) + # doesn't overwrite model.provider/base_url with its stale values. + config["model"] = dict(model) + print(f"Default model set to: {model_name} (via {effective_url})") + else: + if base_url or api_key: + deactivate_provider() + # Even without a model name, persist the endpoint on the caller's config dict. + _caller_model = config.get("model") + if not isinstance(_caller_model, dict): + _caller_model = {"default": _caller_model} if _caller_model else {} + _apply_endpoint(_caller_model) + config["model"] = _caller_model + print("Endpoint saved. Use `/model` in chat or `hermes model` to set a model.") + + # Auto-save to custom_providers so it appears in the menu next time + _save_custom_provider(effective_url, effective_key, model_name or "", context_length=context_length, + name=display_name, api_mode=api_mode, key_env=custom_key_env) + _prune_replaced_custom_model_config_credentials(effective_url, provider_name=display_name) + + +def _configured_model_ids(cfg_models) -> list[str]: + """Model ids from a ``custom_providers[].models`` mapping or list (marker keys skipped).""" + if isinstance(cfg_models, dict): + markers = {"__explicit_model_allowlist__", "__discovered_model_catalog__"} + return [str(m) for m in cfg_models if m not in markers and str(m).strip()] + out: list[str] = [] + if isinstance(cfg_models, list): + for entry in cfg_models: + if isinstance(entry, dict): + model_id = str(entry.get("id") or entry.get("model") or "").strip() + else: + model_id = str(entry).strip() if isinstance(entry, str) else "" + if model_id: + out.append(model_id) + return out + + +def _discover_named_custom_models(provider_info: dict, api_key: str, configured_models: list, explicit_catalog: bool): + """Live catalog probe for a named custom endpoint (native ``/api/tags`` for Ollama). + Returns ``(models, native_catalog_empty)``; persists the live catalog as a side effect.""" + from hermes_cli.config import normalize_extra_headers + from hermes_cli.models import ( + fetch_api_models, fetch_ollama_local_models, _get_ollama_native_headers, _normalize_openai_base_url, + should_use_ollama_native_catalog, + ) + + name, base_url = provider_info["name"], provider_info["base_url"] + api_mode = provider_info.get("api_mode", "") + provider_key = (provider_info.get("provider_key") or "").strip() + print("Fetching available models...") + fetch_kwargs = {"timeout": 8.0} + if api_mode: + fetch_kwargs["api_mode"] = api_mode + native_catalog_provider = "ollama" if provider_key.lower() == "ollama" or name.strip().lower() == "ollama" else "custom" + extra_headers = normalize_extra_headers(provider_info.get("extra_headers")) or {} + candidate_headers = _get_ollama_native_headers(base_url, api_key=api_key) + for key in tuple(candidate_headers): + if any(key.lower() == existing.lower() for existing in extra_headers): + del candidate_headers[key] + candidate_headers.update(extra_headers) + caller_has_authorization = any(key.lower() == "authorization" for key in extra_headers) + if api_key and not caller_has_authorization: + for key in tuple(candidate_headers): + if key.lower() == "authorization": + del candidate_headers[key] + candidate_headers["Authorization"] = f"Bearer {api_key}" + use_native = should_use_ollama_native_catalog(native_catalog_provider, base_url, headers=candidate_headers or None) + native_headers_arg = candidate_headers or None if use_native else (extra_headers or None) + native_catalog_empty = False + if use_native: + if explicit_catalog and configured_models: + live_models = configured_models + else: + live_models = fetch_ollama_local_models(base_url, timeout=8.0, headers=native_headers_arg) + native_catalog_empty = live_models == [] + if live_models is None: + live_models = fetch_api_models(api_key, _normalize_openai_base_url(base_url), headers=native_headers_arg, **fetch_kwargs) + native_catalog_empty = False + else: + live_models = fetch_api_models(api_key, base_url, headers=native_headers_arg, **fetch_kwargs) + models = configured_models if explicit_catalog else [] if native_catalog_empty else (live_models or configured_models) + # Persist the live catalog to the custom_providers entry so no-probe surfaces + # (dashboard, desktop, ACP) show the full list; mirrors model_switch.py's + # _save_discovered_models_to_config. A failed save is non-fatal. + if live_models: + try: + from hermes_cli.model_switch import _save_discovered_models_to_config + + _save_discovered_models_to_config(base_url, live_models, api_mode=api_mode, headers=extra_headers or None) + except Exception: + pass + return models, native_catalog_empty + + +def _pick_named_custom_model(name: str, models: list, saved_model: str): + """Searchable radiolist over *models* (numbered prompt without curses); None = cancelled.""" + default_idx = models.index(saved_model) if saved_model and saved_model in models else 0 + print(f"Found {len(models)} model(s):\n") + try: + from hermes_cli.curses_ui import curses_radiolist + + menu_items = [f"{m} (current)" if m == saved_model else m for m in models] + ["Cancel"] + idx = curses_radiolist(f"Select model from {name}:", menu_items, selected=default_idx, cancel_returns=-1, searchable=True) + print() + except (ImportError, NotImplementedError, OSError, subprocess.SubprocessError): + for i, m in enumerate(models, 1): + print(f" {i}. {m}{' (current)' if m == saved_model else ''}") + _say(f" {len(models) + 1}. Cancel", "") + try: + val = input(f"Choice [1-{len(models) + 1}]: ").strip() + if not val: + print("Cancelled.") + return None + idx = int(val) - 1 + except (ValueError, KeyboardInterrupt, EOFError): + print("\nCancelled.") + return None + if idx < 0 or idx >= len(models): + print("Cancelled.") + return None + return models[idx] + + +def _model_flow_named_custom(config, provider_info): + """Handle a named custom provider from config.yaml custom_providers list. + + Probes the endpoint's model catalog (native ``/api/tags`` for endpoints + conservatively identified as Ollama); a previously saved model is pre-selected + and used as the fallback when probing fails. + """ + from hermes_cli.main import _custom_provider_api_key_config_value, _custom_provider_base_url_config_value, _save_custom_provider + from hermes_cli.auth import _save_model_choice + from hermes_cli.config import load_config, save_config + from hermes_cli.model_switch import _entry_models_discovered, _models_config_is_allowlist + + name = provider_info["name"] + base_url = provider_info["base_url"] + api_mode = provider_info.get("api_mode", "") + api_key = provider_info.get("api_key", "") + key_env = provider_info.get("key_env", "") + saved_model = provider_info.get("model", "") + provider_key = (provider_info.get("provider_key") or "").strip() + + # Resolve key from env var if api_key not set directly + if not api_key and key_env: + api_key = os.environ.get(key_env, "") + config_api_key = _custom_provider_api_key_config_value(provider_info, api_key) + + # ``discover_models: false`` (default True) uses the configured ``models:`` list + # verbatim and skips the live probe, so operators can restrict the picker to the + # subset their plan serves. Same semantics as the slash-command picker. + discover = provider_info.get("discover_models", True) + if isinstance(discover, str): + discover = discover.lower() not in {"false", "no", "0"} + cfg_models = provider_info.get("models", {}) + explicit_catalog = _models_config_is_allowlist(cfg_models, _entry_models_discovered(provider_info)) + configured_models = _configured_model_ids(cfg_models) + + print(f" Provider: {name}") + print(f" URL: {base_url}") + if saved_model: + print(f" Current: {saved_model}") + print() + + native_catalog_empty = False + if not discover: + # Never probe. The active model is a usable sole choice, not a catalog. + models = configured_models or ([saved_model] if saved_model else []) + print(f"Using configured models (discover_models: false): {len(models)}") + else: + models, native_catalog_empty = _discover_named_custom_models(provider_info, api_key, configured_models, explicit_catalog) + + if models: + model_name = _pick_named_custom_model(name, models, saved_model) + if model_name is None: + return + elif saved_model and not native_catalog_empty: + print("Could not fetch models from endpoint.") + model_name = _ask(f"Model name [{saved_model}]: ") + if model_name is None: + return + model_name = model_name or saved_model + else: + print("Could not fetch models from endpoint. Enter model name manually.") + model_name = _ask("Model name: ") + if model_name is None: + return + if not model_name: + print("No model specified. Cancelled.") + return + + # Activate and save the model to the custom_providers entry + _save_model_choice(model_name) + cfg, model = _load_config_model_section() + if provider_key: + model["provider"] = custom_provider_slug(name, provider_key) + model.pop("base_url", None) + model.pop("api_key", None) + else: + model["provider"] = "custom" + model["base_url"] = _custom_provider_base_url_config_value(provider_info, base_url) + if config_api_key: + model["api_key"] = config_api_key + # Apply api_mode from custom_providers entry, or clear stale value + if api_mode: + model["api_mode"] = api_mode + else: + model.pop("api_mode", None) # let runtime auto-detect from URL + _commit_model_config(cfg) + + # Persist the selected model back to whichever schema owns this endpoint. + if provider_key: + cfg = load_config() + providers_cfg = cfg.get("providers") + provider_entry = providers_cfg.get(provider_key) if isinstance(providers_cfg, dict) else None + if isinstance(provider_entry, dict): + provider_entry["default_model"] = model_name + # Only persist an inline api_key when the user originally had one + # (literal or ``${VAR}``). Entries relying on ``key_env`` must not get + # a synthesized api_key — the runtime resolves key_env directly and + # writing it would downgrade credential hygiene. + had_inline_api_key = bool( + str(provider_info.get("api_key_ref", "") or "").strip() or str(provider_info.get("api_key", "") or "").strip() + ) + if had_inline_api_key and config_api_key and not str(provider_entry.get("api_key", "") or "").strip(): + provider_entry["api_key"] = config_api_key + if key_env and not str(provider_entry.get("key_env", "") or "").strip(): + provider_entry["key_env"] = key_env + cfg["providers"] = providers_cfg + save_config(cfg) + else: + # Save model name to the custom_providers entry for next time + _save_custom_provider(base_url, config_api_key, model_name, api_mode=api_mode) + + print(f"\n✅ Model set to: {model_name}") + print(f" Provider: {name} ({base_url})")