返回 ViMax
provider_presets.py
根目录 / utils / provider_presets.py
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
101 lines PYTHON