| 1 | """ |
| 2 | Provider preset system for ViMax chat model configuration. |
| 3 | |
| 4 | Supports auto-detection and resolution of LLM provider settings, |
| 5 | allowing users to specify a provider name (e.g., ``minimax``) instead |
| 6 | of manually configuring base_url and model details. |
| 7 | """ |
| 8 | |
| 9 | import os |
| 10 | import logging |
| 11 | from typing import Dict, Any, Optional |
| 12 | |
| 13 | logger = logging.getLogger(__name__) |
| 14 | |
| 15 | # --------------------------------------------------------------------------- |
| 16 | # Provider presets |
| 17 | # --------------------------------------------------------------------------- |
| 18 | |
| 19 | PROVIDER_PRESETS: Dict[str, Dict[str, Any]] = { |
| 20 | "minimax": { |
| 21 | "base_url": "https://api.minimax.io/v1", |
| 22 | "env_key": "MINIMAX_API_KEY", |
| 23 | "default_model": "MiniMax-M3", |
| 24 | "models": [ |
| 25 | "MiniMax-M3", |
| 26 | "MiniMax-M2.7", |
| 27 | "MiniMax-M2.7-highspeed", |
| 28 | ], |
| 29 | "temperature_range": (0.0, 1.0), |
| 30 | }, |
| 31 | } |
| 32 | |
| 33 | |
| 34 | def resolve_chat_model_config(init_args: Dict[str, Any]) -> Dict[str, Any]: |
| 35 | """Resolve provider presets and return final ``init_chat_model`` kwargs. |
| 36 | |
| 37 | If ``model_provider`` matches a known preset (e.g. ``minimax``), the |
| 38 | returned dict will have: |
| 39 | |
| 40 | * ``model_provider`` rewritten to ``"openai"`` (OpenAI-compatible API) |
| 41 | * ``base_url`` filled in from the preset when not already set |
| 42 | * ``api_key`` sourced from the environment when not already set |
| 43 | * ``model`` defaulted to the preset's default model when not already set |
| 44 | * ``temperature`` clamped to the provider's supported range |
| 45 | |
| 46 | For unknown providers the dict is returned unchanged. |
| 47 | """ |
| 48 | args = dict(init_args) # shallow copy |
| 49 | provider = args.get("model_provider", "openai") |
| 50 | |
| 51 | preset = PROVIDER_PRESETS.get(provider) |
| 52 | if preset is None: |
| 53 | return args |
| 54 | |
| 55 | # base_url |
| 56 | if not args.get("base_url"): |
| 57 | args["base_url"] = preset["base_url"] |
| 58 | |
| 59 | # api_key – fall back to env var |
| 60 | if not args.get("api_key"): |
| 61 | env_key = preset.get("env_key", "") |
| 62 | env_val = os.environ.get(env_key, "") |
| 63 | if env_val: |
| 64 | args["api_key"] = env_val |
| 65 | logger.info("Using %s API key from environment variable %s", provider, env_key) |
| 66 | |
| 67 | # default model |
| 68 | if not args.get("model"): |
| 69 | args["model"] = preset["default_model"] |
| 70 | logger.info("Defaulting to model %s for provider %s", args["model"], provider) |
| 71 | |
| 72 | # temperature clamping |
| 73 | temp_range = preset.get("temperature_range") |
| 74 | if temp_range and "temperature" in args and args["temperature"] is not None: |
| 75 | lo, hi = temp_range |
| 76 | original = args["temperature"] |
| 77 | args["temperature"] = max(lo, min(hi, original)) |
| 78 | if args["temperature"] != original: |
| 79 | logger.warning( |
| 80 | "Clamped temperature %.2f -> %.2f for provider %s", |
| 81 | original, args["temperature"], provider, |
| 82 | ) |
| 83 | |
| 84 | # rewrite to openai-compatible provider for LangChain |
| 85 | args["model_provider"] = "openai" |
| 86 | |
| 87 | return args |
| 88 | |
| 89 | |
| 90 | def detect_provider_from_env() -> Optional[str]: |
| 91 | """Return the name of a provider whose API key is found in the environment. |
| 92 | |
| 93 | Checks ``PROVIDER_PRESETS`` in definition order and returns the first |
| 94 | match, or ``None`` if no key is set. |
| 95 | """ |
| 96 | for name, preset in PROVIDER_PRESETS.items(): |
| 97 | env_key = preset.get("env_key", "") |
| 98 | if env_key and os.environ.get(env_key): |
| 99 | return name |
| 100 | return None |
| 101 |