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.
This commit is contained in:
@@ -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 <base>/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://<resource>.openai.azure.com/openai/v1 or https://<resource>.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.
|
||||
|
||||
|
||||
255
hermes_cli/model_setup_flows_azure.py
Normal file
255
hermes_cli/model_setup_flows_azure.py
Normal file
@@ -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 <base>/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://<resource>.openai.azure.com/openai/v1 or https://<resource>.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)",
|
||||
"")
|
||||
224
hermes_cli/model_setup_flows_bedrock.py
Normal file
224
hermes_cli/model_setup_flows_bedrock.py
Normal file
@@ -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))
|
||||
@@ -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."""
|
||||
|
||||
409
hermes_cli/model_setup_flows_custom.py
Normal file
409
hermes_cli/model_setup_flows_custom.py
Normal file
@@ -0,0 +1,409 @@
|
||||
"""Custom OpenAI-compatible endpoint wizards: the ad-hoc ``custom`` flow and the
|
||||
``custom_providers`` / ``providers.<key>`` 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})")
|
||||
Reference in New Issue
Block a user